From 72960d10e93d1cb1bb6ef7b2fd3be82cb9aa7368 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 18 Aug 2026 14:45:05 -0700 Subject: [PATCH] fix(proxy): address review findings on project ITPM/OTPM quotas - scale batch output-token reservations by the row's n / best_of candidate count - parse client-supplied output caps defensively instead of 500ing on unparseable values - exclude project IO descriptors from the first should_rate_limit pass when TPM reservation is disabled so their buckets are not double-charged --- litellm/proxy/hooks/batch_rate_limiter.py | 12 ++- .../hooks/parallel_request_limiter_v3.py | 47 +++++++---- .../proxy/hooks/test_batch_file_validation.py | 25 ++++++ .../proxy/hooks/test_tpm_concurrent.py | 78 +++++++++++++++++++ 4 files changed, 144 insertions(+), 18 deletions(-) diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index 4052a9b04fa..aa3ac6b820b 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -429,12 +429,20 @@ class _PROXY_BatchRateLimiter(CustomLogger): ), None, ) + candidate_count: Final = max( + ( + v + for v in (body.get("n"), body.get("best_of")) + if isinstance(v, int) and not isinstance(v, bool) and v > 1 + ), + default=1, + ) if explicit_cap is not None: try: - return max(0, int(explicit_cap)) + return max(0, int(explicit_cap)) * candidate_count except (TypeError, ValueError): pass - return self.parallel_request_limiter.no_max_tokens_output_floor(min_configured_otpm_limit) + return self.parallel_request_limiter.no_max_tokens_output_floor(min_configured_otpm_limit) * candidate_count @staticmethod def _has_applicable_batch_rate_limits( diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 751b301a93d..7cd56e66000 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -559,6 +559,15 @@ def _call_id_from_callback_kwargs(kwargs: object) -> str | None: return call_id if isinstance(call_id, str) else None +def _parse_output_cap_value(raw_value: object) -> int | None: + if isinstance(raw_value, bool) or not isinstance(raw_value, (int, float, str)): + return None + try: + return int(float(raw_value)) + except (ValueError, OverflowError): + return None + + class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): def __init__( self, @@ -707,20 +716,18 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): if call_type in GOOGLE_GENAI_NATIVE_CALL_TYPES: config: Final = data.get("config") if "config" in data else data.get("generationConfig") google_cap_values: Final = tuple( - int(raw_value) + parsed for field in ("maxOutputTokens", "max_output_tokens") if isinstance(config, dict) - for raw_value in (config.get(field),) - if isinstance(raw_value, (int, float, str)) + for parsed in (_parse_output_cap_value(config.get(field)),) + if parsed is not None ) return max(google_cap_values, default=None) if call_type in RESPONSES_API_CALL_TYPES: - value: Final = data.get("max_output_tokens") - if value is None: + responses_cap: Final = _parse_output_cap_value(data.get("max_output_tokens")) + if responses_cap is None: return None - if not isinstance(value, (int, float, str)): - return None - return max(RESPONSES_API_MIN_OUTPUT_TOKENS, int(value)) + return max(RESPONSES_API_MIN_OUTPUT_TOKENS, responses_cap) if call_type in EMBEDDING_API_CALL_TYPES: return None fields: Final = ( @@ -729,10 +736,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): else ("max_tokens", "max_completion_tokens", "max_output_tokens") ) output_cap_values: Final = tuple( - int(raw_value) - for field in fields - for raw_value in (data.get(field),) - if isinstance(raw_value, (int, float, str)) + parsed for field in fields for parsed in (_parse_output_cap_value(data.get(field)),) if parsed is not None ) return max(output_cap_values, default=None) @@ -1209,7 +1213,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): async def should_rate_limit( self, - descriptors: list[RateLimitDescriptor], + descriptors: Sequence[RateLimitDescriptor], parent_otel_span: Span | None = None, read_only: bool = False, skip_tpm_check: bool = False, @@ -1335,7 +1339,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): def _collect_windowed_keys_and_gauges( self, - descriptors: list[RateLimitDescriptor], + descriptors: Sequence[RateLimitDescriptor], skip_tpm_check: bool, ) -> tuple[list[str], dict[str, WindowKeyMetadata], list[ParallelRequestGauge]]: """ @@ -3456,7 +3460,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): # in-flight request would pre-inflate the :tokens counter by 1, # shrinking the effective TPM budget by N and causing # false-positive 429s under bursts. When reservation is disabled, - # this pass enforces TPM directly from the post-call counters. + # this pass enforces TPM directly from the post-call counters -- + # except for project ITPM/OTPM descriptors, which are excluded + # then because _reserve_project_io_tokens_or_raise below charges + # them unconditionally and counting them here too would + # double-charge every request. parallel_counter_keys: Final = [ self.create_rate_limit_keys(d["key"], d["value"], "max_parallel_requests") for d in descriptors @@ -3464,8 +3472,15 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ] parallel_slot_id: Final = uuid.uuid4().hex if parallel_counter_keys else None + first_pass_descriptors: Final = ( + descriptors + if self.tpm_reservation_enabled + else tuple( + d for d in descriptors if d["key"] not in (PROJECT_ITPM_DESCRIPTOR_KEY, PROJECT_OTPM_DESCRIPTOR_KEY) + ) + ) response: Final = await self.should_rate_limit( - descriptors=descriptors, + descriptors=first_pass_descriptors, parent_otel_span=user_api_key_dict.parent_otel_span, skip_tpm_check=self.tpm_reservation_enabled, parallel_slot_id=parallel_slot_id, 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 e5211973ec3..3a9d3ff72aa 100644 --- a/tests/test_litellm/proxy/hooks/test_batch_file_validation.py +++ b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py @@ -2060,3 +2060,28 @@ def test_estimate_entry_output_tokens_prefers_max_tokens_over_max_output_tokens( } assert rate_limiter._estimate_entry_output_tokens(entry, None) == 50 + + +@pytest.mark.parametrize( + ("body_extra", "expected"), + [ + ({"max_tokens": 40, "n": 10}, 400), + ({"max_tokens": 40, "best_of": 5}, 200), + ({"max_tokens": 40, "n": 3, "best_of": 5}, 200), + ({"max_tokens": 40, "n": 0}, 40), + ({"max_tokens": 40, "n": -2}, 40), + ({"n": 3}, 2997), + ], +) +def test_estimate_entry_output_tokens_multiplies_candidate_count(body_extra, expected): + """A row generating n / best_of candidates consumes that many completions' + worth of output tokens, so the OTPM reservation must scale with the + effective candidate count. Pre-fix a `max_tokens: 40, n: 10` row consumed + up to 400 output tokens while reserving only 40.""" + rate_limiter = _output_estimator() + entry = { + "url": "/v1/chat/completions", + "body": {"model": "gpt-4o", "messages": [], **body_extra}, + } + + assert rate_limiter._estimate_entry_output_tokens(entry, None) == expected diff --git a/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py b/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py index fd7e0626af9..e12610b769c 100644 --- a/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py +++ b/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py @@ -3393,6 +3393,84 @@ def test_split_quota_helpers_handle_non_mapping_inputs(rate_limiter): assert handler._build_io_token_reservation_ops(object(), object()) == () +@pytest.mark.parametrize( + ("data", "call_type", "expected"), + [ + ({"max_tokens": "30.0"}, "", 30), + ({"max_tokens": "not-a-number"}, "", None), + ({"max_tokens": True}, "", None), + ({"max_output_tokens": "30.0"}, "responses", 30), + ({"max_output_tokens": "nan"}, "responses", None), + ({"generationConfig": {"maxOutputTokens": "12.5"}}, "agenerate_content", 12), + ({"generationConfig": {"maxOutputTokens": "oops"}}, "agenerate_content", None), + ], +) +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.""" + handler, _cache = rate_limiter + + assert handler._get_explicit_output_cap(data, call_type) == expected + + +@pytest.mark.asyncio +async def test_project_io_counters_not_double_charged_when_reservation_disabled( + monkeypatch, +): + """With LITELLM_TPM_TOKEN_RESERVATION_ENABLED=false the first + should_rate_limit pass used to +1 every ITPM/OTPM counter on top of the + full reservation _reserve_project_io_tokens_or_raise always makes, + permanently inflating each bucket by one token per request.""" + monkeypatch.setenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", "false") + cache = DualCache() + handler = RateLimitHandler(internal_usage_cache=InternalUsageCache(cache)) + assert handler.tpm_reservation_enabled is False + + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-io-no-reservation"), + project_id="proj-io-no-reservation", + project_metadata={ + "model_itpm_limit": {"bedrock_mantle/claude-opus": 1000000}, + "model_otpm_limit": {"bedrock_mantle/claude-opus": 1000000}, + }, + ) + data = { + "model": "bedrock_mantle/claude-opus", + "messages": [{"role": "user", "content": "hello"}], + "max_tokens": 50, + } + + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data=data, + call_type="", + ) + + stash = get_request_stash() + assert stash is not None + assert stash.itpm_reserved_tokens > 0 + assert stash.otpm_reserved_tokens > 0 + + for descriptor_key, reserved in ( + ("model_per_project_itpm", stash.itpm_reserved_tokens), + ("model_per_project_otpm", stash.otpm_reserved_tokens), + ): + counter_key = handler.create_rate_limit_keys( + key=descriptor_key, + value="proj-io-no-reservation:bedrock_mantle/claude-opus", + rate_limit_type="tokens", + ) + cached = await cache.async_get_cache(key=counter_key, local_only=True) + assert int(cached or 0) == reserved, ( + f"{descriptor_key} counter {cached} != reserved {reserved}: " + "first-pass should_rate_limit double-charged the bucket" + ) + + @pytest.mark.parametrize( ("call_type", "data"), [