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] 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 = {}