mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
c1c23bf39a
commit
72960d10e9
4 changed files with 144 additions and 18 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
[
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue