diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index d7858d71eb3..bef8925e8e9 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -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 diff --git a/litellm/integrations/websearch_interception/transformation.py b/litellm/integrations/websearch_interception/transformation.py index e44ec35c3a2..e016899e0c3 100644 --- a/litellm/integrations/websearch_interception/transformation.py +++ b/litellm/integrations/websearch_interception/transformation.py @@ -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 diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py index 8b21569546e..a7362a94312 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py @@ -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( diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index d248d2862e8..7bcefc1dd87 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -269,6 +269,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): "logprobs", "top_logprobs", "modalities", + "audio", "parallel_tool_calls", "web_search_options", ] diff --git a/litellm/llms/vertex_ai/videos/transformation.py b/litellm/llms/vertex_ai/videos/transformation.py index 8cdccc4cd64..60852c1bf02 100644 --- a/litellm/llms/vertex_ai/videos/transformation.py +++ b/litellm/llms/vertex_ai/videos/transformation.py @@ -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" diff --git a/litellm/main.py b/litellm/main.py index 8b239c454f4..cb3ddc2f401 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -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, ) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 7e90c64efdc..7484de33ce4 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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]: """ diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 6e1a252be73..16f8f835430 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -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={ diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index e8877b4fff7..48c837c1e4f 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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( diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index de1609baf62..53513f7f522 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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] diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 133f8ec136d..3e2378ada60 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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 diff --git a/litellm/proxy/hooks/model_max_budget_limiter.py b/litellm/proxy/hooks/model_max_budget_limiter.py index b8c073dd061..5e48ef2879e 100644 --- a/litellm/proxy/hooks/model_max_budget_limiter.py +++ b/litellm/proxy/hooks/model_max_budget_limiter.py @@ -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( diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index d2312a00c3b..7ebf9a4caee 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -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 diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index b489369071f..5b56133f1ce 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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}") diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index 1ef85562335..31615a768d7 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -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 diff --git a/litellm/types/videos/main.py b/litellm/types/videos/main.py index 65c4cfe0e00..8e595db39fc 100644 --- a/litellm/types/videos/main.py +++ b/litellm/types/videos/main.py @@ -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] diff --git a/litellm/utils.py b/litellm/utils.py index 7ac828aefc1..81e772b1765 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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, diff --git a/pyproject.toml b/pyproject.toml index c848bc3e789..8eb433fb064 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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" ] diff --git a/tests/local_testing/test_get_optional_params_embeddings.py b/tests/local_testing/test_get_optional_params_embeddings.py index 055be487551..8a94c8f4682 100644 --- a/tests/local_testing/test_get_optional_params_embeddings.py +++ b/tests/local_testing/test_get_optional_params_embeddings.py @@ -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, + ) diff --git a/tests/logging_callback_tests/langfuse_expected_request_body/embedding_with_vllm.json b/tests/logging_callback_tests/langfuse_expected_request_body/embedding_with_vllm.json new file mode 100644 index 00000000000..f8621adadbb --- /dev/null +++ b/tests/logging_callback_tests/langfuse_expected_request_body/embedding_with_vllm.json @@ -0,0 +1,4 @@ +{ + "model": "BAAI/bge-small-en-v1.5", + "input": ["Hello from litellm!"] +} diff --git a/tests/logging_callback_tests/test_langfuse_e2e_test.py b/tests/logging_callback_tests/test_langfuse_e2e_test.py index 866e95a702c..369988ed405 100644 --- a/tests/logging_callback_tests/test_langfuse_e2e_test.py +++ b/tests/logging_callback_tests/test_langfuse_e2e_test.py @@ -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""" diff --git a/tests/mcp_tests/test_proxy_mcp_e2e.py b/tests/mcp_tests/test_proxy_mcp_e2e.py index 2b8cde54710..5aff63a51ef 100644 --- a/tests/mcp_tests/test_proxy_mcp_e2e.py +++ b/tests/mcp_tests/test_proxy_mcp_e2e.py @@ -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" + diff --git a/tests/proxy_unit_tests/test_key_generate_prisma.py b/tests/proxy_unit_tests/test_key_generate_prisma.py index c3f68762810..ed528f21e0d 100644 --- a/tests/proxy_unit_tests/test_key_generate_prisma.py +++ b/tests/proxy_unit_tests/test_key_generate_prisma.py @@ -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 diff --git a/tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py b/tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py index 352db384c88..030d452e55f 100644 --- a/tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py +++ b/tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py @@ -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 diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_thinking.py b/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_thinking.py new file mode 100644 index 00000000000..8093ce6fc12 --- /dev/null +++ b/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_thinking.py @@ -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 == {} diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py index 1ea1374cfb3..839d032c436 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py @@ -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 # ===================================================================== diff --git a/tests/test_litellm/llms/vertex_ai/videos/test_vertex_video_transformation.py b/tests/test_litellm/llms/vertex_ai/videos/test_vertex_video_transformation.py index 7ae344e4999..f0c5899e047 100644 --- a/tests/test_litellm/llms/vertex_ai/videos/test_vertex_video_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/videos/test_vertex_video_transformation.py @@ -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") + diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index b5f2197407a..1fcdeb627d0 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -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, diff --git a/tests/test_litellm/proxy/auth/test_custom_auth_end_user_budget.py b/tests/test_litellm/proxy/auth/test_custom_auth_end_user_budget.py new file mode 100644 index 00000000000..73a97188424 --- /dev/null +++ b/tests/test_litellm/proxy/auth/test_custom_auth_end_user_budget.py @@ -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 diff --git a/tests/test_litellm/proxy/auth/test_mcp_ip_filtering.py b/tests/test_litellm/proxy/auth/test_mcp_ip_filtering.py index 50e51fbe035..3b13ef3641f 100644 --- a/tests/test_litellm/proxy/auth/test_mcp_ip_filtering.py +++ b/tests/test_litellm/proxy/auth/test_mcp_ip_filtering.py @@ -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 diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index fb71adc1085..7565e901ecd 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -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 + + diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py index 30257cd2da8..24f45cc5c91 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py @@ -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 -