From 2c4fed10351660bce488fafa8a4ed017ef0fd0cd Mon Sep 17 00:00:00 2001 From: epistoteles Date: Thu, 21 May 2026 21:19:05 +0200 Subject: [PATCH 01/83] Add missing Databricks model pricing data Add cost data for 14 Databricks models that were missing from the model_prices_and_context_window.json file: - databricks-gpt-5-4, gpt-5-2, gpt-5-4-nano, gpt-5-4-mini - databricks-gpt-5-2-codex, gpt-5-3-codex - databricks-gpt-5-1-codex-mini, gpt-5-1-codex-max - databricks-gemini-3-pro, gemini-3-flash, gemini-3-1-pro, gemini-3-1-flash-lite - databricks-claude-sonnet-4-6, claude-opus-4-6 Pricing sourced from: https://www.databricks.com/product/pricing/proprietary-foundation-model-serving --- model_prices_and_context_window.json | 227 +++++++++++++++++++++++++++ 1 file changed, 227 insertions(+) diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 31a5993a240..6ff71018f18 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -11243,6 +11243,26 @@ "supports_minimal_reasoning_effort": true, "supports_tool_choice": true }, + "databricks/databricks-claude-opus-4-6": { + "input_cost_per_token": 5.00003e-06, + "input_dbu_cost_per_token": 7.1429e-05, + "litellm_provider": "databricks", + "max_input_tokens": 200000, + "max_output_tokens": 64000, + "max_tokens": 64000, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 2.5000010000000002e-05, + "output_dbu_cost_per_token": 0.000357143, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving", + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_reasoning": true, + "supports_minimal_reasoning_effort": true, + "supports_tool_choice": true + }, "databricks/databricks-claude-sonnet-4": { "input_cost_per_token": 2.9999900000000002e-06, "input_dbu_cost_per_token": 4.2857e-05, @@ -11300,6 +11320,25 @@ "supports_reasoning": true, "supports_tool_choice": true }, + "databricks/databricks-claude-sonnet-4-6": { + "input_cost_per_token": 2.9999900000000002e-06, + "input_dbu_cost_per_token": 4.2857e-05, + "litellm_provider": "databricks", + "max_input_tokens": 200000, + "max_output_tokens": 64000, + "max_tokens": 64000, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 1.5000020000000002e-05, + "output_dbu_cost_per_token": 0.000214286, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving", + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_reasoning": true, + "supports_tool_choice": true + }, "databricks/databricks-gemini-2-5-flash": { "input_cost_per_token": 3.0001999999999996e-07, "input_dbu_cost_per_token": 4.285999999999999e-06, @@ -11334,6 +11373,74 @@ "supports_function_calling": true, "supports_tool_choice": true }, + "databricks/databricks-gemini-3-1-flash-lite": { + "input_cost_per_token": 3.1248e-07, + "input_dbu_cost_per_token": 4.464e-06, + "litellm_provider": "databricks", + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_tokens": 65536, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 1.87502e-06, + "output_dbu_cost_per_token": 2.6786e-05, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving", + "supports_function_calling": true, + "supports_tool_choice": true + }, + "databricks/databricks-gemini-3-1-pro": { + "input_cost_per_token": 2.49998e-06, + "input_dbu_cost_per_token": 3.5714e-05, + "litellm_provider": "databricks", + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_tokens": 65536, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 1.5000020000000002e-05, + "output_dbu_cost_per_token": 0.000214286, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving", + "supports_function_calling": true, + "supports_tool_choice": true + }, + "databricks/databricks-gemini-3-flash": { + "input_cost_per_token": 6.2503e-07, + "input_dbu_cost_per_token": 8.929e-06, + "litellm_provider": "databricks", + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_tokens": 65536, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 3.74997e-06, + "output_dbu_cost_per_token": 5.3571e-05, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving", + "supports_function_calling": true, + "supports_tool_choice": true + }, + "databricks/databricks-gemini-3-pro": { + "input_cost_per_token": 2.49998e-06, + "input_dbu_cost_per_token": 3.5714e-05, + "litellm_provider": "databricks", + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_tokens": 65536, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 1.5000020000000002e-05, + "output_dbu_cost_per_token": 0.000214286, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving", + "supports_function_calling": true, + "supports_tool_choice": true + }, "databricks/databricks-gemma-3-12b": { "input_cost_per_token": 1.5000999999999998e-07, "input_dbu_cost_per_token": 2.1429999999999996e-06, @@ -11379,6 +11486,126 @@ "output_dbu_cost_per_token": 0.000142857, "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving" }, + "databricks/databricks-gpt-5-1-codex-max": { + "input_cost_per_token": 1.24999e-06, + "input_dbu_cost_per_token": 1.7857e-05, + "litellm_provider": "databricks", + "max_input_tokens": 272000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 9.999990000000002e-06, + "output_dbu_cost_per_token": 0.000142857, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving" + }, + "databricks/databricks-gpt-5-1-codex-mini": { + "input_cost_per_token": 2.4997e-07, + "input_dbu_cost_per_token": 3.571e-06, + "litellm_provider": "databricks", + "max_input_tokens": 272000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 1.99997e-06, + "output_dbu_cost_per_token": 2.8571e-05, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving" + }, + "databricks/databricks-gpt-5-2": { + "input_cost_per_token": 1.75e-06, + "input_dbu_cost_per_token": 2.5e-05, + "litellm_provider": "databricks", + "max_input_tokens": 272000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 1.4e-05, + "output_dbu_cost_per_token": 0.0002, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving" + }, + "databricks/databricks-gpt-5-2-codex": { + "input_cost_per_token": 1.75e-06, + "input_dbu_cost_per_token": 2.5e-05, + "litellm_provider": "databricks", + "max_input_tokens": 272000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 1.4e-05, + "output_dbu_cost_per_token": 0.0002, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving" + }, + "databricks/databricks-gpt-5-3-codex": { + "input_cost_per_token": 1.75e-06, + "input_dbu_cost_per_token": 2.5e-05, + "litellm_provider": "databricks", + "max_input_tokens": 272000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 1.4e-05, + "output_dbu_cost_per_token": 0.0002, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving" + }, + "databricks/databricks-gpt-5-4": { + "input_cost_per_token": 2.49998e-06, + "input_dbu_cost_per_token": 3.5714e-05, + "litellm_provider": "databricks", + "max_input_tokens": 272000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 1.5000020000000002e-05, + "output_dbu_cost_per_token": 0.000214286, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving" + }, + "databricks/databricks-gpt-5-4-mini": { + "input_cost_per_token": 7.4998e-07, + "input_dbu_cost_per_token": 1.0714e-05, + "litellm_provider": "databricks", + "max_input_tokens": 272000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 4.50002e-06, + "output_dbu_cost_per_token": 6.4286e-05, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving" + }, + "databricks/databricks-gpt-5-4-nano": { + "input_cost_per_token": 1.9999e-07, + "input_dbu_cost_per_token": 2.857e-06, + "litellm_provider": "databricks", + "max_input_tokens": 272000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 1.24999e-06, + "output_dbu_cost_per_token": 1.7857e-05, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving" + }, "databricks/databricks-gpt-5-mini": { "input_cost_per_token": 2.4997000000000006e-07, "input_dbu_cost_per_token": 3.571e-06, From c62e1238d648607f9544a98bcba4b4ec8ad5e369 Mon Sep 17 00:00:00 2001 From: Chenlu Ji Date: Mon, 6 Jul 2026 22:31:48 -0700 Subject: [PATCH 02/83] feat(tinyfish): surface response headers + top-level response extras Follow-up to #31411 (superseded and merged as #31997). Two related fixes so LiteLLM callers see what TinyFish actually returns, plus small correctness cleanups. ## Response headers surfaced on _hidden_params TinyFish sets useful response headers (x-request-id on every response, retry-after and x-ratelimit-limit on 429s). Previously these were only accessible via BaseLLMException.headers on error paths; on the success path they were dropped entirely. Fix: stash headers on both LiteLLM-conventional channels, matching the pattern used by Gemini / Volcengine / Manus / ChatGPT / OpenAI-responses providers. - `_hidden_params["headers"]` -- raw dict from httpx, all keys lowercased. - `_hidden_params["additional_headers"]` -- passed through process_response_headers, which prefixes any x-litellm-* provider header with `llm_provider-` so downstream LiteLLM code that trusts bare x-litellm-* markers can't be spoofed (values still survive under the prefixed key for observability). ## Top-level response extras (query, total_results, page, future fields) transform_search_response was building a fresh SearchResponse from just `results`, silently dropping every top-level field TinyFish's response carries beyond `results` / `object`. Fix: mutate parsed.results to its truncated slice and return the same SearchResponse instance rather than reconstructing. Every field pydantic populated during model_validate -- declared attributes AND extras (query, total_results, page, parameter_warnings, and any future TinyFish additions) -- survives regardless of which storage bucket holds it. Robust against upstream schema evolution: if LiteLLM later promotes a field from extras to declared, this code needs no change. ## Code cleanup - List-valued custom params JSON-encoded on the wire (matching the existing dict handling), so callers can pass a natural Python list for JSON-array wire params. - URL-encodable-params adapter accepts float in addition to str / int / bool; server-side rejection of a wrong-typed float now surfaces cleanly with `TinyFish Search:` attribution + docs link. - Assorted comment / docstring / test-fixture hygiene (no logic changes). ## Tests 70 unit + integration tests pass locally. Live-tested against production TinyFish with 6 diverse queries (basic / max_results / country=US / language=ja / domain filter / fetch={"format":"html"}) -- all 6 pass every expected-behavior check. --- .../llms/tinyfish/search/transformation.py | 74 +++++--- tests/search_tests/test_tinyfish_search.py | 61 ++++++- .../llms/tinyfish/test_tinyfish_search.py | 165 +++++++++++++++--- 3 files changed, 250 insertions(+), 50 deletions(-) diff --git a/litellm/llms/tinyfish/search/transformation.py b/litellm/llms/tinyfish/search/transformation.py index cef5f9cd02e..aea0dfe8b8e 100644 --- a/litellm/llms/tinyfish/search/transformation.py +++ b/litellm/llms/tinyfish/search/transformation.py @@ -14,6 +14,7 @@ import httpx from pydantic import TypeAdapter, ValidationError from litellm._logging import verbose_logger +from litellm.litellm_core_utils.core_helpers import process_response_headers from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.base_llm.search.transformation import ( @@ -22,13 +23,13 @@ from litellm.llms.base_llm.search.transformation import ( ) from litellm.secret_managers.main import get_secret_str -_UrlEncodableParams = TypeAdapter(dict[str, str | int | bool]) +_UrlEncodableParams = TypeAdapter(dict[str, str | int | float | bool]) _StrList = TypeAdapter(list[str]) _StrFrozenSet = TypeAdapter(frozenset[str]) _TINYFISH_PARAMS_KEY = "_tinyfish_params" _TINYFISH_DOCS_URL = "https://docs.tinyfish.ai/search-api" -_TINYFISH_RESULT_CAP = 10 # TinyFish's natural per-page SERP ceiling +_TINYFISH_RESULT_CAP = 10 # Client-side truncation cap for max_results class TinyfishSearchConfig(BaseSearchConfig): @@ -94,16 +95,16 @@ class TinyfishSearchConfig(BaseSearchConfig): TinyFish equivalents: - ``query`` (str or list[str]) → ``query`` (list joined by spaces) - ``country`` → ``location`` - - ``search_domain_filter`` (list[str]) → folded into the query as - ``() (site:a OR site:b ...)`` (TinyFish has no first-class - field today; see ML-2084 for the planned ``include_domains``) + - ``search_domain_filter`` (list[str]) → folded into the query using + search operators - ``max_results`` → not sent on the wire; stashed on ``self._caller_max_results`` for client-side response truncation (TinyFish doesn't honor it server-side) - ``max_tokens_per_page`` → silently dropped (no TinyFish equivalent) Any other ``optional_params`` keys are forwarded to TinyFish as-is. - dict/list values are JSON-encoded so they survive ``urlencode``. + dict and list values are JSON-encoded so structured payloads survive + ``urlencode``. Returns: ``{_TINYFISH_PARAMS_KEY: }``. @@ -144,14 +145,12 @@ class TinyfishSearchConfig(BaseSearchConfig): supported_perplexity = _StrFrozenSet.validate_python(raw_supported) for param, value in optional_params.items(): if param not in supported_perplexity and param not in request_data: - # `fetch` expects a JSON-encoded object on the wire; accept the - # natural Python dict form and serialize here so callers don't - # have to pre-stringify. - if isinstance(value, dict): + # Serialize dicts/lists as JSON so structured params survive urlencode. + if isinstance(value, (dict, list)): value = json.dumps(value, separators=(",", ":")) # `urlencode` would render Python bool as "True"/"False" - # (capitalized). ux-labs validators require lowercase - # "true"/"false" (e.g. `include_thumbnail`); normalize here. + # (capitalized). TinyFish Search's bool params require lowercase + # "true"/"false" strings on the wire; normalize here. elif isinstance(value, bool): value = "true" if value else "false" request_data[param] = value @@ -167,17 +166,35 @@ class TinyfishSearchConfig(BaseSearchConfig): """ Transform a TinyFish response to LiteLLM's unified ``SearchResponse``. - Mappings (per-result): - - ``title`` → ``SearchResult.title`` (defaults to ``""`` if missing/null) - - ``url`` → ``SearchResult.url`` (defaults to ``""``) - - ``snippet`` → ``SearchResult.snippet`` (defaults to ``""``) - - all other per-result fields (``position``, ``site_name``, - ``thumbnail_url``, ``fetch``, ``fetch_error``, ...) ride through as - extras on ``SearchResult`` via its ``extra="allow"`` config. + Per-result field handling: + - ``title``, ``url``, ``snippet`` are declared on ``SearchResult`` and + populated by ``SearchResponse.model_validate`` when present. Missing + or ``None`` values are defaulted to ``""`` beforehand by + ``_default_missing_result_fields`` so a degraded result flows through + instead of failing the whole call. + - All undeclared per-result fields (``position``, ``site_name``, and + any others TinyFish returns) ride through as extras via + ``SearchResult``'s ``extra="allow"`` config — accessible as + attributes on the result object or enumerable via + ``result.model_extra``. - Top-level ``parameter_warnings`` (see ML-2085) is read when present and - each entry is re-fired via ``verbose_logger.warning``. Absent or - malformed entries are silently skipped — never throws. + Top-level ``parameter_warnings`` is read when present and each entry + is re-fired via ``verbose_logger.warning``. Absent or malformed + entries are silently skipped — never throws. + + Top-level extras (``query``, ``total_results``, ``page``, and any + future TinyFish additions) ride through via + ``SearchResponse.extra="allow"``. The validated response is returned + in place after truncating ``results`` to the caller's ``max_results``, + so every field pydantic populated survives regardless of which + storage bucket (declared attribute or ``__pydantic_extra__``) holds it. + + TinyFish response headers (e.g. ``x-request-id``, ``retry-after``, + ``x-ratelimit-limit`` — httpx normalizes header names to lowercase) + are stashed on ``response._hidden_params["headers"]`` (raw) and + ``response._hidden_params["additional_headers"]`` (sanitized via + ``process_response_headers``) so callers can correlate a search with + server-side logs. Error paths routed through ``self._wrap_error`` for uniform ``"TinyFish Search: . See for details."`` wrapping: @@ -223,7 +240,12 @@ class TinyfishSearchConfig(BaseSearchConfig): _emit_parameter_warnings(parsed) max_results = self._caller_max_results or _TINYFISH_RESULT_CAP - return SearchResponse(results=list(parsed.results[:max_results])) + # Truncate in place so all pydantic-populated fields survive — declared and extras. + parsed.results = list(parsed.results[:max_results]) + raw_headers = dict(raw_response.headers) + parsed._hidden_params["headers"] = raw_headers + parsed._hidden_params["additional_headers"] = process_response_headers(raw_headers) + return parsed def _wrap_error( self, @@ -243,9 +265,9 @@ class TinyfishSearchConfig(BaseSearchConfig): carry the ``TinyFish Search:`` prefix — the bare error already names the host in the URL, so attribution is implicit there. """ - # ux-labs frontend wraps every error body as {"error": {"code", "message", "details"?}}. + # TinyFish Search wraps every error body as {"error": {"code", "message", "details"?}}. # Best-effort unwrap to surface the inner message; fall back to the raw body - # for non-ux-labs responses (CDN HTML pages, other JSON envelopes, plain text). + # for other envelope shapes (CDN HTML pages, other JSON envelopes, plain text). inner_message = error_message try: body: object = json.loads(error_message) # any-ok: json.loads -> Any @@ -290,7 +312,7 @@ def _default_missing_result_fields(raw_json: object) -> None: def _emit_parameter_warnings(parsed: SearchResponse) -> None: - """Re-fire TinyFish-side ``parameter_warnings`` (see ML-2085) as warnings. + """Re-fire TinyFish-side ``parameter_warnings`` as warnings. Defensive: skip silently on any shape we don't recognize so a malformed entry (or an early/partial rollout of the field) never throws. diff --git a/tests/search_tests/test_tinyfish_search.py b/tests/search_tests/test_tinyfish_search.py index aca28544513..becb8287a29 100644 --- a/tests/search_tests/test_tinyfish_search.py +++ b/tests/search_tests/test_tinyfish_search.py @@ -35,11 +35,16 @@ MOCK_TINYFISH_RESPONSE = { def _make_mock_response( - json_data: dict, status_code: int = 200, request_url: str | None = None + json_data: dict, + status_code: int = 200, + request_url: str | None = None, + headers: dict | None = None, ) -> MagicMock: mock = MagicMock() mock.status_code = status_code mock.json.return_value = json_data + # httpx.Headers normalizes keys to lowercase — mirror production behavior. + mock.headers = httpx.Headers(headers or {}) if request_url: mock.request = MagicMock() mock.request.url = httpx.URL(request_url) @@ -163,7 +168,7 @@ class TestTinyfishSearch: @pytest.mark.asyncio async def test_fetch_param_round_trip(self): - # End-to-end check: caller passes `fetch=...` (JSON-encoded tf-fetch + # End-to-end check: caller passes `fetch=...` (JSON-encoded fetch # config); param reaches TinyFish on the request side and the nested # `fetch` object on each result surfaces back to the SearchResult on the # response side. No LiteLLM-side support code is required. @@ -235,6 +240,58 @@ class TestTinyfishSearch: assert result.results[0].title == "Result 0" assert result.results[2].title == "Result 2" + @pytest.mark.asyncio + async def test_top_level_extras_surface_end_to_end(self): + # Envelope extras (`query`, `total_results`, `page`) must survive the + # full asearch dispatch — proves LiteLLM's entry-point plumbing outside + # our transformer doesn't accidentally strip them. + os.environ["TINYFISH_API_KEY"] = "sk-tinyfish-test" + + mock_response = _make_mock_response(MOCK_TINYFISH_RESPONSE) + + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get", + new_callable=AsyncMock, + ) as mock_get: + mock_get.return_value = mock_response + + response = await litellm.asearch( + query="web automation tools", + search_provider="tinyfish", + ) + + assert getattr(response, "query", None) == "web automation tools" + assert getattr(response, "total_results", None) == 2 + assert getattr(response, "page", None) == 0 + + @pytest.mark.asyncio + async def test_response_headers_surface_end_to_end(self): + # Response headers must land on `_hidden_params` after the full + # asearch dispatch (both raw and sanitized channels). + os.environ["TINYFISH_API_KEY"] = "sk-tinyfish-test" + + mock_response = _make_mock_response( + MOCK_TINYFISH_RESPONSE, + headers={"X-Request-ID": "req-e2e-1"}, + ) + + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get", + new_callable=AsyncMock, + ) as mock_get: + mock_get.return_value = mock_response + + response = await litellm.asearch( + query="test", + search_provider="tinyfish", + ) + + raw = response._hidden_params["headers"] + add = response._hidden_params["additional_headers"] + # httpx lowercases; both channels agree on the value. + assert raw["x-request-id"] == "req-e2e-1" + assert add["llm_provider-x-request-id"] == "req-e2e-1" + @pytest.mark.asyncio async def test_empty_results(self): os.environ["TINYFISH_API_KEY"] = "sk-tinyfish-test" diff --git a/tests/test_litellm/llms/tinyfish/test_tinyfish_search.py b/tests/test_litellm/llms/tinyfish/test_tinyfish_search.py index 58363e3baea..2dcccb8ea7e 100644 --- a/tests/test_litellm/llms/tinyfish/test_tinyfish_search.py +++ b/tests/test_litellm/llms/tinyfish/test_tinyfish_search.py @@ -47,7 +47,9 @@ def _make_mock_response( mock = MagicMock() mock.status_code = status_code - mock.headers = headers or {} + # httpx.Headers normalizes keys to lowercase — mirror production so tests + # assert what callers actually see. + mock.headers = httpx.Headers(headers or {}) if json_data is not None: mock.json.return_value = json_data mock.text = text if text is not None else _json.dumps(json_data) @@ -222,7 +224,7 @@ class TestTransformSearchRequest: assert param not in result["_tinyfish_params"] def test_arbitrary_param_passed_through(self): - # `fetch` is a TinyFish-specific param (JSON-encoded tf-fetch config). + # `fetch` is a TinyFish-specific param (JSON-encoded fetch config). # The passthrough loop should forward it verbatim without LiteLLM needing # to know about it. config = TinyfishSearchConfig() @@ -237,26 +239,49 @@ class TestTransformSearchRequest: config = TinyfishSearchConfig() result = config.transform_search_request( query="test", - optional_params={"fetch": {"format": "html", "fetch_path": "fast"}}, - ) - assert ( - result["_tinyfish_params"]["fetch"] - == '{"format":"html","fetch_path":"fast"}' + optional_params={"fetch": {"format": "html"}}, ) + assert result["_tinyfish_params"]["fetch"] == '{"format":"html"}' def test_bool_param_serialized_as_lowercase(self): - # urlencode renders Python bool as capitalized "True"/"False"; ux-labs - # rejects those (e.g. include_thumbnail must be literal "true"/"false"). - # Normalize before passing through. + # urlencode renders Python bool as capitalized "True"/"False"; TinyFish + # Search's bool params require lowercase "true"/"false" strings on the + # wire. Normalize before passing through. config = TinyfishSearchConfig() true_result = config.transform_search_request( - query="test", optional_params={"include_thumbnail": True} + query="test", optional_params={"some_bool_param": True} ) false_result = config.transform_search_request( - query="test", optional_params={"include_thumbnail": False} + query="test", optional_params={"some_bool_param": False} ) - assert true_result["_tinyfish_params"]["include_thumbnail"] == "true" - assert false_result["_tinyfish_params"]["include_thumbnail"] == "false" + assert true_result["_tinyfish_params"]["some_bool_param"] == "true" + assert false_result["_tinyfish_params"]["some_bool_param"] == "false" + + def test_float_param_passes_through(self): + # Float values pass the urlencode adapter and land on the wire as + # their decimal string form. If TinyFish's server rejects a float + # for a param it expects as int, the server's 400 response is + # attributed via _wrap_error (`TinyFish Search: ...`) — better than + # a client-side pydantic ValidationError with no context. + config = TinyfishSearchConfig() + result = config.transform_search_request( + query="test", + optional_params={"some_float_param": 0.5}, + ) + assert result["_tinyfish_params"]["some_float_param"] == 0.5 + + def test_list_param_auto_json_encoded(self): + # TinyFish Search's JSON-array params arrive on the wire as JSON- + # encoded strings. Accept the natural Python list form and serialize + # so the caller doesn't have to pre-stringify. Params whose wire + # format is a plain comma-separated string are the caller's + # responsibility to pass as a Python str. + config = TinyfishSearchConfig() + result = config.transform_search_request( + query="test", + optional_params={"some_list_param": ["a.example", "b.example"]}, + ) + assert result["_tinyfish_params"]["some_list_param"] == '["a.example","b.example"]' def test_pre_stringified_param_passed_unchanged(self): # If the caller already JSON-encoded, don't re-encode. @@ -422,10 +447,104 @@ class TestTransformSearchResponse: assert getattr(first, "position", None) == 1 assert getattr(first, "site_name", None) == "tinyfish.ai" + def test_top_level_extras_flow_through(self): + # TinyFish returns `query`, `total_results`, `page` at the envelope + # level. These must ride through to the caller via SearchResponse's + # extra="allow" so pagination logic, echo checks, etc. work. + config = TinyfishSearchConfig() + mock_response = _make_mock_response(MOCK_TINYFISH_RESPONSE) + result = config.transform_search_response( + raw_response=mock_response, logging_obj=None + ) + assert getattr(result, "query", None) == "web automation tools" + assert getattr(result, "total_results", None) == 2 + assert getattr(result, "page", None) == 0 + + def test_top_level_future_extras_flow_through(self): + # Any future TinyFish top-level field must ride through unchanged + # (design contract: no LiteLLM code change needed for new fields). + config = TinyfishSearchConfig() + body = { + "results": [ + {"title": "x", "url": "https://x", "snippet": "x"}, + ], + "query": "test", + "example_int_extra": 123, # hypothetical future field + "example_str_extra": "value", # hypothetical future field + "example_id_extra": "abc-def", # hypothetical future field + } + result = config.transform_search_response( + raw_response=_make_mock_response(body), logging_obj=None + ) + assert getattr(result, "example_int_extra", None) == 123 + assert getattr(result, "example_str_extra", None) == "value" + assert getattr(result, "example_id_extra", None) == "abc-def" + + def test_response_headers_stashed_on_hidden_params(self): + # TinyFish Search sets X-Request-ID on every success response. Confirm it + # lands on both `_hidden_params["headers"]` (raw) and + # `_hidden_params["additional_headers"]` (sanitized/prefixed). + # httpx.Headers lowercases every key, so assertions use lowercase. + config = TinyfishSearchConfig() + mock_response = _make_mock_response( + MOCK_TINYFISH_RESPONSE, + headers={"X-Request-ID": "req-abc-123", "Content-Type": "application/json"}, + ) + result = config.transform_search_response( + raw_response=mock_response, logging_obj=None + ) + # Raw copy — httpx has normalized keys to lowercase. + assert result._hidden_params["headers"]["x-request-id"] == "req-abc-123" + # process_response_headers prefixes non-OpenAI-standard keys with "llm_provider-". + assert result._hidden_params["additional_headers"]["llm_provider-x-request-id"] == "req-abc-123" + + def test_response_headers_future_headers_flow_through(self): + # "Accept extra": any header TinyFish Search adds later must ride + # through without a LiteLLM code change. + config = TinyfishSearchConfig() + mock_response = _make_mock_response( + MOCK_TINYFISH_RESPONSE, + headers={ + "X-Request-ID": "req-1", + "X-Example-Header-A": "value-a", # hypothetical future header + "X-Example-Header-B": "value-b", # hypothetical future header + }, + ) + result = config.transform_search_response( + raw_response=mock_response, logging_obj=None + ) + raw = result._hidden_params["headers"] + # httpx lowercases header names on read. + assert raw["x-example-header-a"] == "value-a" + assert raw["x-example-header-b"] == "value-b" + + def test_response_headers_strips_x_litellm_spoof(self): + # A provider setting `x-litellm-*` in its response must not be able to + # spoof LiteLLM-internal markers via _hidden_params["additional_headers"]. + # The raw copy preserves the header (opt-in debug view); the sanitized + # copy prefixes it with `llm_provider-` so bare `x-litellm-*` markers + # can't be spoofed (values still survive under the prefixed key for + # observability). + config = TinyfishSearchConfig() + mock_response = _make_mock_response( + MOCK_TINYFISH_RESPONSE, + headers={"x-litellm-attempted-fallbacks": "spoofed", "X-Request-ID": "r1"}, + ) + result = config.transform_search_response( + raw_response=mock_response, logging_obj=None + ) + # Raw view still has the spoof. + assert result._hidden_params["headers"]["x-litellm-attempted-fallbacks"] == "spoofed" + # Sanitized view: the spoof survives only under the llm_provider- prefix + # (never under the bare x-litellm-* key that LiteLLM downstream trusts). + additional = result._hidden_params["additional_headers"] + assert "x-litellm-attempted-fallbacks" not in additional + assert additional.get("llm_provider-x-litellm-attempted-fallbacks") == "spoofed" + def test_fetch_field_rides_through_to_search_result(self): - # Mirrors browser-search's per-result `fetch` nested object (see - # api/src/parser.rs SearchResult.fetch). Confirms `fetch=...` requests - # surface their content to LiteLLM callers without provider changes. + # Mirrors TinyFish Search's per-result `fetch` nested object. + # Confirms `fetch=...` requests surface their content to LiteLLM + # callers without provider changes. config = TinyfishSearchConfig() fetched = { "results": [ @@ -568,7 +687,7 @@ class TestTransformSearchResponse: class TestErrorHandling: def test_4xx_response_raises_with_attribution_and_unwrapped_message(self): - # Reproduces ux-labs' error envelope shape for an INVALID_INPUT response. + # Reproduces TinyFish Search's error envelope shape for an INVALID_INPUT response. config = TinyfishSearchConfig() body = { "error": { @@ -590,7 +709,7 @@ class TestErrorHandling: def test_429_preserves_status_code_and_headers(self): config = TinyfishSearchConfig() - body = {"error": {"code": "RATE_LIMIT_EXCEEDED", "message": "60 rpm"}} + body = {"error": {"code": "RATE_LIMIT_EXCEEDED", "message": "rate limit exceeded"}} mock_response = _make_mock_response( body, status_code=429, headers={"Retry-After": "60"} ) @@ -600,10 +719,12 @@ class TestErrorHandling: ) assert getattr(exc_info.value, "status_code", None) == 429 headers = getattr(exc_info.value, "headers", {}) or {} - assert headers.get("Retry-After") == "60" + # httpx lowercases; the exception carries the same dict shape. + assert headers.get("retry-after") == "60" - def test_5xx_with_non_ux_labs_body_falls_back_to_raw_text(self): - # Cloudflare-style JSON or any other envelope: unwrap fails, fall back to raw. + def test_5xx_with_non_tinyfish_envelope_shape_falls_back_to_raw_text(self): + # A JSON body that doesn't match TinyFish Search's error envelope shape: + # unwrap fails, fall back to the raw body text. config = TinyfishSearchConfig() body = {"errors": [{"code": "10000", "message": "Internal"}]} mock_response = _make_mock_response(body, status_code=502) From 3b843708b097d6c4ff96d374900282ddf5142f0d Mon Sep 17 00:00:00 2001 From: shivam Date: Sat, 18 Jul 2026 20:17:57 +0000 Subject: [PATCH 03/83] fix(bedrock): degrade gracefully on malformed tool-call arguments split_concatenated_json_objects re-raised JSONDecodeError on genuinely malformed (non-concatenated) tool-call arguments, which propagated out of _convert_to_bedrock_tool_call_invoke and turned every replayed Bedrock conversation into a 500. Catch the decode error, keep whatever complete objects parsed, log a warning, and let the caller fall back to input={} so the conversation continues. Fixes #18667 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../prompt_templates/common_utils.py | 27 +++++++++--- ...ore_utils_prompt_templates_common_utils.py | 30 +++++++++++-- ...llm_core_utils_prompt_templates_factory.py | 43 +++++++++++++++++++ 3 files changed, 89 insertions(+), 11 deletions(-) diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index 538d5f650ef..db3856ce86b 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -1679,16 +1679,19 @@ def split_concatenated_json_objects(raw: str) -> List[Dict[str, Any]]: This helper uses ``json.JSONDecoder.raw_decode()`` to walk the string and extract each JSON object individually. + The walk degrades gracefully: if the string is malformed or truncated + (e.g. a stream that ended mid-tool-call), whatever complete objects were + parsed before the bad tail are returned and the remainder is discarded + with a warning, rather than raising. The sole caller + (``_convert_to_bedrock_tool_call_invoke``) treats an empty result as + ``input={}`` so the conversation can continue instead of hard-failing. + Returns ------- list[dict] A list of parsed dicts – one per JSON object found. If *raw* is - empty or whitespace-only, an empty list is returned. - - Raises - ------ - json.JSONDecodeError - If the string contains text that cannot be parsed as JSON at all. + empty, whitespace-only, or wholly unparseable, an empty list is + returned. """ import json @@ -1708,7 +1711,17 @@ def split_concatenated_json_objects(raw: str) -> List[Dict[str, Any]]: if idx >= length: break - obj, end_idx = decoder.raw_decode(raw, idx) + try: + obj, end_idx = decoder.raw_decode(raw, idx) + except json.JSONDecodeError as e: + verbose_logger.warning( + "split_concatenated_json_objects: discarding unparseable tool-call " + "arguments tail after %d complete object(s); error=%s at char %d", + len(results), + e, + idx, + ) + break if isinstance(obj, dict): results.append(obj) else: diff --git a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py index 1b1db634ed2..6d14d283b84 100644 --- a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py +++ b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py @@ -251,10 +251,32 @@ def test_split_concatenated_json_non_dict_value(): assert result == [{}] -def test_split_concatenated_json_invalid_raises(): - """Completely invalid JSON raises JSONDecodeError.""" - with pytest.raises(json.JSONDecodeError): - split_concatenated_json_objects("not json at all") +def test_split_concatenated_json_wholly_invalid_returns_empty(): + """ + Wholly unparseable JSON degrades to an empty list instead of raising. + + Regression for https://github.com/BerriAI/litellm/issues/18667: a raise + here propagated out of `_convert_to_bedrock_tool_call_invoke` and turned + every replayed conversation into a 500. + """ + assert split_concatenated_json_objects("not json at all") == [] + + +def test_split_concatenated_json_malformed_object_returns_empty(): + """ + A single malformed object (missing comma between keys) degrades to an + empty list rather than raising `Expecting ',' delimiter`. + """ + assert split_concatenated_json_objects('{"location": "Boston" "unit": "celsius"}') == [] + + +def test_split_concatenated_json_salvages_prefix_before_truncated_tail(): + """ + Complete objects parsed before an unparseable/truncated tail are kept; + only the bad tail is discarded. + """ + result = split_concatenated_json_objects('{"a": 1}{"b": 2}{"c":') + assert result == [{"a": 1}, {"b": 2}] # --------------------------------------------------------------------------- diff --git a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index bcda88ea609..17f680df4d9 100644 --- a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -2286,6 +2286,49 @@ def test_bedrock_tool_call_invoke_non_dict_arguments(): assert result[0]["toolUse"]["input"] == {} +def test_bedrock_tool_call_invoke_malformed_json_does_not_raise(): + """ + Regression for https://github.com/BerriAI/litellm/issues/18667. + + When the model emits malformed JSON in tool-call arguments (here a + missing comma between keys), replaying that history must NOT raise + `Unable to convert openai tool calls ... Expecting ',' delimiter`. + It degrades to an empty-object input so the conversation can continue. + """ + tool_calls = [ + { + "id": "toolu_abc123", + "type": "function", + "function": { + "name": "get_weather", + "arguments": '{"location": "Boston" "unit": "celsius"}', + }, + } + ] + result = _convert_to_bedrock_tool_call_invoke(tool_calls) + assert len(result) == 1 + assert result[0]["toolUse"]["toolUseId"] == "toolu_abc123" + assert result[0]["toolUse"]["name"] == "get_weather" + assert result[0]["toolUse"]["input"] == {} + + +def test_bedrock_tool_call_invoke_salvages_valid_prefix_before_truncated_tail(): + """ + A valid leading object followed by a truncated tail keeps the valid + object rather than dropping everything or raising. + """ + tool_calls = [ + { + "id": "call_partial", + "type": "function", + "function": {"name": "shell", "arguments": '{"cmd": "ls"}{"cmd":'}, + } + ] + result = _convert_to_bedrock_tool_call_invoke(tool_calls) + assert len(result) == 1 + assert result[0]["toolUse"]["input"] == {"cmd": "ls"} + + def test_make_valid_bedrock_tool_name_preserves_hyphens(): assert make_valid_bedrock_tool_name("my-tool") == "my-tool" assert ( From c2d8a4e4263e45bee96e94f7091071140bf79d83 Mon Sep 17 00:00:00 2001 From: shivam Date: Sat, 18 Jul 2026 20:28:35 +0000 Subject: [PATCH 04/83] chore(bedrock): clarify tool-call decode warning to avoid double char reference Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/litellm_core_utils/prompt_templates/common_utils.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index db3856ce86b..9dcbca954b3 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -1716,10 +1716,10 @@ def split_concatenated_json_objects(raw: str) -> List[Dict[str, Any]]: except json.JSONDecodeError as e: verbose_logger.warning( "split_concatenated_json_objects: discarding unparseable tool-call " - "arguments tail after %d complete object(s); error=%s at char %d", + "arguments tail after %d complete object(s); decode_start=%d error=%s", len(results), - e, idx, + e, ) break if isinstance(obj, dict): From 9dbe61aa6d8915f40a030ee11f522345afb9cc6a Mon Sep 17 00:00:00 2001 From: Shivi Jain Date: Fri, 31 Jul 2026 16:49:07 +0530 Subject: [PATCH 05/83] 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; From 292161f766ca7ac88cf3cdbef0c4b599a5576a88 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 7 Aug 2026 23:25:10 -0700 Subject: [PATCH 06/83] fix(proxy): read through to the DB on registry misses so just-created models, guardrails, and agents resolve on sibling replicas --- .../proxy/agent_endpoints/a2a_endpoints.py | 15 +- litellm/proxy/agent_endpoints/a2a_routing.py | 6 +- .../common_utils/registry_read_through.py | 143 +++++++++++ .../proxy/guardrails/guardrail_endpoints.py | 6 +- ...model_access_group_management_endpoints.py | 30 ++- litellm/proxy/route_llm_request.py | 232 ++++++++++-------- ruff.toml | 2 +- .../test_registry_read_through.py | 231 +++++++++++++++++ .../test_access_group_management.py | 96 ++++++++ .../proxy/test_route_a2a_models.py | 75 ++++++ .../proxy/test_route_llm_request.py | 121 +++++++++ 11 files changed, 843 insertions(+), 114 deletions(-) create mode 100644 litellm/proxy/common_utils/registry_read_through.py create mode 100644 tests/test_litellm/proxy/common_utils/test_registry_read_through.py diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py index 27780aeb994..a4e1ac126d9 100644 --- a/litellm/proxy/agent_endpoints/a2a_endpoints.py +++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py @@ -152,14 +152,13 @@ def _jsonrpc_error( ) -def _get_agent(agent_id: str): +async def _get_agent(agent_id: str) -> "AgentResponse | None": """Look up an agent by ID or name. Returns None if not found.""" - from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry + from litellm.proxy.common_utils.registry_read_through import ( + get_agent_with_read_through, + ) - agent = global_agent_registry.get_agent_by_id(agent_id=agent_id) - if agent is None: - agent = global_agent_registry.get_agent_by_name(agent_name=agent_id) - return agent + return await get_agent_with_read_through(agent_id) def _enforce_inbound_trace_id(agent: Any, request: Request) -> None: @@ -531,7 +530,7 @@ async def get_agent_card( ) try: - agent: Final = _get_agent(agent_id) + agent: Final = await _get_agent(agent_id) if agent is None: raise HTTPException(status_code=404, detail=f"Agent '{agent_id}' not found") @@ -645,7 +644,7 @@ async def invoke_agent_a2a( params.pop(key) # Find the agent - agent: Final = _get_agent(agent_id) + agent: Final = await _get_agent(agent_id) if agent is None: return _jsonrpc_error(request_id, -32000, f"Agent '{agent_id}' not found", 404) diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py index 038b6b4a840..2228735d805 100644 --- a/litellm/proxy/agent_endpoints/a2a_routing.py +++ b/litellm/proxy/agent_endpoints/a2a_routing.py @@ -25,10 +25,12 @@ async def route_a2a_agent_request( Returns None if not an A2A request (allows normal routing to continue). """ # Import here to avoid circular imports - from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry from litellm.proxy.agent_endpoints.auth.agent_permission_handler import ( AgentRequestHandler, ) + from litellm.proxy.common_utils.registry_read_through import ( + get_agent_with_read_through, + ) from litellm.proxy.route_llm_request import ( ROUTE_ENDPOINT_MAPPING, ProxyModelNotFoundError, @@ -44,7 +46,7 @@ async def route_a2a_agent_request( agent_name: Final = model_name[4:] # Look up agent in registry - agent: Final = global_agent_registry.get_agent_by_name(agent_name) + agent: Final = await get_agent_with_read_through(agent_name) if agent is None: verbose_proxy_logger.error("[A2A] Agent '%s' not found in registry", agent_name) route_name = ROUTE_ENDPOINT_MAPPING.get(route_type, route_type) diff --git a/litellm/proxy/common_utils/registry_read_through.py b/litellm/proxy/common_utils/registry_read_through.py new file mode 100644 index 00000000000..b78106205d4 --- /dev/null +++ b/litellm/proxy/common_utils/registry_read_through.py @@ -0,0 +1,143 @@ +"""Read-through recovery for in-memory registries in multi-replica deployments. + +A management write (POST /model/new, /guardrails, /v1/agents) lands on one +replica and reaches Postgres, but sibling replicas only refresh their in-memory +registries on the periodic config reload or the Redis config-sync resync, both +of which lag by seconds. A request that uses the new object immediately can +land on a sibling that has never heard of it and fail with a 400/404. + +On a registry miss, callers here fetch the missing object from the DB and load +it into the local registry before giving up. A short negative-result TTL keeps +repeated lookups of genuinely unknown names from hammering the DB. +""" + +import asyncio +from collections.abc import Awaitable, Callable +from typing import TYPE_CHECKING, Final + +from litellm._logging import verbose_proxy_logger +from litellm.caching.in_memory_cache import InMemoryCache + +if TYPE_CHECKING: + from litellm.integrations.custom_guardrail import CustomGuardrail + from litellm.types.agents import AgentResponse + +READ_THROUGH_MISS_TTL_SECONDS: Final = 2.0 + + +class RegistryReadThrough: + __slots__ = ("_lock", "_miss_ttl_seconds", "_recent_misses", "_resync") + + def __init__( + self, + resync: Callable[[str], Awaitable[bool]], + miss_ttl_seconds: float = READ_THROUGH_MISS_TTL_SECONDS, + ) -> None: + self._resync = resync + self._miss_ttl_seconds = miss_ttl_seconds + self._lock = asyncio.Lock() + self._recent_misses = InMemoryCache(max_size_in_memory=1000) + + async def attempt(self, key: str) -> bool: + if self._recent_misses.get_cache(key) is not None: + return False + async with self._lock: + if self._recent_misses.get_cache(key) is not None: + return False + try: + found: Final = await self._resync(key) + except Exception as e: # noqa: BLE001 # a failed read-through must surface the original miss error, not a 500 + verbose_proxy_logger.warning("registry read-through for %r failed: %s", key, e) + return False + if not found: + self._recent_misses.set_cache(key, True, ttl=self._miss_ttl_seconds) + return found + + +def _db_backed_registries_enabled() -> bool: + from litellm.proxy import proxy_server + + return proxy_server.prisma_client is not None and proxy_server.store_model_in_db is True + + +async def _resync_model_deployments(model_name: str) -> bool: + from litellm.proxy import proxy_server + from litellm.repositories.model_repository import ModelRepository + + if not _db_backed_registries_enabled(): + return False + prisma_client: Final = proxy_server.prisma_client + assert prisma_client is not None + rows: Final = await ModelRepository(prisma_client).table.find_many( + where={"OR": [{"model_name": model_name}, {"model_id": model_name}]} + ) + if not rows: + return False + if proxy_server.llm_router is None: + await proxy_server.proxy_config.add_deployment( + prisma_client=prisma_client, proxy_logging_obj=proxy_server.proxy_logging_obj + ) + return proxy_server.llm_router is not None + proxy_server.proxy_config._add_deployment(db_models=rows) + proxy_server.llm_model_list = proxy_server.llm_router.get_model_list() + return True + + +async def _resync_guardrails(guardrail_name: str) -> bool: + from litellm.proxy import proxy_server + + if not _db_backed_registries_enabled(): + return False + prisma_client: Final = proxy_server.prisma_client + assert prisma_client is not None + await proxy_server.proxy_config._init_guardrails_in_db(prisma_client=prisma_client) + return _initialized_guardrail(guardrail_name) is not None + + +async def _resync_agents(agent_id_or_name: str) -> bool: + from litellm.proxy import proxy_server + + if not _db_backed_registries_enabled(): + return False + prisma_client: Final = proxy_server.prisma_client + assert prisma_client is not None + await proxy_server.proxy_config._init_agents_in_db(prisma_client=prisma_client) + return _agent_from_registry(agent_id_or_name) is not None + + +model_registry_read_through: Final = RegistryReadThrough(resync=_resync_model_deployments) +guardrail_registry_read_through: Final = RegistryReadThrough(resync=_resync_guardrails) +agent_registry_read_through: Final = RegistryReadThrough(resync=_resync_agents) + + +def _agent_from_registry(agent_id_or_name: str) -> "AgentResponse | None": + from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry + + by_id: Final = global_agent_registry.get_agent_by_id(agent_id=agent_id_or_name) + if by_id is not None: + return by_id + return global_agent_registry.get_agent_by_name(agent_name=agent_id_or_name) + + +async def get_agent_with_read_through(agent_id_or_name: str) -> "AgentResponse | None": + agent: Final = _agent_from_registry(agent_id_or_name) + if agent is not None: + return agent + if not await agent_registry_read_through.attempt(agent_id_or_name): + return None + return _agent_from_registry(agent_id_or_name) + + +def _initialized_guardrail(guardrail_name: str) -> "CustomGuardrail | None": + from litellm.proxy.guardrails import guardrail_endpoints + + return guardrail_endpoints.GUARDRAIL_REGISTRY.get_initialized_guardrail_callback(guardrail_name=guardrail_name) + + +async def get_initialized_guardrail_with_read_through(guardrail_name: str) -> "CustomGuardrail | None": + active: Final = _initialized_guardrail(guardrail_name) + if active is not None: + return active + if not await guardrail_registry_read_through.attempt(guardrail_name): + return None + return _initialized_guardrail(guardrail_name) diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index 761d8aabc8a..dff70ccf68d 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -2244,8 +2244,12 @@ async def apply_guardrail( litellm_logging_obj = None start_time: Final = datetime.now(timezone.utc) + from litellm.proxy.common_utils.registry_read_through import ( + get_initialized_guardrail_with_read_through, + ) + try: - active_guardrail: Final[CustomGuardrail | None] = GUARDRAIL_REGISTRY.get_initialized_guardrail_callback( + active_guardrail: Final[CustomGuardrail | None] = await get_initialized_guardrail_with_read_through( guardrail_name=request.guardrail_name ) if active_guardrail is None: diff --git a/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py b/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py index 7051f705a03..75d33c6c40a 100644 --- a/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py @@ -7,7 +7,10 @@ Endpoints here: import json from collections.abc import Mapping, Sequence -from typing import Any, Final +from typing import TYPE_CHECKING, Any, Final + +if TYPE_CHECKING: + from litellm.router import Router from fastapi import APIRouter, Depends, HTTPException @@ -52,6 +55,23 @@ def validate_models_exist(model_names: list[str], llm_router) -> tuple[bool, lis return (len(missing) == 0, missing) +async def _missing_models_after_read_through( + model_names: Sequence[str], llm_router: "Router | None" +) -> tuple[str, ...]: + from litellm.proxy import proxy_server + from litellm.proxy.common_utils.registry_read_through import ( + model_registry_read_through, + ) + + _, missing = validate_models_exist(model_names=list(model_names), llm_router=llm_router) + if not missing: + return () + for name in missing: + await model_registry_read_through.attempt(name) + _, still_missing = validate_models_exist(model_names=list(model_names), llm_router=proxy_server.llm_router) + return tuple(still_missing) + + def add_access_group_to_deployment(model_info: dict[str, Any], access_group: str) -> tuple[dict[str, Any], bool]: """ Add an access group to a deployment's model_info. @@ -369,12 +389,12 @@ async def create_model_group( # Validate model_names exist in router (only if using model_names path) if not use_model_ids and has_model_names: assert data.model_names is not None - all_valid, missing_models = validate_models_exist( + missing_models: Final = await _missing_models_after_read_through( model_names=data.model_names, llm_router=llm_router, ) - if not all_valid: + if missing_models: raise HTTPException( status_code=400, detail={"error": f"Model(s) not found: {', '.join(missing_models)}"}, @@ -633,12 +653,12 @@ async def update_access_group( # Validation: Check if all new models exist (only if using model_names path) if not use_model_ids and has_model_names: assert data.model_names is not None - all_valid, missing_models = validate_models_exist( + missing_models: Final = await _missing_models_after_read_through( model_names=data.model_names, llm_router=llm_router, ) - if not all_valid: + if missing_models: raise HTTPException( status_code=400, detail={"error": f"Model(s) not found: {', '.join(missing_models)}"}, diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index dd8deed57f1..657e0cbcafa 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -313,112 +313,150 @@ async def add_shared_session_to_data(data: dict) -> None: pass +RouteType = Literal[ + "acompletion", + "atext_completion", + "aembedding", + "aimage_generation", + "aspeech", + "atranscription", + "amoderation", + "arerank", + "aresponses", + "aget_responses", + "adelete_responses", + "acancel_responses", + "acompact_responses", + "acreate_response_reply", + "alist_input_items", + "_arealtime", # private function for realtime API + "acreate_realtime_client_secret", + "arealtime_calls", + "acreate_realtime_transcription_session", + "_aresponses_websocket", # private function for responses WebSocket mode + "aimage_edit", + "agenerate_content", + "agenerate_content_stream", + "allm_passthrough_route", + "acreate_batch", + "aretrieve_batch", + "alist_batches", + "afile_content", + "afile_retrieve", + "acreate_fine_tuning_job", + "acancel_fine_tuning_job", + "alist_fine_tuning_jobs", + "aretrieve_fine_tuning_job", + "avector_store_search", + "avector_store_create", + "avector_store_retrieve", + "avector_store_list", + "avector_store_update", + "avector_store_delete", + "avector_store_file_create", + "avector_store_file_list", + "avector_store_file_retrieve", + "avector_store_file_content", + "avector_store_file_update", + "avector_store_file_delete", + "aocr", + "asearch", + "avideo_generation", + "avideo_list", + "avideo_status", + "avideo_content", + "avideo_remix", + "avideo_create_character", + "avideo_get_character", + "avideo_edit", + "avideo_extension", + "acreate_container", + "alist_containers", + "aretrieve_container", + "adelete_container", + "aupload_container_file", + "alist_container_files", + "aretrieve_container_file", + "adelete_container_file", + "aretrieve_container_file_content", + "acreate_skill", + "alist_skills", + "aget_skill", + "adelete_skill", + "aingest", + "anthropic_messages", + "acreate_interaction", + "aget_interaction", + "adelete_interaction", + "acancel_interaction", + "acreate_agent", + "alist_agents", + "aget_agent", + "adelete_agent", + "alist_agent_versions", + "asend_message", + "call_mcp_tool", + "acancel_batch", + "afile_delete", + "acreate_eval", + "alist_evals", + "aget_eval", + "aupdate_eval", + "adelete_eval", + "acancel_eval", + "acreate_run", + "alist_runs", + "aget_run", + "acancel_run", + "adelete_run", +] + + async def route_request( data: dict, llm_router: LitellmRouter | None, user_model: str | None, - route_type: Literal[ - "acompletion", - "atext_completion", - "aembedding", - "aimage_generation", - "aspeech", - "atranscription", - "amoderation", - "arerank", - "aresponses", - "aget_responses", - "adelete_responses", - "acancel_responses", - "acompact_responses", - "acreate_response_reply", - "alist_input_items", - "_arealtime", # private function for realtime API - "acreate_realtime_client_secret", - "arealtime_calls", - "acreate_realtime_transcription_session", - "_aresponses_websocket", # private function for responses WebSocket mode - "aimage_edit", - "agenerate_content", - "agenerate_content_stream", - "allm_passthrough_route", - "acreate_batch", - "aretrieve_batch", - "alist_batches", - "afile_content", - "afile_retrieve", - "acreate_fine_tuning_job", - "acancel_fine_tuning_job", - "alist_fine_tuning_jobs", - "aretrieve_fine_tuning_job", - "avector_store_search", - "avector_store_create", - "avector_store_retrieve", - "avector_store_list", - "avector_store_update", - "avector_store_delete", - "avector_store_file_create", - "avector_store_file_list", - "avector_store_file_retrieve", - "avector_store_file_content", - "avector_store_file_update", - "avector_store_file_delete", - "aocr", - "asearch", - "avideo_generation", - "avideo_list", - "avideo_status", - "avideo_content", - "avideo_remix", - "avideo_create_character", - "avideo_get_character", - "avideo_edit", - "avideo_extension", - "acreate_container", - "alist_containers", - "aretrieve_container", - "adelete_container", - "aupload_container_file", - "alist_container_files", - "aretrieve_container_file", - "adelete_container_file", - "aretrieve_container_file_content", - "acreate_skill", - "alist_skills", - "aget_skill", - "adelete_skill", - "aingest", - "anthropic_messages", - "acreate_interaction", - "aget_interaction", - "adelete_interaction", - "acancel_interaction", - "acreate_agent", - "alist_agents", - "aget_agent", - "adelete_agent", - "alist_agent_versions", - "asend_message", - "call_mcp_tool", - "acancel_batch", - "afile_delete", - "acreate_eval", - "alist_evals", - "aget_eval", - "aupdate_eval", - "adelete_eval", - "acancel_eval", - "acreate_run", - "alist_runs", - "aget_run", - "acancel_run", - "adelete_run", - ], + route_type: RouteType, user_api_key_dict: UserAPIKeyAuth | None = None, ): """ Common helper to route the request """ + try: + return await _route_request_single_attempt( + data=data, + llm_router=llm_router, + user_model=user_model, + route_type=route_type, + user_api_key_dict=user_api_key_dict, + ) + except ProxyModelNotFoundError: + requested_model: Final = data.get("model", "") + if not isinstance(requested_model, str) or not requested_model: + raise + from litellm.proxy import proxy_server + from litellm.proxy.common_utils.registry_read_through import ( + model_registry_read_through, + ) + + if not await model_registry_read_through.attempt(requested_model): + raise + return await _route_request_single_attempt( + data=data, + llm_router=proxy_server.llm_router, + user_model=user_model, + route_type=route_type, + user_api_key_dict=user_api_key_dict, + ) + + +async def _route_request_single_attempt( # noqa: ANN202 # returns unawaited provider coroutines; the inferred union keeps route_request's callers typed + data: dict, # noqa: LIT001 # request body is the proxy-wide mutable dict contract shared with route_request + llm_router: LitellmRouter | None, + user_model: str | None, + route_type: RouteType, + user_api_key_dict: UserAPIKeyAuth | None = None, +): raise_if_required_body_param_missing(route_type=route_type, data=data) await add_shared_session_to_data(data) diff --git a/ruff.toml b/ruff.toml index 095e3e24c52..bd3f8334e94 100644 --- a/ruff.toml +++ b/ruff.toml @@ -6,7 +6,7 @@ lint.extend-select = ["T20", "PGH004", "RUF008", "RUF009", "RUF100"] # litellm's own ruff config both rely on suppressions this config can't see. lint.external = [ # Enforced by the strict-rule gate (scripts/ruff_strict_gate.py + ruff-strict.toml) - "C901", "TID251", + "ANN202", "C901", "TID251", # Enforced by upstream litellm's ruff config, but not run in this repo's CI "PLC0415", "E402", "BLE001", "ARG002", "S102", "S324", "S606", "D401", "F403", "F405", ] diff --git a/tests/test_litellm/proxy/common_utils/test_registry_read_through.py b/tests/test_litellm/proxy/common_utils/test_registry_read_through.py new file mode 100644 index 00000000000..e1c8f031579 --- /dev/null +++ b/tests/test_litellm/proxy/common_utils/test_registry_read_through.py @@ -0,0 +1,231 @@ +import asyncio +from typing import Final + +import pytest + +from litellm.proxy.common_utils.registry_read_through import RegistryReadThrough + + +class ResyncSpy: + def __init__(self, found: bool = True, error: Exception | None = None) -> None: + self.found = found + self.error = error + self.calls: list[str] = [] + + async def __call__(self, key: str) -> bool: + self.calls.append(key) + if self.error is not None: + raise self.error + return self.found + + +@pytest.mark.asyncio +async def test_attempt_returns_true_when_resync_finds_object(): + spy: Final = ResyncSpy(found=True) + read_through: Final = RegistryReadThrough(resync=spy) + + assert await read_through.attempt("new-model") is True + assert spy.calls == ["new-model"] + + +@pytest.mark.asyncio +async def test_attempt_found_key_is_not_negative_cached(): + spy: Final = ResyncSpy(found=True) + read_through: Final = RegistryReadThrough(resync=spy) + + assert await read_through.attempt("new-model") is True + assert await read_through.attempt("new-model") is True + assert spy.calls == ["new-model", "new-model"] + + +@pytest.mark.asyncio +async def test_missing_key_is_negative_cached_within_ttl(): + spy: Final = ResyncSpy(found=False) + read_through: Final = RegistryReadThrough(resync=spy, miss_ttl_seconds=60.0) + + assert await read_through.attempt("ghost-model") is False + assert await read_through.attempt("ghost-model") is False + assert spy.calls == ["ghost-model"] + + +@pytest.mark.asyncio +async def test_negative_cache_expires_and_resync_runs_again(): + spy: Final = ResyncSpy(found=False) + read_through: Final = RegistryReadThrough(resync=spy, miss_ttl_seconds=0.05) + + assert await read_through.attempt("ghost-model") is False + await asyncio.sleep(0.1) + assert await read_through.attempt("ghost-model") is False + assert spy.calls == ["ghost-model", "ghost-model"] + + +@pytest.mark.asyncio +async def test_resync_exception_returns_false_without_negative_caching(): + spy: Final = ResyncSpy(error=RuntimeError("db down")) + read_through: Final = RegistryReadThrough(resync=spy) + + assert await read_through.attempt("new-model") is False + assert await read_through.attempt("new-model") is False + assert spy.calls == ["new-model", "new-model"] + + +@pytest.mark.asyncio +async def test_concurrent_attempts_for_missing_key_resync_once(): + class SlowResyncSpy(ResyncSpy): + async def __call__(self, key: str) -> bool: + await asyncio.sleep(0.05) + return await super().__call__(key) + + spy: Final = SlowResyncSpy(found=False) + read_through: Final = RegistryReadThrough(resync=spy, miss_ttl_seconds=60.0) + + results: Final = await asyncio.gather(*(read_through.attempt("ghost-model") for _ in range(5))) + assert results == [False] * 5 + assert spy.calls == ["ghost-model"] + + +@pytest.mark.asyncio +async def test_distinct_keys_do_not_share_negative_cache(): + spy: Final = ResyncSpy(found=False) + read_through: Final = RegistryReadThrough(resync=spy, miss_ttl_seconds=60.0) + + assert await read_through.attempt("ghost-a") is False + assert await read_through.attempt("ghost-b") is False + assert spy.calls == ["ghost-a", "ghost-b"] + + +class FakeAgentRow: + def __init__(self, agent_id: str, agent_name: str) -> None: + self.agent_id = agent_id + self.agent_name = agent_name + self.object_permission = None + self.spend = 0.0 + + def __iter__(self): + return iter( + { + "agent_id": self.agent_id, + "agent_name": self.agent_name, + "agent_card_params": {"name": self.agent_name, "url": "http://db-agent"}, + "litellm_params": {}, + }.items() + ) + + +@pytest.fixture +def clean_agent_registry(): + from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry + + original_agents: Final = list(global_agent_registry.agent_list) + original_config_agents: Final = getattr(global_agent_registry, "config_agents", ()) + global_agent_registry.agent_list = [] + global_agent_registry.config_agents = () + try: + yield global_agent_registry + finally: + global_agent_registry.agent_list = original_agents + global_agent_registry.config_agents = original_config_agents + + +@pytest.mark.asyncio +async def test_get_agent_with_read_through_recovers_agent_created_on_sibling_replica( + clean_agent_registry, monkeypatch +): + from unittest.mock import AsyncMock, MagicMock + + import litellm.proxy.proxy_server as proxy_server + from litellm.proxy.common_utils.registry_read_through import get_agent_with_read_through + + agent_id: Final = "read-through-db-agent-id" + prisma_client: Final = MagicMock() + prisma_client.db.litellm_agentstable.find_many = AsyncMock( + return_value=[FakeAgentRow(agent_id, "read-through-db-agent")] + ) + monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) + monkeypatch.setattr(proxy_server, "store_model_in_db", True) + + assert clean_agent_registry.get_agent_by_id(agent_id=agent_id) is None + agent: Final = await get_agent_with_read_through(agent_id) + + assert agent is not None + assert agent.agent_id == agent_id + + +@pytest.mark.asyncio +async def test_get_agent_with_read_through_returns_none_for_unknown_agent(clean_agent_registry, monkeypatch): + from unittest.mock import AsyncMock, MagicMock + + import litellm.proxy.proxy_server as proxy_server + from litellm.proxy.common_utils.registry_read_through import get_agent_with_read_through + + prisma_client: Final = MagicMock() + prisma_client.db.litellm_agentstable.find_many = AsyncMock(return_value=[]) + monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) + monkeypatch.setattr(proxy_server, "store_model_in_db", True) + + assert await get_agent_with_read_through("agent-nobody-created") is None + + +class FakeGuardrailRow: + def __init__(self, guardrail_id: str, guardrail_name: str) -> None: + self.guardrail_id = guardrail_id + self.guardrail_name = guardrail_name + + def __iter__(self): + return iter( + { + "guardrail_id": self.guardrail_id, + "guardrail_name": self.guardrail_name, + "litellm_params": { + "guardrail": "litellm_content_filter", + "mode": "pre_call", + "default_on": True, + "blocked_words": [{"keyword": "secret", "action": "BLOCK"}], + }, + "guardrail_info": {}, + }.items() + ) + + +@pytest.mark.asyncio +async def test_get_guardrail_with_read_through_recovers_guardrail_created_on_sibling_replica(monkeypatch): + from unittest.mock import AsyncMock, MagicMock + + import litellm.proxy.proxy_server as proxy_server + from litellm.proxy.common_utils.registry_read_through import ( + get_initialized_guardrail_with_read_through, + ) + from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER + + guardrail_id: Final = "read-through-db-guardrail-id" + guardrail_name: Final = "read-through-db-guardrail" + prisma_client: Final = MagicMock() + prisma_client.db.litellm_guardrailstable.find_many = AsyncMock( + return_value=[FakeGuardrailRow(guardrail_id, guardrail_name)] + ) + monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) + monkeypatch.setattr(proxy_server, "store_model_in_db", True) + + try: + guardrail: Final = await get_initialized_guardrail_with_read_through(guardrail_name=guardrail_name) + assert guardrail is not None + assert guardrail.guardrail_name == guardrail_name + finally: + IN_MEMORY_GUARDRAIL_HANDLER.delete_in_memory_guardrail(guardrail_id) + + +@pytest.mark.asyncio +async def test_get_guardrail_with_read_through_returns_none_for_unknown_guardrail(monkeypatch): + from unittest.mock import AsyncMock, MagicMock + + import litellm.proxy.proxy_server as proxy_server + from litellm.proxy.common_utils.registry_read_through import ( + get_initialized_guardrail_with_read_through, + ) + + prisma_client: Final = MagicMock() + prisma_client.db.litellm_guardrailstable.find_many = AsyncMock(return_value=[]) + monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) + monkeypatch.setattr(proxy_server, "store_model_in_db", True) + + assert await get_initialized_guardrail_with_read_through(guardrail_name="guardrail-nobody-created") is None diff --git a/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py b/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py index 3240ad20edb..1722d8c377b 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py +++ b/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py @@ -430,3 +430,99 @@ async def test_delete_access_group_ignores_models_that_were_already_dead(): assert response.models_updated == 1 mock_prisma.db.litellm_proxymodeltable.update.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_create_access_group_read_through_recovers_model_created_on_sibling_replica(): + """Regression: an access group referencing a model that another replica just wrote + to the DB must be created instead of 400ing until the periodic config reload.""" + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.management_endpoints.model_access_group_management_endpoints import ( + create_model_group, + ) + from litellm.types.proxy.management_endpoints.model_management_endpoints import ( + NewModelGroupRequest, + ) + + from types import SimpleNamespace + + model_name = "e2e-ag-sibling-replica-model" + db_row = SimpleNamespace( + model_id=f"{model_name}-id", + model_name=model_name, + litellm_params={"model": "openai/gpt-4o", "api_key": "fake", "mock_response": "hi"}, + model_info={}, + blocked=False, + ) + + mock_router = Router( + model_list=[ + { + "model_name": "some-other-model", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "fake"}, + } + ] + ) + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(side_effect=[[db_row], [], [db_row]]) + mock_prisma.db.litellm_proxymodeltable.update = AsyncMock() + + with ( + patch("litellm.proxy.proxy_server.llm_router", mock_router), + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch("litellm.proxy.proxy_server.store_model_in_db", True), + patch( + "litellm.proxy.management_endpoints.model_access_group_management_endpoints.clear_cache", + new=AsyncMock(return_value=None), + ), + ): + response = await create_model_group( + data=NewModelGroupRequest(access_group="replica-lag-group", model_names=[model_name]), + user_api_key_dict=UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN), + ) + + assert response.models_updated == 1 + assert response.model_names == [model_name] + assert mock_prisma.db.litellm_proxymodeltable.find_many.await_args_list[0].kwargs["where"] == { + "OR": [{"model_name": model_name}, {"model_id": model_name}] + } + + +@pytest.mark.asyncio +async def test_create_access_group_model_missing_everywhere_still_400s(): + from fastapi import HTTPException + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.management_endpoints.model_access_group_management_endpoints import ( + create_model_group, + ) + from litellm.types.proxy.management_endpoints.model_management_endpoints import ( + NewModelGroupRequest, + ) + + model_name = "e2e-ag-model-nobody-created" + mock_router = Router( + model_list=[ + { + "model_name": "some-other-model", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "fake"}, + } + ] + ) + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + + with ( + patch("litellm.proxy.proxy_server.llm_router", mock_router), + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch("litellm.proxy.proxy_server.store_model_in_db", True), + ): + with pytest.raises(HTTPException) as exc_info: + await create_model_group( + data=NewModelGroupRequest(access_group="ghost-group", model_names=[model_name]), + user_api_key_dict=UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN), + ) + + assert exc_info.value.status_code == 400 + assert model_name in str(exc_info.value.detail) diff --git a/tests/test_litellm/proxy/test_route_a2a_models.py b/tests/test_litellm/proxy/test_route_a2a_models.py index 616fa62cda5..22f99d03d21 100644 --- a/tests/test_litellm/proxy/test_route_a2a_models.py +++ b/tests/test_litellm/proxy/test_route_a2a_models.py @@ -49,6 +49,7 @@ async def test_route_a2a_model_bypasses_router(): ) mock_registry = Mock() + mock_registry.get_agent_by_id = Mock(return_value=None) mock_registry.get_agent_by_name = Mock(return_value=mock_agent) # Mock litellm.acompletion to verify it's called @@ -104,3 +105,77 @@ async def test_route_non_a2a_model_raises_error_if_not_in_router(): user_model=None, route_type="acompletion", ) + + +class _DbAgentRow: + def __init__(self, agent_id: str, agent_name: str) -> None: + self.agent_id = agent_id + self.agent_name = agent_name + self.object_permission = None + self.spend = 0.0 + + def __iter__(self): + return iter( + { + "agent_id": self.agent_id, + "agent_name": self.agent_name, + "agent_card_params": {"name": self.agent_name, "url": "http://sibling-db-agent.example.com"}, + "litellm_params": {}, + }.items() + ) + + +def _router_without_models(): + mock_router = Mock() + mock_router.model_names = [] + mock_router.deployment_names = [] + mock_router.has_model_id = Mock(return_value=False) + mock_router.model_group_alias = None + mock_router.router_general_settings = Mock(pass_through_all_models=False) + mock_router.default_deployment = None + mock_router.pattern_router = Mock(patterns=[]) + mock_router.map_team_model = Mock(return_value=None) + return mock_router + + +@pytest.mark.asyncio +async def test_route_a2a_model_read_through_recovers_agent_created_on_sibling_replica(monkeypatch): + import litellm.proxy.proxy_server as proxy_server + from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry + + agent_name = "a2a-sibling-replica-agent" + prisma_client = Mock() + prisma_client.db.litellm_agentstable.find_many = AsyncMock( + return_value=[_DbAgentRow("a2a-sibling-replica-agent-id", agent_name)] + ) + monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) + monkeypatch.setattr(proxy_server, "store_model_in_db", True) + + original_agents = list(global_agent_registry.agent_list) + original_config_agents = getattr(global_agent_registry, "config_agents", ()) + global_agent_registry.agent_list = [] + global_agent_registry.config_agents = () + + data = { + "model": f"a2a/{agent_name}", + "messages": [{"role": "user", "content": "Hello"}], + } + mock_acompletion = AsyncMock(return_value={"id": "read-through-response"}) + + try: + with patch("litellm.acompletion", mock_acompletion): + await route_request( + data=data, + llm_router=_router_without_models(), + user_model=None, + route_type="acompletion", + ) + finally: + global_agent_registry.agent_list = original_agents + global_agent_registry.config_agents = original_config_agents + + mock_acompletion.assert_called_once() + call_kwargs = mock_acompletion.call_args.kwargs + assert call_kwargs["model"] == f"a2a/{agent_name}" + assert call_kwargs["api_base"] == "http://sibling-db-agent.example.com" + prisma_client.db.litellm_agentstable.find_many.assert_awaited() diff --git a/tests/test_litellm/proxy/test_route_llm_request.py b/tests/test_litellm/proxy/test_route_llm_request.py index 3ae0e1e7d18..ebd52c448bc 100644 --- a/tests/test_litellm/proxy/test_route_llm_request.py +++ b/tests/test_litellm/proxy/test_route_llm_request.py @@ -1091,3 +1091,124 @@ async def test_route_request_rejects_chat_completion_without_messages(): assert exc_info.value.status_code == 400 assert exc_info.value.param == "messages" llm_router.acompletion.assert_not_called() + + +class FakeProxyModelTable: + def __init__(self, rows): + self.rows = rows + self.find_many_wheres = [] + + async def find_many(self, where=None, **kwargs): + self.find_many_wheres.append(where) + return list(self.rows) + + +def _fake_prisma_client_with_models(rows): + from types import SimpleNamespace + + table = FakeProxyModelTable(rows) + return SimpleNamespace(db=SimpleNamespace(litellm_proxymodeltable=table)), table + + +def _db_model_row(model_name: str, mock_response: str): + from types import SimpleNamespace + + return SimpleNamespace( + model_id=f"{model_name}-id", + model_name=model_name, + litellm_params={"model": "openai/gpt-4o", "api_key": "fake", "mock_response": mock_response}, + model_info={}, + blocked=False, + ) + + +@pytest.mark.asyncio +async def test_route_request_read_through_recovers_model_created_on_sibling_replica(monkeypatch): + """Regression: a model written to the DB by another replica must be served on + first request instead of 400ing until the periodic config reload.""" + import litellm + import litellm.proxy.proxy_server as proxy_server + + model_name = "e2e-sibling-replica-model" + router = litellm.Router( + model_list=[ + { + "model_name": "some-other-model", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "fake"}, + } + ] + ) + fake_prisma, table = _fake_prisma_client_with_models([_db_model_row(model_name, "hello-from-db")]) + monkeypatch.setattr(proxy_server, "prisma_client", fake_prisma) + monkeypatch.setattr(proxy_server, "store_model_in_db", True) + monkeypatch.setattr(proxy_server, "llm_router", router) + + llm_call = await route_request( + data={"model": model_name, "messages": [{"role": "user", "content": "hi"}]}, + llm_router=router, + user_model=None, + route_type="acompletion", + ) + response = await llm_call + + assert response.choices[0].message.content == "hello-from-db" + assert len(table.find_many_wheres) == 1 + assert table.find_many_wheres[0] == {"OR": [{"model_name": model_name}, {"model_id": model_name}]} + + +@pytest.mark.asyncio +async def test_route_request_unknown_model_raises_and_hits_db_once_within_ttl(monkeypatch): + import litellm + import litellm.proxy.proxy_server as proxy_server + + model_name = "e2e-model-nobody-created" + router = litellm.Router( + model_list=[ + { + "model_name": "some-other-model", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "fake"}, + } + ] + ) + fake_prisma, table = _fake_prisma_client_with_models([]) + monkeypatch.setattr(proxy_server, "prisma_client", fake_prisma) + monkeypatch.setattr(proxy_server, "store_model_in_db", True) + monkeypatch.setattr(proxy_server, "llm_router", router) + + data = {"model": model_name, "messages": [{"role": "user", "content": "hi"}]} + with pytest.raises(ProxyModelNotFoundError): + await route_request(data=data, llm_router=router, user_model=None, route_type="acompletion") + with pytest.raises(ProxyModelNotFoundError): + await route_request(data=data, llm_router=router, user_model=None, route_type="acompletion") + + assert len(table.find_many_wheres) == 1 + + +@pytest.mark.asyncio +async def test_route_request_read_through_disabled_without_store_model_in_db(monkeypatch): + import litellm + import litellm.proxy.proxy_server as proxy_server + + model_name = "e2e-config-only-proxy-model" + router = litellm.Router( + model_list=[ + { + "model_name": "some-other-model", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "fake"}, + } + ] + ) + fake_prisma, table = _fake_prisma_client_with_models([_db_model_row(model_name, "should-not-load")]) + monkeypatch.setattr(proxy_server, "prisma_client", fake_prisma) + monkeypatch.setattr(proxy_server, "store_model_in_db", False) + monkeypatch.setattr(proxy_server, "llm_router", router) + + with pytest.raises(ProxyModelNotFoundError): + await route_request( + data={"model": model_name, "messages": [{"role": "user", "content": "hi"}]}, + llm_router=router, + user_model=None, + route_type="acompletion", + ) + + assert table.find_many_wheres == [] From 312d12fe0c98409b07b65fa7a677db0981b60cb5 Mon Sep 17 00:00:00 2001 From: Shivi Jain Date: Fri, 14 Aug 2026 21:39:34 +0530 Subject: [PATCH 07/83] fix(proxy): enforce project ITPM/OTPM quota on every Responses WebSocket frame The connection-level pre-call hook only ran once per WebSocket connection, so a project caller could send unlimited high-token response.create frames after a single minimal reservation. Adds enforce_project_io_token_quota_for_frame to the v3 rate limiter and wires it into both the native and managed WebSocket handlers via a duck-typed litellm.callbacks lookup, so the SDK layer stays free of proxy imports. A rejected frame gets an error event; the connection stays open for the client to retry. Also fixes the RET504 and BLE001 strict-lint-budget violations the litellm_internal_staging merge introduced in parallel_request_limiter_v3.py, which were failing the lint check. --- litellm/llms/custom_httpx/llm_http_handler.py | 28 +++ .../hooks/parallel_request_limiter_v3.py | 40 ++++- litellm/responses/streaming_iterator.py | 112 +++++++++++- .../hooks/test_parallel_request_limiter_v3.py | 56 ++++++ .../test_responses_websocket_all_providers.py | 165 ++++++++++++++++++ 5 files changed, 396 insertions(+), 5 deletions(-) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 721b9545ac1..e4c3a29b956 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -252,6 +252,30 @@ def _has_pre_call_deployment_hook(logging_obj: LiteLLMLoggingObj) -> bool: return False +def _collect_ws_project_quota_callbacks() -> list: + """Duck-type discover proxy hooks exposing per-frame project ITPM/OTPM + enforcement, so the Responses WebSocket loop can charge every + ``response.create`` frame, not just the connection's first one. + + Uses duck-typing on ``litellm.callbacks`` (rather than importing the + proxy hook directly) to avoid a layering violation (SDK importing from + the proxy layer). + """ + try: + import litellm as _litellm + + return [ + cb for cb in _litellm.callbacks if callable(getattr(cb, "enforce_project_io_token_quota_for_frame", None)) + ] + except Exception as exc: # noqa: BLE001 - discovery must not block the connection + verbose_logger.warning( + "Responses WebSocket: failed to collect project quota callbacks — " + "per-frame ITPM/OTPM enforcement will be skipped. Error: %s", + exc, + ) + return [] + + class BaseLLMHTTPHandler: async def _make_common_async_call( self, @@ -6168,6 +6192,8 @@ class BaseLLMHTTPHandler: - Uses ManagedResponsesWebSocketHandler which makes HTTP streaming calls - Forwards events over the websocket connection """ + _ws_quota_callbacks: Final = _collect_ws_project_quota_callbacks() + if responses_api_provider_config is None or not responses_api_provider_config.supports_native_websocket(): from litellm.responses.streaming_iterator import ( ManagedResponsesWebSocketHandler, @@ -6184,6 +6210,7 @@ class BaseLLMHTTPHandler: timeout=timeout, custom_llm_provider=custom_llm_provider, first_message=first_message, + quota_callbacks=_ws_quota_callbacks, **kwargs, ) await handler.run() @@ -6304,6 +6331,7 @@ class BaseLLMHTTPHandler: first_message=first_message, guardrail_callbacks=_ws_guardrail_callbacks, output_guardrail_callbacks=_ws_output_guardrail_callbacks, + quota_callbacks=_ws_quota_callbacks, authorized_model=model, ) await streaming.bidirectional_forward() diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 1c7ae1ea9f8..7a97718247e 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -680,7 +680,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): 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( + return 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, @@ -690,7 +690,6 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): 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: @@ -2171,6 +2170,41 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): assert itpm_response is not None return itpm_response, itpm_reserved, 0 + async def enforce_project_io_token_quota_for_frame( + self, + user_api_key_dict: UserAPIKeyAuth, + requested_model: str | None, + estimated_input_tokens: int, + estimated_output_tokens: int, + ) -> None: + """Reserve one WebSocket ``response.create`` frame's tokens against + the caller's project ITPM/OTPM quota. + + The Responses WebSocket connection-level pre-call hook only runs once + per connection, but a connection accepts many ``response.create`` + frames over its lifetime. Without this, a project caller could send + unlimited high-token generations after a single minimal reservation. + There is no per-frame post-call hook to reconcile against, so -- + like the batch rate limiter -- this charges the estimate immediately + and never refunds it. + """ + descriptors: Final[list[RateLimitDescriptor]] = [] + self._add_project_io_token_rate_limit_descriptors_from_metadata( + user_api_key_dict=user_api_key_dict, + requested_model=requested_model, + descriptors=descriptors, + ) + if not descriptors: + return + response, _itpm_reserved, _otpm_reserved = await self.reserve_io_tokens( + descriptors=descriptors, + estimated_input_tokens=estimated_input_tokens, + estimated_output_tokens=estimated_output_tokens, + parent_otel_span=user_api_key_dict.parent_otel_span, + ) + if response["overall_code"] == "OVER_LIMIT": + self._handle_rate_limit_error(response, descriptors, requested_model) + def create_organization_rate_limit_descriptor( self, user_api_key_dict: UserAPIKeyAuth, requested_model: str | None = None ) -> list[RateLimitDescriptor]: @@ -3858,7 +3892,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ], ) continue - except Exception as e: + except Exception as e: # noqa: BLE001 - any Redis/Lua failure degrades to the plain increment fallback, never a 500 verbose_proxy_logger.warning( "Window-guarded token adjustment failed for %s: %s", operation["key"], diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 25e5fcb6976..af561c9092b 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -20,7 +20,7 @@ from litellm.constants import ( LITELLM_MAX_STREAMING_DURATION_SECONDS, STREAM_SSE_DONE_STRING, ) -from litellm.exceptions import MidStreamFallbackError +from litellm.exceptions import MidStreamFallbackError, RateLimitError from litellm.litellm_core_utils.asyncify import run_async_function from litellm.litellm_core_utils.core_helpers import process_response_headers from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -1326,6 +1326,79 @@ def _build_synthetic_response_events( from litellm._logging import verbose_logger +# Conservative per-frame output-token floor used when a response.create +# frame omits max_output_tokens, so a project OTPM quota can't be bypassed +# by simply never declaring an output cap. +_FRAME_NO_MAX_OUTPUT_TOKENS_FLOOR: Final = 1024 + +# Rough chars-per-token ratio for estimating a frame's input tokens without +# resolving a real per-model tokenizer, matching the conservative estimate +# the proxy's own rate limiter uses for the same purpose. +_FRAME_CHARS_PER_TOKEN_ESTIMATE: Final = 4 + + +def _extract_frame_quota_estimate_inputs(msg_obj: Mapping[str, object]) -> tuple[int, int | None]: + """Extract a rough input-token count and any explicit max_output_tokens + from a ``response.create`` frame, handling both wire shapes: + flat: {"type": "response.create", "input": ..., "max_output_tokens": ...} + nested: {"type": "response.create", "response": {"input": ..., "max_output_tokens": ...}} + """ + nested: Final = msg_obj.get("response") + params: Final[Mapping[str, object]] = ( + nested if _is_json_object(nested) and nested else {k: v for k, v in msg_obj.items() if k != "type"} + ) + text_parts: list[str] = [] # mutable-ok: local accumulator built in one pass, not shared + + def _collect_text(value: object) -> None: + if isinstance(value, str): + text_parts.append(value) + elif _is_json_array(value): + for item in value: + if isinstance(item, str): + text_parts.append(item) + elif _is_json_object(item): + _collect_text(item.get("content")) + _collect_text(item.get("text")) + + _collect_text(params.get("input")) + _collect_text(params.get("instructions")) + total_chars: Final = sum(len(part) for part in text_parts) + estimated_input_tokens: Final = max(1, total_chars // _FRAME_CHARS_PER_TOKEN_ESTIMATE) if total_chars else 0 + + max_output_tokens = params.get("max_output_tokens") + return estimated_input_tokens, max_output_tokens if isinstance(max_output_tokens, int) else None + + +async def _enforce_frame_project_quota( + quota_callbacks: Sequence[Any], + user_api_key_dict: UserAPIKeyAuth | None, + model: str | None, + raw_message: str, +) -> None: + """Charge one response.create frame's estimated tokens against every + registered project ITPM/OTPM quota callback, in isolation from PII + masking / logging so a malformed frame still reaches those callbacks.""" + if not quota_callbacks: + return + try: + msg_obj = json.loads(raw_message) + except (json.JSONDecodeError, TypeError): + return + if not _is_json_object(msg_obj) or msg_obj.get("type") != "response.create": + return + estimated_input_tokens, explicit_max_output_tokens = _extract_frame_quota_estimate_inputs(msg_obj) + estimated_output_tokens = ( + explicit_max_output_tokens if explicit_max_output_tokens is not None else _FRAME_NO_MAX_OUTPUT_TOKENS_FLOOR + ) + for callback in quota_callbacks: + await callback.enforce_project_io_token_quota_for_frame( + user_api_key_dict=user_api_key_dict, + requested_model=model, + estimated_input_tokens=estimated_input_tokens, + estimated_output_tokens=estimated_output_tokens, + ) + + RESPONSES_WS_LOGGED_EVENT_TYPES: Final = [ "response.created", "response.completed", @@ -1360,6 +1433,7 @@ class ResponsesWebSocketStreaming: first_message: str | None = None, guardrail_callbacks: list[Any] | None = None, output_guardrail_callbacks: list[PresidioGuardrailCallback] | None = None, + quota_callbacks: list[Any] | None = None, authorized_model: str | None = None, ): self.websocket = websocket @@ -1372,6 +1446,7 @@ class ResponsesWebSocketStreaming: self.first_message = first_message self.guardrail_callbacks: list[Any] = guardrail_callbacks or [] self.output_guardrail_callbacks: list[PresidioGuardrailCallback] = output_guardrail_callbacks or [] + self.quota_callbacks: list[Any] = quota_callbacks or [] # Model name authorized at connection time; enforced on every # response.create frame to prevent deployment-substitution attacks. self.authorized_model: str | None = authorized_model @@ -1781,10 +1856,31 @@ class ResponsesWebSocketStreaming: return json.dumps(evt_obj) if modified else response_str + async def _enforce_or_reject_frame(self, message: str) -> bool: + """Run the per-frame project quota check. + + On rejection, sends an ``error`` event to the client and reports that + the frame must be dropped instead of forwarded, so the connection + stays open for the client to retry once the window resets. + """ + try: + await _enforce_frame_project_quota( + self.quota_callbacks, self.user_api_key_dict, self.authorized_model, message + ) + except RateLimitError as e: + try: + await self.websocket.send_text( + json.dumps({"type": "error", "error": {"type": "rate_limit_exceeded", "message": str(e)}}) + ) + except Exception: # noqa: BLE001, S110 - best-effort notification, client may already be gone + pass + return False + return True + async def client_to_backend(self) -> None: """Forward response.create events from client to backend.""" try: - if self.first_message is not None: + if self.first_message is not None and await self._enforce_or_reject_frame(self.first_message): masked_first: Final = await self._mask_response_create(self.first_message) self._store_input(masked_first) self._store_event(masked_first) @@ -1792,6 +1888,8 @@ class ResponsesWebSocketStreaming: while True: message = await self.websocket.receive_text() + if not await self._enforce_or_reject_frame(message): + continue masked = await self._mask_response_create(message) self._store_input(masked) self._store_event(masked) @@ -1871,6 +1969,7 @@ class ManagedResponsesWebSocketHandler: timeout: float | None = None, custom_llm_provider: str | None = None, first_message: str | None = None, + quota_callbacks: list[Any] | None = None, **kwargs: object, ) -> None: self.websocket = websocket @@ -1887,6 +1986,7 @@ class ManagedResponsesWebSocketHandler: self.custom_llm_provider = custom_llm_provider self._connection_provider = self._resolve_provider(model) or custom_llm_provider self.first_message = first_message + self.quota_callbacks: list[Any] = quota_callbacks or [] # Carry through safe pass-through kwargs (e.g. extra_headers) self.extra_kwargs: dict[str, object] = {k: v for k, v in kwargs.items() if k not in _MANAGED_WS_SKIP_KWARGS} # In-memory session history: response_id → full accumulated message list. @@ -2292,6 +2392,14 @@ class ManagedResponsesWebSocketHandler: verbose_logger.debug("ManagedResponsesWS: error sending warmup ack: %s", exc) return + try: + await _enforce_frame_project_quota( + self.quota_callbacks, self.user_api_key_dict, self.model_group or self.model, raw_message + ) + except RateLimitError as e: + await self._send_error(str(e), error_type="rate_limit_exceeded") + return + call_kwargs: Final = self._build_base_call_kwargs(msg_obj) call_kwargs["stream"] = True 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 ac0396c11c1..4ec0183f33c 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 @@ -3246,6 +3246,62 @@ async def test_project_model_itpm_and_tpm_limits_coexist_v3(): assert "model_per_project_otpm" in descriptor_keys +@pytest.mark.asyncio +async def test_enforce_project_io_token_quota_for_frame_blocks_over_limit_otpm(): + """VERIA regression: the Responses WebSocket connection-level pre-call + hook only runs once, but a connection accepts many response.create + frames. enforce_project_io_token_quota_for_frame is the per-frame check + that closes that gap; it must reserve against the caller's project OTPM + limit and reject once a frame's estimated output tokens exceed it.""" + _api_key = hash_token("sk-ws-frame-otpm") + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + user_api_key_dict = UserAPIKeyAuth( + api_key=_api_key, + project_id="proj-mantle-ws", + project_metadata={"model_otpm_limit": {"gpt-4o": 50}}, + ) + + await handler.enforce_project_io_token_quota_for_frame( + user_api_key_dict=user_api_key_dict, + requested_model="gpt-4o", + estimated_input_tokens=1, + estimated_output_tokens=30, + ) + + with pytest.raises(HTTPException) as exc: + await handler.enforce_project_io_token_quota_for_frame( + user_api_key_dict=user_api_key_dict, + requested_model="gpt-4o", + estimated_input_tokens=1, + estimated_output_tokens=30, + ) + + assert exc.value.status_code == 429 + assert "model_per_project_otpm" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_enforce_project_io_token_quota_for_frame_noop_without_project_limits(): + """A key with no project ITPM/OTPM configured must never be blocked by + the per-frame check (no descriptors to reserve against).""" + _api_key = hash_token("sk-ws-frame-no-limits") + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + user_api_key_dict = UserAPIKeyAuth(api_key=_api_key) + + await handler.enforce_project_io_token_quota_for_frame( + user_api_key_dict=user_api_key_dict, + requested_model="gpt-4o", + estimated_input_tokens=10_000_000, + estimated_output_tokens=10_000_000, + ) + + @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 diff --git a/tests/test_litellm/responses/test_responses_websocket_all_providers.py b/tests/test_litellm/responses/test_responses_websocket_all_providers.py index 4509abc7749..2d523bfdeb3 100644 --- a/tests/test_litellm/responses/test_responses_websocket_all_providers.py +++ b/tests/test_litellm/responses/test_responses_websocket_all_providers.py @@ -1030,6 +1030,171 @@ class TestWebSocketErrorHandling: assert "Invalid JSON" in error_event +class TestWebSocketProjectQuotaEnforcement: + """VERIA regression: the connection-level pre-call hook only runs once, + but a WebSocket connection accepts many response.create frames. Every + frame must be checked against any registered project ITPM/OTPM quota + callback, not just the first one.""" + + @pytest.mark.asyncio + async def test_managed_handler_blocks_frame_rejected_by_quota_callback(self, monkeypatch): + from unittest.mock import AsyncMock, MagicMock + + import litellm + from litellm.exceptions import RateLimitError + from litellm.litellm_core_utils.litellm_logging import Logging + from litellm.responses.streaming_iterator import ( + ManagedResponsesWebSocketHandler, + ) + + aresponses_called = False + + async def fake_aresponses(*args, **kwargs): + nonlocal aresponses_called + aresponses_called = True + + monkeypatch.setattr(litellm, "aresponses", fake_aresponses) + + quota_callback = MagicMock() + quota_callback.enforce_project_io_token_quota_for_frame = AsyncMock( + side_effect=RateLimitError(message="project OTPM exceeded", llm_provider="", model="") + ) + + mock_websocket = MagicMock() + mock_websocket.send_text = AsyncMock() + mock_logging_obj = Logging( + model="test-model", + messages=[], + stream=True, + call_type="aresponses", + start_time=0, + litellm_call_id="test-id", + function_id="test-func", + ) + handler = ManagedResponsesWebSocketHandler( + websocket=mock_websocket, + model="test-model", + logging_obj=mock_logging_obj, + quota_callbacks=[quota_callback], + ) + + await handler._process_response_create(json.dumps({"type": "response.create", "input": "hi"})) + + quota_callback.enforce_project_io_token_quota_for_frame.assert_awaited_once() + assert aresponses_called is False + mock_websocket.send_text.assert_called_once() + error_event = mock_websocket.send_text.call_args[0][0] + assert "rate_limit_exceeded" in error_event + + @pytest.mark.asyncio + async def test_managed_handler_forwards_frame_allowed_by_quota_callback(self, monkeypatch): + from unittest.mock import AsyncMock, MagicMock + + import litellm + from litellm.litellm_core_utils.litellm_logging import Logging + from litellm.responses.streaming_iterator import ( + ManagedResponsesWebSocketHandler, + ) + + aresponses_called = False + + async def fake_aresponses(*args, **kwargs): + nonlocal aresponses_called + aresponses_called = True + + async def _empty(): + return + yield + + return _empty() + + monkeypatch.setattr(litellm, "aresponses", fake_aresponses) + + quota_callback = MagicMock() + quota_callback.enforce_project_io_token_quota_for_frame = AsyncMock(return_value=None) + + mock_websocket = MagicMock() + mock_websocket.send_text = AsyncMock() + mock_logging_obj = Logging( + model="test-model", + messages=[], + stream=True, + call_type="aresponses", + start_time=0, + litellm_call_id="test-id", + function_id="test-func", + ) + handler = ManagedResponsesWebSocketHandler( + websocket=mock_websocket, + model="test-model", + logging_obj=mock_logging_obj, + quota_callbacks=[quota_callback], + ) + + await handler._process_response_create(json.dumps({"type": "response.create", "input": "hi"})) + + quota_callback.enforce_project_io_token_quota_for_frame.assert_awaited_once() + assert aresponses_called is True + + @pytest.mark.asyncio + async def test_native_handler_blocks_frame_rejected_by_quota_callback(self): + from unittest.mock import AsyncMock, MagicMock + + from litellm.exceptions import RateLimitError + from litellm.responses.streaming_iterator import ResponsesWebSocketStreaming + + quota_callback = MagicMock() + quota_callback.enforce_project_io_token_quota_for_frame = AsyncMock( + side_effect=RateLimitError(message="project OTPM exceeded", llm_provider="", model="") + ) + + mock_backend_ws = MagicMock() + mock_backend_ws.send = AsyncMock() + mock_websocket = MagicMock() + mock_websocket.send_text = AsyncMock() + + handler = ResponsesWebSocketStreaming( + websocket=mock_websocket, + backend_ws=mock_backend_ws, + logging_obj=MagicMock(), + authorized_model="gpt-4o", + quota_callbacks=[quota_callback], + ) + + allowed = await handler._enforce_or_reject_frame( + json.dumps({"type": "response.create", "input": "hi"}) + ) + + assert allowed is False + mock_backend_ws.send.assert_not_called() + mock_websocket.send_text.assert_called_once() + assert "rate_limit_exceeded" in mock_websocket.send_text.call_args[0][0] + + @pytest.mark.asyncio + async def test_native_handler_forwards_frame_allowed_by_quota_callback(self): + from unittest.mock import AsyncMock, MagicMock + + from litellm.responses.streaming_iterator import ResponsesWebSocketStreaming + + quota_callback = MagicMock() + quota_callback.enforce_project_io_token_quota_for_frame = AsyncMock(return_value=None) + + handler = ResponsesWebSocketStreaming( + websocket=MagicMock(), + backend_ws=MagicMock(), + logging_obj=MagicMock(), + authorized_model="gpt-4o", + quota_callbacks=[quota_callback], + ) + + allowed = await handler._enforce_or_reject_frame( + json.dumps({"type": "response.create", "input": "hi"}) + ) + + assert allowed is True + quota_callback.enforce_project_io_token_quota_for_frame.assert_awaited_once() + + class TestNativeWebSocketGuardrails: @pytest.mark.asyncio async def test_response_create_injects_authorized_model(self): From 2b23295f82d29cac7eb97474ea0d53b992c3f511 Mon Sep 17 00:00:00 2001 From: Shivi Jain Date: Fri, 14 Aug 2026 23:31:18 +0530 Subject: [PATCH 08/83] fix(proxy): reconcile project quota reservations --- basedpyright-code-budget.json | 6 +- litellm/llms/custom_httpx/llm_http_handler.py | 26 +- litellm/proxy/hooks/batch_rate_limiter.py | 50 +- .../hooks/parallel_request_limiter_v3.py | 589 +++++++++--------- litellm/responses/streaming_iterator.py | 59 +- .../proxy/hooks/test_tpm_concurrent.py | 10 +- type-discipline-budget.json | 8 +- 7 files changed, 407 insertions(+), 341 deletions(-) diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index 521b4315e6e..67b0575a0cb 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -57,7 +57,7 @@ "limit": 5707 }, "reportMissingTypeArgument": { - "limit": 15642 + "limit": 15641 }, "reportMissingTypeStubs": { "limit": 40 @@ -108,7 +108,7 @@ "limit": 39237 }, "reportUnknownParameterType": { - "limit": 19969 + "limit": 19968 }, "reportUnknownVariableType": { "limit": 30881 @@ -132,7 +132,7 @@ "limit": 27 }, "reportUnusedClass": { - "limit": 23 + "limit": 22 }, "reportUnusedFunction": { "limit": 139 diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index e4c3a29b956..8627d1797a5 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -2,7 +2,7 @@ import asyncio import json import os import ssl -from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping +from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping, Sequence from contextlib import asynccontextmanager from functools import lru_cache from types import ModuleType @@ -69,6 +69,7 @@ from litellm.llms.custom_httpx.http_handler import ( from litellm.responses.streaming_iterator import ( BaseResponsesAPIStreamingIterator, MockResponsesAPIStreamingIterator, + ProjectQuotaCallback, ResponsesAPIStreamingIterator, ResponsesWebSocketStreaming, SyncResponsesAPIStreamingIterator, @@ -252,7 +253,7 @@ def _has_pre_call_deployment_hook(logging_obj: LiteLLMLoggingObj) -> bool: return False -def _collect_ws_project_quota_callbacks() -> list: +def _collect_ws_project_quota_callbacks() -> tuple[ProjectQuotaCallback, ...]: """Duck-type discover proxy hooks exposing per-frame project ITPM/OTPM enforcement, so the Responses WebSocket loop can charge every ``response.create`` frame, not just the connection's first one. @@ -261,19 +262,16 @@ def _collect_ws_project_quota_callbacks() -> list: proxy hook directly) to avoid a layering violation (SDK importing from the proxy layer). """ - try: - import litellm as _litellm + import litellm as _litellm - return [ - cb for cb in _litellm.callbacks if callable(getattr(cb, "enforce_project_io_token_quota_for_frame", None)) - ] - except Exception as exc: # noqa: BLE001 - discovery must not block the connection - verbose_logger.warning( - "Responses WebSocket: failed to collect project quota callbacks — " - "per-frame ITPM/OTPM enforcement will be skipped. Error: %s", - exc, - ) - return [] + callbacks: Final = cast( # cast-ok: callback registry is inspected before protocol use + Sequence[object], _litellm.callbacks + ) + return tuple( + cast(ProjectQuotaCallback, callback) # cast-ok: required callback method is callable + for callback in callbacks + if callable(getattr(callback, "enforce_project_io_token_quota_for_frame", None)) + ) class BaseLLMHTTPHandler: diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index 3091b0f6973..efef246a7a6 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -18,11 +18,12 @@ Quick summary: """ import json -from collections.abc import Iterable +from collections.abc import Iterable, Mapping, Sequence +from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn from fastapi import HTTPException -from pydantic import BaseModel +from pydantic import BaseModel, TypeAdapter import litellm from litellm._logging import verbose_proxy_logger @@ -77,6 +78,9 @@ else: RateLimitDescriptor = dict[str, Any] +_BATCH_BODY_ADAPTER: Final = TypeAdapter(dict[str, object]) + + class BatchFileUsage(BaseModel): """ Internal model for batch file usage tracking, used for batch rate limiting @@ -214,7 +218,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): tpm_limit_type=None, model_has_failures=False, ) - self.parallel_request_limiter._add_project_io_token_rate_limit_descriptors_from_metadata( + self.parallel_request_limiter.add_project_io_token_rate_limit_descriptors_from_metadata( user_api_key_dict=user_api_key_dict, requested_model=self._get_batch_routing_model(data), descriptors=descriptors, @@ -314,7 +318,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): def _estimate_entry_output_tokens( self, - entry: dict, + entry: Mapping[str, object], min_configured_otpm_limit: int | None, ) -> int: """Conservative per-row output-token estimate for the project OTPM reservation. @@ -325,16 +329,21 @@ class _PROXY_BatchRateLimiter(CustomLogger): that omits ``max_tokens`` can't be used to bypass OTPM the way an unbounded streaming request could. """ - body: Final = entry.get("body", {}) or {} + raw_body: Final = entry.get("body") + body: Final[Mapping[str, object]] = ( + MappingProxyType(_BATCH_BODY_ADAPTER.validate_python(raw_body)) + if isinstance(raw_body, Mapping) + else MappingProxyType({}) # mutable-ok: immediately frozen empty fallback + ) if body.get("input") is not None and body.get("messages") is None and body.get("prompt") is None: return 0 # embeddings: no output tokens - explicit_cap = body.get("max_tokens", body.get("max_completion_tokens")) + explicit_cap: Final = body.get("max_tokens", body.get("max_completion_tokens")) if explicit_cap is not None: try: return max(0, int(explicit_cap)) except (TypeError, ValueError): pass - return self.parallel_request_limiter._no_max_tokens_output_floor(min_configured_otpm_limit) + return self.parallel_request_limiter.no_max_tokens_output_floor(min_configured_otpm_limit) @staticmethod def _has_applicable_batch_rate_limits( @@ -446,7 +455,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): f"Limit resets at: {reset_time_formatted}" ) else: # tokens - batch_token_count = ( + batch_token_count: Final = ( batch_usage.output_tokens if descriptor.get("key") == PROJECT_OTPM_DESCRIPTOR_KEY else batch_usage.total_tokens @@ -496,8 +505,8 @@ class _PROXY_BatchRateLimiter(CustomLogger): data=data, ) - increments: Final[list[dict[Literal["requests", "tokens"], int]]] = [ - { + increments: Final = [ # mutable-ok: atomic limiter API requires mutable increment records + { # mutable-ok: atomic limiter API requires mutable increment records "requests": batch_usage.request_count, "tokens": batch_usage.output_tokens if d.get("key") == PROJECT_OTPM_DESCRIPTOR_KEY @@ -530,7 +539,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai", user_api_key_dict: UserAPIKeyAuth | None = None, data: dict | None = None, - descriptors: list["RateLimitDescriptor"] | None = None, + descriptors: Sequence["RateLimitDescriptor"] | None = None, ) -> BatchFileUsage: """ Count number of requests and tokens in a batch input file. @@ -545,13 +554,14 @@ class _PROXY_BatchRateLimiter(CustomLogger): Returns: BatchFileUsage with total_tokens, output_tokens, and request_count """ - otpm_limits: Final = [ + otpm_limits: Final = tuple( int(v) - for d in (descriptors or []) + for d in (descriptors or ()) if d.get("key") == PROJECT_OTPM_DESCRIPTOR_KEY - for v in [(d.get("rate_limit") or {}).get("tokens_per_unit")] + for rate_limit in (d.get("rate_limit"),) + for v in (rate_limit.get("tokens_per_unit") if rate_limit is not None else None,) if v is not None - ] + ) min_configured_otpm_limit: Final = min(otpm_limits) if otpm_limits else None try: # Check if this is a managed file (base64 encoded unified file ID) @@ -605,7 +615,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): # the counter can't measure is estimated, not hard-rejected. models: Final[set] = set() total_tokens = 0 - output_tokens = 0 + output_tokens = 0 # rebind-ok: accumulated per JSONL row in the loop below request_count = 0 for raw_line in _iter_batch_input_lines(file_content_bytes): request_count += 1 @@ -613,9 +623,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): entry = json.loads(raw_line) except Exception: total_tokens += _estimate_batch_entry_tokens(raw_line) - output_tokens += self.parallel_request_limiter._no_max_tokens_output_floor( - min_configured_otpm_limit - ) + output_tokens += self.parallel_request_limiter.no_max_tokens_output_floor(min_configured_otpm_limit) continue if isinstance(entry, dict): model = (entry.get("body") or {}).get("model") @@ -623,9 +631,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): models.add(model) output_tokens += self._estimate_entry_output_tokens(entry, min_configured_otpm_limit) else: - output_tokens += self.parallel_request_limiter._no_max_tokens_output_floor( - min_configured_otpm_limit - ) + output_tokens += self.parallel_request_limiter.no_max_tokens_output_floor(min_configured_otpm_limit) try: total_tokens += _count_entry_tokens(entry) except Exception: diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 7a97718247e..f858fd2af98 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -22,7 +22,7 @@ from typing import ( TypedDict, ) -from typing_extensions import NotRequired +from typing_extensions import NotRequired, ReadOnly from litellm import DualCache from litellm._logging import verbose_proxy_logger @@ -69,6 +69,7 @@ else: Span = Any InternalUsageCache = Any + BATCH_RATE_LIMITER_SCRIPT: Final = """ local results = {} local now = tonumber(ARGV[1]) @@ -413,7 +414,7 @@ class RateLimitStatus(TypedDict): class RateLimitResponse(TypedDict): overall_code: str statuses: list[RateLimitStatus] - reservation_windows: NotRequired[frozenset[tuple[str, str, Literal["redis", "local"]]]] + reservation_windows: NotRequired[ReadOnly[frozenset[tuple[str, str, Literal["redis", "local"]]]]] class ReservationAwareIncrementOperation(RedisPipelineIncrementOperation): @@ -452,6 +453,7 @@ class AtomicCounterMeta(TypedDict): class AtomicCounterState(TypedDict): window_expired: bool current: int + window_start: ReadOnly[str] DescriptorAtomicGroup: TypeAlias = tuple[list[str], list[int], list[AtomicCounterMeta]] @@ -639,7 +641,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return self._time_provider() @staticmethod - def _no_max_tokens_output_floor( + def no_max_tokens_output_floor( min_configured_tpm_limit: int | None, ) -> int: """Output-budget floor used when the request omits max_tokens. @@ -670,7 +672,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): data: object, call_type: str | None, ) -> Mapping[str, object] | None: - contents = data.get("contents") if isinstance(data, dict) else None + contents: Final = data.get("contents") if isinstance(data, dict) else None if ( not isinstance(data, dict) or call_type not in GOOGLE_GENAI_NATIVE_CALL_TYPES @@ -679,7 +681,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return None from litellm.google_genai.adapters.transformation import GoogleGenAIAdapter - config = data.get("config") if "config" in data else data.get("generationConfig") + config: Final = data.get("config") if "config" in data else data.get("generationConfig") return GoogleGenAIAdapter().translate_generate_content_to_completion( model=data.get("model") if isinstance(data.get("model"), str) else "", contents=contents, @@ -696,15 +698,17 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): 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]) + config: Final = data.get("config") if "config" in data else data.get("generationConfig") + google_cap_values: Final = tuple( + int(raw_value) for field in ("maxOutputTokens", "max_output_tokens") - if isinstance(config, dict) and isinstance(config.get(field), (int, float, str)) + if isinstance(config, dict) + for raw_value in (config.get(field),) + if isinstance(raw_value, (int, float, str)) ) - return max(values, default=None) + return max(google_cap_values, default=None) if call_type in RESPONSES_API_CALL_TYPES: - value = data.get("max_output_tokens") + value: Final = data.get("max_output_tokens") if value is None: return None if not isinstance(value, (int, float, str)): @@ -712,13 +716,18 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return max(RESPONSES_API_MIN_OUTPUT_TOKENS, int(value)) if call_type in EMBEDDING_API_CALL_TYPES: return None - fields = ( + fields: Final = ( ("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) + output_cap_values: Final = tuple( + int(raw_value) + for field in fields + for raw_value in (data.get(field),) + if isinstance(raw_value, (int, float, str)) + ) + return max(output_cap_values, default=None) @classmethod def _has_explicit_output_cap(cls, data: object, call_type: str | None) -> bool: @@ -733,18 +742,18 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): def _get_output_candidate_count(data: object, call_type: str | None = None) -> int: if not isinstance(data, dict): return 1 - config = ( + config: Final = ( (data.get("config") if "config" in data else data.get("generationConfig")) if call_type in GOOGLE_GENAI_NATIVE_CALL_TYPES else None ) - candidate_values = ( + candidate_values: Final = ( 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 + candidate_count = 1 # rebind-ok: running maximum across candidate-count aliases for value in candidate_values: try: candidate_count = max(candidate_count, int(value or 1)) @@ -774,31 +783,34 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): """ 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) + base_capped_floor: Final = _PROXY_MaxParallelRequestsHandler_v3.no_max_tokens_output_floor(min_configured_limit) + capped_floor: Final = ( + max(base_capped_floor, RESPONSES_API_MIN_OUTPUT_TOKENS) + if call_type in RESPONSES_API_CALL_TYPES + else base_capped_floor + ) + baseline_floor: Final = DEFAULT_MAX_TOKENS_ESTIMATE // _TPM_FLOOR_FRACTION + is_embedding: Final = _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 - effective_cap = max(capped_floor, configured_output_tokens or 0) + effective_cap: Final = max(capped_floor, configured_output_tokens or 0) 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) + config_field: Final = "config" if "config" in data or "generationConfig" not in data else "generationConfig" + config: Final = 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 + data[config_field] = { # rebind-ok: routed request needs cap # mutable-ok: downstream needs dict **(config or {}), # mutable-ok: downstream native routing requires a mutable request config "maxOutputTokens": effective_cap, } return - cap_field = "max_output_tokens" if call_type in RESPONSES_API_CALL_TYPES else "max_tokens" - existing_cap = data.get(cap_field) + cap_field: Final = "max_output_tokens" if call_type in RESPONSES_API_CALL_TYPES else "max_tokens" + existing_cap: Final = data.get(cap_field) if existing_cap is None or effective_cap < existing_cap: - data[cap_field] = effective_cap + data[cap_field] = effective_cap # rebind-ok: downstream routing requires the bounded output cap def _estimate_tokens_for_request( self, @@ -878,69 +890,57 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): 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") + selected_fields: Final[tuple[object | None, object | None, object | None]] = ( + (None, None, estimable_data.get("input")) + if call_type in RESPONSES_API_CALL_TYPES or call_type in EMBEDDING_API_CALL_TYPES + else (None, estimable_data.get("prompt"), None) + if call_type in TEXT_COMPLETION_API_CALL_TYPES + else (estimable_data.get("messages"), None, None) + if call_type + else ( + estimable_data.get("messages"), + estimable_data.get("prompt"), + estimable_data.get("input"), + ) + ) + messages, prompt, input_text = selected_fields - 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 (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 + total_chars: Final = ( + len(get_str_from_messages(messages)) + if isinstance(messages, list) and messages + else len(prompt) + if isinstance(prompt, str) + else sum(len(str(item)) for item in prompt) + if isinstance(prompt, list) + else len(input_text) + if isinstance(input_text, str) + else sum(len(str(item)) for item in input_text) + if isinstance(input_text, list) + else 0 + ) estimated_input_tokens: Final = max(1, total_chars // DEFAULT_CHARS_PER_TOKEN) if total_chars > 0 else 0 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, is_embedding): - case (_, True): - max_tokens_estimate = 0 - case (mt, _) if mt is not None: - max_tokens_estimate = mt - case _ if total_chars == 0 and configured_output_tokens is None: - # Fully contentless request (no messages, prompt, or input). - # Don't apply the conservative output-budget floor here — it - # would over-reserve and could push small TPM limits into a - # false 429. The caller floors at 1 so backpressure still - # applies once the counter is at limit. - max_tokens_estimate = 0 - case _: - # No max_tokens specified — reserve at least the input size with a - # conservative floor so a stream of small concurrent requests can't - # collectively bypass the limit. Cap the floor by a fraction of - # 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 = 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 = ( - configured_output_tokens - if configured_output_tokens is not None - else max(estimated_input_tokens, output_floor) - ) + base_output_floor: Final = self.no_max_tokens_output_floor(min_configured_tpm_limit) + output_floor: Final = ( + max(base_output_floor, RESPONSES_API_MIN_OUTPUT_TOKENS) + if call_type in RESPONSES_API_CALL_TYPES + else base_output_floor + ) + max_tokens_estimate: Final = ( + 0 + if is_embedding or (explicit_max_tokens is None and total_chars == 0 and configured_output_tokens is None) + else explicit_max_tokens + if explicit_max_tokens is not None + else configured_output_tokens + if configured_output_tokens is not None + else max(estimated_input_tokens, output_floor) + ) - max_tokens_estimate *= self._get_output_candidate_count(data, call_type) - return estimated_input_tokens, max_tokens_estimate + return estimated_input_tokens, max_tokens_estimate * self._get_output_candidate_count(data, call_type) def _is_redis_cluster(self) -> bool: """ @@ -1755,7 +1755,10 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): mid-loop, refund applied increments and fall back to in-memory. """ if not descriptor_groups: - return RateLimitResponse(overall_code="OK", statuses=[]) + return RateLimitResponse( + overall_code="OK", + statuses=[], # mutable-ok: response contract requires a status list + ) applied: Final[list[list[AtomicCounterMeta]]] = [] statuses: Final[list[RateLimitStatus]] = [] raw: list[CacheCounterValue] @@ -1946,7 +1949,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ], ) descriptor_state.append( - { + { # mutable-ok: local atomic-counter state is updated during pass two "window_expired": window_expired, "current": current_counter, "window_start": str(now_int if window_expired else int(window_start)), @@ -2096,21 +2099,18 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): configured, or if the reservation failed), for the caller to stash for post-call reconciliation. """ - itpm_descriptors = [ # mutable-ok: atomic limiter API requires lists + itpm_descriptors: Final = [ # 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 + otpm_descriptors: Final = [ # 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( + itpm_response: Final = ( + 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 @@ -2118,12 +2118,15 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ], parent_otel_span=parent_otel_span, ) - if itpm_response["overall_code"] == "OVER_LIMIT": - return itpm_response, 0, 0 - itpm_reserved = estimated_input_tokens + if itpm_descriptors + else None + ) + if itpm_response is not None and itpm_response["overall_code"] == "OVER_LIMIT": + return itpm_response, 0, 0 + itpm_reserved: Final = estimated_input_tokens if itpm_response is not None else 0 if otpm_descriptors: - otpm_response = await self.atomic_check_and_increment_by_n( + otpm_response: Final = 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 @@ -2142,7 +2145,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): parent_otel_span=parent_otel_span, ) return otpm_response, 0, 0 - statuses = ( + statuses: Final = ( [ # mutable-ok: response contract uses a list *itpm_response["statuses"], *otpm_response["statuses"], @@ -2172,7 +2175,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): async def enforce_project_io_token_quota_for_frame( self, - user_api_key_dict: UserAPIKeyAuth, + user_api_key_dict: UserAPIKeyAuth | None, requested_model: str | None, estimated_input_tokens: int, estimated_output_tokens: int, @@ -2188,8 +2191,10 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): like the batch rate limiter -- this charges the estimate immediately and never refunds it. """ - descriptors: Final[list[RateLimitDescriptor]] = [] - self._add_project_io_token_rate_limit_descriptors_from_metadata( + if user_api_key_dict is None: + return + descriptors: Final[list[RateLimitDescriptor]] = [] # mutable-ok: descriptor helper appends in place + self.add_project_io_token_rate_limit_descriptors_from_metadata( user_api_key_dict=user_api_key_dict, requested_model=requested_model, descriptors=descriptors, @@ -2911,7 +2916,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) ) - def _add_project_io_token_rate_limit_descriptors_from_metadata( + def add_project_io_token_rate_limit_descriptors_from_metadata( self, user_api_key_dict: UserAPIKeyAuth, requested_model: str | None, @@ -2926,22 +2931,22 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): if requested_model is None or user_api_key_dict.project_id is None: return - itpm_limit_for_project_model = ( + itpm_limit_for_project_model: Final = ( 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 = ( + otpm_limit_for_project_model: Final = ( 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) + model_itpm_limit: Final = itpm_limit_for_project_model.get(requested_model) + model_otpm_limit: Final = 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}" + descriptor_value: Final = f"{user_api_key_dict.project_id}:{requested_model}" if model_itpm_limit is not None: descriptors.append( RateLimitDescriptor( @@ -3026,10 +3031,10 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): """ 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 + input_audio: Final = block.get("input_audio") + b64_data: Final = 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 + decoded_bytes: Final = len(b64_data) * 3 // 4 return max(decoded_bytes // _AUDIO_BYTES_PER_TOKEN, DEFAULT_AUDIO_TOKEN_ESTIMATE) return DEFAULT_AUDIO_TOKEN_ESTIMATE @@ -3042,17 +3047,15 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): """ 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 + return sum( + cls._estimate_audio_block_tokens(block) + for message in messages + if isinstance(message, dict) + for content in (message.get("content"),) + if isinstance(content, list) + for block in content + if isinstance(block, dict) and block.get("type") == "input_audio" + ) @staticmethod def _strip_audio_content_blocks(messages: object) -> object: @@ -3066,7 +3069,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): """ if not isinstance(messages, list): return messages - sanitized = [] # mutable-ok: token_counter requires a list of message dicts + sanitized: Final[list[object]] = [] # mutable-ok: token_counter requires a list of message dicts for message in messages: if not isinstance(message, dict): sanitized.append(message) @@ -3127,7 +3130,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): @classmethod def _contains_image_content(cls, value: object) -> bool: if isinstance(value, dict): - media_type = value.get("media_type") or value.get("mime_type") + media_type: Final = 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/")) @@ -3159,9 +3162,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): @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 + documents: Final = data.get("documents") + document_items: Final[Sequence[object]] = documents if isinstance(documents, list) else () # pyright: ignore[reportUnknownVariableType] # rerank documents are validated runtime JSON + input_parts: Final[tuple[object, ...]] = ( # pyright: ignore[reportUnknownVariableType] # list narrowing preserves unknown JSON element types data.get("query"), *document_items, ) @@ -3199,39 +3202,45 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): 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") + is_responses_request: Final = call_type in RESPONSES_API_CALL_TYPES + translated_request: Final = ( + None if is_responses_request else self._translate_google_genai_native_request(data, call_type) + ) + is_embedding_request: Final = self._is_embedding_request(data, call_type) + embedding_text: Final = data.get("input") if is_embedding_request else None + pretokenized_input_tokens: Final = ( + self._count_pretokenized_embedding_input(embedding_text) if is_embedding_request else None + ) + if pretokenized_input_tokens is not None: + return pretokenized_input_tokens - audio_token_estimate = self._estimate_audio_content_tokens(messages) - countable_messages = self._strip_audio_content_blocks(messages) if audio_token_estimate > 0 else messages + prompt: Final = data.get("prompt") + fallback_text: Final = prompt if prompt is not None else data.get("input") + selected_inputs: Final[tuple[object | None, object | None, object | None, object | None]] = ( + (self._responses_input_to_chat_messages(data), None, data.get("tools"), data.get("tool_choice")) + if is_responses_request + else ( + translated_request.get("messages"), + None, + translated_request.get("tools"), + translated_request.get("tool_choice"), + ) + if translated_request is not None + else (None, embedding_text, data.get("tools"), data.get("tool_choice")) + if is_embedding_request + else (None, self._rerank_input_to_text(data), data.get("tools"), data.get("tool_choice")) + if call_type in RERANK_API_CALL_TYPES + else (None, prompt, data.get("tools"), data.get("tool_choice")) + if call_type in TEXT_COMPLETION_API_CALL_TYPES + else (data.get("messages"), fallback_text, data.get("tools"), data.get("tool_choice")) + ) + messages, selected_text, countable_tools, countable_tool_choice = selected_inputs + + audio_token_estimate: Final = self._estimate_audio_content_tokens(messages) + countable_messages: Final = self._strip_audio_content_blocks(messages) if audio_token_estimate > 0 else messages try: - estimate = max( + estimate: Final = max( 0, int( token_counter( @@ -3245,7 +3254,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ), ) return estimate + audio_token_estimate - except Exception: # noqa: BLE001 - any tokenizer/model-resolution/transform failure degrades to the cheap estimate, never a 500 + except Exception: # noqa: BLE001 # tokenizer failures degrade to the cheap estimate 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) @@ -3272,14 +3281,14 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): """ if not isinstance(data, dict): return - stash = claim_request_stash_for_data(data) - io_token_descriptors = [ # mutable-ok: reservation API requires descriptor lists + stash: Final = claim_request_stash_for_data(data) + io_token_descriptors: Final = [ # 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 + configured_otpm_limits: Final = [ # mutable-ok: min calculation materializes validated limits int(v) for d in io_token_descriptors if d["key"] == PROJECT_OTPM_DESCRIPTOR_KEY @@ -3290,8 +3299,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ] 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 + min_configured_otpm_limit: Final = min(configured_otpm_limits) if configured_otpm_limits else None + configured_itpm_limits: Final = [ # mutable-ok: min calculation materializes validated limits int(v) for d in io_token_descriptors if d["key"] == PROJECT_ITPM_DESCRIPTOR_KEY @@ -3302,14 +3311,14 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ] if v is not None ] - min_configured_itpm_limit = min(configured_itpm_limits) if configured_itpm_limits else None + min_configured_itpm_limit: Final = min(configured_itpm_limits) if configured_itpm_limits else None - _, estimated_output_tokens = self._estimate_input_and_output_tokens( + _, raw_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 = ( + raw_estimated_input_tokens: Final = ( min_configured_itpm_limit if min_configured_itpm_limit is not None and ( @@ -3319,9 +3328,12 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) 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) + estimated_input_tokens: Final = max(raw_estimated_input_tokens, 1) + estimated_output_tokens: Final = ( + raw_estimated_output_tokens + if self._has_explicit_output_cap(data, call_type) + else max(raw_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 @@ -3353,7 +3365,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): parent_otel_span=user_api_key_dict.parent_otel_span, ) stash.reservation_released = True - acquisition = stash.parallel_slot + acquisition: Final = stash.parallel_slot if acquisition is not None: await self._release_parallel_request_slots( acquisition=acquisition, @@ -3367,9 +3379,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) if itpm_reserved > 0: - itpm_scopes = [ # mutable-ok: request stash freezes the collected scopes + itpm_scopes: Final = tuple( (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( @@ -3378,9 +3390,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): if "model_per_project_itpm" in counter_key ) if otpm_reserved > 0: - otpm_scopes = [ # mutable-ok: request stash freezes the collected scopes + otpm_scopes: Final = tuple( (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( @@ -3473,7 +3485,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): requested_model=requested_model, descriptors=descriptors, ) - self._add_project_io_token_rate_limit_descriptors_from_metadata( + self.add_project_io_token_rate_limit_descriptors_from_metadata( user_api_key_dict=user_api_key_dict, requested_model=requested_model, descriptors=descriptors, @@ -3546,8 +3558,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): # 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 + tpm_reservation_scopes: Sequence[tuple[str, str]] = () # rebind-ok: set after successful reservation + tpm_reservation_amount = 0 # rebind-ok: set after successful reservation if has_tpm_limits and self.tpm_reservation_enabled: min_configured_tpm_limit: Final = min(configured_tpm_limits) @@ -3633,8 +3645,10 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) is not None ) - tpm_reservation_scopes = tuple(stash.reserved_scopes) - tpm_reservation_amount = estimated_tokens + tpm_reservation_scopes = tuple( # rebind-ok: record successful reservation scopes + stash.reserved_scopes + ) + tpm_reservation_amount = estimated_tokens # rebind-ok: record successful reservation amount # Merge TPM statuses into the stored rate-limit response # so x-ratelimit-{key}-remaining-tokens / -limit-tokens @@ -3648,7 +3662,6 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): verbose_proxy_logger.debug( "TPM tokens reserved: %s for model %s", estimated_tokens, requested_model ) - await self._reserve_project_io_tokens_or_raise( descriptors=descriptors, data=data, @@ -3734,7 +3747,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return total_tokens @staticmethod - def _aggregate_only_total_tokens(usage: Usage | dict | None) -> int: + def _aggregate_only_total_tokens(usage: Usage | ResponseAPIUsage | Mapping[str, object] | None) -> int: """Total for usage that carries no input/output split, else 0. A source that can only report one number for the whole request (a @@ -3744,24 +3757,43 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): uncharged, which is how pass-through traffic slips past a TPM limit it is supposed to share. """ - if isinstance(usage, Usage): - prompt_tokens, completion_tokens, total_tokens = ( - usage.prompt_tokens or 0, - usage.completion_tokens or 0, - usage.total_tokens or 0, - ) - elif isinstance(usage, dict): - prompt_tokens, completion_tokens, total_tokens = ( - usage.get("prompt_tokens") or 0, - usage.get("completion_tokens") or 0, + if usage is None: + return 0 + token_counts: Final = ( + (usage.prompt_tokens or 0, usage.completion_tokens or 0, usage.total_tokens or 0) + if isinstance(usage, Usage) + else (usage.input_tokens or 0, usage.output_tokens or 0, usage.total_tokens or 0) + if isinstance(usage, ResponseAPIUsage) + else ( + usage.get("prompt_tokens") or usage.get("input_tokens") or 0, + usage.get("completion_tokens") or usage.get("output_tokens") or 0, usage.get("total_tokens") or 0, ) - else: - return 0 - if prompt_tokens or completion_tokens: + ) + prompt_tokens, completion_tokens, total_tokens = token_counts + if prompt_tokens or completion_tokens or not isinstance(total_tokens, int): return 0 return total_tokens + @staticmethod + def _response_usage( + response_obj: object, + ) -> Usage | ResponseAPIUsage | Mapping[str, object] | None: + if isinstance(response_obj, (Usage, ResponseAPIUsage)): + return response_obj + if isinstance( + response_obj, + (ModelResponse, EmbeddingResponse, TextCompletionResponse, BaseLiteLLMOpenAIResponseObject), + ): + usage: Final = getattr(response_obj, "usage", None) + return usage if isinstance(usage, (Usage, ResponseAPIUsage, dict)) else None + if isinstance(response_obj, dict): + nested_usage: Final = response_obj.get("usage") + if isinstance(nested_usage, (Usage, ResponseAPIUsage, dict)): + return nested_usage + return response_obj + return None + async def _execute_token_increment_script( self, pipeline_operations: list["RedisPipelineIncrementOperation"], @@ -3884,15 +3916,18 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): if self.window_guarded_token_increment_script is not None: try: await self.window_guarded_token_increment_script( - keys=[window_key, operation["key"]], - args=[ + keys=[ # mutable-ok: Redis script interface requires a key list + window_key, + operation["key"], + ], + args=[ # mutable-ok: Redis script interface requires an argument list expected_window_start, operation["increment_value"], operation["ttl"] or 0, ], ) continue - except Exception as e: # noqa: BLE001 - any Redis/Lua failure degrades to the plain increment fallback, never a 500 + except Exception as e: # noqa: BLE001 # Redis failures use the plain increment fallback verbose_proxy_logger.warning( "Window-guarded token adjustment failed for %s: %s", operation["key"], @@ -3978,16 +4013,16 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): 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 + rerank_tokens: Final = 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 + input_tokens: Final = rerank_tokens.get("input_tokens") or 0 # pyright: ignore[reportUnknownMemberType] # token fields are typed integers despite the generic get overload + output_tokens: Final = 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 + billed_units: Final = 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 + total_tokens: Final = 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 @@ -4003,68 +4038,57 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): 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) + rerank_usage: Final = 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 + usage: Final = self._response_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 + prompt_tokens: Final = usage.prompt_tokens or 0 + completion_tokens: Final = usage.completion_tokens or 0 + cached_tokens: Final = ( + getattr(usage.prompt_tokens_details, "cached_tokens", 0) or 0 + if usage.prompt_tokens_details is not None + else 0 ) - 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 - if prompt_tokens == 0 and completion_tokens == 0: - return 0, 0, False - return max(0, prompt_tokens - cached_tokens), completion_tokens, True + if isinstance(usage, ResponseAPIUsage): + response_input_tokens: Final = usage.input_tokens or 0 + response_output_tokens: Final = usage.output_tokens or 0 + response_cached_tokens: Final = ( + usage.input_tokens_details.cached_tokens or 0 if usage.input_tokens_details is not None else 0 + ) + if response_input_tokens == 0 and response_output_tokens == 0: + return 0, 0, False + return max(0, response_input_tokens - response_cached_tokens), response_output_tokens, True + + if isinstance(usage, Mapping): + raw_prompt_tokens: Final = usage.get("prompt_tokens") or usage.get("input_tokens") or 0 + raw_completion_tokens: Final = usage.get("completion_tokens") or usage.get("output_tokens") or 0 + mapped_prompt_tokens: Final = raw_prompt_tokens if isinstance(raw_prompt_tokens, int) else 0 + mapped_completion_tokens: Final = raw_completion_tokens if isinstance(raw_completion_tokens, int) else 0 + prompt_details: Final = usage.get("prompt_tokens_details") or usage.get("input_tokens_details") + raw_cached_tokens: Final = ( + (prompt_details.get("cached_tokens", 0) if isinstance(prompt_details, dict) else 0) + or usage.get("cache_read_input_tokens") + or 0 + ) + mapped_cached_tokens: Final = raw_cached_tokens if isinstance(raw_cached_tokens, int) else 0 + if mapped_prompt_tokens == 0 and mapped_completion_tokens == 0: + return 0, 0, False + return max(0, mapped_prompt_tokens - mapped_cached_tokens), mapped_completion_tokens, True + + return 0, 0, False def _build_io_token_reservation_ops( self, kwargs: object, response_obj: object, - ) -> list[RedisPipelineIncrementOperation] | tuple[ReservationAwareIncrementOperation, ...]: + ) -> Sequence[RedisPipelineIncrementOperation]: """ Reconcile project ITPM/OTPM reservations to actual usage on success: ITPM to billable input tokens, OTPM to actual completion tokens. @@ -4075,25 +4099,33 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): """ if not isinstance(kwargs, dict): return () - stash = get_request_stash_for_call(_call_id_from_callback_kwargs(kwargs)) + stash: Final = 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 + itpm_reserved: Final = stash.itpm_reserved_tokens + otpm_reserved: Final = 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 + response_usage: Final = self._resolve_io_token_reconcile_usage(response_obj) + combined_usage: Final = self._resolve_io_token_reconcile_usage(kwargs.get("combined_usage_object")) + aggregate_total: Final = self._aggregate_only_total_tokens( + self._response_usage(response_obj) + ) or self._aggregate_only_total_tokens(self._response_usage(kwargs.get("combined_usage_object"))) + + if not response_usage[2] and not combined_usage[2] and aggregate_total <= 0 and not stash.reservation_released: + return () + resolved_usage: Final = ( + response_usage + if response_usage[2] + else combined_usage + if combined_usage[2] + else (aggregate_total, aggregate_total, True) + if aggregate_total > 0 + else (itpm_reserved, otpm_reserved, False) + ) + billable_input, completion_tokens, _ = resolved_usage if stash.reservation_released or ( not stash.itpm_reserved_window_identities and not stash.otpm_reserved_window_identities @@ -4110,24 +4142,28 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): 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( + itpm_ops: Final[Sequence[ReservationAwareIncrementOperation]] = ( + 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( + if itpm_reserved > 0 + else () + ) + otpm_ops: Final[Sequence[ReservationAwareIncrementOperation]] = ( + 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, ) + if otpm_reserved > 0 + else () + ) return tuple((*itpm_ops, *otpm_ops)) def _collect_tpm_scope_targets( @@ -4522,14 +4558,11 @@ 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. - 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 + reserved_tokens, itpm_reserved, otpm_reserved = ( + (0, 0, 0) + if stash is None or stash.reservation_released + else (stash.reserved_tokens, stash.itpm_reserved_tokens, 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) diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index af561c9092b..e678fba2852 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -50,6 +50,16 @@ if TYPE_CHECKING: ) +class ProjectQuotaCallback(Protocol): + async def enforce_project_io_token_quota_for_frame( + self, + user_api_key_dict: UserAPIKeyAuth | None, + requested_model: str | None, + estimated_input_tokens: int, + estimated_output_tokens: int, + ) -> None: ... + + @lru_cache(maxsize=1) def _get_openai_response_types(): from litellm.types.llms import openai as openai_types @@ -1345,11 +1355,19 @@ def _extract_frame_quota_estimate_inputs(msg_obj: Mapping[str, object]) -> tuple """ nested: Final = msg_obj.get("response") params: Final[Mapping[str, object]] = ( - nested if _is_json_object(nested) and nested else {k: v for k, v in msg_obj.items() if k != "type"} + nested + if _is_json_object(nested) and nested + else MappingProxyType( # mutable-ok: immediately frozen filtered frame + {k: v for k, v in msg_obj.items() if k != "type"} + ) ) - text_parts: list[str] = [] # mutable-ok: local accumulator built in one pass, not shared - - def _collect_text(value: object) -> None: + text_parts: Final[list[str]] = [] # mutable-ok: local accumulator built in one pass, not shared + pending: Final[list[object]] = [ # mutable-ok: explicit worklist avoids recursion + params.get("input"), + params.get("instructions"), + ] + while pending: + value = pending.pop() if isinstance(value, str): text_parts.append(value) elif _is_json_array(value): @@ -1357,20 +1375,17 @@ def _extract_frame_quota_estimate_inputs(msg_obj: Mapping[str, object]) -> tuple if isinstance(item, str): text_parts.append(item) elif _is_json_object(item): - _collect_text(item.get("content")) - _collect_text(item.get("text")) - - _collect_text(params.get("input")) - _collect_text(params.get("instructions")) + pending.append(item.get("content")) + pending.append(item.get("text")) total_chars: Final = sum(len(part) for part in text_parts) estimated_input_tokens: Final = max(1, total_chars // _FRAME_CHARS_PER_TOKEN_ESTIMATE) if total_chars else 0 - max_output_tokens = params.get("max_output_tokens") + max_output_tokens: Final = params.get("max_output_tokens") return estimated_input_tokens, max_output_tokens if isinstance(max_output_tokens, int) else None async def _enforce_frame_project_quota( - quota_callbacks: Sequence[Any], + quota_callbacks: Sequence[ProjectQuotaCallback], user_api_key_dict: UserAPIKeyAuth | None, model: str | None, raw_message: str, @@ -1387,7 +1402,7 @@ async def _enforce_frame_project_quota( if not _is_json_object(msg_obj) or msg_obj.get("type") != "response.create": return estimated_input_tokens, explicit_max_output_tokens = _extract_frame_quota_estimate_inputs(msg_obj) - estimated_output_tokens = ( + estimated_output_tokens: Final = ( explicit_max_output_tokens if explicit_max_output_tokens is not None else _FRAME_NO_MAX_OUTPUT_TOKENS_FLOOR ) for callback in quota_callbacks: @@ -1433,7 +1448,7 @@ class ResponsesWebSocketStreaming: first_message: str | None = None, guardrail_callbacks: list[Any] | None = None, output_guardrail_callbacks: list[PresidioGuardrailCallback] | None = None, - quota_callbacks: list[Any] | None = None, + quota_callbacks: Sequence[ProjectQuotaCallback] | None = None, authorized_model: str | None = None, ): self.websocket = websocket @@ -1446,7 +1461,7 @@ class ResponsesWebSocketStreaming: self.first_message = first_message self.guardrail_callbacks: list[Any] = guardrail_callbacks or [] self.output_guardrail_callbacks: list[PresidioGuardrailCallback] = output_guardrail_callbacks or [] - self.quota_callbacks: list[Any] = quota_callbacks or [] + self.quota_callbacks: tuple[ProjectQuotaCallback, ...] = tuple(quota_callbacks) if quota_callbacks else () # Model name authorized at connection time; enforced on every # response.create frame to prevent deployment-substitution attacks. self.authorized_model: str | None = authorized_model @@ -1870,9 +1885,17 @@ class ResponsesWebSocketStreaming: except RateLimitError as e: try: await self.websocket.send_text( - json.dumps({"type": "error", "error": {"type": "rate_limit_exceeded", "message": str(e)}}) + json.dumps( # mutable-ok: WebSocket wire payload requires JSON objects + { # mutable-ok: WebSocket wire payload requires JSON objects + "type": "error", + "error": { # mutable-ok: nested WebSocket error object + "type": "rate_limit_exceeded", + "message": str(e), + }, + } + ) ) - except Exception: # noqa: BLE001, S110 - best-effort notification, client may already be gone + except Exception: # noqa: BLE001, S110 # client may already be gone pass return False return True @@ -1969,7 +1992,7 @@ class ManagedResponsesWebSocketHandler: timeout: float | None = None, custom_llm_provider: str | None = None, first_message: str | None = None, - quota_callbacks: list[Any] | None = None, + quota_callbacks: Sequence[ProjectQuotaCallback] | None = None, **kwargs: object, ) -> None: self.websocket = websocket @@ -1986,7 +2009,7 @@ class ManagedResponsesWebSocketHandler: self.custom_llm_provider = custom_llm_provider self._connection_provider = self._resolve_provider(model) or custom_llm_provider self.first_message = first_message - self.quota_callbacks: list[Any] = quota_callbacks or [] + self.quota_callbacks: tuple[ProjectQuotaCallback, ...] = tuple(quota_callbacks) if quota_callbacks else () # Carry through safe pass-through kwargs (e.g. extra_headers) self.extra_kwargs: dict[str, object] = {k: v for k, v in kwargs.items() if k not in _MANAGED_WS_SKIP_KWARGS} # In-memory session history: response_id → full accumulated message list. diff --git a/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py b/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py index 66fc84ab1e0..e55185cfa67 100644 --- a/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py +++ b/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py @@ -3582,18 +3582,24 @@ async def test_streaming_combined_usage_reconciles_project_io_reservations( assert [operation["increment_value"] for operation in otpm_adjustments] == [-45] -def test_aggregate_only_combined_usage_keeps_project_io_reservations(rate_limiter): +def test_aggregate_only_combined_usage_reconciles_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")} ) + stash.otpm_reserved_tokens = 80 + stash.otpm_reserved_scopes = frozenset( + {(PROJECT_OTPM_DESCRIPTOR_KEY, "project:model")} + ) kwargs = { "combined_usage_object": Usage(total_tokens=55), } - assert handler._build_io_token_reservation_ops(kwargs, object()) == () + operations = handler._build_io_token_reservation_ops(kwargs, object()) + + assert [operation["increment_value"] for operation in operations] == [-45, -25] def test_raw_split_usage_dict_reconciles_project_io_tokens(rate_limiter): diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 894d99c92e0..21dc0bc2791 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,9 +1,9 @@ { "LIT001": { - "limit": 22941 + "limit": 22936 }, "LIT002": { - "limit": 27139 + "limit": 27133 }, "LIT003": { "limit": 269 @@ -27,10 +27,10 @@ "limit": 0 }, "LIT010": { - "limit": 16716 + "limit": 16701 }, "LIT011": { - "limit": 5596 + "limit": 5595 }, "LIT012": { "limit": 4519 From a6dc447470f41baea1383b30fcee4c10cf6bdc04 Mon Sep 17 00:00:00 2001 From: Shivi Jain Date: Sun, 16 Aug 2026 16:24:23 +0530 Subject: [PATCH 09/83] fix(proxy): stop Responses batch rows from bypassing project OTPM Embeddings rows were identified by body shape (has `input`, no `messages`/`prompt`), which also matches a `/v1/responses` batch row and reserved zero output tokens for it -- letting a project caller run large Responses generations against a quota-limited model without consuming OTPM. Classify embeddings by the row's own `url` instead, and read `max_output_tokens` as a Responses output cap alongside `max_tokens`/`max_completion_tokens`. Co-authored-by: Cursor --- litellm/proxy/hooks/batch_rate_limiter.py | 30 ++++++- .../proxy/hooks/test_batch_file_validation.py | 80 +++++++++++++++++++ 2 files changed, 106 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index 3b53d792ce6..4052a9b04fa 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -395,18 +395,40 @@ class _PROXY_BatchRateLimiter(CustomLogger): Batch completion never reconciles actual usage back into the rate limiter, so this pre-call estimate is the only OTPM enforcement a batch gets. Mirrors the real-time no-``max_tokens`` floor so a row - that omits ``max_tokens`` can't be used to bypass OTPM the way an + that omits an output cap can't be used to bypass OTPM the way an unbounded streaming request could. + + Embeddings rows are identified by the row's own ``url`` (the OpenAI + batch schema puts the target route there, e.g. ``/v1/embeddings``), + never by body shape: a `/v1/responses` row also carries `body.input` + with no `messages`/`prompt`, so guessing from body shape alone would + misclassify a token-generating Responses row as a zero-output + embeddings row and let it skip the OTPM reservation entirely. """ + url: Final = entry.get("url") + if isinstance(url, str) and "embeddings" in url: + return 0 # embeddings: no output tokens raw_body: Final = entry.get("body") body: Final[Mapping[str, object]] = ( MappingProxyType(_BATCH_BODY_ADAPTER.validate_python(raw_body)) if isinstance(raw_body, Mapping) else MappingProxyType({}) # mutable-ok: immediately frozen empty fallback ) - if body.get("input") is not None and body.get("messages") is None and body.get("prompt") is None: - return 0 # embeddings: no output tokens - explicit_cap: Final = body.get("max_tokens", body.get("max_completion_tokens")) + # `max_tokens`/`max_completion_tokens` cap chat completions; `/v1/responses` + # rows cap output with `max_output_tokens` instead -- omitting it here + # would fall through to the floor estimate for every capped Responses row. + explicit_cap: Final = next( + ( + v + for v in ( + body.get("max_tokens"), + body.get("max_completion_tokens"), + body.get("max_output_tokens"), + ) + if v is not None + ), + None, + ) if explicit_cap is not None: try: return max(0, int(explicit_cap)) diff --git a/tests/test_litellm/proxy/hooks/test_batch_file_validation.py b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py index ee37cf7d302..e5211973ec3 100644 --- a/tests/test_litellm/proxy/hooks/test_batch_file_validation.py +++ b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py @@ -1980,3 +1980,83 @@ async def test_count_input_file_usage_collects_models_after_malformed_line(): ) assert exc.value.status_code == 403 + + +# --------------------------------------------------------------------------- +# VERIA-Low regression: Responses batch rows must not bypass project OTPM +# --------------------------------------------------------------------------- + + +def _output_estimator(): + """A `_PROXY_BatchRateLimiter` whose output-token floor is observable: + the no-`max_tokens` floor mock returns a distinctive sentinel so tests can + tell "floor was used" apart from "an explicit cap was read".""" + from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + + limiter = MagicMock() + limiter.no_max_tokens_output_floor.return_value = 999 + return _PROXY_BatchRateLimiter( + internal_usage_cache=MagicMock(), + parallel_request_limiter=limiter, + ) + + +def test_estimate_entry_output_tokens_zero_for_embeddings_url(): + """A real `/v1/embeddings` row reserves zero output tokens.""" + rate_limiter = _output_estimator() + entry = { + "url": "/v1/embeddings", + "body": {"model": "text-embedding-3-small", "input": "hello world"}, + } + + assert rate_limiter._estimate_entry_output_tokens(entry, None) == 0 + + +def test_estimate_entry_output_tokens_does_not_zero_responses_row_with_input(): + """Pre-fix: a `/v1/responses` row carries `body.input` with no `messages`/ + `prompt`, so the old body-shape heuristic misclassified it as embeddings + and reserved zero output tokens -- a project caller could submit large + Responses generations against a quota-limited model without consuming + OTPM. The row's own `url` (not body shape) must decide this.""" + rate_limiter = _output_estimator() + entry = { + "url": "/v1/responses", + "body": {"model": "gpt-4o", "input": "write me an essay"}, + } + + # No explicit cap on the row, so it must fall back to the no-max-tokens + # floor -- never straight to zero. + assert rate_limiter._estimate_entry_output_tokens(entry, None) == 999 + rate_limiter.parallel_request_limiter.no_max_tokens_output_floor.assert_called_once_with(None) + + +def test_estimate_entry_output_tokens_uses_max_output_tokens_for_responses(): + """`/v1/responses` caps output with `max_output_tokens`, not `max_tokens`/ + `max_completion_tokens`. Pre-fix this field was never inspected, so a + capped Responses row still fell through to the (possibly larger) floor + estimate instead of the caller's own declared cap.""" + rate_limiter = _output_estimator() + entry = { + "url": "/v1/responses", + "body": {"model": "gpt-4o", "input": "hi", "max_output_tokens": 123}, + } + + assert rate_limiter._estimate_entry_output_tokens(entry, None) == 123 + rate_limiter.parallel_request_limiter.no_max_tokens_output_floor.assert_not_called() + + +def test_estimate_entry_output_tokens_prefers_max_tokens_over_max_output_tokens(): + """When a row somehow carries both fields, the chat-style cap wins first -- + `max_output_tokens` is only consulted once the chat-style caps are absent.""" + rate_limiter = _output_estimator() + entry = { + "url": "/v1/chat/completions", + "body": { + "model": "gpt-4o", + "messages": [], + "max_tokens": 50, + "max_output_tokens": 500, + }, + } + + assert rate_limiter._estimate_entry_output_tokens(entry, None) == 50 From 89e563d3dabdede841f48a717c74fb21794223a8 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 18 Aug 2026 13:37:03 -0700 Subject: [PATCH 10/83] fix(tinyfish): keep hidden-params header stashing within lint budgets --- litellm/llms/tinyfish/search/transformation.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/litellm/llms/tinyfish/search/transformation.py b/litellm/llms/tinyfish/search/transformation.py index ae11b4bcc57..b688dc2cd01 100644 --- a/litellm/llms/tinyfish/search/transformation.py +++ b/litellm/llms/tinyfish/search/transformation.py @@ -240,11 +240,11 @@ class TinyfishSearchConfig(BaseSearchConfig): _emit_parameter_warnings(parsed) max_results: Final = self._caller_max_results or _TINYFISH_RESULT_CAP - # Truncate in place so all pydantic-populated fields survive — declared and extras. - parsed.results = list(parsed.results[:max_results]) - raw_headers = dict(raw_response.headers) - parsed._hidden_params["headers"] = raw_headers - parsed._hidden_params["additional_headers"] = process_response_headers(raw_headers) + parsed.results = parsed.results[:max_results] + raw_headers: Final = dict(raw_response.headers) + hidden: Final = parsed._hidden_params # pyright: ignore[reportPrivateUsage] # sole hidden-params channel + hidden["headers"] = raw_headers + hidden["additional_headers"] = process_response_headers(raw_headers) return parsed def _wrap_error( From c1c23bf39ad82e3f494282bd9e75864e0f867039 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 18 Aug 2026 13:51:47 -0700 Subject: [PATCH 11/83] fix(proxy): reserve measured input tokens for multimodal project ITPM Image, file, video, and previous_response_id requests reserved the whole project ITPM limit up front, so any window with existing usage rejected them and one in-flight multimodal request blocked the entire project. Reserve the token_counter estimate instead, like every other request; post-call reconciliation already charges actual usage. --- basedpyright-code-budget.json | 6 +- .../hooks/parallel_request_limiter_v3.py | 65 +------ .../proxy/hooks/test_tpm_concurrent.py | 165 ++++++------------ type-discipline-budget.json | 4 +- 4 files changed, 63 insertions(+), 177 deletions(-) diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index 5eb72a5c243..7ac52925459 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -57,7 +57,7 @@ "limit": 5681 }, "reportMissingTypeArgument": { - "limit": 15607 + "limit": 15606 }, "reportMissingTypeStubs": { "limit": 40 @@ -108,7 +108,7 @@ "limit": 39154 }, "reportUnknownParameterType": { - "limit": 19946 + "limit": 19945 }, "reportUnknownVariableType": { "limit": 30772 @@ -132,7 +132,7 @@ "limit": 27 }, "reportUnusedClass": { - "limit": 23 + "limit": 22 }, "reportUnusedFunction": { "limit": 139 diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 80d58abee7d..751b301a93d 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -3119,47 +3119,6 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): 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: Final = 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): @@ -3312,33 +3271,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): if v is not None ] min_configured_otpm_limit: Final = min(configured_otpm_limits) if configured_otpm_limits else None - configured_itpm_limits: Final = [ # 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: Final = min(configured_itpm_limits) if configured_itpm_limits else None - _, raw_estimated_output_tokens = self._estimate_input_and_output_tokens( data=data, min_configured_tpm_limit=min_configured_otpm_limit, call_type=call_type, ) - raw_estimated_input_tokens: Final = ( - 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) + raw_estimated_input_tokens: Final = self._estimate_precise_input_tokens( + data=data, model=requested_model, call_type=call_type ) estimated_input_tokens: Final = max(raw_estimated_input_tokens, 1) estimated_output_tokens: Final = ( diff --git a/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py b/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py index e04ea66db5e..fd7e0626af9 100644 --- a/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py +++ b/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py @@ -2386,13 +2386,34 @@ async def test_itpm_rejects_large_audio_payload_that_would_pass_flat_estimate( { "role": "user", "content": [ + {"type": "text", "text": "describe this"}, { "type": "image_url", "image_url": { "url": "https://example.com/high-resolution.png", "detail": "high", }, - } + }, + ], + } + ] + }, + ), + ( + "acompletion", + { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "summarize this"}, + { + "type": "file", + "file": { + "filename": "document.pdf", + "file_data": "data:application/pdf;base64,dGVzdA==", + }, + }, ], } ] @@ -2405,29 +2426,43 @@ async def test_itpm_rejects_large_audio_payload_that_would_pass_flat_estimate( { "role": "user", "content": [ + {"type": "input_text", "text": "describe this"}, { "type": "input_image", "image_url": "https://example.com/high-resolution.png", "detail": "high", - } + }, ], } ] }, ), + ( + "aresponses", + {"input": "continue", "previous_response_id": "resp-123"}, + ), ], ) -async def test_image_content_reserves_full_project_itpm( +async def test_multimodal_requests_reserve_measured_project_itpm_not_full_limit( rate_limiter, call_type, request_data, ): + """ + Regression: image, file, and previous_response_id requests used to + reserve the project's whole ITPM limit up front. Because the atomic + check is ``current + increment > limit``, that made every such request + 429 as soon as the window carried any usage at all and, while in + flight, blocked every other request for the same project + model. They + now reserve the token_counter estimate like everything else, so two + multimodal requests fit in the same window. + """ handler, cache = rate_limiter model = "bedrock_mantle/claude-opus" - project_itpm_limit = 1_000 + project_itpm_limit = 10_000 user_api_key_dict = UserAPIKeyAuth( - api_key=hash_token("sk-high-resolution-image"), - project_id="project-high-resolution-image", + api_key=hash_token("sk-multimodal-measured"), + project_id="project-multimodal-measured", project_metadata={"model_itpm_limit": {model: project_itpm_limit}}, ) @@ -2437,10 +2472,22 @@ async def test_image_content_reserves_full_project_itpm( data={"model": model, **request_data}, call_type=call_type, ) + first_stash = get_request_stash() + assert first_stash is not None + first_reservation = first_stash.itpm_reserved_tokens + assert 0 < first_reservation < project_itpm_limit // 2 - stash = get_request_stash() - assert stash is not None - assert stash.itpm_reserved_tokens == project_itpm_limit + _request_stash.set(None) + 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, + ) + second_stash = get_request_stash() + assert second_stash is not None + assert second_stash is not first_stash + assert second_stash.itpm_reserved_tokens == first_reservation @pytest.mark.asyncio @@ -3191,95 +3238,6 @@ def test_anthropic_messages_usage_reconciles_split_project_quota(rate_limiter): 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", @@ -3469,18 +3427,7 @@ def test_split_quota_multimodal_guards_handle_non_mapping_inputs(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 diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 88a807e39b3..7a099462256 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,6 +1,6 @@ { "LIT001": { - "limit": 22897 + "limit": 22895 }, "LIT002": { "limit": 26889 @@ -27,7 +27,7 @@ "limit": 0 }, "LIT010": { - "limit": 16687 + "limit": 16661 }, "LIT011": { "limit": 5590 From d2fbaff2c948e5beb0fe02623640e1fc9c21cf69 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 18 Aug 2026 14:05:28 -0700 Subject: [PATCH 12/83] fix(proxy): record estimated input tokens in spend logs for dispatched failed requests Failure rows in the spend log only carried token counts when a broken stream stashed recovered partial usage; non-stream requests that reached the provider and then failed (timeouts, provider 4xx/5xx) logged 0/0/0 even though the provider billed the input tokens. Estimate the input side in post_call_failure_hook with the same tokenizer fallback interrupted streams use, gated to requests that were actually dispatched (first_api_call_start_time set and no litellm_no_upstream_llm_call marker), and pin response_cost to 0.0 so failed requests never bill spend. Recovered partial-stream usage still wins over the estimate. --- litellm/proxy/utils.py | 67 +++++++-- tests/test_litellm/proxy/test_proxy_utils.py | 139 +++++++++++++++++++ 2 files changed, 197 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 2ad7180bd5f..3457ae0f352 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -40,7 +40,7 @@ from litellm.proxy._types import ( from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error from litellm.types.guardrails import GuardrailEventHooks from litellm.types.proxy.model_listing import ModelInfoResponse -from litellm.types.utils import CallTypes, CallTypesLiteral, ModelInfo +from litellm.types.utils import CallTypes, CallTypesLiteral, ModelInfo, Usage try: from litellm_enterprise.enterprise_callbacks.send_emails.base_email import ( @@ -403,6 +403,52 @@ def _exception_changes_request_flow(exc: BaseException) -> bool: return isinstance(exc, (SensitiveDataRouteException, ModifyResponseException)) +def _count_request_input_tokens(model: str, request_input: object) -> int: + if isinstance(request_input, str): + return litellm.token_counter(model=model, text=request_input) + if not isinstance(request_input, list) or not request_input: + return 0 + text_entries: Final = tuple(entry for entry in request_input if isinstance(entry, str)) + if len(text_entries) == len(request_input): + return litellm.token_counter(model=model, text="".join(text_entries)) + return litellm.token_counter(model=model, messages=request_input) + + +def _estimate_dispatched_failure_usage(model: str, request_input: object) -> Usage | None: + """A request that failed after dispatch consumed provider-billed input + tokens, but no provider usage ever came back. Estimate the input side with + the same tokenizer fallback interrupted streams use, so the spend log's + failure row records what was sent instead of zero.""" + try: + input_tokens: Final = _count_request_input_tokens(model=model, request_input=request_input) + except Exception: + return None + if input_tokens <= 0: + return None + return Usage(prompt_tokens=input_tokens, completion_tokens=0, total_tokens=input_tokens) + + +def _failure_usage_to_lift(model_call_details: Mapping[str, object], dispatched: bool) -> tuple[object, object] | None: + """A stream that broke mid-flight still billed the provider for the chunks + already delivered; the streaming handler stashes that recovered usage and + cost in model_call_details, so prefer it. Otherwise a request that was + dispatched to a provider and failed without upstream usage gets an + estimated input-side Usage with zero cost. Returns the + (combined_usage_object, response_cost) pair to lift, or None.""" + recovered_usage: Final = model_call_details.get("combined_usage_object") + if recovered_usage is not None: + return recovered_usage, model_call_details.get("response_cost") + if not dispatched or model_call_details.get(LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL): + return None + estimated_usage: Final = _estimate_dispatched_failure_usage( + model=str(model_call_details.get("model") or ""), + request_input=model_call_details.get("messages"), + ) + if estimated_usage is None: + return None + return estimated_usage, 0.0 + + @dataclass(frozen=True) class _CallbackCapabilities: """Cached per-hook capability flags derived from ``litellm.callbacks``. @@ -2190,15 +2236,18 @@ class ProxyLogging: if _first_handoff is not None: request_data["first_api_call_start_time"] = _first_handoff - # A stream that broke mid-flight still billed the provider for the - # chunks already delivered; the streaming handler stashes that - # recovered usage and cost here. Lift them onto request_data so the + # Lift recovered partial-stream usage, or an estimated input-side + # usage for a dispatched failure, onto request_data so the # failure-path spend callbacks (which run after the logging object - # is popped) record the real partial spend instead of zero. - _recovered_usage: Final = _model_call_details.get("combined_usage_object") - if _recovered_usage is not None: - request_data["combined_usage_object"] = _recovered_usage - request_data["response_cost"] = _model_call_details.get("response_cost") + # is popped) record real token counts instead of zero. + _usage_to_lift: Final = _failure_usage_to_lift( + model_call_details=_model_call_details, + dispatched=_first_handoff is not None, + ) + if _usage_to_lift is not None: + _lifted_usage, _lifted_cost = _usage_to_lift + request_data["combined_usage_object"] = _lifted_usage + request_data["response_cost"] = _lifted_cost # Remove before callbacks iterate — not serialisable request_data.pop("litellm_logging_obj", None) diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index 1504c3c3103..70baf157ab9 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -478,6 +478,145 @@ class TestPostCallFailureHookLiftsRecoveredPartialSpend: assert "response_cost" not in request_data +class TestPostCallFailureHookEstimatesDispatchedInputTokens: + """A non-stream request that failed after dispatch (timeout, provider + error) consumed provider-billed input tokens but recovered no usage. + post_call_failure_hook must estimate the input side onto request_data so + the spend log's failure row records what was sent instead of zero, while + never charging spend for the failure (LIT-5690). + """ + + async def _run(self, request_data): + from unittest.mock import AsyncMock, patch + + from litellm.proxy._types import UserAPIKeyAuth + + proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + proxy_logging_obj.alert_types = [] + with patch.object(proxy_logging_obj, "update_request_status", new=AsyncMock()): + await proxy_logging_obj.post_call_failure_hook( + request_data=request_data, + original_exception=Exception("boom"), + user_api_key_dict=UserAPIKeyAuth(), + ) + + def _logging_obj(self, model_call_details): + logging_obj = MagicMock() + logging_obj.model_call_details = model_call_details + return logging_obj + + @pytest.mark.asyncio + async def test_dispatched_failure_estimates_input_tokens_with_zero_cost(self): + from datetime import datetime + + from litellm.types.utils import Usage + + request_data = { + "litellm_logging_obj": self._logging_obj( + { + "first_api_call_start_time": datetime.now(), + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "count these input tokens please"}], + } + ), + "metadata": {}, + "response_cost": 123.0, + } + await self._run(request_data) + + estimated = request_data["combined_usage_object"] + assert isinstance(estimated, Usage) + assert estimated.prompt_tokens > 0 + assert estimated.completion_tokens == 0 + assert estimated.total_tokens == estimated.prompt_tokens + assert request_data["response_cost"] == 0.0 + + @pytest.mark.asyncio + async def test_failure_before_dispatch_stays_zero(self): + request_data = { + "litellm_logging_obj": self._logging_obj( + { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "never dispatched"}], + } + ), + "metadata": {}, + } + await self._run(request_data) + + assert "combined_usage_object" not in request_data + assert "response_cost" not in request_data + + @pytest.mark.asyncio + async def test_proxy_only_error_never_dispatched_stays_zero(self): + from datetime import datetime + + from litellm.constants import LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL + + request_data = { + "litellm_logging_obj": self._logging_obj( + { + "first_api_call_start_time": datetime.now(), + "model": "no-such-model", + "messages": [{"role": "user", "content": "hi"}], + LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL: True, + } + ), + "metadata": {}, + } + await self._run(request_data) + + assert "combined_usage_object" not in request_data + assert "response_cost" not in request_data + + @pytest.mark.asyncio + async def test_recovered_partial_usage_wins_over_estimate(self): + from datetime import datetime + + from litellm.types.utils import Usage + + recovered_usage = Usage(prompt_tokens=30, completion_tokens=7, total_tokens=37) + request_data = { + "litellm_logging_obj": self._logging_obj( + { + "first_api_call_start_time": datetime.now(), + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "mid-stream failure"}], + "combined_usage_object": recovered_usage, + "response_cost": 3.5e-05, + } + ), + "metadata": {}, + } + await self._run(request_data) + + assert request_data["combined_usage_object"] is recovered_usage + assert request_data["response_cost"] == 3.5e-05 + + @pytest.mark.asyncio + async def test_dispatched_failure_with_text_completion_prompt(self): + from datetime import datetime + + from litellm.types.utils import Usage + + request_data = { + "litellm_logging_obj": self._logging_obj( + { + "first_api_call_start_time": datetime.now(), + "model": "gpt-3.5-turbo", + "messages": "a plain text-completion prompt string", + } + ), + "metadata": {}, + } + await self._run(request_data) + + estimated = request_data["combined_usage_object"] + assert isinstance(estimated, Usage) + assert estimated.prompt_tokens > 0 + assert estimated.completion_tokens == 0 + + from typing import cast import litellm From 803113c63af3c543473c42c91bb1846974f782ab Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 18 Aug 2026 14:21:47 -0700 Subject: [PATCH 13/83] fix(proxy): estimate failed-request input tokens on /v1/messages and count system prompts The Anthropic messages endpoint's exception handler passed the raw request body dict to the failure hook, but request setup had already replaced the processor's dict with one carrying the logging object, so failure rows for /v1/messages never lifted recovered or estimated usage. Pass the processor's dict instead. The input-side estimate only counted the messages list, missing the Anthropic top-level system prompt (string or text-block list) and the Responses API instructions field, which live in optional_params. Count them too. --- .../proxy/anthropic_endpoints/endpoints.py | 4 +- litellm/proxy/utils.py | 42 ++++++++++-- .../anthropic_endpoints/test_endpoints.py | 35 ++++++++++ tests/test_litellm/proxy/test_proxy_utils.py | 68 +++++++++++++++++++ 4 files changed, 140 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/anthropic_endpoints/endpoints.py b/litellm/proxy/anthropic_endpoints/endpoints.py index a48ef0f08bb..f742965ade2 100644 --- a/litellm/proxy/anthropic_endpoints/endpoints.py +++ b/litellm/proxy/anthropic_endpoints/endpoints.py @@ -179,7 +179,7 @@ async def anthropic_response( await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, - request_data=data, + request_data=base_llm_response_processor.data, ) body: Final = AnthropicExceptionMapping.transform_to_anthropic_error( status_code=e.status_code, @@ -189,7 +189,7 @@ async def anthropic_response( return JSONResponse(status_code=e.status_code, content=body) except Exception as e: await proxy_logging_obj.post_call_failure_hook( - user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data + user_api_key_dict=user_api_key_dict, original_exception=e, request_data=base_llm_response_processor.data ) verbose_proxy_logger.exception("litellm.proxy.proxy_server.anthropic_response(): Exception occured - %s", e) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 3457ae0f352..5bfc1c2d1e1 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -403,24 +403,45 @@ def _exception_changes_request_flow(exc: BaseException) -> bool: return isinstance(exc, (SensitiveDataRouteException, ModifyResponseException)) -def _count_request_input_tokens(model: str, request_input: object) -> int: +def _prompt_block_text(block: object) -> str: + if isinstance(block, str): + return block + if not isinstance(block, dict): + return "" + block_text: Final = block.get("text") + return block_text if isinstance(block_text, str) else "" + + +def _system_prompt_text(system_input: object) -> str: + if isinstance(system_input, str): + return system_input + if not isinstance(system_input, list): + return "" + return "".join(_prompt_block_text(block) for block in system_input) + + +def _count_request_input_tokens(model: str, request_input: object, system_input: object) -> int: + system_text: Final = _system_prompt_text(system_input) + system_tokens: Final = litellm.token_counter(model=model, text=system_text) if system_text else 0 if isinstance(request_input, str): - return litellm.token_counter(model=model, text=request_input) + return system_tokens + litellm.token_counter(model=model, text=request_input) if not isinstance(request_input, list) or not request_input: - return 0 + return system_tokens text_entries: Final = tuple(entry for entry in request_input if isinstance(entry, str)) if len(text_entries) == len(request_input): - return litellm.token_counter(model=model, text="".join(text_entries)) - return litellm.token_counter(model=model, messages=request_input) + return system_tokens + litellm.token_counter(model=model, text="".join(text_entries)) + return system_tokens + litellm.token_counter(model=model, messages=request_input) -def _estimate_dispatched_failure_usage(model: str, request_input: object) -> Usage | None: +def _estimate_dispatched_failure_usage(model: str, request_input: object, system_input: object) -> Usage | None: """A request that failed after dispatch consumed provider-billed input tokens, but no provider usage ever came back. Estimate the input side with the same tokenizer fallback interrupted streams use, so the spend log's failure row records what was sent instead of zero.""" try: - input_tokens: Final = _count_request_input_tokens(model=model, request_input=request_input) + input_tokens: Final = _count_request_input_tokens( + model=model, request_input=request_input, system_input=system_input + ) except Exception: return None if input_tokens <= 0: @@ -440,9 +461,16 @@ def _failure_usage_to_lift(model_call_details: Mapping[str, object], dispatched: return recovered_usage, model_call_details.get("response_cost") if not dispatched or model_call_details.get(LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL): return None + optional_params: Final = model_call_details.get("optional_params") + system_input: Final = ( + (optional_params.get("system") or optional_params.get("instructions")) + if isinstance(optional_params, dict) + else None + ) estimated_usage: Final = _estimate_dispatched_failure_usage( model=str(model_call_details.get("model") or ""), request_input=model_call_details.get("messages"), + system_input=system_input, ) if estimated_usage is None: return None diff --git a/tests/test_litellm/proxy/anthropic_endpoints/test_endpoints.py b/tests/test_litellm/proxy/anthropic_endpoints/test_endpoints.py index 0a427df0cb7..9a90daeccb7 100644 --- a/tests/test_litellm/proxy/anthropic_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/anthropic_endpoints/test_endpoints.py @@ -164,6 +164,41 @@ class TestProxyExceptionPassthrough: mock_logging.post_call_failure_hook.assert_awaited_once() +class TestFailureHookRequestData: + @pytest.mark.asyncio + async def test_failure_hook_gets_post_setup_data_with_logging_obj(self): + """Request setup replaces the processor's data dict (adding the logging + object the failure hook needs to lift token usage from); the exception + handler must pass that replaced dict, not the raw request body dict.""" + import litellm.proxy.anthropic_endpoints.endpoints as ep + import litellm.proxy.proxy_server as proxy_server + from litellm.proxy._types import ProxyException, UserAPIKeyAuth + + captured = {} + + async def fake_process(self, **kwargs): + self.data = {**self.data, "litellm_logging_obj": "logging-obj-sentinel"} + captured["processor_data"] = self.data + raise RuntimeError("provider timeout") + + with ( + patch.object(ep, "_read_request_body", new=AsyncMock(return_value={"model": "claude-sonnet"})), + patch.object(ep.ProxyBaseLLMRequestProcessing, "base_process_llm_request", new=fake_process), + patch.object(proxy_server, "proxy_logging_obj") as mock_logging, + ): + mock_logging.post_call_failure_hook = AsyncMock() + with pytest.raises(ProxyException): + await ep.anthropic_response( + fastapi_response=MagicMock(), + request=MagicMock(), + user_api_key_dict=UserAPIKeyAuth(), + ) + + hook_request_data = mock_logging.post_call_failure_hook.await_args.kwargs["request_data"] + assert hook_request_data is captured["processor_data"] + assert hook_request_data["litellm_logging_obj"] == "logging-obj-sentinel" + + class TestEventLoggingBatchEndpoint: """Test the stubbed event logging batch endpoint""" diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index 70baf157ab9..0455a806c0b 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -616,6 +616,74 @@ class TestPostCallFailureHookEstimatesDispatchedInputTokens: assert estimated.prompt_tokens > 0 assert estimated.completion_tokens == 0 + def _dispatched_request_data(self, messages, optional_params): + from datetime import datetime + + return { + "litellm_logging_obj": self._logging_obj( + { + "first_api_call_start_time": datetime.now(), + "model": "gpt-3.5-turbo", + "messages": messages, + "optional_params": optional_params, + } + ), + "metadata": {}, + } + + @pytest.mark.asyncio + async def test_anthropic_system_prompt_counted_in_estimate(self): + import litellm as litellm_module + from litellm.types.utils import Usage + + system_prompt = "You are a verbose historian who narrates every fact in exhaustive detail." + messages = [{"role": "user", "content": "write a short essay"}] + request_data = self._dispatched_request_data(messages, {"system": system_prompt, "max_tokens": 100}) + await self._run(request_data) + + estimated = request_data["combined_usage_object"] + assert isinstance(estimated, Usage) + expected = litellm_module.token_counter(model="gpt-3.5-turbo", messages=messages) + litellm_module.token_counter( + model="gpt-3.5-turbo", text=system_prompt + ) + assert estimated.prompt_tokens == expected + + @pytest.mark.asyncio + async def test_anthropic_system_text_blocks_counted_in_estimate(self): + import litellm as litellm_module + from litellm.types.utils import Usage + + system_blocks = [ + {"type": "text", "text": "part one of the system prompt. "}, + {"type": "text", "text": "part two of the system prompt."}, + ] + messages = [{"role": "user", "content": "write a short essay"}] + request_data = self._dispatched_request_data(messages, {"system": system_blocks}) + await self._run(request_data) + + estimated = request_data["combined_usage_object"] + assert isinstance(estimated, Usage) + expected = litellm_module.token_counter(model="gpt-3.5-turbo", messages=messages) + litellm_module.token_counter( + model="gpt-3.5-turbo", text="part one of the system prompt. part two of the system prompt." + ) + assert estimated.prompt_tokens == expected + + @pytest.mark.asyncio + async def test_responses_instructions_counted_in_estimate(self): + import litellm as litellm_module + from litellm.types.utils import Usage + + instructions = "Answer every question as a meticulous archivist." + request_data = self._dispatched_request_data("summarize the archive", {"instructions": instructions}) + await self._run(request_data) + + estimated = request_data["combined_usage_object"] + assert isinstance(estimated, Usage) + expected = litellm_module.token_counter( + model="gpt-3.5-turbo", text="summarize the archive" + ) + litellm_module.token_counter(model="gpt-3.5-turbo", text=instructions) + assert estimated.prompt_tokens == expected + from typing import cast From 3a4d3a01af14cb9a1058ff57466869b0735f23d6 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 18 Aug 2026 14:36:57 -0700 Subject: [PATCH 14/83] fix(proxy): only estimate failed-request input tokens for call types whose input is countable --- litellm/proxy/utils.py | 31 ++++++++++++++++++++ tests/test_litellm/proxy/test_proxy_utils.py | 28 +++++++++++++++++- 2 files changed, 58 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 5bfc1c2d1e1..1465e76d01f 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -449,6 +449,35 @@ def _estimate_dispatched_failure_usage(model: str, request_input: object, system return Usage(prompt_tokens=input_tokens, completion_tokens=0, total_tokens=input_tokens) +_INPUT_ESTIMABLE_CALL_TYPES: Final = frozenset( + call_type.value + for call_type in ( + CallTypes.completion, + CallTypes.acompletion, + CallTypes.text_completion, + CallTypes.atext_completion, + CallTypes.anthropic_messages, + CallTypes.aanthropic_messages, + CallTypes.responses, + CallTypes.aresponses, + CallTypes.embedding, + CallTypes.aembedding, + CallTypes.moderation, + CallTypes.amoderation, + CallTypes.image_generation, + CallTypes.aimage_generation, + CallTypes.speech, + CallTypes.aspeech, + CallTypes.rerank, + CallTypes.arerank, + CallTypes.generate_content, + CallTypes.agenerate_content, + CallTypes.generate_content_stream, + CallTypes.agenerate_content_stream, + ) +) + + def _failure_usage_to_lift(model_call_details: Mapping[str, object], dispatched: bool) -> tuple[object, object] | None: """A stream that broke mid-flight still billed the provider for the chunks already delivered; the streaming handler stashes that recovered usage and @@ -461,6 +490,8 @@ def _failure_usage_to_lift(model_call_details: Mapping[str, object], dispatched: return recovered_usage, model_call_details.get("response_cost") if not dispatched or model_call_details.get(LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL): return None + if str(model_call_details.get("call_type")) not in _INPUT_ESTIMABLE_CALL_TYPES: + return None optional_params: Final = model_call_details.get("optional_params") system_input: Final = ( (optional_params.get("system") or optional_params.get("instructions")) diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index 0455a806c0b..96e2c74e5d4 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -517,6 +517,7 @@ class TestPostCallFailureHookEstimatesDispatchedInputTokens: "first_api_call_start_time": datetime.now(), "model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "count these input tokens please"}], + "call_type": "acompletion", } ), "metadata": {}, @@ -582,6 +583,7 @@ class TestPostCallFailureHookEstimatesDispatchedInputTokens: "first_api_call_start_time": datetime.now(), "model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "mid-stream failure"}], + "call_type": "acompletion", "combined_usage_object": recovered_usage, "response_cost": 3.5e-05, } @@ -605,6 +607,7 @@ class TestPostCallFailureHookEstimatesDispatchedInputTokens: "first_api_call_start_time": datetime.now(), "model": "gpt-3.5-turbo", "messages": "a plain text-completion prompt string", + "call_type": "atext_completion", } ), "metadata": {}, @@ -616,7 +619,7 @@ class TestPostCallFailureHookEstimatesDispatchedInputTokens: assert estimated.prompt_tokens > 0 assert estimated.completion_tokens == 0 - def _dispatched_request_data(self, messages, optional_params): + def _dispatched_request_data(self, messages, optional_params, call_type="acompletion"): from datetime import datetime return { @@ -626,11 +629,34 @@ class TestPostCallFailureHookEstimatesDispatchedInputTokens: "model": "gpt-3.5-turbo", "messages": messages, "optional_params": optional_params, + "call_type": call_type, } ), "metadata": {}, } + @pytest.mark.asyncio + async def test_embedding_string_list_input_counted_in_estimate(self): + import litellm as litellm_module + from litellm.types.utils import Usage + + embedding_input = ["first embedding text", "second embedding text"] + request_data = self._dispatched_request_data(embedding_input, {}, call_type="aembedding") + await self._run(request_data) + + estimated = request_data["combined_usage_object"] + assert isinstance(estimated, Usage) + expected = litellm_module.token_counter(model="gpt-3.5-turbo", text="".join(embedding_input)) + assert estimated.prompt_tokens == expected + + @pytest.mark.asyncio + async def test_transcription_checksum_not_estimated(self): + request_data = self._dispatched_request_data("a1b2c3d4e5f6a7b8c9d0e1f2a3b4c5d6", {}, call_type="atranscription") + await self._run(request_data) + + assert "combined_usage_object" not in request_data + assert "response_cost" not in request_data + @pytest.mark.asyncio async def test_anthropic_system_prompt_counted_in_estimate(self): import litellm as litellm_module From 72960d10e93d1cb1bb6ef7b2fd3be82cb9aa7368 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 18 Aug 2026 14:45:05 -0700 Subject: [PATCH 15/83] fix(proxy): address review findings on project ITPM/OTPM quotas - scale batch output-token reservations by the row's n / best_of candidate count - parse client-supplied output caps defensively instead of 500ing on unparseable values - exclude project IO descriptors from the first should_rate_limit pass when TPM reservation is disabled so their buckets are not double-charged --- litellm/proxy/hooks/batch_rate_limiter.py | 12 ++- .../hooks/parallel_request_limiter_v3.py | 47 +++++++---- .../proxy/hooks/test_batch_file_validation.py | 25 ++++++ .../proxy/hooks/test_tpm_concurrent.py | 78 +++++++++++++++++++ 4 files changed, 144 insertions(+), 18 deletions(-) diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index 4052a9b04fa..aa3ac6b820b 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -429,12 +429,20 @@ class _PROXY_BatchRateLimiter(CustomLogger): ), None, ) + candidate_count: Final = max( + ( + v + for v in (body.get("n"), body.get("best_of")) + if isinstance(v, int) and not isinstance(v, bool) and v > 1 + ), + default=1, + ) if explicit_cap is not None: try: - return max(0, int(explicit_cap)) + return max(0, int(explicit_cap)) * candidate_count except (TypeError, ValueError): pass - return self.parallel_request_limiter.no_max_tokens_output_floor(min_configured_otpm_limit) + return self.parallel_request_limiter.no_max_tokens_output_floor(min_configured_otpm_limit) * candidate_count @staticmethod def _has_applicable_batch_rate_limits( diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 751b301a93d..7cd56e66000 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -559,6 +559,15 @@ def _call_id_from_callback_kwargs(kwargs: object) -> str | None: return call_id if isinstance(call_id, str) else None +def _parse_output_cap_value(raw_value: object) -> int | None: + if isinstance(raw_value, bool) or not isinstance(raw_value, (int, float, str)): + return None + try: + return int(float(raw_value)) + except (ValueError, OverflowError): + return None + + class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): def __init__( self, @@ -707,20 +716,18 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): if call_type in GOOGLE_GENAI_NATIVE_CALL_TYPES: config: Final = data.get("config") if "config" in data else data.get("generationConfig") google_cap_values: Final = tuple( - int(raw_value) + parsed for field in ("maxOutputTokens", "max_output_tokens") if isinstance(config, dict) - for raw_value in (config.get(field),) - if isinstance(raw_value, (int, float, str)) + for parsed in (_parse_output_cap_value(config.get(field)),) + if parsed is not None ) return max(google_cap_values, default=None) if call_type in RESPONSES_API_CALL_TYPES: - value: Final = data.get("max_output_tokens") - if value is None: + responses_cap: Final = _parse_output_cap_value(data.get("max_output_tokens")) + if responses_cap is None: return None - if not isinstance(value, (int, float, str)): - return None - return max(RESPONSES_API_MIN_OUTPUT_TOKENS, int(value)) + return max(RESPONSES_API_MIN_OUTPUT_TOKENS, responses_cap) if call_type in EMBEDDING_API_CALL_TYPES: return None fields: Final = ( @@ -729,10 +736,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): else ("max_tokens", "max_completion_tokens", "max_output_tokens") ) output_cap_values: Final = tuple( - int(raw_value) - for field in fields - for raw_value in (data.get(field),) - if isinstance(raw_value, (int, float, str)) + parsed for field in fields for parsed in (_parse_output_cap_value(data.get(field)),) if parsed is not None ) return max(output_cap_values, default=None) @@ -1209,7 +1213,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): async def should_rate_limit( self, - descriptors: list[RateLimitDescriptor], + descriptors: Sequence[RateLimitDescriptor], parent_otel_span: Span | None = None, read_only: bool = False, skip_tpm_check: bool = False, @@ -1335,7 +1339,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): def _collect_windowed_keys_and_gauges( self, - descriptors: list[RateLimitDescriptor], + descriptors: Sequence[RateLimitDescriptor], skip_tpm_check: bool, ) -> tuple[list[str], dict[str, WindowKeyMetadata], list[ParallelRequestGauge]]: """ @@ -3456,7 +3460,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): # in-flight request would pre-inflate the :tokens counter by 1, # shrinking the effective TPM budget by N and causing # false-positive 429s under bursts. When reservation is disabled, - # this pass enforces TPM directly from the post-call counters. + # this pass enforces TPM directly from the post-call counters -- + # except for project ITPM/OTPM descriptors, which are excluded + # then because _reserve_project_io_tokens_or_raise below charges + # them unconditionally and counting them here too would + # double-charge every request. parallel_counter_keys: Final = [ self.create_rate_limit_keys(d["key"], d["value"], "max_parallel_requests") for d in descriptors @@ -3464,8 +3472,15 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ] parallel_slot_id: Final = uuid.uuid4().hex if parallel_counter_keys else None + first_pass_descriptors: Final = ( + descriptors + if self.tpm_reservation_enabled + else tuple( + d for d in descriptors if d["key"] not in (PROJECT_ITPM_DESCRIPTOR_KEY, PROJECT_OTPM_DESCRIPTOR_KEY) + ) + ) response: Final = await self.should_rate_limit( - descriptors=descriptors, + descriptors=first_pass_descriptors, parent_otel_span=user_api_key_dict.parent_otel_span, skip_tpm_check=self.tpm_reservation_enabled, parallel_slot_id=parallel_slot_id, diff --git a/tests/test_litellm/proxy/hooks/test_batch_file_validation.py b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py index e5211973ec3..3a9d3ff72aa 100644 --- a/tests/test_litellm/proxy/hooks/test_batch_file_validation.py +++ b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py @@ -2060,3 +2060,28 @@ def test_estimate_entry_output_tokens_prefers_max_tokens_over_max_output_tokens( } assert rate_limiter._estimate_entry_output_tokens(entry, None) == 50 + + +@pytest.mark.parametrize( + ("body_extra", "expected"), + [ + ({"max_tokens": 40, "n": 10}, 400), + ({"max_tokens": 40, "best_of": 5}, 200), + ({"max_tokens": 40, "n": 3, "best_of": 5}, 200), + ({"max_tokens": 40, "n": 0}, 40), + ({"max_tokens": 40, "n": -2}, 40), + ({"n": 3}, 2997), + ], +) +def test_estimate_entry_output_tokens_multiplies_candidate_count(body_extra, expected): + """A row generating n / best_of candidates consumes that many completions' + worth of output tokens, so the OTPM reservation must scale with the + effective candidate count. Pre-fix a `max_tokens: 40, n: 10` row consumed + up to 400 output tokens while reserving only 40.""" + rate_limiter = _output_estimator() + entry = { + "url": "/v1/chat/completions", + "body": {"model": "gpt-4o", "messages": [], **body_extra}, + } + + assert rate_limiter._estimate_entry_output_tokens(entry, None) == expected diff --git a/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py b/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py index fd7e0626af9..e12610b769c 100644 --- a/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py +++ b/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py @@ -3393,6 +3393,84 @@ def test_split_quota_helpers_handle_non_mapping_inputs(rate_limiter): assert handler._build_io_token_reservation_ops(object(), object()) == () +@pytest.mark.parametrize( + ("data", "call_type", "expected"), + [ + ({"max_tokens": "30.0"}, "", 30), + ({"max_tokens": "not-a-number"}, "", None), + ({"max_tokens": True}, "", None), + ({"max_output_tokens": "30.0"}, "responses", 30), + ({"max_output_tokens": "nan"}, "responses", None), + ({"generationConfig": {"maxOutputTokens": "12.5"}}, "agenerate_content", 12), + ({"generationConfig": {"maxOutputTokens": "oops"}}, "agenerate_content", None), + ], +) +def test_get_explicit_output_cap_tolerates_unparseable_values( + rate_limiter, data, call_type, expected +): + """A client-supplied cap the proxy cannot parse must fall back to the + no-cap output estimate instead of raising ValueError and 500ing the + request before it ever reaches the provider.""" + handler, _cache = rate_limiter + + assert handler._get_explicit_output_cap(data, call_type) == expected + + +@pytest.mark.asyncio +async def test_project_io_counters_not_double_charged_when_reservation_disabled( + monkeypatch, +): + """With LITELLM_TPM_TOKEN_RESERVATION_ENABLED=false the first + should_rate_limit pass used to +1 every ITPM/OTPM counter on top of the + full reservation _reserve_project_io_tokens_or_raise always makes, + permanently inflating each bucket by one token per request.""" + monkeypatch.setenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", "false") + cache = DualCache() + handler = RateLimitHandler(internal_usage_cache=InternalUsageCache(cache)) + assert handler.tpm_reservation_enabled is False + + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-io-no-reservation"), + project_id="proj-io-no-reservation", + project_metadata={ + "model_itpm_limit": {"bedrock_mantle/claude-opus": 1000000}, + "model_otpm_limit": {"bedrock_mantle/claude-opus": 1000000}, + }, + ) + data = { + "model": "bedrock_mantle/claude-opus", + "messages": [{"role": "user", "content": "hello"}], + "max_tokens": 50, + } + + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data=data, + call_type="", + ) + + stash = get_request_stash() + assert stash is not None + assert stash.itpm_reserved_tokens > 0 + assert stash.otpm_reserved_tokens > 0 + + for descriptor_key, reserved in ( + ("model_per_project_itpm", stash.itpm_reserved_tokens), + ("model_per_project_otpm", stash.otpm_reserved_tokens), + ): + counter_key = handler.create_rate_limit_keys( + key=descriptor_key, + value="proj-io-no-reservation:bedrock_mantle/claude-opus", + rate_limit_type="tokens", + ) + cached = await cache.async_get_cache(key=counter_key, local_only=True) + assert int(cached or 0) == reserved, ( + f"{descriptor_key} counter {cached} != reserved {reserved}: " + "first-pass should_rate_limit double-charged the bucket" + ) + + @pytest.mark.parametrize( ("call_type", "data"), [ From 96cee087bebad1fa215c8ce1051e31ba730e5f06 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 18 Aug 2026 14:49:12 -0700 Subject: [PATCH 16/83] test(e2e): pin auto-router tag-split, alias pricing, heuristic scope, and Responses routing regressions --- tests/e2e/coverage_registry/reliability.yaml | 9 + tests/e2e/models.py | 4 + .../test_auto_router_regressions_e2e.py | 632 ++++++++++++++++++ 3 files changed, 645 insertions(+) create mode 100644 tests/e2e/router/test_auto_router_regressions_e2e.py diff --git a/tests/e2e/coverage_registry/reliability.yaml b/tests/e2e/coverage_registry/reliability.yaml index ebbfd3415a5..5a1d437eac9 100644 --- a/tests/e2e/coverage_registry/reliability.yaml +++ b/tests/e2e/coverage_registry/reliability.yaml @@ -18,6 +18,15 @@ - {id: reliability.routing.usage_based.picks_under_tpm, module: reliability, tier: P0, behavior: routing, variant: usage_based, assertions: [picks_under_tpm], exercised_on: [chat_completions, messages], source: "router_strategy/lowest_tpm_rpm_v2.py", rationale: "Routes to lowest-TPM deployment; prevents over-allocation"} - {id: reliability.routing.least_busy.picks_lowest_traffic, module: reliability, tier: P1, behavior: routing, variant: least_busy, assertions: [picks_lowest_traffic], exercised_on: [chat_completions, messages], source: "router_strategy/least_busy.py", rationale: "Fewest in-flight requests"} - {id: reliability.routing.complexity_llm_classifier.routes_by_llm_tier, module: reliability, tier: P1, behavior: routing, variant: complexity_llm_classifier, assertions: [routes_by_llm_tier], exercised_on: [chat_completions], source: "router_strategy/complexity_router/complexity_router.py", fail_before_fix: proven, rationale: "v2 auto-router LLM complexity classifier runs over the proxy and routes by semantic tier instead of silently crashing on absent litellm_metadata and falling back to heuristic scoring"} +- {id: reliability.routing.tagged_marker.request_tag_selects_marker, module: reliability, tier: P0, behavior: routing, variant: tagged_marker, assertions: [request_tag_selects_marker], exercised_on: [chat_completions], source: "litellm/router.py:11445", rationale: "Tagged request selects the tagged strategy marker under a shared model_name instead of the plain deployment registered first (GitHub issue #36619)"} +- {id: reliability.routing.tagged_marker.untagged_request_served_by_plain_deployment, module: reliability, tier: P0, behavior: routing, variant: tagged_marker, assertions: [untagged_request_served_by_plain_deployment], exercised_on: [chat_completions, messages, responses], source: "litellm/router.py:11445", rationale: "Untagged requests to a shared model_name are served by the plain deployment on every call, never captured or errored by the tagged marker (GitHub issue #36620)"} +- {id: reliability.routing.tagged_marker.header_tag_selects_marker, module: reliability, tier: P0, behavior: routing, variant: tagged_marker, assertions: [header_tag_selects_marker], exercised_on: [messages], source: "litellm/router.py:11445", rationale: "A request tagged only via the x-litellm-tags header selects the tagged marker on Anthropic-native /v1/messages (GitHub issue #36621)"} +- {id: reliability.routing.tagged_marker.untagged_tier_deployments_still_served, module: reliability, tier: P1, behavior: routing, variant: tagged_marker, assertions: [untagged_tier_deployments_still_served], exercised_on: [chat_completions, messages], source: "litellm/router_strategy/tag_based_routing.py:433", rationale: "Routing tags the marker consumed no longer constrain deployment selection inside the routed tier group, so untagged tier deployments serve the rewrite (GitHub issue #36621)"} +- {id: reliability.routing.tagged_marker.tag_semantics_stay_strict, module: reliability, tier: P1, behavior: routing, variant: tagged_marker, assertions: [tag_semantics_stay_strict], exercised_on: [chat_completions], source: "litellm/router_strategy/tag_based_routing.py:299", rationale: "Tag consumption must not loosen strict semantics: a tagged call aimed straight at an untagged deployment still gets the 401 tags-configuration denial"} +- {id: reliability.routing.tagged_marker.responses_input_routes_through_marker, module: reliability, tier: P0, behavior: routing, variant: tagged_marker, assertions: [responses_input_routes_through_marker], exercised_on: [responses], source: "litellm/router.py:11489", rationale: "Tagged /v1/responses (header or litellm_metadata.tags, string or list input) routes through the marker to its tier, extending the GitHub issues #36620/#36621 tag split to the Responses surface"} +- {id: reliability.routing.semantic_auto_router.responses_input_routed, module: reliability, tier: P0, behavior: routing, variant: semantic_auto_router, assertions: [responses_input_routed], exercised_on: [responses], source: "litellm/router_strategy/auto_router/auto_router.py:131", fail_before_fix: proven, rationale: "/v1/responses input is resolved into messages for the semantic auto-router pre-routing hook instead of failing 400 Unmapped LLM provider auto_router (GitHub PR #37333)"} +- {id: reliability.routing.strategy_alias.custom_pricing_ignored, module: reliability, tier: P1, behavior: routing, variant: strategy_alias, assertions: [custom_pricing_ignored], exercised_on: [chat_completions], source: "litellm/router.py:11489", rationale: "Custom pricing on a strategy-router alias never prices the routed request; spend logs at the routed tier deployment's own rate (GitHub PR #36691)"} +- {id: reliability.routing.complexity_heuristic.scores_current_ask_only, module: reliability, tier: P1, behavior: routing, variant: complexity_heuristic, assertions: [scores_current_ask_only], exercised_on: [chat_completions], source: "router_strategy/complexity_router/complexity_router.py:942", rationale: "The heuristic complexity classifier scores the caller's current ask only, so a keyword-heavy agent system prompt cannot inflate the tier (GitHub PR #36721)"} - {id: reliability.cache.exact.returns_cached, module: reliability, tier: P1, behavior: cache, variant: exact, assertions: [returns_cached], exercised_on: [chat_completions, messages, embeddings], source: "litellm/caching/caching.py", rationale: "Response cache returns cached on exact match"} - {id: reliability.cache.prompt_caching_model_select.returns_cached, module: reliability, tier: P1, behavior: cache, variant: prompt_caching_model_select, assertions: [returns_cached], exercised_on: [chat_completions], source: "router_utils/prompt_caching_cache.py", rationale: "Selects model supporting prompt caching for cacheable prefix"} - {id: reliability.circuit_breaker.redis.trips_then_recovers, module: reliability, tier: P0, behavior: circuit_breaker, variant: redis, assertions: [trips_then_recovers], exercised_on: [chat_completions, messages, embeddings], source: "litellm/caching/redis_cache.py:99", rationale: "Redis breaker CLOSED->OPEN->HALF_OPEN; guards all cache/rate-limit ops"} diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 734e63a94e6..957da605546 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -744,6 +744,10 @@ class LiteLLMParamsBody(BaseModel): extra_headers: dict[str, str] | None = None use_in_pass_through: bool | None = None complexity_router_config: dict[str, object] | None = None + auto_router_config: str | None = None + auto_router_default_model: str | None = None + auto_router_embedding_model: str | None = None + tags: list[str] | None = None mock_response: str | None = None timeout: float | None = None tpm: int | None = None diff --git a/tests/e2e/router/test_auto_router_regressions_e2e.py b/tests/e2e/router/test_auto_router_regressions_e2e.py new file mode 100644 index 00000000000..ee68d3fa1bb --- /dev/null +++ b/tests/e2e/router/test_auto_router_regressions_e2e.py @@ -0,0 +1,632 @@ +"""Live e2e regression pins for strategy-router (auto-router) routing. + +A strategy marker (an ``auto_router/complexity_router`` deployment) and a plain +deployment can share one ``model_name``, split by tags once +``enable_tag_filtering`` is on: tagged requests route through the marker to its +tier models, untagged requests go to the plain deployment. That split, and the +strategy-router alias behaviors around it, regressed repeatedly; each test here +pins one fixed behavior: + +- GitHub issue #36619: a tagged request selects the tagged marker under a + shared name even when a plain deployment was registered first. +- GitHub issue #36620: untagged requests keep being served by the plain + deployment on every call, never captured or 400'd by the tagged marker. +- GitHub issue #36621: a request tagged via the ``x-litellm-tags`` header + routes through the marker even when the tier deployments carry no tags + (the marker consumes the routing tags before deployment selection), while a + tagged call aimed straight at an untagged deployment stays denied. +- GitHub issues #36620/#36621 on /v1/responses: the same tag split holds for + string and list input, whether the tag arrives in litellm_metadata or the + x-litellm-tags header. +- GitHub PR #37333: /v1/responses input is resolved into messages for a + semantic ``auto_router`` deployment's pre-routing hook; such requests used + to fail with 400 "Unmapped LLM provider auto_router" because only chat + messages fed the route matcher. +- GitHub PR #36691: custom pricing on the marker alias never prices the routed + request; spend logs at the routed tier deployment's own rate. +- GitHub PR #36721: the heuristic complexity classifier scores the caller's + current ask only, so a large agent system prompt cannot inflate the tier. + +Every deployment is registered via /model/new (stage has no static config for +these) and ``enable_tag_filtering`` is flipped through /config/update and +restored on teardown, mirroring TestRouterSettings in the management suite. +The served deployment is always read back from the spend log's ``model``, +which stores either the registered alias or the provider-prefixed form. +""" + +import json +import os +import time +from collections.abc import Iterator +from dataclasses import dataclass +from typing import Final + +import pytest +from pydantic import BaseModel, ConfigDict, Field + +from e2e_config import unique_marker +from e2e_http import AnthropicHeaders, AuthHeaders, NoBody, UnauthorizedError, unwrap +from lifecycle import ResourceManager +from models import ( + AnthropicMessagesBody, + AnthropicMessagesResponse, + ChatBody, + ChatMessage, + ChatMetadata, + KeyGenerateBody, + LiteLLMParamsBody, + SpendLogRow, +) +from proxy_client import ProxyClient + +pytestmark = pytest.mark.e2e + +PLAIN_MODEL = "anthropic/claude-sonnet-5" +CHEAP_MODEL = "anthropic/claude-haiku-4-5" +STRONG_MODEL = "openai/gpt-5.6" +MAX_TOKENS = 16 +PLAIN_SERVED = frozenset({PLAIN_MODEL, "claude-sonnet-5"}) +CHEAP_SERVED = frozenset({CHEAP_MODEL, "claude-haiku-4-5"}) +EMBEDDING_MODEL = "openai/text-embedding-3-small" +SEMANTIC_ROUTE_UTTERANCE = "summarize this quarterly revenue report into three bullet points" + +KEYWORD_HEAVY_SYSTEM_PROMPT = ( + "You are the principal architecture assistant for a distributed systems platform. " + "Analyze every request step by step: design the algorithm, prove its correctness, " + "evaluate time and space complexity, and reason about concurrency, consistency, and " + "fault tolerance tradeoffs. When asked, refactor and debug multi-threaded code, " + "optimize database query plans, derive mathematical proofs, and explain the theorem " + "or lemma behind each optimization. Think through edge cases rigorously before answering. " +) * 4 + + +class TaggedAuthHeaders(AuthHeaders): + x_litellm_tags: str | None = Field(default=None, serialization_alias="x-litellm-tags") + + +class TaggedAnthropicHeaders(AnthropicHeaders): + x_litellm_tags: str | None = Field(default=None, serialization_alias="x-litellm-tags") + + +class ResponsesTagMetadata(BaseModel): + tags: list[str] + + +class ResponsesInputItem(BaseModel): + role: str + content: str + + +class ResponsesBody(BaseModel): + model: str + input: str | list[ResponsesInputItem] + max_output_tokens: int | None = None + litellm_metadata: ResponsesTagMetadata | None = None + + +class ResponsesApiResponse(BaseModel): + """Minimal /v1/responses answer shape; routing is proven from spend logs, + so only the fields the assertions read are modeled.""" + + model_config = ConfigDict(extra="allow") + id: str | None = None + status: str | None = None + model: str | None = None + + +class RouterSettingsPatch(BaseModel): + enable_tag_filtering: bool + + +class ConfigUpdateBody(BaseModel): + router_settings: RouterSettingsPatch + + +class ConfigUpdateResponse(BaseModel): + message: str + + +class RouterCurrentValues(BaseModel): + enable_tag_filtering: bool | None = None + + +class RouterSettingsResponse(BaseModel): + current_values: RouterCurrentValues + + +@dataclass(frozen=True, slots=True) +class TagSplitDeployments: + """Scenario A mirrors the customer-shaped config from GitHub issue #36619: + plain deployment registered first, tier deployment and marker both tagged. + Scenario B flips both axes for GitHub issue #36621: marker registered first + and its tier deployment left untagged, so routing depends neither on + registration order nor on tier deployments carrying tags.""" + + tag_a: str + shared_a: str + tier_a: str + tag_b: str + shared_b: str + tier_b: str + + +@dataclass(frozen=True, slots=True) +class ZeroPricedAlias: + alias: str + tier: str + + +@dataclass(frozen=True, slots=True) +class HeuristicSplit: + alias: str + cheap: str + strong: str + + +@dataclass(frozen=True, slots=True) +class SemanticAutoRouter: + marker: str + target: str + fallback: str + embedding: str + + +def _provider_key(env_var: str) -> str: + return os.environ.get(env_var) or f"os.environ/{env_var}" + + +def _uniform_tier_config(tier_model: str) -> dict[str, object]: + return { + "classifier_type": "heuristic", + "tiers": {"SIMPLE": tier_model, "MEDIUM": tier_model, "COMPLEX": tier_model, "REASONING": tier_model}, + } + + +def _read_tag_filtering(proxy: ProxyClient) -> bool | None: + return unwrap( + proxy.transport.get( + "/router/settings", + headers=proxy.transport.master, + params=NoBody(), + response_type=RouterSettingsResponse, + ) + ).current_values.enable_tag_filtering + + +def _write_tag_filtering(proxy: ProxyClient, enabled: bool) -> None: + response: Final = unwrap( + proxy.transport.post( + "/config/update", + headers=proxy.transport.master, + json=ConfigUpdateBody(router_settings=RouterSettingsPatch(enable_tag_filtering=enabled)), + response_type=ConfigUpdateResponse, + ) + ) + assert "success" in response.message.lower(), ( + f"/config/update reported {response.message!r}, expected a success message" + ) + + +def _await_tag_filtering(proxy: ProxyClient, expected: bool) -> None: + deadline: Final = time.monotonic() + proxy.poll_timeout + while time.monotonic() < deadline: + if _read_tag_filtering(proxy) is expected: + return + time.sleep(proxy.poll_interval) + raise AssertionError( + f"GET /router/settings never reported enable_tag_filtering={expected} after /config/update" + ) + + +def _key_for(proxy: ProxyClient, resources: ResourceManager, models: list[str]) -> str: + key: Final = proxy.generate_key(KeyGenerateBody(models=models, user_id="e2e-auto-router-regressions")) + resources.defer(lambda: proxy.delete_key(key)) + return key + + +def _hello_chat_body(model: str, tags: list[str] | None = None) -> ChatBody: + return ChatBody( + model=model, + messages=[ChatMessage(role="user", content=f"say hello {unique_marker()}")], + max_tokens=MAX_TOKENS, + metadata=ChatMetadata(tags=tags) if tags is not None else None, + ) + + +def _hello_messages_body(model: str) -> AnthropicMessagesBody: + return AnthropicMessagesBody( + model=model, + messages=[ChatMessage(role="user", content=f"say hello {unique_marker()}")], + max_tokens=MAX_TOKENS, + ) + + +def _assert_served_only_by(rows: list[SpendLogRow], allowed: frozenset[str], context: str) -> None: + served: Final = tuple(row.model for row in rows) + assert served and all(model in allowed for model in served), ( + f"{context}: expected every request to be served by one of {sorted(allowed)}, spend logs show {served}" + ) + + +@pytest.fixture(scope="module") +def tag_filtering(proxy: ProxyClient) -> Iterator[None]: + """enable_tag_filtering is what splits tagged from untagged traffic in every + scenario here. /config/update is the only write path for router_settings; + the original value is restored on teardown so the shared proxy keeps its + configuration for the rest of the run.""" + original: Final = bool(_read_tag_filtering(proxy)) + _write_tag_filtering(proxy, True) + _await_tag_filtering(proxy, True) + try: + yield + finally: + _write_tag_filtering(proxy, original) + _await_tag_filtering(proxy, original) + + +@pytest.fixture(scope="module") +def split(proxy: ProxyClient, tag_filtering: None) -> Iterator[TagSplitDeployments]: + marker: Final = unique_marker() + deployments: Final = TagSplitDeployments( + tag_a=f"e2e-split-a-{marker}", + shared_a=f"e2e-autoroute-a-{marker}", + tier_a=f"e2e-tier-a-{marker}", + tag_b=f"e2e-split-b-{marker}", + shared_b=f"e2e-autoroute-b-{marker}", + tier_b=f"e2e-tier-b-{marker}", + ) + anthropic_key: Final = _provider_key("ANTHROPIC_API_KEY") + marker_params_a: Final = LiteLLMParamsBody( + model="auto_router/complexity_router", + complexity_router_config=_uniform_tier_config(deployments.tier_a), + tags=[deployments.tag_a], + ) + marker_params_b: Final = LiteLLMParamsBody( + model="auto_router/complexity_router", + complexity_router_config=_uniform_tier_config(deployments.tier_b), + tags=[deployments.tag_b], + ) + registrations: Final[tuple[tuple[str, LiteLLMParamsBody], ...]] = ( + (deployments.shared_a, LiteLLMParamsBody(model=PLAIN_MODEL, api_key=anthropic_key)), + (deployments.tier_a, LiteLLMParamsBody(model=CHEAP_MODEL, api_key=anthropic_key, tags=[deployments.tag_a])), + (deployments.shared_a, marker_params_a), + (deployments.shared_b, marker_params_b), + (deployments.tier_b, LiteLLMParamsBody(model=CHEAP_MODEL, api_key=anthropic_key)), + (deployments.shared_b, LiteLLMParamsBody(model=PLAIN_MODEL, api_key=anthropic_key)), + ) + created: Final = tuple(proxy.create_model(name, params) for name, params in registrations) + try: + yield deployments + finally: + for model_id in created: + proxy.delete_model(model_id) + + +@pytest.fixture(scope="module") +def zero_priced_alias(proxy: ProxyClient) -> Iterator[ZeroPricedAlias]: + marker: Final = unique_marker() + named: Final = ZeroPricedAlias(alias=f"e2e-priced-alias-{marker}", tier=f"e2e-priced-tier-{marker}") + alias_params: Final = LiteLLMParamsBody( + model="auto_router/complexity_router", + complexity_router_config=_uniform_tier_config(named.tier), + input_cost_per_token=0.0, + output_cost_per_token=0.0, + ) + registrations: Final[tuple[tuple[str, LiteLLMParamsBody], ...]] = ( + (named.tier, LiteLLMParamsBody(model=CHEAP_MODEL, api_key=_provider_key("ANTHROPIC_API_KEY"))), + (named.alias, alias_params), + ) + created: Final = tuple(proxy.create_model(name, params) for name, params in registrations) + try: + yield named + finally: + for model_id in created: + proxy.delete_model(model_id) + + +@pytest.fixture(scope="module") +def heuristic_split(proxy: ProxyClient) -> Iterator[HeuristicSplit]: + marker: Final = unique_marker() + named: Final = HeuristicSplit( + alias=f"e2e-heuristic-router-{marker}", + cheap=f"e2e-heuristic-cheap-{marker}", + strong=f"e2e-heuristic-strong-{marker}", + ) + config: Final[dict[str, object]] = { + "classifier_type": "heuristic", + "token_thresholds": {"simple": 15, "complex": 400}, + "tiers": {"SIMPLE": named.cheap, "MEDIUM": named.strong, "COMPLEX": named.strong, "REASONING": named.strong}, + } + registrations: Final[tuple[tuple[str, LiteLLMParamsBody], ...]] = ( + (named.cheap, LiteLLMParamsBody(model=CHEAP_MODEL, api_key=_provider_key("ANTHROPIC_API_KEY"))), + (named.strong, LiteLLMParamsBody(model=STRONG_MODEL, api_key=_provider_key("OPENAI_API_KEY"))), + (named.alias, LiteLLMParamsBody(model="auto_router/complexity_router", complexity_router_config=config)), + ) + created: Final = tuple(proxy.create_model(name, params) for name, params in registrations) + try: + yield named + finally: + for model_id in created: + proxy.delete_model(model_id) + + +@pytest.fixture(scope="module") +def semantic_auto_router(proxy: ProxyClient) -> Iterator[SemanticAutoRouter]: + marker: Final = unique_marker() + named: Final = SemanticAutoRouter( + marker=f"e2e-semantic-router-{marker}", + target=f"e2e-semantic-target-{marker}", + fallback=f"e2e-semantic-fallback-{marker}", + embedding=f"e2e-semantic-embedding-{marker}", + ) + router_config: Final = json.dumps( + {"routes": [{"name": named.target, "utterances": [SEMANTIC_ROUTE_UTTERANCE], "score_threshold": 0.3}]} + ) + marker_params: Final = LiteLLMParamsBody( + model=f"auto_router/{named.marker}", + auto_router_config=router_config, + auto_router_default_model=named.fallback, + auto_router_embedding_model=named.embedding, + ) + registrations: Final[tuple[tuple[str, LiteLLMParamsBody], ...]] = ( + (named.embedding, LiteLLMParamsBody(model=EMBEDDING_MODEL, api_key=_provider_key("OPENAI_API_KEY"))), + (named.target, LiteLLMParamsBody(model=CHEAP_MODEL, api_key=_provider_key("ANTHROPIC_API_KEY"))), + (named.fallback, LiteLLMParamsBody(model=PLAIN_MODEL, api_key=_provider_key("ANTHROPIC_API_KEY"))), + (named.marker, marker_params), + ) + created: Final = tuple(proxy.create_model(name, params) for name, params in registrations) + try: + yield named + finally: + for model_id in created: + proxy.delete_model(model_id) + + +class TestTagSplitRouting: + @pytest.mark.covers("reliability.routing.tagged_marker.request_tag_selects_marker") + def test_body_tagged_chat_routes_through_the_marker_to_its_tier( + self, proxy: ProxyClient, resources: ResourceManager, split: TagSplitDeployments + ) -> None: + """Pins GitHub issue #36619: with tag filtering on, a chat request whose + body metadata tags match the tagged marker under a shared model name is + answered by the marker's tier deployment, not by the plain deployment + that was registered under the name first.""" + key: Final = _key_for(proxy, resources, [split.shared_a, split.tier_a]) + chat: Final = unwrap(proxy.chat(key, _hello_chat_body(split.shared_a, tags=[split.tag_a]))) + assert chat.choices, "tagged chat through the shared name returned no choices" + rows: Final = proxy.poll_logs_for_key(key, min_rows=1) + _assert_served_only_by(rows, CHEAP_SERVED | {split.tier_a}, "body-tagged chat on the shared name") + + @pytest.mark.covers("reliability.routing.tagged_marker.untagged_request_served_by_plain_deployment") + def test_untagged_chat_is_always_served_by_the_plain_deployment( + self, proxy: ProxyClient, resources: ResourceManager, split: TagSplitDeployments + ) -> None: + """Pins GitHub issue #36620: untagged chat requests to the shared name + succeed on every call and are all served by the plain deployment; the + tagged marker never captures them, so no intermittent auto-router + errors and no tier hijacking.""" + key: Final = _key_for(proxy, resources, [split.shared_a, split.tier_a]) + for _ in range(5): + chat = unwrap(proxy.chat(key, _hello_chat_body(split.shared_a))) + assert chat.choices, "untagged chat through the shared name returned no choices" + rows: Final = proxy.poll_logs_for_key(key, min_rows=5) + _assert_served_only_by(rows, PLAIN_SERVED | {split.shared_a}, "untagged chat on the shared name") + + @pytest.mark.covers("reliability.routing.tagged_marker.untagged_request_served_by_plain_deployment") + def test_untagged_messages_is_served_by_the_plain_deployment( + self, proxy: ProxyClient, resources: ResourceManager, split: TagSplitDeployments + ) -> None: + """Pins GitHub issue #36620 on the /v1/messages surface: an untagged + Anthropic-native request to the shared name is served by the plain + deployment, not captured by the tagged marker.""" + key: Final = _key_for(proxy, resources, [split.shared_a, split.tier_a]) + answer: Final = unwrap(proxy.messages(key, _hello_messages_body(split.shared_a))) + assert answer.content or answer.choices, "untagged /v1/messages returned neither content nor choices" + rows: Final = proxy.poll_logs_for_key(key, min_rows=1) + _assert_served_only_by(rows, PLAIN_SERVED | {split.shared_a}, "untagged /v1/messages on the shared name") + + +class TestUntaggedTierDeployments: + @pytest.mark.covers("reliability.routing.tagged_marker.header_tag_selects_marker") + def test_header_tagged_messages_routes_through_the_marker_to_an_untagged_tier( + self, proxy: ProxyClient, resources: ResourceManager, split: TagSplitDeployments + ) -> None: + """Pins GitHub issue #36621: a /v1/messages request tagged only via the + x-litellm-tags header selects the tagged marker, and the rewrite still + lands on the tier deployment even though that deployment carries no + tags, because the marker consumed the routing tags.""" + key: Final = _key_for(proxy, resources, [split.shared_b, split.tier_b]) + headers: Final = TaggedAnthropicHeaders(authorization=f"Bearer {key}", x_litellm_tags=split.tag_b) + answer: Final = unwrap( + proxy.transport.post( + "/v1/messages", + headers=headers, + json=_hello_messages_body(split.shared_b), + response_type=AnthropicMessagesResponse, + ) + ) + assert answer.content or answer.choices, "header-tagged /v1/messages returned neither content nor choices" + rows: Final = proxy.poll_logs_for_key(key, min_rows=1) + _assert_served_only_by(rows, CHEAP_SERVED | {split.tier_b}, "header-tagged /v1/messages on the shared name") + + @pytest.mark.covers("reliability.routing.tagged_marker.untagged_tier_deployments_still_served") + def test_body_tagged_chat_reaches_the_untagged_tier_after_marker_rewrite( + self, proxy: ProxyClient, resources: ResourceManager, split: TagSplitDeployments + ) -> None: + """Pins the tag-consumption half of GitHub issue #36621: after the + tagged marker rewrites the request to its tier model, the consumed + routing tags no longer constrain deployment selection, so the untagged + tier deployment serves the request instead of a strict-tag denial.""" + key: Final = _key_for(proxy, resources, [split.shared_b, split.tier_b]) + chat: Final = unwrap(proxy.chat(key, _hello_chat_body(split.shared_b, tags=[split.tag_b]))) + assert chat.choices, "body-tagged chat through the marker-first shared name returned no choices" + rows: Final = proxy.poll_logs_for_key(key, min_rows=1) + _assert_served_only_by(rows, CHEAP_SERVED | {split.tier_b}, "body-tagged chat with untagged tier") + + @pytest.mark.covers("reliability.routing.tagged_marker.tag_semantics_stay_strict") + def test_tagged_call_straight_at_an_untagged_deployment_stays_denied( + self, proxy: ProxyClient, resources: ResourceManager, split: TagSplitDeployments + ) -> None: + """The tag-consumption fix must not loosen strict tag semantics: a + tagged request aimed directly at an untagged deployment (no marker + involved) is still rejected with the 401 tags-configuration error.""" + key: Final = _key_for(proxy, resources, [split.tier_b]) + result: Final = proxy.chat(key, _hello_chat_body(split.tier_b, tags=[split.tag_b])) + assert isinstance(result, UnauthorizedError), ( + f"expected the tagged direct call to an untagged deployment to be denied with 401, got {result}" + ) + + +class TestResponsesApiTagRouting: + @pytest.mark.covers("reliability.routing.tagged_marker.responses_input_routes_through_marker") + def test_header_tagged_responses_with_string_input_routes_to_the_tier( + self, proxy: ProxyClient, resources: ResourceManager, split: TagSplitDeployments + ) -> None: + """Pins the /v1/responses surface of the tag split (GitHub issues + #36620/#36621): a /v1/responses request with string input, tagged via + the x-litellm-tags header, succeeds and routes through the tagged + marker to its tier.""" + key: Final = _key_for(proxy, resources, [split.shared_a, split.tier_a]) + headers: Final = TaggedAuthHeaders(authorization=f"Bearer {key}", x_litellm_tags=split.tag_a) + body: Final = ResponsesBody( + model=split.shared_a, input=f"say hello {unique_marker()}", max_output_tokens=64 + ) + answer: Final = unwrap( + proxy.transport.post("/v1/responses", headers=headers, json=body, response_type=ResponsesApiResponse) + ) + assert answer.id, "header-tagged /v1/responses returned no response id" + rows: Final = proxy.poll_logs_for_key(key, min_rows=1) + _assert_served_only_by(rows, CHEAP_SERVED | {split.tier_a}, "header-tagged /v1/responses string input") + + @pytest.mark.covers("reliability.routing.tagged_marker.responses_input_routes_through_marker") + def test_body_tagged_responses_with_list_input_routes_to_the_tier( + self, proxy: ProxyClient, resources: ResourceManager, split: TagSplitDeployments + ) -> None: + """Pins the body-tag and list-input combination of the same split: + /v1/responses with litellm_metadata.tags and structured input items + routes through the tagged marker to its tier.""" + key: Final = _key_for(proxy, resources, [split.shared_a, split.tier_a]) + body: Final = ResponsesBody( + model=split.shared_a, + input=[ResponsesInputItem(role="user", content=f"say hello {unique_marker()}")], + max_output_tokens=64, + litellm_metadata=ResponsesTagMetadata(tags=[split.tag_a]), + ) + answer: Final = unwrap( + proxy.transport.post( + "/v1/responses", + headers=proxy.transport.bearer(key), + json=body, + response_type=ResponsesApiResponse, + ) + ) + assert answer.id, "body-tagged /v1/responses returned no response id" + rows: Final = proxy.poll_logs_for_key(key, min_rows=1) + _assert_served_only_by(rows, CHEAP_SERVED | {split.tier_a}, "body-tagged /v1/responses list input") + + @pytest.mark.covers("reliability.routing.tagged_marker.untagged_request_served_by_plain_deployment") + def test_untagged_responses_is_served_by_the_plain_deployment( + self, proxy: ProxyClient, resources: ResourceManager, split: TagSplitDeployments + ) -> None: + """Pins the untagged half of the /v1/responses tag split: an untagged + request to the shared name is served by the plain deployment, matching + the chat and messages surfaces.""" + key: Final = _key_for(proxy, resources, [split.shared_a, split.tier_a]) + body: Final = ResponsesBody( + model=split.shared_a, input=f"say hello {unique_marker()}", max_output_tokens=64 + ) + answer: Final = unwrap( + proxy.transport.post( + "/v1/responses", + headers=proxy.transport.bearer(key), + json=body, + response_type=ResponsesApiResponse, + ) + ) + assert answer.id, "untagged /v1/responses returned no response id" + rows: Final = proxy.poll_logs_for_key(key, min_rows=1) + _assert_served_only_by(rows, PLAIN_SERVED | {split.shared_a}, "untagged /v1/responses on the shared name") + + +class TestStrategyAliasPricing: + @pytest.mark.covers("reliability.routing.strategy_alias.custom_pricing_ignored") + def test_zero_priced_alias_still_logs_spend_at_the_tier_rate( + self, proxy: ProxyClient, resources: ResourceManager, zero_priced_alias: ZeroPricedAlias + ) -> None: + """Pins GitHub PR #36691: custom pricing registered on a strategy-router + alias never prices the routed request. The alias here carries explicit + zero pricing, so any zero-spend row would prove the alias pricing was + applied; the routed tier deployment's real rate must produce spend > 0.""" + key: Final = _key_for(proxy, resources, [zero_priced_alias.alias, zero_priced_alias.tier]) + chat: Final = unwrap(proxy.chat(key, _hello_chat_body(zero_priced_alias.alias))) + assert chat.choices, "chat through the zero-priced alias returned no choices" + rows: Final = proxy.poll_logs_for_key( + key, min_rows=1, predicate=lambda logged: all((row.spend or 0.0) > 0.0 for row in logged) + ) + _assert_served_only_by(rows, CHEAP_SERVED | {zero_priced_alias.tier}, "chat through the zero-priced alias") + priced: Final = tuple((row.model, row.spend) for row in rows) + assert all((row.spend or 0.0) > 0.0 for row in rows), ( + f"expected spend at the tier deployment's own rate, got zero-spend rows: {priced}" + ) + + +class TestComplexityHeuristicScope: + @pytest.mark.covers("reliability.routing.complexity_heuristic.scores_current_ask_only") + def test_trivial_ask_behind_keyword_heavy_system_prompt_stays_on_the_cheap_tier( + self, proxy: ProxyClient, resources: ResourceManager, heuristic_split: HeuristicSplit + ) -> None: + """Pins GitHub PR #36721: the heuristic complexity classifier scores the + caller's current ask alone. The trivial ask scores SIMPLE on its own, + while the accompanying ~2KB agent system prompt is packed with enough + reasoning and complexity keywords that scoring the combined text lands + in REASONING; only ask-only scoring keeps this on the cheap tier.""" + key: Final = _key_for( + proxy, resources, [heuristic_split.alias, heuristic_split.cheap, heuristic_split.strong] + ) + body: Final = ChatBody( + model=heuristic_split.alias, + messages=[ + ChatMessage(role="system", content=KEYWORD_HEAVY_SYSTEM_PROMPT), + ChatMessage(role="user", content=f"hi {unique_marker()}"), + ], + max_tokens=MAX_TOKENS, + ) + chat: Final = unwrap(proxy.chat(key, body)) + assert chat.choices, "chat through the heuristic router returned no choices" + rows: Final = proxy.poll_logs_for_key(key, min_rows=1) + _assert_served_only_by( + rows, CHEAP_SERVED | {heuristic_split.cheap}, "trivial ask behind a keyword-heavy system prompt" + ) + + +class TestSemanticAutoRouterResponses: + @pytest.mark.covers("reliability.routing.semantic_auto_router.responses_input_routed") + def test_responses_input_reaches_the_semantic_auto_router( + self, proxy: ProxyClient, resources: ResourceManager, semantic_auto_router: SemanticAutoRouter + ) -> None: + """Pins GitHub PR #37333: /v1/responses input is resolved into messages + for the semantic auto-router's pre-routing hook, so the marker embeds + the input, matches its route, and the target deployment serves the + request; before the fix the hook saw no messages and the request + failed with 400 "Unmapped LLM provider auto_router".""" + key: Final = _key_for( + proxy, + resources, + [semantic_auto_router.marker, semantic_auto_router.target, semantic_auto_router.fallback], + ) + body: Final = ResponsesBody( + model=semantic_auto_router.marker, input=SEMANTIC_ROUTE_UTTERANCE, max_output_tokens=64 + ) + answer: Final = unwrap( + proxy.transport.post( + "/v1/responses", + headers=proxy.transport.bearer(key), + json=body, + response_type=ResponsesApiResponse, + ) + ) + assert answer.id, "/v1/responses through the semantic auto-router returned no response id" + rows: Final = proxy.poll_logs_for_key(key, min_rows=1) + _assert_served_only_by( + rows, CHEAP_SERVED | {semantic_auto_router.target}, "semantic auto-router /v1/responses string input" + ) From e9355a7fe9d9e72b894e1dd1fcf82977af57d5df Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 18 Aug 2026 14:49:42 -0700 Subject: [PATCH 17/83] fix(proxy): hand the embeddings failure hook the post-setup request data --- litellm/proxy/proxy_server.py | 12 +----- tests/test_litellm/proxy/test_proxy_server.py | 43 +++++++++++++++++++ 2 files changed, 45 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 5d4a306a73e..b9a2be82969 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -10197,11 +10197,9 @@ async def embeddings( """ global proxy_logging_obj - data: Any = {} + data: Final = await _read_request_body(request=request) + base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: - # Use shared request body reading helper (same as chat/completions) - data = await _read_request_body(request=request) - ### HANDLE TOKEN ARRAY INPUT DECODING ### # This must happen BEFORE base_process_llm_request() since it modifies the input router_model_names: Final = llm_router.model_names if llm_router is not None else [] @@ -10245,10 +10243,6 @@ async def embeddings( if hasattr(user_api_key_dict, "agent_id") and user_api_key_dict.agent_id is not None: data["metadata"]["agent_id"] = user_api_key_dict.agent_id - # Use unified request processor (same as chat/completions and responses) - base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data) - - # Process the request with all optimizations (shared sessions, network tuning, etc.) response: Final = await base_llm_response_processor.base_process_llm_request( request=request, fastapi_response=fastapi_response, @@ -10270,8 +10264,6 @@ async def embeddings( return response except Exception as e: - # Use unified error handler - base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data) raise await base_llm_response_processor._handle_llm_api_exception( e=e, user_api_key_dict=user_api_key_dict, diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 5545ee92e84..7fdfbea843f 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -11098,3 +11098,46 @@ async def test_moderations_reraises_proxy_exception_unwrapped(): assert exc_info.value.code == "400" assert exc_info.value.param == "metadata" mock_logging.post_call_failure_hook.assert_awaited_once() + + +class TestEmbeddingsFailureHookRequestData: + @pytest.mark.asyncio + async def test_failure_hook_gets_post_setup_data_with_logging_obj(self): + """Request setup replaces the processor's data dict (adding the logging + object the failure hook needs to lift token usage from); the embeddings + exception handler must pass that replaced dict, not the raw request body + dict it was rebuilt from.""" + from litellm.proxy._types import ProxyException + + captured = {} + logging_obj_sentinel = MagicMock() + + async def fake_process(self, **kwargs): + self.data = {**self.data, "litellm_logging_obj": logging_obj_sentinel} + captured["processor_data"] = self.data + raise RuntimeError("provider timeout") + + with ( + patch.object( + proxy_server_module, + "_read_request_body", + new=AsyncMock(return_value={"model": "my-embed", "input": "hello"}), + ), + patch.object( + proxy_server_module.ProxyBaseLLMRequestProcessing, + "base_process_llm_request", + new=fake_process, + ), + patch.object(proxy_server_module, "proxy_logging_obj") as mock_logging, + ): + mock_logging.post_call_failure_hook = AsyncMock(return_value=None) + with pytest.raises(ProxyException): + await proxy_server_module.embeddings( + request=MagicMock(), + fastapi_response=MagicMock(), + user_api_key_dict=UserAPIKeyAuth(), + ) + + hook_request_data = mock_logging.post_call_failure_hook.await_args.kwargs["request_data"] + assert hook_request_data is captured["processor_data"] + assert hook_request_data["litellm_logging_obj"] is logging_obj_sentinel From 69ea1c6599214f96e51848403ba024f9186b6ac6 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 18 Aug 2026 15:04:18 -0700 Subject: [PATCH 18/83] fix(proxy): coerce batch candidate counts like the live limiter path --- litellm/proxy/hooks/batch_rate_limiter.py | 9 +-------- litellm/proxy/hooks/parallel_request_limiter_v3.py | 6 +++--- .../proxy/hooks/test_batch_file_validation.py | 7 +++++++ tests/test_litellm/proxy/hooks/test_tpm_concurrent.py | 2 +- 4 files changed, 12 insertions(+), 12 deletions(-) diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index aa3ac6b820b..404de776396 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -429,14 +429,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): ), None, ) - candidate_count: Final = max( - ( - v - for v in (body.get("n"), body.get("best_of")) - if isinstance(v, int) and not isinstance(v, bool) and v > 1 - ), - default=1, - ) + candidate_count: Final = self.parallel_request_limiter.get_output_candidate_count(body) if explicit_cap is not None: try: return max(0, int(explicit_cap)) * candidate_count diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 7cd56e66000..4a8da065fcf 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -750,8 +750,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): 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): + def get_output_candidate_count(data: object, call_type: str | None = None) -> int: + if not isinstance(data, Mapping): return 1 config: Final = ( (data.get("config") if "config" in data else data.get("generationConfig")) @@ -951,7 +951,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): else max(estimated_input_tokens, output_floor) ) - return estimated_input_tokens, max_tokens_estimate * self._get_output_candidate_count(data, call_type) + return estimated_input_tokens, max_tokens_estimate * self.get_output_candidate_count(data, call_type) def _is_redis_cluster(self) -> bool: """ diff --git a/tests/test_litellm/proxy/hooks/test_batch_file_validation.py b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py index 3a9d3ff72aa..31afb5c6019 100644 --- a/tests/test_litellm/proxy/hooks/test_batch_file_validation.py +++ b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py @@ -1992,9 +1992,13 @@ def _output_estimator(): the no-`max_tokens` floor mock returns a distinctive sentinel so tests can tell "floor was used" apart from "an explicit cap was read".""" from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + _PROXY_MaxParallelRequestsHandler_v3, + ) limiter = MagicMock() limiter.no_max_tokens_output_floor.return_value = 999 + limiter.get_output_candidate_count = _PROXY_MaxParallelRequestsHandler_v3.get_output_candidate_count return _PROXY_BatchRateLimiter( internal_usage_cache=MagicMock(), parallel_request_limiter=limiter, @@ -2070,6 +2074,9 @@ def test_estimate_entry_output_tokens_prefers_max_tokens_over_max_output_tokens( ({"max_tokens": 40, "n": 3, "best_of": 5}, 200), ({"max_tokens": 40, "n": 0}, 40), ({"max_tokens": 40, "n": -2}, 40), + ({"max_tokens": 40, "n": 5.0}, 200), + ({"max_tokens": 40, "n": "10"}, 400), + ({"max_tokens": 40, "n": "not-a-number"}, 40), ({"n": 3}, 2997), ], ) diff --git a/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py b/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py index e12610b769c..e08b18c1ae9 100644 --- a/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py +++ b/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py @@ -3384,7 +3384,7 @@ def test_split_quota_helpers_handle_non_mapping_inputs(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_output_candidate_count(object()) == 1 assert ( handler._get_explicit_output_cap({"max_output_tokens": []}, "responses") is None ) From d4db4b1379b943fef632a929f165ed08c2913c15 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 18 Aug 2026 15:14:05 -0700 Subject: [PATCH 19/83] test(e2e): pin marker alias connection params staying off the routed tier --- tests/e2e/coverage_registry/reliability.yaml | 1 + .../test_auto_router_regressions_e2e.py | 47 +++++++++++++++++++ 2 files changed, 48 insertions(+) diff --git a/tests/e2e/coverage_registry/reliability.yaml b/tests/e2e/coverage_registry/reliability.yaml index 5a1d437eac9..b50551ec105 100644 --- a/tests/e2e/coverage_registry/reliability.yaml +++ b/tests/e2e/coverage_registry/reliability.yaml @@ -24,6 +24,7 @@ - {id: reliability.routing.tagged_marker.untagged_tier_deployments_still_served, module: reliability, tier: P1, behavior: routing, variant: tagged_marker, assertions: [untagged_tier_deployments_still_served], exercised_on: [chat_completions, messages], source: "litellm/router_strategy/tag_based_routing.py:433", rationale: "Routing tags the marker consumed no longer constrain deployment selection inside the routed tier group, so untagged tier deployments serve the rewrite (GitHub issue #36621)"} - {id: reliability.routing.tagged_marker.tag_semantics_stay_strict, module: reliability, tier: P1, behavior: routing, variant: tagged_marker, assertions: [tag_semantics_stay_strict], exercised_on: [chat_completions], source: "litellm/router_strategy/tag_based_routing.py:299", rationale: "Tag consumption must not loosen strict semantics: a tagged call aimed straight at an untagged deployment still gets the 401 tags-configuration denial"} - {id: reliability.routing.tagged_marker.responses_input_routes_through_marker, module: reliability, tier: P0, behavior: routing, variant: tagged_marker, assertions: [responses_input_routes_through_marker], exercised_on: [responses], source: "litellm/router.py:11489", rationale: "Tagged /v1/responses (header or litellm_metadata.tags, string or list input) routes through the marker to its tier, extending the GitHub issues #36620/#36621 tag split to the Responses surface"} +- {id: reliability.routing.tagged_marker.alias_connection_params_stay_with_tier, module: reliability, tier: P0, behavior: routing, variant: tagged_marker, assertions: [alias_connection_params_stay_with_tier], exercised_on: [chat_completions], source: "litellm/router.py:11567", rationale: "An api_key or api_base on the marker alias is never forwarded onto the routed request; the tier deployment calls its provider with its own credential (GitHub PR #36626)"} - {id: reliability.routing.semantic_auto_router.responses_input_routed, module: reliability, tier: P0, behavior: routing, variant: semantic_auto_router, assertions: [responses_input_routed], exercised_on: [responses], source: "litellm/router_strategy/auto_router/auto_router.py:131", fail_before_fix: proven, rationale: "/v1/responses input is resolved into messages for the semantic auto-router pre-routing hook instead of failing 400 Unmapped LLM provider auto_router (GitHub PR #37333)"} - {id: reliability.routing.strategy_alias.custom_pricing_ignored, module: reliability, tier: P1, behavior: routing, variant: strategy_alias, assertions: [custom_pricing_ignored], exercised_on: [chat_completions], source: "litellm/router.py:11489", rationale: "Custom pricing on a strategy-router alias never prices the routed request; spend logs at the routed tier deployment's own rate (GitHub PR #36691)"} - {id: reliability.routing.complexity_heuristic.scores_current_ask_only, module: reliability, tier: P1, behavior: routing, variant: complexity_heuristic, assertions: [scores_current_ask_only], exercised_on: [chat_completions], source: "router_strategy/complexity_router/complexity_router.py:942", rationale: "The heuristic complexity classifier scores the caller's current ask only, so a keyword-heavy agent system prompt cannot inflate the tier (GitHub PR #36721)"} diff --git a/tests/e2e/router/test_auto_router_regressions_e2e.py b/tests/e2e/router/test_auto_router_regressions_e2e.py index ee68d3fa1bb..a23b50f2a73 100644 --- a/tests/e2e/router/test_auto_router_regressions_e2e.py +++ b/tests/e2e/router/test_auto_router_regressions_e2e.py @@ -26,6 +26,9 @@ pins one fixed behavior: request; spend logs at the routed tier deployment's own rate. - GitHub PR #36721: the heuristic complexity classifier scores the caller's current ask only, so a large agent system prompt cannot inflate the tier. +- GitHub PR #36626: connection params on the marker alias (``api_key``, + ``api_base``) stay with the alias; the routed tier calls its provider with + its own credentials. Every deployment is registered via /model/new (stage has no static config for these) and ``enable_tag_filtering`` is flipped through /config/update and @@ -171,6 +174,12 @@ class SemanticAutoRouter: embedding: str +@dataclass(frozen=True, slots=True) +class CredentialedAlias: + alias: str + tier: str + + def _provider_key(env_var: str) -> str: return os.environ.get(env_var) or f"os.environ/{env_var}" @@ -382,6 +391,27 @@ def semantic_auto_router(proxy: ProxyClient) -> Iterator[SemanticAutoRouter]: proxy.delete_model(model_id) +@pytest.fixture(scope="module") +def credentialed_alias(proxy: ProxyClient) -> Iterator[CredentialedAlias]: + marker: Final = unique_marker() + named: Final = CredentialedAlias(alias=f"e2e-cred-alias-{marker}", tier=f"e2e-cred-tier-{marker}") + alias_params: Final = LiteLLMParamsBody( + model="auto_router/complexity_router", + complexity_router_config=_uniform_tier_config(named.tier), + api_key=f"sk-alias-never-used-{marker}", + ) + registrations: Final[tuple[tuple[str, LiteLLMParamsBody], ...]] = ( + (named.tier, LiteLLMParamsBody(model=CHEAP_MODEL, api_key=_provider_key("ANTHROPIC_API_KEY"))), + (named.alias, alias_params), + ) + created: Final = tuple(proxy.create_model(name, params) for name, params in registrations) + try: + yield named + finally: + for model_id in created: + proxy.delete_model(model_id) + + class TestTagSplitRouting: @pytest.mark.covers("reliability.routing.tagged_marker.request_tag_selects_marker") def test_body_tagged_chat_routes_through_the_marker_to_its_tier( @@ -630,3 +660,20 @@ class TestSemanticAutoRouterResponses: _assert_served_only_by( rows, CHEAP_SERVED | {semantic_auto_router.target}, "semantic auto-router /v1/responses string input" ) + + +class TestAliasParamForwarding: + @pytest.mark.covers("reliability.routing.tagged_marker.alias_connection_params_stay_with_tier") + def test_alias_api_key_never_overrides_the_tier_credential( + self, proxy: ProxyClient, resources: ResourceManager, credentialed_alias: CredentialedAlias + ) -> None: + """Pins GitHub PR #36626: an api_key set on the marker alias entry is + never forwarded onto the routed request, so the tier deployment calls + its provider with its own credential. Before the fix the alias's key + was copied into the request, overriding the tier's credential, and + every routed call failed provider auth.""" + key: Final = _key_for(proxy, resources, [credentialed_alias.alias, credentialed_alias.tier]) + chat: Final = unwrap(proxy.chat(key, _hello_chat_body(credentialed_alias.alias))) + assert chat.choices, "chat through the credentialed alias returned no choices" + rows: Final = proxy.poll_logs_for_key(key, min_rows=1) + _assert_served_only_by(rows, CHEAP_SERVED | {credentialed_alias.tier}, "chat through the credentialed alias") From 645b87fae1c79eff67d08ee7033286ddd0ac7fbd Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 18 Aug 2026 22:40:27 +0000 Subject: [PATCH 20/83] fix(types): map nested prompt_tokens_details.cache_creation_input_tokens to cache_write_tokens Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/types/utils.py | 11 ++++- .../test_dashscope_cost_calculator.py | 41 +++++++++++++++++++ tests/test_litellm/types/test_types_utils.py | 23 +++++++++++ 3 files changed, 74 insertions(+), 1 deletion(-) diff --git a/litellm/types/utils.py b/litellm/types/utils.py index cd2ef9dde2c..b97fc2b3047 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -1627,8 +1627,17 @@ class PromptTokensDetailsWrapper( def __init__(self, *args, **kwargs) -> None: super().__init__(*args, **kwargs) + nested_cache_creation_input_tokens: Final = (self.model_extra or {}).get("cache_creation_input_tokens") self.cache_write_tokens = ( - self.cache_write_tokens if self.cache_write_tokens is not None else self.cache_creation_tokens + self.cache_write_tokens + if self.cache_write_tokens is not None + else ( + self.cache_creation_tokens + if self.cache_creation_tokens is not None + else ( + nested_cache_creation_input_tokens if isinstance(nested_cache_creation_input_tokens, int) else None + ) + ) ) if self.character_count is None: del self.character_count diff --git a/tests/test_litellm/llms/dashscope/test_dashscope_cost_calculator.py b/tests/test_litellm/llms/dashscope/test_dashscope_cost_calculator.py index 6f5aaabae06..510776ddfdf 100644 --- a/tests/test_litellm/llms/dashscope/test_dashscope_cost_calculator.py +++ b/tests/test_litellm/llms/dashscope/test_dashscope_cost_calculator.py @@ -271,6 +271,47 @@ class TestDashscopeCostCalculator: assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10) + def test_dashscope_nested_cache_creation_input_tokens_bill_at_cache_write_rate(self): + """ + Regression (LIT-5757): DashScope nests cache_creation_input_tokens inside + prompt_tokens_details; those tokens must bill at the tier's cache-creation + rate instead of being folded into text tokens at the input rate. + """ + self._register_tiered_model( + "dashscope/qwen-nested-cache-write-test", + [ + { + "range": [0, 128000], + "input_cost_per_token": 4e-07, + "cache_read_input_token_cost": 1.6e-07, + "cache_creation_input_token_cost": 5e-07, + "output_cost_per_token": 1.6e-06, + } + ], + ) + + usage = Usage( + prompt_tokens=2059, + completion_tokens=201, + total_tokens=2260, + prompt_tokens_details={ + "cached_tokens": 0, + "text_tokens": 2059, + "cache_type": "ephemeral", + "cache_creation_input_tokens": 2048, + "cache_creation": {"ephemeral_5m_input_tokens": 2048}, + }, + completion_tokens_details={"reasoning_tokens": 170}, + ) + + prompt_cost, _ = dashscope_cost_per_token( + model="qwen-nested-cache-write-test", usage=usage + ) + + assert math.isclose( + prompt_cost, (2048 * 5e-07) + (11 * 4e-07), rel_tol=1e-10 + ) + def test_dashscope_tiered_cache_creation_falls_back_to_tier_input_rate(self): """ Tiers without a cache_creation_input_token_cost bill cache-creation tokens at diff --git a/tests/test_litellm/types/test_types_utils.py b/tests/test_litellm/types/test_types_utils.py index cd5e8dda012..52547e064ac 100644 --- a/tests/test_litellm/types/test_types_utils.py +++ b/tests/test_litellm/types/test_types_utils.py @@ -75,6 +75,29 @@ def test_usage_dump(): assert new_usage.prompt_tokens_details.web_search_requests == 1 +def test_prompt_tokens_details_maps_nested_cache_creation_input_tokens(): + """Regression (LIT-5757): DashScope nests the Anthropic-spelled + cache_creation_input_tokens inside prompt_tokens_details. It must populate + the canonical cache_write_tokens/cache_creation_tokens pair, without + overriding an explicitly provided canonical value.""" + from litellm.types.utils import PromptTokensDetailsWrapper + + nested = PromptTokensDetailsWrapper( + cached_tokens=0, text_tokens=2059, cache_creation_input_tokens=2048 + ) + assert nested.cache_write_tokens == 2048 + assert nested.cache_creation_tokens == 2048 + + explicit = PromptTokensDetailsWrapper( + cache_write_tokens=100, cache_creation_input_tokens=2048 + ) + assert explicit.cache_write_tokens == 100 + assert explicit.cache_creation_tokens == 100 + + non_int = PromptTokensDetailsWrapper(cache_creation_input_tokens=None) + assert not hasattr(non_int, "cache_write_tokens") + + def test_usage_server_tool_use_dict_is_coerced_and_round_trips(): from litellm.types.utils import ServerToolUse, Usage From 3c34c344597b9fb3d1d30e3e0f97924e47476a4b Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 18 Aug 2026 15:42:02 -0700 Subject: [PATCH 21/83] fix(proxy): guard candidate-count and batch cap coercion against float overflow --- litellm/proxy/hooks/batch_rate_limiter.py | 2 +- litellm/proxy/hooks/parallel_request_limiter_v3.py | 2 +- tests/test_litellm/proxy/hooks/test_batch_file_validation.py | 2 ++ tests/test_litellm/proxy/hooks/test_tpm_concurrent.py | 1 + 4 files changed, 5 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index 404de776396..d6229fb80a6 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -433,7 +433,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): if explicit_cap is not None: try: return max(0, int(explicit_cap)) * candidate_count - except (TypeError, ValueError): + except (TypeError, ValueError, OverflowError): pass return self.parallel_request_limiter.no_max_tokens_output_floor(min_configured_otpm_limit) * candidate_count diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 4a8da065fcf..2d799ded752 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -768,7 +768,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): for value in candidate_values: try: candidate_count = max(candidate_count, int(value or 1)) - except (TypeError, ValueError): + except (TypeError, ValueError, OverflowError): continue return candidate_count diff --git a/tests/test_litellm/proxy/hooks/test_batch_file_validation.py b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py index 31afb5c6019..1ce1a2f3e51 100644 --- a/tests/test_litellm/proxy/hooks/test_batch_file_validation.py +++ b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py @@ -2077,6 +2077,8 @@ def test_estimate_entry_output_tokens_prefers_max_tokens_over_max_output_tokens( ({"max_tokens": 40, "n": 5.0}, 200), ({"max_tokens": 40, "n": "10"}, 400), ({"max_tokens": 40, "n": "not-a-number"}, 40), + ({"max_tokens": 40, "n": 1e309}, 40), + ({"max_tokens": 1e309, "n": 3}, 2997), ({"n": 3}, 2997), ], ) diff --git a/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py b/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py index e08b18c1ae9..8d03857c917 100644 --- a/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py +++ b/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py @@ -3385,6 +3385,7 @@ def test_split_quota_helpers_handle_non_mapping_inputs(rate_limiter): 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_output_candidate_count({"n": 1e309}) == 1 assert ( handler._get_explicit_output_cap({"max_output_tokens": []}, "responses") is None ) From 0585c45cd935954c723eb6d5c2a45c9f346ec7d6 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 18 Aug 2026 22:52:02 +0000 Subject: [PATCH 22/83] fix(types): avoid mutable dict literal in nested cache token lookup Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/types/utils.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/litellm/types/utils.py b/litellm/types/utils.py index b97fc2b3047..82f60f94656 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -1627,7 +1627,10 @@ class PromptTokensDetailsWrapper( def __init__(self, *args, **kwargs) -> None: super().__init__(*args, **kwargs) - nested_cache_creation_input_tokens: Final = (self.model_extra or {}).get("cache_creation_input_tokens") + extra_fields: Final = self.model_extra + nested_cache_creation_input_tokens: Final = ( + extra_fields.get("cache_creation_input_tokens") if extra_fields is not None else None + ) self.cache_write_tokens = ( self.cache_write_tokens if self.cache_write_tokens is not None From 42ddc5c5359bb32762926577fca8fcb1e5b3836d Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 18 Aug 2026 16:18:23 -0700 Subject: [PATCH 23/83] fix(proxy): estimate image message tokens without fetching the image url --- litellm/proxy/utils.py | 4 ++- tests/test_litellm/proxy/test_proxy_utils.py | 28 ++++++++++++++++++++ 2 files changed, 31 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 1465e76d01f..c45b4ad17c0 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -430,7 +430,9 @@ def _count_request_input_tokens(model: str, request_input: object, system_input: text_entries: Final = tuple(entry for entry in request_input if isinstance(entry, str)) if len(text_entries) == len(request_input): return system_tokens + litellm.token_counter(model=model, text="".join(text_entries)) - return system_tokens + litellm.token_counter(model=model, messages=request_input) + return system_tokens + litellm.token_counter( + model=model, messages=request_input, use_default_image_token_count=True + ) def _estimate_dispatched_failure_usage(model: str, request_input: object, system_input: object) -> Usage | None: diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index 96e2c74e5d4..b70f93054d2 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -635,6 +635,34 @@ class TestPostCallFailureHookEstimatesDispatchedInputTokens: "metadata": {}, } + @pytest.mark.asyncio + async def test_image_message_estimated_without_fetching_image(self): + import litellm as litellm_module + from litellm.types.utils import Usage + + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "describe this image"}, + { + "type": "image_url", + "image_url": {"url": "http://127.0.0.1:1/unreachable.png", "detail": "high"}, + }, + ], + } + ] + request_data = self._dispatched_request_data(messages, {}) + await self._run(request_data) + + estimated = request_data["combined_usage_object"] + assert isinstance(estimated, Usage) + expected = litellm_module.token_counter( + model="gpt-3.5-turbo", messages=messages, use_default_image_token_count=True + ) + assert estimated.prompt_tokens == expected + assert estimated.prompt_tokens > 0 + @pytest.mark.asyncio async def test_embedding_string_list_input_counted_in_estimate(self): import litellm as litellm_module From 4a43b5080015203785cbca8a0559449ae6bd6b6e Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 18 Aug 2026 23:19:38 +0000 Subject: [PATCH 24/83] test: add Final annotations to LIT-5757 regression test variables Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/test_litellm/types/test_types_utils.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/tests/test_litellm/types/test_types_utils.py b/tests/test_litellm/types/test_types_utils.py index 52547e064ac..672aa84cc73 100644 --- a/tests/test_litellm/types/test_types_utils.py +++ b/tests/test_litellm/types/test_types_utils.py @@ -1,5 +1,6 @@ import os import sys +from typing import Final import pytest @@ -82,19 +83,19 @@ def test_prompt_tokens_details_maps_nested_cache_creation_input_tokens(): overriding an explicitly provided canonical value.""" from litellm.types.utils import PromptTokensDetailsWrapper - nested = PromptTokensDetailsWrapper( + nested: Final = PromptTokensDetailsWrapper( cached_tokens=0, text_tokens=2059, cache_creation_input_tokens=2048 ) assert nested.cache_write_tokens == 2048 assert nested.cache_creation_tokens == 2048 - explicit = PromptTokensDetailsWrapper( + explicit: Final = PromptTokensDetailsWrapper( cache_write_tokens=100, cache_creation_input_tokens=2048 ) assert explicit.cache_write_tokens == 100 assert explicit.cache_creation_tokens == 100 - non_int = PromptTokensDetailsWrapper(cache_creation_input_tokens=None) + non_int: Final = PromptTokensDetailsWrapper(cache_creation_input_tokens=None) assert not hasattr(non_int, "cache_write_tokens") From 47f3cf804ed917db460a59605c0af419a6f9c54a Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 18 Aug 2026 16:19:43 -0700 Subject: [PATCH 25/83] fix(router): honor request-level tag filtering in pre-routing strategy selection Key and team router_settings set enable_tag_filtering on the request kwargs, and get_deployments_for_tag already treats that as authoritative, but _select_pre_routing_strategy only consulted the router-wide flag, so tagged auto-router markers still captured untagged requests from keys that enabled filtering. The e2e auto-router module now enables tag filtering through key-level router_settings instead of flipping /config/update module-wide, which was denying concurrently running tagged requests from other suites on the shared per-build CI proxy. --- litellm/router.py | 9 +- tests/e2e/models.py | 13 ++- .../test_auto_router_regressions_e2e.py | 109 ++++-------------- tests/test_litellm/test_router.py | 15 +++ 4 files changed, 52 insertions(+), 94 deletions(-) diff --git a/litellm/router.py b/litellm/router.py index 8b9c4b0db1a..85eb2dac51e 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -11454,8 +11454,10 @@ class Router: deployment the strategy was registered from via its (model_name, tags) pair. - With tag filtering enabled, strategies that all carry real tags matching - none of the request's do not capture it when the name also has plain + With tag filtering enabled, router-wide or by the request's + enable_tag_filtering (which the proxy sets from key/team + router_settings), strategies that all carry real tags matching none of + the request's do not capture it when the name also has plain deployments: returning None hands the request to ordinary tag-aware deployment selection. """ @@ -11478,8 +11480,9 @@ class Router: for tagged in candidates: if "default" in tagged.tags: return tagged + request_scoped_filtering: Final = request_kwargs.get("enable_tag_filtering") is True if ( - self.enable_tag_filtering + (self.enable_tag_filtering or request_scoped_filtering) and all(tagged.tags for tagged in candidates) and self._model_name_has_plain_deployments(model) ): diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 957da605546..df5cb841fad 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -73,6 +73,7 @@ class KeyGenerateBody(BaseModel): allowed_passthrough_routes: list[str] | None = None metadata: KeyMetadata | None = None object_permission: ObjectPermission | None = None + router_settings: "RouterSettingsOverride | None" = None class KeyGenerateResponse(BaseModel): @@ -234,16 +235,18 @@ class ChatBody(BaseModel): class RouterSettingsOverride(BaseModel): - """Per-request `router_settings_override` in a /chat/completions body: the - reliability knobs (fallbacks by trigger, retry count) the reliability suite - drives per call instead of via static router config. Serialized exclude_none, so - an override sets only the strategies a test exercises. Each fallbacks map is - model_name -> the ordered fallback model_names to try.""" + """Router settings a test scopes below the global config: sent per request as + `router_settings_override` in a /chat/completions body (the reliability suite's + fallback and retry knobs) or stored on a key as `router_settings` at + /key/generate (the auto-router suite's tag filtering switch). Serialized + exclude_none, so an override sets only the knobs a test exercises. Each + fallbacks map is model_name -> the ordered fallback model_names to try.""" fallbacks: list[dict[str, list[str]]] | None = None context_window_fallbacks: list[dict[str, list[str]]] | None = None content_policy_fallbacks: list[dict[str, list[str]]] | None = None num_retries: int | None = None + enable_tag_filtering: bool | None = None class ReliabilityChatBody(ChatBody): diff --git a/tests/e2e/router/test_auto_router_regressions_e2e.py b/tests/e2e/router/test_auto_router_regressions_e2e.py index a23b50f2a73..c6ef9cda05d 100644 --- a/tests/e2e/router/test_auto_router_regressions_e2e.py +++ b/tests/e2e/router/test_auto_router_regressions_e2e.py @@ -31,15 +31,15 @@ pins one fixed behavior: its own credentials. Every deployment is registered via /model/new (stage has no static config for -these) and ``enable_tag_filtering`` is flipped through /config/update and -restored on teardown, mirroring TestRouterSettings in the management suite. +these) and ``enable_tag_filtering`` is enabled through key-level +``router_settings`` on the keys the tag tests mint, so the switch rides only +this module's own requests and the rest of the suite is never filtered. The served deployment is always read back from the spend log's ``model``, which stores either the registered alias or the provider-prefixed form. """ import json import os -import time from collections.abc import Iterator from dataclasses import dataclass from typing import Final @@ -48,7 +48,7 @@ import pytest from pydantic import BaseModel, ConfigDict, Field from e2e_config import unique_marker -from e2e_http import AnthropicHeaders, AuthHeaders, NoBody, UnauthorizedError, unwrap +from e2e_http import AnthropicHeaders, AuthHeaders, UnauthorizedError, unwrap from lifecycle import ResourceManager from models import ( AnthropicMessagesBody, @@ -58,6 +58,7 @@ from models import ( ChatMetadata, KeyGenerateBody, LiteLLMParamsBody, + RouterSettingsOverride, SpendLogRow, ) from proxy_client import ProxyClient @@ -117,26 +118,6 @@ class ResponsesApiResponse(BaseModel): model: str | None = None -class RouterSettingsPatch(BaseModel): - enable_tag_filtering: bool - - -class ConfigUpdateBody(BaseModel): - router_settings: RouterSettingsPatch - - -class ConfigUpdateResponse(BaseModel): - message: str - - -class RouterCurrentValues(BaseModel): - enable_tag_filtering: bool | None = None - - -class RouterSettingsResponse(BaseModel): - current_values: RouterCurrentValues - - @dataclass(frozen=True, slots=True) class TagSplitDeployments: """Scenario A mirrors the customer-shaped config from GitHub issue #36619: @@ -191,44 +172,16 @@ def _uniform_tier_config(tier_model: str) -> dict[str, object]: } -def _read_tag_filtering(proxy: ProxyClient) -> bool | None: - return unwrap( - proxy.transport.get( - "/router/settings", - headers=proxy.transport.master, - params=NoBody(), - response_type=RouterSettingsResponse, - ) - ).current_values.enable_tag_filtering - - -def _write_tag_filtering(proxy: ProxyClient, enabled: bool) -> None: - response: Final = unwrap( - proxy.transport.post( - "/config/update", - headers=proxy.transport.master, - json=ConfigUpdateBody(router_settings=RouterSettingsPatch(enable_tag_filtering=enabled)), - response_type=ConfigUpdateResponse, +def _key_for( + proxy: ProxyClient, resources: ResourceManager, models: list[str], tag_filtering: bool = False +) -> str: + key: Final = proxy.generate_key( + KeyGenerateBody( + models=models, + user_id="e2e-auto-router-regressions", + router_settings=RouterSettingsOverride(enable_tag_filtering=True) if tag_filtering else None, ) ) - assert "success" in response.message.lower(), ( - f"/config/update reported {response.message!r}, expected a success message" - ) - - -def _await_tag_filtering(proxy: ProxyClient, expected: bool) -> None: - deadline: Final = time.monotonic() + proxy.poll_timeout - while time.monotonic() < deadline: - if _read_tag_filtering(proxy) is expected: - return - time.sleep(proxy.poll_interval) - raise AssertionError( - f"GET /router/settings never reported enable_tag_filtering={expected} after /config/update" - ) - - -def _key_for(proxy: ProxyClient, resources: ResourceManager, models: list[str]) -> str: - key: Final = proxy.generate_key(KeyGenerateBody(models=models, user_id="e2e-auto-router-regressions")) resources.defer(lambda: proxy.delete_key(key)) return key @@ -258,23 +211,7 @@ def _assert_served_only_by(rows: list[SpendLogRow], allowed: frozenset[str], con @pytest.fixture(scope="module") -def tag_filtering(proxy: ProxyClient) -> Iterator[None]: - """enable_tag_filtering is what splits tagged from untagged traffic in every - scenario here. /config/update is the only write path for router_settings; - the original value is restored on teardown so the shared proxy keeps its - configuration for the rest of the run.""" - original: Final = bool(_read_tag_filtering(proxy)) - _write_tag_filtering(proxy, True) - _await_tag_filtering(proxy, True) - try: - yield - finally: - _write_tag_filtering(proxy, original) - _await_tag_filtering(proxy, original) - - -@pytest.fixture(scope="module") -def split(proxy: ProxyClient, tag_filtering: None) -> Iterator[TagSplitDeployments]: +def split(proxy: ProxyClient) -> Iterator[TagSplitDeployments]: marker: Final = unique_marker() deployments: Final = TagSplitDeployments( tag_a=f"e2e-split-a-{marker}", @@ -421,7 +358,7 @@ class TestTagSplitRouting: body metadata tags match the tagged marker under a shared model name is answered by the marker's tier deployment, not by the plain deployment that was registered under the name first.""" - key: Final = _key_for(proxy, resources, [split.shared_a, split.tier_a]) + key: Final = _key_for(proxy, resources, [split.shared_a, split.tier_a], tag_filtering=True) chat: Final = unwrap(proxy.chat(key, _hello_chat_body(split.shared_a, tags=[split.tag_a]))) assert chat.choices, "tagged chat through the shared name returned no choices" rows: Final = proxy.poll_logs_for_key(key, min_rows=1) @@ -435,7 +372,7 @@ class TestTagSplitRouting: succeed on every call and are all served by the plain deployment; the tagged marker never captures them, so no intermittent auto-router errors and no tier hijacking.""" - key: Final = _key_for(proxy, resources, [split.shared_a, split.tier_a]) + key: Final = _key_for(proxy, resources, [split.shared_a, split.tier_a], tag_filtering=True) for _ in range(5): chat = unwrap(proxy.chat(key, _hello_chat_body(split.shared_a))) assert chat.choices, "untagged chat through the shared name returned no choices" @@ -449,7 +386,7 @@ class TestTagSplitRouting: """Pins GitHub issue #36620 on the /v1/messages surface: an untagged Anthropic-native request to the shared name is served by the plain deployment, not captured by the tagged marker.""" - key: Final = _key_for(proxy, resources, [split.shared_a, split.tier_a]) + key: Final = _key_for(proxy, resources, [split.shared_a, split.tier_a], tag_filtering=True) answer: Final = unwrap(proxy.messages(key, _hello_messages_body(split.shared_a))) assert answer.content or answer.choices, "untagged /v1/messages returned neither content nor choices" rows: Final = proxy.poll_logs_for_key(key, min_rows=1) @@ -465,7 +402,7 @@ class TestUntaggedTierDeployments: x-litellm-tags header selects the tagged marker, and the rewrite still lands on the tier deployment even though that deployment carries no tags, because the marker consumed the routing tags.""" - key: Final = _key_for(proxy, resources, [split.shared_b, split.tier_b]) + key: Final = _key_for(proxy, resources, [split.shared_b, split.tier_b], tag_filtering=True) headers: Final = TaggedAnthropicHeaders(authorization=f"Bearer {key}", x_litellm_tags=split.tag_b) answer: Final = unwrap( proxy.transport.post( @@ -487,7 +424,7 @@ class TestUntaggedTierDeployments: tagged marker rewrites the request to its tier model, the consumed routing tags no longer constrain deployment selection, so the untagged tier deployment serves the request instead of a strict-tag denial.""" - key: Final = _key_for(proxy, resources, [split.shared_b, split.tier_b]) + key: Final = _key_for(proxy, resources, [split.shared_b, split.tier_b], tag_filtering=True) chat: Final = unwrap(proxy.chat(key, _hello_chat_body(split.shared_b, tags=[split.tag_b]))) assert chat.choices, "body-tagged chat through the marker-first shared name returned no choices" rows: Final = proxy.poll_logs_for_key(key, min_rows=1) @@ -500,7 +437,7 @@ class TestUntaggedTierDeployments: """The tag-consumption fix must not loosen strict tag semantics: a tagged request aimed directly at an untagged deployment (no marker involved) is still rejected with the 401 tags-configuration error.""" - key: Final = _key_for(proxy, resources, [split.tier_b]) + key: Final = _key_for(proxy, resources, [split.tier_b], tag_filtering=True) result: Final = proxy.chat(key, _hello_chat_body(split.tier_b, tags=[split.tag_b])) assert isinstance(result, UnauthorizedError), ( f"expected the tagged direct call to an untagged deployment to be denied with 401, got {result}" @@ -516,7 +453,7 @@ class TestResponsesApiTagRouting: #36620/#36621): a /v1/responses request with string input, tagged via the x-litellm-tags header, succeeds and routes through the tagged marker to its tier.""" - key: Final = _key_for(proxy, resources, [split.shared_a, split.tier_a]) + key: Final = _key_for(proxy, resources, [split.shared_a, split.tier_a], tag_filtering=True) headers: Final = TaggedAuthHeaders(authorization=f"Bearer {key}", x_litellm_tags=split.tag_a) body: Final = ResponsesBody( model=split.shared_a, input=f"say hello {unique_marker()}", max_output_tokens=64 @@ -535,7 +472,7 @@ class TestResponsesApiTagRouting: """Pins the body-tag and list-input combination of the same split: /v1/responses with litellm_metadata.tags and structured input items routes through the tagged marker to its tier.""" - key: Final = _key_for(proxy, resources, [split.shared_a, split.tier_a]) + key: Final = _key_for(proxy, resources, [split.shared_a, split.tier_a], tag_filtering=True) body: Final = ResponsesBody( model=split.shared_a, input=[ResponsesInputItem(role="user", content=f"say hello {unique_marker()}")], @@ -561,7 +498,7 @@ class TestResponsesApiTagRouting: """Pins the untagged half of the /v1/responses tag split: an untagged request to the shared name is served by the plain deployment, matching the chat and messages surfaces.""" - key: Final = _key_for(proxy, resources, [split.shared_a, split.tier_a]) + key: Final = _key_for(proxy, resources, [split.shared_a, split.tier_a], tag_filtering=True) body: Final = ResponsesBody( model=split.shared_a, input=f"say hello {unique_marker()}", max_output_tokens=64 ) diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index b3c348a1221..16a309ebb4c 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -7674,6 +7674,21 @@ class TestTaggedAutoRouterOnSharedModelName: assert response is not None assert response.model == "gemini-flash" + @pytest.mark.asyncio + async def test_request_level_tag_filtering_from_key_settings_bypasses_the_marker(self): + router = self._router(marker_tags=["route"], include_plain_sibling=True, enable_tag_filtering=False) + + assert await self._hook_response(router, {"enable_tag_filtering": True}) is None + + @pytest.mark.asyncio + async def test_globally_disabled_filtering_still_lets_the_marker_capture_untagged_requests(self): + router = self._router(marker_tags=["route"], include_plain_sibling=True, enable_tag_filtering=False) + + response = await self._hook_response(router, {}) + + assert response is not None + assert response.model == "gemini-flash" + @pytest.mark.asyncio async def test_marker_only_alias_still_captures_untagged_requests(self): router = self._router(marker_tags=["route"], include_plain_sibling=False, enable_tag_filtering=True) From a30e1f6e3d718f2b20fd6f62073fd9c9a6b4ab7d Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 18 Aug 2026 16:41:40 -0700 Subject: [PATCH 26/83] fix(mcp): bind tool existence check to the selected server --- .../mcp_server/mcp_server_manager.py | 45 ++++++------ .../mcp_server/test_mcp_server.py | 2 +- .../mcp_server/test_mcp_server_manager.py | 72 +++++++++++++++++++ 3 files changed, 95 insertions(+), 24 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 7fff6c12fe0..9835c7f01ed 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -2321,23 +2321,30 @@ class MCPServerManager: openapi_key_prefix: Final = prefix_root + MCP_TOOL_PREFIX_SEPARATOR global_mcp_tool_registry.unregister_tools_with_prefix(openapi_key_prefix) - owned_raw: Final[set[str]] = set() - for p in iter_known_server_prefixes(server): - if p: - owned_raw.add(p) - if server.name: - owned_raw.add(server.name) + owned_normalized: Final = self._owned_mapping_values(server) - owned_normalized: Final = {normalize_server_name(x) for x in owned_raw} - - stale_mapping_keys: Final[list[str]] = [] - for tool_name, mapped_server in list(self.tool_name_to_mcp_server_name_mapping.items()): - if mapped_server in owned_raw or normalize_server_name(str(mapped_server)) in owned_normalized: - stale_mapping_keys.append(tool_name) + stale_mapping_keys: Final = tuple( + tool_name + for tool_name, mapped_server in self.tool_name_to_mcp_server_name_mapping.items() + if normalize_server_name(str(mapped_server)) in owned_normalized + ) for key in stale_mapping_keys: del self.tool_name_to_mcp_server_name_mapping[key] + def _owned_mapping_values(self, server: MCPServer) -> frozenset[str]: + return frozenset( + normalize_server_name(value) for value in (*iter_known_server_prefixes(server), server.name) if value + ) + + def _server_exposes_tool(self, server: MCPServer, tool_name: str) -> bool: + owned: Final = self._owned_mapping_values(server) + mapped_owners: Final = ( + self.tool_name_to_mcp_server_name_mapping.get(spelling) + for spelling in iter_known_tool_name_spellings(tool_name, server) + ) + return any(owner is not None and normalize_server_name(owner) in owned for owner in mapped_owners) + def remove_server(self, mcp_server: LiteLLM_MCPServerTable): """ Remove a server from the registry @@ -5463,13 +5470,8 @@ class MCPServerManager: if mcp_server is None: raise ValueError(f"Tool {name} not found") - if resolved_by_server_name_only: - tool_known: Final = ( - name in self.tool_name_to_mcp_server_name_mapping - or prefixed_tool_name in self.tool_name_to_mcp_server_name_mapping - ) - if not tool_known: - raise ValueError(f"Tool {name} not found") + if resolved_by_server_name_only and not self._server_exposes_tool(mcp_server, name): + raise ValueError(f"Tool {name} not found") return mcp_server @@ -5847,10 +5849,7 @@ class MCPServerManager: if matched is not None: matched_prefix, original_tool_name = matched matched_server: Final = prefix_to_server.get(matched_prefix) - if matched_server is not None and ( - original_tool_name in self.tool_name_to_mcp_server_name_mapping - or tool_name in self.tool_name_to_mcp_server_name_mapping - ): + if matched_server is not None and self._server_exposes_tool(matched_server, original_tool_name): return matched_server return None diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 3392203dbab..20aa19b1d32 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -6406,7 +6406,7 @@ async def test_execute_mcp_tool_rest_server_id_injects_requested_server_credenti with ( patch.dict( mcp_module.global_mcp_server_manager.tool_name_to_mcp_server_name_mapping, - {"echo": collision_server.name}, + {"echo": collision_server.name, "echo_requested-echo": requested_server.name}, ), patch.object( mcp_module.global_mcp_server_manager, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index cd1ef7320dd..043458f599b 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -4833,6 +4833,78 @@ class TestMCPServerManager: with pytest.raises(ValueError, match="Tool missing_tool not found"): manager._resolve_mcp_server_for_tool_call("github", "missing_tool") + @staticmethod + def _manager_with_deepwiki_and_huggingface() -> MCPServerManager: + manager = MCPServerManager() + deepwiki = MCPServer(server_id="deepwiki-id", name="deepwiki", server_name="deepwiki", transport=MCPTransport.http) + huggingface = MCPServer( + server_id="huggingface-id", name="huggingface", server_name="huggingface", transport=MCPTransport.http + ) + manager.registry = {"deepwiki-id": deepwiki, "huggingface-id": huggingface} + manager.tool_name_to_mcp_server_name_mapping.update( + { + "read_wiki_structure": "deepwiki", + "deepwiki-read_wiki_structure": "deepwiki", + "hub_repo_search": "huggingface", + "huggingface-hub_repo_search": "huggingface", + } + ) + return manager + + def test_resolve_mcp_server_for_tool_call_rejects_tool_exposed_only_by_another_server(self): + manager = self._manager_with_deepwiki_and_huggingface() + + with pytest.raises(ValueError, match="Tool read_wiki_structure not found"): + manager._resolve_mcp_server_for_tool_call("huggingface", "read_wiki_structure") + with pytest.raises(ValueError, match="Tool hub_repo_search not found"): + manager._resolve_mcp_server_for_tool_call("deepwiki", "hub_repo_search") + + assert manager._resolve_mcp_server_for_tool_call("deepwiki", "read_wiki_structure") is manager.registry["deepwiki-id"] + assert manager._resolve_mcp_server_for_tool_call("huggingface", "hub_repo_search") is manager.registry["huggingface-id"] + + def test_get_mcp_server_from_tool_name_rejects_other_servers_prefix(self): + manager = self._manager_with_deepwiki_and_huggingface() + + assert manager._get_mcp_server_from_tool_name("huggingface-read_wiki_structure") is None + assert manager._get_mcp_server_from_tool_name("deepwiki-hub_repo_search") is None + assert manager._get_mcp_server_from_tool_name("deepwiki-read_wiki_structure") is manager.registry["deepwiki-id"] + assert manager._get_mcp_server_from_tool_name("huggingface-hub_repo_search") is manager.registry["huggingface-id"] + + def test_resolve_mcp_server_for_tool_call_shared_bare_name_resolves_via_own_prefixed_spelling(self): + manager = MCPServerManager() + zapier = MCPServer(server_id="zapier-id", name="zapier", alias="zapier-alias", transport=MCPTransport.http) + other = MCPServer(server_id="other-id", name="other", server_name="other", transport=MCPTransport.http) + manager.registry = {"zapier-id": zapier, "other-id": other} + manager.tool_name_to_mcp_server_name_mapping.update( + { + "create_zap": "other", + "other-create_zap": "other", + "zapier-alias-create_zap": "zapier-alias", + } + ) + + assert manager._resolve_mcp_server_for_tool_call("zapier", "create_zap") is zapier + assert manager._resolve_mcp_server_for_tool_call("other", "create_zap") is other + + def test_remove_server_drops_only_its_own_tool_mapping_rows(self): + manager = self._manager_with_deepwiki_and_huggingface() + + manager.remove_server( + LiteLLM_MCPServerTable( + server_id="huggingface-id", + alias="huggingface", + url="https://huggingface.co/mcp", + transport=MCPTransport.http, + created_at=datetime.now(), + updated_at=datetime.now(), + ) + ) + + assert manager.tool_name_to_mcp_server_name_mapping == { + "read_wiki_structure": "deepwiki", + "deepwiki-read_wiki_structure": "deepwiki", + } + @pytest.mark.asyncio async def test_resolve_oauth2_headers_skipped_when_not_user_oauth(self): """Returns input headers unchanged when server does not need user OAuth.""" From 75a569b3e2a62dca695a58551bfee23ea9e84713 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 18 Aug 2026 17:04:25 -0700 Subject: [PATCH 27/83] fix(databricks): sync the packaged cost map backup with the new Databricks entries --- ...odel_prices_and_context_window_backup.json | 227 ++++++++++++++++++ 1 file changed, 227 insertions(+) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 409022016b0..e1e89cab7cf 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -14456,6 +14456,26 @@ "supports_tool_choice": true, "supports_output_config": true }, + "databricks/databricks-claude-opus-4-6": { + "input_cost_per_token": 5.00003e-06, + "input_dbu_cost_per_token": 7.1429e-05, + "litellm_provider": "databricks", + "max_input_tokens": 200000, + "max_output_tokens": 64000, + "max_tokens": 64000, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 2.5000010000000002e-05, + "output_dbu_cost_per_token": 0.000357143, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving", + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_reasoning": true, + "supports_minimal_reasoning_effort": true, + "supports_tool_choice": true + }, "databricks/databricks-claude-sonnet-4": { "input_cost_per_token": 2.9999900000000002e-06, "input_dbu_cost_per_token": 4.2857e-05, @@ -14513,6 +14533,25 @@ "supports_reasoning": true, "supports_tool_choice": true }, + "databricks/databricks-claude-sonnet-4-6": { + "input_cost_per_token": 2.9999900000000002e-06, + "input_dbu_cost_per_token": 4.2857e-05, + "litellm_provider": "databricks", + "max_input_tokens": 200000, + "max_output_tokens": 64000, + "max_tokens": 64000, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 1.5000020000000002e-05, + "output_dbu_cost_per_token": 0.000214286, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving", + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_reasoning": true, + "supports_tool_choice": true + }, "databricks/databricks-gemini-2-5-flash": { "input_cost_per_token": 3.0001999999999996e-07, "input_dbu_cost_per_token": 4.285999999999999e-06, @@ -14547,6 +14586,74 @@ "supports_function_calling": true, "supports_tool_choice": true }, + "databricks/databricks-gemini-3-1-flash-lite": { + "input_cost_per_token": 3.1248e-07, + "input_dbu_cost_per_token": 4.464e-06, + "litellm_provider": "databricks", + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_tokens": 65536, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 1.87502e-06, + "output_dbu_cost_per_token": 2.6786e-05, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving", + "supports_function_calling": true, + "supports_tool_choice": true + }, + "databricks/databricks-gemini-3-1-pro": { + "input_cost_per_token": 2.49998e-06, + "input_dbu_cost_per_token": 3.5714e-05, + "litellm_provider": "databricks", + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_tokens": 65536, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 1.5000020000000002e-05, + "output_dbu_cost_per_token": 0.000214286, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving", + "supports_function_calling": true, + "supports_tool_choice": true + }, + "databricks/databricks-gemini-3-flash": { + "input_cost_per_token": 6.2503e-07, + "input_dbu_cost_per_token": 8.929e-06, + "litellm_provider": "databricks", + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_tokens": 65536, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 3.74997e-06, + "output_dbu_cost_per_token": 5.3571e-05, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving", + "supports_function_calling": true, + "supports_tool_choice": true + }, + "databricks/databricks-gemini-3-pro": { + "input_cost_per_token": 2.49998e-06, + "input_dbu_cost_per_token": 3.5714e-05, + "litellm_provider": "databricks", + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_tokens": 65536, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 1.5000020000000002e-05, + "output_dbu_cost_per_token": 0.000214286, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving", + "supports_function_calling": true, + "supports_tool_choice": true + }, "databricks/databricks-gemma-3-12b": { "input_cost_per_token": 1.5000999999999998e-07, "input_dbu_cost_per_token": 2.1429999999999996e-06, @@ -14592,6 +14699,126 @@ "output_dbu_cost_per_token": 0.000142857, "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving" }, + "databricks/databricks-gpt-5-1-codex-max": { + "input_cost_per_token": 1.24999e-06, + "input_dbu_cost_per_token": 1.7857e-05, + "litellm_provider": "databricks", + "max_input_tokens": 272000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 9.999990000000002e-06, + "output_dbu_cost_per_token": 0.000142857, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving" + }, + "databricks/databricks-gpt-5-1-codex-mini": { + "input_cost_per_token": 2.4997e-07, + "input_dbu_cost_per_token": 3.571e-06, + "litellm_provider": "databricks", + "max_input_tokens": 272000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 1.99997e-06, + "output_dbu_cost_per_token": 2.8571e-05, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving" + }, + "databricks/databricks-gpt-5-2": { + "input_cost_per_token": 1.75e-06, + "input_dbu_cost_per_token": 2.5e-05, + "litellm_provider": "databricks", + "max_input_tokens": 272000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 1.4e-05, + "output_dbu_cost_per_token": 0.0002, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving" + }, + "databricks/databricks-gpt-5-2-codex": { + "input_cost_per_token": 1.75e-06, + "input_dbu_cost_per_token": 2.5e-05, + "litellm_provider": "databricks", + "max_input_tokens": 272000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 1.4e-05, + "output_dbu_cost_per_token": 0.0002, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving" + }, + "databricks/databricks-gpt-5-3-codex": { + "input_cost_per_token": 1.75e-06, + "input_dbu_cost_per_token": 2.5e-05, + "litellm_provider": "databricks", + "max_input_tokens": 272000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 1.4e-05, + "output_dbu_cost_per_token": 0.0002, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving" + }, + "databricks/databricks-gpt-5-4": { + "input_cost_per_token": 2.49998e-06, + "input_dbu_cost_per_token": 3.5714e-05, + "litellm_provider": "databricks", + "max_input_tokens": 272000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 1.5000020000000002e-05, + "output_dbu_cost_per_token": 0.000214286, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving" + }, + "databricks/databricks-gpt-5-4-mini": { + "input_cost_per_token": 7.4998e-07, + "input_dbu_cost_per_token": 1.0714e-05, + "litellm_provider": "databricks", + "max_input_tokens": 272000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 4.50002e-06, + "output_dbu_cost_per_token": 6.4286e-05, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving" + }, + "databricks/databricks-gpt-5-4-nano": { + "input_cost_per_token": 1.9999e-07, + "input_dbu_cost_per_token": 2.857e-06, + "litellm_provider": "databricks", + "max_input_tokens": 272000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "metadata": { + "notes": "Input/output cost per token is dbu cost * $0.070. Number provided for reference, '*_dbu_cost_per_token' used in actual calculation." + }, + "mode": "chat", + "output_cost_per_token": 1.24999e-06, + "output_dbu_cost_per_token": 1.7857e-05, + "source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving" + }, "databricks/databricks-gpt-5-mini": { "input_cost_per_token": 2.4997000000000006e-07, "input_dbu_cost_per_token": 3.571e-06, From c33b3a32a60c8f86ed20c1e1940b92e32585de12 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Tue, 18 Aug 2026 17:05:40 -0700 Subject: [PATCH 28/83] feat(bedrock): add a config toggle to disable agent-runtime pass-through (#37386) * feat(bedrock): add a config toggle to disable agent-runtime pass-through The /bedrock pass-through dispatches agents, knowledge bases, flows, rerank, retrieveAndGenerate, generateQuery and optimize-prompt to bedrock-agent-runtime, so an operator who only wants to expose model invoke and converse has no way to narrow that surface Adds general_settings.disable_bedrock_agent_runtime_passthrough. When set, those routes are rejected with a 403 before credentials are fetched or the request is signed. Plain bedrock-runtime model pass-through is unaffected, and the setting defaults to off, so existing deployments behave exactly as before The branch is inverted to an early return for the non-agent-runtime case so the toggle can reject outright instead of falling through to model extraction, which would surface a confusing 400 about an unparseable model * style(bedrock): drop redundant docstrings from the agent-runtime toggle --- .../llm_passthrough_endpoints.py | 23 +++- .../test_llm_pass_through_endpoints.py | 114 ++++++++++++++++++ 2 files changed, 133 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 635767f4db7..c5ab7f1fc63 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -50,7 +50,7 @@ from litellm.proxy.vector_store_endpoints.utils import ( get_litellm_managed_vector_store, is_allowed_to_call_vector_store_endpoint, ) -from litellm.secret_managers.main import get_secret_str +from litellm.secret_managers.main import get_secret_str, str_to_bool from litellm.types.passthrough_endpoints.pass_through_endpoints import ( LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, @@ -1016,15 +1016,21 @@ async def bedrock_proxy_route( raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") aws_region_name: Final = litellm.utils.get_secret(secret_name="AWS_REGION_NAME") - if _is_bedrock_agent_runtime_route(endpoint=endpoint): # handle bedrock agents - base_target_url: Final = f"https://bedrock-agent-runtime.{aws_region_name}.amazonaws.com" - else: + if not _is_bedrock_agent_runtime_route(endpoint=endpoint): return await bedrock_llm_proxy_route( endpoint=endpoint, request=request, fastapi_response=fastapi_response, user_api_key_dict=user_api_key_dict, ) + + if _is_bedrock_agent_runtime_passthrough_disabled(): + raise HTTPException( + status_code=403, + detail="bedrock-agent-runtime pass-through is disabled on this proxy.", + ) + + base_target_url: Final = f"https://bedrock-agent-runtime.{aws_region_name}.amazonaws.com" encoded_endpoint = httpx.URL(endpoint).path # Ensure endpoint starts with '/' for proper URL construction @@ -1292,6 +1298,15 @@ def _is_bedrock_agent_runtime_route(endpoint: str) -> bool: return False +def _is_bedrock_agent_runtime_passthrough_disabled() -> bool: + from litellm.proxy.proxy_server import general_settings + + setting: Final = general_settings.get("disable_bedrock_agent_runtime_passthrough") + if isinstance(setting, str): + return str_to_bool(setting) is True + return setting is True + + @router.api_route( "/assemblyai/{endpoint:path}", methods=["GET", "POST", "PUT", "DELETE", "PATCH"], diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 050070e2fcf..b56a8da7c66 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -1,7 +1,10 @@ +import contextlib import json import os import sys import traceback +from collections.abc import Mapping +from types import MappingProxyType, SimpleNamespace from typing import Final from unittest import mock from unittest.mock import AsyncMock, MagicMock, Mock, patch @@ -22,6 +25,7 @@ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( _join_url_paths, azure_proxy_route, bedrock_llm_proxy_route, + bedrock_proxy_route, create_pass_through_route, cursor_proxy_route, get_azure_ai_search_index_from_endpoint, @@ -1730,6 +1734,116 @@ class TestBedrockLLMProxyRoute: assert "Blocked by guardrail" in str(exc_info.value.detail) +class TestBedrockAgentRuntimePassthroughToggle: + AGENT_RUNTIME_ENDPOINT: Final = "knowledgebases/KB1234567/retrieve" + MODEL_ENDPOINT: Final = "model/us.anthropic.claude-sonnet-4-5-20250929-v1:0/converse" + DISABLED: Final = MappingProxyType({"disable_bedrock_agent_runtime_passthrough": True}) + + @staticmethod + def _mock_request() -> Mock: + request: Final = Mock() + request.method = "POST" + request.state = SimpleNamespace() + request.json = AsyncMock(return_value={"retrievalQuery": {"text": "hi"}}) # mutable-ok: must be json.dumps-able + return request + + @contextlib.contextmanager + def _patched_dispatch(self, general_settings: Mapping[str, object]): + from botocore.credentials import Credentials + + bedrock_llm: Final = Mock() + bedrock_llm.get_credentials = Mock(return_value=Credentials("ak", "sk")) + forwarder: Final = AsyncMock(return_value="forwarded") + + with ( + patch("litellm.proxy.proxy_server.general_settings", general_settings), + patch("litellm.utils.get_secret", return_value="us-east-1"), + patch("litellm.llms.bedrock.chat.BedrockConverseLLM", return_value=bedrock_llm), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_request_copy", + Mock(), + ), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route", + return_value=forwarder, + ) as create_route, + ): + yield create_route, forwarder + + @pytest.mark.asyncio + async def test_agent_runtime_dispatch_allowed_by_default(self): + with self._patched_dispatch(MappingProxyType({})) as (create_route, forwarder): + result: Final = await bedrock_proxy_route( + endpoint=self.AGENT_RUNTIME_ENDPOINT, + request=self._mock_request(), + fastapi_response=Mock(), + user_api_key_dict=UserAPIKeyAuth(), + ) + + assert result == "forwarded" + forwarder.assert_awaited_once() + assert "bedrock-agent-runtime.us-east-1.amazonaws.com" in create_route.call_args.kwargs["target"] + + @pytest.mark.asyncio + @pytest.mark.parametrize("value", (True, "true", "True")) + async def test_agent_runtime_dispatch_rejected_when_disabled(self, value: bool | str): + settings: Final = MappingProxyType({"disable_bedrock_agent_runtime_passthrough": value}) + + with self._patched_dispatch(settings) as (create_route, forwarder): + with pytest.raises(HTTPException) as exc_info: + await bedrock_proxy_route( + endpoint=self.AGENT_RUNTIME_ENDPOINT, + request=self._mock_request(), + fastapi_response=Mock(), + user_api_key_dict=UserAPIKeyAuth(), + ) + + assert exc_info.value.status_code == 403 + assert "bedrock-agent-runtime pass-through is disabled" in str(exc_info.value.detail) + create_route.assert_not_called() + forwarder.assert_not_awaited() + + @pytest.mark.asyncio + async def test_model_invoke_still_routed_when_agent_runtime_disabled(self): + with ( + patch("litellm.proxy.proxy_server.general_settings", self.DISABLED), + patch("litellm.utils.get_secret", return_value="us-east-1"), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_request_copy", + Mock(), + ), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.bedrock_llm_proxy_route", + new=AsyncMock(return_value="llm-route"), + ) as llm_route, + ): + result: Final = await bedrock_proxy_route( + endpoint=self.MODEL_ENDPOINT, + request=self._mock_request(), + fastapi_response=Mock(), + user_api_key_dict=UserAPIKeyAuth(), + ) + + assert result == "llm-route" + llm_route.assert_awaited_once() + + @pytest.mark.asyncio + @pytest.mark.parametrize("value", (False, "false", None, "", "yes")) + async def test_agent_runtime_dispatch_allowed_for_non_true_values(self, value: object): + settings: Final = MappingProxyType({"disable_bedrock_agent_runtime_passthrough": value}) + + with self._patched_dispatch(settings) as (create_route, forwarder): + result: Final = await bedrock_proxy_route( + endpoint=self.AGENT_RUNTIME_ENDPOINT, + request=self._mock_request(), + fastapi_response=Mock(), + user_api_key_dict=UserAPIKeyAuth(), + ) + + assert result == "forwarded" + create_route.assert_called_once() + + class TestLLMPassthroughFactoryProxyRoute: @pytest.mark.asyncio async def test_llm_passthrough_factory_proxy_route_success(self): From 4285ffd82b52f401a0f7104b748ea1ae316b26b5 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 18 Aug 2026 17:09:17 -0700 Subject: [PATCH 29/83] fix(databricks): drop the minimal reasoning effort flag from the claude-opus-4-6 entry --- litellm/model_prices_and_context_window_backup.json | 1 - model_prices_and_context_window.json | 1 - 2 files changed, 2 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index e1e89cab7cf..3ffd65c4d94 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -14473,7 +14473,6 @@ "supports_assistant_prefill": true, "supports_function_calling": true, "supports_reasoning": true, - "supports_minimal_reasoning_effort": true, "supports_tool_choice": true }, "databricks/databricks-claude-sonnet-4": { diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index e1e89cab7cf..3ffd65c4d94 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -14473,7 +14473,6 @@ "supports_assistant_prefill": true, "supports_function_calling": true, "supports_reasoning": true, - "supports_minimal_reasoning_effort": true, "supports_tool_choice": true }, "databricks/databricks-claude-sonnet-4": { From 3f15dc32871c8026295e0fdbac4be7710263861c Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 18 Aug 2026 17:10:00 -0700 Subject: [PATCH 30/83] fix(mcp): attach per-user BYOK credential when listing tools for non-oauth2 auth types (#34787) Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/_experimental/mcp_server/server.py | 3 + .../mcp_server/test_mcp_server.py | 69 +++++++++++++++++++ 2 files changed, 72 insertions(+) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 4184fad009c..d4c5b198fbc 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -2013,6 +2013,9 @@ if MCP_AVAILABLE: prefetched_creds=_prefetched_oauth_creds, ) + if server.is_byok and server.auth_type != MCPAuth.oauth2 and server_auth_header is None: + server_auth_header = await _get_byok_credential(server, user_api_key_auth) + try: tools: Final = await global_mcp_server_manager._get_tools_from_server( server=server, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 3392203dbab..a693559c204 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -8430,3 +8430,72 @@ class TestListFiltersHonorThePrefixBoundary: assert listed == callable_, f"grants={grants!r} listed={listed} callable={callable_}" assert listed is expected, f"grants={grants!r} expected={expected} got={listed}" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "auth_type", + [MCPAuth.authorization, MCPAuth.bearer_token, MCPAuth.api_key, MCPAuth.basic, MCPAuth.token], +) +async def test_list_tools_injects_byok_credential_for_non_oauth2_auth_types(auth_type): + """Regression for BYOK servers on a non-oauth2 auth_type: the stored per-user credential must be + attached when listing tools, otherwise the upstream 401 is absorbed and the server lists nothing.""" + try: + from litellm.proxy._experimental.mcp_server.server import ( + _get_tools_from_mcp_servers, + set_auth_context, + ) + except ImportError: + pytest.skip("MCP server not available") + + user_api_key_auth = UserAPIKeyAuth(api_key="test_key", user_id="byok_user") + set_auth_context(user_api_key_auth) + + server = MagicMock() + server.server_id = "byok_server" + server.name = "byok" + server.alias = "byok" + server.server_name = "byok" + server.auth_type = auth_type + server.is_byok = True + server.allowed_tools = None + server.disallowed_tools = None + server.extra_headers = None + server.tool_name_to_display_name = None + server.tool_name_to_description = None + + seen_auth_headers = [] + + async def mock_get_tools_from_server(server, mcp_auth_header=None, add_prefix=False, **kwargs): + seen_auth_headers.append(mcp_auth_header) + tool = MagicMock() + tool.name = f"{server.alias}-toolA" if add_prefix else "toolA" + tool.description = "desc" + tool.inputSchema = {} + return [tool] + + mock_manager = MagicMock() + mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=[server.server_id]) + mock_manager.get_mcp_server_by_id = MagicMock(return_value=server) + mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0) + mock_manager._get_tools_from_server = mock_get_tools_from_server + + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + mock_manager, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._get_byok_credential", + AsyncMock(return_value="personal-api-key"), + ), + ): + listing = await _get_tools_from_mcp_servers( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=None, + mcp_servers=None, + mcp_server_auth_headers=None, + ) + + assert seen_auth_headers == ["personal-api-key"] + assert [tool.name for tool in listing.tools] == ["byok-toolA"] From 9018a9503724989ae430d8ca5f07a2c26f7c3dd0 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 18 Aug 2026 17:35:40 -0700 Subject: [PATCH 31/83] test(mcp): build fixture mapping state without in-place mutation --- .../mcp_server/test_mcp_server_manager.py | 26 ++++++++----------- 1 file changed, 11 insertions(+), 15 deletions(-) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 043458f599b..14400ef6376 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -4841,14 +4841,12 @@ class TestMCPServerManager: server_id="huggingface-id", name="huggingface", server_name="huggingface", transport=MCPTransport.http ) manager.registry = {"deepwiki-id": deepwiki, "huggingface-id": huggingface} - manager.tool_name_to_mcp_server_name_mapping.update( - { - "read_wiki_structure": "deepwiki", - "deepwiki-read_wiki_structure": "deepwiki", - "hub_repo_search": "huggingface", - "huggingface-hub_repo_search": "huggingface", - } - ) + manager.tool_name_to_mcp_server_name_mapping = { + "read_wiki_structure": "deepwiki", + "deepwiki-read_wiki_structure": "deepwiki", + "hub_repo_search": "huggingface", + "huggingface-hub_repo_search": "huggingface", + } return manager def test_resolve_mcp_server_for_tool_call_rejects_tool_exposed_only_by_another_server(self): @@ -4875,13 +4873,11 @@ class TestMCPServerManager: zapier = MCPServer(server_id="zapier-id", name="zapier", alias="zapier-alias", transport=MCPTransport.http) other = MCPServer(server_id="other-id", name="other", server_name="other", transport=MCPTransport.http) manager.registry = {"zapier-id": zapier, "other-id": other} - manager.tool_name_to_mcp_server_name_mapping.update( - { - "create_zap": "other", - "other-create_zap": "other", - "zapier-alias-create_zap": "zapier-alias", - } - ) + manager.tool_name_to_mcp_server_name_mapping = { + "create_zap": "other", + "other-create_zap": "other", + "zapier-alias-create_zap": "zapier-alias", + } assert manager._resolve_mcp_server_for_tool_call("zapier", "create_zap") is zapier assert manager._resolve_mcp_server_for_tool_call("other", "create_zap") is other From 564ea1cf735b874c36468b73fae3298983290e98 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Tue, 18 Aug 2026 17:36:10 -0700 Subject: [PATCH 32/83] feat(ui): add success, warning and info status tokens (#37393) The dashboard had no shared tokens for non-destructive status colours, so components reached for raw Tailwind shades instead. Add --success, --warning and --info alongside the existing --destructive, in both :root and .dark, and register them in @theme inline so the usual utilities resolve. Light values are picked for legibility as foreground text rather than by copying a fixed shade number. Tailwind's ramps are not perceptually aligned across hues, so amber-600 and green-600 sit at 66.6% and 62.7% lightness and fail WCAG AA on white (3.19:1 and 3.22:1). green-700, amber-700 and blue-600 land at 52.7%, 55.5% and 54.6%, the same band as --destructive at 57.7%, and clear AA. Dark mode uses the -400 shades, matching --destructive. The .dark values are populated even though nothing can apply that class yet. They are the artifact the later theme switch work will turn on. Alert moves its info and warning variants onto the tokens. The tint is /5 rather than /10 because /10 drops both below AA. The error variant keeps its existing shades: it involves no new token, and its current 9.21:1 is better than anything the token form would give it. --- ui/litellm-dashboard/src/app/globals.css | 9 +++++++++ ui/litellm-dashboard/src/components/shared/Alert.tsx | 5 ++--- 2 files changed, 11 insertions(+), 3 deletions(-) diff --git a/ui/litellm-dashboard/src/app/globals.css b/ui/litellm-dashboard/src/app/globals.css index dddeb576720..e3292261176 100644 --- a/ui/litellm-dashboard/src/app/globals.css +++ b/ui/litellm-dashboard/src/app/globals.css @@ -95,6 +95,9 @@ --accent: oklch(0.967 0.003 264.542); --accent-foreground: oklch(0.21 0.034 264.665); --destructive: oklch(0.577 0.245 27.325); + --success: oklch(0.527 0.154 150.069); + --warning: oklch(0.555 0.163 48.998); + --info: oklch(0.546 0.245 262.881); --border: oklch(0.928 0.006 264.531); --input: oklch(0.928 0.006 264.531); --ring: oklch(0.707 0.022 261.325); @@ -130,6 +133,9 @@ --accent: oklch(0.278 0.033 256.848); --accent-foreground: oklch(0.985 0.002 247.839); --destructive: oklch(0.704 0.191 22.216); + --success: oklch(0.792 0.209 151.711); + --warning: oklch(0.828 0.189 84.429); + --info: oklch(0.707 0.165 254.624); --border: oklch(1 0 0 / 10%); --input: oklch(1 0 0 / 15%); --ring: oklch(0.551 0.027 264.364); @@ -168,6 +174,9 @@ --color-accent: var(--accent); --color-accent-foreground: var(--accent-foreground); --color-destructive: var(--destructive); + --color-success: var(--success); + --color-warning: var(--warning); + --color-info: var(--info); --color-border: var(--border); --color-input: var(--input); --color-ring: var(--ring); diff --git a/ui/litellm-dashboard/src/components/shared/Alert.tsx b/ui/litellm-dashboard/src/components/shared/Alert.tsx index 6acf5a987d9..718957b4f08 100644 --- a/ui/litellm-dashboard/src/components/shared/Alert.tsx +++ b/ui/litellm-dashboard/src/components/shared/Alert.tsx @@ -9,9 +9,8 @@ const alertVariants = cva({ variant: { default: "bg-card text-card-foreground", destructive: "bg-card text-destructive *:data-[slot=alert-description]:text-destructive/90 *:[svg]:text-current", - info: "border-blue-200 bg-blue-50 text-blue-900 *:data-[slot=alert-description]:text-blue-800 *:[svg]:text-blue-600", - warning: - "border-amber-200 bg-amber-50 text-amber-900 *:data-[slot=alert-description]:text-amber-800 *:[svg]:text-amber-600", + info: "border-info/20 bg-info/5 text-info *:[svg]:text-current", + warning: "border-warning/20 bg-warning/5 text-warning *:[svg]:text-current", error: "border-red-200 bg-red-50 text-red-900 *:data-[slot=alert-description]:text-red-800 *:[svg]:text-red-600", }, }, From bb8324c1193d8a017501490d9b8ea087af02447b Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Tue, 18 Aug 2026 17:38:12 -0700 Subject: [PATCH 33/83] refactor(ui): drop @tremor/react and the theming scaffolding it needed (#37394) The last tremor component import left the dashboard when the primitive sweep merged, so the package, its v3 compatibility shim, its @theme token block and the palette safelist it needed at runtime all have no consumer. Removing the safelist is what shrinks the shipped stylesheet: tremor built class names at runtime, so Tailwind had to emit every bg/text/border/ring/ stroke/fill utility across 22 palettes and 11 shades in case one was used. Nothing in the app constructs a class name that way any more, so the scanner finds every utility on its own. The date-fns overrides pin also goes. It only existed because tremor and react-day-picker@8 peered on date-fns 3 while Base UI wanted 4, and the lockfile still resolves a single hoisted 4.4.0 without it. --- ui/litellm-dashboard/eslint-suppressions.json | 21 - ui/litellm-dashboard/package-lock.json | 400 ------------ ui/litellm-dashboard/package.json | 2 - .../PriceDataManagementTab.test.tsx | 7 +- .../_components/impact_popover.test.tsx | 10 - .../policies/_components/index.test.tsx | 27 - ui/litellm-dashboard/src/app/globals.css | 63 -- .../src/app/tremor-v3-compat.css | 615 ------------------ .../common_components/chartUtils.test.tsx | 384 ----------- .../common_components/chartUtils.tsx | 104 --- .../KeyInfoView.handleKeyUpdate.test.tsx | 64 -- .../src/components/view_user_spend.tsx | 12 +- ui/litellm-dashboard/tests/setupTests.ts | 39 +- 13 files changed, 7 insertions(+), 1741 deletions(-) delete mode 100644 ui/litellm-dashboard/src/app/tremor-v3-compat.css delete mode 100644 ui/litellm-dashboard/src/components/common_components/chartUtils.test.tsx delete mode 100644 ui/litellm-dashboard/src/components/common_components/chartUtils.tsx diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index ed72e784a35..ca61a0c15af 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -1055,11 +1055,6 @@ "count": 1 } }, - "src/app/(dashboard)/policies/_components/impact_popover.test.tsx": { - "react/display-name": { - "count": 1 - } - }, "src/app/(dashboard)/policies/_components/impact_popover.tsx": { "local/filename-pascal-case": { "count": 1 @@ -1073,11 +1068,6 @@ "count": 1 } }, - "src/app/(dashboard)/policies/_components/index.test.tsx": { - "react/display-name": { - "count": 1 - } - }, "src/app/(dashboard)/policies/_components/index.tsx": { "local/filename-pascal-case": { "count": 1 @@ -2063,14 +2053,6 @@ "count": 1 } }, - "src/components/common_components/chartUtils.tsx": { - "local/filename-pascal-case": { - "count": 1 - }, - "no-nested-ternary": { - "count": 1 - } - }, "src/components/common_components/check_openapi_schema.tsx": { "local/filename-pascal-case": { "count": 1 @@ -3062,9 +3044,6 @@ "tests/setupTests.ts": { "@typescript-eslint/no-this-alias": { "count": 1 - }, - "react/display-name": { - "count": 1 } } } \ No newline at end of file diff --git a/ui/litellm-dashboard/package-lock.json b/ui/litellm-dashboard/package-lock.json index 93ee979f37d..186d38234d5 100644 --- a/ui/litellm-dashboard/package-lock.json +++ b/ui/litellm-dashboard/package-lock.json @@ -18,7 +18,6 @@ "@tanstack/react-pacer": "0.22.1", "@tanstack/react-query": "5.100.7", "@tanstack/react-table": "8.21.3", - "@tremor/react": "3.18.7", "@types/papaparse": "5.5.2", "antd": "5.29.3", "cva": "1.0.0-beta.4", @@ -1460,87 +1459,12 @@ "@floating-ui/utils": "^0.2.11" } }, - "node_modules/@floating-ui/react": { - "version": "0.19.2", - "resolved": "https://registry.npmjs.org/@floating-ui/react/-/react-0.19.2.tgz", - "integrity": "sha512-JyNk4A0Ezirq8FlXECvRtQOX/iBe5Ize0W/pLkrZjfHW9GUV7Xnq6zm6fyZuQzaHHqEnVizmvlA96e1/CkZv+w==", - "license": "MIT", - "dependencies": { - "@floating-ui/react-dom": "^1.3.0", - "aria-hidden": "^1.1.3", - "tabbable": "^6.0.1" - }, - "peerDependencies": { - "react": ">=16.8.0", - "react-dom": ">=16.8.0" - } - }, - "node_modules/@floating-ui/react-dom": { - "version": "1.3.0", - "resolved": "https://registry.npmjs.org/@floating-ui/react-dom/-/react-dom-1.3.0.tgz", - "integrity": "sha512-htwHm67Ji5E/pROEAr7f8IKFShuiCKHwUC/UY4vC3I5jiSvGFAYnSYiZO5MlGmads+QqvUkR9ANHEguGrDv72g==", - "license": "MIT", - "dependencies": { - "@floating-ui/dom": "^1.2.1" - }, - "peerDependencies": { - "react": ">=16.8.0", - "react-dom": ">=16.8.0" - } - }, "node_modules/@floating-ui/utils": { "version": "0.2.11", "resolved": "https://registry.npmjs.org/@floating-ui/utils/-/utils-0.2.11.tgz", "integrity": "sha512-RiB/yIh78pcIxl6lLMG0CgBXAZ2Y0eVHqMPYugu+9U0AeT6YBeiJpf7lbdJNIugFP5SIjwNRgo4DhR1Qxi26Gg==", "license": "MIT" }, - "node_modules/@headlessui/react": { - "version": "2.2.0", - "resolved": "https://registry.npmjs.org/@headlessui/react/-/react-2.2.0.tgz", - "integrity": "sha512-RzCEg+LXsuI7mHiSomsu/gBJSjpupm6A1qIZ5sWjd7JhARNlMiSA4kKfJpCKwU9tE+zMRterhhrP74PvfJrpXQ==", - "license": "MIT", - "dependencies": { - "@floating-ui/react": "^0.26.16", - "@react-aria/focus": "^3.17.1", - "@react-aria/interactions": "^3.21.3", - "@tanstack/react-virtual": "^3.8.1" - }, - "engines": { - "node": ">=10" - }, - "peerDependencies": { - "react": "^18 || ^19 || ^19.0.0-rc", - "react-dom": "^18 || ^19 || ^19.0.0-rc" - } - }, - "node_modules/@headlessui/react/node_modules/@floating-ui/react": { - "version": "0.26.28", - "resolved": "https://registry.npmjs.org/@floating-ui/react/-/react-0.26.28.tgz", - "integrity": "sha512-yORQuuAtVpiRjpMhdc0wJj06b9JFjrYF4qp96j++v2NBpbi6SEGF7donUJ3TMieerQ6qVkAv1tgr7L4r5roTqw==", - "license": "MIT", - "dependencies": { - "@floating-ui/react-dom": "^2.1.2", - "@floating-ui/utils": "^0.2.8", - "tabbable": "^6.0.0" - }, - "peerDependencies": { - "react": ">=16.8.0", - "react-dom": ">=16.8.0" - } - }, - "node_modules/@headlessui/react/node_modules/@floating-ui/react-dom": { - "version": "2.1.8", - "resolved": "https://registry.npmjs.org/@floating-ui/react-dom/-/react-dom-2.1.8.tgz", - "integrity": "sha512-cC52bHwM/n/CxS87FH0yWdngEZrjdtLW/qVruo68qg+prK7ZQ4YGdut2GyDVpoGeAYe/h899rVeOVm6Oi40k2A==", - "license": "MIT", - "dependencies": { - "@floating-ui/dom": "^1.7.6" - }, - "peerDependencies": { - "react": ">=16.8.0", - "react-dom": ">=16.8.0" - } - }, "node_modules/@headlessui/tailwindcss": { "version": "0.2.2", "resolved": "https://registry.npmjs.org/@headlessui/tailwindcss/-/tailwindcss-0.2.2.tgz", @@ -2141,33 +2065,6 @@ "url": "https://opencollective.com/libvips" } }, - "node_modules/@internationalized/date": { - "version": "3.12.1", - "resolved": "https://registry.npmjs.org/@internationalized/date/-/date-3.12.1.tgz", - "integrity": "sha512-6IedsVWXyq4P9Tj+TxuU8WGWM70hYLl12nbYU8jkikVpa6WXapFazPUcHUMDMoWftIDE2ILDkFFte6W2nFCkRQ==", - "license": "Apache-2.0", - "dependencies": { - "@swc/helpers": "^0.5.0" - } - }, - "node_modules/@internationalized/number": { - "version": "3.6.6", - "resolved": "https://registry.npmjs.org/@internationalized/number/-/number-3.6.6.tgz", - "integrity": "sha512-iFgmQaXHE0vytNfpLZWOC2mEJCBRzcUxt53Xf/yCXG93lRvqas237i3r7X4RKMwO3txiyZD4mQjKAByFv6UGSQ==", - "license": "Apache-2.0", - "dependencies": { - "@swc/helpers": "^0.5.0" - } - }, - "node_modules/@internationalized/string": { - "version": "3.2.8", - "resolved": "https://registry.npmjs.org/@internationalized/string/-/string-3.2.8.tgz", - "integrity": "sha512-NdbMQUSfXLYIQol5VyMtinm9pZDciiMfN7RtmSuSB78io1hqwJ0naYfxyW6vgxWBkzWymQa/3uLDlbfmshtCaA==", - "license": "Apache-2.0", - "dependencies": { - "@swc/helpers": "^0.5.0" - } - }, "node_modules/@istanbuljs/schema": { "version": "0.1.6", "resolved": "https://registry.npmjs.org/@istanbuljs/schema/-/schema-0.1.6.tgz", @@ -2876,44 +2773,6 @@ "react-dom": ">=16.9.0" } }, - "node_modules/@react-aria/focus": { - "version": "3.22.0", - "resolved": "https://registry.npmjs.org/@react-aria/focus/-/focus-3.22.0.tgz", - "integrity": "sha512-ZfDOVuVhqDsM9mkNji3QUZ/d40JhlVgXrDkrfXylM1035QCrcTHN7m2DpbE95sU2A8EQb4wikvt5jM6K/73BPg==", - "license": "Apache-2.0", - "dependencies": { - "@swc/helpers": "^0.5.0", - "react-aria": "3.48.0" - }, - "peerDependencies": { - "react": "^16.8.0 || ^17.0.0-rc.1 || ^18.0.0 || ^19.0.0-rc.1", - "react-dom": "^16.8.0 || ^17.0.0-rc.1 || ^18.0.0 || ^19.0.0-rc.1" - } - }, - "node_modules/@react-aria/interactions": { - "version": "3.28.0", - "resolved": "https://registry.npmjs.org/@react-aria/interactions/-/interactions-3.28.0.tgz", - "integrity": "sha512-OXwdU1EWFdMxmr/K1CXNGJzmNlCClByb+PuCaqUyzBymHPCGVhawirLIon/CrIN5psh3AiWpHSh4H0WeJdVpng==", - "license": "Apache-2.0", - "dependencies": { - "@react-types/shared": "^3.34.0", - "@swc/helpers": "^0.5.0", - "react-aria": "3.48.0" - }, - "peerDependencies": { - "react": "^16.8.0 || ^17.0.0-rc.1 || ^18.0.0 || ^19.0.0-rc.1", - "react-dom": "^16.8.0 || ^17.0.0-rc.1 || ^18.0.0 || ^19.0.0-rc.1" - } - }, - "node_modules/@react-types/shared": { - "version": "3.34.0", - "resolved": "https://registry.npmjs.org/@react-types/shared/-/shared-3.34.0.tgz", - "integrity": "sha512-gp6xo/s2lX54AlTjOiqwDnxA7UW79BNvI9dB9pr3LZTzRKCd1ZA+ZbgKw/ReIiWuvvVw/8QFJpnqeeFyLocMcQ==", - "license": "Apache-2.0", - "peerDependencies": { - "react": "^16.8.0 || ^17.0.0-rc.1 || ^18.0.0 || ^19.0.0-rc.1" - } - }, "node_modules/@redocly/ajv": { "version": "8.11.2", "resolved": "https://registry.npmjs.org/@redocly/ajv/-/ajv-8.11.2.tgz", @@ -3839,23 +3698,6 @@ "react-dom": ">=16.8" } }, - "node_modules/@tanstack/react-virtual": { - "version": "3.13.24", - "resolved": "https://registry.npmjs.org/@tanstack/react-virtual/-/react-virtual-3.13.24.tgz", - "integrity": "sha512-aIJvz5OSkhNIhZIpYivrxrPTKYsjW9Uzy+sP/mx0S3sev2HyvPb7xmjbYvokzEpfgYHy/HjzJ2zFAETuUfgCpg==", - "license": "MIT", - "dependencies": { - "@tanstack/virtual-core": "3.14.0" - }, - "funding": { - "type": "github", - "url": "https://github.com/sponsors/tannerlinsley" - }, - "peerDependencies": { - "react": "^16.8.0 || ^17.0.0 || ^18.0.0 || ^19.0.0", - "react-dom": "^16.8.0 || ^17.0.0 || ^18.0.0 || ^19.0.0" - } - }, "node_modules/@tanstack/store": { "version": "0.11.0", "resolved": "https://registry.npmjs.org/@tanstack/store/-/store-0.11.0.tgz", @@ -3879,16 +3721,6 @@ "url": "https://github.com/sponsors/tannerlinsley" } }, - "node_modules/@tanstack/virtual-core": { - "version": "3.14.0", - "resolved": "https://registry.npmjs.org/@tanstack/virtual-core/-/virtual-core-3.14.0.tgz", - "integrity": "sha512-JLANqGy/D6k4Ujmh8Tr25lGimuOXNiaVyXaCAZS0W+1390sADdGnyUdSWNIfd49gebtIxGMij4IktRVzrdr12Q==", - "license": "MIT", - "funding": { - "type": "github", - "url": "https://github.com/sponsors/tannerlinsley" - } - }, "node_modules/@testing-library/dom": { "version": "10.4.1", "resolved": "https://registry.npmjs.org/@testing-library/dom/-/dom-10.4.1.tgz", @@ -3978,93 +3810,6 @@ "@testing-library/dom": ">=7.21.4" } }, - "node_modules/@tremor/react": { - "version": "3.18.7", - "resolved": "https://registry.npmjs.org/@tremor/react/-/react-3.18.7.tgz", - "integrity": "sha512-nmqvf/1m0GB4LXc7v2ftdfSLoZhy5WLrhV6HNf0SOriE6/l8WkYeWuhQq8QsBjRi94mUIKLJ/VC3/Y/pj6VubQ==", - "license": "Apache 2.0", - "dependencies": { - "@floating-ui/react": "^0.19.2", - "@headlessui/react": "2.2.0", - "date-fns": "^3.6.0", - "react-day-picker": "^8.10.1", - "react-transition-state": "^2.1.2", - "recharts": "^2.13.3", - "tailwind-merge": "^2.5.2" - }, - "peerDependencies": { - "react": "^18.0.0", - "react-dom": ">=16.6.0" - } - }, - "node_modules/@tremor/react/node_modules/eventemitter3": { - "version": "4.0.7", - "resolved": "https://registry.npmjs.org/eventemitter3/-/eventemitter3-4.0.7.tgz", - "integrity": "sha512-8guHBZCwKnFhYdHr2ysuRWErTwhoN2X8XELRlrRwpmfeY2jjuUN4taQMsULKUVo1K4DvZl+0pgfyoysHxvmvEw==", - "license": "MIT" - }, - "node_modules/@tremor/react/node_modules/react-is": { - "version": "18.3.1", - "resolved": "https://registry.npmjs.org/react-is/-/react-is-18.3.1.tgz", - "integrity": "sha512-/LLMVyas0ljjAtoYiPqYiL8VWXzUUdThrmU5+n20DZv+a+ClRoevUzw5JxU+Ieh5/c87ytoTBV9G1FiKfNJdmg==", - "license": "MIT" - }, - "node_modules/@tremor/react/node_modules/recharts": { - "version": "2.15.4", - "resolved": "https://registry.npmjs.org/recharts/-/recharts-2.15.4.tgz", - "integrity": "sha512-UT/q6fwS3c1dHbXv2uFgYJ9BMFHu3fwnd7AYZaEQhXuYQ4hgsxLvsUXzGdKeZrW5xopzDCvuA2N41WJ88I7zIw==", - "deprecated": "1.x and 2.x branches are no longer active. Bump to Recharts v3 to receive latest features and bugfixes. See https://github.com/recharts/recharts/wiki/3.0-migration-guide", - "license": "MIT", - "dependencies": { - "clsx": "^2.0.0", - "eventemitter3": "^4.0.1", - "lodash": "^4.17.21", - "react-is": "^18.3.1", - "react-smooth": "^4.0.4", - "recharts-scale": "^0.4.4", - "tiny-invariant": "^1.3.1", - "victory-vendor": "^36.6.8" - }, - "engines": { - "node": ">=14" - }, - "peerDependencies": { - "react": "^16.0.0 || ^17.0.0 || ^18.0.0 || ^19.0.0", - "react-dom": "^16.0.0 || ^17.0.0 || ^18.0.0 || ^19.0.0" - } - }, - "node_modules/@tremor/react/node_modules/tailwind-merge": { - "version": "2.6.1", - "resolved": "https://registry.npmjs.org/tailwind-merge/-/tailwind-merge-2.6.1.tgz", - "integrity": "sha512-Oo6tHdpZsGpkKG88HJ8RR1rg/RdnEkQEfMoEk2x1XRI3F1AxeU+ijRXpiVUF4UbLfcxxRGw6TbUINKYdWVsQTQ==", - "license": "MIT", - "funding": { - "type": "github", - "url": "https://github.com/sponsors/dcastil" - } - }, - "node_modules/@tremor/react/node_modules/victory-vendor": { - "version": "36.9.2", - "resolved": "https://registry.npmjs.org/victory-vendor/-/victory-vendor-36.9.2.tgz", - "integrity": "sha512-PnpQQMuxlwYdocC8fIJqVXvkeViHYzotI+NJrCuav0ZYFoq912ZHBk3mCeuj+5/VpodOjPe1z0Fk2ihgzlXqjQ==", - "license": "MIT AND ISC", - "dependencies": { - "@types/d3-array": "^3.0.3", - "@types/d3-ease": "^3.0.0", - "@types/d3-interpolate": "^3.0.1", - "@types/d3-scale": "^4.0.2", - "@types/d3-shape": "^3.1.0", - "@types/d3-time": "^3.0.0", - "@types/d3-timer": "^3.0.0", - "d3-array": "^3.1.6", - "d3-ease": "^3.0.1", - "d3-interpolate": "^3.0.1", - "d3-scale": "^4.0.2", - "d3-shape": "^3.1.0", - "d3-time": "^3.0.0", - "d3-timer": "^3.0.1" - } - }, "node_modules/@tybys/wasm-util": { "version": "0.10.3", "resolved": "https://registry.npmjs.org/@tybys/wasm-util/-/wasm-util-0.10.3.tgz", @@ -5203,18 +4948,6 @@ "dev": true, "license": "Python-2.0" }, - "node_modules/aria-hidden": { - "version": "1.2.6", - "resolved": "https://registry.npmjs.org/aria-hidden/-/aria-hidden-1.2.6.tgz", - "integrity": "sha512-ik3ZgC9dY/lYVVM++OISsaYDeg1tb0VtP5uL3ouh1koGOaUMDPpbFIei4JkFimWUFPn90sbMNMXQAIVOlnYKJA==", - "license": "MIT", - "dependencies": { - "tslib": "^2.0.0" - }, - "engines": { - "node": ">=10" - } - }, "node_modules/aria-query": { "version": "5.3.0", "resolved": "https://registry.npmjs.org/aria-query/-/aria-query-5.3.0.tgz", @@ -6314,16 +6047,6 @@ "dev": true, "license": "MIT" }, - "node_modules/dom-helpers": { - "version": "5.2.1", - "resolved": "https://registry.npmjs.org/dom-helpers/-/dom-helpers-5.2.1.tgz", - "integrity": "sha512-nRCa7CK3VTrM2NmGkIy4cbK7IZlgBE/PYMn55rrXefr5xXDP0LdtfPnblFDoVdcAfslJ7or6iqAUnx0CCGIWQA==", - "license": "MIT", - "dependencies": { - "@babel/runtime": "^7.8.7", - "csstype": "^3.0.2" - } - }, "node_modules/dunder-proto": { "version": "1.0.1", "resolved": "https://registry.npmjs.org/dunder-proto/-/dunder-proto-1.0.1.tgz", @@ -7201,15 +6924,6 @@ "dev": true, "license": "MIT" }, - "node_modules/fast-equals": { - "version": "5.4.1", - "resolved": "https://registry.npmjs.org/fast-equals/-/fast-equals-5.4.1.tgz", - "integrity": "sha512-DjlFSM5Pk9cGcL0q5QXl66eGzx0N6szNgaswwc5ZphlBohjTVJSnGgI+rJVOgOi65qUoQnDZN4nDqi33udtydQ==", - "license": "MIT", - "engines": { - "node": ">=6.0.0" - } - }, "node_modules/fast-glob": { "version": "3.3.1", "resolved": "https://registry.npmjs.org/fast-glob/-/fast-glob-3.3.1.tgz", @@ -9229,12 +8943,6 @@ "url": "https://github.com/sponsors/sindresorhus" } }, - "node_modules/lodash": { - "version": "4.18.1", - "resolved": "https://registry.npmjs.org/lodash/-/lodash-4.18.1.tgz", - "integrity": "sha512-dMInicTPVE8d1e5otfwmmjlxkZoUpiVLwyeTdUsi/Caj/gfzzblBcCE5sRHV/AsjuCmxWrte2TNGSYuCeCq+0Q==", - "license": "MIT" - }, "node_modules/lodash.merge": { "version": "4.6.2", "resolved": "https://registry.npmjs.org/lodash.merge/-/lodash.merge-4.6.2.tgz", @@ -11868,27 +11576,6 @@ "node": ">=0.10.0" } }, - "node_modules/react-aria": { - "version": "3.48.0", - "resolved": "https://registry.npmjs.org/react-aria/-/react-aria-3.48.0.tgz", - "integrity": "sha512-jQjd4rBEIMqecBaAKYJbVGK6EqIHLa5znVQ7jwFyK5vCyljoj6KhgtiahmcIPsG5vG5vEDLw+ba+bEWn6A2P4w==", - "license": "Apache-2.0", - "dependencies": { - "@internationalized/date": "^3.12.1", - "@internationalized/number": "^3.6.6", - "@internationalized/string": "^3.2.8", - "@react-types/shared": "^3.34.0", - "@swc/helpers": "^0.5.0", - "aria-hidden": "^1.2.3", - "clsx": "^2.0.0", - "react-stately": "3.46.0", - "use-sync-external-store": "^1.6.0" - }, - "peerDependencies": { - "react": "^16.8.0 || ^17.0.0-rc.1 || ^18.0.0 || ^19.0.0-rc.1", - "react-dom": "^16.8.0 || ^17.0.0-rc.1 || ^18.0.0 || ^19.0.0-rc.1" - } - }, "node_modules/react-copy-to-clipboard": { "version": "5.1.1", "resolved": "https://registry.npmjs.org/react-copy-to-clipboard/-/react-copy-to-clipboard-5.1.1.tgz", @@ -11902,20 +11589,6 @@ "react": ">=15.3.0" } }, - "node_modules/react-day-picker": { - "version": "8.10.2", - "resolved": "https://registry.npmjs.org/react-day-picker/-/react-day-picker-8.10.2.tgz", - "integrity": "sha512-LK68OTbHB3oJNhl9cA0qVizzp3o26w61YSjAFkYi67N86iro32wx86kSNeFU/hq+gI8m1yzWhnomMLfZ041RzQ==", - "license": "MIT", - "funding": { - "type": "individual", - "url": "https://github.com/sponsors/gpbl" - }, - "peerDependencies": { - "date-fns": "^2.28.0 || ^3.0.0", - "react": "^16.8.0 || ^17.0.0 || ^18.0.0 || ^19.0.0" - } - }, "node_modules/react-dom": { "version": "18.3.1", "resolved": "https://registry.npmjs.org/react-dom/-/react-dom-18.3.1.tgz", @@ -12013,38 +11686,6 @@ } } }, - "node_modules/react-smooth": { - "version": "4.0.4", - "resolved": "https://registry.npmjs.org/react-smooth/-/react-smooth-4.0.4.tgz", - "integrity": "sha512-gnGKTpYwqL0Iii09gHobNolvX4Kiq4PKx6eWBCYYix+8cdw+cGo3do906l1NBPKkSWx1DghC1dlWG9L2uGd61Q==", - "license": "MIT", - "dependencies": { - "fast-equals": "^5.0.1", - "prop-types": "^15.8.1", - "react-transition-group": "^4.4.5" - }, - "peerDependencies": { - "react": "^16.8.0 || ^17.0.0 || ^18.0.0 || ^19.0.0", - "react-dom": "^16.8.0 || ^17.0.0 || ^18.0.0 || ^19.0.0" - } - }, - "node_modules/react-stately": { - "version": "3.46.0", - "resolved": "https://registry.npmjs.org/react-stately/-/react-stately-3.46.0.tgz", - "integrity": "sha512-OdxhWvHgs2L4OJGIs7hnuTr5WjjMM6enhNEAMRqiekhF8+ITvA2LRwNftOZwcogaoCslGYq5S2VQTQwnm0GbCA==", - "license": "Apache-2.0", - "dependencies": { - "@internationalized/date": "^3.12.1", - "@internationalized/number": "^3.6.6", - "@internationalized/string": "^3.2.8", - "@react-types/shared": "^3.34.0", - "@swc/helpers": "^0.5.0", - "use-sync-external-store": "^1.6.0" - }, - "peerDependencies": { - "react": "^16.8.0 || ^17.0.0-rc.1 || ^18.0.0 || ^19.0.0-rc.1" - } - }, "node_modules/react-syntax-highlighter": { "version": "15.6.6", "resolved": "https://registry.npmjs.org/react-syntax-highlighter/-/react-syntax-highlighter-15.6.6.tgz", @@ -12062,32 +11703,6 @@ "react": ">= 0.14.0" } }, - "node_modules/react-transition-group": { - "version": "4.4.5", - "resolved": "https://registry.npmjs.org/react-transition-group/-/react-transition-group-4.4.5.tgz", - "integrity": "sha512-pZcd1MCJoiKiBR2NRxeCRg13uCXbydPnmB4EOeRrY7480qNWO8IIgQG6zlDkm6uRMsURXPuKq0GWtiM59a5Q6g==", - "license": "BSD-3-Clause", - "dependencies": { - "@babel/runtime": "^7.5.5", - "dom-helpers": "^5.0.1", - "loose-envify": "^1.4.0", - "prop-types": "^15.6.2" - }, - "peerDependencies": { - "react": ">=16.6.0", - "react-dom": ">=16.6.0" - } - }, - "node_modules/react-transition-state": { - "version": "2.3.3", - "resolved": "https://registry.npmjs.org/react-transition-state/-/react-transition-state-2.3.3.tgz", - "integrity": "sha512-wsIyg07ohlWEAYDZHvuXh/DY7mxlcLb0iqVv2aMXJ0gwgPVKNWKhOyNyzuJy/tt/6urSq0WT6BBZ/tdpybaAsQ==", - "license": "MIT", - "peerDependencies": { - "react": ">=16.8.0", - "react-dom": ">=16.8.0" - } - }, "node_modules/recharts": { "version": "3.9.2", "resolved": "https://registry.npmjs.org/recharts/-/recharts-3.9.2.tgz", @@ -12118,15 +11733,6 @@ "react-is": "^16.8.0 || ^17.0.0 || ^18.0.0 || ^19.0.0" } }, - "node_modules/recharts-scale": { - "version": "0.4.5", - "resolved": "https://registry.npmjs.org/recharts-scale/-/recharts-scale-0.4.5.tgz", - "integrity": "sha512-kivNFO+0OcUNu7jQquLXAxz1FIwZj8nrj+YkOKc5694NbjCvcT6aSZiIzNzd2Kul4o4rTto8QVR9lMNtxD4G1w==", - "license": "MIT", - "dependencies": { - "decimal.js-light": "^2.4.1" - } - }, "node_modules/redent": { "version": "3.0.0", "resolved": "https://registry.npmjs.org/redent/-/redent-3.0.0.tgz", @@ -13200,12 +12806,6 @@ "dev": true, "license": "MIT" }, - "node_modules/tabbable": { - "version": "6.4.0", - "resolved": "https://registry.npmjs.org/tabbable/-/tabbable-6.4.0.tgz", - "integrity": "sha512-05PUHKSNE8ou2dwIxTngl4EzcnsCDZGJ/iCLtDflR/SHB/ny14rXc+qU5P4mG9JkusiV7EivzY9Mhm55AzAvCg==", - "license": "MIT" - }, "node_modules/tailwind-merge": { "version": "3.4.0", "resolved": "https://registry.npmjs.org/tailwind-merge/-/tailwind-merge-3.4.0.tgz", diff --git a/ui/litellm-dashboard/package.json b/ui/litellm-dashboard/package.json index 483ab39c336..9407ff25a1f 100644 --- a/ui/litellm-dashboard/package.json +++ b/ui/litellm-dashboard/package.json @@ -31,7 +31,6 @@ "@tanstack/react-pacer": "0.22.1", "@tanstack/react-query": "5.100.7", "@tanstack/react-table": "8.21.3", - "@tremor/react": "3.18.7", "@types/papaparse": "5.5.2", "antd": "5.29.3", "cva": "1.0.0-beta.4", @@ -103,7 +102,6 @@ "axios": "1.13.6", "postcss": "8.5.23", "esbuild": "0.28.1", - "date-fns": "^4.4.0", "sharp": "^0.35.0" }, "engines": { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/PriceDataManagementTab.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/PriceDataManagementTab.test.tsx index 282cc2722db..8b34d61ebad 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/PriceDataManagementTab.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/PriceDataManagementTab.test.tsx @@ -3,11 +3,6 @@ import { render } from "@testing-library/react"; import { describe, expect, it, vi } from "vitest"; import PriceDataManagementTab from "./PriceDataManagementTab"; -// Deliberately do NOT mock @tremor/react. These tab components render standalone -// (inside antd Tabs / directly as a route page), no longer inside a Tremor -// . A Tremor root renders nothing without that context, so -// this asserts the component's content is visible on its own — reverting the root -// back to makes the title disappear and fails this test. vi.mock("@/components/price_data_reload", () => ({ default: () =>
reload
})); vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ default: () => ({ accessToken: "sk-test" }) })); vi.mock("@/app/(dashboard)/hooks/models/useModelCostMap", () => ({ @@ -15,7 +10,7 @@ vi.mock("@/app/(dashboard)/hooks/models/useModelCostMap", () => ({ })); describe("PriceDataManagementTab", () => { - it("renders its content standalone, without a Tremor TabGroup ancestor", () => { + it("renders its content standalone, without a tab-panel ancestor", () => { const { getByText } = render(); expect(getByText("Price Data Management")).toBeInTheDocument(); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/impact_popover.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/impact_popover.test.tsx index 66f03577e20..69c471a7c3d 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/impact_popover.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/impact_popover.test.tsx @@ -30,16 +30,6 @@ vi.mock("@heroicons/react/outline", () => ({ }, })); -vi.mock("@tremor/react", async (importOriginal) => { - const actual = await importOriginal(); - return { - ...actual, - Icon: React.forwardRef(({ icon: _icon, ...props }, ref) => ( -