mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge branch 'BerriAI:main' into main
This commit is contained in:
commit
9c255bca4e
32 changed files with 1840 additions and 146 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -269,6 +269,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
"logprobs",
|
||||
"top_logprobs",
|
||||
"modalities",
|
||||
"audio",
|
||||
"parallel_tool_calls",
|
||||
"web_search_options",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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={
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,4 @@
|
|||
{
|
||||
"model": "BAAI/bge-small-en-v1.5",
|
||||
"input": ["Hello from litellm!"]
|
||||
}
|
||||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 == {}
|
||||
|
|
@ -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
|
||||
# =====================================================================
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue