From 9dbe61aa6d8915f40a030ee11f522345afb9cc6a Mon Sep 17 00:00:00 2001 From: Shivi Jain Date: Fri, 31 Jul 2026 16:49:07 +0530 Subject: [PATCH] feat(proxy): add project-level ITPM and OTPM quotas Add model_itpm_limit and model_otpm_limit to project create and update requests, storing both quota maps in project metadata without a database migration Reserve input and output tokens independently before provider dispatch, expose separate project rate-limit headers, and reconcile counters across successful calls, failures, retries, fallbacks, streaming, caching, and cancellation Harden token estimation for pre-tokenized embeddings, multimodal inputs, Responses API requests, native Gemini requests, multiple candidates, and conflicting output-cap aliases Reject negative output caps, preserve conservative reservations when usage is missing or zero, bind reconciliation and refunds to the reservation window, prevent double refunds or negative counters, update generated API types, and add regression coverage --- litellm/proxy/_types.py | 8 +- litellm/proxy/auth/auth_utils.py | 2 +- .../proxy/hooks/dynamic_rate_limiter_v3.py | 4 +- .../hooks/parallel_request_limiter_v3.py | 1561 ++++++++++- .../hooks/test_parallel_request_limiter_v3.py | 132 +- .../proxy/hooks/test_tpm_concurrent.py | 2444 ++++++++++++++++- tests/test_litellm/proxy/test_proxy_types.py | 18 + ui/litellm-dashboard/src/lib/http/schema.d.ts | 16 + 8 files changed, 4095 insertions(+), 90 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 7d6829aca70..b4aab95e2d4 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1,7 +1,7 @@ import enum import json import os -from collections.abc import Callable +from collections.abc import Callable, Mapping from datetime import datetime from typing import TYPE_CHECKING, Any, Final, Literal, Union @@ -2903,6 +2903,8 @@ class NewProjectRequest(LiteLLM_BudgetTable): models: list[str] = [] model_rpm_limit: dict | None = None model_tpm_limit: dict | None = None + model_itpm_limit: Mapping[str, int] | None = None + model_otpm_limit: Mapping[str, int] | None = None blocked: bool = False object_permission: LiteLLM_ObjectPermissionBase | None = None @@ -2935,6 +2937,8 @@ class UpdateProjectRequest(LiteLLM_BudgetTable): models: list[str] | None = None model_rpm_limit: dict | None = None model_tpm_limit: dict | None = None + model_itpm_limit: Mapping[str, int] | None = None + model_otpm_limit: Mapping[str, int] | None = None blocked: bool | None = None budget_id: str | None = None object_permission: LiteLLM_ObjectPermissionBase | None = None @@ -4072,6 +4076,8 @@ class PassThroughEndpointLoggingTypedDict(TypedDict): LiteLLM_ManagementEndpoint_MetadataFields: Final = [ "model_rpm_limit", "model_tpm_limit", + "model_itpm_limit", + "model_otpm_limit", "mcp_rpm_limit", "tag_rpm_limit", "rpm_limit_type", diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 87db8ed3b5e..21c41e08ecc 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -994,7 +994,7 @@ def get_key_model_tpm_limit( def get_model_rate_limit_from_metadata( user_api_key_dict: UserAPIKeyAuth, metadata_accessor_key: Literal["team_metadata", "organization_metadata", "project_metadata"], - rate_limit_key: Literal["model_rpm_limit", "model_tpm_limit"], + rate_limit_key: Literal["model_rpm_limit", "model_tpm_limit", "model_itpm_limit", "model_otpm_limit"], ) -> dict[str, int] | None: if getattr(user_api_key_dict, metadata_accessor_key): return getattr(user_api_key_dict, metadata_accessor_key).get(rate_limit_key) diff --git a/litellm/proxy/hooks/dynamic_rate_limiter_v3.py b/litellm/proxy/hooks/dynamic_rate_limiter_v3.py index 4492f42782c..de8834449de 100644 --- a/litellm/proxy/hooks/dynamic_rate_limiter_v3.py +++ b/litellm/proxy/hooks/dynamic_rate_limiter_v3.py @@ -454,7 +454,9 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): parent_otel_span=user_api_key_dict.parent_otel_span, ) - verbose_proxy_logger.debug("Atomic check+increment response: %s", json.dumps(atomic_response, indent=2)) + verbose_proxy_logger.debug( + "Atomic check+increment response: %s", json.dumps(atomic_response, indent=2, default=list) + ) if atomic_response["overall_code"] == "OVER_LIMIT": resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(model) diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 2725da1ee12..e6cb54393ab 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -8,11 +8,22 @@ import asyncio import binascii import os import uuid -from collections.abc import Callable +from collections.abc import Callable, Mapping, Sequence, Set from contextvars import ContextVar from dataclasses import dataclass, field from datetime import datetime -from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict, Union, cast +from typing import ( + TYPE_CHECKING, + Any, + Final, + Literal, + Protocol, + TypedDict, + Union, + cast, +) + +from typing_extensions import NotRequired from litellm import DualCache from litellm._logging import verbose_proxy_logger @@ -34,11 +45,12 @@ from litellm.proxy.common_utils.proxy_rate_limit_error import ( ) from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit from litellm.types.caching import RedisPipelineIncrementOperation -from litellm.types.llms.openai import BaseLiteLLMOpenAIResponseObject +from litellm.types.llms.openai import BaseLiteLLMOpenAIResponseObject, ResponseAPIUsage from litellm.types.utils import ( CallTypes, EmbeddingResponse, ModelResponse, + RerankResponse, TextCompletionResponse, Usage, ) @@ -109,7 +121,8 @@ CHECK_AND_INCREMENT_BY_N_SCRIPT: Final = """ -- ARGV[(i-1)*4 + 3] = ttl_seconds (counter TTL when window resets) -- ARGV[(i-1)*4 + 4] = window_size_seconds (sliding-window length) -- --- Return on success: { 0, new_counter_1, new_counter_2, ... } +-- Return on success: +-- { 0, new_counter_1, window_start_1, new_counter_2, window_start_2, ... } -- Return on over-limit: { 1, descriptor_index, current_counter, limit } local time_reply = redis.call('TIME') local now = tonumber(time_reply[1]) @@ -146,7 +159,7 @@ for i = 1, descriptor_count do return { 1, i, current_counter, limit } end - descriptor_state[i] = { window_expired, current_counter } + descriptor_state[i] = { window_expired, current_counter, window_start } end -- Pass 2: all checks passed. Apply increments. @@ -160,8 +173,10 @@ for i = 1, descriptor_count do local window_size = tonumber(ARGV[arg_base + 3]) local window_expired = descriptor_state[i][1] + local active_window_start if window_expired then + active_window_start = now redis.call('SET', window_key, tostring(now)) redis.call('SET', counter_key, increment) redis.call('EXPIRE', window_key, window_size) @@ -170,6 +185,7 @@ for i = 1, descriptor_count do end table.insert(results, increment) else + active_window_start = tonumber(descriptor_state[i][3]) local new_counter = redis.call('INCRBY', counter_key, increment) local current_ttl = redis.call('TTL', counter_key) if current_ttl == -1 and ttl > 0 then @@ -177,11 +193,39 @@ for i = 1, descriptor_count do end table.insert(results, new_counter) end + table.insert(results, active_window_start) end return results """ +WINDOW_GUARDED_TOKEN_INCREMENT_SCRIPT: Final = """ +local results = {} +for i = 1, #KEYS, 2 do + local window_key = KEYS[i] + local counter_key = KEYS[i + 1] + local arg_base = ((i - 1) / 2) * 3 + 1 + local expected_window_start = ARGV[arg_base] + local increment = tonumber(ARGV[arg_base + 1]) + local ttl = tonumber(ARGV[arg_base + 2]) + local active_window_start = redis.call('GET', window_key) + + if active_window_start and active_window_start == expected_window_start then + local new_counter = redis.call('INCRBY', counter_key, increment) + local current_ttl = redis.call('TTL', counter_key) + if current_ttl == -1 and ttl > 0 then + redis.call('EXPIRE', counter_key, ttl) + end + table.insert(results, 1) + table.insert(results, new_counter) + else + table.insert(results, 0) + table.insert(results, tonumber(redis.call('GET', counter_key) or 0)) + end +end +return results +""" + PARALLEL_ACQUIRE_SCRIPT: Final = """ -- Atomic check-and-acquire for the max_parallel_requests concurrency gauge. -- Each gauge key is a sorted set of per-request slot ids scored by acquire @@ -286,6 +330,38 @@ DEFAULT_CHARS_PER_TOKEN: Final = 4 # (baseline floor) and to the smallest configured TPM limit (capped floor for # small per-tenant TPM caps). _TPM_FLOOR_FRACTION: Final = 4 +# Both embeddings and the Responses API put their prompt in data["input"], +# but only embeddings have no output tokens. Every "is this an embedding" +# check on data["input"] must exclude these call types, or a Responses call +# gets misclassified as an embedding and skips output-token reservation/caps. +RESPONSES_API_CALL_TYPES: Final = ("aresponses", "responses") +EMBEDDING_API_CALL_TYPES: Final = ("aembedding", "embedding") +TEXT_COMPLETION_API_CALL_TYPES: Final = ("atext_completion", "text_completion") +RERANK_API_CALL_TYPES: Final = (CallTypes.rerank.value, CallTypes.arerank.value) +GOOGLE_GENAI_NATIVE_CALL_TYPES: Final = ( + CallTypes.generate_content.value, + CallTypes.agenerate_content.value, + CallTypes.generate_content_stream.value, + 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. +PROJECT_ITPM_DESCRIPTOR_KEY: Final = "model_per_project_itpm" +PROJECT_OTPM_DESCRIPTOR_KEY: Final = "model_per_project_otpm" # How long an acquired slot counts toward the in-flight total before it is # considered leaked (worker crashed without any release callback firing) and # pruned. Also the longest request duration the gauge can track: a request @@ -328,6 +404,13 @@ class RateLimitStatus(TypedDict): class RateLimitResponse(TypedDict): overall_code: str statuses: list[RateLimitStatus] + reservation_windows: NotRequired[frozenset[tuple[str, str, Literal["redis", "local"]]]] + + +class ReservationAwareIncrementOperation(RedisPipelineIncrementOperation): + window_key: NotRequired[str] + expected_window_start: NotRequired[str] + reservation_backend: NotRequired[Literal["redis", "local"]] class RateLimitResponseWithDescriptors(TypedDict): @@ -335,6 +418,10 @@ class RateLimitResponseWithDescriptors(TypedDict): response: RateLimitResponse +class _RateLimitDescriptorSink(Protocol): + def append(self, descriptor: RateLimitDescriptor, /) -> None: ... + + @dataclass(slots=True) class RequestRateLimiterStash: """ @@ -364,6 +451,16 @@ class RequestRateLimiterStash: reserved_tokens: int = 0 reserved_model: str | None = None reserved_scopes: frozenset[tuple[str, str]] = field(default_factory=frozenset) + itpm_reserved_tokens: int = 0 + itpm_reserved_scopes: frozenset[tuple[str, str]] = field(default_factory=frozenset) + itpm_reserved_window_identities: frozenset[tuple[str, str, Literal["redis", "local"]]] = field( + default_factory=frozenset + ) + otpm_reserved_tokens: int = 0 + otpm_reserved_scopes: frozenset[tuple[str, str]] = field(default_factory=frozenset) + otpm_reserved_window_identities: frozenset[tuple[str, str, Literal["redis", "local"]]] = field( + default_factory=frozenset + ) reservation_released: bool = False @@ -426,6 +523,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): self.check_and_increment_by_n_script = ( self.internal_usage_cache.dual_cache.redis_cache.async_register_script(CHECK_AND_INCREMENT_BY_N_SCRIPT) ) + self.window_guarded_token_increment_script = ( + self.internal_usage_cache.dual_cache.redis_cache.async_register_script( + WINDOW_GUARDED_TOKEN_INCREMENT_SCRIPT + ) + ) self.parallel_acquire_script = self.internal_usage_cache.dual_cache.redis_cache.async_register_script( PARALLEL_ACQUIRE_SCRIPT ) @@ -439,6 +541,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): self.batch_rate_limiter_script = None self.token_increment_script = None self.check_and_increment_by_n_script = None + self.window_guarded_token_increment_script = None self.parallel_acquire_script = None self.parallel_release_script = None self.parallel_count_script = None @@ -505,18 +608,163 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return baseline return min(baseline, max(1, min_configured_tpm_limit // _TPM_FLOOR_FRACTION)) + @staticmethod + def _is_embedding_request(data: object, call_type: str | None) -> bool: + if call_type in EMBEDDING_API_CALL_TYPES: + return True + if call_type in RESPONSES_API_CALL_TYPES: + return False + if call_type: + return False + if not isinstance(data, dict): + return False + return data.get("input") is not None + + @staticmethod + def _translate_google_genai_native_request( + data: object, + call_type: str | None, + ) -> Mapping[str, object] | None: + contents = data.get("contents") if isinstance(data, dict) else None + if ( + not isinstance(data, dict) + or call_type not in GOOGLE_GENAI_NATIVE_CALL_TYPES + or not isinstance(contents, (dict, list)) + ): + return None + from litellm.google_genai.adapters.transformation import GoogleGenAIAdapter + + config = data.get("config") if "config" in data else data.get("generationConfig") + translated_request = GoogleGenAIAdapter().translate_generate_content_to_completion( + model=data.get("model") if isinstance(data.get("model"), str) else "", + contents=contents, + config=config if isinstance(config, dict) else None, + systemInstruction=data.get("systemInstruction"), + system_instruction=data.get("system_instruction"), + tools=data.get("tools"), + toolConfig=data.get("toolConfig"), + tool_config=data.get("tool_config"), + ) + return translated_request + + @staticmethod + def _get_explicit_output_cap(data: object, call_type: str | None) -> int | None: + if not isinstance(data, dict): + return None + if call_type in GOOGLE_GENAI_NATIVE_CALL_TYPES: + config = data.get("config") if "config" in data else data.get("generationConfig") + values = tuple( + int(config[field]) + for field in ("maxOutputTokens", "max_output_tokens") + if isinstance(config, dict) and isinstance(config.get(field), (int, float, str)) + ) + return max(values, default=None) + if call_type in RESPONSES_API_CALL_TYPES: + value = data.get("max_output_tokens") + if value is None: + return None + if not isinstance(value, (int, float, str)): + return None + return max(RESPONSES_API_MIN_OUTPUT_TOKENS, int(value)) + if call_type in EMBEDDING_API_CALL_TYPES: + return None + fields = ( + ("max_tokens", "max_completion_tokens") + if call_type + else ("max_tokens", "max_completion_tokens", "max_output_tokens") + ) + values = tuple(int(data[field]) for field in fields if isinstance(data.get(field), (int, float, str))) + return max(values, default=None) + + @classmethod + def _has_explicit_output_cap(cls, data: object, call_type: str | None) -> bool: + """Whether the caller explicitly set an output-token cap. + + Checked via ``is not None`` (not truthiness) so an explicit 0 -- + a legitimate zero-output request -- counts as explicit. + """ + return cls._get_explicit_output_cap(data, call_type) is not None + + @staticmethod + def _get_output_candidate_count(data: object, call_type: str | None = None) -> int: + if not isinstance(data, dict): + return 1 + config = ( + (data.get("config") if "config" in data else data.get("generationConfig")) + if call_type in GOOGLE_GENAI_NATIVE_CALL_TYPES + else None + ) + candidate_values = ( + data.get("n"), + data.get("best_of"), + config.get("candidateCount") if isinstance(config, dict) else None, + config.get("candidate_count") if isinstance(config, dict) else None, + ) + candidate_count = 1 + for value in candidate_values: + try: + candidate_count = max(candidate_count, int(value or 1)) + except (TypeError, ValueError): + continue + return candidate_count + + @staticmethod + def _apply_implicit_output_cap( + data: object, + min_configured_limit: int | None, + call_type: str | None, + ) -> None: + """Hard-cap generation length when the request has no explicit cap. + + Guards against an unbounded response overshooting a small TPM/OTPM + budget before post-call reconciliation runs. Skips requests that + already set an explicit cap and embeddings, which have no generation + budget. The Responses API only honors ``max_output_tokens`` (its + underlying chat-completion transformation ignores ``max_tokens``), so + the cap must be written to that field for Responses call types. + """ + if not isinstance(data, dict): + return + capped_floor = _PROXY_MaxParallelRequestsHandler_v3._no_max_tokens_output_floor(min_configured_limit) + if call_type in RESPONSES_API_CALL_TYPES: + capped_floor = max(capped_floor, RESPONSES_API_MIN_OUTPUT_TOKENS) + baseline_floor = DEFAULT_MAX_TOKENS_ESTIMATE // _TPM_FLOOR_FRACTION + is_embedding = _PROXY_MaxParallelRequestsHandler_v3._is_embedding_request(data, call_type) + if ( + capped_floor >= baseline_floor + or _PROXY_MaxParallelRequestsHandler_v3._has_explicit_output_cap(data, call_type) + or is_embedding + ): + return + if call_type in GOOGLE_GENAI_NATIVE_CALL_TYPES: + config_field = "config" if "config" in data or "generationConfig" not in data else "generationConfig" + config = data.get(config_field) + if config is None or isinstance(config, dict): + data[config_field] = { # mutable-ok: downstream native routing requires a mutable request config + **(config or {}), # mutable-ok: downstream native routing requires a mutable request config + "maxOutputTokens": capped_floor, + } + return + cap_field = "max_output_tokens" if call_type in RESPONSES_API_CALL_TYPES else "max_tokens" + existing_cap = data.get(cap_field) + if existing_cap is None or capped_floor < existing_cap: + data[cap_field] = capped_floor + def _estimate_tokens_for_request( self, data: dict, model: str | None = None, min_configured_tpm_limit: int | None = None, + call_type: str | None = None, ) -> int: """ Estimate total tokens this request will consume so we can reserve them upfront (input + output budget): estimated = input_tokens + max_tokens. - Supports chat (messages), completions (prompt), and embeddings (input). + Supports chat (messages), completions (prompt), embeddings (input), + and the Responses API (also `input`, disambiguated from embeddings + via ``call_type``). ``min_configured_tpm_limit`` is the smallest ``tokens_per_unit`` among the TPM-bearing descriptors this request will be charged against. When @@ -524,34 +772,87 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): fraction of that limit so small TPM caps remain usable. Omit to preserve the unconstrained floor. """ - messages = data.get("messages") - prompt: Final = data.get("prompt") - input_text: Final = data.get("input") # embeddings + estimated_input_tokens, max_tokens_estimate = self._estimate_input_and_output_tokens( + data=data, + min_configured_tpm_limit=min_configured_tpm_limit, + call_type=call_type, + ) + total_estimated: Final = estimated_input_tokens + max_tokens_estimate + + verbose_proxy_logger.debug( + "TPM reservation estimate: input=%s, max_tokens=%s, total=%s", + estimated_input_tokens, + max_tokens_estimate, + total_estimated, + ) + + return total_estimated + + def _estimate_input_and_output_tokens( + self, + data: object, + min_configured_tpm_limit: int | None = None, + call_type: str | None = None, + ) -> tuple[int, int]: + """ + Estimate input tokens and output (max_tokens) budget separately, so + callers needing independent ITPM/OTPM reservations (rather than one + combined TPM reservation) can use each half on its own. + + ``min_configured_tpm_limit`` is the smallest ``tokens_per_unit`` among + the TPM-bearing descriptors this request will be charged against. When + provided, the no-``max_tokens`` output-budget floor is capped at a + fraction of that limit so small TPM caps remain usable. Omit to + preserve the unconstrained floor. + + ``call_type`` disambiguates embeddings from the Responses API: both + put their prompt in ``data["input"]``, but only embeddings have no + output tokens. Unset (the default) preserves the historical + "any `input` means zero output" behavior for callers that don't have + a call type to pass. + """ + if not isinstance(data, dict): + return 0, 0 + translated_data: Final = self._translate_google_genai_native_request(data, call_type) + estimable_data: Final = translated_data if translated_data is not None else data + messages = estimable_data.get("messages") + prompt = estimable_data.get("prompt") + input_text = estimable_data.get("input") + + if call_type in RESPONSES_API_CALL_TYPES or call_type in EMBEDDING_API_CALL_TYPES: + messages = None + prompt = None + elif call_type in TEXT_COMPLETION_API_CALL_TYPES: + messages = None + input_text = None + elif call_type: + prompt = None + input_text = None match (messages, prompt, input_text): - case (messages, _, _) if messages: - total_chars = len(get_str_from_messages(messages)) - case (_, str() as p, _): - total_chars = len(p) - case (_, list() as p, _): - total_chars = sum(len(str(item)) for item in p) - case (_, _, str() as t): - total_chars = len(t) - case (_, _, list() as t): - total_chars = sum(len(str(item)) for item in t) + case (selected_messages, _, _) if selected_messages: + total_chars = len(get_str_from_messages(selected_messages)) + case (_, str() as selected_prompt, _): + total_chars = len(selected_prompt) + case (_, list() as selected_prompt, _): + total_chars = sum(len(str(item)) for item in selected_prompt) + case (_, _, str() as selected_input): + total_chars = len(selected_input) + case (_, _, list() as selected_input): + total_chars = sum(len(str(item)) for item in selected_input) case _: total_chars = 0 estimated_input_tokens: Final = max(1, total_chars // DEFAULT_CHARS_PER_TOKEN) if total_chars > 0 else 0 - explicit_max_tokens: Final = data.get("max_tokens") or data.get("max_completion_tokens") + explicit_max_tokens: Final = self._get_explicit_output_cap(data, call_type) + is_embedding: Final = self._is_embedding_request(data, call_type) - match (explicit_max_tokens, input_text): - case (mt, _) if mt is not None: - max_tokens_estimate = int(mt) - case (_, embeddings_input) if embeddings_input: - # Embeddings have no output tokens + match (explicit_max_tokens, is_embedding): + case (_, True): max_tokens_estimate = 0 + case (mt, _) if mt is not None: + max_tokens_estimate = mt case _ if total_chars == 0: # Fully contentless request (no messages, prompt, or input). # Don't apply the conservative output-budget floor here — it @@ -566,20 +867,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): # the smallest TPM limit this request will be charged against, # so a small per-tenant TPM cap can't be tripped by the floor # alone. - output_floor: Final = self._no_max_tokens_output_floor(min_configured_tpm_limit) + output_floor = self._no_max_tokens_output_floor(min_configured_tpm_limit) + if call_type in RESPONSES_API_CALL_TYPES: + output_floor = max(output_floor, RESPONSES_API_MIN_OUTPUT_TOKENS) max_tokens_estimate = max(estimated_input_tokens, output_floor) - total_estimated: Final = estimated_input_tokens + max_tokens_estimate - - verbose_proxy_logger.debug( - "TPM reservation estimate: input=%s, max_tokens=%s (explicit=%s), total=%s", - estimated_input_tokens, - max_tokens_estimate, - explicit_max_tokens is not None, - total_estimated, - ) - - return total_estimated + max_tokens_estimate *= self._get_output_candidate_count(data, call_type) + return estimated_input_tokens, max_tokens_estimate def _is_redis_cluster(self) -> bool: """ @@ -1367,6 +1661,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): descriptor i, refund descriptors 0..i-1's increments. On Lua failure mid-loop, refund applied increments and fall back to in-memory. """ + if not descriptor_groups: + return RateLimitResponse(overall_code="OK", statuses=[]) applied: Final[list[list[dict[str, Any]]]] = [] statuses: Final[list[RateLimitStatus]] = [] @@ -1400,10 +1696,16 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): if response["overall_code"] == "OVER_LIMIT": await self._refund_applied_descriptor_groups(applied) return response + if len(descriptor_groups) == 1: + return response applied.append(meta) statuses.extend(response["statuses"]) - return RateLimitResponse(overall_code="OK", statuses=statuses) + return RateLimitResponse( + overall_code="OK", + statuses=statuses, + reservation_windows=frozenset(), + ) async def _refund_applied_descriptor_groups( self, @@ -1471,7 +1773,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) statuses: Final[list[RateLimitStatus]] = [] - for meta, new_counter in zip(per_counter_meta, raw[1:]): + for index, meta in enumerate(per_counter_meta): + new_counter = raw[1 + index * 2] statuses.append( RateLimitStatus( code="OK", @@ -1481,7 +1784,18 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): descriptor_key=meta["descriptor_key"], ) ) - return RateLimitResponse(overall_code="OK", statuses=statuses) + return RateLimitResponse( + overall_code="OK", + statuses=statuses, + reservation_windows=frozenset( + ( + meta["counter_key"], + str(int(raw[2 + index * 2])), + "redis", + ) + for index, meta in enumerate(per_counter_meta) + ), + ) async def _atomic_check_and_increment_in_memory( self, @@ -1539,7 +1853,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) ], ) - descriptor_state.append({"window_expired": window_expired, "current": current_counter}) + descriptor_state.append( + { + "window_expired": window_expired, + "current": current_counter, + "window_start": str(now_int if window_expired else int(window_start)), + } + ) # Pass 2: apply increments. statuses: Final[list[RateLimitStatus]] = [] @@ -1569,7 +1889,14 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): descriptor_key=meta["descriptor_key"], ) ) - return RateLimitResponse(overall_code="OK", statuses=statuses) + return RateLimitResponse( + overall_code="OK", + statuses=statuses, + reservation_windows=frozenset( + (meta["counter_key"], state["window_start"], "local") + for meta, state in zip(per_counter_meta, descriptor_state) + ), + ) async def reserve_tpm_tokens( self, @@ -1586,6 +1913,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): TPM-only descriptor/increment list and delegates the all-or-nothing atomicity (Lua on Redis, asyncio-locked DualCache otherwise) to the shared primitive. + + Excludes project ITPM/OTPM descriptors -- those are reserved + separately (different estimate per bucket) via ``reserve_io_tokens``. """ tpm_descriptors: Final[list[RateLimitDescriptor]] = [ RateLimitDescriptor( @@ -1597,7 +1927,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ), ) for d in descriptors - if (d.get("rate_limit") or {}).get("tokens_per_unit") is not None + if d["key"] not in (PROJECT_ITPM_DESCRIPTOR_KEY, PROJECT_OTPM_DESCRIPTOR_KEY) + and (d.get("rate_limit") or {}).get("tokens_per_unit") is not None # mutable-ok: optional descriptor ] if not tpm_descriptors: return RateLimitResponse(overall_code="OK", statuses=[]) @@ -1611,6 +1942,142 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): parent_otel_span=parent_otel_span, ) + async def _refund_reserved_tokens( + self, + scopes: Sequence[tuple[str, str]], + amount: int, + reservation_windows: frozenset[tuple[str, str, Literal["redis", "local"]]] = frozenset(), + parent_otel_span: Span | None = None, + ) -> None: + """ + Directly decrement previously-reserved token counters for ``scopes`` + by ``amount``. Used to roll back a reservation that already + succeeded once a *different* bucket in the same request turns out to + be over its limit (e.g. ITPM reserved fine, OTPM then hits its + limit -- the ITPM reservation must not be left inflated). + """ + if amount <= 0 or not scopes: + return + if not reservation_windows: + await self.async_increment_tokens_with_ttl_preservation( + pipeline_operations=self._build_reservation_aware_tpm_ops( + targets=scopes, + reserved_scopes=frozenset(scopes), + actual_tokens=0, + reserved_tokens=amount, + ), + parent_otel_span=parent_otel_span, + ) + return + pipeline_operations: Final = self._build_project_reservation_ops( + targets=scopes, + reserved_scopes=frozenset(scopes), + actual_tokens=0, + reserved_tokens=amount, + reservation_window_identities=reservation_windows, + ) + await self.async_increment_reservation_aware_tokens( + pipeline_operations=pipeline_operations, + parent_otel_span=parent_otel_span, + ) + + async def reserve_io_tokens( + self, + descriptors: Sequence[RateLimitDescriptor], + estimated_input_tokens: int, + estimated_output_tokens: int, + parent_otel_span: Span | None = None, + ) -> tuple[RateLimitResponse, int, int]: + """ + Reserve ``estimated_input_tokens`` against project ITPM descriptors + and ``estimated_output_tokens`` against project OTPM descriptors. + + ITPM and OTPM are reserved from different-sized estimates, so unlike + same-size TPM descriptors they can't share a single + ``atomic_check_and_increment_by_n`` call -- each bucket gets its own + all-or-nothing atomic call. If the OTPM reservation is over limit + after ITPM already succeeded, the ITPM reservation this call made is + rolled back before returning, so a partial reservation never leaks. + + Returns ``(response, itpm_reserved, otpm_reserved)`` -- the latter two + are the amounts actually reserved (0 if that bucket wasn't + configured, or if the reservation failed), for the caller to stash + for post-call reconciliation. + """ + itpm_descriptors = [ # mutable-ok: atomic limiter API requires lists + d for d in descriptors if d["key"] == PROJECT_ITPM_DESCRIPTOR_KEY + ] + otpm_descriptors = [ # mutable-ok: atomic limiter API requires lists + d for d in descriptors if d["key"] == PROJECT_OTPM_DESCRIPTOR_KEY + ] + + if not itpm_descriptors and not otpm_descriptors: + return RateLimitResponse(overall_code="OK", statuses=[]), 0, 0 # mutable-ok: response contract uses a list + + itpm_response: RateLimitResponse | None = None + itpm_reserved = 0 + + if itpm_descriptors: + itpm_response = await self.atomic_check_and_increment_by_n( + descriptors=itpm_descriptors, + increments=[ # mutable-ok: atomic limiter API requires mutable increment records + {"tokens": estimated_input_tokens} # mutable-ok: atomic limiter increment record + for _ in itpm_descriptors + ], + parent_otel_span=parent_otel_span, + ) + if itpm_response["overall_code"] == "OVER_LIMIT": + return itpm_response, 0, 0 + itpm_reserved = estimated_input_tokens + + if otpm_descriptors: + otpm_response = await self.atomic_check_and_increment_by_n( + descriptors=otpm_descriptors, + increments=[ # mutable-ok: atomic limiter API requires mutable increment records + {"tokens": estimated_output_tokens} # mutable-ok: atomic limiter increment record + for _ in otpm_descriptors + ], + parent_otel_span=parent_otel_span, + ) + if otpm_response["overall_code"] == "OVER_LIMIT": + if itpm_reserved > 0: + await self._refund_reserved_tokens( + scopes=[ # mutable-ok: reservation rollback accepts collected scopes + (d["key"], d["value"]) for d in itpm_descriptors + ], + amount=itpm_reserved, + reservation_windows=itpm_response.get("reservation_windows", frozenset()), + parent_otel_span=parent_otel_span, + ) + return otpm_response, 0, 0 + statuses = ( + [ # mutable-ok: response contract uses a list + *itpm_response["statuses"], + *otpm_response["statuses"], + ] + if itpm_response is not None + else otpm_response["statuses"] + ) + return ( + RateLimitResponse( + overall_code="OK", + statuses=statuses, + reservation_windows=( + ( + itpm_response.get("reservation_windows", frozenset()) + if itpm_response is not None + else frozenset() + ) + | otpm_response.get("reservation_windows", frozenset()) + ), + ), + itpm_reserved, + estimated_output_tokens, + ) + + assert itpm_response is not None + return itpm_response, itpm_reserved, 0 + def create_organization_rate_limit_descriptor( self, user_api_key_dict: UserAPIKeyAuth, requested_model: str | None = None ) -> list[RateLimitDescriptor]: @@ -2317,6 +2784,62 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) ) + def _add_project_io_token_rate_limit_descriptors_from_metadata( + self, + user_api_key_dict: UserAPIKeyAuth, + requested_model: str | None, + descriptors: _RateLimitDescriptorSink, + ) -> None: + """Add project-scoped ITPM/OTPM descriptors from project_metadata. + + Enforced independently of, and alongside, the combined ``model_per_project`` + TPM descriptor above -- these give Bedrock Mantle-style separate input/output + token quotas at the project level. + """ + if requested_model is None or user_api_key_dict.project_id is None: + return + + itpm_limit_for_project_model = ( + get_model_rate_limit_from_metadata(user_api_key_dict, "project_metadata", "model_itpm_limit") + or {} # mutable-ok: metadata helper returns an optional mapping + ) + otpm_limit_for_project_model = ( + get_model_rate_limit_from_metadata(user_api_key_dict, "project_metadata", "model_otpm_limit") + or {} # mutable-ok: metadata helper returns an optional mapping + ) + + model_itpm_limit = itpm_limit_for_project_model.get(requested_model) + model_otpm_limit = otpm_limit_for_project_model.get(requested_model) + + if model_itpm_limit is None and model_otpm_limit is None: + return + + descriptor_value = f"{user_api_key_dict.project_id}:{requested_model}" + if model_itpm_limit is not None: + descriptors.append( + RateLimitDescriptor( + key=PROJECT_ITPM_DESCRIPTOR_KEY, + value=descriptor_value, + rate_limit={ # mutable-ok: descriptor TypedDict requires a runtime dict + "requests_per_unit": None, + "tokens_per_unit": model_itpm_limit, + "window_size": self.window_size, + }, + ) + ) + if model_otpm_limit is not None: + descriptors.append( + RateLimitDescriptor( + key=PROJECT_OTPM_DESCRIPTOR_KEY, + value=descriptor_value, + rate_limit={ # mutable-ok: descriptor TypedDict requires a runtime dict + "requests_per_unit": None, + "tokens_per_unit": model_otpm_limit, + "window_size": self.window_size, + }, + ) + ) + def _handle_rate_limit_error( self, response: RateLimitResponse, @@ -2361,6 +2884,396 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): llm_provider=llm_provider, ) + @staticmethod + def _estimate_audio_block_tokens(block: object) -> int: + """ + 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``), + assuming the lowest reasonable audio bitrate so we never under-reserve + for higher-quality recordings of the same duration. + + When no payload is present (reference-only block or missing ``data``), + falls back to ``DEFAULT_AUDIO_TOKEN_ESTIMATE``. + """ + if not isinstance(block, dict): + return DEFAULT_AUDIO_TOKEN_ESTIMATE + input_audio = block.get("input_audio") + b64_data = input_audio.get("data") if isinstance(input_audio, dict) else None + if b64_data and isinstance(b64_data, str): + decoded_bytes = len(b64_data) * 3 // 4 + return max(decoded_bytes // _AUDIO_BYTES_PER_TOKEN, DEFAULT_AUDIO_TOKEN_ESTIMATE) + return DEFAULT_AUDIO_TOKEN_ESTIMATE + + @classmethod + def _estimate_audio_content_tokens(cls, messages: object) -> int: + """ + Sum of per-block audio token estimates across all ``messages``. + Returns 0 when there are no ``input_audio`` blocks, which the caller + uses to skip the (relatively expensive) strip pass. + """ + if not isinstance(messages, list): + return 0 + total = 0 + for message in messages: + content = message.get("content") if isinstance(message, dict) else None + if not isinstance(content, list): + continue + total += sum( + cls._estimate_audio_block_tokens(block) + for block in content + if isinstance(block, dict) and block.get("type") == "input_audio" + ) + return total + + @staticmethod + 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. + """ + if not isinstance(messages, list): + return messages + sanitized = [] # mutable-ok: token_counter requires a list of message dicts + for message in messages: + if not isinstance(message, dict): + sanitized.append(message) + continue + content = message.get("content") + if not isinstance(content, list): + sanitized.append(message) + continue + filtered_content = [ # mutable-ok: token_counter requires list content blocks + block for block in content if not (isinstance(block, dict) and block.get("type") == "input_audio") + ] + sanitized.append( # mutable-ok: token_counter requires mutable message dicts + {**message, "content": filtered_content} # mutable-ok: token_counter requires message dicts + ) + return sanitized + + @staticmethod + def _responses_input_to_chat_messages(data: object) -> Sequence[object]: + """ + Convert a Responses API ``input`` (string or list of input items) into + chat-completion-style messages via the standard LiteLLM transformation + (the same one guardrails use, e.g. ``purview_dlp.py``), so multimodal + ``input_image``/``input_text`` content blocks get counted by + ``token_counter``'s ``messages`` path instead of silently contributing + zero tokens via its ``text`` path, which only joins plain strings. + """ + from litellm.responses.litellm_completion_transformation.transformation import ( + LiteLLMCompletionResponsesConfig, + ) + + if not isinstance(data, dict): + return () + return LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages( + input=data.get("input") or "", + responses_api_request=data, + ) + + @classmethod + def _contains_responses_file_reference(cls, value: object) -> bool: + if isinstance(value, dict): + return value.get("type") == "input_file" or any( + cls._contains_responses_file_reference(child) for child in value.values() + ) + if isinstance(value, list): + return any(cls._contains_responses_file_reference(child) for child in value) + return False + + @classmethod + def _contains_unmeasurable_chat_media(cls, value: object) -> bool: + if isinstance(value, dict): + return value.get("type") in ("document", "file", "video_url") or any( + cls._contains_unmeasurable_chat_media(child) for child in value.values() + ) + if isinstance(value, list): + return any(cls._contains_unmeasurable_chat_media(child) for child in value) + return False + + @classmethod + def _contains_image_content(cls, value: object) -> bool: + if isinstance(value, dict): + media_type = value.get("media_type") or value.get("mime_type") + return ( + value.get("type") in ("image", "image_url", "input_image") + or (isinstance(media_type, str) and media_type.startswith("image/")) + or any(cls._contains_image_content(child) for child in value.values()) + ) + if isinstance(value, list): + return any(cls._contains_image_content(child) for child in value) + return False + + @classmethod + def _requires_conservative_responses_input_reservation(cls, data: object, call_type: str | None) -> bool: + if not isinstance(data, dict): + return False + return call_type in RESPONSES_API_CALL_TYPES and ( + data.get("previous_response_id") is not None or cls._contains_responses_file_reference(data.get("input")) + ) + + @staticmethod + def _count_pretokenized_embedding_input(value: object) -> int | None: + if not isinstance(value, list): + return None + if all(isinstance(token, int) for token in value): + return len(value) + if all( + isinstance(token_ids, list) and all(isinstance(token, int) for token in token_ids) for token_ids in value + ): + return sum(len(token_ids) for token_ids in value) + return None + + @staticmethod + def _rerank_input_to_text(data: Mapping[str, object]) -> str: + documents = data.get("documents") + document_items: Sequence[object] = documents if isinstance(documents, list) else () # pyright: ignore[reportUnknownVariableType] # rerank documents are validated runtime JSON + input_parts: tuple[object, ...] = ( # pyright: ignore[reportUnknownVariableType] # list narrowing preserves unknown JSON element types + data.get("query"), + *document_items, + ) + return "\n".join( + str(part) # pyright: ignore[reportUnknownArgumentType] # accepted document dicts have provider-defined fields + for part in input_parts # pyright: ignore[reportUnknownVariableType] # runtime JSON list elements remain unknown after list narrowing + if isinstance(part, (str, dict)) + ) + + def _estimate_precise_input_tokens(self, data: object, model: str | None, call_type: str | None = None) -> int: + """ + Model-aware input token estimate for the project ITPM reservation, + using ``litellm.token_counter`` -- the same approach the + deployment-level itpm/otpm check uses in + ``io_token_rate_limit_check.py``. Unlike the cheap char-count + estimate the combined-TPM path uses, this accounts for image/tool + content and derives per-``input_audio``-block estimates from the + base64 payload size (assuming the lowest reasonable bitrate so + longer recordings always reserve proportionally more), so a burst + of multimodal, tool-heavy, or audio-heavy requests can't each + reserve only the one-token floor and blow past ITPM before + post-call reconciliation catches up. + + For the Responses API, ``input`` is converted to chat messages first + (via ``_responses_input_to_chat_messages``) so its own multimodal + content blocks are counted the same way; ``token_counter``'s ``text`` + argument can only see plain strings in a list, not content blocks. + + Falls back to the cheap char-count estimate if ``token_counter`` + can't resolve a tokenizer for this model (e.g. an unrecognized + custom model name) or otherwise raises -- the audio add-on still + applies on top of the fallback. + """ + from litellm import token_counter + + if not isinstance(data, dict): + return 0 + selected_text = None + countable_tools = data.get("tools") + countable_tool_choice = data.get("tool_choice") + if call_type in RESPONSES_API_CALL_TYPES: + messages = self._responses_input_to_chat_messages(data) + elif (translated_request := self._translate_google_genai_native_request(data, call_type)) is not None: + messages = translated_request.get("messages") + countable_tools = translated_request.get("tools") + countable_tool_choice = translated_request.get("tool_choice") + elif self._is_embedding_request(data, call_type): + messages = None + selected_text = data.get("input") + pretokenized_input_tokens = self._count_pretokenized_embedding_input(selected_text) + if pretokenized_input_tokens is not None: + return pretokenized_input_tokens + elif call_type in RERANK_API_CALL_TYPES: + messages = None + selected_text = self._rerank_input_to_text(data) # pyright: ignore[reportUnknownArgumentType] # proxy request bodies are runtime-validated JSON + elif call_type in TEXT_COMPLETION_API_CALL_TYPES: + messages = None + selected_text = data.get("prompt") + else: + messages = data.get("messages") + if messages is None: + selected_text = data.get("prompt") + if messages is None and selected_text is None: + selected_text = data.get("input") + + audio_token_estimate = self._estimate_audio_content_tokens(messages) + countable_messages = self._strip_audio_content_blocks(messages) if audio_token_estimate > 0 else messages + + try: + estimate = max( + 0, + int( + token_counter( + model=model or "", + messages=countable_messages, + text=selected_text, + tools=countable_tools, + tool_choice=countable_tool_choice, + use_default_image_token_count=True, + ) + ), + ) + return estimate + audio_token_estimate + except Exception: # noqa: BLE001 - any tokenizer/model-resolution/transform failure degrades to the cheap estimate, never a 500 + if call_type in RERANK_API_CALL_TYPES and isinstance(selected_text, str): + return max(0, len(selected_text) // DEFAULT_CHARS_PER_TOKEN) + estimated_input_tokens, _ = self._estimate_input_and_output_tokens(data=data, call_type=call_type) + return estimated_input_tokens + audio_token_estimate + + async def _reserve_project_io_tokens_or_raise( + self, + descriptors: Sequence[RateLimitDescriptor], + data: object, + requested_model: str | None, + user_api_key_dict: UserAPIKeyAuth, + tpm_reservation_scopes: Sequence[tuple[str, str]], + tpm_reservation_amount: int, + call_type: str | None = None, + ) -> None: + """ + Reserve project-scoped ITPM/OTPM tokens (Bedrock Mantle-style + separate input/output token buckets), independently of -- and, when + both are configured, in addition to -- the combined-TPM reservation + the caller already made. Raises (via ``_handle_rate_limit_error``) on + an over-limit reservation, first rolling back the combined-TPM + reservation named by ``tpm_reservation_scopes``/``tpm_reservation_amount`` + if one was made, so a partial reservation never leaks. + """ + if not isinstance(data, dict): + return + stash = claim_request_stash_for_data(data) + io_token_descriptors = [ # mutable-ok: reservation API requires descriptor lists + d for d in descriptors if d["key"] in (PROJECT_ITPM_DESCRIPTOR_KEY, PROJECT_OTPM_DESCRIPTOR_KEY) + ] + if not io_token_descriptors: + return + + configured_otpm_limits = [ # mutable-ok: min calculation materializes validated limits + int(v) + for d in io_token_descriptors + if d["key"] == PROJECT_OTPM_DESCRIPTOR_KEY + for v in [ # mutable-ok: comprehension binds the optional descriptor value + (d.get("rate_limit") or {}).get( # mutable-ok: optional descriptor fallback + "tokens_per_unit" + ) + ] + if v is not None + ] + min_configured_otpm_limit = min(configured_otpm_limits) if configured_otpm_limits else None + configured_itpm_limits = [ # mutable-ok: min calculation materializes validated limits + int(v) + for d in io_token_descriptors + if d["key"] == PROJECT_ITPM_DESCRIPTOR_KEY + for v in [ # mutable-ok: comprehension binds the optional descriptor value + (d.get("rate_limit") or {}).get( # mutable-ok: optional descriptor fallback + "tokens_per_unit" + ) + ] + if v is not None + ] + min_configured_itpm_limit = min(configured_itpm_limits) if configured_itpm_limits else None + + _, estimated_output_tokens = self._estimate_input_and_output_tokens( + data=data, + min_configured_tpm_limit=min_configured_otpm_limit, + call_type=call_type, + ) + estimated_input_tokens = ( + min_configured_itpm_limit + if min_configured_itpm_limit is not None + and ( + self._requires_conservative_responses_input_reservation(data, call_type) + or self._contains_unmeasurable_chat_media(data.get("messages")) + or self._contains_image_content(data) + ) + else self._estimate_precise_input_tokens(data=data, model=requested_model, call_type=call_type) + ) + estimated_input_tokens = max(estimated_input_tokens, 1) + if not self._has_explicit_output_cap(data, call_type): + estimated_output_tokens = max(estimated_output_tokens, 1) + + # Hard-cap generation length so an unbounded response can't overshoot + # the OTPM budget before post-call reconciliation runs, mirroring the + # combined-TPM floor cap in the caller. + self._apply_implicit_output_cap( + data=data, + min_configured_limit=min_configured_otpm_limit, + call_type=call_type, + ) + + io_response, itpm_reserved, otpm_reserved = await self.reserve_io_tokens( + descriptors=io_token_descriptors, + estimated_input_tokens=estimated_input_tokens, + estimated_output_tokens=estimated_output_tokens, + parent_otel_span=user_api_key_dict.parent_otel_span, + ) + + if io_response["overall_code"] == "OVER_LIMIT": + # A combined-TPM reservation may have already succeeded above for + # this same request; refund it too, or its counter stays inflated + # until the window's TTL expires. Mark it released so the + # ProxyRateLimitError we're about to raise doesn't get refunded + # a second time when async_post_call_failure_hook sees the same + # (still-stashed) reservation and refunds it again. + if tpm_reservation_amount > 0: + await self._refund_reserved_tokens( + scopes=tpm_reservation_scopes, + amount=tpm_reservation_amount, + parent_otel_span=user_api_key_dict.parent_otel_span, + ) + stash.reservation_released = True + acquisition = stash.parallel_slot + if acquisition is not None: + await self._release_parallel_request_slots( + acquisition=acquisition, + parent_otel_span=user_api_key_dict.parent_otel_span, + ) + stash.parallel_slot = None + self._handle_rate_limit_error( + response=io_response, + descriptors=descriptors, + requested_model=requested_model, + ) + + if itpm_reserved > 0: + itpm_scopes = [ # mutable-ok: request stash freezes the collected scopes + (d["key"], d["value"]) for d in io_token_descriptors if d["key"] == PROJECT_ITPM_DESCRIPTOR_KEY + ] + stash.itpm_reserved_tokens = itpm_reserved + stash.itpm_reserved_scopes = frozenset(itpm_scopes) + stash.itpm_reserved_window_identities = frozenset( + (counter_key, window_start, backend) + for counter_key, window_start, backend in io_response.get("reservation_windows", frozenset()) + if "model_per_project_itpm" in counter_key + ) + if otpm_reserved > 0: + otpm_scopes = [ # mutable-ok: request stash freezes the collected scopes + (d["key"], d["value"]) for d in io_token_descriptors if d["key"] == PROJECT_OTPM_DESCRIPTOR_KEY + ] + stash.otpm_reserved_tokens = otpm_reserved + stash.otpm_reserved_scopes = frozenset(otpm_scopes) + stash.otpm_reserved_window_identities = frozenset( + (counter_key, window_start, backend) + for counter_key, window_start, backend in io_response.get("reservation_windows", frozenset()) + if "model_per_project_otpm" in counter_key + ) + + if stash.rate_limit_response is not None: + stash.rate_limit_response["statuses"].extend(io_response["statuses"]) + elif io_response["statuses"]: + stash.rate_limit_response = io_response + + verbose_proxy_logger.debug( + "ITPM/OTPM tokens reserved: itpm=%s, otpm=%s for model %s", + itpm_reserved, + otpm_reserved, + requested_model, + ) + async def async_pre_call_hook( self, user_api_key_dict: UserAPIKeyAuth, @@ -2433,6 +3346,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): requested_model=requested_model, descriptors=descriptors, ) + self._add_project_io_token_rate_limit_descriptors_from_metadata( + user_api_key_dict=user_api_key_dict, + requested_model=requested_model, + descriptors=descriptors, + ) # Org Level Rate Limits descriptors.extend(self.create_organization_rate_limit_descriptor(user_api_key_dict, requested_model)) @@ -2489,28 +3407,33 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): configured_tpm_limits: Final = [ int(v) for d in descriptors + if d["key"] not in (PROJECT_ITPM_DESCRIPTOR_KEY, PROJECT_OTPM_DESCRIPTOR_KEY) for v in [(d.get("rate_limit") or {}).get("tokens_per_unit")] if v is not None ] has_tpm_limits: Final = bool(configured_tpm_limits) + # Populated on a successful combined-TPM reservation below, so the + # project ITPM/OTPM block further down can roll it back if a + # different bucket in the same request subsequently hits its + # limit. Stays empty/0 whenever no combined-TPM reservation was + # made (or it was over limit, in which case execution never + # reaches the ITPM/OTPM block -- `_handle_rate_limit_error` raises). + tpm_reservation_scopes: Sequence[tuple[str, str]] = () + tpm_reservation_amount = 0 + if has_tpm_limits and self.tpm_reservation_enabled: min_configured_tpm_limit: Final = min(configured_tpm_limits) # When the configured TPM cap is small enough to constrain the - # no-max_tokens floor, also hard-cap the model output via - # data["max_tokens"] so concurrent unbounded generations can't - # spend past the limit before post-call reconciliation runs. - # Skip when the request already sets max_tokens or has no - # generation budget at all (embeddings). - capped_floor: Final = self._no_max_tokens_output_floor(min_configured_tpm_limit) - baseline_floor: Final = DEFAULT_MAX_TOKENS_ESTIMATE // _TPM_FLOOR_FRACTION - has_explicit_max_tokens: Final = ( - data.get("max_tokens") is not None or data.get("max_completion_tokens") is not None + # no-max_tokens floor, also hard-cap the model output so + # concurrent unbounded generations can't spend past the limit + # before post-call reconciliation runs. + self._apply_implicit_output_cap( + data=data, + min_configured_limit=min_configured_tpm_limit, + call_type=call_type, ) - is_embedding: Final = data.get("input") is not None - if capped_floor < baseline_floor and not has_explicit_max_tokens and not is_embedding: - data["max_tokens"] = capped_floor # Floor at 1 token so contentless requests (/responses, # tool-call continuations, empty messages) still flow @@ -2524,6 +3447,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): data=data, model=requested_model, min_configured_tpm_limit=min_configured_tpm_limit, + call_type=call_type, ), 1, ) @@ -2557,8 +3481,14 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): stash.reserved_scopes = frozenset( (d["key"], d["value"]) for d in descriptors - if (d.get("rate_limit") or {}).get("tokens_per_unit") is not None + if d["key"] not in (PROJECT_ITPM_DESCRIPTOR_KEY, PROJECT_OTPM_DESCRIPTOR_KEY) + and (d.get("rate_limit") or {}).get( # mutable-ok: optional descriptor fallback + "tokens_per_unit" + ) + is not None ) + tpm_reservation_scopes = tuple(stash.reserved_scopes) + tpm_reservation_amount = estimated_tokens # Merge TPM statuses into the stored rate-limit response # so x-ratelimit-{key}-remaining-tokens / -limit-tokens @@ -2573,6 +3503,16 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): "TPM tokens reserved: %s for model %s", estimated_tokens, requested_model ) + await self._reserve_project_io_tokens_or_raise( + descriptors=descriptors, + data=data, + requested_model=requested_model, + user_api_key_dict=user_api_key_dict, + tpm_reservation_scopes=tpm_reservation_scopes, + tpm_reservation_amount=tpm_reservation_amount, + call_type=call_type, + ) + def _create_pipeline_operations( self, key: str, @@ -2751,6 +3691,113 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): litellm_parent_otel_span=parent_otel_span, ) + async def _apply_local_window_guarded_token_increments( + self, + operations: Sequence[ReservationAwareIncrementOperation], + parent_otel_span: Span | None = None, + ) -> None: + async with self._check_and_increment_lock: + for operation in operations: + window_key = operation.get("window_key") + expected_window_start = operation.get("expected_window_start") + if window_key is None or expected_window_start is None: + continue + active_window_start = await self.internal_usage_cache.async_get_cache( + key=window_key, + litellm_parent_otel_span=parent_otel_span, + local_only=True, + ) + if active_window_start is None or str(active_window_start) != expected_window_start: + continue + current_counter = ( + await self.internal_usage_cache.async_get_cache( + key=operation["key"], + litellm_parent_otel_span=parent_otel_span, + local_only=True, + ) + or 0 + ) + await self.internal_usage_cache.async_set_cache( + key=operation["key"], + value=float(current_counter) + operation["increment_value"], + ttl=operation["ttl"], + litellm_parent_otel_span=parent_otel_span, + local_only=True, + ) + + async def _apply_redis_window_guarded_token_increments( + self, + operations: Sequence[ReservationAwareIncrementOperation], + parent_otel_span: Span | None = None, + ) -> None: + for operation in operations: + window_key = operation.get("window_key") + expected_window_start = operation.get("expected_window_start") + if window_key is None or expected_window_start is None: + continue + if self.window_guarded_token_increment_script is not None: + try: + await self.window_guarded_token_increment_script( + keys=[window_key, operation["key"]], + args=[ + expected_window_start, + operation["increment_value"], + operation["ttl"] or 0, + ], + ) + continue + except Exception as e: + verbose_proxy_logger.warning( + "Window-guarded token adjustment failed for %s: %s", + operation["key"], + e, + ) + if operation["increment_value"] > 0: + await self.internal_usage_cache.async_increment_cache( + key=operation["key"], + value=operation["increment_value"], + litellm_parent_otel_span=parent_otel_span, + ttl=operation["ttl"], + ) + + async def async_increment_reservation_aware_tokens( + self, + pipeline_operations: Sequence[ReservationAwareIncrementOperation], + parent_otel_span: Span | None = None, + ) -> None: + for operation in pipeline_operations: + if operation.get("window_key") is None or operation.get("expected_window_start") is None: + await self.internal_usage_cache.async_increment_cache( + key=operation["key"], + value=operation["increment_value"], + litellm_parent_otel_span=parent_otel_span, + ttl=operation["ttl"], + ) + local_guarded_operations: Final = tuple( + operation + for operation in pipeline_operations + if operation.get("window_key") is not None + and operation.get("expected_window_start") is not None + and operation.get("reservation_backend") == "local" + ) + redis_guarded_operations: Final = tuple( + operation + for operation in pipeline_operations + if operation.get("window_key") is not None + and operation.get("expected_window_start") is not None + and operation.get("reservation_backend") != "local" + ) + if local_guarded_operations: + await self._apply_local_window_guarded_token_increments( + operations=local_guarded_operations, + parent_otel_span=parent_otel_span, + ) + if redis_guarded_operations: + await self._apply_redis_window_guarded_token_increments( + operations=redis_guarded_operations, + parent_otel_span=parent_otel_span, + ) + def get_rate_limit_type(self) -> Literal["output", "input", "total"]: from litellm.proxy.proxy_server import general_settings @@ -2780,6 +3827,163 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): merged[f"{prefix}-limit-{status['rate_limit_type']}"] = status["current_limit"] return merged + @staticmethod + def _resolve_rerank_token_usage(response_obj: object) -> tuple[int, int, bool] | None: + if not isinstance(response_obj, RerankResponse) or response_obj.meta is None: + return None + + rerank_tokens = response_obj.meta.get("tokens") # pyright: ignore[reportUnknownMemberType] # TypedDict's optional generic metadata widens get overloads + if rerank_tokens is not None: + input_tokens = rerank_tokens.get("input_tokens") or 0 # pyright: ignore[reportUnknownMemberType] # token fields are typed integers despite the generic get overload + output_tokens = rerank_tokens.get("output_tokens") or 0 # pyright: ignore[reportUnknownMemberType] # token fields are typed integers despite the generic get overload + if input_tokens or output_tokens: + return max(0, input_tokens), max(0, output_tokens), True + + billed_units = response_obj.meta.get("billed_units") # pyright: ignore[reportUnknownMemberType] # TypedDict's optional generic metadata widens get overloads + if billed_units is not None: + total_tokens = billed_units.get("total_tokens") or 0 # pyright: ignore[reportUnknownMemberType] # billed total is a typed integer despite the generic get overload + if total_tokens: + return max(0, total_tokens), 0, True + return None + + def _resolve_io_token_reconcile_usage( + self, + response_obj: object, + ) -> tuple[int, int, bool]: + """ + Resolve ``(billable_input_tokens, completion_tokens, usage_resolved)`` + for ITPM/OTPM reconciliation. Cache-read tokens are excluded from + billable input -- Bedrock Mantle doesn't count them toward ITPM -- + but they're untouched everywhere else (cost/usage logging still sees + the full prompt token count). + """ + rerank_usage = self._resolve_rerank_token_usage(response_obj) + if rerank_usage is not None: + return rerank_usage + + usage: object | None = None + if isinstance(response_obj, (Usage, ResponseAPIUsage)): + usage = response_obj + elif isinstance( + response_obj, + (ModelResponse, EmbeddingResponse, TextCompletionResponse, BaseLiteLLMOpenAIResponseObject), + ): + usage = getattr(response_obj, "usage", None) + elif isinstance(response_obj, dict): + usage = response_obj.get("usage") + if usage is None and any( + key in response_obj + for key in ( + "prompt_tokens", + "completion_tokens", + "input_tokens", + "output_tokens", + ) + ): + usage = response_obj + + if isinstance(usage, Usage): + prompt_tokens = usage.prompt_tokens or 0 + completion_tokens = usage.completion_tokens or 0 + cached_tokens = 0 + if usage.prompt_tokens_details is not None: + cached_tokens = getattr(usage.prompt_tokens_details, "cached_tokens", 0) or 0 + elif isinstance(usage, ResponseAPIUsage): + # Responses API usage uses input_tokens/output_tokens instead of + # prompt_tokens/completion_tokens. + prompt_tokens = usage.input_tokens or 0 + completion_tokens = usage.output_tokens or 0 + cached_tokens = 0 + if usage.input_tokens_details is not None: + cached_tokens = usage.input_tokens_details.cached_tokens or 0 + elif isinstance(usage, dict): + prompt_tokens = usage.get("prompt_tokens") or usage.get("input_tokens") or 0 + completion_tokens = usage.get("completion_tokens") or usage.get("output_tokens") or 0 + prompt_details = ( + usage.get("prompt_tokens_details") + or usage.get("input_tokens_details") + or {} # mutable-ok: usage details are optional mappings + ) + cached_tokens = ( + (prompt_details.get("cached_tokens", 0) or 0) if isinstance(prompt_details, dict) else 0 + ) or (usage.get("cache_read_input_tokens") or 0) + else: + return 0, 0, False + + if prompt_tokens == 0 and completion_tokens == 0: + return 0, 0, False + return max(0, prompt_tokens - cached_tokens), completion_tokens, True + + def _build_io_token_reservation_ops( + self, + kwargs: object, + response_obj: object, + ) -> list[RedisPipelineIncrementOperation] | tuple[ReservationAwareIncrementOperation, ...]: + """ + Reconcile project ITPM/OTPM reservations to actual usage on success: + ITPM to billable input tokens, OTPM to actual completion tokens. + Reuses ``_build_reservation_aware_tpm_ops``'s delta pattern -- ITPM/OTPM + are stored in the same ":tokens" cache bucket as combined TPM, just + under distinct scope keys, so the reservation-aware increment math is + identical; only the usage fields being reconciled against differ. + """ + if not isinstance(kwargs, dict): + return () + stash = get_request_stash_for_call(_call_id_from_callback_kwargs(kwargs)) + if stash is None: + return () + + itpm_reserved = stash.itpm_reserved_tokens + otpm_reserved = stash.otpm_reserved_tokens + if itpm_reserved <= 0 and otpm_reserved <= 0: + return () + + billable_input, completion_tokens, usage_resolved = self._resolve_io_token_reconcile_usage(response_obj) + if not usage_resolved: + billable_input, completion_tokens, usage_resolved = self._resolve_io_token_reconcile_usage( + kwargs.get("combined_usage_object") + ) + if not usage_resolved: + if not stash.reservation_released: + return () + billable_input = itpm_reserved + completion_tokens = otpm_reserved + + if stash.reservation_released or ( + not stash.itpm_reserved_window_identities and not stash.otpm_reserved_window_identities + ): + return self._build_reservation_aware_tpm_ops( + targets=tuple(stash.itpm_reserved_scopes), + reserved_scopes=frozenset() if stash.reservation_released else stash.itpm_reserved_scopes, + actual_tokens=billable_input, + reserved_tokens=0 if stash.reservation_released else itpm_reserved, + ) + self._build_reservation_aware_tpm_ops( + targets=tuple(stash.otpm_reserved_scopes), + reserved_scopes=frozenset() if stash.reservation_released else stash.otpm_reserved_scopes, + actual_tokens=completion_tokens, + reserved_tokens=0 if stash.reservation_released else otpm_reserved, + ) + + itpm_ops: Sequence[ReservationAwareIncrementOperation] = () + if itpm_reserved > 0: + itpm_ops = self._build_project_reservation_ops( + targets=tuple(stash.itpm_reserved_scopes), + reserved_scopes=frozenset() if stash.reservation_released else stash.itpm_reserved_scopes, + actual_tokens=billable_input, + reserved_tokens=itpm_reserved, + reservation_window_identities=stash.itpm_reserved_window_identities, + ) + otpm_ops: Sequence[ReservationAwareIncrementOperation] = () + if otpm_reserved > 0: + otpm_ops = self._build_project_reservation_ops( + targets=tuple(stash.otpm_reserved_scopes), + reserved_scopes=frozenset() if stash.reservation_released else stash.otpm_reserved_scopes, + actual_tokens=completion_tokens, + reserved_tokens=otpm_reserved, + reservation_window_identities=stash.otpm_reserved_window_identities, + ) + return tuple((*itpm_ops, *otpm_ops)) + def _collect_tpm_scope_targets( self, standard_logging_metadata: dict[str, Any], @@ -2844,8 +4048,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): def _build_reservation_aware_tpm_ops( self, - targets: list[tuple[str, str]], - reserved_scopes: frozenset[tuple[str, str]], + targets: Sequence[tuple[str, str]], + reserved_scopes: Set[tuple[str, str]], actual_tokens: int, reserved_tokens: int, ) -> list[RedisPipelineIncrementOperation]: @@ -2878,6 +4082,66 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) return ops + def _build_project_reservation_op( + self, + scope: tuple[str, str], + reserved_scopes: Set[tuple[str, str]], + actual_tokens: int, + reserved_tokens: int, + reservation_window_identities: frozenset[tuple[str, str, Literal["redis", "local"]]], + ) -> ReservationAwareIncrementOperation | None: + scope_key, scope_value = scope + is_reserved_scope: Final = scope in reserved_scopes + increment: Final = actual_tokens - reserved_tokens if is_reserved_scope else actual_tokens + if increment == 0: + return None + counter_key: Final = self.create_rate_limit_keys(scope_key, scope_value, "tokens") + window_identity: Final = next( + ( + (window_start, backend) + for identity_counter_key, window_start, backend in reservation_window_identities + if identity_counter_key == counter_key + ), + None, + ) + if not is_reserved_scope or window_identity is None: + return ReservationAwareIncrementOperation( + key=counter_key, + increment_value=increment, + ttl=self.window_size, + ) + return ReservationAwareIncrementOperation( + key=counter_key, + increment_value=increment, + ttl=self.window_size, + window_key=f"{{{scope_key}:{scope_value}}}:window", + expected_window_start=window_identity[0], + reservation_backend=window_identity[1], + ) + + def _build_project_reservation_ops( + self, + targets: Sequence[tuple[str, str]], + reserved_scopes: Set[tuple[str, str]], + actual_tokens: int, + reserved_tokens: int, + reservation_window_identities: frozenset[tuple[str, str, Literal["redis", "local"]]], + ) -> tuple[ReservationAwareIncrementOperation, ...]: + return tuple( + operation + for scope in targets + if ( + operation := self._build_project_reservation_op( + scope=scope, + reserved_scopes=reserved_scopes, + actual_tokens=actual_tokens, + reserved_tokens=reserved_tokens, + reservation_window_identities=reservation_window_identities, + ) + ) + is not None + ) + def _build_success_event_pipeline_operations( self, kwargs: Any, @@ -2994,12 +4258,26 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): response_obj=response_obj, rate_limit_type=rate_limit_type, ) - if pipeline_operations: await self.async_increment_tokens_with_ttl_preservation( pipeline_operations=pipeline_operations, parent_otel_span=litellm_parent_otel_span, ) + io_token_operations: Final = self._build_io_token_reservation_ops( + kwargs=kwargs, + response_obj=response_obj, + ) + if io_token_operations: + if isinstance(io_token_operations, list): + await self.async_increment_tokens_with_ttl_preservation( + pipeline_operations=io_token_operations, + parent_otel_span=litellm_parent_otel_span, + ) + else: + await self.async_increment_reservation_aware_tokens( + pipeline_operations=io_token_operations, + parent_otel_span=litellm_parent_otel_span, + ) except Exception as e: verbose_proxy_logger.exception("Error in rate limit success event: %s", e) @@ -3092,9 +4370,15 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): # already released it (proxy-level rejection that also bubbles up # here as an LLM-error callback). max_parallel_requests is its # own counter and is always decremented per call. - reserved_tokens = 0 - if stash is not None and not stash.reservation_released: + if stash is None or stash.reservation_released: + reserved_tokens = 0 + itpm_reserved = 0 + otpm_reserved = 0 + else: reserved_tokens = stash.reserved_tokens + itpm_reserved = stash.itpm_reserved_tokens + otpm_reserved = stash.otpm_reserved_tokens + if stash is not None and reserved_tokens > 0: verbose_proxy_logger.debug("Releasing reserved TPM tokens on failure: %s", reserved_tokens) # Refund only against the scopes the reservation actually @@ -3111,12 +4395,64 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) ) + # Refund project ITPM/OTPM reservations the same way -- full + # refund, since a failed call has no billable usage to reconcile + # against. + itpm_operations: Final = ( + self._build_project_reservation_ops( + targets=tuple(stash.itpm_reserved_scopes), + reserved_scopes=stash.itpm_reserved_scopes, + actual_tokens=0, + reserved_tokens=itpm_reserved, + reservation_window_identities=stash.itpm_reserved_window_identities, + ) + if stash is not None and itpm_reserved > 0 and stash.itpm_reserved_window_identities + else self._build_reservation_aware_tpm_ops( + targets=tuple(stash.itpm_reserved_scopes), + reserved_scopes=stash.itpm_reserved_scopes, + actual_tokens=0, + reserved_tokens=itpm_reserved, + ) + if stash is not None and itpm_reserved > 0 + else () + ) + + otpm_operations: Final = ( + self._build_project_reservation_ops( + targets=tuple(stash.otpm_reserved_scopes), + reserved_scopes=stash.otpm_reserved_scopes, + actual_tokens=0, + reserved_tokens=otpm_reserved, + reservation_window_identities=stash.otpm_reserved_window_identities, + ) + if stash is not None and otpm_reserved > 0 and stash.otpm_reserved_window_identities + else self._build_reservation_aware_tpm_ops( + targets=tuple(stash.otpm_reserved_scopes), + reserved_scopes=stash.otpm_reserved_scopes, + actual_tokens=0, + reserved_tokens=otpm_reserved, + ) + if stash is not None and otpm_reserved > 0 + else () + ) + if pipeline_operations: await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline( increment_list=pipeline_operations, litellm_parent_otel_span=litellm_parent_otel_span, ) - if stash is not None and reserved_tokens > 0: + for project_operations in (itpm_operations, otpm_operations): + if isinstance(project_operations, list): + await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline( + increment_list=project_operations, + litellm_parent_otel_span=litellm_parent_otel_span, + ) + elif project_operations: + await self.async_increment_reservation_aware_tokens( + pipeline_operations=project_operations, + parent_otel_span=litellm_parent_otel_span, + ) + if stash is not None and (reserved_tokens > 0 or itpm_reserved > 0 or otpm_reserved > 0): stash.reservation_released = True except Exception as e: verbose_proxy_logger.exception("Error in rate limit failure event: %s", e) @@ -3194,19 +4530,19 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): traceback_str: str | None = None, ) -> None: """ - Release the parallel-request slot and any TPM reservation when the - request is rejected after the pre-call hook acquired them but before - the LLM call ran (e.g. a downstream guardrail/auth hook raised). - Without this, those resources are stranded — async_log_failure_event - is a litellm completion-level callback and never fires for proxy-side - rejections, so a leaked slot would occupy the gauge for the full - PARALLEL_REQUEST_SLOT_TTL_SECONDS. + Release the parallel-request slot and any TPM/ITPM/OTPM reservation + when the request is rejected after the pre-call hook acquired them + but before the LLM call ran (e.g. a downstream guardrail/auth hook + raised). Without this, those resources are stranded — + async_log_failure_event is a litellm completion-level callback and + never fires for proxy-side rejections, so a leaked slot would occupy + the gauge for the full PARALLEL_REQUEST_SLOT_TTL_SECONDS. Idempotent: the slot release clears the stashed acquisition (and slot - removal is a no-op ZREM on a second run), and the TPM refund is - guarded by the stash's ``reservation_released`` flag — if both this - hook and async_log_failure_event end up running in the same flow, only - the first release/refund applies. + removal is a no-op ZREM on a second run), and the TPM/ITPM/OTPM + refund is guarded by the stash's ``reservation_released`` flag — if + both this hook and async_log_failure_event end up running in the same + flow, only the first release/refund applies. """ try: stash: Final = get_request_stash() @@ -3222,23 +4558,80 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): if stash.reservation_released: return reserved_tokens: Final = stash.reserved_tokens - if reserved_tokens <= 0: + itpm_reserved: Final = stash.itpm_reserved_tokens + otpm_reserved: Final = stash.otpm_reserved_tokens + if reserved_tokens <= 0 and itpm_reserved <= 0 and otpm_reserved <= 0: return - ops: Final = self._build_reservation_aware_tpm_ops( - targets=list(stash.reserved_scopes), - reserved_scopes=stash.reserved_scopes, - actual_tokens=0, - reserved_tokens=reserved_tokens, - ) - if ops: - verbose_proxy_logger.debug( - "Releasing reserved TPM tokens on proxy-level rejection: %s", reserved_tokens + combined_ops: Final = ( + self._build_reservation_aware_tpm_ops( + targets=tuple(stash.reserved_scopes), + reserved_scopes=stash.reserved_scopes, + actual_tokens=0, + reserved_tokens=reserved_tokens, ) + if reserved_tokens > 0 + else () + ) + itpm_ops: Final = ( + self._build_project_reservation_ops( + targets=tuple(stash.itpm_reserved_scopes), + reserved_scopes=stash.itpm_reserved_scopes, + actual_tokens=0, + reserved_tokens=itpm_reserved, + reservation_window_identities=stash.itpm_reserved_window_identities, + ) + if itpm_reserved > 0 and stash.itpm_reserved_window_identities + else self._build_reservation_aware_tpm_ops( + targets=tuple(stash.itpm_reserved_scopes), + reserved_scopes=stash.itpm_reserved_scopes, + actual_tokens=0, + reserved_tokens=itpm_reserved, + ) + if itpm_reserved > 0 + else () + ) + otpm_ops: Final = ( + self._build_project_reservation_ops( + targets=tuple(stash.otpm_reserved_scopes), + reserved_scopes=stash.otpm_reserved_scopes, + actual_tokens=0, + reserved_tokens=otpm_reserved, + reservation_window_identities=stash.otpm_reserved_window_identities, + ) + if otpm_reserved > 0 and stash.otpm_reserved_window_identities + else self._build_reservation_aware_tpm_ops( + targets=tuple(stash.otpm_reserved_scopes), + reserved_scopes=stash.otpm_reserved_scopes, + actual_tokens=0, + reserved_tokens=otpm_reserved, + ) + if otpm_reserved > 0 + else () + ) + if combined_ops or itpm_ops or otpm_ops: + verbose_proxy_logger.debug( + "Releasing reserved tokens on proxy-level rejection: tpm=%s, itpm=%s, otpm=%s", + reserved_tokens, + itpm_reserved, + otpm_reserved, + ) + if combined_ops: await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline( - increment_list=ops, + increment_list=combined_ops, litellm_parent_otel_span=user_api_key_dict.parent_otel_span, ) + for project_ops in (itpm_ops, otpm_ops): + if isinstance(project_ops, list): + await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline( + increment_list=project_ops, + litellm_parent_otel_span=user_api_key_dict.parent_otel_span, + ) + elif project_ops: + await self.async_increment_reservation_aware_tokens( + pipeline_operations=project_ops, + parent_otel_span=user_api_key_dict.parent_otel_span, + ) stash.reservation_released = True except Exception as e: verbose_proxy_logger.exception("Error releasing TPM reservation on post-call failure: %s", e) diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index 1c29c287c3a..872b447ad2d 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -3114,6 +3114,136 @@ async def test_project_model_rate_limits_not_triggered_for_other_model_v3(): ), f"model_per_project should not be added for unrelated model, got: {descriptor_keys}" +@pytest.mark.asyncio +async def test_project_model_itpm_otpm_limits_enforced_v3(): + """ + Project-level model_itpm_limit/model_otpm_limit must produce distinct + Bedrock Mantle-style input and output token descriptors. + """ + _api_key = hash_token("sk-project-io-test") + local_cache = DualCache() + parallel_request_handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + + captured_descriptors = [] + + async def mock_should_rate_limit(descriptors, **kwargs): + captured_descriptors.extend(descriptors) + return {"overall_code": "OK", "statuses": []} + + parallel_request_handler.should_rate_limit = mock_should_rate_limit + + user_api_key_dict = UserAPIKeyAuth( + api_key=_api_key, + project_id="proj-mantle", + project_metadata={ + "model_itpm_limit": {"bedrock_mantle/claude-opus": 20000000}, + "model_otpm_limit": {"bedrock_mantle/claude-opus": 4000000}, + }, + ) + + await parallel_request_handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data={"model": "bedrock_mantle/claude-opus"}, + call_type="", + ) + + descriptor_keys = [d["key"] for d in captured_descriptors] + assert "model_per_project_itpm" in descriptor_keys + assert "model_per_project_otpm" in descriptor_keys + assert "model_per_project" not in descriptor_keys + + itpm_descriptor = next( + d for d in captured_descriptors if d["key"] == "model_per_project_itpm" + ) + otpm_descriptor = next( + d for d in captured_descriptors if d["key"] == "model_per_project_otpm" + ) + assert itpm_descriptor["value"] == "proj-mantle:bedrock_mantle/claude-opus" + assert itpm_descriptor["rate_limit"]["tokens_per_unit"] == 20000000 + assert otpm_descriptor["value"] == "proj-mantle:bedrock_mantle/claude-opus" + assert otpm_descriptor["rate_limit"]["tokens_per_unit"] == 4000000 + + +@pytest.mark.asyncio +async def test_project_model_itpm_otpm_limits_not_triggered_for_other_model_v3(): + """Split project limits must not apply to an unrelated model.""" + _api_key = hash_token("sk-project-io-test-2") + local_cache = DualCache() + parallel_request_handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + + captured_descriptors = [] + + async def mock_should_rate_limit(descriptors, **kwargs): + captured_descriptors.extend(descriptors) + return {"overall_code": "OK", "statuses": []} + + parallel_request_handler.should_rate_limit = mock_should_rate_limit + + user_api_key_dict = UserAPIKeyAuth( + api_key=_api_key, + project_id="proj-mantle", + project_metadata={ + "model_itpm_limit": {"bedrock_mantle/claude-opus": 20000000}, + }, + ) + + await parallel_request_handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data={"model": "gpt-4"}, + call_type="", + ) + + descriptor_keys = [d["key"] for d in captured_descriptors] + assert "model_per_project_itpm" not in descriptor_keys + assert "model_per_project_otpm" not in descriptor_keys + + +@pytest.mark.asyncio +async def test_project_model_itpm_and_tpm_limits_coexist_v3(): + """Combined project TPM and split ITPM/OTPM limits are enforced together.""" + _api_key = hash_token("sk-project-io-test-3") + local_cache = DualCache() + parallel_request_handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + + captured_descriptors = [] + + async def mock_should_rate_limit(descriptors, **kwargs): + captured_descriptors.extend(descriptors) + return {"overall_code": "OK", "statuses": []} + + parallel_request_handler.should_rate_limit = mock_should_rate_limit + + user_api_key_dict = UserAPIKeyAuth( + api_key=_api_key, + project_id="proj-mantle", + project_metadata={ + "model_tpm_limit": {"bedrock_mantle/claude-opus": 1000}, + "model_itpm_limit": {"bedrock_mantle/claude-opus": 20000000}, + "model_otpm_limit": {"bedrock_mantle/claude-opus": 4000000}, + }, + ) + + await parallel_request_handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data={"model": "bedrock_mantle/claude-opus"}, + call_type="", + ) + + descriptor_keys = [d["key"] for d in captured_descriptors] + assert "model_per_project" in descriptor_keys + assert "model_per_project_itpm" in descriptor_keys + assert "model_per_project_otpm" in descriptor_keys + + @pytest.mark.asyncio async def test_pre_call_hook_keeps_internal_stash_out_of_request_body(): """Regression for #27001 / #35197: the limiter's per-request bookkeeping @@ -3190,7 +3320,7 @@ async def test_responses_route_body_untouched_by_pre_call_hook(caller_metadata): _api_key = hash_token("sk-responses-regression") user_api_key_dict = UserAPIKeyAuth( api_key=_api_key, - tpm_limit=1000, + tpm_limit=100000, rpm_limit=5, ) local_cache = DualCache() diff --git a/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py b/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py index f7bd37b412a..66fc84ab1e0 100644 --- a/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py +++ b/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py @@ -15,7 +15,7 @@ Redis. """ import asyncio -from datetime import datetime +from datetime import datetime, timedelta from typing import Any, Dict import pytest @@ -23,15 +23,25 @@ import pytest from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + PROJECT_ITPM_DESCRIPTOR_KEY, + PROJECT_OTPM_DESCRIPTOR_KEY, + _AUDIO_BYTES_PER_TOKEN, _PROXY_MaxParallelRequestsHandler_v3 as RateLimitHandler, ) from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + _call_id_from_callback_kwargs, _request_stash, get_or_create_request_stash, get_request_stash, ) from litellm.proxy.utils import InternalUsageCache, hash_token -from litellm.types.utils import ModelResponse, Usage +from litellm.types.llms.openai import ( + InputTokensDetails, + ResponseAPIUsage, + ResponsesAPIResponse, +) +from litellm.types.rerank import RerankResponse +from litellm.types.utils import ModelResponse, PromptTokensDetailsWrapper, Usage @pytest.fixture @@ -582,6 +592,47 @@ async def test_estimate_tokens_uses_max_tokens_when_explicit(rate_limiter): assert estimate == 4 + 25 +@pytest.mark.asyncio +async def test_estimate_tokens_honors_explicit_zero_max_tokens(rate_limiter): + """ + Regression for a Greptile finding: explicit_max_tokens was resolved via + `data.get("max_tokens") or data.get("max_completion_tokens") or + data.get("max_output_tokens")`, so an explicit 0 in the first field was + falsy and fell through to the next field (or the no-max_tokens floor), + silently discarding a caller's explicit zero-output request. + """ + handler, _cache = rate_limiter + + estimate = handler._estimate_tokens_for_request( + data={ + "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}" + ) + + +@pytest.mark.asyncio +async def test_estimate_tokens_honors_explicit_zero_max_output_tokens_for_responses( + rate_limiter, +): + handler, _cache = rate_limiter + + estimate = handler._estimate_tokens_for_request( + data={ + "input": "describe this image in detail", # 29 chars ~ 7 tokens + "max_output_tokens": 0, + }, + min_configured_tpm_limit=40, + call_type="aresponses", + ) + assert estimate == 23 + + @pytest.mark.asyncio async def test_estimate_tokens_zero_for_empty_embeddings(rate_limiter): """Embeddings have no output budget — reservation should equal input only.""" @@ -1197,5 +1248,2394 @@ async def test_small_tpm_cap_preserves_explicit_max_tokens(rate_limiter): assert data["max_tokens"] == 500 +@pytest.mark.asyncio +async def test_project_otpm_reservation_prevents_concurrent_bypass(rate_limiter): + """ + Bedrock Mantle-style OTPM: with a 100 OTPM limit and 5 concurrent + requests each reserving 50+ output tokens, upfront reservation must + reject the late arrivals -- not let all 5 through. Exercises the + in-memory fallback in ``atomic_check_and_increment_by_n`` for the + project-scoped ITPM/OTPM descriptors specifically. + """ + handler, cache = rate_limiter + + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-otpm-bypass"), + project_id="proj-mantle-bypass", + project_metadata={ + "model_otpm_limit": {"bedrock_mantle/claude-opus": 100}, + }, + ) + + request_data = { + "model": "bedrock_mantle/claude-opus", + "messages": [{"role": "user", "content": "hello"}], + "max_tokens": 50, + } + + async def make_request(request_id: int) -> Dict[str, Any]: + data = request_data.copy() + try: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data=data, + call_type="", + ) + return {"request_id": request_id, "success": True} + except Exception as e: + return { + "request_id": request_id, + "success": False, + "status_code": getattr(e, "status_code", None), + } + + 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 + ] + + assert len(rate_limited) > 0, ( + f"Expected some OTPM-rate-limited requests but all {len(successful)} succeeded." + ) + + +@pytest.mark.asyncio +async def test_project_otpm_rejects_multiple_completion_candidates(rate_limiter): + handler, cache = rate_limiter + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-otpm-multiple-candidates"), + project_id="proj-multiple-candidates", + project_metadata={ + "model_otpm_limit": {"bedrock_mantle/claude-opus": 500}, + }, + ) + data = { + "model": "bedrock_mantle/claude-opus", + "messages": [{"role": "user", "content": "hello"}], + "max_tokens": 100, + "n": 10, + } + + with pytest.raises(Exception) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data=data, + call_type="acompletion", + ) + + assert getattr(exc_info.value, "status_code", None) == 429 + + +@pytest.mark.asyncio +async def test_project_otpm_reserves_largest_conflicting_output_cap(rate_limiter): + handler, cache = rate_limiter + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-otpm-conflicting-caps"), + project_id="proj-conflicting-caps", + project_metadata={ + "model_otpm_limit": {"bedrock_mantle/claude-opus": 50}, + }, + ) + data = { + "model": "bedrock_mantle/claude-opus", + "messages": [{"role": "user", "content": "hello"}], + "max_tokens": 1, + "max_completion_tokens": 100, + } + + with pytest.raises(Exception) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data=data, + call_type="acompletion", + ) + + assert getattr(exc_info.value, "status_code", None) == 429 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "call_type", + ["agenerate_content", "agenerate_content_stream"], +) +@pytest.mark.parametrize("config_field", ["config", "generationConfig"]) +async def test_project_otpm_rejects_google_genai_native_output_cap( + rate_limiter, + call_type, + config_field, +): + handler, cache = rate_limiter + model = "gemini/gemini-3-flash-preview" + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-google-genai-native-otpm"), + project_id="project-google-genai-native-otpm", + project_metadata={"model_otpm_limit": {model: 50}}, + ) + + with pytest.raises(Exception) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data={ + "model": model, + "contents": [{"role": "user", "parts": [{"text": "Hello"}]}], + config_field: {"maxOutputTokens": 100}, + }, + call_type=call_type, + ) + + assert getattr(exc_info.value, "status_code", None) == 429 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "call_type", + ["agenerate_content", "agenerate_content_stream"], +) +@pytest.mark.parametrize("candidate_count_field", ["candidateCount", "candidate_count"]) +async def test_project_otpm_rejects_google_genai_native_candidate_count( + rate_limiter, + call_type, + candidate_count_field, +): + handler, cache = rate_limiter + model = "gemini/gemini-3-flash-preview" + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-google-genai-native-candidate-count"), + project_id="project-google-genai-native-candidate-count", + project_metadata={"model_otpm_limit": {model: 150}}, + ) + + with pytest.raises(Exception) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data={ + "model": model, + "contents": [{"role": "user", "parts": [{"text": "Hello"}]}], + "config": { + "maxOutputTokens": 50, + candidate_count_field: 4, + }, + }, + call_type=call_type, + ) + + assert getattr(exc_info.value, "status_code", None) == 429 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "call_type", + ["agenerate_content", "agenerate_content_stream"], +) +@pytest.mark.parametrize("config_field", [None, "config", "generationConfig"]) +async def test_project_otpm_injects_google_genai_native_output_cap( + rate_limiter, + call_type, + config_field, +): + handler, cache = rate_limiter + model = "gemini/gemini-3-flash-preview" + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-google-genai-native-implicit-otpm"), + project_id="project-google-genai-native-implicit-otpm", + project_metadata={"model_otpm_limit": {model: 40}}, + ) + data = { + "model": model, + "contents": [{"role": "user", "parts": [{"text": "Hello"}]}], + } + if config_field is not None: + data[config_field] = {"temperature": 0} + + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data=data, + call_type=call_type, + ) + + stash = get_request_stash() + assert stash is not None + assert stash.otpm_reserved_tokens == 10 + expected_config_field = config_field or "config" + assert data[expected_config_field]["maxOutputTokens"] == 10 + assert "max_tokens" not in data + + +@pytest.mark.asyncio +async def test_project_otpm_over_limit_rolls_back_itpm_reservation(rate_limiter): + """ + When ITPM reserves fine but OTPM is then over limit, the ITPM + reservation this same pre-call already made must be rolled back -- + otherwise it leaks until the window's TTL, silently shrinking the ITPM + budget for every other request in that minute. + """ + handler, cache = rate_limiter + + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-otpm-rollback"), + project_id="proj-mantle-rollback", + project_metadata={ + "model_itpm_limit": {"bedrock_mantle/claude-opus": 1000000}, + "model_otpm_limit": {"bedrock_mantle/claude-opus": 10}, + }, + ) + + itpm_counter_key = handler.create_rate_limit_keys( + key="model_per_project_itpm", + value="proj-mantle-rollback:bedrock_mantle/claude-opus", + rate_limit_type="tokens", + ) + + data = { + "model": "bedrock_mantle/claude-opus", + "messages": [{"role": "user", "content": "hello"}], + "max_tokens": 500, # blows past the 10-token OTPM limit + } + + with pytest.raises(Exception) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data=data, + call_type="", + ) + 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}" + ) + + +@pytest.mark.asyncio +async def test_project_itpm_reconciled_on_success_excludes_cached_tokens(rate_limiter): + """ + On success, ITPM reconciles to billable input tokens (prompt_tokens + minus cached_tokens) -- not raw prompt_tokens. Cached prompt-read tokens + are free under Bedrock Mantle and must not count against the ITPM quota, + even though they still appear in usage/cost logging elsewhere. + """ + handler, _cache = rate_limiter + + itpm_scope = ("model_per_project_itpm", "proj-mantle:model") + otpm_scope = ("model_per_project_otpm", "proj-mantle:model") + + stash = get_or_create_request_stash() + stash.itpm_reserved_tokens = 100 + stash.itpm_reserved_scopes = frozenset({itpm_scope}) + stash.otpm_reserved_tokens = 60 + stash.otpm_reserved_scopes = frozenset({otpm_scope}) + mock_kwargs = {} + + mock_response = ModelResponse( + id="test", + object="chat.completion", + created=int(datetime.now().timestamp()), + model="bedrock_mantle/claude-opus", + usage=Usage( + prompt_tokens=80, + completion_tokens=40, + total_tokens=120, + prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=30), + ), + choices=[], + ) + + increments = [] + + async def mock_increment(increment_list, **kwargs): + 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 + ) + + await handler.async_log_success_event( + kwargs=mock_kwargs, + response_obj=mock_response, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + itpm_adjustments = [i for i in increments if "model_per_project_itpm" in i["key"]] + otpm_adjustments = [i for i in increments if "model_per_project_otpm" in i["key"]] + + # billable_input = 80 - 30 cached = 50; delta = 50 - 100 reserved = -50 + assert any(i["increment"] == -50 for i in itpm_adjustments), ( + f"Expected a -50 ITPM adjustment (50 billable - 100 reserved), got: {itpm_adjustments}" + ) + # delta = 40 actual completion - 60 reserved = -20 + assert any(i["increment"] == -20 for i in otpm_adjustments), ( + f"Expected a -20 OTPM adjustment (40 actual - 60 reserved), got: {otpm_adjustments}" + ) + + +@pytest.mark.asyncio +async def test_project_reconciliation_does_not_decrement_later_window(): + current_time = datetime(2026, 8, 5, 12, 0, 0) + cache = DualCache() + handler = RateLimitHandler( + internal_usage_cache=InternalUsageCache(cache), + time_provider=lambda: current_time, + ) + handler.window_size = 60 + scope = (PROJECT_ITPM_DESCRIPTOR_KEY, "project:model") + descriptor = { + "key": scope[0], + "value": scope[1], + "rate_limit": {"tokens_per_unit": 1000, "window_size": 60}, + } + + reservation = await handler.atomic_check_and_increment_by_n( + descriptors=[descriptor], + 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 + ) + 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} + ) + + current_time += timedelta(seconds=61) + later_reservation = await handler.atomic_check_and_increment_by_n( + descriptors=[descriptor], + increments=[{"tokens": 20}], + ) + assert window_identity not in later_reservation["reservation_windows"] + + await handler.async_log_success_event( + kwargs={}, + response_obj=ModelResponse( + usage=Usage(prompt_tokens=10, completion_tokens=0, total_tokens=10) + ), + start_time=current_time, + end_time=current_time, + ) + + assert float(await cache.async_get_cache(key=counter_key, local_only=True) or 0) == 20 + + +@pytest.mark.asyncio +async def test_project_reconciliation_decrements_its_active_window(rate_limiter): + handler, cache = rate_limiter + scope = (PROJECT_ITPM_DESCRIPTOR_KEY, "project:model") + descriptor = { + "key": scope[0], + "value": scope[1], + "rate_limit": {"tokens_per_unit": 1000, "window_size": 60}, + } + reservation = await handler.atomic_check_and_increment_by_n( + descriptors=[descriptor], + 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 + ) + 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} + ) + + await handler.async_log_success_event( + kwargs={}, + response_obj=ModelResponse( + usage=Usage(prompt_tokens=10, completion_tokens=0, total_tokens=10) + ), + start_time=datetime.now(), + end_time=datetime.now(), + ) + + assert float(await cache.async_get_cache(key=counter_key, local_only=True) or 0) == 10 + + +@pytest.mark.asyncio +async def test_redis_window_guard_uses_reservation_identity_and_never_falls_back_negative( + rate_limiter, +): + handler, _cache = rate_limiter + calls = [] + + async def failing_guard(*, keys, args): + calls.append((keys, args)) + raise RuntimeError("redis unavailable") + + unguarded_calls = [] + + async def capture_unguarded(pipeline_operations, **_kwargs): + unguarded_calls.extend(pipeline_operations) + + handler.window_guarded_token_increment_script = failing_guard + handler.async_increment_tokens_with_ttl_preservation = capture_unguarded + await handler.async_increment_reservation_aware_tokens( + pipeline_operations=[ + { + "key": "{model_per_project_itpm:project:model}:tokens", + "increment_value": -90, + "ttl": 60, + "window_key": "{model_per_project_itpm:project:model}:window", + "expected_window_start": "1234", + "reservation_backend": "redis", + } + ] + ) + + assert calls == [ + ( + [ + "{model_per_project_itpm:project:model}:window", + "{model_per_project_itpm:project:model}:tokens", + ], + ["1234", -90, 60], + ) + ] + assert unguarded_calls == [] + + +@pytest.mark.asyncio +async def test_atomic_lua_response_carries_redis_window_identity(rate_limiter): + handler, _cache = rate_limiter + counter_key = "{model_per_project_itpm:project:model}:tokens" + meta = [ + { + "descriptor_key": PROJECT_ITPM_DESCRIPTOR_KEY, + "current_limit": 100, + "rate_limit_type": "tokens", + "counter_key": counter_key, + } + ] + + async def successful_reservation(*, keys, args): + return [0, 25, 1234] + + handler.check_and_increment_by_n_script = successful_reservation + assert await handler._atomic_lua_per_descriptor([]) == { + "overall_code": "OK", + "statuses": [], + } + + response = await handler._atomic_lua_per_descriptor( + descriptor_groups=[ + ( + [ + "{model_per_project_itpm:project:model}:window", + counter_key, + ], + [100, 25, 60, 60], + meta, + ) + ] + ) + + assert response["statuses"][0]["limit_remaining"] == 75 + assert response["reservation_windows"] == frozenset( + {(counter_key, "1234", "redis")} + ) + + +@pytest.mark.asyncio +async def test_project_itpm_otpm_released_on_failure(rate_limiter): + """On failure, the full ITPM and OTPM reservations must be refunded.""" + handler, _cache = rate_limiter + + itpm_scope = ("model_per_project_itpm", "proj-mantle:model") + otpm_scope = ("model_per_project_otpm", "proj-mantle:model") + + stash = get_or_create_request_stash() + stash.itpm_reserved_tokens = 100 + stash.itpm_reserved_scopes = frozenset({itpm_scope}) + stash.otpm_reserved_tokens = 60 + stash.otpm_reserved_scopes = frozenset({otpm_scope}) + mock_kwargs = {} + + increments = [] + + async def mock_increment(increment_list, **kwargs): + 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 + ) + + await handler.async_log_failure_event( + kwargs=mock_kwargs, + response_obj=None, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + itpm_releases = [i for i in increments if "model_per_project_itpm" in i["key"]] + otpm_releases = [i for i in increments if "model_per_project_otpm" in i["key"]] + + assert any(i["increment"] == -100 for i in itpm_releases), itpm_releases + assert any(i["increment"] == -60 for i in otpm_releases), otpm_releases + + +@pytest.mark.asyncio +async def test_proxy_rejection_refunds_itpm_otpm_by_their_own_amount_not_combined( + rate_limiter, +): + """ + Regression for a Greptile-flagged bug: when a project configures both a + combined model_tpm_limit and split model_itpm_limit/model_otpm_limit for + the same model, async_post_call_failure_hook's proxy-side refund path + used to decrement every token descriptor -- including the ITPM/OTPM + ones -- by the flat combined reservation amount, instead of each + bucket's own reserved amount. That drives the split counters negative + (or under-refunds them) instead of returning them to exactly zero. + """ + handler, cache = rate_limiter + + api_key = hash_token("sk-mixed-tpm-io") + user_api_key_dict = UserAPIKeyAuth( + api_key=api_key, + project_id="proj-mixed", + project_metadata={ + "model_tpm_limit": {"bedrock_mantle/claude-opus": 100000}, + "model_itpm_limit": {"bedrock_mantle/claude-opus": 100000}, + "model_otpm_limit": {"bedrock_mantle/claude-opus": 100000}, + }, + ) + + data = { + "model": "bedrock_mantle/claude-opus", + "messages": [ + {"role": "user", "content": "hello there, this is a test message"} + ], + "max_tokens": 60, + } + + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data=data, + call_type="", + ) + + tpm_counter_key = handler.create_rate_limit_keys( + key="model_per_project", + value="proj-mixed:bedrock_mantle/claude-opus", + rate_limit_type="tokens", + ) + itpm_counter_key = handler.create_rate_limit_keys( + key="model_per_project_itpm", + value="proj-mixed:bedrock_mantle/claude-opus", + rate_limit_type="tokens", + ) + otpm_counter_key = handler.create_rate_limit_keys( + key="model_per_project_otpm", + value="proj-mixed:bedrock_mantle/claude-opus", + 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 + ) + assert tpm_reserved > 0 and itpm_reserved > 0 and otpm_reserved > 0 + + await handler.async_post_call_failure_hook( + request_data=data, + original_exception=Exception("guardrail rejected"), + 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 + ) + + 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}" + ) + + +@pytest.mark.asyncio +async def test_proxy_rejection_refunds_itpm_otpm_only_reservation_with_no_combined_tpm( + rate_limiter, +): + """ + Regression for the second half of the same bug: with only + model_itpm_limit/model_otpm_limit configured (no model_tpm_limit), the + combined reserved_tokens is 0, and the proxy-side refund path used to + return immediately on that -- leaking the ITPM/OTPM reservations until + the rate-limit window's TTL expired. + """ + handler, cache = rate_limiter + + api_key = hash_token("sk-io-only") + user_api_key_dict = UserAPIKeyAuth( + api_key=api_key, + project_id="proj-io-only", + project_metadata={ + "model_itpm_limit": {"bedrock_mantle/claude-opus": 100000}, + "model_otpm_limit": {"bedrock_mantle/claude-opus": 100000}, + }, + ) + + data = { + "model": "bedrock_mantle/claude-opus", + "messages": [ + {"role": "user", "content": "hello there, this is a test message"} + ], + "max_tokens": 60, + } + + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data=data, + call_type="", + ) + + itpm_counter_key = handler.create_rate_limit_keys( + key="model_per_project_itpm", + value="proj-io-only:bedrock_mantle/claude-opus", + rate_limit_type="tokens", + ) + otpm_counter_key = handler.create_rate_limit_keys( + key="model_per_project_otpm", + 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 + ) + + await handler.async_post_call_failure_hook( + request_data=data, + original_exception=Exception("guardrail rejected"), + 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}" + ) + + +@pytest.mark.asyncio +async def test_otpm_rejection_does_not_double_refund_combined_tpm(rate_limiter): + """ + Regression for a High-severity review finding: when the project ITPM + reservation succeeds but OTPM is then over limit, + _reserve_project_io_tokens_or_raise rolls back the combined-TPM + reservation that already succeeded earlier in the same pre-call, then + raises. If it doesn't also mark that reservation released, + async_post_call_failure_hook -- which fires next in the real request + lifecycle, since raising from async_pre_call_hook triggers it -- sees + the same still-stashed reservation and refunds it a second time, + driving the combined TPM counter negative and letting a caller push + past the project's real TPM budget by repeatedly triggering OTPM + rejections. + """ + handler, cache = rate_limiter + + api_key = hash_token("sk-double-refund") + user_api_key_dict = UserAPIKeyAuth( + api_key=api_key, + project_id="proj-double-refund", + project_metadata={ + "model_tpm_limit": {"bedrock_mantle/claude-opus": 100000}, + "model_itpm_limit": {"bedrock_mantle/claude-opus": 100000}, + "model_otpm_limit": {"bedrock_mantle/claude-opus": 5}, + }, + ) + + data = { + "model": "bedrock_mantle/claude-opus", + "messages": [ + {"role": "user", "content": "hello there, this is a test message"} + ], + "max_tokens": 60, # blows past the 5-token OTPM limit + } + + tpm_counter_key = handler.create_rate_limit_keys( + key="model_per_project", + value="proj-double-refund:bedrock_mantle/claude-opus", + rate_limit_type="tokens", + ) + + with pytest.raises(Exception) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data=data, + call_type="", + ) + 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}" + ) + + # 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. + await handler.async_post_call_failure_hook( + request_data=data, + original_exception=exc_info.value, + 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 + ) + assert tpm_after_failure_hook == 0, ( + f"combined TPM counter went negative from a double refund: {tpm_after_failure_hook}" + ) + + +@pytest.mark.parametrize( + "embedding_input", + [ + list(range(51)), + [list(range(25)), list(range(26))], + ], +) +@pytest.mark.asyncio +async def test_project_itpm_rejects_pretokenized_embedding_input( + rate_limiter, + embedding_input, +): + handler, cache = rate_limiter + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-pretokenized-embedding-itpm"), + project_id="proj-pretokenized-embedding", + project_metadata={ + "model_itpm_limit": {"text-embedding-3-small": 50}, + }, + ) + data = { + "model": "text-embedding-3-small", + "input": embedding_input, + } + + with pytest.raises(Exception) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data=data, + call_type="aembedding", + ) + + assert getattr(exc_info.value, "status_code", None) == 429 + + +@pytest.mark.asyncio +async def test_responses_api_not_misclassified_as_embedding_for_output_estimate( + rate_limiter, +): + """ + Regression for a High-severity review finding: the Responses API also + puts its prompt in data["input"], the same field embeddings use, so the + output-token estimate treated every Responses call as an embedding and + reserved zero output tokens. call_type now disambiguates the two: the + same input-only payload gets zero output tokens for an embedding call + but a real floor for a Responses API call. + """ + handler, _cache = rate_limiter + + data = {"input": "describe this image in detail"} + + _, 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" + ) + assert responses_output_estimate > 0, ( + "Responses API call was misclassified as an embedding and reserved zero output tokens" + ) + + +@pytest.mark.parametrize( + ("data", "call_type", "expected_output_tokens"), + [ + ( + { + "messages": [{"role": "user", "content": "hello"}], + "max_tokens": 100, + "n": 10, + }, + "acompletion", + 1000, + ), + ( + { + "prompt": "hello", + "max_tokens": 100, + "n": 2, + "best_of": 5, + }, + "text_completion", + 500, + ), + ( + { + "messages": [{"role": "user", "content": "hello"}], + "max_tokens": 100, + "n": 0, + "best_of": "invalid", + }, + "acompletion", + 100, + ), + ], +) +def test_output_estimate_accounts_for_completion_candidates( + rate_limiter, + data, + call_type, + expected_output_tokens, +): + handler, _cache = rate_limiter + + _, estimated_output_tokens = handler._estimate_input_and_output_tokens( + data=data, + call_type=call_type, + ) + + assert estimated_output_tokens == expected_output_tokens + + +@pytest.mark.asyncio +async def test_responses_api_usage_reconciles_using_input_output_tokens_fields( + rate_limiter, +): + """ + Regression for the other half of the same finding: ResponseAPIUsage + exposes input_tokens/output_tokens, not prompt_tokens/completion_tokens. + Before this fix, _resolve_io_token_reconcile_usage couldn't resolve + Responses API usage at all, so the reservation was silently kept as-is + instead of being trued up to the much larger actual usage. + """ + handler, _cache = rate_limiter + + itpm_scope = ("model_per_project_itpm", "proj-responses:model") + otpm_scope = ("model_per_project_otpm", "proj-responses:model") + + stash = get_or_create_request_stash() + stash.itpm_reserved_tokens = 10 + stash.itpm_reserved_scopes = frozenset({itpm_scope}) + stash.otpm_reserved_tokens = 10 + stash.otpm_reserved_scopes = frozenset({otpm_scope}) + mock_kwargs = {} + + mock_response = ResponsesAPIResponse( + id="resp_test", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage(input_tokens=80, output_tokens=400, total_tokens=480), + ) + + increments = [] + + async def mock_increment(increment_list, **kwargs): + 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 + ) + + await handler.async_log_success_event( + kwargs=mock_kwargs, + response_obj=mock_response, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + itpm_adjustments = [i for i in increments if "model_per_project_itpm" in i["key"]] + otpm_adjustments = [i for i in increments if "model_per_project_otpm" in i["key"]] + + # delta = 80 actual input - 10 reserved = +70 + assert any(i["increment"] == 70 for i in itpm_adjustments), ( + f"ITPM reservation was never trued up to actual Responses API usage: {itpm_adjustments}" + ) + # delta = 400 actual output - 10 reserved = +390 + assert any(i["increment"] == 390 for i in otpm_adjustments), ( + f"OTPM reservation was never trued up to actual Responses API usage: {otpm_adjustments}" + ) + + +@pytest.mark.asyncio +async def test_itpm_reservation_accounts_for_audio_content_not_just_text(rate_limiter): + """ + Regression for the audio half of a Medium-severity review finding: + litellm.token_counter has no per-type handling for `input_audio` + content blocks (unlike images, which it does count via + use_default_image_token_count), so it silently contributes 0 tokens for + them. Without DEFAULT_AUDIO_TOKEN_ESTIMATE, a burst of audio-heavy + requests with minimal text would each reserve only the one-token floor + and blow past the project ITPM limit before post-call reconciliation. + """ + handler, cache = rate_limiter + + api_key = hash_token("sk-audio-itpm") + user_api_key_dict = UserAPIKeyAuth( + api_key=api_key, + project_id="proj-audio", + project_metadata={ + # Tighter than DEFAULT_AUDIO_TOKEN_ESTIMATE (300), but far bigger + # than the handful of tokens the bare text "hi" would cost. + "model_itpm_limit": {"bedrock_mantle/claude-opus": 50}, + }, + ) + + data = { + "model": "bedrock_mantle/claude-opus", + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "hi"}, + { + "type": "input_audio", + "input_audio": {"data": "base64-audio-bytes", "format": "wav"}, + }, + ], + } + ], + } + + with pytest.raises(Exception) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data=data, + call_type="", + ) + assert getattr(exc_info.value, "status_code", None) == 429, ( + "Expected the audio content to push the ITPM reservation over the " + "50-token limit; if this doesn't raise, audio content isn't being " + "counted again." + ) + + +def test_audio_token_estimate_scales_with_payload_size(): + """ + Regression for veria-ai Low finding: audio token reservation was flat + 300 per block regardless of duration. A short clip and a long clip both + reserved the same amount, letting a caller hide long audio in one block + 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 + DEFAULT_AUDIO_TOKEN_ESTIMATE so reference-only blocks and genuinely + short clips still get a non-trivial reservation. + + To exceed the floor the decoded payload must be > 300 * 1600 = 480 000 + bytes. We synthesise a fake b64-length string of 650 000 chars + (decoded ≈ 487 500 bytes → 304 tokens) to avoid actually allocating + and encoding ~480 kB of audio in every test run. + """ + large_b64 = "A" * 650_000 + very_large_b64 = "A" * 12_900_000 + small_b64 = "A" * 1_000 + + large_block = { + "type": "input_audio", + "input_audio": {"data": large_b64, "format": "wav"}, + } + small_block = { + "type": "input_audio", + "input_audio": {"data": small_b64, "format": "wav"}, + } + very_large_block = { + "type": "input_audio", + "input_audio": {"data": very_large_b64, "format": "wav"}, + } + 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 + ) + small_estimate = RateLimitHandler._estimate_audio_block_tokens(small_block) + no_data_estimate = RateLimitHandler._estimate_audio_block_tokens(no_data_block) + + assert large_estimate > small_estimate, ( + 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 > 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}" + ) + assert small_estimate >= 300, ( + f"Small payload must be floored at DEFAULT_AUDIO_TOKEN_ESTIMATE=300; got {small_estimate}" + ) + + +@pytest.mark.asyncio +async def test_itpm_rejects_large_audio_payload_that_would_pass_flat_estimate( + rate_limiter, +): + """ + Regression: a caller placing a long audio clip in one block previously + reserved only 300 tokens (the flat estimate). With the size-proportional + estimate, the same clip now reserves proportionally more and must trip + the ITPM limit when the limit is tuned to exactly expose the difference. + + 1 100 000 b64 chars → decoded ≈ 825 000 bytes → 825 000 // 1600 ≈ 515 + tokens > the 400-token limit. The flat estimate (300) would have passed. + """ + handler, cache = rate_limiter + + large_b64 = "A" * 1_100_000 + + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-large-audio"), + project_id="proj-large-audio", + project_metadata={ + "model_itpm_limit": {"bedrock_mantle/claude-opus": 400}, + }, + ) + + data = { + "model": "bedrock_mantle/claude-opus", + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "transcribe this"}, + { + "type": "input_audio", + "input_audio": {"data": large_b64, "format": "wav"}, + }, + ], + } + ], + } + + with pytest.raises(Exception) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data=data, + call_type="", + ) + assert getattr(exc_info.value, "status_code", None) == 429, ( + "Large audio payload must exceed the 400-token ITPM limit under the " + "size-proportional estimate; the old flat-rate estimate (300 tokens) " + "would have passed this limit silently" + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("call_type", "request_data"), + [ + ( + "acompletion", + { + "messages": [ + { + "role": "user", + "content": [ + { + "type": "image_url", + "image_url": { + "url": "https://example.com/high-resolution.png", + "detail": "high", + }, + } + ], + } + ] + }, + ), + ( + "aresponses", + { + "input": [ + { + "role": "user", + "content": [ + { + "type": "input_image", + "image_url": "https://example.com/high-resolution.png", + "detail": "high", + } + ], + } + ] + }, + ), + ], +) +async def test_image_content_reserves_full_project_itpm( + rate_limiter, + call_type, + request_data, +): + handler, cache = rate_limiter + model = "bedrock_mantle/claude-opus" + project_itpm_limit = 1_000 + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-high-resolution-image"), + project_id="project-high-resolution-image", + project_metadata={"model_itpm_limit": {model: project_itpm_limit}}, + ) + + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data={"model": model, **request_data}, + call_type=call_type, + ) + + stash = get_request_stash() + assert stash is not None + assert stash.itpm_reserved_tokens == project_itpm_limit + + +@pytest.mark.asyncio +async def test_itpm_otpm_reservation_is_kept_on_stream_disconnect(rate_limiter): + handler, cache = rate_limiter + + api_key = hash_token("sk-disconnect-test") + user_api_key_dict = UserAPIKeyAuth( + api_key=api_key, + project_id="proj-disconnect", + project_metadata={ + "model_itpm_limit": {"bedrock_mantle/claude-opus": 1000}, + "model_otpm_limit": {"bedrock_mantle/claude-opus": 500}, + }, + ) + + data: Dict[str, Any] = { + "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, ( + "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"]} + ) + + 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 + ) + + 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 + assert stash.reservation_released is False + + +@pytest.mark.asyncio +async def test_responses_api_otpm_output_cap_applied_not_skipped_as_embedding( + rate_limiter, +): + """ + Regression for a Greptile P1 finding: _reserve_project_io_tokens_or_raise + classified any request with data["input"] set as an embedding (no output + tokens), which also misclassifies the Responses API -- it puts its prompt + in "input" too, but does generate output. That skipped the output cap + applied whenever the configured OTPM limit is small enough to need it, + letting an unbounded Responses generation blow past OTPM before + post-call reconciliation catches up. + + The cap must land on data["max_output_tokens"], not data["max_tokens"]: + the Responses-to-chat-completion transformation only reads + max_output_tokens, so a max_tokens cap is silently dropped before + provider dispatch (a second Greptile finding on the same code path). + """ + handler, cache = rate_limiter + + api_key = hash_token("sk-responses-otpm-cap") + user_api_key_dict = UserAPIKeyAuth( + api_key=api_key, + project_id="proj-responses-otpm", + project_metadata={ + "model_otpm_limit": {"bedrock_mantle/claude-opus": 40}, + }, + ) + + data: Dict[str, Any] = { + "model": "bedrock_mantle/claude-opus", + "input": "describe this image in detail", + } + + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data=data, + call_type="aresponses", + ) + + assert data.get("max_output_tokens") is not None, ( + "Responses call was misclassified as an embedding and skipped the OTPM output cap" + ) + assert data["max_output_tokens"] == 16 + assert data.get("max_tokens") is None, ( + "OTPM output cap was written to max_tokens, which the Responses transformation ignores" + ) + + +@pytest.mark.asyncio +async def test_explicit_zero_output_responses_call_reserves_effective_provider_minimum( + rate_limiter, +): + handler, cache = rate_limiter + + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-responses-zero-output"), + project_id="proj-responses-zero-output", + project_metadata={ + "model_otpm_limit": {"bedrock_mantle/claude-opus": 5}, + }, + ) + + with pytest.raises(Exception) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data={ + "model": "bedrock_mantle/claude-opus", + "input": "describe this image in detail", + "max_output_tokens": 0, + }, + call_type="aresponses", + ) + + assert getattr(exc_info.value, "status_code", None) == 429 + + +@pytest.mark.asyncio +async def test_responses_api_combined_tpm_output_cap_applied_not_skipped_as_embedding( + rate_limiter, +): + """ + Regression for the same misclassification bug in the combined-TPM + output-cap block of async_pre_call_hook (a second, independent + `is_embedding = data.get("input") is not None` check). A project with + only a combined model_tpm_limit (no split itpm/otpm) configured small + enough to need the output cap must still apply it to a Responses call, + and must write it to max_output_tokens for the same reason as the OTPM + case above. + """ + handler, cache = rate_limiter + + api_key = hash_token("sk-responses-tpm-cap") + user_api_key_dict = UserAPIKeyAuth( + api_key=api_key, + project_id="proj-responses-tpm", + project_metadata={ + "model_tpm_limit": {"bedrock_mantle/claude-opus": 40}, + }, + ) + + data: Dict[str, Any] = { + "model": "bedrock_mantle/claude-opus", + "input": "describe this image in detail", + } + + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data=data, + call_type="aresponses", + ) + + assert data.get("max_output_tokens") is not None, ( + "Responses call was misclassified as an embedding and skipped the combined-TPM output cap" + ) + assert data["max_output_tokens"] == 16 + assert data.get("max_tokens") is None, ( + "combined-TPM output cap was written to max_tokens, which the Responses transformation ignores" + ) + + +@pytest.mark.asyncio +async def test_responses_api_multimodal_input_counts_image_content(rate_limiter): + """ + Regression for a Low-severity veria-ai finding: the Responses API's + `input` is commonly a list of message/content-block dicts, but + litellm.token_counter's `text` argument only joins plain string entries + in a list and silently drops everything else -- so an `input_image` + block contributed ~0 tokens to the ITPM estimate instead of the real + image token count. _estimate_precise_input_tokens now converts Responses + `input` to chat messages first (via the standard + transform_responses_api_input_to_messages helper) so image content is + counted the same way a chat completion's image content already is. + """ + handler, _cache = rate_limiter + + text_only_estimate = handler._estimate_precise_input_tokens( + data={"input": "hi"}, + model="bedrock_mantle/claude-opus", + call_type="aresponses", + ) + + multimodal_estimate = handler._estimate_precise_input_tokens( + data={ + "input": [ + { + "role": "user", + "content": [ + {"type": "input_text", "text": "hi"}, + { + "type": "input_image", + "image_url": "https://example.com/some-image.png", + }, + ], + } + ], + }, + model="bedrock_mantle/claude-opus", + call_type="aresponses", + ) + + assert multimodal_estimate > text_only_estimate + 100, ( + "Responses API input_image content block was not counted; got " + f"text_only={text_only_estimate}, multimodal={multimodal_estimate}" + ) + + +@pytest.mark.asyncio +async def test_refund_reserved_tokens_noop_when_amount_zero(rate_limiter): + """_refund_reserved_tokens returns immediately without calling Redis when amount=0.""" + handler, _cache = rate_limiter + + calls = [] + + async def mock_increment(pipeline_operations, **kwargs): + calls.extend(pipeline_operations) + + handler.async_increment_tokens_with_ttl_preservation = mock_increment + + await handler._refund_reserved_tokens( + scopes=[("api_key", "sk-test")], + amount=0, + ) + + assert not calls, "No Redis ops expected when amount is zero" + + +@pytest.mark.asyncio +async def test_reserve_io_tokens_noop_when_no_itpm_otpm_descriptors(rate_limiter): + """reserve_io_tokens returns OK immediately when no ITPM/OTPM descriptors present.""" + handler, _cache = rate_limiter + + non_io_descriptor = { + "key": "api_key", + "value": "sk-test", + "rate_limit": {"tokens_per_unit": 1000, "window_size": 60}, + } + response, itpm_reserved, otpm_reserved = await handler.reserve_io_tokens( + descriptors=[non_io_descriptor], + estimated_input_tokens=50, + estimated_output_tokens=50, + ) + + assert response["overall_code"] == "OK" + assert itpm_reserved == 0 + assert otpm_reserved == 0 + + +@pytest.mark.asyncio +async def test_reserve_io_tokens_itpm_only_no_otpm(rate_limiter): + """When only ITPM descriptors are present (no OTPM), returns itpm_reserved with otpm=0.""" + handler, cache = rate_limiter + + itpm_descriptor = { + "key": PROJECT_ITPM_DESCRIPTOR_KEY, + "value": "proj-a:model", + "rate_limit": {"tokens_per_unit": 10000, "window_size": 60}, + } + response, itpm_reserved, otpm_reserved = await handler.reserve_io_tokens( + descriptors=[itpm_descriptor], + estimated_input_tokens=100, + estimated_output_tokens=50, + ) + + assert response["overall_code"] == "OK" + assert itpm_reserved == 100 + assert otpm_reserved == 0 + + +def test_strip_audio_content_blocks_passthrough_non_list_messages(): + """Non-list input is returned unchanged (early return on line 2605).""" + result = RateLimitHandler._strip_audio_content_blocks("not a list") + assert result == "not a list" + + +def test_strip_audio_content_blocks_passthrough_non_dict_message(): + """Non-dict entries in the message list are appended unchanged.""" + messages = ["plain string message"] + result = RateLimitHandler._strip_audio_content_blocks(messages) + assert result == ["plain string message"] + + +def test_strip_audio_content_blocks_passthrough_non_list_content(): + """Messages with non-list content (e.g. plain string) pass through unchanged.""" + messages = [{"role": "user", "content": "hello"}] + result = RateLimitHandler._strip_audio_content_blocks(messages) + assert result == [{"role": "user", "content": "hello"}] + + +@pytest.mark.asyncio +async def test_otpm_rejection_releases_stashed_parallel_slot(rate_limiter): + """ + When OTPM is over limit and a parallel slot was already acquired, the + disconnect cleanup path in _reserve_project_io_tokens_or_raise must + release that slot. Exercises lines 2773-2777. + """ + handler, cache = rate_limiter + + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-otpm-slot"), + project_id="proj-slot", + project_metadata={"model_otpm_limit": {"m": 5}}, + ) + + data: Dict[str, Any] = { + "model": "m", + "messages": [{"role": "user", "content": "hello"}], + "max_tokens": 50, + } + + slot_released = [] + + async def mock_release(acquisition, parent_otel_span=None): + slot_released.append(acquisition) + + handler._release_parallel_request_slots = mock_release + + stash = get_or_create_request_stash() + stash.parallel_slot = { + "slot_id": "test-slot-id", + "counter_keys": ["some-key"], + } + + otpm_descriptor = { + "key": PROJECT_OTPM_DESCRIPTOR_KEY, + "value": "proj-slot:m", + "rate_limit": {"tokens_per_unit": 5, "window_size": 60}, + } + + with pytest.raises(Exception) as exc_info: + await handler._reserve_project_io_tokens_or_raise( + descriptors=[otpm_descriptor], + data=data, + requested_model="m", + user_api_key_dict=user_api_key_dict, + tpm_reservation_scopes=[], + tpm_reservation_amount=0, + ) + assert getattr(exc_info.value, "status_code", None) == 429 + assert slot_released, "Parallel slot must be released when OTPM rejects" + assert stash.parallel_slot is None + + +@pytest.mark.asyncio +async def test_itpm_only_status_stored_when_no_prior_rate_limit_response(rate_limiter): + """ + When only ITPM is configured (no combined TPM/RPM to pre-populate + the request stash), a successful ITPM reservation must store its status + there so post-call headers can read it. + """ + handler, cache = rate_limiter + + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-itpm-only-store"), + project_id="proj-store", + ) + + data: Dict[str, Any] = {"model": "m", "messages": []} + + itpm_descriptor = { + "key": PROJECT_ITPM_DESCRIPTOR_KEY, + "value": "proj-store:m", + "rate_limit": {"tokens_per_unit": 100000, "window_size": 60}, + } + + await handler._reserve_project_io_tokens_or_raise( + descriptors=[itpm_descriptor], + data=data, + requested_model="m", + user_api_key_dict=user_api_key_dict, + tpm_reservation_scopes=[], + tpm_reservation_amount=0, + ) + + 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.get("statuses"), "Stored response must contain statuses" + + +def test_resolve_io_token_usage_responses_api_with_cached_tokens(rate_limiter): + """ + ResponsesAPIResponse whose usage.input_tokens_details.cached_tokens is set + subtracts the cached portion from billable input. Covers line 3501. + """ + handler, _cache = rate_limiter + + response_obj = ResponsesAPIResponse( + id="resp_cached", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage( + input_tokens=100, + output_tokens=50, + total_tokens=150, + input_tokens_details=InputTokensDetails(cached_tokens=25), + ), + ) + 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}" + assert completion_tokens == 50 + + +def test_resolve_io_token_usage_dict_format(rate_limiter): + """ + Dict-shaped usage on a ModelResponse (older SDK versions or raw dicts in + the usage field) is parsed correctly. Covers lines 3502-3506. + """ + handler, _cache = rate_limiter + + response_obj = ModelResponse.model_construct( + usage={ + "prompt_tokens": 80, + "completion_tokens": 40, + "prompt_tokens_details": {"cached_tokens": 20}, + } + ) + 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}" + assert completion_tokens == 40 + + +def test_resolve_io_token_usage_unknown_type_returns_unresolved(rate_limiter): + """ + A ModelResponse whose usage attribute is not a Usage, ResponseAPIUsage, + or dict (e.g. a plain int) returns (0, 0, False) so the reservation is + kept rather than guessed. Covers lines 3507-3508. + """ + handler, _cache = rate_limiter + + response_obj = ModelResponse.model_construct(usage=42) + billable_input, completion_tokens, resolved = ( + handler._resolve_io_token_reconcile_usage(response_obj) + ) + + assert resolved is False + assert billable_input == 0 + assert completion_tokens == 0 + + +@pytest.mark.parametrize( + ("combined_usage", "expected_increments"), + [ + (None, ()), + ( + Usage(prompt_tokens=40, completion_tokens=15, total_tokens=55), + (-60, -45), + ), + ], +) +def test_zero_usage_keeps_reservations_unless_measured_fallback_exists( + rate_limiter, + combined_usage, + expected_increments, +): + handler, _cache = rate_limiter + itpm_scope = (PROJECT_ITPM_DESCRIPTOR_KEY, "project:model") + otpm_scope = (PROJECT_OTPM_DESCRIPTOR_KEY, "project:model") + stash = get_or_create_request_stash() + stash.itpm_reserved_tokens = 100 + stash.itpm_reserved_scopes = frozenset({itpm_scope}) + 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) + ) + + operations = handler._build_io_token_reservation_ops(kwargs, response_obj) + + assert tuple(operation["increment_value"] for operation in operations) == expected_increments + + +@pytest.mark.parametrize( + ("usage", "expected_increments"), + [ + ( + Usage(prompt_tokens=40, completion_tokens=15, total_tokens=55), + (40, 15), + ), + ( + Usage(prompt_tokens=0, completion_tokens=0, total_tokens=0), + (100, 60), + ), + ], +) +def test_retry_success_charges_released_project_io_reservations( + rate_limiter, + usage, + expected_increments, +): + handler, _cache = rate_limiter + itpm_scope = (PROJECT_ITPM_DESCRIPTOR_KEY, "project:model") + otpm_scope = (PROJECT_OTPM_DESCRIPTOR_KEY, "project:model") + stash = get_or_create_request_stash() + stash.itpm_reserved_tokens = 100 + stash.itpm_reserved_scopes = frozenset({itpm_scope}) + stash.otpm_reserved_tokens = 60 + stash.otpm_reserved_scopes = frozenset({otpm_scope}) + stash.reservation_released = True + + operations = handler._build_io_token_reservation_ops( + {}, + ModelResponse(usage=usage), + ) + + assert tuple(operation["increment_value"] for operation in operations) == expected_increments + + +@pytest.mark.asyncio +async def test_build_io_token_reservation_ops_skips_unresolvable_usage(rate_limiter): + """ + When response_obj has no parseable usage, _build_io_token_reservation_ops + returns [] to keep the reservation as-is rather than zeroing it out on a + bad guess. Covers line 3538. + """ + handler, _cache = rate_limiter + + itpm_scope = ("model_per_project_itpm", "proj-b:model") + stash = get_or_create_request_stash() + stash.itpm_reserved_tokens = 50 + stash.itpm_reserved_scopes = frozenset({itpm_scope}) + mock_kwargs = {} + + ops = handler._build_io_token_reservation_ops( + kwargs=mock_kwargs, + response_obj=object(), + ) + + assert not ops, f"Expected empty ops for unresolvable usage, got {ops}" + + +@pytest.mark.asyncio +async def test_post_call_failure_skips_rpm_only_descriptor_in_tpm_refund(rate_limiter): + """ + async_post_call_failure_hook skips descriptors without tokens_per_unit + (e.g. an RPM-only api_key scope) when building the combined-TPM refund ops, + so a key with rpm_limit but no tpm_limit doesn't receive a spurious refund + that would drive its counter negative. Covers the continue guard at line 4250. + """ + handler, cache = rate_limiter + + api_key = hash_token("sk-rpm-only-desc") + user_api_key_dict = UserAPIKeyAuth( + api_key=api_key, + rpm_limit=100, + project_id="proj-rpm-only-desc", + project_metadata={"model_tpm_limit": {"gpt-3.5-turbo": 100000}}, + ) + + data = { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "hello"}], + "max_tokens": 20, + } + + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data=data, + call_type="", + ) + + 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, + original_exception=Exception("rejected"), + 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 + ) + assert api_key_tokens_after >= 0, ( + f"RPM-only api_key scope must not receive a negative TPM refund; got {api_key_tokens_after}" + ) + + +@pytest.mark.asyncio +async def test_max_output_tokens_prevents_cap_injection(rate_limiter): + """ + Regression for veria-ai comment: when a Responses API request supplies + max_output_tokens (the canonical Responses output bound) but not max_tokens + or max_completion_tokens, the has_explicit_max_tokens check was False, so + the code injected data["max_tokens"] = capped_floor and silently truncated + the response. + + With the fix, max_output_tokens is included in the explicit-cap check and + data["max_tokens"] must NOT be injected when max_output_tokens is already + set. + """ + handler, cache = rate_limiter + + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-max-output-tokens"), + project_id="proj-responses-max-output", + project_metadata={ + "model_otpm_limit": {"mock-model": 100}, + }, + ) + + data: dict = { + "model": "mock-model", + "input": "Summarise the document", + "max_output_tokens": 80, + "litellm_call_id": "test-max-output-tokens", + "metadata": {}, + } + + try: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data=data, + call_type="responses", + ) + except Exception: + pass + + assert "max_tokens" not in data, ( + "data['max_tokens'] must not be injected when max_output_tokens is already " + "set; the cap injection was overriding the caller's explicit output bound" + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("call_type", "request_data", "cap_field", "reserved_tokens"), + [ + ("aresponses", {"input": "hello", "max_tokens": 1}, "max_output_tokens", 16), + ( + "acompletion", + { + "messages": [{"role": "user", "content": "hello"}], + "max_output_tokens": 1, + }, + "max_tokens", + 10, + ), + ], +) +async def test_output_reservation_ignores_cap_fields_from_other_endpoints( + rate_limiter, + call_type, + request_data, + cap_field, + reserved_tokens, +): + handler, cache = rate_limiter + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token(f"sk-{call_type}"), + project_id=f"project-{call_type}", + project_metadata={"model_otpm_limit": {"model": 40}}, + ) + data = {"model": "model", **request_data} + + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data=data, + call_type=call_type, + ) + + assert data[cap_field] == reserved_tokens + stash = get_request_stash() + assert stash is not None + assert stash.otpm_reserved_tokens == reserved_tokens + + +def test_responses_input_is_counted_even_when_messages_is_present(rate_limiter): + handler, _cache = rate_limiter + small_estimate = handler._estimate_precise_input_tokens( + data={"input": "short", "messages": [{"role": "user", "content": "ignored"}]}, + model="", + call_type="aresponses", + ) + large_estimate = handler._estimate_precise_input_tokens( + data={"input": "large input " * 500, "messages": []}, + model="", + call_type="aresponses", + ) + + assert large_estimate > small_estimate + + +def test_anthropic_messages_usage_reconciles_split_project_quota(rate_limiter): + handler, _cache = rate_limiter + + billable_input, output_tokens, resolved = handler._resolve_io_token_reconcile_usage( + { + "usage": { + "input_tokens": 100, + "output_tokens": 25, + "cache_read_input_tokens": 30, + } + } + ) + + assert resolved is True + assert billable_input == 70 + assert output_tokens == 25 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "request_data", + [ + {"input": "continue", "previous_response_id": "resp-123"}, + { + "input": [ + { + "role": "user", + "content": [{"type": "input_file", "file_id": "file-123"}], + } + ] + }, + ], +) +async def test_unmeasurable_responses_input_reserves_full_project_itpm( + rate_limiter, request_data +): + handler, cache = rate_limiter + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-unmeasurable-input"), + project_id="project-unmeasurable-input", + project_metadata={"model_itpm_limit": {"model": 100}}, + ) + data = {"model": "model", **request_data} + + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data=data, + call_type="aresponses", + ) + + stash = get_request_stash() + assert stash is not None + assert stash.itpm_reserved_tokens == 100 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "media_block", + [ + { + "type": "document", + "source": { + "type": "base64", + "media_type": "application/pdf", + "data": "dGVzdA==", + }, + }, + { + "type": "file", + "file": { + "filename": "document.pdf", + "file_data": "data:application/pdf;base64,dGVzdA==", + }, + }, + { + "type": "video_url", + "video_url": {"url": "https://example.com/video.mp4"}, + }, + ], +) +async def test_unmeasurable_chat_media_reserves_full_project_itpm( + rate_limiter, + media_block, +): + handler, cache = rate_limiter + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-unmeasurable-chat-media"), + project_id="project-unmeasurable-chat-media", + project_metadata={"model_itpm_limit": {"model": 100}}, + ) + + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data={ + "model": "model", + "messages": [{"role": "user", "content": [media_block]}], + }, + call_type="acompletion", + ) + + stash = get_request_stash() + assert stash is not None + assert stash.itpm_reserved_tokens == 100 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "call_type", + ["agenerate_content", "agenerate_content_stream"], +) +async def test_google_genai_native_contents_reserve_project_itpm( + rate_limiter, + call_type, +): + handler, cache = rate_limiter + model = "gemini/gemini-2.5-flash" + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-google-genai-native-itpm"), + project_id="project-google-genai-native-itpm", + project_metadata={"model_itpm_limit": {model: 10_000}}, + ) + + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data={ + "model": model, + "contents": [ + { + "role": "user", + "parts": [{"text": "Gemini quota input " * 200}], + } + ], + }, + call_type=call_type, + ) + + stash = get_request_stash() + assert stash is not None + assert stash.itpm_reserved_tokens > 100 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("call_type", ["rerank", "arerank"]) +async def test_rerank_query_and_documents_enforce_project_itpm( + rate_limiter, + monkeypatch, + call_type, +): + handler, cache = rate_limiter + captured = {} + + def token_counter(**kwargs): + captured.update(kwargs) + return 101 + + monkeypatch.setattr("litellm.token_counter", token_counter) + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token(f"sk-{call_type}-itpm"), + project_id=f"project-{call_type}-itpm", + project_metadata={"model_itpm_limit": {"rerank-model": 100}}, + ) + + with pytest.raises(Exception) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data={ + "model": "rerank-model", + "query": "Which document is most relevant?", + "documents": ["first document", {"text": "second document"}], + }, + call_type=call_type, + ) + + assert getattr(exc_info.value, "status_code", None) == 429 + 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( + rate_limiter, + monkeypatch, +): + handler, _cache = rate_limiter + data = { + "query": "query text", + "documents": ["first document", "second document"], + } + + def token_counter(**_kwargs): + raise ValueError("tokenizer unavailable") + + 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 + + +@pytest.mark.parametrize( + ("response_obj", "expected"), + [ + ( + RerankResponse( + meta={"tokens": {"input_tokens": 42, "output_tokens": 3}} + ), + (42, 3, True), + ), + ( + RerankResponse( + meta={ + "tokens": {"input_tokens": 0, "output_tokens": 0}, + "billed_units": {"total_tokens": 57}, + } + ), + (57, 0, True), + ), + ( + RerankResponse( + meta={ + "tokens": {"input_tokens": 0, "output_tokens": 0}, + "billed_units": {"total_tokens": 0}, + } + ), + (0, 0, False), + ), + ], +) +def test_rerank_usage_reconciles_project_split_token_quota( + rate_limiter, + response_obj, + expected, +): + handler, _cache = rate_limiter + + assert handler._resolve_io_token_reconcile_usage(response_obj) == expected + + +def test_split_quota_helpers_handle_non_mapping_inputs(rate_limiter): + handler, _cache = rate_limiter + + assert _call_id_from_callback_kwargs(object()) is None + assert handler._is_embedding_request(object(), None) is False + assert handler._get_explicit_output_cap(object(), None) is None + assert handler._get_output_candidate_count(object()) == 1 + 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()) == () + + +@pytest.mark.parametrize( + ("call_type", "data"), + [ + ( + "text_completion", + { + "messages": [{"role": "user", "content": "ignored"}], + "prompt": "abcd", + "input": "ignored", + "max_tokens": 1, + }, + ), + (None, {"prompt": "abcd", "max_tokens": 1}), + (None, {"prompt": ["abcd", "efgh"], "max_tokens": 1}), + ], +) +def test_split_token_estimate_selects_endpoint_input(rate_limiter, call_type, data): + handler, _cache = rate_limiter + + estimated_input, estimated_output = handler._estimate_input_and_output_tokens( + data=data, + call_type=call_type, + ) + + assert estimated_input > 0 + assert estimated_output == 1 + + +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._contains_unmeasurable_chat_media(object()) is False + assert handler._contains_image_content(object()) is False + assert handler._contains_image_content( + {"inline_data": {"mime_type": "image/png", "data": "dGVzdA=="}} + ) + assert handler._responses_input_to_chat_messages(object()) == () + assert ( + handler._requires_conservative_responses_input_reservation( + object(), "responses" + ) + is False + ) + assert handler._estimate_precise_input_tokens(object(), model=None) == 0 + + +@pytest.mark.parametrize( + ("call_type", "data", "expected_text"), + [ + ("embedding", {"input": "embedding input"}, "embedding input"), + ( + "embedding", + {"input": ["first embedding", "second embedding"]}, + ["first embedding", "second embedding"], + ), + ("text_completion", {"prompt": "completion prompt"}, "completion prompt"), + ], +) +def test_precise_input_estimate_selects_endpoint_text( + rate_limiter, + monkeypatch, + call_type, + data, + expected_text, +): + handler, _cache = rate_limiter + captured = {} + + def token_counter(**kwargs): + captured.update(kwargs) + return 7 + + monkeypatch.setattr("litellm.token_counter", token_counter) + + assert ( + handler._estimate_precise_input_tokens(data, model="test", call_type=call_type) + == 7 + ) + assert captured["messages"] is None + assert captured["text"] == expected_text + + +@pytest.mark.asyncio +async def test_project_io_reservation_ignores_non_mapping_request_data(rate_limiter): + handler, _cache = rate_limiter + + await handler._reserve_project_io_tokens_or_raise( + descriptors=[], + data=object(), + requested_model=None, + user_api_key_dict=UserAPIKeyAuth(), + tpm_reservation_scopes=(), + tpm_reservation_amount=0, + ) + + +@pytest.mark.asyncio +async def test_streaming_combined_usage_reconciles_project_io_reservations( + rate_limiter, +): + handler, _cache = rate_limiter + itpm_scope = (PROJECT_ITPM_DESCRIPTOR_KEY, "project:model") + otpm_scope = (PROJECT_OTPM_DESCRIPTOR_KEY, "project:model") + stash = get_or_create_request_stash() + stash.itpm_reserved_tokens = 100 + stash.itpm_reserved_scopes = frozenset({itpm_scope}) + stash.otpm_reserved_tokens = 60 + stash.otpm_reserved_scopes = frozenset({otpm_scope}) + kwargs = { + "combined_usage_object": Usage( + prompt_tokens=40, + completion_tokens=15, + total_tokens=55, + ), + } + increments = [] + + async def capture_increments(increment_list, **_kwargs): + increments.extend(increment_list) + + handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = ( + capture_increments + ) + + await handler.async_log_success_event( + kwargs=kwargs, + response_obj={"response": "stream body"}, + start_time=datetime.now(), + 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"] + ] + assert [operation["increment_value"] for operation in itpm_adjustments] == [-60] + assert [operation["increment_value"] for operation in otpm_adjustments] == [-45] + + +def test_aggregate_only_combined_usage_keeps_project_io_reservations(rate_limiter): + 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")} + ) + kwargs = { + "combined_usage_object": Usage(total_tokens=55), + } + + assert handler._build_io_token_reservation_ops(kwargs, object()) == () + + +def test_raw_split_usage_dict_reconciles_project_io_tokens(rate_limiter): + handler, _cache = rate_limiter + + assert handler._resolve_io_token_reconcile_usage( + { + "input_tokens": 30, + "output_tokens": 12, + "input_tokens_details": {"cached_tokens": 5}, + } + ) == (25, 12, True) + + +@pytest.mark.asyncio +async def test_post_call_success_hook_contains_header_merge_failures( + rate_limiter, monkeypatch +): + handler, _cache = rate_limiter + response = ModelResponse() + response._hidden_params = {} + + def raise_on_merge(**_kwargs): + raise RuntimeError("header merge failed") + + monkeypatch.setattr( + handler, + "_merge_ratelimit_statuses_into_additional_headers", + raise_on_merge, + ) + + await handler.async_post_call_success_hook( + data={ + "litellm_proxy_rate_limit_response": { + "overall_code": "OK", + "statuses": (), + } + }, + user_api_key_dict=UserAPIKeyAuth(), + response=response, + ) + + if __name__ == "__main__": pytest.main([__file__, "-v", "-s"]) diff --git a/tests/test_litellm/proxy/test_proxy_types.py b/tests/test_litellm/proxy/test_proxy_types.py index 5354de182a0..4a93e9ac7ba 100644 --- a/tests/test_litellm/proxy/test_proxy_types.py +++ b/tests/test_litellm/proxy/test_proxy_types.py @@ -159,3 +159,21 @@ def test_update_key_request_requires_key_or_key_alias(): by_alias = UpdateKeyRequest(key_alias="my-alias") assert by_alias.key is None assert by_alias.key_alias == "my-alias" + + +@pytest.mark.parametrize("request_type", ["new", "update"]) +def test_project_io_token_limits_are_stored_in_metadata(request_type): + from litellm.proxy._types import NewProjectRequest, UpdateProjectRequest + + limits = { + "model_itpm_limit": {"bedrock_mantle/openai.gpt-oss-120b": 20_000_000}, + "model_otpm_limit": {"bedrock_mantle/openai.gpt-oss-120b": 4_000_000}, + } + request = ( + NewProjectRequest(team_id="team-1", **limits) + if request_type == "new" + else UpdateProjectRequest(project_id="project-1", **limits) + ) + + assert request.metadata == limits + assert request.model_dump(exclude_none=True)["metadata"] == limits diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index df85decc676..fd3eb6b16df 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -28772,10 +28772,18 @@ export interface components { metadata?: { [key: string]: unknown; } | null; + /** Model Itpm Limit */ + model_itpm_limit?: { + [key: string]: number; + } | null; /** Model Max Budget */ model_max_budget?: { [key: string]: unknown; } | null; + /** Model Otpm Limit */ + model_otpm_limit?: { + [key: string]: number; + } | null; /** Model Rpm Limit */ model_rpm_limit?: { [key: string]: unknown; @@ -33623,10 +33631,18 @@ export interface components { metadata?: { [key: string]: unknown; } | null; + /** Model Itpm Limit */ + model_itpm_limit?: { + [key: string]: number; + } | null; /** Model Max Budget */ model_max_budget?: { [key: string]: unknown; } | null; + /** Model Otpm Limit */ + model_otpm_limit?: { + [key: string]: number; + } | null; /** Model Rpm Limit */ model_rpm_limit?: { [key: string]: unknown;