From 7b1c639bd73d07755d9021cf5c8fcd50b47f367f Mon Sep 17 00:00:00 2001 From: Mihidum Hettiyahandi <55163074+mihidumh@users.noreply.github.com> Date: Thu, 27 Aug 2026 14:53:23 +1000 Subject: [PATCH 1/6] fix(token_counter): support OpenAI 'input_audio' content blocks token_counter raised 'Invalid content item type: input_audio' on audio understanding payloads ({'type': 'input_audio', ...}). Callers that swallow the error fail open: Router._pre_call_checks returns the unfiltered deployment list (context-window guard skipped), the prompt-caching deployment check skips (no cache-affinity pinning), trim_messages returns over-budget conversations untrimmed, and /utils/token_counter 500s outright. parallel_request_limiter_v3 already documents the gap and strips/estimates audio blocks locally; every other call site still hit the raise. Estimate audio blocks the same way that limiter does -- decoded-base64 byte count at a conservative low bitrate, floored per block -- with the constants promoted to litellm/constants.py (core cannot import from proxy hooks). Fixes #38459 Co-Authored-By: Claude Fable 5 --- litellm/constants.py | 9 +++ litellm/litellm_core_utils/token_counter.py | 28 ++++++- .../litellm_core_utils/test_token_counter.py | 77 +++++++++++++++++++ 3 files changed, 111 insertions(+), 3 deletions(-) diff --git a/litellm/constants.py b/litellm/constants.py index a292b654778..6f6469cff7c 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -85,6 +85,15 @@ DEFAULT_REPLICATE_POLLING_RETRIES: Final = int(os.getenv("DEFAULT_REPLICATE_POLL DEFAULT_REPLICATE_POLLING_DELAY_SECONDS: Final = int(os.getenv("DEFAULT_REPLICATE_POLLING_DELAY_SECONDS", 1)) DEFAULT_IMAGE_TOKEN_COUNT: Final = int(os.getenv("DEFAULT_IMAGE_TOKEN_COUNT", 250)) HF_CONFIG_FETCH_TIMEOUT_SECONDS: Final = 10.0 +# Token estimate for one `input_audio` content block (audio understanding). +# The real cost is the provider's server-side audio tokenization and cannot be +# derived exactly client-side. With a base64 payload the estimate comes from +# the decoded byte count at a conservative low bitrate -- 8 kHz mono PCM-16 +# (16 000 bytes/s) at 10 tokens/s -- so equal-duration higher-quality audio is +# never under-estimated; a payload-less (reference-only) block gets the flat +# per-block floor. parallel_request_limiter_v3 reserves with the same numbers. +DEFAULT_AUDIO_TOKEN_ESTIMATE: Final = 300 +AUDIO_BYTES_PER_TOKEN: Final = 1600 # Maximum wall-clock seconds a streaming response is allowed to run. # Streams exceeding this duration are terminated with a Timeout error. diff --git a/litellm/litellm_core_utils/token_counter.py b/litellm/litellm_core_utils/token_counter.py index cdd2d0654be..0b532270415 100644 --- a/litellm/litellm_core_utils/token_counter.py +++ b/litellm/litellm_core_utils/token_counter.py @@ -16,6 +16,8 @@ import litellm from litellm import verbose_logger from litellm._lazy_imports import _get_default_encoding from litellm.constants import ( + AUDIO_BYTES_PER_TOKEN, + DEFAULT_AUDIO_TOKEN_ESTIMATE, DEFAULT_IMAGE_HEIGHT, DEFAULT_IMAGE_TOKEN_COUNT, DEFAULT_IMAGE_WIDTH, @@ -866,6 +868,25 @@ def _count_anthropic_content( return tokens +def _count_input_audio_content_block(c: Mapping[str, object]) -> int: + """ + Estimate tokens for an OpenAI ``input_audio`` content block (audio + understanding), e.g. {"type": "input_audio", "input_audio": {"data": + "", "format": "wav"}}. The real token cost is the provider's + server-side audio tokenization and cannot be derived exactly client-side; + when the block carries a base64 payload, derive the estimate from the + decoded byte count at a conservative low bitrate, otherwise use the flat + per-block floor -- the same numbers ``parallel_request_limiter_v3`` + already uses for its audio reservations (issue #38459). + """ + input_audio: Final = c.get("input_audio") + b64_data: Final = input_audio.get("data") if isinstance(input_audio, dict) else None + if isinstance(b64_data, str) and b64_data: + decoded_bytes: Final = len(b64_data) * 3 // 4 + return max(decoded_bytes // AUDIO_BYTES_PER_TOKEN, DEFAULT_AUDIO_TOKEN_ESTIMATE) + return DEFAULT_AUDIO_TOKEN_ESTIMATE + + def _count_content_list( count_function: TokenCounterFunction, content_list: str @@ -921,8 +942,7 @@ def _count_content_list( # Claude extended thinking content block # Count the thinking text and skip the opaque blobs (signature, redacted data) thinking_text = str(c.get("thinking", "")) - if thinking_text: - num_tokens += count_function(thinking_text) + num_tokens += count_function(thinking_text) if thinking_text else 0 elif c["type"] == "tool_reference": # Anthropic tool-search reference block: a lightweight pointer to # a deferred tool, e.g. {"type": "tool_reference", "tool_name": ...}. @@ -934,13 +954,15 @@ def _count_content_list( tool_name = str(c.get("tool_name") or "") if tool_name: num_tokens += count_function(tool_name) + elif c["type"] == "input_audio": + num_tokens += _count_input_audio_content_block(c) else: content_type = c.get("type", type(c).__name__) if isinstance(c, dict) else type(c).__name__ raise ValueError( f"Invalid content item type: {content_type}. " f"Expected str or dict with 'type' field " f"(text, image_url, image, document, file, tool_use, tool_result, thinking, redacted_thinking, " - f"tool_reference)." + f"tool_reference, input_audio)." ) return num_tokens except Exception as e: diff --git a/tests/unit/litellm_core_utils/test_token_counter.py b/tests/unit/litellm_core_utils/test_token_counter.py index c71b1496bdd..dc0a9e223d5 100644 --- a/tests/unit/litellm_core_utils/test_token_counter.py +++ b/tests/unit/litellm_core_utils/test_token_counter.py @@ -1633,3 +1633,80 @@ def test_token_counter_uses_the_tokenizer_of_each_model_family_and_of_a_custom_t "custom": expected["Xenova/llama-3-tokenizer"], "requested": sorted(served), } + + +def test_token_counter_with_input_audio_content_block(): + """ + Regression test for issue #38459: a message containing an OpenAI + `input_audio` content block (audio understanding) must NOT raise from + token_counter. Before the fix the raise poisoned the whole message and + every caller that swallows counter errors failed open (router + context-window pre-call check, prompt-caching deployment check), while + /utils/token_counter returned HTTP 500. + + The estimate mirrors parallel_request_limiter_v3's audio reservation: + decoded-base64 byte count at AUDIO_BYTES_PER_TOKEN, floored at + DEFAULT_AUDIO_TOKEN_ESTIMATE per block. + """ + from litellm.constants import AUDIO_BYTES_PER_TOKEN, DEFAULT_AUDIO_TOKEN_ESTIMATE + + small_b64 = "UklGRiQAAABXQVZFZm10IBAAAAABAAEAQB8AAIA+AAACABAAZGF0YQAAAAA=" + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What does the audio say?"}, + {"type": "input_audio", "input_audio": {"data": small_b64, "format": "wav"}}, + ], + } + ] + + tokens = token_counter_new(model="gpt-4o-audio-preview", messages=messages) + assert tokens >= DEFAULT_AUDIO_TOKEN_ESTIMATE, f"Expected at least the per-block floor, got {tokens}" + + # a large payload must scale the estimate (decoded bytes / AUDIO_BYTES_PER_TOKEN) + big_b64 = "A" * (AUDIO_BYTES_PER_TOKEN * 4000) # decoded ~3000 * AUDIO_BYTES_PER_TOKEN bytes + tokens_big = token_counter_new( + model="gpt-4o-audio-preview", + messages=[ + { + "role": "user", + "content": [{"type": "input_audio", "input_audio": {"data": big_b64, "format": "wav"}}], + } + ], + ) + assert tokens_big >= 2900, f"large audio payload must scale the estimate, got {tokens_big}" + assert tokens_big > tokens + + # a payload-less (reference-only) block must not raise and gets the floor + tokens_bare = token_counter_new( + model="gpt-4o-audio-preview", + messages=[{"role": "user", "content": [{"type": "input_audio", "input_audio": {"format": "wav"}}]}], + ) + assert tokens_bare >= DEFAULT_AUDIO_TOKEN_ESTIMATE + + +def test_trim_messages_with_input_audio_content_block(): + """Companion to issue #38459, same shape as the `file`-block case + (#28409): trim_messages swallows the token_counter error and silently + returns an over-budget conversation UNTRIMMED. With the fix, trimming + must actually happen.""" + small_b64 = "UklGRiQAAABXQVZFZm10IBAAAAABAAEAQB8AAIA+AAACABAAZGF0YQAAAAA=" + messages = [ + {"role": "user", "content": "filler message " * 200}, + {"role": "user", "content": "filler message " * 200}, + { + "role": "user", + "content": [ + {"type": "text", "text": "What does the audio say?"}, + {"type": "input_audio", "input_audio": {"data": small_b64, "format": "wav"}}, + ], + }, + ] + + trimmed = litellm.utils.trim_messages(messages, model="gpt-4o-audio-preview", max_tokens=500) + + assert trimmed is not None + assert len(trimmed) < len(messages), ( + "trim_messages must actually trim an over-budget conversation containing an input_audio block" + ) From 88f36b5b818a813da341b7e8f8ba8cd45424bc7f Mon Sep 17 00:00:00 2001 From: Mihidum Hettiyahandi <55163074+mihidumh@users.noreply.github.com> Date: Thu, 27 Aug 2026 14:56:07 +1000 Subject: [PATCH 2/6] fix(batch): floor input_audio rows at the size-based estimate so audio payloads cannot evade TPM limits Before input_audio blocks were countable, a batch row carrying one RAISED inside token_counter and the batch rate limiter fell back to its conservative size-based estimate (raw bytes / 4), which covers the base64 payload. Making the row countable replaced that with the counter's audio estimate (decoded bytes / AUDIO_BYTES_PER_TOKEN, a deliberately low assumed bitrate), so a row carrying a large audio payload reserved ~500x less and a crafted batch could slide under the TPM limit -- the same loophole class raised and fixed for file blocks on #33659. Restore the conservatism at the rate-limiter call site: when a row's messages carry an input_audio block, take max(counted, size-based estimate). Plain rows keep the measured count. Live (non-batch) requests are unaffected -- parallel_request_limiter_v3 strips audio blocks before counting and applies its own identical estimate. Co-Authored-By: Claude Fable 5 --- litellm/litellm_core_utils/token_counter.py | 22 +++++ litellm/proxy/hooks/batch_rate_limiter.py | 12 +++ .../proxy/hooks/test_batch_file_validation.py | 96 +++++++++++++++++++ 3 files changed, 130 insertions(+) diff --git a/litellm/litellm_core_utils/token_counter.py b/litellm/litellm_core_utils/token_counter.py index 0b532270415..9f8870767f1 100644 --- a/litellm/litellm_core_utils/token_counter.py +++ b/litellm/litellm_core_utils/token_counter.py @@ -887,6 +887,28 @@ def _count_input_audio_content_block(c: Mapping[str, object]) -> int: return DEFAULT_AUDIO_TOKEN_ESTIMATE +def messages_contain_input_audio_content_blocks(messages: object) -> bool: + """ + True when any message carries an OpenAI ``input_audio`` content block. + + Callers that use ``token_counter`` for reservations or rate limits must + check this first: an audio block's contribution is a size-derived + ESTIMATE at a deliberately low assumed bitrate (see + ``_count_input_audio_content_block``), not a measurement, so the count + for such a message must not be trusted as an upper bound. + """ + if not isinstance(messages, list): + return False + for message in messages: + content = message.get("content") if isinstance(message, dict) else None + if not isinstance(content, list): + continue + for content_item in content: + if isinstance(content_item, dict) and content_item.get("type") == "input_audio": + return True + return False + + def _count_content_list( count_function: TokenCounterFunction, content_list: str diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index 22a17bd4cd8..140783a4a4d 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -37,6 +37,7 @@ from litellm.batches.batch_utils import ( from litellm.constants import BATCH_TPD_DESCRIPTOR_SUFFIX, BATCH_TPD_WINDOW_SECONDS from litellm.exceptions import RateLimitErrorCategory from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.token_counter import messages_contain_input_audio_content_blocks from litellm.proxy._types import ( ProxyErrorTypes, ProxyException, @@ -971,6 +972,17 @@ class _PROXY_BatchRateLimiter(CustomLogger): try: entry_total_tokens = _count_entry_tokens(entry) + # An `input_audio` block's contribution is a size-derived + # estimate at a deliberately low assumed bitrate (see + # token_counter._count_input_audio_content_block), far + # below this fallback's raw-bytes estimate. Before such + # blocks were countable (#38459) an audio row RAISED + # inside token_counter and fell back to the size-based + # estimate; floor it there again so a row carrying a + # large base64 audio payload cannot slide the batch + # under the TPM limit. + if messages_contain_input_audio_content_blocks((entry.get("body") or {}).get("messages")): + entry_total_tokens = max(entry_total_tokens, _estimate_batch_entry_tokens(raw_line)) except Exception: entry_total_tokens = _estimate_batch_entry_tokens(raw_line) total_tokens += entry_total_tokens diff --git a/tests/test_litellm/proxy/hooks/test_batch_file_validation.py b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py index 4ef94f0965b..1b2bedf070f 100644 --- a/tests/test_litellm/proxy/hooks/test_batch_file_validation.py +++ b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py @@ -2286,3 +2286,99 @@ async def test_disable_flag_still_skips_batch_processing_with_enqueued_limits(): assert result is data afile_content_mock.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_input_audio_row_reserves_size_based_floor(): + """A chat row carrying an OpenAI `input_audio` content block must reserve + at least the size-based estimate (serialized bytes / 4). + + token_counter's audio contribution is a size-derived estimate at a + deliberately low assumed bitrate (decoded bytes / AUDIO_BYTES_PER_TOKEN), + far below the raw-bytes fallback. Before audio blocks were countable + (#38459) such a row RAISED inside token_counter and fell back to the + size-based estimate; the floor restores exactly that conservatism so a + large base64 audio payload cannot slide the batch under the TPM limit. + """ + import json as _json + + from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + + prl = MagicMock() + prl.no_max_tokens_output_floor.return_value = 0 + rate_limiter = _PROXY_BatchRateLimiter( + internal_usage_cache=MagicMock(), + parallel_request_limiter=prl, + ) + + blob = "A" * 400_000 + audio_row_bytes = _json.dumps( + { + "custom_id": "row-1", + "method": "POST", + "url": "/v1/chat/completions", + "body": { + "model": "gpt-4o-audio-preview", + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What does the audio say?"}, + {"type": "input_audio", "input_audio": {"data": blob, "format": "wav"}}, + ], + } + ], + }, + } + ).encode("utf-8") + fake_content = MagicMock() + fake_content.content = audio_row_bytes + + with patch( # test-quality-ok: mirrors the file's established harness — the download, not an HTTP boundary + "litellm.afile_content", + new=AsyncMock(return_value=fake_content), + ): + usage = await rate_limiter.count_input_file_usage( + file_id="file-not-managed", + custom_llm_provider="openai", + user_api_key_dict=None, + ) + + size_based_floor = len(audio_row_bytes) // 4 + assert usage.request_count == 1 + assert usage.total_tokens >= size_based_floor, ( + f"audio-block row must reserve at least the size-based estimate " + f"({size_based_floor} tokens for {len(audio_row_bytes)} bytes), got " + f"{usage.total_tokens} — a large audio payload would evade TPM limits" + ) + + # Control: a plain-text row must NOT be floored at its serialized size — + # measured text rows keep the (smaller) real token count. + text_row_bytes = _json.dumps( + { + "custom_id": "row-1", + "method": "POST", + "url": "/v1/chat/completions", + "body": { + "model": "gpt-4o-audio-preview", + "messages": [{"role": "user", "content": "What does the audio say?"}], + }, + } + ).encode("utf-8") + fake_text_content = MagicMock() + fake_text_content.content = text_row_bytes + + with patch( # test-quality-ok: mirrors the file's established harness — the download, not an HTTP boundary + "litellm.afile_content", + new=AsyncMock(return_value=fake_text_content), + ): + text_usage = await rate_limiter.count_input_file_usage( + file_id="file-not-managed", + custom_llm_provider="openai", + user_api_key_dict=None, + ) + + assert text_usage.request_count == 1 + assert text_usage.total_tokens < len(text_row_bytes) // 4, ( + "plain-text rows must keep the measured token count, not the size-based floor" + ) From 7bc172dd4f196ddf680c3d6c955957cb60a5ae8d Mon Sep 17 00:00:00 2001 From: Mihidum Hettiyahandi <55163074+mihidumh@users.noreply.github.com> Date: Wed, 9 Sep 2026 09:44:41 +1000 Subject: [PATCH 3/6] refactor(proxy): import the audio token-estimate constants instead of duplicating them parallel_request_limiter_v3 kept its own copies of DEFAULT_AUDIO_TOKEN_ESTIMATE and _AUDIO_BYTES_PER_TOKEN next to the shared values in litellm.constants, so a future edit to one could silently make token counting and limiter reservations disagree. Import the shared ones and drop the local copies; the strip-pass docstring now says why the blocks are still stripped (token_counter counts them since #38459, so stripping keeps the audio contribution added exactly once). --- .../hooks/parallel_request_limiter_v3.py | 33 +- .../proxy/hooks/test_tpm_concurrent.py | 480 +++++------------- 2 files changed, 147 insertions(+), 366 deletions(-) diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index e4b782d5ff3..93f09e94ebb 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -34,7 +34,12 @@ from typing_extensions import NotRequired, ReadOnly from litellm import DualCache from litellm._logging import verbose_proxy_logger from litellm.caching.redis_cache import log_redis_failure -from litellm.constants import DYNAMIC_RATE_LIMIT_ERROR_THRESHOLD_PER_MINUTE, INTERNAL_CALL_ORIGIN_METADATA_KEY +from litellm.constants import ( + AUDIO_BYTES_PER_TOKEN, + DEFAULT_AUDIO_TOKEN_ESTIMATE, + DYNAMIC_RATE_LIMIT_ERROR_THRESHOLD_PER_MINUTE, + INTERNAL_CALL_ORIGIN_METADATA_KEY, +) from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.prompt_templates.common_utils import ( get_str_from_messages, @@ -446,18 +451,6 @@ GOOGLE_GENAI_NATIVE_CALL_TYPES: Final = ( CallTypes.agenerate_content_stream.value, ) RESPONSES_API_MIN_OUTPUT_TOKENS: Final = 16 -# litellm.token_counter has no per-type handling for "input_audio" content -# blocks (unlike images, which use use_default_image_token_count) -- it -# silently contributes 0 tokens for them. When the block carries a base64 -# payload, the estimate is derived from the decoded byte count; when the -# block is a reference without a payload (or the payload is missing), this -# flat per-block floor is used instead. -DEFAULT_AUDIO_TOKEN_ESTIMATE: Final = 300 -# Conservative bytes-per-token assumption for size-based audio estimation: -# equivalent to 8 kHz mono PCM-16 (16 000 bytes/s) at 10 tokens/s. Choosing -# the lowest reasonable bitrate means we never under-reserve for higher- -# quality audio recorded at the same wall-clock duration. -_AUDIO_BYTES_PER_TOKEN: Final = 1600 # Descriptor "key" values for project-scoped ITPM/OTPM. Distinct from # "model_per_project" (the combined-TPM descriptor) so both can be enforced # on the same project+model simultaneously without colliding on cache keys. @@ -3361,7 +3354,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): Token estimate for one ``input_audio`` content block. When the block carries a base64 ``data`` payload, the estimate comes - from the decoded byte count (``len(b64) * 3 // 4 // _AUDIO_BYTES_PER_TOKEN``), + from the decoded byte count (``len(b64) * 3 // 4 // AUDIO_BYTES_PER_TOKEN``), assuming the lowest reasonable audio bitrate so we never under-reserve for higher-quality recordings of the same duration. @@ -3374,7 +3367,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): b64_data: Final = input_audio.get("data") if isinstance(input_audio, dict) else None if b64_data and isinstance(b64_data, str): decoded_bytes: Final = len(b64_data) * 3 // 4 - return max(decoded_bytes // _AUDIO_BYTES_PER_TOKEN, DEFAULT_AUDIO_TOKEN_ESTIMATE) + return max(decoded_bytes // AUDIO_BYTES_PER_TOKEN, DEFAULT_AUDIO_TOKEN_ESTIMATE) return DEFAULT_AUDIO_TOKEN_ESTIMATE @classmethod @@ -3400,11 +3393,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): def _strip_audio_content_blocks(messages: object) -> object: """ Drop ``input_audio`` content blocks before passing ``messages`` to - ``token_counter``, which raises ``ValueError`` on them (no per-type - handling, unlike images). The audio contribution is added back - separately via ``DEFAULT_AUDIO_TOKEN_ESTIMATE`` so the rest of the - message (text/images/tools) still gets counted accurately instead of - the whole call falling back to the cheap char-count estimate. + ``token_counter``. It counts them with the same size-derived estimate + since #38459 (it used to raise); stripping keeps the audio contribution + added exactly once, by ``_estimate_audio_content_tokens``, so the rest + of the message (text/images/tools) is counted accurately without ever + double-counting the audio. """ if not isinstance(messages, list): return messages diff --git a/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py b/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py index 42c1f489bdd..5146b53eeca 100644 --- a/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py +++ b/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py @@ -21,12 +21,12 @@ from typing import Any, Dict import pytest from litellm.caching.caching import DualCache +from litellm.constants import AUDIO_BYTES_PER_TOKEN from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.parallel_request_limiter_v3 import ( PROJECT_ITPM_DESCRIPTOR_KEY, PROJECT_OTPM_DESCRIPTOR_KEY, RateLimitedModel, - _AUDIO_BYTES_PER_TOKEN, _PROXY_MaxParallelRequestsHandler_v3 as RateLimitHandler, ) from litellm.proxy.hooks.parallel_request_limiter_v3 import ( @@ -151,7 +151,7 @@ async def test_no_leak_on_over_limit_rejection(rate_limiter): f"estimated={estimated}, limit={user_api_key_dict.tpm_limit}" ) - with pytest.raises(Exception, match='Limit type: tokens\\. Current limit') as exc_info: + with pytest.raises(Exception, match="Limit type: tokens\\. Current limit") as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=cache, @@ -168,8 +168,7 @@ async def test_no_leak_on_over_limit_rejection(rate_limiter): cached_value = await cache.async_get_cache(key=counter_key, local_only=True) cached_int = int(cached_value or 0) assert cached_int < estimated, ( - f"Reservation leaked: counter={cached_int} after rejection of an " - f"estimated_tokens={estimated} reservation." + f"Reservation leaked: counter={cached_int} after rejection of an estimated_tokens={estimated} reservation." ) @@ -218,9 +217,7 @@ async def test_token_adjustment_on_success(rate_limiter): } ) - handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( - mock_increment - ) + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_increment await handler.async_log_success_event( kwargs=mock_kwargs, @@ -232,8 +229,7 @@ async def test_token_adjustment_on_success(rate_limiter): token_adjustments = [i for i in increments if "tokens" in i["key"]] assert any(i["increment"] == -50 for i in token_adjustments), ( - f"Expected a -50 token adjustment (50 actual - 100 reserved) but got: " - f"{token_adjustments}" + f"Expected a -50 token adjustment (50 actual - 100 reserved) but got: {token_adjustments}" ) @@ -270,9 +266,7 @@ async def test_token_release_on_failure(rate_limiter): } ) - handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( - mock_increment - ) + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_increment await handler.async_log_failure_event( kwargs=mock_kwargs, @@ -283,9 +277,9 @@ async def test_token_release_on_failure(rate_limiter): token_releases = [i for i in increments if "tokens" in i["key"]] - assert any( - i["increment"] == -100 for i in token_releases - ), f"Expected the full reservation (-100) to be released, got: {token_releases}" + assert any(i["increment"] == -100 for i in token_releases), ( + f"Expected the full reservation (-100) to be released, got: {token_releases}" + ) @pytest.mark.asyncio @@ -329,9 +323,7 @@ async def test_model_scope_refund_targets_reserved_model(rate_limiter): for op in increment_list: increments.append({"key": op["key"], "increment": op["increment_value"]}) - handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( - mock_increment - ) + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_increment await handler.async_log_failure_event( kwargs=mock_kwargs, @@ -352,8 +344,7 @@ async def test_model_scope_refund_targets_reserved_model(rate_limiter): f"{[i['key'] for i in increments]}" ) assert matching[0]["increment"] == -100, ( - f"Expected full -100 refund on model_per_team counter, got " - f"{matching[0]['increment']}" + f"Expected full -100 refund on model_per_team counter, got {matching[0]['increment']}" ) @@ -374,9 +365,7 @@ async def test_should_rate_limit_does_not_inflate_tokens_counter(rate_limiter): tpm_limit=10_000, ) - tokens_counter_key = handler.create_rate_limit_keys( - key="api_key", value=api_key, rate_limit_type="tokens" - ) + tokens_counter_key = handler.create_rate_limit_keys(key="api_key", value=api_key, rate_limit_type="tokens") data = { "model": "gpt-3.5-turbo", @@ -393,9 +382,7 @@ async def test_should_rate_limit_does_not_inflate_tokens_counter(rate_limiter): call_type="", ) - cached = int( - await cache.async_get_cache(key=tokens_counter_key, local_only=True) or 0 - ) + cached = int(await cache.async_get_cache(key=tokens_counter_key, local_only=True) or 0) # The :tokens counter should reflect ONLY the reservation amount — not # an additional +1 from the should_rate_limit pre-pass. @@ -487,9 +474,7 @@ async def test_org_scope_refund_on_failure(rate_limiter): for op in increment_list: increments.append({"key": op["key"], "increment": op["increment_value"]}) - handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( - mock_increment - ) + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_increment await handler.async_log_failure_event( kwargs=mock_kwargs, @@ -498,17 +483,13 @@ async def test_org_scope_refund_on_failure(rate_limiter): end_time=datetime.now(), ) - expected_org_key = handler.create_rate_limit_keys( - key="organization", value=org_id, rate_limit_type="tokens" - ) + expected_org_key = handler.create_rate_limit_keys(key="organization", value=org_id, rate_limit_type="tokens") matching = [i for i in increments if i["key"] == expected_org_key] assert matching, ( f"Expected a refund on the org tokens counter ({expected_org_key}) " f"but got keys: {[i['key'] for i in increments]}" ) - assert ( - matching[0]["increment"] == -100 - ), f"Expected full -100 refund on org counter, got {matching[0]['increment']}" + assert matching[0]["increment"] == -100, f"Expected full -100 refund on org counter, got {matching[0]['increment']}" @pytest.mark.asyncio @@ -551,9 +532,7 @@ async def test_org_scope_reconciled_on_success(rate_limiter): for op in increment_list: increments.append({"key": op["key"], "increment": op["increment_value"]}) - handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( - mock_increment - ) + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_increment await handler.async_log_success_event( kwargs=mock_kwargs, @@ -562,17 +541,14 @@ async def test_org_scope_reconciled_on_success(rate_limiter): end_time=datetime.now(), ) - expected_org_key = handler.create_rate_limit_keys( - key="organization", value=org_id, rate_limit_type="tokens" - ) + expected_org_key = handler.create_rate_limit_keys(key="organization", value=org_id, rate_limit_type="tokens") matching = [i for i in increments if i["key"] == expected_org_key] assert matching, ( f"Expected a reconciliation op on the org tokens counter " f"({expected_org_key}), got keys: {[i['key'] for i in increments]}" ) assert matching[0]["increment"] == -50, ( - f"Expected -50 delta on org counter (50 actual - 100 reserved), got " - f"{matching[0]['increment']}" + f"Expected -50 delta on org counter (50 actual - 100 reserved), got {matching[0]['increment']}" ) @@ -583,9 +559,7 @@ async def test_estimate_tokens_uses_max_tokens_when_explicit(rate_limiter): estimate = handler._estimate_tokens_for_request( data={ - "messages": [ - {"role": "user", "content": "abcd" * 4} - ], # 16 chars ~ 4 tokens + "messages": [{"role": "user", "content": "abcd" * 4}], # 16 chars ~ 4 tokens "max_tokens": 25, } ) @@ -606,15 +580,11 @@ async def test_estimate_tokens_honors_explicit_zero_max_tokens(rate_limiter): estimate = handler._estimate_tokens_for_request( data={ - "messages": [ - {"role": "user", "content": "abcd" * 4} - ], # 16 chars ~ 4 tokens + "messages": [{"role": "user", "content": "abcd" * 4}], # 16 chars ~ 4 tokens "max_tokens": 0, } ) - assert estimate == 4, ( - f"expected input-only reservation (4) for an explicit max_tokens=0, got {estimate}" - ) + assert estimate == 4, f"expected input-only reservation (4) for an explicit max_tokens=0, got {estimate}" @pytest.mark.asyncio @@ -660,9 +630,7 @@ async def test_contentless_request_reserves_minimum(rate_limiter): api_key = hash_token("sk-contentless") user_api_key_dict = UserAPIKeyAuth(api_key=api_key, tpm_limit=2) - counter_key = handler.create_rate_limit_keys( - key="api_key", value=api_key, rate_limit_type="tokens" - ) + counter_key = handler.create_rate_limit_keys(key="api_key", value=api_key, rate_limit_type="tokens") # Two contentless requests should consume two slots of the 2-token # budget. The third must 429. @@ -674,19 +642,14 @@ async def test_contentless_request_reserves_minimum(rate_limiter): data=data, call_type="", ) - assert ( - get_request_stash().reserved_tokens == 1 - ), "Contentless request should reserve the floor of 1 token" + assert get_request_stash().reserved_tokens == 1, "Contentless request should reserve the floor of 1 token" - counter_after_two = int( - await cache.async_get_cache(key=counter_key, local_only=True) or 0 - ) + counter_after_two = int(await cache.async_get_cache(key=counter_key, local_only=True) or 0) assert counter_after_two == 2, ( - f"After two contentless requests at the floor, the api_key tokens " - f"counter should be 2, got {counter_after_two}" + f"After two contentless requests at the floor, the api_key tokens counter should be 2, got {counter_after_two}" ) - with pytest.raises(Exception, match='Limit type: tokens\\. Current limit') as exc_info: + with pytest.raises(Exception, match="Limit type: tokens\\. Current limit") as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=cache, @@ -694,8 +657,7 @@ async def test_contentless_request_reserves_minimum(rate_limiter): call_type="", ) assert getattr(exc_info.value, "status_code", None) == 429, ( - "Third contentless request must be rate-limited; pre-fix it would " - "have bypassed the TPM check entirely." + "Third contentless request must be rate-limited; pre-fix it would have bypassed the TPM check entirely." ) @@ -772,12 +734,8 @@ async def test_reservation_released_on_proxy_rejection(rate_limiter): reserved = get_request_stash().reserved_tokens assert reserved > 0 - counter_key = handler.create_rate_limit_keys( - key="api_key", value=api_key, rate_limit_type="tokens" - ) - counter_after_reserve = int( - await cache.async_get_cache(key=counter_key, local_only=True) or 0 - ) + counter_key = handler.create_rate_limit_keys(key="api_key", value=api_key, rate_limit_type="tokens") + counter_after_reserve = int(await cache.async_get_cache(key=counter_key, local_only=True) or 0) assert counter_after_reserve == reserved # Simulate a downstream guardrail rejecting the request. @@ -787,16 +745,12 @@ async def test_reservation_released_on_proxy_rejection(rate_limiter): user_api_key_dict=user_api_key_dict, ) - counter_after_release = int( - await cache.async_get_cache(key=counter_key, local_only=True) or 0 - ) + counter_after_release = int(await cache.async_get_cache(key=counter_key, local_only=True) or 0) assert counter_after_release == 0, ( - f"Reservation leaked: counter={counter_after_release} after " - f"proxy-level rejection refund (expected 0)." + f"Reservation leaked: counter={counter_after_release} after proxy-level rejection refund (expected 0)." ) assert get_request_stash().reservation_released is True, ( - "Released flag must be set to prevent " - "async_log_failure_event from double-refunding." + "Released flag must be set to prevent async_log_failure_event from double-refunding." ) @@ -817,9 +771,7 @@ async def test_reservation_release_idempotent(rate_limiter): for op in increment_list: increments.append({"key": op["key"], "increment": op["increment_value"]}) - handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( - mock_increment - ) + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_increment # Both hooks read the same per-request ContextVar stash: the # post-call-failure-hook flips reservation_released on it, and the @@ -900,9 +852,7 @@ async def test_unreserved_scopes_charged_actual_not_delta_on_success(rate_limite for op in increment_list: increments.append({"key": op["key"], "increment": op["increment_value"]}) - handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( - mock_increment - ) + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_increment await handler.async_log_success_event( kwargs=mock_kwargs, @@ -911,19 +861,14 @@ async def test_unreserved_scopes_charged_actual_not_delta_on_success(rate_limite end_time=datetime.now(), ) - api_key_token_key = handler.create_rate_limit_keys( - key="api_key", value=api_key, rate_limit_type="tokens" - ) - team_token_key = handler.create_rate_limit_keys( - key="team", value=team_id, rate_limit_type="tokens" - ) + api_key_token_key = handler.create_rate_limit_keys(key="api_key", value=api_key, rate_limit_type="tokens") + team_token_key = handler.create_rate_limit_keys(key="team", value=team_id, rate_limit_type="tokens") api_key_ops = [i for i in increments if i["key"] == api_key_token_key] team_ops = [i for i in increments if i["key"] == team_token_key] assert api_key_ops and api_key_ops[0]["increment"] == -50, ( - f"Reserved api_key scope must reconcile via delta (50-100=-50), " - f"got {api_key_ops}" + f"Reserved api_key scope must reconcile via delta (50-100=-50), got {api_key_ops}" ) assert team_ops and team_ops[0]["increment"] == 50, ( f"Unreserved team scope must be charged full actual (+50), not the " @@ -962,9 +907,7 @@ async def test_unreserved_scopes_not_refunded_on_failure(rate_limiter): for op in increment_list: increments.append({"key": op["key"], "increment": op["increment_value"]}) - handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( - mock_increment - ) + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_increment await handler.async_log_failure_event( kwargs=mock_kwargs, @@ -973,23 +916,16 @@ async def test_unreserved_scopes_not_refunded_on_failure(rate_limiter): end_time=datetime.now(), ) - team_token_key = handler.create_rate_limit_keys( - key="team", value=team_id, rate_limit_type="tokens" - ) - api_key_token_key = handler.create_rate_limit_keys( - key="api_key", value=api_key, rate_limit_type="tokens" - ) + team_token_key = handler.create_rate_limit_keys(key="team", value=team_id, rate_limit_type="tokens") + api_key_token_key = handler.create_rate_limit_keys(key="api_key", value=api_key, rate_limit_type="tokens") team_ops = [i for i in increments if i["key"] == team_token_key] api_key_ops = [i for i in increments if i["key"] == api_key_token_key] - assert not team_ops, ( - f"Unreserved team scope must NOT be refunded (would drift negative), " - f"got {team_ops}" + assert not team_ops, f"Unreserved team scope must NOT be refunded (would drift negative), got {team_ops}" + assert api_key_ops and api_key_ops[0]["increment"] == -100, ( + f"Reserved api_key scope must be refunded -100, got {api_key_ops}" ) - assert ( - api_key_ops and api_key_ops[0]["increment"] == -100 - ), f"Reserved api_key scope must be refunded -100, got {api_key_ops}" @pytest.mark.asyncio @@ -1024,9 +960,7 @@ async def test_token_rate_limit_headers_present_in_stored_response(rate_limiter) ) response = get_request_stash().rate_limit_response - assert isinstance( - response, dict - ), "Expected the stashed rate-limit response to be set after pre-call" + assert isinstance(response, dict), "Expected the stashed rate-limit response to be set after pre-call" statuses = response.get("statuses") or [] token_statuses = [s for s in statuses if s.get("rate_limit_type") == "tokens"] @@ -1038,8 +972,7 @@ async def test_token_rate_limit_headers_present_in_stored_response(rate_limiter) f"statuses: {[(s.get('descriptor_key'), s.get('rate_limit_type')) for s in statuses]}" ) assert request_statuses, ( - "RPM rate-limit status was clobbered by the TPM merge — both must " - "coexist in the stored response." + "RPM rate-limit status was clobbered by the TPM merge — both must coexist in the stored response." ) # The token status carries the limit and a positive remaining budget. @@ -1067,9 +1000,7 @@ async def test_estimate_tokens_floor_caps_at_smallest_configured_tpm(rate_limite ) # input ~= 5//4 = 1 token; output floor capped at 1000//4 = 250; # total ~= 251 (well under 1000). - assert ( - estimate <= 1000 // 2 - ), f"With TPM=1000, reservation must stay well under the limit; got {estimate}" + assert estimate <= 1000 // 2, f"With TPM=1000, reservation must stay well under the limit; got {estimate}" assert estimate >= 1, "Estimate must be at least the call-site floor of 1" @@ -1139,10 +1070,7 @@ async def test_small_tpm_cap_admits_no_max_tokens_request(rate_limiter): reserved = get_request_stash().reserved_tokens assert reserved > 0, "Reservation should have been stashed" - assert reserved <= 1000 // 2, ( - f"Capped floor must keep the reservation well under the 1000 TPM " - f"cap; got {reserved}" - ) + assert reserved <= 1000 // 2, f"Capped floor must keep the reservation well under the 1000 TPM cap; got {reserved}" @pytest.mark.asyncio @@ -1177,8 +1105,7 @@ async def test_small_tpm_cap_injects_matching_max_tokens(rate_limiter): ) assert data.get("max_tokens") == 1000 // 4, ( - f"Capped floor must be written to max_tokens to bound the actual " - f"model output; got {data.get('max_tokens')}" + f"Capped floor must be written to max_tokens to bound the actual model output; got {data.get('max_tokens')}" ) @@ -1211,10 +1138,7 @@ async def test_large_tpm_cap_does_not_inject_max_tokens(rate_limiter): call_type="", ) - assert "max_tokens" not in data, ( - f"Large TPM caps should leave max_tokens alone; got " - f"{data.get('max_tokens')}" - ) + assert "max_tokens" not in data, f"Large TPM caps should leave max_tokens alone; got {data.get('max_tokens')}" @pytest.mark.asyncio @@ -1294,13 +1218,9 @@ async def test_project_otpm_reservation_prevents_concurrent_bypass(rate_limiter) results = await asyncio.gather(*[make_request(i) for i in range(5)]) successful = [r for r in results if r["success"]] - rate_limited = [ - r for r in results if not r["success"] and r.get("status_code") == 429 - ] + rate_limited = [r for r in results if not r["success"] and r.get("status_code") == 429] - assert len(rate_limited) > 0, ( - f"Expected some OTPM-rate-limited requests but all {len(successful)} succeeded." - ) + assert len(rate_limited) > 0, f"Expected some OTPM-rate-limited requests but all {len(successful)} succeeded." @pytest.mark.asyncio @@ -1320,7 +1240,7 @@ async def test_project_otpm_rejects_multiple_completion_candidates(rate_limiter) "n": 10, } - with pytest.raises(Exception, match='Rate limit exceeded for model_per_project_otpm') as exc_info: + with pytest.raises(Exception, match="Rate limit exceeded for model_per_project_otpm") as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=cache, @@ -1348,7 +1268,7 @@ async def test_project_otpm_reserves_largest_conflicting_output_cap(rate_limiter "max_completion_tokens": 100, } - with pytest.raises(Exception, match='Rate limit exceeded for model_per_project_otpm') as exc_info: + with pytest.raises(Exception, match="Rate limit exceeded for model_per_project_otpm") as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=cache, @@ -1378,7 +1298,7 @@ async def test_project_otpm_rejects_google_genai_native_output_cap( project_metadata={"model_otpm_limit": {model: 50}}, ) - with pytest.raises(Exception, match='Rate limit exceeded for model_per_project_otpm') as exc_info: + with pytest.raises(Exception, match="Rate limit exceeded for model_per_project_otpm") as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=cache, @@ -1412,7 +1332,7 @@ async def test_project_otpm_rejects_google_genai_native_candidate_count( project_metadata={"model_otpm_limit": {model: 150}}, ) - with pytest.raises(Exception, match='Rate limit exceeded for model_per_project_otpm') as exc_info: + with pytest.raises(Exception, match="Rate limit exceeded for model_per_project_otpm") as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=cache, @@ -1501,7 +1421,7 @@ async def test_project_otpm_over_limit_rolls_back_itpm_reservation(rate_limiter) "max_tokens": 500, # blows past the 10-token OTPM limit } - with pytest.raises(Exception, match='Rate limit exceeded for model_per_project_otpm') as exc_info: + with pytest.raises(Exception, match="Rate limit exceeded for model_per_project_otpm") as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=cache, @@ -1511,9 +1431,7 @@ async def test_project_otpm_over_limit_rolls_back_itpm_reservation(rate_limiter) assert getattr(exc_info.value, "status_code", None) == 429 cached_value = await cache.async_get_cache(key=itpm_counter_key, local_only=True) - assert int(cached_value or 0) == 0, ( - f"ITPM reservation leaked after OTPM rejection: counter={cached_value}" - ) + assert int(cached_value or 0) == 0, f"ITPM reservation leaked after OTPM rejection: counter={cached_value}" @pytest.mark.asyncio @@ -1556,9 +1474,7 @@ async def test_project_itpm_reconciled_on_success_excludes_cached_tokens(rate_li for op in increment_list: increments.append({"key": op["key"], "increment": op["increment_value"]}) - handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( - mock_increment - ) + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_increment await handler.async_log_success_event( kwargs=mock_kwargs, @@ -1601,17 +1517,11 @@ async def test_project_reconciliation_does_not_decrement_later_window(): increments=[{"tokens": 100}], ) counter_key = handler.create_rate_limit_keys(*scope, rate_limit_type="tokens") - window_identity = next( - identity - for identity in reservation["reservation_windows"] - if identity[0] == counter_key - ) + window_identity = next(identity for identity in reservation["reservation_windows"] if identity[0] == counter_key) stash = get_or_create_request_stash() stash.itpm_reserved_tokens = 100 stash.itpm_reserved_scopes = frozenset({scope}) - stash.itpm_reserved_window_identities = frozenset( - {window_identity} - ) + stash.itpm_reserved_window_identities = frozenset({window_identity}) current_time += timedelta(seconds=61) later_reservation = await handler.atomic_check_and_increment_by_n( @@ -1622,9 +1532,7 @@ async def test_project_reconciliation_does_not_decrement_later_window(): await handler.async_log_success_event( kwargs={}, - response_obj=ModelResponse( - usage=Usage(prompt_tokens=10, completion_tokens=0, total_tokens=10) - ), + response_obj=ModelResponse(usage=Usage(prompt_tokens=10, completion_tokens=0, total_tokens=10)), start_time=current_time, end_time=current_time, ) @@ -1646,23 +1554,15 @@ async def test_project_reconciliation_decrements_its_active_window(rate_limiter) increments=[{"tokens": 100}], ) counter_key = handler.create_rate_limit_keys(*scope, rate_limit_type="tokens") - window_identity = next( - identity - for identity in reservation["reservation_windows"] - if identity[0] == counter_key - ) + window_identity = next(identity for identity in reservation["reservation_windows"] if identity[0] == counter_key) stash = get_or_create_request_stash() stash.itpm_reserved_tokens = 100 stash.itpm_reserved_scopes = frozenset({scope}) - stash.itpm_reserved_window_identities = frozenset( - {window_identity} - ) + stash.itpm_reserved_window_identities = frozenset({window_identity}) await handler.async_log_success_event( kwargs={}, - response_obj=ModelResponse( - usage=Usage(prompt_tokens=10, completion_tokens=0, total_tokens=10) - ), + response_obj=ModelResponse(usage=Usage(prompt_tokens=10, completion_tokens=0, total_tokens=10)), start_time=datetime.now(), end_time=datetime.now(), ) @@ -1750,9 +1650,7 @@ async def test_atomic_lua_response_carries_redis_window_identity(rate_limiter): ) assert response["statuses"][0]["limit_remaining"] == 75 - assert response["reservation_windows"] == frozenset( - {(counter_key, "1234", "redis")} - ) + assert response["reservation_windows"] == frozenset({(counter_key, "1234", "redis")}) @pytest.mark.asyncio @@ -1776,9 +1674,7 @@ async def test_project_itpm_otpm_released_on_failure(rate_limiter): for op in increment_list: increments.append({"key": op["key"], "increment": op["increment_value"]}) - handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( - mock_increment - ) + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_increment await handler.async_log_failure_event( kwargs=mock_kwargs, @@ -1822,9 +1718,7 @@ async def test_proxy_rejection_refunds_itpm_otpm_by_their_own_amount_not_combine data = { "model": "bedrock_mantle/claude-opus", - "messages": [ - {"role": "user", "content": "hello there, this is a test message"} - ], + "messages": [{"role": "user", "content": "hello there, this is a test message"}], "max_tokens": 60, } @@ -1851,15 +1745,9 @@ async def test_proxy_rejection_refunds_itpm_otpm_by_their_own_amount_not_combine rate_limit_type="tokens", ) - tpm_reserved = int( - await cache.async_get_cache(key=tpm_counter_key, local_only=True) or 0 - ) - itpm_reserved = int( - await cache.async_get_cache(key=itpm_counter_key, local_only=True) or 0 - ) - otpm_reserved = int( - await cache.async_get_cache(key=otpm_counter_key, local_only=True) or 0 - ) + tpm_reserved = int(await cache.async_get_cache(key=tpm_counter_key, local_only=True) or 0) + itpm_reserved = int(await cache.async_get_cache(key=itpm_counter_key, local_only=True) or 0) + otpm_reserved = int(await cache.async_get_cache(key=otpm_counter_key, local_only=True) or 0) assert tpm_reserved > 0 and itpm_reserved > 0 and otpm_reserved > 0 await handler.async_post_call_failure_hook( @@ -1868,23 +1756,13 @@ async def test_proxy_rejection_refunds_itpm_otpm_by_their_own_amount_not_combine user_api_key_dict=user_api_key_dict, ) - tpm_after = int( - await cache.async_get_cache(key=tpm_counter_key, local_only=True) or 0 - ) - itpm_after = int( - await cache.async_get_cache(key=itpm_counter_key, local_only=True) or 0 - ) - otpm_after = int( - await cache.async_get_cache(key=otpm_counter_key, local_only=True) or 0 - ) + tpm_after = int(await cache.async_get_cache(key=tpm_counter_key, local_only=True) or 0) + itpm_after = int(await cache.async_get_cache(key=itpm_counter_key, local_only=True) or 0) + otpm_after = int(await cache.async_get_cache(key=otpm_counter_key, local_only=True) or 0) assert tpm_after == 0, f"combined TPM counter leaked: {tpm_after}" - assert itpm_after == 0, ( - f"ITPM counter corrupted by combined-amount refund: {itpm_after}" - ) - assert otpm_after == 0, ( - f"OTPM counter corrupted by combined-amount refund: {otpm_after}" - ) + assert itpm_after == 0, f"ITPM counter corrupted by combined-amount refund: {itpm_after}" + assert otpm_after == 0, f"OTPM counter corrupted by combined-amount refund: {otpm_after}" @pytest.mark.asyncio @@ -1912,9 +1790,7 @@ async def test_proxy_rejection_refunds_itpm_otpm_only_reservation_with_no_combin data = { "model": "bedrock_mantle/claude-opus", - "messages": [ - {"role": "user", "content": "hello there, this is a test message"} - ], + "messages": [{"role": "user", "content": "hello there, this is a test message"}], "max_tokens": 60, } @@ -1935,12 +1811,8 @@ async def test_proxy_rejection_refunds_itpm_otpm_only_reservation_with_no_combin value="proj-io-only:bedrock_mantle/claude-opus", rate_limit_type="tokens", ) - assert ( - int(await cache.async_get_cache(key=itpm_counter_key, local_only=True) or 0) > 0 - ) - assert ( - int(await cache.async_get_cache(key=otpm_counter_key, local_only=True) or 0) > 0 - ) + assert int(await cache.async_get_cache(key=itpm_counter_key, local_only=True) or 0) > 0 + assert int(await cache.async_get_cache(key=otpm_counter_key, local_only=True) or 0) > 0 await handler.async_post_call_failure_hook( request_data=data, @@ -1948,18 +1820,10 @@ async def test_proxy_rejection_refunds_itpm_otpm_only_reservation_with_no_combin user_api_key_dict=user_api_key_dict, ) - itpm_after = int( - await cache.async_get_cache(key=itpm_counter_key, local_only=True) or 0 - ) - otpm_after = int( - await cache.async_get_cache(key=otpm_counter_key, local_only=True) or 0 - ) - assert itpm_after == 0, ( - f"ITPM-only reservation leaked on proxy rejection: {itpm_after}" - ) - assert otpm_after == 0, ( - f"OTPM-only reservation leaked on proxy rejection: {otpm_after}" - ) + itpm_after = int(await cache.async_get_cache(key=itpm_counter_key, local_only=True) or 0) + otpm_after = int(await cache.async_get_cache(key=otpm_counter_key, local_only=True) or 0) + assert itpm_after == 0, f"ITPM-only reservation leaked on proxy rejection: {itpm_after}" + assert otpm_after == 0, f"OTPM-only reservation leaked on proxy rejection: {otpm_after}" @pytest.mark.asyncio @@ -1992,9 +1856,7 @@ async def test_otpm_rejection_does_not_double_refund_combined_tpm(rate_limiter): data = { "model": "bedrock_mantle/claude-opus", - "messages": [ - {"role": "user", "content": "hello there, this is a test message"} - ], + "messages": [{"role": "user", "content": "hello there, this is a test message"}], "max_tokens": 60, # blows past the 5-token OTPM limit } @@ -2004,7 +1866,7 @@ async def test_otpm_rejection_does_not_double_refund_combined_tpm(rate_limiter): rate_limit_type="tokens", ) - with pytest.raises(Exception, match='Rate limit exceeded for model_per_project_otpm') as exc_info: + with pytest.raises(Exception, match="Rate limit exceeded for model_per_project_otpm") as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=cache, @@ -2013,12 +1875,8 @@ async def test_otpm_rejection_does_not_double_refund_combined_tpm(rate_limiter): ) assert getattr(exc_info.value, "status_code", None) == 429 - tpm_after_pre_call = int( - await cache.async_get_cache(key=tpm_counter_key, local_only=True) or 0 - ) - assert tpm_after_pre_call == 0, ( - f"combined TPM reservation not rolled back: {tpm_after_pre_call}" - ) + tpm_after_pre_call = int(await cache.async_get_cache(key=tpm_counter_key, local_only=True) or 0) + assert tpm_after_pre_call == 0, f"combined TPM reservation not rolled back: {tpm_after_pre_call}" # In the real request lifecycle, async_post_call_failure_hook fires next # for a pre-call rejection. It must not refund the same reservation again. @@ -2028,9 +1886,7 @@ async def test_otpm_rejection_does_not_double_refund_combined_tpm(rate_limiter): user_api_key_dict=user_api_key_dict, ) - tpm_after_failure_hook = int( - await cache.async_get_cache(key=tpm_counter_key, local_only=True) or 0 - ) + tpm_after_failure_hook = int(await cache.async_get_cache(key=tpm_counter_key, local_only=True) or 0) assert tpm_after_failure_hook == 0, ( f"combined TPM counter went negative from a double refund: {tpm_after_failure_hook}" ) @@ -2061,7 +1917,7 @@ async def test_project_itpm_rejects_pretokenized_embedding_input( "input": embedding_input, } - with pytest.raises(Exception, match='Rate limit exceeded for model_per_project_itpm') as exc_info: + with pytest.raises(Exception, match="Rate limit exceeded for model_per_project_itpm") as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=cache, @@ -2088,14 +1944,10 @@ async def test_responses_api_not_misclassified_as_embedding_for_output_estimate( data = {"input": "describe this image in detail"} - _, embedding_output_estimate = handler._estimate_input_and_output_tokens( - data=data, call_type="aembedding" - ) + _, embedding_output_estimate = handler._estimate_input_and_output_tokens(data=data, call_type="aembedding") assert embedding_output_estimate == 0 - _, responses_output_estimate = handler._estimate_input_and_output_tokens( - data=data, call_type="aresponses" - ) + _, responses_output_estimate = handler._estimate_input_and_output_tokens(data=data, call_type="aresponses") assert responses_output_estimate > 0, ( "Responses API call was misclassified as an embedding and reserved zero output tokens" ) @@ -2187,9 +2039,7 @@ async def test_responses_api_usage_reconciles_using_input_output_tokens_fields( for op in increment_list: increments.append({"key": op["key"], "increment": op["increment_value"]}) - handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( - mock_increment - ) + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_increment await handler.async_log_success_event( kwargs=mock_kwargs, @@ -2251,7 +2101,7 @@ async def test_itpm_reservation_accounts_for_audio_content_not_just_text(rate_li ], } - with pytest.raises(Exception, match='Rate limit exceeded for model_per_project_itpm') as exc_info: + with pytest.raises(Exception, match="Rate limit exceeded for model_per_project_itpm") as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=cache, @@ -2273,7 +2123,7 @@ def test_audio_token_estimate_scales_with_payload_size(): to exhaust ITPM quota while reserving almost nothing. The estimate must now grow proportionally with the base64 payload size - (len(b64) * 3 // 4 // _AUDIO_BYTES_PER_TOKEN), floored at + (len(b64) * 3 // 4 // AUDIO_BYTES_PER_TOKEN), floored at DEFAULT_AUDIO_TOKEN_ESTIMATE so reference-only blocks and genuinely short clips still get a non-trivial reservation. @@ -2301,9 +2151,7 @@ def test_audio_token_estimate_scales_with_payload_size(): no_data_block = {"type": "input_audio", "input_audio": {"format": "wav"}} large_estimate = RateLimitHandler._estimate_audio_block_tokens(large_block) - very_large_estimate = RateLimitHandler._estimate_audio_block_tokens( - very_large_block - ) + very_large_estimate = RateLimitHandler._estimate_audio_block_tokens(very_large_block) small_estimate = RateLimitHandler._estimate_audio_block_tokens(small_block) no_data_estimate = RateLimitHandler._estimate_audio_block_tokens(no_data_block) @@ -2311,7 +2159,7 @@ def test_audio_token_estimate_scales_with_payload_size(): f"Large payload ({large_estimate}) must reserve more than small payload " f"({small_estimate}); flat-rate bug is back" ) - assert very_large_estimate == len(very_large_b64) * 3 // 4 // _AUDIO_BYTES_PER_TOKEN + assert very_large_estimate == len(very_large_b64) * 3 // 4 // AUDIO_BYTES_PER_TOKEN assert very_large_estimate > 6_000 assert no_data_estimate >= 300, ( f"Reference-only block (no data) must use the DEFAULT_AUDIO_TOKEN_ESTIMATE floor; got {no_data_estimate}" @@ -2362,7 +2210,7 @@ async def test_itpm_rejects_large_audio_payload_that_would_pass_flat_estimate( ], } - with pytest.raises(Exception, match='Rate limit exceeded for model_per_project_itpm') as exc_info: + with pytest.raises(Exception, match="Rate limit exceeded for model_per_project_itpm") as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=cache, @@ -2520,39 +2368,21 @@ async def test_itpm_otpm_reservation_is_kept_on_stream_disconnect(rate_limiter): stash = get_request_stash() assert stash is not None - assert stash.itpm_reserved_tokens > 0, ( - "pre-call hook must stash an ITPM reservation" - ) - assert stash.otpm_reserved_tokens > 0, ( - "pre-call hook must stash an OTPM reservation" - ) + assert stash.itpm_reserved_tokens > 0, "pre-call hook must stash an ITPM reservation" + assert stash.otpm_reserved_tokens > 0, "pre-call hook must stash an OTPM reservation" increment_calls: list[dict] = [] async def mock_increment(increment_list, litellm_parent_otel_span=None): for op in increment_list: - increment_calls.append( - {"key": op["key"], "increment": op["increment_value"]} - ) + increment_calls.append({"key": op["key"], "increment": op["increment_value"]}) - handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( - mock_increment - ) + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_increment - await handler.async_release_max_parallel_requests_on_disconnect( - user_api_key_dict=user_api_key_dict - ) + await handler.async_release_max_parallel_requests_on_disconnect(user_api_key_dict=user_api_key_dict) - itpm_refunds = [ - c - for c in increment_calls - if "model_per_project_itpm" in c["key"] and c["increment"] < 0 - ] - otpm_refunds = [ - c - for c in increment_calls - if "model_per_project_otpm" in c["key"] and c["increment"] < 0 - ] + itpm_refunds = [c for c in increment_calls if "model_per_project_itpm" in c["key"] and c["increment"] < 0] + otpm_refunds = [c for c in increment_calls if "model_per_project_otpm" in c["key"] and c["increment"] < 0] assert not itpm_refunds assert not otpm_refunds @@ -2623,7 +2453,7 @@ async def test_explicit_zero_output_responses_call_reserves_effective_provider_m }, ) - with pytest.raises(Exception, match='Rate limit exceeded for model_per_project_otpm') as exc_info: + with pytest.raises(Exception, match="Rate limit exceeded for model_per_project_otpm") as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=cache, @@ -2851,7 +2681,7 @@ async def test_otpm_rejection_releases_stashed_parallel_slot(rate_limiter): "rate_limit": {"tokens_per_unit": 5, "window_size": 60}, } - with pytest.raises(Exception, match='Rate limit exceeded for model_per_project_otpm') as exc_info: + with pytest.raises(Exception, match="Rate limit exceeded for model_per_project_otpm") as exc_info: await handler._reserve_project_io_tokens_or_raise( descriptors=[otpm_descriptor], data=data, @@ -2899,9 +2729,7 @@ async def test_itpm_only_status_stored_when_no_prior_rate_limit_response(rate_li stash = get_request_stash() assert stash is not None stored = stash.rate_limit_response - assert stored is not None, ( - "ITPM status must be stored in litellm_proxy_rate_limit_response" - ) + assert stored is not None, "ITPM status must be stored in litellm_proxy_rate_limit_response" assert stored.get("statuses"), "Stored response must contain statuses" @@ -2923,9 +2751,7 @@ def test_resolve_io_token_usage_responses_api_with_cached_tokens(rate_limiter): input_tokens_details=InputTokensDetails(cached_tokens=25), ), ) - billable_input, completion_tokens, resolved = ( - handler._resolve_io_token_reconcile_usage(response_obj) - ) + billable_input, completion_tokens, resolved = handler._resolve_io_token_reconcile_usage(response_obj) assert resolved is True assert billable_input == 75, f"Expected 100 - 25 cached = 75, got {billable_input}" @@ -2946,9 +2772,7 @@ def test_resolve_io_token_usage_dict_format(rate_limiter): "prompt_tokens_details": {"cached_tokens": 20}, } ) - billable_input, completion_tokens, resolved = ( - handler._resolve_io_token_reconcile_usage(response_obj) - ) + billable_input, completion_tokens, resolved = handler._resolve_io_token_reconcile_usage(response_obj) assert resolved is True assert billable_input == 60, f"Expected 80 - 20 cached = 60, got {billable_input}" @@ -2964,9 +2788,7 @@ def test_resolve_io_token_usage_unknown_type_returns_unresolved(rate_limiter): handler, _cache = rate_limiter response_obj = ModelResponse.model_construct(usage=42) - billable_input, completion_tokens, resolved = ( - handler._resolve_io_token_reconcile_usage(response_obj) - ) + billable_input, completion_tokens, resolved = handler._resolve_io_token_reconcile_usage(response_obj) assert resolved is False assert billable_input == 0 @@ -2997,9 +2819,7 @@ def test_zero_usage_keeps_reservations_unless_measured_fallback_exists( stash.otpm_reserved_tokens = 60 stash.otpm_reserved_scopes = frozenset({otpm_scope}) kwargs = {} if combined_usage is None else {"combined_usage_object": combined_usage} - response_obj = ModelResponse( - usage=Usage(prompt_tokens=0, completion_tokens=0, total_tokens=0) - ) + response_obj = ModelResponse(usage=Usage(prompt_tokens=0, completion_tokens=0, total_tokens=0)) operations = handler._build_io_token_reservation_ops(kwargs, response_obj) @@ -3096,9 +2916,7 @@ async def test_post_call_failure_skips_rpm_only_descriptor_in_tpm_refund(rate_li call_type="", ) - rpm_tokens_key = handler.create_rate_limit_keys( - key="api_key", value=api_key, rate_limit_type="tokens" - ) + rpm_tokens_key = handler.create_rate_limit_keys(key="api_key", value=api_key, rate_limit_type="tokens") await handler.async_post_call_failure_hook( request_data=data, @@ -3106,9 +2924,7 @@ async def test_post_call_failure_skips_rpm_only_descriptor_in_tpm_refund(rate_li user_api_key_dict=user_api_key_dict, ) - api_key_tokens_after = int( - await cache.async_get_cache(key=rpm_tokens_key, local_only=True) or 0 - ) + api_key_tokens_after = int(await cache.async_get_cache(key=rpm_tokens_key, local_only=True) or 0) assert api_key_tokens_after >= 0, ( f"RPM-only api_key scope must not receive a negative TPM refund; got {api_key_tokens_after}" ) @@ -3297,7 +3113,7 @@ async def test_rerank_query_and_documents_enforce_project_itpm( project_metadata={"model_itpm_limit": {"rerank-model": 100}}, ) - with pytest.raises(Exception, match='Rate limit exceeded for model_per_project_itpm') as exc_info: + with pytest.raises(Exception, match="Rate limit exceeded for model_per_project_itpm") as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=cache, @@ -3310,11 +3126,7 @@ async def test_rerank_query_and_documents_enforce_project_itpm( ) assert getattr(exc_info.value, "status_code", None) == 429 - assert captured["text"] == ( - "Which document is most relevant?\n" - "first document\n" - "{'text': 'second document'}" - ) + assert captured["text"] == ("Which document is most relevant?\nfirst document\n{'text': 'second document'}") def test_rerank_input_estimate_falls_back_to_character_count( @@ -3333,20 +3145,21 @@ def test_rerank_input_estimate_falls_back_to_character_count( monkeypatch.setattr("litellm.token_counter", token_counter) rerank_text = handler._rerank_input_to_text(data) - assert handler._estimate_precise_input_tokens( - data, - model="custom-rerank-model", - call_type="rerank", - ) == len(rerank_text) // 4 + assert ( + handler._estimate_precise_input_tokens( + data, + model="custom-rerank-model", + call_type="rerank", + ) + == len(rerank_text) // 4 + ) @pytest.mark.parametrize( ("response_obj", "expected"), [ ( - RerankResponse( - meta={"tokens": {"input_tokens": 42, "output_tokens": 3}} - ), + RerankResponse(meta={"tokens": {"input_tokens": 42, "output_tokens": 3}}), (42, 3, True), ), ( @@ -3387,9 +3200,7 @@ def test_split_quota_helpers_handle_non_mapping_inputs(rate_limiter): assert handler._get_explicit_output_cap(object(), None) is None assert handler.get_output_candidate_count(object()) == 1 assert handler.get_output_candidate_count({"n": 1e309}) == 1 - assert ( - handler._get_explicit_output_cap({"max_output_tokens": []}, "responses") is None - ) + assert handler._get_explicit_output_cap({"max_output_tokens": []}, "responses") is None assert handler._apply_implicit_output_cap(object(), 100, "responses") is None assert handler._estimate_input_and_output_tokens(object()) == (0, 0) assert handler._build_io_token_reservation_ops(object(), object()) == () @@ -3407,9 +3218,7 @@ def test_split_quota_helpers_handle_non_mapping_inputs(rate_limiter): ({"generationConfig": {"maxOutputTokens": "oops"}}, "agenerate_content", None), ], ) -def test_get_explicit_output_cap_tolerates_unparseable_values( - rate_limiter, data, call_type, expected -): +def test_get_explicit_output_cap_tolerates_unparseable_values(rate_limiter, data, call_type, expected): """A client-supplied cap the proxy cannot parse must fall back to the no-cap output estimate instead of raising ValueError and 500ing the request before it ever reaches the provider.""" @@ -3504,9 +3313,7 @@ def test_split_token_estimate_selects_endpoint_input(rate_limiter, call_type, da def test_split_quota_multimodal_guards_handle_non_mapping_inputs(rate_limiter): handler, _cache = rate_limiter - assert handler._estimate_audio_block_tokens( - object() - ) == handler._estimate_audio_block_tokens({}) + assert handler._estimate_audio_block_tokens(object()) == handler._estimate_audio_block_tokens({}) assert handler._responses_input_to_chat_messages(object()) == () assert handler._estimate_precise_input_tokens(object(), model=None) == 0 @@ -3539,10 +3346,7 @@ def test_precise_input_estimate_selects_endpoint_text( monkeypatch.setattr("litellm.token_counter", token_counter) - assert ( - handler._estimate_precise_input_tokens(data, model="test", call_type=call_type) - == 7 - ) + assert handler._estimate_precise_input_tokens(data, model="test", call_type=call_type) == 7 assert captured["messages"] is None assert captured["text"] == expected_text @@ -3585,9 +3389,7 @@ async def test_streaming_combined_usage_reconciles_project_io_reservations( async def capture_increments(increment_list, **_kwargs): increments.extend(increment_list) - handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( - capture_increments - ) + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = capture_increments await handler.async_log_success_event( kwargs=kwargs, @@ -3596,16 +3398,8 @@ async def test_streaming_combined_usage_reconciles_project_io_reservations( end_time=datetime.now(), ) - itpm_adjustments = [ - operation - for operation in increments - if PROJECT_ITPM_DESCRIPTOR_KEY in operation["key"] - ] - otpm_adjustments = [ - operation - for operation in increments - if PROJECT_OTPM_DESCRIPTOR_KEY in operation["key"] - ] + itpm_adjustments = [operation for operation in increments if PROJECT_ITPM_DESCRIPTOR_KEY in operation["key"]] + otpm_adjustments = [operation for operation in increments if PROJECT_OTPM_DESCRIPTOR_KEY in operation["key"]] assert [operation["increment_value"] for operation in itpm_adjustments] == [-60] assert [operation["increment_value"] for operation in otpm_adjustments] == [-45] @@ -3614,13 +3408,9 @@ def test_aggregate_only_combined_usage_reconciles_project_io_reservations(rate_l handler, _cache = rate_limiter stash = get_or_create_request_stash() stash.itpm_reserved_tokens = 100 - stash.itpm_reserved_scopes = frozenset( - {(PROJECT_ITPM_DESCRIPTOR_KEY, "project:model")} - ) + stash.itpm_reserved_scopes = frozenset({(PROJECT_ITPM_DESCRIPTOR_KEY, "project:model")}) stash.otpm_reserved_tokens = 80 - stash.otpm_reserved_scopes = frozenset( - {(PROJECT_OTPM_DESCRIPTOR_KEY, "project:model")} - ) + stash.otpm_reserved_scopes = frozenset({(PROJECT_OTPM_DESCRIPTOR_KEY, "project:model")}) kwargs = { "combined_usage_object": Usage(total_tokens=55), } @@ -3643,9 +3433,7 @@ def test_raw_split_usage_dict_reconciles_project_io_tokens(rate_limiter): @pytest.mark.asyncio -async def test_post_call_success_hook_contains_header_merge_failures( - rate_limiter, monkeypatch -): +async def test_post_call_success_hook_contains_header_merge_failures(rate_limiter, monkeypatch): handler, _cache = rate_limiter response = ModelResponse() response._hidden_params = {} From 2ca3143b34b10a203b414faea6366e945df74f99 Mon Sep 17 00:00:00 2001 From: Mihidum Hettiyahandi <55163074+mihidumh@users.noreply.github.com> Date: Wed, 9 Sep 2026 09:44:41 +1000 Subject: [PATCH 4/6] fix(proxy): floor budget reservation for input_audio requests at the serialised fallback Before #38459 an input_audio request raised inside token_counter and the budget reservation counted it from the serialised messages (base64 payload included). Making the blocks countable turned that into a size-derived estimate at a deliberately low bitrate, priced at the text rate, so a large or highly compressed audio payload could be admitted against a budget more cheaply than before. Floor audio-bearing requests at the serialised fallback, the same compensate-in-the-caller pattern the batch limiter uses; text-only requests are untouched. --- litellm/proxy/spend_tracking/input_tokens.py | 18 +++++++-- .../proxy/spend_tracking/test_input_tokens.py | 39 +++++++++++++++++++ 2 files changed, 54 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/spend_tracking/input_tokens.py b/litellm/proxy/spend_tracking/input_tokens.py index 6c7083fb6db..a80bbb79d4a 100644 --- a/litellm/proxy/spend_tracking/input_tokens.py +++ b/litellm/proxy/spend_tracking/input_tokens.py @@ -19,6 +19,7 @@ from typing import Final import litellm from litellm._logging import verbose_proxy_logger +from litellm.litellm_core_utils.token_counter import messages_contain_input_audio_content_blocks from litellm.rust_bridge import runtime from litellm.rust_bridge.catalog import Route, RouteContext from litellm.rust_bridge.token_counter import ( @@ -128,15 +129,26 @@ def _count_input_tokens_for_models( def count_input_tokens_for_model(request_body: dict, model: str) -> int | None: try: if "messages" in request_body: + messages: Final = request_body.get("messages") try: - return litellm.token_counter( + counted: Final = litellm.token_counter( model=model, - messages=request_body.get("messages") or (), + messages=messages or (), tools=request_body.get("tools"), tool_choice=request_body.get("tool_choice"), ) except ValueError: - return _count_text_tokens(model=model, text=request_body.get("messages")) + return _count_text_tokens(model=model, text=messages) + # An ``input_audio`` block counts as a size-derived estimate at a + # deliberately low assumed bitrate, and this reservation prices it + # at the text rate. Before such blocks were countable (#38459) an + # audio request RAISED in token_counter and reserved from the + # serialised-messages fallback above; floor it there again so a + # large or highly compressed audio payload cannot be admitted + # against a budget more cheaply than it was before. + if messages_contain_input_audio_content_blocks(messages): + return max(counted, _count_text_tokens(model=model, text=messages)) + return counted if "prompt" in request_body: return _count_text_tokens(model=model, text=request_body.get("prompt")) if "input" in request_body: diff --git a/tests/test_litellm/proxy/spend_tracking/test_input_tokens.py b/tests/test_litellm/proxy/spend_tracking/test_input_tokens.py index a1da6c79cd7..3439e0a16af 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_input_tokens.py +++ b/tests/test_litellm/proxy/spend_tracking/test_input_tokens.py @@ -2,6 +2,7 @@ from __future__ import annotations +import base64 from types import MappingProxyType from typing import Final @@ -9,6 +10,7 @@ import pytest from litellm.proxy.spend_tracking.input_tokens import ( TOKENIZE_OFF_EVENT_LOOP_MIN_CHARS, + _count_text_tokens, count_input_tokens, count_input_tokens_for_model, ) @@ -27,3 +29,40 @@ async def test_large_input_is_still_counted() -> None: assert counts[CL100K_MODEL] == count_input_tokens_for_model(request_body=request_body, model=CL100K_MODEL) assert isinstance(counts, MappingProxyType) + + +def test_input_audio_requests_reserve_at_least_the_serialised_fallback() -> None: + """ + Budget reservation counts an ``input_audio`` block as a size-derived + estimate at a deliberately low bitrate, priced at the text rate. Before + #38459 the same request raised inside ``token_counter`` and reserved + from the serialised-messages fallback, which tokenises the base64 + payload itself. Floor audio-bearing requests at that fallback so a + caller cannot be admitted against a budget more cheaply than before + the blocks became countable (compressed audio carries far more duration + per byte than the estimate assumes). + """ + model: Final = "gpt-4o-audio-preview" + audio_b64: Final = base64.b64encode(bytes(range(256)) * 400).decode() + audio_messages: Final = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "Transcribe this recording."}, + {"type": "input_audio", "input_audio": {"data": audio_b64, "format": "mp3"}}, + ], + } + ] + text_messages: Final = [{"role": "user", "content": "Transcribe this recording."}] + + audio_count: Final = count_input_tokens_for_model(request_body={"messages": audio_messages}, model=model) + fallback: Final = _count_text_tokens(model=model, text=audio_messages) + text_count: Final = count_input_tokens_for_model(request_body={"messages": text_messages}, model=model) + + assert audio_count is not None, "an audio-bearing request must still be countable" + assert fallback > 0, "the serialised fallback must see the base64 payload" + assert audio_count >= fallback, ( + f"audio request reserved {audio_count} tokens, below the pre-#38459 fallback of {fallback}" + ) + assert text_count is not None + assert text_count < fallback, "a text-only request must not be floored" From 8197b319e766a6dc73e3e3018675d7bb223139e9 Mon Sep 17 00:00:00 2001 From: Mihidum Hettiyahandi <55163074+mihidumh@users.noreply.github.com> Date: Thu, 17 Sep 2026 09:21:32 +1000 Subject: [PATCH 5/6] test(proxy): keep test_tpm_concurrent.py to the constant import change The earlier refactor commit reformatted the whole test file, which the CI format check (litellm/** only) does not ask for. Keep only the import of AUDIO_BYTES_PER_TOKEN from litellm.constants and its two uses, so the PR diff for this file is 3 lines instead of +134/-346. --- .../proxy/hooks/test_tpm_concurrent.py | 474 +++++++++++++----- 1 file changed, 343 insertions(+), 131 deletions(-) diff --git a/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py b/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py index 5146b53eeca..82a4e269c3b 100644 --- a/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py +++ b/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py @@ -151,7 +151,7 @@ async def test_no_leak_on_over_limit_rejection(rate_limiter): f"estimated={estimated}, limit={user_api_key_dict.tpm_limit}" ) - with pytest.raises(Exception, match="Limit type: tokens\\. Current limit") as exc_info: + with pytest.raises(Exception, match='Limit type: tokens\\. Current limit') as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=cache, @@ -168,7 +168,8 @@ async def test_no_leak_on_over_limit_rejection(rate_limiter): cached_value = await cache.async_get_cache(key=counter_key, local_only=True) cached_int = int(cached_value or 0) assert cached_int < estimated, ( - f"Reservation leaked: counter={cached_int} after rejection of an estimated_tokens={estimated} reservation." + f"Reservation leaked: counter={cached_int} after rejection of an " + f"estimated_tokens={estimated} reservation." ) @@ -217,7 +218,9 @@ async def test_token_adjustment_on_success(rate_limiter): } ) - handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_increment + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( + mock_increment + ) await handler.async_log_success_event( kwargs=mock_kwargs, @@ -229,7 +232,8 @@ async def test_token_adjustment_on_success(rate_limiter): token_adjustments = [i for i in increments if "tokens" in i["key"]] assert any(i["increment"] == -50 for i in token_adjustments), ( - f"Expected a -50 token adjustment (50 actual - 100 reserved) but got: {token_adjustments}" + f"Expected a -50 token adjustment (50 actual - 100 reserved) but got: " + f"{token_adjustments}" ) @@ -266,7 +270,9 @@ async def test_token_release_on_failure(rate_limiter): } ) - handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_increment + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( + mock_increment + ) await handler.async_log_failure_event( kwargs=mock_kwargs, @@ -277,9 +283,9 @@ async def test_token_release_on_failure(rate_limiter): token_releases = [i for i in increments if "tokens" in i["key"]] - assert any(i["increment"] == -100 for i in token_releases), ( - f"Expected the full reservation (-100) to be released, got: {token_releases}" - ) + assert any( + i["increment"] == -100 for i in token_releases + ), f"Expected the full reservation (-100) to be released, got: {token_releases}" @pytest.mark.asyncio @@ -323,7 +329,9 @@ async def test_model_scope_refund_targets_reserved_model(rate_limiter): for op in increment_list: increments.append({"key": op["key"], "increment": op["increment_value"]}) - handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_increment + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( + mock_increment + ) await handler.async_log_failure_event( kwargs=mock_kwargs, @@ -344,7 +352,8 @@ async def test_model_scope_refund_targets_reserved_model(rate_limiter): f"{[i['key'] for i in increments]}" ) assert matching[0]["increment"] == -100, ( - f"Expected full -100 refund on model_per_team counter, got {matching[0]['increment']}" + f"Expected full -100 refund on model_per_team counter, got " + f"{matching[0]['increment']}" ) @@ -365,7 +374,9 @@ async def test_should_rate_limit_does_not_inflate_tokens_counter(rate_limiter): tpm_limit=10_000, ) - tokens_counter_key = handler.create_rate_limit_keys(key="api_key", value=api_key, rate_limit_type="tokens") + tokens_counter_key = handler.create_rate_limit_keys( + key="api_key", value=api_key, rate_limit_type="tokens" + ) data = { "model": "gpt-3.5-turbo", @@ -382,7 +393,9 @@ async def test_should_rate_limit_does_not_inflate_tokens_counter(rate_limiter): call_type="", ) - cached = int(await cache.async_get_cache(key=tokens_counter_key, local_only=True) or 0) + cached = int( + await cache.async_get_cache(key=tokens_counter_key, local_only=True) or 0 + ) # The :tokens counter should reflect ONLY the reservation amount — not # an additional +1 from the should_rate_limit pre-pass. @@ -474,7 +487,9 @@ async def test_org_scope_refund_on_failure(rate_limiter): for op in increment_list: increments.append({"key": op["key"], "increment": op["increment_value"]}) - handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_increment + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( + mock_increment + ) await handler.async_log_failure_event( kwargs=mock_kwargs, @@ -483,13 +498,17 @@ async def test_org_scope_refund_on_failure(rate_limiter): end_time=datetime.now(), ) - expected_org_key = handler.create_rate_limit_keys(key="organization", value=org_id, rate_limit_type="tokens") + expected_org_key = handler.create_rate_limit_keys( + key="organization", value=org_id, rate_limit_type="tokens" + ) matching = [i for i in increments if i["key"] == expected_org_key] assert matching, ( f"Expected a refund on the org tokens counter ({expected_org_key}) " f"but got keys: {[i['key'] for i in increments]}" ) - assert matching[0]["increment"] == -100, f"Expected full -100 refund on org counter, got {matching[0]['increment']}" + assert ( + matching[0]["increment"] == -100 + ), f"Expected full -100 refund on org counter, got {matching[0]['increment']}" @pytest.mark.asyncio @@ -532,7 +551,9 @@ async def test_org_scope_reconciled_on_success(rate_limiter): for op in increment_list: increments.append({"key": op["key"], "increment": op["increment_value"]}) - handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_increment + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( + mock_increment + ) await handler.async_log_success_event( kwargs=mock_kwargs, @@ -541,14 +562,17 @@ async def test_org_scope_reconciled_on_success(rate_limiter): end_time=datetime.now(), ) - expected_org_key = handler.create_rate_limit_keys(key="organization", value=org_id, rate_limit_type="tokens") + expected_org_key = handler.create_rate_limit_keys( + key="organization", value=org_id, rate_limit_type="tokens" + ) matching = [i for i in increments if i["key"] == expected_org_key] assert matching, ( f"Expected a reconciliation op on the org tokens counter " f"({expected_org_key}), got keys: {[i['key'] for i in increments]}" ) assert matching[0]["increment"] == -50, ( - f"Expected -50 delta on org counter (50 actual - 100 reserved), got {matching[0]['increment']}" + f"Expected -50 delta on org counter (50 actual - 100 reserved), got " + f"{matching[0]['increment']}" ) @@ -559,7 +583,9 @@ async def test_estimate_tokens_uses_max_tokens_when_explicit(rate_limiter): estimate = handler._estimate_tokens_for_request( data={ - "messages": [{"role": "user", "content": "abcd" * 4}], # 16 chars ~ 4 tokens + "messages": [ + {"role": "user", "content": "abcd" * 4} + ], # 16 chars ~ 4 tokens "max_tokens": 25, } ) @@ -580,11 +606,15 @@ async def test_estimate_tokens_honors_explicit_zero_max_tokens(rate_limiter): estimate = handler._estimate_tokens_for_request( data={ - "messages": [{"role": "user", "content": "abcd" * 4}], # 16 chars ~ 4 tokens + "messages": [ + {"role": "user", "content": "abcd" * 4} + ], # 16 chars ~ 4 tokens "max_tokens": 0, } ) - assert estimate == 4, f"expected input-only reservation (4) for an explicit max_tokens=0, got {estimate}" + assert estimate == 4, ( + f"expected input-only reservation (4) for an explicit max_tokens=0, got {estimate}" + ) @pytest.mark.asyncio @@ -630,7 +660,9 @@ async def test_contentless_request_reserves_minimum(rate_limiter): api_key = hash_token("sk-contentless") user_api_key_dict = UserAPIKeyAuth(api_key=api_key, tpm_limit=2) - counter_key = handler.create_rate_limit_keys(key="api_key", value=api_key, rate_limit_type="tokens") + counter_key = handler.create_rate_limit_keys( + key="api_key", value=api_key, rate_limit_type="tokens" + ) # Two contentless requests should consume two slots of the 2-token # budget. The third must 429. @@ -642,14 +674,19 @@ async def test_contentless_request_reserves_minimum(rate_limiter): data=data, call_type="", ) - assert get_request_stash().reserved_tokens == 1, "Contentless request should reserve the floor of 1 token" + assert ( + get_request_stash().reserved_tokens == 1 + ), "Contentless request should reserve the floor of 1 token" - counter_after_two = int(await cache.async_get_cache(key=counter_key, local_only=True) or 0) + counter_after_two = int( + await cache.async_get_cache(key=counter_key, local_only=True) or 0 + ) assert counter_after_two == 2, ( - f"After two contentless requests at the floor, the api_key tokens counter should be 2, got {counter_after_two}" + f"After two contentless requests at the floor, the api_key tokens " + f"counter should be 2, got {counter_after_two}" ) - with pytest.raises(Exception, match="Limit type: tokens\\. Current limit") as exc_info: + with pytest.raises(Exception, match='Limit type: tokens\\. Current limit') as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=cache, @@ -657,7 +694,8 @@ async def test_contentless_request_reserves_minimum(rate_limiter): call_type="", ) assert getattr(exc_info.value, "status_code", None) == 429, ( - "Third contentless request must be rate-limited; pre-fix it would have bypassed the TPM check entirely." + "Third contentless request must be rate-limited; pre-fix it would " + "have bypassed the TPM check entirely." ) @@ -734,8 +772,12 @@ async def test_reservation_released_on_proxy_rejection(rate_limiter): reserved = get_request_stash().reserved_tokens assert reserved > 0 - counter_key = handler.create_rate_limit_keys(key="api_key", value=api_key, rate_limit_type="tokens") - counter_after_reserve = int(await cache.async_get_cache(key=counter_key, local_only=True) or 0) + counter_key = handler.create_rate_limit_keys( + key="api_key", value=api_key, rate_limit_type="tokens" + ) + counter_after_reserve = int( + await cache.async_get_cache(key=counter_key, local_only=True) or 0 + ) assert counter_after_reserve == reserved # Simulate a downstream guardrail rejecting the request. @@ -745,12 +787,16 @@ async def test_reservation_released_on_proxy_rejection(rate_limiter): user_api_key_dict=user_api_key_dict, ) - counter_after_release = int(await cache.async_get_cache(key=counter_key, local_only=True) or 0) + counter_after_release = int( + await cache.async_get_cache(key=counter_key, local_only=True) or 0 + ) assert counter_after_release == 0, ( - f"Reservation leaked: counter={counter_after_release} after proxy-level rejection refund (expected 0)." + f"Reservation leaked: counter={counter_after_release} after " + f"proxy-level rejection refund (expected 0)." ) assert get_request_stash().reservation_released is True, ( - "Released flag must be set to prevent async_log_failure_event from double-refunding." + "Released flag must be set to prevent " + "async_log_failure_event from double-refunding." ) @@ -771,7 +817,9 @@ async def test_reservation_release_idempotent(rate_limiter): for op in increment_list: increments.append({"key": op["key"], "increment": op["increment_value"]}) - handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_increment + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( + mock_increment + ) # Both hooks read the same per-request ContextVar stash: the # post-call-failure-hook flips reservation_released on it, and the @@ -852,7 +900,9 @@ async def test_unreserved_scopes_charged_actual_not_delta_on_success(rate_limite for op in increment_list: increments.append({"key": op["key"], "increment": op["increment_value"]}) - handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_increment + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( + mock_increment + ) await handler.async_log_success_event( kwargs=mock_kwargs, @@ -861,14 +911,19 @@ async def test_unreserved_scopes_charged_actual_not_delta_on_success(rate_limite end_time=datetime.now(), ) - api_key_token_key = handler.create_rate_limit_keys(key="api_key", value=api_key, rate_limit_type="tokens") - team_token_key = handler.create_rate_limit_keys(key="team", value=team_id, rate_limit_type="tokens") + api_key_token_key = handler.create_rate_limit_keys( + key="api_key", value=api_key, rate_limit_type="tokens" + ) + team_token_key = handler.create_rate_limit_keys( + key="team", value=team_id, rate_limit_type="tokens" + ) api_key_ops = [i for i in increments if i["key"] == api_key_token_key] team_ops = [i for i in increments if i["key"] == team_token_key] assert api_key_ops and api_key_ops[0]["increment"] == -50, ( - f"Reserved api_key scope must reconcile via delta (50-100=-50), got {api_key_ops}" + f"Reserved api_key scope must reconcile via delta (50-100=-50), " + f"got {api_key_ops}" ) assert team_ops and team_ops[0]["increment"] == 50, ( f"Unreserved team scope must be charged full actual (+50), not the " @@ -907,7 +962,9 @@ async def test_unreserved_scopes_not_refunded_on_failure(rate_limiter): for op in increment_list: increments.append({"key": op["key"], "increment": op["increment_value"]}) - handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_increment + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( + mock_increment + ) await handler.async_log_failure_event( kwargs=mock_kwargs, @@ -916,16 +973,23 @@ async def test_unreserved_scopes_not_refunded_on_failure(rate_limiter): end_time=datetime.now(), ) - team_token_key = handler.create_rate_limit_keys(key="team", value=team_id, rate_limit_type="tokens") - api_key_token_key = handler.create_rate_limit_keys(key="api_key", value=api_key, rate_limit_type="tokens") + team_token_key = handler.create_rate_limit_keys( + key="team", value=team_id, rate_limit_type="tokens" + ) + api_key_token_key = handler.create_rate_limit_keys( + key="api_key", value=api_key, rate_limit_type="tokens" + ) team_ops = [i for i in increments if i["key"] == team_token_key] api_key_ops = [i for i in increments if i["key"] == api_key_token_key] - assert not team_ops, f"Unreserved team scope must NOT be refunded (would drift negative), got {team_ops}" - assert api_key_ops and api_key_ops[0]["increment"] == -100, ( - f"Reserved api_key scope must be refunded -100, got {api_key_ops}" + assert not team_ops, ( + f"Unreserved team scope must NOT be refunded (would drift negative), " + f"got {team_ops}" ) + assert ( + api_key_ops and api_key_ops[0]["increment"] == -100 + ), f"Reserved api_key scope must be refunded -100, got {api_key_ops}" @pytest.mark.asyncio @@ -960,7 +1024,9 @@ async def test_token_rate_limit_headers_present_in_stored_response(rate_limiter) ) response = get_request_stash().rate_limit_response - assert isinstance(response, dict), "Expected the stashed rate-limit response to be set after pre-call" + assert isinstance( + response, dict + ), "Expected the stashed rate-limit response to be set after pre-call" statuses = response.get("statuses") or [] token_statuses = [s for s in statuses if s.get("rate_limit_type") == "tokens"] @@ -972,7 +1038,8 @@ async def test_token_rate_limit_headers_present_in_stored_response(rate_limiter) f"statuses: {[(s.get('descriptor_key'), s.get('rate_limit_type')) for s in statuses]}" ) assert request_statuses, ( - "RPM rate-limit status was clobbered by the TPM merge — both must coexist in the stored response." + "RPM rate-limit status was clobbered by the TPM merge — both must " + "coexist in the stored response." ) # The token status carries the limit and a positive remaining budget. @@ -1000,7 +1067,9 @@ async def test_estimate_tokens_floor_caps_at_smallest_configured_tpm(rate_limite ) # input ~= 5//4 = 1 token; output floor capped at 1000//4 = 250; # total ~= 251 (well under 1000). - assert estimate <= 1000 // 2, f"With TPM=1000, reservation must stay well under the limit; got {estimate}" + assert ( + estimate <= 1000 // 2 + ), f"With TPM=1000, reservation must stay well under the limit; got {estimate}" assert estimate >= 1, "Estimate must be at least the call-site floor of 1" @@ -1070,7 +1139,10 @@ async def test_small_tpm_cap_admits_no_max_tokens_request(rate_limiter): reserved = get_request_stash().reserved_tokens assert reserved > 0, "Reservation should have been stashed" - assert reserved <= 1000 // 2, f"Capped floor must keep the reservation well under the 1000 TPM cap; got {reserved}" + assert reserved <= 1000 // 2, ( + f"Capped floor must keep the reservation well under the 1000 TPM " + f"cap; got {reserved}" + ) @pytest.mark.asyncio @@ -1105,7 +1177,8 @@ async def test_small_tpm_cap_injects_matching_max_tokens(rate_limiter): ) assert data.get("max_tokens") == 1000 // 4, ( - f"Capped floor must be written to max_tokens to bound the actual model output; got {data.get('max_tokens')}" + f"Capped floor must be written to max_tokens to bound the actual " + f"model output; got {data.get('max_tokens')}" ) @@ -1138,7 +1211,10 @@ async def test_large_tpm_cap_does_not_inject_max_tokens(rate_limiter): call_type="", ) - assert "max_tokens" not in data, f"Large TPM caps should leave max_tokens alone; got {data.get('max_tokens')}" + assert "max_tokens" not in data, ( + f"Large TPM caps should leave max_tokens alone; got " + f"{data.get('max_tokens')}" + ) @pytest.mark.asyncio @@ -1218,9 +1294,13 @@ async def test_project_otpm_reservation_prevents_concurrent_bypass(rate_limiter) results = await asyncio.gather(*[make_request(i) for i in range(5)]) successful = [r for r in results if r["success"]] - rate_limited = [r for r in results if not r["success"] and r.get("status_code") == 429] + rate_limited = [ + r for r in results if not r["success"] and r.get("status_code") == 429 + ] - assert len(rate_limited) > 0, f"Expected some OTPM-rate-limited requests but all {len(successful)} succeeded." + assert len(rate_limited) > 0, ( + f"Expected some OTPM-rate-limited requests but all {len(successful)} succeeded." + ) @pytest.mark.asyncio @@ -1240,7 +1320,7 @@ async def test_project_otpm_rejects_multiple_completion_candidates(rate_limiter) "n": 10, } - with pytest.raises(Exception, match="Rate limit exceeded for model_per_project_otpm") as exc_info: + with pytest.raises(Exception, match='Rate limit exceeded for model_per_project_otpm') as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=cache, @@ -1268,7 +1348,7 @@ async def test_project_otpm_reserves_largest_conflicting_output_cap(rate_limiter "max_completion_tokens": 100, } - with pytest.raises(Exception, match="Rate limit exceeded for model_per_project_otpm") as exc_info: + with pytest.raises(Exception, match='Rate limit exceeded for model_per_project_otpm') as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=cache, @@ -1298,7 +1378,7 @@ async def test_project_otpm_rejects_google_genai_native_output_cap( project_metadata={"model_otpm_limit": {model: 50}}, ) - with pytest.raises(Exception, match="Rate limit exceeded for model_per_project_otpm") as exc_info: + with pytest.raises(Exception, match='Rate limit exceeded for model_per_project_otpm') as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=cache, @@ -1332,7 +1412,7 @@ async def test_project_otpm_rejects_google_genai_native_candidate_count( project_metadata={"model_otpm_limit": {model: 150}}, ) - with pytest.raises(Exception, match="Rate limit exceeded for model_per_project_otpm") as exc_info: + with pytest.raises(Exception, match='Rate limit exceeded for model_per_project_otpm') as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=cache, @@ -1421,7 +1501,7 @@ async def test_project_otpm_over_limit_rolls_back_itpm_reservation(rate_limiter) "max_tokens": 500, # blows past the 10-token OTPM limit } - with pytest.raises(Exception, match="Rate limit exceeded for model_per_project_otpm") as exc_info: + with pytest.raises(Exception, match='Rate limit exceeded for model_per_project_otpm') as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=cache, @@ -1431,7 +1511,9 @@ async def test_project_otpm_over_limit_rolls_back_itpm_reservation(rate_limiter) assert getattr(exc_info.value, "status_code", None) == 429 cached_value = await cache.async_get_cache(key=itpm_counter_key, local_only=True) - assert int(cached_value or 0) == 0, f"ITPM reservation leaked after OTPM rejection: counter={cached_value}" + assert int(cached_value or 0) == 0, ( + f"ITPM reservation leaked after OTPM rejection: counter={cached_value}" + ) @pytest.mark.asyncio @@ -1474,7 +1556,9 @@ async def test_project_itpm_reconciled_on_success_excludes_cached_tokens(rate_li for op in increment_list: increments.append({"key": op["key"], "increment": op["increment_value"]}) - handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_increment + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( + mock_increment + ) await handler.async_log_success_event( kwargs=mock_kwargs, @@ -1517,11 +1601,17 @@ async def test_project_reconciliation_does_not_decrement_later_window(): increments=[{"tokens": 100}], ) counter_key = handler.create_rate_limit_keys(*scope, rate_limit_type="tokens") - window_identity = next(identity for identity in reservation["reservation_windows"] if identity[0] == counter_key) + window_identity = next( + identity + for identity in reservation["reservation_windows"] + if identity[0] == counter_key + ) stash = get_or_create_request_stash() stash.itpm_reserved_tokens = 100 stash.itpm_reserved_scopes = frozenset({scope}) - stash.itpm_reserved_window_identities = frozenset({window_identity}) + stash.itpm_reserved_window_identities = frozenset( + {window_identity} + ) current_time += timedelta(seconds=61) later_reservation = await handler.atomic_check_and_increment_by_n( @@ -1532,7 +1622,9 @@ async def test_project_reconciliation_does_not_decrement_later_window(): await handler.async_log_success_event( kwargs={}, - response_obj=ModelResponse(usage=Usage(prompt_tokens=10, completion_tokens=0, total_tokens=10)), + response_obj=ModelResponse( + usage=Usage(prompt_tokens=10, completion_tokens=0, total_tokens=10) + ), start_time=current_time, end_time=current_time, ) @@ -1554,15 +1646,23 @@ async def test_project_reconciliation_decrements_its_active_window(rate_limiter) increments=[{"tokens": 100}], ) counter_key = handler.create_rate_limit_keys(*scope, rate_limit_type="tokens") - window_identity = next(identity for identity in reservation["reservation_windows"] if identity[0] == counter_key) + window_identity = next( + identity + for identity in reservation["reservation_windows"] + if identity[0] == counter_key + ) stash = get_or_create_request_stash() stash.itpm_reserved_tokens = 100 stash.itpm_reserved_scopes = frozenset({scope}) - stash.itpm_reserved_window_identities = frozenset({window_identity}) + stash.itpm_reserved_window_identities = frozenset( + {window_identity} + ) await handler.async_log_success_event( kwargs={}, - response_obj=ModelResponse(usage=Usage(prompt_tokens=10, completion_tokens=0, total_tokens=10)), + response_obj=ModelResponse( + usage=Usage(prompt_tokens=10, completion_tokens=0, total_tokens=10) + ), start_time=datetime.now(), end_time=datetime.now(), ) @@ -1650,7 +1750,9 @@ async def test_atomic_lua_response_carries_redis_window_identity(rate_limiter): ) assert response["statuses"][0]["limit_remaining"] == 75 - assert response["reservation_windows"] == frozenset({(counter_key, "1234", "redis")}) + assert response["reservation_windows"] == frozenset( + {(counter_key, "1234", "redis")} + ) @pytest.mark.asyncio @@ -1674,7 +1776,9 @@ async def test_project_itpm_otpm_released_on_failure(rate_limiter): for op in increment_list: increments.append({"key": op["key"], "increment": op["increment_value"]}) - handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_increment + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( + mock_increment + ) await handler.async_log_failure_event( kwargs=mock_kwargs, @@ -1718,7 +1822,9 @@ async def test_proxy_rejection_refunds_itpm_otpm_by_their_own_amount_not_combine data = { "model": "bedrock_mantle/claude-opus", - "messages": [{"role": "user", "content": "hello there, this is a test message"}], + "messages": [ + {"role": "user", "content": "hello there, this is a test message"} + ], "max_tokens": 60, } @@ -1745,9 +1851,15 @@ async def test_proxy_rejection_refunds_itpm_otpm_by_their_own_amount_not_combine rate_limit_type="tokens", ) - tpm_reserved = int(await cache.async_get_cache(key=tpm_counter_key, local_only=True) or 0) - itpm_reserved = int(await cache.async_get_cache(key=itpm_counter_key, local_only=True) or 0) - otpm_reserved = int(await cache.async_get_cache(key=otpm_counter_key, local_only=True) or 0) + tpm_reserved = int( + await cache.async_get_cache(key=tpm_counter_key, local_only=True) or 0 + ) + itpm_reserved = int( + await cache.async_get_cache(key=itpm_counter_key, local_only=True) or 0 + ) + otpm_reserved = int( + await cache.async_get_cache(key=otpm_counter_key, local_only=True) or 0 + ) assert tpm_reserved > 0 and itpm_reserved > 0 and otpm_reserved > 0 await handler.async_post_call_failure_hook( @@ -1756,13 +1868,23 @@ async def test_proxy_rejection_refunds_itpm_otpm_by_their_own_amount_not_combine user_api_key_dict=user_api_key_dict, ) - tpm_after = int(await cache.async_get_cache(key=tpm_counter_key, local_only=True) or 0) - itpm_after = int(await cache.async_get_cache(key=itpm_counter_key, local_only=True) or 0) - otpm_after = int(await cache.async_get_cache(key=otpm_counter_key, local_only=True) or 0) + tpm_after = int( + await cache.async_get_cache(key=tpm_counter_key, local_only=True) or 0 + ) + itpm_after = int( + await cache.async_get_cache(key=itpm_counter_key, local_only=True) or 0 + ) + otpm_after = int( + await cache.async_get_cache(key=otpm_counter_key, local_only=True) or 0 + ) assert tpm_after == 0, f"combined TPM counter leaked: {tpm_after}" - assert itpm_after == 0, f"ITPM counter corrupted by combined-amount refund: {itpm_after}" - assert otpm_after == 0, f"OTPM counter corrupted by combined-amount refund: {otpm_after}" + assert itpm_after == 0, ( + f"ITPM counter corrupted by combined-amount refund: {itpm_after}" + ) + assert otpm_after == 0, ( + f"OTPM counter corrupted by combined-amount refund: {otpm_after}" + ) @pytest.mark.asyncio @@ -1790,7 +1912,9 @@ async def test_proxy_rejection_refunds_itpm_otpm_only_reservation_with_no_combin data = { "model": "bedrock_mantle/claude-opus", - "messages": [{"role": "user", "content": "hello there, this is a test message"}], + "messages": [ + {"role": "user", "content": "hello there, this is a test message"} + ], "max_tokens": 60, } @@ -1811,8 +1935,12 @@ async def test_proxy_rejection_refunds_itpm_otpm_only_reservation_with_no_combin value="proj-io-only:bedrock_mantle/claude-opus", rate_limit_type="tokens", ) - assert int(await cache.async_get_cache(key=itpm_counter_key, local_only=True) or 0) > 0 - assert int(await cache.async_get_cache(key=otpm_counter_key, local_only=True) or 0) > 0 + assert ( + int(await cache.async_get_cache(key=itpm_counter_key, local_only=True) or 0) > 0 + ) + assert ( + int(await cache.async_get_cache(key=otpm_counter_key, local_only=True) or 0) > 0 + ) await handler.async_post_call_failure_hook( request_data=data, @@ -1820,10 +1948,18 @@ async def test_proxy_rejection_refunds_itpm_otpm_only_reservation_with_no_combin user_api_key_dict=user_api_key_dict, ) - itpm_after = int(await cache.async_get_cache(key=itpm_counter_key, local_only=True) or 0) - otpm_after = int(await cache.async_get_cache(key=otpm_counter_key, local_only=True) or 0) - assert itpm_after == 0, f"ITPM-only reservation leaked on proxy rejection: {itpm_after}" - assert otpm_after == 0, f"OTPM-only reservation leaked on proxy rejection: {otpm_after}" + itpm_after = int( + await cache.async_get_cache(key=itpm_counter_key, local_only=True) or 0 + ) + otpm_after = int( + await cache.async_get_cache(key=otpm_counter_key, local_only=True) or 0 + ) + assert itpm_after == 0, ( + f"ITPM-only reservation leaked on proxy rejection: {itpm_after}" + ) + assert otpm_after == 0, ( + f"OTPM-only reservation leaked on proxy rejection: {otpm_after}" + ) @pytest.mark.asyncio @@ -1856,7 +1992,9 @@ async def test_otpm_rejection_does_not_double_refund_combined_tpm(rate_limiter): data = { "model": "bedrock_mantle/claude-opus", - "messages": [{"role": "user", "content": "hello there, this is a test message"}], + "messages": [ + {"role": "user", "content": "hello there, this is a test message"} + ], "max_tokens": 60, # blows past the 5-token OTPM limit } @@ -1866,7 +2004,7 @@ async def test_otpm_rejection_does_not_double_refund_combined_tpm(rate_limiter): rate_limit_type="tokens", ) - with pytest.raises(Exception, match="Rate limit exceeded for model_per_project_otpm") as exc_info: + with pytest.raises(Exception, match='Rate limit exceeded for model_per_project_otpm') as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=cache, @@ -1875,8 +2013,12 @@ async def test_otpm_rejection_does_not_double_refund_combined_tpm(rate_limiter): ) assert getattr(exc_info.value, "status_code", None) == 429 - tpm_after_pre_call = int(await cache.async_get_cache(key=tpm_counter_key, local_only=True) or 0) - assert tpm_after_pre_call == 0, f"combined TPM reservation not rolled back: {tpm_after_pre_call}" + tpm_after_pre_call = int( + await cache.async_get_cache(key=tpm_counter_key, local_only=True) or 0 + ) + assert tpm_after_pre_call == 0, ( + f"combined TPM reservation not rolled back: {tpm_after_pre_call}" + ) # In the real request lifecycle, async_post_call_failure_hook fires next # for a pre-call rejection. It must not refund the same reservation again. @@ -1886,7 +2028,9 @@ async def test_otpm_rejection_does_not_double_refund_combined_tpm(rate_limiter): user_api_key_dict=user_api_key_dict, ) - tpm_after_failure_hook = int(await cache.async_get_cache(key=tpm_counter_key, local_only=True) or 0) + tpm_after_failure_hook = int( + await cache.async_get_cache(key=tpm_counter_key, local_only=True) or 0 + ) assert tpm_after_failure_hook == 0, ( f"combined TPM counter went negative from a double refund: {tpm_after_failure_hook}" ) @@ -1917,7 +2061,7 @@ async def test_project_itpm_rejects_pretokenized_embedding_input( "input": embedding_input, } - with pytest.raises(Exception, match="Rate limit exceeded for model_per_project_itpm") as exc_info: + with pytest.raises(Exception, match='Rate limit exceeded for model_per_project_itpm') as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=cache, @@ -1944,10 +2088,14 @@ async def test_responses_api_not_misclassified_as_embedding_for_output_estimate( data = {"input": "describe this image in detail"} - _, embedding_output_estimate = handler._estimate_input_and_output_tokens(data=data, call_type="aembedding") + _, embedding_output_estimate = handler._estimate_input_and_output_tokens( + data=data, call_type="aembedding" + ) assert embedding_output_estimate == 0 - _, responses_output_estimate = handler._estimate_input_and_output_tokens(data=data, call_type="aresponses") + _, responses_output_estimate = handler._estimate_input_and_output_tokens( + data=data, call_type="aresponses" + ) assert responses_output_estimate > 0, ( "Responses API call was misclassified as an embedding and reserved zero output tokens" ) @@ -2039,7 +2187,9 @@ async def test_responses_api_usage_reconciles_using_input_output_tokens_fields( for op in increment_list: increments.append({"key": op["key"], "increment": op["increment_value"]}) - handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_increment + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( + mock_increment + ) await handler.async_log_success_event( kwargs=mock_kwargs, @@ -2101,7 +2251,7 @@ async def test_itpm_reservation_accounts_for_audio_content_not_just_text(rate_li ], } - with pytest.raises(Exception, match="Rate limit exceeded for model_per_project_itpm") as exc_info: + with pytest.raises(Exception, match='Rate limit exceeded for model_per_project_itpm') as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=cache, @@ -2151,7 +2301,9 @@ def test_audio_token_estimate_scales_with_payload_size(): no_data_block = {"type": "input_audio", "input_audio": {"format": "wav"}} large_estimate = RateLimitHandler._estimate_audio_block_tokens(large_block) - very_large_estimate = RateLimitHandler._estimate_audio_block_tokens(very_large_block) + very_large_estimate = RateLimitHandler._estimate_audio_block_tokens( + very_large_block + ) small_estimate = RateLimitHandler._estimate_audio_block_tokens(small_block) no_data_estimate = RateLimitHandler._estimate_audio_block_tokens(no_data_block) @@ -2210,7 +2362,7 @@ async def test_itpm_rejects_large_audio_payload_that_would_pass_flat_estimate( ], } - with pytest.raises(Exception, match="Rate limit exceeded for model_per_project_itpm") as exc_info: + with pytest.raises(Exception, match='Rate limit exceeded for model_per_project_itpm') as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=cache, @@ -2368,21 +2520,39 @@ async def test_itpm_otpm_reservation_is_kept_on_stream_disconnect(rate_limiter): stash = get_request_stash() assert stash is not None - assert stash.itpm_reserved_tokens > 0, "pre-call hook must stash an ITPM reservation" - assert stash.otpm_reserved_tokens > 0, "pre-call hook must stash an OTPM reservation" + assert stash.itpm_reserved_tokens > 0, ( + "pre-call hook must stash an ITPM reservation" + ) + assert stash.otpm_reserved_tokens > 0, ( + "pre-call hook must stash an OTPM reservation" + ) increment_calls: list[dict] = [] async def mock_increment(increment_list, litellm_parent_otel_span=None): for op in increment_list: - increment_calls.append({"key": op["key"], "increment": op["increment_value"]}) + increment_calls.append( + {"key": op["key"], "increment": op["increment_value"]} + ) - handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_increment + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( + mock_increment + ) - await handler.async_release_max_parallel_requests_on_disconnect(user_api_key_dict=user_api_key_dict) + await handler.async_release_max_parallel_requests_on_disconnect( + user_api_key_dict=user_api_key_dict + ) - itpm_refunds = [c for c in increment_calls if "model_per_project_itpm" in c["key"] and c["increment"] < 0] - otpm_refunds = [c for c in increment_calls if "model_per_project_otpm" in c["key"] and c["increment"] < 0] + itpm_refunds = [ + c + for c in increment_calls + if "model_per_project_itpm" in c["key"] and c["increment"] < 0 + ] + otpm_refunds = [ + c + for c in increment_calls + if "model_per_project_otpm" in c["key"] and c["increment"] < 0 + ] assert not itpm_refunds assert not otpm_refunds @@ -2453,7 +2623,7 @@ async def test_explicit_zero_output_responses_call_reserves_effective_provider_m }, ) - with pytest.raises(Exception, match="Rate limit exceeded for model_per_project_otpm") as exc_info: + with pytest.raises(Exception, match='Rate limit exceeded for model_per_project_otpm') as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=cache, @@ -2681,7 +2851,7 @@ async def test_otpm_rejection_releases_stashed_parallel_slot(rate_limiter): "rate_limit": {"tokens_per_unit": 5, "window_size": 60}, } - with pytest.raises(Exception, match="Rate limit exceeded for model_per_project_otpm") as exc_info: + with pytest.raises(Exception, match='Rate limit exceeded for model_per_project_otpm') as exc_info: await handler._reserve_project_io_tokens_or_raise( descriptors=[otpm_descriptor], data=data, @@ -2729,7 +2899,9 @@ async def test_itpm_only_status_stored_when_no_prior_rate_limit_response(rate_li stash = get_request_stash() assert stash is not None stored = stash.rate_limit_response - assert stored is not None, "ITPM status must be stored in litellm_proxy_rate_limit_response" + assert stored is not None, ( + "ITPM status must be stored in litellm_proxy_rate_limit_response" + ) assert stored.get("statuses"), "Stored response must contain statuses" @@ -2751,7 +2923,9 @@ def test_resolve_io_token_usage_responses_api_with_cached_tokens(rate_limiter): input_tokens_details=InputTokensDetails(cached_tokens=25), ), ) - billable_input, completion_tokens, resolved = handler._resolve_io_token_reconcile_usage(response_obj) + billable_input, completion_tokens, resolved = ( + handler._resolve_io_token_reconcile_usage(response_obj) + ) assert resolved is True assert billable_input == 75, f"Expected 100 - 25 cached = 75, got {billable_input}" @@ -2772,7 +2946,9 @@ def test_resolve_io_token_usage_dict_format(rate_limiter): "prompt_tokens_details": {"cached_tokens": 20}, } ) - billable_input, completion_tokens, resolved = handler._resolve_io_token_reconcile_usage(response_obj) + billable_input, completion_tokens, resolved = ( + handler._resolve_io_token_reconcile_usage(response_obj) + ) assert resolved is True assert billable_input == 60, f"Expected 80 - 20 cached = 60, got {billable_input}" @@ -2788,7 +2964,9 @@ def test_resolve_io_token_usage_unknown_type_returns_unresolved(rate_limiter): handler, _cache = rate_limiter response_obj = ModelResponse.model_construct(usage=42) - billable_input, completion_tokens, resolved = handler._resolve_io_token_reconcile_usage(response_obj) + billable_input, completion_tokens, resolved = ( + handler._resolve_io_token_reconcile_usage(response_obj) + ) assert resolved is False assert billable_input == 0 @@ -2819,7 +2997,9 @@ def test_zero_usage_keeps_reservations_unless_measured_fallback_exists( stash.otpm_reserved_tokens = 60 stash.otpm_reserved_scopes = frozenset({otpm_scope}) kwargs = {} if combined_usage is None else {"combined_usage_object": combined_usage} - response_obj = ModelResponse(usage=Usage(prompt_tokens=0, completion_tokens=0, total_tokens=0)) + response_obj = ModelResponse( + usage=Usage(prompt_tokens=0, completion_tokens=0, total_tokens=0) + ) operations = handler._build_io_token_reservation_ops(kwargs, response_obj) @@ -2916,7 +3096,9 @@ async def test_post_call_failure_skips_rpm_only_descriptor_in_tpm_refund(rate_li call_type="", ) - rpm_tokens_key = handler.create_rate_limit_keys(key="api_key", value=api_key, rate_limit_type="tokens") + rpm_tokens_key = handler.create_rate_limit_keys( + key="api_key", value=api_key, rate_limit_type="tokens" + ) await handler.async_post_call_failure_hook( request_data=data, @@ -2924,7 +3106,9 @@ async def test_post_call_failure_skips_rpm_only_descriptor_in_tpm_refund(rate_li user_api_key_dict=user_api_key_dict, ) - api_key_tokens_after = int(await cache.async_get_cache(key=rpm_tokens_key, local_only=True) or 0) + api_key_tokens_after = int( + await cache.async_get_cache(key=rpm_tokens_key, local_only=True) or 0 + ) assert api_key_tokens_after >= 0, ( f"RPM-only api_key scope must not receive a negative TPM refund; got {api_key_tokens_after}" ) @@ -3113,7 +3297,7 @@ async def test_rerank_query_and_documents_enforce_project_itpm( project_metadata={"model_itpm_limit": {"rerank-model": 100}}, ) - with pytest.raises(Exception, match="Rate limit exceeded for model_per_project_itpm") as exc_info: + with pytest.raises(Exception, match='Rate limit exceeded for model_per_project_itpm') as exc_info: await handler.async_pre_call_hook( user_api_key_dict=user_api_key_dict, cache=cache, @@ -3126,7 +3310,11 @@ async def test_rerank_query_and_documents_enforce_project_itpm( ) assert getattr(exc_info.value, "status_code", None) == 429 - assert captured["text"] == ("Which document is most relevant?\nfirst document\n{'text': 'second document'}") + assert captured["text"] == ( + "Which document is most relevant?\n" + "first document\n" + "{'text': 'second document'}" + ) def test_rerank_input_estimate_falls_back_to_character_count( @@ -3145,21 +3333,20 @@ def test_rerank_input_estimate_falls_back_to_character_count( monkeypatch.setattr("litellm.token_counter", token_counter) rerank_text = handler._rerank_input_to_text(data) - assert ( - handler._estimate_precise_input_tokens( - data, - model="custom-rerank-model", - call_type="rerank", - ) - == len(rerank_text) // 4 - ) + assert handler._estimate_precise_input_tokens( + data, + model="custom-rerank-model", + call_type="rerank", + ) == len(rerank_text) // 4 @pytest.mark.parametrize( ("response_obj", "expected"), [ ( - RerankResponse(meta={"tokens": {"input_tokens": 42, "output_tokens": 3}}), + RerankResponse( + meta={"tokens": {"input_tokens": 42, "output_tokens": 3}} + ), (42, 3, True), ), ( @@ -3200,7 +3387,9 @@ def test_split_quota_helpers_handle_non_mapping_inputs(rate_limiter): assert handler._get_explicit_output_cap(object(), None) is None assert handler.get_output_candidate_count(object()) == 1 assert handler.get_output_candidate_count({"n": 1e309}) == 1 - assert handler._get_explicit_output_cap({"max_output_tokens": []}, "responses") is None + assert ( + handler._get_explicit_output_cap({"max_output_tokens": []}, "responses") is None + ) assert handler._apply_implicit_output_cap(object(), 100, "responses") is None assert handler._estimate_input_and_output_tokens(object()) == (0, 0) assert handler._build_io_token_reservation_ops(object(), object()) == () @@ -3218,7 +3407,9 @@ def test_split_quota_helpers_handle_non_mapping_inputs(rate_limiter): ({"generationConfig": {"maxOutputTokens": "oops"}}, "agenerate_content", None), ], ) -def test_get_explicit_output_cap_tolerates_unparseable_values(rate_limiter, data, call_type, expected): +def test_get_explicit_output_cap_tolerates_unparseable_values( + rate_limiter, data, call_type, expected +): """A client-supplied cap the proxy cannot parse must fall back to the no-cap output estimate instead of raising ValueError and 500ing the request before it ever reaches the provider.""" @@ -3313,7 +3504,9 @@ def test_split_token_estimate_selects_endpoint_input(rate_limiter, call_type, da def test_split_quota_multimodal_guards_handle_non_mapping_inputs(rate_limiter): handler, _cache = rate_limiter - assert handler._estimate_audio_block_tokens(object()) == handler._estimate_audio_block_tokens({}) + assert handler._estimate_audio_block_tokens( + object() + ) == handler._estimate_audio_block_tokens({}) assert handler._responses_input_to_chat_messages(object()) == () assert handler._estimate_precise_input_tokens(object(), model=None) == 0 @@ -3346,7 +3539,10 @@ def test_precise_input_estimate_selects_endpoint_text( monkeypatch.setattr("litellm.token_counter", token_counter) - assert handler._estimate_precise_input_tokens(data, model="test", call_type=call_type) == 7 + assert ( + handler._estimate_precise_input_tokens(data, model="test", call_type=call_type) + == 7 + ) assert captured["messages"] is None assert captured["text"] == expected_text @@ -3389,7 +3585,9 @@ async def test_streaming_combined_usage_reconciles_project_io_reservations( async def capture_increments(increment_list, **_kwargs): increments.extend(increment_list) - handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = capture_increments + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( + capture_increments + ) await handler.async_log_success_event( kwargs=kwargs, @@ -3398,8 +3596,16 @@ async def test_streaming_combined_usage_reconciles_project_io_reservations( end_time=datetime.now(), ) - itpm_adjustments = [operation for operation in increments if PROJECT_ITPM_DESCRIPTOR_KEY in operation["key"]] - otpm_adjustments = [operation for operation in increments if PROJECT_OTPM_DESCRIPTOR_KEY in operation["key"]] + itpm_adjustments = [ + operation + for operation in increments + if PROJECT_ITPM_DESCRIPTOR_KEY in operation["key"] + ] + otpm_adjustments = [ + operation + for operation in increments + if PROJECT_OTPM_DESCRIPTOR_KEY in operation["key"] + ] assert [operation["increment_value"] for operation in itpm_adjustments] == [-60] assert [operation["increment_value"] for operation in otpm_adjustments] == [-45] @@ -3408,9 +3614,13 @@ def test_aggregate_only_combined_usage_reconciles_project_io_reservations(rate_l handler, _cache = rate_limiter stash = get_or_create_request_stash() stash.itpm_reserved_tokens = 100 - stash.itpm_reserved_scopes = frozenset({(PROJECT_ITPM_DESCRIPTOR_KEY, "project:model")}) + stash.itpm_reserved_scopes = frozenset( + {(PROJECT_ITPM_DESCRIPTOR_KEY, "project:model")} + ) stash.otpm_reserved_tokens = 80 - stash.otpm_reserved_scopes = frozenset({(PROJECT_OTPM_DESCRIPTOR_KEY, "project:model")}) + stash.otpm_reserved_scopes = frozenset( + {(PROJECT_OTPM_DESCRIPTOR_KEY, "project:model")} + ) kwargs = { "combined_usage_object": Usage(total_tokens=55), } @@ -3433,7 +3643,9 @@ def test_raw_split_usage_dict_reconciles_project_io_tokens(rate_limiter): @pytest.mark.asyncio -async def test_post_call_success_hook_contains_header_merge_failures(rate_limiter, monkeypatch): +async def test_post_call_success_hook_contains_header_merge_failures( + rate_limiter, monkeypatch +): handler, _cache = rate_limiter response = ModelResponse() response._hidden_params = {} From c8c9ef2690e57c5747b8578afcb7b242b3d7dcd3 Mon Sep 17 00:00:00 2001 From: Mihidum Hettiyahandi <55163074+mihidumh@users.noreply.github.com> Date: Thu, 17 Sep 2026 09:21:32 +1000 Subject: [PATCH 6/6] refactor(batch): read the row body without a mutable fallback literal `(entry.get("body") or {})` counts as mutable-collection construction (LIT002), and main has since lowered that ceiling. Check the body is a dict instead; behaviour is unchanged. --- litellm/proxy/hooks/batch_rate_limiter.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index 140783a4a4d..12585916e16 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -981,7 +981,10 @@ class _PROXY_BatchRateLimiter(CustomLogger): # estimate; floor it there again so a row carrying a # large base64 audio payload cannot slide the batch # under the TPM limit. - if messages_contain_input_audio_content_blocks((entry.get("body") or {}).get("messages")): + entry_body = entry.get("body") if isinstance(entry, dict) else None + if isinstance(entry_body, dict) and messages_contain_input_audio_content_blocks( + entry_body.get("messages") + ): entry_total_tokens = max(entry_total_tokens, _estimate_batch_entry_tokens(raw_line)) except Exception: entry_total_tokens = _estimate_batch_entry_tokens(raw_line)