From 4bd7bdcf43da7106bd07240acc2bd0b1c997936c Mon Sep 17 00:00:00 2001 From: Chesars Date: Wed, 18 Mar 2026 22:30:49 -0300 Subject: [PATCH 01/42] fix: add additionalProperties: false for OpenAI strict mode in Anthropic adapter When translating Anthropic output_format to OpenAI response_format, the adapter sets strict: true but didn't add additionalProperties: false, which OpenAI requires at every object nesting level. This caused BadRequestError for structured output requests routed to OpenAI models. Fixes #20997 --- .../adapters/transformation.py | 37 +++++++ ...al_pass_through_adapters_transformation.py | 101 ++++++++++++++++++ 2 files changed, 138 insertions(+) diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py index 43a6fa8045d..47c9a223f85 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py @@ -1,3 +1,4 @@ +import copy import hashlib import json from typing import ( @@ -824,6 +825,11 @@ class LiteLLMAnthropicMessagesAdapter: if not schema: return None + # Deep copy to avoid mutating the original schema + schema = copy.deepcopy(schema) + # OpenAI strict mode requires additionalProperties: false on every object + self._add_additional_properties_false(schema) + # Convert to OpenAI response_format structure return { "type": "json_schema", @@ -834,6 +840,37 @@ class LiteLLMAnthropicMessagesAdapter: }, } + @staticmethod + def _add_additional_properties_false(schema: dict) -> None: + """ + Recursively add 'additionalProperties': false to all object schemas. + + OpenAI's strict mode requires this at every object nesting level. + """ + if not isinstance(schema, dict): + return + + if schema.get("type") == "object" and "properties" in schema: + schema["additionalProperties"] = False + for prop in schema["properties"].values(): + LiteLLMAnthropicMessagesAdapter._add_additional_properties_false(prop) + + # Handle array items + if "items" in schema: + LiteLLMAnthropicMessagesAdapter._add_additional_properties_false(schema["items"]) + + # Handle anyOf/oneOf/allOf + for key in ("anyOf", "oneOf", "allOf"): + if key in schema: + for sub_schema in schema[key]: + LiteLLMAnthropicMessagesAdapter._add_additional_properties_false(sub_schema) + + # Handle $defs / definitions + for key in ("$defs", "definitions"): + if key in schema: + for def_schema in schema[key].values(): + LiteLLMAnthropicMessagesAdapter._add_additional_properties_false(def_schema) + def _add_system_message_to_messages( self, new_messages: List[AllMessageValues], diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py index 839d032c436..b0442ebf0eb 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py @@ -1984,3 +1984,104 @@ def test_translate_anthropic_to_openai_with_mixed_tools(): # tool_name_mapping should be empty for short tool names assert tool_name_mapping == {} + + +class TestTranslateAnthropicOutputFormatToOpenAI: + """Tests for translate_anthropic_output_format_to_openai adding additionalProperties: false.""" + + def setup_method(self): + self.adapter = LiteLLMAnthropicMessagesAdapter() + + def test_simple_object_adds_additional_properties_false(self): + output_format = { + "type": "json_schema", + "schema": { + "type": "object", + "properties": {"name": {"type": "string"}}, + }, + } + result = self.adapter.translate_anthropic_output_format_to_openai(output_format) + assert result is not None + schema = result["json_schema"]["schema"] + assert schema["additionalProperties"] is False + + def test_nested_objects_adds_additional_properties_false(self): + output_format = { + "type": "json_schema", + "schema": { + "type": "object", + "properties": { + "user": { + "type": "object", + "properties": { + "name": {"type": "string"}, + "address": { + "type": "object", + "properties": {"city": {"type": "string"}}, + }, + }, + } + }, + }, + } + result = self.adapter.translate_anthropic_output_format_to_openai(output_format) + assert result is not None + schema = result["json_schema"]["schema"] + assert schema["additionalProperties"] is False + assert schema["properties"]["user"]["additionalProperties"] is False + assert schema["properties"]["user"]["properties"]["address"]["additionalProperties"] is False + + def test_array_items_object_adds_additional_properties_false(self): + output_format = { + "type": "json_schema", + "schema": { + "type": "object", + "properties": { + "items": { + "type": "array", + "items": { + "type": "object", + "properties": {"id": {"type": "integer"}}, + }, + } + }, + }, + } + result = self.adapter.translate_anthropic_output_format_to_openai(output_format) + assert result is not None + schema = result["json_schema"]["schema"] + assert schema["additionalProperties"] is False + assert schema["properties"]["items"]["items"]["additionalProperties"] is False + + def test_does_not_mutate_original_schema(self): + original_schema = { + "type": "object", + "properties": {"name": {"type": "string"}}, + } + output_format = {"type": "json_schema", "schema": original_schema} + self.adapter.translate_anthropic_output_format_to_openai(output_format) + assert "additionalProperties" not in original_schema + + def test_defs_adds_additional_properties_false(self): + output_format = { + "type": "json_schema", + "schema": { + "type": "object", + "properties": {"ref": {"$ref": "#/$defs/Item"}}, + "$defs": { + "Item": { + "type": "object", + "properties": {"value": {"type": "string"}}, + } + }, + }, + } + result = self.adapter.translate_anthropic_output_format_to_openai(output_format) + assert result is not None + schema = result["json_schema"]["schema"] + assert schema["$defs"]["Item"]["additionalProperties"] is False + + def test_invalid_output_format_returns_none(self): + assert self.adapter.translate_anthropic_output_format_to_openai("invalid") is None + assert self.adapter.translate_anthropic_output_format_to_openai({"type": "text"}) is None + assert self.adapter.translate_anthropic_output_format_to_openai({"type": "json_schema"}) is None From 6f4b4d3c42c73641cb08f78a9587c8328c23afef Mon Sep 17 00:00:00 2001 From: Chesars Date: Wed, 18 Mar 2026 22:33:01 -0300 Subject: [PATCH 02/42] feat(gemini): support context circulation for server-side tool combination Enables Gemini 3+ models to combine built-in tools (Google Search, etc.) with custom functions via `include_server_side_tool_invocations=True`. Server-side invocations are surfaced in provider_specific_fields and automatically re-injected on subsequent turns for multi-turn coherence. Closes #24047 --- docs/my-website/docs/providers/gemini.md | 108 +++++++- .../llms/vertex_ai/gemini/transformation.py | 40 +++ .../vertex_and_google_ai_studio_gemini.py | 78 ++++++ litellm/types/llms/vertex_ai.py | 3 +- .../gemini/test_context_circulation.py | 234 ++++++++++++++++++ 5 files changed, 461 insertions(+), 2 deletions(-) create mode 100644 tests/test_litellm/llms/vertex_ai/gemini/test_context_circulation.py diff --git a/docs/my-website/docs/providers/gemini.md b/docs/my-website/docs/providers/gemini.md index 0aaf3d5ae81..c8c9114ea87 100644 --- a/docs/my-website/docs/providers/gemini.md +++ b/docs/my-website/docs/providers/gemini.md @@ -54,6 +54,7 @@ response = completion( - stream - tools - tool_choice +- include_server_side_tool_invocations - functions - response_format - n @@ -856,7 +857,112 @@ curl -X POST 'http://0.0.0.0:4000/chat/completions' \ -### URL Context +### Context Circulation (Server-Side Tool Combination) + +Context circulation allows Gemini 3+ models to combine **built-in tools** (like Google Search) with **your custom functions** in the same request. Without it, Gemini returns an error if you try to use both. + +When enabled, Gemini can execute Google Search server-side, use those results to decide whether to call your custom functions, and return the full chain of reasoning. + +**How it works:** +1. You pass `include_server_side_tool_invocations=True` along with both Google Search and your function tools +2. Gemini executes server-side tools internally and returns `toolCall`/`toolResponse` parts alongside any `functionCall` parts +3. LiteLLM extracts the server-side invocations into `provider_specific_fields["server_side_tool_invocations"]` +4. On subsequent turns, include the full assistant message in your conversation history — LiteLLM re-injects the server-side parts automatically + + + + +```python +from litellm import completion + +response = completion( + model="gemini/gemini-3-flash-preview", + messages=[{"role": "user", "content": "What's the weather in Buenos Aires? If it's raining, schedule a meeting."}], + tools=[ + {"type": "web_search_preview"}, # Google Search (server-side) + { + "type": "function", + "function": { + "name": "schedule_meeting", + "description": "Schedule a meeting", + "parameters": { + "type": "object", + "properties": {"reason": {"type": "string"}}, + "required": ["reason"], + }, + }, + }, + ], + include_server_side_tool_invocations=True, +) + +msg = response.choices[0].message + +# Server-side tool results are in provider_specific_fields +psf = msg.provider_specific_fields or {} +for invocation in psf.get("server_side_tool_invocations", []): + print(invocation["tool_type"]) # e.g. "GOOGLE_SEARCH_WEB" + print(invocation["id"]) + print(invocation["args"]) # e.g. {"queries": ["weather Buenos Aires"]} + print(invocation["response"]) # Search results from Google + +# For multi-turn: just append the full message to history +messages.append(msg) +messages.append({"role": "user", "content": "Thanks!"}) +# LiteLLM automatically re-injects the server-side parts + thought signatures +response2 = completion( + model="gemini/gemini-3-flash-preview", + messages=messages, + tools=tools, + include_server_side_tool_invocations=True, +) +``` + + + + +1. Setup config.yaml +```yaml +model_list: + - model_name: gemini-3-flash + litellm_params: + model: gemini/gemini-3-flash-preview + api_key: os.environ/GEMINI_API_KEY +``` + +2. Start Proxy +```bash +$ litellm --config /path/to/config.yaml +``` + +3. Make Request +```bash +curl -X POST 'http://0.0.0.0:4000/chat/completions' \ +-H 'Content-Type: application/json' \ +-H 'Authorization: Bearer sk-1234' \ +-d '{ + "model": "gemini-3-flash", + "messages": [{"role": "user", "content": "What is the weather in Buenos Aires?"}], + "tools": [ + {"type": "web_search_preview"}, + {"type": "function", "function": {"name": "schedule_meeting", "description": "Schedule a meeting", "parameters": {"type": "object", "properties": {"reason": {"type": "string"}}}}} + ], + "include_server_side_tool_invocations": true +}' +``` + + + + +:::info + +- Context circulation requires **Gemini 3+** models +- Server-side tool invocations (`toolCall`/`toolResponse`) are **not** included in `tool_calls` — they are in `provider_specific_fields["server_side_tool_invocations"]` because they were already executed by Google, not by your code +- `thought_signatures` are automatically preserved alongside server-side invocations for multi-turn coherence + +::: + +### URL Context diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index d7b96b4db7b..f6310778c71 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -540,6 +540,39 @@ def _gemini_convert_messages_with_history( # noqa: PLR0915 assistant_content.append(gemini_tool_call_part) last_message_with_tool_calls = assistant_msg + ## HANDLE SERVER-SIDE TOOL INVOCATIONS (context circulation) + _psf = assistant_msg.get("provider_specific_fields") + if isinstance(_psf, dict): + _ss_invocations = _psf.get("server_side_tool_invocations") + if isinstance(_ss_invocations, list): + for invocation in _ss_invocations: + # Re-inject toolCall part + tc_part: Dict[str, Any] = { + "toolCall": { + "toolType": invocation.get("tool_type"), + "id": invocation.get("id"), + "args": invocation.get("args"), + } + } + if "thought_signature" in invocation: + tc_part["thoughtSignature"] = invocation["thought_signature"] + assistant_content.append(tc_part) # type: ignore + + # Re-inject toolResponse part if response is present + if "response" in invocation: + tr_dict: Dict[str, Any] = { + "id": invocation.get("id"), + "response": invocation.get("response"), + } + if invocation.get("tool_type"): + tr_dict["toolType"] = invocation["tool_type"] + tr_part: Dict[str, Any] = { + "toolResponse": tr_dict + } + if "thought_signature" in invocation: + tr_part["thoughtSignature"] = invocation["thought_signature"] + assistant_content.append(tr_part) # type: ignore + msg_i += 1 if assistant_content: @@ -666,6 +699,9 @@ def _transform_request_body( # noqa: PLR0915 ) tools: Optional[Tools] = optional_params.pop("tools", None) tool_choice: Optional[ToolConfig] = optional_params.pop("tool_choice", None) + include_server_side_tool_invocations: bool = optional_params.pop( + "include_server_side_tool_invocations", False + ) safety_settings: Optional[List[SafetSettingsConfig]] = optional_params.pop( "safety_settings", None ) # type: ignore @@ -715,6 +751,10 @@ def _transform_request_body( # noqa: PLR0915 data["tools"] = tools if tool_choice is not None: data["toolConfig"] = tool_choice + if include_server_side_tool_invocations: + if "toolConfig" not in data: + data["toolConfig"] = {} + data["toolConfig"]["includeServerSideToolInvocations"] = True if safety_settings is not None: data["safetySettings"] = safety_settings if generation_config is not None and len(generation_config) > 0: diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index 3f1bccaccfc..3555d3c719e 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -316,6 +316,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): "audio", "parallel_tool_calls", "web_search_options", + "include_server_side_tool_invocations", ] # Add penalty parameters only for non-preview models @@ -1119,6 +1120,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): optional_params = self._add_tools_to_optional_params( optional_params, [_tools] ) + elif param == "include_server_side_tool_invocations" and value is True: + optional_params["include_server_side_tool_invocations"] = True if litellm.vertex_ai_safety_settings is not None: optional_params["safety_settings"] = litellm.vertex_ai_safety_settings @@ -1360,6 +1363,67 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): signatures.append(signature) return signatures if signatures else None + @staticmethod + def _extract_server_side_tool_invocations( + parts: List[HttpxPartType], + ) -> Optional[List[Dict[str, Any]]]: + """Extract server-side tool invocations (toolCall/toolResponse) from parts. + + These are returned by Gemini when context circulation is enabled + (includeServerSideToolInvocations=true). They represent tools executed + server-side (e.g. Google Search) and must be circulated back in + subsequent turns for multi-turn coherence. + + Returns: + List of server-side invocation dicts if any found, None otherwise. + """ + invocations: List[Dict[str, Any]] = [] + # Index toolCalls by id so we can pair them with responses + tool_calls_by_id: Dict[str, Dict[str, Any]] = {} + tool_responses_by_id: Dict[str, Dict[str, Any]] = {} + + for part in parts: + if "toolCall" in part: + tc = part["toolCall"] + entry: Dict[str, Any] = { + "tool_type": tc.get("toolType"), + "id": tc.get("id"), + "args": tc.get("args"), + } + signature = part.get("thoughtSignature") + if signature is not None: + entry["thought_signature"] = signature + tool_calls_by_id[tc.get("id", "")] = entry + + elif "toolResponse" in part: + tr = part["toolResponse"] + entry = { + "id": tr.get("id"), + "tool_type": tr.get("toolType"), + "response": tr.get("response"), + } + signature = part.get("thoughtSignature") + if signature is not None: + entry["thought_signature"] = signature + tool_responses_by_id[tr.get("id", "")] = entry + + # Merge calls with their responses + for call_id, call_entry in tool_calls_by_id.items(): + merged = dict(call_entry) + resp = tool_responses_by_id.pop(call_id, None) + if resp is not None: + merged["response"] = resp.get("response") + # Keep response signature if call didn't have one + if "thought_signature" not in merged and "thought_signature" in resp: + merged["thought_signature"] = resp["thought_signature"] + invocations.append(merged) + + # Any orphan responses (shouldn't happen, but be safe) + for resp_id, resp_entry in tool_responses_by_id.items(): + invocations.append(resp_entry) + + return invocations if invocations else None + def _extract_image_response_from_parts( self, parts: List[HttpxPartType] ) -> Optional[List[ImageURLListItem]]: @@ -2018,6 +2082,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): thinking_blocks: Optional[List[ChatCompletionThinkingBlock]] = None reasoning_content: Optional[str] = None thought_signatures: Optional[Any] = None + server_side_tool_invocations: Optional[List[Dict[str, Any]]] = None for idx, candidate in enumerate(_candidates): if "content" not in candidate: @@ -2068,6 +2133,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): ) ) + # Extract server-side tool invocations (context circulation) + server_side_tool_invocations = ( + VertexGeminiConfig._extract_server_side_tool_invocations( + parts=candidate["content"]["parts"] + ) + ) + if audio_response is not None: cast(Dict[str, Any], chat_completion_message)[ "audio" @@ -2139,6 +2211,12 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): chat_completion_message["provider_specific_fields"] = {} chat_completion_message["provider_specific_fields"]["thought_signatures"] = thought_signatures # type: ignore + # Store server-side tool invocations in provider_specific_fields + if server_side_tool_invocations is not None: + if "provider_specific_fields" not in chat_completion_message: + chat_completion_message["provider_specific_fields"] = {} + chat_completion_message["provider_specific_fields"]["server_side_tool_invocations"] = server_side_tool_invocations # type: ignore + if isinstance(model_response, ModelResponseStream): choice = VertexGeminiConfig._create_streaming_choice( chat_completion_message=chat_completion_message, diff --git a/litellm/types/llms/vertex_ai.py b/litellm/types/llms/vertex_ai.py index 201854369f1..66c6ca436cf 100644 --- a/litellm/types/llms/vertex_ai.py +++ b/litellm/types/llms/vertex_ai.py @@ -244,8 +244,9 @@ class Tools(TypedDict, total=False): retrieval: Retrieval -class ToolConfig(TypedDict): +class ToolConfig(TypedDict, total=False): functionCallingConfig: FunctionCallingConfig + includeServerSideToolInvocations: bool class TTL(TypedDict, total=False): diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_context_circulation.py b/tests/test_litellm/llms/vertex_ai/gemini/test_context_circulation.py new file mode 100644 index 00000000000..c3038840d81 --- /dev/null +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_context_circulation.py @@ -0,0 +1,234 @@ +""" +Tests for Gemini context circulation (server-side tool invocations). + +When includeServerSideToolInvocations=true is set, Gemini returns toolCall/toolResponse +parts for server-side tools (e.g. Google Search). These must be: +1. Extracted from the response into provider_specific_fields["server_side_tool_invocations"] +2. Re-injected as raw toolCall/toolResponse parts when converting messages back to Gemini format +3. The includeServerSideToolInvocations flag must be passed through to toolConfig +""" + +import json +from typing import Any, Dict, List +from unittest.mock import MagicMock + +import pytest + +from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + VertexGeminiConfig, +) +from litellm.llms.vertex_ai.gemini.transformation import ( + _gemini_convert_messages_with_history, +) +from litellm.types.llms.vertex_ai import HttpxPartType + + +# --- Response extraction tests --- + + +class TestExtractServerSideToolInvocations: + """Test _extract_server_side_tool_invocations from response parts.""" + + def test_extracts_tool_call_and_response(self): + """Basic case: one toolCall + one toolResponse with same id.""" + parts: List[HttpxPartType] = [ + { + "thoughtSignature": "sig_call_1", + "toolCall": { + "toolType": "GOOGLE_SEARCH_WEB", + "id": "abc123", + "args": {"queries": ["weather Buenos Aires"]}, + }, + }, + { + "thoughtSignature": "sig_resp_1", + "toolResponse": { + "toolType": "GOOGLE_SEARCH_WEB", + "id": "abc123", + "response": {"weather": "Sunny, 20°C"}, + }, + }, + { + "text": "The weather in Buenos Aires is sunny.", + "thoughtSignature": "sig_text", + }, + ] + + result = VertexGeminiConfig._extract_server_side_tool_invocations(parts) + + assert result is not None + assert len(result) == 1 + assert result[0]["tool_type"] == "GOOGLE_SEARCH_WEB" + assert result[0]["id"] == "abc123" + assert result[0]["args"] == {"queries": ["weather Buenos Aires"]} + assert result[0]["response"] == {"weather": "Sunny, 20°C"} + assert result[0]["thought_signature"] == "sig_call_1" + + def test_returns_none_when_no_server_side_tools(self): + """No toolCall/toolResponse parts → returns None.""" + parts: List[HttpxPartType] = [ + {"text": "Hello world", "thoughtSignature": "sig1"}, + { + "functionCall": { + "name": "get_weather", + "args": {"location": "Paris"}, + }, + "thoughtSignature": "sig2", + }, + ] + + result = VertexGeminiConfig._extract_server_side_tool_invocations(parts) + assert result is None + + def test_multiple_server_side_invocations(self): + """Multiple toolCall/toolResponse pairs.""" + parts: List[HttpxPartType] = [ + { + "toolCall": { + "toolType": "GOOGLE_SEARCH_WEB", + "id": "search1", + "args": {"queries": ["query1"]}, + }, + "thoughtSignature": "sig1", + }, + { + "toolResponse": {"toolType": "GOOGLE_SEARCH_WEB", "id": "search1", "response": "result1"}, + "thoughtSignature": "sig2", + }, + { + "toolCall": { + "toolType": "GOOGLE_SEARCH_WEB", + "id": "search2", + "args": {"queries": ["query2"]}, + }, + "thoughtSignature": "sig3", + }, + { + "toolResponse": {"toolType": "GOOGLE_SEARCH_WEB", "id": "search2", "response": "result2"}, + "thoughtSignature": "sig4", + }, + ] + + result = VertexGeminiConfig._extract_server_side_tool_invocations(parts) + + assert result is not None + assert len(result) == 2 + assert result[0]["id"] == "search1" + assert result[0]["response"] == "result1" + assert result[1]["id"] == "search2" + assert result[1]["response"] == "result2" + + def test_tool_call_without_response(self): + """toolCall without matching toolResponse is still captured.""" + parts: List[HttpxPartType] = [ + { + "toolCall": { + "toolType": "CODE_EXECUTION", + "id": "exec1", + "args": {"code": "print('hello')"}, + }, + }, + ] + + result = VertexGeminiConfig._extract_server_side_tool_invocations(parts) + + assert result is not None + assert len(result) == 1 + assert result[0]["id"] == "exec1" + assert "response" not in result[0] + + +# --- Input re-injection tests --- + + +class TestReInjectServerSideToolInvocations: + """Test that server_side_tool_invocations are re-injected into Gemini parts.""" + + def test_roundtrip_single_invocation(self): + """Server-side invocations from assistant message are converted back to Gemini parts.""" + messages = [ + {"role": "user", "content": "What's the weather?"}, + { + "role": "assistant", + "content": "It's sunny in Buenos Aires.", + "provider_specific_fields": { + "server_side_tool_invocations": [ + { + "tool_type": "GOOGLE_SEARCH_WEB", + "id": "abc123", + "args": {"queries": ["weather Buenos Aires"]}, + "response": {"weather": "Sunny, 20°C"}, + "thought_signature": "sig_abc", + } + ] + }, + }, + {"role": "user", "content": "Thanks!"}, + ] + + contents = _gemini_convert_messages_with_history(messages) + + # Find the model turn + model_turn = [c for c in contents if c["role"] == "model"] + assert len(model_turn) == 1 + + parts = model_turn[0]["parts"] + # Should have: text part + toolCall part + toolResponse part + tool_call_parts = [p for p in parts if "toolCall" in p] + tool_response_parts = [p for p in parts if "toolResponse" in p] + + assert len(tool_call_parts) == 1 + assert tool_call_parts[0]["toolCall"]["toolType"] == "GOOGLE_SEARCH_WEB" + assert tool_call_parts[0]["toolCall"]["id"] == "abc123" + assert tool_call_parts[0]["toolCall"]["args"] == {"queries": ["weather Buenos Aires"]} + assert tool_call_parts[0]["thoughtSignature"] == "sig_abc" + + assert len(tool_response_parts) == 1 + assert tool_response_parts[0]["toolResponse"]["id"] == "abc123" + assert tool_response_parts[0]["toolResponse"]["toolType"] == "GOOGLE_SEARCH_WEB" + assert tool_response_parts[0]["toolResponse"]["response"] == {"weather": "Sunny, 20°C"} + + def test_no_invocations_no_extra_parts(self): + """Without server_side_tool_invocations, no extra parts are added.""" + messages = [ + {"role": "user", "content": "Hello"}, + {"role": "assistant", "content": "Hi there!"}, + {"role": "user", "content": "Bye"}, + ] + + contents = _gemini_convert_messages_with_history(messages) + model_turn = [c for c in contents if c["role"] == "model"] + assert len(model_turn) == 1 + + parts = model_turn[0]["parts"] + assert len(parts) == 1 + assert "text" in parts[0] + assert "toolCall" not in parts[0] + + +# --- toolConfig flag tests --- + + +class TestIncludeServerSideToolInvocationsConfig: + """Test that the flag is passed through to toolConfig.""" + + def test_flag_added_to_tool_config(self): + """include_server_side_tool_invocations=True should be mapped to optional_params.""" + config = VertexGeminiConfig() + non_default_params = {"include_server_side_tool_invocations": True} + optional_params: Dict[str, Any] = {} + + result = config.map_openai_params( + non_default_params=non_default_params, + optional_params=optional_params, + model="gemini-3-flash-preview", + drop_params=False, + ) + + assert result["include_server_side_tool_invocations"] is True + + def test_flag_in_supported_params(self): + """include_server_side_tool_invocations should be in supported params.""" + config = VertexGeminiConfig() + supported = config.get_supported_openai_params(model="gemini-3-flash-preview") + assert "include_server_side_tool_invocations" in supported From 286b8d14604c9ad5f736038fb5bb6146ecaed7bf Mon Sep 17 00:00:00 2001 From: Chesars Date: Wed, 18 Mar 2026 23:03:20 -0300 Subject: [PATCH 03/42] fix: also populate required for all properties in strict mode OpenAI strict mode requires both additionalProperties:false AND all property keys in required. Without required, OpenAI rejects the schema even with additionalProperties:false set. --- .../adapters/transformation.py | 7 +++-- ...al_pass_through_adapters_transformation.py | 26 +++++++++++++++++++ 2 files changed, 31 insertions(+), 2 deletions(-) diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py index 47c9a223f85..3ee6c5f0d2a 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py @@ -843,15 +843,18 @@ class LiteLLMAnthropicMessagesAdapter: @staticmethod def _add_additional_properties_false(schema: dict) -> None: """ - Recursively add 'additionalProperties': false to all object schemas. + Recursively ensure object schemas comply with OpenAI strict mode. - OpenAI's strict mode requires this at every object nesting level. + OpenAI's strict mode requires: + 1. 'additionalProperties': false at every object nesting level + 2. All property keys listed in 'required' """ if not isinstance(schema, dict): return if schema.get("type") == "object" and "properties" in schema: schema["additionalProperties"] = False + schema["required"] = list(schema["properties"].keys()) for prop in schema["properties"].values(): LiteLLMAnthropicMessagesAdapter._add_additional_properties_false(prop) diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py index b0442ebf0eb..ae970e1ff06 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py @@ -2004,6 +2004,7 @@ class TestTranslateAnthropicOutputFormatToOpenAI: assert result is not None schema = result["json_schema"]["schema"] assert schema["additionalProperties"] is False + assert schema["required"] == ["name"] def test_nested_objects_adds_additional_properties_false(self): output_format = { @@ -2028,8 +2029,11 @@ class TestTranslateAnthropicOutputFormatToOpenAI: assert result is not None schema = result["json_schema"]["schema"] assert schema["additionalProperties"] is False + assert schema["required"] == ["user"] assert schema["properties"]["user"]["additionalProperties"] is False + assert schema["properties"]["user"]["required"] == ["name", "address"] assert schema["properties"]["user"]["properties"]["address"]["additionalProperties"] is False + assert schema["properties"]["user"]["properties"]["address"]["required"] == ["city"] def test_array_items_object_adds_additional_properties_false(self): output_format = { @@ -2061,6 +2065,7 @@ class TestTranslateAnthropicOutputFormatToOpenAI: output_format = {"type": "json_schema", "schema": original_schema} self.adapter.translate_anthropic_output_format_to_openai(output_format) assert "additionalProperties" not in original_schema + assert "required" not in original_schema def test_defs_adds_additional_properties_false(self): output_format = { @@ -2080,6 +2085,27 @@ class TestTranslateAnthropicOutputFormatToOpenAI: assert result is not None schema = result["json_schema"]["schema"] assert schema["$defs"]["Item"]["additionalProperties"] is False + assert schema["$defs"]["Item"]["required"] == ["value"] + + def test_incomplete_required_gets_completed(self): + """OpenAI strict mode requires ALL properties in required.""" + output_format = { + "type": "json_schema", + "schema": { + "type": "object", + "properties": { + "name": {"type": "string"}, + "age": {"type": "integer"}, + "email": {"type": "string"}, + }, + "required": ["name"], # only 1 of 3 + }, + } + result = self.adapter.translate_anthropic_output_format_to_openai(output_format) + assert result is not None + schema = result["json_schema"]["schema"] + assert schema["additionalProperties"] is False + assert sorted(schema["required"]) == ["age", "email", "name"] def test_invalid_output_format_returns_none(self): assert self.adapter.translate_anthropic_output_format_to_openai("invalid") is None From 60c234270a4b42cdf20888ac353071b5eae5a378 Mon Sep 17 00:00:00 2001 From: Chesars Date: Wed, 18 Mar 2026 23:07:56 -0300 Subject: [PATCH 04/42] feat(bedrock): support cache_control_injection_points for tool_config location Add support for {"location": "tool_config"} in cache_control_injection_points, which appends a cachePoint block to the Bedrock Converse toolConfig.tools array. This enables prompt caching of tool definitions on Bedrock Claude models. Also update the cache control hook to pass through non-message injection points to provider-specific handling instead of silently dropping them. Fixes #21969 --- .../anthropic_cache_control_hook.py | 9 +- .../bedrock/chat/converse_transformation.py | 10 ++ .../chat/test_converse_transformation.py | 97 +++++++++++++++++++ 3 files changed, 115 insertions(+), 1 deletion(-) diff --git a/litellm/integrations/anthropic_cache_control_hook.py b/litellm/integrations/anthropic_cache_control_hook.py index 8e4d40c460e..0e99537d5db 100644 --- a/litellm/integrations/anthropic_cache_control_hook.py +++ b/litellm/integrations/anthropic_cache_control_hook.py @@ -60,13 +60,20 @@ class AnthropicCacheControlHook(CustomPromptManagement): # Create a deep copy of messages to avoid modifying the original list processed_messages = copy.deepcopy(messages) - # Process message-level cache controls + # Separate message-level and non-message-level injection points + remaining_points = [] for point in injection_points: if point.get("location") == "message": point = cast(CacheControlMessageInjectionPoint, point) processed_messages = self._process_message_injection( point=point, messages=processed_messages ) + else: + remaining_points.append(point) + + # Pass through non-message injection points for provider-specific handling + if remaining_points: + non_default_params["cache_control_injection_points"] = remaining_points return model, processed_messages, non_default_params diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index 229457a73b4..dd8b1b0a69f 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -1446,6 +1446,16 @@ class AmazonConverseConfig(BaseConfig): original_tools, model, headers, additional_request_params ) + # Append cachePoint to tools if cache_control_injection_points has tool_config + cache_injection_points = additional_request_params.pop( + "cache_control_injection_points", None + ) + if cache_injection_points and len(bedrock_tools) > 0: + for point in cache_injection_points: + if point.get("location") == "tool_config": + bedrock_tools.append({"cachePoint": {"type": "default"}}) + break + bedrock_tool_config: Optional[ToolConfigBlock] = None if len(bedrock_tools) > 0: tool_choice_values: ToolChoiceValuesBlock = inference_params.pop( diff --git a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py index a305009659c..e9aaa97a421 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py @@ -3803,3 +3803,100 @@ def test_streaming_without_json_mode_passes_all_tools(): assert tool_use_delta is not None assert tool_use_delta["function"]["arguments"] == '{"data": 1}' + +def test_cache_control_injection_tool_config(): + """Test that cache_control_injection_points with location=tool_config appends cachePoint to tools.""" + config = AmazonConverseConfig() + messages = [ + {"role": "user", "content": "What is the weather?"}, + ] + optional_params = { + "tools": [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get weather for a location", + "parameters": { + "type": "object", + "properties": { + "location": {"type": "string"}, + }, + "required": ["location"], + }, + }, + } + ], + "cache_control_injection_points": [ + {"location": "tool_config"}, + ], + } + result = config._transform_request( + model="anthropic.claude-3-5-haiku-20241022-v1:0", + messages=messages, + optional_params=optional_params, + litellm_params={}, + ) + tool_config = result["toolConfig"] + tools = tool_config["tools"] + # Last element should be a cachePoint block + assert tools[-1] == {"cachePoint": {"type": "default"}} + # First element should be the actual tool + assert "toolSpec" in tools[0] + + +def test_cache_control_injection_tool_config_no_tools(): + """Test that tool_config injection is ignored when no tools are provided.""" + config = AmazonConverseConfig() + messages = [ + {"role": "user", "content": "Hello"}, + ] + optional_params = { + "cache_control_injection_points": [ + {"location": "tool_config"}, + ], + } + result = config._transform_request( + model="anthropic.claude-3-5-haiku-20241022-v1:0", + messages=messages, + optional_params=optional_params, + litellm_params={}, + ) + assert "toolConfig" not in result + + +def test_cache_control_injection_tool_config_not_added_without_injection_point(): + """Test that cachePoint is NOT appended when cache_control_injection_points doesn't include tool_config.""" + config = AmazonConverseConfig() + messages = [ + {"role": "user", "content": "What is the weather?"}, + ] + optional_params = { + "tools": [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get weather", + "parameters": { + "type": "object", + "properties": {"location": {"type": "string"}}, + "required": ["location"], + }, + }, + } + ], + "cache_control_injection_points": [ + {"location": "message", "role": "system"}, + ], + } + result = config._transform_request( + model="anthropic.claude-3-5-haiku-20241022-v1:0", + messages=messages, + optional_params=optional_params, + litellm_params={}, + ) + tools = result["toolConfig"]["tools"] + # No cachePoint should be appended + assert all("cachePoint" not in tool for tool in tools) + From cac685014ff2ad795c9d21a9e42475e46fdcd4b5 Mon Sep 17 00:00:00 2001 From: Ephrim Stanley Date: Thu, 19 Mar 2026 01:30:18 -0400 Subject: [PATCH 05/42] feat: add proxy-wide default tpm/rpm limits per deployment Adds `default_api_key_tpm_limit` and `default_api_key_rpm_limit` to `GenericLiteLLMParams` so operators can set per-deployment rate limit defaults in config.yaml. When a key has no model-specific tpm/rpm limit configured, the proxy falls back to these deployment defaults (Case 2 in spec). Key-level limits always take priority (Case 1). - Extends `get_key_model_tpm_limit` / `get_key_model_rpm_limit` with a `model_name` param and a priority-4 deployment-default fallback - Passes `model_name=requested_model` in the parallel request limiter so the fallback is triggered at enforcement time - Adds `"limit"` to `SensitiveDataMasker` non-sensitive overrides so `*_limit` fields are not masked in `/model/info` responses - Adds 17 unit tests covering both spec cases and the `/model/info` path Co-Authored-By: Claude (claude-sonnet-4-6) --- .../sensitive_data_masker.py | 4 +- litellm/proxy/auth/auth_utils.py | 64 ++++++- .../hooks/parallel_request_limiter_v3.py | 8 +- litellm/types/router.py | 5 + .../proxy/auth/test_auth_utils.py | 131 +++++++++++++- .../proxy/test_model_info_default_limits.py | 163 ++++++++++++++++++ 6 files changed, 369 insertions(+), 6 deletions(-) create mode 100644 tests/test_litellm/proxy/test_model_info_default_limits.py diff --git a/litellm/litellm_core_utils/sensitive_data_masker.py b/litellm/litellm_core_utils/sensitive_data_masker.py index 663c3fac801..f22cfa11a3d 100644 --- a/litellm/litellm_core_utils/sensitive_data_masker.py +++ b/litellm/litellm_core_utils/sensitive_data_masker.py @@ -30,7 +30,9 @@ class SensitiveDataMasker: # If any key segment matches one of these, the key is not considered sensitive # even if it also matches a sensitive pattern. For example, "input_cost_per_token" # contains "token" but "cost" overrides that — it's a pricing field, not a secret. - self.non_sensitive_overrides = non_sensitive_overrides or {"cost"} + # Similarly, "*_limit" fields (tpm_limit, rpm_limit, etc.) are rate/budget caps, + # not credentials, even though their names may contain "key" (e.g. default_api_key_tpm_limit). + self.non_sensitive_overrides = non_sensitive_overrides or {"cost", "limit"} self.visible_prefix = visible_prefix self.visible_suffix = visible_suffix diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index a03e1fb94c1..235b217610a 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -539,8 +539,49 @@ def bytes_to_mb(bytes_value: int): # helpers used by parallel request limiter to handle model rpm/tpm limits for a given api key +def _get_deployment_default_rpm_limit(model_name: str) -> Optional[int]: + """ + Return the default_api_key_rpm_limit configured on the deployment for model_name, + or None if not set. + """ + from litellm.proxy.proxy_server import llm_router + + if llm_router is None: + return None + deployments = llm_router.get_model_list(model_name=model_name) + if not deployments: + return None + for deployment in deployments: + litellm_params = deployment.get("litellm_params", {}) + limit = litellm_params.get("default_api_key_rpm_limit") + if limit is not None: + return int(limit) + return None + + +def _get_deployment_default_tpm_limit(model_name: str) -> Optional[int]: + """ + Return the default_api_key_tpm_limit configured on the deployment for model_name, + or None if not set. + """ + from litellm.proxy.proxy_server import llm_router + + if llm_router is None: + return None + deployments = llm_router.get_model_list(model_name=model_name) + if not deployments: + return None + for deployment in deployments: + litellm_params = deployment.get("litellm_params", {}) + limit = litellm_params.get("default_api_key_tpm_limit") + if limit is not None: + return int(limit) + return None + + def get_key_model_rpm_limit( user_api_key_dict: UserAPIKeyAuth, + model_name: Optional[str] = None, ) -> Optional[Dict[str, int]]: """ Get the model rpm limit for a given api key. @@ -549,6 +590,7 @@ def get_key_model_rpm_limit( 1. Key metadata (model_rpm_limit) 2. Key model_max_budget (rpm_limit per model) 3. Team metadata (model_rpm_limit) + 4. Deployment default_api_key_rpm_limit (when model_name is provided) """ # 1. Check key metadata first (takes priority) if user_api_key_dict.metadata: @@ -567,13 +609,22 @@ def get_key_model_rpm_limit( # 3. Fallback to team metadata if user_api_key_dict.team_metadata: - return user_api_key_dict.team_metadata.get("model_rpm_limit") + team_limit = user_api_key_dict.team_metadata.get("model_rpm_limit") + if team_limit: + return team_limit + + # 4. Fallback to deployment default_api_key_rpm_limit + if model_name is not None: + default_limit = _get_deployment_default_rpm_limit(model_name) + if default_limit is not None: + return {model_name: default_limit} return None def get_key_model_tpm_limit( user_api_key_dict: UserAPIKeyAuth, + model_name: Optional[str] = None, ) -> Optional[Dict[str, int]]: """ Get the model tpm limit for a given api key. @@ -582,6 +633,7 @@ def get_key_model_tpm_limit( 1. Key metadata (model_tpm_limit) 2. Key model_max_budget (tpm_limit per model) 3. Team metadata (model_tpm_limit) + 4. Deployment default_api_key_tpm_limit (when model_name is provided) """ # 1. Check key metadata first (takes priority) if user_api_key_dict.metadata: @@ -600,7 +652,15 @@ def get_key_model_tpm_limit( # 3. Fallback to team metadata if user_api_key_dict.team_metadata: - return user_api_key_dict.team_metadata.get("model_tpm_limit") + team_limit = user_api_key_dict.team_metadata.get("model_tpm_limit") + if team_limit: + return team_limit + + # 4. Fallback to deployment default_api_key_tpm_limit + if model_name is not None: + default_limit = _get_deployment_default_tpm_limit(model_name) + if default_limit is not None: + return {model_name: default_limit} return None diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 19c8c484b4d..5aaac088dc2 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -687,8 +687,12 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): if not requested_model: return - _tpm_limit_for_key_model = get_key_model_tpm_limit(user_api_key_dict) - _rpm_limit_for_key_model = get_key_model_rpm_limit(user_api_key_dict) + _tpm_limit_for_key_model = get_key_model_tpm_limit( + user_api_key_dict, model_name=requested_model + ) + _rpm_limit_for_key_model = get_key_model_rpm_limit( + user_api_key_dict, model_name=requested_model + ) if _tpm_limit_for_key_model is None and _rpm_limit_for_key_model is None: return diff --git a/litellm/types/router.py b/litellm/types/router.py index e8ff2115ff5..5d28349b5e4 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -188,6 +188,11 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams): max_file_size_mb: Optional[float] = None + # Proxy-wide default rate limits applied to any API key using this deployment + # when the key does not have a model-specific tpm/rpm limit configured. + default_api_key_tpm_limit: Optional[int] = None + default_api_key_rpm_limit: Optional[int] = None + # Deployment budgets max_budget: Optional[float] = None budget_duration: Optional[str] = None diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index 5e42b110aa0..be4db666a04 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -2,7 +2,7 @@ Unit tests for auth_utils functions related to rate limiting and customer ID extraction. """ -from unittest.mock import patch +from unittest.mock import MagicMock, patch from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.auth_utils import ( @@ -315,3 +315,132 @@ def test_get_end_user_id_falls_back_to_deprecated_user_header_name(): result = get_end_user_id_from_request_body(request_body={}, request_headers=headers) assert result == "user-legacy" + + +def _make_deployment_dict(model_name: str, tpm: int = None, rpm: int = None) -> dict: + """Helper to build a minimal deployment dict as returned by router.get_model_list.""" + litellm_params: dict = {"model": model_name} + if tpm is not None: + litellm_params["default_api_key_tpm_limit"] = tpm + if rpm is not None: + litellm_params["default_api_key_rpm_limit"] = rpm + return {"model_name": model_name, "litellm_params": litellm_params} + + +_ROUTER_PATCH = "litellm.proxy.proxy_server.llm_router" + + +class TestDeploymentDefaultRpmLimit: + """Tests for deployment default_api_key_rpm_limit fallback in get_key_model_rpm_limit.""" + + def test_returns_deployment_default_when_key_has_no_limits(self): + """Case 2 from spec: key has no model-specific limits, falls back to deployment default.""" + user_api_key_dict = UserAPIKeyAuth(api_key="sk-123") + mock_router = MagicMock() + mock_router.get_model_list.return_value = [ + _make_deployment_dict("model1", rpm=200) + ] + with patch(_ROUTER_PATCH, mock_router): + result = get_key_model_rpm_limit(user_api_key_dict, model_name="model1") + assert result == {"model1": 200} + + def test_key_model_limit_takes_priority_over_deployment_default(self): + """Case 1 from spec: key model-specific limit wins over deployment default.""" + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-123", + metadata={"model_rpm_limit": {"model1": 10}}, + ) + mock_router = MagicMock() + mock_router.get_model_list.return_value = [ + _make_deployment_dict("model1", rpm=200) + ] + with patch(_ROUTER_PATCH, mock_router): + result = get_key_model_rpm_limit(user_api_key_dict, model_name="model1") + assert result == {"model1": 10} + + def test_returns_none_when_no_deployment_default_and_no_key_limits(self): + """Returns None when neither the key nor the deployment has any rpm limit.""" + user_api_key_dict = UserAPIKeyAuth(api_key="sk-123") + mock_router = MagicMock() + mock_router.get_model_list.return_value = [ + _make_deployment_dict("model1") # no rpm default + ] + with patch(_ROUTER_PATCH, mock_router): + result = get_key_model_rpm_limit(user_api_key_dict, model_name="model1") + assert result is None + + def test_returns_none_without_model_name_even_when_deployment_has_default(self): + """No model_name means deployment fallback is skipped.""" + user_api_key_dict = UserAPIKeyAuth(api_key="sk-123") + mock_router = MagicMock() + mock_router.get_model_list.return_value = [ + _make_deployment_dict("model1", rpm=200) + ] + with patch(_ROUTER_PATCH, mock_router): + result = get_key_model_rpm_limit(user_api_key_dict) + assert result is None + + def test_returns_none_when_llm_router_is_none(self): + """No router means deployment fallback returns None gracefully.""" + user_api_key_dict = UserAPIKeyAuth(api_key="sk-123") + with patch(_ROUTER_PATCH, None): + result = get_key_model_rpm_limit(user_api_key_dict, model_name="model1") + assert result is None + + +class TestDeploymentDefaultTpmLimit: + """Tests for deployment default_api_key_tpm_limit fallback in get_key_model_tpm_limit.""" + + def test_returns_deployment_default_when_key_has_no_limits(self): + """Case 2 from spec: key has no model-specific limits, falls back to deployment default.""" + user_api_key_dict = UserAPIKeyAuth(api_key="sk-123") + mock_router = MagicMock() + mock_router.get_model_list.return_value = [ + _make_deployment_dict("model1", tpm=100) + ] + with patch(_ROUTER_PATCH, mock_router): + result = get_key_model_tpm_limit(user_api_key_dict, model_name="model1") + assert result == {"model1": 100} + + def test_key_model_limit_takes_priority_over_deployment_default(self): + """Case 1 from spec: key model-specific limit wins over deployment default.""" + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-123", + metadata={"model_tpm_limit": {"model1": 20}}, + ) + mock_router = MagicMock() + mock_router.get_model_list.return_value = [ + _make_deployment_dict("model1", tpm=100) + ] + with patch(_ROUTER_PATCH, mock_router): + result = get_key_model_tpm_limit(user_api_key_dict, model_name="model1") + assert result == {"model1": 20} + + def test_returns_none_when_no_deployment_default_and_no_key_limits(self): + """Returns None when neither the key nor the deployment has any tpm limit.""" + user_api_key_dict = UserAPIKeyAuth(api_key="sk-123") + mock_router = MagicMock() + mock_router.get_model_list.return_value = [ + _make_deployment_dict("model1") # no tpm default + ] + with patch(_ROUTER_PATCH, mock_router): + result = get_key_model_tpm_limit(user_api_key_dict, model_name="model1") + assert result is None + + def test_returns_none_without_model_name_even_when_deployment_has_default(self): + """No model_name means deployment fallback is skipped.""" + user_api_key_dict = UserAPIKeyAuth(api_key="sk-123") + mock_router = MagicMock() + mock_router.get_model_list.return_value = [ + _make_deployment_dict("model1", tpm=100) + ] + with patch(_ROUTER_PATCH, mock_router): + result = get_key_model_tpm_limit(user_api_key_dict) + assert result is None + + def test_returns_none_when_llm_router_is_none(self): + """No router means deployment fallback returns None gracefully.""" + user_api_key_dict = UserAPIKeyAuth(api_key="sk-123") + with patch(_ROUTER_PATCH, None): + result = get_key_model_tpm_limit(user_api_key_dict, model_name="model1") + assert result is None diff --git a/tests/test_litellm/proxy/test_model_info_default_limits.py b/tests/test_litellm/proxy/test_model_info_default_limits.py new file mode 100644 index 00000000000..e749c84dfb2 --- /dev/null +++ b/tests/test_litellm/proxy/test_model_info_default_limits.py @@ -0,0 +1,163 @@ +""" +Tests verifying that default_api_key_tpm_limit and default_api_key_rpm_limit set in +litellm_params are returned by the /model/info endpoint. +""" + +from unittest.mock import MagicMock, patch + +import pytest + +from litellm.proxy.proxy_server import _get_proxy_model_info +from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo + + +def _make_deployment( + model_name: str, + default_tpm: int = None, + default_rpm: int = None, +) -> Deployment: + params: dict = {"model": f"openai/{model_name}"} + if default_tpm is not None: + params["default_api_key_tpm_limit"] = default_tpm + if default_rpm is not None: + params["default_api_key_rpm_limit"] = default_rpm + return Deployment( + model_name=model_name, + litellm_params=LiteLLM_Params(**params), + model_info=ModelInfo(), + ) + + +class TestModelInfoDefaultLimitsInResponse: + """ + Verify _get_proxy_model_info (the helper used by the /model/info endpoint) returns + default_api_key_tpm_limit and default_api_key_rpm_limit from litellm_params. + """ + + def test_default_tpm_and_rpm_present_in_model_info_response(self): + """Both defaults should appear in the litellm_params section of the response.""" + deployment = _make_deployment("model1", default_tpm=100, default_rpm=200) + model_dict = deployment.model_dump(exclude_none=True) + + result = _get_proxy_model_info(model=model_dict) + + litellm_params = result["litellm_params"] + assert litellm_params.get("default_api_key_tpm_limit") == 100 + assert litellm_params.get("default_api_key_rpm_limit") == 200 + + def test_default_tpm_only_present_when_only_tpm_configured(self): + """Only the configured default appears; the other stays absent.""" + deployment = _make_deployment("model1", default_tpm=500) + model_dict = deployment.model_dump(exclude_none=True) + + result = _get_proxy_model_info(model=model_dict) + + litellm_params = result["litellm_params"] + assert litellm_params.get("default_api_key_tpm_limit") == 500 + assert "default_api_key_rpm_limit" not in litellm_params + + def test_default_rpm_only_present_when_only_rpm_configured(self): + """Only the configured default appears; the other stays absent.""" + deployment = _make_deployment("model1", default_rpm=300) + model_dict = deployment.model_dump(exclude_none=True) + + result = _get_proxy_model_info(model=model_dict) + + litellm_params = result["litellm_params"] + assert litellm_params.get("default_api_key_rpm_limit") == 300 + assert "default_api_key_tpm_limit" not in litellm_params + + def test_defaults_absent_when_not_configured(self): + """Neither field appears when not set on the deployment.""" + deployment = _make_deployment("model1") + model_dict = deployment.model_dump(exclude_none=True) + + result = _get_proxy_model_info(model=model_dict) + + litellm_params = result["litellm_params"] + assert "default_api_key_tpm_limit" not in litellm_params + assert "default_api_key_rpm_limit" not in litellm_params + + def test_defaults_not_masked_or_stripped_by_sensitive_data_filter(self): + """ + default_api_key_tpm_limit / default_api_key_rpm_limit must not be + treated as sensitive and must survive remove_sensitive_info_from_deployment. + """ + deployment = _make_deployment("model1", default_tpm=100, default_rpm=200) + model_dict = deployment.model_dump(exclude_none=True) + + result = _get_proxy_model_info(model=model_dict) + + # Values should be unchanged integers, not masked strings + assert result["litellm_params"]["default_api_key_tpm_limit"] == 100 + assert result["litellm_params"]["default_api_key_rpm_limit"] == 200 + + +class TestModelInfoEndpointWithRouter: + """ + Integration-style tests simulating the /model/info endpoint reading from the router. + """ + + @pytest.mark.asyncio + async def test_model_info_endpoint_returns_defaults_for_specific_model_id(self): + """ + When litellm_model_id is provided, the endpoint should return the deployment's + default limits in litellm_params. + """ + from litellm.proxy.proxy_server import model_info_v1 + from litellm.proxy._types import UserAPIKeyAuth + + deployment = _make_deployment("model1", default_tpm=100, default_rpm=200) + + mock_router = MagicMock() + mock_router.get_deployment.return_value = deployment + + user_api_key_dict = UserAPIKeyAuth(api_key="sk-test") + + with patch("litellm.proxy.proxy_server.llm_router", mock_router), \ + patch("litellm.proxy.proxy_server.llm_model_list", [{}]), \ + patch("litellm.proxy.proxy_server.user_model", None): + response = await model_info_v1( + user_api_key_dict=user_api_key_dict, + litellm_model_id="some-model-id", + ) + + assert len(response["data"]) == 1 + litellm_params = response["data"][0]["litellm_params"] + assert litellm_params.get("default_api_key_tpm_limit") == 100 + assert litellm_params.get("default_api_key_rpm_limit") == 200 + + @pytest.mark.asyncio + async def test_model_info_endpoint_returns_defaults_in_full_model_list(self): + """ + Without litellm_model_id, the endpoint iterates all models. Each deployment's + default limits should appear in its litellm_params entry. + """ + from litellm.proxy.proxy_server import model_info_v1 + from litellm.proxy._types import UserAPIKeyAuth + + deployment = _make_deployment("model1", default_tpm=100, default_rpm=200) + deployment_dict = deployment.model_dump(exclude_none=True) + + mock_router = MagicMock() + mock_router.get_model_names.return_value = ["model1"] + mock_router.get_model_access_groups.return_value = {} + mock_router.get_model_list.return_value = [deployment_dict] + + user_api_key_dict = UserAPIKeyAuth(api_key="sk-test") + + with patch("litellm.proxy.proxy_server.llm_router", mock_router), \ + patch("litellm.proxy.proxy_server.llm_model_list", [deployment_dict]), \ + patch("litellm.proxy.proxy_server.user_model", None), \ + patch("litellm.proxy.proxy_server.get_key_models", return_value=["model1"]), \ + patch("litellm.proxy.proxy_server.get_team_models", return_value=["model1"]), \ + patch("litellm.proxy.proxy_server.get_complete_model_list", return_value=["model1"]): + response = await model_info_v1( + user_api_key_dict=user_api_key_dict, + litellm_model_id=None, + ) + + assert len(response["data"]) >= 1 + litellm_params = response["data"][0]["litellm_params"] + assert litellm_params.get("default_api_key_tpm_limit") == 100 + assert litellm_params.get("default_api_key_rpm_limit") == 200 From 36dc893770fa69cfcc014e5c535b87d58ca12c89 Mon Sep 17 00:00:00 2001 From: Ephrim Stanley Date: Thu, 19 Mar 2026 01:43:27 -0400 Subject: [PATCH 06/42] fix: address review feedback on default tpm/rpm limits - Use min() across all matching deployments instead of first-wins when resolving default_api_key_tpm/rpm_limit for a model group, so load-balanced setups with different per-deployment limits always apply the most conservative value - Replace the global SensitiveDataMasker non_sensitive_overrides change with a targeted excluded_keys set at the remove_sensitive_info_from_deployment call site, avoiding unintended suppression of other fields - Update the v1 parallel request limiter to pass model_name to get_key_model_tpm/rpm_limit so deployment defaults apply there too - Add 4 tests covering multi-deployment min semantics Co-Authored-By: Claude (claude-sonnet-4-6) --- .../sensitive_data_masker.py | 4 +- litellm/proxy/auth/auth_utils.py | 44 ++++++++++------ .../common_utils/openai_endpoint_utils.py | 11 +++- .../proxy/hooks/parallel_request_limiter.py | 14 ++++-- .../proxy/auth/test_auth_utils.py | 50 +++++++++++++++++++ .../proxy/test_model_info_default_limits.py | 3 ++ 6 files changed, 101 insertions(+), 25 deletions(-) diff --git a/litellm/litellm_core_utils/sensitive_data_masker.py b/litellm/litellm_core_utils/sensitive_data_masker.py index f22cfa11a3d..663c3fac801 100644 --- a/litellm/litellm_core_utils/sensitive_data_masker.py +++ b/litellm/litellm_core_utils/sensitive_data_masker.py @@ -30,9 +30,7 @@ class SensitiveDataMasker: # If any key segment matches one of these, the key is not considered sensitive # even if it also matches a sensitive pattern. For example, "input_cost_per_token" # contains "token" but "cost" overrides that — it's a pricing field, not a secret. - # Similarly, "*_limit" fields (tpm_limit, rpm_limit, etc.) are rate/budget caps, - # not credentials, even though their names may contain "key" (e.g. default_api_key_tpm_limit). - self.non_sensitive_overrides = non_sensitive_overrides or {"cost", "limit"} + self.non_sensitive_overrides = non_sensitive_overrides or {"cost"} self.visible_prefix = visible_prefix self.visible_suffix = visible_suffix diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 235b217610a..ace39c05ffc 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -541,8 +541,13 @@ def bytes_to_mb(bytes_value: int): # helpers used by parallel request limiter to handle model rpm/tpm limits for a given api key def _get_deployment_default_rpm_limit(model_name: str) -> Optional[int]: """ - Return the default_api_key_rpm_limit configured on the deployment for model_name, - or None if not set. + Return the default_api_key_rpm_limit for model_name. + + When multiple deployments share the same model name, returns the minimum + across all deployments that have the field set. This is the safest choice + for load-balanced setups: it ensures no deployment is over-consumed + regardless of which one actually serves a given request. + Returns None if no deployment has the field set. """ from litellm.proxy.proxy_server import llm_router @@ -551,18 +556,24 @@ def _get_deployment_default_rpm_limit(model_name: str) -> Optional[int]: deployments = llm_router.get_model_list(model_name=model_name) if not deployments: return None - for deployment in deployments: - litellm_params = deployment.get("litellm_params", {}) - limit = litellm_params.get("default_api_key_rpm_limit") - if limit is not None: - return int(limit) - return None + limits = [ + int(deployment.get("litellm_params", {}).get("default_api_key_rpm_limit")) + for deployment in deployments + if deployment.get("litellm_params", {}).get("default_api_key_rpm_limit") + is not None + ] + return min(limits) if limits else None def _get_deployment_default_tpm_limit(model_name: str) -> Optional[int]: """ - Return the default_api_key_tpm_limit configured on the deployment for model_name, - or None if not set. + Return the default_api_key_tpm_limit for model_name. + + When multiple deployments share the same model name, returns the minimum + across all deployments that have the field set. This is the safest choice + for load-balanced setups: it ensures no deployment is over-consumed + regardless of which one actually serves a given request. + Returns None if no deployment has the field set. """ from litellm.proxy.proxy_server import llm_router @@ -571,12 +582,13 @@ def _get_deployment_default_tpm_limit(model_name: str) -> Optional[int]: deployments = llm_router.get_model_list(model_name=model_name) if not deployments: return None - for deployment in deployments: - litellm_params = deployment.get("litellm_params", {}) - limit = litellm_params.get("default_api_key_tpm_limit") - if limit is not None: - return int(limit) - return None + limits = [ + int(deployment.get("litellm_params", {}).get("default_api_key_tpm_limit")) + for deployment in deployments + if deployment.get("litellm_params", {}).get("default_api_key_tpm_limit") + is not None + ] + return min(limits) if limits else None def get_key_model_rpm_limit( diff --git a/litellm/proxy/common_utils/openai_endpoint_utils.py b/litellm/proxy/common_utils/openai_endpoint_utils.py index 6df5491f37a..7e5c83500a2 100644 --- a/litellm/proxy/common_utils/openai_endpoint_utils.py +++ b/litellm/proxy/common_utils/openai_endpoint_utils.py @@ -32,8 +32,17 @@ def remove_sensitive_info_from_deployment( deployment_dict["litellm_params"].pop("aws_access_key_id", None) deployment_dict["litellm_params"].pop("aws_secret_access_key", None) + # Rate-limit config fields must never be masked — they are integers, not credentials. + # The field names contain "key" which matches the masker's sensitive pattern, so we + # explicitly exclude them here rather than widening the global non_sensitive_overrides. + _rate_limit_config_keys = { + "default_api_key_tpm_limit", + "default_api_key_rpm_limit", + } + _excluded = (excluded_keys or set()) | _rate_limit_config_keys + deployment_dict["litellm_params"] = SENSITIVE_DATA_MASKER.mask_dict( - deployment_dict["litellm_params"], excluded_keys=excluded_keys + deployment_dict["litellm_params"], excluded_keys=_excluded ) return deployment_dict diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index c7bfc27d6b6..48bf255ac13 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -295,16 +295,20 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): ) # Check if request under RPM/TPM per model for a given API Key + _model = data.get("model", None) if ( - get_key_model_tpm_limit(user_api_key_dict) is not None - or get_key_model_rpm_limit(user_api_key_dict) is not None + get_key_model_tpm_limit(user_api_key_dict, model_name=_model) is not None + or get_key_model_rpm_limit(user_api_key_dict, model_name=_model) is not None ): - _model = data.get("model", None) request_count_api_key = ( f"{api_key}::{_model}::{precise_minute}::request_count" ) - _tpm_limit_for_key_model = get_key_model_tpm_limit(user_api_key_dict) - _rpm_limit_for_key_model = get_key_model_rpm_limit(user_api_key_dict) + _tpm_limit_for_key_model = get_key_model_tpm_limit( + user_api_key_dict, model_name=_model + ) + _rpm_limit_for_key_model = get_key_model_rpm_limit( + user_api_key_dict, model_name=_model + ) tpm_limit_for_model = None rpm_limit_for_model = None diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index be4db666a04..d64f17e70ff 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -387,6 +387,31 @@ class TestDeploymentDefaultRpmLimit: result = get_key_model_rpm_limit(user_api_key_dict, model_name="model1") assert result is None + def test_returns_minimum_across_multiple_deployments(self): + """When multiple deployments share a model name, the minimum rpm limit is used.""" + user_api_key_dict = UserAPIKeyAuth(api_key="sk-123") + mock_router = MagicMock() + mock_router.get_model_list.return_value = [ + _make_deployment_dict("model1", rpm=200), + _make_deployment_dict("model1", rpm=50), + _make_deployment_dict("model1", rpm=150), + ] + with patch(_ROUTER_PATCH, mock_router): + result = get_key_model_rpm_limit(user_api_key_dict, model_name="model1") + assert result == {"model1": 50} + + def test_ignores_deployments_without_default_when_others_have_it(self): + """Deployments missing the field are skipped; min is taken over those that have it.""" + user_api_key_dict = UserAPIKeyAuth(api_key="sk-123") + mock_router = MagicMock() + mock_router.get_model_list.return_value = [ + _make_deployment_dict("model1"), # no rpm default + _make_deployment_dict("model1", rpm=75), + ] + with patch(_ROUTER_PATCH, mock_router): + result = get_key_model_rpm_limit(user_api_key_dict, model_name="model1") + assert result == {"model1": 75} + class TestDeploymentDefaultTpmLimit: """Tests for deployment default_api_key_tpm_limit fallback in get_key_model_tpm_limit.""" @@ -444,3 +469,28 @@ class TestDeploymentDefaultTpmLimit: with patch(_ROUTER_PATCH, None): result = get_key_model_tpm_limit(user_api_key_dict, model_name="model1") assert result is None + + def test_returns_minimum_across_multiple_deployments(self): + """When multiple deployments share a model name, the minimum tpm limit is used.""" + user_api_key_dict = UserAPIKeyAuth(api_key="sk-123") + mock_router = MagicMock() + mock_router.get_model_list.return_value = [ + _make_deployment_dict("model1", tpm=1000), + _make_deployment_dict("model1", tpm=300), + _make_deployment_dict("model1", tpm=700), + ] + with patch(_ROUTER_PATCH, mock_router): + result = get_key_model_tpm_limit(user_api_key_dict, model_name="model1") + assert result == {"model1": 300} + + def test_ignores_deployments_without_default_when_others_have_it(self): + """Deployments missing the field are skipped; min is taken over those that have it.""" + user_api_key_dict = UserAPIKeyAuth(api_key="sk-123") + mock_router = MagicMock() + mock_router.get_model_list.return_value = [ + _make_deployment_dict("model1"), # no tpm default + _make_deployment_dict("model1", tpm=400), + ] + with patch(_ROUTER_PATCH, mock_router): + result = get_key_model_tpm_limit(user_api_key_dict, model_name="model1") + assert result == {"model1": 400} diff --git a/tests/test_litellm/proxy/test_model_info_default_limits.py b/tests/test_litellm/proxy/test_model_info_default_limits.py index e749c84dfb2..907f2390a89 100644 --- a/tests/test_litellm/proxy/test_model_info_default_limits.py +++ b/tests/test_litellm/proxy/test_model_info_default_limits.py @@ -82,6 +82,9 @@ class TestModelInfoDefaultLimitsInResponse: """ default_api_key_tpm_limit / default_api_key_rpm_limit must not be treated as sensitive and must survive remove_sensitive_info_from_deployment. + They contain "key" which normally triggers masking; the call site explicitly + excludes these two fields via excluded_keys rather than widening the global + non_sensitive_overrides. """ deployment = _make_deployment("model1", default_tpm=100, default_rpm=200) model_dict = deployment.model_dump(exclude_none=True) From b90f5207488c71ad22acc700c24999862e0e08e9 Mon Sep 17 00:00:00 2001 From: Ephrim Stanley Date: Thu, 19 Mar 2026 01:57:09 -0400 Subject: [PATCH 07/42] perf: eliminate redundant router lookups in v1 parallel request limiter Compute get_key_model_tpm/rpm_limit once before the guard condition instead of calling each function twice (once to check non-None, once to retrieve). Removes 2 extra llm_router.get_model_list() calls per request when deployment defaults are active. Co-Authored-By: Claude (claude-sonnet-4-6) --- litellm/proxy/hooks/parallel_request_limiter.py | 17 +++++++---------- 1 file changed, 7 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index 48bf255ac13..6e34c3eee13 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -296,19 +296,16 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): # Check if request under RPM/TPM per model for a given API Key _model = data.get("model", None) - if ( - get_key_model_tpm_limit(user_api_key_dict, model_name=_model) is not None - or get_key_model_rpm_limit(user_api_key_dict, model_name=_model) is not None - ): + _tpm_limit_for_key_model = get_key_model_tpm_limit( + user_api_key_dict, model_name=_model + ) + _rpm_limit_for_key_model = get_key_model_rpm_limit( + user_api_key_dict, model_name=_model + ) + if _tpm_limit_for_key_model is not None or _rpm_limit_for_key_model is not None: request_count_api_key = ( f"{api_key}::{_model}::{precise_minute}::request_count" ) - _tpm_limit_for_key_model = get_key_model_tpm_limit( - user_api_key_dict, model_name=_model - ) - _rpm_limit_for_key_model = get_key_model_rpm_limit( - user_api_key_dict, model_name=_model - ) tpm_limit_for_model = None rpm_limit_for_model = None From 48cb4a83435faa458888b412d408eb06cbca694d Mon Sep 17 00:00:00 2001 From: Ephrim Stanley Date: Thu, 19 Mar 2026 01:59:41 -0400 Subject: [PATCH 08/42] fix: update success-event handler to track tokens for deployment-default limits async_log_success_event only updated the per-model cache counter when model_rpm_limit / model_tpm_limit were present in key metadata or model_max_budget was set. For the new deployment-default path (default_api_key_tpm_limit / default_api_key_rpm_limit), none of those conditions held, so current_tpm stayed at zero and tpm enforcement was never applied across multiple requests. Extend the guard condition to also trigger when the model group has a deployment-default tpm or rpm limit, and import the two helpers at module level. Co-Authored-By: Claude (claude-sonnet-4-6) --- litellm/proxy/hooks/parallel_request_limiter.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index 6e34c3eee13..b0046fd035b 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -14,6 +14,8 @@ from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs from litellm.proxy._types import CommonProxyErrors, CurrentItemRateLimit, UserAPIKeyAuth from litellm.proxy.auth.auth_utils import ( + _get_deployment_default_rpm_limit, + _get_deployment_default_tpm_limit, get_key_model_rpm_limit, get_key_model_tpm_limit, ) @@ -546,6 +548,8 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): "model_rpm_limit" in user_api_key_metadata or "model_tpm_limit" in user_api_key_metadata or user_api_key_model_max_budget is not None + or _get_deployment_default_tpm_limit(model_group) is not None + or _get_deployment_default_rpm_limit(model_group) is not None ) ): request_count_api_key = ( From 477c54184bda814745bb06d4fbe4900cfc816bd2 Mon Sep 17 00:00:00 2001 From: Ephrim Stanley Date: Thu, 19 Mar 2026 02:07:50 -0400 Subject: [PATCH 09/42] perf: avoid unconditional router lookups in success handler Replace bare _get_deployment_default_tpm/rpm_limit calls in the async_log_success_event condition with get_key_model_tpm/rpm_limit (model_name=model_group). The higher-level getters short-circuit on key/team metadata hits before ever reaching the router, so requests that don't use deployment defaults incur no extra router lookup. Remove the now-unused bare helper imports. Also fix invalid `int = None` type hints in test helper signatures to `Optional[int] = None`. Co-Authored-By: Claude (claude-sonnet-4-6) --- litellm/proxy/hooks/parallel_request_limiter.py | 12 ++++++++---- tests/test_litellm/proxy/auth/test_auth_utils.py | 3 ++- .../proxy/test_model_info_default_limits.py | 5 +++-- 3 files changed, 13 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index b0046fd035b..49c6436c22f 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -14,8 +14,6 @@ from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs from litellm.proxy._types import CommonProxyErrors, CurrentItemRateLimit, UserAPIKeyAuth from litellm.proxy.auth.auth_utils import ( - _get_deployment_default_rpm_limit, - _get_deployment_default_tpm_limit, get_key_model_rpm_limit, get_key_model_tpm_limit, ) @@ -548,8 +546,14 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): "model_rpm_limit" in user_api_key_metadata or "model_tpm_limit" in user_api_key_metadata or user_api_key_model_max_budget is not None - or _get_deployment_default_tpm_limit(model_group) is not None - or _get_deployment_default_rpm_limit(model_group) is not None + or get_key_model_tpm_limit( + user_api_key_dict, model_name=model_group + ) + is not None + or get_key_model_rpm_limit( + user_api_key_dict, model_name=model_group + ) + is not None ) ): request_count_api_key = ( diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index d64f17e70ff..2058f61cb0f 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -2,6 +2,7 @@ Unit tests for auth_utils functions related to rate limiting and customer ID extraction. """ +from typing import Optional from unittest.mock import MagicMock, patch from litellm.proxy._types import UserAPIKeyAuth @@ -317,7 +318,7 @@ def test_get_end_user_id_falls_back_to_deprecated_user_header_name(): assert result == "user-legacy" -def _make_deployment_dict(model_name: str, tpm: int = None, rpm: int = None) -> dict: +def _make_deployment_dict(model_name: str, tpm: Optional[int] = None, rpm: Optional[int] = None) -> dict: """Helper to build a minimal deployment dict as returned by router.get_model_list.""" litellm_params: dict = {"model": model_name} if tpm is not None: diff --git a/tests/test_litellm/proxy/test_model_info_default_limits.py b/tests/test_litellm/proxy/test_model_info_default_limits.py index 907f2390a89..8b85531785a 100644 --- a/tests/test_litellm/proxy/test_model_info_default_limits.py +++ b/tests/test_litellm/proxy/test_model_info_default_limits.py @@ -3,6 +3,7 @@ Tests verifying that default_api_key_tpm_limit and default_api_key_rpm_limit set litellm_params are returned by the /model/info endpoint. """ +from typing import Optional from unittest.mock import MagicMock, patch import pytest @@ -13,8 +14,8 @@ from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo def _make_deployment( model_name: str, - default_tpm: int = None, - default_rpm: int = None, + default_tpm: Optional[int] = None, + default_rpm: Optional[int] = None, ) -> Deployment: params: dict = {"model": f"openai/{model_name}"} if default_tpm is not None: From b20c448188acabe3df75597d1c76a65b07d61149 Mon Sep 17 00:00:00 2001 From: chengyongru <61816729+chengyongru@users.noreply.github.com> Date: Thu, 19 Mar 2026 15:34:47 +0800 Subject: [PATCH 10/42] fix(openai): handle missing 'id' field in streaming chunks for MiniMax (#23931) - Change chunk["id"] to chunk.get("id") for compatibility with MiniMax - ModelResponseStream auto-generates id when None is passed - Add regression test test_chunk_parser_without_id_field --- .../llms/openai/chat/gpt_transformation.py | 2 +- .../chat/test_openai_gpt_transformation.py | 39 +++++++++++++++++++ 2 files changed, 40 insertions(+), 1 deletion(-) diff --git a/litellm/llms/openai/chat/gpt_transformation.py b/litellm/llms/openai/chat/gpt_transformation.py index 63beb82ded8..34a23222c2a 100644 --- a/litellm/llms/openai/chat/gpt_transformation.py +++ b/litellm/llms/openai/chat/gpt_transformation.py @@ -806,7 +806,7 @@ class OpenAIChatCompletionStreamingHandler(BaseModelResponseIterator): choices = self._map_reasoning_to_reasoning_content(choices) kwargs = { - "id": chunk["id"], + "id": chunk.get("id"), "object": "chat.completion.chunk", "created": chunk.get("created"), "model": chunk.get("model"), diff --git a/tests/test_litellm/llms/openai/chat/test_openai_gpt_transformation.py b/tests/test_litellm/llms/openai/chat/test_openai_gpt_transformation.py index 086d01f65b4..5d0b1ec8565 100644 --- a/tests/test_litellm/llms/openai/chat/test_openai_gpt_transformation.py +++ b/tests/test_litellm/llms/openai/chat/test_openai_gpt_transformation.py @@ -281,6 +281,45 @@ class TestOpenAIChatCompletionStreamingHandler: # Verify that reasoning_content is not set (it should be deleted by Delta.__init__) assert not hasattr(parsed_chunk.choices[0].delta, "reasoning_content") + def test_chunk_parser_without_id_field(self): + """ + Test that chunk_parser works when chunk is missing the 'id' field. + + Some OpenAI-compatible providers (e.g., MiniMax) return streaming chunks + without an 'id' field in certain cases. This should not raise KeyError. + + Regression test for: KeyError: 'id' when using MiniMax m2.5 model + """ + handler = OpenAIChatCompletionStreamingHandler( + streaming_response=None, sync_stream=True + ) + + # Simulate a chunk without 'id' field (as returned by MiniMax) + chunk = { + "object": "chat.completion.chunk", + "created": 1769511767, + "model": "minimax/m2.5", + "choices": [ + { + "delta": { + "content": "Hello", + "role": "assistant", + }, + "finish_reason": None, + "index": 0, + } + ], + } + + # Parse the chunk - should not raise KeyError + parsed_chunk = handler.chunk_parser(chunk) + + # Verify that content is present and id was auto-generated + assert parsed_chunk.choices[0].delta.content == "Hello" + assert parsed_chunk.choices[0].delta.role == "assistant" + # ModelResponseStream auto-generates an id when None is passed + assert parsed_chunk.id is not None + class TestPromptCacheKeyIntegration: """Tests for prompt_cache_key support""" From e19a717b53bd6129c350ebf0c7c63a592ec5ad08 Mon Sep 17 00:00:00 2001 From: superpoussin22 Date: Thu, 19 Mar 2026 09:22:10 +0100 Subject: [PATCH 11/42] Add IF NOT EXISTS to index creation in migration --- .../20260318140652_add_index_to_team_table/migration.sql | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260318140652_add_index_to_team_table/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260318140652_add_index_to_team_table/migration.sql index 494aaf6238f..89121d636f4 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260318140652_add_index_to_team_table/migration.sql +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260318140652_add_index_to_team_table/migration.sql @@ -1,9 +1,9 @@ -- CreateIndex -CREATE INDEX "LiteLLM_TeamTable_organization_id_idx" ON "LiteLLM_TeamTable"("organization_id"); +CREATE INDEX IF NOT EXISTS "LiteLLM_TeamTable_organization_id_idx" ON "LiteLLM_TeamTable"("organization_id"); -- CreateIndex -CREATE INDEX "LiteLLM_TeamTable_team_alias_idx" ON "LiteLLM_TeamTable"("team_alias"); +CREATE INDEX IF NOT EXISTS "LiteLLM_TeamTable_team_alias_idx" ON "LiteLLM_TeamTable"("team_alias"); -- CreateIndex -CREATE INDEX "LiteLLM_TeamTable_created_at_idx" ON "LiteLLM_TeamTable"("created_at"); +CREATE INDEX IF NOT EXISTS "LiteLLM_TeamTable_created_at_idx" ON "LiteLLM_TeamTable"("created_at"); From e562c1d0640283e4820deff27ee3933646a29f65 Mon Sep 17 00:00:00 2001 From: Ephrim Stanley Date: Thu, 19 Mar 2026 07:26:43 -0400 Subject: [PATCH 12/42] refactor: consolidate duplicate helpers and eliminate success-handler double lookup - Merge _get_deployment_default_rpm_limit and _get_deployment_default_tpm_limit into a single _get_deployment_default_limit(model_name, field) helper; the two thin wrappers are preserved for callers but share one implementation - Compute _success_tpm_limit / _success_rpm_limit once before the guard condition in async_log_success_event, eliminating the previous two unconditional get_key_model_* calls (each of which could hit llm_router.get_model_list) - Replace fragile llm_model_list=[{}] sentinel in test with [] Co-Authored-By: Claude (claude-sonnet-4-6) --- litellm/proxy/auth/auth_utils.py | 46 ++++++------------- .../proxy/hooks/parallel_request_limiter.py | 20 ++++---- .../proxy/test_model_info_default_limits.py | 2 +- 3 files changed, 26 insertions(+), 42 deletions(-) diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index ace39c05ffc..b4e0093b91a 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -539,15 +539,14 @@ def bytes_to_mb(bytes_value: int): # helpers used by parallel request limiter to handle model rpm/tpm limits for a given api key -def _get_deployment_default_rpm_limit(model_name: str) -> Optional[int]: +def _get_deployment_default_limit(model_name: str, field: str) -> Optional[int]: """ - Return the default_api_key_rpm_limit for model_name. + Return the minimum value of `field` across all deployments for model_name, + or None if no deployment has the field set. - When multiple deployments share the same model name, returns the minimum - across all deployments that have the field set. This is the safest choice - for load-balanced setups: it ensures no deployment is over-consumed - regardless of which one actually serves a given request. - Returns None if no deployment has the field set. + When multiple deployments share the same model name, taking the minimum is + the safest choice for load-balanced setups: it ensures no deployment is + over-consumed regardless of which one actually serves a given request. """ from litellm.proxy.proxy_server import llm_router @@ -557,38 +556,19 @@ def _get_deployment_default_rpm_limit(model_name: str) -> Optional[int]: if not deployments: return None limits = [ - int(deployment.get("litellm_params", {}).get("default_api_key_rpm_limit")) + int(deployment.get("litellm_params", {}).get(field)) for deployment in deployments - if deployment.get("litellm_params", {}).get("default_api_key_rpm_limit") - is not None + if deployment.get("litellm_params", {}).get(field) is not None ] return min(limits) if limits else None +def _get_deployment_default_rpm_limit(model_name: str) -> Optional[int]: + return _get_deployment_default_limit(model_name, "default_api_key_rpm_limit") + + def _get_deployment_default_tpm_limit(model_name: str) -> Optional[int]: - """ - Return the default_api_key_tpm_limit for model_name. - - When multiple deployments share the same model name, returns the minimum - across all deployments that have the field set. This is the safest choice - for load-balanced setups: it ensures no deployment is over-consumed - regardless of which one actually serves a given request. - Returns None if no deployment has the field set. - """ - from litellm.proxy.proxy_server import llm_router - - if llm_router is None: - return None - deployments = llm_router.get_model_list(model_name=model_name) - if not deployments: - return None - limits = [ - int(deployment.get("litellm_params", {}).get("default_api_key_tpm_limit")) - for deployment in deployments - if deployment.get("litellm_params", {}).get("default_api_key_tpm_limit") - is not None - ] - return min(limits) if limits else None + return _get_deployment_default_limit(model_name, "default_api_key_tpm_limit") def get_key_model_rpm_limit( diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index 49c6436c22f..55e89e02d67 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -539,6 +539,16 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): # Update usage - model group + API Key # ------------ model_group = get_model_group_from_litellm_kwargs(kwargs) + _success_tpm_limit = ( + get_key_model_tpm_limit(user_api_key_dict, model_name=model_group) + if model_group is not None + else None + ) + _success_rpm_limit = ( + get_key_model_rpm_limit(user_api_key_dict, model_name=model_group) + if model_group is not None + else None + ) if ( user_api_key is not None and model_group is not None @@ -546,14 +556,8 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): "model_rpm_limit" in user_api_key_metadata or "model_tpm_limit" in user_api_key_metadata or user_api_key_model_max_budget is not None - or get_key_model_tpm_limit( - user_api_key_dict, model_name=model_group - ) - is not None - or get_key_model_rpm_limit( - user_api_key_dict, model_name=model_group - ) - is not None + or _success_tpm_limit is not None + or _success_rpm_limit is not None ) ): request_count_api_key = ( diff --git a/tests/test_litellm/proxy/test_model_info_default_limits.py b/tests/test_litellm/proxy/test_model_info_default_limits.py index 8b85531785a..d9ebd554edc 100644 --- a/tests/test_litellm/proxy/test_model_info_default_limits.py +++ b/tests/test_litellm/proxy/test_model_info_default_limits.py @@ -119,7 +119,7 @@ class TestModelInfoEndpointWithRouter: user_api_key_dict = UserAPIKeyAuth(api_key="sk-test") with patch("litellm.proxy.proxy_server.llm_router", mock_router), \ - patch("litellm.proxy.proxy_server.llm_model_list", [{}]), \ + patch("litellm.proxy.proxy_server.llm_model_list", []), \ patch("litellm.proxy.proxy_server.user_model", None): response = await model_info_v1( user_api_key_dict=user_api_key_dict, From ae0769b1dfb44b4f4ec9a676e04c2e778b426957 Mon Sep 17 00:00:00 2001 From: Ephrim Stanley Date: Thu, 19 Mar 2026 07:40:47 -0400 Subject: [PATCH 13/42] fix: guard empty-dict team limits and malformed int in deployment default limits - Change `if team_limit:` to `if team_limit is not None:` in both get_key_model_rpm_limit and get_key_model_tpm_limit so that an explicitly-empty team rate-limit map ({}) is returned as-is instead of silently falling through to deployment defaults (P1 fix). - Replace the bare `int()` list comprehension in _get_deployment_default_limit with a loop that catches ValueError/TypeError so malformed config strings do not raise an unhandled exception during request handling (P2 fix). - Add corresponding unit tests for both edge cases. Co-Authored-By: Claude (claude-sonnet-4-6) --- litellm/proxy/auth/auth_utils.py | 17 +++--- .../proxy/auth/test_auth_utils.py | 54 +++++++++++++++++++ 2 files changed, 64 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index b4e0093b91a..7d3427ed4c1 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -555,11 +555,14 @@ def _get_deployment_default_limit(model_name: str, field: str) -> Optional[int]: deployments = llm_router.get_model_list(model_name=model_name) if not deployments: return None - limits = [ - int(deployment.get("litellm_params", {}).get(field)) - for deployment in deployments - if deployment.get("litellm_params", {}).get(field) is not None - ] + limits = [] + for deployment in deployments: + raw = deployment.get("litellm_params", {}).get(field) + if raw is not None: + try: + limits.append(int(raw)) + except (ValueError, TypeError): + pass return min(limits) if limits else None @@ -602,7 +605,7 @@ def get_key_model_rpm_limit( # 3. Fallback to team metadata if user_api_key_dict.team_metadata: team_limit = user_api_key_dict.team_metadata.get("model_rpm_limit") - if team_limit: + if team_limit is not None: return team_limit # 4. Fallback to deployment default_api_key_rpm_limit @@ -645,7 +648,7 @@ def get_key_model_tpm_limit( # 3. Fallback to team metadata if user_api_key_dict.team_metadata: team_limit = user_api_key_dict.team_metadata.get("model_tpm_limit") - if team_limit: + if team_limit is not None: return team_limit # 4. Fallback to deployment default_api_key_tpm_limit diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index 2058f61cb0f..b66c081a943 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -71,6 +71,19 @@ class TestGetKeyModelRpmLimit: assert result is None + def test_team_metadata_empty_rpm_dict_falls_through_to_deployment_default(self): + """Explicitly empty team model_rpm_limit ({}) should be returned as-is, not fallen through.""" + # An empty dict is a valid team limit map (no per-model limits configured). + # It should be returned directly rather than falling through to deployment defaults, + # so a team with an empty map is treated as unconstrained at the team level. + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-123", + team_metadata={"model_rpm_limit": {}}, + ) + result = get_key_model_rpm_limit(user_api_key_dict) + assert result == {} + + class TestGetKeyModelTpmLimit: """Tests for get_key_model_tpm_limit function.""" @@ -137,6 +150,33 @@ class TestGetKeyModelTpmLimit: assert result == {"gpt-4": 10000} + def test_team_metadata_empty_tpm_dict_falls_through_to_deployment_default(self): + """Explicitly empty team model_tpm_limit ({}) should be returned as-is, not fallen through.""" + # An empty dict is a valid team limit map (no per-model limits configured). + # It should be returned directly rather than falling through to deployment defaults, + # so a team with an empty map is treated as unconstrained at the team level. + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-123", + team_metadata={"model_tpm_limit": {}}, + ) + result = get_key_model_tpm_limit(user_api_key_dict) + assert result == {} + + + def test_skips_deployments_with_malformed_limit_value(self): + """Deployments with non-integer-parseable limit values are skipped without raising.""" + user_api_key_dict = UserAPIKeyAuth(api_key="sk-123") + mock_router = MagicMock() + mock_router.get_model_list.return_value = [ + {"model_name": "model1", "litellm_params": {"default_api_key_tpm_limit": "not-a-number"}}, + _make_deployment_dict("model1", tpm=500), + ] + with patch(_ROUTER_PATCH, mock_router): + result = get_key_model_tpm_limit(user_api_key_dict, model_name="model1") + # The malformed deployment is skipped; the valid one provides 500 + assert result == {"model1": 500} + + class TestGetCustomerIdFromStandardHeaders: """Tests for _get_customer_id_from_standard_headers helper function.""" @@ -414,6 +454,20 @@ class TestDeploymentDefaultRpmLimit: assert result == {"model1": 75} + def test_skips_deployments_with_malformed_limit_value(self): + """Deployments with non-integer-parseable limit values are skipped without raising.""" + user_api_key_dict = UserAPIKeyAuth(api_key="sk-123") + mock_router = MagicMock() + mock_router.get_model_list.return_value = [ + {"model_name": "model1", "litellm_params": {"default_api_key_rpm_limit": "not-a-number"}}, + _make_deployment_dict("model1", rpm=100), + ] + with patch(_ROUTER_PATCH, mock_router): + result = get_key_model_rpm_limit(user_api_key_dict, model_name="model1") + # The malformed deployment is skipped; the valid one provides 100 + assert result == {"model1": 100} + + class TestDeploymentDefaultTpmLimit: """Tests for deployment default_api_key_tpm_limit fallback in get_key_model_tpm_limit.""" From 001501fb31dd77e6e911fdd565eec6f24a7ae26f Mon Sep 17 00:00:00 2001 From: michelligabriele Date: Thu, 19 Mar 2026 15:56:11 +0100 Subject: [PATCH 14/42] fix(proxy): defer logging until post-call guardrails complete guardrail_information is None in StandardLoggingPayload because logging fires before post-call guardrails write to metadata. Non-streaming: wrapper_async stores a closure instead of calling create_task immediately. The proxy fires it in a try/finally after post_call_success_hook so the SLP is built with guardrail info. Streaming: a closure on logging_obj is called by CSW.__anext__ at stream end. The closure runs only guardrail hooks (not all callbacks) on the assembled response, then fires both logging handlers. This avoids behavioral changes for non-guardrail callbacks on streaming. --- .../docs/proxy/guardrails/custom_guardrail.md | 12 +- .../litellm_core_utils/streaming_handler.py | 36 +- litellm/proxy/common_request_processing.py | 407 ++++++++--- litellm/utils.py | 39 +- .../test_deferred_guardrail_logging.py | 687 ++++++++++++++++++ 5 files changed, 1056 insertions(+), 125 deletions(-) create mode 100644 tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py diff --git a/docs/my-website/docs/proxy/guardrails/custom_guardrail.md b/docs/my-website/docs/proxy/guardrails/custom_guardrail.md index c9115cf8265..638cae9c835 100644 --- a/docs/my-website/docs/proxy/guardrails/custom_guardrail.md +++ b/docs/my-website/docs/proxy/guardrails/custom_guardrail.md @@ -117,6 +117,14 @@ guardrails: ::: +:::note Streaming and post_call guardrails + +For **streaming responses**, `post_call` guardrails run on the fully assembled response **after** all chunks have been delivered to the client. This means `post_call` guardrails on streaming are **audit-only** — they can inspect and log the complete response, but cannot block content delivery. Guardrail results are recorded in `guardrail_information` within the logging payload for compliance and auditing. + +To filter or block streaming content in real-time, use `async_post_call_streaming_iterator_hook` instead, which processes chunks as they arrive. + +::: +
Advanced: Multiple modes with individual event hooks @@ -655,8 +663,8 @@ class myCustomGuardrail(CustomGuardrail): | `apply_guardrail` | Simple method to check and optionally modify text | ✅ | INPUT or OUTPUT | ✅ | ✅ | ✅ | | `async_pre_call_hook` | A hook that runs before the LLM API call | ✅ | INPUT | ✅ | ❌ | ✅ | | `async_moderation_hook` | A hook that runs during the LLM API call| ✅ | INPUT | ❌ | ❌ | ✅ | -| `async_post_call_success_hook` | A hook that runs after a successful LLM API call| ✅ | INPUT, OUTPUT | ❌ | ✅ | ✅ | -| `async_post_call_streaming_iterator_hook` | A hook that processes streaming responses | ✅ | OUTPUT | ❌ | ✅ | ✅ | +| `async_post_call_success_hook` | A hook that runs after a successful LLM API call. For streaming, runs on the assembled response after delivery (audit-only, cannot block). | ✅ | INPUT, OUTPUT | ❌ | ✅ | ✅ (non-streaming only) | +| `async_post_call_streaming_iterator_hook` | A hook that processes streaming responses in real-time (can filter/block chunks) | ✅ | OUTPUT | ❌ | ✅ | ✅ | ## Frequently Asked Questions diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index 6e991e6911b..ccba18088ec 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -2136,22 +2136,36 @@ class CustomStreamWrapper: self.sent_stream_usage = True return response - asyncio.create_task( - self.logging_obj.async_success_handler( + _deferred_cb = getattr( + self.logging_obj, + "_on_deferred_stream_complete", + None, + ) + if _deferred_cb is not None: + # Proxy has post-call guardrails — let the closure + # run guardrails on the assembled response, then + # fire logging with guardrail_information populated. + self.logging_obj._on_deferred_stream_complete = None # type: ignore[attr-defined] + asyncio.create_task( + _deferred_cb(complete_streaming_response, cache_hit) + ) + else: + asyncio.create_task( + self.logging_obj.async_success_handler( + complete_streaming_response, + cache_hit=cache_hit, + start_time=None, + end_time=None, + ) + ) + + executor.submit( + self.logging_obj.success_handler, complete_streaming_response, cache_hit=cache_hit, start_time=None, end_time=None, ) - ) - - executor.submit( - self.logging_obj.success_handler, - complete_streaming_response, - cache_hit=cache_hit, - start_time=None, - end_time=None, - ) raise StopAsyncIteration # Re-raise StopIteration else: diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 72765aab7da..eb2a376ef16 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -45,7 +45,9 @@ from litellm.proxy.common_utils.callback_utils import ( from litellm.proxy.dd_span_tagger import DDSpanTagger from litellm.proxy.route_llm_request import route_request from litellm.proxy.utils import ProxyLogging +from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.router import Router +from litellm.types.guardrails import GuardrailEventHooks from litellm.types.utils import ServerToolUse # Type alias for streaming chunk serializer (chunk after hooks + cost injection -> wire format) @@ -801,7 +803,7 @@ class ProxyBaseLLMRequestProcessing: json.dumps(self.data, indent=4, default=str), ) - async def base_process_llm_request( + async def base_process_llm_request( # noqa: PLR0915 self, request: Request, fastapi_response: Response, @@ -926,6 +928,26 @@ class ProxyBaseLLMRequestProcessing: llm_router=llm_router, ) + # Defer async logging when post-call guardrails are configured so the + # StandardLoggingPayload is built after guardrails write to metadata. + # Cache the result to avoid scanning litellm.callbacks twice. + _has_post_call_guardrails = self._has_post_call_guardrails() + + # Non-streaming: defer the create_task in wrapper_async so the + # SLP is built after guardrails write to metadata. Streaming + # uses a separate closure mechanism (see below). + # + # Edge case: if _is_streaming_request is False but the response + # turns out to be a CustomStreamWrapper (rare provider behavior), + # wrapper_async exits early before the _defer_async_logging block + # so _enqueue_deferred_logging is never stored — the finally + # block is a no-op. The CSW path handles this correctly via + # _on_deferred_stream_complete, which fires its own logging. + if _has_post_call_guardrails and not self._is_streaming_request( + data=self.data, is_streaming_request=is_streaming_request + ): + logging_obj._defer_async_logging = True # type: ignore + tasks = [] # Start the moderation check (during_call_hook) as early as possible # This gives it a head start to mask/validate input while the proxy handles routing @@ -962,124 +984,181 @@ class ProxyBaseLLMRequestProcessing: response = responses[1] - hidden_params = getattr(response, "_hidden_params", {}) or {} - model_id = self._get_model_id_from_response(hidden_params, self.data) + try: + hidden_params = getattr(response, "_hidden_params", {}) or {} + model_id = self._get_model_id_from_response(hidden_params, self.data) - cache_key, api_base, response_cost = ( - hidden_params.get("cache_key", None) or "", - hidden_params.get("api_base", None) or "", - hidden_params.get("response_cost", None) or "", - ) - fastest_response_batch_completion, additional_headers = ( - hidden_params.get("fastest_response_batch_completion", None), - hidden_params.get("additional_headers", {}) or {}, - ) - - # Post Call Processing - if llm_router is not None: - self.data["deployment"] = llm_router.get_deployment(model_id=model_id) - asyncio.create_task( - proxy_logging_obj.update_request_status( - litellm_call_id=self.data.get("litellm_call_id", ""), status="success" + cache_key, api_base, response_cost = ( + hidden_params.get("cache_key", None) or "", + hidden_params.get("api_base", None) or "", + hidden_params.get("response_cost", None) or "", ) - ) - if self._is_streaming_request( - data=self.data, is_streaming_request=is_streaming_request - ) or self._is_streaming_response( - response - ): # use generate_responses to stream responses - custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers( - user_api_key_dict=user_api_key_dict, - call_id=logging_obj.litellm_call_id, - model_id=model_id, - cache_key=cache_key, - api_base=api_base, - version=version, - response_cost=response_cost, - model_region=getattr(user_api_key_dict, "allowed_model_region", ""), - fastest_response_batch_completion=fastest_response_batch_completion, - request_data=self.data, - hidden_params=hidden_params, - litellm_logging_obj=logging_obj, - **additional_headers, + fastest_response_batch_completion, additional_headers = ( + hidden_params.get("fastest_response_batch_completion", None), + hidden_params.get("additional_headers", {}) or {}, ) - # Call response headers hook for streaming success - callback_headers = await proxy_logging_obj.post_call_response_headers_hook( - data=self.data, - user_api_key_dict=user_api_key_dict, - response=response, - request_headers=dict(request.headers), + # Post Call Processing + if llm_router is not None: + self.data["deployment"] = llm_router.get_deployment(model_id=model_id) + asyncio.create_task( + proxy_logging_obj.update_request_status( + litellm_call_id=self.data.get("litellm_call_id", ""), status="success" + ) ) - if callback_headers: - custom_headers.update(callback_headers) + if self._is_streaming_request( + data=self.data, is_streaming_request=is_streaming_request + ) or self._is_streaming_response( + response + ): # use generate_responses to stream responses + custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers( + user_api_key_dict=user_api_key_dict, + call_id=logging_obj.litellm_call_id, + model_id=model_id, + cache_key=cache_key, + api_base=api_base, + version=version, + response_cost=response_cost, + model_region=getattr(user_api_key_dict, "allowed_model_region", ""), + fastest_response_batch_completion=fastest_response_batch_completion, + request_data=self.data, + hidden_params=hidden_params, + litellm_logging_obj=logging_obj, + **additional_headers, + ) - # Preserve the original client-requested model (pre-alias mapping) for downstream - # streaming generators. Pre-call processing can rewrite `self.data["model"]` for - # aliasing/routing, but the OpenAI-compatible response `model` field should reflect - # what the client sent. - if requested_model_from_client: - self.data[ - "_litellm_client_requested_model" - ] = requested_model_from_client - if route_type == "allm_passthrough_route": - # Check if response is an async generator - if self._is_streaming_response(response): - if asyncio.iscoroutine(response): - generator = await response - else: - generator = response + # Call response headers hook for streaming success + callback_headers = await proxy_logging_obj.post_call_response_headers_hook( + data=self.data, + user_api_key_dict=user_api_key_dict, + response=response, + request_headers=dict(request.headers), + ) + if callback_headers: + custom_headers.update(callback_headers) - # For passthrough routes, stream directly without error parsing - # since we're dealing with raw binary data (e.g., AWS event streams) - return StreamingResponse( - content=generator, - status_code=status.HTTP_200_OK, - headers=custom_headers, - ) - else: - # Traditional HTTP response with aiter_bytes - return StreamingResponse( - content=response.aiter_bytes(), - status_code=response.status_code, - headers=custom_headers, - ) - elif route_type == "anthropic_messages": - # Check if response is actually a streaming response (async generator) - # Non-streaming responses (dict) should be returned directly - # This handles cases like websearch_interception agentic loop - # which returns a non-streaming dict even for streaming requests - if self._is_streaming_response(response): - selected_data_generator = ( - ProxyBaseLLMRequestProcessing.async_sse_data_generator( - response=response, - user_api_key_dict=user_api_key_dict, - request_data=self.data, - proxy_logging_obj=proxy_logging_obj, + # Preserve the original client-requested model (pre-alias mapping) for downstream + # streaming generators. Pre-call processing can rewrite `self.data["model"]` for + # aliasing/routing, but the OpenAI-compatible response `model` field should reflect + # what the client sent. + if requested_model_from_client: + self.data[ + "_litellm_client_requested_model" + ] = requested_model_from_client + + # Streaming: attach a closure that CSW.__anext__ will call + # at stream end instead of firing logging directly. The + # closure runs ONLY guardrail hooks (not all callbacks) on + # the assembled response so guardrail_information is + # populated, then fires both logging handlers. + # Only for CustomStreamWrapper — raw async generators from + # passthrough routes bypass CSW and would orphan the closure. + from litellm.litellm_core_utils.streaming_handler import ( + CustomStreamWrapper, + ) + + if _has_post_call_guardrails and isinstance( + response, CustomStreamWrapper + ): + # Intentionally a live reference (not a copy) — mirrors + # ProxyLogging.post_call_success_hook which also mutates + # data["guardrail_to_apply"] during iteration. + _captured_data = self.data + _captured_user_api_key_dict = user_api_key_dict + _captured_logging_obj = logging_obj + + async def _on_deferred_stream_complete( + assembled_response, cache_hit + ): + await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails( + captured_data=_captured_data, + captured_user_api_key_dict=_captured_user_api_key_dict, + captured_logging_obj=_captured_logging_obj, + assembled_response=assembled_response, + cache_hit=cache_hit, ) + + logging_obj._on_deferred_stream_complete = _on_deferred_stream_complete # type: ignore[attr-defined] + + if route_type == "allm_passthrough_route": + # Check if response is an async generator + if self._is_streaming_response(response): + if asyncio.iscoroutine(response): + generator = await response + else: + generator = response + + # For passthrough routes, stream directly without error parsing + # since we're dealing with raw binary data (e.g., AWS event streams) + return StreamingResponse( + content=generator, + status_code=status.HTTP_200_OK, + headers=custom_headers, + ) + else: + # Traditional HTTP response with aiter_bytes + return StreamingResponse( + content=response.aiter_bytes(), + status_code=response.status_code, + headers=custom_headers, + ) + elif route_type == "anthropic_messages": + # Check if response is actually a streaming response (async generator) + # Non-streaming responses (dict) should be returned directly + # This handles cases like websearch_interception agentic loop + # which returns a non-streaming dict even for streaming requests + if self._is_streaming_response(response): + selected_data_generator = ( + ProxyBaseLLMRequestProcessing.async_sse_data_generator( + response=response, + user_api_key_dict=user_api_key_dict, + request_data=self.data, + proxy_logging_obj=proxy_logging_obj, + ) + ) + return await create_response( + generator=selected_data_generator, + media_type="text/event-stream", + headers=custom_headers, + ) + # Non-streaming response - fall through to normal response handling + elif select_data_generator: + selected_data_generator = select_data_generator( + response=response, + user_api_key_dict=user_api_key_dict, + request_data=self.data, ) return await create_response( generator=selected_data_generator, media_type="text/event-stream", headers=custom_headers, ) - # Non-streaming response - fall through to normal response handling - elif select_data_generator: - selected_data_generator = select_data_generator( - response=response, - user_api_key_dict=user_api_key_dict, - request_data=self.data, - ) - return await create_response( - generator=selected_data_generator, - media_type="text/event-stream", - headers=custom_headers, - ) - ### CALL HOOKS ### - modify outgoing data - response = await proxy_logging_obj.post_call_success_hook( - data=self.data, user_api_key_dict=user_api_key_dict, response=response - ) + ### CALL HOOKS ### - modify outgoing data + # If we reach here with a streaming closure still set, it means + # no early-return route consumed the CSW (hypothetical fallthrough). + # Clear the closure so guardrails run inline as before — this + # preserves blocking behavior and avoids double invocation. + if getattr(logging_obj, "_on_deferred_stream_complete", None): + logging_obj._on_deferred_stream_complete = None # type: ignore[attr-defined] + response = await proxy_logging_obj.post_call_success_hook( + data=self.data, user_api_key_dict=user_api_key_dict, response=response + ) + finally: + # Enqueue deferred logging after post-call guardrails have written + # guardrail_information to metadata. The finally block ensures + # logging fires even if a guardrail raises. + # For streaming early-returns: no closure is stored (wrapper_async + # returns before the deferred block), so _enqueue_fn is None — no-op. + _enqueue_fn = getattr(logging_obj, "_enqueue_deferred_logging", None) + if _enqueue_fn is not None: + logging_obj._enqueue_deferred_logging = None # type: ignore[attr-defined] + try: + _enqueue_fn() + except Exception as e: + verbose_proxy_logger.exception( + "Error firing deferred logging: %s", e + ) # Always return the client-requested model name (not provider-prefixed internal identifiers) # for OpenAI-compatible responses. @@ -1217,6 +1296,126 @@ class ProxyBaseLLMRequestProcessing: return True return False + @staticmethod + def _has_post_call_guardrails() -> bool: + """ + Check if any registered callback is a post-call guardrail. + + Uses the global litellm.callbacks list rather than per-request + should_run_guardrail() — intentionally conservative so that the + check is simple and stateless. The deferral path produces + identical logging output, just fires it slightly later, so + false-positives are harmless. + """ + for cb in litellm.callbacks: + if isinstance(cb, CustomGuardrail) and cb._event_hook_is_event_type( + GuardrailEventHooks.post_call + ): + return True + return False + + @staticmethod + async def _run_deferred_stream_guardrails( + captured_data: dict, + captured_user_api_key_dict: "UserAPIKeyAuth", + captured_logging_obj: Any, + assembled_response: Any, + cache_hit: Any, + ) -> None: + """ + Run only post-call guardrail hooks on an assembled streaming response, + then fire both async and sync logging handlers. + + Called by CSW.__anext__ at stream end via a closure stored on + logging_obj._on_deferred_stream_complete. + + This is audit-only — content has already been delivered to the client. + Blocking guardrails that raise HTTPException cannot prevent content + delivery for streaming. Per-chunk filtering should use + async_post_call_streaming_hook instead. + + Extracted as a static method so tests can call the production + implementation directly rather than reimplementing the closure. + """ + from litellm.litellm_core_utils.thread_pool_executor import executor + from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import ( + UnifiedLLMGuardrails, + ) + from litellm.proxy.proxy_server import llm_router as _global_llm_router + from litellm.proxy.utils import _check_and_merge_model_level_guardrails + + _response = assembled_response + _unified_guardrail = UnifiedLLMGuardrails() + guardrail_data = _check_and_merge_model_level_guardrails( + data=captured_data, llm_router=_global_llm_router + ) + for cb in litellm.callbacks: + if not isinstance(cb, CustomGuardrail): + continue + if not cb.should_run_guardrail( + data=guardrail_data, + event_type=GuardrailEventHooks.post_call, + ): + continue + try: + guardrail_result = None + if "apply_guardrail" in type(cb).__dict__: + captured_data["guardrail_to_apply"] = cb + guardrail_result = ( + await _unified_guardrail.async_post_call_success_hook( + user_api_key_dict=captured_user_api_key_dict, + data=captured_data, + response=_response, + ) + ) + else: + guardrail_result = await cb.async_post_call_success_hook( + user_api_key_dict=captured_user_api_key_dict, + data=captured_data, + response=_response, + ) + if guardrail_result is not None: + _response = guardrail_result + except Exception as e: + verbose_proxy_logger.exception( + "Error running post-call guardrail %s on streaming response: %s", + getattr(cb, "guardrail_name", type(cb).__name__), + e, + ) + if isinstance(e, HTTPException) and hasattr( + captured_logging_obj, "model_call_details" + ): + captured_logging_obj.model_call_details.setdefault( + "metadata", {} + )["guardrail_blocked"] = True + + try: + asyncio.create_task( + captured_logging_obj.async_success_handler( + _response, + cache_hit=cache_hit, + start_time=None, + end_time=None, + ) + ) + except Exception as e: + verbose_proxy_logger.exception( + "Error in deferred streaming async logging: %s", e, + ) + + try: + executor.submit( + captured_logging_obj.success_handler, + _response, + cache_hit=cache_hit, + start_time=None, + end_time=None, + ) + except Exception as e: + verbose_proxy_logger.exception( + "Error in deferred streaming sync logging: %s", e, + ) + async def _handle_llm_api_exception( self, e: Exception, diff --git a/litellm/utils.py b/litellm/utils.py index 81d749ab821..0fda994aea1 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -1944,15 +1944,38 @@ def client(original_function): # noqa: PLR0915 ) # LOG SUCCESS - handle streaming success logging in the _next_ object - asyncio.create_task( - _client_async_logging_helper( - logging_obj=logging_obj, - result=result, - start_time=start_time, - end_time=end_time, - is_completion_with_fallbacks=is_completion_with_fallbacks, + # NOTE: streaming requests return early (before this point) via + # CustomStreamWrapper, so this block is non-streaming only. + if getattr(logging_obj, "_defer_async_logging", False): + # Proxy has post-call guardrails that must complete before the + # SLP is built. Store a closure the proxy will call after + # post_call_success_hook so guardrail_information is in metadata. + # Only create_task is deferred; sync callbacks fire immediately + # (below, outside the if/else) for billing/rate-limiting. + def _enqueue_deferred_logging() -> None: + asyncio.create_task( + _client_async_logging_helper( + logging_obj=logging_obj, + result=result, + start_time=start_time, + end_time=end_time, + is_completion_with_fallbacks=is_completion_with_fallbacks, + ) + ) + + logging_obj._enqueue_deferred_logging = _enqueue_deferred_logging # type: ignore + else: + asyncio.create_task( + _client_async_logging_helper( + logging_obj=logging_obj, + result=result, + start_time=start_time, + end_time=end_time, + is_completion_with_fallbacks=is_completion_with_fallbacks, + ) ) - ) + + # Sync callbacks always fire immediately regardless of deferral logging_obj.handle_sync_success_callbacks_for_async_calls( result=result, start_time=start_time, diff --git a/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py b/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py new file mode 100644 index 00000000000..82e389da651 --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py @@ -0,0 +1,687 @@ +""" +Tests for deferred logging with post-call guardrails. + +When post-call guardrails are configured, the async logging task is deferred +until after guardrails complete. This ensures the StandardLoggingPayload +is built with guardrail_information populated. + +Non-streaming: create_task in wrapper_async is replaced by a closure that + the proxy fires in a try/finally after post_call_success_hook. + +Streaming: a closure on logging_obj is called by CSW.__anext__ at stream end. + The closure runs ONLY guardrail hooks (not all callbacks), then fires + both logging handlers. +""" + +import asyncio +import os +import sys +from typing import Any +from unittest.mock import MagicMock, patch + +import pytest +from starlette.exceptions import HTTPException + +sys.path.insert(0, os.path.abspath("../../../..")) + +import litellm +from litellm.caching.caching import DualCache +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.integrations.custom_logger import CustomLogger +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing +from litellm.proxy.utils import ProxyLogging +from litellm.types.guardrails import GuardrailEventHooks + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +class PostCallGuardrail(CustomGuardrail): + """A post-call guardrail.""" + + def __init__(self): + super().__init__( + guardrail_name="post-call", + default_on=True, + event_hook=GuardrailEventHooks.post_call, + ) + + async def async_post_call_success_hook( + self, data: dict, user_api_key_dict: UserAPIKeyAuth, response: Any + ) -> Any: + return response + + +class PreCallGuardrail(CustomGuardrail): + """A pre-call-only guardrail — should NOT trigger deferral.""" + + def __init__(self): + super().__init__( + guardrail_name="pre-call", + default_on=True, + event_hook=GuardrailEventHooks.pre_call, + ) + + +class AllEventsGuardrail(CustomGuardrail): + """A guardrail with event_hook=None (runs on all events).""" + + def __init__(self): + super().__init__( + guardrail_name="all-events", + default_on=True, + event_hook=None, + ) + + +# --------------------------------------------------------------------------- +# 1. _has_post_call_guardrails detection +# --------------------------------------------------------------------------- + + +class TestHasPostCallGuardrails: + def test_returns_true_for_post_call_guardrail(self): + with patch("litellm.callbacks", [PostCallGuardrail()]): + assert ProxyBaseLLMRequestProcessing._has_post_call_guardrails() is True + + def test_returns_true_for_event_hook_none(self): + """event_hook=None means 'all events', including post_call.""" + with patch("litellm.callbacks", [AllEventsGuardrail()]): + assert ProxyBaseLLMRequestProcessing._has_post_call_guardrails() is True + + def test_returns_false_for_pre_call_only(self): + with patch("litellm.callbacks", [PreCallGuardrail()]): + assert ProxyBaseLLMRequestProcessing._has_post_call_guardrails() is False + + def test_returns_false_for_no_callbacks(self): + with patch("litellm.callbacks", []): + assert ProxyBaseLLMRequestProcessing._has_post_call_guardrails() is False + + def test_ignores_non_guardrail_callbacks(self): + """String callbacks and CustomLogger instances are not guardrails.""" + with patch("litellm.callbacks", ["langfuse", CustomLogger()]): + assert ProxyBaseLLMRequestProcessing._has_post_call_guardrails() is False + + def test_returns_true_for_list_with_post_call(self): + """event_hook as a list containing post_call should trigger deferral.""" + + class ListGuardrail(CustomGuardrail): + def __init__(self): + super().__init__( + guardrail_name="list-post", + default_on=True, + event_hook=[GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call], + ) + + with patch("litellm.callbacks", [ListGuardrail()]): + assert ProxyBaseLLMRequestProcessing._has_post_call_guardrails() is True + + def test_returns_false_for_list_without_post_call(self): + """event_hook as a list without post_call should not trigger deferral.""" + + class ListGuardrail(CustomGuardrail): + def __init__(self): + super().__init__( + guardrail_name="list-pre", + default_on=True, + event_hook=[GuardrailEventHooks.pre_call], + ) + + with patch("litellm.callbacks", [ListGuardrail()]): + assert ProxyBaseLLMRequestProcessing._has_post_call_guardrails() is False + + +# --------------------------------------------------------------------------- +# 2. Non-streaming: deferral flag → closure stored, create_task skipped +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_deferred_flag_stores_and_executes_closure(): + """ + When _defer_async_logging is True on logging_obj: + 1. wrapper_async stores a callable closure instead of calling create_task + 2. Calling the closure fires create_task + 3. Sync callbacks fire immediately (not deferred) + """ + mock_logging_obj = MagicMock() + mock_logging_obj._defer_async_logging = True + mock_logging_obj._enqueue_deferred_logging = None + + await litellm.acompletion( + model="gpt-3.5-turbo", + messages=[{"role": "user", "content": "hi"}], + mock_response="Hello!", + litellm_logging_obj=mock_logging_obj, + ) + + # Closure was stored + enqueue_fn = mock_logging_obj._enqueue_deferred_logging + assert callable(enqueue_fn), "Closure should be stored on logging_obj" + + # Sync callbacks fired immediately + mock_logging_obj.handle_sync_success_callbacks_for_async_calls.assert_called_once() + + # Calling the closure fires create_task + created_tasks = [] + real_create_task = asyncio.create_task + + def tracking_create_task(coro): + task = real_create_task(coro) + created_tasks.append(task) + return task + + with patch("asyncio.create_task", side_effect=tracking_create_task): + enqueue_fn() + + assert len(created_tasks) >= 1, "Closure should fire asyncio.create_task" + + for task in created_tasks: + if not task.done(): + task.cancel() + try: + await task + except (asyncio.CancelledError, Exception): + pass + + +# --------------------------------------------------------------------------- +# 3. Non-streaming regression: without flag, create_task fires normally +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_no_flag_fires_create_task_normally(): + """Without _defer_async_logging, wrapper_async calls create_task as before.""" + created_tasks = [] + real_create_task = asyncio.create_task + + def tracking_create_task(coro): + task = real_create_task(coro) + created_tasks.append(task) + return task + + with patch("asyncio.create_task", side_effect=tracking_create_task): + await litellm.acompletion( + model="gpt-3.5-turbo", + messages=[{"role": "user", "content": "hi"}], + mock_response="Hello!", + ) + + assert len(created_tasks) >= 1 + + for task in created_tasks: + if not task.done(): + task.cancel() + try: + await task + except (asyncio.CancelledError, Exception): + pass + + +# --------------------------------------------------------------------------- +# 4. Non-streaming: deferred logging fires even if guardrail raises +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_deferred_logging_fires_on_guardrail_exception(): + """ + If post_call_success_hook raises (e.g., guardrail blocks content), + the deferred logging closure must still fire (via try/finally). + """ + enqueue_called = False + + def mock_enqueue(): + nonlocal enqueue_called + enqueue_called = True + + class BlockingGuardrail(CustomGuardrail): + def __init__(self): + super().__init__( + guardrail_name="blocker", + default_on=True, + event_hook=GuardrailEventHooks.post_call, + ) + + async def async_post_call_success_hook( + self, data: dict, user_api_key_dict: UserAPIKeyAuth, response: Any + ) -> Any: + raise HTTPException(status_code=400, detail="Content blocked") + + guardrail = BlockingGuardrail() + + logging_obj = MagicMock() + logging_obj._enqueue_deferred_logging = mock_enqueue + + with patch("litellm.callbacks", [guardrail]): + proxy_logging = ProxyLogging(user_api_key_cache=DualCache()) + + with pytest.raises(HTTPException): + try: + await proxy_logging.post_call_success_hook( + data={"model": "gpt-4", "metadata": {}}, + response=MagicMock(), + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + ) + finally: + # Mirrors the proxy's finally block + _enqueue_fn = getattr(logging_obj, "_enqueue_deferred_logging", None) + if _enqueue_fn is not None: + logging_obj._enqueue_deferred_logging = None + _enqueue_fn() + + assert enqueue_called is True + assert logging_obj._enqueue_deferred_logging is None + + +# --------------------------------------------------------------------------- +# 5. Streaming: closure defers logging at stream end +# --------------------------------------------------------------------------- + + +class TestDeferredStreamingClosure: + @pytest.mark.asyncio + async def test_streaming_closure_defers_logging(self): + """When _on_deferred_stream_complete is set, CSW calls the closure + instead of firing async_success_handler directly.""" + mock_logging_obj = MagicMock() + callback_called = False + callback_args = {} + + async def mock_callback(assembled_response, cache_hit): + nonlocal callback_called, callback_args + callback_called = True + callback_args = {"response": assembled_response, "cache_hit": cache_hit} + + mock_logging_obj._on_deferred_stream_complete = mock_callback + + resp = await litellm.acompletion( + model="gpt-3.5-turbo", + messages=[{"role": "user", "content": "hi"}], + mock_response="Hello!", + stream=True, + litellm_logging_obj=mock_logging_obj, + ) + async for _ in resp: + pass + + await asyncio.sleep(0) + + assert callback_called is True, "Closure should be called at stream end" + assert callback_args["response"] is not None + assert mock_logging_obj._on_deferred_stream_complete is None + + @pytest.mark.asyncio + async def test_streaming_no_closure_fires_normally(self): + """Regression: without closure, CSW fires logging immediately.""" + created_tasks = [] + real_create_task = asyncio.create_task + + def tracking_create_task(coro): + task = real_create_task(coro) + created_tasks.append(task) + return task + + resp = await litellm.acompletion( + model="gpt-3.5-turbo", + messages=[{"role": "user", "content": "hi"}], + mock_response="Hello!", + stream=True, + ) + with patch("asyncio.create_task", side_effect=tracking_create_task): + async for _ in resp: + pass + + assert len(created_tasks) >= 1 + for task in created_tasks: + if not task.done(): + task.cancel() + try: + await task + except (asyncio.CancelledError, Exception): + pass + + @pytest.mark.asyncio + async def test_closure_runs_only_guardrail_hooks(self): + """The closure must call only CustomGuardrail hooks, not all callbacks. + This is the key v2 change — PR #23929 called post_call_success_hook + which ran ALL callbacks, causing behavioral changes for streaming.""" + guardrail_called = False + logger_called = False + + class TrackingGuardrail(CustomGuardrail): + def __init__(self): + super().__init__( + guardrail_name="tracker", + default_on=True, + event_hook=GuardrailEventHooks.post_call, + ) + + async def async_post_call_success_hook( + self, data: dict, user_api_key_dict: UserAPIKeyAuth, response: Any + ) -> Any: + nonlocal guardrail_called + guardrail_called = True + return response + + class TrackingLogger(CustomLogger): + async def async_post_call_success_hook( + self, user_api_key_dict, data, response + ): + nonlocal logger_called + logger_called = True + return response + + mock_logging_obj = MagicMock() + mock_logging_obj.model_call_details = {"metadata": {}} + + async def track_async_success(*args, **kwargs): + pass + + mock_logging_obj.async_success_handler = track_async_success + + tracking_guardrail = TrackingGuardrail() + tracking_logger = TrackingLogger() + + # Use the real production static method via a thin closure + _captured_data = {"model": "gpt-4", "metadata": {}} + _captured_user_api_key_dict = UserAPIKeyAuth(api_key="test") + + async def _on_deferred_stream_complete(assembled_response, cache_hit): + await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails( + captured_data=_captured_data, + captured_user_api_key_dict=_captured_user_api_key_dict, + captured_logging_obj=mock_logging_obj, + assembled_response=assembled_response, + cache_hit=cache_hit, + ) + + mock_logging_obj._on_deferred_stream_complete = _on_deferred_stream_complete + + with patch("litellm.callbacks", [tracking_guardrail, tracking_logger]): + resp = await litellm.acompletion( + model="gpt-3.5-turbo", + messages=[{"role": "user", "content": "hi"}], + mock_response="Hello!", + stream=True, + litellm_logging_obj=mock_logging_obj, + ) + async for _ in resp: + pass + + await asyncio.sleep(0) + await asyncio.sleep(0) + + assert guardrail_called is True, "Guardrail hook should be called" + assert logger_called is False, "Non-guardrail logger should NOT be called by closure" + + @pytest.mark.asyncio + async def test_closure_passes_guardrail_modified_response_to_logging(self): + """The closure passes the guardrail-modified response to logging handlers.""" + mock_logging_obj = MagicMock() + modified_response = MagicMock() + logged_response = None + + async def mock_async_success(*args, **kwargs): + nonlocal logged_response + logged_response = args[0] if args else None + + mock_logging_obj.async_success_handler = mock_async_success + + async def closure(assembled_response, cache_hit): + # Simulate guardrail modifying the response + asyncio.create_task( + mock_logging_obj.async_success_handler( + modified_response, cache_hit=cache_hit, start_time=None, end_time=None + ) + ) + + mock_logging_obj._on_deferred_stream_complete = closure + + resp = await litellm.acompletion( + model="gpt-3.5-turbo", + messages=[{"role": "user", "content": "hi"}], + mock_response="Hello!", + stream=True, + litellm_logging_obj=mock_logging_obj, + ) + async for _ in resp: + pass + + await asyncio.sleep(0) + await asyncio.sleep(0) + + assert logged_response is modified_response + + @pytest.mark.asyncio + async def test_closure_logs_even_on_guardrail_exception(self): + """If the guardrail raises HTTPException, logging still fires + and guardrail_blocked is set in metadata.""" + logging_called = False + + async def mock_async_success(*args, **kwargs): + nonlocal logging_called + logging_called = True + + mock_logging_obj = MagicMock() + mock_logging_obj.model_call_details = {"metadata": {}} + mock_logging_obj.async_success_handler = mock_async_success + + async def closure(assembled_response, cache_hit): + _response = assembled_response + try: + raise HTTPException(status_code=400, detail="Blocked") + except Exception as e: + if isinstance(e, HTTPException) and hasattr( + mock_logging_obj, "model_call_details" + ): + mock_logging_obj.model_call_details.setdefault( + "metadata", {} + )["guardrail_blocked"] = True + + asyncio.create_task( + mock_logging_obj.async_success_handler( + _response, cache_hit=cache_hit, start_time=None, end_time=None + ) + ) + + mock_logging_obj._on_deferred_stream_complete = closure + + resp = await litellm.acompletion( + model="gpt-3.5-turbo", + messages=[{"role": "user", "content": "hi"}], + mock_response="Hello!", + stream=True, + litellm_logging_obj=mock_logging_obj, + ) + async for _ in resp: + pass + + await asyncio.sleep(0) + await asyncio.sleep(0) + + assert logging_called is True + assert mock_logging_obj.model_call_details["metadata"].get( + "guardrail_blocked" + ) is True + + @pytest.mark.asyncio + async def test_transient_error_does_not_set_guardrail_blocked(self): + """Transient errors (not HTTPException) should NOT set guardrail_blocked.""" + mock_logging_obj = MagicMock() + mock_logging_obj.model_call_details = {"metadata": {}} + + async def mock_async_success(*args, **kwargs): + pass + + mock_logging_obj.async_success_handler = mock_async_success + + async def closure(assembled_response, cache_hit): + try: + raise ConnectionError("Network timeout") + except Exception as e: + if isinstance(e, HTTPException) and hasattr( + mock_logging_obj, "model_call_details" + ): + mock_logging_obj.model_call_details.setdefault( + "metadata", {} + )["guardrail_blocked"] = True + + asyncio.create_task( + mock_logging_obj.async_success_handler( + assembled_response, cache_hit=cache_hit, start_time=None, end_time=None + ) + ) + + mock_logging_obj._on_deferred_stream_complete = closure + + resp = await litellm.acompletion( + model="gpt-3.5-turbo", + messages=[{"role": "user", "content": "hi"}], + mock_response="Hello!", + stream=True, + litellm_logging_obj=mock_logging_obj, + ) + async for _ in resp: + pass + + await asyncio.sleep(0) + + assert mock_logging_obj.model_call_details["metadata"].get( + "guardrail_blocked" + ) is not True + + @pytest.mark.asyncio + async def test_production_closure_integration(self): + """Integration test: calls the real _run_deferred_stream_guardrails + static method and verifies it calls guardrail hooks and passes + the modified response to logging.""" + hook_called = False + logged_response = None + modified_response = MagicMock() + + mock_logging_obj = MagicMock() + mock_logging_obj.model_call_details = {"metadata": {}} + + async def track_async_success(*args, **kwargs): + nonlocal logged_response + logged_response = args[0] if args else None + + mock_logging_obj.async_success_handler = track_async_success + + class TestGuardrail(CustomGuardrail): + def __init__(self): + super().__init__( + guardrail_name="test", + default_on=True, + event_hook=GuardrailEventHooks.post_call, + ) + + async def async_post_call_success_hook( + self, data: dict, user_api_key_dict: UserAPIKeyAuth, response: Any + ) -> Any: + nonlocal hook_called + hook_called = True + return modified_response + + guardrail = TestGuardrail() + + async def _on_deferred_stream_complete(assembled_response, cache_hit): + await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails( + captured_data={"model": "gpt-4", "metadata": {}}, + captured_user_api_key_dict=UserAPIKeyAuth(api_key="test"), + captured_logging_obj=mock_logging_obj, + assembled_response=assembled_response, + cache_hit=cache_hit, + ) + + mock_logging_obj._on_deferred_stream_complete = _on_deferred_stream_complete + + with patch("litellm.callbacks", [guardrail]): + resp = await litellm.acompletion( + model="gpt-3.5-turbo", + messages=[{"role": "user", "content": "hi"}], + mock_response="Hello!", + stream=True, + litellm_logging_obj=mock_logging_obj, + ) + async for _ in resp: + pass + + await asyncio.sleep(0) + await asyncio.sleep(0) + + assert hook_called is True, \ + "Production closure must call guardrail hook" + assert logged_response is modified_response, \ + "Production closure must pass guardrail-modified response to logging" + + @pytest.mark.asyncio + async def test_apply_guardrail_path_uses_unified_guardrail(self): + """Guardrails that define apply_guardrail should be dispatched through + UnifiedLLMGuardrails.async_post_call_success_hook via the real + _run_deferred_stream_guardrails static method.""" + from litellm.types.utils import GenericGuardrailAPIInputs + + unified_hook_called = False + + class ApplyGuardrailType(CustomGuardrail): + def __init__(self): + super().__init__( + guardrail_name="apply-type", + default_on=True, + event_hook=GuardrailEventHooks.post_call, + ) + + async def apply_guardrail( + self, inputs, request_data, input_type, logging_obj=None + ) -> GenericGuardrailAPIInputs: + nonlocal unified_hook_called + unified_hook_called = True + return inputs + + mock_logging_obj = MagicMock() + mock_logging_obj.model_call_details = {"metadata": {}} + logged_response = None + + async def track_async_success(*args, **kwargs): + nonlocal logged_response + logged_response = args[0] if args else None + + mock_logging_obj.async_success_handler = track_async_success + + guardrail = ApplyGuardrailType() + + async def _on_deferred_stream_complete(assembled_response, cache_hit): + await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails( + captured_data={"model": "gpt-4", "metadata": {}}, + captured_user_api_key_dict=UserAPIKeyAuth(api_key="test"), + captured_logging_obj=mock_logging_obj, + assembled_response=assembled_response, + cache_hit=cache_hit, + ) + + mock_logging_obj._on_deferred_stream_complete = _on_deferred_stream_complete + + with patch("litellm.callbacks", [guardrail]): + resp = await litellm.acompletion( + model="gpt-3.5-turbo", + messages=[{"role": "user", "content": "hi"}], + mock_response="Hello!", + stream=True, + litellm_logging_obj=mock_logging_obj, + ) + async for _ in resp: + pass + + await asyncio.sleep(0) + await asyncio.sleep(0) + + assert unified_hook_called is True, \ + "apply_guardrail guardrails must be dispatched through UnifiedLLMGuardrails" + assert logged_response is not None, \ + "Logging must fire after unified guardrail path" From 4b8c532ba8cbdb9b55fc006c9e218e3dd669d5e3 Mon Sep 17 00:00:00 2001 From: michelligabriele Date: Thu, 19 Mar 2026 17:39:06 +0100 Subject: [PATCH 15/42] fix(proxy): pass guardrail_data to hooks in streaming deferred path Use the merged guardrail_data dict (from _check_and_merge_model_level_guardrails) for hook invocations in _run_deferred_stream_guardrails, instead of the original captured_data. This ensures model-level non-default guardrails are visible to inner should_run_guardrail re-checks inside UnifiedLLMGuardrails. Rewrite three hand-crafted closure tests to exercise the production _run_deferred_stream_guardrails exception-handling path. Add three new tests that use deep-copy mocks to prove hooks receive the merged dict. --- litellm/proxy/common_request_processing.py | 6 +- .../test_deferred_guardrail_logging.py | 395 ++++++++++++++---- 2 files changed, 309 insertions(+), 92 deletions(-) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index eb2a376ef16..1db8327482d 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1360,18 +1360,18 @@ class ProxyBaseLLMRequestProcessing: try: guardrail_result = None if "apply_guardrail" in type(cb).__dict__: - captured_data["guardrail_to_apply"] = cb + guardrail_data["guardrail_to_apply"] = cb guardrail_result = ( await _unified_guardrail.async_post_call_success_hook( user_api_key_dict=captured_user_api_key_dict, - data=captured_data, + data=guardrail_data, response=_response, ) ) else: guardrail_result = await cb.async_post_call_success_hook( user_api_key_dict=captured_user_api_key_dict, - data=captured_data, + data=guardrail_data, response=_response, ) if guardrail_result is not None: diff --git a/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py b/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py index 82e389da651..f5c9eeba1c6 100644 --- a/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py +++ b/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py @@ -20,7 +20,7 @@ from typing import Any from unittest.mock import MagicMock, patch import pytest -from starlette.exceptions import HTTPException +from fastapi import HTTPException sys.path.insert(0, os.path.abspath("../../../..")) @@ -421,139 +421,141 @@ class TestDeferredStreamingClosure: @pytest.mark.asyncio async def test_closure_passes_guardrail_modified_response_to_logging(self): - """The closure passes the guardrail-modified response to logging handlers.""" - mock_logging_obj = MagicMock() - modified_response = MagicMock() + """The production _run_deferred_stream_guardrails must pass the + guardrail-modified response to async_success_handler.""" logged_response = None + modified_response = MagicMock() - async def mock_async_success(*args, **kwargs): + mock_logging_obj = MagicMock() + mock_logging_obj.model_call_details = {"metadata": {}} + + async def track_async_success(*args, **kwargs): nonlocal logged_response logged_response = args[0] if args else None - mock_logging_obj.async_success_handler = mock_async_success + mock_logging_obj.async_success_handler = track_async_success - async def closure(assembled_response, cache_hit): - # Simulate guardrail modifying the response - asyncio.create_task( - mock_logging_obj.async_success_handler( - modified_response, cache_hit=cache_hit, start_time=None, end_time=None + class ModifyingGuardrail(CustomGuardrail): + def __init__(self): + super().__init__( + guardrail_name="modifier", + default_on=True, + event_hook=GuardrailEventHooks.post_call, ) + + async def async_post_call_success_hook( + self, data: dict, user_api_key_dict: UserAPIKeyAuth, response: Any + ) -> Any: + return modified_response + + guardrail = ModifyingGuardrail() + + with patch("litellm.callbacks", [guardrail]): + await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails( + captured_data={"model": "gpt-4", "metadata": {}}, + captured_user_api_key_dict=UserAPIKeyAuth(api_key="test"), + captured_logging_obj=mock_logging_obj, + assembled_response=MagicMock(), + cache_hit=False, ) - mock_logging_obj._on_deferred_stream_complete = closure - - resp = await litellm.acompletion( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hi"}], - mock_response="Hello!", - stream=True, - litellm_logging_obj=mock_logging_obj, - ) - async for _ in resp: - pass - await asyncio.sleep(0) await asyncio.sleep(0) - assert logged_response is modified_response + assert logged_response is modified_response, \ + "Logging must receive the guardrail-modified response" @pytest.mark.asyncio async def test_closure_logs_even_on_guardrail_exception(self): - """If the guardrail raises HTTPException, logging still fires - and guardrail_blocked is set in metadata.""" + """If a guardrail raises HTTPException, the production + _run_deferred_stream_guardrails must still fire logging + and set guardrail_blocked in metadata.""" logging_called = False - async def mock_async_success(*args, **kwargs): + class BlockingGuardrail(CustomGuardrail): + def __init__(self): + super().__init__( + guardrail_name="blocker", + default_on=True, + event_hook=GuardrailEventHooks.post_call, + ) + + async def async_post_call_success_hook( + self, data: dict, user_api_key_dict: UserAPIKeyAuth, response: Any + ) -> Any: + raise HTTPException(status_code=400, detail="Blocked") + + mock_logging_obj = MagicMock() + mock_logging_obj.model_call_details = {"metadata": {}} + + async def track_async_success(*args, **kwargs): nonlocal logging_called logging_called = True - mock_logging_obj = MagicMock() - mock_logging_obj.model_call_details = {"metadata": {}} - mock_logging_obj.async_success_handler = mock_async_success + mock_logging_obj.async_success_handler = track_async_success - async def closure(assembled_response, cache_hit): - _response = assembled_response - try: - raise HTTPException(status_code=400, detail="Blocked") - except Exception as e: - if isinstance(e, HTTPException) and hasattr( - mock_logging_obj, "model_call_details" - ): - mock_logging_obj.model_call_details.setdefault( - "metadata", {} - )["guardrail_blocked"] = True + guardrail = BlockingGuardrail() - asyncio.create_task( - mock_logging_obj.async_success_handler( - _response, cache_hit=cache_hit, start_time=None, end_time=None - ) + with patch("litellm.callbacks", [guardrail]): + await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails( + captured_data={"model": "gpt-4", "metadata": {}}, + captured_user_api_key_dict=UserAPIKeyAuth(api_key="test"), + captured_logging_obj=mock_logging_obj, + assembled_response=MagicMock(), + cache_hit=False, ) - mock_logging_obj._on_deferred_stream_complete = closure - - resp = await litellm.acompletion( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hi"}], - mock_response="Hello!", - stream=True, - litellm_logging_obj=mock_logging_obj, - ) - async for _ in resp: - pass - await asyncio.sleep(0) await asyncio.sleep(0) - assert logging_called is True + assert logging_called is True, \ + "Logging must fire even when guardrail raises HTTPException" assert mock_logging_obj.model_call_details["metadata"].get( "guardrail_blocked" - ) is True + ) is True, "guardrail_blocked must be set for HTTPException" @pytest.mark.asyncio async def test_transient_error_does_not_set_guardrail_blocked(self): - """Transient errors (not HTTPException) should NOT set guardrail_blocked.""" + """Transient errors (not HTTPException) should NOT set + guardrail_blocked. Uses the production _run_deferred_stream_guardrails.""" + + class TransientErrorGuardrail(CustomGuardrail): + def __init__(self): + super().__init__( + guardrail_name="transient", + default_on=True, + event_hook=GuardrailEventHooks.post_call, + ) + + async def async_post_call_success_hook( + self, data: dict, user_api_key_dict: UserAPIKeyAuth, response: Any + ) -> Any: + raise ConnectionError("Network timeout") + mock_logging_obj = MagicMock() mock_logging_obj.model_call_details = {"metadata": {}} - async def mock_async_success(*args, **kwargs): + async def track_async_success(*args, **kwargs): pass - mock_logging_obj.async_success_handler = mock_async_success + mock_logging_obj.async_success_handler = track_async_success - async def closure(assembled_response, cache_hit): - try: - raise ConnectionError("Network timeout") - except Exception as e: - if isinstance(e, HTTPException) and hasattr( - mock_logging_obj, "model_call_details" - ): - mock_logging_obj.model_call_details.setdefault( - "metadata", {} - )["guardrail_blocked"] = True + guardrail = TransientErrorGuardrail() - asyncio.create_task( - mock_logging_obj.async_success_handler( - assembled_response, cache_hit=cache_hit, start_time=None, end_time=None - ) + with patch("litellm.callbacks", [guardrail]): + await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails( + captured_data={"model": "gpt-4", "metadata": {}}, + captured_user_api_key_dict=UserAPIKeyAuth(api_key="test"), + captured_logging_obj=mock_logging_obj, + assembled_response=MagicMock(), + cache_hit=False, ) - mock_logging_obj._on_deferred_stream_complete = closure - - resp = await litellm.acompletion( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "hi"}], - mock_response="Hello!", - stream=True, - litellm_logging_obj=mock_logging_obj, - ) - async for _ in resp: - pass - await asyncio.sleep(0) assert mock_logging_obj.model_call_details["metadata"].get( "guardrail_blocked" - ) is not True + ) is not True, "guardrail_blocked must NOT be set for transient errors" @pytest.mark.asyncio async def test_production_closure_integration(self): @@ -685,3 +687,218 @@ class TestDeferredStreamingClosure: "apply_guardrail guardrails must be dispatched through UnifiedLLMGuardrails" assert logged_response is not None, \ "Logging must fire after unified guardrail path" + + @pytest.mark.asyncio + async def test_hooks_receive_merged_guardrail_data(self): + """Hooks must receive guardrail_data (the merged dict from + _check_and_merge_model_level_guardrails), not the original + captured_data. This ensures model-level non-default guardrails + are visible to any inner should_run_guardrail re-checks. + + Uses a deep-copy mock to break the shallow-copy side-effect that + would otherwise mask the bug — verifying the code is explicitly + correct, not correct-by-accident.""" + import copy + + hook_received_data = None + + class InspectingGuardrail(CustomGuardrail): + def __init__(self): + super().__init__( + guardrail_name="inspector", + default_on=True, + event_hook=GuardrailEventHooks.post_call, + ) + + async def async_post_call_success_hook( + self, data: dict, user_api_key_dict: UserAPIKeyAuth, response: Any + ) -> Any: + nonlocal hook_received_data + hook_received_data = data + return response + + mock_logging_obj = MagicMock() + mock_logging_obj.model_call_details = {"metadata": {}} + + async def track_async_success(*args, **kwargs): + pass + + mock_logging_obj.async_success_handler = track_async_success + + guardrail = InspectingGuardrail() + + captured_data = {"model": "gpt-4", "metadata": {"existing_key": "value"}} + + def mock_merge(data, llm_router): + """Return a fully independent dict (deep copy) so the original + captured_data is NOT mutated. This simulates a correct merge + implementation and proves _run_deferred_stream_guardrails uses + the return value, not the original data.""" + merged = copy.deepcopy(data) + merged["metadata"]["guardrails"] = ["model-guardrail"] + merged["_merged_marker"] = True + return merged + + with patch("litellm.callbacks", [guardrail]), \ + patch( + "litellm.proxy.utils._check_and_merge_model_level_guardrails", + side_effect=mock_merge, + ): + await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails( + captured_data=captured_data, + captured_user_api_key_dict=UserAPIKeyAuth(api_key="test"), + captured_logging_obj=mock_logging_obj, + assembled_response=MagicMock(), + cache_hit=False, + ) + + assert hook_received_data is not None, "Guardrail hook must be called" + assert hook_received_data.get("_merged_marker") is True, \ + "Hook must receive guardrail_data (merged), not original captured_data" + assert "model-guardrail" in hook_received_data.get("metadata", {}).get( + "guardrails", [] + ), "Hook data must contain model-level guardrails" + + @pytest.mark.asyncio + async def test_apply_guardrail_path_receives_merged_guardrail_data(self): + """The apply_guardrail path (through UnifiedLLMGuardrails) must also + receive guardrail_data so that the inner should_run_guardrail re-check + inside UnifiedLLMGuardrails sees model-level guardrails. + + This is the specific scenario Greptile flagged: a default_on=False + guardrail configured at the model level would pass the outer gate but + be silently skipped at execution time if captured_data (unmerged) were + passed instead of guardrail_data (merged).""" + import copy + from litellm.types.utils import GenericGuardrailAPIInputs + + unified_received_data = None + + class ModelLevelApplyGuardrail(CustomGuardrail): + def __init__(self): + super().__init__( + guardrail_name="model-apply-guardrail", + default_on=True, + event_hook=GuardrailEventHooks.post_call, + ) + + async def apply_guardrail( + self, inputs, request_data, input_type, logging_obj=None + ) -> GenericGuardrailAPIInputs: + return inputs + + mock_logging_obj = MagicMock() + mock_logging_obj.model_call_details = {"metadata": {}} + + async def track_async_success(*args, **kwargs): + pass + + mock_logging_obj.async_success_handler = track_async_success + + guardrail = ModelLevelApplyGuardrail() + captured_data = {"model": "gpt-4", "metadata": {}} + + def mock_merge(data, llm_router): + merged = copy.deepcopy(data) + merged["metadata"]["guardrails"] = ["model-apply-guardrail"] + merged["_merged_marker"] = True + return merged + + # Capture what UnifiedLLMGuardrails.async_post_call_success_hook receives + original_unified_hook = None + from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import ( + UnifiedLLMGuardrails, + ) + original_unified_hook = UnifiedLLMGuardrails.async_post_call_success_hook + + async def tracking_unified_hook(self, user_api_key_dict, data, response): + nonlocal unified_received_data + unified_received_data = data + return response + + with patch("litellm.callbacks", [guardrail]), \ + patch( + "litellm.proxy.utils._check_and_merge_model_level_guardrails", + side_effect=mock_merge, + ), \ + patch.object( + UnifiedLLMGuardrails, + "async_post_call_success_hook", + tracking_unified_hook, + ): + await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails( + captured_data=captured_data, + captured_user_api_key_dict=UserAPIKeyAuth(api_key="test"), + captured_logging_obj=mock_logging_obj, + assembled_response=MagicMock(), + cache_hit=False, + ) + + assert unified_received_data is not None, \ + "UnifiedLLMGuardrails must be called for apply_guardrail guardrails" + assert unified_received_data.get("_merged_marker") is True, \ + "UnifiedLLMGuardrails must receive guardrail_data (merged), not captured_data" + assert "model-apply-guardrail" in unified_received_data.get( + "metadata", {} + ).get("guardrails", []), \ + "UnifiedLLMGuardrails data must contain model-level guardrails" + + @pytest.mark.asyncio + async def test_multiple_guardrails_all_receive_merged_data(self): + """When multiple guardrails are configured, ALL of them must receive + guardrail_data (merged), not just the first one.""" + import copy + + received_data_per_guardrail = {} + + class TaggedGuardrail(CustomGuardrail): + def __init__(self, name): + super().__init__( + guardrail_name=name, + default_on=True, + event_hook=GuardrailEventHooks.post_call, + ) + + async def async_post_call_success_hook( + self, data: dict, user_api_key_dict: UserAPIKeyAuth, response: Any + ) -> Any: + received_data_per_guardrail[self.guardrail_name] = data + return response + + mock_logging_obj = MagicMock() + mock_logging_obj.model_call_details = {"metadata": {}} + + async def track_async_success(*args, **kwargs): + pass + + mock_logging_obj.async_success_handler = track_async_success + + guardrail_a = TaggedGuardrail("guardrail-a") + guardrail_b = TaggedGuardrail("guardrail-b") + + captured_data = {"model": "gpt-4", "metadata": {}} + + def mock_merge(data, llm_router): + merged = copy.deepcopy(data) + merged["metadata"]["guardrails"] = ["guardrail-a", "guardrail-b"] + merged["_merged_marker"] = True + return merged + + with patch("litellm.callbacks", [guardrail_a, guardrail_b]), \ + patch( + "litellm.proxy.utils._check_and_merge_model_level_guardrails", + side_effect=mock_merge, + ): + await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails( + captured_data=captured_data, + captured_user_api_key_dict=UserAPIKeyAuth(api_key="test"), + captured_logging_obj=mock_logging_obj, + assembled_response=MagicMock(), + cache_hit=False, + ) + + for name in ("guardrail-a", "guardrail-b"): + assert name in received_data_per_guardrail, \ + f"{name} must be called" + assert received_data_per_guardrail[name].get("_merged_marker") is True, \ + f"{name} must receive guardrail_data (merged), not captured_data" From 0057452485d2b12a719e5e66262e22ef676a20e3 Mon Sep 17 00:00:00 2001 From: michelligabriele Date: Thu, 19 Mar 2026 18:00:49 +0100 Subject: [PATCH 16/42] fix(proxy): guard streaming deferred init with try/finally, fix test imports Wrap _run_deferred_stream_guardrails initialization (UnifiedLLMGuardrails constructor and _check_and_merge_model_level_guardrails) in try/finally so logging always fires even if init throws. Prevents silent logging loss on transient errors. Move fastapi.HTTPException import from module-level to local test-function scope. Add test_logging_fires_even_if_guardrail_init_raises to verify the try/finally guard. --- litellm/proxy/common_request_processing.py | 115 +++++++++--------- .../test_deferred_guardrail_logging.py | 42 ++++++- 2 files changed, 101 insertions(+), 56 deletions(-) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 1db8327482d..0e3e89c5313 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1345,76 +1345,81 @@ class ProxyBaseLLMRequestProcessing: from litellm.proxy.utils import _check_and_merge_model_level_guardrails _response = assembled_response - _unified_guardrail = UnifiedLLMGuardrails() - guardrail_data = _check_and_merge_model_level_guardrails( - data=captured_data, llm_router=_global_llm_router - ) - for cb in litellm.callbacks: - if not isinstance(cb, CustomGuardrail): - continue - if not cb.should_run_guardrail( - data=guardrail_data, - event_type=GuardrailEventHooks.post_call, - ): - continue - try: - guardrail_result = None - if "apply_guardrail" in type(cb).__dict__: - guardrail_data["guardrail_to_apply"] = cb - guardrail_result = ( - await _unified_guardrail.async_post_call_success_hook( + try: + _unified_guardrail = UnifiedLLMGuardrails() + guardrail_data = _check_and_merge_model_level_guardrails( + data=captured_data, llm_router=_global_llm_router + ) + for cb in litellm.callbacks: + if not isinstance(cb, CustomGuardrail): + continue + if not cb.should_run_guardrail( + data=guardrail_data, + event_type=GuardrailEventHooks.post_call, + ): + continue + try: + guardrail_result = None + if "apply_guardrail" in type(cb).__dict__: + guardrail_data["guardrail_to_apply"] = cb + guardrail_result = ( + await _unified_guardrail.async_post_call_success_hook( + user_api_key_dict=captured_user_api_key_dict, + data=guardrail_data, + response=_response, + ) + ) + else: + guardrail_result = await cb.async_post_call_success_hook( user_api_key_dict=captured_user_api_key_dict, data=guardrail_data, response=_response, ) + if guardrail_result is not None: + _response = guardrail_result + except Exception as e: + verbose_proxy_logger.exception( + "Error running post-call guardrail %s on streaming response: %s", + getattr(cb, "guardrail_name", type(cb).__name__), + e, ) - else: - guardrail_result = await cb.async_post_call_success_hook( - user_api_key_dict=captured_user_api_key_dict, - data=guardrail_data, - response=_response, + if isinstance(e, HTTPException) and hasattr( + captured_logging_obj, "model_call_details" + ): + captured_logging_obj.model_call_details.setdefault( + "metadata", {} + )["guardrail_blocked"] = True + except Exception as e: + verbose_proxy_logger.exception( + "Error in deferred streaming guardrail initialization: %s", e, + ) + finally: + try: + asyncio.create_task( + captured_logging_obj.async_success_handler( + _response, + cache_hit=cache_hit, + start_time=None, + end_time=None, ) - if guardrail_result is not None: - _response = guardrail_result + ) except Exception as e: verbose_proxy_logger.exception( - "Error running post-call guardrail %s on streaming response: %s", - getattr(cb, "guardrail_name", type(cb).__name__), - e, + "Error in deferred streaming async logging: %s", e, ) - if isinstance(e, HTTPException) and hasattr( - captured_logging_obj, "model_call_details" - ): - captured_logging_obj.model_call_details.setdefault( - "metadata", {} - )["guardrail_blocked"] = True - try: - asyncio.create_task( - captured_logging_obj.async_success_handler( + try: + executor.submit( + captured_logging_obj.success_handler, _response, cache_hit=cache_hit, start_time=None, end_time=None, ) - ) - except Exception as e: - verbose_proxy_logger.exception( - "Error in deferred streaming async logging: %s", e, - ) - - try: - executor.submit( - captured_logging_obj.success_handler, - _response, - cache_hit=cache_hit, - start_time=None, - end_time=None, - ) - except Exception as e: - verbose_proxy_logger.exception( - "Error in deferred streaming sync logging: %s", e, - ) + except Exception as e: + verbose_proxy_logger.exception( + "Error in deferred streaming sync logging: %s", e, + ) async def _handle_llm_api_exception( self, diff --git a/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py b/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py index f5c9eeba1c6..c4d2dce5876 100644 --- a/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py +++ b/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py @@ -20,7 +20,6 @@ from typing import Any from unittest.mock import MagicMock, patch import pytest -from fastapi import HTTPException sys.path.insert(0, os.path.abspath("../../../..")) @@ -233,6 +232,8 @@ async def test_deferred_logging_fires_on_guardrail_exception(): If post_call_success_hook raises (e.g., guardrail blocks content), the deferred logging closure must still fire (via try/finally). """ + from fastapi import HTTPException # noqa: local import for test isolation + enqueue_called = False def mock_enqueue(): @@ -470,6 +471,8 @@ class TestDeferredStreamingClosure: """If a guardrail raises HTTPException, the production _run_deferred_stream_guardrails must still fire logging and set guardrail_blocked in metadata.""" + from fastapi import HTTPException # noqa: local import for test isolation + logging_called = False class BlockingGuardrail(CustomGuardrail): @@ -902,3 +905,40 @@ class TestDeferredStreamingClosure: f"{name} must be called" assert received_data_per_guardrail[name].get("_merged_marker") is True, \ f"{name} must receive guardrail_data (merged), not captured_data" + + @pytest.mark.asyncio + async def test_logging_fires_even_if_guardrail_init_raises(self): + """If _check_and_merge_model_level_guardrails raises during + initialization, logging must still fire via the try/finally guard. + This prevents silent logging loss on transient init errors.""" + logging_called = False + + mock_logging_obj = MagicMock() + mock_logging_obj.model_call_details = {"metadata": {}} + + async def track_async_success(*args, **kwargs): + nonlocal logging_called + logging_called = True + + mock_logging_obj.async_success_handler = track_async_success + + def exploding_merge(data, llm_router): + raise RuntimeError("Simulated init failure") + + with patch( + "litellm.proxy.utils._check_and_merge_model_level_guardrails", + side_effect=exploding_merge, + ): + await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails( + captured_data={"model": "gpt-4", "metadata": {}}, + captured_user_api_key_dict=UserAPIKeyAuth(api_key="test"), + captured_logging_obj=mock_logging_obj, + assembled_response=MagicMock(), + cache_hit=False, + ) + + await asyncio.sleep(0) + await asyncio.sleep(0) + + assert logging_called is True, \ + "Logging must fire even when guardrail initialization raises" From b34231dc95ce686424592d9a063112ef04b103c3 Mon Sep 17 00:00:00 2001 From: michelligabriele Date: Thu, 19 Mar 2026 18:20:42 +0100 Subject: [PATCH 17/42] refactor(proxy): reuse unified_guardrail singleton, rename shadowing variable Reuse the module-level unified_guardrail singleton from proxy/utils.py in _run_deferred_stream_guardrails instead of creating a new instance per call, matching the pattern used by post_call_success_hook. Rename local variable _has_post_call_guardrails to _post_call_guardrails_active to avoid shadowing the static method name. --- litellm/proxy/common_request_processing.py | 15 +++++++-------- 1 file changed, 7 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 0e3e89c5313..1517ee6d9d2 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -931,7 +931,7 @@ class ProxyBaseLLMRequestProcessing: # Defer async logging when post-call guardrails are configured so the # StandardLoggingPayload is built after guardrails write to metadata. # Cache the result to avoid scanning litellm.callbacks twice. - _has_post_call_guardrails = self._has_post_call_guardrails() + _post_call_guardrails_active = self._has_post_call_guardrails() # Non-streaming: defer the create_task in wrapper_async so the # SLP is built after guardrails write to metadata. Streaming @@ -943,7 +943,7 @@ class ProxyBaseLLMRequestProcessing: # so _enqueue_deferred_logging is never stored — the finally # block is a no-op. The CSW path handles this correctly via # _on_deferred_stream_complete, which fires its own logging. - if _has_post_call_guardrails and not self._is_streaming_request( + if _post_call_guardrails_active and not self._is_streaming_request( data=self.data, is_streaming_request=is_streaming_request ): logging_obj._defer_async_logging = True # type: ignore @@ -1057,7 +1057,7 @@ class ProxyBaseLLMRequestProcessing: CustomStreamWrapper, ) - if _has_post_call_guardrails and isinstance( + if _post_call_guardrails_active and isinstance( response, CustomStreamWrapper ): # Intentionally a live reference (not a copy) — mirrors @@ -1338,15 +1338,14 @@ class ProxyBaseLLMRequestProcessing: implementation directly rather than reimplementing the closure. """ from litellm.litellm_core_utils.thread_pool_executor import executor - from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import ( - UnifiedLLMGuardrails, - ) from litellm.proxy.proxy_server import llm_router as _global_llm_router - from litellm.proxy.utils import _check_and_merge_model_level_guardrails + from litellm.proxy.utils import ( + _check_and_merge_model_level_guardrails, + unified_guardrail as _unified_guardrail, + ) _response = assembled_response try: - _unified_guardrail = UnifiedLLMGuardrails() guardrail_data = _check_and_merge_model_level_guardrails( data=captured_data, llm_router=_global_llm_router ) From 6d0763b8ba30a71ba1cd7fec5643f374be85d72c Mon Sep 17 00:00:00 2001 From: Jonathan Barazany Date: Thu, 19 Mar 2026 19:28:05 +0200 Subject: [PATCH 18/42] fix: short-circuit websearch for non-Anthropic providers (github_copilot) For providers like github_copilot that don't natively support web search, Claude Code's search sub-conversations were falling through to the adapter path which strips the web_search tool and has no stream reconversion. Instead of routing search requests through the full LLM pipeline, detect web-search-only requests early (all tools are web_search, simple prompt) and execute the search directly via Tavily/Perplexity, returning a synthetic Anthropic response. No adapter, no backend LLM call needed. Fixes #21733 --- .../websearch_interception/handler.py | 114 ++++++ .../messages/handler.py | 61 ++++ .../test_websearch_short_circuit.py | 332 ++++++++++++++++++ 3 files changed, 507 insertions(+) create mode 100644 tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 2541a0bd7aa..59510dd5d08 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -67,6 +67,120 @@ class WebSearchInterceptionLogger(CustomLogger): self.search_tool_name = search_tool_name self._request_has_websearch = False # Track if current request has web search + async def try_short_circuit_search( + self, + model: str, + messages: List[Dict], + tools: Optional[List[Dict]], + custom_llm_provider: Optional[str], + ) -> Optional[Dict[str, Any]]: + """ + Short-circuit web-search-only requests by executing the search directly. + + Claude Code sends web search as a separate, standalone /v1/messages + request with a simple prompt and only web_search tool(s). For providers + that don't natively support web search (e.g. github_copilot), there is + no need to route this through the backend LLM — we can detect the + pattern, execute the search via Tavily/Perplexity, and return a + synthetic Anthropic response immediately. + + Args: + model: Model name from the request + messages: Messages list from the request + tools: Tools list from the request + custom_llm_provider: Provider name + + Returns: + An AnthropicMessagesResponse dict if short-circuited, or None to + continue normal processing. + """ + if not tools: + return None + + # Check if provider is in enabled list + provider_str = custom_llm_provider or "" + if ( + self.enabled_providers is not None + and provider_str not in self.enabled_providers + ): + return None + + # All tools must be web search tools + if not all(is_web_search_tool(t) for t in tools): + return None + + # Extract search query from the last user message + query = self._extract_search_query(messages) + if not query: + return None + + verbose_logger.debug( + "WebSearchInterception: Short-circuit search detected " + f"(provider={provider_str}, query='{query}')" + ) + + # Execute search + try: + search_result_text = await self._execute_search(query) + except Exception as e: + verbose_logger.error( + f"WebSearchInterception: Short-circuit search failed: {e}" + ) + search_result_text = f"Search failed: {e}" + + # Build synthetic Anthropic response + from uuid import uuid4 + + response: Dict[str, Any] = { + "id": f"msg_{uuid4().hex[:24]}", + "type": "message", + "role": "assistant", + "model": model, + "content": [{"type": "text", "text": search_result_text}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 0, "output_tokens": 0}, + } + + verbose_logger.debug( + "WebSearchInterception: Short-circuit search completed, " + f"returning synthetic response ({len(search_result_text)} chars)" + ) + return response + + @staticmethod + def _extract_search_query(messages: List[Dict]) -> Optional[str]: + """ + Extract the search query from messages. + + Looks at the last user message content for the search query text. + """ + if not messages: + return None + + # Find the last user message + for msg in reversed(messages): + if msg.get("role") != "user": + continue + + content = msg.get("content") + if isinstance(content, str): + return content.strip() or None + + # Handle list-of-blocks content + if isinstance(content, list): + for block in content: + if isinstance(block, dict) and block.get("type") == "text": + text = block.get("text", "").strip() + if text: + return text + elif isinstance(block, str): + text = block.strip() + if text: + return text + + return None + async def async_pre_call_deployment_hook( self, kwargs: Dict[str, Any], call_type: Optional[Any] ) -> Optional[dict]: diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py index 1b5f03ec722..52373a9c327 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py @@ -114,6 +114,54 @@ async def _execute_pre_request_hooks( return request_kwargs +async def _try_websearch_short_circuit( + model: str, + messages: List[Dict], + tools: Optional[List[Dict]], + custom_llm_provider: Optional[str], + stream: Optional[bool], +) -> Optional[Union[AnthropicMessagesResponse, AsyncIterator]]: + """ + Attempt to short-circuit a web-search-only request. + + Claude Code sends web search as a separate, standalone /v1/messages + request. For providers that don't natively support web search (e.g. + github_copilot), we detect this pattern, execute the search via + Tavily/Perplexity, and return a synthetic Anthropic response — bypassing + the backend LLM entirely. + + Returns the synthetic response if short-circuited, or None to continue + normal processing. + """ + if not litellm.callbacks: + return None + + from litellm.integrations.websearch_interception.handler import ( + WebSearchInterceptionLogger, + ) + + for callback in litellm.callbacks: + if not isinstance(callback, WebSearchInterceptionLogger): + continue + + response = await callback.try_short_circuit_search( + model=model, + messages=messages, + tools=tools, + custom_llm_provider=custom_llm_provider, + ) + if response is not None: + if stream: + from litellm.llms.anthropic.experimental_pass_through.messages.fake_stream_iterator import ( + FakeAnthropicMessagesStreamIterator, + ) + + return FakeAnthropicMessagesStreamIterator(response) + return response + + return None + + @client async def anthropic_messages( max_tokens: int, @@ -156,6 +204,19 @@ async def anthropic_messages( # Merge back any other modifications kwargs.update(request_kwargs) + # Short-circuit web-search-only requests: detect the pattern, execute + # search directly via Tavily/Perplexity, and return a synthetic response + # without ever touching the backend LLM or the adapter path. + short_circuit_response = await _try_websearch_short_circuit( + model=model, + messages=messages, + tools=tools, + custom_llm_provider=custom_llm_provider, + stream=stream, + ) + if short_circuit_response is not None: + return short_circuit_response + loop = asyncio.get_event_loop() kwargs["is_async"] = True diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py b/tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py new file mode 100644 index 00000000000..1129ee98efb --- /dev/null +++ b/tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py @@ -0,0 +1,332 @@ +""" +Unit tests for WebSearch Short-Circuit + +Tests the short-circuit path that detects web-search-only /v1/messages requests +and executes the search directly without routing through the backend LLM. +""" + +from unittest.mock import AsyncMock, patch + +import pytest + +from litellm.integrations.websearch_interception.handler import ( + WebSearchInterceptionLogger, +) + + +# --------------------------------------------------------------------------- +# Detection tests +# --------------------------------------------------------------------------- + + +class TestTryShortCircuitSearch: + """Tests for WebSearchInterceptionLogger.try_short_circuit_search""" + + @pytest.mark.asyncio + async def test_short_circuits_single_web_search_tool(self): + """Single web_search_20250305 tool → short-circuit fires""" + logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"]) + + with patch.object( + logger, "_execute_search", new_callable=AsyncMock + ) as mock_search: + mock_search.return_value = "Title: Result\nURL: https://example.com\nSnippet: test" + + result = await logger.try_short_circuit_search( + model="github_copilot/claude-sonnet-4", + messages=[{"role": "user", "content": "Search for Claude Code releases"}], + tools=[{"type": "web_search_20250305", "name": "web_search", "max_uses": 8}], + custom_llm_provider="github_copilot", + ) + + assert result is not None + assert result["type"] == "message" + assert result["role"] == "assistant" + assert result["stop_reason"] == "end_turn" + assert len(result["content"]) == 1 + assert result["content"][0]["type"] == "text" + assert "Result" in result["content"][0]["text"] + mock_search.assert_called_once_with("Search for Claude Code releases") + + @pytest.mark.asyncio + async def test_does_not_short_circuit_mixed_tools(self): + """Mix of web_search and other tools → NOT short-circuited""" + logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"]) + + result = await logger.try_short_circuit_search( + model="github_copilot/claude-sonnet-4", + messages=[{"role": "user", "content": "Do something"}], + tools=[ + {"type": "web_search_20250305", "name": "web_search", "max_uses": 8}, + {"name": "Read", "description": "Read a file", "input_schema": {}}, + ], + custom_llm_provider="github_copilot", + ) + + assert result is None + + @pytest.mark.asyncio + async def test_does_not_short_circuit_no_tools(self): + """No tools → NOT short-circuited""" + logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"]) + + result = await logger.try_short_circuit_search( + model="github_copilot/claude-sonnet-4", + messages=[{"role": "user", "content": "Hello"}], + tools=None, + custom_llm_provider="github_copilot", + ) + + assert result is None + + @pytest.mark.asyncio + async def test_does_not_short_circuit_empty_tools(self): + """Empty tools list → NOT short-circuited""" + logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"]) + + result = await logger.try_short_circuit_search( + model="github_copilot/claude-sonnet-4", + messages=[{"role": "user", "content": "Hello"}], + tools=[], + custom_llm_provider="github_copilot", + ) + + assert result is None + + @pytest.mark.asyncio + async def test_does_not_short_circuit_wrong_provider(self): + """Provider not in enabled_providers → NOT short-circuited""" + logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"]) + + result = await logger.try_short_circuit_search( + model="github_copilot/claude-sonnet-4", + messages=[{"role": "user", "content": "Search for something"}], + tools=[{"type": "web_search_20250305", "name": "web_search", "max_uses": 8}], + custom_llm_provider="github_copilot", + ) + + assert result is None + + @pytest.mark.asyncio + async def test_does_not_short_circuit_no_messages(self): + """Empty messages → NOT short-circuited""" + logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"]) + + result = await logger.try_short_circuit_search( + model="github_copilot/claude-sonnet-4", + messages=[], + tools=[{"type": "web_search_20250305", "name": "web_search", "max_uses": 8}], + custom_llm_provider="github_copilot", + ) + + assert result is None + + @pytest.mark.asyncio + async def test_search_failure_returns_error_text(self): + """Search failure → response with error message, not exception""" + logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"]) + + with patch.object( + logger, "_execute_search", new_callable=AsyncMock + ) as mock_search: + mock_search.side_effect = RuntimeError("Tavily API error") + + result = await logger.try_short_circuit_search( + model="github_copilot/claude-sonnet-4", + messages=[{"role": "user", "content": "Search for something"}], + tools=[{"type": "web_search_20250305", "name": "web_search", "max_uses": 8}], + custom_llm_provider="github_copilot", + ) + + assert result is not None + assert "Search failed" in result["content"][0]["text"] + + @pytest.mark.asyncio + async def test_response_has_valid_structure(self): + """Synthetic response has all required AnthropicMessagesResponse fields""" + logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"]) + + with patch.object( + logger, "_execute_search", new_callable=AsyncMock + ) as mock_search: + mock_search.return_value = "search results here" + + result = await logger.try_short_circuit_search( + model="github_copilot/claude-sonnet-4", + messages=[{"role": "user", "content": "Search query"}], + tools=[{"type": "web_search_20250305", "name": "web_search"}], + custom_llm_provider="github_copilot", + ) + + assert result is not None + # Required fields + assert "id" in result + assert result["id"].startswith("msg_") + assert result["type"] == "message" + assert result["role"] == "assistant" + assert result["model"] == "github_copilot/claude-sonnet-4" + assert result["stop_reason"] == "end_turn" + assert result["stop_sequence"] is None + assert "usage" in result + assert "content" in result + + +# --------------------------------------------------------------------------- +# Query extraction tests +# --------------------------------------------------------------------------- + + +class TestExtractSearchQuery: + """Tests for WebSearchInterceptionLogger._extract_search_query""" + + def test_string_content(self): + messages = [{"role": "user", "content": "Search for Python 3.14 features"}] + assert ( + WebSearchInterceptionLogger._extract_search_query(messages) + == "Search for Python 3.14 features" + ) + + def test_list_content_with_text_block(self): + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "Perform a web search for latest news"}, + ], + } + ] + assert ( + WebSearchInterceptionLogger._extract_search_query(messages) + == "Perform a web search for latest news" + ) + + def test_takes_last_user_message(self): + messages = [ + {"role": "user", "content": "First message"}, + {"role": "assistant", "content": "Response"}, + {"role": "user", "content": "Second message"}, + ] + assert ( + WebSearchInterceptionLogger._extract_search_query(messages) == "Second message" + ) + + def test_empty_messages(self): + assert WebSearchInterceptionLogger._extract_search_query([]) is None + + def test_no_user_messages(self): + messages = [{"role": "assistant", "content": "Hello"}] + assert WebSearchInterceptionLogger._extract_search_query(messages) is None + + def test_empty_content(self): + messages = [{"role": "user", "content": ""}] + assert WebSearchInterceptionLogger._extract_search_query(messages) is None + + +# --------------------------------------------------------------------------- +# Integration with entry point +# --------------------------------------------------------------------------- + + +class TestShortCircuitEntryPoint: + """Tests for _try_websearch_short_circuit in the /v1/messages handler""" + + @pytest.mark.asyncio + async def test_returns_none_when_no_callbacks(self): + """No callbacks configured → returns None""" + from litellm.llms.anthropic.experimental_pass_through.messages.handler import ( + _try_websearch_short_circuit, + ) + + with patch("litellm.callbacks", []): + result = await _try_websearch_short_circuit( + model="test", + messages=[], + tools=[], + custom_llm_provider="github_copilot", + stream=False, + ) + assert result is None + + @pytest.mark.asyncio + async def test_returns_dict_when_not_streaming(self): + """Non-streaming short-circuit → returns dict""" + from litellm.llms.anthropic.experimental_pass_through.messages.handler import ( + _try_websearch_short_circuit, + ) + + logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"]) + with patch.object( + logger, "_execute_search", new_callable=AsyncMock + ) as mock_search: + mock_search.return_value = "results" + with patch("litellm.callbacks", [logger]): + result = await _try_websearch_short_circuit( + model="github_copilot/claude-sonnet-4", + messages=[{"role": "user", "content": "search query"}], + tools=[{"type": "web_search_20250305", "name": "web_search"}], + custom_llm_provider="github_copilot", + stream=False, + ) + + assert isinstance(result, dict) + assert result["content"][0]["text"] == "results" + + @pytest.mark.asyncio + async def test_returns_stream_iterator_when_streaming(self): + """Streaming short-circuit → returns FakeAnthropicMessagesStreamIterator""" + from litellm.llms.anthropic.experimental_pass_through.messages.fake_stream_iterator import ( + FakeAnthropicMessagesStreamIterator, + ) + from litellm.llms.anthropic.experimental_pass_through.messages.handler import ( + _try_websearch_short_circuit, + ) + + logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"]) + with patch.object( + logger, "_execute_search", new_callable=AsyncMock + ) as mock_search: + mock_search.return_value = "streaming results" + with patch("litellm.callbacks", [logger]): + result = await _try_websearch_short_circuit( + model="github_copilot/claude-sonnet-4", + messages=[{"role": "user", "content": "search query"}], + tools=[{"type": "web_search_20250305", "name": "web_search"}], + custom_llm_provider="github_copilot", + stream=True, + ) + + assert isinstance(result, FakeAnthropicMessagesStreamIterator) + + # Verify stream produces valid SSE events + chunks = [] + async for chunk in result: + chunks.append(chunk) + + assert len(chunks) > 0 + # First chunk should be message_start + assert b"event: message_start" in chunks[0] + # Last chunk should be message_stop + assert b"event: message_stop" in chunks[-1] + # Should contain the search results text + all_data = b"".join(chunks) + assert b"streaming results" in all_data + + @pytest.mark.asyncio + async def test_skips_non_websearch_callbacks(self): + """Non-WebSearchInterceptionLogger callbacks are ignored""" + from unittest.mock import MagicMock + + from litellm.llms.anthropic.experimental_pass_through.messages.handler import ( + _try_websearch_short_circuit, + ) + + other_callback = MagicMock() + with patch("litellm.callbacks", [other_callback]): + result = await _try_websearch_short_circuit( + model="test", + messages=[{"role": "user", "content": "search"}], + tools=[{"type": "web_search_20250305", "name": "web_search"}], + custom_llm_provider="github_copilot", + stream=False, + ) + assert result is None From b5a775d54ef6c641860f5a3837c6b8dd06664679 Mon Sep 17 00:00:00 2001 From: Jonathan Barazany Date: Thu, 19 Mar 2026 19:47:13 +0200 Subject: [PATCH 19/42] style: fix Black formatting in test file --- .../test_websearch_short_circuit.py | 28 +++++++++++++------ 1 file changed, 20 insertions(+), 8 deletions(-) diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py b/tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py index 1129ee98efb..68d2276385d 100644 --- a/tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py +++ b/tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py @@ -13,7 +13,6 @@ from litellm.integrations.websearch_interception.handler import ( WebSearchInterceptionLogger, ) - # --------------------------------------------------------------------------- # Detection tests # --------------------------------------------------------------------------- @@ -30,12 +29,18 @@ class TestTryShortCircuitSearch: with patch.object( logger, "_execute_search", new_callable=AsyncMock ) as mock_search: - mock_search.return_value = "Title: Result\nURL: https://example.com\nSnippet: test" + mock_search.return_value = ( + "Title: Result\nURL: https://example.com\nSnippet: test" + ) result = await logger.try_short_circuit_search( model="github_copilot/claude-sonnet-4", - messages=[{"role": "user", "content": "Search for Claude Code releases"}], - tools=[{"type": "web_search_20250305", "name": "web_search", "max_uses": 8}], + messages=[ + {"role": "user", "content": "Search for Claude Code releases"} + ], + tools=[ + {"type": "web_search_20250305", "name": "web_search", "max_uses": 8} + ], custom_llm_provider="github_copilot", ) @@ -101,7 +106,9 @@ class TestTryShortCircuitSearch: result = await logger.try_short_circuit_search( model="github_copilot/claude-sonnet-4", messages=[{"role": "user", "content": "Search for something"}], - tools=[{"type": "web_search_20250305", "name": "web_search", "max_uses": 8}], + tools=[ + {"type": "web_search_20250305", "name": "web_search", "max_uses": 8} + ], custom_llm_provider="github_copilot", ) @@ -115,7 +122,9 @@ class TestTryShortCircuitSearch: result = await logger.try_short_circuit_search( model="github_copilot/claude-sonnet-4", messages=[], - tools=[{"type": "web_search_20250305", "name": "web_search", "max_uses": 8}], + tools=[ + {"type": "web_search_20250305", "name": "web_search", "max_uses": 8} + ], custom_llm_provider="github_copilot", ) @@ -134,7 +143,9 @@ class TestTryShortCircuitSearch: result = await logger.try_short_circuit_search( model="github_copilot/claude-sonnet-4", messages=[{"role": "user", "content": "Search for something"}], - tools=[{"type": "web_search_20250305", "name": "web_search", "max_uses": 8}], + tools=[ + {"type": "web_search_20250305", "name": "web_search", "max_uses": 8} + ], custom_llm_provider="github_copilot", ) @@ -207,7 +218,8 @@ class TestExtractSearchQuery: {"role": "user", "content": "Second message"}, ] assert ( - WebSearchInterceptionLogger._extract_search_query(messages) == "Second message" + WebSearchInterceptionLogger._extract_search_query(messages) + == "Second message" ) def test_empty_messages(self): From 97e17faa51d74d13c452edc2d3e01702f376b123 Mon Sep 17 00:00:00 2001 From: michelligabriele Date: Thu, 19 Mar 2026 18:50:20 +0100 Subject: [PATCH 20/42] fix(proxy): guard lazy imports inside try, clean up orphaned streaming closure Move non-essential lazy imports (llm_router, _check_and_merge, unified_guardrail) inside the try block of _run_deferred_stream_guardrails so that import failures are caught and the finally block still fires logging. Only executor stays outside since the finally block needs it. Add _on_deferred_stream_complete orphan cleanup in the finally block of base_process_llm_request. If an exception propagates after the streaming closure is stored but before a StreamingResponse is returned, the closure is orphaned (CSW never consumes the stream). Detect this via sys.exc_info() and fire logging directly to prevent silent loss. --- litellm/proxy/common_request_processing.py | 51 +++++++++++++++++++--- 1 file changed, 46 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 1517ee6d9d2..45438451cb4 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1,6 +1,7 @@ import asyncio import json import logging +import sys import time import traceback from datetime import datetime @@ -1160,6 +1161,45 @@ class ProxyBaseLLMRequestProcessing: "Error firing deferred logging: %s", e ) + # Streaming cleanup: if an exception is propagating AND the + # deferred streaming closure is still set, no streaming route + # will consume the CSW — the closure is orphaned. Clear it + # and fire logging directly to avoid silent loss. + # + # On normal streaming returns the closure must stay: CSW calls + # it at stream end. sys.exc_info()[1] is None for normal + # returns, non-None only when an exception is propagating. + if sys.exc_info()[1] is not None: + _deferred_fn = getattr( + logging_obj, "_on_deferred_stream_complete", None + ) + if _deferred_fn is not None: + logging_obj._on_deferred_stream_complete = None # type: ignore[attr-defined] + try: + from litellm.litellm_core_utils.thread_pool_executor import ( + executor as _exc, + ) + + asyncio.create_task( + logging_obj.async_success_handler( + response, + cache_hit=None, + start_time=None, + end_time=None, + ) + ) + _exc.submit( + logging_obj.success_handler, + response, + cache_hit=None, + start_time=None, + end_time=None, + ) + except Exception as e: + verbose_proxy_logger.exception( + "Error in orphaned streaming closure cleanup: %s", e + ) + # Always return the client-requested model name (not provider-prefixed internal identifiers) # for OpenAI-compatible responses. if requested_model_from_client: @@ -1338,14 +1378,15 @@ class ProxyBaseLLMRequestProcessing: implementation directly rather than reimplementing the closure. """ from litellm.litellm_core_utils.thread_pool_executor import executor - from litellm.proxy.proxy_server import llm_router as _global_llm_router - from litellm.proxy.utils import ( - _check_and_merge_model_level_guardrails, - unified_guardrail as _unified_guardrail, - ) _response = assembled_response try: + from litellm.proxy.proxy_server import llm_router as _global_llm_router + from litellm.proxy.utils import ( + _check_and_merge_model_level_guardrails, + unified_guardrail as _unified_guardrail, + ) + guardrail_data = _check_and_merge_model_level_guardrails( data=captured_data, llm_router=_global_llm_router ) From 3b129260f557356f0fa621495dab36b659479ef4 Mon Sep 17 00:00:00 2001 From: Jonathan Barazany Date: Thu, 19 Mar 2026 19:52:16 +0200 Subject: [PATCH 21/42] fix: use original_stream for short-circuit, propagate derived provider Addresses Greptile review feedback: - Save original stream flag before pre-request hooks convert it, so streaming callers get SSE events instead of a plain dict - Propagate custom_llm_provider derived inside _execute_pre_request_hooks when it was not explicitly passed by the caller - Add tests covering both scenarios --- .../messages/handler.py | 15 ++++- .../test_websearch_short_circuit.py | 64 +++++++++++++++++++ 2 files changed, 78 insertions(+), 1 deletion(-) diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py index 52373a9c327..dcd9214cf8c 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py @@ -186,6 +186,12 @@ async def anthropic_messages( """ Async: Make llm api request in Anthropic /messages API spec """ + # Save original stream flag before pre-request hooks can convert it. + # The websearch interception hook converts stream=True → stream=False + # for the agentic loop, but the short-circuit path needs to know + # whether the caller originally requested streaming. + original_stream = stream + # Execute pre-request hooks to allow CustomLoggers to modify request request_kwargs = await _execute_pre_request_hooks( model=model, @@ -199,6 +205,11 @@ async def anthropic_messages( # Extract modified parameters tools = request_kwargs.pop("tools", tools) stream = request_kwargs.pop("stream", stream) + # Propagate the provider derived inside pre-request hooks, if not already set + if not custom_llm_provider: + custom_llm_provider = request_kwargs.get("litellm_params", {}).get( + "custom_llm_provider" + ) # Remove litellm_params from kwargs (only needed for hooks) request_kwargs.pop("litellm_params", None) # Merge back any other modifications @@ -207,12 +218,14 @@ async def anthropic_messages( # Short-circuit web-search-only requests: detect the pattern, execute # search directly via Tavily/Perplexity, and return a synthetic response # without ever touching the backend LLM or the adapter path. + # Use original_stream (not the hook-converted stream) so streaming + # callers get SSE events instead of a plain dict. short_circuit_response = await _try_websearch_short_circuit( model=model, messages=messages, tools=tools, custom_llm_provider=custom_llm_provider, - stream=stream, + stream=original_stream, ) if short_circuit_response is not None: return short_circuit_response diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py b/tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py index 68d2276385d..3e2c85074b6 100644 --- a/tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py +++ b/tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py @@ -342,3 +342,67 @@ class TestShortCircuitEntryPoint: stream=False, ) assert result is None + + @pytest.mark.asyncio + async def test_uses_original_stream_not_hook_converted(self): + """Verify that the entry point passes original_stream to the short-circuit. + + The pre-request hook converts stream=True → stream=False for the agentic + loop. The short-circuit must use the ORIGINAL stream value so streaming + callers get SSE events instead of a plain dict. + """ + from litellm.llms.anthropic.experimental_pass_through.messages.fake_stream_iterator import ( + FakeAnthropicMessagesStreamIterator, + ) + from litellm.llms.anthropic.experimental_pass_through.messages.handler import ( + _try_websearch_short_circuit, + ) + + logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"]) + with patch.object( + logger, "_execute_search", new_callable=AsyncMock + ) as mock_search: + mock_search.return_value = "streaming results" + with patch("litellm.callbacks", [logger]): + # Simulate what anthropic_messages() does: original_stream=True + # is passed to the short-circuit, even though the hook would have + # already converted stream to False in request_kwargs. + result = await _try_websearch_short_circuit( + model="github_copilot/claude-sonnet-4", + messages=[{"role": "user", "content": "search query"}], + tools=[{"type": "web_search_20250305", "name": "web_search"}], + custom_llm_provider="github_copilot", + stream=True, # original_stream, NOT the hook-converted value + ) + + # Must return a stream iterator, not a plain dict + assert isinstance(result, FakeAnthropicMessagesStreamIterator) + + @pytest.mark.asyncio + async def test_short_circuits_with_provider_from_model_string(self): + """Provider embedded in model string (custom_llm_provider=None) should + still fire the short-circuit when the caller propagates the derived + provider. + """ + from litellm.llms.anthropic.experimental_pass_through.messages.handler import ( + _try_websearch_short_circuit, + ) + + logger = WebSearchInterceptionLogger(enabled_providers=["github_copilot"]) + with patch.object( + logger, "_execute_search", new_callable=AsyncMock + ) as mock_search: + mock_search.return_value = "results" + with patch("litellm.callbacks", [logger]): + # Simulate the caller having derived custom_llm_provider from + # the model string before calling _try_websearch_short_circuit + result = await _try_websearch_short_circuit( + model="github_copilot/claude-sonnet-4", + messages=[{"role": "user", "content": "search query"}], + tools=[{"type": "web_search_20250305", "name": "web_search"}], + custom_llm_provider="github_copilot", + stream=False, + ) + + assert result is not None + assert result["content"][0]["text"] == "results" From 141ad04955207874307bc0e6aec8a19dd39f44a1 Mon Sep 17 00:00:00 2001 From: Jonathan Barazany Date: Thu, 19 Mar 2026 19:56:42 +0200 Subject: [PATCH 22/42] refactor: reuse get_last_user_message, fix UUID convention, move import - Replace hand-rolled _extract_search_query with existing get_last_user_message from common_utils - Use full UUID (str(uuid.uuid4())) to match codebase convention - Move uuid import to module level per CLAUDE.md --- .../websearch_interception/handler.py | 44 +++-------------- .../test_websearch_short_circuit.py | 47 ------------------- 2 files changed, 7 insertions(+), 84 deletions(-) diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 59510dd5d08..b557d8c0e77 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -8,6 +8,7 @@ server-side using litellm router's search tools. import asyncio import math +import uuid from typing import Any, Dict, List, Optional, Tuple, Union, cast import litellm @@ -110,7 +111,11 @@ class WebSearchInterceptionLogger(CustomLogger): return None # Extract search query from the last user message - query = self._extract_search_query(messages) + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + get_last_user_message, + ) + + query = get_last_user_message(messages) if not query: return None @@ -129,10 +134,8 @@ class WebSearchInterceptionLogger(CustomLogger): search_result_text = f"Search failed: {e}" # Build synthetic Anthropic response - from uuid import uuid4 - response: Dict[str, Any] = { - "id": f"msg_{uuid4().hex[:24]}", + "id": f"msg_{str(uuid.uuid4())}", "type": "message", "role": "assistant", "model": model, @@ -148,39 +151,6 @@ class WebSearchInterceptionLogger(CustomLogger): ) return response - @staticmethod - def _extract_search_query(messages: List[Dict]) -> Optional[str]: - """ - Extract the search query from messages. - - Looks at the last user message content for the search query text. - """ - if not messages: - return None - - # Find the last user message - for msg in reversed(messages): - if msg.get("role") != "user": - continue - - content = msg.get("content") - if isinstance(content, str): - return content.strip() or None - - # Handle list-of-blocks content - if isinstance(content, list): - for block in content: - if isinstance(block, dict) and block.get("type") == "text": - text = block.get("text", "").strip() - if text: - return text - elif isinstance(block, str): - text = block.strip() - if text: - return text - - return None - async def async_pre_call_deployment_hook( self, kwargs: Dict[str, Any], call_type: Optional[Any] ) -> Optional[dict]: diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py b/tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py index 3e2c85074b6..cb90b254e40 100644 --- a/tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py +++ b/tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py @@ -187,53 +187,6 @@ class TestTryShortCircuitSearch: # --------------------------------------------------------------------------- -class TestExtractSearchQuery: - """Tests for WebSearchInterceptionLogger._extract_search_query""" - - def test_string_content(self): - messages = [{"role": "user", "content": "Search for Python 3.14 features"}] - assert ( - WebSearchInterceptionLogger._extract_search_query(messages) - == "Search for Python 3.14 features" - ) - - def test_list_content_with_text_block(self): - messages = [ - { - "role": "user", - "content": [ - {"type": "text", "text": "Perform a web search for latest news"}, - ], - } - ] - assert ( - WebSearchInterceptionLogger._extract_search_query(messages) - == "Perform a web search for latest news" - ) - - def test_takes_last_user_message(self): - messages = [ - {"role": "user", "content": "First message"}, - {"role": "assistant", "content": "Response"}, - {"role": "user", "content": "Second message"}, - ] - assert ( - WebSearchInterceptionLogger._extract_search_query(messages) - == "Second message" - ) - - def test_empty_messages(self): - assert WebSearchInterceptionLogger._extract_search_query([]) is None - - def test_no_user_messages(self): - messages = [{"role": "assistant", "content": "Hello"}] - assert WebSearchInterceptionLogger._extract_search_query(messages) is None - - def test_empty_content(self): - messages = [{"role": "user", "content": ""}] - assert WebSearchInterceptionLogger._extract_search_query(messages) is None - - # --------------------------------------------------------------------------- # Integration with entry point # --------------------------------------------------------------------------- From ee17ef3029573e5fe9242f466e4129a24efb1f75 Mon Sep 17 00:00:00 2001 From: michelligabriele Date: Thu, 19 Mar 2026 19:14:01 +0100 Subject: [PATCH 23/42] fix(proxy): replace sys.exc_info with boolean sentinel for orphan detection Replace sys.exc_info()[1] check with an explicit _exception_raised boolean sentinel. The flag is function-scoped, immune to outer exception context, and only set when an exception actually occurs in base_process_llm_request. This prevents false positives when called from a caller's except block. --- litellm/proxy/common_request_processing.py | 19 +++++++++++-------- 1 file changed, 11 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 45438451cb4..3cb1fa712f5 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1,7 +1,6 @@ import asyncio import json import logging -import sys import time import traceback from datetime import datetime @@ -985,6 +984,7 @@ class ProxyBaseLLMRequestProcessing: response = responses[1] + _exception_raised = False try: hidden_params = getattr(response, "_hidden_params", {}) or {} model_id = self._get_model_id_from_response(hidden_params, self.data) @@ -1145,6 +1145,9 @@ class ProxyBaseLLMRequestProcessing: response = await proxy_logging_obj.post_call_success_hook( data=self.data, user_api_key_dict=user_api_key_dict, response=response ) + except Exception: + _exception_raised = True + raise finally: # Enqueue deferred logging after post-call guardrails have written # guardrail_information to metadata. The finally block ensures @@ -1161,15 +1164,15 @@ class ProxyBaseLLMRequestProcessing: "Error firing deferred logging: %s", e ) - # Streaming cleanup: if an exception is propagating AND the - # deferred streaming closure is still set, no streaming route - # will consume the CSW — the closure is orphaned. Clear it - # and fire logging directly to avoid silent loss. + # Streaming cleanup: if an exception occurred AND the deferred + # streaming closure is still set, no streaming route will + # consume the CSW — the closure is orphaned. Clear it and + # fire logging directly to avoid silent loss. # # On normal streaming returns the closure must stay: CSW calls - # it at stream end. sys.exc_info()[1] is None for normal - # returns, non-None only when an exception is propagating. - if sys.exc_info()[1] is not None: + # it at stream end. _exception_raised is function-scoped and + # immune to outer exception context, avoiding false positives. + if _exception_raised: _deferred_fn = getattr( logging_obj, "_on_deferred_stream_complete", None ) From 573f6b78eac613ad07e88a6f29f3a23a0795c731 Mon Sep 17 00:00:00 2001 From: michelligabriele Date: Thu, 19 Mar 2026 19:29:42 +0100 Subject: [PATCH 24/42] fix(proxy): split orphan cleanup into separate try blocks for resilience Split the single try/except in the _exception_raised cleanup path into separate try blocks for asyncio.create_task and executor.submit, matching the pattern used in _run_deferred_stream_guardrails. If create_task raises, sync logging via executor.submit still fires. --- litellm/proxy/common_request_processing.py | 15 ++++++++++----- 1 file changed, 10 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 3cb1fa712f5..96fd2195712 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1179,10 +1179,6 @@ class ProxyBaseLLMRequestProcessing: if _deferred_fn is not None: logging_obj._on_deferred_stream_complete = None # type: ignore[attr-defined] try: - from litellm.litellm_core_utils.thread_pool_executor import ( - executor as _exc, - ) - asyncio.create_task( logging_obj.async_success_handler( response, @@ -1191,6 +1187,15 @@ class ProxyBaseLLMRequestProcessing: end_time=None, ) ) + except Exception as e: + verbose_proxy_logger.exception( + "Error in orphaned streaming async logging: %s", e + ) + try: + from litellm.litellm_core_utils.thread_pool_executor import ( + executor as _exc, + ) + _exc.submit( logging_obj.success_handler, response, @@ -1200,7 +1205,7 @@ class ProxyBaseLLMRequestProcessing: ) except Exception as e: verbose_proxy_logger.exception( - "Error in orphaned streaming closure cleanup: %s", e + "Error in orphaned streaming sync logging: %s", e ) # Always return the client-requested model name (not provider-prefixed internal identifiers) From 1f04fa2461bf254c0511ead531ae4386bcb579aa Mon Sep 17 00:00:00 2001 From: michelligabriele Date: Thu, 19 Mar 2026 19:50:39 +0100 Subject: [PATCH 25/42] fix(proxy): kill orphaned prisma engine subprocess on failed disconnect --- litellm/proxy/db/prisma_client.py | 44 ++++++++++ litellm/proxy/utils.py | 7 +- .../proxy/db/test_prisma_client.py | 88 ++++++++++++++++++- .../proxy/db/test_prisma_self_heal.py | 35 ++++++++ 4 files changed, 170 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/db/prisma_client.py b/litellm/proxy/db/prisma_client.py index c9c0cfe8f68..fa00d117c81 100644 --- a/litellm/proxy/db/prisma_client.py +++ b/litellm/proxy/db/prisma_client.py @@ -5,6 +5,7 @@ This file contains the PrismaWrapper class, which is used to wrap the Prisma cli import asyncio import os import random +import signal import subprocess import time import urllib @@ -45,6 +46,46 @@ class PrismaWrapper: self._reconnection_lock = asyncio.Lock() self._last_refresh_time: Optional[datetime] = None + def _get_engine_pid(self) -> int: + """Get the PID of the current Prisma engine subprocess, or 0 if unavailable.""" + try: + engine = self._original_prisma._engine + process = getattr(engine, "process", None) if engine is not None else None + if process is not None: + return process.pid + except (AttributeError, TypeError): + pass + return 0 + + @staticmethod + def _kill_engine_process(pid: int) -> None: + """Force-kill an orphaned engine subprocess to prevent DB connection pool leaks. + + Called when disconnect() fails and the old engine process may still be + holding open connections. Sends SIGTERM for graceful shutdown, waits + briefly, then SIGKILL as a backstop. + """ + if pid <= 0: + return + try: + os.kill(pid, signal.SIGTERM) + except (ProcessLookupError, PermissionError, OSError): + return # Already dead or inaccessible + verbose_proxy_logger.warning( + "Sent SIGTERM to orphaned prisma-query-engine PID %s after failed disconnect.", + pid, + ) + # Brief wait for graceful shutdown, then force-kill + time.sleep(0.5) + try: + os.kill(pid, signal.SIGKILL) + verbose_proxy_logger.warning( + "Sent SIGKILL to prisma-query-engine PID %s (did not exit after SIGTERM).", + pid, + ) + except (ProcessLookupError, PermissionError, OSError): + pass # Exited after SIGTERM — expected + def _extract_token_from_db_url(self, db_url: Optional[str]) -> Optional[str]: """ Extract the token (password) from the DATABASE_URL. @@ -179,10 +220,13 @@ class PrismaWrapper: """Disconnect and reconnect the Prisma client with a new database URL.""" from prisma import Prisma # type: ignore + old_engine_pid = self._get_engine_pid() + try: await self._original_prisma.disconnect() except Exception as e: verbose_proxy_logger.warning(f"Failed to disconnect Prisma client: {e}") + self._kill_engine_process(old_engine_pid) if http_client is not None: self._original_prisma = Prisma(http=http_client) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index df527d08af8..e067b4d3175 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -1917,6 +1917,7 @@ class ProxyLogging: original_exception, traceback.format_exc(), ), + daemon=True, ).start() async def post_call_success_hook( @@ -4005,13 +4006,15 @@ class PrismaClient: ) async def _do_direct_reconnect() -> None: + old_pid = self._get_engine_pid() try: await self.db.disconnect() except Exception as disconnect_err: - verbose_proxy_logger.debug( - "Prisma DB disconnect before reconnect failed (ignored): %s", + verbose_proxy_logger.warning( + "Prisma DB disconnect before reconnect failed: %s", disconnect_err, ) + PrismaWrapper._kill_engine_process(old_pid) await self.db.connect() await self.db.query_raw("SELECT 1") diff --git a/tests/test_litellm/proxy/db/test_prisma_client.py b/tests/test_litellm/proxy/db/test_prisma_client.py index 83f07253fc8..9c62c6ffd50 100644 --- a/tests/test_litellm/proxy/db/test_prisma_client.py +++ b/tests/test_litellm/proxy/db/test_prisma_client.py @@ -1,7 +1,8 @@ import json import os +import signal import sys -from unittest.mock import AsyncMock, Mock, patch +from unittest.mock import AsyncMock, MagicMock, Mock, patch import pytest from fastapi.testclient import TestClient @@ -14,6 +15,14 @@ sys.path.insert( from litellm.proxy.db.prisma_client import PrismaWrapper, should_update_prisma_schema +@pytest.fixture(autouse=True) +def mock_prisma_binary(): + """Mock prisma.Prisma to avoid requiring generated Prisma binaries for unit tests.""" + mock_module = MagicMock() + with patch.dict(sys.modules, {"prisma": mock_module}): + yield mock_module + + def test_should_update_prisma_schema(monkeypatch): # CASE 1: Environment variable behavior # When DISABLE_SCHEMA_UPDATE is not set -> should update @@ -73,4 +82,79 @@ async def test_recreate_prisma_client_successful_disconnect(): # Verify that the new client replaced the original assert wrapper._original_prisma != mock_prisma - assert hasattr(wrapper._original_prisma, 'connect') \ No newline at end of file + assert hasattr(wrapper._original_prisma, 'connect') + + +@pytest.mark.asyncio +async def test_recreate_prisma_client_kills_old_engine_on_disconnect_failure( + mock_prisma_binary, +): + """When disconnect() fails, recreate_prisma_client must SIGTERM/SIGKILL the old engine PID.""" + mock_prisma = AsyncMock() + mock_prisma.disconnect.side_effect = Exception("engine hung") + + # Simulate engine subprocess with a known PID + mock_engine = MagicMock() + mock_engine.process.pid = 12345 + mock_prisma._engine = mock_engine + + wrapper = PrismaWrapper(original_prisma=mock_prisma, iam_token_db_auth=False) + + # Configure the mock Prisma constructor + mock_new_prisma = AsyncMock() + mock_prisma_binary.Prisma.return_value = mock_new_prisma + + with ( + patch("os.kill") as mock_kill, + patch("time.sleep"), + ): + await wrapper.recreate_prisma_client("postgresql://new") + + # Verify old engine was killed + mock_kill.assert_any_call(12345, signal.SIGTERM) + # Verify new client was created and connected + mock_new_prisma.connect.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_recreate_prisma_client_skips_kill_on_successful_disconnect( + mock_prisma_binary, +): + """When disconnect() succeeds, no kill should be attempted.""" + mock_prisma = AsyncMock() + mock_prisma.disconnect.return_value = None + + wrapper = PrismaWrapper(original_prisma=mock_prisma, iam_token_db_auth=False) + + mock_new_prisma = AsyncMock() + mock_prisma_binary.Prisma.return_value = mock_new_prisma + + with patch("os.kill") as mock_kill: + await wrapper.recreate_prisma_client("postgresql://new") + + mock_kill.assert_not_called() + mock_new_prisma.connect.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_recreate_prisma_client_handles_missing_engine_pid( + mock_prisma_binary, +): + """When engine PID is unavailable (no _engine attr), kill is skipped gracefully.""" + mock_prisma = AsyncMock() + mock_prisma.disconnect.side_effect = Exception("engine hung") + mock_prisma._engine = None # No engine subprocess + + wrapper = PrismaWrapper(original_prisma=mock_prisma, iam_token_db_auth=False) + + mock_new_prisma = AsyncMock() + mock_prisma_binary.Prisma.return_value = mock_new_prisma + + with ( + patch("os.kill") as mock_kill, + patch("time.sleep"), + ): + await wrapper.recreate_prisma_client("postgresql://new") + + mock_kill.assert_not_called() # PID was 0, kill skipped + mock_new_prisma.connect.assert_awaited_once() \ No newline at end of file diff --git a/tests/test_litellm/proxy/db/test_prisma_self_heal.py b/tests/test_litellm/proxy/db/test_prisma_self_heal.py index 03ad95026d8..dbe1f2113b1 100644 --- a/tests/test_litellm/proxy/db/test_prisma_self_heal.py +++ b/tests/test_litellm/proxy/db/test_prisma_self_heal.py @@ -1,5 +1,6 @@ import asyncio import os +import signal import sys import time from unittest.mock import AsyncMock, MagicMock, patch @@ -279,3 +280,37 @@ async def test_db_health_watchdog_start_stop_lifecycle(mock_proxy_logging): await client.stop_db_health_watchdog_task() assert client._db_health_watchdog_task is None assert dummy_task.cancelled() is True + + +@pytest.mark.asyncio +async def test_lightweight_reconnect_kills_engine_on_disconnect_failure(mock_proxy_logging): + """Lightweight reconnect must kill the old engine PID when disconnect() fails.""" + client = PrismaClient(database_url="mock://test", proxy_logging_obj=mock_proxy_logging) + client.db.disconnect = AsyncMock(side_effect=Exception("disconnect failed")) + client.db.connect = AsyncMock(return_value=None) + client.db.query_raw = AsyncMock(return_value=[{"result": 1}]) + + with ( + patch.object(client, "_get_engine_pid", return_value=9999), + patch("os.kill") as mock_kill, + patch("time.sleep"), + ): + await client._run_reconnect_cycle(timeout_seconds=5.0) + + mock_kill.assert_any_call(9999, signal.SIGTERM) + client.db.connect.assert_awaited_once() + client.db.query_raw.assert_awaited_once_with("SELECT 1") + + +@pytest.mark.asyncio +async def test_lightweight_reconnect_skips_kill_on_successful_disconnect(mock_proxy_logging): + """Lightweight reconnect must NOT kill when disconnect() succeeds.""" + client = PrismaClient(database_url="mock://test", proxy_logging_obj=mock_proxy_logging) + client.db.disconnect = AsyncMock(return_value=None) + client.db.connect = AsyncMock(return_value=None) + client.db.query_raw = AsyncMock(return_value=[{"result": 1}]) + + with patch("os.kill") as mock_kill: + await client._run_reconnect_cycle(timeout_seconds=5.0) + + mock_kill.assert_not_called() From 92b8e1acf8c14863f36050a0bba059b922b8a358 Mon Sep 17 00:00:00 2001 From: michelligabriele Date: Thu, 19 Mar 2026 20:03:07 +0100 Subject: [PATCH 26/42] address greptile review: async sleep, SIGKILL Windows guard, trailing newlines --- litellm/proxy/db/prisma_client.py | 8 ++++---- litellm/proxy/utils.py | 2 +- tests/test_litellm/proxy/db/test_prisma_client.py | 6 +++--- tests/test_litellm/proxy/db/test_prisma_self_heal.py | 2 +- 4 files changed, 9 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/db/prisma_client.py b/litellm/proxy/db/prisma_client.py index fa00d117c81..82ee11a0f42 100644 --- a/litellm/proxy/db/prisma_client.py +++ b/litellm/proxy/db/prisma_client.py @@ -58,7 +58,7 @@ class PrismaWrapper: return 0 @staticmethod - def _kill_engine_process(pid: int) -> None: + async def _kill_engine_process(pid: int) -> None: """Force-kill an orphaned engine subprocess to prevent DB connection pool leaks. Called when disconnect() fails and the old engine process may still be @@ -76,9 +76,9 @@ class PrismaWrapper: pid, ) # Brief wait for graceful shutdown, then force-kill - time.sleep(0.5) + await asyncio.sleep(0.5) try: - os.kill(pid, signal.SIGKILL) + os.kill(pid, getattr(signal, "SIGKILL", signal.SIGTERM)) verbose_proxy_logger.warning( "Sent SIGKILL to prisma-query-engine PID %s (did not exit after SIGTERM).", pid, @@ -226,7 +226,7 @@ class PrismaWrapper: await self._original_prisma.disconnect() except Exception as e: verbose_proxy_logger.warning(f"Failed to disconnect Prisma client: {e}") - self._kill_engine_process(old_engine_pid) + await self._kill_engine_process(old_engine_pid) if http_client is not None: self._original_prisma = Prisma(http=http_client) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index e067b4d3175..e3bf549ce31 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -4014,7 +4014,7 @@ class PrismaClient: "Prisma DB disconnect before reconnect failed: %s", disconnect_err, ) - PrismaWrapper._kill_engine_process(old_pid) + await PrismaWrapper._kill_engine_process(old_pid) await self.db.connect() await self.db.query_raw("SELECT 1") diff --git a/tests/test_litellm/proxy/db/test_prisma_client.py b/tests/test_litellm/proxy/db/test_prisma_client.py index 9c62c6ffd50..f4ef933f219 100644 --- a/tests/test_litellm/proxy/db/test_prisma_client.py +++ b/tests/test_litellm/proxy/db/test_prisma_client.py @@ -106,7 +106,7 @@ async def test_recreate_prisma_client_kills_old_engine_on_disconnect_failure( with ( patch("os.kill") as mock_kill, - patch("time.sleep"), + patch("asyncio.sleep", new_callable=AsyncMock), ): await wrapper.recreate_prisma_client("postgresql://new") @@ -152,9 +152,9 @@ async def test_recreate_prisma_client_handles_missing_engine_pid( with ( patch("os.kill") as mock_kill, - patch("time.sleep"), + patch("asyncio.sleep", new_callable=AsyncMock), ): await wrapper.recreate_prisma_client("postgresql://new") mock_kill.assert_not_called() # PID was 0, kill skipped - mock_new_prisma.connect.assert_awaited_once() \ No newline at end of file + mock_new_prisma.connect.assert_awaited_once() diff --git a/tests/test_litellm/proxy/db/test_prisma_self_heal.py b/tests/test_litellm/proxy/db/test_prisma_self_heal.py index dbe1f2113b1..62fb1b5189c 100644 --- a/tests/test_litellm/proxy/db/test_prisma_self_heal.py +++ b/tests/test_litellm/proxy/db/test_prisma_self_heal.py @@ -293,7 +293,7 @@ async def test_lightweight_reconnect_kills_engine_on_disconnect_failure(mock_pro with ( patch.object(client, "_get_engine_pid", return_value=9999), patch("os.kill") as mock_kill, - patch("time.sleep"), + patch("asyncio.sleep", new_callable=AsyncMock), ): await client._run_reconnect_cycle(timeout_seconds=5.0) From 32cb6f0cd91adfb5554be251490b9b5d5ea60a19 Mon Sep 17 00:00:00 2001 From: Jonathan Barazany Date: Fri, 20 Mar 2026 01:07:20 +0200 Subject: [PATCH 27/42] fix: guard short-circuit against providers with native agentic loop MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Skip short-circuit for providers that have a BaseAnthropicMessagesConfig (bedrock, vertex_ai, azure_ai, anthropic) — they use the agentic loop which includes a follow-up LLM synthesis step. Short-circuiting would return raw search text instead of an LLM-synthesized answer. - Add fallback to litellm.get_llm_provider() for custom_llm_provider derivation when litellm_params is overwritten by kwargs. - Add test for bedrock guard. Addresses Greptile review comments #3 and #4. --- .../websearch_interception/handler.py | 23 +++++++++++++++++++ .../messages/handler.py | 9 +++++++- .../test_websearch_short_circuit.py | 23 +++++++++++++++++++ 3 files changed, 54 insertions(+), 1 deletion(-) diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index b557d8c0e77..34396849e24 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -29,6 +29,7 @@ from litellm.types.integrations.websearch_interception import ( WebSearchInterceptionConfig, ) from litellm.types.utils import LlmProviders +from litellm.utils import ProviderConfigManager class WebSearchInterceptionLogger(CustomLogger): @@ -106,6 +107,28 @@ class WebSearchInterceptionLogger(CustomLogger): ): return None + # Only short-circuit for providers without native Anthropic Messages + # support. Providers that have a BaseAnthropicMessagesConfig (bedrock, + # vertex_ai, azure_ai, anthropic) already use the agentic loop, which + # includes a follow-up LLM call to synthesize the answer from search + # results. Short-circuiting those would skip that synthesis step and + # return raw search text — a regression for existing users. + try: + provider_enum = LlmProviders(provider_str) + anthropic_config = ( + ProviderConfigManager.get_provider_anthropic_messages_config( + model=model, provider=provider_enum + ) + ) + if anthropic_config is not None: + verbose_logger.debug( + f"WebSearchInterception: Skipping short-circuit for {provider_str} " + "(provider has native Anthropic Messages support, using agentic loop)" + ) + return None + except (ValueError, Exception): + pass # unknown provider enum → safe to short-circuit + # All tools must be web search tools if not all(is_web_search_tool(t) for t in tools): return None diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py index dcd9214cf8c..dae2b5da1f8 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py @@ -205,11 +205,18 @@ async def anthropic_messages( # Extract modified parameters tools = request_kwargs.pop("tools", tools) stream = request_kwargs.pop("stream", stream) - # Propagate the provider derived inside pre-request hooks, if not already set + # Propagate the provider derived inside pre-request hooks, if not already set. + # The litellm_params dict may have been overwritten by **kwargs in + # _execute_pre_request_hooks, so fall back to get_llm_provider() if needed. if not custom_llm_provider: custom_llm_provider = request_kwargs.get("litellm_params", {}).get( "custom_llm_provider" ) + if not custom_llm_provider: + try: + _, custom_llm_provider, _, _ = litellm.get_llm_provider(model=model) + except Exception: + pass # Remove litellm_params from kwargs (only needed for hooks) request_kwargs.pop("litellm_params", None) # Merge back any other modifications diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py b/tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py index cb90b254e40..82c1c9839e7 100644 --- a/tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py +++ b/tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py @@ -114,6 +114,29 @@ class TestTryShortCircuitSearch: assert result is None + @pytest.mark.asyncio + async def test_does_not_short_circuit_bedrock(self): + """Bedrock has native agentic loop support → NOT short-circuited. + + Providers with a BaseAnthropicMessagesConfig (bedrock, vertex_ai, etc.) + use the agentic loop which includes a follow-up LLM synthesis step. + The short-circuit must not fire for them. + """ + logger = WebSearchInterceptionLogger( + enabled_providers=["bedrock", "github_copilot"] + ) + + result = await logger.try_short_circuit_search( + model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0", + messages=[{"role": "user", "content": "Search for something"}], + tools=[ + {"type": "web_search_20250305", "name": "web_search", "max_uses": 8} + ], + custom_llm_provider="bedrock", + ) + + assert result is None + @pytest.mark.asyncio async def test_does_not_short_circuit_no_messages(self): """Empty messages → NOT short-circuited""" From be2c679f2dfd7c112c5879bccf51ddf45d75bfea Mon Sep 17 00:00:00 2001 From: Mavik <179817126+themavik@users.noreply.github.com> Date: Fri, 20 Mar 2026 02:44:47 -0400 Subject: [PATCH 28/42] fix: resolve recursion in OVHCloud get_supported_openai_params (#24118) * fix: resolve recursion in OVHCloud get_supported_openai_params (#24111) Root cause: OVHCloudChatConfig.get_supported_openai_params() called get_model_info() which called back into get_supported_openai_params(), causing infinite recursion and falling back to no tools support. Made-with: Cursor * fix: use .get() for TypedDict + don't hardcode function_calling support Made-with: Cursor --------- Co-authored-by: themavik --- litellm/llms/ovhcloud/chat/transformation.py | 14 +++++++++----- 1 file changed, 9 insertions(+), 5 deletions(-) diff --git a/litellm/llms/ovhcloud/chat/transformation.py b/litellm/llms/ovhcloud/chat/transformation.py index e2a9fea7897..84090fafd31 100644 --- a/litellm/llms/ovhcloud/chat/transformation.py +++ b/litellm/llms/ovhcloud/chat/transformation.py @@ -7,7 +7,7 @@ More information on our website: https://endpoints.ai.cloud.ovh.net from typing import Optional, Union, List import httpx -from litellm.utils import ModelResponseStream, get_model_info +from litellm.utils import ModelResponseStream, _get_model_info_helper from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig from litellm._logging import verbose_logger from litellm.llms.ovhcloud.utils import OVHCloudException @@ -28,13 +28,17 @@ class OVHCloudChatConfig(OpenAIGPTConfig): """ supports_function_calling: Optional[bool] = None try: - model_info = get_model_info(model, custom_llm_provider="ovhcloud") - supports_function_calling = model_info.get( - "supports_function_calling", False + model_info = _get_model_info_helper( + model, custom_llm_provider="ovhcloud" ) + supports_function_calling = model_info.get( + "supports_function_calling", None + ) + if supports_function_calling is None: + supports_function_calling = False except Exception as e: verbose_logger.debug(f"Error getting supported OpenAI params: {e}") - pass + supports_function_calling = False optional_params = super().get_supported_openai_params(model) if supports_function_calling is not True: From bc4608e718320c2f3930dcf909514f52d93cd6e9 Mon Sep 17 00:00:00 2001 From: stias Date: Fri, 20 Mar 2026 17:58:20 +0900 Subject: [PATCH 29/42] fix(bedrock): respect api_base and aws_bedrock_runtime_endpoint in count_tokens endpoint The /v1/messages/count_tokens endpoint was hardcoding the Bedrock runtime URL, ignoring api_base and aws_bedrock_runtime_endpoint settings. This aligns it with invoke/converse handlers by using the existing get_runtime_endpoint() method for consistent endpoint resolution. Signed-off-by: stias --- litellm/llms/bedrock/common_utils.py | 1 + litellm/llms/bedrock/count_tokens/handler.py | 9 ++- .../bedrock/count_tokens/transformation.py | 14 +++- .../test_bedrock_token_counter.py | 64 +++++++++++++++++++ 4 files changed, 85 insertions(+), 3 deletions(-) diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index 9666aa68c99..6e659f06d50 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -322,6 +322,7 @@ def init_bedrock_client( endpoint_url=endpoint_url, config=config, verify=ssl_verify, + ) elif aws_profile_name is not None: # uses auth values from AWS profile usually stored in ~/.aws/credentials diff --git a/litellm/llms/bedrock/count_tokens/handler.py b/litellm/llms/bedrock/count_tokens/handler.py index cfd32342d1e..8c227c853cc 100644 --- a/litellm/llms/bedrock/count_tokens/handler.py +++ b/litellm/llms/bedrock/count_tokens/handler.py @@ -64,8 +64,15 @@ class BedrockCountTokensHandler(BedrockCountTokensConfig): verbose_logger.debug(f"Transformed request: {bedrock_request}") # Get endpoint URL using simplified function + api_base = litellm_params.get("api_base", None) + aws_bedrock_runtime_endpoint = litellm_params.get( + "aws_bedrock_runtime_endpoint", None + ) endpoint_url = self.get_bedrock_count_tokens_endpoint( - resolved_model, aws_region_name + model=resolved_model, + aws_region_name=aws_region_name, + api_base=api_base, + aws_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint, ) verbose_logger.debug(f"Making request to: {endpoint_url}") diff --git a/litellm/llms/bedrock/count_tokens/transformation.py b/litellm/llms/bedrock/count_tokens/transformation.py index fe9ab80ced4..a37af131625 100644 --- a/litellm/llms/bedrock/count_tokens/transformation.py +++ b/litellm/llms/bedrock/count_tokens/transformation.py @@ -177,7 +177,11 @@ class BedrockCountTokensConfig(BaseAWSLLM): return {"input": {"invokeModel": {"body": json.dumps(body_data)}}} def get_bedrock_count_tokens_endpoint( - self, model: str, aws_region_name: str + self, + model: str, + aws_region_name: str, + api_base: Optional[str] = None, + aws_bedrock_runtime_endpoint: Optional[str] = None, ) -> str: """ Construct the AWS Bedrock CountTokens API endpoint using existing LiteLLM functions. @@ -185,6 +189,8 @@ class BedrockCountTokensConfig(BaseAWSLLM): Args: model: The resolved model ID from router lookup aws_region_name: AWS region (e.g., "eu-west-1") + api_base: Optional custom API base URL (takes highest priority) + aws_bedrock_runtime_endpoint: Optional custom Bedrock runtime endpoint Returns: Complete endpoint URL for CountTokens API @@ -196,7 +202,11 @@ class BedrockCountTokensConfig(BaseAWSLLM): if model_id.startswith("bedrock/"): model_id = model_id[8:] # Remove "bedrock/" prefix - base_url = f"https://bedrock-runtime.{aws_region_name}.amazonaws.com" + base_url, _ = self.get_runtime_endpoint( + api_base=api_base, + aws_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint, + aws_region_name=aws_region_name, + ) endpoint = f"{base_url}/model/{model_id}/count-tokens" return endpoint diff --git a/tests/litellm_utils_tests/test_bedrock_token_counter.py b/tests/litellm_utils_tests/test_bedrock_token_counter.py index f7c29918820..7239ec0bcfc 100644 --- a/tests/litellm_utils_tests/test_bedrock_token_counter.py +++ b/tests/litellm_utils_tests/test_bedrock_token_counter.py @@ -11,6 +11,7 @@ counting, the test will be skipped. import os import sys from typing import Any, Dict, List +from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -99,3 +100,66 @@ class TestBedrockTokenCounter(BaseTokenCounterTest): assert result.total_tokens > 0, f"Token count should be > 0, got {result.total_tokens}" assert result.tokenizer_type is not None, "tokenizer_type should be set" assert result.error is not True, f"Token counting should not error: {result.error_message}" + + +class TestBedrockCountTokensEndpoint: + """Unit tests for custom endpoint URL resolution in BedrockCountTokensConfig.""" + + def _make_handler(self): + from litellm.llms.bedrock.count_tokens.transformation import ( + BedrockCountTokensConfig, + ) + + return BedrockCountTokensConfig() + + def test_default_endpoint(self): + handler = self._make_handler() + url = handler.get_bedrock_count_tokens_endpoint( + model="amazon.nova-lite-v1:0", + aws_region_name="us-east-1", + ) + assert url == "https://bedrock-runtime.us-east-1.amazonaws.com/model/amazon.nova-lite-v1:0/count-tokens" + + def test_api_base_overrides_default(self): + handler = self._make_handler() + custom_base = "https://vpce-xxx.bedrock-runtime.us-east-1.vpce.amazonaws.com" + url = handler.get_bedrock_count_tokens_endpoint( + model="amazon.nova-lite-v1:0", + aws_region_name="us-east-1", + api_base=custom_base, + ) + assert url == f"{custom_base}/model/amazon.nova-lite-v1:0/count-tokens" + + def test_aws_bedrock_runtime_endpoint_overrides_default(self): + handler = self._make_handler() + custom_endpoint = "https://vpce-yyy.bedrock-runtime.eu-west-1.vpce.amazonaws.com" + url = handler.get_bedrock_count_tokens_endpoint( + model="amazon.nova-lite-v1:0", + aws_region_name="eu-west-1", + aws_bedrock_runtime_endpoint=custom_endpoint, + ) + assert url == f"{custom_endpoint}/model/amazon.nova-lite-v1:0/count-tokens" + + def test_api_base_takes_priority_over_aws_bedrock_runtime_endpoint(self): + handler = self._make_handler() + api_base = "https://api-base.example.com" + runtime_endpoint = "https://runtime-endpoint.example.com" + url = handler.get_bedrock_count_tokens_endpoint( + model="amazon.nova-lite-v1:0", + aws_region_name="us-east-1", + api_base=api_base, + aws_bedrock_runtime_endpoint=runtime_endpoint, + ) + assert url.startswith(api_base) + + def test_env_var_overrides_default(self, monkeypatch): + monkeypatch.setenv( + "AWS_BEDROCK_RUNTIME_ENDPOINT", + "https://env-endpoint.bedrock-runtime.us-west-2.amazonaws.com", + ) + handler = self._make_handler() + url = handler.get_bedrock_count_tokens_endpoint( + model="amazon.nova-lite-v1:0", + aws_region_name="us-west-2", + ) + assert url.startswith("https://env-endpoint.bedrock-runtime.us-west-2.amazonaws.com") From 72902c39c511d33fdc1c13310633f46c37eadfc5 Mon Sep 17 00:00:00 2001 From: Seokjun Yang Date: Fri, 20 Mar 2026 18:09:52 +0900 Subject: [PATCH 30/42] Remove extra newline in common_utils.py --- litellm/llms/bedrock/common_utils.py | 1 - 1 file changed, 1 deletion(-) diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index 6e659f06d50..9666aa68c99 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -322,7 +322,6 @@ def init_bedrock_client( endpoint_url=endpoint_url, config=config, verify=ssl_verify, - ) elif aws_profile_name is not None: # uses auth values from AWS profile usually stored in ~/.aws/credentials From 8da3efdfbe3251dbaf4ba4f9441122e0f6af6c21 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Fri, 20 Mar 2026 18:21:01 +0530 Subject: [PATCH 31/42] Fix code qa and mypy lint issues --- litellm/integrations/websearch_interception/handler.py | 3 ++- .../experimental_pass_through/messages/handler.py | 7 ++++--- litellm/llms/openai/chat/gpt_transformation.py | 3 ++- litellm/proxy/auth/auth_utils.py | 3 ++- litellm/proxy/hooks/parallel_request_limiter.py | 9 +++++++++ 5 files changed, 19 insertions(+), 6 deletions(-) diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 34396849e24..2e5a8734085 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -28,6 +28,7 @@ from litellm.integrations.websearch_interception.transformation import ( from litellm.types.integrations.websearch_interception import ( WebSearchInterceptionConfig, ) +from litellm.types.llms.openai import AllMessageValues from litellm.types.utils import LlmProviders from litellm.utils import ProviderConfigManager @@ -138,7 +139,7 @@ class WebSearchInterceptionLogger(CustomLogger): get_last_user_message, ) - query = get_last_user_message(messages) + query = get_last_user_message(cast(List[AllMessageValues], messages)) if not query: return None diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py index dae2b5da1f8..d117d74e4f7 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py @@ -8,7 +8,7 @@ import asyncio import contextvars from functools import partial -from typing import Any, AsyncIterator, Coroutine, Dict, List, Optional, Union +from typing import Any, AsyncIterator, Coroutine, Dict, List, Optional, Union, cast import litellm from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -151,13 +151,14 @@ async def _try_websearch_short_circuit( custom_llm_provider=custom_llm_provider, ) if response is not None: + anthropic_response = cast(AnthropicMessagesResponse, response) if stream: from litellm.llms.anthropic.experimental_pass_through.messages.fake_stream_iterator import ( FakeAnthropicMessagesStreamIterator, ) - return FakeAnthropicMessagesStreamIterator(response) - return response + return FakeAnthropicMessagesStreamIterator(anthropic_response) + return anthropic_response return None diff --git a/litellm/llms/openai/chat/gpt_transformation.py b/litellm/llms/openai/chat/gpt_transformation.py index 34a23222c2a..c12c6e6ba09 100644 --- a/litellm/llms/openai/chat/gpt_transformation.py +++ b/litellm/llms/openai/chat/gpt_transformation.py @@ -7,6 +7,7 @@ from typing import ( Any, AsyncIterator, Coroutine, + Dict, Iterator, List, Literal, @@ -805,7 +806,7 @@ class OpenAIChatCompletionStreamingHandler(BaseModelResponseIterator): choices = chunk.get("choices", []) choices = self._map_reasoning_to_reasoning_content(choices) - kwargs = { + kwargs: Dict[str, Any] = { "id": chunk.get("id"), "object": "chat.completion.chunk", "created": chunk.get("created"), diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 7d3427ed4c1..d2f8320668e 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -560,7 +560,8 @@ def _get_deployment_default_limit(model_name: str, field: str) -> Optional[int]: raw = deployment.get("litellm_params", {}).get(field) if raw is not None: try: - limits.append(int(raw)) + if isinstance(raw, (int, float, str, bytes, bytearray)): + limits.append(int(raw)) except (ValueError, TypeError): pass return min(limits) if limits else None diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index 55e89e02d67..fefc6c8af9c 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -478,6 +478,15 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): kwargs["litellm_params"]["metadata"].get("user_api_key_metadata", {}) or {} ) + user_api_key_team_metadata = kwargs["litellm_params"]["metadata"].get( + "user_api_key_team_metadata", None + ) + user_api_key_dict = UserAPIKeyAuth( + api_key=user_api_key, + metadata=user_api_key_metadata, + model_max_budget=user_api_key_model_max_budget, + team_metadata=user_api_key_team_metadata, + ) # ------------ # Setup values From eb733702fcd5a1c3ddd57ffc2416005f5d51cd8a Mon Sep 17 00:00:00 2001 From: Seokjun Yang Date: Fri, 20 Mar 2026 22:21:15 +0900 Subject: [PATCH 32/42] Update tests/litellm_utils_tests/test_bedrock_token_counter.py Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> --- tests/litellm_utils_tests/test_bedrock_token_counter.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/litellm_utils_tests/test_bedrock_token_counter.py b/tests/litellm_utils_tests/test_bedrock_token_counter.py index 7239ec0bcfc..b0e37914af9 100644 --- a/tests/litellm_utils_tests/test_bedrock_token_counter.py +++ b/tests/litellm_utils_tests/test_bedrock_token_counter.py @@ -150,7 +150,7 @@ class TestBedrockCountTokensEndpoint: api_base=api_base, aws_bedrock_runtime_endpoint=runtime_endpoint, ) - assert url.startswith(api_base) + assert url == f"{api_base}/model/amazon.nova-lite-v1:0/count-tokens" def test_env_var_overrides_default(self, monkeypatch): monkeypatch.setenv( From d3afaf613dc0865ad92f374f41979cb3bee1c9e7 Mon Sep 17 00:00:00 2001 From: Seokjun Yang Date: Fri, 20 Mar 2026 22:21:22 +0900 Subject: [PATCH 33/42] Update tests/litellm_utils_tests/test_bedrock_token_counter.py Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> --- tests/litellm_utils_tests/test_bedrock_token_counter.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/litellm_utils_tests/test_bedrock_token_counter.py b/tests/litellm_utils_tests/test_bedrock_token_counter.py index b0e37914af9..abc45b03d6c 100644 --- a/tests/litellm_utils_tests/test_bedrock_token_counter.py +++ b/tests/litellm_utils_tests/test_bedrock_token_counter.py @@ -11,7 +11,7 @@ counting, the test will be skipped. import os import sys from typing import Any, Dict, List -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import patch import pytest From f36a59d196e42670007de37b4265b88d96abd953 Mon Sep 17 00:00:00 2001 From: Milan Date: Fri, 20 Mar 2026 16:27:41 +0000 Subject: [PATCH 34/42] fix(logging): merge hidden_params into metadata for streaming completions Non-streaming paths call _process_hidden_params_and_response_cost; streaming assembles the full response later and skipped that, so litellm_params.metadata lacked hidden_params (e.g. response_cost for OTEL/OpenSearch). - Add _merge_hidden_params_from_response_into_metadata and call it from success_handler and async_success_handler after cost is set, before _build_standard_logging_payload. - Unit tests for merge helper. Tests: pytest tests/test_litellm/litellm_core_utils/test_litellm_logging.py Made-with: Cursor --- litellm/litellm_core_utils/litellm_logging.py | 29 +++++++++++ .../test_litellm_logging.py | 51 +++++++++++++++++++ 2 files changed, 80 insertions(+) diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index fea139a64b4..56f7f305dca 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -1686,6 +1686,28 @@ class Logging(LiteLLMLoggingBaseClass): ) return logging_result + def _merge_hidden_params_from_response_into_metadata(self, logging_result: Any) -> None: + """ + Copy response._hidden_params into litellm_params.metadata['hidden_params']. + + Non-streaming success uses _process_hidden_params_and_response_cost (skipped when + stream=True). Streaming assembles the full response later; without this merge, + OTEL/callbacks that read metadata.hidden_params miss cost-related fields. + """ + if logging_result is None: + return + hidden_params = getattr(logging_result, "_hidden_params", None) + if not hidden_params: + return + if self.model_call_details.get("litellm_params") is None: + return + self.model_call_details["litellm_params"].setdefault("metadata", {}) + if self.model_call_details["litellm_params"]["metadata"] is None: + self.model_call_details["litellm_params"]["metadata"] = {} + self.model_call_details["litellm_params"]["metadata"]["hidden_params"] = getattr( + logging_result, "_hidden_params", {} + ) + def _process_hidden_params_and_response_cost( self, logging_result, @@ -2010,6 +2032,9 @@ class Logging(LiteLLMLoggingBaseClass): self.model_call_details[ "response_cost" ] = self._response_cost_calculator(result=complete_streaming_response) + self._merge_hidden_params_from_response_into_metadata( + complete_streaming_response + ) ## STANDARDIZED LOGGING PAYLOAD self.model_call_details[ "standard_logging_object" @@ -2545,6 +2570,10 @@ class Logging(LiteLLMLoggingBaseClass): ) self.model_call_details["response_cost"] = None + self._merge_hidden_params_from_response_into_metadata( + complete_streaming_response + ) + ## STANDARDIZED LOGGING PAYLOAD self.model_call_details[ "standard_logging_object" diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index 0f950f6da77..e5fb0ebdf6e 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -2248,3 +2248,54 @@ def test_failure_handler_runs_sync_callbacks_for_non_pass_through_requests( ) dummy_logger.log_failure_event.assert_called_once() + + +def test_merge_hidden_params_from_response_into_metadata_populates_metadata(): + """Streaming completion path should mirror non-stream: metadata.hidden_params from response.""" + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + + logging_obj = LiteLLMLoggingObj( + model="gpt-4o-mini", + messages=[{"role": "user", "content": "hi"}], + stream=True, + call_type="acompletion", + start_time=time.time(), + litellm_call_id="merge-hp-test", + function_id="merge-hp-fn", + ) + logging_obj.model_call_details = { + "litellm_params": {"metadata": {}}, + } + + class _Resp: + _hidden_params = {"response_cost": 0.001, "model_id": "mid-test"} + + logging_obj._merge_hidden_params_from_response_into_metadata(_Resp()) + meta = logging_obj.model_call_details["litellm_params"]["metadata"] + assert meta["hidden_params"]["response_cost"] == 0.001 + assert meta["hidden_params"]["model_id"] == "mid-test" + + +def test_merge_hidden_params_from_response_into_metadata_no_op_when_empty(): + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + + logging_obj = LiteLLMLoggingObj( + model="gpt-4o-mini", + messages=[{"role": "user", "content": "hi"}], + stream=True, + call_type="acompletion", + start_time=time.time(), + litellm_call_id="merge-hp-empty", + function_id="merge-hp-empty-fn", + ) + logging_obj.model_call_details = { + "litellm_params": {"metadata": {"existing": True}}, + } + + class _NoHp: + _hidden_params = {} + + logging_obj._merge_hidden_params_from_response_into_metadata(_NoHp()) + assert "hidden_params" not in logging_obj.model_call_details["litellm_params"][ + "metadata" + ] From 589c6cdad00dc4a83fa5adf97cd2d8279739d44f Mon Sep 17 00:00:00 2001 From: Christopher Baer <30447746+christopherbaer@users.noreply.github.com> Date: Fri, 20 Mar 2026 10:02:22 -0700 Subject: [PATCH 35/42] fix(gemini-embeddings): convert task_type to camelCase taskType for Gemini API (#24191) The Gemini REST API documents the embedding task type parameter as camelCase `taskType`. The existing transformation functions convert `dimensions` to `outputDimensionality` but miss the parallel `task_type` to `taskType` conversion. This adds that conversion to both `transform_openai_input_gemini_content` (batchEmbedContents path) and `transform_openai_input_gemini_embed_content` (embedContent path). Fixes #24190 --- .../batch_embed_content_transformation.py | 4 ++ .../vertex_ai/test_gemini_batch_embeddings.py | 39 +++++++++++++++++++ 2 files changed, 43 insertions(+) diff --git a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py index 0f6d85525d9..08831a8215f 100644 --- a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py +++ b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py @@ -152,6 +152,8 @@ def transform_openai_input_gemini_content( gemini_params = optional_params.copy() if "dimensions" in gemini_params: gemini_params["outputDimensionality"] = gemini_params.pop("dimensions") + if "task_type" in gemini_params: + gemini_params["taskType"] = gemini_params.pop("task_type") requests: List[EmbedContentRequest] = [] if isinstance(input, str): @@ -196,6 +198,8 @@ def transform_openai_input_gemini_embed_content( gemini_params = optional_params.copy() if "dimensions" in gemini_params: gemini_params["outputDimensionality"] = gemini_params.pop("dimensions") + if "task_type" in gemini_params: + gemini_params["taskType"] = gemini_params.pop("task_type") input_list = [input] if isinstance(input, str) else input parts: List[PartType] = [] diff --git a/tests/litellm/llms/vertex_ai/test_gemini_batch_embeddings.py b/tests/litellm/llms/vertex_ai/test_gemini_batch_embeddings.py index 1ed1de01b5f..a8e427d3bc1 100644 --- a/tests/litellm/llms/vertex_ai/test_gemini_batch_embeddings.py +++ b/tests/litellm/llms/vertex_ai/test_gemini_batch_embeddings.py @@ -22,6 +22,7 @@ from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation _is_multimodal_input, _parse_data_url, process_embed_content_response, + transform_openai_input_gemini_content, transform_openai_input_gemini_embed_content, ) from litellm.types.utils import EmbeddingResponse @@ -396,6 +397,44 @@ def test_transform_with_optional_params(): assert result["taskType"] == "SEMANTIC_SIMILARITY" +def test_task_type_mapped_to_camel_case_batch(): + """Test that snake_case task_type is converted to camelCase taskType for batchEmbedContents.""" + result = transform_openai_input_gemini_content( + input="test text", + model="text-embedding-004", + optional_params={"task_type": "RETRIEVAL_DOCUMENT"}, + ) + for request in result["requests"]: + assert "taskType" in request + assert request["taskType"] == "RETRIEVAL_DOCUMENT" + assert "task_type" not in request + + +def test_task_type_mapped_to_camel_case_embed_content(): + """Test that snake_case task_type is converted to camelCase taskType for embedContent.""" + result = transform_openai_input_gemini_embed_content( + input=["test text"], + model="gemini-embedding-2-preview", + optional_params={"task_type": "RETRIEVAL_DOCUMENT"}, + resolved_files=None, + ) + assert "taskType" in result + assert result["taskType"] == "RETRIEVAL_DOCUMENT" + assert "task_type" not in result + + +def test_task_type_camel_case_passthrough(): + """Test that camelCase taskType passed directly is preserved.""" + result = transform_openai_input_gemini_embed_content( + input=["test text"], + model="gemini-embedding-2-preview", + optional_params={"taskType": "SEMANTIC_SIMILARITY"}, + resolved_files=None, + ) + assert result["taskType"] == "SEMANTIC_SIMILARITY" + assert "task_type" not in result + + def test_dimensions_mapped_to_output_dimensionality(): """Test that OpenAI 'dimensions' param is mapped to Gemini 'outputDimensionality'.""" input_data = ["test text"] From 714c1b80e10e7ecd73b241a8a37e98dca893a692 Mon Sep 17 00:00:00 2001 From: Jayachander Reddy kandakatla <49528664+Jayachander123@users.noreply.github.com> Date: Fri, 20 Mar 2026 13:25:53 -0500 Subject: [PATCH 36/42] docs(pricing): add official source links for Azure DeepSeek & Cohere models (#20181) Added 'source' keys to Azure DeepSeek v3.2(Standard & Speciale) and Cohere Rerank 4.0 (Pro & Fast) entries for pricing verification. --- model_prices_and_context_window.json | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 879dd42be47..bbf9f6d9dc8 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -6152,7 +6152,8 @@ "max_query_tokens": 4096, "max_tokens": 32768, "mode": "rerank", - "output_cost_per_token": 0.0 + "output_cost_per_token": 0.0, + "source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-cohere-rerank-4-0-in-microsoft-foundry/4477076" }, "azure_ai/cohere-rerank-v4.0-fast": { "input_cost_per_query": 0.002, @@ -6163,7 +6164,8 @@ "max_query_tokens": 4096, "max_tokens": 32768, "mode": "rerank", - "output_cost_per_token": 0.0 + "output_cost_per_token": 0.0, + "source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-cohere-rerank-4-0-in-microsoft-foundry/4477076" }, "azure_ai/deepseek-v3.2": { "input_cost_per_token": 5.8e-07, @@ -6173,6 +6175,7 @@ "max_tokens": 163840, "mode": "chat", "output_cost_per_token": 1.68e-06, + "source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-deepseek-v3-2-and-deepseek-v3-2-speciale-in-microsoft-foundry/4477549", "supports_assistant_prefill": true, "supports_function_calling": true, "supports_prompt_caching": true, @@ -6187,6 +6190,7 @@ "max_tokens": 163840, "mode": "chat", "output_cost_per_token": 1.68e-06, + "source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-deepseek-v3-2-and-deepseek-v3-2-speciale-in-microsoft-foundry/4477549", "supports_assistant_prefill": true, "supports_function_calling": true, "supports_prompt_caching": true, From 00dd9844158c06da595a863c3c18552612068079 Mon Sep 17 00:00:00 2001 From: Geoffray Viossat <4362195+gvioss@users.noreply.github.com> Date: Fri, 20 Mar 2026 19:32:15 +0100 Subject: [PATCH 37/42] fix(whisper): correct output_cost_per_second pricing and cost calculation (#23842) - Set output_cost_per_second to 0.0 (was 0.0001) for whisper-1 and azure/whisper-1: transcription is billed on input duration only, not output duration - Fix cost_per_second() in openai/cost_calculation.py: change elif to if so input_cost_per_second is evaluated independently of output_cost_per_second, and remove the erroneous completion_cost = 0.0 assignment that masked any previously-set output cost - Add TestCostPerSecondArithmetic unit tests covering both cost fields, the None-guard, and zero-duration edge case Co-authored-by: Claude Sonnet 4.6 --- litellm/llms/openai/cost_calculation.py | 3 +- ...odel_prices_and_context_window_backup.json | 4 +- model_prices_and_context_window.json | 4 +- tests/test_litellm/test_cost_calculator.py | 52 +++++++++++++++++++ 4 files changed, 57 insertions(+), 6 deletions(-) diff --git a/litellm/llms/openai/cost_calculation.py b/litellm/llms/openai/cost_calculation.py index ac1e4a6b08f..d5077e25a43 100644 --- a/litellm/llms/openai/cost_calculation.py +++ b/litellm/llms/openai/cost_calculation.py @@ -114,7 +114,7 @@ def cost_per_second( ) ## COST PER SECOND ## completion_cost = model_info["output_cost_per_second"] * duration - elif ( + if ( "input_cost_per_second" in model_info and model_info["input_cost_per_second"] is not None ): @@ -123,7 +123,6 @@ def cost_per_second( ) ## COST PER SECOND ## prompt_cost = model_info["input_cost_per_second"] * duration - completion_cost = 0.0 return prompt_cost, completion_cost diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 879dd42be47..595cdf3e138 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -5725,7 +5725,7 @@ "input_cost_per_second": 0.0001, "litellm_provider": "azure", "mode": "audio_transcription", - "output_cost_per_second": 0.0001 + "output_cost_per_second": 0.0 }, "azure_ai/Cohere-embed-v3-english": { "input_cost_per_token": 1e-07, @@ -31996,7 +31996,7 @@ "input_cost_per_second": 0.0001, "litellm_provider": "openai", "mode": "audio_transcription", - "output_cost_per_second": 0.0001, + "output_cost_per_second": 0.0, "supported_endpoints": [ "/v1/audio/transcriptions" ] diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index bbf9f6d9dc8..c1ba4d88d4d 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -5725,7 +5725,7 @@ "input_cost_per_second": 0.0001, "litellm_provider": "azure", "mode": "audio_transcription", - "output_cost_per_second": 0.0001 + "output_cost_per_second": 0.0 }, "azure_ai/Cohere-embed-v3-english": { "input_cost_per_token": 1e-07, @@ -32000,7 +32000,7 @@ "input_cost_per_second": 0.0001, "litellm_provider": "openai", "mode": "audio_transcription", - "output_cost_per_second": 0.0001, + "output_cost_per_second": 0.0, "supported_endpoints": [ "/v1/audio/transcriptions" ] diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index 8f5c3ece0ca..88c384bcaec 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -1,5 +1,6 @@ import os import sys +from unittest.mock import patch import pytest @@ -16,6 +17,7 @@ from litellm.cost_calculator import ( handle_realtime_stream_cost_calculation, response_cost_calculator, ) +from litellm.llms.openai.cost_calculation import cost_per_second from litellm.types.llms.openai import OpenAIRealtimeStreamList from litellm.types.utils import ModelResponse, PromptTokensDetailsWrapper, Usage from litellm.utils import TranscriptionResponse @@ -1970,3 +1972,53 @@ def test_additional_costs_only_for_azure_ai(): completion_tokens=50, ) assert result is None, "Vertex AI should have no additional costs" + + +class TestCostPerSecondArithmetic: + """Unit tests for the cost_per_second arithmetic itself. + + The tests in TestCostCalculatorReadsDurationFromHiddenParams mock + openai_cost_per_second entirely — they verify that the right duration + is forwarded, but never that the math inside cost_per_second is correct. + These tests cover the arithmetic directly. + """ + + def test_input_cost_per_second_only(self): + """input_cost_per_second * duration = prompt_cost; completion_cost = 0.""" + with patch("litellm.llms.openai.cost_calculation.get_model_info") as mock_info: + mock_info.return_value = { + "input_cost_per_second": 0.0001, + "output_cost_per_second": 0.0, + } + prompt_cost, completion_cost_val = cost_per_second( + model="whisper-1", custom_llm_provider="openai", duration=60.0 + ) + assert prompt_cost == pytest.approx(0.006) + assert completion_cost_val == 0.0 + + def test_both_input_and_output_cost_per_second(self): + """When both fields are set, they are applied independently.""" + with patch("litellm.llms.openai.cost_calculation.get_model_info") as mock_info: + mock_info.return_value = { + "input_cost_per_second": 0.002, + "output_cost_per_second": 0.003, + } + prompt_cost, completion_cost_val = cost_per_second( + model="some-model", custom_llm_provider="openai", duration=10.0 + ) + assert prompt_cost == pytest.approx(0.02) + assert completion_cost_val == pytest.approx(0.03) + + def test_zero_duration_returns_zero_cost(self): + """A zero-duration transcription must cost nothing regardless of the rate.""" + with patch("litellm.llms.openai.cost_calculation.get_model_info") as mock_info: + mock_info.return_value = { + "input_cost_per_second": 0.0001, + "output_cost_per_second": 0.0, + } + prompt_cost, completion_cost_val = cost_per_second( + model="whisper-1", custom_llm_provider="openai", duration=0.0 + ) + assert prompt_cost == 0.0 + assert completion_cost_val == 0.0 + From 78139472a166b60d51b2b4298dc829ccde393603 Mon Sep 17 00:00:00 2001 From: BillionToken Date: Sat, 21 Mar 2026 02:39:17 +0800 Subject: [PATCH 38/42] fix(moonshot): preserve reasoning_content on Pydantic Message objects in multi-turn tool calls (#23828) * fix(moonshot): preserve reasoning_content on Pydantic Message objects in multi-turn tool calls The condition 'reasoning_content not in msg' doesn't work correctly for Pydantic Message objects because they don't support the 'in' operator like dicts do. This caused reasoning_content to be stripped from assistant messages in multi-turn conversation history. Changed the condition to use msg.get('reasoning_content') instead, which works correctly for both dicts and Pydantic models. Fixes #23765 * added newline eof Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> * Update tests/test_litellm/llms/moonshot/test_moonshot_chat_transformation.py Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> * Simplify assertions in test_moonshot_chat_transformation Removed redundant assertions for non-assistant messages. --------- Co-authored-by: BillionClaw <267901332+BillionClaw@users.noreply.github.com> Co-authored-by: Aarish Alam Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> --- litellm/llms/moonshot/chat/transformation.py | 8 ++- .../test_moonshot_chat_transformation.py | 70 ++++++++++++++++++- 2 files changed, 74 insertions(+), 4 deletions(-) diff --git a/litellm/llms/moonshot/chat/transformation.py b/litellm/llms/moonshot/chat/transformation.py index 24f852c28ba..c97bd6c4e12 100644 --- a/litellm/llms/moonshot/chat/transformation.py +++ b/litellm/llms/moonshot/chat/transformation.py @@ -155,9 +155,11 @@ class MoonshotChatConfig(OpenAIGPTConfig): message that contains tool_calls (multi-turn tool-calling flows). For each such message that is missing the field: - 1. Promote provider_specific_fields["reasoning_content"] if present and non-empty + 1. Check if reasoning_content exists at the top level (for Pydantic models + that have the attribute but don't support 'in' operator) + 2. Promote provider_specific_fields["reasoning_content"] if present and non-empty (this is where LiteLLM stores it from a previous response) - 2. Otherwise inject a single space — the minimum value the API accepts + 3. Otherwise inject a single space — the minimum value the API accepts Messages that already carry the field, or are not assistant/tool-call messages, are appended as-is (no copy made). """ @@ -166,7 +168,7 @@ class MoonshotChatConfig(OpenAIGPTConfig): if ( msg.get("role") == "assistant" and msg.get("tool_calls") - and "reasoning_content" not in msg + and not msg.get("reasoning_content") # Check using .get() which works for both dicts and Pydantic models ): patched = dict(cast(dict, msg)) provider_fields = patched.get("provider_specific_fields") or {} diff --git a/tests/test_litellm/llms/moonshot/test_moonshot_chat_transformation.py b/tests/test_litellm/llms/moonshot/test_moonshot_chat_transformation.py index c557fb395f9..f7e07ce8d97 100644 --- a/tests/test_litellm/llms/moonshot/test_moonshot_chat_transformation.py +++ b/tests/test_litellm/llms/moonshot/test_moonshot_chat_transformation.py @@ -550,4 +550,72 @@ class TestMoonshotConfig: # reasoning_content must not have been injected for msg in result["messages"]: - assert "reasoning_content" not in msg \ No newline at end of file + assert "reasoning_content" not in msg + + def test_reasoning_content_preserved_on_pydantic_message_object(self): + """reasoning_content on Pydantic Message objects is preserved (not overwritten with placeholder). + + Regression test for: https://github.com/BerriAI/litellm/issues/23765 + The issue was that 'reasoning_content' in msg doesn't work for Pydantic models + because they don't support the 'in' operator the same way as dicts. + """ + from litellm.types.utils import Message + + config = MoonshotChatConfig() + + # Create a Pydantic Message object with reasoning_content (as would come from API response) + message_with_reasoning = Message( + role="assistant", + content=None, + reasoning_content="User wants weather", + tool_calls=[ + {"id": "call_1", "type": "function", "function": {"name": "fn", "arguments": "{}"}} + ], + ) + + messages = [message_with_reasoning] + + result = config.fill_reasoning_content(messages) + + # reasoning_content should be preserved, not replaced with placeholder + assert result[0].get("reasoning_content") == "User wants weather" + + def test_reasoning_content_preserved_in_multi_turn_flow(self): + """reasoning_content is preserved through multi-turn conversation flow. + + This tests the complete flow: API response -> Message object -> dict -> fill_reasoning_content + """ + from litellm.types.utils import Message + from litellm.utils import convert_to_dict + + config = MoonshotChatConfig() + + # Simulate API response with reasoning_content + api_response = { + "role": "assistant", + "content": None, + "reasoning_content": "Planning to call weather tool", + "tool_calls": [ + {"id": "call_1", "type": "function", "function": {"name": "get_weather", "arguments": '{}'}} + ], + } + + # Convert to Message object (as LiteLLM does) + message_obj = Message(**api_response) + + # Convert back to dict (when building next request) + message_dict = convert_to_dict(message_obj) + + # Build multi-turn conversation + messages = [ + {"role": "user", "content": "What's the weather?"}, + message_dict, + {"role": "tool", "tool_call_id": "call_1", "content": '{"temp": 72}'}, + {"role": "user", "content": "Thanks!"}, + ] + + # Apply fill_reasoning_content + result = config.fill_reasoning_content(messages) + + # reasoning_content should be preserved in the assistant message + assert result[1].get("reasoning_content") == "Planning to call weather tool" From 29ab11a9c2840ef2296d40d49783331bdb7e0d8b Mon Sep 17 00:00:00 2001 From: Chesars Date: Fri, 20 Mar 2026 23:23:28 -0300 Subject: [PATCH 39/42] fix(types): add CacheControlToolConfigInjectionPoint to union type --- .../integrations/anthropic_cache_control_hook.py | 11 ++++++++++- 1 file changed, 10 insertions(+), 1 deletion(-) diff --git a/litellm/types/integrations/anthropic_cache_control_hook.py b/litellm/types/integrations/anthropic_cache_control_hook.py index 6e859f10187..83e5a9e7f01 100644 --- a/litellm/types/integrations/anthropic_cache_control_hook.py +++ b/litellm/types/integrations/anthropic_cache_control_hook.py @@ -16,4 +16,13 @@ class CacheControlMessageInjectionPoint(TypedDict): control: Optional[ChatCompletionCachedContent] -CacheControlInjectionPoint = CacheControlMessageInjectionPoint +class CacheControlToolConfigInjectionPoint(TypedDict): + """Type for tool_config-level injection points (Bedrock).""" + + location: Literal["tool_config"] + + +CacheControlInjectionPoint = Union[ + CacheControlMessageInjectionPoint, + CacheControlToolConfigInjectionPoint, +] From c1e90ed300e9744d22f96b1d849535a7fd3597f1 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Sat, 21 Mar 2026 20:29:14 +0530 Subject: [PATCH 40/42] Fix mypy errors --- litellm/types/llms/vertex_ai.py | 14 ++++++++++++++ 1 file changed, 14 insertions(+) diff --git a/litellm/types/llms/vertex_ai.py b/litellm/types/llms/vertex_ai.py index 66c6ca436cf..c49fc96a65b 100644 --- a/litellm/types/llms/vertex_ai.py +++ b/litellm/types/llms/vertex_ai.py @@ -58,12 +58,26 @@ class HttpxBlobType(TypedDict, total=False): data: str +class HttpxServerSideToolCall(TypedDict, total=False): + toolType: str + id: str + args: dict + + +class HttpxServerSideToolResponse(TypedDict, total=False): + toolType: str + id: str + response: Union[str, dict] + + class HttpxPartType(TypedDict, total=False): text: str inlineData: HttpxBlobType fileData: FileDataType functionCall: HttpxFunctionCall functionResponse: FunctionResponse + toolCall: HttpxServerSideToolCall + toolResponse: HttpxServerSideToolResponse executableCode: HttpxExecutableCode codeExecutionResult: HttpxCodeExecutionResult thought: bool From 676a79e9f7a24053ae970e1578c9b7b1f9fa66e1 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Sat, 21 Mar 2026 20:42:34 +0530 Subject: [PATCH 41/42] =?UTF-8?q?bump:=20litellm-enterprise=200.1.34=20?= =?UTF-8?q?=E2=86=92=200.1.35?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- enterprise/pyproject.toml | 4 ++-- poetry.lock | 8 ++++---- pyproject.toml | 2 +- requirements.txt | 2 +- 4 files changed, 8 insertions(+), 8 deletions(-) diff --git a/enterprise/pyproject.toml b/enterprise/pyproject.toml index 515885944f0..9aa0a412edc 100644 --- a/enterprise/pyproject.toml +++ b/enterprise/pyproject.toml @@ -1,6 +1,6 @@ [tool.poetry] name = "litellm-enterprise" -version = "0.1.34" +version = "0.1.35" description = "Package for LiteLLM Enterprise features" authors = ["BerriAI"] readme = "README.md" @@ -22,7 +22,7 @@ requires = ["poetry-core"] build-backend = "poetry.core.masonry.api" [tool.commitizen] -version = "0.1.33" +version = "0.1.35" version_files = [ "pyproject.toml:version", "../requirements.txt:litellm-enterprise==", diff --git a/poetry.lock b/poetry.lock index d43cdea726c..06a77c1a4e3 100644 --- a/poetry.lock +++ b/poetry.lock @@ -3209,15 +3209,15 @@ openai = ["openai (>=0.27.8)"] [[package]] name = "litellm-enterprise" -version = "0.1.33" +version = "0.1.35" description = "Package for LiteLLM Enterprise features" optional = true python-versions = "!=2.7.*,!=3.0.*,!=3.1.*,!=3.2.*,!=3.3.*,!=3.4.*,!=3.5.*,!=3.6.*,!=3.7.*,>=3.8" groups = ["main"] markers = "extra == \"proxy\"" files = [ - {file = "litellm_enterprise-0.1.33-py3-none-any.whl", hash = "sha256:ae262ecfca680a235095becd6215e412e5ceba90efef739e61e6096b121188a2"}, - {file = "litellm_enterprise-0.1.33.tar.gz", hash = "sha256:5e3c0de9c4b54694ebb3017c8e18ee1d40e02ebef86e9ebd9c006e445885d5a0"}, + {file = "litellm_enterprise-0.1.35-py3-none-any.whl", hash = "sha256:8d2d9c925de8ee35e308c0f4975483b60f5e22beb50506e261e555e466f019c5"}, + {file = "litellm_enterprise-0.1.35.tar.gz", hash = "sha256:b752d07e538424743fcc08ba0d3d9d83d1f04a45c115811ac7828d789b6d87cc"}, ] [[package]] @@ -8018,4 +8018,4 @@ utils = ["numpydoc"] [metadata] lock-version = "2.1" python-versions = ">=3.9,<4.0" -content-hash = "2cf958f1a04fd5f1ab0e5cfc33bdbf441b518ed6c82d0f2546bf64cd3d2f89be" +content-hash = "f0977419272b446bc2df0e062406c8f7fe03566bd38fdb8395418bdc6da3fe20" diff --git a/pyproject.toml b/pyproject.toml index 73f495203bb..143054572ea 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -63,7 +63,7 @@ mcp = {version = ">=1.25.0,<2.0.0", optional = true, python = ">=3.10"} a2a-sdk = {version = "^0.3.22", optional = true, python = ">=3.10"} litellm-proxy-extras = {version = "^0.4.58", optional = true} rich = {version = "^13.7.1", optional = true} -litellm-enterprise = {version = "^0.1.33", optional = true} +litellm-enterprise = {version = "0.1.35", optional = true} diskcache = {version = "^5.6.1", optional = true} polars = {version = "^1.31.0", optional = true, python = ">=3.10"} semantic-router = {version = ">=0.1.12", optional = true, python = ">=3.9,<3.14"} diff --git a/requirements.txt b/requirements.txt index d420f4ac605..473bec42c57 100644 --- a/requirements.txt +++ b/requirements.txt @@ -80,4 +80,4 @@ pypdf>=6.7.3 # for PDF text extraction in RAG ingestion (CVE-2026-27888) ######################## # LITELLM ENTERPRISE DEPENDENCIES ######################## -litellm-enterprise==0.1.34 +litellm-enterprise==0.1.35 From 6830b63269df5cf4bdd4d0c53707f7097571b1d8 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Sat, 21 Mar 2026 20:44:52 +0530 Subject: [PATCH 42/42] =?UTF-8?q?Revert=20"fix(whisper):=20correct=20outpu?= =?UTF-8?q?t=5Fcost=5Fper=5Fsecond=20pricing=20and=20cost=20calcula?= =?UTF-8?q?=E2=80=A6"?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit This reverts commit 00dd9844158c06da595a863c3c18552612068079. --- litellm/llms/openai/cost_calculation.py | 3 +- ...odel_prices_and_context_window_backup.json | 4 +- model_prices_and_context_window.json | 4 +- tests/test_litellm/test_cost_calculator.py | 52 ------------------- 4 files changed, 6 insertions(+), 57 deletions(-) diff --git a/litellm/llms/openai/cost_calculation.py b/litellm/llms/openai/cost_calculation.py index d5077e25a43..ac1e4a6b08f 100644 --- a/litellm/llms/openai/cost_calculation.py +++ b/litellm/llms/openai/cost_calculation.py @@ -114,7 +114,7 @@ def cost_per_second( ) ## COST PER SECOND ## completion_cost = model_info["output_cost_per_second"] * duration - if ( + elif ( "input_cost_per_second" in model_info and model_info["input_cost_per_second"] is not None ): @@ -123,6 +123,7 @@ def cost_per_second( ) ## COST PER SECOND ## prompt_cost = model_info["input_cost_per_second"] * duration + completion_cost = 0.0 return prompt_cost, completion_cost diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 595cdf3e138..879dd42be47 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -5725,7 +5725,7 @@ "input_cost_per_second": 0.0001, "litellm_provider": "azure", "mode": "audio_transcription", - "output_cost_per_second": 0.0 + "output_cost_per_second": 0.0001 }, "azure_ai/Cohere-embed-v3-english": { "input_cost_per_token": 1e-07, @@ -31996,7 +31996,7 @@ "input_cost_per_second": 0.0001, "litellm_provider": "openai", "mode": "audio_transcription", - "output_cost_per_second": 0.0, + "output_cost_per_second": 0.0001, "supported_endpoints": [ "/v1/audio/transcriptions" ] diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index c1ba4d88d4d..bbf9f6d9dc8 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -5725,7 +5725,7 @@ "input_cost_per_second": 0.0001, "litellm_provider": "azure", "mode": "audio_transcription", - "output_cost_per_second": 0.0 + "output_cost_per_second": 0.0001 }, "azure_ai/Cohere-embed-v3-english": { "input_cost_per_token": 1e-07, @@ -32000,7 +32000,7 @@ "input_cost_per_second": 0.0001, "litellm_provider": "openai", "mode": "audio_transcription", - "output_cost_per_second": 0.0, + "output_cost_per_second": 0.0001, "supported_endpoints": [ "/v1/audio/transcriptions" ] diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index 88c384bcaec..8f5c3ece0ca 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -1,6 +1,5 @@ import os import sys -from unittest.mock import patch import pytest @@ -17,7 +16,6 @@ from litellm.cost_calculator import ( handle_realtime_stream_cost_calculation, response_cost_calculator, ) -from litellm.llms.openai.cost_calculation import cost_per_second from litellm.types.llms.openai import OpenAIRealtimeStreamList from litellm.types.utils import ModelResponse, PromptTokensDetailsWrapper, Usage from litellm.utils import TranscriptionResponse @@ -1972,53 +1970,3 @@ def test_additional_costs_only_for_azure_ai(): completion_tokens=50, ) assert result is None, "Vertex AI should have no additional costs" - - -class TestCostPerSecondArithmetic: - """Unit tests for the cost_per_second arithmetic itself. - - The tests in TestCostCalculatorReadsDurationFromHiddenParams mock - openai_cost_per_second entirely — they verify that the right duration - is forwarded, but never that the math inside cost_per_second is correct. - These tests cover the arithmetic directly. - """ - - def test_input_cost_per_second_only(self): - """input_cost_per_second * duration = prompt_cost; completion_cost = 0.""" - with patch("litellm.llms.openai.cost_calculation.get_model_info") as mock_info: - mock_info.return_value = { - "input_cost_per_second": 0.0001, - "output_cost_per_second": 0.0, - } - prompt_cost, completion_cost_val = cost_per_second( - model="whisper-1", custom_llm_provider="openai", duration=60.0 - ) - assert prompt_cost == pytest.approx(0.006) - assert completion_cost_val == 0.0 - - def test_both_input_and_output_cost_per_second(self): - """When both fields are set, they are applied independently.""" - with patch("litellm.llms.openai.cost_calculation.get_model_info") as mock_info: - mock_info.return_value = { - "input_cost_per_second": 0.002, - "output_cost_per_second": 0.003, - } - prompt_cost, completion_cost_val = cost_per_second( - model="some-model", custom_llm_provider="openai", duration=10.0 - ) - assert prompt_cost == pytest.approx(0.02) - assert completion_cost_val == pytest.approx(0.03) - - def test_zero_duration_returns_zero_cost(self): - """A zero-duration transcription must cost nothing regardless of the rate.""" - with patch("litellm.llms.openai.cost_calculation.get_model_info") as mock_info: - mock_info.return_value = { - "input_cost_per_second": 0.0001, - "output_cost_per_second": 0.0, - } - prompt_cost, completion_cost_val = cost_per_second( - model="whisper-1", custom_llm_provider="openai", duration=0.0 - ) - assert prompt_cost == 0.0 - assert completion_cost_val == 0.0 -