mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge pull request #24174 from BerriAI/litellm_oss_staging_03_19_2026
Litellm oss staging 03 19 2026
This commit is contained in:
commit
ea02c7cc15
48 changed files with 3737 additions and 174 deletions
|
|
@ -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' \
|
|||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
### 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
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
```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,
|
||||
)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="proxy" label="PROXY">
|
||||
|
||||
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
|
||||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
:::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
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
:::
|
||||
|
||||
<details>
|
||||
<summary>Advanced: Multiple modes with individual event hooks</summary>
|
||||
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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==",
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 {}
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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] = []
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 = (
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
8
poetry.lock
generated
8
poetry.lock
generated
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -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"
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -550,4 +550,72 @@ class TestMoonshotConfig:
|
|||
|
||||
# reasoning_content must not have been injected
|
||||
for msg in result["messages"]:
|
||||
assert "reasoning_content" not in msg
|
||||
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="<thinking>User wants weather</thinking>",
|
||||
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") == "<thinking>User wants weather</thinking>"
|
||||
|
||||
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": "<thinking>Planning to call weather tool</thinking>",
|
||||
"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") == "<thinking>Planning to call weather tool</thinking>"
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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')
|
||||
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()
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
167
tests/test_litellm/proxy/test_model_info_default_limits.py
Normal file
167
tests/test_litellm/proxy/test_model_info_default_limits.py
Normal file
|
|
@ -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
|
||||
Loading…
Add table
Reference in a new issue