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/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/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/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");
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/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py
index 2541a0bd7aa..2e5a8734085 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
@@ -27,7 +28,9 @@ 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
class WebSearchInterceptionLogger(CustomLogger):
@@ -67,6 +70,111 @@ 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
+
+ # 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
+
+ # Extract search query from the last user message
+ from litellm.litellm_core_utils.prompt_templates.common_utils import (
+ get_last_user_message,
+ )
+
+ query = get_last_user_message(cast(List[AllMessageValues], 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
+ response: Dict[str, Any] = {
+ "id": f"msg_{str(uuid.uuid4())}",
+ "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
+
async def async_pre_call_deployment_hook(
self, kwargs: Dict[str, Any], call_type: Optional[Any]
) -> Optional[dict]:
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/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py
index 96e70845b28..e99bae8ece8 100644
--- a/litellm/litellm_core_utils/streaming_handler.py
+++ b/litellm/litellm_core_utils/streaming_handler.py
@@ -2150,22 +2150,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/llms/anthropic/experimental_pass_through/adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py
index 3fda05172b6..b9a8d1de488 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 (
@@ -833,6 +834,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",
@@ -843,6 +849,40 @@ class LiteLLMAnthropicMessagesAdapter:
},
}
+ @staticmethod
+ def _add_additional_properties_false(schema: dict) -> None:
+ """
+ Recursively ensure object schemas comply with OpenAI strict mode.
+
+ 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)
+
+ # 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/litellm/llms/anthropic/experimental_pass_through/messages/handler.py b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py
index 1b5f03ec722..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
@@ -114,6 +114,55 @@ 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:
+ anthropic_response = cast(AnthropicMessagesResponse, response)
+ if stream:
+ from litellm.llms.anthropic.experimental_pass_through.messages.fake_stream_iterator import (
+ FakeAnthropicMessagesStreamIterator,
+ )
+
+ return FakeAnthropicMessagesStreamIterator(anthropic_response)
+ return anthropic_response
+
+ return None
+
+
@client
async def anthropic_messages(
max_tokens: int,
@@ -138,6 +187,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,
@@ -151,11 +206,38 @@ 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.
+ # 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
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.
+ # 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=original_stream,
+ )
+ if short_circuit_response is not None:
+ return short_circuit_response
+
loop = asyncio.get_event_loop()
kwargs["is_async"] = True
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/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/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/litellm/llms/openai/chat/gpt_transformation.py b/litellm/llms/openai/chat/gpt_transformation.py
index 63beb82ded8..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,8 +806,8 @@ class OpenAIChatCompletionStreamingHandler(BaseModelResponseIterator):
choices = chunk.get("choices", [])
choices = self._map_reasoning_to_reasoning_content(choices)
- kwargs = {
- "id": chunk["id"],
+ kwargs: Dict[str, Any] = {
+ "id": chunk.get("id"),
"object": "chat.completion.chunk",
"created": chunk.get("created"),
"model": chunk.get("model"),
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:
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/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/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py
index a03e1fb94c1..d2f8320668e 100644
--- a/litellm/proxy/auth/auth_utils.py
+++ b/litellm/proxy/auth/auth_utils.py
@@ -539,8 +539,45 @@ 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_limit(model_name: str, field: str) -> Optional[int]:
+ """
+ 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, 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
+
+ if llm_router is None:
+ return None
+ deployments = llm_router.get_model_list(model_name=model_name)
+ if not deployments:
+ return None
+ limits = []
+ for deployment in deployments:
+ raw = deployment.get("litellm_params", {}).get(field)
+ if raw is not None:
+ try:
+ if isinstance(raw, (int, float, str, bytes, bytearray)):
+ limits.append(int(raw))
+ except (ValueError, TypeError):
+ pass
+ 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 _get_deployment_default_limit(model_name, "default_api_key_tpm_limit")
+
+
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 +586,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 +605,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 is not None:
+ 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 +629,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 +648,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 is not None:
+ 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/common_request_processing.py b/litellm/proxy/common_request_processing.py
index e5a31c36719..bad70d30da3 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.
+ _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
+ # 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 _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
+
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,229 @@ class ProxyBaseLLMRequestProcessing:
response = responses[1]
- hidden_params = getattr(response, "_hidden_params", {}) or {}
- model_id = self._get_model_id_from_response(hidden_params, self.data)
+ _exception_raised = False
+ 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 _post_call_guardrails_active 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
+ )
+ except Exception:
+ _exception_raised = True
+ raise
+ 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
+ )
+
+ # 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. _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
+ )
+ if _deferred_fn is not None:
+ logging_obj._on_deferred_stream_complete = None # type: ignore[attr-defined]
+ try:
+ asyncio.create_task(
+ logging_obj.async_success_handler(
+ response,
+ cache_hit=None,
+ start_time=None,
+ 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,
+ cache_hit=None,
+ start_time=None,
+ end_time=None,
+ )
+ except Exception as e:
+ verbose_proxy_logger.exception(
+ "Error in orphaned streaming sync logging: %s", e
+ )
# Always return the client-requested model name (not provider-prefixed internal identifiers)
# for OpenAI-compatible responses.
@@ -1217,6 +1344,131 @@ 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
+
+ _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
+ )
+ 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,
+ )
+ 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,
+ )
+ )
+ 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/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/db/prisma_client.py b/litellm/proxy/db/prisma_client.py
index c9c0cfe8f68..82ee11a0f42 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
+ 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
+ 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
+ await asyncio.sleep(0.5)
+ try:
+ 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,
+ )
+ 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}")
+ 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/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py
index c7bfc27d6b6..fefc6c8af9c 100644
--- a/litellm/proxy/hooks/parallel_request_limiter.py
+++ b/litellm/proxy/hooks/parallel_request_limiter.py
@@ -295,16 +295,17 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
)
# Check if request under RPM/TPM per model for a given API Key
- 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
- ):
- _model = data.get("model", None)
+ _model = data.get("model", 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)
- _rpm_limit_for_key_model = get_key_model_rpm_limit(user_api_key_dict)
tpm_limit_for_model = None
rpm_limit_for_model = None
@@ -477,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
@@ -538,6 +548,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
@@ -545,6 +565,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 _success_tpm_limit is not None
+ or _success_rpm_limit is not None
)
):
request_count_api_key = (
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/proxy/utils.py b/litellm/proxy/utils.py
index 351bc23915a..7954f0b6460 100644
--- a/litellm/proxy/utils.py
+++ b/litellm/proxy/utils.py
@@ -1922,6 +1922,7 @@ class ProxyLogging:
original_exception,
traceback.format_exc(),
),
+ daemon=True,
).start()
async def post_call_success_hook(
@@ -4010,13 +4011,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,
)
+ await PrismaWrapper._kill_engine_process(old_pid)
await self.db.connect()
await self.db.query_raw("SELECT 1")
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,
+]
diff --git a/litellm/types/llms/vertex_ai.py b/litellm/types/llms/vertex_ai.py
index 201854369f1..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
@@ -244,8 +258,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/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/litellm/utils.py b/litellm/utils.py
index c8272586dad..c674190ba8c 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/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,
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
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"]
diff --git a/tests/litellm_utils_tests/test_bedrock_token_counter.py b/tests/litellm_utils_tests/test_bedrock_token_counter.py
index f7c29918820..abc45b03d6c 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 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 == f"{api_base}/model/amazon.nova-lite-v1:0/count-tokens"
+
+ 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")
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..82c1c9839e7
--- /dev/null
+++ b/tests/test_litellm/integrations/websearch_interception/test_websearch_short_circuit.py
@@ -0,0 +1,384 @@
+"""
+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_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"""
+ 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
+# ---------------------------------------------------------------------------
+
+
+# ---------------------------------------------------------------------------
+# 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
+
+ @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"
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"
+ ]
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..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
@@ -1984,3 +1984,130 @@ 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
+ assert schema["required"] == ["name"]
+
+ 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["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 = {
+ "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
+ assert "required" 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
+ 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
+ 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
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)
+
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"
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"""
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
diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py
index 5e42b110aa0..b66c081a943 100644
--- a/tests/test_litellm/proxy/auth/test_auth_utils.py
+++ b/tests/test_litellm/proxy/auth/test_auth_utils.py
@@ -2,7 +2,8 @@
Unit tests for auth_utils functions related to rate limiting and customer ID extraction.
"""
-from unittest.mock import patch
+from typing import Optional
+from unittest.mock import MagicMock, patch
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_utils import (
@@ -70,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."""
@@ -136,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."""
@@ -315,3 +356,196 @@ 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: 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:
+ 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
+
+ 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}
+
+
+ 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."""
+
+ 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
+
+ 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/db/test_prisma_client.py b/tests/test_litellm/proxy/db/test_prisma_client.py
index 83f07253fc8..f4ef933f219 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("asyncio.sleep", new_callable=AsyncMock),
+ ):
+ 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("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()
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..62fb1b5189c 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("asyncio.sleep", new_callable=AsyncMock),
+ ):
+ 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()
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..c4d2dce5876
--- /dev/null
+++ b/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py
@@ -0,0 +1,944 @@
+"""
+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
+
+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).
+ """
+ from fastapi import HTTPException # noqa: local import for test isolation
+
+ 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 production _run_deferred_stream_guardrails must pass the
+ guardrail-modified response to async_success_handler."""
+ 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 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,
+ )
+
+ await asyncio.sleep(0)
+ await asyncio.sleep(0)
+
+ 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 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):
+ 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.async_success_handler = track_async_success
+
+ guardrail = BlockingGuardrail()
+
+ 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,
+ )
+
+ await asyncio.sleep(0)
+ await asyncio.sleep(0)
+
+ 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, "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. 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 track_async_success(*args, **kwargs):
+ pass
+
+ mock_logging_obj.async_success_handler = track_async_success
+
+ guardrail = TransientErrorGuardrail()
+
+ 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,
+ )
+
+ await asyncio.sleep(0)
+
+ assert mock_logging_obj.model_call_details["metadata"].get(
+ "guardrail_blocked"
+ ) is not True, "guardrail_blocked must NOT be set for transient errors"
+
+ @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"
+
+ @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"
+
+ @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"
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..d9ebd554edc
--- /dev/null
+++ b/tests/test_litellm/proxy/test_model_info_default_limits.py
@@ -0,0 +1,167 @@
+"""
+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 typing import Optional
+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: Optional[int] = None,
+ default_rpm: Optional[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.
+ 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)
+
+ 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