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
This commit is contained in:
mateo-berri 2026-08-18 14:45:05 -07:00
parent c1c23bf39a
commit 72960d10e9
4 changed files with 144 additions and 18 deletions

View file

@ -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(

View file

@ -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,

View file

@ -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

View file

@ -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"),
[