Merge branch 'BerriAI:main' into main

This commit is contained in:
Esteban Zeller 2026-02-26 15:46:03 -03:00 • committed by GitHub
commit 9c255bca4e
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
32 changed files with 1840 additions and 146 deletions

View file

@ -299,12 +299,54 @@ class WebSearchInterceptionLogger(CustomLogger):
f"WebSearchInterception: Detected {len(tool_calls)} WebSearch tool call(s), executing agentic loop"
)
# Return tools dict with tool calls
# Extract thinking blocks from response content.
# When extended thinking is enabled, the model response includes
# thinking/redacted_thinking blocks that must be preserved and
# prepended to the follow-up assistant message.
thinking_blocks: List[Dict] = []
if isinstance(response, dict):
content = response.get("content", [])
else:
content = getattr(response, "content", []) or []
for block in content:
if isinstance(block, dict):
block_type = block.get("type")
else:
block_type = getattr(block, "type", None)
if block_type in ("thinking", "redacted_thinking"):
if isinstance(block, dict):
thinking_blocks.append(block)
else:
# Convert object to dict using getattr, matching the
# pattern in _detect_from_non_streaming_response
thinking_block_dict: Dict = {"type": block_type}
if block_type == "thinking":
thinking_block_dict["thinking"] = getattr(
block, "thinking", ""
)
thinking_block_dict["signature"] = getattr(
block, "signature", ""
)
else: # redacted_thinking
thinking_block_dict["data"] = getattr(
block, "data", ""
)
thinking_blocks.append(thinking_block_dict)
if thinking_blocks:
verbose_logger.debug(
f"WebSearchInterception: Extracted {len(thinking_blocks)} thinking block(s) from response"
)
# Return tools dict with tool calls and thinking blocks
tools_dict = {
"tool_calls": tool_calls,
"tool_type": "websearch",
"provider": custom_llm_provider,
"response_format": "anthropic",
"thinking_blocks": thinking_blocks,
}
return True, tools_dict
@ -387,6 +429,7 @@ class WebSearchInterceptionLogger(CustomLogger):
"""
tool_calls = tools["tool_calls"]
thinking_blocks = tools.get("thinking_blocks", [])
verbose_logger.debug(
f"WebSearchInterception: Executing agentic loop for {len(tool_calls)} search(es)"
@ -396,6 +439,7 @@ class WebSearchInterceptionLogger(CustomLogger):
model=model,
messages=messages,
tool_calls=tool_calls,
thinking_blocks=thinking_blocks,
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
logging_obj=logging_obj,
stream=stream,
@ -442,6 +486,7 @@ class WebSearchInterceptionLogger(CustomLogger):
model: str,
messages: List[Dict],
tool_calls: List[Dict],
thinking_blocks: List[Dict],
anthropic_messages_optional_request_params: Dict,
logging_obj: Any,
stream: bool,
@ -495,6 +540,7 @@ class WebSearchInterceptionLogger(CustomLogger):
assistant_message, user_message = WebSearchTransformation.transform_response(
tool_calls=tool_calls,
search_results=final_search_results,
thinking_blocks=thinking_blocks,
)
# Make follow-up request with search results

View file

@ -4,7 +4,7 @@ WebSearch Tool Transformation
Transforms between Anthropic/OpenAI tool_use format and LiteLLM search format.
"""
import json
from typing import Any, Dict, List, Tuple, Union
from typing import Any, Dict, List, Optional, Tuple, Union
from litellm._logging import verbose_logger
from litellm.constants import LITELLM_WEB_SEARCH_TOOL_NAME
@ -224,6 +224,7 @@ class WebSearchTransformation:
tool_calls: List[Dict],
search_results: List[str],
response_format: str = "anthropic",
thinking_blocks: Optional[List[Dict]] = None,
) -> Tuple[Dict, Union[Dict, List[Dict]]]:
"""
Transform LiteLLM search results to Anthropic/OpenAI tool_result format.
@ -235,6 +236,10 @@ class WebSearchTransformation:
tool_calls: List of tool_use/tool_calls dicts from transform_request
search_results: List of search result strings (one per tool_call)
response_format: Response format - "anthropic" or "openai" (default: "anthropic")
thinking_blocks: Optional list of thinking/redacted_thinking blocks
from the model's response. When present, prepended to the
assistant message content (required by Anthropic API when
thinking is enabled).
Returns:
(assistant_message, user_or_tool_messages):
@ -247,19 +252,29 @@ class WebSearchTransformation:
)
else:
return WebSearchTransformation._transform_response_anthropic(
tool_calls, search_results
tool_calls, search_results, thinking_blocks=thinking_blocks
)
@staticmethod
def _transform_response_anthropic(
tool_calls: List[Dict],
search_results: List[str],
thinking_blocks: Optional[List[Dict]] = None,
) -> Tuple[Dict, Dict]:
"""Transform to Anthropic format (single user message with tool_result blocks)"""
# Build assistant message with tool_use blocks
assistant_message = {
"role": "assistant",
"content": [
# Build assistant message content
assistant_content: List[Dict] = []
# Prepend thinking blocks if present.
# When extended thinking is enabled, Anthropic requires the assistant
# message to start with thinking/redacted_thinking blocks before any
# tool_use blocks. Same pattern as anthropic_messages_pt in factory.py.
if thinking_blocks:
assistant_content.extend(thinking_blocks)
# Add tool_use blocks
assistant_content.extend(
[
{
"type": "tool_use",
"id": tc["id"],
@ -267,7 +282,12 @@ class WebSearchTransformation:
"input": tc["input"],
}
for tc in tool_calls
],
]
)
assistant_message = {
"role": "assistant",
"content": assistant_content,
}
# Build user message with tool_result blocks

View file

@ -1106,19 +1106,19 @@ class LiteLLMAnthropicMessagesAdapter:
# extract usage
usage: Usage = getattr(response, "usage")
uncached_input_tokens = usage.prompt_tokens or 0
cached_tokens = 0
if hasattr(usage, "prompt_tokens_details") and usage.prompt_tokens_details:
cached_tokens = getattr(usage.prompt_tokens_details, "cached_tokens", 0) or 0
uncached_input_tokens -= cached_tokens
anthropic_usage = AnthropicUsage(
input_tokens=uncached_input_tokens,
output_tokens=usage.completion_tokens or 0,
)
# Add cache tokens if available (for prompt caching support)
if hasattr(usage, "_cache_creation_input_tokens") and usage._cache_creation_input_tokens > 0:
anthropic_usage["cache_creation_input_tokens"] = usage._cache_creation_input_tokens
if hasattr(usage, "_cache_read_input_tokens") and usage._cache_read_input_tokens > 0:
anthropic_usage["cache_read_input_tokens"] = usage._cache_read_input_tokens
if cached_tokens > 0:
anthropic_usage["cache_read_input_tokens"] = cached_tokens
translated_obj = AnthropicMessagesResponse(
id=response.id,
@ -1271,19 +1271,19 @@ class LiteLLMAnthropicMessagesAdapter:
litellm_usage_chunk = None
if litellm_usage_chunk is not None:
uncached_input_tokens = litellm_usage_chunk.prompt_tokens or 0
cached_tokens = 0
if hasattr(litellm_usage_chunk, "prompt_tokens_details") and litellm_usage_chunk.prompt_tokens_details:
cached_tokens = getattr(litellm_usage_chunk.prompt_tokens_details, "cached_tokens", 0) or 0
uncached_input_tokens -= cached_tokens
usage_delta = UsageDelta(
input_tokens=uncached_input_tokens,
output_tokens=litellm_usage_chunk.completion_tokens or 0,
)
# Add cache tokens if available (for prompt caching support)
if hasattr(litellm_usage_chunk, "_cache_creation_input_tokens") and litellm_usage_chunk._cache_creation_input_tokens > 0:
usage_delta["cache_creation_input_tokens"] = litellm_usage_chunk._cache_creation_input_tokens
if hasattr(litellm_usage_chunk, "_cache_read_input_tokens") and litellm_usage_chunk._cache_read_input_tokens > 0:
usage_delta["cache_read_input_tokens"] = litellm_usage_chunk._cache_read_input_tokens
if cached_tokens > 0:
usage_delta["cache_read_input_tokens"] = cached_tokens
else:
usage_delta = UsageDelta(input_tokens=0, output_tokens=0)
return MessageBlockDelta(

View file

@ -269,6 +269,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
"logprobs",
"top_logprobs",
"modalities",
"audio",
"parallel_tool_calls",
"web_search_options",
]

View file

@ -119,6 +119,12 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase):
# Map input_reference to image (will be processed in transform_video_create_request)
if "input_reference" in video_create_optional_params:
mapped_params["image"] = video_create_optional_params["input_reference"]
elif "image" in video_create_optional_params:
mapped_params["image"] = video_create_optional_params["image"]
# Pass through a provider-specific parameters block if provided directly
if "parameters" in video_create_optional_params:
mapped_params["parameters"] = video_create_optional_params["parameters"]
# Map size to aspectRatio
if "size" in video_create_optional_params:
@ -263,23 +269,49 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase):
instance_dict: Dict[str, Any] = {"prompt": prompt}
params_copy = video_create_optional_request_params.copy()
# Check if user wants to provide full instance dict
if "instances" in params_copy and isinstance(params_copy["instances"], dict):
# Replace/merge with user-provided instance
instance_dict.update(params_copy["instances"])
params_copy.pop("instances")
elif "image" in params_copy and params_copy["image"] is not None:
image_data = _convert_image_to_vertex_format(params_copy["image"])
image = params_copy["image"]
if isinstance(image, dict):
# Already in Vertex format e.g. {"gcsUri": "gs://..."} or
# {"bytesBase64Encoded": "...", "mimeType": "..."}
image_data = image
elif isinstance(image, str) and image.startswith("gs://"):
# Bare GCS URI — Vertex AI accepts gcsUri natively, no download needed
image_data = {"gcsUri": image}
elif isinstance(image, str):
raise ValueError(
f"Unsupported image value '{image}'. "
"Provide a GCS URI (gs://...), a dict with 'gcsUri' or "
"'bytesBase64Encoded'/'mimeType', or a binary file-like object."
)
else:
# File-like object — encode to base64
image_data = _convert_image_to_vertex_format(image)
instance_dict["image"] = image_data
params_copy.pop("image")
# Extract a nested "parameters" block that map_openai_params may have placed
# inside params_copy (e.g. from provider-specific pass-through). Merging it
# flat prevents the double-nesting bug:
# {"parameters": {"parameters": {...}}} ← wrong
# {"parameters": {...}} ← correct
nested_params = params_copy.pop("parameters", None)
vertex_params: Dict[str, Any] = {}
if isinstance(nested_params, dict):
vertex_params.update(nested_params)
vertex_params.update(params_copy)
# Build request data directly (TypedDict doesn't have model_dump)
request_data: Dict[str, Any] = {"instances": [instance_dict]}
# Only add parameters if there are any
if params_copy:
request_data["parameters"] = params_copy
if vertex_params:
request_data["parameters"] = vertex_params
# Append :predictLongRunning endpoint to api_base
url = f"{api_base}:predictLongRunning"

View file

@ -4680,12 +4680,16 @@ def embedding( # noqa: PLR0915
if dynamic_api_key is not None:
api_key = dynamic_api_key
allowed_openai_params: Optional[List[str]] = kwargs.get(
"allowed_openai_params", None
)
optional_params = get_optional_params_embeddings(
model=model,
user=user,
dimensions=dimensions,
encoding_format=encoding_format,
custom_llm_provider=custom_llm_provider,
allowed_openai_params=allowed_openai_params,
**non_default_params,
)

View file

@ -756,14 +756,30 @@ class MCPServerManager:
Returns server_ids unchanged when client_ip is None (no filtering).
"""
filtered, _ = self.filter_server_ids_by_ip_with_info(server_ids, client_ip)
return filtered
def filter_server_ids_by_ip_with_info(
self, server_ids: List[str], client_ip: Optional[str]
) -> Tuple[List[str], int]:
"""
Filter server IDs by client IP — external callers only see public servers.
Returns (filtered_ids, ip_blocked_count) where ip_blocked_count is the number
of servers that were blocked because the client IP is not allowed to access them.
Returns server_ids unchanged (with 0 blocked) when client_ip is None.
"""
if client_ip is None:
return server_ids
return [
sid
for sid in server_ids
if (s := self.get_mcp_server_by_id(sid)) is not None
and self._is_server_accessible_from_ip(s, client_ip)
]
return server_ids, 0
allowed = []
blocked = 0
for sid in server_ids:
s = self.get_mcp_server_by_id(sid)
if s is not None and self._is_server_accessible_from_ip(s, client_ip):
allowed.append(sid)
elif s is not None:
blocked += 1
return allowed, blocked
async def get_tools_for_server(self, server_id: str) -> List[MCPTool]:
"""

View file

@ -10,9 +10,9 @@ from litellm.proxy._experimental.mcp_server.ui_session_utils import (
)
from litellm.proxy._experimental.mcp_server.utils import merge_mcp_headers
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
from litellm.types.mcp import MCPAuth
from litellm.types.utils import CallTypes
@ -283,8 +283,10 @@ if MCP_AVAILABLE:
)
allowed_server_ids_set.update(servers)
allowed_server_ids = global_mcp_server_manager.filter_server_ids_by_ip(
list(allowed_server_ids_set), _rest_client_ip
allowed_server_ids, _ip_blocked_count = (
global_mcp_server_manager.filter_server_ids_by_ip_with_info(
list(allowed_server_ids_set), _rest_client_ip
)
)
list_tools_result = []
@ -293,6 +295,26 @@ if MCP_AVAILABLE:
# If server_id is specified, only query that specific server
if server_id:
if server_id not in allowed_server_ids:
_server = global_mcp_server_manager.get_mcp_server_by_id(server_id)
if (
_server is not None
and _rest_client_ip is not None
and not global_mcp_server_manager._is_server_accessible_from_ip(
_server, _rest_client_ip
)
):
raise HTTPException(
status_code=403,
detail={
"error": "ip_filtering",
"message": (
f"MCP server '{server_id}' is not accessible from your IP address "
f"({_rest_client_ip}). This server is restricted to internal "
"networks only. To make it externally accessible, set "
"'available_on_public_internet: true' in the server configuration."
),
},
)
raise HTTPException(
status_code=403,
detail={
@ -330,6 +352,19 @@ if MCP_AVAILABLE:
}
else:
if not allowed_server_ids:
if _ip_blocked_count > 0:
raise HTTPException(
status_code=403,
detail={
"error": "ip_filtering",
"message": (
f"No MCP tools are available for your IP address ({_rest_client_ip}). "
f"{_ip_blocked_count} server(s) are restricted to internal networks only. "
"To make servers externally accessible, set "
"'available_on_public_internet: true' in the server configuration."
),
},
)
raise HTTPException(
status_code=403,
detail={

View file

@ -771,8 +771,8 @@ if MCP_AVAILABLE:
user_api_key_auth
)
)
allowed_mcp_server_ids = (
global_mcp_server_manager.filter_server_ids_by_ip(
allowed_mcp_server_ids, _ip_blocked = (
global_mcp_server_manager.filter_server_ids_by_ip_with_info(
allowed_mcp_server_ids, client_ip
)
)
@ -780,6 +780,16 @@ if MCP_AVAILABLE:
"MCP IP filter: client_ip=%s, allowed_server_ids=%s",
client_ip, allowed_mcp_server_ids,
)
if _ip_blocked > 0:
verbose_logger.debug(
"MCP IP filtering: %d server(s) are not accessible from client IP %s "
"because they are restricted to internal networks. "
"No tools from those servers will be returned. "
"To expose a server externally, set 'available_on_public_internet: true' "
"in its configuration.",
_ip_blocked,
client_ip,
)
allowed_mcp_servers: List[MCPServer] = []
for allowed_mcp_server_id in allowed_mcp_server_ids:
mcp_server = global_mcp_server_manager.get_mcp_server_by_id(

View file

@ -1106,7 +1106,9 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase):
raise ValueError("args is required for stdio transport")
elif transport in [MCPTransport.http, MCPTransport.sse]:
if not values.get("url") and not values.get("spec_path"):
raise ValueError("url or spec_path is required for HTTP/SSE transport")
raise ValueError(
"url or spec_path is required for HTTP/SSE transport"
)
return values
@model_validator(mode="before")
@ -1158,7 +1160,9 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase):
raise ValueError("args is required for stdio transport")
elif transport in [MCPTransport.http, MCPTransport.sse]:
if not values.get("url") and not values.get("spec_path"):
raise ValueError("url or spec_path is required for HTTP/SSE transport")
raise ValueError(
"url or spec_path is required for HTTP/SSE transport"
)
return values
@ -1409,12 +1413,12 @@ class NewCustomerRequest(BudgetNewRequest):
blocked: bool = False # allow/disallow requests for this end-user
budget_id: Optional[str] = None # give either a budget_id or max_budget
spend: Optional[float] = None
allowed_model_region: Optional[AllowedModelRegion] = (
None # require all user requests to use models in this specific region
)
default_model: Optional[str] = (
None # if no equivalent model in allowed region - default all requests to this model
)
allowed_model_region: Optional[
AllowedModelRegion
] = None # require all user requests to use models in this specific region
default_model: Optional[
str
] = None # if no equivalent model in allowed region - default all requests to this model
object_permission: Optional[LiteLLM_ObjectPermissionBase] = None
@model_validator(mode="before")
@ -1437,12 +1441,12 @@ class UpdateCustomerRequest(LiteLLMPydanticObjectBase):
blocked: bool = False # allow/disallow requests for this end-user
max_budget: Optional[float] = None
budget_id: Optional[str] = None # give either a budget_id or max_budget
allowed_model_region: Optional[AllowedModelRegion] = (
None # require all user requests to use models in this specific region
)
default_model: Optional[str] = (
None # if no equivalent model in allowed region - default all requests to this model
)
allowed_model_region: Optional[
AllowedModelRegion
] = None # require all user requests to use models in this specific region
default_model: Optional[
str
] = None # if no equivalent model in allowed region - default all requests to this model
object_permission: Optional[LiteLLM_ObjectPermissionBase] = None
@ -2268,6 +2272,7 @@ class LiteLLM_VerificationTokenView(LiteLLM_VerificationToken):
end_user_tpm_limit: Optional[int] = None
end_user_rpm_limit: Optional[int] = None
end_user_max_budget: Optional[float] = None
end_user_model_max_budget: Optional[dict] = None
# Organization Params
organization_max_budget: Optional[float] = None
@ -3056,7 +3061,9 @@ class SpendLogsMetadata(TypedDict):
str
] # S3/GCS object key for cold storage retrieval
litellm_overhead_time_ms: Optional[float] # LiteLLM overhead time in milliseconds
attempted_retries: Optional[int] # Number of retries attempted (0 = first attempt succeeded)
attempted_retries: Optional[
int
] # Number of retries attempted (0 = first attempt succeeded)
max_retries: Optional[int] # Max retries configured for this request
cost_breakdown: Optional[
CostBreakdown
@ -4117,10 +4124,10 @@ class SpendUpdateQueueItem(TypedDict, total=False):
class ToolDiscoveryQueueItem(TypedDict, total=False):
tool_name: str
origin: Optional[str] # MCP server name or "user_defined"
origin: Optional[str] # MCP server name or "user_defined"
created_by: Optional[str]
key_hash: Optional[str] # hash of virtual key that triggered discovery
team_id: Optional[str] # team that triggered discovery
key_hash: Optional[str] # hash of virtual key that triggered discovery
team_id: Optional[str] # team that triggered discovery
key_alias: Optional[str] # human-readable key alias
@ -4144,6 +4151,7 @@ class LiteLLM_ManagedObjectTable(LiteLLMPydanticObjectBase):
class LiteLLM_ManagedVectorStoreTable(LiteLLMPydanticObjectBase):
"""Table for managing vector stores with target_model_names support."""
unified_resource_id: str
resource_object: Optional[Any] = None # VectorStoreCreateResponse
model_mappings: Dict[str, str]

View file

@ -183,6 +183,9 @@ def _apply_budget_limits_to_end_user_params(
if budget_info.max_budget is not None:
end_user_params["end_user_max_budget"] = budget_info.max_budget
if budget_info.model_max_budget is not None:
end_user_params["end_user_model_max_budget"] = budget_info.model_max_budget
verbose_proxy_logger.debug(f"Applied budget limits to end user {end_user_id}")
@ -241,9 +244,20 @@ def update_valid_token_with_end_user_params(
valid_token: UserAPIKeyAuth, end_user_params: dict
) -> UserAPIKeyAuth:
valid_token.end_user_id = end_user_params.get("end_user_id")
valid_token.end_user_tpm_limit = end_user_params.get("end_user_tpm_limit")
valid_token.end_user_rpm_limit = end_user_params.get("end_user_rpm_limit")
valid_token.allowed_model_region = end_user_params.get("allowed_model_region")
# Only overwrite token fields when the DB-derived value is not None.
# This prevents DB lookups (where the budget table has no value set)
# from silently clearing values that a custom auth function may have
# already set on the token.
if end_user_params.get("end_user_tpm_limit") is not None:
valid_token.end_user_tpm_limit = end_user_params["end_user_tpm_limit"]
if end_user_params.get("end_user_rpm_limit") is not None:
valid_token.end_user_rpm_limit = end_user_params["end_user_rpm_limit"]
if end_user_params.get("allowed_model_region") is not None:
valid_token.allowed_model_region = end_user_params["allowed_model_region"]
if end_user_params.get("end_user_model_max_budget") is not None:
valid_token.end_user_model_max_budget = end_user_params[
"end_user_model_max_budget"
]
return valid_token
@ -498,13 +512,29 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
request=request, api_key=api_key, user_custom_auth=user_custom_auth
)
if response is not None and isinstance(response, UserAPIKeyAuth):
return UserAPIKeyAuth.model_validate(response)
validated = UserAPIKeyAuth.model_validate(response)
validated = await _run_post_custom_auth_checks(
valid_token=validated,
request=request,
request_data=request_data,
route=route,
parent_otel_span=parent_otel_span,
)
return validated
elif response is not None and isinstance(response, str):
api_key = response
custom_auth_api_key = True
elif user_custom_auth is not None:
response = await user_custom_auth(request=request, api_key=api_key) # type: ignore
return UserAPIKeyAuth.model_validate(response)
validated = UserAPIKeyAuth.model_validate(response)
validated = await _run_post_custom_auth_checks(
valid_token=validated,
request=request,
request_data=request_data,
route=route,
parent_otel_span=parent_otel_span,
)
return validated
### LITELLM-DEFINED AUTH FUNCTION ###
#### IF JWT ####
@ -1210,6 +1240,21 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
model=current_model,
)
# Check 5b. End-user model max budget
end_user_mmb = valid_token.end_user_model_max_budget
if (
end_user_mmb is not None
and isinstance(end_user_mmb, dict)
and len(end_user_mmb) > 0
and current_model is not None
and valid_token.end_user_id is not None
):
await model_max_budget_limiter.is_end_user_within_model_budget(
end_user_id=valid_token.end_user_id,
end_user_model_max_budget=end_user_mmb,
model=current_model,
)
# Check 6: Additional Common Checks across jwt + key auth
if valid_token.team_id is not None:
try:
@ -1501,3 +1546,218 @@ def _update_key_budget_with_temp_budget_increase(
temp_budget_increase = _get_temp_budget_increase(valid_token) or 0.0
valid_token.max_budget = valid_token.max_budget + temp_budget_increase
return valid_token
async def _lookup_end_user_and_apply_budget(
valid_token: UserAPIKeyAuth,
route: str,
parent_otel_span: Optional[Span],
prisma_client,
user_api_key_cache,
proxy_logging_obj,
):
"""Look up end_user from DB and apply budget limits to valid_token."""
end_user_object = None
try:
end_user_object = await get_end_user_object(
end_user_id=valid_token.end_user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
route=route,
)
if end_user_object is not None:
end_user_params = {
"end_user_id": valid_token.end_user_id,
"allowed_model_region": end_user_object.allowed_model_region,
}
if end_user_object.litellm_budget_table is not None:
_apply_budget_limits_to_end_user_params(
end_user_params=end_user_params,
budget_info=end_user_object.litellm_budget_table,
end_user_id=valid_token.end_user_id,
)
valid_token = update_valid_token_with_end_user_params(
valid_token=valid_token, end_user_params=end_user_params
)
elif litellm.max_end_user_budget_id is not None:
from litellm.proxy.auth.auth_checks import get_default_end_user_budget
default_budget = await get_default_end_user_budget(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
)
if default_budget is not None:
end_user_params = {"end_user_id": valid_token.end_user_id}
_apply_budget_limits_to_end_user_params(
end_user_params=end_user_params,
budget_info=default_budget,
end_user_id=valid_token.end_user_id,
)
valid_token = update_valid_token_with_end_user_params(
valid_token=valid_token, end_user_params=end_user_params
)
except Exception as e:
if isinstance(e, litellm.BudgetExceededError):
raise e
verbose_proxy_logger.debug(f"Unable to find user in db. Error - {str(e)}")
return valid_token, end_user_object
async def _run_post_custom_auth_checks(
valid_token: UserAPIKeyAuth,
request: Request,
request_data: dict,
route: str,
parent_otel_span: Optional[Span],
) -> UserAPIKeyAuth:
from litellm.proxy.proxy_server import (
prisma_client,
user_api_key_cache,
proxy_logging_obj,
general_settings,
llm_router,
model_max_budget_limiter,
)
# 1. Look up end_user object from DB if end_user_id is set
end_user_object = None
if valid_token.end_user_id is not None:
valid_token, end_user_object = await _lookup_end_user_and_apply_budget(
valid_token=valid_token,
route=route,
parent_otel_span=parent_otel_span,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
# 2. Check token expiry
if valid_token.expires is not None:
current_time = datetime.now(timezone.utc)
if isinstance(valid_token.expires, datetime):
expiry_time = valid_token.expires
else:
expiry_time = datetime.fromisoformat(valid_token.expires)
if (
expiry_time.tzinfo is None
or expiry_time.tzinfo.utcoffset(expiry_time) is None
):
expiry_time = expiry_time.replace(tzinfo=timezone.utc)
if expiry_time < current_time:
raise ProxyException(
message=f"Authentication Error - Expired Key. Key Expiry time {expiry_time} and current time {current_time}",
type=ProxyErrorTypes.expired_key,
code=400,
param=abbreviate_api_key(api_key=valid_token.token)
if valid_token.token
else "",
)
current_model = request_data.get("model", None)
# 3. Check key-level model_max_budget
max_budget_per_model = valid_token.model_max_budget
if (
max_budget_per_model is not None
and isinstance(max_budget_per_model, dict)
and len(max_budget_per_model) > 0
and current_model is not None
and valid_token.token is not None
):
await model_max_budget_limiter.is_key_within_model_budget(
user_api_key_dict=valid_token,
model=current_model,
)
# 4. Check end-user model_max_budget
end_user_mmb = valid_token.end_user_model_max_budget
if (
end_user_mmb is not None
and isinstance(end_user_mmb, dict)
and len(end_user_mmb) > 0
and current_model is not None
and valid_token.end_user_id is not None
):
await model_max_budget_limiter.is_end_user_within_model_budget(
end_user_id=valid_token.end_user_id,
end_user_model_max_budget=end_user_mmb,
model=current_model,
)
# 5. Look up user object if user_id is set
user_object = None
if valid_token.user_id is not None:
try:
user_object = await get_user_object(
user_id=valid_token.user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
except Exception:
# If user_role is PROXY_ADMIN on the token, create a synthetic user object
# so that admin route checks pass for custom auth
if valid_token.user_role == LitellmUserRoles.PROXY_ADMIN:
user_object = LiteLLM_UserTable(
user_id=valid_token.user_id,
user_role=LitellmUserRoles.PROXY_ADMIN,
spend=0.0,
)
# 6. Run common checks
if valid_token.team_id is not None:
try:
_team_obj = await get_team_object(
team_id=valid_token.team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
except HTTPException:
_team_obj = LiteLLM_TeamTableCachedObj(
team_id=valid_token.team_id,
max_budget=valid_token.team_max_budget,
soft_budget=valid_token.team_soft_budget,
spend=valid_token.team_spend,
tpm_limit=valid_token.team_tpm_limit,
rpm_limit=valid_token.team_rpm_limit,
blocked=valid_token.team_blocked,
models=valid_token.team_models,
metadata=valid_token.team_metadata,
object_permission_id=valid_token.team_object_permission_id,
)
else:
_team_obj = None
_project_obj = None
if valid_token.project_id is not None:
_project_obj = await get_project_object(
project_id=valid_token.project_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
_ = await common_checks(
request=request,
request_body=request_data,
team_object=_team_obj,
user_object=user_object,
end_user_object=end_user_object,
general_settings=general_settings,
global_proxy_spend=None,
route=route,
llm_router=llm_router,
proxy_logging_obj=proxy_logging_obj,
valid_token=valid_token,
skip_budget_checks=False,
project_object=_project_obj,
)
return valid_token

View file

@ -15,6 +15,7 @@ from litellm.types.utils import (
)
VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX = "virtual_key_spend"
END_USER_SPEND_CACHE_KEY_PREFIX = "end_user_model_spend"
class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
@ -83,6 +84,81 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
return True
async def is_end_user_within_model_budget(
self,
end_user_id: str,
end_user_model_max_budget: dict,
model: str,
) -> bool:
"""
Check if the end_user is within the model budget
Raises:
BudgetExceededError: If the end_user has exceeded the model budget
"""
internal_model_max_budget: GenericBudgetConfigType = {}
for _model, _budget_info in end_user_model_max_budget.items():
internal_model_max_budget[_model] = BudgetConfig(**_budget_info)
verbose_proxy_logger.debug(
"end_user internal_model_max_budget %s",
json.dumps(internal_model_max_budget, indent=4, default=str),
)
# check if current model is in internal_model_max_budget
_current_model_budget_info = self._get_request_model_budget_config(
model=model, internal_model_max_budget=internal_model_max_budget
)
if _current_model_budget_info is None:
verbose_proxy_logger.debug(
f"Model {model} not found in end_user_model_max_budget"
)
return True
# check if current model is within budget
if (
_current_model_budget_info.max_budget
and _current_model_budget_info.max_budget > 0
):
_current_spend = await self._get_end_user_spend_for_model(
end_user_id=end_user_id,
model=model,
key_budget_config=_current_model_budget_info,
)
if (
_current_spend is not None
and _current_model_budget_info.max_budget is not None
and _current_spend > _current_model_budget_info.max_budget
):
raise litellm.BudgetExceededError(
message=f"LiteLLM End User: {end_user_id}, exceeded budget for model={model}",
current_cost=_current_spend,
max_budget=_current_model_budget_info.max_budget,
)
return True
async def _get_end_user_spend_for_model(
self,
end_user_id: str,
model: str,
key_budget_config: BudgetConfig,
) -> Optional[float]:
# 1. model: directly look up `model`
end_user_model_spend_cache_key = f"{END_USER_SPEND_CACHE_KEY_PREFIX}:{end_user_id}:{model}:{key_budget_config.budget_duration}"
_current_spend = await self.dual_cache.async_get_cache(
key=end_user_model_spend_cache_key,
)
if _current_spend is None:
# 2. If 1, does not exist, check if passed as {custom_llm_provider}/model
end_user_model_spend_cache_key = f"{END_USER_SPEND_CACHE_KEY_PREFIX}:{end_user_id}:{self._get_model_without_custom_llm_provider(model)}:{key_budget_config.budget_duration}"
_current_spend = await self.dual_cache.async_get_cache(
key=end_user_model_spend_cache_key,
)
return _current_spend
async def _get_virtual_key_spend_for_model(
self,
user_api_key_hash: Optional[str],
@ -163,46 +239,77 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
user_api_key_model_max_budget: Optional[dict] = _metadata.get(
"user_api_key_model_max_budget", None
)
user_api_key_end_user_model_max_budget: Optional[dict] = _metadata.get(
"user_api_key_end_user_model_max_budget", None
)
if (
user_api_key_model_max_budget is None
or len(user_api_key_model_max_budget) == 0
) and (
user_api_key_end_user_model_max_budget is None
or len(user_api_key_end_user_model_max_budget) == 0
):
verbose_proxy_logger.debug(
"Not running _PROXY_VirtualKeyModelMaxBudgetLimiter.async_log_success_event because user_api_key_model_max_budget is None or empty. `user_api_key_model_max_budget`=%s",
user_api_key_model_max_budget,
"Not running _PROXY_VirtualKeyModelMaxBudgetLimiter.async_log_success_event because user_api_key_model_max_budget and user_api_key_end_user_model_max_budget are None or empty."
)
return
response_cost: float = standard_logging_payload.get("response_cost", 0)
model = standard_logging_payload.get("model")
virtual_key = standard_logging_payload.get("metadata", {}).get(
"user_api_key_hash"
)
end_user_id = standard_logging_payload.get(
"end_user"
) or standard_logging_payload.get("metadata", {}).get(
"user_api_key_end_user_id"
)
if virtual_key is None or model is None:
if model is None:
return
# Resolve per-model budget config (same logic as is_key_within_model_budget)
internal_model_max_budget: GenericBudgetConfigType = {}
for _model, _budget_info in user_api_key_model_max_budget.items():
internal_model_max_budget[_model] = BudgetConfig(**_budget_info)
key_budget_config = self._get_request_model_budget_config(
model=model, internal_model_max_budget=internal_model_max_budget
)
if key_budget_config is None or not key_budget_config.budget_duration:
verbose_proxy_logger.debug(
"Not incrementing model spend: no budget config or budget_duration for model=%s",
model,
if (
virtual_key is not None
and user_api_key_model_max_budget is not None
and len(user_api_key_model_max_budget) > 0
):
internal_model_max_budget: GenericBudgetConfigType = {}
for _model, _budget_info in user_api_key_model_max_budget.items():
internal_model_max_budget[_model] = BudgetConfig(**_budget_info)
key_budget_config = self._get_request_model_budget_config(
model=model, internal_model_max_budget=internal_model_max_budget
)
return
if key_budget_config is not None and key_budget_config.budget_duration:
virtual_spend_key = f"{VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX}:{virtual_key}:{model}:{key_budget_config.budget_duration}"
virtual_start_time_key = f"virtual_key_budget_start_time:{virtual_key}"
await self._increment_spend_for_key(
budget_config=key_budget_config,
spend_key=virtual_spend_key,
start_time_key=virtual_start_time_key,
response_cost=response_cost,
)
if (
end_user_id is not None
and user_api_key_end_user_model_max_budget is not None
and len(user_api_key_end_user_model_max_budget) > 0
):
internal_model_max_budget: GenericBudgetConfigType = {}
for _model, _budget_info in user_api_key_end_user_model_max_budget.items():
internal_model_max_budget[_model] = BudgetConfig(**_budget_info)
key_budget_config = self._get_request_model_budget_config(
model=model, internal_model_max_budget=internal_model_max_budget
)
if key_budget_config is not None and key_budget_config.budget_duration:
end_user_spend_key = f"{END_USER_SPEND_CACHE_KEY_PREFIX}:{end_user_id}:{model}:{key_budget_config.budget_duration}"
end_user_start_time_key = f"end_user_budget_start_time:{end_user_id}"
await self._increment_spend_for_key(
budget_config=key_budget_config,
spend_key=end_user_spend_key,
start_time_key=end_user_start_time_key,
response_cost=response_cost,
)
virtual_spend_key = f"{VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX}:{virtual_key}:{model}:{key_budget_config.budget_duration}"
virtual_start_time_key = f"virtual_key_budget_start_time:{virtual_key}"
await self._increment_spend_for_key(
budget_config=key_budget_config,
spend_key=virtual_spend_key,
start_time_key=virtual_start_time_key,
response_cost=response_cost,
)
verbose_proxy_logger.debug(
"current state of in memory cache %s",
json.dumps(

View file

@ -1047,6 +1047,9 @@ async def add_litellm_data_to_request( # noqa: PLR0915
data[_metadata_variable_name][
"user_api_key_model_max_budget"
] = user_api_key_dict.model_max_budget
data[_metadata_variable_name][
"user_api_key_end_user_model_max_budget"
] = user_api_key_dict.end_user_model_max_budget
# User spend, budget - used by prometheus.py
# Follow same pattern as team and API key budgets

View file

@ -4109,13 +4109,23 @@ async def list_keys(
dependencies=[Depends(user_api_key_auth)],
)
@management_endpoint_wrapper
async def key_aliases() -> Dict[str, List[str]]:
async def key_aliases(
page: int = Query(1, ge=1, description="Page number"),
size: int = Query(50, ge=1, le=100, description="Page size"),
search: Optional[str] = Query(
None, description="Search key aliases (case-insensitive partial match)"
),
) -> Dict[str, Any]:
"""
Lists all key aliases
Lists key aliases with pagination and optional search.
Returns:
{
"aliases": List[str]
"aliases": List[str],
"total_count": int,
"current_page": int,
"total_pages": int,
"size": int,
}
"""
try:
@ -4127,36 +4137,55 @@ async def key_aliases() -> Dict[str, List[str]]:
verbose_proxy_logger.error("Database not connected")
raise Exception("Database not connected")
where: Dict[str, Any] = {}
try:
where.update(_get_condition_to_filter_out_ui_session_tokens())
except NameError:
# Helper may not exist in some builds; ignore if missing
pass
# Build a parameterized WHERE clause to avoid loading full rows into
# memory. Raw SQL is used because the Prisma client wrapper does not
# support column-level SELECT projection on find_many.
#
# $1 is always UI_SESSION_TOKEN_TEAM_ID (filters out UI session tokens).
query_params: List[Any] = [UI_SESSION_TOKEN_TEAM_ID]
where_parts = [
"key_alias IS NOT NULL",
"key_alias != ''",
"(team_id IS NULL OR team_id != $1)",
]
if search:
query_params.append(f"%{search}%")
where_parts.append(f"key_alias ILIKE ${len(query_params)}")
rows = await prisma_client.db.litellm_verificationtoken.find_many(
where=where,
order=[{"key_alias": "asc"}],
where_sql = " AND ".join(where_parts)
count_sql = (
f'SELECT COUNT(*) AS count FROM "LiteLLM_VerificationToken" WHERE {where_sql}'
)
count_rows = await prisma_client.db.query_raw(count_sql, *query_params)
total_count = int(count_rows[0]["count"]) if count_rows else 0
aliases_params = query_params + [size, (page - 1) * size]
limit_idx = len(aliases_params) - 1
offset_idx = len(aliases_params)
aliases_sql = (
f"SELECT key_alias"
f' FROM "LiteLLM_VerificationToken"'
f" WHERE {where_sql}"
f" ORDER BY key_alias ASC"
f" LIMIT ${limit_idx} OFFSET ${offset_idx}"
)
alias_rows = await prisma_client.db.query_raw(aliases_sql, *aliases_params)
aliases: List[str] = [row["key_alias"] for row in alias_rows if row.get("key_alias")]
total_pages = -(-total_count // size) if total_count > 0 else 0
verbose_proxy_logger.debug(
f"key_aliases: page={page}, size={size}, search={search!r}, "
f"total_count={total_count}, total_pages={total_pages}"
)
seen = set()
aliases: List[str] = []
for row in rows:
alias = getattr(row, "key_alias", None)
if alias is None and isinstance(row, dict):
alias = row.get("key_alias")
if not alias:
continue
alias_str = str(alias).strip()
if alias_str and alias_str not in seen:
seen.add(alias_str)
aliases.append(alias_str)
verbose_proxy_logger.debug(f"Returning {len(aliases)} key aliases")
return {"aliases": aliases}
return {
"aliases": aliases,
"total_count": total_count,
"current_page": page,
"total_pages": total_pages,
"size": size,
}
except Exception as e:
verbose_proxy_logger.exception(f"Error in key_aliases: {e}")

View file

@ -1,5 +1,6 @@
import hashlib
import json
import os
import secrets
from datetime import datetime
from datetime import datetime as dt
@ -10,7 +11,10 @@ from pydantic import BaseModel
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import MAX_STRING_LENGTH_PROMPT_IN_DB, REDACTED_BY_LITELM_STRING
from litellm.constants import (
MAX_STRING_LENGTH_PROMPT_IN_DB as DEFAULT_MAX_STRING_LENGTH_PROMPT_IN_DB,
)
from litellm.constants import REDACTED_BY_LITELM_STRING
from litellm.litellm_core_utils.core_helpers import (
get_litellm_metadata_from_kwargs,
reconstruct_model_name,
@ -30,6 +34,20 @@ from litellm.types.utils import (
from litellm.utils import get_end_user_id_for_cost_tracking
def _get_max_string_length_prompt_in_db() -> int:
"""
Resolve prompt truncation threshold at runtime so values loaded later via
proxy config environment_variables are honored.
"""
max_length_str = os.getenv("MAX_STRING_LENGTH_PROMPT_IN_DB")
if max_length_str is None:
return DEFAULT_MAX_STRING_LENGTH_PROMPT_IN_DB
try:
return int(max_length_str)
except (TypeError, ValueError):
return DEFAULT_MAX_STRING_LENGTH_PROMPT_IN_DB
def _is_master_key(api_key: str, _master_key: Optional[str]) -> bool:
if _master_key is None:
return False
@ -609,6 +627,7 @@ def _get_messages_for_spend_logs_payload(
def _sanitize_request_body_for_spend_logs_payload(
request_body: dict,
visited: Optional[set] = None,
max_string_length_prompt_in_db: Optional[int] = None,
) -> dict:
"""
Recursively sanitize request body to prevent logging large base64 strings or other large values.
@ -618,6 +637,8 @@ def _sanitize_request_body_for_spend_logs_payload(
if visited is None:
visited = set()
if max_string_length_prompt_in_db is None:
max_string_length_prompt_in_db = _get_max_string_length_prompt_in_db()
# Get the object's memory address to track visited objects
obj_id = id(request_body)
@ -627,27 +648,29 @@ def _sanitize_request_body_for_spend_logs_payload(
def _sanitize_value(value: Any) -> Any:
if isinstance(value, dict):
return _sanitize_request_body_for_spend_logs_payload(value, visited)
return _sanitize_request_body_for_spend_logs_payload(
value, visited, max_string_length_prompt_in_db
)
elif isinstance(value, list):
return [_sanitize_value(item) for item in value]
elif isinstance(value, str):
if len(value) > MAX_STRING_LENGTH_PROMPT_IN_DB:
if len(value) > max_string_length_prompt_in_db:
# Keep 35% from beginning and 65% from end (end is usually more important)
# This split ensures we keep more context from the end of conversations
start_ratio = 0.35
end_ratio = 0.65
# Calculate character distribution
start_chars = int(MAX_STRING_LENGTH_PROMPT_IN_DB * start_ratio)
end_chars = int(MAX_STRING_LENGTH_PROMPT_IN_DB * end_ratio)
start_chars = int(max_string_length_prompt_in_db * start_ratio)
end_chars = int(max_string_length_prompt_in_db * end_ratio)
# Ensure we don't exceed the total limit
total_keep = start_chars + end_chars
if total_keep > MAX_STRING_LENGTH_PROMPT_IN_DB:
end_chars = MAX_STRING_LENGTH_PROMPT_IN_DB - start_chars
if total_keep > max_string_length_prompt_in_db:
end_chars = max_string_length_prompt_in_db - start_chars
# If the string length is less than what we want to keep, just truncate normally
if len(value) <= MAX_STRING_LENGTH_PROMPT_IN_DB:
if len(value) <= max_string_length_prompt_in_db:
return value
# Calculate how many characters are being skipped

View file

@ -1,7 +1,8 @@
from typing import Any, Dict, List, Literal, Optional
from typing_extensions import TypedDict
from pydantic import BaseModel
from typing_extensions import TypedDict
from litellm.types.utils import FileTypes
@ -72,6 +73,8 @@ class VideoCreateOptionalRequestParams(TypedDict, total=False):
Params here: https://platform.openai.com/docs/api-reference/videos/create
"""
input_reference: Optional[FileTypes] # File reference for input image
image: Optional[Any] # Image for image-to-video; dict with gcsUri/bytesBase64Encoded, or file-like object
parameters: Optional[Dict[str, Any]] # Provider-specific parameters block passed directly to the API
model: Optional[str]
seconds: Optional[str]
size: Optional[str]

View file

@ -3117,6 +3117,7 @@ def get_optional_params_embeddings( # noqa: PLR0915
custom_llm_provider="",
drop_params: Optional[bool] = None,
additional_drop_params: Optional[List[str]] = None,
allowed_openai_params: Optional[List[str]] = None,
**kwargs,
):
# Lazy load get_supported_openai_params
@ -3131,6 +3132,7 @@ def get_optional_params_embeddings( # noqa: PLR0915
drop_params = passed_params.pop("drop_params", None)
additional_drop_params = passed_params.pop("additional_drop_params", None)
allowed_openai_params = passed_params.pop("allowed_openai_params", None) or []
# Remove function objects from passed_params to avoid JSON serialization errors
passed_params.pop("get_supported_openai_params", None)
@ -3188,11 +3190,11 @@ def get_optional_params_embeddings( # noqa: PLR0915
## raise exception if non-default value passed for non-openai/azure embedding calls
elif custom_llm_provider == "openai":
# 'dimensions` is only supported in `text-embedding-3` and later models
if (
model is not None
and "text-embedding-3" not in model
and "dimensions" in non_default_params.keys()
and "dimensions" not in (allowed_openai_params or [])
):
raise UnsupportedParamsError(
status_code=500,

View file

@ -1,6 +1,6 @@
[tool.poetry]
name = "litellm"
version = "1.81.15"
version = "1.81.16"
description = "Library to easily interface with LLM API providers"
authors = ["BerriAI"]
license = "MIT"
@ -183,7 +183,7 @@ requires = ["poetry-core", "wheel"]
build-backend = "poetry.core.masonry.api"
[tool.commitizen]
version = "1.81.15"
version = "1.81.16"
version_files = [
"pyproject.toml:^version"
]

View file

@ -69,3 +69,40 @@ def test_bedrock_embed_v2_with_drop_params():
)
print(f"received optional_params: {optional_params}")
assert optional_params == {"dimensions": 512, "embeddingTypes": ["binary"]}
def test_openai_non_text_embedding_3_with_allowed_openai_params():
"""
Test that `dimensions` is allowed for non-text-embedding-3 OpenAI models
when `allowed_openai_params=["dimensions"]` is passed. Without this flag,
an UnsupportedParamsError would be raised.
"""
model, custom_llm_provider, _, _ = get_llm_provider(
model="openai/nvidia/llama-3.2-nv-embedqa-1b-v2"
)
optional_params = get_optional_params_embeddings(
model=model,
dimensions=1024,
custom_llm_provider=custom_llm_provider,
allowed_openai_params=["dimensions"],
)
print(f"received optional_params: {optional_params}")
assert optional_params.get("dimensions") == 1024
def test_openai_non_text_embedding_3_without_allowed_openai_params_raises():
"""
Test that passing `dimensions` to a non-text-embedding-3 OpenAI model
without `allowed_openai_params` still raises UnsupportedParamsError.
"""
from litellm.exceptions import UnsupportedParamsError
model, custom_llm_provider, _, _ = get_llm_provider(
model="openai/nvidia/llama-3.2-nv-embedqa-1b-v2"
)
with pytest.raises(UnsupportedParamsError):
get_optional_params_embeddings(
model=model,
dimensions=1024,
custom_llm_provider=custom_llm_provider,
)

View file

@ -0,0 +1,4 @@
{
"model": "BAAI/bge-small-en-v1.5",
"input": ["Hello from litellm!"]
}

View file

@ -4,9 +4,11 @@ import json
import logging
import os
import sys
from typing import Any, Optional
from unittest.mock import MagicMock, patch
import threading
from typing import Any, Optional
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
logging.basicConfig(level=logging.DEBUG)
sys.path.insert(0, os.path.abspath("../.."))
@ -14,6 +16,7 @@ sys.path.insert(0, os.path.abspath("../.."))
import litellm
from litellm import completion
from litellm.caching import InMemoryCache
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
litellm.num_retries = 3
litellm.success_callback = ["langfuse"]
@ -445,6 +448,57 @@ class TestLangfuseLogging:
setup["mock_post"], "completion_with_vertex_call.json", setup["trace_id"]
)
@pytest.mark.asyncio
async def test_langfuse_logging_vllm_embedding(self, mock_setup):
"""
Test that the request sent to the vllm embedding endpoint is correct.
Verifies the request body matches the expected JSON fixture,
including that the hosted_vllm/ prefix is stripped from the model name
and that no unexpected fields (e.g. encoding_format) are included.
"""
setup = mock_setup
vllm_response_data = {
"object": "list",
"data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3]}],
"model": "BAAI/bge-small-en-v1.5",
"usage": {"prompt_tokens": 10, "total_tokens": 10},
}
mock_vllm_response = httpx.Response(
status_code=200,
json=vllm_response_data,
)
mock_async_client = AsyncHTTPHandler()
mock_async_client.post = AsyncMock(return_value=mock_vllm_response)
with patch("httpx.Client.post", setup["mock_post"]):
await litellm.aembedding(
model="hosted_vllm/BAAI/bge-small-en-v1.5",
input=["Hello from litellm!"],
api_base="http://my-fake-vllm.com/v1",
metadata={"trace_id": setup["trace_id"]},
client=mock_async_client,
)
# Verify the request sent to vllm matches the expected JSON fixture
assert mock_async_client.post.call_count == 1
actual_vllm_request = mock_async_client.post.call_args.kwargs["json"]
pwd = os.path.dirname(os.path.realpath(__file__))
expected_body_path = os.path.join(
pwd, "langfuse_expected_request_body", "embedding_with_vllm.json"
)
with open(expected_body_path, "r") as f:
expected_vllm_request = json.load(f)
assert actual_vllm_request == expected_vllm_request, (
f"vllm request body mismatch:\n"
f"actual: {json.dumps(actual_vllm_request, indent=2)}\n"
f"expected: {json.dumps(expected_vllm_request, indent=2)}"
)
@pytest.mark.asyncio
async def test_langfuse_logging_with_router(self, mock_setup):
"""Test Langfuse logging with router"""

View file

@ -234,3 +234,62 @@ class TestProxyMcpSimpleConnections:
)
assert stdio_result == "5"
assert streamable_result == "9"
class TestProxyMcpStatelessBehavior:
"""
Verify that the LiteLLM MCP proxy operates in stateless mode.
When StreamableHTTPSessionManager is configured with stateless=True,
independent clients must be able to connect, list tools, and call tools
without sharing or inheriting session state from other clients.
With stateless=False this fails because the server tracks sessions and
expects clients to supply an mcp-session-id header obtained from a
prior handshake — breaking clients that don't manage session IDs.
Regression test for https://github.com/BerriAI/litellm/issues/20242
"""
@pytest.mark.asyncio
async def test_independent_clients_no_shared_session(
self, proxy_server_url: str
) -> None:
"""Two independent clients connect and operate without sharing session state."""
async with asyncio.timeout(30):
# --- Client A: connect, initialize, call tool ---
async with streamablehttp_client(
url=f"{proxy_server_url}/mcp",
headers={
"Authorization": PROXY_AUTHORIZATION_HEADER,
"x-mcp-servers": "math_stdio",
},
) as (read_a, write_a, _get_sid_a):
async with ClientSession(read_a, write_a) as session_a:
await session_a.initialize()
result_a = await session_a.call_tool(
"add", arguments={"a": 10, "b": 20}
)
assert result_a.content
text_a = getattr(result_a.content[0], "text", None)
assert text_a == "30"
# --- Client B: completely independent connection ---
async with streamablehttp_client(
url=f"{proxy_server_url}/mcp",
headers={
"Authorization": PROXY_AUTHORIZATION_HEADER,
"x-mcp-servers": "math_stdio",
},
) as (read_b, write_b, _get_sid_b):
async with ClientSession(read_b, write_b) as session_b:
await session_b.initialize()
tools = await session_b.list_tools()
assert any(t.name.endswith("add") for t in tools.tools)
result_b = await session_b.call_tool(
"add", arguments={"a": 100, "b": 200}
)
assert result_b.content
text_b = getattr(result_b.content[0], "text", None)
assert text_b == "300"

View file

@ -3668,9 +3668,10 @@ async def test_list_keys(prisma_client):
async def test_key_aliases(prisma_client):
"""
Test the key_aliases function:
- Returns a list
- Returns a paginated response
- Includes alias from a newly created key
- Aliases are unique and sorted
- Aliases are sorted
- Pagination and search params work correctly
"""
import asyncio
import uuid
@ -3682,10 +3683,16 @@ async def test_key_aliases(prisma_client):
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
await litellm.proxy.proxy_server.prisma_client.connect()
# Basic call
response = await key_aliases()
# Basic call - check pagination response shape
response = await key_aliases(page=1, size=50)
assert "aliases" in response
assert isinstance(response["aliases"], list)
assert "total_count" in response
assert "current_page" in response
assert "total_pages" in response
assert "size" in response
assert response["current_page"] == 1
assert response["size"] == 50
# Create a new user (and key) with a unique alias
unique_id = str(uuid.uuid4())
@ -3704,17 +3711,22 @@ async def test_key_aliases(prisma_client):
# Allow async DB writes to settle
await asyncio.sleep(2)
# Call again and validate
response_after = await key_aliases()
# Call again and validate alias is present
response_after = await key_aliases(page=1, size=50)
aliases = response_after["aliases"]
# Contains the new alias
assert test_alias in aliases
# Unique & sorted (endpoint dedupes and orders ascending)
assert len(aliases) == len(set(aliases))
assert aliases == sorted(aliases)
# Search by partial alias
partial = test_alias[:10]
search_response = await key_aliases(page=1, size=50, search=partial)
assert test_alias in search_response["aliases"]
# Search with no match
no_match_response = await key_aliases(page=1, size=50, search="__no_match_xyz__")
assert len(no_match_response["aliases"]) == 0
assert no_match_response["total_count"] == 0
@pytest.mark.skip(reason="Requires reliable external DB connection (prisma).")
@pytest.mark.asyncio

View file

@ -158,3 +158,109 @@ async def test_async_log_success_event_uses_per_model_budget_duration(budget_lim
f"{VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX}:{virtual_key}:{model}:{budget_duration}"
)
assert call_kwargs["response_cost"] == 0.05
# Test is_end_user_within_model_budget
@pytest.mark.asyncio
async def test_is_end_user_within_model_budget(budget_limiter):
# Test when model is within budget
with patch.object(
budget_limiter, "_get_end_user_spend_for_model", return_value=50.0
):
assert (
await budget_limiter.is_end_user_within_model_budget(
"test-user",
{"gpt-4": {"budget_limit": 100.0, "time_period": "1d"}},
"gpt-4",
)
is True
)
# Test when model exceeds budget
with patch.object(
budget_limiter, "_get_end_user_spend_for_model", return_value=150.0
):
with pytest.raises(litellm.BudgetExceededError):
await budget_limiter.is_end_user_within_model_budget(
"test-user",
{"gpt-4": {"budget_limit": 100.0, "time_period": "1d"}},
"gpt-4",
)
# Test model not in budget config
assert (
await budget_limiter.is_end_user_within_model_budget(
"test-user",
{"gpt-4": {"budget_limit": 100.0, "time_period": "1d"}},
"non-existent",
)
is True
)
# Test _get_end_user_spend_for_model
@pytest.mark.asyncio
async def test_get_end_user_spend_for_model(budget_limiter):
budget_config = GenericBudgetInfo(budget_limit=100.0, time_period="1d")
# Mock cache get
with patch.object(budget_limiter.dual_cache, "async_get_cache", return_value=50.0):
spend = await budget_limiter._get_end_user_spend_for_model(
end_user_id="test-user", model="gpt-4", key_budget_config=budget_config
)
assert spend == 50.0
# Test with provider prefix
spend = await budget_limiter._get_end_user_spend_for_model(
end_user_id="test-user",
model="openai/gpt-4",
key_budget_config=budget_config,
)
assert spend == 50.0
@pytest.mark.asyncio
async def test_async_log_success_event_uses_end_user_model_budget_duration(
budget_limiter,
):
"""
async_log_success_event must use the per-model budget_duration for the end user cache key
"""
from litellm.proxy.hooks.model_max_budget_limiter import (
END_USER_SPEND_CACHE_KEY_PREFIX,
)
end_user_id = "test-user"
model = "gpt-4"
budget_duration = "1d"
user_api_key_end_user_model_max_budget = {
model: {"budget_limit": 100.0, "time_period": budget_duration},
}
kwargs = {
"standard_logging_object": {
"response_cost": 0.05,
"model": model,
"end_user": end_user_id,
"metadata": {"user_api_key_end_user_id": end_user_id},
},
"litellm_params": {
"metadata": {
"user_api_key_end_user_model_max_budget": user_api_key_end_user_model_max_budget
},
},
}
with patch.object(
budget_limiter,
"_increment_spend_for_key",
new_callable=AsyncMock,
) as mock_increment:
await budget_limiter.async_log_success_event(
kwargs, response_obj=None, start_time=None, end_time=None
)
mock_increment.assert_awaited_once()
call_kwargs = mock_increment.call_args.kwargs
spend_key = call_kwargs["spend_key"]
assert spend_key == (
f"{END_USER_SPEND_CACHE_KEY_PREFIX}:{end_user_id}:{model}:{budget_duration}"
)
assert call_kwargs["response_cost"] == 0.05

View file

@ -0,0 +1,327 @@
"""
Unit tests for WebSearch Interception with Extended Thinking
Tests that the websearch interception agentic loop correctly handles
thinking/redacted_thinking blocks when extended thinking is enabled.
"""
from unittest.mock import Mock
import pytest
from litellm.integrations.websearch_interception.handler import (
WebSearchInterceptionLogger,
)
from litellm.integrations.websearch_interception.transformation import (
WebSearchTransformation,
)
class TestTransformResponseWithThinking:
"""Tests for _transform_response_anthropic with thinking blocks."""
def test_thinking_blocks_prepended_to_assistant_message(self):
"""Test that thinking blocks are prepended before tool_use blocks."""
tool_calls = [
{
"id": "toolu_01",
"type": "tool_use",
"name": "litellm_web_search",
"input": {"query": "latest news"},
}
]
search_results = [
"Title: News\nURL: https://example.com\nSnippet: Latest news"
]
thinking_blocks = [
{
"type": "thinking",
"thinking": "Let me search for that.",
"signature": "sig123",
},
{"type": "redacted_thinking", "data": "abc123"},
]
assistant_msg, user_msg = (
WebSearchTransformation._transform_response_anthropic(
tool_calls=tool_calls,
search_results=search_results,
thinking_blocks=thinking_blocks,
)
)
# Verify thinking blocks come first
content = assistant_msg["content"]
assert len(content) == 3 # 2 thinking + 1 tool_use
assert content[0]["type"] == "thinking"
assert content[0]["thinking"] == "Let me search for that."
assert content[1]["type"] == "redacted_thinking"
assert content[1]["data"] == "abc123"
assert content[2]["type"] == "tool_use"
assert content[2]["id"] == "toolu_01"
def test_no_thinking_blocks_backward_compat(self):
"""Test that transform works without thinking blocks (backward compat)."""
tool_calls = [
{
"id": "toolu_01",
"type": "tool_use",
"name": "litellm_web_search",
"input": {"query": "test"},
}
]
search_results = ["Search result text"]
# No thinking_blocks param (default None)
assistant_msg, _ = (
WebSearchTransformation._transform_response_anthropic(
tool_calls=tool_calls,
search_results=search_results,
)
)
content = assistant_msg["content"]
assert len(content) == 1
assert content[0]["type"] == "tool_use"
def test_empty_thinking_blocks_list(self):
"""Test that an empty thinking_blocks list behaves like None."""
tool_calls = [
{
"id": "toolu_01",
"type": "tool_use",
"name": "litellm_web_search",
"input": {"query": "test"},
}
]
search_results = ["Search result text"]
assistant_msg, _ = (
WebSearchTransformation._transform_response_anthropic(
tool_calls=tool_calls,
search_results=search_results,
thinking_blocks=[],
)
)
content = assistant_msg["content"]
assert len(content) == 1
assert content[0]["type"] == "tool_use"
def test_transform_response_passes_thinking_to_anthropic(self):
"""Test that transform_response routes thinking_blocks correctly."""
tool_calls = [
{
"id": "toolu_01",
"type": "tool_use",
"name": "litellm_web_search",
"input": {"query": "test"},
}
]
search_results = ["Search result"]
thinking_blocks = [
{
"type": "thinking",
"thinking": "Reasoning here.",
"signature": "sig",
},
]
assistant_msg, _ = WebSearchTransformation.transform_response(
tool_calls=tool_calls,
search_results=search_results,
response_format="anthropic",
thinking_blocks=thinking_blocks,
)
content = assistant_msg["content"]
assert content[0]["type"] == "thinking"
assert content[1]["type"] == "tool_use"
def test_transform_response_openai_ignores_thinking(self):
"""Test that OpenAI format is unaffected by thinking_blocks param."""
tool_calls = [
{
"id": "call_01",
"type": "function",
"name": "litellm_web_search",
"function": {
"name": "litellm_web_search",
"arguments": {"query": "test"},
},
"input": {"query": "test"},
}
]
search_results = ["Search result"]
thinking_blocks = [
{
"type": "thinking",
"thinking": "Should not appear.",
"signature": "sig",
},
]
assistant_msg, _ = WebSearchTransformation.transform_response(
tool_calls=tool_calls,
search_results=search_results,
response_format="openai",
thinking_blocks=thinking_blocks,
)
# OpenAI format uses tool_calls key, not content — thinking is irrelevant
assert "tool_calls" in assistant_msg
assert "content" not in assistant_msg
class TestAgenticLoopThinkingExtraction:
"""Tests for thinking block extraction in async_should_run_agentic_loop."""
@pytest.mark.asyncio
async def test_extracts_thinking_blocks_from_dict_response(self):
"""Test extraction of thinking blocks from dict-style response."""
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"])
response = {
"content": [
{
"type": "thinking",
"thinking": "Let me think...",
"signature": "sig1",
},
{"type": "redacted_thinking", "data": "redacted_data"},
{
"type": "tool_use",
"id": "toolu_01",
"name": "litellm_web_search",
"input": {"query": "latest news"},
},
]
}
should_run, tools_dict = await logger.async_should_run_agentic_loop(
response=response,
model="bedrock/claude",
messages=[],
tools=[{"name": "WebSearch"}],
stream=False,
custom_llm_provider="bedrock",
kwargs={},
)
assert should_run is True
assert len(tools_dict["tool_calls"]) == 1
assert len(tools_dict["thinking_blocks"]) == 2
assert tools_dict["thinking_blocks"][0]["type"] == "thinking"
assert tools_dict["thinking_blocks"][0]["thinking"] == "Let me think..."
assert tools_dict["thinking_blocks"][1]["type"] == "redacted_thinking"
assert tools_dict["thinking_blocks"][1]["data"] == "redacted_data"
@pytest.mark.asyncio
async def test_extracts_thinking_blocks_from_object_response(self):
"""Test extraction of thinking blocks from non-dict response objects.
In practice, the Anthropic pass-through always returns plain dicts
(TypedDict(**raw_json) produces a dict). This test covers the safety
branch for non-dict response objects.
"""
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"])
# Simulate object-style response blocks
thinking_block = Mock()
thinking_block.type = "thinking"
thinking_block.thinking = "Reasoning..."
thinking_block.signature = "sig"
redacted_block = Mock()
redacted_block.type = "redacted_thinking"
redacted_block.data = "abc"
tool_block = Mock()
tool_block.type = "tool_use"
tool_block.name = "litellm_web_search"
tool_block.id = "toolu_01"
tool_block.input = {"query": "test"}
response = Mock()
response.content = [thinking_block, redacted_block, tool_block]
should_run, tools_dict = await logger.async_should_run_agentic_loop(
response=response,
model="bedrock/claude",
messages=[],
tools=[{"name": "WebSearch"}],
stream=False,
custom_llm_provider="bedrock",
kwargs={},
)
assert should_run is True
assert len(tools_dict["thinking_blocks"]) == 2
# Verify getattr-based conversion produced correct dicts
assert tools_dict["thinking_blocks"][0] == {
"type": "thinking",
"thinking": "Reasoning...",
"signature": "sig",
}
assert tools_dict["thinking_blocks"][1] == {
"type": "redacted_thinking",
"data": "abc",
}
@pytest.mark.asyncio
async def test_no_thinking_blocks_when_thinking_disabled(self):
"""Test that thinking_blocks is empty when response has no thinking."""
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"])
response = {
"content": [
{
"type": "tool_use",
"id": "toolu_01",
"name": "litellm_web_search",
"input": {"query": "test"},
},
]
}
should_run, tools_dict = await logger.async_should_run_agentic_loop(
response=response,
model="bedrock/claude",
messages=[],
tools=[{"name": "WebSearch"}],
stream=False,
custom_llm_provider="bedrock",
kwargs={},
)
assert should_run is True
assert tools_dict["thinking_blocks"] == []
@pytest.mark.asyncio
async def test_thinking_blocks_not_extracted_when_no_tool_calls(self):
"""Test that no extraction happens when no websearch tool calls found."""
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"])
response = {
"content": [
{
"type": "thinking",
"thinking": "Just thinking...",
"signature": "sig",
},
{"type": "text", "text": "Here is my response."},
]
}
should_run, tools_dict = await logger.async_should_run_agentic_loop(
response=response,
model="bedrock/claude",
messages=[],
tools=[{"name": "WebSearch"}],
stream=False,
custom_llm_provider="bedrock",
kwargs={},
)
assert should_run is False
assert tools_dict == {}

View file

@ -1813,6 +1813,51 @@ def test_translate_openai_response_to_anthropic_input_tokens_no_cache():
assert anthropic_response["usage"]["output_tokens"] == 50
def test_translate_openai_response_to_anthropic_cache_tokens_from_prompt_tokens_details():
"""
OpenAI/Azure providers set prompt_tokens_details.cached_tokens but not
_cache_read_input_tokens. The adapter should populate cache_read_input_tokens
from prompt_tokens_details.cached_tokens directly.
"""
from litellm.types.utils import PromptTokensDetailsWrapper
# OpenAI-style usage: only prompt_tokens_details, no cache_read_input_tokens kwarg
usage = Usage(
prompt_tokens=100,
completion_tokens=50,
total_tokens=150,
prompt_tokens_details=PromptTokensDetailsWrapper(
cached_tokens=30
),
)
response = ModelResponse(
id="test-id",
choices=[
Choices(
index=0,
finish_reason="stop",
message=Message(
role="assistant",
content="Test response",
),
)
],
model="gpt-4o-2024-08-06",
usage=usage,
)
adapter = LiteLLMAnthropicMessagesAdapter()
anthropic_response = adapter.translate_openai_response_to_anthropic(
response=response,
tool_name_mapping=None,
)
assert anthropic_response["usage"]["input_tokens"] == 70
assert anthropic_response["usage"]["output_tokens"] == 50
assert anthropic_response["usage"]["cache_read_input_tokens"] == 30
# =====================================================================
# Web Search Tool Transformation Tests
# =====================================================================

View file

@ -1,19 +1,20 @@
"""
Tests for Vertex AI (Veo) video generation transformation.
"""
import base64
import json
import os
import pytest
from unittest.mock import Mock, MagicMock, patch
from unittest.mock import MagicMock, Mock, patch
import httpx
import base64
import pytest
from litellm.llms.vertex_ai.videos.transformation import (
VertexAIVideoConfig,
_convert_image_to_vertex_format,
)
from litellm.types.videos.main import VideoObject
from litellm.types.router import GenericLiteLLMParams
from litellm.types.videos.main import VideoObject
class TestVertexAIVideoConfig:
@ -548,3 +549,170 @@ class TestConvertImageToVertexFormat:
decoded = base64.b64decode(result["bytesBase64Encoded"])
assert decoded == fake_image_data
class TestImageAndParametersPassthrough:
"""
Tests that image (gcsUri / bare gs:// / file-like) and a pre-built
parameters dict are correctly forwarded through map_openai_params and
transform_video_create_request.
"""
def setup_method(self):
self.config = VertexAIVideoConfig()
self.api_base = (
"https://us-central1-aiplatform.googleapis.com/v1/projects/"
"test-project/locations/us-central1/publishers/google/models/veo-002"
)
# ------------------------------------------------------------------ #
# map_openai_params #
# ------------------------------------------------------------------ #
def test_map_openai_params_passes_image_dict(self):
"""image dict (gcsUri format) is forwarded as-is."""
image = {"gcsUri": "gs://my-bucket/boardwalk.jpg"}
mapped = self.config.map_openai_params(
video_create_optional_params={"image": image},
model="veo-002",
drop_params=False,
)
assert mapped["image"] == image
def test_map_openai_params_passes_parameters_dict(self):
"""A pre-built parameters dict is forwarded as-is."""
params = {"sampleCount": 1, "videoLengthSeconds": 5, "aspectRatio": "16:9"}
mapped = self.config.map_openai_params(
video_create_optional_params={"parameters": params},
model="veo-002",
drop_params=False,
)
assert mapped["parameters"] == params
def test_map_openai_params_input_reference_takes_priority_over_image(self):
"""input_reference wins over a directly passed image key."""
mock_file = Mock()
image_dict = {"gcsUri": "gs://my-bucket/other.jpg"}
mapped = self.config.map_openai_params(
video_create_optional_params={
"input_reference": mock_file,
"image": image_dict,
},
model="veo-002",
drop_params=False,
)
assert mapped["image"] is mock_file
# ------------------------------------------------------------------ #
# transform_video_create_request – image forms #
# ------------------------------------------------------------------ #
def test_transform_request_image_gcs_uri_dict(self):
"""image passed as {"gcsUri": "gs://..."} is placed in instances as-is."""
image = {"gcsUri": "gs://my-bucket/boardwalk.jpg"}
data, _, url = self.config.transform_video_create_request(
model="veo-002",
prompt="Cinematic drone shot",
api_base=self.api_base,
video_create_optional_request_params={"image": image},
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert data["instances"][0]["image"] == image
assert url.endswith(":predictLongRunning")
def test_transform_request_image_bare_gs_uri_string(self):
"""A bare gs:// string is wrapped in {"gcsUri": ...} without downloading."""
gs_uri = "gs://my-bucket/boardwalk.jpg"
data, _, _ = self.config.transform_video_create_request(
model="veo-002",
prompt="Cinematic drone shot",
api_base=self.api_base,
video_create_optional_request_params={"image": gs_uri},
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert data["instances"][0]["image"] == {"gcsUri": gs_uri}
def test_transform_request_image_bytes_base64_dict(self):
"""image already in bytesBase64Encoded format is passed through unchanged."""
image = {"bytesBase64Encoded": "abc123", "mimeType": "image/jpeg"}
data, _, _ = self.config.transform_video_create_request(
model="veo-002",
prompt="Cinematic drone shot",
api_base=self.api_base,
video_create_optional_request_params={"image": image},
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert data["instances"][0]["image"] == image
# ------------------------------------------------------------------ #
# transform_video_create_request – parameters dict #
# ------------------------------------------------------------------ #
def test_transform_request_parameters_dict_not_double_nested(self):
"""A pre-built parameters dict becomes request_data["parameters"] directly."""
params = {"sampleCount": 1, "videoLengthSeconds": 5, "aspectRatio": "16:9"}
data, _, _ = self.config.transform_video_create_request(
model="veo-002",
prompt="Cinematic drone shot",
api_base=self.api_base,
video_create_optional_request_params={"parameters": params},
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert data["parameters"] == params
# Must NOT be double-nested
assert "parameters" not in data["parameters"]
# ------------------------------------------------------------------ #
# Full user scenario #
# ------------------------------------------------------------------ #
def test_transform_request_full_user_scenario(self):
"""
Reproduces the exact user request:
image: {"gcsUri": "gs://your-bucket-name/path/to/boardwalk.jpg"}
parameters: {"sampleCount": 1, "videoLengthSeconds": 5,
"aspectRatio": "16:9", "storageUri": "gs://test/outputs/"}
"""
image = {"gcsUri": "gs://your-bucket-name/path/to/boardwalk.jpg"}
parameters = {
"sampleCount": 1,
"videoLengthSeconds": 5,
"aspectRatio": "16:9",
"storageUri": "gs://test/outputs/",
}
# Simulate the full pipeline: map_openai_params → transform_video_create_request
mapped = self.config.map_openai_params(
video_create_optional_params={"image": image, "parameters": parameters},
model="veo-3.1-generate-preview",
drop_params=False,
)
data, _, url = self.config.transform_video_create_request(
model="veo-3.1-generate-preview",
prompt="Cinematic drone shot moving forward along the beach boardwalk",
api_base=self.api_base,
video_create_optional_request_params=mapped,
litellm_params=GenericLiteLLMParams(),
headers={},
)
# instances contains prompt + image
assert len(data["instances"]) == 1
instance = data["instances"][0]
assert instance["prompt"] == "Cinematic drone shot moving forward along the beach boardwalk"
assert instance["image"] == image
# parameters block is correct and not double-nested
assert data["parameters"] == parameters
assert "parameters" not in data["parameters"]
assert url.endswith(":predictLongRunning")

View file

@ -486,7 +486,7 @@ async def test_get_tools_from_mcp_servers_continues_when_one_server_fails():
working_server if server_id == "working_server" else failing_server
)
# Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering)
mock_manager.filter_server_ids_by_ip = lambda server_ids, client_ip: server_ids
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0)
async def mock_get_tools_from_server(
server,
@ -588,7 +588,7 @@ async def test_get_tools_from_mcp_servers_handles_all_servers_failing():
failing_server1 if server_id == "failing_server1" else failing_server2
)
# Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering)
mock_manager.filter_server_ids_by_ip = lambda server_ids, client_ip: server_ids
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0)
async def mock_get_tools_from_server(
server,
@ -1035,7 +1035,7 @@ async def test_list_tools_single_server_unprefixed_names():
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1"])
mock_manager.get_mcp_server_by_id = MagicMock(return_value=server)
# Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering)
mock_manager.filter_server_ids_by_ip = lambda server_ids, client_ip: server_ids
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0)
async def mock_get_tools_from_server(
server,
@ -1113,7 +1113,7 @@ async def test_list_tools_multiple_servers_prefixed_names():
server1 if server_id == "server1" else server2
)
# Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering)
mock_manager.filter_server_ids_by_ip = lambda server_ids, client_ip: server_ids
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0)
async def mock_get_tools_from_server(
server,
@ -1364,7 +1364,7 @@ async def test_list_tools_filters_by_key_team_permissions():
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1"])
mock_manager.get_mcp_server_by_id = lambda server_id: server
# Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering)
mock_manager.filter_server_ids_by_ip = lambda server_ids, client_ip: server_ids
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0)
async def mock_get_tools_from_server(
server,
@ -1471,7 +1471,7 @@ async def test_list_tools_with_team_tool_permissions_inheritance():
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1"])
mock_manager.get_mcp_server_by_id = lambda server_id: server
# Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering)
mock_manager.filter_server_ids_by_ip = lambda server_ids, client_ip: server_ids
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0)
async def mock_get_tools_from_server(
server,
@ -1563,7 +1563,7 @@ async def test_list_tools_with_no_tool_permissions_shows_all():
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1"])
mock_manager.get_mcp_server_by_id = lambda server_id: server
# Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering)
mock_manager.filter_server_ids_by_ip = lambda server_ids, client_ip: server_ids
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0)
async def mock_get_tools_from_server(
server,
@ -1658,7 +1658,7 @@ async def test_list_tools_strips_prefix_when_matching_permissions():
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["gitmcp_server"])
mock_manager.get_mcp_server_by_id = MagicMock(return_value=server)
# Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering)
mock_manager.filter_server_ids_by_ip = lambda server_ids, client_ip: server_ids
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0)
async def mock_get_tools_from_server(
server,

View file

@ -0,0 +1,119 @@
import pytest
from unittest.mock import AsyncMock, patch
import litellm
from litellm.proxy.auth.user_api_key_auth import (
_run_post_custom_auth_checks,
update_valid_token_with_end_user_params,
)
from litellm.proxy._types import UserAPIKeyAuth
@pytest.mark.asyncio
async def test_custom_auth_run_post_custom_auth_checks_without_end_user_id():
# Test backwards compatibility
valid_token = UserAPIKeyAuth(token="test_token")
with patch(
"litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock
) as mock_common:
mock_common.return_value = True
result = await _run_post_custom_auth_checks(
valid_token=valid_token,
request=None,
request_data={},
route="/v1/chat/completions",
parent_otel_span=None,
)
assert result.token == "test_token"
assert getattr(result, "end_user_id", None) is None
mock_common.assert_awaited_once()
@pytest.mark.asyncio
async def test_custom_auth_run_post_custom_auth_checks_with_end_user_budget_exceeded():
valid_token = UserAPIKeyAuth(
token="test_token",
end_user_id="test_user",
end_user_model_max_budget={
"gpt-4": {"budget_limit": 10.0, "time_period": "1d"}
},
)
request_data = {"model": "gpt-4"}
with patch(
"litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock
):
with patch(
"litellm.proxy.proxy_server.model_max_budget_limiter.is_end_user_within_model_budget",
new_callable=AsyncMock,
) as mock_budget_check:
mock_budget_check.side_effect = litellm.BudgetExceededError(
message="Exceeded budget", current_cost=20.0, max_budget=10.0
)
with pytest.raises(litellm.BudgetExceededError):
await _run_post_custom_auth_checks(
valid_token=valid_token,
request=None,
request_data=request_data,
route="/v1/chat/completions",
parent_otel_span=None,
)
mock_budget_check.assert_awaited_once()
def test_update_valid_token_does_not_override_custom_auth_values_with_none():
"""
Greptile feedback: if custom auth sets end_user_model_max_budget on the token,
but the DB end_user has no model_max_budget in their budget table, the DB lookup
should NOT clear the custom-auth-provided value.
"""
custom_auth_budget = {"gpt-4": {"budget_limit": 5.0, "time_period": "1d"}}
valid_token = UserAPIKeyAuth(
token="test_token",
end_user_id="user_1",
end_user_tpm_limit=100,
end_user_rpm_limit=50,
end_user_model_max_budget=custom_auth_budget,
)
# Simulate DB lookup that found the end_user but budget table has no limits set
end_user_params = {
"end_user_id": "user_1",
"allowed_model_region": None,
# No tpm_limit, rpm_limit, or model_max_budget from DB
}
result = update_valid_token_with_end_user_params(valid_token, end_user_params)
# Custom-auth-provided values should be preserved, not cleared to None
assert result.end_user_tpm_limit == 100
assert result.end_user_rpm_limit == 50
assert result.end_user_model_max_budget == custom_auth_budget
assert result.end_user_id == "user_1"
def test_update_valid_token_db_values_override_custom_auth_when_set():
"""
When the DB budget table has explicit values, they should override
whatever the custom auth function set (DB is source of truth).
"""
valid_token = UserAPIKeyAuth(
token="test_token",
end_user_id="user_1",
end_user_tpm_limit=100,
end_user_model_max_budget={"gpt-4": {"budget_limit": 5.0, "time_period": "1d"}},
)
db_budget = {"gpt-4": {"budget_limit": 20.0, "time_period": "1d"}}
end_user_params = {
"end_user_id": "user_1",
"end_user_tpm_limit": 500,
"end_user_model_max_budget": db_budget,
}
result = update_valid_token_with_end_user_params(valid_token, end_user_params)
# DB values should win
assert result.end_user_tpm_limit == 500
assert result.end_user_model_max_budget == db_budget

View file

@ -91,3 +91,58 @@ class TestMCPServerIPFiltering:
result = manager.filter_server_ids_by_ip(["priv"], client_ip=None)
assert result == ["priv"]
class TestFilterServerIdsByIpWithInfo:
"""Tests that filter_server_ids_by_ip_with_info returns accurate block counts."""
@patch("litellm.public_mcp_servers", [])
@patch("litellm.proxy.proxy_server.general_settings", {})
def test_external_ip_reports_blocked_count(self):
pub = _make_server("pub", available_on_public_internet=True)
priv = _make_server("priv", available_on_public_internet=False)
manager = _make_manager([pub, priv])
allowed, blocked = manager.filter_server_ids_by_ip_with_info(
["pub", "priv"], client_ip="8.8.8.8"
)
assert allowed == ["pub"]
assert blocked == 1
@patch("litellm.public_mcp_servers", [])
@patch("litellm.proxy.proxy_server.general_settings", {})
def test_internal_ip_reports_zero_blocked(self):
pub = _make_server("pub", available_on_public_internet=True)
priv = _make_server("priv", available_on_public_internet=False)
manager = _make_manager([pub, priv])
allowed, blocked = manager.filter_server_ids_by_ip_with_info(
["pub", "priv"], client_ip="192.168.1.1"
)
assert allowed == ["pub", "priv"]
assert blocked == 0
@patch("litellm.public_mcp_servers", [])
@patch("litellm.proxy.proxy_server.general_settings", {})
def test_no_ip_returns_all_with_zero_blocked(self):
priv = _make_server("priv", available_on_public_internet=False)
manager = _make_manager([priv])
allowed, blocked = manager.filter_server_ids_by_ip_with_info(
["priv"], client_ip=None
)
assert allowed == ["priv"]
assert blocked == 0
@patch("litellm.public_mcp_servers", [])
@patch("litellm.proxy.proxy_server.general_settings", {})
def test_all_private_external_ip_reports_all_blocked(self):
priv1 = _make_server("priv1", available_on_public_internet=False)
priv2 = _make_server("priv2", available_on_public_internet=False)
manager = _make_manager([priv1, priv2])
allowed, blocked = manager.filter_server_ids_by_ip_with_info(
["priv1", "priv2"], client_ip="1.2.3.4"
)
assert allowed == []
assert blocked == 2

View file

@ -45,6 +45,7 @@ from litellm.proxy.management_endpoints.key_management_endpoints import (
check_team_key_model_specific_limits,
delete_verification_tokens,
generate_key_helper_fn,
key_aliases,
list_keys,
prepare_key_update_data,
reset_key_spend_fn,
@ -6210,3 +6211,97 @@ async def test_generate_key_helper_fn_agent_id():
assert key_data.get("agent_id") == "test-agent-456", (
f"Expected agent_id='test-agent-456' in key_data, got: {key_data.get('agent_id')}"
)
@pytest.mark.asyncio
async def test_key_aliases_response_shape():
"""Test that key_aliases returns the correct paginated response shape."""
mock_prisma_client = AsyncMock()
mock_prisma_client.db.query_raw = AsyncMock(
side_effect=[
[{"count": 2}],
[{"key_alias": "alias-alpha"}, {"key_alias": "alias-beta"}],
]
)
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client):
result = await key_aliases(page=1, size=50, search=None)
assert result["aliases"] == ["alias-alpha", "alias-beta"]
assert result["total_count"] == 2
assert result["current_page"] == 1
assert result["total_pages"] == 1
assert result["size"] == 50
# Both SQL calls must filter out null/empty aliases
count_sql = mock_prisma_client.db.query_raw.call_args_list[0].args[0]
aliases_sql = mock_prisma_client.db.query_raw.call_args_list[1].args[0]
assert "key_alias IS NOT NULL" in count_sql
assert "key_alias IS NOT NULL" in aliases_sql
@pytest.mark.asyncio
async def test_key_aliases_pagination_skip_take():
"""Test that LIMIT and OFFSET are correctly derived from page and size."""
mock_prisma_client = AsyncMock()
mock_prisma_client.db.query_raw = AsyncMock(
side_effect=[
[{"count": 120}],
[],
]
)
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client):
result = await key_aliases(page=3, size=25, search=None)
assert result["current_page"] == 3
assert result["size"] == 25
assert result["total_count"] == 120
assert result["total_pages"] == 5 # ceil(120 / 25)
# aliases query params: [UI_SESSION_TOKEN_TEAM_ID, size=25, offset=50]
aliases_call_args = mock_prisma_client.db.query_raw.call_args_list[1].args
assert aliases_call_args[-2] == 25 # LIMIT = size
assert aliases_call_args[-1] == 50 # OFFSET = (3 - 1) * 25
@pytest.mark.asyncio
async def test_key_aliases_search_filter():
"""Test that the search param adds a case-insensitive ILIKE condition."""
mock_prisma_client = AsyncMock()
mock_prisma_client.db.query_raw = AsyncMock(
side_effect=[
[{"count": 0}],
[],
]
)
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client):
await key_aliases(page=1, size=50, search="my-key")
count_call = mock_prisma_client.db.query_raw.call_args_list[0]
count_sql = count_call.args[0]
count_params = count_call.args[1:]
assert "ILIKE" in count_sql
assert "%my-key%" in count_params
@pytest.mark.asyncio
async def test_key_aliases_no_search_omits_ilike_filter():
"""Test that without a search term no ILIKE condition is added."""
mock_prisma_client = AsyncMock()
mock_prisma_client.db.query_raw = AsyncMock(
side_effect=[
[{"count": 0}],
[],
]
)
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client):
await key_aliases(page=1, size=50, search=None)
count_sql = mock_prisma_client.db.query_raw.call_args_list[0].args[0]
assert "ILIKE" not in count_sql

View file

@ -161,6 +161,21 @@ def test_sanitize_request_body_for_spend_logs_payload_mixed_types():
assert len(sanitized["nested"]["dict"]["key"]) == expected_length
def test_sanitize_request_body_for_spend_logs_payload_uses_runtime_env_override(
monkeypatch: pytest.MonkeyPatch,
):
from litellm.constants import MAX_STRING_LENGTH_PROMPT_IN_DB
override_max = max(MAX_STRING_LENGTH_PROMPT_IN_DB + 1000, 6000)
test_string = "a" * (MAX_STRING_LENGTH_PROMPT_IN_DB + 500)
# Simulate config-loaded env var being set after module import.
monkeypatch.setenv("MAX_STRING_LENGTH_PROMPT_IN_DB", str(override_max))
sanitized = _sanitize_request_body_for_spend_logs_payload({"text": test_string})
assert sanitized["text"] == test_string
def test_sanitize_request_body_for_spend_logs_payload_circular_reference():
# Create a circular reference
a: dict[str, Any] = {}
@ -1349,4 +1364,3 @@ def test_get_logging_payload_includes_request_duration_ms():
)
assert payload["request_duration_ms"] == 3000