From dbc0d23c1ecea6450e248b3d2b846c28a25b2869 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 29 Jul 2026 18:38:29 -0700 Subject: [PATCH 1/8] fix(vertex_ai): skip context caching when the cached block ends on a model turn --- .../context_caching/transformation.py | 13 +- .../vertex_ai_context_caching.py | 15 ++ .../test_vertex_ai_context_caching.py | 129 ++++++++++++++++++ 3 files changed, 156 insertions(+), 1 deletion(-) diff --git a/litellm/llms/vertex_ai/context_caching/transformation.py b/litellm/llms/vertex_ai/context_caching/transformation.py index f0ce3323ef6..a74e0c97abc 100644 --- a/litellm/llms/vertex_ai/context_caching/transformation.py +++ b/litellm/llms/vertex_ai/context_caching/transformation.py @@ -5,7 +5,7 @@ Why separate file? Make it easy to see how transformation works """ import re -from typing import List, Optional, Tuple, Literal +from typing import List, Optional, Sequence, Tuple, Literal from litellm.types.llms.openai import AllMessageValues from litellm.types.llms.vertex_ai import CachedContentRequestBody @@ -152,6 +152,17 @@ def separate_cached_messages( return cached_messages, non_cached_messages +def cached_messages_end_on_supported_turn(cached_messages: Sequence[AllMessageValues]) -> bool: + """ + The cachedContents API rejects contents ending on a model turn, which is how it + classifies both assistant messages and tool results, with HTTP 400 + "Requests ending with a model turn are not supported". + """ + if not cached_messages: + return False + return cached_messages[-1].get("role") not in ("assistant", "tool", "function") + + def transform_openai_messages_to_gemini_context_caching( model: str, messages: List[AllMessageValues], diff --git a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py index 0bf3715f798..fe4cd4ec451 100644 --- a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py +++ b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py @@ -22,6 +22,7 @@ from litellm.types.llms.vertex_ai import ( from ..common_utils import VertexAIError, get_vertex_base_url from ..vertex_llm_base import VertexBase from .transformation import ( + cached_messages_end_on_supported_turn, separate_cached_messages, transform_openai_messages_to_gemini_context_caching, ) @@ -308,6 +309,13 @@ class ContextCachingEndpoints(VertexBase): if len(cached_messages) == 0: return messages, optional_params, None + if not cached_messages_end_on_supported_turn(cached_messages): + verbose_logger.debug( + "Vertex AI context caching: cached message block ends on an assistant or " + "tool turn, which the cachedContents API rejects. Skipping context caching." + ) + return messages, optional_params, None + # Gemini requires a minimum of 1024 tokens for context caching. # Skip caching if the cached content is too small to avoid API errors. if not is_prompt_caching_valid_prompt( @@ -459,6 +467,13 @@ class ContextCachingEndpoints(VertexBase): if len(cached_messages) == 0: return messages, optional_params, None + if not cached_messages_end_on_supported_turn(cached_messages): + verbose_logger.debug( + "Vertex AI context caching: cached message block ends on an assistant or " + "tool turn, which the cachedContents API rejects. Skipping context caching." + ) + return messages, optional_params, None + # Gemini requires a minimum of 1024 tokens for context caching. # Skip caching if the cached content is too small to avoid API errors. if not is_prompt_caching_valid_prompt( diff --git a/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py b/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py index cf75964ddb7..1aa724e551e 100644 --- a/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py +++ b/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py @@ -1452,6 +1452,135 @@ class TestContextCachingEndpoints: # Restart the patcher so teardown_method can stop it cleanly self._token_check_patcher.start() + def _model_turn_final_messages(self, final_cached_role): + tool_call = { + "id": "call_abc123", + "type": "function", + "function": {"name": "get_weather", "arguments": '{"location": "Boston"}'}, + } + cached_tail = ( + [ + { + "role": "tool", + "tool_call_id": "call_abc123", + "content": "72F and sunny", + "cache_control": {"type": "ephemeral"}, + } + ] + if final_cached_role == "tool" + else [] + ) + return [ + { + "role": "user", + "content": [ + { + "type": "text", + "text": "Use the weather tool for every answer.", + "cache_control": {"type": "ephemeral"}, + } + ], + }, + { + "role": "assistant", + "content": "", + "tool_calls": [tool_call], + "cache_control": {"type": "ephemeral"}, + }, + *cached_tail, + {"role": "user", "content": "What is the weather in Boston?"}, + ] + + @pytest.mark.parametrize("final_cached_role", ["assistant", "tool"]) + def test_check_and_create_cache_skips_when_cached_block_ends_on_model_turn( + self, final_cached_role + ): + """The cachedContents API rejects contents ending on an assistant or tool turn + with HTTP 400 "Requests ending with a model turn are not supported", so the + request must proceed uncached instead of failing. + """ + all_messages = self._model_turn_final_messages(final_cached_role) + optional_params = self.sample_optional_params.copy() + + result = self.context_caching.check_and_create_cache( + messages=all_messages, + optional_params=optional_params, + api_key="test_key", + api_base=None, + model="gemini-3.6-flash", + client=self.mock_client, + timeout=30.0, + logging_obj=self.mock_logging, + cached_content=None, + custom_llm_provider="vertex_ai", + vertex_project="test_project", + vertex_location="us-central1", + vertex_auth_header="test_token", + ) + + messages, returned_params, returned_cache = result + assert messages == all_messages + assert returned_cache is None + assert "tools" in returned_params + self.mock_client.get.assert_not_called() + self.mock_client.post.assert_not_called() + + @pytest.mark.parametrize("final_cached_role", ["assistant", "tool"]) + @pytest.mark.asyncio + async def test_async_check_and_create_cache_skips_when_cached_block_ends_on_model_turn( + self, final_cached_role + ): + """Async variant: an unsupported terminal turn skips caching instead of failing.""" + all_messages = self._model_turn_final_messages(final_cached_role) + optional_params = self.sample_optional_params.copy() + + result = await self.context_caching.async_check_and_create_cache( + messages=all_messages, + optional_params=optional_params, + api_key="test_key", + api_base=None, + model="gemini-3.6-flash", + client=self.mock_async_client, + timeout=30.0, + logging_obj=self.mock_logging, + cached_content=None, + custom_llm_provider="vertex_ai", + vertex_project="test_project", + vertex_location="us-central1", + vertex_auth_header="test_token", + ) + + messages, returned_params, returned_cache = result + assert messages == all_messages + assert returned_cache is None + assert "tools" in returned_params + self.mock_async_client.get.assert_not_called() + self.mock_async_client.post.assert_not_called() + + +def test_cached_messages_end_on_supported_turn(): + from litellm.llms.vertex_ai.context_caching.transformation import ( + cached_messages_end_on_supported_turn, + ) + + assert ( + cached_messages_end_on_supported_turn( + [{"role": "assistant", "content": "hi"}, {"role": "user", "content": "hello"}] + ) + is True + ) + assert cached_messages_end_on_supported_turn([{"role": "system", "content": "be brief"}]) is True + assert cached_messages_end_on_supported_turn([{"role": "assistant", "content": "hi"}]) is False + assert ( + cached_messages_end_on_supported_turn([{"role": "tool", "tool_call_id": "x", "content": "y"}]) + is False + ) + assert ( + cached_messages_end_on_supported_turn([{"role": "function", "name": "f", "content": "y"}]) + is False + ) + assert cached_messages_end_on_supported_turn([]) is False + class TestCheckCachePagination: """Test pagination logic in check_cache and async_check_cache methods.""" From e0946ccf0d02bb558cd6aab98815ae30d525890f Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 30 Jul 2026 02:43:54 +0000 Subject: [PATCH 2/8] fix(pricing): correct gpt-5.4-mini and gpt-5.4-nano token limits gpt-5.4-mini and gpt-5.4-nano are 400K-context models (272K input, 128K output), but their cost map entries carried gpt-5.4's 1.05M window. The router's pre-call context window check therefore admitted prompts far past what the models accept, so oversized requests were dispatched to the provider and failed there instead of being caught locally or routed through context_window_fallbacks. The azure_ai entries also inherited gpt-5.4's above-272K tiered pricing. OpenAI applies that surcharge to the 1.05M-window models only, so those keys are removed. Limits per OpenAI's model reference and Azure AI Foundry's model table: gpt-5.4-mini and gpt-5.4-nano are 400,000 context / 272,000 input / 128,000 output --- ...odel_prices_and_context_window_backup.json | 40 ++------- model_prices_and_context_window.json | 48 +++-------- .../test_gpt_5_4_model_metadata.py | 81 +++++++++++++++++++ 3 files changed, 101 insertions(+), 68 deletions(-) create mode 100644 tests/test_litellm/test_gpt_5_4_model_metadata.py diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 87d9b6afc18..cc418c9c428 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -3454,22 +3454,16 @@ }, "azure_ai/gpt-5.4-mini": { "cache_read_input_token_cost": 7.5e-08, - "cache_read_input_token_cost_above_272k_tokens": 1.5e-07, "cache_read_input_token_cost_priority": 1.5e-07, - "cache_read_input_token_cost_above_272k_tokens_priority": 3e-07, "input_cost_per_token": 7.5e-07, - "input_cost_per_token_above_272k_tokens": 1.5e-06, "input_cost_per_token_priority": 1.5e-06, - "input_cost_per_token_above_272k_tokens_priority": 3e-06, "litellm_provider": "azure_ai", - "max_input_tokens": 400000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 4.5e-06, - "output_cost_per_token_above_272k_tokens": 6.75e-06, "output_cost_per_token_priority": 9e-06, - "output_cost_per_token_above_272k_tokens_priority": 1.35e-05, "source": "https://ai.azure.com/catalog/models/gpt-5.4-mini", "supported_endpoints": [ "/v1/chat/completions", @@ -3500,22 +3494,16 @@ }, "azure_ai/gpt-5.4-mini-2026-03-17": { "cache_read_input_token_cost": 7.5e-08, - "cache_read_input_token_cost_above_272k_tokens": 1.5e-07, "cache_read_input_token_cost_priority": 1.5e-07, - "cache_read_input_token_cost_above_272k_tokens_priority": 3e-07, "input_cost_per_token": 7.5e-07, - "input_cost_per_token_above_272k_tokens": 1.5e-06, "input_cost_per_token_priority": 1.5e-06, - "input_cost_per_token_above_272k_tokens_priority": 3e-06, "litellm_provider": "azure_ai", - "max_input_tokens": 400000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 4.5e-06, - "output_cost_per_token_above_272k_tokens": 6.75e-06, "output_cost_per_token_priority": 9e-06, - "output_cost_per_token_above_272k_tokens_priority": 1.35e-05, "source": "https://ai.azure.com/catalog/models/gpt-5.4-mini", "supported_endpoints": [ "/v1/chat/completions", @@ -3546,22 +3534,16 @@ }, "azure_ai/gpt-5.4-nano": { "cache_read_input_token_cost": 2e-08, - "cache_read_input_token_cost_above_272k_tokens": 4e-08, "cache_read_input_token_cost_priority": 4e-08, - "cache_read_input_token_cost_above_272k_tokens_priority": 8e-08, "input_cost_per_token": 2e-07, - "input_cost_per_token_above_272k_tokens": 4e-07, "input_cost_per_token_priority": 4e-07, - "input_cost_per_token_above_272k_tokens_priority": 8e-07, "litellm_provider": "azure_ai", - "max_input_tokens": 400000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 1.25e-06, - "output_cost_per_token_above_272k_tokens": 1.875e-06, "output_cost_per_token_priority": 2.5e-06, - "output_cost_per_token_above_272k_tokens_priority": 3.75e-06, "source": "https://ai.azure.com/catalog/models/gpt-5.4-nano", "supported_endpoints": [ "/v1/chat/completions", @@ -3592,22 +3574,16 @@ }, "azure_ai/gpt-5.4-nano-2026-03-17": { "cache_read_input_token_cost": 2e-08, - "cache_read_input_token_cost_above_272k_tokens": 4e-08, "cache_read_input_token_cost_priority": 4e-08, - "cache_read_input_token_cost_above_272k_tokens_priority": 8e-08, "input_cost_per_token": 2e-07, - "input_cost_per_token_above_272k_tokens": 4e-07, "input_cost_per_token_priority": 4e-07, - "input_cost_per_token_above_272k_tokens_priority": 8e-07, "litellm_provider": "azure_ai", - "max_input_tokens": 400000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 1.25e-06, - "output_cost_per_token_above_272k_tokens": 1.875e-06, "output_cost_per_token_priority": 2.5e-06, - "output_cost_per_token_above_272k_tokens_priority": 3.75e-06, "source": "https://ai.azure.com/catalog/models/gpt-5.4-nano", "supported_endpoints": [ "/v1/chat/completions", @@ -7201,7 +7177,7 @@ "cache_read_input_token_cost": 7.5e-08, "input_cost_per_token": 7.5e-07, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -7236,7 +7212,7 @@ "cache_read_input_token_cost": 7.5e-08, "input_cost_per_token": 7.5e-07, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -7271,7 +7247,7 @@ "cache_read_input_token_cost": 2e-08, "input_cost_per_token": 2e-07, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -7306,7 +7282,7 @@ "cache_read_input_token_cost": 2e-08, "input_cost_per_token": 2e-07, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 0edd3bd5f30..c4628fecdb8 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -3454,22 +3454,16 @@ }, "azure_ai/gpt-5.4-mini": { "cache_read_input_token_cost": 7.5e-08, - "cache_read_input_token_cost_above_272k_tokens": 1.5e-07, "cache_read_input_token_cost_priority": 1.5e-07, - "cache_read_input_token_cost_above_272k_tokens_priority": 3e-07, "input_cost_per_token": 7.5e-07, - "input_cost_per_token_above_272k_tokens": 1.5e-06, "input_cost_per_token_priority": 1.5e-06, - "input_cost_per_token_above_272k_tokens_priority": 3e-06, "litellm_provider": "azure_ai", - "max_input_tokens": 400000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 4.5e-06, - "output_cost_per_token_above_272k_tokens": 6.75e-06, "output_cost_per_token_priority": 9e-06, - "output_cost_per_token_above_272k_tokens_priority": 1.35e-05, "source": "https://ai.azure.com/catalog/models/gpt-5.4-mini", "supported_endpoints": [ "/v1/chat/completions", @@ -3500,22 +3494,16 @@ }, "azure_ai/gpt-5.4-mini-2026-03-17": { "cache_read_input_token_cost": 7.5e-08, - "cache_read_input_token_cost_above_272k_tokens": 1.5e-07, "cache_read_input_token_cost_priority": 1.5e-07, - "cache_read_input_token_cost_above_272k_tokens_priority": 3e-07, "input_cost_per_token": 7.5e-07, - "input_cost_per_token_above_272k_tokens": 1.5e-06, "input_cost_per_token_priority": 1.5e-06, - "input_cost_per_token_above_272k_tokens_priority": 3e-06, "litellm_provider": "azure_ai", - "max_input_tokens": 400000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 4.5e-06, - "output_cost_per_token_above_272k_tokens": 6.75e-06, "output_cost_per_token_priority": 9e-06, - "output_cost_per_token_above_272k_tokens_priority": 1.35e-05, "source": "https://ai.azure.com/catalog/models/gpt-5.4-mini", "supported_endpoints": [ "/v1/chat/completions", @@ -3546,22 +3534,16 @@ }, "azure_ai/gpt-5.4-nano": { "cache_read_input_token_cost": 2e-08, - "cache_read_input_token_cost_above_272k_tokens": 4e-08, "cache_read_input_token_cost_priority": 4e-08, - "cache_read_input_token_cost_above_272k_tokens_priority": 8e-08, "input_cost_per_token": 2e-07, - "input_cost_per_token_above_272k_tokens": 4e-07, "input_cost_per_token_priority": 4e-07, - "input_cost_per_token_above_272k_tokens_priority": 8e-07, "litellm_provider": "azure_ai", - "max_input_tokens": 400000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 1.25e-06, - "output_cost_per_token_above_272k_tokens": 1.875e-06, "output_cost_per_token_priority": 2.5e-06, - "output_cost_per_token_above_272k_tokens_priority": 3.75e-06, "source": "https://ai.azure.com/catalog/models/gpt-5.4-nano", "supported_endpoints": [ "/v1/chat/completions", @@ -3592,22 +3574,16 @@ }, "azure_ai/gpt-5.4-nano-2026-03-17": { "cache_read_input_token_cost": 2e-08, - "cache_read_input_token_cost_above_272k_tokens": 4e-08, "cache_read_input_token_cost_priority": 4e-08, - "cache_read_input_token_cost_above_272k_tokens_priority": 8e-08, "input_cost_per_token": 2e-07, - "input_cost_per_token_above_272k_tokens": 4e-07, "input_cost_per_token_priority": 4e-07, - "input_cost_per_token_above_272k_tokens_priority": 8e-07, "litellm_provider": "azure_ai", - "max_input_tokens": 400000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 1.25e-06, - "output_cost_per_token_above_272k_tokens": 1.875e-06, "output_cost_per_token_priority": 2.5e-06, - "output_cost_per_token_above_272k_tokens_priority": 3.75e-06, "source": "https://ai.azure.com/catalog/models/gpt-5.4-nano", "supported_endpoints": [ "/v1/chat/completions", @@ -7201,7 +7177,7 @@ "cache_read_input_token_cost": 7.5e-08, "input_cost_per_token": 7.5e-07, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -7236,7 +7212,7 @@ "cache_read_input_token_cost": 7.5e-08, "input_cost_per_token": 7.5e-07, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -7271,7 +7247,7 @@ "cache_read_input_token_cost": 2e-08, "input_cost_per_token": 2e-07, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -7306,7 +7282,7 @@ "cache_read_input_token_cost": 2e-08, "input_cost_per_token": 2e-07, "litellm_provider": "azure", - "max_input_tokens": 1050000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -24364,7 +24340,7 @@ "input_cost_per_token_batches": 3.75e-07, "input_cost_per_token_priority": 1.5e-06, "litellm_provider": "openai", - "max_input_tokens": 1050000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -24410,7 +24386,7 @@ "input_cost_per_token_batches": 3.75e-07, "input_cost_per_token_priority": 1.5e-06, "litellm_provider": "openai", - "max_input_tokens": 1050000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -24454,7 +24430,7 @@ "input_cost_per_token_flex": 1e-07, "input_cost_per_token_batches": 1e-07, "litellm_provider": "openai", - "max_input_tokens": 1050000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -24497,7 +24473,7 @@ "input_cost_per_token_flex": 1e-07, "input_cost_per_token_batches": 1e-07, "litellm_provider": "openai", - "max_input_tokens": 1050000, + "max_input_tokens": 272000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", diff --git a/tests/test_litellm/test_gpt_5_4_model_metadata.py b/tests/test_litellm/test_gpt_5_4_model_metadata.py new file mode 100644 index 00000000000..f93e6187dcb --- /dev/null +++ b/tests/test_litellm/test_gpt_5_4_model_metadata.py @@ -0,0 +1,81 @@ +import json +from functools import lru_cache +from pathlib import Path + +import pytest + +REPO_ROOT = Path(__file__).parents[2] +MAIN_PATH = REPO_ROOT / "model_prices_and_context_window.json" +BACKUP_PATH = REPO_ROOT / "litellm" / "model_prices_and_context_window_backup.json" + +DOCUMENTED_MAX_INPUT_TOKENS = 272000 +DOCUMENTED_MAX_OUTPUT_TOKENS = 128000 + +SMALL_MODEL_NAMES = ( + "gpt-5.4-mini", + "gpt-5.4-mini-2026-03-17", + "gpt-5.4-nano", + "gpt-5.4-nano-2026-03-17", +) +SMALL_MODELS = tuple(f"{prefix}{name}" for prefix in ("", "azure/", "azure_ai/") for name in SMALL_MODEL_NAMES) + +STANDARD_PRICING = { + "gpt-5.4-mini": (7.5e-07, 4.5e-06, 7.5e-08), + "gpt-5.4-nano": (2e-07, 1.25e-06, 2e-08), +} + +LONG_CONTEXT_MODELS = ("gpt-5.4", "gpt-5.4-pro") + + +@lru_cache(maxsize=2) +def _load(path: Path) -> dict[str, dict[str, object]]: + with open(path) as f: + return json.load(f) + + +def _pricing_key(model: str) -> str: + return "gpt-5.4-nano" if "nano" in model else "gpt-5.4-mini" + + +@pytest.mark.parametrize("model", SMALL_MODELS) +def test_gpt_5_4_small_models_use_documented_token_limits(model: str) -> None: + """gpt-5.4-mini/nano are 400K-window models: 272K in, 128K out, not gpt-5.4's 1.05M window.""" + info = _load(MAIN_PATH).get(model) + assert info is not None, f"{model} not found in model_prices_and_context_window.json" + + assert info["max_input_tokens"] == DOCUMENTED_MAX_INPUT_TOKENS + assert info["max_output_tokens"] == DOCUMENTED_MAX_OUTPUT_TOKENS + assert info["max_tokens"] == DOCUMENTED_MAX_OUTPUT_TOKENS + + +@pytest.mark.parametrize("model", SMALL_MODELS) +def test_gpt_5_4_small_models_have_no_long_context_surcharge(model: str) -> None: + """OpenAI prices prompts above 272K at 2x input / 1.5x output for the 1.05M-window models only.""" + info = _load(MAIN_PATH)[model] + assert [key for key in info if "above_272k" in key] == [] + + +@pytest.mark.parametrize("model", SMALL_MODELS) +def test_gpt_5_4_small_models_standard_pricing(model: str) -> None: + info = _load(MAIN_PATH)[model] + input_cost, output_cost, cache_read_cost = STANDARD_PRICING[_pricing_key(model)] + + assert info["input_cost_per_token"] == input_cost + assert info["output_cost_per_token"] == output_cost + assert info["cache_read_input_token_cost"] == cache_read_cost + + +@pytest.mark.parametrize("model", LONG_CONTEXT_MODELS) +def test_gpt_5_4_long_context_models_keep_surcharge(model: str) -> None: + """The mini/nano correction must leave gpt-5.4 and gpt-5.4-pro tiered pricing intact.""" + info = _load(MAIN_PATH)[model] + + assert info["input_cost_per_token_above_272k_tokens"] == pytest.approx(info["input_cost_per_token"] * 2) + assert info["output_cost_per_token_above_272k_tokens"] == pytest.approx(info["output_cost_per_token"] * 1.5) + + +@pytest.mark.parametrize("model", SMALL_MODELS) +def test_gpt_5_4_small_models_backup_matches_main(model: str) -> None: + assert _load(BACKUP_PATH).get(model) == _load(MAIN_PATH).get(model), ( + f"{model} differs between main and backup model cost maps" + ) From 819dc7812af842b7b7844dc7c794c302d8e1075a Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 29 Jul 2026 19:47:26 -0700 Subject: [PATCH 3/8] fix(vertex_ai): evaluate cached-block terminal turn after system extraction --- .../context_caching/transformation.py | 11 ++++-- .../vertex_ai_context_caching.py | 10 +++-- .../test_vertex_ai_context_caching.py | 38 +++++++++++++++---- 3 files changed, 43 insertions(+), 16 deletions(-) diff --git a/litellm/llms/vertex_ai/context_caching/transformation.py b/litellm/llms/vertex_ai/context_caching/transformation.py index a74e0c97abc..36c78974aca 100644 --- a/litellm/llms/vertex_ai/context_caching/transformation.py +++ b/litellm/llms/vertex_ai/context_caching/transformation.py @@ -156,11 +156,14 @@ def cached_messages_end_on_supported_turn(cached_messages: Sequence[AllMessageVa """ The cachedContents API rejects contents ending on a model turn, which is how it classifies both assistant messages and tool results, with HTTP 400 - "Requests ending with a model turn are not supported". + "Requests ending with a model turn are not supported". System messages are + extracted into system_instruction before contents are built, so the terminal + turn is the last non-system message. """ - if not cached_messages: - return False - return cached_messages[-1].get("role") not in ("assistant", "tool", "function") + non_system_messages = tuple(message for message in cached_messages if message.get("role") != "system") + if not non_system_messages: + return bool(cached_messages) + return non_system_messages[-1].get("role") not in ("assistant", "tool", "function") def transform_openai_messages_to_gemini_context_caching( diff --git a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py index fe4cd4ec451..f8774e33ca4 100644 --- a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py +++ b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py @@ -311,8 +311,9 @@ class ContextCachingEndpoints(VertexBase): if not cached_messages_end_on_supported_turn(cached_messages): verbose_logger.debug( - "Vertex AI context caching: cached message block ends on an assistant or " - "tool turn, which the cachedContents API rejects. Skipping context caching." + "Vertex AI context caching: cached message block ends on a model turn once " + "system messages are extracted, which the cachedContents API rejects. " + "Skipping context caching." ) return messages, optional_params, None @@ -469,8 +470,9 @@ class ContextCachingEndpoints(VertexBase): if not cached_messages_end_on_supported_turn(cached_messages): verbose_logger.debug( - "Vertex AI context caching: cached message block ends on an assistant or " - "tool turn, which the cachedContents API rejects. Skipping context caching." + "Vertex AI context caching: cached message block ends on a model turn once " + "system messages are extracted, which the cachedContents API rejects. " + "Skipping context caching." ) return messages, optional_params, None diff --git a/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py b/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py index 1aa724e551e..ad890d0c7ea 100644 --- a/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py +++ b/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py @@ -1458,18 +1458,24 @@ class TestContextCachingEndpoints: "type": "function", "function": {"name": "get_weather", "arguments": '{"location": "Boston"}'}, } - cached_tail = ( - [ + cached_tail = { + "assistant": [], + "tool": [ { "role": "tool", "tool_call_id": "call_abc123", "content": "72F and sunny", "cache_control": {"type": "ephemeral"}, } - ] - if final_cached_role == "tool" - else [] - ) + ], + "system": [ + { + "role": "system", + "content": "Tool results are authoritative.", + "cache_control": {"type": "ephemeral"}, + } + ], + }[final_cached_role] return [ { "role": "user", @@ -1491,7 +1497,7 @@ class TestContextCachingEndpoints: {"role": "user", "content": "What is the weather in Boston?"}, ] - @pytest.mark.parametrize("final_cached_role", ["assistant", "tool"]) + @pytest.mark.parametrize("final_cached_role", ["assistant", "tool", "system"]) def test_check_and_create_cache_skips_when_cached_block_ends_on_model_turn( self, final_cached_role ): @@ -1525,7 +1531,7 @@ class TestContextCachingEndpoints: self.mock_client.get.assert_not_called() self.mock_client.post.assert_not_called() - @pytest.mark.parametrize("final_cached_role", ["assistant", "tool"]) + @pytest.mark.parametrize("final_cached_role", ["assistant", "tool", "system"]) @pytest.mark.asyncio async def test_async_check_and_create_cache_skips_when_cached_block_ends_on_model_turn( self, final_cached_role @@ -1571,6 +1577,22 @@ def test_cached_messages_end_on_supported_turn(): ) assert cached_messages_end_on_supported_turn([{"role": "system", "content": "be brief"}]) is True assert cached_messages_end_on_supported_turn([{"role": "assistant", "content": "hi"}]) is False + assert ( + cached_messages_end_on_supported_turn( + [ + {"role": "user", "content": "hello"}, + {"role": "assistant", "content": "hi"}, + {"role": "system", "content": "be brief"}, + ] + ) + is False + ) + assert ( + cached_messages_end_on_supported_turn( + [{"role": "system", "content": "be brief"}, {"role": "user", "content": "hello"}] + ) + is True + ) assert ( cached_messages_end_on_supported_turn([{"role": "tool", "tool_call_id": "x", "content": "y"}]) is False From a8952499232b9e6ea92fc9392352bf4e1ea49ad6 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 29 Jul 2026 20:25:36 -0700 Subject: [PATCH 4/8] refactor(bedrock): remove the dead BedrockLLM invoke code path --- basedpyright-code-budget.json | 30 +- litellm/llms/bedrock/chat/__init__.py | 1 - litellm/llms/bedrock/chat/invoke_handler.py | 951 ------------------ litellm/llms/bedrock/common_utils.py | 4 +- litellm/main.py | 2 +- ruff-strict-budget.json | 30 +- .../test_secret_manager.py | 5 +- .../test_bedrock_completion.py | 127 +-- .../llms/bedrock/chat/test_invoke_handler.py | 33 - tests/test_litellm/test_ssl_verify_unit.py | 18 - type-discipline-budget.json | 8 +- 11 files changed, 43 insertions(+), 1166 deletions(-) diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index db3c2502e94..3f89ff179eb 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -1,12 +1,12 @@ { "reportAny": { - "limit": 33216 + "limit": 33171 }, "reportArgumentType": { - "limit": 2648 + "limit": 2645 }, "reportAssignmentType": { - "limit": 330 + "limit": 329 }, "reportAttributeAccessIssue": { "limit": 516 @@ -18,7 +18,7 @@ "limit": 59 }, "reportDeprecated": { - "limit": 326 + "limit": 325 }, "reportDuplicateImport": { "limit": 42 @@ -54,10 +54,10 @@ "limit": 0 }, "reportMissingParameterType": { - "limit": 5893 + "limit": 5869 }, "reportMissingTypeArgument": { - "limit": 15886 + "limit": 15864 }, "reportMissingTypeStubs": { "limit": 41 @@ -72,7 +72,7 @@ "limit": 0 }, "reportOptionalMemberAccess": { - "limit": 1085 + "limit": 1079 }, "reportOptionalOperand": { "limit": 0 @@ -84,13 +84,13 @@ "limit": 77 }, "reportPrivateUsage": { - "limit": 2438 + "limit": 2437 }, "reportRedeclaration": { "limit": 12 }, "reportReturnType": { - "limit": 225 + "limit": 221 }, "reportTypedDictNotRequiredAccess": { "limit": 27 @@ -99,19 +99,19 @@ "limit": 0 }, "reportUnknownArgumentType": { - "limit": 45567 + "limit": 45522 }, "reportUnknownLambdaType": { "limit": 113 }, "reportUnknownMemberType": { - "limit": 40525 + "limit": 40479 }, "reportUnknownParameterType": { - "limit": 20384 + "limit": 20341 }, "reportUnknownVariableType": { - "limit": 32099 + "limit": 32052 }, "reportUnnecessaryCast": { "limit": 177 @@ -123,7 +123,7 @@ "limit": 7 }, "reportUnnecessaryIsInstance": { - "limit": 1206 + "limit": 1205 }, "reportUntypedBaseClass": { "limit": 165 @@ -138,7 +138,7 @@ "limit": 206 }, "reportUnusedImport": { - "limit": 1005 + "limit": 1003 }, "reportUnusedVariable": { "limit": 1297 diff --git a/litellm/llms/bedrock/chat/__init__.py b/litellm/llms/bedrock/chat/__init__.py index c1323b9192a..37dcb270743 100644 --- a/litellm/llms/bedrock/chat/__init__.py +++ b/litellm/llms/bedrock/chat/__init__.py @@ -5,7 +5,6 @@ from .invoke_handler import ( AmazonAnthropicClaudeStreamDecoder, AmazonDeepSeekR1StreamDecoder, AWSEventStreamDecoder, - BedrockLLM, ) diff --git a/litellm/llms/bedrock/chat/invoke_handler.py b/litellm/llms/bedrock/chat/invoke_handler.py index 4c256be1ab8..c28627d5aec 100644 --- a/litellm/llms/bedrock/chat/invoke_handler.py +++ b/litellm/llms/bedrock/chat/invoke_handler.py @@ -1,19 +1,10 @@ -""" -TODO: DELETE FILE. Bedrock LLM is no longer used. Goto `litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py` -""" - -import copy -import time import types -from functools import partial from typing import ( AsyncIterator, - Callable, Iterator, Optional, Tuple, cast, - get_args, ) import httpx # type: ignore @@ -25,16 +16,6 @@ from litellm.caching.caching import InMemoryCache from litellm.constants import RESPONSE_FORMAT_TOOL_NAME from litellm.litellm_core_utils.core_helpers import map_finish_reason from litellm.litellm_core_utils.litellm_logging import Logging -from litellm.litellm_core_utils.logging_utils import track_llm_api_timing -from litellm.litellm_core_utils.prompt_templates.factory import ( - cohere_message_pt, - construct_tool_use_system_prompt, - contains_tag, - custom_prompt, - extract_between_tags, - parse_xml_params, - prompt_factory, -) from litellm.llms.anthropic.chat.handler import ( ModelResponseIterator as AnthropicModelResponseIterator, ) @@ -64,12 +45,9 @@ from litellm.types.utils import ( StreamingChoices, Usage, ) -from litellm.utils import CustomStreamWrapper, get_secret -from ..base_aws_llm import BaseAWSLLM from ..common_utils import ( BedrockError, - ModelResponseIterator, build_bedrock_stream_error, get_bedrock_response_stream_shape, get_bedrock_tool_name, @@ -77,9 +55,6 @@ from ..common_utils import ( bedrock_tool_name_mappings: InMemoryCache = InMemoryCache(max_size_in_memory=50, default_ttl=600) from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig -from litellm.llms.bedrock.chat.invoke_transformations.amazon_openai_transformation import ( - AmazonBedrockOpenAIConfig, -) converse_config = AmazonConverseConfig() @@ -351,932 +326,6 @@ def make_sync_call( raise BedrockError(status_code=500, message=str(e)) -class BedrockLLM(BaseAWSLLM): - """ - Example call - - ``` - curl --location --request POST 'https://bedrock-runtime.{aws_region_name}.amazonaws.com/model/{bedrock_model_name}/invoke' \ - --header 'Content-Type: application/json' \ - --header 'Accept: application/json' \ - --user "$AWS_ACCESS_KEY_ID":"$AWS_SECRET_ACCESS_KEY" \ - --aws-sigv4 "aws:amz:us-east-1:bedrock" \ - --data-raw '{ - "prompt": "Hi", - "temperature": 0, - "p": 0.9, - "max_tokens": 4096 - }' - ``` - """ - - def __init__(self) -> None: - super().__init__() - - @staticmethod - def is_claude_messages_api_model(model: str) -> bool: - """ - Check if the model uses the Claude Messages API (Claude 3+). - - Handles: - - Regional prefixes: eu.anthropic.claude-*, us.anthropic.claude-* - - Claude 3 models: claude-3-haiku, claude-3-sonnet, claude-3-opus, claude-3-5-*, claude-3-7-* - - Claude 4 models: claude-opus-4, claude-sonnet-4, claude-haiku-4 - """ - # Normalize model string to lowercase for matching - model_lower = model.lower() - - # Claude 3+ indicators (all use Messages API) - messages_api_indicators = [ - "claude-3", # Claude 3.x models - "claude-opus-4", # Claude Opus 4 - "claude-sonnet-4", # Claude Sonnet 4 - "claude-haiku-4", # Claude Haiku 4 - ] - - return any(indicator in model_lower for indicator in messages_api_indicators) - - def convert_messages_to_prompt(self, model, messages, provider, custom_prompt_dict) -> Tuple[str, Optional[list]]: - # handle anthropic prompts and amazon titan prompts - prompt = "" - chat_history: Optional[list] = None - ## CUSTOM PROMPT - if model in custom_prompt_dict: - # check if the model has a registered custom prompt - model_prompt_details = custom_prompt_dict[model] - prompt = custom_prompt( - role_dict=model_prompt_details["roles"], - initial_prompt_value=model_prompt_details.get("initial_prompt_value", ""), - final_prompt_value=model_prompt_details.get("final_prompt_value", ""), - messages=messages, - ) - return prompt, None - ## ELSE - if provider == "anthropic" or provider == "amazon": - prompt = prompt_factory(model=model, messages=messages, custom_llm_provider="bedrock") - elif provider == "mistral": - prompt = prompt_factory(model=model, messages=messages, custom_llm_provider="bedrock") - elif provider == "meta" or provider == "llama": - prompt = prompt_factory(model=model, messages=messages, custom_llm_provider="bedrock") - elif provider == "openai": - # OpenAI uses messages directly, no prompt conversion needed - # Return empty prompt as it won't be used - prompt = "" - elif provider == "cohere": - prompt, chat_history = cohere_message_pt(messages=messages) - else: - prompt = "" - for message in messages: - if "role" in message: - if message["role"] == "user": - prompt += f"{message['content']}" - else: - prompt += f"{message['content']}" - else: - prompt += f"{message['content']}" - return prompt, chat_history # type: ignore - - def process_response( - self, - model: str, - response: httpx.Response, - model_response: ModelResponse, - stream: Optional[bool], - logging_obj: Logging, - optional_params: dict, - api_key: str, - data: Union[dict, str], - messages: List, - print_verbose, - encoding, - ) -> Union[ModelResponse, CustomStreamWrapper]: - provider = self.get_bedrock_invoke_provider(model) - ## LOGGING - logging_obj.post_call( - input=messages, - api_key=api_key, - original_response=response.text, - additional_args={"complete_input_dict": data}, - ) - print_verbose(f"raw model_response: {response.text}") - - ## RESPONSE OBJECT - try: - completion_response = response.json() - except Exception: - raise BedrockError(message=response.text, status_code=422) - - outputText: Optional[str] = None - try: - if provider == "cohere": - if "text" in completion_response: - outputText = completion_response["text"] # type: ignore - elif "generations" in completion_response: - outputText = completion_response["generations"][0]["text"] - model_response.choices[0].finish_reason = map_finish_reason( - completion_response["generations"][0]["finish_reason"] - ) - elif provider == "anthropic": - if self.is_claude_messages_api_model(model): - json_schemas: dict = {} - _is_function_call = False - ## Handle Tool Calling - if "tools" in optional_params: - _is_function_call = True - for tool in optional_params["tools"]: - json_schemas[tool["function"]["name"]] = tool["function"].get("parameters", None) - outputText = completion_response.get("content")[0].get("text", None) - if outputText is not None and contains_tag("invoke", outputText): # OUTPUT PARSE FUNCTION CALL - function_name = extract_between_tags("tool_name", outputText)[0] - function_arguments_str = extract_between_tags("invoke", outputText)[0].strip() - function_arguments_str = f"{function_arguments_str}" - function_arguments = parse_xml_params( - function_arguments_str, - json_schema=json_schemas.get( - function_name, None - ), # check if we have a json schema for this function name) - ) - _message = litellm.Message( - tool_calls=[ - { - "id": f"call_{uuid.uuid4()}", - "type": "function", - "function": { - "name": function_name, - "arguments": json.dumps(function_arguments), - }, - } - ], - content=None, - ) - model_response.choices[0].message = _message # type: ignore - model_response._hidden_params["original_response"] = ( - outputText # allow user to access raw anthropic tool calling response - ) - if _is_function_call is True and stream is not None and stream is True: - print_verbose("INSIDE BEDROCK STREAMING TOOL CALLING CONDITION BLOCK") - # return an iterator - streaming_model_response = ModelResponseStream() - streaming_model_response.choices[0].finish_reason = getattr( - model_response.choices[0], "finish_reason", "stop" - ) - # streaming_model_response.choices = [litellm.utils.StreamingChoices()] - streaming_choice = litellm.utils.StreamingChoices() - streaming_choice.index = model_response.choices[0].index - _tool_calls = [] - print_verbose(f"type of model_response.choices[0]: {type(model_response.choices[0])}") - print_verbose(f"type of streaming_choice: {type(streaming_choice)}") - if isinstance(model_response.choices[0], litellm.Choices): - if getattr( - model_response.choices[0].message, "tool_calls", None - ) is not None and isinstance(model_response.choices[0].message.tool_calls, list): - for tool_call in model_response.choices[0].message.tool_calls: - _tool_call = {**tool_call.dict(), "index": 0} - _tool_calls.append(_tool_call) - delta_obj = Delta( - content=getattr(model_response.choices[0].message, "content", None), - role=model_response.choices[0].message.role, - tool_calls=_tool_calls, - ) - streaming_choice.delta = delta_obj - streaming_model_response.choices = [streaming_choice] - completion_stream = ModelResponseIterator(model_response=streaming_model_response) - print_verbose( - "Returns anthropic CustomStreamWrapper with 'cached_response' streaming object" - ) - return litellm.CustomStreamWrapper( - completion_stream=completion_stream, - model=model, - custom_llm_provider="cached_response", - logging_obj=logging_obj, - ) - - model_response.choices[0].finish_reason = map_finish_reason( - completion_response.get("stop_reason", "") - ) - _usage = litellm.Usage( - prompt_tokens=completion_response["usage"]["input_tokens"], - completion_tokens=completion_response["usage"]["output_tokens"], - total_tokens=completion_response["usage"]["input_tokens"] - + completion_response["usage"]["output_tokens"], - ) - setattr(model_response, "usage", _usage) - else: - outputText = completion_response["completion"] - - model_response.choices[0].finish_reason = completion_response["stop_reason"] - elif provider == "ai21": - outputText = completion_response.get("completions")[0].get("data").get("text") - elif provider == "meta" or provider == "llama": - outputText = completion_response["generation"] - elif provider == "openai": - # OpenAI imported models use OpenAI Chat Completions format - if "choices" in completion_response and len(completion_response["choices"]) > 0: - choice = completion_response["choices"][0] - if "message" in choice: - outputText = choice["message"].get("content") - elif "text" in choice: # fallback for completion format - outputText = choice["text"] - - # Set finish reason - if "finish_reason" in choice: - model_response.choices[0].finish_reason = map_finish_reason(choice["finish_reason"]) - - # Set usage if available - if "usage" in completion_response: - usage = completion_response["usage"] - _usage = litellm.Usage( - prompt_tokens=usage.get("prompt_tokens", 0), - completion_tokens=usage.get("completion_tokens", 0), - total_tokens=usage.get("total_tokens", 0), - ) - setattr(model_response, "usage", _usage) - elif provider == "mistral": - outputText = completion_response["outputs"][0]["text"] - model_response.choices[0].finish_reason = completion_response["outputs"][0]["stop_reason"] - else: # amazon titan - outputText = completion_response.get("results")[0].get("outputText") - except Exception as e: - raise BedrockError( - message="Error processing={}, Received error={}".format(response.text, str(e)), - status_code=422, - ) - - try: - if ( - outputText is not None - and len(outputText) > 0 - and hasattr(model_response.choices[0], "message") - and getattr(model_response.choices[0].message, "tool_calls", None) # type: ignore - is None - ): - model_response.choices[0].message.content = outputText # type: ignore - elif ( - hasattr(model_response.choices[0], "message") - and getattr(model_response.choices[0].message, "tool_calls", None) # type: ignore - is not None - ): - pass - else: - raise Exception() - except Exception as e: - raise BedrockError( - message="Error parsing received text={}.\nError-{}".format(outputText, str(e)), - status_code=response.status_code, - ) - - if stream and provider == "ai21": - streaming_model_response = ModelResponseStream() - streaming_model_response.choices[0].finish_reason = model_response.choices[ # type: ignore - 0 - ].finish_reason - # streaming_model_response.choices = [litellm.utils.StreamingChoices()] - streaming_choice = litellm.utils.StreamingChoices() - streaming_choice.index = model_response.choices[0].index - delta_obj = litellm.utils.Delta( - content=getattr(model_response.choices[0].message, "content", None), # type: ignore - role=model_response.choices[0].message.role, # type: ignore - ) - streaming_choice.delta = delta_obj - streaming_model_response.choices = [streaming_choice] - mri = ModelResponseIterator(model_response=streaming_model_response) - return CustomStreamWrapper( - completion_stream=mri, - model=model, - custom_llm_provider="cached_response", - logging_obj=logging_obj, - ) - - ## CALCULATING USAGE - bedrock returns usage in the headers - # Skip if usage was already set (e.g., from JSON response for OpenAI provider) - if not hasattr(model_response, "usage") or getattr(model_response, "usage", None) is None: - bedrock_input_tokens = response.headers.get("x-amzn-bedrock-input-token-count", None) - bedrock_output_tokens = response.headers.get("x-amzn-bedrock-output-token-count", None) - - prompt_tokens = int(bedrock_input_tokens or litellm.token_counter(messages=messages)) - - completion_tokens = int( - bedrock_output_tokens - or litellm.token_counter( - text=model_response.choices[0].message.content, # type: ignore - count_response_tokens=True, - ) - ) - - model_response.created = int(time.time()) - model_response.model = model - usage = Usage( - prompt_tokens=prompt_tokens, - completion_tokens=completion_tokens, - total_tokens=prompt_tokens + completion_tokens, - ) - setattr(model_response, "usage", usage) - else: - # Ensure created and model are set even if usage was already set - model_response.created = int(time.time()) - model_response.model = model - - return model_response - - def completion( - self, - model: str, - messages: list, - api_base: Optional[str], - custom_prompt_dict: dict, - model_response: ModelResponse, - print_verbose: Callable, - encoding, - logging_obj: Logging, - optional_params: dict, - acompletion: bool, - timeout: Optional[Union[float, httpx.Timeout]], - litellm_params=None, - logger_fn=None, - extra_headers: Optional[dict] = None, - client: Optional[Union[AsyncHTTPHandler, HTTPHandler]] = None, - ) -> Union[ModelResponse, CustomStreamWrapper]: - try: - from botocore.credentials import Credentials - except ImportError: - raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") - - ## SETUP ## - stream = optional_params.pop("stream", None) - stream_chunk_size = optional_params.pop("stream_chunk_size", None) - - provider = self.get_bedrock_invoke_provider(model) - modelId = self.get_bedrock_model_id( - model=model, - provider=provider, - optional_params=optional_params, - ) - - ## CREDENTIALS ## - # pop aws_secret_access_key, aws_access_key_id, aws_session_token, aws_region_name from kwargs, since completion calls fail with them - aws_secret_access_key = optional_params.pop("aws_secret_access_key", None) - aws_access_key_id = optional_params.pop("aws_access_key_id", None) - aws_session_token = optional_params.pop("aws_session_token", None) - aws_region_name = optional_params.pop("aws_region_name", None) - aws_role_name = optional_params.pop("aws_role_name", None) - aws_session_name = optional_params.pop("aws_session_name", None) - aws_profile_name = optional_params.pop("aws_profile_name", None) - aws_bedrock_runtime_endpoint = optional_params.pop( - "aws_bedrock_runtime_endpoint", None - ) # https://bedrock-runtime.{region_name}.amazonaws.com - aws_web_identity_token = optional_params.pop("aws_web_identity_token", None) - aws_sts_endpoint = optional_params.pop("aws_sts_endpoint", None) - ssl_verify = optional_params.pop("ssl_verify", None) - - ### SET REGION NAME ### - if aws_region_name is None: - # check env # - litellm_aws_region_name = get_secret("AWS_REGION_NAME", None) - - if litellm_aws_region_name is not None and isinstance(litellm_aws_region_name, str): - aws_region_name = litellm_aws_region_name - - standard_aws_region_name = get_secret("AWS_REGION", None) - if standard_aws_region_name is not None and isinstance(standard_aws_region_name, str): - aws_region_name = standard_aws_region_name - - if aws_region_name is None: - aws_region_name = "us-west-2" - - credentials: Credentials = self.get_credentials( - aws_access_key_id=aws_access_key_id, - aws_secret_access_key=aws_secret_access_key, - aws_session_token=aws_session_token, - aws_region_name=aws_region_name, - aws_session_name=aws_session_name, - aws_profile_name=aws_profile_name, - aws_role_name=aws_role_name, - aws_web_identity_token=aws_web_identity_token, - aws_sts_endpoint=aws_sts_endpoint, - ssl_verify=ssl_verify, - ) - - ### SET RUNTIME ENDPOINT ### - endpoint_url, proxy_endpoint_url = self.get_runtime_endpoint( - api_base=api_base, - aws_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint, - aws_region_name=aws_region_name, - ) - - if (stream is not None and stream is True) and provider != "ai21": - endpoint_url = f"{endpoint_url}/model/{modelId}/invoke-with-response-stream" - proxy_endpoint_url = f"{proxy_endpoint_url}/model/{modelId}/invoke-with-response-stream" - else: - endpoint_url = f"{endpoint_url}/model/{modelId}/invoke" - proxy_endpoint_url = f"{proxy_endpoint_url}/model/{modelId}/invoke" - - if acompletion and provider == "anthropic" and self.is_claude_messages_api_model(model): - if isinstance(client, HTTPHandler): - client = None - return self._async_anthropic_messages_completion( - model=model, - messages=messages, - endpoint_url=endpoint_url, - proxy_endpoint_url=proxy_endpoint_url, - credentials=credentials, - aws_region_name=aws_region_name, - model_response=model_response, - print_verbose=print_verbose, - encoding=encoding, - logging_obj=logging_obj, - optional_params=optional_params, - stream=stream, - litellm_params=litellm_params, - logger_fn=logger_fn, - extra_headers=extra_headers, - timeout=timeout, - client=client, - stream_chunk_size=stream_chunk_size, - ) # type: ignore[return-value] - - prompt, chat_history = self.convert_messages_to_prompt(model, messages, provider, custom_prompt_dict) - inference_params = copy.deepcopy(optional_params) - json_schemas: dict = {} - if provider == "cohere": - if model.startswith("cohere.command-r"): - ## LOAD CONFIG - config = litellm.AmazonCohereChatConfig().get_config() - for k, v in config.items(): - if ( - k not in inference_params - ): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in - inference_params[k] = v - _data = {"message": prompt, **inference_params} - if chat_history is not None: - _data["chat_history"] = chat_history - data = json.dumps(_data) - else: - ## LOAD CONFIG - config = litellm.AmazonCohereConfig.get_config() - for k, v in config.items(): - if ( - k not in inference_params - ): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in - inference_params[k] = v - if stream is True: - inference_params["stream"] = True # cohere requires stream = True in inference params - data = json.dumps({"prompt": prompt, **inference_params}) - elif provider == "anthropic": - if self.is_claude_messages_api_model(model): - # Separate system prompt from rest of message - system_prompt_idx: list[int] = [] - system_messages: list[str] = [] - for idx, message in enumerate(messages): - if message["role"] == "system": - system_messages.append(message["content"]) - system_prompt_idx.append(idx) - if len(system_prompt_idx) > 0: - inference_params["system"] = "\n".join(system_messages) - messages = [i for j, i in enumerate(messages) if j not in system_prompt_idx] - # Format rest of message according to anthropic guidelines - messages = prompt_factory(model=model, messages=messages, custom_llm_provider="anthropic_xml") # type: ignore - ## LOAD CONFIG - config = litellm.AmazonAnthropicClaudeConfig.get_config() - for k, v in config.items(): - if ( - k not in inference_params - ): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in - inference_params[k] = v - ## Handle Tool Calling - if "tools" in inference_params: - _is_function_call = True - for tool in inference_params["tools"]: - json_schemas[tool["function"]["name"]] = tool["function"].get("parameters", None) - tool_calling_system_prompt = construct_tool_use_system_prompt(tools=inference_params["tools"]) - inference_params["system"] = ( - inference_params.get("system", "\n") + tool_calling_system_prompt - ) # add the anthropic tool calling prompt to the system prompt - inference_params.pop("tools") - data = json.dumps({"messages": messages, **inference_params}) - else: - ## LOAD CONFIG - config = litellm.AmazonAnthropicConfig.get_config() - for k, v in config.items(): - if ( - k not in inference_params - ): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in - inference_params[k] = v - data = json.dumps({"prompt": prompt, **inference_params}) - elif provider == "ai21": - ## LOAD CONFIG - config = litellm.AmazonAI21Config.get_config() - for k, v in config.items(): - if ( - k not in inference_params - ): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in - inference_params[k] = v - - data = json.dumps({"prompt": prompt, **inference_params}) - elif provider == "mistral": - ## LOAD CONFIG - config = litellm.AmazonMistralConfig.get_config() - for k, v in config.items(): - if ( - k not in inference_params - ): # completion(top_k=3) > amazon_config(top_k=3) <- allows for dynamic variables to be passed in - inference_params[k] = v - - data = json.dumps({"prompt": prompt, **inference_params}) - elif provider == "amazon": # amazon titan - ## LOAD CONFIG - config = litellm.AmazonTitanConfig.get_config() - for k, v in config.items(): - if ( - k not in inference_params - ): # completion(top_k=3) > amazon_config(top_k=3) <- allows for dynamic variables to be passed in - inference_params[k] = v - - data = json.dumps( - { - "inputText": prompt, - "textGenerationConfig": inference_params, - } - ) - elif provider == "meta" or provider == "llama": - ## LOAD CONFIG - config = litellm.AmazonLlamaConfig.get_config() - for k, v in config.items(): - if ( - k not in inference_params - ): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in - inference_params[k] = v - data = json.dumps({"prompt": prompt, **inference_params}) - elif provider == "openai": - ## OpenAI imported models use OpenAI Chat Completions format (messages-based) - # Use AmazonBedrockOpenAIConfig for proper OpenAI transformation - openai_config = AmazonBedrockOpenAIConfig() - supported_params = openai_config.get_supported_openai_params(model=model) - - # Filter to only supported OpenAI params - filtered_params = {k: v for k, v in inference_params.items() if k in supported_params} - - # OpenAI uses messages format, not prompt - data = json.dumps({"messages": messages, **filtered_params}) - else: - ## LOGGING - logging_obj.pre_call( - input=messages, - api_key="", - additional_args={ - "complete_input_dict": inference_params, - }, - ) - raise BedrockError( - status_code=404, - message="Bedrock Invoke HTTPX: Unknown provider={}, model={}. Try calling via converse route - `bedrock/converse/`.".format( - provider, model - ), - ) - - ## COMPLETION CALL - - headers = {"Content-Type": "application/json"} - if extra_headers is not None: - headers = {"Content-Type": "application/json", **extra_headers} - prepped = self.get_request_headers( - credentials=credentials, - aws_region_name=aws_region_name, - extra_headers=extra_headers, - endpoint_url=endpoint_url, - data=data, - headers=headers, - ) - - ## LOGGING - logging_obj.pre_call( - input=messages, - api_key="", - additional_args={ - "complete_input_dict": data, - "api_base": proxy_endpoint_url, - "headers": prepped.headers, - }, - ) - - ### ROUTING (ASYNC, STREAMING, SYNC) - if acompletion: - if isinstance(client, HTTPHandler): - client = None - if stream is True and provider != "ai21": - return self.async_streaming( - model=model, - messages=messages, - data=data, - api_base=proxy_endpoint_url, - model_response=model_response, - print_verbose=print_verbose, - encoding=encoding, - logging_obj=logging_obj, - optional_params=optional_params, - stream=True, - litellm_params=litellm_params, - logger_fn=logger_fn, - headers=prepped.headers, - timeout=timeout, - client=client, - stream_chunk_size=stream_chunk_size, - ) # type: ignore - ### ASYNC COMPLETION - return self.async_completion( - model=model, - messages=messages, - data=data, - api_base=proxy_endpoint_url, - model_response=model_response, - print_verbose=print_verbose, - encoding=encoding, - logging_obj=logging_obj, - optional_params=optional_params, - stream=stream, # type: ignore - litellm_params=litellm_params, - logger_fn=logger_fn, - headers=prepped.headers, - timeout=timeout, - client=client, - ) # type: ignore - - if client is None or isinstance(client, AsyncHTTPHandler): - _params = {} - if timeout is not None: - if isinstance(timeout, float) or isinstance(timeout, int): - timeout = httpx.Timeout(timeout) - _params["timeout"] = timeout - self.client = _get_httpx_client(_params) # type: ignore - else: - self.client = client - if (stream is not None and stream is True) and provider != "ai21": - response = self.client.post( - url=proxy_endpoint_url, - headers=prepped.headers, # type: ignore - data=data, - stream=stream, - logging_obj=logging_obj, - ) - - if response.status_code != 200: - raise BedrockError(status_code=response.status_code, message=str(response.read())) - - decoder = AWSEventStreamDecoder(model=model) - - completion_stream = decoder.iter_bytes(response.iter_bytes(chunk_size=stream_chunk_size)) - streaming_response = CustomStreamWrapper( - completion_stream=completion_stream, - model=model, - custom_llm_provider="bedrock", - logging_obj=logging_obj, - ) - - ## LOGGING - logging_obj.post_call( - input=messages, - api_key="", - original_response=streaming_response, - additional_args={"complete_input_dict": data}, - ) - return streaming_response - - try: - response = self.client.post( - url=proxy_endpoint_url, - headers=dict(prepped.headers), - data=data, - logging_obj=logging_obj, - ) - response.raise_for_status() - except httpx.HTTPStatusError as err: - error_code = err.response.status_code - raise BedrockError(status_code=error_code, message=err.response.text) - except httpx.TimeoutException: - raise BedrockError(status_code=408, message="Timeout error occurred.") - - return self.process_response( - model=model, - response=response, - model_response=model_response, - stream=stream, - logging_obj=logging_obj, - optional_params=optional_params, - api_key="", - data=data, - messages=messages, - print_verbose=print_verbose, - encoding=encoding, - ) - - async def _async_anthropic_messages_completion( - self, - model: str, - messages: list, - endpoint_url: str, - proxy_endpoint_url: str, - credentials, - aws_region_name: str, - model_response: ModelResponse, - print_verbose: Callable, - encoding, - logging_obj: Logging, - optional_params: dict, - stream, - litellm_params=None, - logger_fn=None, - extra_headers: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[AsyncHTTPHandler] = None, - stream_chunk_size: Optional[int] = None, - ) -> Union[ModelResponse, CustomStreamWrapper]: - transformed_request = await litellm.AmazonAnthropicClaudeConfig().async_transform_request( - model=model, - messages=messages, - optional_params=optional_params, - litellm_params=litellm_params or {}, - headers=extra_headers or {}, - ) - data = json.dumps(transformed_request) - - headers = {"Content-Type": "application/json"} - if extra_headers is not None: - headers = {"Content-Type": "application/json", **extra_headers} - prepped = self.get_request_headers( - credentials=credentials, - aws_region_name=aws_region_name, - extra_headers=extra_headers, - endpoint_url=endpoint_url, - data=data, - headers=headers, - ) - - logging_obj.pre_call( - input=messages, - api_key="", - additional_args={ - "complete_input_dict": data, - "api_base": proxy_endpoint_url, - "headers": prepped.headers, - }, - ) - - if stream is True: - return await self.async_streaming( - model=model, - messages=messages, - data=data, - api_base=proxy_endpoint_url, - model_response=model_response, - print_verbose=print_verbose, - encoding=encoding, - logging_obj=logging_obj, - optional_params=optional_params, - stream=True, - litellm_params=litellm_params, - logger_fn=logger_fn, - headers=prepped.headers, - timeout=timeout, - client=client, - stream_chunk_size=stream_chunk_size, - ) - return await self.async_completion( - model=model, - messages=messages, - data=data, - api_base=proxy_endpoint_url, - model_response=model_response, - print_verbose=print_verbose, - encoding=encoding, - logging_obj=logging_obj, - optional_params=optional_params, - stream=stream, # type: ignore - litellm_params=litellm_params, - logger_fn=logger_fn, - headers=prepped.headers, - timeout=timeout, - client=client, - ) - - async def async_completion( - self, - model: str, - messages: list, - api_base: str, - model_response: ModelResponse, - print_verbose: Callable, - data: str, - timeout: Optional[Union[float, httpx.Timeout]], - encoding, - logging_obj: Logging, - stream, - optional_params: dict, - litellm_params=None, - logger_fn=None, - headers={}, - client: Optional[AsyncHTTPHandler] = None, - ) -> Union[ModelResponse, CustomStreamWrapper]: - if client is None: - _params = {} - if timeout is not None: - if isinstance(timeout, float) or isinstance(timeout, int): - timeout = httpx.Timeout(timeout) - _params["timeout"] = timeout - client = get_async_httpx_client(params=_params, llm_provider=litellm.LlmProviders.BEDROCK) # type: ignore - else: - client = client # type: ignore - - try: - response = await client.post( - api_base, - headers=headers, - data=data, - timeout=timeout, - logging_obj=logging_obj, - ) - response.raise_for_status() - except httpx.HTTPStatusError as err: - error_code = err.response.status_code - raise BedrockError(status_code=error_code, message=err.response.text) - except httpx.TimeoutException: - raise BedrockError(status_code=408, message="Timeout error occurred.") - - return self.process_response( - model=model, - response=response, - model_response=model_response, - stream=stream if isinstance(stream, bool) else False, - logging_obj=logging_obj, - api_key="", - data=data, - messages=messages, - print_verbose=print_verbose, - optional_params=optional_params, - encoding=encoding, - ) - - @track_llm_api_timing() # for streaming, we need to instrument the function calling the wrapper - async def async_streaming( - self, - model: str, - messages: list, - api_base: str, - model_response: ModelResponse, - print_verbose: Callable, - data: str, - timeout: Optional[Union[float, httpx.Timeout]], - encoding, - logging_obj: Logging, - stream, - optional_params: dict, - litellm_params=None, - logger_fn=None, - headers={}, - client: Optional[AsyncHTTPHandler] = None, - stream_chunk_size: Optional[int] = None, - ) -> CustomStreamWrapper: - # The call is not made here; instead, we prepare the necessary objects for the stream. - - streaming_response = CustomStreamWrapper( - completion_stream=None, - make_call=partial( - make_call, - client=client, - api_base=api_base, - headers=headers, - data=data, # type: ignore - model=model, - messages=messages, - logging_obj=logging_obj, - fake_stream=True if "ai21" in api_base else False, - stream_chunk_size=stream_chunk_size, - ), - model=model, - custom_llm_provider="bedrock", - logging_obj=logging_obj, - ) - return streaming_response - - @staticmethod - def _get_provider_from_model_path( - model_path: str, - ) -> Optional[litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL]: - """ - Helper function to get the provider from a model path with format: provider/model-name - - Args: - model_path (str): The model path (e.g., 'llama/arn:aws:bedrock:us-east-1:086734376398:imported-model/r4c4kewx2s0n' or 'anthropic/model-name') - - Returns: - Optional[str]: The provider name, or None if no valid provider found - """ - parts = model_path.split("/") - if len(parts) >= 1: - provider = parts[0] - if provider in get_args(litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL): - return cast(litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL, provider) - return None - - class AWSEventStreamDecoder: def __init__(self, model: str, json_mode: Optional[bool] = False) -> None: from botocore.parsers import EventStreamJSONParser diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index 5114677ffc0..93998f0610e 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -1109,8 +1109,10 @@ def get_bedrock_chat_config(model: str): Returns: The appropriate Bedrock config class instance """ + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + bedrock_route = BedrockModelInfo.get_bedrock_route(model) - bedrock_invoke_provider = litellm.BedrockLLM.get_bedrock_invoke_provider(model=model) + bedrock_invoke_provider = BaseAWSLLM.get_bedrock_invoke_provider(model=model) base_model = BedrockModelInfo.get_base_model(model) # Handle explicit routes first diff --git a/litellm/main.py b/litellm/main.py index dc3ec469a1b..acdec7385da 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -207,7 +207,7 @@ from .llms.azure.chat.o_series_handler import AzureOpenAIO1ChatCompletion from .llms.azure.completion.handler import AzureTextCompletion from .llms.azure_ai.anthropic.handler import AzureAnthropicChatCompletion from .llms.azure_ai.embed import AzureAIEmbedding -from .llms.bedrock.chat import BedrockConverseLLM, BedrockLLM +from .llms.bedrock.chat import BedrockConverseLLM from .llms.bedrock.embed.embedding import BedrockEmbedding from .llms.bedrock.image_edit.handler import BedrockImageEdit from .llms.bedrock.image_generation.image_handler import BedrockImageGeneration diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index f3b4fce97d3..3b5ec5b0dee 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -1,6 +1,6 @@ { "ANN001": { - "limit": 3142 + "limit": 3118 }, "ANN002": { "limit": 69 @@ -33,7 +33,7 @@ "limit": 4 }, "B006": { - "limit": 190 + "limit": 188 }, "B008": { "limit": 505 @@ -42,7 +42,7 @@ "limit": 84 }, "B010": { - "limit": 197 + "limit": 194 }, "B018": { "limit": 5 @@ -60,7 +60,7 @@ "limit": 4 }, "BLE001": { - "limit": 2902 + "limit": 2899 }, "C401": { "limit": 11 @@ -81,7 +81,7 @@ "limit": 4 }, "C901": { - "limit": 316 + "limit": 314 }, "D419": { "limit": 9 @@ -180,7 +180,7 @@ "limit": 34 }, "PLR1714": { - "limit": 265 + "limit": 261 }, "PLR1730": { "limit": 10 @@ -189,7 +189,7 @@ "limit": 4 }, "PLW0127": { - "limit": 44 + "limit": 43 }, "PLW0133": { "limit": 4 @@ -222,7 +222,7 @@ "limit": 38 }, "RET504": { - "limit": 717 + "limit": 716 }, "RUF010": { "limit": 874 @@ -261,7 +261,7 @@ "limit": 24 }, "SIM101": { - "limit": 63 + "limit": 61 }, "SIM102": { "limit": 324 @@ -273,7 +273,7 @@ "limit": 6 }, "SIM114": { - "limit": 113 + "limit": 111 }, "SIM115": { "limit": 5 @@ -288,7 +288,7 @@ "limit": 4 }, "SIM210": { - "limit": 12 + "limit": 11 }, "SIM211": { "limit": 4 @@ -309,7 +309,7 @@ "limit": 2652 }, "TRY002": { - "limit": 548 + "limit": 547 }, "TRY004": { "limit": 98 @@ -324,7 +324,7 @@ "limit": 883 }, "UP006": { - "limit": 12147 + "limit": 12146 }, "UP007": { "limit": 2526 @@ -348,7 +348,7 @@ "limit": 5 }, "UP032": { - "limit": 629 + "limit": 626 }, "UP034": { "limit": 4 @@ -363,6 +363,6 @@ "limit": 105 }, "UP045": { - "limit": 17824 + "limit": 17806 } } diff --git a/tests/litellm_utils_tests/test_secret_manager.py b/tests/litellm_utils_tests/test_secret_manager.py index 0a2419d0bea..0f95fd75c53 100644 --- a/tests/litellm_utils_tests/test_secret_manager.py +++ b/tests/litellm_utils_tests/test_secret_manager.py @@ -19,7 +19,8 @@ sys.path.insert( import pytest import litellm from litellm.llms.azure.azure import get_azure_ad_token_from_oidc -from litellm.llms.bedrock.chat import BedrockConverseLLM, BedrockLLM +from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM +from litellm.llms.bedrock.chat import BedrockConverseLLM from litellm.secret_managers.aws_secret_manager_v2 import AWSSecretsManagerV2 from litellm.secret_managers.main import ( get_secret, @@ -160,7 +161,7 @@ def test_oidc_circle_v1_with_amazon(): aws_role_name = "arn:aws:iam::335785316107:role/litellm-github-unit-tests-circleci-v1-assume-only" aws_web_identity_token = "oidc/circleci/" - bllm = BedrockLLM() + bllm = BaseAWSLLM() creds = bllm.get_credentials( aws_region_name="ca-west-1", aws_web_identity_token=aws_web_identity_token, diff --git a/tests/llm_translation/test_bedrock_completion.py b/tests/llm_translation/test_bedrock_completion.py index f4c307e9c8a..893b9e9c666 100644 --- a/tests/llm_translation/test_bedrock_completion.py +++ b/tests/llm_translation/test_bedrock_completion.py @@ -33,7 +33,7 @@ from litellm import ( completion_cost, embedding, ) -from litellm.llms.bedrock.chat import BedrockLLM +from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.litellm_core_utils.prompt_templates.factory import _bedrock_tools_pt from base_llm_unit_tests import BaseLLMChatTest, BaseAnthropicChatTest @@ -225,7 +225,7 @@ def bedrock_session_token_creds(): aws_region_name = os.environ["AWS_REGION_NAME"] aws_session_token = os.environ.get("AWS_SESSION_TOKEN") - bllm = BedrockLLM() + bllm = BaseAWSLLM() if aws_session_token is not None: # For local testing creds = bllm.get_credentials( @@ -3573,89 +3573,6 @@ def test_bedrock_openai_model_id_extraction(): print(f"✓ Model ID extracted and encoded: {model_id}") -def test_bedrock_openai_convert_messages_to_prompt(): - """ - Test that convert_messages_to_prompt returns empty string for OpenAI models. - """ - from litellm.llms.bedrock.chat.invoke_handler import BedrockLLM - - bedrock_llm = BedrockLLM() - messages = [ - {"role": "system", "content": "You are helpful"}, - {"role": "user", "content": "Hello"}, - ] - - prompt, chat_history = bedrock_llm.convert_messages_to_prompt( - model="test-model", messages=messages, provider="openai", custom_prompt_dict={} - ) - - # OpenAI models use messages directly, no prompt conversion - assert prompt == "" - assert chat_history is None - print("✓ convert_messages_to_prompt returns empty for OpenAI") - - -def test_bedrock_openai_response_parsing(): - """ - Test that OpenAI responses are correctly parsed. - """ - from litellm.llms.bedrock.chat.invoke_handler import BedrockLLM - from litellm import ModelResponse - from unittest.mock import Mock - import json - - bedrock_llm = BedrockLLM() - - # Mock OpenAI-style response - openai_response = { - "choices": [ - { - "message": { - "content": "The capital of France is Paris.", - "role": "assistant", - }, - "finish_reason": "stop", - "index": 0, - } - ], - "usage": {"prompt_tokens": 10, "completion_tokens": 8, "total_tokens": 18}, - } - - mock_response = Mock() - mock_response.json.return_value = openai_response - mock_response.text = json.dumps(openai_response) - mock_response.status_code = 200 - mock_response.headers = {} - - model_response = ModelResponse() - mock_logging = Mock() - - result = bedrock_llm.process_response( - model="openai/arn:aws:bedrock:us-east-1:123:imported-model/test", - response=mock_response, - model_response=model_response, - stream=False, - logging_obj=mock_logging, - optional_params={}, - api_key="", - data={}, - messages=[{"role": "user", "content": "What is the capital of France?"}], - print_verbose=lambda x: None, - encoding=None, - ) - - # Verify response content - assert result.choices[0].message.content == "The capital of France is Paris." - assert result.choices[0].finish_reason == "stop" - - # Verify usage - assert result.usage.prompt_tokens == 10 - assert result.usage.completion_tokens == 8 - assert result.usage.total_tokens == 18 - - print("✓ OpenAI response parsing works correctly") - - def test_bedrock_openai_request_transformation(): """ Test that the request is correctly transformed for OpenAI models. @@ -3845,46 +3762,6 @@ def test_bedrock_openai_multiple_message_types(): print("✓ Multiple message types handled correctly") -def test_bedrock_openai_error_handling(): - """ - Test that errors from OpenAI models are properly handled. - """ - from litellm.llms.bedrock.chat.invoke_handler import BedrockLLM - from litellm import ModelResponse - from litellm.llms.bedrock.common_utils import BedrockError - from unittest.mock import Mock - import json - - bedrock_llm = BedrockLLM() - - # Mock error response - mock_response = Mock() - mock_response.json.side_effect = Exception("Invalid JSON") - mock_response.text = "Invalid response" - mock_response.status_code = 422 - - model_response = ModelResponse() - mock_logging = Mock() - - with pytest.raises(BedrockError) as exc_info: - bedrock_llm.process_response( - model="openai/arn:aws:bedrock:us-east-1:123:imported-model/test", - response=mock_response, - model_response=model_response, - stream=False, - logging_obj=mock_logging, - optional_params={}, - api_key="", - data={}, - messages=[], - print_verbose=lambda x: None, - encoding=None, - ) - - assert exc_info.value.status_code == 422 - print("✓ Error handling works correctly") - - # ============================================================================ # Nova Grounding (web_search_options) Unit Tests (Mocked) # ============================================================================ diff --git a/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py b/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py index 61987d25d9c..ee50b9db015 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py +++ b/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py @@ -8,14 +8,11 @@ sys.path.insert( 0, os.path.abspath("../../../../..") ) # Adds the parent directory to the system path -import litellm from litellm.llms.bedrock.chat.invoke_handler import ( AWSEventStreamDecoder, - BedrockLLM, make_call, make_sync_call, ) -from litellm.llms.custom_httpx.http_handler import HTTPHandler def test_transform_thinking_blocks_with_redacted_content(): @@ -296,33 +293,3 @@ def test_make_sync_call_honors_explicit_stream_chunk_size(): response.iter_bytes.assert_called_once_with(chunk_size=2048) - -def test_legacy_bedrock_llm_streaming_does_not_rechunk_by_default(): - mock_response = MagicMock() - mock_response.status_code = 200 - mock_response.iter_bytes = MagicMock(return_value=iter([])) - client = HTTPHandler() - client.post = MagicMock(return_value=mock_response) - - BedrockLLM().completion( - model="cohere.command-text-v14", - messages=[{"role": "user", "content": "hi"}], - api_base=None, - custom_prompt_dict={}, - model_response=litellm.ModelResponse(), - print_verbose=lambda *args, **kwargs: None, - encoding=litellm.encoding, - logging_obj=MagicMock(), - optional_params={ - "stream": True, - "aws_access_key_id": "fake", - "aws_secret_access_key": "fake", - "aws_region_name": "us-east-1", - }, - acompletion=False, - timeout=None, - litellm_params={}, - client=client, - ) - - mock_response.iter_bytes.assert_called_once_with(chunk_size=None) diff --git a/tests/test_litellm/test_ssl_verify_unit.py b/tests/test_litellm/test_ssl_verify_unit.py index 7cc15703a3b..c39362c01a2 100644 --- a/tests/test_litellm/test_ssl_verify_unit.py +++ b/tests/test_litellm/test_ssl_verify_unit.py @@ -17,7 +17,6 @@ sys.path.insert(0, str(Path(__file__).parent)) import litellm.proxy.guardrails.guardrail_hooks.aim.aim as _aim_module import litellm.proxy.guardrails.guardrail_hooks.cato_networks.cato_networks as _cato_networks_module from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM -from litellm.llms.bedrock.chat.invoke_handler import BedrockLLM from litellm.proxy.guardrails.guardrail_hooks.aim.aim import AimGuardrail from litellm.proxy.guardrails.guardrail_hooks.cato_networks.cato_networks import CatoNetworksGuardrail @@ -87,23 +86,6 @@ class TestBaseAWSLLMSSLVerify: assert True # If we got here without error, parameter was accepted -class TestBedrockLLMSSLVerify: - """Test SSL verification parameter handling in BedrockLLM.""" - - def test_bedrock_llm_accepts_ssl_verify_in_optional_params(self): - """Test that BedrockLLM can receive ssl_verify in optional_params.""" - # This is a simple test to verify the parameter is accepted - # The actual propagation is tested in integration tests - bedrock_llm = BedrockLLM() - - # Verify the class exists and can be instantiated - assert bedrock_llm is not None - - # Verify _get_ssl_verify method exists and works - result = bedrock_llm._get_ssl_verify(ssl_verify="/path/to/cert.pem") - assert result == "/path/to/cert.pem" - - class TestAimGuardrailSSLVerify: """Test SSL verification parameter handling in AimGuardrail.""" diff --git a/type-discipline-budget.json b/type-discipline-budget.json index d56d5a6e305..25cc1621d54 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,9 +1,9 @@ { "LIT001": { - "limit": 23287 + "limit": 23267 }, "LIT002": { - "limit": 27473 + "limit": 27434 }, "LIT003": { "limit": 292 @@ -15,7 +15,7 @@ "limit": 0 }, "LIT006": { - "limit": 1109 + "limit": 1108 }, "LIT007": { "limit": 0 @@ -24,6 +24,6 @@ "limit": 1004 }, "LIT009": { - "limit": 2495 + "limit": 2474 } } From a71d8d887dbe7fb137da4a558fc6adc6c36bc954 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 29 Jul 2026 20:59:08 -0700 Subject: [PATCH 5/8] test(bedrock): port the openai-route invoke tests onto the live config --- .../test_bedrock_completion.py | 61 +++++++++++++++++++ 1 file changed, 61 insertions(+) diff --git a/tests/llm_translation/test_bedrock_completion.py b/tests/llm_translation/test_bedrock_completion.py index 893b9e9c666..8ab2feaf896 100644 --- a/tests/llm_translation/test_bedrock_completion.py +++ b/tests/llm_translation/test_bedrock_completion.py @@ -3573,6 +3573,50 @@ def test_bedrock_openai_model_id_extraction(): print(f"✓ Model ID extracted and encoded: {model_id}") +def test_bedrock_openai_response_parsing(): + from litellm.llms.bedrock.chat.invoke_transformations.amazon_openai_transformation import ( + AmazonBedrockOpenAIConfig, + ) + + openai_response = { + "choices": [ + { + "message": { + "content": "The capital of France is Paris.", + "role": "assistant", + }, + "finish_reason": "stop", + "index": 0, + } + ], + "usage": {"prompt_tokens": 10, "completion_tokens": 8, "total_tokens": 18}, + } + + mock_response = Mock() + mock_response.json.return_value = openai_response + mock_response.text = json.dumps(openai_response) + mock_response.status_code = 200 + mock_response.headers = {} + + result = AmazonBedrockOpenAIConfig().transform_response( + model="openai/arn:aws:bedrock:us-east-1:123:imported-model/test", + raw_response=mock_response, + model_response=ModelResponse(), + logging_obj=Mock(), + request_data={}, + messages=[{"role": "user", "content": "What is the capital of France?"}], + optional_params={}, + litellm_params={}, + encoding=None, + ) + + assert result.choices[0].message.content == "The capital of France is Paris." + assert result.choices[0].finish_reason == "stop" + assert result.usage.prompt_tokens == 10 + assert result.usage.completion_tokens == 8 + assert result.usage.total_tokens == 18 + + def test_bedrock_openai_request_transformation(): """ Test that the request is correctly transformed for OpenAI models. @@ -3762,6 +3806,23 @@ def test_bedrock_openai_multiple_message_types(): print("✓ Multiple message types handled correctly") +def test_bedrock_openai_error_handling(): + from litellm.llms.bedrock.chat.invoke_transformations.amazon_openai_transformation import ( + AmazonBedrockOpenAIConfig, + ) + from litellm.llms.bedrock.common_utils import BedrockError + + error = AmazonBedrockOpenAIConfig().get_error_class( + error_message="ValidationException: bad request", + status_code=422, + headers={}, + ) + + assert isinstance(error, BedrockError) + assert error.status_code == 422 + assert "ValidationException: bad request" in str(error) + + # ============================================================================ # Nova Grounding (web_search_options) Unit Tests (Mocked) # ============================================================================ From 8a300929614cce9b31f6caaa748049590b63ffff Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 29 Jul 2026 23:53:15 -0700 Subject: [PATCH 6/8] test(batches): exercise real GCS validation for vertex batch output reads --- .../test_litellm/batches/test_batch_utils.py | 93 +++++++++++++++++++ 1 file changed, 93 insertions(+) diff --git a/tests/test_litellm/batches/test_batch_utils.py b/tests/test_litellm/batches/test_batch_utils.py index 70c89f70ceb..bf871623a94 100644 --- a/tests/test_litellm/batches/test_batch_utils.py +++ b/tests/test_litellm/batches/test_batch_utils.py @@ -18,7 +18,9 @@ import json import os import sys +import httpx import pytest +import respx sys.path.insert(0, os.path.abspath("../../../..")) @@ -711,6 +713,97 @@ async def test_output_file_content_vertex_unified_file_id_extracts_gcs_uri(monke assert captured["custom_llm_provider"] == "vertex_ai" +def _vertex_predictions_row(custom_id, prompt_tokens, completion_tokens): + return { + "request": { + "contents": [{"role": "user", "parts": [{"text": "hi"}]}], + "labels": {"litellm_custom_id": custom_id}, + }, + "status": "", + "response": { + "candidates": [ + { + "content": {"role": "model", "parts": [{"text": "ok"}]}, + "finishReason": "STOP", + } + ], + "usageMetadata": { + "promptTokenCount": prompt_tokens, + "candidatesTokenCount": completion_tokens, + "totalTokenCount": prompt_tokens + completion_tokens, + }, + "modelVersion": "gemini-3.6-flash", + }, + "processed_time": "2026-07-30T00:00:00.000000+00:00", + } + + +@pytest.fixture +def respx_interceptable_httpx_client(monkeypatch): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + yield + litellm.in_memory_llm_clients_cache.flush_cache() + + +@pytest.mark.asyncio +@respx.mock +async def test_output_file_content_vertex_managed_uri_accepted_by_real_validation(respx_interceptable_httpx_client): + managed_output_uri = ( + "gs://litellm-bucket/litellm-vertex-files/publishers/google/models/" + "gemini-3.6-flash/abc-123/prediction-model/predictions.jsonl" + ) + rows = [ + _vertex_predictions_row("request-1", 10, 5), + _vertex_predictions_row("request-2", 20, 10), + ] + route = respx.get(url__regex=r"https://storage\.googleapis\.com/storage/v1/b/litellm-bucket/o/.*").mock( + return_value=httpx.Response(200, content=_vertex_jsonl(rows)) + ) + + result = await bu._get_batch_output_file_content_as_dictionary( + _batch(managed_output_uri), + custom_llm_provider="vertex_ai", + litellm_params={ + "api_key": "test-token", + "vertex_project": "proj-1", + "vertex_location": "us-central1", + "gcs_bucket_name": "litellm-bucket", + }, + ) + + assert route.call_count == 1 + request = route.calls.last.request + assert request.url.raw_path == ( + b"/storage/v1/b/litellm-bucket/o/" + b"litellm-vertex-files%2Fpublishers%2Fgoogle%2Fmodels%2Fgemini-3.6-flash" + b"%2Fabc-123%2Fprediction-model%2Fpredictions.jsonl?alt=media" + ) + assert [row["custom_id"] for row in result] == ["request-1", "request-2"] + assert all(row["response"]["status_code"] == 200 for row in result) + assert all(row["response"]["body"]["model"] == "gemini-3.6-flash" for row in result) + assert [row["response"]["body"]["usage"]["prompt_tokens"] for row in result] == [10, 20] + assert [row["response"]["body"]["usage"]["completion_tokens"] for row in result] == [5, 10] + + +@pytest.mark.asyncio +@respx.mock +async def test_output_file_content_vertex_foreign_bucket_rejected_by_real_validation(): + with pytest.raises(Exception, match="does not match the configured storage bucket"): + await bu._get_batch_output_file_content_as_dictionary( + _batch("gs://attacker-bucket/litellm-vertex-files/x/predictions.jsonl"), + custom_llm_provider="vertex_ai", + litellm_params={ + "api_key": "test-token", + "vertex_project": "proj-1", + "vertex_location": "us-central1", + "gcs_bucket_name": "litellm-bucket", + }, + ) + + assert respx.mock.calls.call_count == 0 + + @pytest.mark.asyncio async def test_handle_completed_vertex_batch_computes_cost_usage_and_models(monkeypatch): import litellm.files.main as files_main From 76cf3bf6acaf900c0e00ba93d53c6745930e5cbd Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 30 Jul 2026 13:48:43 +0000 Subject: [PATCH 7/8] chore(typing): clear basedpyright Any errors in proxy auth, repositories, and openai transforms Replace `Model(**untyped_dict)` construction with `Model.model_validate(...)` at the hot Any seams, and give the repository layer a real record type instead of `Any`. reportAny 22710 -> 21448, reportExplicitAny 7283 -> 7269, with every other rule at or below its baseline repo-wide. --- basedpyright-code-budget.json | 10 ++-- litellm/llms/openai/openai.py | 24 ++++---- .../llms/openai/responses/transformation.py | 8 +-- .../mcp_server/mcp_server_manager.py | 2 +- litellm/proxy/auth/auth_checks.py | 4 +- litellm/proxy/auth/oauth2_proxy_hook.py | 19 +++--- litellm/proxy/auth/resolvers/store.py | 2 +- .../management_endpoints/common_utils.py | 4 +- .../mcp_management_endpoints.py | 4 +- litellm/proxy/management_helpers/utils.py | 8 +-- litellm/proxy/proxy_server.py | 8 +-- litellm/repositories/base_repository.py | 54 ++++++++++------- .../repositories/organization_repository.py | 6 +- litellm/repositories/project_repository.py | 9 +-- litellm/repositories/team_repository.py | 56 +++++++++--------- .../verification_token_repository.py | 58 +++++++++---------- ruff-strict-budget.json | 4 +- .../repositories/test_repositories.py | 19 ++++-- type-discipline-budget.json | 4 +- 19 files changed, 161 insertions(+), 142 deletions(-) diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index 3f89ff179eb..f4e506afcde 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -1,6 +1,6 @@ { "reportAny": { - "limit": 33171 + "limit": 31909 }, "reportArgumentType": { "limit": 2645 @@ -24,7 +24,7 @@ "limit": 42 }, "reportExplicitAny": { - "limit": 10228 + "limit": 10214 }, "reportFunctionMemberAccess": { "limit": 11 @@ -90,7 +90,7 @@ "limit": 12 }, "reportReturnType": { - "limit": 221 + "limit": 219 }, "reportTypedDictNotRequiredAccess": { "limit": 27 @@ -99,7 +99,7 @@ "limit": 0 }, "reportUnknownArgumentType": { - "limit": 45522 + "limit": 45366 }, "reportUnknownLambdaType": { "limit": 113 @@ -111,7 +111,7 @@ "limit": 20341 }, "reportUnknownVariableType": { - "limit": 32052 + "limit": 32051 }, "reportUnnecessaryCast": { "limit": 177 diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index 6b191144a11..8fc7e6d0ebd 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -1626,7 +1626,7 @@ class OpenAIFilesAPI(BaseLLM): openai_client: AsyncOpenAI, ) -> OpenAIFileObject: response = await openai_client.files.create(**create_file_data) # type: ignore[arg-type] - return OpenAIFileObject(**response.model_dump()) + return OpenAIFileObject.model_validate(response.model_dump()) def create_file( self, @@ -1662,7 +1662,7 @@ class OpenAIFilesAPI(BaseLLM): create_file_data=create_file_data, openai_client=openai_client ) response = cast(OpenAI, openai_client).files.create(**create_file_data) # type: ignore[arg-type] - return OpenAIFileObject(**response.model_dump()) + return OpenAIFileObject.model_validate(response.model_dump()) async def afile_content( self, @@ -1986,7 +1986,7 @@ class OpenAIBatchesAPI(BaseLLM): openai_client: AsyncOpenAI, ) -> LiteLLMBatch: response = await openai_client.batches.create(**create_batch_data) # type: ignore[arg-type] - return LiteLLMBatch(**response.model_dump()) + return LiteLLMBatch.model_validate(response.model_dump()) def create_batch( self, @@ -2023,7 +2023,7 @@ class OpenAIBatchesAPI(BaseLLM): ) response = cast(OpenAI, openai_client).batches.create(**create_batch_data) # type: ignore[arg-type] - return LiteLLMBatch(**response.model_dump()) + return LiteLLMBatch.model_validate(response.model_dump()) async def aretrieve_batch( self, @@ -2032,7 +2032,7 @@ class OpenAIBatchesAPI(BaseLLM): ) -> LiteLLMBatch: verbose_logger.debug("retrieving batch, args= %s", retrieve_batch_data) response = await openai_client.batches.retrieve(**retrieve_batch_data) # type: ignore[arg-type] - return LiteLLMBatch(**response.model_dump()) + return LiteLLMBatch.model_validate(response.model_dump()) def retrieve_batch( self, @@ -2068,7 +2068,7 @@ class OpenAIBatchesAPI(BaseLLM): retrieve_batch_data=retrieve_batch_data, openai_client=openai_client ) response = cast(OpenAI, openai_client).batches.retrieve(**retrieve_batch_data) # type: ignore[arg-type] - return LiteLLMBatch(**response.model_dump()) + return LiteLLMBatch.model_validate(response.model_dump()) async def acancel_batch( self, @@ -2077,7 +2077,7 @@ class OpenAIBatchesAPI(BaseLLM): ) -> LiteLLMBatch: verbose_logger.debug("async cancelling batch, args= %s", cancel_batch_data) response = await openai_client.batches.cancel(**cancel_batch_data) - return LiteLLMBatch(**response.model_dump()) + return LiteLLMBatch.model_validate(response.model_dump()) def cancel_batch( self, @@ -2117,7 +2117,7 @@ class OpenAIBatchesAPI(BaseLLM): if not isinstance(openai_client, OpenAI): raise ValueError("OpenAI client is not an instance of OpenAI. Make sure you passed a sync OpenAI client.") response = openai_client.batches.cancel(**cancel_batch_data) - return LiteLLMBatch(**response.model_dump()) + return LiteLLMBatch.model_validate(response.model_dump()) async def alist_batches( self, @@ -2477,9 +2477,9 @@ class OpenAIAssistantsAPI(BaseLLM): response_obj: Optional[OpenAIMessage] = None if getattr(thread_message, "status", None) is None: thread_message.status = "completed" - response_obj = OpenAIMessage(**thread_message.dict()) + response_obj = OpenAIMessage.model_validate(thread_message.dict()) else: - response_obj = OpenAIMessage(**thread_message.dict()) + response_obj = OpenAIMessage.model_validate(thread_message.dict()) return response_obj # fmt: off @@ -2556,9 +2556,9 @@ class OpenAIAssistantsAPI(BaseLLM): response_obj: Optional[OpenAIMessage] = None if getattr(thread_message, "status", None) is None: thread_message.status = "completed" - response_obj = OpenAIMessage(**thread_message.dict()) + response_obj = OpenAIMessage.model_validate(thread_message.dict()) else: - response_obj = OpenAIMessage(**thread_message.dict()) + response_obj = OpenAIMessage.model_validate(thread_message.dict()) return response_obj async def async_get_messages( diff --git a/litellm/llms/openai/responses/transformation.py b/litellm/llms/openai/responses/transformation.py index 3c2ae238a0b..dc4e98e6216 100644 --- a/litellm/llms/openai/responses/transformation.py +++ b/litellm/llms/openai/responses/transformation.py @@ -280,7 +280,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): raw_response_headers = dict(raw_response.headers) processed_headers = process_response_headers(raw_response_headers) try: - response = ResponsesAPIResponse(**raw_response_json) + response = ResponsesAPIResponse.model_validate(raw_response_json) except Exception: verbose_logger.debug(f"Error constructing ResponsesAPIResponse: {raw_response_json}, using model_construct") response = ResponsesAPIResponse.model_construct(**raw_response_json) @@ -506,7 +506,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): raise OpenAIError(message=raw_response.text, status_code=raw_response.status_code) raw_response_headers = dict(raw_response.headers) processed_headers = process_response_headers(raw_response_headers) - response = ResponsesAPIResponse(**raw_response_json) + response = ResponsesAPIResponse.model_validate(raw_response_json) response._hidden_params["additional_headers"] = processed_headers response._hidden_params["headers"] = raw_response_headers @@ -588,7 +588,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): raw_response_headers = dict(raw_response.headers) processed_headers = process_response_headers(raw_response_headers) - response = ResponsesAPIResponse(**raw_response_json) + response = ResponsesAPIResponse.model_validate(raw_response_json) response._hidden_params["additional_headers"] = processed_headers response._hidden_params["headers"] = raw_response_headers @@ -647,7 +647,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): processed_headers = process_response_headers(raw_response_headers) try: - response = ResponsesAPIResponse(**raw_response_json) + response = ResponsesAPIResponse.model_validate(raw_response_json) except Exception: verbose_logger.debug(f"Error constructing ResponsesAPIResponse: {raw_response_json}, using model_construct") response = ResponsesAPIResponse.model_construct(**raw_response_json) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 3e0775ac09e..f61ac4866b0 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -5322,7 +5322,7 @@ class MCPServerManager: ] } ) - db_mcp_servers = [LiteLLM_MCPServerTable(**r.model_dump()) for r in raw_rows] + db_mcp_servers = [LiteLLM_MCPServerTable.model_validate(r.model_dump()) for r in raw_rows] verbose_logger.info(f"Found {len(db_mcp_servers)} MCP servers in database") previous_registry = self.registry diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index c46bc110ca8..9d023292074 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -2434,7 +2434,7 @@ class ExperimentalUIJWTToken: if decrypted_token is None: return None try: - return UserAPIKeyAuth(**json.loads(decrypted_token)) + return UserAPIKeyAuth.model_validate(json.loads(decrypted_token)) except Exception as e: raise Exception(f"Invalid hash key. Hash key={hashed_token}. Decrypted token={decrypted_token}. Error: {e}") @@ -2553,7 +2553,7 @@ async def get_key_object( code=status.HTTP_401_UNAUTHORIZED, ) - _response = UserAPIKeyAuth(**_valid_token.model_dump(exclude_none=True)) + _response = UserAPIKeyAuth.model_validate(_valid_token.model_dump(exclude_none=True)) # Load object_permission if object_permission_id exists but object_permission is not loaded if _response.object_permission_id and not _response.object_permission: diff --git a/litellm/proxy/auth/oauth2_proxy_hook.py b/litellm/proxy/auth/oauth2_proxy_hook.py index 2b0593d3618..ca6a7ee4b1d 100644 --- a/litellm/proxy/auth/oauth2_proxy_hook.py +++ b/litellm/proxy/auth/oauth2_proxy_hook.py @@ -1,4 +1,5 @@ -from typing import Any, Dict, FrozenSet +from collections.abc import Mapping +from typing import Dict, FrozenSet, List, Union from fastapi import Request @@ -83,21 +84,17 @@ async def handle_oauth2_proxy_request(request: Request) -> UserAPIKeyAuth: "(signature-validated) instead of header-trust." ) - auth_data: Dict[str, Any] = {} - for key, header in oauth2_config_mappings.items(): - value = request.headers.get(header) - if not value: - continue - if key == "models": - auth_data[key] = [model.strip() for model in value.split(",")] - else: - auth_data[key] = value + auth_data: Mapping[str, Union[str, List[str]]] = { + key: [model.strip() for model in value.split(",")] if key == "models" else value + for key, header in oauth2_config_mappings.items() + if (value := request.headers.get(header)) + } verbose_proxy_logger.debug( "Auth data before creating UserAPIKeyAuth object: keys=%s", list(auth_data.keys()), ) - user_api_key_auth = UserAPIKeyAuth(**auth_data) + user_api_key_auth = UserAPIKeyAuth.model_validate(auth_data) verbose_proxy_logger.debug( "UserAPIKeyAuth object created with keys: %s", list(user_api_key_auth.__fields_set__), diff --git a/litellm/proxy/auth/resolvers/store.py b/litellm/proxy/auth/resolvers/store.py index ad9fd234163..7c2bd324064 100644 --- a/litellm/proxy/auth/resolvers/store.py +++ b/litellm/proxy/auth/resolvers/store.py @@ -118,7 +118,7 @@ class IdentityStore: if from_db is None: raise KeyNotFoundError(hashed_token) - key = UserAPIKeyAuth(**from_db.model_dump(exclude_none=True)) + key = UserAPIKeyAuth.model_validate(from_db.model_dump(exclude_none=True)) if key.object_permission_id and not key.object_permission: try: diff --git a/litellm/proxy/management_endpoints/common_utils.py b/litellm/proxy/management_endpoints/common_utils.py index 8162babef40..877130c2066 100644 --- a/litellm/proxy/management_endpoints/common_utils.py +++ b/litellm/proxy/management_endpoints/common_utils.py @@ -216,7 +216,7 @@ async def _user_has_admin_privileges( teams = await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": user_obj.teams}}) for team in teams: - team_obj = LiteLLM_TeamTable(**team.model_dump()) + team_obj = LiteLLM_TeamTable.model_validate(team.model_dump()) if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj): return True @@ -288,7 +288,7 @@ async def _team_admin_can_invite_user( for team in teams if _is_user_team_admin( user_api_key_dict=user_api_key_dict, - team_obj=LiteLLM_TeamTable(**team.model_dump()), + team_obj=LiteLLM_TeamTable.model_validate(team.model_dump()), ) ] if not admin_team_ids: diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 1205d23ce02..64cc13a5543 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -459,7 +459,7 @@ if MCP_AVAILABLE: payload_dict: dict[str, Any] = loaded try: - return MCPServer(**payload_dict) + return MCPServer.model_validate(payload_dict) except Exception as e: verbose_proxy_logger.debug(f"Invalid temporary MCP server payload in Redis cache: {str(e)}") return None @@ -704,7 +704,7 @@ if MCP_AVAILABLE: except AttributeError: payload_dict = payload.dict() # type: ignore[attr-defined] payload_dict["credentials"] = inherited_credentials - return NewMCPServerRequest(**payload_dict) + return NewMCPServerRequest.model_validate(payload_dict) def _build_temporary_mcp_server_record( payload: NewMCPServerRequest, diff --git a/litellm/proxy/management_helpers/utils.py b/litellm/proxy/management_helpers/utils.py index 7c52b04c4eb..14ba1dbfde1 100644 --- a/litellm/proxy/management_helpers/utils.py +++ b/litellm/proxy/management_helpers/utils.py @@ -308,7 +308,7 @@ async def add_new_member( ) await _append_team_id_if_absent(prisma_client, new_member.user_id, team_id) if _returned_user is not None: - returned_user = LiteLLM_UserTable(**_returned_user.model_dump()) + returned_user = LiteLLM_UserTable.model_validate(_returned_user.model_dump()) elif new_member.user_email is not None: new_user_defaults = get_new_internal_user_defaults(user_id=str(uuid.uuid4()), user_email=new_member.user_email) ## user email is not unique acc. to prisma schema -> future improvement @@ -323,11 +323,11 @@ async def add_new_member( _returned_user = await prisma_client.insert_data(data=new_user_defaults, table_name="user") # type: ignore if _returned_user is not None: - returned_user = LiteLLM_UserTable(**_returned_user.model_dump()) + returned_user = LiteLLM_UserTable.model_validate(_returned_user.model_dump()) elif len(existing_user_row) == 1: user_info = existing_user_row[0] await _append_team_id_if_absent(prisma_client, user_info.user_id, team_id) - returned_user = LiteLLM_UserTable(**user_info.model_dump()) + returned_user = LiteLLM_UserTable.model_validate(user_info.model_dump()) elif len(existing_user_row) > 1: raise HTTPException( status_code=400, @@ -354,7 +354,7 @@ async def add_new_member( include={"litellm_budget_table": True}, ) - returned_team_membership = LiteLLM_TeamMembership(**_returned_team_membership.model_dump()) + returned_team_membership = LiteLLM_TeamMembership.model_validate(_returned_team_membership.model_dump()) if returned_user is None: raise Exception("Unable to update user table with membership information!") diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 18a927e7a44..ba41083e9b8 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -5398,7 +5398,7 @@ class ProxyConfig: # decrypt values for k, v in _litellm_params.items(): _litellm_params[k] = self._resolve_db_litellm_param(key=k, value=v) - _litellm_params = LiteLLM_Params(**_litellm_params) + _litellm_params = LiteLLM_Params.model_validate(_litellm_params) else: verbose_proxy_logger.error( @@ -5429,7 +5429,7 @@ class ProxyConfig: # decrypt values for k, v in _litellm_params.items(): _litellm_params[k] = self._resolve_db_litellm_param(key=k, value=v) - _litellm_params = LiteLLM_Params(**_litellm_params) + _litellm_params = LiteLLM_Params.model_validate(_litellm_params) else: verbose_proxy_logger.error( f"Invalid model added to proxy db. Invalid litellm params. litellm_params={_litellm_params}" @@ -13063,7 +13063,7 @@ def _get_model_group_info( _model_group_info = llm_router.get_model_group_info(model_group=model) if _model_group_info is not None: - model_groups.append(ModelGroupInfoProxy(**_model_group_info.model_dump())) + model_groups.append(ModelGroupInfoProxy.model_validate(_model_group_info.model_dump())) else: model_group_info = ModelGroupInfoProxy( model_group=model, @@ -14782,7 +14782,7 @@ async def update_config_general_settings( ) try: - ConfigGeneralSettings(**{data.field_name: data.field_value}) + ConfigGeneralSettings.model_validate({data.field_name: data.field_value}) except Exception: raise HTTPException( status_code=400, diff --git a/litellm/repositories/base_repository.py b/litellm/repositories/base_repository.py index 40aeb6df3de..755e4595c01 100644 --- a/litellm/repositories/base_repository.py +++ b/litellm/repositories/base_repository.py @@ -3,38 +3,58 @@ Base repository class with common functionality. """ from abc import ABC, abstractmethod -from typing import Any, Dict, Generic, List, Optional, Type, TypeVar +from collections.abc import Iterable, Mapping, Sequence +from typing import Any, Dict, Generic, List, Optional, Protocol, Tuple, Type, TypeVar, Union, runtime_checkable from pydantic import BaseModel T = TypeVar("T", bound=BaseModel) -def _record_to_dict(record: Any) -> Dict[str, Any]: - if isinstance(record, dict): - return record - if hasattr(record, "model_dump") and callable(record.model_dump): +@runtime_checkable +class SupportsModelDump(Protocol): + def model_dump(self) -> Dict[str, object]: ... + + +@runtime_checkable +class SupportsDict(Protocol): + def dict(self) -> Dict[str, object]: ... + + +DbRecord = Union[ + Mapping[str, object], + SupportsModelDump, + SupportsDict, + Sequence[Tuple[str, object]], +] + + +def record_to_dict(record: DbRecord) -> Mapping[str, object]: + """Project a database record into a mapping of column name to value.""" + if isinstance(record, SupportsModelDump): return record.model_dump() - if hasattr(record, "dict") and callable(record.dict): + if isinstance(record, SupportsDict): return record.dict() - return dict(record) + if isinstance(record, Mapping): + return record + return {key: value for key, value in record} class BaseRepository(ABC, Generic[T]): """Abstract base class for all repositories.""" - def __init__(self, prisma_client: Any): + def __init__(self, prisma_client: Any): # any-ok: PrismaClient is an untyped runtime wrapper self._prisma_client = prisma_client @property - def prisma_client(self) -> Any: + def prisma_client(self) -> Any: # any-ok: PrismaClient is an untyped runtime wrapper if self._prisma_client is None: raise RuntimeError("No DB Connected. See - https://docs.litellm.ai/docs/proxy/virtual_keys") return self._prisma_client @property @abstractmethod - def table(self) -> Any: + def table(self) -> Any: # any-ok: Prisma table actions are reached through the untyped client wrapper """Return the Prisma table for this repository.""" ... @@ -44,21 +64,15 @@ class BaseRepository(ABC, Generic[T]): """Return the domain model class for this repository.""" ... - def _to_model(self, record: Any) -> Optional[T]: + def _to_model(self, record: Optional[DbRecord]) -> Optional[T]: """Convert a database record to a domain model.""" if record is None: return None - return self.model_class(**_record_to_dict(record)) + return self.model_class.model_validate(record_to_dict(record)) - def _to_model_list(self, records: List[Any]) -> List[T]: + def _to_model_list(self, records: Iterable[Optional[DbRecord]]) -> List[T]: """Convert a list of database records to domain models.""" - result: List[T] = [] - for r in records: - if r is not None: - model = self._to_model(r) - if model is not None: - result.append(model) - return result + return [model for record in records if record is not None and (model := self._to_model(record)) is not None] async def find_by_id(self, id_value: str, id_field: str = "id") -> Optional[T]: """Find a record by its primary key.""" diff --git a/litellm/repositories/organization_repository.py b/litellm/repositories/organization_repository.py index d5f8c990001..99c4a881736 100644 --- a/litellm/repositories/organization_repository.py +++ b/litellm/repositories/organization_repository.py @@ -26,10 +26,8 @@ class OrganizationRepository(BaseRepository[LiteLLM_OrganizationTable]): async def find_by_alias(self, organization_alias: str) -> Optional[LiteLLM_OrganizationTable]: """Find an organization by alias.""" - records = await self.table.find_many(where={"organization_alias": organization_alias}) - if records: - return self._to_model(records[0]) - return None + organizations = await self.find_many(where={"organization_alias": organization_alias}) + return organizations[0] if organizations else None async def create_organization( self, diff --git a/litellm/repositories/project_repository.py b/litellm/repositories/project_repository.py index 86faaf2e13c..27cb346e1b1 100644 --- a/litellm/repositories/project_repository.py +++ b/litellm/repositories/project_repository.py @@ -24,15 +24,12 @@ class ProjectRepository(BaseRepository[LiteLLM_ProjectTable]): async def find_by_alias(self, project_alias: str) -> Optional[LiteLLM_ProjectTable]: """Find a project by alias.""" - records = await self.table.find_many(where={"project_alias": project_alias}) - if records: - return self._to_model(records[0]) - return None + projects = await self.find_many(where={"project_alias": project_alias}) + return projects[0] if projects else None async def find_by_team_id(self, team_id: str) -> List[LiteLLM_ProjectTable]: """Find all projects belonging to a team.""" - records = await self.table.find_many(where={"team_id": team_id}) - return self._to_model_list(records) + return await self.find_many(where={"team_id": team_id}) async def create_project( self, diff --git a/litellm/repositories/team_repository.py b/litellm/repositories/team_repository.py index 68875bd7972..25437cfe49a 100644 --- a/litellm/repositories/team_repository.py +++ b/litellm/repositories/team_repository.py @@ -3,55 +3,59 @@ Team repository for database operations on LiteLLM_TeamTable. """ import json +from collections.abc import Mapping from datetime import datetime from typing import TYPE_CHECKING, Any, Dict, List, Optional, Type from pydantic import TypeAdapter from litellm.models.team import LiteLLM_TeamTable, Member -from litellm.repositories.base_repository import BaseRepository +from litellm.repositories.base_repository import ( + BaseRepository, + DbRecord, + record_to_dict, +) if TYPE_CHECKING: from prisma import Prisma _MEMBERS_WITH_ROLES_ADAPTER = TypeAdapter(list[Member]) +_JSON_ENCODED_TEAM_FIELDS = ( + "metadata", + "model_spend", + "model_max_budget", + "router_settings", + "budget_limits", + "members_with_roles", +) class TeamRepository(BaseRepository[LiteLLM_TeamTable]): """Repository for team database operations.""" @property - def table(self) -> Any: + def table(self) -> Any: # any-ok: PrismaClient.db is an untyped runtime wrapper return self.prisma_client.db.litellm_teamtable @property - def deleted_table(self) -> Any: + def deleted_table(self) -> Any: # any-ok: PrismaClient.db is an untyped runtime wrapper return self.prisma_client.db.litellm_deletedteamtable @property def model_class(self) -> Type[LiteLLM_TeamTable]: return LiteLLM_TeamTable - def _to_model(self, record: Any) -> Optional[LiteLLM_TeamTable]: + def _to_model(self, record: Optional[DbRecord]) -> Optional[LiteLLM_TeamTable]: """Convert a database record to a Team model.""" if record is None: return None - data = record.dict() if hasattr(record, "dict") else dict(record) + data = { + field: json.loads(value) if field in _JSON_ENCODED_TEAM_FIELDS and isinstance(value, str) else value + for field, value in record_to_dict(record).items() + } - json_fields = [ - "metadata", - "model_spend", - "model_max_budget", - "router_settings", - "budget_limits", - "members_with_roles", - ] - for field in json_fields: - if isinstance(data.get(field), str): - data[field] = json.loads(data[field]) - - return LiteLLM_TeamTable(**data) + return LiteLLM_TeamTable.model_validate(data) async def get_members_with_roles_locked(self, tx: "Prisma", team_id: str) -> List[Member]: """Return the team's members_with_roles, locking the row FOR UPDATE. @@ -103,8 +107,8 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): organization_id: Optional[str] = None, admins: Optional[List[str]] = None, members: Optional[List[str]] = None, - members_with_roles: Optional[Dict[str, Any]] = None, - metadata: Optional[Dict[str, Any]] = None, + members_with_roles: Optional[Mapping[str, object]] = None, + metadata: Optional[Mapping[str, object]] = None, max_budget: Optional[float] = None, soft_budget: Optional[float] = None, models: Optional[List[str]] = None, @@ -115,7 +119,7 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): object_permission_id: Optional[str] = None, ) -> LiteLLM_TeamTable: """Create a new team.""" - data: Dict[str, Any] = {"team_id": team_id} + data: Dict[str, object] = {"team_id": team_id} if team_alias is not None: data["team_alias"] = team_alias if organization_id is not None: @@ -154,8 +158,8 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): organization_id: Optional[str] = None, admins: Optional[List[str]] = None, members: Optional[List[str]] = None, - members_with_roles: Optional[Dict[str, Any]] = None, - metadata: Optional[Dict[str, Any]] = None, + members_with_roles: Optional[Mapping[str, object]] = None, + metadata: Optional[Mapping[str, object]] = None, max_budget: Optional[float] = None, soft_budget: Optional[float] = None, models: Optional[List[str]] = None, @@ -167,7 +171,7 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): object_permission_id: Optional[str] = None, ) -> Optional[LiteLLM_TeamTable]: """Update a team.""" - data: Dict[str, Any] = {} + data: Dict[str, object] = {} if team_alias is not None: data["team_alias"] = team_alias if organization_id is not None: @@ -228,9 +232,9 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): return team - def _build_archive_data(self, team: LiteLLM_TeamTable) -> Dict[str, Any]: + def _build_archive_data(self, team: LiteLLM_TeamTable) -> Dict[str, object]: """Build archive data dict with only columns that exist in LiteLLM_DeletedTeamTable.""" - data: Dict[str, Any] = {"team_id": team.team_id} + data: Dict[str, object] = {"team_id": team.team_id} if team.team_alias is not None: data["team_alias"] = team.team_alias if team.organization_id is not None: diff --git a/litellm/repositories/verification_token_repository.py b/litellm/repositories/verification_token_repository.py index 19352c1b3c4..f7795f15fd5 100644 --- a/litellm/repositories/verification_token_repository.py +++ b/litellm/repositories/verification_token_repository.py @@ -3,14 +3,18 @@ VerificationToken repository for database operations on LiteLLM_VerificationToke """ import json -from collections.abc import Iterator, Mapping +from collections.abc import Mapping from datetime import datetime -from typing import TYPE_CHECKING, Any, Protocol +from typing import TYPE_CHECKING, Any from litellm.models.verification_token import ( LiteLLM_VerificationToken, ) -from litellm.repositories.base_repository import BaseRepository +from litellm.repositories.base_repository import ( + BaseRepository, + DbRecord, + record_to_dict, +) if TYPE_CHECKING: from prisma.models import ( @@ -19,11 +23,17 @@ if TYPE_CHECKING: from litellm.proxy.utils import PrismaClient - -class _DictConvertible(Protocol): - def dict(self) -> dict[str, object]: ... - - def __iter__(self) -> Iterator[tuple[str, object]]: ... +_JSON_ENCODED_TOKEN_FIELDS = ( + "aliases", + "config", + "permissions", + "metadata", + "model_spend", + "model_max_budget", + "router_settings", + "budget_limits", + "litellm_budget_table", +) class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]): @@ -46,31 +56,21 @@ class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]): def model_class(self) -> type[LiteLLM_VerificationToken]: return LiteLLM_VerificationToken - def _to_model(self, record: _DictConvertible | None) -> LiteLLM_VerificationToken | None: + def _to_model(self, record: DbRecord | None) -> LiteLLM_VerificationToken | None: """Convert a database record to a VerificationToken model.""" if record is None: return None - data = record.dict() if hasattr(record, "dict") else dict(record) - - json_fields = [ - "aliases", - "config", - "permissions", - "metadata", - "model_spend", - "model_max_budget", - "router_settings", - "budget_limits", - "litellm_budget_table", - ] - for field in json_fields: - value = data.get(field) - if isinstance(value, str): - data[field] = json.loads(value) - - if data.get("org_id") is None and data.get("organization_id") is not None: - data["org_id"] = data["organization_id"] + decoded = { + field: json.loads(value) if field in _JSON_ENCODED_TOKEN_FIELDS and isinstance(value, str) else value + for field, value in record_to_dict(record).items() + } + organization_id = decoded.get("organization_id") + data = ( + decoded + if decoded.get("org_id") is not None or organization_id is None + else {**decoded, "org_id": organization_id} + ) return LiteLLM_VerificationToken.model_validate(data) diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 3b5ec5b0dee..f1c205c2425 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -24,7 +24,7 @@ "limit": 130 }, "ANN401": { - "limit": 2013 + "limit": 2010 }, "ASYNC230": { "limit": 14 @@ -324,7 +324,7 @@ "limit": 883 }, "UP006": { - "limit": 12146 + "limit": 12142 }, "UP007": { "limit": 2526 diff --git a/tests/test_litellm/repositories/test_repositories.py b/tests/test_litellm/repositories/test_repositories.py index 6308faf8fc7..c923b722991 100644 --- a/tests/test_litellm/repositories/test_repositories.py +++ b/tests/test_litellm/repositories/test_repositories.py @@ -203,23 +203,32 @@ class TestBaseRepository: assert len(budgets) == 1 def test_record_to_dict_branches(self): - from litellm.repositories.base_repository import _record_to_dict + from litellm.repositories.base_repository import record_to_dict - assert _record_to_dict({"a": 1}) == {"a": 1} + assert record_to_dict({"a": 1}) == {"a": 1} class WithModelDump: def model_dump(self): return {"src": "model_dump"} - assert _record_to_dict(WithModelDump()) == {"src": "model_dump"} + assert record_to_dict(WithModelDump()) == {"src": "model_dump"} class WithDict: def dict(self): return {"src": "dict"} - assert _record_to_dict(WithDict()) == {"src": "dict"} + assert record_to_dict(WithDict()) == {"src": "dict"} - assert _record_to_dict([("k", "v")]) == {"k": "v"} + assert record_to_dict([("k", "v")]) == {"k": "v"} + + class WithBoth: + def model_dump(self): + return {"src": "model_dump"} + + def dict(self): + return {"src": "dict"} + + assert record_to_dict(WithBoth()) == {"src": "model_dump"} class TestBudgetRepository: diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 25cc1621d54..69b2506dc08 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,9 +1,9 @@ { "LIT001": { - "limit": 23267 + "limit": 23261 }, "LIT002": { - "limit": 27434 + "limit": 27433 }, "LIT003": { "limit": 292 From 59118ae5b648666055e240f39b6301601a580920 Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Thu, 30 Jul 2026 10:29:40 -0700 Subject: [PATCH 8/8] feat(cookbook): add a Grafana dashboard for the OTel GenAI metrics (#35159) The existing dashboards in this folder chart the litellm_* Prometheus metrics. Nothing charted the gen_ai.* metrics the OpenTelemetry v2 integration emits, and Grafana's own prebuilt GenAI dashboards cannot: twenty of their twenty-two panels filter on telemetry_sdk_name="openlit", a label LiteLLM does not carry and has no setting to add. Ten panels over the six gen_ai instruments: spend, tokens, request count and p95 duration as stats, then request rate, spend per hour, tokens per minute split by input and output, and p95 duration, time to first token, and provider generation time by model. Template variables for data source, service, and model. Verified against a live Grafana Cloud stack with real traffic across three models. The readme documents the attribute filter the panels depend on, since the default attribute set gives nearly every request its own series and makes every rate-based panel read zero. --- .../grafana_dashboard.json | 523 ++++++++++++++++++ .../dashboard_genai_otel/readme.md | 35 ++ .../grafana_dashboard/readme.md | 4 + 3 files changed, 562 insertions(+) create mode 100644 cookbook/litellm_proxy_server/grafana_dashboard/dashboard_genai_otel/grafana_dashboard.json create mode 100644 cookbook/litellm_proxy_server/grafana_dashboard/dashboard_genai_otel/readme.md diff --git a/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_genai_otel/grafana_dashboard.json b/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_genai_otel/grafana_dashboard.json new file mode 100644 index 00000000000..70608a2ffe8 --- /dev/null +++ b/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_genai_otel/grafana_dashboard.json @@ -0,0 +1,523 @@ +{ + "annotations": { + "list": [] + }, + "editable": true, + "fiscalYearStartMonth": 0, + "graphTooltip": 0, + "links": [], + "panels": [ + { + "type": "stat", + "title": "Requests", + "gridPos": { + "h": 4, + "w": 6, + "x": 0, + "y": 0 + }, + "datasource": { + "type": "prometheus", + "uid": "${datasource}" + }, + "fieldConfig": { + "defaults": { + "unit": "short", + "decimals": 0, + "color": { + "mode": "fixed", + "fixedColor": "blue" + } + }, + "overrides": [] + }, + "options": { + "reduceOptions": { + "calcs": [ + "lastNotNull" + ], + "fields": "", + "values": false + }, + "colorMode": "background", + "graphMode": "none" + }, + "targets": [ + { + "refId": "A", + "editorMode": "code", + "instant": true, + "expr": "sum(increase(gen_ai_client_operation_duration_seconds_count{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__range]))" + } + ], + "id": 1 + }, + { + "type": "stat", + "title": "Spend", + "description": "LiteLLM's computed cost for the selected window, from gen_ai.usage.cost", + "gridPos": { + "h": 4, + "w": 6, + "x": 6, + "y": 0 + }, + "datasource": { + "type": "prometheus", + "uid": "${datasource}" + }, + "fieldConfig": { + "defaults": { + "unit": "currencyUSD", + "decimals": 4, + "color": { + "mode": "fixed", + "fixedColor": "green" + } + }, + "overrides": [] + }, + "options": { + "reduceOptions": { + "calcs": [ + "lastNotNull" + ], + "fields": "", + "values": false + }, + "colorMode": "background", + "graphMode": "none" + }, + "targets": [ + { + "refId": "A", + "editorMode": "code", + "instant": true, + "expr": "sum(increase(gen_ai_usage_cost_USD_sum{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__range]))" + } + ], + "id": 2 + }, + { + "type": "stat", + "title": "Tokens", + "gridPos": { + "h": 4, + "w": 6, + "x": 12, + "y": 0 + }, + "datasource": { + "type": "prometheus", + "uid": "${datasource}" + }, + "fieldConfig": { + "defaults": { + "unit": "short", + "decimals": 0, + "color": { + "mode": "fixed", + "fixedColor": "purple" + } + }, + "overrides": [] + }, + "options": { + "reduceOptions": { + "calcs": [ + "lastNotNull" + ], + "fields": "", + "values": false + }, + "colorMode": "background", + "graphMode": "none" + }, + "targets": [ + { + "refId": "A", + "editorMode": "code", + "instant": true, + "expr": "sum(increase(gen_ai_client_token_usage_sum{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__range]))" + } + ], + "id": 3 + }, + { + "type": "stat", + "title": "p95 request duration", + "gridPos": { + "h": 4, + "w": 6, + "x": 18, + "y": 0 + }, + "datasource": { + "type": "prometheus", + "uid": "${datasource}" + }, + "fieldConfig": { + "defaults": { + "unit": "s", + "decimals": 2, + "color": { + "mode": "fixed", + "fixedColor": "orange" + } + }, + "overrides": [] + }, + "options": { + "reduceOptions": { + "calcs": [ + "lastNotNull" + ], + "fields": "", + "values": false + }, + "colorMode": "background", + "graphMode": "none" + }, + "targets": [ + { + "refId": "A", + "editorMode": "code", + "instant": true, + "expr": "histogram_quantile(0.95, sum by (le) (rate(gen_ai_client_operation_duration_seconds_bucket{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__range])))" + } + ], + "id": 4 + }, + { + "type": "timeseries", + "title": "Request rate by model", + "gridPos": { + "h": 8, + "w": 12, + "x": 0, + "y": 4 + }, + "datasource": { + "type": "prometheus", + "uid": "${datasource}" + }, + "fieldConfig": { + "defaults": { + "unit": "reqpm", + "custom": { + "lineWidth": 2, + "fillOpacity": 8, + "showPoints": "never" + } + }, + "overrides": [] + }, + "options": { + "legend": { + "displayMode": "list", + "placement": "bottom" + }, + "tooltip": { + "mode": "multi", + "sort": "desc" + } + }, + "targets": [ + { + "refId": "A", + "editorMode": "code", + "legendFormat": "{{gen_ai_request_model}}", + "expr": "sum by (gen_ai_request_model) (rate(gen_ai_client_operation_duration_seconds_count{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])) * 60" + } + ], + "id": 5 + }, + { + "type": "timeseries", + "title": "Spend rate by model", + "description": "USD per hour, derived from the gen_ai.usage.cost histogram", + "gridPos": { + "h": 8, + "w": 12, + "x": 12, + "y": 4 + }, + "datasource": { + "type": "prometheus", + "uid": "${datasource}" + }, + "fieldConfig": { + "defaults": { + "unit": "currencyUSD", + "custom": { + "lineWidth": 2, + "fillOpacity": 8, + "showPoints": "never" + } + }, + "overrides": [] + }, + "options": { + "legend": { + "displayMode": "list", + "placement": "bottom" + }, + "tooltip": { + "mode": "multi", + "sort": "desc" + } + }, + "targets": [ + { + "refId": "A", + "editorMode": "code", + "legendFormat": "{{gen_ai_request_model}}", + "expr": "sum by (gen_ai_request_model) (rate(gen_ai_usage_cost_USD_sum{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])) * 3600" + } + ], + "id": 6 + }, + { + "type": "timeseries", + "title": "Tokens per minute by model and type", + "description": "gen_ai.client.token.usage split by the gen_ai.token.type attribute", + "gridPos": { + "h": 8, + "w": 12, + "x": 0, + "y": 12 + }, + "datasource": { + "type": "prometheus", + "uid": "${datasource}" + }, + "fieldConfig": { + "defaults": { + "unit": "short", + "custom": { + "lineWidth": 2, + "fillOpacity": 8, + "showPoints": "never" + } + }, + "overrides": [] + }, + "options": { + "legend": { + "displayMode": "list", + "placement": "bottom" + }, + "tooltip": { + "mode": "multi", + "sort": "desc" + } + }, + "targets": [ + { + "refId": "A", + "editorMode": "code", + "legendFormat": "{{gen_ai_request_model}} {{gen_ai_token_type}}", + "expr": "sum by (gen_ai_request_model, gen_ai_token_type) (rate(gen_ai_client_token_usage_sum{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])) * 60" + } + ], + "id": 7 + }, + { + "type": "timeseries", + "title": "p95 request duration by model", + "gridPos": { + "h": 8, + "w": 12, + "x": 12, + "y": 12 + }, + "datasource": { + "type": "prometheus", + "uid": "${datasource}" + }, + "fieldConfig": { + "defaults": { + "unit": "s", + "custom": { + "lineWidth": 2, + "fillOpacity": 0, + "showPoints": "never" + } + }, + "overrides": [] + }, + "options": { + "legend": { + "displayMode": "list", + "placement": "bottom" + }, + "tooltip": { + "mode": "multi", + "sort": "desc" + } + }, + "targets": [ + { + "refId": "A", + "editorMode": "code", + "legendFormat": "{{gen_ai_request_model}}", + "expr": "histogram_quantile(0.95, sum by (le, gen_ai_request_model) (rate(gen_ai_client_operation_duration_seconds_bucket{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])))" + } + ], + "id": 8 + }, + { + "type": "timeseries", + "title": "p95 time to first token (streaming)", + "description": "gen_ai.server.time_to_first_token, recorded only for streaming requests", + "gridPos": { + "h": 8, + "w": 12, + "x": 0, + "y": 20 + }, + "datasource": { + "type": "prometheus", + "uid": "${datasource}" + }, + "fieldConfig": { + "defaults": { + "unit": "s", + "custom": { + "lineWidth": 2, + "fillOpacity": 0, + "showPoints": "never" + } + }, + "overrides": [] + }, + "options": { + "legend": { + "displayMode": "list", + "placement": "bottom" + }, + "tooltip": { + "mode": "multi", + "sort": "desc" + } + }, + "targets": [ + { + "refId": "A", + "editorMode": "code", + "legendFormat": "{{gen_ai_request_model}}", + "expr": "histogram_quantile(0.95, sum by (le, gen_ai_request_model) (rate(gen_ai_server_time_to_first_token_seconds_bucket{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])))" + } + ], + "id": 9 + }, + { + "type": "timeseries", + "title": "p95 provider generation time", + "description": "gen_ai.client.response.duration, upstream generation time excluding LiteLLM overhead", + "gridPos": { + "h": 8, + "w": 12, + "x": 12, + "y": 20 + }, + "datasource": { + "type": "prometheus", + "uid": "${datasource}" + }, + "fieldConfig": { + "defaults": { + "unit": "s", + "custom": { + "lineWidth": 2, + "fillOpacity": 0, + "showPoints": "never" + } + }, + "overrides": [] + }, + "options": { + "legend": { + "displayMode": "list", + "placement": "bottom" + }, + "tooltip": { + "mode": "multi", + "sort": "desc" + } + }, + "targets": [ + { + "refId": "A", + "editorMode": "code", + "legendFormat": "{{gen_ai_request_model}}", + "expr": "histogram_quantile(0.95, sum by (le, gen_ai_request_model) (rate(gen_ai_client_response_duration_seconds_bucket{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])))" + } + ], + "id": 10 + } + ], + "preload": false, + "refresh": "30s", + "schemaVersion": 42, + "tags": [ + "litellm", + "genai", + "opentelemetry" + ], + "templating": { + "list": [ + { + "name": "datasource", + "label": "Prometheus", + "type": "datasource", + "query": "prometheus", + "current": {}, + "hide": 0 + }, + { + "name": "service", + "label": "Service", + "type": "query", + "datasource": { + "type": "prometheus", + "uid": "${datasource}" + }, + "query": "label_values(gen_ai_client_operation_duration_seconds_count, service_name)", + "refresh": 2, + "includeAll": true, + "multi": true, + "current": { + "text": "All", + "value": "$__all" + } + }, + { + "name": "model", + "label": "Model", + "type": "query", + "datasource": { + "type": "prometheus", + "uid": "${datasource}" + }, + "query": "label_values(gen_ai_client_operation_duration_seconds_count{service_name=~\"$service\"}, gen_ai_request_model)", + "refresh": 2, + "includeAll": true, + "multi": true, + "current": { + "text": "All", + "value": "$__all" + } + } + ] + }, + "time": { + "from": "now-1h", + "to": "now" + }, + "timepicker": {}, + "timezone": "browser", + "title": "LiteLLM GenAI (OpenTelemetry)", + "uid": "litellm-genai-otel", + "weekStart": "" +} diff --git a/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_genai_otel/readme.md b/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_genai_otel/readme.md new file mode 100644 index 00000000000..c51f0166462 --- /dev/null +++ b/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_genai_otel/readme.md @@ -0,0 +1,35 @@ +# LiteLLM GenAI dashboard (OpenTelemetry metrics) + +Dashboard for the `gen_ai.*` metrics the OpenTelemetry v2 integration emits, as opposed to the `litellm_*` Prometheus metrics the other dashboards in this folder chart. + +Import `grafana_dashboard.json` from **Dashboards > New > Import** and pick your Prometheus data source. Panels: request count, spend, token count, p95 duration, request rate by model, spend rate per hour by model, tokens per minute split by input and output, p95 duration by model, p95 time to first token, and p95 provider generation time. Template variables for data source, service, and model. + +## Pre-requisites + +Metrics are off by default. In the proxy environment: + +```shell +LITELLM_OTEL_V2=true +LITELLM_OTEL_INTEGRATION_ENABLE_METRICS=true +OTEL_EXPORTER="otlp_http" +OTEL_ENDPOINT="" +``` + +You also need the metric attribute filter, or the panels will plot flat lines at zero. LiteLLM's default attribute set includes per-request fields, so nearly every request lands in its own time series with a single sample, and `rate()` has nothing to compute over: + +```yaml title="config.yaml" +callback_settings: + otel: + attributes: + include_list: + - gen_ai.operation.name + - gen_ai.system + - gen_ai.request.model + - gen_ai.framework +``` + +See [Grafana Cloud](https://docs.litellm.ai/docs/observability/grafana_cloud) for the full setup, and [OpenTelemetry v2](https://docs.litellm.ai/docs/observability/opentelemetry_v2#metrics) for the metric reference. + +## Note on Grafana's AI Observability integration + +Grafana Cloud ships prebuilt GenAI dashboards that query these same metric names, so they look like a drop-in alternative to this one. They are not: twenty of their twenty-two panels filter on `telemetry_sdk_name="openlit"`, a label LiteLLM does not carry and cannot be configured to add, so those panels stay empty. diff --git a/cookbook/litellm_proxy_server/grafana_dashboard/readme.md b/cookbook/litellm_proxy_server/grafana_dashboard/readme.md index 81235c308f2..a1564a406e0 100644 --- a/cookbook/litellm_proxy_server/grafana_dashboard/readme.md +++ b/cookbook/litellm_proxy_server/grafana_dashboard/readme.md @@ -2,6 +2,10 @@ This folder contains the `json` for creating Grafana Dashboards +## [LiteLLM GenAI Dashboard (OpenTelemetry)](./dashboard_genai_otel) + +Charts the `gen_ai.*` metrics from the OpenTelemetry v2 integration: spend, tokens, request rate, and latency percentiles by model. Separate from the dashboards below, which chart the `litellm_*` Prometheus metrics. + ## [LiteLLM v2 Dashboard](./dashboard_v2) grafana_1