From b49661064f90f75645ca7cf5c3a0fddbc5369daa Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 01:44:13 +0000 Subject: [PATCH 01/41] feat(cost-map): add bedrock_mantle rows for claude opus 5.5 and sonnet 5.5 (#43647) * feat(cost-map): add bedrock_mantle rows for claude opus 5.5 and sonnet 5.5 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * revert(cost-map): keep bedrock_mantle claude 5.5 change to cost map rows only Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: kerry Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- ...odel_prices_and_context_window_backup.json | 152 ++++++++++++++++++ model_prices_and_context_window.json | 152 ++++++++++++++++++ tests/unit/test_cost_calculator.py | 61 +++++-- 3 files changed, 354 insertions(+), 11 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 140e0cf0071..58b4e9fb65a 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -58680,6 +58680,87 @@ "input_cost_per_token_batches": 5e-07, "output_cost_per_token_batches": 2.5e-06 }, + "bedrock_mantle/anthropic.claude-opus-5-5": { + "bedrock_converse_supports_strict_tools": false, + "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, + "cache_creation_input_token_cost": 5.5e-06, + "cache_creation_input_token_cost_above_1hr": 8.8e-06, + "cache_read_input_token_cost": 2.2e-07, + "input_cost_per_token": 4.4e-06, + "litellm_provider": "bedrock_mantle", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 2.2e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": false, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "supports_parallel_tool_use_config": true, + "prompt_cache_min_tokens": 512, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-anthropic-claude-opus-5-5.html", + "thinking_always_on": true, + "supports_forced_tool_use": false + }, + "bedrock_mantle/anthropic.claude-sonnet-5-5": { + "bedrock_converse_supports_strict_tools": false, + "cache_creation_input_token_cost": 2.75e-06, + "cache_creation_input_token_cost_above_1hr": 4.4e-06, + "cache_read_input_token_cost": 2.2e-07, + "input_cost_per_token": 2.2e-06, + "litellm_provider": "bedrock_mantle", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1.1e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": false, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh", + "supports_parallel_tool_use_config": true, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-anthropic-claude-sonnet-5-5.html" + }, "us.xai.grok-4.6": { "input_cost_per_token": 2.2e-06, "output_cost_per_token": 6.6e-06, @@ -65298,6 +65379,77 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "bedrock_mantle/us-gov-west-1/anthropic.claude-opus-5-5": { + "bedrock_converse_supports_strict_tools": false, + "cache_creation_input_token_cost": 6e-06, + "cache_creation_input_token_cost_above_1hr": 9.6e-06, + "cache_read_input_token_cost": 2.4e-07, + "input_cost_per_token": 4.8e-06, + "litellm_provider": "bedrock_mantle", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 2.4e-05, + "prompt_cache_min_tokens": 512, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_mid_conversation_system": true, + "supports_native_structured_output": false, + "supports_output_config": true, + "supports_parallel_tool_use_config": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "thinking_always_on": true, + "supports_forced_tool_use": false, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-anthropic-claude-opus-5-5.html" + }, + "bedrock_mantle/us-gov-west-1/anthropic.claude-sonnet-5-5": { + "bedrock_converse_supports_strict_tools": false, + "bedrock_output_config_effort_ceiling": "xhigh", + "cache_creation_input_token_cost": 3e-06, + "cache_creation_input_token_cost_above_1hr": 4.8e-06, + "cache_read_input_token_cost": 2.4e-07, + "input_cost_per_token": 2.4e-06, + "litellm_provider": "bedrock_mantle", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1.2e-05, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_mid_conversation_system": true, + "supports_native_structured_output": false, + "supports_output_config": true, + "supports_parallel_tool_use_config": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-anthropic-claude-sonnet-5-5.html" + }, "bedrock_mantle/us-gov-east-1/openai.gpt-5.4": { "litellm_provider": "bedrock_mantle", "max_input_tokens": 1050000, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 140e0cf0071..58b4e9fb65a 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -58680,6 +58680,87 @@ "input_cost_per_token_batches": 5e-07, "output_cost_per_token_batches": 2.5e-06 }, + "bedrock_mantle/anthropic.claude-opus-5-5": { + "bedrock_converse_supports_strict_tools": false, + "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, + "cache_creation_input_token_cost": 5.5e-06, + "cache_creation_input_token_cost_above_1hr": 8.8e-06, + "cache_read_input_token_cost": 2.2e-07, + "input_cost_per_token": 4.4e-06, + "litellm_provider": "bedrock_mantle", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 2.2e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": false, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "supports_parallel_tool_use_config": true, + "prompt_cache_min_tokens": 512, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-anthropic-claude-opus-5-5.html", + "thinking_always_on": true, + "supports_forced_tool_use": false + }, + "bedrock_mantle/anthropic.claude-sonnet-5-5": { + "bedrock_converse_supports_strict_tools": false, + "cache_creation_input_token_cost": 2.75e-06, + "cache_creation_input_token_cost_above_1hr": 4.4e-06, + "cache_read_input_token_cost": 2.2e-07, + "input_cost_per_token": 2.2e-06, + "litellm_provider": "bedrock_mantle", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1.1e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_mid_conversation_system": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "supports_native_structured_output": false, + "supports_max_reasoning_effort": true, + "supports_output_config": true, + "bedrock_output_config_effort_ceiling": "xhigh", + "supports_parallel_tool_use_config": true, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-anthropic-claude-sonnet-5-5.html" + }, "us.xai.grok-4.6": { "input_cost_per_token": 2.2e-06, "output_cost_per_token": 6.6e-06, @@ -65298,6 +65379,77 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "bedrock_mantle/us-gov-west-1/anthropic.claude-opus-5-5": { + "bedrock_converse_supports_strict_tools": false, + "cache_creation_input_token_cost": 6e-06, + "cache_creation_input_token_cost_above_1hr": 9.6e-06, + "cache_read_input_token_cost": 2.4e-07, + "input_cost_per_token": 4.8e-06, + "litellm_provider": "bedrock_mantle", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 2.4e-05, + "prompt_cache_min_tokens": 512, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_mid_conversation_system": true, + "supports_native_structured_output": false, + "supports_output_config": true, + "supports_parallel_tool_use_config": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "thinking_always_on": true, + "supports_forced_tool_use": false, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-anthropic-claude-opus-5-5.html" + }, + "bedrock_mantle/us-gov-west-1/anthropic.claude-sonnet-5-5": { + "bedrock_converse_supports_strict_tools": false, + "bedrock_output_config_effort_ceiling": "xhigh", + "cache_creation_input_token_cost": 3e-06, + "cache_creation_input_token_cost_above_1hr": 4.8e-06, + "cache_read_input_token_cost": 2.4e-07, + "input_cost_per_token": 2.4e-06, + "litellm_provider": "bedrock_mantle", + "supports_tool_search": true, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1.2e-05, + "supports_forced_tool_use": false, + "thinking_always_on": true, + "prompt_cache_min_tokens": 512, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_mid_conversation_system": true, + "supports_native_structured_output": false, + "supports_output_config": true, + "supports_parallel_tool_use_config": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_sampling_params": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-anthropic-claude-sonnet-5-5.html" + }, "bedrock_mantle/us-gov-east-1/openai.gpt-5.4": { "litellm_provider": "bedrock_mantle", "max_input_tokens": 1050000, diff --git a/tests/unit/test_cost_calculator.py b/tests/unit/test_cost_calculator.py index 62ef9f11c2e..adfedf61d45 100644 --- a/tests/unit/test_cost_calculator.py +++ b/tests/unit/test_cost_calculator.py @@ -3682,26 +3682,46 @@ def test_completion_cost_mantle_native_messages_prices_claude_from_the_bedrock_r ) == pytest.approx(expected) -def test_completion_cost_mantle_native_messages_prices_haiku_from_the_mantle_row(_local_model_cost_map): - """Mantle serves Anthropic's un-versioned haiku id, which has no bare Bedrock row (Bedrock's carries - the -20251001-v1:0 suffix), and Claude Code sends every small-fast-model call to it. Both the plain - and the region-prefixed deployment names must price from bedrock_mantle/anthropic.claude-haiku-4-5 - instead of billing $0.""" +@pytest.mark.parametrize( + "response_model,mantle_row,deployment_models", + [ + ( + "claude-haiku-4-5", + "bedrock_mantle/anthropic.claude-haiku-4-5", + ( + "bedrock_mantle/anthropic.claude-haiku-4-5", + "bedrock_mantle/us-east-2/anthropic.claude-haiku-4-5", + ), + ), + ( + "claude-opus-5-5", + "bedrock_mantle/anthropic.claude-opus-5-5", + ("bedrock_mantle/anthropic.claude-opus-5-5",), + ), + ( + "claude-sonnet-5-5", + "bedrock_mantle/anthropic.claude-sonnet-5-5", + ("bedrock_mantle/anthropic.claude-sonnet-5-5",), + ), + ], +) +def test_completion_cost_mantle_native_messages_prices_unversioned_claude_from_the_mantle_row( + _local_model_cost_map, response_model, mantle_row, deployment_models +): + """Mantle serves Anthropic's un-versioned Claude ids; the plain and region-prefixed deployment + names must price from the model's own bedrock_mantle/ row instead of billing $0.""" response = litellm.ModelResponse( id="msg_x", choices=[{"index": 0, "message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}], - model="claude-haiku-4-5", + model=response_model, usage={"prompt_tokens": 100, "completion_tokens": 10, "total_tokens": 110}, ) - row = litellm.model_cost["bedrock_mantle/anthropic.claude-haiku-4-5"] + row = litellm.model_cost[mantle_row] expected = 100 * row["input_cost_per_token"] + 10 * row["output_cost_per_token"] assert expected > 0 - for model in ( - "bedrock_mantle/anthropic.claude-haiku-4-5", - "bedrock_mantle/us-east-2/anthropic.claude-haiku-4-5", - ): + for model in deployment_models: assert litellm.completion_cost( completion_response=response, model=model, @@ -3709,6 +3729,25 @@ def test_completion_cost_mantle_native_messages_prices_haiku_from_the_mantle_row ) == pytest.approx(expected), model +@pytest.mark.parametrize("model", ["anthropic.claude-opus-5-5", "anthropic.claude-sonnet-5-5"]) +def test_cost_per_token_gov_region_prices_mantle_claude_on_the_gov_row(_local_model_cost_map, model): + """A bedrock_mantle/ deployment in us-gov-west-1 must price from the + bedrock_mantle/us-gov-west-1/ row.""" + + prompt_cost, completion_cost = litellm.cost_per_token( + model=f"bedrock_mantle/{model}", + prompt_tokens=38, + completion_tokens=20, + custom_llm_provider="bedrock_mantle", + region_name="us-gov-west-1", + ) + gov = litellm.model_cost[f"bedrock_mantle/us-gov-west-1/{model}"] + + assert prompt_cost + completion_cost == pytest.approx( + 38 * gov["input_cost_per_token"] + 20 * gov["output_cost_per_token"] + ) + + def test_completion_cost_legacy_mantle_route_prices_after_router_registration(local_model_cost_map): """The proxy registers every deployment under its provider-prefixed key at boot. A bedrock/mantle/ deployment must resolve to the bare Bedrock row there, otherwise the boot From 3e21e5e348e652a9b8a5c5c8b144c229c0bd9018 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 01:48:22 +0000 Subject: [PATCH 02/41] chore(codeowners): drop UI, migration, and CODEOWNERS self owners (#43653) * chore(codeowners): drop UI and migration code owners Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore(codeowners): drop CODEOWNERS self-owner line Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yuneng Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .github/CODEOWNERS | 8 -------- 1 file changed, 8 deletions(-) diff --git a/.github/CODEOWNERS b/.github/CODEOWNERS index 70a50d7f06e..582cf0f5217 100644 --- a/.github/CODEOWNERS +++ b/.github/CODEOWNERS @@ -1,10 +1,2 @@ -/ui/ @yuneng-berri @ryan-crabbe-berri -/litellm/proxy/_experimental/out/ @yuneng-berri @ryan-crabbe-berri -/ui/Dockerfile -/ui/nginx.conf -/ui/litellm-dashboard/src/lib/http/schema.d.ts -/ui/litellm-dashboard/tsconfig.tsbuildinfo /model_prices_and_context_window.json @mateo-berri @ryan-crabbe-berri @kerry-berri /litellm/model_prices_and_context_window_backup.json @mateo-berri @ryan-crabbe-berri @kerry-berri -/litellm-proxy-extras/litellm_proxy_extras/migrations/ @yuneng-berri @ryan-crabbe-berri -/.github/CODEOWNERS @yuneng-berri From 39d14bd8557737dbdc963748f46a909c005a0c24 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 28 Sep 2026 19:08:16 -0700 Subject: [PATCH 03/41] feat(fireworks_ai): route and list the auto, auto-instant and firerouter routers (#43641) * feat(fireworks_ai): route and list the auto, auto-instant and firerouter routers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(fireworks_ai): drive the router request test through an httpx MockTransport Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(fireworks_ai): let custom firerouter/ IDs inherit the firerouter row's capabilities Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(fireworks_ai): integration coverage for router short names forwarding tool_choice and reasoning_effort Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(fireworks_ai): assert tool definitions reach the router upstream Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: kerry Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../llms/fireworks_ai/chat/transformation.py | 9 +++ litellm/llms/fireworks_ai/common_utils.py | 3 +- ...odel_prices_and_context_window_backup.json | 27 +++++++ model_prices_and_context_window.json | 27 +++++++ .../test_fireworks_ai_router_slug_wire.py | 74 +++++++++++++++++ .../test_fireworks_ai_chat_transformation.py | 81 +++++++++++++++++++ .../test_fireworks_ai_common_utils.py | 7 ++ 7 files changed, 227 insertions(+), 1 deletion(-) diff --git a/litellm/llms/fireworks_ai/chat/transformation.py b/litellm/llms/fireworks_ai/chat/transformation.py index 196022b3558..15049e71bbc 100644 --- a/litellm/llms/fireworks_ai/chat/transformation.py +++ b/litellm/llms/fireworks_ai/chat/transformation.py @@ -40,6 +40,7 @@ from ...openai.chat.gpt_transformation import ( OpenAIGPTConfig, ) from ..common_utils import ( + FIREROUTER, FireworksAIException, FireworksAIMixin, resolve_fireworks_resource_name, @@ -574,12 +575,20 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig): short_name = short_name.removeprefix("accounts/fireworks/models/") return short_name + @staticmethod + def _firerouter_family_cost_keys(model: str) -> tuple[str, ...]: + firerouter_resource: Final = f"accounts/fireworks/routers/{FIREROUTER}" + if not resolve_fireworks_resource_name(model).startswith(f"{firerouter_resource}/"): + return () + return (f"fireworks_ai/{firerouter_resource}",) + def _get_model_cost_capability_exact(self, model: str, capability: str) -> bool | None: short_name: Final = self._short_model_name(model) candidate_keys: Final = ( model, f"fireworks_ai/{short_name}", f"fireworks_ai/accounts/fireworks/models/{short_name}", + *self._firerouter_family_cost_keys(model), ) for candidate_key in candidate_keys: model_info = litellm.model_cost.get(candidate_key) diff --git a/litellm/llms/fireworks_ai/common_utils.py b/litellm/llms/fireworks_ai/common_utils.py index ae52b89aa58..8352690235d 100644 --- a/litellm/llms/fireworks_ai/common_utils.py +++ b/litellm/llms/fireworks_ai/common_utils.py @@ -60,6 +60,7 @@ def resolve_fireworks_api_key(api_key: str | None) -> str | None: AZURE_FOUNDRY_FIREWORKS_MODEL_ID_PREFIX: Final = "FW-" FIREROUTER: Final = "firerouter" +ROUTER_SHORT_NAMES: Final = frozenset({FIREROUTER, "auto", "auto-instant"}) def resolve_fireworks_resource_name(model: str) -> str: @@ -68,7 +69,7 @@ def resolve_fireworks_resource_name(model: str) -> str: return stripped if stripped.startswith(("routers/", "models/")): return f"accounts/fireworks/{stripped}" - if stripped.endswith("-fast") or stripped == FIREROUTER or stripped.startswith(f"{FIREROUTER}/"): + if stripped.endswith("-fast") or stripped in ROUTER_SHORT_NAMES or stripped.startswith(f"{FIREROUTER}/"): return f"accounts/fireworks/routers/{stripped}" return f"accounts/fireworks/models/{stripped}" diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 58b4e9fb65a..7bb83a95bd4 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -64178,6 +64178,33 @@ "supports_tool_choice": true, "supports_vision": false }, + "fireworks_ai/accounts/fireworks/routers/auto": { + "litellm_provider": "fireworks_ai", + "mode": "chat", + "source": "https://docs.fireworks.ai/nexus/firerouter#example-router-ids", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "fireworks_ai/accounts/fireworks/routers/auto-instant": { + "litellm_provider": "fireworks_ai", + "mode": "chat", + "source": "https://docs.fireworks.ai/nexus/firerouter#example-router-ids", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "fireworks_ai/accounts/fireworks/routers/firerouter": { + "litellm_provider": "fireworks_ai", + "mode": "chat", + "source": "https://docs.fireworks.ai/nexus/firerouter#example-router-ids", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, "fireworks_ai/glm-5p3-fast": { "cache_read_input_token_cost": 3.9e-07, "input_cost_per_token": 2.1e-06, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 58b4e9fb65a..7bb83a95bd4 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -64178,6 +64178,33 @@ "supports_tool_choice": true, "supports_vision": false }, + "fireworks_ai/accounts/fireworks/routers/auto": { + "litellm_provider": "fireworks_ai", + "mode": "chat", + "source": "https://docs.fireworks.ai/nexus/firerouter#example-router-ids", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "fireworks_ai/accounts/fireworks/routers/auto-instant": { + "litellm_provider": "fireworks_ai", + "mode": "chat", + "source": "https://docs.fireworks.ai/nexus/firerouter#example-router-ids", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "fireworks_ai/accounts/fireworks/routers/firerouter": { + "litellm_provider": "fireworks_ai", + "mode": "chat", + "source": "https://docs.fireworks.ai/nexus/firerouter#example-router-ids", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, "fireworks_ai/glm-5p3-fast": { "cache_read_input_token_cost": 3.9e-07, "input_cost_per_token": 2.1e-06, diff --git a/tests/integration/providers/test_fireworks_ai_router_slug_wire.py b/tests/integration/providers/test_fireworks_ai_router_slug_wire.py index 55c83945ec7..f6262cc21f1 100644 --- a/tests/integration/providers/test_fireworks_ai_router_slug_wire.py +++ b/tests/integration/providers/test_fireworks_ai_router_slug_wire.py @@ -44,6 +44,35 @@ def _catalog_cost(model: str, field: str) -> float: _ROUTED_MODEL: Final = _pick_routed_model() +_FIREWORKS_MODEL_PREFIX: Final = "fireworks_ai/accounts/fireworks/models/" + + +def _pick_open_model_key() -> str: + catalog: Final = _COST_MAP.validate_json(_COST_MAP_PATH.read_bytes()) + return next( + key + for key, entry in catalog.items() + if key.startswith(_FIREWORKS_MODEL_PREFIX) + and _positive_rate(entry, "input_cost_per_token") + and _positive_rate(entry, "output_cost_per_token") + ) + + +_SERVED_OPEN_MODEL_KEY: Final = _pick_open_model_key() +_ROUTERS_ACCEPTING_TOOL_CHOICE_AND_REASONING: Final = ( + "auto", + "auto-instant", + "firerouter", + "firerouter/opus", + "firerouter/auto", +) +_WEATHER_TOOL: Final = { + "type": "function", + "function": { + "name": "get_weather", + "parameters": {"type": "object", "properties": {"city": {"type": "string"}}, "required": ["city"]}, + }, +} def _approx(value: float) -> object: @@ -197,3 +226,48 @@ def test_fireworks_firerouter_claude_leg_is_charged_at_the_routed_models_own_rat spend: Final = rows[0]["spend"] assert isinstance(spend, (int, float, str)) assert float(spend) == _approx(expected_cost) + + +@pytest.mark.parametrize("router", _ROUTERS_ACCEPTING_TOOL_CHOICE_AND_REASONING) +def test_fireworks_router_forwards_tool_choice_and_reasoning_and_bills_the_served_open_model( + gateway: Gateway, router: str +) -> None: + identity: Final = f"fw-{router.replace('/', '-')}-{uuid.uuid4().hex}" + served_resource: Final = _SERVED_OPEN_MODEL_KEY.removeprefix("fireworks_ai/") + + def respond(request: Request) -> Reply: + body: Final = _provider_body(request, "/chat/completions") + assert body["model"] == f"accounts/fireworks/routers/{router}", body + assert body["tools"] == [_WEATHER_TOOL], body + assert body["tool_choice"] == "any", body + assert body["reasoning_effort"] == "low", body + return Reply(body=_chat_completion(identity, served_resource, 23, 41)) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"fireworks_ai/{router}", api_base=wire.url, api_key=_API_KEY) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": _PROMPT}], + "tools": [_WEATHER_TOOL], + "tool_choice": "required", + "reasoning_effort": "low", + }, + ) + assert response.status_code == 200, response.text + expected_cost: Final = 23 * _catalog_cost(_SERVED_OPEN_MODEL_KEY, "input_cost_per_token") + 41 * _catalog_cost( + _SERVED_OPEN_MODEL_KEY, "output_cost_per_token" + ) + assert expected_cost > 0 + assert float(response.headers["x-litellm-response-cost"]) == _approx(expected_cost) + assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/chat/completions")] + rows: Final = eventually( + lambda: read_rows('SELECT spend FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (identity,)), + lambda values: len(values) == 1, + seconds=70, + ) + spend: Final = rows[0]["spend"] + assert isinstance(spend, (int, float, str)) + assert float(spend) == _approx(expected_cost) diff --git a/tests/unit/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py b/tests/unit/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py index 3f740edf834..2298bed2b76 100644 --- a/tests/unit/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py +++ b/tests/unit/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py @@ -1,10 +1,13 @@ import json +from typing import Final from unittest.mock import MagicMock, patch +import httpx import pytest import litellm from litellm.constants import SESSION_ID_GENERATED_METADATA_KEY +from litellm.llms.custom_httpx.http_handler import HTTPHandler from litellm.llms.fireworks_ai.chat.transformation import FireworksAIConfig from litellm.llms.fireworks_ai.common_utils import get_fireworks_session_id from litellm.types.utils import ( @@ -1781,8 +1784,86 @@ def test_streaming_preserves_selected_model_for_private_accounting(): [ ("deepseek-r1", "fireworks_ai/accounts/fireworks/models/deepseek-r1"), ("glm-5p3-fast", "fireworks_ai/accounts/fireworks/routers/glm-5p3-fast"), + ("auto", "fireworks_ai/accounts/fireworks/routers/auto"), ("accounts/fireworks/models/deepseek-r1", "fireworks_ai/accounts/fireworks/models/deepseek-r1"), ], ) def test_get_model_cost_key_resolves_short_names_to_long_keys(model: str, expected: str) -> None: assert FireworksAIConfig().get_model_cost_key(model) == expected + + +_LISTED_ROUTERS = ("auto", "auto-instant", "firerouter") + + +@pytest.mark.parametrize("router", _LISTED_ROUTERS) +def test_listed_router_short_name_resolves_to_its_catalog_row_and_accepts_tool_choice_and_reasoning( + router: str, +) -> None: + info = litellm.get_model_info(model=f"fireworks_ai/{router}") + params = FireworksAIConfig().get_supported_openai_params(router) + + assert info["key"] == f"fireworks_ai/accounts/fireworks/routers/{router}" + assert {"tools", "tool_choice", "reasoning_effort"} <= set(params), params + + +@pytest.mark.parametrize( + "router", + [ + "firerouter/opus", + "firerouter/auto", + "firerouter/auto-instant", + "firerouter/kimi-k3/glm-5p3", + "fireworks_ai/firerouter/opus", + "accounts/fireworks/routers/firerouter/opus", + ], +) +def test_custom_firerouter_id_accepts_the_same_tool_choice_and_reasoning_params_as_firerouter(router: str) -> None: + params: Final = FireworksAIConfig().get_supported_openai_params(router) + + assert {"tools", "tool_choice", "reasoning_effort"} <= set(params), params + + +@pytest.mark.parametrize("model", ["firerouter-v2", "models/firerouter-opus", "routers/firerouter-opus"]) +def test_names_that_only_start_with_firerouter_do_not_inherit_the_firerouter_row(model: str) -> None: + params: Final = FireworksAIConfig().get_supported_openai_params(model) + + assert "tool_choice" not in params, params + + +class _RecordingChatHandler: + def __init__(self, reply: dict[str, object]) -> None: + self.reply: Final = reply + self.request_body: dict[str, object] | None = None + + def __call__(self, request: httpx.Request) -> httpx.Response: + self.request_body = json.loads(request.content) + return httpx.Response(200, json=self.reply, request=request) + + +@pytest.mark.parametrize("router", _LISTED_ROUTERS) +def test_listed_router_request_is_sent_to_the_router_resource_and_billed_at_the_served_models_rate(router: str) -> None: + served_model: Final = "glm-5p3-flash" + handler: Final = _RecordingChatHandler( + { + "id": f"chat-{router}", + "object": "chat.completion", + "created": 1, + "model": served_model, + "choices": [{"index": 0, "message": {"role": "assistant", "content": "pong"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 23, "completion_tokens": 41, "total_tokens": 64}, + } + ) + + response: Final = litellm.completion( + model=f"fireworks_ai/{router}", + messages=[{"role": "user", "content": "ping"}], + api_key="fw-test-key", + client=HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(handler))), + ) + + served_info: Final = litellm.model_cost[f"fireworks_ai/{served_model}"] + expected_cost: Final = 23 * served_info["input_cost_per_token"] + 41 * served_info["output_cost_per_token"] + assert handler.request_body is not None + assert handler.request_body["model"] == f"accounts/fireworks/routers/{router}" + assert expected_cost > 0 + assert response._hidden_params["response_cost"] == pytest.approx(expected_cost) diff --git a/tests/unit/llms/fireworks_ai/test_fireworks_ai_common_utils.py b/tests/unit/llms/fireworks_ai/test_fireworks_ai_common_utils.py index e505f2ae8a6..7ebbd0c6a8a 100644 --- a/tests/unit/llms/fireworks_ai/test_fireworks_ai_common_utils.py +++ b/tests/unit/llms/fireworks_ai/test_fireworks_ai_common_utils.py @@ -18,6 +18,13 @@ from litellm.llms.fireworks_ai.common_utils import resolve_fireworks_resource_na ("fireworks_ai/firerouter", "accounts/fireworks/routers/firerouter"), ("firerouter/kimi-k3/deepseek-v4", "accounts/fireworks/routers/firerouter/kimi-k3/deepseek-v4"), ("firerouter-v2", "accounts/fireworks/models/firerouter-v2"), + ("auto", "accounts/fireworks/routers/auto"), + ("fireworks_ai/auto", "accounts/fireworks/routers/auto"), + ("auto-instant", "accounts/fireworks/routers/auto-instant"), + ("fireworks_ai/auto-instant", "accounts/fireworks/routers/auto-instant"), + ("firerouter/auto", "accounts/fireworks/routers/firerouter/auto"), + ("autoglm-9b", "accounts/fireworks/models/autoglm-9b"), + ("auto-v2", "accounts/fireworks/models/auto-v2"), ( "accounts/fireworks/routers/glm-latest", "accounts/fireworks/routers/glm-latest", From 118ce3cc916d78490ff9ae9721fc87ef78d05e84 Mon Sep 17 00:00:00 2001 From: ishaan-berri <155045088+ishaan-berri@users.noreply.github.com> Date: Mon, 28 Sep 2026 19:40:04 -0700 Subject: [PATCH 04/41] feat: add model leaderboard page (#43649) * feat(proxy): add model leaderboard analytics * feat: add model insights task and range constants * feat: record task type from task tags in model usage rollup * feat: serve 365 days of model insights by UTC date * test: cover task tag resolution in model usage rollup * test: update model insights range limit test to 365 days * chore: regenerate dashboard api types for model insights * feat: add model insights aggregation helpers * test: cover model insights aggregation helpers * feat: redesign model leaderboard with stacked bars, treemap and ranking * test: update model leaderboard view test * feat: mark model leaderboard as beta in sidebar * chore: sync schema.prisma copies from root * fix: only treat task: prefixed tags as model insight tasks * feat: add metric type for model insights ranking * fix: rank model insights by selected metric and scope detail queries to ranked deployments * test: plain tags are not model insight tasks * test: cover metric ranking, deployment scoping and rollup round trip * fix: build model insights weeks and halves from the requested date range * test: cover empty weeks and range-based change comparison * fix: refetch by metric, show load errors and ignore stale responses * test: cover metric refetch and error state * feat: define model insight tasks in a JSON file * feat: return task labels and categories from model insights * feat: load model insight tasks from JSON * refactor: validate rollup task tags against the JSON task list * feat: serve the task list with model insights * refactor: drop hardcoded task list from constants * build: ship model insight tasks JSON in the wheel * test: cover model insight task JSON * refactor: take task labels and categories from the API * test: pass task info to task tile builder * refactor: color treemap by API-provided category * test: include tasks in model leaderboard fixture * fix: make daily model usage migration idempotent * feat: bound the model insights task query size * fix: compute task breakdown independent of the chart metric * test: task breakdown is stable across chart metrics * chore: regenerate lazy openapi snapshot for model insights * chore: regenerate dashboard api types for model insights * fix: keep previous ranking dimmed while a new metric loads * test: cover stale metric state in model leaderboard * refactor: drop task row cap constant * fix: return the full task breakdown instead of a truncated one * test: task query is not truncated * feat: add task summary types for model insights * feat: summarise tasks server-side on a separate model insights endpoint * test: cover the model insights tasks endpoint * chore: regenerate lazy openapi snapshot for model insights tasks * chore: regenerate dashboard api types for model insights tasks * refactor: drop client-side task aggregation * test: remove client-side task aggregation tests * feat: load task breakdown separately from the chart metric * test: task breakdown is not refetched on chart metric change --------- Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> --- .../migration.sql | 19 + .../litellm_proxy_extras/schema.prisma | 20 + litellm/constants.py | 4 + litellm/proxy/_lazy_features.py | 5 + litellm/proxy/_lazy_openapi_snapshot.json | 457 ++++++++++++++++++ litellm/proxy/db/db_spend_update_writer.py | 10 + litellm/proxy/db/model_insights_tasks.py | 14 + litellm/proxy/db/model_usage_rollup.py | 86 ++++ .../model_insights_endpoints.py | 220 +++++++++ litellm/proxy/model_insights_tasks.json | 22 + litellm/proxy/schema.prisma | 20 + litellm/repositories/__init__.py | 2 + litellm/repositories/table_repositories.py | 4 + litellm/types/model_insights.py | 47 ++ pyproject.toml | 1 + schema.prisma | 20 + .../proxy/db/test_model_insights_tasks.py | 19 + .../proxy/db/test_model_usage_rollup.py | 89 ++++ .../test_model_insights_endpoints.py | 233 +++++++++ .../src/app/(dashboard)/legacyPageRoutes.ts | 1 + .../_components/ModelInsightsView.test.tsx | 148 ++++++ .../_components/ModelInsightsView.tsx | 352 ++++++++++++++ .../_components/modelInsightsData.test.ts | 92 ++++ .../_components/modelInsightsData.ts | 132 +++++ .../app/(dashboard)/model-insights/page.tsx | 9 + .../src/components/leftnav.tsx | 11 + .../src/components/page_metadata.ts | 1 + ui/litellm-dashboard/src/lib/http/schema.d.ts | 187 +++++++ 28 files changed, 2225 insertions(+) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260928000000_add_daily_model_usage/migration.sql create mode 100644 litellm/proxy/db/model_insights_tasks.py create mode 100644 litellm/proxy/db/model_usage_rollup.py create mode 100644 litellm/proxy/management_endpoints/model_insights_endpoints.py create mode 100644 litellm/proxy/model_insights_tasks.json create mode 100644 litellm/types/model_insights.py create mode 100644 tests/test_litellm/proxy/db/test_model_insights_tasks.py create mode 100644 tests/test_litellm/proxy/db/test_model_usage_rollup.py create mode 100644 tests/test_litellm/proxy/management_endpoints/test_model_insights_endpoints.py create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.test.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/modelInsightsData.test.ts create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/modelInsightsData.ts create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/model-insights/page.tsx diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260928000000_add_daily_model_usage/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260928000000_add_daily_model_usage/migration.sql new file mode 100644 index 00000000000..1395296ea61 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260928000000_add_daily_model_usage/migration.sql @@ -0,0 +1,19 @@ +CREATE TABLE IF NOT EXISTS "LiteLLM_DailyModelUsage" ( + "date" TEXT NOT NULL, + "model_group" TEXT NOT NULL, + "model" TEXT NOT NULL, + "custom_llm_provider" TEXT NOT NULL, + "task_type" TEXT NOT NULL, + "spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0, + "prompt_tokens" BIGINT NOT NULL DEFAULT 0, + "completion_tokens" BIGINT NOT NULL DEFAULT 0, + "request_count" BIGINT NOT NULL DEFAULT 0, + "successful_requests" BIGINT NOT NULL DEFAULT 0, + "failed_requests" BIGINT NOT NULL DEFAULT 0, + "created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "updated_at" TIMESTAMP(3) NOT NULL, + CONSTRAINT "LiteLLM_DailyModelUsage_pkey" PRIMARY KEY ("date", "model_group", "model", "custom_llm_provider", "task_type") +); + +CREATE INDEX IF NOT EXISTS "LiteLLM_DailyModelUsage_date_idx" ON "LiteLLM_DailyModelUsage"("date"); +CREATE INDEX IF NOT EXISTS "LiteLLM_DailyModelUsage_model_group_idx" ON "LiteLLM_DailyModelUsage"("model_group"); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 7edc565879a..03e59257f76 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -1260,6 +1260,26 @@ model LiteLLM_DailyToolSpend { @@id([date, tool_name]) } +model LiteLLM_DailyModelUsage { + date String + model_group String + model String + custom_llm_provider String + task_type String + spend Float @default(0.0) + prompt_tokens BigInt @default(0) + completion_tokens BigInt @default(0) + request_count BigInt @default(0) + successful_requests BigInt @default(0) + failed_requests BigInt @default(0) + created_at DateTime @default(now()) + updated_at DateTime @updatedAt + + @@id([date, model_group, model, custom_llm_provider, task_type]) + @@index([date]) + @@index([model_group]) +} + // Gateway request counts recorded at the ASGI edge by // BillableRequestMetricsMiddleware. This is the source of truth for SGR // (successful gateway requests): it counts what the proxy actually answered, diff --git a/litellm/constants.py b/litellm/constants.py index 10c943656f7..39c10d71709 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1777,6 +1777,10 @@ SCHEDULED_JOB_SHUTDOWN_CANCEL_TIMEOUT_SECONDS: Final = float( os.getenv("SCHEDULED_JOB_SHUTDOWN_CANCEL_TIMEOUT_SECONDS", "5") ) TOOL_SPEND_TOP_TOOLS: Final = 100 +MODEL_INSIGHTS_TOP_MODELS: Final = 10 +MODEL_INSIGHTS_MAX_RANGE_DAYS: Final = 365 +MODEL_INSIGHTS_DEFAULT_TASK: Final = "uncategorized" +MODEL_INSIGHTS_TASK_TAG_PREFIX: Final = "task:" SPEND_LOG_PARTITION_INTERVAL: Final = os.getenv("SPEND_LOG_PARTITION_INTERVAL", "day") SPEND_LOG_PARTITION_PRECREATE_AHEAD: Final = int(os.getenv("SPEND_LOG_PARTITION_PRECREATE_AHEAD", 7)) SPEND_LOG_WRITE_BATCH_MAX_BYTES: Final = max(1, int(os.getenv("SPEND_LOG_WRITE_BATCH_MAX_BYTES", 2_000_000))) diff --git a/litellm/proxy/_lazy_features.py b/litellm/proxy/_lazy_features.py index 5be87a8bf4d..98cf3a4ba23 100644 --- a/litellm/proxy/_lazy_features.py +++ b/litellm/proxy/_lazy_features.py @@ -128,6 +128,11 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = ( module_path="litellm.proxy.management_endpoints.tool_management_endpoints", path_prefixes=("/v1/tool", "/tool"), ), + LazyFeature( + name="model_insights", + module_path="litellm.proxy.management_endpoints.model_insights_endpoints", + path_prefixes=("/model-insights",), + ), LazyFeature( name="search_tools", module_path="litellm.proxy.search_endpoints.search_tool_management", diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 3e0d623375c..75dce43c84a 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -40700,6 +40700,463 @@ } } }, + "model_insights": { + "components": { + "schemas": { + "HTTPValidationError": { + "properties": { + "detail": { + "items": { + "$ref": "#/components/schemas/ValidationError" + }, + "title": "Detail", + "type": "array" + } + }, + "title": "HTTPValidationError", + "type": "object" + }, + "ModelInsightDailyMetric": { + "properties": { + "completion_tokens": { + "title": "Completion Tokens", + "type": "integer" + }, + "date": { + "title": "Date", + "type": "string" + }, + "failed_requests": { + "title": "Failed Requests", + "type": "integer" + }, + "model": { + "title": "Model", + "type": "string" + }, + "model_group": { + "title": "Model Group", + "type": "string" + }, + "prompt_tokens": { + "title": "Prompt Tokens", + "type": "integer" + }, + "provider": { + "title": "Provider", + "type": "string" + }, + "requests": { + "title": "Requests", + "type": "integer" + }, + "spend": { + "title": "Spend", + "type": "number" + }, + "successful_requests": { + "title": "Successful Requests", + "type": "integer" + } + }, + "required": [ + "model_group", + "model", + "provider", + "spend", + "prompt_tokens", + "completion_tokens", + "requests", + "successful_requests", + "failed_requests", + "date" + ], + "title": "ModelInsightDailyMetric", + "type": "object" + }, + "ModelInsightMetric": { + "properties": { + "completion_tokens": { + "title": "Completion Tokens", + "type": "integer" + }, + "failed_requests": { + "title": "Failed Requests", + "type": "integer" + }, + "model": { + "title": "Model", + "type": "string" + }, + "model_group": { + "title": "Model Group", + "type": "string" + }, + "prompt_tokens": { + "title": "Prompt Tokens", + "type": "integer" + }, + "provider": { + "title": "Provider", + "type": "string" + }, + "requests": { + "title": "Requests", + "type": "integer" + }, + "spend": { + "title": "Spend", + "type": "number" + }, + "successful_requests": { + "title": "Successful Requests", + "type": "integer" + } + }, + "required": [ + "model_group", + "model", + "provider", + "spend", + "prompt_tokens", + "completion_tokens", + "requests", + "successful_requests", + "failed_requests" + ], + "title": "ModelInsightMetric", + "type": "object" + }, + "ModelInsightTaskSummary": { + "properties": { + "category": { + "title": "Category", + "type": "string" + }, + "label": { + "title": "Label", + "type": "string" + }, + "leader": { + "title": "Leader", + "type": "string" + }, + "provider": { + "title": "Provider", + "type": "string" + }, + "share": { + "title": "Share", + "type": "number" + }, + "task_type": { + "title": "Task Type", + "type": "string" + }, + "value": { + "title": "Value", + "type": "number" + } + }, + "required": [ + "task_type", + "label", + "category", + "value", + "share", + "leader", + "provider" + ], + "title": "ModelInsightTaskSummary", + "type": "object" + }, + "ModelInsightTasksResponse": { + "properties": { + "end_date": { + "title": "End Date", + "type": "string" + }, + "start_date": { + "title": "Start Date", + "type": "string" + }, + "tasks": { + "items": { + "$ref": "#/components/schemas/ModelInsightTaskSummary" + }, + "title": "Tasks", + "type": "array" + } + }, + "required": [ + "start_date", + "end_date", + "tasks" + ], + "title": "ModelInsightTasksResponse", + "type": "object" + }, + "ModelInsightsResponse": { + "properties": { + "daily": { + "items": { + "$ref": "#/components/schemas/ModelInsightDailyMetric" + }, + "title": "Daily", + "type": "array" + }, + "end_date": { + "title": "End Date", + "type": "string" + }, + "start_date": { + "title": "Start Date", + "type": "string" + }, + "top_models": { + "items": { + "$ref": "#/components/schemas/ModelInsightMetric" + }, + "title": "Top Models", + "type": "array" + } + }, + "required": [ + "start_date", + "end_date", + "daily", + "top_models" + ], + "title": "ModelInsightsResponse", + "type": "object" + }, + "ValidationError": { + "properties": { + "ctx": { + "title": "Context", + "type": "object" + }, + "input": { + "title": "Input" + }, + "loc": { + "items": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "integer" + } + ] + }, + "title": "Location", + "type": "array" + }, + "msg": { + "title": "Message", + "type": "string" + }, + "type": { + "title": "Error Type", + "type": "string" + } + }, + "required": [ + "loc", + "msg", + "type" + ], + "title": "ValidationError", + "type": "object" + } + } + }, + "paths": { + "/model-insights": { + "get": { + "operationId": "get_model_insights_model_insights_get", + "parameters": [ + { + "description": "YYYY-MM-DD, defaults to 365 days ago", + "in": "query", + "name": "start_date", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "description": "YYYY-MM-DD, defaults to 365 days ago", + "title": "Start Date" + } + }, + { + "description": "YYYY-MM-DD, defaults to today", + "in": "query", + "name": "end_date", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "description": "YYYY-MM-DD, defaults to today", + "title": "End Date" + } + }, + { + "description": "Metric the top models are ranked by", + "in": "query", + "name": "metric", + "required": false, + "schema": { + "default": "tokens", + "description": "Metric the top models are ranked by", + "enum": [ + "requests", + "spend", + "tokens" + ], + "title": "Metric", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ModelInsightsResponse" + } + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Get Model Insights", + "tags": [ + "model_insights" + ] + } + }, + "/model-insights/tasks": { + "get": { + "operationId": "get_model_insight_tasks_model_insights_tasks_get", + "parameters": [ + { + "description": "YYYY-MM-DD, defaults to 365 days ago", + "in": "query", + "name": "start_date", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "description": "YYYY-MM-DD, defaults to 365 days ago", + "title": "Start Date" + } + }, + { + "description": "YYYY-MM-DD, defaults to today", + "in": "query", + "name": "end_date", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "description": "YYYY-MM-DD, defaults to today", + "title": "End Date" + } + }, + { + "description": "Metric task shares are computed from", + "in": "query", + "name": "metric", + "required": false, + "schema": { + "default": "spend", + "description": "Metric task shares are computed from", + "enum": [ + "requests", + "spend", + "tokens" + ], + "title": "Metric", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ModelInsightTasksResponse" + } + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Get Model Insight Tasks", + "tags": [ + "model_insights" + ] + } + } + } + }, "policies": { "components": { "schemas": { diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 17e6152bef6..72553e82283 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -1148,6 +1148,16 @@ class DBSpendUpdateWriter: traceback.format_exc(), ) + try: + from litellm.proxy.db.model_usage_rollup import increment_daily_model_usage + + await increment_daily_model_usage(prisma_client=prisma_client, payload=payload_copy) + except Exception: + verbose_proxy_logger.debug( + "_batch_database_updates: increment_daily_model_usage failed: %s", + traceback.format_exc(), + ) + async def _update_key_db( self, response_cost: float | None, diff --git a/litellm/proxy/db/model_insights_tasks.py b/litellm/proxy/db/model_insights_tasks.py new file mode 100644 index 00000000000..865965dcf75 --- /dev/null +++ b/litellm/proxy/db/model_insights_tasks.py @@ -0,0 +1,14 @@ +import json +from functools import lru_cache +from pathlib import Path +from typing import Final + +from litellm.types.model_insights import ModelInsightTask + +_TASKS_FILE: Final = Path(__file__).resolve().parent.parent / "model_insights_tasks.json" + + +@lru_cache(maxsize=1) +def load_model_insight_tasks() -> dict[str, ModelInsightTask]: + raw: Final = json.loads(_TASKS_FILE.read_text()) + return {name: ModelInsightTask(task_type=name, **entry) for name, entry in raw.items()} diff --git a/litellm/proxy/db/model_usage_rollup.py b/litellm/proxy/db/model_usage_rollup.py new file mode 100644 index 00000000000..808c3528051 --- /dev/null +++ b/litellm/proxy/db/model_usage_rollup.py @@ -0,0 +1,86 @@ +from datetime import datetime +from typing import Final + +from pydantic import TypeAdapter, ValidationError + +from litellm.constants import ( + INTERNAL_CALL_ORIGIN_METADATA_KEY, + MODEL_INSIGHTS_DEFAULT_TASK, + MODEL_INSIGHTS_TASK_TAG_PREFIX, +) +from litellm.proxy._types import SpendLogsPayload +from litellm.proxy.db.model_insights_tasks import load_model_insight_tasks +from litellm.proxy.utils import PrismaClient +from litellm.repositories.table_repositories import DailyModelUsageRepository + +_METADATA: Final = TypeAdapter(dict[str, object]) +_TAGS: Final = TypeAdapter(list[object]) + + +def model_usage_task_type(request_tags: str) -> str: + try: + tags: Final = _TAGS.validate_json(request_tags) + except ValidationError: + return MODEL_INSIGHTS_DEFAULT_TASK + for tag in tags: + if isinstance(tag, str) and tag.startswith(MODEL_INSIGHTS_TASK_TAG_PREFIX): + task = tag.removeprefix(MODEL_INSIGHTS_TASK_TAG_PREFIX) + if task in load_model_insight_tasks(): + return task + return MODEL_INSIGHTS_DEFAULT_TASK + + +def _is_internal_call(metadata: str) -> bool: + try: + decoded: Final = _METADATA.validate_json(metadata) + except ValidationError: + return False + return bool(decoded.get(INTERNAL_CALL_ORIGIN_METADATA_KEY)) + + +def _date_from_start_time(start_time: datetime | str) -> str | None: + if isinstance(start_time, datetime): + return start_time.date().isoformat() + return start_time[:10] if len(start_time) >= 10 else None + + +async def increment_daily_model_usage(prisma_client: PrismaClient, payload: SpendLogsPayload) -> None: + date: Final = _date_from_start_time(payload["startTime"]) + if date is None or _is_internal_call(payload["metadata"]): + return + + model: Final = payload["model"] or "unknown" + model_group: Final = payload["model_group"] or model + provider: Final = payload["custom_llm_provider"] or "unknown" + task_type: Final = model_usage_task_type(payload["request_tags"]) + successful: Final = 1 if payload["status"] == "success" else 0 + failed: Final = 1 - successful + key: Final = { + "date": date, + "model_group": model_group, + "model": model, + "custom_llm_provider": provider, + "task_type": task_type, + } + await DailyModelUsageRepository(prisma_client).table.upsert( + where={"date_model_group_model_custom_llm_provider_task_type": key}, + data={ + "create": { + **key, + "spend": payload["spend"], + "prompt_tokens": payload["prompt_tokens"], + "completion_tokens": payload["completion_tokens"], + "request_count": 1, + "successful_requests": successful, + "failed_requests": failed, + }, + "update": { + "spend": {"increment": payload["spend"]}, + "prompt_tokens": {"increment": payload["prompt_tokens"]}, + "completion_tokens": {"increment": payload["completion_tokens"]}, + "request_count": {"increment": 1}, + "successful_requests": {"increment": successful}, + "failed_requests": {"increment": failed}, + }, + }, + ) diff --git a/litellm/proxy/management_endpoints/model_insights_endpoints.py b/litellm/proxy/management_endpoints/model_insights_endpoints.py new file mode 100644 index 00000000000..dbdaa59d7d4 --- /dev/null +++ b/litellm/proxy/management_endpoints/model_insights_endpoints.py @@ -0,0 +1,220 @@ +from collections.abc import Mapping +from datetime import date, datetime, timedelta, timezone +from typing import Annotated, Final + +from fastapi import APIRouter, Depends, HTTPException, Query +from pydantic import BaseModel, Field, TypeAdapter + +from litellm.constants import MODEL_INSIGHTS_DEFAULT_TASK, MODEL_INSIGHTS_MAX_RANGE_DAYS, MODEL_INSIGHTS_TOP_MODELS +from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.db.model_insights_tasks import load_model_insight_tasks +from litellm.repositories.table_repositories import DailyModelUsageRepository +from litellm.types.model_insights import ( + ModelInsightDailyMetric, + ModelInsightMetric, + ModelInsightsMetric, + ModelInsightsResponse, + ModelInsightTask, + ModelInsightTasksResponse, + ModelInsightTaskSummary, +) + +router: Final = APIRouter() + + +class _Sums(BaseModel): + spend: float = 0.0 + prompt_tokens: int = 0 + completion_tokens: int = 0 + request_count: int = 0 + successful_requests: int = 0 + failed_requests: int = 0 + + +class _GroupedModel(BaseModel): + model_group: str + model: str + custom_llm_provider: str + sums: _Sums = Field(alias="_sum") + + +class _GroupedDaily(_GroupedModel): + date: str + + +class _GroupedTask(_GroupedModel): + task_type: str + + +_MODEL_ROWS: Final = TypeAdapter(list[_GroupedModel]) +_DAILY_ROWS: Final = TypeAdapter(list[_GroupedDaily]) +_TASK_ROWS: Final = TypeAdapter(list[_GroupedTask]) +_UNCATEGORIZED_TASK: Final = ModelInsightTask( + task_type=MODEL_INSIGHTS_DEFAULT_TASK, label="Uncategorized", category="General" +) +_SUM_FIELDS: Final = { + "spend": True, + "prompt_tokens": True, + "completion_tokens": True, + "request_count": True, + "successful_requests": True, + "failed_requests": True, +} + + +def _parse_date(value: str | None, fallback: date) -> date: + if value is None: + return fallback + try: + return date.fromisoformat(value) + except ValueError as exc: + raise HTTPException(status_code=400, detail="Dates must use YYYY-MM-DD") from exc + + +def _metric(row: _GroupedModel) -> ModelInsightMetric: + return ModelInsightMetric( + model_group=row.model_group, + model=row.model, + provider=row.custom_llm_provider, + spend=row.sums.spend, + prompt_tokens=row.sums.prompt_tokens, + completion_tokens=row.sums.completion_tokens, + requests=row.sums.request_count, + successful_requests=row.sums.successful_requests, + failed_requests=row.sums.failed_requests, + ) + + +def _rank_value(row: _GroupedModel, metric: ModelInsightsMetric) -> float: + if metric == "requests": + return row.sums.request_count + if metric == "spend": + return row.sums.spend + return row.sums.prompt_tokens + row.sums.completion_tokens + + +def _top_model_rows(rows: list[_GroupedModel], metric: ModelInsightsMetric) -> list[_GroupedModel]: + return sorted(rows, key=lambda row: _rank_value(row, metric), reverse=True)[:MODEL_INSIGHTS_TOP_MODELS] + + +def _deployment_filter(rows: list[_GroupedModel]) -> list[dict[str, str]]: + return [ + {"model_group": row.model_group, "model": row.model, "custom_llm_provider": row.custom_llm_provider} + for row in rows + ] + + +def _daily_metric(row: _GroupedDaily) -> ModelInsightDailyMetric: + return ModelInsightDailyMetric(date=row.date, **_metric(row).model_dump()) + + +def _summarize_tasks(rows: list[_GroupedTask], metric: ModelInsightsMetric) -> list[ModelInsightTaskSummary]: + catalog: Final = load_model_insight_tasks() + totals: Final[dict[str, float]] = {} + leaders: Final[dict[str, _GroupedTask]] = {} + for row in rows: + value = _rank_value(row, metric) + totals[row.task_type] = totals.get(row.task_type, 0.0) + value + leader = leaders.get(row.task_type) + if leader is None or value > _rank_value(leader, metric): + leaders[row.task_type] = row + grand: Final = sum(totals.values()) + return [ + ModelInsightTaskSummary( + **(catalog.get(task) or _UNCATEGORIZED_TASK).model_copy(update={"task_type": task}).model_dump(), + value=value, + share=value / grand * 100 if grand else 0.0, + leader=leaders[task].model_group, + provider=leaders[task].custom_llm_provider, + ) + for task, value in sorted(totals.items(), key=lambda item: item[1], reverse=True) + ] + + +def _resolve_window( + user_api_key_dict: UserAPIKeyAuth, start_date: str | None, end_date: str | None +) -> tuple[date, date, Mapping[str, object], DailyModelUsageRepository]: + from litellm.proxy.proxy_server import prisma_client + + if user_api_key_dict.user_role not in (LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY): + raise HTTPException(status_code=403, detail="Only proxy admins can view deployment-wide model insights") + if prisma_client is None: + raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value) + + end_day: Final = _parse_date(end_date, datetime.now(timezone.utc).date()) + start_day: Final = _parse_date(start_date, end_day - timedelta(days=MODEL_INSIGHTS_MAX_RANGE_DAYS - 1)) + if start_day > end_day or (end_day - start_day).days >= MODEL_INSIGHTS_MAX_RANGE_DAYS: + raise HTTPException( + status_code=400, detail=f"Date range must be between 1 and {MODEL_INSIGHTS_MAX_RANGE_DAYS} days" + ) + date_window: Final[Mapping[str, object]] = {"date": {"gte": start_day.isoformat(), "lte": end_day.isoformat()}} + return start_day, end_day, date_window, DailyModelUsageRepository(prisma_client) + + +@router.get( + "/model-insights", + tags=["model insights"], + dependencies=[Depends(user_api_key_auth)], + response_model=ModelInsightsResponse, +) +async def get_model_insights( + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + start_date: Annotated[str | None, Query(description="YYYY-MM-DD, defaults to 365 days ago")] = None, + end_date: Annotated[str | None, Query(description="YYYY-MM-DD, defaults to today")] = None, + metric: Annotated[ModelInsightsMetric, Query(description="Metric the top models are ranked by")] = "tokens", +) -> ModelInsightsResponse: + start_day, end_day, date_window, repository = _resolve_window(user_api_key_dict, start_date, end_date) + table: Final = repository.table + grouped_model_rows: Final = _MODEL_ROWS.validate_python( + await table.group_by( + by=["model_group", "model", "custom_llm_provider"], + sum=_SUM_FIELDS, + where=date_window, + ) + ) + model_rows: Final = _top_model_rows(grouped_model_rows, metric) + selected_window: Final = {**date_window, "OR": _deployment_filter(model_rows)} + daily_rows: Final = _DAILY_ROWS.validate_python( + await table.group_by( + by=["date", "model_group", "model", "custom_llm_provider"], + sum=_SUM_FIELDS, + where=selected_window, + order={"date": "asc"}, + ) + if model_rows + else [] + ) + return ModelInsightsResponse( + start_date=start_day.isoformat(), + end_date=end_day.isoformat(), + top_models=[_metric(row) for row in model_rows], + daily=[_daily_metric(row) for row in daily_rows], + ) + + +@router.get( + "/model-insights/tasks", + tags=["model insights"], + dependencies=[Depends(user_api_key_auth)], + response_model=ModelInsightTasksResponse, +) +async def get_model_insight_tasks( + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + start_date: Annotated[str | None, Query(description="YYYY-MM-DD, defaults to 365 days ago")] = None, + end_date: Annotated[str | None, Query(description="YYYY-MM-DD, defaults to today")] = None, + metric: Annotated[ModelInsightsMetric, Query(description="Metric task shares are computed from")] = "spend", +) -> ModelInsightTasksResponse: + start_day, end_day, date_window, repository = _resolve_window(user_api_key_dict, start_date, end_date) + task_rows: Final = _TASK_ROWS.validate_python( + await repository.table.group_by( + by=["task_type", "model_group", "model", "custom_llm_provider"], + sum=_SUM_FIELDS, + where=date_window, + ) + ) + return ModelInsightTasksResponse( + start_date=start_day.isoformat(), + end_date=end_day.isoformat(), + tasks=_summarize_tasks(task_rows, metric), + ) diff --git a/litellm/proxy/model_insights_tasks.json b/litellm/proxy/model_insights_tasks.json new file mode 100644 index 00000000000..17f9dabd24e --- /dev/null +++ b/litellm/proxy/model_insights_tasks.json @@ -0,0 +1,22 @@ +{ + "classification": {"label": "Classification", "category": "General"}, + "content_writing": {"label": "Content Writing", "category": "General"}, + "roleplay_fiction": {"label": "Roleplay & Fiction", "category": "General"}, + "conversation": {"label": "Conversation", "category": "General"}, + "research_reports": {"label": "Research & Reports", "category": "General"}, + "qa_knowledge": {"label": "Q&A & Knowledge", "category": "General"}, + "customer_support": {"label": "Customer Support", "category": "General"}, + "summarization": {"label": "Summarization", "category": "General"}, + "translation": {"label": "Translation", "category": "General"}, + "workflow_execution": {"label": "Workflow Execution", "category": "Agent"}, + "multi_step_planning": {"label": "Multi-step Planning", "category": "Agent"}, + "tool_dispatch": {"label": "Tool Dispatch", "category": "Agent"}, + "code_generation": {"label": "Code Generation", "category": "Code"}, + "debugging": {"label": "Debugging", "category": "Code"}, + "code_review": {"label": "Code Review", "category": "Code"}, + "frontend_ui": {"label": "Frontend & UI", "category": "Code"}, + "file_io": {"label": "File I/O", "category": "Code"}, + "shell_execution": {"label": "Shell Execution", "category": "Code"}, + "data_extraction": {"label": "Data Extraction", "category": "Data"}, + "data_transformation": {"label": "Data Transformation", "category": "Data"} +} diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 7edc565879a..03e59257f76 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -1260,6 +1260,26 @@ model LiteLLM_DailyToolSpend { @@id([date, tool_name]) } +model LiteLLM_DailyModelUsage { + date String + model_group String + model String + custom_llm_provider String + task_type String + spend Float @default(0.0) + prompt_tokens BigInt @default(0) + completion_tokens BigInt @default(0) + request_count BigInt @default(0) + successful_requests BigInt @default(0) + failed_requests BigInt @default(0) + created_at DateTime @default(now()) + updated_at DateTime @updatedAt + + @@id([date, model_group, model, custom_llm_provider, task_type]) + @@index([date]) + @@index([model_group]) +} + // Gateway request counts recorded at the ASGI edge by // BillableRequestMetricsMiddleware. This is the source of truth for SGR // (successful gateway requests): it counts what the proxy actually answered, diff --git a/litellm/repositories/__init__.py b/litellm/repositories/__init__.py index dcf9ddfc32a..7ffdcfa5ce6 100644 --- a/litellm/repositories/__init__.py +++ b/litellm/repositories/__init__.py @@ -30,6 +30,7 @@ from litellm.repositories.table_repositories import ( ConfigOverridesRepository, DailyGuardrailMetricsRepository, DailyGuardrailUsageUnitsRepository, + DailyModelUsageRepository, DailyPolicyMetricsRepository, DailyTagSpendRepository, DailyToolSpendRepository, @@ -105,6 +106,7 @@ __all__ = [ "CredentialsRepository", "DailyGuardrailMetricsRepository", "DailyGuardrailUsageUnitsRepository", + "DailyModelUsageRepository", "DailyPolicyMetricsRepository", "DailyTagSpendRepository", "DailyToolSpendRepository", diff --git a/litellm/repositories/table_repositories.py b/litellm/repositories/table_repositories.py index 1ad7a735d96..ab68f1a2bc7 100644 --- a/litellm/repositories/table_repositories.py +++ b/litellm/repositories/table_repositories.py @@ -212,6 +212,10 @@ class DailyToolSpendRepository(PrismaTableRepository["prisma_models.LiteLLM_Dail table_name = "litellm_dailytoolspend" +class DailyModelUsageRepository(PrismaTableRepository["prisma_models.LiteLLM_DailyModelUsage"]): + table_name = "litellm_dailymodelusage" + + class SpendLogGuardrailIndexRepository(PrismaTableRepository["prisma_models.LiteLLM_SpendLogGuardrailIndex"]): table_name = "litellm_spendlogguardrailindex" diff --git a/litellm/types/model_insights.py b/litellm/types/model_insights.py new file mode 100644 index 00000000000..6b7939386a7 --- /dev/null +++ b/litellm/types/model_insights.py @@ -0,0 +1,47 @@ +from typing import Literal + +from pydantic import BaseModel + +ModelInsightsMetric = Literal["requests", "spend", "tokens"] + + +class ModelInsightMetric(BaseModel): + model_group: str + model: str + provider: str + spend: float + prompt_tokens: int + completion_tokens: int + requests: int + successful_requests: int + failed_requests: int + + +class ModelInsightDailyMetric(ModelInsightMetric): + date: str + + +class ModelInsightTask(BaseModel): + task_type: str + label: str + category: str + + +class ModelInsightTaskSummary(ModelInsightTask): + value: float + share: float + leader: str + provider: str + + +class ModelInsightsResponse(BaseModel): + start_date: str + end_date: str + daily: list[ModelInsightDailyMetric] + top_models: list[ModelInsightMetric] + + +class ModelInsightTasksResponse(BaseModel): + start_date: str + end_date: str + tasks: list[ModelInsightTaskSummary] diff --git a/pyproject.toml b/pyproject.toml index 28b00379cc7..fb21d8fa23b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -317,6 +317,7 @@ include = [ "litellm/proxy/_experimental/out/**", "litellm/router_strategy/complexity_router/artifacts/*.json", "litellm/router_strategy/complexity_router/fuse_presets.json", + "litellm/proxy/model_insights_tasks.json", "litellm/proxy/client/cli/commands/codex_base_instructions.md", ] exclude = [ diff --git a/schema.prisma b/schema.prisma index 7edc565879a..03e59257f76 100644 --- a/schema.prisma +++ b/schema.prisma @@ -1260,6 +1260,26 @@ model LiteLLM_DailyToolSpend { @@id([date, tool_name]) } +model LiteLLM_DailyModelUsage { + date String + model_group String + model String + custom_llm_provider String + task_type String + spend Float @default(0.0) + prompt_tokens BigInt @default(0) + completion_tokens BigInt @default(0) + request_count BigInt @default(0) + successful_requests BigInt @default(0) + failed_requests BigInt @default(0) + created_at DateTime @default(now()) + updated_at DateTime @updatedAt + + @@id([date, model_group, model, custom_llm_provider, task_type]) + @@index([date]) + @@index([model_group]) +} + // Gateway request counts recorded at the ASGI edge by // BillableRequestMetricsMiddleware. This is the source of truth for SGR // (successful gateway requests): it counts what the proxy actually answered, diff --git a/tests/test_litellm/proxy/db/test_model_insights_tasks.py b/tests/test_litellm/proxy/db/test_model_insights_tasks.py new file mode 100644 index 00000000000..5c0786deaf5 --- /dev/null +++ b/tests/test_litellm/proxy/db/test_model_insights_tasks.py @@ -0,0 +1,19 @@ +from litellm.proxy.db.model_insights_tasks import load_model_insight_tasks +from litellm.proxy.db.model_usage_rollup import model_usage_task_type + + +def test_every_task_has_a_label_and_a_category() -> None: + tasks = load_model_insight_tasks() + + assert tasks + for name, task in tasks.items(): + assert task.task_type == name + assert task.label + assert task.category in {"General", "Agent", "Code", "Data"} + + +def test_tasks_in_the_json_file_are_the_ones_the_rollup_accepts() -> None: + for name in load_model_insight_tasks(): + assert model_usage_task_type(f'["task:{name}"]') == name + + assert model_usage_task_type('["task:not_in_the_file"]') == "uncategorized" diff --git a/tests/test_litellm/proxy/db/test_model_usage_rollup.py b/tests/test_litellm/proxy/db/test_model_usage_rollup.py new file mode 100644 index 00000000000..f54856129dc --- /dev/null +++ b/tests/test_litellm/proxy/db/test_model_usage_rollup.py @@ -0,0 +1,89 @@ +from datetime import datetime, timezone +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy.db.model_usage_rollup import increment_daily_model_usage, model_usage_task_type + + +def test_model_usage_task_type_reads_task_tag_or_defaults() -> None: + assert model_usage_task_type('["team-a", "task:classification"]') == "classification" + assert model_usage_task_type('["task:made-up"]') == "uncategorized" + assert model_usage_task_type('["debugging"]') == "uncategorized" + assert model_usage_task_type("[]") == "uncategorized" + assert model_usage_task_type("not json") == "uncategorized" + + +@pytest.mark.asyncio +async def test_increment_daily_model_usage_uses_atomic_prisma_upsert() -> None: + table = MagicMock() + table.upsert = AsyncMock() + prisma_client = MagicMock() + prisma_client.db.litellm_dailymodelusage = table + payload = { + "request_id": "request-1", + "call_type": "acompletion", + "api_key": "key", + "spend": 0.25, + "total_tokens": 30, + "prompt_tokens": 10, + "completion_tokens": 20, + "startTime": datetime(2026, 9, 28, tzinfo=timezone.utc), + "endTime": datetime(2026, 9, 28, tzinfo=timezone.utc), + "completionStartTime": None, + "model": "openai/gpt-5.4-mini", + "model_id": None, + "model_group": "fast-chat", + "mcp_namespaced_tool_name": None, + "agent_id": None, + "api_base": "", + "user": "user", + "metadata": "{}", + "cache_hit": "False", + "cache_key": "", + "request_tags": "[]", + "team_id": None, + "organization_id": None, + "end_user": None, + "requester_ip_address": None, + "custom_llm_provider": "openai", + "messages": None, + "response": None, + "proxy_server_request": None, + "session_id": None, + "request_duration_ms": 20, + "status": "success", + "litellm_call_id": None, + } + + await increment_daily_model_usage(prisma_client, payload) + + call = table.upsert.await_args.kwargs + assert call["data"]["create"]["request_count"] == 1 + assert call["data"]["update"]["completion_tokens"] == {"increment": 20} + assert call["data"]["create"]["task_type"] == "uncategorized" + + +@pytest.mark.asyncio +async def test_increment_daily_model_usage_records_task_from_request_tags() -> None: + table = MagicMock() + table.upsert = AsyncMock() + prisma_client = MagicMock() + prisma_client.db.litellm_dailymodelusage = table + payload = { + "call_type": "acompletion", + "spend": 0.1, + "prompt_tokens": 1, + "completion_tokens": 2, + "startTime": datetime(2026, 9, 28, tzinfo=timezone.utc), + "model": "gpt-5", + "model_group": "gpt-5", + "metadata": "{}", + "request_tags": '["task:debugging"]', + "custom_llm_provider": "openai", + "status": "success", + } + + await increment_daily_model_usage(prisma_client, payload) + + assert table.upsert.await_args.kwargs["data"]["create"]["task_type"] == "debugging" diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_insights_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_insights_endpoints.py new file mode 100644 index 00000000000..af58c2d9884 --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/test_model_insights_endpoints.py @@ -0,0 +1,233 @@ +from datetime import datetime, timezone +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.db.model_usage_rollup import increment_daily_model_usage +from litellm.proxy.management_endpoints.model_insights_endpoints import router + + +def _override_auth() -> UserAPIKeyAuth: + return UserAPIKeyAuth(api_key="sk-test", user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + + +def _grouped_row(*, prompt_tokens: str = "100", completion_tokens: str = "200", **dimensions: str) -> dict[str, object]: + return { + **dimensions, + "_sum": { + "spend": 1.25, + "prompt_tokens": prompt_tokens, + "completion_tokens": completion_tokens, + "request_count": "3", + "successful_requests": "3", + "failed_requests": "0", + }, + } + + +def test_model_insights_reads_only_bounded_rollup() -> None: + model = _grouped_row(model_group="fast-chat", model="openai/gpt-5.4-mini", custom_llm_provider="openai") + prompt_heavy_model = _grouped_row( + prompt_tokens="500", + completion_tokens="10", + model_group="long-context", + model="anthropic/claude-sonnet-4-5", + custom_llm_provider="anthropic", + ) + daily = _grouped_row( + date="2026-09-28", + model_group="fast-chat", + model="openai/gpt-5.4-mini", + custom_llm_provider="openai", + ) + table = MagicMock() + table.group_by = AsyncMock(side_effect=[[model, prompt_heavy_model], [daily]]) + prisma = MagicMock() + prisma.db.litellm_dailymodelusage = table + prisma.db.query_raw = AsyncMock() + prisma.db.litellm_spendlogs.find_many = AsyncMock() + app = FastAPI() + app.include_router(router) + app.dependency_overrides[user_api_key_auth] = _override_auth + + with patch("litellm.proxy.proxy_server.prisma_client", prisma): + response = TestClient(app).get("/model-insights?start_date=2026-09-09&end_date=2026-09-28") + + assert response.status_code == 200 + assert response.json()["top_models"][0]["model_group"] == "long-context" + assert "by_task" not in response.json() + assert table.group_by.await_count == 2 + prisma.db.query_raw.assert_not_awaited() + prisma.db.litellm_spendlogs.find_many.assert_not_awaited() + + +def test_model_insights_rejects_ranges_over_365_days() -> None: + prisma = MagicMock() + app = FastAPI() + app.include_router(router) + app.dependency_overrides[user_api_key_auth] = _override_auth + + with patch("litellm.proxy.proxy_server.prisma_client", prisma): + response = TestClient(app).get("/model-insights?start_date=2025-09-01&end_date=2026-09-28") + + assert response.status_code == 400 + + +def _call(table: MagicMock, query: str, path: str = "/model-insights") -> object: + prisma = MagicMock() + prisma.db.litellm_dailymodelusage = table + app = FastAPI() + app.include_router(router) + app.dependency_overrides[user_api_key_auth] = _override_auth + with patch("litellm.proxy.proxy_server.prisma_client", prisma): + return TestClient(app).get(f"{path}?start_date=2026-09-01&end_date=2026-09-28&{query}") + + +def test_model_insights_ranks_top_models_by_selected_metric() -> None: + token_heavy = _grouped_row( + prompt_tokens="9000", completion_tokens="9000", model_group="big", model="m1", custom_llm_provider="openai" + ) + request_heavy = _grouped_row( + prompt_tokens="1", completion_tokens="1", model_group="busy", model="m2", custom_llm_provider="openai" + ) + request_heavy["_sum"]["request_count"] = "500" + table = MagicMock() + table.group_by = AsyncMock(side_effect=[[token_heavy, request_heavy], []]) + + by_requests = _call(table, "metric=requests").json() + by_tokens = _call( + MagicMock(group_by=AsyncMock(side_effect=[[token_heavy, request_heavy], []])), "metric=tokens" + ).json() + + assert by_requests["top_models"][0]["model_group"] == "busy" + assert by_tokens["top_models"][0]["model_group"] == "big" + + +def test_model_insights_scopes_daily_to_ranked_deployments() -> None: + ranked = _grouped_row(model_group="shared", model="m1", custom_llm_provider="openai") + table = MagicMock() + table.group_by = AsyncMock(side_effect=[[ranked], []]) + + _call(table, "metric=tokens") + + daily_where = table.group_by.await_args_list[1].kwargs["where"] + assert daily_where["OR"] == [{"model_group": "shared", "model": "m1", "custom_llm_provider": "openai"}] + assert "model_group" not in daily_where + + +def _task_rows() -> list[dict[str, object]]: + def row(task: str, group: str, requests: str, spend: float) -> dict[str, object]: + base = _grouped_row(task_type=task, model_group=group, model=group, custom_llm_provider="openai") + base["_sum"].update({"request_count": requests, "spend": spend}) # type: ignore[union-attr] + return base + + return [ + row("debugging", "big", "1", 9.0), + row("debugging", "busy", "50", 1.0), + row("classification", "busy", "10", 1.0), + ] + + +def test_model_insight_tasks_are_summarised_on_the_server() -> None: + table = MagicMock(group_by=AsyncMock(return_value=_task_rows())) + + body = _call(table, "metric=spend", path="/model-insights/tasks").json() + + assert [(t["task_type"], t["label"], t["category"], t["leader"]) for t in body["tasks"]] == [ + ("debugging", "Debugging", "Code", "big"), + ("classification", "Classification", "General", "busy"), + ] + assert [round(t["share"], 1) for t in body["tasks"]] == [90.9, 9.1] + assert "OR" not in table.group_by.await_args.kwargs["where"] + assert "take" not in table.group_by.await_args.kwargs + + +def test_model_insight_tasks_leader_follows_the_selected_metric() -> None: + by_spend = _call(MagicMock(group_by=AsyncMock(return_value=_task_rows())), "metric=spend", "/model-insights/tasks") + by_requests = _call( + MagicMock(group_by=AsyncMock(return_value=_task_rows())), "metric=requests", "/model-insights/tasks" + ) + + assert by_spend.json()["tasks"][0]["leader"] == "big" + assert by_requests.json()["tasks"][0]["leader"] == "busy" + + +def test_model_insight_tasks_unknown_task_shows_as_uncategorized() -> None: + row = _grouped_row(task_type="uncategorized", model_group="a", model="a", custom_llm_provider="openai") + body = _call(MagicMock(group_by=AsyncMock(return_value=[row])), "metric=spend", "/model-insights/tasks").json() + + assert [(t["label"], t["category"]) for t in body["tasks"]] == [("Uncategorized", "General")] + + +def test_model_insight_tasks_require_an_admin() -> None: + app = FastAPI() + app.include_router(router) + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + api_key="sk-test", user_id="u", user_role=LitellmUserRoles.INTERNAL_USER + ) + with patch("litellm.proxy.proxy_server.prisma_client", MagicMock()): + assert TestClient(app).get("/model-insights/tasks").status_code == 403 + + +def test_model_insights_rejects_unknown_metric() -> None: + assert _call(MagicMock(group_by=AsyncMock()), "metric=bogus").status_code == 422 + + +class _InMemoryUsageTable: + def __init__(self) -> None: + self.rows: dict[tuple[str, ...], dict[str, float]] = {} + + async def upsert(self, where: dict, data: dict) -> None: + key_fields = where["date_model_group_model_custom_llm_provider_task_type"] + key = tuple(key_fields.values()) + if key not in self.rows: + self.rows[key] = {**key_fields, **{k: v for k, v in data["create"].items() if k not in key_fields}} + return + for field, change in data["update"].items(): + self.rows[key][field] += change["increment"] + + async def group_by(self, by: list[str], sum: dict, where: dict, **_: object) -> list[dict]: + grouped: dict[tuple, dict] = {} + for row in self.rows.values(): + if not where["date"]["gte"] <= row["date"] <= where["date"]["lte"]: + continue + if where.get("OR") and not any(all(row[k] == v for k, v in option.items()) for option in where["OR"]): + continue + bucket = grouped.setdefault(tuple(row[k] for k in by), {**{k: row[k] for k in by}, "_sum": {}}) + for field in sum: + bucket["_sum"][field] = bucket["_sum"].get(field, 0) + row[field] + return list(grouped.values()) + + +@pytest.mark.asyncio +async def test_model_insights_reads_back_what_the_rollup_wrote() -> None: + table = _InMemoryUsageTable() + prisma = MagicMock() + prisma.db.litellm_dailymodelusage = table + payload = { + "call_type": "acompletion", + "spend": 0.5, + "prompt_tokens": 10, + "completion_tokens": 20, + "startTime": datetime(2026, 9, 28, tzinfo=timezone.utc), + "model": "gpt-5", + "model_group": "gpt-5", + "metadata": "{}", + "request_tags": '["task:debugging"]', + "custom_llm_provider": "openai", + "status": "success", + } + + await increment_daily_model_usage(prisma, payload) + await increment_daily_model_usage(prisma, {**payload, "request_tags": "[]"}) + + body = _call(table, "metric=requests").json() + + assert [(m["model_group"], m["requests"], m["prompt_tokens"]) for m in body["top_models"]] == [("gpt-5", 2, 20)] + tasks = _call(table, "metric=requests", path="/model-insights/tasks").json()["tasks"] + assert sorted((t["task_type"], t["value"]) for t in tasks) == [("debugging", 1), ("uncategorized", 1)] + assert [(d["date"], d["requests"]) for d in body["daily"]] == [("2026-09-28", 2)] diff --git a/ui/litellm-dashboard/src/app/(dashboard)/legacyPageRoutes.ts b/ui/litellm-dashboard/src/app/(dashboard)/legacyPageRoutes.ts index cf943b331b9..bfc1b1ba4a8 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/legacyPageRoutes.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/legacyPageRoutes.ts @@ -35,6 +35,7 @@ const LEGACY_PAGE_ROUTES: ReadonlyMap = new Map( new_usage: "usage", usage: "old-usage", "cost-optimization": "cost-optimization", + "model-insights": "model-insights", agents: "agents", "router-settings": "router-settings", users: "users", diff --git a/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.test.tsx new file mode 100644 index 00000000000..67e57a794a1 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.test.tsx @@ -0,0 +1,148 @@ +import { render, screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import type React from "react"; +import { beforeEach, describe, expect, it, vi } from "vitest"; + +import ModelInsightsView from "./ModelInsightsView"; +import { apiClient } from "@/components/networking"; + +vi.mock("@/components/networking", () => ({ apiClient: { get: vi.fn() } })); +vi.mock("@/components/ui/chart", () => ({ + ChartContainer: ({ children }: { children: React.ReactNode }) =>
{children}
, + ChartTooltip: () => null, + ChartTooltipContent: () => null, +})); +vi.mock("recharts", () => ({ + Bar: () => null, + BarChart: ({ children }: { children: React.ReactNode }) =>
{children}
, + CartesianGrid: () => null, + Treemap: () => null, + XAxis: () => null, + YAxis: () => null, +})); + +const metrics = { + model_group: "fast-chat", + model: "openai/gpt-5.4-mini", + provider: "openai", + spend: 2.5, + prompt_tokens: 1000, + completion_tokens: 2000, + requests: 12, + successful_requests: 12, + failed_requests: 0, +}; + +const response = { + start_date: "2025-09-29", + end_date: "2026-09-28", + top_models: [metrics], + daily: [{ ...metrics, date: "2026-09-28" }], +}; + +const taskResponse = { + start_date: "2025-09-29", + end_date: "2026-09-28", + tasks: [ + { + task_type: "code_generation", + label: "Code Generation", + category: "Code", + value: 2.5, + share: 100, + leader: "fast-chat", + provider: "openai", + }, + ], +}; + +const mockApi = (tasks: unknown = taskResponse) => + vi + .mocked(apiClient.get) + .mockImplementation((path: string) => + path === "/model-insights/tasks" ? (tasks as Promise) : Promise.resolve(response), + ); + +describe("ModelInsightsView", () => { + beforeEach(() => { + vi.mocked(apiClient.get).mockReset(); + mockApi(Promise.resolve(taskResponse)); + }); + + it("shows the ranking with share and the task legend from the API response", async () => { + render(); + + expect(await screen.findByText("fast-chat")).toBeInTheDocument(); + expect(screen.getByText("by openai")).toBeInTheDocument(); + expect(await screen.findByText("Code")).toBeInTheDocument(); + expect(screen.getAllByText("100.0%")).toHaveLength(2); + expect(screen.getByRole("tab", { name: "tokens" })).toHaveAttribute("aria-selected", "true"); + expect(screen.getByRole("tab", { name: "log" })).toBeInTheDocument(); + expect(apiClient.get).toHaveBeenCalledWith("/model-insights", { + accessToken: "token", + query: { metric: "tokens" }, + }); + }); + + it("refetches with the selected metric so top models are ranked by it", async () => { + render(); + await screen.findByText("fast-chat"); + + await userEvent.click(screen.getByRole("tab", { name: "requests" })); + + await waitFor(() => + expect(apiClient.get).toHaveBeenCalledWith("/model-insights", { + accessToken: "token", + query: { metric: "requests" }, + }), + ); + }); + + it("does not refetch the task breakdown when the chart metric changes", async () => { + render(); + await screen.findByText("Code"); + const taskCalls = () => + vi.mocked(apiClient.get).mock.calls.filter(([path]) => path === "/model-insights/tasks").length; + const before = taskCalls(); + + await userEvent.click(screen.getByRole("tab", { name: "requests" })); + await waitFor(() => + expect(apiClient.get).toHaveBeenCalledWith("/model-insights", { + accessToken: "token", + query: { metric: "requests" }, + }), + ); + + expect(taskCalls()).toBe(before); + }); + + it("shows the API error instead of loading forever", async () => { + vi.mocked(apiClient.get).mockRejectedValue(new Error("Only proxy admins can view deployment-wide model insights")); + render(); + + expect(await screen.findByText("Could not load model insights")).toBeInTheDocument(); + expect(screen.getByText("Only proxy admins can view deployment-wide model insights")).toBeInTheDocument(); + }); + + it("keeps the previous ranking, dimmed, until the new metric's data arrives", async () => { + render(); + await screen.findByText("fast-chat"); + let resolve: (value: typeof response) => void = () => {}; + vi.mocked(apiClient.get).mockImplementation((path: string) => + path === "/model-insights/tasks" + ? Promise.resolve(taskResponse) + : new Promise((done) => (resolve = done as typeof resolve)), + ); + + await userEvent.click(screen.getByRole("tab", { name: "spend" })); + + expect( + screen.getByText("Share of tokens, with the change between the first and second half of the period"), + ).toBeInTheDocument(); + + resolve(response); + expect( + await screen.findByText("Share of spend, with the change between the first and second half of the period"), + ).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.tsx new file mode 100644 index 00000000000..5f195c9383b --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.tsx @@ -0,0 +1,352 @@ +"use client"; + +import React from "react"; +import { Bar, BarChart, CartesianGrid, Treemap, XAxis, YAxis } from "recharts"; +import { ArrowDownRight, ArrowUpRight, BarChart3, Layers, Minus } from "lucide-react"; + +import { apiClient } from "@/components/networking"; +import { extractErrorMessage } from "@/utils/errorUtils"; +import { ProviderLogo } from "@/components/molecules/models/ProviderLogo"; +import { PageHeader } from "@/components/shared/PageHeader"; +import { Alert, AlertDescription, AlertTitle } from "@/components/ui/alert"; +import { Card, CardContent, CardDescription, CardHeader, CardTitle } from "@/components/ui/card"; +import { ChartConfig, ChartContainer, ChartTooltip, ChartTooltipContent } from "@/components/ui/chart"; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; +import { Skeleton } from "@/components/ui/skeleton"; +import { Tabs, TabsList, TabsTrigger } from "@/components/ui/tabs"; +import { + buildWeeklySeries, + formatMetric, + Metric, + ModelInsightsResponse, + ModelInsightTasksResponse, + TaskSummary, + modelOrder, + rankModels, + RankedModel, +} from "./modelInsightsData"; + +const PALETTE = [ + "#ec4899", + "#a855f7", + "#f59e0b", + "#3b82f6", + "#10b981", + "#ef4444", + "#14b8a6", + "#84cc16", + "#6366f1", + "#f97316", +]; +const FALLBACK_COLOR = "#64748b"; +const CATEGORY_COLORS: Record = { + General: "#ee8650", + Agent: "#7666e4", + Code: "#5fb074", + Data: "#3b82f6", +}; +const SCALES = ["linear", "log"] as const; +const METRIC_LABELS: Record = { requests: "requests", spend: "spend", tokens: "tokens" }; +const RANKING_ROWS = 5; + +type Scale = (typeof SCALES)[number]; + +const formatDelta = (value: number) => `${value > 0 ? "+" : ""}${value.toFixed(1)}`; + +const DeltaBadge = ({ value }: { value: number }) => { + if (Math.abs(value) < 0.05) { + return ( + + 0.0 + + ); + } + const up = value > 0; + const Icon = up ? ArrowUpRight : ArrowDownRight; + return ( + + {formatDelta(value)} + + ); +}; + +const RankingRow = ({ model, rank }: { model: RankedModel; rank: number }) => ( +
  • + {rank} + +
    +

    {model.model_group}

    +

    by {model.provider}

    +
    +
    +

    {model.share.toFixed(1)}%

    + +
    +
  • +); + +type TileProps = TaskSummary & { x: number; y: number; width: number; height: number; index: number }; + +const TaskTileContent = ({ x, y, width, height, category, label, leader }: TileProps) => { + if (width <= 0 || height <= 0) return null; + const color = CATEGORY_COLORS[category] ?? FALLBACK_COLOR; + const fits = width > 90 && height > 44; + return ( + + + {fits && ( + <> + + {label} + + + {leader} + + + )} + + ); +}; + +export default function ModelInsightsView({ accessToken }: { accessToken: string | null }) { + const [loaded, setLoaded] = React.useState<{ metric: Metric; response: ModelInsightsResponse } | null>(null); + const [metric, setMetric] = React.useState("tokens"); + const [scale, setScale] = React.useState("linear"); + const [taskMetric, setTaskMetric] = React.useState("spend"); + const [taskData, setTaskData] = React.useState(null); + const [taskError, setTaskError] = React.useState(null); + const [error, setError] = React.useState(null); + + React.useEffect(() => { + if (!accessToken) return; + let cancelled = false; + apiClient + .get("/model-insights", { accessToken, query: { metric } }) + .then((response) => { + if (cancelled) return; + setError(null); + setLoaded({ metric, response }); + }) + .catch((err: unknown) => { + if (!cancelled) setError(extractErrorMessage(err)); + }); + return () => { + cancelled = true; + }; + }, [accessToken, metric]); + + React.useEffect(() => { + if (!accessToken) return; + let cancelled = false; + apiClient + .get("/model-insights/tasks", { accessToken, query: { metric: taskMetric } }) + .then((response) => { + if (cancelled) return; + setTaskError(null); + setTaskData(response); + }) + .catch((err: unknown) => { + if (!cancelled) setTaskError(extractErrorMessage(err)); + }); + return () => { + cancelled = true; + }; + }, [accessToken, taskMetric]); + + const data = loaded?.response ?? null; + const shown = loaded?.metric ?? metric; + const isStale = loaded !== null && loaded.metric !== metric; + const range = React.useMemo(() => ({ start: data?.start_date ?? "", end: data?.end_date ?? "" }), [data]); + const models = React.useMemo(() => (data ? modelOrder(data.daily, shown) : []), [data, shown]); + const series = React.useMemo( + () => (data ? buildWeeklySeries(data.daily, models, shown, range) : []), + [data, models, shown, range], + ); + const ranking = React.useMemo( + () => (data ? rankModels(data.top_models, data.daily, shown, range) : []), + [data, shown, range], + ); + const tiles = React.useMemo(() => taskData?.tasks ?? [], [taskData]); + const categoryShares = React.useMemo( + () => + [...new Set(tiles.map((tile) => tile.category))].map((category) => ({ + category, + share: tiles.filter((tile) => tile.category === category).reduce((sum, tile) => sum + tile.share, 0), + })), + [tiles], + ); + + if (error) { + return ( +
    + + Could not load model insights + {error} + +
    + ); + } + + if (!data) { + return ( +
    + + +
    + ); + } + + const chartConfig = Object.fromEntries( + models.map((model, index) => [model, { label: model, color: PALETTE[index % PALETTE.length] }]), + ) satisfies ChartConfig; + + return ( +
    + } + title="Model Leaderboard" + subtitle={`See which models your gateway used from ${data.start_date} through ${data.end_date}`} + /> + + + +
    + Top models + Weekly {METRIC_LABELS[shown]} across your gateway +
    +
    + setMetric(value as Metric)}> + + {(["requests", "spend", "tokens"] as const).map((value) => ( + + {value} + + ))} + + + setScale(value as Scale)}> + + {SCALES.map((value) => ( + + {value} + + ))} + + +
    +
    + + + + + + formatMetric(Number(value), shown)} + /> + } /> + {models.map((model, index) => ( + + ))} + + + +
    + + + + Leaderboard + + Share of {METRIC_LABELS[shown]}, with the change between the first and second half of the period + + + +
      + {ranking.slice(0, RANKING_ROWS).map((model, index) => ( + + ))} +
    +
      + {ranking.slice(RANKING_ROWS).map((model, index) => ( + + ))} +
    +
    +
    + + + +
    + + Top models by task + + + Each task's share of {METRIC_LABELS[taskMetric]}, labelled with its leading model + +
    + +
    + + {taskError && ( + + Could not load tasks + {taskError} + + )} + + ({ ...tile, name: tile.task_type }))} + dataKey="value" + isAnimationActive={false} + content={} + /> + +
      + {categoryShares.map(({ category, share }) => ( +
    • + + {category} + {share.toFixed(1)}% +
    • + ))} +
    +
    +
    + + + + Cost per session + Session cost is not estimated from request counts + + +

    + Add a stable session_id to requests to unlock accurate session-level model comparisons in a future bounded + session rollup +

    +
    +
    +
    + ); +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/modelInsightsData.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/modelInsightsData.test.ts new file mode 100644 index 00000000000..e23196bad68 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/modelInsightsData.test.ts @@ -0,0 +1,92 @@ +import { describe, expect, it } from "vitest"; + +import { buildWeeklySeries, DailyMetric, formatMetric, modelOrder, rankModels } from "./modelInsightsData"; + +const row = (over: Partial): DailyMetric => ({ + model_group: "a", + model: "a", + provider: "openai", + date: "2026-01-01", + spend: 0, + prompt_tokens: 0, + completion_tokens: 0, + requests: 0, + successful_requests: 0, + failed_requests: 0, + ...over, +}); + +describe("buildWeeklySeries", () => { + const range = { start: "2026-01-01", end: "2026-01-15" }; + + it("sums days into 7-day buckets per model", () => { + const rows = [ + row({ date: "2026-01-01", requests: 1 }), + row({ date: "2026-01-07", requests: 2 }), + row({ date: "2026-01-08", requests: 4 }), + row({ date: "2026-01-02", model_group: "b", requests: 8 }), + ]; + expect(buildWeeklySeries(rows, ["a", "b"], "requests", range)).toEqual([ + { date: "2026-01-01", a: 3, b: 8 }, + { date: "2026-01-08", a: 4, b: 0 }, + { date: "2026-01-15", a: 0, b: 0 }, + ]); + }); + + it("keeps weeks with no usage as zero instead of dropping them", () => { + const rows = [row({ date: "2026-01-01", requests: 1 }), row({ date: "2026-01-15", requests: 2 })]; + expect(buildWeeklySeries(rows, ["a"], "requests", range).map((week) => [week.date, week.a])).toEqual([ + ["2026-01-01", 1], + ["2026-01-08", 0], + ["2026-01-15", 2], + ]); + }); +}); + +describe("modelOrder", () => { + it("orders models by the selected metric, largest first", () => { + const rows = [row({ model_group: "a", spend: 1, requests: 9 }), row({ model_group: "b", spend: 5, requests: 1 })]; + expect(modelOrder(rows, "spend")).toEqual(["b", "a"]); + expect(modelOrder(rows, "requests")).toEqual(["a", "b"]); + }); +}); + +describe("rankModels", () => { + const range = { start: "2026-01-01", end: "2026-01-10" }; + const totals = [row({ model_group: "a", requests: 40 }), row({ model_group: "b", requests: 40 })]; + + it("computes share and the change in share between the first and second half of the range", () => { + const daily = [ + row({ date: "2026-01-01", model_group: "a", requests: 30 }), + row({ date: "2026-01-01", model_group: "b", requests: 10 }), + row({ date: "2026-01-10", model_group: "a", requests: 10 }), + row({ date: "2026-01-10", model_group: "b", requests: 30 }), + ]; + const ranked = rankModels(totals, daily, "requests", range); + expect(ranked.find((m) => m.model_group === "a")).toMatchObject({ share: 50, delta: -50 }); + expect(ranked.find((m) => m.model_group === "b")).toMatchObject({ share: 50, delta: 50 }); + }); + + it("splits at the middle of the range, not the middle of the days that had usage", () => { + const daily = [ + row({ date: "2026-01-01", model_group: "a", requests: 10 }), + row({ date: "2026-01-02", model_group: "b", requests: 10 }), + row({ date: "2026-01-03", model_group: "b", requests: 10 }), + ]; + const ranked = rankModels(totals, daily, "requests", range); + expect(ranked.find((m) => m.model_group === "a")?.delta).toBe(0); + }); + + it("shows no change when one half of the range has no usage to compare against", () => { + const daily = [row({ date: "2026-01-10", model_group: "a", requests: 10 })]; + const ranked = rankModels(totals, daily, "requests", range); + expect(ranked.map((m) => m.delta)).toEqual([0, 0]); + }); +}); + +describe("formatMetric", () => { + it("formats spend as currency and counts compactly", () => { + expect(formatMetric(12.5, "spend")).toBe("$12.50"); + expect(formatMetric(1_500_000, "tokens")).toBe("1.5M"); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/modelInsightsData.ts b/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/modelInsightsData.ts new file mode 100644 index 00000000000..3cadf0dc95c --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/modelInsightsData.ts @@ -0,0 +1,132 @@ +export type Metric = "requests" | "spend" | "tokens"; + +export type ModelMetric = { + model_group: string; + model: string; + provider: string; + spend: number; + prompt_tokens: number; + completion_tokens: number; + requests: number; + successful_requests: number; + failed_requests: number; +}; +export type DailyMetric = ModelMetric & { date: string }; +export type ModelInsightsResponse = { + start_date: string; + end_date: string; + daily: DailyMetric[]; + top_models: ModelMetric[]; +}; +export type TaskSummary = { + task_type: string; + label: string; + category: string; + value: number; + share: number; + leader: string; + provider: string; +}; +export type ModelInsightTasksResponse = { start_date: string; end_date: string; tasks: TaskSummary[] }; + +export type RankedModel = { model_group: string; provider: string; share: number; delta: number }; +const DAY_MS = 86_400_000; +const WEEK_DAYS = 7; + +export const metricValue = (row: ModelMetric, metric: Metric) => { + if (metric === "requests") return row.requests; + if (metric === "spend") return row.spend; + return row.prompt_tokens + row.completion_tokens; +}; + +const COMPACT_SPEND_FROM = 10_000; + +export const formatMetric = (value: number, metric: Metric) => { + if (metric === "spend") { + const compact = value >= COMPACT_SPEND_FROM; + const options: Intl.NumberFormatOptions = { + style: "currency", + currency: "USD", + notation: compact ? "compact" : "standard", + maximumFractionDigits: compact ? 1 : 2, + }; + return new Intl.NumberFormat("en-US", options).format(value); + } + return new Intl.NumberFormat("en-US", { notation: "compact", maximumFractionDigits: 1 }).format(value); +}; + +const toDay = (date: string) => Date.parse(`${date}T00:00:00Z`); +const isoDay = (ms: number) => new Date(ms).toISOString().slice(0, 10); + +export type DateRange = { start: string; end: string }; + +export const modelOrder = (rows: DailyMetric[], metric: Metric) => { + const totals = new Map(); + for (const row of rows) totals.set(row.model_group, (totals.get(row.model_group) ?? 0) + metricValue(row, metric)); + return [...totals.entries()].sort((a, b) => b[1] - a[1]).map(([model]) => model); +}; + +export const buildWeeklySeries = (rows: DailyMetric[], models: string[], metric: Metric, range: DateRange) => { + const weekMs = WEEK_DAYS * DAY_MS; + const origin = toDay(range.start); + const weekCount = Math.floor((toDay(range.end) - origin) / weekMs) + 1; + const buckets = Array.from({ length: weekCount }, (_, week) => ({ + date: isoDay(origin + week * weekMs), + ...Object.fromEntries(models.map((model) => [model, 0])), + })) as Record[]; + for (const row of rows) { + const bucket = buckets[Math.floor((toDay(row.date) - origin) / weekMs)]; + if (bucket) bucket[row.model_group] = Number(bucket[row.model_group] ?? 0) + metricValue(row, metric); + } + return buckets; +}; + +const shareByModel = (rows: { model_group: string; provider: string }[], values: number[]) => { + const totals = new Map(); + rows.forEach((row, index) => { + const current = totals.get(row.model_group) ?? { provider: row.provider, value: 0 }; + totals.set(row.model_group, { provider: row.provider, value: current.value + values[index] }); + }); + const grand = [...totals.values()].reduce((sum, entry) => sum + entry.value, 0); + return { totals, grand }; +}; + +const halfShares = (daily: DailyMetric[], metric: Metric, range: DateRange) => { + const midpoint = isoDay(toDay(range.start) + Math.floor((toDay(range.end) - toDay(range.start)) / 2 + DAY_MS / 2)); + const share = (rows: DailyMetric[]) => { + const { totals, grand } = shareByModel( + rows, + rows.map((row) => metricValue(row, metric)), + ); + return { + hasUsage: grand > 0, + of: (model: string) => (grand === 0 ? 0 : ((totals.get(model)?.value ?? 0) / grand) * 100), + }; + }; + return { + earlier: share(daily.filter((row) => row.date < midpoint)), + later: share(daily.filter((row) => row.date >= midpoint)), + }; +}; + +export const rankModels = ( + rows: ModelMetric[], + daily: DailyMetric[], + metric: Metric, + range: DateRange, +): RankedModel[] => { + const { totals, grand } = shareByModel( + rows, + rows.map((row) => metricValue(row, metric)), + ); + const { earlier, later } = halfShares(daily, metric, range); + const comparable = earlier.hasUsage && later.hasUsage; + return [...totals.entries()] + .sort((a, b) => b[1].value - a[1].value) + .map(([model_group, entry]) => ({ + model_group, + provider: entry.provider, + share: grand === 0 ? 0 : (entry.value / grand) * 100, + delta: comparable ? later.of(model_group) - earlier.of(model_group) : 0, + })); +}; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/model-insights/page.tsx b/ui/litellm-dashboard/src/app/(dashboard)/model-insights/page.tsx new file mode 100644 index 00000000000..01a673ce10f --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/model-insights/page.tsx @@ -0,0 +1,9 @@ +"use client"; + +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import ModelInsightsView from "./_components/ModelInsightsView"; + +export default function ModelInsightsPage() { + const { accessToken } = useAuthorized(); + return ; +} diff --git a/ui/litellm-dashboard/src/components/leftnav.tsx b/ui/litellm-dashboard/src/components/leftnav.tsx index 9d772f45153..824e8e14a1c 100644 --- a/ui/litellm-dashboard/src/components/leftnav.tsx +++ b/ui/litellm-dashboard/src/components/leftnav.tsx @@ -204,6 +204,17 @@ const menuGroups: MenuGroup[] = [ roles: [...all_admin_roles, ...internalUserRoles], label: "Usage", }, + { + key: "model-insights", + page: "model-insights", + icon: , + roles: all_admin_roles, + label: ( + + Model Leaderboard + + ), + }, { key: "cost-optimization", page: "cost-optimization", diff --git a/ui/litellm-dashboard/src/components/page_metadata.ts b/ui/litellm-dashboard/src/components/page_metadata.ts index 0f2ff639bb3..6f76be5fb07 100644 --- a/ui/litellm-dashboard/src/components/page_metadata.ts +++ b/ui/litellm-dashboard/src/components/page_metadata.ts @@ -20,6 +20,7 @@ export const pageDescriptions: Record = { "vector-stores": "Manage vector databases for embeddings", new_usage: "View usage analytics and metrics", "cost-optimization": "Track and configure cost-saving features: prompt compression, caching, and auto routing", + "model-insights": "Model Leaderboard: compare usage, spend, tokens, and task mix across this gateway", logs: "Access request and response logs", "guardrails-monitor": "Monitor guardrail performance and view logs", users: "Manage internal user accounts and permissions", diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 6b4087a4664..93f4718a586 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -9189,6 +9189,40 @@ export interface paths { patch: operations["mistral_proxy_route_mistral__endpoint__patch"]; trace?: never; }; + "/model-insights": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** Get Model Insights */ + get: operations["get_model_insights_model_insights_get"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/model-insights/tasks": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** Get Model Insight Tasks */ + get: operations["get_model_insight_tasks_model_insights_tasks_get"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/model/block": { parameters: { query?: never; @@ -35974,6 +36008,87 @@ export interface components { /** Id */ id: string; }; + /** ModelInsightDailyMetric */ + ModelInsightDailyMetric: { + /** Completion Tokens */ + completion_tokens: number; + /** Date */ + date: string; + /** Failed Requests */ + failed_requests: number; + /** Model */ + model: string; + /** Model Group */ + model_group: string; + /** Prompt Tokens */ + prompt_tokens: number; + /** Provider */ + provider: string; + /** Requests */ + requests: number; + /** Spend */ + spend: number; + /** Successful Requests */ + successful_requests: number; + }; + /** ModelInsightMetric */ + ModelInsightMetric: { + /** Completion Tokens */ + completion_tokens: number; + /** Failed Requests */ + failed_requests: number; + /** Model */ + model: string; + /** Model Group */ + model_group: string; + /** Prompt Tokens */ + prompt_tokens: number; + /** Provider */ + provider: string; + /** Requests */ + requests: number; + /** Spend */ + spend: number; + /** Successful Requests */ + successful_requests: number; + }; + /** ModelInsightTaskSummary */ + ModelInsightTaskSummary: { + /** Category */ + category: string; + /** Label */ + label: string; + /** Leader */ + leader: string; + /** Provider */ + provider: string; + /** Share */ + share: number; + /** Task Type */ + task_type: string; + /** Value */ + value: number; + }; + /** ModelInsightTasksResponse */ + ModelInsightTasksResponse: { + /** End Date */ + end_date: string; + /** Start Date */ + start_date: string; + /** Tasks */ + tasks: components["schemas"]["ModelInsightTaskSummary"][]; + }; + /** ModelInsightsResponse */ + ModelInsightsResponse: { + /** Daily */ + daily: components["schemas"]["ModelInsightDailyMetric"][]; + /** End Date */ + end_date: string; + /** Start Date */ + start_date: string; + /** Top Models */ + top_models: components["schemas"]["ModelInsightMetric"][]; + }; /** ModelParams */ ModelParams: { /** Litellm Params */ @@ -59332,6 +59447,78 @@ export interface operations { }; }; }; + get_model_insights_model_insights_get: { + parameters: { + query?: { + /** @description YYYY-MM-DD, defaults to 365 days ago */ + start_date?: string | null; + /** @description YYYY-MM-DD, defaults to today */ + end_date?: string | null; + /** @description Metric the top models are ranked by */ + metric?: "requests" | "spend" | "tokens"; + }; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["ModelInsightsResponse"]; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + get_model_insight_tasks_model_insights_tasks_get: { + parameters: { + query?: { + /** @description YYYY-MM-DD, defaults to 365 days ago */ + start_date?: string | null; + /** @description YYYY-MM-DD, defaults to today */ + end_date?: string | null; + /** @description Metric task shares are computed from */ + metric?: "requests" | "spend" | "tokens"; + }; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["ModelInsightTasksResponse"]; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; block_model_model_block_post: { parameters: { query?: never; From 319b08b4b1368a036c616ca2e3eded58f287aebf Mon Sep 17 00:00:00 2001 From: hsm207 Date: Tue, 29 Sep 2026 06:29:55 +0200 Subject: [PATCH 05/41] fix(google_genai): preserve proxy_server_request in completion adapter (#43536) --- litellm/google_genai/adapters/handler.py | 2 + litellm/types/google_genai/adapters.py | 1 + .../google_genai/test_google_genai_handler.py | 98 +++++++++++-------- 3 files changed, 61 insertions(+), 40 deletions(-) diff --git a/litellm/google_genai/adapters/handler.py b/litellm/google_genai/adapters/handler.py index 8df71504850..ed06aca0809 100644 --- a/litellm/google_genai/adapters/handler.py +++ b/litellm/google_genai/adapters/handler.py @@ -45,6 +45,8 @@ class GenerateContentToCompletionHandler: # Forward extra_headers for providers that require custom headers (e.g., github_copilot) if "extra_headers" in extra_kwargs: completion_kwargs["extra_headers"] = extra_kwargs["extra_headers"] + if "proxy_server_request" in extra_kwargs: + completion_kwargs["proxy_server_request"] = extra_kwargs["proxy_server_request"] if stream: completion_kwargs["stream"] = stream diff --git a/litellm/types/google_genai/adapters.py b/litellm/types/google_genai/adapters.py index 172a45b4cbc..771b362cae3 100644 --- a/litellm/types/google_genai/adapters.py +++ b/litellm/types/google_genai/adapters.py @@ -19,3 +19,4 @@ class GenerateContentCompletionKwargs(TypedDict, total=False): stream: bool metadata: dict[str, object] extra_headers: dict[str, str] | None + proxy_server_request: dict[str, object] | None diff --git a/tests/unit/google_genai/test_google_genai_handler.py b/tests/unit/google_genai/test_google_genai_handler.py index 5361d91718d..68f24c1c0f5 100644 --- a/tests/unit/google_genai/test_google_genai_handler.py +++ b/tests/unit/google_genai/test_google_genai_handler.py @@ -2,11 +2,11 @@ """ Test to verify the Google GenAI generate_content handler functionality """ + from unittest.mock import AsyncMock, MagicMock, patch import pytest - from litellm.google_genai.adapters.handler import GenerateContentToCompletionHandler from litellm.google_genai.adapters.transformation import GoogleGenAIAdapter @@ -49,9 +49,7 @@ async def test_stream_response_when_stream_requested_async(): """ # Mock a stream response mock_stream = MagicMock() - mock_stream.__aiter__ = AsyncMock( - return_value=iter([]) - ) # Return an empty async iterator + mock_stream.__aiter__ = AsyncMock(return_value=iter([])) # Return an empty async iterator # Mock the GoogleGenAIAdapter's translate_completion_output_params_streaming method with patch.object( @@ -61,13 +59,11 @@ async def test_stream_response_when_stream_requested_async(): ) as mock_translate: with patch("litellm.acompletion", return_value=mock_stream): # Call the handler with stream=True - result = ( - await GenerateContentToCompletionHandler.async_generate_content_handler( - model="gemini-pro", - contents=[{"role": "user", "parts": [{"text": "Hello"}]}], - litellm_params={}, # Empty dict for params - stream=True, - ) + result = await GenerateContentToCompletionHandler.async_generate_content_handler( + model="gemini-pro", + contents=[{"role": "user", "parts": [{"text": "Hello"}]}], + litellm_params={}, # Empty dict for params + stream=True, ) # Verify that translate_completion_output_params_streaming was called @@ -93,9 +89,7 @@ def test_stream_transformation_error_sync(): # Patch litellm.completion directly to prevent real API calls with patch("litellm.completion", return_value=mock_stream): # Call the handler with stream=True and expect a ValueError - with pytest.raises( - ValueError, match="Failed to transform streaming response" - ): + with pytest.raises(ValueError, match="Failed to transform streaming response"): GenerateContentToCompletionHandler.generate_content_handler( model="gemini-pro", contents=[{"role": "user", "parts": [{"text": "Hello"}]}], @@ -125,9 +119,7 @@ async def test_stream_transformation_error_async(): # Use AsyncMock for async function mock_litellm.acompletion = AsyncMock(return_value=mock_stream) # Call the handler with stream=True and expect a ValueError - with pytest.raises( - ValueError, match="Failed to transform streaming response" - ): + with pytest.raises(ValueError, match="Failed to transform streaming response"): await GenerateContentToCompletionHandler.async_generate_content_handler( model="gemini-pro", contents=[{"role": "user", "parts": [{"text": "Hello"}]}], @@ -153,11 +145,7 @@ def test_citation_metadata_transformation(): "candidates": [ { "content": { - "parts": [ - { - "text": "This is a video analysis response with citation metadata." - } - ], + "parts": [{"text": "This is a video analysis response with citation metadata."}], "role": "model", }, "finishReason": "STOP", @@ -232,28 +220,58 @@ def test_citation_metadata_transformation(): citation_metadata = candidate.citationMetadata # Check that citations field exists - assert hasattr( - citation_metadata, "citations" - ), "citations field should exist after transformation" + assert hasattr(citation_metadata, "citations"), "citations field should exist after transformation" # Verify the citations data is preserved - if ( - hasattr(citation_metadata, "citations") - and citation_metadata.citations - ): - assert ( - len(citation_metadata.citations) == 2 - ), "Should have 2 citations" - assert ( - citation_metadata.citations[0]["uri"] - == "https://example.com/video-source" - ) - assert ( - citation_metadata.citations[1]["uri"] - == "https://another-source.com/reference" - ) + if hasattr(citation_metadata, "citations") and citation_metadata.citations: + assert len(citation_metadata.citations) == 2, "Should have 2 citations" + assert citation_metadata.citations[0]["uri"] == "https://example.com/video-source" + assert citation_metadata.citations[1]["uri"] == "https://another-source.com/reference" print("✅ Citation metadata transformation test passed!") except Exception as e: pytest.fail(f"Citation metadata transformation failed: {e}") + + +@pytest.mark.asyncio +async def test_generate_content_adapter_preserves_proxy_server_request(): + """ + Ensure GenerateContentToCompletionHandler forwards proxy_server_request + to the downstream completion call so proxy spend logging captures the request body. + """ + from litellm.types.router import GenericLiteLLMParams + from litellm.types.utils import Choices, Message, ModelResponse + + handler = GenerateContentToCompletionHandler() + + dummy_proxy_request: dict[str, object] = { + "url": "http://localhost:4000/v1beta/models/gemini-2.0-flash:generateContent", + "method": "POST", + "headers": {"content-type": "application/json"}, + "body": {"contents": [{"role": "user", "parts": [{"text": "Hello, world!"}]}]}, + } + + gemini_data: list[dict[str, object]] = [{"role": "user", "parts": [{"text": "Hello, world!"}]}] + + mock_response = ModelResponse(choices=[Choices(message=Message(content="Hi!", role="assistant"))]) + + with patch( + "litellm.google_genai.adapters.handler.litellm.acompletion", + new_callable=AsyncMock, + ) as mock_acompletion: + mock_acompletion.return_value = mock_response + + await handler.async_generate_content_handler( + model="gemini-2.0-flash", + contents=gemini_data, + litellm_params=GenericLiteLLMParams(), + proxy_server_request=dummy_proxy_request, + metadata={"source": "unit_test"}, + ) + + assert mock_acompletion.called, "Inner acompletion was not called" + called_kwargs = mock_acompletion.call_args.kwargs + + assert "proxy_server_request" in called_kwargs, "proxy_server_request was dropped from completion_kwargs" + assert called_kwargs["proxy_server_request"] == dummy_proxy_request From 60fca8298e69fbd91ae6cb8cc54488b23acb8b78 Mon Sep 17 00:00:00 2001 From: Ankit Jha Date: Tue, 29 Sep 2026 10:00:56 +0530 Subject: [PATCH 06/41] fix(otel): send cache and reasoning tokens in langfuse usage_details (#43553) * fix(otel): send cache and reasoning tokens in langfuse usage_details The OTel V2 Langfuse mapper only sent input, output and total, so cache reads, cache writes and reasoning tokens never reached Langfuse. Emit them as input_cached_tokens, input_cache_creation and output_reasoning_tokens, and send input/output net of those buckets so Langfuse does not price the same tokens twice. Fixes #43542 * fix(otel): drop redundant comments from the usage_details change --- litellm/integrations/otel/mappers/langfuse.py | 8 ++- litellm/integrations/otel/model/payloads.py | 19 ++++++ .../otel/test_otel_v2_vendor_mappers.py | 68 +++++++++++++++++++ 3 files changed, 93 insertions(+), 2 deletions(-) diff --git a/litellm/integrations/otel/mappers/langfuse.py b/litellm/integrations/otel/mappers/langfuse.py index 9aff944cff0..68860931b76 100644 --- a/litellm/integrations/otel/mappers/langfuse.py +++ b/litellm/integrations/otel/mappers/langfuse.py @@ -56,9 +56,13 @@ class LangfuseMapper: "presence_penalty": lambda rp: rp.presence_penalty, "seed": lambda rp: rp.seed, } + # Langfuse prices every key, and litellm's prompt/completion counts include cache and reasoning tokens _USAGE_FIELDS: dict[str, Callable[[LLMUsage], AttrValue | None]] = { - "input": lambda u: u.input_tokens, - "output": lambda u: u.output_tokens, + "input": lambda u: u.uncached_input_tokens, + "input_cached_tokens": lambda u: u.cache_read_input_tokens or None, + "input_cache_creation": lambda u: u.cache_creation_input_tokens or None, + "output": lambda u: u.non_reasoning_output_tokens, + "output_reasoning_tokens": lambda u: u.reasoning_tokens or None, "total": lambda u: u.total_tokens, } diff --git a/litellm/integrations/otel/model/payloads.py b/litellm/integrations/otel/model/payloads.py index 7e47abfb20d..c007eda7707 100644 --- a/litellm/integrations/otel/model/payloads.py +++ b/litellm/integrations/otel/model/payloads.py @@ -124,6 +124,20 @@ class LLMUsage: total_tokens: int | None = None cache_creation_input_tokens: int | None = None cache_read_input_tokens: int | None = None + reasoning_tokens: int | None = None + + @property + def uncached_input_tokens(self) -> int | None: + if self.input_tokens is None: + return None + cached: Final = (self.cache_read_input_tokens or 0) + (self.cache_creation_input_tokens or 0) + return max(self.input_tokens - cached, 0) + + @property + def non_reasoning_output_tokens(self) -> int | None: + if self.output_tokens is None: + return None + return max(self.output_tokens - (self.reasoning_tokens or 0), 0) @classmethod def from_standard_logging_payload(cls, payload: StandardLoggingPayload) -> LLMUsage: @@ -135,6 +149,10 @@ class LLMUsage: prompt_details: Final[Mapping[str, object]] = ( raw_details if isinstance(raw_details, Mapping) else MappingProxyType({}) ) + raw_completion_details: Final = usage_object.get("completion_tokens_details") + completion_details: Final[Mapping[str, object]] = ( + raw_completion_details if isinstance(raw_completion_details, Mapping) else MappingProxyType({}) + ) return cls( input_tokens=as_int(payload.get("prompt_tokens")), output_tokens=as_int(payload.get("completion_tokens")), @@ -150,6 +168,7 @@ class LLMUsage: prompt_details.get("cached_tokens"), usage_object.get("prompt_cache_hit_tokens"), ), + reasoning_tokens=_cache_token_value(completion_details.get("reasoning_tokens")), ) diff --git a/tests/unit/integrations/otel/test_otel_v2_vendor_mappers.py b/tests/unit/integrations/otel/test_otel_v2_vendor_mappers.py index 1e2ae24a329..cdff9c960f3 100644 --- a/tests/unit/integrations/otel/test_otel_v2_vendor_mappers.py +++ b/tests/unit/integrations/otel/test_otel_v2_vendor_mappers.py @@ -137,6 +137,74 @@ def test_langfuse_mapper_observation_attrs(): assert attrs["langfuse.trace.metadata.team_id"] == "t1" +def _langfuse_usage_details(usage_object: Mapping[str, object]) -> dict[str, object]: + payload: Final = { + "call_type": "acompletion", + "custom_llm_provider": "openai", + "model": "gpt-4o", + "prompt_tokens": usage_object["prompt_tokens"], + "completion_tokens": usage_object["completion_tokens"], + "total_tokens": usage_object["total_tokens"], + "metadata": {"usage_object": usage_object}, + } + attrs: Final = LangfuseMapper().map(LLMCallSpanData.from_standard_logging_payload(payload)) + return json.loads(attrs["langfuse.observation.usage_details"]) + + +def test_langfuse_usage_details_split_openai_cached_and_reasoning_tokens(): + usage: Final = _langfuse_usage_details( + { + "prompt_tokens": 100, + "completion_tokens": 50, + "total_tokens": 150, + "prompt_tokens_details": {"cached_tokens": 60}, + "completion_tokens_details": {"reasoning_tokens": 30}, + } + ) + assert usage == { + "input": 40, + "input_cached_tokens": 60, + "output": 20, + "output_reasoning_tokens": 30, + "total": 150, + } + + +def test_langfuse_usage_details_split_anthropic_cache_read_and_creation_tokens(): + usage: Final = _langfuse_usage_details( + { + "prompt_tokens": 1000, + "completion_tokens": 40, + "total_tokens": 1040, + "cache_read_input_tokens": 800, + "cache_creation_input_tokens": 150, + "prompt_tokens_details": {"cached_tokens": 800, "cache_creation_tokens": 150}, + } + ) + assert usage == { + "input": 50, + "input_cached_tokens": 800, + "input_cache_creation": 150, + "output": 40, + "total": 1040, + } + + +def test_langfuse_usage_details_omit_zero_cache_and_reasoning_counts(): + usage: Final = _langfuse_usage_details( + { + "prompt_tokens": 12, + "completion_tokens": 8, + "total_tokens": 20, + "cache_read_input_tokens": 0, + "cache_creation_input_tokens": 0, + "prompt_tokens_details": {"cached_tokens": 0}, + "completion_tokens_details": {"reasoning_tokens": 0}, + } + ) + assert usage == {"input": 12, "output": 8, "total": 20} + + def test_langfuse_mapper_names_the_trace_from_the_caller(): named = LangfuseMapper().map(_llm_call(trace=TraceControls(name="nightly-eval"))) assert named["langfuse.trace.name"] == "nightly-eval" From 0fb93ed9deffd0b2ea790815ed57c2aab4eb3665 Mon Sep 17 00:00:00 2001 From: Chase Date: Mon, 28 Sep 2026 21:35:08 -0700 Subject: [PATCH 07/41] fix(vertex_ai): forward the per-turn-control beta for per-message output_config (#43558) Claude Code attaches output_config to mid-conversation system messages and sends the per-turn-control-2026-07-01 beta with it. The Vertex beta map dropped that beta, so Vertex rejected the body with 'messages.N.output_config: Extra inputs are not permitted'. Forward the beta for vertex_ai, the way azure_ai already does, and add it on the Vertex Messages path whenever a message carries output_config. --- litellm/anthropic_beta_headers_config.json | 2 +- .../transformation.py | 4 ++ ...est_anthropic_messages_per_turn_control.py | 7 +-- ...artner_models_anthropic_messages_config.py | 54 +++++++++++++++++++ 4 files changed, 63 insertions(+), 4 deletions(-) diff --git a/litellm/anthropic_beta_headers_config.json b/litellm/anthropic_beta_headers_config.json index 71e7081b440..3a938007633 100644 --- a/litellm/anthropic_beta_headers_config.json +++ b/litellm/anthropic_beta_headers_config.json @@ -194,7 +194,7 @@ "mcp-servers-2025-12-04": null, "output-128k-2025-02-19": null, "structured-output-2024-03-01": null, - "per-turn-control-2026-07-01": null, + "per-turn-control-2026-07-01": "per-turn-control-2026-07-01", "prompt-caching-scope-2026-01-05": null, "skills-2025-10-02": null, "structured-outputs-2025-11-13": null, diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py index 38376ea17c3..be2dacd23c2 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py @@ -3,6 +3,7 @@ from typing import Any, Final from litellm.llms.anthropic.common_utils import AnthropicModelInfo from litellm.llms.anthropic.pass_through.messages.transformation import ( AnthropicMessagesConfig, + _messages_carry_output_config, ) from litellm.types.llms.anthropic import ( ANTHROPIC_BETA_HEADER_VALUES, @@ -111,6 +112,9 @@ class VertexAIPartnerModelsAnthropicMessagesConfig(AnthropicMessagesConfig, Vert if optional_params.get("safeguards") is not None: beta_values.add(ANTHROPIC_BETA_HEADER_VALUES.DANGEROUS_TOOL_USE_2026_09_03.value) + if _messages_carry_output_config(messages): + beta_values.add(ANTHROPIC_BETA_HEADER_VALUES.PER_TURN_CONTROL_2026_07_01.value) + if beta_values: headers["anthropic-beta"] = ",".join(beta_values) diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_per_turn_control.py b/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_per_turn_control.py index 4197192e4af..ef6a6e72d07 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_per_turn_control.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_per_turn_control.py @@ -95,15 +95,16 @@ def test_added_per_turn_control_beta_survives_the_anthropic_allowlist(): assert PER_TURN_CONTROL in _betas(filtered) -@pytest.mark.parametrize("provider", ["bedrock", "bedrock_converse", "vertex_ai", "databricks"]) +@pytest.mark.parametrize("provider", ["bedrock", "bedrock_converse", "databricks"]) def test_per_turn_control_beta_is_dropped_for_providers_without_it(provider): filtered = update_headers_with_filtered_beta(headers={"anthropic-beta": PER_TURN_CONTROL}, provider=provider) assert "anthropic-beta" not in filtered -def test_per_turn_control_beta_is_forwarded_for_azure_ai(): - filtered = update_headers_with_filtered_beta(headers={"anthropic-beta": PER_TURN_CONTROL}, provider="azure_ai") +@pytest.mark.parametrize("provider", ["azure_ai", "vertex_ai"]) +def test_per_turn_control_beta_is_forwarded_for_providers_with_it(provider): + filtered = update_headers_with_filtered_beta(headers={"anthropic-beta": PER_TURN_CONTROL}, provider=provider) assert _betas(filtered) == {PER_TURN_CONTROL} diff --git a/tests/unit/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py b/tests/unit/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py index e3ae891f0d9..67a32cc82bc 100644 --- a/tests/unit/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py +++ b/tests/unit/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py @@ -126,6 +126,60 @@ def test_no_safeguards_leaves_dangerous_tool_use_beta_header_out(): assert "dangerous-tool-use-2026-09-03" not in updated_headers.get("anthropic-beta", "") +def _validate_vertex_headers(client_headers, messages): + config = VertexAIPartnerModelsAnthropicMessagesConfig() + litellm_params = { + "vertex_ai_project": "test-project", + "vertex_ai_location": "global", + "vertex_credentials": "{}", + } + + with ( + patch.object(config, "_ensure_access_token", return_value=("token", "test-project")), + patch.object(config, "get_complete_vertex_url", return_value="https://mock-url"), + ): + updated_headers, _ = config.validate_anthropic_messages_environment( + headers=client_headers, + model="claude-opus-5-5", + messages=messages, + optional_params={"max_tokens": 64}, + litellm_params=litellm_params, + api_base=None, + ) + return updated_headers + + +@pytest.mark.parametrize( + "client_headers", + [{"anthropic-beta": "per-turn-control-2026-07-01"}, {}], + ids=["client_sends_beta", "client_omits_beta"], +) +def test_per_message_output_config_reaches_vertex_with_per_turn_control_beta(client_headers, monkeypatch): + """Vertex rejects a message-level `output_config` as an extra input unless the per-turn-control beta is present, so the beta must survive the Vertex beta filter.""" + from litellm import anthropic_beta_headers_manager + from litellm.anthropic_beta_headers_manager import update_headers_with_filtered_beta + + monkeypatch.setenv("LITELLM_LOCAL_ANTHROPIC_BETA_HEADERS", "True") + monkeypatch.setattr(anthropic_beta_headers_manager, "_BETA_HEADERS_CONFIG", None) + + messages = [ + {"role": "user", "content": [{"type": "text", "text": "Hello"}]}, + {"role": "system", "content": [{"type": "text", "text": "# Environment"}], "output_config": {"effort": "low"}}, + ] + + filtered = update_headers_with_filtered_beta( + headers=_validate_vertex_headers(client_headers, messages), provider="vertex_ai" + ) + + assert filtered["anthropic-beta"].split(",").count("per-turn-control-2026-07-01") == 1 + + +def test_no_per_message_output_config_leaves_per_turn_control_beta_out(): + headers = _validate_vertex_headers({}, [{"role": "user", "content": "Hello"}]) + + assert "per-turn-control-2026-07-01" not in headers.get("anthropic-beta", "") + + def test_web_search_header_not_added_without_tool(): """Test that beta header is NOT added when web search tool is not present""" config = VertexAIPartnerModelsAnthropicMessagesConfig() From 7b2cbf6e7f6c22d8d057858187e88c73e5025883 Mon Sep 17 00:00:00 2001 From: Flexomatic81 Date: Tue, 29 Sep 2026 06:40:01 +0200 Subject: [PATCH 08/41] fix(cost-map): add tool calling and reasoning flags, correct max output for nebius DeepSeek-V4.1-Flash (#43588) --- litellm/model_prices_and_context_window_backup.json | 6 ++++-- model_prices_and_context_window.json | 6 ++++-- 2 files changed, 8 insertions(+), 4 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 7bb83a95bd4..fc48c17b506 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -39699,11 +39699,13 @@ "input_cost_per_token": 3e-07, "litellm_provider": "nebius", "max_input_tokens": 1048576, - "max_output_tokens": 1048576, - "max_tokens": 1048576, + "max_output_tokens": 384000, + "max_tokens": 384000, "mode": "chat", "output_cost_per_token": 1.2e-06, "source": "https://tokenfactory.nebius.com/endpoints?modals=endpoint-details&model-id=deepseek-ai/DeepSeek-V4.1-Flash", + "supports_function_calling": true, + "supports_reasoning": true, "supports_vision": true }, "nebius/MiniMaxAI/MiniMax-M2.5": { diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 7bb83a95bd4..fc48c17b506 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -39699,11 +39699,13 @@ "input_cost_per_token": 3e-07, "litellm_provider": "nebius", "max_input_tokens": 1048576, - "max_output_tokens": 1048576, - "max_tokens": 1048576, + "max_output_tokens": 384000, + "max_tokens": 384000, "mode": "chat", "output_cost_per_token": 1.2e-06, "source": "https://tokenfactory.nebius.com/endpoints?modals=endpoint-details&model-id=deepseek-ai/DeepSeek-V4.1-Flash", + "supports_function_calling": true, + "supports_reasoning": true, "supports_vision": true }, "nebius/MiniMaxAI/MiniMax-M2.5": { From 85dc7cb62efa52794982af90164eb19a41140b80 Mon Sep 17 00:00:00 2001 From: fedaeho <39611158+fedaeho@users.noreply.github.com> Date: Tue, 29 Sep 2026 13:46:33 +0900 Subject: [PATCH 09/41] fix(proxy): resolve model_group_alias in the zero-cost budget predicate (#43512) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `_is_model_cost_zero()` reads a group's cost through `Router.get_model_group_info()`, which resolves `model_group_alias`, and then gates that on `_is_cost_explicitly_configured()`, which scanned `Router.model_list` for an exact `model_name` match. Alias names live only in `Router.model_group_alias` and are never `model_name` entries, so the scan found nothing and returned False. That False means "the zero cost was defaulted, not configured" (the sparse auto-registration gate added for #24770), so a model priced explicitly at 0 had budget enforced against it when requested through an alias, while the same deployment under its own name was exempt. Both names route to the same deployment and add nothing to spend. The two lookups in one function disagreeing is the bug, so they now share one resolution: `_is_cost_explicitly_configured()` resolves through `Router.get_model_list()`, the same alias-aware path `get_model_group_info()` takes. That also reaches a deployment which prices itself through its `model_info` block, whose cost-map entry lands under the deployment id. `_group_declares_explicit_cost()` was an alias-aware copy of this function, wired only into `model_has_no_cost_mapping()` and never into the budget path; its body is what `_is_cost_explicitly_configured()` now carries, and both callers share it so the two cannot drift apart again. `_has_ptu_flat_cost()` scanned `model_list` the same way and runs after the gate above, so resolving one without the other would let an aliased PTU group — explicit zero per-token price alongside a flat capacity cost — pass as free. It resolves the same way now. Tests cover the predicate and the request path it feeds: over-budget requests through `_should_skip_budget_checks()` into `common_checks()` for an aliased free model (allowed) and an aliased paid model (refused), the predicate for free, paid, PTU, hidden and dangling aliases, and `model_has_no_cost_mapping()` through an alias so the other caller of the shared check stays covered. Unchanged: priced groups (the predicate returns False before the gate), unmapped groups whose zero cost was defaulted (#24770), hidden aliases and aliases pointing at a nonexistent group (`get_model_group_info()` returns None for both, so the cost is unknown and budget is enforced), and non-aliased PTU groups. Co-authored-by: Claude Opus 5 (1M context) --- litellm/proxy/auth/auth_checks.py | 43 +++--- .../proxy/auth/test_auth_checks.py | 24 ++++ .../test_unmapped_model_budget_enforcement.py | 136 ++++++++++++++++++ .../test_zero_cost_model_budget_bypass.py | 81 +++++++++++ 4 files changed, 257 insertions(+), 27 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 12d420141f1..51cc70c010b 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -567,10 +567,12 @@ def _has_ptu_flat_cost(model: str, llm_router: "Router") -> bool: Such a deployment carries an explicit zero per-token price so the flat cost is not charged twice, which otherwise reads here as a free model and waives every budget check for it. + + Resolved through ``Router.get_model_list()``, which includes ``model_group_alias``, because + this runs after the explicit-cost gate: resolving that gate alone would let an aliased PTU + group through as free. """ - for deployment in llm_router.model_list: - if deployment.get("model_name") != model: - continue + for deployment in llm_router.get_model_list(model_name=model) or (): model_info = deployment.get("model_info") or _NO_MODEL_INFO if model_info.get("ptu_count") is not None and model_info.get("cost_per_ptu_per_hour") is not None: return True @@ -586,14 +588,19 @@ def _is_cost_explicitly_configured(model: str, llm_router: "Router") -> bool: cost map, it creates a sparse entry like {"id": ""} with no cost fields. _get_model_info_helper() then defaults missing costs to 0. This function detects that scenario by checking the raw model_cost entry. + + The group is resolved through ``Router.get_model_list()``, the same resolution + ``get_model_group_info()`` applies when the caller reads the cost a few lines earlier, so the + two lookups cannot disagree: names defined in ``Router.model_group_alias`` are not + ``model_name`` entries in ``Router.model_list``, and scanning that list by exact name reported + every aliased group as unconfigured. It also reaches a deployment that prices itself through + its ``model_info`` block, whose entry lands in the cost map under the deployment id. """ - for deployment in llm_router.model_list: - if deployment.get("model_name") != model: - continue - model_id = deployment.get("model_info", {}).get("id") + for deployment in llm_router.get_model_list(model_name=model) or (): + model_id = (deployment.get("model_info") or _EMPTY_COST_ENTRY).get("id") if model_id is None: continue - raw_entry = litellm.model_cost.get(model_id, {}) + raw_entry = litellm.model_cost.get(model_id, _EMPTY_COST_ENTRY) if "input_cost_per_token" in raw_entry or "output_cost_per_token" in raw_entry: return True return False @@ -648,24 +655,6 @@ def _model_group_has_pricing(model: str, llm_router: "Router") -> bool: return False -def _group_declares_explicit_cost(model: str, llm_router: "Router") -> bool: - """ - Alias-aware counterpart to ``_is_cost_explicitly_configured``, which resolves the model group - the same way ``_model_group_has_pricing`` does. A deployment that prices itself through its - ``model_info`` block lands in the cost map under its deployment id rather than in its - litellm_params, and reaching that entry through the router's own resolution keeps an alias - pointing at such a group from being read as unpriced. - """ - for deployment in llm_router.get_model_list(model_name=model) or (): - model_id = (deployment.get("model_info") or _EMPTY_COST_ENTRY).get("id") - if model_id is None: - continue - raw_entry = litellm.model_cost.get(model_id, _EMPTY_COST_ENTRY) - if "input_cost_per_token" in raw_entry or "output_cost_per_token" in raw_entry: - return True - return False - - def model_has_no_cost_mapping(model: str | None, llm_router: Router | None) -> bool: if not model or llm_router is None: return False @@ -676,7 +665,7 @@ def model_has_no_cost_mapping(model: str | None, llm_router: Router | None) -> b if _model_group_has_pricing(model=model, llm_router=llm_router): return False - return not _group_declares_explicit_cost(model=model, llm_router=llm_router) + return not _is_cost_explicitly_configured(model=model, llm_router=llm_router) def _unpriced_models_in_request(model: str | list[str] | None, llm_router: Router | None) -> tuple[str, ...]: diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index f014e9c26d1..dabd97cff0b 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -8620,6 +8620,30 @@ def test_model_has_no_cost_mapping_unpriced_model_is_true(): assert model_has_no_cost_mapping(model="unpriced-group", llm_router=router) is True +def test_model_has_no_cost_mapping_resolves_model_group_alias(): + """This helper and the zero-cost budget predicate share one explicit-cost check, so the + alias resolution it depends on has to keep working for both.""" + from litellm.proxy.auth.auth_checks import model_has_no_cost_mapping + from litellm.router import Router + + router = Router( + model_list=[ + { + "model_name": "priced-group", + "litellm_params": {"model": "gpt-3.5-turbo", "api_key": "sk-test"}, + }, + { + "model_name": "unpriced-group", + "litellm_params": {"model": UNPRICED_UNDERLYING_MODEL, "api_key": "sk-test"}, + }, + ], + model_group_alias={"priced-alias": "priced-group", "unpriced-alias": "unpriced-group"}, + ) + + assert model_has_no_cost_mapping(model="priced-alias", llm_router=router) is False + assert model_has_no_cost_mapping(model="unpriced-alias", llm_router=router) is True + + def test_model_has_no_cost_mapping_no_model_or_router_is_false(): from litellm.proxy.auth.auth_checks import model_has_no_cost_mapping diff --git a/tests/test_litellm/proxy/auth/test_unmapped_model_budget_enforcement.py b/tests/test_litellm/proxy/auth/test_unmapped_model_budget_enforcement.py index bbe343bcede..7665008a6a6 100644 --- a/tests/test_litellm/proxy/auth/test_unmapped_model_budget_enforcement.py +++ b/tests/test_litellm/proxy/auth/test_unmapped_model_budget_enforcement.py @@ -190,6 +190,142 @@ class TestUnmappedModelBudgetEnforcement: assert "input_cost_per_token" not in litellm.model_cost.get("alias-id", {}) assert _is_model_cost_zero(model="smart-router", llm_router=router) is False + def test_model_group_alias_to_free_model_bypasses_budget(self): + """A zero-cost group reached through model_group_alias bypasses budget, like its own name. + + Both names route to the same deployment and add nothing to spend, so refusing one of + them denies a request on spend it cannot produce. + """ + router = Router( + model_list=[ + { + "model_name": "free-model", + "litellm_params": { + "model": "ollama/llama2", + "api_base": "http://localhost:11434", + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + }, + "model_info": {"id": "free-model-id"}, + }, + ], + model_group_alias={"free-model-alias": "free-model"}, + ) + + assert _is_model_cost_zero(model="free-model", llm_router=router) is True + assert _is_model_cost_zero(model="free-model-alias", llm_router=router) is True, ( + "An alias pointing at an explicitly-zero-cost group must be read as free, like its own name" + ) + + def test_model_group_alias_item_form_bypasses_budget(self): + """The dict alias form ({"model": ..., "hidden": False}) resolves like the string form.""" + router = Router( + model_list=[ + { + "model_name": "free-model", + "litellm_params": { + "model": "ollama/llama2", + "api_base": "http://localhost:11434", + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + }, + "model_info": {"id": "free-model-id"}, + }, + ], + model_group_alias={"free-model-alias": {"model": "free-model", "hidden": False}}, + ) + + assert _is_model_cost_zero(model="free-model-alias", llm_router=router) is True + + def test_model_group_alias_to_paid_model_enforces_budget(self): + """An alias does not turn a priced group into a free one.""" + router = Router( + model_list=[ + { + "model_name": "paid-model", + "litellm_params": {"model": "gpt-3.5-turbo", "api_key": "sk-fake"}, + "model_info": {"id": "paid-model-id"}, + }, + ], + model_group_alias={"paid-model-alias": "paid-model"}, + ) + + assert _is_model_cost_zero(model="paid-model-alias", llm_router=router) is False + + def test_model_group_alias_to_ptu_flat_cost_enforces_budget(self): + """A PTU group keeps budget enforced through an alias. + + Its explicit zero per-token price exists so the flat capacity cost is not charged twice, + so the PTU check has to resolve the alias too — resolving only the explicit-cost gate + would let this through as free. + """ + router = Router( + model_list=[ + { + "model_name": "ptu-model", + "litellm_params": { + "model": "azure/ptu-deployment", + "api_base": "https://fake.openai.azure.com", + "api_key": "sk-fake", + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + }, + "model_info": { + "id": "ptu-model-id", + "ptu_count": 100, + "cost_per_ptu_per_hour": 2.0, + }, + }, + ], + model_group_alias={"ptu-model-alias": "ptu-model"}, + ) + + assert _is_model_cost_zero(model="ptu-model", llm_router=router) is False + assert _is_model_cost_zero(model="ptu-model-alias", llm_router=router) is False, ( + "An aliased PTU group must not be read as free" + ) + + def test_hidden_model_group_alias_enforces_budget(self): + """A hidden alias keeps budget enforced: get_model_group_info() returns None for it, + so the cost is unknown before the configuration gate is reached.""" + router = Router( + model_list=[ + { + "model_name": "free-model", + "litellm_params": { + "model": "ollama/llama2", + "api_base": "http://localhost:11434", + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + }, + "model_info": {"id": "free-model-id"}, + }, + ], + model_group_alias={"hidden-alias": {"model": "free-model", "hidden": True}}, + ) + + assert _is_model_cost_zero(model="hidden-alias", llm_router=router) is False + + def test_dangling_model_group_alias_enforces_budget(self): + """An alias pointing at a group that does not exist keeps budget enforced.""" + router = Router( + model_list=[ + { + "model_name": "free-model", + "litellm_params": { + "model": "ollama/llama2", + "api_base": "http://localhost:11434", + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + }, + "model_info": {"id": "free-model-id"}, + }, + ], + model_group_alias={"dangling-alias": "model-that-does-not-exist"}, + ) + + assert _is_model_cost_zero(model="dangling-alias", llm_router=router) is False + def test_handles_router_without_zero_cost_cache_attribute(self): """Tolerate router-like objects (e.g. ``MagicMock`` stand-ins) that do not expose ``_zero_cost_cache`` — the auth check must still diff --git a/tests/unit/proxy/test_zero_cost_model_budget_bypass.py b/tests/unit/proxy/test_zero_cost_model_budget_bypass.py index 51a7cb2ee9d..56133f2d35b 100644 --- a/tests/unit/proxy/test_zero_cost_model_budget_bypass.py +++ b/tests/unit/proxy/test_zero_cost_model_budget_bypass.py @@ -588,3 +588,84 @@ class TestEdgeCases: request=MagicMock(), ) assert result is True + + +class TestOverBudgetRequestThroughModelGroupAlias: + """The whole path a request takes, not just the predicate. + + `user_api_key_auth._should_skip_budget_checks()` derives the exemption from the requested + model name and `common_checks()` enforces the budgets with it, so a break anywhere between + alias resolution and enforcement shows up here. See + https://github.com/BerriAI/litellm/issues/35369. + """ + + ROUTE = "/v1/chat/completions" + + @staticmethod + def _router() -> Router: + return Router( + model_list=[ + { + "model_name": "free-model", + "litellm_params": { + "model": "ollama/llama2", + "api_base": "http://localhost:11434", + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + }, + "model_info": {"id": "free-model-id"}, + }, + { + "model_name": "paid-model", + "litellm_params": {"model": "gpt-3.5-turbo", "api_key": "sk-test"}, + "model_info": {"id": "paid-model-id"}, + }, + ], + model_group_alias={"free-model-alias": "free-model", "paid-model-alias": "paid-model"}, + ) + + async def _request(self, model: str, proxy_logging) -> bool: + """Run one over-budget request for `model`, deriving the exemption the way auth does.""" + from litellm.proxy.auth.user_api_key_auth import _should_skip_budget_checks + + router = self._router() + request_data = {"model": model} + skip_budget_checks = _should_skip_budget_checks( + request_data=request_data, route=self.ROUTE, request=None, llm_router=router + ) + return await common_checks( + request_body=request_data, + team_object=None, + user_object=LiteLLM_UserTable(user_id="test-user", spend=100.0, max_budget=50.0), + end_user_object=None, + global_proxy_spend=None, + general_settings={}, + route=self.ROUTE, + llm_router=router, + proxy_logging_obj=proxy_logging, + valid_token=UserAPIKeyAuth(token="test-token", user_id="test-user"), + request=MagicMock(), + skip_budget_checks=skip_budget_checks, + ) + + @pytest.mark.asyncio + async def test_over_budget_request_for_aliased_free_model_is_allowed(self, mock_proxy_logging): + assert await self._request("free-model-alias", mock_proxy_logging) is True + + @pytest.mark.asyncio + async def test_over_budget_request_for_free_model_is_allowed(self, mock_proxy_logging): + """The same deployment under its own name, so the alias is the only difference above.""" + assert await self._request("free-model", mock_proxy_logging) is True + + @pytest.mark.asyncio + async def test_over_budget_request_for_aliased_paid_model_is_blocked(self, mock_proxy_logging): + with pytest.raises(litellm.BudgetExceededError) as exc_info: + await self._request("paid-model-alias", mock_proxy_logging) + + assert exc_info.value.current_cost == 100.0 + assert exc_info.value.max_budget == 50.0 + + @pytest.mark.asyncio + async def test_over_budget_request_for_paid_model_is_blocked(self, mock_proxy_logging): + with pytest.raises(litellm.BudgetExceededError): + await self._request("paid-model", mock_proxy_logging) From 7f95b5f3615b383ae08551156ddfad6af87acfcf Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 02:15:12 -0700 Subject: [PATCH 10/41] refactor: clean up fresh tech debt from 2026-09-28 (#43674) * refactor: clean up fresh tech debt from 2026-09-28 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor: group leaderboard rows in one pass and wrap docstring at 120 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/tool_catalog_guard.py | 16 ++++++------- litellm/proxy/auth/auth_checks.py | 7 +++--- litellm/proxy/db/model_usage_rollup.py | 16 ++++++++----- .../model_insights_endpoints.py | 24 ++++++++++++------- .../rust_bridge/callbacks_legacy_python.py | 13 ++++++---- .../test_model_insights_endpoints.py | 2 +- 6 files changed, 46 insertions(+), 32 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/tool_catalog_guard.py b/litellm/proxy/_experimental/mcp_server/tool_catalog_guard.py index 58227640fba..b3c70a33d0e 100644 --- a/litellm/proxy/_experimental/mcp_server/tool_catalog_guard.py +++ b/litellm/proxy/_experimental/mcp_server/tool_catalog_guard.py @@ -39,10 +39,10 @@ class _ScanKwargs(TypedDict): server_name: ReadOnly[str] mcp_rate_limit_server_name: ReadOnly[str] user_api_key_auth: ReadOnly[UserAPIKeyAuth | None] - user_api_key_user_id: ReadOnly[object] - user_api_key_team_id: ReadOnly[object] - user_api_key_end_user_id: ReadOnly[object] - user_api_key_hash: ReadOnly[object] + user_api_key_user_id: ReadOnly[str | None] + user_api_key_team_id: ReadOnly[str | None] + user_api_key_end_user_id: ReadOnly[str | None] + user_api_key_hash: ReadOnly[str | None] headers: ReadOnly[Mapping[str, str]] mcp_tool_description: ReadOnly[str] mcp_input_schema: ReadOnly[Mapping[str, object]] @@ -208,10 +208,10 @@ async def _guarded_catalog_entry( "server_name": server.name, "mcp_rate_limit_server_name": server.alias or server.server_name or server.name, "user_api_key_auth": user_api_key_auth, - "user_api_key_user_id": getattr(user_api_key_auth, "user_id", None), - "user_api_key_team_id": getattr(user_api_key_auth, "team_id", None), - "user_api_key_end_user_id": getattr(user_api_key_auth, "end_user_id", None), - "user_api_key_hash": getattr(user_api_key_auth, "api_key", None), + "user_api_key_user_id": user_api_key_auth.user_id if user_api_key_auth else None, + "user_api_key_team_id": user_api_key_auth.team_id if user_api_key_auth else None, + "user_api_key_end_user_id": user_api_key_auth.end_user_id if user_api_key_auth else None, + "user_api_key_hash": user_api_key_auth.api_key if user_api_key_auth else None, "headers": logging_safe_mcp_headers(raw_headers), "mcp_tool_description": tool.description or "", "mcp_input_schema": tool.input_schema, diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 51cc70c010b..8fbeaf18460 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -591,10 +591,9 @@ def _is_cost_explicitly_configured(model: str, llm_router: "Router") -> bool: The group is resolved through ``Router.get_model_list()``, the same resolution ``get_model_group_info()`` applies when the caller reads the cost a few lines earlier, so the - two lookups cannot disagree: names defined in ``Router.model_group_alias`` are not - ``model_name`` entries in ``Router.model_list``, and scanning that list by exact name reported - every aliased group as unconfigured. It also reaches a deployment that prices itself through - its ``model_info`` block, whose entry lands in the cost map under the deployment id. + two lookups cannot disagree, including for names defined in ``Router.model_group_alias``. + It also reaches a deployment that prices itself through its ``model_info`` block, whose entry + lands in the cost map under the deployment id. """ for deployment in llm_router.get_model_list(model_name=model) or (): model_id = (deployment.get("model_info") or _EMPTY_COST_ENTRY).get("id") diff --git a/litellm/proxy/db/model_usage_rollup.py b/litellm/proxy/db/model_usage_rollup.py index 808c3528051..acd9130da30 100644 --- a/litellm/proxy/db/model_usage_rollup.py +++ b/litellm/proxy/db/model_usage_rollup.py @@ -22,12 +22,16 @@ def model_usage_task_type(request_tags: str) -> str: tags: Final = _TAGS.validate_json(request_tags) except ValidationError: return MODEL_INSIGHTS_DEFAULT_TASK - for tag in tags: - if isinstance(tag, str) and tag.startswith(MODEL_INSIGHTS_TASK_TAG_PREFIX): - task = tag.removeprefix(MODEL_INSIGHTS_TASK_TAG_PREFIX) - if task in load_model_insight_tasks(): - return task - return MODEL_INSIGHTS_DEFAULT_TASK + return next( + ( + task + for tag in tags + if isinstance(tag, str) + and tag.startswith(MODEL_INSIGHTS_TASK_TAG_PREFIX) + and (task := tag.removeprefix(MODEL_INSIGHTS_TASK_TAG_PREFIX)) in load_model_insight_tasks() + ), + MODEL_INSIGHTS_DEFAULT_TASK, + ) def _is_internal_call(metadata: str) -> bool: diff --git a/litellm/proxy/management_endpoints/model_insights_endpoints.py b/litellm/proxy/management_endpoints/model_insights_endpoints.py index dbdaa59d7d4..0c6c7d1227d 100644 --- a/litellm/proxy/management_endpoints/model_insights_endpoints.py +++ b/litellm/proxy/management_endpoints/model_insights_endpoints.py @@ -1,3 +1,5 @@ +import functools +import itertools from collections.abc import Mapping from datetime import date, datetime, timedelta, timezone from typing import Annotated, Final @@ -111,14 +113,20 @@ def _daily_metric(row: _GroupedDaily) -> ModelInsightDailyMetric: def _summarize_tasks(rows: list[_GroupedTask], metric: ModelInsightsMetric) -> list[ModelInsightTaskSummary]: catalog: Final = load_model_insight_tasks() - totals: Final[dict[str, float]] = {} - leaders: Final[dict[str, _GroupedTask]] = {} - for row in rows: - value = _rank_value(row, metric) - totals[row.task_type] = totals.get(row.task_type, 0.0) + value - leader = leaders.get(row.task_type) - if leader is None or value > _rank_value(leader, metric): - leaders[row.task_type] = row + first_seen: Final = {task: index for index, task in enumerate(dict.fromkeys(row.task_type for row in rows))} + by_task: Final = { + task: tuple(group) + for task, group in itertools.groupby( + sorted(rows, key=lambda row: first_seen[row.task_type]), key=lambda row: row.task_type + ) + } + totals: Final = { + task: functools.reduce(lambda total, row: total + _rank_value(row, metric), task_rows, 0.0) + for task, task_rows in by_task.items() + } + leaders: Final = { + task: max(task_rows, key=lambda row: _rank_value(row, metric)) for task, task_rows in by_task.items() + } grand: Final = sum(totals.values()) return [ ModelInsightTaskSummary( diff --git a/litellm/rust_bridge/callbacks_legacy_python.py b/litellm/rust_bridge/callbacks_legacy_python.py index 4f1f2b3c9fd..25513666c43 100644 --- a/litellm/rust_bridge/callbacks_legacy_python.py +++ b/litellm/rust_bridge/callbacks_legacy_python.py @@ -16,6 +16,7 @@ from dataclasses import dataclass from typing import ( TYPE_CHECKING, Final, + Literal, Protocol, cast, # noqa: TID251 # bounded compatibility calls into legacy Python integrations ) @@ -231,6 +232,10 @@ def defer_success(logger: LoggingSurface, pending: object) -> None: setattr(logger, "_native_pending_logging", pending) +def _cache_hit(logger: LoggingSurface) -> Literal[True] | None: + return True if logger.model_call_details.get("cache_hit") is True else None + + def sync_success_for_async_call( logger: LoggingSurface, response: object, start: datetime.datetime, end: datetime.datetime ) -> None: @@ -238,7 +243,7 @@ def sync_success_for_async_call( result=response, start_time=start, end_time=end, - cache_hit=True if logger.model_call_details.get("cache_hit") is True else None, + cache_hit=_cache_hit(logger), ) @@ -265,16 +270,14 @@ def submit_success(logger: LoggingSurface, response: object, start: datetime.dat response, start, end, - cache_hit=True if logger.model_call_details.get("cache_hit") is True else None, + cache_hit=_cache_hit(logger), ) def async_success_handler( logger: LoggingSurface, response: object, start: datetime.datetime, end: datetime.datetime ) -> Coroutine[object, object, None]: - return logger.async_success_handler( - response, start, end, cache_hit=True if logger.model_call_details.get("cache_hit") is True else None - ) + return logger.async_success_handler(response, start, end, cache_hit=_cache_hit(logger)) def enqueue_logging(coroutine: Coroutine[object, object, None]) -> None: diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_insights_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_insights_endpoints.py index af58c2d9884..2cb66771e72 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_insights_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_insights_endpoints.py @@ -122,7 +122,7 @@ def test_model_insights_scopes_daily_to_ranked_deployments() -> None: def _task_rows() -> list[dict[str, object]]: def row(task: str, group: str, requests: str, spend: float) -> dict[str, object]: base = _grouped_row(task_type=task, model_group=group, model=group, custom_llm_provider="openai") - base["_sum"].update({"request_count": requests, "spend": spend}) # type: ignore[union-attr] + base["_sum"].update({"request_count": requests, "spend": spend}) return base return [ From 9dfa42dcde44a54bb5ec2b99cfbb0d469e8b91c8 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 06:12:58 -0700 Subject: [PATCH 11/41] refactor(types): replace Any with proven types in 7 files (#43704) * refactor(types): replace Any with proven types in 11 files Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(types): revert Any changes that broke existing callers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(types): drop prompt factory helper wrappers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): cover typing sweep surfaces Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): tighten sweep audit tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/litellm_core_utils/litellm_logging.py | 12 +- .../litellm_core_utils/streaming_handler.py | 2 +- litellm/llms/custom_httpx/llm_http_handler.py | 10 +- .../vertex_and_google_ai_studio_gemini.py | 11 +- .../mcp_server/mcp_server_manager.py | 3 +- .../guardrail_hooks/xecguard/xecguard.py | 2 +- litellm/utils.py | 7 +- .../mcp/test_mcp_tool_permission_merge.py | 50 ++++ .../observability/test_xecguard_wire.py | 84 +++++++ .../test_image_gen_drop_params_wire.py | 49 ++++ .../test_openai_stream_text_usage_wire.py | 81 +++++++ .../test_vertex_gemini_function_call_wire.py | 229 ++++++++++++++++++ .../spend/test_chaos_burst_spend_once.py | 56 +++++ 13 files changed, 574 insertions(+), 22 deletions(-) create mode 100644 tests/integration/mcp/test_mcp_tool_permission_merge.py create mode 100644 tests/integration/observability/test_xecguard_wire.py create mode 100644 tests/integration/providers/test_image_gen_drop_params_wire.py create mode 100644 tests/integration/providers/test_openai_stream_text_usage_wire.py create mode 100644 tests/integration/providers/test_vertex_gemini_function_call_wire.py create mode 100644 tests/integration/spend/test_chaos_burst_spend_once.py diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 5af669591f7..06cbfd4fc04 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -2432,7 +2432,7 @@ class Logging(LiteLLMLoggingBaseClass): await invalidate_baseline_cache(self, reason, completed=completed) def _build_standard_logging_payload( - self, init_response_obj: object, start_time: Any, end_time: Any + self, init_response_obj: object, start_time: dt_object, end_time: dt_object ) -> StandardLoggingPayload | None: """Build StandardLoggingPayload and accumulate its construction time.""" _start: Final = time.time() @@ -2732,7 +2732,7 @@ class Logging(LiteLLMLoggingBaseClass): def success_handler( self, - result: Any = None, # heterogeneous response object; varies by call type (ANN401 ignored, see ruff-strict.toml) + result: object = None, # heterogeneous response object; varies by call type (ANN401 ignored, see ruff-strict.toml) start_time: datetime.datetime | None = None, end_time: datetime.datetime | None = None, cache_hit: bool | None = None, @@ -3171,7 +3171,7 @@ class Logging(LiteLLMLoggingBaseClass): async def async_success_handler( self, - result: Any = None, # heterogeneous response object; varies by call type (ANN401 ignored, see ruff-strict.toml) + result: object = None, # heterogeneous response object; varies by call type (ANN401 ignored, see ruff-strict.toml) start_time: datetime.datetime | None = None, end_time: datetime.datetime | None = None, cache_hit: bool | None = None, @@ -3189,7 +3189,7 @@ class Logging(LiteLLMLoggingBaseClass): async def _async_success_handler_body( self, - result: Any = None, # heterogeneous response object; varies by call type (ANN401 ignored, see ruff-strict.toml) + result: object = None, # heterogeneous response object; varies by call type (ANN401 ignored, see ruff-strict.toml) start_time: datetime.datetime | None = None, end_time: datetime.datetime | None = None, cache_hit: bool | None = None, @@ -4296,7 +4296,7 @@ class Logging(LiteLLMLoggingBaseClass): ) return result - def _handle_a2a_response_logging(self, result: Any) -> Any: + def _handle_a2a_response_logging(self, result: Any) -> object: """ Handles logging for A2A (Agent-to-Agent) responses. @@ -5705,7 +5705,7 @@ class StandardLoggingPayloadSetup: @staticmethod def get_standard_logging_metadata( - metadata: dict[str, Any] | None, + metadata: Mapping[str, object] | None, litellm_params: dict | None = None, prompt_integration: str | None = None, applied_guardrails: list[str] | None = None, diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index fa4650aec4f..946d19c028f 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -816,7 +816,7 @@ class CustomStreamWrapper: self, completion_obj: dict[str, Any], model_response: ModelResponseStream, - response_obj: dict[str, Any], + response_obj: Mapping[str, object], ) -> bool: if ( "content" in completion_obj diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 60d12337447..8d65aa7b0ca 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -201,6 +201,7 @@ if TYPE_CHECKING: from aiohttp import ClientSession from websockets.asyncio.client import ClientConnection + from litellm.google_genai.streaming_iterator import AsyncGoogleGenAIGenerateContentStreamingIterator from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj from litellm.litellm_core_utils.tokenizer import Encoding as Tokenizer @@ -209,6 +210,7 @@ if TYPE_CHECKING: ) from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig from litellm.proxy._types import UserAPIKeyAuth + from litellm.types.google_genai.main import GenerateContentResponse from litellm.types.llms.openai_evals import ( CancelEvalResponse, CancelRunResponse, @@ -401,7 +403,8 @@ def _decoded_body_headers(response: httpx.Response) -> httpx.Headers: `aiter_bytes` yields the decoded body, so the upstream transfer headers only describe the bytes on the wire when no content-encoding was applied. """ - if response.headers.get("content-encoding", "identity").lower() == "identity": + headers: Final[Mapping[str, str]] = response.headers + if headers.get("content-encoding", "identity").lower() == "identity": return response.headers return httpx.Headers( [ @@ -3291,7 +3294,8 @@ class BaseLLMHTTPHandler: """ if upload_url_location == "headers": # Google Cloud Storage style - URL in X-Goog-Upload-URL header - upload_url = response.headers.get("X-Goog-Upload-URL") + upload_headers: Final[Mapping[str, str]] = response.headers + upload_url = upload_headers.get("X-Goog-Upload-URL") return upload_url, None else: # Response body style (e.g., Manus, S3 presigned URLs) @@ -11594,7 +11598,7 @@ class BaseLLMHTTPHandler: stream: bool = False, litellm_metadata: dict[str, object] | None = None, system_instruction: object | None = None, - ) -> Any: + ) -> "AsyncGoogleGenAIGenerateContentStreamingIterator | GenerateContentResponse": """ Async version of the generate content handler. Uses async HTTP client to make requests. diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index 941ec4ad419..b5f32d57061 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -1571,12 +1571,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): gemini_call_id = part["functionCall"].get("id") if is_function_call is True: - function_dict: dict[str, Any] = dict(_function_chunk) - if thought_signature: - if "provider_specific_fields" not in function_dict: - function_dict["provider_specific_fields"] = {} - function_dict["provider_specific_fields"]["thought_signature"] = thought_signature - function = cast(ChatCompletionToolCallFunctionChunk, function_dict) + function = ( + {**_function_chunk, "provider_specific_fields": {"thought_signature": thought_signature}} + if thought_signature + else {**_function_chunk} + ) else: _tool_response_chunk: ChatCompletionToolCallChunk = { "id": f"call_{uuid.uuid4().hex[:28]}", diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 3e3b53f387d..490d3072955 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -28,7 +28,6 @@ from contextlib import asynccontextmanager from dataclasses import dataclass, replace from functools import lru_cache from itertools import chain, groupby -from operator import itemgetter from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Generic, Literal, TypeAlias, TypedDict, TypeVar, cast from urllib.parse import ParseResult, urlparse @@ -6988,7 +6987,7 @@ class MCPServerManager: ) return { server_id: list(dict.fromkeys(chain.from_iterable(tools for _, tools in group))) - for server_id, group in groupby(sorted(expanded, key=itemgetter(0)), key=itemgetter(0)) + for server_id, group in groupby(sorted(expanded, key=lambda pair: pair[0]), key=lambda pair: pair[0]) } def get_mcp_server_by_name(self, server_name: str, client_ip: str | None = None) -> MCPServer | None: diff --git a/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py b/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py index f4330ad6aa9..ddf9cace8b9 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py +++ b/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py @@ -474,7 +474,7 @@ class XecGuardGuardrail(CustomGuardrail): return "\n".join(text_parts) or None @staticmethod - def _extract_choice_content(choice: Any) -> Any: + def _extract_choice_content(choice: Any) -> object: if hasattr(choice, "message"): message = choice.message elif isinstance(choice, dict): diff --git a/litellm/utils.py b/litellm/utils.py index 13a46840431..09b5067339d 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -3595,7 +3595,7 @@ def get_optional_params_image_gen( passed_params.pop("provider_config", None) passed_params.pop("drop_params", None) drop_params = normalize_drop_params(drop_params) - additional_drop_params = passed_params.pop("additional_drop_params", None) + passed_params.pop("additional_drop_params", None) passed_params.pop("kwargs") special_params: Final[Mapping[str, object]] = kwargs for k, v in special_params.items(): @@ -4434,11 +4434,12 @@ def get_optional_params( store: bool | None = None, prompt_cache_key: str | None = None, base_model: str | None = None, - **kwargs, + **kwargs: object, ): drop_params = normalize_drop_params(drop_params) # rebind-ok: config and DB deployments pass "true" as a string passed_params: Final = locals().copy() - special_params: Final = passed_params.pop("kwargs") + passed_params.pop("kwargs") + special_params: Final = kwargs # Remove base_model from passed_params so it doesn't interfere with # non_default_params / _check_valid_arg — it's a routing hint, not an # OpenAI param. diff --git a/tests/integration/mcp/test_mcp_tool_permission_merge.py b/tests/integration/mcp/test_mcp_tool_permission_merge.py new file mode 100644 index 00000000000..40e57a6cd21 --- /dev/null +++ b/tests/integration/mcp/test_mcp_tool_permission_merge.py @@ -0,0 +1,50 @@ +import uuid +from typing import Final + +from integration._support.client import Gateway, eventually +from integration._support.mcp import ( + call_tool, + mcp_peer, + register_mcp, + tool_names, +) + + +def test_tool_permissions_merge_when_keys_resolve_to_same_server(gateway: Gateway) -> None: + with mcp_peer() as first, mcp_peer() as second, gateway.scenario() as scenario: + shared_alias: Final = "merge" + uuid.uuid4().hex[:8] + other_alias: Final = "other" + uuid.uuid4().hex[:8] + first_id: Final = register_mcp(scenario, first, shared_alias) + second_id: Final = register_mcp(scenario, second, other_alias) + key: Final = scenario.key( + object_permission={ + "mcp_servers": [first_id, second_id], + "mcp_tool_permissions": { + shared_alias: ["add"], + first_id: ["multiply", "add"], + second_id: ["add"], + }, + } + ) + + first_names: Final = eventually( + lambda: tool_names(gateway, key, first_id), + lambda names: set(names) != set(), + seconds=15, + ) + assert set(first_names) == {"add", "multiply"}, first_names + assert set(tool_names(gateway, key, second_id)) == {"add"} + + first.drain() + add: Final = call_tool(gateway, key, first_id, first_names["add"], {"a": 1, "b": 2}) + assert add.status_code == 200 and add.json()["isError"] is False, add.text + assert add.json()["content"][0]["text"] == "3" + multiply: Final = call_tool(gateway, key, first_id, first_names["multiply"], {"a": 2, "b": 3}) + assert multiply.status_code == 200 and multiply.json()["isError"] is False, multiply.text + assert multiply.json()["content"][0]["text"] == "6" + fail_name: Final = f"{shared_alias}-fail" + denied: Final = call_tool(gateway, key, first_id, fail_name, {}) + assert denied.status_code == 403, denied.text + detail: Final = denied.json()["detail"]["error"] + assert "is not allowed for your key/team" in detail and "fail" in detail, detail + assert len(tuple(item for item in first.drain() if item["body"].get("method") == "tools/call")) == 2 diff --git a/tests/integration/observability/test_xecguard_wire.py b/tests/integration/observability/test_xecguard_wire.py new file mode 100644 index 00000000000..df80ca83577 --- /dev/null +++ b/tests/integration/observability/test_xecguard_wire.py @@ -0,0 +1,84 @@ +import json +import uuid +from pathlib import Path +from typing import Final + +import yaml +from integration._support.client import Gateway +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) + + +def test_xecguard_post_call_scan_reaches_vendor_and_call_succeeds(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "xecguard" + uuid.uuid4().hex + + def vendor(request: Request) -> Reply: + assert request.method == "POST" + assert request.target == "/xecguard/v1/scan" + assert request.headers["authorization"] == "Bearer synthetic-xecguard-key" + body: Final = _JSON_OBJECT.validate_json(request.body) + assert body["model"] == "xecguard_v2" + assert body["scan_type"] in ("input", "response") + assert any(message.get("content") == "hi" for message in body.get("messages", [])), body + return Reply(body=json.dumps({"decision": "SAFE", "violations": []}).encode()) + + def provider(request: Request) -> Reply: + assert request.target == "/chat/completions" + return Reply( + body=json.dumps( + { + "id": "chatcmpl-xec", + "object": "chat.completion", + "created": 1700000000, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "permitted"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8}, + } + ).encode() + ) + + with wire_server(vendor) as policy, wire_server(provider) as upstream: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["guardrails"] = [ + { + "guardrail_name": identity, + "litellm_params": { + "guardrail": "xecguard", + "mode": "post_call", + "default_on": True, + "api_base": policy.url, + "api_key": "synthetic-xecguard-key", + }, + } + ] + path: Final = tmp_path / "xecguard.yaml" + path.write_text(yaml.safe_dump(config)) + with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=upstream.url, + api_key="synthetic-openai-key", + ) + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "max_tokens": 16, + "messages": [{"role": "user", "content": "hi"}], + }, + ) + assert response.status_code == 200, response.text + assert response.json()["choices"][0]["message"]["content"] == "permitted" + scans: Final = tuple(request for request in policy.drain() if request.target == "/xecguard/v1/scan") + assert scans, "post-call xecguard scan never reached the vendor" + assert len(upstream.drain()) == 1 diff --git a/tests/integration/providers/test_image_gen_drop_params_wire.py b/tests/integration/providers/test_image_gen_drop_params_wire.py new file mode 100644 index 00000000000..7addfefeb36 --- /dev/null +++ b/tests/integration/providers/test_image_gen_drop_params_wire.py @@ -0,0 +1,49 @@ +import json +from typing import Final + +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) + + +def test_image_generation_additional_drop_params_reaches_provider_body(gateway: Gateway) -> None: + def respond(request: Request) -> Reply: + assert request.method == "POST" + assert request.target == "/images/generations" + body: Final = _JSON_OBJECT.validate_json(request.body) + assert "style" not in body, body + assert body["model"] == "dall-e-3" + assert body["prompt"] == "a scripted cat" + assert body["size"] == "1024x1024" + return Reply( + body=json.dumps( + { + "created": 1700000000, + "data": [{"b64_json": "aW1n", "revised_prompt": None, "url": None}], + } + ).encode() + ) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/dall-e-3", + api_base=wire.url, + api_key="synthetic-image-key", + additional_drop_params=["style"], + ) + response: Final = gateway.client.post( + "/v1/images/generations", + json={ + "model": model, + "prompt": "a scripted cat", + "size": "1024x1024", + "style": "vivid", + }, + headers={"Authorization": f"Bearer {gateway.key}"}, + timeout=30, + ) + assert response.status_code == 200, response.text + assert response.json()["data"][0]["b64_json"] == "aW1n" + assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/images/generations")] diff --git a/tests/integration/providers/test_openai_stream_text_usage_wire.py b/tests/integration/providers/test_openai_stream_text_usage_wire.py new file mode 100644 index 00000000000..735d1a904bb --- /dev/null +++ b/tests/integration/providers/test_openai_stream_text_usage_wire.py @@ -0,0 +1,81 @@ +import json +from collections.abc import Mapping +from typing import Final + +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) + +_IDENTITY: Final = "chatcmpl-stream-usage" + + +def _frame(delta: Mapping[str, JsonValue], finish: str | None = None) -> bytes: + return ( + b"data: " + + json.dumps( + { + "id": _IDENTITY, + "object": "chat.completion.chunk", + "created": 1, + "model": "gpt-4o-mini", + "choices": [{"index": 0, "delta": delta, "finish_reason": finish}], + } + ).encode() + + b"\n\n" + ) + + +def test_streaming_chat_assembles_text_and_final_usage(gateway: Gateway) -> None: + def respond(request: Request) -> Reply: + assert request.target == "/chat/completions" + body: Final = _JSON_OBJECT.validate_json(request.body) + assert body["stream"] is True, body + assert body["stream_options"]["include_usage"] is True, body + usage: Final = json.dumps( + { + "id": _IDENTITY, + "object": "chat.completion.chunk", + "created": 1, + "model": "gpt-4o-mini", + "choices": [], + "usage": {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15}, + } + ) + return Reply( + content_type="text/event-stream", + chunks=[ + _frame({"role": "assistant", "content": "Hello "}), + _frame({"content": "world"}), + _frame({}, finish="stop"), + b"data: " + usage.encode() + b"\n\n", + b"data: [DONE]\n\n", + ], + ) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(api_base=wire.url) + response: Final = gateway.client.post( + "/chat/completions", + json={ + "model": model, + "stream": True, + "stream_options": {"include_usage": True}, + "messages": [{"role": "user", "content": "hi"}], + }, + headers={"Authorization": f"Bearer {gateway.key}"}, + timeout=30, + ) + assert response.status_code == 200, response.text + chunks: Final = tuple( + json.loads(line[6:]) + for line in response.text.splitlines() + if line.startswith("data: ") and line != "data: [DONE]" + ) + text: Final = "".join(choice["delta"].get("content", "") for chunk in chunks for choice in chunk["choices"]) + assert text == "Hello world" + usages: Final = tuple(chunk["usage"] for chunk in chunks if chunk.get("usage")) + assert len(usages) == 1 + assert usages[0]["prompt_tokens"] == 11 and usages[0]["completion_tokens"] == 4 + assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/chat/completions")] diff --git a/tests/integration/providers/test_vertex_gemini_function_call_wire.py b/tests/integration/providers/test_vertex_gemini_function_call_wire.py new file mode 100644 index 00000000000..fef5e31f9c7 --- /dev/null +++ b/tests/integration/providers/test_vertex_gemini_function_call_wire.py @@ -0,0 +1,229 @@ +import json +from typing import Final + +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import rsa +from integration._support.client import Gateway, Scenario +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +_BACKEND: Final = "gemini-3.7-flash" +_PROJECT: Final = "scripted-project" +_LOCATION: Final = "us-central1" +_MODEL_PATH: Final = f"/v1/projects/{_PROJECT}/locations/{_LOCATION}/publishers/google/models/{_BACKEND}" +_SIGNATURE: Final = "sig-4f2a" +_ARGS: Final = {"city": "Paris"} +_FUNCTIONS: Final = [ + { + "name": "get_weather", + "description": "Return the weather for a city", + "parameters": { + "type": "object", + "properties": {"city": {"type": "string"}}, + "required": ["city"], + }, + } +] +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) + + +def _service_account_json(token_url: str) -> str: + private_key: Final = ( + rsa.generate_private_key(public_exponent=65537, key_size=2048) + .private_bytes( + serialization.Encoding.PEM, + serialization.PrivateFormat.PKCS8, + serialization.NoEncryption(), + ) + .decode() + ) + return json.dumps( + { + "type": "service_account", + "project_id": _PROJECT, + "private_key_id": "scripted", + "private_key": private_key, + "client_email": f"scripted@{_PROJECT}.iam.gserviceaccount.com", + "client_id": "0", + "auth_uri": f"{token_url}/_oauth/authorize", + "token_uri": f"{token_url}/_oauth/token", + } + ) + + +def _candidate(*, with_signature: bool) -> dict[str, JsonValue]: + part: Final = { + "functionCall": {"name": "get_weather", "args": _ARGS, "id": "fc-1"}, + **({"thoughtSignature": _SIGNATURE} if with_signature else {}), + } + return { + "candidates": [ + { + "content": {"role": "model", "parts": [part]}, + "finishReason": "STOP", + } + ], + "usageMetadata": {"promptTokenCount": 11, "candidatesTokenCount": 7, "totalTokenCount": 18}, + "modelVersion": _BACKEND, + } + + +def _model(gateway: Gateway, scenario: Scenario, wire_url: str) -> str: + return scenario.model( + model=f"vertex_ai/{_BACKEND}", + api_base=f"{wire_url}{_MODEL_PATH}", + api_key=None, + vertex_project=_PROJECT, + vertex_location=_LOCATION, + vertex_credentials=_service_account_json(gateway.upstream_url.rstrip("/")), + ) + + +def _non_streaming_call(gateway: Gateway, model: str) -> dict[str, JsonValue]: + response: Final = gateway.client.post( + "/v1/chat/completions", + json={ + "model": model, + "functions": _FUNCTIONS, + "messages": [{"role": "user", "content": "weather?"}], + }, + headers={"Authorization": f"Bearer {gateway.key}"}, + timeout=30, + ) + assert response.status_code == 200, response.text + return response.json() + + +def _streaming_call(gateway: Gateway, model: str) -> tuple[dict[str, JsonValue], ...]: + with gateway.client.stream( + "POST", + "/v1/chat/completions", + json={ + "model": model, + "functions": _FUNCTIONS, + "messages": [{"role": "user", "content": "weather?"}], + "stream": True, + }, + headers={"Authorization": f"Bearer {gateway.key}"}, + timeout=30, + ) as response: + assert response.status_code == 200, response.read() + lines: Final = tuple(line for line in response.iter_lines() if line.startswith("data: ")) + assert lines[-1] == "data: [DONE]", lines[-3:] + return tuple(_JSON_OBJECT.validate_json(line.removeprefix("data: ").encode()) for line in lines[:-1]) + + +def _function_call_of(response: dict[str, JsonValue]) -> dict[str, JsonValue]: + message: Final = response["choices"][0]["message"] + assert isinstance(message, dict) + call: Final = message["function_call"] + assert isinstance(call, dict) + return call + + +def test_vertex_gemini_function_call_thought_signature_is_returned_non_streaming(gateway: Gateway) -> None: + def respond(request: Request) -> Reply: + assert request.target == f"{_MODEL_PATH}:generateContent" + return Reply(body=json.dumps(_candidate(with_signature=True)).encode()) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = _model(gateway, scenario, wire.url) + call: Final = _function_call_of(_non_streaming_call(gateway, model)) + assert call["name"] == "get_weather" + assert json.loads(str(call["arguments"])) == _ARGS + assert call.get("provider_specific_fields") == {"thought_signature": _SIGNATURE} + + +def test_vertex_gemini_function_call_without_signature_has_no_provider_fields_non_streaming( + gateway: Gateway, +) -> None: + def respond(request: Request) -> Reply: + assert request.target == f"{_MODEL_PATH}:generateContent" + return Reply(body=json.dumps(_candidate(with_signature=False)).encode()) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = _model(gateway, scenario, wire.url) + call: Final = _function_call_of(_non_streaming_call(gateway, model)) + assert call["name"] == "get_weather" + assert json.loads(str(call["arguments"])) == _ARGS + assert "provider_specific_fields" not in call + assert "thought_signature" not in json.dumps(call) + + +def test_vertex_gemini_function_call_thought_signature_is_returned_streaming(gateway: Gateway) -> None: + def respond(request: Request) -> Reply: + assert request.target == f"{_MODEL_PATH}:streamGenerateContent?alt=sse" + payload: Final = json.dumps(_candidate(with_signature=True)) + return Reply(content_type="text/event-stream", chunks=[f"data: {payload}\n\n".encode()]) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = _model(gateway, scenario, wire.url) + chunks: Final = _streaming_call(gateway, model) + function_calls: Final = tuple( + choice["delta"]["function_call"] + for chunk in chunks + for choice in chunk.get("choices", ()) + if choice.get("delta", {}).get("function_call") + ) + assert function_calls, "no function_call delta received" + merged: Final = "".join(str(call.get("arguments", "")) for call in function_calls) + assert json.loads(merged) == _ARGS + assert function_calls[-1].get("provider_specific_fields") == {"thought_signature": _SIGNATURE} + + +def test_vertex_gemini_function_call_without_signature_has_no_provider_fields_streaming( + gateway: Gateway, +) -> None: + def respond(request: Request) -> Reply: + assert request.target == f"{_MODEL_PATH}:streamGenerateContent?alt=sse" + payload: Final = json.dumps(_candidate(with_signature=False)) + return Reply(content_type="text/event-stream", chunks=[f"data: {payload}\n\n".encode()]) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = _model(gateway, scenario, wire.url) + chunks: Final = _streaming_call(gateway, model) + function_calls: Final = tuple( + choice["delta"]["function_call"] + for chunk in chunks + for choice in chunk.get("choices", ()) + if choice.get("delta", {}).get("function_call") + ) + assert function_calls, "no function_call delta received" + assert all("provider_specific_fields" not in call for call in function_calls) + assert "thought_signature" not in json.dumps(function_calls) + + +def test_vertex_gemini_kwargs_extra_param_reaches_generation_config(gateway: Gateway) -> None: + def respond(request: Request) -> Reply: + assert request.target == f"{_MODEL_PATH}:generateContent" + body: Final = _JSON_OBJECT.validate_json(request.body) + assert body["generationConfig"]["top_k"] == 3, body + return Reply( + body=json.dumps( + { + "candidates": [ + { + "content": {"role": "model", "parts": [{"text": "done"}]}, + "finishReason": "STOP", + } + ], + "usageMetadata": {"promptTokenCount": 4, "candidatesTokenCount": 2, "totalTokenCount": 6}, + "modelVersion": _BACKEND, + } + ).encode() + ) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = _model(gateway, scenario, wire.url) + response: Final = gateway.client.post( + "/v1/chat/completions", + json={ + "model": model, + "top_k": 3, + "messages": [{"role": "user", "content": "hi"}], + }, + headers={"Authorization": f"Bearer {gateway.key}"}, + timeout=30, + ) + assert response.status_code == 200, response.text + assert response.json()["choices"][0]["message"]["content"] == "done" diff --git a/tests/integration/spend/test_chaos_burst_spend_once.py b/tests/integration/spend/test_chaos_burst_spend_once.py new file mode 100644 index 00000000000..77b08b1d559 --- /dev/null +++ b/tests/integration/spend/test_chaos_burst_spend_once.py @@ -0,0 +1,56 @@ +import uuid +from concurrent.futures import ThreadPoolExecutor +from typing import Final + +import httpx +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows + +_BURST: Final = 24 + + +def test_burst_with_partial_upstream_failures_logs_each_success_once(gateway: Gateway) -> None: + with ( + httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream, + gateway.scenario() as scenario, + ): + provider_model: Final = f"burst-{uuid.uuid4().hex}" + model: Final = scenario.model(model=f"openai/{provider_model}", input_cost_per_token=0, output_cost_per_token=0) + statuses: Final = [500] + [200, 200, 200] * (_BURST // 4 + 2) + + def remove_script() -> None: + response: Final = upstream.delete(f"/__scripts/{provider_model}") + assert response.status_code in (200, 404), response.text + + scenario.cleanups.callback(remove_script) + configured: Final = upstream.post(f"/__scripts/{provider_model}", json={"statuses": statuses}) + assert configured.status_code == 200, configured.text + upstream.get("/__observations").raise_for_status() + + def attempt(index: int) -> httpx.Response: + return gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": f"burst {index}"}]}, + ) + + with ThreadPoolExecutor(max_workers=_BURST) as pool: + responses: Final = tuple(pool.map(attempt, range(_BURST))) + + succeeded: Final = tuple(response.json()["id"] for response in responses if response.status_code == 200) + assert len(succeeded) > 0, [response.status_code for response in responses] + assert len(set(succeeded)) == len(succeeded), "duplicate response id in burst" + assert all(response.status_code in (200, 429, 500) for response in responses), [ + response.status_code for response in responses + ] + + rows: Final = eventually( + lambda: read_rows( + 'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id = ANY(%s)', + (list(succeeded),), + ), + lambda values: len(values) == len(succeeded), + seconds=90, + ) + landed: Final = [row["request_id"] for row in rows] + assert sorted(landed) == sorted(succeeded), "a successful burst id did not land exactly once" From 684a1edd44efa3a7c7f0395ccfa1bf9803017ea2 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 13:15:37 +0000 Subject: [PATCH 12/41] docs(security): point readers to the security announcements mailing list signup (#43713) * docs(security): point readers to the security announcements mailing list signup Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * docs(security): formalize the security announcements wording Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * docs(security): tighten the best-effort sentence Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: oliver Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- security.md | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/security.md b/security.md index cb5eda7ee22..c73379a9c01 100644 --- a/security.md +++ b/security.md @@ -1,5 +1,10 @@ # Data Privacy and Security +## Security Announcements + +LiteLLM maintains a security announcements mailing list that is open to anyone. Subscribers receive advance notice, typically one to two days, before we release a fix for a particularly severe vulnerability or for any vulnerability exploitable by an unauthenticated attacker. This notice is provided on a best-effort basis + +To subscribe, visit [https://berriai.github.io/security-announce-signup/](https://berriai.github.io/security-announce-signup/) ## Security Vulnerability Reporting Guidelines From 66db132627fc89f2d92f41202e3c59a237b4030e Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 09:19:33 -0700 Subject: [PATCH 13/41] refactor(rust): add shared llms wire type derives (#43730) --- litellm-rust/Cargo.lock | 38 +-- litellm-rust/Cargo.toml | 3 +- .../crates/callbacks-legacy-python/AGENTS.md | 1 + .../crates/callbacks-legacy-python/Cargo.toml | 1 - .../callbacks-legacy-python/src/adapter.rs | 69 ++--- .../crates/callbacks-legacy-python/src/lib.rs | 8 + .../src/test_support.rs | 2 +- litellm-rust/crates/core-utils/Cargo.toml | 3 +- .../crates/core-utils/src/core_helpers.rs | 2 +- .../src/get_provider_specific_headers.rs | 2 +- .../src/prompt_templates/factory.rs | 2 +- .../crates/core-utils/src/serde_compat.rs | 195 -------------- litellm-rust/crates/core/AGENTS.md | 6 +- litellm-rust/crates/core/Cargo.toml | 2 +- .../core/src/chat_completions/handler.rs | 5 +- .../crates/core/src/chat_completions/mod.rs | 2 +- .../core/src/chat_completions/prepare.rs | 2 +- .../crates/core/src/chat_completions/route.rs | 2 +- .../crates/core/src/chat_completions/types.rs | 2 +- .../crates/core/src/messages/AGENTS.md | 2 +- .../crates/core/src/messages/common_utils.rs | 4 +- .../crates/core/src/messages/handler.rs | 24 +- litellm-rust/crates/core/src/messages/mod.rs | 10 +- .../crates/core/src/messages/prepare.rs | 16 +- .../crates/core/src/messages/route.rs | 6 +- .../crates/core/src/messages/types.rs | 26 +- litellm-rust/crates/core/src/ocr/client.rs | 5 +- litellm-rust/crates/core/src/ocr/document.rs | 6 +- litellm-rust/crates/core/src/ocr/handler.rs | 6 +- litellm-rust/crates/core/src/ocr/prepare.rs | 3 +- .../crates/core/src/ocr/provider_config.rs | 4 +- litellm-rust/crates/core/src/ocr/route.rs | 3 +- litellm-rust/crates/core/src/ocr/types.rs | 7 +- litellm-rust/crates/core/src/ocr/wire.rs | 6 +- .../crates/core/src/responses/route.rs | 2 +- .../crates/core/src/responses/types.rs | 2 +- .../crates/core/src/responses/websocket.rs | 2 +- litellm-rust/crates/core/tests/caching.rs | 8 +- .../crates/core/tests/chat_completions.rs | 2 +- .../crates/core/tests/messages/main.rs | 8 +- .../crates/core/tests/messages/request.rs | 6 +- .../crates/core/tests/messages/response.rs | 11 +- .../crates/core/tests/messages/stream.rs | 8 +- litellm-rust/crates/core/tests/ocr/main.rs | 7 +- .../crates/gateway-inference/Cargo.toml | 2 +- .../crates/gateway-inference/src/messages.rs | 2 +- .../crates/gateway-inference/src/ocr.rs | 9 +- .../crates/gateway-inference/tests/ocr.rs | 5 +- .../crates/{types => llms-types}/AGENTS.md | 17 +- .../crates/{types => llms-types}/Cargo.toml | 4 +- .../src/formats}/audio_transcription.rs | 3 +- .../crates/llms-types/src/formats/batches.rs | 36 +++ .../src/formats/chat_completions.rs | 223 ++++++++++++++++ .../src/formats}/messages/AGENTS.md | 0 .../llms-types/src/formats/messages/mod.rs | 11 + .../src/formats/messages/request.rs} | 92 ++++--- .../src/formats/messages/response.rs} | 9 +- .../src/formats}/messages/streaming.rs | 17 +- .../crates/llms-types/src/formats/mod.rs | 6 + .../crates/llms-types/src/formats/ocr.rs | 152 +++++++++++ .../llms-types/src/formats/responses/mod.rs | 4 + .../src/formats/responses/response.rs} | 3 +- .../formats}/responses/streaming_websocket.rs | 10 +- litellm-rust/crates/llms-types/src/headers.rs | 17 ++ litellm-rust/crates/llms-types/src/lib.rs | 11 + .../src/providers}/anthropic.rs | 0 .../crates/llms-types/src/providers/mod.rs | 1 + .../{types => llms-types}/src/recognized.rs | 3 +- .../crates/llms-types/src/serde_compat.rs | 113 ++++++++ .../tests/messages_request.rs} | 2 +- .../tests/messages_streaming.rs | 2 +- litellm-rust/crates/llms-types/tests/ocr.rs | 104 ++++++++ .../crates/llms-types/tests/serde_compat.rs | 86 ++++++ .../crates/llms-types/tests/wire_type.rs | 32 +++ litellm-rust/crates/llms/AGENTS.md | 4 +- litellm-rust/crates/llms/Cargo.toml | 2 +- .../crates/llms/src/anthropic/AGENTS.md | 2 +- .../src/anthropic/batches/transformation.rs | 59 +---- .../crates/llms/src/anthropic/chat/handler.rs | 14 +- .../llms/src/anthropic/chat/transformation.rs | 28 +- .../crates/llms/src/anthropic/common_utils.rs | 71 ++--- .../anthropic/count_tokens/transformation.rs | 18 +- .../llms/src/anthropic/messages/AGENTS.md | 4 +- .../llms/src/anthropic/messages/handler.rs | 22 +- .../llms/src/anthropic/messages/thinking.rs | 83 +++--- .../src/anthropic/messages/transformation.rs | 55 ++-- .../ocr/analyze_transformation.rs | 4 +- .../llms/src/aws_textract/ocr/common_utils.rs | 6 +- .../src/aws_textract/ocr/transformation.rs | 4 +- .../llms/src/azure_ai/messages/AGENTS.md | 2 +- .../src/azure_ai/messages/transformation.rs | 38 +-- .../ocr/cohere_parse_transformation.rs | 6 +- .../document_intelligence/transformation.rs | 15 +- .../llms/src/azure_ai/ocr/transformation.rs | 9 +- .../audio_transcription/transformation.rs | 2 +- .../llms/src/base_llm/chat/streaming.rs | 2 +- .../llms/src/base_llm/chat/transformation.rs | 5 +- .../llms/src/base_llm/messages/AGENTS.md | 2 +- .../src/base_llm/messages/normalization.rs | 13 +- .../llms/src/base_llm/messages/streaming.rs | 6 +- .../src/base_llm/messages/transformation.rs | 26 +- .../crates/llms/src/base_llm/ocr/document.rs | 3 +- .../crates/llms/src/base_llm/ocr/handler.rs | 5 +- .../llms/src/base_llm/ocr/transformation.rs | 250 +----------------- .../src/base_llm/responses/transformation.rs | 5 +- .../src/bedrock/audio_transcription/mod.rs | 2 +- .../bedrock/chat/converse_transformation.rs | 9 +- .../llms/src/bedrock/chat/invoke_handler.rs | 4 +- .../llms/src/bedrock/messages/AGENTS.md | 2 +- .../anthropic_claude3_transformation.rs | 24 +- .../llms/src/cohere/ocr/transformation.rs | 18 +- .../llms/src/mistral/ocr/transformation.rs | 8 +- .../src/openai/responses/transformation.rs | 5 +- .../src/openai_like/chat/transformation.rs | 10 +- .../llms/src/reducto/ocr/transformation.rs | 10 +- .../vertex_ai/ocr/deepseek_transformation.rs | 14 +- .../llms/src/vertex_ai/ocr/transformation.rs | 4 +- .../tests/anthropic_chat_transformation.rs | 2 +- .../tests/bedrock_converse_transformation.rs | 2 +- .../llms/tests/messages_normalization.rs | 4 +- .../tests/openai_like_chat_transformation.rs | 2 +- litellm-rust/crates/model-catalog/Cargo.toml | 4 +- .../crates/model-catalog/src/model_info.rs | 2 +- litellm-rust/crates/python-bridge/Cargo.toml | 2 +- .../crates/python-bridge/src/marshal.rs | 4 +- .../crates/python-bridge/src/routes/AGENTS.md | 2 +- .../src/routes/chat_completions.rs | 6 +- .../python-bridge/src/routes/messages/host.rs | 6 +- .../python-bridge/src/routes/messages/mod.rs | 4 +- .../crates/python-bridge/src/routes/mod.rs | 4 +- .../python-bridge/src/routes/ocr/host.rs | 3 +- .../python-bridge/src/routes/ocr/mod.rs | 12 +- .../python-bridge/src/routes/ocr/project.rs | 2 +- .../python-bridge/src/routes/responses.rs | 4 +- .../crates/token-counter/src/counter.rs | 12 +- .../crates/token-counter/src/types.rs | 4 +- litellm-rust/crates/types/src/lib.rs | 14 - .../types/src/llms/anthropic_messages/mod.rs | 2 - litellm-rust/crates/types/src/llms/mod.rs | 3 - litellm-rust/crates/types/src/llms/openai.rs | 132 --------- litellm-rust/crates/types/src/messages/mod.rs | 1 - .../crates/types/src/responses/mod.rs | 2 - litellm-rust/crates/types/src/utils.rs | 108 -------- 143 files changed, 1404 insertions(+), 1321 deletions(-) rename litellm-rust/crates/{types => llms-types}/AGENTS.md (81%) rename litellm-rust/crates/{types => llms-types}/Cargo.toml (77%) rename litellm-rust/crates/{types/src => llms-types/src/formats}/audio_transcription.rs (72%) create mode 100644 litellm-rust/crates/llms-types/src/formats/batches.rs create mode 100644 litellm-rust/crates/llms-types/src/formats/chat_completions.rs rename litellm-rust/crates/{types/src => llms-types/src/formats}/messages/AGENTS.md (100%) create mode 100644 litellm-rust/crates/llms-types/src/formats/messages/mod.rs rename litellm-rust/crates/{types/src/llms/anthropic_messages/anthropic_request.rs => llms-types/src/formats/messages/request.rs} (89%) rename litellm-rust/crates/{types/src/llms/anthropic_messages/anthropic_response.rs => llms-types/src/formats/messages/response.rs} (92%) rename litellm-rust/crates/{types/src => llms-types/src/formats}/messages/streaming.rs (89%) create mode 100644 litellm-rust/crates/llms-types/src/formats/mod.rs create mode 100644 litellm-rust/crates/llms-types/src/formats/ocr.rs create mode 100644 litellm-rust/crates/llms-types/src/formats/responses/mod.rs rename litellm-rust/crates/{types/src/responses/main.rs => llms-types/src/formats/responses/response.rs} (67%) rename litellm-rust/crates/{types/src => llms-types/src/formats}/responses/streaming_websocket.rs (94%) create mode 100644 litellm-rust/crates/llms-types/src/headers.rs create mode 100644 litellm-rust/crates/llms-types/src/lib.rs rename litellm-rust/crates/{types/src/llms => llms-types/src/providers}/anthropic.rs (100%) create mode 100644 litellm-rust/crates/llms-types/src/providers/mod.rs rename litellm-rust/crates/{types => llms-types}/src/recognized.rs (91%) create mode 100644 litellm-rust/crates/llms-types/src/serde_compat.rs rename litellm-rust/crates/{types/tests/anthropic_request.rs => llms-types/tests/messages_request.rs} (95%) rename litellm-rust/crates/{types => llms-types}/tests/messages_streaming.rs (94%) create mode 100644 litellm-rust/crates/llms-types/tests/ocr.rs create mode 100644 litellm-rust/crates/llms-types/tests/serde_compat.rs create mode 100644 litellm-rust/crates/llms-types/tests/wire_type.rs delete mode 100644 litellm-rust/crates/types/src/lib.rs delete mode 100644 litellm-rust/crates/types/src/llms/anthropic_messages/mod.rs delete mode 100644 litellm-rust/crates/types/src/llms/mod.rs delete mode 100644 litellm-rust/crates/types/src/llms/openai.rs delete mode 100644 litellm-rust/crates/types/src/messages/mod.rs delete mode 100644 litellm-rust/crates/types/src/responses/mod.rs delete mode 100644 litellm-rust/crates/types/src/utils.rs diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 1f3790c7b61..8d189c8c515 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -3619,7 +3619,6 @@ dependencies = [ "litellm-auth", "litellm-host", "litellm-host-python", - "litellm-types", "proptest", "pyo3", "rstest", @@ -3658,9 +3657,9 @@ dependencies = [ "litellm-host-native", "litellm-http", "litellm-llms", + "litellm-llms-types", "litellm-secrets", "litellm-tracing", - "litellm-types", "mime_guess", "moka", "rand 0.8.7", @@ -3688,13 +3687,12 @@ name = "litellm-core-utils" version = "0.1.0" dependencies = [ "fancy-regex 0.19.2", + "litellm-llms-types", "litellm-tracing", - "litellm-types", "rstest", "serde", "serde_json", "serde_path_to_error", - "serde_with", "strum", "thiserror 2.0.19", "url", @@ -3824,9 +3822,9 @@ dependencies = [ "litellm-host-http", "litellm-http", "litellm-llms", + "litellm-llms-types", "litellm-router", "litellm-secrets", - "litellm-types", "rstest", "serde", "serde_json", @@ -3999,9 +3997,9 @@ dependencies = [ "litellm-framing", "litellm-host", "litellm-http", + "litellm-llms-types", "litellm-python-compat", "litellm-secrets", - "litellm-types", "reqwest 0.12.28", "rstest", "serde", @@ -4015,13 +4013,26 @@ dependencies = [ "url", ] +[[package]] +name = "litellm-llms-types" +version = "0.1.0" +dependencies = [ + "macro_rules_attribute", + "rstest", + "schemars 1.2.2", + "serde", + "serde_json", + "serde_with", + "strum", +] + [[package]] name = "litellm-model-catalog" version = "0.1.0" dependencies = [ "indexmap 2.14.0", "jsonschema", - "litellm-types", + "litellm-llms-types", "rstest", "schemars 1.2.2", "serde", @@ -4059,12 +4070,12 @@ dependencies = [ "litellm-host-python", "litellm-http", "litellm-llms", + "litellm-llms-types", "litellm-secrets", "litellm-secrets-aws", "litellm-secrets-types", "litellm-token-counter", "litellm-tracing", - "litellm-types", "pyo3", "pyo3-async-runtimes", "qdrant-client", @@ -4354,17 +4365,6 @@ dependencies = [ "tracing-subscriber", ] -[[package]] -name = "litellm-types" -version = "0.1.0" -dependencies = [ - "rstest", - "schemars 1.2.2", - "serde", - "serde_json", - "strum", -] - [[package]] name = "litemap" version = "0.8.2" diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index 32919e23927..53aaf7a4d52 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -39,7 +39,7 @@ litellm-secrets-azure = { path = "crates/secrets-azure" } litellm-secrets-cyberark = { path = "crates/secrets-cyberark" } litellm-http = { path = "crates/http" } litellm-llms = { path = "crates/llms" } -litellm-types = { path = "crates/types" } +litellm-llms-types = { path = "crates/llms-types" } litellm-core-utils = { path = "crates/core-utils" } litellm-db = { path = "crates/db" } litellm-db-testing = { path = "crates/db-testing" } @@ -74,6 +74,7 @@ proptest = "1.7.0" pyo3 = "0.29.2" pyo3-async-runtimes = { version = "0.29.0", features = ["tokio-runtime"] } rand = "0.8" +macro_rules_attribute = "0.2.3" schemars = "1" reqwest = { version = "0.12", default-features = false, features = ["json", "multipart", "rustls-tls", "http2", "stream"] } qdrant-client = { version = "1.19.0", default-features = false } diff --git a/litellm-rust/crates/callbacks-legacy-python/AGENTS.md b/litellm-rust/crates/callbacks-legacy-python/AGENTS.md index e76a15099dc..de99fe17a4b 100644 --- a/litellm-rust/crates/callbacks-legacy-python/AGENTS.md +++ b/litellm-rust/crates/callbacks-legacy-python/AGENTS.md @@ -7,6 +7,7 @@ - The enum only shrinks: when Rust owns a subsystem, delete its group rather than adding a Rust path beside it - Calling a user's own callback directly is permanent Python surface and gets its own type outside `LegacyPython` - `PublicCall` is the caller's call as `Logging` sees it: the positional arguments, the keyword view as the call rewrites it (setup, deployment hook, preflight) and the bound request object backing omitted keywords; shared bridge composition hands it to `LegacyLogging`; routes use the neutral call boundary +- `LoggingOperation` selects legacy logging entrypoints and response handling. It belongs here rather than in shared inference data contracts - `setup` reuses a `Logging` passed as `litellm_logging_obj` (the proxy and Router) and otherwise builds one through `function_setup`; which callbacks run is `Logging`'s decision, never this crate's - Callbacks receive the caller's own objects and may mutate them; this crate alone carries that obligation - Retain complete boundary arguments, opaque values, aliases, omitted/default distinctions and deliberate copies; preserve the deployment-hook kwargs view diff --git a/litellm-rust/crates/callbacks-legacy-python/Cargo.toml b/litellm-rust/crates/callbacks-legacy-python/Cargo.toml index 8ee795092b4..ed5e0fb9691 100644 --- a/litellm-rust/crates/callbacks-legacy-python/Cargo.toml +++ b/litellm-rust/crates/callbacks-legacy-python/Cargo.toml @@ -6,7 +6,6 @@ license.workspace = true repository.workspace = true [dependencies] -litellm-types.workspace = true litellm-host.workspace = true litellm-host-python.workspace = true diff --git a/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs b/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs index a2505588761..21f563d9f3b 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs @@ -2,8 +2,8 @@ //! raises is answered with the same `Logging` calls, in the same order, as the Python //! `@client` path makes them. +use crate::LoggingOperation; use litellm_host_python::PythonOwned; -use litellm_types::Operation; use litellm_host::{ interceptors::{RawResponse, RequestContext, WireRequest}, @@ -45,7 +45,7 @@ struct LoggedRequest { } pub struct LegacyLogging { - operation: Operation, + operation: LoggingOperation, call: PublicCall, logger: Option, start: Py, @@ -68,7 +68,12 @@ fn is_cancellation(py: Python<'_>, error: &PyErr) -> bool { } impl LegacyLogging { - pub fn new(py: Python<'_>, operation: Operation, call: PublicCall, asynchronous: bool) -> Self { + pub fn new( + py: Python<'_>, + operation: LoggingOperation, + call: PublicCall, + asynchronous: bool, + ) -> Self { Self { operation, call, @@ -87,32 +92,34 @@ impl LegacyLogging { fn call_type(&self) -> &'static str { match (self.operation, self.asynchronous) { - (Operation::Completion, false) => "completion", - (Operation::Completion, true) => "acompletion", - (Operation::Responses, false) => "responses", - (Operation::Responses, true) => "aresponses", - (Operation::Messages, _) => "anthropic_messages", - (Operation::Ocr, false) => "ocr", - (Operation::Ocr, true) => "aocr", + (LoggingOperation::Completion, false) => "completion", + (LoggingOperation::Completion, true) => "acompletion", + (LoggingOperation::Responses, false) => "responses", + (LoggingOperation::Responses, true) => "aresponses", + (LoggingOperation::Messages, _) => "anthropic_messages", + (LoggingOperation::Ocr, false) => "ocr", + (LoggingOperation::Ocr, true) => "aocr", } } fn input_description(&self) -> &'static str { match self.operation { - Operation::Completion => "Chat completions", - Operation::Responses => "Responses", - Operation::Messages => "Messages", - Operation::Ocr => "OCR document processing", + LoggingOperation::Completion => "Chat completions", + LoggingOperation::Responses => "Responses", + LoggingOperation::Messages => "Messages", + LoggingOperation::Ocr => "OCR document processing", } } fn stream_billing(&self) -> Option { match self.operation { - Operation::Messages => Some(PassThroughStream { + LoggingOperation::Messages => Some(PassThroughStream { url_route: "/v1/messages", endpoint_type: "anthropic", }), - Operation::Completion | Operation::Responses | Operation::Ocr => None, + LoggingOperation::Completion | LoggingOperation::Responses | LoggingOperation::Ocr => { + None + } } } @@ -643,16 +650,16 @@ kwargs = {'logger': logger, 'document': document} } #[rstest] - #[case::sync_completion(litellm_types::Operation::Completion, false, "completion")] - #[case::async_completion(litellm_types::Operation::Completion, true, "acompletion")] - #[case::sync_responses(litellm_types::Operation::Responses, false, "responses")] - #[case::async_responses(litellm_types::Operation::Responses, true, "aresponses")] - #[case::sync_messages(litellm_types::Operation::Messages, false, "anthropic_messages")] - #[case::async_messages(litellm_types::Operation::Messages, true, "anthropic_messages")] - #[case::sync_ocr(litellm_types::Operation::Ocr, false, "ocr")] - #[case::async_ocr(litellm_types::Operation::Ocr, true, "aocr")] + #[case::sync_completion(crate::LoggingOperation::Completion, false, "completion")] + #[case::async_completion(crate::LoggingOperation::Completion, true, "acompletion")] + #[case::sync_responses(crate::LoggingOperation::Responses, false, "responses")] + #[case::async_responses(crate::LoggingOperation::Responses, true, "aresponses")] + #[case::sync_messages(crate::LoggingOperation::Messages, false, "anthropic_messages")] + #[case::async_messages(crate::LoggingOperation::Messages, true, "anthropic_messages")] + #[case::sync_ocr(crate::LoggingOperation::Ocr, false, "ocr")] + #[case::async_ocr(crate::LoggingOperation::Ocr, true, "aocr")] fn operation_selects_the_legacy_setup_and_deployment_hook_contract( - #[case] operation: litellm_types::Operation, + #[case] operation: crate::LoggingOperation, #[case] asynchronous: bool, #[case] expected: &str, ) { @@ -1088,12 +1095,12 @@ check = lambda: None } #[rstest] - #[case::completion(litellm_types::Operation::Completion, "Chat completions")] - #[case::responses(litellm_types::Operation::Responses, "Responses")] - #[case::messages(litellm_types::Operation::Messages, "Messages")] - #[case::ocr(litellm_types::Operation::Ocr, "OCR document processing")] + #[case::completion(crate::LoggingOperation::Completion, "Chat completions")] + #[case::responses(crate::LoggingOperation::Responses, "Responses")] + #[case::messages(crate::LoggingOperation::Messages, "Messages")] + #[case::ocr(crate::LoggingOperation::Ocr, "OCR document processing")] fn prepared_arguments_replace_the_legacy_view_without_losing_callback_aliases( - #[case] operation: litellm_types::Operation, + #[case] operation: crate::LoggingOperation, #[case] description: &str, ) { Python::initialize(); @@ -1763,7 +1770,7 @@ assert logger.calls[1][1] is response Python::attach(|py| { let locals = namespace(py, c"first = b'first'\nlast = b'last'\nresponse = None"); let mut logging = LegacyLogging { - operation: litellm_types::Operation::Messages, + operation: crate::LoggingOperation::Messages, ..logged(py, &locals, true) }; logging diff --git a/litellm-rust/crates/callbacks-legacy-python/src/lib.rs b/litellm-rust/crates/callbacks-legacy-python/src/lib.rs index bce186380b8..38c6b1aedbd 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/lib.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/lib.rs @@ -20,5 +20,13 @@ pub(crate) use callbacks::{LegacyCallbacks, is_internal_call}; pub(crate) use logger::{DeploymentHooks, PythonLogger, finalize, setup}; pub use mapping::{CallBoundary, CallbackMapping, Dispatch, callback_mappings}; +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum LoggingOperation { + Completion, + Responses, + Messages, + Ocr, +} + #[cfg(test)] mod test_support; diff --git a/litellm-rust/crates/callbacks-legacy-python/src/test_support.rs b/litellm-rust/crates/callbacks-legacy-python/src/test_support.rs index 39a879f9ff6..46c93369100 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/test_support.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/test_support.rs @@ -189,5 +189,5 @@ pub(crate) fn legacy_call( .map(|kwargs| kwargs.cast_into::().unwrap()) .unwrap_or_else(|| PyDict::new(py)); let call = PublicCall::capture(&request, &PyTuple::empty(py), &kwargs).unwrap(); - LegacyLogging::new(py, litellm_types::Operation::Ocr, call, asynchronous) + LegacyLogging::new(py, crate::LoggingOperation::Ocr, call, asynchronous) } diff --git a/litellm-rust/crates/core-utils/Cargo.toml b/litellm-rust/crates/core-utils/Cargo.toml index 22196979781..1feea5fcc08 100644 --- a/litellm-rust/crates/core-utils/Cargo.toml +++ b/litellm-rust/crates/core-utils/Cargo.toml @@ -8,11 +8,10 @@ repository.workspace = true [dependencies] fancy-regex.workspace = true litellm-tracing.workspace = true -litellm-types.workspace = true +litellm-llms-types.workspace = true serde.workspace = true serde_json.workspace = true serde_path_to_error = "0.1" -serde_with.workspace = true strum.workspace = true thiserror.workspace = true url.workspace = true diff --git a/litellm-rust/crates/core-utils/src/core_helpers.rs b/litellm-rust/crates/core-utils/src/core_helpers.rs index 9f00a0a5efe..ada1c3ceb1a 100644 --- a/litellm-rust/crates/core-utils/src/core_helpers.rs +++ b/litellm-rust/crates/core-utils/src/core_helpers.rs @@ -2,7 +2,7 @@ use std::time::{SystemTime, UNIX_EPOCH}; -use litellm_types::utils::{ChatCompletionsUsage, PromptTokensDetails}; +use litellm_llms_types::formats::chat_completions::{ChatCompletionsUsage, PromptTokensDetails}; /// OpenAI finish reasons, mirroring Python's `_FINISH_REASON_MAP` for the /// reasons the providers on this route can emit. Python warns and falls back to diff --git a/litellm-rust/crates/core-utils/src/get_provider_specific_headers.rs b/litellm-rust/crates/core-utils/src/get_provider_specific_headers.rs index bfcd448e2d8..c6597161a59 100644 --- a/litellm-rust/crates/core-utils/src/get_provider_specific_headers.rs +++ b/litellm-rust/crates/core-utils/src/get_provider_specific_headers.rs @@ -1,4 +1,4 @@ -use litellm_types::utils::{ProviderSpecificHeader, ProviderSpecificHeaders}; +use litellm_llms_types::headers::{ProviderSpecificHeader, ProviderSpecificHeaders}; use serde_json::{Map, Value}; pub fn get_provider_specific_headers( diff --git a/litellm-rust/crates/core-utils/src/prompt_templates/factory.rs b/litellm-rust/crates/core-utils/src/prompt_templates/factory.rs index 63ef79c0fa2..10a0d719e9d 100644 --- a/litellm-rust/crates/core-utils/src/prompt_templates/factory.rs +++ b/litellm-rust/crates/core-utils/src/prompt_templates/factory.rs @@ -10,7 +10,7 @@ //! `_bedrock_converse_messages_pt` for the text-only surface this route //! accepts; anything richer is declined upstream by the capability gate. -use litellm_types::llms::openai::{ChatMessage, ChatMessageContent}; +use litellm_llms_types::formats::chat_completions::{ChatMessage, ChatMessageContent}; use strum::IntoStaticStr; pub const EMPTY_TEXT_PLACEHOLDER: &str = diff --git a/litellm-rust/crates/core-utils/src/serde_compat.rs b/litellm-rust/crates/core-utils/src/serde_compat.rs index e3aaa2d8ead..3e4d82d3a3e 100644 --- a/litellm-rust/crates/core-utils/src/serde_compat.rs +++ b/litellm-rust/crates/core-utils/src/serde_compat.rs @@ -1,12 +1,3 @@ -use serde::{ - Deserializer, - de::{Error, Visitor}, -}; -use serde_with::DeserializeAs; - -pub struct LaxI64; -pub struct FiniteF64; - pub fn parse_str_bool(value: &str) -> Option { let token = value.trim_matches(|character: char| { character.is_whitespace() || matches!(character, '\u{1c}'..='\u{1f}') @@ -22,129 +13,12 @@ pub fn parse_redis_bool(value: &str) -> bool { value == "1" || value.eq_ignore_ascii_case("true") || value.eq_ignore_ascii_case("yes") } -impl<'de> DeserializeAs<'de, i64> for LaxI64 { - fn deserialize_as>(deserializer: D) -> Result { - deserializer.deserialize_any(Self) - } -} - -impl<'de> Visitor<'de> for LaxI64 { - type Value = i64; - - fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - formatter.write_str("an integer in the i64 range") - } - - fn visit_i64(self, value: i64) -> Result { - Ok(value) - } - - fn visit_u64(self, value: u64) -> Result { - i64::try_from(value).map_err(E::custom) - } - - fn visit_f64(self, value: f64) -> Result { - integral_float(value).ok_or_else(|| E::custom("expected an integer in the i64 range")) - } - - fn visit_str(self, value: &str) -> Result { - integer_string(value.trim()) - .ok_or_else(|| E::custom("expected an integer in the i64 range")) - } - - fn visit_bool(self, value: bool) -> Result { - Ok(i64::from(value)) - } -} - -impl<'de> DeserializeAs<'de, f64> for FiniteF64 { - fn deserialize_as>(deserializer: D) -> Result { - deserializer.deserialize_any(Self) - } -} - -impl<'de> Visitor<'de> for FiniteF64 { - type Value = f64; - - fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - formatter.write_str("a finite number") - } - - fn visit_i64(self, value: i64) -> Result { - Ok(value as f64) - } - - fn visit_u64(self, value: u64) -> Result { - Ok(value as f64) - } - - fn visit_f64(self, value: f64) -> Result { - value - .is_finite() - .then_some(value) - .ok_or_else(|| E::custom("expected a finite number")) - } - - fn visit_str(self, value: &str) -> Result { - self.visit_f64(value.trim().parse::().map_err(E::custom)?) - } - - fn visit_bool(self, value: bool) -> Result { - Ok(f64::from(value)) - } -} - -fn integer_string(value: &str) -> Option { - let integer = match value.split_once('.') { - Some((integer, fraction)) => { - if fraction.is_empty() || !fraction.bytes().all(|byte| byte == b'0') { - return None; - } - integer - } - None => value, - }; - if integer.starts_with('_') || integer.ends_with('_') || integer.contains("__") { - return None; - } - let digits = integer.strip_prefix(['+', '-']).unwrap_or(integer); - if digits.is_empty() - || digits.starts_with('_') - || !digits - .bytes() - .all(|byte| byte.is_ascii_digit() || byte == b'_') - { - return None; - } - integer.replace('_', "").parse().ok() -} - -fn integral_float(value: f64) -> Option { - (value.is_finite() - && value.fract() == 0.0 - && value >= i64::MIN as f64 - && value < -(i64::MIN as f64)) - .then_some(value as i64) -} - #[cfg(test)] mod tests { use rstest::rstest; - use serde::{Deserialize, Serialize}; - use serde_json::json; - use serde_with::serde_as; use super::*; - #[serde_as] - #[derive(Debug, Deserialize, Serialize, PartialEq)] - struct Numbers { - #[serde_as(deserialize_as = "Option>")] - integers: Option>, - #[serde_as(deserialize_as = "Option")] - float: Option, - } - #[rstest] #[case::trimmed_true(" True ", Some(true))] #[case::control_whitespace_true("\u{1c}TRUE\u{1f}", Some(true))] @@ -160,73 +34,4 @@ mod tests { ) { assert_eq!(parse_str_bool(input), expected, "{input:?}"); } - - #[test] - fn adapters_compose_and_serialize_as_numbers() { - let numbers: Numbers = serde_json::from_value(json!({ - "integers": ["9007199254740993.0", "1_000", " +2.000 ", 3.0, true], - "float": " 1.5 " - })) - .unwrap(); - assert_eq!( - serde_json::to_value(numbers).unwrap(), - json!({ - "integers": [9_007_199_254_740_993_i64, 1000, 2, 3, 1], "float": 1.5 - }) - ); - for input in [json!({}), json!({"integers": null, "float": null})] { - assert_eq!( - serde_json::from_value::(input).unwrap(), - Numbers { - integers: None, - float: None, - } - ); - } - } - - #[test] - fn integer_bounds_and_invalid_values_are_checked() { - for input in [ - json!(i64::MIN), - json!(i64::MAX), - json!(i64::MAX.to_string()), - ] { - assert!(serde_json::from_value::(json!({"integers": [input]})).is_ok()); - } - for input in [ - json!(u64::MAX), - json!(9_223_372_036_854_775_808_u64), - json!(9_223_372_036_854_775_808.0), - json!("-9223372036854775809"), - json!("1.0000000000000001"), - json!("1e3"), - json!("2."), - json!(".0"), - json!("_2"), - json!("2__0"), - json!(2.5), - json!(null), - json!({}), - ] { - assert!(serde_json::from_value::(json!({"integers": [input]})).is_err()); - } - } - - #[test] - fn floats_reject_nonfinite_and_invalid_values() { - for input in [ - json!("NaN"), - json!("inf"), - json!("-inf"), - json!("1e999"), - json!([]), - ] { - assert!(serde_json::from_value::(json!({"float": input})).is_err()); - } - for (input, expected) in [(json!(2), 2.0), (json!(2.5), 2.5), (json!(true), 1.0)] { - let numbers: Numbers = serde_json::from_value(json!({"float": input})).unwrap(); - assert_eq!(numbers.float, Some(expected)); - } - } } diff --git a/litellm-rust/crates/core/AGENTS.md b/litellm-rust/crates/core/AGENTS.md index 217fdfc5e11..ec96239beac 100644 --- a/litellm-rust/crates/core/AGENTS.md +++ b/litellm-rust/crates/core/AGENTS.md @@ -10,11 +10,11 @@ Responses WebSocket sessions remain separate from the HTTP call driver because a ## Crate layering -For Messages, Responses, Chat Completions, OCR, and other API formats, `core/src//` owns orchestration. Shared API data contracts belong in `litellm-types`, adapter contracts and shared transformation machinery in `llms/src/base_llm//`, and provider policy in `llms/src///`. A repeated format directory name does not imply interchangeable responsibilities. Select concrete adapters here, then invoke their contracts instead of applying one provider's policy to every call. Route types describe call envelopes and execution state, not duplicate public payload schemas +For Messages, Responses, Chat Completions, OCR, and other API formats, `core/src//` owns orchestration. Shared API data contracts belong in `litellm-llms-types`, adapter contracts and shared transformation machinery in `llms/src/base_llm//`, and provider policy in `llms/src///`. A repeated format directory name does not imply interchangeable responsibilities. Select concrete adapters here, then invoke their contracts instead of applying one provider's policy to every call. Route types describe call envelopes and execution state, not duplicate public payload schemas -Each crate mirrors one top-level Python package, so a Rust path reads as its Python path with the crate name in place of the package directory. Dependencies only point down: +Crates separate API data, transformations, transport, and orchestration. Python package names identify counterparts, not ownership. Dependencies only point down: -- `litellm-types` mirrors `litellm/types/`: pure serde data, no I/O +- `litellm-llms-types` owns shared inference API contracts, grouped by format: pure serde data and shape validation, no I/O - `litellm-core-utils` mirrors `litellm/litellm_core_utils/`: pure helpers (provider resolution, prompt factory, call arguments, settings lookup and layer merge), no network I/O - `litellm-http` is Rust-only and route-neutral: settings resolution, the pooled `reqwest` clients, TLS, proxies, the SSRF-safe media fetcher, request and header helpers, and transport errors. Python's `litellm/llms/custom_httpx/` is split by responsibility instead of mirrored: its transport half lives here, its OCR handler in `litellm-llms` - `litellm-llms` mirrors `litellm/llms/`: `base_llm//transformation.rs`, `//transformation.rs`, and `base_llm/ocr/handler.rs` (the OCR request handler) diff --git a/litellm-rust/crates/core/Cargo.toml b/litellm-rust/crates/core/Cargo.toml index 56f9a0c4163..8410aff1d6a 100644 --- a/litellm-rust/crates/core/Cargo.toml +++ b/litellm-rust/crates/core/Cargo.toml @@ -11,7 +11,7 @@ litellm-cache-response.workspace = true litellm-framing.workspace = true tokio-util = { version = "0.7", features = ["codec"] } litellm-secrets.workspace = true -litellm-types.workspace = true +litellm-llms-types.workspace = true litellm-core-utils.workspace = true litellm-host.workspace = true bytes.workspace = true diff --git a/litellm-rust/crates/core/src/chat_completions/handler.rs b/litellm-rust/crates/core/src/chat_completions/handler.rs index 8cfee9b59cf..9f9d48cb177 100644 --- a/litellm-rust/crates/core/src/chat_completions/handler.rs +++ b/litellm-rust/crates/core/src/chat_completions/handler.rs @@ -1,5 +1,4 @@ -use litellm_host::lifecycle::ExecutionEvent; -use litellm_host::observation::ObservationSender; +use litellm_host::{lifecycle::ExecutionEvent, observation::ObservationSender}; use std::time::Duration; use litellm_auth::AuthServices; @@ -9,7 +8,7 @@ use litellm_llms::base_llm::{ auth::{Authenticated, resolve_auth}, chat::transformation::ProviderChatResponseData, }; -use litellm_types::utils::ChatCompletionsResponse; +use litellm_llms_types::formats::chat_completions::ChatCompletionsResponse; use serde_json::Value; use super::Error; diff --git a/litellm-rust/crates/core/src/chat_completions/mod.rs b/litellm-rust/crates/core/src/chat_completions/mod.rs index a64249fa185..a26648b88ef 100644 --- a/litellm-rust/crates/core/src/chat_completions/mod.rs +++ b/litellm-rust/crates/core/src/chat_completions/mod.rs @@ -5,7 +5,7 @@ pub use crate::error::RouteError as Error; mod common_utils; pub(crate) mod handler; mod prepare; -use litellm_types::utils::ChatCompletionsResponse; +use litellm_llms_types::formats::chat_completions::ChatCompletionsResponse; use prepare::{prepare_provider_request, resolve_request}; use crate::chat_completions::types::ChatCompletionsRequest; diff --git a/litellm-rust/crates/core/src/chat_completions/prepare.rs b/litellm-rust/crates/core/src/chat_completions/prepare.rs index ec832b3e59a..fd4d8700fa2 100644 --- a/litellm-rust/crates/core/src/chat_completions/prepare.rs +++ b/litellm-rust/crates/core/src/chat_completions/prepare.rs @@ -2,8 +2,8 @@ use litellm_auth::SecretValue; use litellm_core_utils::settings::Lookup; use litellm_http::request::with_default_headers; use litellm_llms::base_llm::{auth::ValidatedEnvironment, chat::transformation::BaseConfig}; +use litellm_llms_types::formats::chat_completions::ChatMessage; use litellm_secrets::source::Secrets; -use litellm_types::llms::openai::ChatMessage; use serde_json::Value; use super::{ diff --git a/litellm-rust/crates/core/src/chat_completions/route.rs b/litellm-rust/crates/core/src/chat_completions/route.rs index d43d2bf9eef..9da86dfa27d 100644 --- a/litellm-rust/crates/core/src/chat_completions/route.rs +++ b/litellm-rust/crates/core/src/chat_completions/route.rs @@ -5,7 +5,7 @@ use litellm_host::{ call::{CallOutput, HostedMachine, hosted_call}, protocol::Protocol, }; -use litellm_types::utils::ChatCompletionsResponse; +use litellm_llms_types::formats::chat_completions::ChatCompletionsResponse; use super::{ ChatCompletionsRoute, Error, diff --git a/litellm-rust/crates/core/src/chat_completions/types.rs b/litellm-rust/crates/core/src/chat_completions/types.rs index 73d6378fc92..5d7d04804f7 100644 --- a/litellm-rust/crates/core/src/chat_completions/types.rs +++ b/litellm-rust/crates/core/src/chat_completions/types.rs @@ -3,7 +3,7 @@ use std::time::Duration; use litellm_auth::SecretValue; use litellm_llms::base_llm::{auth::ValidatedEnvironment, chat::transformation::BaseConfig}; -use litellm_types::llms::openai::ChatMessage; +use litellm_llms_types::formats::chat_completions::ChatMessage; use serde_json::{Map, Value}; /// A `/chat/completions` call as it crosses into the core. diff --git a/litellm-rust/crates/core/src/messages/AGENTS.md b/litellm-rust/crates/core/src/messages/AGENTS.md index 0bea24a65ce..0feff9c30c2 100644 --- a/litellm-rust/crates/core/src/messages/AGENTS.md +++ b/litellm-rust/crates/core/src/messages/AGENTS.md @@ -1,4 +1,4 @@ -This directory owns provider-independent Messages call orchestration: the entrypoint, call envelopes, provider selection, credential resolution, transport coordination, hooks, and stream lifecycle. Shared API data contracts belong in `litellm-types::messages`, adapter contracts and execution inputs in `llms/src/base_llm/messages`, and provider implementations in `llms/src//messages` +This directory owns provider-independent Messages call orchestration: the entrypoint, call envelopes, provider selection, credential resolution, transport coordination, hooks, and stream lifecycle. Shared API data contracts belong in `litellm-llms-types::formats::messages`, adapter contracts and execution inputs in `llms/src/base_llm/messages`, and provider implementations in `llms/src//messages` Select concrete provider adapters and invoke their contracts. Delegate authentication policy, beta selection, payload rewriting, and response interpretation to those adapters. Keep provider policy out of request preparation and transport handlers. Calling a concrete provider helper for every provider is still a policy dependency diff --git a/litellm-rust/crates/core/src/messages/common_utils.rs b/litellm-rust/crates/core/src/messages/common_utils.rs index 8f5ad05d3dc..98a92c90dba 100644 --- a/litellm-rust/crates/core/src/messages/common_utils.rs +++ b/litellm-rust/crates/core/src/messages/common_utils.rs @@ -3,7 +3,7 @@ pub(super) use litellm_http::request::truncate_error_body; use litellm_llms::{ anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG, azure_ai::messages::transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG, - base_llm::messages::transformation::BaseAnthropicMessagesConfig, + base_llm::messages::transformation::BaseMessagesConfig, bedrock::messages::invoke_transformations::anthropic_claude3_transformation::BEDROCK_ANTHROPIC_MESSAGES_CONFIG, }; use serde_json::{Map, Value}; @@ -30,7 +30,7 @@ impl MessagesProvider { .into() } - pub(crate) fn config(self) -> &'static dyn BaseAnthropicMessagesConfig { + pub(crate) fn config(self) -> &'static dyn BaseMessagesConfig { match self { Self::Anthropic => &ANTHROPIC_MESSAGES_CONFIG, Self::AzureAi => &AZURE_ANTHROPIC_MESSAGES_CONFIG, diff --git a/litellm-rust/crates/core/src/messages/handler.rs b/litellm-rust/crates/core/src/messages/handler.rs index 4d379e27ca4..5194156b6eb 100644 --- a/litellm-rust/crates/core/src/messages/handler.rs +++ b/litellm-rust/crates/core/src/messages/handler.rs @@ -1,5 +1,4 @@ -use litellm_host::lifecycle::ExecutionEvent; -use litellm_host::observation::ObservationSender; +use litellm_host::{lifecycle::ExecutionEvent, observation::ObservationSender}; use std::time::Duration; use bytes::Bytes; @@ -11,15 +10,16 @@ use litellm_llms::base_llm::{ auth::{Authenticated, resolve_auth}, messages::{ streaming::{ByteStream, StreamDecoder, encode_anthropic_sse}, - transformation::BaseAnthropicMessagesConfig, + transformation::BaseMessagesConfig, }, }; +use litellm_llms_types::formats::messages::MessagesResponse; use litellm_tracing::ByteChunk; -use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; use serde_json::Value; use super::{ - Error, MessagesResponse, common_utils::truncate_error_body, prepare::ProviderMessagesRequest, + Error, MessagesCallResponse, common_utils::truncate_error_body, + prepare::ProviderMessagesRequest, }; use crate::{constants::MESSAGES_TIMEOUT_SECS, outbound::outbound_request}; @@ -31,7 +31,7 @@ pub(super) async fn execute( cache_options: Option, interceptors: &impl Interceptors, observers: Option<&ObservationSender>, -) -> Result { +) -> Result { let ProviderMessagesRequest { provider, url, @@ -110,7 +110,7 @@ pub(super) async fn execute( .await .map_err(Error::post_call)?; decode_response(config, &body.model, &text) - .map(|message| MessagesResponse::Complete(Box::new(message))) + .map(|message| MessagesCallResponse::Complete(Box::new(message))) }, ) .await @@ -158,10 +158,10 @@ async fn provider_error(response: reqwest::Response) -> Error { } fn decode_response( - config: &dyn BaseAnthropicMessagesConfig, + config: &dyn BaseMessagesConfig, model: &str, text: &str, -) -> Result { +) -> Result { let response = serde_json::from_str(text).map_err(|err| { Error::InvalidResponse(litellm_llms::ErrorDetail::invalid( "messages response JSON", @@ -177,7 +177,7 @@ fn streaming_response( response: reqwest::Response, decoder: Option, provider: &'static str, -) -> MessagesResponse { +) -> MessagesCallResponse { let headers = response .headers() .iter() @@ -194,7 +194,7 @@ fn streaming_response( .boxed(), Some(decode) => decoded_chunks(response, decode, provider), }; - MessagesResponse::Stream { + MessagesCallResponse::Stream { head: super::route::MessagesStreamHead { headers }, chunks, } @@ -268,7 +268,7 @@ mod tests { .send() .await .unwrap(); - let MessagesResponse::Stream { mut chunks, .. } = + let MessagesCallResponse::Stream { mut chunks, .. } = streaming_response(response, Some(anthropic_sse_event_stream), "test") else { panic!("a streaming response returns chunks"); diff --git a/litellm-rust/crates/core/src/messages/mod.rs b/litellm-rust/crates/core/src/messages/mod.rs index 23e0a3fb624..8374e2ca32d 100644 --- a/litellm-rust/crates/core/src/messages/mod.rs +++ b/litellm-rust/crates/core/src/messages/mod.rs @@ -10,7 +10,7 @@ use litellm_secrets::source::SecretSource; use std::sync::Arc; pub use crate::error::RouteError as Error; -pub use types::{MessagesCall, MessagesResponse, MessagesShaping, messages_body}; +pub use types::{MessagesCall, MessagesCallResponse, MessagesShaping, messages_body}; #[derive(Clone)] pub struct MessagesRoute { @@ -103,7 +103,7 @@ impl MessagesRoute { call: MessagesCall, interceptors: &impl litellm_host::interceptors::Interceptors, options: impl Into, - ) -> Result { + ) -> Result { let crate::CallOptions { cache: cache_options, observers, @@ -129,7 +129,7 @@ impl MessagesRoute { cache_options: Option, interceptors: &impl litellm_host::interceptors::Interceptors, observers: Option<&ObservationSender>, - ) -> Result { + ) -> Result { crate::diagnostic::call(async { self.run_provider(call, cache_options, interceptors, observers) .await @@ -143,10 +143,10 @@ impl MessagesRoute { cache_options: Option, interceptors: &impl litellm_host::interceptors::Interceptors, observers: Option<&ObservationSender>, - ) -> Result { + ) -> Result { let request = prepare::prepare(call, self.secrets.as_ref()).await?; crate::diagnostic::provider(&request.body.model, request.provider.as_str()); - let execute: futures_util::future::BoxFuture<'_, Result> = + let execute: futures_util::future::BoxFuture<'_, Result> = Box::pin(handler::execute( &self.http, &self.auth, diff --git a/litellm-rust/crates/core/src/messages/prepare.rs b/litellm-rust/crates/core/src/messages/prepare.rs index 13e77d4649b..b8e0a40b230 100644 --- a/litellm-rust/crates/core/src/messages/prepare.rs +++ b/litellm-rust/crates/core/src/messages/prepare.rs @@ -9,8 +9,8 @@ use litellm_http::request::with_default_headers; use litellm_llms::base_llm::{ auth::ValidatedEnvironment, messages::context::MessagesTransformContext, }; +use litellm_llms_types::formats::messages::MessagesRequest; use litellm_secrets::source::SecretSource; -use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest; use super::{ Error, MessagesCall, @@ -27,7 +27,7 @@ struct ResolvedProvider { pub(super) struct ProviderMessagesRequest { pub(super) provider: MessagesProvider, pub(super) url: String, - pub(super) body: AnthropicMessagesRequest, + pub(super) body: MessagesRequest, pub(super) environment: ValidatedEnvironment, pub(super) timeout: Option, /// The caller's own credential, reported to the host beside the wire request. @@ -79,7 +79,7 @@ fn prepare_provider_request( let env_lookup = |key: &str| secrets.get(key); let sanitized = config.shape_request( - AnthropicMessagesRequest { model, ..body }, + MessagesRequest { model, ..body }, shaping.reasoning_auto_summary, )?; let trimmed = without_additional_drop_params(sanitized, &shaping.additional_drop_params)?; @@ -124,9 +124,9 @@ fn prepare_provider_request( } fn without_additional_drop_params( - request: AnthropicMessagesRequest, + request: MessagesRequest, paths: &[String], -) -> Result { +) -> Result { if paths.is_empty() { return Ok(request); } @@ -134,7 +134,7 @@ fn without_additional_drop_params( let trimmed = paths .iter() .fold(params, |params, path| delete_nested_value(params, path)); - Ok(AnthropicMessagesRequest { + Ok(MessagesRequest { params: serde_json::from_value(trimmed).map_err(invalid_request)?, ..request }) @@ -143,7 +143,7 @@ fn without_additional_drop_params( #[cfg(test)] mod tests { use litellm_llms::base_llm::auth::resolve_auth; - use litellm_types::utils::ProviderSpecificHeaders; + use litellm_llms_types::headers::ProviderSpecificHeaders; use rstest::{fixture, rstest}; use serde_json::{Map, Value, json}; @@ -155,7 +155,7 @@ mod tests { MessagesShaping::default() } - fn body(value: Value) -> AnthropicMessagesRequest { + fn body(value: Value) -> MessagesRequest { serde_json::from_value(value).unwrap() } diff --git a/litellm-rust/crates/core/src/messages/route.rs b/litellm-rust/crates/core/src/messages/route.rs index 1d2f95da957..85aa6c0995a 100644 --- a/litellm-rust/crates/core/src/messages/route.rs +++ b/litellm-rust/crates/core/src/messages/route.rs @@ -5,11 +5,11 @@ use litellm_host::{ call::{HostedCompletion, HostedMachine, hosted_call}, protocol::Protocol, }; -use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; +use litellm_llms_types::formats::messages::MessagesResponse; use super::{Error, MessagesCall}; -pub type MessagesOutput = HostedCompletion>; +pub type MessagesOutput = HostedCompletion>; /// The upstream response as the caller sees it at stream hand-off, before any chunk. pub struct MessagesStreamHead { @@ -19,7 +19,7 @@ pub struct MessagesStreamHead { pub struct Messages; impl Protocol for Messages { - type Response = Box; + type Response = Box; type Error = Error; type Request = MessagesCall; type HostCall = Infallible; diff --git a/litellm-rust/crates/core/src/messages/types.rs b/litellm-rust/crates/core/src/messages/types.rs index bc77b1dbded..6736e9178ba 100644 --- a/litellm-rust/crates/core/src/messages/types.rs +++ b/litellm-rust/crates/core/src/messages/types.rs @@ -2,12 +2,10 @@ use std::time::Duration; use bytes::Bytes; use litellm_host::call::CallOutput; -use litellm_llms::base_llm::messages::context::MessagesModelCapabilities as AnthropicModelCapabilities; -use litellm_types::{ - llms::anthropic_messages::{ - anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse, - }, - utils::ProviderSpecificHeaders, +use litellm_llms::base_llm::messages::context::MessagesModelCapabilities; +use litellm_llms_types::{ + formats::messages::{MessagesRequest, MessagesResponse}, + headers::ProviderSpecificHeaders, }; use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; @@ -15,7 +13,7 @@ use serde_json::{Map, Value}; use super::Error; pub struct MessagesCall { - pub body: AnthropicMessagesRequest, + pub body: MessagesRequest, pub api_key: Option, pub api_base: Option, pub custom_llm_provider: Option, @@ -25,7 +23,7 @@ pub struct MessagesCall { pub shaping: MessagesShaping, } -pub fn messages_body(body: Map) -> Result { +pub fn messages_body(body: Map) -> Result { serde_json::from_value(Value::Object(body)).map_err(invalid_request) } @@ -33,13 +31,13 @@ pub(super) fn invalid_request(err: serde_json::Error) -> Error { Error::InvalidRequest(format!("invalid Anthropic messages request: {err}").into()) } -pub type MessagesResponse = - CallOutput, super::route::MessagesStreamHead, Bytes, Error>; +pub type MessagesCallResponse = + CallOutput, super::route::MessagesStreamHead, Bytes, Error>; #[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] pub struct MessagesShaping { #[serde(default)] - pub capabilities: AnthropicModelCapabilities, + pub capabilities: MessagesModelCapabilities, #[serde(default)] pub drop_params: bool, #[serde(default)] @@ -76,9 +74,9 @@ mod tests { #[case::partial_capabilities( json!({"capabilities": {"supports_reasoning": true}}), MessagesShaping { - capabilities: AnthropicModelCapabilities { + capabilities: MessagesModelCapabilities { supports_reasoning: true, - ..AnthropicModelCapabilities::default() + ..MessagesModelCapabilities::default() }, ..MessagesShaping::default() }, @@ -100,7 +98,7 @@ mod tests { "additional_drop_params": ["metadata.user_id", "thinking"] }), MessagesShaping { - capabilities: AnthropicModelCapabilities { + capabilities: MessagesModelCapabilities { supports_reasoning: true, supports_adaptive_thinking: true, thinking_always_on: false, diff --git a/litellm-rust/crates/core/src/ocr/client.rs b/litellm-rust/crates/core/src/ocr/client.rs index df1fd1cda92..c193d9174f7 100644 --- a/litellm-rust/crates/core/src/ocr/client.rs +++ b/litellm-rust/crates/core/src/ocr/client.rs @@ -2,9 +2,8 @@ use litellm_host::observation::ObservationSender; use std::sync::Arc; use litellm_host::interceptors::Interceptors; -use litellm_llms::base_llm::ocr::{ - error::Error, handler::OcrClient, transformation::LiteLLMOcrResponse, -}; +use litellm_llms::base_llm::ocr::{error::Error, handler::OcrClient}; +use litellm_llms_types::formats::ocr::LiteLLMOcrResponse; use super::{ handler::perform_ocr_request, diff --git a/litellm-rust/crates/core/src/ocr/document.rs b/litellm-rust/crates/core/src/ocr/document.rs index ce4170323d8..4be1eb932eb 100644 --- a/litellm-rust/crates/core/src/ocr/document.rs +++ b/litellm-rust/crates/core/src/ocr/document.rs @@ -1,10 +1,8 @@ use std::{collections::BTreeMap as Map, io::Read, path::Path}; use base64::{Engine, engine::general_purpose::STANDARD}; -use litellm_llms::base_llm::ocr::{ - error::Error, - transformation::{OCR_INLINE_MAX_BYTES, OcrDocument}, -}; +use litellm_llms::base_llm::ocr::{error::Error, transformation::OCR_INLINE_MAX_BYTES}; +use litellm_llms_types::formats::ocr::OcrDocument; use crate::ocr::types::OcrDocumentInput; diff --git a/litellm-rust/crates/core/src/ocr/handler.rs b/litellm-rust/crates/core/src/ocr/handler.rs index 7aed179e07a..106b7162f2d 100644 --- a/litellm-rust/crates/core/src/ocr/handler.rs +++ b/litellm-rust/crates/core/src/ocr/handler.rs @@ -1,12 +1,12 @@ use futures_util::future::BoxFuture; use litellm_host::interceptors::{Interceptors, RawResponse, RequestContext, WireRequest}; -use litellm_host::lifecycle::ExecutionEvent; -use litellm_host::observation::ObservationSender; +use litellm_host::{lifecycle::ExecutionEvent, observation::ObservationSender}; use litellm_llms::base_llm::ocr::{ error::Error, handler::{CallHooks, OcrClient}, - transformation::{LiteLLMOcrResponse, PreparedOcrRequest}, + transformation::PreparedOcrRequest, }; +use litellm_llms_types::formats::ocr::LiteLLMOcrResponse; use serde_json::Value; use super::{arguments::is_secret_param, prepare::prepare_request, provider_config::OcrConfigKind}; diff --git a/litellm-rust/crates/core/src/ocr/prepare.rs b/litellm-rust/crates/core/src/ocr/prepare.rs index c2b401a7ab9..0f4e217c074 100644 --- a/litellm-rust/crates/core/src/ocr/prepare.rs +++ b/litellm-rust/crates/core/src/ocr/prepare.rs @@ -85,12 +85,13 @@ mod tests { base_llm::ocr::{ error::Error, handler::{CallHooks, OcrClient}, - transformation::{BaseOcrConfig, OcrResponseFormat}, + transformation::BaseOcrConfig, }, cohere::ocr::transformation::CohereParseConfig, mistral::ocr::transformation::MistralOcrConfig, vertex_ai::ocr::transformation::VertexAiOcrConfig, }; + use litellm_llms_types::formats::ocr::OcrResponseFormat; use serde_json::{Value, json}; use super::*; diff --git a/litellm-rust/crates/core/src/ocr/provider_config.rs b/litellm-rust/crates/core/src/ocr/provider_config.rs index 1e27b83c4c5..6e934b53e82 100644 --- a/litellm-rust/crates/core/src/ocr/provider_config.rs +++ b/litellm-rust/crates/core/src/ocr/provider_config.rs @@ -14,8 +14,7 @@ use litellm_llms::{ error::Error, handler::{self, CallHooks, OcrClient}, transformation::{ - BaseOcrConfig, LiteLLMOcrResponse, OcrCredentialInputs, OcrDocument, OcrResponseFormat, - PreparedOcrRequest, ResolvedOcrCredentials, + BaseOcrConfig, OcrCredentialInputs, PreparedOcrRequest, ResolvedOcrCredentials, }, }, cohere::ocr::transformation::CohereParseConfig, @@ -25,6 +24,7 @@ use litellm_llms::{ deepseek_transformation::VertexAIDeepSeekOCRConfig, transformation::VertexAiOcrConfig, }, }; +use litellm_llms_types::formats::ocr::{LiteLLMOcrResponse, OcrDocument, OcrResponseFormat}; macro_rules! with_config { ($kind:expr, $config:ident => $body:expr) => { diff --git a/litellm-rust/crates/core/src/ocr/route.rs b/litellm-rust/crates/core/src/ocr/route.rs index b576585049a..f6f6533929c 100644 --- a/litellm-rust/crates/core/src/ocr/route.rs +++ b/litellm-rust/crates/core/src/ocr/route.rs @@ -6,7 +6,8 @@ use litellm_host::{ protocol::Protocol, protocol::Reply, }; -use litellm_llms::base_llm::ocr::{error::Error, transformation::LiteLLMOcrResponse}; +use litellm_llms::base_llm::ocr::error::Error; +use litellm_llms_types::formats::ocr::LiteLLMOcrResponse; use crate::ocr::types::{LiteLLMOcrRequest, OcrDocumentInput}; diff --git a/litellm-rust/crates/core/src/ocr/types.rs b/litellm-rust/crates/core/src/ocr/types.rs index 20a21e43676..fda65e3284e 100644 --- a/litellm-rust/crates/core/src/ocr/types.rs +++ b/litellm-rust/crates/core/src/ocr/types.rs @@ -5,10 +5,9 @@ use litellm_auth::{InputSource, SecretValue, TokenProviderHandle}; use litellm_core_utils::call_arguments::CallArguments; use litellm_llms::base_llm::ocr::{ error::Error, - transformation::{ - OcrCredentialInputs, OcrDocument, OcrResponseFormat, OcrTransportConfig, response_format, - }, + transformation::{OcrCredentialInputs, OcrTransportConfig, response_format}, }; +use litellm_llms_types::formats::ocr::{OcrDocument, OcrResponseFormat}; use serde_json::{Map, Value}; use super::provider_config::{OcrConfigKind, resolve_provider_config}; @@ -222,7 +221,7 @@ mod tests { use super::*; fn document() -> OcrDocument { - OcrDocument::try_from( + serde_json::from_value( json!({"type":"document_url","document_url":"data:application/pdf;base64,YWJj"}), ) .unwrap() diff --git a/litellm-rust/crates/core/src/ocr/wire.rs b/litellm-rust/crates/core/src/ocr/wire.rs index b9c60f57e3c..ca83e26e3e7 100644 --- a/litellm-rust/crates/core/src/ocr/wire.rs +++ b/litellm-rust/crates/core/src/ocr/wire.rs @@ -1,10 +1,8 @@ use std::{collections::BTreeMap, time::Duration}; use litellm_auth::{InputSource, SecretValue}; -use litellm_llms::base_llm::ocr::{ - error::Error, - transformation::{OcrDocument, decode_request_value}, -}; +use litellm_llms::base_llm::ocr::{error::Error, transformation::decode_request_value}; +use litellm_llms_types::formats::ocr::OcrDocument; use serde::Deserialize; use serde_json::{Map, Value}; diff --git a/litellm-rust/crates/core/src/responses/route.rs b/litellm-rust/crates/core/src/responses/route.rs index cc641a22acc..2cdd143ee2b 100644 --- a/litellm-rust/crates/core/src/responses/route.rs +++ b/litellm-rust/crates/core/src/responses/route.rs @@ -5,7 +5,7 @@ use litellm_host::{ call::{HostedMachine, hosted_call}, protocol::Protocol, }; -use litellm_types::responses::main::ResponsesApiResponse; +use litellm_llms_types::formats::responses::ResponsesApiResponse; use super::{ Error, ResponsesRoute, diff --git a/litellm-rust/crates/core/src/responses/types.rs b/litellm-rust/crates/core/src/responses/types.rs index ce634a9862f..18c64dc8178 100644 --- a/litellm-rust/crates/core/src/responses/types.rs +++ b/litellm-rust/crates/core/src/responses/types.rs @@ -5,7 +5,7 @@ use litellm_host::call::CallOutput; use litellm_llms::base_llm::{ auth::ValidatedEnvironment, responses::transformation::BaseResponsesApiConfig, }; -use litellm_types::responses::main::ResponsesApiResponse; +use litellm_llms_types::formats::responses::ResponsesApiResponse; use serde_json::{Map, Value}; use super::Error; diff --git a/litellm-rust/crates/core/src/responses/websocket.rs b/litellm-rust/crates/core/src/responses/websocket.rs index 69165186d25..4ceff787a66 100644 --- a/litellm-rust/crates/core/src/responses/websocket.rs +++ b/litellm-rust/crates/core/src/responses/websocket.rs @@ -2,7 +2,7 @@ use std::{collections::HashMap, sync::Arc, time::Duration}; use futures_util::{SinkExt, StreamExt}; use litellm_http::websocket::{UpstreamWebSocket, connect_upstream}; -use litellm_types::responses::streaming_websocket::ResponsesWsEventType; +use litellm_llms_types::formats::responses::streaming_websocket::ResponsesWsEventType; use tokio::sync::Mutex; use tokio_tungstenite::tungstenite::{ Message, diff --git a/litellm-rust/crates/core/tests/caching.rs b/litellm-rust/crates/core/tests/caching.rs index 77c6b4cde1d..51fab8b163b 100644 --- a/litellm-rust/crates/core/tests/caching.rs +++ b/litellm-rust/crates/core/tests/caching.rs @@ -411,7 +411,7 @@ async fn responses_refetches_instead_of_deserializing_another_api_response( #[case] poisoned: Value, ) { use litellm_core::responses::route::Responses; - use litellm_types::responses::main::ResponsesApiResponse; + use litellm_llms_types::formats::responses::ResponsesApiResponse; let cache: Arc = Arc::new(InvalidEntryCache( ResponseCache::new(Arc::new(InMemoryCache::default())), @@ -461,7 +461,7 @@ async fn messages_cache_identity_includes_provider_native_parameters( #[case] changed: Value, ) { use litellm_core::messages::route::Messages; - use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; + use litellm_llms_types::formats::messages::MessagesResponse; let calls = AtomicUsize::new(0); for (value, expected_call) in [(original.clone(), 0), (changed, 1), (original, 0)] { @@ -487,7 +487,7 @@ async fn messages_cache_identity_includes_provider_native_parameters( None, || async { let call = calls.fetch_add(1, Ordering::SeqCst); - Ok(Box::new(serde_json::from_value::(json!({ + Ok(Box::new(serde_json::from_value::(json!({ "id":call.to_string(), "type":"message", "role":"assistant", "model":"test", "content":[{"type":"text","text":format!("answer {call}")}], "stop_reason":"end_turn", "stop_sequence":null @@ -843,7 +843,7 @@ async fn responses_cache_only_reuses_completed_responses( #[case] expected_calls: usize, ) { use litellm_core::responses::route::Responses; - use litellm_types::responses::main::ResponsesApiResponse; + use litellm_llms_types::formats::responses::ResponsesApiResponse; let calls = AtomicUsize::new(0); for _ in 0..2 { diff --git a/litellm-rust/crates/core/tests/chat_completions.rs b/litellm-rust/crates/core/tests/chat_completions.rs index fa9bd731809..b08fec41d3a 100644 --- a/litellm-rust/crates/core/tests/chat_completions.rs +++ b/litellm-rust/crates/core/tests/chat_completions.rs @@ -7,7 +7,7 @@ use std::time::Duration; use litellm_core::chat_completions::{Error, types::ChatCompletionsRequest}; use litellm_http::transport::Error as TransportError; -use litellm_types::utils::ChatCompletionsResponse; +use litellm_llms_types::formats::chat_completions::ChatCompletionsResponse; use rstest::{fixture, rstest}; use serde_json::{Map, Value, json}; use wiremock::ResponseTemplate; diff --git a/litellm-rust/crates/core/tests/messages/main.rs b/litellm-rust/crates/core/tests/messages/main.rs index 05e9aadd351..dd9689cf673 100644 --- a/litellm-rust/crates/core/tests/messages/main.rs +++ b/litellm-rust/crates/core/tests/messages/main.rs @@ -8,10 +8,8 @@ use litellm_core::messages::{ route::{Messages, MessagesMachine, MessagesOutput}, }; use litellm_http::{HttpSettings, Resolution}; +use litellm_llms_types::formats::messages::{MessagesRequest, MessagesResponse}; use litellm_secrets::source::SecretSource; -use litellm_types::llms::anthropic_messages::{ - anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse, -}; use rstest::fixture; use serde_json::{Map, Value, json}; use wiremock::ResponseTemplate; @@ -35,7 +33,7 @@ fn object(value: Value) -> Map { map } -fn body(value: Value) -> AnthropicMessagesRequest { +fn body(value: Value) -> MessagesRequest { serde_json::from_value(value).unwrap() } @@ -116,7 +114,7 @@ async fn run(call: MessagesCall) -> Result { run_with(Arc::new(RecordingSecrets::empty()), call).await } -async fn run_message(call: MessagesCall) -> AnthropicMessagesResponse { +async fn run_message(call: MessagesCall) -> MessagesResponse { match run(call).await.expect("messages call succeeds") { MessagesOutput::Complete(message) => *message, MessagesOutput::StreamEnded | MessagesOutput::Detached => { diff --git a/litellm-rust/crates/core/tests/messages/request.rs b/litellm-rust/crates/core/tests/messages/request.rs index b76895b9f1a..6a01be2b4f4 100644 --- a/litellm-rust/crates/core/tests/messages/request.rs +++ b/litellm-rust/crates/core/tests/messages/request.rs @@ -1,6 +1,8 @@ use litellm_llms::base_llm::messages::context::{MessagesModelCapabilities, SupportedEffortTiers}; -use litellm_types::llms::anthropic::{AnthropicBeta, BetaSet}; -use litellm_types::utils::{ProviderSpecificHeader, ProviderSpecificHeaders}; +use litellm_llms_types::{ + headers::{ProviderSpecificHeader, ProviderSpecificHeaders}, + providers::anthropic::{AnthropicBeta, BetaSet}, +}; use rstest::rstest; use super::*; diff --git a/litellm-rust/crates/core/tests/messages/response.rs b/litellm-rust/crates/core/tests/messages/response.rs index 7ef669599fb..42d3596e56a 100644 --- a/litellm-rust/crates/core/tests/messages/response.rs +++ b/litellm-rust/crates/core/tests/messages/response.rs @@ -1,4 +1,4 @@ -use litellm_core::messages::{MessagesResponse, messages_body}; +use litellm_core::messages::{MessagesCallResponse, messages_body}; use litellm_host::{ interceptors::{ExecutionFacts, ResultSource}, lifecycle::ExecutionEvent, @@ -33,7 +33,7 @@ async fn calls_defer_execution_until_polled( let request = host.request().unwrap(); let observer: Option = with_observer.then(|| host.events.0.sender.clone()); - let future: BoxFuture<'_, Result> = if with_hooks { + let future: BoxFuture<'_, Result> = if with_hooks { Box::pin(route.execute(request, &host, observer)) } else { Box::pin(route.execute(request, &(), observer)) @@ -43,7 +43,7 @@ async fn calls_defer_execution_until_polled( assert!(host.events.0.lock().unwrap().is_empty()); assert!(received(&upstream).await.is_empty()); - let MessagesResponse::Complete(response) = future.await.unwrap() else { + let MessagesCallResponse::Complete(response) = future.await.unwrap() else { panic!("expected a completed message"); }; assert_eq!( @@ -289,7 +289,7 @@ async fn the_facade_sends_through_the_injected_http_pool_configuration(call: Mes .await .expect("messages request succeeds"); - let MessagesResponse::Complete(message) = response else { + let MessagesCallResponse::Complete(message) = response else { panic!("a non-streaming request returns a message"); }; assert_eq!(message.id, "msg_1"); @@ -380,7 +380,8 @@ async fn builder_preserves_dependencies_and_optional_cache( api_base: Some(upstream.uri()), ..super::call() }; - let MessagesResponse::Complete(response) = route.execute(request, &(), None).await.unwrap() + let MessagesCallResponse::Complete(response) = + route.execute(request, &(), None).await.unwrap() else { panic!("expected a completed message"); }; diff --git a/litellm-rust/crates/core/tests/messages/stream.rs b/litellm-rust/crates/core/tests/messages/stream.rs index 0fb85920077..ad5ae5a8765 100644 --- a/litellm-rust/crates/core/tests/messages/stream.rs +++ b/litellm-rust/crates/core/tests/messages/stream.rs @@ -6,7 +6,7 @@ use std::{ use bytes::Bytes; use futures_util::{StreamExt, TryStreamExt}; use litellm_core::messages::{ - MessagesResponse, + MessagesCallResponse, route::{Messages, MessagesStreamHead}, }; use litellm_tracing::{Logger, Metadata, Record, Sink}; @@ -353,7 +353,7 @@ async fn the_sdk_returns_stream_headers_and_every_sse_byte( .await .unwrap(); - let MessagesResponse::Stream { head, chunks } = response else { + let MessagesCallResponse::Stream { head, chunks } = response else { panic!("a streaming request returns a stream"); }; for (name, value) in UPSTREAM_HEADERS { @@ -407,7 +407,7 @@ async fn dropping_the_sdk_stream_closes_the_unfinished_upstream( .expect("messages() returns before the upstream finishes") .unwrap(); - let MessagesResponse::Stream { mut chunks, .. } = response else { + let MessagesCallResponse::Stream { mut chunks, .. } = response else { panic!("a streaming request returns a stream"); }; if read_chunk { @@ -442,7 +442,7 @@ async fn the_sdk_yields_a_body_error_once_after_delivered_chunks(call: MessagesC .await .unwrap(); - let MessagesResponse::Stream { mut chunks, .. } = response else { + let MessagesCallResponse::Stream { mut chunks, .. } = response else { panic!("a streaming request returns a stream"); }; assert_eq!( diff --git a/litellm-rust/crates/core/tests/ocr/main.rs b/litellm-rust/crates/core/tests/ocr/main.rs index bf0752707fb..aa18dc8df24 100644 --- a/litellm-rust/crates/core/tests/ocr/main.rs +++ b/litellm-rust/crates/core/tests/ocr/main.rs @@ -9,11 +9,8 @@ use litellm_host::{ interceptors::{RequestContext, WireRequest}, lifecycle::CallEvent, }; -use litellm_llms::base_llm::ocr::{ - error::Error, - settings::OcrSettings, - transformation::{LiteLLMOcrResponse, OcrDocument}, -}; +use litellm_llms::base_llm::ocr::{error::Error, settings::OcrSettings}; +use litellm_llms_types::formats::ocr::{LiteLLMOcrResponse, OcrDocument}; use serde_json::{Map, Value, json}; use std::sync::Mutex; use wiremock::{MockServer, ResponseTemplate}; diff --git a/litellm-rust/crates/gateway-inference/Cargo.toml b/litellm-rust/crates/gateway-inference/Cargo.toml index e890679ff23..c854f0ea1ad 100644 --- a/litellm-rust/crates/gateway-inference/Cargo.toml +++ b/litellm-rust/crates/gateway-inference/Cargo.toml @@ -19,7 +19,7 @@ litellm-http.workspace = true litellm-llms.workspace = true litellm-router.workspace = true litellm-secrets.workspace = true -litellm-types.workspace = true +litellm-llms-types.workspace = true serde.workspace = true serde_json.workspace = true thiserror.workspace = true diff --git a/litellm-rust/crates/gateway-inference/src/messages.rs b/litellm-rust/crates/gateway-inference/src/messages.rs index 6be5921ac6d..0f1d1d31689 100644 --- a/litellm-rust/crates/gateway-inference/src/messages.rs +++ b/litellm-rust/crates/gateway-inference/src/messages.rs @@ -12,7 +12,7 @@ use axum::{ }; use litellm_core::messages::{MessagesCall, messages_body, route::Messages}; use litellm_host_http::Sse; -use litellm_types::utils::{ProviderSpecificHeader, ProviderSpecificHeaders}; +use litellm_llms_types::headers::{ProviderSpecificHeader, ProviderSpecificHeaders}; use serde_json::{Map, Value}; use crate::{Deployment, Error, Gateway, JsonObject, RequestId, request}; diff --git a/litellm-rust/crates/gateway-inference/src/ocr.rs b/litellm-rust/crates/gateway-inference/src/ocr.rs index 3a62e006dce..2d8f62887e5 100644 --- a/litellm-rust/crates/gateway-inference/src/ocr.rs +++ b/litellm-rust/crates/gateway-inference/src/ocr.rs @@ -4,7 +4,8 @@ use std::sync::Arc; use axum::{Json, extract::State, http::HeaderMap, response::IntoResponse}; use litellm_auth::SecretValue; use litellm_core::ocr::types::{LiteLLMOcrRequest, OcrConnectionInputs, OcrDocumentInput}; -use litellm_llms::base_llm::ocr::transformation::OcrDocument; +use litellm_llms::base_llm::ocr::transformation::decode_request_value; +use litellm_llms_types::formats::ocr::OcrDocument; use serde_json::Value; use crate::{ @@ -42,7 +43,11 @@ async fn handle( file_name: upload.file_name, mime_type: upload.mime_type, }, - None => OcrDocument::try_from(body.get("document").cloned().unwrap_or_default())?.into(), + None => decode_request_value::( + body.get("document").cloned().unwrap_or_default(), + "document", + )? + .into(), }; let format = body .get("req_format") diff --git a/litellm-rust/crates/gateway-inference/tests/ocr.rs b/litellm-rust/crates/gateway-inference/tests/ocr.rs index dba8099b685..fbe523addf0 100644 --- a/litellm-rust/crates/gateway-inference/tests/ocr.rs +++ b/litellm-rust/crates/gateway-inference/tests/ocr.rs @@ -2,7 +2,8 @@ mod support; use axum::{body::Body, http::Request}; use litellm_gateway_inference::Error; -use litellm_llms::base_llm::ocr::{error::Error as OcrError, transformation::OcrDocument}; +use litellm_llms::base_llm::ocr::{error::Error as OcrError, transformation::decode_request_value}; +use litellm_llms_types::formats::ocr::OcrDocument; use rstest::rstest; use serde_json::{Value, json}; use tower::ServiceExt; @@ -146,7 +147,7 @@ async fn malformed_multipart_uses_an_openai_error_envelope( #[rstest] #[case::missing_document( "/v1/ocr", "mistral/test-ocr", "", - Error::Ocr(OcrDocument::try_from(Value::Null).unwrap_err()), + Error::Ocr(decode_request_value::(Value::Null, "document").unwrap_err()), )] #[case::empty_document( "/v1/ocr", diff --git a/litellm-rust/crates/types/AGENTS.md b/litellm-rust/crates/llms-types/AGENTS.md similarity index 81% rename from litellm-rust/crates/types/AGENTS.md rename to litellm-rust/crates/llms-types/AGENTS.md index 4b0792a6316..8590050283c 100644 --- a/litellm-rust/crates/types/AGENTS.md +++ b/litellm-rust/crates/llms-types/AGENTS.md @@ -1,15 +1,23 @@ The same ownership rule applies to Messages, Responses, Chat Completions, OCR, and other API formats. This crate owns their shared API data contracts. Adapter contracts and shared transformation machinery belong in `llms/src/base_llm//`, provider policy in `llms/src///`, and call orchestration in `core/src//`. A provider originating a format, or several providers using a type, does not change these responsibilities. Existing model locations outside this crate are not exceptions to this rule for new shared API contracts -- `litellm-types` owns shared API data contracts and their serialization +- `litellm-llms-types` owns shared API data contracts and their serialization - A type belongs here when it describes a request, response, event, or value that consumers must agree on independently of how a call executes - Being public, serializable, or used by several crates is not sufficient - These are intended boundaries, not a claim that every existing item follows them -- Organize public contracts by API format: `messages`, `chat_completions`, and `responses` - - Use names such as `litellm_types::messages::MessagesRequest`, without an Anthropic prefix solely because Anthropic designed Messages - - Existing `llms::openai`, `llms::anthropic_messages`, and chat types under `utils` are legacy locations, not patterns for new modules +- Organize public API contracts under `formats`: `messages`, `chat_completions`, `responses`, `ocr`, `audio_transcription`, and `batches` + - Use names such as `litellm_llms_types::formats::messages::MessagesRequest`, without an Anthropic prefix solely because Anthropic designed Messages - Keep one canonical definition and import path when moving a contract, updating consumers together instead of adding duplicate models or compatibility re-exports +- Keep shared provider-specific wire types and extensions under `providers` + - Provider types may reuse format types; format types must not depend on provider types + - A field belonging to an API format stays under `formats` even when provider support varies. Including it in a type does not promise provider support + - Add a typed provider extension when a consumer needs to interpret or construct it. Keep adapter-only projections in `llms` until a shared public data contract is needed + - Keep one authoritative representation of each field, preserving unknown fields without duplicating typed values in an extension map + - Provider capability checks, defaults, authentication, header selection, and transformations remain in `llms` + +- Keep format-independent data helpers such as `headers`, `recognized`, and `serde_compat` at the crate root + - Shared request/response bodies, message and content-block enums, usage records, tool-call chunks, stream-event payloads, and protocol error bodies belong here - This includes LiteLLM's normalized response contracts and extensions, not just exact upstream schemas - `ChatCompletionsResponse` currently represents the response handed to the host, so replacing it with a supposedly more complete upstream schema must not silently change that contract @@ -31,6 +39,7 @@ The same ownership rule applies to Messages, Responses, Chat Completions, OCR, a - Provider config traits, `MessagesTransformContext`, `MessagesModelCapabilities`, `ThinkingBudgets`, `StreamShape`, and transformer state belong in `llms` - Catalog records and pricing belong in `model-catalog`, which may reuse wire enums such as `ReasoningEffort` - Host hooks, Python objects, credentials, clients, timeouts, and routing decisions do not become API payload types merely because they cross a crate boundary + - Legacy logging operation selection belongs in `callbacks-legacy-python`, not this crate - Stream-event data belongs here, but live streams, decoders, framing, buffering, and stream lifecycle decisions do not - Keep SSE and AWS framing in `framer`, provider decoding and conversion in `llms`, and call orchestration in `core` diff --git a/litellm-rust/crates/types/Cargo.toml b/litellm-rust/crates/llms-types/Cargo.toml similarity index 77% rename from litellm-rust/crates/types/Cargo.toml rename to litellm-rust/crates/llms-types/Cargo.toml index e356c8e127d..2d880b87faf 100644 --- a/litellm-rust/crates/types/Cargo.toml +++ b/litellm-rust/crates/llms-types/Cargo.toml @@ -1,5 +1,5 @@ [package] -name = "litellm-types" +name = "litellm-llms-types" version = "0.1.0" edition.workspace = true license.workspace = true @@ -9,9 +9,11 @@ repository.workspace = true schema = ["dep:schemars"] [dependencies] +macro_rules_attribute.workspace = true schemars = { workspace = true, optional = true } serde.workspace = true serde_json.workspace = true +serde_with.workspace = true strum.workspace = true [dev-dependencies] diff --git a/litellm-rust/crates/types/src/audio_transcription.rs b/litellm-rust/crates/llms-types/src/formats/audio_transcription.rs similarity index 72% rename from litellm-rust/crates/types/src/audio_transcription.rs rename to litellm-rust/crates/llms-types/src/formats/audio_transcription.rs index 151c3a9d098..e00ecb0b5fb 100644 --- a/litellm-rust/crates/types/src/audio_transcription.rs +++ b/litellm-rust/crates/llms-types/src/formats/audio_transcription.rs @@ -1,7 +1,6 @@ -use serde::{Deserialize, Serialize}; use serde_json::Value; -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] pub struct AudioTranscriptionResponseData { pub text: String, } diff --git a/litellm-rust/crates/llms-types/src/formats/batches.rs b/litellm-rust/crates/llms-types/src/formats/batches.rs new file mode 100644 index 00000000000..9749b042a36 --- /dev/null +++ b/litellm-rust/crates/llms-types/src/formats/batches.rs @@ -0,0 +1,36 @@ +#[macro_rules_attribute::apply(wire_type)] +#[derive(Copy, Eq)] +#[serde(rename_all = "snake_case")] +pub enum BatchStatus { + InProgress, + Cancelling, + Completed, +} + +#[macro_rules_attribute::apply(wire_type)] +#[derive(Eq)] +pub struct BatchRequestCounts { + pub total: u64, + pub completed: u64, + pub failed: u64, +} + +#[macro_rules_attribute::apply(wire_type)] +#[derive(Eq)] +pub struct BatchResponse { + pub id: String, + pub object: String, + pub endpoint: String, + pub input_file_id: String, + pub completion_window: String, + pub status: BatchStatus, + pub output_file_id: String, + pub created_at: i64, + pub in_progress_at: Option, + pub expires_at: Option, + pub completed_at: Option, + pub expired_at: Option, + pub cancelling_at: Option, + pub cancelled_at: Option, + pub request_counts: BatchRequestCounts, +} diff --git a/litellm-rust/crates/llms-types/src/formats/chat_completions.rs b/litellm-rust/crates/llms-types/src/formats/chat_completions.rs new file mode 100644 index 00000000000..31b5046469a --- /dev/null +++ b/litellm-rust/crates/llms-types/src/formats/chat_completions.rs @@ -0,0 +1,223 @@ +use serde_json::{Map, Value}; +use strum::IntoStaticStr; + +/// Reasoning effort level accepted or applied by the model. +#[macro_rules_attribute::apply(wire_type)] +#[derive(Copy, Eq, IntoStaticStr)] +#[serde(rename_all = "snake_case")] +#[strum(serialize_all = "snake_case")] +pub enum ReasoningEffort { + None, + Minimal, + Low, + Medium, + High, + Xhigh, + Max, +} + +impl ReasoningEffort { + pub const ALL: [Self; 7] = [ + Self::None, + Self::Minimal, + Self::Low, + Self::Medium, + Self::High, + Self::Xhigh, + Self::Max, + ]; + + pub fn as_str(self) -> &'static str { + self.into() + } + + pub fn parse(value: &str) -> Option { + Self::ALL + .into_iter() + .find(|effort| effort.as_str() == value) + } +} + +#[macro_rules_attribute::apply(wire_type)] +#[serde(untagged)] +pub enum ChatMessageContent { + Text(String), + Parts(Vec), +} + +#[macro_rules_attribute::apply(wire_type)] +pub struct ChatMessage { + pub role: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub content: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub name: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(wire_type)] +pub struct ChatCompletionToolCallFunctionChunk { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub name: Option, + pub arguments: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub provider_specific_fields: Option>, +} + +#[macro_rules_attribute::apply(wire_type)] +pub struct ChatCompletionToolCallChunk { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub id: Option, + #[serde(rename = "type")] + pub tool_type: String, + pub function: ChatCompletionToolCallFunctionChunk, + pub index: i64, +} + +#[macro_rules_attribute::apply(wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ChatCompletionThinkingBlock { + Thinking { + #[serde(default, skip_serializing_if = "Option::is_none")] + thinking: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + signature: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + cache_control: Option, + }, + RedactedThinking { + #[serde(default, skip_serializing_if = "Option::is_none")] + data: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + cache_control: Option, + }, +} + +/// OpenAI `usage`, including the `prompt_tokens_details` split LiteLLM's Python +/// path reports so cost tracking sees the same numbers on either path. +#[macro_rules_attribute::apply(wire_type)] +#[derive(Default)] +pub struct PromptTokensDetails { + pub cached_tokens: u64, + pub cache_creation_tokens: u64, + pub text_tokens: u64, +} + +#[macro_rules_attribute::apply(wire_type)] +#[derive(Default)] +pub struct ChatCompletionsUsage { + pub prompt_tokens: u64, + pub completion_tokens: u64, + pub total_tokens: u64, + pub prompt_tokens_details: PromptTokensDetails, +} + +#[macro_rules_attribute::apply(wire_type)] +pub struct ChatCompletionsChoiceMessage { + pub role: String, + // Whether an empty turn is `None` or `""` is the provider's choice, not a + // shared invariant: Anthropic's transform ends on `merged_text or None` + // while Converse assigns the joined string unconditionally. Each config + // mirrors its own, so keep this optional and serialize it even when None. + pub content: Option, +} + +#[macro_rules_attribute::apply(wire_type)] +pub struct ChatCompletionsChoice { + pub index: u64, + pub message: ChatCompletionsChoiceMessage, + pub finish_reason: String, +} + +/// The normalized response handed back to the host. +/// +/// There is deliberately no `id`: Python mints the `chatcmpl-…` id on the +/// `ModelResponse` it already created, and echoing the provider's own id here +/// would change it. Pinned by `response_carries_no_id` in the Anthropic chat transformation tests. +#[macro_rules_attribute::apply(wire_type)] +pub struct ChatCompletionsResponse { + pub created: u64, + pub model: String, + pub choices: Vec, + pub usage: ChatCompletionsUsage, +} + +#[macro_rules_attribute::apply(wire_type)] +#[derive(Default)] +pub struct ChatCompletionDelta { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub content: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub role: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub tool_calls: Option>, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub reasoning_content: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub thinking_blocks: Option>, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub provider_specific_fields: Option>, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(wire_type)] +pub struct ChatCompletionStreamingChoice { + pub index: u64, + pub delta: ChatCompletionDelta, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub finish_reason: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub logprobs: Option, +} + +#[macro_rules_attribute::apply(wire_type)] +pub struct ChatCompletionChunk { + pub id: String, + pub created: u64, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub model: Option, + pub object: String, + pub choices: Vec, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub usage: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub provider_specific_fields: Option>, +} + +#[cfg(test)] +mod tests { + use rstest::rstest; + + use super::*; + + #[rstest] + fn reasoning_effort_names_match_the_wire_and_parse_back( + #[values( + ReasoningEffort::None, + ReasoningEffort::Minimal, + ReasoningEffort::Low, + ReasoningEffort::Medium, + ReasoningEffort::High, + ReasoningEffort::Xhigh, + ReasoningEffort::Max + )] + effort: ReasoningEffort, + ) { + assert_eq!( + serde_json::to_value(effort).unwrap(), + Value::String(effort.as_str().to_string()) + ); + assert_eq!(ReasoningEffort::parse(effort.as_str()), Some(effort)); + assert!(ReasoningEffort::ALL.contains(&effort)); + } + + #[rstest] + #[case::unknown("ultra")] + #[case::uppercase("HIGH")] + #[case::empty("")] + fn reasoning_effort_parse_rejects(#[case] value: &str) { + assert_eq!(ReasoningEffort::parse(value), None); + } +} diff --git a/litellm-rust/crates/types/src/messages/AGENTS.md b/litellm-rust/crates/llms-types/src/formats/messages/AGENTS.md similarity index 100% rename from litellm-rust/crates/types/src/messages/AGENTS.md rename to litellm-rust/crates/llms-types/src/formats/messages/AGENTS.md diff --git a/litellm-rust/crates/llms-types/src/formats/messages/mod.rs b/litellm-rust/crates/llms-types/src/formats/messages/mod.rs new file mode 100644 index 00000000000..219e0ae63a0 --- /dev/null +++ b/litellm-rust/crates/llms-types/src/formats/messages/mod.rs @@ -0,0 +1,11 @@ +mod request; +mod response; +pub mod streaming; + +pub use request::{ + AdaptiveThinking, CacheControl, ContentBlock, ContentBlockType, ContextEdit, ContextManagement, + DisabledThinking, EffortLevel, EnabledThinking, Message, MessageContent, + MessagesOptionalParams, MessagesRequest, MessagesTool, OutputConfig, Speed, SystemPrompt, + ThinkingConfig, ThinkingDisplay, +}; +pub use response::MessagesResponse; diff --git a/litellm-rust/crates/types/src/llms/anthropic_messages/anthropic_request.rs b/litellm-rust/crates/llms-types/src/formats/messages/request.rs similarity index 89% rename from litellm-rust/crates/types/src/llms/anthropic_messages/anthropic_request.rs rename to litellm-rust/crates/llms-types/src/formats/messages/request.rs index 118848bee0a..d14e9afd0c3 100644 --- a/litellm-rust/crates/types/src/llms/anthropic_messages/anthropic_request.rs +++ b/litellm-rust/crates/llms-types/src/formats/messages/request.rs @@ -1,26 +1,25 @@ -use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; use strum::IntoStaticStr; -use crate::{llms::openai::ReasoningEffort, recognized::Recognized}; +use crate::formats::chat_completions::ReasoningEffort; +use crate::recognized::Recognized; -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] #[serde(untagged)] pub enum SystemPrompt { Text(String), Blocks(Vec), } -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] #[serde(untagged)] pub enum MessageContent { Text(String), Blocks(Vec), } -#[derive( - Clone, Debug, PartialEq, Eq, Serialize, Deserialize, strum::Display, strum::EnumString, -)] +#[macro_rules_attribute::apply(wire_type)] +#[derive(Eq, strum::Display, strum::EnumString)] #[serde(from = "String", into = "String")] #[strum(serialize_all = "snake_case")] pub enum ContentBlockType { @@ -49,7 +48,8 @@ impl From for String { } } -#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] +#[derive(Default)] pub struct ContentBlock { #[serde(rename = "type", default, skip_serializing_if = "Option::is_none")] pub block_type: Option, @@ -93,7 +93,8 @@ impl ContentBlock { } } -#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] +#[derive(Default)] pub struct CacheControl { #[serde(rename = "type", skip_serializing_if = "Option::is_none")] pub cache_type: Option, @@ -105,15 +106,16 @@ pub struct CacheControl { pub extra: Map, } -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -pub struct AnthropicMessage { +#[macro_rules_attribute::apply(wire_type)] +pub struct Message { pub role: String, pub content: MessageContent, #[serde(flatten)] pub extra: Map, } -#[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq, Hash, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] +#[derive(Copy, Hash, IntoStaticStr, Eq)] #[serde(rename_all = "lowercase")] #[strum(serialize_all = "lowercase")] pub enum EffortLevel { @@ -142,7 +144,8 @@ impl From for ReasoningEffort { } } -#[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] +#[derive(Copy, IntoStaticStr, Eq)] #[serde(rename_all = "lowercase")] #[strum(serialize_all = "lowercase")] pub enum Speed { @@ -158,9 +161,9 @@ impl Speed { /// The tools whose presence changes how the request is sent. Every other tool, custom or /// server, deserializes as `Recognized::Unrecognized` and passes through verbatim. -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] #[serde(tag = "type")] -pub enum AnthropicTool { +pub enum MessagesTool { #[serde(rename = "advisor_20260301")] Advisor { #[serde(flatten)] @@ -178,7 +181,7 @@ pub enum AnthropicTool { }, } -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] #[serde(tag = "type")] pub enum ContextEdit { #[serde(rename = "compact_20260112")] @@ -198,7 +201,8 @@ pub enum ContextEdit { }, } -#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] +#[derive(Default)] pub struct ContextManagement { #[serde(default, skip_serializing_if = "Option::is_none")] pub edits: Option>>, @@ -206,7 +210,8 @@ pub struct ContextManagement { pub extra: Map, } -#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] +#[derive(Default)] pub struct OutputConfig { #[serde(default, skip_serializing_if = "Option::is_none")] pub effort: Option>, @@ -222,7 +227,8 @@ impl OutputConfig { } } -#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] +#[derive(Copy, Eq)] #[serde(rename_all = "lowercase")] pub enum ThinkingDisplay { Summarized, @@ -230,7 +236,8 @@ pub enum ThinkingDisplay { Updates, } -#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] +#[derive(Default)] pub struct EnabledThinking { #[serde(default, skip_serializing_if = "Option::is_none")] pub budget_tokens: Option>, @@ -240,7 +247,8 @@ pub struct EnabledThinking { pub extra: Map, } -#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] +#[derive(Default)] pub struct AdaptiveThinking { #[serde(default, skip_serializing_if = "Option::is_none")] pub display: Option>, @@ -248,13 +256,14 @@ pub struct AdaptiveThinking { pub extra: Map, } -#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] +#[derive(Default)] pub struct DisabledThinking { #[serde(flatten)] pub extra: Map, } -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] #[serde(tag = "type", rename_all = "lowercase")] pub enum ThinkingConfig { Enabled(EnabledThinking), @@ -278,16 +287,17 @@ impl ThinkingConfig { } } -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -pub struct AnthropicMessagesRequest { +#[macro_rules_attribute::apply(wire_type)] +pub struct MessagesRequest { pub model: String, - pub messages: Vec, + pub messages: Vec, #[serde(flatten)] - pub params: AnthropicMessagesOptionalParams, + pub params: MessagesOptionalParams, } -#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] -pub struct AnthropicMessagesOptionalParams { +#[macro_rules_attribute::apply(wire_type)] +#[derive(Default)] +pub struct MessagesOptionalParams { #[serde(skip_serializing_if = "Option::is_none")] pub max_tokens: Option, #[serde(skip_serializing_if = "Option::is_none")] @@ -305,7 +315,7 @@ pub struct AnthropicMessagesOptionalParams { #[serde(skip_serializing_if = "Option::is_none")] pub top_k: Option, #[serde(skip_serializing_if = "Option::is_none")] - pub tools: Option>>, + pub tools: Option>>, #[serde(skip_serializing_if = "Option::is_none")] pub tool_choice: Option, #[serde(skip_serializing_if = "Option::is_none")] @@ -334,7 +344,7 @@ pub struct AnthropicMessagesOptionalParams { pub extra: Map, } -impl AnthropicMessage { +impl Message { pub fn blocks(&self) -> &[ContentBlock] { match &self.content { MessageContent::Blocks(blocks) => blocks, @@ -357,7 +367,7 @@ mod tests { use super::*; - fn round_trip(value: &Value) -> Value { + fn round_trip(value: &Value) -> Value { let parsed: T = serde_json::from_value(value.clone()).unwrap(); serde_json::to_value(parsed).unwrap() } @@ -396,7 +406,7 @@ mod tests { "stream": true, "safeguards": [{"type": "dangerous_tool_use"}] }); - let request: AnthropicMessagesRequest = serde_json::from_value(body.clone()).unwrap(); + let request: MessagesRequest = serde_json::from_value(body.clone()).unwrap(); assert_eq!( ( @@ -432,7 +442,7 @@ mod tests { #[case] message: Value, #[case] expected: Vec, ) { - let message: AnthropicMessage = serde_json::from_value(message).unwrap(); + let message: Message = serde_json::from_value(message).unwrap(); assert_eq!(message.blocks(), expected.as_slice()); } @@ -440,7 +450,7 @@ mod tests { #[case::replaces_string_content(json!({"role": "assistant", "content": "old", "name": "kept"}))] #[case::replaces_block_content(json!({"role": "assistant", "content": [{"type": "text", "text": "old"}], "name": "kept"}))] fn with_blocks_replaces_content_and_keeps_the_rest(#[case] message: Value) { - let message: AnthropicMessage = serde_json::from_value(message).unwrap(); + let message: Message = serde_json::from_value(message).unwrap(); assert_eq!( serde_json::to_value(message.with_blocks(vec![ContentBlock::text("new")])).unwrap(), json!({"role": "assistant", "content": [{"type": "text", "text": "new"}], "name": "kept"}) @@ -506,7 +516,7 @@ mod tests { "context_management": [{"type": "compaction", "compact_threshold": 5}] }))] fn request_round_trips_unchanged(#[case] request: Value) { - assert_eq!(round_trip::(&request), request); + assert_eq!(round_trip::(&request), request); } #[rstest] @@ -555,15 +565,15 @@ mod tests { #[rstest] #[case::advisor( json!({"type": "advisor_20260301", "name": "advisor"}), - Recognized::Known(AnthropicTool::Advisor { extra: Map::from_iter([("name".to_string(), json!("advisor"))]) }) + Recognized::Known(MessagesTool::Advisor { extra: Map::from_iter([("name".to_string(), json!("advisor"))]) }) )] #[case::regex_tool_search( json!({"type": "tool_search_tool_regex_20251119"}), - Recognized::Known(AnthropicTool::ToolSearchRegex { extra: Map::new() }) + Recognized::Known(MessagesTool::ToolSearchRegex { extra: Map::new() }) )] #[case::bm25_tool_search( json!({"type": "tool_search_tool_bm25_20251119"}), - Recognized::Known(AnthropicTool::ToolSearchBm25 { extra: Map::new() }) + Recognized::Known(MessagesTool::ToolSearchBm25 { extra: Map::new() }) )] #[case::custom_tool_without_a_type( json!({"name": "advisor", "input_schema": {}}), @@ -576,10 +586,10 @@ mod tests { #[case::not_an_object(json!("advisor_20260301"), Recognized::Unrecognized(json!("advisor_20260301")))] fn tools_are_recognized_by_their_exact_type( #[case] tool: Value, - #[case] expected: Recognized, + #[case] expected: Recognized, ) { assert_eq!( - serde_json::from_value::>(tool).unwrap(), + serde_json::from_value::>(tool).unwrap(), expected ); } diff --git a/litellm-rust/crates/types/src/llms/anthropic_messages/anthropic_response.rs b/litellm-rust/crates/llms-types/src/formats/messages/response.rs similarity index 92% rename from litellm-rust/crates/types/src/llms/anthropic_messages/anthropic_response.rs rename to litellm-rust/crates/llms-types/src/formats/messages/response.rs index 0a2653f352f..2d8e1c054fa 100644 --- a/litellm-rust/crates/types/src/llms/anthropic_messages/anthropic_response.rs +++ b/litellm-rust/crates/llms-types/src/formats/messages/response.rs @@ -1,8 +1,7 @@ -use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -pub struct AnthropicMessagesResponse { +#[macro_rules_attribute::apply(wire_type)] +pub struct MessagesResponse { pub id: String, #[serde(rename = "type")] pub message_type: String, @@ -31,8 +30,8 @@ mod tests { stop_sequence: Option<&str>, usage: Option, container: Option, - ) -> AnthropicMessagesResponse { - AnthropicMessagesResponse { + ) -> MessagesResponse { + MessagesResponse { id: "msg_1".to_string(), message_type: "message".to_string(), role: "assistant".to_string(), diff --git a/litellm-rust/crates/types/src/messages/streaming.rs b/litellm-rust/crates/llms-types/src/formats/messages/streaming.rs similarity index 89% rename from litellm-rust/crates/types/src/messages/streaming.rs rename to litellm-rust/crates/llms-types/src/formats/messages/streaming.rs index f77fdb01aa3..abdcfa26a8c 100644 --- a/litellm-rust/crates/types/src/messages/streaming.rs +++ b/litellm-rust/crates/llms-types/src/formats/messages/streaming.rs @@ -1,7 +1,7 @@ -use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; -#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] +#[derive(Default)] pub struct MessagesStreamUsage { #[serde(default, skip_serializing_if = "Option::is_none")] pub input_tokens: Option, @@ -17,7 +17,7 @@ pub struct MessagesStreamUsage { pub extra: Map, } -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] pub struct MessagesStreamMessage { pub id: String, #[serde(rename = "type")] @@ -32,7 +32,7 @@ pub struct MessagesStreamMessage { pub extra: Map, } -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] #[serde(tag = "type", rename_all = "snake_case")] pub enum MessagesContentBlockDelta { TextDelta { @@ -56,7 +56,7 @@ pub enum MessagesContentBlockDelta { }, } -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] pub struct MessagesContentBlock { #[serde(rename = "type")] pub block_type: String, @@ -82,7 +82,8 @@ pub struct MessagesContentBlock { pub extra: Map, } -#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] +#[derive(Default)] pub struct MessagesDelta { #[serde(default, skip_serializing_if = "Option::is_none")] pub stop_reason: Option, @@ -96,7 +97,7 @@ pub struct MessagesDelta { pub extra: Map, } -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] pub struct MessagesStreamError { #[serde(rename = "type")] pub error_type: String, @@ -107,7 +108,7 @@ pub struct MessagesStreamError { pub extra: Map, } -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] #[serde(tag = "type", rename_all = "snake_case")] pub enum MessagesStreamEvent { MessageStart { diff --git a/litellm-rust/crates/llms-types/src/formats/mod.rs b/litellm-rust/crates/llms-types/src/formats/mod.rs new file mode 100644 index 00000000000..53f2577090b --- /dev/null +++ b/litellm-rust/crates/llms-types/src/formats/mod.rs @@ -0,0 +1,6 @@ +pub mod audio_transcription; +pub mod batches; +pub mod chat_completions; +pub mod messages; +pub mod ocr; +pub mod responses; diff --git a/litellm-rust/crates/llms-types/src/formats/ocr.rs b/litellm-rust/crates/llms-types/src/formats/ocr.rs new file mode 100644 index 00000000000..b491f5d82f1 --- /dev/null +++ b/litellm-rust/crates/llms-types/src/formats/ocr.rs @@ -0,0 +1,152 @@ +use std::collections::BTreeMap; + +use serde_json::{Map, Value}; +use serde_with::serde_as; + +use crate::serde_compat::{FiniteF64, LaxI64}; + +#[macro_rules_attribute::apply(wire_type)] +#[serde(tag = "type")] +pub enum OcrDocument { + #[serde(rename = "document_url")] + DocumentUrl { + document_url: String, + #[serde(flatten)] + extra_fields: BTreeMap>, + }, + #[serde(rename = "image_url")] + ImageUrl { + image_url: String, + #[serde(flatten)] + extra_fields: BTreeMap>, + }, +} + +impl OcrDocument { + pub fn source(&self) -> &str { + match self { + Self::DocumentUrl { document_url, .. } => document_url, + Self::ImageUrl { image_url, .. } => image_url, + } + } + + pub fn is_remote(&self) -> bool { + let source = self.source(); + source.starts_with("http://") || source.starts_with("https://") + } + + pub fn with_source(self, source: String) -> Self { + match self { + Self::DocumentUrl { extra_fields, .. } => Self::DocumentUrl { + document_url: source, + extra_fields, + }, + Self::ImageUrl { extra_fields, .. } => Self::ImageUrl { + image_url: source, + extra_fields, + }, + } + } +} + +#[macro_rules_attribute::apply(wire_type)] +#[derive(Copy, Default, Eq)] +#[serde(rename_all = "lowercase")] +pub enum OcrResponseFormat { + #[default] + Litellm, + Native, +} + +#[serde_as] +#[macro_rules_attribute::apply(wire_type)] +#[derive(Default)] +pub struct OcrPageDimensions { + #[serde_as(deserialize_as = "Option")] + pub dpi: Option, + #[serde_as(deserialize_as = "Option")] + pub height: Option, + #[serde_as(deserialize_as = "Option")] + pub width: Option, +} + +#[macro_rules_attribute::apply(wire_type)] +#[derive(Default)] +pub struct OcrPageImage { + pub image_base64: Option, + pub bbox: Option>, + #[serde(flatten)] + pub extra_fields: Map, +} + +#[serde_as] +#[macro_rules_attribute::apply(wire_type)] +#[derive(Default)] +pub struct OcrPage { + #[serde_as(deserialize_as = "LaxI64")] + pub index: i64, + pub markdown: String, + pub images: Option>, + pub dimensions: Option, + #[serde(flatten)] + pub extra_fields: Map, +} + +#[serde_as] +#[macro_rules_attribute::apply(wire_type)] +#[derive(Default)] +pub struct OcrUsageInfo { + #[serde_as(deserialize_as = "Option")] + pub pages_processed: Option, + #[serde_as(deserialize_as = "Option")] + pub pages_processed_annotation: Option, + #[serde_as(deserialize_as = "Option")] + pub credits: Option, + #[serde_as(deserialize_as = "Option")] + pub doc_size_bytes: Option, + #[serde(flatten)] + pub extra_fields: Map, +} + +#[macro_rules_attribute::apply(wire_type)] +pub struct LiteLLMOcrResponse { + pub pages: Vec, + pub model: String, + pub document_annotation: Option, + pub usage_info: Option, + pub content: Option, + pub tables: Option>>, + #[serde(rename = "keyValuePairs")] + pub key_value_pairs: Option>>, + #[serde(default = "ocr_object")] + pub object: String, + #[serde(flatten)] + pub extra_fields: Map, + #[serde(skip_serializing_if = "Option::is_none")] + pub provider_native_response: Option>, +} + +impl LiteLLMOcrResponse { + pub fn new(model: impl Into, pages: Vec) -> Self { + Self { + pages, + model: model.into(), + document_annotation: None, + usage_info: None, + content: None, + tables: None, + key_value_pairs: None, + object: ocr_object(), + extra_fields: Map::new(), + provider_native_response: None, + } + } + + pub fn into_json(self) -> Value { + serde_json::to_value(self).expect("OCR response fields are JSON-compatible") + } +} + +fn ocr_object() -> String { + "ocr".into() +} diff --git a/litellm-rust/crates/llms-types/src/formats/responses/mod.rs b/litellm-rust/crates/llms-types/src/formats/responses/mod.rs new file mode 100644 index 00000000000..0aefd8a8698 --- /dev/null +++ b/litellm-rust/crates/llms-types/src/formats/responses/mod.rs @@ -0,0 +1,4 @@ +mod response; +pub mod streaming_websocket; + +pub use response::ResponsesApiResponse; diff --git a/litellm-rust/crates/types/src/responses/main.rs b/litellm-rust/crates/llms-types/src/formats/responses/response.rs similarity index 67% rename from litellm-rust/crates/types/src/responses/main.rs rename to litellm-rust/crates/llms-types/src/formats/responses/response.rs index 548dcd8d75e..7017d0fa4e4 100644 --- a/litellm-rust/crates/types/src/responses/main.rs +++ b/litellm-rust/crates/llms-types/src/formats/responses/response.rs @@ -1,7 +1,6 @@ -use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] pub struct ResponsesApiResponse { pub id: String, pub model: String, diff --git a/litellm-rust/crates/types/src/responses/streaming_websocket.rs b/litellm-rust/crates/llms-types/src/formats/responses/streaming_websocket.rs similarity index 94% rename from litellm-rust/crates/types/src/responses/streaming_websocket.rs rename to litellm-rust/crates/llms-types/src/formats/responses/streaming_websocket.rs index cee1e4f0c03..75858b45223 100644 --- a/litellm-rust/crates/types/src/responses/streaming_websocket.rs +++ b/litellm-rust/crates/llms-types/src/formats/responses/streaming_websocket.rs @@ -2,6 +2,8 @@ use serde::{Deserialize, Deserializer, Serialize, Serializer}; use serde_json::{Map, Value}; #[derive(Clone, Debug, PartialEq, Eq, strum::AsRefStr)] +#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] +#[cfg_attr(feature = "schema", schemars(with = "String"))] pub enum ResponsesWsEventType { #[strum(serialize = "response.create")] ResponseCreate, @@ -52,7 +54,7 @@ impl<'de> Deserialize<'de> for ResponsesWsEventType { } } -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] pub struct ResponsesWsEvent { #[serde(rename = "type")] pub event_type: ResponsesWsEventType, @@ -78,7 +80,8 @@ impl ResponsesWsEvent { } } -#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] +#[derive(Eq)] pub struct ResponsesErrorFrame { #[serde(rename = "type")] pub frame_type: &'static str, @@ -97,7 +100,8 @@ impl ResponsesErrorFrame { } } -#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] +#[derive(Eq)] pub struct ResponsesErrorBody { #[serde(rename = "type")] pub error_type: &'static str, diff --git a/litellm-rust/crates/llms-types/src/headers.rs b/litellm-rust/crates/llms-types/src/headers.rs new file mode 100644 index 00000000000..bf4f42b493d --- /dev/null +++ b/litellm-rust/crates/llms-types/src/headers.rs @@ -0,0 +1,17 @@ +use serde_json::{Map, Value}; + +#[macro_rules_attribute::apply(wire_type)] +#[derive(Default)] +pub struct ProviderSpecificHeader { + #[serde(default)] + pub custom_llm_provider: String, + #[serde(default)] + pub extra_headers: Map, +} + +#[macro_rules_attribute::apply(wire_type)] +#[serde(untagged)] +pub enum ProviderSpecificHeaders { + One(ProviderSpecificHeader), + Many(Vec), +} diff --git a/litellm-rust/crates/llms-types/src/lib.rs b/litellm-rust/crates/llms-types/src/lib.rs new file mode 100644 index 00000000000..116c11c0f88 --- /dev/null +++ b/litellm-rust/crates/llms-types/src/lib.rs @@ -0,0 +1,11 @@ +macro_rules_attribute::attribute_alias! { + #[apply(wire_type)] = + #[derive(Clone, Debug, PartialEq, serde::Serialize, serde::Deserialize)] + #[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]; +} + +pub mod formats; +pub mod headers; +pub mod providers; +pub mod recognized; +pub mod serde_compat; diff --git a/litellm-rust/crates/types/src/llms/anthropic.rs b/litellm-rust/crates/llms-types/src/providers/anthropic.rs similarity index 100% rename from litellm-rust/crates/types/src/llms/anthropic.rs rename to litellm-rust/crates/llms-types/src/providers/anthropic.rs diff --git a/litellm-rust/crates/llms-types/src/providers/mod.rs b/litellm-rust/crates/llms-types/src/providers/mod.rs new file mode 100644 index 00000000000..e529997219e --- /dev/null +++ b/litellm-rust/crates/llms-types/src/providers/mod.rs @@ -0,0 +1 @@ +pub mod anthropic; diff --git a/litellm-rust/crates/types/src/recognized.rs b/litellm-rust/crates/llms-types/src/recognized.rs similarity index 91% rename from litellm-rust/crates/types/src/recognized.rs rename to litellm-rust/crates/llms-types/src/recognized.rs index d82b51f9fde..148d65381a5 100644 --- a/litellm-rust/crates/types/src/recognized.rs +++ b/litellm-rust/crates/llms-types/src/recognized.rs @@ -1,7 +1,6 @@ -use serde::{Deserialize, Serialize}; use serde_json::Value; -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[macro_rules_attribute::apply(wire_type)] #[serde(untagged)] pub enum Recognized { Known(T), diff --git a/litellm-rust/crates/llms-types/src/serde_compat.rs b/litellm-rust/crates/llms-types/src/serde_compat.rs new file mode 100644 index 00000000000..ffa86b7aec8 --- /dev/null +++ b/litellm-rust/crates/llms-types/src/serde_compat.rs @@ -0,0 +1,113 @@ +use serde::{ + Deserializer, + de::{Error, Visitor}, +}; +use serde_with::DeserializeAs; + +pub struct LaxI64; +pub struct FiniteF64; + +impl<'de> DeserializeAs<'de, i64> for LaxI64 { + fn deserialize_as>(deserializer: D) -> Result { + deserializer.deserialize_any(Self) + } +} + +impl<'de> Visitor<'de> for LaxI64 { + type Value = i64; + + fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str("an integer in the i64 range") + } + + fn visit_i64(self, value: i64) -> Result { + Ok(value) + } + + fn visit_u64(self, value: u64) -> Result { + i64::try_from(value).map_err(E::custom) + } + + fn visit_f64(self, value: f64) -> Result { + integral_float(value).ok_or_else(|| E::custom("expected an integer in the i64 range")) + } + + fn visit_str(self, value: &str) -> Result { + integer_string(value.trim()) + .ok_or_else(|| E::custom("expected an integer in the i64 range")) + } + + fn visit_bool(self, value: bool) -> Result { + Ok(i64::from(value)) + } +} + +impl<'de> DeserializeAs<'de, f64> for FiniteF64 { + fn deserialize_as>(deserializer: D) -> Result { + deserializer.deserialize_any(Self) + } +} + +impl<'de> Visitor<'de> for FiniteF64 { + type Value = f64; + + fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str("a finite number") + } + + fn visit_i64(self, value: i64) -> Result { + Ok(value as f64) + } + + fn visit_u64(self, value: u64) -> Result { + Ok(value as f64) + } + + fn visit_f64(self, value: f64) -> Result { + value + .is_finite() + .then_some(value) + .ok_or_else(|| E::custom("expected a finite number")) + } + + fn visit_str(self, value: &str) -> Result { + self.visit_f64(value.trim().parse::().map_err(E::custom)?) + } + + fn visit_bool(self, value: bool) -> Result { + Ok(f64::from(value)) + } +} + +fn integer_string(value: &str) -> Option { + let integer = match value.split_once('.') { + Some((integer, fraction)) => { + if fraction.is_empty() || !fraction.bytes().all(|byte| byte == b'0') { + return None; + } + integer + } + None => value, + }; + if integer.starts_with('_') || integer.ends_with('_') || integer.contains("__") { + return None; + } + let digits = integer.strip_prefix(['+', '-']).unwrap_or(integer); + if digits.is_empty() + || digits.starts_with('_') + || !digits + .bytes() + .all(|byte| byte.is_ascii_digit() || byte == b'_') + { + return None; + } + integer.replace('_', "").parse().ok() +} + +fn integral_float(value: f64) -> Option { + (value.is_finite() + && value.fract() == 0.0 + && value >= i64::MIN as f64 + && value < -(i64::MIN as f64)) + .then_some(value as i64) +} diff --git a/litellm-rust/crates/types/tests/anthropic_request.rs b/litellm-rust/crates/llms-types/tests/messages_request.rs similarity index 95% rename from litellm-rust/crates/types/tests/anthropic_request.rs rename to litellm-rust/crates/llms-types/tests/messages_request.rs index b66eec7d948..4ebc196fb12 100644 --- a/litellm-rust/crates/types/tests/anthropic_request.rs +++ b/litellm-rust/crates/llms-types/tests/messages_request.rs @@ -1,4 +1,4 @@ -use litellm_types::llms::anthropic_messages::anthropic_request::{ContentBlock, ContentBlockType}; +use litellm_llms_types::formats::messages::{ContentBlock, ContentBlockType}; use rstest::rstest; use serde_json::{Value, json}; diff --git a/litellm-rust/crates/types/tests/messages_streaming.rs b/litellm-rust/crates/llms-types/tests/messages_streaming.rs similarity index 94% rename from litellm-rust/crates/types/tests/messages_streaming.rs rename to litellm-rust/crates/llms-types/tests/messages_streaming.rs index c06ea4c2357..5aebb1c052c 100644 --- a/litellm-rust/crates/types/tests/messages_streaming.rs +++ b/litellm-rust/crates/llms-types/tests/messages_streaming.rs @@ -1,4 +1,4 @@ -use litellm_types::messages::streaming::MessagesStreamEvent; +use litellm_llms_types::formats::messages::streaming::MessagesStreamEvent; use rstest::rstest; use serde_json::{Value, json}; diff --git a/litellm-rust/crates/llms-types/tests/ocr.rs b/litellm-rust/crates/llms-types/tests/ocr.rs new file mode 100644 index 00000000000..48c816819f2 --- /dev/null +++ b/litellm-rust/crates/llms-types/tests/ocr.rs @@ -0,0 +1,104 @@ +use litellm_llms_types::formats::ocr::{LiteLLMOcrResponse, OcrDocument, OcrPage}; +use rstest::rstest; +use serde_json::{Map, Value, json}; + +#[rstest] +#[case::missing_page_fields(json!({"pages": [{}]}))] +#[case::invalid_markdown(json!({"pages": [{"index": 0, "markdown": false}]}))] +#[case::invalid_image_bounds(json!({"pages": [{"index": 0, "markdown": "", "images": [{"bbox": []}]}]}))] +#[case::fractional_page_count(json!({"usage_info": {"pages_processed": 1.5}}))] +#[case::invalid_table(json!({"tables": [false]}))] +#[case::invalid_key_value_pair(json!({"keyValuePairs": [[]]}))] +#[case::invalid_native_response(json!({"provider_native_response": []}))] +fn normalized_response_rejects_invalid_shared_fields(#[case] fields: Value) { + let payload: Map = json!({"model": "model", "pages": []}) + .as_object() + .unwrap() + .iter() + .chain(fields.as_object().unwrap()) + .map(|(key, value)| (key.clone(), value.clone())) + .collect(); + assert!(serde_json::from_value::(Value::Object(payload)).is_err()); +} + +#[rstest] +fn document_rejects_non_string_provider_fields() { + assert!( + serde_json::from_value::(json!({ + "type": "image_url", "image_url": "https://example.com/image", "detail": 42 + })) + .is_err() + ); +} + +#[rstest] +#[case::large_integer(json!("9007199254740993.0"), 9_007_199_254_740_993)] +#[case::signed_decimal(json!("+2.000"), 2)] +#[case::separator(json!("1_000"), 1000)] +#[case::boolean(json!(true), 1)] +#[case::integral_float(json!(2.0), 2)] +fn numeric_coercion_preserves_integer_precision(#[case] value: Value, #[case] expected: i64) { + let page: OcrPage = serde_json::from_value(json!({"index": value, "markdown": ""})).unwrap(); + assert_eq!(page.index, expected); + assert_eq!( + serde_json::to_value(page).unwrap()["index"], + json!(expected) + ); +} + +#[rstest] +#[case::exponent(json!("1e2"))] +#[case::missing_integer(json!(".0"))] +#[case::missing_fraction(json!("2."))] +#[case::leading_separator(json!("_2"))] +#[case::repeated_separator(json!("2__0"))] +#[case::fractional_float(json!(2.5))] +#[case::null(json!(null))] +fn page_index_rejects_invalid_integers(#[case] value: Value) { + assert!(serde_json::from_value::(json!({"index": value, "markdown": ""})).is_err()); +} + +#[rstest] +#[case::document_url("document_url", "document_name", "application/pdf")] +#[case::image_url("image_url", "detail", "image/png")] +fn document_variants_preserve_provider_fields_when_rewriting_sources( + #[case] kind: &str, + #[case] field: &str, + #[case] mime_type: &str, + #[values(json!("kept"), Value::Null)] extra: Value, +) { + let original = "https://example.com/input"; + let replacement = format!("data:{mime_type};base64,AA=="); + let document: OcrDocument = + serde_json::from_value(json!({"type": kind, kind: original, field: extra})).unwrap(); + assert_eq!(document.source(), original); + assert!(document.is_remote()); + let rewritten = document.with_source(replacement.clone()); + assert!(!rewritten.is_remote()); + assert_eq!( + serde_json::to_value(rewritten).unwrap(), + json!({"type": kind, kind: replacement, field: extra}) + ); +} + +#[rstest] +#[case::absent_native(None)] +#[case::present_native(Some(Map::from_iter([("native".into(), json!({"nested": [null, 1]}))])))] +fn response_serialization_preserves_extensions_and_native_presence( + #[case] native: Option>, +) { + let response = LiteLLMOcrResponse { + extra_fields: Map::from_iter([("provider_field".into(), json!("kept"))]), + provider_native_response: native.clone(), + ..LiteLLMOcrResponse::new("model", vec![]) + }; + let serialized = response.into_json(); + assert_eq!(serialized["provider_field"], "kept"); + assert_eq!( + serialized.get("provider_native_response").cloned(), + native.clone().map(Value::Object) + ); + let decoded: LiteLLMOcrResponse = serde_json::from_value(serialized.clone()).unwrap(); + assert_eq!(decoded.provider_native_response, native); + assert_eq!(decoded.into_json(), serialized); +} diff --git a/litellm-rust/crates/llms-types/tests/serde_compat.rs b/litellm-rust/crates/llms-types/tests/serde_compat.rs new file mode 100644 index 00000000000..76dba17a241 --- /dev/null +++ b/litellm-rust/crates/llms-types/tests/serde_compat.rs @@ -0,0 +1,86 @@ +use litellm_llms_types::serde_compat::{FiniteF64, LaxI64}; +use rstest::rstest; +use serde::{Deserialize, Serialize}; +use serde_json::{Value, json}; +use serde_with::serde_as; + +#[serde_as] +#[derive(Debug, Deserialize, Serialize, PartialEq)] +struct Numbers { + #[serde_as(deserialize_as = "Option>")] + integers: Option>, + #[serde_as(deserialize_as = "Option")] + float: Option, +} + +#[rstest] +fn adapters_compose_and_serialize_as_numbers() { + let numbers: Numbers = serde_json::from_value(json!({ + "integers": ["9007199254740993.0", "1_000", " +2.000 ", 3.0, true], + "float": " 1.5 " + })) + .unwrap(); + assert_eq!( + serde_json::to_value(numbers).unwrap(), + json!({"integers": [9_007_199_254_740_993_i64, 1000, 2, 3, 1], "float": 1.5}) + ); +} + +#[rstest] +#[case::missing(json!({}))] +#[case::null(json!({"integers": null, "float": null}))] +fn optional_adapters_accept_missing_and_null_fields(#[case] input: Value) { + assert_eq!( + serde_json::from_value::(input).unwrap(), + Numbers { + integers: None, + float: None + } + ); +} + +#[rstest] +#[case::minimum(json!(i64::MIN), i64::MIN)] +#[case::maximum(json!(i64::MAX), i64::MAX)] +#[case::maximum_string(json!(i64::MAX.to_string()), i64::MAX)] +fn integers_preserve_bounds(#[case] input: Value, #[case] expected: i64) { + let numbers: Numbers = serde_json::from_value(json!({"integers": [input]})).unwrap(); + assert_eq!(numbers.integers, Some(vec![expected])); +} + +#[rstest] +#[case::unsigned_maximum(json!(u64::MAX))] +#[case::above_maximum(json!(9_223_372_036_854_775_808_u64))] +#[case::float_above_maximum(json!(9_223_372_036_854_775_808.0))] +#[case::below_minimum(json!("-9223372036854775809"))] +#[case::precise_fraction(json!("1.0000000000000001"))] +#[case::exponent(json!("1e3"))] +#[case::missing_fraction(json!("2."))] +#[case::missing_integer(json!(".0"))] +#[case::leading_separator(json!("_2"))] +#[case::repeated_separator(json!("2__0"))] +#[case::fraction(json!(2.5))] +#[case::null(json!(null))] +#[case::object(json!({}))] +fn integers_reject_invalid_values(#[case] input: Value) { + assert!(serde_json::from_value::(json!({"integers": [input]})).is_err()); +} + +#[rstest] +#[case::nan(json!("NaN"))] +#[case::positive_infinity(json!("inf"))] +#[case::negative_infinity(json!("-inf"))] +#[case::overflow(json!("1e999"))] +#[case::array(json!([]))] +fn floats_reject_nonfinite_and_invalid_values(#[case] input: Value) { + assert!(serde_json::from_value::(json!({"float": input})).is_err()); +} + +#[rstest] +#[case::integer(json!(2), 2.0)] +#[case::float(json!(2.5), 2.5)] +#[case::boolean(json!(true), 1.0)] +fn floats_accept_finite_numbers(#[case] input: Value, #[case] expected: f64) { + let numbers: Numbers = serde_json::from_value(json!({"float": input})).unwrap(); + assert_eq!(numbers.float, Some(expected)); +} diff --git a/litellm-rust/crates/llms-types/tests/wire_type.rs b/litellm-rust/crates/llms-types/tests/wire_type.rs new file mode 100644 index 00000000000..75755b9f40a --- /dev/null +++ b/litellm-rust/crates/llms-types/tests/wire_type.rs @@ -0,0 +1,32 @@ +use litellm_llms_types::formats::chat_completions::ChatMessage; +use rstest::rstest; +use serde_json::json; + +#[rstest] +fn wire_type_preserves_serialization() { + let message = ChatMessage { + role: "user".to_owned(), + content: None, + name: None, + extra: Default::default(), + }; + + assert_eq!( + serde_json::to_value(message).unwrap(), + json!({"role": "user"}) + ); +} + +#[cfg(feature = "schema")] +#[rstest] +fn wire_type_supports_schema_generation() { + let schema = schemars::schema_for!(ChatMessage); + + assert!( + schema + .to_value() + .get("properties") + .and_then(serde_json::Value::as_object) + .is_some_and(|properties| properties.contains_key("role")) + ); +} diff --git a/litellm-rust/crates/llms/AGENTS.md b/litellm-rust/crates/llms/AGENTS.md index 6ecbf7e8a52..c48dc7962c8 100644 --- a/litellm-rust/crates/llms/AGENTS.md +++ b/litellm-rust/crates/llms/AGENTS.md @@ -12,7 +12,7 @@ Use trait defaults for unchanged inherited behavior and explicit delegation for Use named `#[rstest]` cases for independent input/output scenarios instead of loops or repeated calls in one test. Inject reusable setup with `#[fixture]` arguments and use `#[with(...)]` for fixture overrides. Keep assertions about the same result together -Base OCR currently keeps response models next to `BaseOcrConfig` in `src/base_llm/ocr/transformation.rs`. This is legacy placement, not an exception to the shared API contract ownership in `litellm-types`. Rust context/environment types support the runtime. `BaseOcrConfig::prepare_request` corresponds to Python's HTTP-handler preparation rather than a `BaseOCRConfig` method, and `validate_request_body` is a Rust-only hook. `src/base_llm/ocr/error.rs` and `src/base_llm/ocr/document.rs` are Rust-only: the OCR error taxonomy shared with the route, and inline-document helpers shared by several providers +Shared OCR document and response contracts live in `litellm-llms-types::formats::ocr`. `BaseOcrConfig` and decoding into adapter errors remain in `src/base_llm/ocr/transformation.rs`. Rust context/environment types support the runtime. `BaseOcrConfig::prepare_request` corresponds to Python's HTTP-handler preparation rather than a `BaseOCRConfig` method, and `validate_request_body` is a Rust-only hook. `src/base_llm/ocr/error.rs` and `src/base_llm/ocr/document.rs` are Rust-only: the OCR error taxonomy shared with the route, and inline-document helpers shared by several providers For Mistral, `async_transform_ocr_request` uses the base default in both languages. `resolve_headers` and `build_ocr_url` implement the respective environment and URL operations, and `normalize_response` implements the typed part of response transformation. Existing auth key/header handling and top-level response-extra preservation differ between languages; layout refactors must preserve those behaviors and verify them with the existing tests @@ -22,7 +22,7 @@ Azure Messages maps to `llms/azure_ai/anthropic/messages_transformation.py`; Bed ## Provider and format boundaries -The same ownership rule applies to Messages, Responses, Chat Completions, OCR, and other API formats. `litellm-types` owns shared API data contracts. `llms/src/base_llm//` owns provider adapter contracts and shared transformation machinery. `llms/src///` owns provider implementations and policy. `core/src//` owns call orchestration. Repeating a format name identifies the API each layer handles, not duplicate ownership of its schema. These boundaries also apply between modules in the same crate +The same ownership rule applies to Messages, Responses, Chat Completions, OCR, and other API formats. `litellm-llms-types` owns shared API data contracts. `llms/src/base_llm//` owns provider adapter contracts and shared transformation machinery. `llms/src///` owns provider implementations and policy. `core/src//` owns call orchestration. Repeating a format name identifies the API each layer handles, not duplicate ownership of its schema. These boundaries also apply between modules in the same crate A provider adapter may explicitly reuse another provider's transformation helper when that policy applies to its backend, such as Bedrock's Claude adapter using Anthropic payload shaping. Reuse across hosts of the same model family does not make the policy format-wide. Keep provider policy out of shared trait defaults and generic normalization, and keep shared execution contexts limited to inputs the adapter contract actually needs. Pure payload rewrites belong with transformations, not transport handlers diff --git a/litellm-rust/crates/llms/Cargo.toml b/litellm-rust/crates/llms/Cargo.toml index beff99bc73a..cac52454108 100644 --- a/litellm-rust/crates/llms/Cargo.toml +++ b/litellm-rust/crates/llms/Cargo.toml @@ -9,7 +9,7 @@ repository.workspace = true test-support = ["litellm-http/test-support"] [dependencies] -litellm-types.workspace = true +litellm-llms-types.workspace = true litellm-core-utils.workspace = true litellm-auth = { workspace = true, features = ["aws", "azure", "gcp"] } litellm-auth-aws.workspace = true diff --git a/litellm-rust/crates/llms/src/anthropic/AGENTS.md b/litellm-rust/crates/llms/src/anthropic/AGENTS.md index 52c01a911d0..a226d4b56ce 100644 --- a/litellm-rust/crates/llms/src/anthropic/AGENTS.md +++ b/litellm-rust/crates/llms/src/anthropic/AGENTS.md @@ -3,6 +3,6 @@ - Put behavior specific to the Messages API in `messages/` - Keep generic HTTP mechanics in `litellm-http`, configuration lookup in the existing settings utilities, and credential application in the shared auth layer - Choose authentication policy and required headers here, then let shared infrastructure apply those decisions -- Consume shared API contracts from `litellm-types`. Do not define public Messages protocol types under this provider +- Consume shared API contracts from `litellm-llms-types`. Do not define public Messages protocol types under this provider - Preserve Python's concepts and observable behavior where useful, without mechanically reproducing its class hierarchy, helpers, or file structure - `ReplayedWebSearchResult` and `ReplayedWebSearchContent` are private partial models for replay flattening, not complete public protocol contracts. Keep them private while they serve that transformation diff --git a/litellm-rust/crates/llms/src/anthropic/batches/transformation.rs b/litellm-rust/crates/llms/src/anthropic/batches/transformation.rs index b65db1644bb..209fc7a0058 100644 --- a/litellm-rust/crates/llms/src/anthropic/batches/transformation.rs +++ b/litellm-rust/crates/llms/src/anthropic/batches/transformation.rs @@ -1,4 +1,5 @@ -use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; +use litellm_llms_types::formats::batches::{BatchRequestCounts, BatchResponse, BatchStatus}; +use litellm_llms_types::formats::messages::MessagesResponse; use serde::{Deserialize, Serialize}; use serde_json::Value; use time::OffsetDateTime; @@ -45,46 +46,8 @@ struct BatchResultRecord { #[derive(Deserialize)] #[serde(tag = "type", rename_all = "snake_case")] enum BatchResult { - Succeeded { - message: Box, - }, - Errored { - error: Value, - }, -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)] -#[serde(rename_all = "snake_case")] -pub enum BatchStatus { - InProgress, - Cancelling, - Completed, -} - -#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub struct BatchRequestCounts { - pub total: u64, - pub completed: u64, - pub failed: u64, -} - -#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub struct LiteLlmMessageBatch { - pub id: String, - pub object: String, - pub endpoint: String, - pub input_file_id: String, - pub completion_window: String, - pub status: BatchStatus, - pub output_file_id: String, - pub created_at: i64, - pub in_progress_at: Option, - pub expires_at: Option, - pub completed_at: Option, - pub expired_at: Option, - pub cancelling_at: Option, - pub cancelled_at: Option, - pub request_counts: BatchRequestCounts, + Succeeded { message: Box }, + Errored { error: Value }, } pub trait AnthropicBatchesConfig { @@ -100,7 +63,7 @@ pub trait AnthropicBatchesConfig { &self, response: AnthropicMessageBatch, now: i64, - ) -> Result; + ) -> Result; fn retrieve_batch_url( &self, @@ -115,9 +78,9 @@ pub trait AnthropicBatchesConfig { &self, response: AnthropicMessageBatch, now: i64, - ) -> LiteLlmMessageBatch; + ) -> BatchResponse; - fn transform_batch_results(&self, body: &str) -> Result, Error>; + fn transform_batch_results(&self, body: &str) -> Result, Error>; } pub struct AnthropicBatchesTransformation; @@ -172,7 +135,7 @@ impl AnthropicBatchesConfig for AnthropicBatchesTransformation { &self, _response: AnthropicMessageBatch, _now: i64, - ) -> Result { + ) -> Result { Err(Error::Unsupported("Anthropic message batch creation")) } @@ -200,7 +163,7 @@ impl AnthropicBatchesConfig for AnthropicBatchesTransformation { &self, response: AnthropicMessageBatch, now: i64, - ) -> LiteLlmMessageBatch { + ) -> BatchResponse { let created_at = timestamp(response.created_at.as_deref()); let ended_at = timestamp(response.ended_at.as_deref()); let expires_at = timestamp(response.expires_at.as_deref()); @@ -221,7 +184,7 @@ impl AnthropicBatchesConfig for AnthropicBatchesTransformation { failed: response.request_counts.errored, }; - LiteLlmMessageBatch { + BatchResponse { id: response.id.clone(), object: "batch".into(), endpoint: "/v1/messages".into(), @@ -248,7 +211,7 @@ impl AnthropicBatchesConfig for AnthropicBatchesTransformation { } } - fn transform_batch_results(&self, body: &str) -> Result, Error> { + fn transform_batch_results(&self, body: &str) -> Result, Error> { body.lines() .filter(|line| !line.trim().is_empty()) .enumerate() diff --git a/litellm-rust/crates/llms/src/anthropic/chat/handler.rs b/litellm-rust/crates/llms/src/anthropic/chat/handler.rs index 1b6dd26f4af..eaee006f228 100644 --- a/litellm-rust/crates/llms/src/anthropic/chat/handler.rs +++ b/litellm-rust/crates/llms/src/anthropic/chat/handler.rs @@ -1,11 +1,13 @@ use std::collections::HashMap; -use litellm_types::messages::streaming::{ - MessagesContentBlock, MessagesContentBlockDelta, MessagesStreamEvent, MessagesStreamUsage, -}; -use litellm_types::{ - llms::openai::{ChatCompletionThinkingBlock, ChatCompletionToolCallChunk}, - utils::{ChatCompletionChunk, ChatCompletionsUsage}, +use litellm_llms_types::formats::{ + chat_completions::{ + ChatCompletionChunk, ChatCompletionThinkingBlock, ChatCompletionToolCallChunk, + ChatCompletionsUsage, + }, + messages::streaming::{ + MessagesContentBlock, MessagesContentBlockDelta, MessagesStreamEvent, MessagesStreamUsage, + }, }; use serde_json::Value; diff --git a/litellm-rust/crates/llms/src/anthropic/chat/transformation.rs b/litellm-rust/crates/llms/src/anthropic/chat/transformation.rs index 922e2377eb5..ecc6cfaf83d 100644 --- a/litellm-rust/crates/llms/src/anthropic/chat/transformation.rs +++ b/litellm-rust/crates/llms/src/anthropic/chat/transformation.rs @@ -3,9 +3,8 @@ use litellm_core_utils::{ core_helpers::{finish_reason_for, unix_now, usage_from_parts}, prompt_templates::factory::{Conversation, build_conversation}, }; -use litellm_types::{ - llms::openai::ChatMessage, - utils::{ChatCompletionsChoice, ChatCompletionsChoiceMessage, ChatCompletionsResponse}, +use litellm_llms_types::formats::chat_completions::{ + ChatCompletionsChoice, ChatCompletionsChoiceMessage, ChatCompletionsResponse, ChatMessage, }; use serde::Deserialize; use serde_json::{Map, Value, json}; @@ -50,16 +49,16 @@ const SUPPORTED_PARAMS: &[(&str, &str)] = &[ ]; #[derive(Deserialize)] -struct MessageResponse { +struct TextResponseProjection { model: String, - content: Vec, - usage: MessageUsage, + content: Vec, + usage: ResponseUsageProjection, stop_reason: Option, } #[derive(Deserialize)] #[serde(tag = "type", rename_all = "snake_case")] -enum ContentBlock { +enum TextResponseBlock { Text { text: String, }, @@ -68,7 +67,7 @@ enum ContentBlock { } #[derive(Deserialize)] -struct MessageUsage { +struct ResponseUsageProjection { input_tokens: u64, output_tokens: u64, #[serde(default)] @@ -124,16 +123,17 @@ impl BaseConfig for AnthropicConfig { _model: &str, response: ProviderChatResponseData, ) -> Result { - let body: MessageResponse = serde_json::from_value(response.body).map_err(|error| { - Error::InvalidResponse(crate::ErrorDetail::invalid("messages response", error)) - })?; + let body: TextResponseProjection = + serde_json::from_value(response.body).map_err(|error| { + Error::InvalidResponse(crate::ErrorDetail::invalid("messages response", error)) + })?; // The route declines tool and thinking requests, so a non-text block // means the response carries something this path never asked for. // Decline rather than silently dropping it; the host falls back. if body .content .iter() - .any(|block| matches!(block, ContentBlock::Other)) + .any(|block| matches!(block, TextResponseBlock::Other)) { return Err(Error::Unsupported("non-text response content block")); } @@ -141,8 +141,8 @@ impl BaseConfig for AnthropicConfig { .content .into_iter() .map(|block| match block { - ContentBlock::Text { text } => text, - ContentBlock::Other => String::new(), + TextResponseBlock::Text { text } => text, + TextResponseBlock::Other => String::new(), }) .collect(); diff --git a/litellm-rust/crates/llms/src/anthropic/common_utils.rs b/litellm-rust/crates/llms/src/anthropic/common_utils.rs index 528c33d2fd1..6e2fa3b8785 100644 --- a/litellm-rust/crates/llms/src/anthropic/common_utils.rs +++ b/litellm-rust/crates/llms/src/anthropic/common_utils.rs @@ -4,14 +4,13 @@ use litellm_core_utils::settings::resolve_non_empty; use litellm_http::request::{ has_header, header_value, header_values, with_header, without_headers, }; -use litellm_types::llms::{ - anthropic::{AnthropicBeta, BetaSet}, - anthropic_messages::anthropic_request::{ - AnthropicMessage, AnthropicTool, ContentBlock, ContentBlockType, EffortLevel, - MessageContent, +use litellm_llms_types::{ + formats::messages::{ + ContentBlock, ContentBlockType, EffortLevel, Message, MessageContent, MessagesTool, }, + providers::anthropic::{AnthropicBeta, BetaSet}, + recognized::Recognized, }; -use litellm_types::recognized::Recognized; use serde::Deserialize; use serde_json::Value; @@ -229,36 +228,30 @@ pub fn optionally_handle_anthropic_oauth(headers: Headers, api_key: Option<&str> OauthHandling::Untouched(headers) } -pub fn is_tool_search_used(tools: Option<&[Recognized]>) -> bool { +pub fn is_tool_search_used(tools: Option<&[Recognized]>) -> bool { tools.into_iter().flatten().any(|tool| { matches!( tool, Recognized::Known( - AnthropicTool::ToolSearchRegex { .. } | AnthropicTool::ToolSearchBm25 { .. } + MessagesTool::ToolSearchRegex { .. } | MessagesTool::ToolSearchBm25 { .. } ) ) }) } -pub fn has_advisor_tool(tools: Option<&[Recognized]>) -> bool { +pub fn has_advisor_tool(tools: Option<&[Recognized]>) -> bool { tools .into_iter() .flatten() - .any(|tool| matches!(tool, Recognized::Known(AnthropicTool::Advisor { .. }))) + .any(|tool| matches!(tool, Recognized::Known(MessagesTool::Advisor { .. }))) } -pub fn requires_native_compaction_beta( - compaction: Option<&Value>, - messages: &[AnthropicMessage], -) -> bool { +pub fn requires_native_compaction_beta(compaction: Option<&Value>, messages: &[Message]) -> bool { compaction.is_some() - || messages - .iter() - .flat_map(AnthropicMessage::blocks) - .any(|block| { - block.is_type(ContentBlockType::Compaction) - && block.signature.as_deref().is_some_and(|s| !s.is_empty()) - }) + || messages.iter().flat_map(Message::blocks).any(|block| { + block.is_type(ContentBlockType::Compaction) + && block.signature.as_deref().is_some_and(|s| !s.is_empty()) + }) } fn is_blank(text: Option<&str>) -> bool { @@ -273,10 +266,7 @@ pub fn is_empty_thinking_block(block: &ContentBlock) -> bool { block.is_type(ContentBlockType::Thinking) && is_blank(block.thinking.as_deref()) } -fn retain_blocks( - messages: Vec, - keep: impl Fn(&ContentBlock) -> bool, -) -> Vec { +fn retain_blocks(messages: Vec, keep: impl Fn(&ContentBlock) -> bool) -> Vec { messages .into_iter() .filter_map(|message| match message.content { @@ -293,7 +283,7 @@ fn retain_blocks( .collect() } -pub fn strip_empty_content_blocks(messages: Vec) -> Vec { +pub fn strip_empty_content_blocks(messages: Vec) -> Vec { retain_blocks(messages, |block| { !is_empty_text_block(block) && !is_empty_thinking_block(block) }) @@ -350,11 +340,11 @@ fn sanitize_tool_use_id_block(block: ContentBlock) -> ContentBlock { } } -pub fn sanitize_tool_use_ids(messages: Vec) -> Vec { +pub fn sanitize_tool_use_ids(messages: Vec) -> Vec { messages .into_iter() .map(|message| match message.content { - MessageContent::Blocks(blocks) => AnthropicMessage { + MessageContent::Blocks(blocks) => Message { content: MessageContent::Blocks( blocks.into_iter().map(sanitize_tool_use_id_block).collect(), ), @@ -365,11 +355,11 @@ pub fn sanitize_tool_use_ids(messages: Vec) -> Vec) -> Vec { +pub fn strip_provider_specific_fields(messages: Vec) -> Vec { messages .into_iter() .map(|message| match message.content { - MessageContent::Blocks(blocks) => AnthropicMessage { + MessageContent::Blocks(blocks) => Message { content: MessageContent::Blocks( blocks .into_iter() @@ -395,7 +385,7 @@ pub fn is_encrypted_reasoning_block(block: &ContentBlock) -> bool { field.is_some_and(|value| value.starts_with(ENCRYPTED_REASONING_SIGNATURE_PREFIX)) } -pub fn strip_encrypted_reasoning_blocks(messages: Vec) -> Vec { +pub fn strip_encrypted_reasoning_blocks(messages: Vec) -> Vec { retain_blocks(messages, |block| !is_encrypted_reasoning_block(block)) } @@ -405,7 +395,7 @@ fn is_advisor_use(block: &ContentBlock) -> bool { && block.id.as_deref().is_some_and(|id| !id.is_empty()) } -pub fn strip_advisor_blocks(messages: Vec) -> Vec { +pub fn strip_advisor_blocks(messages: Vec) -> Vec { messages .into_iter() .map(|message| { @@ -588,13 +578,11 @@ fn flatten_web_search_results_in_blocks(blocks: Vec) -> Vec, -) -> Vec { +pub fn flatten_unencrypted_web_search_results(messages: Vec) -> Vec { messages .into_iter() .map(|message| match message.content { - MessageContent::Blocks(blocks) => AnthropicMessage { + MessageContent::Blocks(blocks) => Message { content: MessageContent::Blocks(flatten_web_search_results_in_blocks(blocks)), ..message }, @@ -619,11 +607,8 @@ mod tests { EffortLevel::Max, ]; - fn apply( - sanitizer: fn(Vec) -> Vec, - messages: Value, - ) -> Value { - let parsed: Vec = serde_json::from_value(messages).unwrap(); + fn apply(sanitizer: fn(Vec) -> Vec, messages: Value) -> Value { + let parsed: Vec = serde_json::from_value(messages).unwrap(); serde_json::to_value(sanitizer(parsed)).unwrap() } @@ -631,11 +616,11 @@ mod tests { serde_json::from_value(value).unwrap() } - fn history(messages: Value) -> Vec { + fn history(messages: Value) -> Vec { serde_json::from_value(messages).unwrap() } - fn tools(value: Option) -> Option>> { + fn tools(value: Option) -> Option>> { value.map(|tools| serde_json::from_value(tools).unwrap()) } diff --git a/litellm-rust/crates/llms/src/anthropic/count_tokens/transformation.rs b/litellm-rust/crates/llms/src/anthropic/count_tokens/transformation.rs index 9fa831b8b66..4d892ca198c 100644 --- a/litellm-rust/crates/llms/src/anthropic/count_tokens/transformation.rs +++ b/litellm-rust/crates/llms/src/anthropic/count_tokens/transformation.rs @@ -1,4 +1,4 @@ -use litellm_types::llms::anthropic_messages::anthropic_request::{AnthropicMessage, SystemPrompt}; +use litellm_llms_types::formats::messages::{Message, SystemPrompt}; use serde::{Deserialize, Serialize}; use serde_json::Value; @@ -10,7 +10,7 @@ const TOKEN_COUNTING_BETA: &str = "token-counting-2024-11-01"; #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] pub struct AnthropicCountTokensRequest { pub model: String, - pub messages: Vec, + pub messages: Vec, #[serde(skip_serializing_if = "Option::is_none")] pub tools: Option>, #[serde(skip_serializing_if = "Option::is_none")] @@ -25,12 +25,12 @@ pub struct AnthropicCountTokensResponse { pub trait AnthropicCountTokensConfig { fn endpoint(&self) -> &'static str; - fn validate_request(&self, model: &str, messages: &[AnthropicMessage]) -> Result<(), Error>; + fn validate_request(&self, model: &str, messages: &[Message]) -> Result<(), Error>; fn transform_request( &self, model: &str, - messages: Vec, + messages: Vec, tools: Option>, system: Option, ) -> Result; @@ -51,7 +51,7 @@ impl AnthropicCountTokensConfig for AnthropicCountTokensTransformation { fn transform_request( &self, model: &str, - messages: Vec, + messages: Vec, tools: Option>, system: Option, ) -> Result { @@ -65,7 +65,7 @@ impl AnthropicCountTokensConfig for AnthropicCountTokensTransformation { }) } - fn validate_request(&self, model: &str, messages: &[AnthropicMessage]) -> Result<(), Error> { + fn validate_request(&self, model: &str, messages: &[Message]) -> Result<(), Error> { if model.is_empty() { return Err(Error::MissingField("model")); } @@ -92,13 +92,13 @@ impl AnthropicCountTokensConfig for AnthropicCountTokensTransformation { #[cfg(test)] mod tests { - use litellm_types::llms::anthropic_messages::anthropic_request::MessageContent; + use litellm_llms_types::formats::messages::MessageContent; use serde_json::{Map, json}; use super::*; - fn message() -> AnthropicMessage { - AnthropicMessage { + fn message() -> Message { + Message { role: "user".into(), content: MessageContent::Text("hello".into()), extra: Map::new(), diff --git a/litellm-rust/crates/llms/src/anthropic/messages/AGENTS.md b/litellm-rust/crates/llms/src/anthropic/messages/AGENTS.md index 53cd7e7b95e..87d5c9a0936 100644 --- a/litellm-rust/crates/llms/src/anthropic/messages/AGENTS.md +++ b/litellm-rust/crates/llms/src/anthropic/messages/AGENTS.md @@ -1,9 +1,9 @@ -This directory owns Anthropic's implementation of the Messages adapter contract in `base_llm/messages`. Shared Messages API data contracts belong in `litellm-types::messages`, and call orchestration belongs in `core/src/messages`. Sharing the `llms` crate with `base_llm/messages` does not erase this boundary +This directory owns Anthropic's implementation of the Messages adapter contract in `base_llm/messages`. Shared Messages API data contracts belong in `litellm-llms-types::formats::messages`, and call orchestration belongs in `core/src/messages`. Sharing the `llms` crate with `base_llm/messages` does not erase this boundary Payload shaping, metadata filtering, tool-ID rewriting, web-search replay handling, thinking translation, and beta selection are provider policy. Keep them here or in Anthropic helpers shared by its operations. Pure payload shaping belongs with transformations, even if an existing file is named `handler.rs` Bedrock and Azure adapters may explicitly reuse these helpers where Anthropic policy applies to their Claude backend. That reuse does not make the policy part of the shared Messages contract or a default for every provider. Shared `base_llm` code must never depend on this implementation -`web_search_result`, `web_search_tool_result_error`, and encrypted-content fields are protocol data owned by `litellm-types`. Keep those schemas separate from decisions about flattening, encrypted results, beta requirements, and model capabilities +`web_search_result`, `web_search_tool_result_error`, and encrypted-content fields are protocol data owned by `litellm-llms-types`. Keep those schemas separate from decisions about flattening, encrypted results, beta requirements, and model capabilities Protocol reference: [Messages API](https://platform.claude.com/docs/en/api/http/messages/create) diff --git a/litellm-rust/crates/llms/src/anthropic/messages/handler.rs b/litellm-rust/crates/llms/src/anthropic/messages/handler.rs index a7aa7b13d22..6e8de65d6fa 100644 --- a/litellm-rust/crates/llms/src/anthropic/messages/handler.rs +++ b/litellm-rust/crates/llms/src/anthropic/messages/handler.rs @@ -1,7 +1,7 @@ -use litellm_types::{ - llms::anthropic_messages::anthropic_request::{ - AdaptiveThinking, AnthropicMessage, AnthropicMessagesOptionalParams, - AnthropicMessagesRequest, EnabledThinking, ThinkingConfig, ThinkingDisplay, +use litellm_llms_types::{ + formats::messages::{ + AdaptiveThinking, EnabledThinking, Message, MessagesOptionalParams, MessagesRequest, + ThinkingConfig, ThinkingDisplay, }, recognized::Recognized, }; @@ -16,12 +16,12 @@ use crate::{ }; pub fn shape_anthropic_messages_request( - request: AnthropicMessagesRequest, + request: MessagesRequest, reasoning_auto_summary: bool, -) -> Result { - Ok(AnthropicMessagesRequest { +) -> Result { + Ok(MessagesRequest { messages: sanitize_anthropic_messages(request.messages), - params: AnthropicMessagesOptionalParams { + params: MessagesOptionalParams { metadata: request .params .metadata @@ -35,7 +35,7 @@ pub fn shape_anthropic_messages_request( }) } -fn sanitize_anthropic_messages(messages: Vec) -> Vec { +fn sanitize_anthropic_messages(messages: Vec) -> Vec { strip_provider_specific_fields(flatten_unencrypted_web_search_results( sanitize_tool_use_ids(strip_empty_content_blocks(messages)), )) @@ -100,11 +100,11 @@ mod tests { use super::*; - fn messages(value: Value) -> Vec { + fn messages(value: Value) -> Vec { serde_json::from_value(value).unwrap() } - fn request(body: Value) -> AnthropicMessagesRequest { + fn request(body: Value) -> MessagesRequest { serde_json::from_value(body).unwrap() } diff --git a/litellm-rust/crates/llms/src/anthropic/messages/thinking.rs b/litellm-rust/crates/llms/src/anthropic/messages/thinking.rs index de55dd5e864..2162c39c229 100644 --- a/litellm-rust/crates/llms/src/anthropic/messages/thinking.rs +++ b/litellm-rust/crates/llms/src/anthropic/messages/thinking.rs @@ -1,14 +1,14 @@ -use litellm_python_compat::{json::from_json, repr::repr, truthy::truthy}; -use litellm_types::{ - llms::{ - anthropic_messages::anthropic_request::{ - AnthropicMessagesOptionalParams, AnthropicMessagesRequest, EffortLevel, OutputConfig, - ThinkingConfig, ThinkingDisplay, +use litellm_llms_types::{ + formats::{ + chat_completions::ReasoningEffort, + messages::{ + EffortLevel, MessagesOptionalParams, MessagesRequest, OutputConfig, ThinkingConfig, + ThinkingDisplay, }, - openai::ReasoningEffort, }, recognized::Recognized, }; +use litellm_python_compat::{json::from_json, repr::repr, truthy::truthy}; use serde_json::Value; use crate::base_llm::messages::context::{ @@ -84,11 +84,11 @@ fn fit_budget_to_max_tokens(budget_tokens: u64, max_tokens: Option) -> Opti (max_tokens > ANTHROPIC_MIN_THINKING_BUDGET_TOKENS).then(|| budget_tokens.min(max_tokens - 1)) } -fn known_thinking(request: &AnthropicMessagesRequest) -> Option<&ThinkingConfig> { +fn known_thinking(request: &MessagesRequest) -> Option<&ThinkingConfig> { request.params.thinking.as_ref().and_then(Recognized::known) } -fn known_effort(request: &AnthropicMessagesRequest) -> Option<&Recognized> { +fn known_effort(request: &MessagesRequest) -> Option<&Recognized> { request .params .output_config @@ -141,14 +141,14 @@ fn legacy_reasoning_effort( } fn translate_reasoning_effort( - request: AnthropicMessagesRequest, + request: MessagesRequest, context: &ThinkingContext, -) -> Result { +) -> Result { let Some(reasoning_effort) = request.params.reasoning_effort else { return Ok(request); }; - let request = AnthropicMessagesRequest { - params: AnthropicMessagesOptionalParams { + let request = MessagesRequest { + params: MessagesOptionalParams { reasoning_effort: None, ..request.params }, @@ -165,8 +165,8 @@ fn translate_reasoning_effort( output_effort(effort), budget_for_effort(&context.budgets, effort), ) else { - return Ok(AnthropicMessagesRequest { - params: AnthropicMessagesOptionalParams { + return Ok(MessagesRequest { + params: MessagesOptionalParams { thinking: None, output_config: None, ..request.params @@ -180,8 +180,8 @@ fn translate_reasoning_effort( return Err(unsupported_effort(level, &request.model)); } let adaptive = ThinkingConfig::adaptive(Some(ThinkingDisplay::Summarized)); - return Ok(AnthropicMessagesRequest { - params: AnthropicMessagesOptionalParams { + return Ok(MessagesRequest { + params: MessagesOptionalParams { thinking: Some( request .params @@ -198,8 +198,8 @@ fn translate_reasoning_effort( return Ok(request); }; let enabled = ThinkingConfig::enabled(budget); - Ok(AnthropicMessagesRequest { - params: AnthropicMessagesOptionalParams { + Ok(MessagesRequest { + params: MessagesOptionalParams { thinking: Some( request .params @@ -212,17 +212,14 @@ fn translate_reasoning_effort( }) } -fn drop_disabled_thinking( - request: AnthropicMessagesRequest, - context: &ThinkingContext, -) -> AnthropicMessagesRequest { +fn drop_disabled_thinking(request: MessagesRequest, context: &ThinkingContext) -> MessagesRequest { if !context.capabilities.thinking_always_on || !matches!(known_thinking(&request), Some(ThinkingConfig::Disabled(_))) { return request; } - AnthropicMessagesRequest { - params: AnthropicMessagesOptionalParams { + MessagesRequest { + params: MessagesOptionalParams { thinking: None, ..request.params }, @@ -231,9 +228,9 @@ fn drop_disabled_thinking( } fn translate_legacy_thinking_for_adaptive_model( - request: AnthropicMessagesRequest, + request: MessagesRequest, context: &ThinkingContext, -) -> AnthropicMessagesRequest { +) -> MessagesRequest { let capabilities = &context.capabilities; if !capabilities.supports_adaptive_thinking || capabilities.supports_legacy_thinking { return request; @@ -248,8 +245,8 @@ fn translate_legacy_thinking_for_adaptive_model( .copied() .unwrap_or(0); let level = effort_for_budget(&context.budgets, budget, capabilities); - AnthropicMessagesRequest { - params: AnthropicMessagesOptionalParams { + MessagesRequest { + params: MessagesOptionalParams { thinking: Some(Recognized::Known(ThinkingConfig::adaptive(None))), output_config: with_default_effort(request.params.output_config, level), ..request.params @@ -259,9 +256,9 @@ fn translate_legacy_thinking_for_adaptive_model( } fn translate_adaptive_effort_for_non_adaptive_model( - request: AnthropicMessagesRequest, + request: MessagesRequest, context: &ThinkingContext, -) -> Result { +) -> Result { let capabilities = &context.capabilities; if capabilities.supports_adaptive_thinking { return Ok(request); @@ -276,8 +273,8 @@ fn translate_adaptive_effort_for_non_adaptive_model( _ => true, }; if supports_effort_param(capabilities) && (!adaptive_thinking || level_accepted) { - return Ok(AnthropicMessagesRequest { - params: AnthropicMessagesOptionalParams { + return Ok(MessagesRequest { + params: MessagesOptionalParams { thinking: if adaptive_thinking { None } else { @@ -293,8 +290,8 @@ fn translate_adaptive_effort_for_non_adaptive_model( } else { None }; - Ok(AnthropicMessagesRequest { - params: AnthropicMessagesOptionalParams { + Ok(MessagesRequest { + params: MessagesOptionalParams { thinking: budget .and_then(|budget| fit_budget_to_max_tokens(budget, request.params.max_tokens)) .map(|budget| Recognized::Known(ThinkingConfig::enabled(budget))), @@ -306,9 +303,9 @@ fn translate_adaptive_effort_for_non_adaptive_model( } fn drop_incompatible_temperature_for_thinking( - request: AnthropicMessagesRequest, + request: MessagesRequest, context: &ThinkingContext, -) -> AnthropicMessagesRequest { +) -> MessagesRequest { if context.capabilities.supports_adaptive_thinking { return request; } @@ -321,8 +318,8 @@ fn drop_incompatible_temperature_for_thinking( if !pinned || !(thinking_enabled || effort_enabled) { return request; } - AnthropicMessagesRequest { - params: AnthropicMessagesOptionalParams { + MessagesRequest { + params: MessagesOptionalParams { temperature: None, ..request.params }, @@ -331,9 +328,9 @@ fn drop_incompatible_temperature_for_thinking( } pub fn translate_thinking( - request: AnthropicMessagesRequest, + request: MessagesRequest, context: &ThinkingContext, -) -> Result { +) -> Result { let request = translate_reasoning_effort(request, context)?; let request = drop_disabled_thinking(request, context); let request = translate_legacy_thinking_for_adaptive_model(request, context); @@ -350,7 +347,7 @@ mod tests { const EFFORT_CHOICES: &str = "'none', 'minimal', 'low', 'medium', 'high', 'xhigh', 'max'"; - fn request(fields: Value) -> AnthropicMessagesRequest { + fn request(fields: Value) -> MessagesRequest { let mut body = serde_json::json!({"model": "claude", "messages": [{"role": "user", "content": "Hello"}]}); body.as_object_mut() .unwrap() @@ -368,7 +365,7 @@ mod tests { fn translate( capabilities: MessagesModelCapabilities, fields: Value, - ) -> Result { + ) -> Result { translate_thinking(request(fields), &context(capabilities)) } diff --git a/litellm-rust/crates/llms/src/anthropic/messages/transformation.rs b/litellm-rust/crates/llms/src/anthropic/messages/transformation.rs index 0a33cd08e3a..b9cf6c37272 100644 --- a/litellm-rust/crates/llms/src/anthropic/messages/transformation.rs +++ b/litellm-rust/crates/llms/src/anthropic/messages/transformation.rs @@ -1,12 +1,9 @@ use litellm_auth::CredentialPlacement; -use litellm_types::{ - llms::{ - anthropic::{AnthropicBeta, BetaSet}, - anthropic_messages::anthropic_request::{ - AnthropicMessage, AnthropicMessagesOptionalParams, AnthropicMessagesRequest, - ContextEdit, ContextManagement, Speed, - }, +use litellm_llms_types::{ + formats::messages::{ + ContextEdit, ContextManagement, Message, MessagesOptionalParams, MessagesRequest, Speed, }, + providers::anthropic::{AnthropicBeta, BetaSet}, recognized::Recognized, }; use serde_json::{Map, Value, json}; @@ -24,7 +21,7 @@ use crate::{ }, base_llm::{ auth::AuthScheme, - messages::transformation::{BaseAnthropicMessagesConfig, Headers, ValidatedEnvironment}, + messages::transformation::{BaseMessagesConfig, Headers, ValidatedEnvironment}, }, }; @@ -37,12 +34,12 @@ pub struct AnthropicMessagesConfig; pub const ANTHROPIC_MESSAGES_CONFIG: AnthropicMessagesConfig = AnthropicMessagesConfig; -impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig { +impl BaseMessagesConfig for AnthropicMessagesConfig { fn shape_request( &self, - request: AnthropicMessagesRequest, + request: MessagesRequest, reasoning_auto_summary: bool, - ) -> Result { + ) -> Result { shape_anthropic_messages_request(request, reasoning_auto_summary) } @@ -57,9 +54,9 @@ impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig { fn transform_anthropic_messages_request( &self, - request: AnthropicMessagesRequest, + request: MessagesRequest, context: &MessagesTransformContext, - ) -> Result { + ) -> Result { transform_messages_request(request, context) } @@ -112,15 +109,15 @@ impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig { DEFAULT_HEADERS } - fn request_headers(&self, headers: Headers, request: &AnthropicMessagesRequest) -> Headers { + fn request_headers(&self, headers: Headers, request: &MessagesRequest) -> Headers { update_headers_with_anthropic_beta(headers, request) } } pub(crate) fn transform_messages_request( - request: AnthropicMessagesRequest, + request: MessagesRequest, context: &MessagesTransformContext, -) -> Result { +) -> Result { if request.params.max_tokens.is_none() { return Err(Error::MissingField("max_tokens")); } @@ -136,9 +133,9 @@ pub(crate) fn transform_messages_request( } else { strip_advisor_blocks(request.messages) }; - Ok(AnthropicMessagesRequest { + Ok(MessagesRequest { messages: strip_encrypted_reasoning_blocks(messages), - params: AnthropicMessagesOptionalParams { + params: MessagesOptionalParams { context_management, ..request.params }, @@ -148,12 +145,12 @@ pub(crate) fn transform_messages_request( pub(crate) fn update_headers_with_anthropic_beta( headers: Headers, - request: &AnthropicMessagesRequest, + request: &MessagesRequest, ) -> Headers { merge_beta_headers(headers, feature_betas(request)) } -fn feature_betas(request: &AnthropicMessagesRequest) -> BetaSet { +fn feature_betas(request: &MessagesRequest) -> BetaSet { let params = &request.params; let tools = params.tools.as_deref(); [ @@ -192,7 +189,7 @@ fn context_management_betas( .chain(other.then_some(AnthropicBeta::ContextManagement20250627)) } -fn uses_structured_output(params: &AnthropicMessagesOptionalParams) -> bool { +fn uses_structured_output(params: &MessagesOptionalParams) -> bool { params.output_format.is_some() || params .output_config @@ -201,7 +198,7 @@ fn uses_structured_output(params: &AnthropicMessagesOptionalParams) -> bool { .is_some_and(|config| config.format.is_some()) } -fn messages_carry_output_config(messages: &[AnthropicMessage]) -> bool { +fn messages_carry_output_config(messages: &[Message]) -> bool { messages .iter() .any(|message| message.extra.contains_key("output_config")) @@ -217,9 +214,9 @@ fn unsupported_param(model: &str, param: &str, value: &str, hint: &str) -> Error } fn drop_unsupported_params( - request: AnthropicMessagesRequest, + request: MessagesRequest, context: &MessagesTransformContext, -) -> Result { +) -> Result { let capabilities = &context.thinking.capabilities; let model = request.model.clone(); let reject = |param: &str, value: String, hint: &str| -> Result<(), Error> { @@ -237,8 +234,8 @@ fn drop_unsupported_params( _ => params.speed.clone(), }; if capabilities.supports_sampling_params { - return Ok(AnthropicMessagesRequest { - params: AnthropicMessagesOptionalParams { speed, ..params }, + return Ok(MessagesRequest { + params: MessagesOptionalParams { speed, ..params }, ..request }); } @@ -259,8 +256,8 @@ fn drop_unsupported_params( if let Some(top_k) = params.top_k { reject("top_k", json!(top_k).to_string(), "")?; } - Ok(AnthropicMessagesRequest { - params: AnthropicMessagesOptionalParams { + Ok(MessagesRequest { + params: MessagesOptionalParams { speed, temperature, top_p: None, @@ -366,7 +363,7 @@ mod tests { ) } - fn request(fields: Value) -> AnthropicMessagesRequest { + fn request(fields: Value) -> MessagesRequest { serde_json::from_value(body(fields)).unwrap() } diff --git a/litellm-rust/crates/llms/src/aws_textract/ocr/analyze_transformation.rs b/litellm-rust/crates/llms/src/aws_textract/ocr/analyze_transformation.rs index 1defec654bd..5506c86c17f 100644 --- a/litellm-rust/crates/llms/src/aws_textract/ocr/analyze_transformation.rs +++ b/litellm-rust/crates/llms/src/aws_textract/ocr/analyze_transformation.rs @@ -12,10 +12,10 @@ use crate::base_llm::ocr::{ error::Error, handler::OcrClient, transformation::{ - BaseOcrConfig, LiteLLMOcrResponse, OcrDocument, OcrRequestContext, OcrResponseFormat, - PreparedOcrRequest, decode_and_normalize_response, + BaseOcrConfig, OcrRequestContext, PreparedOcrRequest, decode_and_normalize_response, }, }; +use litellm_llms_types::formats::ocr::{LiteLLMOcrResponse, OcrDocument, OcrResponseFormat}; const DEFAULT_FEATURE_TYPES: [FeatureType; 2] = [FeatureType::Layout, FeatureType::Tables]; diff --git a/litellm-rust/crates/llms/src/aws_textract/ocr/common_utils.rs b/litellm-rust/crates/llms/src/aws_textract/ocr/common_utils.rs index 678104be982..90f5b97322f 100644 --- a/litellm-rust/crates/llms/src/aws_textract/ocr/common_utils.rs +++ b/litellm-rust/crates/llms/src/aws_textract/ocr/common_utils.rs @@ -7,11 +7,9 @@ use strum::{EnumString, IntoStaticStr, VariantNames}; use crate::base_llm::ocr::{ document::{InlineDocument, inline_remote_document}, error::Error, - transformation::{ - LiteLLMOcrResponse, OcrDocument, OcrEnvironment, OcrPage, OcrRequestContext, OcrUsageInfo, - PreparedOcrRequest, - }, + transformation::{OcrEnvironment, OcrRequestContext, PreparedOcrRequest}, }; +use litellm_llms_types::formats::ocr::{LiteLLMOcrResponse, OcrDocument, OcrPage, OcrUsageInfo}; const TEXTRACT_SERVICE: &str = "textract"; const AWS_JSON_CONTENT_TYPE: &str = "application/x-amz-json-1.1"; diff --git a/litellm-rust/crates/llms/src/aws_textract/ocr/transformation.rs b/litellm-rust/crates/llms/src/aws_textract/ocr/transformation.rs index 3b4f8e5a7d9..6bef577e6f8 100644 --- a/litellm-rust/crates/llms/src/aws_textract/ocr/transformation.rs +++ b/litellm-rust/crates/llms/src/aws_textract/ocr/transformation.rs @@ -10,10 +10,10 @@ use crate::base_llm::ocr::{ error::Error, handler::OcrClient, transformation::{ - BaseOcrConfig, LiteLLMOcrResponse, OcrDocument, OcrRequestContext, OcrResponseFormat, - PreparedOcrRequest, decode_and_normalize_response, + BaseOcrConfig, OcrRequestContext, PreparedOcrRequest, decode_and_normalize_response, }, }; +use litellm_llms_types::formats::ocr::{LiteLLMOcrResponse, OcrDocument, OcrResponseFormat}; #[derive(Debug, Deserialize, Serialize)] pub struct DetectDocumentTextRequest { diff --git a/litellm-rust/crates/llms/src/azure_ai/messages/AGENTS.md b/litellm-rust/crates/llms/src/azure_ai/messages/AGENTS.md index 6d3bbb866bc..a0d382bd064 100644 --- a/litellm-rust/crates/llms/src/azure_ai/messages/AGENTS.md +++ b/litellm-rust/crates/llms/src/azure_ai/messages/AGENTS.md @@ -1,3 +1,3 @@ -This directory owns Azure's Messages adapter: its endpoints, authentication policy, headers, and transformations. Implement the shared adapter contract from `base_llm/messages`, consume API data contracts from `litellm-types::messages`, and leave call orchestration to `core/src/messages` +This directory owns Azure's Messages adapter: its endpoints, authentication policy, headers, and transformations. Implement the shared adapter contract from `base_llm/messages`, consume API data contracts from `litellm-llms-types::formats::messages`, and leave call orchestration to `core/src/messages` The Claude adapter may explicitly reuse payload policy from `anthropic/messages` when it applies to Azure's Claude backend. Keep Azure-specific differences here. Sharing that helper does not make Anthropic policy a format-wide default or justify a dependency from `base_llm/messages` on provider implementations diff --git a/litellm-rust/crates/llms/src/azure_ai/messages/transformation.rs b/litellm-rust/crates/llms/src/azure_ai/messages/transformation.rs index 464fecc05a7..ffd079b2d04 100644 --- a/litellm-rust/crates/llms/src/azure_ai/messages/transformation.rs +++ b/litellm-rust/crates/llms/src/azure_ai/messages/transformation.rs @@ -1,8 +1,8 @@ use litellm_auth::{CredentialPlacement, SecretValue}; use litellm_http::request::{has_bearer_auth, has_header}; -use litellm_types::llms::anthropic_messages::anthropic_request::{ - AnthropicMessage, AnthropicMessagesOptionalParams, AnthropicMessagesRequest, CacheControl, - ContentBlock, MessageContent, SystemPrompt, +use litellm_llms_types::formats::messages::{ + CacheControl, ContentBlock, Message, MessageContent, MessagesOptionalParams, MessagesRequest, + SystemPrompt, }; use crate::{ @@ -21,7 +21,7 @@ use crate::{ messages::{ context::MessagesTransformContext, normalization::fold_system_role_messages, - transformation::{BaseAnthropicMessagesConfig, MESSAGES_PATH_SUFFIX}, + transformation::{BaseMessagesConfig, MESSAGES_PATH_SUFFIX}, }, }, }; @@ -34,12 +34,12 @@ pub struct AzureAnthropicMessagesConfig; pub const AZURE_ANTHROPIC_MESSAGES_CONFIG: AzureAnthropicMessagesConfig = AzureAnthropicMessagesConfig; -impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig { +impl BaseMessagesConfig for AzureAnthropicMessagesConfig { fn shape_request( &self, - request: AnthropicMessagesRequest, + request: MessagesRequest, reasoning_auto_summary: bool, - ) -> Result { + ) -> Result { shape_anthropic_messages_request(request, reasoning_auto_summary) } @@ -54,18 +54,18 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig { fn transform_anthropic_messages_request( &self, - request: AnthropicMessagesRequest, + request: MessagesRequest, context: &MessagesTransformContext, - ) -> Result { + ) -> Result { let request = fold_system_role_messages(request); transform_messages_request( - AnthropicMessagesRequest { + MessagesRequest { messages: request .messages .into_iter() .map(strip_scope_from_message) .collect(), - params: AnthropicMessagesOptionalParams { + params: MessagesOptionalParams { system: request.params.system.map(strip_scope_from_system), ..request.params }, @@ -105,7 +105,7 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig { DEFAULT_HEADERS } - fn request_headers(&self, headers: Headers, request: &AnthropicMessagesRequest) -> Headers { + fn request_headers(&self, headers: Headers, request: &MessagesRequest) -> Headers { update_headers_with_anthropic_beta(headers, request) } } @@ -148,8 +148,8 @@ fn strip_scope_from_system(system: SystemPrompt) -> SystemPrompt { } } -fn strip_scope_from_message(message: AnthropicMessage) -> AnthropicMessage { - AnthropicMessage { +fn strip_scope_from_message(message: Message) -> Message { + Message { content: match message.content { MessageContent::Blocks(blocks) => { MessageContent::Blocks(blocks.into_iter().map(strip_scope_from_block).collect()) @@ -162,7 +162,7 @@ fn strip_scope_from_message(message: AnthropicMessage) -> AnthropicMessage { #[cfg(test)] mod tests { - use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; + use litellm_llms_types::formats::messages::MessagesResponse; use rstest::rstest; use serde_json::json; @@ -171,11 +171,11 @@ mod tests { use super::*; use crate::base_llm::messages::context::MessagesModelCapabilities; - fn request_from(value: serde_json::Value) -> AnthropicMessagesRequest { + fn request_from(value: serde_json::Value) -> MessagesRequest { serde_json::from_value(value).expect("valid request") } - fn to_value(request: AnthropicMessagesRequest) -> serde_json::Value { + fn to_value(request: MessagesRequest) -> serde_json::Value { serde_json::to_value(request).expect("serializable request") } @@ -518,7 +518,7 @@ mod tests { #[test] fn transform_request_rejects_non_object_body() { - let err = serde_json::from_value::(json!("bad")) + let err = serde_json::from_value::(json!("bad")) .expect_err("non-object body should error"); assert!(err.is_data()); } @@ -576,7 +576,7 @@ mod tests { #[test] fn transform_response_passes_through() { - let response: AnthropicMessagesResponse = serde_json::from_value(json!({ + let response: MessagesResponse = serde_json::from_value(json!({ "id": "msg_1", "type": "message", "role": "assistant", diff --git a/litellm-rust/crates/llms/src/azure_ai/ocr/cohere_parse_transformation.rs b/litellm-rust/crates/llms/src/azure_ai/ocr/cohere_parse_transformation.rs index 3f9b97c3149..72e321547e3 100644 --- a/litellm-rust/crates/llms/src/azure_ai/ocr/cohere_parse_transformation.rs +++ b/litellm-rust/crates/llms/src/azure_ai/ocr/cohere_parse_transformation.rs @@ -6,15 +6,13 @@ use crate::{ document::{inline_remote_document, validate_inline_document}, error::Error, handler::OcrClient, - transformation::{ - BaseOcrConfig, LiteLLMOcrResponse, OcrDocument, OcrRequestContext, OcrResponseFormat, - PreparedOcrRequest, - }, + transformation::{BaseOcrConfig, OcrRequestContext, PreparedOcrRequest}, }, cohere::ocr::transformation::{ CohereOptions, CohereParseConfig, CohereRequest, validate_document, }, }; +use litellm_llms_types::formats::ocr::{LiteLLMOcrResponse, OcrDocument, OcrResponseFormat}; pub const AZURE_COHERE_PARSE_PATH: [&str; 4] = ["providers", "cohere", "v2", "parse"]; diff --git a/litellm-rust/crates/llms/src/azure_ai/ocr/document_intelligence/transformation.rs b/litellm-rust/crates/llms/src/azure_ai/ocr/document_intelligence/transformation.rs index 5ee5ab3be94..ab273d9dbe9 100644 --- a/litellm-rust/crates/llms/src/azure_ai/ocr/document_intelligence/transformation.rs +++ b/litellm-rust/crates/llms/src/azure_ai/ocr/document_intelligence/transformation.rs @@ -3,11 +3,8 @@ use std::{collections::BTreeSet, time::Duration}; use base64::{Engine, engine::general_purpose::STANDARD}; use litellm_auth::{InputSource, Sourced}; use litellm_auth_azure::{AzureAuthInputs, SECRET_NAMES as AZURE_AUTH_SECRET_NAMES}; -use litellm_core_utils::{ - call_arguments::CallArguments, - serde_compat::{FiniteF64, LaxI64}, - url_utils::ApiUrl, -}; +use litellm_core_utils::{call_arguments::CallArguments, url_utils::ApiUrl}; +use litellm_llms_types::serde_compat::{FiniteF64, LaxI64}; use reqwest::Url; use serde::{Deserialize, Deserializer, Serialize}; use serde_json::{Map, Value}; @@ -20,12 +17,14 @@ use crate::base_llm::ocr::{ handler::{CallHooks, OcrClient, read_json_response}, settings::OcrSettings, transformation::{ - BaseOcrConfig, DecodedOcrResponse, LiteLLMOcrResponse, OCR_INLINE_MAX_BYTES, - OCR_POLL_RETRY_SECS, OcrConnection, OcrCredentialInputs, OcrDocument, OcrPage, - OcrPageDimensions, OcrResponseContext, OcrResponseFormat, OcrUsageInfo, PreparedOcrRequest, + BaseOcrConfig, DecodedOcrResponse, OCR_INLINE_MAX_BYTES, OCR_POLL_RETRY_SECS, + OcrConnection, OcrCredentialInputs, OcrResponseContext, PreparedOcrRequest, ResolvedOcrCredentials, decode_and_normalize_response, decode_response, }, }; +use litellm_llms_types::formats::ocr::{ + LiteLLMOcrResponse, OcrDocument, OcrPage, OcrPageDimensions, OcrResponseFormat, OcrUsageInfo, +}; const AZURE_DI_SUBSCRIPTION_HEADER: &str = "Ocp-Apim-Subscription-Key"; const AZURE_DI_DEFAULT_WIDTH: f64 = 8.5; diff --git a/litellm-rust/crates/llms/src/azure_ai/ocr/transformation.rs b/litellm-rust/crates/llms/src/azure_ai/ocr/transformation.rs index 21fc6e65207..83d88bbb3cb 100644 --- a/litellm-rust/crates/llms/src/azure_ai/ocr/transformation.rs +++ b/litellm-rust/crates/llms/src/azure_ai/ocr/transformation.rs @@ -1,6 +1,5 @@ use litellm_auth::{InputSource, Sourced}; -use litellm_auth_azure::AzureAuthInputs; -use litellm_auth_azure::SECRET_NAMES as AZURE_AUTH_SECRET_NAMES; +use litellm_auth_azure::{AzureAuthInputs, SECRET_NAMES as AZURE_AUTH_SECRET_NAMES}; use litellm_core_utils::{call_arguments::CallArguments, params::OpaqueParams, url_utils::ApiUrl}; use serde_json::Value; @@ -9,13 +8,11 @@ use crate::{ document::{inline_remote_document, validate_inline_document}, error::Error, handler::OcrClient, - transformation::{ - BaseOcrConfig, LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrRequestContext, - OcrResponseFormat, PreparedOcrRequest, - }, + transformation::{BaseOcrConfig, OcrConnection, OcrRequestContext, PreparedOcrRequest}, }, mistral::ocr::transformation::{MistralOcrConfig, MistralOcrRequest}, }; +use litellm_llms_types::formats::ocr::{LiteLLMOcrResponse, OcrDocument, OcrResponseFormat}; pub const AZURE_AI_OCR_PATH: [&str; 4] = ["providers", "mistral", "azure", "ocr"]; diff --git a/litellm-rust/crates/llms/src/base_llm/audio_transcription/transformation.rs b/litellm-rust/crates/llms/src/base_llm/audio_transcription/transformation.rs index 6520a215b5e..3d56d4192ab 100644 --- a/litellm-rust/crates/llms/src/base_llm/audio_transcription/transformation.rs +++ b/litellm-rust/crates/llms/src/base_llm/audio_transcription/transformation.rs @@ -1,4 +1,4 @@ -use litellm_types::audio_transcription::AudioTranscriptionResponseData; +use litellm_llms_types::formats::audio_transcription::AudioTranscriptionResponseData; use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; diff --git a/litellm-rust/crates/llms/src/base_llm/chat/streaming.rs b/litellm-rust/crates/llms/src/base_llm/chat/streaming.rs index b9d715bcd68..2b4a8d084e6 100644 --- a/litellm-rust/crates/llms/src/base_llm/chat/streaming.rs +++ b/litellm-rust/crates/llms/src/base_llm/chat/streaming.rs @@ -1,7 +1,7 @@ use std::collections::HashMap; use futures_util::{StreamExt, stream::BoxStream}; -use litellm_types::utils::ChatCompletionChunk; +use litellm_llms_types::formats::chat_completions::ChatCompletionChunk; use crate::{ Error, diff --git a/litellm-rust/crates/llms/src/base_llm/chat/transformation.rs b/litellm-rust/crates/llms/src/base_llm/chat/transformation.rs index cdf6d47d8f8..30bff76651e 100644 --- a/litellm-rust/crates/llms/src/base_llm/chat/transformation.rs +++ b/litellm-rust/crates/llms/src/base_llm/chat/transformation.rs @@ -1,6 +1,5 @@ -use litellm_types::{ - llms::openai::{ChatMessage, ChatMessageContent}, - utils::ChatCompletionsResponse, +use litellm_llms_types::formats::chat_completions::{ + ChatCompletionsResponse, ChatMessage, ChatMessageContent, }; use serde_json::{Map, Value}; diff --git a/litellm-rust/crates/llms/src/base_llm/messages/AGENTS.md b/litellm-rust/crates/llms/src/base_llm/messages/AGENTS.md index 06d051eb521..228b2853b66 100644 --- a/litellm-rust/crates/llms/src/base_llm/messages/AGENTS.md +++ b/litellm-rust/crates/llms/src/base_llm/messages/AGENTS.md @@ -1,4 +1,4 @@ -This directory owns the shared Messages provider adapter contract, its execution inputs such as `MessagesTransformContext`, and provider-independent transformation machinery. Public request, response, content-block, and event schemas belong in `litellm-types::messages`. Call orchestration belongs in `core/src/messages`, and provider implementations belong in `llms/src//messages` +This directory owns the shared Messages provider adapter contract, its execution inputs such as `MessagesTransformContext`, and provider-independent transformation machinery. Public request, response, content-block, and event schemas belong in `litellm-llms-types::formats::messages`. Call orchestration belongs in `core/src/messages`, and provider implementations belong in `llms/src//messages` Do not import provider implementations or embed their policy in shared trait defaults, normalization, or context defaults. A context carries inputs the shared adapter contract needs, not every provider's settings. Thinking-budget choices and model-specific restrictions do not become format rules merely because several providers host Claude diff --git a/litellm-rust/crates/llms/src/base_llm/messages/normalization.rs b/litellm-rust/crates/llms/src/base_llm/messages/normalization.rs index bcbc08778ef..bddd829304e 100644 --- a/litellm-rust/crates/llms/src/base_llm/messages/normalization.rs +++ b/litellm-rust/crates/llms/src/base_llm/messages/normalization.rs @@ -1,6 +1,5 @@ -use litellm_types::llms::anthropic_messages::anthropic_request::{ - AnthropicMessage, AnthropicMessagesOptionalParams, AnthropicMessagesRequest, ContentBlock, - MessageContent, SystemPrompt, +use litellm_llms_types::formats::messages::{ + ContentBlock, Message, MessageContent, MessagesOptionalParams, MessagesRequest, SystemPrompt, }; const SYSTEM_ROLE: &str = "system"; @@ -20,12 +19,12 @@ fn system_into_blocks(system: Option) -> Vec { } } -pub fn fold_system_role_messages(request: AnthropicMessagesRequest) -> AnthropicMessagesRequest { +pub fn fold_system_role_messages(request: MessagesRequest) -> MessagesRequest { if !request.messages.iter().any(|msg| msg.role == SYSTEM_ROLE) { return request; } - let (system_messages, chat_messages): (Vec, Vec) = request + let (system_messages, chat_messages): (Vec, Vec) = request .messages .into_iter() .partition(|msg| msg.role == SYSTEM_ROLE); @@ -39,9 +38,9 @@ pub fn fold_system_role_messages(request: AnthropicMessagesRequest) -> Anthropic ) .collect(); - AnthropicMessagesRequest { + MessagesRequest { messages: chat_messages, - params: AnthropicMessagesOptionalParams { + params: MessagesOptionalParams { system: (!folded_system.is_empty()).then_some(SystemPrompt::Blocks(folded_system)), ..request.params }, diff --git a/litellm-rust/crates/llms/src/base_llm/messages/streaming.rs b/litellm-rust/crates/llms/src/base_llm/messages/streaming.rs index afd2bdae6bc..0989d297d42 100644 --- a/litellm-rust/crates/llms/src/base_llm/messages/streaming.rs +++ b/litellm-rust/crates/llms/src/base_llm/messages/streaming.rs @@ -1,7 +1,7 @@ use bytes::Bytes; use futures_util::{StreamExt, stream::BoxStream}; use litellm_framing::{frames, sse::SseCodec}; -use litellm_types::messages::streaming::MessagesStreamEvent; +use litellm_llms_types::formats::messages::streaming::MessagesStreamEvent; use crate::Error; pub use crate::base_llm::base_model_iterator::ByteStream; @@ -38,7 +38,9 @@ pub fn encode_anthropic_sse(event: &MessagesStreamEvent) -> Result #[cfg(test)] mod tests { use futures_util::{StreamExt, TryStreamExt, stream}; - use litellm_types::messages::streaming::{MessagesContentBlockDelta, MessagesStreamUsage}; + use litellm_llms_types::formats::messages::streaming::{ + MessagesContentBlockDelta, MessagesStreamUsage, + }; use serde_json::json; use super::*; diff --git a/litellm-rust/crates/llms/src/base_llm/messages/transformation.rs b/litellm-rust/crates/llms/src/base_llm/messages/transformation.rs index 58ac85026ab..2f2d3bcf909 100644 --- a/litellm-rust/crates/llms/src/base_llm/messages/transformation.rs +++ b/litellm-rust/crates/llms/src/base_llm/messages/transformation.rs @@ -1,6 +1,4 @@ -use litellm_types::llms::anthropic_messages::{ - anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse, -}; +use litellm_llms_types::formats::messages::{MessagesRequest, MessagesResponse}; use super::context::MessagesTransformContext; @@ -9,12 +7,12 @@ use crate::{Error, base_llm::messages::streaming::StreamDecoder}; pub const MESSAGES_PATH_SUFFIX: &str = "/v1/messages"; -pub trait BaseAnthropicMessagesConfig: Sync { +pub trait BaseMessagesConfig: Sync { fn shape_request( &self, - request: AnthropicMessagesRequest, + request: MessagesRequest, _reasoning_auto_summary: bool, - ) -> Result { + ) -> Result { Ok(request) } @@ -36,17 +34,17 @@ pub trait BaseAnthropicMessagesConfig: Sync { fn transform_anthropic_messages_request( &self, - request: AnthropicMessagesRequest, + request: MessagesRequest, _context: &MessagesTransformContext, - ) -> Result { + ) -> Result { Ok(request) } fn transform_anthropic_messages_response( &self, _model: &str, - response: AnthropicMessagesResponse, - ) -> Result { + response: MessagesResponse, + ) -> Result { Ok(response) } @@ -74,7 +72,7 @@ pub trait BaseAnthropicMessagesConfig: Sync { &[("content-type", "application/json")] } - fn request_headers(&self, headers: Headers, _request: &AnthropicMessagesRequest) -> Headers { + fn request_headers(&self, headers: Headers, _request: &MessagesRequest) -> Headers { headers } } @@ -87,7 +85,7 @@ mod tests { struct DefaultsConfig; - impl BaseAnthropicMessagesConfig for DefaultsConfig { + impl BaseMessagesConfig for DefaultsConfig { fn secret_names(&self) -> &'static [&'static str] { &[] } @@ -117,7 +115,7 @@ mod tests { #[test] fn default_request_headers_are_the_given_headers() { - let request: AnthropicMessagesRequest = serde_json::from_value(serde_json::json!({ + let request: MessagesRequest = serde_json::from_value(serde_json::json!({ "model": "claude", "max_tokens": 16, "speed": "fast", @@ -134,7 +132,7 @@ mod tests { #[case::disabled(false)] #[case::enabled(true)] fn default_shaping_preserves_provider_policy_inputs(#[case] reasoning_auto_summary: bool) { - let request: AnthropicMessagesRequest = serde_json::from_value(serde_json::json!({ + let request: MessagesRequest = serde_json::from_value(serde_json::json!({ "model": "test-model", "metadata": {"user_id": 7, "extra": "keep"}, "thinking": {"type": "enabled", "budget_tokens": 64}, diff --git a/litellm-rust/crates/llms/src/base_llm/ocr/document.rs b/litellm-rust/crates/llms/src/base_llm/ocr/document.rs index 9bcaad353ab..b32ea4bf73a 100644 --- a/litellm-rust/crates/llms/src/base_llm/ocr/document.rs +++ b/litellm-rust/crates/llms/src/base_llm/ocr/document.rs @@ -8,8 +8,9 @@ use reqwest::Url; use crate::base_llm::ocr::{ error::Error, - transformation::{OCR_INLINE_MAX_BYTES, OCR_MAX_FETCH_REDIRECTS, OcrConnection, OcrDocument}, + transformation::{OCR_INLINE_MAX_BYTES, OCR_MAX_FETCH_REDIRECTS, OcrConnection}, }; +use litellm_llms_types::formats::ocr::OcrDocument; pub struct InlineDocument<'a>(DataUrl<'a>); diff --git a/litellm-rust/crates/llms/src/base_llm/ocr/handler.rs b/litellm-rust/crates/llms/src/base_llm/ocr/handler.rs index f34af90df0e..a4c2d465cbd 100644 --- a/litellm-rust/crates/llms/src/base_llm/ocr/handler.rs +++ b/litellm-rust/crates/llms/src/base_llm/ocr/handler.rs @@ -17,10 +17,11 @@ use crate::base_llm::ocr::{ error::Error, settings::OcrSettings, transformation::{ - BaseOcrConfig, DecodedOcrResponse, LiteLLMOcrResponse, OcrDocument, OcrResponseContext, - PreparedOcrRequest, decode_request_value, decode_response, + BaseOcrConfig, DecodedOcrResponse, OcrResponseContext, PreparedOcrRequest, + decode_request_value, decode_response, }, }; +use litellm_llms_types::formats::ocr::{LiteLLMOcrResponse, OcrDocument}; use litellm_secrets::source::SecretSource; /// The route's view of one call, handed to provider code that has to reach the diff --git a/litellm-rust/crates/llms/src/base_llm/ocr/transformation.rs b/litellm-rust/crates/llms/src/base_llm/ocr/transformation.rs index 3f1b260bb13..2da0e1abbfc 100644 --- a/litellm-rust/crates/llms/src/base_llm/ocr/transformation.rs +++ b/litellm-rust/crates/llms/src/base_llm/ocr/transformation.rs @@ -1,19 +1,15 @@ +use litellm_llms_types::formats::ocr::{LiteLLMOcrResponse, OcrDocument, OcrResponseFormat}; use std::{collections::BTreeMap, future::Future, sync::Arc, time::Duration}; use litellm_auth::{InputSource, SecretValue, Sourced, TokenProviderHandle}; -use litellm_core_utils::{ - call_arguments::CallArguments, - serde_compat::{FiniteF64, LaxI64}, - settings::ProcessEnvironment, -}; +use litellm_core_utils::{call_arguments::CallArguments, settings::ProcessEnvironment}; use litellm_http::outbound::{OutboundRequest, RequestSigner}; use litellm_secrets::source::Secrets; use serde::{ - Deserialize, Serialize, + Serialize, de::{DeserializeOwned, IntoDeserializer}, }; use serde_json::{Map, Value}; -use serde_with::serde_as; use crate::base_llm::ocr::{ error::Error, @@ -26,66 +22,6 @@ pub const OCR_INLINE_MAX_BYTES: usize = 50 * 1024 * 1024; pub const OCR_MAX_FETCH_REDIRECTS: usize = 10; pub const OCR_POLL_RETRY_SECS: u64 = 2; -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -#[serde(tag = "type")] -pub enum OcrDocument { - #[serde(rename = "document_url")] - DocumentUrl { - document_url: String, - #[serde(flatten)] - extra_fields: BTreeMap>, - }, - #[serde(rename = "image_url")] - ImageUrl { - image_url: String, - #[serde(flatten)] - extra_fields: BTreeMap>, - }, -} - -impl OcrDocument { - pub fn source(&self) -> &str { - match self { - Self::DocumentUrl { document_url, .. } => document_url, - Self::ImageUrl { image_url, .. } => image_url, - } - } - - pub fn is_remote(&self) -> bool { - let source = self.source(); - source.starts_with("http://") || source.starts_with("https://") - } - - pub fn with_source(self, source: String) -> Self { - match self { - Self::DocumentUrl { extra_fields, .. } => Self::DocumentUrl { - document_url: source, - extra_fields, - }, - Self::ImageUrl { extra_fields, .. } => Self::ImageUrl { - image_url: source, - extra_fields, - }, - } - } -} - -impl TryFrom for OcrDocument { - type Error = Error; - - fn try_from(value: Value) -> Result { - decode_request_value(value, "document") - } -} - -#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)] -#[serde(rename_all = "lowercase")] -pub enum OcrResponseFormat { - #[default] - Litellm, - Native, -} - #[derive(Clone, Default)] pub struct OcrCredentialInputs { pub api_key: Option>, @@ -249,95 +185,6 @@ pub fn response_format(optional_params: &CallArguments) -> Result, - #[serde_as(deserialize_as = "Option")] - pub height: Option, - #[serde_as(deserialize_as = "Option")] - pub width: Option, -} - -#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] -pub struct OcrPageImage { - pub image_base64: Option, - pub bbox: Option>, - #[serde(flatten)] - pub extra_fields: Map, -} - -#[serde_as] -#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] -pub struct OcrPage { - #[serde_as(deserialize_as = "LaxI64")] - pub index: i64, - pub markdown: String, - pub images: Option>, - pub dimensions: Option, - #[serde(flatten)] - pub extra_fields: Map, -} - -#[serde_as] -#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] -pub struct OcrUsageInfo { - #[serde_as(deserialize_as = "Option")] - pub pages_processed: Option, - #[serde_as(deserialize_as = "Option")] - pub pages_processed_annotation: Option, - #[serde_as(deserialize_as = "Option")] - pub credits: Option, - #[serde_as(deserialize_as = "Option")] - pub doc_size_bytes: Option, - #[serde(flatten)] - pub extra_fields: Map, -} - -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -pub struct LiteLLMOcrResponse { - pub pages: Vec, - pub model: String, - pub document_annotation: Option, - pub usage_info: Option, - pub content: Option, - pub tables: Option>>, - #[serde(rename = "keyValuePairs")] - pub key_value_pairs: Option>>, - #[serde(default = "ocr_object")] - pub object: String, - #[serde(flatten)] - pub extra_fields: Map, - #[serde(skip_serializing_if = "Option::is_none")] - pub provider_native_response: Option>, -} - -impl LiteLLMOcrResponse { - pub fn new(model: impl Into, pages: Vec) -> Self { - Self { - pages, - model: model.into(), - document_annotation: None, - usage_info: None, - content: None, - tables: None, - key_value_pairs: None, - object: ocr_object(), - extra_fields: Map::new(), - provider_native_response: None, - } - } - - pub fn into_json(self) -> Value { - serde_json::to_value(self).expect("OCR response fields are JSON-compatible") - } -} - -fn ocr_object() -> String { - "ocr".into() -} - #[derive(Debug)] pub struct DecodedOcrResponse { pub data: T, @@ -591,7 +438,6 @@ pub fn decode_and_normalize_response( #[cfg(test)] mod tests { - use serde_json::json; use super::*; @@ -620,94 +466,4 @@ mod tests { Duration::from_secs(5) ); } - - #[test] - fn normalized_response_rejects_invalid_shared_fields() { - for fields in [ - json!({"pages":[{}]}), - json!({"pages":[{"index":0,"markdown":false}]}), - json!({"pages":[{"index":0,"markdown":"","images":[{"bbox":[]}]}]}), - json!({"usage_info":{"pages_processed":1.5}}), - json!({"tables":[false]}), - json!({"keyValuePairs":[[]]}), - json!({"provider_native_response":[]}), - ] { - let payload: Map = json!({"model":"model", "pages":[]}) - .as_object() - .unwrap() - .iter() - .chain(fields.as_object().unwrap()) - .map(|(key, value)| (key.clone(), value.clone())) - .collect(); - assert!(serde_json::from_value::(Value::Object(payload)).is_err()); - } - assert!( - serde_json::from_value::(json!({ - "type":"image_url", "image_url":"https://example.com/image", "detail":42 - })) - .is_err() - ); - } - - #[test] - fn numeric_coercion_preserves_integer_precision_and_rejects_fractional_values() { - for (value, expected) in [ - (json!("9007199254740993.0"), 9_007_199_254_740_993), - (json!("+2.000"), 2), - (json!("1_000"), 1000), - (json!(true), 1), - (json!(2.0), 2), - ] { - let page: OcrPage = - serde_json::from_value(json!({"index":value,"markdown":""})).unwrap(); - assert_eq!(page.index, expected); - } - for value in [ - json!("1e2"), - json!(".0"), - json!("2."), - json!("_2"), - json!("2__0"), - json!(2.5), - json!(null), - ] { - assert!( - serde_json::from_value::(json!({"index":value,"markdown":""})).is_err() - ); - } - } - - #[rstest::rstest] - #[case::document_url("document_url", "document_name", "application/pdf")] - #[case::image_url("image_url", "detail", "image/png")] - fn document_variants_preserve_provider_fields_when_rewriting_sources( - #[case] kind: &str, - #[case] field: &str, - #[case] mime_type: &str, - #[values(json!("kept"), Value::Null)] extra: Value, - ) { - let original = "https://example.com/input"; - let replacement = format!("data:{mime_type};base64,AA=="); - let document: OcrDocument = - serde_json::from_value(json!({"type": kind, kind: original, field: extra})).unwrap(); - assert_eq!(document.source(), original); - assert_eq!( - serde_json::to_value(document.with_source(replacement.clone())).unwrap(), - json!({"type": kind, kind: replacement, field: extra}) - ); - } - - #[test] - fn response_serialization_flattens_extra_fields_and_omits_absent_native_response() { - let response = LiteLLMOcrResponse { - extra_fields: json!({"provider_field":"kept"}) - .as_object() - .unwrap() - .clone(), - ..LiteLLMOcrResponse::new("model", vec![]) - }; - let serialized = response.into_json(); - assert_eq!(serialized["provider_field"], "kept"); - assert!(serialized.get("provider_native_response").is_none()); - } } diff --git a/litellm-rust/crates/llms/src/base_llm/responses/transformation.rs b/litellm-rust/crates/llms/src/base_llm/responses/transformation.rs index 3263672edee..30899692fb6 100644 --- a/litellm-rust/crates/llms/src/base_llm/responses/transformation.rs +++ b/litellm-rust/crates/llms/src/base_llm/responses/transformation.rs @@ -1,5 +1,6 @@ -use litellm_types::responses::main::ResponsesApiResponse; -use litellm_types::responses::streaming_websocket::ResponsesWsEvent; +use litellm_llms_types::formats::responses::{ + ResponsesApiResponse, streaming_websocket::ResponsesWsEvent, +}; use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; diff --git a/litellm-rust/crates/llms/src/bedrock/audio_transcription/mod.rs b/litellm-rust/crates/llms/src/bedrock/audio_transcription/mod.rs index f1a1a54828f..cee906e77a4 100644 --- a/litellm-rust/crates/llms/src/bedrock/audio_transcription/mod.rs +++ b/litellm-rust/crates/llms/src/bedrock/audio_transcription/mod.rs @@ -4,7 +4,7 @@ use litellm_auth_aws::{ resolve_bedrock_region, }; use litellm_core_utils::core_helpers::json_type_name; -use litellm_types::audio_transcription::AudioTranscriptionResponseData; +use litellm_llms_types::formats::audio_transcription::AudioTranscriptionResponseData; use serde::Deserialize; use serde_json::{Map, Value, json}; use strum::IntoStaticStr; diff --git a/litellm-rust/crates/llms/src/bedrock/chat/converse_transformation.rs b/litellm-rust/crates/llms/src/bedrock/chat/converse_transformation.rs index d2a4f2a0f46..3a13e388a4b 100644 --- a/litellm-rust/crates/llms/src/bedrock/chat/converse_transformation.rs +++ b/litellm-rust/crates/llms/src/bedrock/chat/converse_transformation.rs @@ -8,12 +8,9 @@ use litellm_core_utils::{ core_helpers::{finish_reason_for, unix_now, usage_from_parts}, prompt_templates::factory::{Conversation, TurnRole, build_conversation}, }; -use litellm_types::{ - llms::openai::{ChatMessage, ChatMessageContent}, - utils::{ - ChatCompletionsChoice, ChatCompletionsChoiceMessage, ChatCompletionsResponse, - ChatCompletionsUsage, - }, +use litellm_llms_types::formats::chat_completions::{ + ChatCompletionsChoice, ChatCompletionsChoiceMessage, ChatCompletionsResponse, + ChatCompletionsUsage, ChatMessage, ChatMessageContent, }; use serde::Deserialize; use serde_json::{Map, Value, json}; diff --git a/litellm-rust/crates/llms/src/bedrock/chat/invoke_handler.rs b/litellm-rust/crates/llms/src/bedrock/chat/invoke_handler.rs index 99d91d931c2..f5bdfb7fe30 100644 --- a/litellm-rust/crates/llms/src/bedrock/chat/invoke_handler.rs +++ b/litellm-rust/crates/llms/src/bedrock/chat/invoke_handler.rs @@ -5,7 +5,7 @@ use litellm_framing::{ aws_event_stream::{AwsEventStreamCodec, Message}, frames, }; -use litellm_types::messages::streaming::MessagesStreamEvent; +use litellm_llms_types::formats::messages::streaming::MessagesStreamEvent; use serde::Deserialize; use serde_json::Value; @@ -103,7 +103,7 @@ mod tests { use base64::engine::general_purpose::STANDARD; use bytes::Bytes; use futures_util::TryStreamExt; - use litellm_types::messages::streaming::MessagesContentBlockDelta; + use litellm_llms_types::formats::messages::streaming::MessagesContentBlockDelta; use super::*; use crate::base_llm::messages::streaming::anthropic_sse_event_stream; diff --git a/litellm-rust/crates/llms/src/bedrock/messages/AGENTS.md b/litellm-rust/crates/llms/src/bedrock/messages/AGENTS.md index dc6072be982..6744ba4c72e 100644 --- a/litellm-rust/crates/llms/src/bedrock/messages/AGENTS.md +++ b/litellm-rust/crates/llms/src/bedrock/messages/AGENTS.md @@ -1,3 +1,3 @@ -This directory owns Bedrock's Messages adapter: its endpoints, authentication policy, wire adaptation, and response decoding. Implement the shared adapter contract from `base_llm/messages`, consume API data contracts from `litellm-types::messages`, and leave call orchestration to `core/src/messages` +This directory owns Bedrock's Messages adapter: its endpoints, authentication policy, wire adaptation, and response decoding. Implement the shared adapter contract from `base_llm/messages`, consume API data contracts from `litellm-llms-types::formats::messages`, and leave call orchestration to `core/src/messages` The Claude adapter may explicitly reuse payload policy from `anthropic/messages` when it applies to Bedrock's Claude backend. Keep Bedrock-specific differences here. Sharing that helper does not make Anthropic policy a format-wide default or justify a dependency from `base_llm/messages` on provider implementations diff --git a/litellm-rust/crates/llms/src/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.rs b/litellm-rust/crates/llms/src/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.rs index 48192fa50cb..beae62ed065 100644 --- a/litellm-rust/crates/llms/src/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.rs +++ b/litellm-rust/crates/llms/src/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.rs @@ -1,7 +1,9 @@ use std::convert::Infallible; -use crate::anthropic::messages::handler::shape_anthropic_messages_request; -use crate::base_llm::messages::context::MessagesTransformContext; +use crate::{ + anthropic::messages::handler::shape_anthropic_messages_request, + base_llm::messages::context::MessagesTransformContext, +}; use futures_util::StreamExt; use litellm_auth::{CredentialPlacement, SecretValue}; use litellm_auth_aws::{ @@ -12,8 +14,10 @@ use litellm_auth_aws::{ }, resolve_bedrock_region, }; -use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest; -use litellm_types::messages::streaming::{MessagesStreamEvent, MessagesStreamUsage}; +use litellm_llms_types::formats::messages::{ + MessagesRequest, + streaming::{MessagesStreamEvent, MessagesStreamUsage}, +}; use serde_json::{Map, Value}; use crate::{ @@ -23,7 +27,7 @@ use crate::{ base_model_iterator::{StreamError, StreamTransformer, transform_stream}, messages::{ streaming::{ByteStream, EventStream, StreamDecoder}, - transformation::{BaseAnthropicMessagesConfig, Headers, ValidatedEnvironment}, + transformation::{BaseMessagesConfig, Headers, ValidatedEnvironment}, }, }, bedrock::chat::invoke_handler::{decode_invoke_anthropic_chunk, invoke_chunk_stream}, @@ -84,12 +88,12 @@ fn invoke_url( format!("{}/model/{model_id}/{path}", endpoint.trim_end_matches('/')) } -impl BaseAnthropicMessagesConfig for AmazonAnthropicClaudeMessagesConfig { +impl BaseMessagesConfig for AmazonAnthropicClaudeMessagesConfig { fn shape_request( &self, - request: AnthropicMessagesRequest, + request: MessagesRequest, reasoning_auto_summary: bool, - ) -> Result { + ) -> Result { shape_anthropic_messages_request(request, reasoning_auto_summary) } @@ -113,9 +117,9 @@ impl BaseAnthropicMessagesConfig for AmazonAnthropicClaudeMessagesConfig { fn transform_anthropic_messages_request( &self, - _request: AnthropicMessagesRequest, + _request: MessagesRequest, _context: &MessagesTransformContext, - ) -> Result { + ) -> Result { Err(Error::Unsupported( "Bedrock invoke messages request shaping", )) diff --git a/litellm-rust/crates/llms/src/cohere/ocr/transformation.rs b/litellm-rust/crates/llms/src/cohere/ocr/transformation.rs index c0bb4c60563..67b81ec6feb 100644 --- a/litellm-rust/crates/llms/src/cohere/ocr/transformation.rs +++ b/litellm-rust/crates/llms/src/cohere/ocr/transformation.rs @@ -1,8 +1,8 @@ use litellm_core_utils::{ call_arguments::{CallArguments, parse_options}, - serde_compat::LaxI64, url_utils::ApiUrl, }; +use litellm_llms_types::serde_compat::LaxI64; use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; use serde_with::serde_as; @@ -12,11 +12,13 @@ use crate::base_llm::ocr::{ error::Error, handler::OcrClient, transformation::{ - BaseOcrConfig, LiteLLMOcrResponse, OCR_INLINE_MAX_BYTES, OcrConnection, OcrDocument, - OcrPage, OcrPageImage, OcrResponseFormat, OcrUsageInfo, PreparedOcrRequest, + BaseOcrConfig, OCR_INLINE_MAX_BYTES, OcrConnection, PreparedOcrRequest, decode_and_normalize_response, decode_response_value, }, }; +use litellm_llms_types::formats::ocr::{ + LiteLLMOcrResponse, OcrDocument, OcrPage, OcrPageImage, OcrResponseFormat, OcrUsageInfo, +}; const COHERE_PARSE_API_BASE: &str = "https://api.cohere.com"; const COHERE_API_KEY_ENV: &str = "COHERE_API_KEY"; @@ -561,10 +563,10 @@ mod tests { #[rstest] fn response_types_documented_block_variants( #[values( - crate::base_llm::ocr::transformation::OcrResponseFormat::Litellm, - crate::base_llm::ocr::transformation::OcrResponseFormat::Native + litellm_llms_types::formats::ocr::OcrResponseFormat::Litellm, + litellm_llms_types::formats::ocr::OcrResponseFormat::Native )] - response_format: crate::base_llm::ocr::transformation::OcrResponseFormat, + response_format: litellm_llms_types::formats::ocr::OcrResponseFormat, ) { let payload = json!({ "pages": [{ @@ -634,10 +636,10 @@ mod tests { Some(1) ); match response_format { - crate::base_llm::ocr::transformation::OcrResponseFormat::Litellm => { + litellm_llms_types::formats::ocr::OcrResponseFormat::Litellm => { assert!(normalized.provider_native_response.is_none()); } - crate::base_llm::ocr::transformation::OcrResponseFormat::Native => { + litellm_llms_types::formats::ocr::OcrResponseFormat::Native => { assert_eq!( normalized.provider_native_response.as_ref(), payload.as_object() diff --git a/litellm-rust/crates/llms/src/mistral/ocr/transformation.rs b/litellm-rust/crates/llms/src/mistral/ocr/transformation.rs index 149e8056789..29d4d5c1610 100644 --- a/litellm-rust/crates/llms/src/mistral/ocr/transformation.rs +++ b/litellm-rust/crates/llms/src/mistral/ocr/transformation.rs @@ -6,10 +6,12 @@ use crate::base_llm::ocr::{ error::Error, handler::OcrClient, transformation::{ - BaseOcrConfig, LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrPage, OcrResponseFormat, - OcrUsageInfo, PreparedOcrRequest, decode_and_normalize_response, + BaseOcrConfig, OcrConnection, PreparedOcrRequest, decode_and_normalize_response, }, }; +use litellm_llms_types::formats::ocr::{ + LiteLLMOcrResponse, OcrDocument, OcrPage, OcrResponseFormat, OcrUsageInfo, +}; const MISTRAL_OCR_API_BASE: &str = "https://api.mistral.ai/v1"; @@ -326,7 +328,7 @@ mod tests { .transform_ocr_response( "model", raw, - crate::base_llm::ocr::transformation::OcrResponseFormat::Native, + litellm_llms_types::formats::ocr::OcrResponseFormat::Native, ) .unwrap(); assert_eq!(response.pages[0].index, 2); diff --git a/litellm-rust/crates/llms/src/openai/responses/transformation.rs b/litellm-rust/crates/llms/src/openai/responses/transformation.rs index ecb5f2f3a65..959be2a9a4a 100644 --- a/litellm-rust/crates/llms/src/openai/responses/transformation.rs +++ b/litellm-rust/crates/llms/src/openai/responses/transformation.rs @@ -1,5 +1,6 @@ -use litellm_types::responses::main::ResponsesApiResponse; -use litellm_types::responses::streaming_websocket::ResponsesWsEvent; +use litellm_llms_types::formats::responses::{ + ResponsesApiResponse, streaming_websocket::ResponsesWsEvent, +}; use serde_json::{Map, Value}; use litellm_auth::{CredentialPlacement, SecretValue}; diff --git a/litellm-rust/crates/llms/src/openai_like/chat/transformation.rs b/litellm-rust/crates/llms/src/openai_like/chat/transformation.rs index 2e81f396ba5..f1ec8dc8b1d 100644 --- a/litellm-rust/crates/llms/src/openai_like/chat/transformation.rs +++ b/litellm-rust/crates/llms/src/openai_like/chat/transformation.rs @@ -6,9 +6,9 @@ use litellm_auth::{CredentialPlacement, SecretValue}; use litellm_core_utils::core_helpers::unix_now; -use litellm_types::{ - llms::openai::ChatMessage, - utils::{ChatCompletionsChoice, ChatCompletionsChoiceMessage, ChatCompletionsResponse}, +use litellm_llms_types::formats::chat_completions::{ + ChatCompletionsChoice, ChatCompletionsChoiceMessage, ChatCompletionsResponse, + ChatCompletionsUsage, ChatMessage, PromptTokensDetails, }; use serde_json::{Map, Value, json}; @@ -156,11 +156,11 @@ impl BaseConfig for OpenAILikeChatConfig { .unwrap_or(model) .to_string(), choices, - usage: litellm_types::utils::ChatCompletionsUsage { + usage: ChatCompletionsUsage { prompt_tokens: field("prompt_tokens"), completion_tokens: field("completion_tokens"), total_tokens: field("total_tokens"), - prompt_tokens_details: litellm_types::utils::PromptTokensDetails { + prompt_tokens_details: PromptTokensDetails { cached_tokens: details .and_then(|d| d.get("cached_tokens")) .and_then(Value::as_u64) diff --git a/litellm-rust/crates/llms/src/reducto/ocr/transformation.rs b/litellm-rust/crates/llms/src/reducto/ocr/transformation.rs index 147056dab8d..1979b442936 100644 --- a/litellm-rust/crates/llms/src/reducto/ocr/transformation.rs +++ b/litellm-rust/crates/llms/src/reducto/ocr/transformation.rs @@ -14,11 +14,13 @@ use crate::base_llm::ocr::{ error::Error, handler::{CallHooks, OcrClient, build_http_request, guardrail_document}, transformation::{ - BaseOcrConfig, LiteLLMOcrResponse, OCR_INLINE_MAX_BYTES, OcrConnection, OcrDocument, - OcrPage, OcrRequestContext, OcrResponseFormat, OcrUsageInfo, PreparedOcrRequest, + BaseOcrConfig, OCR_INLINE_MAX_BYTES, OcrConnection, OcrRequestContext, PreparedOcrRequest, decode_and_normalize_response, }, }; +use litellm_llms_types::formats::ocr::{ + LiteLLMOcrResponse, OcrDocument, OcrPage, OcrResponseFormat, OcrUsageInfo, +}; const REDUCTO_API_BASE: &str = "https://platform.reducto.ai"; const REDUCTO_API_KEY_ENV: &str = "REDUCTO_API_KEY"; @@ -72,9 +74,9 @@ struct ReductoResult { #[serde_with::serde_as] #[derive(Clone, Debug, Default, Deserialize)] struct ReductoUsage { - #[serde_as(deserialize_as = "Option")] + #[serde_as(deserialize_as = "Option")] pub num_pages: Option, - #[serde_as(deserialize_as = "Option")] + #[serde_as(deserialize_as = "Option")] pub credits: Option, } diff --git a/litellm-rust/crates/llms/src/vertex_ai/ocr/deepseek_transformation.rs b/litellm-rust/crates/llms/src/vertex_ai/ocr/deepseek_transformation.rs index 86231d50f9c..2c341eb684e 100644 --- a/litellm-rust/crates/llms/src/vertex_ai/ocr/deepseek_transformation.rs +++ b/litellm-rust/crates/llms/src/vertex_ai/ocr/deepseek_transformation.rs @@ -8,11 +8,14 @@ use crate::base_llm::ocr::{ error::Error, handler::OcrClient, transformation::{ - BaseOcrConfig, LiteLLMOcrResponse, OcrDocument, OcrPage, OcrPageDimensions, OcrPageImage, - OcrRequestContext, OcrResponseFormat, OcrUsageInfo, PreparedOcrRequest, - decode_and_normalize_response, decode_response_value, + BaseOcrConfig, OcrRequestContext, PreparedOcrRequest, decode_and_normalize_response, + decode_response_value, }, }; +use litellm_llms_types::formats::ocr::{ + LiteLLMOcrResponse, OcrDocument, OcrPage, OcrPageDimensions, OcrPageImage, OcrResponseFormat, + OcrUsageInfo, +}; const DEFAULT_API_BASE: &str = "https://aiplatform.googleapis.com"; const MODEL_PREFIX: &str = "deepseek-ai/"; @@ -85,7 +88,7 @@ enum DeepSeekContent { #[derive(Deserialize)] struct DeepSeekPage { #[serde(default)] - #[serde_as(deserialize_as = "litellm_core_utils::serde_compat::LaxI64")] + #[serde_as(deserialize_as = "litellm_llms_types::serde_compat::LaxI64")] index: i64, #[serde(default)] markdown: String, @@ -424,7 +427,8 @@ mod tests { DeepSeekOcrParams, DeepSeekOcrResponse, VertexAIDeepSeekOCRConfig, normalize_response, provider_model, }; - use crate::base_llm::ocr::transformation::{BaseOcrConfig, OcrDocument}; + use crate::base_llm::ocr::transformation::BaseOcrConfig; + use litellm_llms_types::formats::ocr::OcrDocument; fn document() -> OcrDocument { serde_json::from_value(json!({"type":"image_url","image_url":"gs://bucket/a.png"})).unwrap() diff --git a/litellm-rust/crates/llms/src/vertex_ai/ocr/transformation.rs b/litellm-rust/crates/llms/src/vertex_ai/ocr/transformation.rs index 58e2f6cb0ad..4656b2534b6 100644 --- a/litellm-rust/crates/llms/src/vertex_ai/ocr/transformation.rs +++ b/litellm-rust/crates/llms/src/vertex_ai/ocr/transformation.rs @@ -9,12 +9,12 @@ use crate::{ error::Error, handler::OcrClient, transformation::{ - BaseOcrConfig, LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrEnvironment, - OcrRequestContext, OcrResponseFormat, PreparedOcrRequest, + BaseOcrConfig, OcrConnection, OcrEnvironment, OcrRequestContext, PreparedOcrRequest, }, }, mistral::ocr::transformation::{MistralOcrConfig, MistralOcrRequest}, }; +use litellm_llms_types::formats::ocr::{LiteLLMOcrResponse, OcrDocument, OcrResponseFormat}; const DEFAULT_LOCATION: &str = "us-central1"; diff --git a/litellm-rust/crates/llms/tests/anthropic_chat_transformation.rs b/litellm-rust/crates/llms/tests/anthropic_chat_transformation.rs index e6e213a4efe..37e0ed80cc0 100644 --- a/litellm-rust/crates/llms/tests/anthropic_chat_transformation.rs +++ b/litellm-rust/crates/llms/tests/anthropic_chat_transformation.rs @@ -6,7 +6,7 @@ use litellm_llms::{ chat::transformation::{BaseConfig, ProviderChatResponseData, Unsupported}, }, }; -use litellm_types::{llms::openai::ChatMessage, utils::ChatCompletionsResponse}; +use litellm_llms_types::formats::chat_completions::{ChatCompletionsResponse, ChatMessage}; use rstest::rstest; use serde_json::{Map, Value, json}; diff --git a/litellm-rust/crates/llms/tests/bedrock_converse_transformation.rs b/litellm-rust/crates/llms/tests/bedrock_converse_transformation.rs index 704f0602e69..aa6920d58f6 100644 --- a/litellm-rust/crates/llms/tests/bedrock_converse_transformation.rs +++ b/litellm-rust/crates/llms/tests/bedrock_converse_transformation.rs @@ -7,7 +7,7 @@ use litellm_llms::{ }, bedrock::chat::converse_transformation::BEDROCK_CHAT_COMPLETIONS_CONFIG, }; -use litellm_types::{llms::openai::ChatMessage, utils::ChatCompletionsResponse}; +use litellm_llms_types::formats::chat_completions::{ChatCompletionsResponse, ChatMessage}; use rstest::rstest; use serde_json::{Map, Value, json}; diff --git a/litellm-rust/crates/llms/tests/messages_normalization.rs b/litellm-rust/crates/llms/tests/messages_normalization.rs index a3ccc0a6f95..27dba22c662 100644 --- a/litellm-rust/crates/llms/tests/messages_normalization.rs +++ b/litellm-rust/crates/llms/tests/messages_normalization.rs @@ -1,5 +1,5 @@ use litellm_llms::base_llm::messages::normalization::fold_system_role_messages; -use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest; +use litellm_llms_types::formats::messages::MessagesRequest; use rstest::rstest; use serde_json::{Value, json}; @@ -14,7 +14,7 @@ fn folding_preserves_block_fields_order_and_unrelated_request_fields( let cache_control = json!({"type": "ephemeral", "scope": "global", "future": true}); let folded_block = json!({"type": "text", "text": "second", "cache_control": cache_control}); let user = json!({"role": "user", "content": "hello", "future_message": 42}); - let request: AnthropicMessagesRequest = serde_json::from_value(json!({ + let request: MessagesRequest = serde_json::from_value(json!({ "model": "test-model", "max_tokens": 64, "system": system, diff --git a/litellm-rust/crates/llms/tests/openai_like_chat_transformation.rs b/litellm-rust/crates/llms/tests/openai_like_chat_transformation.rs index b91794c75ab..1c18873772b 100644 --- a/litellm-rust/crates/llms/tests/openai_like_chat_transformation.rs +++ b/litellm-rust/crates/llms/tests/openai_like_chat_transformation.rs @@ -6,7 +6,7 @@ use litellm_llms::{ }, openai_like::chat::transformation::OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG, }; -use litellm_types::{llms::openai::ChatMessage, utils::ChatCompletionsResponse}; +use litellm_llms_types::formats::chat_completions::{ChatCompletionsResponse, ChatMessage}; use rstest::rstest; use serde_json::{Map, Value, json}; diff --git a/litellm-rust/crates/model-catalog/Cargo.toml b/litellm-rust/crates/model-catalog/Cargo.toml index 94a69c94fdf..570e68ca4fd 100644 --- a/litellm-rust/crates/model-catalog/Cargo.toml +++ b/litellm-rust/crates/model-catalog/Cargo.toml @@ -6,10 +6,10 @@ license.workspace = true repository.workspace = true [features] -schema = ["dep:schemars", "litellm-types/schema"] +schema = ["dep:schemars", "litellm-llms-types/schema"] [dependencies] -litellm-types.workspace = true +litellm-llms-types.workspace = true indexmap = { version = "2.14.0", features = ["serde"] } schemars = { workspace = true, optional = true } diff --git a/litellm-rust/crates/model-catalog/src/model_info.rs b/litellm-rust/crates/model-catalog/src/model_info.rs index c46a7e57104..aa543439885 100644 --- a/litellm-rust/crates/model-catalog/src/model_info.rs +++ b/litellm-rust/crates/model-catalog/src/model_info.rs @@ -1,6 +1,6 @@ use crate::capabilities::{AudioFormat, InputModality, Mode, OutputModality, VertexAiAudioApi}; use crate::pricing::{OffPeakPricing, SearchContextCostPerQuery, TieredRate, WebSearchBillingUnit}; -use litellm_types::llms::openai::ReasoningEffort; +use litellm_llms_types::formats::chat_completions::ReasoningEffort; use serde::{Deserialize, Serialize}; use serde_json::Value; use std::collections::BTreeMap; diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index 329fb63c8e7..f8ed125f229 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -46,7 +46,7 @@ litellm-http.workspace = true litellm-llms.workspace = true litellm-secrets = { workspace = true, features = ["aws", "azure", "google", "hashicorp", "cyberark"] } litellm-secrets-types.workspace = true -litellm-types.workspace = true +litellm-llms-types.workspace = true litellm-host-python.workspace = true litellm-token-counter = { path = "../token-counter", default-features = false } pyo3.workspace = true diff --git a/litellm-rust/crates/python-bridge/src/marshal.rs b/litellm-rust/crates/python-bridge/src/marshal.rs index fe5d551a931..7858b695edf 100644 --- a/litellm-rust/crates/python-bridge/src/marshal.rs +++ b/litellm-rust/crates/python-bridge/src/marshal.rs @@ -175,9 +175,9 @@ mod tests { #[serde_with::serde_as] #[derive(Debug, serde::Deserialize, serde::Serialize, PartialEq)] struct Numbers { - #[serde_as(deserialize_as = "Option>")] + #[serde_as(deserialize_as = "Option>")] integers: Option>, - #[serde_as(deserialize_as = "Option")] + #[serde_as(deserialize_as = "Option")] float: Option, } diff --git a/litellm-rust/crates/python-bridge/src/routes/AGENTS.md b/litellm-rust/crates/python-bridge/src/routes/AGENTS.md index 76578447bba..c2afed49b45 100644 --- a/litellm-rust/crates/python-bridge/src/routes/AGENTS.md +++ b/litellm-rust/crates/python-bridge/src/routes/AGENTS.md @@ -8,6 +8,6 @@ Before execution starts, perform only admission checks needed to select native e The host driver owns sequencing and terminal events; the bridge supplies fallible resource composition without exposing route types to the driver. An unstarted async call performs no resource setup. Setup errors after start follow the terminal failure contract and never authorize fallback or provider replay -Use the shared `run_public_call` boundary with hooks supplied by bridge composition. `callbacks-legacy-python` owns legacy argument sharing and `Logging` dispatch behind `PublicCall` and `LegacyLogging`. Route bindings identify their neutral `Operation` and may retain the request needed for projection, but must not duplicate the legacy callback contract +Use the shared `run_public_call` boundary with hooks supplied by bridge composition. `callbacks-legacy-python` owns legacy argument sharing and `Logging` dispatch behind `PublicCall` and `LegacyLogging`. Route bindings supply `callbacks-legacy-python::LoggingOperation` when composing legacy logging and may retain the request needed for projection, but must not duplicate the legacy callback contract Regression tests must observe that an unstarted call does no setup, hook and preflight rewrites affect resource configuration, setup failures reach the selected failure handler once, and provider work is not replayed. Retain existing read-point and object-identity guarantees while changing setup timing diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs index 37d3420a285..fdd7be58a35 100644 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs @@ -4,7 +4,7 @@ use pyo3::types::{PyDict, PyTuple}; use crate::logger::{run_async, run_sync}; use litellm_core::chat_completions::{ChatCompletionsRoute, Error, types::ChatCompletionsRequest}; -use litellm_types::utils::ChatCompletionsResponse; +use litellm_llms_types::formats::chat_completions::ChatCompletionsResponse; use pyo3::prelude::*; use serde_json::{Map, Value}; @@ -137,7 +137,7 @@ fn run_public( asynchronous: bool, ) -> PyResult> { use super::inference::InferenceHost; - use litellm_types::Operation; + use litellm_callbacks_legacy_python::LoggingOperation; let host = InferenceHost::new( request.clone().unbind(), "litellm.rust_bridge.chat_completions.route_host", @@ -150,7 +150,7 @@ fn run_public( crate::cache::admit_native(py, &kwargs, cache_call_type)?; let (arguments, hooks) = crate::routes::call_hooks( py, - Operation::Completion, + LoggingOperation::Completion, &request, &args, &kwargs, diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs index d03be3ebb49..2f67151374d 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs @@ -8,7 +8,7 @@ use litellm_core::messages::{ }; use litellm_host_python::{InvokeError, PythonBinding, from_py, lookup, to_py}; use litellm_http::transport::Error as TransportError; -use litellm_types::utils::ProviderSpecificHeaders; +use litellm_llms_types::headers::ProviderSpecificHeaders; use pyo3::{ exceptions::{PyException, PyValueError}, gc::{PyTraverseError, PyVisit}, @@ -247,9 +247,7 @@ impl PythonBinding for MessagesPythonHost { fn encode_response( &mut self, py: Python<'_>, - response: Box< - litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse, - >, + response: Box, ) -> PyResult> { py.import(ROUTE_HOST_MODULE)? .getattr("response")? diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs index 1bf2b1e0ba1..7ed5375f265 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs @@ -1,7 +1,7 @@ mod host; use host::MessagesPythonHost; -use litellm_types::Operation; +use litellm_callbacks_legacy_python::LoggingOperation; use pyo3::{ prelude::*, types::{PyDict, PyTuple}, @@ -16,7 +16,7 @@ fn run_messages( ) -> PyResult> { let (arguments, hooks) = crate::routes::call_hooks( py, - Operation::Messages, + LoggingOperation::Messages, &request, &args, &kwargs, diff --git a/litellm-rust/crates/python-bridge/src/routes/mod.rs b/litellm-rust/crates/python-bridge/src/routes/mod.rs index 6542983016e..0ea10c52c08 100644 --- a/litellm-rust/crates/python-bridge/src/routes/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/mod.rs @@ -7,10 +7,10 @@ pub(crate) mod ocr; pub(crate) mod responses; pub(crate) mod token_counter; +use litellm_callbacks_legacy_python::LoggingOperation; use litellm_callbacks_legacy_python::{LegacyLogging, PublicCall}; use litellm_host::{call::HostedCompletion, machine::Machine, protocol::Protocol}; use litellm_host_python::{HookChain, PythonBinding, PythonCallHooks, PythonHostCalls}; -use litellm_types::Operation; use pyo3::{ prelude::*, types::{PyDict, PyTuple}, @@ -18,7 +18,7 @@ use pyo3::{ fn call_hooks( py: Python<'_>, - operation: Operation, + operation: LoggingOperation, request: &Bound<'_, PyAny>, args: &Bound<'_, PyTuple>, kwargs: &Bound<'_, PyDict>, diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs index 03a982f8117..28317317544 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs @@ -2,7 +2,8 @@ use litellm_auth::ResolvedCredential; use litellm_core::ocr::route::{Ocr, OcrCall, OcrOp}; use litellm_host_python::{InvokeError, PythonBinding, missing_state, to_py}; use litellm_host_python::{PythonHostCalls, PythonOwned}; -use litellm_llms::base_llm::ocr::{error::Error, transformation::LiteLLMOcrResponse}; +use litellm_llms::base_llm::ocr::error::Error; +use litellm_llms_types::formats::ocr::LiteLLMOcrResponse; use pyo3::{ exceptions::{PyBaseException, PyException}, gc::{PyTraverseError, PyVisit}, diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs index 2c732c3b1a3..041d5c7d0b3 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -4,11 +4,11 @@ mod host; mod project; use host::OcrPythonHost; +use litellm_callbacks_legacy_python::LoggingOperation; use litellm_core::ocr::provider_config; use litellm_core_utils::settings::ProcessEnvironment; use litellm_host_python::to_py; use litellm_llms::base_llm::ocr::settings::OcrSettings; -use litellm_types::Operation; use pyo3::{ prelude::*, types::{PyDict, PyTuple}, @@ -36,8 +36,14 @@ fn run_ocr( kwargs: Bound<'_, PyDict>, asynchronous: bool, ) -> PyResult> { - let (arguments, hooks) = - crate::routes::call_hooks(py, Operation::Ocr, &request, &args, &kwargs, asynchronous)?; + let (arguments, hooks) = crate::routes::call_hooks( + py, + LoggingOperation::Ocr, + &request, + &args, + &kwargs, + asynchronous, + )?; crate::routes::run_public_call( py, arguments, diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs index be43a1b7711..959993f5493 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs @@ -157,7 +157,7 @@ pub(super) fn project_request( #[cfg(test)] mod tests { - use litellm_llms::base_llm::ocr::transformation::OcrDocument; + use litellm_llms_types::formats::ocr::OcrDocument; use pyo3::exceptions::PyValueError; use super::*; diff --git a/litellm-rust/crates/python-bridge/src/routes/responses.rs b/litellm-rust/crates/python-bridge/src/routes/responses.rs index 9b21ef13324..1c2b685854a 100644 --- a/litellm-rust/crates/python-bridge/src/routes/responses.rs +++ b/litellm-rust/crates/python-bridge/src/routes/responses.rs @@ -20,7 +20,7 @@ fn run_public( asynchronous: bool, ) -> PyResult> { use super::inference::InferenceHost; - use litellm_types::Operation; + use litellm_callbacks_legacy_python::LoggingOperation; let host = InferenceHost::new( request.clone().unbind(), "litellm.rust_bridge.responses.route_host", @@ -71,7 +71,7 @@ fn run_public( crate::cache::admit_native(py, &kwargs, cache_call_type)?; let (arguments, hooks) = crate::routes::call_hooks( py, - Operation::Responses, + LoggingOperation::Responses, &request, &args, &kwargs, diff --git a/litellm-rust/crates/token-counter/src/counter.rs b/litellm-rust/crates/token-counter/src/counter.rs index ce08e225be4..6370d9d74f8 100644 --- a/litellm-rust/crates/token-counter/src/counter.rs +++ b/litellm-rust/crates/token-counter/src/counter.rs @@ -4,8 +4,8 @@ use crate::Error; use crate::python_json; use crate::tools::format_function_definitions; use crate::types::{ - ContentBlock, ContentItem, CountableRequest, Message, MessageContent, TextValue, ToolChoice, - ToolDefinition, + ContentItem, CountableContentBlock, CountableRequest, Message, MessageContent, TextValue, + ToolChoice, ToolDefinition, }; const TOKENS_PER_MESSAGE: usize = 3; @@ -130,20 +130,20 @@ impl TokenCounter { fn count_content_item(&self, item: &ContentItem) -> Result { match item { ContentItem::Text(text) => self.count_text(text), - ContentItem::Block(ContentBlock::Text { text }) => self.count_text(text), - ContentItem::Block(ContentBlock::Thinking { thinking }) => { + ContentItem::Block(CountableContentBlock::Text { text }) => self.count_text(text), + ContentItem::Block(CountableContentBlock::Thinking { thinking }) => { if thinking.is_empty() { return Ok(0); } self.count_text(thinking) } - ContentItem::Block(ContentBlock::ToolReference { tool_name }) => { + ContentItem::Block(CountableContentBlock::ToolReference { tool_name }) => { match tool_name.as_deref().filter(|name| !name.is_empty()) { Some(name) => self.count_text(name), None => Ok(0), } } - ContentItem::Block(ContentBlock::Unsupported) => Err(Error::ContentBlock), + ContentItem::Block(CountableContentBlock::Unsupported) => Err(Error::ContentBlock), } } diff --git a/litellm-rust/crates/token-counter/src/types.rs b/litellm-rust/crates/token-counter/src/types.rs index d25554beaac..c1236f94d9a 100644 --- a/litellm-rust/crates/token-counter/src/types.rs +++ b/litellm-rust/crates/token-counter/src/types.rs @@ -158,12 +158,12 @@ pub(crate) enum MessageContent { #[serde(untagged)] pub(crate) enum ContentItem { Text(String), - Block(ContentBlock), + Block(CountableContentBlock), } #[derive(Clone, Debug, Deserialize, PartialEq)] #[serde(tag = "type")] -pub(crate) enum ContentBlock { +pub(crate) enum CountableContentBlock { #[serde(rename = "text")] Text { text: String }, #[serde(rename = "thinking")] diff --git a/litellm-rust/crates/types/src/lib.rs b/litellm-rust/crates/types/src/lib.rs deleted file mode 100644 index 4460e60d51c..00000000000 --- a/litellm-rust/crates/types/src/lib.rs +++ /dev/null @@ -1,14 +0,0 @@ -pub mod audio_transcription; -pub mod llms; -pub mod messages; -pub mod recognized; -pub mod responses; -pub mod utils; - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum Operation { - Completion, - Responses, - Messages, - Ocr, -} diff --git a/litellm-rust/crates/types/src/llms/anthropic_messages/mod.rs b/litellm-rust/crates/types/src/llms/anthropic_messages/mod.rs deleted file mode 100644 index 2b6ada1f22e..00000000000 --- a/litellm-rust/crates/types/src/llms/anthropic_messages/mod.rs +++ /dev/null @@ -1,2 +0,0 @@ -pub mod anthropic_request; -pub mod anthropic_response; diff --git a/litellm-rust/crates/types/src/llms/mod.rs b/litellm-rust/crates/types/src/llms/mod.rs deleted file mode 100644 index 19ce0bb77ef..00000000000 --- a/litellm-rust/crates/types/src/llms/mod.rs +++ /dev/null @@ -1,3 +0,0 @@ -pub mod anthropic; -pub mod anthropic_messages; -pub mod openai; diff --git a/litellm-rust/crates/types/src/llms/openai.rs b/litellm-rust/crates/types/src/llms/openai.rs deleted file mode 100644 index ee8c882c40c..00000000000 --- a/litellm-rust/crates/types/src/llms/openai.rs +++ /dev/null @@ -1,132 +0,0 @@ -use serde::{Deserialize, Serialize}; -use serde_json::{Map, Value}; -use strum::IntoStaticStr; - -/// Reasoning effort level accepted or applied by the model. -#[derive(Clone, Copy, Debug, Deserialize, Eq, IntoStaticStr, PartialEq, Serialize)] -#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] -#[serde(rename_all = "snake_case")] -#[strum(serialize_all = "snake_case")] -pub enum ReasoningEffort { - None, - Minimal, - Low, - Medium, - High, - Xhigh, - Max, -} - -impl ReasoningEffort { - pub const ALL: [Self; 7] = [ - Self::None, - Self::Minimal, - Self::Low, - Self::Medium, - Self::High, - Self::Xhigh, - Self::Max, - ]; - - pub fn as_str(self) -> &'static str { - self.into() - } - - pub fn parse(value: &str) -> Option { - Self::ALL - .into_iter() - .find(|effort| effort.as_str() == value) - } -} - -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -#[serde(untagged)] -pub enum ChatMessageContent { - Text(String), - Parts(Vec), -} - -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -pub struct ChatMessage { - pub role: String, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub content: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub name: Option, - #[serde(flatten)] - pub extra: Map, -} - -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -pub struct ChatCompletionToolCallFunctionChunk { - #[serde(default, skip_serializing_if = "Option::is_none")] - pub name: Option, - pub arguments: String, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub provider_specific_fields: Option>, -} - -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -pub struct ChatCompletionToolCallChunk { - #[serde(default, skip_serializing_if = "Option::is_none")] - pub id: Option, - #[serde(rename = "type")] - pub tool_type: String, - pub function: ChatCompletionToolCallFunctionChunk, - pub index: i64, -} - -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -#[serde(tag = "type", rename_all = "snake_case")] -pub enum ChatCompletionThinkingBlock { - Thinking { - #[serde(default, skip_serializing_if = "Option::is_none")] - thinking: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - signature: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - cache_control: Option, - }, - RedactedThinking { - #[serde(default, skip_serializing_if = "Option::is_none")] - data: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - cache_control: Option, - }, -} - -#[cfg(test)] -mod tests { - use rstest::rstest; - - use super::*; - - #[rstest] - fn reasoning_effort_names_match_the_wire_and_parse_back( - #[values( - ReasoningEffort::None, - ReasoningEffort::Minimal, - ReasoningEffort::Low, - ReasoningEffort::Medium, - ReasoningEffort::High, - ReasoningEffort::Xhigh, - ReasoningEffort::Max - )] - effort: ReasoningEffort, - ) { - assert_eq!( - serde_json::to_value(effort).unwrap(), - Value::String(effort.as_str().to_string()) - ); - assert_eq!(ReasoningEffort::parse(effort.as_str()), Some(effort)); - assert!(ReasoningEffort::ALL.contains(&effort)); - } - - #[rstest] - #[case::unknown("ultra")] - #[case::uppercase("HIGH")] - #[case::empty("")] - fn reasoning_effort_parse_rejects(#[case] value: &str) { - assert_eq!(ReasoningEffort::parse(value), None); - } -} diff --git a/litellm-rust/crates/types/src/messages/mod.rs b/litellm-rust/crates/types/src/messages/mod.rs deleted file mode 100644 index 7bf4fc46291..00000000000 --- a/litellm-rust/crates/types/src/messages/mod.rs +++ /dev/null @@ -1 +0,0 @@ -pub mod streaming; diff --git a/litellm-rust/crates/types/src/responses/mod.rs b/litellm-rust/crates/types/src/responses/mod.rs deleted file mode 100644 index 578373421e6..00000000000 --- a/litellm-rust/crates/types/src/responses/mod.rs +++ /dev/null @@ -1,2 +0,0 @@ -pub mod main; -pub mod streaming_websocket; diff --git a/litellm-rust/crates/types/src/utils.rs b/litellm-rust/crates/types/src/utils.rs deleted file mode 100644 index af0ba2c01c9..00000000000 --- a/litellm-rust/crates/types/src/utils.rs +++ /dev/null @@ -1,108 +0,0 @@ -use serde::{Deserialize, Serialize}; -use serde_json::{Map, Value}; - -use crate::llms::openai::{ChatCompletionThinkingBlock, ChatCompletionToolCallChunk}; - -#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] -pub struct ProviderSpecificHeader { - #[serde(default)] - pub custom_llm_provider: String, - #[serde(default)] - pub extra_headers: Map, -} - -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -#[serde(untagged)] -pub enum ProviderSpecificHeaders { - One(ProviderSpecificHeader), - Many(Vec), -} - -/// OpenAI `usage`, including the `prompt_tokens_details` split LiteLLM's Python -/// path reports so cost tracking sees the same numbers on either path. -#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] -pub struct PromptTokensDetails { - pub cached_tokens: u64, - pub cache_creation_tokens: u64, - pub text_tokens: u64, -} - -#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] -pub struct ChatCompletionsUsage { - pub prompt_tokens: u64, - pub completion_tokens: u64, - pub total_tokens: u64, - pub prompt_tokens_details: PromptTokensDetails, -} - -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -pub struct ChatCompletionsChoiceMessage { - pub role: String, - // Whether an empty turn is `None` or `""` is the provider's choice, not a - // shared invariant: Anthropic's transform ends on `merged_text or None` - // while Converse assigns the joined string unconditionally. Each config - // mirrors its own, so keep this optional and serialize it even when None. - pub content: Option, -} - -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -pub struct ChatCompletionsChoice { - pub index: u64, - pub message: ChatCompletionsChoiceMessage, - pub finish_reason: String, -} - -/// The normalized response handed back to the host. -/// -/// There is deliberately no `id`: Python mints the `chatcmpl-…` id on the -/// `ModelResponse` it already created, and echoing the provider's own id here -/// would change it. Pinned by `response_carries_no_id` in the Anthropic chat transformation tests. -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -pub struct ChatCompletionsResponse { - pub created: u64, - pub model: String, - pub choices: Vec, - pub usage: ChatCompletionsUsage, -} - -#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] -pub struct ChatCompletionDelta { - #[serde(default, skip_serializing_if = "Option::is_none")] - pub content: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub role: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub tool_calls: Option>, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub reasoning_content: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub thinking_blocks: Option>, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub provider_specific_fields: Option>, - #[serde(flatten)] - pub extra: Map, -} - -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -pub struct ChatCompletionStreamingChoice { - pub index: u64, - pub delta: ChatCompletionDelta, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub finish_reason: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub logprobs: Option, -} - -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -pub struct ChatCompletionChunk { - pub id: String, - pub created: u64, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub model: Option, - pub object: String, - pub choices: Vec, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub usage: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub provider_specific_fields: Option>, -} From d46304900f283f34226a5021935efe67de6730b5 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 09:22:05 -0700 Subject: [PATCH 14/41] fix(router): stream /v1/messages lifecycle frames live when no fallback can take over (#43600) * fix(router): stream anthropic messages lifecycle frames live when no fallback can take over The /v1/messages streaming wrapper buffered message_start and content_block_start until the first content_block_delta and dropped pings behind buffered frames unconditionally, even for requests no fallback could ever recover. With adaptive thinking on Bedrock or Vertex the client saw no bytes for the whole thinking pass and hit read timeouts. Buffering now applies only while a fallback can still take over (generic or refusal chain resolving), and a ping is always forwarded live since it carries no lifecycle and keeps the connection alive. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(router): mirror every dispatcher fallback path in the anthropic stream gate Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(router): skip already-tried order levels in the anthropic stream gate Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(router): keep a transport-split ping behind buffered lifecycle frames instead of forwarding its head live Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore(router): credit the #39566 branch this fix supersedes Co-authored-by: Radu Swigler Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: mateo Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: yassin Co-authored-by: Radu Swigler --- .../messages/streaming_iterator.py | 29 +- litellm/router.py | 141 +++++--- ..._anthropic_messages_live_lifecycle_wire.py | 82 +++++ .../messages/test_streaming_iterator.py | 20 ++ tests/unit/test_router/test_router.py | 312 ++++++++++++++++-- 5 files changed, 496 insertions(+), 88 deletions(-) create mode 100644 tests/integration/messages_endpoint/providers/anthropic/test_anthropic_messages_live_lifecycle_wire.py diff --git a/litellm/llms/anthropic/pass_through/messages/streaming_iterator.py b/litellm/llms/anthropic/pass_through/messages/streaming_iterator.py index 81d51cc40d5..89e214efa8b 100644 --- a/litellm/llms/anthropic/pass_through/messages/streaming_iterator.py +++ b/litellm/llms/anthropic/pass_through/messages/streaming_iterator.py @@ -40,22 +40,31 @@ def _is_message_stop_chunk(chunk: object) -> bool: def is_anthropic_ping_chunk(chunk: object) -> bool: """ - Whether a chunk is a pure ``ping`` keepalive frame. It carries no content - and can recur indefinitely on a slow-starting or idle connection, so a - mid-stream fallback wrapper drops it outright while still deciding - whether to commit to the primary stream, rather than buffering it. + Whether a chunk is made only of whole ``ping`` keepalive frames. A ping + carries no content or lifecycle, so a mid-stream fallback wrapper can + forward it live while still deciding whether to commit to the primary + stream, without risking two overlapping message lifecycles on the wire. A physical transport chunk that coalesces a ping with any other SSE event (``message_start``, ``content_block_delta``, ``event: error``, ...) - is NOT a pure ping - dropping it whole would discard those events - so - only a chunk whose every ``event:`` line is ``event: ping`` qualifies. + is NOT a pure ping, and neither is a fragment of a ping frame split + across two reads, or a chunk that opens with the tail of an earlier + frame: forwarding either live would interleave it with frames still + held back for a fallback. Only a chunk that begins with ``event: ping``, + ends on a frame boundary, and whose every ``event:`` line is + ``event: ping`` qualifies. """ if isinstance(chunk, dict): return chunk.get("type") == "ping" - if isinstance(chunk, (bytes, bytearray)): - event_lines: Final = tuple(line for line in chunk.splitlines() if line.startswith(b"event:")) - return bool(event_lines) and all(line == b"event: ping" for line in event_lines) - return False + if not isinstance(chunk, (bytes, bytearray)): + return False + event_lines: Final = tuple(line for line in chunk.splitlines() if line.startswith(b"event:")) + return ( + bool(event_lines) + and all(line == b"event: ping" for line in event_lines) + and chunk.startswith(b"event: ping") + and chunk.endswith((b"\n\n", b"\r\n\r\n")) + ) def is_anthropic_content_delta_chunk(chunk: object) -> bool: diff --git a/litellm/router.py b/litellm/router.py index cfed5dc81c0..86ba5112435 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -518,31 +518,20 @@ def _with_router_resolved_session_model(session: object, model_name: str) -> Map # Router._aanthropic_messages_streaming_iterator buffers lifecycle chunks -# until real content commits the primary stream; a hostile or slow-starting -# upstream that never emits content or an error could otherwise grow that -# buffer without bound, so hitting this cap forces an early commit instead. +# until real content commits the primary stream, and only while a fallback +# can still take over; a hostile or slow-starting upstream that never emits +# content or an error could otherwise grow that buffer without bound, so +# hitting this cap forces an early commit instead. MAX_BUFFERED_PRE_CONTENT_ANTHROPIC_CHUNKS: Final = 200 -def _anthropic_stream_should_drop_pre_content_ping(chunk: object, has_generated_content: bool) -> bool: - """A `ping` keepalive seen before any real content is dropped outright - it recurs indefinitely on a - slow-starting connection and carries nothing worth buffering toward a possible fallback.""" +def _anthropic_stream_forwards_ping_live(chunk: object, has_generated_content: bool) -> bool: + """A `ping` keepalive reaches the client live whenever the stream has not committed: it carries no + lifecycle, so it cannot create overlapping lifecycles on the wire, and it keeps the connection alive + while lifecycle frames sit buffered for a possible fallback during a long thinking pass.""" from litellm.llms.anthropic.pass_through.messages.streaming_iterator import is_anthropic_ping_chunk - if has_generated_content: - return False - return is_anthropic_ping_chunk(chunk) - - -def _anthropic_stream_forwards_ping_live(chunk: object, has_generated_content: bool, buffered_chunk_count: int) -> bool: - """A `ping` that no lifecycle frame precedes reaches the client live: a fallback's own message_start can still - follow it without overlapping lifecycles, and AgenticAnthropicStreamingIterator's hold-back keepalive is exactly - such a ping.""" - from litellm.llms.anthropic.pass_through.messages.streaming_iterator import is_anthropic_ping_chunk - - if has_generated_content or buffered_chunk_count: - return False - return is_anthropic_ping_chunk(chunk) + return not has_generated_content and is_anthropic_ping_chunk(chunk) def _is_retriable_anthropic_status(status_code: int) -> bool: @@ -5457,14 +5446,19 @@ class Router: Lifecycle/bookkeeping frames (message_start, content_block_start, ping, ...) do not by themselves disqualify a fallback attempt - - Anthropic routinely sends message_start before an overload error - - but they are BUFFERED rather than forwarded immediately, since - forwarding one and then appending a fallback attempt's own - message_start would produce two overlapping message lifecycles on - one SSE stream. Buffered frames are flushed, in order, the moment - real content arrives (the primary attempt has committed by then - anyway) or once the stream ends without ever producing content or - an error. + Anthropic routinely sends message_start before an overload error. + When a fallback can still take over they are BUFFERED rather than + forwarded immediately, since forwarding one and then appending a + fallback attempt's own message_start would produce two overlapping + message lifecycles on one SSE stream; a `ping` carries no lifecycle, + so it is forwarded live even while lifecycle frames sit buffered, + keeping the connection alive during a long thinking pass. Buffered + frames are flushed, in order, the moment real content arrives (the + primary attempt has committed by then anyway) or once the stream + ends without ever producing content or an error. When no fallback + can take over the request is already committed, so every frame, + including pings and provider error frames, is forwarded live and + verbatim instead. """ from litellm.llms.anthropic.pass_through.messages.streaming_iterator import ( aclose_if_supported, @@ -5481,34 +5475,33 @@ class Router: from litellm.exceptions import MidStreamFallbackError # Lifecycle/bookkeeping frames (message_start, content_block_start, - # ping, ...) are held back rather than forwarded immediately: - # Anthropic routinely sends message_start before an overload - # error, and once a byte reaches the client a fallback attempt - # can only append its OWN message_start, producing two - # overlapping message lifecycles on one SSE stream. Buffered - # frames are flushed the moment real content (content_block_delta) + # ...) are held back rather than forwarded immediately, but only + # while a fallback can still take over: Anthropic routinely sends + # message_start before an overload error, and once a byte reaches + # the client a fallback attempt can only append its OWN + # message_start, producing two overlapping message lifecycles on + # one SSE stream. A `ping` keepalive carries no lifecycle, so it + # is forwarded live even behind buffered frames, keeping the + # connection alive through a long thinking pass. Buffered frames + # are flushed the moment real content (content_block_delta) # arrives - at that point the primary attempt has committed and a # clean retry is no longer possible anyway - or once the primary - # stream ends without ever producing content. A `ping` keepalive - # that nothing precedes is forwarded live (it is how a hold-back - # turn keeps its connection alive); one behind buffered frames is - # dropped outright rather than buffered, since it can recur - # indefinitely on a slow-starting connection and carries nothing - # worth preserving; hitting MAX_BUFFERED_PRE_CONTENT_ANTHROPIC_CHUNKS - # forces the same early commit as real content arriving, so a - # hostile or pathological upstream can't grow the buffer forever. - has_generated_content = False # rebind-ok: set once real content is seen, or the buffer cap is hit - buffered_lifecycle_chunks: tuple[bytes, ...] = () # rebind-ok: flushed once committed or on decline + # stream ends without ever producing content. Hitting + # MAX_BUFFERED_PRE_CONTENT_ANTHROPIC_CHUNKS forces the same early + # commit as real content arriving, so a hostile or pathological + # upstream can't grow the buffer forever. With no fallback able + # to take over there is nothing to buffer for, so every frame, + # including pings and provider error frames, is forwarded live. model: Final = cast(str, initial_kwargs.get("model")) # cast-ok: kwargs always carries the model group + has_generated_content = not self._anthropic_messages_stream_can_fall_back( # rebind-ok: set once real content is seen, the buffer cap is hit, or no fallback can take over + model, initial_kwargs + ) + buffered_lifecycle_chunks: tuple[bytes, ...] = () # rebind-ok: flushed once committed or on decline try: async for chunk in source_iterator: - if _anthropic_stream_forwards_ping_live( - chunk, has_generated_content, len(buffered_lifecycle_chunks) - ): + if _anthropic_stream_forwards_ping_live(chunk, has_generated_content): yield chunk continue - if _anthropic_stream_should_drop_pre_content_ping(chunk, has_generated_content): - continue if _anthropic_stream_commits_now(chunk, has_generated_content, len(buffered_lifecycle_chunks)): has_generated_content = True # A transport can split one SSE data line across byte chunks, so pre-content @@ -8447,6 +8440,56 @@ class Router: ) return has_unattempted_fallback_target(resolved, kwargs) + def _anthropic_messages_order_levels(self, model_group: str, kwargs: Mapping[str, Any]) -> tuple[int, ...]: + """ + The distinct deployment order levels the fallback dispatcher would see for this request, + computed the same way: the tier a pre-routing hook selected wins over the requested group. + """ + request_team_id: Final[str | None] = (kwargs.get("metadata", {}) or {}).get("user_api_key_team_id") + order_model_group: Final = get_pre_routing_selection(kwargs) or model_group + all_deployments: Final = self.get_model_list(model_name=order_model_group, team_id=request_team_id) or () + return tuple( + sorted( + { + litellm.utils._get_deployment_order(d) + for d in all_deployments + if litellm.utils._get_deployment_order(d) is not None + } + ) + ) + + def _anthropic_messages_stream_can_fall_back(self, model_group: str, kwargs: Mapping[str, Any]) -> bool: + """ + Whether async_function_with_fallbacks_common_utils could still route a + MidStreamFallbackError somewhere for this request (order levels, weighted + failover, content-policy or generic fallbacks), which is the only case where + holding lifecycle frames back from the client buys a clean retry. Errs toward + True whenever a dispatcher path might reach a fallback. + """ + if fallbacks_disabled_for_request(kwargs): + return False + if self.enable_weighted_failover: + return True + order_levels: Final = self._anthropic_messages_order_levels(model_group, kwargs) + if len(order_levels) > 1: + current_target: Final = kwargs.get("_target_order") + skip_up_to: Final = current_target if current_target is not None else order_levels[0] + if any(o > skip_up_to for o in order_levels): + return True + content_policy_fallbacks: Final = kwargs.get("content_policy_fallbacks", self.content_policy_fallbacks) + if content_policy_fallbacks is not None and self._has_content_policy_fallback(model_group, kwargs): + return True + fallbacks: Final = kwargs.get("fallbacks", self.fallbacks) + if not fallbacks: + return False + if _check_non_standard_fallback_format(fallbacks=fallbacks): + return True + resolved, _ = get_fallback_model_group_for_lookup_groups( + fallbacks=fallbacks, + lookup_groups=fallback_lookup_groups(kwargs, model_group), + ) + return has_unattempted_fallback_target(resolved, kwargs) + def _should_raise_content_policy_error(self, model: str, response: ModelResponse, kwargs: dict) -> bool: """ Determines if a content policy error should be raised. diff --git a/tests/integration/messages_endpoint/providers/anthropic/test_anthropic_messages_live_lifecycle_wire.py b/tests/integration/messages_endpoint/providers/anthropic/test_anthropic_messages_live_lifecycle_wire.py new file mode 100644 index 00000000000..cb7043c0362 --- /dev/null +++ b/tests/integration/messages_endpoint/providers/anthropic/test_anthropic_messages_live_lifecycle_wire.py @@ -0,0 +1,82 @@ +import json +import threading +import uuid +from typing import Final + +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server + +_MODEL: Final = "claude-sonnet-4-5-20250929" +_API_KEY: Final = "synthetic-anthropic-key" + + +def _sse(event: str, payload: dict[str, object]) -> bytes: + return f"event: {event}\ndata: {json.dumps(payload)}\n\n".encode() + + +def test_messages_stream_message_start_reaches_client_before_content_without_fallback( + gateway: Gateway, +) -> None: + """With no fallback able to take over, the proxy must not hold lifecycle + frames back for a retry that cannot happen: message_start reaches the + client while the upstream is still thinking.""" + gate: Final = threading.Event() + head: Final = _sse("message_start", {"type": "message_start", "message": {"id": "msg_live_1"}}) + _sse( + "content_block_start", + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + ) + tail: Final = ( + _sse( + "content_block_delta", + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "Hello"}}, + ) + + _sse("content_block_stop", {"type": "content_block_stop", "index": 0}) + + _sse( + "message_delta", + {"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 3}}, + ) + + _sse("message_stop", {"type": "message_stop"}) + ) + prompt: Final = "live-lifecycle-" + uuid.uuid4().hex + + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/v1/messages" + assert request.headers["x-api-key"] == _API_KEY + body: Final = json.loads(request.body) + assert body["model"] == _MODEL + assert body["stream"] is True + assert body["messages"] == [{"role": "user", "content": prompt}] + return Reply(content_type="text/event-stream", chunks=(head, tail), gate_after_first=gate) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=_API_KEY) + with gateway.client.stream( + "POST", + "/v1/messages", + json={ + "model": model, + "max_tokens": 16, + "stream": True, + "messages": [{"role": "user", "content": prompt}], + }, + headers={"Authorization": f"Bearer {gateway.key}"}, + ) as response: + assert response.status_code == 200, response.read().decode() + lines = response.iter_lines() + first_event: Final = next( + json.loads(line.removeprefix("data: ")) for line in lines if line.startswith("data: ") + ) + assert first_event["type"] == "message_start" + gate.set() + events: Final = (first_event,) + tuple( + json.loads(line.removeprefix("data: ")) for line in lines if line.startswith("data: ") + ) + assert tuple(event["type"] for event in events) == ( + "message_start", + "content_block_start", + "content_block_delta", + "content_block_stop", + "message_delta", + "message_stop", + ), f"observed events: {events!r}" + assert [request.target for request in wire.drain()] == ["/v1/messages"] diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_streaming_iterator.py b/tests/unit/llms/anthropic/pass_through/messages/test_streaming_iterator.py index e4efc62f364..39c5b8048c8 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_streaming_iterator.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_streaming_iterator.py @@ -19,6 +19,7 @@ from litellm.llms.anthropic.pass_through.messages.streaming_iterator import ( _is_provider_error_chunk, anthropic_messages_response_as_sse_events, is_anthropic_content_delta_chunk, + is_anthropic_ping_chunk, parse_anthropic_error_event, ) @@ -171,6 +172,25 @@ def test_is_message_stop_chunk(): assert _is_message_stop_chunk("message_stop") is False +@pytest.mark.parametrize( + ("chunk", "expected"), + [ + (b'event: ping\ndata: {"type": "ping"}\n\n', True), + (b'event: ping\r\ndata: {"type": "ping"}\r\n\r\n', True), + (b'event: ping\ndata: {"type": "ping"}\n\nevent: ping\ndata: {"type": "ping"}\n\n', True), + ({"type": "ping"}, True), + (b'event: ping\ndata: {"ty', False), + (b'pe": "ping"}\n\n', False), + (b'pe": "message_start"}}\n\nevent: ping\ndata: {"type": "ping"}\n\n', False), + (b'event: ping\ndata: {"type": "ping"}\n\nevent: content_block_delta\ndata: {}\n\n', False), + ({"type": "message_start"}, False), + ("event: ping", False), + ], +) +def test_is_anthropic_ping_chunk_only_matches_whole_ping_frames(chunk: object, expected: bool): + assert is_anthropic_ping_chunk(chunk) is expected, chunk + + def test_is_message_stop_chunk_ignores_substring_in_payload(): """ Regression: a `content_block_delta` frame whose payload happens to contain diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index 3dc96e4844b..d4f9924dd13 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -45,7 +45,6 @@ from litellm.router import ( _anthropic_stream_forwards_ping_live, _anthropic_stream_raised_error_status, _anthropic_stream_should_decline_fallback, - _anthropic_stream_should_drop_pre_content_ping, _is_retriable_anthropic_status, _responses_stream_holds_event, ) @@ -4170,7 +4169,7 @@ def _make_router_with_fallback(primary="gpt-4", secondary="gpt-3.5-turbo"): class _InjectedFallbackRouter(Router): def __init__(self, fallback_response: object) -> None: - super().__init__(model_list=[]) + super().__init__(model_list=[], fallbacks=[{"primary": ["fallback"]}]) self._fallback_response: Final = fallback_response async def async_function_with_fallbacks_common_utils( @@ -13696,7 +13695,8 @@ def _anthropic_messages_make_wrapper() -> FallbackAwareAnthropicMessagesStream: return FallbackAwareAnthropicMessagesStream(_anthropic_messages_empty_generator(), object()) -def _anthropic_messages_make_router() -> Router: +def _anthropic_messages_make_router(**router_kwargs) -> Router: + router_kwargs.setdefault("fallbacks", [{"primary": ["fallback"]}]) return Router( model_list=[ { @@ -13712,7 +13712,8 @@ def _anthropic_messages_make_router() -> Router: "model": "bedrock/anthropic.claude-sonnet-4-5", }, }, - ] + ], + **router_kwargs, ) @@ -13900,24 +13901,286 @@ async def test_anthropic_messages_content_coalesced_with_error_in_one_physical_c @pytest.mark.asyncio -async def test_anthropic_messages_ping_behind_buffered_lifecycle_frame_is_dropped(): - """Bugbot regression: a `ping` keepalive behind buffered lifecycle frames - carries no content and is dropped outright rather than buffered - - otherwise a slow-starting connection sending many pings could grow the - pre-content buffer without bound.""" - router = _anthropic_messages_make_router() +async def test_anthropic_messages_ping_behind_buffered_lifecycle_frame_is_forwarded_live(): + """A `ping` behind buffered lifecycle frames still reaches the client + live: it carries no lifecycle, so it cannot create overlapping + lifecycles, and it keeps the connection alive while a fallback-able + stream holds message_start back through a long thinking pass.""" + router = _anthropic_messages_make_router(fallbacks=[{"primary": ["fallback"]}]) + content_released = asyncio.Event() + + async def source(): + yield _anthropic_messages_message_start_chunk() + yield _anthropic_messages_ping_chunk() + await content_released.wait() + yield _anthropic_messages_content_chunk("hi") + + wrapped = await router._aanthropic_messages_streaming_iterator(response=source(), initial_kwargs={"model": "primary"}) + + assert await asyncio.wait_for(wrapped.__anext__(), timeout=1) == _anthropic_messages_ping_chunk() + content_released.set() + assert [chunk async for chunk in wrapped] == [ + _anthropic_messages_message_start_chunk(), + _anthropic_messages_content_chunk("hi"), + ] + + +@pytest.mark.asyncio +async def test_anthropic_messages_split_ping_stays_in_order_behind_buffered_lifecycle_frame(): + """A ping the transport splits across two reads is not a whole frame, so + neither fragment may jump ahead of the buffered message_start: yielding + the head live and flushing the tail behind message_start would splice a + lifecycle frame into the middle of the ping on the wire.""" + router = _anthropic_messages_make_router(fallbacks=[{"primary": ["fallback"]}]) + ping_head, ping_tail = b'event: ping\ndata: {"ty', b'pe": "ping"}\n\n' source = _AnthropicMessagesFakeByteStream( - [ - _anthropic_messages_message_start_chunk(), - _anthropic_messages_ping_chunk(), - _anthropic_messages_content_chunk("hi"), - ] + [_anthropic_messages_message_start_chunk(), ping_head, ping_tail, _anthropic_messages_content_chunk("hi")] ) wrapped = await router._aanthropic_messages_streaming_iterator(response=source, initial_kwargs={"model": "primary"}) - collected = [chunk async for chunk in wrapped] - assert collected == [_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("hi")] + assert [chunk async for chunk in wrapped] == [ + _anthropic_messages_message_start_chunk(), + ping_head, + ping_tail, + _anthropic_messages_content_chunk("hi"), + ] + + +@pytest.mark.asyncio +async def test_anthropic_messages_no_fallback_message_start_reaches_client_before_content(): + """With no fallback able to take over, the stream is committed from the + first frame: message_start reaches the client live instead of waiting + behind the buffer for content that may be a whole thinking pass away.""" + router = _anthropic_messages_make_router(fallbacks=None) + content_released = asyncio.Event() + + async def source(): + yield _anthropic_messages_message_start_chunk() + await content_released.wait() + yield _anthropic_messages_content_chunk("hi") + + wrapped = await router._aanthropic_messages_streaming_iterator(response=source(), initial_kwargs={"model": "primary"}) + + assert await asyncio.wait_for(wrapped.__anext__(), timeout=1) == _anthropic_messages_message_start_chunk() + content_released.set() + assert [chunk async for chunk in wrapped] == [_anthropic_messages_content_chunk("hi")] + + +@pytest.mark.asyncio +async def test_anthropic_messages_disabled_fallbacks_message_start_reaches_client_before_content(): + """A router with fallbacks configured cannot take over a request that + opted out with disable_fallbacks=True, so its lifecycle frames reach + the client live exactly like a no-fallback router's.""" + router = _anthropic_messages_make_router(fallbacks=[{"primary": ["fallback"]}]) + content_released = asyncio.Event() + + async def source(): + yield _anthropic_messages_message_start_chunk() + await content_released.wait() + yield _anthropic_messages_content_chunk("hi") + + wrapped = await router._aanthropic_messages_streaming_iterator( + response=source(), initial_kwargs={"model": "primary", "disable_fallbacks": True} + ) + + assert await asyncio.wait_for(wrapped.__anext__(), timeout=1) == _anthropic_messages_message_start_chunk() + content_released.set() + assert [chunk async for chunk in wrapped] == [_anthropic_messages_content_chunk("hi")] + + +@pytest.mark.asyncio +async def test_anthropic_messages_no_fallback_error_frame_reaches_client_verbatim(): + """With no fallback able to take over, a retriable provider error frame + is forwarded verbatim instead of triggering a fallback that does not + exist, and the frames already received stay in order ahead of it.""" + router = _anthropic_messages_make_router(fallbacks=None) + source = _AnthropicMessagesFakeByteStream( + [_anthropic_messages_message_start_chunk(), _anthropic_messages_overloaded_error_chunk()] + ) + + with patch.object( + router, + "async_function_with_fallbacks_common_utils", + new=AsyncMock(return_value=_AnthropicMessagesFallbackByteStream([])), + ) as mock_fallback: + wrapped = await router._aanthropic_messages_streaming_iterator( + response=source, initial_kwargs={"model": "primary"} + ) + collected = [chunk async for chunk in wrapped] + + assert collected == [_anthropic_messages_message_start_chunk(), _anthropic_messages_overloaded_error_chunk()] + mock_fallback.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_anthropic_messages_default_wildcard_fallback_still_buffers_lifecycle_frames(): + """A "*" default fallback can take over for any group, so lifecycle + frames are still held back until real content commits the primary.""" + router = _anthropic_messages_make_router(fallbacks=[{"*": ["fallback"]}]) + content_released = asyncio.Event() + + async def source(): + yield _anthropic_messages_message_start_chunk() + await content_released.wait() + yield _anthropic_messages_content_chunk("hi") + + wrapped = await router._aanthropic_messages_streaming_iterator(response=source(), initial_kwargs={"model": "primary"}) + + pending = asyncio.ensure_future(wrapped.__anext__()) + await asyncio.sleep(0.2) + assert not pending.done() + content_released.set() + assert await asyncio.wait_for(pending, timeout=1) == _anthropic_messages_message_start_chunk() + assert [chunk async for chunk in wrapped] == [_anthropic_messages_content_chunk("hi")] + + +def _anthropic_messages_two_order_primary_model_list() -> list: + return [ + { + "model_name": "primary", + "litellm_params": {"model": "anthropic/claude-sonnet-4-5", "api_key": "sk-test", "order": 1}, + }, + { + "model_name": "primary", + "litellm_params": {"model": "bedrock/anthropic.claude-sonnet-4-5", "order": 2}, + }, + { + "model_name": "fallback", + "litellm_params": {"model": "bedrock/anthropic.claude-sonnet-4-5"}, + }, + ] + + +@pytest.mark.parametrize( + "router_kwargs,request_kwargs,expected", + [ + pytest.param({"fallbacks": None}, {"model": "primary"}, False, id="no-fallbacks"), + pytest.param({"fallbacks": [{"primary": ["fallback"]}]}, {"model": "primary"}, True, id="group-fallback"), + pytest.param({"fallbacks": [{"other": ["fallback"]}]}, {"model": "primary"}, False, id="unrelated-group"), + pytest.param( + {"fallbacks": [{"*": ["fallback"]}]}, + {"model": "primary", "fallbacks": None}, + False, + id="wildcard-overridden-by-request-none", + ), + pytest.param({"fallbacks": [{"*": ["fallback"]}]}, {"model": "primary"}, True, id="wildcard"), + pytest.param({"fallbacks": None}, {"model": "primary", "fallbacks": [{"model": "fallback"}]}, True, id="request-dict-fallback"), + pytest.param({"fallbacks": None}, {"model": "primary", "fallbacks": ["fallback"]}, True, id="request-list-fallback"), + pytest.param( + {"fallbacks": [{"primary": ["fallback"]}]}, + {"model": "primary", "disable_fallbacks": True}, + False, + id="disable-fallbacks", + ), + pytest.param( + {"fallbacks": None, "content_policy_fallbacks": [{"primary": ["fallback"]}]}, + {"model": "primary"}, + True, + id="content-policy-fallback", + ), + pytest.param({"fallbacks": None, "enable_weighted_failover": True}, {"model": "primary"}, True, id="weighted-failover"), + ], +) +def test_anthropic_messages_stream_can_fall_back_direct_call(router_kwargs, request_kwargs, expected): + router = _anthropic_messages_make_router(**router_kwargs) + assert router._anthropic_messages_stream_can_fall_back("primary", request_kwargs) is expected + + +@pytest.mark.parametrize( + "orders,expected", + [ + pytest.param([1, 2], True, id="distinct-orders-can-fall-back"), + pytest.param([1, 1], False, id="same-order-cannot-fall-back"), + ], +) +def test_anthropic_messages_stream_can_fall_back_order_levels(orders, expected): + router = Router( + model_list=[ + { + "model_name": "primary", + "litellm_params": {"model": "anthropic/claude-sonnet-4-5", "api_key": "sk-test", "order": order}, + } + for order in orders + ], + fallbacks=None, + ) + assert router._anthropic_messages_stream_can_fall_back("primary", {"model": "primary"}) is expected + + +@pytest.mark.parametrize( + "request_kwargs,expected", + [ + pytest.param({"model": "primary"}, True, id="no-target-order"), + pytest.param({"model": "primary", "_target_order": 1}, True, id="higher-order-remains"), + pytest.param({"model": "primary", "_target_order": 2}, False, id="top-order-no-order-fallback"), + pytest.param( + {"model": "primary", "_target_order": 2, "fallbacks": [{"primary": ["fallback"]}]}, + True, + id="top-order-external-fallback", + ), + ], +) +def test_anthropic_messages_stream_can_fall_back_order_target(request_kwargs, expected): + router = Router(model_list=_anthropic_messages_two_order_primary_model_list(), fallbacks=None) + assert router._anthropic_messages_stream_can_fall_back("primary", request_kwargs) is expected + + +def test_anthropic_messages_order_levels_direct_call(): + router = Router( + model_list=[ + { + "model_name": "primary", + "litellm_params": {"model": "anthropic/claude-sonnet-4-5", "api_key": "sk-test", "order": order}, + } + for order in (2, 1, None) + ], + fallbacks=None, + ) + assert router._anthropic_messages_order_levels("primary", {"model": "primary"}) == (1, 2) + + +@pytest.mark.asyncio +async def test_anthropic_messages_order_fallback_still_buffers_lifecycle_frames(): + """Two order levels in one group are a real fallback target for the + dispatcher, so lifecycle frames stay buffered until content commits.""" + router = Router(model_list=_anthropic_messages_two_order_primary_model_list(), fallbacks=None) + content_released = asyncio.Event() + + async def source(): + yield _anthropic_messages_message_start_chunk() + await content_released.wait() + yield _anthropic_messages_content_chunk("hi") + + wrapped = await router._aanthropic_messages_streaming_iterator(response=source(), initial_kwargs={"model": "primary"}) + + pending = asyncio.ensure_future(wrapped.__anext__()) + await asyncio.sleep(0.2) + assert not pending.done() + content_released.set() + assert await asyncio.wait_for(pending, timeout=1) == _anthropic_messages_message_start_chunk() + assert [chunk async for chunk in wrapped] == [_anthropic_messages_content_chunk("hi")] + + +@pytest.mark.asyncio +async def test_anthropic_messages_request_fallbacks_none_forwards_message_start_live(): + """A per-request fallbacks=None override disables the router's wildcard + fallback, so lifecycle frames reach the client live before content.""" + router = _anthropic_messages_make_router(fallbacks=[{"*": ["fallback"]}]) + content_released = asyncio.Event() + + async def source(): + yield _anthropic_messages_message_start_chunk() + await content_released.wait() + yield _anthropic_messages_content_chunk("hi") + + wrapped = await router._aanthropic_messages_streaming_iterator( + response=source(), initial_kwargs={"model": "primary", "fallbacks": None} + ) + + assert await asyncio.wait_for(wrapped.__anext__(), timeout=1) == _anthropic_messages_message_start_chunk() + content_released.set() + assert [chunk async for chunk in wrapped] == [_anthropic_messages_content_chunk("hi")] @pytest.mark.asyncio @@ -14320,21 +14583,12 @@ def test_merge_fallback_hidden_params_direct_call(): } -def test_anthropic_stream_should_drop_pre_content_ping_direct_call(): - ping = _anthropic_messages_ping_chunk() - content = _anthropic_messages_content_chunk("hi") - assert _anthropic_stream_should_drop_pre_content_ping(ping, has_generated_content=False) is True - assert _anthropic_stream_should_drop_pre_content_ping(ping, has_generated_content=True) is False - assert _anthropic_stream_should_drop_pre_content_ping(content, has_generated_content=False) is False - - def test_anthropic_stream_forwards_ping_live_direct_call(): ping = _anthropic_messages_ping_chunk() content = _anthropic_messages_content_chunk("hi") - assert _anthropic_stream_forwards_ping_live(ping, has_generated_content=False, buffered_chunk_count=0) is True - assert _anthropic_stream_forwards_ping_live(ping, has_generated_content=False, buffered_chunk_count=1) is False - assert _anthropic_stream_forwards_ping_live(ping, has_generated_content=True, buffered_chunk_count=0) is False - assert _anthropic_stream_forwards_ping_live(content, has_generated_content=False, buffered_chunk_count=0) is False + assert _anthropic_stream_forwards_ping_live(ping, has_generated_content=False) is True + assert _anthropic_stream_forwards_ping_live(ping, has_generated_content=True) is False + assert _anthropic_stream_forwards_ping_live(content, has_generated_content=False) is False def test_anthropic_stream_error_is_gateway_verdict_direct_call(): From d2cbc94fc6baa20964782aea2e368c6187aeab80 Mon Sep 17 00:00:00 2001 From: "berriai-litellm-provider-info-sync[bot]" <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 10:09:11 -0700 Subject: [PATCH 15/41] feat(cost-map): add baseten DeepSeek-V4.1-Flash-Fast (#43735) Price-Sync: litellm-providers Co-authored-by: berriai-litellm-provider-info-sync[bot] <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> --- .../model_prices_and_context_window_backup.json | 17 +++++++++++++++++ model_prices_and_context_window.json | 17 +++++++++++++++++ 2 files changed, 34 insertions(+) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index fc48c17b506..7247f4e677d 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -78773,5 +78773,22 @@ "supports_reasoning": true, "supports_tool_choice": true, "supports_vision": true + }, + "baseten/deepseek-ai/DeepSeek-V4.1-Flash-Fast": { + "cache_read_input_token_cost": 1.4e-07, + "input_cost_per_token": 6e-07, + "litellm_provider": "baseten", + "max_input_tokens": 1048576, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "chat", + "output_cost_per_token": 2.4e-06, + "source": "https://inference.baseten.co/v1/models", + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true } } diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index fc48c17b506..7247f4e677d 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -78773,5 +78773,22 @@ "supports_reasoning": true, "supports_tool_choice": true, "supports_vision": true + }, + "baseten/deepseek-ai/DeepSeek-V4.1-Flash-Fast": { + "cache_read_input_token_cost": 1.4e-07, + "input_cost_per_token": 6e-07, + "litellm_provider": "baseten", + "max_input_tokens": 1048576, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "chat", + "output_cost_per_token": 2.4e-06, + "source": "https://inference.baseten.co/v1/models", + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true } } From 0c553f0398bb964aec8f3df4d3c1a9af6b31916c Mon Sep 17 00:00:00 2001 From: Itai Modiano Date: Tue, 29 Sep 2026 20:20:40 +0300 Subject: [PATCH 16/41] feat(guardrails): send a configured gateway_name from noma_v2 to Noma (#43678) * feat(guardrails): send a configured gateway_name from noma_v2 to Noma The noma_v2 guardrail accepts a gateway_name param, falling back to the NOMA_GATEWAY_NAME env var. The value is stripped, and when it is non-empty it goes out as a top-level gateway_name field on /litellm/guardrail. The param works for both guardrail: noma_v2 and guardrail: noma with use_v2, and it is appended after the existing constructor params so positional callers keep their meaning * chore(ui): regenerate OpenAPI snapshot and dashboard types for gateway_name The new noma_v2 gateway_name param shows up in the proxy OpenAPI spec, so the lazy snapshot and the generated dashboard types need regenerating * Update litellm/proxy/guardrails/guardrail_hooks/noma/noma_v2.py Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> --------- Co-authored-by: Claude Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> --- litellm/proxy/_lazy_openapi_snapshot.json | 12 +++ .../guardrail_hooks/noma/__init__.py | 1 + .../guardrail_hooks/noma/noma_v2.py | 6 ++ litellm/types/guardrails.py | 4 + .../proxy/guardrails/guardrail_hooks/noma.py | 4 + .../guardrail_hooks/test_noma_v2.py | 1 + tests/unit/proxy/guardrails/__init__.py | 0 .../guardrails/guardrail_hooks/__init__.py | 0 .../guardrail_hooks/noma/__init__.py | 0 .../guardrail_hooks/noma/test_noma_v2.py | 99 +++++++++++++++++++ ui/litellm-dashboard/src/lib/http/schema.d.ts | 5 + 11 files changed, 132 insertions(+) create mode 100644 tests/unit/proxy/guardrails/__init__.py create mode 100644 tests/unit/proxy/guardrails/guardrail_hooks/__init__.py create mode 100644 tests/unit/proxy/guardrails/guardrail_hooks/noma/__init__.py create mode 100644 tests/unit/proxy/guardrails/guardrail_hooks/noma/test_noma_v2.py diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 75dce43c84a..bb063f8f77a 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -12133,6 +12133,18 @@ "description": "Whether to fail the request if the guardrail encounters an error. Implemented by guardrail='model_armor', 'generic_guardrail_api' and 'crowdstrike_aidr'. True (default) raises the error. False logs a critical error and lets the request proceed, so only a valid guardrail response can block or modify it.", "title": "Fail On Error" }, + "gateway_name": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "description": "noma_v2 only: name of this gateway, used as the gateway_host label on Noma scans", + "title": "Gateway Name" + }, "grounding_check": { "anyOf": [ { diff --git a/litellm/proxy/guardrails/guardrail_hooks/noma/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/noma/__init__.py index 0375e9f2bce..f82aaab4c0d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/noma/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/noma/__init__.py @@ -42,6 +42,7 @@ def initialize_guardrail_v2(litellm_params: "LitellmParams", guardrail: "Guardra api_key=litellm_params.api_key, api_base=litellm_params.api_base, application_id=litellm_params.application_id, + gateway_name=litellm_params.gateway_name, monitor_mode=litellm_params.monitor_mode, block_failures=litellm_params.block_failures, event_hook=litellm_params.mode, diff --git a/litellm/proxy/guardrails/guardrail_hooks/noma/noma_v2.py b/litellm/proxy/guardrails/guardrail_hooks/noma/noma_v2.py index 292f395053b..8b1fcda7f47 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/noma/noma_v2.py +++ b/litellm/proxy/guardrails/guardrail_hooks/noma/noma_v2.py @@ -62,6 +62,7 @@ class NomaV2Guardrail(CustomGuardrail): application_id: str | None = None, monitor_mode: bool | None = None, block_failures: bool | None = None, + gateway_name: str | None = None, **kwargs: Any, ) -> None: self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback) @@ -69,6 +70,9 @@ class NomaV2Guardrail(CustomGuardrail): self.api_key = api_key or os.environ.get("NOMA_API_KEY") self.api_base = (api_base or os.environ.get("NOMA_API_BASE") or _DEFAULT_API_BASE).rstrip("/") self.application_id = application_id or os.environ.get("NOMA_APPLICATION_ID") + self.gateway_name = self._get_non_empty_str(gateway_name) or self._get_non_empty_str( + os.environ.get("NOMA_GATEWAY_NAME") + ) if monitor_mode is None: self.monitor_mode = os.environ.get("NOMA_MONITOR_MODE", "false").lower() == "true" else: @@ -166,6 +170,8 @@ class NomaV2Guardrail(CustomGuardrail): } if application_id: payload["application_id"] = application_id + if self.gateway_name: + payload["gateway_name"] = self.gateway_name return payload @staticmethod diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index 579a3f6322f..46026c12d24 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -783,6 +783,10 @@ class NomaGuardrailConfigModel(BaseModel): default=None, description="Application ID for Noma Security. Defaults to 'litellm' if not provided", ) + gateway_name: str | None = Field( + default=None, + description="noma_v2 only: name of this gateway, used as the gateway_host label on Noma scans", + ) monitor_mode: bool | None = Field( default=None, description="If True, logs violations without blocking. Defaults to False if not provided", diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/noma.py b/litellm/types/proxy/guardrails/guardrail_hooks/noma.py index 880a9beb333..ef22f73810f 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/noma.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/noma.py @@ -39,6 +39,10 @@ class NomaV2GuardrailConfigModel(GuardrailConfigModel): default=None, description="The Noma Application ID. Reads from NOMA_APPLICATION_ID env var if None.", ) + gateway_name: str | None = Field( + default=None, + description="Gateway name, used as the gateway_host label on Noma scans. Falls back to NOMA_GATEWAY_NAME.", + ) monitor_mode: bool | None = Field( default=None, description="When true, run guardrail checks in monitor mode.", diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_noma_v2.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_noma_v2.py index 2533cf0e8c8..180cdbe5bb5 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_noma_v2.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_noma_v2.py @@ -39,6 +39,7 @@ class TestNomaV2Configuration: assert "api_key" in noma_v2_params assert "api_base" in noma_v2_params assert "application_id" in noma_v2_params + assert "gateway_name" in noma_v2_params assert "monitor_mode" in noma_v2_params assert "block_failures" in noma_v2_params diff --git a/tests/unit/proxy/guardrails/__init__.py b/tests/unit/proxy/guardrails/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/__init__.py b/tests/unit/proxy/guardrails/guardrail_hooks/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/noma/__init__.py b/tests/unit/proxy/guardrails/guardrail_hooks/noma/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/noma/test_noma_v2.py b/tests/unit/proxy/guardrails/guardrail_hooks/noma/test_noma_v2.py new file mode 100644 index 00000000000..6e536f95251 --- /dev/null +++ b/tests/unit/proxy/guardrails/guardrail_hooks/noma/test_noma_v2.py @@ -0,0 +1,99 @@ +import json + +import httpx +import pytest +import respx + +import litellm +from litellm.proxy.guardrails.guardrail_hooks.noma import ( + NomaV2Guardrail, + guardrail_initializer_registry, +) +from litellm.types.guardrails import LitellmParams + +_API_BASE = "https://noma.example.test" + + +@pytest.fixture(autouse=True) +def _fresh_httpx_client(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + monkeypatch.setattr(litellm, "in_memory_llm_clients_cache", None) + monkeypatch.delenv("NOMA_GATEWAY_NAME", raising=False) + + +def _guardrail(gateway_name: str | None) -> NomaV2Guardrail: + return NomaV2Guardrail( + api_base=_API_BASE, + gateway_name=gateway_name, + guardrail_name="noma-guard", + event_hook="pre_call", + default_on=True, + ) + + +async def _scan_body(guardrail: NomaV2Guardrail, respx_mock: respx.MockRouter) -> dict[str, object]: + route = respx_mock.post(f"{_API_BASE}/litellm/guardrail").respond(json={"action": "NONE"}) + await guardrail.apply_guardrail(inputs={"texts": ["hello"]}, request_data={"metadata": {}}, input_type="request") + assert route.call_count == 1 + return json.loads(route.calls.last.request.content) + + +@pytest.mark.asyncio +@pytest.mark.parametrize(("guardrail_type", "extra_params"), [("noma_v2", {}), ("noma", {"use_v2": True})]) +async def test_gateway_name_from_guardrail_config_reaches_noma( + guardrail_type: str, extra_params: dict[str, bool], respx_mock: respx.MockRouter +) -> None: + litellm_params = LitellmParams( + guardrail=guardrail_type, + mode="pre_call", + api_base=_API_BASE, + gateway_name="prod-us-east", + **extra_params, + ) + guardrail = guardrail_initializer_registry[guardrail_type](litellm_params, {"guardrail_name": "noma-guard"}) + + assert (await _scan_body(guardrail, respx_mock))["gateway_name"] == "prod-us-east" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("configured", "env_value", "expected"), + [ + (None, "env-gateway", "env-gateway"), + ("config-gateway", "env-gateway", "config-gateway"), + (" config-gateway ", None, "config-gateway"), + ], +) +async def test_gateway_name_resolution( + configured: str | None, + env_value: str | None, + expected: str, + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + if env_value is not None: + monkeypatch.setenv("NOMA_GATEWAY_NAME", env_value) + + assert (await _scan_body(_guardrail(configured), respx_mock))["gateway_name"] == expected + + +@pytest.mark.asyncio +@pytest.mark.parametrize("configured", [None, "", " "]) +async def test_unset_or_blank_gateway_name_is_left_out(configured: str | None, respx_mock: respx.MockRouter) -> None: + assert "gateway_name" not in await _scan_body(_guardrail(configured), respx_mock) + + +@pytest.mark.asyncio +async def test_positional_args_keep_their_meaning_after_gateway_name_was_added(respx_mock: respx.MockRouter) -> None: + guardrail = NomaV2Guardrail("test-api-key", _API_BASE, "test-app", False, True) + + body = await _scan_body(guardrail, respx_mock) + + assert body["monitor_mode"] is False + assert body["application_id"] == "test-app" + assert "gateway_name" not in body + respx_mock.post(f"{_API_BASE}/litellm/guardrail").respond(status_code=503) + with pytest.raises(httpx.HTTPStatusError): + await guardrail.apply_guardrail( + inputs={"texts": ["hello"]}, request_data={"metadata": {}}, input_type="request" + ) diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 93f4718a586..077e6a6ecee 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -34348,6 +34348,11 @@ export interface components { * @default true */ fail_on_error: boolean | null; + /** + * Gateway Name + * @description noma_v2 only: name of this gateway, used as the gateway_host label on Noma scans + */ + gateway_name?: string | null; /** * Grounding Check * @description Enable grounding verification to ensure output is grounded in provided context. From a3552c451bf2b21aab31e7dd11b2f1757b970d46 Mon Sep 17 00:00:00 2001 From: "berriai-litellm-provider-info-sync[bot]" <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 10:34:04 -0700 Subject: [PATCH 17/41] chore(cost-map): add openai gpt-6.1-sol from the pricing page (#43738) Price-Sync: litellm-providers Co-authored-by: berriai-litellm-provider-info-sync[bot] <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> Co-authored-by: kerry Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- ...odel_prices_and_context_window_backup.json | 76 +++++++++++++++++++ model_prices_and_context_window.json | 76 +++++++++++++++++++ 2 files changed, 152 insertions(+) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 7247f4e677d..4257649ab41 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -78790,5 +78790,81 @@ "supports_response_schema": true, "supports_tool_choice": true, "supports_vision": true + }, + "gpt-6.1-sol": { + "cache_creation_input_token_cost": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5e-06, + "cache_creation_input_token_cost_above_272k_tokens_batches": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens_flex": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens_priority": 1e-05, + "cache_creation_input_token_cost_batches": 1.25e-06, + "cache_creation_input_token_cost_flex": 1.25e-06, + "cache_creation_input_token_cost_priority": 5e-06, + "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_above_272k_tokens": 2e-07, + "cache_read_input_token_cost_above_272k_tokens_batches": 1e-07, + "cache_read_input_token_cost_above_272k_tokens_flex": 1e-07, + "cache_read_input_token_cost_above_272k_tokens_priority": 4e-07, + "cache_read_input_token_cost_batches": 5e-08, + "cache_read_input_token_cost_flex": 5e-08, + "cache_read_input_token_cost_priority": 2e-07, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "input_cost_per_token_above_272k_tokens_batches": 2e-06, + "input_cost_per_token_above_272k_tokens_flex": 2e-06, + "input_cost_per_token_above_272k_tokens_priority": 8e-06, + "input_cost_per_token_batches": 1e-06, + "input_cost_per_token_flex": 1e-06, + "input_cost_per_token_priority": 4e-06, + "litellm_provider": "openai", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1e-05, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "output_cost_per_token_above_272k_tokens_batches": 7.5e-06, + "output_cost_per_token_above_272k_tokens_flex": 7.5e-06, + "output_cost_per_token_above_272k_tokens_priority": 3e-05, + "output_cost_per_token_batches": 5e-06, + "output_cost_per_token_flex": 5e-06, + "output_cost_per_token_priority": 2e-05, + "regional_processing_uplift_multiplier_eu": 1.1, + "regional_processing_uplift_multiplier_us": 1.1, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://developers.openai.com/api/docs/pricing", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/batch", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": false, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_cache_breakpoint": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true } } diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 7247f4e677d..4257649ab41 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -78790,5 +78790,81 @@ "supports_response_schema": true, "supports_tool_choice": true, "supports_vision": true + }, + "gpt-6.1-sol": { + "cache_creation_input_token_cost": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5e-06, + "cache_creation_input_token_cost_above_272k_tokens_batches": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens_flex": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens_priority": 1e-05, + "cache_creation_input_token_cost_batches": 1.25e-06, + "cache_creation_input_token_cost_flex": 1.25e-06, + "cache_creation_input_token_cost_priority": 5e-06, + "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_above_272k_tokens": 2e-07, + "cache_read_input_token_cost_above_272k_tokens_batches": 1e-07, + "cache_read_input_token_cost_above_272k_tokens_flex": 1e-07, + "cache_read_input_token_cost_above_272k_tokens_priority": 4e-07, + "cache_read_input_token_cost_batches": 5e-08, + "cache_read_input_token_cost_flex": 5e-08, + "cache_read_input_token_cost_priority": 2e-07, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "input_cost_per_token_above_272k_tokens_batches": 2e-06, + "input_cost_per_token_above_272k_tokens_flex": 2e-06, + "input_cost_per_token_above_272k_tokens_priority": 8e-06, + "input_cost_per_token_batches": 1e-06, + "input_cost_per_token_flex": 1e-06, + "input_cost_per_token_priority": 4e-06, + "litellm_provider": "openai", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1e-05, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "output_cost_per_token_above_272k_tokens_batches": 7.5e-06, + "output_cost_per_token_above_272k_tokens_flex": 7.5e-06, + "output_cost_per_token_above_272k_tokens_priority": 3e-05, + "output_cost_per_token_batches": 5e-06, + "output_cost_per_token_flex": 5e-06, + "output_cost_per_token_priority": 2e-05, + "regional_processing_uplift_multiplier_eu": 1.1, + "regional_processing_uplift_multiplier_us": 1.1, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://developers.openai.com/api/docs/pricing", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/batch", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": false, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_cache_breakpoint": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true } } From b3dcf8208daaa11ab814dfb152e1701cdabd2601 Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Tue, 29 Sep 2026 10:49:30 -0700 Subject: [PATCH 18/41] test(integration): callback credential canary slots C1-C3 and D5 (#43630) * test(integration): credential canary suite harness Adds tests/integration/security with canary generation and search, sweeps over the database, GET routes, client responses, sink doubles and Redis, an owned proxy rig, a sweep sensitivity self-test and the config deployment api_key slot. Registers the security group in run.py, the manifest and the CircleCI integration matrix. * test(integration): widen canary route sweep and harden the rig Enumerate lazily registered feature routers, call parameterized routes with placeholder ids, fail on routes that return no response, skip provider pass-through routes, add an explicit admin-only route allowance, let the sink double use a configurable token, inflate gzip members anywhere in a blob, sweep Redis before the route walk, and trap outbound connections from the owned proxy. * test(integration): descend into any decoded value that can still hold an encoded canary * test(integration): bound canary decoding by depth and decoded bytes * test(integration): scope log-table and spend-log reads to the scenario window * test(integration): sweep spend-log rows in the scenario date window * test(integration): keep spend-log date window summarized * test(integration): resolve deployment ids, scope paginated log lists, key allowances by slot * test(integration): expect 404 from the caller-scoped team membership route * test(integration): use the rig's own master key and expect 404 from submission lookups * test(integration): check the overridden rig key without assuming the default key is unknown * test(integration): callback credential canary slots C1-C3 and D5 Team callback, team callback_settings, config default_team_settings and key metadata.logging Langfuse secrets, a team Datadog dd_api_key, and request-body Langfuse keys (allow_client_side_credentials) must reach only their sink. Each scenario checks its sink received the canary as auth and that the marker is visible at the stored body, the Logs drawer route and the sink. Adds a unit test that the stored request body snapshot carries no callback parameter. * test(integration): give the callback sink waits a wider bound * test(integration): sweep provider requests for callback credentials --- .../integration/security/_callback_traffic.py | 173 ++++++++ tests/integration/security/_canary.py | 6 + .../security/test_callback_credentials.py | 401 ++++++++++++++++++ .../test_body_snapshot_callback_params.py | 36 ++ 4 files changed, 616 insertions(+) create mode 100644 tests/integration/security/_callback_traffic.py create mode 100644 tests/integration/security/test_callback_credentials.py create mode 100644 tests/test_litellm/proxy/test_body_snapshot_callback_params.py diff --git a/tests/integration/security/_callback_traffic.py b/tests/integration/security/_callback_traffic.py new file mode 100644 index 00000000000..662d84d2019 --- /dev/null +++ b/tests/integration/security/_callback_traffic.py @@ -0,0 +1,173 @@ +"""Traffic matrix and sink doubles for the callback credential slots. + +- ``upstream(request)``: provider double for every endpoint in ``ENDPOINTS``: OpenAI chat (plain + and SSE) and OpenAI Responses (``/v1/messages`` reaches it as chat). A body carrying + ``PROVIDER_4XX`` gets HTTP 400 and one carrying ``PROVIDER_5XX`` gets HTTP 500. The sensitivity + marker found in the body is echoed back. +- ``langfuse_sink`` / ``datadog_sink``: Langfuse OTLP ingest and Datadog intake doubles. +- ``send(gateway, key, endpoint, model, text, extra)``: one client call per endpoint. +- ``spend_request_id(marker)``: the spend row written for the request carrying ``marker``. +- ``wait_for_sink(recorder, marker)``: bounded wait until a sink received the marker (gzip aware). +""" + +from __future__ import annotations + +import json +import re +import uuid +from collections.abc import Mapping +from typing import Final + +import httpx +from integration._support.client import Gateway, eventually, string_value +from integration._support.database import read_rows +from integration._support.wire import Reply, Request +from integration.security._canary import Canary, find_canary +from integration.security._sinks import PROVIDER_4XX, Recorder +from pydantic import JsonValue + +PROVIDER_5XX: Final = "canary-provider-5xx" +ENDPOINTS: Final = ("chat", "chat_stream", "messages", "responses") +OUTCOMES: Final = ("success", "provider_4xx", "provider_5xx") +EXPECTED_STATUS: Final = {"success": 200, "provider_4xx": 400, "provider_5xx": 500} +LANGFUSE_PUBLIC_KEY: Final = "pk-lf-canary-public" +_MARKER: Final = re.compile(rb"lkc-M0-[0-9a-f]{32}") + + +def _echo(body: bytes) -> str: + found: Final = _MARKER.search(body) + return "echo " + (found.group().decode() if found else "none") + + +def _failure(body: bytes) -> Reply | None: + if PROVIDER_4XX.encode() in body: + return Reply( + status=400, + body=b'{"error":{"type":"invalid_request_error","code":"canary_rejected","message":"rejected"}}', + ) + if PROVIDER_5XX.encode() in body: + return Reply(status=500, body=b'{"error":{"type":"server_error","message":"upstream exploded"}}') + return None + + +def _chat(body: Mapping[str, JsonValue], text: str) -> Reply: + identity: Final = f"chatcmpl-{uuid.uuid4().hex}" + usage: Final = {"prompt_tokens": 7, "completion_tokens": 3, "total_tokens": 10} + if body.get("stream") is True: + chunks: Final = ( + {"choices": [{"index": 0, "delta": {"role": "assistant", "content": text}, "finish_reason": None}]}, + {"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}]}, + {"choices": [], "usage": usage}, + ) + events: Final = b"".join( + b"data: " + + json.dumps( + {"id": identity, "object": "chat.completion.chunk", "created": 1, "model": "gpt-4o-mini", **chunk} + ).encode() + + b"\n\n" + for chunk in chunks + ) + return Reply(body=events + b"data: [DONE]\n\n", content_type="text/event-stream") + return Reply( + body=json.dumps( + { + "id": identity, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [{"index": 0, "message": {"role": "assistant", "content": text}, "finish_reason": "stop"}], + "usage": usage, + } + ).encode() + ) + + +def _responses(text: str) -> Reply: + return Reply( + body=json.dumps( + { + "id": f"resp_{uuid.uuid4().hex}", + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [ + { + "type": "message", + "id": f"msg_{uuid.uuid4().hex}", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": text, "annotations": []}], + } + ], + "parallel_tool_calls": True, + "tool_choice": "auto", + "tools": [], + "usage": {"input_tokens": 7, "output_tokens": 3, "total_tokens": 10}, + } + ).encode() + ) + + +def upstream(request: Request) -> Reply: + failure: Final = _failure(request.body) + if failure is not None: + return failure + text: Final = _echo(request.body) + if request.target.split("?", 1)[0].endswith("/responses"): + return _responses(text) + return _chat(json.loads(request.body or b"{}"), text) + + +def langfuse_sink(request: Request) -> Reply: + if request.method == "GET" and request.target.startswith("/api/public/projects"): + return Reply(body=b'{"data":[{"id":"canary-project","name":"canary"}]}') + return Reply(body=b"", content_type="application/x-protobuf") + + +def datadog_sink(request: Request) -> Reply: + return Reply(status=202, body=b"{}") + + +def body_for(endpoint: str, model: str, text: str) -> dict[str, JsonValue]: + if endpoint == "responses": + return {"model": model, "input": text} + if endpoint == "messages": + return {"model": model, "max_tokens": 16, "messages": [{"role": "user", "content": text}]} + return { + "model": model, + "messages": [{"role": "user", "content": text}], + **({"stream": True, "stream_options": {"include_usage": True}} if endpoint == "chat_stream" else {}), + } + + +def send( + gateway: Gateway, key: str, endpoint: str, model: str, text: str, extra: Mapping[str, JsonValue] | None = None +) -> httpx.Response: + path: Final = {"responses": "/v1/responses", "messages": "/v1/messages"}.get(endpoint, "/v1/chat/completions") + return gateway.request("POST", path, {**body_for(endpoint, model, text), **(extra or {})}, key=key) + + +def outcome_text(slot: str, marker: Canary, outcome: str) -> str: + trigger: Final = {"success": "", "provider_4xx": f" {PROVIDER_4XX}", "provider_5xx": f" {PROVIDER_5XX}"}[outcome] + return f"slot {slot} {marker.value}{trigger}" + + +def spend_request_id(marker: Canary) -> str: + rows: Final = eventually( + lambda: read_rows( + 'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE proxy_server_request::text LIKE %s', + (f"%{marker.core}%",), + ), + lambda found: len(found) >= 1, + seconds=70, + ) + return string_value(rows[0]["request_id"]) + + +def wait_for_sink(recorder: Recorder, marker: Canary, seconds: float = 90) -> tuple[Request, ...]: + return eventually( + lambda: tuple(request for request in recorder.requests() if find_canary(request.body, (marker,))), + bool, + seconds=seconds, + ) diff --git a/tests/integration/security/_canary.py b/tests/integration/security/_canary.py index 9787bba3414..9b8bc7be673 100644 --- a/tests/integration/security/_canary.py +++ b/tests/integration/security/_canary.py @@ -83,6 +83,12 @@ SLOTS: Final = MappingProxyType( "A1": Slot("A1", "Virtual key raw value, set as a custom key through /key/generate", prefix="sk-"), "A2": Slot("A2", "Proxy master key from the LITELLM_MASTER_KEY environment variable", prefix="sk-"), "B1": Slot("B1", "Deployment api_key declared in the proxy config.yaml model_list"), + "C1": Slot( + "C1", "Team callback langfuse_secret_key (team callback API, config team settings, callback_settings)" + ), + "C2": Slot("C2", "Key-level callback langfuse_secret_key in key metadata.logging"), + "C3": Slot("C3", "Team callback dd_api_key for the Datadog sink"), + "D5": Slot("D5", "Request-supplied langfuse_secret_key in the request body"), "B2": Slot("B2", "Deployment api_key added through /model/new and stored encrypted"), "B3": Slot("B3", "Credentials table api_key referenced by a deployment's litellm_credential_name"), "B4": Slot("B4", "Deployment aws_secret_access_key added through /model/new"), diff --git a/tests/integration/security/test_callback_credentials.py b/tests/integration/security/test_callback_credentials.py new file mode 100644 index 00000000000..3bc549ea646 --- /dev/null +++ b/tests/integration/security/test_callback_credentials.py @@ -0,0 +1,401 @@ +"""Slots C1, C2, C3 and D5: callback credentials must reach only their sink. + +C1 is the team callback ``langfuse_secret_key`` (team callback API, the deprecated team +``metadata.callback_settings`` and the config ``default_team_settings``), C2 the key-level +``metadata.logging`` Langfuse key, C3 a team callback ``dd_api_key`` for Datadog, and D5 a +``langfuse_secret_key`` the caller sends in the request body (``langfuse_host`` in a body is +rejected without an admin opt-in, so D5 runs on its own proxy with +``general_settings.allow_client_side_credentials`` on). + +Positive control: the owning sink double must receive the request's marker under an auth +header built from the canary (Langfuse ``Basic pk:sk``, Datadog ``DD-API-KEY``), or the test +fails before sweeping. Sensitivity control: the marker must be seen in the stored request body, +the Logs drawer route and the owning sink. Then no sweep may find the canary anywhere else, +including every request the provider double received (swept as the ``provider`` sink, with no +header allowance; the provider's own key is slot B1, which these tests do not search for). +""" + +from __future__ import annotations + +import base64 +from collections.abc import Callable, Iterator, Mapping +from contextlib import contextmanager +from dataclasses import dataclass +from datetime import UTC, datetime +from pathlib import Path +from typing import Final +from urllib.parse import quote + +import pytest +from integration._support.client import Scenario +from integration._support.wire import Request, wire_server +from integration.security._callback_traffic import ( + ENDPOINTS, + EXPECTED_STATUS, + LANGFUSE_PUBLIC_KEY, + OUTCOMES, + datadog_sink, + langfuse_sink, + outcome_text, + send, + spend_request_id, + upstream, + wait_for_sink, +) +from integration.security._canary import MARKER, Canary, canary, find_canary +from integration.security._sinks import CONFIG_MODEL, GENERIC_SINK, Caller, Recorder, Rig, canary_rig +from integration.security._sweeps import assert_marker_seen, assert_no_hits, record_route_sweep, sweep_all +from pydantic import JsonValue + +LANGFUSE: Final = "langfuse" +DATADOG: Final = "datadog" +PROVIDER: Final = "provider" +BOTH: Final = "success_and_failure" + + +@dataclass(frozen=True, slots=True) +class CallbackRig: + rig: Rig + langfuse: Recorder + datadog: Recorder + + def sinks(self) -> dict[str, tuple[Request, ...]]: + return { + **{name: sink.requests() for name, sink in self.rig.sinks.items()}, + LANGFUSE: self.langfuse.requests(), + DATADOG: self.datadog.requests(), + PROVIDER: self.rig.provider.requests(), + } + + def datadog_port(self) -> str: + return self.datadog.url.rsplit(":", 1)[1] + + +@contextmanager +def callback_rig( + root: Path, configure: Callable[[dict[str, object], str, str], None] | None = None +) -> Iterator[CallbackRig]: + with ( + wire_server(langfuse_sink) as langfuse, + wire_server(datadog_sink) as datadog, + canary_rig( + root, + configure=(lambda config, provider: configure(config, provider, langfuse.url)) if configure else None, + environment={"LANGFUSE_FLUSH_INTERVAL": "1"}, + upstream=upstream, + ) as rig, + ): + yield CallbackRig(rig, Recorder(langfuse), Recorder(datadog)) + + +def _allow_client_side_credentials(config: dict[str, object], _provider: str, _langfuse: str) -> None: + settings: Final = config["general_settings"] + assert isinstance(settings, dict) + settings["allow_client_side_credentials"] = True + + +@pytest.fixture(scope="module") +def client_side(tmp_path_factory: pytest.TempPathFactory) -> Iterator[CallbackRig]: + with callback_rig(tmp_path_factory.mktemp("canary-client-side"), _allow_client_side_credentials) as value: + yield value + + +@pytest.fixture(scope="module") +def shared(tmp_path_factory: pytest.TempPathFactory) -> Iterator[CallbackRig]: + with callback_rig(tmp_path_factory.mktemp("canary-callbacks")) as value: + yield value + + +def langfuse_vars(secret: Canary, host: str) -> dict[str, JsonValue]: + return {"langfuse_public_key": LANGFUSE_PUBLIC_KEY, "langfuse_secret_key": secret.value, "langfuse_host": host} + + +def caller( + scenario: Scenario, + *, + team_id: str | None = None, + team_metadata: Mapping[str, JsonValue] | None = None, + key_metadata: Mapping[str, JsonValue] | None = None, +) -> Caller: + team: Final = scenario.team( + **({"team_id": team_id} if team_id else {}), **({"metadata": dict(team_metadata)} if team_metadata else {}) + ) + user: Final = scenario.member(team) + key: Final = scenario.key( + team_id=team, user_id=user, models=[CONFIG_MODEL], **({"metadata": dict(key_metadata)} if key_metadata else {}) + ) + return Caller(team, user, key) + + +def langfuse_control(secret: Canary) -> Callable[[CallbackRig, Canary], None]: + expected: Final = "Basic " + base64.b64encode(f"{LANGFUSE_PUBLIC_KEY}:{secret.value}".encode()).decode() + + def check(rig: CallbackRig, marker: Canary) -> None: + delivered: Final = wait_for_sink(rig.langfuse, marker) + assert {request.headers.get("authorization") for request in delivered} == {expected}, ( + f"Positive control: the Langfuse double never received the {secret.slot} canary as its Basic auth" + ) + + return check + + +def datadog_control(secret: Canary) -> Callable[[CallbackRig, Canary], None]: + def check(rig: CallbackRig, marker: Canary) -> None: + delivered: Final = wait_for_sink(rig.datadog, marker) + assert {request.headers.get("dd-api-key") for request in delivered} == {secret.value}, ( + "Positive control: the Datadog double never received the C3 canary as DD-API-KEY" + ) + + return check + + +def run_scenario( + cb: CallbackRig, + scenario: Scenario, + who: Caller, + secret: Canary, + endpoint: str, + outcome: str, + *, + control: Callable[[CallbackRig, Canary], None], + sink: str, + own_header: tuple[str, str], + node: str, + extra: Mapping[str, JsonValue] | None = None, +) -> None: + marker: Final = canary(MARKER) + started: Final = datetime.now(UTC) + response: Final = send( + cb.rig.proxy, who.key, endpoint, CONFIG_MODEL, outcome_text(secret.slot, marker, outcome), extra + ) + assert response.status_code == EXPECTED_STATUS[outcome], response.text + control(cb, marker) + request_id: Final = spend_request_id(marker) + wait_for_sink(cb.rig.sinks[GENERIC_SINK], marker) + + report: Final = sweep_all( + cb.rig.proxy, + (marker, secret), + responses=(response,), + sinks=cb.sinks(), + ids={ + "request_id": request_id, + "team_id": who.team_id, + "user_id": who.user_id, + "model_id": cb.rig.model_id, + "model": CONFIG_MODEL, + }, + callers=who.callers(cb.rig), + own_headers={**cb.rig.own_headers, sink: own_header}, + since=started, + ) + record_route_sweep(report.routes, node) + assert_marker_seen( + report, + { + "S1": "LiteLLM_SpendLogs.proxy_server_request", + "S2": f"GET /spend/logs/ui/{quote(request_id, safe='')} as admin -> 200", + "S4": f"{sink}[", + }, + ) + assert_marker_seen(report, {"S2": f"GET /spend/logs?request_id={quote(request_id, safe='')} as admin -> 200"}) + assert_marker_seen(report, {"S4": f"{PROVIDER}["}) + assert_no_hits(report.credential_hits(), f"slot {secret.slot}, {endpoint}, {outcome}") + + +MATRIX: Final = [ + pytest.param(endpoint, outcome, id=f"{endpoint}-{outcome}") for endpoint in ENDPOINTS for outcome in OUTCOMES +] + + +@pytest.mark.timeout(240) # full S1/S2 walk: every table and ~430 GET routes as two callers +@pytest.mark.parametrize(("endpoint", "outcome"), MATRIX) +def test_c1_team_callback_api_langfuse_secret_reaches_only_langfuse( + shared: CallbackRig, endpoint: str, outcome: str, request: pytest.FixtureRequest +) -> None: + secret: Final = canary("C1") + with shared.rig.proxy.scenario() as scenario: + who: Final = caller(scenario) + shared.rig.proxy.post( + f"/team/{who.team_id}/callback", + { + "callback_name": "langfuse", + "callback_type": BOTH, + "callback_vars": langfuse_vars(secret, shared.langfuse.url), + }, + ) + run_scenario( + shared, + scenario, + who, + secret, + endpoint, + outcome, + control=langfuse_control(secret), + sink=LANGFUSE, + own_header=("authorization", "C1"), + node=request.node.nodeid, + ) + + +@pytest.mark.timeout(240) # full S1/S2 walk: every table and ~430 GET routes as two callers +@pytest.mark.parametrize("endpoint", ENDPOINTS) +def test_c1_deprecated_team_callback_settings_langfuse_secret_reaches_only_langfuse( + shared: CallbackRig, endpoint: str, request: pytest.FixtureRequest +) -> None: + secret: Final = canary("C1") + settings: Final = { + "success_callback": ["langfuse"], + "failure_callback": ["langfuse"], + "callback_vars": langfuse_vars(secret, shared.langfuse.url), + } + with shared.rig.proxy.scenario() as scenario: + who: Final = caller(scenario, team_metadata={"callback_settings": settings}) + run_scenario( + shared, + scenario, + who, + secret, + endpoint, + "success", + control=langfuse_control(secret), + sink=LANGFUSE, + own_header=("authorization", "C1"), + node=request.node.nodeid, + ) + + +@pytest.mark.timeout(240) # full S1/S2 walk: every table and ~430 GET routes as two callers +@pytest.mark.parametrize("endpoint", ENDPOINTS) +def test_c1_config_default_team_settings_langfuse_secret_reaches_only_langfuse( + tmp_path: Path, endpoint: str, request: pytest.FixtureRequest +) -> None: + """The team callback comes from ``litellm_settings.default_team_settings`` in config.yaml.""" + secret: Final = canary("C1") + team_id: Final = f"canary-config-team-{secret.core[:12]}" + + def configure(config: dict[str, object], _provider: str, langfuse_url: str) -> None: + settings: Final = config["litellm_settings"] + assert isinstance(settings, dict) + settings["default_team_settings"] = [ + { + "team_id": team_id, + "success_callback": ["langfuse"], + "failure_callback": ["langfuse"], + "langfuse_public_key": LANGFUSE_PUBLIC_KEY, + "langfuse_secret": secret.value, + "langfuse_host": langfuse_url, + } + ] + + with callback_rig(tmp_path, configure) as cb, cb.rig.proxy.scenario() as scenario: + who: Final = caller(scenario, team_id=team_id) + run_scenario( + cb, + scenario, + who, + secret, + endpoint, + "success", + control=langfuse_control(secret), + sink=LANGFUSE, + own_header=("authorization", "C1"), + node=request.node.nodeid, + ) + + +@pytest.mark.timeout(240) # full S1/S2 walk: every table and ~430 GET routes as two callers +@pytest.mark.parametrize(("endpoint", "outcome"), MATRIX) +def test_c2_key_logging_langfuse_secret_reaches_only_langfuse( + shared: CallbackRig, endpoint: str, outcome: str, request: pytest.FixtureRequest +) -> None: + secret: Final = canary("C2") + logging: Final = [ + { + "callback_name": "langfuse", + "callback_type": BOTH, + "callback_vars": langfuse_vars(secret, shared.langfuse.url), + } + ] + with shared.rig.proxy.scenario() as scenario: + who: Final = caller(scenario, key_metadata={"logging": logging}) + run_scenario( + shared, + scenario, + who, + secret, + endpoint, + outcome, + control=langfuse_control(secret), + sink=LANGFUSE, + own_header=("authorization", "C2"), + node=request.node.nodeid, + ) + + +@pytest.mark.timeout(240) # full S1/S2 walk: every table and ~430 GET routes as two callers +@pytest.mark.parametrize(("endpoint", "outcome"), MATRIX) +def test_c3_team_callback_datadog_api_key_reaches_only_datadog( + shared: CallbackRig, endpoint: str, outcome: str, request: pytest.FixtureRequest +) -> None: + secret: Final = canary("C3") + with shared.rig.proxy.scenario() as scenario: + who: Final = caller(scenario) + shared.rig.proxy.post( + f"/team/{who.team_id}/callback", + { + "callback_name": "datadog", + "callback_type": BOTH, + "callback_vars": { + "dd_api_key": secret.value, + "dd_agent_host": "127.0.0.1", + "dd_agent_port": shared.datadog_port(), + }, + }, + ) + run_scenario( + shared, + scenario, + who, + secret, + endpoint, + outcome, + control=datadog_control(secret), + sink=DATADOG, + own_header=("dd-api-key", "C3"), + node=request.node.nodeid, + ) + + +@pytest.mark.timeout(240) # full S1/S2 walk: every table and ~430 GET routes as two callers +@pytest.mark.parametrize(("endpoint", "outcome"), MATRIX) +def test_d5_request_body_langfuse_secret_reaches_only_langfuse( + client_side: CallbackRig, endpoint: str, outcome: str, request: pytest.FixtureRequest +) -> None: + secret: Final = canary("D5") + with client_side.rig.proxy.scenario() as scenario: + who: Final = caller(scenario) + run_scenario( + client_side, + scenario, + who, + secret, + endpoint, + outcome, + control=langfuse_control(secret), + sink=LANGFUSE, + own_header=("authorization", "D5"), + node=request.node.nodeid, + extra={ + **langfuse_vars(secret, client_side.langfuse.url), + "success_callback": ["langfuse"], + "failure_callback": ["langfuse"], + }, + ) + + +def test_find_canary_sees_the_langfuse_basic_auth_header() -> None: + """The Langfuse positive control and own-header rule depend on decoding ``Basic pk:sk``.""" + secret: Final = canary("C1") + header: Final = "Basic " + base64.b64encode(f"{LANGFUSE_PUBLIC_KEY}:{secret.value}".encode()).decode() + assert [match.slot for match in find_canary(header, (secret,))] == ["C1"] diff --git a/tests/test_litellm/proxy/test_body_snapshot_callback_params.py b/tests/test_litellm/proxy/test_body_snapshot_callback_params.py new file mode 100644 index 00000000000..b79521fc119 --- /dev/null +++ b/tests/test_litellm/proxy/test_body_snapshot_callback_params.py @@ -0,0 +1,36 @@ +"""The stored request body never carries callback parameters. + +Every ``StandardCallbackDynamicParams`` key and ``litellm_trusted_callback_vars`` is set on the +request dict with a unique value, the body snapshot is refreshed, and none of the keys or values +may be in ``proxy_server_request["body"]``. A control key proves the snapshot was rebuilt. +""" + +from __future__ import annotations + +import json +import uuid +from typing import Final + +from litellm.proxy.litellm_pre_call_utils import refresh_proxy_server_request_body_snapshot +from litellm.types.utils import TRUSTED_CALLBACK_VARS_FIELD, StandardCallbackDynamicParams + + +def test_body_snapshot_excludes_every_callback_dynamic_param_and_the_trusted_vars() -> None: + core: Final = uuid.uuid4().hex + params: Final = {name: f"lkc-{name}-{core}" for name in StandardCallbackDynamicParams.__annotations__} + control: Final = f"control-{uuid.uuid4().hex}" + data: Final = { + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": control}], + **params, + TRUSTED_CALLBACK_VARS_FIELD: dict(params), + "proxy_server_request": {"url": "http://proxy/v1/chat/completions", "body": {}}, + } + + refresh_proxy_server_request_body_snapshot(data) + + body: Final = data["proxy_server_request"]["body"] + assert control in json.dumps(body), "Sensitivity control: the snapshot was not rebuilt from the request" + present: Final = sorted({*params, TRUSTED_CALLBACK_VARS_FIELD} & set(body)) + assert present == [], f"Callback parameters copied into the stored request body: {present}" + assert core not in json.dumps(body, default=str) From cede93e826b2c352de62dcc3bbe725f9728d0352 Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Tue, 29 Sep 2026 10:52:13 -0700 Subject: [PATCH 19/41] test(integration): request-path credential canary slots D1-D4 (#43307) * test(integration): credential canary suite harness Adds tests/integration/security with canary generation and search, sweeps over the database, GET routes, client responses, sink doubles and Redis, an owned proxy rig, a sweep sensitivity self-test and the config deployment api_key slot. Registers the security group in run.py, the manifest and the CircleCI integration matrix. * test(integration): widen canary route sweep and harden the rig Enumerate lazily registered feature routers, call parameterized routes with placeholder ids, fail on routes that return no response, skip provider pass-through routes, add an explicit admin-only route allowance, let the sink double use a configurable token, inflate gzip members anywhere in a blob, sweep Redis before the route walk, and trap outbound connections from the owned proxy. * test(integration): descend into any decoded value that can still hold an encoded canary * test(integration): bound canary decoding by depth and decoded bytes * test(integration): scope log-table and spend-log reads to the scenario window * test(integration): sweep spend-log rows in the scenario date window * test(integration): keep spend-log date window summarized * test(integration): request-path credential canary slots D1-D4 * test(integration): read the Logs drawer and spend-log filter for failed request rows * test(integration): check the marker in each failed row's spend-log filter; run header slots on chat-family routes * test(integration): run D2 on embeddings again; only the client-header slot runs on chat-family routes * test(integration): resolve deployment ids, scope paginated log lists, key allowances by slot * test(integration): pass the slot deployment's model_info id to the route sweep * test(integration): expect 404 from the caller-scoped team membership route * test(integration): use the rig's own master key and expect 404 from submission lookups * test(integration): check the overridden rig key without assuming the default key is unknown --- tests/integration/security/_canary.py | 4 + .../security/test_request_path_slots.py | 507 ++++++++++++++++++ 2 files changed, 511 insertions(+) create mode 100644 tests/integration/security/test_request_path_slots.py diff --git a/tests/integration/security/_canary.py b/tests/integration/security/_canary.py index 9b8bc7be673..2d97c6fde0b 100644 --- a/tests/integration/security/_canary.py +++ b/tests/integration/security/_canary.py @@ -105,6 +105,10 @@ SLOTS: Final = MappingProxyType( "H1": Slot("H1", "Pass-through endpoint credential header resolved from os.environ"), "H2": Slot("H2", "Vector store api_key declared in the proxy config.yaml vector_store_registry"), "H2S": Slot("H2S", "Search tool api_key declared in the proxy config.yaml search_tools"), + "D1": Slot("D1", "Client-side api_key in the request body"), + "D2": Slot("D2", "Client x-api-key header forwarded as the provider key"), + "D3": Slot("D3", "Client x- header forwarded to the provider"), + "D4": Slot("D4", "Anthropic OAuth token in the client Authorization header", prefix="sk-ant-oat01-"), } ) diff --git a/tests/integration/security/test_request_path_slots.py b/tests/integration/security/test_request_path_slots.py new file mode 100644 index 00000000000..dce22eedeea --- /dev/null +++ b/tests/integration/security/test_request_path_slots.py @@ -0,0 +1,507 @@ +"""Request-path slots D1 to D4: a credential the client sends with the request reaches only the provider. + +Each slot is a credential the proxy receives on the request itself and must hand to the provider +without keeping a copy: + +- D1: ``api_key`` in the request body. +- D2: ``x-api-key`` forwarded with ``general_settings.forward_llm_provider_auth_headers``. +- D3: an ``x-goog-api-key`` client header forwarded with + ``litellm_settings.model_group_settings.forward_client_headers_to_llm_api`` (with + ``forward_llm_provider_auth_headers`` on, which lets a provider auth header through). +- D4: an Anthropic OAuth token (``Authorization: Bearer sk-ant-oat...``) sent next to + ``x-litellm-api-key``, forwarded to an Anthropic deployment. + +A test is one slot on one route. It sends three requests carrying the same canary: one the +provider answers, one it rejects with a 4xx and one it fails with a 5xx, because failure logging +takes a different path. Positive control: every provider request of every outcome must carry the +canary where the slot delivers it. Sensitivity control: the marker sent in the same requests must +be in the spend-log row of every outcome and in a sink event of every outcome, and each sweep must +report it where stored prompts belong. The route sweep fills its request-id routes with the +successful row, so the Logs drawer and the spend-log filter are also read for each failed row, +and the marker must show in both. Then no sweep may find the slot's canary anywhere. + +The requests go one at a time, and each waits for its sink event before the next is sent. The +``generic_api`` logger clears its whole queue after a batch POST, so an event queued while a POST +is in flight would be dropped, and the sweep would then miss that outcome's callback payload. + +One owned proxy per slot serves every route of that slot. The canary travels on the request and +never in the config, so a fresh core per test needs no fresh proxy; the config only turns the +slot's setting on. Rows and sink events left by earlier tests carry other cores, which the sweeps +of a later test do not search for. +""" + +from __future__ import annotations + +import json +import uuid +from collections.abc import Callable, Iterator, Mapping +from dataclasses import dataclass +from datetime import UTC, datetime +from types import MappingProxyType +from typing import Final +from urllib.parse import quote, urlencode + +import httpx +import pytest +from integration._support.client import Scenario, eventually, string_value +from integration._support.database import read_rows +from integration._support.wire import Reply, Request +from integration.security._canary import MARKER, Canary, canary, find_canary +from integration.security._sinks import GENERIC_SINK, PROVIDER_4XX, Caller, Rig, canary_rig +from integration.security._sweeps import Hit, assert_marker_seen, assert_no_hits, record_route_sweep, sweep_all + +PROVIDER_5XX: Final = "canary-provider-5xx" +OPENAI_MODEL: Final = "canary-request-openai" +ANTHROPIC_MODEL: Final = "canary-request-anthropic" +FORWARDED_HEADER: Final = "x-goog-api-key" +DEPLOYMENT_KEY: Final = "canary-deployment-placeholder-key" +OUTCOMES: Final = MappingProxyType({"success": 200, "provider_4xx": 400, "provider_5xx": 500}) + + +@dataclass(frozen=True, slots=True) +class Route: + """A client route: its path, the field that carries the prompt text, and fixed extra fields.""" + + path: str + text_field: str + extra: Mapping[str, object] = MappingProxyType({}) + + def body(self, model: str, text: str) -> dict[str, object]: + prompt: Final[object] = [{"role": "user", "content": text}] if self.text_field == "messages" else text + return {"model": model, self.text_field: prompt, **self.extra} + + +ROUTES: Final = MappingProxyType( + { + "chat": Route("/v1/chat/completions", "messages"), + "chat_stream": Route("/v1/chat/completions", "messages", MappingProxyType({"stream": True})), + "messages": Route("/v1/messages", "messages", MappingProxyType({"max_tokens": 16})), + "messages_stream": Route("/v1/messages", "messages", MappingProxyType({"max_tokens": 16, "stream": True})), + "responses": Route("/v1/responses", "input"), + "embeddings": Route("/v1/embeddings", "input"), + } +) + + +@dataclass(frozen=True, slots=True) +class RequestSlot: + """How a slot's canary rides the request, where the provider must receive it, and its setting.""" + + model: str + routes: tuple[str, ...] + body: Callable[[Canary], Mapping[str, object]] + headers: Callable[[Canary, str], Mapping[str, str]] + delivered: Callable[[Request], str | None] + expected: Callable[[Canary], str] + configure: Callable[[dict[str, object]], None] + + +def _no_body(_canary: Canary) -> Mapping[str, object]: + return {} + + +def _bearer_key(_canary: Canary, key: str) -> Mapping[str, str]: + return {"Authorization": f"Bearer {key}"} + + +def _authorization(request: Request) -> str | None: + return request.headers.get("authorization") + + +def _bearer(value: Canary) -> str: + return f"Bearer {value.value}" + + +def _no_setting(_config: dict[str, object]) -> None: + return None + + +def _forward_provider_auth(config: dict[str, object]) -> None: + general: Final = config["general_settings"] + assert isinstance(general, dict) + general["forward_llm_provider_auth_headers"] = True + + +def _forward_client_headers(config: dict[str, object]) -> None: + _forward_provider_auth(config) + settings: Final = config["litellm_settings"] + assert isinstance(settings, dict) + settings["model_group_settings"] = {"forward_client_headers_to_llm_api": [OPENAI_MODEL]} + + +OPENAI_ROUTES: Final = ("chat", "chat_stream", "messages", "responses", "embeddings") +# The client-header forwarding slot (forward_client_headers_to_llm_api) runs on the chat-family routes. +CLIENT_HEADER_ROUTES: Final = ("chat", "chat_stream", "messages", "responses") +ANTHROPIC_ROUTES: Final = ("messages", "messages_stream", "chat", "responses") + +REQUEST_SLOTS: Final = MappingProxyType( + { + "D1": RequestSlot( + OPENAI_MODEL, + OPENAI_ROUTES, + lambda value: {"api_key": value.value}, + _bearer_key, + _authorization, + _bearer, + _no_setting, + ), + "D2": RequestSlot( + OPENAI_MODEL, + OPENAI_ROUTES, + _no_body, + lambda value, key: {"Authorization": f"Bearer {key}", "x-api-key": value.value}, + _authorization, + _bearer, + _forward_provider_auth, + ), + "D3": RequestSlot( + OPENAI_MODEL, + CLIENT_HEADER_ROUTES, + _no_body, + lambda value, key: {"Authorization": f"Bearer {key}", FORWARDED_HEADER: value.value}, + lambda request: request.headers.get(FORWARDED_HEADER), + lambda value: value.value, + _forward_client_headers, + ), + "D4": RequestSlot( + ANTHROPIC_MODEL, + ANTHROPIC_ROUTES, + _no_body, + lambda value, key: {"Authorization": f"Bearer {value.value}", "x-litellm-api-key": key}, + _authorization, + _bearer, + _no_setting, + ), + } +) + + +def _sse(events: tuple[tuple[str | None, dict[str, object]], ...], done: bool) -> tuple[bytes, ...]: + frames: Final = tuple( + (f"event: {name}\n" if name else "").encode() + b"data: " + json.dumps(data).encode() + b"\n\n" + for name, data in events + ) + return (*frames, b"data: [DONE]\n\n") if done else frames + + +def _anthropic_reply(stream: bool) -> Reply: + message: Final = { + "id": f"msg_{uuid.uuid4().hex}", + "type": "message", + "role": "assistant", + "model": "claude-sonnet-4-5", + "content": [{"type": "text", "text": "ok"}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 7, "output_tokens": 3}, + } + if not stream: + return Reply(body=json.dumps(message).encode()) + events: Final = ( + ("message_start", {"type": "message_start", "message": {**message, "content": [], "stop_reason": None}}), + ( + "content_block_start", + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + ), + ( + "content_block_delta", + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "ok"}}, + ), + ("content_block_stop", {"type": "content_block_stop", "index": 0}), + ( + "message_delta", + {"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 3}}, + ), + ("message_stop", {"type": "message_stop"}), + ) + return Reply(content_type="text/event-stream", chunks=_sse(events, done=False)) + + +def _chat_reply(stream: bool) -> Reply: + identity: Final = f"chatcmpl-{uuid.uuid4().hex}" + if not stream: + return Reply( + body=json.dumps( + { + "id": identity, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"} + ], + "usage": {"prompt_tokens": 7, "completion_tokens": 3, "total_tokens": 10}, + } + ).encode() + ) + base: Final = {"id": identity, "object": "chat.completion.chunk", "created": 1, "model": "gpt-4o-mini"} + events: Final = ( + ( + None, + {**base, "choices": [{"index": 0, "delta": {"role": "assistant", "content": "ok"}, "finish_reason": None}]}, + ), + (None, {**base, "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}]}), + (None, {**base, "choices": [], "usage": {"prompt_tokens": 7, "completion_tokens": 3, "total_tokens": 10}}), + ) + return Reply(content_type="text/event-stream", chunks=_sse(events, done=True)) + + +def _responses_reply() -> Reply: + return Reply( + body=json.dumps( + { + "id": f"resp_{uuid.uuid4().hex}", + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [ + { + "type": "message", + "id": f"msg_{uuid.uuid4().hex}", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "ok", "annotations": []}], + } + ], + "parallel_tool_calls": True, + "tool_choice": "auto", + "tools": [], + "usage": {"input_tokens": 7, "output_tokens": 3, "total_tokens": 10}, + } + ).encode() + ) + + +def _embeddings_reply() -> Reply: + return Reply( + body=json.dumps( + { + "object": "list", + "data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3]}], + "model": "text-embedding-3-small", + "usage": {"prompt_tokens": 3, "total_tokens": 3}, + } + ).encode() + ) + + +def _error(status: int, anthropic: bool) -> Reply: + kind: Final = "invalid_request_error" if status < 500 else "api_error" + body: Final = ( + {"type": "error", "error": {"type": kind, "message": "rejected"}} + if anthropic + else {"error": {"type": kind, "code": "canary_rejected", "message": "rejected"}} + ) + return Reply(status=status, body=json.dumps(body).encode()) + + +def provider_upstream(request: Request) -> Reply: + """OpenAI chat, responses and embeddings plus Anthropic messages; fails on the outcome triggers.""" + anthropic: Final = request.target.startswith("/v1/messages") + if PROVIDER_5XX.encode() in request.body: + return _error(500, anthropic) + if PROVIDER_4XX.encode() in request.body: + return _error(400, anthropic) + stream: Final = json.loads(request.body or b"{}").get("stream") is True + if anthropic: + return _anthropic_reply(stream) + if request.target.startswith("/v1/responses"): + return _responses_reply() + if request.target.startswith("/v1/embeddings"): + return _embeddings_reply() + return _chat_reply(stream) + + +def _configure(slot: RequestSlot) -> Callable[[dict[str, object], str], None]: + def configure(config: dict[str, object], provider_url: str) -> None: + models: Final = config["model_list"] + assert isinstance(models, list) + models.extend( + ( + { + "model_name": OPENAI_MODEL, + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_base": provider_url + "/v1", + "api_key": DEPLOYMENT_KEY, + }, + }, + { + "model_name": ANTHROPIC_MODEL, + "litellm_params": { + "model": "anthropic/claude-sonnet-4-5", + "api_base": provider_url, + "api_key": DEPLOYMENT_KEY, + }, + }, + ) + ) + slot.configure(config) + + return configure + + +@pytest.fixture(scope="module") +def rig(request: pytest.FixtureRequest, tmp_path_factory: pytest.TempPathFactory) -> Iterator[Rig]: + """One owned proxy per slot, shared by every route of that slot (see the module docstring).""" + slot_id: Final = str(request.param) + with canary_rig( + tmp_path_factory.mktemp(f"canary-{slot_id}"), + configure=_configure(REQUEST_SLOTS[slot_id]), + upstream=provider_upstream, + ) as value: + yield value + + +def _caller(scenario: Scenario, model: str) -> Caller: + team: Final = scenario.team() + user: Final = scenario.user(user_role="internal_user") + scenario.gateway.post("/team/member_add", {"team_id": team, "member": {"user_id": user, "role": "user"}}) + return Caller(team, user, scenario.key(team_id=team, user_id=user, models=[model])) + + +def _deployment_id(rig: Rig, model: str) -> str: + """The router's ``model_info.id`` for the slot's deployment, for the ``{model_id}`` routes.""" + data: Final = rig.proxy.get("/model/info").get("data") + assert isinstance(data, list), data + found: Final = tuple( + info["id"] + for entry in data + if isinstance(entry, dict) + and entry.get("model_name") == model + and isinstance(info := entry.get("model_info"), dict) + and isinstance(info.get("id"), str) + ) + assert len(found) == 1, f"expected one {model} deployment in /model/info, got {found}" + return str(found[0]) + + +def _trigger(outcome: str) -> str: + return {"success": "", "provider_4xx": f" {PROVIDER_4XX}", "provider_5xx": f" {PROVIDER_5XX}"}[outcome] + + +def _tag(outcome: str, marker: Canary) -> str: + return f"{outcome} {marker.value}" + + +def _spend_rows(marker: Canary) -> list[dict[str, object]]: + return [ + dict(row) + for row in read_rows( + 'SELECT request_id, status FROM "LiteLLM_SpendLogs" WHERE proxy_server_request::text LIKE %s', + (f"%{marker.core}%",), + ) + ] + + +def _request_row_hits( + rig: Rig, callers: Mapping[str, str], request_ids: tuple[str, ...], canaries: tuple[Canary, ...] +) -> tuple[Hit, ...]: + """S2 for the rows the route sweep does not fill in: the Logs drawer and the spend-log filter per row.""" + found: Final[list[Hit]] = [] # mutable-ok: accumulated across rows and callers + for request_id in request_ids: + for path in ( + f"/spend/logs/ui/{quote(request_id, safe='')}", + f"/spend/logs?{urlencode({'request_id': request_id})}", + ): + for label, key in callers.items(): + response = rig.proxy.client.get(path, headers={"Authorization": f"Bearer {key}"}) + where = f"GET {path} as {label} -> {response.status_code}" + found.extend( + Hit("S2", where, match.slot, match.encoding) for match in find_canary(response.content, canaries) + ) + return tuple(found) + + +CASES: Final = tuple( + pytest.param(slot_id, slot_id, route, id=f"{slot_id}-{route}") + for slot_id, slot in REQUEST_SLOTS.items() + for route in slot.routes +) + + +@pytest.mark.timeout(240) # three requests, then the full S1/S2 walk as two callers +@pytest.mark.parametrize(("rig", "slot_id", "route"), CASES, indirect=["rig"], scope="module") +def test_request_credential_reaches_only_the_provider( + rig: Rig, slot_id: str, route: str, request: pytest.FixtureRequest +) -> None: + slot: Final = REQUEST_SLOTS[slot_id] + endpoint: Final = ROUTES[route] + credential: Final = canary(slot_id) + marker: Final = canary(MARKER) + started: Final = datetime.now(UTC) + with rig.proxy.scenario() as scenario: + caller: Final = _caller(scenario, slot.model) + responses: Final[list[httpx.Response]] = [] + for outcome, status in OUTCOMES.items(): + response = rig.proxy.client.post( + endpoint.path, + json={ + **endpoint.body(slot.model, f"slot {slot_id} {_tag(outcome, marker)}{_trigger(outcome)}"), + **slot.body(credential), + }, + headers=dict(slot.headers(credential, caller.key)), + ) + responses.append(response) + assert response.status_code == status, f"{outcome}: {response.status_code} {response.text}" + for name, sink in rig.sinks.items(): + assert eventually( + lambda sink=sink, outcome=outcome: sink.carrying(_tag(outcome, marker)), + bool, + seconds=30, + return_last_on_timeout=True, + ), f"Sensitivity control: {name} never received the {outcome} event" + + for outcome in OUTCOMES: + delivered = rig.provider.carrying(_tag(outcome, marker)) + assert delivered and all(slot.delivered(each) == slot.expected(credential) for each in delivered), ( + f"Positive control: the provider double never received the {slot_id} canary for {outcome}: " + f"{[dict(each.headers) for each in delivered]}" + ) + + rows: Final = eventually(lambda: _spend_rows(marker), lambda found: len(found) == len(OUTCOMES), seconds=70) + assert sorted(string_value(row["status"]) for row in rows) == ["failure", "failure", "success"], rows + request_id: Final = next(string_value(row["request_id"]) for row in rows if row["status"] == "success") + failed_ids: Final = tuple(string_value(row["request_id"]) for row in rows if row["status"] == "failure") + + report: Final = sweep_all( + rig.proxy, + (marker, credential), + responses=tuple(responses), + sinks={name: sink.requests() for name, sink in rig.sinks.items()}, + ids={ + "request_id": request_id, + "team_id": caller.team_id, + "user_id": caller.user_id, + "model_id": _deployment_id(rig, slot.model), + "model": slot.model, + }, + callers=caller.callers(rig), + own_headers=rig.own_headers, + since=started, + ) + record_route_sweep(report.routes, request.node.nodeid) + assert_marker_seen( + report, + { + "S1": "LiteLLM_SpendLogs.proxy_server_request", + "S2": f"GET /spend/logs/ui/{quote(request_id, safe='')} as admin -> 200", + "S4": f"{GENERIC_SINK}[", + }, + ) + assert_marker_seen(report, {"S2": f"GET /spend/logs?{urlencode({'request_id': request_id})} as admin -> 200"}) + failure_rows: Final = _request_row_hits(rig, caller.callers(rig), failed_ids, (marker, credential)) + for failed_id in failed_ids: + for path in ( + f"/spend/logs/ui/{quote(failed_id, safe='')}", + f"/spend/logs?{urlencode({'request_id': failed_id})}", + ): + where = f"GET {path} as admin -> 200" + assert any(hit.slot == MARKER and hit.location == where for hit in failure_rows), ( + f"Sensitivity control: the marker is missing from {where}" + ) + assert_no_hits( + (*report.credential_hits(), *(hit for hit in failure_rows if hit.slot != MARKER)), + f"slot {slot_id}, {endpoint.path} ({route})", + ) From 336c7c08496c1938824e52361b738edbeda6b7a6 Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Tue, 29 Sep 2026 10:57:28 -0700 Subject: [PATCH 20/41] test(integration): sweep proxy logs, metrics, a Datadog intake and the Logs drawer for credential canaries (#43306) * test(integration): credential canary suite harness Adds tests/integration/security with canary generation and search, sweeps over the database, GET routes, client responses, sink doubles and Redis, an owned proxy rig, a sweep sensitivity self-test and the config deployment api_key slot. Registers the security group in run.py, the manifest and the CircleCI integration matrix. * test(integration): widen canary route sweep and harden the rig Enumerate lazily registered feature routers, call parameterized routes with placeholder ids, fail on routes that return no response, skip provider pass-through routes, add an explicit admin-only route allowance, let the sink double use a configurable token, inflate gzip members anywhere in a blob, sweep Redis before the route walk, and trap outbound connections from the owned proxy. * test(integration): descend into any decoded value that can still hold an encoded canary * test(integration): bound canary decoding by depth and decoded bytes * test(integration): scope log-table and spend-log reads to the scenario window * test(integration): sweep spend-log rows in the scenario date window * test(integration): keep spend-log date window summarized * test(integration): sweep proxy logs, metrics, a gzip Datadog intake and the Logs drawer for credential canaries * test(e2e): treat an unset prompt-storage setting as unset and restore it * test(integration): name the Datadog sink slot G1d * test(e2e): search the Logs page for base64 forms of the deployment key * test(integration): resolve deployment ids, scope paginated log lists, key allowances by slot * test(integration): pass the resolved deployment id to the Datadog route sweep * test(integration): expect 404 from the caller-scoped team membership route * test(integration): use the rig's own master key and expect 404 from submission lookups * test(integration): check the overridden rig key without assuming the default key is unknown --- .../tests/integrationCritical/expected.json | 3 +- .../logsDrawerCredentialCanary.spec.ts | 224 ++++++++++++++++++ tests/integration/security/_canary.py | 1 + .../integration/security/test_datadog_sink.py | 169 +++++++++++++ .../integration/security/test_metrics_text.py | 89 +++++++ tests/integration/security/test_proxy_logs.py | 84 +++++++ 6 files changed, 569 insertions(+), 1 deletion(-) create mode 100644 tests/e2e/ui/tests/integrationCritical/logsDrawerCredentialCanary.spec.ts create mode 100644 tests/integration/security/test_datadog_sink.py create mode 100644 tests/integration/security/test_metrics_text.py create mode 100644 tests/integration/security/test_proxy_logs.py diff --git a/tests/e2e/ui/tests/integrationCritical/expected.json b/tests/e2e/ui/tests/integrationCritical/expected.json index b73c04acbe0..1614b188188 100644 --- a/tests/e2e/ui/tests/integrationCritical/expected.json +++ b/tests/e2e/ui/tests/integrationCritical/expected.json @@ -7,5 +7,6 @@ "tests/e2e/ui/tests/integrationCritical/mcpUserEnvVars.spec.ts::a server with two per-user variables reports the remaining gap until both are saved", "tests/e2e/ui/tests/integrationCritical/mcpUserEnvVars.spec.ts::a server without per-user variables shows no credential row", "tests/e2e/ui/tests/integrationCritical/mcpUserEnvVars.spec.ts::clearing credentials for a server deleted underneath the modal reports the failure without losing the page", - "tests/e2e/ui/tests/integrationCritical/costOptimizationModelGroups.spec.ts::cache leakage by model merges a deployment's resolved and requested model names into its model group" + "tests/e2e/ui/tests/integrationCritical/costOptimizationModelGroups.spec.ts::cache leakage by model merges a deployment's resolved and requested model names into its model group", + "tests/e2e/ui/tests/integrationCritical/logsDrawerCredentialCanary.spec.ts::the Logs drawer renders the stored request without the deployment api_key" ] diff --git a/tests/e2e/ui/tests/integrationCritical/logsDrawerCredentialCanary.spec.ts b/tests/e2e/ui/tests/integrationCritical/logsDrawerCredentialCanary.spec.ts new file mode 100644 index 00000000000..1d86dee107d --- /dev/null +++ b/tests/e2e/ui/tests/integrationCritical/logsDrawerCredentialCanary.spec.ts @@ -0,0 +1,224 @@ +import { test, expect, type APIRequestContext } from "@playwright/test"; +import { randomUUID } from "node:crypto"; +import { Page } from "../../fixtures/pages"; +import { dismissFeedbackPopup, navigateToPage } from "../../helpers/navigation"; + +/** + * Credential canary S8: what the Logs page renders for a request, including any client-side + * merge, never shows the deployment api_key that served it. + * + * A deployment is registered with a fresh canary api_key pointing at the owned upstream. The + * upstream must receive that canary as its bearer (positive control). The request carries a + * marker in its message content with stored prompts on, and the drawer must render that marker + * in both its pretty view and its raw request JSON view (sensitivity control: the stored request + * really reached the page) while the page's DOM holds no copy of the canary core in either view, + * raw or base64-encoded. + */ +const unhex = (): string => randomUUID().replaceAll("-", ""); + +/** + * The forms the canary core can take on the page: raw (JSON and percent encoding leave a hex + * core unchanged), and base64 in the standard and URL-safe alphabets at each of the three byte + * alignments it can start at. Each base64 form keeps only the characters that depend on core + * bytes alone, so it matches whatever bytes precede or follow the core. + */ +const canaryForms = (core: string): ReadonlyMap => { + const forms = new Map([["raw", core]]); + for (let offset = 0; offset < 3; offset++) { + const bytes = Buffer.concat([Buffer.alloc(offset), Buffer.from(core)]); + const first = offset === 0 ? 0 : 4; + const last = Math.floor(bytes.length / 3) * 4; + const text = bytes.toString("base64").slice(first, last); + forms.set(`base64@${offset}`, text); + forms.set( + `base64url@${offset}`, + text.replaceAll("+", "-").replaceAll("/", "_"), + ); + } + return forms; +}; + +/** The names of the canary forms found in ``text``; the raw form ignores case. */ +const foundForms = ( + text: string, + forms: ReadonlyMap, +): string[] => + [...forms] + .filter(([name, needle]) => + name === "raw" + ? text.toLowerCase().includes(needle) + : text.includes(needle), + ) + .map(([name]) => name); + +test("the Logs drawer renders the stored request without the deployment api_key", async ({ + page, + request, +}) => { + const master = process.env.LITELLM_MASTER_KEY ?? "sk-integration-master"; + const upstream = ( + process.env.INTEGRATION_UPSTREAM_URL ?? "http://127.0.0.1:8190" + ).replace(/\/+$/, ""); + const auth = { Authorization: `Bearer ${master}` }; + const canaryCore = unhex(); + const deploymentKey = `lkc-B1-${canaryCore}`; + const forms = canaryForms(canaryCore); + for (const prefix of ["", "k", "k:"]) { + const encoded = Buffer.from(`${prefix}${deploymentKey}`).toString("base64"); + expect( + foundForms(`Basic ${encoded}`, forms), + `the decoder misses base64 after a ${prefix.length}-byte prefix`, + ).not.toEqual([]); + } + const marker = `lkc-M0-${unhex()}`; + const model = `canary-drawer-${unhex()}`; + + const post = async (api: APIRequestContext, path: string, data: object) => { + const response = await api.post(path, { headers: auth, data }); + expect(response.status(), `POST ${path}: ${await response.text()}`).toBe( + 200, + ); + return response.json(); + }; + + const setting = await request.get( + "/config/field/info?field_name=store_prompts_in_spend_logs", + { headers: auth }, + ); + // A fresh database has no stored value, and the route answers 400 "... is not set". + const settingText = await setting.text(); + expect( + setting.status() === 200 || settingText.includes("is not set"), + settingText, + ).toBe(true); + const promptsStored: boolean | null = + setting.status() === 200 + ? JSON.parse(settingText).field_value === true + : null; + let modelId = ""; + try { + await post(request, "/config/update", { + general_settings: { store_prompts_in_spend_logs: true }, + }); + const created = await post(request, "/model/new", { + model_name: model, + litellm_params: { + model: "openai/gpt-4o-mini", + api_key: deploymentKey, + api_base: `${upstream}/v1`, + }, + }); + modelId = created.model_id; + let requestId = ""; + await expect + .poll( + async () => { + const response = await request.post("/v1/chat/completions", { + headers: auth, + data: { + model, + messages: [{ role: "user", content: `drawer ${marker}` }], + }, + }); + if (response.status() === 200) requestId = (await response.json()).id; + return response.status(); + }, + { + timeout: 30_000, + message: "the new deployment never served the request", + }, + ) + .toBe(200); + + const observed = await request.get(`${upstream}/__observations`); + const delivered = ( + (await observed.json()).requests as { + authorization: string; + body: unknown; + }[] + ).filter((entry) => JSON.stringify(entry.body).includes(marker)); + expect( + delivered.map((entry) => entry.authorization), + "Positive control: the upstream never received the deployment key", + ).toEqual([`Bearer ${deploymentKey}`]); + + await expect + .poll( + async () => { + const response = await request.get( + `/spend/logs/ui/${encodeURIComponent(requestId)}`, + { headers: auth }, + ); + return response.status() === 200 + ? JSON.stringify(await response.json()).includes(marker) + : false; + }, + { + timeout: 70_000, + message: `the stored request for ${requestId} never carried the marker`, + }, + ) + .toBe(true); + + await page.goto("/ui/login"); + await page.getByPlaceholder("Enter your username").fill("admin"); + await page.getByPlaceholder("Enter your password").fill(master); + await page.getByRole("button", { name: "Login", exact: true }).click(); + await expect(page).toHaveURL( + (url) => + url.pathname.startsWith("/ui") && !url.pathname.includes("login"), + ); + await navigateToPage(page, Page.Logs); + await dismissFeedbackPopup(page); + + const search = page + .getByTestId("datatable-search") + .filter({ visible: true }); + await expect(search).toBeVisible({ timeout: 20_000 }); + await search.fill(requestId); + const row = page + .locator("table") + .filter({ visible: true }) + .first() + .locator("tbody tr") + .filter({ hasText: requestId }); + await expect(row).toHaveCount(1, { timeout: 30_000 }); + await row.click(); + + const drawer = page.getByRole("dialog").first(); + await expect(drawer.getByText("Request & Response")).toBeVisible({ + timeout: 20_000, + }); + await expect( + drawer.getByText(marker, { exact: false }).first(), + ).toBeVisible({ timeout: 20_000 }); + expect( + foundForms(await page.content(), forms), + "the drawer's pretty view holds the deployment api_key", + ).toEqual([]); + + await drawer.getByRole("tab", { name: "JSON", exact: true }).click(); + await drawer.getByRole("tab", { name: "Request", exact: true }).click(); + const requestJson = drawer + .getByRole("tabpanel") + .filter({ hasText: marker }) + .last(); + await expect(requestJson).toBeVisible({ timeout: 20_000 }); + expect( + foundForms(await page.content(), forms), + "the drawer's request JSON holds the deployment api_key", + ).toEqual([]); + } finally { + if (modelId) await post(request, "/model/delete", { id: modelId }); + if (promptsStored === null) { + await post(request, "/config/field/delete", { + config_type: "general_settings", + field_name: "store_prompts_in_spend_logs", + }); + } else { + await post(request, "/config/update", { + general_settings: { store_prompts_in_spend_logs: promptsStored }, + }); + } + } +}); diff --git a/tests/integration/security/_canary.py b/tests/integration/security/_canary.py index 2d97c6fde0b..e58bda0936d 100644 --- a/tests/integration/security/_canary.py +++ b/tests/integration/security/_canary.py @@ -83,6 +83,7 @@ SLOTS: Final = MappingProxyType( "A1": Slot("A1", "Virtual key raw value, set as a custom key through /key/generate", prefix="sk-"), "A2": Slot("A2", "Proxy master key from the LITELLM_MASTER_KEY environment variable", prefix="sk-"), "B1": Slot("B1", "Deployment api_key declared in the proxy config.yaml model_list"), + "G1d": Slot("G1d", "Logging sink credential read from the proxy environment (DD_API_KEY)"), "C1": Slot( "C1", "Team callback langfuse_secret_key (team callback API, config team settings, callback_settings)" ), diff --git a/tests/integration/security/test_datadog_sink.py b/tests/integration/security/test_datadog_sink.py new file mode 100644 index 00000000000..be8f866e260 --- /dev/null +++ b/tests/integration/security/test_datadog_sink.py @@ -0,0 +1,169 @@ +"""Slot G1d through a Datadog intake double: the sink key reaches only its own auth header. + +The owned proxy enables the ``datadog`` callback with ``DD_API_KEY`` set to a fresh G1d canary +and ``DD_BASE_URL`` pointed at a local intake double. Datadog batches are gzip-compressed JSON +(a single event sent on the sync path is plain JSON), so the double inflates ``Content-Encoding: +gzip`` bodies, requires JSON log events, answers 202 like the real intake, and records the bytes +exactly as received for S4 (``find_canary`` inflates them). Events the route sweep itself +produces are swept again after it. + +Positive control: the intake double must receive ``DD-API-KEY: `` on the batch +carrying the scenario's marker, and the provider double ``Authorization: Bearer ``. +Sensitivity control: the marker must be found inside the gzip body (encoding ``gzip``), in the +stored spend row, on the Logs drawer route and in the generic sink. Then S1 to S5 plus the +intake double may not hold B1 or G1d anywhere, except G1d in the intake's own ``dd-api-key`` and +on the proxy admin's callback settings route (``ADMIN_ONLY_ALLOWANCES``). That route's gate for +everyone else is asserted directly: the internal user gets 401, and a ``proxy_admin_viewer`` +must read ``DD_API_KEY`` as ``REDACTED``. Routes are swept as the admin, the internal user and +that admin viewer. +""" + +from __future__ import annotations + +import gzip +import json +from collections.abc import Iterator +from datetime import UTC, datetime +from pathlib import Path +from typing import Final + +import pytest +from integration._support.client import eventually, string_value +from integration._support.wire import Reply, Request, wire_server +from integration.security._canary import MARKER, Canary, canary +from integration.security._sinks import CONFIG_MODEL, GENERIC_SINK, Recorder, Rig, canary_rig, settle, team_caller +from integration.security._sweeps import ( + assert_marker_seen, + assert_no_hits, + record_route_sweep, + sweep_all, + sweep_sink, +) + +DATADOG_SINK: Final = "datadog" +DATADOG_KEY_HEADER: Final = "dd-api-key" +CALLBACK_SETTINGS_ROUTE: Final = "/get/config/callbacks" + + +def inflated(request: Request) -> bytes: + """The body as Datadog reads it: batches are gzip-compressed, single sync events are not.""" + return gzip.decompress(request.body) if request.headers.get("content-encoding") == "gzip" else request.body + + +def datadog_intake(request: Request) -> Reply: + assert request.target == "/api/v2/logs", request.target + events: Final = json.loads(inflated(request)) + assert isinstance(events, (list, dict)) and events, events + return Reply(status=202, body=b"{}") + + +def enable_datadog(config: dict[str, object], _provider_url: str) -> None: + settings: Final = config["litellm_settings"] + assert isinstance(settings, dict) + settings["callbacks"] = [*settings["callbacks"], DATADOG_SINK] + + +@pytest.fixture +def intake() -> Iterator[Recorder]: + with wire_server(datadog_intake) as wire: + yield Recorder(wire) + + +@pytest.fixture +def g1() -> Canary: + return canary("G1d") + + +@pytest.fixture +def rig(tmp_path: Path, intake: Recorder, g1: Canary) -> Iterator[Rig]: + environment: Final = {"DD_API_KEY": g1.value, "DD_SITE": "datadog.invalid", "DD_BASE_URL": intake.url} + with canary_rig(tmp_path, configure=enable_datadog, environment=environment) as value: + yield value + + +def carrying_inflated(intake: Recorder, marker: Canary) -> tuple[Request, ...]: + """Gzip batches whose inflated body holds ``marker``.""" + return tuple( + request + for request in intake.requests() + if request.headers.get("content-encoding") == "gzip" and marker.core.encode() in inflated(request) + ) + + +@pytest.mark.timeout(240) # full S1/S2 walk: every table and ~400 GET routes as three callers +def test_datadog_api_key_reaches_only_its_own_header( + rig: Rig, intake: Recorder, g1: Canary, request: pytest.FixtureRequest +) -> None: + b1: Final = rig.canaries["B1"] + marker: Final = canary(MARKER) + started: Final = datetime.now(UTC) + with rig.proxy.scenario() as scenario: + caller: Final = team_caller(scenario) + response: Final = rig.proxy.request( + "POST", + "/v1/chat/completions", + {"model": CONFIG_MODEL, "messages": [{"role": "user", "content": f"slot G1d {marker.value}"}]}, + key=caller.key, + ) + assert response.status_code == 200, response.text + assert [request.headers.get("authorization") for request in rig.provider.carrying(marker.value)] == [ + f"Bearer {b1.value}" + ], "Positive control: the provider double never received the B1 canary" + request_id: Final = string_value(response.json()["id"]) + settle(rig, request_id, marker) + batches: Final = eventually(lambda: carrying_inflated(intake, marker), bool, seconds=30) + assert {batch.headers.get(DATADOG_KEY_HEADER) for batch in batches} == {g1.value}, ( + "Positive control: the Datadog intake double never received the G1d canary" + ) + assert all(marker.core.encode() not in batch.body for batch in batches), "Datadog body was not compressed" + + denied: Final = rig.proxy.request("GET", CALLBACK_SETTINGS_ROUTE, key=caller.key) + assert denied.status_code == 401, f"internal_user read the callback settings: {denied.text}" + viewer: Final = scenario.key(user_id=scenario.user(user_role="proxy_admin_viewer")) + settings: Final = rig.proxy.request("GET", CALLBACK_SETTINGS_ROUTE, key=viewer) + assert settings.status_code == 200, settings.text + datadog_variables: Final = [ + entry["variables"] for entry in settings.json()["callbacks"] if entry["name"] == DATADOG_SINK + ] + assert datadog_variables and all(variables["DD_API_KEY"] == "REDACTED" for variables in datadog_variables), ( + f"The admin viewer's callback settings did not redact DD_API_KEY: {datadog_variables}" + ) + + swept: Final = intake.requests() + report: Final = sweep_all( + rig.proxy, + (marker, b1, g1), + responses=(response,), + sinks={**{name: sink.requests() for name, sink in rig.sinks.items()}, DATADOG_SINK: swept}, + ids={ + "request_id": request_id, + "team_id": caller.team_id, + "user_id": caller.user_id, + "model_id": rig.model_id, + "model": CONFIG_MODEL, + }, + callers={**caller.callers(rig), "admin_viewer": viewer}, + own_headers={**rig.own_headers, DATADOG_SINK: (DATADOG_KEY_HEADER, "G1d")}, + since=started, + ) + record_route_sweep(report.routes, request.node.nodeid) + assert_marker_seen( + report, + { + "S1": "LiteLLM_SpendLogs.proxy_server_request", + "S2": f"GET /spend/logs/ui/{request_id} as admin -> 200", + "S4": f"{GENERIC_SINK}[", + }, + ) + assert_marker_seen(report, {"S2": f"GET /spend/logs?request_id={request_id} as admin -> 200"}) + assert any( + hit.slot == MARKER and hit.location.startswith(f"{DATADOG_SINK}[") and hit.encoding == "gzip" + for hit in report.hits + ), f"Sensitivity control: S4 never inflated the marker out of the Datadog body: {report.marker_locations()}" + late: Final = sweep_sink( + f"{DATADOG_SINK} after the route sweep", + intake.requests()[len(swept) :], + (b1, g1), + own_header=(DATADOG_KEY_HEADER, "G1d"), + ) + assert_no_hits((*report.credential_hits(), *late), "slots B1 and G1d, Datadog intake") diff --git a/tests/integration/security/test_metrics_text.py b/tests/integration/security/test_metrics_text.py new file mode 100644 index 00000000000..e5a0fdc339e --- /dev/null +++ b/tests/integration/security/test_metrics_text.py @@ -0,0 +1,89 @@ +"""S7: the Prometheus ``/metrics/`` text never carries a credential canary. + +Metric label values come from request fields (caller, model, route, user agent, exception +class), so a credential copied into one of them would be served to every scraper. The owned +proxy enables the ``prometheus`` callback, sends one successful and one provider-rejected chat +completion, and searches the whole scrape. + +Positive control: the provider double must receive ``Authorization: Bearer `` for +both requests (their content carries the fresh marker, so neither is served from the response +cache). Sensitivity control: both requests send the marker as their ``User-Agent``, +which the proxy exports as the ``user_agent`` label, so the scrape must carry the marker on +the success and the failure series before the credential search counts. +""" + +from __future__ import annotations + +from collections.abc import Iterator +from pathlib import Path +from typing import Final + +import pytest +from integration._support.client import eventually +from integration.security._canary import MARKER, Canary, canary, find_canary +from integration.security._sinks import CONFIG_MODEL, PROVIDER_4XX, Rig, canary_rig, team_caller +from integration.security._sweeps import Hit, assert_no_hits + +METRICS_ROUTE: Final = "/metrics/" + + +def enable_prometheus(config: dict[str, object], _provider_url: str) -> None: + settings: Final = config["litellm_settings"] + assert isinstance(settings, dict) + settings["callbacks"] = [*settings["callbacks"], "prometheus"] + + +def sweep_metrics(text: str, canaries: tuple[Canary, ...]) -> tuple[Hit, ...]: + """Every canary in the scrape, attributed to the series line that holds it.""" + if not find_canary(text, canaries): + return () + return tuple( + Hit("S7", f"GET {METRICS_ROUTE} line {number}: {line[:160]!r}", match.slot, match.encoding) + for number, line in enumerate(text.splitlines(), start=1) + for match in find_canary(line, canaries) + ) + + +@pytest.fixture +def rig(tmp_path: Path) -> Iterator[Rig]: + with canary_rig(tmp_path, configure=enable_prometheus) as value: + yield value + + +def test_metrics_text_carries_no_credential(rig: Rig) -> None: + b1: Final = rig.canaries["B1"] + marker: Final = canary(MARKER) + agent: Final = f"canary-agent/{marker.value}" + with rig.proxy.scenario() as scenario: + caller: Final = team_caller(scenario) + responses: Final = tuple( + rig.proxy.request( + "POST", + "/v1/chat/completions", + {"model": CONFIG_MODEL, "messages": [{"role": "user", "content": text}]}, + key=caller.key, + headers={"User-Agent": agent}, + ) + for text in (f"slot B1 metrics {marker.value}", f"slot B1 metrics {marker.value} {PROVIDER_4XX}") + ) + assert [response.status_code for response in responses] == [200, 400], [r.text for r in responses] + delivered: Final = rig.provider.carrying(marker.value) + assert [request.headers.get("authorization") for request in delivered] == [f"Bearer {b1.value}"] * 2, ( + "Positive control: the provider double never received the B1 canary" + ) + + def scrape() -> str: + response: Final = rig.proxy.request("GET", METRICS_ROUTE) + assert response.status_code == 200, response.text + return response.text + + def both_outcomes_exported(text: str) -> bool: + lines: Final = text.splitlines() + return all( + any(marker.core in line and f'status_code="{status}"' in line for line in lines) + for status in ("200", "400") + ) + + hits: Final = sweep_metrics(eventually(scrape, both_outcomes_exported, seconds=30), (marker, b1)) + assert any(hit.slot == MARKER for hit in hits), "Sensitivity control: the scrape never carried the marker" + assert_no_hits(tuple(hit for hit in hits if hit.slot != MARKER), "slot B1, metrics text") diff --git a/tests/integration/security/test_proxy_logs.py b/tests/integration/security/test_proxy_logs.py new file mode 100644 index 00000000000..514bd69bb38 --- /dev/null +++ b/tests/integration/security/test_proxy_logs.py @@ -0,0 +1,84 @@ +"""S6: the owned proxy's own stdout and stderr never carry a credential canary. + +Each leg boots its own proxy (slot B1 lives in its config), sends one successful and one +provider-rejected chat completion, stops the proxy so every buffered write reaches the log +file, and then searches the whole captured log. The ``default`` leg runs with ``LITELLM_LOG`` +unset, the level an operator gets out of the box; the ``debug`` leg runs with +``LITELLM_LOG=DEBUG``, which prints request data, router decisions and provider calls. + +Positive control: the provider double must receive ``Authorization: Bearer `` for +both requests. Sensitivity control: the provider double echoes the rejected message in its +error text, and the proxy logs that error at every level, so the marker must be found in the +log; a capture that misses the log file or reads it before the writes land fails there. +""" + +from __future__ import annotations + +import json +from pathlib import Path +from types import MappingProxyType +from typing import Final + +import pytest +from integration._support.client import string_value +from integration._support.wire import Reply, Request +from integration.security._canary import MARKER, Canary, canary, find_canary +from integration.security._sinks import CONFIG_MODEL, PROVIDER_4XX, canary_rig, chat_upstream, settle, team_caller +from integration.security._sweeps import Hit, assert_no_hits + +LEGS: Final = MappingProxyType({"default": MappingProxyType({}), "debug": MappingProxyType({"LITELLM_LOG": "DEBUG"})}) + + +def echoing_upstream(request: Request) -> Reply: + """``chat_upstream``, except a rejection repeats the rejected message in its error text.""" + body: Final = json.loads(request.body or b"{}") + text: Final = str((body.get("messages") or [{}])[-1].get("content", "")) + if PROVIDER_4XX not in text: + return chat_upstream(request) + return Reply( + status=400, + body=json.dumps( + {"error": {"type": "invalid_request_error", "code": "canary_rejected", "message": f"rejected: {text}"}} + ).encode(), + ) + + +def sweep_log(path: Path, canaries: tuple[Canary, ...]) -> tuple[Hit, ...]: + """Every canary in the captured log, attributed to the line that holds it.""" + data: Final = path.read_bytes() + if not find_canary(data, canaries): + return () + return tuple( + Hit("S6", f"{path.name} line {number}: {line[:160]!r}", match.slot, match.encoding) + for number, line in enumerate(data.splitlines(), start=1) + for match in find_canary(line, canaries) + ) + + +@pytest.mark.parametrize("leg", tuple(LEGS)) +def test_proxy_log_carries_no_credential(leg: str, tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("LITELLM_LOG", raising=False) + marker: Final = canary(MARKER) + with canary_rig(tmp_path, environment=LEGS[leg], upstream=echoing_upstream) as rig: + b1: Final = rig.canaries["B1"] + with rig.proxy.scenario() as scenario: + caller: Final = team_caller(scenario) + responses: Final = tuple( + rig.proxy.request( + "POST", + "/v1/chat/completions", + {"model": CONFIG_MODEL, "messages": [{"role": "user", "content": f"slot B1 {suffix}"}]}, + key=caller.key, + ) + for suffix in (marker.value, f"{marker.value} {PROVIDER_4XX}") + ) + assert [response.status_code for response in responses] == [200, 400], [r.text for r in responses] + delivered: Final = rig.provider.carrying(marker.value) + assert [request.headers.get("authorization") for request in delivered] == [f"Bearer {b1.value}"] * 2, ( + "Positive control: the provider double never received the B1 canary" + ) + settle(rig, string_value(responses[0].json()["id"]), marker) + log: Final = rig.owned.log + hits: Final = sweep_log(log, (marker, b1)) + assert any(hit.slot == MARKER for hit in hits), f"Sensitivity control: the marker never reached {log}" + assert_no_hits(tuple(hit for hit in hits if hit.slot != MARKER), f"slot B1, proxy log, {leg} level") From bda2763f2c722d617171e71deb331d3de7d345e7 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 18:06:14 +0000 Subject: [PATCH 21/41] chore(cost-map): add azure and openrouter gpt-6.1-sol rows (#43744) Co-authored-by: kerry Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- ...odel_prices_and_context_window_backup.json | 245 ++++++++++++++++++ model_prices_and_context_window.json | 245 ++++++++++++++++++ 2 files changed, 490 insertions(+) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 4257649ab41..75241031730 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -3917,6 +3917,55 @@ "supports_web_search": true, "supports_xhigh_reasoning_effort": true }, + "azure_ai/gpt-6.1-sol": { + "cache_creation_input_token_cost": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5e-06, + "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_above_272k_tokens": 2e-07, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "litellm_provider": "azure_ai", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1e-05, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://learn.microsoft.com/en-us/azure/ai-foundry/openai/concepts/models", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_cache_breakpoint": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, "azure_ai/gpt-5.5": { "deprecation_date": "2027-10-26", "cache_read_input_token_cost": 5e-07, @@ -8509,6 +8558,102 @@ "supports_web_search": true, "supports_xhigh_reasoning_effort": true }, + "azure/gpt-6.1-sol": { + "cache_creation_input_token_cost": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5e-06, + "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_above_272k_tokens": 2e-07, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "litellm_provider": "azure", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1e-05, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://learn.microsoft.com/en-us/azure/ai-foundry/openai/concepts/models", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, + "azure/gpt-6.1-sol-2026-09-29": { + "cache_creation_input_token_cost": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5e-06, + "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_above_272k_tokens": 2e-07, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "litellm_provider": "azure", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1e-05, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://learn.microsoft.com/en-us/azure/ai-foundry/openai/concepts/models", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, "azure/gpt-chat-latest": { "cache_read_input_token_cost": 5e-07, "deprecation_date": "2026-12-02", @@ -77430,6 +77575,106 @@ "supports_vision": true, "supports_web_search": true }, + "openrouter/openai/gpt-6.1-sol": { + "cache_creation_input_token_cost": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5e-06, + "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_above_272k_tokens": 2e-07, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "litellm_provider": "openrouter", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1e-05, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": false, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, + "openrouter/openai/gpt-6.1-sol-pro": { + "cache_creation_input_token_cost": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5e-06, + "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_above_272k_tokens": 2e-07, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "litellm_provider": "openrouter", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1e-05, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": false, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, + "openrouter/openai/gpt-6.1-sol-pro:batch": { + "cache_creation_input_token_cost": 1.25e-06, + "cache_creation_input_token_cost_above_272k_tokens": 2.5e-06, + "cache_read_input_token_cost": 5e-08, + "cache_read_input_token_cost_above_272k_tokens": 1e-07, + "input_cost_per_token": 1e-06, + "input_cost_per_token_above_272k_tokens": 2e-06, + "litellm_provider": "openrouter", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-06, + "output_cost_per_token_above_272k_tokens": 7.5e-06, + "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": false, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, + "openrouter/openai/gpt-6.1-sol:batch": { + "cache_creation_input_token_cost": 1.25e-06, + "cache_creation_input_token_cost_above_272k_tokens": 2.5e-06, + "cache_read_input_token_cost": 5e-08, + "cache_read_input_token_cost_above_272k_tokens": 1e-07, + "input_cost_per_token": 1e-06, + "input_cost_per_token_above_272k_tokens": 2e-06, + "litellm_provider": "openrouter", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-06, + "output_cost_per_token_above_272k_tokens": 7.5e-06, + "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": false, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, "openrouter/openai/gpt-oss-20b:batch": { "input_cost_per_token": 2.4e-08, "litellm_provider": "openrouter", diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 4257649ab41..75241031730 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -3917,6 +3917,55 @@ "supports_web_search": true, "supports_xhigh_reasoning_effort": true }, + "azure_ai/gpt-6.1-sol": { + "cache_creation_input_token_cost": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5e-06, + "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_above_272k_tokens": 2e-07, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "litellm_provider": "azure_ai", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1e-05, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://learn.microsoft.com/en-us/azure/ai-foundry/openai/concepts/models", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_cache_breakpoint": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, "azure_ai/gpt-5.5": { "deprecation_date": "2027-10-26", "cache_read_input_token_cost": 5e-07, @@ -8509,6 +8558,102 @@ "supports_web_search": true, "supports_xhigh_reasoning_effort": true }, + "azure/gpt-6.1-sol": { + "cache_creation_input_token_cost": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5e-06, + "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_above_272k_tokens": 2e-07, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "litellm_provider": "azure", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1e-05, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://learn.microsoft.com/en-us/azure/ai-foundry/openai/concepts/models", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, + "azure/gpt-6.1-sol-2026-09-29": { + "cache_creation_input_token_cost": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5e-06, + "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_above_272k_tokens": 2e-07, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "litellm_provider": "azure", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1e-05, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://learn.microsoft.com/en-us/azure/ai-foundry/openai/concepts/models", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, "azure/gpt-chat-latest": { "cache_read_input_token_cost": 5e-07, "deprecation_date": "2026-12-02", @@ -77430,6 +77575,106 @@ "supports_vision": true, "supports_web_search": true }, + "openrouter/openai/gpt-6.1-sol": { + "cache_creation_input_token_cost": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5e-06, + "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_above_272k_tokens": 2e-07, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "litellm_provider": "openrouter", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1e-05, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": false, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, + "openrouter/openai/gpt-6.1-sol-pro": { + "cache_creation_input_token_cost": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5e-06, + "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_above_272k_tokens": 2e-07, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "litellm_provider": "openrouter", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1e-05, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": false, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, + "openrouter/openai/gpt-6.1-sol-pro:batch": { + "cache_creation_input_token_cost": 1.25e-06, + "cache_creation_input_token_cost_above_272k_tokens": 2.5e-06, + "cache_read_input_token_cost": 5e-08, + "cache_read_input_token_cost_above_272k_tokens": 1e-07, + "input_cost_per_token": 1e-06, + "input_cost_per_token_above_272k_tokens": 2e-06, + "litellm_provider": "openrouter", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-06, + "output_cost_per_token_above_272k_tokens": 7.5e-06, + "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": false, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, + "openrouter/openai/gpt-6.1-sol:batch": { + "cache_creation_input_token_cost": 1.25e-06, + "cache_creation_input_token_cost_above_272k_tokens": 2.5e-06, + "cache_read_input_token_cost": 5e-08, + "cache_read_input_token_cost_above_272k_tokens": 1e-07, + "input_cost_per_token": 1e-06, + "input_cost_per_token_above_272k_tokens": 2e-06, + "litellm_provider": "openrouter", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-06, + "output_cost_per_token_above_272k_tokens": 7.5e-06, + "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": false, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, "openrouter/openai/gpt-oss-20b:batch": { "input_cost_per_token": 2.4e-08, "litellm_provider": "openrouter", From 5a5e56393829e4e10c7b5eba7d8c63f67fd34d71 Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Tue, 29 Sep 2026 11:11:12 -0700 Subject: [PATCH 22/41] test(proxy): classify every credential-bearing param for the canary suite (#43298) * test(proxy): classify every credential-bearing param for the canary suite * test(proxy): classify gcs_path_service_account as secret, run registry in auth-checks shard, check slot ids at import * test(proxy): use one generic slot id for callback and request-body credential params * test(proxy): name a canary slot only for params an integration test plants * test(proxy): classify the SigNoz callback params * test(proxy): move the slot sync note into the module docstring * test(proxy): model unplanted credential params as their own classification * test(security): classify request-body api_key as unplanted until D1 exists; check registry slots against the harness * test(security): classify request-body api_key under slot D1 * test(security): classify Langfuse and Datadog callback secrets under slots C1 and C3 --- .circleci/scripts/unit_selection.sh | 1 + .../proxy/test_credential_slot_registry.py | 219 ++++++++++++++++++ 2 files changed, 220 insertions(+) create mode 100644 tests/unit/proxy/test_credential_slot_registry.py diff --git a/.circleci/scripts/unit_selection.sh b/.circleci/scripts/unit_selection.sh index 6510b3fd4b5..4cd2b69dc47 100755 --- a/.circleci/scripts/unit_selection.sh +++ b/.circleci/scripts/unit_selection.sh @@ -89,6 +89,7 @@ legacy_paths() { proxy-db-auth-checks) echo tests/unit/proxy/auth/test_auth_checks.py echo tests/unit/proxy/auth/test_user_api_key_auth.py + echo tests/unit/proxy/test_credential_slot_registry.py echo tests/unit/proxy/test_deprecated_key_grace_period.py ;; proxy-db-budgets) echo tests/unit/proxy/auth/test_default_end_user_budget_simple.py diff --git a/tests/unit/proxy/test_credential_slot_registry.py b/tests/unit/proxy/test_credential_slot_registry.py new file mode 100644 index 00000000000..98ddf38c661 --- /dev/null +++ b/tests/unit/proxy/test_credential_slot_registry.py @@ -0,0 +1,219 @@ +"""Every credential-bearing param is classified for the credential canary suite. + +These tests fail until a param is classified below as one of: + +- ``Secret()``: an integration test in ``tests/integration/security`` plants a canary in + exactly this param under that slot id. +- ``Unplanted()``: the param can carry a credential, but no integration test plants a canary + in it yet. This is a classification only. +- ``NotSecret()``: the param cannot carry a credential. + +``CANARY_SLOTS`` mirrors ``SLOTS`` in ``tests/integration/security/_canary.py``, limited to the ids +whose test plants a canary in one of these params. +""" + +import re +from collections.abc import Iterable, Mapping +from dataclasses import dataclass +from pathlib import Path +from types import MappingProxyType +from typing import Final + +from litellm.proxy.auth.auth_utils import is_request_body_safe +from litellm.types.router import LiteLLM_Params, LiteLLMParamsTypedDict +from litellm.types.utils import CustomPricingLiteLLMParams, StandardCallbackDynamicParams + +CANARY_SLOTS: Final[Mapping[str, str]] = MappingProxyType( + { + "B1": "deployment api_key in config.yaml", + "B4": "deployment aws_secret_access_key added through /model/new", + "B4v": "deployment vertex_credentials added through /model/new", + "C1": "team callback langfuse_secret / langfuse_secret_key", + "C3": "team callback dd_api_key for the Datadog sink", + "D1": "client-side api_key in the request body", + } +) + +THIS_FILE: Final = "tests/unit/proxy/test_credential_slot_registry.py" + +HARNESS_FILE: Final = Path(__file__).resolve().parents[2] / "integration" / "security" / "_canary.py" + +CREDENTIAL_NAME: Final = re.compile(r"(?:^|_)(?:key|secret|token|password|credential)") +"""Matches a name segment that starts with a credential word. Anchoring on a segment start keeps +``valkey_host`` and the other ``valkey_*`` settings out, and still matches ``aws_access_key_id``.""" + +PRICING_FIELDS: Final = frozenset(CustomPricingLiteLLMParams.model_fields) +"""Excluded from the name match: in ``input_cost_per_token`` and friends, token is a billing unit.""" + + +@dataclass(frozen=True) +class Secret: + slot: str + + def __post_init__(self) -> None: + if self.slot not in CANARY_SLOTS: + raise ValueError(f"Secret({self.slot!r}) names no slot in CANARY_SLOTS") + + +@dataclass(frozen=True) +class Unplanted: + pass + + +@dataclass(frozen=True) +class NotSecret: + reason: str + + +Classification = Secret | Unplanted | NotSecret + +CALLBACK_PARAM_CLASSIFICATION: Final[Mapping[str, Classification]] = MappingProxyType( + { + "langfuse_public_key": NotSecret("public half of the Langfuse key pair, an identifier"), + "langfuse_secret": Secret("C1"), + "langfuse_secret_key": Secret("C1"), + "langfuse_host": NotSecret("sink endpoint URL"), + "langfuse_environment": NotSecret("environment label"), + "langfuse_span_scope": NotSecret("span scope setting"), + "langfuse_prompt_version": NotSecret("prompt version number"), + "gcs_bucket_name": NotSecret("bucket name"), + "gcs_path_service_account": Unplanted(), + "langsmith_api_key": Unplanted(), + "langsmith_project": NotSecret("project name"), + "langsmith_base_url": NotSecret("sink endpoint URL"), + "langsmith_sampling_rate": NotSecret("sampling rate"), + "langsmith_tenant_id": NotSecret("tenant identifier"), + "humanloop_api_key": Unplanted(), + "arize_api_key": Unplanted(), + "arize_space_key": Unplanted(), + "arize_space_id": NotSecret("space identifier"), + "arize_success_sampling_rate": NotSecret("sampling rate"), + "arize_error_sampling_rate": NotSecret("sampling rate"), + "posthog_api_key": Unplanted(), + "posthog_api_url": NotSecret("sink endpoint URL"), + "wandb_api_key": Unplanted(), + "weave_project_id": NotSecret("project identifier"), + "dd_api_key": Secret("C3"), + "dd_site": NotSecret("sink site name"), + "dd_agent_host": NotSecret("agent host name"), + "dd_agent_port": NotSecret("agent port"), + "newrelic_api_key": Unplanted(), + "newrelic_region": NotSecret("region name"), + "signoz_ingestion_key": Unplanted(), + "signoz_ingestion_endpoint": NotSecret("sink endpoint URL"), + "turn_off_message_logging": NotSecret("boolean logging switch"), + "litellm_disabled_callbacks": NotSecret("list of callback names"), + } +) + +DEPLOYMENT_PARAM_CLASSIFICATION: Final[Mapping[str, Classification]] = MappingProxyType( + { + "api_key": Secret("B1"), + "azure_ad_token": Unplanted(), + "client_secret": Unplanted(), + "azure_password": Unplanted(), + "vertex_credentials": Secret("B4v"), + "aws_access_key_id": Unplanted(), + "aws_secret_access_key": Secret("B4"), + "aws_session_token": Unplanted(), + "aws_web_identity_token": Unplanted(), + "s3_access_key_id": Unplanted(), + "s3_secret_access_key": Unplanted(), + "s3_encryption_key_id": NotSecret("KMS key identifier, not key material"), + "litellm_credential_name": NotSecret("name of a credentials table entry, not a credential"), + "default_api_key_tpm_limit": NotSecret("rate limit number"), + "default_api_key_rpm_limit": NotSecret("rate limit number"), + "valkey_password": Unplanted(), + } +) + +REQUEST_BODY_PARAM_CLASSIFICATION: Final[Mapping[str, Classification]] = MappingProxyType( + { + "api_key": Secret("D1"), + "aws_access_key_id": Unplanted(), + "aws_secret_access_key": Unplanted(), + "aws_session_token": Unplanted(), + "azure_password": Unplanted(), + "client_secret": Unplanted(), + "s3_access_key_id": Unplanted(), + "s3_secret_access_key": Unplanted(), + "valkey_password": Unplanted(), + "s3_encryption_key_id": NotSecret("KMS key identifier, not key material"), + "litellm_credential_name": NotSecret("name of a credentials table entry, not a credential"), + "default_api_key_tpm_limit": NotSecret("rate limit number"), + "default_api_key_rpm_limit": NotSecret("rate limit number"), + } +) + + +def _credential_named(names: Iterable[str]) -> frozenset[str]: + return frozenset(name for name in names if CREDENTIAL_NAME.search(name)) - PRICING_FIELDS + + +def _deployment_param_names() -> frozenset[str]: + return ( + frozenset(LiteLLM_Params.model_fields) + | LiteLLMParamsTypedDict.__required_keys__ + | LiteLLMParamsTypedDict.__optional_keys__ + ) + + +def _callback_param_names() -> frozenset[str]: + return StandardCallbackDynamicParams.__required_keys__ | StandardCallbackDynamicParams.__optional_keys__ + + +def _accepted_in_request_body(param: str) -> bool: + try: + return is_request_body_safe({"model": "m", param: "v"}, general_settings={}, llm_router=None, model="m") + except ValueError: + return False + + +def _assert_classified( + source: str, names: frozenset[str], mapping: Mapping[str, Classification], mapping_name: str +) -> None: + unclassified: Final = sorted(names - mapping.keys()) + stale: Final = sorted(mapping.keys() - names) + assert not unclassified, ( + f"{source} has params with no credential classification: {unclassified}. " + f"Add each to {mapping_name} in {THIS_FILE} as Secret('') if it can hold a credential " + "and an integration test plants it under a slot in CANARY_SLOTS, as Unplanted() if it can hold a credential " + "but no integration test plants it yet, " + "or as NotSecret('') if it cannot." + ) + assert not stale, f"{mapping_name} in {THIS_FILE} classifies params {source} no longer has: {stale}. Remove them." + + +def test_every_callback_dynamic_param_is_classified(): + _assert_classified( + "StandardCallbackDynamicParams", + _callback_param_names(), + CALLBACK_PARAM_CLASSIFICATION, + "CALLBACK_PARAM_CLASSIFICATION", + ) + + +def test_every_credential_named_deployment_param_is_classified(): + _assert_classified( + "LiteLLM_Params / LiteLLMParamsTypedDict", + _credential_named(_deployment_param_names()), + DEPLOYMENT_PARAM_CLASSIFICATION, + "DEPLOYMENT_PARAM_CLASSIFICATION", + ) + + +def test_every_credential_named_param_a_client_may_send_is_classified(): + candidates: Final = _credential_named(_deployment_param_names() | _callback_param_names()) + _assert_classified( + "is_request_body_safe with default settings", + frozenset(name for name in candidates if _accepted_in_request_body(name)), + REQUEST_BODY_PARAM_CLASSIFICATION, + "REQUEST_BODY_PARAM_CLASSIFICATION", + ) + + +def test_every_canary_slot_exists_in_the_harness(): + harness_slots: Final = frozenset(re.findall(r'^\s+"(\w+)": Slot\(', HARNESS_FILE.read_text(), re.MULTILINE)) + assert harness_slots, f"found no Slot(...) entries in {HARNESS_FILE}" + missing: Final = sorted(CANARY_SLOTS.keys() - harness_slots) + assert not missing, f"CANARY_SLOTS in {THIS_FILE} names slots {HARNESS_FILE.name} does not define: {missing}" From 0fe4028cd962b0401aea89585415e393c514ea94 Mon Sep 17 00:00:00 2001 From: "berriai-litellm-provider-info-sync[bot]" <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 11:14:14 -0700 Subject: [PATCH 23/41] fix(cost-map): lower fireworks up-to-4b size tier to the pricing page price (#43740) Price-Sync: litellm-providers Co-authored-by: berriai-litellm-provider-info-sync[bot] <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> --- litellm/model_prices_and_context_window_backup.json | 5 +++-- model_prices_and_context_window.json | 5 +++-- 2 files changed, 6 insertions(+), 4 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 75241031730..c49c012cc78 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -25469,9 +25469,10 @@ "output_cost_per_token": 5e-07 }, "fireworks-ai-up-to-4b": { - "input_cost_per_token": 2e-07, + "input_cost_per_token": 1e-07, "litellm_provider": "fireworks_ai", - "output_cost_per_token": 2e-07 + "output_cost_per_token": 1e-07, + "source": "https://docs.fireworks.ai/serverless/pricing" }, "fireworks_ai/WhereIsAI/UAE-Large-V1": { "input_cost_per_token": 1.6e-08, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 75241031730..c49c012cc78 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -25469,9 +25469,10 @@ "output_cost_per_token": 5e-07 }, "fireworks-ai-up-to-4b": { - "input_cost_per_token": 2e-07, + "input_cost_per_token": 1e-07, "litellm_provider": "fireworks_ai", - "output_cost_per_token": 2e-07 + "output_cost_per_token": 1e-07, + "source": "https://docs.fireworks.ai/serverless/pricing" }, "fireworks_ai/WhereIsAI/UAE-Large-V1": { "input_cost_per_token": 1.6e-08, From abc85c26517c98b0a140737814b803fe45bcf4ff Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 11:27:14 -0700 Subject: [PATCH 24/41] fix(cost_calculator): bill chat per-second pricing once with a new cost_per_second field (#43614) * feat(cost_calculator): add cost_per_second for chat per-second pricing Keep legacy input_cost_per_second and output_cost_per_second as aliases for chat, completion, embedding and responses. When both legacy fields are set, input_cost_per_second wins Move Bedrock commitment rows to cost_per_second so they bill once Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(cost_calculator): drop legacy per-second fields from chat paths Keep Azure chat token pricing generic and update inert Voxtral rates and SageMaker examples Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(cost_calculator): recognize output-only per-second rates Include output_cost_per_second when checking whether a deployment cost entry has pricing so output-only legacy aliases remain attached to the deployment during cost selection Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(pricing): cover cost_per_second and legacy per-second aliases through the proxy Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(cost_calculator): drop output_cost_per_second as a chat per-second alias Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * feat(cost_calculator): restore output_cost_per_second as a chat per-second fallback Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(cost-map): keep input_cost_per_second on bedrock commitment rows for older clients Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: kerry Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- cookbook/misc/config.yaml | 2 +- .../crates/model-catalog/src/model_info.rs | 2 + litellm/cost_calculator.py | 22 +- .../litellm_core_utils/get_litellm_params.py | 2 + litellm/llms/azure/cost_calculation.py | 24 --- litellm/main.py | 18 +- ...odel_prices_and_context_window_backup.json | 74 ++++--- litellm/router.py | 2 +- litellm/types/router.py | 1 + litellm/types/utils.py | 3 + litellm/utils.py | 1 + model_prices_and_context_window.json | 74 ++++--- model_prices_and_context_window.schema.json | 4 + proxy_server_config.yaml | 2 +- .../pricing/test_per_second_pricing.py | 196 ++++++++++++++++++ tests/local_testing/test_embedding.py | 4 +- tests/local_testing/test_sagemaker.py | 14 +- .../proxy/auth/test_auth_checks.py | 2 +- .../proxy/test_pricing_field_strip.py | 1 + .../test_zero_cost_diagnostic.py | 2 +- .../test_response_metadata.py | 8 +- .../test_get_litellm_params.py | 8 +- .../test_litellm_logging.py | 4 +- .../router_strategy/test_complexity_router.py | 8 +- tests/unit/test_cost_calculator.py | 77 ++++++- .../test_register_model_custom_pricing.py | 19 ++ tests/unit/test_utils.py | 2 + tests/unit/types/test_router.py | 1 + ui/litellm-dashboard/src/lib/http/schema.d.ts | 4 + 29 files changed, 441 insertions(+), 140 deletions(-) create mode 100644 tests/integration/pricing/test_per_second_pricing.py diff --git a/cookbook/misc/config.yaml b/cookbook/misc/config.yaml index 27a6332a882..a485bf825fc 100644 --- a/cookbook/misc/config.yaml +++ b/cookbook/misc/config.yaml @@ -24,7 +24,7 @@ model_list: - model_name: sagemaker-completion-model litellm_params: model: sagemaker/berri-benchmarking-Llama-2-70b-chat-hf-4 - input_cost_per_second: 0.000420 + cost_per_second: 0.000420 - model_name: text-embedding-ada-002 litellm_params: model: azure/azure-embedding-model diff --git a/litellm-rust/crates/model-catalog/src/model_info.rs b/litellm-rust/crates/model-catalog/src/model_info.rs index aa543439885..75f16e0c00d 100644 --- a/litellm-rust/crates/model-catalog/src/model_info.rs +++ b/litellm-rust/crates/model-catalog/src/model_info.rs @@ -125,6 +125,8 @@ pub struct ModelInfo { pub computer_use_input_cost_per_1k_tokens: Option, #[serde(skip_serializing_if = "Option::is_none")] pub computer_use_output_cost_per_1k_tokens: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub cost_per_second: Option, /// Reasoning effort the provider applies when the request omits reasoning_effort. Gates whether a non-default temperature or the top_p/logprobs sampling params are accepted, which hold only when the effort resolves to 'none'. #[serde(skip_serializing_if = "Option::is_none")] pub default_reasoning_effort: Option, diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index a279b9f0903..2f76b84f5e3 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -351,19 +351,27 @@ def _per_second_pricing_cost( return None if _has_token_or_tiered_pricing(model_info) or not _bills_wall_clock_seconds(model_info): return None + cost_per_second: Final = model_info.get("cost_per_second") input_cost_per_second: Final = model_info.get("input_cost_per_second") output_cost_per_second: Final = model_info.get("output_cost_per_second") - if input_cost_per_second is None and output_cost_per_second is None: + resolved_cost_per_second: Final = ( + cost_per_second + if cost_per_second is not None + else input_cost_per_second + if input_cost_per_second is not None + else output_cost_per_second + ) + if resolved_cost_per_second is None: return None + seconds: Final = (response_time_ms or 0.0) / 1000 verbose_logger.debug( - "For model=%s - input_cost_per_second: %s; output_cost_per_second: %s; response time: %s", + "For model=%s - cost_per_second: %s; response time: %s", model, - input_cost_per_second, - output_cost_per_second, + resolved_cost_per_second, response_time_ms, ) - return (input_cost_per_second or 0.0) * seconds, (output_cost_per_second or 0.0) * seconds + return resolved_cost_per_second * seconds, 0.0 def cost_per_token( @@ -790,7 +798,9 @@ def _get_hidden_str_for_cost_calc(hidden_params: object, key: str) -> str | None return value if isinstance(value, str) and value else None -_NON_TOKEN_RATE_FIELDS: Final = frozenset({"input_cost_per_second", "input_cost_per_query", "tiered_pricing"}) +_NON_TOKEN_RATE_FIELDS: Final = frozenset( + {"cost_per_second", "input_cost_per_second", "output_cost_per_second", "input_cost_per_query", "tiered_pricing"} +) def _cost_map_entry_prices_anything(entry: Mapping[str, object]) -> bool: diff --git a/litellm/litellm_core_utils/get_litellm_params.py b/litellm/litellm_core_utils/get_litellm_params.py index f28259a1b7f..b8441d2bc6d 100644 --- a/litellm/litellm_core_utils/get_litellm_params.py +++ b/litellm/litellm_core_utils/get_litellm_params.py @@ -155,6 +155,7 @@ def get_litellm_params( allm_passthrough_route=None, preset_cache_key=None, no_log=None, + cost_per_second: float | None = None, input_cost_per_second=None, input_cost_per_token=None, output_cost_per_token=None, @@ -216,6 +217,7 @@ def get_litellm_params( "preset_cache_key": preset_cache_key, "no-log": no_log or kwargs.get("no-log"), "stream_response": {}, # litellm_call_id: ModelResponse Dict + "cost_per_second": cost_per_second, "input_cost_per_token": input_cost_per_token, "input_cost_per_second": input_cost_per_second, "output_cost_per_token": output_cost_per_token, diff --git a/litellm/llms/azure/cost_calculation.py b/litellm/llms/azure/cost_calculation.py index 8dc809507d5..057e9dbb9d9 100644 --- a/litellm/llms/azure/cost_calculation.py +++ b/litellm/llms/azure/cost_calculation.py @@ -3,12 +3,8 @@ Helper util for handling azure openai-specific cost calculation - e.g.: prompt caching, audio tokens """ -from typing import Final - -from litellm._logging import verbose_logger from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token from litellm.types.utils import Usage -from litellm.utils import get_model_info def cost_per_token( @@ -27,26 +23,6 @@ def cost_per_token( Returns: Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd """ - ## GET MODEL INFO - model_info: Final = get_model_info(model=model, custom_llm_provider="azure") - - ## Speech / Audio cost calculation (cost per second for TTS models) - if ( - "output_cost_per_second" in model_info - and model_info["output_cost_per_second"] is not None - and response_time_ms is not None - ): - verbose_logger.debug( - "For model=%s - output_cost_per_second: %s; response time: %s", - model, - model_info.get("output_cost_per_second"), - response_time_ms, - ) - ## COST PER SECOND ## - prompt_cost: Final = 0.0 - completion_cost: Final = model_info["output_cost_per_second"] * response_time_ms / 1000 - return prompt_cost, completion_cost - ## Use generic cost calculator for all other cases ## This properly handles: text tokens, audio tokens, cached tokens, reasoning tokens, etc. return generic_cost_per_token( diff --git a/litellm/main.py b/litellm/main.py index 6c85adf3ae8..769eac79488 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -5353,6 +5353,7 @@ def completion( ### CUSTOM MODEL COST ### input_cost_per_token: Final = kwargs.get("input_cost_per_token", None) output_cost_per_token: Final = kwargs.get("output_cost_per_token", None) + cost_per_second: Final = kwargs.get("cost_per_second", None) input_cost_per_second: Final = kwargs.get("input_cost_per_second", None) output_cost_per_second: Final = kwargs.get("output_cost_per_second", None) ### CUSTOM PROMPT TEMPLATE ### @@ -5514,8 +5515,11 @@ def completion( ### REGISTER CUSTOM MODEL PRICING -- IF GIVEN ### if ( - input_cost_per_token is not None and output_cost_per_token is not None - ) or input_cost_per_second is not None: + (input_cost_per_token is not None and output_cost_per_token is not None) + or input_cost_per_second is not None + or output_cost_per_second is not None + or cost_per_second is not None + ): _register_custom_pricing_for_request( model=model, custom_llm_provider=custom_llm_provider, @@ -5657,6 +5661,7 @@ def completion( proxy_server_request=proxy_server_request, preset_cache_key=preset_cache_key, no_log=no_log, + cost_per_second=cost_per_second, input_cost_per_second=input_cost_per_second, input_cost_per_token=input_cost_per_token, output_cost_per_second=output_cost_per_second, @@ -6354,7 +6359,9 @@ def embedding( ### CUSTOM MODEL COST ### input_cost_per_token: Final = kwargs.get("input_cost_per_token", None) output_cost_per_token: Final = kwargs.get("output_cost_per_token", None) + cost_per_second: Final = kwargs.get("cost_per_second", None) input_cost_per_second: Final = kwargs.get("input_cost_per_second", None) + output_cost_per_second: Final = kwargs.get("output_cost_per_second", None) openai_params: Final = [ "user", "dimensions", @@ -6395,7 +6402,12 @@ def embedding( ) ### REGISTER CUSTOM MODEL PRICING -- IF GIVEN ### - if (input_cost_per_token is not None and output_cost_per_token is not None) or input_cost_per_second is not None: + if ( + (input_cost_per_token is not None and output_cost_per_token is not None) + or input_cost_per_second is not None + or output_cost_per_second is not None + or cost_per_second is not None + ): _register_custom_pricing_for_request( model=model, custom_llm_provider=custom_llm_provider, diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index c49c012cc78..9ae270243d2 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -12682,43 +12682,43 @@ "source": "https://developers.openai.com/api/docs/pricing" }, "bedrock/*/1-month-commitment/cohere.command-light-text-v14": { + "cost_per_second": 0.001902, "input_cost_per_second": 0.001902, "litellm_provider": "bedrock", "max_input_tokens": 4096, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", - "output_cost_per_second": 0.001902, "supports_tool_choice": true }, "bedrock/*/1-month-commitment/cohere.command-text-v14": { + "cost_per_second": 0.011, "input_cost_per_second": 0.011, "litellm_provider": "bedrock", "max_input_tokens": 4096, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", - "output_cost_per_second": 0.011, "supports_tool_choice": true }, "bedrock/*/6-month-commitment/cohere.command-light-text-v14": { + "cost_per_second": 0.0011416, "input_cost_per_second": 0.0011416, "litellm_provider": "bedrock", "max_input_tokens": 4096, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", - "output_cost_per_second": 0.0011416, "supports_tool_choice": true }, "bedrock/*/6-month-commitment/cohere.command-text-v14": { + "cost_per_second": 0.0066027, "input_cost_per_second": 0.0066027, "litellm_provider": "bedrock", "max_input_tokens": 4096, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", - "output_cost_per_second": 0.0066027, "supports_tool_choice": true }, "bedrock/guardrails": { @@ -12737,61 +12737,61 @@ "source": "https://aws.amazon.com/bedrock/pricing/" }, "bedrock/ap-northeast-1/1-month-commitment/anthropic.claude-instant-v1": { + "cost_per_second": 0.01475, "input_cost_per_second": 0.01475, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.01475, "supports_tool_choice": true }, "bedrock/ap-northeast-1/1-month-commitment/anthropic.claude-v1": { + "cost_per_second": 0.0455, "input_cost_per_second": 0.0455, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, - "mode": "chat", - "output_cost_per_second": 0.0455 + "mode": "chat" }, "bedrock/ap-northeast-1/1-month-commitment/anthropic.claude-v2:1": { + "cost_per_second": 0.0455, "input_cost_per_second": 0.0455, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.0455, "supports_tool_choice": true }, "bedrock/ap-northeast-1/6-month-commitment/anthropic.claude-instant-v1": { + "cost_per_second": 0.008194, "input_cost_per_second": 0.008194, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.008194, "supports_tool_choice": true }, "bedrock/ap-northeast-1/6-month-commitment/anthropic.claude-v1": { + "cost_per_second": 0.02527, "input_cost_per_second": 0.02527, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, - "mode": "chat", - "output_cost_per_second": 0.02527 + "mode": "chat" }, "bedrock/ap-northeast-1/6-month-commitment/anthropic.claude-v2:1": { + "cost_per_second": 0.02527, "input_cost_per_second": 0.02527, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.02527, "supports_tool_choice": true }, "bedrock/ap-northeast-1/anthropic.claude-instant-v1": { @@ -13241,61 +13241,61 @@ "source": "https://aws.amazon.com/bedrock/pricing/" }, "bedrock/eu-central-1/1-month-commitment/anthropic.claude-instant-v1": { + "cost_per_second": 0.01635, "input_cost_per_second": 0.01635, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.01635, "supports_tool_choice": true }, "bedrock/eu-central-1/1-month-commitment/anthropic.claude-v1": { + "cost_per_second": 0.0415, "input_cost_per_second": 0.0415, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, - "mode": "chat", - "output_cost_per_second": 0.0415 + "mode": "chat" }, "bedrock/eu-central-1/1-month-commitment/anthropic.claude-v2:1": { + "cost_per_second": 0.0415, "input_cost_per_second": 0.0415, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.0415, "supports_tool_choice": true }, "bedrock/eu-central-1/6-month-commitment/anthropic.claude-instant-v1": { + "cost_per_second": 0.009083, "input_cost_per_second": 0.009083, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.009083, "supports_tool_choice": true }, "bedrock/eu-central-1/6-month-commitment/anthropic.claude-v1": { + "cost_per_second": 0.02305, "input_cost_per_second": 0.02305, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, - "mode": "chat", - "output_cost_per_second": 0.02305 + "mode": "chat" }, "bedrock/eu-central-1/6-month-commitment/anthropic.claude-v2:1": { + "cost_per_second": 0.02305, "input_cost_per_second": 0.02305, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.02305, "supports_tool_choice": true }, "bedrock/eu-central-1/anthropic.claude-instant-v1": { @@ -13737,61 +13737,61 @@ "source": "https://aws.amazon.com/bedrock/pricing/" }, "bedrock/us-east-1/1-month-commitment/anthropic.claude-instant-v1": { + "cost_per_second": 0.011, "input_cost_per_second": 0.011, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.011, "supports_tool_choice": true }, "bedrock/us-east-1/1-month-commitment/anthropic.claude-v1": { + "cost_per_second": 0.0175, "input_cost_per_second": 0.0175, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, - "mode": "chat", - "output_cost_per_second": 0.0175 + "mode": "chat" }, "bedrock/us-east-1/1-month-commitment/anthropic.claude-v2:1": { + "cost_per_second": 0.0175, "input_cost_per_second": 0.0175, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.0175, "supports_tool_choice": true }, "bedrock/us-east-1/6-month-commitment/anthropic.claude-instant-v1": { + "cost_per_second": 0.00611, "input_cost_per_second": 0.00611, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.00611, "supports_tool_choice": true }, "bedrock/us-east-1/6-month-commitment/anthropic.claude-v1": { + "cost_per_second": 0.00972, "input_cost_per_second": 0.00972, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, - "mode": "chat", - "output_cost_per_second": 0.00972 + "mode": "chat" }, "bedrock/us-east-1/6-month-commitment/anthropic.claude-v2:1": { + "cost_per_second": 0.00972, "input_cost_per_second": 0.00972, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.00972, "supports_tool_choice": true }, "bedrock/us-east-1/anthropic.claude-instant-v1": { @@ -14385,61 +14385,61 @@ "output_cost_per_token": 6e-07 }, "bedrock/us-west-2/1-month-commitment/anthropic.claude-instant-v1": { + "cost_per_second": 0.011, "input_cost_per_second": 0.011, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.011, "supports_tool_choice": true }, "bedrock/us-west-2/1-month-commitment/anthropic.claude-v1": { + "cost_per_second": 0.0175, "input_cost_per_second": 0.0175, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, - "mode": "chat", - "output_cost_per_second": 0.0175 + "mode": "chat" }, "bedrock/us-west-2/1-month-commitment/anthropic.claude-v2:1": { + "cost_per_second": 0.0175, "input_cost_per_second": 0.0175, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.0175, "supports_tool_choice": true }, "bedrock/us-west-2/6-month-commitment/anthropic.claude-instant-v1": { + "cost_per_second": 0.00611, "input_cost_per_second": 0.00611, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.00611, "supports_tool_choice": true }, "bedrock/us-west-2/6-month-commitment/anthropic.claude-v1": { + "cost_per_second": 0.00972, "input_cost_per_second": 0.00972, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, - "mode": "chat", - "output_cost_per_second": 0.00972 + "mode": "chat" }, "bedrock/us-west-2/6-month-commitment/anthropic.claude-v2:1": { + "cost_per_second": 0.00972, "input_cost_per_second": 0.00972, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.00972, "supports_tool_choice": true }, "bedrock/us-west-2/anthropic.claude-instant-v1": { @@ -38527,7 +38527,6 @@ }, "mistral/voxtral-small-2507": { "cache_read_input_token_cost": 1e-08, - "input_cost_per_second": 6.666666666666667e-05, "input_cost_per_token": 1e-07, "litellm_provider": "mistral", "max_input_tokens": 32768, @@ -38543,7 +38542,6 @@ }, "mistral/voxtral-small-latest": { "cache_read_input_token_cost": 1e-08, - "input_cost_per_second": 6.666666666666667e-05, "input_cost_per_token": 1e-07, "litellm_provider": "mistral", "max_input_tokens": 32768, diff --git a/litellm/router.py b/litellm/router.py index 86ba5112435..842cd9de378 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -8801,7 +8801,7 @@ class Router: return if any( model_info.get(field) is not None - for field in ("input_cost_per_token", "input_cost_per_second", "tiered_pricing") + for field in ("input_cost_per_token", "input_cost_per_second", "cost_per_second", "tiered_pricing") ): return try: diff --git a/litellm/types/router.py b/litellm/types/router.py index d545f7ae639..2ab1a1185ed 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -606,6 +606,7 @@ class LiteLLMParamsTypedDict(TypedDict, total=False): ## CUSTOM PRICING ## input_cost_per_token: float | None output_cost_per_token: float | None + cost_per_second: ReadOnly[float | None] input_cost_per_second: float | None output_cost_per_second: float | None output_cost_per_second_480p: ReadOnly[float | None] diff --git a/litellm/types/utils.py b/litellm/types/utils.py index dba15bc99a5..e32e3b74ec6 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -329,6 +329,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False): input_cost_per_video_per_second: float | None # only for vertex ai models input_cost_per_audio_token_batches: ReadOnly[float | None] input_cost_per_image_token_batches: ReadOnly[float | None] + cost_per_second: ReadOnly[float | None] input_cost_per_second: float | None # for OpenAI Speech models input_cost_per_token_batches: float | None input_cost_per_video_token_batches: ReadOnly[float | None] @@ -2784,6 +2785,7 @@ class LoggedLiteLLMParams(TypedDict, total=False): acompletion: bool | None preset_cache_key: str | None no_log: bool | None + cost_per_second: ReadOnly[float | None] input_cost_per_second: float | None input_cost_per_token: float | None output_cost_per_token: float | None @@ -3709,6 +3711,7 @@ class MirroredPricingParams(BaseModel): class CustomPricingLiteLLMParams(MirroredPricingParams): ## CUSTOM PRICING ## + cost_per_second: float | None = None input_cost_per_second: float | None = None output_cost_per_second: float | None = None output_cost_per_second_1080p: float | None = None diff --git a/litellm/utils.py b/litellm/utils.py index 09b5067339d..ffd507fad45 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -6168,6 +6168,7 @@ def _get_model_info_helper( ), input_cost_per_token_above_512k_tokens=_model_info.get("input_cost_per_token_above_512k_tokens", None), input_cost_per_query=_model_info.get("input_cost_per_query", None), + cost_per_second=_model_info.get("cost_per_second", None), input_cost_per_second=_model_info.get("input_cost_per_second", None), input_cost_per_audio_token=_model_info.get("input_cost_per_audio_token", None), input_cost_per_image_token=_model_info.get("input_cost_per_image_token", None), diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index c49c012cc78..9ae270243d2 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -12682,43 +12682,43 @@ "source": "https://developers.openai.com/api/docs/pricing" }, "bedrock/*/1-month-commitment/cohere.command-light-text-v14": { + "cost_per_second": 0.001902, "input_cost_per_second": 0.001902, "litellm_provider": "bedrock", "max_input_tokens": 4096, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", - "output_cost_per_second": 0.001902, "supports_tool_choice": true }, "bedrock/*/1-month-commitment/cohere.command-text-v14": { + "cost_per_second": 0.011, "input_cost_per_second": 0.011, "litellm_provider": "bedrock", "max_input_tokens": 4096, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", - "output_cost_per_second": 0.011, "supports_tool_choice": true }, "bedrock/*/6-month-commitment/cohere.command-light-text-v14": { + "cost_per_second": 0.0011416, "input_cost_per_second": 0.0011416, "litellm_provider": "bedrock", "max_input_tokens": 4096, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", - "output_cost_per_second": 0.0011416, "supports_tool_choice": true }, "bedrock/*/6-month-commitment/cohere.command-text-v14": { + "cost_per_second": 0.0066027, "input_cost_per_second": 0.0066027, "litellm_provider": "bedrock", "max_input_tokens": 4096, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", - "output_cost_per_second": 0.0066027, "supports_tool_choice": true }, "bedrock/guardrails": { @@ -12737,61 +12737,61 @@ "source": "https://aws.amazon.com/bedrock/pricing/" }, "bedrock/ap-northeast-1/1-month-commitment/anthropic.claude-instant-v1": { + "cost_per_second": 0.01475, "input_cost_per_second": 0.01475, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.01475, "supports_tool_choice": true }, "bedrock/ap-northeast-1/1-month-commitment/anthropic.claude-v1": { + "cost_per_second": 0.0455, "input_cost_per_second": 0.0455, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, - "mode": "chat", - "output_cost_per_second": 0.0455 + "mode": "chat" }, "bedrock/ap-northeast-1/1-month-commitment/anthropic.claude-v2:1": { + "cost_per_second": 0.0455, "input_cost_per_second": 0.0455, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.0455, "supports_tool_choice": true }, "bedrock/ap-northeast-1/6-month-commitment/anthropic.claude-instant-v1": { + "cost_per_second": 0.008194, "input_cost_per_second": 0.008194, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.008194, "supports_tool_choice": true }, "bedrock/ap-northeast-1/6-month-commitment/anthropic.claude-v1": { + "cost_per_second": 0.02527, "input_cost_per_second": 0.02527, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, - "mode": "chat", - "output_cost_per_second": 0.02527 + "mode": "chat" }, "bedrock/ap-northeast-1/6-month-commitment/anthropic.claude-v2:1": { + "cost_per_second": 0.02527, "input_cost_per_second": 0.02527, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.02527, "supports_tool_choice": true }, "bedrock/ap-northeast-1/anthropic.claude-instant-v1": { @@ -13241,61 +13241,61 @@ "source": "https://aws.amazon.com/bedrock/pricing/" }, "bedrock/eu-central-1/1-month-commitment/anthropic.claude-instant-v1": { + "cost_per_second": 0.01635, "input_cost_per_second": 0.01635, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.01635, "supports_tool_choice": true }, "bedrock/eu-central-1/1-month-commitment/anthropic.claude-v1": { + "cost_per_second": 0.0415, "input_cost_per_second": 0.0415, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, - "mode": "chat", - "output_cost_per_second": 0.0415 + "mode": "chat" }, "bedrock/eu-central-1/1-month-commitment/anthropic.claude-v2:1": { + "cost_per_second": 0.0415, "input_cost_per_second": 0.0415, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.0415, "supports_tool_choice": true }, "bedrock/eu-central-1/6-month-commitment/anthropic.claude-instant-v1": { + "cost_per_second": 0.009083, "input_cost_per_second": 0.009083, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.009083, "supports_tool_choice": true }, "bedrock/eu-central-1/6-month-commitment/anthropic.claude-v1": { + "cost_per_second": 0.02305, "input_cost_per_second": 0.02305, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, - "mode": "chat", - "output_cost_per_second": 0.02305 + "mode": "chat" }, "bedrock/eu-central-1/6-month-commitment/anthropic.claude-v2:1": { + "cost_per_second": 0.02305, "input_cost_per_second": 0.02305, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.02305, "supports_tool_choice": true }, "bedrock/eu-central-1/anthropic.claude-instant-v1": { @@ -13737,61 +13737,61 @@ "source": "https://aws.amazon.com/bedrock/pricing/" }, "bedrock/us-east-1/1-month-commitment/anthropic.claude-instant-v1": { + "cost_per_second": 0.011, "input_cost_per_second": 0.011, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.011, "supports_tool_choice": true }, "bedrock/us-east-1/1-month-commitment/anthropic.claude-v1": { + "cost_per_second": 0.0175, "input_cost_per_second": 0.0175, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, - "mode": "chat", - "output_cost_per_second": 0.0175 + "mode": "chat" }, "bedrock/us-east-1/1-month-commitment/anthropic.claude-v2:1": { + "cost_per_second": 0.0175, "input_cost_per_second": 0.0175, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.0175, "supports_tool_choice": true }, "bedrock/us-east-1/6-month-commitment/anthropic.claude-instant-v1": { + "cost_per_second": 0.00611, "input_cost_per_second": 0.00611, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.00611, "supports_tool_choice": true }, "bedrock/us-east-1/6-month-commitment/anthropic.claude-v1": { + "cost_per_second": 0.00972, "input_cost_per_second": 0.00972, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, - "mode": "chat", - "output_cost_per_second": 0.00972 + "mode": "chat" }, "bedrock/us-east-1/6-month-commitment/anthropic.claude-v2:1": { + "cost_per_second": 0.00972, "input_cost_per_second": 0.00972, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.00972, "supports_tool_choice": true }, "bedrock/us-east-1/anthropic.claude-instant-v1": { @@ -14385,61 +14385,61 @@ "output_cost_per_token": 6e-07 }, "bedrock/us-west-2/1-month-commitment/anthropic.claude-instant-v1": { + "cost_per_second": 0.011, "input_cost_per_second": 0.011, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.011, "supports_tool_choice": true }, "bedrock/us-west-2/1-month-commitment/anthropic.claude-v1": { + "cost_per_second": 0.0175, "input_cost_per_second": 0.0175, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, - "mode": "chat", - "output_cost_per_second": 0.0175 + "mode": "chat" }, "bedrock/us-west-2/1-month-commitment/anthropic.claude-v2:1": { + "cost_per_second": 0.0175, "input_cost_per_second": 0.0175, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.0175, "supports_tool_choice": true }, "bedrock/us-west-2/6-month-commitment/anthropic.claude-instant-v1": { + "cost_per_second": 0.00611, "input_cost_per_second": 0.00611, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.00611, "supports_tool_choice": true }, "bedrock/us-west-2/6-month-commitment/anthropic.claude-v1": { + "cost_per_second": 0.00972, "input_cost_per_second": 0.00972, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, - "mode": "chat", - "output_cost_per_second": 0.00972 + "mode": "chat" }, "bedrock/us-west-2/6-month-commitment/anthropic.claude-v2:1": { + "cost_per_second": 0.00972, "input_cost_per_second": 0.00972, "litellm_provider": "bedrock", "max_input_tokens": 100000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_second": 0.00972, "supports_tool_choice": true }, "bedrock/us-west-2/anthropic.claude-instant-v1": { @@ -38527,7 +38527,6 @@ }, "mistral/voxtral-small-2507": { "cache_read_input_token_cost": 1e-08, - "input_cost_per_second": 6.666666666666667e-05, "input_cost_per_token": 1e-07, "litellm_provider": "mistral", "max_input_tokens": 32768, @@ -38543,7 +38542,6 @@ }, "mistral/voxtral-small-latest": { "cache_read_input_token_cost": 1e-08, - "input_cost_per_second": 6.666666666666667e-05, "input_cost_per_token": 1e-07, "litellm_provider": "mistral", "max_input_tokens": 32768, diff --git a/model_prices_and_context_window.schema.json b/model_prices_and_context_window.schema.json index e893b6265fa..c4b424a26b4 100644 --- a/model_prices_and_context_window.schema.json +++ b/model_prices_and_context_window.schema.json @@ -249,6 +249,10 @@ "comment": { "type": "string" }, + "cost_per_second": { + "type": "number", + "minimum": 0 + }, "default_reasoning_effort": { "type": "string", "description": "Reasoning effort the provider applies when the request omits reasoning_effort. Gates whether a non-default temperature or the top_p/logprobs sampling params are accepted, which hold only when the effort resolves to 'none'.", diff --git a/proxy_server_config.yaml b/proxy_server_config.yaml index 24e26ea8e22..be6dd20647d 100644 --- a/proxy_server_config.yaml +++ b/proxy_server_config.yaml @@ -31,7 +31,7 @@ model_list: - model_name: sagemaker-completion-model litellm_params: model: sagemaker/berri-benchmarking-Llama-2-70b-chat-hf-4 - input_cost_per_second: 0.000420 + cost_per_second: 0.000420 - model_name: text-embedding-ada-002 litellm_params: model: openai/text-embedding-3-small diff --git a/tests/integration/pricing/test_per_second_pricing.py b/tests/integration/pricing/test_per_second_pricing.py new file mode 100644 index 00000000000..ad44a631054 --- /dev/null +++ b/tests/integration/pricing/test_per_second_pricing.py @@ -0,0 +1,196 @@ +import json +import uuid +from collections.abc import Mapping +from typing import Final + +import httpx +import pytest +from pydantic import JsonValue + +from tests.integration._support.client import JSON_OBJECT, Gateway, eventually, object_value, string_value +from tests.integration._support.database import read_rows +from tests.integration._support.upstream import delete_scenario, register_scenario +from tests.integration.cost_calculation.cost_tracking_case import SseResponse + +RATE: Final = 0.5 +FRAME_DELAY_MS: Final = 300 +CONTENT: Final = ("one", " two", " three", " four") +PRICING_FIELDS: Final = frozenset({"cost_per_second", "input_cost_per_second", "output_cost_per_second"}) +PER_SECOND_CONFIGURATIONS: Final[tuple[tuple[str, Mapping[str, JsonValue]], ...]] = ( + ("new_field", {"cost_per_second": RATE}), + ("legacy_input", {"input_cost_per_second": RATE}), + ("legacy_output", {"output_cost_per_second": RATE}), + ("legacy_both", {"input_cost_per_second": RATE, "output_cost_per_second": 0.25}), + ( + "all_three", + {"cost_per_second": RATE, "input_cost_per_second": 0.25, "output_cost_per_second": 0.125}, + ), +) + + +def _sse_chunk(delta: dict[str, JsonValue], finish_reason: str | None) -> str: + payload: Final = { + "id": "$REQUEST_ID", + "object": "chat.completion.chunk", + "created": 1, + "model": "integration-per-second", + "choices": [{"index": 0, "delta": delta, "finish_reason": finish_reason}], + } + return f"data: {json.dumps(payload)}" + + +def _sse_frames() -> tuple[str, ...]: + content_frames: Final = tuple(_sse_chunk({"content": content}, None) for content in CONTENT) + usage_payload: Final = { + "id": "$REQUEST_ID", + "object": "chat.completion.chunk", + "created": 1, + "model": "integration-per-second", + "choices": [], + "usage": {"prompt_tokens": 20, "completion_tokens": 20, "total_tokens": 40}, + } + usage_frame: Final = f"data: {json.dumps(usage_payload)}" + return (*content_frames, _sse_chunk({}, "stop"), usage_frame, "data: [DONE]") + + +def _stream_content(event: dict[str, JsonValue]) -> str: + choices: Final = event.get("choices") + if not isinstance(choices, list) or not choices: + return "" + delta: Final = object_value(object_value(choices[0])["delta"]) + content: Final = delta.get("content") + return content if isinstance(content, str) else "" + + +def _clear_observations(upstream: httpx.Client) -> None: + response: Final = upstream.get("/__observations") + assert response.status_code == 200, response.text + + +def _observed_request_body(upstream: httpx.Client) -> dict[str, JsonValue]: + observations: Final = JSON_OBJECT.validate_json(upstream.get("/__observations").content)["requests"] + assert isinstance(observations, list) + assert len(observations) == 1 + return object_value(object_value(observations[0])["body"]) + + +@pytest.mark.parametrize( + ("pricing_case", "pricing"), + PER_SECOND_CONFIGURATIONS, + ids=("new_field", "legacy_input", "legacy_output", "legacy_both", "all_three"), +) +def test_chat_per_second_pricing_is_charged_once_and_not_forwarded( + gateway: Gateway, pricing_case: str, pricing: Mapping[str, JsonValue] +) -> None: + with gateway.scenario() as scenario: + scenario_id: Final = f"per-second-{pricing_case}-{uuid.uuid4().hex}" + key: Final = scenario.key() + model: Final = scenario.model( + model=f"openai/integration-per-second-{uuid.uuid4().hex}", + api_key=scenario_id, + api_base=f"{gateway.upstream_url.rstrip('/')}/v1", + **pricing, + ) + with httpx.Client(base_url=gateway.upstream_url, trust_env=False) as upstream: + _clear_observations(upstream) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "price this request"}]}, + key=key, + ) + body: Final = _observed_request_body(upstream) + assert response.status_code == 200, f"{pricing_case}: {response.text}" + response_cost: Final = float(response.headers.get("x-litellm-response-cost", "0")) + duration_ms: Final = float(response.headers.get("x-litellm-response-duration-ms", "0")) + assert response_cost > 0, f"{pricing_case}: cost={response_cost}, duration_ms={duration_ms}, body={body}" + assert response_cost == pytest.approx(RATE * duration_ms / 1000, rel=1e-3), ( + f"{pricing_case}: cost={response_cost}, duration_ms={duration_ms}, body={body}" + ) + assert not PRICING_FIELDS.intersection(body), body + + request_id: Final = string_value(object_value(response.json())["id"]) + rows: Final = eventually( + lambda: read_rows( + 'SELECT spend FROM "LiteLLM_SpendLogs" WHERE request_id = %s', + (request_id,), + ), + lambda values: len(values) == 1, + seconds=70, + ) + assert float(str(rows[0]["spend"])) == pytest.approx(response_cost, rel=1e-3) + + +@pytest.mark.parametrize( + ("pricing_case", "pricing"), + PER_SECOND_CONFIGURATIONS, + ids=("new_field", "legacy_input", "legacy_output", "legacy_both", "all_three"), +) +def test_streaming_chat_per_second_pricing_covers_the_full_stream( + gateway: Gateway, pricing_case: str, pricing: Mapping[str, JsonValue] +) -> None: + with gateway.scenario() as scenario: + scenario_id: Final = f"per-second-stream-{pricing_case}-{uuid.uuid4().hex}" + frames: Final = _sse_frames() + handle: Final = register_scenario( + scenario_id, + SseResponse(content_type="text/event-stream", frames=frames, frame_delay_ms=FRAME_DELAY_MS), + ) + scenario.cleanups.callback(delete_scenario, handle) + key: Final = scenario.key() + model: Final = scenario.model( + model=f"openai/integration-per-second-{uuid.uuid4().hex}", + api_key=scenario_id, + api_base=handle.api_base(), + **pricing, + ) + with httpx.Client(base_url=gateway.upstream_url, trust_env=False) as upstream: + _clear_observations(upstream) + with gateway.client.stream( + "POST", + "/v1/chat/completions", + json={ + "model": model, + "messages": [{"role": "user", "content": "price this streamed request"}], + "stream": True, + "stream_options": {"include_usage": True}, + }, + headers={"Authorization": f"Bearer {key}"}, + ) as response: + stream_lines: Final = tuple(response.iter_lines()) + assert response.status_code == 200, "\n".join(stream_lines) + body: Final = _observed_request_body(upstream) + + events: Final = tuple( + JSON_OBJECT.validate_json(line.removeprefix("data: ")) + for line in stream_lines + if line.startswith("data: ") and line != "data: [DONE]" + ) + assert len(events) == len(frames) - 1, events + assert "".join(_stream_content(event) for event in events) == "".join(CONTENT), events + usage: Final = object_value(events[-1]["usage"]) + assert usage["total_tokens"] == 40, events[-1] + request_id: Final = string_value(events[0]["id"]) + rows: Final = eventually( + lambda: read_rows( + 'SELECT spend, request_duration_ms, ' + 'CAST(EXTRACT(EPOCH FROM ("endTime" - "startTime")) * 1000 AS DOUBLE PRECISION) ' + 'AS elapsed_duration_ms ' + 'FROM "LiteLLM_SpendLogs" WHERE request_id = %s', + (request_id,), + ), + lambda values: len(values) == 1, + seconds=70, + ) + spend: Final = float(str(rows[0]["spend"])) + request_duration_ms: Final = float(str(rows[0]["request_duration_ms"])) + elapsed_duration_ms: Final = float(str(rows[0]["elapsed_duration_ms"])) + assert spend == pytest.approx(RATE * request_duration_ms / 1000, rel=5e-2), ( + f"spend={spend}, request_duration_ms={request_duration_ms}, " + f"endTime-startTime duration_ms={elapsed_duration_ms}, body={body}" + ) + total_frame_delay_seconds: Final = (len(frames) - 1) * FRAME_DELAY_MS / 1000 + assert spend >= RATE * total_frame_delay_seconds * 0.95, ( + f"spend={spend}, total frame delay={total_frame_delay_seconds}s, body={body}" + ) + assert not PRICING_FIELDS.intersection(body), body diff --git a/tests/local_testing/test_embedding.py b/tests/local_testing/test_embedding.py index acbc4f20405..19885f891c0 100644 --- a/tests/local_testing/test_embedding.py +++ b/tests/local_testing/test_embedding.py @@ -713,7 +713,7 @@ def test_sagemaker_embeddings(): response = litellm.embedding( model="sagemaker/berri-benchmarking-gpt-j-6b-fp16", input=["good morning from litellm", "this is another item"], - input_cost_per_second=0.000420, + cost_per_second=0.000420, ) print(f"response: {response}") cost = completion_cost(completion_response=response) @@ -731,7 +731,7 @@ async def test_sagemaker_aembeddings(): response = await litellm.aembedding( model="sagemaker/berri-benchmarking-gpt-j-6b-fp16", input=["good morning from litellm", "this is another item"], - input_cost_per_second=0.000420, + cost_per_second=0.000420, ) print(f"response: {response}") cost = completion_cost(completion_response=response) diff --git a/tests/local_testing/test_sagemaker.py b/tests/local_testing/test_sagemaker.py index a01c8c217c6..bcbe230bc0a 100644 --- a/tests/local_testing/test_sagemaker.py +++ b/tests/local_testing/test_sagemaker.py @@ -55,7 +55,7 @@ async def test_completion_sagemaker(sync_mode): ], temperature=0.2, max_tokens=80, - input_cost_per_second=0.000420, + cost_per_second=0.000420, ) else: response = await litellm.acompletion( @@ -65,7 +65,7 @@ async def test_completion_sagemaker(sync_mode): ], temperature=0.2, max_tokens=80, - input_cost_per_second=0.000420, + cost_per_second=0.000420, ) # Add any assertions here to check the response print(response) @@ -169,7 +169,7 @@ async def test_completion_sagemaker_stream(sync_mode, model): temperature=0.2, stream=True, max_tokens=80, - input_cost_per_second=0.000420, + cost_per_second=0.000420, ) for idx, chunk in enumerate(response): @@ -187,7 +187,7 @@ async def test_completion_sagemaker_stream(sync_mode, model): stream=True, temperature=0.2, max_tokens=80, - input_cost_per_second=0.000420, + cost_per_second=0.000420, ) print("streaming response") @@ -280,7 +280,7 @@ async def test_acompletion_sagemaker_non_stream(): ], temperature=0.2, max_tokens=80, - input_cost_per_second=0.000420, + cost_per_second=0.000420, ) # Print what was called on the mock @@ -340,7 +340,7 @@ async def test_completion_sagemaker_non_stream(): ], temperature=0.2, max_tokens=80, - input_cost_per_second=0.000420, + cost_per_second=0.000420, ) # Print what was called on the mock @@ -457,7 +457,7 @@ async def test_completion_sagemaker_non_stream_with_aws_params(): ], temperature=0.2, max_tokens=80, - input_cost_per_second=0.000420, + cost_per_second=0.000420, aws_access_key_id="gm", aws_secret_access_key="s", aws_region_name="us-west-5", diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index dabd97cff0b..d3cb8e4a645 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -8697,7 +8697,7 @@ def test_model_has_no_cost_mapping_non_token_price_from_litellm_params_is_false( assert model_has_no_cost_mapping(model="custom-tts", llm_router=router) is False -@pytest.mark.parametrize("cost_field", ["input_cost_per_second", "input_cost_per_token"]) +@pytest.mark.parametrize("cost_field", ["cost_per_second", "input_cost_per_second", "input_cost_per_token"]) def test_model_has_no_cost_mapping_explicit_zero_price_is_false(cost_field): from litellm.proxy.auth.auth_checks import model_has_no_cost_mapping from litellm.router import Router diff --git a/tests/test_litellm/proxy/test_pricing_field_strip.py b/tests/test_litellm/proxy/test_pricing_field_strip.py index a84c6ba2b8a..a0e25e91f37 100644 --- a/tests/test_litellm/proxy/test_pricing_field_strip.py +++ b/tests/test_litellm/proxy/test_pricing_field_strip.py @@ -65,6 +65,7 @@ class TestStripClientPricingOverrides: for field in ( "input_cost_per_token", "output_cost_per_token", + "cost_per_second", "input_cost_per_second", "cache_creation_input_token_cost", ): diff --git a/tests/unit/litellm_core_utils/llm_cost_calc/test_zero_cost_diagnostic.py b/tests/unit/litellm_core_utils/llm_cost_calc/test_zero_cost_diagnostic.py index 0e453e3f5eb..921dde8bf7b 100644 --- a/tests/unit/litellm_core_utils/llm_cost_calc/test_zero_cost_diagnostic.py +++ b/tests/unit/litellm_core_utils/llm_cost_calc/test_zero_cost_diagnostic.py @@ -11,7 +11,7 @@ from litellm.litellm_core_utils.llm_cost_calc.zero_cost_diagnostic import ( ) from litellm.types.utils import CompletionTokensDetailsWrapper, PromptTokensDetailsWrapper, Usage -PER_SECOND_ENTRY: Final = {"input_cost_per_second": 0.00042, "output_cost_per_second": 0.00042} +PER_SECOND_ENTRY: Final = {"cost_per_second": 0.00042} FREE_ENTRY: Final = {"input_cost_per_token": 0, "output_cost_per_token": 0, "cache_read_input_token_cost": 2e-08} PRICED_ENTRY: Final = {"input_cost_per_token": 1e-06, "output_cost_per_token": 2e-06} TEXT_USAGE: Final = Usage(prompt_tokens=10, completion_tokens=20, total_tokens=30) diff --git a/tests/unit/litellm_core_utils/llm_response_utils/test_response_metadata.py b/tests/unit/litellm_core_utils/llm_response_utils/test_response_metadata.py index 6f297e6e06a..554447f4273 100644 --- a/tests/unit/litellm_core_utils/llm_response_utils/test_response_metadata.py +++ b/tests/unit/litellm_core_utils/llm_response_utils/test_response_metadata.py @@ -607,8 +607,7 @@ def test_update_response_metadata_prices_per_second_deployment_from_its_stamped_ litellm.register_model( model_cost={ deployment_id: { - "input_cost_per_second": 0.02, - "output_cost_per_second": 0.04, + "cost_per_second": 0.02, "litellm_provider": "openai", "mode": "chat", } @@ -627,8 +626,7 @@ def test_update_response_metadata_prices_per_second_deployment_from_its_stamped_ logging_obj.update_environment_variables( model="gpt-5.4-nano", litellm_params={ - "input_cost_per_second": 0.02, - "output_cost_per_second": 0.04, + "cost_per_second": 0.02, "metadata": {"model_info": {"id": deployment_id}}, }, optional_params={}, @@ -650,4 +648,4 @@ def test_update_response_metadata_prices_per_second_deployment_from_its_stamped_ ) assert result._response_ms == pytest.approx(2000) - assert result._hidden_params["response_cost"] == pytest.approx((0.02 + 0.04) * 2) + assert result._hidden_params["response_cost"] == pytest.approx(0.02 * 2) diff --git a/tests/unit/litellm_core_utils/test_get_litellm_params.py b/tests/unit/litellm_core_utils/test_get_litellm_params.py index 9b5771092ac..19a3323ce53 100644 --- a/tests/unit/litellm_core_utils/test_get_litellm_params.py +++ b/tests/unit/litellm_core_utils/test_get_litellm_params.py @@ -21,7 +21,13 @@ from litellm.litellm_core_utils.get_litellm_params import ( from litellm.types.litellm_params import ControlOptions NAMED_PRICE_PARAMS: Final = frozenset( - {"input_cost_per_token", "output_cost_per_token", "input_cost_per_second", "output_cost_per_second"} + { + "input_cost_per_token", + "output_cost_per_token", + "cost_per_second", + "input_cost_per_second", + "output_cost_per_second", + } ) diff --git a/tests/unit/litellm_core_utils/test_litellm_logging.py b/tests/unit/litellm_core_utils/test_litellm_logging.py index 2fc747e1b48..ed614b93a77 100644 --- a/tests/unit/litellm_core_utils/test_litellm_logging.py +++ b/tests/unit/litellm_core_utils/test_litellm_logging.py @@ -495,7 +495,7 @@ class TestZeroCostDiagnostic: DEPLOYMENT_ID: Final = "lit7898-query-only-priced-deployment" MODEL_GROUP: Final = "query-only-priced-chat" QUERY_ONLY_PRICING: Final = {"input_cost_per_query": 0.00042} - PER_SECOND_PRICING: Final = {"input_cost_per_second": 0.00042, "output_cost_per_second": 0.00042} + PER_SECOND_PRICING: Final = {"cost_per_second": 0.00042} FREE_PRICING: Final = {"input_cost_per_token": 0, "output_cost_per_token": 0} @pytest.fixture(params=["query_only", "free"]) @@ -845,7 +845,7 @@ class TestZeroCostDiagnostic: response: Final = self._response(usage) response._response_ms = 1000.0 with caplog.at_level(logging.WARNING, logger="LiteLLM"): - assert logging_obj._response_cost_calculator(result=response) == pytest.approx(0.00084) + assert logging_obj._response_cost_calculator(result=response) == pytest.approx(0.00042) assert logging_obj.model_call_details["zero_cost_diagnostic"] is None assert self._zero_cost_warnings(caplog) == [] diff --git a/tests/unit/router_strategy/test_complexity_router.py b/tests/unit/router_strategy/test_complexity_router.py index 2d66524326f..333524ffffc 100644 --- a/tests/unit/router_strategy/test_complexity_router.py +++ b/tests/unit/router_strategy/test_complexity_router.py @@ -5043,6 +5043,7 @@ class TestRouterPreRoutingAliasOverrides: "model": "auto_router/complexity_router", "input_cost_per_token": 0.0, "output_cost_per_token": 0.0, + "cost_per_second": 0.0, "input_cost_per_second": 0.0, "drop_params": True, "complexity_router_config": {"tiers": {"SIMPLE": "gpt-4o-mini"}}, @@ -5064,7 +5065,12 @@ class TestRouterPreRoutingAliasOverrides: assert result is not None # Non-pricing alias params still carry over. assert request_kwargs["drop_params"] is True - for field in ("input_cost_per_token", "output_cost_per_token", "input_cost_per_second"): + for field in ( + "input_cost_per_token", + "output_cost_per_token", + "cost_per_second", + "input_cost_per_second", + ): assert field not in request_kwargs @pytest.mark.asyncio diff --git a/tests/unit/test_cost_calculator.py b/tests/unit/test_cost_calculator.py index adfedf61d45..c1939a23e74 100644 --- a/tests/unit/test_cost_calculator.py +++ b/tests/unit/test_cost_calculator.py @@ -3037,9 +3037,9 @@ def test_completion_cost_logs_cache_and_reasoning_breakdown_for_custom_pricing() @pytest.mark.parametrize("custom_llm_provider", ["together_ai", "openai", "anthropic", "bedrock", "azure"]) def test_cost_per_token_per_second_pricing(monkeypatch, custom_llm_provider: str): """ - Models priced by duration (input/output_cost_per_second) with no per-token rates + Models priced by input/output duration rates with no per-token rates must be billed as cost_per_second * response_time_ms / 1000 in cost_per_token, - whether or not the provider has its own cost calculator. + using only the input rate even when both are set, whether or not the provider has its own calculator. """ monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) @@ -3064,11 +3064,40 @@ def test_cost_per_token_per_second_pricing(monkeypatch, custom_llm_provider: str response_time_ms=1500.0, ) - assert prompt_cost == pytest.approx(0.02 * 1.5) - assert completion_cost_value == pytest.approx(0.04 * 1.5) + assert (prompt_cost, completion_cost_value) == pytest.approx((0.02 * 1.5, 0.0)) -def test_cost_per_token_keeps_token_pricing_when_per_second_rates_are_also_set(monkeypatch): +def test_azure_chat_uses_token_rates_when_output_cost_per_second_is_set( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + + model: Final = "test-azure-chat-token-and-output-second-pricing" + litellm.register_model( + model_cost={ + model: { + "input_cost_per_token": 1e-6, + "output_cost_per_token": 2e-6, + "output_cost_per_second": 0.4, + "litellm_provider": "azure", + "mode": "chat", + } + } + ) + + cost: Final = cost_per_token( + model=model, + custom_llm_provider="azure", + prompt_tokens=10, + completion_tokens=20, + response_time_ms=1500.0, + ) + + assert cost == pytest.approx((10 * 1e-6, 20 * 2e-6)) + + +def test_cost_per_token_ignores_cost_per_second_when_token_pricing_is_set(monkeypatch): monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) @@ -3078,8 +3107,7 @@ def test_cost_per_token_keeps_token_pricing_when_per_second_rates_are_also_set(m model: { "input_cost_per_token": 1e-6, "output_cost_per_token": 2e-6, - "input_cost_per_second": 0.02, - "output_cost_per_second": 0.04, + "cost_per_second": 0.02, "litellm_provider": "openai", "mode": "chat", } @@ -3098,6 +3126,39 @@ def test_cost_per_token_keeps_token_pricing_when_per_second_rates_are_also_set(m assert completion_cost_value == pytest.approx(20 * 2e-6) +@pytest.mark.parametrize( + ("pricing_fields", "expected_rate"), + [ + ({"cost_per_second": 0.02}, 0.02), + ({"output_cost_per_second": 0.04}, 0.04), + ( + {"cost_per_second": 0.05, "input_cost_per_second": 0.02, "output_cost_per_second": 0.04}, + 0.05, + ), + ({"input_cost_per_second": 0.02}, 0.02), + ], +) +def test_cost_per_token_resolves_per_second_rate_precedence( + monkeypatch, pricing_fields: dict[str, float], expected_rate: float +): + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + + model: Final = "test-chat-per-second-rate-precedence" + entry: Final = {**pricing_fields, "litellm_provider": "together_ai", "mode": "chat"} + litellm.register_model( + model_cost={model: entry} + ) + + assert cost_per_token( + model=model, + custom_llm_provider="together_ai", + prompt_tokens=10, + completion_tokens=20, + response_time_ms=1500.0, + ) == pytest.approx((expected_rate * 1.5, 0.0)) + + def _logging_obj_with_call_window(duration_ms: float) -> Logging: start_time: Final = datetime.datetime(2026, 9, 21, 12, 0, 0) logging_obj: Final = Logging( @@ -3160,7 +3221,7 @@ def test_completion_cost_per_second_deployment_bills_the_call_duration( litellm_logging_obj=_logging_obj_with_call_window(logged_duration_ms), ) - assert cost == pytest.approx((0.02 + 0.04) * expected_seconds) + assert cost == pytest.approx(0.02 * expected_seconds) @pytest.mark.parametrize("mode", ["audio_transcription", "audio_speech", "video_generation", "realtime"]) diff --git a/tests/unit/test_register_model_custom_pricing.py b/tests/unit/test_register_model_custom_pricing.py index 452a15334ef..fa4fee8f6d8 100644 --- a/tests/unit/test_register_model_custom_pricing.py +++ b/tests/unit/test_register_model_custom_pricing.py @@ -11,6 +11,7 @@ calculations for DB-sourced models with prompt caching pricing. import copy import os +from typing import Final import pytest @@ -993,3 +994,21 @@ def test_completion_cost_applies_off_peak_only_deployment_pricing(): finally: _restore_model_cost_entries(original_entries) del router + + +def test_completion_registers_cost_per_second_pricing(): + model_key: Final = "openai/test-cost-per-second-registration" + original_entries: Final = _snapshot_model_cost_entries([model_key]) + + try: + litellm.completion( + model=model_key, + messages=[{"role": "user", "content": "hello"}], + api_key="fake-key", + cost_per_second=0.02, + mock_response="hello back", + ) + + assert litellm.model_cost[model_key]["cost_per_second"] == 0.02 + finally: + _restore_model_cost_entries(original_entries) diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index 2c612aa350c..ab7ae5ab3b8 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -648,6 +648,7 @@ def validate_model_cost_values(model_data, exceptions=None): "output_cost_per_image_4K", "input_cost_per_pixel", "output_cost_per_pixel", + "cost_per_second", "input_cost_per_second", "output_cost_per_second", "output_cost_per_second_480p", @@ -829,6 +830,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "input_cost_per_pixel": {"type": "number"}, "input_cost_per_query": {"type": "number"}, "input_cost_per_request": {"type": "number"}, + "cost_per_second": {"type": "number"}, "input_cost_per_second": {"type": "number"}, "input_cost_per_token": {"type": "number"}, "input_cost_per_token_above_128k_tokens": {"type": "number"}, diff --git a/tests/unit/types/test_router.py b/tests/unit/types/test_router.py index 4d4c326d1ca..4881b094cd6 100644 --- a/tests/unit/types/test_router.py +++ b/tests/unit/types/test_router.py @@ -40,6 +40,7 @@ def test_custom_pricing_params_keeps_every_field_it_had(): "output_cost_per_character", "cache_read_input_token_cost", "cache_creation_input_token_cost", + "cost_per_second", "input_cost_per_second", "cache_read_input_token_cost_flex", "input_cost_per_character_above_128k_tokens", diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 077e6a6ecee..dc0365a14b1 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -33013,6 +33013,8 @@ export interface components { complexity_router_default_model?: string | null; /** Configurable Clientside Auth Params */ configurable_clientside_auth_params?: (string | components["schemas"]["ConfigurableClientsideParamsCustomAuth-Input"])[] | null; + /** Cost Per Second */ + cost_per_second?: number | null; /** Custom Llm Provider */ custom_llm_provider?: string | null; /** Default Api Key Rpm Limit */ @@ -46842,6 +46844,8 @@ export interface components { complexity_router_default_model?: string | null; /** Configurable Clientside Auth Params */ configurable_clientside_auth_params?: (string | components["schemas"]["ConfigurableClientsideParamsCustomAuth-Input"])[] | null; + /** Cost Per Second */ + cost_per_second?: number | null; /** Custom Llm Provider */ custom_llm_provider?: string | null; /** Default Api Key Rpm Limit */ From f4a7c04d992b48f8cbdd0945087d506739bed9fc Mon Sep 17 00:00:00 2001 From: "berriai-litellm-provider-info-sync[bot]" <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 19:05:57 +0000 Subject: [PATCH 25/41] chore(cost-map): add openai gpt-6-astra ultrafast tier prices from the pricing page (#43745) * chore(cost-map): add openai gpt-6-astra ultrafast tier prices from the pricing page Price-Sync: litellm-providers * feat(cost): support openai ultrafast tier fields in the model catalog --------- Co-authored-by: berriai-litellm-provider-info-sync[bot] <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> Co-authored-by: Kerry --- .../crates/model-catalog/src/model_info.rs | 24 +++++++++++++ ...odel_prices_and_context_window_backup.json | 8 +++++ model_prices_and_context_window.json | 8 +++++ model_prices_and_context_window.schema.json | 36 +++++++++++++++++++ 4 files changed, 76 insertions(+) diff --git a/litellm-rust/crates/model-catalog/src/model_info.rs b/litellm-rust/crates/model-catalog/src/model_info.rs index 75f16e0c00d..380f6713d7a 100644 --- a/litellm-rust/crates/model-catalog/src/model_info.rs +++ b/litellm-rust/crates/model-catalog/src/model_info.rs @@ -57,6 +57,9 @@ pub struct ModelInfo { /// Priority service-tier rate for the same-named base field. #[serde(skip_serializing_if = "Option::is_none")] pub cache_creation_input_token_cost_above_272k_tokens_priority: Option, + /// Ultrafast service-tier rate for the same-named base field. + #[serde(skip_serializing_if = "Option::is_none")] + pub cache_creation_input_token_cost_above_272k_tokens_ultrafast: Option, #[serde(skip_serializing_if = "Option::is_none")] pub cache_creation_input_token_cost_batches: Option, /// Flex service-tier rate for the same-named base field. @@ -65,6 +68,9 @@ pub struct ModelInfo { /// Priority service-tier rate for the same-named base field. #[serde(skip_serializing_if = "Option::is_none")] pub cache_creation_input_token_cost_priority: Option, + /// Ultrafast service-tier rate for the same-named base field. + #[serde(skip_serializing_if = "Option::is_none")] + pub cache_creation_input_token_cost_ultrafast: Option, #[serde(skip_serializing_if = "Option::is_none")] pub cache_read_input_audio_token_cost: Option, #[serde(skip_serializing_if = "Option::is_none")] @@ -101,6 +107,9 @@ pub struct ModelInfo { /// Priority service-tier rate for the same-named base field. #[serde(skip_serializing_if = "Option::is_none")] pub cache_read_input_token_cost_above_272k_tokens_priority: Option, + /// Ultrafast service-tier rate for the same-named base field. + #[serde(skip_serializing_if = "Option::is_none")] + pub cache_read_input_token_cost_above_272k_tokens_ultrafast: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. #[serde(skip_serializing_if = "Option::is_none")] pub cache_read_input_token_cost_above_512k_tokens: Option, @@ -115,6 +124,9 @@ pub struct ModelInfo { /// Priority service-tier rate for the same-named base field. #[serde(skip_serializing_if = "Option::is_none")] pub cache_read_input_token_cost_priority: Option, + /// Ultrafast service-tier rate for the same-named base field. + #[serde(skip_serializing_if = "Option::is_none")] + pub cache_read_input_token_cost_ultrafast: Option, #[serde(skip_serializing_if = "Option::is_none")] pub citation_cost_per_token: Option, #[serde(skip_serializing_if = "Option::is_none")] @@ -213,6 +225,9 @@ pub struct ModelInfo { /// Priority service-tier rate for the same-named base field. #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_token_above_272k_tokens_priority: Option, + /// Ultrafast service-tier rate for the same-named base field. + #[serde(skip_serializing_if = "Option::is_none")] + pub input_cost_per_token_above_272k_tokens_ultrafast: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_token_above_512k_tokens: Option, @@ -230,6 +245,9 @@ pub struct ModelInfo { /// Priority service-tier rate for the same-named base field. #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_token_priority: Option, + /// Ultrafast service-tier rate for the same-named base field. + #[serde(skip_serializing_if = "Option::is_none")] + pub input_cost_per_token_ultrafast: Option, #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_video_per_second: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. @@ -362,6 +380,9 @@ pub struct ModelInfo { /// Priority service-tier rate for the same-named base field. #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_token_above_272k_tokens_priority: Option, + /// Ultrafast service-tier rate for the same-named base field. + #[serde(skip_serializing_if = "Option::is_none")] + pub output_cost_per_token_above_272k_tokens_ultrafast: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_token_above_512k_tokens: Option, @@ -377,6 +398,9 @@ pub struct ModelInfo { /// Priority service-tier rate for the same-named base field. #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_token_priority: Option, + /// Ultrafast service-tier rate for the same-named base field. + #[serde(skip_serializing_if = "Option::is_none")] + pub output_cost_per_token_ultrafast: Option, #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_video_per_second: Option, #[serde(skip_serializing_if = "Option::is_none")] diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 9ae270243d2..c147f4bf94b 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -33870,16 +33870,22 @@ "cache_read_input_token_cost_above_272k_tokens_batches": 1e-06, "cache_creation_input_token_cost_batches": 6.25e-06, "cache_creation_input_token_cost_above_272k_tokens_batches": 1.25e-05, + "cache_creation_input_token_cost_above_272k_tokens_ultrafast": 0.00015, + "cache_creation_input_token_cost_ultrafast": 7.5e-05, + "cache_read_input_token_cost_above_272k_tokens_ultrafast": 1.2e-05, "cache_read_input_token_cost_flex": 5e-07, "cache_read_input_token_cost_priority": 2e-06, + "cache_read_input_token_cost_ultrafast": 6e-06, "input_cost_per_token": 1e-05, "input_cost_per_token_above_272k_tokens": 2e-05, "input_cost_per_token_above_272k_tokens_flex": 1e-05, "input_cost_per_token_above_272k_tokens_priority": 4e-05, "input_cost_per_token_batches": 5e-06, "input_cost_per_token_above_272k_tokens_batches": 1e-05, + "input_cost_per_token_above_272k_tokens_ultrafast": 0.00012, "input_cost_per_token_flex": 5e-06, "input_cost_per_token_priority": 2e-05, + "input_cost_per_token_ultrafast": 6e-05, "litellm_provider": "openai", "max_input_tokens": 922000, "max_output_tokens": 128000, @@ -33891,8 +33897,10 @@ "output_cost_per_token_above_272k_tokens_priority": 0.00015, "output_cost_per_token_batches": 2.5e-05, "output_cost_per_token_above_272k_tokens_batches": 3.75e-05, + "output_cost_per_token_above_272k_tokens_ultrafast": 0.00045, "output_cost_per_token_flex": 2.5e-05, "output_cost_per_token_priority": 0.0001, + "output_cost_per_token_ultrafast": 0.0003, "regional_processing_uplift_multiplier_eu": 1.1, "regional_processing_uplift_multiplier_us": 1.1, "search_context_cost_per_query": { diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 9ae270243d2..c147f4bf94b 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -33870,16 +33870,22 @@ "cache_read_input_token_cost_above_272k_tokens_batches": 1e-06, "cache_creation_input_token_cost_batches": 6.25e-06, "cache_creation_input_token_cost_above_272k_tokens_batches": 1.25e-05, + "cache_creation_input_token_cost_above_272k_tokens_ultrafast": 0.00015, + "cache_creation_input_token_cost_ultrafast": 7.5e-05, + "cache_read_input_token_cost_above_272k_tokens_ultrafast": 1.2e-05, "cache_read_input_token_cost_flex": 5e-07, "cache_read_input_token_cost_priority": 2e-06, + "cache_read_input_token_cost_ultrafast": 6e-06, "input_cost_per_token": 1e-05, "input_cost_per_token_above_272k_tokens": 2e-05, "input_cost_per_token_above_272k_tokens_flex": 1e-05, "input_cost_per_token_above_272k_tokens_priority": 4e-05, "input_cost_per_token_batches": 5e-06, "input_cost_per_token_above_272k_tokens_batches": 1e-05, + "input_cost_per_token_above_272k_tokens_ultrafast": 0.00012, "input_cost_per_token_flex": 5e-06, "input_cost_per_token_priority": 2e-05, + "input_cost_per_token_ultrafast": 6e-05, "litellm_provider": "openai", "max_input_tokens": 922000, "max_output_tokens": 128000, @@ -33891,8 +33897,10 @@ "output_cost_per_token_above_272k_tokens_priority": 0.00015, "output_cost_per_token_batches": 2.5e-05, "output_cost_per_token_above_272k_tokens_batches": 3.75e-05, + "output_cost_per_token_above_272k_tokens_ultrafast": 0.00045, "output_cost_per_token_flex": 2.5e-05, "output_cost_per_token_priority": 0.0001, + "output_cost_per_token_ultrafast": 0.0003, "regional_processing_uplift_multiplier_eu": 1.1, "regional_processing_uplift_multiplier_us": 1.1, "search_context_cost_per_query": { diff --git a/model_prices_and_context_window.schema.json b/model_prices_and_context_window.schema.json index c4b424a26b4..cdf023e71ef 100644 --- a/model_prices_and_context_window.schema.json +++ b/model_prices_and_context_window.schema.json @@ -133,6 +133,11 @@ "minimum": 0, "description": "Priority service-tier rate for the same-named base field." }, + "cache_creation_input_token_cost_above_272k_tokens_ultrafast": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, "cache_creation_input_token_cost_above_32k_tokens": { "type": "number", "minimum": 0, @@ -152,6 +157,10 @@ "minimum": 0, "description": "Priority service-tier rate for the same-named base field." }, + "cache_creation_input_token_cost_ultrafast": { + "type": "number", + "minimum": 0 + }, "cache_read_input_audio_token_cost": { "type": "number", "minimum": 0 @@ -210,6 +219,11 @@ "minimum": 0, "description": "Priority service-tier rate for the same-named base field." }, + "cache_read_input_token_cost_above_272k_tokens_ultrafast": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, "cache_read_input_token_cost_above_32k_tokens": { "type": "number", "minimum": 0, @@ -238,6 +252,10 @@ "minimum": 0, "description": "Priority service-tier rate for the same-named base field." }, + "cache_read_input_token_cost_ultrafast": { + "type": "number", + "minimum": 0 + }, "citation_cost_per_token": { "type": "number", "minimum": 0 @@ -404,6 +422,11 @@ "minimum": 0, "description": "Priority service-tier rate for the same-named base field." }, + "input_cost_per_token_above_272k_tokens_ultrafast": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, "input_cost_per_token_above_32k_tokens": { "type": "number", "minimum": 0, @@ -437,6 +460,10 @@ "minimum": 0, "description": "Priority service-tier rate for the same-named base field." }, + "input_cost_per_token_ultrafast": { + "type": "number", + "minimum": 0 + }, "input_cost_per_video_per_second": { "type": "number", "minimum": 0 @@ -770,6 +797,11 @@ "minimum": 0, "description": "Priority service-tier rate for the same-named base field." }, + "output_cost_per_token_above_272k_tokens_ultrafast": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, "output_cost_per_token_above_32k_tokens": { "type": "number", "minimum": 0, @@ -799,6 +831,10 @@ "minimum": 0, "description": "Priority service-tier rate for the same-named base field." }, + "output_cost_per_token_ultrafast": { + "type": "number", + "minimum": 0 + }, "output_cost_per_video_per_second": { "type": "number", "minimum": 0 From 3b2a447fae76c016b5eb982fd7f31bd359465ab4 Mon Sep 17 00:00:00 2001 From: tin-berri Date: Tue, 29 Sep 2026 12:40:19 -0700 Subject: [PATCH 26/41] fix(autorouter): compare historical and new savings consistently (#43348) * fix(autorouter): compare historical and new savings consistently * fix(autorouter): reject comparisons if request counts changed * fix(router): restore eligible LLM and classification breakdown * fix(router): avoid ambiguous baseline labels for partial comparisons --- litellm/models/autorouter_session.py | 15 +- .../proxy/db/autorouter_savings_comparison.py | 147 ++++++++++++++++++ litellm/proxy/db/autorouter_session_rollup.py | 20 ++- .../auto_router_endpoints.py | 114 ++++++++++++-- .../auto_router_endpoints.py | 27 ++-- .../spend/test_autorouter_session_rollup.py | 65 ++++++++ .../test_auto_router_endpoints.py | 39 +++-- .../AutoRouterBenchmarksTab.test.tsx | 87 ++++++----- .../_components/AutoRouterBenchmarksTab.tsx | 64 ++++---- ...KeyAutoRouterUsageTab.integration.test.tsx | 2 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 38 +++-- 11 files changed, 482 insertions(+), 136 deletions(-) create mode 100644 litellm/proxy/db/autorouter_savings_comparison.py diff --git a/litellm/models/autorouter_session.py b/litellm/models/autorouter_session.py index ddce2b5ef81..9df9af18dff 100644 --- a/litellm/models/autorouter_session.py +++ b/litellm/models/autorouter_session.py @@ -34,15 +34,12 @@ class LiteLLM_AutoRouterSession(LiteLLMPydanticObjectBase): @property def baseline_model(self) -> str | None: - """The baseline most covered turns were priced against, or None when none were estimated. - - A router reconfigured mid-session leaves turns priced against two baselines; the row keeps both - counts, and the label is the one that priced the most money-carrying turns rather than whatever the - router is configured with now. - """ - if not self.savings_estimated_baseline_models: + """A recorded baseline label when excluded turns cannot change the selected model.""" + if not self.baseline_models: + return None + if self.savings_estimated_turns < self.turns and len(self.baseline_models) > 1: return None return max( - self.savings_estimated_baseline_models, - key=lambda model: (self.savings_estimated_baseline_models[model], model), + self.baseline_models, + key=lambda model: (self.baseline_models[model], model), ) diff --git a/litellm/proxy/db/autorouter_savings_comparison.py b/litellm/proxy/db/autorouter_savings_comparison.py new file mode 100644 index 00000000000..041496d63f2 --- /dev/null +++ b/litellm/proxy/db/autorouter_savings_comparison.py @@ -0,0 +1,147 @@ +from collections.abc import Mapping +from contextlib import AbstractAsyncContextManager +from datetime import timedelta +from math import isclose +from types import MappingProxyType +from typing import TYPE_CHECKING, Final, Protocol, cast + +from pydantic import BaseModel, ConfigDict, TypeAdapter + +from litellm._logging import verbose_proxy_logger +from litellm.constants import MAX_SPENDLOG_ROWS_TO_QUERY +from litellm.proxy.db.autorouter_session_rollup import AUTOROUTER_SESSION_WINDOW_SQL +from litellm.proxy.db.create_views import SupportsRawQueries + +if TYPE_CHECKING: + from litellm.proxy.utils import PrismaClient + + +class SessionSavingsComparison(BaseModel): + model_config = ConfigDict(frozen=True, allow_inf_nan=False) + + router_name: str + router_type: str + turns: int + estimated_turns: int + actual_spend: float + classifier_cost: float | None + saved_spend: float + complete: bool + + def coverage_fields(self, recorded_savings: float, recorded_turns: int) -> Mapping[str, float | int]: + if self.turns != recorded_turns or not self.complete: + return MappingProxyType({}) + if not isclose(self.saved_spend, recorded_savings, rel_tol=1e-9, abs_tol=1e-9): + return MappingProxyType({}) + return MappingProxyType( + { + "savings_estimated_turns": self.estimated_turns, + "savings_estimated_actual_spend": self.actual_spend, + "savings_estimated_saved_spend": self.saved_spend, + } + ) + + +class _ReadTransactions(Protocol): + def tx(self, *, timeout: timedelta, max_wait: timedelta) -> AbstractAsyncContextManager[SupportsRawQueries]: ... + + +_COMPARISONS: Final = TypeAdapter(tuple[SessionSavingsComparison, ...]) + + +async def historical_session_comparisons( + prisma_client: "PrismaClient", + start_date: str, + end_date: str, + api_key: str | None, + user_id: str | None, + session_id: str | None = None, +) -> Mapping[tuple[str, str], SessionSavingsComparison]: + try: + reader: Final = cast(_ReadTransactions, prisma_client.read_db) # cast-ok: untyped Prisma transaction delegate + async with reader.tx(timeout=timedelta(seconds=3), max_wait=timedelta(seconds=1)) as transaction: + await transaction.execute_raw("SET TRANSACTION READ ONLY") + await transaction.execute_raw("SET LOCAL statement_timeout = 2000") + rows: Final = await transaction.query_raw( + HISTORICAL_SESSION_COMPARISONS_SQL, + start_date, + end_date, + api_key, + user_id, + session_id, + ) + comparisons: Final = _COMPARISONS.validate_python(rows or ()) + return MappingProxyType({(row.router_name, row.router_type): row for row in comparisons}) + except Exception: # noqa: BLE001 # missing retained logs must not discard recorded dollar savings + verbose_proxy_logger.warning("Historical auto-router cost comparison unavailable; preserving recorded savings") + return MappingProxyType({}) + + +HISTORICAL_SESSION_COMPARISONS_SQL: Final = f""" +WITH {AUTOROUTER_SESSION_WINDOW_SQL}, scoped AS MATERIALIZED ( + SELECT * FROM windowed WHERE $5::text IS NULL OR session_id = $5::text +), limited_logs AS MATERIALIZED ( + SELECT session.api_key, session.session_id, session.router_name, session.router_type, session.comparison_user_id, + session.classifier_cost_recorded_turns = session.turns AS classifier_cost_tracked, + logs.spend, logs.prompt_tokens + logs.completion_tokens AS tokens, + logs.metadata::jsonb -> 'routing_decision' AS decision, + logs.metadata::jsonb -> 'autorouter_savings' AS savings, + logs.metadata::jsonb -> 'autorouter_savings_estimate' AS estimate + FROM scoped AS session JOIN "LiteLLM_SpendLogs" AS logs + ON logs.api_key = session.api_key + AND CASE WHEN char_length(logs.session_id) > 256 + THEN 'sha256:' || encode(sha256(convert_to(logs.session_id, 'UTF8')), 'hex') + ELSE logs.session_id END = session.session_id + AND (session.comparison_user_id IS NULL OR logs."user" = session.comparison_user_id) + AND logs."startTime" BETWEEN session.first_turn_at AND session.last_turn_at + AND COALESCE(logs.metadata::jsonb #>> '{{routing_decision,router_model_name}}', logs.model_group) + = session.router_name + WHERE session.savings_estimated_turns < session.turns + AND logs.status = 'success' AND COALESCE(logs.metadata::jsonb ->> 'internal_call_origin', '') = '' + LIMIT {MAX_SPENDLOG_ROWS_TO_QUERY + 1} +), facts AS ( + SELECT *, + CASE WHEN jsonb_typeof(decision -> 'classifier_cost') = 'number' + THEN (decision ->> 'classifier_cost')::float8 + WHEN classifier_cost_tracked THEN 0 END AS classifier, + CASE WHEN jsonb_typeof(savings) = 'number' AND ( + estimate IS NULL OR estimate = 'null'::jsonb OR ( + jsonb_typeof(estimate -> 'version') = 'number' AND estimate ->> 'version' IN ('1', '2', '3') + AND estimate ->> 'status' = 'estimated' + ) + ) THEN savings::text::float8 END AS saved + FROM limited_logs +), compared AS ( + SELECT api_key, session_id, router_name, router_type, comparison_user_id, + COUNT(*) AS turns, SUM(spend + COALESCE(classifier, 0)) AS spend, SUM(tokens) AS total_tokens, + COUNT(saved) AS estimated_turns, + COALESCE(SUM(spend + COALESCE(classifier, 0)) FILTER (WHERE saved IS NOT NULL), 0)::float8 AS actual_spend, + CASE WHEN COUNT(saved) = COUNT(classifier) FILTER (WHERE saved IS NOT NULL) + THEN COALESCE(SUM(classifier) FILTER (WHERE saved IS NOT NULL), 0)::float8 + END AS estimated_classifier_cost, + COALESCE(SUM(saved), 0)::float8 AS saved_spend + FROM facts GROUP BY 1, 2, 3, 4, 5 +), reconciled AS ( + SELECT session.*, logs.estimated_turns, logs.actual_spend, logs.estimated_classifier_cost, + COALESCE((SELECT COUNT(*) FROM limited_logs) <= {MAX_SPENDLOG_ROWS_TO_QUERY} + AND logs.turns = session.turns AND logs.total_tokens = session.total_tokens + AND ABS(logs.spend - session.spend) <= GREATEST(1e-9, ABS(session.spend) * 1e-9) + AND ABS(logs.saved_spend - session.saved_spend) <= GREATEST(1e-9, ABS(session.saved_spend) * 1e-9), FALSE + ) AS recovered + FROM scoped AS session LEFT JOIN compared AS logs + ON logs.api_key = session.api_key AND logs.session_id = session.session_id + AND logs.router_name = session.router_name AND logs.router_type = session.router_type + AND logs.comparison_user_id IS NOT DISTINCT FROM session.comparison_user_id +) +SELECT router_name, router_type, + SUM(turns)::bigint AS turns, + SUM(CASE WHEN recovered THEN estimated_turns ELSE savings_estimated_turns END)::bigint AS estimated_turns, + SUM(CASE WHEN recovered THEN actual_spend ELSE savings_estimated_actual_spend END)::float8 AS actual_spend, + CASE WHEN BOOL_AND(CASE WHEN recovered THEN estimated_classifier_cost IS NOT NULL + ELSE savings_estimated_turns = turns AND classifier_cost_recorded_turns = turns END) + THEN SUM(CASE WHEN recovered THEN estimated_classifier_cost ELSE classifier_cost END)::float8 + END AS classifier_cost, + SUM(saved_spend)::float8 AS saved_spend, + BOOL_AND(recovered OR savings_estimated_turns = turns) AS complete +FROM reconciled GROUP BY router_name, router_type +""" diff --git a/litellm/proxy/db/autorouter_session_rollup.py b/litellm/proxy/db/autorouter_session_rollup.py index dd08cfd1bef..b762a40f344 100644 --- a/litellm/proxy/db/autorouter_session_rollup.py +++ b/litellm/proxy/db/autorouter_session_rollup.py @@ -7,8 +7,8 @@ on the prisma client. The spend-log flush job drains the queue into key and user session rollups with one atomic statement per turn: each upsert classifies the turn (same model, first visit, return to a model the session already used, out of order) against the row's own columns, so nothing is read before the write and concurrent -pods compose. The benchmarks endpoint aggregates these rows and never touches -LiteLLM_SpendLogs. +pods compose. The benchmarks endpoint aggregates these rows and can recover matching historical +costs from retained spend logs when estimate coverage predates these columns. """ from __future__ import annotations @@ -45,20 +45,24 @@ _SESSION_COLUMNS: Final = """ savings_estimated_baseline_models """ -AUTOROUTER_BENCHMARKS_SQL: Final = f""" -WITH windowed AS ( - SELECT {_SESSION_COLUMNS} FROM "LiteLLM_AutoRouterSession" +AUTOROUTER_SESSION_WINDOW_SQL: Final = f""" +windowed AS ( + SELECT {_SESSION_COLUMNS}, NULL::text AS comparison_user_id FROM "LiteLLM_AutoRouterSession" WHERE $4::text IS NULL AND last_turn_at >= $1::timestamp AND first_turn_at < $2::timestamp AND ($3::text IS NULL OR api_key = $3::text) UNION ALL - SELECT {_SESSION_COLUMNS} FROM "LiteLLM_AutoRouterUserSession" + SELECT {_SESSION_COLUMNS}, user_id AS comparison_user_id FROM "LiteLLM_AutoRouterUserSession" WHERE (($4::text IS NOT NULL AND user_id = $4::text) OR ($4::text IS NULL AND api_key = '')) AND last_turn_at >= $1::timestamp AND first_turn_at < $2::timestamp AND ($3::text IS NULL OR api_key = $3::text) -), +) +""" + +AUTOROUTER_BENCHMARKS_SQL: Final = f""" +WITH {AUTOROUTER_SESSION_WINDOW_SQL}, tier_maps AS ( SELECT router_name, router_type, jsonb_object_agg(tier, tier_turns) AS tier_turns FROM ( @@ -95,6 +99,8 @@ SELECT COALESCE(SUM(saved_spend), 0)::float8 AS saved_spend, COALESCE(SUM(savings_estimated_turns), 0)::int AS savings_estimated_turns, COALESCE(SUM(savings_estimated_actual_spend), 0)::float8 AS savings_estimated_actual_spend, + CASE WHEN BOOL_AND(savings_estimated_turns = turns AND classifier_cost_recorded_turns = turns) + THEN SUM(classifier_cost)::float8 END AS savings_estimated_classifier_cost, COALESCE(SUM(savings_estimated_saved_spend), 0)::float8 AS savings_estimated_saved_spend, COALESCE(SUM(classifier_cost), 0)::float8 AS classifier_cost, COALESCE(SUM(classifier_cost_recorded_turns), 0)::int AS classifier_cost_recorded_turns, diff --git a/litellm/proxy/management_endpoints/auto_router_endpoints.py b/litellm/proxy/management_endpoints/auto_router_endpoints.py index 9708161397a..e3bb2b0b6cc 100644 --- a/litellm/proxy/management_endpoints/auto_router_endpoints.py +++ b/litellm/proxy/management_endpoints/auto_router_endpoints.py @@ -8,6 +8,7 @@ POST /auto_router/validate_complexity_router_config - Dry-run the complexity-rou from collections.abc import Mapping, Sequence from datetime import datetime, timedelta, timezone from itertools import chain, groupby +from math import isclose from types import MappingProxyType from typing import TYPE_CHECKING, Annotated, Final, Protocol from uuid import uuid4 @@ -31,6 +32,7 @@ from litellm.proxy.auth.auth_checks import ( can_key_call_resolved_model, ) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.db.autorouter_savings_comparison import historical_session_comparisons from litellm.proxy.db.autorouter_session_rollup import ( AUTOROUTER_BENCHMARKS_SQL, bounded_session_id, @@ -651,7 +653,9 @@ class _SessionAggRow(BaseModel): saved_spend: float savings_estimated_turns: int = 0 savings_estimated_actual_spend: float = 0.0 + savings_estimated_classifier_cost: float | None = None savings_estimated_saved_spend: float = 0.0 + savings_comparison_complete: bool = True classifier_cost: float classifier_cost_recorded_turns: int session_seconds: float @@ -679,18 +683,25 @@ def _cache_bucket(turns: int, hits: int) -> AutoRouterCacheBucket: def _savings_cohort( - turns: int, estimated_turns: int, actual_spend: float, saved_spend: float + turns: int, estimated_turns: int, actual_spend: float, saved_spend: float, recorded_savings: float ) -> tuple[float | None, float | None]: - if turns > 0 and estimated_turns == 0: + if turns > 0 and estimated_turns == 0 and recorded_savings == 0: return None, None - return saved_spend, actual_spend + saved_spend + if not isclose(saved_spend, recorded_savings, rel_tol=1e-9, abs_tol=1e-9): + return recorded_savings, None + return recorded_savings, actual_spend + recorded_savings def _benchmark_totals(row: _SessionAggRow) -> AutoRouterBenchmarkTotals: return_misses: Final = row.return_turns - row.return_hits - saved_spend, baseline_spend = _savings_cohort( - row.turns, row.savings_estimated_turns, row.savings_estimated_actual_spend, row.savings_estimated_saved_spend + saved_spend, compared_baseline = _savings_cohort( + row.turns, + row.savings_estimated_turns, + row.savings_estimated_actual_spend, + row.savings_estimated_saved_spend, + row.saved_spend, ) + baseline_spend: Final = compared_baseline if row.savings_comparison_complete else None sessions: Final = row.sessions return AutoRouterBenchmarkTotals( sessions=sessions, @@ -701,13 +712,12 @@ def _benchmark_totals(row: _SessionAggRow) -> AutoRouterBenchmarkTotals: spend=row.spend, savings_estimated_turns=row.savings_estimated_turns, savings_estimated_actual_spend=row.savings_estimated_actual_spend, + savings_estimated_classifier_cost=row.savings_estimated_classifier_cost if baseline_spend is not None else None, saved_spend=saved_spend, classifier_cost=row.classifier_cost if row.classifier_cost_recorded_turns == row.turns else None, baseline_spend=baseline_spend, saved_pct=_pct(saved_spend, baseline_spend) if saved_spend is not None and baseline_spend is not None else None, - saved_per_session=(row.savings_estimated_saved_spend / sessions if sessions else 0.0) - if row.savings_estimated_turns == row.turns - else None, + saved_per_session=(saved_spend / sessions if sessions else 0.0) if saved_spend is not None else None, cache=AutoRouterCacheStats( coverage_pct=_pct(row.covered_turns, row.turns), hit_rate_pct=_pct(row.cache_hits, row.covered_turns), @@ -739,6 +749,7 @@ def _benchmark_group(row: _SessionAggRow) -> AutoRouterBenchmarkGroup: saved_spend=totals.saved_spend, savings_estimated_turns=totals.savings_estimated_turns, savings_estimated_actual_spend=totals.savings_estimated_actual_spend, + savings_estimated_classifier_cost=totals.savings_estimated_classifier_cost, classifier_cost=totals.classifier_cost, baseline_spend=totals.baseline_spend, saved_pct=totals.saved_pct, @@ -772,7 +783,13 @@ def _summed_agg_row(rows: Sequence[_SessionAggRow]) -> _SessionAggRow: saved_spend=sum(row.saved_spend for row in rows), savings_estimated_turns=sum(row.savings_estimated_turns for row in rows), savings_estimated_actual_spend=sum(row.savings_estimated_actual_spend for row in rows), + savings_estimated_classifier_cost=( + sum(row.savings_estimated_classifier_cost or 0.0 for row in rows) + if all(row.savings_estimated_classifier_cost is not None for row in rows) + else None + ), savings_estimated_saved_spend=sum(row.savings_estimated_saved_spend for row in rows), + savings_comparison_complete=all(row.savings_comparison_complete for row in rows), classifier_cost=sum(row.classifier_cost for row in rows), classifier_cost_recorded_turns=sum(row.classifier_cost_recorded_turns for row in rows), session_seconds=sum(row.session_seconds for row in rows), @@ -847,8 +864,8 @@ async def get_auto_router_benchmarks( Benchmarks for the auto-router dashboard: session shape, savings against the configured baseline, and prompt-caching behaviour bucketed by what the router did. - Reads session rollups folded once per request at spend-write time, so this endpoint - never scans LiteLLM_SpendLogs. A user filter selects only turns attributed to that + Reads session rollups folded once per request at spend-write time, with bounded + retained-log recovery for historical comparisons. A user filter selects only turns attributed to that internal user when written; older key-only history remains outside user views. A session is in the window when it overlaps it: its last turn is on or after start_date and its first turn is on or before end_date. Overall hit rate is over telemetry-bearing turns; each bucket's hit rate is @@ -882,7 +899,44 @@ async def get_auto_router_benchmarks( api_key, user_id, ) - rows: Final = _SESSION_AGG_ROWS.validate_python(raw_rows or ()) + recorded_rows: Final = _SESSION_AGG_ROWS.validate_python(raw_rows or ()) + comparisons: Final = ( + await historical_session_comparisons( + prisma_client, + start_day.isoformat(), + (end_day + timedelta(days=1)).isoformat(), + api_key, + user_id, + ) + if any(row.savings_estimated_turns < row.turns for row in recorded_rows) + else MappingProxyType({}) + ) + covered_rows: Final = tuple( + row.model_copy( + update={ + **comparison.coverage_fields(row.saved_spend, row.turns), + "savings_estimated_classifier_cost": comparison.classifier_cost, + "savings_comparison_complete": comparison.complete and comparison.turns == row.turns, + } + ) + if (comparison := comparisons.get((row.router_name, row.router_type))) + else row.model_copy(update={"savings_comparison_complete": row.savings_estimated_turns == row.turns}) + for row in recorded_rows + ) + rows: Final = tuple( + row.model_copy( + update={ + "savings_comparison_complete": row.savings_comparison_complete + and isclose( + row.saved_spend, + row.savings_estimated_saved_spend, + rel_tol=1e-9, + abs_tol=1e-9, + ), + } + ) + for row in covered_rows + ) groups: Final = ( *(_benchmark_group(row) for row in rows), *_idle_router_groups(llm_router, frozenset((row.router_name, row.router_type) for row in rows)), @@ -920,15 +974,43 @@ async def get_auto_router_session( if prisma_client is None: raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value) - row: Final = await AutoRouterSessionRepository(prisma_client).find_latest_for_key( + recorded: Final = await AutoRouterSessionRepository(prisma_client).find_latest_for_key( user_api_key_dict.api_key, bounded_session_id(session_id) ) - if row is None: + if recorded is None: raise HTTPException( status_code=404, detail=f"No auto-routed turns recorded for session {session_id!r} under this key" ) - saved_spend, baseline_spend = _savings_cohort( - row.turns, row.savings_estimated_turns, row.savings_estimated_actual_spend, row.savings_estimated_saved_spend + comparisons: Final = ( + await historical_session_comparisons( + prisma_client, + recorded.first_turn_at.isoformat(), + (recorded.last_turn_at + timedelta(microseconds=1)).isoformat(), + user_api_key_dict.api_key, + None, + bounded_session_id(session_id), + ) + if recorded.savings_estimated_turns < recorded.turns + else MappingProxyType({}) + ) + comparison: Final = comparisons.get((recorded.router_name, recorded.router_type)) + row: Final = ( + recorded.model_copy(update=comparison.coverage_fields(recorded.saved_spend, recorded.turns)) + if comparison + else recorded + ) + saved_spend, compared_baseline = _savings_cohort( + row.turns, + row.savings_estimated_turns, + row.savings_estimated_actual_spend, + row.savings_estimated_saved_spend, + row.saved_spend, + ) + baseline_spend: Final = ( + compared_baseline + if row.savings_estimated_turns == row.turns + or (comparison and comparison.complete and comparison.turns == row.turns) + else None ) return AutoRouterSessionResponse( session_id=session_id, @@ -943,7 +1025,7 @@ async def get_auto_router_session( baseline_spend=baseline_spend if row.savings_estimated_turns == row.turns else None, savings_estimated_baseline_spend=baseline_spend, baseline_model=row.baseline_model, - baseline_models=row.savings_estimated_baseline_models, + baseline_models=row.baseline_models, ) diff --git a/litellm/types/management_endpoints/auto_router_endpoints.py b/litellm/types/management_endpoints/auto_router_endpoints.py index e191470ec6e..75a80beac5c 100644 --- a/litellm/types/management_endpoints/auto_router_endpoints.py +++ b/litellm/types/management_endpoints/auto_router_endpoints.py @@ -224,19 +224,24 @@ class AutoRouterBenchmarkTotals(BaseModel): "subtotal recording, and zero for an empty window" ) savings_estimated_turns: int = Field( - description="Turns covered by the current savings estimator; legacy estimates are excluded" + description="Requests with a matching savings comparison, including historical recorded estimates" ) savings_estimated_actual_spend: float = Field( description="Actual spend, including classifier cost, for covered turns only" ) + savings_estimated_classifier_cost: float | None = Field( + default=None, + description="Classifier cost included in the matching historical and newer savings comparison; " + "null when classification costs for those requests are unavailable", + ) saved_spend: float | None = Field( - description="Signed savings for covered turns only; null when traffic has no current estimates" + description="Recorded historical savings plus newer estimates; null when traffic has no recorded savings estimates" ) baseline_spend: float | None = Field(description="Estimated single-model cost for covered turns only") - saved_pct: float | None = Field(description="Covered savings over covered baseline spend, as a percentage") - saved_per_session: float | None = Field( - description="Average session savings; unavailable unless every turn is covered" + saved_pct: float | None = Field( + description="Total recorded savings over the matching historical and current baseline; null when costs are unavailable" ) + saved_per_session: float | None = Field(description="Recorded savings per session, including historical estimates") cache: AutoRouterCacheStats @@ -268,12 +273,14 @@ class AutoRouterSessionResponse(BaseModel): last_model: str = Field(description="The deployment model the most recent turn was routed to") spend: float = Field(description="What the session's routed traffic actually cost, classifier calls included") savings_estimated_turns: int = Field( - description="Turns covered by the current savings estimator; legacy estimates are excluded" + description="Requests with a matching savings comparison, including historical recorded estimates" ) savings_estimated_actual_spend: float = Field( description="Actual spend, including classifier cost, for covered turns only" ) - saved_spend: float | None = Field(description="Estimated savings for covered turns only, net of classifier cost") + saved_spend: float | None = Field( + description="Recorded historical savings plus newer estimates, net of classifier cost" + ) baseline_spend: float | None = Field( description="Estimated single-model cost; unavailable unless every turn is covered" ) @@ -281,14 +288,14 @@ class AutoRouterSessionResponse(BaseModel): description="Estimated single-model cost for covered turns only" ) baseline_model: str | None = Field( - description="The savings baseline most covered turns were priced against, recorded turn by " + description="The savings baseline recorded by most session turns, including historical turns, recorded turn by " "turn, so it still names the counterfactual after the router is reconfigured or removed. None when no " "turn recorded one: rows from before the baseline was recorded, and adaptive and quality routers, " "which derive no baseline and so report no savings" ) baseline_models: Mapping[str, int] = Field( - description="Covered turns priced against each baseline model; more than one entry means the router's " - "baseline changed mid-session and baseline_spend mixes both" + description="Session turns recording each baseline model; more than one entry means the router's " + "baseline changed mid-session; these counts do not imply savings coverage" ) diff --git a/tests/proxy_behavior/spend/test_autorouter_session_rollup.py b/tests/proxy_behavior/spend/test_autorouter_session_rollup.py index 77549b527d8..9ef42f5dc7a 100644 --- a/tests/proxy_behavior/spend/test_autorouter_session_rollup.py +++ b/tests/proxy_behavior/spend/test_autorouter_session_rollup.py @@ -6,6 +6,7 @@ tests/test_litellm/proxy/db/test_autorouter_session_rollup.py. """ import asyncio +import json import time import uuid from datetime import datetime, timedelta, timezone @@ -24,6 +25,10 @@ from litellm.proxy.db.autorouter_session_rollup import ( flush_autorouter_turn_transactions, ) from litellm.proxy.db.db_transaction_queue.spend_log_cleanup import SpendLogCleanup +from litellm.proxy.db.autorouter_savings_comparison import ( + HISTORICAL_SESSION_COMPARISONS_SQL, + SessionSavingsComparison, +) pytestmark = pytest.mark.asyncio(loop_scope="session") @@ -91,6 +96,66 @@ async def _row(db, key: str, session_id: str = "s1", router: str = "auto-1") -> return rows[0] +@pytest.mark.parametrize("historical_saved, damaged, user_id, split_sessions, current_classifier", [ + (29.5, None, None, False, 0.2), (29.5, None, "owner", False, 0.2), (0.0, None, None, False, 0.2), + (-3.0, None, None, False, 0.2), (29.5, "missing", None, False, 0.2), (29.5, "cost", None, False, 0.2), + (0.0, "missing", None, False, 0.2), (29.5, None, None, True, 0.2), (29.5, None, None, False, 0.0), +]) +async def test_historical_and_new_savings_compare_matching_costs_and_exclude_unknown_requests( + db: Prisma, historical_saved: float, damaged: str | None, user_id: str | None, split_sessions: bool, + current_classifier: float, +) -> None: + async with db.tx() as tx: + for table in ("LiteLLM_AutoRouterSession", "LiteLLM_AutoRouterUserSession", "LiteLLM_SpendLogs"): + await tx.execute_raw(f'CREATE TEMP TABLE "{table}" (LIKE public."{table}" INCLUDING ALL) ON COMMIT DROP') + for name, spend, saved, classifier, estimated in ( + ("historical", 9.0, historical_saved, 0.1, False), + ("current", 1.0, 0.5, current_classifier, True), + ("unknown", 99.0, 0.0, 3.0, False), + ): + session_id: Final = "s2" if split_sessions and name == "current" else "s1" + await _turn(tx, "key", "model", T0, spend=spend, saved=saved, classifier_cost=classifier, + estimated=estimated, session_id=session_id) + metadata: Final = { + "routing_decision": {"router_model_name": "auto-1", **({"classifier_cost": classifier} if classifier else {})}, + "autorouter_savings": saved if name != "unknown" else None, + **({"autorouter_savings_estimate": { + "version": 3, "status": "estimated" if estimated else "unknown", + }} if name != "historical" else {}), + } + await tx.execute_raw('''INSERT INTO "LiteLLM_SpendLogs" + (request_id,api_key,session_id,model,"user","startTime","endTime",call_type, + spend,prompt_tokens,completion_tokens,status,metadata) + VALUES ($1,'key',$5,'model','owner',$2::timestamp,$2::timestamp,'acompletion', + $3::float8,100,0,'success',$4::jsonb) + ''', name, T0.isoformat(), spend - classifier, json.dumps(metadata), session_id) + await tx.execute_raw('''INSERT INTO "LiteLLM_AutoRouterUserSession" + (user_id,api_key,session_id,router_name,router_type,first_turn_at,last_turn_at,last_model, + turns,total_tokens,spend,saved_spend,savings_estimated_turns,savings_estimated_actual_spend, + savings_estimated_saved_spend) + SELECT 'owner',api_key,session_id,router_name,router_type,first_turn_at,last_turn_at,last_model, + turns,total_tokens,spend,saved_spend,savings_estimated_turns,savings_estimated_actual_spend, + savings_estimated_saved_spend FROM "LiteLLM_AutoRouterSession" + ''') + if damaged == "missing": + await tx.execute_raw('DELETE FROM "LiteLLM_SpendLogs" WHERE request_id = \'historical\'') + elif damaged == "cost": + await tx.execute_raw('UPDATE "LiteLLM_SpendLogs" SET spend = 1 WHERE request_id = \'historical\'') + rows: Final = await tx.query_raw( + HISTORICAL_SESSION_COMPARISONS_SQL, "2026-08-01", "2026-08-02", "key", user_id, None, + ) + comparison: Final = SessionSavingsComparison.model_validate(rows[0]) + assert comparison.saved_spend == historical_saved + 0.5 + assert comparison.complete is (damaged is None) + assert comparison.classifier_cost == (pytest.approx(0.1 + current_classifier) if damaged is None else None) + assert comparison.coverage_fields(historical_saved + 0.5, 4) == {} + assert comparison.coverage_fields(historical_saved + 0.5, 3) == ({ + "savings_estimated_turns": 2, + "savings_estimated_actual_spend": 10.0, + "savings_estimated_saved_spend": historical_saved + 0.5, + } if damaged is None else {}) + + async def test_every_turn_lands_in_exactly_one_bucket(db): key = f"k-{uuid.uuid4()}" await _turn(db, key, "A", T0, ttl=300) diff --git a/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py index ff3d19e8637..385b2b1cc5b 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py @@ -676,6 +676,7 @@ class TestAutoRouterBenchmarks: saved_spend=30.0, savings_estimated_turns=40, savings_estimated_actual_spend=10.0, + savings_estimated_classifier_cost=0.4, savings_estimated_saved_spend=30.0, classifier_cost=0.4, classifier_cost_recorded_turns=40, @@ -701,6 +702,7 @@ class TestAutoRouterBenchmarks: assert totals.avg_tokens_per_session == 1000.0 assert totals.baseline_spend == 40.0 assert totals.saved_pct == 75.0 + assert totals.savings_estimated_classifier_cost == 0.4 assert totals.saved_per_session == 7.5 assert totals.cache.coverage_pct == 95.0 assert totals.cache.hit_rate_pct == pytest.approx(73.7) @@ -720,7 +722,7 @@ class TestAutoRouterBenchmarks: assert totals.classifier_cost == 0.4 @pytest.mark.parametrize("estimated_turns", [0, 4]) - def test_savings_compare_only_the_current_estimated_cohort(self, estimated_turns: int) -> None: + def test_recorded_savings_survive_when_historical_comparison_costs_are_missing(self, estimated_turns: int) -> None: from litellm.proxy.management_endpoints.auto_router_endpoints import _benchmark_totals row: Final = self.ROW.model_copy( @@ -733,10 +735,11 @@ class TestAutoRouterBenchmarks: totals: Final = _benchmark_totals(row) assert totals.spend == 10.0 assert totals.savings_estimated_turns == estimated_turns - assert totals.saved_spend == (-0.5 if estimated_turns else None) - assert totals.baseline_spend == (1.5 if estimated_turns else None) - assert totals.saved_pct == (pytest.approx(-33.3) if estimated_turns else None) - assert totals.saved_per_session is None + assert totals.saved_spend == 30.0 + assert totals.baseline_spend is None + assert totals.savings_estimated_classifier_cost is None + assert totals.saved_pct is None + assert totals.saved_per_session == 7.5 def test_an_empty_window_folds_to_zeros(self): from litellm.proxy.management_endpoints.auto_router_endpoints import ( @@ -765,6 +768,7 @@ class TestAutoRouterBenchmarks: "spend": 0.0, "savings_estimated_turns": 10, "savings_estimated_actual_spend": 0.0, + "savings_estimated_classifier_cost": 0.0, } ) summed = _summed_agg_row([self.ROW, other]) @@ -773,6 +777,9 @@ class TestAutoRouterBenchmarks: assert summed.turns == 50 assert totals.avg_turns_per_session == 10.0 assert totals.spend == 10.0 + assert totals.savings_estimated_classifier_cost == 0.4 + unknown_cost = other.model_copy(update={"savings_estimated_classifier_cost": None}) + assert _benchmark_totals(_summed_agg_row([self.ROW, unknown_cost])).savings_estimated_classifier_cost is None def test_tier_names_stay_scoped_to_the_router_type_that_recorded_them(self): quality = self.ROW.model_copy( @@ -1128,13 +1135,13 @@ class TestAutoRouterSession: "turns": turns, "last_model": "anthropic/claude-sonnet-5", "spend": spend, - "saved_spend": (0.24 if turns == 3 else -0.04) if estimated else None, + "saved_spend": 0.24, "savings_estimated_turns": 3 if estimated else 0, "savings_estimated_actual_spend": 0.14 if estimated else 0.0, "baseline_spend": pytest.approx(0.38) if turns == 3 else None, - "savings_estimated_baseline_spend": pytest.approx(0.38 if turns == 3 else 0.1) if estimated else None, - "baseline_model": "anthropic/claude-opus-5" if estimated else None, - "baseline_models": {"anthropic/claude-opus-5": 3} if estimated else {}, + "savings_estimated_baseline_spend": pytest.approx(0.38) if turns == 3 else None, + "baseline_model": "anthropic/claude-opus-5", + "baseline_models": {"anthropic/claude-opus-5": 3}, } @pytest.mark.asyncio @@ -1168,11 +1175,10 @@ class TestAutoRouterSession: assert response.router_name == "new-auto" @pytest.mark.asyncio - async def test_a_reconfigured_router_keeps_the_label_the_money_was_priced_against( - self, monkeypatch: pytest.MonkeyPatch + @pytest.mark.parametrize("mixed", [False, True]) + async def test_session_preserves_historical_baseline_labels( + self, monkeypatch: pytest.MonkeyPatch, mixed: bool ): - # The proxy's router now prices against a different baseline, but the row's money was priced - # against opus for two of three turns, and the label says so; the full split is on the response. from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_session priced = {"anthropic/claude-opus-5": 2, "anthropic/claude-sonnet-5": 1} @@ -1183,14 +1189,15 @@ class TestAutoRouterSession: **self.ROW, "api_key": ADMIN.api_key, "session_id": "s", - "baseline_models": {"old-baseline": 100}, + "baseline_models": {"old-baseline": 100, **({"unknown-baseline": 200} if mixed else {})}, + "savings_estimated_turns": 1, "savings_estimated_baseline_models": priced, } ], ) response = await get_auto_router_session(user_api_key_dict=ADMIN, session_id="s") - assert response.baseline_model == "anthropic/claude-opus-5" - assert response.baseline_models == priced + assert response.baseline_model == (None if mixed else "old-baseline") + assert response.baseline_models == {"old-baseline": 100, **({"unknown-baseline": 200} if mixed else {})} @pytest.mark.asyncio async def test_an_oversized_client_session_id_is_bounded_like_the_writer_bounded_it( diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.test.tsx index a144630cdd0..5e8533c8b82 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.test.tsx @@ -158,35 +158,49 @@ describe("AutoRouterBenchmarksTab", () => { }); it.each([ - { estimatedTurns: 0, saved: null, pct: null }, - { estimatedTurns: 10, saved: -0.5, pct: -33.3 }, - { estimatedTurns: 10, saved: 0, pct: 0 }, - ])("preserves costs for $estimatedTurns estimated turns with savings $saved", ({ estimatedTurns, saved, pct }) => { - const cohort = { + { estimatedTurns: 0, actual: 0, saved: null, pct: null }, + { estimatedTurns: 0, actual: 0, saved: 30, pct: null }, + { estimatedTurns: 10, actual: 2, saved: -0.5, pct: -33.3 }, + { estimatedTurns: 10, actual: 2, saved: 0, pct: 0 }, + { estimatedTurns: 40, actual: 10, saved: 30, pct: 75 }, + ])("compares matching old and new requests with savings $saved", ({ estimatedTurns, actual, saved, pct }) => { + const comparison = { + spend: actual + 99, savings_estimated_turns: estimatedTurns, - savings_estimated_actual_spend: estimatedTurns ? 2 : 0, + savings_estimated_actual_spend: actual, + savings_estimated_classifier_cost: 0.1, saved_spend: saved, - baseline_spend: estimatedTurns ? 2 + (saved ?? 0) : null, + baseline_spend: estimatedTurns ? actual + (saved ?? 0) : null, saved_pct: pct, saved_per_session: null, }; - const partial = totals(cohort); - mockHook({ data: response([], partial) }); + mockHook({ + data: response([], totals(comparison)), + }); renderTab(); - expect(screen.getByText("Estimated savings on covered turns")).toBeInTheDocument(); - expect(screen.getByText(`${estimatedTurns} of 3,073 turns estimated`)).toBeInTheDocument(); - expect(screen.getByText("$359.86")).toBeInTheDocument(); - expect(screen.getByText("Actual spend on covered turns")).toBeInTheDocument(); - expect(screen.getByText("Estimated baseline spend on covered turns")).toBeInTheDocument(); - expect(screen.getAllByText("Unavailable")).toHaveLength(estimatedTurns ? 1 : 3); - if (saved === 0) { - expect(screen.getByText("0%")).toBeInTheDocument(); - expect(screen.getAllByText("$2.00")).toHaveLength(2); - } else if (estimatedTurns) { - expect(screen.getByText("-$0.5000")).toBeInTheDocument(); - expect(screen.getByText("+33%")).toBeInTheDocument(); - } else { - expect(screen.queryByText("+0%")).not.toBeInTheDocument(); + expect(screen.getByText("Total estimated savings")).toBeInTheDocument(); + expect(screen.getAllByRole("definition").map((row) => row.textContent)).toEqual( + estimatedTurns + ? [ + `$${actual.toFixed(2)}`, + `$${(actual - 0.1).toFixed(2)}`, + "$0.1000", + `$${(actual + (saved ?? 0)).toFixed(2)}`, + ] + : ["Unavailable", "Unavailable", "Unavailable", "Unavailable"], + ); + expect(screen.queryByText("Actual spend on covered turns")).not.toBeInTheDocument(); + expect(screen.getByLabelText("question-circle")).toBeInTheDocument(); + if (estimatedTurns) { + expect(screen.getByText(`Savings based on ${estimatedTurns} of 3,073 requests`)).toBeInTheDocument(); + const sign = pct && pct > 0 ? "-" : "+"; + const badge = pct === 0 ? "0%" : `${sign}${Math.abs(pct ?? 0).toFixed(0)}%`; + expect(screen.getByText(badge)).toBeInTheDocument(); + } else if (saved != null) { + expect(screen.getByText("$30.00")).toBeInTheDocument(); + expect( + screen.getByText("Historical savings are included. Matching cost details are unavailable."), + ).toBeInTheDocument(); } }); @@ -216,7 +230,7 @@ describe("AutoRouterBenchmarksTab", () => { expect(screen.getByText("-86%")).toBeInTheDocument(); expect(screen.getByText("Actual auto-router spend")).toBeInTheDocument(); expect(screen.getByText("$359.86")).toBeInTheDocument(); - expect(screen.getByText("Estimated spend at highest-tier model")).toBeInTheDocument(); + expect(screen.getByText("Estimated baseline spend")).toBeInTheDocument(); expect(screen.getByText("$2,534.45")).toBeInTheDocument(); expect(screen.getByText("32.7")).toBeInTheDocument(); expect(screen.getByText("2.1h")).toBeInTheDocument(); @@ -242,17 +256,20 @@ describe("AutoRouterBenchmarksTab", () => { expect(screen.getAllByText("$10,126.28").length).toBeGreaterThan(0); }); - it.each([null, undefined])("keeps totals when the classification breakdown is %s", (classifier_cost) => { - const stats = totals({ classifier_cost }); - mockHook({ data: response([group(stats)], stats) }); - renderTab(); + it.each([null, undefined])( + "keeps eligible totals when the classification breakdown is %s", + (savings_estimated_classifier_cost) => { + const stats = totals({ savings_estimated_turns: 30, savings_estimated_classifier_cost }); + mockHook({ data: response([group(stats)], stats) }); + renderTab(); - expect(screen.getAllByText("Unavailable")).toHaveLength(2); - expect(screen.queryByText(/\/ 1K turns/)).not.toBeInTheDocument(); - expect(screen.getByText("$359.86")).toBeInTheDocument(); - expect(screen.getByText("$2,174.59")).toBeInTheDocument(); - expect(screen.getByText(/some usage predates classification-cost tracking/)).toBeInTheDocument(); - }); + expect(screen.getAllByText("Unavailable")).toHaveLength(2); + expect(screen.queryByText(/\/ 1K turns/)).not.toBeInTheDocument(); + expect(screen.getByText("$359.86")).toBeInTheDocument(); + expect(screen.getByText("$2,174.59")).toBeInTheDocument(); + expect(screen.getByText(/some usage predates classification-cost tracking/)).toBeInTheDocument(); + }, + ); it("pairs the savings with the session count it was earned over, in its own tile", () => { mockHook({ data: response([group(), group({ router_name: "gpt-auto" })]) }); @@ -275,7 +292,7 @@ describe("AutoRouterBenchmarksTab", () => { "Actual auto-router spend", "LLM spend", "Classification cost($2.00 / 1K turns)", - "Estimated spend at highest-tier model", + "Estimated baseline spend", ]); expect(values).toEqual(["$359.86", "$353.71", "$6.15", "$2,534.45"]); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.tsx index 063598bd46e..f532c2e4650 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutoRouterBenchmarksTab.tsx @@ -11,7 +11,7 @@ import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@ import { Separator } from "@/components/ui/separator"; import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table"; import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"; -import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip"; +import { SimpleTooltip, Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip"; import { ApiError } from "@/lib/http/client"; import { formatNumberWithCommas } from "@/utils/dataUtils"; @@ -52,15 +52,17 @@ const Metric: React.FC<{ label: string; value: string; hint?: string }> = ({ lab ); -const SpendRow: React.FC<{ label: string; value: string; hint?: string; subdued?: boolean }> = ({ +const SpendRow: React.FC<{ label: string; value: string; hint?: string; subdued?: boolean; tooltip?: string }> = ({ label, value, hint, subdued, + tooltip, }) => (
    {label} + {tooltip && } {hint && {hint}}
    = ({ view }) => { const stats = view.stats; - const cheaper = stats.saved_spend != null && stats.saved_spend >= 0; + const cheaper = stats.saved_pct != null && stats.saved_pct >= 0; const completeCoverage = stats.savings_estimated_turns === stats.turns; + const coveredClassifierCost = + stats.savings_estimated_classifier_cost ?? (completeCoverage ? stats.classifier_cost : null); + const classifierCost = stats.baseline_spend == null ? null : coveredClassifierCost; return (

    - {completeCoverage ? "Total estimated savings" : "Estimated savings on covered turns"} + Total estimated savings

    @@ -91,53 +96,57 @@ const HeroCard: React.FC<{ view: BenchmarkView }> = ({ view }) => { variant="secondary" className={`h-6 px-2.5 text-sm ${cheaper ? "bg-success/10 text-success" : "bg-destructive/10 text-destructive"}`} > - {stats.saved_spend !== 0 && (cheaper ? "-" : "+")} + {stats.saved_pct !== 0 && (cheaper ? "-" : "+")} {Math.abs(stats.saved_pct).toFixed(0)}% )}

    -

    - {stats.savings_estimated_turns.toLocaleString()} of {stats.turns.toLocaleString()} turns estimated -

    - {!completeCoverage && ( + {stats.baseline_spend != null && !completeCoverage && (

    - Turns without a current estimate are excluded, including older estimates. + Savings based on {stats.savings_estimated_turns.toLocaleString()} of {stats.turns.toLocaleString()}{" "} + requests +

    + )} + {stats.saved_spend != null && stats.baseline_spend == null && ( +

    + Historical savings are included. Matching cost details are unavailable.

    )}
    - +
    - {stats.classifier_cost == null && ( + {stats.baseline_spend != null && classifierCost == null && (

    Breakdown unavailable because some usage predates classification-cost tracking.

    )} - {!completeCoverage && ( - - )}
    @@ -307,12 +316,11 @@ const BenchmarksBody: React.FC = ({ isPending, error, data,

    - Compares covered turns with the estimated cost of using the router's highest-tier baseline model. Estimates - use registered requests since tracking began, matching cache prefixes and expiry, and the actual response - length. Total actual spend includes every turn; savings and baseline spend include only turns with a current - estimate, including turns with zero savings. Savings are net of recorded LLM classification cost. Classification - cost per 1K turns is averaged over all auto-router turns, including those that skip classification. The range - counts whole sessions that overlap it, so totals can differ from savings views that group usage by UTC day. + Savings, actual spend, and baseline compare the same historical and newer requests with recorded estimates, + including zero or negative savings. Requests without estimates are excluded. Savings are net of recorded LLM + classification cost. If historical cost details are unavailable, recorded savings remain visible without a + baseline or percentage. The range counts whole sessions that overlap it, so totals can differ from savings views + that group usage by UTC day.

    diff --git a/ui/litellm-dashboard/src/components/templates/KeyAutoRouterUsageTab.integration.test.tsx b/ui/litellm-dashboard/src/components/templates/KeyAutoRouterUsageTab.integration.test.tsx index 5429c688e17..807bd2f4f15 100644 --- a/ui/litellm-dashboard/src/components/templates/KeyAutoRouterUsageTab.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/templates/KeyAutoRouterUsageTab.integration.test.tsx @@ -92,7 +92,7 @@ describe("KeyAutoRouterUsageTab", () => { expect(screen.getByText("Classification cost")).toBeInTheDocument(); expect(screen.getByText("$0.2500")).toBeInTheDocument(); expect(screen.getByText("($62.50 / 1K turns)")).toBeInTheDocument(); - expect(screen.getByText("Estimated spend at highest-tier model")).toBeInTheDocument(); + expect(screen.getByText("Estimated baseline spend")).toBeInTheDocument(); expect(screen.getByText("$10.00")).toBeInTheDocument(); expect(screen.getByText("Auto-router prompt caching")).toBeInTheDocument(); expect(screen.getAllByText("50.0%").length).toBeGreaterThan(0); diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index dc0365a14b1..ecfabf33e87 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -1263,8 +1263,8 @@ export interface paths { * @description Benchmarks for the auto-router dashboard: session shape, savings against the configured * baseline, and prompt-caching behaviour bucketed by what the router did. * - * Reads session rollups folded once per request at spend-write time, so this endpoint - * never scans LiteLLM_SpendLogs. A user filter selects only turns attributed to that + * Reads session rollups folded once per request at spend-write time, with bounded + * retained-log recovery for historical comparisons. A user filter selects only turns attributed to that * internal user when written; older key-only history remains outside user views. A session * is in the window when it overlaps it: its last turn is on or after start_date and its first turn is on or before * end_date. Overall hit rate is over telemetry-bearing turns; each bucket's hit rate is @@ -24899,17 +24899,17 @@ export interface components { router_type: string; /** * Saved Pct - * @description Covered savings over covered baseline spend, as a percentage + * @description Total recorded savings over the matching historical and current baseline; null when costs are unavailable */ saved_pct: number | null; /** * Saved Per Session - * @description Average session savings; unavailable unless every turn is covered + * @description Recorded savings per session, including historical estimates */ saved_per_session: number | null; /** * Saved Spend - * @description Signed savings for covered turns only; null when traffic has no current estimates + * @description Recorded historical savings plus newer estimates; null when traffic has no recorded savings estimates */ saved_spend: number | null; /** @@ -24917,9 +24917,14 @@ export interface components { * @description Actual spend, including classifier cost, for covered turns only */ savings_estimated_actual_spend: number; + /** + * Savings Estimated Classifier Cost + * @description Classifier cost included in the matching historical and newer savings comparison; null when classification costs for those requests are unavailable + */ + savings_estimated_classifier_cost?: number | null; /** * Savings Estimated Turns - * @description Turns covered by the current savings estimator; legacy estimates are excluded + * @description Requests with a matching savings comparison, including historical recorded estimates */ savings_estimated_turns: number; /** Sessions */ @@ -24963,17 +24968,17 @@ export interface components { classifier_cost: number | null; /** * Saved Pct - * @description Covered savings over covered baseline spend, as a percentage + * @description Total recorded savings over the matching historical and current baseline; null when costs are unavailable */ saved_pct: number | null; /** * Saved Per Session - * @description Average session savings; unavailable unless every turn is covered + * @description Recorded savings per session, including historical estimates */ saved_per_session: number | null; /** * Saved Spend - * @description Signed savings for covered turns only; null when traffic has no current estimates + * @description Recorded historical savings plus newer estimates; null when traffic has no recorded savings estimates */ saved_spend: number | null; /** @@ -24981,9 +24986,14 @@ export interface components { * @description Actual spend, including classifier cost, for covered turns only */ savings_estimated_actual_spend: number; + /** + * Savings Estimated Classifier Cost + * @description Classifier cost included in the matching historical and newer savings comparison; null when classification costs for those requests are unavailable + */ + savings_estimated_classifier_cost?: number | null; /** * Savings Estimated Turns - * @description Turns covered by the current savings estimator; legacy estimates are excluded + * @description Requests with a matching savings comparison, including historical recorded estimates */ savings_estimated_turns: number; /** Sessions */ @@ -25260,12 +25270,12 @@ export interface components { AutoRouterSessionResponse: { /** * Baseline Model - * @description The savings baseline most covered turns were priced against, recorded turn by turn, so it still names the counterfactual after the router is reconfigured or removed. None when no turn recorded one: rows from before the baseline was recorded, and adaptive and quality routers, which derive no baseline and so report no savings + * @description The savings baseline recorded by most session turns, including historical turns, recorded turn by turn, so it still names the counterfactual after the router is reconfigured or removed. None when no turn recorded one: rows from before the baseline was recorded, and adaptive and quality routers, which derive no baseline and so report no savings */ baseline_model: string | null; /** * Baseline Models - * @description Covered turns priced against each baseline model; more than one entry means the router's baseline changed mid-session and baseline_spend mixes both + * @description Session turns recording each baseline model; more than one entry means the router's baseline changed mid-session; these counts do not imply savings coverage */ baseline_models: { [key: string]: number; @@ -25292,7 +25302,7 @@ export interface components { router_type: string; /** * Saved Spend - * @description Estimated savings for covered turns only, net of classifier cost + * @description Recorded historical savings plus newer estimates, net of classifier cost */ saved_spend: number | null; /** @@ -25307,7 +25317,7 @@ export interface components { savings_estimated_baseline_spend: number | null; /** * Savings Estimated Turns - * @description Turns covered by the current savings estimator; legacy estimates are excluded + * @description Requests with a matching savings comparison, including historical recorded estimates */ savings_estimated_turns: number; /** Session Id */ From 1bfa3d4fa62d49d37dbf381bacc8cb5e569ef955 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 12:45:22 -0700 Subject: [PATCH 27/41] fix(model-prices): align Azure, Bedrock, Copilot, Gemini, Groq, OpenAI and OpenRouter entries with official docs (#43598) * fix(model-prices): correct azure/eu/gpt-6-astra to Data Zone rates Co-authored-by: rain <1504569896@qq.com> Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(model-prices): align groq, gemini and openai entries with official docs Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(model-prices): roll in verified Vertex, Gemini, OpenRouter and Azure AI registry fixes Absorbs the fields from #43609, #43666, #43671 and #43644 that match the provider's own docs or price API today, and adds a cost test for the azure/eu/gpt-6-astra Data Zone tiers Co-authored-by: bunnysayzz Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(model-prices): add Copilot, Bedrock Kimi K3, Gemini Robotics and OpenRouter values from official sources Co-authored-by: Michal Formanek Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: rain <1504569896@qq.com> Co-authored-by: bunnysayzz Co-authored-by: Michal Formanek --- ...odel_prices_and_context_window_backup.json | 266 +++++++++++------- model_prices_and_context_window.json | 266 +++++++++++------- .../llm_cost_calc/test_llm_cost_calc_utils.py | 29 ++ 3 files changed, 363 insertions(+), 198 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index c147f4bf94b..81c0c045a30 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -12260,7 +12260,7 @@ "azure_ai/deepseek-v3.2": { "input_cost_per_token": 5.8e-07, "litellm_provider": "azure_ai", - "max_input_tokens": 163840, + "max_input_tokens": 128000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -12368,7 +12368,7 @@ "azure_ai/grok-4": { "input_cost_per_token": 3e-06, "litellm_provider": "azure_ai", - "max_input_tokens": 131072, + "max_input_tokens": 262000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", @@ -12490,7 +12490,7 @@ "azure_ai/grok-code-fast-1": { "input_cost_per_token": 2e-07, "litellm_provider": "azure_ai", - "max_input_tokens": 131072, + "max_input_tokens": 256000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", @@ -27518,7 +27518,8 @@ "search_context_size_high": 0.035 }, "gemini_native_audio": true, - "input_cost_per_image_token": 3e-06 + "input_cost_per_image_token": 3e-06, + "input_cost_per_video_token": 3e-06 }, "gemini-live-2.5-flash-preview-native-audio-09-2025": { "input_cost_per_audio_token": 3e-06, @@ -28418,7 +28419,7 @@ "max_output_tokens": 65536, "max_tokens": 65536, "mode": "chat", - "output_cost_per_reasoning_token": 1e-05, + "output_cost_per_reasoning_token": 5e-06, "output_cost_per_token": 5e-06, "output_cost_per_token_batches": 2.5e-06, "search_context_cost_per_query": { @@ -29015,7 +29016,7 @@ "image" ], "supports_function_calling": false, - "supports_prompt_caching": true, + "supports_prompt_caching": false, "supports_response_schema": false, "supports_system_messages": true, "supports_vision": true, @@ -29026,7 +29027,7 @@ "search_context_size_high": 0.014 }, "web_search_billing_unit": "per_query", - "supports_reasoning": false + "supports_reasoning": true }, "gemini/nano-banana-pro-preview": { "input_cost_per_image": 0.0011, @@ -29104,8 +29105,8 @@ "image" ], "supports_function_calling": false, - "supports_prompt_caching": true, - "supports_reasoning": false, + "supports_prompt_caching": false, + "supports_reasoning": true, "supports_response_schema": false, "supports_system_messages": true, "supports_vision": true, @@ -29115,7 +29116,8 @@ "search_context_size_medium": 0.014, "search_context_size_high": 0.014 }, - "web_search_billing_unit": "per_query" + "web_search_billing_unit": "per_query", + "supports_pdf_input": true }, "gemini/gemini-3.1-flash-lite-image": { "input_cost_per_image": 0.00028, @@ -29146,8 +29148,9 @@ "image" ], "supports_function_calling": false, + "supports_pdf_input": true, "supports_prompt_caching": false, - "supports_reasoning": false, + "supports_reasoning": true, "supports_response_schema": false, "supports_system_messages": true, "supports_vision": true, @@ -29159,17 +29162,15 @@ "input_cost_per_token": 2e-06, "input_cost_per_token_batches": 1e-06, "litellm_provider": "gemini", - "max_input_tokens": 65536, - "max_output_tokens": 32768, - "max_tokens": 32768, - "mode": "image_generation", - "output_cost_per_image": 0.134, - "output_cost_per_image_token": 0.00012, + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", "output_cost_per_token": 1.2e-05, "rpm": 1000, "tpm": 4000000, "output_cost_per_token_batches": 6e-06, - "source": "https://ai.google.dev/gemini-api/docs/pricing", + "source": "https://ai.google.dev/gemini-api/docs/models/deep-research-pro-preview-12-2025", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -29177,11 +29178,12 @@ ], "supported_modalities": [ "text", - "image" + "image", + "audio", + "video" ], "supported_output_modalities": [ - "text", - "image" + "text" ], "supports_function_calling": false, "supports_prompt_caching": true, @@ -29193,7 +29195,8 @@ "search_context_size_low": 0.035, "search_context_size_medium": 0.035, "search_context_size_high": 0.035 - } + }, + "supports_pdf_input": true }, "gemini/gemini-2.5-flash-lite": { "cache_read_input_audio_token_cost": 3e-08, @@ -30785,11 +30788,15 @@ ] }, "github_copilot/claude-haiku-4.5": { + "cache_creation_input_token_cost": 1.25e-06, + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token": 1e-06, "litellm_provider": "github_copilot", "max_input_tokens": 128000, "max_output_tokens": 16000, "max_tokens": 16000, "mode": "chat", + "output_cost_per_token": 5e-06, "supported_endpoints": [ "/v1/chat/completions" ], @@ -30838,11 +30845,15 @@ "supports_vision": true }, "github_copilot/claude-sonnet-4": { + "cache_creation_input_token_cost": 3.75e-06, + "cache_read_input_token_cost": 3e-07, + "input_cost_per_token": 3e-06, "litellm_provider": "github_copilot", "max_input_tokens": 128000, "max_output_tokens": 16000, "max_tokens": 16000, "mode": "chat", + "output_cost_per_token": 1.5e-05, "supported_endpoints": [ "/v1/chat/completions" ], @@ -31003,11 +31014,14 @@ "supports_vision": true }, "github_copilot/gpt-5-mini": { + "cache_read_input_token_cost": 2.5e-08, + "input_cost_per_token": 2.5e-07, "litellm_provider": "github_copilot", "max_input_tokens": 128000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", + "output_cost_per_token": 2e-06, "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, @@ -31058,11 +31072,14 @@ "supports_vision": true }, "github_copilot/gpt-5.3-codex": { + "cache_read_input_token_cost": 1.75e-07, + "input_cost_per_token": 1.75e-06, "litellm_provider": "github_copilot", "max_input_tokens": 128000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", + "output_cost_per_token": 1.4e-05, "supported_endpoints": [ "/v1/responses" ], @@ -32745,6 +32762,25 @@ "audio" ] }, + "gpt-4o-mini-tts-2025-03-20": { + "input_cost_per_token": 6e-07, + "litellm_provider": "openai", + "mode": "audio_speech", + "output_cost_per_audio_token": 1.2e-05, + "output_cost_per_second": 0.00025, + "output_cost_per_token": 1e-05, + "source": "https://developers.openai.com/api/docs/models/gpt-4o-mini-tts", + "supported_endpoints": [ + "/v1/audio/speech" + ], + "supported_modalities": [ + "text", + "audio" + ], + "supported_output_modalities": [ + "audio" + ] + }, "gpt-4o-search-preview": { "cache_read_input_token_cost": 1.25e-06, "input_cost_per_token": 2.5e-06, @@ -42065,8 +42101,8 @@ "input_cost_per_token_cache_hit": 2e-08, "litellm_provider": "openrouter", "max_input_tokens": 163840, - "max_output_tokens": 65536, - "max_tokens": 65536, + "max_output_tokens": 147456, + "max_tokens": 147456, "mode": "chat", "output_cost_per_token": 4.1e-07, "source": "https://openrouter.ai/api/v1/models", @@ -42146,14 +42182,14 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4.1-flash": { - "cache_read_input_token_cost": 3.135e-08, - "input_cost_per_token": 3.483e-08, + "cache_read_input_token_cost": 6e-09, + "input_cost_per_token": 3e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, - "max_output_tokens": 384000, - "max_tokens": 384000, + "max_output_tokens": 943718, + "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 6e-07, + "output_cost_per_token": 1.2e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -42173,7 +42209,7 @@ "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 3.5e-06, + "output_cost_per_token": 4.2e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -42527,13 +42563,13 @@ "max_output_tokens": 8000 }, "openrouter/minimax/minimax-m2": { - "input_cost_per_token": 2.55e-07, + "input_cost_per_token": 3e-07, "litellm_provider": "openrouter", "max_input_tokens": 204800, - "max_output_tokens": 131072, - "max_tokens": 131072, + "max_output_tokens": 176947, + "max_tokens": 176947, "mode": "chat", - "output_cost_per_token": 1.02e-06, + "output_cost_per_token": 1.2e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -42739,14 +42775,14 @@ "supports_web_search": false }, "openrouter/nvidia/nemotron-3.5-lightning": { - "cache_read_input_token_cost": 4e-08, - "input_cost_per_token": 8e-08, + "cache_read_input_token_cost": 3e-08, + "input_cost_per_token": 6e-08, "litellm_provider": "openrouter", "max_input_tokens": 262144, - "max_output_tokens": 235929, - "max_tokens": 235929, + "max_output_tokens": 32768, + "max_tokens": 32768, "mode": "chat", - "output_cost_per_token": 2e-07, + "output_cost_per_token": 1.6e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -43199,14 +43235,14 @@ "supports_web_search": true }, "openrouter/openai/gpt-5.6-sol-pro": { - "input_cost_per_token": 2e-06, - "output_cost_per_token": 1e-05, - "cache_read_input_token_cost": 2e-07, - "cache_creation_input_token_cost": 2.5e-06, - "cache_creation_input_token_cost_above_272k_tokens": 5e-06, - "input_cost_per_token_above_272k_tokens": 4e-06, - "output_cost_per_token_above_272k_tokens": 1.5e-05, - "cache_read_input_token_cost_above_272k_tokens": 4e-07, + "input_cost_per_token": 4e-06, + "output_cost_per_token": 2e-05, + "cache_read_input_token_cost": 4e-07, + "cache_creation_input_token_cost": 5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 1e-05, + "input_cost_per_token_above_272k_tokens": 8e-06, + "output_cost_per_token_above_272k_tokens": 3e-05, + "cache_read_input_token_cost_above_272k_tokens": 8e-07, "litellm_provider": "openrouter", "max_input_tokens": 1050000, "max_output_tokens": 128000, @@ -43264,6 +43300,7 @@ "supports_web_search": false }, "openrouter/openai/gpt-oss-20b": { + "cache_read_input_token_cost": 9e-09, "input_cost_per_token": 1.8e-08, "litellm_provider": "openrouter", "max_input_tokens": 131072, @@ -49891,7 +49928,8 @@ "supports_tool_choice": true, "supports_vision": true, "prompt_cache_min_tokens": 1024, - "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", + "cache_creation_input_token_cost_batches": 1.88e-06 }, "vertex_ai/claude-sonnet-5": { "deprecation_date": "2026-12-24", @@ -49999,7 +50037,8 @@ "supports_vision": true, "supports_native_streaming": true, "prompt_cache_min_tokens": 1024, - "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", + "cache_creation_input_token_cost_batches": 1.88e-06 }, "vertex_ai/mistralai/codestral-2@001": { "input_cost_per_token": 3e-07, @@ -60514,6 +60553,9 @@ "input_cost_per_audio_token": 1e-06, "input_cost_per_token": 1e-06, "litellm_provider": "gemini", + "max_input_tokens": 131072, + "max_output_tokens": 65536, + "max_tokens": 65536, "mode": "chat", "output_cost_per_token": 5e-06, "search_context_cost_per_query": { @@ -60536,6 +60578,7 @@ ], "supports_audio_input": true, "supports_function_calling": true, + "supports_reasoning": true, "supports_video_input": true, "supports_vision": true, "supports_web_search": true, @@ -64006,7 +64049,7 @@ "groq/qwen/qwen3.8-27b": { "input_cost_per_token": 8e-07, "litellm_provider": "groq", - "max_input_tokens": 131042, + "max_input_tokens": 131072, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", @@ -67212,13 +67255,13 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-flash-vision-exp": { - "input_cost_per_token": 4.4e-07, - "output_cost_per_token": 1.32e-06, - "cache_read_input_token_cost": 1.4e-08, + "input_cost_per_token": 2.156e-07, + "output_cost_per_token": 6.468e-07, + "cache_read_input_token_cost": 6.86e-09, "litellm_provider": "openrouter", "max_input_tokens": 1048576, - "max_output_tokens": 943718, - "max_tokens": 943718, + "max_output_tokens": 262144, + "max_tokens": 262144, "mode": "chat", "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, @@ -67232,13 +67275,13 @@ "supports_web_search": false }, "openrouter/z-ai/glm-5.3": { - "input_cost_per_token": 3.556e-07, - "output_cost_per_token": 2.574e-06, - "cache_read_input_token_cost": 6.604e-08, + "input_cost_per_token": 1.4e-06, + "output_cost_per_token": 4.4e-06, + "cache_read_input_token_cost": 2.6e-07, "litellm_provider": "openrouter", "max_input_tokens": 1310720, - "max_output_tokens": 943718, - "max_tokens": 943718, + "max_output_tokens": 943717, + "max_tokens": 943717, "mode": "chat", "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, @@ -67370,7 +67413,7 @@ }, "openrouter/deepseek/deepseek-v4-flash-0731": { "cache_read_input_token_cost": 1.6e-08, - "input_cost_per_token": 2.1e-08, + "input_cost_per_token": 1.8e-08, "litellm_provider": "openrouter", "max_input_tokens": 1310720, "max_output_tokens": 943718, @@ -67964,9 +68007,9 @@ "supports_web_search": false }, "openrouter/moonshotai/kimi-k2.6": { - "input_cost_per_token": 9.5e-07, - "output_cost_per_token": 4e-06, - "cache_read_input_token_cost": 1.6e-07, + "input_cost_per_token": 6.5e-07, + "output_cost_per_token": 3.41e-06, + "cache_read_input_token_cost": 1.5e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 235929, @@ -69207,12 +69250,12 @@ }, "openrouter/qwen/qwen3-30b-a3b": { "deprecation_date": "2026-10-09", - "input_cost_per_token": 1.3e-07, - "output_cost_per_token": 5.2e-07, + "input_cost_per_token": 1.2e-07, + "output_cost_per_token": 5e-07, "litellm_provider": "openrouter", "max_input_tokens": 131072, - "max_output_tokens": 8192, - "max_tokens": 8192, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, @@ -69247,8 +69290,8 @@ }, "openrouter/qwen/qwen3-14b": { "deprecation_date": "2026-10-09", - "input_cost_per_token": 2.275e-07, - "output_cost_per_token": 9.1e-07, + "input_cost_per_token": 1.2e-07, + "output_cost_per_token": 2.4e-07, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 16384, @@ -69878,7 +69921,9 @@ "vertex_ai/gemini-2.5-flash-native-audio": { "deprecation_date": "2026-12-13", "input_cost_per_audio_token": 3e-06, + "input_cost_per_image_token": 3e-06, "input_cost_per_token": 5e-07, + "input_cost_per_video_token": 3e-06, "litellm_provider": "vertex_ai", "mode": "realtime", "output_cost_per_audio_token": 1.2e-05, @@ -70678,17 +70723,17 @@ }, "azure/eu/gpt-6-astra": { "deprecation_date": "2028-01-11", - "cache_creation_input_token_cost": 1.375e-05, - "cache_creation_input_token_cost_above_272k_tokens": 2.75e-05, - "cache_read_input_token_cost": 1.1e-06, - "cache_read_input_token_cost_above_272k_tokens": 2.2e-06, - "input_cost_per_token": 1.1e-05, - "input_cost_per_token_above_272k_tokens": 2.2e-05, + "cache_creation_input_token_cost": 1.5e-05, + "cache_creation_input_token_cost_above_272k_tokens": 3e-05, + "cache_read_input_token_cost": 1.2e-06, + "cache_read_input_token_cost_above_272k_tokens": 2.4e-06, + "input_cost_per_token": 1.2e-05, + "input_cost_per_token_above_272k_tokens": 2.4e-05, "litellm_provider": "azure", "mode": "chat", - "output_cost_per_token": 5.5e-05, - "output_cost_per_token_above_272k_tokens": 8.25e-05, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "output_cost_per_token": 6e-05, + "output_cost_per_token_above_272k_tokens": 9e-05, + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'swedencentral'%20and%20priceType%20eq%20'Consumption'", "supports_reasoning": true }, "azure/eu/gpt-6-luna": { @@ -72821,17 +72866,17 @@ "supports_web_search": true }, "openrouter/~x-ai/grok-latest": { - "cache_read_input_token_cost": 4e-07, - "cache_read_input_token_cost_above_200k_tokens": 8e-07, - "input_cost_per_token": 1.6e-06, - "input_cost_per_token_above_200k_tokens": 3.2e-06, + "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_above_200k_tokens": 1e-06, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_200k_tokens": 4e-06, "litellm_provider": "openrouter", "max_input_tokens": 500000, "max_output_tokens": 450000, "max_tokens": 450000, "mode": "chat", - "output_cost_per_token": 4.8e-06, - "output_cost_per_token_above_200k_tokens": 9.6e-06, + "output_cost_per_token": 6e-06, + "output_cost_per_token_above_200k_tokens": 1.2e-05, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -74106,7 +74151,7 @@ "cache_read_input_token_cost": 4.2e-09, "input_cost_per_token": 2.1e-08, "litellm_provider": "openrouter", - "max_input_tokens": 131072, + "max_input_tokens": 262144, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", @@ -76134,7 +76179,7 @@ "cache_read_input_token_cost": 1.7e-07, "input_cost_per_token": 1e-06, "litellm_provider": "openrouter", - "max_input_tokens": 1048576, + "max_input_tokens": 524288, "max_output_tokens": 471859, "max_tokens": 471859, "mode": "chat", @@ -76154,7 +76199,7 @@ "cache_read_input_token_cost": 1e-07, "input_cost_per_token": 4.5e-07, "litellm_provider": "openrouter", - "max_input_tokens": 1048576, + "max_input_tokens": 524288, "max_output_tokens": 262144, "max_tokens": 262144, "mode": "chat", @@ -76371,6 +76416,7 @@ "supports_web_search": false }, "openrouter/prism-ml/ternary-bonsai-2-27b": { + "cache_read_input_token_cost": 3.75e-08, "input_cost_per_token": 7.5e-08, "litellm_provider": "openrouter", "max_input_tokens": 262144, @@ -76431,17 +76477,17 @@ "supports_web_search": false }, "openrouter/x-ai/grok-4.7": { - "cache_read_input_token_cost": 4e-07, - "cache_read_input_token_cost_above_200k_tokens": 8e-07, - "input_cost_per_token": 1.6e-06, - "input_cost_per_token_above_200k_tokens": 3.2e-06, + "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_above_200k_tokens": 1e-06, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_200k_tokens": 4e-06, "litellm_provider": "openrouter", "max_input_tokens": 500000, "max_output_tokens": 450000, "max_tokens": 450000, "mode": "chat", - "output_cost_per_token": 4.8e-06, - "output_cost_per_token_above_200k_tokens": 9.6e-06, + "output_cost_per_token": 6e-06, + "output_cost_per_token_above_200k_tokens": 1.2e-05, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -76454,16 +76500,16 @@ "supports_web_search": true }, "moonshotai.kimi-k3": { - "cache_creation_input_token_cost": 3.75e-06, - "cache_read_input_token_cost": 3e-07, - "input_cost_per_token": 3e-06, + "cache_creation_input_token_cost": 4.125e-06, + "cache_read_input_token_cost": 3.3e-07, + "input_cost_per_token": 3.3e-06, "litellm_provider": "bedrock_converse", "max_input_tokens": 1000000, "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 1.5e-05, - "source": "https://aws.amazon.com/bedrock/pricing/", + "output_cost_per_token": 1.65e-05, + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrock/current/us-east-1/index.json", "supports_audio_input": false, "supports_function_calling": true, "supports_prompt_caching": true, @@ -79026,6 +79072,28 @@ "supports_tool_choice": true, "supports_vision": true }, + "openrouter/anthropic/claude-sonnet-5.5:batch": { + "cache_creation_input_token_cost": 1.25e-06, + "cache_creation_input_token_cost_above_1hr": 2e-06, + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token": 1e-06, + "litellm_provider": "openrouter", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-06, + "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": false, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, "baseten/deepseek-ai/DeepSeek-V4.1-Flash-Fast": { "cache_read_input_token_cost": 1.4e-07, "input_cost_per_token": 6e-07, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index c147f4bf94b..81c0c045a30 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -12260,7 +12260,7 @@ "azure_ai/deepseek-v3.2": { "input_cost_per_token": 5.8e-07, "litellm_provider": "azure_ai", - "max_input_tokens": 163840, + "max_input_tokens": 128000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -12368,7 +12368,7 @@ "azure_ai/grok-4": { "input_cost_per_token": 3e-06, "litellm_provider": "azure_ai", - "max_input_tokens": 131072, + "max_input_tokens": 262000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", @@ -12490,7 +12490,7 @@ "azure_ai/grok-code-fast-1": { "input_cost_per_token": 2e-07, "litellm_provider": "azure_ai", - "max_input_tokens": 131072, + "max_input_tokens": 256000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", @@ -27518,7 +27518,8 @@ "search_context_size_high": 0.035 }, "gemini_native_audio": true, - "input_cost_per_image_token": 3e-06 + "input_cost_per_image_token": 3e-06, + "input_cost_per_video_token": 3e-06 }, "gemini-live-2.5-flash-preview-native-audio-09-2025": { "input_cost_per_audio_token": 3e-06, @@ -28418,7 +28419,7 @@ "max_output_tokens": 65536, "max_tokens": 65536, "mode": "chat", - "output_cost_per_reasoning_token": 1e-05, + "output_cost_per_reasoning_token": 5e-06, "output_cost_per_token": 5e-06, "output_cost_per_token_batches": 2.5e-06, "search_context_cost_per_query": { @@ -29015,7 +29016,7 @@ "image" ], "supports_function_calling": false, - "supports_prompt_caching": true, + "supports_prompt_caching": false, "supports_response_schema": false, "supports_system_messages": true, "supports_vision": true, @@ -29026,7 +29027,7 @@ "search_context_size_high": 0.014 }, "web_search_billing_unit": "per_query", - "supports_reasoning": false + "supports_reasoning": true }, "gemini/nano-banana-pro-preview": { "input_cost_per_image": 0.0011, @@ -29104,8 +29105,8 @@ "image" ], "supports_function_calling": false, - "supports_prompt_caching": true, - "supports_reasoning": false, + "supports_prompt_caching": false, + "supports_reasoning": true, "supports_response_schema": false, "supports_system_messages": true, "supports_vision": true, @@ -29115,7 +29116,8 @@ "search_context_size_medium": 0.014, "search_context_size_high": 0.014 }, - "web_search_billing_unit": "per_query" + "web_search_billing_unit": "per_query", + "supports_pdf_input": true }, "gemini/gemini-3.1-flash-lite-image": { "input_cost_per_image": 0.00028, @@ -29146,8 +29148,9 @@ "image" ], "supports_function_calling": false, + "supports_pdf_input": true, "supports_prompt_caching": false, - "supports_reasoning": false, + "supports_reasoning": true, "supports_response_schema": false, "supports_system_messages": true, "supports_vision": true, @@ -29159,17 +29162,15 @@ "input_cost_per_token": 2e-06, "input_cost_per_token_batches": 1e-06, "litellm_provider": "gemini", - "max_input_tokens": 65536, - "max_output_tokens": 32768, - "max_tokens": 32768, - "mode": "image_generation", - "output_cost_per_image": 0.134, - "output_cost_per_image_token": 0.00012, + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", "output_cost_per_token": 1.2e-05, "rpm": 1000, "tpm": 4000000, "output_cost_per_token_batches": 6e-06, - "source": "https://ai.google.dev/gemini-api/docs/pricing", + "source": "https://ai.google.dev/gemini-api/docs/models/deep-research-pro-preview-12-2025", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -29177,11 +29178,12 @@ ], "supported_modalities": [ "text", - "image" + "image", + "audio", + "video" ], "supported_output_modalities": [ - "text", - "image" + "text" ], "supports_function_calling": false, "supports_prompt_caching": true, @@ -29193,7 +29195,8 @@ "search_context_size_low": 0.035, "search_context_size_medium": 0.035, "search_context_size_high": 0.035 - } + }, + "supports_pdf_input": true }, "gemini/gemini-2.5-flash-lite": { "cache_read_input_audio_token_cost": 3e-08, @@ -30785,11 +30788,15 @@ ] }, "github_copilot/claude-haiku-4.5": { + "cache_creation_input_token_cost": 1.25e-06, + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token": 1e-06, "litellm_provider": "github_copilot", "max_input_tokens": 128000, "max_output_tokens": 16000, "max_tokens": 16000, "mode": "chat", + "output_cost_per_token": 5e-06, "supported_endpoints": [ "/v1/chat/completions" ], @@ -30838,11 +30845,15 @@ "supports_vision": true }, "github_copilot/claude-sonnet-4": { + "cache_creation_input_token_cost": 3.75e-06, + "cache_read_input_token_cost": 3e-07, + "input_cost_per_token": 3e-06, "litellm_provider": "github_copilot", "max_input_tokens": 128000, "max_output_tokens": 16000, "max_tokens": 16000, "mode": "chat", + "output_cost_per_token": 1.5e-05, "supported_endpoints": [ "/v1/chat/completions" ], @@ -31003,11 +31014,14 @@ "supports_vision": true }, "github_copilot/gpt-5-mini": { + "cache_read_input_token_cost": 2.5e-08, + "input_cost_per_token": 2.5e-07, "litellm_provider": "github_copilot", "max_input_tokens": 128000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", + "output_cost_per_token": 2e-06, "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, @@ -31058,11 +31072,14 @@ "supports_vision": true }, "github_copilot/gpt-5.3-codex": { + "cache_read_input_token_cost": 1.75e-07, + "input_cost_per_token": 1.75e-06, "litellm_provider": "github_copilot", "max_input_tokens": 128000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", + "output_cost_per_token": 1.4e-05, "supported_endpoints": [ "/v1/responses" ], @@ -32745,6 +32762,25 @@ "audio" ] }, + "gpt-4o-mini-tts-2025-03-20": { + "input_cost_per_token": 6e-07, + "litellm_provider": "openai", + "mode": "audio_speech", + "output_cost_per_audio_token": 1.2e-05, + "output_cost_per_second": 0.00025, + "output_cost_per_token": 1e-05, + "source": "https://developers.openai.com/api/docs/models/gpt-4o-mini-tts", + "supported_endpoints": [ + "/v1/audio/speech" + ], + "supported_modalities": [ + "text", + "audio" + ], + "supported_output_modalities": [ + "audio" + ] + }, "gpt-4o-search-preview": { "cache_read_input_token_cost": 1.25e-06, "input_cost_per_token": 2.5e-06, @@ -42065,8 +42101,8 @@ "input_cost_per_token_cache_hit": 2e-08, "litellm_provider": "openrouter", "max_input_tokens": 163840, - "max_output_tokens": 65536, - "max_tokens": 65536, + "max_output_tokens": 147456, + "max_tokens": 147456, "mode": "chat", "output_cost_per_token": 4.1e-07, "source": "https://openrouter.ai/api/v1/models", @@ -42146,14 +42182,14 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4.1-flash": { - "cache_read_input_token_cost": 3.135e-08, - "input_cost_per_token": 3.483e-08, + "cache_read_input_token_cost": 6e-09, + "input_cost_per_token": 3e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, - "max_output_tokens": 384000, - "max_tokens": 384000, + "max_output_tokens": 943718, + "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 6e-07, + "output_cost_per_token": 1.2e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -42173,7 +42209,7 @@ "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 3.5e-06, + "output_cost_per_token": 4.2e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -42527,13 +42563,13 @@ "max_output_tokens": 8000 }, "openrouter/minimax/minimax-m2": { - "input_cost_per_token": 2.55e-07, + "input_cost_per_token": 3e-07, "litellm_provider": "openrouter", "max_input_tokens": 204800, - "max_output_tokens": 131072, - "max_tokens": 131072, + "max_output_tokens": 176947, + "max_tokens": 176947, "mode": "chat", - "output_cost_per_token": 1.02e-06, + "output_cost_per_token": 1.2e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -42739,14 +42775,14 @@ "supports_web_search": false }, "openrouter/nvidia/nemotron-3.5-lightning": { - "cache_read_input_token_cost": 4e-08, - "input_cost_per_token": 8e-08, + "cache_read_input_token_cost": 3e-08, + "input_cost_per_token": 6e-08, "litellm_provider": "openrouter", "max_input_tokens": 262144, - "max_output_tokens": 235929, - "max_tokens": 235929, + "max_output_tokens": 32768, + "max_tokens": 32768, "mode": "chat", - "output_cost_per_token": 2e-07, + "output_cost_per_token": 1.6e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -43199,14 +43235,14 @@ "supports_web_search": true }, "openrouter/openai/gpt-5.6-sol-pro": { - "input_cost_per_token": 2e-06, - "output_cost_per_token": 1e-05, - "cache_read_input_token_cost": 2e-07, - "cache_creation_input_token_cost": 2.5e-06, - "cache_creation_input_token_cost_above_272k_tokens": 5e-06, - "input_cost_per_token_above_272k_tokens": 4e-06, - "output_cost_per_token_above_272k_tokens": 1.5e-05, - "cache_read_input_token_cost_above_272k_tokens": 4e-07, + "input_cost_per_token": 4e-06, + "output_cost_per_token": 2e-05, + "cache_read_input_token_cost": 4e-07, + "cache_creation_input_token_cost": 5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 1e-05, + "input_cost_per_token_above_272k_tokens": 8e-06, + "output_cost_per_token_above_272k_tokens": 3e-05, + "cache_read_input_token_cost_above_272k_tokens": 8e-07, "litellm_provider": "openrouter", "max_input_tokens": 1050000, "max_output_tokens": 128000, @@ -43264,6 +43300,7 @@ "supports_web_search": false }, "openrouter/openai/gpt-oss-20b": { + "cache_read_input_token_cost": 9e-09, "input_cost_per_token": 1.8e-08, "litellm_provider": "openrouter", "max_input_tokens": 131072, @@ -49891,7 +49928,8 @@ "supports_tool_choice": true, "supports_vision": true, "prompt_cache_min_tokens": 1024, - "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", + "cache_creation_input_token_cost_batches": 1.88e-06 }, "vertex_ai/claude-sonnet-5": { "deprecation_date": "2026-12-24", @@ -49999,7 +50037,8 @@ "supports_vision": true, "supports_native_streaming": true, "prompt_cache_min_tokens": 1024, - "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing" + "source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing", + "cache_creation_input_token_cost_batches": 1.88e-06 }, "vertex_ai/mistralai/codestral-2@001": { "input_cost_per_token": 3e-07, @@ -60514,6 +60553,9 @@ "input_cost_per_audio_token": 1e-06, "input_cost_per_token": 1e-06, "litellm_provider": "gemini", + "max_input_tokens": 131072, + "max_output_tokens": 65536, + "max_tokens": 65536, "mode": "chat", "output_cost_per_token": 5e-06, "search_context_cost_per_query": { @@ -60536,6 +60578,7 @@ ], "supports_audio_input": true, "supports_function_calling": true, + "supports_reasoning": true, "supports_video_input": true, "supports_vision": true, "supports_web_search": true, @@ -64006,7 +64049,7 @@ "groq/qwen/qwen3.8-27b": { "input_cost_per_token": 8e-07, "litellm_provider": "groq", - "max_input_tokens": 131042, + "max_input_tokens": 131072, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", @@ -67212,13 +67255,13 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-flash-vision-exp": { - "input_cost_per_token": 4.4e-07, - "output_cost_per_token": 1.32e-06, - "cache_read_input_token_cost": 1.4e-08, + "input_cost_per_token": 2.156e-07, + "output_cost_per_token": 6.468e-07, + "cache_read_input_token_cost": 6.86e-09, "litellm_provider": "openrouter", "max_input_tokens": 1048576, - "max_output_tokens": 943718, - "max_tokens": 943718, + "max_output_tokens": 262144, + "max_tokens": 262144, "mode": "chat", "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, @@ -67232,13 +67275,13 @@ "supports_web_search": false }, "openrouter/z-ai/glm-5.3": { - "input_cost_per_token": 3.556e-07, - "output_cost_per_token": 2.574e-06, - "cache_read_input_token_cost": 6.604e-08, + "input_cost_per_token": 1.4e-06, + "output_cost_per_token": 4.4e-06, + "cache_read_input_token_cost": 2.6e-07, "litellm_provider": "openrouter", "max_input_tokens": 1310720, - "max_output_tokens": 943718, - "max_tokens": 943718, + "max_output_tokens": 943717, + "max_tokens": 943717, "mode": "chat", "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, @@ -67370,7 +67413,7 @@ }, "openrouter/deepseek/deepseek-v4-flash-0731": { "cache_read_input_token_cost": 1.6e-08, - "input_cost_per_token": 2.1e-08, + "input_cost_per_token": 1.8e-08, "litellm_provider": "openrouter", "max_input_tokens": 1310720, "max_output_tokens": 943718, @@ -67964,9 +68007,9 @@ "supports_web_search": false }, "openrouter/moonshotai/kimi-k2.6": { - "input_cost_per_token": 9.5e-07, - "output_cost_per_token": 4e-06, - "cache_read_input_token_cost": 1.6e-07, + "input_cost_per_token": 6.5e-07, + "output_cost_per_token": 3.41e-06, + "cache_read_input_token_cost": 1.5e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 235929, @@ -69207,12 +69250,12 @@ }, "openrouter/qwen/qwen3-30b-a3b": { "deprecation_date": "2026-10-09", - "input_cost_per_token": 1.3e-07, - "output_cost_per_token": 5.2e-07, + "input_cost_per_token": 1.2e-07, + "output_cost_per_token": 5e-07, "litellm_provider": "openrouter", "max_input_tokens": 131072, - "max_output_tokens": 8192, - "max_tokens": 8192, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, @@ -69247,8 +69290,8 @@ }, "openrouter/qwen/qwen3-14b": { "deprecation_date": "2026-10-09", - "input_cost_per_token": 2.275e-07, - "output_cost_per_token": 9.1e-07, + "input_cost_per_token": 1.2e-07, + "output_cost_per_token": 2.4e-07, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 16384, @@ -69878,7 +69921,9 @@ "vertex_ai/gemini-2.5-flash-native-audio": { "deprecation_date": "2026-12-13", "input_cost_per_audio_token": 3e-06, + "input_cost_per_image_token": 3e-06, "input_cost_per_token": 5e-07, + "input_cost_per_video_token": 3e-06, "litellm_provider": "vertex_ai", "mode": "realtime", "output_cost_per_audio_token": 1.2e-05, @@ -70678,17 +70723,17 @@ }, "azure/eu/gpt-6-astra": { "deprecation_date": "2028-01-11", - "cache_creation_input_token_cost": 1.375e-05, - "cache_creation_input_token_cost_above_272k_tokens": 2.75e-05, - "cache_read_input_token_cost": 1.1e-06, - "cache_read_input_token_cost_above_272k_tokens": 2.2e-06, - "input_cost_per_token": 1.1e-05, - "input_cost_per_token_above_272k_tokens": 2.2e-05, + "cache_creation_input_token_cost": 1.5e-05, + "cache_creation_input_token_cost_above_272k_tokens": 3e-05, + "cache_read_input_token_cost": 1.2e-06, + "cache_read_input_token_cost_above_272k_tokens": 2.4e-06, + "input_cost_per_token": 1.2e-05, + "input_cost_per_token_above_272k_tokens": 2.4e-05, "litellm_provider": "azure", "mode": "chat", - "output_cost_per_token": 5.5e-05, - "output_cost_per_token_above_272k_tokens": 8.25e-05, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "output_cost_per_token": 6e-05, + "output_cost_per_token_above_272k_tokens": 9e-05, + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'swedencentral'%20and%20priceType%20eq%20'Consumption'", "supports_reasoning": true }, "azure/eu/gpt-6-luna": { @@ -72821,17 +72866,17 @@ "supports_web_search": true }, "openrouter/~x-ai/grok-latest": { - "cache_read_input_token_cost": 4e-07, - "cache_read_input_token_cost_above_200k_tokens": 8e-07, - "input_cost_per_token": 1.6e-06, - "input_cost_per_token_above_200k_tokens": 3.2e-06, + "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_above_200k_tokens": 1e-06, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_200k_tokens": 4e-06, "litellm_provider": "openrouter", "max_input_tokens": 500000, "max_output_tokens": 450000, "max_tokens": 450000, "mode": "chat", - "output_cost_per_token": 4.8e-06, - "output_cost_per_token_above_200k_tokens": 9.6e-06, + "output_cost_per_token": 6e-06, + "output_cost_per_token_above_200k_tokens": 1.2e-05, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -74106,7 +74151,7 @@ "cache_read_input_token_cost": 4.2e-09, "input_cost_per_token": 2.1e-08, "litellm_provider": "openrouter", - "max_input_tokens": 131072, + "max_input_tokens": 262144, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", @@ -76134,7 +76179,7 @@ "cache_read_input_token_cost": 1.7e-07, "input_cost_per_token": 1e-06, "litellm_provider": "openrouter", - "max_input_tokens": 1048576, + "max_input_tokens": 524288, "max_output_tokens": 471859, "max_tokens": 471859, "mode": "chat", @@ -76154,7 +76199,7 @@ "cache_read_input_token_cost": 1e-07, "input_cost_per_token": 4.5e-07, "litellm_provider": "openrouter", - "max_input_tokens": 1048576, + "max_input_tokens": 524288, "max_output_tokens": 262144, "max_tokens": 262144, "mode": "chat", @@ -76371,6 +76416,7 @@ "supports_web_search": false }, "openrouter/prism-ml/ternary-bonsai-2-27b": { + "cache_read_input_token_cost": 3.75e-08, "input_cost_per_token": 7.5e-08, "litellm_provider": "openrouter", "max_input_tokens": 262144, @@ -76431,17 +76477,17 @@ "supports_web_search": false }, "openrouter/x-ai/grok-4.7": { - "cache_read_input_token_cost": 4e-07, - "cache_read_input_token_cost_above_200k_tokens": 8e-07, - "input_cost_per_token": 1.6e-06, - "input_cost_per_token_above_200k_tokens": 3.2e-06, + "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_above_200k_tokens": 1e-06, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_200k_tokens": 4e-06, "litellm_provider": "openrouter", "max_input_tokens": 500000, "max_output_tokens": 450000, "max_tokens": 450000, "mode": "chat", - "output_cost_per_token": 4.8e-06, - "output_cost_per_token_above_200k_tokens": 9.6e-06, + "output_cost_per_token": 6e-06, + "output_cost_per_token_above_200k_tokens": 1.2e-05, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -76454,16 +76500,16 @@ "supports_web_search": true }, "moonshotai.kimi-k3": { - "cache_creation_input_token_cost": 3.75e-06, - "cache_read_input_token_cost": 3e-07, - "input_cost_per_token": 3e-06, + "cache_creation_input_token_cost": 4.125e-06, + "cache_read_input_token_cost": 3.3e-07, + "input_cost_per_token": 3.3e-06, "litellm_provider": "bedrock_converse", "max_input_tokens": 1000000, "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 1.5e-05, - "source": "https://aws.amazon.com/bedrock/pricing/", + "output_cost_per_token": 1.65e-05, + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrock/current/us-east-1/index.json", "supports_audio_input": false, "supports_function_calling": true, "supports_prompt_caching": true, @@ -79026,6 +79072,28 @@ "supports_tool_choice": true, "supports_vision": true }, + "openrouter/anthropic/claude-sonnet-5.5:batch": { + "cache_creation_input_token_cost": 1.25e-06, + "cache_creation_input_token_cost_above_1hr": 2e-06, + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token": 1e-06, + "litellm_provider": "openrouter", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-06, + "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": false, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, "baseten/deepseek-ai/DeepSeek-V4.1-Flash-Fast": { "cache_read_input_token_cost": 1.4e-07, "input_cost_per_token": 6e-07, diff --git a/tests/unit/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/unit/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py index 0afd989272e..088247c2ea4 100644 --- a/tests/unit/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py +++ b/tests/unit/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py @@ -1145,6 +1145,35 @@ def test_generic_cost_per_token_gpt54_above_272k_tokens(_local_model_cost_map): assert round(completion_cost, 10) == round(expected_completion, 10) +@pytest.mark.parametrize( + ("prompt_tokens", "input_rate", "cache_read_rate", "output_rate"), + [ + (100_000, 1.2e-05, 1.2e-06, 6e-05), + (300_000, 2.4e-05, 2.4e-06, 9e-05), + ], +) +def test_generic_cost_per_token_azure_eu_gpt_6_astra_tiers( + _local_model_cost_map, prompt_tokens, input_rate, cache_read_rate, output_rate +): + """azure/eu/gpt-6-astra bills Azure's Data Zone rates, doubling input and cache read past 272K.""" + cached_tokens = 20_000 + completion_tokens = 1_000 + usage = Usage( + prompt_tokens=prompt_tokens, + completion_tokens=completion_tokens, + total_tokens=prompt_tokens + completion_tokens, + prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=cached_tokens), + ) + prompt_cost, completion_cost = generic_cost_per_token( + model="azure/eu/gpt-6-astra", + usage=usage, + custom_llm_provider="azure", + ) + expected_prompt = (prompt_tokens - cached_tokens) * input_rate + cached_tokens * cache_read_rate + assert prompt_cost == pytest.approx(expected_prompt) + assert completion_cost == pytest.approx(completion_tokens * output_rate) + + def test_generic_cost_per_token_minimax_m3_above_512k_tokens(_local_model_cost_map): """MiniMax-M3: prompts >512K input tokens priced at 2x input, output, and cache read.""" model = "minimax/MiniMax-M3" From fb74957ddd486d08aa4eda600f46fafdc8b5f7c5 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 12:50:08 -0700 Subject: [PATCH 28/41] fix(guardrails): enable explicit PANW MCP output scanning (#43109) * fix(guardrails): declare post_mcp_call for PANW Prisma AIRS Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(guardrails): exercise post_mcp_call_hook dispatch in PANW tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(model-catalog): add fal_ai resolution-tiered image cost fields Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(guardrails): keep post_mcp_call opt-in for PANW Prisma AIRS Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(model-catalog): add fal_ai resolution-tiered image cost keys to cost map schema Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: joshua Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> --- .../panw_prisma_airs/panw_prisma_airs.py | 1 + .../guardrail_hooks/test_panw_prisma_airs.py | 87 +++++++++++++++++++ .../proxy/guardrails/test_init_guardrails.py | 34 +++++++- 3 files changed, 121 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py index cd538ad8c8d..e4822195bec 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py +++ b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py @@ -2012,4 +2012,5 @@ class PanwPrismaAirsHandler(CustomGuardrail): GuardrailEventHooks.logging_only, GuardrailEventHooks.pre_mcp_call, GuardrailEventHooks.during_mcp_call, + GuardrailEventHooks.post_mcp_call, ] diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py index 1f52fa224ee..5db3e11ac06 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py @@ -21,7 +21,9 @@ from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest from fastapi import HTTPException +from mcp.types import CallToolResult, TextContent +import litellm from litellm.caching import DualCache from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.proxy._types import UserAPIKeyAuth @@ -29,6 +31,7 @@ from litellm.proxy.guardrails.guardrail_hooks.panw_prisma_airs import ( PanwPrismaAirsHandler, initialize_guardrail, ) +from litellm.proxy.utils import ProxyLogging from litellm.types.guardrails import GuardrailEventHooks, LitellmParams from litellm.types.utils import ( ChatCompletionCustomToolCallPayload, @@ -2025,6 +2028,30 @@ class TestPanwAirsShouldRunGuardrail: True, id="explicit_pre_mcp_call_mode", ), + pytest.param( + True, + "post_call", + _simple_data(), + GuardrailEventHooks.post_mcp_call, + False, + id="post_call_mode_does_not_run_for_post_mcp_call", + ), + pytest.param( + True, + "post_mcp_call", + _simple_data(), + GuardrailEventHooks.post_mcp_call, + True, + id="explicit_post_mcp_call_mode", + ), + pytest.param( + True, + "post_mcp_call", + _simple_data(), + GuardrailEventHooks.post_call, + False, + id="post_mcp_call_mode_does_not_run_for_regular_post_call", + ), pytest.param( True, "pre_call", @@ -2048,6 +2075,66 @@ class TestPanwAirsShouldRunGuardrail: assert handler.should_run_guardrail(data, query_event) is expected +class TestPanwAirsPostMcpCall: + """Explicit MCP output scans use the existing AIRS response contract.""" + + @pytest.mark.asyncio + @pytest.mark.parametrize("action", ["allow", "block", "mask"]) + async def test_post_mcp_call_scans_tool_result(self, monkeypatch: pytest.MonkeyPatch, action: str) -> None: + original: Final = "ssn 123-45-6789" + masked: Final = "ssn ***********" + + def respond(request: httpx.Request) -> httpx.Response: + payload: Final = json.loads(request.content) + assert request.url.path.endswith("/v1/scan/sync/request") + assert payload["contents"] == [{"response": original}] + assert payload["ai_profile"] == {"profile_name": "test_profile"} + return httpx.Response( + 200, + json={ + "action": "block" if action == "block" else "allow", + "category": "malicious" if action == "block" else "benign", + "scan_id": "s1", + "report_id": "r1", + "profile_name": "test_profile", + **({"response_masked_data": {"data": masked}} if action == "mask" else {}), + }, + ) + + transport_handler: Final = MagicMock(side_effect=respond) + http_client: Final = AsyncHTTPHandler(transport=httpx.MockTransport(transport_handler)) + handler: Final = make_handler( + event_hook="post_mcp_call", + default_on=True, + mask_response_content=True, + http_client=http_client, + ) + monkeypatch.setattr(litellm, "callbacks", [handler]) + proxy_logging: Final = ProxyLogging(user_api_key_cache=DualCache()) + result: Final = CallToolResult(content=[TextContent(type="text", text=original)], isError=False) + try: + if action == "block": + with pytest.raises(HTTPException) as exc_info: + await proxy_logging.post_mcp_call_hook( + response=result, + request_data={"litellm_call_id": "c1"}, + user_api_key_dict=None, + ) + assert exc_info.value.status_code == 400 + transport_handler.assert_called_once() + return + returned: Final = await proxy_logging.post_mcp_call_hook( + response=result, + request_data={"litellm_call_id": "c1"}, + user_api_key_dict=None, + ) + transport_handler.assert_called_once() + assert returned.model_dump(by_alias=True)["isError"] is False + assert returned.content == [TextContent(type="text", text=masked if action == "mask" else original)] + finally: + await http_client.client.aclose() + + class TestPanwAirsToolEventIsResponseFix: """Tests for Bug A fix: tool_event scans must not set is_response metadata.""" diff --git a/tests/test_litellm/proxy/guardrails/test_init_guardrails.py b/tests/test_litellm/proxy/guardrails/test_init_guardrails.py index 79d91db902c..fcd7e537937 100644 --- a/tests/test_litellm/proxy/guardrails/test_init_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/test_init_guardrails.py @@ -1,5 +1,5 @@ import json -from typing import Literal +from typing import Final, Literal from unittest.mock import MagicMock, patch import pytest @@ -11,6 +11,38 @@ from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2 from litellm.types.guardrails import Mode, SupportedGuardrailIntegrations +def test_init_guardrails_v2_registers_panw_mcp_output_scanner(monkeypatch: pytest.MonkeyPatch) -> None: + import litellm + from litellm.proxy.guardrails import guardrail_registry + from litellm.proxy.guardrails.guardrail_hooks.panw_prisma_airs import PanwPrismaAirsHandler + from litellm.types.guardrails import GuardrailEventHooks + + monkeypatch.setenv("LITELLM_STRICT_GUARDRAIL_MODES", "true") + monkeypatch.setattr(guardrail_registry, "IN_MEMORY_GUARDRAIL_HANDLER", InMemoryGuardrailHandler()) + init_guardrails_v2( + all_guardrails=[ + { + "guardrail_name": "panw-mcp-output", + "litellm_params": { + "guardrail": "panw_prisma_airs", + "mode": "post_mcp_call", + "default_on": True, + "api_key": "test-panw-key", + "profile_name": "test-profile", + }, + } + ] + ) + scanners: Final = tuple( + callback + for callback in litellm.callbacks + if isinstance(callback, PanwPrismaAirsHandler) and callback.guardrail_name == "panw-mcp-output" + ) + assert len(scanners) == 1, "PANW MCP output scanning must be registered at startup" + assert scanners[0].should_run_guardrail({}, GuardrailEventHooks.post_mcp_call) is True + assert scanners[0].should_run_guardrail({}, GuardrailEventHooks.post_call) is False + + def test_initialize_presidio_guardrail(): """ Test that initialize_guardrail correctly uses registered initializers From 273489824aa45024e471a21d54a490132c8db388 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 19:54:12 +0000 Subject: [PATCH 29/41] refactor(rust): orchestrate Messages route execution (#43719) Co-authored-by: Yujong Lee Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm-rust/crates/cache-response/AGENTS.md | 2 +- litellm-rust/crates/cache-response/src/lib.rs | 4 +- .../crates/cache-response/src/service.rs | 44 ++-- .../crates/cache-response/tests/service.rs | 17 +- litellm-rust/crates/core/AGENTS.md | 4 +- litellm-rust/crates/core/src/caching.rs | 182 +++++++++------ .../core/src/chat_completions/handler.rs | 2 +- .../crates/core/src/chat_completions/mod.rs | 2 +- .../crates/core/src/chat_completions/route.rs | 2 +- litellm-rust/crates/core/src/context.rs | 48 ++++ litellm-rust/crates/core/src/lib.rs | 7 +- .../crates/core/src/messages/handler.rs | 208 ++++++++++-------- litellm-rust/crates/core/src/messages/mod.rs | 147 ++++--------- .../crates/core/src/messages/route.rs | 10 +- .../crates/core/src/responses/handler.rs | 2 +- litellm-rust/crates/core/src/responses/mod.rs | 4 +- litellm-rust/crates/core/tests/caching.rs | 16 +- .../crates/core/tests/messages/host.rs | 176 ++++++++++++++- .../crates/core/tests/messages/response.rs | 126 ++++++++--- litellm-rust/crates/core/tests/support/mod.rs | 10 +- .../crates/gateway-inference/src/caching.rs | 14 +- .../gateway-inference/src/chat_completions.rs | 2 +- .../crates/gateway-inference/src/lib.rs | 6 +- .../crates/gateway-inference/src/messages.rs | 2 +- .../crates/gateway-inference/src/responses.rs | 2 +- .../python-bridge/src/cache/native/v2.rs | 20 +- .../python-bridge/src/cache/selection.rs | 11 +- .../src/routes/chat_completions.rs | 2 +- .../python-bridge/src/routes/messages/mod.rs | 16 +- .../python-bridge/src/routes/responses.rs | 2 +- 30 files changed, 711 insertions(+), 379 deletions(-) create mode 100644 litellm-rust/crates/core/src/context.rs diff --git a/litellm-rust/crates/cache-response/AGENTS.md b/litellm-rust/crates/cache-response/AGENTS.md index d86fe6cc588..4dbcc65d403 100644 --- a/litellm-rust/crates/cache-response/AGENTS.md +++ b/litellm-rust/crates/cache-response/AGENTS.md @@ -24,6 +24,6 @@ Keep unary caching independent of stream-only methods. Store streams only after Test each contract in its owner: storage capabilities in backend tests, envelopes and freshness here, reuse and replay in core, Python callback and fallback behavior at the bridge, and HTTP behavior at the gateway. Run backend contract checks and Python response-codec fixtures before exposing a new backend -`ScopedCache` requires an explicit shared or isolated scope at construction. `CacheOptions` has no default sharing policy. Callers may override policy per invocation without replacing the attached service. Versioned native envelopes reject incompatible API surfaces and versions as misses; this envelope is distinct from the legacy Python response codec +`ScopedCache` requires an explicit shared or isolated scope at construction. Per-call `CachePolicy` controls reads, writes, expiry, and freshness without replacing the attached scope or service. `CacheOptions` binds that policy to an explicit scope for storage requests and has no default sharing policy. Versioned native envelopes reject incompatible API surfaces and versions as misses; this envelope is distinct from the legacy Python response codec Response storage is not the source of budget or rate-limit coordination dependencies. Keep counters, reservations, and atomic admission operations out of `ResponseCacheService`, including when both services happen to use Redis diff --git a/litellm-rust/crates/cache-response/src/lib.rs b/litellm-rust/crates/cache-response/src/lib.rs index ebabcf70c9f..78de27f2d9f 100644 --- a/litellm-rust/crates/cache-response/src/lib.rs +++ b/litellm-rust/crates/cache-response/src/lib.rs @@ -17,6 +17,6 @@ pub use exact::{ConnectionProbe, ExactResponseCache}; pub use response::{ResponseCache, ResponseCacheRequest}; pub use service::{ - CacheOptions, CacheScope, ResponseCacheConfig, ResponseCacheService, ResponseEnvelope, - ScopedCache, + CacheOptions, CachePolicy, CacheScope, ResponseCacheConfig, ResponseCacheService, + ResponseEnvelope, ScopedCache, }; diff --git a/litellm-rust/crates/cache-response/src/service.rs b/litellm-rust/crates/cache-response/src/service.rs index 51359a9a8d6..0bdf948ec48 100644 --- a/litellm-rust/crates/cache-response/src/service.rs +++ b/litellm-rust/crates/cache-response/src/service.rs @@ -73,32 +73,35 @@ pub enum CacheScope { Isolated(String), } -#[derive(Clone)] -pub struct CacheOptions { +#[derive(Clone, Copy, Default)] +pub struct CachePolicy { pub caching: Option, pub no_cache: bool, pub no_store: bool, pub ttl: Option, pub max_age: Option, +} + +impl CachePolicy { + pub fn enabled(&self) -> bool { + self.caching != Some(false) && !(self.no_cache && self.no_store) + } +} + +#[derive(Clone)] +pub struct CacheOptions { + pub policy: CachePolicy, pub scope: CacheScope, } impl CacheOptions { pub fn new(scope: CacheScope) -> Self { Self { - caching: None, - no_cache: false, - no_store: false, - ttl: None, - max_age: None, + policy: CachePolicy::default(), scope, } } - pub fn enabled(&self) -> bool { - self.caching != Some(false) && !(self.no_cache && self.no_store) - } - pub fn request(self, namespace: &str, surface: &str, mut input: Value) -> ResponseCacheRequest { input.sort_all_objects(); let scope = match self.scope { @@ -128,13 +131,15 @@ impl CacheOptions { supported_call_type: true, native_backend: true, default_on: true, - caching: self.caching, - no_cache: self.no_cache, - no_store: self.no_store, + caching: self.policy.caching, + no_cache: self.policy.no_cache, + no_store: self.policy.no_store, ..Default::default() }, - context: ExactCacheContext { ttl: self.ttl }, - max_age: self.max_age, + context: ExactCacheContext { + ttl: self.policy.ttl, + }, + max_age: self.policy.max_age, } } } @@ -171,7 +176,10 @@ impl ScopedCache { Self { service, scope } } - pub fn options(&self, overrides: Option) -> CacheOptions { - overrides.unwrap_or_else(|| CacheOptions::new(self.scope.clone())) + pub fn options(&self, policy: Option) -> CacheOptions { + CacheOptions { + policy: policy.unwrap_or_default(), + scope: self.scope.clone(), + } } } diff --git a/litellm-rust/crates/cache-response/tests/service.rs b/litellm-rust/crates/cache-response/tests/service.rs index d4532776719..be1cf1f8ea7 100644 --- a/litellm-rust/crates/cache-response/tests/service.rs +++ b/litellm-rust/crates/cache-response/tests/service.rs @@ -129,11 +129,20 @@ async fn isolated_policy_controls_actual_entry_reuse( #[case] first: &str, #[case] second: &str, #[case] hit: bool, + #[values(false, true)] override_policy: bool, ) { - use litellm_cache_response::{CacheOptions, CacheScope}; - let service = ResponseCache::new(Arc::new(InMemoryCache::::default())); - let request = - |scope| CacheOptions::new(scope).request("test", "messages", json!({"prompt":"hello"})); + use litellm_cache_response::{CachePolicy, CacheScope, ScopedCache}; + let service = Arc::new(ResponseCache::new(Arc::new( + InMemoryCache::::default(), + ))); + let request = |scope| { + ScopedCache::new(service.clone(), scope) + .options(override_policy.then_some(CachePolicy { + ttl: Some(Duration::from_secs(30)), + ..CachePolicy::default() + })) + .request("test", "messages", json!({"prompt":"hello"})) + }; service .async_store( &request(CacheScope::Isolated(first.into())), diff --git a/litellm-rust/crates/core/AGENTS.md b/litellm-rust/crates/core/AGENTS.md index ec96239beac..6cb07e6dbfc 100644 --- a/litellm-rust/crates/core/AGENTS.md +++ b/litellm-rust/crates/core/AGENTS.md @@ -36,7 +36,9 @@ Not here: serving HTTP (axum routes, extractors), config file reading, rollout s ## Response caching and accounting boundary -Attach a `litellm_cache_response::ScopedCache` with `route.with_cache(cache)`. Cached and uncached routes use the same `execute` and `machine` methods. `CallOptions` carries per-call cache overrides and observation; attaching a service does not change the execution contract +Attach a `litellm_cache_response::ScopedCache` with `route.with_cache(cache)`. Cached and uncached routes use the same `execute` and `machine` methods. `CallOptions` carries a scope-free `CachePolicy` and observation; per-call policy never replaces the attached scope or service + +Messages groups per-call dependencies in `CallContext` and explicitly sequences cache lookup, provider execution, result acceptance, and cache storage. Provider transport does not own cache orchestration. Stream capture remains in the shared cache implementation Core owns request identity, typed response reconstruction and stream capture/replay. `cache-response` owns cache policy, namespacing, scope encoding, versioned envelopes and freshness. The SDK explicitly chooses shared scope. The gateway derives isolated scope from authenticated identity before attaching its service diff --git a/litellm-rust/crates/core/src/caching.rs b/litellm-rust/crates/core/src/caching.rs index d182ba94543..980e0e6d34e 100644 --- a/litellm-rust/crates/core/src/caching.rs +++ b/litellm-rust/crates/core/src/caching.rs @@ -1,5 +1,6 @@ use std::{ future::Future, + marker::PhantomData, sync::Arc, time::{Duration, SystemTime, UNIX_EPOCH}, }; @@ -7,7 +8,8 @@ use std::{ use bytes::{Bytes, BytesMut}; use futures_util::{StreamExt, TryStreamExt, stream}; use litellm_cache_response::{ - CacheOptions, ResponseCacheRequest, ResponseCacheService, ResponseEnvelope, cache_key, + CacheOptions, CachePolicy, ResponseCacheRequest, ResponseCacheService, ResponseEnvelope, + ScopedCache, cache_key, }; use litellm_host::{ call::{CallOutput, OutputOf}, @@ -77,7 +79,7 @@ impl CacheSession { options: Option, request: &CacheRequest, ) -> Option { - let options = options.filter(CacheOptions::enabled)?; + let options = options.filter(|options| options.policy.enabled())?; let service = service?; let input = request.input.clone(); let request = options.request(&service.config().namespace, P::SURFACE, input); @@ -198,81 +200,137 @@ where let identity = request.identity.clone(); crate::diagnostic::provider(&identity.model, &identity.provider); let session = CacheSession::prepare::

    (cache, options, &request); - let hit = match &session { - Some(session) => session.lookup::

    ().await.and_then(|entry| { - let output = match entry { - CachedOutput::Response(response) => Some(CallOutput::Complete(response)), - CachedOutput::Stream(data) => P::replay(Bytes::from(data)), - }; - output.map(|output| (output, cache_key(&session.request.key))) - }), - None => None, + let cache = CallCache::

    { + session, + protocol: PhantomData, }; + let hit = cache.lookup().await; let (output, source) = match hit { - Some((output, key)) => (output, ResultSource::Cache { key }), + Some(hit) => hit, None => (provider().await?, ResultSource::Provider), }; - let from_provider = source == ResultSource::Provider; publish( ExecutionFacts { provider: identity, - source, + source: source.clone(), }, interceptors, observers, ) .await?; - let Some(session) = - session.filter(|session| from_provider && session.request.controls.writes()) - else { - return Ok(output); - }; - match output { - CallOutput::Complete(response) => { - session.store_response::

    (&response).await; - Ok(CallOutput::Complete(response)) - } - CallOutput::Stream { head, chunks } => { - let captured = stream::try_unfold( - (chunks, Some(Vec::::new()), session), - |(mut chunks, captured, session)| async move { - match chunks.try_next().await? { - Some(chunk) => { - let captured = captured.and_then(|mut data| { - let bytes = P::bytes(&chunk); - if data.len().saturating_add(bytes.len()) - > session.service.config().max_entry_bytes - { - return None; - } - data.extend_from_slice(bytes); - Some(data) - }); - Ok(Some((chunk, (chunks, captured, session)))) - } - None => { - if let Some(data) = captured - && let Ok(text) = String::from_utf8(data) - && successful_stream(&text, P::TERMINAL_EVENT) - && let Ok(entry) = serde_json::to_value(ResponseEnvelope::new( - P::SURFACE, - CachedOutput::::Stream(text), - )) - { - session.store(entry).await; - } - Ok::<_, RouteError>(None) - } - } - }, - ) - .boxed(); - Ok(CallOutput::Stream { - head, - chunks: captured, + Ok(cache.finish(output, &source).await) +} + +pub(crate) struct CallCache

    { + session: Option, + protocol: PhantomData

    , +} + +impl CallCache

    { + pub(crate) fn from_wire( + cache: Option<&ScopedCache>, + policy: CachePolicy, + identity: &ProviderIdentity, + wire: &WireRequest, + ) -> Self { + let session = cache.and_then(|cache| { + if !policy.enabled() { + return None; + } + let options = cache.options(Some(policy)); + let request = CacheRequest::from_wire(identity.clone(), Some(wire)); + Some(CacheSession { + request: options.request( + &cache.service.config().namespace, + P::SURFACE, + request.input, + ), + service: cache.service.clone(), }) + }); + Self { + session, + protocol: PhantomData, } } + + pub(crate) async fn lookup(&self) -> Option<(OutputOf

    , ResultSource)> + where + P::Response: DeserializeOwned, + { + let session = self.session.as_ref()?; + let output = match session.lookup::

    ().await? { + CachedOutput::Response(response) => CallOutput::Complete(response), + CachedOutput::Stream(data) => P::replay(Bytes::from(data))?, + }; + Some(( + output, + ResultSource::Cache { + key: cache_key(&session.request.key), + }, + )) + } + + pub(crate) async fn finish(self, output: OutputOf

    , source: &ResultSource) -> OutputOf

    + where + P::Response: Serialize, + { + let Some(session) = self.session.filter(|session| { + *source == ResultSource::Provider && session.request.controls.writes() + }) else { + return output; + }; + match output { + CallOutput::Complete(response) => { + session.store_response::

    (&response).await; + CallOutput::Complete(response) + } + CallOutput::Stream { head, chunks } => CallOutput::Stream { + head, + chunks: capture_stream::

    (chunks, session), + }, + } + } +} + +fn capture_stream( + chunks: futures_util::stream::BoxStream<'static, Result>, + session: CacheSession, +) -> futures_util::stream::BoxStream<'static, Result> { + stream::try_unfold( + (chunks, Some(Vec::::new()), session), + |(mut chunks, captured, session)| async move { + match chunks.try_next().await? { + Some(chunk) => { + let captured = captured.and_then(|mut data| { + let bytes = P::bytes(&chunk); + if data.len().saturating_add(bytes.len()) + > session.service.config().max_entry_bytes + { + return None; + } + data.extend_from_slice(bytes); + Some(data) + }); + Ok(Some((chunk, (chunks, captured, session)))) + } + None => { + if let Some(data) = captured + && let Ok(text) = String::from_utf8(data) + && successful_stream(&text, P::TERMINAL_EVENT) + && let Ok(entry) = serde_json::to_value(ResponseEnvelope::new( + P::SURFACE, + CachedOutput::::Stream(text), + )) + { + session.store(entry).await; + } + Ok::<_, RouteError>(None) + } + } + }, + ) + .boxed() } fn now() -> Duration { diff --git a/litellm-rust/crates/core/src/chat_completions/handler.rs b/litellm-rust/crates/core/src/chat_completions/handler.rs index 9f9d48cb177..b5148cce7af 100644 --- a/litellm-rust/crates/core/src/chat_completions/handler.rs +++ b/litellm-rust/crates/core/src/chat_completions/handler.rs @@ -22,7 +22,7 @@ pub(super) async fn execute( auth: &AuthServices, request: ProviderChatCompletionsRequest, cache: Option, - cache_options: Option, + cache_options: Option, interceptors: &impl Interceptors, observers: Option<&ObservationSender>, ) -> Result { diff --git a/litellm-rust/crates/core/src/chat_completions/mod.rs b/litellm-rust/crates/core/src/chat_completions/mod.rs index a26648b88ef..00aadb509f4 100644 --- a/litellm-rust/crates/core/src/chat_completions/mod.rs +++ b/litellm-rust/crates/core/src/chat_completions/mod.rs @@ -67,7 +67,7 @@ impl ChatCompletionsRoute { async fn run( &self, request: ChatCompletionsRequest<'_>, - cache_options: Option, + cache_options: Option, interceptors: &impl litellm_host::interceptors::Interceptors, observers: Option<&ObservationSender>, ) -> Result { diff --git a/litellm-rust/crates/core/src/chat_completions/route.rs b/litellm-rust/crates/core/src/chat_completions/route.rs index 9da86dfa27d..41a47b3bf70 100644 --- a/litellm-rust/crates/core/src/chat_completions/route.rs +++ b/litellm-rust/crates/core/src/chat_completions/route.rs @@ -55,7 +55,7 @@ impl ChatCompletionsRoute { pub(super) async fn run_call( &self, call: ChatCompletionsCall, - cache_options: Option, + cache_options: Option, interceptors: &impl litellm_host::interceptors::Interceptors, observers: Option<&ObservationSender>, ) -> Result { diff --git a/litellm-rust/crates/core/src/context.rs b/litellm-rust/crates/core/src/context.rs new file mode 100644 index 00000000000..caadf66cfb6 --- /dev/null +++ b/litellm-rust/crates/core/src/context.rs @@ -0,0 +1,48 @@ +use litellm_cache_response::CachePolicy; +use litellm_host::{ + interceptors::{ExecutionFacts, Interceptors, RawResponse}, + lifecycle::{CallEvent, ExecutionEvent}, + observation::ObservationSender, +}; + +use crate::{CallOptions, RouteError}; + +pub(crate) struct CallContext<'a, I> { + pub interceptors: &'a I, + pub observers: Option, + pub cache: CachePolicy, +} + +impl<'a, I: Interceptors> CallContext<'a, I> { + pub fn new(interceptors: &'a I, options: CallOptions) -> Self { + Self { + interceptors, + observers: options.observers, + cache: options.cache.unwrap_or_default(), + } + } + + pub async fn result_ready(&self, facts: ExecutionFacts) -> Result<(), RouteError> { + if let Some(observers) = &self.observers { + observers.emit(CallEvent::Execution(ExecutionEvent::ResultReady { + facts: facts.clone(), + })); + } + self.interceptors.result_ready(facts).await + } + + pub async fn response_received(&self, body: &str) -> Result<(), RouteError> { + let raw = RawResponse { + body: body.to_owned(), + }; + if let Some(observers) = &self.observers { + observers.emit(CallEvent::Execution( + ExecutionEvent::ProviderResponseReceived { raw: raw.clone() }, + )); + } + self.interceptors + .after_provider_response(raw) + .await + .map_err(RouteError::post_call) + } +} diff --git a/litellm-rust/crates/core/src/lib.rs b/litellm-rust/crates/core/src/lib.rs index dbdfc63e929..1fd38df191f 100644 --- a/litellm-rust/crates/core/src/lib.rs +++ b/litellm-rust/crates/core/src/lib.rs @@ -1,3 +1,4 @@ +mod context; mod diagnostic; pub mod audio_transcription; @@ -16,7 +17,7 @@ pub use error::RouteError; #[derive(Clone, Default)] pub struct CallOptions { - pub cache: Option, + pub cache: Option, pub observers: Option, } @@ -29,8 +30,8 @@ impl From> for CallOptions } } -impl From for CallOptions { - fn from(cache: litellm_cache_response::CacheOptions) -> Self { +impl From for CallOptions { + fn from(cache: litellm_cache_response::CachePolicy) -> Self { Self { cache: Some(cache), observers: None, diff --git a/litellm-rust/crates/core/src/messages/handler.rs b/litellm-rust/crates/core/src/messages/handler.rs index 5194156b6eb..49df46ef512 100644 --- a/litellm-rust/crates/core/src/messages/handler.rs +++ b/litellm-rust/crates/core/src/messages/handler.rs @@ -1,10 +1,8 @@ -use litellm_host::{lifecycle::ExecutionEvent, observation::ObservationSender}; use std::time::Duration; use bytes::Bytes; use futures_util::{StreamExt, TryStreamExt, stream::BoxStream}; -use litellm_auth::AuthServices; -use litellm_host::interceptors::{Interceptors, RawResponse, RequestContext, WireRequest}; +use litellm_host::interceptors::{Interceptors, ProviderIdentity, RequestContext, WireRequest}; use litellm_http::transport::Error as TransportError; use litellm_llms::base_llm::{ auth::{Authenticated, resolve_auth}, @@ -18,102 +16,130 @@ use litellm_tracing::ByteChunk; use serde_json::Value; use super::{ - Error, MessagesCallResponse, common_utils::truncate_error_body, + Error, MessagesCallResponse, MessagesRoute, common_utils::truncate_error_body, prepare::ProviderMessagesRequest, }; -use crate::{constants::MESSAGES_TIMEOUT_SECS, outbound::outbound_request}; +use crate::{constants::MESSAGES_TIMEOUT_SECS, context::CallContext, outbound::outbound_request}; -pub(super) async fn execute( - http: &litellm_http::Client, - auth: &AuthServices, - request: ProviderMessagesRequest, - cache: Option, - cache_options: Option, - interceptors: &impl Interceptors, - observers: Option<&ObservationSender>, -) -> Result { - let ProviderMessagesRequest { - provider, - url, - body, - environment, - timeout, - api_key, - } = request; - let stream = body.params.stream == Some(true); - let context = RequestContext { - model: body.model.clone(), - custom_llm_provider: provider.as_str().to_string(), - optional_params: serde_json::to_value(&body.params).map_err(serialize_failure)?, - secret_fields: Vec::new(), - api_key, - }; - let authenticated = resolve_auth(auth, environment, &|key| std::env::var(key).ok()).await?; - let identity = litellm_host::interceptors::ProviderIdentity { - model: context.model.clone(), - provider: context.custom_llm_provider.clone(), - }; - let wire = interceptors - .before_provider_request( - WireRequest { - url, - headers: authenticated.headers, - body: serde_json::to_value(&body).map_err(serialize_failure)?, - }, - context, - ) - .await?; - let cache = cache.filter(|_| authenticated.signer.is_none()); - let cache_request = - crate::caching::CacheRequest::from_wire(identity, cache.as_ref().map(|_| &wire)); - crate::caching::execute_streaming::( - cache_request, - cache.as_ref().map(|cache| cache.service.clone()), - cache.as_ref().map(|cache| cache.options(cache_options)), - interceptors, - observers, - || async move { - let provider_name = provider.as_str(); - log_request_body(provider_name, stream, &wire.body); - let response = send( - http, - Authenticated { - headers: wire.headers, - signer: authenticated.signer, +pub(super) struct ProviderCall { + pub identity: ProviderIdentity, + pub wire: WireRequest, + provider: super::common_utils::MessagesProvider, + signer: Option, + timeout: Option, + stream: bool, +} + +impl ProviderCall { + pub fn cacheable(&self) -> bool { + self.signer.is_none() + } +} + +impl MessagesRoute { + pub(super) async fn prepare_outbound( + &self, + request: ProviderMessagesRequest, + context: &CallContext<'_, impl Interceptors>, + ) -> Result { + let ProviderMessagesRequest { + provider, + url, + body, + environment, + timeout, + api_key, + } = request; + let request_context = RequestContext { + model: body.model.clone(), + custom_llm_provider: provider.as_str().to_string(), + optional_params: serde_json::to_value(&body.params).map_err(serialize_failure)?, + secret_fields: Vec::new(), + api_key, + }; + let authenticated = + resolve_auth(&self.auth, environment, &|key| std::env::var(key).ok()).await?; + let identity = ProviderIdentity { + model: request_context.model.clone(), + provider: request_context.custom_llm_provider.clone(), + }; + let wire = context + .interceptors + .before_provider_request( + WireRequest { + url, + headers: authenticated.headers, + body: serde_json::to_value(&body).map_err(serialize_failure)?, }, - &wire.url, - &wire.body, - timeout, + request_context, ) .await?; - if !response.status().is_success() { - return Err(provider_error(response).await); - } - let config = provider.config(); - if stream { - return Ok(streaming_response( - response, - config.stream_decoder(), - provider_name, + let stream = match wire.body.get("stream") { + None | Some(Value::Null) => false, + Some(Value::Bool(stream)) => *stream, + Some(value) => { + return Err(Error::InvalidRequest( + litellm_llms::ErrorDetail::InvalidValue { + field: "stream", + expected: "a boolean", + actual: value.clone(), + }, )); } - let text = response.text().await.map_err(network)?; - log_response_body(&text); - let raw = RawResponse { body: text.clone() }; - if let Some(observers) = observers { - observers.emit(litellm_host::lifecycle::CallEvent::Execution( - ExecutionEvent::ProviderResponseReceived { raw: raw.clone() }, - )); - } - interceptors - .after_provider_response(raw) - .await - .map_err(Error::post_call)?; - decode_response(config, &body.model, &text) - .map(|message| MessagesCallResponse::Complete(Box::new(message))) - }, - ) - .await + }; + Ok(ProviderCall { + identity, + wire, + provider, + signer: authenticated.signer, + timeout, + stream, + }) + } + + pub(super) async fn call_provider( + &self, + request: ProviderCall, + context: &CallContext<'_, impl Interceptors>, + ) -> Result { + let ProviderCall { + identity, + wire, + provider, + signer, + timeout, + stream, + } = request; + let provider_name = provider.as_str(); + log_request_body(provider_name, stream, &wire.body); + let response = send( + &self.http, + Authenticated { + headers: wire.headers, + signer, + }, + &wire.url, + &wire.body, + timeout, + ) + .await?; + if !response.status().is_success() { + return Err(provider_error(response).await); + } + let config = provider.config(); + if stream { + return Ok(streaming_response( + response, + config.stream_decoder(), + provider_name, + )); + } + let text = response.text().await.map_err(network)?; + log_response_body(&text); + context.response_received(&text).await?; + decode_response(config, &identity.model, &text) + .map(|message| MessagesCallResponse::Complete(Box::new(message))) + } } fn serialize_failure(err: serde_json::Error) -> Error { diff --git a/litellm-rust/crates/core/src/messages/mod.rs b/litellm-rust/crates/core/src/messages/mod.rs index 8374e2ca32d..8d0586cc0d9 100644 --- a/litellm-rust/crates/core/src/messages/mod.rs +++ b/litellm-rust/crates/core/src/messages/mod.rs @@ -1,11 +1,14 @@ -use litellm_host::observation::ObservationSender; mod common_utils; mod handler; mod prepare; pub mod route; mod types; +use futures_util::FutureExt; use litellm_auth::AuthServices; +use litellm_host::interceptors::{ExecutionFacts, Interceptors, ResultSource}; + +use crate::{caching::CallCache, context::CallContext}; use litellm_secrets::source::SecretSource; use std::sync::Arc; @@ -20,74 +23,18 @@ pub struct MessagesRoute { cache: Option, } -#[must_use] -#[derive(Clone, Default)] -pub struct MessagesRouteBuilder { - http: Http, - auth: Auth, - secrets: Secrets, - cache: Option, -} - -impl MessagesRouteBuilder { - pub fn with_http( - self, - http: litellm_http::Client, - ) -> MessagesRouteBuilder { - MessagesRouteBuilder { - http, - auth: self.auth, - secrets: self.secrets, - cache: self.cache, - } - } - - pub fn with_auth( - self, - auth: Arc, - ) -> MessagesRouteBuilder, Secrets> { - MessagesRouteBuilder { - http: self.http, - auth, - secrets: self.secrets, - cache: self.cache, - } - } - - pub fn with_secrets( - self, - secrets: Arc, - ) -> MessagesRouteBuilder> { - MessagesRouteBuilder { - http: self.http, - auth: self.auth, - secrets, - cache: self.cache, - } - } - - pub fn with_cache(self, cache: litellm_cache_response::ScopedCache) -> Self { - Self { - cache: Some(cache), - ..self - } - } -} - -impl MessagesRouteBuilder, Arc> { - pub fn build(self) -> MessagesRoute { - MessagesRoute { - http: self.http, - auth: self.auth, - secrets: self.secrets, - cache: self.cache, - } - } -} - impl MessagesRoute { - pub fn builder() -> MessagesRouteBuilder { - MessagesRouteBuilder::default() + pub fn new( + http: litellm_http::Client, + auth: Arc, + secrets: Arc, + ) -> Self { + Self { + http, + auth, + secrets, + cache: None, + } } #[must_use] @@ -104,15 +51,9 @@ impl MessagesRoute { interceptors: &impl litellm_host::interceptors::Interceptors, options: impl Into, ) -> Result { - let crate::CallOptions { - cache: cache_options, - observers, - } = options.into(); - litellm_host::lifecycle::observe_call( - observers.clone(), - self.run(call, cache_options, interceptors, observers.as_ref()), - ) - .await + let context = CallContext::new(interceptors, options.into()); + litellm_host::lifecycle::observe_call(context.observers.clone(), self.run(call, context)) + .await } #[tracing::instrument(name = "litellm.route", skip_all, fields( @@ -126,36 +67,34 @@ impl MessagesRoute { async fn run( &self, call: MessagesCall, - cache_options: Option, - interceptors: &impl litellm_host::interceptors::Interceptors, - observers: Option<&ObservationSender>, + context: CallContext<'_, impl Interceptors>, ) -> Result { crate::diagnostic::call(async { - self.run_provider(call, cache_options, interceptors, observers) - .await + let prepared = prepare::prepare(call, self.secrets.as_ref()).await?; + crate::diagnostic::provider(&prepared.body.model, prepared.provider.as_str()); + let request = self.prepare_outbound(prepared, &context).boxed().await?; + let cache = CallCache::::from_wire( + self.cache.as_ref().filter(|_| request.cacheable()), + context.cache, + &request.identity, + &request.wire, + ); + let identity = request.identity.clone(); + let (output, source) = match cache.lookup().await { + Some(hit) => hit, + None => ( + self.call_provider(request, &context).await?, + ResultSource::Provider, + ), + }; + context + .result_ready(ExecutionFacts { + provider: identity, + source: source.clone(), + }) + .await?; + Ok(cache.finish(output, &source).await) }) .await } - - async fn run_provider( - &self, - call: MessagesCall, - cache_options: Option, - interceptors: &impl litellm_host::interceptors::Interceptors, - observers: Option<&ObservationSender>, - ) -> Result { - let request = prepare::prepare(call, self.secrets.as_ref()).await?; - crate::diagnostic::provider(&request.body.model, request.provider.as_str()); - let execute: futures_util::future::BoxFuture<'_, Result> = - Box::pin(handler::execute( - &self.http, - &self.auth, - request, - self.cache.clone(), - cache_options, - interceptors, - observers, - )); - execute.await - } } diff --git a/litellm-rust/crates/core/src/messages/route.rs b/litellm-rust/crates/core/src/messages/route.rs index 85aa6c0995a..56060fd9d1c 100644 --- a/litellm-rust/crates/core/src/messages/route.rs +++ b/litellm-rust/crates/core/src/messages/route.rs @@ -43,8 +43,14 @@ impl super::MessagesRoute { request, observers, move |call, _, interceptors, observers| async move { - self.run(call, cache_options, &interceptors, observers.as_ref()) - .await + let context = crate::context::CallContext::new( + &interceptors, + crate::CallOptions { + cache: cache_options, + observers, + }, + ); + self.run(call, context).await }, ) } diff --git a/litellm-rust/crates/core/src/responses/handler.rs b/litellm-rust/crates/core/src/responses/handler.rs index b90e6af594a..a4b89c19d8a 100644 --- a/litellm-rust/crates/core/src/responses/handler.rs +++ b/litellm-rust/crates/core/src/responses/handler.rs @@ -16,7 +16,7 @@ pub(super) async fn execute( auth: &litellm_auth::AuthServices, request: ProviderResponsesRequest, cache: Option, - cache_options: Option, + cache_options: Option, interceptors: &impl Interceptors, observers: Option<&ObservationSender>, ) -> Result { diff --git a/litellm-rust/crates/core/src/responses/mod.rs b/litellm-rust/crates/core/src/responses/mod.rs index f388df25c7a..fb48050184f 100644 --- a/litellm-rust/crates/core/src/responses/mod.rs +++ b/litellm-rust/crates/core/src/responses/mod.rs @@ -71,7 +71,7 @@ impl ResponsesRoute { async fn run( &self, call: ResponsesCall, - cache_options: Option, + cache_options: Option, interceptors: &impl litellm_host::interceptors::Interceptors, observers: Option<&ObservationSender>, ) -> Result { @@ -85,7 +85,7 @@ impl ResponsesRoute { async fn run_provider( &self, call: ResponsesCall, - cache_options: Option, + cache_options: Option, interceptors: &impl Interceptors, observers: Option<&ObservationSender>, ) -> Result { diff --git a/litellm-rust/crates/core/tests/caching.rs b/litellm-rust/crates/core/tests/caching.rs index 51fab8b163b..4f4f6e20ad6 100644 --- a/litellm-rust/crates/core/tests/caching.rs +++ b/litellm-rust/crates/core/tests/caching.rs @@ -12,8 +12,8 @@ use bytes::Bytes; use futures_util::{StreamExt, TryStreamExt, stream}; use litellm_cache_memory::InMemoryCache; use litellm_cache_response::{ - CacheOptions, CacheScope, ResponseCache, ResponseCacheConfig, ResponseCacheService, - ResponseEnvelope, + CacheOptions, CachePolicy, CacheScope, ResponseCache, ResponseCacheConfig, + ResponseCacheService, ResponseEnvelope, }; use litellm_core::{ RouteError, @@ -117,9 +117,9 @@ async fn call( #[rstest] #[case::normal(CacheOptions::new(CacheScope::Shared), true, true)] -#[case::no_cache(CacheOptions { no_cache: true, ..CacheOptions::new(CacheScope::Shared) }, false, true)] -#[case::no_store(CacheOptions { no_store: true, ..CacheOptions::new(CacheScope::Shared) }, true, false)] -#[case::disabled(CacheOptions { caching: Some(false), ..CacheOptions::new(CacheScope::Shared) }, false, false)] +#[case::no_cache(CacheOptions { policy: CachePolicy { no_cache: true, ..CachePolicy::default() }, ..CacheOptions::new(CacheScope::Shared) }, false, true)] +#[case::no_store(CacheOptions { policy: CachePolicy { no_store: true, ..CachePolicy::default() }, ..CacheOptions::new(CacheScope::Shared) }, true, false)] +#[case::disabled(CacheOptions { policy: CachePolicy { caching: Some(false), ..CachePolicy::default() }, ..CacheOptions::new(CacheScope::Shared) }, false, false)] #[tokio::test] async fn cache_controls_apply_to_both_reads_and_writes( cache: Arc, @@ -723,9 +723,9 @@ async fn unary_call( #[rstest] #[case::normal(CacheOptions::new(CacheScope::Shared), true, true)] -#[case::no_cache(CacheOptions { no_cache: true, ..CacheOptions::new(CacheScope::Shared) }, false, true)] -#[case::no_store(CacheOptions { no_store: true, ..CacheOptions::new(CacheScope::Shared) }, true, false)] -#[case::disabled(CacheOptions { caching: Some(false), ..CacheOptions::new(CacheScope::Shared) }, false, false)] +#[case::no_cache(CacheOptions { policy: CachePolicy { no_cache: true, ..CachePolicy::default() }, ..CacheOptions::new(CacheScope::Shared) }, false, true)] +#[case::no_store(CacheOptions { policy: CachePolicy { no_store: true, ..CachePolicy::default() }, ..CacheOptions::new(CacheScope::Shared) }, true, false)] +#[case::disabled(CacheOptions { policy: CachePolicy { caching: Some(false), ..CachePolicy::default() }, ..CacheOptions::new(CacheScope::Shared) }, false, false)] #[tokio::test] async fn unary_cache_controls_do_not_change_the_shared_service( cache: Arc, diff --git a/litellm-rust/crates/core/tests/messages/host.rs b/litellm-rust/crates/core/tests/messages/host.rs index 6c6bf144238..46ac2a634c5 100644 --- a/litellm-rust/crates/core/tests/messages/host.rs +++ b/litellm-rust/crates/core/tests/messages/host.rs @@ -1,9 +1,9 @@ use litellm_host::lifecycle::ExecutionEvent; use std::sync::Mutex; -use litellm_core::messages::route::Messages; +use litellm_core::messages::{MessagesCallResponse, route::Messages}; use litellm_host::{ - interceptors::{RequestContext, WireRequest}, + interceptors::{ExecutionFacts, RequestContext, ResultSource, WireRequest}, lifecycle::CallEvent, }; use litellm_llms::base_llm::messages::context::MessagesModelCapabilities as AnthropicModelCapabilities; @@ -20,6 +20,8 @@ struct RecordingHost { rewrite: Rewrite, events: super::support::Observations, optional_params: Mutex>, + facts: Mutex>, + reject_result: bool, } impl RecordingHost { @@ -29,6 +31,8 @@ impl RecordingHost { rewrite, events: super::support::Observations::default(), optional_params: Mutex::new(Vec::new()), + facts: Mutex::new(Vec::new()), + reject_result: false, } } @@ -73,6 +77,14 @@ impl litellm_host::lifecycle::CallObserver for RecordingHost { impl litellm_host::interceptors::Interceptors<::Error> for RecordingHost { + async fn result_ready(&self, facts: ExecutionFacts) -> Result<(), Error> { + self.facts.lock().unwrap().push(facts); + if self.reject_result { + return Err(Error::Unsupported("result rejected")); + } + Ok(()) + } + async fn before_provider_request( &self, wire: WireRequest, @@ -98,6 +110,91 @@ impl litellm_host::interceptors::Interceptors< Ok(()), + Ok(MessagesCallResponse::Stream { chunks, .. }) => { + chunks.try_collect::>().await.map(|_| ()) + } + Err(error) => Err(error), + } + }; + assert_eq!( + result, + if reject { + Err(Error::Unsupported("result rejected")) + } else { + Ok(()) + } + ); + assert_eq!(received(&upstream).await.len(), expected_requests); + let facts = host.facts.lock().unwrap(); + assert_eq!(facts.len(), 1); + assert_eq!( + matches!(facts[0].source, ResultSource::Cache { .. }), + cached + ); + } +} + async fn run_through(host: &RecordingHost) -> Result { litellm_host_native::in_process::run_hosted( machine(Arc::new(RecordingSecrets::empty()))(host.request()?), @@ -143,6 +240,81 @@ async fn what_before_send_returns_is_what_the_provider_receives(call: MessagesCa assert_eq!(request.header("x-api-key"), Some("sk-ant")); } +#[rstest] +#[case::enable(false, json!(true), Some(true))] +#[case::disable(true, json!(false), Some(false))] +#[case::null(true, Value::Null, Some(false))] +#[case::invalid(false, json!("true"), None)] +#[tokio::test] +async fn response_mode_follows_the_intercepted_request( + call: MessagesCall, + traces: TraceCapture, + #[case] original_stream: bool, + #[case] rewritten_stream: Value, + #[case] expected_stream: Option, +) { + use futures_util::TryStreamExt; + + let sse = "event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n"; + let response = if expected_stream == Some(true) { + ResponseTemplate::new(200).set_body_raw(sse, "text/event-stream") + } else { + message_response() + }; + let upstream = upstream([response]).await; + let rewrite = rewritten_stream.clone(); + let host = RecordingHost::new( + authenticated( + with_fields(call, json!({"stream": original_stream})), + upstream.uri(), + ), + Box::new(move |wire| { + let mut body = wire.body; + body["stream"] = rewrite.clone(); + Ok(WireRequest { body, ..wire }) + }), + ); + let result = traces + .logger() + .instrument(async { + let output = messages_route(no_secrets()) + .execute(host.request()?, &host, None) + .await?; + match output { + MessagesCallResponse::Stream { chunks, .. } => { + assert_eq!(expected_stream, Some(true)); + assert_eq!( + chunks.try_collect::>().await?.concat(), + sse.as_bytes() + ); + } + MessagesCallResponse::Complete(message) => { + assert_eq!(expected_stream, Some(false)); + assert_eq!(*message, serde_json::from_value(message_body()).unwrap()); + } + } + Ok::<_, Error>(()) + }) + .await; + + let summaries = traces.summaries("litellm.route"); + assert_eq!(summaries.len(), 1); + let Some(expected_stream) = expected_stream else { + assert!(matches!(result, Err(Error::InvalidRequest(_)))); + assert!(received(&upstream).await.is_empty()); + assert_eq!(summaries[0]["outcome"], "failure"); + return; + }; + result.unwrap(); + assert_eq!( + only_request(&upstream).await.json()["stream"], + rewritten_stream + ); + assert_eq!(host.raw_responses().len(), usize::from(!expected_stream)); + assert_eq!(summaries[0]["stream"], expected_stream); + assert_eq!(summaries[0]["outcome"], "success"); +} + #[rstest] #[tokio::test] async fn a_before_send_failure_never_sends(call: MessagesCall) { diff --git a/litellm-rust/crates/core/tests/messages/response.rs b/litellm-rust/crates/core/tests/messages/response.rs index 42d3596e56a..6d91e5e242d 100644 --- a/litellm-rust/crates/core/tests/messages/response.rs +++ b/litellm-rust/crates/core/tests/messages/response.rs @@ -269,25 +269,22 @@ async fn the_facade_sends_through_the_injected_http_pool_configuration(call: Mes }; let resources = support::resources(); - let response = litellm_core::messages::MessagesRoute::builder() - .with_http(provider_http( - &resources, - &Resolution::from(&settings).config, - )) - .with_auth(resources.auth) - .with_secrets(no_secrets()) - .build() - .execute( - MessagesCall { - api_key: Some("sk-ant".into()), - api_base: Some(base), - ..call - }, - &(), - None, - ) - .await - .expect("messages request succeeds"); + let response = litellm_core::messages::MessagesRoute::new( + provider_http(&resources, &Resolution::from(&settings).config), + resources.auth, + no_secrets(), + ) + .execute( + MessagesCall { + api_key: Some("sk-ant".into()), + api_base: Some(base), + ..call + }, + &(), + None, + ) + .await + .expect("messages request succeeds"); let MessagesCallResponse::Complete(message) = response else { panic!("a non-streaming request returns a message"); @@ -345,7 +342,7 @@ async fn message_route_summary_excludes_payload_diagnostics( #[case::uncached(false, 2)] #[case::cached(true, 1)] #[tokio::test] -async fn builder_preserves_dependencies_and_optional_cache( +async fn route_uses_injected_dependencies_and_optional_cache( #[case] caching: bool, #[case] expected_requests: usize, ) { @@ -355,9 +352,13 @@ async fn builder_preserves_dependencies_and_optional_cache( let upstream = upstream([message_response(), message_response()]).await; let resources = resources(); - let builder = MessagesRoute::builder(); - let builder = if caching { - builder.with_cache(ScopedCache::new( + let route = MessagesRoute::new( + provider_http(&resources, &http_config()), + resources.auth.clone(), + Arc::new(RecordingSecrets::new([("ANTHROPIC_API_KEY", "route-key")])), + ); + let route = if caching { + route.with_cache(ScopedCache::new( Arc::new(ResponseCache::new(Arc::new(InMemoryCache::new( Some(100), Some(Duration::from_secs(60)), @@ -365,16 +366,8 @@ async fn builder_preserves_dependencies_and_optional_cache( CacheScope::Shared, )) } else { - builder + route }; - let route = builder - .with_secrets(Arc::new(RecordingSecrets::new([( - "ANTHROPIC_API_KEY", - "builder-key", - )]))) - .with_auth(resources.auth.clone()) - .with_http(provider_http(&resources, &http_config())) - .build(); for _ in 0..2 { let request = MessagesCall { api_base: Some(upstream.uri()), @@ -392,5 +385,72 @@ async fn builder_preserves_dependencies_and_optional_cache( } let requests = received(&upstream).await; assert_eq!(requests.len(), expected_requests); - assert_eq!(requests[0].header("x-api-key"), Some("builder-key")); + assert_eq!(requests[0].header("x-api-key"), Some("route-key")); +} + +#[rstest] +#[tokio::test] +async fn cache_overrides_preserve_the_routes_isolated_scope(call: MessagesCall) { + use litellm_cache_memory::InMemoryCache; + use litellm_cache_response::{CachePolicy, CacheScope, ResponseCache, ScopedCache}; + + let first_body = message_body(); + let second_body = Value::Object( + first_body + .as_object() + .unwrap() + .iter() + .map(|(key, value)| { + ( + key.clone(), + if key == "id" { + json!("msg_second") + } else { + value.clone() + }, + ) + }) + .collect(), + ); + let upstream = upstream([ + json_response(first_body.clone()), + json_response(second_body.clone()), + ]) + .await; + let service = Arc::new(ResponseCache::new(Arc::new(InMemoryCache::new( + Some(100), + Some(Duration::from_secs(60)), + )))); + let first = messages_route(no_secrets()).with_cache(ScopedCache::new( + service.clone(), + CacheScope::Isolated("first".into()), + )); + let second = messages_route(no_secrets()).with_cache(ScopedCache::new( + service, + CacheScope::Isolated("second".into()), + )); + for (route, expected) in [ + (&first, &first_body), + (&second, &second_body), + (&first, &first_body), + (&second, &second_body), + ] { + let request = MessagesCall { + body: call.body.clone(), + api_key: Some("same-key".into()), + api_base: Some(upstream.uri()), + ..super::call() + }; + let override_options = CachePolicy { + ttl: Some(Duration::from_secs(30)), + ..CachePolicy::default() + }; + let MessagesCallResponse::Complete(response) = + route.execute(request, &(), override_options).await.unwrap() + else { + panic!("expected a completed message"); + }; + assert_eq!(response.id, expected["id"].as_str().unwrap()); + } + assert_eq!(received(&upstream).await.len(), 2); } diff --git a/litellm-rust/crates/core/tests/support/mod.rs b/litellm-rust/crates/core/tests/support/mod.rs index 5ba1eb3ca46..1dd53114293 100644 --- a/litellm-rust/crates/core/tests/support/mod.rs +++ b/litellm-rust/crates/core/tests/support/mod.rs @@ -44,11 +44,11 @@ pub fn provider_http( pub fn messages_route(secrets: Arc) -> litellm_core::messages::MessagesRoute { let resources = resources(); - litellm_core::messages::MessagesRoute::builder() - .with_http(provider_http(&resources, &http_config())) - .with_auth(resources.auth) - .with_secrets(secrets) - .build() + litellm_core::messages::MessagesRoute::new( + provider_http(&resources, &http_config()), + resources.auth, + secrets, + ) } pub fn chat_completions_route() -> litellm_core::chat_completions::ChatCompletionsRoute { diff --git a/litellm-rust/crates/gateway-inference/src/caching.rs b/litellm-rust/crates/gateway-inference/src/caching.rs index 020f942ec19..5472baf115f 100644 --- a/litellm-rust/crates/gateway-inference/src/caching.rs +++ b/litellm-rust/crates/gateway-inference/src/caching.rs @@ -1,6 +1,6 @@ use std::time::Duration; -use litellm_cache_response::{CacheOptions, CacheScope}; +use litellm_cache_response::{CacheOptions, CachePolicy, CacheScope}; use litellm_gateway_auth::AuthenticatedRequest; use serde::Deserialize; use serde_json::{Map, Value}; @@ -38,11 +38,13 @@ pub(crate) fn prepare( .map_err(|error| Error::InvalidBody(error.to_string()))?; let caller = identity.caller(); let options = CacheOptions { - caching, - no_cache: controls.no_cache, - no_store: controls.no_store, - ttl: controls.ttl.map(duration).transpose()?, - max_age: controls.max_age.map(duration).transpose()?, + policy: CachePolicy { + caching, + no_cache: controls.no_cache, + no_store: controls.no_store, + ttl: controls.ttl.map(duration).transpose()?, + max_age: controls.max_age.map(duration).transpose()?, + }, scope: CacheScope::Isolated( serde_json::json!([ caller.principal().authority(), diff --git a/litellm-rust/crates/gateway-inference/src/chat_completions.rs b/litellm-rust/crates/gateway-inference/src/chat_completions.rs index 27b8e856b7e..c9fabca2522 100644 --- a/litellm-rust/crates/gateway-inference/src/chat_completions.rs +++ b/litellm-rust/crates/gateway-inference/src/chat_completions.rs @@ -69,7 +69,7 @@ async fn handle( extra_headers: None, timeout: deployment.timeout, }, - cache_options, + cache_options.policy, ), (), headers.clone(), diff --git a/litellm-rust/crates/gateway-inference/src/lib.rs b/litellm-rust/crates/gateway-inference/src/lib.rs index b669fb1ced1..a8a70ffcefc 100644 --- a/litellm-rust/crates/gateway-inference/src/lib.rs +++ b/litellm-rust/crates/gateway-inference/src/lib.rs @@ -68,11 +68,7 @@ impl Gateway { auth.clone(), secrets.clone(), ), - messages: MessagesRoute::builder() - .with_http(provider.clone()) - .with_auth(auth.clone()) - .with_secrets(secrets.clone()) - .build(), + messages: MessagesRoute::new(provider.clone(), auth.clone(), secrets.clone()), responses: ResponsesRoute::new(provider, auth.clone(), secrets.clone()), ocr: OcrRoute::new(OcrClient::new( &resources.pool, diff --git a/litellm-rust/crates/gateway-inference/src/messages.rs b/litellm-rust/crates/gateway-inference/src/messages.rs index 0f1d1d31689..1a08946044b 100644 --- a/litellm-rust/crates/gateway-inference/src/messages.rs +++ b/litellm-rust/crates/gateway-inference/src/messages.rs @@ -54,7 +54,7 @@ async fn handle( }; let call = project(deployment, body, headers)?; - let machine = route.machine(call, cache_options); + let machine = route.machine(call, cache_options.policy); let stream = Sse::::new(Json, |error| Bytes::from(Error::from(error).sse_frame())); let headers = crate::caching::CacheHeaders::default(); diff --git a/litellm-rust/crates/gateway-inference/src/responses.rs b/litellm-rust/crates/gateway-inference/src/responses.rs index 5b324d74172..7a690e4e4c0 100644 --- a/litellm-rust/crates/gateway-inference/src/responses.rs +++ b/litellm-rust/crates/gateway-inference/src/responses.rs @@ -38,7 +38,7 @@ pub(crate) async fn create( extra_headers: None, timeout: deployment.timeout, }; - let machine = route.machine(call, cache_options); + let machine = route.machine(call, cache_options.policy); let stream = Sse::::new(Json, |error| { let error = Error::from(error); Bytes::from(format!( diff --git a/litellm-rust/crates/python-bridge/src/cache/native/v2.rs b/litellm-rust/crates/python-bridge/src/cache/native/v2.rs index e355d0d698a..654de75d6bb 100644 --- a/litellm-rust/crates/python-bridge/src/cache/native/v2.rs +++ b/litellm-rust/crates/python-bridge/src/cache/native/v2.rs @@ -330,15 +330,17 @@ pub(in crate::cache) fn configured( Ok(( Some(cache.service.clone()), litellm_cache_response::CacheOptions { - caching: kwargs - .get_item("caching")? - .filter(|value| !value.is_none()) - .map(|value| value.extract()) - .transpose()?, - no_cache: boolean("no-cache")?, - no_store: boolean("no-store")?, - ttl: seconds("ttl")?, - max_age: seconds("s-max-age")?.or(seconds("s-maxage")?), + policy: litellm_cache_response::CachePolicy { + caching: kwargs + .get_item("caching")? + .filter(|value| !value.is_none()) + .map(|value| value.extract()) + .transpose()?, + no_cache: boolean("no-cache")?, + no_store: boolean("no-store")?, + ttl: seconds("ttl")?, + max_age: seconds("s-max-age")?.or(seconds("s-maxage")?), + }, scope: litellm_cache_response::CacheScope::Shared, }, )) diff --git a/litellm-rust/crates/python-bridge/src/cache/selection.rs b/litellm-rust/crates/python-bridge/src/cache/selection.rs index 723ff703f64..e6e9f8d2d4f 100644 --- a/litellm-rust/crates/python-bridge/src/cache/selection.rs +++ b/litellm-rust/crates/python-bridge/src/cache/selection.rs @@ -1,5 +1,7 @@ use super::{native, python}; -use litellm_cache_response::{CacheOptions, CacheScope, ResponseCacheService, ScopedCache}; +use litellm_cache_response::{ + CacheOptions, CachePolicy, CacheScope, ResponseCacheService, ScopedCache, +}; use litellm_host::{ machine::{HostServices, MachineFault}, protocol::Protocol, @@ -154,8 +156,11 @@ pub(crate) fn configure( .map(|value| value.unwrap_or(false)) }; let options = CacheOptions { - no_cache: boolean("no-cache")?, - no_store: boolean("no-store")?, + policy: CachePolicy { + no_cache: boolean("no-cache")?, + no_store: boolean("no-store")?, + ..CachePolicy::default() + }, ..CacheOptions::new(CacheScope::Shared) }; let namespace = cache diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs index fdd7be58a35..bf1d1645c0a 100644 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs @@ -175,7 +175,7 @@ fn run_public( )), None => route, }; - Ok(route.machine(request, cache_options)) + Ok(route.machine(request, cache_options.policy)) }, host::ChatCompletionsPythonHost(host), hooks, diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs index 7ed5375f265..4838c973a34 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs @@ -26,14 +26,12 @@ fn run_messages( py, arguments, move |py, arguments, request| { - let builder = litellm_core::messages::MessagesRoute::builder() - .with_http( - crate::http::provider_client(py, arguments, asynchronous)? - .map_err(crate::http::client_error)?, - ) - .with_auth(crate::http::resources().auth.clone()) - .with_secrets(crate::secrets::source(py)?); - let route = builder.build(); + let route = litellm_core::messages::MessagesRoute::new( + crate::http::provider_client(py, arguments, asynchronous)? + .map_err(crate::http::client_error)?, + crate::http::resources().auth.clone(), + crate::secrets::source(py)?, + ); Ok(litellm_host::call::hosted_call( request, None, @@ -51,7 +49,7 @@ fn run_messages( call, &interceptors, litellm_core::CallOptions { - cache: Some(options), + cache: Some(options.policy), observers, }, ) diff --git a/litellm-rust/crates/python-bridge/src/routes/responses.rs b/litellm-rust/crates/python-bridge/src/routes/responses.rs index 1c2b685854a..bcddaa2bf8d 100644 --- a/litellm-rust/crates/python-bridge/src/routes/responses.rs +++ b/litellm-rust/crates/python-bridge/src/routes/responses.rs @@ -96,7 +96,7 @@ fn run_public( )), None => route, }; - Ok(route.machine(request, cache_options)) + Ok(route.machine(request, cache_options.policy)) }, host::ResponsesPythonHost(host), hooks, From e814532033608d6505ff03ca9cc818d69445f2bf Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 12:54:17 -0700 Subject: [PATCH 30/41] fix(streaming): keep the served service_tier on streamed chunks and spend rows (#42870) * fix(streaming): keep the provider's served service_tier on streamed chunks and spend rows Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(streaming): satisfy type-discipline and strict ruff budgets Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(streaming): stamp the served service_tier on every Responses bridge chunk Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(anthropic-adapter): expose streamed chunks so disconnects bill partial spend Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(service-tier): cover anthropic and responses served-tier billing paths Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(anthropic-adapter): return a chunks-exposing stream so disconnects bill partial spend Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(service-tier): bill disconnects through the router's anthropic stream wrapper Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * style: apply ruff format to the anthropic stream changes Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(coverage): ignore delegating properties the ast scan cannot see Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * style: keep the cast-ok reasons on the cast call line Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): cover served service_tier billing for streamed chat and messages, complete and disconnected Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(anthropic-cache): delegate chunks/messages/model through the messages stream cache writer Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(streaming): keep service_tier on OpenAI-compatible parsed chunks Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(streaming): parameterize delegated chunks and messages types Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(tests): follow the anthropic pass_through rename after merging main Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(anthropic): drain the logging worker between response cache tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(spend): cover azure, databricks, responses bridge and gemini served tiers in the stream billing integration test Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(databricks): keep the served service_tier on streamed chunks and bill it Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(databricks): type the served service_tier chunk without a loose kwargs dict Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(cost): bill the served service_tier over the requested one Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(cost): drop explanatory comment from the tier resolution Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: kerry --- .../transformation.py | 29 +- litellm/cost_calculator.py | 68 +- litellm/litellm_core_utils/litellm_logging.py | 1 - .../streaming_chunk_builder_utils.py | 14 + .../litellm_core_utils/streaming_handler.py | 8 + .../adapters/streaming_iterator.py | 60 ++ .../pass_through/adapters/transformation.py | 4 +- .../pass_through/messages/response_cache.py | 22 +- .../llms/databricks/chat/transformation.py | 11 + litellm/llms/databricks/cost_calculator.py | 3 +- .../llms/openai/chat/gpt_transformation.py | 3 + litellm/proxy/proxy_server.py | 3 +- litellm/router.py | 18 + .../router_code_coverage.py | 3 + .../coverage_registry/llm_conversational.yaml | 1 + .../coverage_registry/quota_management.yaml | 3 + tests/e2e/models.py | 11 + tests/e2e/proxy_client.py | 4 + .../test_service_tier_pricing_e2e.py | 208 +++++- .../spend/test_service_tier_stream_billing.py | 660 ++++++++++++++++++ .../proxy_server/test_streaming_helpers.py | 10 + .../proxy/test_common_request_processing.py | 56 ++ ...responses_transformation_transformation.py | 36 + .../test_litellm_logging.py | 52 ++ .../test_streaming_chunk_builder_utils.py | 31 + .../test_streaming_handler.py | 47 ++ .../test_streaming_iterator_sse_stream.py | 87 +++ .../messages/test_response_cache.py | 41 ++ .../test_databricks_chat_transformation.py | 10 + .../test_databricks_cost_calculator.py | 27 +- .../chat/test_openai_gpt_transformation.py | 27 + tests/unit/test_cost_calculator.py | 116 ++- 32 files changed, 1621 insertions(+), 53 deletions(-) create mode 100644 tests/integration/spend/test_service_tier_stream_billing.py create mode 100644 tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_sse_stream.py diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index 31af5a144eb..71d3f1e900e 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -1334,6 +1334,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): ): super().__init__(streaming_response, sync_stream, json_mode) self._chat_completion_id: str | None = None + self._served_service_tier: str | None = None self._tool_call_index_map: dict[int, int] = {} # mutable-ok: per-stream accumulator state def _handle_string_chunk( @@ -1598,6 +1599,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(response_data.get("usage")) provider_metadata: Final = _provider_metadata(response_data) + served_service_tier: Final = response_data.get("service_tier") return ModelResponseStream( choices=[ StreamingChoices( @@ -1611,6 +1613,11 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): ], usage=usage, provider_specific_fields=dict(provider_metadata) or None, # mutable-ok: field is typed dict + **( + MappingProxyType({"service_tier": served_service_tier}) + if isinstance(served_service_tier, str) + else MappingProxyType({}) + ), ) else: pass @@ -1639,12 +1646,28 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): ModelResponseStream: OpenAI-formatted streaming chunk """ verbose_logger.debug("Chat provider: transform_streaming_response called with chunk: %s", chunk) - return self._with_stream_scoped_id( - OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream( - chunk, tool_call_index_map=self._tool_call_index_map + self._remember_served_service_tier(chunk) + return self._with_served_service_tier( + self._with_stream_scoped_id( + OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream( + chunk, tool_call_index_map=self._tool_call_index_map + ) ) ) + def _remember_served_service_tier(self, chunk: dict[str, object]) -> None: + response_payload: Final = chunk.get("response") + if not isinstance(response_payload, dict): + return + served_tier: Final = response_payload.get("service_tier") + if isinstance(served_tier, str) and served_tier: + self._served_service_tier = served_tier + + def _with_served_service_tier(self, chunk: "ModelResponseStream") -> "ModelResponseStream": + if self._served_service_tier is not None and chunk.model_dump().get("service_tier") is None: + setattr(chunk, "service_tier", self._served_service_tier) # noqa: B010 # pydantic extra, not a declared field + return chunk + def _with_stream_scoped_id(self, chunk: "ModelResponseStream") -> "ModelResponseStream": if self._chat_completion_id is None: self._chat_completion_id = chunk.id diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 2f76b84f5e3..6b3e739ac4f 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -26,6 +26,7 @@ from litellm.litellm_core_utils.llm_cost_calc.usage_object_transformation import TranscriptionUsageObjectTransformation, ) from litellm.litellm_core_utils.llm_cost_calc.utils import ( + _SERVICE_TIER_TO_COST_KEY_SUFFIX, BilledTokenRates, CostCalculatorUtils, _generic_cost_per_character, @@ -704,7 +705,7 @@ def cost_per_token( data_residency=data_residency, ) elif custom_llm_provider == "databricks": - return databricks_cost_per_token(model=model, usage=usage_block) + return databricks_cost_per_token(model=model, usage=usage_block, service_tier=service_tier) elif custom_llm_provider == "fireworks_ai": return fireworks_ai_cost_per_token(model=model, usage=usage_block) elif custom_llm_provider == "azure": @@ -969,6 +970,37 @@ def _normalize_service_tier(service_tier: object) -> str | None: return service_tier +_BASE_PRICING_SERVICE_TIERS: Final[frozenset[str]] = frozenset({"default", "standard"}) + + +def _resolve_billable_service_tier(requested: object, served: object) -> str | None: + """Served tier wins when it names a priced tier or explicitly says base pricing; otherwise the request decides.""" + served_lower: Final = served.lower() if isinstance(served, str) else None + if served_lower is not None and served_lower in _SERVICE_TIER_TO_COST_KEY_SUFFIX: + return served_lower + if served_lower in _BASE_PRICING_SERVICE_TIERS: + return None + return _normalize_service_tier(requested) + + +def _served_service_tier(completion_response: object, usage_object: Usage | None) -> str | None: + """Find the tier the provider actually served: response, then usage, then Gemini trafficType.""" + response_tier: Final = _extract_service_tier(completion_response) + if isinstance(response_tier, str): + return response_tier + usage_tier: Final = _extract_service_tier(usage_object) + if isinstance(usage_tier, str): + return usage_tier + hidden_params: Final = getattr(completion_response, "_hidden_params", None) + if hidden_params is None: + return None + provider_specific: Final = hidden_params.get("provider_specific_fields") or {} + raw_traffic_type: Final = provider_specific.get("traffic_type") + if not raw_traffic_type: + return None + return _map_traffic_type_to_service_tier(raw_traffic_type) or "default" + + def _extract_service_tier(source: object) -> str | None: """Read a raw ``service_tier`` off a response body or usage object, dict or pydantic model alike.""" if isinstance(source, BaseModel): @@ -1388,23 +1420,14 @@ def completion_cost( ) rerank_billed_units: RerankBilledUnits | None = None - # Extract service_tier from optional_params if not provided directly - if service_tier is None and optional_params is not None: - service_tier = optional_params.get("service_tier") - - service_tier = _normalize_service_tier(service_tier) - - # Extract service_tier from completion_response if not provided - if service_tier is None and completion_response is not None: - service_tier = _extract_service_tier(completion_response) - - service_tier = _normalize_service_tier(service_tier) - - # Extract service_tier from usage object if not provided - if service_tier is None and cost_per_token_usage_object is not None: - service_tier = _extract_service_tier(cost_per_token_usage_object) - - service_tier = _normalize_service_tier(service_tier) + explicit_tier: Final = _normalize_service_tier(service_tier) + if explicit_tier is not None: + service_tier = explicit_tier + else: + service_tier = _resolve_billable_service_tier( # rebind-ok: resolved from request then response + requested=optional_params.get("service_tier") if optional_params is not None else None, + served=_served_service_tier(completion_response, cost_per_token_usage_object), + ) explicit_pricing: Final = custom_pricing is True or base_model is not None selected_model: Final = _select_model_name_for_cost_calc( @@ -1494,15 +1517,6 @@ def completion_cost( custom_llm_provider = hidden_params.get("custom_llm_provider", custom_llm_provider or None) region_name = hidden_params.get("region_name", region_name) - # For Gemini/Vertex AI responses, trafficType is stored in - # provider_specific_fields. Map it to the service_tier used - # by the cost key lookup (_priority / _flex suffixes) so that - # ON_DEMAND_PRIORITY requests are billed at priority prices. - if service_tier is None: - provider_specific = hidden_params.get("provider_specific_fields") or {} - raw_traffic_type = provider_specific.get("traffic_type") - if raw_traffic_type: - service_tier = _map_traffic_type_to_service_tier(raw_traffic_type) else: if model is None: raise ValueError( diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 06cbfd4fc04..c393fa3caef 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -1884,7 +1884,6 @@ class Logging(LiteLLMLoggingBaseClass): "standard_built_in_tools_params": self.standard_built_in_tools_params, "router_model_id": router_model_id, "litellm_logging_obj": self, - "service_tier": (self.optional_params.get("service_tier") if self.optional_params else None), "data_residency": ( self.litellm_params.get("data_residency") if hasattr(self, "litellm_params") and self.litellm_params diff --git a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py index 67684a230e3..bdf53013224 100644 --- a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py +++ b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py @@ -109,6 +109,7 @@ class _BaseChunk(TypedDict, total=False): created: ReadOnly[int] model: ReadOnly[str] system_fingerprint: ReadOnly[str | None] + service_tier: ReadOnly[str | None] choices: ReadOnly[Required[Sequence[StreamingChoices]]] _hidden_params: ReadOnly[_ChunkHiddenParams] @@ -369,6 +370,13 @@ class ChunkProcessor: # Fall back to first chunk's model if no different model found return first_chunk_model + @staticmethod + def _get_service_tier_from_chunks(chunks: Sequence["_BaseChunk"]) -> str | None: + return next( + (tier for chunk in reversed(chunks) if isinstance(tier := chunk.get("service_tier"), str) and tier), + None, + ) + def build_base_response(self, chunks: Sequence["_BaseChunk"]) -> ModelResponse: chunk = self.first_chunk id: Final = ChunkProcessor._get_chunk_id(chunks) @@ -378,6 +386,7 @@ class ChunkProcessor: # Get the actual model - for Azure Model Router, this finds the real model from later chunks model: Final = ChunkProcessor._get_model_from_chunks(chunks, first_chunk_model) system_fingerprint: Final = chunk.get("system_fingerprint", None) + service_tier: Final = ChunkProcessor._get_service_tier_from_chunks(chunks) role: Final = ChunkProcessor._get_role_from_chunks(chunks) finish_reason = "stop" @@ -399,6 +408,11 @@ class ChunkProcessor: "created": created, "model": model, "system_fingerprint": system_fingerprint, + **( + MappingProxyType({"service_tier": service_tier}) + if service_tier is not None + else MappingProxyType({}) + ), "choices": [ { "index": 0, diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index 946d19c028f..d2853a625c9 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -72,6 +72,12 @@ def _next_sync_or_exhausted(it: Any) -> object: return _SYNC_ITER_EXHAUSTED +def _stamp_served_service_tier(response: ModelResponseStream, complete_streaming_response: ModelResponse) -> None: + served_tier: Final = complete_streaming_response.model_dump().get("service_tier") + if isinstance(served_tier, str) and served_tier: + setattr(response, "service_tier", served_tier) # noqa: B010 # pydantic extra, not a declared field + + def is_async_iterable(obj: object) -> bool: """ Check if an object is an async iterable (can be used with 'async for'). @@ -1876,6 +1882,7 @@ class CustomStreamWrapper: "usage", getattr(complete_streaming_response, "usage"), ) + _stamp_served_service_tier(response, complete_streaming_response) try: _cache_copy = complete_streaming_response.model_copy(deep=True) _log_copy = complete_streaming_response.model_copy(deep=True) @@ -2127,6 +2134,7 @@ class CustomStreamWrapper: "usage", getattr(complete_streaming_response, "usage"), ) + _stamp_served_service_tier(response, complete_streaming_response) try: _copy = complete_streaming_response.model_copy(deep=True) except RuntimeError: diff --git a/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py b/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py index 12eee663ca5..38380cc056d 100644 --- a/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py +++ b/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py @@ -11,6 +11,7 @@ from typing import ( Final, Literal, Protocol, + cast, get_args, ) @@ -35,6 +36,7 @@ from litellm.types.utils import AdapterCompletionStreamWrapper, Delta if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObject + from litellm.types.llms.openai import AllMessageValues from litellm.types.utils import ModelResponseStream @@ -115,6 +117,18 @@ class _CombinedChunkSplitter: self._async_iter: AsyncIterator[ModelResponseStream] | None = None self._buffer: deque[ModelResponseStream] = deque() + @property + def chunks(self) -> "list[ModelResponseStream] | None": + return cast( # cast-ok: chunks is a list of ModelResponseStream on the inner stream + "list[ModelResponseStream] | None", getattr(self._stream, "chunks", None) + ) + + @property + def messages(self) -> "list[AllMessageValues] | None": + return cast( # cast-ok: messages is a list of AllMessageValues on the inner stream + "list[AllMessageValues] | None", getattr(self._stream, "messages", None) + ) + @staticmethod def _is_combined(chunk: "ModelResponseStream") -> bool: """True if ``chunk`` carries response content AND a finish_reason.""" @@ -351,6 +365,18 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): text="", ) + @property + def chunks(self) -> "list[ModelResponseStream] | None": + return cast( # cast-ok: chunks is a list of ModelResponseStream on the inner stream + "list[ModelResponseStream] | None", getattr(self.completion_stream, "chunks", None) + ) + + @property + def messages(self) -> "list[AllMessageValues] | None": + return cast( # cast-ok: messages is a list of AllMessageValues on the inner stream + "list[AllMessageValues] | None", getattr(self.completion_stream, "messages", None) + ) + def _merge_usage_into_held_stop_reason_chunk(self, chunk: Any) -> MessageBlockDelta: """Merge usage data from ``chunk`` into the held ``message_delta`` chunk. @@ -1173,3 +1199,37 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): return True return False + + +class AnthropicSSEStream(AsyncIterator[bytes]): + """ + AsyncIterator[bytes] view of AnthropicStreamWrapper returned to callers of + translate_completion_output_params_streaming. Keeps the wrapper reachable so + the proxy's disconnect-time partial billing can read the inner chat stream's + collected chunks, messages, and model; a bare async generator would hide them. + """ + + def __init__(self, anthropic_wrapper: AnthropicStreamWrapper) -> None: + self._anthropic_wrapper = anthropic_wrapper + self._byte_stream: Final[AsyncIterator[bytes]] = anthropic_wrapper.async_anthropic_sse_wrapper() + self._hidden_params: dict[ + str, object + ] = {} # mutable-ok: the proxy merges provider headers onto _hidden_params in place + + @property + def chunks(self) -> "list[ModelResponseStream] | None": + return self._anthropic_wrapper.chunks + + @property + def messages(self) -> "list[AllMessageValues] | None": + return self._anthropic_wrapper.messages + + @property + def model(self) -> str: + return self._anthropic_wrapper.model + + async def __anext__(self) -> bytes: + return await self._byte_stream.__anext__() + + async def aclose(self) -> None: + await self._byte_stream.aclose() diff --git a/litellm/llms/anthropic/pass_through/adapters/transformation.py b/litellm/llms/anthropic/pass_through/adapters/transformation.py index 2bb081bd0a4..040c8f0e170 100644 --- a/litellm/llms/anthropic/pass_through/adapters/transformation.py +++ b/litellm/llms/anthropic/pass_through/adapters/transformation.py @@ -201,7 +201,7 @@ from litellm.types.llms.openai import ( from litellm.types.utils import Choices, ModelResponse, StreamingChoices, Usage from litellm.utils import supports_mid_conversation_system -from .streaming_iterator import AnthropicStreamWrapper +from .streaming_iterator import AnthropicSSEStream, AnthropicStreamWrapper if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObject @@ -341,7 +341,7 @@ class AnthropicAdapter: ) # Return the SSE-wrapped version for proper event formatting. if is_async: - return anthropic_wrapper.async_anthropic_sse_wrapper() + return AnthropicSSEStream(anthropic_wrapper) return anthropic_wrapper.anthropic_sse_wrapper() diff --git a/litellm/llms/anthropic/pass_through/messages/response_cache.py b/litellm/llms/anthropic/pass_through/messages/response_cache.py index 1a8b041e674..7b9435d45ac 100644 --- a/litellm/llms/anthropic/pass_through/messages/response_cache.py +++ b/litellm/llms/anthropic/pass_through/messages/response_cache.py @@ -1,7 +1,7 @@ import re from collections.abc import AsyncIterator, Mapping, Sequence from types import MappingProxyType -from typing import TYPE_CHECKING, Final +from typing import TYPE_CHECKING, Final, cast import litellm from litellm._logging import verbose_logger @@ -17,6 +17,8 @@ from litellm.llms.anthropic.pass_through.messages.streaming_iterator import ( if TYPE_CHECKING: from litellm.caching.caching_handler import LLMCachingHandler from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.types.llms.openai import AllMessageValues + from litellm.types.utils import ModelResponseStream CACHED_STREAM_EVENTS_KEY: Final = "litellm_cached_anthropic_sse_events" @@ -51,6 +53,24 @@ class AnthropicMessagesStreamCacheWriter: def has_buffered_provider_output(self) -> bool: return getattr(self.stream, "has_buffered_provider_output", False) is True + @property + def chunks(self) -> "list[ModelResponseStream] | None": + return cast( # cast-ok: chunks is a list of ModelResponseStream on the inner stream + "list[ModelResponseStream] | None", getattr(self.stream, "chunks", None) + ) + + @property + def messages(self) -> "list[AllMessageValues] | None": + return cast( # cast-ok: messages is a list of AllMessageValues on the inner stream + "list[AllMessageValues] | None", getattr(self.stream, "messages", None) + ) + + @property + def model(self) -> str | None: + return cast( # cast-ok: model is a str on the inner stream + "str | None", getattr(self.stream, "model", None) + ) + def __aiter__(self) -> "AnthropicMessagesStreamCacheWriter": return self diff --git a/litellm/llms/databricks/chat/transformation.py b/litellm/llms/databricks/chat/transformation.py index 538904b34e6..d669f2acc6d 100644 --- a/litellm/llms/databricks/chat/transformation.py +++ b/litellm/llms/databricks/chat/transformation.py @@ -778,6 +778,17 @@ class DatabricksChatResponseIterator(BaseModelResponseIterator): ) choice["delta"]["thinking_blocks"] = thinking_blocks translated_choices.append(choice) + service_tier: Final = chunk.get("service_tier") + if isinstance(service_tier, str) and service_tier: + return ModelResponseStream( + id=chunk["id"], + object="chat.completion.chunk", + created=chunk["created"], + model=chunk["model"], + choices=translated_choices, + usage=chunk.get("usage"), + service_tier=service_tier, + ) return ModelResponseStream( id=chunk["id"], object="chat.completion.chunk", diff --git a/litellm/llms/databricks/cost_calculator.py b/litellm/llms/databricks/cost_calculator.py index 64166e6fc11..2bb5b99f0ad 100644 --- a/litellm/llms/databricks/cost_calculator.py +++ b/litellm/llms/databricks/cost_calculator.py @@ -30,7 +30,7 @@ def _registry_key(model: str) -> str: ) -def cost_per_token(model: str, usage: Usage) -> tuple[float, float]: +def cost_per_token(model: str, usage: Usage, service_tier: str | None = None) -> tuple[float, float]: """ Calculates the cost per token for a given model, prompt tokens, and completion tokens. @@ -45,4 +45,5 @@ def cost_per_token(model: str, usage: Usage) -> tuple[float, float]: model=_registry_key(model), usage=usage, custom_llm_provider="databricks", + service_tier=service_tier, ) diff --git a/litellm/llms/openai/chat/gpt_transformation.py b/litellm/llms/openai/chat/gpt_transformation.py index 62351d8e39a..3b38825c83d 100644 --- a/litellm/llms/openai/chat/gpt_transformation.py +++ b/litellm/llms/openai/chat/gpt_transformation.py @@ -890,6 +890,9 @@ class OpenAIChatCompletionStreamingHandler(BaseModelResponseIterator): } if "usage" in chunk and chunk["usage"] is not None: kwargs["usage"] = chunk["usage"] + service_tier: Final = chunk.get("service_tier") + if isinstance(service_tier, str) and service_tier: + kwargs["service_tier"] = service_tier return ModelResponseStream(**kwargs) except Exception as e: raise e diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 7cd116c23d7..6a221f2bfed 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -9296,9 +9296,10 @@ def _fast_serialize_simple_model_response_stream( "object": getattr(chunk, "object", None), "created": getattr(chunk, "created", None), "model": model, + "service_tier": getattr(chunk, "service_tier", None), "choices": [choice_dict], } - for top_level_key in ("id", "object", "created"): + for top_level_key in ("id", "object", "created", "service_tier"): if payload[top_level_key] is None: payload.pop(top_level_key) return orjson.dumps(payload) diff --git a/litellm/router.py b/litellm/router.py index 842cd9de378..86a67a8d5ca 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -638,6 +638,24 @@ class FallbackAwareAnthropicMessagesStream: def has_buffered_provider_output(self) -> bool: return getattr(self._source_iterator, "has_buffered_provider_output", False) is True + @property + def chunks(self) -> list[ModelResponseStream] | None: + return cast( # cast-ok: chunks is a list of ModelResponseStream on the inner stream + "list[ModelResponseStream] | None", getattr(self._source_iterator, "chunks", None) + ) + + @property + def messages(self) -> list[AllMessageValues] | None: + return cast( # cast-ok: messages is a list of AllMessageValues on the inner stream + "list[AllMessageValues] | None", getattr(self._source_iterator, "messages", None) + ) + + @property + def model(self) -> str | None: + return cast( # cast-ok: model is a str on the inner stream + "str | None", getattr(self._source_iterator, "model", None) + ) + def adopt_fallback_source(self, fallback_response: object) -> None: self._source_iterator = fallback_response self.fallback_headers_adopted = True diff --git a/tests/code_coverage_tests/router_code_coverage.py b/tests/code_coverage_tests/router_code_coverage.py index 06e5b020836..7c247ae3303 100644 --- a/tests/code_coverage_tests/router_code_coverage.py +++ b/tests/code_coverage_tests/router_code_coverage.py @@ -81,6 +81,9 @@ ignored_function_names = [ "_merge_tools_from_deployment", # Tested indirectly via _update_kwargs_with_deployment (test files lack "router" in name) "_invalidate_access_groups_cache", # Tested indirectly via set_model_list, upsert_model etc. (test files lack "router" in name) "has_buffered_provider_output", # Property, so its reads in test_router.py are never an ast.Call + "chunks", # Property on FallbackAwareAnthropicMessagesStream, so its reads in tests are never an ast.Call + "messages", # Property on FallbackAwareAnthropicMessagesStream, so its reads in tests are never an ast.Call + "model", # Property on FallbackAwareAnthropicMessagesStream, so its reads in tests are never an ast.Call "_request_header", # Tested through Claude Code session routing in test_router.py "_claude_code_session_router_cache_key", # Tested through Claude Code session routing in test_router.py "_delete_claude_code_session_router_binding", # Tested through Redis cleanup failure in test_router.py diff --git a/tests/e2e/coverage_registry/llm_conversational.yaml b/tests/e2e/coverage_registry/llm_conversational.yaml index cd52e563d69..8de8d5875b4 100644 --- a/tests/e2e/coverage_registry/llm_conversational.yaml +++ b/tests/e2e/coverage_registry/llm_conversational.yaml @@ -13,6 +13,7 @@ - {id: llm.chat_completions.openai.vision.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: openai, capability: vision, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "gpt-4o vision; high usage"} - {id: llm.chat_completions.openai.prompt_cache_5m.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: openai, capability: prompt_cache_5m, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Prompt caching cost optimization"} - {id: llm.chat_completions.openai.service_tier.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: openai, capability: service_tier, streaming: nonstream, assertions: [works], source: "OpenAI service_tier param", rationale: "OpenAI scale-tier request option is forwarded and echoed"} +- {id: llm.chat_completions.openai.service_tier.stream.echoes_served_tier, module: llm, tier: P1, subject_endpoint: chat_completions, route: openai, capability: service_tier, streaming: stream, assertions: [works], source: "litellm_core_utils/streaming_handler.py", fail_before_fix: proven, rationale: "Every relayed stream chunk carries the service_tier OpenAI stamped on it, so a streaming caller can see which tier served the request"} - {id: llm.chat_completions.openai.thinking.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: openai, capability: thinking, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "o-series reasoning; emerging"} - {id: llm.chat_completions.openai.structured_output.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: openai, capability: structured_output, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "response_schema extraction"} - {id: llm.chat_completions.anthropic.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: anthropic, capability: basic, streaming: nonstream, assertions: [works], source: "proxy_server.py:8455", rationale: "P0 route translated to Anthropic"} diff --git a/tests/e2e/coverage_registry/quota_management.yaml b/tests/e2e/coverage_registry/quota_management.yaml index 1051bf0bda9..163de67fc41 100644 --- a/tests/e2e/coverage_registry/quota_management.yaml +++ b/tests/e2e/coverage_registry/quota_management.yaml @@ -61,6 +61,9 @@ - {id: quota_management.spend_tracking.stream_cache_read.bills_cache_read_rate, module: quota_management, tier: P1, behavior: spend_tracking, variant: stream_cache_read, assertions: [bills_cache_read_rate], exercised_on: [chat_completions], source: "litellm_core_utils/streaming_chunk_builder_utils.py", rationale: "A streamed call's reassembled usage keeps the cached-token detail so cache reads bill at the cache-read discount, not full input price (#34812)"} - {id: quota_management.spend_tracking.messages_bridge.keeps_cache_tokens, module: quota_management, tier: P1, behavior: spend_tracking, variant: messages_bridge, assertions: [keeps_cache_tokens], exercised_on: [messages], source: "llms/anthropic/pass_through/responses_adapters/handler.py", rationale: "A /v1/messages request served by a Responses-only OpenAI model keeps its cache-read tokens and their discounted billing across the bridge (#34957)"} - {id: quota_management.spend_tracking.service_tier.bills_tier_rates, module: quota_management, tier: P1, behavior: spend_tracking, variant: service_tier, assertions: [bills_tier_rates], exercised_on: [chat_completions], source: "cost_calculator.py", rationale: "A priority service_tier call bills input, output, and reasoning at the deployment's *_priority rates and records the tier on the row (#35923, #35925)"} +- {id: quota_management.spend_tracking.service_tier_stream.records_served_tier, module: quota_management, tier: P1, behavior: spend_tracking, variant: service_tier_stream, assertions: [records_served_tier], exercised_on: [chat_completions], source: "litellm_core_utils/streaming_chunk_builder_utils.py", fail_before_fix: proven, rationale: "A streamed call with no service_tier requested bills at the rates of the tier OpenAI stamps on its chunks and records that served tier on the row; the reassembled stream dropped the provider tier so the row recorded none and priced at the default rates"} +- {id: quota_management.spend_tracking.service_tier_stream.responses_records_served_tier, module: quota_management, tier: P1, behavior: spend_tracking, variant: service_tier_stream, assertions: [records_served_tier], exercised_on: [responses], source: "responses/streaming_iterator.py", rationale: "A streamed /v1/responses call bills at the tier carried on the response.completed event's inner response and records that served tier on the spend row"} +- {id: quota_management.spend_tracking.service_tier_stream.messages_records_served_tier, module: quota_management, tier: P1, behavior: spend_tracking, variant: service_tier_stream, assertions: [records_served_tier], exercised_on: [messages], source: "llms/anthropic/pass_through/adapters/streaming_iterator.py", rationale: "A streamed /v1/messages call on an OpenAI-backed deployment bills at the tier OpenAI served; the Anthropic wire format has no tier field, so the spend row is the only record of it"} - {id: quota_management.spend_tracking.cost_headers.additive_components, module: quota_management, tier: P1, behavior: spend_tracking, variant: cost_headers, assertions: [additive_components], exercised_on: [chat_completions], source: "proxy/common_request_processing.py", rationale: "The x-litellm-response-cost-* component headers sum to the total, input covers only fresh tokens, and reasoning stays a subset of output (#36965)"} - {id: quota_management.spend_tracking.passthrough_stream.injects_usage_cost, module: quota_management, tier: P1, behavior: spend_tracking, variant: passthrough_stream, assertions: [injects_usage_cost], exercised_on: [openai_passthrough], source: "proxy/pass_through_endpoints/streaming_handler.py", rationale: "With include_cost_in_streaming_usage on, the /openai passthrough's final streaming usage frame carries the proxy-computed cost (#36503). Uncovered: the flag is only settable in litellm_settings, and the shared e2e stack does not turn it on yet"} - {id: quota_management.spend_tracking.websearch_interception.bills_under_request_session, module: quota_management, tier: P1, behavior: spend_tracking, variant: websearch_interception, assertions: [bills_under_request_session], exercised_on: [messages], source: "integrations/websearch_interception/handler.py", fail_before_fix: proven, rationale: "A web_search server tool the proxy intercepts into litellm.asearch writes its own asearch spend row, and that row carries the parent request's session_id so the session view counts the search and its cost next to the turn that triggered it (LIT-8063)"} diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 6e66529ec8f..0383da48c43 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -568,6 +568,17 @@ class AnthropicMessagesBody(BaseModel): cache: dict[str, bool] | None = {"no-cache": True} +class ResponsesStreamBody(BaseModel): + """POST /v1/responses body in the subset the spend tests stream with. + `input` stays a plain string: the tests only drive single-turn prompts.""" + + model: str + input: str + stream: bool = True + max_output_tokens: int | None = None + cache: dict[str, bool] | None = {"no-cache": True} + + class CountTokensBody(BaseModel): """POST /v1/messages/count_tokens body: the /v1/messages shape minus max_tokens (the endpoint only counts the prompt).""" diff --git a/tests/e2e/proxy_client.py b/tests/e2e/proxy_client.py index 4ea83e4b0d3..bd87828db2e 100644 --- a/tests/e2e/proxy_client.py +++ b/tests/e2e/proxy_client.py @@ -84,6 +84,7 @@ from models import ( OcrResponse, RerankBody, RerankResponse, + ResponsesStreamBody, RouterCurrentValues, RouterSettingsResponse, SearchToolCreateBody, @@ -969,6 +970,9 @@ class ProxyClient: def messages_stream(self, key: str, body: AnthropicMessagesBody) -> StreamingResponse: return self.transport.stream("/v1/messages", headers=self.transport.bearer(key), json=body) + def responses_stream(self, key: str, body: ResponsesStreamBody) -> StreamingResponse: + return self.transport.stream("/v1/responses", headers=self.transport.bearer(key), json=body) + def embed(self, key: str, body: EmbedBody) -> Result[EmbedResponse]: return self.transport.post( "/embeddings", diff --git a/tests/e2e/quota_management/spend_tracking/test_service_tier_pricing_e2e.py b/tests/e2e/quota_management/spend_tracking/test_service_tier_pricing_e2e.py index 770c5699b4e..bf68fb68a60 100644 --- a/tests/e2e/quota_management/spend_tracking/test_service_tier_pricing_e2e.py +++ b/tests/e2e/quota_management/spend_tracking/test_service_tier_pricing_e2e.py @@ -13,27 +13,45 @@ priority processing and served the default tier, the test fails there instead of producing a vacuous rate comparison. Reasoning is requested explicitly with `reasoning_effort`, so the reasoning-rate assertion rests on a parameter the test sets rather than on whatever the model happens to do by default. + +The streaming cases pin the served-tier contract: OpenAI stamps the tier it actually +used on every stream chunk, and that echo is what the caller sees and what the bill +must be computed on. The request sets no service_tier, so the only place the tier +can come from is the provider's response. The spend row must record the served tier +and price input at that tier's rate, and every chunk the proxy relays must carry the +same service_tier the provider sent. """ -import pytest +import json +import pytest from cost_rows import ( approx_equal, assert_fresh_tokens_billed_at, assert_total_is_sum_of_components, poll_cost_row, + poll_cost_row_where, register_priced_model, ) -from e2e_config import unique_marker +from e2e_config import CHEAP_OPENAI_MODEL, unique_marker from e2e_http import unwrap from lifecycle import ResourceManager -from models import ChatBody, ChatMessage, LiteLLMParamsBody +from models import ( + AnthropicMessagesBody, + ChatBody, + ChatMessage, + ChatStreamOptions, + LiteLLMParamsBody, + ResponsesStreamBody, +) +from pydantic import BaseModel from spend_e2e_client import SpendClient pytestmark = pytest.mark.e2e BACKEND = "openai/gpt-5.6-luna" OPENAI_API_KEY = "os.environ/OPENAI_API_KEY" +STREAM_BACKEND = f"openai/{CHEAP_OPENAI_MODEL}" INPUT_RATE = 4e-05 OUTPUT_RATE = 8e-05 @@ -42,6 +60,40 @@ PRIORITY_OUTPUT_RATE = 1.6e-04 REASONING_EFFORT = "high" +TIER_INPUT_RATES = {"default": INPUT_RATE, "priority": PRIORITY_INPUT_RATE} + + +class _StreamChunk(BaseModel): + id: str | None = None + service_tier: str | None = None + + +class _CompletedResponseObject(BaseModel): + id: str | None = None + service_tier: str | None = None + + +class _ResponsesStreamEvent(BaseModel): + type: str | None = None + response: _CompletedResponseObject | None = None + + +class _MessagesStreamEvent(BaseModel): + type: str | None = None + + +def _stream_chunks(events: list[str]) -> list[_StreamChunk]: + return [_StreamChunk.model_validate_json(event) for event in events if event.strip() != "[DONE]"] + + +def _served_tier(chunks: list[_StreamChunk]) -> str: + tiers = {chunk.service_tier for chunk in chunks if chunk.service_tier} + assert len(tiers) == 1, ( + f"the relayed stream carried {tiers or 'no'} service tier(s) across {len(chunks)} chunks; OpenAI stamps " + "the served tier on every chat chunk, so exactly one tier must reach the caller" + ) + return tiers.pop() + class TestServiceTierPricing: @pytest.mark.covers("quota_management.spend_tracking.service_tier.bills_tier_rates") @@ -83,8 +135,7 @@ class TestServiceTierPricing: ) ) assert chat.service_tier == "priority", ( - f"OpenAI served tier {chat.service_tier!r} instead of priority; " - "tier billing was never exercised" + f"OpenAI served tier {chat.service_tier!r} instead of priority; tier billing was never exercised" ) assert chat.id, f"chat response carried no id: {chat}" @@ -119,3 +170,150 @@ class TestServiceTierPricing: ) assert_total_is_sum_of_components(row) + + @pytest.mark.covers("quota_management.spend_tracking.service_tier_stream.records_served_tier") + def test_streamed_call_records_and_bills_the_served_tier( + self, client: SpendClient, resources: ResourceManager, scoped_key: str + ) -> None: + model = register_priced_model( + client.proxy, + resources, + "tier-priced-stream", + LiteLLMParamsBody( + model=BACKEND, + api_key=OPENAI_API_KEY, + input_cost_per_token=INPUT_RATE, + output_cost_per_token=OUTPUT_RATE, + input_cost_per_token_priority=PRIORITY_INPUT_RATE, + output_cost_per_token_priority=PRIORITY_OUTPUT_RATE, + ), + ) + + result = client.proxy.chat_stream( + scoped_key, + ChatBody( + model=model, + messages=[ChatMessage(role="user", content=f"{unique_marker()} reply with one word")], + max_completion_tokens=64, + stream=True, + ), + ) + assert result.ok and result.stream_events, ( + f"streamed chat failed (status {result.status_code}): {result.body[:300]}" + ) + chunks = _stream_chunks(result.stream_events) + served_tier = _served_tier(chunks) + assert served_tier in TIER_INPUT_RATES, f"no custom rate registered for served tier {served_tier!r}" + stream_id = chunks[0].id + assert stream_id, f"first stream chunk carried no id: {result.stream_events[0][:200]}" + + row = poll_cost_row(client.proxy, stream_id) + assert row is not None, f"no spend row with a cost breakdown landed for {stream_id}" + assert row.breakdown.service_tier == served_tier, ( + f"the provider served tier {served_tier!r} on every chunk but the bill records " + f"pricing basis {row.breakdown.service_tier!r}" + ) + assert_fresh_tokens_billed_at(row, TIER_INPUT_RATES[served_tier]) + assert_total_is_sum_of_components(row) + + @pytest.mark.covers("llm.chat_completions.openai.service_tier.stream.echoes_served_tier") + def test_every_streamed_chunk_carries_the_served_tier( + self, client: SpendClient, resources: ResourceManager, scoped_key: str + ) -> None: + model = register_priced_model( + client.proxy, resources, "tier-echo-stream", LiteLLMParamsBody(model=STREAM_BACKEND, api_key=OPENAI_API_KEY) + ) + result = client.proxy.chat_stream( + scoped_key, + ChatBody( + model=model, + messages=[ChatMessage(role="user", content=f"{unique_marker()} reply with one word")], + max_completion_tokens=64, + stream=True, + stream_options=ChatStreamOptions(include_usage=True), + ), + ) + assert result.ok and result.stream_events, ( + f"streamed chat failed (status {result.status_code}): {result.body[:300]}" + ) + chunks = _stream_chunks(result.stream_events) + served_tier = _served_tier(chunks) + missing = [ + json.loads(event) for event, chunk in zip(result.stream_events, chunks) if chunk.service_tier is None + ] + assert not missing, ( + f"{len(missing)} of {len(chunks)} relayed chunks dropped the provider's service_tier " + f"{served_tier!r}: {missing}" + ) + + @pytest.mark.covers("quota_management.spend_tracking.service_tier_stream.responses_records_served_tier") + def test_responses_stream_records_the_served_tier( + self, client: SpendClient, resources: ResourceManager, scoped_key: str + ) -> None: + model = register_priced_model( + client.proxy, + resources, + "tier-responses-stream", + LiteLLMParamsBody(model=STREAM_BACKEND, api_key=OPENAI_API_KEY), + ) + + result = client.proxy.responses_stream( + scoped_key, + ResponsesStreamBody(model=model, input=f"{unique_marker()} reply with one word"), + ) + assert result.ok and result.stream_events, ( + f"streamed responses call failed (status {result.status_code}): {result.body[:300]}" + ) + + events = [_ResponsesStreamEvent.model_validate_json(event) for event in result.stream_events] + completed = next((event for event in reversed(events) if event.type == "response.completed"), None) + assert completed is not None and completed.response is not None, ( + f"no response.completed event in the stream: {[e.type for e in events]}" + ) + served_tier = completed.response.service_tier + assert served_tier, f"response.completed carried no service_tier: {completed.response}" + assert served_tier in TIER_INPUT_RATES, f"no custom rate registered for served tier {served_tier!r}" + + row = poll_cost_row_where(client.proxy, scoped_key, lambda r: r.spend is not None and r.spend > 0) + assert row is not None, f"no spend row with a cost breakdown landed for the streamed responses call on {model}" + assert row.breakdown.service_tier == served_tier, ( + f"response.completed served tier {served_tier!r} but the bill records " + f"pricing basis {row.breakdown.service_tier!r}" + ) + + @pytest.mark.covers("quota_management.spend_tracking.service_tier_stream.messages_records_served_tier") + def test_messages_stream_records_the_served_tier( + self, client: SpendClient, resources: ResourceManager, scoped_key: str + ) -> None: + model = register_priced_model( + client.proxy, + resources, + "tier-messages-stream", + LiteLLMParamsBody(model=STREAM_BACKEND, api_key=OPENAI_API_KEY), + ) + + result = client.proxy.messages_stream( + scoped_key, + AnthropicMessagesBody( + model=model, + messages=[ChatMessage(role="user", content=f"{unique_marker()} reply with one word")], + max_tokens=64, + stream=True, + ), + ) + assert result.ok and result.stream_events, ( + f"streamed messages call failed (status {result.status_code}): {result.body[:300]}" + ) + + events = [_MessagesStreamEvent.model_validate_json(event) for event in result.stream_events] + assert any(event.type == "message_delta" for event in events), ( + f"the anthropic stream emitted no message_delta: {[e.type for e in events]}" + ) + + row = poll_cost_row_where(client.proxy, scoped_key, lambda r: r.spend is not None and r.spend > 0) + assert row is not None, f"no spend row with a cost breakdown landed for the streamed messages call on {model}" + served_tier = row.breakdown.service_tier + assert served_tier in TIER_INPUT_RATES and served_tier is not None, ( + "the anthropic wire format carries no service_tier, so the bill is the only record of " + f"the tier OpenAI served; the row recorded pricing basis {served_tier!r}" + ) diff --git a/tests/integration/spend/test_service_tier_stream_billing.py b/tests/integration/spend/test_service_tier_stream_billing.py new file mode 100644 index 00000000000..4791f941d61 --- /dev/null +++ b/tests/integration/spend/test_service_tier_stream_billing.py @@ -0,0 +1,660 @@ +"""Served service_tier drives billing on streamed calls, complete and disconnected. + +The scripted upstream answers OpenAI-compatible /chat/completions with SSE chunks +that carry service_tier "priority" and terminal usage. The deployment registers +distinct default and *_priority rates, so a bill computed on the wrong tier cannot +match the hand-computed expectation. /v1/messages deployments on hosted_vllm have +no anthropic-messages provider config, so they take the chat adapter: the +streamed response is an AnthropicStreamWrapper under AnthropicSSEStream, wrapped +by AnthropicMessagesStreamCacheWriter when litellm.cache is on and then by the +router's FallbackAwareAnthropicMessagesStream; each layer must delegate the +inner stream's chunks for disconnect billing to find them. + +Azure streams run the same OpenAI chunk path against /openai/deployments, so the +served tier must reach the spend row there too (LIT-2850). Databricks streams go +through DatabricksChatResponseIterator.chunk_parser and the databricks branch of +cost_per_token (LIT-8121). The responses bridge relays Responses API SSE as chat +chunks, so the served tier remembered from response.created must land on both +the chunks and the row. Gemini reports capacity as usageMetadata.trafficType, which maps to +service_tier "flex" and the *_flex rates (LIT-6287, LIT-6292). +""" + +import json +from collections.abc import Callable +from hashlib import sha256 +from typing import Final +from uuid import uuid4 + +import pytest +from integration._support.client import Gateway, Scenario, eventually, object_value +from integration._support.database import read_rows +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue + +PROMPT_TOKENS: Final = 30 +COMPLETION_TOKENS: Final = 40 +INPUT_RATE: Final = 0.001 +OUTPUT_RATE: Final = 0.002 +PRIORITY_INPUT_RATE: Final = 0.01 +PRIORITY_OUTPUT_RATE: Final = 0.02 +EXPECTED_FULL_SPEND: Final = PROMPT_TOKENS * PRIORITY_INPUT_RATE + COMPLETION_TOKENS * PRIORITY_OUTPUT_RATE +FLEX_INPUT_RATE: Final = 0.0005 +FLEX_OUTPUT_RATE: Final = 0.001 +EXPECTED_FLEX_SPEND: Final = PROMPT_TOKENS * FLEX_INPUT_RATE + COMPLETION_TOKENS * FLEX_OUTPUT_RATE + + +def _sse_frame(payload: dict[str, JsonValue]) -> bytes: + return f"data: {json.dumps(payload, separators=(',', ':'))}\n\n".encode() + + +def _chat_chunk(request_id: str, upstream_model: str, content: str, served_tier: str) -> dict[str, JsonValue]: + return { + "id": request_id, + "object": "chat.completion.chunk", + "created": 1, + "model": upstream_model, + "service_tier": served_tier, + "choices": [{"index": 0, "delta": {"role": "assistant", "content": content}, "finish_reason": None}], + } + + +def _respond_for( + request_id: str, + prompt: str, + *, + expected_target: str = "/v1/chat/completions", + pause: float = 0.4, + served_tier: str = "priority", + expected_requested_tier: str | None = None, +) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + if request.target == "/v1/models": + return Reply( + body=json.dumps({"object": "list", "data": [{"id": "gpt-4o-mini", "object": "model"}]}).encode() + ) + assert request.target.startswith(expected_target), request.target + body: Final = json.loads(request.body) + assert body["messages"] == [{"role": "user", "content": prompt}], body + if expected_requested_tier is not None: + assert body.get("service_tier") == expected_requested_tier, body + upstream_model: Final = str(body["model"]) + terminal: Final[dict[str, JsonValue]] = { + "id": request_id, + "object": "chat.completion.chunk", + "created": 1, + "model": upstream_model, + "service_tier": served_tier, + "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}], + "usage": { + "prompt_tokens": PROMPT_TOKENS, + "completion_tokens": COMPLETION_TOKENS, + "total_tokens": PROMPT_TOKENS + COMPLETION_TOKENS, + }, + } + return Reply( + content_type="text/event-stream", + chunks=( + _sse_frame(_chat_chunk(request_id, upstream_model, "first", served_tier)), + _sse_frame(_chat_chunk(request_id, upstream_model, "second", served_tier)), + _sse_frame(_chat_chunk(request_id, upstream_model, "third", served_tier)), + _sse_frame(terminal), + b"data: [DONE]\n\n", + ), + pause_between_chunks=pause, + ) + + return respond + + +def _tiered_model( + scenario: Scenario, + wire: Wire, + *, + litellm_model: str, + api_base: str | None = None, + **extra: JsonValue, +) -> str: + return scenario.model( + model=litellm_model, + api_base=api_base or f"{wire.url}/v1", + input_cost_per_token=INPUT_RATE, + output_cost_per_token=OUTPUT_RATE, + input_cost_per_token_priority=PRIORITY_INPUT_RATE, + output_cost_per_token_priority=PRIORITY_OUTPUT_RATE, + input_cost_per_token_flex=FLEX_INPUT_RATE, + output_cost_per_token_flex=FLEX_OUTPUT_RATE, + **extra, + ) + + +def _events(lines: list[str]) -> list[dict[str, JsonValue]]: + return [ + object_value(json.loads(line.removeprefix("data:"))) + for line in lines + if line.startswith("data:") and line.removeprefix("data:").strip() != "[DONE]" + ] + + +def _rows_for_key(key: str) -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT request_id, status, prompt_tokens, completion_tokens, spend, metadata FROM "LiteLLM_SpendLogs" ' + "WHERE api_key=%s", + (sha256(key.encode()).hexdigest(),), + ) + + +def _single_spend_row(key: str) -> dict[str, JsonValue]: + rows: Final = eventually(lambda: _rows_for_key(key), lambda values: len(values) == 1, seconds=70) + return rows[0] + + +def _cost_breakdown(row: dict[str, JsonValue]) -> dict[str, JsonValue]: + metadata: Final = row["metadata"] + parsed: Final = json.loads(metadata) if isinstance(metadata, str) else object_value(metadata) + return object_value(parsed["cost_breakdown"]) + + +@pytest.mark.timeout(120) +def test_completed_chat_stream_bills_the_served_tier(gateway: Gateway) -> None: + prompt: Final = f"tier control {uuid4().hex[:8]}" + request_id: Final = f"chatcmpl-{uuid4().hex[:8]}" + with ( + wire_server(_respond_for(request_id, prompt)) as wire, + gateway.scenario() as scenario, + ): + model: Final = _tiered_model(scenario, wire, litellm_model="openai/gpt-4o-mini") + key: Final = scenario.key(models=[model]) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}], "stream": True}, + key=key, + ) + assert response.status_code == 200, response.text + chunks: Final = _events(list(response.iter_lines())) + + assert len(chunks) == 4, chunks + tiers: Final = {chunk.get("service_tier") for chunk in chunks} + assert tiers == {"priority"}, f"every relayed chunk must carry the served tier: {tiers}" + + row: Final = _single_spend_row(key) + assert row["status"] == "success", row + assert row["request_id"] == request_id, row + assert row["prompt_tokens"] == PROMPT_TOKENS, row + assert row["completion_tokens"] == COMPLETION_TOKENS, row + assert float(str(row["spend"])) == pytest.approx(EXPECTED_FULL_SPEND), row + breakdown: Final = _cost_breakdown(row) + assert breakdown["service_tier"] == "priority", breakdown + assert len(wire.drain()) == 1 + + +@pytest.mark.timeout(120) +def test_disconnected_chat_stream_bills_partial_usage_at_the_served_tier(gateway: Gateway) -> None: + prompt: Final = f"tier control {uuid4().hex[:8]}" + request_id: Final = f"chatcmpl-{uuid4().hex[:8]}" + with ( + wire_server(_respond_for(request_id, prompt, pause=2.0)) as wire, + gateway.scenario() as scenario, + ): + model: Final = _tiered_model(scenario, wire, litellm_model="openai/gpt-4o-mini") + key: Final = scenario.key(models=[model]) + with gateway.client.stream( + "POST", + "/v1/chat/completions", + json={ + "model": model, + "messages": [{"role": "user", "content": prompt}], + "stream": True, + }, + headers={"Authorization": f"Bearer {key}"}, + ) as response: + assert response.status_code == 200, response.read().decode() + first_event: Final = next(line for line in response.iter_lines() if line.startswith("data:")) + assert object_value(json.loads(first_event.removeprefix("data:")))["id"] == request_id + + row: Final = _single_spend_row(key) + assert row["status"] == "success", row + assert int(row["prompt_tokens"]) > 0, row + assert int(row["completion_tokens"]) == 1, row + assert float(str(row["spend"])) == pytest.approx( + int(row["prompt_tokens"]) * PRIORITY_INPUT_RATE + PRIORITY_OUTPUT_RATE + ), row + breakdown: Final = _cost_breakdown(row) + assert breakdown["service_tier"] == "priority", breakdown + assert len(wire.drain()) == 1 + + +@pytest.mark.timeout(120) +def test_completed_messages_stream_bills_the_served_tier(gateway: Gateway) -> None: + prompt: Final = f"tier control {uuid4().hex[:8]}" + with ( + wire_server(_respond_for(f"chatcmpl-{uuid4().hex[:8]}", prompt)) as wire, + gateway.scenario() as scenario, + ): + model: Final = _tiered_model(scenario, wire, litellm_model="hosted_vllm/gpt-4o-mini") + key: Final = scenario.key(models=[model]) + with gateway.client.stream( + "POST", + "/v1/messages", + json={ + "model": model, + "messages": [{"role": "user", "content": prompt}], + "max_tokens": COMPLETION_TOKENS, + "stream": True, + }, + headers={"Authorization": f"Bearer {key}"}, + ) as response: + assert response.status_code == 200, response.read().decode() + events: Final = _events(list(response.iter_lines())) + + assert events[0]["type"] == "message_start", events + assert any(event["type"] == "message_delta" for event in events), events + + row: Final = _single_spend_row(key) + assert row["status"] == "success", row + assert row["prompt_tokens"] == PROMPT_TOKENS, row + assert row["completion_tokens"] == COMPLETION_TOKENS, row + assert float(str(row["spend"])) == pytest.approx(EXPECTED_FULL_SPEND), row + breakdown: Final = _cost_breakdown(row) + assert breakdown["service_tier"] == "priority", breakdown + assert len(wire.drain()) == 1 + + +@pytest.mark.timeout(120) +def test_disconnected_messages_stream_bills_partial_usage_at_the_served_tier(gateway: Gateway) -> None: + prompt: Final = f"tier control {uuid4().hex[:8]}" + with ( + wire_server(_respond_for(f"chatcmpl-{uuid4().hex[:8]}", prompt, pause=2.0)) as wire, + gateway.scenario() as scenario, + ): + model: Final = _tiered_model(scenario, wire, litellm_model="hosted_vllm/gpt-4o-mini") + key: Final = scenario.key(models=[model]) + with gateway.client.stream( + "POST", + "/v1/messages", + json={ + "model": model, + "messages": [{"role": "user", "content": prompt}], + "max_tokens": COMPLETION_TOKENS, + "stream": True, + }, + headers={"Authorization": f"Bearer {key}"}, + ) as response: + assert response.status_code == 200, response.read().decode() + first_event: Final = next(line for line in response.iter_lines() if line.startswith("data:")) + assert object_value(json.loads(first_event.removeprefix("data:")))["type"] == "message_start", first_event + + row: Final = _single_spend_row(key) + assert row["status"] == "success", row + assert float(str(row["spend"])) > 0, row + assert int(row["completion_tokens"]) < COMPLETION_TOKENS, row + breakdown: Final = _cost_breakdown(row) + assert breakdown["service_tier"] == "priority", breakdown + assert len(wire.drain()) == 1 + + +def _responses_frame(event: str, payload: dict[str, JsonValue]) -> bytes: + return f"event: {event}\ndata: {json.dumps(payload, separators=(',', ':'))}\n\n".encode() + + +def _respond_responses_for(response_id: str, prompt: str) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + assert request.target == "/v1/responses", request.target + body: Final = json.loads(request.body) + assert prompt in json.dumps(body["input"]), body["input"] + assert body["stream"] is True, body + upstream_model: Final = str(body["model"]) + text: Final = "firstsecondthird" + response_payload: Final[dict[str, JsonValue]] = { + "id": response_id, + "object": "response", + "model": upstream_model, + "status": "in_progress", + "service_tier": "priority", + "output": [], + } + message_item: Final[dict[str, JsonValue]] = { + "type": "message", + "id": "msg_1", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": text, "annotations": []}], + } + return Reply( + content_type="text/event-stream", + chunks=( + _responses_frame("response.created", {"type": "response.created", "response": response_payload}), + _responses_frame( + "response.output_item.added", + { + "type": "response.output_item.added", + "output_index": 0, + "item": { + "type": "message", + "id": "msg_1", + "status": "in_progress", + "role": "assistant", + "content": [], + }, + }, + ), + _responses_frame( + "response.content_part.added", + { + "type": "response.content_part.added", + "item_id": "msg_1", + "output_index": 0, + "content_index": 0, + "part": {"type": "output_text", "text": ""}, + }, + ), + *( + _responses_frame( + "response.output_text.delta", + { + "type": "response.output_text.delta", + "item_id": "msg_1", + "output_index": 0, + "content_index": 0, + "delta": delta, + }, + ) + for delta in ("first", "second", "third") + ), + _responses_frame( + "response.output_text.done", + { + "type": "response.output_text.done", + "item_id": "msg_1", + "output_index": 0, + "content_index": 0, + "text": text, + }, + ), + _responses_frame( + "response.content_part.done", + { + "type": "response.content_part.done", + "item_id": "msg_1", + "output_index": 0, + "content_index": 0, + "part": {"type": "output_text", "text": text}, + }, + ), + _responses_frame( + "response.output_item.done", + {"type": "response.output_item.done", "output_index": 0, "item": message_item}, + ), + _responses_frame( + "response.completed", + { + "type": "response.completed", + "response": { + **response_payload, + "status": "completed", + "output": [message_item], + "usage": { + "input_tokens": PROMPT_TOKENS, + "output_tokens": COMPLETION_TOKENS, + "total_tokens": PROMPT_TOKENS + COMPLETION_TOKENS, + }, + }, + }, + ), + ), + ) + + return respond + + +def _gemini_chunk(text: str) -> dict[str, JsonValue]: + return {"candidates": [{"index": 0, "content": {"role": "model", "parts": [{"text": text}]}}]} + + +def _respond_gemini_for(prompt: str) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + assert request.target.startswith("/models/gemini-2.5-flash:streamGenerateContent"), request.target + body: Final = json.loads(request.body) + assert prompt in json.dumps(body["contents"]), body["contents"] + terminal: Final[dict[str, JsonValue]] = { + "candidates": [{"index": 0, "content": {"role": "model", "parts": [{"text": ""}]}, "finishReason": "STOP"}], + "usageMetadata": { + "promptTokenCount": PROMPT_TOKENS, + "candidatesTokenCount": COMPLETION_TOKENS, + "totalTokenCount": PROMPT_TOKENS + COMPLETION_TOKENS, + "trafficType": "ON_DEMAND_FLEX", + }, + } + return Reply( + content_type="text/event-stream", + chunks=( + _sse_frame(_gemini_chunk("first")), + _sse_frame(_gemini_chunk("second")), + _sse_frame(_gemini_chunk("third")), + _sse_frame(terminal), + ), + ) + + return respond + + +@pytest.mark.timeout(120) +def test_azure_chat_stream_bills_the_served_tier(gateway: Gateway) -> None: + prompt: Final = f"tier control {uuid4().hex[:8]}" + request_id: Final = f"chatcmpl-{uuid4().hex[:8]}" + with ( + wire_server( + _respond_for(request_id, prompt, expected_target="/openai/deployments/gpt-4o-mini/chat/completions") + ) as wire, + gateway.scenario() as scenario, + ): + model: Final = _tiered_model( + scenario, + wire, + litellm_model="azure/gpt-4o-mini", + api_base=wire.url, + api_version="2024-10-21", + ) + key: Final = scenario.key(models=[model]) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}], "stream": True}, + key=key, + ) + assert response.status_code == 200, response.text + chunks: Final = _events(list(response.iter_lines())) + + assert len(chunks) == 4, chunks + tiers: Final = {chunk.get("service_tier") for chunk in chunks} + assert tiers == {"priority"}, f"every relayed chunk must carry the served tier: {tiers}" + + row: Final = _single_spend_row(key) + assert row["status"] == "success", row + assert row["prompt_tokens"] == PROMPT_TOKENS, row + assert row["completion_tokens"] == COMPLETION_TOKENS, row + assert float(str(row["spend"])) == pytest.approx(EXPECTED_FULL_SPEND), row + breakdown: Final = _cost_breakdown(row) + assert breakdown["service_tier"] == "priority", breakdown + assert len(wire.drain()) == 1 + + +@pytest.mark.timeout(120) +def test_databricks_chat_stream_bills_the_served_tier(gateway: Gateway) -> None: + prompt: Final = f"tier control {uuid4().hex[:8]}" + request_id: Final = f"chatcmpl-{uuid4().hex[:8]}" + with ( + wire_server(_respond_for(request_id, prompt, expected_target="/serving-endpoints/chat/completions")) as wire, + gateway.scenario() as scenario, + ): + model: Final = _tiered_model( + scenario, + wire, + litellm_model="databricks/dbrx-instruct", + api_base=f"{wire.url}/serving-endpoints", + ) + key: Final = scenario.key(models=[model]) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}], "stream": True}, + key=key, + ) + assert response.status_code == 200, response.text + chunks: Final = _events(list(response.iter_lines())) + + assert len(chunks) == 4, chunks + tiers: Final = {chunk.get("service_tier") for chunk in chunks} + assert tiers == {"priority"}, f"every relayed chunk must carry the served tier: {tiers}" + + row: Final = _single_spend_row(key) + assert row["status"] == "success", row + assert row["prompt_tokens"] == PROMPT_TOKENS, row + assert row["completion_tokens"] == COMPLETION_TOKENS, row + assert float(str(row["spend"])) == pytest.approx(EXPECTED_FULL_SPEND), row + breakdown: Final = _cost_breakdown(row) + assert breakdown["service_tier"] == "priority", breakdown + assert len(wire.drain()) == 1 + + +@pytest.mark.timeout(120) +def test_responses_bridge_stream_bills_the_served_tier(gateway: Gateway) -> None: + prompt: Final = f"tier control {uuid4().hex[:8]}" + with ( + wire_server(_respond_responses_for(f"resp_{uuid4().hex[:8]}", prompt)) as wire, + gateway.scenario() as scenario, + ): + model: Final = _tiered_model(scenario, wire, litellm_model="openai/responses/gpt-4o-mini") + key: Final = scenario.key(models=[model]) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}], "stream": True}, + key=key, + ) + assert response.status_code == 200, response.text + chunks: Final = _events(list(response.iter_lines())) + + assert len(chunks) >= 4, chunks + tiers: Final = {chunk.get("service_tier") for chunk in chunks} + assert tiers == {"priority"}, f"every relayed chunk must carry the served tier: {tiers}" + + row: Final = _single_spend_row(key) + assert row["status"] == "success", row + assert row["prompt_tokens"] == PROMPT_TOKENS, row + assert row["completion_tokens"] == COMPLETION_TOKENS, row + assert float(str(row["spend"])) == pytest.approx(EXPECTED_FULL_SPEND), row + breakdown: Final = _cost_breakdown(row) + assert breakdown["service_tier"] == "priority", breakdown + assert len(wire.drain()) == 1 + + +@pytest.mark.timeout(120) +def test_gemini_chat_stream_bills_the_flex_tier(gateway: Gateway) -> None: + prompt: Final = f"tier control {uuid4().hex[:8]}" + with ( + wire_server(_respond_gemini_for(prompt)) as wire, + gateway.scenario() as scenario, + ): + model: Final = _tiered_model( + scenario, + wire, + litellm_model="gemini/gemini-2.5-flash", + api_base=wire.url, + ) + key: Final = scenario.key(models=[model]) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}], "stream": True}, + key=key, + ) + assert response.status_code == 200, response.text + chunks: Final = _events(list(response.iter_lines())) + assert len(chunks) >= 2, chunks + + row: Final = _single_spend_row(key) + assert row["status"] == "success", row + assert row["prompt_tokens"] == PROMPT_TOKENS, row + assert row["completion_tokens"] == COMPLETION_TOKENS, row + assert float(str(row["spend"])) == pytest.approx(EXPECTED_FLEX_SPEND), row + breakdown: Final = _cost_breakdown(row) + assert breakdown["service_tier"] == "flex", breakdown + assert len(wire.drain()) == 1 + + +@pytest.mark.timeout(120) +def test_requested_priority_downgraded_to_default_bills_base_rates(gateway: Gateway) -> None: + prompt: Final = f"tier control {uuid4().hex[:8]}" + request_id: Final = f"chatcmpl-{uuid4().hex[:8]}" + with ( + wire_server( + _respond_for(request_id, prompt, served_tier="default", expected_requested_tier="priority") + ) as wire, + gateway.scenario() as scenario, + ): + model: Final = _tiered_model(scenario, wire, litellm_model="openai/gpt-4o-mini") + key: Final = scenario.key(models=[model]) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": prompt}], + "stream": True, + "service_tier": "priority", + }, + key=key, + ) + assert response.status_code == 200, response.text + chunks: Final = _events(list(response.iter_lines())) + + assert len(chunks) == 4, chunks + tiers: Final = {chunk.get("service_tier") for chunk in chunks} + assert tiers == {"default"}, f"every relayed chunk must carry the served tier: {tiers}" + + row: Final = _single_spend_row(key) + assert row["status"] == "success", row + assert row["request_id"] == request_id, row + assert row["prompt_tokens"] == PROMPT_TOKENS, row + assert row["completion_tokens"] == COMPLETION_TOKENS, row + assert float(str(row["spend"])) == pytest.approx( + PROMPT_TOKENS * INPUT_RATE + COMPLETION_TOKENS * OUTPUT_RATE + ), row + breakdown: Final = _cost_breakdown(row) + assert breakdown.get("service_tier") != "priority", breakdown + assert len(wire.drain()) == 1 + + +@pytest.mark.timeout(120) +def test_requested_priority_with_auto_echo_bills_priority(gateway: Gateway) -> None: + prompt: Final = f"tier control {uuid4().hex[:8]}" + request_id: Final = f"chatcmpl-{uuid4().hex[:8]}" + with ( + wire_server(_respond_for(request_id, prompt, served_tier="auto")) as wire, + gateway.scenario() as scenario, + ): + model: Final = _tiered_model(scenario, wire, litellm_model="openai/gpt-4o-mini") + key: Final = scenario.key(models=[model]) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": prompt}], + "stream": True, + "service_tier": "priority", + }, + key=key, + ) + assert response.status_code == 200, response.text + chunks: Final = _events(list(response.iter_lines())) + assert len(chunks) == 4, chunks + + row: Final = _single_spend_row(key) + assert row["status"] == "success", row + assert row["request_id"] == request_id, row + assert float(str(row["spend"])) == pytest.approx(EXPECTED_FULL_SPEND), row + breakdown: Final = _cost_breakdown(row) + assert breakdown["service_tier"] == "priority", breakdown + assert len(wire.drain()) == 1 diff --git a/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py b/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py index 86dd356e5f5..92de00a4a3f 100644 --- a/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py +++ b/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py @@ -2058,3 +2058,13 @@ async def test_queue_request_stream_is_untouched_while_keepalives_are_unconfigur assert not any(chunk.startswith(b": ping") for chunk in chunks) assert chunks[-1] == b"data: [DONE]\n\n" + + +def test_fast_serialize_simple_model_response_stream_keeps_served_service_tier(): + chunk = _simple_chunk() + chunk.service_tier = "priority" + + result = _fast_serialize_simple_model_response_stream(chunk) + + assert result is not None + assert json.loads(result)["service_tier"] == "priority" diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index c17f41a8b8f..8485c286a30 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -64,6 +64,7 @@ from litellm.proxy._types import ProxyErrorTypes, ProxyException from litellm.proxy._types import UserAPIKeyAuth as ProxyUserAPIKeyAuth from litellm.proxy.utils import ProxyLogging from litellm.router import Router +from litellm.router_utils.add_retry_fallback_headers import prepare_response_for_header_attachment def test_attach_guardrail_information_copies_recorded_entries_onto_model_response(): @@ -7337,6 +7338,61 @@ class TestStreamingClientDisconnectBilling: assert standard_logging_object["total_tokens"] > 0 assert standard_logging_object["response_cost"] >= 0.002 + @pytest.mark.asyncio + async def test_disconnect_bills_partial_spend_for_anthropic_adapter_stream(self): + """ + The proxy's cleanup gets the FallbackAwareAnthropicMessagesStream the + router returns for /v1/messages; its chunks/messages must delegate + through the translate_completion_output_params_streaming result to the + inner chat stream's collected chunks or a disconnect bills nothing. + """ + from litellm.llms.anthropic.pass_through.adapters.streaming_iterator import ( + AnthropicSSEStream, + ) + from litellm.llms.anthropic.pass_through.adapters.transformation import ( + AnthropicAdapter, + ) + from litellm.router import FallbackAwareAnthropicMessagesStream + + async def _sse_frames() -> AsyncGenerator[bytes, None]: + yield b"event: message_start\n\n" + + recorder = _RecordingSuccessLogger() + original_callbacks = litellm.callbacks + litellm.callbacks = [recorder] + try: + response = await self._start_partial_stream() + setattr(response.chunks[-1], "service_tier", "priority") # noqa: B010 # pydantic extra, not a declared field + source_iterator: Final = AnthropicAdapter().translate_completion_output_params_streaming( + response, + model=response.model or "gpt-4o-mini", + is_async=True, + litellm_logging_obj=response.logging_obj, + ) + assert isinstance(source_iterator, AnthropicSSEStream) + streamed: Final = prepare_response_for_header_attachment( + FallbackAwareAnthropicMessagesStream(_sse_frames(), source_iterator) + ) + + billed: Final = await _bill_partial_streamed_spend_on_disconnect( + {"litellm_logging_obj": response.logging_obj}, + streamed, + ) + + for _ in range(50): + if recorder.success_events: + break + await asyncio.sleep(0.1) + await asyncio.sleep(0.5) + finally: + litellm.callbacks = original_callbacks + + assert billed is True + assert len(recorder.success_events) == 1 + partial_response: Final = recorder.success_events[0]["response_obj"] + assert getattr(partial_response, "service_tier") == "priority" + assert partial_response.usage.total_tokens > 0 + @pytest.mark.asyncio async def test_completed_stream_does_not_double_bill_on_late_disconnect(self): recorder = _RecordingSuccessLogger() diff --git a/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py b/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py index d0f9bad795d..282b84104a6 100644 --- a/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py +++ b/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py @@ -4352,3 +4352,39 @@ def test_map_optional_params_verbosity_merges_into_text(): verbosity_only_request, ) assert verbosity_only_request["text"] == {"verbosity": "low"} + + +def test_response_completed_carries_the_served_service_tier(): + from litellm.completion_extras.litellm_responses_transformation.transformation import ( + OpenAiResponsesToChatCompletionStreamIterator, + ) + + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) + + result = iterator.chunk_parser( + { + "type": "response.completed", + "response": {"id": "resp_1", "status": "completed", "output": [], "service_tier": "default"}, + } + ) + + assert result.model_dump()["service_tier"] == "default" + + +def test_every_bridged_chunk_after_response_created_carries_the_served_service_tier(): + from litellm.completion_extras.litellm_responses_transformation.transformation import ( + OpenAiResponsesToChatCompletionStreamIterator, + ) + + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) + events = [ + {"type": "response.created", "response": {"id": "resp_1", "status": "in_progress", "service_tier": "default"}}, + {"type": "response.output_item.added", "output_index": 0, "item": {"type": "message"}}, + {"type": "response.output_text.delta", "output_index": 0, "delta": "Hi"}, + {"type": "response.output_item.done", "output_index": 0, "item": {"type": "message"}}, + {"type": "response.completed", "response": {"id": "resp_1", "status": "completed", "output": []}}, + ] + + relayed = [iterator.chunk_parser(event).model_dump().get("service_tier") for event in events] + + assert relayed == ["default"] * len(events), relayed diff --git a/tests/unit/litellm_core_utils/test_litellm_logging.py b/tests/unit/litellm_core_utils/test_litellm_logging.py index ed614b93a77..60b7ed32399 100644 --- a/tests/unit/litellm_core_utils/test_litellm_logging.py +++ b/tests/unit/litellm_core_utils/test_litellm_logging.py @@ -8811,6 +8811,58 @@ async def test_async_failure_handler_delivers_failure_payload_to_custom_logger() assert events.empty() +def test_responses_completed_event_bills_the_served_service_tier(): + """The served service_tier on response.completed's inner ResponsesAPIResponse + must reach the cost calculator, so a priority-served stream prices at the + priority rates instead of the default tier's.""" + logging_obj: Final = LitellmLogging( + model="openai/gpt-5.1", + messages=[{"role": "user", "content": "hi"}], + stream=True, + call_type="aresponses", + start_time=time.time(), + litellm_call_id="resp-served-tier", + function_id="resp-served-tier", + ) + logging_obj.update_environment_variables( + model="openai/gpt-5.1", + user="", + optional_params={}, + litellm_params={}, + custom_llm_provider="openai", + ) + inner: Final = ResponsesAPIResponse( + id="resp-served-tier", + created_at=1, + object="response", + status="completed", + model="gpt-5.1", + output=[], + usage=ResponseAPIUsage(input_tokens=10, output_tokens=20, total_tokens=30), + service_tier="priority", + ) + event: Final = ResponseCompletedEvent(type="response.completed", response=inner) + + cost: Final = logging_obj._response_cost_calculator(result=event) # pyright: ignore[reportPrivateUsage] # parity with the suite's own direct calls + + billed_response: Final = ModelResponse( + model="gpt-5.1", + usage=litellm.Usage(prompt_tokens=10, completion_tokens=20, total_tokens=30), + ) + tier_cost: Final = litellm.completion_cost( + completion_response=billed_response, + model="openai/gpt-5.1", + service_tier="priority", + ) + default_cost: Final = litellm.completion_cost( + completion_response=billed_response, + model="openai/gpt-5.1", + ) + + assert cost == pytest.approx(tier_cost) + assert cost > default_cost + + def _image_logging_obj() -> LitellmLogging: logging_obj = LitellmLogging( model="gpt-image-2", diff --git a/tests/unit/litellm_core_utils/test_streaming_chunk_builder_utils.py b/tests/unit/litellm_core_utils/test_streaming_chunk_builder_utils.py index aaf877df364..83e53b2d80a 100644 --- a/tests/unit/litellm_core_utils/test_streaming_chunk_builder_utils.py +++ b/tests/unit/litellm_core_utils/test_streaming_chunk_builder_utils.py @@ -1820,3 +1820,34 @@ def test_calculate_usage_keeps_a_reported_count_over_a_later_chunks_zero() -> No ) assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == (5, 17, 22) + + +def _tier_chunk(content: str, service_tier: str | None, finish_reason: str | None = None) -> ModelResponseStream: + return ModelResponseStream( + id="chatcmpl-tier", + created=1, + model="gpt-4.1-mini", + object="chat.completion.chunk", + choices=[StreamingChoices(finish_reason=finish_reason, index=0, delta=Delta(content=content, role=None))], + **({"service_tier": service_tier} if service_tier is not None else {}), + ) + + +def test_stream_chunk_builder_records_the_last_service_tier_the_provider_stamped(): + chunks = [ + _tier_chunk("Hel", "auto"), + _tier_chunk("lo", None), + _tier_chunk("", "default", finish_reason="stop"), + ] + + response = stream_chunk_builder(chunks=chunks) + + assert response is not None + assert response.model_dump()["service_tier"] == "default" + + +def test_stream_chunk_builder_omits_service_tier_when_no_chunk_carried_one(): + response = stream_chunk_builder(chunks=[_tier_chunk("Hi", None, finish_reason="stop")]) + + assert response is not None + assert "service_tier" not in response.model_dump() diff --git a/tests/unit/litellm_core_utils/test_streaming_handler.py b/tests/unit/litellm_core_utils/test_streaming_handler.py index 62d8b0e203f..d07e8822eb0 100644 --- a/tests/unit/litellm_core_utils/test_streaming_handler.py +++ b/tests/unit/litellm_core_utils/test_streaming_handler.py @@ -4983,3 +4983,50 @@ async def test_async_stream_without_usage_counts_tokens_off_the_event_loop(): assert chunks[-1].usage.prompt_tokens > 100_000 assert chunks[-1].usage.completion_tokens > 100_000 assert_loop_stayed_free(took, lags) + + +@pytest.mark.parametrize("sync_mode", [True, False]) +@pytest.mark.asyncio +async def test_openai_stream_relays_the_served_service_tier_on_every_chunk_including_usage( + logging_obj: Logging, sync_mode: bool +): + from litellm.utils import ModelResponseListIterator + + def _chunk(content: str, finish_reason: str | None, usage: Usage | None, choices: bool = True): + return ModelResponseStream( + id="chatcmpl-tier", + created=1742056047, + model="gpt-4.1-mini", + choices=[StreamingChoices(finish_reason=finish_reason, index=0, delta=Delta(content=content))] + if choices + else [], + usage=usage, + service_tier="default", + ) + + logging_obj.update_environment_variables( + model="gpt-4.1-mini", + optional_params={"stream_options": {"include_usage": True}}, + litellm_params={}, + custom_llm_provider="openai", + ) + wrapper = CustomStreamWrapper( + completion_stream=ModelResponseListIterator( + model_responses=[ + _chunk("Hi", None, None), + _chunk("", "stop", None), + _chunk("", None, Usage(prompt_tokens=10, completion_tokens=1, total_tokens=11), choices=False), + ] + ), + model="gpt-4.1-mini", + custom_llm_provider="openai", + logging_obj=logging_obj, + stream_options={"include_usage": True}, + ) + + relayed = ( + [chunk.model_dump() for chunk in wrapper] if sync_mode else [chunk.model_dump() async for chunk in wrapper] + ) + + assert [chunk.get("service_tier") for chunk in relayed] == ["default"] * len(relayed), relayed + assert relayed[-1]["usage"]["total_tokens"] == 11 diff --git a/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_sse_stream.py b/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_sse_stream.py new file mode 100644 index 00000000000..fbbbc579d94 --- /dev/null +++ b/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_sse_stream.py @@ -0,0 +1,87 @@ +""" +Tests for AnthropicSSEStream, the object translate_completion_output_params_streaming +hands to the proxy for /v1/messages streaming. It must emit the same SSE bytes as +the wrapper's async_anthropic_sse_wrapper, propagate aclose into it, and expose the +wrapper's chunks/messages/model so disconnect-time partial billing can read them. +""" + +from typing import Final +from unittest.mock import MagicMock + +import pytest + +from litellm.llms.anthropic.pass_through.adapters.streaming_iterator import ( + AnthropicSSEStream, + AnthropicStreamWrapper, +) +from litellm.types.utils import Delta, StreamingChoices + + +def _make_chunk(delta: Delta, finish_reason: str | None = None) -> MagicMock: + chunk = MagicMock() + chunk.choices = [StreamingChoices(finish_reason=finish_reason, index=0, delta=delta, logprobs=None)] + chunk.usage = None + chunk._hidden_params = {} + return chunk + + +class _AsyncStream: + def __init__(self, items: list[MagicMock]): + self._it = iter(items) + self.chunks = list(items) + self.messages: list[dict] = [{"role": "user", "content": "hi"}] + + def __aiter__(self): + return self + + async def __anext__(self): + try: + return next(self._it) + except StopIteration: + raise StopAsyncIteration + + +def _streamed_events() -> AnthropicSSEStream: + upstream: Final = _AsyncStream( + [ + _make_chunk(Delta(content="Once")), + _make_chunk(Delta(content=" upon"), finish_reason="stop"), + ] + ) + wrapper: Final = AnthropicStreamWrapper(completion_stream=upstream, model="gpt-4o-mini") + wrapper._message_id = "msg_test" + return AnthropicSSEStream(wrapper) + + +@pytest.mark.asyncio +async def test_sse_stream_yields_identical_bytes_to_the_wrappers_sse_wrapper(): + upstream_a: Final = _AsyncStream( + [_make_chunk(Delta(content="Once")), _make_chunk(Delta(content=" upon"), finish_reason="stop")] + ) + wrapper_a: Final = AnthropicStreamWrapper(completion_stream=upstream_a, model="gpt-4o-mini") + wrapper_a._message_id = "msg_test" + expected: Final = [event async for event in wrapper_a.async_anthropic_sse_wrapper()] + + actual: Final = [event async for event in _streamed_events()] + + assert actual == expected + + +@pytest.mark.asyncio +async def test_sse_stream_aclose_ends_the_wrapped_stream(): + stream: Final = _streamed_events() + + first: Final = await stream.__anext__() + assert first.startswith(b"event: message_start") + await stream.aclose() + with pytest.raises(StopAsyncIteration): + await stream.__anext__() + + +def test_sse_stream_exposes_chunks_messages_and_model(): + stream: Final = _streamed_events() + + assert stream.model == "gpt-4o-mini" + assert stream.messages == [{"role": "user", "content": "hi"}] + chunks: Final = stream.chunks + assert isinstance(chunks, list) and len(chunks) == 2 diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_response_cache.py b/tests/unit/llms/anthropic/pass_through/messages/test_response_cache.py index e55e73ed43f..a8a8eba0bf7 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_response_cache.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_response_cache.py @@ -10,6 +10,7 @@ import litellm from litellm._internal_context import in_post_response_phase from litellm.caching.caching import Cache, LiteLLMCacheType from litellm.caching.caching_handler import LLMCachingHandler +from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER from litellm.llms.anthropic.pass_through.messages import handler from litellm.llms.anthropic.pass_through.messages.response_cache import ( AnthropicMessagesStreamCacheWriter, @@ -63,6 +64,12 @@ async def _collect(stream: AsyncIterator[bytes]) -> List[bytes]: return [chunk async for chunk in stream] +@pytest.fixture(autouse=True) +async def _drain_logging_worker(): + yield + await GLOBAL_LOGGING_WORKER.flush() + + @pytest.fixture def local_cache(): previous_cache = litellm.cache @@ -282,6 +289,40 @@ class _HeldBackStream: raise StopAsyncIteration +class _AttributedStream: + """Stream stub carrying the billing attributes the disconnect helper reads.""" + + def __init__(self, chunks: list) -> None: + self.chunks = [object()] + self.messages = [{"role": "user", "content": "hi"}] + self.model = "gpt-4o-mini" + self._pending = list(chunks) + + def __aiter__(self) -> "_AttributedStream": + return self + + async def __anext__(self) -> bytes: + if not self._pending: + raise StopAsyncIteration + return self._pending.pop(0) + + +@pytest.mark.asyncio +async def test_cache_writer_exposes_inner_stream_billing_attributes(request_kwargs): + caching_handler = LLMCachingHandler( + original_function=handler.anthropic_messages, + request_kwargs=dict(request_kwargs), + start_time=datetime.datetime.now(), + ) + inner = _AttributedStream(STREAM_EVENTS) + writer = AnthropicMessagesStreamCacheWriter(stream=inner, caching_handler=caching_handler) + + assert writer.chunks is inner.chunks + assert writer.messages is inner.messages + assert writer.model == "gpt-4o-mini" + assert await _collect(writer) == STREAM_EVENTS + + @pytest.mark.asyncio async def test_stream_cache_write_runs_in_post_response_phase(request_kwargs, monkeypatch): """Every event, message_stop included, is already with the client when the stream write diff --git a/tests/unit/llms/databricks/chat/test_databricks_chat_transformation.py b/tests/unit/llms/databricks/chat/test_databricks_chat_transformation.py index 9cd17bd3580..1cbc9eeb897 100644 --- a/tests/unit/llms/databricks/chat/test_databricks_chat_transformation.py +++ b/tests/unit/llms/databricks/chat/test_databricks_chat_transformation.py @@ -883,3 +883,13 @@ def test_completion_merges_system_messages_when_one_has_empty_content(respx_mock {"role": "system", "content": "You are terse."}, {"role": "user", "content": "Hello"}, ] + + +def test_chunk_parser_relays_the_served_service_tier(): + iterator = DatabricksChatResponseIterator(streaming_response=None, sync_stream=True) + + with_tier: Final = iterator.chunk_parser({**_streaming_chunk(), "service_tier": "priority"}) + assert with_tier.model_dump()["service_tier"] == "priority" + + without_tier: Final = iterator.chunk_parser(_streaming_chunk()) + assert getattr(without_tier, "service_tier", None) is None diff --git a/tests/unit/llms/databricks/test_databricks_cost_calculator.py b/tests/unit/llms/databricks/test_databricks_cost_calculator.py index 494b99c1d11..7120a130462 100644 --- a/tests/unit/llms/databricks/test_databricks_cost_calculator.py +++ b/tests/unit/llms/databricks/test_databricks_cost_calculator.py @@ -156,8 +156,6 @@ def test_uncached_request_bills_every_prompt_token_at_the_input_rate(local_model assert completion_cost == pytest.approx(200 * info["output_cost_per_token"]) - - @pytest.mark.parametrize("model", NEW_MODELS) def test_new_models_carry_cache_pricing(local_model_cost_map: None, model: str) -> None: info: Final = _model_info(model) @@ -232,3 +230,28 @@ def test_sonnet_5_ships_standard_rates_not_introductory(local_model_cost_map: No for field in PRICE_FIELDS: assert sonnet_5[field] == pytest.approx(sonnet_4_6[field]), field + + +def test_cost_per_token_bills_the_served_priority_tier( + local_model_cost_map: None, monkeypatch: pytest.MonkeyPatch +) -> None: + rates: Final = { + "input_cost_per_token": 0.001, + "output_cost_per_token": 0.002, + "input_cost_per_token_priority": 0.01, + "output_cost_per_token_priority": 0.02, + "litellm_provider": "databricks", + "mode": "chat", + } + monkeypatch.setitem(litellm.model_cost, "databricks/dbrx-tiered-test", rates) + usage: Final = Usage(prompt_tokens=30, completion_tokens=40, total_tokens=70) + + prompt_cost, completion_cost = cost_per_token( + model="databricks/dbrx-tiered-test", usage=usage, service_tier="priority" + ) + assert prompt_cost == pytest.approx(30 * 0.01) + assert completion_cost == pytest.approx(40 * 0.02) + + prompt_cost, completion_cost = cost_per_token(model="databricks/dbrx-tiered-test", usage=usage) + assert prompt_cost == pytest.approx(30 * 0.001) + assert completion_cost == pytest.approx(40 * 0.002) diff --git a/tests/unit/llms/openai/chat/test_openai_gpt_transformation.py b/tests/unit/llms/openai/chat/test_openai_gpt_transformation.py index 53c5b9d7cbc..85a04778e5c 100644 --- a/tests/unit/llms/openai/chat/test_openai_gpt_transformation.py +++ b/tests/unit/llms/openai/chat/test_openai_gpt_transformation.py @@ -248,6 +248,33 @@ class TestOpenAIChatCompletionStreamingHandler: assert result.usage.completion_tokens == 350 assert result.usage.total_tokens == 14147 + def test_chunk_parser_preserves_service_tier(self): + """OpenAI-compatible upstreams serve a service_tier on every streamed + chunk; chunk_parser must keep it on the emitted ModelResponseStream so + disconnect billing and the reassembled response see the served tier.""" + handler = OpenAIChatCompletionStreamingHandler( + streaming_response=None, sync_stream=True + ) + + tiered_chunk = { + "id": "gen-123", + "created": 1234567890, + "model": "openai/gpt-4o-mini", + "object": "chat.completion.chunk", + "choices": [ + { + "index": 0, + "delta": {"role": "assistant", "content": ""}, + "finish_reason": None, + } + ], + "service_tier": "priority", + } + plain_chunk = {key: value for key, value in tiered_chunk.items() if key != "service_tier"} + + assert handler.chunk_parser(tiered_chunk).model_dump().get("service_tier") == "priority" + assert handler.chunk_parser(plain_chunk).model_dump().get("service_tier") is None + def test_chunk_parser_raises_on_in_body_error_payload(self): """vLLM/sglang return HTTP 200 streams whose body carries the error, e.g. data: {"error": {..., "code": 400}}. chunk_parser must surface it diff --git a/tests/unit/test_cost_calculator.py b/tests/unit/test_cost_calculator.py index c1939a23e74..2fe09e6de98 100644 --- a/tests/unit/test_cost_calculator.py +++ b/tests/unit/test_cost_calculator.py @@ -1948,7 +1948,7 @@ def test_completion_cost_extracts_service_tier_from_usage(_local_model_cost_map) def test_completion_cost_service_tier_priority(_local_model_cost_map): - """Test that service_tier extraction follows priority: optional_params > completion_response > usage.""" + """Test that the served tier wins over the requested tier: response > usage > request.""" from litellm import completion_cost # Test with gpt-5-nano which has flex pricing @@ -1965,7 +1965,7 @@ def test_completion_cost_service_tier_priority(_local_model_cost_map): ) setattr(response, "service_tier", "priority") - # Test that optional_params takes priority over response and usage + # A request-level tier loses to the tier the response actually served cost_from_params = completion_cost( completion_response=response, model=model, @@ -1973,20 +1973,18 @@ def test_completion_cost_service_tier_priority(_local_model_cost_map): optional_params={"service_tier": "flex"}, ) - # Test that response takes priority over usage when optional_params is not provided - completion_cost( + # Response takes priority over usage + cost_served_priority = completion_cost( completion_response=response, model=model, custom_llm_provider="openai", ) - # Test that usage is used when neither optional_params nor response have service_tier - # Create a new response without service_tier attribute + # Create a new response without service_tier attribute so it falls back to usage response_no_tier = ModelResponse( usage=usage, model=model, ) - # Don't set service_tier on response, so it will fall back to usage cost_from_usage = completion_cost( completion_response=response_no_tier, @@ -1994,12 +1992,13 @@ def test_completion_cost_service_tier_priority(_local_model_cost_map): custom_llm_provider="openai", ) - # All should use flex pricing (from different sources) assert cost_from_params > 0, "Cost from params should be greater than 0" assert cost_from_usage > 0, "Cost from usage should be greater than 0" - # Costs should be similar (all using flex) - assert abs(cost_from_params - cost_from_usage) < 1e-6, "Costs from params and usage should be similar (both flex)" + # Requested flex is ignored once the response reports served priority + assert cost_from_params == pytest.approx(cost_served_priority), ( + "request-level service_tier must defer to the served tier on the response" + ) def test_completion_cost_service_tier_for_bedrock(_local_model_cost_map): @@ -5468,3 +5467,100 @@ def test_completion_cost_is_zero_when_explicit_rates_are_zero(monkeypatch: pytes ) assert cost == 0.0 + + +@pytest.mark.parametrize( + ("requested", "served", "expected"), + [ + (None, "priority", "priority"), + ("priority", "flex", "flex"), + ("priority", "default", None), + ("priority", "standard", None), + ("priority", "auto", "priority"), + ("priority", "scale", "priority"), + ("priority", None, "priority"), + ("auto", None, None), + (None, "Priority", "priority"), + ("flex", "on_demand", "flex"), + ], +) +def test_resolve_billable_service_tier(requested: object, served: object, expected: str | None) -> None: + from litellm.cost_calculator import _resolve_billable_service_tier + + assert _resolve_billable_service_tier(requested=requested, served=served) == expected + + +def _served_tier_cost_model(monkeypatch: pytest.MonkeyPatch) -> str: + model: Final = "served-tier-cost-model" + monkeypatch.setitem( + litellm.model_cost, + model, + { + "input_cost_per_token": 0.001, + "output_cost_per_token": 0.002, + "input_cost_per_token_priority": 0.01, + "output_cost_per_token_priority": 0.02, + "litellm_provider": "openai", + "mode": "chat", + }, + ) + return model + + +def test_completion_cost_bills_base_when_served_default_overrides_requested_priority( + _local_model_cost_map: None, monkeypatch: pytest.MonkeyPatch +) -> None: + model: Final = _served_tier_cost_model(monkeypatch) + response: Final = ModelResponse( + model=model, + usage=Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150), + ) + setattr(response, "service_tier", "default") + + cost: Final = completion_cost( + completion_response=response, + model=model, + custom_llm_provider="openai", + optional_params={"service_tier": "priority"}, + ) + + assert cost == pytest.approx(100 * 0.001 + 50 * 0.002) + + +def test_completion_cost_bills_priority_when_served_tier_overrides_missing_request( + _local_model_cost_map: None, monkeypatch: pytest.MonkeyPatch +) -> None: + model: Final = _served_tier_cost_model(monkeypatch) + response: Final = ModelResponse( + model=model, + usage=Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150), + ) + setattr(response, "service_tier", "priority") + + cost: Final = completion_cost( + completion_response=response, + model=model, + custom_llm_provider="openai", + ) + + assert cost == pytest.approx(100 * 0.01 + 50 * 0.02) + + +def test_completion_cost_bills_base_when_gemini_serves_on_demand( + _local_model_cost_map: None, monkeypatch: pytest.MonkeyPatch +) -> None: + model: Final = _served_tier_cost_model(monkeypatch) + response: Final = ModelResponse( + model=model, + usage=Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150), + ) + response._hidden_params["provider_specific_fields"] = {"traffic_type": "ON_DEMAND"} + + cost: Final = completion_cost( + completion_response=response, + model=model, + custom_llm_provider="openai", + optional_params={"service_tier": "priority"}, + ) + + assert cost == pytest.approx(100 * 0.001 + 50 * 0.002) From 27c110cb71e5b3e25cd5bac11e91b7eca2d77364 Mon Sep 17 00:00:00 2001 From: "berriai-litellm-provider-info-sync[bot]" <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 13:00:01 -0700 Subject: [PATCH 31/41] feat(bedrock): add openai gpt-6.1-sol global and base rows (#43758) Price-Sync: litellm-providers Co-authored-by: berriai-litellm-provider-info-sync[bot] <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> --- ...odel_prices_and_context_window_backup.json | 67 +++++++++++++++++++ model_prices_and_context_window.json | 67 +++++++++++++++++++ 2 files changed, 134 insertions(+) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 81c0c045a30..8836c7b2b16 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -79186,5 +79186,72 @@ "supports_vision": true, "supports_web_search": true, "supports_xhigh_reasoning_effort": true + }, + "global.openai.gpt-6.1-sol": { + "cache_creation_input_token_cost": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5e-06, + "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_above_272k_tokens": 2e-07, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1e-05, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "source": "https://developers.openai.com/api/docs/pricing", + "supported_endpoints": [ + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true + }, + "openai.gpt-6.1-sol": { + "cache_creation_input_token_cost": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5e-06, + "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_above_272k_tokens": 2e-07, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1e-05, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "source": "https://developers.openai.com/api/docs/pricing", + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true } } diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 81c0c045a30..8836c7b2b16 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -79186,5 +79186,72 @@ "supports_vision": true, "supports_web_search": true, "supports_xhigh_reasoning_effort": true + }, + "global.openai.gpt-6.1-sol": { + "cache_creation_input_token_cost": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5e-06, + "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_above_272k_tokens": 2e-07, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1e-05, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "source": "https://developers.openai.com/api/docs/pricing", + "supported_endpoints": [ + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true + }, + "openai.gpt-6.1-sol": { + "cache_creation_input_token_cost": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5e-06, + "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_above_272k_tokens": 2e-07, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1e-05, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "source": "https://developers.openai.com/api/docs/pricing", + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true } } From d2a574b79113b383e1088e84f0f7f8c41dbe1204 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 13:52:20 -0700 Subject: [PATCH 32/41] perf(router): fetch cooldown state and usage counters in one Redis round trip (#43320) * perf(router): fetch cooldown state and usage counters in one Redis round trip The cooldown filter (CooldownCache) and usage-based-routing-v2 selection (LowestTPMLoggingHandler_v2) each issued their own MGET on every request because they live in different objects. RoutingReadBatch fetches both key sets through DualCache.async_batch_get_cache_shared while the healthy deployments are resolved and hands the usage slice to the strategy, so selection does not read again. Each cache keeps its own memory tier, throttling, reservation rollback and circuit-breaker handling, and the strategy falls back to its own read when the prefetch does not cover its keys. simple-shuffle keeps reading only cooldowns. aresponses no longer issues a second, blocking response-cache read from the worker thread that runs the sync wrapper. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(caching): keep per-cache tier failures inside the shared batch read Wrap the memory-tier prepare and backfill steps of DualCache.async_batch_get_cache_shared so a failing tier degrades that cache's read to None the way async_batch_get_cache does, instead of escaping into routing. Drop the aresponses sync-cache guard: for native Responses models the worker-thread read is the one whose key matches the write, so skipping it broke cached /v1/responses replays. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(router): rename usage key builder so the async cache-call check reads it as a key helper Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(alerting): narrow daily-report cache values before numeric comparison Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(caching): type the shared batch-read helpers and merge Redis results without mutation Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * style: fix import sort in test_dual_cache Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(caching): flatten shared batch read keys without a stacked comprehension Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/caching/dual_cache.py | 172 ++++++++++++----- .../SlackAlerting/slack_alerting.py | 10 +- litellm/router.py | 26 ++- litellm/router_strategy/lowest_tpm_rpm_v2.py | 70 ++++--- litellm/router_utils/cooldown_cache.py | 8 +- litellm/router_utils/routing_read_batch.py | 72 +++++++ tests/unit/caching/test_dual_cache.py | 160 +++++++++++++--- .../router_strategy/test_lowest_tpm_rpm.py | 28 +++ .../router_utils/test_routing_read_batch.py | 178 ++++++++++++++++++ 9 files changed, 622 insertions(+), 102 deletions(-) create mode 100644 litellm/router_utils/routing_read_batch.py create mode 100644 tests/unit/router_utils/test_routing_read_batch.py diff --git a/litellm/caching/dual_cache.py b/litellm/caching/dual_cache.py index 66be77dbb40..3e13848db02 100644 --- a/litellm/caching/dual_cache.py +++ b/litellm/caching/dual_cache.py @@ -8,9 +8,11 @@ Has 4 primary methods: - async_get_cache """ +import itertools import logging import time from collections.abc import Sequence +from dataclasses import dataclass from threading import Lock from typing import TYPE_CHECKING, Any, Final @@ -47,6 +49,16 @@ class LimitedSizeOrderedDict(OrderedDict): super().__setitem__(key, value) +@dataclass(frozen=True) +class PendingBatchRead: + """A batch read that has consulted the in-memory tier and reserved its Redis keys, but not hit Redis yet.""" + + keys: list[str] + result: list[object | None] + redis_keys: list[str] + previous_access_times: dict[str, float | None] + + class DualCache(BaseCache): """ DualCache is a cache implementation that updates both Redis and an in-memory cache simultaneously. @@ -301,6 +313,37 @@ class DualCache(BaseCache): else: self.last_redis_batch_access_time[key] = previous_time + async def _prepare_batch_get(self, keys: list[str], local_only: bool, **kwargs: object) -> PendingBatchRead: + result: list[object | None] = [None] * len(keys) + if self.in_memory_cache is not None: + in_memory_result: Final = await self.in_memory_cache.async_batch_get_cache(keys, **kwargs) + + if in_memory_result is not None: + result = in_memory_result + + redis_keys: list[str] = [] + previous_access_times: dict[str, float | None] = {} + if None in result and self.redis_cache is not None and local_only is False: + redis_keys, previous_access_times = self._reserve_redis_batch_keys(time.time(), keys, result) + return PendingBatchRead( + keys=keys, result=result, redis_keys=redis_keys, previous_access_times=previous_access_times + ) + + async def _apply_batch_get( + self, pending: PendingBatchRead, redis_result: dict[str, object] | None, **kwargs: object + ) -> list[object | None]: + if redis_result is None or all(v is None for v in redis_result.values()): + return pending.result + + merged: Final[list[object | None]] = [ + redis_result.get(key, value) for key, value in zip(pending.keys, pending.result) + ] + if self.in_memory_cache is not None: + for key, value in redis_result.items(): + if value is not None: + await self.in_memory_cache.async_set_cache(key, value, **self._backfill_kwargs(kwargs)) + return merged + async def async_batch_get_cache( self, keys: list, @@ -309,51 +352,22 @@ class DualCache(BaseCache): **kwargs, ): try: - result = [None] * len(keys) - if self.in_memory_cache is not None: - in_memory_result: Final = await self.in_memory_cache.async_batch_get_cache(keys, **kwargs) - - if in_memory_result is not None: - result = in_memory_result - - if None in result and self.redis_cache is not None and local_only is False: - """ - - for the none values in the result - - check the redis cache - """ - current_time: Final = time.time() - sublist_keys, previous_access_times = self._reserve_redis_batch_keys(current_time, keys, result) - - # Only hit Redis if enough time has passed since last access. - if len(sublist_keys) > 0: - try: - # If not found in in-memory cache, try fetching from Redis - redis_result: Final = await self.redis_cache.async_batch_get_cache( - sublist_keys, parent_otel_span=parent_otel_span - ) - except Exception as e: - # Do not throttle subsequent callers if the Redis read fails. - self._rollback_redis_batch_key_reservations(previous_access_times) - if isinstance(e, RedisCircuitBreakerOpenError): - verbose_logger.debug("LiteLLM Cache: async_batch_get_cache served from memory only: %s", e) - return result - raise - - # Short-circuit if redis_result is None or contains only None values - if redis_result is None or all(v is None for v in redis_result.values()): - return result - - # Pre-compute key-to-index mapping for O(1) lookup - key_to_index: Final = {key: i for i, key in enumerate(keys)} - - # Update both result and in-memory cache in a single loop - for key, value in redis_result.items(): - result[key_to_index[key]] = value - - if value is not None and self.in_memory_cache is not None: - await self.in_memory_cache.async_set_cache(key, value, **self._backfill_kwargs(kwargs)) - - return result + pending: Final = await self._prepare_batch_get(keys, local_only, **kwargs) + # Only hit Redis for keys memory could not serve and enough time has passed since last access. + if not pending.redis_keys or self.redis_cache is None: + return pending.result + try: + redis_result: Final = await self.redis_cache.async_batch_get_cache( + pending.redis_keys, parent_otel_span=parent_otel_span + ) + except Exception as e: + # Do not throttle subsequent callers if the Redis read fails. + self._rollback_redis_batch_key_reservations(pending.previous_access_times) + if isinstance(e, RedisCircuitBreakerOpenError): + verbose_logger.debug("LiteLLM Cache: async_batch_get_cache served from memory only: %s", e) + return pending.result + raise + return await self._apply_batch_get(pending, redis_result, **kwargs) except Exception as e: log_redis_failure( verbose_logger, @@ -363,6 +377,74 @@ class DualCache(BaseCache): with_traceback=True, ) + @staticmethod + async def async_batch_get_cache_shared( + reads: Sequence[tuple["DualCache", list[str]]], + parent_otel_span: Span | None = None, + ) -> list[list[object | None] | None]: + """ + `async_batch_get_cache` for several caches in one Redis round trip. + + Each cache still serves what it can from its own in-memory tier, applies its own Redis read + throttle and backfills its own memory; only the Redis MGET is shared. A failed MGET is reported + to every cache that took part in it exactly as its own failed `async_batch_get_cache` would be: + None when the read raised, the in-memory result when the circuit breaker is open. A cache whose + Redis client is not the one the first cache uses falls back to its own read. + """ + results: Final[list[list[object | None] | None]] = [None] * len(reads) + shared_redis: Final = reads[0][0].redis_cache if reads else None + pendings: Final[list[tuple[int, DualCache, PendingBatchRead]]] = [] + for index, (cache, keys) in enumerate(reads): + if shared_redis is None or cache.redis_cache is not shared_redis: + results[index] = await cache.async_batch_get_cache(keys=keys, parent_otel_span=parent_otel_span) + continue + try: + pending = await cache._prepare_batch_get(keys, local_only=False) + except Exception as e: + DualCache._log_shared_batch_get_failure(e) + continue + pendings.append((index, cache, pending)) + results[index] = pending.result + + redis_keys: Final = list( + dict.fromkeys(itertools.chain.from_iterable(pending.redis_keys for _, _, pending in pendings)) + ) + if shared_redis is None or not redis_keys: + return results + try: + redis_result: Final = await shared_redis.async_batch_get_cache( + redis_keys, parent_otel_span=parent_otel_span + ) + except Exception as e: + for index, cache, pending in pendings: + cache._rollback_redis_batch_key_reservations(pending.previous_access_times) + if pending.redis_keys and not isinstance(e, RedisCircuitBreakerOpenError): + results[index] = None + if isinstance(e, RedisCircuitBreakerOpenError): + verbose_logger.debug("LiteLLM Cache: async_batch_get_cache_shared served from memory only: %s", e) + else: + DualCache._log_shared_batch_get_failure(e) + return results + + for index, cache, pending in pendings: + own_result = {key: redis_result[key] for key in pending.redis_keys if key in redis_result} + try: + results[index] = await cache._apply_batch_get(pending, own_result) + except Exception as e: + results[index] = None + DualCache._log_shared_batch_get_failure(e) + return results + + @staticmethod + def _log_shared_batch_get_failure(e: Exception) -> None: + log_redis_failure( + verbose_logger, + logging.ERROR, + "LiteLLM Cache: exception in async_batch_get_cache_shared", + e, + with_traceback=True, + ) + async def async_set_cache(self, key, value, local_only: bool = False, **kwargs): print_verbose(f"async set cache: cache key: {key}; local_only: {local_only}; value: {value}") try: diff --git a/litellm/integrations/SlackAlerting/slack_alerting.py b/litellm/integrations/SlackAlerting/slack_alerting.py index 7c608aac8d9..50f63316a62 100644 --- a/litellm/integrations/SlackAlerting/slack_alerting.py +++ b/litellm/integrations/SlackAlerting/slack_alerting.py @@ -376,8 +376,12 @@ class SlackAlerting(CustomBatchLogger): if combined_metrics_values is None: return False + metric_values: Final[list[float | None]] = [ + val if isinstance(val, (int, float)) else None for val in combined_metrics_values + ] + all_none = True - for val in combined_metrics_values: + for val in metric_values: if val is not None and val > 0: all_none = False break @@ -385,8 +389,8 @@ class SlackAlerting(CustomBatchLogger): if all_none: return False - failed_request_values: Final = combined_metrics_values[: len(failed_request_keys)] # # [1, 2, None, ..] - latency_values: Final = combined_metrics_values[len(failed_request_keys) :] + failed_request_values: Final = metric_values[: len(failed_request_keys)] # # [1, 2, None, ..] + latency_values: Final = metric_values[len(failed_request_keys) :] # find top 5 failed ## Replace None values with a placeholder value (-1 in this case) diff --git a/litellm/router.py b/litellm/router.py index 86a67a8d5ca..a98631b7f97 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -134,7 +134,7 @@ from litellm.router_strategy.least_busy import LeastBusyLoggingHandler from litellm.router_strategy.lowest_cost import LowestCostLoggingHandler from litellm.router_strategy.lowest_latency import LowestLatencyLoggingHandler from litellm.router_strategy.lowest_tpm_rpm import LowestTPMLoggingHandler -from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2 +from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2, PrefetchedUsage from litellm.router_strategy.simple_shuffle import simple_shuffle from litellm.router_strategy.tag_based_routing import ( _get_tags_from_request_kwargs, @@ -259,6 +259,7 @@ from litellm.router_utils.routing_groups import ( parse_routing_groups, validate_routing_strategy, ) +from litellm.router_utils.routing_read_batch import RoutingReadBatch from litellm.scheduler import FlowItem, Scheduler from litellm.types.litellm_params import RoutingStrategyName from litellm.types.llms.openai import ( @@ -1789,6 +1790,7 @@ class Router: messages: list[dict[str, str]] | None, input: str | list | None, request_kwargs: dict | None, + prefetched_usage: PrefetchedUsage | None = None, ) -> Any | None: """ Asks the strategy selector for a deployment. Caller handles @@ -1814,6 +1816,14 @@ class Router: messages=messages, input=input, ) + case "usage-based-routing-v2" if isinstance(selector, LowestTPMLoggingHandler_v2): + return await selector.async_get_available_deployments( + model_group=model, + healthy_deployments=healthy_deployments, + messages=messages, + input=input, + prefetched_usage=prefetched_usage, + ) case "usage-based-routing-v2" | "cost-based-routing": return await selector.async_get_available_deployments( model_group=model, @@ -12925,6 +12935,7 @@ class Router: specific_deployment: bool | None = False, parent_otel_span: Span | None = None, health_check_probe: bool = False, + routing_read_batch: RoutingReadBatch | None = None, ) -> list[dict] | dict: """ Get the healthy deployments for a model. @@ -12977,8 +12988,14 @@ class Router: health_check_probe=health_check_probe, ) - cooldown_deployments: Final = await _async_get_cooldown_deployments( - litellm_router_instance=self, parent_otel_span=parent_otel_span + cooldown_deployments: Final = ( + await _async_get_cooldown_deployments(litellm_router_instance=self, parent_otel_span=parent_otel_span) + if routing_read_batch is None + else await routing_read_batch.async_get_cooldown_deployments( + litellm_router_instance=self, + healthy_deployments=healthy_deployments, + parent_otel_span=parent_otel_span, + ) ) if verbose_router_logger.isEnabledFor(logging.DEBUG): verbose_router_logger.debug("cooldown deployments: %s", cooldown_deployments) @@ -13256,6 +13273,7 @@ class Router: # the hook can replace `model` and routing-group lookup must key # off the final model name. strategy, strategy_selector = self._get_routing_context(model, request_kwargs) + routing_read_batch: Final = RoutingReadBatch.for_strategy(strategy, strategy_selector) healthy_deployments: Final = await self.async_get_healthy_deployments( model=model, @@ -13264,6 +13282,7 @@ class Router: input=input, specific_deployment=specific_deployment, parent_otel_span=parent_otel_span, + routing_read_batch=routing_read_batch, ) if isinstance(healthy_deployments, dict): await self._async_override_selector_pre_call_check( @@ -13294,6 +13313,7 @@ class Router: messages=messages, input=input, request_kwargs=request_kwargs, + prefetched_usage=routing_read_batch.prefetched_usage if routing_read_batch is not None else None, ) if deployment is None: exception: Final = await async_raise_no_deployment_exception( diff --git a/litellm/router_strategy/lowest_tpm_rpm_v2.py b/litellm/router_strategy/lowest_tpm_rpm_v2.py index a2acce5fcb5..909b47833cb 100644 --- a/litellm/router_strategy/lowest_tpm_rpm_v2.py +++ b/litellm/router_strategy/lowest_tpm_rpm_v2.py @@ -1,7 +1,8 @@ #### What this does #### # identifies lowest tpm deployment import random -from collections.abc import Sequence +from collections.abc import Mapping, Sequence +from dataclasses import dataclass from typing import TYPE_CHECKING, Any, Final import httpx @@ -31,6 +32,26 @@ class RoutingArgs(LiteLLMPydanticObjectBase): ttl: int = 1 * 60 # 1min (RPM/TPM expire key) +@dataclass(frozen=True) +class PrefetchedUsage: + """ + tpm/rpm counter values another read of this request already fetched from the router cache. + + `values` is None when that read failed, which is what `async_batch_get_cache` returns on failure. + """ + + keys: frozenset[str] + values: Mapping[str, object] | None + + def covers(self, keys: Sequence[str]) -> bool: + return self.keys.issuperset(keys) + + def values_for(self, keys: Sequence[str]) -> list[object | None] | None: + if self.values is None: + return None + return [self.values.get(key) for key in keys] + + class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): """ Updated version of TPM/RPM Logging. @@ -412,17 +433,35 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): else: return None + def usage_counter_keys(self, healthy_deployments: list) -> tuple[list[str], list[str]]: + """The `::tpm:` and `::rpm:` counter keys selection reads.""" + current_minute: Final = get_utc_datetime().strftime("%H-%M") + + tpm_keys: Final[list[str]] = [] + rpm_keys: Final[list[str]] = [] + for m in healthy_deployments: + if isinstance(m, dict): + id = m.get("model_info", {}).get( + "id" + ) # a deployment should always have an 'id'. this is set in router.py + deployment_name = m.get("litellm_params", {}).get("model") + tpm_keys.append(f"{id}:{deployment_name}:tpm:{current_minute}") + rpm_keys.append(f"{id}:{deployment_name}:rpm:{current_minute}") + return tpm_keys, rpm_keys + async def async_get_available_deployments( self, model_group: str, healthy_deployments: list, messages: list[dict[str, str]] | None = None, input: str | list | None = None, + prefetched_usage: PrefetchedUsage | None = None, ): """ Async implementation of get deployments. - Reduces time to retrieve the tpm/rpm values from cache + Reduces time to retrieve the tpm/rpm values from cache. `prefetched_usage` skips the cache + read when it already holds this request's counters (see `RoutingReadBatch`). """ # get list of potential deployments verbose_router_logger.debug( @@ -431,28 +470,15 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): healthy_deployments, ) - dt: Final = get_utc_datetime() - current_minute: Final = dt.strftime("%H-%M") - - tpm_keys: Final = [] - rpm_keys: Final = [] - for m in healthy_deployments: - if isinstance(m, dict): - id = m.get("model_info", {}).get( - "id" - ) # a deployment should always have an 'id'. this is set in router.py - deployment_name = m.get("litellm_params", {}).get("model") - tpm_key = f"{id}:{deployment_name}:tpm:{current_minute}" - rpm_key = f"{id}:{deployment_name}:rpm:{current_minute}" - - tpm_keys.append(tpm_key) - rpm_keys.append(rpm_key) - + tpm_keys, rpm_keys = self.usage_counter_keys(healthy_deployments) combined_tpm_rpm_keys: Final = tpm_keys + rpm_keys - combined_tpm_rpm_values: Final = await self.router_cache.async_batch_get_cache( - keys=combined_tpm_rpm_keys - ) # [1, 2, None, ..] + if prefetched_usage is not None and prefetched_usage.covers(combined_tpm_rpm_keys): + combined_tpm_rpm_values = prefetched_usage.values_for(combined_tpm_rpm_keys) + else: + combined_tpm_rpm_values = await self.router_cache.async_batch_get_cache( + keys=combined_tpm_rpm_keys + ) # [1, 2, None, ..] if combined_tpm_rpm_values is not None: tpm_values = combined_tpm_rpm_values[: len(tpm_keys)] diff --git a/litellm/router_utils/cooldown_cache.py b/litellm/router_utils/cooldown_cache.py index ef29f7d8fd3..187215d3d16 100644 --- a/litellm/router_utils/cooldown_cache.py +++ b/litellm/router_utils/cooldown_cache.py @@ -4,7 +4,7 @@ Wrapper around router cache. Meant to handle model cooldown logic import functools import time -from collections.abc import Mapping +from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any, Final from typing_extensions import TypedDict @@ -163,6 +163,12 @@ class CooldownCache: keys: Final = [CooldownCache.get_cooldown_cache_key(model_id) for model_id in model_ids] results: Final = await self.cooldown_store.async_batch_get_cache(keys=keys, parent_otel_span=parent_otel_span) + return self.active_cooldowns_from_results(model_ids, results) + + def active_cooldowns_from_results( + self, model_ids: list[str], results: Sequence[object] | None + ) -> list[tuple[str, CooldownCacheValue]]: + """The cooldowns still active in a `cooldown_store` batch read of `get_cooldown_cache_key(model_id)` per id.""" active_cooldowns: Final[list[tuple[str, CooldownCacheValue]]] = [] if results is None or all(v is None for v in results): diff --git a/litellm/router_utils/routing_read_batch.py b/litellm/router_utils/routing_read_batch.py new file mode 100644 index 00000000000..e9410d0e586 --- /dev/null +++ b/litellm/router_utils/routing_read_batch.py @@ -0,0 +1,72 @@ +""" +One Redis round trip for the reads a request needs before a deployment can be picked. + +The cooldown filter (`CooldownCache`, its own `DualCache`) and usage-based selection +(`LowestTPMLoggingHandler_v2`, the router cache) each issue their own MGET because they live in +different objects. `RoutingReadBatch` fetches both key sets in one +`DualCache.async_batch_get_cache_shared` while the healthy deployments are being resolved and hands +the usage slice to the strategy, so selection does not read again. +""" + +from typing import TYPE_CHECKING, Any, Final + +from litellm._logging import verbose_router_logger +from litellm.caching.dual_cache import DualCache +from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2, PrefetchedUsage +from litellm.router_utils.cooldown_cache import CooldownCache + +if TYPE_CHECKING: + from opentelemetry.trace import Span as _Span + + from litellm.router import Router as _Router + + LitellmRouter = _Router + Span = _Span +else: + LitellmRouter = Any + Span = Any + + +class RoutingReadBatch: + def __init__(self, usage_selector: LowestTPMLoggingHandler_v2) -> None: + self.usage_selector: Final = usage_selector + self.prefetched_usage: PrefetchedUsage | None = None + + @staticmethod + def for_strategy(strategy: str | None, selector: object) -> "RoutingReadBatch | None": + if strategy == "usage-based-routing-v2" and isinstance(selector, LowestTPMLoggingHandler_v2): + return RoutingReadBatch(usage_selector=selector) + return None + + async def async_get_cooldown_deployments( + self, + litellm_router_instance: LitellmRouter, + healthy_deployments: list, + parent_otel_span: Span | None, + ) -> list[str]: + """ + `_async_get_cooldown_deployments`, with the strategy's tpm/rpm counters for + `healthy_deployments` fetched in the same MGET and kept as `prefetched_usage`. + """ + model_ids: Final = litellm_router_instance.get_model_ids() + cooldown_keys: Final = [CooldownCache.get_cooldown_cache_key(model_id) for model_id in model_ids] + tpm_keys, rpm_keys = self.usage_selector.usage_counter_keys(healthy_deployments) + usage_keys: Final = tpm_keys + rpm_keys + + cooldown_results, usage_values = await DualCache.async_batch_get_cache_shared( + [ + (litellm_router_instance.cooldown_cache.cooldown_store, cooldown_keys), + (self.usage_selector.router_cache, usage_keys), + ], + parent_otel_span=parent_otel_span, + ) + self.prefetched_usage = PrefetchedUsage( + keys=frozenset(usage_keys), + values=None if usage_values is None else dict(zip(usage_keys, usage_values)), + ) + + cooldown_models: Final = litellm_router_instance.cooldown_cache.active_cooldowns_from_results( + model_ids, cooldown_results + ) + verbose_router_logger.debug("retrieve cooldown models: %s", cooldown_models) + return [model_id for model_id, _ in cooldown_models] diff --git a/tests/unit/caching/test_dual_cache.py b/tests/unit/caching/test_dual_cache.py index 5f59de9cca5..eb2f19ac377 100644 --- a/tests/unit/caching/test_dual_cache.py +++ b/tests/unit/caching/test_dual_cache.py @@ -6,18 +6,16 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest -from litellm.constants import DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE from litellm.caching.dual_cache import DualCache from litellm.caching.in_memory_cache import InMemoryCache from litellm.caching.redis_cache import RedisCache, _redis_circuit_breaker_guard, _redis_circuit_breaker_guard_sync +from litellm.constants import DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE from litellm.types.caching import RedisPipelineIncrementOperation @pytest.mark.asyncio async def test_dual_cache_async_batch_get_cache_coalesces_concurrent_redis_reads(): - dual_cache = DualCache( - redis_cache=MagicMock(spec=RedisCache), default_redis_batch_cache_expiry=10 - ) + dual_cache = DualCache(redis_cache=MagicMock(spec=RedisCache), default_redis_batch_cache_expiry=10) keys = ["shared_a", "shared_b"] start_gate = asyncio.Event() @@ -44,9 +42,7 @@ async def test_dual_cache_async_batch_get_cache_coalesces_concurrent_redis_reads @pytest.mark.asyncio async def test_dual_cache_async_batch_get_cache_rolls_back_redis_reservation_on_error(): - dual_cache = DualCache( - redis_cache=MagicMock(spec=RedisCache), default_redis_batch_cache_expiry=10 - ) + dual_cache = DualCache(redis_cache=MagicMock(spec=RedisCache), default_redis_batch_cache_expiry=10) keys = ["shared_a", "shared_b"] with patch.object( @@ -116,9 +112,7 @@ def test_dual_cache_batch_get_cache_only_reads_missing_keys_from_redis(): def test_dual_cache_batch_get_cache_throttles_repeat_redis_reads(): mock_redis = _redis_mock_for_sync_batch({"absent_key": None}) - dual_cache = DualCache( - in_memory_cache=InMemoryCache(), redis_cache=mock_redis, default_redis_batch_cache_expiry=10 - ) + dual_cache = DualCache(in_memory_cache=InMemoryCache(), redis_cache=mock_redis, default_redis_batch_cache_expiry=10) first = dual_cache.batch_get_cache(keys=["absent_key"]) second = dual_cache.batch_get_cache(keys=["absent_key"]) @@ -131,9 +125,7 @@ def test_dual_cache_batch_get_cache_throttles_repeat_redis_reads(): def test_dual_cache_batch_get_cache_rolls_back_redis_reservation_on_error(): mock_redis = MagicMock(spec=RedisCache) mock_redis.batch_get_cache.side_effect = RuntimeError("redis unavailable") - dual_cache = DualCache( - in_memory_cache=InMemoryCache(), redis_cache=mock_redis, default_redis_batch_cache_expiry=10 - ) + dual_cache = DualCache(in_memory_cache=InMemoryCache(), redis_cache=mock_redis, default_redis_batch_cache_expiry=10) first_result = dual_cache.batch_get_cache(keys=["shared_a"]) second_result = dual_cache.batch_get_cache(keys=["shared_a"]) @@ -146,9 +138,7 @@ def test_dual_cache_batch_get_cache_rolls_back_redis_reservation_on_error(): def test_dual_cache_batch_get_cache_returns_memory_only_when_redis_read_is_throttled(): mock_redis = _redis_mock_for_sync_batch({"throttled_key": "redis_value"}) - dual_cache = DualCache( - in_memory_cache=InMemoryCache(), redis_cache=mock_redis, default_redis_batch_cache_expiry=10 - ) + dual_cache = DualCache(in_memory_cache=InMemoryCache(), redis_cache=mock_redis, default_redis_batch_cache_expiry=10) dual_cache.last_redis_batch_access_time["throttled_key"] = time.time() result = dual_cache.batch_get_cache(keys=["throttled_key"]) @@ -257,9 +247,7 @@ async def test_dual_cache_batch_redis_backfill_injects_default_in_memory_ttl(): default_in_memory_ttl, same as the single-key path.""" in_memory_cache = InMemoryCache(default_ttl=600) mock_redis = MagicMock(spec=RedisCache) - mock_redis.async_batch_get_cache = AsyncMock( - return_value={"batch_backfill_key": "redis_value"} - ) + mock_redis.async_batch_get_cache = AsyncMock(return_value={"batch_backfill_key": "redis_value"}) dual_cache = DualCache( in_memory_cache=in_memory_cache, redis_cache=mock_redis, @@ -371,9 +359,7 @@ async def test_circuit_breaker_open_skips_redis(): class FakeRedis: def __init__(self): - self._circuit_breaker = RedisCircuitBreaker( - failure_threshold=3, recovery_timeout=60 - ) + self._circuit_breaker = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60) self._circuit_breaker._state = "open" self._circuit_breaker._opened_at = time.time() self.call_count = 0 @@ -426,9 +412,7 @@ def test_circuit_breaker_half_open_concurrent_calls_are_fast_failed(): # All subsequent concurrent callers: HALF_OPEN → fast-fail (return True) for _ in range(10): - assert ( - cb.is_open() is True - ), "concurrent callers should be fast-failed in HALF_OPEN" + assert cb.is_open() is True, "concurrent callers should be fast-failed in HALF_OPEN" def test_circuit_breaker_disabled_never_opens(): @@ -472,9 +456,7 @@ async def test_circuit_breaker_disabled_guard_always_calls_method(): class FakeRedis: def __init__(self): - self._circuit_breaker = RedisCircuitBreaker( - failure_threshold=1, recovery_timeout=60, enabled=False - ) + self._circuit_breaker = RedisCircuitBreaker(failure_threshold=1, recovery_timeout=60, enabled=False) self.call_count = 0 @_redis_circuit_breaker_guard @@ -791,3 +773,125 @@ async def test_async_delete_cache_keys_on_empty_list_touches_no_backend(): await dual_cache.async_delete_cache_keys([]) redis_cache.delete_cache_keys.assert_not_awaited() + + +def _recording_redis(values: dict) -> MagicMock: + redis = MagicMock(spec=RedisCache) + redis.async_batch_get_cache = AsyncMock( + side_effect=lambda key_list, parent_otel_span=None: {key: values.get(key) for key in key_list} + ) + return redis + + +@pytest.mark.asyncio +async def test_shared_batch_read_issues_one_mget_for_two_caches_and_backfills_each_one_separately(): + redis = _recording_redis({"a1": 1, "b2": "x"}) + first = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10) + second = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10) + + results = await DualCache.async_batch_get_cache_shared([(first, ["a1", "a2"]), (second, ["b1", "b2"])]) + + assert results == [[1, None], [None, "x"]] + assert redis.async_batch_get_cache.await_count == 1 + assert redis.async_batch_get_cache.await_args.args[0] == ["a1", "a2", "b1", "b2"] + assert first.in_memory_cache.get_cache("a1") == 1 + assert second.in_memory_cache.get_cache("b2") == "x" + assert first.in_memory_cache.get_cache("b2") is None, "backfill leaked into the other cache" + + +@pytest.mark.asyncio +async def test_shared_batch_read_serves_memory_hits_and_throttles_like_the_separate_reads(): + redis = _recording_redis({"a2": 2, "b1": 3}) + first = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10) + first.in_memory_cache.set_cache("a1", 5) + second = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10) + second.in_memory_cache.set_cache("b1", 3) + + results = await DualCache.async_batch_get_cache_shared([(first, ["a1", "a2"]), (second, ["b1"])]) + + assert results == [[5, 2], [3]] + assert redis.async_batch_get_cache.await_args.args[0] == ["a2"], "memory hits must not hit Redis" + + first.in_memory_cache.delete_cache("a2") + results = await DualCache.async_batch_get_cache_shared([(first, ["a1", "a2"]), (second, ["b1"])]) + + assert results == [[5, None], [3]] + assert redis.async_batch_get_cache.await_count == 1, "a2 was read within the batch expiry, so it is throttled" + + +@pytest.mark.asyncio +async def test_shared_batch_read_failure_degrades_exactly_like_two_failed_reads(): + redis = MagicMock(spec=RedisCache) + redis.async_batch_get_cache = AsyncMock(side_effect=ConnectionError("redis unavailable")) + first = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10) + second = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10) + third = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10) + third.in_memory_cache.set_cache("c1", "memory") + + shared = await DualCache.async_batch_get_cache_shared([(first, ["a1"]), (second, ["b1"]), (third, ["c1"])]) + separate = [ + await first.async_batch_get_cache(keys=["a1"]), + await second.async_batch_get_cache(keys=["b1"]), + await third.async_batch_get_cache(keys=["c1"]), + ] + + assert shared == separate == [None, None, ["memory"]] + assert "a1" not in first.last_redis_batch_access_time + assert "b1" not in second.last_redis_batch_access_time + + +@pytest.mark.asyncio +async def test_shared_batch_read_with_an_open_breaker_keeps_memory_hits_and_releases_reservations(): + first = _dual_cache_with_open_breaker_and_a_memory_hit() + second = DualCache( + in_memory_cache=InMemoryCache(), redis_cache=first.redis_cache, default_redis_batch_cache_expiry=10 + ) + + results = await DualCache.async_batch_get_cache_shared([(first, ["k1", "k2"]), (second, ["k3"])]) + + assert results == [["v1", None], [None]] + assert "k2" not in first.last_redis_batch_access_time + assert "k3" not in second.last_redis_batch_access_time + + +@pytest.mark.asyncio +async def test_shared_batch_read_falls_back_to_a_caches_own_read_when_its_redis_client_differs(): + first_redis = _recording_redis({"a1": 1}) + second_redis = _recording_redis({"b1": 2}) + first = DualCache(in_memory_cache=InMemoryCache(), redis_cache=first_redis, default_redis_batch_cache_expiry=10) + second = DualCache(in_memory_cache=InMemoryCache(), redis_cache=second_redis, default_redis_batch_cache_expiry=10) + memory_only = DualCache(in_memory_cache=InMemoryCache(), redis_cache=None) + memory_only.in_memory_cache.set_cache("m1", "m") + + results = await DualCache.async_batch_get_cache_shared( + [(first, ["a1"]), (second, ["b1"]), (memory_only, ["m1", "m2"])] + ) + + assert results == [[1], [2], ["m", None]] + assert first_redis.async_batch_get_cache.await_args.args[0] == ["a1"] + assert second_redis.async_batch_get_cache.await_args.args[0] == ["b1"] + + +@pytest.mark.asyncio +async def test_shared_batch_read_keeps_a_caches_own_tier_failure_to_itself_like_the_separate_read(): + redis = _recording_redis({"a1": 1, "b1": 2, "c1": 3}) + broken_memory_read = DualCache( + in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10 + ) + broken_memory_read.in_memory_cache.async_batch_get_cache = AsyncMock(side_effect=RuntimeError("memory read")) + broken_backfill = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10) + broken_backfill.in_memory_cache.async_set_cache = AsyncMock(side_effect=RuntimeError("memory write")) + healthy = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10) + + shared = await DualCache.async_batch_get_cache_shared( + [(broken_memory_read, ["a1"]), (broken_backfill, ["b1"]), (healthy, ["c1"])] + ) + broken_backfill.last_redis_batch_access_time.clear() + separate = [ + await broken_memory_read.async_batch_get_cache(keys=["a1"]), + await broken_backfill.async_batch_get_cache(keys=["b1"]), + await healthy.async_batch_get_cache(keys=["c1"]), + ] + + assert shared == separate == [None, None, [3]] + assert redis.async_batch_get_cache.await_args_list[0].args[0] == ["b1", "c1"] diff --git a/tests/unit/router_strategy/test_lowest_tpm_rpm.py b/tests/unit/router_strategy/test_lowest_tpm_rpm.py index 7b13b196d5b..0fa11cda20c 100644 --- a/tests/unit/router_strategy/test_lowest_tpm_rpm.py +++ b/tests/unit/router_strategy/test_lowest_tpm_rpm.py @@ -1,7 +1,12 @@ from datetime import datetime, timedelta from typing import Final +from unittest.mock import AsyncMock + +import pytest from litellm import Router +from litellm.caching.dual_cache import DualCache +from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2, PrefetchedUsage from litellm.types.router import DeploymentTypedDict, LiteLLMParamsTypedDict MODEL_GROUP: Final = "lowest-tpm-router" @@ -52,3 +57,26 @@ def test_usage_based_routing_v1_selects_the_lowest_recorded_tpm() -> None: ) assert deployment["model_info"]["id"] == LOW_USAGE_DEPLOYMENT_ID + + +@pytest.mark.asyncio +async def test_v2_async_selection_uses_prefetched_counters_only_when_they_cover_its_keys(): + router_cache = DualCache() + router_cache.async_batch_get_cache = AsyncMock(return_value=[100, 10, None, None]) # type: ignore[method-assign] + strategy = LowestTPMLoggingHandler_v2(router_cache=router_cache) + deployments = [ + {"model_name": "g", "litellm_params": {"model": "m"}, "model_info": {"id": "a"}}, + {"model_name": "g", "litellm_params": {"model": "m"}, "model_info": {"id": "b"}}, + ] + tpm_keys, rpm_keys = strategy.usage_counter_keys(deployments) + keys = tpm_keys + rpm_keys + + covering = PrefetchedUsage(keys=frozenset(keys), values=dict(zip(keys, [10, 100, None, None]))) + chosen = await strategy.async_get_available_deployments(model_group="g", healthy_deployments=deployments, prefetched_usage=covering) + assert chosen["model_info"]["id"] == "a", "the prefetched counters say a is the lowest" + router_cache.async_batch_get_cache.assert_not_awaited() + + stale = PrefetchedUsage(keys=frozenset(keys[:1]), values={keys[0]: 10}) + chosen = await strategy.async_get_available_deployments(model_group="g", healthy_deployments=deployments, prefetched_usage=stale) + assert chosen["model_info"]["id"] == "b", "counters that do not cover this minute's keys are read again" + router_cache.async_batch_get_cache.assert_awaited_once_with(keys=keys) diff --git a/tests/unit/router_utils/test_routing_read_batch.py b/tests/unit/router_utils/test_routing_read_batch.py new file mode 100644 index 00000000000..73be5fd4a3a --- /dev/null +++ b/tests/unit/router_utils/test_routing_read_batch.py @@ -0,0 +1,178 @@ +""" +One Redis round trip per request for the router's pre-call reads. + +Before `RoutingReadBatch`, `async_get_available_deployment` issued one MGET for the cooldown keys +(`CooldownCache`) and a second one for the tpm/rpm counters (`LowestTPMLoggingHandler_v2`). +""" + +import time +from unittest.mock import AsyncMock, MagicMock + +import pytest + +import litellm +from litellm import Router +from litellm.caching.redis_cache import RedisCache + +_MODEL_GROUP = "claude" +_MESSAGES = [{"role": "user", "content": "ping"}] + + +def _deployment(deployment_id: str) -> dict: + return { + "model_name": _MODEL_GROUP, + "litellm_params": {"model": "anthropic/claude-x", "api_key": "test", "mock_response": "pong"}, + "model_info": {"id": deployment_id}, + } + + +def _redis_answering(values_by_key_prefix: dict[str, object]) -> MagicMock: + """A Redis double that answers each key from its minute-less prefix and records every MGET.""" + + def _mget(key_list, parent_otel_span=None): + return {key: values_by_key_prefix.get(key.rsplit(":", 1)[0], values_by_key_prefix.get(key)) for key in key_list} + + redis = MagicMock(spec=RedisCache) + redis.async_batch_get_cache = AsyncMock(side_effect=_mget) + return redis + + +def _router(redis: MagicMock, routing_strategy: str) -> Router: + router = Router( + model_list=[_deployment("dep-a"), _deployment("dep-b")], + routing_strategy=routing_strategy, + ) + router._update_redis_cache(cache=redis) + return router + + +def _redis_key_families(redis: MagicMock) -> list[list[str]]: + return [ + sorted(key.rsplit(":", 1)[0] if ":tpm:" in key or ":rpm:" in key else key for key in call.args[0]) + for call in redis.async_batch_get_cache.await_args_list + ] + + +def _cooldown(seconds: float) -> dict: + return {"exception_received": "429", "status_code": "429", "timestamp": time.time(), "cooldown_time": seconds} + + +@pytest.mark.asyncio +async def test_usage_based_routing_reads_cooldowns_and_counters_in_one_redis_round_trip(): + redis = _redis_answering({}) + router = _router(redis, "usage-based-routing-v2") + + deployment = await router.async_get_available_deployment( + model=_MODEL_GROUP, request_kwargs={}, messages=_MESSAGES + ) + + assert deployment["model_info"]["id"] in {"dep-a", "dep-b"} + assert _redis_key_families(redis) == [ + [ + "dep-a:anthropic/claude-x:rpm", + "dep-a:anthropic/claude-x:tpm", + "dep-b:anthropic/claude-x:rpm", + "dep-b:anthropic/claude-x:tpm", + "deployment:dep-a:cooldown", + "deployment:dep-b:cooldown", + ] + ], "cooldown state and usage counters must arrive in one MGET" + + +@pytest.mark.asyncio +async def test_simple_shuffle_still_reads_only_cooldowns(): + redis = _redis_answering({}) + router = _router(redis, "simple-shuffle") + + await router.async_get_available_deployment(model=_MODEL_GROUP, request_kwargs={}, messages=_MESSAGES) + + assert _redis_key_families(redis) == [["deployment:dep-a:cooldown", "deployment:dep-b:cooldown"]] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("tpm_a", "tpm_b", "expected"), + [(100, 10, "dep-b"), (10, 100, "dep-a"), (None, 10, "dep-a"), (10, None, "dep-b")], +) +async def test_batched_counters_pick_the_deployment_the_strategy_picks_reading_alone(tpm_a, tpm_b, expected): + counters = {"dep-a:anthropic/claude-x:tpm": tpm_a, "dep-b:anthropic/claude-x:tpm": tpm_b} + routed = _router(_redis_answering(counters), "usage-based-routing-v2") + alone = _router(_redis_answering(counters), "usage-based-routing-v2") + + routed_choice = await routed.async_get_available_deployment( + model=_MODEL_GROUP, request_kwargs={}, messages=_MESSAGES + ) + alone_choice = await alone.lowesttpm_logger_v2.async_get_available_deployments( + model_group=_MODEL_GROUP, healthy_deployments=alone.model_list, messages=_MESSAGES + ) + + assert routed_choice["model_info"]["id"] == alone_choice["model_info"]["id"] == expected + + +@pytest.mark.asyncio +async def test_batched_read_still_excludes_a_cooled_down_deployment(): + redis = _redis_answering( + { + "dep-a:anthropic/claude-x:tpm": 100, + "dep-b:anthropic/claude-x:tpm": 10, + "deployment:dep-b:cooldown": _cooldown(seconds=60), + } + ) + router = _router(redis, "usage-based-routing-v2") + + deployment = await router.async_get_available_deployment( + model=_MODEL_GROUP, request_kwargs={}, messages=_MESSAGES + ) + + assert deployment["model_info"]["id"] == "dep-a", "dep-b has the lowest tpm but is cooling down" + assert redis.async_batch_get_cache.await_count == 1 + + +@pytest.mark.asyncio +async def test_batched_read_ignores_an_expired_cooldown(): + redis = _redis_answering( + { + "dep-a:anthropic/claude-x:tpm": 100, + "dep-b:anthropic/claude-x:tpm": 10, + "deployment:dep-b:cooldown": _cooldown(seconds=-1), + } + ) + router = _router(redis, "usage-based-routing-v2") + + deployment = await router.async_get_available_deployment( + model=_MODEL_GROUP, request_kwargs={}, messages=_MESSAGES + ) + + assert deployment["model_info"]["id"] == "dep-b" + + +@pytest.mark.asyncio +async def test_a_failed_batched_read_degrades_like_the_two_failed_reads_did(): + redis = MagicMock(spec=RedisCache) + redis.async_batch_get_cache = AsyncMock(side_effect=ConnectionError("redis unavailable")) + routed = _router(redis, "usage-based-routing-v2") + alone = _router(redis, "usage-based-routing-v2") + + with pytest.raises(litellm.RateLimitError, match="No deployments available") as routed_error: + await routed.async_get_available_deployment(model=_MODEL_GROUP, request_kwargs={}, messages=_MESSAGES) + with pytest.raises(litellm.RateLimitError, match="No deployments available") as alone_error: + await alone.lowesttpm_logger_v2.async_get_available_deployments( + model_group=_MODEL_GROUP, healthy_deployments=alone.model_list, messages=_MESSAGES + ) + + assert str(routed_error.value) == str(alone_error.value) + assert len(routed.cache.last_redis_batch_access_time) == 0, "a failed read must not throttle the next one" + assert len(routed.cooldown_cache.cooldown_store.last_redis_batch_access_time) == 0 + + +@pytest.mark.asyncio +async def test_a_failed_batched_read_leaves_simple_shuffle_routing(): + redis = MagicMock(spec=RedisCache) + redis.async_batch_get_cache = AsyncMock(side_effect=ConnectionError("redis unavailable")) + router = _router(redis, "simple-shuffle") + + deployment = await router.async_get_available_deployment( + model=_MODEL_GROUP, request_kwargs={}, messages=_MESSAGES + ) + + assert deployment["model_info"]["id"] in {"dep-a", "dep-b"} From 2d034bb35b7f8371404dca11cf68c46c3e0c88ac Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 15:05:33 -0700 Subject: [PATCH 33/41] perf(proxy): hold one spend counter batch across admission and across post-call accounting (#43369) Auth's spend counter MGET scope spans common checks, model budget check and reservation; reservation increments go out as one pipeline; post-call reconcile adjustments ride the ordinary increment pipeline and update_cache uses one batched read. Over-budget reservation counters are charged one at a time so a rejection never touches the counters after it; post-call counter keys are derived from ids without validating a UserAPIKeyAuth. Resolves LIT-8881 Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/caching/dual_cache.py | 14 +- litellm/proxy/auth/user_api_key_auth.py | 63 ++-- .../proxy/hooks/proxy_track_cost_callback.py | 64 +++- litellm/proxy/proxy_server.py | 135 +++++++-- .../spend_tracking/budget_reservation.py | 251 +++++++++++----- .../spend_tracking/spend_counter_batch.py | 60 ++-- .../proxy/auth/test_user_api_key_auth.py | 64 ++++ .../hooks/test_proxy_track_cost_callback.py | 2 + .../proxy/proxy_server/test_spend_counters.py | 33 ++- .../test_budget_reservation_redis_failure.py | 25 +- .../test_spend_counter_batch.py | 277 ++++++++++++++++-- .../proxy/test_budget_reservation.py | 99 +++++-- tests/test_litellm/proxy/test_proxy_server.py | 29 +- tests/unit/proxy/auth/test_jwt.py | 5 +- 14 files changed, 874 insertions(+), 247 deletions(-) diff --git a/litellm/caching/dual_cache.py b/litellm/caching/dual_cache.py index 3e13848db02..1d7afcbee8f 100644 --- a/litellm/caching/dual_cache.py +++ b/litellm/caching/dual_cache.py @@ -313,7 +313,9 @@ class DualCache(BaseCache): else: self.last_redis_batch_access_time[key] = previous_time - async def _prepare_batch_get(self, keys: list[str], local_only: bool, **kwargs: object) -> PendingBatchRead: + async def _prepare_batch_get( + self, keys: list[str], local_only: bool, throttle_redis: bool = True, **kwargs: object + ) -> PendingBatchRead: result: list[object | None] = [None] * len(keys) if self.in_memory_cache is not None: in_memory_result: Final = await self.in_memory_cache.async_batch_get_cache(keys, **kwargs) @@ -324,7 +326,10 @@ class DualCache(BaseCache): redis_keys: list[str] = [] previous_access_times: dict[str, float | None] = {} if None in result and self.redis_cache is not None and local_only is False: - redis_keys, previous_access_times = self._reserve_redis_batch_keys(time.time(), keys, result) + if throttle_redis: + redis_keys, previous_access_times = self._reserve_redis_batch_keys(time.time(), keys, result) + else: + redis_keys = [key for key, value in zip(keys, result) if value is None] return PendingBatchRead( keys=keys, result=result, redis_keys=redis_keys, previous_access_times=previous_access_times ) @@ -349,10 +354,13 @@ class DualCache(BaseCache): keys: list, parent_otel_span: Span | None = None, local_only: bool = False, + throttle_redis: bool = True, **kwargs, ): + """With ``throttle_redis`` False every key memory cannot serve is read from Redis, exactly as a per-key + ``async_get_cache`` would read it, instead of skipping keys that missed within ``redis_batch_cache_expiry``.""" try: - pending: Final = await self._prepare_batch_get(keys, local_only, **kwargs) + pending: Final = await self._prepare_batch_get(keys, local_only, throttle_redis, **kwargs) # Only hit Redis for keys memory could not serve and enough time has passed since last access. if not pending.redis_keys or self.redis_cache is None: return pending.result diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index e3ce9bcd850..ed4d63fb9fc 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -3006,41 +3006,40 @@ async def _run_centralized_common_checks( skip_budget_checks=skip_budget_checks, project_object=project_object, ) + if not skip_budget_checks: + await _check_team_model_budget( + valid_token=user_api_key_auth_obj, + model_max_budget_limiter=model_max_budget_limiter, + models=_get_model_names_for_budget_checks( + model=_get_model_from_request_context( + request_data=request_data, + route=route, + request=request, + llm_router=llm_router, + team_id=user_api_key_auth_obj.team_id, + ) + ), + ) + + await _reserve_budget_after_common_checks( + user_api_key_auth_obj=user_api_key_auth_obj, + request=request, + request_data=request_data, + route=route, + llm_router=llm_router, + team_object=team_object, + user_object=user_object, + end_user_id=end_user_id, + end_user_object=end_user_object, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + skip_budget_checks=skip_budget_checks, + general_settings=general_settings, + ) finally: release_spend_counter_batch() - if not skip_budget_checks: - await _check_team_model_budget( - valid_token=user_api_key_auth_obj, - model_max_budget_limiter=model_max_budget_limiter, - models=_get_model_names_for_budget_checks( - model=_get_model_from_request_context( - request_data=request_data, - route=route, - request=request, - llm_router=llm_router, - team_id=user_api_key_auth_obj.team_id, - ) - ), - ) - - await _reserve_budget_after_common_checks( - user_api_key_auth_obj=user_api_key_auth_obj, - request=request, - request_data=request_data, - route=route, - llm_router=llm_router, - team_object=team_object, - user_object=user_object, - end_user_id=end_user_id, - end_user_object=end_user_object, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - proxy_logging_obj=proxy_logging_obj, - skip_budget_checks=skip_budget_checks, - general_settings=general_settings, - ) - async def _noop_none() -> None: """Sentinel coroutine for asyncio.gather when a fetch is unnecessary diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 0178465739b..05995d22293 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -30,6 +30,7 @@ from litellm.proxy.db.db_spend_update_writer import ( get_llm_router, ) from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup +from litellm.proxy.spend_tracking.spend_counter_batch import post_call_counter_keys, spend_counter_batch_scope from litellm.proxy.spend_tracking.spend_event import ( ObjectMapping, SpendEventBuildError, @@ -695,10 +696,67 @@ async def _update_database_and_spend_counters( model_access_groups: Sequence[str] | None = None, project_id: str | None = None, ) -> bool: + """The reservation is reconciled before the spend is persisted, from its own read. One spend counter batch then + spans the database write and the counter update, so the post-call counters are read with a single MGET after the + write and their increments leave in a single pipeline.""" + from litellm.proxy.proxy_server import spend_counter_cache + from litellm.proxy.spend_tracking.budget_reservation import get_reserved_counter_keys + if budget_reservation is not None: await _reconcile_budget_reservation_before_db_update( budget_reservation=budget_reservation, response_cost=response_cost ) + counter_keys: Final = frozenset( + get_reserved_counter_keys(budget_reservation=budget_reservation) + ) | post_call_counter_keys( + token=user_api_key, + team_id=team_id, + user_id=user_id, + org_id=org_id, + end_user_id=end_user_id, + tags=request_tags, + model_access_groups=model_access_groups, + project_id=project_id, + ) + with spend_counter_batch_scope(spend_counter_cache.redis_cache, counter_keys=counter_keys): + return await _update_database_and_spend_counters_in_batch( + proxy_logging_obj=proxy_logging_obj, + increment_spend_counters=increment_spend_counters, + user_api_key=user_api_key, + user_id=user_id, + end_user_id=end_user_id, + team_id=team_id, + org_id=org_id, + kwargs=kwargs, + completion_response=completion_response, + start_time=start_time, + end_time=end_time, + response_cost=response_cost, + budget_reservation=budget_reservation, + request_tags=request_tags, + model_access_groups=model_access_groups, + project_id=project_id, + ) + + +async def _update_database_and_spend_counters_in_batch( + proxy_logging_obj: "ProxyLogging", + increment_spend_counters: _IncrementSpendCounters, + user_api_key: str | None, + user_id: str | None, + end_user_id: str | None, + team_id: str | None, + org_id: str | None, + kwargs: dict, + completion_response: object, + start_time: datetime | None, + end_time: datetime | None, + response_cost: float, + budget_reservation: dict | None, + request_tags: list[str] | None, + model_access_groups: Sequence[str] | None, + project_id: str | None, +) -> bool: try: charged: Final = await proxy_logging_obj.db_spend_update_writer.update_database( token=user_api_key, @@ -762,11 +820,13 @@ async def _reconcile_budget_reservation_before_db_update( budget_reservation: dict, # mutable-ok: reconcile_budget_reservation stamps applied_adjustment on the caller's shared reservation dict response_cost: float, ) -> None: + """Reseeds the reserved counters that were flushed since reservation; the adjustments themselves are written by ``increment_spend_counters`` in the same pipeline as its increments, or by + the release / invalidation that runs when the spend write fails.""" from litellm.proxy.spend_tracking.budget_reservation import reconcile_budget_reservation try: - await reconcile_budget_reservation( - budget_reservation=budget_reservation, actual_cost=response_cost, finalize=False + _ = await reconcile_budget_reservation( + budget_reservation=budget_reservation, actual_cost=response_cost, finalize=False, apply_consistent=False ) except Exception: # noqa: BLE001 # a failed reconcile must not block the spend write; the counters are dropped instead verbose_proxy_logger.warning( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 6a221f2bfed..b4a497ea1e4 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -759,6 +759,7 @@ from litellm.proxy.shutdown.scheduled_jobs import ( from litellm.proxy.spend_tracking.budget_reservation import ( get_budget_window_start, release_unbound_budget_reservation, + stamp_budget_reservation_actual_cost, ) from litellm.proxy.spend_tracking.spend_capture_rate import ( run_scheduled_spend_capture_rate_check, @@ -2846,13 +2847,16 @@ async def _repair_stale_spend_counter(counter_key: str, db_spend: float) -> None if spend_counter_cache.redis_cache is not None: forget_spend_counter(counter_key) try: - await spend_counter_cache.redis_cache.async_set_max(key=counter_key, value=db_spend) + repaired: Final = await spend_counter_cache.redis_cache.async_set_max(key=counter_key, value=db_spend) except Exception: verbose_proxy_logger.debug( "Unable to repair stale spend counter %s in Redis", counter_key, exc_info=True, ) + return + if repaired is not None: + record_spend_counter_value(counter_key, repaired) async def reseed_spend_counter_from_db(counter_key: str) -> bool: @@ -3049,13 +3053,17 @@ async def _increment_spend_counters_batched( model_access_groups: Sequence[str] | None, project_id: str | None = None, ): - """Runs inside one spend counter batch: the reservation reconcile and the warm checks share a single MGET.""" - reserved_counter_keys: Final = await _reconcile_budget_reservation_for_counter_update( + """Runs inside one spend counter batch: the reservation reconcile and the warm checks share a single MGET, and + the reconcile adjustments go out in the same INCRBYFLOAT pipeline as the counter increments.""" + reservation_update: Final = await _reconcile_budget_reservation_for_counter_update( budget_reservation=budget_reservation, response_cost=response_cost, ) + reserved_counter_keys: Final = reservation_update.reserved_counter_keys if response_cost is None or response_cost == 0: + await _apply_spend_counter_increments(pending=reservation_update.pending) + stamp_budget_reservation_actual_cost(budget_reservation=budget_reservation, actual_cost=response_cost) if budget_reservation is not None: budget_reservation["finalized"] = True return @@ -3276,7 +3284,8 @@ async def _increment_spend_counters_batched( for item in scope if not isinstance(item, BaseException) ) - await _apply_spend_counter_increments(pending=pending) + await _apply_spend_counter_increments(pending=reservation_update.pending + pending) + stamp_budget_reservation_actual_cost(budget_reservation=budget_reservation, actual_cost=response_cost) if scope_errors: raise scope_errors[0] @@ -3284,12 +3293,21 @@ async def _increment_spend_counters_batched( budget_reservation["finalized"] = True +@dataclass(frozen=True, slots=True) +class _ReservationCounterUpdate: + """The reserved counters the direct increment must skip, and the adjustments that settle them on the actual + cost, still to be written; both empty when the reservation could not be reconciled and was dropped.""" + + reserved_counter_keys: frozenset[str] = frozenset() + pending: tuple[PendingSpendIncrement, ...] = () + + async def _reconcile_budget_reservation_for_counter_update( budget_reservation: dict | None, response_cost: float | None, -) -> set[str]: +) -> _ReservationCounterUpdate: if budget_reservation is None or budget_reservation.get("finalized") is True: - return set() + return _ReservationCounterUpdate() from litellm.proxy.spend_tracking.budget_reservation import ( get_reserved_counter_keys, @@ -3299,10 +3317,11 @@ async def _reconcile_budget_reservation_for_counter_update( reserved_counter_keys: Final = get_reserved_counter_keys(budget_reservation=budget_reservation) try: - await reconcile_budget_reservation( + pending: Final = await reconcile_budget_reservation( budget_reservation=budget_reservation, actual_cost=response_cost or 0.0, finalize=False, + apply_consistent=False, ) except Exception: verbose_proxy_logger.warning( @@ -3315,8 +3334,8 @@ async def _reconcile_budget_reservation_for_counter_update( verbose_proxy_logger.exception( "Failed to invalidate reserved counters after reservation reconciliation failed" ) - return set() - return reserved_counter_keys + return _ReservationCounterUpdate() + return _ReservationCounterUpdate(reserved_counter_keys=frozenset(reserved_counter_keys), pending=pending) async def _prepare_end_user_and_tag_spend_increments( @@ -3702,31 +3721,81 @@ async def _apply_spend_counter_increments(pending: Sequence[PendingSpendIncremen raise -async def increment_spend_counters_pipeline(pending: Sequence[PendingSpendIncrement]) -> None: - """One INCRBYFLOAT+EXPIRE pipeline for every pending counter; on failure every counter is invalidated - before the error propagates, so no caller can read a half-applied batch.""" +async def increment_spend_counters_pipeline(pending: Sequence[PendingSpendIncrement]) -> tuple[float | None, ...]: + """One INCRBYFLOAT+EXPIRE pipeline for every pending counter, returning each counter's new value in order; on + failure every counter is invalidated before the error propagates, so no caller can read a half-applied batch.""" + if spend_counter_cache.redis_cache is None: + return await run_spend_counter_pipeline(pending=pending) + try: + return await run_spend_counter_pipeline(pending=pending) + except Exception: + await asyncio.gather(*(_invalidate_spend_counter(counter_key=item.counter_key) for item in pending)) + raise + + +async def run_spend_counter_pipeline(pending: Sequence[PendingSpendIncrement]) -> tuple[float | None, ...]: + """The pipeline behind ``increment_spend_counters_pipeline`` without its invalidation: the caller decides what + happens to counters whose increment may or may not have landed when the pipeline fails.""" if not pending: - return + return () redis_cache: Final = spend_counter_cache.redis_cache if redis_cache is None: - for item in pending: - await SpendCounterReseed.increment_in_memory( - spend_counter_cache=spend_counter_cache, counter_key=item.counter_key, increment=item.increment - ) - return + return tuple( + [ + await SpendCounterReseed.increment_in_memory( + spend_counter_cache=spend_counter_cache, counter_key=item.counter_key, increment=item.increment + ) + for item in pending + ] + ) ttl: Final = redis_cache.get_ttl() increment_list: Final = [ # mutable-ok: async_increment_pipeline signature requires list[RedisPipelineIncrementOperation] RedisPipelineIncrementOperation(key=item.counter_key, increment_value=item.increment, ttl=ttl) for item in pending ] - try: - results: Final = await redis_cache.async_increment_pipeline(increment_list=increment_list) - except Exception: - await asyncio.gather(*(_invalidate_spend_counter(counter_key=item.counter_key) for item in pending)) - raise + results: Final = await redis_cache.async_increment_pipeline(increment_list=increment_list) for item, current_value in zip(pending, results or ()): spend_counter_cache.in_memory_cache.set_cache(key=item.counter_key, value=current_value) record_spend_counter_value(item.counter_key, float(current_value)) + return tuple(float(current_value) for current_value in results or ()) + + +def _update_cache_read_keys( + user_id: str | None, + end_user_id: str | None, + team_id: str | None, + tags: Sequence[object] | None, + response_cost: float | None, +) -> tuple[str, ...]: + if response_cost is None: + return () + user_keys: tuple[str, ...] = (user_id, GLOBAL_PROXY_SPEND_CACHE_KEY) if user_id is not None else () + end_user_keys: tuple[str, ...] = (end_user_cache_key(end_user_id),) if end_user_id is not None else () + team_keys: tuple[str, ...] = (f"team_id:{team_id}",) if team_id is not None else () + tag_keys: tuple[str, ...] = tuple(tag_cache_key(tag) for tag in tags or () if isinstance(tag, str) and tag) + return user_keys + end_user_keys + team_keys + tag_keys + + +async def _read_update_cache_values(keys: Sequence[str], parent_otel_span: Span | None) -> Mapping[str, object]: + """One batched read for every object ``update_cache`` refreshes; a failed read leaves them all untouched, + exactly as a failed per-object GET left that object untouched.""" + if not keys: + return MappingProxyType({}) + try: + values: Final = await user_api_key_cache.async_batch_get_cache( + keys=list(keys), parent_otel_span=parent_otel_span, throttle_redis=False + ) + except Exception as e: + verbose_proxy_logger.warning( + "Spend tracking - failed to read cached spend objects. Budget enforcement may use stale spend values. " + "keys=%s - %s", + keys, + str(e), + ) + return MappingProxyType({}) + if values is None: + return MappingProxyType({}) + return MappingProxyType({key: value for key, value in zip(keys, values) if value is not None}) async def update_cache( @@ -3745,6 +3814,12 @@ async def update_cache( """ values_to_update_in_cache: Final[list[tuple[str, object]]] = [] + cached_values: Final = await _read_update_cache_values( + keys=_update_cache_read_keys( + user_id=user_id, end_user_id=end_user_id, team_id=team_id, tags=tags, response_cost=response_cost + ), + parent_otel_span=parent_otel_span, + ) ### UPDATE KEY SPEND ### async def _update_key_cache(token: str, response_cost: float): @@ -3810,7 +3885,7 @@ async def update_cache( # Fetch the existing cost for the given user if _id is None: continue - cached_user = await user_api_key_cache.async_get_cache(key=_id) + cached_user = cached_values.get(_id) if cached_user is None: # do nothing if there is no cache value return @@ -3833,11 +3908,11 @@ async def update_cache( ) ) ## UPDATE GLOBAL PROXY ## - global_proxy_spend: Final = await user_api_key_cache.async_get_cache(key=GLOBAL_PROXY_SPEND_CACHE_KEY) - if global_proxy_spend is None: + global_proxy_spend: Final = cached_values.get(GLOBAL_PROXY_SPEND_CACHE_KEY) + if not isinstance(global_proxy_spend, (int, float)): # do nothing if not in cache return - elif response_cost is not None and global_proxy_spend is not None: + elif response_cost is not None: increment: Final = global_proxy_spend + response_cost values_to_update_in_cache.append((GLOBAL_PROXY_SPEND_CACHE_KEY, increment)) except Exception as e: @@ -3859,7 +3934,7 @@ async def update_cache( _id: Final = end_user_cache_key(end_user_id) try: # Fetch the existing cost for the given user - cached_end_user: Final = await user_api_key_cache.async_get_cache(key=_id) + cached_end_user: Final = cached_values.get(_id) if cached_end_user is None: # if user does not exist in LiteLLM_UserTable, create a new user # do nothing if end-user not in api key cache @@ -3900,7 +3975,7 @@ async def update_cache( _id: Final = f"team_id:{team_id}" try: - cached_team: Final = await user_api_key_cache.async_get_cache(key=_id) + cached_team: Final = cached_values.get(_id) if cached_team is None: # do nothing if team not in api key cache return @@ -3950,7 +4025,7 @@ async def update_cache( cache_key = tag_cache_key(tag_name) # Fetch the existing tag object from cache - cached_tag = await user_api_key_cache.async_get_cache(key=cache_key) + cached_tag = cached_values.get(cache_key) if cached_tag is None: # do nothing if tag not in api key cache continue diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index e28fa2c06a4..c094e91c6c0 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -4,7 +4,7 @@ import asyncio import json import math import time -from collections.abc import Mapping, Sequence +from collections.abc import AsyncIterator, Mapping, Sequence from dataclasses import dataclass from datetime import datetime, timedelta, timezone from types import MappingProxyType @@ -290,7 +290,6 @@ async def reserve_budget_for_request( raw_body=raw_body, ) - current_spend_by_counter_key: Final[dict[str, float]] = {} reservation_cost = estimate_request_max_cost( request_body=request_body, route=route, @@ -306,46 +305,17 @@ async def reserve_budget_for_request( applied_entries: Final[list[dict[str, float | str]]] = [] try: with _counters_batch_scope(frozenset(counter.counter_key for counter in counters)): - for counter in counters: - entry = _counter_to_reservation_entry( - counter=counter, - reserved_cost=reservation_cost, - ) - applied_entries.append(entry) - try: - reserved_value = await _reserve_counter( - counter=counter, - reservation_cost=reservation_cost, - ) - except _CounterReservationUnavailable as exc: - if exc.touched_counter and not exc.counter_invalidated: - await _release_applied_entries_best_effort( - entries=[entry], - default_reserved_cost=reservation_cost, - ) - applied_entries.remove(entry) - if fail_closed_budget_enforcement: - _raise_reservation_unavailable(counter_key=counter.counter_key) - continue - - if reserved_value is not None: - current_spend = reserved_value - else: - cached_spend = current_spend_by_counter_key.get(counter.counter_key) - if cached_spend is None: - cached_spend = await _get_current_counter_value(counter=counter) - current_spend = cached_spend + reservation_cost - if current_spend > counter.max_budget: - reservation_cost = await _apply_over_budget_reservation_policy( - counter=counter, - valid_token=valid_token, - entry=entry, - applied_entries=applied_entries, - reservation_cost=reservation_cost, - current_spend=current_spend, - fail_closed_budget_enforcement=fail_closed_budget_enforcement, - ) - continue + reservable: Final = await _initialize_reservation_counters( + counters=counters, + fail_closed_budget_enforcement=fail_closed_budget_enforcement, + ) + reservation_cost = await _reserve_reservable_counters( + reservable=reservable, + valid_token=valid_token, + applied_entries=applied_entries, + reservation_cost=reservation_cost, + fail_closed_budget_enforcement=fail_closed_budget_enforcement, + ) except Exception: await _release_applied_entries_best_effort( entries=applied_entries, @@ -381,19 +351,39 @@ async def reconcile_budget_reservation( budget_reservation: dict | None, actual_cost: float | None, finalize: bool = True, -) -> None: + apply_consistent: bool = True, +) -> tuple[PendingSpendIncrement, ...]: + """Settle every reserved counter on ``actual_cost``. With ``apply_consistent`` False the adjustments for + counters that still hold the reservation are returned instead of written, so the caller can pipeline them with + its own increments and then call ``stamp_budget_reservation_actual_cost``.""" if not budget_reservation or budget_reservation.get("finalized") is True: - return + return () reserved_cost: Final = float(budget_reservation.get("reserved_cost") or 0.0) actual: Final = float(actual_cost or 0.0) - await _set_reserved_entries_actual_cost( + pending: Final = await _set_reserved_entries_actual_cost( entries=budget_reservation.get("entries") or [], actual_cost=actual, default_reserved_cost=reserved_cost, + apply_consistent=apply_consistent, ) if finalize: budget_reservation["finalized"] = True + return pending + + +def stamp_budget_reservation_actual_cost(budget_reservation: dict | None, actual_cost: float | None) -> None: + """Record that every reserved counter now holds ``actual_cost``, once the adjustments handed back by + ``reconcile_budget_reservation(apply_consistent=False)`` have been written.""" + if not budget_reservation: + return + reserved_cost: Final = float(budget_reservation.get("reserved_cost") or 0.0) + actual: Final = float(actual_cost or 0.0) + for entry in budget_reservation.get("entries") or []: + if "counter_key" in entry: + entry["applied_adjustment"] = actual - _get_entry_reserved_cost( + entry=entry, default_reserved_cost=reserved_cost + ) async def release_budget_reservation(budget_reservation: dict | None) -> None: @@ -917,18 +907,40 @@ def _coerce_window(window: object) -> Mapping[str, object]: return dumped if isinstance(dumped, Mapping) else {} -async def _reserve_counter( - counter: _BudgetCounter, - reservation_cost: float, -) -> float | None: +async def _initialize_reservation_counters( + counters: Sequence[_BudgetCounter], + fail_closed_budget_enforcement: bool, +) -> tuple[_BudgetCounter, ...]: + """The counters whose current value is loaded, in order; one that cannot be loaded is skipped (or rejects the + request under fail-closed enforcement) exactly as it was when each counter was reserved on its own.""" + return tuple([counter async for counter in _loaded_reservation_counters(counters, fail_closed_budget_enforcement)]) + + +async def _loaded_reservation_counters( + counters: Sequence[_BudgetCounter], fail_closed_budget_enforcement: bool +) -> AsyncIterator[_BudgetCounter]: + for counter in counters: + if await _reservation_counter_loaded(counter, fail_closed_budget_enforcement): + yield counter + + +async def _reservation_counter_loaded(counter: _BudgetCounter, fail_closed_budget_enforcement: bool) -> bool: + try: + await _initialize_reservation_counter(counter=counter) + except _CounterReservationUnavailable: + if fail_closed_budget_enforcement: + _raise_reservation_unavailable(counter_key=counter.counter_key) + return False + return True + + +async def _initialize_reservation_counter(counter: _BudgetCounter) -> None: from litellm.proxy.proxy_server import ( _ensure_spend_counter_initialized, _ensure_window_spend_counter_initialized, - _increment_spend_counter_cache, _invalidate_spend_counter, ) - attempted_increment = False try: if counter.source_cache_key is not None: await _ensure_spend_counter_initialized( @@ -949,13 +961,6 @@ async def _reserve_counter( counter.counter_key, ) raise _CounterReservationUnavailable - - attempted_increment = True - reserved_value: Final = await _increment_spend_counter_cache( - counter_key=counter.counter_key, - increment=reservation_cost, - ) - return float(reserved_value) if reserved_value is not None else None except _CounterReservationUnavailable: raise except Exception: @@ -964,20 +969,121 @@ async def _reserve_counter( counter.counter_key, exc_info=True, ) - counter_invalidated = False try: await _invalidate_spend_counter(counter_key=counter.counter_key) - counter_invalidated = True except Exception: verbose_proxy_logger.warning( "Failed to invalidate spend counter after budget reservation failure for %s", counter.counter_key, exc_info=True, ) - raise _CounterReservationUnavailable( - touched_counter=attempted_increment, - counter_invalidated=counter_invalidated, + raise _CounterReservationUnavailable + + +async def _reserve_reservable_counters( + reservable: Sequence[_BudgetCounter], + valid_token: UserAPIKeyAuth | None, + applied_entries: list[dict[str, float | str]], + reservation_cost: float, + fail_closed_budget_enforcement: bool, +) -> float: + """Charge the counters group by group (see ``_reservation_groups``), settling the over-budget policy on each + group before the next is charged, and hand back the reservation cost the policy left standing.""" + current_spend_by_counter_key: Final = { + counter.counter_key: await _get_current_counter_value(counter=counter) for counter in reservable + } + for group in _reservation_groups( + counters=reservable, + current_spend_by_counter_key=current_spend_by_counter_key, + reservation_cost=reservation_cost, + ): + charged_cost = reservation_cost + entries = tuple(_counter_to_reservation_entry(counter=counter, reserved_cost=charged_cost) for counter in group) + applied_entries.extend(entries) + reserved_values = await _reserve_counters(counters=group, entries=entries, reservation_cost=charged_cost) + if reserved_values is None: + for entry in entries: + applied_entries.remove(entry) + if fail_closed_budget_enforcement: + _raise_reservation_unavailable(counter_key=group[0].counter_key) + continue + for counter, entry, reserved_value in zip(group, entries, reserved_values): + if entry not in applied_entries: + continue + if reserved_value is not None: + current_spend = reserved_value - (charged_cost - reservation_cost) + else: + current_spend = current_spend_by_counter_key[counter.counter_key] + reservation_cost + if current_spend > counter.max_budget: + reservation_cost = await _apply_over_budget_reservation_policy( + counter=counter, + valid_token=valid_token, + entry=entry, + applied_entries=applied_entries, + reservation_cost=reservation_cost, + current_spend=current_spend, + fail_closed_budget_enforcement=fail_closed_budget_enforcement, + ) + return reservation_cost + + +def _reservation_groups( + counters: Sequence[_BudgetCounter], + current_spend_by_counter_key: Mapping[str, float], + reservation_cost: float, +) -> tuple[tuple[_BudgetCounter, ...], ...]: + """Every counter the batch read says still has room for the estimate is charged in one pipeline. As soon as one + does not, the counters are charged one at a time so the over-budget policy settles each before the next is + touched, and a rejection charges nothing after it.""" + if not counters: + return () + if all( + current_spend_by_counter_key[counter.counter_key] + reservation_cost <= counter.max_budget + for counter in counters + ): + return (tuple(counters),) + return tuple((counter,) for counter in counters) + + +async def _reserve_counters( + counters: Sequence[_BudgetCounter], + entries: Sequence[dict[str, float | str]], + reservation_cost: float, +) -> tuple[float | None, ...] | None: + """One INCRBYFLOAT pipeline reserves every counter. When it fails each counter is dropped, and one that cannot + be dropped is released instead in case its increment landed, so nothing is left to release by the caller.""" + from litellm.proxy.proxy_server import _invalidate_spend_counter, run_spend_counter_pipeline + + if not counters: + return () + try: + reserved: Final = await run_spend_counter_pipeline( + pending=tuple( + PendingSpendIncrement(counter_key=counter.counter_key, increment=reservation_cost) + for counter in counters + ) ) + except Exception: + verbose_proxy_logger.warning( + "Skipping budget reservation for %s because spend counter reservation failed", + tuple(counter.counter_key for counter in counters), + exc_info=True, + ) + for counter, entry in zip(counters, entries): + try: + await _invalidate_spend_counter(counter_key=counter.counter_key) + except Exception: + verbose_proxy_logger.warning( + "Failed to invalidate spend counter after budget reservation failure for %s", + counter.counter_key, + exc_info=True, + ) + await _release_applied_entries_best_effort( + entries=[entry], # mutable-ok: the release takes the reservation's list of entries + default_reserved_cost=reservation_cost, + ) + return None + return tuple(reserved) + (None,) * (len(counters) - len(reserved)) async def _get_current_counter_value(counter: _BudgetCounter) -> float: @@ -1026,9 +1132,11 @@ async def _set_reserved_entries_actual_cost( actual_cost: float, default_reserved_cost: float, reseed_on_inconsistent: bool = True, -) -> None: - """Every reserved counter is read from one MGET and the consistent adjustments go out in one pipeline. - A counter that was flushed or reseeded since reservation is settled on its own after the pipeline.""" + apply_consistent: bool = True, +) -> tuple[PendingSpendIncrement, ...]: + """Every reserved counter is read from one MGET and the consistent adjustments go out in one pipeline, or are + returned unwritten when ``apply_consistent`` is False. A counter that was flushed or reseeded since reservation + is settled on its own after the pipeline.""" from litellm.proxy.proxy_server import increment_spend_counters_pipeline with _counters_batch_scope(frozenset(str(entry["counter_key"]) for entry in entries if "counter_key" in entry)): @@ -1055,15 +1163,16 @@ async def _set_reserved_entries_actual_cost( f"Cannot resize budget reservation against inconsistent counter {inconsistent[0].counter_key}" ) applicable: Final = tuple(item for item, ok in zip(adjustments, consistent) if ok) - await increment_spend_counters_pipeline( - pending=tuple( - PendingSpendIncrement(counter_key=item.counter_key, increment=item.adjustment) for item in applicable - ) + applicable_pending: Final = tuple( + PendingSpendIncrement(counter_key=item.counter_key, increment=item.adjustment) for item in applicable ) + if apply_consistent: + await increment_spend_counters_pipeline(pending=applicable_pending) for item in inconsistent: await _reseed_reserved_entry(item=item, actual_cost=actual_cost) - for item in adjustments: + for item in adjustments if apply_consistent else inconsistent: item.entry["applied_adjustment"] = item.target_adjustment + return () if apply_consistent else applicable_pending async def _reseed_reserved_entry(item: _EntryAdjustment, actual_cost: float) -> None: diff --git a/litellm/proxy/spend_tracking/spend_counter_batch.py b/litellm/proxy/spend_tracking/spend_counter_batch.py index ddb074ae023..a6694895a27 100644 --- a/litellm/proxy/spend_tracking/spend_counter_batch.py +++ b/litellm/proxy/spend_tracking/spend_counter_batch.py @@ -144,25 +144,43 @@ def release_spend_counter_batch() -> None: batch.close() -def _iter_admission_counter_keys(token: UserAPIKeyAuth, end_user_id: str | None) -> Iterator[str]: - if token.token is not None: - yield f"spend:key:{token.token}" - if token.team_id is not None: - yield f"spend:team:{token.team_id}" - if token.user_id is not None: - yield f"spend:team_member:{token.user_id}:{token.team_id}" - if token.user_id is not None: - yield f"spend:user:{token.user_id}" - if end_user_id is not None: +def _iter_entity_counter_keys( + token: object, + team_id: object, + user_id: object, + org_id: object, + project_id: object, + end_user_id: object, +) -> Iterator[str]: + """Only string ids name a counter; anything else (None, or an unresolved placeholder in synthetic + logging payloads) simply has no counter to bind.""" + if isinstance(token, str): + yield f"spend:key:{token}" + if isinstance(team_id, str): + yield f"spend:team:{team_id}" + if isinstance(user_id, str): + yield f"spend:team_member:{user_id}:{team_id}" + if isinstance(user_id, str): + yield f"spend:user:{user_id}" + if isinstance(end_user_id, str): yield f"spend:end_user:{end_user_id}" - if token.org_id is not None: - yield f"spend:org:{token.org_id}" - if token.project_id is not None: - yield project_spend_counter_key(token.project_id) + if isinstance(org_id, str): + yield f"spend:org:{org_id}" + if isinstance(project_id, str): + yield project_spend_counter_key(project_id) def admission_counter_keys(token: UserAPIKeyAuth, end_user_id: str | None) -> frozenset[str]: - return frozenset(_iter_admission_counter_keys(token, end_user_id)) + return frozenset( + _iter_entity_counter_keys( + token=token.token, + team_id=token.team_id, + user_id=token.user_id, + org_id=token.org_id, + project_id=token.project_id, + end_user_id=end_user_id, + ) + ) def post_call_counter_keys( @@ -176,9 +194,15 @@ def post_call_counter_keys( project_id: str | None = None, ) -> frozenset[str]: """Every counter ``increment_spend_counters`` warm-checks, except budget windows which bind on read.""" - entity_keys: Final = admission_counter_keys( - UserAPIKeyAuth(token=token, team_id=team_id, user_id=user_id, org_id=org_id, project_id=project_id), - end_user_id, + entity_keys: Final = frozenset( + _iter_entity_counter_keys( + token=token, + team_id=team_id, + user_id=user_id, + org_id=org_id, + project_id=project_id, + end_user_id=end_user_id, + ) ) tag_keys: Final = frozenset(f"spend:tag:{tag}" for tag in tags or () if tag and isinstance(tag, str)) group_keys: Final = frozenset( diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index 470db99108a..78b281d5c78 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -9290,3 +9290,67 @@ async def test_websocket_auth_hands_the_reservation_to_the_socket_state(): assert result.budget_reservation == reservation assert websocket.state.budget_reservation is reservation assert websocket.scope["state"]["budget_reservation"] is reservation + + +@pytest.mark.asyncio +async def test_admission_and_budget_reservation_read_the_key_spend_counter_with_one_redis_mget(): + from fastapi import Request + from starlette.datastructures import URL + + import litellm.proxy.proxy_server as _proxy_server_mod + from litellm.proxy.spend_tracking.spend_counter_batch import ( + read_batched_spend_counter, + spend_counter_batch_scope, + ) + + token = UserAPIKeyAuth(api_key="sk-test", token="hashed", max_budget=10.0) + request = Request(scope={"type": "http"}) + request._url = URL(url="/chat/completions") + reads: list[tuple[str, tuple[float | None, bool] | None]] = [] + + async def _admission_reads_spend(**kwargs): + reads.append(("admission", await read_batched_spend_counter("spend:key:hashed"))) + + async def _reservation_reads_spend(**kwargs): + reads.append(("reservation", await read_batched_spend_counter("spend:key:hashed"))) + + redis = MagicMock() + redis.async_batch_get_cache = AsyncMock(return_value={"spend:key:hashed": 4.0}) + attrs = { + **_proxy_attrs_for_centralized_checks(user_custom_auth=None), + "prisma_client": MagicMock(), + "spend_counter_cache": MagicMock(redis_cache=redis), + } + originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs} + try: + for k, v in attrs.items(): + setattr(_proxy_server_mod, k, v) + with ( + patch( # test-quality-ok: authorization has its own tests above; this one checks the shared counter read + "litellm.proxy.auth.user_api_key_auth.common_checks", + new=AsyncMock(side_effect=_admission_reads_spend), + ), + patch( # test-quality-ok: the reservation helper imports reserve_budget_for_request in its body + "litellm.proxy.spend_tracking.budget_reservation.reserve_budget_for_request", + side_effect=_reservation_reads_spend, + ), + spend_counter_batch_scope(redis), + ): + await _run_centralized_common_checks( + user_api_key_auth_obj=token, + request=request, + request_data={"model": "gpt-5.4-mini", "messages": [{"role": "user", "content": "hi"}]}, + route="/chat/completions", + ) + reads.append(("after admission", await read_batched_spend_counter("spend:key:hashed"))) + finally: + for k, v in attrs.items(): + setattr(_proxy_server_mod, k, originals[k]) + + assert reads == [ + ("admission", (4.0, True)), + ("reservation", (4.0, True)), + ("after admission", None), + ], "admission and reservation share one snapshot, and read-then-write callers go to Redis once it closes" + assert redis.async_batch_get_cache.await_count == 1 + assert "spend:key:hashed" in redis.async_batch_get_cache.await_args.kwargs["key_list"] diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index b5e594db701..a5b2d8b0b8d 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -726,6 +726,7 @@ async def test_update_database_and_spend_counters_reconciles_reservation_before_ budget_reservation=budget_reservation, actual_cost=0.2, finalize=False, + apply_consistent=False, ) increment_spend_counters.assert_awaited_once() assert increment_spend_counters.await_args.kwargs["budget_reservation"] is budget_reservation @@ -771,6 +772,7 @@ async def test_update_database_and_spend_counters_releases_reservation_when_db_u budget_reservation=budget_reservation, actual_cost=0.2, finalize=False, + apply_consistent=False, ) mock_release_budget_reservation.assert_awaited_once_with( budget_reservation=budget_reservation, diff --git a/tests/test_litellm/proxy/proxy_server/test_spend_counters.py b/tests/test_litellm/proxy/proxy_server/test_spend_counters.py index 0731c233fef..ad86c3c5267 100644 --- a/tests/test_litellm/proxy/proxy_server/test_spend_counters.py +++ b/tests/test_litellm/proxy/proxy_server/test_spend_counters.py @@ -72,6 +72,7 @@ def _make_spend_counter_cache( def _make_user_api_key_cache(get_value=None, get_side_effect=None): cache = MagicMock() cache.async_get_cache = AsyncMock(return_value=get_value, side_effect=get_side_effect) + cache.async_batch_get_cache = AsyncMock(side_effect=lambda keys, **_: [get_value for _ in keys]) cache.async_set_cache_pipeline = AsyncMock() return cache @@ -633,7 +634,7 @@ async def test_increment_spend_counters_skips_reserved_counter_keys(monkeypatch) reserved = {"spend:key:hashed-tok", "spend:org:org1"} monkeypatch.setattr(br, "get_reserved_counter_keys", MagicMock(return_value=set(reserved))) - monkeypatch.setattr(br, "reconcile_budget_reservation", AsyncMock()) + monkeypatch.setattr(br, "reconcile_budget_reservation", AsyncMock(return_value=())) recorded: dict[str, float] = {} @@ -888,7 +889,8 @@ async def test_increment_spend_counters_pipeline_failure_invalidates_all_counter @pytest.mark.asyncio async def test_reconcile_budget_reservation_for_counter_update_returns_empty_set_when_none(): result = await ps._reconcile_budget_reservation_for_counter_update(budget_reservation=None, response_cost=1.0) - assert result == set() + assert result.reserved_counter_keys == frozenset() + assert result.pending == () @pytest.mark.asyncio @@ -917,7 +919,8 @@ async def test_reconcile_budget_reservation_for_counter_update_failure_invalidat budget_reservation={"foo": "bar"}, response_cost=1.0 ) - assert result == set() + assert result.reserved_counter_keys == frozenset() + assert result.pending == () assert fake_invalidate.called is True @@ -941,7 +944,8 @@ async def test_reconcile_budget_reservation_for_counter_update_finalized_reserva response_cost=1.0, ) - assert result == set() + assert result.reserved_counter_keys == frozenset() + assert result.pending == () fake_reconcile.assert_not_awaited() @@ -1531,16 +1535,15 @@ async def test_update_cache_no_cached_entities_schedules_pipeline_flush(monkeypa tags=["x"], ) - observed = { - "lookups": fake_user_cache.async_get_cache.call_count, - "got_user": True, - "got_team": True, - } - assert normalize(observed) == { - "lookups": 4, - "got_user": True, - "got_team": True, - } + assert fake_user_cache.async_get_cache.await_count == 0 + fake_user_cache.async_batch_get_cache.assert_awaited_once() + assert fake_user_cache.async_batch_get_cache.await_args.kwargs["keys"] == [ + "u1", + f"{ps.litellm_proxy_admin_name}:spend", + "end_user_id:eu1", + "team_id:t1", + "tag:x", + ] @pytest.mark.asyncio @@ -1548,7 +1551,7 @@ async def test_update_cache_user_cache_failure_invalid_state_is_swallowed(monkey """An inner _update_user_cache raising must not propagate — update_cache catches and logs, the public coroutine still completes normally.""" fake_user_cache = MagicMock() - fake_user_cache.async_get_cache = AsyncMock(side_effect=RuntimeError("cache down")) + fake_user_cache.async_batch_get_cache = AsyncMock(side_effect=RuntimeError("cache down")) fake_user_cache.async_set_cache_pipeline = AsyncMock() monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache) diff --git a/tests/test_litellm/proxy/spend_tracking/test_budget_reservation_redis_failure.py b/tests/test_litellm/proxy/spend_tracking/test_budget_reservation_redis_failure.py index 6165af4920d..e0a74d50a6c 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_budget_reservation_redis_failure.py +++ b/tests/test_litellm/proxy/spend_tracking/test_budget_reservation_redis_failure.py @@ -9,8 +9,10 @@ gives up, but ``increment_spend_counters`` still treats the counter as lands in the enforced counter, so budgets stop gating until the next cold reseed pulls a lagging value from the DB. -The fix makes the reconcile path fall back to the direct increment when it -fails, so the actual cost is always written to the shared counter. +The reconcile adjustment and the direct increment now leave in one pipeline, so +a failure either writes the actual cost or drops the counter (and surfaces the +error) for the next read to reseed from the DB; it never leaves the reserved +estimate in place as if it were reconciled. """ import pytest @@ -84,13 +86,14 @@ async def test_direct_increment_runs_when_reservation_reconcile_hits_redis_failu ], } - await proxy_server.increment_spend_counters( - token=hashed_token, - team_id=None, - user_id=None, - response_cost=response_cost, - budget_reservation=budget_reservation, - ) + with pytest.raises(Exception, match="Redis timeout"): + await proxy_server.increment_spend_counters( + token=hashed_token, + team_id=None, + user_id=None, + response_cost=response_cost, + budget_reservation=budget_reservation, + ) - enforced_spend = await flaky_redis.async_get_cache(key=counter_key) - assert enforced_spend == response_cost + assert await flaky_redis.async_get_cache(key=counter_key) is None + assert proxy_server.spend_counter_cache.in_memory_cache.get_cache(key=counter_key) is None diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_counter_batch.py b/tests/test_litellm/proxy/spend_tracking/test_spend_counter_batch.py index 3e4b817fab8..1fddfaaa766 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_counter_batch.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_counter_batch.py @@ -8,6 +8,7 @@ from unittest.mock import AsyncMock, MagicMock import pytest +import litellm import litellm.proxy.proxy_server as ps from litellm.caching.redis_cache import RedisCache from litellm.proxy._types import UserAPIKeyAuth @@ -17,6 +18,7 @@ from litellm.proxy.spend_tracking.spend_counter_batch import ( active_spend_counter_batch, admission_counter_keys, bind_admission_counter_keys, + post_call_counter_keys, release_spend_counter_batch, spend_counter_batch_scope, ) @@ -86,6 +88,21 @@ def test_admission_counter_keys_cover_every_entity_the_checks_read(): ) +def test_post_call_counter_keys_skip_ids_that_are_not_strings(): + """A synthetic logging payload (batch cost polling, tests) can carry placeholders where the ids belong; those + have no counter, and deriving the key set must never raise inside the cost callback.""" + placeholder = object() + assert post_call_counter_keys( + token=placeholder, # pyright: ignore[reportArgumentType] # synthetic payload placeholder, not an id + team_id="team", + user_id=None, + org_id=placeholder, # pyright: ignore[reportArgumentType] # synthetic payload placeholder, not an id + end_user_id="eu", + tags=[placeholder, "t1"], + model_access_groups=None, + ) == {"spend:team:team", "spend:end_user:eu", "spend:tag:t1"} + + @pytest.mark.asyncio async def test_bound_counters_share_one_mget_and_a_clean_miss_is_authoritative(): redis = CountingRedis({"spend:key:hashed": 1.5, "spend:team:team": 2.5}) @@ -407,7 +424,7 @@ def _reservation(reserved_cost: float, counter_keys: frozenset[str] = RESERVED_K @pytest.mark.asyncio -async def test_post_call_with_a_reservation_costs_one_mget_one_reconcile_pipeline_one_increment_pipeline(monkeypatch): +async def test_post_call_with_a_reservation_costs_one_mget_and_one_pipeline_for_reconcile_and_increments(monkeypatch): redis = CountingRedis({key: 1.0 for key in POST_CALL_KEYS}) monkeypatch.setattr(ps, "spend_counter_cache", _spend_counter_cache(redis)) monkeypatch.setattr(ps, "prisma_client", None) @@ -425,10 +442,9 @@ async def test_post_call_with_a_reservation_costs_one_mget_one_reconcile_pipelin budget_reservation=reservation, ) - assert [c.split()[0] for c in redis.commands] == ["MGET", "PIPELINE", "PIPELINE"], redis.commands + assert [c.split()[0] for c in redis.commands] == ["MGET", "PIPELINE"], redis.commands assert set(redis.commands[0].split()[1:]) == POST_CALL_KEYS, "reconcile and warm checks share the MGET" - assert set(redis.commands[1].split()[1:]) == RESERVED_KEYS - assert set(redis.commands[2].split()[1:]) == POST_CALL_KEYS - RESERVED_KEYS + assert set(redis.commands[1].split()[1:]) == POST_CALL_KEYS, "reconcile adjustments ride the increment pipeline" assert {key: round(redis.store[key], 6) for key in POST_CALL_KEYS} == { key: (1.1 if key in RESERVED_KEYS else 1.5) for key in POST_CALL_KEYS } @@ -436,6 +452,27 @@ async def test_post_call_with_a_reservation_costs_one_mget_one_reconcile_pipelin assert reservation["finalized"] is True +@pytest.mark.asyncio +async def test_a_stale_counter_repair_updates_the_open_batch_instead_of_forcing_a_second_mget(monkeypatch): + redis = CountingRedis({"spend:key:hashed": 1.0, "spend:team:team": 1.0}) + + async def set_max(key: str, value: float, **kwargs: object) -> float: + redis.commands.append(f"SETMAX {key} {value}") + redis.store[key] = max(float(str(redis.store.get(key, 0.0))), value) + return float(str(redis.store[key])) + + redis.async_set_max = set_max + monkeypatch.setattr(ps, "spend_counter_cache", _spend_counter_cache(redis)) + + with spend_counter_batch_scope(redis, counter_keys=frozenset({"spend:key:hashed", "spend:team:team"})): + assert await ps.read_spend_counter_cache_value(counter_key="spend:team:team") == (1.0, True) + await ps._repair_stale_spend_counter(counter_key="spend:team:team", db_spend=4.0) + assert await ps.read_spend_counter_cache_value(counter_key="spend:team:team") == (4.0, True) + assert await ps.read_spend_counter_cache_value(counter_key="spend:key:hashed") == (1.0, True) + + assert [c.split()[0] for c in redis.commands] == ["MGET", "SETMAX"], redis.commands + + @pytest.mark.asyncio async def test_reconcile_settles_a_flushed_counter_on_its_own_after_the_shared_pipeline(monkeypatch): from litellm.proxy.spend_tracking.budget_reservation import reconcile_budget_reservation @@ -478,37 +515,35 @@ async def test_pre_call_resize_against_an_inconsistent_counter_writes_nothing_an @pytest.mark.asyncio -async def test_a_failed_reconcile_pipeline_invalidates_every_reserved_counter_and_falls_back(monkeypatch): +async def test_a_failed_post_call_pipeline_invalidates_every_counter_it_carried_and_stamps_nothing(monkeypatch): redis = CountingRedis({key: 1.0 for key in POST_CALL_KEYS}) redis.async_delete_cache = AsyncMock() - reconcile_pipeline_failed = False async def _pipeline(increment_list: Sequence[Mapping[str, object]], **kwargs: object) -> list[float]: - nonlocal reconcile_pipeline_failed - if not reconcile_pipeline_failed: - reconcile_pipeline_failed = True - raise ConnectionError("redis down") - return await CountingRedis.async_increment_pipeline(redis, increment_list, **kwargs) + raise ConnectionError("redis down") redis.async_increment_pipeline = _pipeline # pyright: ignore[reportAttributeAccessIssue] # instance override monkeypatch.setattr(ps, "spend_counter_cache", _spend_counter_cache(redis)) monkeypatch.setattr(ps, "prisma_client", None) reservation = _reservation(reserved_cost=0.4) - await ps.increment_spend_counters( - token="hashed", - team_id="team", - user_id="user", - org_id="org", - end_user_id="eu", - response_cost=0.5, - budget_reservation=reservation, - ) + with pytest.raises(ConnectionError): + await ps.increment_spend_counters( + token="hashed", + team_id="team", + user_id="user", + org_id="org", + end_user_id="eu", + response_cost=0.5, + budget_reservation=reservation, + ) - assert {call.kwargs["key"] for call in redis.async_delete_cache.await_args_list} == RESERVED_KEYS + assert [c.split()[0] for c in redis.commands] == ["MGET"], redis.commands + assert {call.kwargs["key"] for call in redis.async_delete_cache.await_args_list} == RESERVED_KEYS | { + "spend:user:user" + } assert all("applied_adjustment" not in entry for entry in reservation["entries"]) - assert redis.commands[-1].split()[0] == "PIPELINE" - assert set(redis.commands[-1].split()[1:]) == RESERVED_KEYS | {"spend:user:user"} + assert {key: redis.store[key] for key in POST_CALL_KEYS} == {key: 1.0 for key in POST_CALL_KEYS} def test_a_scope_opened_inside_an_open_scope_joins_its_batch_and_a_closed_one_gets_its_own(): @@ -525,3 +560,199 @@ def test_a_scope_opened_inside_an_open_scope_joins_its_batch_and_a_closed_one_ge assert inner is not outer assert inner is not None and inner.counter_keys == {"spend:key:c"} assert active_spend_counter_batch() is outer + + +@pytest.mark.asyncio +async def test_reservation_inside_the_admission_scope_reuses_its_mget_and_reserves_in_one_pipeline(monkeypatch): + from litellm.caching.dual_cache import DualCache + from litellm.proxy._types import LiteLLM_TeamTableCachedObj + from litellm.proxy.spend_tracking.budget_reservation import reserve_budget_for_request + + redis = CountingRedis({"spend:key:hashed": 1.0, "spend:team:team": 2.0}) + redis.default_ttl = 3600 + monkeypatch.setattr(ps, "spend_counter_cache", _spend_counter_cache(redis)) + monkeypatch.setattr(ps, "prisma_client", None) + monkeypatch.setattr("litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", lambda **_: 0.5) + token = UserAPIKeyAuth(token="hashed", team_id="team", max_budget=10.0) + + with spend_counter_batch_scope(redis, counter_keys=admission_counter_keys(token, end_user_id=None)): + reservation = await reserve_budget_for_request( + request_body={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hi"}]}, + route="/chat/completions", + llm_router=None, + valid_token=token, + team_object=LiteLLM_TeamTableCachedObj(team_id="team", max_budget=20.0), + user_object=None, + prisma_client=None, + user_api_key_cache=DualCache(), + proxy_logging_obj=MagicMock(), + ) + + assert reservation is not None + assert [c.split()[0] for c in redis.commands] == ["MGET", "PIPELINE"], redis.commands + assert set(redis.commands[0].split()[1:]) == {"spend:key:hashed", "spend:team:team"} + assert redis.commands[1] == "PIPELINE spend:key:hashed spend:team:team" + assert redis.store == {"spend:key:hashed": 1.5, "spend:team:team": 2.5} + assert [entry["counter_key"] for entry in reservation["entries"]] == ["spend:key:hashed", "spend:team:team"] + + +@pytest.mark.asyncio +async def test_a_failed_reservation_pipeline_drops_every_counter_and_reserves_nothing(monkeypatch): + from litellm.caching.dual_cache import DualCache + from litellm.proxy._types import LiteLLM_TeamTableCachedObj + from litellm.proxy.spend_tracking.budget_reservation import reserve_budget_for_request + + redis = CountingRedis({"spend:key:hashed": 1.0, "spend:team:team": 2.0}) + redis.default_ttl = 3600 + redis.async_delete_cache = AsyncMock() + + async def _pipeline(increment_list: Sequence[Mapping[str, object]], **kwargs: object) -> list[float]: + raise ConnectionError("redis down") + + redis.async_increment_pipeline = _pipeline # pyright: ignore[reportAttributeAccessIssue] # instance override + monkeypatch.setattr(ps, "spend_counter_cache", _spend_counter_cache(redis)) + monkeypatch.setattr(ps, "prisma_client", None) + monkeypatch.setattr("litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", lambda **_: 0.5) + token = UserAPIKeyAuth(token="hashed", team_id="team", max_budget=10.0) + + reservation = await reserve_budget_for_request( + request_body={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hi"}]}, + route="/chat/completions", + llm_router=None, + valid_token=token, + team_object=LiteLLM_TeamTableCachedObj(team_id="team", max_budget=20.0), + user_object=None, + prisma_client=None, + user_api_key_cache=DualCache(), + proxy_logging_obj=MagicMock(), + ) + + assert reservation is None + assert {call.kwargs["key"] for call in redis.async_delete_cache.await_args_list} == { + "spend:key:hashed", + "spend:team:team", + } + assert redis.store == {"spend:key:hashed": 1.0, "spend:team:team": 2.0} + + +@pytest.mark.asyncio +async def test_post_call_lifecycle_reads_the_counters_after_the_db_update_and_writes_one_pipeline(monkeypatch): + from litellm.proxy.hooks.proxy_track_cost_callback import _update_database_and_spend_counters + + redis = CountingRedis({key: 1.0 for key in POST_CALL_KEYS}) + monkeypatch.setattr(ps, "spend_counter_cache", _spend_counter_cache(redis)) + monkeypatch.setattr(ps, "prisma_client", None) + proxy_logging_obj = MagicMock() + + async def _update_database(**kwargs: object) -> bool: + redis.commands.append("DB") + return True + + proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock(side_effect=_update_database) + reservation = _reservation(reserved_cost=0.4) + + charged = await _update_database_and_spend_counters( + proxy_logging_obj=proxy_logging_obj, + increment_spend_counters=ps.increment_spend_counters, + user_api_key="hashed", + user_id="user", + end_user_id="eu", + team_id="team", + org_id="org", + kwargs={}, + completion_response=None, + start_time=None, + end_time=None, + response_cost=0.5, + budget_reservation=reservation, + request_tags=["prod"], + model_access_groups=["premium"], + ) + + assert charged is True + proxy_logging_obj.db_spend_update_writer.update_database.assert_awaited_once() + assert [c.split()[0] for c in redis.commands] == ["MGET", "DB", "MGET", "PIPELINE"], redis.commands + assert set(redis.commands[0].split()[1:]) == RESERVED_KEYS + assert set(redis.commands[2].split()[1:]) == POST_CALL_KEYS + assert set(redis.commands[3].split()[1:]) == POST_CALL_KEYS + assert {key: round(redis.store[key], 6) for key in POST_CALL_KEYS} == { + key: (1.1 if key in RESERVED_KEYS else 1.5) for key in POST_CALL_KEYS + } + assert [round(entry["applied_adjustment"], 6) for entry in reservation["entries"]] == [0.1] * len(RESERVED_KEYS) + assert reservation["finalized"] is True + assert active_spend_counter_batch() is None + + +def _reservation_fixture(monkeypatch, redis: CountingRedis) -> None: + redis.default_ttl = 3600 + monkeypatch.setattr(ps, "spend_counter_cache", _spend_counter_cache(redis)) + monkeypatch.setattr(ps, "prisma_client", None) + monkeypatch.setattr("litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", lambda **_: 0.5) + + +async def _reserve(redis: CountingRedis, token: UserAPIKeyAuth, team_max_budget: float) -> dict | None: + from litellm.caching.dual_cache import DualCache + from litellm.proxy._types import LiteLLM_TeamTableCachedObj + from litellm.proxy.spend_tracking.budget_reservation import reserve_budget_for_request + + with spend_counter_batch_scope(redis, counter_keys=admission_counter_keys(token, end_user_id=None)): + return await reserve_budget_for_request( + request_body={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hi"}]}, + route="/chat/completions", + llm_router=None, + valid_token=token, + team_object=LiteLLM_TeamTableCachedObj(team_id="team", max_budget=team_max_budget), + user_object=None, + prisma_client=None, + user_api_key_cache=DualCache(), + proxy_logging_obj=MagicMock(), + ) + + +@pytest.mark.asyncio +async def test_a_rejected_counter_is_charged_alone_so_the_counters_after_it_are_never_touched(monkeypatch): + """Only counters the admission MGET says still fit the estimate share the reservation pipeline; a counter that + does not is charged on its own first, so its rejection never inflates a sibling counter, not even briefly.""" + redis = CountingRedis({"spend:key:hashed": 10.0, "spend:team:team": 2.0}) + _reservation_fixture(monkeypatch, redis) + token = UserAPIKeyAuth(token="hashed", team_id="team", max_budget=10.0) + + with pytest.raises(litellm.BudgetExceededError): + await _reserve(redis, token, team_max_budget=20.0) + + writes = [c for c in redis.commands if not c.startswith("MGET")] + assert writes and all("spend:team:team" not in c for c in writes), redis.commands + assert redis.store == {"spend:key:hashed": 10.0, "spend:team:team": 2.0} + + +@pytest.mark.asyncio +async def test_a_resized_reservation_is_carried_at_its_resized_cost_to_the_counters_charged_after_it(monkeypatch): + redis = CountingRedis({"spend:key:hashed": 9.8, "spend:team:team": 2.0}) + _reservation_fixture(monkeypatch, redis) + token = UserAPIKeyAuth(token="hashed", team_id="team", max_budget=10.0) + + reservation = await _reserve(redis, token, team_max_budget=2.1) + + assert reservation is not None + assert reservation["reserved_cost"] == pytest.approx(0.1) + assert redis.store["spend:key:hashed"] == pytest.approx(9.9) + assert redis.store["spend:team:team"] == pytest.approx(2.1) + + +@pytest.mark.asyncio +async def test_update_cache_reads_an_object_redis_gained_right_after_a_batch_read_missed_it(monkeypatch): + """DualCache throttles repeated batch reads of a key that just missed; the per-object GET update_cache used to + issue never did, so its batched read must not either.""" + from litellm.caching.dual_cache import DualCache + + redis = CountingRedis() + cache = DualCache(redis_cache=redis) + monkeypatch.setattr(ps, "user_api_key_cache", cache) + assert await cache.async_batch_get_cache(keys=["team_id:team"]) == [None] + redis.store["team_id:team"] = {"spend": 1.0} + assert await cache.async_batch_get_cache(keys=["team_id:team"]) == [None] + + assert await ps._read_update_cache_values(keys=["team_id:team"], parent_otel_span=None) == { + "team_id:team": {"spend": 1.0} + } + assert redis.commands.count("MGET team_id:team") == 2 diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index 18b046cd83c..c8e4df1030f 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -1880,8 +1880,8 @@ async def test_should_raise_503_when_counter_increment_fails_and_fail_closed( async def test_fail_closed_releases_earlier_counters_before_503( spend_counter_state, ): - """#33923: when a later counter's reservation write fails in strict mode, the - counters that already reserved must be released before the 503 propagates.""" + """#33923: when a later counter cannot be loaded in strict mode, the 503 is raised before any counter is + reserved.""" counter_cache, key_cache = spend_counter_state proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) valid_token = UserAPIKeyAuth( @@ -1915,12 +1915,8 @@ async def test_fail_closed_releases_earlier_counters_before_503( ) assert exc_info.value.status_code == 503 - assert ( - counter_cache.in_memory_cache.get_cache( - key="spend:key:key-budget-fail-closed-release" - ) - == 0.0 - ) + assert counter_cache.in_memory_cache.get_cache(key="spend:key:key-budget-fail-closed-release") is None + assert counter_cache.in_memory_cache.get_cache(key="spend:key:key-budget-fail-closed-release:window:1h") is None @pytest.mark.asyncio @@ -1982,21 +1978,10 @@ async def test_should_release_tracked_entry_when_reservation_fails_after_increme max_budget=1.0, ) - import litellm.proxy.proxy_server as ps - - original_increment_counter = ps._increment_spend_counter_cache - first_increment = True - - async def fail_after_increment(counter_key: str, increment: float): - nonlocal first_increment - if first_increment: - first_increment = False - await counter_cache.async_increment_cache(key=counter_key, value=increment) - raise RuntimeError("lost increment response") - return await original_increment_counter( - counter_key=counter_key, - increment=increment, - ) + async def fail_after_increment(pending): + for item in pending: + await counter_cache.async_increment_cache(key=item.counter_key, value=item.increment) + raise RuntimeError("lost increment response") with ( patch( @@ -2004,7 +1989,7 @@ async def test_should_release_tracked_entry_when_reservation_fails_after_increme return_value=0.5, ), patch( - "litellm.proxy.proxy_server._increment_spend_counter_cache", + "litellm.proxy.proxy_server.run_spend_counter_pipeline", side_effect=fail_after_increment, ), patch( @@ -2596,6 +2581,72 @@ async def test_reconcile_before_db_update_does_not_double_count_when_flush_lands assert reservation["finalized"] is True +class _BatchReadingRedisCache(_ExpiringRedisCache): + async def async_batch_get_cache(self, key_list: Sequence[str], **kwargs: object) -> dict[str, float | None]: + return {key: await self.async_get_cache(key) for key in key_list} + + +@pytest.mark.asyncio +async def test_reserved_counter_deleted_during_spend_write_is_reseeded_instead_of_going_negative( + spend_counter_state, +): + import litellm.proxy.proxy_server as ps + from litellm.proxy.hooks.proxy_track_cost_callback import _update_database_and_spend_counters + + counter_cache, _ = spend_counter_state + counter_key = "spend:key:key-deleted-mid-write" + redis_cache = _BatchReadingRedisCache() + counter_cache.redis_cache = redis_cache + await redis_cache.async_set_cache(counter_key, 0.6) + counter_cache.in_memory_cache.set_cache(key=counter_key, value=0.6) + + async def _delete_counter_while_persisting(**kwargs: object) -> bool: + await redis_cache.async_delete_cache(counter_key) + counter_cache.in_memory_cache.delete_cache(key=counter_key) + return True + + proxy_logging_obj = MagicMock() + proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock(side_effect=_delete_counter_while_persisting) + reservation = { + "reserved_cost": 0.6, + "entries": [ + { + "counter_key": counter_key, + "entity_type": "Key", + "entity_id": "key-deleted-mid-write", + "reserved_cost": 0.6, + "applied_adjustment": 0.0, + } + ], + "finalized": False, + } + + with ( + patch.object( # test-quality-ok: the reseed reads the DB floor through a Prisma client the test has no seam for + ps.SpendCounterReseed, "from_db", AsyncMock(return_value=0.3) + ) + ): + charged = await _update_database_and_spend_counters( + proxy_logging_obj=proxy_logging_obj, + increment_spend_counters=ps.increment_spend_counters, + user_api_key="key-deleted-mid-write", + user_id=None, + end_user_id=None, + team_id=None, + org_id=None, + kwargs={}, + completion_response=None, + start_time=datetime.now(), + end_time=datetime.now(), + response_cost=0.05, + budget_reservation=reservation, + ) + + assert charged is True + assert redis_cache.store[counter_key] == pytest.approx(0.35), redis_cache.store + assert reservation["finalized"] is True + + @pytest.mark.asyncio async def test_should_invalidate_reserved_counters_after_persisted_spend_failure( spend_counter_state, diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index df05cf0987e..23d319159ca 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -6189,7 +6189,7 @@ async def test_tag_cache_update_called(): "spend": 10.0, } - with patch.object(cache, "async_get_cache", new=AsyncMock(return_value=mock_tag_obj)) as mock_get_cache: + with patch.object(cache, "async_batch_get_cache", new=AsyncMock(return_value=[mock_tag_obj])) as mock_get_cache: with patch.object(cache, "async_set_cache_pipeline", new=AsyncMock()) as mock_set_cache: await litellm.proxy.proxy_server.update_cache( token=None, @@ -6203,7 +6203,7 @@ async def test_tag_cache_update_called(): await asyncio.sleep(0.1) - mock_get_cache.assert_awaited_once_with(key="tag:test-tag") + mock_get_cache.assert_awaited_once_with(keys=["tag:test-tag"], parent_otel_span=None, throttle_redis=False) mock_set_cache.assert_awaited_once() call_args = mock_set_cache.call_args @@ -6234,15 +6234,11 @@ async def test_tag_cache_update_multiple_tags(): mock_tag1_obj = {"tag_name": "tag1", "spend": 10.0} mock_tag2_obj = {"tag_name": "tag2", "spend": 20.0} - async def mock_get_cache_side_effect(key): - if key == "tag:tag1": - return mock_tag1_obj - elif key == "tag:tag2": - return mock_tag2_obj - return None + async def mock_get_cache_side_effect(keys, **kwargs): + return [{"tag:tag1": mock_tag1_obj, "tag:tag2": mock_tag2_obj}.get(key) for key in keys] with patch.object( - cache, "async_get_cache", new=AsyncMock(side_effect=mock_get_cache_side_effect) + cache, "async_batch_get_cache", new=AsyncMock(side_effect=mock_get_cache_side_effect) ) as mock_get_cache: with patch.object(cache, "async_set_cache_pipeline", new=AsyncMock()) as mock_set_cache: await litellm.proxy.proxy_server.update_cache( @@ -6257,7 +6253,7 @@ async def test_tag_cache_update_multiple_tags(): await asyncio.sleep(0.1) - assert mock_get_cache.call_count == 2 + mock_get_cache.assert_awaited_once_with(keys=["tag:tag1", "tag:tag2"], parent_otel_span=None, throttle_redis=False) mock_set_cache.assert_awaited_once() call_args = mock_set_cache.call_args @@ -6288,8 +6284,8 @@ async def test_update_cache_pipeline_honors_user_api_key_cache_ttl(): try: with patch.object( cache, - "async_get_cache", - new=AsyncMock(return_value={"tag_name": "active-tag", "spend": 1.0}), + "async_batch_get_cache", + new=AsyncMock(return_value=[{"tag_name": "active-tag", "spend": 1.0}]), ): with patch.object(cache, "async_set_cache_pipeline", new=AsyncMock()) as mock_set_cache: await litellm.proxy.proxy_server.update_cache( @@ -6376,18 +6372,21 @@ async def test_update_cache_global_proxy_spend_scalar_stays_shared(): admin_name = litellm.proxy.proxy_server.litellm_proxy_admin_name global_key = "{}:spend".format(admin_name) - async def fake_get(key, **kwargs): + def fake_get(key): if key == "user-lit": return {"user_id": "user-lit", "spend": 1.0} if key == global_key: return 10.0 return None + async def fake_batch_get(keys, **kwargs): + return [fake_get(key) for key in keys] + original_cache = litellm.proxy.proxy_server.user_api_key_cache cache = DualCache(default_in_memory_ttl=300) setattr(litellm.proxy.proxy_server, "user_api_key_cache", cache) try: - with patch.object(cache, "async_get_cache", new=AsyncMock(side_effect=fake_get)): + with patch.object(cache, "async_batch_get_cache", new=AsyncMock(side_effect=fake_batch_get)): with patch.object(cache, "async_set_cache_pipeline", new=AsyncMock()) as mock_set_cache: await litellm.proxy.proxy_server.update_cache( token=None, @@ -13868,7 +13867,7 @@ async def test_window_spend_row_is_enqueued_even_when_the_counter_was_reserved() } original_reconcile = br.reconcile_budget_reservation - br.reconcile_budget_reservation = AsyncMock(return_value=None) + br.reconcile_budget_reservation = AsyncMock(return_value=()) try: with _window_spend_enqueue_env({"hashed-token": key_obj}) as queue: await increment_spend_counters( diff --git a/tests/unit/proxy/auth/test_jwt.py b/tests/unit/proxy/auth/test_jwt.py index 6ad253f33e8..fd1d8974b48 100644 --- a/tests/unit/proxy/auth/test_jwt.py +++ b/tests/unit/proxy/auth/test_jwt.py @@ -874,8 +874,7 @@ async def test_team_cache_update_called(): cache, ) - with patch.object(cache, "async_get_cache", new=AsyncMock()) as mock_call_cache: - cache.async_get_cache = mock_call_cache + with patch.object(cache, "async_batch_get_cache", new=AsyncMock(return_value=[None])) as mock_call_cache: # Call the function under test await litellm.proxy.proxy_server.update_cache( token=None, @@ -887,7 +886,7 @@ async def test_team_cache_update_called(): ) # type: ignore await asyncio.sleep(3) - mock_call_cache.assert_awaited_once() + mock_call_cache.assert_awaited_once_with(keys=["team_id:1234"], parent_otel_span=None, throttle_redis=False) @pytest.fixture From 24a7e8973870fe05ffc7fb4fb4a8cb646a88a487 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 15:09:22 -0700 Subject: [PATCH 34/41] perf(responses): run aresponses through the async wrapper so the cache is read once (#43769) * perf(responses): run aresponses through the async wrapper so the cache is read once aresponses() sets kwargs["aresponses"] = True and runs the decorated sync responses() on an executor, but _is_async_request() did not recognise that flag, so the sync wrapper did a second cache lookup on the executor thread with a differently ordered cache-key input. Every /v1/responses request paid two cache GETs against two different keys. Recognising aresponses in _is_async_request() leaves the async wrapper as the only cache reader and writer for the async path, one GET per request, same key on read and write Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(caching): let a responses cache entry cover aresponses so responses-only configs keep caching Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/caching/caching_handler.py | 19 ++--- litellm/utils.py | 1 + tests/unit/test_utils.py | 115 +++++++++++++++++++++++++++++ 3 files changed, 123 insertions(+), 12 deletions(-) diff --git a/litellm/caching/caching_handler.py b/litellm/caching/caching_handler.py index 0e4f444224b..1b4f446ee0c 100644 --- a/litellm/caching/caching_handler.py +++ b/litellm/caching/caching_handler.py @@ -1129,11 +1129,8 @@ class LLMCachingHandler: Returns: bool: True if the result should be stored in the cache, False otherwise. """ - return ( - (litellm.cache is not None) - and litellm.cache.supported_call_types is not None - and (str(original_function.__name__) in litellm.cache.supported_call_types) - and (kwargs.get("cache", {}).get("no-store", False) is not True) + return self._is_call_type_supported_by_cache(original_function=original_function) and ( + kwargs.get("cache", {}).get("no-store", False) is not True ) def wrap_streaming_result_for_cache( @@ -1170,13 +1167,11 @@ class LLMCachingHandler: Returns: bool: True if the call type is supported by the cache, False otherwise. """ - if ( - litellm.cache is not None - and litellm.cache.supported_call_types is not None - and str(original_function.__name__) in litellm.cache.supported_call_types - ): - return True - return False + if litellm.cache is None or litellm.cache.supported_call_types is None: + return False + call_type: Final = str(original_function.__name__) + covering_call_types: Final = ("aresponses", "responses") if call_type == "aresponses" else (call_type,) + return any(name in litellm.cache.supported_call_types for name in covering_call_types) async def _add_streaming_response_to_cache(self, processed_chunk: ModelResponse): """ diff --git a/litellm/utils.py b/litellm/utils.py index ffd507fad45..eeccd27c1d8 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -2309,6 +2309,7 @@ def _is_async_request( or kwargs.get("_arealtime", False) is True or kwargs.get("acreate_batch", False) is True or kwargs.get("acreate_fine_tuning_job", False) is True + or kwargs.get("aresponses", False) is True or is_pass_through is True ): return True diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index ab7ae5ab3b8..afc449d22c4 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -31,6 +31,7 @@ from litellm._logging import ( ) from litellm.caching.caching import Cache from litellm.caching.caching_handler import _PENDING_CACHE_WRITES +from litellm.caching.in_memory_cache import InMemoryCache from litellm.constants import DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger @@ -38,6 +39,7 @@ from litellm.litellm_core_utils.get_litellm_params import get_litellm_params from litellm.litellm_core_utils.thread_pool_executor import executor as logging_executor from litellm.llms.base_llm.base_model_iterator import MockResponseIterator from litellm.proxy.utils import is_valid_api_key +from litellm.types.caching import CachingSupportedCallTypes from litellm.types.integrations.custom_logger import HEADROOM_CONVERTED_STREAM_KEY from litellm.types.llms.openai import ResponsesAPIResponse from litellm.types.router import CredentialLiteLLMParams, GenericLiteLLMParams @@ -4744,6 +4746,119 @@ async def test_wrapper_async_replays_cached_converted_responses_stream_as_stream _assert_cache_hit_logged_as_stream(capture, await _wait_for_success_kwargs(capture, count=2)) +class _ReadCountingInMemoryCache(InMemoryCache): + def __init__(self) -> None: + super().__init__() + self.reads = 0 + + def get_cache(self, key: str, **kwargs: object) -> object: + self.reads += 1 + return super().get_cache(key, **kwargs) + + +_NATIVE_RESPONSES_BODY: Final = { + "id": "resp_native_replay", + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-5.6", + "output": [ + { + "type": "message", + "id": "msg_native_replay", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "native body", "annotations": []}], + } + ], + "usage": {"input_tokens": 3, "output_tokens": 4, "total_tokens": 7}, +} + + +def _native_responses_route(stream: bool) -> respx.Route: + if not stream: + return respx.post("https://api.openai.com/v1/responses").respond(json=_NATIVE_RESPONSES_BODY) + sse_body: Final = "".join( + f"event: {event_type}\ndata: {json.dumps({'type': event_type, 'response': _NATIVE_RESPONSES_BODY})}\n\n" + for event_type in ("response.created", "response.completed") + ) + return respx.post("https://api.openai.com/v1/responses").respond( + text=sse_body, headers={"content-type": "text/event-stream"} + ) + + +async def _drain_responses_result(result: object) -> None: + from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator + + if isinstance(result, BaseResponsesAPIStreamingIterator): + assert [event async for event in result][-1].type == "response.completed" + return + assert isinstance(result, ResponsesAPIResponse) + + +async def _wait_for_success_kwargs_with_input( + capture: _SuccessKwargsCapture, input_text: str, count: int +) -> dict[str, object]: + expected_messages: Final = [{"role": "user", "content": input_text}] + + def _logged_messages(kwargs: dict[str, object]) -> object: + standard_logging_object: Final = kwargs.get("standard_logging_object") + return standard_logging_object.get("messages") if isinstance(standard_logging_object, dict) else None + + def _matching() -> tuple[dict[str, object], ...]: + return tuple(kwargs for kwargs in capture.success_kwargs if _logged_messages(kwargs) == expected_messages) + + for _ in range(50): + if len(_matching()) >= count and not _PENDING_CACHE_WRITES: + break + await asyncio.sleep(0.05) + await asyncio.sleep(0.2) + matching: Final = _matching() + assert len(matching) == count + return matching[-1] + + +@pytest.mark.asyncio +@respx.mock +@pytest.mark.parametrize("stream", [False, True], ids=["non_stream", "stream"]) +@pytest.mark.parametrize( + "supported_call_types", + [["aresponses", "responses"], ["responses"]], + ids=["both_call_types", "responses_only"], +) +async def test_wrapper_aresponses_reads_cache_once_and_replays_from_that_read( + monkeypatch: pytest.MonkeyPatch, stream: bool, supported_call_types: list[CachingSupportedCallTypes] +) -> None: + capture: Final = _install_converted_stream_callbacks(monkeypatch) + monkeypatch.setattr(litellm, "callbacks", [capture]) + counting: Final = _ReadCountingInMemoryCache() + monkeypatch.setattr( + litellm, "cache", Cache(type="local", _backend=counting, supported_call_types=supported_call_types) + ) + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + route: Final = _native_responses_route(stream) + request: Final = { + "model": "openai/gpt-5.6", + "input": "read me once", + "stream": stream, + "api_key": "sk-test", + "num_retries": 0, + } + + await _drain_responses_result(await litellm.aresponses(**request)) + await _wait_for_success_kwargs_with_input(capture, request["input"], count=1) + assert counting.reads == 1, "aresponses must look the response cache up once, not again on the executor thread" + + await _drain_responses_result(await litellm.aresponses(**request)) + assert counting.reads == 2 + assert route.call_count == 1, "the single async cache read must hit the key the first call stored" + success_kwargs: Final = await _wait_for_success_kwargs_with_input(capture, request["input"], count=2) + standard_logging_object: Final = success_kwargs["standard_logging_object"] + assert isinstance(standard_logging_object, dict) + assert standard_logging_object["cache_hit"] is True + + def test_function_setup_failure_after_logging_construction_restores_context(monkeypatch): """If function_setup() constructs Logging() (which already mutated trace_id_var/session_id_var in __init__) but then raises before returning, From 92c0d6f5c83efc52e9cbe3c89068d3668e202e68 Mon Sep 17 00:00:00 2001 From: tin-berri Date: Tue, 29 Sep 2026 15:17:14 -0700 Subject: [PATCH 35/41] fix(router): bind Claude Code background sessions to their auto-router (#43767) Claude Code background sessions (claude --bg) stamp x-app: cli-bg on every request, including main-loop turns. The session router binding only accepted x-app: cli, so a background session never bound and its subagents' concrete-model calls bypassed the router. The binding write already requires the requested model to resolve to a pre-routing strategy, so background side calls naming plain models still never bind. Since f6eff1bde0 removed the clear path, the x-app check guarded nothing else. Co-authored-by: Claude Opus 5.5 --- litellm/router.py | 2 -- tests/unit/test_router/test_router.py | 7 ++++--- 2 files changed, 4 insertions(+), 5 deletions(-) diff --git a/litellm/router.py b/litellm/router.py index a98631b7f97..c18dfea1a36 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -13716,8 +13716,6 @@ class Router: self._stamp_or_clear_metadata_key(request_kwargs, "model_group", bound_model) return bound_registered_model - if self._request_header(request_kwargs, "x-app") != "cli": - return registered_model_name if self._select_pre_routing_strategy(registered_model_name, request_kwargs) is None: return registered_model_name await self._claude_code_session_router_cache.async_set_cache( diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index d4f9924dd13..0aedfce3598 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -11515,15 +11515,16 @@ class TestClaudeCodeSubagentSessionRouterBinding: } @pytest.mark.asyncio - async def test_subagent_concrete_model_uses_the_main_sessions_router(self): + @pytest.mark.parametrize("app", ["cli", "cli-bg"]) + async def test_subagent_concrete_model_uses_the_main_sessions_router(self, app): router = self._router() await router.acompletion( model="smart-router", messages=[{"role": "user", "content": "main turn"}], - **self._request_kwargs(), + **self._request_kwargs(app=app), ) - subagent_kwargs = self._request_kwargs(agent_id="agent-1234") + subagent_kwargs = self._request_kwargs(app=app, agent_id="agent-1234") response = await router.acompletion( model="expensive-model", From 56a63b4b296015511a2fdfeb576337548378892b Mon Sep 17 00:00:00 2001 From: "berriai-litellm-provider-info-sync[bot]" <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 22:29:45 +0000 Subject: [PATCH 36/41] feat(bedrock): add openai.gpt-6.1-sol us geo cris and Mantle rows (#43763) Co-authored-by: kerry-berri --- ...odel_prices_and_context_window_backup.json | 73 +++++++++++++++++++ model_prices_and_context_window.json | 73 +++++++++++++++++++ 2 files changed, 146 insertions(+) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 8836c7b2b16..99471d48f56 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -79253,5 +79253,78 @@ "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true + }, + "bedrock_mantle/openai.gpt-6.1-sol": { + "cache_creation_input_token_cost": 2.75e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5.5e-06, + "cache_read_input_token_cost": 1.1e-07, + "cache_read_input_token_cost_above_272k_tokens": 2.2e-07, + "input_cost_per_token": 2.2e-06, + "input_cost_per_token_above_272k_tokens": 4.4e-06, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "responses", + "output_cost_per_token": 1.1e-05, + "output_cost_per_token_above_272k_tokens": 1.65e-05, + "source": "https://developers.openai.com/api/docs/models/gpt-6.1-sol", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "use_openai_responses_path": true + }, + "us.openai.gpt-6.1-sol": { + "cache_creation_input_token_cost": 2.75e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5.5e-06, + "cache_read_input_token_cost": 1.1e-07, + "cache_read_input_token_cost_above_272k_tokens": 2.2e-07, + "input_cost_per_token": 2.2e-06, + "input_cost_per_token_above_272k_tokens": 4.4e-06, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1.1e-05, + "output_cost_per_token_above_272k_tokens": 1.65e-05, + "source": "https://developers.openai.com/api/docs/pricing", + "supported_endpoints": [ + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true } } diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 8836c7b2b16..99471d48f56 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -79253,5 +79253,78 @@ "supports_tool_choice": true, "supports_vision": true, "supports_xhigh_reasoning_effort": true + }, + "bedrock_mantle/openai.gpt-6.1-sol": { + "cache_creation_input_token_cost": 2.75e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5.5e-06, + "cache_read_input_token_cost": 1.1e-07, + "cache_read_input_token_cost_above_272k_tokens": 2.2e-07, + "input_cost_per_token": 2.2e-06, + "input_cost_per_token_above_272k_tokens": 4.4e-06, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "responses", + "output_cost_per_token": 1.1e-05, + "output_cost_per_token_above_272k_tokens": 1.65e-05, + "source": "https://developers.openai.com/api/docs/models/gpt-6.1-sol", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "use_openai_responses_path": true + }, + "us.openai.gpt-6.1-sol": { + "cache_creation_input_token_cost": 2.75e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5.5e-06, + "cache_read_input_token_cost": 1.1e-07, + "cache_read_input_token_cost_above_272k_tokens": 2.2e-07, + "input_cost_per_token": 2.2e-06, + "input_cost_per_token_above_272k_tokens": 4.4e-06, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1.1e-05, + "output_cost_per_token_above_272k_tokens": 1.65e-05, + "source": "https://developers.openai.com/api/docs/pricing", + "supported_endpoints": [ + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true } } From f5a1c9f1f1bf75b02cd3d0e59e6affb6797f51f9 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 15:36:09 -0700 Subject: [PATCH 37/41] fix(proxy): recover session key owners from daily spend for usage attribution (#43642) * fix(proxy): recover daily spend key owners Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): simplify daily spend owner recovery Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * style(proxy): format daily activity metadata Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): cover recovered owner metadata merge Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): bound the daily spend owner lookup with the statement timeout * test(integration): audit the daily activity key owner fallback on every usage route Thirty five integration cells under tests/integration/spend cover the daily spend owner fallback on all nine daily activity routes and /usage/ai/chat: the happy path per route, the unanimity rules (two users, blank and null rows, an owner the user table lacks, live and deleted keys with and without their own user, a spend log alias), a non admin reader, an invalid key, a 5 KB key, a locked LiteLLM_DailyUserSpend, 300 keys of one team, repeated reads, a second user landing between reads, a concurrent burst across the unified endpoints, a killed worker, and a proxy restart The traffic cells ignore the GET /v1/models call the proxy's five minute token limit refresh makes to every registered OpenAI compatible deployment, since it lands on a test's provider wire whenever the refresh instant falls inside the test --------- Co-authored-by: jesus Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- .../common_daily_activity.py | 30 +- .../spend_tracking/key_metadata_recovery.py | 62 ++- tests/integration/_support/daily_activity.py | 237 +++++++++++ .../spend/test_daily_activity_key_owner.py | 196 +++++++++ .../test_daily_activity_key_owner_faults.py | 264 ++++++++++++ .../test_daily_activity_key_owner_traffic.py | 397 ++++++++++++++++++ .../test_common_daily_activity.py | 131 +++++- .../test_key_metadata_recovery.py | 90 ++++ .../src/components/UsagePage/types.ts | 1 + .../src/components/activity_metrics.test.tsx | 22 + .../src/components/activity_metrics.tsx | 33 +- 11 files changed, 1426 insertions(+), 37 deletions(-) create mode 100644 tests/integration/_support/daily_activity.py create mode 100644 tests/integration/spend/test_daily_activity_key_owner.py create mode 100644 tests/integration/spend/test_daily_activity_key_owner_faults.py create mode 100644 tests/integration/spend/test_daily_activity_key_owner_traffic.py diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index cecf3e50f5c..c2a0a41c3e2 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -16,6 +16,7 @@ from litellm.proxy.spend_tracking.key_metadata_recovery import ( recover_cli_session_key_metadata, recover_double_hashed_key_metadata, recover_key_metadata_from_spend_logs, + recover_key_owner_from_daily_spend, ) from litellm.proxy.spend_tracking.ptu_feature_flag import is_ptu_cost_attribution_enabled from litellm.proxy.utils import PrismaClient @@ -468,6 +469,17 @@ def _parse_spend_date(raw: str | None) -> datetime | None: _EMPTY_KEY_METADATA: Final[Mapping[str, _KeyMetadataDict]] = MappingProxyType({}) +def _metadata_with_recovered_owner( + metadata: Mapping[str, _KeyMetadataDict], + key: str, + owner: str, +) -> _KeyMetadataDict: + current: Final = metadata.get(key) + if current is None: + return {"user_id": owner} + return {**current, "user_id": owner} + + async def get_api_key_metadata( prisma_client: PrismaClient, api_keys: AbstractSet[str], @@ -530,7 +542,19 @@ async def get_api_key_metadata( else _EMPTY_KEY_METADATA ) combined: Final = MappingProxyType({**after_token_recovery, **from_spend_logs}) - return await attach_user_details(prisma_client, combined) + ownerless: Final = frozenset( + key + for key in api_keys + if not combined.get(key, {}).get("user_id") and not combined.get(key, {}).get("key_exists") + ) + owners: Final = await recover_key_owner_from_daily_spend(prisma_client, ownerless) + metadata_with_owners: Final[Mapping[str, _KeyMetadataDict]] = MappingProxyType( + { + **combined, + **{key: _metadata_with_recovered_owner(combined, key, owner) for key, owner in owners.items()}, + } + ) + return await attach_user_details(prisma_client, metadata_with_owners) def _adjust_dates_for_timezone( @@ -944,7 +968,7 @@ async def _aggregate_spend_records( record.api_key for record in records if record.api_key and record.api_key != PTU_SENTINEL_API_KEY } - api_key_metadata: dict[str, _KeyMetadataDict] = {} + api_key_metadata: Mapping[str, _KeyMetadataDict] = MappingProxyType({}) if api_keys: api_key_metadata = await get_api_key_metadata( prisma_client, api_keys, _spend_logs_window(frozenset(record.date for record in records)) @@ -1144,7 +1168,7 @@ async def _aggregate_grouping_sets_records( """Async wrapper: fetch api_key_metadata, then dispatch on a worker thread.""" api_keys: Final[set[str]] = {r.api_key for r in records if r.api_key and r.api_key != PTU_SENTINEL_API_KEY} - api_key_metadata: dict[str, _KeyMetadataDict] = {} + api_key_metadata: Mapping[str, _KeyMetadataDict] = MappingProxyType({}) if api_keys: api_key_metadata = await get_api_key_metadata( prisma_client, api_keys, _spend_logs_window(frozenset(r.date for r in records)) diff --git a/litellm/proxy/spend_tracking/key_metadata_recovery.py b/litellm/proxy/spend_tracking/key_metadata_recovery.py index ce96dc62780..0e3e0598d17 100644 --- a/litellm/proxy/spend_tracking/key_metadata_recovery.py +++ b/litellm/proxy/spend_tracking/key_metadata_recovery.py @@ -61,6 +61,13 @@ WHERE COALESCE(key_alias, user_id, team_id) IS NOT NULL GROUP BY api_key """ +_DAILY_USER_SPEND_OWNER_SQL: Final = """ +SELECT api_key, MIN(user_id) AS first_owner, MAX(user_id) AS last_owner +FROM "LiteLLM_DailyUserSpend" +WHERE api_key = ANY($1::text[]) AND user_id IS NOT NULL AND user_id <> '' +GROUP BY api_key +""" + _SPEND_LOG_STATEMENT_TIMEOUT_SQL: Final = f"SET LOCAL statement_timeout = {SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS}" _SPEND_LOG_TRANSACTION_TIMEOUT: Final = timedelta(milliseconds=2 * SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS) @@ -104,8 +111,15 @@ class _SpendLogDigestRow(BaseModel): ) +class _DailyUserSpendOwnerRow(BaseModel): + api_key: str + first_owner: str | None = None + last_owner: str | None = None + + _TOKEN_DIGEST_ROWS: Final = TypeAdapter(tuple[_TokenDigestRow, ...]) _SPEND_LOG_DIGEST_ROWS: Final = TypeAdapter(tuple[_SpendLogDigestRow, ...]) +_DAILY_USER_SPEND_OWNER_ROWS: Final = TypeAdapter(tuple[_DailyUserSpendOwnerRow, ...]) _CACHED_KEY_METADATA: Final = TypeAdapter(KeyMetadataDict) _SPEND_LOG_METADATA_CACHE: Final = InMemoryCache( max_size_in_memory=SPEND_LOG_KEY_METADATA_CACHE_MAX_ITEMS, @@ -113,6 +127,7 @@ _SPEND_LOG_METADATA_CACHE: Final = InMemoryCache( ) _SPEND_LOG_QUERY_LOCK: Final = asyncio.Lock() _EMPTY_KEY_METADATA: Final[Mapping[str, KeyMetadataDict]] = MappingProxyType({}) +_EMPTY_KEY_OWNERS: Final[Mapping[str, str]] = MappingProxyType({}) async def _db_or_empty( @@ -129,6 +144,16 @@ async def _db_or_empty( return None +async def _rows_within_the_statement_timeout( + prisma_client: PrismaClient, + sql: str, + *params: object, +) -> Sequence[Mapping[str, object]]: + async with prisma_client.db.tx(timeout=_SPEND_LOG_TRANSACTION_TIMEOUT) as transaction: + await transaction.execute_raw(_SPEND_LOG_STATEMENT_TIMEOUT_SQL) + return await transaction.query_raw(sql, *params) + + async def _reverse_hash_key_metadata( prisma_client: PrismaClient, sql: str, @@ -152,6 +177,29 @@ async def _reverse_hash_key_metadata( ) +async def recover_key_owner_from_daily_spend( + prisma_client: PrismaClient, + keys: AbstractSet[str], +) -> Mapping[str, str]: + if not keys: + return _EMPTY_KEY_OWNERS + rows: Final = await _db_or_empty( + lambda: _rows_within_the_statement_timeout(prisma_client, _DAILY_USER_SPEND_OWNER_SQL, sorted(keys)), + "Failed daily-spend key owner recovery for %d keys: %s", + len(keys), + ) + if rows is None: + return _EMPTY_KEY_OWNERS + return MappingProxyType( + { + row.api_key: owner + for row in _DAILY_USER_SPEND_OWNER_ROWS.validate_python(rows) + for owner in (_unanimous(row.first_owner, row.last_owner),) + if row.api_key in keys and owner is not None + } + ) + + @dataclass(frozen=True, slots=True) class _UserDetails: email: str | None @@ -309,24 +357,14 @@ def _cached_spend_log_metadata( ) -async def _spend_log_rows_within_the_statement_timeout( - prisma_client: PrismaClient, - digests: AbstractSet[str], - window: tuple[datetime, datetime], -) -> Sequence[Mapping[str, object]]: - start, end = window - async with prisma_client.db.tx(timeout=_SPEND_LOG_TRANSACTION_TIMEOUT) as transaction: - await transaction.execute_raw(_SPEND_LOG_STATEMENT_TIMEOUT_SQL) - return await transaction.query_raw(_SPEND_LOG_ALIAS_SQL, sorted(digests), start, end) - - async def _query_spend_log_metadata( prisma_client: PrismaClient, digests: AbstractSet[str], window: tuple[datetime, datetime], ) -> Mapping[str, KeyMetadataDict] | None: + start, end = window rows: Final = await _db_or_empty( - lambda: _spend_log_rows_within_the_statement_timeout(prisma_client, digests, window), + lambda: _rows_within_the_statement_timeout(prisma_client, _SPEND_LOG_ALIAS_SQL, sorted(digests), start, end), "Failed spend-log alias recovery for %d missing keys: %s", len(digests), ) diff --git a/tests/integration/_support/daily_activity.py b/tests/integration/_support/daily_activity.py new file mode 100644 index 00000000000..debb8c4cdb4 --- /dev/null +++ b/tests/integration/_support/daily_activity.py @@ -0,0 +1,237 @@ +import os +import uuid +from collections.abc import Iterator, Mapping, Sequence +from contextlib import contextmanager +from dataclasses import dataclass +from itertools import chain +from typing import Final + +import httpx +import psycopg +import pytest +from integration._support.client import Gateway, Scenario, object_value +from psycopg import sql +from psycopg.types.json import Jsonb +from pydantic import JsonValue + +USER_SPEND: Final = "LiteLLM_DailyUserSpend" +TEAM_SPEND: Final = "LiteLLM_DailyTeamSpend" +TAG_SPEND: Final = "LiteLLM_DailyTagSpend" +ORGANIZATION_SPEND: Final = "LiteLLM_DailyOrganizationSpend" +END_USER_SPEND: Final = "LiteLLM_DailyEndUserSpend" +AGENT_SPEND: Final = "LiteLLM_DailyAgentSpend" +DAY: Final = "2026-02-03" +AGGREGATED_USER_ACTIVITY: Final = "/user/daily/activity/aggregated" + +INSERT_DAILY_ROW: Final = sql.SQL( + "INSERT INTO {table} (id, {entity}, date, api_key, model, model_group, custom_llm_provider, prompt_tokens," + " completion_tokens, spend, api_requests, successful_requests, failed_requests, updated_at)" + " VALUES (gen_random_uuid()::text, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, now())" +) +DELETE_DAILY_ROWS: Final = sql.SQL("DELETE FROM {table} WHERE api_key = ANY(%s)") +INSERT_SPEND_LOG: Final = ( + 'INSERT INTO "LiteLLM_SpendLogs" (request_id, call_type, api_key, "startTime", "endTime", metadata)' + " VALUES (%s, 'acompletion', %s, %s::timestamp, %s::timestamp, %s)" +) +DELETE_SPEND_LOG: Final = 'DELETE FROM "LiteLLM_SpendLogs" WHERE request_id = %s' +LOCK_TABLE: Final = sql.SQL("LOCK TABLE {table} IN ACCESS EXCLUSIVE MODE") + + +@dataclass(frozen=True, slots=True) +class Route: + path: str + table: str + entity_column: str + entity_filter: str | None + + +ROUTES: Final = ( + Route("/user/daily/activity", USER_SPEND, "user_id", None), + Route(AGGREGATED_USER_ACTIVITY, USER_SPEND, "user_id", None), + Route("/team/daily/activity", TEAM_SPEND, "team_id", "team_ids"), + Route("/team/daily/activity/aggregated", TEAM_SPEND, "team_id", "team_ids"), + Route("/tag/daily/activity", TAG_SPEND, "tag", "tags"), + Route("/organization/daily/activity", ORGANIZATION_SPEND, "organization_id", "organization_ids"), + Route("/customer/daily/activity", END_USER_SPEND, "end_user_id", "end_user_ids"), + Route("/end_user/daily/activity", END_USER_SPEND, "end_user_id", "end_user_ids"), + Route("/agent/daily/activity", AGENT_SPEND, "agent_id", "agent_ids"), +) + + +def user_with_an_email(scenario: Scenario) -> tuple[str, str]: + email: Final = f"integration-{uuid.uuid4().hex}@example.com" + return scenario.user(user_email=email), email + + +def key_no_key_table_holds() -> str: + return f"integration-ownerless-{uuid.uuid4().hex}" + + +def activity_of_key( + gateway: Gateway, path: str, api_key: str, *, reader: str | None = None, **filters: str +) -> httpx.Response: + return gateway.request( + "GET", path, params={"start_date": DAY, "end_date": DAY, "api_key": api_key, **filters}, key=reader + ) + + +@dataclass(frozen=True, slots=True) +class DailyRow: + table: str + entity_column: str + entity: str | None + api_key: str + date: str + model: str + provider: str + prompt_tokens: int + completion_tokens: int + spend: float + successful_requests: int + failed_requests: int + + +def _insert(connection: psycopg.Connection[tuple[object, ...]], row: DailyRow) -> None: + connection.execute( + INSERT_DAILY_ROW.format(table=sql.Identifier(row.table), entity=sql.Identifier(row.entity_column)), + ( + row.entity, + row.date, + row.api_key, + row.model, + row.model, + row.provider, + row.prompt_tokens, + row.completion_tokens, + row.spend, + row.successful_requests + row.failed_requests, + row.successful_requests, + row.failed_requests, + ), + ) + + +def insert_daily_rows(rows: Sequence[DailyRow], *, database_url: str | None = None) -> None: + with psycopg.connect(database_url or os.environ["DATABASE_URL"]) as connection: + for row in rows: + _insert(connection, row) + + +def delete_daily_rows(rows: Sequence[DailyRow], *, database_url: str | None = None) -> None: + with psycopg.connect(database_url or os.environ["DATABASE_URL"]) as connection: + for table in sorted({row.table for row in rows}): + connection.execute( + DELETE_DAILY_ROWS.format(table=sql.Identifier(table)), + (sorted({row.api_key for row in rows if row.table == table}),), + ) + + +@contextmanager +def daily_rows(rows: Sequence[DailyRow], *, database_url: str | None = None) -> Iterator[None]: + insert_daily_rows(rows, database_url=database_url) + try: + yield + finally: + delete_daily_rows(rows, database_url=database_url) + + +@contextmanager +def spend_log_naming_only_an_alias(request_id: str, api_key: str, started: str, alias: str) -> Iterator[None]: + with psycopg.connect(os.environ["DATABASE_URL"]) as connection: + connection.execute( + INSERT_SPEND_LOG, (request_id, api_key, started, started, Jsonb({"user_api_key_alias": alias})) + ) + try: + yield + finally: + with psycopg.connect(os.environ["DATABASE_URL"]) as connection: + connection.execute(DELETE_SPEND_LOG, (request_id,)) + + +@contextmanager +def locked_table(table: str, *, database_url: str | None = None) -> Iterator[None]: + with psycopg.connect(database_url or os.environ["DATABASE_URL"]) as connection: + connection.execute(LOCK_TABLE.format(table=sql.Identifier(table))) + try: + yield + finally: + connection.rollback() + + +def records_of_key(node: JsonValue, api_key: str) -> tuple[JsonValue, ...]: + if isinstance(node, list): + return tuple(chain.from_iterable(records_of_key(item, api_key) for item in node)) + if not isinstance(node, dict): + return () + nested: Final = tuple(chain.from_iterable(records_of_key(value, api_key) for value in node.values())) + return (node[api_key], *nested) if api_key in node else nested + + +def seeded_row(table: str, entity_column: str, entity: str | None, api_key: str, date: str) -> DailyRow: + return DailyRow(table, entity_column, entity, api_key, date, "gpt-4o-mini", "openai", 10, 5, 0.25, 1, 0) + + +def user_row(user: str | None, api_key: str, date: str) -> DailyRow: + return seeded_row(USER_SPEND, "user_id", user, api_key, date) + + +def seeded_metrics(rows: int) -> dict[str, float]: + return { + "spend": 0.25 * rows, + "prompt_tokens": 10 * rows, + "completion_tokens": 5 * rows, + "total_tokens": 15 * rows, + "api_requests": rows, + "successful_requests": rows, + } + + +def key_metadata( + *, + alias: str | None = None, + team: str | None = None, + user: str | None = None, + email: str | None = None, + exists: bool = False, +) -> dict[str, JsonValue]: + return {"key_alias": alias, "team_id": team, "user_id": user, "user_email": email, "key_exists": exists} + + +def counted(metrics: JsonValue) -> dict[str, JsonValue]: + return {name: value for name, value in object_value(metrics).items() if value} + + +def assert_key_reported( + response: httpx.Response, + api_key: str, + date: str, + metadata: Mapping[str, JsonValue], + metrics: Mapping[str, float], +) -> None: + assert response.status_code == 200, response.text + body: Final = object_value(response.json()) + records: Final = tuple(object_value(record) for record in records_of_key(body, api_key)) + assert records, response.text + assert all(record["metadata"] == metadata for record in records), response.text + assert all(counted(record["metrics"]) == pytest.approx(metrics) for record in records), response.text + days: Final = body["results"] + assert isinstance(days, list) and len(days) == 1, response.text + day: Final = object_value(days[0]) + assert day["date"] == date, response.text + assert counted(day["metrics"]) == pytest.approx(metrics), response.text + assert object_value(body["metadata"])["total_spend"] == pytest.approx(metrics["spend"]), response.text + + +def assert_key_owner_and_totals( + response: httpx.Response, + api_key: str, + metadata: Mapping[str, JsonValue], + totals: Mapping[str, float], +) -> None: + assert response.status_code == 200, response.text + body: Final = object_value(response.json()) + records: Final = tuple(object_value(record) for record in records_of_key(body, api_key)) + assert records, response.text + assert all(record["metadata"] == metadata for record in records), response.text + reported: Final = object_value(body["metadata"]) + assert {name: reported[name] for name in totals} == pytest.approx(totals), response.text diff --git a/tests/integration/spend/test_daily_activity_key_owner.py b/tests/integration/spend/test_daily_activity_key_owner.py new file mode 100644 index 00000000000..cec19ce5ea0 --- /dev/null +++ b/tests/integration/spend/test_daily_activity_key_owner.py @@ -0,0 +1,196 @@ +import uuid +from hashlib import sha256 +from typing import Final + +import pytest +from integration._support.client import Gateway, Scenario, string_value +from integration._support.daily_activity import ( + AGGREGATED_USER_ACTIVITY, + DAY, + ROUTES, + USER_SPEND, + Route, + activity_of_key, + assert_key_reported, + daily_rows, + key_metadata, + key_no_key_table_holds, + seeded_metrics, + seeded_row, + spend_log_naming_only_an_alias, + user_row, + user_with_an_email, +) + + +@pytest.mark.parametrize("route", ROUTES, ids=lambda route: route.path.strip("/").replace("/", "_")) +def test_key_missing_from_the_key_tables_is_reported_with_the_one_user_its_daily_spend_names( + gateway: Gateway, route: Route +) -> None: + api_key: Final = key_no_key_table_holds() + entity: Final = f"integration-entity-{uuid.uuid4().hex}" + with gateway.scenario() as scenario: + owner, email = user_with_an_email(scenario) + entity_rows: Final = ( + () if route.table == USER_SPEND else (seeded_row(route.table, route.entity_column, entity, api_key, DAY),) + ) + filters: Final = {} if route.entity_filter is None else {route.entity_filter: entity} + with daily_rows((user_row(owner, api_key, DAY), *entity_rows)): + assert_key_reported( + activity_of_key(gateway, route.path, api_key, **filters), + api_key, + DAY, + key_metadata(user=owner, email=email), + seeded_metrics(1), + ) + + +def test_key_whose_daily_spend_names_two_users_is_reported_with_no_owner(gateway: Gateway) -> None: + api_key: Final = key_no_key_table_holds() + with gateway.scenario() as scenario: + first, _ = user_with_an_email(scenario) + second, _ = user_with_an_email(scenario) + with daily_rows((user_row(first, api_key, DAY), user_row(second, api_key, DAY))): + assert_key_reported( + activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key), + api_key, + DAY, + key_metadata(), + seeded_metrics(2), + ) + + +@pytest.mark.parametrize("unnamed", ["", None], ids=["blank_user", "null_user"]) +def test_daily_spend_rows_naming_no_user_do_not_hide_the_one_user_the_others_name( + gateway: Gateway, unnamed: str | None +) -> None: + api_key: Final = key_no_key_table_holds() + with gateway.scenario() as scenario: + owner, email = user_with_an_email(scenario) + with daily_rows((user_row(owner, api_key, DAY), user_row(unnamed, api_key, DAY))): + assert_key_reported( + activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key), + api_key, + DAY, + key_metadata(user=owner, email=email), + seeded_metrics(2), + ) + + +def test_key_whose_daily_spend_names_no_user_at_all_is_reported_with_no_owner(gateway: Gateway) -> None: + api_key: Final = key_no_key_table_holds() + with daily_rows((user_row("", api_key, DAY), user_row(None, api_key, DAY))): + assert_key_reported( + activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key), + api_key, + DAY, + key_metadata(), + seeded_metrics(2), + ) + + +def test_owner_the_user_table_does_not_hold_is_reported_by_id_with_no_email(gateway: Gateway) -> None: + api_key: Final = key_no_key_table_holds() + owner: Final = f"integration-departed-{uuid.uuid4().hex}" + with daily_rows((user_row(owner, api_key, DAY),)): + assert_key_reported( + activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key), + api_key, + DAY, + key_metadata(user=owner), + seeded_metrics(1), + ) + + +def _stored_form(token: str) -> str: + return sha256(token.encode()).hexdigest() + + +def _deleted_key(gateway: Gateway, scenario: Scenario, alias: str, **fields: str) -> str: + token: Final = string_value(gateway.post("/key/generate", {"key_alias": alias, **fields})["key"]) + scenario.delete_key(token) + return _stored_form(token) + + +def test_live_key_keeps_its_own_user_when_its_daily_spend_names_another(gateway: Gateway) -> None: + alias: Final = f"integration-alias-{uuid.uuid4().hex}" + with gateway.scenario() as scenario: + owner, email = user_with_an_email(scenario) + other, _ = user_with_an_email(scenario) + api_key: Final = _stored_form(scenario.key(user_id=owner, key_alias=alias)) + with daily_rows((user_row(other, api_key, DAY),)): + assert_key_reported( + activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key), + api_key, + DAY, + key_metadata(alias=alias, user=owner, email=email, exists=True), + seeded_metrics(1), + ) + + +def test_live_key_with_no_user_is_not_given_the_user_its_daily_spend_names(gateway: Gateway) -> None: + alias: Final = f"integration-alias-{uuid.uuid4().hex}" + with gateway.scenario() as scenario: + spender, _ = user_with_an_email(scenario) + api_key: Final = _stored_form(scenario.key(key_alias=alias)) + with daily_rows((user_row(spender, api_key, DAY),)): + assert_key_reported( + activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key), + api_key, + DAY, + key_metadata(alias=alias, exists=True), + seeded_metrics(1), + ) + + +def test_deleted_key_keeps_its_own_user_when_its_daily_spend_names_another(gateway: Gateway) -> None: + alias: Final = f"integration-alias-{uuid.uuid4().hex}" + with gateway.scenario() as scenario: + owner, email = user_with_an_email(scenario) + other, _ = user_with_an_email(scenario) + api_key: Final = _deleted_key(gateway, scenario, alias, user_id=owner) + with daily_rows((user_row(other, api_key, DAY),)): + assert_key_reported( + activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key), + api_key, + DAY, + key_metadata(alias=alias, user=owner, email=email), + seeded_metrics(1), + ) + + +def test_deleted_key_with_no_user_keeps_its_alias_and_gains_the_one_user_its_daily_spend_names( + gateway: Gateway, +) -> None: + alias: Final = f"integration-alias-{uuid.uuid4().hex}" + with gateway.scenario() as scenario: + owner, email = user_with_an_email(scenario) + api_key: Final = _deleted_key(gateway, scenario, alias) + with daily_rows((user_row(owner, api_key, DAY),)): + assert_key_reported( + activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key), + api_key, + DAY, + key_metadata(alias=alias, user=owner, email=email), + seeded_metrics(1), + ) + + +def test_key_named_only_by_a_spend_log_alias_keeps_that_alias_and_gains_the_one_user_its_daily_spend_names( + gateway: Gateway, +) -> None: + api_key: Final = sha256(uuid.uuid4().bytes).hexdigest() + alias: Final = f"integration-alias-{uuid.uuid4().hex}" + with gateway.scenario() as scenario: + owner, email = user_with_an_email(scenario) + with ( + spend_log_naming_only_an_alias(f"integration-{uuid.uuid4().hex}", api_key, f"{DAY} 12:00:00", alias), + daily_rows((user_row(owner, api_key, DAY),)), + ): + assert_key_reported( + activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key), + api_key, + DAY, + key_metadata(alias=alias, user=owner, email=email), + seeded_metrics(1), + ) diff --git a/tests/integration/spend/test_daily_activity_key_owner_faults.py b/tests/integration/spend/test_daily_activity_key_owner_faults.py new file mode 100644 index 00000000000..998cd2396ae --- /dev/null +++ b/tests/integration/spend/test_daily_activity_key_owner_faults.py @@ -0,0 +1,264 @@ +import os +import signal +import time +import uuid +from collections.abc import Iterator +from contextlib import contextmanager +from itertools import chain +from pathlib import Path +from typing import Final + +import httpx +import psutil +import pytest +from integration._support.client import Gateway, eventually, object_value +from integration._support.daily_activity import ( + AGGREGATED_USER_ACTIVITY, + DAY, + TEAM_SPEND, + USER_SPEND, + activity_of_key, + assert_key_reported, + daily_rows, + insert_daily_rows, + key_metadata, + key_no_key_table_holds, + locked_table, + records_of_key, + seeded_metrics, + seeded_row, + user_row, + user_with_an_email, +) +from integration._support.database import scratch_database +from integration._support.process import OwnedProxy, group_members, owned_proxy_process + +USER_ACTIVITY: Final = "/user/daily/activity" +TEAM_ACTIVITY: Final = "/team/daily/activity" +AGGREGATED_TEAM_ACTIVITY: Final = "/team/daily/activity/aggregated" +KEYS_OF_ONE_TEAM: Final = 300 +GIVES_UP_WITHIN_SECONDS: Final = 10 +READS_AFTER_THE_WORKER_IS_REPLACED: Final = 6 + + +@contextmanager +def _proxy_on(gateway: Gateway, directory: Path, database_url: str, *, workers: int = 1) -> Iterator[OwnedProxy]: + with owned_proxy_process( + gateway, + directory, + {"DATABASE_URL": database_url}, + remove_environment=("DATABASE_URL_READ_REPLICA",), + workers=workers, + ) as owned: + yield owned + + +def _owner_on(candidate: Gateway) -> tuple[str, str]: + owner: Final = f"integration-{uuid.uuid4().hex}" + email: Final = f"{owner}@example.com" + candidate.post("/user/new", {"user_id": owner, "user_email": email, "auto_create_key": False}) + return owner, email + + +def _read_on_a_new_connection(candidate: Gateway, api_key: str) -> httpx.Response: + return candidate.request( + "GET", + AGGREGATED_USER_ACTIVITY, + params={"start_date": DAY, "end_date": DAY, "api_key": api_key}, + headers={"Connection": "close"}, + ) + + +def _running_children(owned: OwnedProxy) -> tuple[int, ...]: + return tuple( + member.pid + for member in group_members(owned.process.pid) + if member.pid != owned.process.pid and member.is_running() and member.status() != psutil.STATUS_ZOMBIE + ) + + +def test_user_reading_a_key_shared_with_another_user_is_shown_no_owner_and_nothing_of_the_other_user( + gateway: Gateway, +) -> None: + api_key: Final = key_no_key_table_holds() + with gateway.scenario() as scenario: + reader, _ = user_with_an_email(scenario) + other, other_email = user_with_an_email(scenario) + reader_key: Final = scenario.key(user_id=reader) + with daily_rows((user_row(reader, api_key, DAY), user_row(other, api_key, DAY))): + response: Final = activity_of_key(gateway, USER_ACTIVITY, api_key, reader=reader_key) + assert_key_reported(response, api_key, DAY, key_metadata(), seeded_metrics(1)) + assert other not in response.text + assert other_email not in response.text + + +def test_user_reading_a_key_only_they_spent_with_is_shown_themselves_as_its_owner(gateway: Gateway) -> None: + api_key: Final = key_no_key_table_holds() + with gateway.scenario() as scenario: + reader, email = user_with_an_email(scenario) + reader_key: Final = scenario.key(user_id=reader) + with daily_rows((user_row(reader, api_key, DAY),)): + response: Final = activity_of_key(gateway, USER_ACTIVITY, api_key, reader=reader_key) + assert_key_reported(response, api_key, DAY, key_metadata(user=reader, email=email), seeded_metrics(1)) + + +def test_user_reading_a_key_only_another_user_spent_with_is_shown_nothing_of_it(gateway: Gateway) -> None: + api_key: Final = key_no_key_table_holds() + with gateway.scenario() as scenario: + reader, _ = user_with_an_email(scenario) + other, other_email = user_with_an_email(scenario) + reader_key: Final = scenario.key(user_id=reader) + with daily_rows((user_row(other, api_key, DAY),)): + response: Final = activity_of_key(gateway, USER_ACTIVITY, api_key, reader=reader_key) + assert response.status_code == 200, response.text + assert object_value(response.json())["results"] == [], response.text + assert other not in response.text + assert other_email not in response.text + + +def test_invalid_key_is_refused_without_naming_the_owner(gateway: Gateway) -> None: + api_key: Final = key_no_key_table_holds() + with gateway.scenario() as scenario: + owner, email = user_with_an_email(scenario) + with daily_rows((user_row(owner, api_key, DAY),)): + response: Final = activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key, reader="sk-not-a-key") + assert response.status_code == 401, response.text + assert owner not in response.text + assert email not in response.text + + +def test_five_kilobyte_key_is_reported_with_the_one_user_its_daily_spend_names(gateway: Gateway) -> None: + api_key: Final = f"integration-5kb-{uuid.uuid4().hex}-{'k' * 5000}" + with gateway.scenario() as scenario: + owner, email = user_with_an_email(scenario) + with daily_rows((user_row(owner, api_key, DAY),)): + assert_key_reported( + activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key), + api_key, + DAY, + key_metadata(user=owner, email=email), + seeded_metrics(1), + ) + + +def test_key_with_no_daily_spend_is_reported_as_no_activity(gateway: Gateway) -> None: + response: Final = activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, key_no_key_table_holds()) + assert response.status_code == 200, response.text + body: Final = object_value(response.json()) + assert body["results"] == [], response.text + totals: Final = object_value(body["metadata"]) + assert [totals["total_spend"], totals["total_api_requests"]] == [0.0, 0], response.text + + +def test_every_key_of_a_team_is_reported_with_its_own_user(gateway: Gateway) -> None: + team: Final = f"integration-entity-{uuid.uuid4().hex}" + owners: Final = {key_no_key_table_holds(): f"integration-owner-{uuid.uuid4().hex}" for _ in range(KEYS_OF_ONE_TEAM)} + rows: Final = tuple( + chain.from_iterable( + (user_row(owner, api_key, DAY), seeded_row(TEAM_SPEND, "team_id", team, api_key, DAY)) + for api_key, owner in owners.items() + ) + ) + with daily_rows(rows): + response: Final = gateway.request( + "GET", AGGREGATED_TEAM_ACTIVITY, params={"start_date": DAY, "end_date": DAY, "team_ids": team} + ) + assert response.status_code == 200, response.text + body: Final = object_value(response.json()) + days: Final = body["results"] + assert isinstance(days, list) and len(days) == 1, response.text + reported: Final = object_value(object_value(object_value(days[0])["breakdown"])["api_keys"]) + assert {api_key: object_value(record)["metadata"] for api_key, record in reported.items()} == { + api_key: key_metadata(user=owner) for api_key, owner in owners.items() + }, response.text + totals: Final = object_value(body["metadata"]) + assert totals["total_api_requests"] == KEYS_OF_ONE_TEAM, response.text + assert totals["total_spend"] == pytest.approx(0.25 * KEYS_OF_ONE_TEAM), response.text + + +def test_reading_the_same_activity_twice_gives_the_same_answer(gateway: Gateway) -> None: + api_key: Final = key_no_key_table_holds() + with gateway.scenario() as scenario: + owner, _ = user_with_an_email(scenario) + with daily_rows((user_row(owner, api_key, DAY),)): + first: Final = activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key) + second: Final = activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key) + assert [first.status_code, second.status_code] == [200, 200], [first.text, second.text] + assert records_of_key(first.json(), api_key), first.text + assert first.json() == second.json(), [first.text, second.text] + + +def test_key_stops_being_reported_with_an_owner_once_a_second_user_spends_with_it(gateway: Gateway) -> None: + api_key: Final = key_no_key_table_holds() + with gateway.scenario() as scenario: + first, email = user_with_an_email(scenario) + second, _ = user_with_an_email(scenario) + with daily_rows((user_row(first, api_key, DAY),)): + alone: Final = activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key) + with daily_rows((user_row(second, api_key, DAY),)): + shared: Final = activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key) + assert_key_reported(alone, api_key, DAY, key_metadata(user=first, email=email), seeded_metrics(1)) + assert_key_reported(shared, api_key, DAY, key_metadata(), seeded_metrics(2)) + + +@pytest.mark.timeout(300) +def test_owner_lookup_gives_up_while_daily_user_spend_is_locked_and_answers_once_it_is_not( + gateway: Gateway, tmp_path: Path +) -> None: + api_key: Final = key_no_key_table_holds() + team: Final = f"integration-entity-{uuid.uuid4().hex}" + with scratch_database() as database_url, _proxy_on(gateway, tmp_path, database_url) as owned: + owner, email = _owner_on(owned.gateway) + rows: Final = (user_row(owner, api_key, DAY), seeded_row(TEAM_SPEND, "team_id", team, api_key, DAY)) + with daily_rows(rows, database_url=database_url): + with locked_table(USER_SPEND, database_url=database_url): + started: Final = time.monotonic() + locked: Final = activity_of_key(owned.gateway, TEAM_ACTIVITY, api_key, team_ids=team) + waited: Final = time.monotonic() - started + unlocked: Final = activity_of_key(owned.gateway, TEAM_ACTIVITY, api_key, team_ids=team) + assert waited < GIVES_UP_WITHIN_SECONDS, waited + assert_key_reported(locked, api_key, DAY, key_metadata(), seeded_metrics(1)) + assert_key_reported(unlocked, api_key, DAY, key_metadata(user=owner, email=email), seeded_metrics(1)) + + +@pytest.mark.timeout(300) +def test_owner_is_reported_while_a_worker_is_killed_and_after_it_is_replaced(gateway: Gateway, tmp_path: Path) -> None: + api_key: Final = key_no_key_table_holds() + with scratch_database() as database_url, _proxy_on(gateway, tmp_path, database_url, workers=2) as owned: + owner, email = _owner_on(owned.gateway) + with daily_rows((user_row(owner, api_key, DAY),), database_url=database_url): + before: Final = _read_on_a_new_connection(owned.gateway, api_key) + members: Final = tuple( + member for member in group_members(owned.process.pid) if member.pid != owned.process.pid + ) + children: Final = tuple(member.pid for member in members) + workers: Final = tuple( + member.pid for member in members if any("spawn_main" in part for part in member.cmdline()) + ) + assert len(workers) >= 2, workers + os.kill(workers[0], signal.SIGKILL) + during: Final = _read_on_a_new_connection(owned.gateway, api_key) + eventually( + lambda: _running_children(owned), + lambda pids: len(pids) >= len(children) and any(pid not in children for pid in pids), + seconds=30, + ) + after: Final = tuple( + _read_on_a_new_connection(owned.gateway, api_key) for _ in range(READS_AFTER_THE_WORKER_IS_REPLACED) + ) + for response in (before, during, *after): + assert_key_reported(response, api_key, DAY, key_metadata(user=owner, email=email), seeded_metrics(1)) + + +@pytest.mark.timeout(300) +def test_owner_is_reported_again_after_the_proxy_restarts(gateway: Gateway, tmp_path: Path) -> None: + api_key: Final = key_no_key_table_holds() + with scratch_database() as database_url: + with _proxy_on(gateway, tmp_path, database_url) as first: + owner, email = _owner_on(first.gateway) + insert_daily_rows((user_row(owner, api_key, DAY),), database_url=database_url) + before: Final = activity_of_key(first.gateway, AGGREGATED_USER_ACTIVITY, api_key) + with _proxy_on(gateway, tmp_path, database_url) as second: + after: Final = activity_of_key(second.gateway, AGGREGATED_USER_ACTIVITY, api_key) + for response in (before, after): + assert_key_reported(response, api_key, DAY, key_metadata(user=owner, email=email), seeded_metrics(1)) diff --git a/tests/integration/spend/test_daily_activity_key_owner_traffic.py b/tests/integration/spend/test_daily_activity_key_owner_traffic.py new file mode 100644 index 00000000000..e0ec1310485 --- /dev/null +++ b/tests/integration/spend/test_daily_activity_key_owner_traffic.py @@ -0,0 +1,397 @@ +import json +import os +import threading +import uuid +from collections.abc import Iterable +from concurrent.futures import ThreadPoolExecutor +from datetime import UTC, datetime, timedelta +from hashlib import sha256 +from pathlib import Path +from queue import SimpleQueue +from typing import Final + +import httpx +import pytest +from integration._support.client import Gateway, Scenario, eventually +from integration._support.daily_activity import ( + AGGREGATED_USER_ACTIVITY, + DAY, + ROUTES, + USER_SPEND, + Route, + activity_of_key, + assert_key_owner_and_totals, + assert_key_reported, + daily_rows, + key_metadata, + key_no_key_table_holds, + seeded_metrics, + seeded_row, + user_row, + user_with_an_email, +) +from integration._support.database import read_rows, scratch_database +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue + +from litellm.proxy._types import LiteLLM_UserTable +from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken + +REQUESTS_OF_KEY: Final = ( + 'SELECT COALESCE(SUM(api_requests), 0)::int AS requests FROM "LiteLLM_DailyUserSpend" ' + "WHERE api_key=%s AND user_id=%s" +) +UNIFIED_ENDPOINTS: Final = ("/v1/chat/completions", "/v1/messages", "/v1/responses") +REQUESTS_OF_A_BURST: Final = 21 +READS_DURING_A_BURST: Final = 30 +TOKEN_LIMIT_DISCOVERY: Final = ("GET", "/v1/models") +TOOL_CALL: Final = "call_integration_usage" +ANSWER: Final = "One request cost $0.25" +SUMMARY_OF_ONE_SEEDED_ROW: Final = "\n".join( + ( + "Total Spend: $0.2500", + "Total Requests: 1", + "Successful: 1 | Failed: 0", + "Total Tokens: 15", + "", + "Top Models by Spend:", + " - gpt-4o-mini: $0.2500 (1 reqs, 15 tokens)", + "", + "Top Providers by Spend:", + " - openai: $0.2500 (1 reqs)", + ) +) + + +def _chat_completion() -> dict[str, JsonValue]: + return { + "id": f"chatcmpl-{uuid.uuid4().hex}", + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}, + } + + +def _response() -> dict[str, JsonValue]: + return { + "id": f"resp_{uuid.uuid4().hex}", + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [ + { + "type": "message", + "id": f"msg_{uuid.uuid4().hex}", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "ok", "annotations": []}], + } + ], + "usage": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15}, + } + + +def _provider(request: Request) -> Reply: + body: Final = _response() if request.target.endswith("/responses") else _chat_completion() + return Reply(body=json.dumps(body).encode()) + + +def _usage_tool_call() -> dict[str, JsonValue]: + call: Final[dict[str, JsonValue]] = { + "id": TOOL_CALL, + "type": "function", + "function": { + "name": "get_usage_data", + "arguments": json.dumps({"start_date": DAY, "end_date": DAY}), + }, + } + return { + "id": f"chatcmpl-{uuid.uuid4().hex}", + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": None, "tool_calls": [call]}, + "finish_reason": "tool_calls", + } + ], + "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}, + } + + +def _streamed_chunk(delta: dict[str, JsonValue], finish_reason: str | None) -> bytes: + chunk: Final = { + "id": "chatcmpl-integration-usage", + "object": "chat.completion.chunk", + "created": 1, + "model": "gpt-4o-mini", + "choices": [{"index": 0, "delta": delta, "finish_reason": finish_reason}], + } + return f"data: {json.dumps(chunk)}\n\n".encode() + + +def _usage_analyst(request: Request) -> Reply: + if json.loads(request.body).get("stream"): + return Reply( + chunks=( + _streamed_chunk({"role": "assistant", "content": ANSWER}, None), + _streamed_chunk({}, "stop"), + b"data: [DONE]\n\n", + ), + content_type="text/event-stream", + ) + return Reply(body=json.dumps(_usage_tool_call()).encode()) + + +def _sent_for_callers(requests: Iterable[Request]) -> tuple[Request, ...]: + return tuple(request for request in requests if (request.method, request.target) != TOKEN_LIMIT_DISCOVERY) + + +def _priced_model(scenario: Scenario, provider_url: str) -> str: + return scenario.model( + api_base=f"{provider_url}/v1", input_cost_per_token=0.001, output_cost_per_token=0.002, num_retries=0 + ) + + +def _request_body(endpoint: str, model: str, prompt: str) -> dict[str, JsonValue]: + if endpoint == "/v1/chat/completions": + return {"model": model, "messages": [{"role": "user", "content": prompt}]} + if endpoint == "/v1/messages": + return {"model": model, "max_tokens": 16, "messages": [{"role": "user", "content": prompt}]} + return {"model": model, "input": prompt} + + +def _activity_on_route(gateway: Gateway, route: Route, api_key: str, entity: str) -> httpx.Response: + filters: Final = {} if route.entity_filter is None else {route.entity_filter: entity} + return activity_of_key(gateway, route.path, api_key, **filters) + + +def _prompt() -> str: + return f"daily activity owner {uuid.uuid4().hex}" + + +def _totals_of_requests(requests: int) -> dict[str, float]: + return { + "total_spend": 0.02 * requests, + "total_prompt_tokens": 10 * requests, + "total_completion_tokens": 5 * requests, + "total_tokens": 15 * requests, + "total_api_requests": requests, + "total_successful_requests": requests, + "total_failed_requests": 0, + } + + +def _activity_around_today(gateway: Gateway, api_key: str) -> httpx.Response: + today: Final = datetime.now(UTC).date() + return gateway.request( + "GET", + AGGREGATED_USER_ACTIVITY, + params={ + "start_date": str(today - timedelta(days=1)), + "end_date": str(today + timedelta(days=1)), + "timezone": "0", + "api_key": api_key, + }, + ) + + +def _wait_for_requests(api_key: str, user: str, requests: int) -> None: + eventually( + lambda: read_rows(REQUESTS_OF_KEY, (api_key, user)), + lambda rows: rows[0]["requests"] == requests, + seconds=70, + ) + + +def _cli_session_token(user: str, team: str) -> str: + cli_user: Final = LiteLLM_UserTable(user_id=user, user_role="internal_user", teams=[team], models=[]) + return ExperimentalUIJWTToken.get_cli_jwt_auth_token(user_info=cli_user, team_id=team, team_alias="cli-team") + + +def test_key_used_on_every_unified_endpoint_is_reported_with_its_own_alias_and_user(gateway: Gateway) -> None: + chat_prompt, messages_prompt, responses_prompt = _prompt(), _prompt(), _prompt() + with wire_server(_provider) as wire, gateway.scenario() as scenario: + model: Final = _priced_model(scenario, wire.url) + owner, email = user_with_an_email(scenario) + alias: Final = f"integration-alias-{uuid.uuid4().hex}" + key: Final = scenario.key(user_id=owner, key_alias=alias, models=[model]) + stored: Final = sha256(key.encode()).hexdigest() + prompts: Final = (chat_prompt, messages_prompt, responses_prompt) + answers: Final = tuple( + gateway.request("POST", endpoint, _request_body(endpoint, model, prompt), key=key) + for endpoint, prompt in zip(UNIFIED_ENDPOINTS, prompts, strict=True) + ) + assert [answer.status_code for answer in answers] == [200, 200, 200], [answer.text for answer in answers] + received: Final = _sent_for_callers(wire.drain()) + assert [request.target for request in received] == ["/v1/chat/completions", "/v1/responses", "/v1/responses"] + assert [json.loads(request.body)["model"] for request in received] == ["gpt-4o-mini"] * 3 + assert json.loads(received[0].body)["messages"] == [{"role": "user", "content": chat_prompt}] + assert messages_prompt in received[1].body.decode() + assert json.loads(received[2].body)["input"] == responses_prompt + _wait_for_requests(stored, owner, 3) + assert_key_owner_and_totals( + _activity_around_today(gateway, stored), + stored, + key_metadata(alias=alias, user=owner, email=email, exists=True), + _totals_of_requests(3), + ) + + +def test_cli_session_spend_is_reported_with_the_user_and_team_of_the_session( + gateway: Gateway, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", os.environ.get("LITELLM_SALT_KEY", "sk-integration-salt")) + prompt: Final = _prompt() + with wire_server(_provider) as wire, gateway.scenario() as scenario: + model: Final = _priced_model(scenario, wire.url) + owner, email = user_with_an_email(scenario) + team: Final = scenario.team(models=[model], members_with_roles=[{"role": "user", "user_id": owner}]) + answer: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": prompt}]}, + key=_cli_session_token(owner, team), + ) + assert answer.status_code == 200, answer.text + received: Final = _sent_for_callers(wire.drain()) + assert [request.target for request in received] == ["/v1/chat/completions"] + assert json.loads(received[0].body) == { + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": prompt}], + } + stored: Final = f"cli-session-{owner}" + _wait_for_requests(stored, owner, 1) + assert_key_owner_and_totals( + _activity_around_today(gateway, stored), + stored, + key_metadata(alias=stored, team=team, user=owner, email=email), + _totals_of_requests(1), + ) + + +@pytest.mark.timeout(300) +def test_usage_ai_chat_hands_the_model_the_usage_summary_without_any_key_owner( + gateway: Gateway, tmp_path: Path +) -> None: + question: Final = f"what did we spend {uuid.uuid4().hex}" + owner: Final = f"integration-{uuid.uuid4().hex}" + ownerless_key: Final = f"integration-ownerless-{uuid.uuid4().hex}" + with ( + scratch_database() as scratch_url, + wire_server(_usage_analyst) as wire, + owned_proxy( + gateway, + tmp_path, + { + "DATABASE_URL": scratch_url, + "OPENAI_API_BASE": f"{wire.url}/v1", + "OPENAI_BASE_URL": f"{wire.url}/v1", + "OPENAI_API_KEY": "integration-provider-key", + }, + remove_environment=("DATABASE_URL_READ_REPLICA",), + ) as candidate, + ): + candidate.post("/user/new", {"user_id": owner, "user_email": f"{owner}@example.com", "auto_create_key": False}) + with daily_rows((user_row(owner, ownerless_key, DAY),), database_url=scratch_url): + answer: Final = candidate.request( + "POST", + "/usage/ai/chat", + {"messages": [{"role": "user", "content": question}], "model": "openai/gpt-4o-mini"}, + ) + assert answer.status_code == 200, answer.text + tool_call: Final = { + "type": "tool_call", + "tool_name": "get_usage_data", + "tool_label": "global usage data", + "arguments": {"start_date": DAY, "end_date": DAY}, + } + events: Final = [ + json.loads(line.removeprefix("data: ")) for line in answer.text.splitlines() if line.startswith("data: ") + ] + assert events == [ + {"type": "status", "message": "Thinking..."}, + {**tool_call, "status": "running"}, + {**tool_call, "status": "complete"}, + {"type": "status", "message": "Analyzing results..."}, + {"type": "chunk", "content": ANSWER}, + {"type": "done"}, + ], answer.text + asked, analysed = wire.drain() + assert [asked.target, analysed.target] == ["/v1/chat/completions", "/v1/chat/completions"] + assert json.loads(asked.body)["messages"][-1] == {"role": "user", "content": question} + assert json.loads(analysed.body)["messages"][-1] == { + "role": "tool", + "tool_call_id": TOOL_CALL, + "content": SUMMARY_OF_ONE_SEEDED_ROW, + } + assert owner not in analysed.body.decode() + assert ownerless_key not in analysed.body.decode() + + +@pytest.mark.timeout(300) +def test_owner_is_reported_on_every_route_while_a_burst_of_requests_waits_on_the_provider(gateway: Gateway) -> None: + released: Final = threading.Event() + held: Final[SimpleQueue[str]] = SimpleQueue() + + def held_provider(request: Request) -> Reply: + if (request.method, request.target) == TOKEN_LIMIT_DISCOVERY: + return _provider(request) + held.put(request.target) + assert released.wait(timeout=120), "The burst was never released" + return _provider(request) + + api_key: Final = key_no_key_table_holds() + entity: Final = f"integration-entity-{uuid.uuid4().hex}" + prompts: Final = tuple(_prompt() for _ in range(REQUESTS_OF_A_BURST)) + entity_columns: Final = {route.table: route.entity_column for route in ROUTES if route.table != USER_SPEND} + with ( + wire_server(held_provider) as wire, + gateway.scenario() as scenario, + httpx.Client(base_url=gateway.client.base_url, timeout=180, trust_env=False) as patient, + ThreadPoolExecutor(max_workers=REQUESTS_OF_A_BURST) as traffic, + ThreadPoolExecutor(max_workers=READS_DURING_A_BURST) as readers, + ): + model: Final = _priced_model(scenario, wire.url) + owner, email = user_with_an_email(scenario) + key: Final = scenario.key(models=[model]) + rows: Final = ( + user_row(owner, api_key, DAY), + *(seeded_row(table, column, entity, api_key, DAY) for table, column in entity_columns.items()), + ) + try: + with daily_rows(rows): + burst: Final = tuple( + traffic.submit( + patient.post, + UNIFIED_ENDPOINTS[index % len(UNIFIED_ENDPOINTS)], + json=_request_body(UNIFIED_ENDPOINTS[index % len(UNIFIED_ENDPOINTS)], model, prompt), + headers={"Authorization": f"Bearer {key}"}, + ) + for index, prompt in enumerate(prompts) + ) + eventually(held.qsize, lambda waiting: waiting >= REQUESTS_OF_A_BURST, seconds=60) + reads: Final = tuple( + readers.submit(_activity_on_route, gateway, ROUTES[index % len(ROUTES)], api_key, entity) + for index in range(READS_DURING_A_BURST) + ) + activity: Final = tuple(read.result() for read in reads) + still_waiting: Final = [call.done() for call in burst] + finally: + released.set() + answers: Final = tuple(call.result() for call in burst) + received: Final = tuple(request.body.decode() for request in _sent_for_callers(wire.drain())) + assert still_waiting == [False] * REQUESTS_OF_A_BURST + assert [answer.status_code for answer in answers] == [200] * REQUESTS_OF_A_BURST, [ + answer.text for answer in answers + ] + assert [sum(prompt in body for body in received) for prompt in prompts] == [1] * REQUESTS_OF_A_BURST + assert len(received) == REQUESTS_OF_A_BURST, len(received) + for response in activity: + assert_key_reported(response, api_key, DAY, key_metadata(user=owner, email=email), seeded_metrics(1)) diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py index 32856a3bee9..52c374fe5a5 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py @@ -23,6 +23,7 @@ from litellm.proxy.management_endpoints.common_daily_activity import ( update_metrics, ) from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR +from litellm.proxy.utils import hash_token from litellm.types.proxy.management_endpoints.common_daily_activity import ( DailySpendMetadata, SpendMetrics, @@ -505,6 +506,7 @@ async def test_get_api_key_metadata_permanent_miss_never_pages_tokens_or_reads_s mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_usertable.find_many = AsyncMock(return_value=[]) mock_prisma.db.query_raw = AsyncMock(return_value=[]) + recovery_query_raw = _recovery_transaction(mock_prisma) result = await get_api_key_metadata( prisma_client=mock_prisma, @@ -512,9 +514,10 @@ async def test_get_api_key_metadata_permanent_miss_never_pages_tokens_or_reads_s ) assert double_hashed not in result - issued_sql = [call.args[0] for call in mock_prisma.db.query_raw.call_args_list] - assert len(issued_sql) == 2 - assert not any("LiteLLM_SpendLogs" in sql for sql in issued_sql) + assert mock_prisma.db.query_raw.await_count == 2 + ((owner_sql, owner_keys),) = [call.args for call in recovery_query_raw.call_args_list] + assert _DAILY_USER_SPEND in owner_sql + assert owner_keys == [double_hashed] token_lookups = ( mock_prisma.db.litellm_verificationtoken.find_many.call_args_list + mock_prisma.db.litellm_deletedverificationtoken.find_many.call_args_list @@ -522,14 +525,29 @@ async def test_get_api_key_metadata_permanent_miss_never_pages_tokens_or_reads_s assert all("take" not in call.kwargs and "skip" not in call.kwargs for call in token_lookups) -def _spend_log_transaction(mock_prisma: MagicMock, rows: list[dict[str, str | None]]) -> AsyncMock: +_DAILY_USER_SPEND: Final = '"LiteLLM_DailyUserSpend"' +_SPEND_LOGS: Final = '"LiteLLM_SpendLogs"' + + +def _recovery_transaction( + mock_prisma: MagicMock, + spend_log_rows: Sequence[dict[str, str | None]] = (), + daily_spend_owner_rows: Sequence[dict[str, str | None]] = (), +) -> AsyncMock: + async def query_raw(sql: str, *_: object) -> Sequence[dict[str, str | None]]: + return daily_spend_owner_rows if _DAILY_USER_SPEND in sql else spend_log_rows + transaction = MagicMock() transaction.execute_raw = AsyncMock(return_value=0) - transaction.query_raw = AsyncMock(return_value=rows) + transaction.query_raw = AsyncMock(side_effect=query_raw) mock_prisma.db.tx.return_value.__aenter__.return_value = transaction return transaction.query_raw +def _calls_reading(query_raw: AsyncMock, table: str) -> tuple[tuple[object, ...], ...]: + return tuple(call.args for call in query_raw.call_args_list if table in call.args[0]) + + def _spend_log_row(digest: str, key_alias: str, user_id: str) -> dict[str, str | None]: return { "digest": digest, @@ -553,15 +571,17 @@ async def test_get_api_key_metadata_permanent_miss_with_a_window_reads_spend_log mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_usertable.find_many = AsyncMock(return_value=[]) mock_prisma.db.query_raw = AsyncMock(return_value=[]) - spend_log_query_raw = _spend_log_transaction(mock_prisma, []) + recovery_query_raw = _recovery_transaction(mock_prisma) result = await get_api_key_metadata(prisma_client=mock_prisma, api_keys={double_hashed}, spend_logs_window=window) assert double_hashed not in result assert mock_prisma.db.query_raw.await_count == 2 - ((_, digests, start, end),) = [call.args for call in spend_log_query_raw.call_args_list] + ((_, digests, start, end),) = _calls_reading(recovery_query_raw, _SPEND_LOGS) assert digests == [double_hashed] assert (start, end) == window + ((_, owner_keys),) = _calls_reading(recovery_query_raw, _DAILY_USER_SPEND) + assert owner_keys == [double_hashed] @pytest.mark.asyncio @@ -583,7 +603,7 @@ async def test_get_daily_activity_recovers_a_session_key_alias_from_spend_logs_a ) mock_prisma.db.query_raw = AsyncMock(return_value=[]) - spend_log_query_raw = _spend_log_transaction( + spend_log_query_raw = _recovery_transaction( mock_prisma, [_spend_log_row(session_digest, "cli-session-alias", "session-user")] ) @@ -1544,6 +1564,8 @@ async def test_get_daily_activity_aggregated_returns_every_api_key( mock_prisma.db.query_raw = _psycopg_query_raw(_aggregated_postgresql, []) mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_usertable = MagicMock() + mock_prisma.db.litellm_usertable.find_many = AsyncMock(return_value=[]) result = await get_daily_activity_aggregated( prisma_client=mock_prisma, @@ -1599,6 +1621,8 @@ async def test_get_daily_activity_aggregated_explicit_api_key_filter_scopes_resu mock_prisma.db.query_raw = _psycopg_query_raw(_aggregated_postgresql, []) mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_usertable = MagicMock() + mock_prisma.db.litellm_usertable.find_many = AsyncMock(return_value=[]) result = await get_daily_activity_aggregated( prisma_client=mock_prisma, @@ -1651,6 +1675,8 @@ async def test_get_daily_activity_aggregated_model_group_rollups_fall_back_to_mo mock_prisma.db.query_raw = _psycopg_query_raw(_aggregated_postgresql, []) mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_usertable = MagicMock() + mock_prisma.db.litellm_usertable.find_many = AsyncMock(return_value=[]) result = await get_daily_activity_aggregated( prisma_client=mock_prisma, @@ -2580,7 +2606,7 @@ async def test_get_api_key_metadata_resolves_session_key_via_spend_log_window(): ) mock_prisma.db.query_raw = AsyncMock(return_value=[]) - spend_log_query_raw = _spend_log_transaction( + spend_log_query_raw = _recovery_transaction( mock_prisma, [_spend_log_row(session_digest, "cli-session-user-42", "user-42")] ) @@ -2627,3 +2653,90 @@ async def test_get_api_key_metadata_resolves_cli_session_keys_from_the_key_itsel assert result["cli-session-alice"]["key_alias"] == "cli-session-alice" assert result["cli-session-alice"]["user_email"] == "alice@example.com" assert result["cli-session-alice"]["team_id"] == "team-a" + + +@pytest.mark.asyncio +async def test_get_api_key_metadata_recovers_legacy_hashed_jwt_owner_from_daily_spend(): + api_key: Final = f"hashed-jwt-{hash_token('legacy-cli-session-daily-spend-owner')}" + user_id: Final = "legacy-owner" + mock_prisma: Final = MagicMock() + mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_usertable.find_many = AsyncMock( + return_value=[SimpleNamespace(user_id=user_id, user_email="legacy-owner@example.com", teams=[])] + ) + recovery_query_raw: Final = _recovery_transaction( + mock_prisma, + daily_spend_owner_rows=[{"api_key": api_key, "first_owner": user_id, "last_owner": user_id}], + ) + + result: Final = await get_api_key_metadata( + prisma_client=mock_prisma, + api_keys={api_key}, + spend_logs_window=(datetime(2026, 9, 7), datetime(2026, 9, 10)), + ) + + assert result.get(api_key, {}).get("user_id") == user_id + assert result.get(api_key, {}).get("user_email") == "legacy-owner@example.com" + assert len(_calls_reading(recovery_query_raw, _SPEND_LOGS)) == 1 + assert len(_calls_reading(recovery_query_raw, _DAILY_USER_SPEND)) == 1 + + +@pytest.mark.asyncio +async def test_get_api_key_metadata_preserves_deleted_key_metadata_when_recovering_daily_spend_owner(): + api_key: Final = f"hashed-jwt-{hash_token('legacy-cli-session-daily-spend-metadata')}" + user_id: Final = "legacy-owner" + mock_prisma: Final = MagicMock() + deleted_key: Final = MagicMock() + deleted_key.token = api_key + deleted_key.key_alias = "legacy-cli-key" + deleted_key.team_id = "team-legacy" + deleted_key.user_id = None + mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[deleted_key]) + mock_prisma.db.litellm_usertable.find_many = AsyncMock( + return_value=[SimpleNamespace(user_id=user_id, user_email="legacy-owner@example.com", teams=[])] + ) + recovery_query_raw: Final = _recovery_transaction( + mock_prisma, + daily_spend_owner_rows=[{"api_key": api_key, "first_owner": user_id, "last_owner": user_id}], + ) + + result: Final = await get_api_key_metadata(prisma_client=mock_prisma, api_keys={api_key}) + + recovered_metadata: Final = result[api_key] + assert recovered_metadata.get("key_alias") == "legacy-cli-key" + assert recovered_metadata.get("team_id") == "team-legacy" + assert recovered_metadata.get("user_id") == user_id + assert recovered_metadata.get("user_email") == "legacy-owner@example.com" + recovery_query_raw.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_get_api_key_metadata_does_not_recover_daily_spend_owner_for_active_keys(): + api_key: Final = "active-token-value" + mock_prisma: Final = MagicMock() + active_key: Final = MagicMock() + active_key.token = api_key + active_key.key_alias = "active-key-alias" + active_key.team_id = "active-team" + active_key.user_id = "active-owner" + mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[active_key]) + mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_usertable.find_many = AsyncMock( + return_value=[SimpleNamespace(user_id="active-owner", user_email="active-owner@example.com", teams=[])] + ) + recovery_query_raw: Final = _recovery_transaction( + mock_prisma, + daily_spend_owner_rows=[{"api_key": api_key, "first_owner": "other-owner", "last_owner": "other-owner"}], + ) + + result: Final = await get_api_key_metadata(prisma_client=mock_prisma, api_keys={api_key}) + + active_metadata: Final = result[api_key] + assert active_metadata.get("key_alias") == "active-key-alias" + assert active_metadata.get("team_id") == "active-team" + assert active_metadata.get("user_id") == "active-owner" + assert active_metadata.get("user_email") == "active-owner@example.com" + assert active_metadata.get("key_exists") is True + recovery_query_raw.assert_not_awaited() diff --git a/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py b/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py index acd03964bf3..74f3a2248c7 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py +++ b/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py @@ -3,6 +3,7 @@ import time from collections.abc import Sequence from datetime import datetime, timedelta from types import SimpleNamespace +from typing import Final from unittest.mock import AsyncMock, MagicMock import pytest @@ -20,6 +21,7 @@ from litellm.proxy.spend_tracking.key_metadata_recovery import ( recover_cli_session_key_metadata, recover_double_hashed_key_metadata, recover_key_metadata_from_spend_logs, + recover_key_owner_from_daily_spend, ) from litellm.proxy.utils import hash_token @@ -702,3 +704,91 @@ async def test_attach_user_details_leaves_metadata_unchanged_when_a_later_chunk_ assert mock_prisma.db.litellm_usertable.find_many.call_count == 2 assert attached == recovered + + +def _daily_spend_owner_row(api_key: str, first_owner: str, last_owner: str) -> dict[str, str]: + return {"api_key": api_key, "first_owner": first_owner, "last_owner": last_owner} + + +def _daily_spend_transaction(mock_prisma: MagicMock, query_raw: AsyncMock) -> MagicMock: + transaction: Final = MagicMock() + transaction.execute_raw = AsyncMock(return_value=0) + transaction.query_raw = query_raw + mock_prisma.db.tx.return_value.__aenter__.return_value = transaction + return transaction + + +@pytest.mark.asyncio +async def test_recover_key_owner_from_daily_spend_keeps_a_unanimous_owner(): + key: Final = "hashed-jwt-digest-a" + mock_prisma: Final = MagicMock() + _daily_spend_transaction(mock_prisma, AsyncMock(return_value=[_daily_spend_owner_row(key, "owner-a", "owner-a")])) + + result: Final = await recover_key_owner_from_daily_spend(mock_prisma, {key}) + + assert dict(result) == {key: "owner-a"} + + +@pytest.mark.asyncio +async def test_recover_key_owner_from_daily_spend_drops_conflicting_owners(): + key: Final = "hashed-jwt-digest-b" + mock_prisma: Final = MagicMock() + _daily_spend_transaction(mock_prisma, AsyncMock(return_value=[_daily_spend_owner_row(key, "owner-a", "owner-b")])) + + result: Final = await recover_key_owner_from_daily_spend(mock_prisma, {key}) + + assert dict(result) == {} + + +@pytest.mark.asyncio +async def test_recover_key_owner_from_daily_spend_skips_empty_input(): + mock_prisma: Final = MagicMock() + transaction: Final = _daily_spend_transaction(mock_prisma, AsyncMock(return_value=[])) + + result: Final = await recover_key_owner_from_daily_spend(mock_prisma, frozenset()) + + assert dict(result) == {} + transaction.query_raw.assert_not_awaited() + mock_prisma.db.tx.assert_not_called() + + +@pytest.mark.asyncio +async def test_recover_key_owner_from_daily_spend_returns_empty_on_prisma_error(): + mock_prisma: Final = MagicMock() + _daily_spend_transaction(mock_prisma, AsyncMock(side_effect=PrismaError("db down"))) + + result: Final = await recover_key_owner_from_daily_spend(mock_prisma, {"hashed-jwt-digest-c"}) + + assert dict(result) == {} + + +@pytest.mark.asyncio +async def test_recover_key_owner_from_daily_spend_names_no_owner_when_the_lookup_hits_the_statement_timeout(): + mock_prisma: Final = MagicMock() + _daily_spend_transaction( + mock_prisma, AsyncMock(side_effect=PrismaError("canceling statement due to statement timeout")) + ) + + result: Final = await recover_key_owner_from_daily_spend(mock_prisma, {"hashed-jwt-digest-d"}) + + assert dict(result) == {} + + +@pytest.mark.asyncio +async def test_recover_key_owner_from_daily_spend_bounds_the_lookup_with_a_statement_timeout(): + key: Final = "hashed-jwt-digest-e" + mock_prisma: Final = MagicMock() + transaction: Final = _daily_spend_transaction( + mock_prisma, AsyncMock(return_value=[_daily_spend_owner_row(key, "owner-a", "owner-a")]) + ) + + result: Final = await recover_key_owner_from_daily_spend(mock_prisma, {key}) + + assert dict(result) == {key: "owner-a"} + assert [name for name, _, _ in transaction.mock_calls] == ["execute_raw", "query_raw"] + transaction.execute_raw.assert_awaited_once_with( + f"SET LOCAL statement_timeout = {SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS}" + ) + assert mock_prisma.db.tx.call_args.kwargs["timeout"] == timedelta( + milliseconds=2 * SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS + ) diff --git a/ui/litellm-dashboard/src/components/UsagePage/types.ts b/ui/litellm-dashboard/src/components/UsagePage/types.ts index d53db68bb9f..6d0e0642ba2 100644 --- a/ui/litellm-dashboard/src/components/UsagePage/types.ts +++ b/ui/litellm-dashboard/src/components/UsagePage/types.ts @@ -58,6 +58,7 @@ export interface TopApiKeyData { api_key: string; key_alias: string | null; team_id: string | null; + user: string | null; spend: number; requests: number; tokens: number; diff --git a/ui/litellm-dashboard/src/components/activity_metrics.test.tsx b/ui/litellm-dashboard/src/components/activity_metrics.test.tsx index 1bd7655b5e3..569f69adb56 100644 --- a/ui/litellm-dashboard/src/components/activity_metrics.test.tsx +++ b/ui/litellm-dashboard/src/components/activity_metrics.test.tsx @@ -258,10 +258,20 @@ describe("ActivityMetrics", () => { api_key: "key-123", key_alias: "Test Key", team_id: "team1", + user: "owner@example.com", spend: 50.25, requests: 25, tokens: 12500, }, + { + api_key: "key-456", + key_alias: "Owner Alias", + team_id: null, + user: "Owner Alias", + spend: 40.25, + requests: 20, + tokens: 10000, + }, ], }, }; @@ -269,6 +279,9 @@ describe("ActivityMetrics", () => { render(); expect(screen.getByText("Top Virtual Keys by Spend")).toBeInTheDocument(); expect(screen.getByText("Test Key")).toBeInTheDocument(); + expect(screen.getByText("User: owner@example.com")).toBeInTheDocument(); + expect(screen.getByText("Owner Alias")).toBeInTheDocument(); + expect(screen.queryByText("User: Owner Alias")).not.toBeInTheDocument(); }); it("should display API key hash when alias is missing", () => { @@ -280,6 +293,7 @@ describe("ActivityMetrics", () => { api_key: "key-1234567890", key_alias: null, team_id: null, + user: null, spend: 50.25, requests: 25, tokens: 12500, @@ -301,6 +315,7 @@ describe("ActivityMetrics", () => { api_key: "key-123", key_alias: "Test Key", team_id: "team1", + user: null, spend: 50.25, requests: 25, tokens: 12500, @@ -1056,6 +1071,8 @@ describe("processActivityData", () => { metadata: { key_alias: "test-key-1", team_id: "team1", + user_id: "owner-id-1", + user_email: "owner-1@example.com", }, }, "key-2": { @@ -1073,6 +1090,7 @@ describe("processActivityData", () => { metadata: { key_alias: "test-key-2", team_id: "team2", + user_id: "owner-id-2", }, }, }, @@ -1094,6 +1112,10 @@ describe("processActivityData", () => { expect(result["gpt-4"].top_api_keys[0].spend).toBe(60.0); expect(result["gpt-4"].top_api_keys[0].api_key).toBe("key-1"); expect(result["gpt-4"].top_api_keys[1].spend).toBe(40.5); + expect(result["gpt-4"].top_api_keys.map(({ api_key, user }) => [api_key, user])).toEqual([ + ["key-1", "owner-1@example.com"], + ["key-2", "owner-id-2"], + ]); }); it("should limit top_api_keys to 5 entries", () => { diff --git a/ui/litellm-dashboard/src/components/activity_metrics.tsx b/ui/litellm-dashboard/src/components/activity_metrics.tsx index 95315f8ec27..5ddb89703ea 100644 --- a/ui/litellm-dashboard/src/components/activity_metrics.tsx +++ b/ui/litellm-dashboard/src/components/activity_metrics.tsx @@ -102,20 +102,26 @@ const ModelSection = ({

    Top Virtual Keys by Spend

    - {metrics.top_api_keys.map((keyData) => ( -
    -
    -

    {keyData.key_alias || `${keyData.api_key.substring(0, 10)}...`}

    - {keyData.team_id &&

    Team: {keyData.team_id}

    } + {metrics.top_api_keys.map((keyData) => { + const keyLabel = keyData.key_alias || `${keyData.api_key.substring(0, 10)}...`; + return ( +
    +
    +

    {keyLabel}

    + {keyData.team_id &&

    Team: {keyData.team_id}

    } + {keyData.user && keyData.user !== keyLabel && ( +

    User: {keyData.user}

    + )} +
    +
    +

    ${formatNumberWithCommas(keyData.spend, 2)}

    +

    + {keyData.requests.toLocaleString()} requests | {keyData.tokens.toLocaleString()} tokens +

    +
    -
    -

    ${formatNumberWithCommas(keyData.spend, 2)}

    -

    - {keyData.requests.toLocaleString()} requests | {keyData.tokens.toLocaleString()} tokens -

    -
    -
    - ))} + ); + })}
    @@ -585,6 +591,7 @@ export const processActivityData = ( api_key: apiKey, key_alias: keyActivityLabel(keyData.metadata, "") || null, team_id: keyData.metadata.team_id, + user: keyData.metadata.user_email ?? keyData.metadata.user_id ?? null, spend: 0, requests: 0, tokens: 0, From e7460f1cff7fc1b6e323c9e4a969fe6b4609d164 Mon Sep 17 00:00:00 2001 From: "berriai-litellm-provider-info-sync[bot]" <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 16:07:08 -0700 Subject: [PATCH 38/41] chore(cost-map): take azure context limits from models-sold-directly (#43759) Price-Sync: litellm-providers Co-authored-by: berriai-litellm-provider-info-sync[bot] <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> --- ...odel_prices_and_context_window_backup.json | 92 ++++++++++++------- model_prices_and_context_window.json | 92 ++++++++++++------- 2 files changed, 114 insertions(+), 70 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 99471d48f56..9989846b7e0 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -3977,7 +3977,7 @@ "input_cost_per_token_priority": 1e-05, "input_cost_per_token_above_272k_tokens_priority": 2e-05, "litellm_provider": "azure_ai", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -4112,7 +4112,7 @@ "input_cost_per_token_priority": 1e-05, "input_cost_per_token_above_272k_tokens_priority": 2e-05, "litellm_provider": "azure_ai", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -4159,7 +4159,7 @@ "input_cost_per_token_priority": 1e-05, "input_cost_per_token_above_272k_tokens_priority": 2e-05, "litellm_provider": "azure_ai", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -4206,7 +4206,7 @@ "input_cost_per_token_priority": 5e-06, "input_cost_per_token_above_272k_tokens_priority": 1e-05, "litellm_provider": "azure_ai", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -4254,7 +4254,7 @@ "input_cost_per_token_priority": 5e-06, "input_cost_per_token_above_272k_tokens_priority": 1e-05, "litellm_provider": "azure_ai", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -4302,7 +4302,7 @@ "input_cost_per_token_priority": 6e-05, "input_cost_per_token_above_272k_tokens_priority": 0.00012, "litellm_provider": "azure_ai", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", @@ -4348,7 +4348,7 @@ "input_cost_per_token_priority": 6e-05, "input_cost_per_token_above_272k_tokens_priority": 0.00012, "litellm_provider": "azure_ai", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", @@ -4678,12 +4678,13 @@ "input_cost_per_audio_token": 4.4e-05, "input_cost_per_token": 5.5e-06, "litellm_provider": "azure", - "max_input_tokens": 128000, + "max_input_tokens": 16000, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "realtime", "output_cost_per_audio_token": 8e-05, "output_cost_per_token": 2.2e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/foundry-models/concepts/models-sold-directly-by-azure", "supported_modalities": [ "text", "audio" @@ -6319,12 +6320,13 @@ "input_cost_per_audio_token": 4e-05, "input_cost_per_token": 5e-06, "litellm_provider": "azure", - "max_input_tokens": 128000, + "max_input_tokens": 16000, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "realtime", "output_cost_per_audio_token": 8e-05, "output_cost_per_token": 2e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/foundry-models/concepts/models-sold-directly-by-azure", "supported_modalities": [ "text", "audio" @@ -6370,6 +6372,9 @@ "deprecation_date": "2027-05-06", "input_cost_per_second": 0.0002833333333333333, "litellm_provider": "azure", + "max_input_tokens": 32000, + "max_output_tokens": 4096, + "max_tokens": 4096, "mode": "audio_transcription", "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/gpt-realtime-whisper", "supported_endpoints": [ @@ -7424,7 +7429,7 @@ "input_cost_per_token_priority": 5e-06, "input_cost_per_token_above_272k_tokens_priority": 1e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -7480,7 +7485,7 @@ "output_cost_per_token": 1.65e-05, "output_cost_per_token_priority": 3.3e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -7530,7 +7535,7 @@ "output_cost_per_token": 1.65e-05, "output_cost_per_token_priority": 3.3e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -7580,7 +7585,7 @@ "input_cost_per_token_priority": 5e-06, "input_cost_per_token_above_272k_tokens_priority": 1e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -7636,7 +7641,7 @@ "output_cost_per_token": 1.65e-05, "output_cost_per_token_priority": 3.3e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -7686,7 +7691,7 @@ "output_cost_per_token": 1.65e-05, "output_cost_per_token_priority": 3.3e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -7737,7 +7742,7 @@ "input_cost_per_token_batches": 1.5e-05, "input_cost_per_token_flex": 1.5e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", @@ -7786,7 +7791,7 @@ "input_cost_per_token_batches": 1.5e-05, "input_cost_per_token_flex": 1.5e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", @@ -9374,7 +9379,7 @@ "input_cost_per_token_batches": 2.5e-06, "input_cost_per_token_flex": 2.5e-06, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -9433,7 +9438,7 @@ "output_cost_per_token_above_272k_tokens": 4.95e-05, "output_cost_per_token_priority": 8.25e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -9488,7 +9493,7 @@ "output_cost_per_token_above_272k_tokens": 4.95e-05, "output_cost_per_token_priority": 8.25e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -9540,7 +9545,7 @@ "input_cost_per_token_priority": 1.25e-05, "input_cost_per_token_above_272k_tokens_priority": 2e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -9600,7 +9605,7 @@ "input_cost_per_token_above_272k_tokens_priority": 2e-05, "input_cost_per_token_flex": 2.5e-06, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -9659,7 +9664,7 @@ "output_cost_per_token_above_272k_tokens": 4.95e-05, "output_cost_per_token_priority": 8.25e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -9712,7 +9717,7 @@ "output_cost_per_token_above_272k_tokens": 4.95e-05, "output_cost_per_token_priority": 8.25e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -9766,7 +9771,7 @@ "output_cost_per_token_above_272k_tokens": 4.95e-05, "output_cost_per_token_priority": 8.25e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -9819,7 +9824,7 @@ "output_cost_per_token_above_272k_tokens": 4.95e-05, "output_cost_per_token_priority": 8.25e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -11100,12 +11105,13 @@ "input_cost_per_audio_token": 4.4e-05, "input_cost_per_token": 5.5e-06, "litellm_provider": "azure", - "max_input_tokens": 128000, + "max_input_tokens": 16000, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "realtime", "output_cost_per_audio_token": 8e-05, "output_cost_per_token": 2.2e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/foundry-models/concepts/models-sold-directly-by-azure", "supported_modalities": [ "text", "audio" @@ -11504,6 +11510,8 @@ }, "azure_ai/FLUX-1.1-pro": { "litellm_provider": "azure_ai", + "max_input_tokens": 5000, + "max_tokens": 5000, "mode": "image_generation", "output_cost_per_image": 0.04, "source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/black-forest-labs-flux-1-kontext-pro-and-flux1-1-pro-now-available-in-azure-ai-f/4434659", @@ -11513,6 +11521,8 @@ }, "azure_ai/FLUX.1-Kontext-pro": { "litellm_provider": "azure_ai", + "max_input_tokens": 5000, + "max_tokens": 5000, "mode": "image_generation", "output_cost_per_image": 0.04, "source": "https://marketplace.microsoft.com/pt-br/marketplace/apps/cohere.cohere-embed-4-offer?tab=PlansAndPrice", @@ -11922,8 +11932,8 @@ "input_cost_per_token": 2.5e-07, "litellm_provider": "azure_ai", "max_input_tokens": 1000000, - "max_output_tokens": 16384, - "max_tokens": 16384, + "max_output_tokens": 1000000, + "max_tokens": 1000000, "mode": "chat", "output_cost_per_token": 1e-06, "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", @@ -12462,9 +12472,9 @@ "input_cost_per_token": 2e-07, "output_cost_per_token": 5e-07, "litellm_provider": "azure_ai", - "max_input_tokens": 131072, - "max_output_tokens": 131072, - "max_tokens": 131072, + "max_input_tokens": 128000, + "max_output_tokens": 128000, + "max_tokens": 128000, "mode": "chat", "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", "supports_function_calling": true, @@ -12476,9 +12486,9 @@ "input_cost_per_token": 2e-07, "output_cost_per_token": 5e-07, "litellm_provider": "azure_ai", - "max_input_tokens": 131072, - "max_output_tokens": 131072, - "max_tokens": 131072, + "max_input_tokens": 128000, + "max_output_tokens": 128000, + "max_tokens": 128000, "mode": "chat", "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", "supports_function_calling": true, @@ -70560,6 +70570,9 @@ "input_cost_per_token": 5.5e-06, "input_cost_per_token_batches": 2.75e-06, "litellm_provider": "azure", + "max_input_tokens": 128000, + "max_output_tokens": 4096, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1.65e-05, "output_cost_per_token_batches": 8.25e-06, @@ -70850,6 +70863,9 @@ "input_cost_per_token": 2.2e-06, "input_cost_per_token_batches": 1.1e-06, "litellm_provider": "azure", + "max_input_tokens": 200000, + "max_output_tokens": 100000, + "max_tokens": 100000, "mode": "chat", "output_cost_per_token": 8.8e-06, "output_cost_per_token_batches": 4.4e-06, @@ -70870,6 +70886,9 @@ "input_cost_per_token": 1.21e-06, "input_cost_per_token_batches": 6.05e-07, "litellm_provider": "azure", + "max_input_tokens": 200000, + "max_output_tokens": 100000, + "max_tokens": 100000, "mode": "chat", "output_cost_per_token": 4.84e-06, "output_cost_per_token_batches": 2.42e-06, @@ -70998,6 +71017,9 @@ "input_cost_per_token": 5.5e-06, "input_cost_per_token_batches": 2.75e-06, "litellm_provider": "azure", + "max_input_tokens": 128000, + "max_output_tokens": 4096, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1.65e-05, "output_cost_per_token_batches": 8.25e-06, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 99471d48f56..9989846b7e0 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -3977,7 +3977,7 @@ "input_cost_per_token_priority": 1e-05, "input_cost_per_token_above_272k_tokens_priority": 2e-05, "litellm_provider": "azure_ai", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -4112,7 +4112,7 @@ "input_cost_per_token_priority": 1e-05, "input_cost_per_token_above_272k_tokens_priority": 2e-05, "litellm_provider": "azure_ai", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -4159,7 +4159,7 @@ "input_cost_per_token_priority": 1e-05, "input_cost_per_token_above_272k_tokens_priority": 2e-05, "litellm_provider": "azure_ai", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -4206,7 +4206,7 @@ "input_cost_per_token_priority": 5e-06, "input_cost_per_token_above_272k_tokens_priority": 1e-05, "litellm_provider": "azure_ai", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -4254,7 +4254,7 @@ "input_cost_per_token_priority": 5e-06, "input_cost_per_token_above_272k_tokens_priority": 1e-05, "litellm_provider": "azure_ai", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -4302,7 +4302,7 @@ "input_cost_per_token_priority": 6e-05, "input_cost_per_token_above_272k_tokens_priority": 0.00012, "litellm_provider": "azure_ai", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", @@ -4348,7 +4348,7 @@ "input_cost_per_token_priority": 6e-05, "input_cost_per_token_above_272k_tokens_priority": 0.00012, "litellm_provider": "azure_ai", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", @@ -4678,12 +4678,13 @@ "input_cost_per_audio_token": 4.4e-05, "input_cost_per_token": 5.5e-06, "litellm_provider": "azure", - "max_input_tokens": 128000, + "max_input_tokens": 16000, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "realtime", "output_cost_per_audio_token": 8e-05, "output_cost_per_token": 2.2e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/foundry-models/concepts/models-sold-directly-by-azure", "supported_modalities": [ "text", "audio" @@ -6319,12 +6320,13 @@ "input_cost_per_audio_token": 4e-05, "input_cost_per_token": 5e-06, "litellm_provider": "azure", - "max_input_tokens": 128000, + "max_input_tokens": 16000, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "realtime", "output_cost_per_audio_token": 8e-05, "output_cost_per_token": 2e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/foundry-models/concepts/models-sold-directly-by-azure", "supported_modalities": [ "text", "audio" @@ -6370,6 +6372,9 @@ "deprecation_date": "2027-05-06", "input_cost_per_second": 0.0002833333333333333, "litellm_provider": "azure", + "max_input_tokens": 32000, + "max_output_tokens": 4096, + "max_tokens": 4096, "mode": "audio_transcription", "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/gpt-realtime-whisper", "supported_endpoints": [ @@ -7424,7 +7429,7 @@ "input_cost_per_token_priority": 5e-06, "input_cost_per_token_above_272k_tokens_priority": 1e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -7480,7 +7485,7 @@ "output_cost_per_token": 1.65e-05, "output_cost_per_token_priority": 3.3e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -7530,7 +7535,7 @@ "output_cost_per_token": 1.65e-05, "output_cost_per_token_priority": 3.3e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -7580,7 +7585,7 @@ "input_cost_per_token_priority": 5e-06, "input_cost_per_token_above_272k_tokens_priority": 1e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -7636,7 +7641,7 @@ "output_cost_per_token": 1.65e-05, "output_cost_per_token_priority": 3.3e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -7686,7 +7691,7 @@ "output_cost_per_token": 1.65e-05, "output_cost_per_token_priority": 3.3e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -7737,7 +7742,7 @@ "input_cost_per_token_batches": 1.5e-05, "input_cost_per_token_flex": 1.5e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", @@ -7786,7 +7791,7 @@ "input_cost_per_token_batches": 1.5e-05, "input_cost_per_token_flex": 1.5e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", @@ -9374,7 +9379,7 @@ "input_cost_per_token_batches": 2.5e-06, "input_cost_per_token_flex": 2.5e-06, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -9433,7 +9438,7 @@ "output_cost_per_token_above_272k_tokens": 4.95e-05, "output_cost_per_token_priority": 8.25e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -9488,7 +9493,7 @@ "output_cost_per_token_above_272k_tokens": 4.95e-05, "output_cost_per_token_priority": 8.25e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -9540,7 +9545,7 @@ "input_cost_per_token_priority": 1.25e-05, "input_cost_per_token_above_272k_tokens_priority": 2e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -9600,7 +9605,7 @@ "input_cost_per_token_above_272k_tokens_priority": 2e-05, "input_cost_per_token_flex": 2.5e-06, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -9659,7 +9664,7 @@ "output_cost_per_token_above_272k_tokens": 4.95e-05, "output_cost_per_token_priority": 8.25e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -9712,7 +9717,7 @@ "output_cost_per_token_above_272k_tokens": 4.95e-05, "output_cost_per_token_priority": 8.25e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -9766,7 +9771,7 @@ "output_cost_per_token_above_272k_tokens": 4.95e-05, "output_cost_per_token_priority": 8.25e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -9819,7 +9824,7 @@ "output_cost_per_token_above_272k_tokens": 4.95e-05, "output_cost_per_token_priority": 8.25e-05, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -11100,12 +11105,13 @@ "input_cost_per_audio_token": 4.4e-05, "input_cost_per_token": 5.5e-06, "litellm_provider": "azure", - "max_input_tokens": 128000, + "max_input_tokens": 16000, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "realtime", "output_cost_per_audio_token": 8e-05, "output_cost_per_token": 2.2e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/foundry-models/concepts/models-sold-directly-by-azure", "supported_modalities": [ "text", "audio" @@ -11504,6 +11510,8 @@ }, "azure_ai/FLUX-1.1-pro": { "litellm_provider": "azure_ai", + "max_input_tokens": 5000, + "max_tokens": 5000, "mode": "image_generation", "output_cost_per_image": 0.04, "source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/black-forest-labs-flux-1-kontext-pro-and-flux1-1-pro-now-available-in-azure-ai-f/4434659", @@ -11513,6 +11521,8 @@ }, "azure_ai/FLUX.1-Kontext-pro": { "litellm_provider": "azure_ai", + "max_input_tokens": 5000, + "max_tokens": 5000, "mode": "image_generation", "output_cost_per_image": 0.04, "source": "https://marketplace.microsoft.com/pt-br/marketplace/apps/cohere.cohere-embed-4-offer?tab=PlansAndPrice", @@ -11922,8 +11932,8 @@ "input_cost_per_token": 2.5e-07, "litellm_provider": "azure_ai", "max_input_tokens": 1000000, - "max_output_tokens": 16384, - "max_tokens": 16384, + "max_output_tokens": 1000000, + "max_tokens": 1000000, "mode": "chat", "output_cost_per_token": 1e-06, "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", @@ -12462,9 +12472,9 @@ "input_cost_per_token": 2e-07, "output_cost_per_token": 5e-07, "litellm_provider": "azure_ai", - "max_input_tokens": 131072, - "max_output_tokens": 131072, - "max_tokens": 131072, + "max_input_tokens": 128000, + "max_output_tokens": 128000, + "max_tokens": 128000, "mode": "chat", "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", "supports_function_calling": true, @@ -12476,9 +12486,9 @@ "input_cost_per_token": 2e-07, "output_cost_per_token": 5e-07, "litellm_provider": "azure_ai", - "max_input_tokens": 131072, - "max_output_tokens": 131072, - "max_tokens": 131072, + "max_input_tokens": 128000, + "max_output_tokens": 128000, + "max_tokens": 128000, "mode": "chat", "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", "supports_function_calling": true, @@ -70560,6 +70570,9 @@ "input_cost_per_token": 5.5e-06, "input_cost_per_token_batches": 2.75e-06, "litellm_provider": "azure", + "max_input_tokens": 128000, + "max_output_tokens": 4096, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1.65e-05, "output_cost_per_token_batches": 8.25e-06, @@ -70850,6 +70863,9 @@ "input_cost_per_token": 2.2e-06, "input_cost_per_token_batches": 1.1e-06, "litellm_provider": "azure", + "max_input_tokens": 200000, + "max_output_tokens": 100000, + "max_tokens": 100000, "mode": "chat", "output_cost_per_token": 8.8e-06, "output_cost_per_token_batches": 4.4e-06, @@ -70870,6 +70886,9 @@ "input_cost_per_token": 1.21e-06, "input_cost_per_token_batches": 6.05e-07, "litellm_provider": "azure", + "max_input_tokens": 200000, + "max_output_tokens": 100000, + "max_tokens": 100000, "mode": "chat", "output_cost_per_token": 4.84e-06, "output_cost_per_token_batches": 2.42e-06, @@ -70998,6 +71017,9 @@ "input_cost_per_token": 5.5e-06, "input_cost_per_token_batches": 2.75e-06, "litellm_provider": "azure", + "max_input_tokens": 128000, + "max_output_tokens": 4096, + "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1.65e-05, "output_cost_per_token_batches": 8.25e-06, From ffb15f946f586c102bf0359b2fa9ac46b340b658 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 16:42:13 -0700 Subject: [PATCH 39/41] perf(proxy): one request-scoped Redis pipeline for auth, spend, rate-limit and routing reads (#43407) RedisBatch: one pipeline per Redis backend for independently declared operations (MGET, GET, Lua scripts, INCRBYFLOAT, SET, DEL), a future per operation so each owner keeps its own fallback, Redis Cluster hash-slot fallback. A request-scoped batch middleware shares that pipeline across the auth identity reads and write-back, the spend counter MGET, the rate limiter Lua groups and the routing read. A rate-limit denial stands when another pipelined group fails; every pipelined group is refunded on rejection; local cooldowns win over the prefetch. The routing prefetch failure log line strips request line breaks (CodeQL py/log-injection) Resolves LIT-8882 Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/caching/redis_batch.py | 422 ++++++++++++++ litellm/proxy/auth/auth_object_prefetch.py | 23 +- litellm/proxy/common_request_processing.py | 3 + .../hooks/parallel_request_limiter_v3.py | 146 ++++- .../redis_request_batch_middleware.py | 25 + litellm/proxy/proxy_server.py | 2 + .../spend_tracking/spend_counter_batch.py | 37 +- litellm/router.py | 25 +- litellm/router_utils/routing_read_batch.py | 114 +++- .../router_code_coverage.py | 1 + .../hooks/test_parallel_request_limiter_v3.py | 33 ++ tests/unit/caching/test_redis_batch.py | 299 ++++++++++ .../test_request_redis_batch_pre_call.py | 530 ++++++++++++++++++ tests/unit/test_router/test_router.py | 34 ++ 14 files changed, 1667 insertions(+), 27 deletions(-) create mode 100644 litellm/caching/redis_batch.py create mode 100644 litellm/proxy/middleware/redis_request_batch_middleware.py create mode 100644 tests/unit/caching/test_redis_batch.py create mode 100644 tests/unit/caching/test_request_redis_batch_pre_call.py diff --git a/litellm/caching/redis_batch.py b/litellm/caching/redis_batch.py new file mode 100644 index 00000000000..8bd9554ec55 --- /dev/null +++ b/litellm/caching/redis_batch.py @@ -0,0 +1,422 @@ +"""One Redis pipeline for several independent operations, each with its own result and its own failure. + +A ``RedisBatch`` collects MGETs, Lua scripts and increments declared by unrelated callers and sends them +in one ``pipeline(transaction=False)`` round trip. Every declaration returns an awaitable; awaiting one +flushes whatever has been declared so far, so callers keep their existing ``await`` shape and their own +error handling while sharing the wire. Redis Cluster clients run each operation on its own, as before: +a cluster pipeline is per node anyway and the existing per-operation paths already group by slot. +""" + +from __future__ import annotations + +import asyncio +import hashlib +import json +import logging +import time +from collections.abc import Awaitable, Callable, Generator, Mapping, Sequence +from contextvars import ContextVar, Token +from dataclasses import dataclass, field +from datetime import timedelta +from types import MappingProxyType, TracebackType +from typing import Final, Generic, Protocol, TypeVar + +from litellm._logging import verbose_logger +from litellm.caching.redis_cache import ( + RedisCache, + _run_under_circuit_breaker, # pyright: ignore[reportPrivateUsage] # same health signal as every RedisCache method + log_redis_failure, +) +from litellm.caching.redis_cluster_cache import RedisClusterCache +from litellm.types.services import ServiceTypes + +_T = TypeVar("_T") +_ScriptArg = str | bytes | int | float + + +class RegisteredScript(Protocol): + def __call__(self, keys: Sequence[str], args: Sequence[_ScriptArg]) -> Awaitable[object]: ... + + +class _RedisPipeline(Protocol): + def mget(self, keys: Sequence[str]) -> object: ... + def evalsha(self, sha: str, numkeys: int, *keys_and_args: _ScriptArg) -> object: ... + def incrbyfloat(self, name: str, amount: float) -> object: ... + def expire(self, name: str, time: timedelta) -> object: ... + def set(self, name: str, value: str, ex: timedelta | None = None) -> object: ... + async def execute(self, raise_on_error: bool = True) -> list[object]: ... + + +class _Op(Generic[_T]): + """One declared operation: how many pipeline replies it consumes, how to turn them into a result, and + how to run on its own when the batch cannot pipeline (cluster client, or a reply the pipeline cannot + settle, like NOSCRIPT).""" + + __slots__ = ("future",) + + def __init__(self) -> None: + self.future: Final[asyncio.Future[_T]] = asyncio.get_running_loop().create_future() + self.future.add_done_callback(_mark_retrieved) + + def enqueue(self, pipe: _RedisPipeline) -> int: + raise NotImplementedError + + def resolve(self, replies: Sequence[object]) -> _T: + raise NotImplementedError + + async def run_alone(self) -> _T: + raise NotImplementedError + + def settle(self, replies: Sequence[object]) -> Awaitable[None] | None: + """Resolve from pipeline replies; return a coroutine when the op has to be retried on its own.""" + failure: Final = next((reply for reply in replies if isinstance(reply, Exception)), None) + if failure is None: + try: + self.future.set_result(self.resolve(replies)) + except Exception as e: # noqa: BLE001 # a reply this op cannot decode fails this op alone + self.future.set_exception(e) + return None + if _is_missing_script(failure): + return self._settle_alone() + self.future.set_exception(failure) + return None + + async def _settle_alone(self) -> None: + try: + self.future.set_result(await self.run_alone()) + except Exception as e: # noqa: BLE001 # the declaring caller owns the failure of its own operation + self.future.set_exception(e) + + +def _is_missing_script(failure: Exception) -> bool: + """Imported lazily: this module is reachable from a base ``import litellm`` while redis is not a base dependency.""" + from redis.exceptions import NoScriptError + + return isinstance(failure, NoScriptError) + + +def _mark_retrieved(future: asyncio.Future[object]) -> None: + """A caller that stops awaiting (cancelled request) must not leave an 'exception never retrieved' log.""" + if not future.cancelled(): + future.exception() + + +class _MGet(_Op[Mapping[str, object]]): + __slots__ = ("_keys", "_redis_cache") + + def __init__(self, redis_cache: RedisCache, keys: Sequence[str]) -> None: + super().__init__() + self._redis_cache: Final = redis_cache + self._keys: Final[tuple[str, ...]] = tuple(dict.fromkeys(keys)) + + def enqueue(self, pipe: _RedisPipeline) -> int: + pipe.mget(tuple(self._redis_cache.check_and_fix_namespace(key=key) for key in self._keys)) + return 1 + + def resolve(self, replies: Sequence[object]) -> Mapping[str, object]: + values: Final = replies[0] + if not isinstance(values, (list, tuple)): + raise TypeError(f"MGET reply is not a list: {type(values).__name__}") + return MappingProxyType( + {key: self._redis_cache._get_cache_logic(value) for key, value in zip(self._keys, values)} # pyright: ignore[reportPrivateUsage, reportUnknownMemberType, reportUnknownArgumentType] # shared decode with async_batch_get_cache + ) + + async def run_alone(self) -> Mapping[str, object]: + found: Mapping[str, object] = await self._redis_cache.async_batch_get_cache(key_list=list(self._keys)) # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # untyped cache API # mutable-ok: the cache API takes a list + if any(key not in found for key in self._keys): + raise ConnectionError("batch get did not return every key") + return found + + +class _Script(_Op[object]): + __slots__ = ("_args", "_keys", "_redis_cache", "_run", "_sha") + + def __init__( + self, + redis_cache: RedisCache, + source: str, + run: RegisteredScript, + keys: Sequence[str], + args: Sequence[_ScriptArg], + ) -> None: + super().__init__() + self._redis_cache: Final = redis_cache + self._sha: Final = hashlib.sha1(source.encode()).hexdigest() # noqa: S324 # EVALSHA identifies scripts by SHA-1 + self._run: Final = run + self._keys: Final[tuple[str, ...]] = tuple(keys) + self._args: Final[tuple[_ScriptArg, ...]] = tuple(args) + + def enqueue(self, pipe: _RedisPipeline) -> int: + namespaced: Final = tuple(self._redis_cache.check_and_fix_namespace(key=key) for key in self._keys) + pipe.evalsha(self._sha, len(namespaced), *namespaced, *self._args) + return 1 + + def resolve(self, replies: Sequence[object]) -> object: + return replies[0] + + async def run_alone(self) -> object: + return await self._run(keys=self._keys, args=self._args) + + +class _Increment(_Op[float]): + __slots__ = ("_key", "_redis_cache", "_ttl", "_value") + + def __init__(self, redis_cache: RedisCache, key: str, value: float, ttl: int | None) -> None: + super().__init__() + self._redis_cache: Final = redis_cache + self._key: Final = key + self._value: Final = value + self._ttl: Final = ttl + + def enqueue(self, pipe: _RedisPipeline) -> int: + name: Final = self._redis_cache.check_and_fix_namespace(key=self._key) + pipe.incrbyfloat(name, self._value) + if self._ttl is None: + return 1 + pipe.expire(name, timedelta(seconds=self._ttl)) + return 2 + + def resolve(self, replies: Sequence[object]) -> float: + reply: Final = replies[0] + if not isinstance(reply, (int, float, str, bytes)): + raise TypeError(f"INCRBYFLOAT reply is not numeric: {type(reply).__name__}") + return float(reply) + + async def run_alone(self) -> float: + value: object = await self._redis_cache.async_increment(key=self._key, value=self._value, ttl=self._ttl) # pyright: ignore[reportUnknownMemberType] # untyped cache API + if not isinstance(value, (int, float)): + raise TypeError(f"increment did not return a number: {type(value).__name__}") + return float(value) + + +class _Set(_Op[None]): + """SET with the cache's TTL rules, same encoding as ``async_set_cache_pipeline_with_ttls``.""" + + __slots__ = ("_key", "_redis_cache", "_ttl", "_value") + + def __init__(self, redis_cache: RedisCache, key: str, value: object, ttl: float | None) -> None: + super().__init__() + self._redis_cache: Final = redis_cache + self._key: Final = key + self._value: Final = value + self._ttl: Final = ttl + + def enqueue(self, pipe: _RedisPipeline) -> int: + ttl: Final = self._redis_cache.get_ttl(ttl=self._ttl) # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # untyped cache API + pipe.set( + self._redis_cache.check_and_fix_namespace(key=self._key), + json.dumps(self._value), + ex=None if ttl is None else timedelta(seconds=ttl), + ) + return 1 + + def resolve(self, replies: Sequence[object]) -> None: + return None + + async def run_alone(self) -> None: + await self._redis_cache.async_set_cache_pipeline_with_ttls(((self._key, self._value, self._ttl),)) + + +class BatchResult(Generic[_T]): + """Awaitable handle for one declared operation; awaiting it flushes the batch it belongs to.""" + + __slots__ = ("_batch", "_op") + + def __init__(self, batch: RedisBatch, op: _Op[_T]) -> None: + self._batch: Final = batch + self._op: Final = op + + def __await__(self) -> Generator[object, None, _T]: + return self._wait().__await__() + + async def _wait(self) -> _T: + if not self._op.future.done(): + await self._batch.flush() + return self._op.future.result() + + @property + def done(self) -> bool: + return self._op.future.done() + + +@dataclass(slots=True) +class RedisBatch: + """Operations declared here go out in one pipeline the next time any of them is awaited or ``flush`` runs.""" + + redis_cache: RedisCache + name: str = "redis_batch" + _pending: list[_Op[object]] = field(default_factory=list) # mutable-ok: drained by flush + _flush_hooks: list[Callable[[], None]] = field(default_factory=list) # mutable-ok: append-only registry + _lock: asyncio.Lock = field(default_factory=asyncio.Lock) + flushes: int = 0 + + def mget(self, keys: Sequence[str]) -> BatchResult[Mapping[str, object]]: + return self._declare(_MGet(self.redis_cache, keys)) + + def script( + self, source: str, run: RegisteredScript, keys: Sequence[str], args: Sequence[_ScriptArg] + ) -> BatchResult[object]: + return self._declare(_Script(self.redis_cache, source, run, keys, args)) + + def increment(self, key: str, value: float, ttl: int | None = None) -> BatchResult[float]: + return self._declare(_Increment(self.redis_cache, key, value, ttl)) + + def set(self, key: str, value: object, ttl: float | None = None) -> BatchResult[None]: + return self._declare(_Set(self.redis_cache, key, value, ttl)) + + def add_flush_hook(self, hook: Callable[[], None]) -> None: + """Called at the start of every flush so lazily bound readers can declare their keys into the same trip.""" + self._flush_hooks.append(hook) + + @property + def pending(self) -> int: + return len(self._pending) + + def _declare(self, op: _Op[_T]) -> BatchResult[_T]: + self._pending.append(op) # pyright: ignore[reportArgumentType] # heterogeneous ops share the flush loop + return BatchResult(self, op) + + async def flush(self) -> None: + async with self._lock: + for hook in self._flush_hooks: + hook() + ops: Final = tuple(self._pending) + self._pending.clear() + if not ops: + return + self.flushes += 1 + try: + if isinstance(self.redis_cache, RedisClusterCache): + await asyncio.gather(*(op._settle_alone() for op in ops)) # pyright: ignore[reportPrivateUsage] # batch owns its ops + else: + await self._flush_pipeline(ops) + finally: + for op in ops: + if not op.future.done(): + op.future.cancel() + + async def _flush_pipeline(self, ops: Sequence[_Op[object]]) -> None: + start_time: Final = time.time() + widths: list[int] = [] # mutable-ok: filled while enqueuing + + async def run() -> list[object]: + client: Final = self.redis_cache.init_async_client() + async with client.pipeline(transaction=False) as pipe: + widths.extend(op.enqueue(pipe) for op in ops) + return await pipe.execute(raise_on_error=False) + + try: + replies: Final = await _run_under_circuit_breaker(self.redis_cache._circuit_breaker, self.name, run) # pyright: ignore[reportPrivateUsage] # same breaker as the cache's own methods + except Exception as e: # noqa: BLE001 # each declaring caller applies its own Redis fallback + log_redis_failure(verbose_logger, logging.WARNING, f"{self.name}: pipeline of {len(ops)} ops failed", e) + asyncio.create_task( + self.redis_cache.service_logger_obj.async_service_failure_hook( + service=ServiceTypes.REDIS, + duration=time.time() - start_time, + error=e, + call_type=f"{self.name}[{len(ops)}]", + start_time=start_time, + end_time=time.time(), + ) + ) + for op in ops: + op.future.set_exception(e) + return + asyncio.create_task( + self.redis_cache.service_logger_obj.async_service_success_hook( + service=ServiceTypes.REDIS, + duration=time.time() - start_time, + call_type=f"{self.name}[{len(ops)}]", + start_time=start_time, + end_time=time.time(), + ) + ) + retries: list[Awaitable[None]] = [] # mutable-ok: collected while slicing replies + offset = 0 + for op, width in zip(ops, widths): + retry = op.settle(replies[offset : offset + width]) + offset += width + if retry is not None: + retries.append(retry) + if retries: + await asyncio.gather(*retries) + + +def _backend_key(redis_cache: RedisCache) -> object: + """Two ``RedisCache`` instances built from the same connection settings and namespace talk to the same server + under the same key prefix, so the proxy's cache and the router's cache share one pipeline (the router gets its + port as a string, hence the ``str`` comparison); a cache whose settings cannot be compared (a test double) gets + its own.""" + try: + settings: Final = tuple(sorted((str(k), str(v)) for k, v in redis_cache.redis_kwargs.items() if v is not None)) # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType, reportUnknownArgumentType] # untyped cache API + except AttributeError: + return ("instance", id(redis_cache)) + return (type(redis_cache), redis_cache.namespace, settings) + + +class RequestRedisBatches: + """One ``RedisBatch`` per Redis backend for the current request, so readers of different caches that + share a server (the proxy's and the router's) share the pipeline.""" + + __slots__ = ("_batches", "prefetched") + + def __init__(self) -> None: + self._batches: Final[dict[object, RedisBatch]] = {} # mutable-ok: lazily filled per backend + # Reads declared early for a consumer that runs later in the request, keyed by consumer name. + self.prefetched: Final[dict[str, object]] = {} # mutable-ok: armed pre-admission, taken at use + + def batch(self, redis_cache: RedisCache) -> RedisBatch: + key: Final = _backend_key(redis_cache) + batch = self._batches.get(key) + if batch is None: + batch = RedisBatch(redis_cache, name="request_redis_batch") + self._batches[key] = batch + return batch + + async def flush_all(self) -> None: + """Send whatever is still declared (write-backs nobody awaits) before the request scope closes.""" + await asyncio.gather(*(batch.flush() for batch in self._batches.values() if batch.pending)) + + @property + def batches(self) -> tuple[RedisBatch, ...]: + return tuple(self._batches.values()) + + +_active_request_batches: Final[ContextVar[RequestRedisBatches | None]] = ContextVar( + "request_redis_batches", default=None +) + + +def active_request_redis_batch(redis_cache: RedisCache) -> RedisBatch | None: + """The request's batch for this backend, or None outside a ``request_redis_batch_scope``.""" + batches: Final = _active_request_batches.get() + if batches is None: + return None + return batches.batch(redis_cache) + + +def active_request_redis_batches() -> RequestRedisBatches | None: + return _active_request_batches.get() + + +class request_redis_batch_scope: + """Redis reads declared inside share one pipeline per backend; nested scopes join the outer one.""" + + __slots__ = ("_token",) + + def __init__(self) -> None: + self._token: Token[RequestRedisBatches | None] | None = None + + def __enter__(self) -> RequestRedisBatches: + outer: Final = _active_request_batches.get() + if outer is not None: + return outer + batches: Final = RequestRedisBatches() + self._token = _active_request_batches.set(batches) + return batches + + def __exit__( + self, exc_type: type[BaseException] | None, exc: BaseException | None, tb: TracebackType | None + ) -> None: + if self._token is not None: + _active_request_batches.reset(self._token) diff --git a/litellm/proxy/auth/auth_object_prefetch.py b/litellm/proxy/auth/auth_object_prefetch.py index 52e26e885c9..ce55190aa02 100644 --- a/litellm/proxy/auth/auth_object_prefetch.py +++ b/litellm/proxy/auth/auth_object_prefetch.py @@ -13,6 +13,7 @@ from typing import Final, Literal, Protocol, TypeAlias from pydantic import BaseModel, TypeAdapter, ValidationError from litellm._logging import verbose_proxy_logger +from litellm.caching.redis_batch import active_request_redis_batch from litellm.caching.redis_cache import RedisCache from litellm.constants import DEFAULT_IN_MEMORY_TTL from litellm.models.organization import LiteLLM_OrganizationTable @@ -218,11 +219,23 @@ def _set_in_memory(memory: _InMemoryCache, cache_key: str, value: object, ttl: f memory.set_cache(key=cache_key, value=value, ttl=ttl) +async def _read_redis_rows(keys: list[str], redis_cache: RedisCache) -> Mapping[str, object]: + """On the request pipeline when one is open; a failed pipeline reads as a miss, like ``async_batch_get_cache``.""" + batch: Final = active_request_redis_batch(redis_cache) + if batch is None: + return await redis_cache.async_batch_get_cache(key_list=keys) # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # untyped cache API + try: + return await batch.mget(keys) + except Exception as e: # noqa: BLE001 # the DB fill below takes over, as it does after a failed MGET today + verbose_proxy_logger.debug("auth prefetch Redis read failed, filling from the database: %s", e) + return MappingProxyType({}) + + async def _fill_from_redis(entries: Sequence[_CacheEntry], redis_cache: RedisCache, memory: _InMemoryCache) -> None: if not entries: return found: Final = _RowValues.validate_python( - await redis_cache.async_batch_get_cache(key_list=sorted(entry.cache_key for entry in entries)) # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # untyped cache API + await _read_redis_rows(sorted(entry.cache_key for entry in entries), redis_cache) ) for entry, value in ((entry, found.get(entry.cache_key)) for entry in entries): if value is not None: @@ -267,8 +280,14 @@ async def _write_back(entries: Sequence[tuple[_CacheEntry, BaseModel]], cache: U memory: Final[_InMemoryCache] = cache.in_memory_cache for cache_key, payload, ttl in payloads: _set_in_memory(memory, cache_key, payload, cache.default_in_memory_ttl if ttl is None else ttl) - if cache.redis_cache is not None: + if cache.redis_cache is None: + return + batch: Final = active_request_redis_batch(cache.redis_cache) + if batch is None: await cache.redis_cache.async_set_cache_pipeline_with_ttls(payloads) + return + for cache_key, payload, ttl in payloads: # rides the request's next round trip; the scope drains leftovers + batch.set(cache_key, payload, ttl) async def _fill_from_db( diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 97de2488b8d..10724e9e7e6 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -2199,6 +2199,9 @@ class ProxyBaseLLMRequestProcessing: if self._tags_before_guardrails is None: self._tags_before_guardrails = frozenset(get_tags_from_request_body(request_body=self.data)) + prefetch_model = self.data.get("model") + if llm_router is not None and isinstance(prefetch_model, str): + llm_router.arm_routing_read_prefetch(prefetch_model, self.data) self.data = await proxy_logging_obj.pre_call_hook( user_api_key_dict=user_api_key_dict, data=self.data, diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index e4b782d5ff3..589aa7da7b5 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -33,6 +33,7 @@ from typing_extensions import NotRequired, ReadOnly from litellm import DualCache from litellm._logging import verbose_proxy_logger +from litellm.caching.redis_batch import BatchResult, RegisteredScript, active_request_redis_batch from litellm.caching.redis_cache import log_redis_failure from litellm.constants import DYNAMIC_RATE_LIMIT_ERROR_THRESHOLD_PER_MINUTE, INTERNAL_CALL_ORIGIN_METADATA_KEY from litellm.integrations.custom_logger import CustomLogger @@ -474,6 +475,19 @@ CacheCounterValue: TypeAlias = int | float | str | bytes CacheCounterValues: TypeAlias = Sequence[CacheCounterValue | None] + +def _as_counter_values(reply: object) -> list[CacheCounterValue]: + """A Lua reply read back off the pipeline is the same array the script returns when called directly.""" + if not isinstance(reply, (list, tuple)): + raise TypeError(f"rate limiter script reply is not a list: {type(reply).__name__}") + values: Final[list[CacheCounterValue]] = [] # mutable-ok: each element is narrowed before it is kept + for value in reply: # pyright: ignore[reportUnknownVariableType] # raw Redis reply + if not isinstance(value, (int, float, str, bytes)): + raise TypeError(f"rate limiter script reply holds {type(value).__name__}") # pyright: ignore[reportUnknownArgumentType] # raw Redis reply + values.append(value) + return values + + ReservationWindowIdentity: TypeAlias = tuple[str, str, Literal["redis", "local"]] ParallelGaugeCacheValue: TypeAlias = dict[str, object] | int | float | str | bytes @@ -1323,6 +1337,21 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): crc: Final = binascii.crc_hqx(key.encode("utf-8"), 0) return crc % REDIS_CLUSTER_SLOTS + def _pipeline_scripts( + self, + source: str, + run: RegisteredScript, + calls: Sequence[tuple[Sequence[str], Sequence[int]]], + ) -> tuple[BatchResult[object] | None, ...]: + """Declare one Lua call per group on the request's Redis batch, so all groups share one round trip + with whatever else the request declared (the routing read). Returns ``None`` per call when no batch + is open, and the caller runs the script directly as before.""" + redis_cache: Final = self.internal_usage_cache.dual_cache.redis_cache + batch: Final = None if redis_cache is None else active_request_redis_batch(redis_cache) + if batch is None: + return (None,) * len(calls) + return tuple(batch.script(source, run, keys, args) for keys, args in calls) + def _group_keys_by_hash_tag(self, keys: list[str]) -> dict[str, list[str]]: """ Group keys by their Redis hash tag to ensure cluster compatibility. @@ -1404,7 +1433,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) return await self._batch_get_counter_values(keys=keys, parent_otel_span=parent_otel_span, local_only=True) - def _reject_if_rate_limit_unverifiable(self, failed_operation: str, error: Exception) -> None: + def _reject_if_rate_limit_unverifiable(self, failed_operation: str, error: BaseException) -> None: if not self._fail_closed_resolver(): return log_redis_failure( @@ -1436,12 +1465,19 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): key_groups: Final = list(self._group_keys_by_hash_tag(keys_to_fetch).items()) all_cache_values: Final[list[CacheCounterValue | None]] = [] + args: Final = (now_int, self.window_size) + pipelined: Final = self._pipeline_scripts( + BATCH_RATE_LIMITER_SCRIPT, + self.batch_rate_limiter_script, + tuple((group_keys, args) for _tag, group_keys in key_groups), + ) - for index, (hash_tag, group_keys) in enumerate(key_groups): + for index, ((hash_tag, group_keys), group_result) in enumerate(zip(key_groups, pipelined)): try: - group_cache_values: CacheCounterValues = await self.batch_rate_limiter_script( - keys=group_keys, - args=[now_int, self.window_size], # Use integer timestamp + group_cache_values: CacheCounterValues = ( + await self.batch_rate_limiter_script(keys=group_keys, args=args) + if group_result is None + else _as_counter_values(await group_result) ) all_cache_values.extend(group_cache_values) except Exception as e: @@ -1450,6 +1486,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): await self._refund_counter_increments( self._counter_refunds_from_batch_values(applied_keys, all_cache_values) ) + await self._refund_later_pipelined_groups(key_groups[index + 1 :], pipelined[index + 1 :]) self._reject_if_rate_limit_unverifiable("batch_rate_limiter_script", e) log_redis_failure( verbose_proxy_logger, logging.WARNING, f"Redis Lua script failed for hash tag {hash_tag}", e @@ -1464,6 +1501,22 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return all_cache_values + async def _refund_later_pipelined_groups( + self, + key_groups: Sequence[tuple[str, list[str]]], + pipelined: Sequence[BatchResult[object] | None], + ) -> None: + """Groups declared on the request batch ran in the same round trip as the one that failed, so their + increments landed even though the loop never read them.""" + for (_tag, group_keys), group_result in zip(key_groups, pipelined): + if group_result is None: + continue + try: + group_values = _as_counter_values(await group_result) + except Exception: # noqa: BLE001 # a group that failed in Redis incremented nothing to refund + continue + await self._refund_counter_increments(self._counter_refunds_from_batch_values(group_keys, group_values)) + async def should_rate_limit( self, descriptors: Sequence[RateLimitDescriptor], @@ -2061,7 +2114,16 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): reservation_windows: Final[set[ReservationWindowIdentity]] = set() # mutable-ok: filled by the group loop raw: list[CacheCounterValue] - for _idx, (keys, args, meta) in enumerate(descriptor_groups): + pipelined: Final = self._pipeline_scripts( + CHECK_AND_INCREMENT_BY_N_SCRIPT, + self.check_and_increment_by_n_script, # pyright: ignore[reportArgumentType] # sole caller guards it is not None + tuple((keys, args) for keys, args, _meta in descriptor_groups), + ) + batched: Final = tuple(result for result in pipelined if result is not None) + if len(batched) == len(descriptor_groups): + return await self._settle_pipelined_descriptor_groups(descriptor_groups, batched, parent_otel_span) + + for keys, args, meta in descriptor_groups: try: raw = await self.check_and_increment_by_n_script( # pyright: ignore[reportOptionalCall] # sole caller guards it is not None keys=keys, @@ -2105,6 +2167,76 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): reservation_windows=frozenset(reservation_windows), ) + async def _settle_pipelined_descriptor_groups( + self, + descriptor_groups: list[DescriptorAtomicGroup], + results: Sequence[BatchResult[object]], + parent_otel_span: Span | None, + ) -> RateLimitResponse: + """Every group's Lua call left in one pipeline, so each group has already checked and incremented on + its own before any result is read. A failed or over-limit group therefore refunds every group that + incremented, after it as well as before it, where the one-at-a-time loop only unwinds the groups it ran. + A Redis denial stands even when another group failed: the in-memory fallback only replaces a verdict + Redis never gave.""" + replies: Final = await asyncio.gather(*results, return_exceptions=True) + responses: Final = tuple( + self._pipelined_group_response(reply, meta) + for reply, (_keys, _args, meta) in zip(replies, descriptor_groups) + ) + applied: Final[list[tuple[CounterRefund, ...]]] = [] # mutable-ok: filled by the group loop + statuses: Final[list[RateLimitStatus]] = [] # mutable-ok: filled by the group loop + reservation_windows: Final[set[ReservationWindowIdentity]] = set() # mutable-ok: filled by the group loop + for reply, response, (_keys, _args, meta) in zip(replies, responses, descriptor_groups): + if isinstance(response, BaseException) or response["overall_code"] != "OK": + continue + applied.append(self._counter_refunds_from_atomic_response(_as_counter_values(reply), meta)) + statuses.extend(response["statuses"]) + reservation_windows.update(response.get("reservation_windows", frozenset())) + + over_limit: Final = next( + (r for r in responses if not isinstance(r, BaseException) and r["overall_code"] == "OVER_LIMIT"), None + ) + if over_limit is not None: + await self._refund_applied_descriptor_groups(applied) + return over_limit + failure: Final = next((r for r in responses if isinstance(r, BaseException)), None) + if failure is not None: + await self._refund_applied_descriptor_groups(applied) + self._reject_if_rate_limit_unverifiable("check_and_increment_by_n_script", failure) + log_redis_failure( + verbose_proxy_logger, + logging.ERROR, + f"atomic_check_and_increment_by_n: Redis Lua execution failed ({type(failure).__name__}). Refunding " + f"{len(applied)} pipelined descriptors and falling back to in-memory enforcement, counters will " + f"diverge from Redis until window expires (window_size={self.window_size}s)", + failure, + ) + flat_meta: Final = tuple( + itertools.chain.from_iterable(group_meta for _k, _a, group_meta in descriptor_groups) + ) + async with self._check_and_increment_lock: + return await self._atomic_check_and_increment_in_memory( + per_counter_meta=flat_meta, + parent_otel_span=parent_otel_span, + ) + if len(responses) == 1 and not isinstance(responses[0], BaseException): + return responses[0] + return RateLimitResponse( + overall_code="OK", + statuses=statuses, + reservation_windows=frozenset(reservation_windows), + ) + + def _pipelined_group_response( + self, reply: object, per_counter_meta: list[AtomicCounterMeta] + ) -> RateLimitResponse | BaseException: + if isinstance(reply, BaseException): + return reply + try: + return self._build_atomic_response(_as_counter_values(reply), per_counter_meta) + except Exception as e: # noqa: BLE001 # a reply this group cannot read is that group's Lua failure + return e + async def _refund_applied_descriptor_groups( self, applied: Sequence[Sequence[CounterRefund]], @@ -2233,7 +2365,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): async def _atomic_check_and_increment_in_memory( self, - per_counter_meta: list[AtomicCounterMeta], + per_counter_meta: Sequence[AtomicCounterMeta], parent_otel_span: Span | None = None, ) -> RateLimitResponse: """In-memory all-or-nothing check-and-increment. Caller holds lock. diff --git a/litellm/proxy/middleware/redis_request_batch_middleware.py b/litellm/proxy/middleware/redis_request_batch_middleware.py new file mode 100644 index 00000000000..bfb5f79a174 --- /dev/null +++ b/litellm/proxy/middleware/redis_request_batch_middleware.py @@ -0,0 +1,25 @@ +from typing import Final + +from starlette.types import ASGIApp, Receive, Scope, Send + +from litellm.caching.redis_batch import request_redis_batch_scope + +_REQUEST_SCOPES: Final = frozenset({"http", "websocket"}) + + +class RedisRequestBatchMiddleware: + """Opens the request's Redis batch scope so auth, admission and routing reads issued anywhere in the + request (dependencies, the endpoint, tasks it spawns) share one pipeline per Redis backend.""" + + def __init__(self, app: ASGIApp) -> None: + self.app = app + + async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: + if scope["type"] not in _REQUEST_SCOPES: + await self.app(scope, receive, send) + return + with request_redis_batch_scope() as batches: + try: + await self.app(scope, receive, send) + finally: + await batches.flush_all() diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index b4a497ea1e4..3f5e01ed0bc 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -681,6 +681,7 @@ from litellm.proxy.middleware.billable_request_metrics_middleware import ( from litellm.proxy.middleware.budget_reservation_release_middleware import ( BudgetReservationReleaseMiddleware, ) +from litellm.proxy.middleware.redis_request_batch_middleware import RedisRequestBatchMiddleware from litellm.proxy.plugin_routes import ( register_plugins_from_config, ) @@ -2417,6 +2418,7 @@ app.add_middleware( sink_factory=lambda: gateway_request_accumulator if prisma_client is not None else None, ) app.add_middleware(BudgetReservationReleaseMiddleware, release=release_unbound_budget_reservation) +app.add_middleware(RedisRequestBatchMiddleware) app.add_middleware(InFlightRequestsMiddleware) app.add_middleware(SecurityHeadersMiddleware) diff --git a/litellm/proxy/spend_tracking/spend_counter_batch.py b/litellm/proxy/spend_tracking/spend_counter_batch.py index a6694895a27..ae24331c236 100644 --- a/litellm/proxy/spend_tracking/spend_counter_batch.py +++ b/litellm/proxy/spend_tracking/spend_counter_batch.py @@ -10,6 +10,7 @@ from typing import Final from pydantic import TypeAdapter from litellm._logging import verbose_proxy_logger +from litellm.caching.redis_batch import BatchResult, RedisBatch, active_request_redis_batch from litellm.caching.redis_cache import RedisCache from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.user_api_key_cache import ( @@ -30,9 +31,13 @@ class PendingSpendIncrement: class SpendCounterBatch: """Bound counters are read with one MGET on first use; counters bound later join the next MGET. ``async_batch_get_cache`` maps a clean miss to ``None`` and drops keys only when Redis failed, so an absent - key means "read it yourself" and a present ``None`` is an authoritative miss.""" + key means "read it yourself" and a present ``None`` is an authoritative miss. - __slots__ = ("_fetched", "_keys", "_loaded", "_lock", "_open", "_redis_cache") + Inside a ``request_redis_batch_scope`` the MGET rides the request's pipeline instead: the batch's flush + hook declares whatever is bound but unread, so whoever flushes first (the auth object prefetch, usually) + carries the spend counters in the same round trip.""" + + __slots__ = ("_fetched", "_inflight", "_keys", "_loaded", "_lock", "_open", "_redis_cache", "_request_batch") def __init__(self, redis_cache: RedisCache) -> None: self._redis_cache: Final = redis_cache @@ -41,6 +46,10 @@ class SpendCounterBatch: self._keys: frozenset[str] = frozenset() self._fetched: frozenset[str] = frozenset() self._loaded: Mapping[str, float | None] = _NO_VALUES + self._inflight: Final[list[BatchResult[Mapping[str, object]]]] = [] # mutable-ok: drained by _load + self._request_batch: Final[RedisBatch | None] = active_request_redis_batch(redis_cache) + if self._request_batch is not None: + self._request_batch.add_flush_hook(self._declare_pending) @property def counter_keys(self) -> frozenset[str]: @@ -85,6 +94,10 @@ class SpendCounterBatch: async def _load(self) -> Mapping[str, float | None]: async with self._lock: + if self._request_batch is not None: + self._declare_pending() + await self._collect_inflight() + return self._loaded pending: Final = self._keys - self._fetched if pending: self._fetched = self._fetched | pending @@ -92,6 +105,26 @@ class SpendCounterBatch: self._loaded = MappingProxyType({**fetched, **self._loaded}) return self._loaded + def _declare_pending(self) -> None: + """Flush hook: put every bound-but-unread counter on the request pipeline that is about to go out.""" + if self._request_batch is None or not self._open: + return + pending: Final = self._keys - self._fetched + if pending: + self._fetched = self._fetched | pending + self._inflight.append(self._request_batch.mget(sorted(pending))) + + async def _collect_inflight(self) -> None: + results: Final = tuple(self._inflight) + self._inflight.clear() + for result in results: + try: + fetched: Mapping[str, float | None] = _CounterValues.validate_python(await result) + except Exception as e: # noqa: BLE001 # per-key reads take over and apply their own Redis fallback + verbose_proxy_logger.debug("spend counter batch read failed, falling back to per-key reads: %s", e) + continue + self._loaded = MappingProxyType({**fetched, **self._loaded}) + async def _fetch(self, keys: frozenset[str]) -> Mapping[str, float | None]: try: return _CounterValues.validate_python( diff --git a/litellm/router.py b/litellm/router.py index c18dfea1a36..8aaf58d5a3e 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -259,7 +259,7 @@ from litellm.router_utils.routing_groups import ( parse_routing_groups, validate_routing_strategy, ) -from litellm.router_utils.routing_read_batch import RoutingReadBatch +from litellm.router_utils.routing_read_batch import RoutingPrefetch, RoutingReadBatch from litellm.scheduler import FlowItem, Scheduler from litellm.types.litellm_params import RoutingStrategyName from litellm.types.llms.openai import ( @@ -539,6 +539,10 @@ def _is_retriable_anthropic_status(status_code: int) -> bool: return status_code == 429 or status_code >= 500 +def _without_line_breaks(value: object) -> str: + return str(value).replace("\r", "").replace("\n", "") + + def _anthropic_stream_error_is_gateway_verdict(chunk: object) -> bool: """AgenticAnthropicStreamingIterator's own retrieval-failure frame is the gateway's verdict, not a provider failure: another deployment would rerun the same failed hook, so it reaches the client instead of falling back.""" @@ -1730,6 +1734,25 @@ class Router: normalized for normalized in map(self._normalize_strategy, configured) if normalized is not None ) + def arm_routing_read_prefetch(self, model: str, request_kwargs: dict[str, object] | None = None) -> None: + """Declare the cooldown read (and, for usage-based routing, the usage read) that + `async_get_available_deployment` will make for `model` on the request's Redis batch, so admission's + flush carries it. A miss (alias, no batch) costs nothing: routing then reads as it always has.""" + try: + strategy, selector = self._get_routing_context(model, request_kwargs) + usage_selector: Final = ( + selector + if strategy == "usage-based-routing-v2" and isinstance(selector, LowestTPMLoggingHandler_v2) + else None + ) + deployments: Final = self.get_model_list(model_name=model) + if deployments: + RoutingPrefetch.arm(self, usage_selector, deployments) + except Exception as e: # noqa: BLE001 # a prefetch is an optimisation, never a reason to fail the request + verbose_router_logger.debug( + "routing read prefetch not armed for %s: %s", _without_line_breaks(model), _without_line_breaks(e) + ) + def _get_routing_context( self, model: str, request_kwargs: dict | None = None ) -> tuple[str | None, RouterStrategySelector | None]: diff --git a/litellm/router_utils/routing_read_batch.py b/litellm/router_utils/routing_read_batch.py index e9410d0e586..4039d7b1508 100644 --- a/litellm/router_utils/routing_read_batch.py +++ b/litellm/router_utils/routing_read_batch.py @@ -8,10 +8,15 @@ different objects. `RoutingReadBatch` fetches both key sets in one the usage slice to the strategy, so selection does not read again. """ +import itertools +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final from litellm._logging import verbose_router_logger from litellm.caching.dual_cache import DualCache +from litellm.caching.redis_batch import BatchResult, active_request_redis_batches from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2, PrefetchedUsage from litellm.router_utils.cooldown_cache import CooldownCache @@ -27,16 +32,68 @@ else: Span = Any +_PREFETCH_SLOT: Final = "routing_read" + + +@dataclass(frozen=True, slots=True) +class RoutingPrefetch: + """The cooldown and usage keys of a model group, declared on the request's Redis batch before admission + flushes it, so the routing read rides the same round trip as the rate limiter's Lua calls.""" + + keys: frozenset[str] + result: BatchResult[Mapping[str, object]] + + @staticmethod + def arm( + litellm_router_instance: LitellmRouter, + usage_selector: LowestTPMLoggingHandler_v2 | None, + deployments: list, + ) -> None: + request: Final = active_request_redis_batches() + redis_cache: Final = litellm_router_instance.cache.redis_cache + if request is None or redis_cache is None or _PREFETCH_SLOT in request.prefetched: + return + cooldown_keys: Final = tuple( + CooldownCache.get_cooldown_cache_key(model_id) for model_id in litellm_router_instance.get_model_ids() + ) + usage_keys: Final = ( + () if usage_selector is None else tuple(itertools.chain(*usage_selector.usage_counter_keys(deployments))) + ) + keys: Final = (*cooldown_keys, *usage_keys) + request.prefetched[_PREFETCH_SLOT] = RoutingPrefetch( + keys=frozenset(keys), result=request.batch(redis_cache).mget(keys) + ) + + @staticmethod + def armed() -> bool: + request: Final = active_request_redis_batches() + return request is not None and _PREFETCH_SLOT in request.prefetched + + @staticmethod + def take(needed: Sequence[str]) -> "RoutingPrefetch | None": + """The armed prefetch when it covers every key this read needs; taken once, so a retry reads fresh.""" + request: Final = active_request_redis_batches() + if request is None: + return None + armed: Final = request.prefetched.pop(_PREFETCH_SLOT, None) + if isinstance(armed, RoutingPrefetch) and armed.keys.issuperset(needed): + return armed + return None + + class RoutingReadBatch: - def __init__(self, usage_selector: LowestTPMLoggingHandler_v2) -> None: + def __init__(self, usage_selector: LowestTPMLoggingHandler_v2 | None) -> None: self.usage_selector: Final = usage_selector self.prefetched_usage: PrefetchedUsage | None = None @staticmethod def for_strategy(strategy: str | None, selector: object) -> "RoutingReadBatch | None": + """Usage-based routing reads its counters with the cooldown state; every other strategy reads only the + cooldown state, and only through this batch when the request armed a prefetch for it. Otherwise the + router's plain cooldown read stays in charge.""" if strategy == "usage-based-routing-v2" and isinstance(selector, LowestTPMLoggingHandler_v2): return RoutingReadBatch(usage_selector=selector) - return None + return RoutingReadBatch(usage_selector=None) if RoutingPrefetch.armed() else None async def async_get_cooldown_deployments( self, @@ -50,23 +107,50 @@ class RoutingReadBatch: """ model_ids: Final = litellm_router_instance.get_model_ids() cooldown_keys: Final = [CooldownCache.get_cooldown_cache_key(model_id) for model_id in model_ids] - tpm_keys, rpm_keys = self.usage_selector.usage_counter_keys(healthy_deployments) - usage_keys: Final = tpm_keys + rpm_keys - - cooldown_results, usage_values = await DualCache.async_batch_get_cache_shared( - [ - (litellm_router_instance.cooldown_cache.cooldown_store, cooldown_keys), - (self.usage_selector.router_cache, usage_keys), - ], - parent_otel_span=parent_otel_span, - ) - self.prefetched_usage = PrefetchedUsage( - keys=frozenset(usage_keys), - values=None if usage_values is None else dict(zip(usage_keys, usage_values)), + reads: Final[list[tuple[DualCache, list[str]]]] = [ # mutable-ok: the usage read is appended below + (litellm_router_instance.cooldown_cache.cooldown_store, cooldown_keys) + ] + usage_keys: list[str] = [] # mutable-ok: DualCache batch reads take a list + if self.usage_selector is not None: + tpm_keys, rpm_keys = self.usage_selector.usage_counter_keys(healthy_deployments) + usage_keys = tpm_keys + rpm_keys + reads.append((self.usage_selector.router_cache, usage_keys)) + results: Final = await self._read_prefetched(reads) or await DualCache.async_batch_get_cache_shared( + reads, parent_otel_span=parent_otel_span ) + cooldown_results: Final = results[0] + if self.usage_selector is not None: + usage_values: Final = results[1] + self.prefetched_usage = PrefetchedUsage( + keys=frozenset(usage_keys), + values=None if usage_values is None else MappingProxyType(dict(zip(usage_keys, usage_values))), + ) cooldown_models: Final = litellm_router_instance.cooldown_cache.active_cooldowns_from_results( model_ids, cooldown_results ) verbose_router_logger.debug("retrieve cooldown models: %s", cooldown_models) return [model_id for model_id, _ in cooldown_models] + + @staticmethod + async def _read_prefetched( + reads: list[tuple[DualCache, list[str]]], + ) -> list[list[object | None] | None] | None: + """Serve the reads from the request's armed `RoutingPrefetch`, backfilling each cache's memory tier as + its own batch read would. None when nothing usable was armed or the prefetch failed.""" + prefetch: Final = RoutingPrefetch.take(tuple(itertools.chain.from_iterable(keys for _, keys in reads))) + if prefetch is None: + return None + try: + values: Final = await prefetch.result + except Exception as e: # noqa: BLE001 # the shared read below applies the caches' own Redis fallback + verbose_router_logger.debug("routing prefetch failed, reading again: %s", e) + return None + results: Final[list[list[object | None] | None]] = [] # mutable-ok: filled per read below + for cache, keys in reads: + pending = await cache._prepare_batch_get(keys, local_only=True) # pyright: ignore[reportPrivateUsage] # same two-step read as async_batch_get_cache_shared + missed = { # mutable-ok: _apply_batch_get takes a dict + key: values.get(key) for key, local in zip(keys, pending.result) if local is None + } + results.append(await cache._apply_batch_get(pending, missed)) # pyright: ignore[reportPrivateUsage] # same two-step read as async_batch_get_cache_shared + return results diff --git a/tests/code_coverage_tests/router_code_coverage.py b/tests/code_coverage_tests/router_code_coverage.py index 7c247ae3303..df149f6c56a 100644 --- a/tests/code_coverage_tests/router_code_coverage.py +++ b/tests/code_coverage_tests/router_code_coverage.py @@ -90,6 +90,7 @@ ignored_function_names = [ "_resolve_claude_code_session_router", # Tested through Claude Code session routing in test_router.py "_get_claude_code_session_router_binding", # Tested through the two-worker session routing test in test_router.py "_apply_updated_routing_strategy_args", # Tested via update_settings in test_lowest_latency.py (file lacks "router" in name) + "arm_routing_read_prefetch", # Tested in tests/unit/caching/test_request_redis_batch_pre_call.py (file lacks "router" in name) ] 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 9aff2636c42..3be8501f5da 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 @@ -6974,6 +6974,39 @@ async def test_batch_increment_refunds_counters_already_applied_when_a_later_clu assert redis.increments == [] +@pytest.mark.parametrize("fail_closed", [True, False], ids=["fail_closed", "fail_open"]) +@pytest.mark.asyncio +async def test_batch_increment_refunds_pipelined_groups_declared_after_the_one_that_failed(fail_closed): + from unittest.mock import patch + + redis = _ScriptedRedis() + handler = _handler_with_redis(redis, fail_closed=fail_closed) + now = int(time.time()) + groups = {"a": ["{a}:window", "{a}:requests"], "b": ["{b}:window", "{b}:requests"]} + loop = asyncio.get_running_loop() + failed_group = loop.create_future() + failed_group.set_exception(ConnectionError("Error 61 connecting to 127.0.0.1:6379. Connection refused.")) + landed_group = loop.create_future() + landed_group.set_result([now, 1]) + + with ( + patch.object(handler, "_group_keys_by_hash_tag", return_value=groups), + patch.object(handler, "_pipeline_scripts", return_value=[failed_group, landed_group]), + ): + if fail_closed: + with pytest.raises(HTTPException) as exc: + await handler._execute_redis_batch_rate_limiter_script( + keys_to_fetch=[*groups["a"], *groups["b"]], now_int=now + ) + assert exc.value.status_code == 503 + else: + await handler._execute_redis_batch_rate_limiter_script( + keys_to_fetch=[*groups["a"], *groups["b"]], now_int=now + ) + + assert redis.guarded_increments == ([(groups["b"], [str(now), -1, 0])] if fail_closed else []) + + @pytest.mark.parametrize( "limits, request_data, counter_scope", [ diff --git a/tests/unit/caching/test_redis_batch.py b/tests/unit/caching/test_redis_batch.py new file mode 100644 index 00000000000..9433aeac524 --- /dev/null +++ b/tests/unit/caching/test_redis_batch.py @@ -0,0 +1,299 @@ +"""RedisBatch: independent operations share one pipeline, each keeps its own result and failure.""" + +from __future__ import annotations + +import asyncio +import hashlib +import json +from collections.abc import Callable, Sequence +from datetime import timedelta +from typing import Any + +import pytest +from redis.exceptions import NoScriptError + +from litellm._service_logger import ServiceLogging +from litellm.caching.redis_batch import ( + RedisBatch, + active_request_redis_batch, + request_redis_batch_scope, +) +from litellm.caching.redis_cache import RedisCache, RedisCircuitBreaker +from litellm.caching.redis_cluster_cache import RedisClusterCache + +SCRIPT = "return redis.call('GET', KEYS[1])" +SHA = hashlib.sha1(SCRIPT.encode()).hexdigest() # noqa: S324 + + +class FakePipeline: + def __init__(self, reply_for: Callable[[tuple[object, ...]], object], fail: Exception | None) -> None: + self.commands: list[tuple[Any, ...]] = [] + self.reply_for = reply_for + self.fail = fail + self.executed = False + + async def __aenter__(self) -> FakePipeline: + return self + + async def __aexit__(self, *exc: object) -> None: + return None + + def mget(self, keys: Sequence[str]) -> FakePipeline: + self.commands.append(("MGET", *keys)) + return self + + def evalsha(self, sha: str, numkeys: int, *keys_and_args: object) -> FakePipeline: + self.commands.append(("EVALSHA", sha, numkeys, *keys_and_args)) + return self + + def incrbyfloat(self, name: str, amount: float) -> FakePipeline: + self.commands.append(("INCRBYFLOAT", name, amount)) + return self + + def expire(self, name: str, time: timedelta) -> FakePipeline: + self.commands.append(("EXPIRE", name, int(time.total_seconds()))) + return self + + def set(self, name: str, value: str, ex: timedelta | None = None) -> FakePipeline: + self.commands.append(("SET", name, value, None if ex is None else int(ex.total_seconds()))) + return self + + async def execute(self, raise_on_error: bool = True) -> list[Any]: + assert raise_on_error is False + self.executed = True + if self.fail is not None: + raise self.fail + return [self.reply_for(command) for command in self.commands] + + +class FakeClient: + def __init__(self, reply_for: Callable[[tuple[object, ...]], object], fail: Exception | None = None) -> None: + self.pipelines: list[FakePipeline] = [] + self.reply_for = reply_for + self.fail = fail + + def pipeline(self, transaction: bool = True) -> FakePipeline: + assert transaction is False + pipe = FakePipeline(self.reply_for, self.fail) + self.pipelines.append(pipe) + return pipe + + +class FakeRedisCache(RedisCache): + def __init__(self, client: FakeClient, namespace: str | None = None) -> None: # super().__init__ needs a server + self.client = client + self.namespace = namespace + self._circuit_breaker = RedisCircuitBreaker(failure_threshold=5, recovery_timeout=30) + self.service_logger_obj = ServiceLogging() + self.default_ttl = None + self.alone: list[tuple[str, Any]] = [] + self.store: dict[str, Any] = {} + + def init_async_client(self) -> FakeClient: # pyright: ignore[reportIncompatibleMethodOverride] # fake client, no server + return self.client + + async def async_batch_get_cache(self, key_list: Sequence[str], **kwargs: object) -> dict[str, Any]: # pyright: ignore[reportIncompatibleMethodOverride] # records the direct read + self.alone.append(("MGET", tuple(key_list))) + return {key: self.store.get(key) for key in key_list} + + async def async_increment(self, key: str, value: float, ttl: int | None = None, **kwargs: object) -> float: # pyright: ignore[reportIncompatibleMethodOverride] # records the direct write + self.alone.append(("INCRBYFLOAT", key, value)) + self.store[key] = float(self.store.get(key, 0.0)) + value + return self.store[key] + + async def async_set_cache_pipeline_with_ttls(self, cache_list: Sequence[tuple[str, object, float | None]]) -> None: + self.alone.append(("SET_PIPELINE", tuple(cache_list))) + for key, value, _ttl in cache_list: + self.store[key] = value + + +class FakeClusterCache(RedisClusterCache, FakeRedisCache): + def __init__(self, client: FakeClient) -> None: # super().__init__ needs a server + FakeRedisCache.__init__(self, client) + + +def replies(command: tuple[Any, ...]) -> Any: + match command[0]: + case "MGET": + return [json.dumps({"k": key}) if key.endswith("hit") else None for key in command[1:]] + case "EVALSHA": + return [1, 2] + case "INCRBYFLOAT": + return b"3.5" + case "EXPIRE": + return 1 + case "SET": + return True + raise AssertionError(command) + + +def make(fail: Exception | None = None, namespace: str | None = None) -> tuple[FakeRedisCache, FakeClient]: + client = FakeClient(replies, fail) + return FakeRedisCache(client, namespace), client + + +async def run_alone_script(keys: Sequence[str], args: Sequence[Any]) -> object: + return ["alone", *keys, *args] + + +@pytest.mark.asyncio +async def test_one_pipeline_carries_every_declared_operation_and_awaiting_one_flushes_all() -> None: + cache, client = make(namespace="ns") + batch = RedisBatch(cache) + got = batch.mget(["a:hit", "b", "a:hit"]) + script = batch.script(SCRIPT, run_alone_script, ["w"], [7, "x"]) + incr = batch.increment("cnt", 2.5, ttl=60) + plain = batch.increment("cnt2", 1) + assert client.pipelines == [] + + assert await got == {"a:hit": {"k": "ns:a:hit"}, "b": None} + assert script.done and incr.done and plain.done + assert await script == [1, 2] + assert await incr == 3.5 + assert await plain == 3.5 + assert batch.flushes == 1 + assert [pipe.commands for pipe in client.pipelines] == [ + [ + ("MGET", "ns:a:hit", "ns:b"), + ("EVALSHA", SHA, 1, "ns:w", 7, "x"), + ("INCRBYFLOAT", "ns:cnt", 2.5), + ("EXPIRE", "ns:cnt", 60), + ("INCRBYFLOAT", "ns:cnt2", 1), + ] + ] + assert cache.alone == [] + + +@pytest.mark.asyncio +async def test_operations_declared_after_a_flush_go_out_in_the_next_pipeline() -> None: + cache, client = make() + batch = RedisBatch(cache) + await batch.mget(["a"]) + later = batch.increment("cnt", 1) + assert not later.done + assert await later == 3.5 + assert batch.flushes == 2 + assert [pipe.commands for pipe in client.pipelines] == [[("MGET", "a")], [("INCRBYFLOAT", "cnt", 1)]] + + +@pytest.mark.asyncio +async def test_a_failing_reply_fails_only_its_own_operation() -> None: + def reply_for(command: tuple[Any, ...]) -> Any: + if command[0] == "EVALSHA": + return ValueError("script blew up") + return replies(command) + + client = FakeClient(reply_for) + cache = FakeRedisCache(client) + batch = RedisBatch(cache) + got = batch.mget(["a:hit"]) + script = batch.script(SCRIPT, run_alone_script, ["w"], []) + assert await got == {"a:hit": {"k": "a:hit"}} + with pytest.raises(ValueError, match="script blew up"): + await script + assert cache.alone == [] + + +@pytest.mark.asyncio +async def test_a_reply_an_operation_cannot_decode_fails_only_that_operation() -> None: + def reply_for(command: tuple[Any, ...]) -> Any: + if command[0] == "MGET": + return "not-a-list" + return replies(command) + + client = FakeClient(reply_for) + cache = FakeRedisCache(client) + batch = RedisBatch(cache) + got = batch.mget(["a:hit"]) + written = batch.set("w", {"k": 1}) + script = batch.script(SCRIPT, run_alone_script, ["w"], []) + with pytest.raises(TypeError, match="MGET reply is not a list"): + await got + assert await written is None + assert await script == [1, 2] + assert len(client.pipelines) == 1 + + +@pytest.mark.asyncio +async def test_pipeline_failure_fails_every_operation_and_trips_the_breaker() -> None: + cache, _client = make(fail=ConnectionError("redis down")) + batch = RedisBatch(cache) + got = batch.mget(["a"]) + incr = batch.increment("cnt", 1) + with pytest.raises(ConnectionError): + await got + with pytest.raises(ConnectionError): + await incr + assert cache._circuit_breaker._failure_count == 1 # pyright: ignore[reportPrivateUsage] + + +@pytest.mark.asyncio +async def test_noscript_reply_reruns_that_script_through_the_registered_executor() -> None: + def reply_for(command: tuple[Any, ...]) -> Any: + if command[0] == "EVALSHA": + return NoScriptError("NOSCRIPT No matching script") + return replies(command) + + client = FakeClient(reply_for) + cache = FakeRedisCache(client) + batch = RedisBatch(cache) + script = batch.script(SCRIPT, run_alone_script, ["w"], [1]) + incr = batch.increment("cnt", 1) + assert await script == ["alone", "w", 1] + assert await incr == 3.5 + assert batch.flushes == 1 + + +@pytest.mark.asyncio +async def test_cluster_cache_runs_each_operation_on_its_own_path() -> None: + client = FakeClient(replies) + cache = FakeClusterCache(client) + cache.store["a"] = 4 + batch = RedisBatch(cache) + got = batch.mget(["a", "b"]) + incr = batch.increment("cnt", 2) + assert await got == {"a": 4, "b": None} + assert await incr == 2.0 + assert client.pipelines == [] + assert cache.alone == [("MGET", ("a", "b")), ("INCRBYFLOAT", "cnt", 2)] + + +@pytest.mark.asyncio +async def test_flush_hook_lets_a_lazy_reader_join_the_pipeline_that_is_going_out() -> None: + cache, client = make() + batch = RedisBatch(cache) + joined: list[Any] = [] + batch.add_flush_hook(lambda: joined.append(batch.mget(["late"]))) + await batch.mget(["early"]) + assert len(joined) == 1 and joined[0].done + assert await joined[0] == {"late": None} + assert [pipe.commands for pipe in client.pipelines] == [[("MGET", "early"), ("MGET", "late")]] + + +@pytest.mark.asyncio +async def test_concurrent_awaiters_share_one_flush() -> None: + cache, client = make() + batch = RedisBatch(cache) + first = batch.mget(["a"]) + second = batch.mget(["b"]) + results = await asyncio.gather(first._wait(), second._wait()) # pyright: ignore[reportPrivateUsage] + assert results == [{"a": None}, {"b": None}] + assert batch.flushes == 1 + assert len(client.pipelines) == 1 + + +def test_request_scope_hands_out_one_batch_per_backend_and_nests() -> None: + cache_a, _ = make() + cache_b, _ = make() + assert active_request_redis_batch(cache_a) is None + with request_redis_batch_scope() as batches: + first = active_request_redis_batch(cache_a) + assert first is not None + assert active_request_redis_batch(cache_a) is first + assert active_request_redis_batch(cache_b) is not first + with request_redis_batch_scope() as inner: + assert inner is batches + assert active_request_redis_batch(cache_a) is first + assert active_request_redis_batch(cache_a) is first + assert len(batches.batches) == 2 + assert active_request_redis_batch(cache_a) is None diff --git a/tests/unit/caching/test_request_redis_batch_pre_call.py b/tests/unit/caching/test_request_redis_batch_pre_call.py new file mode 100644 index 00000000000..c0834974f26 --- /dev/null +++ b/tests/unit/caching/test_request_redis_batch_pre_call.py @@ -0,0 +1,530 @@ +"""One Redis pipeline per backend for the pre-call reads a request makes: rate limiter Lua groups, the +router's cooldown and usage read, auth identity and spend counters all join the request batch.""" + +from __future__ import annotations + +import asyncio +import hashlib +import json +from typing import Any, Final +from unittest.mock import AsyncMock + +import pytest + +from litellm import Router +from litellm.caching.dual_cache import DualCache +from litellm.caching.redis_batch import active_request_redis_batches, request_redis_batch_scope +from litellm.proxy._types import LiteLLM_UserTable +from litellm.proxy.auth.auth_object_prefetch import _CacheEntry, _write_back +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache +from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + CHECK_AND_INCREMENT_BY_N_SCRIPT, + RateLimitDescriptor, + RateLimitUnverifiableError, + _PROXY_MaxParallelRequestsHandler_v3, +) +from litellm.proxy.utils import InternalUsageCache +from litellm.router_utils.cooldown_cache import CooldownCache +from litellm.router_utils.routing_read_batch import RoutingPrefetch + +from .test_redis_batch import FakeClient, FakeRedisCache + +_MODEL_GROUP = "claude" +_FAR_FUTURE = 4_102_444_800.0 # 2100-01-01, a cooldown stamped then is still active + + +def sha_of(script: str) -> str: + return hashlib.sha1(script.encode()).hexdigest() # noqa: S324 + + +def _limiter(redis_cache: FakeRedisCache, fail_closed: bool = False) -> _PROXY_MaxParallelRequestsHandler_v3: + dual_cache = DualCache() + limiter = _PROXY_MaxParallelRequestsHandler_v3( + internal_usage_cache=InternalUsageCache(dual_cache=dual_cache), + fail_closed_resolver=lambda: fail_closed, + ) + dual_cache.attach_redis_cache(redis_cache) # after init: the fake has no server to register scripts on + limiter.check_and_increment_by_n_script = AsyncMock( + side_effect=AssertionError("descriptor groups must ride the request pipeline") + ) + limiter.window_guarded_token_increment_script = AsyncMock(return_value=[1, 0]) + return limiter + + +def _descriptor(key: str, value: str, rpm: int) -> RateLimitDescriptor: + return {"key": key, "value": value, "rate_limit": {"requests_per_unit": rpm}} + + +def _refunds(limiter: _PROXY_MaxParallelRequestsHandler_v3) -> list[tuple[str, float]]: + refund_script = limiter.window_guarded_token_increment_script + assert isinstance(refund_script, AsyncMock) + return [(call.kwargs["keys"][1], call.kwargs["args"][1]) for call in refund_script.await_args_list] + + +def _lua_ok_replies(command: tuple[Any, ...]) -> Any: + if command[0] == "EVALSHA": + return [0, 1, 1700000000] # OK: one counter, new_counter=1, window_start + if command[0] == "MGET": + return [None for _ in command[1:]] + if command[0] == "SET": + return True + raise AssertionError(command) + + +@pytest.mark.asyncio +async def test_descriptor_lua_calls_share_one_pipeline_and_each_keeps_its_result(): + client = FakeClient(_lua_ok_replies) + limiter = _limiter(FakeRedisCache(client)) + descriptors = [ + _descriptor("api_key", "k1", 10), + _descriptor("model_per_key", "k1:gpt", 5), + _descriptor("team", "t1", 20), + ] + + with request_redis_batch_scope(): + response = await limiter.atomic_check_and_increment_by_n( + descriptors=descriptors, + increments=[{"requests": 1}, {"requests": 1}, {"requests": 1}], + ) + + assert response["overall_code"] == "OK" + assert [s["descriptor_key"] for s in response["statuses"]] == ["api_key", "model_per_key", "team"] + assert len(client.pipelines) == 1 + evalshas = [c for c in client.pipelines[0].commands if c[0] == "EVALSHA"] + assert len(evalshas) == 3 + assert {c[1] for c in evalshas} == {sha_of(CHECK_AND_INCREMENT_BY_N_SCRIPT)} + assert [c[3] for c in evalshas] == ["{api_key:k1}:window", "{model_per_key:k1:gpt}:window", "{team:t1}:window"] + + +@pytest.mark.asyncio +async def test_an_over_limit_descriptor_in_the_pipeline_refunds_the_groups_that_were_applied(): + def replies(command: tuple[Any, ...]) -> Any: + if command[0] == "EVALSHA" and command[3] == "{team:t1}:window": + return [1, 1, 21, 20] # OVER_LIMIT on its first counter + return _lua_ok_replies(command) + + client = FakeClient(replies) + redis_cache = FakeRedisCache(client) + limiter = _limiter(redis_cache) + + with request_redis_batch_scope(): + response = await limiter.atomic_check_and_increment_by_n( + descriptors=[_descriptor("api_key", "k1", 10), _descriptor("team", "t1", 20)], + increments=[{"requests": 1}, {"requests": 1}], + ) + + assert response["overall_code"] == "OVER_LIMIT" + assert response["statuses"][0]["descriptor_key"] == "team" + assert _refunds(limiter) == [("{api_key:k1}:requests", -1.0)] + assert len(client.pipelines) == 1 + + +@pytest.mark.asyncio +async def test_an_over_limit_descriptor_also_refunds_the_groups_the_pipeline_incremented_after_it(): + def replies(command: tuple[Any, ...]) -> Any: + if command[0] == "EVALSHA" and command[3] == "{api_key:k1}:window": + return [1, 1, 11, 10] # OVER_LIMIT on the first group; the later groups already incremented + return _lua_ok_replies(command) + + client = FakeClient(replies) + redis_cache = FakeRedisCache(client) + limiter = _limiter(redis_cache) + + with request_redis_batch_scope(): + response = await limiter.atomic_check_and_increment_by_n( + descriptors=[ + _descriptor("api_key", "k1", 10), + _descriptor("team", "t1", 20), + _descriptor("model_per_key", "k1:gpt", 5), + ], + increments=[{"requests": 1}, {"requests": 1}, {"requests": 1}], + ) + + assert response["overall_code"] == "OVER_LIMIT" + assert response["statuses"][0]["descriptor_key"] == "api_key" + assert _refunds(limiter) == [("{team:t1}:requests", -1.0), ("{model_per_key:k1:gpt}:requests", -1.0)] + assert len(client.pipelines) == 1 + + +@pytest.mark.asyncio +async def test_a_redis_denial_stands_when_another_pipelined_group_fails(): + def replies(command: tuple[Any, ...]) -> Any: + if command[0] == "EVALSHA" and command[3] == "{api_key:k1}:window": + return [1, 1, 11, 10] # OVER_LIMIT + if command[0] == "EVALSHA" and command[3] == "{team:t1}:window": + return ValueError("script blew up") + return _lua_ok_replies(command) + + client = FakeClient(replies) + redis_cache = FakeRedisCache(client) + limiter = _limiter(redis_cache) + + with request_redis_batch_scope(): + response = await limiter.atomic_check_and_increment_by_n( + descriptors=[ + _descriptor("api_key", "k1", 10), + _descriptor("team", "t1", 20), + _descriptor("model_per_key", "k1:gpt", 5), + ], + increments=[{"requests": 1}, {"requests": 1}, {"requests": 1}], + ) + + assert response["overall_code"] == "OVER_LIMIT" # not the in-memory fallback's verdict + assert response["statuses"][0]["descriptor_key"] == "api_key" + assert _refunds(limiter) == [("{model_per_key:k1:gpt}:requests", -1.0)] + assert len(client.pipelines) == 1 + + +@pytest.mark.asyncio +async def test_one_failed_lua_group_refunds_the_other_pipelined_groups_and_falls_back_to_in_memory(): + def replies(command: tuple[Any, ...]) -> Any: + if command[0] == "EVALSHA" and command[3] == "{api_key:k1}:window": + return ValueError("script blew up") + return _lua_ok_replies(command) + + client = FakeClient(replies) + redis_cache = FakeRedisCache(client) + limiter = _limiter(redis_cache) + + with request_redis_batch_scope(): + response = await limiter.atomic_check_and_increment_by_n( + descriptors=[_descriptor("api_key", "k1", 10), _descriptor("team", "t1", 20)], + increments=[{"requests": 1}, {"requests": 1}], + ) + + assert response["overall_code"] == "OK" + assert len(response["statuses"]) == 2 # in-memory enforcement covered both descriptors + assert _refunds(limiter) == [("{team:t1}:requests", -1.0)] + assert len(client.pipelines) == 1 + + +@pytest.mark.parametrize( + "client, refunded", + [ + ( + FakeClient( + lambda command: ( + ValueError("script blew up") + if command[0] == "EVALSHA" and command[3] == "{api_key:k1}:window" + else _lua_ok_replies(command) + ) + ), + [("{team:t1}:requests", -1.0)], + ), + (FakeClient(_lua_ok_replies, fail=ConnectionError("redis down")), []), + ], + ids=["one_group_failed", "pipeline_failed"], +) +@pytest.mark.asyncio +async def test_fail_closed_rejects_when_a_pipelined_lua_group_cannot_be_verified( + client: FakeClient, refunded: list[tuple[str, float]] +): + limiter = _limiter(FakeRedisCache(client), fail_closed=True) + + with request_redis_batch_scope(), pytest.raises(RateLimitUnverifiableError) as exc: + await limiter.atomic_check_and_increment_by_n( + descriptors=[_descriptor("api_key", "k1", 10), _descriptor("team", "t1", 20)], + increments=[{"requests": 1}, {"requests": 1}], + ) + + assert exc.value.status_code == 503 + assert _refunds(limiter) == refunded + assert len(client.pipelines) == 1 + + +@pytest.mark.asyncio +async def test_a_pipeline_failure_refunds_nothing_and_falls_back_to_in_memory_enforcement(): + client = FakeClient(_lua_ok_replies, fail=ConnectionError("redis down")) + limiter = _limiter(FakeRedisCache(client)) + + with request_redis_batch_scope(): + response = await limiter.atomic_check_and_increment_by_n( + descriptors=[_descriptor("api_key", "k1", 10), _descriptor("team", "t1", 20)], + increments=[{"requests": 1}, {"requests": 1}], + ) + + assert response["overall_code"] == "OK" + assert len(response["statuses"]) == 2 + assert len(client.pipelines) == 1 + + +@pytest.mark.asyncio +async def test_without_a_request_scope_descriptor_groups_run_the_script_directly_as_before(): + client = FakeClient(_lua_ok_replies) + limiter = _limiter(FakeRedisCache(client)) + limiter.check_and_increment_by_n_script = AsyncMock(return_value=[0, 1, 1700000000]) + + response = await limiter.atomic_check_and_increment_by_n( + descriptors=[_descriptor("api_key", "k1", 10), _descriptor("team", "t1", 20)], + increments=[{"requests": 1}, {"requests": 1}], + ) + + assert response["overall_code"] == "OK" + assert limiter.check_and_increment_by_n_script.await_count == 2 + assert client.pipelines == [] + + +def _deployment(deployment_id: str) -> dict: + return { + "model_name": _MODEL_GROUP, + "litellm_params": {"model": "anthropic/claude-x", "api_key": "test", "mock_response": "pong"}, + "model_info": {"id": deployment_id}, + } + + +def _router(redis_cache: FakeRedisCache, routing_strategy: str = "usage-based-routing-v2") -> Router: + router = Router(model_list=[_deployment("dep-a"), _deployment("dep-b")], routing_strategy=routing_strategy) + router._update_redis_cache(cache=redis_cache) + return router + + +@pytest.mark.asyncio +async def test_armed_routing_read_rides_the_admission_pipeline_and_routing_issues_no_read_of_its_own(): + client = FakeClient(_lua_ok_replies) + redis_cache = FakeRedisCache(client) + router = _router(redis_cache) + limiter = _limiter(redis_cache) + + with request_redis_batch_scope(): + router.arm_routing_read_prefetch(_MODEL_GROUP, {}) + await limiter.atomic_check_and_increment_by_n( + descriptors=[_descriptor("api_key", "k1", 10), _descriptor("team", "t1", 20)], + increments=[{"requests": 1}, {"requests": 1}], + ) + deployment = await router.async_get_available_deployment( + model=_MODEL_GROUP, messages=[{"role": "user", "content": "ping"}], request_kwargs={} + ) + + assert deployment["model_info"]["id"] in {"dep-a", "dep-b"} + assert len(client.pipelines) == 1 + commands = client.pipelines[0].commands + assert [c[0] for c in commands] == ["MGET", "EVALSHA", "EVALSHA"] + mget_keys = set(commands[0][1:]) + assert {CooldownCache.get_cooldown_cache_key("dep-a"), CooldownCache.get_cooldown_cache_key("dep-b")} <= mget_keys + assert any(":tpm:" in key for key in mget_keys) and any(":rpm:" in key for key in mget_keys) + assert redis_cache.alone == [] + + +@pytest.mark.asyncio +async def test_a_cooldown_recorded_locally_after_the_prefetch_left_still_excludes_its_deployment(): + expired = {"exception_received": "429", "status_code": "429", "timestamp": 0.0, "cooldown_time": 60} + + def replies(command: tuple[Any, ...]) -> Any: + if command[0] == "MGET": # Redis holds a stale cooldown for dep-b and nothing for dep-a + return [ + json.dumps(expired) if key == CooldownCache.get_cooldown_cache_key("dep-b") else None + for key in command[1:] + ] + return _lua_ok_replies(command) + + client = FakeClient(replies) + redis_cache = FakeRedisCache(client) + router = _router(redis_cache) + cooldown_store = router.cooldown_cache.cooldown_store + assert cooldown_store.in_memory_cache is not None + + with request_redis_batch_scope(): + router.arm_routing_read_prefetch(_MODEL_GROUP, {}) + cooldown_store.in_memory_cache.set_cache( + CooldownCache.get_cooldown_cache_key("dep-a"), + {"exception_received": "429", "status_code": "429", "timestamp": _FAR_FUTURE, "cooldown_time": 60}, + ) + picks = { + ( + await router.async_get_available_deployment( + model=_MODEL_GROUP, messages=[{"role": "user", "content": "ping"}], request_kwargs={} + ) + )["model_info"]["id"] + for _ in range(5) + } + + assert picks == {"dep-b"} + assert len(client.pipelines) == 1 + + +@pytest.mark.asyncio +async def test_a_prefetch_that_does_not_cover_the_routing_keys_is_ignored_and_routing_reads_itself(): + client = FakeClient(_lua_ok_replies) + redis_cache = FakeRedisCache(client) + router = _router(redis_cache) + + with request_redis_batch_scope() as request: + router.arm_routing_read_prefetch(_MODEL_GROUP, {}) + armed = request.prefetched["routing_read"] + assert isinstance(armed, RoutingPrefetch) + request.prefetched["routing_read"] = RoutingPrefetch(keys=frozenset({"other"}), result=armed.result) + deployment = await router.async_get_available_deployment( + model=_MODEL_GROUP, messages=[{"role": "user", "content": "ping"}], request_kwargs={} + ) + assert request.prefetched == {} + + assert deployment["model_info"]["id"] in {"dep-a", "dep-b"} + assert len(redis_cache.alone) == 1 # the shared cooldown+usage read, one round trip as in P1 + + +@pytest.mark.asyncio +async def test_a_failed_prefetch_falls_back_to_the_shared_read(): + client = FakeClient(_lua_ok_replies, fail=ConnectionError("redis down")) + redis_cache = FakeRedisCache(client) + router = _router(redis_cache) + + with request_redis_batch_scope(): + router.arm_routing_read_prefetch(_MODEL_GROUP, {}) + deployment = await router.async_get_available_deployment( + model=_MODEL_GROUP, messages=[{"role": "user", "content": "ping"}], request_kwargs={} + ) + + assert deployment["model_info"]["id"] in {"dep-a", "dep-b"} + assert len(redis_cache.alone) == 1 + + +@pytest.mark.asyncio +async def test_arming_outside_a_request_scope_is_a_no_op(): + redis_cache = FakeRedisCache(FakeClient(_lua_ok_replies)) + router = _router(redis_cache) + router.arm_routing_read_prefetch(_MODEL_GROUP, {}) + assert active_request_redis_batches() is None + + +@pytest.mark.asyncio +async def test_simple_shuffle_prefetches_only_its_cooldown_read_into_the_admission_pipeline(): + client = FakeClient(_lua_ok_replies) + redis_cache = FakeRedisCache(client) + router = _router(redis_cache, routing_strategy="simple-shuffle") + limiter = _limiter(redis_cache) + + with request_redis_batch_scope(): + router.arm_routing_read_prefetch(_MODEL_GROUP, {}) + await limiter.atomic_check_and_increment_by_n( + descriptors=[_descriptor("api_key", "k1", 10)], + increments=[{"requests": 1}], + ) + deployment = await router.async_get_available_deployment( + model=_MODEL_GROUP, messages=[{"role": "user", "content": "ping"}], request_kwargs={} + ) + + assert deployment["model_info"]["id"] in {"dep-a", "dep-b"} + assert len(client.pipelines) == 1 + commands = client.pipelines[0].commands + assert [c[0] for c in commands] == ["MGET", "EVALSHA"] + assert set(commands[0][1:]) == { + CooldownCache.get_cooldown_cache_key("dep-a"), + CooldownCache.get_cooldown_cache_key("dep-b"), + } + assert redis_cache.alone == [] + + shuffle = Router(model_list=[_deployment("dep-a")], routing_strategy="simple-shuffle") + shuffle._update_redis_cache(cache=redis_cache) + with request_redis_batch_scope() as request: + shuffle.arm_routing_read_prefetch(_MODEL_GROUP, {}) + armed = request.prefetched["routing_read"] + assert isinstance(armed, RoutingPrefetch) + assert armed.keys == {CooldownCache.get_cooldown_cache_key("dep-a")} # no usage counters for shuffle + + +@pytest.mark.asyncio +async def test_two_backends_flush_concurrently_one_pipeline_each(): + a_client, b_client = FakeClient(_lua_ok_replies), FakeClient(_lua_ok_replies) + a, b = FakeRedisCache(a_client), FakeRedisCache(b_client) + with request_redis_batch_scope() as request: + ra = request.batch(a).mget(["x", "y"]) + rb = request.batch(b).mget(["x"]) + await asyncio.gather(ra, rb) + assert len(a_client.pipelines) == 1 and len(b_client.pipelines) == 1 + + +@pytest.mark.asyncio +async def test_a_single_lua_group_rides_the_pipeline_with_the_armed_routing_read(): + client = FakeClient(_lua_ok_replies) + redis_cache = FakeRedisCache(client) + router = _router(redis_cache) + limiter = _limiter(redis_cache) + + with request_redis_batch_scope(): + router.arm_routing_read_prefetch(_MODEL_GROUP, {}) + await limiter.atomic_check_and_increment_by_n( + descriptors=[_descriptor("api_key", "k1", 10)], + increments=[{"requests": 1}], + ) + await router.async_get_available_deployment( + model=_MODEL_GROUP, messages=[{"role": "user", "content": "ping"}], request_kwargs={} + ) + + assert len(client.pipelines) == 1 + assert [c[0] for c in client.pipelines[0].commands] == ["MGET", "EVALSHA"] + assert redis_cache.alone == [] + + +class _SameServerCache(FakeRedisCache): + def __init__(self, client: FakeClient, namespace: str | None = None, **redis_kwargs: object) -> None: + super().__init__(client, namespace) + self.redis_kwargs = redis_kwargs + + +@pytest.mark.asyncio +async def test_caches_built_from_the_same_connection_settings_share_the_request_pipeline(): + client = FakeClient(_lua_ok_replies) + proxy_cache = _SameServerCache(client, host="r", port=6379, db=0) + router_cache = _SameServerCache(FakeClient(_lua_ok_replies), port="6379", host="r", db=0, password=None) + other_cache = _SameServerCache(FakeClient(_lua_ok_replies), host="r", port=6380, db=0) + with request_redis_batch_scope() as request: + assert request.batch(proxy_cache) is request.batch(router_cache) + assert request.batch(proxy_cache) is not request.batch(other_cache) + a = request.batch(proxy_cache).mget(["a"]) + b = request.batch(router_cache).mget(["b"]) + await asyncio.gather(a, b) + assert len(client.pipelines) == 1 + assert [c[0] for c in client.pipelines[0].commands] == ["MGET", "MGET"] + + +@pytest.mark.asyncio +async def test_caches_on_one_server_with_different_namespaces_keep_their_own_key_prefix(): + proxy_client, router_client = FakeClient(_lua_ok_replies), FakeClient(_lua_ok_replies) + proxy_cache = _SameServerCache(proxy_client, namespace="proxy", host="r", port=6379, db=0) + router_cache = _SameServerCache(router_client, namespace="router", host="r", port=6379, db=0) + with request_redis_batch_scope() as request: + await asyncio.gather(request.batch(proxy_cache).mget(["a"]), request.batch(router_cache).mget(["b"])) + sent: Final = tuple( + tuple(command for pipe in client.pipelines for command in pipe.commands) + for client in (proxy_client, router_client) + ) + assert sent == ((("MGET", "proxy:a"),), (("MGET", "router:b"),)), "each cache reads under its own namespace" + + +def _user_entry() -> tuple[_CacheEntry, LiteLLM_UserTable]: + entry = _CacheEntry("user-1", "user_row", LiteLLM_UserTable, 42) + return entry, LiteLLM_UserTable(user_id="user-1", max_budget=None, spend=0.0) + + +@pytest.mark.asyncio +async def test_auth_write_back_rides_the_next_round_trip_and_the_scope_drains_what_nobody_awaited(): + client = FakeClient(_lua_ok_replies) + redis_cache = FakeRedisCache(client) + cache = UserApiKeyCache(redis_cache=redis_cache) + with request_redis_batch_scope() as request: + await _write_back([_user_entry()], cache) + assert client.pipelines == [] # not sent yet: the SET waits for the next round trip + await request.batch(redis_cache).mget(["spend:key:k1"]) + assert len(client.pipelines) == 1 + kinds = [c[0] for c in client.pipelines[0].commands] + assert kinds == ["MGET", "SET"] or kinds == ["SET", "MGET"] + set_command = next(c for c in client.pipelines[0].commands if c[0] == "SET") + assert set_command[1] == "user-1" and set_command[3] == 42 + assert json.loads(set_command[2])["user_id"] == "user-1" + assert cache.in_memory_cache.get_cache("user-1") is not None + + await _write_back([_user_entry()], cache) + assert len(client.pipelines) == 1 + await request.flush_all() + assert len(client.pipelines) == 2 + assert [c[0] for c in client.pipelines[1].commands] == ["SET"] + + +@pytest.mark.asyncio +async def test_auth_write_back_outside_a_scope_writes_through_as_before(): + redis_cache = FakeRedisCache(FakeClient(_lua_ok_replies)) + cache = UserApiKeyCache(redis_cache=redis_cache) + await _write_back([_user_entry()], cache) + assert [(op[0], [(key, ttl) for key, _value, ttl in op[1]]) for op in redis_cache.alone] == [ + ("SET_PIPELINE", [("user-1", 42)]) + ] diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index 0aedfce3598..e4e65f8904c 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -47,6 +47,7 @@ from litellm.router import ( _anthropic_stream_should_decline_fallback, _is_retriable_anthropic_status, _responses_stream_holds_event, + _without_line_breaks, ) from litellm.router_strategy import simple_shuffle from litellm.router_utils.client_initalization_utils import MaxParallelRequestsLimit @@ -18803,3 +18804,36 @@ async def test_a_guardrail_verdict_is_neither_retried_nor_fallen_back(verdict: E await router.acompletion(model="primary", messages=[{"role": "user", "content": "hi"}]) assert [c.kwargs["metadata"]["model_group"] for c in mock_acompletion.call_args_list] == ["primary"] + + +@pytest.mark.parametrize( + ("value", "expected"), + [ + ("gpt-4\r\nERROR forged entry\n", "gpt-4ERROR forged entry"), + (RuntimeError("no deployments\r\nfor gpt-4"), "no deploymentsfor gpt-4"), + ("gpt-4", "gpt-4"), + ], +) +def test_without_line_breaks_drops_every_cr_and_lf_from_the_logged_value(value: object, expected: str) -> None: + assert _without_line_breaks(value) == expected + + +def test_a_failed_routing_read_prefetch_logs_the_request_model_without_its_line_breaks(monkeypatch, caplog) -> None: + router = litellm.Router( + model_list=[{"model_name": "gpt-4", "litellm_params": {"model": "openai/gpt-4", "api_key": "k"}}] + ) + forged_model: Final = "gpt-4\r\nERROR forged entry\n" + + def fail_lookup(model_name: str | None = None, team_id: str | None = None) -> None: + raise RuntimeError(f"no deployments for {model_name}") + + monkeypatch.setattr(router, "get_model_list", fail_lookup) + caplog.clear() + + with caplog.at_level(logging.DEBUG, logger="LiteLLM Router"): + router.arm_routing_read_prefetch(forged_model, {}) + + messages: Final = [r.getMessage() for r in caplog.records if "routing read prefetch not armed" in r.getMessage()] + assert messages == [ + "routing read prefetch not armed for gpt-4ERROR forged entry: no deployments for gpt-4ERROR forged entry" + ] From c129ea4fc99a583fb212150b659d46818baeb7b7 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 17:25:39 -0700 Subject: [PATCH 40/41] fix(mcp): scope OpenAPI listings to the exact server prefix and drop upstream OAuth metadata when a server is saved (#43608) * fix(mcp): key discovery caches per caller correctly and drop stale caches on server updates Discovery-list cache identity now uses the hashed token instead of the raw api_key and treats MCPJWTSigner-signed servers as per caller. Server definition changes also drop the cached upstream OAuth metadata. OpenAPI listings look tools up under the normalized registry prefix with the separator, so an overlapping sibling prefix no longer leaks into the list. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): keep the discovery cache digest call unchanged so CodeQL matches the existing alert Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): guard OAuth metadata cache writes with a per-server generation and drop unproven per-caller discovery keys An upstream metadata fetch that started before a server edit could store its stale reply after invalidate_oauth_metadata_cache ran. Invalidation now bumps a per-server generation and the fetch only stores when the generation it captured before I/O is unchanged. The MCPJWTSigner-based per-caller discovery classification and the api_key to token key change had no reproduction (the signer only injects on tools/list, and UserAPIKeyAuth hashes api_key in place), so both go back to the merge-base behavior. Integration coverage under tests/integration/mcp: overlapping OpenAPI aliases, a config-declared server name with a space, OAuth metadata refetch after a save, and the in-flight stale-write race Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): keep OAuth metadata generations only while a fetch is in flight Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): count queued OAuth metadata fetchers so invalidation survives lock handoff Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): keep a held OAuth metadata lock registered even when no fetcher slot claims it Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): prove a peer worker drops stale upstream OAuth metadata after a save elsewhere Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yucheng Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/discoverable_endpoints.py | 85 +++++++--- .../mcp_server/mcp_server_manager.py | 26 +-- tests/integration/mcp/test_mcp_management.py | 52 ++++++ .../mcp/test_oauth_configuration.py | 139 +++++++++++++++- .../mcp_server/test_discoverable_endpoints.py | 148 ++++++++++++++++++ .../mcp_server/test_mcp_server_manager.py | 54 +++++++ 6 files changed, 470 insertions(+), 34 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index d42c1c6b879..7a0f59c3c2b 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -3,7 +3,8 @@ import html as _html import json import secrets import time -from collections.abc import Callable, Mapping +from collections.abc import AsyncIterator, Callable, Mapping +from contextlib import asynccontextmanager from datetime import datetime, timezone from typing import TYPE_CHECKING, Any, Final, Literal, Optional from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse @@ -107,6 +108,14 @@ _OAUTH_METADATA_CACHE_MAX_SIZE: Final = 128 # Per-(server_id, resource_url) async locks so concurrent discovery requests # coalesce onto a single upstream fetch instead of issuing N parallel calls. _OAUTH_METADATA_FETCH_LOCKS: Final[dict[tuple[str, str], asyncio.Lock]] = {} +# Callers inside ``_oauth_metadata_fetch_slot`` per cache key, lock waiters included. ``Lock.locked()`` +# reads False between one holder's release and the next waiter's wake-up, so it cannot tell an +# idle lock from one being handed off. +_OAUTH_METADATA_FETCHERS: Final[dict[tuple[str, str], int]] = {} +# Per-server_id generation, bumped on invalidation so a fetch that started before the server +# definition changed cannot repopulate the cache with the stale reply. Only servers with a fetch +# in flight carry an entry; the rest are pruned with the cache. +_OAUTH_METADATA_GENERATIONS: Final[dict[str, int]] = {} router: Final = APIRouter( tags=["mcp"], @@ -130,13 +139,52 @@ def _prune_oauth_metadata_cache(now: float | None = None) -> None: for cache_key in cache_keys_by_expiry[:overflow]: _OAUTH_METADATA_CACHE.pop(cache_key, None) - # Drop locks whose cache entry has been evicted and that aren't currently - # held; held locks stay so in-flight callers continue to coalesce. + # Drop locks whose cache entry has been evicted and that nobody holds or + # waits on; the rest stay so in-flight callers continue to coalesce. for cache_key in list(_OAUTH_METADATA_FETCH_LOCKS): - if cache_key in _OAUTH_METADATA_CACHE: + if cache_key in _OAUTH_METADATA_CACHE or not _oauth_metadata_lock_idle(cache_key): continue - lock = _OAUTH_METADATA_FETCH_LOCKS.get(cache_key) - if lock is None or lock.locked(): + _OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None) + + for server_id in [sid for sid in _OAUTH_METADATA_GENERATIONS if not _oauth_metadata_fetch_in_flight(sid)]: + _OAUTH_METADATA_GENERATIONS.pop(server_id, None) + + +def _oauth_metadata_fetch_in_flight(server_id: str) -> bool: + return any(cache_key[0] == server_id for cache_key in _OAUTH_METADATA_FETCHERS) + + +def _oauth_metadata_lock_idle(cache_key: tuple[str, str]) -> bool: + if cache_key in _OAUTH_METADATA_FETCHERS: + return False + lock: Final = _OAUTH_METADATA_FETCH_LOCKS.get(cache_key) + return lock is None or not lock.locked() + + +@asynccontextmanager +async def _oauth_metadata_fetch_slot(cache_key: tuple[str, str]) -> AsyncIterator[None]: + _OAUTH_METADATA_FETCHERS[cache_key] = _OAUTH_METADATA_FETCHERS.get(cache_key, 0) + 1 + try: + async with _OAUTH_METADATA_FETCH_LOCKS.setdefault(cache_key, asyncio.Lock()): + yield + finally: + remaining: Final = _OAUTH_METADATA_FETCHERS.get(cache_key, 0) - 1 + if remaining > 0: + _OAUTH_METADATA_FETCHERS[cache_key] = remaining + else: + _OAUTH_METADATA_FETCHERS.pop(cache_key, None) + + +def invalidate_oauth_metadata_cache(server_id: str) -> None: + """Drop cached upstream IdP metadata for a server whose definition changed.""" + if _oauth_metadata_fetch_in_flight(server_id): + _OAUTH_METADATA_GENERATIONS[server_id] = _OAUTH_METADATA_GENERATIONS.get(server_id, 0) + 1 + else: + _OAUTH_METADATA_GENERATIONS.pop(server_id, None) + for cache_key in [key for key in _OAUTH_METADATA_CACHE if key[0] == server_id]: + del _OAUTH_METADATA_CACHE[cache_key] + for cache_key in [key for key in _OAUTH_METADATA_FETCH_LOCKS if key[0] == server_id]: + if not _oauth_metadata_lock_idle(cache_key): continue _OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None) @@ -2360,12 +2408,19 @@ async def fetch_upstream_oauth_protected_resource( if cached is not None and cached[0] > now: return cached[1] - lock: Final = _OAUTH_METADATA_FETCH_LOCKS.setdefault(cache_key, asyncio.Lock()) - async with lock: + async with _oauth_metadata_fetch_slot(cache_key): now = time.time() cached = _OAUTH_METADATA_CACHE.get(cache_key) if cached is not None and cached[0] > now: return cached[1] + generation: Final = _OAUTH_METADATA_GENERATIONS.get(mcp_server.server_id, 0) + + def store(payload: dict | None, ttl_seconds: int) -> None: + if _OAUTH_METADATA_GENERATIONS.get(mcp_server.server_id, 0) != generation: + return + stored_at: Final = time.time() + _OAUTH_METADATA_CACHE[cache_key] = (stored_at + ttl_seconds, payload) + _prune_oauth_metadata_cache(stored_at) host_base: Final = f"{upstream.scheme}://{upstream.netloc}" candidates: Final = [f"{host_base}/.well-known/oauth-protected-resource"] @@ -2407,12 +2462,7 @@ async def fetch_upstream_oauth_protected_resource( ) continue if isinstance(payload, dict): - now = time.time() - _OAUTH_METADATA_CACHE[cache_key] = ( - now + _OAUTH_METADATA_CACHE_TTL_SECONDS, - payload, - ) - _prune_oauth_metadata_cache(now) + store(payload, _OAUTH_METADATA_CACHE_TTL_SECONDS) return payload if len(network_errors) == len(candidates): @@ -2421,12 +2471,7 @@ async def fetch_upstream_oauth_protected_resource( # Negative-result caching: when no candidate yielded a usable payload, # remember that for a shorter TTL so we don't re-fetch on every # subsequent discovery request (and so the per-key lock can be pruned). - now = time.time() - _OAUTH_METADATA_CACHE[cache_key] = ( - now + _OAUTH_METADATA_NEGATIVE_CACHE_TTL_SECONDS, - None, - ) - _prune_oauth_metadata_cache(now) + store(None, _OAUTH_METADATA_NEGATIVE_CACHE_TTL_SECONDS) return None diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 490d3072955..c0792c32de2 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -2673,7 +2673,7 @@ class MCPServerManager: self._assign_unique_short_prefix(new_server) _warn_legacy_delegate_auth_if_applicable(new_server, source="config") _warn_config_id_jag_server_outruns_sso(new_server) - self._invalidate_discovery_lists(server_id) + self._invalidate_server_definition_caches(server_id) self.config_mcp_servers[server_id] = new_server self._set_oauth_discovery_deferred( server_id, @@ -2877,7 +2877,7 @@ class MCPServerManager: global_mcp_tool_registry, ) - self._invalidate_discovery_lists(server.server_id) + self._invalidate_server_definition_caches(server.server_id) prefix_root: Final = normalize_server_name(get_server_prefix(server)) if server.spec_path and prefix_root: openapi_key_prefix: Final = prefix_root + MCP_TOOL_PREFIX_SEPARATOR @@ -3285,7 +3285,7 @@ class MCPServerManager: # env_vars_are_encrypted=False. new_server: Final = await self.build_mcp_server_from_table(mcp_server, env_vars_are_encrypted=False) self._assign_unique_short_prefix(new_server) - self._invalidate_discovery_lists(mcp_server.server_id) + self._invalidate_server_definition_caches(mcp_server.server_id) self.registry[mcp_server.server_id] = new_server await self._maybe_register_openapi_tools(new_server) self.prime_oauth_metadata_discovery(new_server) @@ -3322,7 +3322,7 @@ class MCPServerManager: previous_server=self.registry[mcp_server.server_id], ) self._assign_unique_short_prefix(new_server) - self._invalidate_discovery_lists(mcp_server.server_id) + self._invalidate_server_definition_caches(mcp_server.server_id) self.registry[mcp_server.server_id] = new_server await self._maybe_register_openapi_tools(new_server) self.prime_oauth_metadata_discovery(new_server) @@ -4504,16 +4504,16 @@ class MCPServerManager: if server.spec_path: # OpenAPI tools were stored in the registry under the prefix # active at registration time — fetch by that same prefix. - registered_prefix: Final = f"{get_server_prefix(server)}{MCP_TOOL_PREFIX_SEPARATOR}" + registry_prefix: Final = normalize_server_name(get_server_prefix(server)) + MCP_TOOL_PREFIX_SEPARATOR registered: Final = global_mcp_tool_registry.convert_tools_to_mcp_sdk_tool_type( - global_mcp_tool_registry.list_tools(tool_prefix=get_server_prefix(server)) + global_mcp_tool_registry.list_tools(tool_prefix=registry_prefix) ) registered_names: Final = MappingProxyType( - {t.name.removeprefix(registered_prefix): t.name for t in registered} + {t.name.removeprefix(registry_prefix): t.name for t in registered} ) guarded_openapi: Final = await self._guard_tool_catalog( server=server, - tools=[t.model_copy(update={"name": t.name.removeprefix(registered_prefix)}) for t in registered], + tools=[t.model_copy(update={"name": t.name.removeprefix(registry_prefix)}) for t in registered], proxy_logging_obj=proxy_logging_obj, user_api_key_auth=user_api_key_auth, raw_headers=raw_headers, @@ -4582,6 +4582,14 @@ class MCPServerManager: self._resource_discovery_cache.invalidate(server_id) self._template_discovery_cache.invalidate(server_id) + def _invalidate_server_definition_caches(self, server_id: str) -> None: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( # noqa: PLC0415 # lazy: discoverable_endpoints lazily imports this module's manager singleton + invalidate_oauth_metadata_cache, + ) + + self._invalidate_discovery_lists(server_id) + invalidate_oauth_metadata_cache(server_id) + def _discovery_key( self, server: MCPServer, @@ -6792,7 +6800,7 @@ class MCPServerManager: for server_id in previous_registry.keys() | registered_registry.keys(): if previous_registry.get(server_id) != registered_registry.get(server_id): - self._invalidate_discovery_lists(server_id) + self._invalidate_server_definition_caches(server_id) self.registry = registered_registry _warn_on_shared_identifier_prefixes(registered_registry.values()) # A discovery task may have published into ``previous_registry`` while diff --git a/tests/integration/mcp/test_mcp_management.py b/tests/integration/mcp/test_mcp_management.py index 66c30a62bde..67cdbbff5a4 100644 --- a/tests/integration/mcp/test_mcp_management.py +++ b/tests/integration/mcp/test_mcp_management.py @@ -1,3 +1,4 @@ +import itertools import uuid from pathlib import Path from typing import Final @@ -7,10 +8,13 @@ import yaml from integration._support.client import Gateway, eventually from integration._support.mcp import ( McpCaller, + McpPeer, call_tool, delete_mcp, forget_mcp, + listed_tools, mcp_peer, + openapi_peer, register_mcp, tool_calls, tool_names, @@ -189,6 +193,54 @@ def test_duplicate_alias_is_rejected_so_tool_prefixes_cannot_collide(gateway: Ga scenario.cleanups.callback(forget_mcp, gateway, winner) +def _openapi_server_lists_and_calls_only_its_own_tools( + gateway: Gateway, key: str, peer: McpPeer, identity: str +) -> None: + listed: Final = set(listed_tools(gateway, key, identity)) + assert listed == {"getpet", "createpet"}, (identity, listed) + peer.drain() + called: Final = call_tool(gateway, key, identity, "getpet", {"petId": "7"}) + assert called.status_code == 200, called.text + assert [(item["method"], item["path"]) for item in peer.drain()] == [("GET", "/pets/7")], identity + + +def test_openapi_listing_is_scoped_to_the_exact_alias_when_aliases_overlap(gateway: Gateway) -> None: + with openapi_peer() as short, openapi_peer() as long, gateway.scenario() as scenario: + stem: Final = "pet" + uuid.uuid4().hex[:8] + servers: Final = tuple( + (peer, alias, register_mcp(scenario, peer, alias)) + for peer, alias in ((short, stem), (long, stem + "store")) + ) + key: Final = scenario.key(object_permission={"mcp_servers": [identity for _, _, identity in servers]}) + for peer, _, identity in servers: + _openapi_server_lists_and_calls_only_its_own_tools(gateway, key, peer, identity) + aggregate: Final = McpCaller(gateway, key, "mcp").list_tools() + assert aggregate.ok, aggregate.raw + assert sorted(aggregate.tools) == sorted( + f"{prefix}-{tool}" for prefix, tool in itertools.product((stem, stem + "store"), ("getpet", "createpet")) + ), aggregate.tools + assert all(peer.drain() == () for peer, _, _ in servers), "listing must not reach any OpenAPI upstream" + + +def test_config_declared_openapi_server_with_a_space_in_its_name_lists_its_tools( + gateway: Gateway, tmp_path: Path +) -> None: + with openapi_peer() as peer: + config: Final = yaml.safe_load((Path(__file__).resolve().parents[1] / "proxy_config.yaml").read_text()) + name: Final = "pet store " + uuid.uuid4().hex[:8] + config["mcp_servers"] = {name: peer.registration()} + path: Final = tmp_path / "openapi-space.yaml" + path.write_text(yaml.safe_dump(config)) + with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario: + identity: Final = next(i for i, s in _servers(candidate).items() if s["server_name"] == name) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + _openapi_server_lists_and_calls_only_its_own_tools(candidate, key, peer, identity) + aggregate: Final = McpCaller(candidate, key, "mcp").list_tools() + assert aggregate.ok, aggregate.raw + prefix: Final = name.replace(" ", "_") + assert sorted(aggregate.tools) == [f"{prefix}-createpet", f"{prefix}-getpet"], aggregate.tools + + def test_invalid_registrations_are_rejected(gateway: Gateway) -> None: with mcp_peer() as peer, gateway.scenario() as scenario: alias: Final = "mgmt" + uuid.uuid4().hex[:8] diff --git a/tests/integration/mcp/test_oauth_configuration.py b/tests/integration/mcp/test_oauth_configuration.py index 4c46c706054..fe2b1069f04 100644 --- a/tests/integration/mcp/test_oauth_configuration.py +++ b/tests/integration/mcp/test_oauth_configuration.py @@ -1,17 +1,23 @@ import json import queue +import threading import uuid -from urllib.parse import parse_qs, urlsplit -from typing import Final, Literal +from collections.abc import Callable +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass, field from pathlib import Path +from typing import Final, Literal +from urllib.parse import parse_qs, urlsplit import pytest - -from integration._support.client import Gateway, eventually +from integration._support.client import Gateway, Scenario, eventually from integration._support.database import read_rows from integration._support.mcp import McpPeer, call_tool, mcp_peer, register_mcp, tool_names from integration._support.process import owned_proxy -from integration._support.wire import Reply, Request, wire_server +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import TypeAdapter + +_Upstream = Callable[[Request], Reply] @pytest.mark.covers("other.mcp.oauth.discovery_cannot_erase_configured_authorization_endpoint") @@ -104,6 +110,129 @@ def test_partial_discovery_and_unrelated_edit_keep_actual_authorization_destinat assert updated.status_code == 202, updated.text +@dataclass(frozen=True, slots=True) +class _Hold: + armed: threading.Event = field(default_factory=threading.Event) + released: threading.Event = field(default_factory=threading.Event) + + +def _idp_upstream(origin: Callable[[], str], moved: threading.Event, hold: _Hold | None = None) -> _Upstream: + def issuer() -> str: + return origin() + ("/idp-after" if moved.is_set() else "/idp-before") + + def respond(request: Request) -> Reply: + if "oauth-authorization-server" in request.target or "openid-configuration" in request.target: + current: Final = issuer() + return Reply( + body=json.dumps( + { + "issuer": current, + "authorization_endpoint": current + "/authorize", + "token_endpoint": current + "/token", + } + ).encode() + ) + if request.target.startswith("/.well-known/oauth-protected-resource"): + body: Final = json.dumps({"resource": origin() + "/mcp", "authorization_servers": [issuer()]}).encode() + if hold is not None and hold.armed.is_set(): + assert hold.released.wait(timeout=15), "the held upstream metadata reply was never released" + return Reply(body=body) + return Reply(status=404, body=b'{"error":"unexpected"}') + + return respond + + +def _register_pass_through(scenario: Scenario, wire: Wire, alias: str) -> str: + return register_mcp(scenario, McpPeer(wire.url + "/mcp", queue.Queue()), alias, auth_type="true_passthrough") + + +def _wire_requests(wire: Wire, seen: list[Request]) -> Callable[[], tuple[Request, ...]]: + def observed() -> tuple[Request, ...]: + seen.extend(wire.drain()) + return tuple(seen) + + return observed + + +def _registration_discovery_settled(requests: tuple[Request, ...]) -> bool: + return any( + "oauth-authorization-server" in item.target or "openid-configuration" in item.target for item in requests + ) + + +def _advertised_authorization_servers(gateway: Gateway, alias: str) -> tuple[str, ...]: + response: Final = gateway.client.get(f"/.well-known/oauth-protected-resource/{alias}/mcp") + assert response.status_code == 200, response.text + return tuple(TypeAdapter(list[str]).validate_python(response.json()["authorization_servers"])) + + +def _eventually_advertises(gateway: Gateway, alias: str, issuer: str) -> None: + eventually( + lambda: gateway.client.get(f"/.well-known/oauth-protected-resource/{alias}/mcp"), + lambda response: response.status_code == 200 and response.json()["authorization_servers"] == [issuer], + seconds=40, + ) + + +def test_saving_a_pass_through_server_refetches_its_upstream_oauth_metadata(gateway: Gateway) -> None: + moved: Final = threading.Event() + with wire_server(_idp_upstream(lambda: wire.url, moved)) as wire, gateway.scenario() as scenario: + alias: Final = "pt" + uuid.uuid4().hex[:8] + identity: Final = _register_pass_through(scenario, wire, alias) + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-before",) + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-before",) + moved.set() + wire.drain() + saved: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "description": "IdP moved"}) + assert saved.status_code == 202, saved.text + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-after",) + assert any(request.target.startswith("/.well-known/oauth-protected-resource") for request in wire.drain()), ( + "the save must send protected-resource discovery back to the upstream" + ) + + +def test_peer_worker_stops_advertising_the_old_idp_after_a_save_on_another_worker( + gateway: Gateway, peer: Gateway +) -> None: + moved: Final = threading.Event() + with wire_server(_idp_upstream(lambda: wire.url, moved)) as wire, gateway.scenario() as scenario: + alias: Final = "pt" + uuid.uuid4().hex[:8] + identity: Final = _register_pass_through(scenario, wire, alias) + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-before",) + _eventually_advertises(peer, alias, wire.url + "/idp-before") + moved.set() + saved: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "description": "IdP moved"}) + assert saved.status_code == 202, saved.text + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-after",) + _eventually_advertises(peer, alias, wire.url + "/idp-after") + + +def test_metadata_fetched_before_a_save_cannot_repopulate_the_cache_after_it(gateway: Gateway) -> None: + moved: Final = threading.Event() + hold: Final = _Hold() + with ( + wire_server(_idp_upstream(lambda: wire.url, moved, hold)) as wire, + gateway.scenario() as scenario, + ThreadPoolExecutor(max_workers=1) as pool, + ): + alias: Final = "pt" + uuid.uuid4().hex[:8] + identity: Final = _register_pass_through(scenario, wire, alias) + seen: Final[list[Request]] = [] + observed: Final = _wire_requests(wire, seen) + eventually(observed, _registration_discovery_settled, seconds=10) + settled: Final = len(seen) + hold.armed.set() + stale: Final = pool.submit(_advertised_authorization_servers, gateway, alias) + eventually(observed, lambda requests: len(requests) > settled, seconds=10) + assert seen[settled].target.startswith("/.well-known/oauth-protected-resource"), seen[settled:] + moved.set() + saved: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "description": "IdP moved"}) + assert saved.status_code == 202, saved.text + hold.released.set() + assert stale.result(timeout=30) == (wire.url + "/idp-before",) + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-after",) + + @pytest.mark.covers("other.mcp.oauth.same_url_credentials_are_isolated_by_user_and_server") @pytest.mark.parametrize("transition", ("revoke", "expire")) def test_same_url_oauth_credentials_and_revocation_are_isolated_by_user_and_server( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index f9a0075e530..b1f0b3fa67e 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -12611,3 +12611,151 @@ async def test_identity_bound_authorize_unrelated_bearer_uses_browser_session( proxy_server.prisma_client.db.litellm_mcpusercredentials.upsert.assert_not_called() proxy_server.prisma_client.db.litellm_usertable.create.assert_not_called() proxy_server.prisma_client.db.litellm_teamtable.create.assert_not_called() + + +@pytest.mark.asyncio +async def test_update_server_drops_cached_upstream_oauth_metadata(): + from litellm.proxy._experimental.mcp_server import discoverable_endpoints + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + from litellm.proxy._types import LiteLLM_MCPServerTable, MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + manager = MCPServerManager() + server = MCPServer( + server_id="oauth-cache-server", + name="oauth_cache_server", + url="http://old-upstream/mcp", + transport=MCPTransport.http, + ) + manager.registry[server.server_id] = server + stale_key: Final = (server.server_id, server.url) + other_key: Final = ("other-server", "http://other/mcp") + discoverable_endpoints._OAUTH_METADATA_CACHE[stale_key] = (time.time() + 300, {"iss": "old-idp"}) + discoverable_endpoints._OAUTH_METADATA_CACHE[other_key] = (time.time() + 300, {"iss": "other"}) + try: + await manager.update_server( + LiteLLM_MCPServerTable( + server_id=server.server_id, + server_name=server.name, + url="http://new-upstream/mcp", + transport=MCPTransport.http, + ) + ) + assert stale_key not in discoverable_endpoints._OAUTH_METADATA_CACHE + assert other_key in discoverable_endpoints._OAUTH_METADATA_CACHE + finally: + discoverable_endpoints._OAUTH_METADATA_CACHE.pop(stale_key, None) + discoverable_endpoints._OAUTH_METADATA_CACHE.pop(other_key, None) + + +@pytest.mark.asyncio +async def test_metadata_fetched_before_invalidation_does_not_repopulate_the_cache(): + import asyncio + + from litellm.proxy._experimental.mcp_server import discoverable_endpoints + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + fetch_upstream_oauth_protected_resource, + invalidate_oauth_metadata_cache, + ) + from litellm.proxy._types import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="stale-write-server", name="stale_write", url="http://upstream/mcp", transport=MCPTransport.http + ) + cache_key: Final = (server.server_id, server.url) + started: Final = asyncio.Event() + release: Final = asyncio.Event() + + async def slow_get(url: str, headers: dict[str, str]) -> MagicMock: + started.set() + await release.wait() + return MagicMock(status_code=200, json=MagicMock(return_value={"authorization_servers": ["old-idp"]})) + + client = MagicMock() + client.get = slow_get + discoverable_endpoints._OAUTH_METADATA_CACHE.pop(cache_key, None) + try: + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=client, + ): + in_flight: Final = asyncio.create_task(fetch_upstream_oauth_protected_resource(server)) + await started.wait() + invalidate_oauth_metadata_cache(server.server_id) + release.set() + assert await in_flight == {"authorization_servers": ["old-idp"]} + assert cache_key not in discoverable_endpoints._OAUTH_METADATA_CACHE + discoverable_endpoints._prune_oauth_metadata_cache() + assert server.server_id not in discoverable_endpoints._OAUTH_METADATA_GENERATIONS + finally: + discoverable_endpoints._OAUTH_METADATA_CACHE.pop(cache_key, None) + discoverable_endpoints._OAUTH_METADATA_GENERATIONS.pop(server.server_id, None) + + +@pytest.mark.asyncio +async def test_fetch_waiting_on_a_lock_handoff_stays_tracked_through_invalidation(): + import asyncio + + from litellm.proxy._experimental.mcp_server import discoverable_endpoints + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + fetch_upstream_oauth_protected_resource, + invalidate_oauth_metadata_cache, + ) + from litellm.proxy._types import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="handoff-server", name="handoff", url="http://upstream/mcp", transport=MCPTransport.http + ) + cache_key: Final = (server.server_id, server.url) + started: Final = asyncio.Event() + release: Final = asyncio.Event() + + async def slow_get(url: str, headers: dict[str, str]) -> MagicMock: + started.set() + await release.wait() + return MagicMock(status_code=200, json=MagicMock(return_value={"authorization_servers": ["pre-save-idp"]})) + + client = MagicMock() + client.get = slow_get + discoverable_endpoints._OAUTH_METADATA_CACHE.pop(cache_key, None) + try: + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=client, + ): + async with discoverable_endpoints._oauth_metadata_fetch_slot(cache_key): + shared_lock: Final = discoverable_endpoints._OAUTH_METADATA_FETCH_LOCKS[cache_key] + waiting: Final = asyncio.create_task(fetch_upstream_oauth_protected_resource(server)) + for _ in range(3): + await asyncio.sleep(0) + assert not started.is_set() and not waiting.done() + invalidate_oauth_metadata_cache(server.server_id) + assert discoverable_endpoints._OAUTH_METADATA_FETCH_LOCKS.get(cache_key) is shared_lock + assert discoverable_endpoints._oauth_metadata_fetch_in_flight(server.server_id) + await started.wait() + invalidate_oauth_metadata_cache(server.server_id) + release.set() + assert await waiting == {"authorization_servers": ["pre-save-idp"]} + assert cache_key not in discoverable_endpoints._OAUTH_METADATA_CACHE + assert not discoverable_endpoints._oauth_metadata_fetch_in_flight(server.server_id) + finally: + discoverable_endpoints._OAUTH_METADATA_CACHE.pop(cache_key, None) + discoverable_endpoints._OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None) + discoverable_endpoints._OAUTH_METADATA_FETCHERS.pop(cache_key, None) + discoverable_endpoints._OAUTH_METADATA_GENERATIONS.pop(server.server_id, None) + + +def test_invalidating_an_idle_server_leaves_no_generation_behind(): + from litellm.proxy._experimental.mcp_server import discoverable_endpoints + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import invalidate_oauth_metadata_cache + + server_ids: Final = tuple(f"churned-server-{i}" for i in range(50)) + try: + for server_id in server_ids: + invalidate_oauth_metadata_cache(server_id) + assert not set(server_ids) & set(discoverable_endpoints._OAUTH_METADATA_GENERATIONS) + finally: + for server_id in server_ids: + discoverable_endpoints._OAUTH_METADATA_GENERATIONS.pop(server_id, None) 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 0a7ea012b98..a1dc0e779da 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 @@ -14064,6 +14064,60 @@ def test_discovery_cache_keys_isolate_user_dependent_auth(auth_type: MCPAuth) -> assert "second" not in str(second) +def _register_local_tool(name: str, description: str) -> None: + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + async def _handler(**kwargs): + return None + + global_mcp_tool_registry.register_tool( + name=name, description=description, input_schema={"type": "object"}, handler=_handler + ) + + +def _openapi_server(name: str) -> MCPServer: + return MCPServer( + server_id=f"{name}-id", name=name, alias=name, transport=MCPTransport.http, url=None, spec_path="/spec.yaml" + ) + + +@pytest.mark.asyncio +async def test_openapi_listing_ignores_overlapping_server_prefix() -> None: + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + manager: Final = MCPServerManager() + manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + for prefix in ("pet-", "petstore-"): + global_mcp_tool_registry.unregister_tools_with_prefix(prefix) + _register_local_tool("pet-list", "Local pet tool") + _register_local_tool("petstore-list", "Foreign petstore tool") + try: + prefixed: Final = await manager._get_tools_from_server(server=_openapi_server("pet"), add_prefix=True) + bare: Final = await manager._get_tools_from_server(server=_openapi_server("pet"), add_prefix=False) + finally: + for prefix in ("pet-", "petstore-"): + global_mcp_tool_registry.unregister_tools_with_prefix(prefix) + + assert [t.name for t in prefixed] == ["pet-list"] + assert [t.name for t in bare] == ["list"] + + +@pytest.mark.asyncio +async def test_openapi_listing_finds_tools_registered_under_the_normalized_prefix() -> None: + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + manager: Final = MCPServerManager() + manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + global_mcp_tool_registry.unregister_tools_with_prefix("pet_store-") + _register_local_tool("pet_store-list", "Pet store tool") + try: + listed: Final = await manager._get_tools_from_server(server=_openapi_server("pet store"), add_prefix=False) + finally: + global_mcp_tool_registry.unregister_tools_with_prefix("pet_store-") + + assert [t.name for t in listed] == ["list"] + + @pytest.mark.asyncio async def test_discovery_cache_retries_cancelled_fetches() -> None: from litellm.proxy._experimental.mcp_server.mcp_server_manager import _DiscoveryCache From 9525452d3706e3b0a39e2e7fa5f4c12b50439c78 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 29 Sep 2026 17:26:05 -0700 Subject: [PATCH 41/41] perf(proxy): one post-call Redis pipeline per backend for spend, rate-limit, routing and response-cache writes (#43779) Post-call owners declare into one request-scoped RedisBatch per Redis backend: spend counter increments and reservation reconciliation, rate-limit token Lua updates and refunds, parallel-slot release (freed locally at once), deployment TPM, and compatible async response-cache SETs. The batch is sent once the success and failure callbacks have run, or on a deadline, and pending batches are drained at shutdown before Redis disconnects. nx writes, non-Redis caches and calls outside a request stay direct; numeric string TTLs keep the direct-path coercion. Resolves LIT-8883 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: yassin --- litellm/caching/caching.py | 48 +- litellm/caching/dual_cache.py | 90 ++- litellm/caching/redis_batch.py | 107 ++- litellm/litellm_core_utils/litellm_logging.py | 3 + .../hooks/parallel_request_limiter_v3.py | 118 +++- .../proxy/hooks/proxy_track_cost_callback.py | 14 + litellm/proxy/proxy_server.py | 84 ++- litellm/router_strategy/lowest_tpm_rpm_v2.py | 2 +- .../test_request_redis_batch_post_call.py | 661 ++++++++++++++++++ 9 files changed, 1091 insertions(+), 36 deletions(-) create mode 100644 tests/unit/caching/test_request_redis_batch_post_call.py diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index e157730779b..85fd56c01e3 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -8,6 +8,7 @@ # Thank you users! We ❤️ you! - Krrish & Ishaan import ast +import asyncio import hashlib import json import logging @@ -30,10 +31,11 @@ from litellm.types.utils import EmbeddingResponse, is_litellm_owned_kwarg from .azure_blob_cache import AzureBlobCache from .base_cache import BaseCache from .disk_cache import DiskCache -from .dual_cache import DualCache # noqa: F401 +from .dual_cache import DualCache from .gcs_cache import GCSCache from .in_memory_cache import InMemoryCache from .qdrant_semantic_cache import QdrantSemanticCache +from .redis_batch import active_post_call_redis_batch from .redis_cache import RedisCache, log_redis_failure from .redis_cluster_cache import RedisClusterCache from .redis_semantic_cache import RedisSemanticCache @@ -68,6 +70,15 @@ def print_verbose(print_statement): pass +def _ttl_seconds(raw: object) -> int | None: + if not isinstance(raw, (int, float, str)): + return None + try: + return int(raw) + except ValueError: + return None + + class CacheMode(str, Enum): default_on = "default_on" default_off = "default_off" @@ -759,6 +770,8 @@ class Cache: await self.batch_cache_write(result, **kwargs) else: cache_key, cached_data, kwargs = self._add_cache_logic(result=result, **kwargs) + if await self._defer_set_to_post_call_batch(cache_key, cached_data, kwargs, dynamic_cache_object): + return if dynamic_cache_object is not None: await dynamic_cache_object.async_set_cache(cache_key, cached_data, **kwargs) else: @@ -766,6 +779,39 @@ class Cache: except Exception as e: self._log_add_cache_failure(e) + async def _defer_set_to_post_call_batch( + self, + cache_key: str, + cached_data: object, + kwargs: Mapping[str, object], + dynamic_cache_object: BaseCache | None, + ) -> bool: + """A plain SET on the Redis response cache rides the request's post-call pipeline with the counters, + instead of its own round trip. Anything with SET options keeps the direct path.""" + if kwargs.get("nx"): + return False + ttl: Final = _ttl_seconds(kwargs.get("ttl")) + if isinstance(dynamic_cache_object, DualCache): + deferred: Final = await dynamic_cache_object.async_set_cache_post_call(cache_key, cached_data, ttl) + if deferred is None: + return False + deferred.on_settled(self._log_deferred_add_cache_failure) + return True + if dynamic_cache_object is not None or not isinstance(self.cache, RedisCache): + return False + batch: Final = active_post_call_redis_batch(self.cache) + if batch is None: + return False + batch.set(cache_key, cached_data, ttl).on_settled(self._log_deferred_add_cache_failure) + return True + + def _log_deferred_add_cache_failure(self, future: asyncio.Future[None]) -> None: + if future.cancelled(): + return + failure: Final = future.exception() + if isinstance(failure, Exception): + self._log_add_cache_failure(failure) + def _convert_to_cached_embedding( self, embedding_response: Any, diff --git a/litellm/caching/dual_cache.py b/litellm/caching/dual_cache.py index 1d7afcbee8f..996273d558a 100644 --- a/litellm/caching/dual_cache.py +++ b/litellm/caching/dual_cache.py @@ -8,23 +8,23 @@ Has 4 primary methods: - async_get_cache """ +import asyncio import itertools import logging import time -from collections.abc import Sequence +from collections.abc import Mapping, Sequence from dataclasses import dataclass from threading import Lock from typing import TYPE_CHECKING, Any, Final -if TYPE_CHECKING: - from litellm.types.caching import RedisPipelineIncrementOperation - import litellm from litellm._logging import print_verbose, verbose_logger from litellm.constants import DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE +from litellm.types.caching import RedisPipelineIncrementOperation from .base_cache import BaseCache from .in_memory_cache import DEFAULT_MAX_SIZE_IN_MEMORY, InMemoryCache +from .redis_batch import BatchResult, RedisBatch, active_post_call_redis_batch from .redis_cache import RedisCache, RedisCircuitBreakerOpenError, log_redis_failure if TYPE_CHECKING: @@ -59,6 +59,24 @@ class PendingBatchRead: previous_access_times: dict[str, float | None] +@dataclass(frozen=True, slots=True) +class DeclaredBatchRead: + """A ``async_batch_get_cache`` split in two: the memory half done, the Redis half declared on a ``RedisBatch`` + so it rides that batch's next round trip, resolved later with ``async_resolve_batch_get``.""" + + keys: tuple[str, ...] + pending: PendingBatchRead + result: BatchResult[Mapping[str, object]] | None + + +def _log_deferred_increment_failure(future: asyncio.Future[float]) -> None: + if future.cancelled(): + return + failure: Final = future.exception() + if failure is not None: + log_redis_failure(verbose_logger, logging.WARNING, "post-call Redis increment failed", failure) + + class DualCache(BaseCache): """ DualCache is a cache implementation that updates both Redis and an in-memory cache simultaneously. @@ -335,7 +353,7 @@ class DualCache(BaseCache): ) async def _apply_batch_get( - self, pending: PendingBatchRead, redis_result: dict[str, object] | None, **kwargs: object + self, pending: PendingBatchRead, redis_result: Mapping[str, object] | None, **kwargs: object ) -> list[object | None]: if redis_result is None or all(v is None for v in redis_result.values()): return pending.result @@ -349,6 +367,22 @@ class DualCache(BaseCache): await self.in_memory_cache.async_set_cache(key, value, **self._backfill_kwargs(kwargs)) return merged + async def declare_batch_get(self, keys: Sequence[str], batch: RedisBatch) -> DeclaredBatchRead: + pending: Final = await self._prepare_batch_get( + list(keys), # mutable-ok: the shared batch read takes a list + local_only=False, + throttle_redis=False, + ) + return DeclaredBatchRead( + keys=tuple(keys), + pending=pending, + result=batch.mget(pending.redis_keys) if pending.redis_keys else None, + ) + + async def async_resolve_batch_get(self, declared: DeclaredBatchRead) -> list[object | None]: + redis_result: Final = None if declared.result is None else await declared.result + return await self._apply_batch_get(declared.pending, redis_result) + async def async_batch_get_cache( self, keys: list, @@ -468,6 +502,17 @@ class DualCache(BaseCache): verbose_logger, logging.ERROR, "LiteLLM Cache: exception in async add_cache", e, with_traceback=True ) + async def async_set_cache_post_call(self, key: str, value: object, ttl: float | None) -> BatchResult[None] | None: + """Memory now, the Redis SET on the request's post-call pipeline; None when no pipeline is open, so the + caller takes its direct path.""" + batch: Final = None if self.redis_cache is None else active_post_call_redis_batch(self.redis_cache) + if batch is None: + return None + effective_ttl: Final = self.default_in_memory_ttl if ttl is None else ttl + if self.in_memory_cache is not None: + await self.in_memory_cache.async_set_cache(key, value, ttl=effective_ttl) + return batch.set(key, value, effective_ttl) + # async_batch_set_cache async def async_set_cache_pipeline( self, cache_list: Sequence[tuple[str, object]], local_only: bool = False, **kwargs @@ -535,6 +580,41 @@ class DualCache(BaseCache): ) return result + async def async_increment_cache_post_call( + self, + key: str, + value: float, + ttl: int | None, + parent_otel_span: Span | None = None, + ) -> None: + """Memory is incremented now; the Redis increment rides the request's post-call pipeline when one is + open, and runs on its own as ``async_increment_cache`` otherwise.""" + await self.async_increment_cache_pipeline_post_call( + (RedisPipelineIncrementOperation(key=key, increment_value=value, ttl=ttl),), parent_otel_span + ) + + async def async_increment_cache_pipeline_post_call( + self, + increment_list: Sequence["RedisPipelineIncrementOperation"], + parent_otel_span: Span | None = None, + ) -> None: + batch: Final = None if self.redis_cache is None else active_post_call_redis_batch(self.redis_cache) + operations: Final = list(increment_list) # mutable-ok: both increment pipelines take a list + if batch is None: + await self.async_increment_cache_pipeline(operations, parent_otel_span=parent_otel_span) + return + try: + if self.in_memory_cache is not None: + await self.in_memory_cache.async_increment_pipeline( + increment_list=operations, parent_otel_span=parent_otel_span + ) + except Exception as e: # noqa: BLE001 # same tolerance as async_increment_cache_pipeline + log_redis_failure(verbose_logger, logging.WARNING, "in-memory increment failed", e) + for operation in increment_list: + batch.increment(operation["key"], operation["increment_value"], operation["ttl"]).on_settled( + _log_deferred_increment_failure + ) + async def async_increment_cache_pipeline( self, increment_list: list["RedisPipelineIncrementOperation"], diff --git a/litellm/caching/redis_batch.py b/litellm/caching/redis_batch.py index 8bd9554ec55..f6052192685 100644 --- a/litellm/caching/redis_batch.py +++ b/litellm/caching/redis_batch.py @@ -14,6 +14,7 @@ import hashlib import json import logging import time +import weakref from collections.abc import Awaitable, Callable, Generator, Mapping, Sequence from contextvars import ContextVar, Token from dataclasses import dataclass, field @@ -32,6 +33,8 @@ from litellm.types.services import ServiceTypes _T = TypeVar("_T") _ScriptArg = str | bytes | int | float +SettledHook = Callable[[asyncio.Future[_T]], Awaitable[None] | None] # mutable-ok: Callable params +POST_CALL_FLUSH_DEADLINE_SECONDS: Final = 1.0 class RegisteredScript(Protocol): @@ -52,11 +55,24 @@ class _Op(Generic[_T]): how to run on its own when the batch cannot pipeline (cluster client, or a reply the pipeline cannot settle, like NOSCRIPT).""" - __slots__ = ("future",) + __slots__ = ("future", "settled_hooks") def __init__(self) -> None: self.future: Final[asyncio.Future[_T]] = asyncio.get_running_loop().create_future() self.future.add_done_callback(_mark_retrieved) + self.settled_hooks: Final[list[SettledHook[_T]]] = [] # mutable-ok: append-only registry + + async def run_settled_hooks(self) -> None: + for hook in self.settled_hooks: + await self._run_settled_hook(hook) + + async def _run_settled_hook(self, hook: SettledHook[_T]) -> None: + try: + follow_up: Final = hook(self.future) + if follow_up is not None: + await follow_up + except Exception as e: # noqa: BLE001 # one owner's follow-up must not stop the others + verbose_logger.warning("redis batch settled hook failed: %s", e) def enqueue(self, pipe: _RedisPipeline) -> int: raise NotImplementedError @@ -238,6 +254,11 @@ class BatchResult(Generic[_T]): def done(self) -> bool: return self._op.future.done() + def on_settled(self, hook: SettledHook[_T]) -> None: + """For an owner that does not await: runs inside the flush once this operation has its result or + failure (or was cancelled with the pipeline), so the flush completes with the follow-up done.""" + self._op.settled_hooks.append(hook) + @dataclass(slots=True) class RedisBatch: @@ -294,6 +315,7 @@ class RedisBatch: for op in ops: if not op.future.done(): op.future.cancel() + await asyncio.gather(*(op.run_settled_hooks() for op in ops)) async def _flush_pipeline(self, ops: Sequence[_Op[object]]) -> None: start_time: Final = time.time() @@ -354,14 +376,34 @@ def _backend_key(redis_cache: RedisCache) -> object: return (type(redis_cache), redis_cache.namespace, settings) +_open_post_call: Final[weakref.WeakSet[RequestRedisBatches]] = weakref.WeakSet() +"""Requests whose post-call batch still holds declared ops, so a shutdown can send them before Redis goes away.""" + + class RequestRedisBatches: """One ``RedisBatch`` per Redis backend for the current request, so readers of different caches that - share a server (the proxy's and the router's) share the pipeline.""" + share a server (the proxy's and the router's) share the pipeline. - __slots__ = ("_batches", "prefetched") + The post-call batches hold the writes nothing waits on (counters, token scripts, the response cache). + They flush once, when the success or failure callbacks have all run, or at ``post_call_deadline`` + seconds after the first declaration when no callback phase closes them.""" - def __init__(self) -> None: + __slots__ = ( + "__weakref__", + "_batches", + "_deadline", + "_deadline_flush", + "_post_call", + "post_call_deadline", + "prefetched", + ) + + def __init__(self, post_call_deadline: float = POST_CALL_FLUSH_DEADLINE_SECONDS) -> None: self._batches: Final[dict[object, RedisBatch]] = {} # mutable-ok: lazily filled per backend + self._post_call: Final[dict[object, RedisBatch]] = {} # mutable-ok: lazily filled per backend + self.post_call_deadline: Final = post_call_deadline + self._deadline: asyncio.TimerHandle | None = None + self._deadline_flush: asyncio.Task[None] | None = None # Reads declared early for a consumer that runs later in the request, keyed by consumer name. self.prefetched: Final[dict[str, object]] = {} # mutable-ok: armed pre-admission, taken at use @@ -373,14 +415,44 @@ class RequestRedisBatches: self._batches[key] = batch return batch + def post_call(self, redis_cache: RedisCache) -> RedisBatch: + key: Final = _backend_key(redis_cache) + existing: Final = self._post_call.get(key) + batch: Final = ( + existing + if existing is not None + else self._post_call.setdefault(key, RedisBatch(redis_cache, name="post_call_redis_batch")) + ) + if self._deadline is None: + self._deadline = asyncio.get_running_loop().call_later(self.post_call_deadline, self._flush_on_deadline) + _open_post_call.add(self) + return batch + + def _flush_on_deadline(self) -> None: + self._deadline = None + self._deadline_flush = asyncio.ensure_future(self.flush_post_call()) + async def flush_all(self) -> None: """Send whatever is still declared (write-backs nobody awaits) before the request scope closes.""" await asyncio.gather(*(batch.flush() for batch in self._batches.values() if batch.pending)) + async def flush_post_call(self) -> None: + """One pipeline per backend for the post-call writes; the deadline is disarmed since this is that flush.""" + if self._deadline is not None: + self._deadline.cancel() + self._deadline = None + await asyncio.gather(*(batch.flush() for batch in self._post_call.values() if batch.pending)) + if not any(batch.pending for batch in self._post_call.values()): + _open_post_call.discard(self) + @property def batches(self) -> tuple[RedisBatch, ...]: return tuple(self._batches.values()) + @property + def post_call_batches(self) -> tuple[RedisBatch, ...]: + return tuple(self._post_call.values()) + _active_request_batches: Final[ContextVar[RequestRedisBatches | None]] = ContextVar( "request_redis_batches", default=None @@ -399,19 +471,40 @@ def active_request_redis_batches() -> RequestRedisBatches | None: return _active_request_batches.get() +def active_post_call_redis_batch(redis_cache: RedisCache) -> RedisBatch | None: + """The request's post-call batch for this backend, or None outside a ``request_redis_batch_scope``.""" + batches: Final = _active_request_batches.get() + if batches is None: + return None + return batches.post_call(redis_cache) + + +async def flush_post_call_redis_batches() -> None: + """Called where the success and failure callbacks of a request have all run.""" + batches: Final = _active_request_batches.get() + if batches is not None: + await batches.flush_post_call() + + +async def drain_post_call_redis_batches() -> None: + """Sends every post-call batch still waiting on its callbacks or deadline; for the shutdown path.""" + await asyncio.gather(*(batches.flush_post_call() for batches in tuple(_open_post_call))) + + class request_redis_batch_scope: """Redis reads declared inside share one pipeline per backend; nested scopes join the outer one.""" - __slots__ = ("_token",) + __slots__ = ("_post_call_deadline", "_token") - def __init__(self) -> None: + def __init__(self, post_call_deadline: float = POST_CALL_FLUSH_DEADLINE_SECONDS) -> None: self._token: Token[RequestRedisBatches | None] | None = None + self._post_call_deadline: Final = post_call_deadline def __enter__(self) -> RequestRedisBatches: outer: Final = _active_request_batches.get() if outer is not None: return outer - batches: Final = RequestRedisBatches() + batches: Final = RequestRedisBatches(post_call_deadline=self._post_call_deadline) self._token = _active_request_batches.set(batches) return batches diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index c393fa3caef..e292ab7b2ec 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -35,6 +35,7 @@ from litellm._uuid import uuid from litellm.batches.batch_utils import _handle_completed_batch, batch_cost_is_final from litellm.caching.caching import DualCache from litellm.caching.caching_handler import LLMCachingHandler +from litellm.caching.redis_batch import flush_post_call_redis_batches from litellm.constants import ( DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT, DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT, @@ -3552,6 +3553,7 @@ class Logging(LiteLLMLoggingBaseClass): traceback.format_exc(), ) self._handle_callback_failure(callback=callback) + await flush_post_call_redis_batches() def _handle_callback_failure(self, callback: object): """ @@ -3937,6 +3939,7 @@ class Logging(LiteLLMLoggingBaseClass): ) # Track callback logging failures in Prometheus self._handle_callback_failure(callback=callback) + await flush_post_call_redis_batches() def _get_trace_id(self, service_name: Literal["langfuse"]) -> str | None: """ diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 589aa7da7b5..e509f03d458 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -33,7 +33,12 @@ from typing_extensions import NotRequired, ReadOnly from litellm import DualCache from litellm._logging import verbose_proxy_logger -from litellm.caching.redis_batch import BatchResult, RegisteredScript, active_request_redis_batch +from litellm.caching.redis_batch import ( + BatchResult, + RegisteredScript, + active_post_call_redis_batch, + active_request_redis_batch, +) from litellm.caching.redis_cache import log_redis_failure from litellm.constants import DYNAMIC_RATE_LIMIT_ERROR_THRESHOLD_PER_MINUTE, INTERNAL_CALL_ORIGIN_METADATA_KEY from litellm.integrations.custom_logger import CustomLogger @@ -1893,6 +1898,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): self, stash: RequestRateLimiterStash | None, parent_otel_span: Span | None, + *, + in_logging_callback: bool = False, ) -> None: if stash is None: return @@ -1900,7 +1907,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): acquisition: Final = stash.parallel_slot if acquisition is None: return - await self._release_parallel_request_slots(acquisition, parent_otel_span) + deferred: Final = in_logging_callback and await self._defer_parallel_slot_release( + acquisition, parent_otel_span + ) + if not deferred: + await self._release_parallel_request_slots(acquisition, parent_otel_span) stash.parallel_slot = None # rebind-ok: marks this request's slot as released async def _release_parallel_request_slots( @@ -1926,14 +1937,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): keys=counter_keys, args=[slot_id for _ in counter_keys], ) - for counter_key, remaining in zip(counter_keys, raw): - await self.internal_usage_cache.async_set_cache( - key=counter_key, - value=max(0, int(remaining)), - ttl=PARALLEL_REQUEST_SLOT_TTL_SECONDS, - litellm_parent_otel_span=parent_otel_span, - local_only=True, - ) + await self._mirror_released_parallel_slots(counter_keys, raw, parent_otel_span) return except Exception as e: # noqa: BLE001 - any Redis/Lua failure degrades to the in-memory release, never a 500 log_redis_failure( @@ -1942,7 +1946,55 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): "parallel_release_script failed, falling back to in-memory release", e, ) + await self._release_parallel_request_slots_in_memory(counter_keys, slot_id, parent_otel_span) + async def _defer_parallel_slot_release( + self, acquisition: ParallelSlotAcquisition, parent_otel_span: Span | None + ) -> bool: + """Only for a release from the logging callbacks: the response has left and the callbacks' end flushes + the pipeline. A release before the response goes to Redis at once, so another worker's next acquire + never counts a finished request. The local gauge frees the slot at once, so admission on this worker + sees the capacity before the pipeline goes out. The count Redis returns from the pipeline is not + mirrored: by then a newer acquire on this worker may have written a fresher count, and the next + acquire refreshes the gauge anyway.""" + counter_keys: Final = acquisition["counter_keys"] + slot_id: Final = acquisition["slot_id"] + redis_cache: Final = self.internal_usage_cache.dual_cache.redis_cache + script: Final = self.parallel_release_script + batch: Final = None if redis_cache is None else active_post_call_redis_batch(redis_cache) + if batch is None or script is None or not counter_keys or not slot_id: + return False + await self._release_parallel_request_slots_in_memory(counter_keys, slot_id, parent_otel_span) + + async def settle(future: asyncio.Future[object]) -> None: + if future.cancelled() or future.exception() is not None: + log_redis_failure( + verbose_proxy_logger, + logging.WARNING, + "parallel_release_script failed, the slot stays released in memory only", + future.exception() if not future.cancelled() else asyncio.CancelledError(), + ) + + batch.script(PARALLEL_RELEASE_SCRIPT, script, counter_keys, (slot_id,) * len(counter_keys)).on_settled(settle) + return True + + async def _mirror_released_parallel_slots( + self, counter_keys: list[str], remaining_by_key: Sequence[object], parent_otel_span: Span | None + ) -> None: + for counter_key, remaining in zip(counter_keys, remaining_by_key): + if not isinstance(remaining, (int, float, str, bytes)): + continue + await self.internal_usage_cache.async_set_cache( + key=counter_key, + value=max(0, int(remaining)), + ttl=PARALLEL_REQUEST_SLOT_TTL_SECONDS, + litellm_parent_otel_span=parent_otel_span, + local_only=True, + ) + + async def _release_parallel_request_slots_in_memory( + self, counter_keys: list[str], slot_id: str, parent_otel_span: Span | None + ) -> None: async with self._check_and_increment_lock: for counter_key in counter_keys: raw_value: ParallelGaugeCacheValue | None = await self.internal_usage_cache.async_get_cache( @@ -4301,11 +4353,43 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): keys.append(op["key"]) args.extend([op["increment_value"], ttl_value]) + if self._defer_token_increment_script(keys, args, group_operations): + continue await self.token_increment_script( keys=keys, args=args, ) + def _defer_token_increment_script( + self, + keys: list[str], + args: list[int], + group_operations: list["RedisPipelineIncrementOperation"], + ) -> bool: + """Declared into the request's post-call pipeline instead of its own EVALSHA round trip; a failed + script falls back to the plain increment pipeline for its own group, as the direct path does.""" + redis_cache: Final = self.internal_usage_cache.dual_cache.redis_cache + script: Final = self.token_increment_script + batch: Final = None if redis_cache is None else active_post_call_redis_batch(redis_cache) + if batch is None or script is None: + return False + + async def fall_back(future: asyncio.Future[object]) -> None: + if future.cancelled() or future.exception() is None: + return + log_redis_failure( + verbose_proxy_logger, + logging.WARNING, + "TTL preservation failed, falling back to regular pipeline", + future.exception(), + ) + await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline( + increment_list=group_operations, + ) + + batch.script(TOKEN_INCREMENT_SCRIPT, script, keys, args).on_settled(fall_back) + return True + async def async_increment_tokens_with_ttl_preservation( self, pipeline_operations: list["RedisPipelineIncrementOperation"], @@ -4919,7 +5003,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): verbose_proxy_logger.debug("INSIDE parallel request limiter ASYNC SUCCESS LOGGING") stash: Final = get_request_stash_for_call(_call_id_from_callback_kwargs(kwargs)) - await self._release_stashed_parallel_slot(stash, litellm_parent_otel_span) + await self._release_stashed_parallel_slot(stash, litellm_parent_otel_span, in_logging_callback=True) pipeline_operations: Final = self._build_success_event_pipeline_operations( kwargs=kwargs, @@ -5039,7 +5123,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): pipeline_operations: Final[list[RedisPipelineIncrementOperation]] = [] stash: Final = get_request_stash_for_call(_call_id_from_callback_kwargs(kwargs)) - await self._release_stashed_parallel_slot(stash, litellm_parent_otel_span) + await self._release_stashed_parallel_slot(stash, litellm_parent_otel_span, in_logging_callback=True) # Skip the reservation refund if async_post_call_failure_hook # already released it (proxy-level rejection that also bubbles up @@ -5109,15 +5193,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) 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, + await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline_post_call( + pipeline_operations, parent_otel_span=litellm_parent_otel_span ) 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, + await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline_post_call( + project_operations, parent_otel_span=litellm_parent_otel_span ) elif project_operations: await self.async_increment_reservation_aware_tokens( diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 05995d22293..877dbfabe5d 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -285,6 +285,7 @@ class _ProxyDBLogger(CustomLogger): increment_spend_counters, proxy_logging_obj, update_cache, + update_cache_read_keys, ) verbose_proxy_logger.debug("INSIDE _PROXY_track_cost_callback") @@ -378,6 +379,13 @@ class _ProxyDBLogger(CustomLogger): request_tags=tags, model_access_groups=model_access_groups, project_id=project_id, + update_cache_read_keys=update_cache_read_keys( + user_id=user_id, + end_user_id=end_user_id, + team_id=team_id, + tags=tags, + response_cost=response_cost, + ), ) if not charged: return @@ -695,6 +703,7 @@ async def _update_database_and_spend_counters( request_tags: list[str] | None = None, model_access_groups: Sequence[str] | None = None, project_id: str | None = None, + update_cache_read_keys: Sequence[str] = (), ) -> bool: """The reservation is reconciled before the spend is persisted, from its own read. One spend counter batch then spans the database write and the counter update, so the post-call counters are read with a single MGET after the @@ -736,6 +745,7 @@ async def _update_database_and_spend_counters( request_tags=request_tags, model_access_groups=model_access_groups, project_id=project_id, + update_cache_read_keys=update_cache_read_keys, ) @@ -756,7 +766,10 @@ async def _update_database_and_spend_counters_in_batch( request_tags: list[str] | None, model_access_groups: Sequence[str] | None, project_id: str | None, + update_cache_read_keys: Sequence[str], ) -> bool: + from litellm.proxy.proxy_server import arm_update_cache_read + try: charged: Final = await proxy_logging_obj.db_spend_update_writer.update_database( token=user_api_key, @@ -788,6 +801,7 @@ async def _update_database_and_spend_counters_in_batch( await _release_budget_reservation(budget_reservation=budget_reservation) return False + await arm_update_cache_read(update_cache_read_keys) try: await increment_spend_counters( token=user_api_key, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 3f5e01ed0bc..a994293a4d8 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -269,6 +269,12 @@ import litellm._redis from litellm import Router from litellm._logging import _redact_string, verbose_proxy_logger, verbose_router_logger from litellm.caching.caching import DualCache, RedisCache +from litellm.caching.dual_cache import DeclaredBatchRead +from litellm.caching.redis_batch import ( + active_post_call_redis_batch, + active_request_redis_batches, + drain_post_call_redis_batches, +) from litellm.caching.redis_cache import RedisCircuitBreakerOpenError, is_redis_timeout_failure from litellm.caching.redis_cluster_cache import RedisClusterCache from litellm.constants import ( @@ -1112,6 +1118,7 @@ async def proxy_shutdown_event(worker_heartbeat: ProxyWorkerHeartbeat | None = N verbose_proxy_logger.debug("Disconnecting from Prisma") await prisma_client.disconnect() + await drain_post_call_redis_batches() if litellm.cache is not None: await litellm.cache.disconnect() @@ -3715,6 +3722,8 @@ async def _invalidate_spend_counter(counter_key: str): async def _apply_spend_counter_increments(pending: Sequence[PendingSpendIncrement]) -> None: + if _defer_spend_counter_increments(pending): + return try: await increment_spend_counters_pipeline(pending=pending) except Exception as e: @@ -3723,6 +3732,41 @@ async def _apply_spend_counter_increments(pending: Sequence[PendingSpendIncremen raise +def _defer_spend_counter_increments(pending: Sequence[PendingSpendIncrement]) -> bool: + """Post-call increments ride the request's post-call pipeline with the other counters. Each counter's + new value lands in memory when the pipeline settles; a failed one is invalidated so no reader trusts a + counter whose increment may not have applied, as ``increment_spend_counters_pipeline`` does.""" + redis_cache: Final = spend_counter_cache.redis_cache + if redis_cache is None or not pending: + return False + batch: Final = active_post_call_redis_batch(redis_cache) + if batch is None: + return False + ttl: Final = redis_cache.get_ttl() + for item in pending: + batch.increment(item.counter_key, item.increment, ttl).on_settled(_settle_spend_counter_increment(item)) + return True + + +def _settle_spend_counter_increment(item: PendingSpendIncrement) -> Callable[[asyncio.Future[float]], Awaitable[None]]: + async def settle(future: asyncio.Future[float]) -> None: + if not future.cancelled() and future.exception() is None: + current_value: Final = float(future.result()) + spend_counter_cache.in_memory_cache.set_cache(key=item.counter_key, value=current_value) + record_spend_counter_value(item.counter_key, current_value) + return + if future.cancelled(): + if spend_counter_cache.in_memory_cache.get_cache(key=item.counter_key) is not None: + spend_counter_cache.in_memory_cache.increment_cache(key=item.counter_key, value=item.increment) + return + verbose_proxy_logger.warning( + "Spend counter %s increment did not land in the post-call pipeline; invalidating it", item.counter_key + ) + await _invalidate_spend_counter(counter_key=item.counter_key) + + return settle + + async def increment_spend_counters_pipeline(pending: Sequence[PendingSpendIncrement]) -> tuple[float | None, ...]: """One INCRBYFLOAT+EXPIRE pipeline for every pending counter, returning each counter's new value in order; on failure every counter is invalidated before the error propagates, so no caller can read a half-applied batch.""" @@ -3762,7 +3806,7 @@ async def run_spend_counter_pipeline(pending: Sequence[PendingSpendIncrement]) - return tuple(float(current_value) for current_value in results or ()) -def _update_cache_read_keys( +def update_cache_read_keys( user_id: str | None, end_user_id: str | None, team_id: str | None, @@ -3778,13 +3822,45 @@ def _update_cache_read_keys( return user_keys + end_user_keys + team_keys + tag_keys -async def _read_update_cache_values(keys: Sequence[str], parent_otel_span: Span | None) -> Mapping[str, object]: +_UPDATE_CACHE_PREFETCH_SLOT: Final = "update_cache_read" + + +async def arm_update_cache_read(keys: Sequence[str], cache: DualCache | None = None) -> None: + """Declares the ``update_cache`` read on the request pipeline once the spend is persisted, so it rides the same + round trip as the post-call spend counter read instead of its own.""" + request: Final = active_request_redis_batches() + target: Final = user_api_key_cache if cache is None else cache + if request is None or target.redis_cache is None or not keys: + return + request.prefetched[_UPDATE_CACHE_PREFETCH_SLOT] = await target.declare_batch_get( + keys, request.batch(target.redis_cache) + ) + + +async def _take_armed_update_cache_read(keys: Sequence[str], cache: DualCache) -> Mapping[str, object] | None: + request: Final = active_request_redis_batches() + if request is None: + return None + armed: Final = request.prefetched.pop(_UPDATE_CACHE_PREFETCH_SLOT, None) + if not isinstance(armed, DeclaredBatchRead) or armed.keys != tuple(keys): + return None + values: Final = await cache.async_resolve_batch_get(armed) + return MappingProxyType({key: value for key, value in zip(keys, values) if value is not None}) + + +async def _read_update_cache_values( + keys: Sequence[str], parent_otel_span: Span | None, cache: DualCache | None = None +) -> Mapping[str, object]: """One batched read for every object ``update_cache`` refreshes; a failed read leaves them all untouched, exactly as a failed per-object GET left that object untouched.""" if not keys: return MappingProxyType({}) + target: Final = user_api_key_cache if cache is None else cache try: - values: Final = await user_api_key_cache.async_batch_get_cache( + armed: Final = await _take_armed_update_cache_read(keys, target) + if armed is not None: + return armed + values: Final = await target.async_batch_get_cache( keys=list(keys), parent_otel_span=parent_otel_span, throttle_redis=False ) except Exception as e: @@ -3817,7 +3893,7 @@ async def update_cache( values_to_update_in_cache: Final[list[tuple[str, object]]] = [] cached_values: Final = await _read_update_cache_values( - keys=_update_cache_read_keys( + keys=update_cache_read_keys( user_id=user_id, end_user_id=end_user_id, team_id=team_id, tags=tags, response_cost=response_cost ), parent_otel_span=parent_otel_span, diff --git a/litellm/router_strategy/lowest_tpm_rpm_v2.py b/litellm/router_strategy/lowest_tpm_rpm_v2.py index 909b47833cb..6e21d5d1f1f 100644 --- a/litellm/router_strategy/lowest_tpm_rpm_v2.py +++ b/litellm/router_strategy/lowest_tpm_rpm_v2.py @@ -304,7 +304,7 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): # update cache parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs) ## TPM - await self.router_cache.async_increment_cache( + await self.router_cache.async_increment_cache_post_call( key=tpm_key, value=total_tokens, ttl=self.routing_args.ttl, diff --git a/tests/unit/caching/test_request_redis_batch_post_call.py b/tests/unit/caching/test_request_redis_batch_post_call.py new file mode 100644 index 00000000000..2b5d3b3dbbb --- /dev/null +++ b/tests/unit/caching/test_request_redis_batch_post_call.py @@ -0,0 +1,661 @@ +"""One Redis pipeline per backend for the post-call writes of a request: spend counters, rate-limit token +scripts and slot releases, deployment TPM and the response-cache SET all ride the post-call batch, which +goes out once the success/failure callbacks have run (or at the deadline when no callback phase closes it).""" + +from __future__ import annotations + +import asyncio +import datetime +import hashlib +import json +from collections.abc import Awaitable, Callable, Mapping, Sequence +from unittest.mock import AsyncMock, MagicMock + +import pytest + +import litellm +from litellm.caching.caching import Cache +from litellm.caching.dual_cache import DualCache +from litellm.caching.in_memory_cache import InMemoryCache +from litellm.caching.redis_batch import ( + active_post_call_redis_batch, + active_request_redis_batches, + drain_post_call_redis_batches, + flush_post_call_redis_batches, + request_redis_batch_scope, +) +from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging +from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + PARALLEL_RELEASE_SCRIPT, + TOKEN_INCREMENT_SCRIPT, + ParallelSlotAcquisition, + RequestRateLimiterStash, + _PROXY_MaxParallelRequestsHandler_v3, +) +from litellm.proxy.spend_tracking.spend_counter_batch import PendingSpendIncrement +from litellm.proxy.utils import InternalUsageCache +from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2 +from litellm.types.caching import RedisPipelineIncrementOperation +from litellm.types.utils import ModelResponse + +from .test_redis_batch import FakeClient, FakeRedisCache + + +async def _script_outside_the_pipeline(keys: Sequence[str], args: Sequence[object]) -> object: + raise AssertionError("post-call scripts must ride the post-call pipeline") + + +class PostCallFakeRedisCache(FakeRedisCache): + """Records the direct (non-pipelined) writes an owner falls back to.""" + + def async_register_script(self, script: str) -> Callable[..., Awaitable[object]]: + return _script_outside_the_pipeline + + async def async_increment_pipeline( + self, increment_list: list[RedisPipelineIncrementOperation], **kwargs: object + ) -> list[float]: + return [await self.async_increment(op["key"], op["increment_value"]) for op in increment_list] + + async def async_delete_cache(self, key: str, **kwargs: object) -> None: # pyright: ignore[reportIncompatibleMethodOverride] # the fake drops RedisCache's unused kwargs + self.alone.append(("DEL", key)) + self.store.pop(key, None) + + async def async_set_cache(self, key: str, value: object, **kwargs: object) -> None: + self.alone.append(("SET", key, dict(kwargs))) + self.store[key] = value + + +def sha_of(script: str) -> str: + return hashlib.sha1(script.encode()).hexdigest() # noqa: S324 + + +def _ok_replies(command: tuple[object, ...]) -> object: + match command[0]: + case "INCRBYFLOAT": + return b"7.5" + case "EXPIRE": + return 1 + case "SET": + return True + case "EVALSHA": + return [3, 0] + case "MGET": + return [json.dumps({"spend": 1.0}) for _ in command[1:]] + raise AssertionError(command) + + +async def _run_ready_callbacks(client: FakeClient) -> None: + for _ in range(20): + if client.pipelines: + return + await asyncio.sleep(0) + + +def _names(client: FakeClient, index: int = 0) -> list[str]: + return [command[0] for command in client.pipelines[index].commands] + + +def _limiter(redis_cache: FakeRedisCache) -> _PROXY_MaxParallelRequestsHandler_v3: + dual_cache = DualCache() + dual_cache.attach_redis_cache(redis_cache) + return _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(dual_cache=dual_cache)) + + +def _slot_stash(slot_id: str, *counter_keys: str) -> RequestRateLimiterStash: + return RequestRateLimiterStash(parallel_slot=ParallelSlotAcquisition(slot_id=slot_id, counter_keys=list(counter_keys))) + + +def _token_ops(*keys: str) -> list[RedisPipelineIncrementOperation]: + return [RedisPipelineIncrementOperation(key=key, increment_value=10, ttl=60) for key in keys] + + +def _response_cache(redis_cache: FakeRedisCache) -> Cache: + cache = Cache(type="local") + cache.type = "redis" # pyright: ignore[reportAttributeAccessIssue] # the fake stands in for the Redis backend + cache.cache = redis_cache + return cache + + +def _tpm_router(redis_cache: FakeRedisCache) -> tuple[LowestTPMLoggingHandler_v2, DualCache]: + router_cache = DualCache() + router_cache.attach_redis_cache(redis_cache) + return LowestTPMLoggingHandler_v2(router_cache=router_cache, routing_args={"ttl": 60}), router_cache + + +def _tpm_kwargs() -> Mapping[str, object]: + return { + "standard_logging_object": { + "model_group": "gpt", + "model_id": "dep-a", + "hidden_params": {"litellm_model_name": "openai/gpt-4o-mini"}, + "total_tokens": 42, + }, + "litellm_params": {"metadata": {}}, + } + + +@pytest.mark.asyncio +async def test_every_post_call_owner_rides_one_pipeline_that_goes_out_when_the_callbacks_are_done(): + client = FakeClient(_ok_replies) + redis_cache = PostCallFakeRedisCache(client) + limiter = _limiter(redis_cache) + response_cache = _response_cache(redis_cache) + tpm, router_cache = _tpm_router(redis_cache) + + with request_redis_batch_scope(): + await response_cache.async_add_cache( + {"id": "resp"}, messages=[{"role": "user", "content": "hi"}], model="gpt", ttl=120 + ) + await tpm.async_log_success_event(_tpm_kwargs(), None, None, None) + await limiter.async_increment_tokens_with_ttl_preservation(_token_ops("{api_key:k1}:tokens")) + await limiter._release_stashed_parallel_slot( + _slot_stash("slot-1", "{api_key:k1}:parallel"), None, in_logging_callback=True + ) + assert client.pipelines == [] # nothing goes out while the callbacks are still declaring + await flush_post_call_redis_batches() + + assert len(client.pipelines) == 1 + assert _names(client) == ["SET", "INCRBYFLOAT", "EXPIRE", "EVALSHA", "EVALSHA"] + evalshas = [c for c in client.pipelines[0].commands if c[0] == "EVALSHA"] + assert [c[1] for c in evalshas] == [sha_of(TOKEN_INCREMENT_SCRIPT), sha_of(PARALLEL_RELEASE_SCRIPT)] + assert redis_cache.alone == [] + assert ( + await router_cache.in_memory_cache.async_get_cache( + next(k for k in router_cache.in_memory_cache.cache_dict if ":tpm:" in k) + ) + == 42 + ) + + +@pytest.mark.asyncio +async def test_the_response_cache_write_is_the_same_set_the_direct_path_issues(): + client = FakeClient(_ok_replies) + redis_cache = PostCallFakeRedisCache(client) + response_cache = _response_cache(redis_cache) + kwargs = {"messages": [{"role": "user", "content": "hi"}], "model": "gpt", "ttl": 120} + + with request_redis_batch_scope(): + await response_cache.async_add_cache({"id": "resp"}, **kwargs) + await flush_post_call_redis_batches() + + cache_key = response_cache.get_cache_key(**kwargs) + (command,) = client.pipelines[0].commands + assert (command[0], command[1], command[3]) == ("SET", cache_key, 120) + assert json.loads(command[2])["response"] == {"id": "resp"} + + +@pytest.mark.asyncio +async def test_a_chat_response_written_through_the_handler_dual_cache_lands_in_memory_and_rides_the_pipeline(): + client = FakeClient(_ok_replies) + redis_cache = PostCallFakeRedisCache(client) + response_cache = _response_cache(redis_cache) + handler_cache = DualCache(redis_cache=redis_cache, in_memory_cache=InMemoryCache()) + kwargs = {"messages": [{"role": "user", "content": "hi"}], "model": "gpt", "ttl": 120} + + with request_redis_batch_scope(): + await response_cache.async_add_cache('{"id": "resp"}', dynamic_cache_object=handler_cache, **kwargs) + cache_key = response_cache.get_cache_key(**kwargs) + in_memory = await handler_cache.in_memory_cache.async_get_cache(cache_key) + assert in_memory["response"] == '{"id": "resp"}' + assert redis_cache.alone == [] + await flush_post_call_redis_batches() + + (command,) = client.pipelines[0].commands + assert (command[0], command[1], command[3]) == ("SET", cache_key, 120) + + +@pytest.mark.asyncio +async def test_a_failed_operation_fails_only_its_owner_and_the_owner_applies_its_own_fallback(): + def replies(command: tuple[object, ...]) -> object: + if command[0] == "EVALSHA" and command[3] == "{api_key:k1}:tokens": + return Exception("ERR Lua") + return _ok_replies(command) + + client = FakeClient(replies) + redis_cache = PostCallFakeRedisCache(client) + limiter = _limiter(redis_cache) + response_cache = _response_cache(redis_cache) + + with request_redis_batch_scope(): + await response_cache.async_add_cache({"id": "resp"}, messages=[{"role": "user", "content": "hi"}], model="gpt") + await limiter.async_increment_tokens_with_ttl_preservation(_token_ops("{api_key:k1}:tokens")) + await limiter.async_increment_tokens_with_ttl_preservation(_token_ops("{team:t1}:tokens")) + await flush_post_call_redis_batches() + + assert len(client.pipelines) == 1 + # the failed group falls back to the plain increment (memory + Redis), the healthy group does not + assert redis_cache.alone == [("INCRBYFLOAT", "{api_key:k1}:tokens", 10)] + assert await limiter.internal_usage_cache.dual_cache.in_memory_cache.async_get_cache("{api_key:k1}:tokens") == 10 + assert await limiter.internal_usage_cache.dual_cache.in_memory_cache.async_get_cache("{team:t1}:tokens") is None + + +@pytest.mark.asyncio +async def test_a_failed_slot_release_script_releases_the_slot_in_memory(): + def replies(command: tuple[object, ...]) -> object: + if command[0] == "EVALSHA": + return Exception("ERR Lua") + return _ok_replies(command) + + redis_cache = PostCallFakeRedisCache(FakeClient(replies)) + limiter = _limiter(redis_cache) + memory = limiter.internal_usage_cache.dual_cache.in_memory_cache + await memory.async_set_cache("{api_key:k1}:parallel", {"slot-1": 1.0, "slot-2": 1.0}) + + with request_redis_batch_scope(): + await limiter._release_stashed_parallel_slot( + _slot_stash("slot-1", "{api_key:k1}:parallel"), None, in_logging_callback=True + ) + await flush_post_call_redis_batches() + + assert await memory.async_get_cache("{api_key:k1}:parallel") == {"slot-2": 1.0} + + +class DirectScriptFakeRedisCache(PostCallFakeRedisCache): + """Records the release script a pre-response caller runs outside the pipeline.""" + + def async_register_script(self, script: str) -> Callable[..., Awaitable[object]]: + async def run(keys: Sequence[str], args: Sequence[object]) -> object: + self.alone.append(("EVALSHA", tuple(keys), tuple(args))) + return [0 for _ in keys] + + return run + + +@pytest.mark.asyncio +async def test_a_slot_released_before_the_response_reaches_redis_at_once_not_on_the_pipeline(): + client = FakeClient(_ok_replies) + redis_cache = DirectScriptFakeRedisCache(client) + limiter = _limiter(redis_cache) + memory = limiter.internal_usage_cache.dual_cache.in_memory_cache + await memory.async_set_cache("{api_key:k1}:parallel", {"slot-1": 1.0}) + + with request_redis_batch_scope(): + await limiter._release_stashed_parallel_slot(_slot_stash("slot-1", "{api_key:k1}:parallel"), None) + assert redis_cache.alone == [("EVALSHA", ("{api_key:k1}:parallel",), ("slot-1",))] + assert await memory.async_get_cache("{api_key:k1}:parallel") == 0 + await flush_post_call_redis_batches() + + assert client.pipelines == [] + + +@pytest.mark.asyncio +async def test_a_deferred_response_cache_set_without_a_ttl_expires_in_redis_like_the_direct_path(): + client = FakeClient(_ok_replies) + redis_cache = PostCallFakeRedisCache(client) + dual_cache = DualCache(redis_cache=redis_cache, in_memory_cache=InMemoryCache(), default_in_memory_ttl=300) + + await dual_cache.async_set_cache("direct", {"id": "resp"}) + with request_redis_batch_scope(): + await dual_cache.async_set_cache_post_call("deferred", {"id": "resp"}, None) + await flush_post_call_redis_batches() + + (command,) = client.pipelines[0].commands + assert (command[0], command[1], command[3]) == ("SET", "deferred", redis_cache.alone[0][2]["ttl"]) + assert command[3] == 300 + + +@pytest.mark.asyncio +async def test_a_released_slot_is_free_locally_at_once_and_the_older_redis_count_does_not_overwrite_the_gauge(): + def replies(command: tuple[object, ...]) -> object: + if command[0] == "EVALSHA": + return [2] + return _ok_replies(command) + + redis_cache = PostCallFakeRedisCache(FakeClient(replies)) + limiter = _limiter(redis_cache) + memory = limiter.internal_usage_cache.dual_cache.in_memory_cache + await memory.async_set_cache("{api_key:k1}:parallel", {"slot-1": 1.0, "slot-2": 1.0, "slot-3": 1.0}) + + with request_redis_batch_scope(): + await limiter._release_stashed_parallel_slot( + _slot_stash("slot-1", "{api_key:k1}:parallel"), None, in_logging_callback=True + ) + assert await memory.async_get_cache("{api_key:k1}:parallel") == {"slot-2": 1.0, "slot-3": 1.0} + await memory.async_set_cache("{api_key:k1}:parallel", {"slot-2": 1.0, "slot-3": 1.0, "slot-4": 1.0}) + await flush_post_call_redis_batches() + + assert await memory.async_get_cache("{api_key:k1}:parallel") == {"slot-2": 1.0, "slot-3": 1.0, "slot-4": 1.0} + + +@pytest.mark.asyncio +async def test_failure_refunds_ride_the_post_call_pipeline_and_count_in_memory_at_once(): + client = FakeClient(_ok_replies) + dual_cache = DualCache() + dual_cache.attach_redis_cache(PostCallFakeRedisCache(client)) + refund = [RedisPipelineIncrementOperation(key="{api_key:k1}:tokens", increment_value=-500, ttl=60)] + + with request_redis_batch_scope(): + await dual_cache.async_increment_cache_pipeline_post_call(refund) + assert await dual_cache.in_memory_cache.async_get_cache("{api_key:k1}:tokens") == -500 + assert client.pipelines == [] + await flush_post_call_redis_batches() + + assert client.pipelines[0].commands[0] == ("INCRBYFLOAT", "{api_key:k1}:tokens", -500) + + +@pytest.mark.asyncio +async def test_outside_a_request_scope_owners_write_directly_as_before(): + client = FakeClient(_ok_replies) + redis_cache = PostCallFakeRedisCache(client) + dual_cache = DualCache() + dual_cache.attach_redis_cache(redis_cache) + response_cache = _response_cache(redis_cache) + + await dual_cache.async_increment_cache_post_call("dep:tpm", 42, ttl=60) + await response_cache.async_add_cache({"id": "resp"}, messages=[{"role": "user", "content": "hi"}], model="gpt") + + assert client.pipelines == [] + assert redis_cache.alone[0] == ("INCRBYFLOAT", "dep:tpm", 42) + assert active_post_call_redis_batch(redis_cache) is None + + +@pytest.mark.asyncio +async def test_a_set_with_options_keeps_the_direct_path(): + client = FakeClient(_ok_replies) + redis_cache = PostCallFakeRedisCache(client) + response_cache = _response_cache(redis_cache) + + with request_redis_batch_scope(): + await response_cache.async_add_cache( + {"id": "r"}, messages=[{"role": "user", "content": "hi"}], model="gpt", nx=True + ) + await flush_post_call_redis_batches() + + assert client.pipelines == [] + (direct_set,) = redis_cache.alone + assert direct_set[0] == "SET" and direct_set[2]["nx"] is True + + +@pytest.mark.asyncio +async def test_two_backends_get_one_post_call_pipeline_each(): + a_client, b_client = FakeClient(_ok_replies), FakeClient(_ok_replies) + a, b = DualCache(), DualCache() + a.attach_redis_cache(PostCallFakeRedisCache(a_client)) + b.attach_redis_cache(PostCallFakeRedisCache(b_client)) + + with request_redis_batch_scope(): + await a.async_increment_cache_post_call("x", 1, ttl=None) + await b.async_increment_cache_post_call("y", 1, ttl=None) + await a.async_increment_cache_post_call("z", 1, ttl=None) + await flush_post_call_redis_batches() + + assert len(a_client.pipelines) == 1 and len(b_client.pipelines) == 1 + assert [c[1] for c in a_client.pipelines[0].commands if c[0] == "INCRBYFLOAT"] == ["x", "z"] + + +@pytest.mark.asyncio +async def test_a_numeric_string_ttl_reaches_redis_as_the_direct_path_would_send_it(): + client = FakeClient(_ok_replies) + response_cache = _response_cache(PostCallFakeRedisCache(client)) + kwargs = {"messages": [{"role": "user", "content": "hi"}], "model": "gpt", "ttl": "3600"} + + with request_redis_batch_scope(): + await response_cache.async_add_cache({"id": "resp"}, **kwargs) + await flush_post_call_redis_batches() + + (command,) = client.pipelines[0].commands + assert (command[0], command[3]) == ("SET", 3600) + + +@pytest.mark.asyncio +async def test_post_call_writes_still_waiting_on_their_callbacks_are_drained_at_shutdown(): + client = FakeClient(_ok_replies) + dual_cache = DualCache() + dual_cache.attach_redis_cache(PostCallFakeRedisCache(client)) + + with request_redis_batch_scope(post_call_deadline=60) as request: + await dual_cache.async_increment_cache_post_call("x", 1, ttl=None) + await request.flush_all() + assert client.pipelines == [] + + await drain_post_call_redis_batches() + assert len(client.pipelines) == 1 and _names(client) == ["INCRBYFLOAT"] + + await drain_post_call_redis_batches() + assert len(client.pipelines) == 1 + + +@pytest.mark.asyncio +async def test_a_post_call_batch_nobody_closes_goes_out_at_the_deadline(monkeypatch: pytest.MonkeyPatch): + client = FakeClient(_ok_replies) + dual_cache = DualCache() + dual_cache.attach_redis_cache(PostCallFakeRedisCache(client)) + + loop = asyncio.get_running_loop() + armed_at = loop.time() + + with request_redis_batch_scope(post_call_deadline=60) as request: + await dual_cache.async_increment_cache_post_call("x", 1, ttl=None) + await request.flush_all() + await _run_ready_callbacks(client) + assert client.pipelines == [], "the request boundary drains the immediate batch, not the post-call one" + + monkeypatch.setattr(loop, "time", lambda: armed_at + 61) + await _run_ready_callbacks(client) + + assert len(client.pipelines) == 1 and _names(client) == ["INCRBYFLOAT"] + + +@pytest.mark.asyncio +async def test_the_success_handler_closes_the_post_call_batch_after_the_last_callback(monkeypatch): + client = FakeClient(_ok_replies) + dual_cache = DualCache() + dual_cache.attach_redis_cache(PostCallFakeRedisCache(client)) + pipelines_seen_by_callbacks: list[int] = [] + + class Counter(CustomLogger): + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + await dual_cache.async_increment_cache_post_call("counted", 1, ttl=None) + pipelines_seen_by_callbacks.append(len(client.pipelines)) + + monkeypatch.setattr(litellm, "_async_success_callback", []) + logging_obj = LitellmLogging( + model="test-model", + messages=[], + stream=False, + call_type="completion", + start_time=datetime.datetime.now(), + litellm_call_id="post-call", + function_id="post-call", + dynamic_async_success_callbacks=[Counter(), Counter()], + ) + logging_obj.update_environment_variables(litellm_params={"metadata": {}}, optional_params={}) + payload = { + "id": "post-call", + "call_type": "completion", + "metadata": {}, + "model_group": "test-model", + "model_parameters": {}, + } + + with request_redis_batch_scope(): + await logging_obj.async_success_handler(result=ModelResponse(), standard_logging_object=payload) + + assert pipelines_seen_by_callbacks == [0, 0] + assert len(client.pipelines) == 1 and _names(client) == ["INCRBYFLOAT", "INCRBYFLOAT"] + + +@pytest.mark.asyncio +async def test_spend_counter_increments_ride_the_pipeline_and_settle_into_memory(monkeypatch): + from litellm.proxy import proxy_server + + client = FakeClient(_ok_replies) + spend_cache = DualCache() + spend_cache.attach_redis_cache(PostCallFakeRedisCache(client)) + monkeypatch.setattr(proxy_server, "spend_counter_cache", spend_cache) + pending = [PendingSpendIncrement("spend:key:k1", 0.5), PendingSpendIncrement("spend:team:t1", 0.5)] + + with request_redis_batch_scope(): + await proxy_server._apply_spend_counter_increments(pending) + assert client.pipelines == [] + await flush_post_call_redis_batches() + + assert [c for c in client.pipelines[0].commands if c[0] == "INCRBYFLOAT"] == [ + ("INCRBYFLOAT", "spend:key:k1", 0.5), + ("INCRBYFLOAT", "spend:team:t1", 0.5), + ] + assert spend_cache.in_memory_cache.get_cache("spend:key:k1") == 7.5 + + +@pytest.mark.asyncio +async def test_a_spend_counter_whose_increment_failed_is_invalidated_not_trusted(monkeypatch): + from litellm.proxy import proxy_server + + def replies(command: tuple[object, ...]) -> object: + if command[0] == "INCRBYFLOAT" and command[1] == "spend:key:k1": + return Exception("OOM") + return _ok_replies(command) + + redis_cache = PostCallFakeRedisCache(FakeClient(replies)) + spend_cache = DualCache() + spend_cache.attach_redis_cache(redis_cache) + spend_cache.in_memory_cache.set_cache("spend:key:k1", 3.0) + spend_cache.in_memory_cache.set_cache("spend:team:t1", 3.0) + monkeypatch.setattr(proxy_server, "spend_counter_cache", spend_cache) + + with request_redis_batch_scope(): + await proxy_server._apply_spend_counter_increments( + [PendingSpendIncrement("spend:key:k1", 0.5), PendingSpendIncrement("spend:team:t1", 0.5)] + ) + await flush_post_call_redis_batches() + + assert spend_cache.in_memory_cache.get_cache("spend:key:k1") is None + assert redis_cache.alone == [("DEL", "spend:key:k1")] + assert spend_cache.in_memory_cache.get_cache("spend:team:t1") == 7.5 + + +@pytest.mark.asyncio +async def test_a_cancelled_post_call_flush_keeps_the_shared_spend_counter_and_counts_the_spend_locally(monkeypatch): + from litellm.proxy import proxy_server + + redis_cache = PostCallFakeRedisCache( + FakeClient(_ok_replies, fail=asyncio.CancelledError()) # pyright: ignore[reportArgumentType] # a cancel raised mid-pipeline + ) + spend_cache = DualCache() + spend_cache.attach_redis_cache(redis_cache) + spend_cache.in_memory_cache.set_cache("spend:key:k1", 3.0) + monkeypatch.setattr(proxy_server, "spend_counter_cache", spend_cache) + + with request_redis_batch_scope(): + await proxy_server._apply_spend_counter_increments( + [PendingSpendIncrement("spend:key:k1", 0.5), PendingSpendIncrement("spend:team:t1", 0.5)] + ) + with pytest.raises(asyncio.CancelledError): + await flush_post_call_redis_batches() + + assert redis_cache.alone == [], "a cancel says nothing about the shared counter, so Redis keeps it" + assert spend_cache.in_memory_cache.get_cache("spend:key:k1") == 3.5, "the local copy counts the cancelled spend" + assert spend_cache.in_memory_cache.get_cache("spend:team:t1") is None, "an absent local copy is not seeded" + + +@pytest.mark.asyncio +async def test_the_update_cache_read_armed_before_accounting_rides_the_pipeline_of_the_reconcile_read(): + from litellm.proxy.proxy_server import _read_update_cache_values, arm_update_cache_read + + client = FakeClient(_ok_replies) + redis_cache = PostCallFakeRedisCache(client) + cache = DualCache() + cache.attach_redis_cache(redis_cache) + keys = ["user-1", "team_id:t1"] + + with request_redis_batch_scope() as request: + await arm_update_cache_read(keys, cache=cache) + assert client.pipelines == [] + await request.batch(redis_cache).mget(["spend:key:k1"]) # the spend reconcile read of the same request + values = await _read_update_cache_values(keys, None, cache=cache) + + assert len(client.pipelines) == 1 + assert client.pipelines[0].commands == [("MGET", "user-1", "team_id:t1"), ("MGET", "spend:key:k1")] + assert values == {"user-1": {"spend": 1.0}, "team_id:t1": {"spend": 1.0}} + assert redis_cache.alone == [] + assert active_request_redis_batches() is None + + +@pytest.mark.asyncio +async def test_an_update_cache_read_armed_for_other_keys_is_ignored_and_the_read_happens_as_before(): + from litellm.proxy.proxy_server import _read_update_cache_values, arm_update_cache_read + + redis_cache = PostCallFakeRedisCache(FakeClient(_ok_replies)) + redis_cache.store["team_id:t1"] = {"spend": 2.0} + cache = DualCache() + cache.attach_redis_cache(redis_cache) + + with request_redis_batch_scope(): + await arm_update_cache_read(["user-1"], cache=cache) + values = await _read_update_cache_values(["team_id:t1"], None, cache=cache) + + assert values == {"team_id:t1": {"spend": 2.0}} + assert ("MGET", ("team_id:t1",)) in redis_cache.alone + + +@pytest.mark.asyncio +async def test_the_update_cache_read_sees_a_cached_spend_written_while_the_spend_was_persisted(monkeypatch): + from litellm.proxy import proxy_server + from litellm.proxy.hooks.proxy_track_cost_callback import _update_database_and_spend_counters + + cached_user_spend = {"user-1": 1.0} + + def replies(command: tuple[object, ...]) -> object: + if command[0] == "MGET": + return [ + json.dumps({"spend": cached_user_spend[key]}) if key in cached_user_spend else b"0.5" + for key in command[1:] + ] + return _ok_replies(command) + + client = FakeClient(replies) + redis_cache = PostCallFakeRedisCache(client) + spend_cache = DualCache() + spend_cache.attach_redis_cache(redis_cache) + user_cache = DualCache() + user_cache.attach_redis_cache(redis_cache) + monkeypatch.setattr(proxy_server, "spend_counter_cache", spend_cache) + monkeypatch.setattr(proxy_server, "user_api_key_cache", user_cache) + + async def _read_on_the_request_pipeline_then_a_concurrent_callback_writes_the_user(**kwargs: object) -> bool: + request = active_request_redis_batches() + assert request is not None + await request.batch(redis_cache).mget(["key-object"]) + cached_user_spend["user-1"] = 5.0 + return True + + proxy_logging_obj = MagicMock() + proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock( + side_effect=_read_on_the_request_pipeline_then_a_concurrent_callback_writes_the_user + ) + reservation = { + "reserved_cost": 0.5, + "entries": [ + { + "counter_key": "spend:key:k1", + "entity_type": "Key", + "entity_id": "k1", + "reserved_cost": 0.5, + "applied_adjustment": 0.0, + } + ], + "finalized": False, + } + + with request_redis_batch_scope(): + charged = await _update_database_and_spend_counters( + proxy_logging_obj=proxy_logging_obj, + increment_spend_counters=proxy_server.increment_spend_counters, + user_api_key="k1", + user_id="user-1", + end_user_id=None, + team_id=None, + org_id=None, + kwargs={}, + completion_response=None, + start_time=datetime.datetime.now(), + end_time=datetime.datetime.now(), + response_cost=0.2, + budget_reservation=reservation, + update_cache_read_keys=("user-1",), + ) + values = await proxy_server._read_update_cache_values(("user-1",), None) + + assert charged is True + assert values == {"user-1": {"spend": 5.0}}, client.pipelines