diff --git a/.github/workflows/test-unit-proxy-endpoints.yml b/.github/workflows/test-unit-proxy-endpoints.yml index 6b34a08a8e0..0a9513ec024 100644 --- a/.github/workflows/test-unit-proxy-endpoints.yml +++ b/.github/workflows/test-unit-proxy-endpoints.yml @@ -36,11 +36,13 @@ jobs: tests/test_litellm/proxy/a2a tests/test_litellm/proxy/discovery_endpoints tests/test_litellm/proxy/health_endpoints + tests/test_litellm/proxy/shutdown tests/test_litellm/proxy/public_endpoints tests/test_litellm/proxy/prompts tests/test_litellm/proxy/rag_endpoints tests/test_litellm/proxy/realtime_endpoints tests/test_litellm/proxy/ui_crud_endpoints + tests/test_litellm/proxy/utils workers: 2 reruns: 2 artifact-name: proxy-endpoints diff --git a/Makefile b/Makefile index 5dbd308a3e2..a00a90da601 100644 --- a/Makefile +++ b/Makefile @@ -146,7 +146,7 @@ test-unit-proxy-core: install-test-deps $(UV_RUN) pytest tests/test_litellm/proxy/auth tests/test_litellm/proxy/client tests/test_litellm/proxy/db tests/test_litellm/proxy/hooks tests/test_litellm/proxy/policy_engine --tb=short -vv -n 4 --durations=20 test-unit-proxy-misc: install-test-deps - $(UV_RUN) pytest tests/test_litellm/proxy/_experimental tests/test_litellm/proxy/agent_endpoints tests/test_litellm/proxy/anthropic_endpoints tests/test_litellm/proxy/common_utils tests/test_litellm/proxy/discovery_endpoints tests/test_litellm/proxy/experimental tests/test_litellm/proxy/google_endpoints tests/test_litellm/proxy/health_endpoints tests/test_litellm/proxy/image_endpoints tests/test_litellm/proxy/middleware tests/test_litellm/proxy/openai_files_endpoint tests/test_litellm/proxy/pass_through_endpoints tests/test_litellm/proxy/prompts tests/test_litellm/proxy/public_endpoints tests/test_litellm/proxy/response_api_endpoints tests/test_litellm/proxy/spend_tracking tests/test_litellm/proxy/ui_crud_endpoints tests/test_litellm/proxy/vector_store_endpoints tests/test_litellm/proxy/test_*.py --tb=short -vv -n 4 --durations=20 + $(UV_RUN) pytest tests/test_litellm/proxy/_experimental tests/test_litellm/proxy/agent_endpoints tests/test_litellm/proxy/anthropic_endpoints tests/test_litellm/proxy/common_utils tests/test_litellm/proxy/discovery_endpoints tests/test_litellm/proxy/experimental tests/test_litellm/proxy/google_endpoints tests/test_litellm/proxy/health_endpoints tests/test_litellm/proxy/image_endpoints tests/test_litellm/proxy/middleware tests/test_litellm/proxy/openai_files_endpoint tests/test_litellm/proxy/pass_through_endpoints tests/test_litellm/proxy/prompts tests/test_litellm/proxy/public_endpoints tests/test_litellm/proxy/response_api_endpoints tests/test_litellm/proxy/shutdown tests/test_litellm/proxy/spend_tracking tests/test_litellm/proxy/ui_crud_endpoints tests/test_litellm/proxy/vector_store_endpoints tests/test_litellm/proxy/test_*.py --tb=short -vv -n 4 --durations=20 test-unit-integrations: install-test-deps $(UV_RUN) pytest tests/test_litellm/integrations --tb=short -vv -n 4 --durations=20 diff --git a/deploy/charts/litellm-helm/values.yaml b/deploy/charts/litellm-helm/values.yaml index a9cdf28f0e7..6e30a6af444 100644 --- a/deploy/charts/litellm-helm/values.yaml +++ b/deploy/charts/litellm-helm/values.yaml @@ -285,11 +285,31 @@ db: deployStandalone: true # Lifecycle hooks for the LiteLLM container +# +# Prefer the native /health/drain preStop hook over a fixed `sleep`: it marks +# the pod NotReady and blocks only until in-flight requests actually finish +# (bounded by GRACEFUL_SHUTDOWN_TIMEOUT, default 30s), instead of always +# waiting the worst-case duration. The drain runs once (the preStop hook and +# the SIGTERM handler share it), so set terminationGracePeriodSeconds a few +# seconds above GRACEFUL_SHUTDOWN_TIMEOUT to leave room for teardown before +# SIGKILL. +# +# /health/drain is off by default; enable it with +# general_settings.enable_drain_endpoint: true. The kubelet calls preStop +# hooks without proxy credentials, so when the health port is reachable from +# other pods (the common case) also set +# general_settings.drain_endpoint_token (or the DRAIN_ENDPOINT_TOKEN env +# var) and send the same value on the X-Drain-Token header from the hook. +# Calls missing/wrong the token get a 401 and have no side effect. # Example: # lifecycle: # preStop: -# exec: -# command: ["/bin/sh", "-c", "sleep 10"] +# httpGet: +# path: /health/drain +# port: 4000 +# httpHeaders: +# - name: X-Drain-Token +# value: lifecycle: {} # Settings for Bitnami postgresql chart (if db.deployStandalone is true, ignored diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260526120000_add_oauth_passthrough_to_mcp_servers/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260526120000_add_oauth_passthrough_to_mcp_servers/migration.sql new file mode 100644 index 00000000000..3c387891a5e --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260526120000_add_oauth_passthrough_to_mcp_servers/migration.sql @@ -0,0 +1,2 @@ +-- AlterTable +ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "oauth_passthrough" BOOLEAN NOT NULL DEFAULT false; diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 78143fe0411..c4754ef6117 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -325,6 +325,7 @@ model LiteLLM_MCPServerTable { allow_all_keys Boolean @default(false) available_on_public_internet Boolean @default(true) delegate_auth_to_upstream Boolean @default(false) + oauth_passthrough Boolean @default(false) is_byok Boolean @default(false) byok_description String[] @default([]) byok_api_key_help_url String? diff --git a/litellm/__init__.py b/litellm/__init__.py index bae15f0362c..c954f5fd31e 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -278,6 +278,7 @@ ovhcloud_key: Optional[str] = None lemonade_key: Optional[str] = None sap_service_key: Optional[str] = None amazon_nova_api_key: Optional[str] = None +inception_key: Optional[str] = None common_cloud_provider_auth_params: dict = { "params": ["project", "region_name", "token"], "providers": ["vertex_ai", "bedrock", "watsonx", "azure", "vertex_ai_beta"], @@ -551,6 +552,7 @@ cohere_models: Set = set() cohere_chat_models: Set = set() mistral_chat_models: Set = set() text_completion_codestral_models: Set = set() +text_completion_inception_models: Set = set() anthropic_models: Set = set() openrouter_models: Set = set() datarobot_models: Set = set() @@ -628,6 +630,7 @@ publicai_models: Set = set() v0_models: Set = set() morph_models: Set = set() lambda_ai_models: Set = set() +inception_models: Set = set() hyperbolic_models: Set = set() black_forest_labs_models: Set = set() recraft_models: Set = set() @@ -792,6 +795,8 @@ def add_known_models(model_cost_map: Optional[Dict] = None): fireworks_ai_embedding_models.add(key) elif value.get("litellm_provider") == "text-completion-codestral": text_completion_codestral_models.add(key) + elif value.get("litellm_provider") == "text-completion-inception": + text_completion_inception_models.add(key) elif value.get("litellm_provider") == "xai": xai_models.add(key) elif value.get("litellm_provider") == "zai": @@ -878,6 +883,8 @@ def add_known_models(model_cost_map: Optional[Dict] = None): morph_models.add(key) elif value.get("litellm_provider") == "lambda_ai": lambda_ai_models.add(key) + elif value.get("litellm_provider") == "inception": + inception_models.add(key) elif value.get("litellm_provider") == "hyperbolic": hyperbolic_models.add(key) elif value.get("litellm_provider") == "black_forest_labs": @@ -980,6 +987,7 @@ model_list = list( | watsonx_models | gemini_models | text_completion_codestral_models + | text_completion_inception_models | xai_models | zai_models | fal_ai_models @@ -1018,6 +1026,7 @@ model_list = list( | v0_models | morph_models | lambda_ai_models + | inception_models | black_forest_labs_models | recraft_models | cometapi_models @@ -1074,6 +1083,7 @@ models_by_provider: dict = { "fireworks_ai": fireworks_ai_models | fireworks_ai_embedding_models, "aleph_alpha": aleph_alpha_models, "text-completion-codestral": text_completion_codestral_models, + "text-completion-inception": text_completion_inception_models, "xai": xai_models, "zai": zai_models, "fal_ai": fal_ai_models, @@ -1118,6 +1128,7 @@ models_by_provider: dict = { "v0": v0_models, "morph": morph_models, "lambda_ai": lambda_ai_models, + "inception": inception_models, "hyperbolic": hyperbolic_models, "black_forest_labs": black_forest_labs_models, "recraft": recraft_models, @@ -1869,6 +1880,9 @@ if TYPE_CHECKING: from .llms.codestral.completion.transformation import ( CodestralTextCompletionConfig as CodestralTextCompletionConfig, ) + from .llms.inception.completion.transformation import ( + InceptionTextCompletionConfig as InceptionTextCompletionConfig, + ) from .llms.azure.azure import ( AzureOpenAIAssistantsAPIConfig as AzureOpenAIAssistantsAPIConfig, ) @@ -1937,6 +1951,9 @@ if TYPE_CHECKING: from .llms.lambda_ai.chat.transformation import ( LambdaAIChatConfig as LambdaAIChatConfig, ) + from .llms.inception.chat.transformation import ( + InceptionChatConfig as InceptionChatConfig, + ) from .llms.hyperbolic.chat.transformation import ( HyperbolicChatConfig as HyperbolicChatConfig, ) diff --git a/litellm/_lazy_imports_registry.py b/litellm/_lazy_imports_registry.py index 17eb6609292..bdc3289b87c 100644 --- a/litellm/_lazy_imports_registry.py +++ b/litellm/_lazy_imports_registry.py @@ -267,6 +267,7 @@ LLM_CONFIG_NAMES = ( "AIMLChatConfig", "VolcEngineChatConfig", "CodestralTextCompletionConfig", + "InceptionTextCompletionConfig", "AzureOpenAIAssistantsAPIConfig", "HerokuChatConfig", "CometAPIConfig", @@ -310,6 +311,7 @@ LLM_CONFIG_NAMES = ( "MorphChatConfig", "RAGFlowConfig", "LambdaAIChatConfig", + "InceptionChatConfig", "HyperbolicChatConfig", "VercelAIGatewayConfig", "OVHCloudChatConfig", @@ -1040,6 +1042,10 @@ _LLM_CONFIGS_IMPORT_MAP = { ".llms.codestral.completion.transformation", "CodestralTextCompletionConfig", ), + "InceptionTextCompletionConfig": ( + ".llms.inception.completion.transformation", + "InceptionTextCompletionConfig", + ), "AzureOpenAIAssistantsAPIConfig": ( ".llms.azure.azure", "AzureOpenAIAssistantsAPIConfig", @@ -1154,6 +1160,10 @@ _LLM_CONFIGS_IMPORT_MAP = { "MorphChatConfig": (".llms.morph.chat.transformation", "MorphChatConfig"), "RAGFlowConfig": (".llms.ragflow.chat.transformation", "RAGFlowConfig"), "LambdaAIChatConfig": (".llms.lambda_ai.chat.transformation", "LambdaAIChatConfig"), + "InceptionChatConfig": ( + ".llms.inception.chat.transformation", + "InceptionChatConfig", + ), "HyperbolicChatConfig": ( ".llms.hyperbolic.chat.transformation", "HyperbolicChatConfig", diff --git a/litellm/_service_logger.py b/litellm/_service_logger.py index 5531c418799..b290b4340e7 100644 --- a/litellm/_service_logger.py +++ b/litellm/_service_logger.py @@ -371,6 +371,8 @@ class ServiceLogging(CustomLogger): service=ServiceTypes.LITELLM, duration=_duration, call_type=kwargs.get("call_type", "unknown"), + start_time=start_time, + end_time=end_time, ) except Exception as e: raise e diff --git a/litellm/a2a_protocol/litellm_completion_bridge/handler.py b/litellm/a2a_protocol/litellm_completion_bridge/handler.py index 67ffcf4f8f7..52e471ff702 100644 --- a/litellm/a2a_protocol/litellm_completion_bridge/handler.py +++ b/litellm/a2a_protocol/litellm_completion_bridge/handler.py @@ -20,9 +20,20 @@ from litellm.a2a_protocol.litellm_completion_bridge.transformation import ( ) from litellm.a2a_protocol.providers.config_manager import A2AProviderConfigManager +# litellm_params key carrying the authenticated principal (hashed virtual key) so +# A2A provider configs can scope provider-side state (e.g. LangFlow session memory) +# per key instead of trusting the client-supplied A2A contextId. +A2A_USER_API_KEY_HASH_PARAM = "litellm_a2a_user_api_key_hash" + # Agent metadata fields stored in litellm_params that are not valid litellm.acompletion() kwargs _AGENT_ONLY_PARAMS = frozenset( - {"is_public", "agent_name", "agent_id", "agent_card_params"} + { + "is_public", + "agent_name", + "agent_id", + "agent_card_params", + A2A_USER_API_KEY_HASH_PARAM, + } ) @@ -37,6 +48,8 @@ class A2ACompletionBridgeHandler: params: Dict[str, Any], litellm_params: Dict[str, Any], api_base: Optional[str] = None, + *, + _skip_a2a_provider_routing: bool = False, ) -> Dict[str, Any]: """ Handle non-streaming A2A request via litellm.acompletion. @@ -50,25 +63,24 @@ class A2ACompletionBridgeHandler: Returns: A2A SendMessageResponse dict """ - # Get provider config for custom_llm_provider custom_llm_provider = litellm_params.get("custom_llm_provider") - a2a_provider_config = A2AProviderConfigManager.get_provider_config( - custom_llm_provider=custom_llm_provider, - model=litellm_params.get("model"), - ) - - # If provider config exists, use it - if a2a_provider_config is not None: - verbose_logger.info(f"A2A: Using provider config for {custom_llm_provider}") - - response_data = await a2a_provider_config.handle_non_streaming( - request_id=request_id, - params=params, - api_base=api_base, - litellm_params=litellm_params, + if not _skip_a2a_provider_routing: + a2a_provider_config = A2AProviderConfigManager.get_provider_config( + custom_llm_provider=custom_llm_provider, + model=litellm_params.get("model"), ) - return response_data + if a2a_provider_config is not None: + verbose_logger.info( + f"A2A: Using provider config for {custom_llm_provider}" + ) + + return await a2a_provider_config.handle_non_streaming( + request_id=request_id, + params=params, + api_base=api_base, + litellm_params=litellm_params, + ) # Extract message from params message = params.get("message", {}) @@ -137,6 +149,8 @@ class A2ACompletionBridgeHandler: params: Dict[str, Any], litellm_params: Dict[str, Any], api_base: Optional[str] = None, + *, + _skip_a2a_provider_routing: bool = False, ) -> AsyncIterator[Dict[str, Any]]: """ Handle streaming A2A request via litellm.acompletion with stream=True. @@ -156,28 +170,27 @@ class A2ACompletionBridgeHandler: Yields: A2A streaming response events """ - # Get provider config for custom_llm_provider custom_llm_provider = litellm_params.get("custom_llm_provider") - a2a_provider_config = A2AProviderConfigManager.get_provider_config( - custom_llm_provider=custom_llm_provider, - model=litellm_params.get("model"), - ) - - # If provider config exists, use it - if a2a_provider_config is not None: - verbose_logger.info( - f"A2A: Using provider config for {custom_llm_provider} (streaming)" + if not _skip_a2a_provider_routing: + a2a_provider_config = A2AProviderConfigManager.get_provider_config( + custom_llm_provider=custom_llm_provider, + model=litellm_params.get("model"), ) - async for chunk in a2a_provider_config.handle_streaming( - request_id=request_id, - params=params, - api_base=api_base, - litellm_params=litellm_params, - ): - yield chunk + if a2a_provider_config is not None: + verbose_logger.info( + f"A2A: Using provider config for {custom_llm_provider} (streaming)" + ) - return + async for chunk in a2a_provider_config.handle_streaming( + request_id=request_id, + params=params, + api_base=api_base, + litellm_params=litellm_params, + ): + yield chunk + + return # Extract message from params message = params.get("message", {}) diff --git a/litellm/a2a_protocol/main.py b/litellm/a2a_protocol/main.py index 3ad5485dea1..6979e1ac659 100644 --- a/litellm/a2a_protocol/main.py +++ b/litellm/a2a_protocol/main.py @@ -159,7 +159,9 @@ async def _send_message_via_completion_bridge( api_base=api_base, ) - return LiteLLMSendMessageResponse.from_dict(response_dict) + return LiteLLMSendMessageResponse.from_dict( + response_dict, request_id=str(request.id) + ) async def _execute_a2a_send_with_retry( @@ -317,15 +319,6 @@ async def asend_message( ) card_url = getattr(agent_card, "url", None) if agent_card else None - context_id = trace_id or str(uuid.uuid4()) - message = request.params.message - if isinstance(message, dict): - if message.get("context_id") is None: - message["context_id"] = context_id - else: - if getattr(message, "context_id", None) is None: - message.context_id = context_id - a2a_response = await _execute_a2a_send_with_retry( a2a_client=a2a_client, request=request, @@ -338,7 +331,9 @@ async def asend_message( verbose_logger.info(f"A2A send_message completed, request_id={request.id}") # Wrap in LiteLLM response type for _hidden_params support - response = LiteLLMSendMessageResponse.from_a2a_response(a2a_response) + response = LiteLLMSendMessageResponse.from_a2a_response( + a2a_response, request_id=str(request.id) + ) # Calculate token usage from request and response response_dict = a2a_response.model_dump(mode="json", exclude_none=True) diff --git a/litellm/a2a_protocol/providers/config_manager.py b/litellm/a2a_protocol/providers/config_manager.py index ecb8f66bdeb..a421afec184 100644 --- a/litellm/a2a_protocol/providers/config_manager.py +++ b/litellm/a2a_protocol/providers/config_manager.py @@ -48,6 +48,11 @@ class A2AProviderConfigManager: return BedrockAgentCoreA2AConfig() + if custom_llm_provider == "langflow": + from litellm.a2a_protocol.providers.langflow.config import LangFlowA2AConfig + + return LangFlowA2AConfig() + if custom_llm_provider == "watsonx_orchestrate": from litellm.a2a_protocol.providers.watsonx_orchestrate.config import ( WatsonxOrchestrateA2AConfig, diff --git a/litellm/a2a_protocol/providers/langflow/__init__.py b/litellm/a2a_protocol/providers/langflow/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/a2a_protocol/providers/langflow/config.py b/litellm/a2a_protocol/providers/langflow/config.py new file mode 100644 index 00000000000..9302c38126b --- /dev/null +++ b/litellm/a2a_protocol/providers/langflow/config.py @@ -0,0 +1,62 @@ +from typing import Any, AsyncIterator, Dict, Optional + +from litellm.a2a_protocol.litellm_completion_bridge.handler import ( + A2A_USER_API_KEY_HASH_PARAM, + A2ACompletionBridgeHandler, +) +from litellm.a2a_protocol.providers.base import BaseA2AProviderConfig +from litellm.llms.langflow.a2a import merge_a2a_session_into_litellm_params + + +class LangFlowA2AConfig(BaseA2AProviderConfig): + """A2A bridge for LangFlow: scopes contextId to the authenticated key as the + LangFlow session_id, then uses completion.""" + + async def handle_non_streaming( + self, + request_id: str, + params: Dict[str, Any], + api_base: Optional[str] = None, + **kwargs, + ) -> Dict[str, Any]: + litellm_params = kwargs.get("litellm_params") + if not litellm_params: + raise ValueError( + "litellm_params is required for LangFlowA2AConfig " + "(must contain custom_llm_provider and model)" + ) + litellm_params = merge_a2a_session_into_litellm_params( + litellm_params, params, litellm_params.get(A2A_USER_API_KEY_HASH_PARAM) + ) + return await A2ACompletionBridgeHandler.handle_non_streaming( + request_id=request_id, + params=params, + litellm_params=litellm_params, + api_base=api_base, + _skip_a2a_provider_routing=True, + ) + + async def handle_streaming( + self, + request_id: str, + params: Dict[str, Any], + api_base: Optional[str] = None, + **kwargs, + ) -> AsyncIterator[Dict[str, Any]]: + litellm_params = kwargs.get("litellm_params") + if not litellm_params: + raise ValueError( + "litellm_params is required for LangFlowA2AConfig " + "(must contain custom_llm_provider and model)" + ) + litellm_params = merge_a2a_session_into_litellm_params( + litellm_params, params, litellm_params.get(A2A_USER_API_KEY_HASH_PARAM) + ) + async for chunk in A2ACompletionBridgeHandler.handle_streaming( + request_id=request_id, + params=params, + litellm_params=litellm_params, + api_base=api_base, + _skip_a2a_provider_routing=True, + ): + yield chunk diff --git a/litellm/completion_extras/litellm_responses_transformation/handler.py b/litellm/completion_extras/litellm_responses_transformation/handler.py index 2de7bda6467..87c26b776e8 100644 --- a/litellm/completion_extras/litellm_responses_transformation/handler.py +++ b/litellm/completion_extras/litellm_responses_transformation/handler.py @@ -182,6 +182,14 @@ class ResponsesToCompletionBridgeHandler: client=kwargs.get("client"), ) + # Pin the resolved provider so `responses()` doesn't re-run + # `get_llm_provider()` on the model string and strip a second + # provider prefix (see GitHub issue #28505). request_data already + # carries `custom_llm_provider` via the spread of + # `sanitized_litellm_params`; overwriting it on the dict (rather + # than adding an explicit kwarg) avoids the duplicate-keyword + # TypeError that would otherwise fire on the real bridge path. + request_data["custom_llm_provider"] = custom_llm_provider result = responses( **request_data, ) @@ -268,6 +276,13 @@ class ResponsesToCompletionBridgeHandler: except Exception as e: raise e + # Pin the resolved provider so `aresponses()` doesn't re-run + # `get_llm_provider()` on the model string and strip a second + # provider prefix (see GitHub issue #28505). Set on request_data + # rather than passed as a separate kwarg to avoid the duplicate- + # keyword TypeError when `sanitized_litellm_params` already + # carries `custom_llm_provider`. + request_data["custom_llm_provider"] = custom_llm_provider result = await aresponses( **request_data, aresponses=True, diff --git a/litellm/constants.py b/litellm/constants.py index df15050e652..26e25d0cef3 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -585,6 +585,7 @@ LITELLM_CHAT_PROVIDERS = [ "volcengine", "codestral", "text-completion-codestral", + "text-completion-inception", "deepseek", "sambanova", "maritalk", @@ -620,6 +621,7 @@ LITELLM_CHAT_PROVIDERS = [ "oci", "morph", "lambda_ai", + "inception", "vercel_ai_gateway", "wandb", "ovhcloud", @@ -779,6 +781,7 @@ openai_compatible_endpoints: List = [ "https://api.v0.dev/v1", "https://api.morphllm.com/v1", "https://api.lambda.ai/v1", + "https://api.inceptionlabs.ai/v1", "https://api.hyperbolic.xyz/v1", "https://ai-gateway.helicone.ai/", "https://ai-gateway.vercel.sh/v1", @@ -835,6 +838,7 @@ openai_compatible_providers: List = [ "helicone", "morph", "lambda_ai", + "inception", "hyperbolic", "vercel_ai_gateway", "aiml", diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 0dc56b6a3bc..7559fe142c4 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -421,8 +421,16 @@ class MCPClient: return factory - async def list_tools(self) -> List[MCPTool]: - """List available tools from the server.""" + async def list_tools(self, raise_on_error: bool = False) -> List[MCPTool]: + """List available tools from the server. + + Args: + raise_on_error: When True, re-raise exceptions instead of returning + an empty list. Used by the proxy's pass-through MCP flow so it + can surface upstream HTTP 401 responses as a proper 401 to the + MCP client (triggering the upstream OAuth flow) rather than + masking them as "connected, no tools". + """ verbose_logger.debug( f"MCP client listing tools from {self.server_url or 'stdio'}" ) @@ -458,6 +466,8 @@ class MCPClient: "the MCP server may have crashed, disconnected, or timed out" ) + if raise_on_error: + raise # Return empty list instead of raising to allow graceful degradation return [] diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 82a35f2eedd..6d0d73e033d 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -662,6 +662,16 @@ class CustomGuardrail(CustomLogger): request_data["metadata"] = {} _append_guardrail_info(request_data["metadata"]) + # Emit the otel guardrail span here, where every guardrail execution lands, + # rather than relying on a post-call hook that does not fire on every path + # (e.g. a pass-through request that passes its guardrails). + try: + from litellm.integrations.otel.logger import emit_guardrail_span + + emit_guardrail_span(slg) + except Exception: + pass + async def apply_guardrail( self, inputs: GenericGuardrailAPIInputs, diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index cb619ae0204..24780eb4bfc 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -1012,6 +1012,13 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): litellm_params = kwargs.get("litellm_params", {}) or {} _metadata = litellm_params.get("metadata", {}) or {} proxy_span = _metadata.get("litellm_parent_otel_span", None) + + # Fallback: check litellm_metadata (used by /v1/messages and other + # LITELLM_METADATA_ROUTES). + if proxy_span is None: + _litellm_metadata = litellm_params.get("litellm_metadata", {}) or {} + proxy_span = _litellm_metadata.get("litellm_parent_otel_span", None) + if ( proxy_span is not None and getattr(proxy_span, "name", None) == LITELLM_PROXY_REQUEST_SPAN_NAME @@ -2668,6 +2675,10 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): ) def _to_ns(self, dt): + if dt is None: + return int(datetime.now().timestamp() * 1e9) + if isinstance(dt, (int, float)): + return int(dt * 1e9) return int(dt.timestamp() * 1e9) def _get_span_name(self, kwargs): @@ -2714,6 +2725,13 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): _metadata = litellm_params.get("metadata", {}) or {} parent_otel_span = _metadata.get("litellm_parent_otel_span", None) + # Fallback: check litellm_metadata (used by /v1/messages and other + # LITELLM_METADATA_ROUTES that store proxy-internal metadata + # separately from the provider's native "metadata" field). + if parent_otel_span is None: + _litellm_metadata = litellm_params.get("litellm_metadata", {}) or {} + parent_otel_span = _litellm_metadata.get("litellm_parent_otel_span", None) + # Priority 1: Explicit parent span from metadata if parent_otel_span is not None: verbose_logger.debug( @@ -3287,6 +3305,32 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): value=int(status_code), ) + def record_error_attributes_on_span( + self, + span: Optional[Span], + exception: Optional[Exception], + status_code: int, + ) -> None: + """Stamp structured ``error.*`` attributes on the SERVER span from the + exception returned to the client, with ``error.code`` pinned to the real + response status. Idempotent (overwrites); emits no exception event.""" + if span is None or exception is None: + return + from litellm.litellm_core_utils.litellm_logging import ( + StandardLoggingPayloadSetup, + ) + + error_information = StandardLoggingPayloadSetup.get_error_information( + original_exception=exception + ) + error_information["error_code"] = str(status_code) + self._record_exception_on_span( + span=span, + kwargs={ + "standard_logging_object": {"error_information": error_information} + }, + ) + def set_preprocessing_duration_attribute( self, span: Optional[Span], container: Any ) -> None: diff --git a/litellm/integrations/opik/opik_payload_builder/extractors.py b/litellm/integrations/opik/opik_payload_builder/extractors.py index 9779ccddacf..1e3a664acc1 100644 --- a/litellm/integrations/opik/opik_payload_builder/extractors.py +++ b/litellm/integrations/opik/opik_payload_builder/extractors.py @@ -39,20 +39,32 @@ def extract_opik_metadata( standard_logging_metadata: Dict[str, Any], ) -> Dict[str, Any]: """ - Extract and merge Opik metadata from request and requester. + Merge Opik metadata from three sources in increasing priority order: + + 1. user_api_key_auth_metadata– lowest priority (operator-level defaults) + 2. litellm_metadata (request)– overrides auth-key defaults + 3. requester_metadata – highest priority (e.g. proxy header overrides) Args: - litellm_metadata: Metadata from litellm_params - standard_logging_metadata: Metadata from standard_logging_object + litellm_metadata: Metadata from litellm_params.mak + standard_logging_metadata: Metadata from standard_logging_object. Returns: - Merged Opik metadata dictionary + Merged Opik metadata dictionary. """ - opik_meta = litellm_metadata.get("opik", {}).copy() + # Start with auth-key defaults (lowest priority). + auth_meta = standard_logging_metadata.get("user_api_key_auth_metadata") or {} + opik_meta = (auth_meta.get("opik") or {}).copy() + # Request-level values override auth-key defaults. + request_opik = litellm_metadata.get("opik") or {} + opik_meta.update(request_opik) + + # Requester-level values win over everything else. requester_metadata = standard_logging_metadata.get("requester_metadata", {}) or {} requester_opik = requester_metadata.get("opik", {}) or {} - opik_meta.update(requester_opik) + if requester_opik: + opik_meta.update(requester_opik) _logging.verbose_logger.debug( f"litellm_opik_metadata - {json.dumps(opik_meta, default=str)}" diff --git a/litellm/integrations/otel/__init__.py b/litellm/integrations/otel/__init__.py index 42a84a85fbd..da3ce4af3e7 100644 --- a/litellm/integrations/otel/__init__.py +++ b/litellm/integrations/otel/__init__.py @@ -32,20 +32,28 @@ from litellm.integrations.otel.model.payloads import ( LLMCallSpanData, LLMRequestParams, LLMUsage, + MCPToolCallSpanData, ProxyRequestSpanData, ServerInfo, ServiceSpanData, SpanError, + is_mcp_tool_call, ) from litellm.integrations.otel.model.semconv import ( DB, + HTTP, + MCP, + Client, Error, GenAI, GenAIOperation, GenAIProvider, - HTTP, + JsonRpc, LiteLLM, + MCPMethod, Metric, + Network, + NetworkTransport, Server, resolve_operation, resolve_provider, @@ -69,13 +77,19 @@ __all__ = [ "BAGGAGE_PROMOTED_KEYS", "DB", "DEFAULT_BAGGAGE_METADATA_KEYS", + "Client", "Error", "GenAI", "GenAIOperation", "GenAIProvider", "HTTP", + "JsonRpc", "LiteLLM", + "MCP", + "MCPMethod", "Metric", + "Network", + "NetworkTransport", "Server", "resolve_operation", "resolve_provider", @@ -92,11 +106,13 @@ __all__ = [ "LLMCallSpanData", "LLMRequestParams", "LLMUsage", + "MCPToolCallSpanData", "ProxyRequestSpanData", "RequestContext", "RequestIdentity", "ServerInfo", "ServiceSpanData", "SpanError", + "is_mcp_tool_call", "promoted_baggage", ] diff --git a/litellm/integrations/otel/emitter.py b/litellm/integrations/otel/emitter.py index cae6514efdf..7fb7be7ab84 100644 --- a/litellm/integrations/otel/emitter.py +++ b/litellm/integrations/otel/emitter.py @@ -13,6 +13,7 @@ from litellm.integrations.otel.mappers.base import AttributeMapper, SpanData from litellm.integrations.otel.model.payloads import ( GuardrailSpanData, LLMCallSpanData, + MCPToolCallSpanData, ServiceSpanData, ) from litellm.integrations.otel.plumbing.providers import to_otel_span_kind @@ -22,6 +23,7 @@ from litellm.integrations.otel.model.spans import ( SpanRole, guardrail_span_name, llm_call_span_name, + mcp_tool_call_span_name, service_span_name, ) @@ -30,6 +32,7 @@ from litellm.integrations.otel.model.spans import ( # have no builder here. _NAME_BUILDERS: dict[SpanRole, Callable[..., str]] = { SpanRole.LLM_CALL: llm_call_span_name, + SpanRole.MCP_TOOL_CALL: mcp_tool_call_span_name, SpanRole.GUARDRAIL: guardrail_span_name, # DB_CALL and SERVICE are both built from ServiceSpanData; they differ only in # span kind (CLIENT vs INTERNAL) and attribute vocabulary, not in naming. @@ -121,10 +124,14 @@ class SpanEmitter: Return the span, or ``None`` if it was deduplicated away. ``tracer`` overrides the bound tracer for this span, used for per-request routing. """ - # Only LLM-call spans carry a dedup key; LLM-call and service spans - # carry an ``error`` field. ``isinstance`` narrows the type for mypy and - # keeps the engine free of duck-typed attribute reads. - dedup_key = data.identity.call_id if isinstance(data, LLMCallSpanData) else None + # LLM-call and MCP tool-call spans carry a dedup key (their request's + # call id), so a sync+async double-firing coalesces. ``isinstance`` narrows + # the type for mypy and keeps the engine free of duck-typed attribute reads. + dedup_key = ( + data.identity.call_id + if isinstance(data, (LLMCallSpanData, MCPToolCallSpanData)) + else None + ) if self._seen(dedup_key, role): return None span = self.start_span( @@ -160,7 +167,15 @@ class SpanEmitter: span.set_attribute(key, value) error = ( data.error - if isinstance(data, (LLMCallSpanData, ServiceSpanData, GuardrailSpanData)) + if isinstance( + data, + ( + LLMCallSpanData, + MCPToolCallSpanData, + ServiceSpanData, + GuardrailSpanData, + ), + ) else None ) if error and (error.error_type or error.message): diff --git a/litellm/integrations/otel/logger.py b/litellm/integrations/otel/logger.py index d7058b34d50..57738c356f7 100644 --- a/litellm/integrations/otel/logger.py +++ b/litellm/integrations/otel/logger.py @@ -25,14 +25,15 @@ from litellm.integrations.otel.mappers import resolve_mappers from litellm.integrations.otel.model.metadata import ( LLMCallEvent, RequestIdentity, - guardrail_entries_from_request_data, model_from_request_data, ) from litellm.integrations.otel.model.payloads import ( GuardrailSpanData, LLMCallSpanData, + MCPToolCallSpanData, ServiceSpanData, SpanError, + is_mcp_tool_call, ) from litellm.integrations.otel.plumbing.providers import ( build_tracer_provider, @@ -43,7 +44,10 @@ from litellm.integrations.otel.model.spans import SpanRole, span_role_for_servic from litellm.integrations.otel.model.utils import to_ns if TYPE_CHECKING: - from litellm.types.utils import StandardLoggingGuardrailInformation + from litellm.types.utils import ( + StandardLoggingGuardrailInformation, + StandardLoggingPayload, + ) LITELLM_TRACER_NAME = "litellm" @@ -200,11 +204,53 @@ class OpenTelemetryV2(CustomLogger): self._open_llm_calls.popitem(last=False) async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + if self._emit_mcp_tool_call(kwargs, start_time, end_time): + return self._close_llm_call(kwargs, start_time, end_time) async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): + if self._emit_mcp_tool_call(kwargs, start_time, end_time): + return self._close_llm_call(kwargs, start_time, end_time) + def _emit_mcp_tool_call( + self, + kwargs: Mapping[str, Any], + start_time: datetime | float | None, + end_time: datetime | float | None, + ) -> bool: + """Emit an MCP tool-call span when the closed request was a tool call. + + MCP tool calls reach the success/failure callbacks like any other request + (with ``call_type`` ``call_mcp_tool``), but they are not LLM calls and have + no ``pre_call`` carrier — so they get their own CLIENT span here, parented + to the request's server span. Returns whether it handled the event, so the + caller skips the LLM-call path. The whole span is emitted at once (there is + no boundary to open it at), deduped on the call id by the emitter. + """ + raw_payload = kwargs.get("standard_logging_object") + if not raw_payload or not is_mcp_tool_call( + cast(Mapping[str, object], raw_payload) + ): + return False + payload = cast("StandardLoggingPayload", raw_payload) + data = MCPToolCallSpanData.from_standard_logging_payload( + payload, capture_content=self.config.capture_span_content + ) + # A stray LLM carrier from a ``pre_call`` that mis-fired for this id would + # otherwise linger until evicted; drop it so it's neither leaked nor closed + # as a phantom LLM span. + if data.identity.call_id: + self._open_llm_calls.pop(data.identity.call_id, None) + self._emitter.emit( + SpanRole.MCP_TOOL_CALL, + data, + parent_context=resolve_request_span_context(), + start_time_ns=to_ns(start_time), + end_time_ns=to_ns(end_time), + ) + return True + def _close_llm_call( self, kwargs: Mapping[str, Any], @@ -421,46 +467,27 @@ class OpenTelemetryV2(CustomLogger): ) return data - async def async_post_call_success_hook( - self, - data: Mapping[str, Any], - user_api_key_dict: Any, - response: Any, - ) -> Any: - self._emit_guardrail_spans(data) - return response - - async def async_post_call_failure_hook( - self, - request_data: Mapping[str, Any], - original_exception: BaseException | None, - user_api_key_dict: Any, - traceback_str: str | None = None, - ) -> None: - self._emit_guardrail_spans(request_data) - - def _emit_guardrail_spans(self, request_data: Mapping[str, Any]) -> None: + def emit_guardrail_span(self, entry: "StandardLoggingGuardrailInformation") -> None: + # Emitted by the guardrail-recording code the moment a guardrail finishes, + # not from a post-call hook — that hook does not fire on every path (a + # pass-through request that passes its guardrails never reaches it), which + # left passing guardrails without a span. + # # A guardrail is a sibling of the LLM call under the request's root span, - # so parent it to the explicit anchor — not the active span, which on the - # failure path can be the live ``auth`` phase span (post-call failure hooks - # run from inside it on an auth rejection). Emit with the guardrail's actual - # execution window so a pre_call guardrail is placed before the LLM call - # rather than at post-call emission time. - guardrails = guardrail_entries_from_request_data(request_data) - if not guardrails: - return - parent_ctx = resolve_request_span_context() - for entry in guardrails: - data = GuardrailSpanData.from_logging_entry( - cast("StandardLoggingGuardrailInformation", entry) - ) - self._emitter.emit( - SpanRole.GUARDRAIL, - data, - parent_context=parent_ctx, - start_time_ns=to_ns(data.start_time), - end_time_ns=to_ns(data.end_time), - ) + # so parent it to the explicit anchor — never the active span, which during + # a pre_call guardrail can be the live ``auth`` phase span. Emit with the + # guardrail's actual execution window so a pre_call guardrail is placed + # before the LLM call rather than at emission time. One entry in, one span + # out — the module-level entry point routes each entry to this single + # registered logger so a guardrail is never emitted more than once. + data = GuardrailSpanData.from_logging_entry(entry) + self._emitter.emit( + SpanRole.GUARDRAIL, + data, + parent_context=resolve_request_span_context(), + start_time_ns=to_ns(data.start_time), + end_time_ns=to_ns(data.end_time), + ) def create_litellm_proxy_request_started_span( self, start_time: datetime, headers: Mapping[str, str] | None @@ -481,6 +508,26 @@ def _registered_v2_logger() -> "OpenTelemetryV2 | None": return logger if isinstance(logger, OpenTelemetryV2) else None +def emit_guardrail_span(entry: "StandardLoggingGuardrailInformation") -> None: + """Emit a guardrail span on the registered v2 OTel logger. + + Called by the guardrail-recording code the moment a guardrail finishes, so a + span is produced regardless of whether a post-call hook later runs (it does + not on the pass-through allow path). Routes through the single canonical + logger — the same one every other v2 entry point uses — so a guardrail + recorded once yields exactly one span; fanning out across every reachable + ``OpenTelemetryV2`` instance double-emits the same entry. Best-effort: span + emission must never break guardrail evaluation. + """ + logger = _registered_v2_logger() + if logger is None: + return + try: + logger.emit_guardrail_span(entry) + except Exception: + pass + + def seed_request_identity(user_api_key_dict: Any, model: Any = None) -> None: logger = _registered_v2_logger() if logger is not None: diff --git a/litellm/integrations/otel/mappers/base.py b/litellm/integrations/otel/mappers/base.py index e8fb5af9797..dfdaf77a83e 100644 --- a/litellm/integrations/otel/mappers/base.py +++ b/litellm/integrations/otel/mappers/base.py @@ -7,6 +7,7 @@ from typing_extensions import Protocol, runtime_checkable from litellm.integrations.otel.model.payloads import ( GuardrailSpanData, LLMCallSpanData, + MCPToolCallSpanData, ServiceSpanData, ) @@ -21,7 +22,7 @@ AttributeMap = dict[str, AttrValue] # The closed set of span-data types the engine routes through the mapper chain. # Server spans (PROXY_REQUEST + management routes) belong to the mounted FastAPI # instrumentor, not the mapper chain. -SpanData = LLMCallSpanData | GuardrailSpanData | ServiceSpanData +SpanData = LLMCallSpanData | MCPToolCallSpanData | GuardrailSpanData | ServiceSpanData @runtime_checkable diff --git a/litellm/integrations/otel/mappers/genai.py b/litellm/integrations/otel/mappers/genai.py index 57fa51ea1fb..6c61feced4d 100644 --- a/litellm/integrations/otel/mappers/genai.py +++ b/litellm/integrations/otel/mappers/genai.py @@ -14,10 +14,18 @@ from litellm.integrations.otel.mappers.utils import collect, drop_none from litellm.integrations.otel.model.payloads import ( GuardrailSpanData, LLMCallSpanData, + MCPToolCallSpanData, ServiceSpanData, ToolDefinition, ) -from litellm.integrations.otel.model.semconv import DB, Error, GenAI, LiteLLM, Server +from litellm.integrations.otel.model.semconv import ( + DB, + MCP, + Error, + GenAI, + LiteLLM, + Server, +) from litellm.integrations.otel.model.spans import db_system @@ -64,6 +72,18 @@ class GenAIMapper: "parameters": lambda t: t.parameters_json or None, } + _MCP_ATTRS: dict[str, Callable[[MCPToolCallSpanData], AttrValue | None]] = { + GenAI.OPERATION_NAME: lambda d: d.operation.value, + MCP.METHOD_NAME: lambda d: d.method, + MCP.SESSION_ID: lambda d: d.session_id, + GenAI.TOOL_NAME: lambda d: d.tool_name or None, + GenAI.TOOL_CALL_ARGUMENTS: lambda d: d.arguments_json, + GenAI.TOOL_CALL_RESULT: lambda d: d.result_json, + LiteLLM.MCP_SERVER_NAME: lambda d: d.server_name, + LiteLLM.CALL_ID: lambda d: d.identity.call_id or None, + f"{LiteLLM.COST_PREFIX}total": lambda d: d.response_cost, + } + _GUARDRAIL_ATTRS: dict[str, Callable[[GuardrailSpanData], AttrValue | None]] = { LiteLLM.GUARDRAIL_NAME: lambda d: d.guardrail_name, LiteLLM.GUARDRAIL_MODE: lambda d: d.mode, @@ -92,6 +112,8 @@ class GenAIMapper: match data: case LLMCallSpanData(): return self._llm_call(data) + case MCPToolCallSpanData(): + return collect(self._MCP_ATTRS, data) case GuardrailSpanData(): return self._guardrail(data) case ServiceSpanData(): diff --git a/litellm/integrations/otel/model/metadata.py b/litellm/integrations/otel/model/metadata.py index 4663ed59761..4c9cecfef57 100644 --- a/litellm/integrations/otel/model/metadata.py +++ b/litellm/integrations/otel/model/metadata.py @@ -255,26 +255,6 @@ def model_from_request_data(data: object) -> str | None: return None -def guardrail_entries_from_request_data( - request_data: Mapping[str, Any], -) -> list[dict]: - """The guardrail-information dicts buried in ``metadata`` of a post-call dict. - - ``standard_logging_guardrail_information`` is stored as either a single dict - or a list of them; normalize to a list of dicts (dropping non-dict noise) so - the caller just iterates. Empty list when none are present. - """ - metadata = request_data.get("metadata") - if not isinstance(metadata, Mapping): - return [] - info = metadata.get("standard_logging_guardrail_information") - if isinstance(info, Mapping): - return [cast(dict, info)] - if isinstance(info, list): - return [entry for entry in info if isinstance(entry, dict)] - return [] - - def resolve_provider_model(payload: "StandardLoggingPayload") -> str | None: """The model litellm dispatched to the provider, from the payload. diff --git a/litellm/integrations/otel/model/payloads.py b/litellm/integrations/otel/model/payloads.py index 65b50d0fc12..bbef40ba374 100644 --- a/litellm/integrations/otel/model/payloads.py +++ b/litellm/integrations/otel/model/payloads.py @@ -14,6 +14,7 @@ from litellm.integrations.otel.model.metadata import ( ) from litellm.integrations.otel.model.semconv import ( GenAIOperation, + MCPMethod, resolve_operation, resolve_provider, ) @@ -35,11 +36,13 @@ __all__ = [ "LLMCallSpanData", "LLMRequestParams", "LLMUsage", + "MCPToolCallSpanData", "ProxyRequestSpanData", "ServerInfo", "ServiceSpanData", "SpanError", "ToolDefinition", + "is_mcp_tool_call", ] if TYPE_CHECKING: @@ -309,6 +312,77 @@ class LLMCallSpanData: ) +# --- the MCP tool-call model ------------------------------------------------- # + + +@dataclass(frozen=True) +class MCPToolCallSpanData: + """One MCP ``tools/call`` execution, parsed from a closed request's payload. + + The proxy is an MCP *client* to the upstream server it forwards the call to, + so this is a CLIENT span. ``arguments_json``/``result_json`` are the tool's + input/output — sensitive content, so they're only retained when content + capture is enabled, mirroring ``LLMCallSpanData``'s message bodies. + """ + + operation: GenAIOperation + method: str + tool_name: str + server_name: str | None + session_id: str | None + arguments_json: str | None + result_json: str | None + error: SpanError | None + response_cost: float | None + identity: RequestIdentity + + @classmethod + def from_standard_logging_payload( + cls, payload: "StandardLoggingPayload", capture_content: bool = False + ) -> "MCPToolCallSpanData": + meta = _mcp_tool_call_metadata(cast(Mapping[str, object], payload)) + return cls( + operation=resolve_operation(as_str(payload.get("call_type"))), + method=MCPMethod.TOOLS_CALL.value, + tool_name=as_str(meta.get("name")) or "", + server_name=as_str(meta.get("mcp_server_name")), + session_id=as_str(meta.get("mcp_session_id")), + arguments_json=( + _json_or_none(meta.get("arguments")) + if capture_content and meta.get("arguments") is not None + else None + ), + result_json=( + _json_or_none(meta.get("result")) + if capture_content and meta.get("result") is not None + else None + ), + error=_parse_error(payload), + response_cost=as_float(payload.get("response_cost")), + identity=RequestContext.from_standard_logging_payload(payload).identity, + ) + + +def _mcp_tool_call_metadata(payload: Mapping[str, object]) -> Mapping[str, object]: + """The MCP gateway's tool-call metadata, which lives under + ``StandardLoggingPayload.metadata`` (a ``StandardLoggingMetadata`` key), not + at the payload's top level.""" + metadata = payload.get("metadata") + if not isinstance(metadata, Mapping): + return {} + meta = metadata.get("mcp_tool_call_metadata") + return meta if isinstance(meta, Mapping) else {} + + +def is_mcp_tool_call(payload: Mapping[str, object]) -> bool: + """Whether a closed request's payload is an MCP tool call rather than an LLM + call — true when the MCP gateway stamped its tool-call metadata, or the call + type says so on a path that hasn't populated the metadata yet.""" + return bool(_mcp_tool_call_metadata(payload)) or ( + payload.get("call_type") == "call_mcp_tool" + ) + + # --- service event_metadata sanitization ------------------------------------ # # Substrings (case-insensitive) of keys that must never reach a span: secrets, diff --git a/litellm/integrations/otel/model/semconv.py b/litellm/integrations/otel/model/semconv.py index 1c6c30eda0d..7df07f30a01 100644 --- a/litellm/integrations/otel/model/semconv.py +++ b/litellm/integrations/otel/model/semconv.py @@ -16,7 +16,7 @@ class GenAIOperation(str, Enum): GENERATE_CONTENT = "generate_content" CREATE_AGENT = "create_agent" # reserved for future agent spans INVOKE_AGENT = "invoke_agent" # reserved for future agent spans - EXECUTE_TOOL = "execute_tool" # reserved for future tool spans + EXECUTE_TOOL = "execute_tool" # MCP tool-call spans class GenAIProvider(str, Enum): @@ -38,6 +38,16 @@ class GenAIProvider(str, Enum): IBM_WATSONX_AI = "ibm.watsonx.ai" +class MCPMethod(str, Enum): + """Well-known values for ``mcp.method.name`` that litellm's MCP gateway + serves. The value is the JSON-RPC method exactly as it travels on the wire.""" + + TOOLS_CALL = "tools/call" + TOOLS_LIST = "tools/list" + PROMPTS_GET = "prompts/get" + PROMPTS_LIST = "prompts/list" + + class GenAI: """Canonical OTel GenAI span-attribute keys.""" @@ -68,11 +78,68 @@ class GenAI: SYSTEM_INSTRUCTIONS: Final = "gen_ai.system_instructions" OUTPUT_TYPE: Final = "gen_ai.output.type" CONVERSATION_ID: Final = "gen_ai.conversation.id" - # agent / tool (reserved) + # agent (reserved) AGENT_ID: Final = "gen_ai.agent.id" AGENT_NAME: Final = "gen_ai.agent.name" + # tool / tool-call (stamped on MCP tool-call spans). Arguments and result are + # the tool's input/output payloads — sensitive, so they're opt-in and gated by + # the same content-capture mode as prompt/response content. TOOL_NAME: Final = "gen_ai.tool.name" TOOL_CALL_ID: Final = "gen_ai.tool.call.id" + TOOL_CALL_ARGUMENTS: Final = "gen_ai.tool.call.arguments" + TOOL_CALL_RESULT: Final = "gen_ai.tool.call.result" + # prompt (MCP ``prompts/get`` etc.) + PROMPT_NAME: Final = "gen_ai.prompt.name" + + +class MCP: + """OTel GenAI MCP (Model Context Protocol) span-attribute keys. + + ``METHOD_NAME`` is the only key litellm populates from a closed request today; + the rest are part of the convention's vocabulary and are stamped when the + corresponding signal (session, protocol version, resource) is available. + """ + + METHOD_NAME: Final = "mcp.method.name" + SESSION_ID: Final = "mcp.session.id" + PROTOCOL_VERSION: Final = "mcp.protocol.version" + RESOURCE_URI: Final = "mcp.resource.uri" + + +class JsonRpc: + """JSON-RPC keys carried on MCP spans. The error/status code lives in the + ``rpc.*`` namespace per semconv, not ``jsonrpc.*``.""" + + REQUEST_ID: Final = "jsonrpc.request.id" + PROTOCOL_VERSION: Final = "jsonrpc.protocol.version" + RESPONSE_STATUS_CODE: Final = "rpc.response.status_code" + + +class NetworkTransport(str, Enum): + """Well-known values for ``network.transport``.""" + + TCP = "tcp" + UDP = "udp" + QUIC = "quic" + UNIX = "unix" + PIPE = "pipe" + + +class Network: + """OTel network keys, recommended on MCP spans to describe the transport + carrying the JSON-RPC messages (stdio pipe, HTTP, websocket, …).""" + + PROTOCOL_NAME: Final = "network.protocol.name" + PROTOCOL_VERSION: Final = "network.protocol.version" + TRANSPORT: Final = "network.transport" + + +class Client: + """Peer (client) network keys, stamped on MCP *server* spans the same way + ``server.*`` is stamped on client spans.""" + + ADDRESS: Final = "client.address" + PORT: Final = "client.port" class Error: @@ -137,6 +204,10 @@ class LiteLLM: SERVICE_NAME: Final = "litellm.service.name" SERVICE_CALL_TYPE: Final = "litellm.service.call_type" PREPROCESSING_MS: Final = "litellm.preprocessing.duration_ms" + # The logical name of the MCP server a tool call was routed to. There is no + # semconv key for an MCP server's *name* (the convention uses ``server.address`` + # for its network location), so it lives under the vendor namespace. + MCP_SERVER_NAME: Final = "litellm.mcp.server.name" class Metric: @@ -179,6 +250,7 @@ _OPERATION_BY_CALL_TYPE: dict[str, GenAIOperation] = { "aembedding": GenAIOperation.EMBEDDINGS, "responses": GenAIOperation.CHAT, "aresponses": GenAIOperation.CHAT, + "call_mcp_tool": GenAIOperation.EXECUTE_TOOL, } diff --git a/litellm/integrations/otel/model/spans.py b/litellm/integrations/otel/model/spans.py index e4876f4ee58..1adc1d68dde 100644 --- a/litellm/integrations/otel/model/spans.py +++ b/litellm/integrations/otel/model/spans.py @@ -46,6 +46,7 @@ if TYPE_CHECKING: from litellm.integrations.otel.model.payloads import ( GuardrailSpanData, LLMCallSpanData, + MCPToolCallSpanData, ProxyRequestSpanData, ServiceSpanData, ) @@ -54,6 +55,7 @@ if TYPE_CHECKING: class SpanRole(str, Enum): PROXY_REQUEST = "proxy_request" LLM_CALL = "llm_call" + MCP_TOOL_CALL = "mcp_tool_call" GUARDRAIL = "guardrail" DB_CALL = "db_call" SERVICE = "service" @@ -81,6 +83,11 @@ SPAN_REGISTRY: dict[SpanRole, SpanSpec] = { SpanRole.LLM_CALL: SpanSpec( SpanRole.LLM_CALL, LiteLLMSpanKind.CLIENT, parent=SpanRole.PROXY_REQUEST ), + # The proxy is an MCP client to the upstream server it dispatches the tool + # call to, so this is a CLIENT span, sibling of the LLM call under the request. + SpanRole.MCP_TOOL_CALL: SpanSpec( + SpanRole.MCP_TOOL_CALL, LiteLLMSpanKind.CLIENT, parent=SpanRole.PROXY_REQUEST + ), SpanRole.GUARDRAIL: SpanSpec( SpanRole.GUARDRAIL, LiteLLMSpanKind.INTERNAL, parent=SpanRole.PROXY_REQUEST ), @@ -165,6 +172,11 @@ def llm_call_span_name(data: "LLMCallSpanData") -> str: return f"{data.operation.value} {model}".strip() +def mcp_tool_call_span_name(data: "MCPToolCallSpanData") -> str: + """``"{mcp.method.name} {tool}"`` e.g. ``"tools/call get-weather"`` (MCP semconv).""" + return f"{data.method} {data.tool_name}".strip() + + def proxy_request_span_name(data: "ProxyRequestSpanData") -> str: """``"{method} {route}"`` (HTTP semconv).""" return f"{data.http_method} {data.route}".strip() diff --git a/litellm/litellm_core_utils/get_llm_provider_logic.py b/litellm/litellm_core_utils/get_llm_provider_logic.py index ba6d438f16c..a71000f00f8 100644 --- a/litellm/litellm_core_utils/get_llm_provider_logic.py +++ b/litellm/litellm_core_utils/get_llm_provider_logic.py @@ -1,3 +1,4 @@ +import re from typing import Optional, Tuple from urllib.parse import urlparse @@ -71,6 +72,25 @@ def _is_azure_claude_model(model: str) -> bool: return False +_CLAUDE_PATTERN = re.compile(r"^claude-[a-z]+-\d+-\d+(?:-\d{8})?$", re.IGNORECASE) + + +def _matches_claude_model_pattern(model: str) -> bool: + """ + Check if a model string matches the Claude model naming pattern. + + Matches patterns like: + - claude-opus-4-7 + - claude-sonnet-4-6 + - claude-haiku-4-5 + - claude-opus-5-1-20270101 (with optional date suffix) + + This allows future Claude models to be routed to the Anthropic provider + without requiring updates to model_prices_and_context_window.json. + """ + return _CLAUDE_PATTERN.match(model) is not None + + def handle_cohere_chat_model_custom_llm_provider( model: str, custom_llm_provider: Optional[str] = None ) -> Tuple[str, Optional[str]]: @@ -353,6 +373,9 @@ def get_llm_provider( # noqa: PLR0915 elif endpoint == "https://api.lambda.ai/v1": custom_llm_provider = "lambda_ai" dynamic_api_key = get_secret_str("LAMBDA_API_KEY") + elif endpoint == "https://api.inceptionlabs.ai/v1": + custom_llm_provider = "inception" + dynamic_api_key = get_secret_str("INCEPTION_API_KEY") elif endpoint == "https://api.hyperbolic.xyz/v1": custom_llm_provider = "hyperbolic" dynamic_api_key = get_secret_str("HYPERBOLIC_API_KEY") @@ -398,6 +421,9 @@ def get_llm_provider( # noqa: PLR0915 custom_llm_provider = "anthropic_text" else: custom_llm_provider = "anthropic" + ## anthropic - pattern-based matching for future Claude models + elif _matches_claude_model_pattern(model): + custom_llm_provider = "anthropic" ## cohere elif model in litellm.cohere_models or model in litellm.cohere_embedding_models: custom_llm_provider = "cohere" @@ -931,6 +957,13 @@ def _get_openai_compatible_provider_info( # noqa: PLR0915 ) = litellm.LambdaAIChatConfig()._get_openai_compatible_provider_info( api_base, api_key ) + elif custom_llm_provider == "inception": + ( + api_base, + dynamic_api_key, + ) = litellm.InceptionChatConfig()._get_openai_compatible_provider_info( + api_base, api_key + ) elif custom_llm_provider == "hyperbolic": ( api_base, diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index c127b3873a7..abc22713be4 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -5300,8 +5300,12 @@ class StandardLoggingPayloadSetup: tb_lines[:MAXIMUM_TRACEBACK_LINES_TO_LOG] ) # Limit to first 100 lines - # Get additional error details - error_message = str(original_exception) + explicit_message = getattr(original_exception, "message", None) + error_message = ( + explicit_message + if isinstance(explicit_message, str) and explicit_message + else str(original_exception) + ) return StandardLoggingPayloadErrorInformation( error_code=error_status, diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py index 882561ed2e8..f39c942f90f 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/utils.py +++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py @@ -34,6 +34,14 @@ _IMAGE_RESPONSE_CALL_TYPES = frozenset( _VALID_DATA_RESIDENCIES = frozenset(r.value for r in DataResidency) +def _get_token_detail_value(details: object, key: str) -> Optional[int]: + if isinstance(details, dict): + value = details.get(key) + else: + value = getattr(details, key, None) + return value if isinstance(value, int) else None + + def _is_above_128k(tokens: float) -> bool: if tokens > 128000: return True @@ -870,17 +878,47 @@ def calculate_image_response_cost_from_usage( cached_tokens=0, ) + output_tokens_details = getattr(usage, "completion_tokens_details", None) + if output_tokens_details is None: + output_tokens_details = getattr(usage, "output_tokens_details", None) + + if output_tokens_details is None: + completion_tokens_details = CompletionTokensDetailsWrapper( + text_tokens=0, + image_tokens=completion_tokens, + reasoning_tokens=0, + audio_tokens=0, + ) + else: + text_tokens = _get_token_detail_value(output_tokens_details, "text_tokens") or 0 + image_tokens = ( + _get_token_detail_value(output_tokens_details, "image_tokens") or 0 + ) + audio_tokens = ( + _get_token_detail_value(output_tokens_details, "audio_tokens") or 0 + ) + reasoning_tokens = ( + _get_token_detail_value(output_tokens_details, "reasoning_tokens") or 0 + ) + known_output_tokens = ( + text_tokens + image_tokens + audio_tokens + reasoning_tokens + ) + if completion_tokens > known_output_tokens: + text_tokens += completion_tokens - known_output_tokens + + completion_tokens_details = CompletionTokensDetailsWrapper( + text_tokens=text_tokens, + image_tokens=image_tokens, + reasoning_tokens=reasoning_tokens, + audio_tokens=audio_tokens, + ) + normalized_usage = Usage( prompt_tokens=prompt_tokens, completion_tokens=completion_tokens, total_tokens=total_tokens, prompt_tokens_details=prompt_tokens_details, - completion_tokens_details=CompletionTokensDetailsWrapper( - text_tokens=0, - image_tokens=completion_tokens, - reasoning_tokens=0, - audio_tokens=0, - ), + completion_tokens_details=completion_tokens_details, ) prompt_cost, completion_cost = generic_cost_per_token( diff --git a/litellm/litellm_core_utils/logging_worker.py b/litellm/litellm_core_utils/logging_worker.py index 3db3700ee07..294ba8e5dea 100644 --- a/litellm/litellm_core_utils/logging_worker.py +++ b/litellm/litellm_core_utils/logging_worker.py @@ -3,6 +3,7 @@ import asyncio import contextvars +import logging from typing import Coroutine, Optional import atexit from typing_extensions import TypedDict @@ -494,31 +495,43 @@ class LoggingWorker: processed = 0 start_time = loop.time() - while not self._queue.empty() and processed < MAX_ITERATIONS_TO_CLEAR_QUEUE: - if loop.time() - start_time >= MAX_TIME_TO_CLEAR_QUEUE: - self._safe_log( - "warning", - f"[LoggingWorker] atexit: Reached time limit ({MAX_TIME_TO_CLEAR_QUEUE}s), stopping flush", - ) - break + # logging.raiseExceptions is a process-wide global; scope the + # suppression to just the drain loop, where shutdown callbacks may + # log to already-closed handler streams, so other threads keep their + # logging error reporting for as little of the window as possible. + previous_raise_exceptions = logging.raiseExceptions + logging.raiseExceptions = False + try: + while ( + not self._queue.empty() + and processed < MAX_ITERATIONS_TO_CLEAR_QUEUE + ): + if loop.time() - start_time >= MAX_TIME_TO_CLEAR_QUEUE: + self._safe_log( + "warning", + f"[LoggingWorker] atexit: Reached time limit ({MAX_TIME_TO_CLEAR_QUEUE}s), stopping flush", + ) + break - try: - task = self._queue.get_nowait() - except asyncio.QueueEmpty: - break + try: + task = self._queue.get_nowait() + except asyncio.QueueEmpty: + break - # Run the coroutine synchronously in new loop - # Note: We run the coroutine directly, not via create_task, - # since we're in a new event loop context - try: - loop.run_until_complete(task["coroutine"]) - processed += 1 - except Exception: - # Silent failure to not break user's program - pass - finally: - # Clear reference to prevent memory leaks - task = None + # Run the coroutine synchronously in new loop + # Note: We run the coroutine directly, not via create_task, + # since we're in a new event loop context + try: + loop.run_until_complete(task["coroutine"]) + processed += 1 + except Exception: + # Silent failure to not break user's program + pass + finally: + # Clear reference to prevent memory leaks + task = None + finally: + logging.raiseExceptions = previous_raise_exceptions self._safe_log( "info", diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 46e9b43a429..1460dbaf0a9 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -1670,15 +1670,15 @@ def convert_to_gemini_tool_call_result( # noqa: PLR0915 if gemini_call_id: _function_response["id"] = gemini_call_id - # Create part with function_response, and optionally inline_data for images (Computer Use) _part: VertexPartType = {"function_response": _function_response} - # For Computer Use, if we have images/files, we need separate parts: - # - One part with function_response - # - One part per inline_data item - # Gemini's PartType is a oneof, so we can't have both in the same part + # For multimodal function responses, Gemini expects media parts nested + # inside functionResponse.parts instead of sibling content parts. if inline_data_list: - return [_part] + [{"inline_data": d} for d in inline_data_list] + _function_response["parts"] = [ + {"inline_data": inline_data} for inline_data in inline_data_list + ] + return [_part] return _part diff --git a/litellm/litellm_core_utils/safe_json_dumps.py b/litellm/litellm_core_utils/safe_json_dumps.py index 051aa2f27a5..154306d01b8 100644 --- a/litellm/litellm_core_utils/safe_json_dumps.py +++ b/litellm/litellm_core_utils/safe_json_dumps.py @@ -6,10 +6,16 @@ from pydantic import BaseModel from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH +def strip_null_bytes(value: str) -> str: + """Strip NUL bytes, which PostgreSQL text/jsonb columns reject (error 22P05).""" + return value.replace("\x00", "") + + def safe_dumps(data: Any, max_depth: int = DEFAULT_MAX_RECURSE_DEPTH) -> str: """ Recursively serialize data while detecting circular references. If a circular reference is detected then a marker string is returned. + NUL bytes are stripped from strings to prevent PostgreSQL 22P05 errors. """ def _serialize(obj: Any, seen: set, depth: int) -> Any: @@ -17,7 +23,9 @@ def safe_dumps(data: Any, max_depth: int = DEFAULT_MAX_RECURSE_DEPTH) -> str: if depth > max_depth: return "MaxDepthExceeded" # Base-case: if it is a primitive, simply return it. - if isinstance(obj, (str, int, float, bool, type(None))): + if isinstance(obj, str): + return strip_null_bytes(obj) + if isinstance(obj, (int, float, bool, type(None))): return obj # Check for circular reference. if id(obj) in seen: @@ -28,7 +36,7 @@ def safe_dumps(data: Any, max_depth: int = DEFAULT_MAX_RECURSE_DEPTH) -> str: result = {} for k, v in obj.items(): if isinstance(k, (str)): - result[k] = _serialize(v, seen, depth + 1) + result[strip_null_bytes(k)] = _serialize(v, seen, depth + 1) seen.remove(id(obj)) return result elif isinstance(obj, list): @@ -51,7 +59,7 @@ def safe_dumps(data: Any, max_depth: int = DEFAULT_MAX_RECURSE_DEPTH) -> str: else: # Fall back to string conversion for non-serializable objects. try: - return str(obj) + return strip_null_bytes(str(obj)) except Exception: return "Unserializable Object" diff --git a/litellm/llms/custom_httpx/aiohttp_transport.py b/litellm/llms/custom_httpx/aiohttp_transport.py index 132191c946c..62f707b3622 100644 --- a/litellm/llms/custom_httpx/aiohttp_transport.py +++ b/litellm/llms/custom_httpx/aiohttp_transport.py @@ -256,7 +256,10 @@ class LiteLLMAiohttpTransport(AiohttpTransport): from yarl import URL as YarlURL try: - data = request.content + # Coerce an empty body to None so aiohttp does not attach a + # `Content-Type: application/octet-stream` header for bodyless + # requests (e.g. DELETE /responses/{id}), which upstream APIs reject. + data = request.content or None except httpx.RequestNotRead: data = request.stream # type: ignore request.headers.pop("transfer-encoding", None) # handled by aiohttp diff --git a/litellm/llms/gemini/common_utils.py b/litellm/llms/gemini/common_utils.py index bc963d62b5f..42a807983b9 100644 --- a/litellm/llms/gemini/common_utils.py +++ b/litellm/llms/gemini/common_utils.py @@ -1,6 +1,8 @@ import base64 import datetime -from typing import Any, Dict, List, Optional, Union +import json +import math +from typing import Any, Dict, List, Optional, Sequence, Union import httpx @@ -12,6 +14,245 @@ from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import AllMessageValues from litellm.types.utils import TokenCountResponse +GEMINI_IMAGE_ASPECT_RATIOS: Dict[str, float] = { + "1:1": 1 / 1, + "1:4": 1 / 4, + "1:8": 1 / 8, + "2:3": 2 / 3, + "3:2": 3 / 2, + "3:4": 3 / 4, + "4:1": 4 / 1, + "4:3": 4 / 3, + "4:5": 4 / 5, + "5:4": 5 / 4, + "8:1": 8 / 1, + "9:16": 9 / 16, + "16:9": 16 / 9, + "21:9": 21 / 9, +} + +# Supported aspect ratio dimensions from Google Gemini image generation docs: +# https://ai.google.dev/gemini-api/docs/image-generation#aspect_ratios_and_image_size +GEMINI_IMAGE_SIZE_TO_ASPECT_RATIO: Dict[tuple[int, int], str] = { + (512, 512): "1:1", + (1024, 1024): "1:1", + (2048, 2048): "1:1", + (4096, 4096): "1:1", + (256, 1024): "1:4", + (512, 2048): "1:4", + (1024, 4096): "1:4", + (2048, 8192): "1:4", + (192, 1536): "1:8", + (384, 3072): "1:8", + (768, 6144): "1:8", + (1536, 12288): "1:8", + (424, 632): "2:3", + (848, 1264): "2:3", + (1696, 2528): "2:3", + (3392, 5056): "2:3", + (632, 424): "3:2", + (1264, 848): "3:2", + (2528, 1696): "3:2", + (5056, 3392): "3:2", + (448, 600): "3:4", + (896, 1200): "3:4", + (1792, 2400): "3:4", + (3584, 4800): "3:4", + (1024, 256): "4:1", + (2048, 512): "4:1", + (4096, 1024): "4:1", + (8192, 2048): "4:1", + (600, 448): "4:3", + (1200, 896): "4:3", + (2400, 1792): "4:3", + (4800, 3584): "4:3", + (464, 576): "4:5", + (928, 1152): "4:5", + (1856, 2304): "4:5", + (3712, 4608): "4:5", + (576, 464): "5:4", + (1152, 928): "5:4", + (2304, 1856): "5:4", + (4608, 3712): "5:4", + (1536, 192): "8:1", + (3072, 384): "8:1", + (6144, 768): "8:1", + (12288, 1536): "8:1", + (384, 688): "9:16", + (768, 1376): "9:16", + (1536, 2752): "9:16", + (3072, 5504): "9:16", + (688, 384): "16:9", + (1376, 768): "16:9", + (2752, 1536): "16:9", + (5504, 3072): "16:9", + (792, 336): "21:9", + (1584, 672): "21:9", + (3168, 1344): "21:9", + (6336, 2688): "21:9", + (1280, 896): "4:3", + (896, 1280): "3:4", +} + + +def map_openai_size_to_gemini_image_config( + size: str, model: str +) -> Optional[Dict[str, str]]: + dimensions = _parse_openai_image_size(size) + if dimensions is None: + return None + + width, height = dimensions + image_config = { + "aspectRatio": _map_dimensions_to_gemini_aspect_ratio(width, height) + } + image_size = _map_dimensions_to_gemini_image_size(width, height) + if is_gemini_image_model(model): + if supports_gemini_image_size(model): + image_config["imageSize"] = image_size + else: + image_config["imageSize"] = image_size + return image_config + + +def supports_gemini_image_size(model: str) -> bool: + try: + model_info = litellm.get_model_info(model=model) + value = model_info.get("supports_image_size") + if value is not None: + return bool(value) + except Exception: + pass + return "2.5-flash" not in model + + +def is_gemini_image_model(model: str) -> bool: + base_model = model.split("/", 1)[-1] + return "gemini" in base_model + + +def map_openai_image_params_to_gemini( + params: Dict[str, Any], + model: str, + supported_params: Sequence[str], + optional_params: Optional[Dict[str, Any]] = None, + parse_image_config_string: bool = False, +) -> Dict[str, Any]: + optional_params = optional_params or {} + filtered_params = { + key: value for key, value in params.items() if key in supported_params + } + + mapped_params: Dict[str, Any] = {} + + if "n" in filtered_params and "n" not in optional_params: + mapped_params["sampleCount"] = filtered_params["n"] + + if "size" in filtered_params and "size" not in optional_params: + image_config = map_openai_size_to_gemini_image_config( + filtered_params["size"], + model, + ) + if image_config is not None: + if is_gemini_image_model(model): + mapped_params["imageConfig"] = image_config + else: + mapped_params["aspectRatio"] = image_config["aspectRatio"] + if "imageSize" in image_config: + mapped_params["imageSize"] = image_config["imageSize"] + + image_config_param = filtered_params.get("imageConfig") + if isinstance(image_config_param, str) and parse_image_config_string: + try: + image_config_param = json.loads(image_config_param) + except json.JSONDecodeError as exc: + raise litellm.UnsupportedParamsError( + model=model, + message="`imageConfig` must be valid JSON when provided as a string.", + ) from exc + if isinstance(image_config_param, dict): + mapped_params["imageConfig"] = image_config_param + + for key, value in filtered_params.items(): + if key not in ("n", "size", "imageConfig") and key not in optional_params: + mapped_params[key] = value + + return mapped_params + + +def get_gemini_image_generation_config( + model: str, + optional_params: Dict[str, Any], +) -> Dict[str, Any]: + generation_config: Dict[str, Any] = {"response_modalities": ["IMAGE", "TEXT"]} + + image_config: Dict[str, Any] = {} + if isinstance(optional_params.get("imageConfig"), dict): + image_config.update(optional_params["imageConfig"]) + + if not supports_gemini_image_size(model): + image_config.pop("imageSize", None) + + if image_config: + generation_config["imageConfig"] = image_config + + candidate_count = next( + ( + optional_params[key] + for key in ("candidateCount", "candidate_count", "sampleCount", "n") + if optional_params.get(key) is not None + ), + None, + ) + if candidate_count is not None: + generation_config["candidateCount"] = candidate_count + + return generation_config + + +def _parse_openai_image_size(size: str) -> Optional[tuple[int, int]]: + if size == "auto": + return None + + width_str, separator, height_str = size.lower().partition("x") + if not separator: + return None + + try: + width = int(width_str) + height = int(height_str) + except ValueError: + return None + + if width <= 0 or height <= 0: + return None + + return width, height + + +def _map_dimensions_to_gemini_aspect_ratio(width: int, height: int) -> str: + if (width, height) in GEMINI_IMAGE_SIZE_TO_ASPECT_RATIO: + return GEMINI_IMAGE_SIZE_TO_ASPECT_RATIO[(width, height)] + + requested_ratio = width / height + return min( + GEMINI_IMAGE_ASPECT_RATIOS, + key=lambda aspect_ratio: abs( + math.log(GEMINI_IMAGE_ASPECT_RATIOS[aspect_ratio] / requested_ratio) + ), + ) + + +def _map_dimensions_to_gemini_image_size(width: int, height: int) -> str: + effective_square_side = math.sqrt(width * height) + if effective_square_side < 768: + return "512" + if effective_square_side < 1536: + return "1K" + if effective_square_side < 3072: + return "2K" + return "4K" + class GeminiError(BaseLLMException): pass diff --git a/litellm/llms/gemini/image_edit/cost_calculator.py b/litellm/llms/gemini/image_edit/cost_calculator.py index 2e332a7fc00..956edb849a0 100644 --- a/litellm/llms/gemini/image_edit/cost_calculator.py +++ b/litellm/llms/gemini/image_edit/cost_calculator.py @@ -4,8 +4,9 @@ Gemini Image Edit Cost Calculator from typing import Any -import litellm -from litellm.types.utils import ImageResponse +from litellm.llms.gemini.image_generation.cost_calculator import ( + cost_calculator as image_generation_cost_calculator, +) def cost_calculator( @@ -15,20 +16,10 @@ def cost_calculator( """ Gemini image edit cost calculator. - Mirrors image generation pricing: charge per returned image based on - model metadata (`output_cost_per_image`). + Gemini image edits and generations share image response billing behavior: + use provider token usage when present, otherwise fall back to per-image pricing. """ - model_info = litellm.get_model_info( + return image_generation_cost_calculator( model=model, - custom_llm_provider="gemini", + image_response=image_response, ) - - output_cost_per_image: float = model_info.get("output_cost_per_image") or 0.0 - - if not isinstance(image_response, ImageResponse): - raise ValueError( - f"image_response must be of type ImageResponse got type={type(image_response)}" - ) - - num_images = len(image_response.data or []) - return output_cost_per_image * num_images diff --git a/litellm/llms/gemini/image_edit/transformation.py b/litellm/llms/gemini/image_edit/transformation.py index c8aaab0e14e..2316361d6e7 100644 --- a/litellm/llms/gemini/image_edit/transformation.py +++ b/litellm/llms/gemini/image_edit/transformation.py @@ -7,10 +7,22 @@ from httpx._types import RequestFiles from litellm.images.utils import ImageEditRequestUtils from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig +from litellm.llms.gemini.common_utils import ( + get_gemini_image_generation_config, + map_openai_image_params_to_gemini, +) +from litellm.llms.gemini.image_usage_transformation import ( + transform_gemini_image_usage, +) from litellm.secret_managers.main import get_secret_str from litellm.types.images.main import ImageEditOptionalRequestParams from litellm.types.router import GenericLiteLLMParams -from litellm.types.utils import FileTypes, ImageObject, ImageResponse, OpenAIImage +from litellm.types.utils import ( + FileTypes, + ImageObject, + ImageResponse, + OpenAIImage, +) if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj @@ -22,7 +34,7 @@ else: class GeminiImageEditConfig(BaseImageEditConfig): DEFAULT_BASE_URL: str = "https://generativelanguage.googleapis.com/v1beta" - SUPPORTED_PARAMS: List[str] = ["size"] + SUPPORTED_PARAMS: List[str] = ["n", "size", "imageConfig"] def get_supported_openai_params(self, model: str) -> List[str]: return list(self.SUPPORTED_PARAMS) @@ -33,21 +45,12 @@ class GeminiImageEditConfig(BaseImageEditConfig): model: str, drop_params: bool, ) -> Dict[str, Any]: - supported_params = self.get_supported_openai_params(model) - filtered_params = { - key: value - for key, value in image_edit_optional_params.items() - if key in supported_params - } - - mapped_params: Dict[str, Any] = {} - - if "size" in filtered_params: - mapped_params["aspectRatio"] = self._map_size_to_aspect_ratio( - filtered_params["size"] # type: ignore[arg-type] - ) - - return mapped_params + return map_openai_image_params_to_gemini( + params=image_edit_optional_params, # type: ignore[arg-type] + model=model, + supported_params=self.get_supported_openai_params(model), + parse_image_config_string=True, + ) def validate_environment( self, @@ -107,18 +110,10 @@ class GeminiImageEditConfig(BaseImageEditConfig): request_body: Dict[str, Any] = {"contents": contents} - generation_config: Dict[str, Any] = {} - - if "aspectRatio" in image_edit_optional_request_params: - # Move aspectRatio into imageConfig inside generationConfig - if "imageConfig" not in generation_config: - generation_config["imageConfig"] = {} - generation_config["imageConfig"]["aspectRatio"] = ( - image_edit_optional_request_params["aspectRatio"] - ) - - if generation_config: - request_body["generationConfig"] = generation_config + request_body["generationConfig"] = get_gemini_image_generation_config( + model=model, + optional_params=image_edit_optional_request_params, + ) empty_files = cast(RequestFiles, []) return request_body, empty_files @@ -156,18 +151,12 @@ class GeminiImageEditConfig(BaseImageEditConfig): ) model_response.data = cast(List[OpenAIImage], data_list) + if "usageMetadata" in response_json: + model_response.usage = transform_gemini_image_usage( + response_json["usageMetadata"] + ) return model_response - def _map_size_to_aspect_ratio(self, size: str) -> str: - aspect_ratio_map = { - "1024x1024": "1:1", - "1792x1024": "16:9", - "1024x1792": "9:16", - "1280x896": "4:3", - "896x1280": "3:4", - } - return aspect_ratio_map.get(size, "1:1") - def _prepare_inline_image_parts( self, image: Union[FileTypes, List[FileTypes]] ) -> List[Dict[str, Any]]: diff --git a/litellm/llms/gemini/image_generation/transformation.py b/litellm/llms/gemini/image_generation/transformation.py index 9c4cd008b8c..e6770a76bcb 100644 --- a/litellm/llms/gemini/image_generation/transformation.py +++ b/litellm/llms/gemini/image_generation/transformation.py @@ -5,18 +5,21 @@ import httpx from litellm.llms.base_llm.image_generation.transformation import ( BaseImageGenerationConfig, ) +from litellm.llms.gemini.common_utils import ( + get_gemini_image_generation_config, + is_gemini_image_model, + map_openai_image_params_to_gemini, +) +from litellm.llms.gemini.image_usage_transformation import ( + transform_gemini_image_usage, +) from litellm.secret_managers.main import get_secret_str from litellm.types.llms.gemini import GeminiImageGenerationRequest from litellm.types.llms.openai import ( AllMessageValues, OpenAIImageGenerationOptionalParams, ) -from litellm.types.utils import ( - ImageObject, - ImageResponse, - ImageUsage, - ImageUsageInputTokensDetails, -) +from litellm.types.utils import ImageObject, ImageResponse if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj @@ -36,7 +39,10 @@ class GoogleImageGenConfig(BaseImageGenerationConfig): Google AI Imagen API supported parameters https://ai.google.dev/gemini-api/docs/imagen """ - return ["n", "size"] + supported_params = ["n", "size"] + if is_gemini_image_model(model): + supported_params.append("imageConfig") + return supported_params # type: ignore[return-value] def map_openai_params( self, @@ -45,64 +51,11 @@ class GoogleImageGenConfig(BaseImageGenerationConfig): model: str, drop_params: bool, ) -> dict: - supported_params = self.get_supported_openai_params(model) - mapped_params = {} - - for k, v in non_default_params.items(): - if k not in optional_params.keys(): - if k in supported_params: - # Map OpenAI parameters to Google format - if k == "n": - mapped_params["sampleCount"] = v - elif k == "size": - # Map OpenAI size format to Google aspectRatio - mapped_params["aspectRatio"] = self._map_size_to_aspect_ratio(v) - else: - mapped_params[k] = v - return mapped_params - - def _map_size_to_aspect_ratio(self, size: str) -> str: - """ - https://ai.google.dev/gemini-api/docs/image-generation - - """ - aspect_ratio_map = { - "1024x1024": "1:1", - "1792x1024": "16:9", - "1024x1792": "9:16", - "1280x896": "4:3", - "896x1280": "3:4", - } - return aspect_ratio_map.get(size, "1:1") - - def _transform_image_usage(self, usage_metadata: dict) -> ImageUsage: - """ - Transform Gemini usageMetadata to ImageUsage format - """ - input_tokens_details = ImageUsageInputTokensDetails( - image_tokens=0, - text_tokens=0, - ) - - # Extract detailed token counts from promptTokensDetails - tokens_details = usage_metadata.get("promptTokensDetails", []) - for details in tokens_details: - if isinstance(details, dict): - modality = str(details.get("modality", "")).upper() - raw_token_count = details.get( - "tokenCount", details.get("token_count", 0) - ) - token_count = raw_token_count if isinstance(raw_token_count, int) else 0 - if modality == "TEXT": - input_tokens_details.text_tokens += token_count - elif modality == "IMAGE": - input_tokens_details.image_tokens += token_count - - return ImageUsage( - input_tokens=usage_metadata.get("promptTokenCount", 0), - input_tokens_details=input_tokens_details, - output_tokens=usage_metadata.get("candidatesTokenCount", 0), - total_tokens=usage_metadata.get("totalTokenCount", 0), + return map_openai_image_params_to_gemini( + params=non_default_params, + model=model, + supported_params=self.get_supported_openai_params(model), + optional_params=optional_params, ) def get_complete_url( @@ -127,7 +80,7 @@ class GoogleImageGenConfig(BaseImageGenerationConfig): complete_url = complete_url.rstrip("/") # Gemini Flash Image Preview models use generateContent endpoint - if "gemini" in model: + if is_gemini_image_model(model): complete_url = f"{complete_url}/models/{model}:generateContent" else: # All other Imagen models use predict endpoint @@ -179,10 +132,13 @@ class GoogleImageGenConfig(BaseImageGenerationConfig): } """ # For Gemini Flash Image Preview models, use standard Gemini format - if "gemini" in model: + if is_gemini_image_model(model): request_body: dict = { "contents": [{"parts": [{"text": prompt}]}], - "generationConfig": {"response_modalities": ["IMAGE", "TEXT"]}, + "generationConfig": get_gemini_image_generation_config( + model=model, + optional_params=optional_params, + ), } return request_body else: @@ -200,6 +156,9 @@ class GoogleImageGenConfig(BaseImageGenerationConfig): ) return request_body_obj.model_dump(exclude_none=True) + def _transform_image_usage(self, usage_metadata: dict): + return transform_gemini_image_usage(usage_metadata) + def transform_image_generation_response( self, model: str, @@ -229,7 +188,7 @@ class GoogleImageGenConfig(BaseImageGenerationConfig): model_response.data = [] # Handle different response formats based on model - if "gemini" in model: + if is_gemini_image_model(model): # Gemini Flash Image Preview models return in candidates format candidates = response_data.get("candidates", []) for candidate in candidates: @@ -255,7 +214,7 @@ class GoogleImageGenConfig(BaseImageGenerationConfig): # Extract usage metadata for Gemini models if "usageMetadata" in response_data: - model_response.usage = self._transform_image_usage( + model_response.usage = transform_gemini_image_usage( response_data["usageMetadata"] ) else: diff --git a/litellm/llms/gemini/image_usage_transformation.py b/litellm/llms/gemini/image_usage_transformation.py new file mode 100644 index 00000000000..5a55bdeffb1 --- /dev/null +++ b/litellm/llms/gemini/image_usage_transformation.py @@ -0,0 +1,73 @@ +from typing import Any + +from litellm.types.utils import ImageUsage, ImageUsageInputTokensDetails + + +def _get_token_count(details: dict) -> int: + raw_token_count = details.get("tokenCount", details.get("token_count", 0)) + return raw_token_count if isinstance(raw_token_count, int) else 0 + + +def _get_modality_token_details(usage_metadata: dict, *details_keys: str) -> list: + for details_key in details_keys: + details = usage_metadata.get(details_key) + if isinstance(details, list): + return details + return [] + + +def _sum_modality_token_details( + usage_metadata: dict, *details_keys: str +) -> ImageUsageInputTokensDetails: + tokens_details = ImageUsageInputTokensDetails( + image_tokens=0, + text_tokens=0, + ) + + for details in _get_modality_token_details(usage_metadata, *details_keys): + if isinstance(details, dict): + modality = str(details.get("modality", "")).upper() + token_count = _get_token_count(details) + if modality == "TEXT": + tokens_details.text_tokens += token_count + elif modality == "IMAGE": + tokens_details.image_tokens += token_count + + return tokens_details + + +def transform_gemini_image_usage(usage_metadata: dict) -> ImageUsage: + """ + Transform Gemini usageMetadata to ImageUsage format. + """ + input_tokens_details = _sum_modality_token_details( + usage_metadata, "promptTokensDetails", "prompt_tokens_details" + ) + output_tokens = usage_metadata.get("candidatesTokenCount", 0) + output_tokens_details = _sum_modality_token_details( + usage_metadata, "candidatesTokensDetails", "candidates_tokens_details" + ) + + if not _get_modality_token_details( + usage_metadata, "candidatesTokensDetails", "candidates_tokens_details" + ): + output_tokens_details.image_tokens = output_tokens + else: + known_output_tokens = ( + output_tokens_details.text_tokens + output_tokens_details.image_tokens + ) + if output_tokens > known_output_tokens: + output_tokens_details.text_tokens += output_tokens - known_output_tokens + + usage_payload: dict[str, Any] = { + "input_tokens": usage_metadata.get("promptTokenCount", 0), + "input_tokens_details": input_tokens_details, + "output_tokens": output_tokens, + "total_tokens": usage_metadata.get("totalTokenCount", 0), + "prompt_tokens": usage_metadata.get("promptTokenCount", 0), + "prompt_tokens_details": input_tokens_details.model_dump(), + "completion_tokens": output_tokens, + "completion_tokens_details": output_tokens_details.model_dump(), + "output_tokens_details": output_tokens_details.model_dump(), + } + return ImageUsage(**usage_payload) diff --git a/litellm/llms/inception/__init__.py b/litellm/llms/inception/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/inception/chat/__init__.py b/litellm/llms/inception/chat/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/inception/chat/transformation.py b/litellm/llms/inception/chat/transformation.py new file mode 100644 index 00000000000..d591f783a99 --- /dev/null +++ b/litellm/llms/inception/chat/transformation.py @@ -0,0 +1,54 @@ +""" +Translate from OpenAI's `/v1/chat/completions` to Inception's `/v1/chat/completions` + +Inception Labs (https://www.inceptionlabs.ai) serves the Mercury family of +diffusion LLMs through an OpenAI-compatible API, so we only need to point the +OpenAI-like handler at the Inception API base and pick up the Inception API key. +""" + +from typing import List, Optional, Tuple + +import litellm +from litellm.secret_managers.main import get_secret_str + +from ...openai_like.chat.transformation import OpenAILikeChatConfig + + +class InceptionChatConfig(OpenAILikeChatConfig): + """ + Inception is OpenAI-compatible with standard endpoints + """ + + @property + def custom_llm_provider(self) -> Optional[str]: + return "inception" + + def get_supported_openai_params(self, model: str) -> List: + return [ + "max_tokens", + "max_completion_tokens", + "temperature", + "stop", + "tools", + "tool_choice", + "stream", + "stream_options", + "response_format", + "reasoning_effort", + "reasoning_summary", + "reasoning_summary_wait", + "diffusing", + "realtime", + ] + + def _get_openai_compatible_provider_info( + self, api_base: Optional[str], api_key: Optional[str] + ) -> Tuple[Optional[str], Optional[str]]: + passed_api_base = api_base + api_base = api_base or get_secret_str("INCEPTION_API_BASE") or "https://api.inceptionlabs.ai/v1" # type: ignore + dynamic_api_key = api_key + if passed_api_base is None or api_key: + dynamic_api_key = ( + api_key or litellm.inception_key or get_secret_str("INCEPTION_API_KEY") + ) + return api_base, dynamic_api_key diff --git a/litellm/llms/inception/completion/__init__.py b/litellm/llms/inception/completion/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/inception/completion/transformation.py b/litellm/llms/inception/completion/transformation.py new file mode 100644 index 00000000000..1035042f6bf --- /dev/null +++ b/litellm/llms/inception/completion/transformation.py @@ -0,0 +1,43 @@ +""" +Inception fill-in-the-middle (FIM) completions. + +Inception's FIM endpoint is OpenAI text-completion compatible: it takes a +`prompt` (prefix) plus an optional `suffix` and returns standard +`choices[].text`. It is served at `/v1/fim/completions` rather than +`/v1/completions`, so routing points the OpenAI client at the `/v1/fim` base +(see the `text-completion-inception` branch in `main.py`). +""" + +from typing import List + +from litellm.llms.openai.completion.transformation import OpenAITextCompletionConfig + + +class InceptionTextCompletionConfig(OpenAITextCompletionConfig): + def get_supported_openai_params(self, model: str) -> List: + return [ + "suffix", + "max_tokens", + "max_completion_tokens", + "top_p", + "frequency_penalty", + "presence_penalty", + "stop", + "stream", + "stream_options", + ] + + def map_openai_params( + self, + non_default_params: dict, + optional_params: dict, + model: str, + drop_params: bool, + ) -> dict: + supported_params = self.get_supported_openai_params(model) + for param, value in non_default_params.items(): + if param == "max_completion_tokens": + optional_params["max_tokens"] = value + elif param in supported_params: + optional_params[param] = value + return optional_params diff --git a/litellm/llms/langflow/__init__.py b/litellm/llms/langflow/__init__.py new file mode 100644 index 00000000000..d1270fc91f5 --- /dev/null +++ b/litellm/llms/langflow/__init__.py @@ -0,0 +1 @@ +"""LangFlow LLM provider for LiteLLM.""" diff --git a/litellm/llms/langflow/a2a.py b/litellm/llms/langflow/a2a.py new file mode 100644 index 00000000000..dbe3e02401d --- /dev/null +++ b/litellm/llms/langflow/a2a.py @@ -0,0 +1,37 @@ +import hashlib +from typing import Any, Dict, Optional + + +def get_session_id_from_a2a_params(params: Dict[str, Any]) -> Optional[str]: + message = params.get("message", {}) + if isinstance(message, dict): + return message.get("contextId") + return getattr(message, "contextId", None) + + +def scope_session_to_principal(session_id: str, principal: Optional[str]) -> str: + """ + Bind a client-supplied A2A contextId to the authenticated principal. + + Without this, two distinct keys authorized for the same LangFlow agent could + set the same contextId and read/append to each other's LangFlow memory. The + principal is hashed (it is already a hashed token) so the raw value is never + sent to the LangFlow backend, while the original contextId is kept as a + suffix for operator-side correlation. + """ + if not principal: + return session_id + principal_prefix = hashlib.sha256(principal.encode("utf-8")).hexdigest()[:16] + return f"{principal_prefix}-{session_id}" + + +def merge_a2a_session_into_litellm_params( + litellm_params: Dict[str, Any], + params: Dict[str, Any], + principal: Optional[str] = None, +) -> Dict[str, Any]: + merged = dict(litellm_params) + session_id = get_session_id_from_a2a_params(params) + if session_id and "session_id" not in merged: + merged["session_id"] = scope_session_to_principal(session_id, principal) + return merged diff --git a/litellm/llms/langflow/chat/__init__.py b/litellm/llms/langflow/chat/__init__.py new file mode 100644 index 00000000000..286b12e31f1 --- /dev/null +++ b/litellm/llms/langflow/chat/__init__.py @@ -0,0 +1 @@ +"""LangFlow chat transformation.""" diff --git a/litellm/llms/langflow/chat/transformation.py b/litellm/llms/langflow/chat/transformation.py new file mode 100644 index 00000000000..f898163ad02 --- /dev/null +++ b/litellm/llms/langflow/chat/transformation.py @@ -0,0 +1,327 @@ +"""LangFlow run API: POST {api_base}/api/v1/run/{flow_id}""" + +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from urllib.parse import quote + +import httpx + +from litellm._logging import verbose_logger +from litellm.litellm_core_utils.prompt_templates.common_utils import ( + convert_content_list_to_str, +) +from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException +from litellm.types.llms.openai import AllMessageValues +from litellm.types.utils import Choices, Message, ModelResponse, Usage + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler + from litellm.utils import CustomStreamWrapper + + LiteLLMLoggingObj = _LiteLLMLoggingObj +else: + LiteLLMLoggingObj = Any + HTTPHandler = Any + AsyncHTTPHandler = Any + CustomStreamWrapper = Any + + +class LangFlowError(BaseLLMException): + """Exception class for LangFlow API errors.""" + + pass + + +class LangFlowConfig(BaseConfig): + """ + Configuration for the LangFlow API. + + LangFlow is a visual, low-code platform for building AI agents and pipelines. + Each flow has a unique flow_id and is invoked via a simple HTTP endpoint. + """ + + def __init__(self, **kwargs): + super().__init__(**kwargs) + + def _get_openai_compatible_provider_info( + self, + api_base: Optional[str], + api_key: Optional[str], + ) -> Tuple[Optional[str], Optional[str]]: + from litellm.secret_managers.main import get_secret_str + + api_base = ( + api_base or get_secret_str("LANGFLOW_API_BASE") or "http://localhost:7860" + ) + api_key = api_key or get_secret_str("LANGFLOW_API_KEY") + return api_base, api_key + + def get_supported_openai_params(self, model: str) -> List[str]: + return ["stream"] + + def map_openai_params( + self, + non_default_params: dict, + optional_params: dict, + model: str, + drop_params: bool, + ) -> dict: + return optional_params + + def _get_flow_id(self, model: str, optional_params: dict) -> str: + """ + Extract flow_id from the authorized model name only. + + Model format: "langflow/{flow_id}". Request kwargs must not override + flow_id (would allow calling another flow with the same API key). + """ + if optional_params.get("flow_id") is not None: + raise LangFlowError( + status_code=400, + message=( + "flow_id cannot be set via request parameters; " + "use model langflow/{flow_id}" + ), + ) + + flow_id = (model.split("/", 1)[1] if "/" in model else model).strip() + if not flow_id: + raise LangFlowError( + status_code=400, + message="flow_id is required; use model langflow/{flow_id}", + ) + return flow_id + + def get_complete_url( + self, + api_base: Optional[str], + api_key: Optional[str], + model: str, + optional_params: dict, + litellm_params: dict, + stream: Optional[bool] = None, + ) -> str: + if api_base is None: + raise ValueError( + "api_base is required for LangFlow. Set it via LANGFLOW_API_BASE env var or api_base parameter." + ) + + api_base = api_base.rstrip("/") + flow_id = quote(self._get_flow_id(model, optional_params), safe="") + return f"{api_base}/api/v1/run/{flow_id}" + + def _get_last_user_message(self, messages: List[AllMessageValues]) -> str: + """Extract the text of the last user message to use as input_value.""" + for msg in reversed(messages): + if msg.get("role") == "user": + content = msg.get("content", "") + if isinstance(content, list): + content = convert_content_list_to_str(msg) + if not isinstance(content, str): + content = str(content) + return content + + # Fallback: use last message regardless of role + if messages: + content = messages[-1].get("content", "") + if isinstance(content, list): + content = convert_content_list_to_str(messages[-1]) + if not isinstance(content, str): + content = str(content) + return content + + return "" + + def _reject_caller_tweaks(self, params: dict) -> None: + if params.get("tweaks") is not None: + raise LangFlowError( + status_code=400, + message=( + "tweaks cannot be set via request parameters; they would " + "override the operator-configured LangFlow flow components" + ), + ) + + def transform_request( + self, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + litellm_params: dict, + headers: dict, + ) -> dict: + """ + Transform the request to LangFlow format. + + LangFlow request format: + { + "input_value": "", + "input_type": "chat", + "output_type": "chat", + "session_id": "" + } + """ + self._reject_caller_tweaks(optional_params) + + input_value = self._get_last_user_message(messages) + + payload: Dict[str, Any] = { + "input_value": input_value, + "input_type": optional_params.get("input_type", "chat"), + "output_type": optional_params.get("output_type", "chat"), + } + + session_id = optional_params.get("session_id") + if session_id: + payload["session_id"] = session_id + + verbose_logger.debug(f"LangFlow request payload: {payload}") + return payload + + def _extract_content_from_response(self, response_json: dict) -> Optional[str]: + """ + Extract the assistant text from a LangFlow run response. + + Expected structure: + {"outputs": [{"outputs": [{"results": {"message": {"text": "..."}}}]}]} + + Returns None when no message text is present so the caller can surface an + explicit error instead of forwarding a raw JSON blob as the answer. + """ + outputs = response_json.get("outputs", []) + if not (isinstance(outputs, list) and outputs): + return None + + first_output = outputs[0] + if not isinstance(first_output, dict): + return None + + inner_outputs = first_output.get("outputs", []) + if not (isinstance(inner_outputs, list) and inner_outputs): + return None + + first_inner = inner_outputs[0] + if not isinstance(first_inner, dict): + return None + + results = first_inner.get("results", {}) + if isinstance(results, dict): + message = results.get("message", {}) + if isinstance(message, dict) and message.get("text"): + return message["text"] + + outputs_dict = first_inner.get("outputs", {}) + if isinstance(outputs_dict, dict): + for val in outputs_dict.values(): + if isinstance(val, dict): + msg = val.get("message", {}) + if isinstance(msg, dict) and msg.get("text"): + return msg["text"] + + return None + + def transform_response( + self, + model: str, + raw_response: httpx.Response, + model_response: ModelResponse, + logging_obj: LiteLLMLoggingObj, + request_data: dict, + messages: List[AllMessageValues], + optional_params: dict, + litellm_params: dict, + encoding: Any, + api_key: Optional[str] = None, + json_mode: Optional[bool] = None, + ) -> ModelResponse: + try: + response_json = raw_response.json() + except Exception as e: + raise LangFlowError( + message=f"LangFlow returned a non-JSON response: {e}", + status_code=raw_response.status_code, + ) + + verbose_logger.debug(f"LangFlow response: {response_json}") + + content = self._extract_content_from_response(response_json) + if content is None: + raise LangFlowError( + message=( + "Could not extract a message from the LangFlow response; " + "ensure the flow ends in a Chat Output component" + ), + status_code=500, + ) + + message = Message(content=content, role="assistant") + choice = Choices(finish_reason="stop", index=0, message=message) + + model_response.choices = [choice] + model_response.model = model + + try: + from litellm.utils import token_counter + + prompt_tokens = token_counter(model=model, messages=messages) + completion_tokens = token_counter( + model=model, text=content, count_response_tokens=True + ) + usage = Usage( + prompt_tokens=prompt_tokens, + completion_tokens=completion_tokens, + total_tokens=prompt_tokens + completion_tokens, + ) + setattr(model_response, "usage", usage) + except Exception as e: + verbose_logger.warning(f"Failed to calculate token usage: {e}") + + return model_response + + def sign_request( + self, + headers: dict, + optional_params: dict, + request_data: dict, + api_base: str, + api_key: Optional[str] = None, + model: Optional[str] = None, + stream: Optional[bool] = None, + fake_stream: Optional[bool] = None, + ) -> Tuple[dict, Optional[bytes]]: + self._reject_caller_tweaks(request_data) + return headers, None + + def validate_environment( + self, + headers: dict, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + litellm_params: dict, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + ) -> dict: + headers["Content-Type"] = "application/json" + + if api_key: + headers["x-api-key"] = api_key + + return headers + + def get_error_class( + self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] + ) -> BaseLLMException: + return LangFlowError(status_code=status_code, message=error_message) + + @property + def supports_stream_param_in_request_body(self) -> bool: + return False + + def should_fake_stream( + self, + model: Optional[str], + stream: Optional[bool], + custom_llm_provider: Optional[str] = None, + ) -> bool: + return stream is True diff --git a/litellm/llms/openai/responses/transformation.py b/litellm/llms/openai/responses/transformation.py index 5043d25ee37..c18f2216f61 100644 --- a/litellm/llms/openai/responses/transformation.py +++ b/litellm/llms/openai/responses/transformation.py @@ -299,11 +299,8 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): or litellm.openai_key or get_secret_str("OPENAI_API_KEY") ) - headers.update( - { - "Authorization": f"Bearer {api_key}", - } - ) + headers.setdefault("Content-Type", "application/json") + headers["Authorization"] = f"Bearer {api_key}" return headers def get_complete_url( diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index 4f5846cc5b6..ef7bf82bfae 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -996,7 +996,19 @@ def _gemini_convert_messages_with_history( # noqa: PLR0915 excluded_keys=["thoughtSignature"], ): assistant_content.append(gemini_tool_call_part) - last_message_with_tool_calls = assistant_msg + # Only record this as the active tool-call message when it actually + # carries tool calls. The `if` guard above is also entered for a + # text-only assistant message (`assistant_msg.get("tool_calls", []) + # is not None` is True for an empty list), so without this check a + # later assistant message with no tool calls would clobber the + # reference. The following tool result would then be matched against + # an assistant message that has no tool_calls, raising "Missing + # corresponding tool call for tool response message". + if ( + assistant_msg.get("tool_calls") + or assistant_msg.get("function_call") is not None + ): + last_message_with_tool_calls = assistant_msg ## HANDLE SERVER-SIDE TOOL INVOCATIONS (context circulation) _psf = assistant_msg.get("provider_specific_fields") diff --git a/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py b/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py index 14a0a406dff..46dedb3d0a4 100644 --- a/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py +++ b/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py @@ -3,6 +3,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union import httpx from litellm import get_model_info +from litellm.exceptions import BadRequestError from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreConfig from litellm.llms.vertex_ai.vertex_llm_base import VertexBase @@ -16,6 +17,8 @@ from litellm.types.vector_stores import ( VectorStoreSearchOptionalRequestParams, VectorStoreSearchResponse, VectorStoreSearchResult, + VertexSearchDataStoreExtraBody, + VertexSearchEngineExtraBody, ) if TYPE_CHECKING: @@ -26,6 +29,31 @@ else: LiteLLMLoggingObj = Any +# Fields that select which data store / serving config to search. These are +# always determined by the request URL path (vector_store_id / vertex_engine_id), +# so allowing them per request could silently redirect the search to a different +# target. Rejected in both data-store and engine/app modes. +VERTEX_SEARCH_TARGET_SELECTING_FIELDS = frozenset( + { + "branch", + "servingConfig", + "entity", + } +) + +# Allowlists of native Discovery Engine SearchRequest fields callers may forward +# via extra_body, derived from the TypedDicts so the type is the source of truth. +# Engine/app mode is a superset (adds dataStoreSpecs, numResultsPerDataStore), +# since an app fans out across multiple member data stores. +VERTEX_SEARCH_DATASTORE_EXTRA_BODY_FIELDS = frozenset( + VertexSearchDataStoreExtraBody.__annotations__ +) + +VERTEX_SEARCH_ENGINE_EXTRA_BODY_FIELDS = frozenset( + VertexSearchEngineExtraBody.__annotations__ +) + + class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase): """ Configuration for Vertex AI Search API Vector Store @@ -36,6 +64,66 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase): def __init__(self): super().__init__() + @staticmethod + def get_supported_extra_body_fields(is_engine: bool = False) -> frozenset: + """ + Native SearchRequest fields callers may forward via ``extra_body``. + + The set depends on which serving config the request targets: + - engine/app mode (``is_engine=True``): includes multi-store fields such + as ``dataStoreSpecs`` and ``numResultsPerDataStore``. + - data-store mode: the engine-only fields are excluded. + """ + if is_engine: + return VERTEX_SEARCH_ENGINE_EXTRA_BODY_FIELDS + return VERTEX_SEARCH_DATASTORE_EXTRA_BODY_FIELDS + + @classmethod + def _filter_extra_body( + cls, extra_body: Dict[str, Any], is_engine: bool = False + ) -> Dict[str, Any]: + """ + Validate ``extra_body`` against the supported-field allowlist for the + active serving config (engine/app vs data store). + + Raises ``BadRequestError`` (HTTP 400) if the caller includes a + target-selecting field (e.g. ``servingConfig``) or any field not + supported for the active mode, so the request fails loudly instead of + silently searching the wrong target. Engine-only fields + (``dataStoreSpecs``, ``numResultsPerDataStore``) are rejected in + data-store mode where they are meaningless. + """ + supported = cls.get_supported_extra_body_fields(is_engine=is_engine) + filtered = { + key: value for key, value in extra_body.items() if value is not None + } + + target_selecting = set(filtered) & VERTEX_SEARCH_TARGET_SELECTING_FIELDS + if target_selecting: + raise BadRequestError( + message=( + "Vertex AI Search extra_body may not set target-selecting fields " + f"{sorted(target_selecting)}: the data store is scoped by " + "vector_store_id / vertex_engine_id and cannot be overridden per request." + ), + model="vertex_ai/search_api", + llm_provider="vertex_ai", + ) + + unsupported = set(filtered) - supported + if unsupported: + mode = "engine/app" if is_engine else "data store" + raise BadRequestError( + message=( + f"Unsupported Vertex AI Search extra_body fields {sorted(unsupported)} " + f"for {mode} mode. Supported fields: {sorted(supported)}." + ), + model="vertex_ai/search_api", + llm_provider="vertex_ai", + ) + + return filtered + def get_auth_credentials( self, litellm_params: dict ) -> BaseVectorStoreAuthCredentials: @@ -133,23 +221,41 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase): extra_body: Optional[Dict[str, Any]] = None, ) -> Tuple[str, Dict[str, Any]]: """ - Transform search request for Vertex AI RAG API + Transform a search request for the Vertex AI Search (Discovery Engine) API. + + Per-request params pass through to the engine: max_num_results maps to + pageSize, and extra_body fields on the supported allowlist + (`get_supported_extra_body_fields`) are merged in with precedence, so + callers can send native Discovery Engine tuning fields such as filter, + boostSpec, or contentSearchSpec. + + The allowlist depends on the serving config: engine/app mode (when + `vertex_engine_id` is set) additionally accepts multi-store fields like + `dataStoreSpecs` and `numResultsPerDataStore`, while data-store mode + rejects them. Target-selecting fields (e.g. servingConfig, branch) are + rejected in both modes: the target is scoped by the URL path + (vector_store_id / vertex_engine_id) and must not be overridable per + request. """ - # Convert query to string if it's a list if isinstance(query, list): query = " ".join(query) - # Vertex AI RAG API endpoint for retrieving contexts url = f"{api_base}:search" - # Construct full rag corpus path - # Build the request body for Vertex AI Search API - request_body = {"query": query, "pageSize": 10} + is_engine = bool(litellm_params.get("vertex_engine_id")) - ######################################################### - # Update logging object with details of the request - ######################################################### - litellm_logging_obj.model_call_details["query"] = query + request_body: Dict[str, Any] = {"query": query, "pageSize": 10} + max_num_results = vector_store_search_optional_params.get("max_num_results") + if max_num_results is not None: + request_body["pageSize"] = max_num_results + if isinstance(extra_body, dict): + request_body.update( + self._filter_extra_body(extra_body, is_engine=is_engine) + ) + + litellm_logging_obj.model_call_details["query"] = request_body.get( + "query", query + ) return url, request_body diff --git a/litellm/llms/vertex_ai/vertex_model_garden/main.py b/litellm/llms/vertex_ai/vertex_model_garden/main.py index 732d5f90dc2..f54b8d93500 100644 --- a/litellm/llms/vertex_ai/vertex_model_garden/main.py +++ b/litellm/llms/vertex_ai/vertex_model_garden/main.py @@ -114,33 +114,18 @@ class VertexAIModelGardenModels(VertexBase): openai_like_chat_completions = OpenAILikeChatHandler() ## CONSTRUCT API BASE + # Skip _check_custom_proxy: its ":verb" URL construction corrupts a + # user-supplied api_base (e.g. Vertex MG dedicated endpoint), and + # OpenAILikeChatHandler already appends "/chat/completions". stream: bool = optional_params.get("stream", False) or False optional_params["stream"] = stream - default_api_base = create_vertex_url( - vertex_location=vertex_location or "us-central1", - vertex_project=vertex_project or project_id, - stream=stream, - model=model, - ) - - if len(default_api_base.split(":")) > 1: - endpoint = default_api_base.split(":")[-1] - else: - endpoint = "" - - _, api_base = self._check_custom_proxy( - api_base=api_base, - custom_llm_provider="vertex_ai", - gemini_api_key=None, - endpoint=endpoint, - stream=stream, - auth_header=None, - url=default_api_base, - model=model, - vertex_project=vertex_project or project_id, - vertex_location=vertex_location or "us-central1", - vertex_api_version="v1beta1", - ) + if api_base is None: + api_base = create_vertex_url( + vertex_location=vertex_location or "us-central1", + vertex_project=vertex_project or project_id, + stream=stream, + model=model, + ) # Publisher/catalog models: model id must be sent in the JSON body (OpenAPI route). # Single-segment endpoint ids: model is encoded in the URL path; body model stays empty. if not _vertex_model_garden_model_id_in_json_body(model): diff --git a/litellm/main.py b/litellm/main.py index 09c70998cf7..da8624d11b8 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -641,6 +641,7 @@ async def acompletion( # noqa: PLR0915 if ( custom_llm_provider == "text-completion-openai" or custom_llm_provider == "text-completion-codestral" + or custom_llm_provider == "text-completion-inception" ) and isinstance(response, TextCompletionResponse): response = litellm.OpenAITextCompletionConfig().convert_to_chat_model_response_object( response_object=response, @@ -3803,6 +3804,67 @@ def completion( # type: ignore # noqa: PLR0915 ): return _model_response response = _model_response + elif custom_llm_provider == "text-completion-inception": + passed_api_base = ( + api_base + or optional_params.pop("api_base", None) + or optional_params.pop("base_url", None) + ) + api_base = ( + passed_api_base + or get_secret_str("INCEPTION_API_BASE") + or "https://api.inceptionlabs.ai/v1" + ) + # FIM is served at `/v1/fim/completions`; the OpenAI client appends + # `/completions`, so point it at the `/v1/fim` base. + api_base = api_base.rstrip("/") + if not api_base.endswith("/fim"): + api_base += "/fim" + + # Don't forward the server-managed Inception key to a caller-supplied + # api_base; only resolve it for the default/server base, or when the + # caller passes their own key. + if passed_api_base is None or api_key: + api_key = ( + api_key + or litellm.inception_key + or get_secret_str("INCEPTION_API_KEY") + ) + + _response = openai_text_completions.completion( + model=model, + messages=messages, + model_response=model_response, + print_verbose=print_verbose, + api_key=api_key, # type: ignore[arg-type] + custom_llm_provider="text-completion-inception", + api_base=api_base, + acompletion=acompletion, + client=client, + logging_obj=logging, + optional_params=optional_params, + litellm_params=litellm_params, + logger_fn=logger_fn, + timeout=timeout, # type: ignore + ) + + if ( + optional_params.get("stream", False) is False + and acompletion is False + and text_completion is False + ): + _response = litellm.OpenAITextCompletionConfig().convert_to_chat_model_response_object( + response_object=_response, model_response_object=model_response + ) + + if optional_params.get("stream", False) or acompletion is True: + logging.post_call( + input=messages, + api_key=api_key, + original_response=_response, + additional_args={"headers": headers}, + ) + response = _response elif custom_llm_provider in ("sagemaker_chat", "sagemaker_nova"): # boto3 reads keys from .env # sagemaker_chat: HF Messages API endpoints @@ -4503,6 +4565,39 @@ def completion( # type: ignore # noqa: PLR0915 client=client, ) + elif custom_llm_provider == "langflow": + # LangFlow - Visual AI Agent Platform + from litellm.llms.langflow.chat.transformation import LangFlowConfig + + ( + api_base, + api_key, + ) = LangFlowConfig()._get_openai_compatible_provider_info( + api_base=api_base or litellm.api_base, + api_key=api_key or litellm.api_key, + ) + + headers = headers or litellm.headers + + response = base_llm_http_handler.completion( + model=model, + stream=stream, + messages=messages, + acompletion=acompletion, + api_base=api_base, + model_response=model_response, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + custom_llm_provider=custom_llm_provider, + timeout=timeout, + headers=headers, + encoding=_get_encoding(), + api_key=api_key, + logging_obj=logging, + client=client, + ) + else: raise LiteLLMUnknownProvider( model=model, custom_llm_provider=custom_llm_provider diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index a604aafa540..ed6de4fa6b7 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -1075,6 +1075,7 @@ }, "eu.anthropic.claude-opus-4-6-v1": { "cache_creation_input_token_cost": 6.875e-06, + "cache_creation_input_token_cost_above_1hr": 1.1e-05, "cache_read_input_token_cost": 5.5e-07, "input_cost_per_token": 5.5e-06, "litellm_provider": "bedrock_converse", @@ -1104,6 +1105,7 @@ }, "au.anthropic.claude-opus-4-6-v1": { "cache_creation_input_token_cost": 6.875e-06, + "cache_creation_input_token_cost_above_1hr": 1.1e-05, "cache_read_input_token_cost": 5.5e-07, "input_cost_per_token": 5.5e-06, "litellm_provider": "bedrock_converse", @@ -1241,6 +1243,7 @@ }, "eu.anthropic.claude-opus-4-7": { "cache_creation_input_token_cost": 6.875e-06, + "cache_creation_input_token_cost_above_1hr": 1.1e-05, "cache_read_input_token_cost": 5.5e-07, "input_cost_per_token": 5.5e-06, "litellm_provider": "bedrock_converse", @@ -1271,6 +1274,7 @@ }, "au.anthropic.claude-opus-4-7": { "cache_creation_input_token_cost": 6.875e-06, + "cache_creation_input_token_cost_above_1hr": 1.1e-05, "cache_read_input_token_cost": 5.5e-07, "input_cost_per_token": 5.5e-06, "litellm_provider": "bedrock_converse", @@ -1543,6 +1547,7 @@ }, "eu.anthropic.claude-sonnet-4-6": { "cache_creation_input_token_cost": 4.125e-06, + "cache_creation_input_token_cost_above_1hr": 6.6e-06, "cache_read_input_token_cost": 3.3e-07, "input_cost_per_token": 3.3e-06, "litellm_provider": "bedrock_converse", @@ -1571,6 +1576,7 @@ }, "au.anthropic.claude-sonnet-4-6": { "cache_creation_input_token_cost": 4.125e-06, + "cache_creation_input_token_cost_above_1hr": 6.6e-06, "cache_read_input_token_cost": 3.3e-07, "input_cost_per_token": 3.3e-06, "litellm_provider": "bedrock_converse", @@ -1599,6 +1605,7 @@ }, "jp.anthropic.claude-sonnet-4-6": { "cache_creation_input_token_cost": 4.125e-06, + "cache_creation_input_token_cost_above_1hr": 6.6e-06, "cache_read_input_token_cost": 3.3e-07, "input_cost_per_token": 3.3e-06, "litellm_provider": "bedrock_converse", @@ -1995,11 +2002,13 @@ }, "au.anthropic.claude-sonnet-4-5-20250929-v1:0": { "cache_creation_input_token_cost": 4.125e-06, + "cache_creation_input_token_cost_above_1hr": 6.6e-06, "cache_read_input_token_cost": 3.3e-07, "input_cost_per_token": 3.3e-06, "input_cost_per_token_above_200k_tokens": 6.6e-06, "output_cost_per_token_above_200k_tokens": 2.475e-05, "cache_creation_input_token_cost_above_200k_tokens": 8.25e-06, + "cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.32e-05, "cache_read_input_token_cost_above_200k_tokens": 6.6e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 200000, @@ -7494,6 +7503,27 @@ "supports_video_input": true, "supports_vision": true }, + "azure_ai/kimi-k2.6": { + "input_cost_per_token": 9.5e-07, + "litellm_provider": "azure_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 4e-06, + "source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-kimi-k2-6-in-microsoft-foundry/4513125", + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_reasoning": true, + "supports_tool_choice": true, + "supports_vision": true + }, "azure_ai/ministral-3b": { "input_cost_per_token": 4e-08, "litellm_provider": "azure_ai", @@ -12682,7 +12712,8 @@ "litellm_provider": "deepinfra", "mode": "chat", "supports_tool_choice": true, - "supports_function_calling": true + "supports_function_calling": true, + "supports_image_size": false }, "deepinfra/google/gemini-2.5-pro": { "max_tokens": 1000000, @@ -13583,6 +13614,7 @@ }, "eu.anthropic.claude-haiku-4-5-20251001-v1:0": { "cache_creation_input_token_cost": 1.375e-06, + "cache_creation_input_token_cost_above_1hr": 2.2e-06, "cache_read_input_token_cost": 1.1e-07, "input_cost_per_token": 1.1e-06, "deprecation_date": "2026-10-15", @@ -13787,11 +13819,13 @@ }, "eu.anthropic.claude-sonnet-4-5-20250929-v1:0": { "cache_creation_input_token_cost": 4.125e-06, + "cache_creation_input_token_cost_above_1hr": 6.6e-06, "cache_read_input_token_cost": 3.3e-07, "input_cost_per_token": 3.3e-06, "input_cost_per_token_above_200k_tokens": 6.6e-06, "output_cost_per_token_above_200k_tokens": 2.475e-05, "cache_creation_input_token_cost_above_200k_tokens": 8.25e-06, + "cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.32e-05, "cache_read_input_token_cost_above_200k_tokens": 6.6e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 200000, @@ -15006,7 +15040,8 @@ "search_context_size_medium": 0.035, "search_context_size_high": 0.035 }, - "supports_service_tier": true + "supports_service_tier": true, + "supports_image_size": false }, "gemini-2.5-flash-image": { "cache_read_input_token_cost": 3e-08, @@ -15056,7 +15091,8 @@ "supports_vision": true, "supports_web_search": false, "tpm": 8000000, - "supports_service_tier": true + "supports_service_tier": true, + "supports_image_size": false }, "gemini-3-pro-image-preview": { "input_cost_per_image": 0.0011, @@ -15196,10 +15232,16 @@ "supports_service_tier": true }, "gemini-3.1-flash-lite": { - "cache_read_input_token_cost": 4.5e-08, - "cache_read_input_token_cost_per_audio_token": 9e-08, - "input_cost_per_audio_token": 9e-07, - "input_cost_per_token": 4.5e-07, + "cache_read_input_token_cost": 2.5e-08, + "cache_read_input_token_cost_batches": 1.25e-08, + "cache_read_input_token_cost_flex": 1.25e-08, + "cache_read_input_token_cost_per_audio_token": 5e-08, + "cache_read_input_token_cost_priority": 4.5e-08, + "input_cost_per_audio_token": 5e-07, + "input_cost_per_token": 2.5e-07, + "input_cost_per_token_batches": 1.25e-07, + "input_cost_per_token_flex": 1.25e-07, + "input_cost_per_token_priority": 4.5e-07, "litellm_provider": "vertex_ai-language-models", "max_audio_length_hours": 8.4, "max_audio_per_prompt": 1, @@ -15211,9 +15253,12 @@ "max_video_length": 1, "max_videos_per_prompt": 10, "mode": "chat", - "output_cost_per_reasoning_token": 2.7e-06, - "output_cost_per_token": 2.7e-06, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models", + "output_cost_per_reasoning_token": 1.5e-06, + "output_cost_per_token": 1.5e-06, + "output_cost_per_token_batches": 7.5e-07, + "output_cost_per_token_flex": 7.5e-07, + "output_cost_per_token_priority": 2.7e-06, + "source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-3.1-flash-lite", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -15336,7 +15381,8 @@ "search_context_size_medium": 0.035, "search_context_size_high": 0.035 }, - "supports_service_tier": true + "supports_service_tier": true, + "supports_image_size": false }, "gemini-2.5-flash-lite-preview-09-2025": { "cache_read_input_token_cost": 1e-08, @@ -15386,7 +15432,8 @@ "search_context_size_low": 0.035, "search_context_size_medium": 0.035, "search_context_size_high": 0.035 - } + }, + "supports_image_size": false }, "gemini-2.5-flash-preview-09-2025": { "cache_read_input_token_cost": 7.5e-08, @@ -15436,7 +15483,8 @@ "search_context_size_low": 0.035, "search_context_size_medium": 0.035, "search_context_size_high": 0.035 - } + }, + "supports_image_size": false }, "gemini-live-2.5-flash-preview-native-audio-09-2025": { "cache_read_input_token_cost": 7.5e-08, @@ -15587,7 +15635,8 @@ "search_context_size_low": 0.035, "search_context_size_medium": 0.035, "search_context_size_high": 0.035 - } + }, + "supports_image_size": false }, "gemini-2.5-pro": { "cache_read_input_token_cost": 1.25e-07, @@ -16597,7 +16646,8 @@ "search_context_size_medium": 0.035, "search_context_size_high": 0.035 }, - "supports_service_tier": true + "supports_service_tier": true, + "supports_image_size": false }, "gemini/gemini-2.5-flash-image": { "cache_read_input_token_cost": 3e-08, @@ -16653,7 +16703,8 @@ "search_context_size_medium": 0.035, "search_context_size_high": 0.035 }, - "supports_service_tier": true + "supports_service_tier": true, + "supports_image_size": false }, "gemini/gemini-3-pro-image-preview": { "input_cost_per_image": 0.0011, @@ -16832,7 +16883,8 @@ "search_context_size_medium": 0.035, "search_context_size_high": 0.035 }, - "supports_service_tier": true + "supports_service_tier": true, + "supports_image_size": false }, "gemini/gemini-2.5-flash-lite-preview-09-2025": { "cache_read_input_token_cost": 1e-08, @@ -16884,7 +16936,8 @@ "search_context_size_low": 0.035, "search_context_size_medium": 0.035, "search_context_size_high": 0.035 - } + }, + "supports_image_size": false }, "gemini/gemini-2.5-flash-preview-09-2025": { "cache_read_input_token_cost": 7.5e-08, @@ -16936,7 +16989,8 @@ "search_context_size_low": 0.035, "search_context_size_medium": 0.035, "search_context_size_high": 0.035 - } + }, + "supports_image_size": false }, "gemini/gemini-flash-latest": { "cache_read_input_token_cost": 7.5e-08, @@ -17093,7 +17147,8 @@ "search_context_size_low": 0.035, "search_context_size_medium": 0.035, "search_context_size_high": 0.035 - } + }, + "supports_image_size": false }, "gemini/gemini-2.5-flash-preview-tts": { "input_cost_per_token": 3e-07, @@ -17318,10 +17373,16 @@ "supports_service_tier": true }, "gemini/gemini-3.1-flash-lite": { - "cache_read_input_token_cost": 4.5e-08, - "cache_read_input_token_cost_per_audio_token": 9e-08, - "input_cost_per_audio_token": 9e-07, - "input_cost_per_token": 4.5e-07, + "cache_read_input_token_cost": 2.5e-08, + "cache_read_input_token_cost_batches": 1.25e-08, + "cache_read_input_token_cost_flex": 1.25e-08, + "cache_read_input_token_cost_per_audio_token": 5e-08, + "cache_read_input_token_cost_priority": 4.5e-08, + "input_cost_per_audio_token": 5e-07, + "input_cost_per_token": 2.5e-07, + "input_cost_per_token_batches": 1.25e-07, + "input_cost_per_token_flex": 1.25e-07, + "input_cost_per_token_priority": 4.5e-07, "litellm_provider": "gemini", "max_audio_length_hours": 8.4, "max_audio_per_prompt": 1, @@ -17333,10 +17394,13 @@ "max_video_length": 1, "max_videos_per_prompt": 10, "mode": "chat", - "output_cost_per_reasoning_token": 2.7e-06, - "output_cost_per_token": 2.7e-06, + "output_cost_per_reasoning_token": 1.5e-06, + "output_cost_per_token": 1.5e-06, + "output_cost_per_token_batches": 7.5e-07, + "output_cost_per_token_flex": 7.5e-07, + "output_cost_per_token_priority": 2.7e-06, "rpm": 15, - "source": "https://ai.google.dev/gemini-api/docs/pricing", + "source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-3.1-flash-lite", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -18098,23 +18162,22 @@ }, "github_copilot/claude-haiku-4.5": { "litellm_provider": "github_copilot", - "max_input_tokens": 200000, - "max_output_tokens": 32000, - "max_tokens": 32000, + "max_input_tokens": 128000, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "supported_endpoints": [ "/v1/chat/completions" ], "supports_function_calling": true, "supports_parallel_function_calling": true, - "supports_vision": true, - "supports_reasoning": true + "supports_vision": true }, "github_copilot/claude-opus-4.5": { "litellm_provider": "github_copilot", - "max_input_tokens": 200000, - "max_output_tokens": 32000, - "max_tokens": 32000, + "max_input_tokens": 128000, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "supported_endpoints": [ "/v1/chat/completions" @@ -18122,7 +18185,6 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_vision": true, - "supports_reasoning": true, "supports_output_config": true }, "github_copilot/claude-opus-4.6-fast": { @@ -18138,22 +18200,6 @@ "supports_parallel_function_calling": true, "supports_vision": true }, - "github_copilot/claude-opus-4.7": { - "litellm_provider": "github_copilot", - "max_input_tokens": 200000, - "max_output_tokens": 64000, - "max_tokens": 64000, - "mode": "chat", - "supported_endpoints": [ - "/v1/chat/completions", - "/v1/messages" - ], - "supports_vision": true, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_response_schema": true, - "supports_reasoning": true - }, "github_copilot/claude-opus-41": { "litellm_provider": "github_copilot", "max_input_tokens": 80000, @@ -18180,33 +18226,16 @@ }, "github_copilot/claude-sonnet-4.5": { "litellm_provider": "github_copilot", - "max_input_tokens": 200000, - "max_output_tokens": 32000, - "max_tokens": 32000, + "max_input_tokens": 128000, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "supported_endpoints": [ "/v1/chat/completions" ], "supports_function_calling": true, "supports_parallel_function_calling": true, - "supports_vision": true, - "supports_reasoning": true - }, - "github_copilot/claude-sonnet-4.6": { - "litellm_provider": "github_copilot", - "max_input_tokens": 200000, - "max_output_tokens": 32000, - "max_tokens": 32000, - "mode": "chat", - "supported_endpoints": [ - "/v1/chat/completions", - "/v1/messages" - ], - "supports_vision": true, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_response_schema": true, - "supports_reasoning": true + "supports_vision": true }, "github_copilot/gemini-2.5-pro": { "litellm_provider": "github_copilot", @@ -18216,25 +18245,7 @@ "mode": "chat", "supports_function_calling": true, "supports_parallel_function_calling": true, - "supports_vision": true, - "supported_endpoints": [ - "/v1/chat/completions" - ], - "supports_reasoning": true - }, - "github_copilot/gemini-3-flash-preview": { - "litellm_provider": "github_copilot", - "max_input_tokens": 128000, - "max_output_tokens": 64000, - "max_tokens": 64000, - "mode": "chat", - "supported_endpoints": [ - "/v1/chat/completions" - ], - "supports_vision": true, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_reasoning": true + "supports_vision": true }, "github_copilot/gemini-3-pro-preview": { "litellm_provider": "github_copilot", @@ -18246,30 +18257,13 @@ "supports_parallel_function_calling": true, "supports_vision": true }, - "github_copilot/gemini-3.1-pro-preview": { - "litellm_provider": "github_copilot", - "max_input_tokens": 128000, - "max_output_tokens": 64000, - "max_tokens": 64000, - "mode": "chat", - "supported_endpoints": [ - "/v1/chat/completions" - ], - "supports_vision": true, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_reasoning": true - }, "github_copilot/gpt-3.5-turbo": { "litellm_provider": "github_copilot", "max_input_tokens": 16384, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", - "supports_function_calling": true, - "supported_endpoints": [ - "/v1/chat/completions" - ] + "supports_function_calling": true }, "github_copilot/gpt-3.5-turbo-0613": { "litellm_provider": "github_copilot", @@ -18277,10 +18271,7 @@ "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", - "supports_function_calling": true, - "supported_endpoints": [ - "/v1/chat/completions" - ] + "supports_function_calling": true }, "github_copilot/gpt-4": { "litellm_provider": "github_copilot", @@ -18288,22 +18279,7 @@ "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", - "supports_function_calling": true, - "supported_endpoints": [ - "/v1/chat/completions" - ] - }, - "github_copilot/gpt-4-0125-preview": { - "litellm_provider": "github_copilot", - "max_input_tokens": 128000, - "max_output_tokens": 4096, - "max_tokens": 4096, - "mode": "chat", - "supported_endpoints": [ - "/v1/chat/completions" - ], - "supports_function_calling": true, - "supports_parallel_function_calling": true + "supports_function_calling": true }, "github_copilot/gpt-4-0613": { "litellm_provider": "github_copilot", @@ -18311,22 +18287,16 @@ "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", - "supports_function_calling": true, - "supported_endpoints": [ - "/v1/chat/completions" - ] + "supports_function_calling": true }, "github_copilot/gpt-4-o-preview": { "litellm_provider": "github_copilot", - "max_input_tokens": 128000, + "max_input_tokens": 64000, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supported_endpoints": [ - "/v1/chat/completions" - ] + "supports_parallel_function_calling": true }, "github_copilot/gpt-4.1": { "litellm_provider": "github_copilot", @@ -18337,10 +18307,7 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true, - "supported_endpoints": [ - "/v1/chat/completions" - ] + "supports_vision": true }, "github_copilot/gpt-4.1-2025-04-14": { "litellm_provider": "github_copilot", @@ -18351,89 +18318,68 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true, - "supported_endpoints": [ - "/v1/chat/completions" - ] + "supports_vision": true }, "github_copilot/gpt-41-copilot": { "litellm_provider": "github_copilot", - "mode": "chat" + "mode": "completion" }, "github_copilot/gpt-4o": { "litellm_provider": "github_copilot", - "max_input_tokens": 128000, + "max_input_tokens": 64000, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", "supports_function_calling": true, "supports_parallel_function_calling": true, - "supports_vision": true, - "supported_endpoints": [ - "/v1/chat/completions" - ] + "supports_vision": true }, "github_copilot/gpt-4o-2024-05-13": { "litellm_provider": "github_copilot", - "max_input_tokens": 128000, + "max_input_tokens": 64000, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", "supports_function_calling": true, "supports_parallel_function_calling": true, - "supports_vision": true, - "supported_endpoints": [ - "/v1/chat/completions" - ] + "supports_vision": true }, "github_copilot/gpt-4o-2024-08-06": { "litellm_provider": "github_copilot", - "max_input_tokens": 128000, + "max_input_tokens": 64000, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supported_endpoints": [ - "/v1/chat/completions" - ] + "supports_parallel_function_calling": true }, "github_copilot/gpt-4o-2024-11-20": { "litellm_provider": "github_copilot", - "max_input_tokens": 128000, + "max_input_tokens": 64000, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "supports_function_calling": true, "supports_parallel_function_calling": true, - "supports_vision": true, - "supported_endpoints": [ - "/v1/chat/completions" - ] + "supports_vision": true }, "github_copilot/gpt-4o-mini": { "litellm_provider": "github_copilot", - "max_input_tokens": 128000, + "max_input_tokens": 64000, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supported_endpoints": [ - "/v1/chat/completions" - ] + "supports_parallel_function_calling": true }, "github_copilot/gpt-4o-mini-2024-07-18": { "litellm_provider": "github_copilot", - "max_input_tokens": 128000, + "max_input_tokens": 64000, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supported_endpoints": [ - "/v1/chat/completions" - ] + "supports_parallel_function_calling": true }, "github_copilot/gpt-5": { "litellm_provider": "github_copilot", @@ -18452,19 +18398,14 @@ }, "github_copilot/gpt-5-mini": { "litellm_provider": "github_copilot", - "max_input_tokens": 264000, + "max_input_tokens": 128000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true, - "supported_endpoints": [ - "/v1/chat/completions", - "/v1/responses" - ], - "supports_reasoning": true + "supports_vision": true }, "github_copilot/gpt-5.1": { "litellm_provider": "github_copilot", @@ -18497,7 +18438,7 @@ }, "github_copilot/gpt-5.2": { "litellm_provider": "github_copilot", - "max_input_tokens": 264000, + "max_input_tokens": 128000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", @@ -18508,27 +18449,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true, - "supports_reasoning": true - }, - "github_copilot/gpt-5.2-codex": { - "litellm_provider": "github_copilot", - "max_input_tokens": 400000, - "max_output_tokens": 128000, - "max_tokens": 128000, - "mode": "responses", - "supported_endpoints": [ - "/v1/responses" - ], - "supports_vision": true, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_response_schema": true, - "supports_reasoning": true + "supports_vision": true }, "github_copilot/gpt-5.3-codex": { "litellm_provider": "github_copilot", - "max_input_tokens": 400000, + "max_input_tokens": 128000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "responses", @@ -18538,96 +18463,25 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true, - "supports_reasoning": true - }, - "github_copilot/gpt-5.4": { - "litellm_provider": "github_copilot", - "max_input_tokens": 400000, - "max_output_tokens": 128000, - "max_tokens": 128000, - "mode": "chat", - "supported_endpoints": [ - "/v1/chat/completions", - "/v1/responses" - ], - "supports_vision": true, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_response_schema": true, - "supports_reasoning": true - }, - "github_copilot/gpt-5.4-mini": { - "litellm_provider": "github_copilot", - "max_input_tokens": 400000, - "max_output_tokens": 128000, - "max_tokens": 128000, - "mode": "responses", - "supported_endpoints": [ - "/v1/responses" - ], - "supports_vision": true, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_response_schema": true, - "supports_reasoning": true - }, - "github_copilot/gpt-5.5": { - "litellm_provider": "github_copilot", - "max_input_tokens": 400000, - "max_output_tokens": 128000, - "max_tokens": 128000, - "mode": "responses", - "supported_endpoints": [ - "/v1/responses" - ], - "supports_vision": true, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_response_schema": true, - "supports_reasoning": true - }, - "github_copilot/oswe-vscode-prime": { - "litellm_provider": "github_copilot", - "max_input_tokens": 264000, - "max_output_tokens": 64000, - "max_tokens": 64000, - "mode": "chat", - "supported_endpoints": [ - "/v1/chat/completions", - "/v1/responses" - ], - "supports_vision": true, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_response_schema": true + "supports_vision": true }, "github_copilot/text-embedding-3-small": { "litellm_provider": "github_copilot", "max_input_tokens": 8191, "max_tokens": 8191, - "mode": "embedding", - "supported_endpoints": [ - "/v1/embeddings" - ] + "mode": "embedding" }, "github_copilot/text-embedding-3-small-inference": { "litellm_provider": "github_copilot", "max_input_tokens": 8191, "max_tokens": 8191, - "mode": "embedding", - "supported_endpoints": [ - "/v1/embeddings" - ] + "mode": "embedding" }, "github_copilot/text-embedding-ada-002": { "litellm_provider": "github_copilot", "max_input_tokens": 8191, "max_tokens": 8191, - "mode": "embedding", - "supported_endpoints": [ - "/v1/embeddings" - ] + "mode": "embedding" }, "chatgpt/gpt-5.4": { "litellm_provider": "chatgpt", @@ -23278,11 +23132,13 @@ }, "jp.anthropic.claude-sonnet-4-5-20250929-v1:0": { "cache_creation_input_token_cost": 4.125e-06, + "cache_creation_input_token_cost_above_1hr": 6.6e-06, "cache_read_input_token_cost": 3.3e-07, "input_cost_per_token": 3.3e-06, "input_cost_per_token_above_200k_tokens": 6.6e-06, "output_cost_per_token_above_200k_tokens": 2.475e-05, "cache_creation_input_token_cost_above_200k_tokens": 8.25e-06, + "cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.32e-05, "cache_read_input_token_cost_above_200k_tokens": 6.6e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 200000, @@ -23308,6 +23164,7 @@ }, "jp.anthropic.claude-haiku-4-5-20251001-v1:0": { "cache_creation_input_token_cost": 1.375e-06, + "cache_creation_input_token_cost_above_1hr": 2.2e-06, "cache_read_input_token_cost": 1.1e-07, "input_cost_per_token": 1.1e-06, "litellm_provider": "bedrock_converse", @@ -23420,6 +23277,31 @@ "supports_system_messages": true, "supports_tool_choice": true }, + "inception/mercury-2": { + "cache_read_input_token_cost": 2.5e-08, + "input_cost_per_token": 2.5e-07, + "litellm_provider": "inception", + "max_input_tokens": 128000, + "max_output_tokens": 50000, + "max_tokens": 50000, + "mode": "chat", + "output_cost_per_token": 7.5e-07, + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true + }, + "text-completion-inception/mercury-edit-2": { + "cache_read_input_token_cost": 2.5e-08, + "input_cost_per_token": 2.5e-07, + "litellm_provider": "text-completion-inception", + "max_input_tokens": 32000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "completion", + "output_cost_per_token": 7.5e-07 + }, "lambda_ai/deepseek-llama3.3-70b": { "input_cost_per_token": 2e-07, "litellm_provider": "lambda_ai", @@ -26432,6 +26314,32 @@ "supports_vision": true, "supports_web_search": true }, + "oci/meta.llama-3.1-8b-instruct": { + "input_cost_per_token": 7.2e-07, + "litellm_provider": "oci", + "max_input_tokens": 128000, + "max_output_tokens": 4000, + "max_tokens": 4000, + "mode": "chat", + "output_cost_per_token": 7.2e-07, + "source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing", + "supports_function_calling": true, + "supports_response_schema": false, + "supports_native_streaming": true + }, + "oci/meta.llama-3.1-70b-instruct": { + "input_cost_per_token": 7.2e-07, + "litellm_provider": "oci", + "max_input_tokens": 128000, + "max_output_tokens": 4000, + "max_tokens": 4000, + "mode": "chat", + "output_cost_per_token": 7.2e-07, + "source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing", + "supports_function_calling": true, + "supports_response_schema": false, + "supports_native_streaming": true + }, "oci/meta.llama-3.1-405b-instruct": { "input_cost_per_token": 1.068e-05, "litellm_provider": "oci", @@ -26442,7 +26350,8 @@ "output_cost_per_token": 1.068e-05, "source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing", "supports_function_calling": true, - "supports_response_schema": false + "supports_response_schema": false, + "supports_native_streaming": true }, "oci/meta.llama-3.2-90b-vision-instruct": { "input_cost_per_token": 2e-06, @@ -26455,6 +26364,7 @@ "source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing", "supports_function_calling": true, "supports_response_schema": false, + "supports_native_streaming": true, "supports_vision": true }, "oci/meta.llama-3.3-70b-instruct": { @@ -26467,31 +26377,35 @@ "output_cost_per_token": 7.2e-07, "source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing", "supports_function_calling": true, - "supports_response_schema": false + "supports_response_schema": false, + "supports_native_streaming": true }, "oci/meta.llama-4-maverick-17b-128e-instruct-fp8": { "input_cost_per_token": 7.2e-07, "litellm_provider": "oci", - "max_input_tokens": 512000, - "max_output_tokens": 4000, - "max_tokens": 4000, + "max_input_tokens": 1048576, + "max_output_tokens": 8192, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 7.2e-07, "source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing", "supports_function_calling": true, - "supports_response_schema": false + "supports_response_schema": false, + "supports_native_streaming": true, + "supports_vision": true }, "oci/meta.llama-4-scout-17b-16e-instruct": { "input_cost_per_token": 7.2e-07, "litellm_provider": "oci", - "max_input_tokens": 192000, - "max_output_tokens": 4000, - "max_tokens": 4000, + "max_input_tokens": 10485760, + "max_output_tokens": 8192, + "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 7.2e-07, "source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing", "supports_function_calling": true, - "supports_response_schema": false + "supports_response_schema": false, + "supports_native_streaming": true }, "oci/xai.grok-3": { "input_cost_per_token": 3e-06, @@ -26503,7 +26417,8 @@ "output_cost_per_token": 1.5e-05, "source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing", "supports_function_calling": true, - "supports_response_schema": false + "supports_response_schema": false, + "supports_native_streaming": true }, "oci/xai.grok-3-fast": { "input_cost_per_token": 5e-06, @@ -26515,7 +26430,8 @@ "output_cost_per_token": 2.5e-05, "source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing", "supports_function_calling": true, - "supports_response_schema": false + "supports_response_schema": false, + "supports_native_streaming": true }, "oci/xai.grok-3-mini": { "input_cost_per_token": 3e-07, @@ -26527,7 +26443,8 @@ "output_cost_per_token": 5e-07, "source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing", "supports_function_calling": true, - "supports_response_schema": false + "supports_response_schema": false, + "supports_native_streaming": true }, "oci/xai.grok-3-mini-fast": { "input_cost_per_token": 6e-07, @@ -26539,7 +26456,8 @@ "output_cost_per_token": 4e-06, "source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing", "supports_function_calling": true, - "supports_response_schema": false + "supports_response_schema": false, + "supports_native_streaming": true }, "oci/xai.grok-4": { "input_cost_per_token": 3e-06, @@ -26551,7 +26469,8 @@ "output_cost_per_token": 1.5e-05, "source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing", "supports_function_calling": true, - "supports_response_schema": false + "supports_response_schema": false, + "supports_native_streaming": true }, "oci/cohere.command-latest": { "input_cost_per_token": 1.56e-06, @@ -26563,7 +26482,8 @@ "output_cost_per_token": 1.56e-06, "source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/", "supports_function_calling": true, - "supports_response_schema": false + "supports_response_schema": false, + "supports_native_streaming": true }, "oci/cohere.command-a-03-2025": { "input_cost_per_token": 1.56e-06, @@ -26575,7 +26495,8 @@ "output_cost_per_token": 1.56e-06, "source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/", "supports_function_calling": true, - "supports_response_schema": false + "supports_response_schema": false, + "supports_native_streaming": true }, "oci/cohere.command-plus-latest": { "input_cost_per_token": 1.56e-06, @@ -26587,7 +26508,88 @@ "output_cost_per_token": 1.56e-06, "source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/", "supports_function_calling": true, - "supports_response_schema": false + "supports_response_schema": false, + "supports_native_streaming": true + }, + "oci/google.gemini-2.5-flash": { + "input_cost_per_token": 1.5e-07, + "litellm_provider": "oci", + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_token": 6e-07, + "source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing", + "supports_function_calling": true, + "supports_response_schema": true, + "supports_vision": true, + "supports_native_streaming": true, + "supports_image_size": false + }, + "oci/google.gemini-2.5-pro": { + "input_cost_per_token": 1.25e-06, + "litellm_provider": "oci", + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_token": 1e-05, + "source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing", + "supports_function_calling": true, + "supports_response_schema": true, + "supports_vision": true, + "supports_native_streaming": true + }, + "oci/google.gemini-2.5-flash-lite": { + "input_cost_per_token": 7.5e-08, + "litellm_provider": "oci", + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_token": 3e-07, + "source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing", + "supports_function_calling": true, + "supports_response_schema": true, + "supports_vision": true, + "supports_native_streaming": true, + "supports_image_size": false + }, + "oci/cohere.command-a-vision": { + "input_cost_per_token": 1.56e-06, + "litellm_provider": "oci", + "max_input_tokens": 256000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 1.56e-06, + "source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/", + "supports_function_calling": true, + "supports_response_schema": false, + "supports_native_streaming": true, + "supports_vision": true + }, + "oci/cohere.command-a-reasoning": { + "input_cost_per_token": 1.56e-06, + "litellm_provider": "oci", + "max_input_tokens": 256000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 1.56e-06, + "source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/", + "supports_function_calling": false, + "supports_response_schema": false, + "supports_native_streaming": true + }, + "oci/cohere.embed-multilingual-image-v3.0": { + "input_cost_per_token": 1e-07, + "litellm_provider": "oci", + "max_input_tokens": 512, + "mode": "embedding", + "output_vector_size": 1024, + "source": "https://www.oracle.com/cloud/ai/generative-ai/pricing/", + "supports_vision": true }, "oci/cohere.command-a-reasoning-08-2025": { "input_cost_per_token": 1.56e-06, @@ -26663,18 +26665,6 @@ "supports_response_schema": false, "supports_vision": true }, - "oci/meta.llama-3.1-70b-instruct": { - "input_cost_per_token": 7.2e-07, - "litellm_provider": "oci", - "max_input_tokens": 128000, - "max_output_tokens": 4000, - "max_tokens": 4000, - "mode": "chat", - "output_cost_per_token": 7.2e-07, - "source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing", - "supports_function_calling": true, - "supports_response_schema": false - }, "oci/meta.llama-3.3-70b-instruct-fp8-dynamic": { "input_cost_per_token": 7.2e-07, "litellm_provider": "oci", @@ -26792,45 +26782,6 @@ "supports_response_schema": true, "supports_vision": true }, - "oci/google.gemini-2.5-pro": { - "input_cost_per_token": 1.25e-06, - "litellm_provider": "oci", - "max_input_tokens": 1048576, - "max_output_tokens": 65536, - "max_tokens": 65536, - "mode": "chat", - "output_cost_per_token": 1e-05, - "source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing", - "supports_function_calling": true, - "supports_response_schema": true, - "supports_vision": true - }, - "oci/google.gemini-2.5-flash": { - "input_cost_per_token": 1.5e-07, - "litellm_provider": "oci", - "max_input_tokens": 1048576, - "max_output_tokens": 65536, - "max_tokens": 65536, - "mode": "chat", - "output_cost_per_token": 6e-07, - "source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing", - "supports_function_calling": true, - "supports_response_schema": true, - "supports_vision": true - }, - "oci/google.gemini-2.5-flash-lite": { - "input_cost_per_token": 7.5e-08, - "litellm_provider": "oci", - "max_input_tokens": 1048576, - "max_output_tokens": 65536, - "max_tokens": 65536, - "mode": "chat", - "output_cost_per_token": 3e-07, - "source": "https://www.oracle.com/artificial-intelligence/generative-ai/generative-ai-service/pricing", - "supports_function_calling": true, - "supports_response_schema": true, - "supports_vision": true - }, "oci/cohere.embed-english-v3.0": { "input_cost_per_token": 1e-07, "litellm_provider": "oci", @@ -27638,7 +27589,8 @@ "supports_response_schema": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_image_size": false }, "openrouter/google/gemini-2.5-pro": { "input_cost_per_audio_token": 7e-07, @@ -29508,7 +29460,8 @@ "mode": "responses", "supports_web_search": true, "supports_reasoning": false, - "supports_function_calling": true + "supports_function_calling": true, + "supports_image_size": false }, "perplexity/xai/grok-4-1-fast-non-reasoning": { "litellm_provider": "perplexity", @@ -30090,7 +30043,8 @@ "supports_vision": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_response_schema": true + "supports_response_schema": true, + "supports_image_size": false }, "replicate/openai/gpt-oss-120b": { "input_cost_per_token": 1.8e-07, @@ -31909,6 +31863,7 @@ }, "au.anthropic.claude-haiku-4-5-20251001-v1:0": { "cache_creation_input_token_cost": 1.375e-06, + "cache_creation_input_token_cost_above_1hr": 2.2e-06, "cache_read_input_token_cost": 1.1e-07, "input_cost_per_token": 1.1e-06, "litellm_provider": "bedrock_converse", @@ -32788,7 +32743,8 @@ "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, - "supports_response_schema": true + "supports_response_schema": true, + "supports_image_size": false }, "vercel_ai_gateway/google/gemini-2.5-pro": { "input_cost_per_token": 2.5e-06, @@ -33555,6 +33511,7 @@ }, "vertex_ai/claude-haiku-4-5": { "cache_creation_input_token_cost": 1.25e-06, + "cache_creation_input_token_cost_above_1hr": 2e-06, "cache_read_input_token_cost": 1e-07, "input_cost_per_token": 1e-06, "litellm_provider": "vertex_ai-anthropic_models", @@ -33576,6 +33533,7 @@ }, "vertex_ai/claude-haiku-4-5@20251001": { "cache_creation_input_token_cost": 1.25e-06, + "cache_creation_input_token_cost_above_1hr": 2e-06, "cache_read_input_token_cost": 1e-07, "input_cost_per_token": 1e-06, "litellm_provider": "vertex_ai-anthropic_models", @@ -33626,6 +33584,7 @@ }, "vertex_ai/claude-3-7-sonnet@20250219": { "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_1hr": 6e-06, "cache_read_input_token_cost": 3e-07, "deprecation_date": "2026-05-11", "input_cost_per_token": 3e-06, @@ -33725,6 +33684,7 @@ }, "vertex_ai/claude-opus-4": { "cache_creation_input_token_cost": 1.875e-05, + "cache_creation_input_token_cost_above_1hr": 3e-05, "cache_read_input_token_cost": 1.5e-06, "input_cost_per_token": 1.5e-05, "litellm_provider": "vertex_ai-anthropic_models", @@ -33750,6 +33710,7 @@ }, "vertex_ai/claude-opus-4-1": { "cache_creation_input_token_cost": 1.875e-05, + "cache_creation_input_token_cost_above_1hr": 3e-05, "cache_read_input_token_cost": 1.5e-06, "input_cost_per_token": 1.5e-05, "input_cost_per_token_batches": 7.5e-06, @@ -33767,6 +33728,7 @@ }, "vertex_ai/claude-opus-4-1@20250805": { "cache_creation_input_token_cost": 1.875e-05, + "cache_creation_input_token_cost_above_1hr": 3e-05, "cache_read_input_token_cost": 1.5e-06, "input_cost_per_token": 1.5e-05, "input_cost_per_token_batches": 7.5e-06, @@ -33784,6 +33746,7 @@ }, "vertex_ai/claude-opus-4-5": { "cache_creation_input_token_cost": 6.25e-06, + "cache_creation_input_token_cost_above_1hr": 1e-05, "cache_read_input_token_cost": 5e-07, "input_cost_per_token": 5e-06, "litellm_provider": "vertex_ai-anthropic_models", @@ -33810,6 +33773,7 @@ }, "vertex_ai/claude-opus-4-5@20251101": { "cache_creation_input_token_cost": 6.25e-06, + "cache_creation_input_token_cost_above_1hr": 1e-05, "cache_read_input_token_cost": 5e-07, "input_cost_per_token": 5e-06, "litellm_provider": "vertex_ai-anthropic_models", @@ -33837,6 +33801,7 @@ }, "vertex_ai/claude-opus-4-6": { "cache_creation_input_token_cost": 6.25e-06, + "cache_creation_input_token_cost_above_1hr": 1e-05, "cache_read_input_token_cost": 5e-07, "input_cost_per_token": 5e-06, "litellm_provider": "vertex_ai-anthropic_models", @@ -33864,6 +33829,7 @@ }, "vertex_ai/claude-opus-4-6@default": { "cache_creation_input_token_cost": 6.25e-06, + "cache_creation_input_token_cost_above_1hr": 1e-05, "cache_read_input_token_cost": 5e-07, "input_cost_per_token": 5e-06, "litellm_provider": "vertex_ai-anthropic_models", @@ -33891,6 +33857,7 @@ }, "vertex_ai/claude-opus-4-7": { "cache_creation_input_token_cost": 6.25e-06, + "cache_creation_input_token_cost_above_1hr": 1e-05, "cache_read_input_token_cost": 5e-07, "input_cost_per_token": 5e-06, "litellm_provider": "vertex_ai-anthropic_models", @@ -33918,6 +33885,7 @@ }, "vertex_ai/claude-opus-4-7@default": { "cache_creation_input_token_cost": 6.25e-06, + "cache_creation_input_token_cost_above_1hr": 1e-05, "cache_read_input_token_cost": 5e-07, "input_cost_per_token": 5e-06, "litellm_provider": "vertex_ai-anthropic_models", @@ -34001,6 +33969,7 @@ }, "vertex_ai/claude-sonnet-4-5": { "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_1hr": 6e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_token": 3e-06, "input_cost_per_token_above_200k_tokens": 6e-06, @@ -34027,6 +33996,7 @@ }, "vertex_ai/claude-sonnet-4-6": { "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_1hr": 6e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_token": 3e-06, "litellm_provider": "vertex_ai-anthropic_models", @@ -34054,6 +34024,7 @@ }, "vertex_ai/claude-sonnet-4-5@20250929": { "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_1hr": 6e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_token": 3e-06, "input_cost_per_token_above_200k_tokens": 6e-06, @@ -34081,6 +34052,7 @@ }, "vertex_ai/claude-opus-4@20250514": { "cache_creation_input_token_cost": 1.875e-05, + "cache_creation_input_token_cost_above_1hr": 3e-05, "cache_read_input_token_cost": 1.5e-06, "input_cost_per_token": 1.5e-05, "litellm_provider": "vertex_ai-anthropic_models", @@ -34106,6 +34078,7 @@ }, "vertex_ai/claude-sonnet-4": { "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_1hr": 6e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_token": 3e-06, "input_cost_per_token_above_200k_tokens": 6e-06, @@ -34135,6 +34108,7 @@ }, "vertex_ai/claude-sonnet-4@20250514": { "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_1hr": 6e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_token": 3e-06, "input_cost_per_token_above_200k_tokens": 6e-06, @@ -34342,7 +34316,8 @@ "supports_url_context": true, "supports_vision": true, "supports_web_search": false, - "tpm": 8000000 + "tpm": 8000000, + "supports_image_size": false }, "vertex_ai/gemini-3-pro-image-preview": { "input_cost_per_image": 0.0011, @@ -34430,10 +34405,16 @@ "web_search_billing_unit": "per_query" }, "vertex_ai/gemini-3.1-flash-lite": { - "cache_read_input_token_cost": 4.5e-08, - "cache_read_input_token_cost_per_audio_token": 9e-08, - "input_cost_per_audio_token": 9e-07, - "input_cost_per_token": 4.5e-07, + "cache_read_input_token_cost": 2.5e-08, + "cache_read_input_token_cost_batches": 1.25e-08, + "cache_read_input_token_cost_flex": 1.25e-08, + "cache_read_input_token_cost_per_audio_token": 5e-08, + "cache_read_input_token_cost_priority": 4.5e-08, + "input_cost_per_audio_token": 5e-07, + "input_cost_per_token": 2.5e-07, + "input_cost_per_token_batches": 1.25e-07, + "input_cost_per_token_flex": 1.25e-07, + "input_cost_per_token_priority": 4.5e-07, "litellm_provider": "vertex_ai-language-models", "max_audio_length_hours": 8.4, "max_audio_per_prompt": 1, @@ -34445,8 +34426,11 @@ "max_video_length": 1, "max_videos_per_prompt": 10, "mode": "chat", - "output_cost_per_reasoning_token": 2.7e-06, - "output_cost_per_token": 2.7e-06, + "output_cost_per_reasoning_token": 1.5e-06, + "output_cost_per_token": 1.5e-06, + "output_cost_per_token_batches": 7.5e-07, + "output_cost_per_token_flex": 7.5e-07, + "output_cost_per_token_priority": 2.7e-06, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models", "supported_endpoints": [ "/v1/chat/completions", @@ -41151,6 +41135,7 @@ }, "vertex_ai/claude-sonnet-4-6@default": { "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_1hr": 6e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_token": 3e-06, "litellm_provider": "vertex_ai-anthropic_models", diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 2aacab80f57..863e6acd41e 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -7,6 +7,7 @@ from starlette.requests import Request from starlette.types import Scope from litellm._logging import verbose_logger +from litellm.constants import DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL from litellm.proxy._types import ( LiteLLM_TeamTable, ProxyException, @@ -14,6 +15,88 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.auth.ip_address_utils import IPAddressUtils + + +def _parse_mcp_server_names_from_path( + path: str, mcp_servers_header: Optional[List[str]] = None +) -> Optional[List[str]]: + """Resolve the single MCP server name a cold-start passthrough bypass may + target. Delegates parsing to + :meth:`MCPRequestHandler._extract_target_server_names_from_path` so the + names used here always match the names downstream routing uses; returns + ``None`` whenever the bypass must not activate (aggregate ``/mcp``, + multi-server CSV paths, or any other unrecognized path). + + Also fails closed when the ``x-mcp-servers`` header introduces any server + not present in the path-derived target set. Downstream routing for + ``/mcp/...`` paths overrides the header with path-derived names, but a + header/path mismatch here is a sign of a confused or hostile caller — + refuse the cold-start bypass rather than admit anonymously based on the + path while the header advertises a stricter, non-passthrough target.""" + servers = MCPRequestHandler._extract_target_server_names_from_path(path) + if len(servers) != 1: + verbose_logger.debug( + "MCP cold-start: path %r resolved to %r; passthrough 401 bypass " + "requires exactly one target and will not activate", + path, + servers, + ) + return None + if mcp_servers_header is not None and (set(mcp_servers_header) - set(servers)): + verbose_logger.debug( + "MCP cold-start: x-mcp-servers header %r introduces target(s) not " + "in path-derived set %r; passthrough 401 bypass will not activate", + mcp_servers_header, + servers, + ) + return None + return servers + + +def _is_mcp_passthrough_cold_start( + mcp_servers: Optional[List[str]], client_ip: Optional[str] +) -> bool: + """True only when EVERY targeted server is a pass-through server with no + auth headers — the cold-start OAuth discovery case per RFC 9728 / MCP + Authorization spec. Lets the route handler's 401 emitter produce the + spec-compliant WWW-Authenticate challenge instead of surfacing a generic + admission error. + + Uses "all" semantics (mirrors :meth:`MCPRequestHandler._target_servers_use_oauth2`): + one non-passthrough target in a co-targeted set must not flip the bypass + open for the others. Fails closed when any target cannot be resolved.""" + if not mcp_servers: + return False + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + for name in mcp_servers: + server = global_mcp_server_manager.get_mcp_server_by_name( + name, client_ip=client_ip + ) + if server is None or not getattr(server, "is_oauth_passthrough", False): + return False + return True + + +def _is_litellm_auth_admission_error(exc: Exception) -> bool: + if isinstance(exc, HTTPException): + return exc.status_code == 401 + if isinstance(exc, ProxyException): + try: + return int(exc.code) == 401 + except (TypeError, ValueError): + return False + return False + + +def _has_client_supplied_mcp_auth( + mcp_auth_header: Optional[str], + mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]], +) -> bool: + return bool(mcp_auth_header) or bool(mcp_server_auth_headers) class MCPRequestHandler: @@ -37,7 +120,7 @@ class MCPRequestHandler: LITELLM_MCP_ACCESS_GROUPS_HEADER_NAME = SpecialHeaders.mcp_access_groups.value @staticmethod - async def process_mcp_request( + async def process_mcp_request( # noqa: PLR0915 scope: Scope, ) -> Tuple[ UserAPIKeyAuth, @@ -130,7 +213,9 @@ class MCPRequestHandler: elif ( not litellm_api_key and MCPRequestHandler._target_servers_delegate_auth_to_upstream( # noqa: E501 - path=request_route, mcp_servers=mcp_servers + path=request_route, + mcp_servers=mcp_servers, + client_ip=IPAddressUtils.get_mcp_client_ip(request), ) ): # Operator opted this oauth2 server into upstream-delegated auth @@ -172,25 +257,87 @@ class MCPRequestHandler: # than coercing (``int("None")`` would raise ValueError and # rewrite the auth error as a 500). status = e.status_code if isinstance(e, HTTPException) else e.code - if status in ( - 401, - 403, - "401", - "403", - ) and MCPRequestHandler._target_servers_use_oauth2( - path=request_route, mcp_servers=mcp_servers + is_auth_error = status in (401, 403, "401", "403") + is_unauthenticated = status in (401, "401") + client_ip = IPAddressUtils.get_mcp_client_ip(request) + if is_auth_error and MCPRequestHandler._target_servers_use_oauth2( + path=request_route, + mcp_servers=mcp_servers, + client_ip=client_ip, ): verbose_logger.debug( "MCP OAuth2: target server is OAuth2-mode, treating " "Authorization as upstream OAuth2 token passthrough" ) validated_user_api_key_auth = UserAPIKeyAuth() + elif is_unauthenticated: + # Pass-through cold-start return: per RFC 9728 / MCP + # Authorization spec the client completes upstream OAuth + # discovery and returns with ``Authorization: Bearer + # ``. For ``auth_type=none`` passthrough + # servers that bearer is not a LiteLLM key (auth above + # failed) but is meant to be forwarded upstream + # unchanged. Fall back to anonymous admission so the + # caller is not rejected for following the discovery + # flow without also setting ``x-litellm-api-key``. + # Only trigger on 401 (token unrecognized); a 403 means + # the key WAS recognized but is forbidden (e.g. over + # budget / rate limited) and must propagate so those + # controls are not bypassed via anonymous admission. + mcp_servers_from_path = _parse_mcp_server_names_from_path( + request_route, mcp_servers + ) + if ( + mcp_servers_from_path is not None + and not _has_client_supplied_mcp_auth( + mcp_auth_header, + mcp_server_auth_headers, + ) + and _is_mcp_passthrough_cold_start( + mcp_servers_from_path, client_ip=client_ip + ) + ): + verbose_logger.debug( + "MCP pass-through return: target server is " + "passthrough, treating Authorization as " + "upstream OAuth token for delegated auth" + ) + validated_user_api_key_auth = UserAPIKeyAuth() + else: + raise else: raise else: - validated_user_api_key_auth = await user_api_key_auth( - api_key=litellm_api_key, request=request - ) + try: + validated_user_api_key_auth = await user_api_key_auth( + api_key=litellm_api_key, request=request + ) + except (HTTPException, ProxyException) as exc: + # Cold-start MCP OAuth discovery: RFC 9728 / MCP Authorization spec + # require unauthenticated requests to protected resources to receive + # 401 + WWW-Authenticate. Defer to _raise_preemptive_401_for_unauthenticated_servers + # for pass-through servers instead of surfacing a generic admission error. + mcp_servers_from_path = _parse_mcp_server_names_from_path( + request_route, mcp_servers + ) + client_ip = IPAddressUtils.get_mcp_client_ip(request) + if ( + mcp_servers_from_path is not None + and not _has_client_supplied_mcp_auth( + mcp_auth_header, + mcp_server_auth_headers, + ) + and _is_litellm_auth_admission_error(exc) + and _is_mcp_passthrough_cold_start( + mcp_servers_from_path, client_ip=client_ip + ) + ): + verbose_logger.debug( + "MCP pass-through cold start: deferring admission to route 401 emitter" + ) + validated_user_api_key_auth = UserAPIKeyAuth() + else: + raise return ( validated_user_api_key_auth, @@ -262,7 +409,9 @@ class MCPRequestHandler: return [servers_and_path] @staticmethod - def _target_servers_use_oauth2(path: str, mcp_servers: Optional[List[str]]) -> bool: + def _target_servers_use_oauth2( + path: str, mcp_servers: Optional[List[str]], client_ip: Optional[str] + ) -> bool: """ True only when EVERY MCP server the request targets is configured for ``auth_type == oauth2``. If any target is non-OAuth2 — or if the target @@ -291,14 +440,16 @@ class MCPRequestHandler: return False for name in target_names: - server = global_mcp_server_manager.get_mcp_server_by_name(name) + server = global_mcp_server_manager.get_mcp_server_by_name( + name, client_ip=client_ip + ) if server is None or server.auth_type != MCPAuth.oauth2: return False return True @staticmethod def _target_servers_delegate_auth_to_upstream( - path: str, mcp_servers: Optional[List[str]] + path: str, mcp_servers: Optional[List[str]], client_ip: Optional[str] ) -> bool: """ True only when EVERY MCP server the request targets is configured for @@ -328,7 +479,9 @@ class MCPRequestHandler: return False for name in target_names: - server = global_mcp_server_manager.get_mcp_server_by_name(name) + server = global_mcp_server_manager.get_mcp_server_by_name( + name, client_ip=client_ip + ) if server is None or server.auth_type != MCPAuth.oauth2: return False # `is True` is intentional: opt-in must be an explicit boolean @@ -1090,22 +1243,21 @@ class MCPRequestHandler: ) return [] - # Sentinel stored in cache when an org has no object_permission, so we - # don't re-query the DB on every MCP request for that org. - _ORG_NO_PERMISSION_SENTINEL = "__org_no_mcp_permission__" - @staticmethod async def _get_org_object_permission( user_api_key_auth: Optional[UserAPIKeyAuth] = None, ): """ - Get org object_permission, using user_api_key_cache to avoid DB hits on every request. - - Caches both positive results and the absence of an object_permission so that orgs - with no MCP permissions configured (the common default) do not trigger a DB query - on every request. + Get org object_permission via the established ``get_org_object`` / + ``get_object_permission`` helpers so MCP requests share the same + ``user_api_key_cache`` entries as the rest of the proxy. """ - from litellm.proxy.proxy_server import prisma_client, user_api_key_cache + from litellm.proxy.auth.auth_checks import get_object_permission, get_org_object + from litellm.proxy.proxy_server import ( + prisma_client, + proxy_logging_obj, + user_api_key_cache, + ) if not user_api_key_auth or not user_api_key_auth.org_id: return None @@ -1114,45 +1266,25 @@ class MCPRequestHandler: verbose_logger.debug("prisma_client is None") return None - org_id = user_api_key_auth.org_id - cache_key = f"org_object_permission:{org_id}" - - from litellm.proxy._types import LiteLLM_ObjectPermissionTable - try: - cached = await user_api_key_cache.async_get_cache(key=cache_key) - if cached is not None: - # Sentinel means the DB confirmed no object_permission for this org - if cached == MCPRequestHandler._ORG_NO_PERMISSION_SENTINEL: - return None - # Redis deserialises to a plain dict; reconstruct the Pydantic model - # so callers can access .mcp_servers / .mcp_tool_permissions as attrs. - if isinstance(cached, dict): - return LiteLLM_ObjectPermissionTable(**cached) - return cached - - org_row = await prisma_client.db.litellm_organizationtable.find_unique( - where={"organization_id": org_id}, - include={"object_permission": True}, + org_obj = await get_org_object( + org_id=user_api_key_auth.org_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=user_api_key_auth.parent_otel_span, + proxy_logging_obj=proxy_logging_obj, ) - if org_row is None or org_row.object_permission is None: - # Cache the negative result so subsequent calls skip the DB - await user_api_key_cache.async_set_cache( - key=cache_key, - value=MCPRequestHandler._ORG_NO_PERMISSION_SENTINEL, - ) + if org_obj is None or not org_obj.object_permission_id: return None - # Convert raw Prisma model → Pydantic before caching. Caching the - # Pydantic .dict() ensures the value survives a Redis JSON round-trip - # as a plain dict that we can reconstruct above (same pattern used by - # get_end_user_object / get_team_object in auth_checks.py). - obj_perm = LiteLLM_ObjectPermissionTable(**org_row.object_permission.dict()) - await user_api_key_cache.async_set_cache( - key=cache_key, value=obj_perm.dict() + return await get_object_permission( + object_permission_id=org_obj.object_permission_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=user_api_key_auth.parent_otel_span, + proxy_logging_obj=proxy_logging_obj, ) - return obj_perm except Exception as e: verbose_logger.warning(f"Failed to get org object permission: {str(e)}") return None @@ -1273,16 +1405,26 @@ class MCPRequestHandler: ) return [] + # Sentinel stored in cache when an agent has no object_permission, so we + # don't re-query the DB on every MCP request for that agent. + _AGENT_NO_PERMISSION_SENTINEL = "__agent_no_mcp_permission__" + @staticmethod async def _get_agent_object_permission( user_api_key_auth: Optional[UserAPIKeyAuth] = None, ): """ - Fetch the agent's object_permission from the DB (single query). - - Returns the object_permission object or None. + Get agent object_permission via the established ``get_object_permission`` + helper. Caches the ``agent_id -> object_permission_id`` mapping so we + avoid re-reading the agent row on every request, and reuses the shared + ``object_permission_id`` cache populated by the org / team / key paths. """ - from litellm.proxy.proxy_server import prisma_client + from litellm.proxy.auth.auth_checks import get_object_permission + from litellm.proxy.proxy_server import ( + prisma_client, + proxy_logging_obj, + user_api_key_cache, + ) if not user_api_key_auth or not user_api_key_auth.agent_id: return None @@ -1291,15 +1433,42 @@ class MCPRequestHandler: verbose_logger.debug("prisma_client is None") return None + agent_id = user_api_key_auth.agent_id + cache_key = f"agent_object_permission_id:{agent_id}" + try: - agent_row = await prisma_client.db.litellm_agentstable.find_unique( - where={"agent_id": user_api_key_auth.agent_id}, - include={"object_permission": True}, + object_permission_id: Optional[str] = ( + await user_api_key_cache.async_get_cache(key=cache_key) ) - if agent_row is None or agent_row.object_permission is None: + + if object_permission_id == MCPRequestHandler._AGENT_NO_PERMISSION_SENTINEL: return None - return agent_row.object_permission + if object_permission_id is None: + agent_row = await prisma_client.db.litellm_agentstable.find_unique( + where={"agent_id": agent_id}, + ) + object_permission_id = ( + getattr(agent_row, "object_permission_id", None) + if agent_row is not None + else None + ) + await user_api_key_cache.async_set_cache( + key=cache_key, + value=object_permission_id + or MCPRequestHandler._AGENT_NO_PERMISSION_SENTINEL, + ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, + ) + if not object_permission_id: + return None + + return await get_object_permission( + object_permission_id=object_permission_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=user_api_key_auth.parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + ) except Exception as e: verbose_logger.warning(f"Failed to get agent object permission: {str(e)}") return None diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 8324ba641a4..ed374635fea 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -1,8 +1,11 @@ +import asyncio import html as _html import json -from typing import Any, Dict, Optional +import time +from typing import Any, Dict, Optional, Tuple from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse +import httpx from fastapi import APIRouter, Form, HTTPException, Request from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse @@ -26,11 +29,54 @@ from litellm.proxy.utils import get_server_root_path from litellm.types.mcp import MCPAuth from litellm.types.mcp_server.mcp_server_manager import MCPServer +# TTL cache for upstream OAuth metadata fetched from pass-through MCP servers. +# Keeps us from hammering the upstream IdP on each discovery request. +# Keyed by (server_id, resource_url) → (expires_at_epoch, payload). +# A payload of ``None`` is a negative-result entry that prevents repeated +# upstream fetches when the IdP consistently has no metadata to serve. +_OAUTH_METADATA_CACHE: Dict[Tuple[str, str], Tuple[float, Optional[dict]]] = {} +_OAUTH_METADATA_CACHE_TTL_SECONDS = 300 +_OAUTH_METADATA_NEGATIVE_CACHE_TTL_SECONDS = 60 +_OAUTH_METADATA_CACHE_MAX_SIZE = 128 +# Per-(server_id, resource_url) async locks so concurrent discovery requests +# coalesce onto a single upstream fetch instead of issuing N parallel calls. +_OAUTH_METADATA_FETCH_LOCKS: Dict[Tuple[str, str], asyncio.Lock] = {} + router = APIRouter( tags=["mcp"], ) +def _prune_oauth_metadata_cache(now: Optional[float] = None) -> None: + now = now if now is not None else time.time() + expired_cache_keys = [ + cache_key + for cache_key, (expires_at, _payload) in _OAUTH_METADATA_CACHE.items() + if expires_at <= now + ] + for cache_key in expired_cache_keys: + _OAUTH_METADATA_CACHE.pop(cache_key, None) + + if len(_OAUTH_METADATA_CACHE) > _OAUTH_METADATA_CACHE_MAX_SIZE: + overflow = len(_OAUTH_METADATA_CACHE) - _OAUTH_METADATA_CACHE_MAX_SIZE + cache_keys_by_expiry = sorted( + _OAUTH_METADATA_CACHE, + key=lambda cache_key: _OAUTH_METADATA_CACHE[cache_key][0], + ) + for cache_key in cache_keys_by_expiry[:overflow]: + _OAUTH_METADATA_CACHE.pop(cache_key, None) + + # Drop locks whose cache entry has been evicted and that aren't currently + # held; held locks stay so in-flight callers continue to coalesce. + for cache_key in list(_OAUTH_METADATA_FETCH_LOCKS): + if cache_key in _OAUTH_METADATA_CACHE: + continue + lock = _OAUTH_METADATA_FETCH_LOCKS.get(cache_key) + if lock is None or lock.locked(): + continue + _OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None) + + def encode_state_with_base_url( base_url: str, original_state: str, @@ -125,6 +171,17 @@ def _resolve_oauth2_server_for_root_endpoints( return None +def _normalize_for_token_comparison(value: Any) -> str: + """Stringify ``value`` for token-rule comparison. + + Booleans are lower-cased so Python's ``True`` / ``False`` line up with + JSON-style ``"true"`` / ``"false"`` rules from admin config. + """ + if isinstance(value, bool): + return "true" if value else "false" + return str(value) + + def _validate_token_response( token_response: Dict[str, Any], validation_rules: Dict[str, Any], @@ -136,7 +193,9 @@ def _validate_token_response( ``token_response["team"]["enterprise_id"]``). Top-level keys are tried first, then dot-split traversal. All comparisons are string-coerced so that numeric values in the response (e.g. ``"org_id": 12345``) match string rules - (``"org_id": "12345"``). + (``"org_id": "12345"``). Booleans are normalised to JSON-style ``"true"`` / + ``"false"`` so admin rules written as ``{"verified": "true"}`` match upstream + responses of ``{"verified": true}``. """ for key, expected in validation_rules.items(): actual: Any = token_response.get(key) @@ -163,7 +222,9 @@ def _validate_token_response( ), }, ) - if str(actual) != str(expected): + if _normalize_for_token_comparison(actual) != _normalize_for_token_comparison( + expected + ): raise HTTPException( status_code=403, detail={ @@ -400,6 +461,11 @@ async def exchange_token_with_server( headers={"Accept": "application/json"}, data=token_data, ) + if response is None: + raise HTTPException( + status_code=502, + detail="MCP upstream token endpoint returned no response", + ) response.raise_for_status() token_response = response.json() @@ -505,6 +571,11 @@ async def register_client_with_server( headers=headers, json=register_data, ) + if response is None: + raise HTTPException( + status_code=502, + detail="MCP upstream registration endpoint returned no response", + ) response.raise_for_status() token_response = response.json() @@ -766,7 +837,119 @@ async def callback( """ -def _build_oauth_protected_resource_response( +async def fetch_upstream_oauth_protected_resource( + mcp_server: MCPServer, +) -> Optional[dict]: + """Fetch the upstream MCP server's ``.well-known/oauth-protected-resource`` + metadata for a pass-through server. + + Tries host-only first, then falls back to the RFC 9728 §3.1 path-suffix + form (e.g. ``https://host/.well-known/oauth-protected-resource/mcp``) to + cover upstreams that scope metadata per resource path. + + Responses are cached in-process for ~5 minutes keyed on + ``(server_id, resource_url)`` so we do not hammer the IdP. + + Returns the parsed JSON dict on success, or ``None`` if neither form + responds with a 2xx JSON payload. Raises on network/connection errors so + the caller can emit HTTP 502 rather than fabricate a gateway response. + """ + if not mcp_server.url: + return None + + upstream = urlparse(mcp_server.url) + if not upstream.scheme or not upstream.netloc: + return None + + cache_key = (mcp_server.server_id, mcp_server.url) + now = time.time() + _prune_oauth_metadata_cache(now) + cached = _OAUTH_METADATA_CACHE.get(cache_key) + if cached is not None and cached[0] > now: + return cached[1] + + lock = _OAUTH_METADATA_FETCH_LOCKS.setdefault(cache_key, asyncio.Lock()) + async with lock: + now = time.time() + cached = _OAUTH_METADATA_CACHE.get(cache_key) + if cached is not None and cached[0] > now: + return cached[1] + + host_base = f"{upstream.scheme}://{upstream.netloc}" + candidates = [f"{host_base}/.well-known/oauth-protected-resource"] + # RFC 9728 §3.1 path fallback + if upstream.path and upstream.path not in ("", "/"): + candidates.append( + f"{host_base}/.well-known/oauth-protected-resource" + f"{upstream.path.rstrip('/')}" + ) + + async_client = get_async_httpx_client( + llm_provider=httpxSpecialProvider.Oauth2Check + ) + + network_errors: list[Exception] = [] + for candidate in candidates: + try: + response = await async_client.get( + candidate, + headers={"Accept": "application/json"}, + ) + except Exception as exc: + if is_network_error(exc): + network_errors.append(exc) + else: + verbose_logger.warning( + "MCP OAuth metadata fetch for %s raised non-transport " + "%s: %s — treating as no metadata for this candidate", + candidate, + type(exc).__name__, + exc, + ) + continue + if response.status_code == 200: + try: + payload = response.json() + except Exception as exc: + verbose_logger.warning( + "MCP OAuth metadata at %s returned 200 but JSON " + "decode failed (%s: %s) — treating as no metadata", + candidate, + type(exc).__name__, + exc, + ) + continue + if isinstance(payload, dict): + now = time.time() + _OAUTH_METADATA_CACHE[cache_key] = ( + now + _OAUTH_METADATA_CACHE_TTL_SECONDS, + payload, + ) + _prune_oauth_metadata_cache(now) + return payload + + if len(network_errors) == len(candidates): + raise network_errors[-1] + + # Negative-result caching: when no candidate yielded a usable payload, + # remember that for a shorter TTL so we don't re-fetch on every + # subsequent discovery request (and so the per-key lock can be pruned). + now = time.time() + _OAUTH_METADATA_CACHE[cache_key] = ( + now + _OAUTH_METADATA_NEGATIVE_CACHE_TTL_SECONDS, + None, + ) + _prune_oauth_metadata_cache(now) + return None + + +def is_network_error(exc: Exception) -> bool: + """True for transport-layer failures (connection refused, DNS, TLS, timeout) + as opposed to HTTP protocol errors (4xx/5xx with a valid response).""" + return isinstance(exc, httpx.TransportError) + + +async def _build_oauth_protected_resource_response( request: Request, mcp_server_name: Optional[str], use_standard_pattern: bool, @@ -774,6 +957,12 @@ def _build_oauth_protected_resource_response( """ Build OAuth protected resource response with the appropriate URL pattern. + For pass-through MCP servers (``MCPServer.is_oauth_passthrough``), the + gateway proxies the upstream's own ``oauth-protected-resource`` metadata + so that standards-compliant MCP clients discover the **upstream** IdP + instead of the gateway. The ``resource`` field is rewritten to the + gateway's own URL so clients present the bearer token back to the gateway. + Args: request: FastAPI Request object mcp_server_name: Name of the MCP server @@ -813,6 +1002,46 @@ def _build_oauth_protected_resource_response( else: resource_url = f"{request_base_url}/mcp" + # Pass-through branch: proxy the upstream's own metadata so discovery + # directs the client at the real IdP (Okta, Keycloak, …) instead of us. + if mcp_server is not None and mcp_server.is_oauth_passthrough: + try: + upstream_metadata = await fetch_upstream_oauth_protected_resource( + mcp_server + ) + except Exception as exc: + verbose_logger.warning( + "Failed to fetch upstream oauth-protected-resource metadata " + f"for pass-through MCP server {mcp_server.name!r}: {exc}" + ) + raise HTTPException( + status_code=502, + detail=( + "Failed to fetch upstream oauth-protected-resource " + f"metadata for MCP server {mcp_server.name!r}" + ), + ) + + if upstream_metadata is not None: + response = {**upstream_metadata, "resource": resource_url} + return response + + # Upstream responded but with non-200 or non-dict payload. For + # pass-through servers the gateway is NOT the authorization server, + # so we must not fall through to the default gateway metadata — + # that would point clients at the wrong IdP. + verbose_logger.warning( + "Upstream oauth-protected-resource metadata unavailable for " + f"pass-through MCP server {mcp_server.name!r}" + ) + raise HTTPException( + status_code=502, + detail=( + "Upstream oauth-protected-resource metadata unavailable " + f"for MCP server {mcp_server.name!r}" + ), + ) + return { "authorization_servers": [ ( @@ -843,7 +1072,7 @@ async def oauth_protected_resource_mcp_standard(request: Request, mcp_server_nam This endpoint is compliant with MCP specification and works with standard MCP clients like mcp-inspector and VSCode Copilot. """ - return _build_oauth_protected_resource_response( + return await _build_oauth_protected_resource_response( request=request, mcp_server_name=mcp_server_name, use_standard_pattern=True, @@ -868,36 +1097,22 @@ async def oauth_protected_resource_mcp( This endpoint is kept for backward compatibility. New integrations should use the standard MCP pattern (/mcp/{server_name}) instead. """ - return _build_oauth_protected_resource_response( + return await _build_oauth_protected_resource_response( request=request, mcp_server_name=mcp_server_name, use_standard_pattern=False, ) -""" - https://datatracker.ietf.org/doc/html/rfc8414#section-3.1 - RFC 8414: Path-aware OAuth discovery - If the issuer identifier value contains a path component, any - terminating "/" MUST be removed before inserting "/.well-known/" and - the well-known URI suffix between the host component and the path(include root path) - component. -""" - - def _build_oauth_authorization_server_response( request: Request, mcp_server_name: Optional[str], ) -> dict: - """ - Build OAuth authorization server metadata response. + """Build OAuth authorization server metadata response (gateway-as-AS shape). - Args: - request: FastAPI Request object - mcp_server_name: Name of the MCP server - - Returns: - OAuth authorization server metadata dict + Synchronous because the body only does dict construction and synchronous + registry lookups; unlike :func:`_build_oauth_protected_resource_response` + it does not need to await any upstream IO. """ from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( global_mcp_server_manager, diff --git a/litellm/proxy/_experimental/mcp_server/exceptions.py b/litellm/proxy/_experimental/mcp_server/exceptions.py new file mode 100644 index 00000000000..fd8fc3d5e58 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/exceptions.py @@ -0,0 +1,80 @@ +"""Exceptions raised by the LiteLLM MCP proxy.""" + +from typing import Optional + +from fastapi import HTTPException + + +class MCPUpstreamAuthError(Exception): + """Raised when an upstream MCP server returns an authentication failure + (typically HTTP 401) and the gateway should surface it transparently to + the client instead of swallowing it. + + Only relevant for pass-through MCP servers (see + ``MCPServer.is_oauth_passthrough``). The gateway converts this exception + into an HTTP 401 response on single-server routes, preserving any + ``WWW-Authenticate`` challenge emitted by the upstream so standards- + compliant MCP clients can trigger the upstream OAuth flow. + """ + + def __init__( + self, + status_code: int, + www_authenticate: Optional[str], + server_name: str, + ) -> None: + self.status_code = status_code + self.www_authenticate = www_authenticate + self.server_name = server_name + super().__init__(f"Upstream MCP server {server_name!r} returned {status_code}") + + def to_http_exception( + self, + base_url: Optional[str] = None, + request_path: Optional[str] = None, + ) -> HTTPException: + """Convert this upstream-auth error into an ``HTTPException`` that + preserves the upstream status code and any ``WWW-Authenticate`` + challenge, so standards-compliant MCP clients can trigger the + upstream OAuth flow. + + When the upstream 401 omits ``WWW-Authenticate`` (non-compliant per + RFC 7235 §3.1) we fabricate a ``Bearer resource_metadata=`` challenge + that points at the gateway's well-known endpoint for this server, so + MCP clients can still initiate RFC 9728 discovery against the upstream + IdP via the gateway's proxied metadata. Callers must pass ``base_url`` + (the gateway origin, no trailing slash) so the fabricated URI is + absolute as RFC 9728 §3.2 requires; if ``base_url`` is missing we + skip fabrication entirely rather than emit a relative URI that strict + clients reject in the Bearer challenge. + + When ``request_path`` is supplied and matches the legacy + ``/{server_name}/mcp`` MCP transport route, the fabricated URI uses + the matching legacy well-known form + ``/.well-known/oauth-protected-resource/{server_name}/mcp``. Otherwise + we default to the standard form + ``/.well-known/oauth-protected-resource/mcp/{server_name}``. This + keeps the ``resource_metadata`` URI aligned with the resource pattern + the client originally targeted, matching the path-aware behaviour of + ``_get_passthrough_resource_metadata_url`` in ``server.py``. + """ + challenge: Optional[str] = self.www_authenticate + if challenge is None and self.status_code == 401 and base_url: + prefix = base_url.rstrip("/") + if request_path and request_path.startswith(f"/{self.server_name}/mcp"): + resource_metadata_url = ( + f"{prefix}/.well-known/oauth-protected-resource/" + f"{self.server_name}/mcp" + ) + else: + resource_metadata_url = ( + f"{prefix}/.well-known/oauth-protected-resource/" + f"mcp/{self.server_name}" + ) + challenge = f'Bearer resource_metadata="{resource_metadata_url}"' + detail = "Forbidden" if self.status_code == 403 else "Unauthorized" + return HTTPException( + status_code=self.status_code, + detail=detail, + headers={"www-authenticate": challenge} if challenge else None, + ) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index b4678a50b2c..739dc4a2f88 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -48,6 +48,7 @@ from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, ) +from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError from litellm.proxy._experimental.mcp_server.oauth2_token_cache import resolve_mcp_auth from litellm.proxy._experimental.mcp_server.utils import ( MCP_TOOL_PREFIX_SEPARATOR, @@ -118,6 +119,103 @@ _AZURE_ENTRA_HOSTS = { } +def _should_strip_caller_authorization( + mcp_server: MCPServer, + raw_headers: Optional[Dict[str, str]], + user_api_key_auth: Optional[UserAPIKeyAuth], +) -> bool: + """Decide whether the caller's ``Authorization`` header must NOT be + forwarded upstream when populating ``extra_headers`` for an MCP server. + + Centralized so ``_call_regular_mcp_tool`` (this module) and + ``_prepare_mcp_server_headers`` (``server.py``) cannot drift apart on + this security-sensitive decision. + + Strip rules: + - **M2M (client_credentials) servers**: never forward the caller's + ``Authorization`` — the proxy fetches its own upstream token. + - **OAuth pass-through servers**: strip when the ``Authorization`` + header is actually the LiteLLM API key — either because admission + validated it (``user_api_key_auth.api_key`` is set) and the caller + did NOT also supply ``x-litellm-api-key`` to disambiguate, or + because the legacy ``user_api_key_auth is None`` call sites did + not supply an explicit admission header. In the anonymous / + pass-through cold-start case (RFC 9728) the bearer in + ``Authorization`` is the upstream OAuth token and must be + forwarded, so we keep it. + """ + if mcp_server.has_client_credentials: + return True + if not mcp_server.is_oauth_passthrough: + return False + + normalized_raw_headers = { + str(k).lower(): v for k, v in (raw_headers or {}).items() if isinstance(k, str) + } + has_explicit_litellm_admission_header = ( + normalized_raw_headers.get("x-litellm-api-key") is not None + ) + admission_consumed_authorization_as_litellm_key = ( + user_api_key_auth is not None + and bool(getattr(user_api_key_auth, "api_key", None)) + and not has_explicit_litellm_admission_header + ) + return admission_consumed_authorization_as_litellm_key or ( + user_api_key_auth is None and not has_explicit_litellm_admission_header + ) + + +def _extract_upstream_auth_failure( + exc: BaseException, +) -> Optional[Tuple[int, Optional[str]]]: + """Walk the exception tree looking for an HTTP 401/403 response from the + upstream MCP server. + + The MCP SDK wraps transport errors in anyio ``ExceptionGroup`` objects and + may chain through ``__cause__`` / ``__context__``. We inspect all of those + layers for an ``httpx.Response``-bearing exception (typically + ``httpx.HTTPStatusError``) and extract the status code and any upstream + ``WWW-Authenticate`` header. + + Returns ``(status_code, www_authenticate)`` on match, else ``None``. + """ + seen: Set[int] = set() + stack: List[BaseException] = [exc] + while stack: + current = stack.pop() + if id(current) in seen: + continue + seen.add(id(current)) + + response = getattr(current, "response", None) + if response is not None: + status_code = getattr(response, "status_code", None) + if isinstance(status_code, int) and status_code in (401, 403): + www_authenticate: Optional[str] = None + headers = getattr(response, "headers", None) + if headers is not None: + try: + www_authenticate = headers.get("www-authenticate") + except Exception: + www_authenticate = None + return status_code, www_authenticate + + # anyio / PEP 654 ExceptionGroup + sub_exceptions = getattr(current, "exceptions", None) + if sub_exceptions: + stack.extend(sub_exceptions) + + if current.__cause__ is not None: + stack.append(current.__cause__) + if ( + current.__context__ is not None + and current.__context__ is not current.__cause__ + ): + stack.append(current.__context__) + + return None + + def _warn_on_server_name_fields( *, server_id: str, @@ -483,6 +581,7 @@ class MCPServerManager: delegate_auth_to_upstream=bool( server_config.get("delegate_auth_to_upstream", False) ), + oauth_passthrough=bool(server_config.get("oauth_passthrough", False)), # AWS SigV4 fields aws_access_key_id=server_config.get("aws_access_key_id", None), aws_secret_access_key=server_config.get("aws_secret_access_key", None), @@ -881,6 +980,7 @@ class MCPServerManager: delegate_auth_to_upstream=bool( getattr(mcp_server, "delegate_auth_to_upstream", False) ), + oauth_passthrough=bool(getattr(mcp_server, "oauth_passthrough", False)), created_at=getattr(mcp_server, "created_at", None), updated_at=getattr(mcp_server, "updated_at", None), tool_name_to_display_name=_deserialize_json_dict( @@ -1599,7 +1699,9 @@ class MCPServerManager: ] return tools else: - tools = await self._fetch_tools_with_timeout(client, server.name) + tools = await self._fetch_tools_with_timeout( + client, server.name, server=server + ) self._remember_upstream_initialize_instructions(server, client) prefixed_or_original_tools = self._create_prefixed_tools( @@ -1608,6 +1710,11 @@ class MCPServerManager: return prefixed_or_original_tools + except MCPUpstreamAuthError: + # Pass-through 401 must surface to single-server routes so the + # client triggers the upstream OAuth flow. The multi-server + # aggregator catches this explicitly to keep absorbing. + raise except Exception as e: verbose_logger.warning( f"Failed to get tools from server {server.name}: {str(e)}" @@ -2209,7 +2316,10 @@ class MCPServerManager: return None async def _fetch_tools_with_timeout( - self, client: MCPClient, server_name: str + self, + client: MCPClient, + server_name: str, + server: Optional[MCPServer] = None, ) -> List[MCPTool]: """ Fetch tools from MCP client with timeout and error handling. @@ -2217,16 +2327,28 @@ class MCPServerManager: Uses anyio.fail_after() instead of asyncio.wait_for() to avoid conflicts with the MCP SDK's anyio TaskGroup. See GitHub issue #20715 for details. + For pass-through MCP servers (``MCPServer.is_oauth_passthrough``) an + upstream HTTP 401 is converted into :class:`MCPUpstreamAuthError` + instead of being swallowed to an empty tool list. That lets the + single-server HTTP routes surface a proper 401 + ``WWW-Authenticate`` + challenge so standards-compliant MCP clients trigger the upstream + OAuth flow. Non-pass-through servers keep today's swallow-and-log + behaviour so the multi-server ``/mcp`` aggregator doesn't get + tainted by a single bad server. + Args: client: MCP client instance server_name: Name of the server for logging + server: Optional MCPServer; when pass-through, auth errors are + re-raised as :class:`MCPUpstreamAuthError`. Returns: List of tools from the server """ + is_passthrough = bool(server is not None and server.is_oauth_passthrough) try: with anyio.fail_after(MCP_TOOL_LISTING_TIMEOUT): - tools = await client.list_tools() + tools = await client.list_tools(raise_on_error=is_passthrough) verbose_logger.debug(f"Tools from {server_name}: {tools}") return tools except TimeoutError: @@ -2243,6 +2365,19 @@ class MCPServerManager: ) return [] except Exception as e: + if is_passthrough: + auth_info = _extract_upstream_auth_failure(e) + if auth_info is not None: + status_code, www_authenticate = auth_info + verbose_logger.info( + f"Upstream auth failure from pass-through MCP server " + f"{server_name}: HTTP {status_code}" + ) + raise MCPUpstreamAuthError( + status_code=status_code, + www_authenticate=www_authenticate, + server_name=server_name, + ) from e verbose_logger.warning(f"Error listing tools from {server_name}: {str(e)}") return [] @@ -2650,6 +2785,9 @@ class MCPServerManager: "name": name, "arguments": arguments, "server_name": server_name, + "mcp_rate_limit_server_name": server.alias + or server.server_name + or server.name, "user_api_key_auth": user_api_key_auth, "user_api_key_user_id": ( getattr(user_api_key_auth, "user_id", None) @@ -2768,6 +2906,7 @@ class MCPServerManager: proxy_logging_obj: Optional[ProxyLogging], host_progress_callback: Optional[Callable] = None, hook_extra_headers: Optional[Dict[str, str]] = None, + user_api_key_auth: Optional[UserAPIKeyAuth] = None, ) -> CallToolResult: """ Call a regular MCP tool using the MCP client. @@ -2833,13 +2972,16 @@ class MCPServerManager: normalized_raw_headers = { str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str) } + strip_caller_authorization = _should_strip_caller_authorization( + mcp_server=mcp_server, + raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, + ) + for header in mcp_server.extra_headers: if not isinstance(header, str): continue - if ( - mcp_server.has_client_credentials - and header.lower() == "authorization" - ): + if header.lower() == "authorization" and strip_caller_authorization: continue header_value = normalized_raw_headers.get(header.lower()) if header_value is None: @@ -3131,6 +3273,7 @@ class MCPServerManager: proxy_logging_obj=proxy_logging_obj, host_progress_callback=host_progress_callback, hook_extra_headers=hook_result.get("extra_headers"), + user_api_key_auth=user_api_key_auth, ) return await self._gather_openapi_tool_tasks(tasks, proxy_logging_obj) @@ -3162,7 +3305,23 @@ class MCPServerManager: if server.needs_user_oauth_token: # Skip OAuth2 servers that rely on user-provided tokens continue - tools = await self._get_tools_from_server(server) + try: + tools = await self._get_tools_from_server(server) + except MCPUpstreamAuthError as e: + # Pass-through servers expect a user-supplied bearer token; + # at startup we have none, so an upstream 401 is normal. + # Swallow it so we keep mapping the remaining servers. + verbose_logger.debug( + f"Skipping tool name mapping for server {server.name} " + f"due to upstream auth error: {str(e)}" + ) + continue + except Exception as e: + verbose_logger.warning( + f"Failed to get tools from server {server.name} during " + f"tool name mapping initialization: {str(e)}" + ) + continue for tool in tools: # The tool.name here is already prefixed from _get_tools_from_server # Extract original name for mapping @@ -3754,6 +3913,7 @@ class MCPServerManager: allow_all_keys=server.allow_all_keys, available_on_public_internet=server.available_on_public_internet, delegate_auth_to_upstream=server.delegate_auth_to_upstream, + oauth_passthrough=getattr(server, "oauth_passthrough", False), is_byok=server.is_byok, byok_description=server.byok_description, byok_api_key_help_url=server.byok_api_key_help_url, diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 693ca5a7642..e20c9f3a082 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -16,6 +16,7 @@ from typing import ( from fastapi import APIRouter, Depends, HTTPException, Query, Request, status from litellm._logging import verbose_logger +from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError from litellm.proxy._experimental.mcp_server.ui_session_utils import ( build_effective_auth_contexts, ) @@ -46,6 +47,9 @@ if MCP_AVAILABLE: from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( global_mcp_server_manager, ) + from litellm.proxy._experimental.mcp_server.oauth_utils import ( + get_request_base_url, + ) from litellm.proxy._experimental.mcp_server.server import ( ListMCPToolsRestAPIResponseObject, MCPServer, @@ -423,101 +427,6 @@ if MCP_AVAILABLE: allowed_mcp_servers.append(server) return allowed_mcp_servers - async def _list_tools_for_single_server( - server_id: str, - allowed_server_ids: List[str], - rest_client_ip: Optional[str], - mcp_server_auth_headers: dict, - mcp_auth_header: Optional[str], - raw_headers_from_request: dict, - user_api_key_dict: "UserAPIKeyAuth", - ) -> dict: - """ - Resolve and fetch tools for a single specified MCP server. - - Returns the full REST response dict (tools / error / message). - Raises HTTPException on access / IP-filter errors. - """ - # Resolve a server name to its UUID if needed - _name_resolved = None - if server_id not in allowed_server_ids: - _name_resolved = global_mcp_server_manager.get_mcp_server_by_name(server_id) - if _name_resolved is not None and _name_resolved.server_id in set( - allowed_server_ids - ): - server_id = _name_resolved.server_id - - if server_id not in allowed_server_ids: - _server = ( - global_mcp_server_manager.get_mcp_server_by_id(server_id) - or _name_resolved - ) - 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={ - "error": "access_denied", - "message": f"The key is not allowed to access server {server_id}", - }, - ) - - server = global_mcp_server_manager.get_mcp_server_by_id(server_id) - if server is None: - return { - "tools": [], - "error": "server_not_found", - "message": f"Server with id {server_id} not found", - } - - server_auth_header = _get_server_auth_header( - server, mcp_server_auth_headers, mcp_auth_header - ) - user_oauth_extra_headers = await _get_user_oauth_extra_headers( - server, user_api_key_dict - ) - - try: - tools = await _get_tools_for_single_server( - server, - server_auth_header, - raw_headers_from_request, - user_api_key_dict, - extra_headers=user_oauth_extra_headers, - ) - except Exception as e: - verbose_logger.exception(f"Error getting tools from {server.name}: {e}") - return { - "tools": [], - "error": "server_error", - "message": f"Failed to get tools from server {server.name}: {str(e)}", - } - - return { - "tools": tools, - "error": None, - "message": "Successfully retrieved tools", - } - - ######################################################## - async def _list_tools_for_single_server( server_id: str, allowed_server_ids: List[str], @@ -591,6 +500,11 @@ if MCP_AVAILABLE: user_api_key_dict, extra_headers=user_oauth_extra_headers, ) + except MCPUpstreamAuthError: + # Surface the upstream 401/403 to the caller so it can emit the + # matching status code and WWW-Authenticate challenge; that is what + # lets standards-compliant MCP clients run the upstream OAuth flow. + raise except Exception as e: verbose_logger.exception(f"Error getting tools from {server.name}: {e}") return { @@ -757,6 +671,24 @@ if MCP_AVAILABLE: ), } + except MCPUpstreamAuthError as e: + # Surface upstream pass-through 401/403 challenges to the client so + # standards-compliant MCP clients can run the upstream OAuth flow. + raise e.to_http_exception( + base_url=get_request_base_url(request), + request_path=request.scope.get("_original_path") or request.url.path, + ) + except HTTPException as http_exc: + # Internal access/IP 403s keep the legacy error-dict response shape + # so the existing contract stays intact. + verbose_logger.exception( + "HTTPException in list_tool_rest_api: %s", str(http_exc) + ) + return { + "tools": [], + "error": "unexpected_error", + "message": (f"An unexpected error occurred: {http_exc.detail}"), + } except Exception as e: verbose_logger.exception( "Unexpected error in list_tool_rest_api: %s", str(e) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index c17ec13d3ef..6e33a105ec8 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -20,6 +20,7 @@ from typing import ( Dict, List, Optional, + Set, Tuple, Union, cast, @@ -38,6 +39,7 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, ) +from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( get_request_base_url, ) @@ -145,6 +147,19 @@ _SESSION_MANAGERS_INITIALIZED = False _INITIALIZATION_LOCK = asyncio.Lock() +def _mcp_session_id_from_headers( + raw_headers: Optional[Dict[str, str]], +) -> Optional[str]: + """The ``mcp-session-id`` of a stateful MCP session, read case-insensitively + from the request headers. ``None`` for stateless calls (no such header).""" + if not raw_headers: + return None + for key, value in raw_headers.items(): + if isinstance(key, str) and key.lower() == "mcp-session-id": + return value or None + return None + + if MCP_AVAILABLE: from mcp.server import Server from mcp.server.lowlevel.server import NotificationOptions @@ -174,6 +189,7 @@ if MCP_AVAILABLE: ) from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( MCPServerManager, + _should_strip_caller_authorization, global_mcp_server_manager, ) from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( @@ -1086,6 +1102,42 @@ if MCP_AVAILABLE: return allowed_mcp_servers + def _client_has_passthrough_authorization( + server: MCPServer, + oauth2_headers: Optional[Dict[str, str]], + mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]], + ) -> bool: + """True if the incoming request already carries an ``Authorization`` + header the gateway will forward to this pass-through server. + + The client may supply the bearer as either the top-level + ``Authorization`` header (surfaced via ``oauth2_headers``) or a + per-server ``x-mcp-auth-`` style header (surfaced via + ``mcp_server_auth_headers``). Either form skips the pre-emptive 401. + """ + if oauth2_headers: + for k in oauth2_headers.keys(): + if k.lower() == "authorization": + return True + if mcp_server_auth_headers: + for key in (server.alias, server.server_name, server.name): + if not key: + continue + server_headers = None + for k, v in mcp_server_auth_headers.items(): + if k.lower() == key.lower(): + server_headers = v + break + if server_headers is None: + continue + if isinstance(server_headers, str) and server_headers.strip(): + return True + if isinstance(server_headers, dict): + for hk in server_headers.keys(): + if hk.lower() == "authorization": + return True + return False + async def _get_user_oauth_extra_headers_from_db( server: MCPServer, user_api_key_auth: Optional[UserAPIKeyAuth], @@ -1266,6 +1318,7 @@ if MCP_AVAILABLE: mcp_auth_header: Optional[str], oauth2_headers: Optional[Dict[str, str]], raw_headers: Optional[Dict[str, str]], + user_api_key_auth: Optional[UserAPIKeyAuth] = None, ) -> Tuple[Optional[Union[Dict[str, str], str]], Optional[Dict[str, str]]]: """Build auth and extra headers for a server.""" server_auth_header: Optional[Union[Dict[str, str], str]] = None @@ -1298,10 +1351,20 @@ if MCP_AVAILABLE: str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str) } + # Centralized strip decision shared with + # ``MCPServerManager._call_regular_mcp_tool`` so the two + # code paths cannot drift on this security-sensitive choice. + # See ``_should_strip_caller_authorization`` for the rules. + strip_caller_authorization = _should_strip_caller_authorization( + mcp_server=server, + raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, + ) + for header in server.extra_headers: if not isinstance(header, str): continue - if server.has_client_credentials and header.lower() == "authorization": + if header.lower() == "authorization" and strip_caller_authorization: continue header_value = normalized_raw_headers.get(header.lower()) if header_value is None: @@ -1510,6 +1573,7 @@ if MCP_AVAILABLE: mcp_auth_header=mcp_auth_header, oauth2_headers=oauth2_headers, raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, ) # Prefer server-stored per-user OAuth when configured, so a stale @@ -1561,6 +1625,13 @@ if MCP_AVAILABLE: f"Successfully fetched {len(tools)} tools from server {server.name}, {len(filtered_tools)} after filtering" ) return filtered_tools + except MCPUpstreamAuthError: + # Surface upstream 401/403 to the outer handler so the + # client receives a proper WWW-Authenticate challenge + # instead of a silently empty tool list. Without this + # re-raise the broad ``except Exception`` below would + # swallow the auth error. + raise except Exception as e: verbose_logger.exception( f"Error getting tools from server {server.name}: {str(e)}" @@ -1684,6 +1755,7 @@ if MCP_AVAILABLE: mcp_auth_header=mcp_auth_header, oauth2_headers=oauth2_headers, raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, ) try: @@ -1741,6 +1813,7 @@ if MCP_AVAILABLE: mcp_auth_header=mcp_auth_header, oauth2_headers=oauth2_headers, raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, ) try: @@ -1796,6 +1869,7 @@ if MCP_AVAILABLE: mcp_auth_header=mcp_auth_header, oauth2_headers=oauth2_headers, raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, ) try: @@ -2324,6 +2398,7 @@ if MCP_AVAILABLE: name=original_tool_name, # Use original name for logging arguments=arguments, server_name=server_name, + session_id=_mcp_session_id_from_headers(raw_headers), ) ) litellm_logging_obj: Optional[LiteLLMLoggingObj] = kwargs.get( @@ -2631,6 +2706,7 @@ if MCP_AVAILABLE: mcp_auth_header=mcp_auth_header, oauth2_headers=oauth2_headers, raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, ) return await global_mcp_server_manager.get_prompt_from_server( @@ -2681,6 +2757,7 @@ if MCP_AVAILABLE: mcp_auth_header=mcp_auth_header, oauth2_headers=oauth2_headers, raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, ) return await global_mcp_server_manager.read_resource_from_server( @@ -2695,8 +2772,10 @@ if MCP_AVAILABLE: name: str, arguments: Dict[str, Any], server_name: Optional[str], + session_id: Optional[str] = None, ) -> StandardLoggingMCPToolCall: mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name) + namespaced_tool_name = f"{server_name}/{name}" if server_name else name if mcp_server: mcp_info = mcp_server.mcp_info or {} return StandardLoggingMCPToolCall( @@ -2704,13 +2783,15 @@ if MCP_AVAILABLE: arguments=arguments, mcp_server_name=mcp_info.get("server_name"), mcp_server_logo_url=mcp_info.get("logo_url"), - namespaced_tool_name=f"{server_name}/{name}" if server_name else name, + namespaced_tool_name=namespaced_tool_name, + mcp_session_id=session_id, ) else: return StandardLoggingMCPToolCall( name=name, arguments=arguments, - namespaced_tool_name=f"{server_name}/{name}" if server_name else name, + namespaced_tool_name=namespaced_tool_name, + mcp_session_id=session_id, ) async def _handle_managed_mcp_tool( @@ -3136,6 +3217,117 @@ if MCP_AVAILABLE: ) return user_api_key_auth.model_copy(update={"object_permission": updated_op}) + def _get_passthrough_resource_metadata_url(scope: Scope, server_name: str) -> str: + request = StarletteRequest(scope) + base_url = get_request_base_url(request) + _path = scope.get("_original_path") or scope.get("path", "") or "" + + if _path.startswith(f"/{server_name}/mcp"): + return f"{base_url}/.well-known/oauth-protected-resource/{server_name}/mcp" + return f"{base_url}/.well-known/oauth-protected-resource/mcp/{server_name}" + + def _get_passthrough_www_authenticate( + scope: Scope, + server_name: str, + invalid_token: bool = False, + ) -> str: + resource_metadata_url = _get_passthrough_resource_metadata_url( + scope=scope, + server_name=server_name, + ) + params = [] + if invalid_token: + params.append('error="invalid_token"') + params.append(f'resource_metadata="{resource_metadata_url}"') + return "Bearer " + ", ".join(params) + + async def _raise_preemptive_401_for_unauthenticated_servers( + scope: Scope, + mcp_servers: Optional[List[str]], + oauth2_headers: Optional[Dict[str, str]], + mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]], + user_api_key_auth: Optional[UserAPIKeyAuth], + client_ip: Optional[str], + allowed_server_ids: Optional[Set[str]] = None, + ) -> None: + """Fail fast with HTTP 401 for MCP servers that need user auth but + didn't receive it on this request. Covers both gateway-managed OAuth2 + (points clients at the gateway AS metadata) and pass-through OAuth + (points clients at the upstream resource-metadata via our well-known). + + ``allowed_server_ids`` may be passed by callers that have already + narrowed the authorized server set (e.g. toolset scoping); servers + not in that set are skipped so a client targeting a toolset that + excludes a passthrough server is not pushed into an OAuth flow for + a server it will be 403'd on immediately after authentication. + """ + for server_name in mcp_servers or []: + server = global_mcp_server_manager.get_mcp_server_by_name( + server_name, client_ip=client_ip + ) + if ( + server is not None + and allowed_server_ids is not None + and server.server_id not in allowed_server_ids + ): + # Caller's narrowed scope excludes this server — skip the + # preemptive challenge and let downstream authorization + # return 403. + continue + if server and server.auth_type == MCPAuth.oauth2 and not oauth2_headers: + # For per-user OAuth servers, only skip the pre-emptive 401 when + # a stored token actually exists for this user+server pair. + # If no stored token exists, fail fast with 401 so clients can + # kick off PKCE/interactive OAuth flow immediately. + if server.needs_user_oauth_token: + stored_oauth_headers = await _get_user_oauth_extra_headers_from_db( + server=server, + user_api_key_auth=user_api_key_auth, + ) + if stored_oauth_headers: + continue + + request = StarletteRequest(scope) + base_url = get_request_base_url(request) + _path = scope.get("_original_path") or scope.get("path", "") or "" + + # Pick the well-known AS-metadata form that matches the inbound route + # so strict RFC 9728 §3.2 clients can resolve it correctly. + if _path.startswith(f"/mcp/{server_name}"): + _as_url = f"{base_url}/.well-known/oauth-authorization-server/mcp/{server_name}" + else: + _as_url = f"{base_url}/.well-known/oauth-authorization-server/{server_name}" + authorization_uri = f'Bearer authorization_uri="{_as_url}"' + + raise HTTPException( + status_code=401, + detail="Unauthorized", + headers={"www-authenticate": authorization_uri}, + ) + + # Pass-through OAuth: when the admin has opted a server into + # forwarding the client's bearer token (is_oauth_passthrough) and + # the client hasn't supplied one, fail fast with 401 and point + # them at the gateway's oauth-protected-resource well-known URL. + # That endpoint proxies the upstream's metadata so the client + # kicks off OAuth against the real upstream IdP, not the gateway. + if ( + server + and server.is_oauth_passthrough + and not _client_has_passthrough_authorization( + server, oauth2_headers, mcp_server_auth_headers + ) + ): + www_authenticate = _get_passthrough_www_authenticate( + scope=scope, + server_name=server_name, + ) + raise HTTPException( + status_code=401, + detail="Unauthorized", + headers={"www-authenticate": www_authenticate}, + ) + def _get_forwarded_auth_from_scope(scope: Scope) -> Optional[str]: """Return the upstream-bound ``Authorization`` header value, or None. @@ -3248,12 +3440,15 @@ if MCP_AVAILABLE: passthrough_servers = [ srv for srv in allowed_servers - if srv.extra_headers - and any(h.lower() == "authorization" for h in srv.extra_headers) - # Exclude M2M servers: _prepare_mcp_server_headers skips caller - # Authorization when has_client_credentials is set, so probing - # those with the caller's token would send the wrong credential. - and not srv.has_client_credentials + # Restrict to genuine OAuth pass-through servers (auth_type none + + # Authorization in extra_headers). Gateway-managed OAuth2 servers + # must not receive the ``resource_metadata=`` challenge emitted + # below — they require ``authorization_uri=`` pointing at the + # gateway AS metadata. ``is_oauth_passthrough`` already requires + # ``auth_type in (None, MCPAuth.none)``, which is mutually + # exclusive with ``has_client_credentials`` (oauth2 + M2M flow), + # so M2M servers are implicitly excluded here. + if srv.is_oauth_passthrough ] if not passthrough_servers: return @@ -3264,19 +3459,20 @@ if MCP_AVAILABLE: for srv in passthrough_servers ] ) - request = StarletteRequest(scope) - base_url = get_request_base_url(request) for srv, (probe_status, _) in zip(passthrough_servers, probe_results): if probe_status == 401: - # Token is missing or expired — direct the client to re-authorize. - authorization_uri = ( - f"Bearer authorization_uri=" - f"{base_url}/.well-known/oauth-authorization-server/{srv.name}" + # Token is missing or expired: keep pass-through clients on the + # protected-resource discovery flow so they re-authorize against + # the upstream IdP metadata proxied by LiteLLM. + www_authenticate = _get_passthrough_www_authenticate( + scope=scope, + server_name=srv.name, + invalid_token=True, ) raise HTTPException( status_code=401, detail="Unauthorized", - headers={"WWW-Authenticate": authorization_uri}, + headers={"www-authenticate": www_authenticate}, ) if probe_status == 403: # Token is valid but the caller lacks permission — do not hint @@ -3311,39 +3507,6 @@ if MCP_AVAILABLE: verbose_logger.debug( f"MCP server auth headers: {list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None}" ) - # https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response - for server_name in mcp_servers or []: - server = global_mcp_server_manager.get_mcp_server_by_name( - server_name, client_ip=_client_ip - ) - if server and server.auth_type == MCPAuth.oauth2 and not oauth2_headers: - # For per-user OAuth servers, only skip the pre-emptive 401 when - # a stored token actually exists for this user+server pair. - # If no stored token exists, fail fast with 401 so clients can - # kick off PKCE/interactive OAuth flow immediately. - if server.needs_user_oauth_token: - stored_oauth_headers = ( - await _get_user_oauth_extra_headers_from_db( - server=server, - user_api_key_auth=user_api_key_auth, - ) - ) - if stored_oauth_headers: - continue - - request = StarletteRequest(scope) - base_url = get_request_base_url(request) - - authorization_uri = ( - f"Bearer authorization_uri=" - f"{base_url}/.well-known/oauth-authorization-server/{server_name}" - ) - - raise HTTPException( - status_code=401, - detail="Unauthorized", - headers={"www-authenticate": authorization_uri}, - ) # Strip any client-supplied x-mcp-toolset-id to prevent forgery. scope["headers"] = [ @@ -3355,10 +3518,28 @@ if MCP_AVAILABLE: # Apply toolset scope if set server-side via ContextVar (set by # /toolset/{name}/mcp and /{name}/mcp route handlers in proxy_server.py). active_toolset_id = _mcp_active_toolset_id.get() + toolset_allowed_server_ids: Optional[Set[str]] = None if active_toolset_id and user_api_key_auth is not None: user_api_key_auth = await _apply_toolset_scope( user_api_key_auth, active_toolset_id ) + op = user_api_key_auth.object_permission + toolset_allowed_server_ids = set(op.mcp_servers or []) if op else set() + + # https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response + # Must run after toolset scoping so the challenge set is derived + # from the fully-authorized server set: a passthrough server that + # the active toolset excludes should not trigger an OAuth flow + # for a server the caller will be 403'd on after authentication. + await _raise_preemptive_401_for_unauthenticated_servers( + scope=scope, + mcp_servers=mcp_servers, + oauth2_headers=oauth2_headers, + mcp_server_auth_headers=mcp_server_auth_headers, + user_api_key_auth=user_api_key_auth, + client_ip=_client_ip, + allowed_server_ids=toolset_allowed_server_ids, + ) # Pre-flight auth check for pass-through servers. Must run after # toolset scoping so the probe list is derived from the fully-authorized @@ -3589,6 +3770,13 @@ if MCP_AVAILABLE: not in _stateful_session_auth_contexts ): _stateful_session_locks.pop(active_request_session_id, None) + except MCPUpstreamAuthError as e: + # Pass-through server returned 401 — surface it to the client so + # standards-compliant MCP clients trigger the upstream OAuth flow. + raise e.to_http_exception( + base_url=get_request_base_url(StarletteRequest(scope)), + request_path=scope.get("_original_path") or scope.get("path"), + ) except HTTPException: # Re-raise HTTP exceptions to preserve status codes and details raise @@ -3632,6 +3820,50 @@ if MCP_AVAILABLE: verbose_logger.debug( f"MCP server auth headers: {list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None}" ) + + # Strip any client-supplied x-mcp-toolset-id to prevent forgery. + scope["headers"] = [ + (k, v) + for k, v in scope.get("headers", []) + if k.lower() != b"x-mcp-toolset-id" + ] + + # Apply toolset scope if set server-side via ContextVar so the + # downstream probe list matches the fully-authorized server set + # (mirrors the streamable HTTP handler). + active_toolset_id = _mcp_active_toolset_id.get() + toolset_allowed_server_ids: Optional[Set[str]] = None + if active_toolset_id and user_api_key_auth is not None: + user_api_key_auth = await _apply_toolset_scope( + user_api_key_auth, active_toolset_id + ) + op = user_api_key_auth.object_permission + toolset_allowed_server_ids = set(op.mcp_servers or []) if op else set() + + # https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response + # Must run after toolset scoping so the challenge set is derived + # from the fully-authorized server set: a passthrough server that + # the active toolset excludes should not trigger an OAuth flow + # for a server the caller will be 403'd on after authentication. + await _raise_preemptive_401_for_unauthenticated_servers( + scope=scope, + mcp_servers=mcp_servers, + oauth2_headers=oauth2_headers, + mcp_server_auth_headers=mcp_server_auth_headers, + user_api_key_auth=user_api_key_auth, + client_ip=_sse_client_ip, + allowed_server_ids=toolset_allowed_server_ids, + ) + + # Pre-flight auth check for pass-through servers: surface upstream + # 401/403 as a proper challenge before the SSE session commits 200 + # headers, so clients can refresh their OAuth token instead of + # being stuck with a silently empty tool list. Must run after + # toolset scoping so the probe list is derived from the fully- + # authorized server set, not the raw user-supplied names. + await _check_passthrough_upstream_auth( + scope, user_api_key_auth, mcp_servers, _sse_client_ip + ) set_auth_context( user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, @@ -3652,9 +3884,20 @@ if MCP_AVAILABLE: _sse_client_ip, ): await sse_session_manager.handle_request(scope, receive, send) + except MCPUpstreamAuthError as e: + # Pass-through server returned 401 — surface it to the client so + # standards-compliant MCP clients trigger the upstream OAuth flow. + raise e.to_http_exception( + base_url=get_request_base_url(StarletteRequest(scope)), + request_path=scope.get("_original_path") or scope.get("path"), + ) + except HTTPException: + # Re-raise HTTP exceptions to preserve status codes and details + # (e.g. 401 + WWW-Authenticate challenges from OAuth pass-through). + raise except Exception as e: verbose_logger.exception(f"Error handling MCP request: {e}") - # Instead of re-raising, try to send a graceful error response + # Try to send a graceful error response for non-HTTP exceptions try: # Send a proper HTTP error response instead of letting the exception bubble up from starlette.responses import JSONResponse diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 98a17e4be95..feab3b04a5e 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -701,6 +701,10 @@ class LiteLLMRoutes(enum.Enum): "/v2/guardrails/list", "/project/list", "/project/info", + # Read-only search tool routes power the Search Tools UI page. + # Create/update/delete and test_connection stay admin-only. + "/search_tools/list", + "/search_tools/ui/available_providers", ] + spend_tracking_routes + key_management_routes @@ -1050,6 +1054,7 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase): model_config = ConfigDict(protected_namespaces=()) model_rpm_limit: Optional[dict] = None model_tpm_limit: Optional[dict] = None + mcp_rpm_limit: Optional[Dict[str, int]] = None guardrails: Optional[List[str]] = None policies: Optional[List[str]] = None prompts: Optional[List[str]] = None @@ -1290,6 +1295,7 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase): allow_all_keys: bool = False available_on_public_internet: bool = True delegate_auth_to_upstream: bool = False + oauth_passthrough: bool = False is_byok: bool = False byok_description: List[str] = Field(default_factory=list) byok_api_key_help_url: Optional[str] = None @@ -1373,6 +1379,7 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase): allow_all_keys: bool = False available_on_public_internet: bool = True delegate_auth_to_upstream: bool = False + oauth_passthrough: bool = False is_byok: bool = False byok_description: List[str] = Field(default_factory=list) byok_api_key_help_url: Optional[str] = None @@ -1445,6 +1452,7 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase): allow_all_keys: bool = False available_on_public_internet: bool = True delegate_auth_to_upstream: bool = False + oauth_passthrough: bool = False is_byok: bool = False byok_description: List[str] = Field(default_factory=list) byok_api_key_help_url: Optional[str] = None @@ -1851,6 +1859,7 @@ class NewTeamRequest(TeamBase): ] = None # raise an error if 'guaranteed_throughput' is set and we're overallocating tpm model_tpm_limit: Optional[Dict[str, int]] = None + mcp_rpm_limit: Optional[Dict[str, int]] = None team_member_budget: Optional[float] = ( None # allow user to set a budget for all team members ) @@ -1920,6 +1929,7 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase): prompts: Optional[List[str]] = None model_rpm_limit: Optional[Dict[str, int]] = None model_tpm_limit: Optional[Dict[str, int]] = None + mcp_rpm_limit: Optional[Dict[str, int]] = None allowed_vector_store_indexes: Optional[List[AllowedVectorStoreIndexItem]] = None enforced_batch_output_expires_after: Optional[dict] = None enforced_file_expires_after: Optional[dict] = None @@ -2516,7 +2526,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): ) mcp_trusted_proxy_ranges: Optional[List[str]] = Field( None, - description="CIDR ranges of trusted reverse proxies. When set, X-Forwarded-For headers are only trusted from these IPs.", + description="CIDR ranges of trusted reverse proxies. When set, X-Forwarded-For and X-Forwarded-* origin headers are only trusted from these IPs.", ) trusted_proxy_ranges: Optional[List[str]] = Field( None, @@ -3667,6 +3677,7 @@ class ProxyException(Exception): provider_specific_fields: Optional[dict] = None, ): self.message = str(message) + super().__init__(self.message) self.type = type self.param = param self.openai_code = openai_code or code @@ -4285,6 +4296,7 @@ class PassThroughEndpointLoggingTypedDict(TypedDict): LiteLLM_ManagementEndpoint_MetadataFields = [ "model_rpm_limit", "model_tpm_limit", + "mcp_rpm_limit", "rpm_limit_type", "tpm_limit_type", "enforced_params", @@ -4436,6 +4448,72 @@ class JWTRoutingOverride(BaseModel): } +class JWTIssuerConfig(BaseModel): + """ + Issuer-bound JWT validation configuration. + + When a token's unverified `iss` claim matches an entry in + ``LiteLLM_JWTAuth.issuers``, LiteLLM validates it only against that + issuer's JWKS and audience. Tokens whose `iss` does not match any + configured issuer fall back to the global JWT_AUDIENCE/JWT_ISSUER + validation path; `issuers` is additive routing, not an allow-list. + """ + + issuer: str = Field(description="Exact expected JWT issuer (`iss`) value.") + jwks_url: Optional[str] = Field( + default=None, + description="Issuer JWKS URL. If omitted, LiteLLM uses the issuer's OIDC discovery document.", + ) + audience: Optional[Union[str, List[str]]] = Field( + default=None, + description="Expected token audience for this issuer.", + ) + disable_audience_validation: bool = Field( + default=False, + description="Explicitly disable audience validation for this issuer. Use only when the issuer cannot provide an audience suitable for LiteLLM.", + ) + user_id_jwt_field: Optional[str] = Field( + default=None, + description="Issuer-specific claim path to normalize into LiteLLM's user id.", + ) + user_email_jwt_field: Optional[str] = Field( + default=None, + description="Issuer-specific claim path to normalize into LiteLLM's user email.", + ) + team_id_jwt_field: Optional[str] = Field( + default=None, + description="Issuer-specific claim path to normalize into LiteLLM's team id.", + ) + team_ids_jwt_field: Optional[str] = Field( + default=None, + description="Issuer-specific claim path to normalize into LiteLLM's team ids.", + ) + org_id_jwt_field: Optional[str] = Field( + default=None, + description="Issuer-specific claim path to normalize into LiteLLM's organization id.", + ) + end_user_id_jwt_field: Optional[str] = Field( + default=None, + description="Issuer-specific claim path to normalize into LiteLLM's end-user id.", + ) + + model_config = { + "extra": "forbid", + } + + @model_validator(mode="after") + def validate_audience_configured(self) -> "JWTIssuerConfig": + if self.audience is None and not self.disable_audience_validation: + raise ValueError( + f"JWT issuer {self.issuer} must configure audience or set disable_audience_validation=True" + ) + if self.audience is not None and self.disable_audience_validation: + raise ValueError( + f"JWT issuer {self.issuer} cannot set audience and disable_audience_validation=True together" + ) + return self + + class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase): """ A class to define the roles and permissions for a LiteLLM Proxy w/ JWT Auth. @@ -4540,6 +4618,10 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase): default=None, description="Optional claim-based routing overrides for JWT-shaped tokens. Matching rules route requests to oauth2 before default JWT flow.", ) + issuers: Optional[List[JWTIssuerConfig]] = Field( + default=None, + description="Optional issuer-bound JWT validation rules. When a token's `iss` matches a configured issuer, validation uses that issuer's JWKS, audience, and claim mappings. Tokens with an unlisted `iss` fall back to the global JWT_AUDIENCE/JWT_ISSUER validation path — this is additive routing, not an allow-list.", + ) ######################################################### def __init__(self, **kwargs: Any) -> None: diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py index 993d30e3811..9f1403d4328 100644 --- a/litellm/proxy/agent_endpoints/a2a_endpoints.py +++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py @@ -6,12 +6,14 @@ The A2A SDK can point to LiteLLM's URL and invoke agents registered with LiteLLM """ import json -from typing import Any, Dict, List, Optional +from typing import Any, AsyncGenerator, Dict, List, Optional +from urllib.parse import urlparse from fastapi import APIRouter, Depends, HTTPException, Request, Response from fastapi.responses import JSONResponse, StreamingResponse from litellm._logging import verbose_proxy_logger +from litellm.litellm_core_utils.url_utils import SSRFError, validate_url from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.agent_endpoints.utils import merge_agent_headers from litellm.proxy.auth.user_api_key_auth import user_api_key_auth @@ -19,9 +21,66 @@ from litellm.types.utils import all_litellm_params router = APIRouter() +_PASCAL_TO_WIRE: Dict[str, str] = { + "GetTask": "tasks/get", + "ListTasks": "tasks/list", + "CancelTask": "tasks/cancel", + "SubscribeToTask": "tasks/resubscribe", + "CreateTaskPushNotificationConfig": "tasks/pushNotificationConfig/set", + "GetTaskPushNotificationConfig": "tasks/pushNotificationConfig/get", + "ListTaskPushNotificationConfigs": "tasks/pushNotificationConfig/list", + "DeleteTaskPushNotificationConfig": "tasks/pushNotificationConfig/delete", + "GetExtendedAgentCard": "agent/getAuthenticatedExtendedCard", +} + + +def _validate_push_notification_url(url: str) -> None: + parsed = urlparse(url) + if parsed.scheme != "https": + raise HTTPException( + status_code=400, + detail="Push notification URL must use HTTPS", + ) + try: + validate_url(url) + except (SSRFError, ValueError) as e: + raise HTTPException(status_code=400, detail=str(e)) from e + + +def _caller_identity_headers(user_api_key_dict: UserAPIKeyAuth) -> Dict[str, str]: + headers: Dict[str, str] = {} + if user_api_key_dict.user_id: + headers["X-LiteLLM-User-Id"] = user_api_key_dict.user_id + if user_api_key_dict.team_id: + headers["X-LiteLLM-Team-Id"] = user_api_key_dict.team_id + return headers + + +def _forwarding_headers( + user_api_key_dict: UserAPIKeyAuth, + request_data: dict, + agent_extra_headers: Optional[Dict[str, str]], +) -> Optional[Dict[str, str]]: + sanitized = ( + { + k: v + for k, v in agent_extra_headers.items() + if not k.lower().startswith("x-litellm-") + } + if agent_extra_headers + else None + ) + merged = merge_agent_headers(dynamic_headers=sanitized, static_headers=None) or {} + identity = _caller_identity_headers(user_api_key_dict) + trace_id = request_data.get("litellm_trace_id") + if trace_id: + identity["X-LiteLLM-Trace-Id"] = str(trace_id) + merged.update(identity) + return merged or None + def _jsonrpc_error( - request_id: Optional[str], + request_id: Optional[Any], code: int, message: str, status_code: int = 400, @@ -67,9 +126,158 @@ def _enforce_inbound_trace_id(agent: Any, request: Request) -> None: ) +async def _forward_jsonrpc( + agent_url: str, + body: dict, + extra_headers: Optional[Dict[str, str]] = None, +) -> dict: + from litellm.llms.custom_httpx.http_handler import get_async_httpx_client + from litellm.types.llms.custom_http import httpxSpecialProvider + + headers = {"Content-Type": "application/json", **(extra_headers or {})} + handler = get_async_httpx_client( + llm_provider=httpxSpecialProvider.A2A, + params={"timeout": 60.0}, + ) + resp = await handler.post(agent_url, json=body, headers=headers) + try: + result = resp.json() + except Exception: + resp.raise_for_status() + raise + if not resp.is_success and "error" not in result: + resp.raise_for_status() + return result + + +async def _a2a_sse_event_source( + agent_url: str, + body: dict, + request_id: Optional[Any] = None, + extra_headers: Optional[Dict[str, str]] = None, +) -> AsyncGenerator[dict, None]: + """Stream an upstream A2A SSE response as parsed JSON-RPC event dicts. + + Upstream HTTP/JSON-RPC errors are surfaced as a single JSON-RPC error event + so the caller can relay them instead of breaking the stream. + """ + from litellm.llms.custom_httpx.http_handler import get_async_httpx_client + from litellm.types.agents import _normalize_a2a_jsonrpc_response + from litellm.types.llms.custom_http import httpxSpecialProvider + + headers = { + "Content-Type": "application/json", + "Accept": "text/event-stream", + **(extra_headers or {}), + } + handler = get_async_httpx_client( + llm_provider=httpxSpecialProvider.A2A, + params={"timeout": None}, + ) + async_client = handler.client + req = async_client.build_request("POST", agent_url, json=body, headers=headers) + resp = await async_client.send(req, stream=True) + try: + if not resp.is_success: + error_body = await resp.aread() + error_event: Optional[dict] = None + try: + parsed = json.loads(error_body) + if isinstance(parsed, dict) and "error" in parsed: + error_event = _normalize_a2a_jsonrpc_response( + parsed, request_id=request_id + ) + except Exception: + error_event = None + yield error_event or { + "jsonrpc": "2.0", + "id": request_id, + "error": {"code": -32603, "message": resp.reason_phrase}, + } + return + async for line in resp.aiter_lines(): + stripped = line.strip() + if not stripped.startswith("data:"): + continue + payload = stripped[len("data:") :].strip() + if not payload: + continue + try: + yield json.loads(payload) + except Exception: + continue + finally: + await resp.aclose() + + +async def _forward_jsonrpc_sse( + agent_url: str, + body: dict, + request_id: Optional[Any] = None, + extra_headers: Optional[Dict[str, str]] = None, + proxy_logging_obj: Optional[Any] = None, + user_api_key_dict: Optional[Any] = None, + request_data: Optional[dict] = None, +) -> StreamingResponse: + event_source = _a2a_sse_event_source( + agent_url, body, request_id=request_id, extra_headers=extra_headers + ) + + def _serialize_chunk(chunk: Any) -> str: + return f"data: {json.dumps(chunk)}\n\n" + + def _serialize_error(proxy_exc: Any) -> str: + return ( + "data: " + + json.dumps( + { + "jsonrpc": "2.0", + "id": request_id, + "error": { + "code": -32603, + "message": getattr(proxy_exc, "message", str(proxy_exc)), + }, + } + ) + + "\n\n" + ) + + if ( + proxy_logging_obj is not None + and user_api_key_dict is not None + and request_data is not None + ): + # Route streamed events through the shared streaming generator so the + # post-call streaming hook (and therefore agent guardrails) inspects + # tasks/resubscribe output the same way message/stream does. + from litellm.proxy.common_request_processing import ( + ProxyBaseLLMRequestProcessing, + ) + + generator: AsyncGenerator[str, None] = ( + ProxyBaseLLMRequestProcessing.async_streaming_data_generator( + response=event_source, + user_api_key_dict=user_api_key_dict, + request_data=request_data, + proxy_logging_obj=proxy_logging_obj, + serialize_chunk=_serialize_chunk, + serialize_error=_serialize_error, + ) + ) + else: + + async def _passthrough() -> AsyncGenerator[str, None]: + async for chunk in event_source: + yield _serialize_chunk(chunk) + + generator = _passthrough() + + return StreamingResponse(generator, media_type="text/event-stream") + + async def _handle_stream_message( api_base: Optional[str], - request_id: str, + request_id: Any, params: dict, litellm_params: Optional[dict] = None, agent_id: Optional[str] = None, @@ -310,8 +518,6 @@ async def invoke_agent_a2a( # noqa: PLR0915 - message/send: Send a message and get a response - message/stream: Send a message and stream the response """ - from litellm.a2a_protocol import asend_message - from litellm.a2a_protocol.main import A2A_SDK_AVAILABLE from litellm.proxy.agent_endpoints.auth.agent_permission_handler import ( AgentRequestHandler, ) @@ -322,9 +528,11 @@ async def invoke_agent_a2a( # noqa: PLR0915 version, ) - body = {} + body: Dict[str, Any] = {} + request_data: Dict[str, Any] = body try: body = await request.json() + request_data = body verbose_proxy_logger.debug(f"A2A request for agent '{agent_id}': {body}") @@ -334,11 +542,14 @@ async def invoke_agent_a2a( # noqa: PLR0915 body.get("id"), -32600, "Invalid Request: jsonrpc must be '2.0'" ) - request_id = body.get("id") - method = body.get("method") + request_id: Optional[Any] = body.get("id") + method: Optional[str] = body.get("method") params = body.get("params", {}) - if params: + if method: + method = _PASCAL_TO_WIRE.get(method, method) + + if isinstance(params, dict): # extract any litellm params from the params - eg. 'guardrails' # ``metadata`` is intentionally excluded: it's a first-class A2A # ``MessageSendParams`` field that the completion bridge forwards @@ -347,20 +558,12 @@ async def invoke_agent_a2a( # noqa: PLR0915 # silently drop the caller's A2A request-level metadata. params_to_remove = [] for key, value in params.items(): - if key in all_litellm_params and key != "metadata": + if key in all_litellm_params and key not in {"id", "metadata"}: params_to_remove.append(key) body[key] = value for key in params_to_remove: params.pop(key) - if not A2A_SDK_AVAILABLE: - return _jsonrpc_error( - request_id, - -32603, - "Server error: 'a2a' package not installed. Please install 'a2a-sdk'.", - 500, - ) - # Find the agent agent = _get_agent(agent_id) if agent is None: @@ -389,6 +592,19 @@ async def invoke_agent_a2a( # noqa: PLR0915 litellm_params = agent.litellm_params or {} custom_llm_provider = litellm_params.get("custom_llm_provider") + # Hand the authenticated key hash to the completion bridge so provider + # configs can scope provider-side session state per key (e.g. LangFlow + # session memory) instead of trusting the client-supplied A2A contextId. + if custom_llm_provider and user_api_key_dict.api_key: + from litellm.a2a_protocol.litellm_completion_bridge.handler import ( + A2A_USER_API_KEY_HASH_PARAM, + ) + + litellm_params = { + **litellm_params, + A2A_USER_API_KEY_HASH_PARAM: user_api_key_dict.api_key, + } + # URL is required unless using completion bridge with a provider that derives endpoint from model # (e.g., bedrock/agentcore derives endpoint from ARN in model string) if not agent_url and not custom_llm_provider: @@ -428,6 +644,7 @@ async def invoke_agent_a2a( # noqa: PLR0915 route_type="asend_message", version=version, ) + request_data = data # Build merged headers for the backend agent static_headers: Dict[str, str] = dict(agent.static_headers or {}) @@ -440,9 +657,10 @@ async def invoke_agent_a2a( # noqa: PLR0915 # 1. Admin-configured extra_headers: forward named headers from client request if agent.extra_headers: for header_name in agent.extra_headers: - val = normalized.get(header_name.lower()) + header_name_str = str(header_name) + val = normalized.get(header_name_str.lower()) if val is not None: - dynamic_headers[header_name] = val + dynamic_headers[header_name_str] = val # 2. Convention-based forwarding: x-a2a-{agent_id_or_name}-{header_name} # Matches both agent_id (UUID) and agent_name (alias), case-insensitive. @@ -476,10 +694,20 @@ async def invoke_agent_a2a( # noqa: PLR0915 # Route through SDK functions if method == "message/send": + from litellm.a2a_protocol import asend_message + from litellm.a2a_protocol.main import A2A_SDK_AVAILABLE + + if not A2A_SDK_AVAILABLE: + return _jsonrpc_error( + request_id, + -32603, + "Server error: 'a2a' package not installed. Please install 'a2a-sdk'.", + 500, + ) from a2a.types import MessageSendParams, SendMessageRequest a2a_request = SendMessageRequest( - id=request_id, + id=request_id if request_id is not None else "", params=MessageSendParams(**params), ) # Defer spend-log until after post_call_success_hook so guardrail @@ -519,7 +747,7 @@ async def invoke_agent_a2a( # noqa: PLR0915 elif method == "message/stream": return await _handle_stream_message( api_base=agent_url, - request_id=request_id, + request_id=request_id if request_id is not None else "", params=params, litellm_params=litellm_params, agent_id=agent.agent_id, @@ -530,6 +758,106 @@ async def invoke_agent_a2a( # noqa: PLR0915 request_data=data, proxy_logging_obj=proxy_logging_obj, ) + elif method in { + "tasks/get", + "tasks/list", + "tasks/cancel", + "tasks/pushNotificationConfig/set", + "tasks/pushNotificationConfig/get", + "tasks/pushNotificationConfig/list", + "tasks/pushNotificationConfig/delete", + "agent/getAuthenticatedExtendedCard", + }: + if not agent_url: + return _jsonrpc_error( + request_id, -32000, f"Agent '{agent_id}' has no URL configured", 500 + ) + if method == "tasks/pushNotificationConfig/set": + if not isinstance(params, dict): + raise HTTPException( + status_code=400, + detail="params must be an object", + ) + push_config = params.get("pushNotificationConfig", {}) + if "pushNotificationConfig" in params and not isinstance( + push_config, dict + ): + raise HTTPException( + status_code=400, + detail="pushNotificationConfig must be an object", + ) + for callback_url in (params.get("url"), push_config.get("url")): + if not callback_url: + continue + if not isinstance(callback_url, str): + raise HTTPException( + status_code=400, + detail="Push notification URL must be a string", + ) + _validate_push_notification_url(callback_url) + forward_body = { + "jsonrpc": "2.0", + "id": request_id, + "method": method, + "params": params, + } + caller_headers = _forwarding_headers( + user_api_key_dict=user_api_key_dict, + request_data=data, + agent_extra_headers=agent_extra_headers, + ) + result = await _forward_jsonrpc( + agent_url, forward_body, extra_headers=caller_headers + ) + if method == "agent/getAuthenticatedExtendedCard": + if isinstance(result.get("result"), dict) and "url" in result["result"]: + result["result"][ + "url" + ] = f"{str(request.base_url).rstrip('/')}/a2a/{agent_id}" + from litellm.types.agents import LiteLLMSendMessageResponse + + response = LiteLLMSendMessageResponse.from_dict( + result, request_id=request_id + ) + response = await proxy_logging_obj.post_call_success_hook( + user_api_key_dict=user_api_key_dict, + data=data, + response=response, + ) + return JSONResponse( + content=( + response.model_dump(mode="json", exclude_none=True) + if hasattr(response, "model_dump") + else response + ) + ) + + elif method == "tasks/resubscribe": + if not agent_url: + return _jsonrpc_error( + request_id, -32000, f"Agent '{agent_id}' has no URL configured", 500 + ) + forward_body = { + "jsonrpc": "2.0", + "id": request_id, + "method": method, + "params": params, + } + sse_caller_headers = _forwarding_headers( + user_api_key_dict=user_api_key_dict, + request_data=data, + agent_extra_headers=agent_extra_headers, + ) + return await _forward_jsonrpc_sse( + agent_url, + forward_body, + request_id=request_id, + extra_headers=sse_caller_headers, + proxy_logging_obj=proxy_logging_obj, + user_api_key_dict=user_api_key_dict, + request_data=data, + ) + else: return _jsonrpc_error(request_id, -32601, f"Method '{method}' not found") @@ -537,4 +865,12 @@ async def invoke_agent_a2a( # noqa: PLR0915 raise except Exception as e: verbose_proxy_logger.exception(f"Error invoking agent: {e}") + try: + await proxy_logging_obj.post_call_failure_hook( + user_api_key_dict=user_api_key_dict, + original_exception=e, + request_data=request_data, + ) + except Exception: + pass return _jsonrpc_error(body.get("id"), -32603, f"Internal error: {str(e)}", 500) diff --git a/litellm/proxy/agent_endpoints/endpoints.py b/litellm/proxy/agent_endpoints/endpoints.py index 082e314b08d..19dbfe33d32 100644 --- a/litellm/proxy/agent_endpoints/endpoints.py +++ b/litellm/proxy/agent_endpoints/endpoints.py @@ -319,36 +319,34 @@ async def create_agent( Example Request: ```bash - curl -X POST "http://localhost:4000/agents" \\ + curl -X POST "http://localhost:4000/v1/agents" \\ -H "Authorization: Bearer " \\ -H "Content-Type: application/json" \\ -d '{ - "agent": { - "agent_name": "my-custom-agent", - "agent_card_params": { - "protocolVersion": "1.0", - "name": "Hello World Agent", - "description": "Just a hello world agent", - "url": "http://localhost:9999/", - "version": "1.0.0", - "defaultInputModes": ["text"], - "defaultOutputModes": ["text"], - "capabilities": { - "streaming": true - }, - "skills": [ - { - "id": "hello_world", - "name": "Returns hello world", - "description": "just returns hello world", - "tags": ["hello world"], - "examples": ["hi", "hello world"] - } - ] + "agent_name": "my-custom-agent", + "agent_card_params": { + "protocolVersion": "1.0", + "name": "Hello World Agent", + "description": "Just a hello world agent", + "url": "http://localhost:9999/", + "version": "1.0.0", + "defaultInputModes": ["text"], + "defaultOutputModes": ["text"], + "capabilities": { + "streaming": true }, - "litellm_params": { - "make_public": true - } + "skills": [ + { + "id": "hello_world", + "name": "Returns hello world", + "description": "just returns hello world", + "tags": ["hello world"], + "examples": ["hi", "hello world"] + } + ] + }, + "litellm_params": { + "make_public": true } }' ``` @@ -441,7 +439,7 @@ async def get_agent_by_id( Example Request: ```bash - curl -X GET "http://localhost:4000/agents/123e4567-e89b-12d3-a456-426614174000" \\ + curl -X GET "http://localhost:4000/v1/agents/123e4567-e89b-12d3-a456-426614174000" \\ -H "Authorization: Bearer " ``` """ @@ -535,28 +533,26 @@ async def update_agent( Example Request: ```bash - curl -X PUT "http://localhost:4000/agents/123e4567-e89b-12d3-a456-426614174000" \\ + curl -X PUT "http://localhost:4000/v1/agents/123e4567-e89b-12d3-a456-426614174000" \\ -H "Authorization: Bearer " \\ -H "Content-Type: application/json" \\ -d '{ - "agent": { - "agent_name": "updated-agent", - "agent_card_params": { - "protocolVersion": "1.0", - "name": "Updated Agent", - "description": "Updated description", - "url": "http://localhost:9999/", - "version": "1.1.0", - "defaultInputModes": ["text"], - "defaultOutputModes": ["text"], - "capabilities": { - "streaming": true - }, - "skills": [] + "agent_name": "updated-agent", + "agent_card_params": { + "protocolVersion": "1.0", + "name": "Updated Agent", + "description": "Updated description", + "url": "http://localhost:9999/", + "version": "1.1.0", + "defaultInputModes": ["text"], + "defaultOutputModes": ["text"], + "capabilities": { + "streaming": true }, - "litellm_params": { - "make_public": false - } + "skills": [] + }, + "litellm_params": { + "make_public": false } }' ``` @@ -645,28 +641,26 @@ async def patch_agent( Example Request: ```bash - curl -X PUT "http://localhost:4000/agents/123e4567-e89b-12d3-a456-426614174000" \\ + curl -X PATCH "http://localhost:4000/v1/agents/123e4567-e89b-12d3-a456-426614174000" \\ -H "Authorization: Bearer " \\ -H "Content-Type: application/json" \\ -d '{ - "agent": { - "agent_name": "updated-agent", - "agent_card_params": { - "protocolVersion": "1.0", - "name": "Updated Agent", - "description": "Updated description", - "url": "http://localhost:9999/", - "version": "1.1.0", - "defaultInputModes": ["text"], - "defaultOutputModes": ["text"], - "capabilities": { - "streaming": true - }, - "skills": [] + "agent_name": "updated-agent", + "agent_card_params": { + "protocolVersion": "1.0", + "name": "Updated Agent", + "description": "Updated description", + "url": "http://localhost:9999/", + "version": "1.1.0", + "defaultInputModes": ["text"], + "defaultOutputModes": ["text"], + "capabilities": { + "streaming": true }, - "litellm_params": { - "make_public": false - } + "skills": [] + }, + "litellm_params": { + "make_public": false } }' ``` @@ -753,7 +747,7 @@ async def delete_agent( Example Request: ```bash - curl -X DELETE "http://localhost:4000/agents/123e4567-e89b-12d3-a456-426614174000" \\ + curl -X DELETE "http://localhost:4000/v1/agents/123e4567-e89b-12d3-a456-426614174000" \\ -H "Authorization: Bearer " ``` diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 917239e0358..93a64889458 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -546,6 +546,7 @@ async def common_checks( # noqa: PLR0915 route=route, request_headers=_safe_get_request_headers(request=request), request_query_params=_safe_get_request_query_params(request=request), + llm_router=llm_router, ) if route in MODEL_DISCOVERY_ROUTES: @@ -1126,7 +1127,7 @@ async def get_end_user_object( end_user_id: Optional[str], prisma_client: Optional[PrismaClient], user_api_key_cache: UserApiKeyCache, - route: str, + route: Optional[str] = "", parent_otel_span: Optional[Span] = None, proxy_logging_obj: Optional[ProxyLogging] = None, ) -> Optional[LiteLLM_EndUserTable]: @@ -1170,9 +1171,6 @@ async def get_end_user_object( parent_otel_span=parent_otel_span, ) - # Check budget limits - await _check_end_user_budget(end_user_obj=return_obj, route=route) - return return_obj # Fetch from database @@ -1203,14 +1201,9 @@ async def get_end_user_object( model_type=LiteLLM_EndUserTable, ) - # Check budget limits - await _check_end_user_budget(end_user_obj=_response, route=route) - return _response - except Exception as e: - if isinstance(e, litellm.BudgetExceededError): - raise e + except Exception: return None @@ -1307,8 +1300,6 @@ async def _end_user_id_exists_in_db( ) if end_user_obj is not None: return True - except litellm.BudgetExceededError: - raise except Exception as e: verbose_proxy_logger.debug( f"end_user validation: get_end_user_object lookup failed: {e}" @@ -3428,6 +3419,29 @@ async def can_team_call_search_tool( ) +async def can_user_view_search_tool( + search_tool_name: str, + valid_token: UserAPIKeyAuth, + team_object: Optional[LiteLLM_TeamTable], +) -> bool: + """ + Boolean variant of the key + team authorization enforced on /search, used to + scope /search_tools/list so a non-admin caller only sees tools it may invoke. + """ + try: + await can_key_call_search_tool( + search_tool_name=search_tool_name, + valid_token=valid_token, + ) + await can_team_call_search_tool( + search_tool_name=search_tool_name, + team_object=team_object, + ) + except ProxyException: + return False + return True + + async def is_valid_fallback_model( model: str, llm_router: Optional[Router], diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 86265270357..71cf5197dec 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -522,14 +522,20 @@ def get_request_route(request: Request) -> str: if not isinstance(scope, dict): return str(request.url.path) raw_path: str = str(scope.get("path", request.url.path)) - root_path: str = str(scope.get("app_root_path", scope.get("root_path", ""))) + root_path: str = str( + scope.get("app_root_path", scope.get("root_path", "")) + ).rstrip("/") if not isinstance(raw_path, str): return str(request.url.path) - # Only strip root_path when it is a meaningful prefix (not bare "/"). - # Stripping bare "/" would remove the leading slash from every path - # e.g. "/team/new" → "team/new", breaking route matching. - if root_path and root_path != "/" and raw_path.startswith(root_path): - return raw_path[len(root_path) :] + # Strip root_path only when it matches whole path segments — guarding + # against sibling paths like "/apifoo" being truncated under + # root_path="/api". Trailing slashes on root_path are stripped above, + # so bare "/" or "/prefix/" still leave the leading "/" intact. + if root_path and ( + raw_path == root_path or raw_path.startswith(root_path + "/") + ): + stripped = raw_path[len(root_path) :] + return stripped or "/" return raw_path except Exception as e: verbose_proxy_logger.debug( @@ -934,6 +940,40 @@ def get_team_model_tpm_limit( return None +def get_key_mcp_rpm_limit( + user_api_key_dict: UserAPIKeyAuth, +) -> Optional[Dict[str, int]]: + """ + Get the per-MCP-server rpm limit for a given api key. + + Priority order (returns first found): + 1. Key metadata (mcp_rpm_limit) + 2. Team metadata (mcp_rpm_limit) + + The returned dict is keyed by MCP server name (alias if set, else the + configured server name). + """ + if user_api_key_dict.metadata: + result = user_api_key_dict.metadata.get("mcp_rpm_limit") + if result is not None: + return result + + if user_api_key_dict.team_metadata: + team_limit = user_api_key_dict.team_metadata.get("mcp_rpm_limit") + if team_limit is not None: + return team_limit + + return None + + +def get_team_mcp_rpm_limit( + user_api_key_dict: UserAPIKeyAuth, +) -> Optional[Dict[str, int]]: + if user_api_key_dict.team_metadata: + return user_api_key_dict.team_metadata.get("mcp_rpm_limit") + return None + + def get_project_model_rpm_limit( user_api_key_dict: UserAPIKeyAuth, ) -> Optional[Dict[str, int]]: @@ -1244,7 +1284,9 @@ def _route_uses_model_routing_sources(route: str) -> bool: def _extract_models_from_managed_resource_id( - resource_id: Any, resource_id_field: Optional[str] = None + resource_id: Any, + resource_id_field: Optional[str] = None, + llm_router: Optional[Router] = None, ) -> List[str]: if not isinstance(resource_id, str) or not resource_id: return [] @@ -1301,16 +1343,18 @@ def _extract_models_from_managed_resource_id( ) if resource_id_field == "video_id": + model_id = decode_video_id_with_provider(resource_id).get("model_id") _append_model_candidates( candidates=candidates, - value=decode_video_id_with_provider(resource_id).get("model_id"), + value=_resolve_model_id_with_router(model_id, llm_router), ) else: + model_id = decode_character_id_with_provider(resource_id).get( + "model_id" + ) _append_model_candidates( candidates=candidates, - value=decode_character_id_with_provider(resource_id).get( - "model_id" - ), + value=_resolve_model_id_with_router(model_id, llm_router), ) except Exception as e: verbose_proxy_logger.debug( @@ -1320,11 +1364,26 @@ def _extract_models_from_managed_resource_id( return _dedupe_model_candidates(candidates) +def _resolve_model_id_with_router( + model_id: Optional[str], llm_router: Optional[Router] +) -> Optional[str]: + if model_id is None or llm_router is None: + return model_id + try: + return llm_router.resolve_model_name_from_model_id(model_id) or model_id + except Exception as e: + verbose_proxy_logger.debug( + "Unable to resolve model_id from managed resource ID: %s", str(e) + ) + return model_id + + def _extract_model_candidates_from_request( request_data: dict, route: str, request_headers: Optional[Mapping[str, Any]] = None, request_query_params: Optional[Mapping[str, Any]] = None, + llm_router: Optional[Router] = None, ) -> List[str]: candidates: List[str] = [] uses_model_routing_sources = _route_uses_model_routing_sources(route=route) @@ -1374,7 +1433,9 @@ def _extract_model_candidates_from_request( _append_model_candidates( candidates, _extract_models_from_managed_resource_id( - request_data.get(field), resource_id_field=field + request_data.get(field), + resource_id_field=field, + llm_router=llm_router, ), ) @@ -1396,12 +1457,14 @@ def get_model_from_request( route: str, request_headers: Optional[Mapping[str, Any]] = None, request_query_params: Optional[Mapping[str, Any]] = None, + llm_router: Optional[Router] = None, ) -> Optional[Union[str, List[str]]]: candidates = _extract_model_candidates_from_request( request_data=request_data, route=route, request_headers=request_headers, request_query_params=request_query_params, + llm_router=llm_router, ) model = _format_model_candidates(candidates) diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 9838b4ba49b..a9961586f59 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -12,12 +12,12 @@ import fnmatch import hashlib import os import re -from typing import Any, List, Literal, Optional, Set, Tuple, cast +from typing import Any, List, Literal, Optional, Set, Tuple, Union, cast from cryptography import x509 from cryptography.hazmat.backends import default_backend from cryptography.hazmat.primitives import serialization -from fastapi import HTTPException +from fastapi import HTTPException, status import jwt from jwt.api_jwk import PyJWK @@ -29,6 +29,7 @@ from litellm.proxy._types import ( RBAC_ROLES, JWKKeyValue, JWTAuthBuilderResult, + JWTIssuerConfig, JWTKeyItem, LiteLLM_EndUserTable, LiteLLM_JWTAuth, @@ -66,6 +67,10 @@ from .auth_checks import ( ) +class NoMatchingJWTPublicKeyError(Exception): + """Raised when a JWKS endpoint returns no key matching the requested ``kid``.""" + + class JWTHandler: """ - treat the sub id passed in as the user id @@ -91,6 +96,22 @@ class JWTHandler: "ES512", "EdDSA", ] + LITELLM_JWT_ISSUER_CLAIM = "_litellm_jwt_issuer" + LITELLM_USER_ID_CLAIM = "_litellm_user_id" + LITELLM_USER_EMAIL_CLAIM = "_litellm_user_email" + LITELLM_TEAM_ID_CLAIM = "_litellm_team_id" + LITELLM_TEAM_IDS_CLAIM = "_litellm_team_ids" + LITELLM_ORG_ID_CLAIM = "_litellm_org_id" + LITELLM_END_USER_ID_CLAIM = "_litellm_end_user_id" + LITELLM_INTERNAL_CLAIMS = ( + LITELLM_JWT_ISSUER_CLAIM, + LITELLM_USER_ID_CLAIM, + LITELLM_USER_EMAIL_CLAIM, + LITELLM_TEAM_ID_CLAIM, + LITELLM_TEAM_IDS_CLAIM, + LITELLM_ORG_ID_CLAIM, + LITELLM_END_USER_ID_CLAIM, + ) def __init__( self, @@ -213,7 +234,33 @@ class JWTHandler: return True return False + def _is_trusted_issuer_normalized_token(self, token: dict) -> bool: + issuer = token.get(self.LITELLM_JWT_ISSUER_CLAIM) + if not isinstance(issuer, str) or not issuer: + return False + + litellm_jwtauth = getattr(self, "litellm_jwtauth", None) + issuer_configs = getattr(litellm_jwtauth, "issuers", None) or [] + return any(issuer_config.issuer == issuer for issuer_config in issuer_configs) + + def _has_trusted_issuer_normalized_claim(self, token: dict, claim: str) -> bool: + return self._is_trusted_issuer_normalized_token(token=token) and claim in token + def get_team_ids_from_jwt(self, token: dict) -> List[str]: + if self._has_trusted_issuer_normalized_claim( + token=token, claim=self.LITELLM_TEAM_IDS_CLAIM + ): + issuer_team_ids = token.get(self.LITELLM_TEAM_IDS_CLAIM) + if isinstance(issuer_team_ids, list): + return issuer_team_ids + if isinstance(issuer_team_ids, str): + return [issuer_team_ids] + # Issuer-scoped claim exists but has an unexpected type + # (e.g. int/dict from an unusual upstream mapping). Don't silently + # fall through to the global ``team_ids_jwt_field`` path — that + # would read a semantically unrelated claim on the same token. + return [] + if self.litellm_jwtauth.team_ids_jwt_field is not None: team_ids: Optional[List[str]] = get_nested_value( data=token, @@ -242,12 +289,18 @@ class JWTHandler: default-team behavior should still go through ``get_team_id``. """ team_ids: List[str] = list(self.get_team_ids_from_jwt(token)) - if self.litellm_jwtauth.team_id_jwt_field is not None: + singular: Any = None + if self._has_trusted_issuer_normalized_claim( + token=token, claim=self.LITELLM_TEAM_ID_CLAIM + ): + singular = token.get(self.LITELLM_TEAM_ID_CLAIM) + elif self.litellm_jwtauth.team_id_jwt_field is not None: singular = get_nested_value( data=token, key_path=self.litellm_jwtauth.team_id_jwt_field, default=None, ) + if singular is not None: if isinstance(singular, list): for item in singular: if item is None: @@ -262,6 +315,11 @@ class JWTHandler: def get_end_user_id( self, token: dict, default_value: Optional[str] ) -> Optional[str]: + if self._has_trusted_issuer_normalized_claim( + token=token, claim=self.LITELLM_END_USER_ID_CLAIM + ): + return token.get(self.LITELLM_END_USER_ID_CLAIM) + try: if self.litellm_jwtauth.end_user_id_jwt_field is not None: user_id = get_nested_value( @@ -303,6 +361,14 @@ class JWTHandler: return False def get_team_id(self, token: dict, default_value: Optional[str]) -> Optional[str]: + if self._has_trusted_issuer_normalized_claim( + token=token, claim=self.LITELLM_TEAM_ID_CLAIM + ): + team_id = token.get(self.LITELLM_TEAM_ID_CLAIM) + if isinstance(team_id, list): + return team_id[0] if team_id else default_value + return team_id + try: if self.litellm_jwtauth.team_id_jwt_field is not None: # Use a sentinel value to detect if the path actually exists @@ -376,6 +442,11 @@ class JWTHandler: return self.litellm_jwtauth.user_id_upsert def get_user_id(self, token: dict, default_value: Optional[str]) -> Optional[str]: + if self._has_trusted_issuer_normalized_claim( + token=token, claim=self.LITELLM_USER_ID_CLAIM + ): + return token.get(self.LITELLM_USER_ID_CLAIM) + try: if self.litellm_jwtauth.user_id_jwt_field is not None: user_id = get_nested_value( @@ -467,6 +538,11 @@ class JWTHandler: def get_user_email( self, token: dict, default_value: Optional[str] ) -> Optional[str]: + if self._has_trusted_issuer_normalized_claim( + token=token, claim=self.LITELLM_USER_EMAIL_CLAIM + ): + return token.get(self.LITELLM_USER_EMAIL_CLAIM) + try: if self.litellm_jwtauth.user_email_jwt_field is not None: user_email = get_nested_value( @@ -495,6 +571,11 @@ class JWTHandler: return object_id def get_org_id(self, token: dict, default_value: Optional[str]) -> Optional[str]: + if self._has_trusted_issuer_normalized_claim( + token=token, claim=self.LITELLM_ORG_ID_CLAIM + ): + return token.get(self.LITELLM_ORG_ID_CLAIM) + try: if self.litellm_jwtauth.org_id_jwt_field is not None: org_id = get_nested_value( @@ -590,55 +671,77 @@ class JWTHandler: await self.user_api_key_cache.async_set_cache( key=cache_key, value=jwks_uri, - ttl=self.litellm_jwtauth.public_key_ttl, + ttl=self._get_public_key_cache_ttl(), ) return jwks_uri + def _get_public_key_cache_ttl(self) -> float: + litellm_jwtauth = getattr(self, "litellm_jwtauth", None) + if litellm_jwtauth is None: + return 600 + return litellm_jwtauth.public_key_ttl + + async def _get_public_key_from_jwks_url( + self, jwks_url: str, kid: Optional[str] + ) -> dict: + resolved_jwks_url = await self._resolve_jwks_url(jwks_url) + cache_key = f"litellm_jwt_auth_keys_{resolved_jwks_url}" + + cached_keys = await self.user_api_key_cache.async_get_cache(cache_key) + + if cached_keys is None: + response = await self.http_handler.get(resolved_jwks_url) + + try: + response_json = response.json() + except Exception as e: + verbose_proxy_logger.error( + f"Error parsing response: {e}. Original Response: {response.text}" + ) + raise Exception( + f"Error parsing response: {e}. Check server logs for original response." + ) + + if "keys" in response_json: + keys: JWKKeyValue = response_json["keys"] + else: + keys = response_json + + await self.user_api_key_cache.async_set_cache( + key=cache_key, + value=keys, + ttl=self._get_public_key_cache_ttl(), + ) + else: + keys = cached_keys + + public_key = self.parse_keys(keys=keys, kid=kid) + if public_key is not None: + return cast(dict, public_key) + + raise NoMatchingJWTPublicKeyError( + f"No matching public key found. keys={resolved_jwks_url}, kid={kid}" + ) + async def get_public_key(self, kid: Optional[str]) -> dict: keys_url = os.getenv("JWT_PUBLIC_KEY_URL") if keys_url is None: raise Exception("Missing JWT Public Key URL from environment.") - keys_url_list = [url.strip() for url in keys_url.split(",")] + keys_url_list = [url.strip() for url in keys_url.split(",") if url.strip()] for key_url in keys_url_list: - key_url = await self._resolve_jwks_url(key_url) - cache_key = f"litellm_jwt_auth_keys_{key_url}" - - cached_keys = await self.user_api_key_cache.async_get_cache(cache_key) - - if cached_keys is None: - response = await self.http_handler.get(key_url) - - try: - response_json = response.json() - except Exception as e: - verbose_proxy_logger.error( - f"Error parsing response: {e}. Original Response: {response.text}" - ) - raise Exception( - f"Error parsing response: {e}. Check server logs for original response." - ) - - if "keys" in response_json: - keys: JWKKeyValue = response.json()["keys"] - else: - keys = response_json - - await self.user_api_key_cache.async_set_cache( - key=cache_key, - value=keys, - ttl=self.litellm_jwtauth.public_key_ttl, # cache for 10 mins + try: + return await self._get_public_key_from_jwks_url( + jwks_url=key_url, kid=kid + ) + except NoMatchingJWTPublicKeyError as e: + verbose_proxy_logger.debug( + "JWT Auth: No matching public key found at %s: %s", key_url, e ) - else: - keys = cached_keys - public_key = self.parse_keys(keys=keys, kid=kid) - if public_key is not None: - return cast(dict, public_key) - - raise Exception( + raise NoMatchingJWTPublicKeyError( f"No matching public key found. keys={keys_url_list}, kid={kid}" ) @@ -753,6 +856,11 @@ class JWTHandler: minted by other applications that share the same IdP signing keys. When both are unset PyJWT only checks the signature and expiry, which is preserved for backward compatibility but logged once as a warning. + + The warning fires even in mixed deployments that also configure + ``LiteLLM_JWTAuth.issuers``: tokens whose ``iss`` does not match any + configured issuer fall through to this global path, and if env-var + scoping is absent that fallback is itself unscoped. """ audience = os.getenv("JWT_AUDIENCE") issuer = os.getenv("JWT_ISSUER") @@ -782,77 +890,230 @@ class JWTHandler: "options": options or None, } - async def auth_jwt(self, token: str) -> dict: - decode_kwargs = self._build_decode_kwargs() + def _get_configured_issuer(self, token: str) -> Optional[JWTIssuerConfig]: + litellm_jwtauth = getattr(self, "litellm_jwtauth", None) + if litellm_jwtauth is None: + return None + issuer_configs = litellm_jwtauth.issuers + if not issuer_configs: + return None + + claims = self.get_unverified_claims(token=token) + if claims is None: + return None + + issuer = claims.get("iss") + if not isinstance(issuer, str) or not issuer: + return None + + for issuer_config in issuer_configs: + if issuer_config.issuer == issuer: + return issuer_config + + return None + + def _get_jwks_url_for_issuer(self, issuer_config: JWTIssuerConfig) -> str: + if issuer_config.jwks_url: + return issuer_config.jwks_url + # _resolve_jwks_url fetches this OIDC discovery document and follows + # its jwks_uri, matching JWTIssuerConfig.jwks_url's documented fallback. + return f"{issuer_config.issuer.rstrip('/')}/.well-known/openid-configuration" + + def _get_claim_value_for_issuer_mapping(self, token: dict, claim_field: str) -> Any: + """Resolve a mapped claim from ``token``. + + Returns ``None`` when the field is absent or empty so that mapped claims + behave like the global ``litellm_jwtauth`` path — present claims override + the normalised value, missing ones simply leave it ``None``. + """ + sentinel = object() + claim_value = get_nested_value( + data=token, + key_path=claim_field, + default=sentinel, + ) + if claim_value is sentinel or claim_value is None or claim_value == "": + return None + return claim_value + + def _apply_issuer_claim_mappings( + self, token: dict, issuer_config: JWTIssuerConfig + ) -> dict: + normalized: dict = { + k: v for k, v in token.items() if k not in self.LITELLM_INTERNAL_CLAIMS + } + normalized[self.LITELLM_JWT_ISSUER_CLAIM] = issuer_config.issuer + claim_mappings = [ + (issuer_config.user_id_jwt_field, self.LITELLM_USER_ID_CLAIM), + (issuer_config.user_email_jwt_field, self.LITELLM_USER_EMAIL_CLAIM), + (issuer_config.team_id_jwt_field, self.LITELLM_TEAM_ID_CLAIM), + (issuer_config.team_ids_jwt_field, self.LITELLM_TEAM_IDS_CLAIM), + (issuer_config.org_id_jwt_field, self.LITELLM_ORG_ID_CLAIM), + (issuer_config.end_user_id_jwt_field, self.LITELLM_END_USER_ID_CLAIM), + ] + + for source_claim, normalized_claim in claim_mappings: + if source_claim is None: + continue + claim_value = self._get_claim_value_for_issuer_mapping( + token=token, + claim_field=source_claim, + ) + if claim_value is not None: + normalized[normalized_claim] = claim_value + + return normalized + + def _get_jwk_from_public_key(self, public_key: dict) -> dict: + jwk = {} + for key in ["kty", "kid", "n", "e", "x", "y", "crv"]: + if key in public_key: + jwk[key] = public_key[key] + return jwk + + def _get_decode_options( + self, + audience: Optional[Union[str, List[str]]], + issuer: Optional[str] = None, + disable_audience_validation: bool = False, + ) -> Optional[dict]: + # Disabling audience verification must be an explicit choice — never + # an implicit consequence of ``audience`` being None. Otherwise a + # caller that accidentally constructs a config with ``audience=None`` + # (bypassing the model validator) would silently lose audience + # validation. Require callers to opt in via + # ``disable_audience_validation=True``. + if audience is None and not disable_audience_validation: + raise ValueError( + "audience must be provided unless disable_audience_validation=True" + ) + options: dict = {} + if audience is None: + options["verify_aud"] = False + if issuer is None: + options["verify_iss"] = False + return options or None + + def _decode_jwt_with_public_key( + self, + token: str, + public_key: Union[dict, str], + audience: Optional[Union[str, List[str]]], + issuer: Optional[str] = None, + options: Optional[dict] = None, + disable_audience_validation: bool = False, + ) -> dict: + decode_options = ( + options + if options is not None + else self._get_decode_options( + audience=audience, + issuer=issuer, + disable_audience_validation=disable_audience_validation, + ) + ) + + if isinstance(public_key, dict): + public_key_obj = PyJWK.from_dict( + self._get_jwk_from_public_key(public_key=public_key) + ).key + return jwt.decode( + token, + public_key_obj, # type: ignore + algorithms=self.SUPPORTED_JWT_ALGORITHMS, + options=decode_options, # type: ignore[arg-type] + audience=audience, + issuer=issuer, + leeway=self.leeway, + ) + + cert = x509.load_pem_x509_certificate(public_key.encode(), default_backend()) + key = cert.public_key().public_bytes( + serialization.Encoding.PEM, + serialization.PublicFormat.SubjectPublicKeyInfo, + ) + return jwt.decode( + token, + key, + algorithms=self.SUPPORTED_JWT_ALGORITHMS, + audience=audience, + issuer=issuer, + options=decode_options, # type: ignore[arg-type] + leeway=self.leeway, + ) + + async def _auth_jwt_with_issuer( + self, token: str, issuer_config: JWTIssuerConfig, kid: Optional[str] + ) -> dict: + public_key = await self._get_public_key_from_jwks_url( + jwks_url=self._get_jwks_url_for_issuer(issuer_config=issuer_config), + kid=kid, + ) + try: + payload = self._decode_jwt_with_public_key( + token=token, + public_key=public_key, + audience=issuer_config.audience, + issuer=issuer_config.issuer, + disable_audience_validation=issuer_config.disable_audience_validation, + ) + except jwt.ExpiredSignatureError: + raise ProxyException( + message="Token Expired", + type=ProxyErrorTypes.expired_key, + param=None, + code=status.HTTP_401_UNAUTHORIZED, + ) + except Exception as e: + raise Exception(f"Validation fails: {str(e)}") + + return self._apply_issuer_claim_mappings( + token=payload, + issuer_config=issuer_config, + ) + + async def auth_jwt(self, token: str) -> dict: header = jwt.get_unverified_header(token) verbose_proxy_logger.debug("header: %s", header) kid = header.get("kid", None) + issuer_config = self._get_configured_issuer(token=token) + if issuer_config is not None: + return await self._auth_jwt_with_issuer( + token=token, + issuer_config=issuer_config, + kid=kid, + ) + + decode_kwargs = self._build_decode_kwargs() + public_key = await self.get_public_key(kid=kid) - if public_key is not None and isinstance(public_key, dict): - jwk = {} - if "kty" in public_key: - jwk["kty"] = public_key["kty"] - if "kid" in public_key: - jwk["kid"] = public_key["kid"] - if "n" in public_key: - jwk["n"] = public_key["n"] - if "e" in public_key: - jwk["e"] = public_key["e"] - if "x" in public_key: - jwk["x"] = public_key["x"] - if "y" in public_key: - jwk["y"] = public_key["y"] - if "crv" in public_key: - jwk["crv"] = public_key["crv"] - - # parse RSA/EC/OKP keys - public_key_obj = PyJWK.from_dict(jwk).key - + if public_key is not None: try: - # decode the token using the public key - payload = jwt.decode( - token, - public_key_obj, # type: ignore - algorithms=self.SUPPORTED_JWT_ALGORITHMS, - leeway=self.leeway, # allow testing of expired tokens - **decode_kwargs, + payload = self._decode_jwt_with_public_key( + token=token, + public_key=public_key, + audience=decode_kwargs["audience"], + issuer=decode_kwargs["issuer"], + options=decode_kwargs["options"], ) - return payload + return { + k: v + for k, v in payload.items() + if k not in self.LITELLM_INTERNAL_CLAIMS + } except jwt.ExpiredSignatureError: - # the token is expired, do something to refresh it - raise Exception("Token Expired") - except Exception as e: - raise Exception(f"Validation fails: {str(e)}") - elif public_key is not None and isinstance(public_key, str): - try: - cert = x509.load_pem_x509_certificate( - public_key.encode(), default_backend() + raise ProxyException( + message="Token Expired", + type=ProxyErrorTypes.expired_key, + param=None, + code=status.HTTP_401_UNAUTHORIZED, ) - - # Extract public key - key = cert.public_key().public_bytes( - serialization.Encoding.PEM, - serialization.PublicFormat.SubjectPublicKeyInfo, - ) - - # decode the token using the public key - payload = jwt.decode( - token, - key, - algorithms=self.SUPPORTED_JWT_ALGORITHMS, - **decode_kwargs, - ) - return payload - - except jwt.ExpiredSignatureError: - # the token is expired, do something to refresh it - raise Exception("Token Expired") except Exception as e: raise Exception(f"Validation fails: {str(e)}") diff --git a/litellm/proxy/auth/ip_address_utils.py b/litellm/proxy/auth/ip_address_utils.py index 39d3282942f..be0d83dfcdc 100644 --- a/litellm/proxy/auth/ip_address_utils.py +++ b/litellm/proxy/auth/ip_address_utils.py @@ -153,8 +153,9 @@ class IPAddressUtils: verbose_proxy_logger.warning( "use_x_forwarded_for is enabled but mcp_trusted_proxy_ranges " "is not configured. X-Forwarded-* headers will NOT be " - "trusted, so MCP OAuth discovery URLs will use the proxy's " - "literal base URL. Set mcp_trusted_proxy_ranges in " + "trusted, so MCP OAuth discovery URLs and access-control " + "client IPs will use the proxy's literal request values. " + "Set mcp_trusted_proxy_ranges in " "general_settings to your reverse-proxy CIDR(s) to allow " "X-Forwarded-* through." ) @@ -199,17 +200,19 @@ class IPAddressUtils: # If XFF is enabled, validate the request comes from a trusted proxy if use_xff and "x-forwarded-for" in request.headers: - trusted_ranges = general_settings.get("mcp_trusted_proxy_ranges") - if trusted_ranges: - # Validate direct connection is from trusted proxy + if not IPAddressUtils.is_request_from_trusted_proxy( + request, general_settings=general_settings + ): direct_ip = request.client.host if request.client else None - trusted_networks = IPAddressUtils.parse_trusted_proxy_networks( - trusted_ranges - ) - if not IPAddressUtils.is_trusted_proxy(direct_ip, trusted_networks): - # Untrusted source trying to set XFF - ignore XFF, use direct IP + if general_settings.get("mcp_trusted_proxy_ranges"): + # Direct connection isn't in any configured trusted CIDR. verbose_proxy_logger.warning( "XFF header from untrusted IP %s, ignoring", direct_ip ) return direct_ip + # XFF enabled but no trusted proxy ranges configured: the direct + # peer is typically the reverse proxy's own (private) IP, so + # returning it would mis-classify external callers as internal. + # Fail closed for access control. + return "" return _get_request_ip_address(request, use_x_forwarded_for=use_xff) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 2f9a1411400..9d4efbaeeee 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -30,6 +30,7 @@ from litellm.proxy._types import * from litellm.proxy.auth.auth_checks import ( ExperimentalUIJWTToken, _cache_key_object, + _check_end_user_budget, _delete_cache_key_object, _get_user_role, _is_model_cost_zero, @@ -146,12 +147,14 @@ def _get_model_from_request_context( request_data: dict, route: str, request: Optional[Request], + llm_router: Optional[Any] = None, ) -> Optional[Union[str, List[str]]]: return get_model_from_request( request_data=request_data, route=route, request_headers=_safe_get_request_headers(request=request), request_query_params=_safe_get_request_query_params(request=request), + llm_router=llm_router, ) @@ -1034,6 +1037,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 request_data=request_data, route=route, request=request, + llm_router=llm_router, ) skip_budget_checks = False if model is not None and llm_router is not None: @@ -1451,6 +1455,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 request_data=request_data, route=route, request=request, + llm_router=llm_router, ) skip_budget_checks = False if model is not None and llm_router is not None: @@ -1579,6 +1584,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 request_data=request_data, route=route, request=request, + llm_router=llm_router, ) current_models = _get_model_names_for_budget_checks( model=current_model @@ -1757,8 +1763,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 async def _safe_fetch(label: str, awaitable): """Run an awaitable and return its result. Re-raises authentication / authorization failures (HTTPException, ProxyException, - BudgetExceededError — which ``get_end_user_object`` raises for - end-user budget violations) so they propagate to the caller. + BudgetExceededError) so they propagate to the caller. Other exceptions (e.g. transient DB errors fetching context) are swallowed with a debug log and ``None`` is returned so ``common_checks`` can still run against whatever limits are recorded @@ -2159,6 +2164,7 @@ def _should_skip_budget_checks( request_data=request_data, route=route, request=request, + llm_router=llm_router, ) if model is not None and llm_router is not None: return _is_model_cost_zero(model=model, llm_router=llm_router) @@ -2475,6 +2481,7 @@ async def _enforce_key_and_fallback_model_access( request_data=request_data, route=route, request=request, + llm_router=llm_router, ) if model is not None: @@ -2577,6 +2584,14 @@ async def _run_post_custom_auth_checks( user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, ) + # common_checks() enforces the end-user budget, but the centralized + # gate skips it for custom-auth deployments unless + # custom_auth_run_common_checks is set. Enforce it here on that path + # so an over-budget end user can't keep making requests. + if end_user_object is not None and not general_settings.get( + "custom_auth_run_common_checks", False + ): + await _check_end_user_budget(end_user_obj=end_user_object, route=route) # 2. Check token expiry if valid_token.expires is not None: @@ -2616,6 +2631,7 @@ async def _run_post_custom_auth_checks( request_data=request_data, route=route, request=request, + llm_router=llm_router, ) current_models = _get_model_names_for_budget_checks(model=current_model) diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index ba3aee75047..c109f374993 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -2,6 +2,7 @@ import asyncio import copy import logging import os +import secrets import time import traceback from datetime import datetime, timedelta @@ -39,6 +40,7 @@ from litellm.proxy.health_check import ( from litellm.proxy.middleware.in_flight_requests_middleware import ( get_in_flight_requests, ) +from litellm.proxy.shutdown.graceful_shutdown_manager import GracefulShutdownManager #### Health ENDPOINTS #### @@ -1551,6 +1553,50 @@ def _allow_public_health_readiness_details() -> bool: return general_settings.get("allow_public_health_readiness_details") is True +def _drain_endpoint_enabled() -> bool: + from litellm.proxy.proxy_server import general_settings + + return general_settings.get("enable_drain_endpoint") is True + + +def _drain_endpoint_token() -> Optional[str]: + """ + Shared secret required on the X-Drain-Token header to call /health/drain. + + Falls back to the ``DRAIN_ENDPOINT_TOKEN`` env var when unset in + general_settings so the kubelet preStop hook can supply it via + ``valueFrom.secretKeyRef`` without a config reload. + """ + from litellm.proxy.proxy_server import general_settings + + token = general_settings.get("drain_endpoint_token") + if isinstance(token, str) and token: + return token + env_token = os.getenv("DRAIN_ENDPOINT_TOKEN") + if env_token: + return env_token + return None + + +def _authorize_drain_request(request: Request) -> None: + """ + Reject /health/drain calls that don't carry the configured X-Drain-Token. + + When no token is configured the endpoint is treated as already opted-in + (the ``enable_drain_endpoint`` flag is the only gate). Comparison uses + ``secrets.compare_digest`` to avoid timing leaks. + """ + expected = _drain_endpoint_token() + if expected is None: + return + supplied = request.headers.get("x-drain-token") or "" + if not secrets.compare_digest(supplied, expected): + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail="Invalid or missing X-Drain-Token", + ) + + async def _resolve_public_readiness_db(response: Response) -> str: """ Return the db status string for the public probe and flip the response to @@ -1580,6 +1626,10 @@ async def health_readiness(response: Response): credential. Admins can opt into the legacy detailed payload with general_settings.allow_public_health_readiness_details. """ + if GracefulShutdownManager.is_shutting_down(): + response.status_code = status.HTTP_503_SERVICE_UNAVAILABLE + return {"status": "shutting_down"} + if _allow_public_health_readiness_details(): return await _get_health_readiness_details(response=response) @@ -1616,6 +1666,54 @@ async def health_backlog(): return {"in_flight_requests": get_in_flight_requests()} +@router.get( + "/health/drain", + tags=["health"], +) +async def health_drain(request: Request): + """ + Graceful-drain probe for Kubernetes ``preStop`` hooks. + + Disabled by default and returns 404 unless ``general_settings`` sets + ``enable_drain_endpoint: true``. Calling it flips a process-wide + shutting-down flag, so a successful call permanently takes the worker out + of rotation until the pod restarts. + + Because the kubelet calls preStop hooks without proxy credentials, the + endpoint does not require ``user_api_key_auth``. To prevent any + pod-reachable caller from triggering shutdown, set + ``general_settings.drain_endpoint_token`` (or the ``DRAIN_ENDPOINT_TOKEN`` + env var) and supply the same value on the ``X-Drain-Token`` header from + the preStop hook. Calls without the header (or with a wrong value) get a + 401 and have no side effect. + + When enabled, it marks the worker as shutting down (so /health/readiness + and /health/liveliness immediately start returning 503, removing the pod + from service) and blocks until the in-flight request counter drains to + zero or ``GRACEFUL_SHUTDOWN_TIMEOUT`` elapses. Unlike a fixed ``sleep``, + this returns as soon as real in-flight work is done. + + Wire it up as: + + ```yaml + lifecycle: + preStop: + httpGet: + path: /health/drain + port: 4000 + httpHeaders: + - name: X-Drain-Token + value: + ``` + """ + if not _drain_endpoint_enabled(): + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Not Found") + _authorize_drain_request(request) + GracefulShutdownManager.start_shutdown() + drained = await GracefulShutdownManager.wait_for_drain(exclude_self=True) + return {"status": "drained", "drained_requests": drained} + + @router.get( "/health/liveliness", # Historical LiteLLM name; doesn't match k8s terminology but kept for backwards compatibility tags=["health"], @@ -1624,10 +1722,16 @@ async def health_backlog(): "/health/liveness", # Kubernetes has "liveness" probes (https://kubernetes.io/docs/tasks/configure-pod-container/configure-liveness-readiness-startup-probes/#define-a-liveness-command) tags=["health"], ) -async def health_liveliness(): +async def health_liveliness(response: Response): """ - Unprotected endpoint for checking if worker is alive + Unprotected endpoint for checking if worker is alive. + + Returns 503 once graceful shutdown has begun so Kubernetes stops counting + the draining pod as live and terminates it on schedule. """ + if GracefulShutdownManager.is_shutting_down(): + response.status_code = status.HTTP_503_SERVICE_UNAVAILABLE + return {"status": "shutting_down"} return "I'm alive!" diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index d03ad70562a..4343747d104 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -36,7 +36,7 @@ from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.auth_utils import get_model_rate_limit_from_metadata from litellm.types.caching import RedisPipelineIncrementOperation from litellm.types.llms.openai import BaseLiteLLMOpenAIResponseObject -from litellm.types.utils import ModelResponse, Usage +from litellm.types.utils import CallTypes, ModelResponse, Usage if TYPE_CHECKING: from opentelemetry.trace import Span as _Span @@ -1375,6 +1375,79 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) ) + def _add_mcp_per_key_rate_limit_descriptor( + self, + user_api_key_dict: UserAPIKeyAuth, + mcp_server_name: Optional[str], + descriptors: List[RateLimitDescriptor], + ) -> None: + """ + Add a per-MCP-server rpm descriptor for the API key, if a limit is + configured for the server being called. + + MCP tool calls have no token usage, so only requests_per_unit is set; + tokens_per_unit stays None so the TPM reservation path is never engaged. + """ + from litellm.proxy.auth.auth_utils import get_key_mcp_rpm_limit + + if not mcp_server_name or not user_api_key_dict.api_key: + return + + mcp_rpm_limit = get_key_mcp_rpm_limit(user_api_key_dict) + if not mcp_rpm_limit: + return + + server_rpm_limit = mcp_rpm_limit.get(mcp_server_name) + if server_rpm_limit is None: + return + + descriptors.append( + RateLimitDescriptor( + key="mcp_per_key", + value=f"{user_api_key_dict.api_key}:{mcp_server_name}", + rate_limit={ + "requests_per_unit": server_rpm_limit, + "tokens_per_unit": None, + "window_size": self.window_size, + }, + ) + ) + + def _add_mcp_per_team_rate_limit_descriptor( + self, + user_api_key_dict: UserAPIKeyAuth, + mcp_server_name: Optional[str], + descriptors: List[RateLimitDescriptor], + ) -> None: + """ + Add a per-MCP-server rpm descriptor for the team, if a limit is + configured for the server being called. + """ + from litellm.proxy.auth.auth_utils import get_team_mcp_rpm_limit + + if not mcp_server_name or not user_api_key_dict.team_id: + return + + mcp_rpm_limit = get_team_mcp_rpm_limit(user_api_key_dict) + if not mcp_rpm_limit: + return + + server_rpm_limit = mcp_rpm_limit.get(mcp_server_name) + if server_rpm_limit is None: + return + + descriptors.append( + RateLimitDescriptor( + key="mcp_per_team", + value=f"{user_api_key_dict.team_id}:{mcp_server_name}", + rate_limit={ + "requests_per_unit": server_rpm_limit, + "tokens_per_unit": None, + "window_size": self.window_size, + }, + ) + ) + def _should_enforce_rate_limit( self, limit_type: Optional[str], @@ -1533,6 +1606,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): rpm_limit_type: Optional[str], tpm_limit_type: Optional[str], model_has_failures: bool, + call_type: Optional[str] = None, ) -> List[RateLimitDescriptor]: """ Create all rate limit descriptors for the request. @@ -1653,6 +1727,21 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): descriptors=descriptors, ) + # REST MCP calls pass the raw body through this hook before server + # resolution; only the later synthetic hook payload may carry this key. + if call_type == CallTypes.call_mcp_tool.value and "server_id" not in data: + mcp_server_name = data.get("mcp_server_name", None) + self._add_mcp_per_key_rate_limit_descriptor( + user_api_key_dict=user_api_key_dict, + mcp_server_name=mcp_server_name, + descriptors=descriptors, + ) + self._add_mcp_per_team_rate_limit_descriptor( + user_api_key_dict=user_api_key_dict, + mcp_server_name=mcp_server_name, + descriptors=descriptors, + ) + if ( get_team_model_rpm_limit(user_api_key_dict) is not None or get_team_model_tpm_limit(user_api_key_dict) is not None @@ -1983,6 +2072,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): rpm_limit_type=rpm_limit_type, tpm_limit_type=tpm_limit_type, model_has_failures=model_has_failures, + call_type=call_type, ) # Add team model rate limits from team_metadata diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 75eb5cd55ef..7b8f0f72e13 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -386,6 +386,7 @@ async def new_user( - soft_budget: Optional[float] - Get alerts when user crosses given budget, doesn't block requests. - model_max_budget: Optional[dict] - Model-specific max budget for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-budgets-to-keys) - model_rpm_limit: Optional[float] - Model-specific rpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys) + - mcp_rpm_limit: Optional[dict] - Per-MCP-server rpm limit, keyed by MCP server name {"github": 100, "slack": 200}. Enforced for keys and teams only; values set on a user are stored but not enforced per user. - model_tpm_limit: Optional[float] - Model-specific tpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys) - spend: Optional[float] - Amount spent by user. Default is 0. Will be updated by proxy whenever user is used. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"), months ("1mo"). - agent_id: Optional[str] - The agent id associated with the user. @@ -1427,6 +1428,7 @@ async def user_update( - soft_budget: Optional[float] - Get alerts when user crosses given budget, doesn't block requests. - model_max_budget: Optional[dict] - Model-specific max budget for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-budgets-to-keys) - model_rpm_limit: Optional[float] - Model-specific rpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys) + - mcp_rpm_limit: Optional[dict] - Per-MCP-server rpm limit, keyed by MCP server name {"github": 100, "slack": 200}. Enforced for keys and teams only; values set on a user are stored but not enforced per user. - model_tpm_limit: Optional[float] - Model-specific tpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys) - spend: Optional[float] - Amount spent by user. Default is 0. Will be updated by proxy whenever user is used. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"), months ("1mo"). - agent_id: Optional[str] - The agent id associated with the user. diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 0e645013b92..cf90f0661b3 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -894,7 +894,12 @@ async def _common_key_generation_helper( # noqa: PLR0915 user_api_key_dict.user_role is not None and user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value ) - if not _is_proxy_admin: + _org_inherited_from_team = ( + team_table is not None + and team_table.organization_id is not None + and data.organization_id == team_table.organization_id + ) + if not _is_proxy_admin and not _org_inherited_from_team: await _validate_caller_can_assign_key_org( user_api_key_dict=user_api_key_dict, organization_id=data.organization_id, @@ -1388,6 +1393,7 @@ async def generate_key_fn( - model_max_budget: Optional[Dict[str, BudgetConfig]] - Model-specific budgets {"gpt-4": {"budget_limit": 0.0005, "time_period": "30d"}}}. IF null or {} then no model specific budget. - model_rpm_limit: Optional[dict] - key-specific model rpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific rpm limit. - model_tpm_limit: Optional[dict] - key-specific model tpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific tpm limit. + - mcp_rpm_limit: Optional[dict] - key-specific per-MCP-server rpm limit, keyed by MCP server name (alias if set, else the configured name). Example - {"github": 100, "slack": 200}. IF null or {} then no MCP-specific rpm limit. - tpm_limit_type: Optional[str] - Type of tpm limit. Options: "best_effort_throughput" (no error if we're overallocating tpm), "guaranteed_throughput" (raise an error if we're overallocating tpm), "dynamic" (dynamically exceed limit when no 429 errors). Defaults to "best_effort_throughput". - rpm_limit_type: Optional[str] - Type of rpm limit. Options: "best_effort_throughput" (no error if we're overallocating rpm), "guaranteed_throughput" (raise an error if we're overallocating rpm), "dynamic" (dynamically exceed limit when no 429 errors). Defaults to "best_effort_throughput". - allowed_cache_controls: Optional[list] - List of allowed cache control values. Example - ["no-cache", "no-store"]. See all values - https://docs.litellm.ai/docs/proxy/caching#turn-on--off-caching-per-request @@ -1606,6 +1612,7 @@ async def generate_service_account_key_fn( - model_max_budget: Optional[Dict[str, BudgetConfig]] - Model-specific budgets {"gpt-4": {"budget_limit": 0.0005, "time_period": "30d"}}}. IF null or {} then no model specific budget. - model_rpm_limit: Optional[dict] - key-specific model rpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific rpm limit. - model_tpm_limit: Optional[dict] - key-specific model tpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific tpm limit. + - mcp_rpm_limit: Optional[dict] - key-specific per-MCP-server rpm limit, keyed by MCP server name (alias if set, else the configured name). Example - {"github": 100, "slack": 200}. IF null or {} then no MCP-specific rpm limit. - tpm_limit_type: Optional[str] - TPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic" - rpm_limit_type: Optional[str] - RPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic" - allowed_cache_controls: Optional[list] - List of allowed cache control values. Example - ["no-cache", "no-store"]. See all values - https://docs.litellm.ai/docs/proxy/caching#turn-on--off-caching-per-request @@ -2422,6 +2429,7 @@ async def update_key_fn( # noqa: PLR0915 - tpm_limit: Optional[int] - Tokens per minute limit - rpm_limit: Optional[int] - Requests per minute limit - model_rpm_limit: Optional[dict] - Model-specific RPM limits {"gpt-4": 100, "claude-v1": 200} + - mcp_rpm_limit: Optional[dict] - Per-MCP-server RPM limits, keyed by MCP server name {"github": 100, "slack": 200} - model_tpm_limit: Optional[dict] - Model-specific TPM limits {"gpt-4": 100000, "claude-v1": 200000} - tpm_limit_type: Optional[str] - TPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic" - rpm_limit_type: Optional[str] - RPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic" @@ -3401,6 +3409,7 @@ async def generate_key_helper_fn( # noqa: PLR0915 model_max_budget: Optional[dict] = {}, model_rpm_limit: Optional[dict] = None, model_tpm_limit: Optional[dict] = None, + mcp_rpm_limit: Optional[dict] = None, guardrails: Optional[list] = None, policies: Optional[list] = None, prompts: Optional[list] = None, @@ -3479,6 +3488,9 @@ async def generate_key_helper_fn( # noqa: PLR0915 if model_tpm_limit is not None: metadata = metadata or {} metadata["model_tpm_limit"] = model_tpm_limit + if mcp_rpm_limit is not None: + metadata = metadata or {} + metadata["mcp_rpm_limit"] = mcp_rpm_limit if guardrails is not None: metadata = metadata or {} metadata["guardrails"] = guardrails diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index b35e2b6e3fd..4d67df16b0b 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -1542,7 +1542,7 @@ if MCP_AVAILABLE: master_key, algorithms=["HS256"], # UI session cookies may omit exp; don't require it. - options={"verify_exp": False}, + options={"verify_exp": False, "verify_aud": False}, ) if decoded.get("login_method") in ("sso", "username_password"): cookie_key = decoded.get("key", "") diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 8a8e703831b..ae7da0d29f2 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -863,8 +863,9 @@ async def new_team( # noqa: PLR0915 - members_with_roles: List[{"role": "admin" or "user", "user_id": ""}] - A list of users and their roles in the team. Get user_id when making a new user via `/user/new`. - team_member_permissions: Optional[List[str]] - A list of routes that non-admin team members can access. example: ["/key/generate", "/key/update", "/key/delete"] - metadata: Optional[dict] - Metadata for team, store information for team. Example metadata = {"extra_info": "some info"} - - model_rpm_limit: Optional[Dict[str, int]] - The RPM (Requests Per Minute) limit for this team - applied across all keys for this team. + - model_rpm_limit: Optional[Dict[str, int]] - The RPM (Requests Per Minute) limit for this team - applied across all keys for this team. - model_tpm_limit: Optional[Dict[str, int]] - The TPM (Tokens Per Minute) limit for this team - applied across all keys for this team. + - mcp_rpm_limit: Optional[Dict[str, int]] - Per-MCP-server RPM limit for this team, keyed by MCP server name (alias if set, else the configured name). Example: {"github": 100, "slack": 200}. Applied across all keys for this team. - tpm_limit: Optional[int] - The TPM (Tokens Per Minute) limit for this team - all keys with this team_id will have at max this TPM limit - rpm_limit: Optional[int] - The RPM (Requests Per Minute) limit for this team - all keys associated with this team_id will have at max this RPM limit - rpm_limit_type: Optional[Literal["guaranteed_throughput", "best_effort_throughput"]] - The type of RPM limit enforcement. Use "guaranteed_throughput" to raise an error if overallocating RPM, or "best_effort_throughput" for best effort enforcement. diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 9d68132b37d..985785ad77e 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -667,6 +667,34 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): return stream +def _carry_guardrail_logging_info( + request_data: dict, guardrail_data: Optional[dict] +) -> None: + """Copy guardrail logging entries from ``guardrail_data`` onto ``request_data``. + + Post-call guardrails run against a throwaway ``hook_data`` dict (its + ``metadata`` is what ``_init_kwargs_for_pass_through_endpoint`` already + stripped off ``_parsed_body``), so a block records the + ``standard_logging_guardrail_information`` there and not on the dict the + failure handler forwards to ``post_call_failure_hook``. Without this the + otel guardrail span is emitted on allow but missing on block. Carry the + entries over so the failure path matches the unified path. + """ + if guardrail_data is None: + return + source_metadata = guardrail_data.get("metadata") + if not isinstance(source_metadata, dict): + return + entries = source_metadata.get("standard_logging_guardrail_information") + if not entries: + return + + metadata = request_data.get("metadata") + if not isinstance(metadata, dict): + metadata = request_data["metadata"] = {} + metadata.setdefault("standard_logging_guardrail_information", list(entries)) + + async def pass_through_request( # noqa: PLR0915 request: Request, target: str, @@ -718,6 +746,9 @@ async def pass_through_request( # noqa: PLR0915 # kwargs for pass through endpoint, contains metadata, litellm_params, call_type, litellm_call_id, passthrough_logging_payload kwargs: Optional[dict] = None logging_obj: Optional[Logging] = None + # the dict post-call guardrails wrote their logging info into; the failure + # handler reuses it so a guardrail block still surfaces its span/logs + post_call_guardrail_data: Optional[dict] = None ######################################################### try: @@ -1160,6 +1191,7 @@ async def pass_through_request( # noqa: PLR0915 **existing_metadata, "guardrails": guardrails_to_run, } + post_call_guardrail_data = hook_data response_body = await proxy_logging_obj.post_call_success_hook( data=hook_data, user_api_key_dict=user_api_key_dict, @@ -1343,6 +1375,8 @@ async def pass_through_request( # noqa: PLR0915 if "custom_llm_provider" not in request_payload and custom_llm_provider: request_payload["custom_llm_provider"] = custom_llm_provider + _carry_guardrail_logging_info(request_payload, post_call_guardrail_data) + await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index e0f139dee57..a092df6cdf0 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -417,6 +417,7 @@ from litellm.proxy.middleware.prometheus_auth_middleware import PrometheusAuthMi from litellm.proxy.middleware.request_size_limit_middleware import ( RequestSizeLimitMiddleware, ) +from litellm.proxy.shutdown.graceful_shutdown_manager import GracefulShutdownManager from litellm.proxy.ocr_endpoints.endpoints import router as ocr_router from litellm.proxy.openai_files_endpoints.files_endpoints import ( router as openai_files_router, @@ -973,6 +974,11 @@ async def proxy_startup_event(app: FastAPI): # noqa: PLR0915 # End of startup event yield + # Shutdown event - drain in-flight requests before tearing down dependencies + # so SIGTERM (rolling update, scale-down, liveness kill) doesn't drop them. + GracefulShutdownManager.start_shutdown() + await GracefulShutdownManager.wait_for_drain() + # Shutdown event - close shared aiohttp session if shared_aiohttp_session is not None: try: @@ -1265,7 +1271,7 @@ async def openai_exception_handler(request: Request, exc: ProxyException): headers = exc.headers error_dict = exc.to_dict() status_code = int(exc.code) if exc.code else status.HTTP_500_INTERNAL_SERVER_ERROR - _close_dangling_otel_server_span(request, status_code) + _close_dangling_otel_server_span(request, status_code, exc=exc) return JSONResponse( status_code=status_code, content={"error": error_dict}, @@ -1273,7 +1279,9 @@ async def openai_exception_handler(request: Request, exc: ProxyException): ) -def _close_dangling_otel_server_span(request: Request, status_code: int) -> None: +def _close_dangling_otel_server_span( + request: Request, status_code: int, exc: Optional[Exception] = None +) -> None: parent_otel_span = getattr(request.state, "parent_otel_span", None) if parent_otel_span is None: return @@ -1296,6 +1304,10 @@ def _close_dangling_otel_server_span(request: Request, status_code: int) -> None open_telemetry_logger.set_response_status_code_attribute( parent_otel_span, status_code ) + if status_code >= 400: + open_telemetry_logger.record_error_attributes_on_span( + parent_otel_span, exc, status_code + ) parent_otel_span.set_status( Status(StatusCode.ERROR if status_code >= 400 else StatusCode.OK) ) @@ -1312,7 +1324,7 @@ def _close_dangling_otel_server_span(request: Request, status_code: int) -> None async def otel_request_validation_exception_handler( request: Request, exc: RequestValidationError ): - _close_dangling_otel_server_span(request, 422) + _close_dangling_otel_server_span(request, 422, exc=exc) return JSONResponse( status_code=422, content={"detail": jsonable_encoder(exc.errors())}, @@ -1326,7 +1338,7 @@ async def otel_unhandled_exception_handler(request: Request, exc: Exception): verbose_proxy_logger.exception( "Unhandled exception in request: %s", type(exc).__name__ ) - _close_dangling_otel_server_span(request, 500) + _close_dangling_otel_server_span(request, 500, exc=exc) return JSONResponse( status_code=500, content={ @@ -15845,6 +15857,8 @@ async def _mcp_forward_as_path(path_segment: str, request: Request): ) scope = dict(request.scope) + # Preserve the public request path for OAuth challenge URL selection. + scope["_original_path"] = scope.get("path", "") scope["path"] = f"/mcp/{path_segment}" return await _stream_mcp_asgi_response( handle_streamable_http_mcp, scope, request.receive @@ -15992,6 +16006,7 @@ async def dynamic_mcp_route(mcp_server_name: str, request: Request): ) if toolset is not None: scope = dict(request.scope) + scope["_original_path"] = scope.get("path", "") scope["path"] = "/mcp" token = _mcp_active_toolset_id.set(toolset.toolset_id) try: diff --git a/litellm/proxy/public_endpoints/agent_create_fields.json b/litellm/proxy/public_endpoints/agent_create_fields.json index e58bd97cce7..36484cc1065 100644 --- a/litellm/proxy/public_endpoints/agent_create_fields.json +++ b/litellm/proxy/public_endpoints/agent_create_fields.json @@ -7,6 +7,48 @@ "credential_fields": [], "litellm_params_template": {} }, + { + "agent_type": "langflow", + "agent_type_display_name": "LangFlow", + "description": "Connect to LangFlow AI agents via the LangFlow Platform API", + "logo_url": "/ui/assets/logos/langflow.svg", + "model_template": "langflow/{flow_id}", + "credential_fields": [ + { + "key": "flow_id", + "label": "Flow ID", + "placeholder": "your-flow-id", + "tooltip": "The Flow ID from your LangFlow deployment (found in the flow URL or settings)", + "required": true, + "field_type": "text", + "default_value": null, + "include_in_litellm_params": false + }, + { + "key": "api_base", + "label": "LangFlow API Base", + "placeholder": "http://localhost:7860", + "tooltip": "The base URL for your LangFlow server (e.g., http://localhost:7860 or your deployed LangFlow URL)", + "required": true, + "field_type": "text", + "default_value": "http://localhost:7860", + "include_in_litellm_params": true + }, + { + "key": "api_key", + "label": "LangFlow API Key", + "placeholder": null, + "tooltip": "API key for authenticating with your LangFlow server (x-api-key header)", + "required": false, + "field_type": "password", + "default_value": null, + "include_in_litellm_params": true + } + ], + "litellm_params_template": { + "custom_llm_provider": "langflow" + } + }, { "agent_type": "langgraph", "agent_type_display_name": "LangGraph", diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 78143fe0411..c4754ef6117 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -325,6 +325,7 @@ model LiteLLM_MCPServerTable { allow_all_keys Boolean @default(false) available_on_public_internet Boolean @default(true) delegate_auth_to_upstream Boolean @default(false) + oauth_passthrough Boolean @default(false) is_byok Boolean @default(false) byok_description String[] @default([]) byok_api_key_help_url String? diff --git a/litellm/proxy/search_endpoints/search_tool_management.py b/litellm/proxy/search_endpoints/search_tool_management.py index 725e83bf96d..5642fcd10c3 100644 --- a/litellm/proxy/search_endpoints/search_tool_management.py +++ b/litellm/proxy/search_endpoints/search_tool_management.py @@ -3,12 +3,17 @@ CRUD ENDPOINTS FOR SEARCH TOOLS """ from datetime import datetime -from typing import Any, Dict, List, Union +from typing import Any, Dict, List, Optional, Union from fastapi import APIRouter, Depends, HTTPException from pydantic import BaseModel from litellm._logging import verbose_proxy_logger +from litellm.proxy._types import ( + LiteLLM_TeamTable, + LitellmUserRoles, + UserAPIKeyAuth, +) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.search_endpoints.search_tool_registry import SearchToolRegistry from litellm.types.search import ( @@ -41,13 +46,61 @@ def _convert_datetime_to_str(value: Union[datetime, str, None]) -> Union[str, No return value +async def _filter_visible_search_tools( + search_tools: List[SearchToolInfoResponse], + user_api_key_dict: UserAPIKeyAuth, +) -> List[SearchToolInfoResponse]: + """ + Drop search tools the caller is not authorized to invoke, applying the same + key/team object_permission allowlists enforced on /search. Admins see all tools. + """ + if user_api_key_dict.user_role in ( + LitellmUserRoles.PROXY_ADMIN, + LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, + ): + return search_tools + + from litellm.proxy.auth.auth_checks import ( + can_user_view_search_tool, + get_team_object, + ) + from litellm.proxy.proxy_server import ( + prisma_client, + proxy_logging_obj, + user_api_key_cache, + ) + + team_object: Optional[LiteLLM_TeamTable] = None + if user_api_key_dict.team_id: + team_object = await get_team_object( + team_id=user_api_key_dict.team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=user_api_key_dict.parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + ) + + visible: List[SearchToolInfoResponse] = [] + for tool in search_tools: + tool_name = tool.get("search_tool_name") + if tool_name and await can_user_view_search_tool( + search_tool_name=tool_name, + valid_token=user_api_key_dict, + team_object=team_object, + ): + visible.append(tool) + return visible + + @router.get( "/search_tools/list", tags=["Search Tools"], dependencies=[Depends(user_api_key_auth)], response_model=ListSearchToolsResponse, ) -async def list_search_tools(): +async def list_search_tools( + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): """ List all search tools that are available in the database and config file. @@ -114,22 +167,25 @@ async def list_search_tools(): f"Could not get config-defined search tools: {e}" ) - for search_tool in config_search_tools: - tool_name = search_tool.get("search_tool_name") + for config_search_tool in config_search_tools: + tool_name = config_search_tool.get("search_tool_name") if tool_name: - litellm_params_dict = dict(search_tool.get("litellm_params", {})) + litellm_params_dict = dict(config_search_tool.get("litellm_params", {})) masked_litellm_params_dict = _get_masked_values( litellm_params_dict, unmasked_length=4, number_of_asterisks=4, ) + config_tool_info = config_search_tool.get("search_tool_info") search_tool_configs.append( SearchToolInfoResponse( search_tool_id=None, search_tool_name=tool_name, litellm_params=masked_litellm_params_dict, - search_tool_info=search_tool.get("search_tool_info"), + search_tool_info=( + dict(config_tool_info) if config_tool_info else None + ), created_at=None, updated_at=None, is_from_config=True, @@ -142,8 +198,8 @@ async def list_search_tools(): if tool.get("search_tool_name") not in db_tool_names ] - for search_tool in search_tools_from_db: - litellm_params_dict = dict(search_tool.get("litellm_params", {})) + for db_search_tool in search_tools_from_db: + litellm_params_dict = dict(db_search_tool.get("litellm_params", {})) masked_litellm_params_dict = _get_masked_values( litellm_params_dict, unmasked_length=4, @@ -152,17 +208,25 @@ async def list_search_tools(): search_tool_configs.append( SearchToolInfoResponse( - search_tool_id=search_tool.get("search_tool_id"), - search_tool_name=search_tool.get("search_tool_name", ""), + search_tool_id=db_search_tool.get("search_tool_id"), + search_tool_name=db_search_tool.get("search_tool_name", ""), litellm_params=masked_litellm_params_dict, - search_tool_info=search_tool.get("search_tool_info"), - created_at=_convert_datetime_to_str(search_tool.get("created_at")), - updated_at=_convert_datetime_to_str(search_tool.get("updated_at")), + search_tool_info=db_search_tool.get("search_tool_info"), + created_at=_convert_datetime_to_str( + db_search_tool.get("created_at") + ), + updated_at=_convert_datetime_to_str( + db_search_tool.get("updated_at") + ), is_from_config=False, ) ) - return ListSearchToolsResponse(search_tools=search_tool_configs) + visible_search_tools = await _filter_visible_search_tools( + search_tool_configs, user_api_key_dict + ) + + return ListSearchToolsResponse(search_tools=visible_search_tools) except Exception as e: verbose_proxy_logger.exception(f"Error getting search tools: {e}") raise HTTPException(status_code=500, detail=str(e)) diff --git a/litellm/proxy/shutdown/__init__.py b/litellm/proxy/shutdown/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/proxy/shutdown/graceful_shutdown_manager.py b/litellm/proxy/shutdown/graceful_shutdown_manager.py new file mode 100644 index 00000000000..20ffabdd7d8 --- /dev/null +++ b/litellm/proxy/shutdown/graceful_shutdown_manager.py @@ -0,0 +1,174 @@ +""" +Application-level graceful shutdown coordination for the LiteLLM proxy. + +Kubernetes terminates a pod by sending ``SIGTERM`` and, after +``terminationGracePeriodSeconds``, ``SIGKILL``. By default LiteLLM delegates +the signal to uvicorn and tears down immediately, dropping any in-flight +requests (streaming, batch inference, long-lived calls). + +A fixed ``preStop`` sleep can not solve this: it has to be sized for the +*worst-case* request, so it either wastes time on every routine shutdown or is +too short for a long-running request. This manager instead drains based on the +*actual* in-flight request counter (already tracked by +``InFlightRequestsMiddleware``), so a pod terminates as soon as its real +in-flight work is done — and never waits longer than ``GRACEFUL_SHUTDOWN_TIMEOUT``. + +The state is process-scoped (class-level), matching the per-uvicorn-worker +granularity of ``InFlightRequestsMiddleware``. +""" + +import asyncio +import os +import time +from typing import Callable, Optional + +from litellm._logging import verbose_proxy_logger +from litellm.proxy.middleware.in_flight_requests_middleware import ( + get_in_flight_requests, +) + +# Keep below terminationGracePeriodSeconds so the process exits before SIGKILL. +DEFAULT_GRACEFUL_SHUTDOWN_TIMEOUT = 30.0 +_DRAIN_POLL_INTERVAL = 0.1 +_DRAIN_LOG_INTERVAL = 5.0 + + +class GracefulShutdownManager: + """ + Process-scoped singleton that tracks whether the worker is draining and + blocks until in-flight requests reach zero (or a timeout elapses). + """ + + _is_shutting_down: bool = False + _shutdown_started_at: Optional[float] = None + _drain_performed: bool = False + + @classmethod + def is_shutting_down(cls) -> bool: + """Whether this worker has begun graceful shutdown.""" + return cls._is_shutting_down + + @classmethod + def get_timeout(cls) -> float: + """ + Read GRACEFUL_SHUTDOWN_TIMEOUT (seconds) from the environment on each + call so deployments can tune it without code changes. Falls back to the + default on an unset or malformed value. + """ + raw = os.getenv("GRACEFUL_SHUTDOWN_TIMEOUT") + if raw is None: + return DEFAULT_GRACEFUL_SHUTDOWN_TIMEOUT + try: + return float(raw) + except (TypeError, ValueError): + verbose_proxy_logger.warning( + "GRACEFUL_SHUTDOWN_TIMEOUT=%r is not a number; using default %ss", + raw, + DEFAULT_GRACEFUL_SHUTDOWN_TIMEOUT, + ) + return DEFAULT_GRACEFUL_SHUTDOWN_TIMEOUT + + @classmethod + def start_shutdown(cls) -> None: + """ + Mark the worker as draining. Idempotent — repeated calls (e.g. SIGTERM + followed by a preStop hit on /health/drain) do not reset the clock. + """ + if cls._is_shutting_down: + return + cls._is_shutting_down = True + cls._shutdown_started_at = time.monotonic() + verbose_proxy_logger.info( + "graceful_shutdown_started in_flight_requests=%s", + get_in_flight_requests(), + ) + + @classmethod + async def wait_for_drain( + cls, + timeout: Optional[float] = None, + exclude_self: bool = False, + count_fn: Optional[Callable[[], int]] = None, + poll_interval: float = _DRAIN_POLL_INTERVAL, + log_interval: float = _DRAIN_LOG_INTERVAL, + ) -> int: + """ + Poll the in-flight request counter until it reaches the drain target or + ``timeout`` seconds elapse. + + Args: + timeout: Max seconds to wait. Defaults to ``get_timeout()``. + exclude_self: When the caller is itself an in-flight HTTP request + (the /health/drain endpoint), set this so the caller's own + request is not counted as outstanding work. + count_fn: Source of the current in-flight count. Defaults to the + live ``InFlightRequestsMiddleware`` counter; injectable for tests. + poll_interval: Seconds between counter polls. + log_interval: Minimum seconds between ``drain_waiting`` log lines. + + Returns: + Number of requests that drained while waiting (>= 0). + """ + # A preStop /health/drain hook and the lifespan SIGTERM handler both + # drain; once one has run, the other must not wait again, otherwise the + # effective window is 2x the timeout and terminationGracePeriodSeconds + # has to be doubled to avoid a mid-drain SIGKILL. + if cls._drain_performed: + return 0 + cls._drain_performed = True + + if timeout is None: + timeout = cls.get_timeout() + if count_fn is None: + count_fn = get_in_flight_requests + + # The /health/drain HTTP request flows through InFlightRequestsMiddleware + # and so counts itself; treat <=1 as "drained" in that case. + target = 1 if exclude_self else 0 + + start = time.monotonic() + initial = count_fn() + last_log = start + + if timeout <= 0: + return max(0, initial - target) + + while True: + current = count_fn() + if current <= target: + drained = max(0, initial - current) + verbose_proxy_logger.info( + "graceful_shutdown_complete drained_requests=%s elapsed_s=%.2f", + drained, + time.monotonic() - start, + ) + return drained + + elapsed = time.monotonic() - start + if elapsed >= timeout: + verbose_proxy_logger.warning( + "graceful_shutdown_timeout in_flight_requests=%s elapsed_s=%.2f " + "timeout_s=%s — proceeding with teardown", + current, + elapsed, + timeout, + ) + return max(0, initial - current) + + now = time.monotonic() + if now - last_log >= log_interval: + verbose_proxy_logger.info( + "drain_waiting in_flight_requests=%s elapsed_s=%.2f", + current, + elapsed, + ) + last_log = now + + await asyncio.sleep(poll_interval) + + @classmethod + def reset(cls) -> None: + """Reset state. Intended for use in tests.""" + cls._is_shutting_down = False + cls._shutdown_started_at = None + cls._drain_performed = False diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 200a17e2368..eb8af3b073e 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -72,7 +72,7 @@ async def reserve_budget_for_request( return None if route in {"/models", "/v1/models", "/utils/token_counter"}: return None - if get_model_from_request(request_body, route) is None: + if get_model_from_request(request_body, route, llm_router=llm_router) is None: return None counters = await _get_budget_counters( @@ -797,7 +797,7 @@ def estimate_request_max_cost( route: str, llm_router: Optional[Router], ) -> Optional[float]: - model = get_model_from_request(request_body, route) + model = get_model_from_request(request_body, route, llm_router=llm_router) if model is None: return None diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index e2881faca0d..d215294fd04 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -24,7 +24,7 @@ from litellm.litellm_core_utils.core_helpers import ( get_litellm_metadata_from_kwargs, reconstruct_model_name, ) -from litellm.litellm_core_utils.safe_json_dumps import safe_dumps +from litellm.litellm_core_utils.safe_json_dumps import safe_dumps, strip_null_bytes from litellm.proxy._types import SpendLogsMetadata, SpendLogsPayload from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error from litellm.proxy.utils import PrismaClient, hash_token @@ -304,7 +304,7 @@ def get_logging_payload( # noqa: PLR0915 # BUG FIX: Don't overwrite api_key when standard_logging_payload is None # The api_key was already extracted from metadata (line 243) and hashed (lines 256-259) request_tags = ( - json.dumps(metadata.get("tags", [])) + safe_dumps(metadata.get("tags", [])) if isinstance(metadata.get("tags", []), list) else "[]" ) @@ -312,7 +312,7 @@ def get_logging_payload( # noqa: PLR0915 standard_logging_payload is not None and standard_logging_payload.get("request_tags") is not None ): # use 'tags' from standard logging payload instead - request_tags = json.dumps(standard_logging_payload["request_tags"]) + request_tags = safe_dumps(standard_logging_payload["request_tags"]) _model_id = metadata.get("model_info", {}).get("id", "") _model_group = metadata.get("model_group", "") @@ -606,7 +606,7 @@ def _get_messages_for_spend_logs_payload( messages = standard_logging_payload.get("messages") if messages is not None: try: - return json.dumps(messages, default=str) + return safe_dumps(messages) except Exception: return "{}" return "{}" @@ -976,7 +976,7 @@ def _get_proxy_server_request_for_spend_logs_payload( perform_redaction(model_call_details=_request_body, result=None) _request_body = _sanitize_request_body_for_spend_logs_payload(_request_body) - _request_body_json_str = json.dumps(_request_body, default=str) + _request_body_json_str = safe_dumps(_request_body) if LITELLM_TRUNCATED_PAYLOAD_FIELD in _request_body_json_str: verbose_proxy_logger.info( "Spend Log: request body was truncated before storing in DB. %s", @@ -1059,7 +1059,7 @@ def _get_response_for_spend_logs_payload( if sanitized_response is None: return "{}" if isinstance(sanitized_response, str): - result_str = sanitized_response + result_str = strip_null_bytes(sanitized_response) else: result_str = safe_dumps(sanitized_response) if LITELLM_TRUNCATED_PAYLOAD_FIELD in result_str: diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 0e72f47e224..8bd50a50a38 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -643,6 +643,7 @@ class ProxyLogging: "user_api_key_request_route": kwargs.get("user_api_key_request_route"), "mcp_tool_name": request_obj.tool_name, # Keep original for reference "mcp_arguments": request_obj.arguments, # Keep original for reference + "mcp_server_name": kwargs.get("mcp_rate_limit_server_name"), # Raw Bearer token from the original HTTP request — allows guardrails # (e.g. MCPJWTSigner) to independently verify the caller's identity # before re-signing an outbound token (FR-5 verify+re-sign). diff --git a/litellm/router.py b/litellm/router.py index d60c39ca402..7aaf989919c 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -9100,7 +9100,10 @@ class Router: except Exception: pass + # Three mutually exclusive scenarios for the model's metadata: if custom_model_info is not None and litellm_model_name_model_info is not None: + # (1) It has both custom model_info set and exists in the built-in map + # merge with custom overriding built-in model_info = cast( ModelInfo, _update_dictionary( @@ -9109,7 +9112,12 @@ class Router: ), ) elif litellm_model_name_model_info is not None: + # (2) Built-in only — no custom pricing to merge model_info = litellm_model_name_model_info + elif custom_model_info is not None: + # (3) Custom only — model not in built-in cost map yet + # custom_model_info already includes base_model defaults at this point, if applicable + model_info = cast(ModelInfo, custom_model_info) return model_info diff --git a/litellm/types/agents.py b/litellm/types/agents.py index 8556b6bac93..f34631b5600 100644 --- a/litellm/types/agents.py +++ b/litellm/types/agents.py @@ -298,6 +298,23 @@ class MakeAgentsPublicRequest(BaseModel): agent_ids: List[str] +def _normalize_a2a_jsonrpc_response( + response_dict: Dict[str, Any], + request_id: Optional[Any] = None, +) -> Dict[str, Any]: + """ + Ensure JSON-RPC responses include ``id`` when the caller supplied one. + + The a2a SDK may omit ``id`` on error payloads even when the upstream agent + returned it. Backfill from the outbound request id so LiteLLM can surface the + agent error instead of failing Pydantic validation. + """ + normalized = dict(response_dict) + if normalized.get("id") is None and request_id is not None: + normalized["id"] = str(request_id) + return normalized + + class LiteLLMSendMessageResponse(LiteLLMPydanticObjectBase): """ LiteLLM wrapper for A2A SendMessageResponse. @@ -322,31 +339,42 @@ class LiteLLMSendMessageResponse(LiteLLMPydanticObjectBase): @classmethod def from_a2a_response( - cls, response: "SendMessageResponse" + cls, + response: "SendMessageResponse", + request_id: Optional[Any] = None, ) -> "LiteLLMSendMessageResponse": """ Create a LiteLLMSendMessageResponse from an a2a SDK SendMessageResponse. Args: response: The a2a SDK SendMessageResponse + request_id: JSON-RPC request id to backfill when the SDK omits it on errors Returns: LiteLLMSendMessageResponse with _hidden_params support """ - # Convert the a2a response to a dict response_dict = response.model_dump(mode="json", exclude_none=True) - + response_dict = _normalize_a2a_jsonrpc_response( + response_dict, request_id=request_id + ) return cls(**response_dict) @classmethod - def from_dict(cls, response_dict: Dict[str, Any]) -> "LiteLLMSendMessageResponse": + def from_dict( + cls, + response_dict: Dict[str, Any], + request_id: Optional[Any] = None, + ) -> "LiteLLMSendMessageResponse": """ Create a LiteLLMSendMessageResponse from a dict. Args: response_dict: Dict with A2A response structure + request_id: JSON-RPC request id to backfill when missing on error payloads Returns: LiteLLMSendMessageResponse with _hidden_params support """ - return cls(**response_dict) + return cls( + **_normalize_a2a_jsonrpc_response(response_dict, request_id=request_id) + ) diff --git a/litellm/types/images/main.py b/litellm/types/images/main.py index 819f4954589..80e55297c42 100644 --- a/litellm/types/images/main.py +++ b/litellm/types/images/main.py @@ -20,6 +20,7 @@ class ImageEditOptionalRequestParams(TypedDict, total=False): response_format: Optional[Literal["url", "b64_json"]] size: Optional[str] user: Optional[str] + imageConfig: Optional[Dict[str, Any]] class ImageEditRequestParams(ImageEditOptionalRequestParams, total=False): diff --git a/litellm/types/interactions/generated.py b/litellm/types/interactions/generated.py index d546e897891..b38cd8f58b9 100644 --- a/litellm/types/interactions/generated.py +++ b/litellm/types/interactions/generated.py @@ -203,6 +203,8 @@ class Status1(Enum): completed = "completed" failed = "failed" cancelled = "cancelled" + incomplete = "incomplete" + budget_exceeded = "budget_exceeded" class InteractionStatusUpdate(BaseModel): @@ -386,13 +388,13 @@ class ResponseModality(Enum): class Status3(Enum): - UNSPECIFIED = "UNSPECIFIED" - IN_PROGRESS = "IN_PROGRESS" - REQUIRES_ACTION = "REQUIRES_ACTION" - COMPLETED = "COMPLETED" - FAILED = "FAILED" - CANCELLED = "CANCELLED" - INCOMPLETE = "INCOMPLETE" + IN_PROGRESS = "in_progress" + REQUIRES_ACTION = "requires_action" + COMPLETED = "completed" + FAILED = "failed" + CANCELLED = "cancelled" + INCOMPLETE = "incomplete" + BUDGET_EXCEEDED = "budget_exceeded" class ModelOption(RootModel[str]): diff --git a/litellm/types/llms/gemini.py b/litellm/types/llms/gemini.py index 8763544facc..38e6d533449 100644 --- a/litellm/types/llms/gemini.py +++ b/litellm/types/llms/gemini.py @@ -1,7 +1,7 @@ from enum import Enum -from typing import Any, Dict, Iterable, List, Literal, Optional, Union +from typing import Any, Dict, List, Literal, Optional -from typing_extensions import Required, TypedDict +from typing_extensions import TypedDict from .vertex_ai import ( GenerationConfig, @@ -171,6 +171,9 @@ class GeminiImageGenerationParameters(BaseModel): aspectRatio: Optional[str] = None """Aspect ratio for generated images (e.g., '1:1', '16:9', '9:16', '4:3', '3:4')""" + imageSize: Optional[str] = None + """Image size for generated images (e.g., '1K', '2K')""" + personGeneration: Optional[str] = None """Controls person generation in images""" diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 346909f14eb..51e408f8c63 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -1084,6 +1084,7 @@ OpenAIImageGenerationOptionalParams = Literal[ "image_url", "image_prompt_strength", "aspect_ratio", + "imageConfig", ] OpenAIImageEditOptionalParams = Literal[ diff --git a/litellm/types/llms/vertex_ai.py b/litellm/types/llms/vertex_ai.py index a1d53978761..b972ff3c538 100644 --- a/litellm/types/llms/vertex_ai.py +++ b/litellm/types/llms/vertex_ai.py @@ -20,6 +20,7 @@ class FunctionResponse(TypedDict, total=False): id: str name: Required[str] response: Optional[dict] + parts: List["FunctionResponsePartType"] class FunctionCall(TypedDict, total=False): @@ -40,6 +41,11 @@ class BlobType(TypedDict, total=False): data: Required[str] +class FunctionResponsePartType(TypedDict, total=False): + inline_data: BlobType + file_data: FileDataType + + class PartType(TypedDict, total=False): text: str inline_data: BlobType diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 13e325838dc..6aa62c35106 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -68,12 +68,29 @@ class MCPServer(BaseModel): access_groups: Optional[List[str]] = None allow_all_keys: bool = False available_on_public_internet: bool = True - # When True AND auth_type == oauth2, MCP requests targeting this server + # Explicit opt-in to upstream-delegated authentication for ``oauth2`` + # servers. When ``auth_type == oauth2`` and this is ``True``, MCP requests # bypass LiteLLM API-key/SSO auth (and the pre-emptive 401) so the client - # completes PKCE directly with the upstream MCP server. Honored only for - # auth_type=oauth2; ignored for any other auth_type. See - # MCPRequestHandler._target_servers_delegate_auth_to_upstream. + # completes PKCE directly with the upstream MCP server. See + # ``MCPRequestHandler._target_servers_delegate_auth_to_upstream``. + # + # Honored only for ``auth_type == oauth2``; ignored for any other + # ``auth_type``. OAuth pass-through for non-oauth2 servers + # (``auth_type in (None, MCPAuth.none)``) is a separate, explicit opt-in — + # see ``oauth_passthrough`` / ``is_oauth_passthrough``. delegate_auth_to_upstream: bool = False + # Explicit opt-in to OAuth pass-through for non-oauth2 servers. When this + # is ``True`` AND ``auth_type in (None, MCPAuth.none)`` AND ``extra_headers`` + # contains ``Authorization``, the gateway proxies upstream + # ``/.well-known/oauth-protected-resource`` metadata, emits spec-compliant + # 401 challenges when no bearer is supplied, and propagates upstream + # 401/403 responses instead of swallowing them. See ``is_oauth_passthrough``. + # + # Intentionally distinct from ``delegate_auth_to_upstream`` (oauth2-only): + # reusing that flag would silently change behavior for servers that forward + # ``Authorization`` for non-OAuth reasons (e.g. static bearer tokens). Must + # be set explicitly to avoid regressing servers that did not opt in. + oauth_passthrough: bool = False is_byok: bool = False byok_description: List[str] = [] byok_api_key_help_url: Optional[str] = None @@ -139,6 +156,42 @@ class MCPServer(BaseModel): return False + @property + def is_oauth_passthrough(self) -> bool: + """True iff the gateway should transparently forward upstream OAuth + (discovery + 401s) rather than participating as an authorization + server itself. + + A server is pass-through for OAuth purposes when ALL three conditions + hold: + 1. ``auth_type`` is ``None`` or ``MCPAuth.none`` (the gateway does + not manage OAuth for this server). + 2. ``extra_headers`` includes ``Authorization`` — the admin has + opted this server into forwarding the client's bearer token + straight to the upstream MCP server. + 3. ``oauth_passthrough`` is ``True`` — the admin has + explicitly opted into upstream-delegated OAuth semantics for + this server. This is the explicit detection flag: without it, + a server that merely forwards ``Authorization`` (e.g. for + static bearer tokens or custom auth schemes) keeps the + pre-PR behavior and is not treated as OAuth pass-through. + This is deliberately a separate flag from + ``delegate_auth_to_upstream`` (which is oauth2-only) so enabling + pass-through here never changes behavior for oauth2 servers. + + This is intentionally narrower than ``requires_per_user_auth``, + which also covers PATs (``x-api-key``, ``api-key``, ``apikey``). + Those are static credentials, not OAuth bearer tokens, so they + must not trigger upstream OAuth discovery or 401 propagation. + """ + if self.auth_type not in (None, MCPAuth.none): + return False + if not self.extra_headers: + return False + if self.oauth_passthrough is not True: + return False + return any(h.lower() == "authorization" for h in self.extra_headers) + @property def has_token_exchange_config(self) -> bool: """True if this server is configured for OAuth2 token exchange (OBO / RFC 8693).""" diff --git a/litellm/types/utils.py b/litellm/types/utils.py index d208eb29ac1..3dcff2be689 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -148,6 +148,7 @@ class ProviderSpecificModelInfo(TypedDict, total=False): supports_xhigh_reasoning_effort: Optional[bool] supports_max_reasoning_effort: Optional[bool] supports_output_config: Optional[bool] + supports_image_size: Optional[bool] bedrock_output_config_effort_ceiling: Optional[ Literal["low", "medium", "high", "max", "xhigh"] ] @@ -2580,6 +2581,12 @@ class StandardLoggingMCPToolCall(TypedDict, total=False): Cost per query for the MCP server tool call """ + mcp_session_id: Optional[str] + """ + The MCP `mcp-session-id` of the stateful session this tool call ran in, when + the client is driving a stateful session. Absent for stateless calls. + """ + class StandardLoggingVectorStoreRequest(TypedDict, total=False): """ @@ -3294,6 +3301,8 @@ class LlmProviders(str, Enum): V0 = "v0" MORPH = "morph" LAMBDA_AI = "lambda_ai" + INCEPTION = "inception" + TEXT_COMPLETION_INCEPTION = "text-completion-inception" DEEPSEEK = "deepseek" SAMBANOVA = "sambanova" MARITALK = "maritalk" @@ -3358,6 +3367,7 @@ class LlmProviders(str, Enum): AMAZON_NOVA = "amazon_nova" A2A_AGENT = "a2a_agent" LANGGRAPH = "langgraph" + LANGFLOW = "langflow" MINIMAX = "minimax" SYNTHETIC = "synthetic" APERTIS = "apertis" diff --git a/litellm/types/vector_stores.py b/litellm/types/vector_stores.py index ce247fc900f..6adfbf4fd35 100644 --- a/litellm/types/vector_stores.py +++ b/litellm/types/vector_stores.py @@ -112,6 +112,66 @@ class VectorStoreSearchRequest(VectorStoreSearchOptionalRequestParams, total=Fal query: Union[str, List[str]] +class VertexSearchDataStoreExtraBody(TypedDict, total=False): + """ + Native Discovery Engine ``SearchRequest`` fields callers may forward via + ``extra_body`` when searching a Vertex AI Search **data store** serving + config (``.../dataStores/{id}/servingConfigs/default_config``). + + The data store is scoped by the request URL path, so target-selecting + fields (``servingConfig``, ``branch``, ``entity``) are intentionally + omitted and rejected by the transformation layer. Engine/app-only fields + such as ``dataStoreSpecs`` and ``numResultsPerDataStore`` live on + ``VertexSearchEngineExtraBody`` instead. + """ + + query: str + pageSize: int + pageToken: str + offset: int + oneBoxPageSize: int + pageCategories: List[str] + imageQuery: Dict[str, Any] + filter: str + canonicalFilter: str + orderBy: str + userInfo: Dict[str, Any] + languageCode: str + facetSpecs: List[Dict[str, Any]] + boostSpec: Dict[str, Any] + params: Dict[str, Any] + queryExpansionSpec: Dict[str, Any] + spellCorrectionSpec: Dict[str, Any] + userPseudoId: str + contentSearchSpec: Dict[str, Any] + rankingExpression: str + rankingExpressionBackend: str + safeSearch: bool + userLabels: Dict[str, str] + naturalLanguageQueryUnderstandingSpec: Dict[str, Any] + searchAsYouTypeSpec: Dict[str, Any] + displaySpec: Dict[str, Any] + crowdingSpecs: List[Dict[str, Any]] + relevanceThreshold: str + relevanceScoreSpec: Dict[str, Any] + customRankingParams: Dict[str, Any] + + +class VertexSearchEngineExtraBody(VertexSearchDataStoreExtraBody, total=False): + """ + Native Discovery Engine ``SearchRequest`` fields callers may forward via + ``extra_body`` when searching a Vertex AI Search **engine/app** serving + config (``.../engines/{id}/servingConfigs/default_serving_config``). + + Inherits every data-store field and adds fields that only make sense when + an app fans out across multiple member data stores, e.g. ``dataStoreSpecs`` + (per-store scoping/filtering) and ``numResultsPerDataStore``. + """ + + dataStoreSpecs: List[Dict[str, Any]] + numResultsPerDataStore: int + + # Vector Store Creation Types class VectorStoreExpirationPolicy(TypedDict, total=False): """The expiration policy for a vector store""" diff --git a/litellm/utils.py b/litellm/utils.py index 68982ea2b35..7cac830b2c2 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -3147,6 +3147,7 @@ def get_optional_params_image_gen( size: Optional[str] = None, style: Optional[str] = None, user: Optional[str] = None, + imageConfig: Optional[dict] = None, custom_llm_provider: Optional[str] = None, additional_drop_params: Optional[list] = None, provider_config: Optional[BaseImageGenerationConfig] = None, @@ -3183,6 +3184,7 @@ def get_optional_params_image_gen( "size": None, "style": None, "user": None, + "imageConfig": None, } non_default_params = _get_non_default_params( @@ -4547,6 +4549,18 @@ def get_optional_params( # noqa: PLR0915 ), ) + elif custom_llm_provider == "text-completion-inception": + optional_params = litellm.InceptionTextCompletionConfig().map_openai_params( + non_default_params=non_default_params, + optional_params=optional_params, + model=model, + drop_params=( + drop_params + if drop_params is not None and isinstance(drop_params, bool) + else False + ), + ) + elif custom_llm_provider == "databricks": optional_params = litellm.DatabricksConfig().map_openai_params( non_default_params=non_default_params, @@ -6083,6 +6097,7 @@ def _get_model_info_helper( # noqa: PLR0915 "provider_specific_entry", None ), uses_embed_content=_model_info.get("uses_embed_content", None), + supports_image_size=_model_info.get("supports_image_size", None), ) except Exception as e: verbose_logger.debug(f"Error getting model info: {e}") @@ -6637,6 +6652,14 @@ def validate_environment( # noqa: PLR0915 keys_in_environment = True else: missing_keys.append("CODESTRAL_API_KEY") + elif ( + custom_llm_provider == "inception" + or custom_llm_provider == "text-completion-inception" + ): + if "INCEPTION_API_KEY" in os.environ: + keys_in_environment = True + else: + missing_keys.append("INCEPTION_API_KEY") elif custom_llm_provider == "deepseek": if "DEEPSEEK_API_KEY" in os.environ: keys_in_environment = True @@ -8291,6 +8314,7 @@ class ProviderConfigManager: LlmProviders.XAI: (lambda: litellm.XAIChatConfig(), False), LlmProviders.ZAI: (lambda: litellm.ZAIChatConfig(), False), LlmProviders.LAMBDA_AI: (lambda: litellm.LambdaAIChatConfig(), False), + LlmProviders.INCEPTION: (lambda: litellm.InceptionChatConfig(), False), LlmProviders.LLAMA: (lambda: litellm.LlamaAPIConfig(), False), LlmProviders.TEXT_COMPLETION_OPENAI: ( lambda: litellm.OpenAITextCompletionConfig(), @@ -8356,6 +8380,10 @@ class ProviderConfigManager: lambda: litellm.CodestralTextCompletionConfig(), False, ), + LlmProviders.TEXT_COMPLETION_INCEPTION: ( + lambda: litellm.InceptionTextCompletionConfig(), + False, + ), LlmProviders.SAMBANOVA: (lambda: litellm.SambanovaConfig(), False), LlmProviders.MARITALK: (lambda: litellm.MaritalkConfig(), False), LlmProviders.VLLM: (lambda: litellm.VLLMConfig(), False), @@ -8394,6 +8422,10 @@ class ProviderConfigManager: lambda: ProviderConfigManager._get_langgraph_config(), False, ), + LlmProviders.LANGFLOW: ( + lambda: ProviderConfigManager._get_langflow_config(), + False, + ), } @staticmethod @@ -8465,6 +8497,13 @@ class ProviderConfigManager: return LangGraphConfig() + @staticmethod + def _get_langflow_config() -> BaseConfig: + """Get LangFlow config.""" + from litellm.llms.langflow.chat.transformation import LangFlowConfig + + return LangFlowConfig() + @staticmethod def get_provider_chat_config( # noqa: PLR0915 model: str, @@ -8917,6 +8956,8 @@ class ProviderConfigManager: return litellm.FireworksAITextCompletionConfig() elif LlmProviders.TOGETHER_AI == provider: return litellm.TogetherAITextCompletionConfig() + elif LlmProviders.TEXT_COMPLETION_INCEPTION == provider: + return litellm.InceptionTextCompletionConfig() return litellm.OpenAITextCompletionConfig() @staticmethod diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index b253235ec2a..ed6de4fa6b7 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -1075,6 +1075,7 @@ }, "eu.anthropic.claude-opus-4-6-v1": { "cache_creation_input_token_cost": 6.875e-06, + "cache_creation_input_token_cost_above_1hr": 1.1e-05, "cache_read_input_token_cost": 5.5e-07, "input_cost_per_token": 5.5e-06, "litellm_provider": "bedrock_converse", @@ -1104,6 +1105,7 @@ }, "au.anthropic.claude-opus-4-6-v1": { "cache_creation_input_token_cost": 6.875e-06, + "cache_creation_input_token_cost_above_1hr": 1.1e-05, "cache_read_input_token_cost": 5.5e-07, "input_cost_per_token": 5.5e-06, "litellm_provider": "bedrock_converse", @@ -1241,6 +1243,7 @@ }, "eu.anthropic.claude-opus-4-7": { "cache_creation_input_token_cost": 6.875e-06, + "cache_creation_input_token_cost_above_1hr": 1.1e-05, "cache_read_input_token_cost": 5.5e-07, "input_cost_per_token": 5.5e-06, "litellm_provider": "bedrock_converse", @@ -1271,6 +1274,7 @@ }, "au.anthropic.claude-opus-4-7": { "cache_creation_input_token_cost": 6.875e-06, + "cache_creation_input_token_cost_above_1hr": 1.1e-05, "cache_read_input_token_cost": 5.5e-07, "input_cost_per_token": 5.5e-06, "litellm_provider": "bedrock_converse", @@ -1543,6 +1547,7 @@ }, "eu.anthropic.claude-sonnet-4-6": { "cache_creation_input_token_cost": 4.125e-06, + "cache_creation_input_token_cost_above_1hr": 6.6e-06, "cache_read_input_token_cost": 3.3e-07, "input_cost_per_token": 3.3e-06, "litellm_provider": "bedrock_converse", @@ -1571,6 +1576,7 @@ }, "au.anthropic.claude-sonnet-4-6": { "cache_creation_input_token_cost": 4.125e-06, + "cache_creation_input_token_cost_above_1hr": 6.6e-06, "cache_read_input_token_cost": 3.3e-07, "input_cost_per_token": 3.3e-06, "litellm_provider": "bedrock_converse", @@ -1599,6 +1605,7 @@ }, "jp.anthropic.claude-sonnet-4-6": { "cache_creation_input_token_cost": 4.125e-06, + "cache_creation_input_token_cost_above_1hr": 6.6e-06, "cache_read_input_token_cost": 3.3e-07, "input_cost_per_token": 3.3e-06, "litellm_provider": "bedrock_converse", @@ -1995,11 +2002,13 @@ }, "au.anthropic.claude-sonnet-4-5-20250929-v1:0": { "cache_creation_input_token_cost": 4.125e-06, + "cache_creation_input_token_cost_above_1hr": 6.6e-06, "cache_read_input_token_cost": 3.3e-07, "input_cost_per_token": 3.3e-06, "input_cost_per_token_above_200k_tokens": 6.6e-06, "output_cost_per_token_above_200k_tokens": 2.475e-05, "cache_creation_input_token_cost_above_200k_tokens": 8.25e-06, + "cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.32e-05, "cache_read_input_token_cost_above_200k_tokens": 6.6e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 200000, @@ -7494,6 +7503,27 @@ "supports_video_input": true, "supports_vision": true }, + "azure_ai/kimi-k2.6": { + "input_cost_per_token": 9.5e-07, + "litellm_provider": "azure_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 4e-06, + "source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-kimi-k2-6-in-microsoft-foundry/4513125", + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_reasoning": true, + "supports_tool_choice": true, + "supports_vision": true + }, "azure_ai/ministral-3b": { "input_cost_per_token": 4e-08, "litellm_provider": "azure_ai", @@ -12682,7 +12712,8 @@ "litellm_provider": "deepinfra", "mode": "chat", "supports_tool_choice": true, - "supports_function_calling": true + "supports_function_calling": true, + "supports_image_size": false }, "deepinfra/google/gemini-2.5-pro": { "max_tokens": 1000000, @@ -13583,6 +13614,7 @@ }, "eu.anthropic.claude-haiku-4-5-20251001-v1:0": { "cache_creation_input_token_cost": 1.375e-06, + "cache_creation_input_token_cost_above_1hr": 2.2e-06, "cache_read_input_token_cost": 1.1e-07, "input_cost_per_token": 1.1e-06, "deprecation_date": "2026-10-15", @@ -13787,11 +13819,13 @@ }, "eu.anthropic.claude-sonnet-4-5-20250929-v1:0": { "cache_creation_input_token_cost": 4.125e-06, + "cache_creation_input_token_cost_above_1hr": 6.6e-06, "cache_read_input_token_cost": 3.3e-07, "input_cost_per_token": 3.3e-06, "input_cost_per_token_above_200k_tokens": 6.6e-06, "output_cost_per_token_above_200k_tokens": 2.475e-05, "cache_creation_input_token_cost_above_200k_tokens": 8.25e-06, + "cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.32e-05, "cache_read_input_token_cost_above_200k_tokens": 6.6e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 200000, @@ -15006,7 +15040,8 @@ "search_context_size_medium": 0.035, "search_context_size_high": 0.035 }, - "supports_service_tier": true + "supports_service_tier": true, + "supports_image_size": false }, "gemini-2.5-flash-image": { "cache_read_input_token_cost": 3e-08, @@ -15056,7 +15091,8 @@ "supports_vision": true, "supports_web_search": false, "tpm": 8000000, - "supports_service_tier": true + "supports_service_tier": true, + "supports_image_size": false }, "gemini-3-pro-image-preview": { "input_cost_per_image": 0.0011, @@ -15345,7 +15381,8 @@ "search_context_size_medium": 0.035, "search_context_size_high": 0.035 }, - "supports_service_tier": true + "supports_service_tier": true, + "supports_image_size": false }, "gemini-2.5-flash-lite-preview-09-2025": { "cache_read_input_token_cost": 1e-08, @@ -15395,7 +15432,8 @@ "search_context_size_low": 0.035, "search_context_size_medium": 0.035, "search_context_size_high": 0.035 - } + }, + "supports_image_size": false }, "gemini-2.5-flash-preview-09-2025": { "cache_read_input_token_cost": 7.5e-08, @@ -15445,7 +15483,8 @@ "search_context_size_low": 0.035, "search_context_size_medium": 0.035, "search_context_size_high": 0.035 - } + }, + "supports_image_size": false }, "gemini-live-2.5-flash-preview-native-audio-09-2025": { "cache_read_input_token_cost": 7.5e-08, @@ -15596,7 +15635,8 @@ "search_context_size_low": 0.035, "search_context_size_medium": 0.035, "search_context_size_high": 0.035 - } + }, + "supports_image_size": false }, "gemini-2.5-pro": { "cache_read_input_token_cost": 1.25e-07, @@ -16606,7 +16646,8 @@ "search_context_size_medium": 0.035, "search_context_size_high": 0.035 }, - "supports_service_tier": true + "supports_service_tier": true, + "supports_image_size": false }, "gemini/gemini-2.5-flash-image": { "cache_read_input_token_cost": 3e-08, @@ -16662,7 +16703,8 @@ "search_context_size_medium": 0.035, "search_context_size_high": 0.035 }, - "supports_service_tier": true + "supports_service_tier": true, + "supports_image_size": false }, "gemini/gemini-3-pro-image-preview": { "input_cost_per_image": 0.0011, @@ -16841,7 +16883,8 @@ "search_context_size_medium": 0.035, "search_context_size_high": 0.035 }, - "supports_service_tier": true + "supports_service_tier": true, + "supports_image_size": false }, "gemini/gemini-2.5-flash-lite-preview-09-2025": { "cache_read_input_token_cost": 1e-08, @@ -16893,7 +16936,8 @@ "search_context_size_low": 0.035, "search_context_size_medium": 0.035, "search_context_size_high": 0.035 - } + }, + "supports_image_size": false }, "gemini/gemini-2.5-flash-preview-09-2025": { "cache_read_input_token_cost": 7.5e-08, @@ -16945,7 +16989,8 @@ "search_context_size_low": 0.035, "search_context_size_medium": 0.035, "search_context_size_high": 0.035 - } + }, + "supports_image_size": false }, "gemini/gemini-flash-latest": { "cache_read_input_token_cost": 7.5e-08, @@ -17102,7 +17147,8 @@ "search_context_size_low": 0.035, "search_context_size_medium": 0.035, "search_context_size_high": 0.035 - } + }, + "supports_image_size": false }, "gemini/gemini-2.5-flash-preview-tts": { "input_cost_per_token": 3e-07, @@ -23086,11 +23132,13 @@ }, "jp.anthropic.claude-sonnet-4-5-20250929-v1:0": { "cache_creation_input_token_cost": 4.125e-06, + "cache_creation_input_token_cost_above_1hr": 6.6e-06, "cache_read_input_token_cost": 3.3e-07, "input_cost_per_token": 3.3e-06, "input_cost_per_token_above_200k_tokens": 6.6e-06, "output_cost_per_token_above_200k_tokens": 2.475e-05, "cache_creation_input_token_cost_above_200k_tokens": 8.25e-06, + "cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.32e-05, "cache_read_input_token_cost_above_200k_tokens": 6.6e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 200000, @@ -23116,6 +23164,7 @@ }, "jp.anthropic.claude-haiku-4-5-20251001-v1:0": { "cache_creation_input_token_cost": 1.375e-06, + "cache_creation_input_token_cost_above_1hr": 2.2e-06, "cache_read_input_token_cost": 1.1e-07, "input_cost_per_token": 1.1e-06, "litellm_provider": "bedrock_converse", @@ -23228,6 +23277,31 @@ "supports_system_messages": true, "supports_tool_choice": true }, + "inception/mercury-2": { + "cache_read_input_token_cost": 2.5e-08, + "input_cost_per_token": 2.5e-07, + "litellm_provider": "inception", + "max_input_tokens": 128000, + "max_output_tokens": 50000, + "max_tokens": 50000, + "mode": "chat", + "output_cost_per_token": 7.5e-07, + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true + }, + "text-completion-inception/mercury-edit-2": { + "cache_read_input_token_cost": 2.5e-08, + "input_cost_per_token": 2.5e-07, + "litellm_provider": "text-completion-inception", + "max_input_tokens": 32000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "completion", + "output_cost_per_token": 7.5e-07 + }, "lambda_ai/deepseek-llama3.3-70b": { "input_cost_per_token": 2e-07, "litellm_provider": "lambda_ai", @@ -26449,7 +26523,8 @@ "supports_function_calling": true, "supports_response_schema": true, "supports_vision": true, - "supports_native_streaming": true + "supports_native_streaming": true, + "supports_image_size": false }, "oci/google.gemini-2.5-pro": { "input_cost_per_token": 1.25e-06, @@ -26477,7 +26552,8 @@ "supports_function_calling": true, "supports_response_schema": true, "supports_vision": true, - "supports_native_streaming": true + "supports_native_streaming": true, + "supports_image_size": false }, "oci/cohere.command-a-vision": { "input_cost_per_token": 1.56e-06, @@ -27513,7 +27589,8 @@ "supports_response_schema": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_image_size": false }, "openrouter/google/gemini-2.5-pro": { "input_cost_per_audio_token": 7e-07, @@ -29383,7 +29460,8 @@ "mode": "responses", "supports_web_search": true, "supports_reasoning": false, - "supports_function_calling": true + "supports_function_calling": true, + "supports_image_size": false }, "perplexity/xai/grok-4-1-fast-non-reasoning": { "litellm_provider": "perplexity", @@ -29965,7 +30043,8 @@ "supports_vision": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_response_schema": true + "supports_response_schema": true, + "supports_image_size": false }, "replicate/openai/gpt-oss-120b": { "input_cost_per_token": 1.8e-07, @@ -31784,6 +31863,7 @@ }, "au.anthropic.claude-haiku-4-5-20251001-v1:0": { "cache_creation_input_token_cost": 1.375e-06, + "cache_creation_input_token_cost_above_1hr": 2.2e-06, "cache_read_input_token_cost": 1.1e-07, "input_cost_per_token": 1.1e-06, "litellm_provider": "bedrock_converse", @@ -32663,7 +32743,8 @@ "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, - "supports_response_schema": true + "supports_response_schema": true, + "supports_image_size": false }, "vercel_ai_gateway/google/gemini-2.5-pro": { "input_cost_per_token": 2.5e-06, @@ -33430,6 +33511,7 @@ }, "vertex_ai/claude-haiku-4-5": { "cache_creation_input_token_cost": 1.25e-06, + "cache_creation_input_token_cost_above_1hr": 2e-06, "cache_read_input_token_cost": 1e-07, "input_cost_per_token": 1e-06, "litellm_provider": "vertex_ai-anthropic_models", @@ -33451,6 +33533,7 @@ }, "vertex_ai/claude-haiku-4-5@20251001": { "cache_creation_input_token_cost": 1.25e-06, + "cache_creation_input_token_cost_above_1hr": 2e-06, "cache_read_input_token_cost": 1e-07, "input_cost_per_token": 1e-06, "litellm_provider": "vertex_ai-anthropic_models", @@ -33501,6 +33584,7 @@ }, "vertex_ai/claude-3-7-sonnet@20250219": { "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_1hr": 6e-06, "cache_read_input_token_cost": 3e-07, "deprecation_date": "2026-05-11", "input_cost_per_token": 3e-06, @@ -33600,6 +33684,7 @@ }, "vertex_ai/claude-opus-4": { "cache_creation_input_token_cost": 1.875e-05, + "cache_creation_input_token_cost_above_1hr": 3e-05, "cache_read_input_token_cost": 1.5e-06, "input_cost_per_token": 1.5e-05, "litellm_provider": "vertex_ai-anthropic_models", @@ -33625,6 +33710,7 @@ }, "vertex_ai/claude-opus-4-1": { "cache_creation_input_token_cost": 1.875e-05, + "cache_creation_input_token_cost_above_1hr": 3e-05, "cache_read_input_token_cost": 1.5e-06, "input_cost_per_token": 1.5e-05, "input_cost_per_token_batches": 7.5e-06, @@ -33642,6 +33728,7 @@ }, "vertex_ai/claude-opus-4-1@20250805": { "cache_creation_input_token_cost": 1.875e-05, + "cache_creation_input_token_cost_above_1hr": 3e-05, "cache_read_input_token_cost": 1.5e-06, "input_cost_per_token": 1.5e-05, "input_cost_per_token_batches": 7.5e-06, @@ -33659,6 +33746,7 @@ }, "vertex_ai/claude-opus-4-5": { "cache_creation_input_token_cost": 6.25e-06, + "cache_creation_input_token_cost_above_1hr": 1e-05, "cache_read_input_token_cost": 5e-07, "input_cost_per_token": 5e-06, "litellm_provider": "vertex_ai-anthropic_models", @@ -33685,6 +33773,7 @@ }, "vertex_ai/claude-opus-4-5@20251101": { "cache_creation_input_token_cost": 6.25e-06, + "cache_creation_input_token_cost_above_1hr": 1e-05, "cache_read_input_token_cost": 5e-07, "input_cost_per_token": 5e-06, "litellm_provider": "vertex_ai-anthropic_models", @@ -33712,6 +33801,7 @@ }, "vertex_ai/claude-opus-4-6": { "cache_creation_input_token_cost": 6.25e-06, + "cache_creation_input_token_cost_above_1hr": 1e-05, "cache_read_input_token_cost": 5e-07, "input_cost_per_token": 5e-06, "litellm_provider": "vertex_ai-anthropic_models", @@ -33739,6 +33829,7 @@ }, "vertex_ai/claude-opus-4-6@default": { "cache_creation_input_token_cost": 6.25e-06, + "cache_creation_input_token_cost_above_1hr": 1e-05, "cache_read_input_token_cost": 5e-07, "input_cost_per_token": 5e-06, "litellm_provider": "vertex_ai-anthropic_models", @@ -33766,6 +33857,7 @@ }, "vertex_ai/claude-opus-4-7": { "cache_creation_input_token_cost": 6.25e-06, + "cache_creation_input_token_cost_above_1hr": 1e-05, "cache_read_input_token_cost": 5e-07, "input_cost_per_token": 5e-06, "litellm_provider": "vertex_ai-anthropic_models", @@ -33793,6 +33885,7 @@ }, "vertex_ai/claude-opus-4-7@default": { "cache_creation_input_token_cost": 6.25e-06, + "cache_creation_input_token_cost_above_1hr": 1e-05, "cache_read_input_token_cost": 5e-07, "input_cost_per_token": 5e-06, "litellm_provider": "vertex_ai-anthropic_models", @@ -33876,6 +33969,7 @@ }, "vertex_ai/claude-sonnet-4-5": { "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_1hr": 6e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_token": 3e-06, "input_cost_per_token_above_200k_tokens": 6e-06, @@ -33902,6 +33996,7 @@ }, "vertex_ai/claude-sonnet-4-6": { "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_1hr": 6e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_token": 3e-06, "litellm_provider": "vertex_ai-anthropic_models", @@ -33929,6 +34024,7 @@ }, "vertex_ai/claude-sonnet-4-5@20250929": { "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_1hr": 6e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_token": 3e-06, "input_cost_per_token_above_200k_tokens": 6e-06, @@ -33956,6 +34052,7 @@ }, "vertex_ai/claude-opus-4@20250514": { "cache_creation_input_token_cost": 1.875e-05, + "cache_creation_input_token_cost_above_1hr": 3e-05, "cache_read_input_token_cost": 1.5e-06, "input_cost_per_token": 1.5e-05, "litellm_provider": "vertex_ai-anthropic_models", @@ -33981,6 +34078,7 @@ }, "vertex_ai/claude-sonnet-4": { "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_1hr": 6e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_token": 3e-06, "input_cost_per_token_above_200k_tokens": 6e-06, @@ -34010,6 +34108,7 @@ }, "vertex_ai/claude-sonnet-4@20250514": { "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_1hr": 6e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_token": 3e-06, "input_cost_per_token_above_200k_tokens": 6e-06, @@ -34217,7 +34316,8 @@ "supports_url_context": true, "supports_vision": true, "supports_web_search": false, - "tpm": 8000000 + "tpm": 8000000, + "supports_image_size": false }, "vertex_ai/gemini-3-pro-image-preview": { "input_cost_per_image": 0.0011, @@ -41035,6 +41135,7 @@ }, "vertex_ai/claude-sonnet-4-6@default": { "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_1hr": 6e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_token": 3e-06, "litellm_provider": "vertex_ai-anthropic_models", diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index 1e8357a8137..b4f782f9c3e 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -1273,6 +1273,24 @@ "interactions": true } }, + "inception": { + "display_name": "Inception (`inception`)", + "url": "https://docs.litellm.ai/docs/providers/inception", + "endpoints": { + "chat_completions": true, + "messages": true, + "responses": true, + "embeddings": false, + "image_generations": false, + "audio_transcriptions": false, + "audio_speech": false, + "moderations": false, + "batches": false, + "rerank": false, + "a2a": true, + "interactions": true + } + }, "infinity": { "display_name": "Infinity (`infinity`)", "url": "https://docs.litellm.ai/docs/providers/infinity", @@ -2430,6 +2448,24 @@ "interactions": true } }, + "langflow": { + "display_name": "LangFlow (`langflow`)", + "url": "https://docs.litellm.ai/docs/providers/langflow", + "endpoints": { + "chat_completions": true, + "messages": false, + "responses": false, + "embeddings": false, + "image_generations": false, + "audio_transcriptions": false, + "audio_speech": false, + "moderations": false, + "batches": false, + "rerank": false, + "a2a": true, + "interactions": false + } + }, "vertex_ai/agent_engine": { "display_name": "Vertex AI Agent Engine (`vertex_ai/agent_engine`)", "url": "https://docs.litellm.ai/docs/providers/vertex_ai_agent_engine", diff --git a/schema.prisma b/schema.prisma index 78143fe0411..c4754ef6117 100644 --- a/schema.prisma +++ b/schema.prisma @@ -325,6 +325,7 @@ model LiteLLM_MCPServerTable { allow_all_keys Boolean @default(false) available_on_public_internet Boolean @default(true) delegate_auth_to_upstream Boolean @default(false) + oauth_passthrough Boolean @default(false) is_byok Boolean @default(false) byok_description String[] @default([]) byok_api_key_help_url String? diff --git a/tests/_vcr_conftest_common.py b/tests/_vcr_conftest_common.py index d08b87bd580..995c333ad0e 100644 --- a/tests/_vcr_conftest_common.py +++ b/tests/_vcr_conftest_common.py @@ -53,6 +53,7 @@ CASSETTE_CACHE_HIGH_WATER_FRACTION = 0.85 SAFE_BODY_MATCHER_NAME = "safe_body" KEY_FINGERPRINT_MATCHER_NAME = "key_fingerprint" TOLERANT_QUERY_MATCHER_NAME = "tolerant_query" +TOLERANT_PATH_MATCHER_NAME = "tolerant_path" KEY_FINGERPRINT_HEADER = "x-litellm-key-fp" VCR_DIAG_DIR_ENV = "LITELLM_VCR_DIAG_DIR" @@ -411,6 +412,7 @@ def _canonical_body(request) -> tuple[bytes, str]: _VCR_UUID_RE = re.compile( rb"[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}" ) +_VCR_LITELLM_BATCH_JOB_RE = re.compile(rb"litellm-batch-[0-9a-fA-F]{8}") # ISO-8601 timestamps, e.g. ``2026-05-25T03:40:37.262045Z`` / # ``2026-05-25T03:40:37+00:00``. _VCR_ISO_TS_RE = re.compile( @@ -436,6 +438,7 @@ def _normalize_volatile_tokens(body: bytes) -> bytes: if not body: return body body = _VCR_UUID_RE.sub(b"", body) + body = _VCR_LITELLM_BATCH_JOB_RE.sub(b"litellm-batch-", body) body = _VCR_ISO_TS_RE.sub(b"", body) body = _VCR_UNIX_MS_RE.sub(b"", body) body = _VCR_UNIX_FLOAT_RE.sub(b"", body) @@ -1059,6 +1062,54 @@ def _tolerant_query_matcher(r1, r2) -> None: _vcr_matchers.query(r1, r2) +_BEDROCK_MANAGED_S3_PATH_RE = re.compile( + r"(?P(?:^|/)(?:litellm-bedrock-files/[^/?#]+-|litellm-bedrock-files-[^/?#]+-))" + r"[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}" + r"(?P\.jsonl)" +) + + +def _request_path_for_matcher(request) -> str: + path = getattr(request, "path", None) + if path is not None: + return str(path) + + uri = getattr(request, "uri", None) or getattr(request, "url", "") or "" + uri = str(uri) + if not uri: + return "" + if "//" in uri: + rest = uri.split("//", 1)[1] + path_part = "/" + rest.split("/", 1)[1] if "/" in rest else "/" + else: + path_part = uri + return path_part.split("?", 1)[0] + + +def _normalize_volatile_path(path: str) -> str: + return _BEDROCK_MANAGED_S3_PATH_RE.sub( + lambda match: f"{match.group('prefix')}{match.group('suffix')}", + path, + ) + + +def _tolerant_path_matcher(r1, r2) -> None: + """vcrpy's ``path`` matcher, plus LiteLLM-managed Bedrock S3 upload UUIDs. + + Bedrock batch file uploads use object keys like + ``litellm-bedrock-files-{model}-{uuid}.jsonl`` (and older cassettes may + contain ``litellm-bedrock-files/{model}-{uuid}.jsonl``). The UUID is + generated client-side before the S3 PUT, so strict path matching makes + every replay miss even when the JSONL body and all provider semantics are + identical. + """ + path1 = _normalize_volatile_path(_request_path_for_matcher(r1)) + path2 = _normalize_volatile_path(_request_path_for_matcher(r2)) + if path1 == path2: + return + _vcr_matchers.path(r1, r2) + + def vcr_config_dict() -> dict: return { "decode_compressed_response": True, @@ -1069,7 +1120,7 @@ def vcr_config_dict() -> dict: "scheme", "host", "port", - "path", + TOLERANT_PATH_MATCHER_NAME, TOLERANT_QUERY_MATCHER_NAME, KEY_FINGERPRINT_MATCHER_NAME, SAFE_BODY_MATCHER_NAME, @@ -1136,6 +1187,7 @@ def register_persister_if_enabled(vcr) -> None: vcr.register_matcher(SAFE_BODY_MATCHER_NAME, _safe_body_matcher) vcr.register_matcher(KEY_FINGERPRINT_MATCHER_NAME, _key_fingerprint_matcher) vcr.register_matcher(TOLERANT_QUERY_MATCHER_NAME, _tolerant_query_matcher) + vcr.register_matcher(TOLERANT_PATH_MATCHER_NAME, _tolerant_path_matcher) patch_vcrpy_aiohttp_record_path() patch_vcrpy_cassette_load_guard() global _atexit_banner_registered @@ -1647,13 +1699,18 @@ def _is_live_call_host(host: str) -> bool: return False if any(host.endswith(suffix) for suffix in _LIVE_CALL_HOST_SUFFIXES): return True - # AWS Bedrock endpoints are ``bedrock-runtime[-fips].{region}.amazonaws.com`` - # (region between ``bedrock-runtime`` and ``amazonaws.com``), so plain - # suffix matching can't catch them. - if host.endswith(".amazonaws.com") and host.split(".", 1)[0].startswith( - "bedrock-runtime" - ): - return True + if host.endswith(".amazonaws.com"): + first_label = host.split(".", 1)[0] + # AWS Bedrock control/runtime endpoints are + # ``bedrock[-runtime][-fips].{region}.amazonaws.com`` (region between + # the service label and ``amazonaws.com``), so plain suffix matching + # can't catch them. + if first_label.startswith("bedrock"): + return True + # Bedrock batch file upload/download uses real S3. Treat those as part + # of the paid provider path so unmarked batch tests surface as leaks. + if first_label in {"s3", "s3-fips"} or ".s3." in host or ".s3-" in host: + return True return False diff --git a/tests/batches_tests/conftest.py b/tests/batches_tests/conftest.py index ecb606b2cf3..e1899a22b6c 100644 --- a/tests/batches_tests/conftest.py +++ b/tests/batches_tests/conftest.py @@ -1,5 +1,3 @@ -# conftest.py - import asyncio import os import sys @@ -11,6 +9,82 @@ sys.path.insert( ) # Adds the parent directory to the system path import litellm # noqa: E402,F401 +from tests._vcr_conftest_common import ( # noqa: E402,F401 + VerboseReporterState, + _pin_multipart_boundary, + apply_vcr_auto_marker_to_items, + emit_cassette_cache_session_banner, + emit_vcr_classification_summary, + emit_vcr_diagnostic_log, + install_live_call_probe, + record_vcr_outcome, + register_persister_if_enabled, + reset_vcr_diag_dir, + vcr_config_dict, +) + +_verbose_state = VerboseReporterState() + +_CALLBACK_ATTRS = ( + "callbacks", + "success_callback", + "failure_callback", + "_async_success_callback", + "_async_failure_callback", +) + +_SCALAR_ATTRS = ( + "num_retries", + "set_verbose", + "cache", + "allowed_fails", + "disable_aiohttp_transport", + "force_ipv4", + "drop_params", + "modify_params", + "api_base", + "api_key", + "cohere_key", +) + + +@pytest.fixture(scope="module") +def vcr_config(): + return vcr_config_dict() + + +def pytest_recording_configure(config, vcr): + register_persister_if_enabled(vcr) + + +@pytest.hookimpl(hookwrapper=True) +def pytest_runtest_makereport(item, call): + outcome = yield + rep = outcome.get_result() + setattr(item, f"rep_{rep.when}", rep) + + +@pytest.fixture(autouse=True) +def _vcr_outcome_gate(request, vcr): + install_live_call_probe(request, vcr) + yield + record_vcr_outcome(request, vcr) + + +def pytest_configure(config): + _verbose_state.remember_pluginmanager(config) + reset_vcr_diag_dir() + + +def pytest_runtest_logreport(report): + _verbose_state.maybe_emit_verdict(report) + + +def pytest_terminal_summary(terminalreporter, exitstatus, config): + emit_cassette_cache_session_banner(terminalreporter) + emit_vcr_classification_summary(terminalreporter) + emit_vcr_diagnostic_log(terminalreporter) + @pytest.fixture(scope="session") def event_loop(): @@ -20,3 +94,64 @@ def event_loop(): loop = asyncio.new_event_loop() yield loop loop.close() + + +def _copy_litellm_state(): + state = {} + for attr in _CALLBACK_ATTRS: + if hasattr(litellm, attr): + value = getattr(litellm, attr) + state[attr] = value.copy() if isinstance(value, list) else value + for attr in _SCALAR_ATTRS: + if hasattr(litellm, attr): + state[attr] = getattr(litellm, attr) + return state + + +def _restore_litellm_state(state) -> None: + for attr, value in state.items(): + if hasattr(litellm, attr): + setattr(litellm, attr, value) + + +def _reset_litellm_callbacks() -> None: + for attr in _CALLBACK_ATTRS: + if hasattr(litellm, attr): + setattr(litellm, attr, []) + manager = getattr(litellm, "logging_callback_manager", None) + reset = getattr(manager, "_reset_all_callbacks", None) + if callable(reset): + reset() + + +def _clear_logging_queue(loop=None) -> None: + from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER + + if loop is not None and not loop.is_closed() and not loop.is_running(): + loop.run_until_complete(GLOBAL_LOGGING_WORKER.clear_queue()) + return + asyncio.run(GLOBAL_LOGGING_WORKER.clear_queue()) + + +@pytest.fixture(scope="function", autouse=True) +def setup_and_teardown(event_loop): + original_state = _copy_litellm_state() + _clear_logging_queue(event_loop) + _reset_litellm_callbacks() + asyncio.set_event_loop(event_loop) + + yield + + _clear_logging_queue(event_loop) + _reset_litellm_callbacks() + _restore_litellm_state(original_state) + + pending = asyncio.all_tasks(event_loop) + for task in pending: + task.cancel() + if pending: + event_loop.run_until_complete(asyncio.gather(*pending, return_exceptions=True)) + + +def pytest_collection_modifyitems(config, items): + apply_vcr_auto_marker_to_items(items) diff --git a/tests/batches_tests/test_batch_rate_limits.py b/tests/batches_tests/test_batch_rate_limits.py index 46013e19d30..e1fe8782ef9 100644 --- a/tests/batches_tests/test_batch_rate_limits.py +++ b/tests/batches_tests/test_batch_rate_limits.py @@ -51,6 +51,12 @@ def get_expected_batch_file_usage(file_path: str) -> tuple[int, int]: return expected_request_count, expected_total_tokens +def _write_batch_file(tmp_path, file_name: str, content: str) -> str: + path = tmp_path / file_name + path.write_text(content) + return str(path) + + @pytest.mark.asyncio() @pytest.mark.skipif( os.environ.get("OPENAI_API_KEY") is None, @@ -114,7 +120,7 @@ async def test_batch_rate_limits(): @pytest.mark.asyncio() -async def test_batch_rate_limit_single_file(): +async def test_batch_rate_limit_single_file(tmp_path): """ Test batch rate limiting with a single file. @@ -122,8 +128,6 @@ async def test_batch_rate_limit_single_file(): - File with < 200 tokens: should go through - File with > 200 tokens: should hit rate limit """ - import tempfile - CUSTOM_LLM_PROVIDER = "openai" # Setup: Create internal usage cache and rate limiter @@ -152,17 +156,18 @@ async def test_batch_rate_limit_single_file(): {"custom_id": "request-2", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "Hi"}]}} {"custom_id": "request-3", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "Hey"}]}}""" - with tempfile.NamedTemporaryFile(mode="w", suffix=".jsonl", delete=False) as f: - f.write(small_batch_content) - small_file_path = f.name + small_file_path = _write_batch_file( + tmp_path, "small-batch-rate-limit.jsonl", small_batch_content + ) try: # Upload file to OpenAI - file_obj_small = await litellm.acreate_file( - file=open(small_file_path, "rb"), - purpose="batch", - custom_llm_provider=CUSTOM_LLM_PROVIDER, - ) + with open(small_file_path, "rb") as batch_file: + file_obj_small = await litellm.acreate_file( + file=batch_file, + purpose="batch", + custom_llm_provider=CUSTOM_LLM_PROVIDER, + ) print(f"Created small file: {file_obj_small.id}") await asyncio.sleep(1) # Give API time to process @@ -183,8 +188,6 @@ async def test_batch_rate_limit_single_file(): print(f" Actual tokens: {result.get('_batch_token_count')}") except HTTPException as e: pytest.fail(f"Should not have hit rate limit with small file: {e.detail}") - finally: - os.unlink(small_file_path) # Test 2: File with > 200 tokens should hit rate limit print("\n=== Test 2: File over 200 tokens ===") @@ -221,47 +224,45 @@ async def test_batch_rate_limit_single_file(): large_batch_content = "\n".join(requests) - with tempfile.NamedTemporaryFile(mode="w", suffix=".jsonl", delete=False) as f: - f.write(large_batch_content) - large_file_path = f.name + large_file_path = _write_batch_file( + tmp_path, "large-batch-rate-limit.jsonl", large_batch_content + ) - try: - # Upload file to OpenAI + # Upload file to OpenAI + with open(large_file_path, "rb") as batch_file: file_obj_large = await litellm.acreate_file( - file=open(large_file_path, "rb"), + file=batch_file, purpose="batch", custom_llm_provider=CUSTOM_LLM_PROVIDER, ) - print(f"Created large file: {file_obj_large.id}") - await asyncio.sleep(1) # Give API time to process + print(f"Created large file: {file_obj_large.id}") + await asyncio.sleep(1) # Give API time to process - data_over_limit = { - "model": "gpt-3.5-turbo", - "input_file_id": file_obj_large.id, - "custom_llm_provider": CUSTOM_LLM_PROVIDER, - } + data_over_limit = { + "model": "gpt-3.5-turbo", + "input_file_id": file_obj_large.id, + "custom_llm_provider": CUSTOM_LLM_PROVIDER, + } - # Should raise HTTPException with 429 status - with pytest.raises(HTTPException) as exc_info: - await batch_limiter.async_pre_call_hook( - user_api_key_dict=user_api_key_dict, - cache=dual_cache, - data=data_over_limit, - call_type="acreate_batch", - ) + # Should raise HTTPException with 429 status + with pytest.raises(HTTPException) as exc_info: + await batch_limiter.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=dual_cache, + data=data_over_limit, + call_type="acreate_batch", + ) - assert exc_info.value.status_code == 429, "Should return 429 status code" - assert ( - "tokens" in exc_info.value.detail.lower() - ), "Error message should mention tokens" - print(f"✓ File with 250+ tokens correctly rejected (over limit of 200)") - print(f" Error: {exc_info.value.detail}") - finally: - os.unlink(large_file_path) + assert exc_info.value.status_code == 429, "Should return 429 status code" + assert ( + "tokens" in exc_info.value.detail.lower() + ), "Error message should mention tokens" + print(f"✓ File with 250+ tokens correctly rejected (over limit of 200)") + print(f" Error: {exc_info.value.detail}") @pytest.mark.asyncio() -async def test_batch_rate_limit_multiple_requests(): +async def test_batch_rate_limit_multiple_requests(tmp_path): """ Test batch rate limiting with multiple requests. @@ -269,8 +270,6 @@ async def test_batch_rate_limit_multiple_requests(): - Request 1: file with ~100 tokens (should go through, 100/200 used) - Request 2: file with ~105 tokens (should hit limit, 100+105=205 > 200) """ - import tempfile - CUSTOM_LLM_PROVIDER = "openai" # Setup: Create internal usage cache and rate limiter @@ -313,17 +312,18 @@ async def test_batch_rate_limit_multiple_requests(): batch_content_1 = "\n".join(requests_1) - with tempfile.NamedTemporaryFile(mode="w", suffix=".jsonl", delete=False) as f: - f.write(batch_content_1) - file_path_1 = f.name + file_path_1 = _write_batch_file( + tmp_path, "batch-rate-limit-request-1.jsonl", batch_content_1 + ) try: # Upload file to OpenAI - file_obj_1 = await litellm.acreate_file( - file=open(file_path_1, "rb"), - purpose="batch", - custom_llm_provider=CUSTOM_LLM_PROVIDER, - ) + with open(file_path_1, "rb") as batch_file: + file_obj_1 = await litellm.acreate_file( + file=batch_file, + purpose="batch", + custom_llm_provider=CUSTOM_LLM_PROVIDER, + ) print(f"Created file 1: {file_obj_1.id}") await asyncio.sleep(1) # Give API time to process @@ -346,8 +346,6 @@ async def test_batch_rate_limit_multiple_requests(): ) except HTTPException as e: pytest.fail(f"Request 1 should not have hit rate limit: {e.detail}") - finally: - os.unlink(file_path_1) # Request 2: File with ~105+ tokens (total would exceed 200) print("\n=== Request 2: File with ~105 tokens (should hit limit) ===") @@ -371,43 +369,41 @@ async def test_batch_rate_limit_multiple_requests(): batch_content_2 = "\n".join(requests_2) - with tempfile.NamedTemporaryFile(mode="w", suffix=".jsonl", delete=False) as f: - f.write(batch_content_2) - file_path_2 = f.name + file_path_2 = _write_batch_file( + tmp_path, "batch-rate-limit-request-2.jsonl", batch_content_2 + ) - try: - # Upload file to OpenAI + # Upload file to OpenAI + with open(file_path_2, "rb") as batch_file: file_obj_2 = await litellm.acreate_file( - file=open(file_path_2, "rb"), + file=batch_file, purpose="batch", custom_llm_provider=CUSTOM_LLM_PROVIDER, ) - print(f"Created file 2: {file_obj_2.id}") - await asyncio.sleep(1) # Give API time to process + print(f"Created file 2: {file_obj_2.id}") + await asyncio.sleep(1) # Give API time to process - data_request2 = { - "model": "gpt-3.5-turbo", - "input_file_id": file_obj_2.id, - "custom_llm_provider": CUSTOM_LLM_PROVIDER, - } + data_request2 = { + "model": "gpt-3.5-turbo", + "input_file_id": file_obj_2.id, + "custom_llm_provider": CUSTOM_LLM_PROVIDER, + } - # Should raise HTTPException with 429 status - with pytest.raises(HTTPException) as exc_info: - await batch_limiter.async_pre_call_hook( - user_api_key_dict=user_api_key_dict, - cache=dual_cache, - data=data_request2, - call_type="acreate_batch", - ) + # Should raise HTTPException with 429 status + with pytest.raises(HTTPException) as exc_info: + await batch_limiter.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=dual_cache, + data=data_request2, + call_type="acreate_batch", + ) - assert exc_info.value.status_code == 429, "Should return 429 status code" - assert ( - "tokens" in exc_info.value.detail.lower() - ), "Error message should mention tokens" - print(f"✓ Request 2 correctly rejected") - print(f" Error: {exc_info.value.detail}") - finally: - os.unlink(file_path_2) + assert exc_info.value.status_code == 429, "Should return 429 status code" + assert ( + "tokens" in exc_info.value.detail.lower() + ), "Error message should mention tokens" + print(f"✓ Request 2 correctly rejected") + print(f" Error: {exc_info.value.detail}") @pytest.mark.asyncio() @@ -415,7 +411,7 @@ async def test_batch_rate_limit_multiple_requests(): os.environ.get("OPENAI_API_KEY") is None, reason="OPENAI_API_KEY not set - skipping integration test", ) -async def test_batch_rate_limiter_with_managed_files(): +async def test_batch_rate_limiter_with_managed_files(tmp_path): """ Test for GEN-2166: Verify batch rate limiter can read user files when managed files are enabled. @@ -425,7 +421,6 @@ async def test_batch_rate_limiter_with_managed_files(): 3. Rate limiting is enforced (not silently bypassed) 4. No 403 Permission Denied errors occur for files owned by the user """ - import tempfile from unittest.mock import AsyncMock, MagicMock, patch CUSTOM_LLM_PROVIDER = "openai" @@ -472,18 +467,19 @@ async def test_batch_rate_limiter_with_managed_files(): batch_content = "\n".join(requests) - with tempfile.NamedTemporaryFile(mode="w", suffix=".jsonl", delete=False) as f: - f.write(batch_content) - file_path = f.name + file_path = _write_batch_file( + tmp_path, "managed-files-batch-rate-limit.jsonl", batch_content + ) try: # Step 1: Upload file to OpenAI (simulating user upload) print("\n1. Uploading batch input file...") - file_obj = await litellm.acreate_file( - file=open(file_path, "rb"), - purpose="batch", - custom_llm_provider=CUSTOM_LLM_PROVIDER, - ) + with open(file_path, "rb") as batch_file: + file_obj = await litellm.acreate_file( + file=batch_file, + purpose="batch", + custom_llm_provider=CUSTOM_LLM_PROVIDER, + ) print(f" ✓ File uploaded: {file_obj.id}") await asyncio.sleep(1) # Give API time to process @@ -568,12 +564,10 @@ async def test_batch_rate_limiter_with_managed_files(): raise except Exception as e: pytest.fail(f"Unexpected error: {str(e)}") - finally: - os.unlink(file_path) @pytest.mark.asyncio() -async def test_batch_rate_limiter_without_user_context(): +async def test_batch_rate_limiter_without_user_context(tmp_path): """ Test that verifies the bug scenario from GEN-2166. @@ -583,8 +577,6 @@ async def test_batch_rate_limiter_without_user_context(): This test documents the expected behavior with and without user context. """ - import tempfile - CUSTOM_LLM_PROVIDER = "openai" # Setup @@ -596,56 +588,53 @@ async def test_batch_rate_limiter_without_user_context(): # Create a simple batch file batch_content = """{"custom_id": "request-1", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "Hello"}]}}""" - with tempfile.NamedTemporaryFile(mode="w", suffix=".jsonl", delete=False) as f: - f.write(batch_content) - file_path = f.name + file_path = _write_batch_file( + tmp_path, "without-user-context-batch-rate-limit.jsonl", batch_content + ) - try: - # Upload file + # Upload file + with open(file_path, "rb") as batch_file: file_obj = await litellm.acreate_file( - file=open(file_path, "rb"), + file=batch_file, purpose="batch", custom_llm_provider=CUSTOM_LLM_PROVIDER, ) - await asyncio.sleep(1) + await asyncio.sleep(1) - # Test 1: Without user context (old behavior - would fail with managed files) - print("\n=== Test 1: count_input_file_usage WITHOUT user context ===") - try: - usage_without_context = await BATCH_LIMITER.count_input_file_usage( - file_id=file_obj.id, - custom_llm_provider=CUSTOM_LLM_PROVIDER, - user_api_key_dict=None, # Explicitly passing None - ) - print( - f"✓ Works for non-managed files (tokens: {usage_without_context.total_tokens})" - ) - print(" Note: Would fail with 403 for managed files (GEN-2166 bug)") - except Exception as e: - print(f"✗ Failed: {str(e)}") - - # Test 2: With user context (new behavior - works with managed files) - print("\n=== Test 2: count_input_file_usage WITH user context ===") - user_api_key_dict = UserAPIKeyAuth( - api_key="test-key", - user_id="test-user-123", - ) - - usage_with_context = await BATCH_LIMITER.count_input_file_usage( + # Test 1: Without user context (old behavior - would fail with managed files) + print("\n=== Test 1: count_input_file_usage WITHOUT user context ===") + try: + usage_without_context = await BATCH_LIMITER.count_input_file_usage( file_id=file_obj.id, custom_llm_provider=CUSTOM_LLM_PROVIDER, - user_api_key_dict=user_api_key_dict, # Passing user context + user_api_key_dict=None, # Explicitly passing None ) - print(f"✓ Works with user context (tokens: {usage_with_context.total_tokens})") - print(" Note: This fixes GEN-2166 for managed files") + print( + f"✓ Works for non-managed files (tokens: {usage_without_context.total_tokens})" + ) + print(" Note: Would fail with 403 for managed files (GEN-2166 bug)") + except Exception as e: + print(f"✗ Failed: {str(e)}") - # Verify both return the same results - assert usage_with_context.total_tokens == usage_without_context.total_tokens - assert usage_with_context.request_count == usage_without_context.request_count - print("\n✓ Both methods return identical results for non-managed files") + # Test 2: With user context (new behavior - works with managed files) + print("\n=== Test 2: count_input_file_usage WITH user context ===") + user_api_key_dict = UserAPIKeyAuth( + api_key="test-key", + user_id="test-user-123", + ) - finally: - os.unlink(file_path) + usage_with_context = await BATCH_LIMITER.count_input_file_usage( + file_id=file_obj.id, + custom_llm_provider=CUSTOM_LLM_PROVIDER, + user_api_key_dict=user_api_key_dict, # Passing user context + ) + print(f"✓ Works with user context (tokens: {usage_with_context.total_tokens})") + print(" Note: This fixes GEN-2166 for managed files") + + # Verify both return the same results + assert usage_with_context.total_tokens == usage_without_context.total_tokens + assert usage_with_context.request_count == usage_without_context.request_count + print("\n✓ Both methods return identical results for non-managed files") @pytest.mark.asyncio() diff --git a/tests/batches_tests/test_bedrock_files_and_batches.py b/tests/batches_tests/test_bedrock_files_and_batches.py index 5148ea4db91..431d5a2a60c 100644 --- a/tests/batches_tests/test_bedrock_files_and_batches.py +++ b/tests/batches_tests/test_bedrock_files_and_batches.py @@ -1,7 +1,7 @@ # What is this? ## Unit Tests for OpenAI Batches API import asyncio -import json +import json as json_module import os import sys import traceback @@ -19,6 +19,103 @@ from typing import Optional import litellm from unittest.mock import patch, MagicMock import httpx +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + + +_BEDROCK_TEST_AWS_ENV = { + "AWS_ACCESS_KEY_ID": "test-access-key", + "AWS_SECRET_ACCESS_KEY": "test-secret-key", + "AWS_REGION": "us-west-2", + "AWS_DEFAULT_REGION": "us-west-2", +} + + +class _CaptureAsyncHTTPHandler(AsyncHTTPHandler): + def __init__(self): + self.timeout = None + self.event_hooks = None + self.client_alias = "bedrock-test" + self.put_calls = [] + self.post_calls = [] + self.batch_jobs = {} + + async def put( + self, + url: str, + data=None, + json=None, + params=None, + headers=None, + timeout=None, + stream: bool = False, + content=None, + ): + self.put_calls.append( + { + "url": url, + "data": data, + "json": json, + "params": params, + "headers": headers or {}, + "timeout": timeout, + "stream": stream, + "content": content, + } + ) + body = data if data is not None else content + content_bytes = body.encode("utf-8") if isinstance(body, str) else body or b"" + content_length = len(content_bytes) + return httpx.Response( + status_code=200, + headers={"Content-Length": str(content_length)}, + request=httpx.Request("PUT", url), + ) + + async def post( + self, + url: str, + data=None, + json=None, + params=None, + headers=None, + timeout=None, + stream: bool = False, + logging_obj=None, + files=None, + content=None, + ): + self.post_calls.append( + { + "url": url, + "data": data, + "json": json, + "params": params, + "headers": headers or {}, + "timeout": timeout, + "stream": stream, + "content": content, + } + ) + raw = json if json is not None else (data if data is not None else content) + payload = raw if isinstance(raw, dict) else json_module.loads(raw) + job_name = payload["jobName"] + job_arn = f"arn:aws:bedrock:us-west-2:941277531214:model-invocation-job/{job_name}" + self.batch_jobs[job_arn] = { + "jobArn": job_arn, + "jobName": job_name, + "modelId": payload["modelId"], + "roleArn": payload["roleArn"], + "status": "InProgress", + "submitTime": "2026-06-02T03:50:00Z", + "lastModifiedTime": "2026-06-02T03:55:00Z", + "inputDataConfig": payload["inputDataConfig"], + "outputDataConfig": payload["outputDataConfig"], + } + return httpx.Response( + status_code=200, + json={"jobArn": job_arn, "jobName": job_name, "status": "Submitted"}, + request=httpx.Request("POST", url), + ) @pytest.mark.asyncio() @@ -34,12 +131,34 @@ async def test_async_create_file(): file_name = "bedrock_batch_completions.jsonl" _current_dir = os.path.dirname(os.path.abspath(__file__)) file_path = os.path.join(_current_dir, file_name) - file_obj = await litellm.acreate_file( - file=open(file_path, "rb"), - purpose="batch", - custom_llm_provider="bedrock", - s3_bucket_name="litellm-proxy", + capture_client = _CaptureAsyncHTTPHandler() + with ( + patch.dict(os.environ, _BEDROCK_TEST_AWS_ENV), + open(file_path, "rb") as batch_file, + ): + file_obj = await litellm.acreate_file( + file=batch_file, + purpose="batch", + custom_llm_provider="bedrock", + s3_bucket_name="litellm-proxy-941277531214", + client=capture_client, + ) + + assert len(capture_client.put_calls) == 1 + put_call = capture_client.put_calls[0] + assert put_call["url"].startswith( + "https://s3.us-west-2.amazonaws.com/litellm-proxy-941277531214/" ) + assert "/litellm-bedrock-files-us.anthropic.claude-haiku-4-5-20251001-v1-0-" in ( + put_call["url"] + ) + assert put_call["url"].endswith(".jsonl") + assert put_call["headers"]["Authorization"].startswith("AWS4-HMAC-SHA256") + assert "recordId" in put_call["data"] + assert file_obj.id.startswith( + "s3://litellm-proxy-941277531214/litellm-bedrock-files-" + ) + assert file_obj.filename.endswith(".jsonl") @pytest.mark.asyncio() @@ -51,36 +170,54 @@ async def test_async_file_and_batch(): file_name = "bedrock_batch_completions.jsonl" _current_dir = os.path.dirname(os.path.abspath(__file__)) file_path = os.path.join(_current_dir, file_name) - file_obj = await litellm.acreate_file( - file=open(file_path, "rb"), - purpose="batch", - custom_llm_provider="bedrock", - s3_bucket_name="litellm-proxy", - ) - print("CREATED FILE RESPONSE=", file_obj) + capture_client = _CaptureAsyncHTTPHandler() + with patch.dict(os.environ, _BEDROCK_TEST_AWS_ENV): + with open(file_path, "rb") as batch_file: + file_obj = await litellm.acreate_file( + file=batch_file, + purpose="batch", + custom_llm_provider="bedrock", + s3_bucket_name="litellm-proxy-941277531214", + client=capture_client, + ) + assert len(capture_client.put_calls) == 1 + print("CREATED FILE RESPONSE=", file_obj) - # create batch - create_batch_response = await litellm.acreate_batch( - completion_window="24h", - endpoint="/v1/chat/completions", - input_file_id=file_obj.id, - metadata={"key1": "value1", "key2": "value2"}, - custom_llm_provider="bedrock", - ######################################################### - # bedrock specific params - ######################################################### - model="us.anthropic.claude-haiku-4-5-20251001-v1:0", - aws_batch_role_arn="arn:aws:iam::888602223428:role/service-role/AmazonBedrockExecutionRoleForAgents_BB9HNW6V4CV", - ) - print("CREATED BATCH RESPONSE=", create_batch_response) + with patch( + "litellm.llms.custom_httpx.llm_http_handler.get_async_httpx_client", + return_value=capture_client, + ): + # create batch + create_batch_response = await litellm.acreate_batch( + completion_window="24h", + endpoint="/v1/chat/completions", + input_file_id=file_obj.id, + metadata={"key1": "value1", "key2": "value2"}, + custom_llm_provider="bedrock", + ######################################################### + # bedrock specific params + ######################################################### + model="us.anthropic.claude-haiku-4-5-20251001-v1:0", + aws_batch_role_arn="arn:aws:iam::941277531214:role/service-role/AmazonBedrockExecutionRoleForAgents_BB9HNW6V4CV", + ) + assert len(capture_client.post_calls) == 1 + print("CREATED BATCH RESPONSE=", create_batch_response) - # retrieve batch - retrieve_batch_response = await litellm.aretrieve_batch( - batch_id=create_batch_response.id, - custom_llm_provider="bedrock", - model="us.anthropic.claude-haiku-4-5-20251001-v1:0", - ) - print("RETRIEVED BATCH RESPONSE=", retrieve_batch_response) + # retrieve batch + mock_bedrock_client = MagicMock() + mock_bedrock_client.get_model_invocation_job.side_effect = ( + lambda jobIdentifier: capture_client.batch_jobs[jobIdentifier] + ) + with patch("boto3.client", return_value=mock_bedrock_client): + retrieve_batch_response = await litellm.aretrieve_batch( + batch_id=create_batch_response.id, + custom_llm_provider="bedrock", + model="us.anthropic.claude-haiku-4-5-20251001-v1:0", + ) + mock_bedrock_client.get_model_invocation_job.assert_called_once_with( + jobIdentifier=create_batch_response.id + ) + print("RETRIEVED BATCH RESPONSE=", retrieve_batch_response) # Validate the response assert retrieve_batch_response.id == create_batch_response.id @@ -101,52 +238,36 @@ async def test_mock_bedrock_file_url_mapping(): """ print("Testing Bedrock file URL mapping") - captured_put_url = None - - async def mock_async_create_file(transformed_request, **kwargs): - nonlocal captured_put_url - # Capture PUT URL from transformed request - if isinstance(transformed_request, dict) and "url" in transformed_request: - captured_put_url = transformed_request["url"] - - # Call the real method to get actual response - from litellm.files.main import base_llm_http_handler - - return await base_llm_http_handler.__class__.async_create_file( - base_llm_http_handler, transformed_request, **kwargs - ) - - with patch( - "litellm.files.main.base_llm_http_handler.async_create_file", - side_effect=mock_async_create_file, + capture_client = _CaptureAsyncHTTPHandler() + with ( + patch.dict(os.environ, _BEDROCK_TEST_AWS_ENV), + open( + os.path.join(os.path.dirname(__file__), "bedrock_batch_completions.jsonl"), + "rb", + ) as batch_file, ): file_obj = await litellm.acreate_file( - file=open( - os.path.join( - os.path.dirname(__file__), "bedrock_batch_completions.jsonl" - ), - "rb", - ), + file=batch_file, purpose="batch", custom_llm_provider="bedrock", - s3_bucket_name="litellm-proxy", + s3_bucket_name="litellm-proxy-941277531214", + client=capture_client, ) - print(f"PUT URL: {captured_put_url}") - print(f"File ID: {file_obj.id}") + captured_put_url = capture_client.put_calls[0]["url"] + print(f"PUT URL: {captured_put_url}") + print(f"File ID: {file_obj.id}") - # Validate URL was captured and response is correct - assert captured_put_url is not None - assert file_obj.id.startswith("s3://") + # Validate URL was captured and response is correct + assert captured_put_url is not None + assert file_obj.id.startswith("s3://") - # Verify mapping - from litellm.llms.bedrock.files.transformation import BedrockFilesConfig + # Verify mapping + from litellm.llms.bedrock.files.transformation import BedrockFilesConfig - bedrock_config = BedrockFilesConfig() - expected_s3_uri, _ = bedrock_config._convert_https_url_to_s3_uri( - captured_put_url - ) - assert file_obj.id == expected_s3_uri + bedrock_config = BedrockFilesConfig() + expected_s3_uri, _ = bedrock_config._convert_https_url_to_s3_uri(captured_put_url) + assert file_obj.id == expected_s3_uri @pytest.mark.asyncio() @@ -237,8 +358,12 @@ def test_bedrock_batch_with_encryption_key_in_post_request(): mock_response.raise_for_status.return_value = None return mock_response - with patch( - "litellm.llms.custom_httpx.http_handler.HTTPHandler.post", side_effect=mock_post + with ( + patch.dict(os.environ, _BEDROCK_TEST_AWS_ENV), + patch( + "litellm.llms.custom_httpx.http_handler.HTTPHandler.post", + side_effect=mock_post, + ), ): response = litellm.create_batch( completion_window="24h", diff --git a/tests/batches_tests/test_openai_batches_and_files.py b/tests/batches_tests/test_openai_batches_and_files.py index 2e89381cba9..bccb5eaaacb 100644 --- a/tests/batches_tests/test_openai_batches_and_files.py +++ b/tests/batches_tests/test_openai_batches_and_files.py @@ -4,7 +4,6 @@ import asyncio import json import os import sys -import traceback import tempfile from dotenv import load_dotenv @@ -15,12 +14,10 @@ sys.path.insert( import logging import time -import asyncio import pytest from typing import Optional import litellm -from litellm import create_batch, create_file from litellm._logging import verbose_logger import openai @@ -28,7 +25,6 @@ verbose_logger.setLevel(logging.DEBUG) from litellm.integrations.custom_logger import CustomLogger from litellm.types.utils import StandardLoggingPayload -import random import socket import httpx from unittest.mock import patch, MagicMock @@ -49,6 +45,21 @@ skip_if_no_openai_network = pytest.mark.skipif( ) +async def _wait_for_standard_logging_object( + custom_logger: "TestCustomLogger", timeout: float = 15.0 +) -> StandardLoggingPayload: + from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER + + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + await GLOBAL_LOGGING_WORKER.flush() + if custom_logger.standard_logging_object is not None: + return custom_logger.standard_logging_object + await asyncio.sleep(0.25) + assert custom_logger.standard_logging_object is not None + return custom_logger.standard_logging_object + + def load_vertex_ai_credentials(): # Define the path to the vertex_key.json file print("loading vertex ai credentials") @@ -95,7 +106,7 @@ def load_vertex_ai_credentials(): @pytest.mark.parametrize("provider", ["openai"]) # , "azure" @pytest.mark.asyncio @skip_if_no_openai_network -async def test_create_batch(provider): +async def test_create_batch(provider, tmp_path): """ 1. Create File for Batch completion 2. Create Batch Request @@ -108,11 +119,12 @@ async def test_create_batch(provider): _current_dir = os.path.dirname(os.path.abspath(__file__)) file_path = os.path.join(_current_dir, file_name) - file_obj = await litellm.acreate_file( - file=open(file_path, "rb"), - purpose="batch", - custom_llm_provider=provider, - ) + with open(file_path, "rb") as batch_file: + file_obj = await litellm.acreate_file( + file=batch_file, + purpose="batch", + custom_llm_provider=provider, + ) print("Response from creating file=", file_obj) batch_input_file_id = file_obj.id @@ -161,10 +173,8 @@ async def test_create_batch(provider): result = file_content.content - result_file_name = "batch_job_results_furniture.jsonl" - - with open(result_file_name, "wb") as file: - file.write(result) + result_file_path = tmp_path / "batch_job_results_furniture.jsonl" + result_file_path.write_bytes(result) # Cancel Batch - handle race condition where batch may already be completed try: @@ -268,9 +278,8 @@ def cleanup_azure_ft_models(): @pytest.mark.parametrize("provider", ["openai"]) @pytest.mark.asyncio() -@pytest.mark.flaky(retries=3, delay=1) @skip_if_no_openai_network -async def test_async_create_batch(provider): +async def test_async_create_batch(provider, tmp_path): """ 1. Create File for Batch completion 2. Create Batch Request @@ -279,17 +288,16 @@ async def test_async_create_batch(provider): litellm._turn_on_debug() print("Testing async create batch") litellm.logging_callback_manager._reset_all_callbacks() - custom_logger = TestCustomLogger() - litellm.callbacks = [custom_logger, "datadog"] file_name = "openai_batch_completions.jsonl" _current_dir = os.path.dirname(os.path.abspath(__file__)) file_path = os.path.join(_current_dir, file_name) - file_obj = await litellm.acreate_file( - file=open(file_path, "rb"), - purpose="batch", - custom_llm_provider=provider, - ) + with open(file_path, "rb") as batch_file: + file_obj = await litellm.acreate_file( + file=batch_file, + purpose="batch", + custom_llm_provider=provider, + ) print("Response from creating file=", file_obj) await asyncio.sleep(10) @@ -302,6 +310,8 @@ async def test_async_create_batch(provider): "user_api_key_alias": "special_api_key_alias", "user_api_key_team_alias": "special_team_alias", } + custom_logger = TestCustomLogger() + litellm.callbacks = [custom_logger, "datadog"] create_batch_response = await litellm.acreate_batch( completion_window="24h", endpoint="/v1/chat/completions", @@ -325,19 +335,18 @@ async def test_async_create_batch(provider): create_batch_response.input_file_id == batch_input_file_id ), f"Failed to create batch, expected input_file_id to be {batch_input_file_id} but got {create_batch_response.input_file_id}" - await asyncio.sleep(6) # Assert that the create batch event is logged on CustomLogger - assert custom_logger.standard_logging_object is not None + standard_logging_object = await _wait_for_standard_logging_object(custom_logger) print( "standard_logging_object=", - json.dumps(custom_logger.standard_logging_object, indent=4, default=str), + json.dumps(standard_logging_object, indent=4, default=str), ) assert ( - custom_logger.standard_logging_object["metadata"]["user_api_key_alias"] + standard_logging_object["metadata"]["user_api_key_alias"] == extra_metadata_field["user_api_key_alias"] ) assert ( - custom_logger.standard_logging_object["metadata"]["user_api_key_team_alias"] + standard_logging_object["metadata"]["user_api_key_team_alias"] == extra_metadata_field["user_api_key_team_alias"] ) @@ -383,10 +392,8 @@ async def test_async_create_batch(provider): print("all_files_list = ", all_files_list) - result_file_name = "batch_job_results_furniture.jsonl" - - with open(result_file_name, "wb") as file: - file.write(file_content.content) + result_file_path = tmp_path / "batch_job_results_furniture.jsonl" + result_file_path.write_bytes(file_content.content) # Cancel Batch - handle race condition where batch may already be completed try: @@ -407,11 +414,6 @@ async def test_async_create_batch(provider): print(f"Unexpected error during batch cancellation: {e}") raise - if random.randint(1, 3) == 1: - print("Running random cleanup of Azure files and models...") - cleanup_azure_files() - cleanup_azure_ft_models() - mock_file_response = { "kind": "storage#object", diff --git a/tests/litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 8785e450a4b..2a8768df722 100644 --- a/tests/litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -1,8 +1,32 @@ """Tests for MCP OAuth discoverable endpoints""" import pytest +from fastapi import HTTPException from unittest.mock import AsyncMock, MagicMock, patch +TRUSTED_PROXY_IP = "10.0.0.5" +TRUSTED_PROXY_RANGES = ["10.0.0.0/8"] + + +def set_request_from_trusted_proxy(mock_request): + mock_request.client = MagicMock() + mock_request.client.host = TRUSTED_PROXY_IP + + +@pytest.fixture +def trusted_proxy_origin_headers(): + with ( + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.IPAddressUtils.is_request_from_trusted_proxy", + return_value=True, + ), + patch( + "litellm.proxy._experimental.mcp_server.oauth_utils.IPAddressUtils.is_request_from_trusted_proxy", + return_value=True, + ), + ): + yield + @pytest.mark.asyncio async def test_authorize_endpoint_includes_response_type(): @@ -56,7 +80,7 @@ async def test_authorize_endpoint_includes_response_type(): request=mock_request, client_id="test_client_id", mcp_server_name="test_oauth", - redirect_uri="https://client.example.com/callback", + redirect_uri="http://127.0.0.1:60108/callback", state="test_state", ) @@ -154,7 +178,6 @@ async def test_token_endpoint_forwards_code_verifier(): from litellm.types.mcp_server.mcp_server_manager import MCPServer from litellm.proxy._types import MCPTransport from fastapi import Request - import httpx except ImportError: pytest.skip("MCP discoverable endpoints not available") @@ -244,10 +267,15 @@ async def test_register_client_without_mcp_server_name_returns_dummy(): from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( register_client, ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) from fastapi import Request except ImportError: pytest.skip("MCP discoverable endpoints not available") + global_mcp_server_manager.registry.clear() + mock_request = MagicMock(spec=Request) mock_request.base_url = "https://proxy.litellm.example/" mock_request.headers = {} @@ -410,7 +438,9 @@ async def test_register_client_remote_registration_success(): @pytest.mark.asyncio -async def test_authorize_endpoint_respects_x_forwarded_proto(): +async def test_authorize_endpoint_respects_x_forwarded_proto( + trusted_proxy_origin_headers, +): """Test that authorize endpoint uses X-Forwarded-Proto header to construct correct redirect_uri""" try: from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( @@ -449,6 +479,7 @@ async def test_authorize_endpoint_respects_x_forwarded_proto(): mock_request = MagicMock(spec=Request) mock_request.base_url = "http://litellm.example.com/" # HTTP mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy + set_request_from_trusted_proxy(mock_request) # Mock the encryption functions with patch( @@ -461,7 +492,7 @@ async def test_authorize_endpoint_respects_x_forwarded_proto(): request=mock_request, client_id="test_client_id", mcp_server_name="test_oauth", - redirect_uri="https://client.example.com/callback", + redirect_uri="http://127.0.0.1:60108/callback", state="test_state", ) @@ -476,7 +507,9 @@ async def test_authorize_endpoint_respects_x_forwarded_proto(): @pytest.mark.asyncio -async def test_token_endpoint_respects_x_forwarded_proto(): +async def test_token_endpoint_respects_x_forwarded_proto( + trusted_proxy_origin_headers, +): """Test that token endpoint uses X-Forwarded-Proto header for redirect_uri""" try: from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( @@ -515,6 +548,7 @@ async def test_token_endpoint_respects_x_forwarded_proto(): mock_request = MagicMock(spec=Request) mock_request.base_url = "http://litellm-proxy.example.com/" # HTTP mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy + set_request_from_trusted_proxy(mock_request) # Mock httpx client response mock_response = MagicMock() @@ -535,7 +569,7 @@ async def test_token_endpoint_respects_x_forwarded_proto(): mock_get_client.return_value = mock_async_client # Call token endpoint - response = await token_endpoint( + await token_endpoint( request=mock_request, grant_type="authorization_code", code="test_code", @@ -666,7 +700,9 @@ async def test_oauth_protected_resource_legacy_pattern(): @pytest.mark.asyncio -async def test_oauth_protected_resource_respects_x_forwarded_proto(): +async def test_oauth_protected_resource_respects_x_forwarded_proto( + trusted_proxy_origin_headers, +): """Test that oauth_protected_resource_mcp uses X-Forwarded-Proto for URLs""" try: from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( @@ -704,6 +740,7 @@ async def test_oauth_protected_resource_respects_x_forwarded_proto(): mock_request = MagicMock(spec=Request) mock_request.base_url = "http://litellm.example.com/" # HTTP mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy + set_request_from_trusted_proxy(mock_request) # Call the endpoint response = await oauth_protected_resource_mcp( @@ -719,7 +756,9 @@ async def test_oauth_protected_resource_respects_x_forwarded_proto(): @pytest.mark.asyncio -async def test_oauth_authorization_server_respects_x_forwarded_proto(): +async def test_oauth_authorization_server_respects_x_forwarded_proto( + trusted_proxy_origin_headers, +): """Test that oauth_authorization_server_mcp uses X-Forwarded-Proto for URLs""" try: from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( @@ -757,6 +796,7 @@ async def test_oauth_authorization_server_respects_x_forwarded_proto(): mock_request = MagicMock(spec=Request) mock_request.base_url = "http://litellm.example.com/" # HTTP mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy + set_request_from_trusted_proxy(mock_request) # Call the endpoint response = await oauth_authorization_server_mcp( @@ -773,20 +813,28 @@ async def test_oauth_authorization_server_respects_x_forwarded_proto(): @pytest.mark.asyncio -async def test_register_client_respects_x_forwarded_proto(): +async def test_register_client_respects_x_forwarded_proto( + trusted_proxy_origin_headers, +): """Test that register_client uses X-Forwarded-Proto for redirect_uris""" try: from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( register_client, ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) from fastapi import Request except ImportError: pytest.skip("MCP discoverable endpoints not available") + global_mcp_server_manager.registry.clear() + # Mock request with http base_url but X-Forwarded-Proto: https mock_request = MagicMock(spec=Request) mock_request.base_url = "http://proxy.litellm.example/" # HTTP mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy + set_request_from_trusted_proxy(mock_request) with patch( "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", @@ -803,7 +851,9 @@ async def test_register_client_respects_x_forwarded_proto(): @pytest.mark.asyncio -async def test_authorize_endpoint_respects_x_forwarded_host(): +async def test_authorize_endpoint_respects_x_forwarded_host( + trusted_proxy_origin_headers, +): """Test that authorize endpoint uses X-Forwarded-Host and X-Forwarded-Proto to construct correct redirect_uri""" try: from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( @@ -847,6 +897,7 @@ async def test_authorize_endpoint_respects_x_forwarded_host(): "X-Forwarded-Proto": "https", "X-Forwarded-Host": "proxy.example.com", } + set_request_from_trusted_proxy(mock_request) # Mock the encryption functions with patch( @@ -859,7 +910,7 @@ async def test_authorize_endpoint_respects_x_forwarded_host(): request=mock_request, client_id="test_client_id", mcp_server_name="test_oauth", - redirect_uri="https://client.example.com/callback", + redirect_uri="http://127.0.0.1:60108/callback", state="test_state", ) @@ -875,7 +926,9 @@ async def test_authorize_endpoint_respects_x_forwarded_host(): @pytest.mark.asyncio -async def test_token_endpoint_respects_x_forwarded_host(): +async def test_token_endpoint_respects_x_forwarded_host( + trusted_proxy_origin_headers, +): """Test that token endpoint uses X-Forwarded-Host and X-Forwarded-Proto for redirect_uri""" try: from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( @@ -917,6 +970,7 @@ async def test_token_endpoint_respects_x_forwarded_host(): "X-Forwarded-Proto": "https", "X-Forwarded-Host": "proxy.example.com", } + set_request_from_trusted_proxy(mock_request) # Mock httpx client response mock_response = MagicMock() @@ -937,7 +991,7 @@ async def test_token_endpoint_respects_x_forwarded_host(): mock_get_client.return_value = mock_async_client # Call token endpoint - response = await token_endpoint( + await token_endpoint( request=mock_request, grant_type="authorization_code", code="test_code", @@ -1075,7 +1129,12 @@ async def test_token_endpoint_respects_x_forwarded_host(): ], ) def test_get_request_base_url_comprehensive( - base_url, x_forwarded_proto, x_forwarded_host, x_forwarded_port, expected_url + base_url, + x_forwarded_proto, + x_forwarded_host, + x_forwarded_port, + expected_url, + trusted_proxy_origin_headers, ): """Comprehensive test for get_request_base_url with various header combinations""" try: @@ -1089,6 +1148,7 @@ def test_get_request_base_url_comprehensive( # Create mock request mock_request = MagicMock(spec=Request) mock_request.base_url = base_url + set_request_from_trusted_proxy(mock_request) # Build headers dict headers = {} @@ -1116,3 +1176,93 @@ def test_get_request_base_url_comprehensive( f"X-Forwarded-Host={x_forwarded_host}, " f"X-Forwarded-Port={x_forwarded_port}" ) + + +def test_get_request_base_url_ignores_forwarded_headers_from_untrusted_client(): + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + get_request_base_url, + ) + from fastapi import Request + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + mock_request = MagicMock(spec=Request) + mock_request.base_url = "https://gateway.example.com/mcp" + mock_request.headers = { + "X-Forwarded-Proto": "https", + "X-Forwarded-Host": "attacker.example.com", + "X-Forwarded-Port": "443", + } + mock_request.client = MagicMock() + mock_request.client.host = "203.0.113.10" + + with patch( + "litellm.proxy.proxy_server.general_settings", + { + "use_x_forwarded_for": True, + "mcp_trusted_proxy_ranges": TRUSTED_PROXY_RANGES, + }, + create=True, + ): + assert get_request_base_url(mock_request) == "https://gateway.example.com/mcp" + + +def test_validate_trusted_redirect_uri_rejects_spoofed_forwarded_host(): + try: + from litellm.proxy._experimental.mcp_server.oauth_utils import ( + validate_trusted_redirect_uri, + ) + from fastapi import Request + except ImportError: + pytest.skip("MCP OAuth utilities not available") + + mock_request = MagicMock(spec=Request) + mock_request.base_url = "https://gateway.example.com/" + mock_request.headers = { + "X-Forwarded-Proto": "https", + "X-Forwarded-Host": "attacker.example.com", + } + mock_request.client = MagicMock() + mock_request.client.host = "203.0.113.10" + + with ( + patch( + "litellm.proxy.proxy_server.general_settings", + { + "use_x_forwarded_for": True, + "mcp_trusted_proxy_ranges": TRUSTED_PROXY_RANGES, + }, + create=True, + ), + pytest.raises(HTTPException), + ): + validate_trusted_redirect_uri( + mock_request, + "https://attacker.example.com/callback", + ) + + +def test_validate_trusted_redirect_uri_allows_forwarded_origin_from_trusted_proxy( + trusted_proxy_origin_headers, +): + try: + from litellm.proxy._experimental.mcp_server.oauth_utils import ( + validate_trusted_redirect_uri, + ) + from fastapi import Request + except ImportError: + pytest.skip("MCP OAuth utilities not available") + + mock_request = MagicMock(spec=Request) + mock_request.base_url = "http://localhost:4000/" + mock_request.headers = { + "X-Forwarded-Proto": "https", + "X-Forwarded-Host": "proxy.example.com", + } + set_request_from_trusted_proxy(mock_request) + + validate_trusted_redirect_uri( + mock_request, + "https://proxy.example.com/callback", + ) diff --git a/tests/llm_translation/realtime/test_realtime_guardrails_openai.py b/tests/llm_translation/realtime/test_realtime_guardrails_openai.py index 50cedba2ac0..413f5d1ff8b 100644 --- a/tests/llm_translation/realtime/test_realtime_guardrails_openai.py +++ b/tests/llm_translation/realtime/test_realtime_guardrails_openai.py @@ -212,6 +212,8 @@ async def test_text_message_blocked_by_guardrail_no_ai_response(): "policy", "can't repeat", "cannot repeat", + "can't say", + "cannot say", "won't repeat", "can't assist", "can't help", diff --git a/tests/llm_translation/test_gemini.py b/tests/llm_translation/test_gemini.py index bd340aa63be..0a5aebdf91b 100644 --- a/tests/llm_translation/test_gemini.py +++ b/tests/llm_translation/test_gemini.py @@ -17,6 +17,66 @@ from litellm import completion import json +GEMINI_3_IMAGE_SIZE_MAPPINGS = [ + ("512x512", "1:1", "512"), + ("1024x1024", "1:1", "1K"), + ("2048x2048", "1:1", "2K"), + ("4096x4096", "1:1", "4K"), + ("256x1024", "1:4", "512"), + ("512x2048", "1:4", "1K"), + ("1024x4096", "1:4", "2K"), + ("2048x8192", "1:4", "4K"), + ("192x1536", "1:8", "512"), + ("384x3072", "1:8", "1K"), + ("768x6144", "1:8", "2K"), + ("1536x12288", "1:8", "4K"), + ("424x632", "2:3", "512"), + ("848x1264", "2:3", "1K"), + ("1696x2528", "2:3", "2K"), + ("3392x5056", "2:3", "4K"), + ("632x424", "3:2", "512"), + ("1264x848", "3:2", "1K"), + ("2528x1696", "3:2", "2K"), + ("5056x3392", "3:2", "4K"), + ("448x600", "3:4", "512"), + ("896x1200", "3:4", "1K"), + ("1792x2400", "3:4", "2K"), + ("3584x4800", "3:4", "4K"), + ("1024x256", "4:1", "512"), + ("2048x512", "4:1", "1K"), + ("4096x1024", "4:1", "2K"), + ("8192x2048", "4:1", "4K"), + ("600x448", "4:3", "512"), + ("1200x896", "4:3", "1K"), + ("2400x1792", "4:3", "2K"), + ("4800x3584", "4:3", "4K"), + ("464x576", "4:5", "512"), + ("928x1152", "4:5", "1K"), + ("1856x2304", "4:5", "2K"), + ("3712x4608", "4:5", "4K"), + ("576x464", "5:4", "512"), + ("1152x928", "5:4", "1K"), + ("2304x1856", "5:4", "2K"), + ("4608x3712", "5:4", "4K"), + ("1536x192", "8:1", "512"), + ("3072x384", "8:1", "1K"), + ("6144x768", "8:1", "2K"), + ("12288x1536", "8:1", "4K"), + ("384x688", "9:16", "512"), + ("768x1376", "9:16", "1K"), + ("1536x2752", "9:16", "2K"), + ("3072x5504", "9:16", "4K"), + ("688x384", "16:9", "512"), + ("1376x768", "16:9", "1K"), + ("2752x1536", "16:9", "2K"), + ("5504x3072", "16:9", "4K"), + ("792x336", "21:9", "512"), + ("1584x672", "21:9", "1K"), + ("3168x1344", "21:9", "2K"), + ("6336x2688", "21:9", "4K"), +] + + class TestGoogleAIStudioGemini(BaseLLMChatTest): def get_base_completion_call_args(self) -> dict: return {"model": "gemini/gemini-2.5-flash"} @@ -365,6 +425,143 @@ def test_gemini_flash_image_preview_models(model_name: str): ] +@pytest.mark.parametrize( + "model, kwargs, expected_image_config", + [ + ( + "gemini/gemini-3-pro-image-preview", + {"imageConfig": {"aspectRatio": "16:9", "imageSize": "512px"}}, + {"aspectRatio": "16:9", "imageSize": "512px"}, + ), + ( + "gemini/gemini-2.5-flash-image", + {"size": "2048x2048"}, + {"aspectRatio": "1:1"}, + ), + ], +) +def test_gemini_image_generation_forwards_image_config( + model: str, kwargs: dict, expected_image_config: dict +): + from unittest.mock import patch, MagicMock + + with patch( + "litellm.llms.custom_httpx.llm_http_handler.HTTPHandler.post" + ) as mock_post: + mock_http_response = MagicMock() + mock_http_response.json.return_value = { + "candidates": [ + { + "content": { + "parts": [{"inlineData": {"data": "test_base64_image_data"}}] + } + } + ] + } + mock_http_response.status_code = 200 + mock_post.return_value = mock_http_response + + litellm.image_generation( + model=model, + prompt="Generate a simple test image", + api_key="test_api_key", + **kwargs, + ) + + request_data = mock_post.call_args.kwargs.get("json", {}) + assert request_data["generationConfig"]["imageConfig"] == expected_image_config + + +def test_gemini_image_generation_image_config_takes_precedence_over_size(): + from litellm.llms.gemini.image_generation.transformation import GoogleImageGenConfig + + explicit_image_config = {"aspectRatio": "16:9", "imageSize": "2K"} + + mapped_params = GoogleImageGenConfig().map_openai_params( + non_default_params={ + "imageConfig": explicit_image_config, + "size": "768x1376", + }, + optional_params={}, + model="gemini-3-pro-image-preview", + drop_params=False, + ) + + assert mapped_params["imageConfig"] == explicit_image_config + + +def test_gemini_image_generation_ignores_non_dict_image_config(): + from litellm.llms.gemini.image_generation.transformation import GoogleImageGenConfig + + mapped_params = GoogleImageGenConfig().map_openai_params( + non_default_params={ + "size": "768x1376", + "imageConfig": "not-a-dict", + }, + optional_params={}, + model="gemini-3-pro-image-preview", + drop_params=False, + ) + + assert mapped_params["imageConfig"] == {"aspectRatio": "9:16", "imageSize": "1K"} + + +@pytest.mark.parametrize( + "size, expected_aspect_ratio, expected_image_size", + GEMINI_3_IMAGE_SIZE_MAPPINGS, +) +def test_gemini_image_generation_openai_size_maps_to_google_table( + size: str, expected_aspect_ratio: str, expected_image_size: str +): + from litellm.llms.gemini.common_utils import ( + map_openai_size_to_gemini_image_config, + ) + + assert map_openai_size_to_gemini_image_config( + size, "gemini-3-pro-image-preview" + ) == { + "aspectRatio": expected_aspect_ratio, + "imageSize": expected_image_size, + } + + +@pytest.mark.parametrize( + "size, expected_aspect_ratio, expected_image_size", + [ + ("1000x1800", "9:16", "1K"), + ("1800x1000", "16:9", "1K"), + ("3000x3000", "1:1", "2K"), + ("500x500", "1:1", "512"), + ("1280x896", "4:3", "1K"), + ("896x1280", "3:4", "1K"), + ], +) +def test_gemini_image_generation_openai_size_snaps_to_nearest_option( + size: str, expected_aspect_ratio: str, expected_image_size: str +): + from litellm.llms.gemini.common_utils import ( + map_openai_size_to_gemini_image_config, + ) + + assert map_openai_size_to_gemini_image_config( + size, "gemini-3-pro-image-preview" + ) == { + "aspectRatio": expected_aspect_ratio, + "imageSize": expected_image_size, + } + + +@pytest.mark.parametrize("size", ["auto", "invalid", "0x1024", "1024x0"]) +def test_gemini_image_generation_openai_size_auto_uses_google_defaults(size: str): + from litellm.llms.gemini.common_utils import ( + map_openai_size_to_gemini_image_config, + ) + + assert map_openai_size_to_gemini_image_config( + size, "gemini-3-pro-image-preview" + ) is None + + def test_gemini_imagen_models_use_predict_endpoint(): """ Test that Imagen models still use :predict endpoint (not broken by gemini-2.5-flash-image-preview fix) @@ -387,6 +584,7 @@ def test_gemini_imagen_models_use_predict_endpoint(): response = litellm.image_generation( model="gemini/imagen-3.0-generate-001", prompt="Generate a simple test image", + size="1280x896", api_key="test_api_key", ) @@ -410,6 +608,9 @@ def test_gemini_imagen_models_use_predict_endpoint(): request_data = call_args.kwargs.get("json", {}) assert "instances" in request_data assert "parameters" in request_data + assert request_data["parameters"]["aspectRatio"] == "4:3" + assert request_data["parameters"]["imageSize"] == "1K" + assert "imageConfig" not in request_data["parameters"] def test_gemini_thinking(): diff --git a/tests/llm_translation/test_vcr_classification.py b/tests/llm_translation/test_vcr_classification.py index babb3427311..781c37cf9c4 100644 --- a/tests/llm_translation/test_vcr_classification.py +++ b/tests/llm_translation/test_vcr_classification.py @@ -178,9 +178,12 @@ def test_should_distinguish_different_aws_access_keys(): [ ("api.openai.com", True), ("api.anthropic.com", True), + ("bedrock.us-east-1.amazonaws.com", True), ("bedrock-runtime.us-east-1.amazonaws.com", True), ("bedrock-runtime-fips.us-east-1.amazonaws.com", True), ("api.us-east-1.bedrock-runtime.amazonaws.com", False), + ("s3.us-west-2.amazonaws.com", True), + ("litellm-proxy-test.s3.us-west-2.amazonaws.com", True), ("foo.bar.openai.com", True), ("127.0.0.1", False), ("localhost", False), diff --git a/tests/local_testing/conftest.py b/tests/local_testing/conftest.py index abb871789c3..0831313c136 100644 --- a/tests/local_testing/conftest.py +++ b/tests/local_testing/conftest.py @@ -22,13 +22,13 @@ sys.path.insert( ) # Adds the parent directory to the system path import litellm -# ``litellm.model_cost`` is loaded at import time from the URL pinned to -# ``main`` (``LITELLM_MODEL_COST_MAP_URL``). The in-tree backup ships with -# this branch and can include pricing entries that main has not yet picked -# up (e.g. an upstream provider rotates a model id and the test cassette -# records the new name). Backfill any entries that are missing from the -# remote-fetched map so cost-calculator lookups in tests succeed against -# the cassette state the branch is being tested with. +# ``litellm.model_cost`` is loaded at import time from the URL pinned to ``main`` +# (``LITELLM_MODEL_COST_MAP_URL``). The in-tree backup ships with this branch +# and can include pricing entries that ``main`` has not yet picked up (e.g. +# Mistral now returns ``ministral-8b-2512`` from ``mistral-tiny`` and the entry +# was added on this branch). Backfill any entries that are missing from the +# remote-fetched map so cost-calculator lookups in tests succeed against the +# cassette state the branch is being tested with. from litellm.litellm_core_utils.get_model_cost_map import GetModelCostMap for _k, _v in GetModelCostMap.load_local_model_cost_map().items(): diff --git a/tests/local_testing/test_get_llm_provider.py b/tests/local_testing/test_get_llm_provider.py index 14b9e8cd136..46f2f89b3ce 100644 --- a/tests/local_testing/test_get_llm_provider.py +++ b/tests/local_testing/test_get_llm_provider.py @@ -478,3 +478,108 @@ def test_get_llm_provider_use_proxy_arg_true_with_direct_args(): assert key == arg_api_key # Should use the argument key assert base == arg_api_base # Should use the argument base + +# -------- Tests for Claude model pattern matching --------- + + +class TestClaudeModelPatternMatching: + """ + Tests for _matches_claude_model_pattern which routes future Claude models + to the Anthropic provider without requiring model_prices_and_context_window.json updates. + """ + + def test_matches_claude_opus_pattern(self): + """Test claude-opus-X-Y pattern matching.""" + from litellm.litellm_core_utils.get_llm_provider_logic import ( + _matches_claude_model_pattern, + ) + + assert _matches_claude_model_pattern("claude-opus-4-7") is True + assert _matches_claude_model_pattern("claude-opus-4-9") is True + assert _matches_claude_model_pattern("claude-opus-5-1") is True + + def test_matches_claude_sonnet_pattern(self): + """Test claude-sonnet-X-Y pattern matching.""" + from litellm.litellm_core_utils.get_llm_provider_logic import ( + _matches_claude_model_pattern, + ) + + assert _matches_claude_model_pattern("claude-sonnet-4-6") is True + assert _matches_claude_model_pattern("claude-sonnet-5-0") is True + + def test_matches_claude_haiku_pattern(self): + """Test claude-haiku-X-Y pattern matching.""" + from litellm.litellm_core_utils.get_llm_provider_logic import ( + _matches_claude_model_pattern, + ) + + assert _matches_claude_model_pattern("claude-haiku-4-5") is True + assert _matches_claude_model_pattern("claude-haiku-5-0") is True + + def test_matches_claude_with_date_suffix(self): + """Test claude model pattern with date suffix.""" + from litellm.litellm_core_utils.get_llm_provider_logic import ( + _matches_claude_model_pattern, + ) + + assert _matches_claude_model_pattern("claude-opus-5-1-20270101") is True + assert _matches_claude_model_pattern("claude-sonnet-4-7-20260601") is True + assert _matches_claude_model_pattern("claude-haiku-4-6-20251201") is True + + def test_matches_unknown_tier_name(self): + """A tier segment we don't know about today should still route to anthropic. + + The pattern intentionally accepts any ``[a-z]+`` tier rather than a + hard-coded ``opus|sonnet|haiku`` list so a future tier (e.g. a new + "mini" line) is covered without a code change. This guards against a + regression back to hard-coded tier names. + """ + from litellm.litellm_core_utils.get_llm_provider_logic import ( + _matches_claude_model_pattern, + ) + + assert _matches_claude_model_pattern("claude-mini-4-5") is True + assert _matches_claude_model_pattern("claude-neptune-6-0") is True + + def test_rejects_non_claude_models(self): + """Test that non-Claude models are not matched.""" + from litellm.litellm_core_utils.get_llm_provider_logic import ( + _matches_claude_model_pattern, + ) + + assert _matches_claude_model_pattern("gpt-4") is False + assert _matches_claude_model_pattern("mistral-large") is False + assert _matches_claude_model_pattern("llama-3") is False + + def test_rejects_invalid_claude_patterns(self): + """Test that invalid Claude model patterns are not matched.""" + from litellm.litellm_core_utils.get_llm_provider_logic import ( + _matches_claude_model_pattern, + ) + + # Wrong order (variant before name) + assert _matches_claude_model_pattern("claude-4-opus") is False + # Missing version numbers + assert _matches_claude_model_pattern("claude-opus") is False + # Old format (claude-3-opus instead of claude-opus-3) + assert _matches_claude_model_pattern("claude-3-opus-20240229") is False + + def test_get_llm_provider_future_claude_model(self): + """Test that get_llm_provider routes future Claude models to anthropic.""" + model, custom_llm_provider, dynamic_api_key, api_base = ( + litellm.get_llm_provider( + model="claude-opus-4-9", + ) + ) + assert custom_llm_provider == "anthropic" + assert model == "claude-opus-4-9" + + def test_get_llm_provider_future_claude_model_with_date(self): + """Test that get_llm_provider routes future Claude models with date suffix.""" + model, custom_llm_provider, dynamic_api_key, api_base = ( + litellm.get_llm_provider( + model="claude-opus-5-1-20270101", + ) + ) + assert custom_llm_provider == "anthropic" + assert model == "claude-opus-5-1-20270101" diff --git a/tests/local_testing/test_lunary.py b/tests/local_testing/test_lunary.py index d181d24c782..0dbae1b817f 100644 --- a/tests/local_testing/test_lunary.py +++ b/tests/local_testing/test_lunary.py @@ -26,9 +26,6 @@ def test_lunary_logging(): print(e) -test_lunary_logging() - - def test_lunary_template(): import lunary diff --git a/tests/local_testing/test_multiple_deployments.py b/tests/local_testing/test_multiple_deployments.py index f7276d4f14e..72bfd5012c1 100644 --- a/tests/local_testing/test_multiple_deployments.py +++ b/tests/local_testing/test_multiple_deployments.py @@ -49,6 +49,3 @@ def test_multiple_deployments(): except Exception as e: traceback.print_exc() pytest.fail(f"An exception occurred: {e}") - - -test_multiple_deployments() diff --git a/tests/local_testing/test_no_top_level_test_invocations.py b/tests/local_testing/test_no_top_level_test_invocations.py new file mode 100644 index 00000000000..eb1d836a18d --- /dev/null +++ b/tests/local_testing/test_no_top_level_test_invocations.py @@ -0,0 +1,36 @@ +import ast +from pathlib import Path + +LOCAL_TESTING_DIR = Path(__file__).parent + + +def _top_level_test_invocations(tree): + invocations = [] + for node in tree.body: + if not isinstance(node, ast.Expr) or not isinstance(node.value, ast.Call): + continue + func = node.value.func + name = getattr(func, "id", None) or getattr(func, "attr", None) + if name and name.startswith("test_"): + invocations.append((name, node.lineno)) + return invocations + + +def test_no_module_level_test_invocations(): + offenders = [] + for path in sorted(LOCAL_TESTING_DIR.rglob("*.py")): + try: + tree = ast.parse(path.read_text(encoding="utf-8"), filename=str(path)) + except SyntaxError: + continue + for name, lineno in _top_level_test_invocations(tree): + offenders.append( + f"{path.relative_to(LOCAL_TESTING_DIR)}:{lineno} calls {name}()" + ) + + assert not offenders, ( + "Test functions are invoked at module scope, so they run during pytest " + "collection (making network calls and erroring collection for every job " + "that globs this directory). Remove these calls; pytest collects test " + "functions automatically:\n" + "\n".join(offenders) + ) diff --git a/tests/local_testing/test_register_model.py b/tests/local_testing/test_register_model.py index 6b170798874..44fb440bbbd 100644 --- a/tests/local_testing/test_register_model.py +++ b/tests/local_testing/test_register_model.py @@ -1,8 +1,11 @@ #### What this tests #### # This tests calling batch_completions by running 100 messages together +import ast import sys, os import traceback +from pathlib import Path + import pytest sys.path.insert( @@ -62,4 +65,22 @@ def test_update_model_cost_via_completion(): pytest.fail(f"An error occurred: {e}") -test_update_model_cost_via_completion() +def test_no_test_invocation_at_module_scope(): + tree = ast.parse(Path(__file__).read_text()) + defined = { + node.name + for node in tree.body + if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) + } + invoked = [ + node.value.func.id + for node in tree.body + if isinstance(node, ast.Expr) + and isinstance(node.value, ast.Call) + and isinstance(node.value.func, ast.Name) + and node.value.func.id in defined + ] + assert not invoked, ( + f"{invoked} run at import time, so pytest collecting this file fires real " + "provider calls; any failure aborts collection and tears down the whole job" + ) diff --git a/tests/local_testing/test_wandb.py b/tests/local_testing/test_wandb.py index 6cdca40492f..58a9c9f5ddf 100644 --- a/tests/local_testing/test_wandb.py +++ b/tests/local_testing/test_wandb.py @@ -51,9 +51,6 @@ def test_wandb_logging_async(): pass -test_wandb_logging_async() - - def test_wandb_logging(): try: response = completion( diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/mcp_tests/test_mcp_server.py index e65b45fb38b..eea2f2721ab 100644 --- a/tests/mcp_tests/test_mcp_server.py +++ b/tests/mcp_tests/test_mcp_server.py @@ -514,6 +514,69 @@ async def test_sse_mcp_handler_mock(): ) +@pytest.mark.asyncio +async def test_sse_mcp_handler_propagates_passthrough_401(): + """SSE handler must raise 401 + WWW-Authenticate when the upstream + pass-through probe rejects the client's bearer token, instead of letting + the SSE session start and silently return empty tool lists.""" + from fastapi import HTTPException + + from litellm.proxy._types import UserAPIKeyAuth + + mock_scope = { + "type": "http", + "method": "GET", + "path": "/mcp/sse", + "headers": [(b"accept", b"text/event-stream")], + "query_string": b"", + "server": ("localhost", 8000), + "scheme": "http", + } + mock_receive = AsyncMock() + mock_send = AsyncMock() + + mock_auth_result = (UserAPIKeyAuth(), None, None, {}, {}, []) + + challenge = HTTPException( + status_code=401, + detail="Unauthorized", + headers={"WWW-Authenticate": "Bearer authorization_uri=https://example/"}, + ) + + with ( + patch( + "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", + True, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.sse_session_manager", + AsyncMock(), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new=AsyncMock(return_value=mock_auth_result), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.set_auth_context", + ), + patch( + "litellm.proxy._experimental.mcp_server.server._raise_preemptive_401_for_unauthenticated_servers", + new=AsyncMock(), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._check_passthrough_upstream_auth", + new=AsyncMock(side_effect=challenge), + ), + ): + from litellm.proxy._experimental.mcp_server.server import handle_sse_mcp + + with pytest.raises(HTTPException) as excinfo: + await handle_sse_mcp(mock_scope, mock_receive, mock_send) + + assert excinfo.value.status_code == 401 + assert excinfo.value.headers and "WWW-Authenticate" in excinfo.value.headers + + def test_generate_stable_server_id(): """ Test the _generate_stable_server_id method to ensure hash stability across releases. diff --git a/tests/mcp_tests/test_per_user_oauth_cache.py b/tests/mcp_tests/test_per_user_oauth_cache.py index 43e514b32ae..141b906fce9 100644 --- a/tests/mcp_tests/test_per_user_oauth_cache.py +++ b/tests/mcp_tests/test_per_user_oauth_cache.py @@ -183,6 +183,31 @@ class TestValidateTokenResponse: server_id="atlassian", ) + def test_boolean_value_matches_lowercase_string_rule(self): + """Boolean ``True`` in token response must match the JSON-style rule ``"true"``. + + Admin config is typically written as ``{"verified": "true"}`` (lower-case + from JSON / YAML), but the OAuth response returns ``{"verified": true}`` + (Python ``True``). The normaliser must align them. + """ + _validate_token_response = _import_validate() + token_response = {"access_token": "tok", "verified": True} + # Should not raise + _validate_token_response( + token_response=token_response, + validation_rules={"verified": "true"}, + server_id="test", + ) + + def test_boolean_false_matches_lowercase_string_rule(self): + _validate_token_response = _import_validate() + token_response = {"access_token": "tok", "is_admin": False} + _validate_token_response( + token_response=token_response, + validation_rules={"is_admin": "false"}, + server_id="test", + ) + # ── _compute_per_user_token_ttl ────────────────────────────────────────────── diff --git a/tests/pass_through_tests/test_gemini_with_spend.test.js b/tests/pass_through_tests/test_gemini_with_spend.test.js index 989bbc4b8e3..b9a25d3a3ed 100644 --- a/tests/pass_through_tests/test_gemini_with_spend.test.js +++ b/tests/pass_through_tests/test_gemini_with_spend.test.js @@ -32,7 +32,7 @@ describe('Gemini AI Tests', () => { }; const model = genAI.getGenerativeModel({ - model: 'gemini-2.5-flash-lite' + model: 'gemini-3.1-flash-lite' }, requestOptions); const prompt = 'Say "hello test" and nothing else'; @@ -83,7 +83,7 @@ describe('Gemini AI Tests', () => { }; const model = genAI.getGenerativeModel({ - model: 'gemini-2.5-flash-lite' + model: 'gemini-3.1-flash-lite' }, requestOptions); const prompt = 'Say "hello test" and nothing else'; diff --git a/tests/pass_through_tests/test_local_gemini.js b/tests/pass_through_tests/test_local_gemini.js index 0a72ca5cd7b..dc033a51f18 100644 --- a/tests/pass_through_tests/test_local_gemini.js +++ b/tests/pass_through_tests/test_local_gemini.js @@ -1,13 +1,13 @@ const { GoogleGenerativeAI, ModelParams, RequestOptions } = require("@google/generative-ai"); const modelParams = { - model: 'gemini-2.5-flash-lite', + model: 'gemini-3.1-flash-lite', }; const requestOptions = { baseUrl: 'http://127.0.0.1:4000/gemini', customHeaders: { - "tags": "gemini-js-sdk,gemini-2.5-flash-lite" + "tags": "gemini-js-sdk,gemini-3.1-flash-lite" } }; diff --git a/tests/pass_through_tests/test_local_vertex.js b/tests/pass_through_tests/test_local_vertex.js index 149635e2d6f..7cfe31db95b 100644 --- a/tests/pass_through_tests/test_local_vertex.js +++ b/tests/pass_through_tests/test_local_vertex.js @@ -4,7 +4,7 @@ const { VertexAI, RequestOptions } = require('@google-cloud/vertexai'); const vertexAI = new VertexAI({ project: 'litellm-ci-cd', - location: 'us-central1', + location: 'global', apiEndpoint: "127.0.0.1:4000/vertex-ai" }); @@ -20,7 +20,7 @@ const requestOptions = { }; const generativeModel = vertexAI.getGenerativeModel( - { model: 'gemini-2.5-flash-lite' }, + { model: 'gemini-3.1-flash-lite' }, requestOptions ); diff --git a/tests/pass_through_tests/test_vertex.test.js b/tests/pass_through_tests/test_vertex.test.js index 7b5edf6acd7..e0e879c2897 100644 --- a/tests/pass_through_tests/test_vertex.test.js +++ b/tests/pass_through_tests/test_vertex.test.js @@ -56,6 +56,9 @@ beforeAll(() => { loadVertexAiCredentials(); }); +// Configure Jest to retry flaky tests up to 3 times (useful for 429 rate limiting) +jest.retryTimes(3); + // Non-streaming Vertex generateContent can exceed 5s in CI / under load const VERTEX_TEST_TIMEOUT_MS = 30000; @@ -65,7 +68,7 @@ describe('Vertex AI Tests', () => { async () => { const vertexAI = new VertexAI({ project: 'litellm-ci-cd', - location: 'us-central1', + location: 'global', apiEndpoint: "localhost:4000/vertex-ai" }); @@ -78,7 +81,7 @@ describe('Vertex AI Tests', () => { }; const generativeModel = vertexAI.getGenerativeModel( - { model: 'gemini-2.5-flash-lite' }, + { model: 'gemini-3.1-flash-lite' }, requestOptions ); @@ -108,13 +111,13 @@ describe('Vertex AI Tests', () => { async () => { const vertexAI = new VertexAI({ project: 'litellm-ci-cd', - location: 'us-central1', + location: 'global', apiEndpoint: "localhost:4000/vertex-ai" }); const customHeaders = new Headers({"x-litellm-api-key": "sk-1234"}); const requestOptions = {customHeaders: customHeaders}; const generativeModel = vertexAI.getGenerativeModel( - {model: 'gemini-2.5-flash-lite'}, + {model: 'gemini-3.1-flash-lite'}, requestOptions ); const request = {contents: [{role: 'user', parts: [{text: 'What is 2+2?'}]}]}; diff --git a/tests/pass_through_tests/test_vertex_ai.py b/tests/pass_through_tests/test_vertex_ai.py index 73bf03c5000..bf1200489aa 100644 --- a/tests/pass_through_tests/test_vertex_ai.py +++ b/tests/pass_through_tests/test_vertex_ai.py @@ -103,12 +103,12 @@ async def test_basic_vertex_ai_pass_through_with_spendlog(): vertexai.init( project="litellm-ci-cd", - location="us-central1", + location="global", api_endpoint=f"{LITE_LLM_ENDPOINT}/vertex_ai", api_transport="rest", ) - model = GenerativeModel(model_name="gemini-2.5-flash-lite") + model = GenerativeModel(model_name="gemini-3.1-flash-lite") response = model.generate_content("hi") print("response", response) @@ -143,12 +143,12 @@ async def test_basic_vertex_ai_pass_through_streaming_with_spendlog(): vertexai.init( project="litellm-ci-cd", - location="us-central1", + location="global", api_endpoint=f"{LITE_LLM_ENDPOINT}/vertex_ai", api_transport="rest", ) - model = GenerativeModel(model_name="gemini-2.5-flash-lite") + model = GenerativeModel(model_name="gemini-3.1-flash-lite") response = model.generate_content("hi", stream=True) for chunk in response: @@ -182,7 +182,7 @@ async def test_vertex_ai_pass_through_endpoint_context_caching(): vertexai.init( project="litellm-ci-cd", - location="us-central1", + location="global", api_endpoint=f"{LITE_LLM_ENDPOINT}/vertex_ai", api_transport="rest", ) @@ -204,7 +204,7 @@ async def test_vertex_ai_pass_through_endpoint_context_caching(): ] cached_content = caching.CachedContent.create( - model_name="gemini-2.5-flash-lite-001", + model_name="gemini-3.1-flash-lite", system_instruction=system_instruction, contents=contents, ttl=datetime.timedelta(minutes=60), diff --git a/tests/pass_through_tests/test_vertex_with_spend.test.js b/tests/pass_through_tests/test_vertex_with_spend.test.js index 142a1cec8ff..4dee890dc78 100644 --- a/tests/pass_through_tests/test_vertex_with_spend.test.js +++ b/tests/pass_through_tests/test_vertex_with_spend.test.js @@ -71,7 +71,7 @@ describe('Vertex AI Tests', () => { test('should successfully generate non-streaming content with tags', async () => { const vertexAI = new VertexAI({ project: 'litellm-ci-cd', - location: 'us-central1', + location: 'global', apiEndpoint: "127.0.0.1:4000/vertex_ai" }); @@ -85,7 +85,7 @@ describe('Vertex AI Tests', () => { }; const generativeModel = vertexAI.getGenerativeModel( - { model: 'gemini-2.5-flash-lite' }, + { model: 'gemini-3.1-flash-lite' }, requestOptions ); @@ -130,7 +130,7 @@ describe('Vertex AI Tests', () => { test('should successfully generate streaming content with tags', async () => { const vertexAI = new VertexAI({ project: 'litellm-ci-cd', - location: 'us-central1', + location: 'global', apiEndpoint: "127.0.0.1:4000/vertex_ai" }); @@ -144,7 +144,7 @@ describe('Vertex AI Tests', () => { }; const generativeModel = vertexAI.getGenerativeModel( - { model: 'gemini-2.5-flash-lite' }, + { model: 'gemini-3.1-flash-lite' }, requestOptions ); diff --git a/tests/proxy_unit_tests/test_auth_checks.py b/tests/proxy_unit_tests/test_auth_checks.py index d9f4a6e56b8..e7136ecb195 100644 --- a/tests/proxy_unit_tests/test_auth_checks.py +++ b/tests/proxy_unit_tests/test_auth_checks.py @@ -38,8 +38,12 @@ from litellm.proxy.utils import CallInfo @pytest.mark.asyncio async def test_get_end_user_object(customer_spend, customer_budget): """ - Scenario 1: normal - Scenario 2: user over budget + Scenario 1: normal - get_end_user_object returns the cached user + Scenario 2: user over budget - NOTE: budget enforcement now happens in + common_checks() via _check_end_user_budget(), not in get_end_user_object() + + This test verifies that get_end_user_object correctly retrieves the end user + from cache. Budget enforcement is tested separately in test_check_end_user_budget(). """ end_user_id = "my-test-customer" _budget = LiteLLM_BudgetTable(max_budget=customer_budget) @@ -58,31 +62,62 @@ async def test_get_end_user_object(customer_spend, customer_budget): value=end_user_obj, model_type=LiteLLM_EndUserTable, ) + # get_end_user_object only fetches data - it no longer enforces budget + # Budget enforcement happens in common_checks() via _check_end_user_budget() + result = await get_end_user_object( + end_user_id=end_user_id, + prisma_client="RANDOM VALUE", # type: ignore + user_api_key_cache=_cache, + route="/v1/chat/completions", + ) + assert result is not None + assert result.user_id == end_user_id + + +@pytest.mark.parametrize("customer_spend, customer_budget", [(0, 10), (10, 0)]) +@pytest.mark.asyncio +async def test_check_end_user_budget(customer_spend, customer_budget): + """ + Test _check_end_user_budget enforcement: + - Scenario 1: customer_spend=0, customer_budget=10 - should pass (under budget) + - Scenario 2: customer_spend=10, customer_budget=0 - should fail (over budget) + + Note: Budget enforcement for end users happens in common_checks() via + _check_end_user_budget(), not in get_end_user_object(). + """ + from litellm.proxy.auth.auth_checks import _check_end_user_budget + + _budget = LiteLLM_BudgetTable(max_budget=customer_budget) + end_user_obj = LiteLLM_EndUserTable( + user_id="my-test-customer", + spend=customer_spend, + litellm_budget_table=_budget, + blocked=False, + ) + + should_exceed = customer_spend > customer_budget + try: - await get_end_user_object( - end_user_id=end_user_id, - prisma_client="RANDOM VALUE", # type: ignore - user_api_key_cache=_cache, + await _check_end_user_budget( + end_user_obj=end_user_obj, route="/v1/chat/completions", ) - if customer_spend > customer_budget: + if should_exceed: pytest.fail( - "Expected call to fail. Customer Spend={}, Customer Budget={}".format( + "Expected BudgetExceededError. Customer Spend={}, Customer Budget={}".format( customer_spend, customer_budget ) ) - except Exception as e: - if ( - isinstance(e, litellm.BudgetExceededError) - and customer_spend > customer_budget - ): - pass - else: + except litellm.BudgetExceededError as e: + if not should_exceed: pytest.fail( - "Expected call to work. Customer Spend={}, Customer Budget={}, Error={}".format( + "Unexpected BudgetExceededError. Customer Spend={}, Customer Budget={}, Error={}".format( customer_spend, customer_budget, str(e) ) ) + # Verify the error has correct info + assert e.current_cost == customer_spend + assert e.max_budget == customer_budget @pytest.mark.parametrize( diff --git a/tests/proxy_unit_tests/test_default_end_user_budget_simple.py b/tests/proxy_unit_tests/test_default_end_user_budget_simple.py index 970a7ab4718..6170b0a972e 100644 --- a/tests/proxy_unit_tests/test_default_end_user_budget_simple.py +++ b/tests/proxy_unit_tests/test_default_end_user_budget_simple.py @@ -134,9 +134,14 @@ async def test_explicit_budget_not_overridden_by_default(): @pytest.mark.asyncio async def test_budget_enforcement_blocks_over_budget_users(): """ - Core scenario: Budget limits are actually enforced. + Core scenario: Budget limits are actually enforced via _check_end_user_budget. Users who exceed their budget should be blocked. + + Note: Budget enforcement happens in common_checks() via _check_end_user_budget(), + not in get_end_user_object(). get_end_user_object only fetches the user data. """ + from litellm.proxy.auth.auth_checks import _check_end_user_budget + end_user_id = f"test_user_{uuid.uuid4().hex}" default_budget_id = str(uuid.uuid4()) litellm.max_end_user_budget_id = default_budget_id @@ -170,12 +175,23 @@ async def test_budget_enforcement_blocks_over_budget_users(): mock_cache.async_get_cache = AsyncMock(return_value=None) mock_cache.async_set_cache = AsyncMock() - # Should raise BudgetExceededError + # First, get the end user object (this just fetches data, doesn't enforce budget) + result = await get_end_user_object( + end_user_id=end_user_id, + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + route="/chat/completions", + ) + + # Verify user was fetched with default budget applied + assert result is not None + assert result.litellm_budget_table is not None + assert result.litellm_budget_table.max_budget == 10.0 + + # Now test budget enforcement separately via _check_end_user_budget with pytest.raises(litellm.BudgetExceededError) as exc_info: - await get_end_user_object( - end_user_id=end_user_id, - prisma_client=mock_prisma_client, - user_api_key_cache=mock_cache, + await _check_end_user_budget( + end_user_obj=result, route="/chat/completions", ) diff --git a/tests/proxy_unit_tests/test_jwt.py b/tests/proxy_unit_tests/test_jwt.py index 9a8d6d37020..92209e11315 100644 --- a/tests/proxy_unit_tests/test_jwt.py +++ b/tests/proxy_unit_tests/test_jwt.py @@ -2,6 +2,8 @@ # Unit tests for JWT-Auth import asyncio +import base64 +import logging import os import random import sys @@ -21,6 +23,9 @@ from datetime import datetime, timedelta from unittest.mock import AsyncMock, MagicMock, patch import pytest +import jwt +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import rsa from fastapi import Request, HTTPException from fastapi.routing import APIRoute from fastapi.responses import Response @@ -35,7 +40,7 @@ from litellm.proxy._types import ( from litellm.proxy.auth.handle_jwt import JWTHandler, JWTAuthManager from litellm.proxy.management_endpoints.team_endpoints import new_team from litellm.proxy.proxy_server import chat_completion -from typing import Literal +from typing import Literal, Optional public_key = { "kty": "RSA", @@ -1584,3 +1589,524 @@ async def test_auth_jwt_mismatched_key_fails(monkeypatch): with pytest.raises(Exception) as exc: await h.auth_jwt(token) assert "Validation fails" in str(exc.value) + + +def _base64url_encode_bytes(value: bytes) -> str: + return base64.urlsafe_b64encode(value).rstrip(b"=").decode() + + +def _base64url_encode_int(value: int) -> str: + value_bytes = value.to_bytes((value.bit_length() + 7) // 8, "big") + return _base64url_encode_bytes(value=value_bytes) + + +def _get_rsa_key_and_jwk(kid: str): + private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + public_numbers = private_key.public_key().public_numbers() + jwk = { + "kty": "RSA", + "n": _base64url_encode_int(value=public_numbers.n), + "e": _base64url_encode_int(value=public_numbers.e), + "kid": kid, + "alg": "RS256", + "use": "sig", + } + return private_key, jwk + + +def _encode_rsa_jwt( + private_key, + issuer: str, + audience: str, + kid: str, + extra_claims: Optional[dict] = None, +) -> str: + private_key_pem = private_key.private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.PKCS8, + encryption_algorithm=serialization.NoEncryption(), + ) + current_time = int(time.time()) + claims = { + "sub": "test-subject", + "iss": issuer, + "aud": audience, + "iat": current_time, + "exp": current_time + 300, + } + if extra_claims: + claims.update(extra_claims) + + return jwt.encode( + claims, + private_key_pem, + algorithm="RS256", + headers={"kid": kid}, + ) + + +def _get_jwt_handler_with_issuer_keys(issuers: list, keys_by_url: dict) -> JWTHandler: + cache = DualCache() + for jwks_url, keys in keys_by_url.items(): + cache.set_cache( + key=f"litellm_jwt_auth_keys_{jwks_url}", + value=keys, + ) + + jwt_handler = JWTHandler() + jwt_handler.update_environment( + prisma_client=None, + user_api_key_cache=cache, + litellm_jwtauth=LiteLLM_JWTAuth(issuers=issuers), + ) + return jwt_handler + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_validates_selected_issuer_and_maps_claims( + monkeypatch, +): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer_one = "https://issuer-one.example.com" + issuer_two = "https://issuer-two.example.com" + issuer_one_jwks_url = f"{issuer_one}/keys" + issuer_two_jwks_url = f"{issuer_two}/keys" + shared_kid = "shared-kid" + + _, issuer_one_jwk = _get_rsa_key_and_jwk(kid=shared_kid) + issuer_two_private_key, issuer_two_jwk = _get_rsa_key_and_jwk(kid=shared_kid) + + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer_one, + "jwks_url": issuer_one_jwks_url, + "audience": "audience-one", + "user_id_jwt_field": "email", + "user_email_jwt_field": "email", + }, + { + "issuer": issuer_two, + "jwks_url": issuer_two_jwks_url, + "audience": "audience-two", + "user_id_jwt_field": "repository_owner", + "team_id_jwt_field": "repository", + }, + ], + keys_by_url={ + issuer_one_jwks_url: [issuer_one_jwk], + issuer_two_jwks_url: [issuer_two_jwk], + }, + ) + + token = _encode_rsa_jwt( + private_key=issuer_two_private_key, + issuer=issuer_two, + audience="audience-two", + kid=shared_kid, + extra_claims={ + "repository_owner": "example-org", + "repository": "example-org/litellm-fork", + }, + ) + + claims = await jwt_handler.auth_jwt(token=token) + + assert claims[JWTHandler.LITELLM_JWT_ISSUER_CLAIM] == issuer_two + assert jwt_handler.get_user_id(token=claims, default_value=None) == ("example-org") + assert jwt_handler.get_team_id(token=claims, default_value=None) == ( + "example-org/litellm-fork" + ) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_maps_kubernetes_namespace_claim(monkeypatch): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer = "https://oidc.eks.eu-west-1.amazonaws.com/id/test-cluster" + jwks_url = f"{issuer}/keys" + private_key, jwk = _get_rsa_key_and_jwk(kid="k8s-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + "audience": None, + "disable_audience_validation": True, + "user_id_jwt_field": "kubernetes\\.io.namespace", + } + ], + keys_by_url={jwks_url: [jwk]}, + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer=issuer, + audience="kubernetes.default.svc", + kid="k8s-key", + extra_claims={"kubernetes.io": {"namespace": "example-namespace"}}, + ) + + claims = await jwt_handler.auth_jwt(token=token) + + assert ( + jwt_handler.get_user_id(token=claims, default_value=None) == "example-namespace" + ) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_falls_back_to_global_jwks_for_unknown_issuer( + monkeypatch, +): + """Unknown ``iss`` claims fall through to the global ``JWT_PUBLIC_KEY_URL`` + path so adding the new ``issuers`` config to a live deployment doesn't + break tokens minted by issuers that still rely on the legacy global JWKS. + """ + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + configured_issuer = "https://issuer.example.com" + unknown_issuer = "https://unknown-issuer.example.com" + global_jwks_url = "https://global.example.com/keys" + monkeypatch.setenv("JWT_PUBLIC_KEY_URL", global_jwks_url) + + configured_private_key, configured_jwk = _get_rsa_key_and_jwk(kid="configured-key") + unknown_private_key, unknown_jwk = _get_rsa_key_and_jwk(kid="global-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": configured_issuer, + "jwks_url": f"{configured_issuer}/keys", + "audience": "expected-audience", + } + ], + keys_by_url={ + f"{configured_issuer}/keys": [configured_jwk], + global_jwks_url: [unknown_jwk], + }, + ) + token = _encode_rsa_jwt( + private_key=unknown_private_key, + issuer=unknown_issuer, + audience="expected-audience", + kid="global-key", + ) + + claims = await jwt_handler.auth_jwt(token=token) + + assert claims["iss"] == unknown_issuer + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_unknown_issuer_without_global_jwks_rejected( + monkeypatch, +): + """When there is no ``JWT_PUBLIC_KEY_URL`` to fall back to, an unknown + ``iss`` claim still fails — the fallback path raises ``Missing JWT + Public Key URL`` rather than the legacy ``Unsupported JWT issuer``. + """ + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + configured_issuer = "https://issuer.example.com" + private_key, jwk = _get_rsa_key_and_jwk(kid="issuer-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": configured_issuer, + "jwks_url": f"{configured_issuer}/keys", + "audience": "expected-audience", + } + ], + keys_by_url={f"{configured_issuer}/keys": [jwk]}, + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer="https://unknown-issuer.example.com", + audience="expected-audience", + kid="issuer-key", + ) + + with pytest.raises(Exception) as exc: + await jwt_handler.auth_jwt(token=token) + + assert "Missing JWT Public Key URL" in str(exc.value) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_rejects_wrong_audience(monkeypatch): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer = "https://issuer.example.com" + jwks_url = f"{issuer}/keys" + private_key, jwk = _get_rsa_key_and_jwk(kid="issuer-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + "audience": "expected-audience", + } + ], + keys_by_url={jwks_url: [jwk]}, + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer=issuer, + audience="wrong-audience", + kid="issuer-key", + ) + + with pytest.raises(Exception) as exc: + await jwt_handler.auth_jwt(token=token) + + assert "Validation fails" in str(exc.value) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_same_kid_does_not_cross_issuer_keys(monkeypatch): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer_one = "https://issuer-one.example.com" + issuer_two = "https://issuer-two.example.com" + issuer_one_jwks_url = f"{issuer_one}/keys" + issuer_two_jwks_url = f"{issuer_two}/keys" + shared_kid = "shared-kid" + issuer_one_private_key, issuer_one_jwk = _get_rsa_key_and_jwk(kid=shared_kid) + _, issuer_two_jwk = _get_rsa_key_and_jwk(kid=shared_kid) + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer_one, + "jwks_url": issuer_one_jwks_url, + "audience": "audience-one", + }, + { + "issuer": issuer_two, + "jwks_url": issuer_two_jwks_url, + "audience": "audience-two", + }, + ], + keys_by_url={ + issuer_one_jwks_url: [issuer_one_jwk], + issuer_two_jwks_url: [issuer_two_jwk], + }, + ) + token = _encode_rsa_jwt( + private_key=issuer_one_private_key, + issuer=issuer_two, + audience="audience-two", + kid=shared_kid, + ) + + with pytest.raises(Exception) as exc: + await jwt_handler.auth_jwt(token=token) + + assert "Validation fails" in str(exc.value) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_missing_mapped_claim_is_optional(monkeypatch): + """Configured issuer claim mappings are advisory, not mandatory. + + When the token simply omits a mapped field (e.g. a service-to-service token + with no ``email`` claim), JWT auth still succeeds and the normalized claim + is just absent — matching the global ``litellm_jwtauth`` behaviour. + """ + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer = "https://issuer.example.com" + jwks_url = f"{issuer}/keys" + private_key, jwk = _get_rsa_key_and_jwk(kid="issuer-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + "audience": "expected-audience", + "user_id_jwt_field": "email", + } + ], + keys_by_url={jwks_url: [jwk]}, + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer=issuer, + audience="expected-audience", + kid="issuer-key", + ) + + claims = await jwt_handler.auth_jwt(token=token) + + assert claims[JWTHandler.LITELLM_JWT_ISSUER_CLAIM] == issuer + assert JWTHandler.LITELLM_USER_ID_CLAIM not in claims + + +def test_multi_issuer_jwt_requires_audience_unless_explicitly_disabled( + monkeypatch, +): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer = "https://issuer.example.com" + jwks_url = f"{issuer}/keys" + + with pytest.raises(Exception) as exc: + LiteLLM_JWTAuth( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + } + ] + ) + + assert "must configure audience" in str(exc.value) + + +@pytest.mark.asyncio +async def test_global_jwt_ignores_user_supplied_internal_claims(monkeypatch): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_ISSUER", raising=False) + + jwks_url = "https://global-issuer.example.com/keys" + monkeypatch.setenv("JWT_PUBLIC_KEY_URL", jwks_url) + + private_key, jwk = _get_rsa_key_and_jwk(kid="global-key") + cache = DualCache() + cache.set_cache(key=f"litellm_jwt_auth_keys_{jwks_url}", value=[jwk]) + + jwt_handler = JWTHandler() + jwt_handler.update_environment( + prisma_client=None, + user_api_key_cache=cache, + litellm_jwtauth=LiteLLM_JWTAuth( + user_id_jwt_field="email", + user_email_jwt_field="email", + team_id_jwt_field="team.id", + team_ids_jwt_field="teams", + org_id_jwt_field="org.id", + end_user_id_jwt_field="end_user.id", + ), + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer="https://global-issuer.example.com", + audience="some-other-client", + kid="global-key", + extra_claims={ + "email": "real-user@example.com", + "team": {"id": "real-team"}, + "teams": ["real-team", "secondary-team"], + "org": {"id": "real-org"}, + "end_user": {"id": "real-end-user"}, + JWTHandler.LITELLM_JWT_ISSUER_CLAIM: "https://issuer.example.com", + JWTHandler.LITELLM_USER_ID_CLAIM: "victim-user", + JWTHandler.LITELLM_USER_EMAIL_CLAIM: "victim@example.com", + JWTHandler.LITELLM_TEAM_ID_CLAIM: "victim-team", + JWTHandler.LITELLM_TEAM_IDS_CLAIM: ["victim-team"], + JWTHandler.LITELLM_ORG_ID_CLAIM: "victim-org", + JWTHandler.LITELLM_END_USER_ID_CLAIM: "victim-end-user", + }, + ) + + claims = await jwt_handler.auth_jwt(token=token) + + assert jwt_handler.get_user_id(token=claims, default_value=None) == ( + "real-user@example.com" + ) + assert jwt_handler.get_user_email(token=claims, default_value=None) == ( + "real-user@example.com" + ) + assert jwt_handler.get_team_id(token=claims, default_value=None) == "real-team" + assert jwt_handler.get_team_ids_from_jwt(token=claims) == [ + "real-team", + "secondary-team", + ] + assert jwt_handler.get_org_id(token=claims, default_value=None) == "real-org" + assert jwt_handler.get_end_user_id(token=claims, default_value=None) == ( + "real-end-user" + ) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_strips_unmapped_internal_claims(monkeypatch): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer = "https://issuer.example.com" + jwks_url = f"{issuer}/keys" + private_key, jwk = _get_rsa_key_and_jwk(kid="issuer-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + "audience": "expected-audience", + "user_email_jwt_field": "email", + } + ], + keys_by_url={jwks_url: [jwk]}, + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer=issuer, + audience="expected-audience", + kid="issuer-key", + extra_claims={ + "email": "real-user@example.com", + JWTHandler.LITELLM_USER_ID_CLAIM: "victim-user", + JWTHandler.LITELLM_TEAM_ID_CLAIM: "victim-team", + }, + ) + + claims = await jwt_handler.auth_jwt(token=token) + + assert JWTHandler.LITELLM_USER_ID_CLAIM not in claims + assert JWTHandler.LITELLM_TEAM_ID_CLAIM not in claims + assert jwt_handler.get_user_id(token=claims, default_value=None) is None + assert jwt_handler.get_team_id(token=claims, default_value=None) is None + assert jwt_handler.get_user_email(token=claims, default_value=None) == ( + "real-user@example.com" + ) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_does_not_emit_unscoped_global_warning( + monkeypatch, caplog +): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_ISSUER", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + JWTHandler._unscoped_jwt_warning_emitted = False + + issuer = "https://issuer.example.com" + jwks_url = f"{issuer}/keys" + private_key, jwk = _get_rsa_key_and_jwk(kid="issuer-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + "audience": "expected-audience", + } + ], + keys_by_url={jwks_url: [jwk]}, + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer=issuer, + audience="expected-audience", + kid="issuer-key", + ) + + with caplog.at_level(logging.WARNING): + await jwt_handler.auth_jwt(token=token) + + assert "Tokens minted by any application" not in caplog.text + assert JWTHandler._unscoped_jwt_warning_emitted is False diff --git a/tests/proxy_unit_tests/test_proxy_server.py b/tests/proxy_unit_tests/test_proxy_server.py index 9c08175767d..6fdabe64e25 100644 --- a/tests/proxy_unit_tests/test_proxy_server.py +++ b/tests/proxy_unit_tests/test_proxy_server.py @@ -804,6 +804,7 @@ def test_img_gen(mock_aimage_generation, client_no_auth): "prompt": "A cute baby sea otter", "n": 1, "size": "1024x1024", + "imageConfig": {"aspectRatio": "9:16", "imageSize": "1K"}, } response = client_no_auth.post("/v1/images/generations", json=test_data) @@ -813,6 +814,7 @@ def test_img_gen(mock_aimage_generation, client_no_auth): prompt="A cute baby sea otter", n=1, size="1024x1024", + imageConfig={"aspectRatio": "9:16", "imageSize": "1K"}, metadata=mock.ANY, proxy_server_request=mock.ANY, secret_fields=mock.ANY, diff --git a/tests/test_litellm/a2a_protocol/test_send_message_response.py b/tests/test_litellm/a2a_protocol/test_send_message_response.py new file mode 100644 index 00000000000..832aa288c7a --- /dev/null +++ b/tests/test_litellm/a2a_protocol/test_send_message_response.py @@ -0,0 +1,43 @@ +"""Tests for LiteLLMSendMessageResponse JSON-RPC normalization.""" + +from litellm.types.agents import LiteLLMSendMessageResponse + + +def test_from_dict_backfills_id_on_agent_error_response(): + agent_error = { + "jsonrpc": "2.0", + "error": {"code": -32054, "message": "Session not found"}, + } + + response = LiteLLMSendMessageResponse.from_dict( + agent_error, request_id="r1" + ) + + assert response.id == "r1" + assert response.error == {"code": -32054, "message": "Session not found"} + assert response.result is None + + +def test_from_dict_preserves_existing_id(): + payload = { + "id": "upstream-id", + "jsonrpc": "2.0", + "error": {"code": -32001, "message": "Task not found"}, + } + + response = LiteLLMSendMessageResponse.from_dict( + payload, request_id="r1" + ) + + assert response.id == "upstream-id" + + +def test_from_dict_without_request_id_still_requires_id(): + try: + LiteLLMSendMessageResponse.from_dict( + {"jsonrpc": "2.0", "error": {"code": -32054, "message": "x"}} + ) + except Exception as exc: + assert "id" in str(exc).lower() + else: + raise AssertionError("expected validation error when id and request_id missing") diff --git a/tests/test_litellm/completion_extras/__init__.py b/tests/test_litellm/completion_extras/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/completion_extras/test_responses_bridge_provider_propagation.py b/tests/test_litellm/completion_extras/test_responses_bridge_provider_propagation.py new file mode 100644 index 00000000000..b41dbd54b85 --- /dev/null +++ b/tests/test_litellm/completion_extras/test_responses_bridge_provider_propagation.py @@ -0,0 +1,116 @@ +""" +Regression test for https://github.com/BerriAI/litellm/issues/28505 - +the Responses API bridge double-strips the provider prefix from the +model name when a Chat Completions request has both `tools` and +`reasoning_effort`. + +Root cause: the bridge handler called `litellm.responses()` / +`litellm.aresponses()` without passing the already-resolved +`custom_llm_provider`. The downstream call then re-invoked +`get_llm_provider()` with `custom_llm_provider=None`, which stripped +a second provider prefix from a `provider/provider/model` deployment +string. + +This test pins both the sync and async bridge handler call sites: +the resolved `custom_llm_provider` must be forwarded to the underlying +`responses` / `aresponses` call so the provider isn't re-detected. +""" + +from unittest.mock import MagicMock, patch + +import pytest + +from litellm.completion_extras.litellm_responses_transformation.handler import ( + ResponsesToCompletionBridgeHandler, +) + + +def _validated_kwargs(): + return { + "model": "openai/openai/openai/gpt-5.5", + "messages": [{"role": "user", "content": "hi"}], + "optional_params": {}, + "litellm_params": {}, + "headers": {}, + "model_response": MagicMock(), + "logging_obj": MagicMock(), + "custom_llm_provider": "openai", + } + + +def test_sync_completion_forwards_custom_llm_provider(): + handler = ResponsesToCompletionBridgeHandler() + handler.transformation_handler = MagicMock() + handler.transformation_handler.transform_request.return_value = { + "model": "openai/openai/openai/gpt-5.5", + "input": [], + # `_build_sanitized_litellm_params` spreads `custom_llm_provider` from + # `litellm_params` into request_data on the real bridge path. Seed + # it here so the test exercises the overwrite (not an explicit kwarg + # that would TypeError against an already-present key). + "custom_llm_provider": "should-be-overwritten", + } + handler.transformation_handler.transform_response.return_value = ( + _validated_kwargs()["model_response"] + ) + with ( + patch.object( + handler, "validate_input_kwargs", return_value=_validated_kwargs() + ), + patch( + "litellm.responses", + return_value=MagicMock(spec=[]), + ) as mock_responses, + ): + # The handler routes ResponsesAPIResponse through transform_response. + # We just want to verify the kwargs going INTO responses(). + try: + handler.completion(acompletion=False) + except Exception: + # Downstream handling (transform_response, type checks) is not + # the subject of this test. + pass + assert mock_responses.called + kwargs = mock_responses.call_args.kwargs + assert kwargs.get("custom_llm_provider") == "openai", ( + "sync bridge must forward custom_llm_provider to litellm.responses() " + "so the downstream get_llm_provider() call does not re-strip the " + "provider prefix on a provider/provider/model deployment string" + ) + + +@pytest.mark.asyncio +async def test_async_completion_forwards_custom_llm_provider(): + handler = ResponsesToCompletionBridgeHandler() + handler.transformation_handler = MagicMock() + handler.transformation_handler.transform_request.return_value = { + "model": "openai/openai/openai/gpt-5.5", + "input": [], + # `_build_sanitized_litellm_params` spreads `custom_llm_provider` from + # `litellm_params` into request_data on the real bridge path. Seed + # it here so the test exercises the overwrite (not an explicit kwarg + # that would TypeError against an already-present key). + "custom_llm_provider": "should-be-overwritten", + } + + async def _fake_aresponses(**kwargs): + _fake_aresponses.kwargs = kwargs + return MagicMock(spec=[]) + + _fake_aresponses.kwargs = {} + + with ( + patch.object( + handler, "validate_input_kwargs", return_value=_validated_kwargs() + ), + patch("litellm.aresponses", _fake_aresponses), + ): + try: + await handler.acompletion() + except Exception: + pass + assert _fake_aresponses.kwargs.get("custom_llm_provider") == "openai", ( + "async bridge must forward custom_llm_provider to litellm.aresponses() " + "so the downstream get_llm_provider() call does not re-strip the " + "provider prefix on a provider/provider/model deployment string" + ) diff --git a/tests/test_litellm/integrations/open_telemetry/test_otel_exception_handler.py b/tests/test_litellm/integrations/open_telemetry/test_otel_exception_handler.py index 56059c260c3..348ef5082e7 100644 --- a/tests/test_litellm/integrations/open_telemetry/test_otel_exception_handler.py +++ b/tests/test_litellm/integrations/open_telemetry/test_otel_exception_handler.py @@ -18,7 +18,9 @@ from litellm.proxy.proxy_server import ( otel_unhandled_exception_handler, ) -from ._helpers import assert_server_span_attrs +from litellm.integrations._types.open_inference import ErrorAttributes + +from ._helpers import assert_server_span_attrs, get_server_span def _fake_request(parent_otel_span=None): @@ -92,6 +94,33 @@ def test_exception_handler_closes_span( ) +@pytest.mark.parametrize("path", ["/team/list", "/organization/list"]) +def test_openai_exception_handler_stamps_structured_error_on_span( + wired_otel, server_span_factory, path +): + """A ProxyException 401 (invalid/expired key on a management endpoint) must + leave error.type, error.code AND error.message on the SERVER span. Pre-fix, + ProxyException stringified to "" so error.message was dropped — the span + showed an error with no message.""" + msg = "Authentication Error, Invalid proxy server token passed." + request = _fake_request(parent_otel_span=server_span_factory(path)) + exc = ProxyException(message=msg, type="auth_error", param="key", code=401) + + response = asyncio.run(openai_exception_handler(request, exc)) + assert response.status_code == 401 + + assert_server_span_attrs( + wired_otel, + expected_status=401, + expected_url_path=path, + where=f"openai_exception_handler ({path})", + ) + attrs = get_server_span(wired_otel).attributes + assert attrs.get(ErrorAttributes.ERROR_MESSAGE) == msg + assert attrs.get(ErrorAttributes.ERROR_TYPE) == "ProxyException" + assert attrs.get(ErrorAttributes.ERROR_CODE) == "401" + + def test_unhandled_handler_reraises_known_exceptions(wired_otel, server_span_factory): """ProxyException / HTTPException / RequestValidationError have dedicated handlers.""" request = _fake_request(parent_otel_span=server_span_factory("/key/generate")) diff --git a/tests/test_litellm/integrations/opik/test_opik_extractors.py b/tests/test_litellm/integrations/opik/test_opik_extractors.py new file mode 100644 index 00000000000..6f85a1c6090 --- /dev/null +++ b/tests/test_litellm/integrations/opik/test_opik_extractors.py @@ -0,0 +1,84 @@ +from litellm.integrations.opik.opik_payload_builder.extractors import ( + extract_opik_metadata, +) + + +def test_extract_opik_metadata_fills_missing_keys_from_auth_metadata(): + litellm_metadata = {"opik": {"project_name": "my-proj"}} + standard_logging_metadata = { + "user_api_key_auth_metadata": { + "opik": { + "workspace": "auth-workspace", + "project_name": "auth-project", + } + } + } + + result = extract_opik_metadata( + litellm_metadata=litellm_metadata, + standard_logging_metadata=standard_logging_metadata, + ) + + assert result == { + "project_name": "my-proj", + "workspace": "auth-workspace", + } + + +def test_extract_opik_metadata_request_metadata_overrides_auth_metadata(): + litellm_metadata = { + "opik": { + "workspace": "request-workspace", + "thread_id": "request-thread", + } + } + standard_logging_metadata = { + "user_api_key_auth_metadata": { + "opik": { + "workspace": "auth-workspace", + "thread_id": "auth-thread", + "project_name": "auth-project", + } + } + } + + result = extract_opik_metadata( + litellm_metadata=litellm_metadata, + standard_logging_metadata=standard_logging_metadata, + ) + + assert result == { + "workspace": "request-workspace", + "thread_id": "request-thread", + "project_name": "auth-project", + } + + +def test_extract_opik_metadata_requester_metadata_overrides_all_other_sources(): + litellm_metadata = {"opik": {"project_name": "request-project"}} + standard_logging_metadata = { + "user_api_key_auth_metadata": { + "opik": { + "workspace": "auth-workspace", + "project_name": "auth-project", + } + }, + "requester_metadata": { + "opik": { + "workspace": "requester-workspace", + "thread_id": "requester-thread", + "project_name": "requester-project", + } + }, + } + + result = extract_opik_metadata( + litellm_metadata=litellm_metadata, + standard_logging_metadata=standard_logging_metadata, + ) + + assert result == { + "project_name": "requester-project", + "workspace": "requester-workspace", + "thread_id": "requester-thread", + } diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_logger.py b/tests/test_litellm/integrations/otel/test_otel_v2_logger.py index 7fb0e10a247..8dffb71bbf0 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_logger.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_logger.py @@ -238,6 +238,155 @@ def test_idempotent_on_repeat_callback(): assert len(exporter.get_finished_spans()) == 1 +# --------------------------------------------------------------------------- # +# MCP tool-call spans +# --------------------------------------------------------------------------- # + + +def _mcp_payload(**overrides): + payload = { + "call_type": "call_mcp_tool", + "status": "success", + "litellm_call_id": "mcp_1", + "response_cost": 0.01, + "metadata": { + "user_api_key_team_id": "t1", + "mcp_tool_call_metadata": { + "name": "get_weather", + "arguments": {"city": "Paris"}, + "result": {"temp_c": 21}, + "mcp_server_name": "weather-mcp", + "mcp_session_id": "sess-abc123", + }, + }, + "hidden_params": {}, + } + payload.update(overrides) + return payload + + +def _logger_capturing(): + from litellm.integrations.otel.model.config import CaptureMessageContent + + cfg = OpenTelemetryV2Config( + exporter="in_memory", + legacy_compat=False, + capture_message_content=CaptureMessageContent.SPAN_ONLY, + ) + exporter = InMemorySpanExporter() + tracer_provider = providers.build_tracer_provider(cfg, exporter=exporter) + return OpenTelemetryV2(config=cfg, tracer_provider=tracer_provider), exporter + + +def test_mcp_tool_call_emits_client_span(): + """A closed MCP tool call becomes a CLIENT span named ``tools/call {tool}``, + carrying the MCP semconv method/operation and the vendor server name.""" + logger, exporter = _logger() + kwargs = {"standard_logging_object": _mcp_payload()} + asyncio.run(logger.async_log_success_event(kwargs, None, None, None)) + (span,) = exporter.get_finished_spans() + assert span.name == "tools/call get_weather" + assert span.kind is SpanKind.CLIENT + assert span.attributes["mcp.method.name"] == "tools/call" + assert span.attributes["mcp.session.id"] == "sess-abc123" + assert span.attributes[GenAI.OPERATION_NAME] == "execute_tool" + assert span.attributes["gen_ai.tool.name"] == "get_weather" + assert span.attributes[LiteLLM.MCP_SERVER_NAME] == "weather-mcp" + assert span.attributes[LiteLLM.CALL_ID] == "mcp_1" + assert span.status.status_code is StatusCode.UNSET + # Tool I/O is content: withheld while capture is off (the default). + assert "gen_ai.tool.call.arguments" not in span.attributes + assert "gen_ai.tool.call.result" not in span.attributes + + +def test_mcp_tool_call_stateless_omits_session_id(): + """A stateless MCP call carries no ``mcp-session-id``, so the span must omit + ``mcp.session.id`` rather than stamping an empty or ``None`` value.""" + logger, exporter = _logger() + payload = _mcp_payload() + del payload["metadata"]["mcp_tool_call_metadata"]["mcp_session_id"] + asyncio.run( + logger.async_log_success_event( + {"standard_logging_object": payload}, None, None, None + ) + ) + (span,) = exporter.get_finished_spans() + assert "mcp.session.id" not in span.attributes + assert span.attributes["mcp.method.name"] == "tools/call" + + +def test_mcp_tool_call_is_not_logged_as_llm_call(): + """The MCP branch must short-circuit the LLM-call path: even if ``pre_call`` + opened a stray carrier for this id, the result is one MCP span, never an LLM + ``chat`` span.""" + logger, exporter = _logger() + kwargs = {"standard_logging_object": _mcp_payload()} + logger.log_pre_api_call(model="MCP: get_weather", messages=[], kwargs=kwargs) + asyncio.run(logger.async_log_success_event(kwargs, None, None, None)) + (span,) = exporter.get_finished_spans() + assert span.attributes["mcp.method.name"] == "tools/call" + assert "gen_ai.request.model" not in span.attributes + + +def test_mcp_tool_call_captures_io_when_enabled(): + logger, exporter = _logger_capturing() + kwargs = {"standard_logging_object": _mcp_payload()} + asyncio.run(logger.async_log_success_event(kwargs, None, None, None)) + (span,) = exporter.get_finished_spans() + assert '"Paris"' in span.attributes["gen_ai.tool.call.arguments"] + assert "21" in span.attributes["gen_ai.tool.call.result"] + + +def test_mcp_tool_call_failure_marks_error(): + logger, exporter = _logger() + payload = _mcp_payload( + status="failure", + error_information={"error_class": "MCPError", "error_message": "upstream 500"}, + ) + asyncio.run( + logger.async_log_failure_event( + {"standard_logging_object": payload}, None, None, None + ) + ) + (span,) = exporter.get_finished_spans() + assert span.name == "tools/call get_weather" + assert span.status.status_code is StatusCode.ERROR + assert span.attributes["error.type"] == "MCPError" + + +def test_mcp_tool_call_deduped_on_repeat(): + logger, exporter = _logger() + kwargs = {"standard_logging_object": _mcp_payload()} + asyncio.run(logger.async_log_success_event(kwargs, None, None, None)) + asyncio.run(logger.async_log_success_event(kwargs, None, None, None)) + assert len(exporter.get_finished_spans()) == 1 + + +def test_mcp_tool_call_metadata_read_from_nested_metadata_not_top_level(): + """``mcp_tool_call_metadata`` lives under ``StandardLoggingPayload.metadata``; + a top-level copy (the pre-fix shape the reader used to look at) must be ignored + so the reader can't silently regress to producing an empty ``tools/call`` span + with no session id, tool name, or server name.""" + logger, exporter = _logger() + payload = _mcp_payload() + # Move the real metadata to the top level only, mirroring the old buggy read + # location. ``call_type`` still classifies this as an MCP call, so the span is + # emitted, but none of its fields are reachable from the wrong nesting level. + payload["mcp_tool_call_metadata"] = payload["metadata"].pop( + "mcp_tool_call_metadata" + ) + asyncio.run( + logger.async_log_success_event( + {"standard_logging_object": payload}, None, None, None + ) + ) + (span,) = exporter.get_finished_spans() + assert span.name == "tools/call" + assert "mcp.session.id" not in span.attributes + assert "gen_ai.tool.name" not in span.attributes + assert LiteLLM.MCP_SERVER_NAME not in span.attributes + + def test_pre_call_idempotent_keeps_first_span(): """A retried call may re-enter ``pre_call`` with the same call id; the first span (with the true start time) is kept, not replaced.""" @@ -419,17 +568,10 @@ def test_guardrail_span_anchors_to_root_inside_active_phase_span(): SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME ) set_request_root_span(server) - request_data = { - "metadata": { - "standard_logging_guardrail_information": { - "guardrail_name": "my_guard", - "guardrail_status": "success", - } - } - } + entry = {"guardrail_name": "my_guard", "guardrail_status": "success"} with trace.use_span(server, end_on_exit=False): with logger.start_phase_span("auth /chat/completions"): - logger._emit_guardrail_spans(request_data) + logger.emit_guardrail_span(entry) server.end() by_name = {s.name: s for s in exporter.get_finished_spans()} guard = by_name["execute_guardrail my_guard"] @@ -1066,37 +1208,29 @@ def test_boundary_span_closes_without_proxy_fanout(monkeypatch): # --------------------------------------------------------------------------- # -def _guardrail_request_data(*, start, end): +def _guardrail_entry(*, start, end): return { - "metadata": { - "standard_logging_guardrail_information": [ - { - "guardrail_name": "openai-moderation", - "guardrail_mode": "pre_call", - "guardrail_status": "success", - "start_time": start, - "end_time": end, - "duration": end - start, - } - ], - } + "guardrail_name": "openai-moderation", + "guardrail_mode": "pre_call", + "guardrail_status": "success", + "start_time": start, + "end_time": end, + "duration": end - start, } def test_guardrail_span_parents_to_ambient_server_span(): - """The post-call hook runs in the request task with the server span ambient, - so the guardrail span parents to it natively — no span threaded through - metadata. (Auth already finished, so no phase span is active.)""" + """``emit_guardrail_span`` runs in the request task with the server span + ambient, so with no explicit anchor set the guardrail span parents to it. + (Auth already finished, so no phase span is active.)""" logger, exporter = _logger() server = logger._emitter.start_span( SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME ) - data = _guardrail_request_data(start=1000.0, end=1000.5) + entry = _guardrail_entry(start=1000.0, end=1000.5) try: with trace.use_span(server, end_on_exit=False): - asyncio.run( - logger.async_post_call_success_hook(data, _Auth(), {"ok": True}) - ) + logger.emit_guardrail_span(entry) finally: server.end() g = {s.name: s for s in exporter.get_finished_spans()}[ @@ -1107,17 +1241,15 @@ def test_guardrail_span_parents_to_ambient_server_span(): def test_guardrail_span_uses_actual_execution_timestamps(): """A pre_call guardrail's span carries its real start/end (from the logging - entry), so it sorts before the LLM call instead of at post-call emit time.""" + entry), so it sorts before the LLM call instead of at emission time.""" logger, exporter = _logger() server = logger._emitter.start_span( SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME ) - data = _guardrail_request_data(start=1700.0, end=1700.25) + entry = _guardrail_entry(start=1700.0, end=1700.25) try: with trace.use_span(server, end_on_exit=False): - asyncio.run( - logger.async_post_call_success_hook(data, _Auth(), {"ok": True}) - ) + logger.emit_guardrail_span(entry) finally: server.end() g = {s.name: s for s in exporter.get_finished_spans()}[ @@ -1125,3 +1257,60 @@ def test_guardrail_span_uses_actual_execution_timestamps(): ] assert g.start_time == to_ns(1700.0) assert g.end_time == to_ns(1700.25) + + +def test_emit_guardrail_span_anchors_to_root_not_ambient_phase_span(): + """With an explicit request-root anchor set, the guardrail span parents to it + even while a phase span is the active OTel context — the anchor wins over + ambient, so a guardrail emitted mid-``auth`` is a sibling of the LLM call, not + a child of ``auth``.""" + logger, exporter = _logger() + server = logger._emitter.start_span( + SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME + ) + set_request_root_span(server) + entry = _guardrail_entry(start=2000.0, end=2000.1) + with logger.start_phase_span("auth /chat/completions"): + logger.emit_guardrail_span(entry) + server.end() + by_name = {s.name: s for s in exporter.get_finished_spans()} + guard = by_name["execute_guardrail openai-moderation"] + auth_span = by_name["auth /chat/completions"] + assert guard.parent.span_id == server.get_span_context().span_id + assert guard.parent.span_id != auth_span.get_span_context().span_id + + +def test_module_level_emit_guardrail_span_routes_to_registered_logger(monkeypatch): + """The module-level entry point custom_guardrail calls routes the entry to the + single registered v2 logger and emits exactly one span.""" + import litellm.integrations.otel.logger as otel_logger + + logger, exporter = _logger() + monkeypatch.setattr(otel_logger, "_registered_v2_logger", lambda: logger) + + otel_logger.emit_guardrail_span(_guardrail_entry(start=3000.0, end=3000.2)) + + names = [s.name for s in exporter.get_finished_spans()] + assert names.count("execute_guardrail openai-moderation") == 1 + + +def test_module_level_emit_guardrail_span_noop_without_registered_logger(monkeypatch): + """No registered v2 logger (SDK path / OTel not configured) → emitting is a + no-op rather than an error.""" + import litellm.integrations.otel.logger as otel_logger + + monkeypatch.setattr(otel_logger, "_registered_v2_logger", lambda: None) + otel_logger.emit_guardrail_span(_guardrail_entry(start=1.0, end=2.0)) + + +def test_module_level_emit_guardrail_span_swallows_emit_errors(monkeypatch): + """Span emission is best-effort: a logger that raises must never propagate out + of the guardrail-recording path and break guardrail evaluation.""" + import litellm.integrations.otel.logger as otel_logger + + class _Boom: + def emit_guardrail_span(self, entry): + raise RuntimeError("emit blew up") + + monkeypatch.setattr(otel_logger, "_registered_v2_logger", lambda: _Boom()) + otel_logger.emit_guardrail_span(_guardrail_entry(start=1.0, end=2.0)) diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py b/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py index c4e80145c70..20824ca09e6 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py @@ -1,8 +1,6 @@ """Tests for the OTel v2 sources of truth: span registry, semconv keys, config, and the typed StandardLoggingPayload adapter. These need no OTel SDK.""" -import pytest - from litellm.integrations.otel import ( BAGGAGE_PROMOTED_KEYS, DB, @@ -92,11 +90,14 @@ def test_registry_hierarchy_shape(): # guardrail runs before the LLM call exists, so it's a sibling of it. assert set(child_roles(SpanRole.PROXY_REQUEST)) == { SpanRole.LLM_CALL, + SpanRole.MCP_TOOL_CALL, SpanRole.GUARDRAIL, SpanRole.DB_CALL, SpanRole.SERVICE, } assert SPAN_REGISTRY[SpanRole.LLM_CALL].kind is LiteLLMSpanKind.CLIENT + # The proxy is an MCP client to the upstream tool server: CLIENT span. + assert SPAN_REGISTRY[SpanRole.MCP_TOOL_CALL].kind is LiteLLMSpanKind.CLIENT assert SPAN_REGISTRY[SpanRole.PROXY_REQUEST].kind is LiteLLMSpanKind.SERVER assert SPAN_REGISTRY[SpanRole.GUARDRAIL].parent is SpanRole.PROXY_REQUEST # An outbound datastore call is a CLIENT span; an internal service is INTERNAL. @@ -121,14 +122,52 @@ def _all_constants(cls): def test_attribute_keys_are_unique_across_namespaces(): + from litellm.integrations.otel import MCP, Client, JsonRpc, Network + # prefixes are allowed to be substrings; exact keys must not collide. exact = set() - for cls in (GenAI, Error, Server, HTTP, DB): + for cls in (GenAI, Error, Server, HTTP, DB, MCP, JsonRpc, Network, Client): for key in _all_constants(cls): assert key not in exact, f"duplicate attribute key {key}" exact.add(key) +def test_mcp_attribute_vocabulary_is_complete(): + """Every span-attribute key the OTel GenAI MCP semconv defines has a constant. + + Pins the vocabulary so a dropped or renamed key fails here rather than + silently emitting a non-conformant attribute name. + """ + from litellm.integrations.otel import MCP, Client, JsonRpc, Network + + defined = set() + for cls in (GenAI, Error, Server, MCP, JsonRpc, Network, Client): + defined |= _all_constants(cls) + required = { + "mcp.method.name", + "mcp.session.id", + "mcp.protocol.version", + "mcp.resource.uri", + "jsonrpc.request.id", + "jsonrpc.protocol.version", + "rpc.response.status_code", + "gen_ai.operation.name", + "gen_ai.tool.name", + "gen_ai.tool.call.arguments", + "gen_ai.tool.call.result", + "gen_ai.prompt.name", + "error.type", + "server.address", + "server.port", + "client.address", + "client.port", + "network.protocol.name", + "network.protocol.version", + "network.transport", + } + assert required <= defined, f"missing MCP semconv keys: {required - defined}" + + def test_provider_resolution(): assert resolve_provider("openai") == "openai" assert resolve_provider("bedrock") == "aws.bedrock" @@ -143,6 +182,103 @@ def test_operation_resolution(): assert resolve_operation("aembedding") is GenAIOperation.EMBEDDINGS assert resolve_operation("atext_completion") is GenAIOperation.TEXT_COMPLETION assert resolve_operation(None) is GenAIOperation.CHAT + # An MCP tool call is an ``execute_tool`` operation, not a chat completion. + assert resolve_operation("call_mcp_tool") is GenAIOperation.EXECUTE_TOOL + + +# --- MCP tool-call (source of truth #1/#2/#3) ------------------------------- # + + +def _mcp_payload(capture=False, **overrides): + payload = { + "call_type": "call_mcp_tool", + "status": "success", + "litellm_call_id": "mcp_call_1", + "response_cost": 0.01, + "metadata": { + "user_api_key_team_id": "t1", + "mcp_tool_call_metadata": { + "name": "get_weather", + "arguments": {"city": "Paris"}, + "result": {"temp_c": 21}, + "mcp_server_name": "weather-mcp", + "mcp_session_id": "sess-abc123", + }, + }, + "hidden_params": {}, + } + payload.update(overrides) + return payload + + +def test_mcp_method_values_match_wire_format(): + from litellm.integrations.otel import MCP, MCPMethod + + assert MCPMethod.TOOLS_CALL.value == "tools/call" + assert MCPMethod.TOOLS_LIST.value == "tools/list" + assert MCP.METHOD_NAME == "mcp.method.name" + + +def test_mcp_tool_call_adapter_extracts_fields(): + from litellm.integrations.otel import MCPToolCallSpanData + + data = MCPToolCallSpanData.from_standard_logging_payload(_mcp_payload()) + assert data.operation is GenAIOperation.EXECUTE_TOOL + assert data.method == "tools/call" + assert data.tool_name == "get_weather" + assert data.server_name == "weather-mcp" + assert data.session_id == "sess-abc123" + assert data.response_cost == 0.01 + assert data.identity.call_id == "mcp_call_1" + assert data.identity.team_id == "t1" + assert data.error is None + + +def test_mcp_tool_call_content_gated_off_by_default(): + # Arguments and result are sensitive tool I/O: withheld unless content capture + # is explicitly enabled, exactly like prompt/response bodies. + from litellm.integrations.otel import MCPToolCallSpanData + + off = MCPToolCallSpanData.from_standard_logging_payload(_mcp_payload()) + assert off.arguments_json is None and off.result_json is None + + on = MCPToolCallSpanData.from_standard_logging_payload( + _mcp_payload(), capture_content=True + ) + assert on.arguments_json is not None and '"Paris"' in on.arguments_json + assert on.result_json is not None and "21" in on.result_json + + +def test_mcp_tool_call_failure_path(): + from litellm.integrations.otel import MCPToolCallSpanData + + data = MCPToolCallSpanData.from_standard_logging_payload( + _mcp_payload( + status="failure", + error_information={"error_class": "MCPError", "error_message": "boom"}, + ) + ) + assert data.error is not None + assert data.error.error_type == "MCPError" + assert data.error.message == "boom" + + +def test_is_mcp_tool_call_detection(): + from litellm.integrations.otel import is_mcp_tool_call + + assert is_mcp_tool_call(_mcp_payload()) is True + # call_type alone is enough even before the gateway stamps its metadata. + assert is_mcp_tool_call({"call_type": "call_mcp_tool"}) is True + assert is_mcp_tool_call({"call_type": "acompletion"}) is False + assert is_mcp_tool_call({}) is False + + +def test_mcp_tool_call_span_name(): + from litellm.integrations.otel import MCPToolCallSpanData + from litellm.integrations.otel.model.spans import mcp_tool_call_span_name + + data = MCPToolCallSpanData.from_standard_logging_payload(_mcp_payload()) + assert mcp_tool_call_span_name(data) == "tools/call get_weather" # --- typed adapter (source of truth #3) ------------------------------------- # diff --git a/tests/test_litellm/integrations/test_custom_guardrail.py b/tests/test_litellm/integrations/test_custom_guardrail.py index a881044dc18..4956fccb8db 100644 --- a/tests/test_litellm/integrations/test_custom_guardrail.py +++ b/tests/test_litellm/integrations/test_custom_guardrail.py @@ -500,6 +500,65 @@ class TestGuardrailLoggingAggregation: assert info[1]["guardrail_name"] == "test_guardrail" +class TestGuardrailOtelSpanEmission: + """Recording a guardrail emits its otel span inline, so every guardrail + execution produces a span — including the pass-through allow path that never + reaches a post-call hook.""" + + def _make_guardrail(self): + from litellm.types.guardrails import GuardrailEventHooks + + return CustomGuardrail( + guardrail_name="emit_guard", + event_hook=GuardrailEventHooks.pre_call, + ) + + def _record(self, guardrail, request_data): + guardrail.add_standard_logging_guardrail_information_to_request_data( + guardrail_json_response={"result": "ok"}, + request_data=request_data, + guardrail_status="success", + start_time=1.0, + end_time=2.0, + duration=1.0, + ) + + def test_emits_span_for_recorded_entry(self, monkeypatch): + captured = [] + monkeypatch.setattr( + "litellm.integrations.otel.logger.emit_guardrail_span", + captured.append, + ) + + request_data = {"metadata": {}} + self._record(self._make_guardrail(), request_data) + + assert len(captured) == 1 + emitted = captured[0] + recorded = request_data["metadata"]["standard_logging_guardrail_information"][ + -1 + ] + assert emitted is recorded + assert emitted["guardrail_name"] == "emit_guard" + assert emitted["start_time"] == 1.0 + assert emitted["end_time"] == 2.0 + + def test_span_emission_failure_does_not_break_recording(self, monkeypatch): + def _boom(_entry): + raise RuntimeError("otel exporter down") + + monkeypatch.setattr( + "litellm.integrations.otel.logger.emit_guardrail_span", _boom + ) + + request_data = {"metadata": {}} + self._record(self._make_guardrail(), request_data) + + info = request_data["metadata"]["standard_logging_guardrail_information"] + assert len(info) == 1 + assert info[0]["guardrail_name"] == "emit_guard" + + class TestGuardrailSensitiveFieldStripping: """Tests that secret_fields is stripped from guardrail responses before logging. diff --git a/tests/test_litellm/integrations/test_opentelemetry.py b/tests/test_litellm/integrations/test_opentelemetry.py index c4500bd6135..0601f9c0eef 100644 --- a/tests/test_litellm/integrations/test_opentelemetry.py +++ b/tests/test_litellm/integrations/test_opentelemetry.py @@ -1263,7 +1263,6 @@ class TestOpenTelemetry(unittest.TestCase): ) as mock_get_headers, patch.object(otel, "_get_tracer_with_dynamic_headers") as mock_get_tracer, ): - # Test case 1: With dynamic headers mock_get_headers.return_value = { "arize-space-id": "test-space", @@ -1761,6 +1760,31 @@ class TestOpenTelemetry(unittest.TestCase): mock_tracer.start_span.assert_not_called() +class TestOpenTelemetryToNs(unittest.TestCase): + """``_to_ns`` converts a span boundary to epoch nanoseconds. Service spans now + feed it real float/datetime windows, and a missing boundary arrives as + ``None`` — all three shapes must convert without raising the ``AttributeError`` + a bare ``dt.timestamp()`` would on a float or ``None``.""" + + def setUp(self): + self.otel = OpenTelemetry() + + def test_datetime_converts_to_epoch_ns(self): + dt = datetime(2026, 5, 26, 12, 0, 0, tzinfo=timezone.utc) + self.assertEqual(self.otel._to_ns(dt), int(dt.timestamp() * 1e9)) + + def test_float_epoch_seconds_scaled_to_ns(self): + self.assertEqual(self.otel._to_ns(1700.5), 1_700_500_000_000) + + def test_int_epoch_seconds_scaled_to_ns(self): + self.assertEqual(self.otel._to_ns(1700), 1_700_000_000_000) + + @patch("litellm.integrations.opentelemetry.datetime") + def test_none_falls_back_to_current_time(self, mock_datetime): + mock_datetime.now.return_value.timestamp.return_value = 1700.0 + self.assertEqual(self.otel._to_ns(None), 1_700_000_000_000) + + class TestOpenTelemetryHeaderSplitting(unittest.TestCase): """Test suite for _get_headers_dictionary method""" @@ -2643,7 +2667,7 @@ class TestOpenTelemetryExternalSpan(unittest.TestCase): # Verify parent span is still recording after each call self.assertTrue( parent_span.is_recording(), - f"External span should still be recording after completion #{i+1}", + f"External span should still be recording after completion #{i + 1}", ) # Verify all spans have the same trace_id @@ -5145,6 +5169,138 @@ class TestOpenTelemetryPreprocessingDuration(unittest.TestCase): assert "litellm.preprocessing.duration_ms" not in self._attr(span, exp) +class TestGetSpanContextLitellmMetadataFallback(unittest.TestCase): + """ + Tests for _get_span_context() falling back to litellm_metadata. + + On /v1/messages (Anthropic Messages API) and other LITELLM_METADATA_ROUTES, + litellm_parent_otel_span is stored in litellm_params["litellm_metadata"] + instead of litellm_params["metadata"]. _get_span_context() must check + both locations. + + Fixes: https://github.com/BerriAI/litellm/issues/27934 + """ + + def test_span_context_from_metadata(self): + """Parent span is found when stored in litellm_params['metadata'] (OpenAI path).""" + otel = OpenTelemetry() + mock_span = MagicMock() + mock_span.get_span_context.return_value = MagicMock(is_valid=True) + + kwargs = { + "litellm_params": { + "metadata": {"litellm_parent_otel_span": mock_span}, + } + } + + ctx, detected_span = otel._get_span_context(kwargs) + self.assertIsNotNone(ctx) + # Should NOT fall through to "no parent context" path + self.assertIsNone(detected_span) + + def test_span_context_from_litellm_metadata_fallback(self): + """Parent span is found when stored in litellm_params['litellm_metadata'] (Anthropic path).""" + otel = OpenTelemetry() + mock_span = MagicMock() + mock_span.get_span_context.return_value = MagicMock(is_valid=True) + + kwargs = { + "litellm_params": { + "metadata": { + "user_id": "test-user" + }, # Anthropic native metadata, no span + "litellm_metadata": {"litellm_parent_otel_span": mock_span}, + } + } + + ctx, detected_span = otel._get_span_context(kwargs) + self.assertIsNotNone(ctx) + self.assertIsNone(detected_span) + + def test_span_context_metadata_takes_priority(self): + """When both metadata and litellm_metadata have the span, metadata wins.""" + otel = OpenTelemetry() + span_from_metadata = MagicMock(name="span_from_metadata") + span_from_metadata.get_span_context.return_value = MagicMock(is_valid=True) + span_from_litellm_metadata = MagicMock(name="span_from_litellm_metadata") + span_from_litellm_metadata.get_span_context.return_value = MagicMock( + is_valid=True + ) + + kwargs = { + "litellm_params": { + "metadata": {"litellm_parent_otel_span": span_from_metadata}, + "litellm_metadata": { + "litellm_parent_otel_span": span_from_litellm_metadata + }, + } + } + + ctx, detected_span = otel._get_span_context(kwargs) + self.assertIsNotNone(ctx) + self.assertIsNone(detected_span) + # metadata span is found first, so get_span_context on the + # litellm_metadata span should never be called — proving + # metadata takes priority over litellm_metadata. + span_from_litellm_metadata.get_span_context.assert_not_called() + + def test_span_context_no_parent_when_neither_has_span(self): + """When neither metadata nor litellm_metadata has a span, returns (None, None).""" + otel = OpenTelemetry() + + kwargs = { + "litellm_params": { + "metadata": {"user_id": "test-user"}, + "litellm_metadata": {"some_key": "some_value"}, + } + } + + ctx, detected_span = otel._get_span_context(kwargs) + # No parent span in either metadata dict and no active span in test + # context, so both should be None. + self.assertIsNone(ctx) + self.assertIsNone(detected_span) + + +class TestEndProxySpanLitellmMetadataFallback(unittest.TestCase): + """ + Tests for _end_proxy_span_from_kwargs() falling back to litellm_metadata. + + Fixes: https://github.com/BerriAI/litellm/issues/27934 + """ + + def test_end_proxy_span_from_metadata(self): + """Proxy span is found and ended from litellm_params['metadata'].""" + otel = OpenTelemetry() + mock_span = MagicMock() + mock_span.name = "Received Proxy Server Request" + mock_span.is_recording.return_value = True + + kwargs = { + "litellm_params": { + "metadata": {"litellm_parent_otel_span": mock_span}, + } + } + + otel._end_proxy_span_from_kwargs(kwargs, end_time=datetime.now()) + mock_span.end.assert_called_once() + + def test_end_proxy_span_from_litellm_metadata(self): + """Proxy span is found and ended from litellm_params['litellm_metadata'] (fallback).""" + otel = OpenTelemetry() + mock_span = MagicMock() + mock_span.name = "Received Proxy Server Request" + mock_span.is_recording.return_value = True + + kwargs = { + "litellm_params": { + "metadata": {"user_id": "test-user"}, # No span here + "litellm_metadata": {"litellm_parent_otel_span": mock_span}, + } + } + + otel._end_proxy_span_from_kwargs(kwargs, end_time=datetime.now()) + mock_span.end.assert_called_once() class TestOpenTelemetryInferenceIdentityAttributes(unittest.TestCase): """team_metadata, http.route, and both model names (the user-facing model_group alias and the dispatched provider model) must land on the diff --git a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index 129ea237efe..91fd07dcffc 100644 --- a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -22,6 +22,18 @@ from litellm.litellm_core_utils.prompt_templates.factory import ( from litellm.types.llms.openai import ChatCompletionToolMessage +def _get_gemini_function_response_inline_data_parts(result): + assert isinstance(result, list), "expected Gemini parts list" + assert len(result) == 1, "multimodal function responses should stay in one part" + function_response_part = result[0] + assert ( + "inline_data" not in function_response_part + ), "inline_data should be nested under function_response.parts" + function_response = function_response_part["function_response"] + nested_parts = function_response["parts"] + return [part["inline_data"] for part in nested_parts if "inline_data" in part] + + def test_ollama_pt_simple_messages(): """Test basic functionality with simple text messages""" messages = [ @@ -615,8 +627,8 @@ def test_convert_gemini_tool_call_result_with_image_url(): message=message_str_format, last_message_with_tool_calls=last_message_with_tool_calls, ) - # Should have inline_data for the image - assert isinstance(result, list) and any("inline_data" in p for p in result) + inline_parts = _get_gemini_function_response_inline_data_parts(result) + assert len(inline_parts) == 1 # Test with dict image_url format (OpenAI standard) message_dict_format = ChatCompletionToolMessage( @@ -635,7 +647,8 @@ def test_convert_gemini_tool_call_result_with_image_url(): message=message_dict_format, last_message_with_tool_calls=last_message_with_tool_calls, ) - assert isinstance(result2, list) and any("inline_data" in p for p in result2) + inline_parts = _get_gemini_function_response_inline_data_parts(result2) + assert len(inline_parts) == 1 def test_convert_gemini_tool_call_result_with_anthropic_image_block(): @@ -677,11 +690,10 @@ def test_convert_gemini_tool_call_result_with_anthropic_image_block(): message=message, last_message_with_tool_calls=last_message_with_tool_calls, ) - assert isinstance(result, list), "expected a list of parts" - inline_parts = [p for p in result if "inline_data" in p] + inline_parts = _get_gemini_function_response_inline_data_parts(result) assert len(inline_parts) == 1, "expected exactly one inline_data part" - assert inline_parts[0]["inline_data"]["mime_type"] == "image/png" - assert inline_parts[0]["inline_data"]["data"] == tiny_png_b64 + assert inline_parts[0]["mime_type"] == "image/png" + assert inline_parts[0]["data"] == tiny_png_b64 def test_convert_gemini_tool_call_result_with_multiple_anthropic_image_blocks(): @@ -734,12 +746,11 @@ def test_convert_gemini_tool_call_result_with_multiple_anthropic_image_blocks(): message=message, last_message_with_tool_calls=last_message_with_tool_calls, ) - assert isinstance(result, list), "expected a list of parts" - inline_parts = [p for p in result if "inline_data" in p] + inline_parts = _get_gemini_function_response_inline_data_parts(result) assert ( len(inline_parts) == 2 ), f"expected 2 inline_data parts, got {len(inline_parts)}" - mime_types = {p["inline_data"]["mime_type"] for p in inline_parts} + mime_types = {p["mime_type"] for p in inline_parts} assert mime_types == {"image/png", "image/jpeg"} @@ -773,13 +784,12 @@ def test_convert_gemini_tool_call_result_with_data_url_string(): message=message, last_message_with_tool_calls=last_message_with_tool_calls, ) - assert isinstance(result, list), "expected a list of parts" - inline_parts = [p for p in result if "inline_data" in p] + inline_parts = _get_gemini_function_response_inline_data_parts(result) assert ( len(inline_parts) == 1 ), "data-URL image string was not converted to inline_data" - assert inline_parts[0]["inline_data"]["mime_type"] == "image/png" - assert inline_parts[0]["inline_data"]["data"] == tiny_png_b64 + assert inline_parts[0]["mime_type"] == "image/png" + assert inline_parts[0]["data"] == tiny_png_b64 def test_convert_gemini_tool_call_result_with_data_url_extra_params(): @@ -811,12 +821,11 @@ def test_convert_gemini_tool_call_result_with_data_url_extra_params(): message=message, last_message_with_tool_calls=last_message_with_tool_calls, ) - assert isinstance(result, list), "expected a list of parts" - inline_parts = [p for p in result if "inline_data" in p] + inline_parts = _get_gemini_function_response_inline_data_parts(result) assert len(inline_parts) == 1 assert ( - inline_parts[0]["inline_data"]["mime_type"] == "image/png" - ), f"expected clean 'image/png', got '{inline_parts[0]['inline_data']['mime_type']}'" + inline_parts[0]["mime_type"] == "image/png" + ), f"expected clean 'image/png', got '{inline_parts[0]['mime_type']}'" def test_bedrock_tools_unpack_defs(): diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index b64cb7c6905..8acda74a156 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -3078,3 +3078,39 @@ class TestFirstApiCallStartTimeSetOnce: assert obj.model_call_details["api_call_start_time"] > first assert obj.model_call_details["first_api_call_start_time"] == first assert user_meta == {} + + +def test_get_error_information_proxy_exception_preserves_message(): + """ProxyException keeps its text in ``.message`` (str() was empty pre-fix), + so error_information must still surface the message and code.""" + from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup + from litellm.proxy._types import ProxyException + + msg = "Authentication Error, Invalid proxy server token passed." + exc = ProxyException(message=msg, type="auth_error", param="key", code=401) + + info = StandardLoggingPayloadSetup.get_error_information(original_exception=exc) + assert info["error_message"] == msg + assert info["error_class"] == "ProxyException" + assert info["error_code"] == "401" + + +def test_get_error_information_prefers_message_attribute_over_empty_str(): + """error_message must come from a populated ``.message`` even when the + exception's __str__ is empty — guards classes that store the text on + ``.message`` without forwarding it to ``Exception.__init__``.""" + from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup + + class _SilentExc(Exception): + def __init__(self): + self.message = "real failure detail" + self.code = 401 + + def __str__(self): + return "" + + info = StandardLoggingPayloadSetup.get_error_information( + original_exception=_SilentExc() + ) + assert info["error_message"] == "real failure detail" + assert info["error_code"] == "401" diff --git a/tests/test_litellm/litellm_core_utils/test_logging_worker.py b/tests/test_litellm/litellm_core_utils/test_logging_worker.py index a44b821db87..978f22ca2e4 100644 --- a/tests/test_litellm/litellm_core_utils/test_logging_worker.py +++ b/tests/test_litellm/litellm_core_utils/test_logging_worker.py @@ -4,6 +4,8 @@ Tests for the LoggingWorker class to ensure graceful shutdown handling. import asyncio import contextvars +import io +import logging from unittest.mock import AsyncMock, patch import pytest @@ -65,6 +67,57 @@ class TestLoggingWorker: # Verify the queue is empty after clearing assert logging_worker._queue.empty() + def test_flush_on_exit_suppresses_closed_handler_errors(self, capsys): + """Atexit flushing should not print logging errors after streams close.""" + worker = LoggingWorker(timeout=1.0, max_queue_size=10) + worker._queue = asyncio.Queue(maxsize=10) + + stream = io.StringIO() + handler = logging.StreamHandler(stream) + logger = logging.getLogger("test_logging_worker_closed_handler") + logger.addHandler(handler) + logger.setLevel(logging.DEBUG) + logger.propagate = False + + async def log_with_closed_handler(): + logger.debug("flush me during shutdown") + + previous_raise_exceptions = logging.raiseExceptions + logging.raiseExceptions = True + + try: + worker.enqueue(log_with_closed_handler()) + stream.close() + + worker._flush_on_exit() + + captured = capsys.readouterr() + assert "I/O operation on closed file" not in captured.err + finally: + logging.raiseExceptions = previous_raise_exceptions + logger.removeHandler(handler) + + def test_flush_on_exit_swallows_errors_and_drains_remaining(self): + """A failing queued coroutine must not abort the atexit drain of later events.""" + worker = LoggingWorker(timeout=1.0, max_queue_size=10) + worker._queue = asyncio.Queue(maxsize=10) + + processed = [] + + async def raises_during_flush(): + raise RuntimeError("boom during shutdown flush") + + async def records_during_flush(): + processed.append("ran") + + worker.enqueue(raises_during_flush()) + worker.enqueue(records_during_flush()) + + worker._flush_on_exit() + + assert processed == ["ran"] + assert worker._queue.empty() + @pytest.mark.asyncio async def test_worker_handles_cancellation_gracefully(self, logging_worker): """Test that the worker handles cancellation without throwing exceptions.""" diff --git a/tests/test_litellm/litellm_core_utils/test_safe_json_dumps.py b/tests/test_litellm/litellm_core_utils/test_safe_json_dumps.py index c71a229cca5..74574370e46 100644 --- a/tests/test_litellm/litellm_core_utils/test_safe_json_dumps.py +++ b/tests/test_litellm/litellm_core_utils/test_safe_json_dumps.py @@ -8,7 +8,7 @@ sys.path.insert( 0, os.path.abspath("../../..") ) # Adds the parent directory to the system path -from litellm.litellm_core_utils.safe_json_dumps import safe_dumps +from litellm.litellm_core_utils.safe_json_dumps import safe_dumps, strip_null_bytes def test_primitive_types(): @@ -140,6 +140,40 @@ def test_non_standard_dict_keys_complex(): raise e +def test_strip_null_bytes_helper(): + assert strip_null_bytes("hello\x00world") == "helloworld" + assert strip_null_bytes("\x00\x00abc\x00") == "abc" + assert strip_null_bytes("no null here") == "no null here" + + +def test_null_byte_stripped_from_string(): + out = safe_dumps("hello\x00world") + assert "\\u0000" not in out + assert json.loads(out) == "helloworld" + + +def test_null_byte_stripped_in_nested_structure(): + data = { + "messages": [{"role": "user", "content": "bad\x00content"}], + "nested": {"k\x00ey": "v\x00alue"}, + } + out = safe_dumps(data) + assert "\\u0000" not in out + result = json.loads(out) + assert result["messages"][0]["content"] == "badcontent" + assert result["nested"] == {"key": "value"} + + +def test_null_byte_stripped_in_fallback_str(): + class WithNullStr: + def __str__(self): + return "obj\x00repr" + + out = safe_dumps({"obj": WithNullStr()}) + assert "\\u0000" not in out + assert json.loads(out)["obj"] == "objrepr" + + def test_pydantic_base_model(): from pydantic import BaseModel diff --git a/tests/test_litellm/llms/azure_ai/test_azure_ai_kimi_k26_metadata.py b/tests/test_litellm/llms/azure_ai/test_azure_ai_kimi_k26_metadata.py new file mode 100644 index 00000000000..812b9288ca8 --- /dev/null +++ b/tests/test_litellm/llms/azure_ai/test_azure_ai_kimi_k26_metadata.py @@ -0,0 +1,76 @@ +""" +Test Azure AI Kimi K2.6 model metadata. +""" + +import json +from importlib.resources import files + +import pytest + + +@pytest.fixture(scope="module") +def use_local_model_cost_map(): + monkeypatch = pytest.MonkeyPatch() + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + + import litellm + from litellm.utils import _invalidate_model_cost_lowercase_map + + original_model_cost = litellm.model_cost + litellm.model_cost = json.loads( + files("litellm") + .joinpath("model_prices_and_context_window_backup.json") + .read_text(encoding="utf-8") + ) + litellm.get_model_info.cache_clear() + _invalidate_model_cost_lowercase_map() + try: + yield litellm + finally: + litellm.model_cost = original_model_cost + litellm.get_model_info.cache_clear() + _invalidate_model_cost_lowercase_map() + monkeypatch.undo() + + +def test_azure_ai_kimi_k26_model_info(use_local_model_cost_map): + model_info = use_local_model_cost_map.get_model_info(model="azure_ai/kimi-k2.6") + + assert model_info["litellm_provider"] == "azure_ai" + assert model_info["mode"] == "chat" + assert model_info["max_input_tokens"] == 262144 + assert model_info["max_output_tokens"] == 262144 + assert model_info["max_tokens"] == 262144 + assert model_info["input_cost_per_token"] == pytest.approx(9.5e-07) + assert model_info["output_cost_per_token"] == pytest.approx(4e-06) + assert model_info["supports_function_calling"] is True + assert model_info["supports_reasoning"] is True + assert model_info["supports_tool_choice"] is True + assert model_info["supports_vision"] is True + + +def test_azure_ai_kimi_k26_raw_model_cost_entry(use_local_model_cost_map): + model_info = use_local_model_cost_map.model_cost["azure_ai/kimi-k2.6"] + + assert model_info["supported_modalities"] == ["text", "image"] + assert model_info["supported_output_modalities"] == ["text"] + assert model_info["supports_function_calling"] is True + assert model_info["supports_reasoning"] is True + assert model_info["supports_tool_choice"] is True + assert model_info["supports_vision"] is True + + +def test_azure_ai_kimi_k26_cost_per_token(use_local_model_cost_map): + from litellm.llms.azure_ai.cost_calculator import cost_per_token + from litellm.types.utils import Usage + + usage = Usage( + prompt_tokens=1_000_000, + completion_tokens=1_000_000, + total_tokens=2_000_000, + ) + + prompt_cost, completion_cost = cost_per_token(model="kimi-k2.6", usage=usage) + + assert prompt_cost == pytest.approx(0.95) + assert completion_cost == pytest.approx(4.0) diff --git a/tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py b/tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py index 0817d92d6b2..474ffee3304 100644 --- a/tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py +++ b/tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py @@ -262,6 +262,61 @@ async def test_handle_async_request_uses_env_proxy(monkeypatch): assert captured["proxy"] == proxy_url +@pytest.mark.asyncio +async def test_handle_async_request_empty_body_sends_no_data(): + """ + A bodyless request (e.g. DELETE /responses/{id}) must reach aiohttp with + data=None. Passing the empty `b""` httpx content makes aiohttp attach a + `Content-Type: application/octet-stream` header, which providers like + OpenAI reject with `unsupported_content_type`. + """ + captured = {} + + class FakeSession: + def __init__(self): + self.closed = False + try: + self._loop = asyncio.get_running_loop() + except RuntimeError: + self._loop = None + + def request(self, *args, **kwargs): + captured["data"] = kwargs.get("data") + + class Resp: + status = 200 + headers = {} + + async def __aenter__(self): + return self + + async def __aexit__(self, exc_type, exc, tb): + pass + + @property + def content(self): + class C: + async def iter_chunked(self, size): + yield b"" + + return C() + + return Resp() + + transport = LiteLLMAiohttpTransport(client=lambda: FakeSession()) # type: ignore + + empty_request = httpx.Request("DELETE", "http://example.com/responses/resp_123") + await transport.handle_async_request(empty_request) + assert captured["data"] is None + + body_request = httpx.Request( + "POST", "http://example.com/responses", json={"input": "ping"} + ) + await transport.handle_async_request(body_request) + assert captured["data"] == body_request.content + assert captured["data"] + + @pytest.mark.asyncio async def test_handle_async_request_uses_env_proxy_per_url(monkeypatch): """Aiohttp transport should honor HTTP(S)_PROXY env vars unless NO_PROXY matches""" diff --git a/tests/test_litellm/llms/gemini/image_edit/test_gemini_image_edit_transformation.py b/tests/test_litellm/llms/gemini/image_edit/test_gemini_image_edit_transformation.py index 682df923693..9b57e1991de 100644 --- a/tests/test_litellm/llms/gemini/image_edit/test_gemini_image_edit_transformation.py +++ b/tests/test_litellm/llms/gemini/image_edit/test_gemini_image_edit_transformation.py @@ -7,6 +7,8 @@ from unittest.mock import MagicMock import httpx import pytest +import litellm +from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup from litellm.llms.gemini.image_edit.transformation import GeminiImageEditConfig @@ -19,6 +21,7 @@ class TestGeminiImageEditTransformation: def test_map_openai_params(self) -> None: optional_params: Dict[str, object] = { + "n": 2, "size": "1792x1024", "response_format": "b64_json", "quality": "high", @@ -30,20 +33,77 @@ class TestGeminiImageEditTransformation: drop_params=False, ) - assert mapped["aspectRatio"] == "16:9" + assert mapped["imageConfig"] == {"aspectRatio": "16:9"} + assert mapped["sampleCount"] == 2 assert "response_format" not in mapped assert "quality" not in mapped + def test_map_openai_params_with_image_size_for_gemini_3(self) -> None: + optional_params: Dict[str, object] = { + "size": "768x1376", + } + + mapped = self.config.map_openai_params( + image_edit_optional_params=optional_params, # type: ignore[arg-type] + model="gemini-3-pro-image-preview", + drop_params=False, + ) + + assert mapped["imageConfig"] == {"aspectRatio": "9:16", "imageSize": "1K"} + + def test_map_openai_params_forwards_image_config_as_is(self) -> None: + optional_params: Dict[str, object] = { + "size": "1024x1024", + "imageConfig": {"aspectRatio": "16:9", "imageSize": "512px"}, + } + + mapped = self.config.map_openai_params( + image_edit_optional_params=optional_params, # type: ignore[arg-type] + model="gemini-3-pro-image-preview", + drop_params=False, + ) + + assert mapped["imageConfig"] == {"aspectRatio": "16:9", "imageSize": "512px"} + + def test_map_openai_params_parses_form_image_config_json(self) -> None: + optional_params: Dict[str, object] = { + "imageConfig": '{"aspectRatio":"16:9","imageSize":"1K"}', + } + + mapped = self.config.map_openai_params( + image_edit_optional_params=optional_params, # type: ignore[arg-type] + model="gemini-3-pro-image-preview", + drop_params=False, + ) + + assert mapped["imageConfig"] == {"aspectRatio": "16:9", "imageSize": "1K"} + + def test_map_openai_params_rejects_malformed_form_image_config_json( + self, + ) -> None: + optional_params: Dict[str, object] = { + "imageConfig": "{bad", + } + + with pytest.raises(litellm.UnsupportedParamsError) as exc_info: + self.config.map_openai_params( + image_edit_optional_params=optional_params, # type: ignore[arg-type] + model="gemini-3-pro-image-preview", + drop_params=False, + ) + + assert "`imageConfig` must be valid JSON" in str(exc_info.value) + def test_transform_image_edit_request(self) -> None: image_bytes = b"fake_image_data" image = BytesIO(image_bytes) optional_params = { "sampleCount": 2, - "aspectRatio": "16:9", + "imageConfig": {"aspectRatio": "16:9", "imageSize": "2K"}, } request_body, files = self.config.transform_image_edit_request( - model=self.model, + model="gemini-3-pro-image-preview", prompt=self.prompt, image=[image], # Gemini pipeline passes list of images image_edit_optional_request_params=optional_params, @@ -61,7 +121,28 @@ class TestGeminiImageEditTransformation: assert base64.b64decode(inline_data["data"]) == image_bytes generation_config = request_body["generationConfig"] + assert generation_config["candidateCount"] == 2 assert generation_config["imageConfig"]["aspectRatio"] == "16:9" + assert generation_config["imageConfig"]["imageSize"] == "2K" + + def test_transform_image_edit_request_omits_image_size_for_gemini_25(self) -> None: + image = BytesIO(b"fake_image_data") + optional_params = { + "imageConfig": {"aspectRatio": "16:9", "imageSize": "2K"}, + } + + request_body, _ = self.config.transform_image_edit_request( + model=self.model, + prompt=self.prompt, + image=[image], + image_edit_optional_request_params=optional_params, + litellm_params=MagicMock(), + headers={}, + ) + + assert request_body["generationConfig"]["imageConfig"] == { + "aspectRatio": "16:9" + } def test_transform_image_edit_request_multiple_images(self) -> None: image_one = BytesIO(b"image_one") @@ -115,7 +196,16 @@ class TestGeminiImageEditTransformation: ] } }, - ] + ], + "usageMetadata": { + "promptTokenCount": 35, + "candidatesTokenCount": 1716, + "totalTokenCount": 1751, + "promptTokensDetails": [ + {"modality": "TEXT", "tokenCount": 30}, + {"modality": "IMAGE", "tokenCount": 5}, + ], + }, } mock_response = MagicMock(spec=httpx.Response) @@ -138,6 +228,19 @@ class TestGeminiImageEditTransformation: "utf-8" ) + usage = image_response.model_dump()["usage"] + assert usage["input_tokens"] == 35 + assert usage["output_tokens"] == 1716 + assert usage["prompt_tokens"] == 35 + assert usage["completion_tokens"] == 1716 + assert usage["prompt_tokens_details"]["image_tokens"] == 5 + assert usage["completion_tokens_details"]["image_tokens"] == 1716 + + logging_usage = StandardLoggingPayloadSetup.get_usage_as_dict( + response_obj=image_response.model_dump() + ) + assert logging_usage["completion_tokens_details"]["image_tokens"] == 1716 + def test_transform_image_edit_request_without_image_raises(self) -> None: optional_params = {} diff --git a/tests/test_litellm/llms/gemini/test_cost_calculator.py b/tests/test_litellm/llms/gemini/test_cost_calculator.py index 9bb83aa7cff..6d51bcd2c88 100644 --- a/tests/test_litellm/llms/gemini/test_cost_calculator.py +++ b/tests/test_litellm/llms/gemini/test_cost_calculator.py @@ -1,7 +1,23 @@ +import os + import pytest +import litellm from litellm.llms.gemini.cost_calculator import cost_per_web_search_request -from litellm.types.utils import PromptTokensDetailsWrapper, Usage +from litellm.llms.gemini.image_edit.cost_calculator import ( + cost_calculator as gemini_image_edit_cost_calculator, +) +from litellm.llms.gemini.image_generation.cost_calculator import ( + cost_calculator as gemini_image_generation_cost_calculator, +) +from litellm.types.utils import ( + ImageObject, + ImageResponse, + ImageUsage, + ImageUsageInputTokensDetails, + PromptTokensDetailsWrapper, + Usage, +) def _make_usage(web_search_requests: int) -> Usage: @@ -63,3 +79,171 @@ def test_no_usage_details(): usage = Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150) cost = cost_per_web_search_request(usage=usage, model_info=model_info) assert cost == 0.0 + + +def test_gemini_image_edit_cost_prefers_token_usage_metadata(): + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + model = "gemini/gemini-3-pro-image-preview" + model_info = litellm.get_model_info(model=model, custom_llm_provider="gemini") + + input_text_tokens = 20 + input_image_tokens = 1120 + output_image_tokens = 1120 + prompt_tokens = input_text_tokens + input_image_tokens + image_response = ImageResponse( + data=[ImageObject(b64_json="img1"), ImageObject(b64_json="img2")], + usage=ImageUsage( + input_tokens=prompt_tokens, + input_tokens_details=ImageUsageInputTokensDetails( + text_tokens=input_text_tokens, + image_tokens=input_image_tokens, + ), + output_tokens=output_image_tokens, + total_tokens=prompt_tokens + output_image_tokens, + ), + ) + + cost = gemini_image_edit_cost_calculator( + model=model, + image_response=image_response, + ) + + expected_cost = ( + prompt_tokens * model_info["input_cost_per_token"] + + output_image_tokens * model_info["output_cost_per_image_token"] + ) + flat_image_cost = ( + len(image_response.data or []) * model_info["output_cost_per_image"] + ) + assert round(cost, 10) == round(expected_cost, 10) + assert cost != flat_image_cost + + +def test_gemini_image_edit_cost_uses_output_token_details(): + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + model = "gemini/gemini-3-pro-image-preview" + model_info = litellm.get_model_info(model=model, custom_llm_provider="gemini") + + input_text_tokens = 20 + output_text_tokens = 213 + output_image_tokens = 1120 + output_tokens = output_text_tokens + output_image_tokens + image_response = ImageResponse( + data=[ImageObject(b64_json="img1")], + usage=ImageUsage( + input_tokens=input_text_tokens, + input_tokens_details=ImageUsageInputTokensDetails( + text_tokens=input_text_tokens, + image_tokens=0, + ), + output_tokens=output_tokens, + total_tokens=input_text_tokens + output_tokens, + prompt_tokens=input_text_tokens, + completion_tokens=output_tokens, + prompt_tokens_details={ + "text_tokens": input_text_tokens, + "image_tokens": 0, + }, + completion_tokens_details={ + "text_tokens": output_text_tokens, + "image_tokens": output_image_tokens, + }, + output_tokens_details={ + "text_tokens": output_text_tokens, + "image_tokens": output_image_tokens, + }, + ), + ) + + cost = gemini_image_edit_cost_calculator( + model=model, + image_response=image_response, + ) + + expected_cost = ( + input_text_tokens * model_info["input_cost_per_token"] + + output_text_tokens * model_info["output_cost_per_token"] + + output_image_tokens * model_info["output_cost_per_image_token"] + ) + all_output_as_image_cost = ( + input_text_tokens * model_info["input_cost_per_token"] + + (output_text_tokens + output_image_tokens) + * model_info["output_cost_per_image_token"] + ) + assert round(cost, 10) == round(expected_cost, 10) + assert cost != all_output_as_image_cost + + +def test_gemini_image_generation_cost_uses_output_token_details(): + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + model = "gemini/gemini-3-pro-image-preview" + model_info = litellm.get_model_info(model=model, custom_llm_provider="gemini") + + input_text_tokens = 20 + output_text_tokens = 213 + output_image_tokens = 1120 + output_tokens = output_text_tokens + output_image_tokens + image_response = ImageResponse( + data=[ImageObject(b64_json="img1")], + usage=ImageUsage( + input_tokens=input_text_tokens, + input_tokens_details=ImageUsageInputTokensDetails( + text_tokens=input_text_tokens, + image_tokens=0, + ), + output_tokens=output_tokens, + total_tokens=input_text_tokens + output_tokens, + prompt_tokens=input_text_tokens, + completion_tokens=output_tokens, + prompt_tokens_details={ + "text_tokens": input_text_tokens, + "image_tokens": 0, + }, + completion_tokens_details={ + "text_tokens": output_text_tokens, + "image_tokens": output_image_tokens, + }, + output_tokens_details={ + "text_tokens": output_text_tokens, + "image_tokens": output_image_tokens, + }, + ), + ) + + cost = gemini_image_generation_cost_calculator( + model=model, + image_response=image_response, + ) + + expected_cost = ( + input_text_tokens * model_info["input_cost_per_token"] + + output_text_tokens * model_info["output_cost_per_token"] + + output_image_tokens * model_info["output_cost_per_image_token"] + ) + all_output_as_image_cost = ( + input_text_tokens * model_info["input_cost_per_token"] + + (output_text_tokens + output_image_tokens) + * model_info["output_cost_per_image_token"] + ) + assert round(cost, 10) == round(expected_cost, 10) + assert cost != all_output_as_image_cost + + +def test_gemini_image_edit_cost_falls_back_to_flat_image_pricing(): + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + model = "gemini/gemini-3-pro-image-preview" + model_info = litellm.get_model_info(model=model, custom_llm_provider="gemini") + image_response = ImageResponse( + data=[ImageObject(b64_json="img1"), ImageObject(b64_json="img2")] + ) + + cost = gemini_image_edit_cost_calculator( + model=model, + image_response=image_response, + ) + + assert cost == len(image_response.data or []) * model_info["output_cost_per_image"] diff --git a/tests/test_litellm/llms/gemini/test_gemini_image_generation_transformation.py b/tests/test_litellm/llms/gemini/test_gemini_image_generation_transformation.py new file mode 100644 index 00000000000..4610d1b99bf --- /dev/null +++ b/tests/test_litellm/llms/gemini/test_gemini_image_generation_transformation.py @@ -0,0 +1,240 @@ +import httpx + +from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup +from litellm.llms.gemini.image_generation.transformation import GoogleImageGenConfig +from litellm.types.utils import ImageResponse + + +def test_gemini_image_generation_request_uses_shared_generation_config(): + config = GoogleImageGenConfig() + + request = config.transform_image_generation_request( + model="gemini-3.1-flash-image-preview", + prompt="Generate a simple app icon", + optional_params={ + "sampleCount": 2, + "imageConfig": {"aspectRatio": "16:9", "imageSize": "2K"}, + }, + litellm_params={}, + headers={}, + ) + + assert request["contents"][0]["parts"] == [{"text": "Generate a simple app icon"}] + assert request["generationConfig"] == { + "response_modalities": ["IMAGE", "TEXT"], + "imageConfig": {"aspectRatio": "16:9", "imageSize": "2K"}, + "candidateCount": 2, + } + + +def test_gemini_image_generation_map_openai_params_maps_n_size_and_image_config(): + config = GoogleImageGenConfig() + + mapped = config.map_openai_params( + non_default_params={ + "n": 2, + "size": "768x1376", + "imageConfig": {"aspectRatio": "1:1", "imageSize": "512"}, + }, + optional_params={}, + model="gemini-3.1-flash-image-preview", + drop_params=False, + ) + + assert mapped == { + "sampleCount": 2, + "imageConfig": {"aspectRatio": "1:1", "imageSize": "512"}, + } + + +def test_imagen_generation_with_provider_prefix_uses_imagen_params_and_response(): + config = GoogleImageGenConfig() + + mapped = config.map_openai_params( + non_default_params={ + "n": 1, + "size": "1024x1024", + }, + optional_params={}, + model="gemini/imagen-4.0-generate-001", + drop_params=False, + ) + assert mapped == { + "sampleCount": 1, + "aspectRatio": "1:1", + "imageSize": "1K", + } + + request = config.transform_image_generation_request( + model="gemini/imagen-4.0-generate-001", + prompt="Generate a simple app icon", + optional_params=mapped, + litellm_params={}, + headers={}, + ) + assert request == { + "instances": [{"prompt": "Generate a simple app icon"}], + "parameters": { + "sampleCount": 1, + "aspectRatio": "1:1", + "imageSize": "1K", + }, + } + + result = config.transform_image_generation_response( + model="gemini/imagen-4.0-generate-001", + raw_response=httpx.Response( + status_code=200, + json={ + "predictions": [ + { + "bytesBase64Encoded": "fake-imagen-image", + } + ] + }, + ), + model_response=ImageResponse(data=[]), + logging_obj=None, + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + assert result.data is not None + assert result.data[0].b64_json == "fake-imagen-image" + + +def test_imagen_generation_forwards_mapped_openai_size_image_size(): + config = GoogleImageGenConfig() + + mapped = config.map_openai_params( + non_default_params={ + "size": "512x512", + }, + optional_params={}, + model="gemini/imagen-4.0-generate-001", + drop_params=False, + ) + assert mapped == {"aspectRatio": "1:1", "imageSize": "512"} + + request = config.transform_image_generation_request( + model="gemini/imagen-4.0-generate-001", + prompt="Generate a simple app icon", + optional_params=mapped, + litellm_params={}, + headers={}, + ) + + assert request == { + "instances": [{"prompt": "Generate a simple app icon"}], + "parameters": {"aspectRatio": "1:1", "imageSize": "512"}, + } + + +def test_gemini_image_generation_usage_includes_chat_token_details(): + config = GoogleImageGenConfig() + raw_response = httpx.Response( + status_code=200, + json={ + "candidates": [ + { + "content": { + "parts": [ + { + "inlineData": { + "mimeType": "image/png", + "data": "fake-image", + } + } + ] + } + } + ], + "usageMetadata": { + "promptTokenCount": 35, + "candidatesTokenCount": 1716, + "totalTokenCount": 1751, + "promptTokensDetails": [ + {"modality": "TEXT", "tokenCount": 30}, + {"modality": "IMAGE", "tokenCount": 5}, + ], + "candidatesTokensDetails": [ + {"modality": "TEXT", "tokenCount": 213}, + {"modality": "IMAGE", "tokenCount": 1120}, + ], + }, + }, + ) + + result = config.transform_image_generation_response( + model="gemini-3.1-flash-image-preview", + raw_response=raw_response, + model_response=ImageResponse(data=[]), + logging_obj=None, + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + usage = result.model_dump()["usage"] + + assert usage["input_tokens"] == 35 + assert usage["output_tokens"] == 1716 + assert usage["prompt_tokens"] == 35 + assert usage["completion_tokens"] == 1716 + assert usage["prompt_tokens_details"]["image_tokens"] == 5 + assert usage["completion_tokens_details"]["text_tokens"] == 596 + assert usage["completion_tokens_details"]["image_tokens"] == 1120 + assert usage["output_tokens_details"]["text_tokens"] == 596 + assert usage["output_tokens_details"]["image_tokens"] == 1120 + + logging_usage = StandardLoggingPayloadSetup.get_usage_as_dict( + response_obj=result.model_dump() + ) + assert logging_usage["completion_tokens_details"]["text_tokens"] == 596 + assert logging_usage["completion_tokens_details"]["image_tokens"] == 1120 + + +def test_gemini_image_generation_usage_without_output_details_treats_output_as_image(): + config = GoogleImageGenConfig() + raw_response = httpx.Response( + status_code=200, + json={ + "candidates": [ + { + "content": { + "parts": [ + { + "inlineData": { + "mimeType": "image/png", + "data": "fake-image", + } + } + ] + } + } + ], + "usageMetadata": { + "promptTokenCount": 35, + "candidatesTokenCount": 1716, + "totalTokenCount": 1751, + "promptTokensDetails": [{"modality": "TEXT", "tokenCount": 35}], + }, + }, + ) + + result = config.transform_image_generation_response( + model="gemini-3.1-flash-image-preview", + raw_response=raw_response, + model_response=ImageResponse(data=[]), + logging_obj=None, + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + usage = result.model_dump()["usage"] + assert usage["completion_tokens_details"]["text_tokens"] == 0 + assert usage["completion_tokens_details"]["image_tokens"] == 1716 diff --git a/tests/test_litellm/llms/inception/__init__.py b/tests/test_litellm/llms/inception/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/inception/test_inception_chat_transformation.py b/tests/test_litellm/llms/inception/test_inception_chat_transformation.py new file mode 100644 index 00000000000..0750fb9e405 --- /dev/null +++ b/tests/test_litellm/llms/inception/test_inception_chat_transformation.py @@ -0,0 +1,326 @@ +""" +Tests for Inception (Mercury) chat provider integration +""" + +import json +import os +from unittest import mock + +import httpx + +import litellm +from litellm.llms.inception.chat.transformation import InceptionChatConfig + + +def test_inception_config_initialization(): + config = InceptionChatConfig() + assert config.custom_llm_provider == "inception" + + +def test_inception_chat_supports_diffusion_params(): + """The chat config must expose Inception's diffusion-LLM request controls""" + params = InceptionChatConfig().get_supported_openai_params("mercury-2") + for p in ( + "reasoning_effort", + "reasoning_summary", + "reasoning_summary_wait", + "diffusing", + "realtime", + "tools", + "tool_choice", + "response_format", + ): + assert p in params, f"{p} should be a supported chat param" + + +def test_inception_chat_sends_diffusion_params_in_body(): + """reasoning_effort (incl. `instant`) and the diffusion flags reach the request body""" + + captured = {} + + def fake_send(self, request, **kwargs): + captured["body"] = json.loads(request.content.decode()) + return httpx.Response( + status_code=200, + request=request, + headers={"content-type": "application/json"}, + content=json.dumps( + { + "id": "c-1", + "object": "chat.completion", + "created": 1, + "model": "mercury-2", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "hi"}, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 5, + "completion_tokens": 1, + "total_tokens": 6, + }, + } + ).encode(), + ) + + with mock.patch("httpx.Client.send", new=fake_send): + litellm.completion( + model="inception/mercury-2", + messages=[{"role": "user", "content": "hi"}], + api_key="sk-x", + reasoning_effort="instant", + reasoning_summary=True, + reasoning_summary_wait=True, + diffusing=True, + realtime=True, + max_completion_tokens=128, + ) + + body = captured["body"] + assert body["reasoning_effort"] == "instant" + assert body["reasoning_summary"] is True + assert body["reasoning_summary_wait"] is True + assert body["diffusing"] is True + assert body["realtime"] is True + assert body["max_tokens"] == 128 # max_completion_tokens mapped to max_tokens + + +def test_inception_chat_response_surfaces_reasoning_and_usage(): + """reasoning_summary / warning survive, and reasoning_tokens maps to usage details""" + + def fake_send(self, request, **kwargs): + return httpx.Response( + status_code=200, + request=request, + headers={"content-type": "application/json"}, + content=json.dumps( + { + "id": "c-1", + "object": "chat.completion", + "created": 1, + "model": "mercury-2", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "answer"}, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 5, + "completion_tokens": 2, + "total_tokens": 7, + "reasoning_tokens": 4, + "cached_input_tokens": 3, + }, + "reasoning_summary": { + "content": "step by step", + "status": "complete", + }, + "warning": "heads up", + } + ).encode(), + ) + + with mock.patch("httpx.Client.send", new=fake_send): + r = litellm.completion( + model="inception/mercury-2", + messages=[{"role": "user", "content": "hi"}], + api_key="sk-x", + ) + + assert r.reasoning_summary == {"content": "step by step", "status": "complete"} + assert r.warning == "heads up" + assert r.usage.completion_tokens_details.reasoning_tokens == 4 + assert r.usage.model_extra.get("cached_input_tokens") == 3 + + +def test_inception_get_openai_compatible_provider_info(): + config = InceptionChatConfig() + + with mock.patch.dict(os.environ, {}, clear=True): + with mock.patch.object(litellm, "inception_key", None): + api_base, api_key = config._get_openai_compatible_provider_info(None, None) + assert api_base == "https://api.inceptionlabs.ai/v1" + assert api_key is None + + with mock.patch.dict( + os.environ, + { + "INCEPTION_API_KEY": "test-key", + "INCEPTION_API_BASE": "https://custom.inceptionlabs.ai/v1", + }, + ): + api_base, api_key = config._get_openai_compatible_provider_info(None, None) + assert api_base == "https://custom.inceptionlabs.ai/v1" + assert api_key == "test-key" + + with mock.patch.dict( + os.environ, + { + "INCEPTION_API_KEY": "env-key", + "INCEPTION_API_BASE": "https://env.inceptionlabs.ai/v1", + }, + ): + api_base, api_key = config._get_openai_compatible_provider_info( + "https://param.inceptionlabs.ai/v1", "param-key" + ) + assert api_base == "https://param.inceptionlabs.ai/v1" + assert api_key == "param-key" + + +def test_inception_key_module_attr_fallback(): + """litellm.inception_key is used when no param/env key is provided""" + config = InceptionChatConfig() + with mock.patch.dict(os.environ, {}, clear=True): + with mock.patch.object(litellm, "inception_key", "module-attr-key"): + _, api_key = config._get_openai_compatible_provider_info(None, None) + assert api_key == "module-attr-key" + + +def test_inception_does_not_leak_key_to_caller_api_base(): + """ + The server-managed Inception key must not be forwarded to a caller-supplied + api_base. It is only resolved for the default/server base, or when the + caller also supplies their own key. + """ + config = InceptionChatConfig() + with mock.patch.dict( + os.environ, {"INCEPTION_API_KEY": "server-secret"}, clear=True + ): + with mock.patch.object(litellm, "inception_key", "module-secret"): + # caller overrides api_base without a key -> server key withheld + api_base, api_key = config._get_openai_compatible_provider_info( + "https://attacker.example/v1", None + ) + assert api_base == "https://attacker.example/v1" + assert api_key is None + + # caller overrides api_base AND supplies their own key -> used as-is + _, api_key = config._get_openai_compatible_provider_info( + "https://attacker.example/v1", "caller-key" + ) + assert api_key == "caller-key" + + # default/server base -> server-managed key resolved + _, api_key = config._get_openai_compatible_provider_info(None, None) + assert api_key == "module-secret" + + +def test_get_llm_provider_inception(): + from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider + + model, provider, _, _ = get_llm_provider("inception/mercury-2") + assert model == "mercury-2" + assert provider == "inception" + + model, provider, _, api_base = get_llm_provider( + "mercury-2", api_base="https://api.inceptionlabs.ai/v1" + ) + assert model == "mercury-2" + assert provider == "inception" + assert api_base == "https://api.inceptionlabs.ai/v1" + + +def test_inception_in_provider_lists(): + assert "inception" in litellm.openai_compatible_providers + assert "inception" in litellm.provider_list + assert "https://api.inceptionlabs.ai/v1" in litellm.openai_compatible_endpoints + + +def test_inception_model_configuration(): + from litellm import get_model_info + + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + litellm.inception_models = set() + litellm.add_known_models() + + info = get_model_info("inception/mercury-2") + assert info.get("litellm_provider") == "inception" + assert info.get("mode") == "chat" + assert info.get("max_input_tokens") == 128000 + assert info.get("input_cost_per_token") == 2.5e-07 + assert info.get("output_cost_per_token") == 7.5e-07 + assert info.get("cache_read_input_token_cost") == 2.5e-08 + assert info.get("supports_function_calling") is True + assert info.get("supports_tool_choice") is True + assert info.get("supports_response_schema") is True + + +def test_inception_model_list_populated(): + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + litellm.inception_models = set() + litellm.add_known_models() + + assert "inception/mercury-2" in litellm.inception_models + for model in litellm.inception_models: + assert model.startswith("inception/") + + +def test_inception_completion_targets_inception_endpoint(): + """ + End-to-end: a completion routed through the inception provider must hit + Inception's base URL and path, send a Bearer token, strip the + `inception/` prefix from the model name, and forward tool_choice. + """ + + captured = {} + + def fake_send(self, request, **kwargs): + captured["url"] = str(request.url) + captured["auth"] = request.headers.get("authorization") + captured["body"] = json.loads(request.content.decode()) + return httpx.Response( + status_code=200, + request=request, + headers={"content-type": "application/json"}, + content=json.dumps( + { + "id": "cmpl-1", + "object": "chat.completion", + "created": 1, + "model": "mercury-2", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "hi"}, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 5, + "completion_tokens": 1, + "total_tokens": 6, + }, + } + ).encode(), + ) + + tools = [ + { + "type": "function", + "function": { + "name": "f", + "parameters": {"type": "object", "properties": {}}, + }, + } + ] + with mock.patch("httpx.Client.send", new=fake_send): + response = litellm.completion( + model="inception/mercury-2", + messages=[{"role": "user", "content": "hello"}], + api_key="sk-test-fake-123", + tools=tools, + tool_choice="auto", + ) + + assert captured["url"] == "https://api.inceptionlabs.ai/v1/chat/completions" + assert captured["auth"] == "Bearer sk-test-fake-123" + assert captured["body"]["model"] == "mercury-2" + assert captured["body"]["tool_choice"] == "auto" + assert response.choices[0].message.content == "hi" diff --git a/tests/test_litellm/llms/inception/test_inception_completion_transformation.py b/tests/test_litellm/llms/inception/test_inception_completion_transformation.py new file mode 100644 index 00000000000..9b7c8dd3742 --- /dev/null +++ b/tests/test_litellm/llms/inception/test_inception_completion_transformation.py @@ -0,0 +1,300 @@ +""" +Tests for Inception (Mercury) fill-in-the-middle (FIM) provider integration +""" + +import json +import os +from unittest import mock + +import httpx +import pytest + +import litellm +from litellm.llms.inception.completion.transformation import ( + InceptionTextCompletionConfig, +) + + +def _fim_response_bytes(): + return json.dumps( + { + "id": "fim-1", + "object": "text_completion", + "created": 1, + "model": "mercury-edit-2", + "choices": [ + {"text": "a + b", "index": 0, "finish_reason": "stop", "logprobs": None} + ], + "usage": {"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8}, + } + ).encode() + + +def test_inception_fim_supports_suffix_param(): + """The FIM config must keep `suffix` (otherwise FIM requests lose context)""" + config = InceptionTextCompletionConfig() + assert "suffix" in config.get_supported_openai_params("mercury-edit-2") + + mapped = config.map_openai_params( + non_default_params={"suffix": "\n return x", "max_completion_tokens": 50}, + optional_params={}, + model="mercury-edit-2", + drop_params=False, + ) + assert mapped["suffix"] == "\n return x" + assert mapped["max_tokens"] == 50 + + +def test_inception_fim_supported_params_match_schema(): + """FIM exposes the OpenAI subset of Inception's FIMCompletionRequest only""" + params = InceptionTextCompletionConfig().get_supported_openai_params( + "mercury-edit-2" + ) + for p in ("suffix", "top_p", "frequency_penalty", "presence_penalty", "stop"): + assert p in params + # Chat-only sampling controls are not part of Inception's FIM schema + for p in ("temperature", "seed", "logprobs", "n", "user"): + assert p not in params + + +def test_text_completion_inception_in_provider_lists(): + from litellm.types.utils import LlmProviders + + assert LlmProviders.TEXT_COMPLETION_INCEPTION == "text-completion-inception" + assert "text-completion-inception" in litellm.provider_list + + +def test_inception_get_supported_openai_params_dispatch(): + """litellm.get_supported_openai_params routes the FIM provider to our config""" + params = litellm.get_supported_openai_params( + model="mercury-edit-2", custom_llm_provider="text-completion-inception" + ) + assert "suffix" in params + assert "temperature" not in params + + +@pytest.mark.parametrize("provider", ["inception", "text-completion-inception"]) +def test_inception_validate_environment(provider): + model = ( + "inception/mercury-2" + if provider == "inception" + else "text-completion-inception/mercury-edit-2" + ) + + with mock.patch.dict(os.environ, {}, clear=True): + result = litellm.validate_environment(model) + assert result["keys_in_environment"] is False + assert "INCEPTION_API_KEY" in result["missing_keys"] + + with mock.patch.dict(os.environ, {"INCEPTION_API_KEY": "sk-x"}, clear=True): + result = litellm.validate_environment(model) + assert result["keys_in_environment"] is True + + +def test_inception_completion_endpoint_returns_chat_object(): + """ + Calling chat `completion()` with the FIM provider converts the text + completion result into a chat-shaped ModelResponse. + """ + + def fake_send(self, request, **kwargs): + return httpx.Response( + status_code=200, + request=request, + headers={"content-type": "application/json"}, + content=_fim_response_bytes(), + ) + + with mock.patch("httpx.Client.send", new=fake_send): + r = litellm.completion( + model="text-completion-inception/mercury-edit-2", + messages=[{"role": "user", "content": "def add(a, b): return "}], + api_key="sk-x", + ) + + assert r.choices[0].message.content == "a + b" + + +@pytest.mark.asyncio +async def test_inception_fim_async(): + """async FIM path (acompletion) hits Inception's /v1/fim/completions""" + + captured = {} + + async def fake_asend(self, request, **kwargs): + captured["url"] = str(request.url) + return httpx.Response( + status_code=200, + request=request, + headers={"content-type": "application/json"}, + content=_fim_response_bytes(), + ) + + with mock.patch("httpx.AsyncClient.send", new=fake_asend): + r = await litellm.atext_completion( + model="text-completion-inception/mercury-edit-2", + prompt="def add(a, b): return ", + suffix="\n", + api_key="sk-x", + max_tokens=10, + ) + + assert captured["url"] == "https://api.inceptionlabs.ai/v1/fim/completions" + assert r.choices[0].text == "a + b" + + +def test_inception_fim_model_configuration(): + from litellm import get_model_info + + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + litellm.text_completion_inception_models = set() + litellm.add_known_models() + + assert ( + "text-completion-inception/mercury-edit-2" + in litellm.text_completion_inception_models + ) + info = get_model_info("text-completion-inception/mercury-edit-2") + assert info.get("litellm_provider") == "text-completion-inception" + assert info.get("mode") == "completion" + assert info.get("max_input_tokens") == 32000 + + +def test_inception_fim_targets_fim_endpoint(): + """ + End-to-end: a FIM request must hit `/v1/fim/completions` (NOT + `/v1/completions`), carry the `suffix`, and parse the standard `text` field. + """ + + captured = {} + + def fake_send(self, request, **kwargs): + captured["url"] = str(request.url) + captured["auth"] = request.headers.get("authorization") + captured["body"] = json.loads(request.content.decode()) + return httpx.Response( + status_code=200, + request=request, + headers={"content-type": "application/json"}, + content=json.dumps( + { + "id": "fim-1", + "object": "text_completion", + "created": 1, + "model": "mercury-edit-2", + "choices": [ + { + "text": "a + b", + "index": 0, + "finish_reason": "stop", + "logprobs": None, + } + ], + "usage": { + "prompt_tokens": 5, + "completion_tokens": 3, + "total_tokens": 8, + }, + } + ).encode(), + ) + + with mock.patch("httpx.Client.send", new=fake_send): + response = litellm.text_completion( + model="text-completion-inception/mercury-edit-2", + prompt="def add(a, b):\n return ", + suffix="\n", + api_key="sk-fim-fake", + max_tokens=20, + ) + + assert captured["url"] == "https://api.inceptionlabs.ai/v1/fim/completions" + assert captured["auth"] == "Bearer sk-fim-fake" + assert captured["body"]["model"] == "mercury-edit-2" + assert captured["body"]["suffix"] == "\n" + assert "prompt" in captured["body"] + assert response.choices[0].text == "a + b" + + +def test_inception_fim_does_not_leak_global_api_key(): + """ + Regression: the global litellm.api_key (commonly an OpenAI key) must not be + forwarded to Inception. Only an Inception-specific key (param, + litellm.inception_key, or INCEPTION_API_KEY) may be sent to the Inception base. + """ + + captured = {} + + def fake_send(self, request, **kwargs): + captured["auth"] = request.headers.get("authorization") + return httpx.Response( + status_code=200, + request=request, + headers={"content-type": "application/json"}, + content=_fim_response_bytes(), + ) + + with mock.patch.dict( + os.environ, {"INCEPTION_API_KEY": "sk-inception-correct"}, clear=True + ): + with mock.patch.object(litellm, "inception_key", None): + with mock.patch.object(litellm, "api_key", "sk-global-should-not-leak"): + with mock.patch("httpx.Client.send", new=fake_send): + litellm.text_completion( + model="text-completion-inception/mercury-edit-2", + prompt="def add(a, b): return ", + max_tokens=10, + ) + + assert captured["auth"] == "Bearer sk-inception-correct" + + +def test_inception_fim_extra_body_forwards_vllm_params(): + """top_k / repetition_penalty are reachable via extra_body (not OpenAI params)""" + + captured = {} + + def fake_send(self, request, **kwargs): + captured["body"] = json.loads(request.content.decode()) + return httpx.Response( + status_code=200, + request=request, + headers={"content-type": "application/json"}, + content=json.dumps( + { + "id": "f-1", + "object": "text_completion", + "created": 1, + "model": "mercury-edit-2", + "choices": [ + { + "text": "x", + "index": 0, + "finish_reason": "stop", + "logprobs": None, + } + ], + "usage": { + "prompt_tokens": 2, + "completion_tokens": 1, + "total_tokens": 3, + }, + } + ).encode(), + ) + + with mock.patch("httpx.Client.send", new=fake_send): + litellm.text_completion( + model="text-completion-inception/mercury-edit-2", + prompt="def f(", + suffix=")", + api_key="sk-x", + top_p=0.9, + extra_body={"top_k": 40, "repetition_penalty": 1.1}, + ) + + body = captured["body"] + assert body["top_p"] == 0.9 + assert body["top_k"] == 40 + assert body["repetition_penalty"] == 1.1 diff --git a/tests/test_litellm/llms/langflow/chat/test_langflow_chat_transformation.py b/tests/test_litellm/llms/langflow/chat/test_langflow_chat_transformation.py new file mode 100644 index 00000000000..c03919a0659 --- /dev/null +++ b/tests/test_litellm/llms/langflow/chat/test_langflow_chat_transformation.py @@ -0,0 +1,398 @@ +from unittest.mock import MagicMock, patch + +import httpx +import pytest + +from litellm.llms.langflow.chat.transformation import LangFlowConfig, LangFlowError +from litellm.types.utils import LlmProviders, ModelResponse +from litellm.utils import ProviderConfigManager + + +def test_flow_id_cannot_be_overridden_via_optional_params(): + config = LangFlowConfig() + url = config.get_complete_url( + api_base="http://localhost:7860", + api_key=None, + model="langflow/authorized-flow", + optional_params={}, + litellm_params={}, + stream=False, + ) + assert url.endswith("/api/v1/run/authorized-flow") + + with pytest.raises(LangFlowError): + config.get_complete_url( + api_base="http://localhost:7860", + api_key=None, + model="langflow/authorized-flow", + optional_params={"flow_id": "malicious-flow"}, + litellm_params={}, + stream=False, + ) + + +def test_langflow_config_get_complete_url(): + config = LangFlowConfig() + url = config.get_complete_url( + api_base="http://localhost:7860", + api_key=None, + model="langflow/my-flow-id", + optional_params={}, + litellm_params={}, + stream=False, + ) + assert url == "http://localhost:7860/api/v1/run/my-flow-id" + + +def test_langflow_config_get_complete_url_requires_api_base(): + config = LangFlowConfig() + with pytest.raises(ValueError): + config.get_complete_url( + api_base=None, + api_key=None, + model="langflow/my-flow-id", + optional_params={}, + litellm_params={}, + stream=False, + ) + + +def test_langflow_config_flow_id_is_path_segment_encoded(): + config = LangFlowConfig() + url = config.get_complete_url( + api_base="http://localhost:7860", + api_key=None, + model="langflow/../../secret?x=1", + optional_params={}, + litellm_params={}, + stream=False, + ) + assert url == "http://localhost:7860/api/v1/run/..%2F..%2Fsecret%3Fx%3D1" + assert "/api/v1/run/" in url + assert url.rsplit("/api/v1/run/", 1)[1] not in ("..", "../..") + + +@pytest.mark.parametrize("model", ["langflow/", "langflow/ "]) +def test_langflow_config_rejects_empty_flow_id(model): + config = LangFlowConfig() + with pytest.raises(LangFlowError): + config.get_complete_url( + api_base="http://localhost:7860", + api_key=None, + model=model, + optional_params={}, + litellm_params={}, + stream=False, + ) + + +def test_langflow_config_strips_flow_id_whitespace(): + config = LangFlowConfig() + url = config.get_complete_url( + api_base="http://localhost:7860", + api_key=None, + model="langflow/ my-flow-id ", + optional_params={}, + litellm_params={}, + stream=False, + ) + assert url == "http://localhost:7860/api/v1/run/my-flow-id" + + +def test_langflow_config_transform_request_includes_session_id(): + config = LangFlowConfig() + request = config.transform_request( + model="langflow/my-flow-id", + messages=[{"role": "user", "content": "hello"}], + optional_params={"session_id": "sess-abc"}, + litellm_params={}, + headers={}, + ) + + assert request["input_value"] == "hello" + assert request["input_type"] == "chat" + assert request["output_type"] == "chat" + assert request["session_id"] == "sess-abc" + + +def test_langflow_config_transform_request_uses_last_user_message(): + config = LangFlowConfig() + request = config.transform_request( + model="langflow/my-flow-id", + messages=[ + {"role": "system", "content": "be helpful"}, + {"role": "user", "content": "first"}, + {"role": "assistant", "content": "ok"}, + {"role": "user", "content": [{"type": "text", "text": "second"}]}, + ], + optional_params={}, + litellm_params={}, + headers={}, + ) + + assert request["input_value"] == "second" + assert "session_id" not in request + + +def test_langflow_config_transform_request_falls_back_to_last_message(): + config = LangFlowConfig() + request = config.transform_request( + model="langflow/my-flow-id", + messages=[{"role": "assistant", "content": "only assistant"}], + optional_params={}, + litellm_params={}, + headers={}, + ) + + assert request["input_value"] == "only assistant" + + +def test_langflow_config_transform_request_empty_messages(): + config = LangFlowConfig() + request = config.transform_request( + model="langflow/my-flow-id", + messages=[], + optional_params={}, + litellm_params={}, + headers={}, + ) + + assert request["input_value"] == "" + + +def test_langflow_config_rejects_tweaks_from_request_params(): + config = LangFlowConfig() + with pytest.raises(LangFlowError): + config.transform_request( + model="langflow/my-flow-id", + messages=[{"role": "user", "content": "hi"}], + optional_params={"tweaks": {"HttpComponent": {"url": "http://attacker"}}}, + litellm_params={}, + headers={}, + ) + + +def test_langflow_config_rejects_tweaks_from_request_body(): + config = LangFlowConfig() + with pytest.raises(LangFlowError): + config.sign_request( + headers={}, + optional_params={}, + request_data={ + "input_value": "hi", + "tweaks": {"HttpComponent": {"url": "http://attacker"}}, + }, + api_base="http://localhost:7860", + ) + + +def test_langflow_config_sign_request_passes_through_without_tweaks(): + config = LangFlowConfig() + headers, body = config.sign_request( + headers={"x-api-key": "secret"}, + optional_params={}, + request_data={"input_value": "hi"}, + api_base="http://localhost:7860", + ) + assert headers == {"x-api-key": "secret"} + assert body is None + + +def test_langflow_config_validate_environment_sets_api_key_header(): + config = LangFlowConfig() + headers = config.validate_environment( + headers={}, + model="langflow/my-flow-id", + messages=[{"role": "user", "content": "hi"}], + optional_params={}, + litellm_params={}, + api_key="secret", + ) + assert headers["Content-Type"] == "application/json" + assert headers["x-api-key"] == "secret" + + +def test_langflow_extra_body_cannot_inject_tweaks_into_run_payload(): + import json + + import litellm + from litellm.llms.custom_httpx.http_handler import HTTPHandler + + posted_bodies = [] + + def fake_post(*args, **kwargs): + body = kwargs.get("data") + posted_bodies.append(json.loads(body) if isinstance(body, str) else body) + resp = MagicMock(spec=httpx.Response) + resp.status_code = 200 + resp.json.return_value = { + "outputs": [{"outputs": [{"results": {"message": {"text": "hi"}}}]}] + } + resp.headers = {} + resp.text = "{}" + return resp + + with patch.object(HTTPHandler, "post", side_effect=fake_post): + with pytest.raises(Exception): + litellm.completion( + model="langflow/my-flow", + messages=[{"role": "user", "content": "hello"}], + api_base="http://example.com", + api_key="sk-test", + extra_body={"tweaks": {"HttpComponent": {"url": "http://attacker"}}}, + ) + + assert all("tweaks" not in (body or {}) for body in posted_bodies) + + +def test_langflow_config_extract_response(): + config = LangFlowConfig() + content = config._extract_content_from_response( + { + "session_id": "sess-abc", + "outputs": [ + { + "outputs": [ + { + "results": { + "message": {"text": "Hello from LangFlow"}, + } + } + ] + } + ], + } + ) + assert content == "Hello from LangFlow" + + +def test_langflow_config_extract_response_from_outputs_dict(): + config = LangFlowConfig() + content = config._extract_content_from_response( + { + "outputs": [ + { + "outputs": [ + { + "results": {}, + "outputs": { + "message": {"message": {"text": "via outputs dict"}} + }, + } + ] + } + ], + } + ) + assert content == "via outputs dict" + + +def test_langflow_extract_response_returns_none_when_no_message(): + config = LangFlowConfig() + assert config._extract_content_from_response({"outputs": []}) is None + assert config._extract_content_from_response({"detail": "flow failed"}) is None + assert config._extract_content_from_response({"outputs": ["not-a-dict"]}) is None + assert ( + config._extract_content_from_response({"outputs": [{"outputs": ["bad"]}]}) + is None + ) + assert ( + config._extract_content_from_response( + {"outputs": [{"outputs": [{"results": {"message": {"text": ""}}}]}]} + ) + is None + ) + + +def test_langflow_transform_response_builds_model_response_with_usage(): + config = LangFlowConfig() + raw_response = httpx.Response( + status_code=200, + json={ + "session_id": "sess-abc", + "outputs": [ + {"outputs": [{"results": {"message": {"text": "Hello from LangFlow"}}}]} + ], + }, + ) + + result = config.transform_response( + model="langflow/my-flow-id", + raw_response=raw_response, + model_response=ModelResponse(), + logging_obj=None, + request_data={}, + messages=[{"role": "user", "content": "hi"}], + optional_params={}, + litellm_params={}, + encoding=None, + ) + + assert result.choices[0].message.content == "Hello from LangFlow" + assert result.choices[0].finish_reason == "stop" + assert result.model == "langflow/my-flow-id" + assert result.usage.completion_tokens > 0 + assert result.usage.total_tokens == ( + result.usage.prompt_tokens + result.usage.completion_tokens + ) + + +def test_langflow_transform_response_raises_on_unparseable_body(): + config = LangFlowConfig() + raw_response = httpx.Response(status_code=200, json={"detail": "flow failed"}) + + with pytest.raises(LangFlowError): + config.transform_response( + model="langflow/my-flow-id", + raw_response=raw_response, + model_response=ModelResponse(), + logging_obj=None, + request_data={}, + messages=[{"role": "user", "content": "hi"}], + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +def test_langflow_transform_response_raises_on_non_json_body(): + config = LangFlowConfig() + raw_response = httpx.Response( + status_code=200, content=b"not json", headers={"content-type": "text/plain"} + ) + + with pytest.raises(LangFlowError): + config.transform_response( + model="langflow/my-flow-id", + raw_response=raw_response, + model_response=ModelResponse(), + logging_obj=None, + request_data={}, + messages=[{"role": "user", "content": "hi"}], + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +def test_langflow_config_get_error_class(): + config = LangFlowConfig() + err = config.get_error_class(error_message="boom", status_code=503, headers={}) + assert isinstance(err, LangFlowError) + assert err.status_code == 503 + + +def test_langflow_config_stream_behavior_flags(): + config = LangFlowConfig() + assert config.supports_stream_param_in_request_body is False + assert config.should_fake_stream(model="langflow/x", stream=True) is True + assert config.should_fake_stream(model="langflow/x", stream=False) is False + + +def test_langflow_provider_config_registered(): + cfg = ProviderConfigManager.get_provider_chat_config( + model="langflow/flow-1", + provider=LlmProviders.LANGFLOW, + ) + assert cfg is not None + assert cfg.__class__.__name__ == "LangFlowConfig" diff --git a/tests/test_litellm/llms/langflow/test_langflow_a2a.py b/tests/test_litellm/llms/langflow/test_langflow_a2a.py new file mode 100644 index 00000000000..c49ec8d87c2 --- /dev/null +++ b/tests/test_litellm/llms/langflow/test_langflow_a2a.py @@ -0,0 +1,159 @@ +from unittest.mock import AsyncMock, patch + +import pytest + +from litellm.a2a_protocol.litellm_completion_bridge.handler import ( + A2A_USER_API_KEY_HASH_PARAM, +) +from litellm.a2a_protocol.providers.config_manager import A2AProviderConfigManager +from litellm.llms.langflow.a2a import merge_a2a_session_into_litellm_params + + +def test_merge_a2a_session_into_litellm_params(): + merged = merge_a2a_session_into_litellm_params( + {"custom_llm_provider": "langflow", "model": "langflow/flow-1"}, + {"message": {"contextId": "shared-session-99"}}, + ) + assert merged["session_id"] == "shared-session-99" + + +def test_merge_a2a_session_is_scoped_per_principal(): + """The LangFlow session must be bound to the authenticated key so two + distinct keys cannot share memory by reusing the same A2A contextId, while + the same key keeps a stable session across turns.""" + base = {"custom_llm_provider": "langflow", "model": "langflow/flow-1"} + params = {"message": {"contextId": "ctx-1"}} + + key_a = merge_a2a_session_into_litellm_params(base, params, "hash-a")["session_id"] + key_a_again = merge_a2a_session_into_litellm_params(base, params, "hash-a")[ + "session_id" + ] + key_b = merge_a2a_session_into_litellm_params(base, params, "hash-b")["session_id"] + + assert key_a == key_a_again, "same key + contextId must stay on one session" + assert key_a != key_b, "different keys must not collide on the same contextId" + assert key_a != "ctx-1", "raw client contextId must not be used verbatim" + assert key_a.endswith("-ctx-1"), "original contextId kept for correlation" + assert "hash-a" not in key_a, "raw principal must not be sent to LangFlow" + + +def test_merge_a2a_session_without_context_id_is_noop(): + merged = merge_a2a_session_into_litellm_params( + {"custom_llm_provider": "langflow", "model": "langflow/flow-1"}, + {"message": {"role": "user"}}, + ) + assert "session_id" not in merged + + +def test_langflow_a2a_provider_config_registered(): + cfg = A2AProviderConfigManager.get_provider_config( + custom_llm_provider="langflow", + model="langflow/flow-1", + ) + assert cfg is not None + assert cfg.__class__.__name__ == "LangFlowA2AConfig" + + +@pytest.mark.asyncio +async def test_langflow_a2a_config_passes_session_id_to_completion(): + from litellm.a2a_protocol.providers.langflow.config import LangFlowA2AConfig + + mock_response = type( + "R", + (), + { + "choices": [ + type( + "C", + (), + {"message": type("M", (), {"content": "ok"})()}, + )() + ] + }, + )() + + with patch("litellm.acompletion", new_callable=AsyncMock) as mock_acompletion: + mock_acompletion.return_value = mock_response + + await LangFlowA2AConfig().handle_non_streaming( + request_id="req-1", + params={ + "message": { + "role": "user", + "parts": [{"kind": "text", "text": "hi"}], + "contextId": "shared-session-99", + } + }, + litellm_params={ + "custom_llm_provider": "langflow", + "model": "langflow/flow-1", + "api_base": "http://localhost:7860", + }, + api_base="http://localhost:7860", + ) + + assert ( + mock_acompletion.call_args.kwargs.get("session_id") == "shared-session-99" + ) + + +@pytest.mark.asyncio +async def test_langflow_a2a_config_scopes_session_by_authenticated_key(): + from litellm.a2a_protocol.providers.langflow.config import LangFlowA2AConfig + + mock_response = type( + "R", + (), + {"choices": [type("C", (), {"message": type("M", (), {"content": "ok"})()})()]}, + )() + + with patch("litellm.acompletion", new_callable=AsyncMock) as mock_acompletion: + mock_acompletion.return_value = mock_response + + await LangFlowA2AConfig().handle_non_streaming( + request_id="req-1", + params={ + "message": { + "role": "user", + "parts": [{"kind": "text", "text": "hi"}], + "contextId": "ctx-1", + } + }, + litellm_params={ + "custom_llm_provider": "langflow", + "model": "langflow/flow-1", + "api_base": "http://localhost:7860", + A2A_USER_API_KEY_HASH_PARAM: "hashed-key-1", + }, + api_base="http://localhost:7860", + ) + + forwarded = mock_acompletion.call_args.kwargs + assert forwarded.get("session_id") != "ctx-1" + assert forwarded.get("session_id").endswith("-ctx-1") + assert ( + A2A_USER_API_KEY_HASH_PARAM not in forwarded + ), "internal principal param must not leak to the LLM call" + + +@pytest.mark.asyncio +async def test_langflow_a2a_config_requires_litellm_params_non_streaming(): + from litellm.a2a_protocol.providers.langflow.config import LangFlowA2AConfig + + with pytest.raises(ValueError, match="litellm_params is required"): + await LangFlowA2AConfig().handle_non_streaming( + request_id="req-1", + params={"message": {"contextId": "shared-session-99"}}, + ) + + +@pytest.mark.asyncio +async def test_langflow_a2a_config_requires_litellm_params_streaming(): + from litellm.a2a_protocol.providers.langflow.config import LangFlowA2AConfig + + with pytest.raises(ValueError, match="litellm_params is required"): + async for _ in LangFlowA2AConfig().handle_streaming( + request_id="req-1", + params={"message": {"contextId": "shared-session-99"}}, + ): + pass diff --git a/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py b/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py index 4b2e9471fb7..d389b54b3f1 100644 --- a/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py +++ b/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py @@ -247,6 +247,7 @@ class TestOpenAIResponsesAPIConfig: assert "Authorization" in result assert result["Authorization"] == f"Bearer {api_key}" + assert result["Content-Type"] == "application/json" # Test with empty headers headers = {} diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_tool_call_followed_by_text_assistant.py b/tests/test_litellm/llms/vertex_ai/gemini/test_tool_call_followed_by_text_assistant.py new file mode 100644 index 00000000000..bbd12e25f43 --- /dev/null +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_tool_call_followed_by_text_assistant.py @@ -0,0 +1,57 @@ +""" +Regression test for tool-call / tool-result matching in the Gemini message converter. + +When an assistant message that contains tool_calls is followed by a *second* assistant +message that has no tool_calls (e.g. the model emits a short narration turn after the +tool call but before the tool result), the converter used to overwrite its +`last_message_with_tool_calls` reference with the text-only assistant message. The +subsequent tool result could then no longer be matched to its tool call, and conversion +failed with: + + Exception: Missing corresponding tool call for tool response message. + +This happens for any OpenAI-style history with that shape, independent of provider/model. +""" + +import pytest + +from litellm.llms.vertex_ai.gemini.transformation import ( + _gemini_convert_messages_with_history, +) + + +def _messages_with_text_assistant_between_tool_call_and_result(): + return [ + {"role": "user", "content": "list the files"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_abc123", + "type": "function", + "function": {"name": "shell", "arguments": '{"command": ["ls"]}'}, + } + ], + }, + # text-only assistant message in between (no tool_calls) + {"role": "assistant", "content": "Running the command now."}, + {"role": "tool", "tool_call_id": "call_abc123", "content": "math.py"}, + ] + + +def test_tool_result_matches_tool_call_with_text_assistant_in_between(): + messages = _messages_with_text_assistant_between_tool_call_and_result() + + # Should not raise "Missing corresponding tool call for tool response message". + contents = _gemini_convert_messages_with_history(messages=messages) + + # The function response must be present and carry the correct tool name. + function_responses = [ + part["function_response"] + for content in contents + for part in content["parts"] + if isinstance(part, dict) and part.get("function_response") + ] + assert function_responses, f"expected a functionResponse part, got: {contents}" + assert function_responses[0]["name"] == "shell" diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py index 263fb1c6e65..628a6ed4cba 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py @@ -1154,44 +1154,82 @@ def test_convert_tool_response_with_base64_image(): ] } - # Convert tool response (returns list when image is present) + # Convert tool response with nested multimodal functionResponse.parts. result = convert_to_gemini_tool_call_result( tool_message, last_message_with_tool_calls ) - # Verify results - should be a list with 2 parts (function_response + inline_data) - assert isinstance( - result, list - ), f"Expected list when image present, got {type(result)}" - assert len(result) == 2, f"Expected 2 parts, got {len(result)}" - - # Find function_response part and inline_data part - function_response_part = None - inline_data_part = None - for part in result: - if "function_response" in part: - function_response_part = part - elif "inline_data" in part: - inline_data_part = part - - # Check function_response exists - assert function_response_part is not None, "Missing function_response part" - function_response = function_response_part["function_response"] + assert isinstance(result, list), "Should return a parts list when media is present" + assert len(result) == 1, "Should return one function_response part" + result_part = result[0] + assert "function_response" in result_part + assert "inline_data" not in result_part + function_response = result_part["function_response"] assert function_response["name"] == "click_at" assert "response" in function_response # Verify JSON response is parsed correctly assert "url" in function_response["response"] assert function_response["response"]["url"] == "https://example.com" - # Check inline_data exists - assert inline_data_part is not None, "Missing inline_data part" - inline_data: BlobType = inline_data_part["inline_data"] + # Check inline_data is nested under functionResponse.parts. + assert "parts" in function_response + assert len(function_response["parts"]) == 1 + inline_data: BlobType = function_response["parts"][0]["inline_data"] assert "data" in inline_data assert "mime_type" in inline_data assert inline_data["mime_type"] == "image/png" assert inline_data["data"] == test_image_base64 +def test_gemini_history_nests_multimodal_tool_response_parts(): + """Full history conversion should not emit sibling inline_data tool result parts.""" + test_image_base64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg==" + messages = [ + {"role": "user", "content": "Get me an image"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_get_image", + "type": "function", + "function": {"name": "get_image", "arguments": "{}"}, + } + ], + }, + { + "role": "tool", + "tool_call_id": "call_get_image", + "content": [ + {"type": "text", "text": '{"image_ref": "inline"}'}, + { + "type": "image", + "source": { + "type": "base64", + "media_type": "image/png", + "data": test_image_base64, + }, + }, + ], + }, + ] + + contents = _gemini_convert_messages_with_history(messages=messages) + + tool_response_parts = contents[-1]["parts"] + assert len(tool_response_parts) == 1 + assert "inline_data" not in tool_response_parts[0] + function_response = tool_response_parts[0]["function_response"] + assert function_response["parts"] == [ + { + "inline_data": { + "data": test_image_base64, + "mime_type": "image/png", + } + } + ] + + def test_convert_tool_response_with_url_image(): """Test tool response with HTTP URL image (will download and convert).""" import pytest @@ -1225,24 +1263,20 @@ def test_convert_tool_response_with_url_image(): tool_message, last_message_with_tool_calls ) - # Should be a list with 2 parts when image is present assert isinstance( result, list - ), f"Expected list when image present, got {type(result)}" - assert len(result) == 2, f"Expected 2 parts, got {len(result)}" - - # Find parts - function_response_part = next(p for p in result if "function_response" in p) - inline_data_part = next(p for p in result if "inline_data" in p) - - # Check function_response exists - assert function_response_part is not None, "Missing function_response part" - function_response = function_response_part["function_response"] + ), "Should return a parts list when media is present" + assert len(result) == 1, "Should return one function_response part" + result_part = result[0] + assert "function_response" in result_part + assert "inline_data" not in result_part + function_response = result_part["function_response"] assert function_response["name"] == "type_text_at" - # Check inline_data exists (URL should be downloaded and converted) - assert inline_data_part is not None, "Missing inline_data part" - inline_data: BlobType = inline_data_part["inline_data"] + # Check inline_data is nested under functionResponse.parts. + assert "parts" in function_response + assert len(function_response["parts"]) == 1 + inline_data: BlobType = function_response["parts"][0]["inline_data"] assert "data" in inline_data assert "mime_type" in inline_data except Exception as e: @@ -1558,38 +1592,27 @@ def test_convert_tool_response_with_pdf_file(): ] } - # Convert tool response (returns list when file is present) + # Convert tool response with nested multimodal functionResponse.parts. result = convert_to_gemini_tool_call_result( tool_message, last_message_with_tool_calls ) - # Verify results - should be a list with 2 parts (function_response + inline_data) - assert isinstance( - result, list - ), f"Expected list when file present, got {type(result)}" - assert len(result) == 2, f"Expected 2 parts, got {len(result)}" - - # Find function_response part and inline_data part - function_response_part = None - inline_data_part = None - for part in result: - if "function_response" in part: - function_response_part = part - elif "inline_data" in part: - inline_data_part = part - - # Check function_response exists - assert function_response_part is not None, "Missing function_response part" - function_response = function_response_part["function_response"] + assert isinstance(result, list), "Should return a parts list when media is present" + assert len(result) == 1, "Should return one function_response part" + result_part = result[0] + assert "function_response" in result_part + assert "inline_data" not in result_part + function_response = result_part["function_response"] assert function_response["name"] == "analyze_document" assert "response" in function_response # Verify JSON response is parsed correctly assert "status" in function_response["response"] assert function_response["response"]["status"] == "success" - # Check inline_data exists - assert inline_data_part is not None, "Missing inline_data part" - inline_data: BlobType = inline_data_part["inline_data"] + # Check inline_data is nested under functionResponse.parts. + assert "parts" in function_response + assert len(function_response["parts"]) == 1 + inline_data: BlobType = function_response["parts"][0]["inline_data"] assert "data" in inline_data assert "mime_type" in inline_data assert inline_data["mime_type"] == "application/pdf" @@ -1624,21 +1647,13 @@ def test_convert_tool_response_with_input_file_type(): tool_message, last_message_with_tool_calls ) - # Verify results - assert isinstance( - result, list - ), f"Expected list when file present, got {type(result)}" - assert len(result) == 2, f"Expected 2 parts, got {len(result)}" - - # Find inline_data part - inline_data_part = None - for part in result: - if "inline_data" in part: - inline_data_part = part - - # Check inline_data exists - assert inline_data_part is not None, "Missing inline_data part" - assert inline_data_part["inline_data"]["mime_type"] == "application/pdf" + # Check inline_data is nested under functionResponse.parts. + assert isinstance(result, list), "Should return a parts list when media is present" + assert len(result) == 1, "Should return one function_response part" + function_response = result[0]["function_response"] + assert ( + function_response["parts"][0]["inline_data"]["mime_type"] == "application/pdf" + ) def test_convert_tool_response_with_nested_file_object(): @@ -1669,21 +1684,11 @@ def test_convert_tool_response_with_nested_file_object(): tool_message, last_message_with_tool_calls ) - # Verify results - should be a list with 2 parts - assert isinstance( - result, list - ), f"Expected list when file present, got {type(result)}" - assert len(result) == 2, f"Expected 2 parts, got {len(result)}" - - # Find inline_data part - inline_data_part = None - for part in result: - if "inline_data" in part: - inline_data_part = part - - # Check inline_data exists - assert inline_data_part is not None, "Missing inline_data part" - inline_data: BlobType = inline_data_part["inline_data"] + # Check inline_data is nested under functionResponse.parts. + assert isinstance(result, list), "Should return a parts list when media is present" + assert len(result) == 1, "Should return one function_response part" + function_response = result[0]["function_response"] + inline_data: BlobType = function_response["parts"][0]["inline_data"] assert "data" in inline_data assert "mime_type" in inline_data assert inline_data["mime_type"] == "application/pdf" diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_search_vector_store_transformation.py b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_search_vector_store_transformation.py index 5ca71dc08c3..034f85f5a0b 100644 --- a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_search_vector_store_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_search_vector_store_transformation.py @@ -1,5 +1,8 @@ +from types import SimpleNamespace + import pytest +from litellm.exceptions import BadRequestError from litellm.llms.vertex_ai.vector_stores.search_api.transformation import ( VertexSearchAPIVectorStoreConfig, ) @@ -126,3 +129,171 @@ def test_should_raise_when_neither_engine_id_nor_vector_store_id_provided(): "vertex_location": "global", }, ) + + +_ENGINE_BASE = ( + "https://discoveryengine.googleapis.com/v1/projects/p/locations/global/" + "collections/default_collection/engines/app-2/servingConfigs/default_serving_config" +) + +_DATASTORE_BASE = ( + "https://discoveryengine.googleapis.com/v1/projects/p/locations/global/" + "collections/default_collection/dataStores/ds-1/servingConfigs/default_config" +) + + +def _search_request(**overrides): + """Engine/app-mode search request (vertex_engine_id set).""" + kwargs = dict( + vector_store_id="vs", + query="hello", + vector_store_search_optional_params={}, + api_base=_ENGINE_BASE, + litellm_logging_obj=SimpleNamespace(model_call_details={}), + litellm_params={"vertex_engine_id": "app-2"}, + ) + kwargs.update(overrides) + return VertexSearchAPIVectorStoreConfig().transform_search_vector_store_request( + **kwargs + ) + + +def _datastore_search_request(**overrides): + """Data-store-mode search request (no vertex_engine_id).""" + kwargs = dict( + vector_store_id="ds-1", + query="hello", + vector_store_search_optional_params={}, + api_base=_DATASTORE_BASE, + litellm_logging_obj=SimpleNamespace(model_call_details={}), + litellm_params={}, + ) + kwargs.update(overrides) + return VertexSearchAPIVectorStoreConfig().transform_search_vector_store_request( + **kwargs + ) + + +def test_search_request_defaults_to_query_and_pagesize_10(): + url, body = _search_request() + + assert url == _ENGINE_BASE + ":search" + assert body == {"query": "hello", "pageSize": 10} + + +def test_search_request_maps_max_num_results_to_pagesize(): + _, body = _search_request( + vector_store_search_optional_params={"max_num_results": 25} + ) + + assert body["pageSize"] == 25 + + +def test_engine_search_request_forwards_datastorespecs(): + specs = [ + { + "dataStore": "projects/p/locations/global/collections/default_collection/dataStores/ds-beta" + } + ] + + _, body = _search_request(extra_body={"dataStoreSpecs": specs}) + + assert body["dataStoreSpecs"] == specs + + +def test_engine_search_request_forwards_num_results_per_data_store(): + _, body = _search_request(extra_body={"numResultsPerDataStore": 3}) + + assert body["numResultsPerDataStore"] == 3 + + +def test_datastore_search_request_rejects_datastorespecs(): + specs = [{"dataStore": "projects/p/.../dataStores/ds-beta"}] + + with pytest.raises(BadRequestError, match="data store mode"): + _datastore_search_request(extra_body={"dataStoreSpecs": specs}) + + +def test_datastore_search_request_rejects_num_results_per_data_store(): + with pytest.raises(BadRequestError, match="data store mode"): + _datastore_search_request(extra_body={"numResultsPerDataStore": 3}) + + +@pytest.mark.parametrize("field", ["branch", "servingConfig", "entity"]) +def test_search_request_rejects_target_selecting_fields(field): + with pytest.raises(BadRequestError, match="target-selecting"): + _search_request(extra_body={field: "x"}) + + +@pytest.mark.parametrize("field", ["branch", "servingConfig", "entity"]) +def test_datastore_search_request_rejects_target_selecting_fields(field): + with pytest.raises(BadRequestError, match="target-selecting"): + _datastore_search_request(extra_body={field: "x"}) + + +def test_search_request_rejects_unsupported_extra_body_field(): + with pytest.raises(BadRequestError, match="Unsupported Vertex AI Search extra_body"): + _search_request(extra_body={"notARealField": True}) + + +def test_rejected_extra_body_raises_http_400(): + with pytest.raises(BadRequestError) as exc_info: + _search_request(extra_body={"notARealField": True}) + + assert exc_info.value.status_code == 400 + + +def test_search_request_forwards_supported_extra_body_fields(): + _, body = _search_request( + extra_body={ + "filter": 'category: ANY("docs")', + "boostSpec": {"conditionBoostSpecs": []}, + } + ) + + assert body["filter"] == 'category: ANY("docs")' + assert body["boostSpec"] == {"conditionBoostSpecs": []} + assert body["query"] == "hello" + + +def test_datastore_search_request_forwards_supported_extra_body_fields(): + _, body = _datastore_search_request( + extra_body={"filter": 'category: ANY("docs")'} + ) + + assert body["filter"] == 'category: ANY("docs")' + + +def test_search_request_ignores_none_valued_extra_body_fields(): + _, body = _search_request(extra_body={"filter": None}) + + assert "filter" not in body + + +def test_search_request_extra_body_takes_precedence_over_defaults(): + _, body = _search_request( + vector_store_search_optional_params={"max_num_results": 5}, + extra_body={"pageSize": 50, "filter": 'category: ANY("docs")'}, + ) + + assert body["pageSize"] == 50 + assert body["filter"] == 'category: ANY("docs")' + + +def test_search_request_joins_list_query(): + _, body = _search_request(query=["foo", "bar"]) + + assert body["query"] == "foo bar" + + +def test_search_request_logs_effective_query_when_extra_body_overrides_query(): + log = SimpleNamespace(model_call_details={}) + + _, body = _search_request( + query="original", + extra_body={"query": "from-extra-body"}, + litellm_logging_obj=log, + ) + + assert body["query"] == "from-extra-body" + assert log.model_call_details["query"] == "from-extra-body" diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_model_garden_openapi.py b/tests/test_litellm/llms/vertex_ai/test_vertex_model_garden_openapi.py index 91261b63252..0dcaa4c72c2 100644 --- a/tests/test_litellm/llms/vertex_ai/test_vertex_model_garden_openapi.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_model_garden_openapi.py @@ -1,7 +1,13 @@ """Vertex Model Garden: OpenAPI base URL for publisher/model ids vs per-endpoint path.""" +import json +import os +import sys +from unittest.mock import AsyncMock, MagicMock, patch + import pytest +import litellm from litellm.llms.vertex_ai.vertex_model_garden.main import ( _vertex_model_garden_model_id_in_json_body, create_vertex_url, @@ -37,5 +43,198 @@ def test_create_vertex_url_openapi_vs_deployed_endpoint( def test_model_id_in_json_body_heuristic() -> None: - assert _vertex_model_garden_model_id_in_json_body("xai/grok-4.1-fast-reasoning") is True + assert ( + _vertex_model_garden_model_id_in_json_body("xai/grok-4.1-fast-reasoning") + is True + ) assert _vertex_model_garden_model_id_in_json_body("5464397967697903616") is False + + +@pytest.fixture +def _reset_litellm_http_client_cache(): + from litellm import in_memory_llm_clients_cache + + in_memory_llm_clients_cache.flush_cache() + yield + in_memory_llm_clients_cache.flush_cache() + + +@pytest.fixture +def clean_vertex_env(): + saved_env = {} + env_vars_to_clear = [ + "GOOGLE_APPLICATION_CREDENTIALS", + "GOOGLE_CLOUD_PROJECT", + "VERTEXAI_PROJECT", + "VERTEXAI_LOCATION", + "VERTEXAI_CREDENTIALS", + "VERTEX_PROJECT", + "VERTEX_LOCATION", + "VERTEX_AI_PROJECT", + ] + for var in env_vars_to_clear: + if var in os.environ: + saved_env[var] = os.environ[var] + del os.environ[var] + + yield + + for var, value in saved_env.items(): + os.environ[var] = value + + +def _mock_chat_completion_response(model_in_response: str) -> MagicMock: + response = MagicMock() + response.status_code = 200 + response.headers = {} + response.json.return_value = { + "id": "chatcmpl-test", + "object": "chat.completion", + "created": 1234567890, + "model": model_in_response, + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "hi"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + } + return response + + +async def _invoke_model_garden_completion( + *, + model: str, + api_base, + mock_response: MagicMock, +): + """Drive litellm.acompletion through the Vertex Model Garden route and return + the patched AsyncHTTPHandler so callers can inspect the outbound HTTP call.""" + mock_vertexai = MagicMock() + mock_vertexai.preview = MagicMock() + mock_vertexai.preview.language_models = MagicMock() + + with ( + patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler" + ) as mock_http_handler, + patch( + "litellm.llms.vertex_ai.vertex_model_garden.main.VertexAIModelGardenModels._ensure_access_token", + return_value=("fake-token", "test-project"), + ), + patch.dict( + sys.modules, + {"vertexai": mock_vertexai, "vertexai.preview": mock_vertexai.preview}, + ), + ): + mock_http_handler.return_value.post = AsyncMock(return_value=mock_response) + + kwargs = dict( + model=model, + messages=[{"role": "user", "content": "hello"}], + vertex_ai_location="us-central1", + vertex_ai_project="test-project", + ) + if api_base is not None: + kwargs["api_base"] = api_base + + await litellm.acompletion(**kwargs) + + return mock_http_handler + + +@pytest.mark.asyncio +async def test_user_supplied_api_base_passes_through_unchanged( + clean_vertex_env, _reset_litellm_http_client_cache +): + """A user-supplied api_base must reach the OpenAI-like handler unchanged, + with only its own '/chat/completions' suffix appended.""" + user_api_base = "https://my-endpoint.example.com/v1" + mock_http_handler = await _invoke_model_garden_completion( + model="vertex_ai/openai/5464397967697903616", + api_base=user_api_base, + mock_response=_mock_chat_completion_response("5464397967697903616"), + ) + + mock_http_handler.return_value.post.assert_called_once() + call_args = mock_http_handler.return_value.post.call_args + called_url = call_args.kwargs.get("url") or call_args.args[0] + request_body = json.loads(call_args.kwargs["data"]) + + assert called_url == f"{user_api_base}/chat/completions" + assert ":" not in called_url.replace("https://", "") + assert "aiplatform.googleapis.com" not in called_url + assert request_body["model"] == "" + + +@pytest.mark.asyncio +async def test_user_supplied_api_base_passthrough_for_publisher_model( + clean_vertex_env, _reset_litellm_http_client_cache +): + """User-supplied api_base is forwarded unchanged for publisher/catalog + models too; the publisher model id stays in the JSON body.""" + user_api_base = "https://my-endpoint.example.com/v1" + mock_http_handler = await _invoke_model_garden_completion( + model="vertex_ai/openai/xai/grok-4.1-fast-reasoning", + api_base=user_api_base, + mock_response=_mock_chat_completion_response("xai/grok-4.1-fast-reasoning"), + ) + + mock_http_handler.return_value.post.assert_called_once() + call_args = mock_http_handler.return_value.post.call_args + called_url = call_args.kwargs.get("url") or call_args.args[0] + request_body = json.loads(call_args.kwargs["data"]) + + assert called_url == f"{user_api_base}/chat/completions" + assert "aiplatform.googleapis.com" not in called_url + assert request_body["model"] == "xai/grok-4.1-fast-reasoning" + + +@pytest.mark.asyncio +async def test_default_api_base_when_none_provided_single_segment( + clean_vertex_env, _reset_litellm_http_client_cache +): + """With no api_base, single-segment endpoint ids must hit the per-endpoint + Vertex URL and send an empty model field in the body.""" + mock_http_handler = await _invoke_model_garden_completion( + model="vertex_ai/openai/5464397967697903616", + api_base=None, + mock_response=_mock_chat_completion_response("5464397967697903616"), + ) + + mock_http_handler.return_value.post.assert_called_once() + call_args = mock_http_handler.return_value.post.call_args + called_url = call_args.kwargs.get("url") or call_args.args[0] + request_body = json.loads(call_args.kwargs["data"]) + + assert called_url == ( + "https://us-central1-aiplatform.googleapis.com/v1beta1/projects/" + "test-project/locations/us-central1/endpoints/5464397967697903616/chat/completions" + ) + assert request_body["model"] == "" + + +@pytest.mark.asyncio +async def test_default_api_base_when_none_provided_publisher_model( + clean_vertex_env, _reset_litellm_http_client_cache +): + """With no api_base, publisher/catalog models must hit the shared OpenAPI + URL and send the publisher model id in the body.""" + mock_http_handler = await _invoke_model_garden_completion( + model="vertex_ai/openai/xai/grok-4.1-fast-reasoning", + api_base=None, + mock_response=_mock_chat_completion_response("xai/grok-4.1-fast-reasoning"), + ) + + mock_http_handler.return_value.post.assert_called_once() + call_args = mock_http_handler.return_value.post.call_args + called_url = call_args.kwargs.get("url") or call_args.args[0] + request_body = json.loads(call_args.kwargs["data"]) + + assert called_url == ( + "https://us-central1-aiplatform.googleapis.com/v1/projects/" + "test-project/locations/us-central1/endpoints/openapi/chat/completions" + ) + assert request_body["model"] == "xai/grok-4.1-fast-reasoning" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 95c826daa8e..7753378ab4f 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -1,12 +1,9 @@ import json import os import sys -from unittest import mock -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, call as mock_call, patch -import orjson import pytest -from fastapi import FastAPI, Request from fastapi.testclient import TestClient sys.path.insert( @@ -19,7 +16,6 @@ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, ) from litellm.proxy._types import SpecialHeaders, UserAPIKeyAuth -from litellm.proxy.auth.user_api_key_auth import user_api_key_auth @pytest.mark.asyncio @@ -453,7 +449,7 @@ class TestMCPRequestHandler: with patch( "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", side_effect=mock_user_api_key_auth, - ) as mock_auth: + ): # Call the method ( auth_result, @@ -998,6 +994,284 @@ class TestMCPPublicRouteGuard: assert isinstance(auth_result, UserAPIKeyAuth) +@pytest.mark.asyncio +class TestMCPPassthroughColdStartAdmission: + @staticmethod + def _make_passthrough_server(): + server = MagicMock() + server.is_oauth_passthrough = True + return server + + async def test_cold_start_ignores_header_without_path_target(self): + from fastapi import HTTPException + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp", + "headers": [(b"x-mcp-servers", b"passthrough_server")], + } + + async def mock_user_api_key_auth_fails(api_key, request): + raise HTTPException(status_code=401, detail="Invalid API key") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_user_api_key_auth_fails, + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp._is_mcp_passthrough_cold_start" + ) as mock_cold_start, + ): + mock_mgr.get_mcp_server_by_name.return_value = ( + TestMCPPassthroughColdStartAdmission._make_passthrough_server() + ) + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(scope) + + assert exc_info.value.status_code == 401 + # Cold-start admission must not fire for the aggregate ``/mcp`` + # route — only path-targeted routes are eligible for OAuth + # discovery admission. + mock_cold_start.assert_not_called() + + async def test_cold_start_rejects_server_specific_authorization_header(self): + from fastapi import HTTPException + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/passthrough_server", + "headers": [ + ( + b"x-mcp-passthrough_server-authorization", + b"Bearer upstream-token", + ) + ], + } + + async def mock_user_api_key_auth_fails(api_key, request): + raise HTTPException(status_code=401, detail="Invalid API key") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_user_api_key_auth_fails, + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = ( + TestMCPPassthroughColdStartAdmission._make_passthrough_server() + ) + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(scope) + + assert exc_info.value.status_code == 401 + + async def test_cold_start_rejects_legacy_mcp_auth_header(self): + from fastapi import HTTPException + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/passthrough_server", + "headers": [(b"x-mcp-auth", b"Bearer upstream-token")], + } + + async def mock_user_api_key_auth_fails(api_key, request): + raise HTTPException(status_code=401, detail="Invalid API key") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_user_api_key_auth_fails, + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = ( + TestMCPPassthroughColdStartAdmission._make_passthrough_server() + ) + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(scope) + + assert exc_info.value.status_code == 401 + + async def test_cold_start_fails_closed_when_client_ip_hides_server(self): + from fastapi import HTTPException + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/passthrough_server", + "headers": [], + } + + async def mock_user_api_key_auth_fails(api_key, request): + raise HTTPException(status_code=401, detail="Invalid API key") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_user_api_key_auth_fails, + ), + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.IPAddressUtils.get_mcp_client_ip", + return_value="203.0.113.10", + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = None + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(scope) + + assert exc_info.value.status_code == 401 + mock_mgr.get_mcp_server_by_name.assert_any_call( + "passthrough_server", client_ip="203.0.113.10" + ) + + async def test_cold_start_propagates_non_401_http_error(self): + from fastapi import HTTPException + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/passthrough_server", + "headers": [], + } + + async def mock_user_api_key_auth_forbidden(api_key, request): + raise HTTPException(status_code=403, detail="Forbidden") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_user_api_key_auth_forbidden, + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = ( + TestMCPPassthroughColdStartAdmission._make_passthrough_server() + ) + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(scope) + + assert exc_info.value.status_code == 403 + + async def test_cold_start_propagates_non_auth_proxy_exception(self): + from litellm.proxy._types import ProxyException + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/passthrough_server", + "headers": [], + } + + async def mock_user_api_key_auth_server_error(api_key, request): + raise ProxyException( + message="Internal error", + type="server_error", + param=None, + code=500, + ) + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_user_api_key_auth_server_error, + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = ( + TestMCPPassthroughColdStartAdmission._make_passthrough_server() + ) + with pytest.raises(ProxyException): + await MCPRequestHandler.process_mcp_request(scope) + + async def test_cold_start_allows_401_for_path_passthrough_target(self): + from fastapi import HTTPException + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/passthrough_server", + "headers": [], + } + + async def mock_user_api_key_auth_fails(api_key, request): + raise HTTPException(status_code=401, detail="Invalid API key") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_user_api_key_auth_fails, + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = ( + TestMCPPassthroughColdStartAdmission._make_passthrough_server() + ) + auth_result, *_rest = await MCPRequestHandler.process_mcp_request(scope) + + assert isinstance(auth_result, UserAPIKeyAuth) + mock_mgr.get_mcp_server_by_name.assert_any_call( + "passthrough_server", client_ip="" + ) + + async def test_cold_start_allows_proxy_exception_401_for_path_target(self): + from litellm.proxy._types import ProxyException + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/passthrough_server", + "headers": [], + } + + async def mock_user_api_key_auth_fails(api_key, request): + raise ProxyException( + message="Authentication Error", + type="auth_error", + param="api_key", + code=401, + ) + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_user_api_key_auth_fails, + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = ( + TestMCPPassthroughColdStartAdmission._make_passthrough_server() + ) + auth_result, *_rest = await MCPRequestHandler.process_mcp_request(scope) + + assert isinstance(auth_result, UserAPIKeyAuth) + mock_mgr.get_mcp_server_by_name.assert_any_call( + "passthrough_server", client_ip="" + ) + + @pytest.mark.asyncio class TestMCPOAuth2FallbackTargetGating: """ @@ -1009,9 +1283,14 @@ class TestMCPOAuth2FallbackTargetGating: """ @staticmethod - def _make_server(auth_type): + def _make_server(auth_type, is_oauth_passthrough=False): server = MagicMock() server.auth_type = auth_type + # MagicMock would otherwise auto-create truthy stand-ins for any + # attribute access (including ``is_oauth_passthrough``), which + # would silently flip the passthrough fallback gate on. Pin the + # boolean explicitly so non-passthrough fixtures stay non-passthrough. + server.is_oauth_passthrough = is_oauth_passthrough return server async def test_fallback_blocked_when_target_is_not_oauth2(self): @@ -1113,6 +1392,88 @@ class TestMCPOAuth2FallbackTargetGating: auth_result, *_rest = await MCPRequestHandler.process_mcp_request(scope) assert isinstance(auth_result, UserAPIKeyAuth) + async def test_fallback_allowed_when_target_is_passthrough(self): + """ + Cold-start return per RFC 9728 / MCP Authorization spec: client + discovered the upstream IdP via the gateway's protected-resource + metadata, completed OAuth, and is returning with + ``Authorization: Bearer ``. The bearer is not a + LiteLLM key but the target is a pass-through server, so admission + falls back to anonymous and forwards the bearer upstream. + """ + from fastapi import HTTPException + + from litellm.types.mcp import MCPAuth + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/passthrough_server", + "headers": [(b"authorization", b"Bearer upstream-token-xyz")], + } + + async def mock_user_api_key_auth_fails(api_key, request): + raise HTTPException(status_code=401, detail="Invalid API key") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_user_api_key_auth_fails, + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = ( + TestMCPOAuth2FallbackTargetGating._make_server( + auth_type=MCPAuth.none, + is_oauth_passthrough=True, + ) + ) + auth_result, *_rest = await MCPRequestHandler.process_mcp_request(scope) + assert isinstance(auth_result, UserAPIKeyAuth) + assert auth_result.api_key is None + + async def test_fallback_blocked_when_client_ip_hides_oauth2_target(self): + from fastapi import HTTPException + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/hidden_oauth2_server", + "headers": [(b"authorization", b"Bearer upstream-token")], + } + + async def mock_user_api_key_auth_fails(api_key, request): + raise HTTPException(status_code=401, detail="Invalid API key") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_user_api_key_auth_fails, + ), + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.IPAddressUtils.get_mcp_client_ip", + return_value="203.0.113.10", + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = None + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(scope) + + assert exc_info.value.status_code == 401 + # Lookup may run twice — once for the oauth2-target fallback gate + # and once for the passthrough-target fallback gate. Both must + # resolve to ``None`` (hidden by client IP) so neither bypass + # opens. Use ``assert_any_call`` to assert the IP-scoped lookup + # happened without locking the count. + mock_mgr.get_mcp_server_by_name.assert_any_call( + "hidden_oauth2_server", client_ip="203.0.113.10" + ) + async def test_fallback_blocked_when_any_target_in_header_is_not_oauth2(self): """ x-mcp-servers can list multiple targets. If ANY of them is non-OAuth2, @@ -1241,6 +1602,39 @@ class TestMCPDelegateAuthToUpstream: is False ) + def test_build_mcp_server_table_preserves_oauth_passthrough(self): + """Registry → API list rows must expose oauth_passthrough for the UI. + + ``oauth_passthrough`` is the dedicated non-oauth2 pass-through opt-in, + distinct from ``delegate_auth_to_upstream`` (oauth2-only). Both must + round-trip independently so neither flag silently implies the other. + """ + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + MCPServerManager, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + manager = MCPServerManager() + passthrough = MCPServer( + server_id="passthrough-1", + name="passthrough", + transport="http", + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=True, + available_on_public_internet=True, + ) + row = manager._build_mcp_server_table(passthrough) + assert row.oauth_passthrough is True + # The oauth2-only flag must remain independent and default off. + assert row.delegate_auth_to_upstream is False + + not_passthrough = passthrough.model_copy(update={"oauth_passthrough": False}) + assert ( + manager._build_mcp_server_table(not_passthrough).oauth_passthrough is False + ) + async def test_delegate_skips_litellm_auth_with_no_authorization(self): """ oauth2 + delegate_auth_to_upstream=True, no Authorization header at @@ -1806,9 +2200,12 @@ class TestMCPDelegateAuthToUpstream: delegate_auth_to_upstream=True, ) - def lookup_by_name(name): + def lookup_by_name(name, **_kwargs): # Only the *exact* delegated name resolves. Anything else (e.g. # ``delegated_server/extra``) returns None so the bypass fails. + # ``**_kwargs`` accepts the ``client_ip`` kwarg the cold-start + # admission path now forwards (real signature: + # ``get_mcp_server_by_name(name, client_ip=None)``). if name == "delegated_server": return delegate_server return None @@ -1869,7 +2266,10 @@ class TestMCPDelegateAuthToUpstream: auth_type=MCPAuth.api_key, ) - def lookup_by_name(name): + def lookup_by_name(name, **_kwargs): + # ``**_kwargs`` accepts the ``client_ip`` kwarg the cold-start + # admission path now forwards (real signature: + # ``get_mcp_server_by_name(name, client_ip=None)``). return { "delegated_server": delegate_server, "non_delegate_server": non_delegate, @@ -2342,7 +2742,6 @@ class TestMCPAccessGroupsE2E: mock_auth.assert_called_once() -@pytest.mark.asyncio def test_mcp_path_based_server_segregation(monkeypatch): # Import the MCP server FastAPI app and context getter from litellm.proxy._experimental.mcp_server.server import app, get_auth_context @@ -2956,6 +3355,89 @@ class TestAgentMCPPermissions: ) assert sorted(result) == ["tool_a", "tool_b"] + async def test_get_agent_object_permission_uses_shared_helper(self): + """``_get_agent_object_permission`` must resolve the agent's + ``object_permission_id`` and then defer to the shared + ``get_object_permission`` helper so cache entries are shared with the + org / team / key paths.""" + from litellm.caching.dual_cache import DualCache + + cache = DualCache() + agent_row = MagicMock() + agent_row.object_permission_id = "perm-xyz" + prisma_client = MagicMock() + prisma_client.db.litellm_agentstable.find_unique = AsyncMock( + return_value=agent_row + ) + user_api_key_auth = UserAPIKeyAuth( + api_key="test-key", + user_id="test-user", + agent_id="agent-shared", + ) + expected_perm = MagicMock() + + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma_client), + patch("litellm.proxy.proxy_server.user_api_key_cache", cache), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + patch( + "litellm.proxy.auth.auth_checks.get_object_permission", + new_callable=AsyncMock, + return_value=expected_perm, + ) as mock_get_perm, + ): + result = await MCPRequestHandler._get_agent_object_permission( + user_api_key_auth + ) + assert result is expected_perm + mock_get_perm.assert_awaited_once() + assert mock_get_perm.await_args.kwargs["object_permission_id"] == "perm-xyz" + + # Second call: the agent_id -> object_permission_id mapping is + # cached, so the agent row is not re-fetched. + prisma_client.db.litellm_agentstable.find_unique.reset_mock() + await MCPRequestHandler._get_agent_object_permission(user_api_key_auth) + prisma_client.db.litellm_agentstable.find_unique.assert_not_called() + + async def test_get_agent_object_permission_caches_missing_permission(self): + """When the agent has no ``object_permission_id`` the sentinel must be + cached so subsequent requests do not hit the DB again.""" + from litellm.caching.dual_cache import DualCache + + cache = DualCache() + agent_row = MagicMock() + agent_row.object_permission_id = None + prisma_client = MagicMock() + prisma_client.db.litellm_agentstable.find_unique = AsyncMock( + return_value=agent_row + ) + user_api_key_auth = UserAPIKeyAuth( + api_key="test-key", + user_id="test-user", + agent_id="agent-no-perm", + ) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma_client), + patch("litellm.proxy.proxy_server.user_api_key_cache", cache), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + patch( + "litellm.proxy.auth.auth_checks.get_object_permission", + new_callable=AsyncMock, + ) as mock_get_perm, + ): + assert ( + await MCPRequestHandler._get_agent_object_permission(user_api_key_auth) + is None + ) + assert ( + await MCPRequestHandler._get_agent_object_permission(user_api_key_auth) + is None + ) + + mock_get_perm.assert_not_awaited() + prisma_client.db.litellm_agentstable.find_unique.assert_awaited_once() + @pytest.mark.asyncio async def test_tool_permission_servers_included_in_allowed_servers(): diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index c8789e0b0a6..da66d60aed8 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -1515,7 +1515,7 @@ async def test_oauth_protected_resource_returns_empty_scopes_when_none(): mock_request.headers = {} try: - response = _build_oauth_protected_resource_response( + response = await _build_oauth_protected_resource_response( request=mock_request, mcp_server_name="atlassian_mcp", use_standard_pattern=False, @@ -2005,7 +2005,7 @@ async def test_discovery_root_does_not_expose_private_server_for_external_client request=mock_request, mcp_server_name=None, ) - resource_response = _build_oauth_protected_resource_response( + resource_response = await _build_oauth_protected_resource_response( request=mock_request, mcp_server_name=None, use_standard_pattern=False, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py index cbea386a69c..04ff1e4be20 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py @@ -826,3 +826,87 @@ class TestUserAPIKeyAuthJwtClaims: auth.jwt_claims = claims assert auth.jwt_claims == claims assert auth.jwt_claims["groups"] == ["admin"] + + +class TestMcpRateLimitServerNameSurfacing: + """ + The per-MCP-server rate limiter only sees the request `data` dict, so the + server identity must be surfaced into it. These tests pin the contract + between pre_call_tool_check, _convert_mcp_to_llm_format, and the limiter. + """ + + def setup_method(self): + self.proxy_logging = ProxyLogging(user_api_key_cache=MagicMock()) + + def test_convert_mcp_to_llm_format_surfaces_rate_limit_server_name(self): + request_obj = MagicMock() + request_obj.tool_name = "list_repos" + request_obj.arguments = {"org": "acme"} + + result = self.proxy_logging._convert_mcp_to_llm_format( + request_obj, {"mcp_rate_limit_server_name": "github"} + ) + + assert result["mcp_server_name"] == "github" + + def test_convert_mcp_to_llm_format_server_name_none_when_absent(self): + request_obj = MagicMock() + request_obj.tool_name = "list_repos" + request_obj.arguments = {} + + result = self.proxy_logging._convert_mcp_to_llm_format(request_obj, {}) + + assert result["mcp_server_name"] is None + + @pytest.mark.asyncio + async def test_pre_call_tool_check_resolves_alias_for_rate_limit(self): + """ + The rate-limit server key must be the alias when set (falling back to + server_name), matching how an admin keys mcp_rpm_limit in config. + """ + manager = MCPServerManager() + server = MCPServer( + server_id="test-id", + name="gh", + alias="gh", + server_name="github_full_name", + url="https://example.com", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + ) + + captured = {} + + def capture_convert(request_obj, kwargs): + captured["kwargs"] = kwargs + return {"model": "fake"} + + proxy_logging = MagicMock(spec=ProxyLogging) + proxy_logging._create_mcp_request_object_from_kwargs = MagicMock( + return_value=MagicMock() + ) + proxy_logging._convert_mcp_to_llm_format = MagicMock( + side_effect=capture_convert + ) + proxy_logging.pre_call_hook = AsyncMock(return_value=None) + proxy_logging._convert_mcp_hook_response_to_kwargs = MagicMock( + return_value={"arguments": {}} + ) + + with patch.object(manager, "check_allowed_or_banned_tools", return_value=True): + with patch.object( + manager, + "check_tool_permission_for_key_team", + new_callable=AsyncMock, + ): + with patch.object(manager, "validate_allowed_params"): + await manager.pre_call_tool_check( + name="list_repos", + arguments={}, + server_name="github_full_name", + user_api_key_auth=None, + proxy_logging_obj=proxy_logging, + server=server, + ) + + assert captured["kwargs"]["mcp_rate_limit_server_name"] == "gh" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough.py new file mode 100644 index 00000000000..ad78609ee18 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough.py @@ -0,0 +1,474 @@ +"""Unit tests for MCP OAuth passthrough metadata behavior. + +Covers: +- `MCPServer.is_oauth_passthrough` property semantics. +- `/.well-known/oauth-protected-resource/...` pass-through branch (proxies + upstream metadata, normalizes the `resource` field, caches, and surfaces + network errors as HTTP 502). +""" + +import asyncio +import sys +import time +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest +from fastapi import HTTPException, Request + +sys.path.insert(0, "../../../../../") + + +from litellm.proxy._experimental.mcp_server import discoverable_endpoints +from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + _OAUTH_METADATA_CACHE, + _OAUTH_METADATA_FETCH_LOCKS, + _build_oauth_protected_resource_response, +) +from litellm.proxy._types import MCPTransport +from litellm.types.mcp import MCPAuth +from litellm.types.mcp_server.mcp_server_manager import MCPServer + + +@pytest.fixture(autouse=True) +def _mock_mcp_client_ip(): + """Bypass IP-based access control in tests.""" + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints" + ".IPAddressUtils.get_mcp_client_ip", + return_value=None, + ): + yield + + +@pytest.fixture(autouse=True) +def _clear_metadata_cache(): + """Prevent cross-test cache bleed for the oauth-protected-resource TTL cache.""" + _OAUTH_METADATA_CACHE.clear() + _OAUTH_METADATA_FETCH_LOCKS.clear() + yield + _OAUTH_METADATA_CACHE.clear() + _OAUTH_METADATA_FETCH_LOCKS.clear() + + +def _make_request(base_url: str = "https://gateway.example.com/") -> Request: + request = MagicMock(spec=Request) + request.base_url = base_url + request.headers = {} + return request + + +# -------------------------------------------------------------------------- +# is_oauth_passthrough property +# -------------------------------------------------------------------------- + + +def test_is_oauth_passthrough_true_when_none_auth_and_authorization_header(): + server = MCPServer( + server_id="s1", + name="s1", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=True, + ) + assert server.is_oauth_passthrough is True + + +def test_is_oauth_passthrough_true_when_auth_type_none_and_mixed_case_header(): + server = MCPServer( + server_id="s1", + name="s1", + transport=MCPTransport.http, + auth_type=None, + extra_headers=["authorization", "x-request-id"], + oauth_passthrough=True, + ) + assert server.is_oauth_passthrough is True + + +def test_is_oauth_passthrough_false_for_oauth2_server(): + server = MCPServer( + server_id="s1", + name="s1", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + extra_headers=["Authorization"], + oauth_passthrough=True, + ) + assert server.is_oauth_passthrough is False + + +def test_is_oauth_passthrough_false_without_authorization_header(): + server = MCPServer( + server_id="s1", + name="s1", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["x-api-key"], + oauth_passthrough=True, + ) + assert server.is_oauth_passthrough is False + + +def test_is_oauth_passthrough_false_without_extra_headers(): + server = MCPServer( + server_id="s1", + name="s1", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + oauth_passthrough=True, + ) + assert server.is_oauth_passthrough is False + + +def test_is_oauth_passthrough_false_without_oauth_passthrough_flag(): + """The detection flag must be set explicitly. Without it, the legacy + behavior is preserved for servers that forward Authorization for + non-OAuth reasons (static bearer tokens, custom auth schemes).""" + server = MCPServer( + server_id="s1", + name="s1", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + # oauth_passthrough defaults to False + ) + assert server.is_oauth_passthrough is False + + +def test_is_oauth_passthrough_false_when_oauth_passthrough_explicitly_false(): + server = MCPServer( + server_id="s1", + name="s1", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=False, + ) + assert server.is_oauth_passthrough is False + + +def test_is_oauth_passthrough_false_when_only_delegate_auth_to_upstream_set(): + """Regression guard: ``delegate_auth_to_upstream`` is the oauth2-only + PKCE-bypass flag and must NOT, on its own, turn a non-oauth2 server into + an OAuth pass-through server. Pass-through requires the dedicated + ``oauth_passthrough`` opt-in. This protects existing deployments that set + ``delegate_auth_to_upstream`` from silently gaining pass-through behavior. + """ + server = MCPServer( + server_id="s1", + name="s1", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + delegate_auth_to_upstream=True, + # oauth_passthrough intentionally left at its default (False) + ) + assert server.is_oauth_passthrough is False + + +# -------------------------------------------------------------------------- +# _build_oauth_protected_resource_response: pass-through branch +# -------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_oauth_protected_resource_passthrough_proxies_upstream_metadata(): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + global_mcp_server_manager.registry.clear() + passthrough_server = MCPServer( + server_id="passthrough-1", + name="sample_docs", + server_name="sample_docs", + alias="sample_docs", + url="https://upstream.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=True, + ) + global_mcp_server_manager.registry[passthrough_server.server_id] = ( + passthrough_server + ) + + upstream_payload = { + "resource": "https://upstream.example.com/mcp", + "authorization_servers": ["https://okta.example.com/oauth2/default"], + "scopes_supported": ["openid", "profile"], + "bearer_methods_supported": ["header"], + } + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.json.return_value = upstream_payload + mock_client = MagicMock() + mock_client.get = AsyncMock(return_value=mock_response) + + with patch.object( + discoverable_endpoints, "get_async_httpx_client", return_value=mock_client + ): + result = await _build_oauth_protected_resource_response( + request=_make_request(), + mcp_server_name="sample_docs", + use_standard_pattern=True, + ) + + assert result["authorization_servers"] == [ + "https://okta.example.com/oauth2/default" + ] + # resource is normalized to the gateway URL so bearers are sent back to us + assert result["resource"].endswith("/mcp/sample_docs") + assert result["scopes_supported"] == ["openid", "profile"] + + +@pytest.mark.asyncio +async def test_oauth_protected_resource_passthrough_cache_hit(): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + global_mcp_server_manager.registry.clear() + passthrough_server = MCPServer( + server_id="passthrough-2", + name="sample_docs", + server_name="sample_docs", + alias="sample_docs", + url="https://upstream.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=True, + ) + global_mcp_server_manager.registry[passthrough_server.server_id] = ( + passthrough_server + ) + + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.json.return_value = { + "authorization_servers": ["https://okta.example.com"], + } + mock_client = MagicMock() + mock_client.get = AsyncMock(return_value=mock_response) + + with patch.object( + discoverable_endpoints, "get_async_httpx_client", return_value=mock_client + ): + await _build_oauth_protected_resource_response( + request=_make_request(), + mcp_server_name="sample_docs", + use_standard_pattern=True, + ) + await _build_oauth_protected_resource_response( + request=_make_request(), + mcp_server_name="sample_docs", + use_standard_pattern=True, + ) + + assert mock_client.get.await_count == 1 + + +def test_oauth_metadata_cache_prunes_to_max_size(): + now = 1_000_000.0 + max_size = discoverable_endpoints._OAUTH_METADATA_CACHE_MAX_SIZE + + for index in range(max_size + 10): + _OAUTH_METADATA_CACHE[(f"server-{index}", f"https://upstream/{index}")] = ( + now + index + 1, + {"index": index}, + ) + + discoverable_endpoints._prune_oauth_metadata_cache(now) + + assert len(_OAUTH_METADATA_CACHE) == max_size + assert ("server-0", "https://upstream/0") not in _OAUTH_METADATA_CACHE + assert ( + f"server-{max_size + 9}", + f"https://upstream/{max_size + 9}", + ) in _OAUTH_METADATA_CACHE + + +def test_oauth_metadata_fetch_locks_pruned_alongside_cache(): + now = 1_000_000.0 + cached_key = ("server-active", "https://upstream/active") + expired_key = ("server-expired", "https://upstream/expired") + orphan_key = ("server-orphan", "https://upstream/orphan") + + _OAUTH_METADATA_CACHE[cached_key] = (now + 100, {"index": 0}) + _OAUTH_METADATA_CACHE[expired_key] = (now - 1, {"index": 1}) + + _OAUTH_METADATA_FETCH_LOCKS[cached_key] = asyncio.Lock() + _OAUTH_METADATA_FETCH_LOCKS[expired_key] = asyncio.Lock() + _OAUTH_METADATA_FETCH_LOCKS[orphan_key] = asyncio.Lock() + + discoverable_endpoints._prune_oauth_metadata_cache(now) + + assert cached_key in _OAUTH_METADATA_FETCH_LOCKS + assert expired_key not in _OAUTH_METADATA_FETCH_LOCKS + assert orphan_key not in _OAUTH_METADATA_FETCH_LOCKS + + +@pytest.mark.asyncio +async def test_oauth_metadata_fetch_locks_held_lock_preserved_during_prune(): + held_key = ("server-busy", "https://upstream/busy") + held_lock = asyncio.Lock() + _OAUTH_METADATA_FETCH_LOCKS[held_key] = held_lock + + async with held_lock: + discoverable_endpoints._prune_oauth_metadata_cache(time.time()) + assert held_key in _OAUTH_METADATA_FETCH_LOCKS + + +@pytest.mark.asyncio +async def test_oauth_metadata_cache_expired_entry_is_refetched(): + passthrough_server = MCPServer( + server_id="expired-cache-server", + name="sample_docs", + server_name="sample_docs", + alias="sample_docs", + url="https://upstream.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=True, + ) + _OAUTH_METADATA_CACHE[(passthrough_server.server_id, passthrough_server.url)] = ( + 0, + {"authorization_servers": ["https://stale.example.com"]}, + ) + + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.json.return_value = { + "authorization_servers": ["https://fresh.example.com"], + } + mock_client = MagicMock() + mock_client.get = AsyncMock(return_value=mock_response) + + with patch.object( + discoverable_endpoints, "get_async_httpx_client", return_value=mock_client + ): + result = await discoverable_endpoints.fetch_upstream_oauth_protected_resource( + passthrough_server + ) + + assert result == {"authorization_servers": ["https://fresh.example.com"]} + assert mock_client.get.await_count == 1 + + +@pytest.mark.asyncio +async def test_oauth_protected_resource_passthrough_network_error_returns_502(): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + global_mcp_server_manager.registry.clear() + passthrough_server = MCPServer( + server_id="passthrough-3", + name="sample_docs", + server_name="sample_docs", + alias="sample_docs", + url="https://upstream.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=True, + ) + global_mcp_server_manager.registry[passthrough_server.server_id] = ( + passthrough_server + ) + + mock_client = MagicMock() + mock_client.get = AsyncMock(side_effect=httpx.ConnectError("boom")) + + with patch.object( + discoverable_endpoints, "get_async_httpx_client", return_value=mock_client + ): + with pytest.raises(HTTPException) as exc_info: + await _build_oauth_protected_resource_response( + request=_make_request(), + mcp_server_name="sample_docs", + use_standard_pattern=True, + ) + + assert exc_info.value.status_code == 502 + + +@pytest.mark.asyncio +async def test_fetch_upstream_metadata_returns_none_when_not_all_candidates_network_fail(): + passthrough_server = MCPServer( + server_id="passthrough-partial-network", + name="sample_docs", + server_name="sample_docs", + alias="sample_docs", + url="https://upstream.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=True, + ) + + not_found_response = MagicMock() + not_found_response.status_code = 404 + mock_client = MagicMock() + mock_client.get = AsyncMock( + side_effect=[not_found_response, httpx.ConnectError("path fallback failed")] + ) + + with patch.object( + discoverable_endpoints, "get_async_httpx_client", return_value=mock_client + ): + result = await discoverable_endpoints.fetch_upstream_oauth_protected_resource( + passthrough_server + ) + + assert result is None + assert mock_client.get.await_count == 2 + + +@pytest.mark.asyncio +async def test_oauth_protected_resource_gateway_managed_unchanged(): + """Regression guard: OAuth2 servers still advertise the gateway as AS.""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + global_mcp_server_manager.registry.clear() + oauth2_server = MCPServer( + server_id="oauth2-1", + name="keycloak_whoami", + server_name="keycloak_whoami", + alias="keycloak_whoami", + url="https://upstream.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="cid", + client_secret="cs", + authorization_url="https://keycloak/auth", + token_url="https://keycloak/token", + scopes=["read"], + ) + global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server + + # If the code mistakenly fetched upstream metadata for a gateway-managed + # server, this spy would catch it. + mock_client = MagicMock() + mock_client.get = AsyncMock() + + with patch.object( + discoverable_endpoints, "get_async_httpx_client", return_value=mock_client + ): + result = await _build_oauth_protected_resource_response( + request=_make_request(), + mcp_server_name="keycloak_whoami", + use_standard_pattern=True, + ) + + mock_client.get.assert_not_awaited() + assert result["authorization_servers"] == [ + "https://gateway.example.com/keycloak_whoami" + ] + assert result["scopes_supported"] == ["read"] diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_cold_start.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_cold_start.py new file mode 100644 index 00000000000..3e934577a66 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_cold_start.py @@ -0,0 +1,156 @@ +"""Unit tests for MCP OAuth passthrough cold-start route behavior.""" + +import sys + +import pytest + +sys.path.insert(0, "../../../../../") + +from litellm.proxy._types import MCPTransport +from litellm.types.mcp import MCPAuth +from litellm.types.mcp_server.mcp_server_manager import MCPServer + + +def _make_scope(path: str, headers: list = None) -> dict: + """Build a minimal ASGI HTTP scope for testing.""" + raw_headers = [(key.encode(), value.encode()) for key, value in (headers or [])] + return { + "type": "http", + "method": "POST", + "path": path, + "headers": raw_headers, + "query_string": b"", + "server": ("localhost", 4000), + "scheme": "http", + } + + +@pytest.mark.parametrize( + "route,expected_metadata_path", + [ + ( + "/mcp/sample_docs", + "/.well-known/oauth-protected-resource/mcp/sample_docs", + ), + ( + "/sample_docs/mcp", + "/.well-known/oauth-protected-resource/sample_docs/mcp", + ), + ], +) +def test_passthrough_cold_start_emits_401_with_matching_resource_metadata( + route, expected_metadata_path +): + """No auth headers on a passthrough server route emits matching metadata.""" + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + _is_mcp_passthrough_cold_start, + _parse_mcp_server_names_from_path, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + global_mcp_server_manager.registry.clear() + passthrough_server = MCPServer( + server_id="pt-cold-start", + name="sample_docs", + server_name="sample_docs", + alias="sample_docs", + url="https://upstream.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=True, + ) + global_mcp_server_manager.registry[passthrough_server.server_id] = ( + passthrough_server + ) + + if route.startswith("/mcp/"): + scope = _make_scope(route) + else: + scope = _make_scope("/mcp/sample_docs") + scope["_original_path"] = route + + servers = _parse_mcp_server_names_from_path(scope.get("path", "")) + assert _is_mcp_passthrough_cold_start(servers, client_ip=None) is True + + server_name = "sample_docs" + base_url = "http://localhost:4000" + path = scope.get("_original_path") or scope.get("path", "") or "" + if path.startswith(f"/{server_name}/mcp"): + resource_metadata_url = ( + f"{base_url}/.well-known/oauth-protected-resource/{server_name}/mcp" + ) + else: + resource_metadata_url = ( + f"{base_url}/.well-known/oauth-protected-resource/mcp/{server_name}" + ) + + assert resource_metadata_url == f"{base_url}{expected_metadata_path}", ( + f"resource_metadata_url {resource_metadata_url!r} does not match " + f"expected {base_url + expected_metadata_path!r}" + ) + + +def test_is_mcp_passthrough_cold_start_false_for_oauth2_server(): + """Gateway-managed OAuth2 servers must not trigger the cold-start bypass.""" + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + _is_mcp_passthrough_cold_start, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + global_mcp_server_manager.registry.clear() + oauth2_server = MCPServer( + server_id="oauth2-cold", + name="keycloak_whoami", + server_name="keycloak_whoami", + alias="keycloak_whoami", + url="https://upstream.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="cid", + client_secret="cs", + authorization_url="https://keycloak/auth", + token_url="https://keycloak/token", + ) + global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server + + result = _is_mcp_passthrough_cold_start(["keycloak_whoami"], client_ip=None) + assert result is False + + +def test_is_mcp_passthrough_cold_start_false_for_empty_servers(): + """Aggregate /mcp route (no server list) must not trigger bypass.""" + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + _is_mcp_passthrough_cold_start, + ) + + assert _is_mcp_passthrough_cold_start(None, client_ip=None) is False + assert _is_mcp_passthrough_cold_start([], client_ip=None) is False + + +@pytest.mark.parametrize( + "path,expected", + [ + ("/mcp/sample_docs", ["sample_docs"]), + # Server names may contain at most one slash (mirrors + # ``_extract_target_server_names_from_path``), so when more than two + # segments follow ``/mcp/`` the first two are treated as the name. + ("/mcp/sample_docs/tools/list", ["sample_docs/tools"]), + ("/mcp/custom_solutions/user_123", ["custom_solutions/user_123"]), + ("/sample_docs/mcp", ["sample_docs"]), + ("/sample_docs/mcp/tools/list", ["sample_docs"]), + ("/mcp", None), + ("/mcp/", None), + ("/other/path", None), + ], +) +def test_parse_mcp_server_names_from_path(path, expected): + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + _parse_mcp_server_names_from_path, + ) + + assert _parse_mcp_server_names_from_path(path) == expected diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py new file mode 100644 index 00000000000..d900f690c57 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py @@ -0,0 +1,197 @@ +"""Unit tests for MCP OAuth passthrough tool-fetch behavior.""" + +import sys +from unittest.mock import AsyncMock, MagicMock + +import httpx +import pytest + +sys.path.insert(0, "../../../../../") + +from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError +from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + MCPServerManager, + _extract_upstream_auth_failure, +) +from litellm.proxy._types import MCPTransport +from litellm.types.mcp import MCPAuth +from litellm.types.mcp_server.mcp_server_manager import MCPServer + + +def test_extract_upstream_auth_failure_finds_401_in_http_status_error(): + response = httpx.Response( + status_code=401, + headers={"www-authenticate": 'Bearer resource_metadata="https://x"'}, + request=httpx.Request("GET", "https://upstream/mcp"), + ) + exc = httpx.HTTPStatusError("401", request=response.request, response=response) + + result = _extract_upstream_auth_failure(exc) + assert result == (401, 'Bearer resource_metadata="https://x"') + + +def test_extract_upstream_auth_failure_walks_exception_group(): + response = httpx.Response( + status_code=401, + headers={"www-authenticate": "Bearer"}, + request=httpx.Request("GET", "https://upstream/mcp"), + ) + inner = httpx.HTTPStatusError("401", request=response.request, response=response) + + try: + raise ExceptionGroup("wrapped", [inner]) # noqa: F821 (PEP 654, py3.11+) + except Exception as group: + result = _extract_upstream_auth_failure(group) + + assert result == (401, "Bearer") + + +def test_extract_upstream_auth_failure_returns_none_for_non_auth(): + assert _extract_upstream_auth_failure(RuntimeError("boom")) is None + + +@pytest.mark.asyncio +async def test_fetch_tools_from_passthrough_raises_on_upstream_401(): + manager = MCPServerManager() + passthrough_server = MCPServer( + server_id="p1", + name="sample_docs", + url="https://upstream/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=True, + ) + + response = httpx.Response( + status_code=401, + headers={"www-authenticate": 'Bearer resource_metadata="https://upstream"'}, + request=httpx.Request("GET", "https://upstream/mcp"), + ) + upstream_error = httpx.HTTPStatusError( + "401", request=response.request, response=response + ) + + mock_client = MagicMock() + mock_client.list_tools = AsyncMock(side_effect=upstream_error) + + with pytest.raises(MCPUpstreamAuthError) as exc_info: + await manager._fetch_tools_with_timeout( + mock_client, passthrough_server.name, server=passthrough_server + ) + + assert exc_info.value.status_code == 401 + assert exc_info.value.www_authenticate == ( + 'Bearer resource_metadata="https://upstream"' + ) + assert exc_info.value.server_name == "sample_docs" + mock_client.list_tools.assert_awaited_with(raise_on_error=True) + + +@pytest.mark.asyncio +async def test_fetch_tools_from_passthrough_returns_tools_on_success(): + manager = MCPServerManager() + passthrough_server = MCPServer( + server_id="p1", + name="sample_docs", + url="https://upstream/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=True, + ) + + tool = MagicMock() + tool.name = "list_documents" + mock_client = MagicMock() + mock_client.list_tools = AsyncMock(return_value=[tool]) + + tools = await manager._fetch_tools_with_timeout( + mock_client, passthrough_server.name, server=passthrough_server + ) + assert tools == [tool] + + +def test_to_http_exception_preserves_upstream_www_authenticate(): + err = MCPUpstreamAuthError( + status_code=401, + www_authenticate='Bearer resource_metadata="https://upstream/.well-known/oauth-protected-resource"', + server_name="sample_docs", + ) + + http_exc = err.to_http_exception() + assert http_exc.status_code == 401 + assert http_exc.headers == { + "www-authenticate": 'Bearer resource_metadata="https://upstream/.well-known/oauth-protected-resource"' + } + + +def test_to_http_exception_skips_fabrication_when_base_url_missing(): + """Without ``base_url`` we cannot build an RFC 9728 §3.2-compliant absolute + URI, so we omit the fabricated ``WWW-Authenticate`` challenge entirely + instead of emitting a relative URI strict clients reject.""" + err = MCPUpstreamAuthError( + status_code=401, + www_authenticate=None, + server_name="sample_docs", + ) + + http_exc = err.to_http_exception() + assert http_exc.status_code == 401 + assert http_exc.headers is None + + +def test_to_http_exception_fabricates_absolute_resource_metadata_with_base_url(): + err = MCPUpstreamAuthError( + status_code=401, + www_authenticate=None, + server_name="sample_docs", + ) + + http_exc = err.to_http_exception(base_url="https://gateway.example.com/") + assert http_exc.status_code == 401 + assert http_exc.headers == { + "www-authenticate": 'Bearer resource_metadata="https://gateway.example.com/.well-known/oauth-protected-resource/mcp/sample_docs"' + } + + +def test_to_http_exception_skips_challenge_for_non_401_status(): + err = MCPUpstreamAuthError( + status_code=403, + www_authenticate=None, + server_name="sample_docs", + ) + + http_exc = err.to_http_exception() + assert http_exc.status_code == 403 + assert http_exc.headers is None + + +@pytest.mark.asyncio +async def test_fetch_tools_from_gateway_managed_swallows_errors(): + """Regression guard: non-pass-through servers keep returning [] on errors.""" + manager = MCPServerManager() + oauth2_server = MCPServer( + server_id="o1", + name="keycloak_whoami", + url="https://upstream/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + ) + + response = httpx.Response( + status_code=401, + headers={}, + request=httpx.Request("GET", "https://upstream/mcp"), + ) + upstream_error = httpx.HTTPStatusError( + "401", request=response.request, response=response + ) + mock_client = MagicMock() + mock_client.list_tools = AsyncMock(side_effect=upstream_error) + + tools = await manager._fetch_tools_with_timeout( + mock_client, oauth2_server.name, server=oauth2_server + ) + assert tools == [] + mock_client.list_tools.assert_awaited_with(raise_on_error=False) 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 bb0cc860375..6b6c7bc37d5 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 @@ -131,13 +131,129 @@ def test_prepare_mcp_server_headers_case_insensitive_extra_headers(): mcp_server_auth_headers=None, mcp_auth_header=None, oauth2_headers=None, - raw_headers={"authorization": "Bearer token"}, + raw_headers={ + "x-litellm-api-key": "Bearer sk-litellm-key", + "authorization": "Bearer token", + }, ) assert server_auth_header is None assert extra_headers == {"Authorization": "Bearer token"} +def test_prepare_mcp_server_headers_passthrough_strips_authorization_without_admission_header(): + try: + from litellm.proxy._experimental.mcp_server.server import ( + _prepare_mcp_server_headers, + ) + except ImportError: + pytest.skip("MCP server not available") + + server = MCPServer( + server_id="server-passthrough-no-admission", + name="server", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization", "x-request-id"], + oauth_passthrough=True, + ) + + server_auth_header, extra_headers = _prepare_mcp_server_headers( + server=server, + mcp_server_auth_headers=None, + mcp_auth_header=None, + oauth2_headers=None, + raw_headers={ + "authorization": "Bearer sk-litellm-key", + "x-request-id": "req-789", + }, + ) + + assert server_auth_header is None + assert extra_headers == {"x-request-id": "req-789"} + + +def test_prepare_mcp_server_headers_passthrough_forwards_authorization_for_anonymous_admission(): + """Cold-start return per RFC 9728: client admits anonymously through + the pass-through fallback in :meth:`MCPRequestHandler.process_mcp_request` + (``user_api_key_auth.api_key is None``) and the ``Authorization`` bearer + is the upstream OAuth token — it must be forwarded, not stripped.""" + try: + from litellm.proxy._experimental.mcp_server.server import ( + _prepare_mcp_server_headers, + ) + except ImportError: + pytest.skip("MCP server not available") + + from litellm.proxy._types import UserAPIKeyAuth + + server = MCPServer( + server_id="server-passthrough-anon-admission", + name="server", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization", "x-request-id"], + oauth_passthrough=True, + ) + + server_auth_header, extra_headers = _prepare_mcp_server_headers( + server=server, + mcp_server_auth_headers=None, + mcp_auth_header=None, + oauth2_headers=None, + raw_headers={ + "authorization": "Bearer upstream-oauth-token", + "x-request-id": "req-790", + }, + user_api_key_auth=UserAPIKeyAuth(), + ) + + assert server_auth_header is None + assert extra_headers == { + "Authorization": "Bearer upstream-oauth-token", + "x-request-id": "req-790", + } + + +def test_prepare_mcp_server_headers_passthrough_strips_authorization_for_authenticated_admission(): + """When admission validated ``Authorization`` as a LiteLLM key + (``user_api_key_auth.api_key`` is set, no explicit ``x-litellm-api-key``), + the bearer must still be stripped to avoid leaking the gateway key + upstream.""" + try: + from litellm.proxy._experimental.mcp_server.server import ( + _prepare_mcp_server_headers, + ) + except ImportError: + pytest.skip("MCP server not available") + + from litellm.proxy._types import UserAPIKeyAuth + + server = MCPServer( + server_id="server-passthrough-authenticated", + name="server", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization", "x-request-id"], + oauth_passthrough=True, + ) + + server_auth_header, extra_headers = _prepare_mcp_server_headers( + server=server, + mcp_server_auth_headers=None, + mcp_auth_header=None, + oauth2_headers=None, + raw_headers={ + "authorization": "Bearer sk-litellm-key", + "x-request-id": "req-791", + }, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-key"), + ) + + assert server_auth_header is None + assert extra_headers == {"x-request-id": "req-791"} + + def test_prepare_mcp_server_headers_oauth2_m2m_omits_litellm_caller_authorization(): """M2M OAuth must not put caller Bearer (LiteLLM API key) into extra_headers (#23652).""" try: @@ -514,6 +630,7 @@ async def test_mcp_get_prompt_success(): mcp_auth_header=None, oauth2_headers=None, raw_headers=None, + user_api_key_auth=user_api_key_auth, ) mock_manager.get_prompt_from_server.assert_awaited_once_with( server=server, @@ -575,6 +692,7 @@ async def test_mcp_read_resource_success(): mcp_auth_header=None, oauth2_headers=None, raw_headers=None, + user_api_key_auth=user_api_key_auth, ) mock_manager.read_resource_from_server.assert_awaited_once_with( server=server, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 2db9845c765..ec690aef629 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -480,6 +480,167 @@ class TestMCPServerManager: assert captured_extra_headers == {"Authorization": "Bearer token"} assert isinstance(result, CallToolResult) + @pytest.mark.asyncio + async def test_call_regular_mcp_tool_passthrough_strips_authorization_when_admission_consumed_litellm_key( + self, + ): + """OAuth pass-through must not forward the caller's Authorization to upstream + when LiteLLM admission consumed the bearer as its API key — otherwise the + LiteLLM key the caller used for admission would leak upstream.""" + from litellm.proxy._types import UserAPIKeyAuth + + manager = MCPServerManager() + server = MCPServer( + server_id="server-passthrough-call", + name="passthrough-server", + url="https://example.com", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization", "x-request-id"], + oauth_passthrough=True, + ) + + mock_client = AsyncMock() + mock_client.call_tool = AsyncMock( + return_value=CallToolResult(content=[], isError=False) + ) + captured_extra_headers = None + + async def capture_create_mcp_client( + server, mcp_auth_header, extra_headers, stdio_env, subject_token=None + ): # pragma: no cover - helper + nonlocal captured_extra_headers + captured_extra_headers = extra_headers + return mock_client + + manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) + + await manager._call_regular_mcp_tool( + mcp_server=server, + original_tool_name="tool", + arguments={}, + tasks=[], + mcp_auth_header=None, + mcp_server_auth_headers=None, + oauth2_headers=None, + raw_headers={ + "authorization": "Bearer sk-litellm-key", + "x-request-id": "req-123", + }, + proxy_logging_obj=None, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-key"), + ) + + assert captured_extra_headers == {"x-request-id": "req-123"} + + @pytest.mark.asyncio + async def test_call_regular_mcp_tool_passthrough_forwards_authorization_with_admission_header( + self, + ): + """OAuth pass-through forwards Authorization upstream when x-litellm-api-key + provides admission — in that case Authorization carries the upstream OAuth + bearer, not the LiteLLM key.""" + from litellm.proxy._types import UserAPIKeyAuth + + manager = MCPServerManager() + server = MCPServer( + server_id="server-passthrough-call-admission", + name="passthrough-server", + url="https://example.com", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=True, + ) + + mock_client = AsyncMock() + mock_client.call_tool = AsyncMock( + return_value=CallToolResult(content=[], isError=False) + ) + captured_extra_headers = None + + async def capture_create_mcp_client( + server, mcp_auth_header, extra_headers, stdio_env, subject_token=None + ): # pragma: no cover - helper + nonlocal captured_extra_headers + captured_extra_headers = extra_headers + return mock_client + + manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) + + await manager._call_regular_mcp_tool( + mcp_server=server, + original_tool_name="tool", + arguments={}, + tasks=[], + mcp_auth_header=None, + mcp_server_auth_headers=None, + oauth2_headers=None, + raw_headers={ + "x-litellm-api-key": "Bearer sk-litellm-key", + "authorization": "Bearer upstream-oauth-bearer", + }, + proxy_logging_obj=None, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-key"), + ) + + assert captured_extra_headers == { + "Authorization": "Bearer upstream-oauth-bearer" + } + + @pytest.mark.asyncio + async def test_call_regular_mcp_tool_passthrough_forwards_authorization_for_anonymous_admission( + self, + ): + """OAuth pass-through cold-start return (RFC 9728): the caller's only + credential is the upstream bearer in Authorization, and LiteLLM admission + is anonymous (no api_key on user_api_key_auth). Authorization must be + forwarded so the delegated flow can complete.""" + from litellm.proxy._types import UserAPIKeyAuth + + manager = MCPServerManager() + server = MCPServer( + server_id="server-passthrough-call-anon", + name="passthrough-server", + url="https://example.com", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=True, + ) + + mock_client = AsyncMock() + mock_client.call_tool = AsyncMock( + return_value=CallToolResult(content=[], isError=False) + ) + captured_extra_headers = None + + async def capture_create_mcp_client( + server, mcp_auth_header, extra_headers, stdio_env, subject_token=None + ): # pragma: no cover - helper + nonlocal captured_extra_headers + captured_extra_headers = extra_headers + return mock_client + + manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) + + await manager._call_regular_mcp_tool( + mcp_server=server, + original_tool_name="tool", + arguments={}, + tasks=[], + mcp_auth_header=None, + mcp_server_auth_headers=None, + oauth2_headers=None, + raw_headers={"authorization": "Bearer upstream-oauth-bearer"}, + proxy_logging_obj=None, + user_api_key_auth=UserAPIKeyAuth(api_key=None), + ) + + assert captured_extra_headers == { + "Authorization": "Bearer upstream-oauth-bearer" + } + @pytest.mark.asyncio async def test_get_prompts_from_server_success(self): """Ensure prompts are fetched and prefixed when requested.""" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_session_logging.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_session_logging.py new file mode 100644 index 00000000000..790937cc1de --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_session_logging.py @@ -0,0 +1,19 @@ +"""The MCP ``mcp-session-id`` is captured for tool-call logging so the otel span +can carry ``mcp.session.id``. Guards the header read against casing and absence.""" + +from litellm.proxy._experimental.mcp_server.server import _mcp_session_id_from_headers + + +def test_reads_session_id_case_insensitively(): + # Clients send varied casing (``Mcp-Session-Id``, ``mcp-session-id``); all resolve. + assert _mcp_session_id_from_headers({"mcp-session-id": "s1"}) == "s1" + assert _mcp_session_id_from_headers({"Mcp-Session-Id": "s2"}) == "s2" + assert _mcp_session_id_from_headers({"MCP-SESSION-ID": "s3"}) == "s3" + + +def test_stateless_call_has_no_session_id(): + # No header (stateless request) and an empty value both yield None, not "". + assert _mcp_session_id_from_headers({"authorization": "Bearer x"}) is None + assert _mcp_session_id_from_headers({"mcp-session-id": ""}) is None + assert _mcp_session_id_from_headers(None) is None + assert _mcp_session_id_from_headers({}) is None diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index 593facd9279..6433e0f6360 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -544,6 +544,78 @@ class TestListToolsRestAPI: assert result["error"] is None assert result["message"] == "Successfully retrieved tools" + @pytest.mark.parametrize("upstream_status", [401, 403]) + async def test_upstream_auth_failure_surfaces_status_and_challenge( + self, monkeypatch, upstream_status + ): + """A single-server pass-through request whose upstream rejects the token + must surface the upstream status (401 or 403) plus its WWW-Authenticate + challenge, not collapse into a 200 ``unexpected_error`` body.""" + from litellm.proxy._experimental.mcp_server.exceptions import ( + MCPUpstreamAuthError, + ) + + class StubServer: + alias = "server-1" + server_name = "server-1" + name = "passthrough" + allowed_tools = None + mcp_info = {"server_name": "passthrough"} + available_on_public_internet = True + + stub_server = StubServer() + + async def fake_contexts(user_api_key_auth): + return [user_api_key_auth] + + async def fake_get_allowed_mcp_servers(*args, **kwargs): + return ["server-1"] + + challenge = 'Bearer resource_metadata="https://upstream/.well-known"' + + async def fake_get_tools(*args, **kwargs): + raise MCPUpstreamAuthError( + status_code=upstream_status, + www_authenticate=challenge, + server_name="passthrough", + ) + + monkeypatch.setattr( + rest_endpoints, + "build_effective_auth_contexts", + fake_contexts, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_allowed_mcp_servers", + fake_get_allowed_mcp_servers, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_mcp_server_by_id", + lambda server_id: stub_server if server_id == "server-1" else None, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints, + "_get_tools_for_single_server", + fake_get_tools, + raising=False, + ) + + request = _build_request(path="/mcp-rest/tools/list", method="GET") + with pytest.raises(HTTPException) as exc_info: + await rest_endpoints.list_tool_rest_api( + request, + server_id="server-1", + user_api_key_dict=UserAPIKeyAuth(), + ) + + assert exc_info.value.status_code == upstream_status + assert exc_info.value.headers == {"www-authenticate": challenge} + async def test_name_resolution_finds_server_by_uuid(self, monkeypatch): """When server_id is a name string, it should be resolved to its UUID and used for the tools lookup when the UUID is in allowed_server_ids.""" diff --git a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py index 268e6d2dc13..07e878401e0 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py @@ -5,7 +5,9 @@ Tests that invoke_agent_a2a properly integrates with add_litellm_data_to_request """ import json +import socket import sys +from contextlib import ExitStack from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -246,3 +248,1315 @@ async def test_invoke_agent_a2a_handles_none_agent_card_params(): assert body["jsonrpc"] == "2.0" assert body["error"]["code"] == -32000 assert "no URL configured" in body["error"]["message"] + + +@pytest.mark.asyncio +async def test_invoke_agent_a2a_injects_authenticated_key_hash_for_bridge(): + """Completion-bridge agents must receive the authenticated key hash in + litellm_params so provider configs (e.g. LangFlow) can scope provider-side + session memory per key. Regression for cross-key A2A session bleed.""" + from litellm.a2a_protocol.litellm_completion_bridge.handler import ( + A2A_USER_API_KEY_HASH_PARAM, + ) + from litellm.proxy._types import UserAPIKeyAuth + + captured = {} + + async def mock_add_litellm_data(data, **kwargs): + data["proxy_server_request"] = { + "url": "http://localhost:4000/a2a/lf-agent", + "method": "POST", + "headers": {}, + "body": {}, + } + data.setdefault("metadata", {}) + return data + + async def capture_asend_message(**kwargs): + captured.update(kwargs) + resp = MagicMock() + resp.model_dump.return_value = {"jsonrpc": "2.0", "id": "test-id", "result": {}} + return resp + + mock_agent = MagicMock() + mock_agent.agent_id = "lf-agent" + mock_agent.agent_name = "lf-agent" + # No URL: the bridge derives the endpoint from the LangFlow agent config. + mock_agent.agent_card_params = {"name": "LF Agent"} + mock_agent.litellm_params = { + "custom_llm_provider": "langflow", + "model": "langflow/flow-1", + } + mock_agent.static_headers = None + mock_agent.extra_headers = None + + mock_request = MagicMock() + mock_request.headers = {} + mock_request.json = AsyncMock( + return_value={ + "jsonrpc": "2.0", + "id": "test-id", + "method": "message/send", + "params": { + "message": { + "role": "user", + "parts": [{"kind": "text", "text": "hi"}], + "messageId": "msg-1", + "contextId": "ctx-1", + } + }, + } + ) + + mock_user_api_key_dict = UserAPIKeyAuth( + api_key="sk-hashed-123", + user_id="test-user", + team_id="test-team", + ) + + with ( + patch( + "litellm.proxy.agent_endpoints.a2a_endpoints._get_agent", + return_value=mock_agent, + ), + patch( + "litellm.proxy.common_request_processing.add_litellm_data_to_request", + side_effect=mock_add_litellm_data, + ), + patch( + "litellm.proxy.agent_endpoints.auth.agent_permission_handler.AgentRequestHandler.is_agent_allowed", + new=AsyncMock(return_value=True), + ), + patch( + "litellm.a2a_protocol.asend_message", + new=AsyncMock(side_effect=capture_asend_message), + ), + patch("litellm.proxy.proxy_server.general_settings", {}), + patch("litellm.proxy.proxy_server.proxy_config", MagicMock()), + patch("litellm.proxy.proxy_server.version", "1.0.0"), + patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True), + patch.dict(sys.modules, {"a2a": MagicMock(), "a2a.types": MagicMock()}), + ): + from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a + + await invoke_agent_a2a( + agent_id="lf-agent", + request=mock_request, + fastapi_response=MagicMock(), + user_api_key_dict=mock_user_api_key_dict, + ) + + assert ( + captured.get("litellm_params", {}).get(A2A_USER_API_KEY_HASH_PARAM) + == mock_user_api_key_dict.api_key + ), "authenticated key hash was not forwarded to the completion bridge" + + +def _make_agent_mock(url: str = "http://backend-agent:10001") -> MagicMock: + agent = MagicMock() + agent.agent_id = "test-agent" + agent.agent_name = "test-agent" + agent.agent_card_params = {"url": url, "name": "Test Agent"} + agent.litellm_params = {} + agent.static_headers = None + agent.extra_headers = None + return agent + + +def _make_request_mock( + method: str, params: dict, request_id: object = "req-1" +) -> MagicMock: + req = MagicMock() + req.headers = {} + req.json = AsyncMock( + return_value={ + "jsonrpc": "2.0", + "id": request_id, + "method": method, + "params": params, + } + ) + return req + + +def _base_patches(agent: MagicMock): + return [ + patch( + "litellm.proxy.agent_endpoints.a2a_endpoints._get_agent", + return_value=agent, + ), + patch( + "litellm.proxy.agent_endpoints.auth.agent_permission_handler.AgentRequestHandler.is_agent_allowed", + new=AsyncMock(return_value=True), + ), + patch( + "litellm.proxy.common_request_processing.add_litellm_data_to_request", + new=AsyncMock(side_effect=_add_proxy_data), + ), + patch("litellm.proxy.proxy_server.general_settings", {}), + patch("litellm.proxy.proxy_server.proxy_config", MagicMock()), + patch("litellm.proxy.proxy_server.version", "1.0.0"), + ] + + +async def _add_proxy_data(data, **kwargs): + data["proxy_server_request"] = { + "url": "http://localhost:4000", + "method": "POST", + "headers": {}, + "body": {}, + } + data.setdefault("metadata", {}) + return data + + +@pytest.mark.asyncio +@pytest.mark.parametrize("method", ["message/send", "message/stream"]) +async def test_message_methods_preserve_numeric_zero_request_id(method: str): + from fastapi.responses import JSONResponse + from litellm.proxy._types import UserAPIKeyAuth + + class MessageSendParams: + def __init__(self, **kwargs): + self.__dict__.update(kwargs) + + class SendMessageRequest: + def __init__(self, **kwargs): + self.__dict__.update(kwargs) + + agent = _make_agent_mock() + params = { + "message": { + "role": "user", + "parts": [{"kind": "text", "text": "Hello"}], + "messageId": "msg-123", + } + } + mock_request = _make_request_mock(method, params, request_id=0) + user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1") + captured = {} + + async def capture_asend_message(request, **kwargs): + captured["request_id"] = request.id + response = MagicMock() + response.model_dump.return_value = { + "jsonrpc": "2.0", + "id": request.id, + "result": {"status": "success"}, + } + return response + + async def capture_stream_message(**kwargs): + captured["request_id"] = kwargs["request_id"] + return JSONResponse({"jsonrpc": "2.0", "id": kwargs["request_id"]}) + + mock_a2a_types = MagicMock() + mock_a2a_types.MessageSendParams = MessageSendParams + mock_a2a_types.SendMessageRequest = SendMessageRequest + + with ExitStack() as stack: + for p in _base_patches(agent): + stack.enter_context(p) + stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True)) + if method == "message/send": + stack.enter_context( + patch.dict( + sys.modules, + {"a2a": MagicMock(), "a2a.types": mock_a2a_types}, + ) + ) + stack.enter_context( + patch( + "litellm.a2a_protocol.asend_message", + new=AsyncMock(side_effect=capture_asend_message), + ) + ) + else: + stack.enter_context( + patch( + "litellm.proxy.agent_endpoints.a2a_endpoints._handle_stream_message", + new=AsyncMock(side_effect=capture_stream_message), + ) + ) + + from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a + + await invoke_agent_a2a( + agent_id="test-agent", + request=mock_request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + assert captured["request_id"] == 0 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "method,params", + [ + ("tasks/get", {"id": "task-1"}), + ("tasks/list", {"contextId": "ctx-1"}), + ("tasks/cancel", {"id": "task-1"}), + ( + "tasks/pushNotificationConfig/set", + {"taskId": "task-1", "url": "https://webhook.example.com"}, + ), + ("tasks/pushNotificationConfig/get", {"taskId": "task-1", "id": "cfg-1"}), + ("tasks/pushNotificationConfig/list", {"taskId": "task-1"}), + ("tasks/pushNotificationConfig/delete", {"taskId": "task-1", "id": "cfg-1"}), + ], +) +async def test_task_methods_forward_jsonrpc(method: str, params: dict): + from litellm.proxy._types import UserAPIKeyAuth + + upstream_response = { + "jsonrpc": "2.0", + "id": "req-1", + "result": {"id": "task-1", "status": {"state": "completed"}}, + } + agent = _make_agent_mock() + mock_request = _make_request_mock(method, params) + + mock_http_response = MagicMock() + mock_http_response.json.return_value = upstream_response + mock_http_response.is_success = True + mock_http_response.raise_for_status = MagicMock() + + mock_handler = MagicMock() + mock_handler.post = AsyncMock(return_value=mock_http_response) + mock_handler.client = MagicMock() + + user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1") + + with ExitStack() as stack: + for p in _base_patches(agent): + stack.enter_context(p) + stack.enter_context( + patch( + "litellm.llms.custom_httpx.http_handler.get_async_httpx_client", + return_value=mock_handler, + ) + ) + stack.enter_context( + patch( + "litellm.proxy.agent_endpoints.a2a_endpoints.validate_url", + return_value=("https://webhook.example.com", "webhook.example.com"), + ) + ) + + from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a + + response = await invoke_agent_a2a( + agent_id="test-agent", + request=mock_request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + body = json.loads(response.body.decode()) + assert body["jsonrpc"] == "2.0" + assert body["result"]["id"] == "task-1" + + posted = mock_handler.post.call_args + assert posted is not None + forwarded_body = posted.kwargs.get("json") or posted.args[1] + assert forwarded_body["method"] == method + + +@pytest.mark.asyncio +@pytest.mark.parametrize("method", ["tasks/get", "tasks/resubscribe"]) +async def test_task_methods_extract_litellm_params_before_forwarding(method: str): + from litellm.proxy._types import UserAPIKeyAuth + + agent = _make_agent_mock() + params = { + "id": "task-1", + "guardrails": ["guardrail-1"], + "tags": ["tag-1"], + } + mock_request = _make_request_mock(method, params) + user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1") + captured_data = {} + + async def capture_proxy_data(data, **kwargs): + captured_data.update(data) + return await _add_proxy_data(data, **kwargs) + + upstream_response = { + "jsonrpc": "2.0", + "id": "req-1", + "result": {"id": "task-1", "status": {"state": "completed"}}, + } + mock_http_response = MagicMock() + mock_http_response.json.return_value = upstream_response + mock_http_response.is_success = True + + async def fake_aiter_lines(): + yield 'data: {"jsonrpc":"2.0","id":"req-1","result":{"taskId":"task-1"}}' + + mock_resp = AsyncMock() + mock_resp.is_success = True + mock_resp.aiter_lines = fake_aiter_lines + mock_resp.aclose = AsyncMock() + + mock_async_client = MagicMock() + mock_async_client.build_request = MagicMock(return_value=MagicMock()) + mock_async_client.send = AsyncMock(return_value=mock_resp) + + mock_handler = MagicMock() + mock_handler.post = AsyncMock(return_value=mock_http_response) + mock_handler.client = mock_async_client + + with ExitStack() as stack: + for p in _base_patches(agent): + stack.enter_context(p) + stack.enter_context( + patch( + "litellm.proxy.common_request_processing.add_litellm_data_to_request", + new=AsyncMock(side_effect=capture_proxy_data), + ) + ) + stack.enter_context( + patch( + "litellm.llms.custom_httpx.http_handler.get_async_httpx_client", + return_value=mock_handler, + ) + ) + + from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a + + response = await invoke_agent_a2a( + agent_id="test-agent", + request=mock_request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + if method == "tasks/resubscribe": + async for _ in response.body_iterator: + pass + + if method == "tasks/resubscribe": + forwarded_body = mock_async_client.build_request.call_args.kwargs["json"] + else: + forwarded_body = mock_handler.post.call_args.kwargs["json"] + assert forwarded_body["params"] == {"id": "task-1"} + assert captured_data["guardrails"] == ["guardrail-1"] + assert captured_data["tags"] == ["tag-1"] + + +@pytest.mark.asyncio +async def test_subscribe_to_task_returns_sse_stream(): + from litellm.proxy._types import UserAPIKeyAuth + + agent = _make_agent_mock() + mock_request = _make_request_mock("SubscribeToTask", {"id": "task-1"}) + user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1") + + sse_lines = [ + 'data: {"jsonrpc":"2.0","id":"req-1","result":{"taskId":"task-1","status":{"state":"working"}}}', + 'data: {"jsonrpc":"2.0","id":"req-1","result":{"taskId":"task-1","status":{"state":"completed"}}}', + ] + + async def fake_aiter_lines(): + for line in sse_lines: + yield line + + mock_resp = AsyncMock() + mock_resp.is_success = True + mock_resp.aiter_lines = fake_aiter_lines + mock_resp.aclose = AsyncMock() + + mock_async_client = MagicMock() + mock_async_client.build_request = MagicMock(return_value=MagicMock()) + mock_async_client.send = AsyncMock(return_value=mock_resp) + + mock_handler = MagicMock() + mock_handler.client = mock_async_client + mock_handler.post = AsyncMock() + + chunks = [] + with ExitStack() as stack: + for p in _base_patches(agent): + stack.enter_context(p) + stack.enter_context( + patch( + "litellm.llms.custom_httpx.http_handler.get_async_httpx_client", + return_value=mock_handler, + ) + ) + + from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a + + response = await invoke_agent_a2a( + agent_id="test-agent", + request=mock_request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + assert response.media_type == "text/event-stream" + async for chunk in response.body_iterator: + chunks.append(chunk) + + full = "".join(chunks) + assert "working" in full + assert "completed" in full + + +@pytest.mark.asyncio +async def test_subscribe_to_task_calls_pre_call_hook(): + """tasks/resubscribe must run pre_call_hook so guardrails configured on + the agent are enforced before streaming begins.""" + from litellm.proxy._types import UserAPIKeyAuth + + agent = _make_agent_mock() + mock_request = _make_request_mock("tasks/resubscribe", {"id": "task-1"}) + user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1") + + async def fake_aiter_lines(): + yield 'data: {"jsonrpc":"2.0","id":"req-1","result":{"taskId":"task-1","status":{"state":"completed"}}}' + + mock_resp = AsyncMock() + mock_resp.is_success = True + mock_resp.aiter_lines = fake_aiter_lines + mock_resp.aclose = AsyncMock() + + mock_async_client = MagicMock() + mock_async_client.build_request = MagicMock(return_value=MagicMock()) + mock_async_client.send = AsyncMock(return_value=mock_resp) + + mock_handler = MagicMock() + mock_handler.client = mock_async_client + mock_handler.post = AsyncMock() + + async def _passthrough_iterator(response, **kwargs): + async for chunk in response: + yield chunk + + mock_proxy_logging = MagicMock() + mock_proxy_logging.pre_call_hook = AsyncMock( + side_effect=lambda user_api_key_dict, data, call_type: data + ) + mock_proxy_logging.async_post_call_streaming_iterator_hook = _passthrough_iterator + mock_proxy_logging.post_call_failure_hook = AsyncMock(return_value=None) + + with ExitStack() as stack: + for p in _base_patches(agent): + stack.enter_context(p) + stack.enter_context( + patch( + "litellm.llms.custom_httpx.http_handler.get_async_httpx_client", + return_value=mock_handler, + ) + ) + stack.enter_context( + patch( + "litellm.proxy.proxy_server.proxy_logging_obj", + mock_proxy_logging, + ) + ) + + from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a + + response = await invoke_agent_a2a( + agent_id="test-agent", + request=mock_request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + assert response.media_type == "text/event-stream" + async for _ in response.body_iterator: + pass + + mock_proxy_logging.pre_call_hook.assert_awaited_once() + call_kwargs = mock_proxy_logging.pre_call_hook.await_args.kwargs + assert call_kwargs.get("call_type") == "asend_message" + assert call_kwargs.get("user_api_key_dict") == user_api_key_dict + + +@pytest.mark.asyncio +async def test_subscribe_to_task_runs_post_call_streaming_guardrail(): + """tasks/resubscribe must route streamed events through the post-call + streaming hook so output guardrails configured on the agent inspect the + streamed task content. Regression: the SSE path previously returned the raw + upstream stream and bypassed guardrails entirely.""" + import litellm + from litellm.integrations.custom_guardrail import CustomGuardrail + from litellm.proxy._types import UserAPIKeyAuth + + inspected: list = [] + + class _RecordingGuardrail(CustomGuardrail): + async def async_post_call_streaming_hook(self, user_api_key_dict, response): + inspected.append(response) + return response + + guardrail = _RecordingGuardrail( + guardrail_name="record-a2a", default_on=True, event_hook="post_call" + ) + + agent = _make_agent_mock() + mock_request = _make_request_mock("tasks/resubscribe", {"id": "task-1"}) + user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1") + + async def fake_aiter_lines(): + yield ( + 'data: {"jsonrpc":"2.0","id":"req-1","result":' + '{"kind":"message","parts":[{"kind":"text","text":"resubscribe-secret"}]}}' + ) + + mock_resp = AsyncMock() + mock_resp.is_success = True + mock_resp.aiter_lines = fake_aiter_lines + mock_resp.aclose = AsyncMock() + + mock_async_client = MagicMock() + mock_async_client.build_request = MagicMock(return_value=MagicMock()) + mock_async_client.send = AsyncMock(return_value=mock_resp) + + mock_handler = MagicMock() + mock_handler.client = mock_async_client + mock_handler.post = AsyncMock() + + with ExitStack() as stack: + for p in _base_patches(agent): + stack.enter_context(p) + stack.enter_context( + patch( + "litellm.llms.custom_httpx.http_handler.get_async_httpx_client", + return_value=mock_handler, + ) + ) + stack.enter_context(patch.object(litellm, "callbacks", [guardrail])) + + from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a + + response = await invoke_agent_a2a( + agent_id="test-agent", + request=mock_request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + assert response.media_type == "text/event-stream" + async for _ in response.body_iterator: + pass + + assert any("resubscribe-secret" in str(r) for r in inspected), ( + "tasks/resubscribe streamed content was not passed to the post-call " + "streaming guardrail hook" + ) + + +@pytest.mark.asyncio +async def test_task_method_failure_hook_uses_enriched_request_data(): + from litellm.proxy._types import UserAPIKeyAuth + + agent = _make_agent_mock() + mock_request = _make_request_mock("tasks/get", {"id": "task-1"}) + user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1") + + async def add_proxy_data_copy(data, **kwargs): + enriched = dict(data) + enriched["proxy_server_request"] = { + "url": "http://localhost:4000", + "method": "POST", + "headers": {}, + "body": {}, + } + enriched.setdefault("metadata", {}) + return enriched + + mock_handler = MagicMock() + mock_handler.post = AsyncMock(side_effect=RuntimeError("upstream failed")) + + mock_proxy_logging = MagicMock() + mock_proxy_logging.pre_call_hook = AsyncMock( + side_effect=lambda user_api_key_dict, data, call_type: data + ) + mock_proxy_logging.post_call_failure_hook = AsyncMock(return_value=None) + + with ExitStack() as stack: + for p in _base_patches(agent): + stack.enter_context(p) + stack.enter_context( + patch( + "litellm.proxy.common_request_processing.add_litellm_data_to_request", + new=AsyncMock(side_effect=add_proxy_data_copy), + ) + ) + stack.enter_context( + patch( + "litellm.llms.custom_httpx.http_handler.get_async_httpx_client", + return_value=mock_handler, + ) + ) + stack.enter_context( + patch( + "litellm.proxy.proxy_server.proxy_logging_obj", + mock_proxy_logging, + ) + ) + + from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a + + response = await invoke_agent_a2a( + agent_id="test-agent", + request=mock_request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + body = json.loads(response.body.decode()) + assert body["error"]["code"] == -32603 + failure_data = mock_proxy_logging.post_call_failure_hook.await_args.kwargs[ + "request_data" + ] + assert failure_data.get("litellm_call_id") + assert failure_data.get("agent_id") == "test-agent" + + +@pytest.mark.asyncio +async def test_get_extended_agent_card_rewrites_url(): + from litellm.proxy._types import UserAPIKeyAuth + + agent = _make_agent_mock() + mock_request = _make_request_mock("GetExtendedAgentCard", {}) + mock_request.base_url = "http://localhost:4000/" + user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1") + + upstream_card = { + "name": "Test Agent", + "url": "http://backend-agent:10001", + "description": "A test agent", + } + upstream_response = {"jsonrpc": "2.0", "id": "req-1", "result": upstream_card} + + mock_http_response = MagicMock() + mock_http_response.json.return_value = upstream_response + mock_http_response.is_success = True + mock_http_response.raise_for_status = MagicMock() + + mock_handler = MagicMock() + mock_handler.post = AsyncMock(return_value=mock_http_response) + mock_handler.client = MagicMock() + + with ExitStack() as stack: + for p in _base_patches(agent): + stack.enter_context(p) + stack.enter_context( + patch( + "litellm.llms.custom_httpx.http_handler.get_async_httpx_client", + return_value=mock_handler, + ) + ) + + from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a + + response = await invoke_agent_a2a( + agent_id="test-agent", + request=mock_request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + body = json.loads(response.body.decode()) + assert body["result"]["url"] == "http://localhost:4000/a2a/test-agent" + assert body["result"]["name"] == "Test Agent" + + +@pytest.mark.asyncio +async def test_unknown_method_returns_jsonrpc_error(): + from litellm.proxy._types import UserAPIKeyAuth + + agent = _make_agent_mock() + mock_request = _make_request_mock("SomeUnknownMethod", {}) + user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1") + + with ExitStack() as stack: + for p in _base_patches(agent): + stack.enter_context(p) + + from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a + + response = await invoke_agent_a2a( + agent_id="test-agent", + request=mock_request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + body = json.loads(response.body.decode()) + assert body["error"]["code"] == -32601 + assert "SomeUnknownMethod" in body["error"]["message"] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "pascal_method,expected_wire_method", + [ + ("GetTask", "tasks/get"), + ("ListTasks", "tasks/list"), + ("CancelTask", "tasks/cancel"), + ("SubscribeToTask", "tasks/resubscribe"), + ("CreateTaskPushNotificationConfig", "tasks/pushNotificationConfig/set"), + ("GetTaskPushNotificationConfig", "tasks/pushNotificationConfig/get"), + ("ListTaskPushNotificationConfigs", "tasks/pushNotificationConfig/list"), + ("DeleteTaskPushNotificationConfig", "tasks/pushNotificationConfig/delete"), + ("GetExtendedAgentCard", "agent/getAuthenticatedExtendedCard"), + ], +) +async def test_pascal_method_names_normalize_to_wire_format( + pascal_method: str, expected_wire_method: str +): + from litellm.proxy._types import UserAPIKeyAuth + + agent = _make_agent_mock() + mock_request = _make_request_mock(pascal_method, {"id": "task-1"}) + user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1") + + upstream_response = {"jsonrpc": "2.0", "id": "req-1", "result": {"id": "task-1"}} + mock_http_response = MagicMock() + mock_http_response.json.return_value = upstream_response + mock_http_response.is_success = True + mock_http_response.raise_for_status = MagicMock() + + async def _empty_aiter_lines(): + return + yield # make it an async generator + + mock_sse_resp = AsyncMock() + mock_sse_resp.is_success = True + mock_sse_resp.aiter_lines = _empty_aiter_lines + mock_sse_resp.aclose = AsyncMock() + + mock_async_client = MagicMock() + mock_async_client.build_request = MagicMock(return_value=MagicMock()) + mock_async_client.send = AsyncMock(return_value=mock_sse_resp) + + mock_handler = MagicMock() + mock_handler.post = AsyncMock(return_value=mock_http_response) + mock_handler.client = mock_async_client + + with ExitStack() as stack: + for p in _base_patches(agent): + stack.enter_context(p) + stack.enter_context( + patch( + "litellm.llms.custom_httpx.http_handler.get_async_httpx_client", + return_value=mock_handler, + ) + ) + + from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a + + response = await invoke_agent_a2a( + agent_id="test-agent", + request=mock_request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + if expected_wire_method == "tasks/resubscribe": + assert response.media_type == "text/event-stream" + async for _ in response.body_iterator: + pass + else: + body = json.loads(response.body.decode()) + assert "error" not in body, f"Got error: {body}" + + if expected_wire_method != "tasks/resubscribe": + posted = mock_handler.post.call_args + forwarded_body = posted.kwargs.get("json") or posted.args[1] + assert forwarded_body["method"] == expected_wire_method, ( + f"Expected '{expected_wire_method}' forwarded for PascalCase '{pascal_method}', " + f"but got '{forwarded_body['method']}'" + ) + + +@pytest.mark.asyncio +async def test_task_method_upstream_jsonrpc_error_on_http_4xx_is_relayed(): + """When upstream returns HTTP 4xx with a JSON-RPC error body, the error body + must be relayed to the client unchanged, not replaced with a generic string.""" + from litellm.proxy._types import UserAPIKeyAuth + + agent = _make_agent_mock() + mock_request = _make_request_mock("tasks/get", {"id": "nonexistent"}) + user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1") + + upstream_error = { + "jsonrpc": "2.0", + "id": "req-1", + "error": {"code": -32001, "message": "Task not found"}, + } + + mock_http_response = MagicMock() + mock_http_response.json.return_value = upstream_error + mock_http_response.is_success = False + mock_http_response.raise_for_status = MagicMock( + side_effect=Exception("404 Not Found") + ) + + mock_handler = MagicMock() + mock_handler.post = AsyncMock(return_value=mock_http_response) + mock_handler.client = MagicMock() + + with ExitStack() as stack: + for p in _base_patches(agent): + stack.enter_context(p) + stack.enter_context( + patch( + "litellm.llms.custom_httpx.http_handler.get_async_httpx_client", + return_value=mock_handler, + ) + ) + + from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a + + response = await invoke_agent_a2a( + agent_id="test-agent", + request=mock_request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + body = json.loads(response.body.decode()) + assert body["error"]["code"] == -32001 + assert body["error"]["message"] == "Task not found" + + +@pytest.mark.asyncio +async def test_subscribe_to_task_upstream_error_yields_jsonrpc_error_event(): + """When upstream returns a non-2xx response for tasks/resubscribe, the SSE + stream must yield a JSON-RPC error event instead of silently breaking.""" + from litellm.proxy._types import UserAPIKeyAuth + + agent = _make_agent_mock() + mock_request = _make_request_mock("tasks/resubscribe", {"id": "task-1"}) + user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1") + + mock_resp = AsyncMock() + mock_resp.is_success = False + mock_resp.status_code = 404 + mock_resp.reason_phrase = "Not Found" + mock_resp.aread = AsyncMock( + return_value=b'{"jsonrpc":"2.0","error":{"code":-32001,"message":"Task not found"}}' + ) + mock_resp.aclose = AsyncMock() + + mock_async_client = MagicMock() + mock_async_client.build_request = MagicMock(return_value=MagicMock()) + mock_async_client.send = AsyncMock(return_value=mock_resp) + + mock_handler = MagicMock() + mock_handler.client = mock_async_client + mock_handler.post = AsyncMock() + + chunks = [] + with ExitStack() as stack: + for p in _base_patches(agent): + stack.enter_context(p) + stack.enter_context( + patch( + "litellm.llms.custom_httpx.http_handler.get_async_httpx_client", + return_value=mock_handler, + ) + ) + + from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a + + response = await invoke_agent_a2a( + agent_id="test-agent", + request=mock_request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + assert response.media_type == "text/event-stream" + async for chunk in response.body_iterator: + chunks.append(chunk) + + full = "".join(chunks) + body = json.loads(full.removeprefix("data: ").strip()) + assert body["id"] == "req-1" + assert body["error"]["code"] == -32001 + assert body["error"]["message"] == "Task not found" + + +@pytest.mark.asyncio +async def test_forward_jsonrpc_sse_fallback_error_uses_jsonrpc_error_code(): + mock_resp = AsyncMock() + mock_resp.is_success = False + mock_resp.status_code = 503 + mock_resp.reason_phrase = "Service Unavailable" + mock_resp.aread = AsyncMock(return_value=b"upstream unavailable") + mock_resp.aclose = AsyncMock() + + mock_async_client = MagicMock() + mock_async_client.build_request = MagicMock(return_value=MagicMock()) + mock_async_client.send = AsyncMock(return_value=mock_resp) + + mock_handler = MagicMock() + mock_handler.client = mock_async_client + + with patch( + "litellm.llms.custom_httpx.http_handler.get_async_httpx_client", + return_value=mock_handler, + ): + from litellm.proxy.agent_endpoints.a2a_endpoints import _forward_jsonrpc_sse + + response = await _forward_jsonrpc_sse( + agent_url="http://backend-agent:10001", + body={"jsonrpc": "2.0", "id": "req-1", "method": "tasks/resubscribe"}, + request_id="req-1", + ) + + chunks = [] + async for chunk in response.body_iterator: + chunks.append(chunk) + + body = json.loads("".join(chunks).removeprefix("data: ").strip()) + assert body["error"]["code"] == -32603 + assert body["error"]["message"] == "Service Unavailable" + + +@pytest.mark.asyncio +async def test_task_methods_forward_caller_identity_headers(): + """Task operations must forward X-LiteLLM-User-Id and X-LiteLLM-Team-Id so the + upstream agent can scope resources to the authenticated caller.""" + from litellm.proxy._types import UserAPIKeyAuth + + upstream_response = { + "jsonrpc": "2.0", + "id": "req-1", + "result": {"id": "task-1", "status": {"state": "completed"}}, + } + agent = _make_agent_mock() + mock_request = _make_request_mock("tasks/get", {"id": "task-1"}) + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-test", user_id="user-abc", team_id="team-xyz" + ) + + mock_http_response = MagicMock() + mock_http_response.json.return_value = upstream_response + mock_http_response.is_success = True + + mock_handler = MagicMock() + mock_handler.post = AsyncMock(return_value=mock_http_response) + + with ExitStack() as stack: + for p in _base_patches(agent): + stack.enter_context(p) + stack.enter_context( + patch( + "litellm.llms.custom_httpx.http_handler.get_async_httpx_client", + return_value=mock_handler, + ) + ) + + from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a + + await invoke_agent_a2a( + agent_id="test-agent", + request=mock_request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + posted_headers = mock_handler.post.call_args.kwargs.get("headers") or {} + assert posted_headers.get("X-LiteLLM-User-Id") == "user-abc" + assert posted_headers.get("X-LiteLLM-Team-Id") == "team-xyz" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("method", ["tasks/get", "tasks/resubscribe"]) +async def test_task_methods_forward_trace_header(method: str): + from litellm.proxy._types import UserAPIKeyAuth + + agent = _make_agent_mock() + mock_request = _make_request_mock(method, {"id": "task-1"}) + user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1") + + async def add_proxy_data_with_trace(data, **kwargs): + data = await _add_proxy_data(data, **kwargs) + data["litellm_trace_id"] = "trace-123" + return data + + upstream_response = { + "jsonrpc": "2.0", + "id": "req-1", + "result": {"id": "task-1", "status": {"state": "completed"}}, + } + + mock_http_response = MagicMock() + mock_http_response.json.return_value = upstream_response + mock_http_response.is_success = True + + async def fake_aiter_lines(): + yield 'data: {"jsonrpc":"2.0","id":"req-1","result":{"taskId":"task-1"}}' + + mock_resp = AsyncMock() + mock_resp.is_success = True + mock_resp.aiter_lines = fake_aiter_lines + mock_resp.aclose = AsyncMock() + + mock_async_client = MagicMock() + mock_async_client.build_request = MagicMock(return_value=MagicMock()) + mock_async_client.send = AsyncMock(return_value=mock_resp) + + mock_handler = MagicMock() + mock_handler.post = AsyncMock(return_value=mock_http_response) + mock_handler.client = mock_async_client + + with ExitStack() as stack: + for p in _base_patches(agent): + stack.enter_context(p) + stack.enter_context( + patch( + "litellm.proxy.common_request_processing.add_litellm_data_to_request", + new=AsyncMock(side_effect=add_proxy_data_with_trace), + ) + ) + stack.enter_context( + patch( + "litellm.llms.custom_httpx.http_handler.get_async_httpx_client", + return_value=mock_handler, + ) + ) + + from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a + + response = await invoke_agent_a2a( + agent_id="test-agent", + request=mock_request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + if method == "tasks/resubscribe": + async for _ in response.body_iterator: + pass + + if method == "tasks/resubscribe": + forwarded_headers = mock_async_client.build_request.call_args.kwargs["headers"] + else: + forwarded_headers = mock_handler.post.call_args.kwargs["headers"] + assert forwarded_headers.get("X-LiteLLM-Trace-Id") == "trace-123" + + +@pytest.mark.asyncio +async def test_push_notification_config_set_rejects_http_url(): + """tasks/pushNotificationConfig/set must reject non-HTTPS callback URLs to prevent SSRF.""" + from fastapi import HTTPException + + from litellm.proxy._types import UserAPIKeyAuth + + agent = _make_agent_mock() + mock_request = _make_request_mock( + "tasks/pushNotificationConfig/set", + {"taskId": "task-1", "url": "http://internal-webhook.example.com/hook"}, + ) + user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1") + + with ExitStack() as stack: + for p in _base_patches(agent): + stack.enter_context(p) + + from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a + + with pytest.raises(HTTPException) as exc_info: + await invoke_agent_a2a( + agent_id="test-agent", + request=mock_request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + assert exc_info.value.status_code == 400 + assert "HTTPS" in exc_info.value.detail + + +@pytest.mark.asyncio +async def test_push_notification_config_set_rejects_private_ip(): + """tasks/pushNotificationConfig/set must reject callback URLs pointing to private IP ranges.""" + from fastapi import HTTPException + + from litellm.proxy._types import UserAPIKeyAuth + + agent = _make_agent_mock() + mock_request = _make_request_mock( + "tasks/pushNotificationConfig/set", + {"taskId": "task-1", "url": "https://192.168.1.100/hook"}, + ) + user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1") + + with ExitStack() as stack: + for p in _base_patches(agent): + stack.enter_context(p) + + from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a + + with pytest.raises(HTTPException) as exc_info: + await invoke_agent_a2a( + agent_id="test-agent", + request=mock_request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + assert exc_info.value.status_code == 400 + assert "blocked address" in exc_info.value.detail.lower() + + +@pytest.mark.asyncio +async def test_push_notification_config_set_validates_nested_url_when_top_level_present(): + """A safe top-level params.url must not let a private pushNotificationConfig.url bypass SSRF checks. + + Both URL-bearing fields are forwarded to the agent, so both must be validated independently. + """ + from fastapi import HTTPException + + from litellm.proxy._types import UserAPIKeyAuth + + agent = _make_agent_mock() + mock_request = _make_request_mock( + "tasks/pushNotificationConfig/set", + { + "taskId": "task-1", + "url": "https://1.1.1.1/hook", + "pushNotificationConfig": {"url": "https://192.168.1.100/hook"}, + }, + ) + user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1") + + with ExitStack() as stack: + for p in _base_patches(agent): + stack.enter_context(p) + + from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a + + with pytest.raises(HTTPException) as exc_info: + await invoke_agent_a2a( + agent_id="test-agent", + request=mock_request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + assert exc_info.value.status_code == 400 + assert "blocked address" in exc_info.value.detail.lower() + + +def test_push_notification_config_set_rejects_private_dns_resolution(): + from fastapi import HTTPException + + from litellm.proxy.agent_endpoints.a2a_endpoints import ( + _validate_push_notification_url, + ) + + with patch( + "litellm.litellm_core_utils.url_utils.socket.getaddrinfo", + return_value=[ + ( + socket.AF_INET, + socket.SOCK_STREAM, + socket.IPPROTO_TCP, + "", + ("10.0.0.5", 443), + ) + ], + ): + with pytest.raises(HTTPException) as exc_info: + _validate_push_notification_url("https://webhook.example.com/hook") + + assert exc_info.value.status_code == 400 + assert "blocked address" in exc_info.value.detail.lower() + + +@pytest.mark.asyncio +async def test_push_notification_config_set_rejects_null_push_config(): + from fastapi import HTTPException + + from litellm.proxy._types import UserAPIKeyAuth + + agent = _make_agent_mock() + mock_request = _make_request_mock( + "tasks/pushNotificationConfig/set", + {"taskId": "task-1", "pushNotificationConfig": None}, + ) + user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1") + + with ExitStack() as stack: + for p in _base_patches(agent): + stack.enter_context(p) + + from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a + + with pytest.raises(HTTPException) as exc_info: + await invoke_agent_a2a( + agent_id="test-agent", + request=mock_request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + assert exc_info.value.status_code == 400 + assert "pushNotificationConfig must be an object" in exc_info.value.detail + + +@pytest.mark.asyncio +async def test_caller_identity_headers_cannot_be_spoofed_via_forwarded_headers(): + """A client must not be able to override X-LiteLLM-User-Id / X-LiteLLM-Team-Id + by including x-a2a--x-litellm-user-id in their request headers. + The authenticated identity must always win.""" + from litellm.proxy._types import UserAPIKeyAuth + + upstream_response = { + "jsonrpc": "2.0", + "id": "req-1", + "result": {"id": "task-1", "status": {"state": "completed"}}, + } + agent = _make_agent_mock() + mock_request = _make_request_mock("tasks/get", {"id": "task-1"}) + mock_request.headers = { + "x-a2a-test-agent-x-litellm-user-id": "attacker-user", + "x-a2a-test-agent-x-litellm-team-id": "attacker-team", + } + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-test", user_id="real-user", team_id="real-team" + ) + + mock_http_response = MagicMock() + mock_http_response.json.return_value = upstream_response + mock_http_response.is_success = True + + mock_handler = MagicMock() + mock_handler.post = AsyncMock(return_value=mock_http_response) + + with ExitStack() as stack: + for p in _base_patches(agent): + stack.enter_context(p) + stack.enter_context( + patch( + "litellm.llms.custom_httpx.http_handler.get_async_httpx_client", + return_value=mock_handler, + ) + ) + + from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a + + await invoke_agent_a2a( + agent_id="test-agent", + request=mock_request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + posted_headers = mock_handler.post.call_args.kwargs.get("headers") or {} + assert ( + posted_headers.get("X-LiteLLM-User-Id") == "real-user" + ), "authenticated user id must not be overridden by forwarded client headers" + assert ( + posted_headers.get("X-LiteLLM-Team-Id") == "real-team" + ), "authenticated team id must not be overridden by forwarded client headers" diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 6f6c4508333..d5043c775b3 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -3391,30 +3391,44 @@ async def test_resolve_end_user_swallows_db_errors_and_returns_none( @pytest.mark.asyncio -async def test_resolve_end_user_reraises_budget_exceeded( +async def test_resolve_end_user( _validate_flag_on, monkeypatch ): - """BudgetExceededError from get_end_user_object must bubble up so the - auth path enforces spend limits instead of silently dropping the id.""" - import litellm + """Verify that resolve_and_validate_end_user_id does NOT raise BudgetExceededError. + + Note: As of the refactor that moved _check_end_user_budget out of + get_end_user_object, budget enforcement now happens in common_checks(). + + The end-user validation path should return the user ID regardless of budget status. + Budget enforcement for end users happens later in common_checks() via + _check_end_user_budget(), which respects skip_budget_checks for zero-cost models. + + This test verifies that even when get_end_user_object returns a user with a budget, + resolve_and_validate_end_user_id does not block the request - budget enforcement + is deferred to common_checks() where skip_budget_checks logic can be applied. + """ from litellm.proxy.auth import auth_checks from litellm.proxy.auth.auth_checks import resolve_and_validate_end_user_id + # Mock get_end_user_object to return a user with budget info + # (simulating a user who may have exceeded their budget) + mock_end_user = MagicMock() + mock_end_user.user_id = "customer-over-budget" monkeypatch.setattr( auth_checks, "get_end_user_object", - AsyncMock( - side_effect=litellm.BudgetExceededError(current_cost=10.0, max_budget=5.0) - ), + AsyncMock(return_value=mock_end_user), ) cache = _validation_cache() - with pytest.raises(litellm.BudgetExceededError): - await resolve_and_validate_end_user_id( - raw_end_user_id="customer-over-budget", - prisma_client=MagicMock(), - user_api_key_cache=cache, - ) + # resolve_and_validate_end_user_id should return the user ID without raising + # BudgetExceededError - budget enforcement happens in common_checks() + result = await resolve_and_validate_end_user_id( + raw_end_user_id="customer-over-budget", + prisma_client=MagicMock(), + user_api_key_cache=cache, + ) + assert result == "customer-over-budget" @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index 2d40db9017e..d4ca55ca16b 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -14,6 +14,7 @@ from litellm.proxy.auth.auth_utils import ( abbreviate_api_key, check_complete_credentials, get_end_user_id_from_request_body, + get_key_mcp_rpm_limit, get_key_model_rpm_limit, get_key_model_tpm_limit, get_model_from_request, @@ -92,6 +93,22 @@ class TestGetKeyModelRpmLimit: assert result == {} +class TestGetKeyMcpRpmLimit: + def test_empty_dict_limits_are_returned(self): + key_override = UserAPIKeyAuth( + api_key="sk-123", + metadata={"mcp_rpm_limit": {}}, + team_metadata={"mcp_rpm_limit": {"github": 50}}, + ) + assert get_key_mcp_rpm_limit(key_override) == {} + + team_empty = UserAPIKeyAuth( + api_key="sk-123", + team_metadata={"mcp_rpm_limit": {}}, + ) + assert get_key_mcp_rpm_limit(team_empty) == {} + + class TestGetKeyModelTpmLimit: """Tests for get_key_model_tpm_limit function.""" @@ -382,6 +399,62 @@ def test_get_model_from_request_extracts_video_id_model(): ) +def test_get_model_from_request_resolves_video_id_model_with_router(): + from litellm.types.videos.utils import encode_video_id_with_provider + + provider_video_id = ( + "projects/test-project/locations/us-central1/publishers/google/models/" + "veo-3.1-generate-001/operations/operation-id" + ) + video_id = encode_video_id_with_provider( + video_id=provider_video_id, + provider="vertex_ai", + model_id="veo-3.1-generate-001", + ) + llm_router = MagicMock() + llm_router.resolve_model_name_from_model_id.return_value = ( + "gcp/google/veo-3.1-generate-001" + ) + + assert ( + get_model_from_request( + request_data={"video_id": video_id}, + route="/v1/videos/{video_id}", + llm_router=llm_router, + ) + == "gcp/google/veo-3.1-generate-001" + ) + llm_router.resolve_model_name_from_model_id.assert_called_once_with( + "veo-3.1-generate-001" + ) + + +def test_get_model_from_request_resolves_character_id_model_with_router(): + from litellm.types.videos.utils import encode_character_id_with_provider + + character_id = encode_character_id_with_provider( + character_id="character-provider-id", + provider="vertex_ai", + model_id="veo-3.1-generate-001", + ) + llm_router = MagicMock() + llm_router.resolve_model_name_from_model_id.return_value = ( + "gcp/google/veo-3.1-generate-001" + ) + + assert ( + get_model_from_request( + request_data={"character_id": character_id}, + route="/v1/videos/characters/{character_id}", + llm_router=llm_router, + ) + == "gcp/google/veo-3.1-generate-001" + ) + llm_router.resolve_model_name_from_model_id.assert_called_once_with( + "veo-3.1-generate-001" + ) + + def test_get_model_from_request_only_runs_media_decoders_for_matching_fields(): with ( patch( 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 index 4084fa4f3aa..68907de6f2d 100644 --- 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 @@ -5,7 +5,11 @@ 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 +from litellm.proxy._types import ( + LiteLLM_BudgetTable, + LiteLLM_EndUserTable, + UserAPIKeyAuth, +) @pytest.mark.asyncio @@ -88,6 +92,85 @@ async def test_custom_auth_run_post_custom_auth_checks_with_end_user_budget_exce mock_budget_check.assert_awaited_once() +@pytest.mark.asyncio +async def test_custom_auth_enforces_end_user_budget_when_common_checks_skipped(): + # custom-auth deployments with custom_auth_run_common_checks unset skip + # common_checks() (and its end-user budget enforcement) in the centralized + # gate, so the helper must enforce the end-user budget itself. Regression: + # an over-budget end user must be rejected on this path. + valid_token = UserAPIKeyAuth(token="test_token", end_user_id="customer-1") + over_budget_end_user = LiteLLM_EndUserTable( + user_id="customer-1", + blocked=False, + spend=0.0, + litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0), + ) + + async def mock_get_current_spend(counter_key, fallback_spend): + if counter_key == "spend:end_user:customer-1": + return 5.0 + return fallback_spend + + with ( + patch( + "litellm.proxy.auth.user_api_key_auth.get_end_user_object", + new_callable=AsyncMock, + return_value=over_budget_end_user, + ), + patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend), + patch("litellm.proxy.proxy_server.general_settings", {}), + ): + with pytest.raises(litellm.BudgetExceededError): + await _run_post_custom_auth_checks( + valid_token=valid_token, + request=None, + request_data={"model": "gpt-4"}, + route="/v1/chat/completions", + parent_otel_span=None, + ) + + +@pytest.mark.asyncio +async def test_custom_auth_defers_end_user_budget_to_common_checks_when_enabled(): + # With custom_auth_run_common_checks set, the wrapper's common_checks() + # enforces the end-user budget, so the helper must not double-enforce it. + valid_token = UserAPIKeyAuth(token="test_token", end_user_id="customer-1") + end_user_obj = LiteLLM_EndUserTable( + user_id="customer-1", + blocked=False, + spend=0.0, + litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0), + ) + + with ( + patch( + "litellm.proxy.auth.user_api_key_auth.get_end_user_object", + new_callable=AsyncMock, + return_value=end_user_obj, + ), + patch( + "litellm.proxy.auth.user_api_key_auth._check_end_user_budget", + new_callable=AsyncMock, + ) as mock_check, + patch( + "litellm.proxy.auth.user_api_key_auth._enforce_key_and_fallback_model_access", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.proxy_server.general_settings", + {"custom_auth_run_common_checks": True}, + ), + ): + await _run_post_custom_auth_checks( + valid_token=valid_token, + request=None, + request_data={"model": "gpt-4"}, + route="/v1/chat/completions", + parent_otel_span=None, + ) + mock_check.assert_not_awaited() + + 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, diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index 14119f7ad4e..41fc417d979 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -3179,3 +3179,727 @@ def test_build_decode_kwargs_no_warning_when_scoped( if "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage() ] assert matching == [] + + +@pytest.mark.asyncio +async def test_auth_jwt_expired_token_raises_401_jwk_path(): + """An expired JWT (access token) decoded via the JWK/dict public-key path + must raise a ProxyException carrying a 401 status code so the status is + preserved end-to-end (client response + OTel traces). + """ + import jwt as jwt_lib + + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth() + + with ( + patch.object( + jwt_handler, "get_public_key", new_callable=AsyncMock + ) as mock_get_public_key, + patch( + "litellm.proxy.auth.handle_jwt.jwt.get_unverified_header", + return_value={"kid": "test-kid"}, + ), + patch( + "litellm.proxy.auth.handle_jwt.PyJWK.from_dict", + return_value=MagicMock(key="fake-key"), + ), + patch( + "litellm.proxy.auth.handle_jwt.jwt.decode", + side_effect=jwt_lib.ExpiredSignatureError("Signature has expired"), + ), + ): + mock_get_public_key.return_value = {"kty": "RSA", "kid": "test-kid"} + + with pytest.raises(ProxyException) as exc_info: + await jwt_handler.auth_jwt(token="expired.jwt.token") + + assert exc_info.value.code == str(401) + assert exc_info.value.type == ProxyErrorTypes.expired_key.value + assert "Token Expired" in exc_info.value.message + + +@pytest.mark.asyncio +async def test_auth_jwt_expired_token_raises_401_pem_cert_path(): + """Same as above but for the PEM-certificate (string public-key) decode path.""" + import jwt as jwt_lib + + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth() + + mock_cert = MagicMock() + mock_cert.public_key.return_value.public_bytes.return_value = b"fake-key" + + with ( + patch.object( + jwt_handler, "get_public_key", new_callable=AsyncMock + ) as mock_get_public_key, + patch( + "litellm.proxy.auth.handle_jwt.jwt.get_unverified_header", + return_value={"kid": "test-kid"}, + ), + patch( + "litellm.proxy.auth.handle_jwt.x509.load_pem_x509_certificate", + return_value=mock_cert, + ), + patch( + "litellm.proxy.auth.handle_jwt.jwt.decode", + side_effect=jwt_lib.ExpiredSignatureError("Signature has expired"), + ), + ): + mock_get_public_key.return_value = ( + "-----BEGIN CERTIFICATE-----\nfake\n-----END CERTIFICATE-----" + ) + + with pytest.raises(ProxyException) as exc_info: + await jwt_handler.auth_jwt(token="expired.jwt.token") + + assert exc_info.value.code == str(401) + assert exc_info.value.type == ProxyErrorTypes.expired_key.value + assert "Token Expired" in exc_info.value.message + + +def _base64url_encode_int(value: int) -> str: + import base64 + + value_bytes = value.to_bytes((value.bit_length() + 7) // 8, "big") + return base64.urlsafe_b64encode(value_bytes).decode("utf-8").rstrip("=") + + +def _get_rsa_key_and_jwk(kid: str): + from cryptography.hazmat.primitives.asymmetric import rsa + + private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + public_numbers = private_key.public_key().public_numbers() + jwk = { + "kty": "RSA", + "n": _base64url_encode_int(value=public_numbers.n), + "e": _base64url_encode_int(value=public_numbers.e), + "kid": kid, + "alg": "RS256", + "use": "sig", + } + return private_key, jwk + + +def _encode_rsa_jwt( + private_key, + issuer: str, + audience: str, + kid: str, + extra_claims: Optional[dict] = None, +) -> str: + import time + + import jwt + from cryptography.hazmat.primitives import serialization + + private_key_pem = private_key.private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.PKCS8, + encryption_algorithm=serialization.NoEncryption(), + ) + current_time = int(time.time()) + claims = { + "sub": "test-subject", + "iss": issuer, + "aud": audience, + "iat": current_time, + "exp": current_time + 300, + } + if extra_claims: + claims.update(extra_claims) + + return jwt.encode( + claims, + private_key_pem, + algorithm="RS256", + headers={"kid": kid}, + ) + + +def _get_jwt_handler_with_issuer_keys(issuers: list, keys_by_url: dict) -> JWTHandler: + from litellm.caching.dual_cache import DualCache + + cache = DualCache() + for jwks_url, keys in keys_by_url.items(): + cache.set_cache( + key=f"litellm_jwt_auth_keys_{jwks_url}", + value=keys, + ) + + jwt_handler = JWTHandler() + jwt_handler.update_environment( + prisma_client=None, + user_api_key_cache=cache, + litellm_jwtauth=LiteLLM_JWTAuth(issuers=issuers), + ) + return jwt_handler + + +@pytest.mark.asyncio +async def test_get_public_key_fetches_and_caches_jwks_response(): + from unittest.mock import AsyncMock, MagicMock + + from litellm.caching.dual_cache import DualCache + + jwt_handler = JWTHandler() + cache = DualCache() + jwt_handler.update_environment( + prisma_client=None, + user_api_key_cache=cache, + litellm_jwtauth=LiteLLM_JWTAuth(public_key_ttl=123), + ) + expected_key_id = "cached-key" + _, jwk = _get_rsa_key_and_jwk(kid=expected_key_id) + mock_response = MagicMock() + mock_response.json.return_value = {"keys": [jwk]} + jwt_handler.http_handler.get = AsyncMock(return_value=mock_response) + + public_key = await jwt_handler._get_public_key_from_jwks_url( + jwks_url="https://issuer.example.com/keys", + kid=expected_key_id, + ) + + assert public_key == jwk + cached_keys = await cache.async_get_cache( + key="litellm_jwt_auth_keys_https://issuer.example.com/keys" + ) + assert cached_keys == [jwk] + + +@pytest.mark.asyncio +async def test_get_public_key_tries_next_jwks_url_when_kid_missing(monkeypatch): + from litellm.caching.dual_cache import DualCache + + first_jwks_url = "https://first.example.com/keys" + second_jwks_url = "https://second.example.com/keys" + monkeypatch.setenv("JWT_PUBLIC_KEY_URL", f"{first_jwks_url}, {second_jwks_url},,") + _, first_jwk = _get_rsa_key_and_jwk(kid="first-key") + _, second_jwk = _get_rsa_key_and_jwk(kid="second-key") + cache = DualCache() + cache.set_cache(key=f"litellm_jwt_auth_keys_{first_jwks_url}", value=[first_jwk]) + cache.set_cache(key=f"litellm_jwt_auth_keys_{second_jwks_url}", value=[second_jwk]) + jwt_handler = JWTHandler() + jwt_handler.update_environment( + prisma_client=None, + user_api_key_cache=cache, + litellm_jwtauth=LiteLLM_JWTAuth(), + ) + + public_key = await jwt_handler.get_public_key(kid="second-key") + + assert public_key == second_jwk + + +def test_get_jwks_url_for_issuer_falls_back_to_discovery_document(): + jwt_handler = JWTHandler() + issuer_config = LiteLLM_JWTAuth( + issuers=[ + { + "issuer": "https://issuer.example.com/tenant/", + "disable_audience_validation": True, + } + ] + ).issuers[0] + + jwks_url = jwt_handler._get_jwks_url_for_issuer(issuer_config=issuer_config) + + assert ( + jwks_url == "https://issuer.example.com/tenant/.well-known/openid-configuration" + ) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_validates_selected_issuer_and_maps_claims( + monkeypatch, +): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer_one = "https://issuer-one.example.com" + issuer_two = "https://issuer-two.example.com" + issuer_one_jwks_url = f"{issuer_one}/keys" + issuer_two_jwks_url = f"{issuer_two}/keys" + shared_kid = "shared-kid" + + _, issuer_one_jwk = _get_rsa_key_and_jwk(kid=shared_kid) + issuer_two_private_key, issuer_two_jwk = _get_rsa_key_and_jwk(kid=shared_kid) + + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer_one, + "jwks_url": issuer_one_jwks_url, + "audience": "audience-one", + "user_id_jwt_field": "email", + "user_email_jwt_field": "email", + }, + { + "issuer": issuer_two, + "jwks_url": issuer_two_jwks_url, + "audience": "audience-two", + "user_id_jwt_field": "repository_owner", + "team_id_jwt_field": "repository", + }, + ], + keys_by_url={ + issuer_one_jwks_url: [issuer_one_jwk], + issuer_two_jwks_url: [issuer_two_jwk], + }, + ) + + token = _encode_rsa_jwt( + private_key=issuer_two_private_key, + issuer=issuer_two, + audience="audience-two", + kid=shared_kid, + extra_claims={ + "repository_owner": "example-org", + "repository": "example-org/litellm-fork", + }, + ) + + claims = await jwt_handler.auth_jwt(token=token) + + assert claims[JWTHandler.LITELLM_JWT_ISSUER_CLAIM] == issuer_two + assert jwt_handler.get_user_id(token=claims, default_value=None) == "example-org" + assert jwt_handler.get_team_id(token=claims, default_value=None) == ( + "example-org/litellm-fork" + ) + + +@pytest.mark.asyncio +async def test_auth_jwt_issuer_path_expired_token_raises_401(monkeypatch): + """An expired JWT validated through the issuer-scoped path + (_auth_jwt_with_issuer) must raise a ProxyException carrying a 401 so the + status is preserved end-to-end, just like the non-issuer path. + """ + import time + + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer = "https://issuer.example.com" + jwks_url = f"{issuer}/keys" + kid = "expired-kid" + + private_key, jwk = _get_rsa_key_and_jwk(kid=kid) + + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[{"issuer": issuer, "jwks_url": jwks_url, "audience": "my-audience"}], + keys_by_url={jwks_url: [jwk]}, + ) + + token = _encode_rsa_jwt( + private_key=private_key, + issuer=issuer, + audience="my-audience", + kid=kid, + extra_claims={"exp": int(time.time()) - 100}, + ) + + with pytest.raises(ProxyException) as exc_info: + await jwt_handler.auth_jwt(token=token) + + assert exc_info.value.code == str(401) + assert exc_info.value.type == ProxyErrorTypes.expired_key.value + assert "Token Expired" in exc_info.value.message + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_maps_kubernetes_namespace_claim(monkeypatch): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer = "https://oidc.eks.eu-west-1.amazonaws.com/id/test-cluster" + jwks_url = f"{issuer}/keys" + private_key, jwk = _get_rsa_key_and_jwk(kid="k8s-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + "audience": None, + "disable_audience_validation": True, + "user_id_jwt_field": "kubernetes\\.io.namespace", + } + ], + keys_by_url={jwks_url: [jwk]}, + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer=issuer, + audience="kubernetes.default.svc", + kid="k8s-key", + extra_claims={"kubernetes.io": {"namespace": "example-namespace"}}, + ) + + claims = await jwt_handler.auth_jwt(token=token) + + assert ( + jwt_handler.get_user_id(token=claims, default_value=None) == "example-namespace" + ) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_unknown_issuer_falls_back_to_global_jwks(monkeypatch): + """Tokens whose ``iss`` is not in the configured issuers list fall through + to the legacy ``JWT_PUBLIC_KEY_URL`` path so operators can add the new + ``issuers`` list to a live deployment without breaking existing tokens + minted by non-configured IdPs. With no global JWKS configured, the legacy + path surfaces a ``Missing JWT Public Key URL from environment.`` error. + """ + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + configured_issuer = "https://issuer.example.com" + private_key, jwk = _get_rsa_key_and_jwk(kid="issuer-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": configured_issuer, + "jwks_url": f"{configured_issuer}/keys", + "audience": "expected-audience", + } + ], + keys_by_url={f"{configured_issuer}/keys": [jwk]}, + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer="https://unknown-issuer.example.com", + audience="expected-audience", + kid="issuer-key", + ) + + with pytest.raises(Exception) as exc: + await jwt_handler.auth_jwt(token=token) + + assert "Missing JWT Public Key URL from environment." in str(exc.value) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_rejects_wrong_audience(monkeypatch): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer = "https://issuer.example.com" + jwks_url = f"{issuer}/keys" + private_key, jwk = _get_rsa_key_and_jwk(kid="issuer-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + "audience": "expected-audience", + } + ], + keys_by_url={jwks_url: [jwk]}, + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer=issuer, + audience="wrong-audience", + kid="issuer-key", + ) + + with pytest.raises(Exception) as exc: + await jwt_handler.auth_jwt(token=token) + + assert "Validation fails" in str(exc.value) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_same_kid_does_not_cross_issuer_keys(monkeypatch): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer_one = "https://issuer-one.example.com" + issuer_two = "https://issuer-two.example.com" + issuer_one_jwks_url = f"{issuer_one}/keys" + issuer_two_jwks_url = f"{issuer_two}/keys" + shared_kid = "shared-kid" + issuer_one_private_key, issuer_one_jwk = _get_rsa_key_and_jwk(kid=shared_kid) + _, issuer_two_jwk = _get_rsa_key_and_jwk(kid=shared_kid) + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer_one, + "jwks_url": issuer_one_jwks_url, + "audience": "audience-one", + }, + { + "issuer": issuer_two, + "jwks_url": issuer_two_jwks_url, + "audience": "audience-two", + }, + ], + keys_by_url={ + issuer_one_jwks_url: [issuer_one_jwk], + issuer_two_jwks_url: [issuer_two_jwk], + }, + ) + token = _encode_rsa_jwt( + private_key=issuer_one_private_key, + issuer=issuer_two, + audience="audience-two", + kid=shared_kid, + ) + + with pytest.raises(Exception) as exc: + await jwt_handler.auth_jwt(token=token) + + assert "Validation fails" in str(exc.value) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_missing_mapped_claim_leaves_user_id_unset( + monkeypatch, +): + """Mapped issuer claims behave like the global ``litellm_jwtauth`` path — + present claims override the normalised value, missing ones simply leave + the corresponding LiteLLM-internal claim absent (rather than failing the + JWT outright). This keeps multi-issuer auth tolerant of tokens that omit + optional fields like email or org id. + """ + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer = "https://issuer.example.com" + jwks_url = f"{issuer}/keys" + private_key, jwk = _get_rsa_key_and_jwk(kid="issuer-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + "audience": "expected-audience", + "user_id_jwt_field": "email", + } + ], + keys_by_url={jwks_url: [jwk]}, + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer=issuer, + audience="expected-audience", + kid="issuer-key", + ) + + claims = await jwt_handler.auth_jwt(token=token) + + assert claims[jwt_handler.LITELLM_JWT_ISSUER_CLAIM] == issuer + assert jwt_handler.LITELLM_USER_ID_CLAIM not in claims + + +def test_multi_issuer_jwt_requires_audience_unless_explicitly_disabled( + monkeypatch, +): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer = "https://issuer.example.com" + jwks_url = f"{issuer}/keys" + + with pytest.raises(Exception) as exc: + LiteLLM_JWTAuth( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + } + ] + ) + + assert "must configure audience" in str(exc.value) + + +def test_multi_issuer_jwt_rejects_audience_with_disable_audience_validation(): + issuer = "https://issuer.example.com" + jwks_url = f"{issuer}/keys" + + with pytest.raises(Exception) as exc: + LiteLLM_JWTAuth( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + "audience": "some-audience", + "disable_audience_validation": True, + } + ] + ) + + assert "cannot set audience and disable_audience_validation=True together" in str( + exc.value + ) + + +@pytest.mark.asyncio +async def test_global_jwt_ignores_user_supplied_internal_claims(monkeypatch): + from litellm.caching.dual_cache import DualCache + + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_ISSUER", raising=False) + + jwks_url = "https://global-issuer.example.com/keys" + monkeypatch.setenv("JWT_PUBLIC_KEY_URL", jwks_url) + + private_key, jwk = _get_rsa_key_and_jwk(kid="global-key") + cache = DualCache() + cache.set_cache(key=f"litellm_jwt_auth_keys_{jwks_url}", value=[jwk]) + + jwt_handler = JWTHandler() + jwt_handler.update_environment( + prisma_client=None, + user_api_key_cache=cache, + litellm_jwtauth=LiteLLM_JWTAuth( + user_id_jwt_field="email", + user_email_jwt_field="email", + team_id_jwt_field="team.id", + team_ids_jwt_field="teams", + org_id_jwt_field="org.id", + end_user_id_jwt_field="end_user.id", + ), + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer="https://global-issuer.example.com", + audience="some-other-client", + kid="global-key", + extra_claims={ + "email": "real-user@example.com", + "team": {"id": "real-team"}, + "teams": ["real-team", "secondary-team"], + "org": {"id": "real-org"}, + "end_user": {"id": "real-end-user"}, + JWTHandler.LITELLM_JWT_ISSUER_CLAIM: "https://issuer.example.com", + JWTHandler.LITELLM_USER_ID_CLAIM: "victim-user", + JWTHandler.LITELLM_USER_EMAIL_CLAIM: "victim@example.com", + JWTHandler.LITELLM_TEAM_ID_CLAIM: "victim-team", + JWTHandler.LITELLM_TEAM_IDS_CLAIM: ["victim-team"], + JWTHandler.LITELLM_ORG_ID_CLAIM: "victim-org", + JWTHandler.LITELLM_END_USER_ID_CLAIM: "victim-end-user", + }, + ) + + claims = await jwt_handler.auth_jwt(token=token) + + assert jwt_handler.get_user_id(token=claims, default_value=None) == ( + "real-user@example.com" + ) + assert jwt_handler.get_user_email(token=claims, default_value=None) == ( + "real-user@example.com" + ) + assert jwt_handler.get_team_id(token=claims, default_value=None) == "real-team" + assert jwt_handler.get_team_ids_from_jwt(token=claims) == [ + "real-team", + "secondary-team", + ] + assert jwt_handler.get_org_id(token=claims, default_value=None) == "real-org" + assert jwt_handler.get_end_user_id(token=claims, default_value=None) == ( + "real-end-user" + ) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_strips_unmapped_internal_claims(monkeypatch): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer = "https://issuer.example.com" + jwks_url = f"{issuer}/keys" + private_key, jwk = _get_rsa_key_and_jwk(kid="issuer-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + "audience": "expected-audience", + "user_email_jwt_field": "email", + } + ], + keys_by_url={jwks_url: [jwk]}, + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer=issuer, + audience="expected-audience", + kid="issuer-key", + extra_claims={ + "email": "real-user@example.com", + JWTHandler.LITELLM_USER_ID_CLAIM: "victim-user", + JWTHandler.LITELLM_TEAM_ID_CLAIM: "victim-team", + }, + ) + + claims = await jwt_handler.auth_jwt(token=token) + + assert JWTHandler.LITELLM_USER_ID_CLAIM not in claims + assert JWTHandler.LITELLM_TEAM_ID_CLAIM not in claims + assert jwt_handler.get_user_id(token=claims, default_value=None) is None + assert jwt_handler.get_team_id(token=claims, default_value=None) is None + assert jwt_handler.get_user_email(token=claims, default_value=None) == ( + "real-user@example.com" + ) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_does_not_emit_unscoped_global_warning( + monkeypatch, caplog +): + import logging + + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_ISSUER", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + JWTHandler._unscoped_jwt_warning_emitted = False + + issuer = "https://issuer.example.com" + jwks_url = f"{issuer}/keys" + private_key, jwk = _get_rsa_key_and_jwk(kid="issuer-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + "audience": "expected-audience", + } + ], + keys_by_url={jwks_url: [jwk]}, + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer=issuer, + audience="expected-audience", + kid="issuer-key", + ) + + with caplog.at_level(logging.WARNING): + await jwt_handler.auth_jwt(token=token) + + assert "Tokens minted by any application" not in caplog.text + assert JWTHandler._unscoped_jwt_warning_emitted is False + + +def test_build_decode_kwargs_warns_for_unscoped_global_fallback_in_mixed_deployment( + monkeypatch, _reset_unscoped_warning_flag, caplog +): + """The unscoped-fallback warning must fire even when per-issuer configs + are set. In mixed deployments, tokens whose ``iss`` does not match any + configured issuer fall through to the global path; if env-var scoping is + absent that fallback IS unscoped, and the operator needs to be told.""" + import logging + + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_ISSUER", raising=False) + caplog.set_level(logging.WARNING) + + JWTHandler._build_decode_kwargs() + + matching = [ + r + for r in caplog.records + if "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage() + ] + assert len(matching) == 1 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 3b13ef3641f..9444e4ebd2d 100644 --- a/tests/test_litellm/proxy/auth/test_mcp_ip_filtering.py +++ b/tests/test_litellm/proxy/auth/test_mcp_ip_filtering.py @@ -5,8 +5,9 @@ Tests that internal callers see all MCP servers while external callers only see servers with available_on_public_internet=True. """ -import ipaddress -from unittest.mock import patch +from unittest.mock import MagicMock, patch + +from fastapi import Request from litellm.proxy.auth.ip_address_utils import IPAddressUtils from litellm.types.mcp_server.mcp_server_manager import MCPServer @@ -58,6 +59,75 @@ class TestIsInternalIp: assert IPAddressUtils.is_internal_ip("not-an-ip") is False +class TestMCPClientIPExtraction: + def test_fails_closed_when_xff_enabled_without_trusted_proxy_ranges(self): + request = MagicMock(spec=Request) + request.client = MagicMock() + request.client.host = "203.0.113.5" + request.headers = {"x-forwarded-for": "10.0.0.1"} + + result = IPAddressUtils.get_mcp_client_ip( + request, + general_settings={"use_x_forwarded_for": True}, + ) + + # XFF is untrusted (no mcp_trusted_proxy_ranges) so it must be ignored, + # and we must not trust the direct peer either: fail closed so the caller + # is classified as external and is_internal_ip("") is False. + assert result == "" + assert IPAddressUtils.is_internal_ip(result) is False + + def test_private_proxy_peer_does_not_grant_internal_access(self): + # Regression: behind an internal reverse proxy with use_x_forwarded_for + # enabled but mcp_trusted_proxy_ranges unset, the direct peer is the + # proxy's private IP. Returning it would mis-classify an external caller + # as internal and expose available_on_public_internet=false servers. + request = MagicMock(spec=Request) + request.client = MagicMock() + request.client.host = "10.0.0.7" + request.headers = {"x-forwarded-for": "8.8.8.8"} + + result = IPAddressUtils.get_mcp_client_ip( + request, + general_settings={"use_x_forwarded_for": True}, + ) + + assert result == "" + assert IPAddressUtils.is_internal_ip(result) is False + + def test_honours_xff_from_trusted_proxy(self): + request = MagicMock(spec=Request) + request.client = MagicMock() + request.client.host = "10.0.0.5" + request.headers = {"x-forwarded-for": "192.168.1.10"} + + result = IPAddressUtils.get_mcp_client_ip( + request, + general_settings={ + "use_x_forwarded_for": True, + "mcp_trusted_proxy_ranges": ["10.0.0.0/8"], + }, + ) + + assert result == "192.168.1.10" + + def test_ignores_xff_from_untrusted_direct_caller(self): + request = MagicMock(spec=Request) + request.client = MagicMock() + request.client.host = "203.0.113.5" + request.headers = {"x-forwarded-for": "10.0.0.1"} + + result = IPAddressUtils.get_mcp_client_ip( + request, + general_settings={ + "use_x_forwarded_for": True, + "mcp_trusted_proxy_ranges": ["10.0.0.0/8"], + }, + ) + + assert result == "203.0.113.5" + + class TestMCPServerIPFiltering: """Tests that external callers only see public MCP servers.""" diff --git a/tests/test_litellm/proxy/auth/test_route_checks.py b/tests/test_litellm/proxy/auth/test_route_checks.py index ad9295d6b19..197572216d0 100644 --- a/tests/test_litellm/proxy/auth/test_route_checks.py +++ b/tests/test_litellm/proxy/auth/test_route_checks.py @@ -2602,3 +2602,72 @@ def test_legitimate_passthrough_routes_still_classified_as_llm_route(route): assert ( RouteChecks.is_llm_api_route(route=route) is True ), f"{route!r} should be classified as an LLM API route" + + +@pytest.mark.parametrize( + "route", + [ + "/search_tools/list", + "/search_tools/ui/available_providers", + ], +) +def test_internal_user_can_read_search_tools(route): + """Regression for LIT-3150: internal users must be able to view search tools, + the same way they can view vector stores.""" + user_obj = LiteLLM_UserTable( + user_id="test_user", + user_email="user@example.com", + user_role=LitellmUserRoles.INTERNAL_USER.value, + ) + valid_token = UserAPIKeyAuth( + user_id="test_user", + user_role=LitellmUserRoles.INTERNAL_USER.value, + ) + request = MagicMock(spec=Request) + request.query_params = {} + + RouteChecks.non_proxy_admin_allowed_routes_check( + user_obj=user_obj, + _user_role=LitellmUserRoles.INTERNAL_USER.value, + route=route, + request=request, + valid_token=valid_token, + request_data={}, + ) + + +@pytest.mark.parametrize( + "route", + [ + "/search_tools", # create + "/search_tools/abc123", # update / delete / get-by-id + "/search_tools/test_connection", + ], +) +def test_internal_user_blocked_from_search_tool_writes(route): + """Read access must not leak the search-tool management write routes to + internal users; only proxy admins create/update/delete/test them.""" + user_obj = LiteLLM_UserTable( + user_id="test_user", + user_email="user@example.com", + user_role=LitellmUserRoles.INTERNAL_USER.value, + ) + valid_token = UserAPIKeyAuth( + user_id="test_user", + user_role=LitellmUserRoles.INTERNAL_USER.value, + ) + request = MagicMock(spec=Request) + request.query_params = {} + + with pytest.raises(Exception) as exc_info: + RouteChecks.non_proxy_admin_allowed_routes_check( + user_obj=user_obj, + _user_role=LitellmUserRoles.INTERNAL_USER.value, + route=route, + request=request, + valid_token=valid_token, + request_data={}, + ) + assert "Only proxy admin" in str(exc_info.value) + assert f"Route={route}" in str(exc_info.value) + assert "Your role=internal_user" in str(exc_info.value) diff --git a/tests/test_litellm/proxy/conftest.py b/tests/test_litellm/proxy/conftest.py index 20236ebdf45..607315eb246 100644 --- a/tests/test_litellm/proxy/conftest.py +++ b/tests/test_litellm/proxy/conftest.py @@ -14,7 +14,6 @@ import pytest import yaml from fastapi.testclient import TestClient - _PROXY_MODULE_GLOBALS_TO_ISOLATE = ( "master_key", "prisma_client", @@ -49,6 +48,18 @@ def _isolate_proxy_module_globals(): setattr(proxy_server, name, value) +@pytest.fixture(autouse=True) +def _reset_graceful_shutdown_state(): + """Graceful shutdown state is process-scoped; keep it from leaking between tests.""" + from litellm.proxy.shutdown.graceful_shutdown_manager import ( + GracefulShutdownManager, + ) + + GracefulShutdownManager.reset() + yield + GracefulShutdownManager.reset() + + def build_cache_config(enable_cache: bool = True) -> Optional[Dict]: """ Build Redis cache configuration from environment variables. diff --git a/tests/test_litellm/proxy/health_endpoints/test_graceful_shutdown_endpoints.py b/tests/test_litellm/proxy/health_endpoints/test_graceful_shutdown_endpoints.py new file mode 100644 index 00000000000..e20b54f28f5 --- /dev/null +++ b/tests/test_litellm/proxy/health_endpoints/test_graceful_shutdown_endpoints.py @@ -0,0 +1,166 @@ +""" +Behaviour tests for the graceful-shutdown health probes. + +Builds a minimal FastAPI app from the health router plus +InFlightRequestsMiddleware so the probe responses can be asserted without +standing up the full proxy. +""" + +import pytest +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from litellm.proxy.health_endpoints._health_endpoints import router +from litellm.proxy.middleware.in_flight_requests_middleware import ( + InFlightRequestsMiddleware, +) +from litellm.proxy.shutdown.graceful_shutdown_manager import GracefulShutdownManager + + +@pytest.fixture(autouse=True) +def _reset(): + GracefulShutdownManager.reset() + InFlightRequestsMiddleware._in_flight = 0 + yield + GracefulShutdownManager.reset() + InFlightRequestsMiddleware._in_flight = 0 + + +@pytest.fixture +def client(): + app = FastAPI() + app.include_router(router) + app.add_middleware(InFlightRequestsMiddleware) + return TestClient(app) + + +@pytest.fixture +def enable_drain(monkeypatch): + from litellm.proxy import proxy_server + + monkeypatch.setattr( + proxy_server, "general_settings", {"enable_drain_endpoint": True} + ) + + +@pytest.fixture +def enable_drain_with_token(monkeypatch): + from litellm.proxy import proxy_server + + monkeypatch.setattr( + proxy_server, + "general_settings", + {"enable_drain_endpoint": True, "drain_endpoint_token": "secret-123"}, + ) + + +def test_drain_disabled_by_default_returns_404_with_no_side_effect(client, monkeypatch): + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "general_settings", {}) + resp = client.get("/health/drain") + assert resp.status_code == 404 + assert GracefulShutdownManager.is_shutting_down() is False + + +def test_drain_disabled_ignores_token_header(client, monkeypatch): + """A token alone must not bypass the enable flag; otherwise enabling the + token side-channel would silently enable the endpoint.""" + from litellm.proxy import proxy_server + + monkeypatch.setattr( + proxy_server, "general_settings", {"drain_endpoint_token": "secret-123"} + ) + resp = client.get("/health/drain", headers={"X-Drain-Token": "secret-123"}) + assert resp.status_code == 404 + assert GracefulShutdownManager.is_shutting_down() is False + + +def test_drain_when_enabled_without_token_sets_shutting_down_and_returns_drained( + client, enable_drain +): + resp = client.get("/health/drain") + assert resp.status_code == 200 + body = resp.json() + assert body["status"] == "drained" + assert body["drained_requests"] == 0 + assert GracefulShutdownManager.is_shutting_down() is True + + +def test_drain_with_token_configured_rejects_missing_header( + client, enable_drain_with_token +): + resp = client.get("/health/drain") + assert resp.status_code == 401 + assert GracefulShutdownManager.is_shutting_down() is False + + +def test_drain_with_token_configured_rejects_wrong_header( + client, enable_drain_with_token +): + resp = client.get("/health/drain", headers={"X-Drain-Token": "wrong-value"}) + assert resp.status_code == 401 + assert GracefulShutdownManager.is_shutting_down() is False + + +def test_drain_with_token_configured_accepts_correct_header( + client, enable_drain_with_token +): + resp = client.get("/health/drain", headers={"X-Drain-Token": "secret-123"}) + assert resp.status_code == 200 + assert resp.json()["status"] == "drained" + assert GracefulShutdownManager.is_shutting_down() is True + + +def test_drain_with_token_from_env_var(client, enable_drain, monkeypatch): + monkeypatch.setenv("DRAIN_ENDPOINT_TOKEN", "env-token") + resp = client.get("/health/drain") + assert resp.status_code == 401 + resp = client.get("/health/drain", headers={"X-Drain-Token": "env-token"}) + assert resp.status_code == 200 + + +def test_drain_general_settings_token_overrides_env_var(client, monkeypatch): + from litellm.proxy import proxy_server + + monkeypatch.setattr( + proxy_server, + "general_settings", + {"enable_drain_endpoint": True, "drain_endpoint_token": "config-token"}, + ) + monkeypatch.setenv("DRAIN_ENDPOINT_TOKEN", "env-token") + resp = client.get("/health/drain", headers={"X-Drain-Token": "env-token"}) + assert resp.status_code == 401 + resp = client.get("/health/drain", headers={"X-Drain-Token": "config-token"}) + assert resp.status_code == 200 + + +def test_readiness_returns_503_shutting_down_during_drain(client): + GracefulShutdownManager.start_shutdown() + resp = client.get("/health/readiness") + assert resp.status_code == 503 + assert resp.json() == {"status": "shutting_down"} + + +def test_readiness_does_not_report_shutting_down_normally(client): + resp = client.get("/health/readiness") + assert resp.json().get("status") != "shutting_down" + + +def test_liveliness_returns_503_during_drain(client): + GracefulShutdownManager.start_shutdown() + resp = client.get("/health/liveliness") + assert resp.status_code == 503 + assert resp.json() == {"status": "shutting_down"} + + +def test_liveness_alias_returns_503_during_drain(client): + GracefulShutdownManager.start_shutdown() + resp = client.get("/health/liveness") + assert resp.status_code == 503 + + +def test_liveliness_returns_alive_when_not_shutting_down(client): + resp = client.get("/health/liveliness") + assert resp.status_code == 200 + assert resp.json() == "I'm alive!" diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index 3e2eb4b02c2..676f623a5dd 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -2893,3 +2893,230 @@ async def test_pre_call_hook_rejects_caller_supplied_stash_values(): ): leaked = [k for k in _LITELLM_STASH_KEYS if k in channel] assert not leaked, f"caller-supplied stash survived in {channel!r}: {leaked}" + + +# ----------------------- Per-MCP-server rate limiting (v3) ----------------------- + + +def _make_mcp_handler(): + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + return handler, local_cache + + +def _find_descriptor(descriptors, key): + return next((d for d in descriptors if d["key"] == key), None) + + +def _build_mcp_descriptors(handler, user_api_key_dict, data, call_type="call_mcp_tool"): + return handler._create_rate_limit_descriptors( + user_api_key_dict=user_api_key_dict, + data=data, + rpm_limit_type=None, + tpm_limit_type=None, + model_has_failures=False, + call_type=call_type, + ) + + +def test_mcp_per_key_descriptor_created_for_matching_server_v3(): + handler, _ = _make_mcp_handler() + api_key = hash_token("sk-mcp-key") + user_api_key_dict = UserAPIKeyAuth( + api_key=api_key, + metadata={"mcp_rpm_limit": {"github": 5}}, + ) + + descriptors = _build_mcp_descriptors( + handler, user_api_key_dict, {"mcp_server_name": "github"} + ) + + descriptor = _find_descriptor(descriptors, "mcp_per_key") + assert descriptor is not None + assert descriptor["value"] == f"{api_key}:github" + assert descriptor["rate_limit"]["requests_per_unit"] == 5 + # MCP tool calls have no token usage; tokens_per_unit must stay None so the + # TPM reservation path is never engaged (otherwise budget would leak). + assert descriptor["rate_limit"]["tokens_per_unit"] is None + + +def test_mcp_per_key_descriptor_skipped_for_non_matching_server_v3(): + handler, _ = _make_mcp_handler() + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-mcp-key"), + metadata={"mcp_rpm_limit": {"github": 5}}, + ) + + descriptors = _build_mcp_descriptors( + handler, user_api_key_dict, {"mcp_server_name": "slack"} + ) + + assert _find_descriptor(descriptors, "mcp_per_key") is None + + +def test_mcp_descriptor_skipped_for_non_mcp_request_v3(): + """A non-MCP request must not create an MCP descriptor even if the caller + injects mcp_server_name in the body; otherwise an LLM call could consume a + target server's MCP quota and 429 legitimate tool calls.""" + handler, _ = _make_mcp_handler() + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-mcp-key"), + metadata={"mcp_rpm_limit": {"github": 5}}, + ) + + descriptors = _build_mcp_descriptors( + handler, + user_api_key_dict, + {"model": "gpt-4", "mcp_server_name": "github"}, + call_type="completion", + ) + + assert _find_descriptor(descriptors, "mcp_per_key") is None + + +def test_mcp_descriptor_skipped_for_raw_rest_body_v3(): + handler, _ = _make_mcp_handler() + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-mcp-key"), + team_id="team-1", + metadata={"mcp_rpm_limit": {"github": 5}}, + team_metadata={"mcp_rpm_limit": {"github": 3}}, + ) + + descriptors = _build_mcp_descriptors( + handler, + user_api_key_dict, + { + "server_id": "slack", + "name": "demo-tool", + "arguments": {}, + "mcp_server_name": "github", + }, + ) + + assert _find_descriptor(descriptors, "mcp_per_key") is None + assert _find_descriptor(descriptors, "mcp_per_team") is None + + +def test_mcp_per_team_descriptor_created_from_team_metadata_v3(): + handler, _ = _make_mcp_handler() + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-mcp-key"), + team_id="team-1", + team_metadata={"mcp_rpm_limit": {"github": 3}}, + ) + + descriptors = _build_mcp_descriptors( + handler, user_api_key_dict, {"mcp_server_name": "github"} + ) + + descriptor = _find_descriptor(descriptors, "mcp_per_team") + assert descriptor is not None + assert descriptor["value"] == "team-1:github" + assert descriptor["rate_limit"]["requests_per_unit"] == 3 + assert descriptor["rate_limit"]["tokens_per_unit"] is None + + +@pytest.mark.asyncio +async def test_mcp_per_key_rpm_enforced_v3(monkeypatch): + """ + A key configured with mcp_rpm_limit={"github": 2} must allow 2 calls to the + github MCP server within the window and reject the 3rd with a 429, while + calls to a different MCP server are unaffected. + """ + monkeypatch.setenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", "60") + api_key = hash_token("sk-mcp-enforce") + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + + window_starts: Dict[str, int] = {} + request_counts: Dict[str, int] = {} + + async def mock_batch_rate_limiter(*args, **kwargs): + keys = kwargs.get("keys") if kwargs else args[0] + args_list = kwargs.get("args") if kwargs else args[1] + now = args_list[0] + window_size = args_list[1] + results = [] + for i in range(0, len(keys), 2): + window_key = keys[i] + counter_key = keys[i + 1] + prev_window = window_starts.get(window_key) + prev_counter = request_counts.get(counter_key, 0) + if prev_window is None or (now - prev_window) >= window_size: + window_starts[window_key] = now + new_counter = 1 + else: + new_counter = prev_counter + 1 + request_counts[counter_key] = new_counter + results.append(now) + results.append(new_counter) + return results + + handler.batch_rate_limiter_script = mock_batch_rate_limiter + + user_api_key_dict = UserAPIKeyAuth( + api_key=api_key, + metadata={"mcp_rpm_limit": {"github": 2}}, + ) + + for _ in range(2): + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data={"mcp_server_name": "github"}, + call_type="call_mcp_tool", + ) + + with pytest.raises(HTTPException) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data={"mcp_server_name": "github"}, + call_type="call_mcp_tool", + ) + assert exc_info.value.status_code == 429 + + # A different server has no configured limit -> not rate limited. + for _ in range(5): + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data={"mcp_server_name": "slack"}, + call_type="call_mcp_tool", + ) + + # The TPM counter must never be created for an MCP descriptor. + assert not any(":tokens" in key and "github" in key for key in request_counts) + + +def test_get_key_mcp_rpm_limit_precedence(): + from litellm.proxy.auth.auth_utils import ( + get_key_mcp_rpm_limit, + get_team_mcp_rpm_limit, + ) + + # Key metadata takes precedence over team metadata. + key_first = UserAPIKeyAuth( + api_key=hash_token("sk-mcp-key"), + metadata={"mcp_rpm_limit": {"github": 10}}, + team_metadata={"mcp_rpm_limit": {"github": 99}}, + ) + assert get_key_mcp_rpm_limit(key_first) == {"github": 10} + + # Falls back to team metadata when key has none. + team_only = UserAPIKeyAuth( + api_key=hash_token("sk-mcp-key"), + team_metadata={"mcp_rpm_limit": {"github": 7}}, + ) + assert get_key_mcp_rpm_limit(team_only) == {"github": 7} + assert get_team_mcp_rpm_limit(team_only) == {"github": 7} + + # No configuration anywhere. + none_set = UserAPIKeyAuth(api_key=hash_token("sk-mcp-key")) + assert get_key_mcp_rpm_limit(none_set) is None + assert get_team_mcp_rpm_limit(none_set) is None diff --git a/tests/test_litellm/proxy/management_endpoints/search_endpoints/test_search_tool_management.py b/tests/test_litellm/proxy/management_endpoints/search_endpoints/test_search_tool_management.py index 55b4181e92e..ea7e5591f18 100644 --- a/tests/test_litellm/proxy/management_endpoints/search_endpoints/test_search_tool_management.py +++ b/tests/test_litellm/proxy/management_endpoints/search_endpoints/test_search_tool_management.py @@ -1,3 +1,4 @@ +import contextlib import os import sys from datetime import datetime @@ -10,7 +11,12 @@ sys.path.insert( 0, os.path.abspath("../../../../..") ) # Adds the parent directory to the system path -from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy._types import ( + LiteLLM_ObjectPermissionTable, + LiteLLM_TeamTable, + LitellmUserRoles, + UserAPIKeyAuth, +) # Import proxy_server module first to ensure it's initialized import litellm.proxy.proxy_server as ps @@ -603,3 +609,170 @@ async def test_list_search_tools_db_masking_sensitive_values(monkeypatch): assert tool4["litellm_params"]["search_provider"] == "custom" finally: app.dependency_overrides.pop(user_api_key_auth, None) + + +@contextlib.contextmanager +def _mock_search_tool_backend(db_tools): + """Patch the DB registry, prisma client, and config so /search_tools/list + returns exactly ``db_tools`` (no config-defined tools).""" + mock_registry = MagicMock() + mock_registry.get_all_search_tools_from_db = AsyncMock(return_value=db_tools) + mock_proxy_config = MagicMock() + mock_proxy_config.get_config = AsyncMock(return_value={}) + mock_proxy_config.parse_search_tools = MagicMock(return_value=None) + with ( + patch( + "litellm.proxy.search_endpoints.search_tool_management.SEARCH_TOOL_REGISTRY", + mock_registry, + ), + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config), + ): + yield + + +def _scoping_db_tools(): + return [ + { + "search_tool_id": "db-id-1", + "search_tool_name": "db-tool-1", + "litellm_params": { + "search_provider": "perplexity", + "api_key": "pplx-secret-1", + "api_base": "https://api.perplexity.ai", + }, + "search_tool_info": {"description": "Perplexity"}, + "created_at": datetime(2024, 1, 1), + "updated_at": datetime(2024, 1, 1), + }, + { + "search_tool_id": "db-id-2", + "search_tool_name": "db-tool-2", + "litellm_params": { + "search_provider": "tavily", + "api_key": "tvly-secret-2", + "api_base": "https://api.tavily.com", + }, + "search_tool_info": {"description": "Tavily"}, + "created_at": datetime(2024, 1, 1), + "updated_at": datetime(2024, 1, 1), + }, + { + "search_tool_id": "db-id-3", + "search_tool_name": "db-tool-3", + "litellm_params": {"search_provider": "exa", "api_key": "exa-secret-3"}, + "search_tool_info": {"description": "Exa"}, + "created_at": datetime(2024, 1, 1), + "updated_at": datetime(2024, 1, 1), + }, + ] + + +@contextlib.contextmanager +def _override_auth(user): + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + app.dependency_overrides[user_api_key_auth] = lambda: user + try: + yield + finally: + app.dependency_overrides.pop(user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_list_search_tools_scoped_to_key_object_permission(): + """ + Regression: an internal user whose key is restricted to specific search tools + must only see those tools. Before the fix /search_tools/list returned every + configured tool, leaking ids, api_base, and metadata for tools it cannot call. + """ + restricted_user = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="internal_user", + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="op-key", + search_tools=["db-tool-1"], + ), + ) + + with ( + _mock_search_tool_backend(_scoping_db_tools()), + _override_auth(restricted_user), + ): + response = TestClient(app).get("/search_tools/list") + + assert response.status_code == 200 + tools = response.json()["search_tools"] + assert [t["search_tool_name"] for t in tools] == ["db-tool-1"] + leaked = {t["litellm_params"].get("api_base") for t in tools} + assert "https://api.tavily.com" not in leaked + + +@pytest.mark.asyncio +async def test_list_search_tools_unrestricted_internal_user_sees_all(): + """An internal user with no search_tools allowlist is unrestricted and sees every tool.""" + unrestricted_user = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, user_id="internal_user" + ) + + with ( + _mock_search_tool_backend(_scoping_db_tools()), + _override_auth(unrestricted_user), + ): + response = TestClient(app).get("/search_tools/list") + + assert response.status_code == 200 + names = {t["search_tool_name"] for t in response.json()["search_tools"]} + assert names == {"db-tool-1", "db-tool-2", "db-tool-3"} + + +@pytest.mark.asyncio +async def test_list_search_tools_scoped_to_team_object_permission(): + """A team-level search_tools allowlist also scopes the listing for a non-admin caller.""" + team_member = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="internal_user", + team_id="team-1", + ) + team_object = LiteLLM_TeamTable( + team_id="team-1", + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="op-team", + search_tools=["db-tool-2"], + ), + ) + + with ( + _mock_search_tool_backend(_scoping_db_tools()), + patch( + "litellm.proxy.auth.auth_checks.get_team_object", + AsyncMock(return_value=team_object), + ), + _override_auth(team_member), + ): + response = TestClient(app).get("/search_tools/list") + + assert response.status_code == 200 + assert [t["search_tool_name"] for t in response.json()["search_tools"]] == [ + "db-tool-2" + ] + + +@pytest.mark.asyncio +async def test_list_search_tools_admin_with_restricted_key_still_sees_all(): + """Proxy admins bypass search-tool scoping even if their key carries an allowlist.""" + admin_user = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + user_id="admin_user", + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="op-admin", + search_tools=["db-tool-1"], + ), + ) + + with _mock_search_tool_backend(_scoping_db_tools()), _override_auth(admin_user): + response = TestClient(app).get("/search_tools/list") + + assert response.status_code == 200 + names = {t["search_tool_name"] for t in response.json()["search_tools"]} + assert names == {"db-tool-1", "db-tool-2", "db-tool-3"} diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_utils.py b/tests/test_litellm/proxy/management_endpoints/test_common_utils.py index f898763d2cb..d53ea6fa34d 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_utils.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_utils.py @@ -482,6 +482,33 @@ class TestSetObjectMetadataField: _set_object_metadata_field(team, "model_rpm_limit", {"x": 1}) assert team.metadata == {"model_rpm_limit": {"x": 1}} + def test_mcp_rpm_limit_is_hoisted_into_metadata(self): + """ + Per-MCP-server rpm limits are stored in the metadata JSON column, not a + dedicated DB column. The key/team management endpoints rely on + LiteLLM_ManagementEndpoint_MetadataFields to move the request field into + metadata; this regression guards that mcp_rpm_limit is in that list and + round-trips through the same loop the endpoints use. + """ + from litellm.proxy._types import LiteLLM_ManagementEndpoint_MetadataFields + + assert "mcp_rpm_limit" in LiteLLM_ManagementEndpoint_MetadataFields + + from types import SimpleNamespace + + team = LiteLLM_TeamTable(team_id="t1", metadata={}) + mcp_rpm_limit = {"github": 100} + data = SimpleNamespace(mcp_rpm_limit=mcp_rpm_limit) + + with patch( + "litellm.proxy.management_endpoints.common_utils._premium_user_check" + ): + for field in LiteLLM_ManagementEndpoint_MetadataFields: + if getattr(data, field, None) is not None: + _set_object_metadata_field(team, field, getattr(data, field)) + + assert team.metadata["mcp_rpm_limit"] == mcp_rpm_limit + class TestRequireCallerUserIdForNonAdmin: """ 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 8fb242372e3..cda22da6ebd 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 @@ -3094,6 +3094,249 @@ async def test_generate_key_with_object_permission(): assert "object_permission" not in key_data +@pytest.mark.asyncio +async def test_generate_key_team_member_inherits_org_skips_membership_check(): + """Regression: a team member creating a key for an org-scoped team must not + be blocked by the org-membership check. + + When ``organization_id`` is inherited from the key's team (via + ``apply_enterprise_key_management_params`` -> ``add_team_organization_id``), + the caller already passed team-level authorization. Requiring an explicit + ``LiteLLM_OrganizationMembership`` row on top of that broke the normal admin + workflow (admins only add users to teams). This asserts the org-membership + check is skipped when the org id came from the caller's team. + """ + from unittest.mock import AsyncMock, MagicMock, patch + + from litellm.proxy._types import GenerateKeyRequest, LitellmUserRoles + from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _common_key_generation_helper, + ) + + org_id = "org-from-team" + + # Team belongs to an org; caller is a team member but NOT an explicit member + # of that organization (the regression scenario). + mock_team_table = MagicMock() + mock_team_table.organization_id = org_id + mock_team_table.metadata = None + + mock_validate_org = AsyncMock() + mock_generate_key = AsyncMock( + return_value={ + "key": "sk-test-key", + "expires": None, + "user_id": "alice", + "team_id": "team-1", + } + ) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.llm_router", None), + patch("litellm.proxy.proxy_server.premium_user", True), + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.validate_key_mcp_servers_against_team", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.validate_key_search_tools_against_team", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._validate_caller_can_assign_key_org", + mock_validate_org, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.get_org_object", + new_callable=AsyncMock, + return_value=MagicMock(litellm_budget_table=None), + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._check_org_key_limits", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn", + mock_generate_key, + ), + ): + result = await _common_key_generation_helper( + data=GenerateKeyRequest( + user_id="alice", + team_id="team-1", + organization_id=org_id, + ), + user_api_key_dict=UserAPIKeyAuth( + user_id="alice", + user_role=LitellmUserRoles.INTERNAL_USER.value, + ), + litellm_changed_by=None, + team_table=mock_team_table, + ) + + # Key creation proceeded for the team member ... + mock_generate_key.assert_awaited_once() + assert result is not None + # ... and the org-membership check was bypassed because organization_id was + # inherited from the caller's team. + mock_validate_org.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_generate_key_foreign_org_without_team_still_enforces_membership(): + """VERIA-55: a caller assigning a key to an organization that was NOT + inherited from a team must still pass the org-membership check. + + This guards the IDOR fix: ``team_table is None`` (or an org id that does not + match the team) means the org id did not come from team context, so the + explicit membership validation must run. + """ + from unittest.mock import AsyncMock, MagicMock, patch + + from litellm.proxy._types import GenerateKeyRequest, LitellmUserRoles + from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _common_key_generation_helper, + ) + + foreign_org_id = "someone-elses-org" + + mock_validate_org = AsyncMock() + mock_generate_key = AsyncMock( + return_value={ + "key": "sk-test-key", + "expires": None, + "user_id": "alice", + "team_id": None, + } + ) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.llm_router", None), + patch("litellm.proxy.proxy_server.premium_user", True), + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._validate_caller_can_assign_key_org", + mock_validate_org, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.get_org_object", + new_callable=AsyncMock, + return_value=MagicMock(litellm_budget_table=None), + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._check_org_key_limits", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn", + mock_generate_key, + ), + ): + await _common_key_generation_helper( + data=GenerateKeyRequest( + user_id="alice", + organization_id=foreign_org_id, + ), + user_api_key_dict=UserAPIKeyAuth( + user_id="alice", + user_role=LitellmUserRoles.INTERNAL_USER.value, + ), + litellm_changed_by=None, + team_table=None, + ) + + # No team context -> the org-membership check must still run. + mock_validate_org.assert_awaited_once() + assert mock_validate_org.call_args.kwargs["organization_id"] == foreign_org_id + + +@pytest.mark.asyncio +async def test_generate_key_foreign_org_with_mismatched_team_still_enforces_membership(): + """VERIA-55: when a team is present but its organization_id differs from the + organization_id on the key request, the org-membership check must still run.""" + from unittest.mock import AsyncMock, MagicMock, patch + + from litellm.proxy._types import GenerateKeyRequest, LitellmUserRoles + from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _common_key_generation_helper, + ) + + team_org_id = "other-org" + foreign_org_id = "someone-elses-org" + + mock_team_table = MagicMock() + mock_team_table.organization_id = team_org_id + mock_team_table.metadata = None + + mock_validate_org = AsyncMock() + mock_generate_key = AsyncMock( + return_value={ + "key": "sk-test-key", + "expires": None, + "user_id": "alice", + "team_id": "team-1", + } + ) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.llm_router", None), + patch("litellm.proxy.proxy_server.premium_user", True), + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.validate_key_mcp_servers_against_team", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.validate_key_search_tools_against_team", + new_callable=AsyncMock, + ), + patch( + "litellm_enterprise.proxy.management_endpoints.key_management_endpoints.apply_enterprise_key_management_params", + side_effect=lambda data, team_table: data, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._validate_caller_can_assign_key_org", + mock_validate_org, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.get_org_object", + new_callable=AsyncMock, + return_value=MagicMock(litellm_budget_table=None), + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._check_org_key_limits", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn", + mock_generate_key, + ), + ): + await _common_key_generation_helper( + data=GenerateKeyRequest( + user_id="alice", + team_id="team-1", + organization_id=foreign_org_id, + ), + user_api_key_dict=UserAPIKeyAuth( + user_id="alice", + user_role=LitellmUserRoles.INTERNAL_USER.value, + ), + litellm_changed_by=None, + team_table=mock_team_table, + ) + + mock_validate_org.assert_awaited_once() + assert mock_validate_org.call_args.kwargs["organization_id"] == foreign_org_id + + # ============================================ # Organization Key Limit Tests # ============================================ diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index 5d66c184495..a5c8320a5cc 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -1747,6 +1747,37 @@ class TestTemporaryMCPSessionEndpoints: assert isinstance(result, UserAPIKeyAuth) auth_builder_mock.assert_not_called() + def test_mcp_oauth_authorize_token_routes_use_browser_auth_dependency(self): + from fastapi.routing import APIRoute + + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + _mcp_oauth_user_api_key_auth, + router, + ) + + oauth_routes = { + route.path: route + for route in router.routes + if isinstance(route, APIRoute) + and route.path + in { + "/v1/mcp/server/oauth/{server_id}/authorize", + "/v1/mcp/server/oauth/{server_id}/token", + } + } + + assert set(oauth_routes) == { + "/v1/mcp/server/oauth/{server_id}/authorize", + "/v1/mcp/server/oauth/{server_id}/token", + } + for route in oauth_routes.values(): + dependency_names = { + dependant.name + for dependant in route.dependant.dependencies + if dependant.call is _mcp_oauth_user_api_key_auth + } + assert dependency_names == {None, "user_api_key_dict"} + @pytest.mark.asyncio async def test_mcp_authorize_proxies_to_discoverable_endpoint(self): from litellm.proxy.management_endpoints.mcp_management_endpoints import ( diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_carry_guardrail_logging_info.py b/tests/test_litellm/proxy/pass_through_endpoints/test_carry_guardrail_logging_info.py new file mode 100644 index 00000000000..3071812e117 --- /dev/null +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_carry_guardrail_logging_info.py @@ -0,0 +1,68 @@ +"""Unit tests for ``_carry_guardrail_logging_info``. + +This is the helper that lets a passthrough guardrail block still surface its otel +span: it copies ``standard_logging_guardrail_information`` from the post-call +guardrail's (otherwise discarded) ``hook_data`` onto the dict the failure handler +forwards to ``post_call_failure_hook``. No otel dependency here, so these run +everywhere and pin the helper's contract directly. +""" + +from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( + _carry_guardrail_logging_info, +) + +_ENTRY = {"guardrail_name": "block-demo", "guardrail_status": "guardrail_intervened"} + + +def _source(entries): + return {"metadata": {"standard_logging_guardrail_information": entries}} + + +def test_carries_entries_onto_request_without_metadata(): + request_data: dict = {} + _carry_guardrail_logging_info(request_data, _source([_ENTRY])) + assert request_data["metadata"]["standard_logging_guardrail_information"] == [ + _ENTRY + ] + + +def test_carried_list_is_copied_not_shared(): + source = _source([_ENTRY]) + request_data: dict = {} + _carry_guardrail_logging_info(request_data, source) + carried = request_data["metadata"]["standard_logging_guardrail_information"] + assert carried is not source["metadata"]["standard_logging_guardrail_information"] + carried.append({"guardrail_name": "other"}) + assert source["metadata"]["standard_logging_guardrail_information"] == [_ENTRY] + + +def test_existing_metadata_without_guardrail_key_is_populated(): + request_data: dict = {"metadata": {"user_api_key": "sk-x"}} + _carry_guardrail_logging_info(request_data, _source([_ENTRY])) + assert request_data["metadata"]["user_api_key"] == "sk-x" + assert request_data["metadata"]["standard_logging_guardrail_information"] == [ + _ENTRY + ] + + +def test_existing_guardrail_entries_are_not_clobbered(): + existing = [{"guardrail_name": "already-logged"}] + request_data = {"metadata": {"standard_logging_guardrail_information": existing}} + _carry_guardrail_logging_info(request_data, _source([_ENTRY])) + assert ( + request_data["metadata"]["standard_logging_guardrail_information"] is existing + ) + + +def test_noop_when_guardrail_data_is_none(): + request_data: dict = {} + _carry_guardrail_logging_info(request_data, None) + assert request_data == {} + + +def test_noop_when_no_guardrail_entries(): + request_data: dict = {} + _carry_guardrail_logging_info(request_data, {"metadata": {}}) + _carry_guardrail_logging_info(request_data, _source([])) + _carry_guardrail_logging_info(request_data, {}) + assert request_data == {} diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_guardrail_block_otel_span.py b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_guardrail_block_otel_span.py new file mode 100644 index 00000000000..73927e92c15 --- /dev/null +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_guardrail_block_otel_span.py @@ -0,0 +1,191 @@ +"""Regression: a guardrail block on a passthrough endpoint must still emit the +otel guardrail span. + +The span is emitted from the guardrail-recording path the moment a guardrail +finishes (``add_standard_logging_guardrail_information_to_request_data`` -> +``emit_guardrail_span``), routed through the proxy's registered otel V2 logger, +rather than from a post-call hook that does not fire on every path. A block +raises out of the post-call hook before any later hook runs, so the recording +path is the only place the span is reliably produced. These tests drive the real +``pass_through_request`` with a real ``ProxyLogging`` + a real otel V2 logger +registered as the proxy's ``open_telemetry_logger`` and assert the span is +emitted on both allow and block. +""" + +import json +from contextlib import ExitStack +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest +from fastapi import HTTPException + +pytest.importorskip("opentelemetry") + +from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( # noqa: E402 + InMemorySpanExporter, +) + +import litellm # noqa: E402 +from litellm.caching.dual_cache import DualCache # noqa: E402 +from litellm.integrations.custom_guardrail import ( # noqa: E402 + CustomGuardrail, + log_guardrail_information, +) +from litellm.integrations.otel.logger import OpenTelemetryV2 # noqa: E402 +from litellm.integrations.otel.model.config import OpenTelemetryV2Config # noqa: E402 +from litellm.integrations.otel.plumbing import providers # noqa: E402 +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache # noqa: E402 +from litellm.proxy.utils import ProxyLogging # noqa: E402 +from litellm.types.guardrails import GuardrailEventHooks # noqa: E402 + +_PT_MOD = "litellm.proxy.pass_through_endpoints.pass_through_endpoints" +_COLLECT = ( + "litellm.proxy.pass_through_endpoints.passthrough_guardrails." + "PassthroughGuardrailHandler.collect_guardrails" +) +_GUARDRAIL_SPAN = "execute_guardrail block-demo" +_TRIGGER = "BLOCKME" + +# pass_through_endpoints imports proxy_server lazily (inside the request +# function), so importing this at module scope does not require the real +# proxy_server and does not mutate sys.modules. +from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( # noqa: E402 + pass_through_request, +) + + +class _BlockOnTextGuardrail(CustomGuardrail): + """Denies (HTTP 400) when the response carries the trigger word; records its + standard guardrail logging info on both allow and block via the decorator.""" + + @log_guardrail_information + async def async_post_call_success_hook(self, data, user_api_key_dict, response): + if _TRIGGER in json.dumps(response): + raise HTTPException( + status_code=400, detail={"error": "blocked by block-demo guardrail"} + ) + return response + + +def _user_api_key_dict(): + d = MagicMock() + d.api_key = "sk-test" + d.user_id = "user-1" + d.team_id = "team-1" + d.org_id = None + d.metadata = {} + d.team_metadata = {} + d.parent_otel_span = None + d.request_route = "/mock/echo" + return d + + +def _mock_request(): + r = MagicMock() + r.method = "POST" + r.query_params = {} + r.url = "http://testserver/mock/echo" + headers = MagicMock() + headers.copy.return_value = {} + r.headers = headers + return r + + +def _httpx_response(text: str) -> httpx.Response: + body = {"candidates": [{"content": {"role": "model", "parts": [{"text": text}]}}]} + return httpx.Response( + status_code=200, + headers={"content-type": "application/json"}, + content=json.dumps(body).encode("utf-8"), + request=httpx.Request("POST", "https://upstream.example/echo"), + ) + + +def _otel_logger_with_exporter(): + cfg = OpenTelemetryV2Config(exporter="in_memory") + exporter = InMemorySpanExporter() + tracer_provider = providers.build_tracer_provider(cfg, exporter=exporter) + return OpenTelemetryV2(config=cfg, tracer_provider=tracer_provider), exporter + + +def _guardrail_span_names(exporter): + return [ + s.name + for s in exporter.get_finished_spans() + if s.name.startswith("execute_guardrail") + ] + + +async def _drive(response_text: str): + """Run the real pass_through_request with the block-demo guardrail + otel V2 + logger registered, returning (status_code, guardrail_span_names).""" + otel, exporter = _otel_logger_with_exporter() + guardrail = _BlockOnTextGuardrail( + guardrail_name="block-demo", event_hook=[GuardrailEventHooks.post_call] + ) + proxy_logging = ProxyLogging(user_api_key_cache=UserApiKeyCache(DualCache())) + + saved_callbacks = list(litellm.callbacks) + litellm.callbacks = [guardrail, otel] + + mock_async_client_obj = MagicMock() + mock_async_client_obj.client = AsyncMock() + mock_pt_logging = MagicMock() + mock_pt_logging.pass_through_async_success_handler = AsyncMock() + + patches = [ + patch( + f"{_PT_MOD}.HttpPassThroughEndpointHelpers.non_streaming_http_request_handler", + new_callable=AsyncMock, + return_value=_httpx_response(response_text), + ), + patch(f"{_PT_MOD}._is_streaming_response", return_value=False), + patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging), + patch("litellm.proxy.proxy_server.open_telemetry_logger", otel), + patch("litellm.proxy.proxy_server.llm_router", None), + patch(f"{_PT_MOD}.pass_through_endpoint_logging", mock_pt_logging), + patch(f"{_PT_MOD}.get_async_httpx_client", return_value=mock_async_client_obj), + patch(f"{_PT_MOD}._read_request_body", new_callable=AsyncMock, return_value={}), + patch(f"{_PT_MOD}._safe_get_request_headers", return_value={}), + patch(_COLLECT, return_value=["block-demo"]), + ] + try: + with ExitStack() as stack: + for p in patches: + stack.enter_context(p) + try: + result = await pass_through_request( + request=_mock_request(), + target="https://upstream.example/echo", + custom_headers={"Content-Type": "application/json"}, + user_api_key_dict=_user_api_key_dict(), + stream=False, + ) + # A deny (HTTP 4xx) re-raises as ProxyException; an allow returns + # the upstream Response. + status_code = result.status_code + except Exception as e: + status_code = getattr(e, "code", None) or getattr( + e, "status_code", None + ) + return int(status_code), _guardrail_span_names(exporter) + finally: + litellm.callbacks = saved_callbacks + + +@pytest.mark.asyncio +async def test_guardrail_block_emits_otel_guardrail_span(): + status_code, span_names = await _drive(f"{_TRIGGER} please") + assert status_code == 400 + assert span_names == [_GUARDRAIL_SPAN], ( + "guardrail span must be emitted when a passthrough guardrail blocks, " + f"got spans: {span_names}" + ) + + +@pytest.mark.asyncio +async def test_guardrail_allow_emits_otel_guardrail_span(): + status_code, span_names = await _drive("hello world") + assert status_code == 200 + assert span_names == [_GUARDRAIL_SPAN] diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_post_call_guardrails.py b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_post_call_guardrails.py index f061434a971..eafe71e1063 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_post_call_guardrails.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_post_call_guardrails.py @@ -12,6 +12,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest +from fastapi import HTTPException from litellm.integrations.custom_guardrail import ( CustomGuardrail, @@ -216,6 +217,53 @@ class TestPassthroughPostCallGuardrails: assert body["error"]["guardrail_name"] == "rubrik" assert body["error"]["model"] == "gemini-2.0-flash" + @patch(_COLLECT, return_value=["rubrik"]) + async def test_deny_forwards_guardrail_logging_info_to_failure_hook( + self, + mock_collect, + ): + """A post-call guardrail deny (non-ModifyResponseException) records its + standard_logging_guardrail_information on the hook_data dict; the failure + handler must forward that info to post_call_failure_hook so downstream + loggers (e.g. the otel guardrail span) still see it. Regression for the + block path dropping it.""" + mock_response = _make_httpx_response(_GEMINI_RESPONSE) + + def _block(*, data, user_api_key_dict, response): + metadata = data.setdefault("metadata", {}) + metadata.setdefault("standard_logging_guardrail_information", []).append( + {"guardrail_name": "rubrik", "guardrail_status": "guardrail_intervened"} + ) + raise HTTPException(status_code=400, detail={"error": "blocked"}) + + captured = {} + + async def _capture_failure(**kwargs): + captured.update(kwargs) + + mock_proxy_logging = MagicMock() + mock_proxy_logging.pre_call_hook = AsyncMock(return_value={}) + mock_proxy_logging.post_call_success_hook = AsyncMock(side_effect=_block) + mock_proxy_logging.post_call_failure_hook = AsyncMock( + side_effect=_capture_failure + ) + + with _common_patches(mock_proxy_logging, mock_response): + with pytest.raises(Exception): + await pass_through_request( + request=_make_mock_request(), + target="https://example.com/v1/generateContent", + custom_headers={"Content-Type": "application/json"}, + user_api_key_dict=_make_user_api_key_dict(), + stream=False, + ) + + mock_proxy_logging.post_call_failure_hook.assert_awaited_once() + entries = captured["request_data"]["metadata"][ + "standard_logging_guardrail_information" + ] + assert any(e.get("guardrail_name") == "rubrik" for e in entries) + @pytest.mark.asyncio class TestUnifiedGuardrailCallTypeResolution: diff --git a/tests/test_litellm/proxy/shutdown/test_graceful_shutdown_manager.py b/tests/test_litellm/proxy/shutdown/test_graceful_shutdown_manager.py new file mode 100644 index 00000000000..d38852617b9 --- /dev/null +++ b/tests/test_litellm/proxy/shutdown/test_graceful_shutdown_manager.py @@ -0,0 +1,181 @@ +""" +Tests for GracefulShutdownManager. + +These verify the drain logic that lets a pod terminate as soon as its real +in-flight work is done (bounded by GRACEFUL_SHUTDOWN_TIMEOUT) rather than +sleeping for a fixed worst-case duration. +""" + +import time + +import pytest + +from litellm.proxy.shutdown.graceful_shutdown_manager import ( + DEFAULT_GRACEFUL_SHUTDOWN_TIMEOUT, + GracefulShutdownManager, +) + + +@pytest.fixture(autouse=True) +def _reset(): + GracefulShutdownManager.reset() + yield + GracefulShutdownManager.reset() + + +def _counter_that_drains_after(calls_before_zero: int): + """Return a count_fn that reports N in-flight until it has been polled + `calls_before_zero` times, then reports 0.""" + state = {"polls": 0} + + def count_fn() -> int: + state["polls"] += 1 + return 0 if state["polls"] > calls_before_zero else 3 + + return count_fn + + +# ── shutdown flag ─────────────────────────────────────────────────────────── + + +def test_not_shutting_down_by_default(): + assert GracefulShutdownManager.is_shutting_down() is False + + +def test_start_shutdown_sets_flag(): + GracefulShutdownManager.start_shutdown() + assert GracefulShutdownManager.is_shutting_down() is True + + +def test_start_shutdown_is_idempotent_and_does_not_reset_clock(): + GracefulShutdownManager.start_shutdown() + first = GracefulShutdownManager._shutdown_started_at + time.sleep(0.01) + GracefulShutdownManager.start_shutdown() + assert GracefulShutdownManager._shutdown_started_at == first + + +def test_reset_clears_flag(): + GracefulShutdownManager.start_shutdown() + GracefulShutdownManager.reset() + assert GracefulShutdownManager.is_shutting_down() is False + + +# ── timeout config ──────────────────────────────────────────────────────────── + + +def test_timeout_defaults_when_unset(monkeypatch): + monkeypatch.delenv("GRACEFUL_SHUTDOWN_TIMEOUT", raising=False) + assert GracefulShutdownManager.get_timeout() == DEFAULT_GRACEFUL_SHUTDOWN_TIMEOUT + + +def test_timeout_reads_env(monkeypatch): + monkeypatch.setenv("GRACEFUL_SHUTDOWN_TIMEOUT", "5") + assert GracefulShutdownManager.get_timeout() == 5.0 + + +def test_timeout_falls_back_on_garbage(monkeypatch): + monkeypatch.setenv("GRACEFUL_SHUTDOWN_TIMEOUT", "not-a-number") + assert GracefulShutdownManager.get_timeout() == DEFAULT_GRACEFUL_SHUTDOWN_TIMEOUT + + +# ── wait_for_drain ──────────────────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_returns_immediately_when_already_drained(): + start = time.monotonic() + drained = await GracefulShutdownManager.wait_for_drain( + timeout=10, count_fn=lambda: 0 + ) + assert drained == 0 + assert time.monotonic() - start < 0.5 + + +@pytest.mark.asyncio +async def test_waits_until_counter_reaches_zero_then_returns_drained_count(): + count_fn = _counter_that_drains_after(calls_before_zero=3) + drained = await GracefulShutdownManager.wait_for_drain( + timeout=10, count_fn=count_fn + ) + assert drained == 3 + + +@pytest.mark.asyncio +async def test_times_out_when_counter_never_drains(): + start = time.monotonic() + drained = await GracefulShutdownManager.wait_for_drain( + timeout=0.3, count_fn=lambda: 2 + ) + elapsed = time.monotonic() - start + assert 0.3 <= elapsed < 2.0 + assert drained == 0 + + +@pytest.mark.asyncio +async def test_zero_timeout_does_not_block(): + start = time.monotonic() + drained = await GracefulShutdownManager.wait_for_drain( + timeout=0, count_fn=lambda: 5 + ) + assert time.monotonic() - start < 0.2 + assert drained == 5 + + +@pytest.mark.asyncio +async def test_exclude_self_treats_one_inflight_as_drained(): + """The /health/drain request counts itself, so a steady count of 1 must be + treated as fully drained rather than timing out.""" + start = time.monotonic() + drained = await GracefulShutdownManager.wait_for_drain( + timeout=5, exclude_self=True, count_fn=lambda: 1 + ) + assert time.monotonic() - start < 0.5 + assert drained == 0 + + +@pytest.mark.asyncio +async def test_without_exclude_self_one_inflight_blocks_until_timeout(): + start = time.monotonic() + await GracefulShutdownManager.wait_for_drain(timeout=0.3, count_fn=lambda: 1) + assert time.monotonic() - start >= 0.3 + + +@pytest.mark.asyncio +async def test_defaults_to_get_timeout_and_live_counter(monkeypatch): + """With no timeout/count_fn passed, it falls back to get_timeout() and the + live InFlightRequestsMiddleware counter.""" + from litellm.proxy.middleware.in_flight_requests_middleware import ( + InFlightRequestsMiddleware, + ) + + monkeypatch.delenv("GRACEFUL_SHUTDOWN_TIMEOUT", raising=False) + InFlightRequestsMiddleware._in_flight = 0 + drained = await GracefulShutdownManager.wait_for_drain() + assert drained == 0 + + +@pytest.mark.asyncio +async def test_second_drain_is_a_noop_so_window_is_not_doubled(): + """preStop /health/drain and the lifespan SIGTERM handler both drain; the + second call must return immediately rather than waiting another full + timeout (which would require doubling terminationGracePeriodSeconds).""" + await GracefulShutdownManager.wait_for_drain(timeout=0.2, count_fn=lambda: 1) + + start = time.monotonic() + drained = await GracefulShutdownManager.wait_for_drain( + timeout=5, count_fn=lambda: 1 + ) + assert time.monotonic() - start < 0.1 + assert drained == 0 + + +@pytest.mark.asyncio +async def test_emits_periodic_drain_waiting_log_while_waiting(): + """With a zero log interval, the periodic drain_waiting branch runs on each + poll until the counter finally drains.""" + count_fn = _counter_that_drains_after(calls_before_zero=2) + drained = await GracefulShutdownManager.wait_for_drain( + timeout=10, count_fn=count_fn, poll_interval=0, log_interval=0 + ) + assert drained == 3 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 5ca058fc8d9..0c7511589de 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 @@ -300,6 +300,25 @@ def test_get_messages_for_spend_logs_realtime_returns_messages(mock_should_store assert parsed[1]["content"] == "What is the weather today?" +@patch( + "litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs" +) +def test_get_messages_for_spend_logs_strips_null_bytes(mock_should_store): + """Regression for PostgreSQL 22P05: NUL bytes must be stripped from messages.""" + mock_should_store.return_value = True + payload = cast( + StandardLoggingPayload, + { + "call_type": "_arealtime", + "messages": [{"role": "user", "content": "hello\x00world"}], + }, + ) + result = _get_messages_for_spend_logs_payload(payload) + assert "\\u0000" not in result + parsed = json.loads(result) + assert parsed[0]["content"] == "helloworld" + + @patch( "litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs" ) @@ -370,6 +389,21 @@ def test_get_response_for_spend_logs_payload_truncates_large_base64(mock_should_ assert parsed["data"][0]["other_field"] == "value" +@patch( + "litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs" +) +def test_get_response_for_spend_logs_payload_strips_null_bytes(mock_should_store): + """Regression for PostgreSQL 22P05: NUL bytes must be stripped from response.""" + mock_should_store.return_value = True + payload = cast( + StandardLoggingPayload, + {"response": {"content": "answer\x00here"}}, + ) + response_json = _get_response_for_spend_logs_payload(payload) + assert "\\u0000" not in response_json + assert json.loads(response_json)["content"] == "answerhere" + + @patch( "litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs" ) @@ -936,6 +970,36 @@ def test_get_logging_payload_includes_overhead_in_spend_logs_metadata(): ), f"Expected overhead '{test_overhead_ms}', got '{metadata.get('litellm_overhead_time_ms')}'" +@patch("litellm.proxy.proxy_server.master_key", None) +@patch("litellm.proxy.proxy_server.general_settings", {}) +def test_get_logging_payload_strips_null_bytes_from_request_tags(): + """Regression for PostgreSQL 22P05: NUL bytes must be stripped from request_tags.""" + kwargs = { + "model": "gpt-3.5-turbo", + "litellm_params": { + "metadata": { + "user_api_key": "sk-test-key", + "tags": ["clean-tag", "bad\x00tag"], + } + }, + } + + start_time = datetime.datetime.now(timezone.utc) + end_time = datetime.datetime.now(timezone.utc) + + payload = get_logging_payload( + kwargs=kwargs, + response_obj={}, + start_time=start_time, + end_time=end_time, + ) + + request_tags = payload.get("request_tags") + assert request_tags is not None + assert "\\u0000" not in request_tags + assert json.loads(request_tags) == ["clean-tag", "badtag"] + + @patch("litellm.proxy.proxy_server.master_key", None) @patch("litellm.proxy.proxy_server.general_settings", {}) def test_get_logging_payload_handles_missing_overhead_gracefully(): diff --git a/tests/test_litellm/proxy/test_proxy_types.py b/tests/test_litellm/proxy/test_proxy_types.py index dbb952968ed..b4352446382 100644 --- a/tests/test_litellm/proxy/test_proxy_types.py +++ b/tests/test_litellm/proxy/test_proxy_types.py @@ -87,3 +87,22 @@ def test_user_api_key_auth_hashes_authorization_header_form_of_key(): assert from_header.api_key == baseline.api_key assert from_header.token == baseline.token assert not from_header.api_key.lower().startswith("bearer") + + +def test_proxy_exception_str_returns_message(): + """ProxyException must stringify to its message: OTEL's + ``span.record_exception`` and ``str(exc)``-based logging read the string + form, which was empty pre-fix. The OpenAI-mapped fields must stay intact.""" + from litellm.proxy._types import ProxyException + + msg = "Authentication Error, Invalid proxy server token passed." + exc = ProxyException(message=msg, type="auth_error", param="key", code=401) + + assert str(exc) == msg + assert exc.message == msg + assert exc.to_dict() == { + "message": msg, + "type": "auth_error", + "param": "key", + "code": "401", + } diff --git a/tests/test_litellm/proxy/test_utils.py b/tests/test_litellm/proxy/test_utils.py deleted file mode 100644 index 9dfeb27f4cb..00000000000 --- a/tests/test_litellm/proxy/test_utils.py +++ /dev/null @@ -1,22 +0,0 @@ -import pytest - -from litellm.proxy.utils import _get_openapi_url - - -@pytest.mark.parametrize( - "env_vars, expected_url", - [ - ({}, "/openapi.json"), # default case - ({"NO_OPENAPI": "True"}, None), # OpenAPI disabled - ], -) -def test_get_openapi_url(monkeypatch, env_vars, expected_url): - # Clear relevant environment variables - monkeypatch.delenv("NO_OPENAPI", raising=False) - - # Set test environment variables - for key, value in env_vars.items(): - monkeypatch.setenv(key, value) - - result = _get_openapi_url() - assert result == expected_url diff --git a/tests/test_litellm/proxy/utils/__init__.py b/tests/test_litellm/proxy/utils/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/proxy/utils/helpers/__init__.py b/tests/test_litellm/proxy/utils/helpers/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/proxy/utils/helpers/test_error_helpers.py b/tests/test_litellm/proxy/utils/helpers/test_error_helpers.py new file mode 100644 index 00000000000..e73e3c151e0 --- /dev/null +++ b/tests/test_litellm/proxy/utils/helpers/test_error_helpers.py @@ -0,0 +1,173 @@ +import json + +import pytest +from fastapi import HTTPException + +from litellm.proxy._types import ProxyErrorTypes, ProxyException +from litellm.proxy.utils import get_error_message_str, handle_exception_on_proxy + + +def normalize(value): + return value + + +def test_get_error_message_str_happy_path_http_exception_with_string_detail(): + exc = HTTPException(status_code=400, detail="something went wrong") + summary = { + "result": get_error_message_str(exc), + "status_code": exc.status_code, + "is_str": True, + } + assert summary == { + "result": "something went wrong", + "status_code": 400, + "is_str": True, + } + + +def test_get_error_message_str_happy_path_http_exception_with_dict_detail(): + detail = {"error": "bad input", "code": "invalid_request"} + exc = HTTPException(status_code=422, detail=detail) + summary = { + "result": get_error_message_str(exc), + "result_parsed": json.loads(get_error_message_str(exc)), + "status_code": exc.status_code, + } + assert summary == { + "result": json.dumps(detail), + "result_parsed": detail, + "status_code": 422, + } + + +def test_get_error_message_str_happy_path_generic_exception(): + exc = ValueError("boom") + summary = { + "result": get_error_message_str(exc), + "type": type(exc).__name__, + "args": list(exc.args), + } + assert summary == { + "result": "boom", + "type": "ValueError", + "args": ["boom"], + } + + +def test_get_error_message_str_with_runtime_error(): + exc = RuntimeError("runtime explosion") + summary = { + "result": get_error_message_str(exc), + "type": type(exc).__name__, + "matches_str": str(exc) == get_error_message_str(exc), + } + assert summary == { + "result": "runtime explosion", + "type": "RuntimeError", + "matches_str": True, + } + + +def test_get_error_message_str_error_path_none_input_returns_string_none(): + summary = { + "result": get_error_message_str(None), + "is_str": isinstance(get_error_message_str(None), str), + "input": None, + } + assert summary == { + "result": "None", + "is_str": True, + "input": None, + } + + +def test_handle_exception_on_proxy_happy_path_http_exception(): + exc = HTTPException(status_code=403, detail="forbidden") + result = handle_exception_on_proxy(exc) + snapshot = { + "is_proxy_exception": isinstance(result, ProxyException), + "message": result.message, + "type": result.type, + "code": result.code, + } + assert snapshot == { + "is_proxy_exception": True, + "message": "forbidden", + "type": ProxyErrorTypes.internal_server_error.value, + "code": "403", + } + + +def test_handle_exception_on_proxy_happy_path_already_proxy_exception(): + original = ProxyException( + message="already wrapped", + type=ProxyErrorTypes.budget_exceeded.value, + param="key", + code=402, + ) + result = handle_exception_on_proxy(original) + snapshot = { + "is_same_object": result is original, + "message": result.message, + "type": result.type, + "code": result.code, + } + assert snapshot == { + "is_same_object": True, + "message": "already wrapped", + "type": ProxyErrorTypes.budget_exceeded.value, + "code": "402", + } + + +def test_handle_exception_on_proxy_happy_path_generic_exception_defaults_to_500(): + exc = ValueError("kaboom") + result = handle_exception_on_proxy(exc) + snapshot = { + "is_proxy_exception": isinstance(result, ProxyException), + "message": result.message, + "type": result.type, + "code": result.code, + "param": result.param, + } + assert snapshot == { + "is_proxy_exception": True, + "message": "kaboom", + "type": ProxyErrorTypes.internal_server_error.value, + "code": "500", + "param": "None", + } + + +def test_handle_exception_on_proxy_uses_attached_status_code_when_present(): + class _CustomErr(Exception): + status_code = 418 + + exc = _CustomErr("teapot") + result = handle_exception_on_proxy(exc) + snapshot = { + "code": result.code, + "message": result.message, + "type": result.type, + } + assert snapshot == { + "code": "418", + "message": "teapot", + "type": ProxyErrorTypes.internal_server_error.value, + } + + +def test_handle_exception_on_proxy_error_path_none_input_wraps_as_500(): + result = handle_exception_on_proxy(None) + snapshot = { + "is_proxy_exception": isinstance(result, ProxyException), + "message": result.message, + "code": result.code, + "type": result.type, + } + assert snapshot == { + "is_proxy_exception": True, + "message": "None", + "code": "500", + "type": ProxyErrorTypes.internal_server_error.value, + } diff --git a/tests/test_litellm/proxy/utils/helpers/test_guardrail_merge.py b/tests/test_litellm/proxy/utils/helpers/test_guardrail_merge.py new file mode 100644 index 00000000000..117484d61a4 --- /dev/null +++ b/tests/test_litellm/proxy/utils/helpers/test_guardrail_merge.py @@ -0,0 +1,201 @@ +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest + +from litellm.proxy.utils import ( + _check_and_merge_model_level_guardrails, + _merge_guardrails_with_existing, +) + + +def normalize(value): + return value + + +def _router_with_deployment(guardrails): + deployment = SimpleNamespace(litellm_params={"guardrails": guardrails}) + router = MagicMock() + router.get_deployment.return_value = deployment + return router + + +def _router_without_deployment(): + router = MagicMock() + router.get_deployment.return_value = None + return router + + +def test_check_and_merge_model_level_guardrails_happy_path_merges_lists(): + router = _router_with_deployment(["pii-redact", "toxic-filter"]) + data = { + "model": "gpt-4o", + "metadata": { + "model_info": {"id": "deployment-123"}, + "guardrails": ["user-policy"], + }, + } + result = _check_and_merge_model_level_guardrails(data, router) + snapshot = { + "model": result["model"], + "model_info_id": result["metadata"]["model_info"]["id"], + "guardrails_sorted": sorted(result["metadata"]["guardrails"]), + } + assert snapshot == { + "model": "gpt-4o", + "model_info_id": "deployment-123", + "guardrails_sorted": ["pii-redact", "toxic-filter", "user-policy"], + } + + +def test_check_and_merge_model_level_guardrails_returns_data_when_router_none(): + data = {"metadata": {"model_info": {"id": "x"}}, "model": "m", "other": 1} + result = _check_and_merge_model_level_guardrails(data, None) + assert result is data + assert normalize(result) == { + "metadata": {"model_info": {"id": "x"}}, + "model": "m", + "other": 1, + } + + +def test_check_and_merge_model_level_guardrails_returns_data_when_model_id_missing(): + router = _router_with_deployment(["pii"]) + data = {"metadata": {"model_info": {}}, "model": "m", "extra": "v"} + result = _check_and_merge_model_level_guardrails(data, router) + snapshot = { + "is_same_object": result is data, + "metadata": result["metadata"], + "model": result["model"], + "extra": result["extra"], + } + assert snapshot == { + "is_same_object": True, + "metadata": {"model_info": {}}, + "model": "m", + "extra": "v", + } + router.get_deployment.assert_not_called() + + +def test_check_and_merge_model_level_guardrails_returns_data_when_deployment_none(): + router = _router_without_deployment() + data = {"metadata": {"model_info": {"id": "x"}}, "model": "m"} + result = _check_and_merge_model_level_guardrails(data, router) + assert result is data + + +def test_check_and_merge_model_level_guardrails_returns_data_when_guardrails_none(): + router = _router_with_deployment(None) + data = {"metadata": {"model_info": {"id": "x"}}, "model": "m"} + result = _check_and_merge_model_level_guardrails(data, router) + assert result is data + + +def test_check_and_merge_model_level_guardrails_handles_missing_metadata(): + router = _router_with_deployment(["pii"]) + data = {"model": "m"} + result = _check_and_merge_model_level_guardrails(data, router) + snapshot = { + "is_same_object": result is data, + "model": result["model"], + "metadata_present": "metadata" in result, + } + assert snapshot == { + "is_same_object": True, + "model": "m", + "metadata_present": False, + } + + +def test_check_and_merge_model_level_guardrails_raises_when_metadata_is_not_dict(): + router = _router_with_deployment(["pii"]) + data = {"metadata": "not-a-dict", "model": "m"} + with pytest.raises(AttributeError): + _check_and_merge_model_level_guardrails(data, router) + + +def test_merge_guardrails_with_existing_happy_path_combines_lists(): + data = { + "metadata": {"guardrails": ["a", "b"], "user": "u"}, + "model": "m", + } + result = _merge_guardrails_with_existing(data, ["c", "a"]) + snapshot = { + "guardrails_sorted": sorted(result["metadata"]["guardrails"]), + "user": result["metadata"]["user"], + "model": result["model"], + "is_copy": result is not data, + } + assert snapshot == { + "guardrails_sorted": ["a", "b", "c"], + "user": "u", + "model": "m", + "is_copy": True, + } + + +def test_merge_guardrails_with_existing_wraps_scalar_existing_guardrail(): + data = {"metadata": {"guardrails": "single-policy"}} + result = _merge_guardrails_with_existing(data, ["model-policy"]) + snapshot = { + "guardrails_sorted": sorted(result["metadata"]["guardrails"]), + "is_list": isinstance(result["metadata"]["guardrails"], list), + "count": len(result["metadata"]["guardrails"]), + } + assert snapshot == { + "guardrails_sorted": ["model-policy", "single-policy"], + "is_list": True, + "count": 2, + } + + +def test_merge_guardrails_with_existing_wraps_scalar_model_guardrail(): + data = {"metadata": {}} + result = _merge_guardrails_with_existing(data, "model-policy") + snapshot = { + "guardrails": result["metadata"]["guardrails"], + "is_list": isinstance(result["metadata"]["guardrails"], list), + "count": len(result["metadata"]["guardrails"]), + } + assert snapshot == { + "guardrails": ["model-policy"], + "is_list": True, + "count": 1, + } + + +def test_merge_guardrails_with_existing_empty_existing_empty_model_yields_empty(): + data = {"metadata": {"guardrails": None}} + result = _merge_guardrails_with_existing(data, None) + snapshot = { + "guardrails": result["metadata"]["guardrails"], + "is_list": isinstance(result["metadata"]["guardrails"], list), + "count": len(result["metadata"]["guardrails"]), + } + assert snapshot == { + "guardrails": [], + "is_list": True, + "count": 0, + } + + +def test_merge_guardrails_with_existing_creates_metadata_when_missing(): + data = {"model": "m"} + result = _merge_guardrails_with_existing(data, ["g1"]) + snapshot = { + "guardrails": result["metadata"]["guardrails"], + "model_preserved": result["model"], + "original_data_unchanged": "metadata" not in data, + } + assert snapshot == { + "guardrails": ["g1"], + "model_preserved": "m", + "original_data_unchanged": True, + } + + +def test_merge_guardrails_with_existing_raises_on_unhashable_guardrail(): + data = {"metadata": {"guardrails": [{"unhashable": True}]}} + with pytest.raises(TypeError): + _merge_guardrails_with_existing(data, ["g1"]) diff --git a/tests/test_litellm/proxy/utils/helpers/test_misc_helpers.py b/tests/test_litellm/proxy/utils/helpers/test_misc_helpers.py new file mode 100644 index 00000000000..7968fa40655 --- /dev/null +++ b/tests/test_litellm/proxy/utils/helpers/test_misc_helpers.py @@ -0,0 +1,201 @@ +import pytest +from fastapi import HTTPException + +from litellm.proxy.utils import ( + construct_database_url_from_env_vars, + get_prisma_client_or_throw, + is_valid_api_key, +) + + +def normalize(value): + return value + + +def test_get_prisma_client_or_throw_happy_path_returns_client(monkeypatch): + sentinel = object() + import litellm.proxy.proxy_server as ps + + monkeypatch.setattr(ps, "prisma_client", sentinel, raising=False) + result = get_prisma_client_or_throw("some message") + summary = { + "is_sentinel": result is sentinel, + "message_arg": "some message", + "raised": False, + } + assert summary == { + "is_sentinel": True, + "message_arg": "some message", + "raised": False, + } + + +def test_get_prisma_client_or_throw_raises_when_client_none(monkeypatch): + import litellm.proxy.proxy_server as ps + + monkeypatch.setattr(ps, "prisma_client", None, raising=False) + with pytest.raises(HTTPException) as exc_info: + get_prisma_client_or_throw("db not connected") + snapshot = { + "status_code": exc_info.value.status_code, + "is_dict_detail": isinstance(exc_info.value.detail, dict), + "error_message": exc_info.value.detail["error"], + } + assert snapshot == { + "status_code": 500, + "is_dict_detail": True, + "error_message": "db not connected", + } + + +def test_is_valid_api_key_happy_path_sk_prefix(): + summary = { + "result": is_valid_api_key("sk-abc123_XYZ-456"), + "key": "sk-abc123_XYZ-456", + "len": len("sk-abc123_XYZ-456"), + } + assert summary == { + "result": True, + "key": "sk-abc123_XYZ-456", + "len": 17, + } + + +def test_is_valid_api_key_happy_path_hashed_64_hex(): + key = "a" * 64 + summary = { + "result": is_valid_api_key(key), + "key_len": len(key), + "is_hex": True, + } + assert summary == { + "result": True, + "key_len": 64, + "is_hex": True, + } + + +def test_is_valid_api_key_happy_path_mixed_case_hex(): + key = "AbCdEf0123456789" * 4 + summary = { + "result": is_valid_api_key(key), + "key_len": len(key), + "first": key[0], + } + assert summary == { + "result": True, + "key_len": 64, + "first": "A", + } + + +def test_is_valid_api_key_error_path_too_long(): + assert is_valid_api_key("sk-" + "a" * 200) is False + + +def test_is_valid_api_key_error_path_non_string(): + assert is_valid_api_key(12345) is False # type: ignore[arg-type] + + +def test_is_valid_api_key_error_path_invalid_format(): + assert is_valid_api_key("not-a-valid-key-format!!!!") is False + + +def test_is_valid_api_key_error_path_too_short(): + assert is_valid_api_key("sk") is False + + +def test_construct_database_url_from_env_vars_happy_path_full(monkeypatch): + monkeypatch.setenv("DATABASE_HOST", "db.example.com") + monkeypatch.setenv("DATABASE_USERNAME", "user") + monkeypatch.setenv("DATABASE_PASSWORD", "pass") + monkeypatch.setenv("DATABASE_NAME", "litellm") + monkeypatch.delenv("DATABASE_SCHEMA", raising=False) + result = construct_database_url_from_env_vars() + summary = { + "result": result, + "host": "db.example.com", + "scheme": result.split("://", 1)[0] if result else None, + "has_password": "pass" in (result or ""), + } + assert summary == { + "result": "postgresql://user:pass@db.example.com/litellm", + "host": "db.example.com", + "scheme": "postgresql", + "has_password": True, + } + + +def test_construct_database_url_from_env_vars_happy_path_no_password(monkeypatch): + monkeypatch.setenv("DATABASE_HOST", "db.example.com") + monkeypatch.setenv("DATABASE_USERNAME", "user") + monkeypatch.delenv("DATABASE_PASSWORD", raising=False) + monkeypatch.setenv("DATABASE_NAME", "litellm") + monkeypatch.delenv("DATABASE_SCHEMA", raising=False) + result = construct_database_url_from_env_vars() + summary = { + "result": result, + "no_colon_password": ":pass@" not in (result or ""), + "host": "db.example.com", + "user": "user", + } + assert summary == { + "result": "postgresql://user@db.example.com/litellm", + "no_colon_password": True, + "host": "db.example.com", + "user": "user", + } + + +def test_construct_database_url_from_env_vars_special_chars_encoded(monkeypatch): + monkeypatch.setenv("DATABASE_HOST", "db.example.com") + monkeypatch.setenv("DATABASE_USERNAME", "us er@x") + monkeypatch.setenv("DATABASE_PASSWORD", "p@ss/word") + monkeypatch.setenv("DATABASE_NAME", "lite/llm") + monkeypatch.delenv("DATABASE_SCHEMA", raising=False) + result = construct_database_url_from_env_vars() + summary = { + "result": result, + "username_encoded": "us+er%40x" in result, + "password_encoded": "p%40ss%2Fword" in result, + "name_encoded": "lite%2Fllm" in result, + } + assert summary == { + "result": "postgresql://us+er%40x:p%40ss%2Fword@db.example.com/lite%2Fllm", + "username_encoded": True, + "password_encoded": True, + "name_encoded": True, + } + + +def test_construct_database_url_from_env_vars_with_schema(monkeypatch): + monkeypatch.setenv("DATABASE_HOST", "db.example.com") + monkeypatch.setenv("DATABASE_USERNAME", "user") + monkeypatch.setenv("DATABASE_PASSWORD", "pass") + monkeypatch.setenv("DATABASE_NAME", "litellm") + monkeypatch.setenv("DATABASE_SCHEMA", "public") + result = construct_database_url_from_env_vars() + summary = { + "result": result, + "schema_appended": result.endswith("?schema=public"), + "host": "db.example.com", + } + assert summary == { + "result": "postgresql://user:pass@db.example.com/litellm?schema=public", + "schema_appended": True, + "host": "db.example.com", + } + + +def test_construct_database_url_from_env_vars_error_path_missing_host(monkeypatch): + monkeypatch.delenv("DATABASE_HOST", raising=False) + monkeypatch.setenv("DATABASE_USERNAME", "user") + monkeypatch.setenv("DATABASE_NAME", "litellm") + assert construct_database_url_from_env_vars() is None + + +def test_construct_database_url_from_env_vars_error_path_missing_username(monkeypatch): + monkeypatch.setenv("DATABASE_HOST", "db.example.com") + monkeypatch.delenv("DATABASE_USERNAME", raising=False) + monkeypatch.setenv("DATABASE_NAME", "litellm") + assert construct_database_url_from_env_vars() is None diff --git a/tests/test_litellm/proxy/utils/helpers/test_model_access.py b/tests/test_litellm/proxy/utils/helpers/test_model_access.py new file mode 100644 index 00000000000..b8e4013c960 --- /dev/null +++ b/tests/test_litellm/proxy/utils/helpers/test_model_access.py @@ -0,0 +1,406 @@ +from unittest.mock import MagicMock + +import pytest +from fastapi import HTTPException + +import litellm +from litellm import ModelResponse +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.utils import ( + create_model_info_response, + get_available_models_for_user, + is_known_model, + is_known_vector_store_index, + model_dump_with_preserved_fields, + validate_model_access, +) + + +def normalize(value): + return value + + +def _router_with_models(model_names): + router = MagicMock() + router.get_model_names.return_value = model_names + router.get_model_access_groups.return_value = {} + return router + + +def test_is_known_model_happy_path_returns_true_when_in_router(): + router = _router_with_models(["gpt-4o", "claude-haiku"]) + summary = { + "result": is_known_model("gpt-4o", router), + "model": "gpt-4o", + "router_models": ["gpt-4o", "claude-haiku"], + } + assert summary == { + "result": True, + "model": "gpt-4o", + "router_models": ["gpt-4o", "claude-haiku"], + } + + +def test_is_known_model_returns_false_when_not_in_router(): + router = _router_with_models(["gpt-4o"]) + summary = { + "result": is_known_model("claude-haiku", router), + "model": "claude-haiku", + "router_models": ["gpt-4o"], + } + assert summary == { + "result": False, + "model": "claude-haiku", + "router_models": ["gpt-4o"], + } + + +def test_is_known_model_error_path_none_model(): + router = _router_with_models(["gpt-4o"]) + assert is_known_model(None, router) is False + + +def test_is_known_model_error_path_none_router(): + assert is_known_model("gpt-4o", None) is False + + +def test_is_known_vector_store_index_happy_path(monkeypatch): + registry = MagicMock() + registry.get_vector_store_indexes.return_value = ["index-a", "index-b"] + monkeypatch.setattr(litellm, "vector_store_index_registry", registry) + summary = { + "result": is_known_vector_store_index("index-a"), + "indexes": ["index-a", "index-b"], + "input": "index-a", + } + assert summary == { + "result": True, + "indexes": ["index-a", "index-b"], + "input": "index-a", + } + + +def test_is_known_vector_store_index_returns_false_when_missing(monkeypatch): + registry = MagicMock() + registry.get_vector_store_indexes.return_value = ["index-a"] + monkeypatch.setattr(litellm, "vector_store_index_registry", registry) + summary = { + "result": is_known_vector_store_index("missing"), + "indexes": ["index-a"], + "input": "missing", + } + assert summary == { + "result": False, + "indexes": ["index-a"], + "input": "missing", + } + + +def test_is_known_vector_store_index_error_path_no_registry(monkeypatch): + monkeypatch.setattr(litellm, "vector_store_index_registry", None) + assert is_known_vector_store_index("anything") is False + + +def test_create_model_info_response_happy_path_no_metadata(): + result = create_model_info_response(model_id="gpt-4o", provider="openai") + assert result == { + "id": "gpt-4o", + "object": "model", + "created": result["created"], + "owned_by": "openai", + } + snapshot = { + "id": result["id"], + "object": result["object"], + "owned_by": result["owned_by"], + "created_is_int": isinstance(result["created"], int), + "metadata_absent": "metadata" not in result, + } + assert snapshot == { + "id": "gpt-4o", + "object": "model", + "owned_by": "openai", + "created_is_int": True, + "metadata_absent": True, + } + + +def test_create_model_info_response_with_metadata_default_general(monkeypatch): + monkeypatch.setattr( + "litellm.proxy.auth.model_checks.get_all_fallbacks", + lambda **_kwargs: [{"model": "fallback-1"}], + ) + result = create_model_info_response( + model_id="gpt-4o", + provider="openai", + include_metadata=True, + ) + snapshot = { + "id": result["id"], + "owned_by": result["owned_by"], + "object": result["object"], + "fallbacks": result["metadata"]["fallbacks"], + } + assert snapshot == { + "id": "gpt-4o", + "owned_by": "openai", + "object": "model", + "fallbacks": [{"model": "fallback-1"}], + } + + +def test_create_model_info_response_with_explicit_fallback_type(monkeypatch): + captured = {} + + def _capture(model, llm_router, fallback_type): + captured["fallback_type"] = fallback_type + return ["x"] + + monkeypatch.setattr("litellm.proxy.auth.model_checks.get_all_fallbacks", _capture) + result = create_model_info_response( + model_id="gpt-4o", + provider="openai", + include_metadata=True, + fallback_type="context_window", + ) + snapshot = { + "id": result["id"], + "fallbacks": result["metadata"]["fallbacks"], + "captured_fallback_type": captured["fallback_type"], + "owned_by": result["owned_by"], + } + assert snapshot == { + "id": "gpt-4o", + "fallbacks": ["x"], + "captured_fallback_type": "context_window", + "owned_by": "openai", + } + + +def test_create_model_info_response_invalid_fallback_type_raises(): + with pytest.raises(HTTPException) as exc_info: + create_model_info_response( + model_id="gpt-4o", + provider="openai", + include_metadata=True, + fallback_type="bogus", + ) + assert exc_info.value.status_code == 400 + assert "Invalid fallback_type" in str(exc_info.value.detail) + + +def test_validate_model_access_happy_path_single_model_in_list(): + summary = { + "result": validate_model_access("gpt-4o", ["gpt-4o", "claude-haiku"]), + "model": "gpt-4o", + "available": ["gpt-4o", "claude-haiku"], + } + assert summary == { + "result": None, + "model": "gpt-4o", + "available": ["gpt-4o", "claude-haiku"], + } + + +def test_validate_model_access_happy_path_batch_all_accessible(): + summary = { + "result": validate_model_access( + "gpt-4o,claude-haiku", ["gpt-4o", "claude-haiku", "gemini"] + ), + "input": "gpt-4o,claude-haiku", + "available": ["gpt-4o", "claude-haiku", "gemini"], + } + assert summary == { + "result": None, + "input": "gpt-4o,claude-haiku", + "available": ["gpt-4o", "claude-haiku", "gemini"], + } + + +def test_validate_model_access_single_model_not_accessible_raises(): + with pytest.raises(HTTPException) as exc_info: + validate_model_access("missing-model", ["gpt-4o"]) + assert exc_info.value.status_code == 404 + assert "missing-model" in str(exc_info.value.detail) + + +def test_validate_model_access_batch_partial_inaccessible_raises(): + with pytest.raises(HTTPException) as exc_info: + validate_model_access("gpt-4o,unknown-x", ["gpt-4o"]) + assert exc_info.value.status_code == 404 + assert "unknown-x" in str(exc_info.value.detail) + assert "gpt-4o" not in str(exc_info.value.detail).split("not accessible:")[1] + + +def _make_model_response(): + return ModelResponse( + id="resp-123", + choices=[ + { + "message": { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": {"name": "do_thing", "arguments": "{}"}, + } + ], + }, + "index": 0, + "finish_reason": "tool_calls", + } + ], + model="gpt-4o", + ) + + +def test_model_dump_with_preserved_fields_restores_none_content(): + resp = _make_model_response() + result = model_dump_with_preserved_fields(resp) + message = result["choices"][0]["message"] + snapshot = { + "content_is_none": message["content"] is None, + "role": message["role"], + "has_tool_calls": "tool_calls" in message, + "model": result["model"], + } + assert snapshot == { + "content_is_none": True, + "role": "assistant", + "has_tool_calls": True, + "model": "gpt-4o", + } + + +def test_model_dump_with_preserved_fields_no_choices_returns_plain_dump(): + class _Bare: + def model_dump(self, **_kwargs): + return {"id": "x", "object": "y", "extra": "z"} + + bare = _Bare() + result = model_dump_with_preserved_fields(bare) + assert result == {"id": "x", "object": "y", "extra": "z"} + + +def test_model_dump_with_preserved_fields_error_path_invalid_obj_raises(): + with pytest.raises(AttributeError): + model_dump_with_preserved_fields(None) + + +@pytest.mark.asyncio +async def test_get_available_models_for_user_happy_path_returns_complete_list( + monkeypatch, +): + monkeypatch.setattr( + "litellm.proxy.auth.model_checks.get_key_models", + lambda **_k: ["gpt-4o"], + ) + monkeypatch.setattr( + "litellm.proxy.auth.model_checks.get_team_models", + lambda **_k: ["claude-haiku"], + ) + monkeypatch.setattr( + "litellm.proxy.auth.model_checks.get_complete_model_list", + lambda **_k: ["gpt-4o", "claude-haiku", "gemini"], + ) + router = _router_with_models(["gpt-4o", "claude-haiku", "gemini"]) + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-test-key", + user_id="user-1", + team_id=None, + team_models=[], + ) + result = await get_available_models_for_user( + user_api_key_dict=user_api_key_dict, + llm_router=router, + general_settings={}, + user_model=None, + ) + summary = { + "result_sorted": sorted(result), + "count": len(result), + "user_id": user_api_key_dict.user_id, + "router_set": True, + } + assert summary == { + "result_sorted": ["claude-haiku", "gemini", "gpt-4o"], + "count": 3, + "user_id": "user-1", + "router_set": True, + } + + +@pytest.mark.asyncio +async def test_get_available_models_for_user_with_none_router(monkeypatch): + monkeypatch.setattr( + "litellm.proxy.auth.model_checks.get_key_models", + lambda **_k: [], + ) + monkeypatch.setattr( + "litellm.proxy.auth.model_checks.get_team_models", + lambda **_k: [], + ) + monkeypatch.setattr( + "litellm.proxy.auth.model_checks.get_complete_model_list", + lambda **_k: ["user-model"], + ) + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-test-key", + user_id="user-1", + team_id=None, + team_models=[], + ) + result = await get_available_models_for_user( + user_api_key_dict=user_api_key_dict, + llm_router=None, + general_settings={}, + user_model="user-model", + ) + summary = { + "result": result, + "router_is_none": True, + "user_model": "user-model", + "count": len(result), + } + assert summary == { + "result": ["user-model"], + "router_is_none": True, + "user_model": "user-model", + "count": 1, + } + + +@pytest.mark.asyncio +async def test_get_available_models_for_user_error_path_complete_list_raises( + monkeypatch, +): + monkeypatch.setattr( + "litellm.proxy.auth.model_checks.get_key_models", + lambda **_k: [], + ) + monkeypatch.setattr( + "litellm.proxy.auth.model_checks.get_team_models", + lambda **_k: [], + ) + + def _boom(**_kwargs): + raise RuntimeError("downstream failure") + + monkeypatch.setattr( + "litellm.proxy.auth.model_checks.get_complete_model_list", _boom + ) + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-test-key", + user_id="user-1", + team_id=None, + team_models=[], + ) + with pytest.raises(RuntimeError): + await get_available_models_for_user( + user_api_key_dict=user_api_key_dict, + llm_router=None, + general_settings={}, + user_model=None, + ) diff --git a/tests/test_litellm/proxy/utils/helpers/test_month_end_projection.py b/tests/test_litellm/proxy/utils/helpers/test_month_end_projection.py new file mode 100644 index 00000000000..5afe1f4faf8 --- /dev/null +++ b/tests/test_litellm/proxy/utils/helpers/test_month_end_projection.py @@ -0,0 +1,232 @@ +from datetime import date, timedelta + +import pytest + +from litellm.proxy.utils import ( + _get_month_end_date, + _get_projected_spend_over_limit, + _is_projected_spend_over_limit, +) + + +def normalize(value): + return value + + +def _freeze_today(monkeypatch, frozen): + class _FrozenDate(date): + @classmethod + def today(cls): + return frozen + + monkeypatch.setattr("litellm.proxy.utils.date", _FrozenDate) + + +@pytest.mark.parametrize( + "today, expected", + [ + (date(2024, 1, 15), date(2024, 1, 31)), + (date(2024, 2, 1), date(2024, 2, 29)), + (date(2023, 2, 1), date(2023, 2, 28)), + (date(2024, 4, 10), date(2024, 4, 30)), + (date(2024, 12, 1), date(2024, 12, 31)), + ], +) +def test_get_month_end_date_happy_path(today, expected): + result = _get_month_end_date(today) + assert normalize( + { + "year": result.year, + "month": result.month, + "day": result.day, + "expected": expected.isoformat(), + "input": today.isoformat(), + } + ) == { + "year": expected.year, + "month": expected.month, + "day": expected.day, + "expected": expected.isoformat(), + "input": today.isoformat(), + } + + +def test_get_month_end_date_raises_on_non_date_input(): + with pytest.raises(AttributeError): + _get_month_end_date("2024-01-15") + + +def test_is_projected_spend_over_limit_happy_path_under_budget(monkeypatch): + _freeze_today(monkeypatch, date(2024, 1, 11)) + summary = { + "result": _is_projected_spend_over_limit( + current_spend=10.0, soft_budget_limit=1_000_000.0 + ), + "current_spend": 10.0, + "soft_budget_limit": 1_000_000.0, + } + assert summary == { + "result": False, + "current_spend": 10.0, + "soft_budget_limit": 1_000_000.0, + } + + +def test_is_projected_spend_over_limit_happy_path_over_budget(monkeypatch): + _freeze_today(monkeypatch, date(2024, 1, 11)) + summary = { + "result": _is_projected_spend_over_limit( + current_spend=100.0, soft_budget_limit=50.0 + ), + "current_spend": 100.0, + "soft_budget_limit": 50.0, + } + assert summary == { + "result": True, + "current_spend": 100.0, + "soft_budget_limit": 50.0, + } + + +def test_is_projected_spend_over_limit_first_of_month_no_division_by_zero(monkeypatch): + _freeze_today(monkeypatch, date(2024, 1, 1)) + summary = { + "result": _is_projected_spend_over_limit( + current_spend=5.0, soft_budget_limit=10.0 + ), + "current_spend": 5.0, + "soft_budget_limit": 10.0, + } + assert summary == { + "result": True, + "current_spend": 5.0, + "soft_budget_limit": 10.0, + } + + +def test_is_projected_spend_over_limit_none_limit_returns_false(): + assert ( + _is_projected_spend_over_limit(current_spend=10_000.0, soft_budget_limit=None) + is False + ) + + +def test_is_projected_spend_over_limit_raises_when_today_missing(monkeypatch): + class _Broken: + @classmethod + def today(cls): + raise RuntimeError("clock unavailable") + + monkeypatch.setattr("litellm.proxy.utils.date", _Broken) + with pytest.raises(RuntimeError): + _is_projected_spend_over_limit(current_spend=1.0, soft_budget_limit=1.0) + + +def test_get_projected_spend_over_limit_happy_path_over_budget(monkeypatch): + _freeze_today(monkeypatch, date(2024, 1, 11)) + result = _get_projected_spend_over_limit( + current_spend=100.0, soft_budget_limit=50.0 + ) + assert result is not None + projected, exceed_date = result + summary = { + "projected_spend": projected, + "exceed_date": exceed_date.isoformat(), + "current_spend": 100.0, + "soft_budget_limit": 50.0, + } + assert summary == { + "projected_spend": 300.0, + "exceed_date": "2024-01-11", + "current_spend": 100.0, + "soft_budget_limit": 50.0, + } + + +def test_get_projected_spend_over_limit_first_of_month_uses_current_as_daily( + monkeypatch, +): + _freeze_today(monkeypatch, date(2024, 1, 1)) + result = _get_projected_spend_over_limit(current_spend=5.0, soft_budget_limit=10.0) + assert result is not None + projected, exceed_date = result + expected_exceed = date(2024, 1, 1) + timedelta(days=1.0) + summary = { + "projected_spend": projected, + "exceed_date": exceed_date.isoformat(), + "expected_exceed_date": expected_exceed.isoformat(), + "soft_budget_limit": 10.0, + } + assert summary == { + "projected_spend": 155.0, + "exceed_date": expected_exceed.isoformat(), + "expected_exceed_date": expected_exceed.isoformat(), + "soft_budget_limit": 10.0, + } + + +def test_get_projected_spend_over_limit_zero_daily_spend_exceed_today(monkeypatch): + _freeze_today(monkeypatch, date(2024, 1, 11)) + result = _get_projected_spend_over_limit(current_spend=0.0, soft_budget_limit=-1.0) + assert result is not None + projected, exceed_date = result + summary = { + "projected_spend": projected, + "exceed_date": exceed_date.isoformat(), + "soft_budget_limit": -1.0, + } + assert summary == { + "projected_spend": 0.0, + "exceed_date": "2024-01-11", + "soft_budget_limit": -1.0, + } + + +def test_get_projected_spend_over_limit_under_budget_returns_none(monkeypatch): + _freeze_today(monkeypatch, date(2024, 1, 11)) + assert ( + _get_projected_spend_over_limit( + current_spend=1.0, soft_budget_limit=1_000_000.0 + ) + is None + ) + + +def test_get_projected_spend_over_limit_exceed_date_uses_remaining_budget(monkeypatch): + _freeze_today(monkeypatch, date(2024, 1, 11)) + result = _get_projected_spend_over_limit(current_spend=20.0, soft_budget_limit=30.0) + assert result is not None + projected, exceed_date = result + daily = 20.0 / 10 + remaining_budget = 30.0 - 20.0 + expected_exceed = date(2024, 1, 11) + timedelta(days=remaining_budget / daily) + summary = { + "projected_spend": projected, + "exceed_date": exceed_date.isoformat(), + "expected_exceed_date": expected_exceed.isoformat(), + "soft_budget_limit": 30.0, + } + assert summary == { + "projected_spend": 60.0, + "exceed_date": expected_exceed.isoformat(), + "expected_exceed_date": expected_exceed.isoformat(), + "soft_budget_limit": 30.0, + } + + +def test_get_projected_spend_over_limit_none_limit_returns_none(): + assert ( + _get_projected_spend_over_limit(current_spend=1.0, soft_budget_limit=None) + is None + ) + + +def test_get_projected_spend_over_limit_raises_when_today_missing(monkeypatch): + class _Broken: + @classmethod + def today(cls): + raise RuntimeError("clock unavailable") + + monkeypatch.setattr("litellm.proxy.utils.date", _Broken) + with pytest.raises(RuntimeError): + _get_projected_spend_over_limit(current_spend=1.0, soft_budget_limit=1.0) diff --git a/tests/test_litellm/proxy/utils/helpers/test_premium_user_check.py b/tests/test_litellm/proxy/utils/helpers/test_premium_user_check.py new file mode 100644 index 00000000000..0a9539c6dc3 --- /dev/null +++ b/tests/test_litellm/proxy/utils/helpers/test_premium_user_check.py @@ -0,0 +1,77 @@ +import pytest +from fastapi import HTTPException + +from litellm.proxy.utils import _premium_user_check + + +def normalize(value): + return value + + +def test_premium_user_check_happy_path_no_raise_when_premium(monkeypatch): + import litellm.proxy.proxy_server as ps + + monkeypatch.setattr(ps, "premium_user", True, raising=False) + summary = { + "result": _premium_user_check(), + "premium_user": True, + "raised": False, + } + assert summary == { + "result": None, + "premium_user": True, + "raised": False, + } + + +def test_premium_user_check_happy_path_with_feature_no_raise(monkeypatch): + import litellm.proxy.proxy_server as ps + + monkeypatch.setattr(ps, "premium_user", True, raising=False) + summary = { + "result": _premium_user_check(feature="model-routing"), + "premium_user": True, + "feature": "model-routing", + } + assert summary == { + "result": None, + "premium_user": True, + "feature": "model-routing", + } + + +def test_premium_user_check_raises_when_not_premium(monkeypatch): + import litellm.proxy.proxy_server as ps + + monkeypatch.setattr(ps, "premium_user", False, raising=False) + with pytest.raises(HTTPException) as exc_info: + _premium_user_check() + snapshot = { + "status_code": exc_info.value.status_code, + "is_dict_detail": isinstance(exc_info.value.detail, dict), + "has_error_key": "error" in exc_info.value.detail, + } + assert snapshot == { + "status_code": 403, + "is_dict_detail": True, + "has_error_key": True, + } + + +def test_premium_user_check_raises_with_feature_message(monkeypatch): + import litellm.proxy.proxy_server as ps + + monkeypatch.setattr(ps, "premium_user", False, raising=False) + with pytest.raises(HTTPException) as exc_info: + _premium_user_check(feature="custom-callbacks") + error_msg = exc_info.value.detail["error"] + snapshot = { + "status_code": exc_info.value.status_code, + "feature_in_message": "custom-callbacks" in error_msg, + "enterprise_in_message": "LiteLLM Enterprise" in error_msg, + } + assert snapshot == { + "status_code": 403, + "feature_in_message": True, + "enterprise_in_message": True, + } diff --git a/tests/test_litellm/proxy/utils/helpers/test_team_configs.py b/tests/test_litellm/proxy/utils/helpers/test_team_configs.py new file mode 100644 index 00000000000..0e0906892b0 --- /dev/null +++ b/tests/test_litellm/proxy/utils/helpers/test_team_configs.py @@ -0,0 +1,76 @@ +import pytest + +from litellm.proxy.utils import _is_valid_team_configs + + +def normalize(value): + return value + + +def test_is_valid_team_configs_happy_path_allowed_model_mutates_config(): + team_config = {"models": ["gpt-4o", "gpt-4o-mini"], "max_budget": 100.0} + request_data = {"model": "gpt-4o"} + snapshot = { + "result": _is_valid_team_configs( + team_id="team-1", + team_config=team_config, + request_data=request_data, + ), + "models_popped": "models" not in team_config, + "remaining_keys": sorted(team_config.keys()), + } + assert snapshot == { + "result": None, + "models_popped": True, + "remaining_keys": ["max_budget"], + } + + +def test_is_valid_team_configs_no_models_key_is_noop(): + team_config = {"max_budget": 100.0, "tpm_limit": 1000} + request_data = {"model": "anything"} + snapshot = { + "result": _is_valid_team_configs( + team_id="team-1", + team_config=team_config, + request_data=request_data, + ), + "team_config": team_config, + "request_data": request_data, + } + assert snapshot == { + "result": None, + "team_config": {"max_budget": 100.0, "tpm_limit": 1000}, + "request_data": {"model": "anything"}, + } + + +def test_is_valid_team_configs_short_circuits_when_team_id_none(): + team_config = {"models": ["only-this"]} + snapshot = { + "result": _is_valid_team_configs( + team_id=None, + team_config=team_config, + request_data={"model": "anything-else"}, + ), + "team_config_unchanged": team_config, + "models_key_preserved": "models" in team_config, + } + assert snapshot == { + "result": None, + "team_config_unchanged": {"models": ["only-this"]}, + "models_key_preserved": True, + } + + +def test_is_valid_team_configs_raises_on_model_not_in_team_models(): + team_config = {"models": ["gpt-4o"]} + request_data = {"model": "claude-haiku"} + with pytest.raises(Exception) as exc_info: + _is_valid_team_configs( + team_id="team-1", + team_config=team_config, + request_data=request_data, + ) + assert "Invalid model for team team-1" in str(exc_info.value) + assert "claude-haiku" in str(exc_info.value) diff --git a/tests/test_litellm/proxy/utils/helpers/test_to_ns.py b/tests/test_litellm/proxy/utils/helpers/test_to_ns.py new file mode 100644 index 00000000000..64ff6d30f0c --- /dev/null +++ b/tests/test_litellm/proxy/utils/helpers/test_to_ns.py @@ -0,0 +1,59 @@ +from datetime import datetime, timezone + +import pytest + +from litellm.proxy.utils import _to_ns + + +def normalize(value): + return value + + +def test_to_ns_happy_path_utc_epoch(): + dt = datetime(2024, 1, 1, 0, 0, 0, tzinfo=timezone.utc) + expected = int(dt.timestamp() * 1e9) + summary = { + "input_iso": dt.isoformat(), + "result": _to_ns(dt), + "expected": expected, + } + assert summary == { + "input_iso": "2024-01-01T00:00:00+00:00", + "result": expected, + "expected": expected, + } + + +def test_to_ns_happy_path_microsecond_precision(): + dt = datetime(2024, 6, 15, 12, 30, 45, 123456, tzinfo=timezone.utc) + expected = int(dt.timestamp() * 1e9) + summary = { + "input_iso": dt.isoformat(), + "result": _to_ns(dt), + "expected": expected, + } + assert summary == { + "input_iso": "2024-06-15T12:30:45.123456+00:00", + "result": expected, + "expected": expected, + } + + +def test_to_ns_result_is_int(): + dt = datetime(2024, 1, 1, tzinfo=timezone.utc) + result = _to_ns(dt) + summary = { + "type": type(result).__name__, + "is_positive": result > 0, + "result": result, + } + assert summary == { + "type": "int", + "is_positive": True, + "result": int(dt.timestamp() * 1e9), + } + + +def test_to_ns_raises_on_invalid_input(): + with pytest.raises(AttributeError): + _to_ns("2024-01-01T00:00:00") diff --git a/tests/test_litellm/proxy/utils/helpers/test_url_helpers.py b/tests/test_litellm/proxy/utils/helpers/test_url_helpers.py new file mode 100644 index 00000000000..31ea1bdce74 --- /dev/null +++ b/tests/test_litellm/proxy/utils/helpers/test_url_helpers.py @@ -0,0 +1,316 @@ +import pytest + +from litellm.proxy.utils import ( + _get_docs_url, + _get_openapi_url, + _get_redoc_url, + get_custom_url, + get_proxy_base_url, + get_server_root_path, + join_paths, + normalize_route_for_root_path, +) + + +def normalize(value): + return value + + +def _clear_url_env(monkeypatch): + for var in ( + "REDOC_URL", + "NO_REDOC", + "DOCS_URL", + "NO_DOCS", + "OPENAPI_URL", + "NO_OPENAPI", + "PROXY_BASE_URL", + "SERVER_ROOT_PATH", + ): + monkeypatch.delenv(var, raising=False) + + +def test_get_redoc_url_default(monkeypatch): + _clear_url_env(monkeypatch) + summary = { + "result": _get_redoc_url(), + "redoc_url_env": None, + "no_redoc_env": None, + } + assert summary == { + "result": "/redoc", + "redoc_url_env": None, + "no_redoc_env": None, + } + + +def test_get_redoc_url_custom_env(monkeypatch): + _clear_url_env(monkeypatch) + monkeypatch.setenv("REDOC_URL", "/custom-redoc") + summary = { + "result": _get_redoc_url(), + "redoc_url_env": "/custom-redoc", + "default_overridden": True, + } + assert summary == { + "result": "/custom-redoc", + "redoc_url_env": "/custom-redoc", + "default_overridden": True, + } + + +def test_get_redoc_url_disabled_returns_none_error_path(monkeypatch): + _clear_url_env(monkeypatch) + monkeypatch.setenv("NO_REDOC", "True") + assert _get_redoc_url() is None + + +def test_get_docs_url_default(monkeypatch): + _clear_url_env(monkeypatch) + summary = { + "result": _get_docs_url(), + "no_docs": None, + "docs_url": None, + } + assert summary == { + "result": "/", + "no_docs": None, + "docs_url": None, + } + + +def test_get_docs_url_custom_env(monkeypatch): + _clear_url_env(monkeypatch) + monkeypatch.setenv("DOCS_URL", "/api-docs") + summary = { + "result": _get_docs_url(), + "env": "/api-docs", + "default_overridden": True, + } + assert summary == { + "result": "/api-docs", + "env": "/api-docs", + "default_overridden": True, + } + + +def test_get_docs_url_disabled_returns_none_error_path(monkeypatch): + _clear_url_env(monkeypatch) + monkeypatch.setenv("NO_DOCS", "True") + assert _get_docs_url() is None + + +def test_get_openapi_url_default(monkeypatch): + _clear_url_env(monkeypatch) + summary = { + "result": _get_openapi_url(), + "no_openapi": None, + "openapi_url": None, + } + assert summary == { + "result": "/openapi.json", + "no_openapi": None, + "openapi_url": None, + } + + +def test_get_openapi_url_custom_env(monkeypatch): + _clear_url_env(monkeypatch) + monkeypatch.setenv("OPENAPI_URL", "/api-schema") + summary = { + "result": _get_openapi_url(), + "env": "/api-schema", + "default_overridden": True, + } + assert summary == { + "result": "/api-schema", + "env": "/api-schema", + "default_overridden": True, + } + + +def test_get_openapi_url_disabled_returns_none_error_path(monkeypatch): + _clear_url_env(monkeypatch) + monkeypatch.setenv("NO_OPENAPI", "True") + assert _get_openapi_url() is None + + +@pytest.mark.parametrize( + "base, route, expected", + [ + ("https://proxy.example.com", "/v1/chat", "https://proxy.example.com/v1/chat"), + ("https://proxy.example.com/", "/v1/chat", "https://proxy.example.com/v1/chat"), + ("https://proxy.example.com", "v1/chat", "https://proxy.example.com/v1/chat"), + ("https://proxy.example.com", "", "https://proxy.example.com"), + ("", "/v1/chat", "/v1/chat"), + ("", "", "/"), + ], +) +def test_join_paths_happy_path(base, route, expected): + result = join_paths(base, route) + assert { + "input_base": base, + "input_route": route, + "result": result, + "expected": expected, + } == { + "input_base": base, + "input_route": route, + "result": expected, + "expected": expected, + } + + +def test_join_paths_avoids_duplicating_route_suffix(): + summary = { + "result": join_paths("https://api.example.com/v1/chat", "/v1/chat"), + "base": "https://api.example.com/v1/chat", + "route": "/v1/chat", + } + assert summary == { + "result": "https://api.example.com/v1/chat", + "base": "https://api.example.com/v1/chat", + "route": "/v1/chat", + } + + +def test_join_paths_invalid_input_raises(): + with pytest.raises(AttributeError): + join_paths(None, "/v1/chat") + + +def test_get_proxy_base_url_returns_env_when_set(monkeypatch): + _clear_url_env(monkeypatch) + monkeypatch.setenv("PROXY_BASE_URL", "https://litellm.test") + summary = { + "result": get_proxy_base_url(), + "env": "https://litellm.test", + "is_set": True, + } + assert summary == { + "result": "https://litellm.test", + "env": "https://litellm.test", + "is_set": True, + } + + +def test_get_proxy_base_url_error_path_returns_none_when_unset(monkeypatch): + _clear_url_env(monkeypatch) + assert get_proxy_base_url() is None + + +def test_get_server_root_path_returns_env(monkeypatch): + _clear_url_env(monkeypatch) + monkeypatch.setenv("SERVER_ROOT_PATH", "/proxy") + summary = { + "result": get_server_root_path(), + "env": "/proxy", + "is_set": True, + } + assert summary == { + "result": "/proxy", + "env": "/proxy", + "is_set": True, + } + + +def test_get_server_root_path_error_path_default_empty_string(monkeypatch): + _clear_url_env(monkeypatch) + assert get_server_root_path() == "" + + +def test_get_custom_url_with_proxy_base_and_root_and_route(monkeypatch): + _clear_url_env(monkeypatch) + monkeypatch.setenv("PROXY_BASE_URL", "https://api.example.com") + monkeypatch.setenv("SERVER_ROOT_PATH", "/proxy") + result = get_custom_url("https://request.example.com", "/v1/chat") + summary = { + "result": result, + "base_used": "PROXY_BASE_URL", + "root_path": "/proxy", + "route": "/v1/chat", + } + assert summary == { + "result": "https://api.example.com/proxy/v1/chat", + "base_used": "PROXY_BASE_URL", + "root_path": "/proxy", + "route": "/v1/chat", + } + + +def test_get_custom_url_falls_back_to_request_base(monkeypatch): + _clear_url_env(monkeypatch) + result = get_custom_url("https://request.example.com", "/v1/chat") + summary = { + "result": result, + "base_used": "request_base_url", + "root_path": "", + "route": "/v1/chat", + } + assert summary == { + "result": "https://request.example.com/v1/chat", + "base_used": "request_base_url", + "root_path": "", + "route": "/v1/chat", + } + + +def test_get_custom_url_no_route_uses_root_path(monkeypatch): + _clear_url_env(monkeypatch) + monkeypatch.setenv("SERVER_ROOT_PATH", "/proxy") + result = get_custom_url("https://request.example.com", route=None) + summary = { + "result": result, + "base_used": "request_base_url", + "root_path": "/proxy", + "route": None, + } + assert summary == { + "result": "https://request.example.com/proxy", + "base_used": "request_base_url", + "root_path": "/proxy", + "route": None, + } + + +def test_get_custom_url_error_path_invalid_base_raises(monkeypatch): + _clear_url_env(monkeypatch) + with pytest.raises(AttributeError): + get_custom_url(None, "/v1/chat") + + +def test_normalize_route_for_root_path_strips_prefix(monkeypatch): + _clear_url_env(monkeypatch) + monkeypatch.setenv("SERVER_ROOT_PATH", "/proxy") + summary = { + "result": normalize_route_for_root_path("/proxy/v1/chat"), + "root_path": "/proxy", + "input": "/proxy/v1/chat", + } + assert summary == { + "result": "/v1/chat", + "root_path": "/proxy", + "input": "/proxy/v1/chat", + } + + +def test_normalize_route_for_root_path_returns_route_when_no_root(monkeypatch): + _clear_url_env(monkeypatch) + summary = { + "result": normalize_route_for_root_path("/v1/chat"), + "root_path": "", + "input": "/v1/chat", + } + assert summary == { + "result": "/v1/chat", + "root_path": "", + "input": "/v1/chat", + } + + +def test_normalize_route_for_root_path_error_path_when_route_not_under_root( + monkeypatch, +): + _clear_url_env(monkeypatch) + monkeypatch.setenv("SERVER_ROOT_PATH", "/proxy") + assert normalize_route_for_root_path("/other/v1/chat") is None diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/__init__.py b/tests/test_litellm/proxy/utils/prisma_and_spend/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/_harness_smoke_test.py b/tests/test_litellm/proxy/utils/prisma_and_spend/_harness_smoke_test.py new file mode 100644 index 00000000000..2243d46ae7f --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/_harness_smoke_test.py @@ -0,0 +1,84 @@ +"""Self-tests for the prisma_and_spend test harness fixtures. + +Verifies the fixtures themselves do what their docstrings claim. +""" + +from __future__ import annotations + +import asyncio +from typing import Any +from unittest.mock import AsyncMock + +import pytest + +from litellm.proxy.utils import PrismaClient + + +def test_normalize_scrubs_volatile_keys() -> None: + from tests.test_litellm.proxy.utils.prisma_and_spend.conftest import normalize + + out = normalize({"id": 1, "spend": 2.0, "team_id": "t1"}) + assert out == {"id": "", "spend": "", "team_id": "t1"} + + +def test_normalize_recurses_into_lists() -> None: + from tests.test_litellm.proxy.utils.prisma_and_spend.conftest import normalize + + out = normalize([{"id": "x"}, {"team_id": "t"}]) + assert out == [{"id": ""}, {"team_id": "t"}] + + +def test_mock_prisma_client_has_common_tables(mock_prisma_client: Any) -> None: + for table in ( + "litellm_verificationtoken", + "litellm_teamtable", + "litellm_usertable", + "litellm_spendlogs", + "litellm_config", + "litellm_healthchecktable", + ): + assert hasattr(mock_prisma_client.db, table) + + +@pytest.mark.asyncio +async def test_mock_dual_cache_round_trip(mock_dual_cache: Any) -> None: + await mock_dual_cache.async_set_cache("k", "v") + assert await mock_dual_cache.async_get_cache("k") == "v" + await mock_dual_cache.async_delete_cache("k") + assert await mock_dual_cache.async_get_cache("k") is None + + +def test_prisma_client_fixture_is_a_real_prismaclient( + prisma_client: PrismaClient, +) -> None: + assert isinstance(prisma_client, PrismaClient) + assert callable(prisma_client.hash_token) + + +@pytest.mark.asyncio +async def test_fake_clock_advances(fake_clock: Any) -> None: + start = fake_clock.now + await asyncio.sleep(2.5) + assert fake_clock.now == start + 2.5 + assert fake_clock.sleep_calls == [2.5] + + +def test_make_spend_log_row_factory(make_spend_log_row: Any) -> None: + row = make_spend_log_row(request_id="abc", spend=0.5) + assert row["request_id"] == "abc" + assert row["spend"] == 0.5 + + +@pytest.mark.asyncio +async def test_in_memory_smtp_captures(in_memory_smtp: Any) -> None: + factory = in_memory_smtp.server_factory() + conn = factory("smtp.invalid", 25) + with conn: + conn.starttls() + from email.message import EmailMessage + + m = EmailMessage() + m["Subject"] = "S" + m.set_content("

x

", subtype="html") + conn.send_message(m, from_addr="a@b", to_addrs="c@d") + assert len(in_memory_smtp.sent) == 1 diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/conftest.py b/tests/test_litellm/proxy/utils/prisma_and_spend/conftest.py new file mode 100644 index 00000000000..2305a88b6dd --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/conftest.py @@ -0,0 +1,387 @@ +"""Shared fixtures for tests/test_litellm/proxy/utils/prisma_and_spend/. + +All fixtures used by PR2 test files live here. Do NOT add fixtures inside +individual test files; if a fixture is missing, add it here and update the +Notion plan. + +The PrismaClient is exercised against a fully-mocked Prisma stack: the +``prisma.Prisma`` constructor and the writer/reader wrappers are patched +before PrismaClient.__init__ runs so the init code paths execute without +needing a generated Prisma client or a real database. +""" + +from __future__ import annotations + +import asyncio +import sys +from dataclasses import dataclass, field +from email.message import EmailMessage +from pathlib import Path +from typing import Any, Callable, Dict, Iterator, List, Optional +from unittest.mock import AsyncMock, MagicMock + +import pytest + +sys.path.insert(0, str(Path(__file__).resolve().parents[5])) + + +VOLATILE_KEYS = frozenset( + { + "created_at", + "updated_at", + "checked_at", + "started_at", + "request_id", + "id", + "token", + "expires", + "expires_at", + "litellm_call_id", + "created", + "spend", + "last_refreshed_at", + "startTime", + "endTime", + "salt", + } +) + + +def normalize(data: Any, volatile: frozenset = VOLATILE_KEYS) -> Any: + """Recursively replace values for volatile keys with ''.""" + if isinstance(data, dict): + return { + k: ("" if k in volatile else normalize(v, volatile)) + for k, v in data.items() + } + if isinstance(data, list): + return [normalize(v, volatile) for v in data] + return data + + +_PRISMA_TABLES: List[str] = [ + "litellm_verificationtoken", + "litellm_teamtable", + "litellm_usertable", + "litellm_endusertable", + "litellm_organizationtable", + "litellm_proxymodeltable", + "litellm_modeltable", + "litellm_budgettable", + "litellm_spendlogs", + "litellm_config", + "litellm_usernotifications", + "litellm_healthchecktable", + "litellm_dailyuserspend", + "litellm_dailyteamspend", + "litellm_dailytagspend", + "litellm_managed_object_table", + "litellm_credentialstable", + "litellm_mcpservertable", + "litellm_audit_log", + "litellm_invitationlink", + "litellm_session_token_table", + "litellm_passthrough_endpoint_table", + "litellm_cron_job", + "litellm_passthrough_logs", + "litellm_promptstable", + "litellm_guardrailstable", + "litellm_managed_files", + "litellm_mcpusercredentials", + "litellm_objectpermissiontable", + "litellm_organizationmembership", +] + + +def _make_table_mock() -> MagicMock: + table = MagicMock() + table.find_unique = AsyncMock(return_value=None) + table.find_many = AsyncMock(return_value=[]) + table.find_first = AsyncMock(return_value=None) + table.create = AsyncMock() + table.create_many = AsyncMock() + table.update = AsyncMock() + table.update_many = AsyncMock() + table.upsert = AsyncMock() + table.delete = AsyncMock() + table.delete_many = AsyncMock() + table.count = AsyncMock(return_value=0) + table.group_by = AsyncMock(return_value=[]) + table.aggregate = AsyncMock(return_value={}) + return table + + +@pytest.fixture +def mock_prisma_client() -> MagicMock: + """Bare ``db`` mock with all common LiteLLM_* tables stubbed. + + Override individual return values in a test:: + + mock_prisma_client.db.litellm_usertable.find_unique.return_value = user + """ + client = MagicMock(name="MockPrismaClient") + client.db = MagicMock(name="MockPrismaDB") + client.connect = AsyncMock() + client.disconnect = AsyncMock() + client.health_check = AsyncMock(return_value=[{"?column?": 1}]) + client.proxy_logging_obj = MagicMock() + client.proxy_logging_obj.failure_handler = AsyncMock() + client.spend_log_transactions = [] + client._spend_log_transactions_lock = asyncio.Lock() + client.jsonify_object = lambda data: dict(data) + client.db.is_connected = MagicMock(return_value=False) + client.db.connect = AsyncMock() + client.db.disconnect = AsyncMock() + client.db.query_raw = AsyncMock(return_value=[{"?column?": 1}]) + client.db.execute_raw = AsyncMock() + client.db.tx = MagicMock() + client.db.batch_ = MagicMock() + for table_name in _PRISMA_TABLES: + setattr(client.db, table_name, _make_table_mock()) + return client + + +@pytest.fixture +def mock_dual_cache() -> MagicMock: + """In-memory DualCache stand-in. + + Sync and async get/set wired against a private dict. Override or read + ``cache._store`` directly in a test for assertion convenience. + """ + cache = MagicMock(name="MockDualCache") + cache._store: Dict[str, Any] = {} + + def _sync_get(key: str, **_: Any) -> Any: + return cache._store.get(key) + + def _sync_set(key: str, value: Any, **_: Any) -> None: + cache._store[key] = value + + async def _async_get(key: str, **_: Any) -> Any: + return cache._store.get(key) + + async def _async_set(key: str, value: Any, **_: Any) -> None: + cache._store[key] = value + + async def _async_delete(key: str, **_: Any) -> None: + cache._store.pop(key, None) + + cache.get_cache = MagicMock(side_effect=_sync_get) + cache.set_cache = MagicMock(side_effect=_sync_set) + cache.async_get_cache = AsyncMock(side_effect=_async_get) + cache.async_set_cache = AsyncMock(side_effect=_async_set) + cache.async_delete_cache = AsyncMock(side_effect=_async_delete) + return cache + + +@pytest.fixture +def patched_prisma_import(monkeypatch: pytest.MonkeyPatch) -> Iterator[MagicMock]: + """Replace ``prisma.Prisma`` and ``PrismaWrapper`` so PrismaClient.__init__ + runs without a generated client. Yields the fake Prisma instance. + + ``prisma`` raises RuntimeError (not AttributeError) for the missing + ``Prisma`` attribute, so ``monkeypatch.setattr`` can't probe it; assign + directly and restore in teardown. + """ + import prisma as _prisma_pkg + import litellm.proxy.utils as _utils_mod + + fake_prisma = MagicMock(name="FakePrisma") + fake_prisma.is_connected = MagicMock(return_value=False) + fake_prisma.connect = AsyncMock() + fake_prisma.disconnect = AsyncMock() + + fake_prisma_factory = MagicMock(name="FakePrismaFactory", return_value=fake_prisma) + had_prisma_attr = "Prisma" in _prisma_pkg.__dict__ + previous_prisma_attr = _prisma_pkg.__dict__.get("Prisma") + _prisma_pkg.Prisma = fake_prisma_factory # type: ignore[attr-defined] + + fake_wrapper = MagicMock(name="FakePrismaWrapper") + fake_wrapper.is_connected = MagicMock(return_value=False) + fake_wrapper.connect = AsyncMock() + fake_wrapper.disconnect = AsyncMock() + fake_wrapper.query_raw = AsyncMock(return_value=[{"?column?": 1}]) + + def _fake_wrapper_ctor(*args: Any, **kwargs: Any) -> MagicMock: + return fake_wrapper + + monkeypatch.setattr(_utils_mod, "PrismaWrapper", _fake_wrapper_ctor) + fake_prisma.__wrapper__ = fake_wrapper + try: + yield fake_prisma + finally: + if had_prisma_attr: + _prisma_pkg.Prisma = previous_prisma_attr # type: ignore[attr-defined] + else: + try: + del _prisma_pkg.Prisma # type: ignore[attr-defined] + except AttributeError: + pass + + +@pytest.fixture +def prisma_client( + patched_prisma_import: MagicMock, + mock_prisma_client: MagicMock, + monkeypatch: pytest.MonkeyPatch, +) -> Any: + """Wired ``PrismaClient`` whose ``db`` attribute is the table mock. + + The init runs through the real code path (testing the constructor's + config-attribute setup) and is then snapped to the easier-to-assert + table mock for downstream behavior pinning. + """ + monkeypatch.delenv("DATABASE_URL_READ_REPLICA", raising=False) + monkeypatch.delenv("IAM_TOKEN_DB_AUTH", raising=False) + from litellm.proxy.utils import PrismaClient + + proxy_logging_obj = MagicMock(name="MockProxyLogging") + proxy_logging_obj.failure_handler = AsyncMock() + pc = PrismaClient( + database_url="postgresql://test:test@localhost:5432/test", + proxy_logging_obj=proxy_logging_obj, + ) + pc.db = mock_prisma_client.db + return pc + + +@dataclass +class FakeClock: + """Monotonic-time controller for the spend monitor loop. + + Tests advance time via ``clock.advance(seconds)`` while asyncio.sleep + is replaced with a clock-driven no-op. + """ + + now: float = 0.0 + sleep_calls: List[float] = field(default_factory=list) + + def advance(self, seconds: float) -> None: + self.now += seconds + + def time(self) -> float: + return self.now + + async def sleep(self, seconds: float) -> None: + self.sleep_calls.append(seconds) + self.now += seconds + + +@pytest.fixture +def fake_clock(monkeypatch: pytest.MonkeyPatch) -> FakeClock: + """Install a controllable clock + asyncio.sleep replacement.""" + clock = FakeClock() + monkeypatch.setattr("time.time", clock.time) + monkeypatch.setattr("time.monotonic", clock.time) + + async def _fast_sleep(seconds: float, *_: Any, **__: Any) -> None: + clock.sleep_calls.append(seconds) + clock.now += seconds + + monkeypatch.setattr("asyncio.sleep", _fast_sleep) + return clock + + +@pytest.fixture +def make_spend_log_row() -> Callable[..., Dict[str, Any]]: + """Factory for fake LiteLLM_SpendLogs rows.""" + + def _make( + request_id: str = "req-1", + spend: float = 0.01, + model: str = "gpt-4o-mini", + **overrides: Any, + ) -> Dict[str, Any]: + row = { + "request_id": request_id, + "spend": spend, + "model": model, + "user": "user-1", + "team_id": "team-1", + "api_key": "hashed-key", + "startTime": "2026-06-02T00:00:00Z", + "endTime": "2026-06-02T00:00:01Z", + "metadata": {}, + } + row.update(overrides) + return row + + return _make + + +@dataclass +class _SentMessage: + from_addr: Optional[str] + to_addrs: Any + subject: Optional[str] + body: Optional[str] + starttls_called: bool + login_args: Optional[tuple] + + +@dataclass +class InMemorySMTP: + """Captures outbound SMTP traffic for ``send_email`` tests.""" + + sent: List[_SentMessage] = field(default_factory=list) + raise_on_send: Optional[Exception] = None + + def server_factory(self) -> Callable[..., Any]: + outer = self + + class _Conn: + def __init__(self) -> None: + self._starttls_called = False + self._login_args: Optional[tuple] = None + + def __enter__(self) -> "_Conn": + return self + + def __exit__(self, *exc: Any) -> None: + return None + + def starttls(self) -> None: + self._starttls_called = True + + def login(self, user: str, password: str) -> None: + self._login_args = (user, password) + + def send_message( + self, + msg: EmailMessage, + from_addr: Optional[str] = None, + to_addrs: Any = None, + ) -> None: + if outer.raise_on_send is not None: + raise outer.raise_on_send + body = "" + for part in msg.walk(): + if part.get_content_type() == "text/html": + body = part.get_payload(decode=False) or "" + break + outer.sent.append( + _SentMessage( + from_addr=from_addr, + to_addrs=to_addrs, + subject=msg["Subject"], + body=body, + starttls_called=self._starttls_called, + login_args=self._login_args, + ) + ) + + def _factory(*args: Any, **kwargs: Any) -> _Conn: + return _Conn() + + return _factory + + +@pytest.fixture +def in_memory_smtp(monkeypatch: pytest.MonkeyPatch) -> InMemorySMTP: + """Patch ``smtplib.SMTP`` to capture sends in memory. + + Override ``smtp.raise_on_send`` to test the SMTP error path. + """ + smtp = InMemorySMTP() + monkeypatch.setattr("smtplib.SMTP", smtp.server_factory()) + return smtp diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_cache_user_row.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_cache_user_row.py new file mode 100644 index 00000000000..d1270b60b19 --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_cache_user_row.py @@ -0,0 +1,81 @@ +"""Pin ``_cache_user_row``. + +Symbols pinned here: + - ``_cache_user_row`` +""" + +from __future__ import annotations + +from types import SimpleNamespace +from typing import Any +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy.utils import _cache_user_row + + +@pytest.mark.asyncio +async def test_cache_user_row_caches_on_miss( + mock_dual_cache: Any, +) -> None: + user_row = SimpleNamespace( + user_id="u1", spend=2.5, max_budget=10.0, name="Alice" + ) + user_row.model_dump_json = MagicMock( + return_value='{"user_id":"u1","spend":2.5,"max_budget":10.0,"name":"Alice"}' + ) + db = MagicMock() + db.get_data = AsyncMock(return_value=user_row) + + result = await _cache_user_row("u1", mock_dual_cache, db) + cache_key = "u1_user_api_key_user_id" + pinned = { + "result": result, + "cache_value": mock_dual_cache._store[cache_key], + "get_calls": mock_dual_cache.get_cache.call_count, + "set_calls": mock_dual_cache.set_cache.call_count, + "db_called": db.get_data.await_count, + } + assert pinned == { + "result": None, + "cache_value": '{"user_id":"u1","spend":2.5,"max_budget":10.0,"name":"Alice"}', + "get_calls": 1, + "set_calls": 1, + "db_called": 1, + } + + +@pytest.mark.asyncio +async def test_cache_user_row_skips_db_on_cache_hit( + mock_dual_cache: Any, +) -> None: + cache_key = "u-hit_user_api_key_user_id" + mock_dual_cache._store[cache_key] = "cached-blob" + db = MagicMock() + db.get_data = AsyncMock(return_value=None) + result = await _cache_user_row("u-hit", mock_dual_cache, db) + assert result is None + assert db.get_data.await_count == 0 + + +@pytest.mark.asyncio +async def test_cache_user_row_skips_set_when_user_row_lacks_model_dump_json( + mock_dual_cache: Any, +) -> None: + user_row = SimpleNamespace(user_id="u2", spend=1.0) + db = MagicMock() + db.get_data = AsyncMock(return_value=user_row) + await _cache_user_row("u2", mock_dual_cache, db) + assert mock_dual_cache._store == {} + assert mock_dual_cache.set_cache.call_count == 0 + + +@pytest.mark.asyncio +async def test_cache_user_row_propagates_db_error( + mock_dual_cache: Any, +) -> None: + db = MagicMock() + db.get_data = AsyncMock(side_effect=RuntimeError("db down")) + with pytest.raises(RuntimeError, match="db down"): + await _cache_user_row("u3", mock_dual_cache, db) diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_config_param_cache.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_config_param_cache.py new file mode 100644 index 00000000000..761835078f4 --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_config_param_cache.py @@ -0,0 +1,267 @@ +"""Pin the LiteLLM_Config cached-read layer. + +Symbols pinned here: + - ``_ConfigRow`` + - ``_config_cache_key`` + - ``_pack_config_row`` + - ``_unpack_config_row`` + - ``get_config_param`` + - ``invalidate_config_param`` + - ``prefetch_config_params`` +""" + +from __future__ import annotations + +from types import SimpleNamespace +from typing import Any, List +from unittest.mock import AsyncMock, MagicMock + +import pytest + +import litellm.proxy.utils as utils_mod +from litellm.proxy.utils import ( + _config_cache_key, + _ConfigRow, + _pack_config_row, + _unpack_config_row, + get_config_param, + invalidate_config_param, + prefetch_config_params, +) + + +@pytest.fixture(autouse=True) +def _swap_config_cache( + monkeypatch: pytest.MonkeyPatch, mock_dual_cache: Any +) -> Any: + """Replace the module-level cache so tests see a clean store per run.""" + monkeypatch.setattr(utils_mod, "litellm_config_cache", mock_dual_cache) + return mock_dual_cache + + +def test_config_cache_key_uses_documented_prefix() -> None: + actual = { + "key": _config_cache_key("max_budget"), + "another": _config_cache_key("disable_spend_updates"), + "prefix": _config_cache_key("x").split(":")[0], + } + assert actual == { + "key": "litellm_config:param:max_budget", + "another": "litellm_config:param:disable_spend_updates", + "prefix": "litellm_config", + } + + +def test_config_cache_key_error_propagates_from_bad_format() -> None: + class _Boom: + def __format__(self, _spec: str) -> str: + raise ValueError("format failure") + + with pytest.raises(ValueError, match="format failure"): + _config_cache_key(_Boom()) # type: ignore[arg-type] + + +def test_config_row_dataclass_shape() -> None: + row = _ConfigRow(param_name="alpha", param_value={"k": 1}) + assert { + "param_name": row.param_name, + "param_value": row.param_value, + "slots": _ConfigRow.__slots__, + } == { + "param_name": "alpha", + "param_value": {"k": 1}, + "slots": ("param_name", "param_value"), + } + + +def test_config_row_rejects_unknown_attribute() -> None: + row = _ConfigRow("a", 1) + with pytest.raises(AttributeError): + row.something_else = 2 # type: ignore[attr-defined] + + +def test_pack_config_row_returns_dict_for_caching() -> None: + row = SimpleNamespace(param_name="zeta", param_value=[1, 2, 3]) + actual = _pack_config_row(row) + expanded = {**actual, "is_dict": isinstance(actual, dict)} + assert expanded == { + "param_name": "zeta", + "param_value": [1, 2, 3], + "is_dict": True, + } + + +def test_pack_config_row_error_on_missing_attribute() -> None: + bad = SimpleNamespace(param_name="only_name") + with pytest.raises(AttributeError): + _pack_config_row(bad) + + +def test_unpack_config_row_round_trips_dict() -> None: + packed = {"param_name": "alpha", "param_value": "abc"} + unpacked = _unpack_config_row(packed) + assert isinstance(unpacked, _ConfigRow) + actual = { + "param_name": unpacked.param_name, + "param_value": unpacked.param_value, + "from_none": _unpack_config_row(None), + "from_miss_sentinel": _unpack_config_row(utils_mod._CONFIG_CACHE_MISS), + "from_other_type": _unpack_config_row(123), + } + assert actual == { + "param_name": "alpha", + "param_value": "abc", + "from_none": None, + "from_miss_sentinel": None, + "from_other_type": None, + } + + +def test_unpack_config_row_error_on_malformed_dict() -> None: + with pytest.raises(KeyError): + _unpack_config_row({"only_name": "x"}) + + +@pytest.mark.asyncio +async def test_get_config_param_cache_hit_returns_unpacked_row( + _swap_config_cache: Any, +) -> None: + cache_key = _config_cache_key("p1") + await _swap_config_cache.async_set_cache( + cache_key, {"param_name": "p1", "param_value": {"x": 1}} + ) + prisma = MagicMock() + prisma.get_generic_data = AsyncMock() + + row = await get_config_param(prisma, "p1") + actual = { + "type": type(row).__name__, + "param_name": row.param_name, + "param_value": row.param_value, + "db_not_touched": prisma.get_generic_data.await_count == 0, + } + assert actual == { + "type": "_ConfigRow", + "param_name": "p1", + "param_value": {"x": 1}, + "db_not_touched": True, + } + + +@pytest.mark.asyncio +async def test_get_config_param_cache_miss_fetches_from_db_and_caches( + _swap_config_cache: Any, +) -> None: + db_row = SimpleNamespace(param_name="p2", param_value={"y": 2}) + prisma = MagicMock() + prisma.get_generic_data = AsyncMock(return_value=db_row) + + row = await get_config_param(prisma, "p2") + cached = _swap_config_cache._store[_config_cache_key("p2")] + actual = { + "returned": row, + "cached": cached, + "db_called": prisma.get_generic_data.await_count, + "db_args": prisma.get_generic_data.await_args.kwargs, + } + assert actual == { + "returned": db_row, + "cached": {"param_name": "p2", "param_value": {"y": 2}}, + "db_called": 1, + "db_args": {"key": "param_name", "value": "p2", "table_name": "config"}, + } + + +@pytest.mark.asyncio +async def test_get_config_param_caches_negative_lookup_as_miss_sentinel( + _swap_config_cache: Any, +) -> None: + prisma = MagicMock() + prisma.get_generic_data = AsyncMock(return_value=None) + row = await get_config_param(prisma, "absent") + assert row is None + assert _swap_config_cache._store[_config_cache_key("absent")] == ( + utils_mod._CONFIG_CACHE_MISS + ) + + +@pytest.mark.asyncio +async def test_get_config_param_raises_when_db_raises() -> None: + prisma = MagicMock() + prisma.get_generic_data = AsyncMock(side_effect=RuntimeError("db down")) + with pytest.raises(RuntimeError, match="db down"): + await get_config_param(prisma, "p3") + + +@pytest.mark.asyncio +async def test_invalidate_config_param_evicts_from_cache( + _swap_config_cache: Any, +) -> None: + cache_key = _config_cache_key("p4") + await _swap_config_cache.async_set_cache(cache_key, {"param_name": "p4", "param_value": 1}) + await invalidate_config_param("p4") + actual = { + "store_empty": _swap_config_cache._store == {}, + "delete_calls": _swap_config_cache.async_delete_cache.await_count, + "delete_arg": _swap_config_cache.async_delete_cache.await_args.args[0], + } + assert actual == { + "store_empty": True, + "delete_calls": 1, + "delete_arg": "litellm_config:param:p4", + } + + +@pytest.mark.asyncio +async def test_invalidate_config_param_propagates_cache_error( + _swap_config_cache: Any, +) -> None: + _swap_config_cache.async_delete_cache = AsyncMock( + side_effect=ConnectionError("redis down") + ) + with pytest.raises(ConnectionError): + await invalidate_config_param("p5") + + +@pytest.mark.asyncio +async def test_prefetch_config_params_populates_cache_for_each_name( + _swap_config_cache: Any, +) -> None: + rows: List[SimpleNamespace] = [ + SimpleNamespace(param_name="a", param_value={"av": 1}), + SimpleNamespace(param_name="c", param_value=[3]), + ] + prisma = MagicMock() + prisma.db.litellm_config.find_many = AsyncMock(return_value=rows) + await prefetch_config_params(prisma, ["a", "b", "c"]) + actual = { + "a": _swap_config_cache._store[_config_cache_key("a")], + "b": _swap_config_cache._store[_config_cache_key("b")], + "c": _swap_config_cache._store[_config_cache_key("c")], + } + assert actual == { + "a": {"param_name": "a", "param_value": {"av": 1}}, + "b": utils_mod._CONFIG_CACHE_MISS, + "c": {"param_name": "c", "param_value": [3]}, + } + + +@pytest.mark.asyncio +async def test_prefetch_config_params_empty_list_is_noop( + _swap_config_cache: Any, +) -> None: + prisma = MagicMock() + prisma.db.litellm_config.find_many = AsyncMock(return_value=[]) + await prefetch_config_params(prisma, []) + assert prisma.db.litellm_config.find_many.await_count == 0 + assert _swap_config_cache._store == {} + + +@pytest.mark.asyncio +async def test_prefetch_config_params_swallows_db_error_without_caching( + _swap_config_cache: Any, +) -> None: + prisma = MagicMock() + prisma.db.litellm_config.find_many = AsyncMock(side_effect=RuntimeError("boom")) + await prefetch_config_params(prisma, ["a", "b"]) + assert _swap_config_cache._store == {} diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_password_helpers.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_password_helpers.py new file mode 100644 index 00000000000..3c028473479 --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_password_helpers.py @@ -0,0 +1,223 @@ +"""Pin password/token helper behavior. + +Symbols pinned here: + - ``hash_token`` + - ``hash_password`` + - ``verify_password`` + - ``migrate_passwords_to_scrypt_async`` + - ``_hash_token_if_needed`` + - ``PrismaClient._is_sha256_hex`` (a nested helper inside + ``migrate_passwords_to_scrypt_async``; the pin list labels it under the + PrismaClient health cluster as a documentation artifact) +""" + +from __future__ import annotations + +import hashlib +from types import SimpleNamespace +from typing import List +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy.utils import ( + _hash_token_if_needed, + hash_password, + hash_token, + migrate_passwords_to_scrypt_async, + verify_password, +) + + +def test_hash_token_returns_sha256_hex_of_input() -> None: + token = "sk-abcDEF12345" + result = hash_token(token) + expected = hashlib.sha256(token.encode()).hexdigest() + actual = { + "len": len(result), + "hex": all(c in "0123456789abcdef" for c in result), + "hash": result, + "matches_sha256": result == expected, + } + assert actual == { + "len": 64, + "hex": True, + "hash": expected, + "matches_sha256": True, + } + + +def test_hash_token_empty_string_still_hashes() -> None: + result = hash_token("") + assert result == hashlib.sha256(b"").hexdigest() + + +def test_hash_token_raises_for_non_string() -> None: + with pytest.raises(AttributeError): + hash_token(None) # type: ignore[arg-type] + + +def test_hash_password_uses_scrypt_prefix() -> None: + h = hash_password("hunter2") + fields = { + "prefix": h[:7], + "min_length": len(h) > 60, + "verifies_self": verify_password("hunter2", h), + "rejects_other": verify_password("hunter3", h), + } + assert fields == { + "prefix": "scrypt:", + "min_length": True, + "verifies_self": True, + "rejects_other": False, + } + + +def test_hash_password_returns_distinct_hashes_per_call() -> None: + a = hash_password("same-password") + b = hash_password("same-password") + assert a != b + assert verify_password("same-password", a) + assert verify_password("same-password", b) + + +def test_hash_password_error_for_non_string_raises() -> None: + with pytest.raises(AttributeError): + hash_password(None) # type: ignore[arg-type] + + +def test_verify_password_sha256_legacy_path() -> None: + plaintext = "legacy-pass" + sha = hashlib.sha256(plaintext.encode()).hexdigest() + matrix = { + "correct": verify_password(plaintext, sha), + "wrong": verify_password("other", sha), + "non_hex_short": verify_password(plaintext, "not-hex"), + "empty_stored": verify_password(plaintext, ""), + } + assert matrix == { + "correct": True, + "wrong": False, + "non_hex_short": False, + "empty_stored": False, + } + + +def test_verify_password_scrypt_malformed_returns_false() -> None: + assert verify_password("anything", "scrypt:not-base64") is False + + +def test_verify_password_unknown_format_returns_false() -> None: + assert verify_password("x", "plaintext-not-supported") is False + + +def test_hash_token_if_needed_handles_sk_prefix() -> None: + plain = "sk-secret-xyz" + already_hashed = hashlib.sha256(plain.encode()).hexdigest() + not_a_secret = "token-without-sk-prefix" + actual = { + "sk_input_is_hashed": _hash_token_if_needed(plain) == already_hashed, + "non_sk_passthrough": _hash_token_if_needed(not_a_secret) == not_a_secret, + "double_hash_stable": _hash_token_if_needed(already_hashed) == already_hashed, + } + assert actual == { + "sk_input_is_hashed": True, + "non_sk_passthrough": True, + "double_hash_stable": True, + } + + +def test_hash_token_if_needed_error_on_non_string() -> None: + with pytest.raises(AttributeError): + _hash_token_if_needed(None) # type: ignore[arg-type] + + +# --------------------------------------------------------------------------- +# migrate_passwords_to_scrypt_async — pins behavior of the nested +# ``_is_sha256_hex`` helper too: scrypt-prefixed and sha256-hex rows are +# left alone, plaintext rows are upgraded in place. +# --------------------------------------------------------------------------- + + +def _make_user(user_id: str, password) -> SimpleNamespace: + return SimpleNamespace(user_id=user_id, password=password) + + +@pytest.mark.asyncio +async def test_migrate_passwords_skips_when_no_plaintext() -> None: + pc = MagicMock() + pc.db = MagicMock() + sha = hashlib.sha256(b"already-hashed").hexdigest() + pc.db.litellm_usertable.find_many = AsyncMock( + return_value=[ + _make_user("a", "scrypt:abc"), + _make_user("b", sha), + ] + ) + pc.db.litellm_usertable.update = AsyncMock() + + result = await migrate_passwords_to_scrypt_async(pc) + outcome = { + "message": result, + "updates": pc.db.litellm_usertable.update.await_count, + "find_called": pc.db.litellm_usertable.find_many.await_count, + "fetch_filter": pc.db.litellm_usertable.find_many.await_args.kwargs["where"], + } + assert outcome == { + "message": "No plaintext passwords found", + "updates": 0, + "find_called": 1, + "fetch_filter": {"password": {"not": None}}, + } + + +@pytest.mark.asyncio +async def test_migrate_passwords_upgrades_only_plaintext_rows() -> None: + pc = MagicMock() + pc.db = MagicMock() + users: List[SimpleNamespace] = [ + _make_user("plaintext-user-1", "plain-1"), + _make_user("plaintext-user-2", "plain-2"), + _make_user("scrypt-user", "scrypt:already"), + _make_user( + "sha-user", + hashlib.sha256(b"alreadyhashed").hexdigest(), + ), + _make_user("null-pw", None), + ] + pc.db.litellm_usertable.find_many = AsyncMock(return_value=users) + pc.db.litellm_usertable.update = AsyncMock() + + result = await migrate_passwords_to_scrypt_async(pc) + + updated_user_ids = sorted( + call.kwargs["where"]["user_id"] + for call in pc.db.litellm_usertable.update.await_args_list + ) + new_password_prefixes = sorted( + call.kwargs["data"]["password"][:7] + for call in pc.db.litellm_usertable.update.await_args_list + ) + outcome = { + "message": result, + "update_count": pc.db.litellm_usertable.update.await_count, + "updated_ids": updated_user_ids, + "all_scrypt_prefixed": new_password_prefixes, + } + assert outcome == { + "message": "Migrated 2 plaintext passwords to scrypt", + "update_count": 2, + "updated_ids": ["plaintext-user-1", "plaintext-user-2"], + "all_scrypt_prefixed": ["scrypt:", "scrypt:"], + } + + +@pytest.mark.asyncio +async def test_migrate_passwords_raises_on_db_failure() -> None: + pc = MagicMock() + pc.db = MagicMock() + pc.db.litellm_usertable.find_many = AsyncMock( + side_effect=RuntimeError("db unavailable") + ) + with pytest.raises(RuntimeError, match="db unavailable"): + await migrate_passwords_to_scrypt_async(pc) diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_engine_watcher.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_engine_watcher.py new file mode 100644 index 00000000000..7b862eecbd4 --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_engine_watcher.py @@ -0,0 +1,521 @@ +"""Pin ``PrismaClient`` engine watcher methods. + +Symbols pinned here: + - ``PrismaClient._get_engine_pid`` + - ``PrismaClient._is_engine_alive`` + - ``PrismaClient._reap_all_zombies`` + - ``PrismaClient._try_waitpid_watch`` + - ``PrismaClient._waitpid_thread_func`` + - ``PrismaClient._on_engine_death_from_thread`` + - ``PrismaClient._try_pidfd_watch`` + - ``PrismaClient._on_pidfd_readable`` + - ``PrismaClient._poll_engine_proc`` + - ``PrismaClient._cleanup_engine_watcher`` + - ``PrismaClient._start_engine_watcher`` + - ``PrismaClient._stop_engine_watcher`` + +Linux-only tests are skipped on Windows; the production code uses +``waitpid``/``pidfd_open`` which are Unix-only. +""" + +from __future__ import annotations + +import asyncio +import os +import sys +import threading +from typing import Any, Optional +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy.utils import PrismaClient + + +pytestmark = pytest.mark.skipif( + sys.platform == "win32", reason="engine watcher is Unix-only" +) + + +def test_get_engine_pid_extracts_process_pid(prisma_client: PrismaClient) -> None: + fake_engine = MagicMock() + fake_engine.process = MagicMock() + fake_engine.process.pid = 4242 + prisma_client.db._original_prisma = MagicMock() + prisma_client.db._original_prisma._engine = fake_engine + actual = { + "pid": prisma_client._get_engine_pid(), + "engine_attr": prisma_client.db._original_prisma._engine is fake_engine, + "process_pid": fake_engine.process.pid, + } + assert actual == {"pid": 4242, "engine_attr": True, "process_pid": 4242} + + +def test_get_engine_pid_returns_zero_when_engine_attr_missing( + prisma_client: PrismaClient, +) -> None: + prisma_client.db._original_prisma = MagicMock(spec=[]) + assert prisma_client._get_engine_pid() == 0 + + +def test_is_engine_alive_true_when_pid_zero(prisma_client: PrismaClient) -> None: + prisma_client._engine_pid = 0 + pinned = { + "result": prisma_client._is_engine_alive(), + "pid": prisma_client._engine_pid, + "type": type(prisma_client._is_engine_alive()).__name__, + } + assert pinned == {"result": True, "pid": 0, "type": "bool"} + + +def test_is_engine_alive_false_when_process_lookup_fails( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + prisma_client._engine_pid = 99999 + monkeypatch.setattr( + "os.kill", MagicMock(side_effect=ProcessLookupError()) + ) + assert prisma_client._is_engine_alive() is False + + +def test_is_engine_alive_true_on_permission_error( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + prisma_client._engine_pid = 1 + monkeypatch.setattr("os.kill", MagicMock(side_effect=PermissionError())) + assert prisma_client._is_engine_alive() is True + + +def test_reap_all_zombies_returns_set_of_reaped_pids( + monkeypatch: pytest.MonkeyPatch, +) -> None: + calls = iter([(111, 0), (222, 0), (0, 0)]) + + def fake_waitpid(pid: int, flags: int) -> Any: + return next(calls) + + monkeypatch.setattr("os.waitpid", fake_waitpid) + reaped = PrismaClient._reap_all_zombies() + pinned = { + "type": type(reaped).__name__, + "size": len(reaped), + "contains_111": 111 in reaped, + "contains_222": 222 in reaped, + } + assert pinned == {"type": "set", "size": 2, "contains_111": True, "contains_222": True} + + +def test_reap_all_zombies_handles_no_children_error( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr( + "os.waitpid", MagicMock(side_effect=ChildProcessError()) + ) + assert PrismaClient._reap_all_zombies() == set() + + +@pytest.mark.asyncio +async def test_try_waitpid_watch_starts_thread_for_live_child( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr("os.waitpid", MagicMock(return_value=(0, 0))) + + threads: list[threading.Thread] = [] + + real_thread_cls = threading.Thread + + def _capture_thread(*args: Any, **kwargs: Any) -> threading.Thread: + t = real_thread_cls(*args, **kwargs) + threads.append(t) + # Replace start so we don't actually launch the thread. + t.start = MagicMock() # type: ignore[method-assign] + return t + + monkeypatch.setattr("threading.Thread", _capture_thread) + monkeypatch.setattr(prisma_client, "_waitpid_thread_func", MagicMock()) + + result = prisma_client._try_waitpid_watch(7777) + pinned = { + "returned": result, + "threads_made": len(threads), + "wait_thread_set": prisma_client._engine_wait_thread is threads[0], + "thread_name_prefix": threads[0].name.startswith("prisma-engine-waitpid-"), + } + assert pinned == { + "returned": True, + "threads_made": 1, + "wait_thread_set": True, + "thread_name_prefix": True, + } + + +@pytest.mark.asyncio +async def test_try_waitpid_watch_returns_false_for_non_child( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr( + "os.waitpid", MagicMock(side_effect=ChildProcessError()) + ) + assert prisma_client._try_waitpid_watch(123) is False + + +@pytest.mark.asyncio +async def test_try_waitpid_watch_handles_already_dead_pid( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + """If the engine PID is already dead at watch start, _try_waitpid_watch + returns True and schedules a reconnect. + """ + prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) + monkeypatch.setattr("os.waitpid", MagicMock(return_value=(8888, 0))) + monkeypatch.setattr(PrismaClient, "_reap_all_zombies", staticmethod(lambda: set())) + monkeypatch.setattr(prisma_client, "_cleanup_engine_watcher", MagicMock()) + + result = prisma_client._try_waitpid_watch(8888) + # Drain pending tasks so attempt_db_reconnect is awaited and we don't leak. + await asyncio.sleep(0) + pinned = { + "result": result, + "engine_confirmed_dead": prisma_client._engine_confirmed_dead, + "cleanup_called": prisma_client._cleanup_engine_watcher.call_count, + "reconnect_scheduled": prisma_client.attempt_db_reconnect.await_count >= 1, + } + assert pinned == { + "result": True, + "engine_confirmed_dead": True, + "cleanup_called": 1, + "reconnect_scheduled": True, + } + + +def test_waitpid_thread_func_swallows_child_process_error( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr("os.waitpid", MagicMock(side_effect=ChildProcessError())) + loop = MagicMock() + loop.call_soon_threadsafe = MagicMock() + prisma_client._waitpid_thread_func(123, loop) + assert loop.call_soon_threadsafe.call_count == 1 + + +def test_waitpid_thread_func_invokes_on_engine_death_on_normal_exit( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr("os.waitpid", MagicMock(return_value=(123, 0))) + loop = MagicMock() + received: list[Any] = [] + loop.call_soon_threadsafe = lambda fn, pid: received.append((fn, pid)) + prisma_client._waitpid_thread_func(123, loop) + pinned = { + "callbacks_received": len(received), + "callback_target": received[0][0] == prisma_client._on_engine_death_from_thread, + "pid_arg": received[0][1], + "first_tuple_size": len(received[0]), + } + assert pinned == { + "callbacks_received": 1, + "callback_target": True, + "pid_arg": 123, + "first_tuple_size": 2, + } + + +def test_waitpid_thread_func_swallows_loop_runtime_error( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr("os.waitpid", MagicMock(return_value=(123, 0))) + loop = MagicMock() + loop.call_soon_threadsafe = MagicMock(side_effect=RuntimeError("loop closed")) + prisma_client._waitpid_thread_func(123, loop) + + +@pytest.mark.asyncio +async def test_on_engine_death_from_thread_schedules_reconnect( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + prisma_client._engine_pid = 7777 + prisma_client._engine_confirmed_dead = False + prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) + monkeypatch.setattr(PrismaClient, "_reap_all_zombies", staticmethod(lambda: set())) + monkeypatch.setattr(prisma_client, "_cleanup_engine_watcher", MagicMock()) + + prisma_client._on_engine_death_from_thread(7777) + await asyncio.sleep(0) + pinned = { + "confirmed_dead": prisma_client._engine_confirmed_dead, + "cleanup_called": prisma_client._cleanup_engine_watcher.call_count, + "reconnect_called": prisma_client.attempt_db_reconnect.await_count, + "reconnect_reason": prisma_client.attempt_db_reconnect.await_args.kwargs["reason"], + } + assert pinned == { + "confirmed_dead": True, + "cleanup_called": 1, + "reconnect_called": 1, + "reconnect_reason": "engine_process_death", + } + + +def test_on_engine_death_from_thread_ignores_wrong_pid_or_already_dead( + prisma_client: PrismaClient, +) -> None: + prisma_client._engine_pid = 1111 + prisma_client._engine_confirmed_dead = True + prisma_client._cleanup_engine_watcher = MagicMock() + prisma_client._on_engine_death_from_thread(1111) + assert prisma_client._cleanup_engine_watcher.call_count == 0 + + +def test_on_engine_death_from_thread_wrong_pid_does_nothing( + prisma_client: PrismaClient, +) -> None: + prisma_client._engine_pid = 1111 + prisma_client._engine_confirmed_dead = False + prisma_client._cleanup_engine_watcher = MagicMock() + prisma_client._on_engine_death_from_thread(2222) + assert prisma_client._cleanup_engine_watcher.call_count == 0 + assert prisma_client._engine_confirmed_dead is False + + +@pytest.mark.asyncio +async def test_try_pidfd_watch_returns_false_when_pidfd_open_missing( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.delattr("os.pidfd_open", raising=False) + assert prisma_client._try_pidfd_watch(123) is False + + +@pytest.mark.asyncio +async def test_try_pidfd_watch_arms_reader_when_available( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + def fake_pidfd(pid: int, flags: int) -> int: + return 42 + + monkeypatch.setattr("os.pidfd_open", fake_pidfd, raising=False) + loop = asyncio.get_running_loop() + fake_add_reader = MagicMock() + monkeypatch.setattr(loop, "add_reader", fake_add_reader) + + result = prisma_client._try_pidfd_watch(123) + assert result is True + assert prisma_client._engine_pidfd == 42 + assert fake_add_reader.call_args.args[0] == 42 + + +@pytest.mark.asyncio +async def test_try_pidfd_watch_error_returns_false_and_cleans_up( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + def fake_pidfd(pid: int, flags: int) -> int: + raise OSError("ENOSYS") + + monkeypatch.setattr("os.pidfd_open", fake_pidfd, raising=False) + assert prisma_client._try_pidfd_watch(123) is False + assert prisma_client._engine_pidfd == -1 + + +@pytest.mark.asyncio +async def test_on_pidfd_readable_invokes_reconnect_path( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + prisma_client._engine_pid = 4321 + prisma_client._engine_confirmed_dead = False + prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) + monkeypatch.setattr(PrismaClient, "_reap_all_zombies", staticmethod(lambda: set())) + cleanup = MagicMock() + prisma_client._cleanup_engine_watcher = cleanup + + prisma_client._on_pidfd_readable() + await asyncio.sleep(0) + pinned = { + "confirmed_dead": prisma_client._engine_confirmed_dead, + "cleanup_called": cleanup.call_count, + "reconnect_called": prisma_client.attempt_db_reconnect.await_count, + "force_kwarg": prisma_client.attempt_db_reconnect.await_args.kwargs["force"], + } + assert pinned == { + "confirmed_dead": True, + "cleanup_called": 1, + "reconnect_called": 1, + "force_kwarg": True, + } + + +@pytest.mark.asyncio +async def test_on_pidfd_readable_noop_when_already_dead_closes_pidfd( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + """When _engine_confirmed_dead is already True, the reader handler should + not schedule another reconnect and should release the pidfd resource. + """ + closed: list[int] = [] + monkeypatch.setattr("os.close", lambda fd: closed.append(fd)) + loop = asyncio.get_running_loop() + removed: list[int] = [] + monkeypatch.setattr(loop, "remove_reader", lambda fd: removed.append(fd)) + + prisma_client._engine_confirmed_dead = True + prisma_client._engine_pidfd = 99 + prisma_client.attempt_db_reconnect = AsyncMock() + + prisma_client._on_pidfd_readable() + pinned = { + "engine_pidfd": prisma_client._engine_pidfd, + "closed": closed, + "removed": removed, + "reconnect_call_count": prisma_client.attempt_db_reconnect.await_count, + } + assert pinned == { + "engine_pidfd": -1, + "closed": [99], + "removed": [99], + "reconnect_call_count": 0, + } + + +@pytest.mark.asyncio +async def test_poll_engine_proc_detects_death_and_reconnects( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + prisma_client._engine_pid = 555 + prisma_client._watching_engine = True + prisma_client.attempt_db_reconnect = AsyncMock() + monkeypatch.setattr("os.kill", MagicMock(side_effect=ProcessLookupError())) + monkeypatch.setattr(PrismaClient, "_reap_all_zombies", staticmethod(lambda: set())) + prisma_client._cleanup_engine_watcher = MagicMock() + + await prisma_client._poll_engine_proc() + pinned = { + "reconnect_count": prisma_client.attempt_db_reconnect.await_count, + "cleanup_count": prisma_client._cleanup_engine_watcher.call_count, + "confirmed_dead": prisma_client._engine_confirmed_dead, + "reason": prisma_client.attempt_db_reconnect.await_args.kwargs["reason"], + } + assert pinned == { + "reconnect_count": 1, + "cleanup_count": 1, + "confirmed_dead": True, + "reason": "engine_process_death", + } + + +@pytest.mark.asyncio +async def test_poll_engine_proc_returns_on_permission_error( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + prisma_client._engine_pid = 555 + prisma_client._watching_engine = True + monkeypatch.setattr("os.kill", MagicMock(side_effect=PermissionError())) + prisma_client._cleanup_engine_watcher = MagicMock() + await prisma_client._poll_engine_proc() + assert prisma_client._cleanup_engine_watcher.call_count == 1 + + +@pytest.mark.asyncio +async def test_cleanup_engine_watcher_resets_state( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + closed: list[int] = [] + monkeypatch.setattr("os.close", lambda fd: closed.append(fd)) + loop = asyncio.get_running_loop() + removed: list[int] = [] + monkeypatch.setattr(loop, "remove_reader", lambda fd: removed.append(fd)) + + prisma_client._engine_pidfd = 42 + prisma_client._engine_pid = 999 + prisma_client._engine_wait_thread = MagicMock() + prisma_client._watching_engine = True + + prisma_client._cleanup_engine_watcher() + pinned = { + "engine_pidfd": prisma_client._engine_pidfd, + "engine_pid": prisma_client._engine_pid, + "wait_thread": prisma_client._engine_wait_thread, + "watching": prisma_client._watching_engine, + "closed": closed, + "removed": removed, + } + assert pinned == { + "engine_pidfd": -1, + "engine_pid": 0, + "wait_thread": None, + "watching": False, + "closed": [42], + "removed": [42], + } + + +@pytest.mark.asyncio +async def test_cleanup_engine_watcher_swallows_close_error( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr("os.close", MagicMock(side_effect=OSError("bad fd"))) + loop = asyncio.get_running_loop() + monkeypatch.setattr(loop, "remove_reader", MagicMock(side_effect=Exception("boom"))) + prisma_client._engine_pidfd = 99 + prisma_client._cleanup_engine_watcher() + assert prisma_client._engine_pidfd == -1 + + +@pytest.mark.asyncio +async def test_start_engine_watcher_picks_waitpid_when_available( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr(prisma_client, "_get_engine_pid", MagicMock(return_value=12345)) + monkeypatch.setattr(prisma_client, "_try_waitpid_watch", MagicMock(return_value=True)) + pidfd_called = MagicMock(return_value=False) + monkeypatch.setattr(prisma_client, "_try_pidfd_watch", pidfd_called) + await prisma_client._start_engine_watcher() + pinned = { + "engine_pid": prisma_client._engine_pid, + "confirmed_dead_reset": prisma_client._engine_confirmed_dead, + "waitpid_called": prisma_client._try_waitpid_watch.call_count, + "pidfd_skipped": pidfd_called.call_count, + } + assert pinned == { + "engine_pid": 12345, + "confirmed_dead_reset": False, + "waitpid_called": 1, + "pidfd_skipped": 0, + } + + +@pytest.mark.asyncio +async def test_start_engine_watcher_returns_early_when_pid_unknown( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr(prisma_client, "_get_engine_pid", MagicMock(return_value=0)) + monkeypatch.setattr(prisma_client, "_try_waitpid_watch", MagicMock()) + await prisma_client._start_engine_watcher() + assert prisma_client._try_waitpid_watch.call_count == 0 + + +@pytest.mark.asyncio +async def test_start_engine_watcher_falls_back_to_polling_when_no_kernel_apis( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr(prisma_client, "_get_engine_pid", MagicMock(return_value=4242)) + monkeypatch.setattr(prisma_client, "_try_waitpid_watch", MagicMock(return_value=False)) + monkeypatch.setattr(prisma_client, "_try_pidfd_watch", MagicMock(return_value=False)) + monkeypatch.setattr(prisma_client, "_poll_engine_proc", AsyncMock()) + await prisma_client._start_engine_watcher() + await asyncio.sleep(0) + assert prisma_client._watching_engine is True + + +def test_stop_engine_watcher_clears_dead_flag( + prisma_client: PrismaClient, +) -> None: + prisma_client._engine_confirmed_dead = True + prisma_client._cleanup_engine_watcher = MagicMock() + prisma_client._stop_engine_watcher() + assert prisma_client._cleanup_engine_watcher.call_count == 1 + assert prisma_client._engine_confirmed_dead is False + + +def test_stop_engine_watcher_error_in_cleanup_propagates( + prisma_client: PrismaClient, +) -> None: + prisma_client._cleanup_engine_watcher = MagicMock(side_effect=RuntimeError("cleanup boom")) + with pytest.raises(RuntimeError, match="cleanup boom"): + prisma_client._stop_engine_watcher() diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_get_data.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_get_data.py new file mode 100644 index 00000000000..7e7e98d1360 --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_get_data.py @@ -0,0 +1,400 @@ +"""Pin ``PrismaClient`` read-side data operations. + +Symbols pinned here: + - ``PrismaClient.hash_token`` + - ``PrismaClient.jsonify_object`` + - ``PrismaClient.jsonify_team_object`` + - ``PrismaClient.check_view_exists`` + - ``PrismaClient.get_request_status`` + - ``PrismaClient.get_generic_data`` + - ``PrismaClient._query_first_with_cached_plan_fallback`` + - ``PrismaClient.get_data`` +""" + +from __future__ import annotations + +import hashlib +import json +from types import SimpleNamespace +from typing import Any +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi import HTTPException + +from litellm.proxy.utils import PrismaClient + + +def test_hash_token_method_returns_sha256(prisma_client: PrismaClient) -> None: + token = "sk-token-xyz" + actual = { + "result": prisma_client.hash_token(token), + "len": len(prisma_client.hash_token(token)), + "expected": hashlib.sha256(token.encode()).hexdigest(), + "deterministic": prisma_client.hash_token(token) + == prisma_client.hash_token(token), + } + assert actual == { + "result": hashlib.sha256(token.encode()).hexdigest(), + "len": 64, + "expected": hashlib.sha256(token.encode()).hexdigest(), + "deterministic": True, + } + + +def test_hash_token_method_error_on_non_string(prisma_client: PrismaClient) -> None: + with pytest.raises(AttributeError): + prisma_client.hash_token(None) # type: ignore[arg-type] + + +def test_jsonify_object_serializes_nested_dicts(prisma_client: PrismaClient) -> None: + data = { + "metadata": {"a": 1, "b": [2, 3]}, + "models": ["gpt-4o", "gpt-4o-mini"], + "token": "abc", + "spend": 1.23, + } + result = prisma_client.jsonify_object(data) + parsed_meta = json.loads(result["metadata"]) + assert result == { + "metadata": json.dumps(data["metadata"]), + "models": ["gpt-4o", "gpt-4o-mini"], + "token": "abc", + "spend": 1.23, + } + assert parsed_meta == {"a": 1, "b": [2, 3]} + + +def test_jsonify_object_fallback_for_unserializable_dict( + prisma_client: PrismaClient, +) -> None: + class _Bad: + pass + + data = {"metadata": {"x": _Bad()}, "label": "ok", "n": 1} + result = prisma_client.jsonify_object(data) + assert result == { + "metadata": "failed-to-serialize-json", + "label": "ok", + "n": 1, + } + + +def test_jsonify_object_error_on_non_dict(prisma_client: PrismaClient) -> None: + with pytest.raises(AttributeError): + prisma_client.jsonify_object(None) # type: ignore[arg-type] + + +def test_jsonify_team_object_converts_members_to_json_string( + prisma_client: PrismaClient, +) -> None: + data = { + "team_id": "t1", + "members_with_roles": [{"role": "admin", "user_id": "u1"}], + "metadata": {"foo": "bar"}, + "models": ["gpt-4"], + } + result = prisma_client.jsonify_team_object(data) + assert result == { + "team_id": "t1", + "members_with_roles": json.dumps(data["members_with_roles"]), + "metadata": json.dumps(data["metadata"]), + "models": ["gpt-4"], + } + + +def test_jsonify_team_object_error_on_non_dict(prisma_client: PrismaClient) -> None: + with pytest.raises(AttributeError): + prisma_client.jsonify_team_object(None) # type: ignore[arg-type] + + +@pytest.mark.parametrize( + "metadata,expected", + [ + ({"status": "failure"}, "failure"), + ({"status": "success"}, "success"), + ({}, "success"), + ("not-json", "success"), + (json.dumps({"status": "failure"}), "failure"), + ], +) +def test_get_request_status_pins_status_resolution( + prisma_client: PrismaClient, metadata: Any, expected: str +) -> None: + assert prisma_client.get_request_status({"metadata": metadata}) == expected + + +def test_get_request_status_error_returns_success_default( + prisma_client: PrismaClient, +) -> None: + """``get_request_status`` swallows AttributeError / JSONDecodeError and + defaults to ``success`` to avoid blocking the request pipeline. + """ + + class _Broken: + def get(self, *_: Any, **__: Any) -> Any: + raise AttributeError("broken metadata") + + actual = prisma_client.get_request_status({"metadata": _Broken()}) + assert actual == "success" + + +@pytest.mark.asyncio +async def test_get_generic_data_dispatches_by_table( + prisma_client: PrismaClient, +) -> None: + row = SimpleNamespace(user_id="u1", spend=0.5, name="Alice") + prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=row) + result = await prisma_client.get_generic_data( + key="user_id", value="u1", table_name="users" + ) + actual = { + "result_is_row": result is row, + "find_first_count": prisma_client.db.litellm_usertable.find_first.await_count, + "where_kwarg": prisma_client.db.litellm_usertable.find_first.await_args.kwargs[ + "where" + ], + "user_attr": result.user_id, + } + assert actual == { + "result_is_row": True, + "find_first_count": 1, + "where_kwarg": {"user_id": "u1"}, + "user_attr": "u1", + } + + +@pytest.mark.asyncio +async def test_get_generic_data_unknown_table_returns_none( + prisma_client: PrismaClient, +) -> None: + result = await prisma_client.get_generic_data( + key="x", value="y", table_name="bogus" # type: ignore[arg-type] + ) + assert result is None + + +@pytest.mark.asyncio +async def test_get_generic_data_logs_failure_handler_and_raises_on_error( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.litellm_usertable.find_first = AsyncMock( + side_effect=RuntimeError("db boom") + ) + with pytest.raises(RuntimeError, match="db boom"): + await prisma_client.get_generic_data( + key="user_id", value="x", table_name="users" + ) + + +@pytest.mark.asyncio +async def test_query_first_with_cached_plan_fallback_happy_returns_row( + prisma_client: PrismaClient, +) -> None: + expected = {"token": "abc", "team_spend": 1.0, "team_max_budget": 5.0} + prisma_client.db.query_first = AsyncMock(return_value=expected) + result = await prisma_client._query_first_with_cached_plan_fallback( + "SELECT * FROM x WHERE token = $1", "abc" + ) + actual = { + "result": result, + "call_count": prisma_client.db.query_first.await_count, + "args": prisma_client.db.query_first.await_args.args, + "matches": result == expected, + } + assert actual == { + "result": expected, + "call_count": 1, + "args": ("SELECT * FROM x WHERE token = $1", "abc"), + "matches": True, + } + + +@pytest.mark.asyncio +async def test_query_first_with_cached_plan_fallback_retries_on_cached_plan_error( + prisma_client: PrismaClient, +) -> None: + expected = {"token": "abc", "team_spend": 1.0, "team_max_budget": 5.0} + prisma_client.db.query_first = AsyncMock( + side_effect=[ + RuntimeError("cached plan must not change result type"), + expected, + ] + ) + result = await prisma_client._query_first_with_cached_plan_fallback( + "SELECT * FROM x WHERE token = $1", "abc" + ) + assert result == expected + assert prisma_client.db.query_first.await_count == 2 + second_call_sql = prisma_client.db.query_first.await_args_list[1].args[0] + assert "cache_invalidated_" in second_call_sql + + +@pytest.mark.asyncio +async def test_query_first_with_cached_plan_fallback_reraises_non_plan_errors( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.query_first = AsyncMock(side_effect=RuntimeError("totally unrelated")) + with pytest.raises(RuntimeError, match="totally unrelated"): + await prisma_client._query_first_with_cached_plan_fallback("SELECT 1") + + +@pytest.mark.asyncio +async def test_check_view_exists_noop_when_all_views_present( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.query_raw = AsyncMock( + return_value=[ + { + "view_count": 8, + "view_names": [ + "LiteLLM_VerificationTokenView", + "MonthlyGlobalSpend", + "Last30dKeysBySpend", + "Last30dModelsBySpend", + "MonthlyGlobalSpendPerKey", + "MonthlyGlobalSpendPerUserPerKey", + "Last30dTopEndUsersSpend", + "DailyTagSpend", + ], + } + ] + ) + prisma_client.db.execute_raw = AsyncMock() + result = await prisma_client.check_view_exists() + actual = { + "result": result, + "query_raw_calls": prisma_client.db.query_raw.await_count, + "execute_raw_calls": prisma_client.db.execute_raw.await_count, + "view_query_contains_token_view": "LiteLLM_VerificationTokenView" + in prisma_client.db.query_raw.await_args.args[0], + } + assert actual == { + "result": None, + "query_raw_calls": 1, + "execute_raw_calls": 0, + "view_query_contains_token_view": True, + } + + +@pytest.mark.asyncio +async def test_check_view_exists_creates_token_view_when_missing( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.query_raw = AsyncMock( + return_value=[ + { + "view_count": 1, + "view_names": ["DailyTagSpend"], + } + ] + ) + prisma_client.db.execute_raw = AsyncMock() + prisma_client.health_check = AsyncMock(return_value=[{"?column?": 1}]) + result = await prisma_client.check_view_exists() + actual = { + "result": result, + "create_called": prisma_client.db.execute_raw.await_count, + "create_sql_starts_with_create_view": prisma_client.db.execute_raw.await_args.args[ + 0 + ] + .strip() + .startswith('CREATE VIEW "LiteLLM_VerificationTokenView"'), + } + assert actual == { + "result": None, + "create_called": 1, + "create_sql_starts_with_create_view": True, + } + + +@pytest.mark.asyncio +async def test_check_view_exists_raises_when_query_raw_fails( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.query_raw = AsyncMock(side_effect=RuntimeError("db down")) + with pytest.raises(RuntimeError, match="db down"): + await prisma_client.check_view_exists() + + +@pytest.mark.asyncio +async def test_get_data_token_find_unique_returns_record( + prisma_client: PrismaClient, +) -> None: + token = "sk-key-1" + hashed = hashlib.sha256(token.encode()).hexdigest() + record = SimpleNamespace(token=hashed, user_id="u1", expires=None, spend=0.5) + prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=record + ) + + result = await prisma_client.get_data(token=token, table_name="key") + actual = { + "result_is_record": result is record, + "where_arg": prisma_client.db.litellm_verificationtoken.find_unique.await_args.kwargs[ + "where" + ], + "include_arg": prisma_client.db.litellm_verificationtoken.find_unique.await_args.kwargs[ + "include" + ], + "token_field_matches": result.token == hashed, + } + assert actual == { + "result_is_record": True, + "where_arg": {"token": hashed}, + "include_arg": {"litellm_budget_table": True}, + "token_field_matches": True, + } + + +@pytest.mark.asyncio +async def test_get_data_token_find_unique_missing_token_raises_401( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=None) + with pytest.raises(HTTPException) as excinfo: + await prisma_client.get_data(token="sk-missing", table_name="key") + err = excinfo.value + assert "invalid user key" in err.detail + assert err.status_code == 401 + + +@pytest.mark.asyncio +async def test_get_data_user_find_unique_returns_user_row( + prisma_client: PrismaClient, +) -> None: + row = SimpleNamespace( + user_id="u-7", + spend=1.5, + max_budget=10.0, + organization_memberships=[], + ) + prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=row) + result = await prisma_client.get_data(user_id="u-7", table_name="user") + actual = { + "result_is_row": result is row, + "where_arg": prisma_client.db.litellm_usertable.find_unique.await_args.kwargs[ + "where" + ], + "include_arg": prisma_client.db.litellm_usertable.find_unique.await_args.kwargs[ + "include" + ], + "spend": row.spend, + } + assert actual == { + "result_is_row": True, + "where_arg": {"user_id": "u-7"}, + "include_arg": {"organization_memberships": True}, + "spend": 1.5, + } + + +@pytest.mark.asyncio +async def test_get_data_logs_and_raises_on_db_error( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + side_effect=RuntimeError("network split") + ) + with pytest.raises(RuntimeError, match="network split"): + await prisma_client.get_data(token="sk-broken", table_name="key") diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_health.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_health.py new file mode 100644 index 00000000000..220fff1a881 --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_health.py @@ -0,0 +1,292 @@ +"""Pin ``PrismaClient`` health + spend-logs counter helpers. + +Symbols pinned here: + - ``PrismaClient.health_check`` + - ``PrismaClient._get_spend_logs_row_count`` + - ``PrismaClient._set_spend_logs_row_count_in_proxy_state`` + - ``PrismaClient._validate_response_time`` + - ``PrismaClient._clean_details`` + - ``PrismaClient.save_health_check_result`` + - ``PrismaClient.get_health_check_history`` + - ``PrismaClient.get_all_latest_health_checks`` + - ``PrismaClient._is_sha256_hex`` (a nested helper inside + ``migrate_passwords_to_scrypt_async``; the pin list assigns it to this + cluster as a documentation artifact) +""" + +from __future__ import annotations + +from typing import Any +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy.utils import PrismaClient + + +@pytest.mark.asyncio +async def test_health_check_returns_query_raw_result( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.query_raw = AsyncMock(return_value=[{"?column?": 1}]) + result = await prisma_client.health_check() + actual = { + "result": result, + "query_raw_called": prisma_client.db.query_raw.await_count, + "query_sql": prisma_client.db.query_raw.await_args.args[0], + "type": type(result).__name__, + } + assert actual == { + "result": [{"?column?": 1}], + "query_raw_called": 1, + "query_sql": "SELECT 1", + "type": "list", + } + + +@pytest.mark.asyncio +async def test_health_check_raises_when_query_raw_fails( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.query_raw = AsyncMock(side_effect=RuntimeError("connection refused")) + with pytest.raises(RuntimeError, match="connection refused"): + await prisma_client.health_check() + + +@pytest.mark.asyncio +async def test_get_spend_logs_row_count_returns_int_from_pg_class( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.query_raw = AsyncMock(return_value=[{"reltuples": 12345}]) + result = await prisma_client._get_spend_logs_row_count() + actual = { + "result": result, + "query_count": prisma_client.db.query_raw.await_count, + "query_kwargs": prisma_client.db.query_raw.await_args.kwargs, + "type": type(result).__name__, + } + assert actual == { + "result": 12345, + "query_count": 1, + "query_kwargs": { + "query": prisma_client.db.query_raw.await_args.kwargs["query"] + }, + "type": "int", + } + + +@pytest.mark.asyncio +async def test_get_spend_logs_row_count_error_falls_back_to_zero( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.query_raw = AsyncMock(side_effect=RuntimeError("perm denied")) + assert await prisma_client._get_spend_logs_row_count() == 0 + + +@pytest.mark.asyncio +async def test_set_spend_logs_row_count_in_proxy_state_writes_to_state( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + fake_state = MagicMock() + fake_state.set_proxy_state_variable = MagicMock() + + import litellm.proxy.proxy_server as proxy_server_mod + + monkeypatch.setattr(proxy_server_mod, "proxy_state", fake_state, raising=False) + + prisma_client._get_spend_logs_row_count = AsyncMock(return_value=99) + await prisma_client._set_spend_logs_row_count_in_proxy_state() + kwargs = fake_state.set_proxy_state_variable.call_args.kwargs + assert kwargs == {"variable_name": "spend_logs_row_count", "value": 99} + + +@pytest.mark.asyncio +async def test_set_spend_logs_row_count_error_raises_through_backoff( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + fake_state = MagicMock() + fake_state.set_proxy_state_variable = MagicMock(side_effect=RuntimeError("boom")) + import litellm.proxy.proxy_server as proxy_server_mod + + monkeypatch.setattr(proxy_server_mod, "proxy_state", fake_state, raising=False) + + prisma_client._get_spend_logs_row_count = AsyncMock(return_value=1) + with pytest.raises(RuntimeError, match="boom"): + await prisma_client._set_spend_logs_row_count_in_proxy_state() + + +def test_validate_response_time_passes_finite_value(prisma_client: PrismaClient) -> None: + inputs = { + "ok": prisma_client._validate_response_time(123.45), + "none": prisma_client._validate_response_time(None), + "inf": prisma_client._validate_response_time(float("inf")), + "neg_inf": prisma_client._validate_response_time(float("-inf")), + "nan": prisma_client._validate_response_time(float("nan")), + } + assert inputs == { + "ok": 123.45, + "none": None, + "inf": None, + "neg_inf": None, + "nan": None, + } + + +def test_validate_response_time_invalid_string_returns_none( + prisma_client: PrismaClient, +) -> None: + """Non-numeric input is logged and returned as None. The name is the + error hint; the input itself is invalid, not a thrown exception.""" + assert prisma_client._validate_response_time("not-a-float") is None + + +def test_clean_details_round_trips_json(prisma_client: PrismaClient) -> None: + details = {"latency": 1.5, "ok": True, "error": None, "model": "gpt-4o"} + cleaned = prisma_client._clean_details(details) + pinned = { + "cleaned": cleaned, + "is_dict": isinstance(cleaned, dict), + "none_for_non_dict": prisma_client._clean_details("oops"), # type: ignore[arg-type] + "none_for_none": prisma_client._clean_details(None), + } + assert pinned == { + "cleaned": details, + "is_dict": True, + "none_for_non_dict": None, + "none_for_none": None, + } + + +def test_clean_details_invalid_payload_returns_none( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + """When ``safe_dumps`` itself blows up (e.g. an internal exception), the + error path swallows it and returns None. + """ + import litellm.proxy.utils as utils_mod + + def _explode(_: Any) -> str: + raise RuntimeError("safe_dumps broken") + + monkeypatch.setattr(utils_mod, "safe_dumps", _explode) + assert prisma_client._clean_details({"x": 1}) is None + + +@pytest.mark.asyncio +async def test_save_health_check_result_creates_record( + prisma_client: PrismaClient, +) -> None: + expected = MagicMock(name="HealthCheckRow") + prisma_client.db.litellm_healthchecktable.create = AsyncMock(return_value=expected) + result = await prisma_client.save_health_check_result( + model_name="gpt-4o", + status="healthy", + healthy_count=3, + unhealthy_count=0, + response_time_ms=150.0, + details={"latency": 1, "ok": True}, + checked_by="probe", + model_id="m-1", + ) + data = prisma_client.db.litellm_healthchecktable.create.await_args.kwargs["data"] + pinned = { + "returned": result, + "model_name": data["model_name"], + "status": data["status"], + "healthy_count": data["healthy_count"], + "response_time_ms": data["response_time_ms"], + "details": data["details"], + "checked_by": data["checked_by"], + "model_id": data["model_id"], + } + assert pinned == { + "returned": expected, + "model_name": "gpt-4o", + "status": "healthy", + "healthy_count": 3, + "response_time_ms": 150.0, + "details": {"latency": 1, "ok": True}, + "checked_by": "probe", + "model_id": "m-1", + } + + +@pytest.mark.asyncio +async def test_save_health_check_result_db_failure_returns_none( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.litellm_healthchecktable.create = AsyncMock( + side_effect=RuntimeError("db down") + ) + result = await prisma_client.save_health_check_result( + model_name="gpt-4o", status="healthy" + ) + assert result is None + + +@pytest.mark.asyncio +async def test_get_health_check_history_filters_by_model_and_status( + prisma_client: PrismaClient, +) -> None: + rows = [MagicMock(name=f"row-{i}") for i in range(2)] + prisma_client.db.litellm_healthchecktable.find_many = AsyncMock(return_value=rows) + result = await prisma_client.get_health_check_history( + model_name="gpt-4o", limit=5, offset=10, status_filter="healthy" + ) + kwargs = prisma_client.db.litellm_healthchecktable.find_many.await_args.kwargs + actual = { + "result_len": len(result), + "where": kwargs["where"], + "order": kwargs["order"], + "take": kwargs["take"], + "skip": kwargs["skip"], + } + assert actual == { + "result_len": 2, + "where": {"model_name": "gpt-4o", "status": "healthy"}, + "order": {"checked_at": "desc"}, + "take": 5, + "skip": 10, + } + + +@pytest.mark.asyncio +async def test_get_health_check_history_db_error_returns_empty_list( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.litellm_healthchecktable.find_many = AsyncMock( + side_effect=RuntimeError("network down") + ) + assert await prisma_client.get_health_check_history() == [] + + +@pytest.mark.asyncio +async def test_get_all_latest_health_checks_uses_distinct( + prisma_client: PrismaClient, +) -> None: + rows = [MagicMock(name=f"row-{i}") for i in range(3)] + prisma_client.db.litellm_healthchecktable.find_many = AsyncMock(return_value=rows) + result = await prisma_client.get_all_latest_health_checks() + kwargs = prisma_client.db.litellm_healthchecktable.find_many.await_args.kwargs + actual = { + "len": len(result), + "distinct": kwargs["distinct"], + "order_len": len(kwargs["order"]), + "first_order": kwargs["order"][0], + } + assert actual == { + "len": 3, + "distinct": ["model_id", "model_name"], + "order_len": 3, + "first_order": {"model_id": "asc"}, + } + + +@pytest.mark.asyncio +async def test_get_all_latest_health_checks_db_error_returns_empty_list( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.litellm_healthchecktable.find_many = AsyncMock( + side_effect=RuntimeError("oops") + ) + assert await prisma_client.get_all_latest_health_checks() == [] diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_lifecycle.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_lifecycle.py new file mode 100644 index 00000000000..30fd4a74bb0 --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_lifecycle.py @@ -0,0 +1,207 @@ +"""Pin ``PrismaClient`` lifecycle methods. + +Symbols pinned here: + - ``PrismaClient.__init__`` + - ``PrismaClient.writer_db`` + - ``PrismaClient.connect`` + - ``PrismaClient.disconnect`` +""" + +from __future__ import annotations + +import asyncio +from typing import Any +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy.utils import PrismaClient + + +@pytest.mark.asyncio +async def test_prismaclient_init_wires_default_config( + patched_prisma_import: MagicMock, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv("DATABASE_URL_READ_REPLICA", raising=False) + monkeypatch.delenv("IAM_TOKEN_DB_AUTH", raising=False) + monkeypatch.delenv("PRISMA_RECONNECT_COOLDOWN_SECONDS", raising=False) + monkeypatch.delenv("PRISMA_HEALTH_WATCHDOG_INTERVAL_SECONDS", raising=False) + monkeypatch.delenv("PRISMA_HEALTH_WATCHDOG_ENABLED", raising=False) + monkeypatch.delenv("PRISMA_RECONNECT_ESCALATION_THRESHOLD", raising=False) + + proxy_logging = MagicMock() + pc = PrismaClient( + database_url="postgres://x:y@h:5432/db", + proxy_logging_obj=proxy_logging, + ) + pinned = { + "iam_token_db_auth": pc.iam_token_db_auth, + "db_reconnect_cooldown_seconds": pc._db_reconnect_cooldown_seconds, + "db_health_watchdog_interval_seconds": pc._db_health_watchdog_interval_seconds, + "db_health_watchdog_enabled": pc._db_health_watchdog_enabled, + "reconnect_escalation_threshold": pc._reconnect_escalation_threshold, + "consecutive_reconnect_failures": pc._consecutive_reconnect_failures, + "engine_pid": pc._engine_pid, + "watching_engine": pc._watching_engine, + "proxy_logging_obj_set": pc.proxy_logging_obj is proxy_logging, + "db_reconnect_lock_is_lock": isinstance(pc._db_reconnect_lock, asyncio.Lock), + } + assert pinned == { + "iam_token_db_auth": None, + "db_reconnect_cooldown_seconds": 15, + "db_health_watchdog_interval_seconds": 30, + "db_health_watchdog_enabled": True, + "reconnect_escalation_threshold": 3, + "consecutive_reconnect_failures": 0, + "engine_pid": 0, + "watching_engine": False, + "proxy_logging_obj_set": True, + "db_reconnect_lock_is_lock": True, + } + + +def test_prismaclient_init_honors_env_overrides( + patched_prisma_import: MagicMock, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("PRISMA_RECONNECT_COOLDOWN_SECONDS", "42") + monkeypatch.setenv("PRISMA_HEALTH_WATCHDOG_INTERVAL_SECONDS", "60") + monkeypatch.setenv("PRISMA_HEALTH_WATCHDOG_ENABLED", "false") + monkeypatch.setenv("PRISMA_RECONNECT_ESCALATION_THRESHOLD", "7") + monkeypatch.delenv("DATABASE_URL_READ_REPLICA", raising=False) + monkeypatch.delenv("IAM_TOKEN_DB_AUTH", raising=False) + + pc = PrismaClient( + database_url="postgres://x:y@h:5432/db", + proxy_logging_obj=MagicMock(), + ) + pinned = { + "db_reconnect_cooldown_seconds": pc._db_reconnect_cooldown_seconds, + "db_health_watchdog_interval_seconds": pc._db_health_watchdog_interval_seconds, + "db_health_watchdog_enabled": pc._db_health_watchdog_enabled, + "reconnect_escalation_threshold": pc._reconnect_escalation_threshold, + } + assert pinned == { + "db_reconnect_cooldown_seconds": 42, + "db_health_watchdog_interval_seconds": 60, + "db_health_watchdog_enabled": False, + "reconnect_escalation_threshold": 7, + } + + +def test_prismaclient_init_raises_when_prisma_not_generated() -> None: + """If ``from prisma import Prisma`` fails, the init re-raises with the + 'prisma generate' guidance message. + """ + import prisma as _prisma_pkg + + had_prisma_attr = "Prisma" in _prisma_pkg.__dict__ + previous_prisma_attr = _prisma_pkg.__dict__.get("Prisma") + if had_prisma_attr: + del _prisma_pkg.Prisma # type: ignore[attr-defined] + try: + with pytest.raises(Exception, match="prisma generate"): + PrismaClient( + database_url="postgres://x:y@h:5432/db", + proxy_logging_obj=MagicMock(), + ) + finally: + if had_prisma_attr: + _prisma_pkg.Prisma = previous_prisma_attr # type: ignore[attr-defined] + + +def test_writer_db_returns_db_when_no_routing(prisma_client: PrismaClient) -> None: + actual = { + "writer_is_db": prisma_client.writer_db is prisma_client.db, + "type_consistency": type(prisma_client.writer_db) is type(prisma_client.db), + "callable_query_raw": callable(prisma_client.writer_db.query_raw), + } + assert actual == { + "writer_is_db": True, + "type_consistency": True, + "callable_query_raw": True, + } + + +def test_writer_db_unwraps_routing_wrapper(prisma_client: PrismaClient) -> None: + from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper + + inner_writer = MagicMock(name="WriterInsideRouter") + + class _FakeRouting(RoutingPrismaWrapper): # type: ignore[misc] + def __init__(self) -> None: + self._writer = inner_writer + + prisma_client.db = _FakeRouting() + assert prisma_client.writer_db is inner_writer + + +def test_writer_db_error_when_db_attribute_missing(prisma_client: PrismaClient) -> None: + del prisma_client.db + with pytest.raises(AttributeError): + _ = prisma_client.writer_db + + +@pytest.mark.asyncio +async def test_connect_invokes_underlying_when_disconnected( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.is_connected = MagicMock(return_value=False) + prisma_client.db.connect = AsyncMock() + await prisma_client.connect() + actual = { + "connect_called": prisma_client.db.connect.await_count, + "is_connected_called": prisma_client.db.is_connected.call_count, + "no_failure_handler": prisma_client.proxy_logging_obj.failure_handler.await_count, + } + assert actual == { + "connect_called": 1, + "is_connected_called": 1, + "no_failure_handler": 0, + } + + +@pytest.mark.asyncio +async def test_connect_is_noop_when_already_connected( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.is_connected = MagicMock(return_value=True) + prisma_client.db.connect = AsyncMock() + await prisma_client.connect() + assert prisma_client.db.connect.await_count == 0 + + +@pytest.mark.asyncio +async def test_connect_invokes_failure_handler_and_raises_on_error( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.is_connected = MagicMock(return_value=False) + prisma_client.db.connect = AsyncMock(side_effect=RuntimeError("network down")) + with pytest.raises(RuntimeError, match="network down"): + await prisma_client.connect() + + +@pytest.mark.asyncio +async def test_disconnect_calls_underlying(prisma_client: PrismaClient) -> None: + prisma_client.db.disconnect = AsyncMock() + await prisma_client.disconnect() + actual = { + "disconnect_called": prisma_client.db.disconnect.await_count, + "failure_handler_called": prisma_client.proxy_logging_obj.failure_handler.await_count, + "type": type(prisma_client.db.disconnect).__name__, + } + assert actual == { + "disconnect_called": 1, + "failure_handler_called": 0, + "type": "AsyncMock", + } + + +@pytest.mark.asyncio +async def test_disconnect_raises_when_underlying_fails( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.disconnect = AsyncMock(side_effect=RuntimeError("disconnect boom")) + with pytest.raises(RuntimeError, match="disconnect boom"): + await prisma_client.disconnect() diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_reconnect.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_reconnect.py new file mode 100644 index 00000000000..f669e6be88d --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_reconnect.py @@ -0,0 +1,371 @@ +"""Pin ``PrismaClient`` reconnect + watchdog symbols. + +Symbols pinned here: + - ``PrismaClient._run_reconnect_cycle`` + - ``PrismaClient._attempt_reconnect_inside_lock`` + - ``PrismaClient.attempt_db_reconnect`` + - ``PrismaClient.start_db_health_watchdog_task`` + - ``PrismaClient.stop_db_health_watchdog_task`` + - ``PrismaClient._db_health_watchdog_loop`` +""" + +from __future__ import annotations + +import asyncio +from typing import Any +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy.utils import PrismaClient + + +@pytest.mark.asyncio +async def test_run_reconnect_cycle_direct_path_when_engine_alive( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("DATABASE_URL", "postgres://x:y@h:5432/db") + prisma_client._engine_confirmed_dead = False + prisma_client._engine_pid = 0 + prisma_client.db.recreate_prisma_client = AsyncMock() + prisma_client._start_engine_watcher = AsyncMock() + prisma_client._cleanup_engine_watcher = MagicMock() + + writer = MagicMock() + writer.query_raw = AsyncMock(return_value=[{"?column?": 1}]) + monkeypatch.setattr( + PrismaClient, + "writer_db", + property(lambda self: writer), + ) + + await prisma_client._run_reconnect_cycle(timeout_seconds=5) + pinned = { + "recreate_called": prisma_client.db.recreate_prisma_client.await_count, + "start_watcher_called": prisma_client._start_engine_watcher.await_count, + "writer_smoke_test_called": writer.query_raw.await_count, + "engine_confirmed_dead": prisma_client._engine_confirmed_dead, + } + assert pinned == { + "recreate_called": 1, + "start_watcher_called": 1, + "writer_smoke_test_called": 1, + "engine_confirmed_dead": False, + } + + +@pytest.mark.asyncio +async def test_run_reconnect_cycle_heavy_path_when_engine_dead( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("DATABASE_URL", "postgres://x:y@h:5432/db") + prisma_client._engine_confirmed_dead = True + prisma_client._engine_pid = 1234 + prisma_client.db.recreate_prisma_client = AsyncMock() + prisma_client._start_engine_watcher = AsyncMock() + prisma_client._cleanup_engine_watcher = MagicMock() + monkeypatch.setattr(PrismaClient, "_reap_all_zombies", staticmethod(lambda: set())) + + await prisma_client._run_reconnect_cycle(timeout_seconds=5) + pinned = { + "recreate_called": prisma_client.db.recreate_prisma_client.await_count, + "start_watcher_called": prisma_client._start_engine_watcher.await_count, + "cleanup_called": prisma_client._cleanup_engine_watcher.call_count, + "dead_flag_cleared": prisma_client._engine_confirmed_dead, + } + assert pinned == { + "recreate_called": 1, + "start_watcher_called": 1, + "cleanup_called": 1, + "dead_flag_cleared": False, + } + + +@pytest.mark.asyncio +async def test_run_reconnect_cycle_raises_when_database_url_missing( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.delenv("DATABASE_URL", raising=False) + with pytest.raises(RuntimeError, match="DATABASE_URL not set"): + await prisma_client._run_reconnect_cycle(timeout_seconds=1) + + +@pytest.mark.asyncio +async def test_attempt_reconnect_inside_lock_runs_cycle_and_resets_counter( + prisma_client: PrismaClient, +) -> None: + prisma_client._db_last_reconnect_attempt_ts = 0.0 + prisma_client._consecutive_reconnect_failures = 2 + prisma_client._run_reconnect_cycle = AsyncMock() + + ok = await prisma_client._attempt_reconnect_inside_lock( + force=True, reason="test", timeout_seconds=1 + ) + pinned = { + "returned": ok, + "cycle_called": prisma_client._run_reconnect_cycle.await_count, + "failures_reset": prisma_client._consecutive_reconnect_failures, + } + assert pinned == { + "returned": True, + "cycle_called": 1, + "failures_reset": 0, + } + + +@pytest.mark.asyncio +async def test_attempt_reconnect_inside_lock_skips_when_in_cooldown( + prisma_client: PrismaClient, +) -> None: + import time + + prisma_client._db_reconnect_cooldown_seconds = 60 + prisma_client._db_last_reconnect_attempt_ts = time.time() + prisma_client._run_reconnect_cycle = AsyncMock() + + ok = await prisma_client._attempt_reconnect_inside_lock( + force=False, reason="test", timeout_seconds=1 + ) + assert ok is False + assert prisma_client._run_reconnect_cycle.await_count == 0 + + +@pytest.mark.asyncio +async def test_attempt_reconnect_inside_lock_increments_failure_counter_on_error( + prisma_client: PrismaClient, +) -> None: + prisma_client._db_last_reconnect_attempt_ts = 0.0 + prisma_client._consecutive_reconnect_failures = 0 + prisma_client._run_reconnect_cycle = AsyncMock(side_effect=RuntimeError("boom")) + + ok = await prisma_client._attempt_reconnect_inside_lock( + force=True, reason="failing_test", timeout_seconds=1 + ) + assert ok is False + assert prisma_client._consecutive_reconnect_failures == 1 + + +@pytest.mark.asyncio +async def test_attempt_db_reconnect_force_runs_under_lock( + prisma_client: PrismaClient, +) -> None: + prisma_client._db_last_reconnect_attempt_ts = 0.0 + prisma_client._attempt_reconnect_inside_lock = AsyncMock(return_value=True) + + result = await prisma_client.attempt_db_reconnect(reason="explicit", force=True) + args = prisma_client._attempt_reconnect_inside_lock.await_args + pinned = { + "returned": result, + "calls": prisma_client._attempt_reconnect_inside_lock.await_count, + "passed_force": args.args[0], + "passed_reason": args.args[1], + "passed_timeout": args.args[2], + } + assert pinned == { + "returned": True, + "calls": 1, + "passed_force": True, + "passed_reason": "explicit", + "passed_timeout": None, + } + + +@pytest.mark.asyncio +async def test_attempt_db_reconnect_lock_timeout_returns_false( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + """A reconnect attempt that can't acquire the lock within + ``lock_timeout_seconds`` returns False without running the cycle. + + The production code creates an inner task, races it against the + timeout via ``asyncio.wait``, then cancels and awaits the loser. + Under coverage instrumentation on Python 3.11 the CancelledError from + a freshly-cancelled task can outrun the surrounding ``except`` block, + so this test pre-completes the inner task (no cancellation happens) + by replacing ``asyncio.wait`` with a callable that returns the loser + task as still-pending after it's already been completed elsewhere. + """ + completed_task: asyncio.Task[bool] = asyncio.get_running_loop().create_task( + _no_op_returning_true() + ) + # Ensure the inner task has finished before attempt_db_reconnect sees it. + await completed_task + + async def _wait_returns_loser(_tasks: Any, **kwargs: Any) -> Any: + return set(), {completed_task} + + monkeypatch.setattr("asyncio.wait", _wait_returns_loser) + monkeypatch.setattr( + asyncio, + "create_task", + lambda coro, *a, **kw: (coro.close() or completed_task), + ) + + prisma_client._db_last_reconnect_attempt_ts = 0.0 + prisma_client._attempt_reconnect_inside_lock = AsyncMock() + + ok = await prisma_client.attempt_db_reconnect( + reason="lock_busy", + lock_timeout_seconds=0.0, + ) + assert ok is False + assert prisma_client._attempt_reconnect_inside_lock.await_count == 0 + + +async def _no_op_returning_true() -> bool: + return True + + +@pytest.mark.asyncio +async def test_attempt_db_reconnect_skips_in_cooldown_returns_false( + prisma_client: PrismaClient, +) -> None: + import time + + prisma_client._db_reconnect_cooldown_seconds = 60 + prisma_client._db_last_reconnect_attempt_ts = time.time() + ok = await prisma_client.attempt_db_reconnect(reason="cooled_down") + assert ok is False + + +@pytest.mark.asyncio +async def test_start_db_health_watchdog_task_creates_loop_task( + prisma_client: PrismaClient, +) -> None: + prisma_client._db_health_watchdog_enabled = True + prisma_client._db_health_watchdog_task = None + prisma_client._start_engine_watcher = AsyncMock() + prisma_client._db_health_watchdog_loop = AsyncMock(return_value=None) + + await prisma_client.start_db_health_watchdog_task() + task = prisma_client._db_health_watchdog_task + # Yield control so the just-scheduled task actually invokes the loop mock. + await asyncio.sleep(0) + pinned = { + "task_type": type(task).__name__, + "watcher_started": prisma_client._start_engine_watcher.await_count, + "loop_invoked": prisma_client._db_health_watchdog_loop.await_count, + } + if task is not None: + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + assert pinned == { + "task_type": "Task", + "watcher_started": 1, + "loop_invoked": 1, + } + + +@pytest.mark.asyncio +async def test_start_db_health_watchdog_task_disabled_short_circuits( + prisma_client: PrismaClient, +) -> None: + prisma_client._db_health_watchdog_enabled = False + prisma_client._start_engine_watcher = AsyncMock() + await prisma_client.start_db_health_watchdog_task() + assert prisma_client._db_health_watchdog_task is None + assert prisma_client._start_engine_watcher.await_count == 0 + + +@pytest.mark.asyncio +async def test_stop_db_health_watchdog_task_cancels_and_clears( + prisma_client: PrismaClient, +) -> None: + prisma_client._stop_engine_watcher = MagicMock() + + cancel_called = {"n": 0} + + class _FakeTask: + def cancel(self) -> None: + cancel_called["n"] += 1 + + def __await__(self): + return iter([]) + + prisma_client._db_health_watchdog_task = _FakeTask() # type: ignore[assignment] + + await prisma_client.stop_db_health_watchdog_task() + pinned = { + "task_cleared": prisma_client._db_health_watchdog_task, + "engine_stop_called": prisma_client._stop_engine_watcher.call_count, + "cancel_called": cancel_called["n"], + "no_failure": True, + } + assert pinned == { + "task_cleared": None, + "engine_stop_called": 1, + "cancel_called": 1, + "no_failure": True, + } + + +@pytest.mark.asyncio +async def test_stop_db_health_watchdog_task_noop_when_no_task( + prisma_client: PrismaClient, +) -> None: + prisma_client._db_health_watchdog_task = None + prisma_client._stop_engine_watcher = MagicMock(side_effect=RuntimeError("err")) + with pytest.raises(RuntimeError, match="err"): + await prisma_client.stop_db_health_watchdog_task() + + +@pytest.mark.asyncio +async def test_db_health_watchdog_loop_triggers_reconnect_on_timeout( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + """The watchdog loop reconnects when ``wait_for`` raises TimeoutError + or a recognized DB connection error. + """ + prisma_client._db_health_watchdog_interval_seconds = 0 + prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) + + call_count = {"n": 0} + + async def _timeout_then_cancel(*args: Any, **kwargs: Any) -> None: + call_count["n"] += 1 + if call_count["n"] >= 2: + raise asyncio.CancelledError() + raise asyncio.TimeoutError() + + monkeypatch.setattr("asyncio.wait_for", _timeout_then_cancel) + await prisma_client._db_health_watchdog_loop() + pinned = { + "reconnect_called": prisma_client.attempt_db_reconnect.await_count, + "reconnect_reason": prisma_client.attempt_db_reconnect.await_args.kwargs[ + "reason" + ], + "wait_for_calls": call_count["n"], + "loop_exited_clean": True, + } + assert pinned == { + "reconnect_called": 1, + "reconnect_reason": "db_health_watchdog_connection_error", + "wait_for_calls": 2, + "loop_exited_clean": True, + } + + +@pytest.mark.asyncio +async def test_db_health_watchdog_loop_swallows_non_db_errors( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + """A non-DB error during the probe should NOT trigger reconnect; the + loop logs and continues until cancellation. + """ + prisma_client._db_health_watchdog_interval_seconds = 0 + prisma_client.attempt_db_reconnect = AsyncMock() + + call_count = {"n": 0} + + async def _raise_then_cancel(*args: Any, **kwargs: Any) -> None: + call_count["n"] += 1 + if call_count["n"] >= 2: + raise asyncio.CancelledError() + raise ValueError("not a db error") + + monkeypatch.setattr("asyncio.wait_for", _raise_then_cancel) + await prisma_client._db_health_watchdog_loop() + assert prisma_client.attempt_db_reconnect.await_count == 0 diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_writes.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_writes.py new file mode 100644 index 00000000000..4e547b81acc --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_writes.py @@ -0,0 +1,260 @@ +"""Pin ``PrismaClient`` write-side data operations. + +Symbols pinned here: + - ``PrismaClient.insert_data`` + - ``PrismaClient.update_data`` + - ``PrismaClient.delete_data`` +""" + +from __future__ import annotations + +import hashlib +import json +from types import SimpleNamespace +from typing import Any +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi import HTTPException + +from litellm.proxy.utils import PrismaClient + + +@pytest.mark.asyncio +async def test_insert_data_hashes_token_and_upserts(prisma_client: PrismaClient) -> None: + token = "sk-secret-1" + response = SimpleNamespace(token=hashlib.sha256(token.encode()).hexdigest(), + key_alias="alias", user_id="u1") + prisma_client.db.litellm_verificationtoken.upsert = AsyncMock(return_value=response) + data = { + "token": token, + "user_id": "u1", + "team_id": "t1", + "metadata": {"a": 1}, + } + result = await prisma_client.insert_data(data=data, table_name="key") + upsert_kwargs = prisma_client.db.litellm_verificationtoken.upsert.await_args.kwargs + actual = { + "returned": result, + "where": upsert_kwargs["where"], + "include": upsert_kwargs["include"], + "create_token": upsert_kwargs["data"]["create"]["token"], + "create_metadata_serialized": isinstance( + upsert_kwargs["data"]["create"]["metadata"], str + ), + "update_empty": upsert_kwargs["data"]["update"], + } + expected_hash = hashlib.sha256(token.encode()).hexdigest() + assert actual == { + "returned": response, + "where": {"token": expected_hash}, + "include": {"litellm_budget_table": True}, + "create_token": expected_hash, + "create_metadata_serialized": True, + "update_empty": {}, + } + + +@pytest.mark.asyncio +async def test_insert_data_strips_null_budget_limits(prisma_client: PrismaClient) -> None: + prisma_client.db.litellm_verificationtoken.upsert = AsyncMock(return_value=None) + await prisma_client.insert_data( + data={"token": "sk-1", "budget_limits": None}, table_name="key" + ) + create_payload = prisma_client.db.litellm_verificationtoken.upsert.await_args.kwargs[ + "data" + ]["create"] + assert "budget_limits" not in create_payload + + +@pytest.mark.asyncio +async def test_insert_data_team_serializes_members(prisma_client: PrismaClient) -> None: + prisma_client.db.litellm_teamtable.upsert = AsyncMock( + return_value=SimpleNamespace(team_id="t1", team_alias="x", spend=0) + ) + data = { + "team_id": "t1", + "team_alias": "x", + "members_with_roles": [{"role": "admin", "user_id": "u1"}], + } + result = await prisma_client.insert_data(data=data, table_name="team") + create_payload = prisma_client.db.litellm_teamtable.upsert.await_args.kwargs["data"][ + "create" + ] + assert result.team_id == "t1" + assert create_payload["members_with_roles"] == json.dumps(data["members_with_roles"]) + assert create_payload["team_id"] == "t1" + + +@pytest.mark.asyncio +async def test_insert_data_user_organization_fk_raises_400( + prisma_client: PrismaClient, +) -> None: + err = RuntimeError( + "Foreign key constraint failed on the field: `LiteLLM_UserTable_organization_id_fkey (index)`" + ) + prisma_client.db.litellm_usertable.upsert = AsyncMock(side_effect=err) + with pytest.raises(HTTPException) as excinfo: + await prisma_client.insert_data( + data={"user_id": "u1", "organization_id": "org-bad"}, table_name="user" + ) + raised = excinfo.value + assert "Foreign Key Constraint failed" in raised.detail["error"] + assert raised.status_code == 400 + + +@pytest.mark.asyncio +async def test_insert_data_logs_and_raises_generic_error( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.litellm_verificationtoken.upsert = AsyncMock( + side_effect=RuntimeError("write boom") + ) + with pytest.raises(RuntimeError, match="write boom"): + await prisma_client.insert_data(data={"token": "sk-1"}, table_name="key") + + +@pytest.mark.asyncio +async def test_update_data_token_hashes_and_updates( + prisma_client: PrismaClient, +) -> None: + token = "sk-update-1" + response = SimpleNamespace( + token=hashlib.sha256(token.encode()).hexdigest(), + model_dump=lambda: { + "token": hashlib.sha256(token.encode()).hexdigest(), + "spend": 1.0, + "user_id": "u1", + }, + ) + prisma_client.db.litellm_verificationtoken.update = AsyncMock(return_value=response) + result = await prisma_client.update_data( + token=token, + data={"spend": 1.0}, + ) + update_kwargs = prisma_client.db.litellm_verificationtoken.update.await_args.kwargs + hashed = hashlib.sha256(token.encode()).hexdigest() + actual = { + "result": result, + "where": update_kwargs["where"], + "data_token": update_kwargs["data"]["token"], + "data_spend": update_kwargs["data"]["spend"], + } + assert actual == { + "result": { + "token": hashed, + "data": {"token": hashed, "spend": 1.0, "user_id": "u1"}, + }, + "where": {"token": hashed}, + "data_token": hashed, + "data_spend": 1.0, + } + + +@pytest.mark.asyncio +async def test_update_data_user_upsert_returns_user_envelope( + prisma_client: PrismaClient, +) -> None: + row = SimpleNamespace(user_id="u2", spend=2.0) + prisma_client.db.litellm_usertable.upsert = AsyncMock(return_value=row) + result = await prisma_client.update_data( + data={"user_id": "u2", "spend": 2.0}, + table_name="user", + ) + assert result == {"user_id": "u2", "data": row} + + +@pytest.mark.asyncio +async def test_update_data_team_serializes_members_when_list( + prisma_client: PrismaClient, +) -> None: + row = SimpleNamespace(team_id="t9", team_alias="x") + prisma_client.db.litellm_teamtable.upsert = AsyncMock(return_value=row) + members = [{"role": "admin", "user_id": "u1"}] + result = await prisma_client.update_data( + data={"team_id": "t9", "members_with_roles": members}, + update_key_values={"members_with_roles": members}, + table_name="team", + ) + upsert_kwargs = prisma_client.db.litellm_teamtable.upsert.await_args.kwargs + actual = { + "result_team_id": result["team_id"], + "result_data": result["data"], + "create_members": upsert_kwargs["data"]["create"]["members_with_roles"], + "update_members": upsert_kwargs["data"]["update"]["members_with_roles"], + } + assert actual == { + "result_team_id": "t9", + "result_data": row, + "create_members": json.dumps(members), + "update_members": json.dumps(members), + } + + +@pytest.mark.asyncio +async def test_update_data_logs_and_raises_on_error( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.litellm_verificationtoken.update = AsyncMock( + side_effect=RuntimeError("update fail") + ) + with pytest.raises(RuntimeError, match="update fail"): + await prisma_client.update_data(token="sk-x", data={"spend": 1.0}) + + +@pytest.mark.asyncio +async def test_delete_data_hashes_sk_tokens_and_calls_delete_many( + prisma_client: PrismaClient, +) -> None: + deleted = SimpleNamespace(count=2) + prisma_client.db.litellm_verificationtoken.delete_many = AsyncMock( + return_value=deleted + ) + tokens = ["sk-one", "sk-two", "raw-hashed-token"] + result = await prisma_client.delete_data(tokens=tokens) + where = prisma_client.db.litellm_verificationtoken.delete_many.await_args.kwargs[ + "where" + ] + expected_hashes = sorted( + [ + hashlib.sha256(b"sk-one").hexdigest(), + hashlib.sha256(b"sk-two").hexdigest(), + "raw-hashed-token", + ] + ) + actual = { + "deleted_keys_attr": result["deleted_keys"], + "where_keys": list(where.keys()), + "filter_in_sorted": sorted(where["token"]["in"]), + "delete_call_count": prisma_client.db.litellm_verificationtoken.delete_many.await_count, + } + assert actual == { + "deleted_keys_attr": deleted, + "where_keys": ["token"], + "filter_in_sorted": expected_hashes, + "delete_call_count": 1, + } + + +@pytest.mark.asyncio +async def test_delete_data_team_calls_team_delete_many( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.litellm_teamtable.delete_many = AsyncMock() + result = await prisma_client.delete_data( + team_id_list=["t1", "t2"], table_name="team" + ) + where = prisma_client.db.litellm_teamtable.delete_many.await_args.kwargs["where"] + assert result == {"deleted_teams": ["t1", "t2"]} + assert where == {"team_id": {"in": ["t1", "t2"]}} + + +@pytest.mark.asyncio +async def test_delete_data_logs_and_raises_on_error( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.litellm_verificationtoken.delete_many = AsyncMock( + side_effect=RuntimeError("delete fail") + ) + with pytest.raises(RuntimeError, match="delete fail"): + await prisma_client.delete_data(tokens=["sk-x"]) diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_proxy_update_spend.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_proxy_update_spend.py new file mode 100644 index 00000000000..6a4fd516c9b --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_proxy_update_spend.py @@ -0,0 +1,275 @@ +"""Pin ``ProxyUpdateSpend`` behavior. + +Symbols pinned here: + - ``ProxyUpdateSpend.update_end_user_spend`` + - ``ProxyUpdateSpend.update_spend_logs`` + - ``ProxyUpdateSpend.disable_spend_updates`` +""" + +from __future__ import annotations + +import asyncio +from typing import Any, Dict, List +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy.utils import ProxyUpdateSpend + + +class _AsyncCM: + def __init__(self, target: Any) -> None: + self.target = target + + async def __aenter__(self) -> Any: + return self.target + + async def __aexit__(self, *exc: Any) -> None: + return None + + +@pytest.mark.asyncio +async def test_update_end_user_spend_upserts_each_end_user( + mock_prisma_client: Any, +) -> None: + batcher = MagicMock() + batcher.litellm_endusertable.upsert = MagicMock() + transaction = MagicMock() + transaction.batch_ = lambda: _AsyncCM(batcher) + mock_prisma_client.db.tx = lambda timeout: _AsyncCM(transaction) + + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + + end_user_costs: Dict[str, float] = {"u_b": 1.0, "u_a": 0.5} + await ProxyUpdateSpend.update_end_user_spend( + n_retry_times=0, + prisma_client=mock_prisma_client, + proxy_logging_obj=proxy_logging, + end_user_list_transactions=end_user_costs, + ) + calls = batcher.litellm_endusertable.upsert.call_args_list + ordered_ids = [c.kwargs["where"]["user_id"] for c in calls] + creates = [c.kwargs["data"]["create"] for c in calls] + pinned = { + "upsert_count": len(calls), + "ordered_ids": ordered_ids, + "first_create_keys": sorted(creates[0].keys()), + "first_create_user_id": creates[0]["user_id"], + "first_create_spend": creates[0]["spend"], + } + assert pinned == { + "upsert_count": 2, + "ordered_ids": ["u_a", "u_b"], + "first_create_keys": sorted(["user_id", "spend", "blocked"]), + "first_create_user_id": "u_a", + "first_create_spend": 0.5, + } + + +@pytest.mark.asyncio +async def test_update_end_user_spend_retries_on_connection_error( + mock_prisma_client: Any, monkeypatch: pytest.MonkeyPatch +) -> None: + """``DB_CONNECTION_ERROR_TYPES`` failures should be retried with backoff; + once retries are exhausted, ``_raise_failed_update_spend_exception`` is + invoked and the original exception bubbles up. + """ + import httpx + import litellm.proxy.utils as utils_mod + + sleeps: list[float] = [] + + async def _fake_sleep(seconds: float) -> None: + sleeps.append(seconds) + + monkeypatch.setattr(utils_mod.asyncio, "sleep", _fake_sleep) + + err = httpx.ReadError("conn reset") + mock_prisma_client.db.tx = MagicMock(side_effect=err) + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + with pytest.raises(httpx.ReadError): + await ProxyUpdateSpend.update_end_user_spend( + n_retry_times=1, + prisma_client=mock_prisma_client, + proxy_logging_obj=proxy_logging, + end_user_list_transactions={"u": 1.0}, + ) + assert sleeps == [1.0] + + +@pytest.mark.asyncio +async def test_update_end_user_spend_non_connection_error_raises_immediately( + mock_prisma_client: Any, +) -> None: + mock_prisma_client.db.tx = MagicMock(side_effect=RuntimeError("unknown")) + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + with pytest.raises(RuntimeError, match="unknown"): + await ProxyUpdateSpend.update_end_user_spend( + n_retry_times=3, + prisma_client=mock_prisma_client, + proxy_logging_obj=proxy_logging, + end_user_list_transactions={"u": 1.0}, + ) + + +@pytest.mark.asyncio +async def test_update_spend_logs_writes_batches_via_create_many( + mock_prisma_client: Any, make_spend_log_row: Any +) -> None: + logs = [make_spend_log_row(request_id=f"r{i}", spend=float(i)) for i in range(3)] + mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock() + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + await ProxyUpdateSpend.update_spend_logs( + n_retry_times=0, + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + logs_to_process=logs, + ) + kwargs = mock_prisma_client.db.litellm_spendlogs.create_many.await_args.kwargs + pinned = { + "calls": mock_prisma_client.db.litellm_spendlogs.create_many.await_count, + "data_len": len(kwargs["data"]), + "skip_duplicates": kwargs["skip_duplicates"], + "first_request_id": kwargs["data"][0]["request_id"], + } + assert pinned == { + "calls": 1, + "data_len": 3, + "skip_duplicates": True, + "first_request_id": "r0", + } + + +@pytest.mark.asyncio +async def test_update_spend_logs_uses_spend_logs_url_when_set( + mock_prisma_client: Any, + make_spend_log_row: Any, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("SPEND_LOGS_URL", "http://writer.invalid") + writer = MagicMock() + writer.post = AsyncMock(return_value=MagicMock(status_code=200)) + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + logs = [make_spend_log_row(request_id="r1")] + await ProxyUpdateSpend.update_spend_logs( + n_retry_times=0, + prisma_client=mock_prisma_client, + db_writer_client=writer, + proxy_logging_obj=proxy_logging, + logs_to_process=logs, + ) + pinned = { + "post_calls": writer.post.await_count, + "url": writer.post.await_args.kwargs["url"], + "headers": writer.post.await_args.kwargs["headers"], + "create_many_calls": mock_prisma_client.db.litellm_spendlogs.create_many.await_count, + } + assert pinned == { + "post_calls": 1, + "url": "http://writer.invalid/spend/update", + "headers": {"Content-Type": "application/json"}, + "create_many_calls": 0, + } + + +@pytest.mark.asyncio +async def test_update_spend_logs_pops_logs_when_logs_to_process_is_none( + mock_prisma_client: Any, make_spend_log_row: Any +) -> None: + mock_prisma_client.spend_log_transactions = [ + make_spend_log_row(request_id="a"), + make_spend_log_row(request_id="b"), + ] + mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock() + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + await ProxyUpdateSpend.update_spend_logs( + n_retry_times=0, + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + ) + assert mock_prisma_client.spend_log_transactions == [] + assert mock_prisma_client.db.litellm_spendlogs.create_many.await_count == 1 + + +@pytest.mark.asyncio +async def test_update_spend_logs_failure_raises_after_retries( + mock_prisma_client: Any, + make_spend_log_row: Any, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """When all retries exhaust the underlying DB error, the helper raises + via ``_raise_failed_update_spend_exception``. + """ + import httpx + import litellm.proxy.utils as utils_mod + + async def _fake_sleep(_: float) -> None: + return None + + monkeypatch.setattr(utils_mod.asyncio, "sleep", _fake_sleep) + + mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock( + side_effect=httpx.ReadError("network blip") + ) + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + with pytest.raises(httpx.ReadError): + await ProxyUpdateSpend.update_spend_logs( + n_retry_times=1, + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + logs_to_process=[make_spend_log_row(request_id="r1")], + ) + + +def test_disable_spend_updates_reflects_general_settings( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """The static method delegates to ``general_settings['disable_spend_updates']``; + flipping that value toggles the helper's return. + """ + import litellm.proxy.proxy_server as proxy_server_mod + + monkeypatch.setattr( + proxy_server_mod, "general_settings", {"disable_spend_updates": True} + ) + pinned = { + "with_flag_true": ProxyUpdateSpend.disable_spend_updates(), + "type_is_bool": isinstance(ProxyUpdateSpend.disable_spend_updates(), bool), + "method_is_static": isinstance( + ProxyUpdateSpend.__dict__["disable_spend_updates"], staticmethod + ), + } + assert pinned == { + "with_flag_true": True, + "type_is_bool": True, + "method_is_static": True, + } + + +def test_disable_spend_updates_default_false_without_flag( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import litellm.proxy.proxy_server as proxy_server_mod + + monkeypatch.setattr(proxy_server_mod, "general_settings", {}) + assert ProxyUpdateSpend.disable_spend_updates() is False + + +def test_disable_spend_updates_error_when_general_settings_unavailable( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import litellm.proxy.proxy_server as proxy_server_mod + + monkeypatch.delattr(proxy_server_mod, "general_settings", raising=False) + with pytest.raises(ImportError): + ProxyUpdateSpend.disable_spend_updates() diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_send_email.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_send_email.py new file mode 100644 index 00000000000..5028b65705f --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_send_email.py @@ -0,0 +1,105 @@ +"""Pin ``send_email``. + +Symbols pinned here: + - ``send_email`` +""" + +from __future__ import annotations + +from typing import Any + +import pytest + +from litellm.proxy.utils import send_email + + +@pytest.fixture(autouse=True) +def _smtp_env(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("SMTP_HOST", "smtp.invalid") + monkeypatch.setenv("SMTP_PORT", "2525") + monkeypatch.setenv("SMTP_USERNAME", "u") + monkeypatch.setenv("SMTP_PASSWORD", "p") + monkeypatch.setenv("SMTP_SENDER_EMAIL", "from@invalid") + monkeypatch.setenv("SMTP_TLS", "True") + + +@pytest.mark.asyncio +async def test_send_email_dispatches_via_smtp(in_memory_smtp: Any) -> None: + await send_email( + receiver_email="to@invalid", + subject="Hello", + html="

body

", + ) + assert len(in_memory_smtp.sent) == 1 + sent = in_memory_smtp.sent[0] + pinned = { + "from_addr": sent.from_addr, + "to_addrs": sent.to_addrs, + "subject": sent.subject, + "starttls": sent.starttls_called, + "login": sent.login_args, + } + assert pinned == { + "from_addr": "from@invalid", + "to_addrs": "to@invalid", + "subject": "Hello", + "starttls": True, + "login": ("u", "p"), + } + assert "

body

" in sent.body + + +@pytest.mark.asyncio +async def test_send_email_skips_starttls_when_disabled( + in_memory_smtp: Any, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("SMTP_TLS", "False") + await send_email( + receiver_email="to@invalid", + subject="Hi", + html="

x

", + ) + assert in_memory_smtp.sent[0].starttls_called is False + + +@pytest.mark.asyncio +async def test_send_email_error_missing_sender_email( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv("SMTP_SENDER_EMAIL", raising=False) + with pytest.raises(ValueError, match="SMTP_SENDER_EMAIL"): + await send_email( + receiver_email="x@y", subject="s", html="

h

" + ) + + +@pytest.mark.asyncio +async def test_send_email_error_missing_receiver() -> None: + with pytest.raises(ValueError, match="receiver email"): + await send_email(receiver_email=None, subject="s", html="

h

") + + +@pytest.mark.asyncio +async def test_send_email_error_missing_subject() -> None: + with pytest.raises(ValueError, match="subject"): + await send_email(receiver_email="x@y", subject=None, html="

h

") + + +@pytest.mark.asyncio +async def test_send_email_error_missing_html() -> None: + with pytest.raises(ValueError, match="HTML"): + await send_email(receiver_email="x@y", subject="s", html=None) + + +@pytest.mark.asyncio +async def test_send_email_smtp_failure_is_swallowed( + in_memory_smtp: Any, +) -> None: + """SMTP send_message errors are caught and logged; ``send_email`` itself + does not raise so a failing email never blocks the proxy. + """ + in_memory_smtp.raise_on_send = RuntimeError("smtp boom") + await send_email( + receiver_email="to@invalid", subject="Hi", html="

x

" + ) + assert in_memory_smtp.sent == [] diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py new file mode 100644 index 00000000000..a0b3af54750 --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py @@ -0,0 +1,360 @@ +"""Pin module-level spend functions. + +Symbols pinned here: + - ``update_spend`` + - ``update_daily_tag_spend`` + - ``update_spend_logs_job`` + - ``_monitor_spend_logs_queue`` + - ``_raise_failed_update_spend_exception`` +""" + +from __future__ import annotations + +import asyncio +from typing import Any, Dict, List +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy.utils import ( + _monitor_spend_logs_queue, + _raise_failed_update_spend_exception, + update_daily_tag_spend, + update_spend, + update_spend_logs_job, +) + + +@pytest.mark.asyncio +async def test_update_spend_invokes_writer_and_skips_empty_queue( + mock_prisma_client: Any, +) -> None: + proxy_logging = MagicMock() + proxy_logging.db_spend_update_writer = MagicMock() + proxy_logging.db_spend_update_writer.db_update_spend_transaction_handler = AsyncMock() + mock_prisma_client.spend_log_transactions = [] + + await update_spend( + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + ) + handler = proxy_logging.db_spend_update_writer.db_update_spend_transaction_handler + pinned = { + "handler_called": handler.await_count, + "handler_kwargs": handler.await_args.kwargs, + "queue_empty": mock_prisma_client.spend_log_transactions, + } + assert pinned == { + "handler_called": 1, + "handler_kwargs": { + "prisma_client": mock_prisma_client, + "n_retry_times": 3, + "proxy_logging_obj": proxy_logging, + }, + "queue_empty": [], + } + + +@pytest.mark.asyncio +async def test_update_spend_processes_logs_when_queue_nonempty( + mock_prisma_client: Any, make_spend_log_row: Any, monkeypatch: pytest.MonkeyPatch +) -> None: + proxy_logging = MagicMock() + proxy_logging.db_spend_update_writer = MagicMock() + proxy_logging.db_spend_update_writer.db_update_spend_transaction_handler = AsyncMock() + mock_prisma_client.spend_log_transactions = [make_spend_log_row(request_id="r1")] + + import litellm.proxy.utils as utils_mod + + job_mock = AsyncMock() + monkeypatch.setattr(utils_mod, "update_spend_logs_job", job_mock) + + await update_spend( + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + ) + assert job_mock.await_count == 1 + + +@pytest.mark.asyncio +async def test_update_spend_handler_failure_propagates( + mock_prisma_client: Any, +) -> None: + proxy_logging = MagicMock() + proxy_logging.db_spend_update_writer = MagicMock() + proxy_logging.db_spend_update_writer.db_update_spend_transaction_handler = AsyncMock( + side_effect=RuntimeError("handler down") + ) + with pytest.raises(RuntimeError, match="handler down"): + await update_spend( + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + ) + + +@pytest.mark.asyncio +async def test_update_daily_tag_spend_redis_path_when_buffered( + mock_prisma_client: Any, +) -> None: + proxy_logging = MagicMock() + writer = MagicMock() + proxy_logging.db_spend_update_writer = writer + writer.redis_update_buffer = MagicMock() + writer.redis_update_buffer._should_commit_spend_updates_to_redis = MagicMock( + return_value=True + ) + writer._commit_daily_tag_spend_to_db_with_redis = AsyncMock() + writer._commit_daily_tag_spend_to_db = AsyncMock() + + await update_daily_tag_spend( + prisma_client=mock_prisma_client, proxy_logging_obj=proxy_logging + ) + redis_kwargs = writer._commit_daily_tag_spend_to_db_with_redis.await_args.kwargs + pinned = { + "redis_calls": writer._commit_daily_tag_spend_to_db_with_redis.await_count, + "direct_calls": writer._commit_daily_tag_spend_to_db.await_count, + "redis_kwargs_keys": sorted(redis_kwargs.keys()), + "redis_n_retries": redis_kwargs["n_retry_times"], + } + assert pinned == { + "redis_calls": 1, + "direct_calls": 0, + "redis_kwargs_keys": sorted( + ["prisma_client", "n_retry_times", "proxy_logging_obj"] + ), + "redis_n_retries": 3, + } + + +@pytest.mark.asyncio +async def test_update_daily_tag_spend_direct_path_when_no_redis( + mock_prisma_client: Any, +) -> None: + proxy_logging = MagicMock() + writer = MagicMock() + proxy_logging.db_spend_update_writer = writer + writer.redis_update_buffer = MagicMock() + writer.redis_update_buffer._should_commit_spend_updates_to_redis = MagicMock( + return_value=False + ) + writer._commit_daily_tag_spend_to_db_with_redis = AsyncMock() + writer._commit_daily_tag_spend_to_db = AsyncMock() + + await update_daily_tag_spend( + prisma_client=mock_prisma_client, proxy_logging_obj=proxy_logging + ) + assert writer._commit_daily_tag_spend_to_db.await_count == 1 + assert writer._commit_daily_tag_spend_to_db_with_redis.await_count == 0 + + +@pytest.mark.asyncio +async def test_update_daily_tag_spend_logs_and_swallows_errors( + mock_prisma_client: Any, +) -> None: + """A failure in the commit path is logged but not re-raised; this matches + the historical behavior of this site (see plain ``logger.error`` rather + than ``spend_log_error``). + """ + proxy_logging = MagicMock() + proxy_logging.db_spend_update_writer = MagicMock() + proxy_logging.db_spend_update_writer.redis_update_buffer = MagicMock() + proxy_logging.db_spend_update_writer.redis_update_buffer._should_commit_spend_updates_to_redis = MagicMock( + return_value=False + ) + proxy_logging.db_spend_update_writer._commit_daily_tag_spend_to_db = AsyncMock( + side_effect=RuntimeError("commit boom") + ) + await update_daily_tag_spend( + prisma_client=mock_prisma_client, proxy_logging_obj=proxy_logging + ) + + +@pytest.mark.asyncio +async def test_update_spend_logs_job_skips_when_queue_empty( + mock_prisma_client: Any, +) -> None: + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + mock_prisma_client.spend_log_transactions = [] + mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock() + await update_spend_logs_job( + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + ) + assert mock_prisma_client.db.litellm_spendlogs.create_many.await_count == 0 + + +@pytest.mark.asyncio +async def test_update_spend_logs_job_processes_and_clears_queue( + mock_prisma_client: Any, make_spend_log_row: Any, monkeypatch: pytest.MonkeyPatch +) -> None: + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + mock_prisma_client.spend_log_transactions = [ + make_spend_log_row(request_id="r1"), + make_spend_log_row(request_id="r2"), + ] + mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock() + + # Stub auxiliary imports so the test focuses on the spend-logs write path. + import litellm.proxy.guardrails.usage_tracking as guard_mod + import litellm.proxy.db.spend_log_tool_index as tool_mod + + monkeypatch.setattr( + guard_mod, "process_spend_logs_guardrail_usage", AsyncMock(), raising=False + ) + monkeypatch.setattr( + tool_mod, "process_spend_logs_tool_usage", AsyncMock(), raising=False + ) + + await update_spend_logs_job( + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + ) + pinned = { + "create_many_calls": mock_prisma_client.db.litellm_spendlogs.create_many.await_count, + "queue_after": mock_prisma_client.spend_log_transactions, + "first_data_request_id": mock_prisma_client.db.litellm_spendlogs.create_many.await_args.kwargs[ + "data" + ][0]["request_id"], + "skip_duplicates_set": mock_prisma_client.db.litellm_spendlogs.create_many.await_args.kwargs[ + "skip_duplicates" + ], + } + assert pinned == { + "create_many_calls": 1, + "queue_after": [], + "first_data_request_id": "r1", + "skip_duplicates_set": True, + } + + +@pytest.mark.asyncio +async def test_monitor_spend_logs_queue_invokes_job_when_queue_nonempty( + mock_prisma_client: Any, + make_spend_log_row: Any, + monkeypatch: pytest.MonkeyPatch, +) -> None: + import litellm.proxy.utils as utils_mod + import litellm.constants as constants_mod + + monkeypatch.setattr(constants_mod, "SPEND_LOG_QUEUE_POLL_INTERVAL", 0.0, raising=False) + monkeypatch.setattr(constants_mod, "SPEND_LOG_QUEUE_SIZE_THRESHOLD", 1, raising=False) + proxy_logging = MagicMock() + mock_prisma_client.spend_log_transactions = [make_spend_log_row(request_id="r1")] + + cancel_after = {"n": 0} + + async def _fake_job(*args: Any, **kwargs: Any) -> None: + cancel_after["n"] += 1 + if cancel_after["n"] >= 1: + raise asyncio.CancelledError() + + monkeypatch.setattr(utils_mod, "update_spend_logs_job", _fake_job) + + with pytest.raises(asyncio.CancelledError): + await _monitor_spend_logs_queue( + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + ) + assert cancel_after["n"] == 1 + + +@pytest.mark.asyncio +async def test_monitor_spend_logs_queue_swallows_errors_and_backs_off( + mock_prisma_client: Any, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """An exception inside the loop is logged with backoff and the loop + continues running rather than crashing the monitor task. + """ + import litellm.proxy.utils as utils_mod + import litellm.constants as constants_mod + + monkeypatch.setattr(constants_mod, "SPEND_LOG_QUEUE_POLL_INTERVAL", 0.0, raising=False) + + sleep_count = {"n": 0} + + async def _short_sleep(_: float, *args: Any, **kwargs: Any) -> None: + sleep_count["n"] += 1 + if sleep_count["n"] >= 3: + raise asyncio.CancelledError() + + monkeypatch.setattr(utils_mod.asyncio, "sleep", _short_sleep) + proxy_logging = MagicMock() + + bad_lock = MagicMock() + bad_lock.__aenter__ = AsyncMock(side_effect=RuntimeError("lock broken")) + bad_lock.__aexit__ = AsyncMock(return_value=False) + mock_prisma_client._spend_log_transactions_lock = bad_lock + + with pytest.raises(asyncio.CancelledError): + await _monitor_spend_logs_queue( + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + ) + assert sleep_count["n"] == 3 + + +def test_raise_failed_update_spend_exception_emits_failure_handler() -> None: + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + + async def _runner() -> Any: + try: + _raise_failed_update_spend_exception( + e=RuntimeError("boom"), + start_time=0.0, + proxy_logging_obj=proxy_logging, + ) + except RuntimeError as e: + return e + return None + + err = asyncio.run(_runner()) + pinned = { + "raised": str(err), + "failure_handler_called": proxy_logging.failure_handler.call_count, + "call_type": ( + proxy_logging.failure_handler.call_args.kwargs.get("call_type") + if proxy_logging.failure_handler.call_args + else None + ), + "non_blocking_in_traceback": ( + "Non-Blocking" + in proxy_logging.failure_handler.call_args.kwargs["traceback_str"] + if proxy_logging.failure_handler.call_args + else False + ), + } + assert pinned == { + "raised": "boom", + "failure_handler_called": 1, + "call_type": "update_spend", + "non_blocking_in_traceback": True, + } + + +def test_raise_failed_update_spend_exception_raises_original_error() -> None: + """Error path: the function always re-raises the original exception so + the caller can observe the failure. + """ + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + + async def _runner() -> None: + _raise_failed_update_spend_exception( + e=ValueError("specific"), + start_time=0.0, + proxy_logging_obj=proxy_logging, + ) + + with pytest.raises(ValueError, match="specific"): + asyncio.run(_runner()) diff --git a/tests/test_litellm/proxy/utils/proxy_logging/__init__.py b/tests/test_litellm/proxy/utils/proxy_logging/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/proxy/utils/proxy_logging/_harness_smoke_test.py b/tests/test_litellm/proxy/utils/proxy_logging/_harness_smoke_test.py new file mode 100644 index 00000000000..1ec01f8c563 --- /dev/null +++ b/tests/test_litellm/proxy/utils/proxy_logging/_harness_smoke_test.py @@ -0,0 +1,56 @@ +"""Sanity tests for the proxy_logging conftest fixtures. + +Excluded from the pin-check by name. +""" + +from __future__ import annotations + +import pytest + + +def test_normalize_replaces_volatile_keys(normalize_fn): + raw = {"id": 7, "name": "x", "nested": {"created_at": 1, "value": 2}} + expected = {"id": "", "name": "x", "nested": {"created_at": "", "value": 2}} + assert normalize_fn(raw) == expected + + +def test_normalize_handles_lists(normalize_fn): + raw = [{"id": 1}, {"id": 2}] + assert normalize_fn(raw) == [{"id": ""}, {"id": ""}] + + +def test_mock_dual_cache_is_dual_cache(mock_dual_cache): + from litellm.caching.caching import DualCache + + assert isinstance(mock_dual_cache, DualCache) + + +def test_make_user_api_key_auth_returns_correct_type(make_user_api_key_auth): + from litellm.proxy._types import UserAPIKeyAuth + + auth = make_user_api_key_auth() + assert isinstance(auth, UserAPIKeyAuth) + assert auth.user_id == "test-user" + + +def test_make_user_api_key_auth_overrides_apply(make_user_api_key_auth): + auth = make_user_api_key_auth(user_id="custom-id") + assert auth.user_id == "custom-id" + + +def test_proxy_logging_fixture_is_initialized(proxy_logging): + from litellm.proxy.utils import InternalUsageCache, ProxyLogging + + assert isinstance(proxy_logging, ProxyLogging) + assert isinstance(proxy_logging.internal_usage_cache, InternalUsageCache) + assert proxy_logging.proxy_hook_mapping == {} + + +def test_make_mcp_request_obj_default(make_mcp_request_obj): + obj = make_mcp_request_obj() + assert obj.tool_name == "calculator" + assert obj.arguments == {"x": 1, "y": 2} + + +def test_mock_router_has_guardrail_list(mock_router): + assert mock_router.guardrail_list == [] diff --git a/tests/test_litellm/proxy/utils/proxy_logging/conftest.py b/tests/test_litellm/proxy/utils/proxy_logging/conftest.py new file mode 100644 index 00000000000..74508a74e3b --- /dev/null +++ b/tests/test_litellm/proxy/utils/proxy_logging/conftest.py @@ -0,0 +1,136 @@ +"""Shared fixtures for tests/test_litellm/proxy/utils/proxy_logging/. + +All fixtures used by PR1 of the proxy/utils.py behavior-pinning project +live here. Tests should not declare fixtures inline. +""" + +from __future__ import annotations + +import sys +from pathlib import Path +from typing import Any, Dict, Optional +from unittest.mock import MagicMock + +import pytest + +sys.path.insert(0, str(Path(__file__).resolve().parents[5])) + + +VOLATILE_KEYS = frozenset( + { + "created_at", + "updated_at", + "id", + "request_id", + "token", + "expires", + "expires_at", + "litellm_call_id", + "key_alias", + "created", + "start_time", + "end_time", + "duration", + "guardrail_start_time", + "guardrail_end_time", + "guardrail_duration", + } +) + + +def normalize(data: Any, volatile: frozenset = VOLATILE_KEYS) -> Any: + if isinstance(data, dict): + return { + k: ("" if k in volatile else normalize(v, volatile)) + for k, v in data.items() + } + if isinstance(data, list): + return [normalize(v, volatile) for v in data] + return data + + +@pytest.fixture +def mock_dual_cache(): + from litellm.caching.caching import DualCache + + cache = DualCache(default_in_memory_ttl=1) + return cache + + +@pytest.fixture +def mock_router(): + router = MagicMock() + router.guardrail_list = [] + router.get_available_guardrail = MagicMock(return_value={"callback": None}) + return router + + +@pytest.fixture +def mock_callbacks_disabled(monkeypatch): + """Disable all litellm callbacks for the duration of a test.""" + import litellm + + monkeypatch.setattr(litellm, "callbacks", []) + monkeypatch.setattr(litellm, "success_callback", []) + monkeypatch.setattr(litellm, "failure_callback", []) + monkeypatch.setattr(litellm, "_async_success_callback", []) + monkeypatch.setattr(litellm, "_async_failure_callback", []) + yield + + +@pytest.fixture +def make_user_api_key_auth(): + from litellm.proxy._types import UserAPIKeyAuth + + def _make(**overrides) -> UserAPIKeyAuth: + defaults: Dict[str, Any] = { + "api_key": "sk-test-1234", + "user_id": "test-user", + "team_id": "test-team", + "user_role": None, + "max_budget": None, + "spend": 0.0, + } + defaults.update(overrides) + return UserAPIKeyAuth(**defaults) + + return _make + + +@pytest.fixture +def proxy_logging(mock_callbacks_disabled): + """A wired-up ProxyLogging instance backed by a fresh DualCache. + + The fixture leaves it un-started; tests that need ``startup_event`` + should call it explicitly with the deps they want to control. + """ + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.utils import ProxyLogging + + return ProxyLogging(user_api_key_cache=UserApiKeyCache()) + + +@pytest.fixture +def normalize_fn(): + return normalize + + +@pytest.fixture +def make_mcp_request_obj(): + from litellm.types.llms.base import HiddenParams + from litellm.types.mcp import MCPPreCallRequestObject + + def _make( + tool_name: str = "calculator", + arguments: Optional[dict] = None, + server_name: Optional[str] = "math-server", + ) -> MCPPreCallRequestObject: + return MCPPreCallRequestObject( + tool_name=tool_name, + arguments=arguments if arguments is not None else {"x": 1, "y": 2}, + server_name=server_name, + user_api_key_auth={}, + hidden_params=HiddenParams(), + ) + + return _make diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_alerting.py b/tests/test_litellm/proxy/utils/proxy_logging/test_alerting.py new file mode 100644 index 00000000000..cede859cb38 --- /dev/null +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_alerting.py @@ -0,0 +1,262 @@ +"""Pin alerting helpers on ``ProxyLogging``. + +Covers ``failed_tracking_alert``, ``budget_alerts``, ``alerting_handler``, +``failure_handler``. +""" + +from __future__ import annotations + +from datetime import datetime +from typing import Any, Dict +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi import HTTPException + +import litellm +from litellm.proxy._types import AlertType, CallInfo + + +# --------------------------------------------------------------------------- +# failed_tracking_alert +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_failed_tracking_alert_no_op_when_alerting_none(proxy_logging): + proxy_logging.alerting = None + proxy_logging.slack_alerting_instance = MagicMock(failed_tracking_alert=AsyncMock()) + await proxy_logging.failed_tracking_alert(error_message="x", failing_model="m") + proxy_logging.slack_alerting_instance.failed_tracking_alert.assert_not_called() + + +@pytest.mark.asyncio +async def test_failed_tracking_alert_forwards_to_slack(proxy_logging): + proxy_logging.alerting = ["slack"] + captured: Dict[str, Any] = {} + + async def fake_alert(**kwargs): + captured.update(kwargs) + + proxy_logging.slack_alerting_instance = MagicMock(failed_tracking_alert=fake_alert) + await proxy_logging.failed_tracking_alert(error_message="db down", failing_model="gpt-4") + snapshot = { + "error_message": captured["error_message"], + "failing_model": captured["failing_model"], + "captured_keys": sorted(captured.keys()), + } + assert snapshot == { + "error_message": "db down", + "failing_model": "gpt-4", + "captured_keys": ["error_message", "failing_model"], + } + + +@pytest.mark.asyncio +async def test_failed_tracking_alert_slack_error_raises(proxy_logging): + proxy_logging.alerting = ["slack"] + proxy_logging.slack_alerting_instance = MagicMock( + failed_tracking_alert=AsyncMock(side_effect=RuntimeError("slack down")) + ) + with pytest.raises(RuntimeError): + await proxy_logging.failed_tracking_alert(error_message="x", failing_model="m") + + +# --------------------------------------------------------------------------- +# budget_alerts +# --------------------------------------------------------------------------- + + +def _user_info(alert_emails=None): + return CallInfo( + spend=0.0, + max_budget=1.0, + token="tok", + user_id="u1", + team_id="t1", + team_alias=None, + user_email=None, + key_alias=None, + projected_exceeded_date=None, + projected_spend=None, + event_group="user", + event="threshold_crossed", + alert_emails=alert_emails, + ) + + +@pytest.mark.asyncio +async def test_budget_alerts_no_op_when_alerting_off_and_no_emails(proxy_logging): + proxy_logging.alerting = None + proxy_logging.slack_alerting_instance = MagicMock(budget_alerts=AsyncMock()) + proxy_logging.email_logging_instance = MagicMock(budget_alerts=AsyncMock()) + await proxy_logging.budget_alerts(type="user_budget", user_info=_user_info()) + proxy_logging.slack_alerting_instance.budget_alerts.assert_not_called() + proxy_logging.email_logging_instance.budget_alerts.assert_not_called() + + +@pytest.mark.asyncio +async def test_budget_alerts_slack_when_slack_alerting(proxy_logging): + proxy_logging.alerting = ["slack"] + captured: Dict[str, Any] = {} + + async def fake_alert(**kwargs): + captured.update(kwargs) + + proxy_logging.slack_alerting_instance = MagicMock(budget_alerts=fake_alert) + proxy_logging.email_logging_instance = None + user_info = _user_info() + await proxy_logging.budget_alerts(type="user_budget", user_info=user_info) + snapshot = { + "type": captured["type"], + "user_info_is_callinfo": isinstance(captured["user_info"], CallInfo), + "user_id": captured["user_info"].user_id, + } + assert snapshot == {"type": "user_budget", "user_info_is_callinfo": True, "user_id": "u1"} + + +@pytest.mark.asyncio +async def test_budget_alerts_soft_budget_with_alert_emails_bypasses_global(proxy_logging): + proxy_logging.alerting = None + proxy_logging.slack_alerting_instance = MagicMock(budget_alerts=AsyncMock()) + proxy_logging.email_logging_instance = MagicMock(budget_alerts=AsyncMock()) + info = _user_info(alert_emails=["a@b.c"]) + await proxy_logging.budget_alerts(type="soft_budget", user_info=info) + proxy_logging.email_logging_instance.budget_alerts.assert_called_once() + proxy_logging.slack_alerting_instance.budget_alerts.assert_not_called() + + +@pytest.mark.asyncio +async def test_budget_alerts_slack_failure_raises(proxy_logging): + proxy_logging.alerting = ["slack"] + proxy_logging.slack_alerting_instance = MagicMock( + budget_alerts=AsyncMock(side_effect=ConnectionError("slack")) + ) + proxy_logging.email_logging_instance = None + with pytest.raises(ConnectionError): + await proxy_logging.budget_alerts(type="user_budget", user_info=_user_info()) + + +# --------------------------------------------------------------------------- +# alerting_handler +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_alerting_handler_no_op_when_alerting_is_none(proxy_logging): + proxy_logging.alerting = None + proxy_logging.slack_alerting_instance = MagicMock(send_alert=AsyncMock()) + await proxy_logging.alerting_handler(message="x", level="High", alert_type=AlertType.db_exceptions) + proxy_logging.slack_alerting_instance.send_alert.assert_not_called() + + +@pytest.mark.asyncio +async def test_alerting_handler_sends_to_slack(proxy_logging): + proxy_logging.alerting = ["slack"] + captured: Dict[str, Any] = {} + + async def fake_send(**kwargs): + captured.update(kwargs) + + proxy_logging.slack_alerting_instance = MagicMock(send_alert=fake_send) + await proxy_logging.alerting_handler( + message="hi", level="High", alert_type=AlertType.db_exceptions, request_data={"metadata": {}} + ) + snapshot = { + "message": captured["message"], + "level": captured["level"], + "alert_type": captured["alert_type"], + "user_info": captured["user_info"], + } + assert snapshot == { + "message": "hi", + "level": "High", + "alert_type": AlertType.db_exceptions, + "user_info": None, + } + + +@pytest.mark.asyncio +async def test_alerting_handler_sentry_without_sdk_error_raises(proxy_logging, monkeypatch): + proxy_logging.alerting = ["sentry"] + monkeypatch.setattr(litellm.utils, "sentry_sdk_instance", None) + with pytest.raises(Exception, match="SENTRY_DSN"): + await proxy_logging.alerting_handler(message="x", level="Low", alert_type=AlertType.db_exceptions) + + +# --------------------------------------------------------------------------- +# failure_handler +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_failure_handler_skips_when_db_exceptions_not_in_alert_types(proxy_logging): + proxy_logging.alert_types = ["llm_too_slow"] # type: ignore[list-item] + proxy_logging.alerting_handler = AsyncMock() + proxy_logging.service_logging_obj = MagicMock(async_service_failure_hook=AsyncMock()) + await proxy_logging.failure_handler(original_exception=Exception("x"), duration=1.0, call_type="db_read") + proxy_logging.alerting_handler.assert_not_called() + proxy_logging.service_logging_obj.async_service_failure_hook.assert_not_called() + + +@pytest.mark.asyncio +async def test_failure_handler_logs_db_error_and_calls_service_logging(proxy_logging, monkeypatch): + proxy_logging.alert_types = [AlertType.db_exceptions] + proxy_logging.alerting_handler = AsyncMock() + proxy_logging.service_logging_obj = MagicMock(async_service_failure_hook=AsyncMock()) + monkeypatch.setattr(litellm.utils, "capture_exception", None) + await proxy_logging.failure_handler( + original_exception=HTTPException(status_code=500, detail="boom"), + duration=1.5, + call_type="db_write", + ) + call_kwargs = proxy_logging.service_logging_obj.async_service_failure_hook.call_args.kwargs + snapshot = { + "service": call_kwargs["service"].value if hasattr(call_kwargs["service"], "value") else call_kwargs["service"], + "duration": call_kwargs["duration"], + "call_type": call_kwargs["call_type"], + } + assert snapshot == { + "service": "postgres", + "duration": 1.5, + "call_type": "db_write", + } + + +@pytest.mark.asyncio +async def test_failure_handler_with_capture_exception_invoked(proxy_logging, monkeypatch): + proxy_logging.alert_types = [AlertType.db_exceptions] + proxy_logging.alerting_handler = AsyncMock() + proxy_logging.service_logging_obj = MagicMock(async_service_failure_hook=AsyncMock()) + captured: Dict[str, Any] = {} + + def fake_capture(error): + captured["error"] = error + + monkeypatch.setattr(litellm.utils, "capture_exception", fake_capture) + err = RuntimeError("real") + await proxy_logging.failure_handler(original_exception=err, duration=1.0, call_type="db_read") + snapshot = { + "captured_is_input": captured["error"] is err, + "service_failure_called": proxy_logging.service_logging_obj.async_service_failure_hook.called, + "alerting_handler_scheduled": proxy_logging.alerting_handler.called, + } + assert snapshot == { + "captured_is_input": True, + "service_failure_called": True, + "alerting_handler_scheduled": True, + } + + +@pytest.mark.asyncio +async def test_failure_handler_propagates_service_logging_error_raises(proxy_logging, monkeypatch): + proxy_logging.alert_types = [AlertType.db_exceptions] + proxy_logging.alerting_handler = AsyncMock() + proxy_logging.service_logging_obj = MagicMock( + async_service_failure_hook=AsyncMock(side_effect=RuntimeError("svc")) + ) + monkeypatch.setattr(litellm.utils, "capture_exception", None) + with pytest.raises(RuntimeError): + await proxy_logging.failure_handler( + original_exception=Exception("x"), duration=0.0, call_type="db_read" + ) diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_callback_capabilities_class.py b/tests/test_litellm/proxy/utils/proxy_logging/test_callback_capabilities_class.py new file mode 100644 index 00000000000..45b81acbce1 --- /dev/null +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_callback_capabilities_class.py @@ -0,0 +1,338 @@ +"""Pin the ``ProxyLogging`` capability-probe family. + +Covers ``_callback_capabilities`` (the cached deriver), +``has_post_call_response_headers_callbacks``, ``has_streaming_callbacks``, +``has_streaming_chunk_hook_overrides``, ``needs_iterator_wrap``, +``needs_per_chunk_streaming_hook``, ``has_during_call_guardrails``, and +``get_combined_callback_list``. +""" + +from __future__ import annotations + +from typing import Any + +import pytest + +import litellm +from litellm.integrations.custom_logger import CustomLogger +from litellm.proxy.utils import ProxyLogging, _CallbackCapabilities + + +class _PlainLogger(CustomLogger): + pass + + +class _OverridesResponseHeaders(CustomLogger): + async def async_post_call_response_headers_hook(self, *args, **kwargs): # type: ignore[override] + return None + + +class _OverridesIterator(CustomLogger): + async def async_post_call_streaming_iterator_hook(self, *args, **kwargs): # type: ignore[override] + return None + + +class _OverridesPerChunk(CustomLogger): + async def async_post_call_streaming_hook(self, *args, **kwargs): # type: ignore[override] + return None + + +class _OverridesPreCall(CustomLogger): + async def async_pre_call_hook(self, *args, **kwargs): # type: ignore[override] + return None + + +@pytest.fixture(autouse=True) +def _clear_caps_cache(): + ProxyLogging._callback_capabilities_cache.clear() + yield + ProxyLogging._callback_capabilities_cache.clear() + + +def test_callback_capabilities_with_no_callbacks_returns_defaults(mock_callbacks_disabled): + caps = ProxyLogging._callback_capabilities() + snapshot = { + "headers": caps.has_post_call_response_headers, + "iterator": caps.has_iterator_override, + "chunk": caps.has_streaming_chunk_override, + "guardrail": caps.has_guardrail, + "pre_call": caps.has_pre_call_override, + "callbacks": caps.resolved_callbacks, + "overrides": caps.iterator_overrides, + } + assert snapshot == { + "headers": False, + "iterator": False, + "chunk": False, + "guardrail": False, + "pre_call": False, + "callbacks": (), + "overrides": (), + } + + +def test_callback_capabilities_detects_overrides(monkeypatch): + cb1 = _OverridesResponseHeaders() + cb2 = _OverridesIterator() + cb3 = _OverridesPerChunk() + cb4 = _OverridesPreCall() + monkeypatch.setattr(litellm, "callbacks", [cb1, cb2, cb3, cb4]) + + caps = ProxyLogging._callback_capabilities() + snapshot = { + "headers": caps.has_post_call_response_headers, + "iterator": caps.has_iterator_override, + "chunk": caps.has_streaming_chunk_override, + "pre_call": caps.has_pre_call_override, + } + assert snapshot == { + "headers": True, + "iterator": True, + "chunk": True, + "pre_call": True, + } + + +def test_callback_capabilities_caches_result(monkeypatch): + cb = _OverridesResponseHeaders() + monkeypatch.setattr(litellm, "callbacks", [cb]) + first = ProxyLogging._callback_capabilities() + second = ProxyLogging._callback_capabilities() + assert first is second + + +def test_callback_capabilities_invalidates_on_change(monkeypatch): + monkeypatch.setattr(litellm, "callbacks", [_OverridesResponseHeaders()]) + first = ProxyLogging._callback_capabilities() + monkeypatch.setattr(litellm, "callbacks", [_OverridesIterator()]) + second = ProxyLogging._callback_capabilities() + assert first is not second + assert first.has_post_call_response_headers is True + assert second.has_post_call_response_headers is False + assert second.has_iterator_override is True + + +def test_callback_capabilities_callback_resolution_error_raises(monkeypatch): + monkeypatch.setattr(litellm, "callbacks", ["unknown-string"]) + monkeypatch.setattr( + litellm.litellm_core_utils.litellm_logging, + "get_custom_logger_compatible_class", + lambda *a, **kw: (_ for _ in ()).throw(RuntimeError("bad")), + ) + with pytest.raises(RuntimeError): + ProxyLogging._callback_capabilities() + + +# --------------------------------------------------------------------------- +# Individual capability probes +# --------------------------------------------------------------------------- + + +def test_has_post_call_response_headers_callbacks_truth_table(monkeypatch, mock_callbacks_disabled): + """One snapshot covering true + false + cache invalidation.""" + snapshot = { + "empty_returns_false": ProxyLogging.has_post_call_response_headers_callbacks(), + } + monkeypatch.setattr(litellm, "callbacks", [_OverridesResponseHeaders()]) + ProxyLogging._callback_capabilities_cache.clear() + snapshot["override_returns_true"] = ProxyLogging.has_post_call_response_headers_callbacks() + monkeypatch.setattr(litellm, "callbacks", [_PlainLogger()]) + ProxyLogging._callback_capabilities_cache.clear() + snapshot["plain_logger_false"] = ProxyLogging.has_post_call_response_headers_callbacks() + assert snapshot == { + "empty_returns_false": False, + "override_returns_true": True, + "plain_logger_false": False, + } + + +def test_has_post_call_response_headers_callbacks_error_when_bad_callback(monkeypatch): + monkeypatch.setattr(litellm, "callbacks", ["x"]) + monkeypatch.setattr( + litellm.litellm_core_utils.litellm_logging, + "get_custom_logger_compatible_class", + lambda *a, **kw: (_ for _ in ()).throw(RuntimeError("kaboom")), + ) + with pytest.raises(RuntimeError): + ProxyLogging.has_post_call_response_headers_callbacks() + + +def test_has_streaming_callbacks_truth_table(monkeypatch, mock_callbacks_disabled): + snapshot = { + "empty_false": ProxyLogging.has_streaming_callbacks(), + } + monkeypatch.setattr(litellm, "callbacks", [_OverridesIterator()]) + ProxyLogging._callback_capabilities_cache.clear() + snapshot["iterator_override_true"] = ProxyLogging.has_streaming_callbacks() + monkeypatch.setattr(litellm, "callbacks", [_OverridesPerChunk()]) + ProxyLogging._callback_capabilities_cache.clear() + snapshot["per_chunk_override_true"] = ProxyLogging.has_streaming_callbacks() + assert snapshot == { + "empty_false": False, + "iterator_override_true": True, + "per_chunk_override_true": True, + } + + +def test_has_streaming_callbacks_error_when_resolution_fails(monkeypatch): + monkeypatch.setattr(litellm, "callbacks", ["x"]) + monkeypatch.setattr( + litellm.litellm_core_utils.litellm_logging, + "get_custom_logger_compatible_class", + lambda *a, **kw: (_ for _ in ()).throw(ValueError("nope")), + ) + with pytest.raises(ValueError): + ProxyLogging.has_streaming_callbacks() + + +def test_has_streaming_chunk_hook_overrides_truth_table(monkeypatch, mock_callbacks_disabled): + snapshot = { + "empty_false": ProxyLogging.has_streaming_chunk_hook_overrides(), + } + monkeypatch.setattr(litellm, "callbacks", [_OverridesPerChunk()]) + ProxyLogging._callback_capabilities_cache.clear() + snapshot["per_chunk_override_true"] = ProxyLogging.has_streaming_chunk_hook_overrides() + monkeypatch.setattr(litellm, "callbacks", [_OverridesIterator()]) + ProxyLogging._callback_capabilities_cache.clear() + snapshot["only_iterator_false"] = ProxyLogging.has_streaming_chunk_hook_overrides() + assert snapshot == { + "empty_false": False, + "per_chunk_override_true": True, + "only_iterator_false": False, + } + + +def test_has_streaming_chunk_hook_overrides_error_raises(monkeypatch): + monkeypatch.setattr(litellm, "callbacks", ["x"]) + monkeypatch.setattr( + litellm.litellm_core_utils.litellm_logging, + "get_custom_logger_compatible_class", + lambda *a, **kw: (_ for _ in ()).throw(TypeError("nope")), + ) + with pytest.raises(TypeError): + ProxyLogging.has_streaming_chunk_hook_overrides() + + +def test_needs_iterator_wrap_truth_table(proxy_logging, monkeypatch, mock_callbacks_disabled): + snapshot = { + "empty_false": proxy_logging.needs_iterator_wrap(), + } + monkeypatch.setattr(litellm, "callbacks", [_OverridesIterator()]) + ProxyLogging._callback_capabilities_cache.clear() + snapshot["with_iter_override_true"] = proxy_logging.needs_iterator_wrap() + monkeypatch.setattr(litellm, "callbacks", [_OverridesPerChunk()]) + ProxyLogging._callback_capabilities_cache.clear() + snapshot["only_per_chunk_false"] = proxy_logging.needs_iterator_wrap() + assert snapshot == { + "empty_false": False, + "with_iter_override_true": True, + "only_per_chunk_false": False, + } + + +def test_needs_iterator_wrap_error_raises(proxy_logging, monkeypatch): + monkeypatch.setattr(litellm, "callbacks", ["x"]) + monkeypatch.setattr( + litellm.litellm_core_utils.litellm_logging, + "get_custom_logger_compatible_class", + lambda *a, **kw: (_ for _ in ()).throw(RuntimeError("oops")), + ) + with pytest.raises(RuntimeError): + proxy_logging.needs_iterator_wrap() + + +def test_needs_per_chunk_streaming_hook_truth_table(proxy_logging, monkeypatch, mock_callbacks_disabled): + snapshot = { + "empty_false": proxy_logging.needs_per_chunk_streaming_hook(), + } + monkeypatch.setattr(litellm, "callbacks", [_OverridesPerChunk()]) + ProxyLogging._callback_capabilities_cache.clear() + snapshot["per_chunk_override_true"] = proxy_logging.needs_per_chunk_streaming_hook() + monkeypatch.setattr(litellm, "callbacks", [_OverridesIterator()]) + ProxyLogging._callback_capabilities_cache.clear() + snapshot["only_iter_override_false"] = proxy_logging.needs_per_chunk_streaming_hook() + assert snapshot == { + "empty_false": False, + "per_chunk_override_true": True, + "only_iter_override_false": False, + } + + +def test_needs_per_chunk_streaming_hook_error_raises(proxy_logging, monkeypatch): + monkeypatch.setattr(litellm, "callbacks", ["x"]) + monkeypatch.setattr( + litellm.litellm_core_utils.litellm_logging, + "get_custom_logger_compatible_class", + lambda *a, **kw: (_ for _ in ()).throw(KeyError("oops")), + ) + with pytest.raises(KeyError): + proxy_logging.needs_per_chunk_streaming_hook() + + +def test_has_during_call_guardrails_truth_table(monkeypatch, mock_callbacks_disabled): + from litellm.integrations.custom_guardrail import CustomGuardrail + + class _G(CustomGuardrail): + def __init__(self): + super().__init__(guardrail_name="g", event_hook="pre_call") + + snapshot = { + "empty_false": ProxyLogging.has_during_call_guardrails(), + } + monkeypatch.setattr(litellm, "callbacks", [_G()]) + ProxyLogging._callback_capabilities_cache.clear() + snapshot["with_guardrail_true"] = ProxyLogging.has_during_call_guardrails() + monkeypatch.setattr(litellm, "callbacks", [_PlainLogger()]) + ProxyLogging._callback_capabilities_cache.clear() + snapshot["only_plain_logger_false"] = ProxyLogging.has_during_call_guardrails() + assert snapshot == { + "empty_false": False, + "with_guardrail_true": True, + "only_plain_logger_false": False, + } + + +def test_has_during_call_guardrails_resolution_error_raises(monkeypatch): + monkeypatch.setattr(litellm, "callbacks", ["x"]) + monkeypatch.setattr( + litellm.litellm_core_utils.litellm_logging, + "get_custom_logger_compatible_class", + lambda *a, **kw: (_ for _ in ()).throw(RuntimeError("oops")), + ) + with pytest.raises(RuntimeError): + ProxyLogging.has_during_call_guardrails() + + +# --------------------------------------------------------------------------- +# get_combined_callback_list +# --------------------------------------------------------------------------- + + +def test_get_combined_callback_list_matrix(proxy_logging): + snapshot = { + "merge_dedupes_shared": sorted( + proxy_logging.get_combined_callback_list( + dynamic_success_callbacks=["dyn-1", "shared"], + global_callbacks=["glob-1", "shared"], + ) + ), + "none_dynamic_returns_global_copy": proxy_logging.get_combined_callback_list( + dynamic_success_callbacks=None, global_callbacks=["a", "b", "c"] + ), + "empty_both": proxy_logging.get_combined_callback_list( + dynamic_success_callbacks=[], global_callbacks=[] + ), + } + assert snapshot == { + "merge_dedupes_shared": ["dyn-1", "glob-1", "shared"], + "none_dynamic_returns_global_copy": ["a", "b", "c"], + "empty_both": [], + } + + +def test_get_combined_callback_list_unhashable_dynamic_raises(proxy_logging): + with pytest.raises(TypeError): + proxy_logging.get_combined_callback_list( + dynamic_success_callbacks=[{"unhashable": True}], + global_callbacks=[], + ) diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_callback_capabilities_dataclass.py b/tests/test_litellm/proxy/utils/proxy_logging/test_callback_capabilities_dataclass.py new file mode 100644 index 00000000000..931c832732e --- /dev/null +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_callback_capabilities_dataclass.py @@ -0,0 +1,59 @@ +"""Pin the ``_CallbackCapabilities`` dataclass shape and defaults.""" + +from __future__ import annotations + +import dataclasses + +import pytest + +from litellm.proxy.utils import _CallbackCapabilities + + +def test_callback_capabilities_default_values(): + caps = _CallbackCapabilities() + snapshot = { + "has_post_call_response_headers": caps.has_post_call_response_headers, + "has_iterator_override": caps.has_iterator_override, + "has_streaming_chunk_override": caps.has_streaming_chunk_override, + "has_guardrail": caps.has_guardrail, + "has_pre_call_override": caps.has_pre_call_override, + "iterator_overrides": caps.iterator_overrides, + "resolved_callbacks": caps.resolved_callbacks, + } + assert snapshot == { + "has_post_call_response_headers": False, + "has_iterator_override": False, + "has_streaming_chunk_override": False, + "has_guardrail": False, + "has_pre_call_override": False, + "iterator_overrides": (), + "resolved_callbacks": (), + } + + +def test_callback_capabilities_explicit_values_preserved(): + cb1 = object() + cb2 = object() + caps = _CallbackCapabilities( + has_post_call_response_headers=True, + has_iterator_override=True, + has_streaming_chunk_override=False, + has_guardrail=True, + has_pre_call_override=False, + iterator_overrides=((cb1, "override"), (cb2, "apply_guardrail")), + resolved_callbacks=(cb1, cb2), + ) + assert caps.has_post_call_response_headers is True + assert caps.iterator_overrides == ((cb1, "override"), (cb2, "apply_guardrail")) + assert caps.resolved_callbacks == (cb1, cb2) + + +def test_callback_capabilities_is_frozen_error_on_mutation_raises(): + caps = _CallbackCapabilities() + with pytest.raises(dataclasses.FrozenInstanceError): + caps.has_post_call_response_headers = True # type: ignore[misc] + + +def test_callback_capabilities_invalid_field_error_raises(): + with pytest.raises(TypeError): + _CallbackCapabilities(unknown_field=True) # type: ignore[call-arg] diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_during_call_hook.py b/tests/test_litellm/proxy/utils/proxy_logging/test_during_call_hook.py new file mode 100644 index 00000000000..3c5d879c2dc --- /dev/null +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_during_call_hook.py @@ -0,0 +1,86 @@ +"""Pin ``ProxyLogging.during_call_hook``.""" + +from __future__ import annotations + +from typing import Any, Dict +from unittest.mock import AsyncMock, MagicMock + +import pytest + +import litellm +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.proxy.utils import ProxyLogging +from litellm.types.guardrails import GuardrailEventHooks + + +@pytest.fixture(autouse=True) +def _clear_caps_cache(): + ProxyLogging._callback_capabilities_cache.clear() + yield + ProxyLogging._callback_capabilities_cache.clear() + + +def _make_guardrail(name="g1", should_run=True, response=None): + cb = MagicMock(spec=CustomGuardrail) + cb.__class__ = CustomGuardrail + cb.guardrail_name = name + cb.event_hook = GuardrailEventHooks.during_call + cb.use_native_during_call_hook = False + cb.should_run_guardrail = MagicMock(return_value=should_run) + cb.async_moderation_hook = AsyncMock(return_value=response) + return cb + + +@pytest.mark.asyncio +async def test_during_call_hook_no_guardrail_fast_path_returns_data(proxy_logging, make_user_api_key_auth, mock_callbacks_disabled): + data = {"messages": [{"role": "user"}], "model": "m", "temperature": 0.1} + out = await proxy_logging.during_call_hook( + data=data, + user_api_key_dict=make_user_api_key_auth(), + call_type="completion", + ) + assert out is data + + +@pytest.mark.asyncio +async def test_during_call_hook_runs_guardrails_in_parallel(proxy_logging, make_user_api_key_auth, monkeypatch): + g1 = _make_guardrail("a") + g2 = _make_guardrail("b") + monkeypatch.setattr(litellm, "callbacks", [g1, g2]) + data = {"messages": [{"role": "user"}], "model": "m", "temperature": 0.1} + out = await proxy_logging.during_call_hook( + data=data, + user_api_key_dict=make_user_api_key_auth(), + call_type="completion", + ) + snapshot = { + "out_is_data": out is data, + "a_called": g1.async_moderation_hook.called, + "b_called": g2.async_moderation_hook.called, + } + assert snapshot == {"out_is_data": True, "a_called": True, "b_called": True} + + +@pytest.mark.asyncio +async def test_during_call_hook_guardrail_skipped_when_should_not_run(proxy_logging, make_user_api_key_auth, monkeypatch): + g = _make_guardrail("g", should_run=False) + monkeypatch.setattr(litellm, "callbacks", [g]) + await proxy_logging.during_call_hook( + data={"model": "m"}, + user_api_key_dict=make_user_api_key_auth(), + call_type="completion", + ) + g.async_moderation_hook.assert_not_called() + + +@pytest.mark.asyncio +async def test_during_call_hook_guardrail_error_raises(proxy_logging, make_user_api_key_auth, monkeypatch): + g = _make_guardrail("bad") + g.async_moderation_hook = AsyncMock(side_effect=RuntimeError("blocked")) + monkeypatch.setattr(litellm, "callbacks", [g]) + with pytest.raises(RuntimeError): + await proxy_logging.during_call_hook( + data={"model": "m"}, + user_api_key_dict=make_user_api_key_auth(), + call_type="completion", + ) diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_guardrail_pipeline.py b/tests/test_litellm/proxy/utils/proxy_logging/test_guardrail_pipeline.py new file mode 100644 index 00000000000..1ff9fbf8d83 --- /dev/null +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_guardrail_pipeline.py @@ -0,0 +1,559 @@ +"""Pin ProxyLogging guardrail pipeline helpers. + +Covers ``_should_use_guardrail_load_balancing``, ``_execute_guardrail_hook``, +``_execute_guardrail_with_load_balancing``, ``_process_guardrail_callback``, +``_process_prompt_template``, ``_process_guardrail_metadata``, +``_maybe_execute_pipelines``, ``_handle_pipeline_result``, +``_run_guardrail_task_with_enrichment``. +""" + +from __future__ import annotations + +import asyncio +from typing import Any, Dict, List +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from fastapi import HTTPException + +import litellm +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + ModifyResponseException, +) +from litellm.proxy.utils import ProxyLogging +from litellm.types.guardrails import GuardrailEventHooks + + +@pytest.fixture(autouse=True) +def _clear_caps_cache(): + ProxyLogging._callback_capabilities_cache.clear() + yield + ProxyLogging._callback_capabilities_cache.clear() + + +# --------------------------------------------------------------------------- +# _should_use_guardrail_load_balancing +# --------------------------------------------------------------------------- + + +def test_should_use_guardrail_load_balancing_truth_table(proxy_logging): + snapshot = {} + router = MagicMock() + router.guardrail_list = [{"guardrail_name": "g1"}, {"guardrail_name": "g1"}] + with patch("litellm.proxy.proxy_server.llm_router", router): + snapshot["multiple_deployments"] = proxy_logging._should_use_guardrail_load_balancing("g1") + router.guardrail_list = [{"guardrail_name": "g1"}] + with patch("litellm.proxy.proxy_server.llm_router", router): + snapshot["single_deployment"] = proxy_logging._should_use_guardrail_load_balancing("g1") + with patch("litellm.proxy.proxy_server.llm_router", None): + snapshot["no_router"] = proxy_logging._should_use_guardrail_load_balancing("g1") + router.guardrail_list = [{"guardrail_name": "other"}, {"guardrail_name": "other"}] + with patch("litellm.proxy.proxy_server.llm_router", router): + snapshot["unmatched_name"] = proxy_logging._should_use_guardrail_load_balancing("g1") + assert snapshot == { + "multiple_deployments": True, + "single_deployment": False, + "no_router": False, + "unmatched_name": False, + } + + +def test_should_use_guardrail_load_balancing_error_on_bad_guardrail_list(proxy_logging): + router = MagicMock() + router.guardrail_list = "not a list" + with patch("litellm.proxy.proxy_server.llm_router", router): + with pytest.raises((TypeError, AttributeError)): + proxy_logging._should_use_guardrail_load_balancing("g1") + + +# --------------------------------------------------------------------------- +# _execute_guardrail_hook +# --------------------------------------------------------------------------- + + +def _make_guardrail(): + cb = MagicMock(spec=CustomGuardrail) + cb.__class__ = CustomGuardrail + cb.guardrail_name = "g" + cb.event_hook = GuardrailEventHooks.pre_call + cb.use_native_during_call_hook = False + cb.async_pre_call_hook = AsyncMock(return_value={"a": 1, "b": 2, "c": 3}) + cb.async_moderation_hook = AsyncMock(return_value={"x": 1, "y": 2, "z": 3}) + cb.async_post_call_success_hook = AsyncMock(return_value={"p": 1, "q": 2, "r": 3}) + return cb + + +@pytest.mark.asyncio +async def test_execute_guardrail_hook_pre_call(proxy_logging, make_user_api_key_auth): + cb = _make_guardrail() + out = await proxy_logging._execute_guardrail_hook( + callback=cb, + hook_type="pre_call", + data={"model": "m"}, + user_api_key_dict=make_user_api_key_auth(), + call_type="completion", + ) + assert out == {"a": 1, "b": 2, "c": 3} + + +@pytest.mark.asyncio +async def test_execute_guardrail_hook_during_call(proxy_logging, make_user_api_key_auth): + cb = _make_guardrail() + out = await proxy_logging._execute_guardrail_hook( + callback=cb, + hook_type="during_call", + data={"model": "m"}, + user_api_key_dict=make_user_api_key_auth(), + call_type="completion", + ) + assert out == {"x": 1, "y": 2, "z": 3} + + +@pytest.mark.asyncio +async def test_execute_guardrail_hook_post_call(proxy_logging, make_user_api_key_auth): + cb = _make_guardrail() + out = await proxy_logging._execute_guardrail_hook( + callback=cb, + hook_type="post_call", + data={"model": "m"}, + user_api_key_dict=make_user_api_key_auth(), + call_type="completion", + response={"original": True}, + ) + assert out == {"p": 1, "q": 2, "r": 3} + + +@pytest.mark.asyncio +async def test_execute_guardrail_hook_unknown_hook_type_raises(proxy_logging, make_user_api_key_auth): + cb = _make_guardrail() + with pytest.raises(ValueError, match="Unknown hook_type"): + await proxy_logging._execute_guardrail_hook( + callback=cb, + hook_type="weird", # type: ignore[arg-type] + data={}, + user_api_key_dict=make_user_api_key_auth(), + call_type="completion", + ) + + +# --------------------------------------------------------------------------- +# _execute_guardrail_with_load_balancing +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_execute_guardrail_with_load_balancing_routes_through_router( + proxy_logging, make_user_api_key_auth +): + cb = _make_guardrail() + router = MagicMock() + router.get_available_guardrail = MagicMock(return_value={"callback": cb}) + with patch("litellm.proxy.proxy_server.llm_router", router): + out = await proxy_logging._execute_guardrail_with_load_balancing( + guardrail_name="g", + hook_type="pre_call", + data={"model": "m"}, + user_api_key_dict=make_user_api_key_auth(), + call_type="completion", + ) + assert out == {"a": 1, "b": 2, "c": 3} + + +@pytest.mark.asyncio +async def test_execute_guardrail_with_load_balancing_router_none_raises( + proxy_logging, make_user_api_key_auth +): + with patch("litellm.proxy.proxy_server.llm_router", None): + with pytest.raises(ValueError, match="Router not initialized"): + await proxy_logging._execute_guardrail_with_load_balancing( + guardrail_name="g", + hook_type="pre_call", + data={}, + user_api_key_dict=make_user_api_key_auth(), + call_type="completion", + ) + + +@pytest.mark.asyncio +async def test_execute_guardrail_with_load_balancing_no_callback_raises( + proxy_logging, make_user_api_key_auth +): + router = MagicMock() + router.get_available_guardrail = MagicMock(return_value={"callback": None}) + with patch("litellm.proxy.proxy_server.llm_router", router): + with pytest.raises(ValueError, match="No callback found"): + await proxy_logging._execute_guardrail_with_load_balancing( + guardrail_name="g", + hook_type="pre_call", + data={}, + user_api_key_dict=make_user_api_key_auth(), + call_type="completion", + ) + + +# --------------------------------------------------------------------------- +# _process_guardrail_callback +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_process_guardrail_callback_skipped_when_should_run_false( + proxy_logging, make_user_api_key_auth +): + cb = _make_guardrail() + cb.should_run_guardrail = MagicMock(return_value=False) + out = await proxy_logging._process_guardrail_callback( + callback=cb, + data={"model": "m"}, + user_api_key_dict=make_user_api_key_auth(), + call_type="completion", + event_type=GuardrailEventHooks.pre_call, + ) + assert out is None + + +@pytest.mark.asyncio +async def test_process_guardrail_callback_returns_data_on_success( + proxy_logging, make_user_api_key_auth, monkeypatch +): + cb = _make_guardrail() + cb.should_run_guardrail = MagicMock(return_value=True) + proxy_logging._should_use_guardrail_load_balancing = MagicMock(return_value=False) + out = await proxy_logging._process_guardrail_callback( + callback=cb, + data={"model": "m", "messages": [{"role": "user"}], "temperature": 0.1}, + user_api_key_dict=make_user_api_key_auth(), + call_type="completion", + event_type=GuardrailEventHooks.pre_call, + ) + assert out == {"a": 1, "b": 2, "c": 3} + + +@pytest.mark.asyncio +async def test_process_guardrail_callback_enriches_and_reraises_http_exception( + proxy_logging, make_user_api_key_auth, monkeypatch +): + cb = _make_guardrail() + cb.should_run_guardrail = MagicMock(return_value=True) + detail = {"error": "blocked"} + cb.async_pre_call_hook = AsyncMock(side_effect=HTTPException(status_code=400, detail=detail)) + cb.event_hook = "pre_call" + proxy_logging._should_use_guardrail_load_balancing = MagicMock(return_value=False) + + with pytest.raises(HTTPException): + await proxy_logging._process_guardrail_callback( + callback=cb, + data={"model": "m"}, + user_api_key_dict=make_user_api_key_auth(), + call_type="completion", + event_type=GuardrailEventHooks.pre_call, + ) + assert detail["guardrail_name"] == "g" + + +# --------------------------------------------------------------------------- +# _process_guardrail_metadata +# --------------------------------------------------------------------------- + + +def test_process_guardrail_metadata_calls_header_helper(proxy_logging, monkeypatch): + calls: List[Dict[str, Any]] = [] + + def fake_add(request_data, guardrail_name): + calls.append({"data": request_data, "name": guardrail_name}) + + from litellm.proxy.common_utils import callback_utils + + monkeypatch.setattr(callback_utils, "add_guardrail_to_applied_guardrails_header", fake_add) + data = {"metadata": {"guardrails": ["g1", "g2"]}} + proxy_logging._process_guardrail_metadata(data) + snapshot = { + "call_count": len(calls), + "first_name": calls[0]["name"], + "second_name": calls[1]["name"], + "data_passed_is_input": all(c["data"] is data for c in calls), + } + assert snapshot == { + "call_count": 2, + "first_name": "g1", + "second_name": "g2", + "data_passed_is_input": True, + } + + +def test_process_guardrail_metadata_skips_already_applied(proxy_logging, monkeypatch): + calls: List[str] = [] + + def fake_add(request_data, guardrail_name): + calls.append(guardrail_name) + + from litellm.proxy.common_utils import callback_utils + + monkeypatch.setattr(callback_utils, "add_guardrail_to_applied_guardrails_header", fake_add) + data = {"metadata": {"guardrails": ["g1", "g2"], "applied_guardrails": ["g1"]}} + proxy_logging._process_guardrail_metadata(data) + assert calls == ["g2"] + + +def test_process_guardrail_metadata_no_metadata_is_noop(proxy_logging, monkeypatch): + from litellm.proxy.common_utils import callback_utils + + monkeypatch.setattr( + callback_utils, + "add_guardrail_to_applied_guardrails_header", + MagicMock(side_effect=AssertionError("should not be called")), + ) + proxy_logging._process_guardrail_metadata({}) + + +def test_process_guardrail_metadata_invalid_data_raises(proxy_logging): + with pytest.raises(AttributeError): + proxy_logging._process_guardrail_metadata(None) # type: ignore[arg-type] + + +# --------------------------------------------------------------------------- +# _maybe_execute_pipelines +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_maybe_execute_pipelines_no_pipelines_returns_data(proxy_logging, make_user_api_key_auth): + data = {"messages": [{"role": "user"}], "model": "m", "temperature": 0.1} + out = await proxy_logging._maybe_execute_pipelines( + data=data, + user_api_key_dict=make_user_api_key_auth(), + call_type="completion", + event_hook="pre_call", + ) + assert out == {"messages": [{"role": "user"}], "model": "m", "temperature": 0.1} + + +@pytest.mark.asyncio +async def test_maybe_execute_pipelines_skips_pipelines_with_other_mode(proxy_logging, make_user_api_key_auth, monkeypatch): + pipeline = MagicMock() + pipeline.mode = "post_call" # not pre_call + data = {"metadata": {"_guardrail_pipelines": [("p1", pipeline)]}, "model": "m", "messages": []} + executed = MagicMock() + monkeypatch.setattr( + "litellm.proxy.policy_engine.pipeline_executor.PipelineExecutor.execute_steps", executed + ) + out = await proxy_logging._maybe_execute_pipelines( + data=data, + user_api_key_dict=make_user_api_key_auth(), + call_type="completion", + event_hook="pre_call", + ) + executed.assert_not_called() + assert out is data + + +@pytest.mark.asyncio +async def test_maybe_execute_pipelines_blocks_on_block_terminal_action_raises( + proxy_logging, make_user_api_key_auth, monkeypatch +): + pipeline = MagicMock() + pipeline.mode = "pre_call" + pipeline.steps = [] + fake_result = MagicMock() + fake_result.terminal_action = "block" + fake_result.step_results = [] + data = {"metadata": {"_guardrail_pipelines": [("policy-1", pipeline)]}, "messages": [], "model": "m"} + + async def fake_execute_steps(**kwargs): + return fake_result + + monkeypatch.setattr( + "litellm.proxy.policy_engine.pipeline_executor.PipelineExecutor.execute_steps", + fake_execute_steps, + ) + with pytest.raises(HTTPException): + await proxy_logging._maybe_execute_pipelines( + data=data, + user_api_key_dict=make_user_api_key_auth(), + call_type="completion", + event_hook="pre_call", + ) + + +# --------------------------------------------------------------------------- +# _handle_pipeline_result +# --------------------------------------------------------------------------- + + +def test_handle_pipeline_result_allow_with_modifications(): + data = {"a": 1} + result = MagicMock() + result.terminal_action = "allow" + result.modified_data = {"b": 2, "c": 3} + out = ProxyLogging._handle_pipeline_result(result=result, data=data, policy_name="p") + assert out == {"a": 1, "b": 2, "c": 3} + + +def test_handle_pipeline_result_block_raises_http_exception(): + result = MagicMock() + result.terminal_action = "block" + result.step_results = [] + with pytest.raises(HTTPException) as info: + ProxyLogging._handle_pipeline_result(result=result, data={"model": "m"}, policy_name="p") + detail = info.value.detail + snapshot = { + "is_dict": isinstance(detail, dict), + "error_type": detail["error"]["type"], + "policy": detail["error"]["pipeline_context"]["policy"], + } + assert snapshot == { + "is_dict": True, + "error_type": "guardrail_pipeline_error", + "policy": "p", + } + + +def test_handle_pipeline_result_modify_response_raises_modify_exception(): + result = MagicMock() + result.terminal_action = "modify_response" + result.modify_response_message = "filtered" + with pytest.raises(ModifyResponseException): + ProxyLogging._handle_pipeline_result(result=result, data={"model": "m"}, policy_name="p") + + +def test_handle_pipeline_result_unknown_action_returns_data(): + data = {"a": 1, "b": 2, "c": 3} + result = MagicMock() + result.terminal_action = "something_else" + assert ProxyLogging._handle_pipeline_result(result=result, data=data, policy_name="p") is data + + +# --------------------------------------------------------------------------- +# _run_guardrail_task_with_enrichment +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_run_guardrail_task_with_enrichment_passes_result(): + async def task(): + return {"a": 1, "b": 2, "c": 3} + + out = await ProxyLogging._run_guardrail_task_with_enrichment( + callback=MagicMock(guardrail_name="g"), coro=task() + ) + assert out == {"a": 1, "b": 2, "c": 3} + + +@pytest.mark.asyncio +async def test_run_guardrail_task_with_enrichment_enriches_http_exception_raises(): + detail = {"error": "blocked"} + + async def task(): + raise HTTPException(status_code=400, detail=detail) + + cb = MagicMock() + cb.guardrail_name = "presidio" + cb.event_hook = "pre_call" + with pytest.raises(HTTPException): + await ProxyLogging._run_guardrail_task_with_enrichment(callback=cb, coro=task()) + assert detail["guardrail_name"] == "presidio" + + +# --------------------------------------------------------------------------- +# _process_prompt_template +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_process_prompt_template_no_op_when_no_prompt_spec(proxy_logging, monkeypatch): + from litellm.proxy.prompts import prompt_registry + + monkeypatch.setattr( + prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "get_prompt_callback_by_id", lambda *a, **kw: None + ) + monkeypatch.setattr( + prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "get_prompt_by_id", lambda *a, **kw: None + ) + data: Dict[str, Any] = {"messages": [{"role": "user"}], "model": "m", "temperature": 0.1} + await proxy_logging._process_prompt_template( + data=data, + litellm_logging_obj=MagicMock(), + prompt_id="some-id", + prompt_version=1, + call_type="completion", + ) + assert data == {"messages": [{"role": "user"}], "model": "m", "temperature": 0.1} + + +@pytest.mark.asyncio +async def test_process_prompt_template_applies_when_spec_resolves(proxy_logging, monkeypatch): + from litellm.proxy.prompts import prompt_registry + + custom_logger = MagicMock() + prompt_spec = MagicMock() + prompt_spec.litellm_params = MagicMock(prompt_id="resolved-id") + + monkeypatch.setattr( + prompt_registry.IN_MEMORY_PROMPT_REGISTRY, + "get_prompt_callback_by_id", + lambda *a, **kw: custom_logger, + ) + monkeypatch.setattr( + prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "get_prompt_by_id", lambda *a, **kw: prompt_spec + ) + + logging_obj = MagicMock() + logging_obj.async_get_chat_completion_prompt = AsyncMock( + return_value=( + "model-out", + [{"role": "user", "content": "rendered"}], + {"temperature": 0.5, "top_p": 1}, + ) + ) + data: Dict[str, Any] = { + "messages": [{"role": "user", "content": "orig"}], + "model": "m", + "prompt_id": "x", + } + await proxy_logging._process_prompt_template( + data=data, + litellm_logging_obj=logging_obj, + prompt_id="x", + prompt_version=None, + call_type="completion", + ) + snapshot = { + "model": data["model"], + "messages": data["messages"], + "temperature": data["temperature"], + "top_p": data["top_p"], + } + assert snapshot == { + "model": "model-out", + "messages": [{"role": "user", "content": "rendered"}], + "temperature": 0.5, + "top_p": 1, + } + + +@pytest.mark.asyncio +async def test_process_prompt_template_async_get_prompt_error_raises(proxy_logging, monkeypatch): + from litellm.proxy.prompts import prompt_registry + + custom_logger = MagicMock() + prompt_spec = MagicMock() + prompt_spec.litellm_params = MagicMock(prompt_id="x") + monkeypatch.setattr( + prompt_registry.IN_MEMORY_PROMPT_REGISTRY, + "get_prompt_callback_by_id", + lambda *a, **kw: custom_logger, + ) + monkeypatch.setattr( + prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "get_prompt_by_id", lambda *a, **kw: prompt_spec + ) + logging_obj = MagicMock() + logging_obj.async_get_chat_completion_prompt = AsyncMock(side_effect=RuntimeError("bad prompt")) + with pytest.raises(RuntimeError): + await proxy_logging._process_prompt_template( + data={"messages": [], "model": "m", "prompt_id": "x"}, + litellm_logging_obj=logging_obj, + prompt_id="x", + prompt_version=None, + call_type="completion", + ) diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_internal_usage_cache.py b/tests/test_litellm/proxy/utils/proxy_logging/test_internal_usage_cache.py new file mode 100644 index 00000000000..ff0afa45e36 --- /dev/null +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_internal_usage_cache.py @@ -0,0 +1,186 @@ +"""Pin behavior of ``InternalUsageCache``: a thin adapter over ``DualCache``. + +Each method should pass-through to the underlying ``DualCache`` with +exactly the same arguments, mapping ``litellm_parent_otel_span`` to the +``DualCache`` kw it expects. +""" + +from __future__ import annotations + +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.caching.caching import DualCache +from litellm.proxy.utils import InternalUsageCache + + +def _kwargs_snapshot(call): + return dict(call.kwargs) + + +def test_internal_usage_cache_init_stores_dual_cache(): + inner = DualCache(default_in_memory_ttl=1) + cache = InternalUsageCache(dual_cache=inner) + snapshot = { + "is_internal_usage_cache": isinstance(cache, InternalUsageCache), + "dual_cache_is_inner": cache.dual_cache is inner, + "ttl_is_one": inner.default_in_memory_ttl == 1, + } + assert snapshot == { + "is_internal_usage_cache": True, + "dual_cache_is_inner": True, + "ttl_is_one": True, + } + + +def test_internal_usage_cache_init_error_requires_dual_cache(): + with pytest.raises(TypeError): + InternalUsageCache() # type: ignore[call-arg] + + +@pytest.mark.asyncio +async def test_async_get_cache_forwards_args(): + inner = MagicMock() + inner.async_get_cache = AsyncMock(return_value={"hit": True, "value": 42, "source": "redis"}) + cache = InternalUsageCache(dual_cache=inner) + + result = await cache.async_get_cache(key="k", litellm_parent_otel_span="span", local_only=True, extra="x") + forwarded = _kwargs_snapshot(inner.async_get_cache.call_args) + assert forwarded == {"key": "k", "local_only": True, "parent_otel_span": "span", "extra": "x"} + assert result == {"hit": True, "value": 42, "source": "redis"} + + +@pytest.mark.asyncio +async def test_async_get_cache_propagates_underlying_error_raises(): + inner = MagicMock() + inner.async_get_cache = AsyncMock(side_effect=RuntimeError("redis down")) + cache = InternalUsageCache(dual_cache=inner) + with pytest.raises(RuntimeError, match="redis down"): + await cache.async_get_cache(key="k", litellm_parent_otel_span=None) + + +@pytest.mark.asyncio +async def test_async_set_cache_forwards_args(): + inner = MagicMock() + inner.async_set_cache = AsyncMock() + cache = InternalUsageCache(dual_cache=inner) + + await cache.async_set_cache(key="k", value="v", litellm_parent_otel_span="span", local_only=False, ttl=60) + forwarded = _kwargs_snapshot(inner.async_set_cache.call_args) + assert forwarded == { + "key": "k", + "value": "v", + "local_only": False, + "litellm_parent_otel_span": "span", + "ttl": 60, + } + + +@pytest.mark.asyncio +async def test_async_set_cache_propagates_error_raises(): + inner = MagicMock() + inner.async_set_cache = AsyncMock(side_effect=ValueError("bad value")) + cache = InternalUsageCache(dual_cache=inner) + with pytest.raises(ValueError, match="bad value"): + await cache.async_set_cache(key="k", value="v", litellm_parent_otel_span=None) + + +@pytest.mark.asyncio +async def test_async_batch_set_cache_forwards_pipeline(): + inner = MagicMock() + inner.async_set_cache_pipeline = AsyncMock() + cache = InternalUsageCache(dual_cache=inner) + + pairs = [("a", 1), ("b", 2)] + await cache.async_batch_set_cache(cache_list=pairs, litellm_parent_otel_span=None, local_only=True, ttl=10) + forwarded = _kwargs_snapshot(inner.async_set_cache_pipeline.call_args) + assert forwarded == { + "cache_list": pairs, + "local_only": True, + "litellm_parent_otel_span": None, + "ttl": 10, + } + + +@pytest.mark.asyncio +async def test_async_batch_set_cache_propagates_error_raises(): + inner = MagicMock() + inner.async_set_cache_pipeline = AsyncMock(side_effect=ConnectionError("network")) + cache = InternalUsageCache(dual_cache=inner) + with pytest.raises(ConnectionError): + await cache.async_batch_set_cache(cache_list=[], litellm_parent_otel_span=None) + + +@pytest.mark.asyncio +async def test_async_batch_get_cache_forwards_args(): + inner = MagicMock() + inner.async_batch_get_cache = AsyncMock(return_value=[1, 2, 3]) + cache = InternalUsageCache(dual_cache=inner) + result = await cache.async_batch_get_cache(keys=["a", "b", "c"], parent_otel_span="span", local_only=False) + forwarded = _kwargs_snapshot(inner.async_batch_get_cache.call_args) + assert forwarded == {"keys": ["a", "b", "c"], "parent_otel_span": "span", "local_only": False} + assert result == [1, 2, 3] + + +@pytest.mark.asyncio +async def test_async_batch_get_cache_invalid_input_raises(): + inner = MagicMock() + inner.async_batch_get_cache = AsyncMock(side_effect=TypeError("not a list")) + cache = InternalUsageCache(dual_cache=inner) + with pytest.raises(TypeError): + await cache.async_batch_get_cache(keys=None) # type: ignore[arg-type] + + +@pytest.mark.asyncio +async def test_async_increment_cache_forwards_args(): + inner = MagicMock() + inner.async_increment_cache = AsyncMock(return_value=5.0) + cache = InternalUsageCache(dual_cache=inner) + result = await cache.async_increment_cache(key="counter", value=1.5, litellm_parent_otel_span="span") + forwarded = _kwargs_snapshot(inner.async_increment_cache.call_args) + assert forwarded == {"key": "counter", "value": 1.5, "local_only": False, "parent_otel_span": "span"} + assert result == 5.0 + + +@pytest.mark.asyncio +async def test_async_increment_cache_propagates_error_raises(): + inner = MagicMock() + inner.async_increment_cache = AsyncMock(side_effect=OverflowError()) + cache = InternalUsageCache(dual_cache=inner) + with pytest.raises(OverflowError): + await cache.async_increment_cache(key="x", value=1.0, litellm_parent_otel_span=None) + + +def test_set_cache_forwards_args(): + inner = MagicMock() + cache = InternalUsageCache(dual_cache=inner) + cache.set_cache(key="k", value="v", local_only=True, ttl=30) + forwarded = _kwargs_snapshot(inner.set_cache.call_args) + assert forwarded == {"key": "k", "value": "v", "local_only": True, "ttl": 30} + + +def test_set_cache_propagates_error_raises(): + inner = MagicMock() + inner.set_cache = MagicMock(side_effect=RuntimeError("no redis")) + cache = InternalUsageCache(dual_cache=inner) + with pytest.raises(RuntimeError): + cache.set_cache(key="k", value="v") + + +def test_get_cache_forwards_args_and_returns_inner_result(): + inner = MagicMock() + inner.get_cache = MagicMock(return_value={"k": "v", "ttl": 60, "source": "mem"}) + cache = InternalUsageCache(dual_cache=inner) + result = cache.get_cache(key="k", local_only=False) + forwarded = _kwargs_snapshot(inner.get_cache.call_args) + assert forwarded == {"key": "k", "local_only": False} + assert result == {"k": "v", "ttl": 60, "source": "mem"} + + +def test_get_cache_propagates_error_raises(): + inner = MagicMock() + inner.get_cache = MagicMock(side_effect=KeyError("missing")) + cache = InternalUsageCache(dual_cache=inner) + with pytest.raises(KeyError): + cache.get_cache(key="missing") diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_lifecycle.py b/tests/test_litellm/proxy/utils/proxy_logging/test_lifecycle.py new file mode 100644 index 00000000000..e33da672599 --- /dev/null +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_lifecycle.py @@ -0,0 +1,403 @@ +"""Pin ProxyLogging lifecycle: ``__init__``, ``startup_event``, +``update_values``, ``_add_proxy_hooks``, ``get_proxy_hook``, and +``_init_litellm_callbacks``. + +Also covers ``update_request_status`` and ``_convert_user_api_key_auth_to_dict`` +because they are direct dependents on the lifecycle state. +""" + +from __future__ import annotations + +import asyncio +from typing import Any, Dict, List +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +import litellm +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache +from litellm.proxy.utils import ( + InternalUsageCache, + ProxyLogging, +) + + +# --------------------------------------------------------------------------- +# __init__ +# --------------------------------------------------------------------------- + + +def test_proxy_logging_init_sets_default_state(mock_callbacks_disabled): + cache = UserApiKeyCache() + pl = ProxyLogging(user_api_key_cache=cache) + snapshot = { + "internal_usage_cache_type": type(pl.internal_usage_cache).__name__, + "alerting_is_none": pl.alerting is None, + "alerting_threshold": pl.alerting_threshold, + "premium_user": pl.premium_user, + "proxy_hook_mapping": pl.proxy_hook_mapping, + "daily_report_started": pl.daily_report_started, + "hanging_requests_check_started": pl.hanging_requests_check_started, + } + assert snapshot == { + "internal_usage_cache_type": "InternalUsageCache", + "alerting_is_none": True, + "alerting_threshold": 300, + "premium_user": False, + "proxy_hook_mapping": {}, + "daily_report_started": False, + "hanging_requests_check_started": False, + } + + +def test_proxy_logging_init_premium_user_flag(mock_callbacks_disabled): + pl = ProxyLogging(user_api_key_cache=UserApiKeyCache(), premium_user=True) + assert pl.premium_user is True + + +def test_proxy_logging_init_missing_cache_raises(): + with pytest.raises(TypeError): + ProxyLogging() # type: ignore[call-arg] + + +# --------------------------------------------------------------------------- +# update_values +# --------------------------------------------------------------------------- + + +def test_update_values_stores_alerting_state(proxy_logging): + proxy_logging.slack_alerting_instance = MagicMock() + proxy_logging.update_values( + alerting=["slack"], + alerting_threshold=42.0, + alert_types=["llm_too_slow"], + alert_to_webhook_url={"key": "value"}, + ) + snapshot = { + "alerting": proxy_logging.alerting, + "threshold": proxy_logging.alerting_threshold, + "alert_types": proxy_logging.alert_types, + "webhook_url": proxy_logging.alert_to_webhook_url, + } + assert snapshot == { + "alerting": ["slack"], + "threshold": 42.0, + "alert_types": ["llm_too_slow"], + "webhook_url": {"key": "value"}, + } + + +def test_update_values_with_only_redis_cache_does_not_touch_slack(proxy_logging): + proxy_logging.slack_alerting_instance = MagicMock() + redis = MagicMock() + proxy_logging.update_values(redis_cache=redis) + proxy_logging.slack_alerting_instance.update_values.assert_not_called() + assert proxy_logging.internal_usage_cache.dual_cache.redis_cache is redis + + +def test_update_values_with_no_args_is_noop(proxy_logging): + proxy_logging.slack_alerting_instance = MagicMock() + proxy_logging.update_values() + proxy_logging.slack_alerting_instance.update_values.assert_not_called() + + +def test_update_values_invalid_type_for_alerting_raises(proxy_logging): + proxy_logging.slack_alerting_instance = MagicMock( + update_values=MagicMock(side_effect=TypeError("bad type")) + ) + with pytest.raises(TypeError): + proxy_logging.update_values(alerting={"not": "a list"}) # type: ignore[arg-type] + + +# --------------------------------------------------------------------------- +# startup_event +# --------------------------------------------------------------------------- + + +def test_startup_event_initializes_slack_and_callbacks(proxy_logging): + proxy_logging.slack_alerting_instance = MagicMock() + proxy_logging.slack_alerting_instance.alert_types = [] + proxy_logging._init_litellm_callbacks = MagicMock() + proxy_logging.update_values = MagicMock() + + proxy_logging.startup_event(llm_router=None, redis_usage_cache=None) + snapshot = { + "update_called": proxy_logging.update_values.called, + "init_called": proxy_logging._init_litellm_callbacks.called, + "slack_update_called": proxy_logging.slack_alerting_instance.update_values.called, + } + assert snapshot == { + "update_called": True, + "init_called": True, + "slack_update_called": True, + } + + +def test_startup_event_propagates_init_callbacks_failure_raises(proxy_logging): + proxy_logging.slack_alerting_instance = MagicMock() + proxy_logging.slack_alerting_instance.alert_types = [] + proxy_logging._init_litellm_callbacks = MagicMock(side_effect=RuntimeError("boom")) + + with pytest.raises(RuntimeError, match="boom"): + proxy_logging.startup_event(llm_router=None, redis_usage_cache=None) + + +# --------------------------------------------------------------------------- +# _add_proxy_hooks +# --------------------------------------------------------------------------- + + +def test_add_proxy_hooks_registers_callbacks(proxy_logging, monkeypatch): + """Patch ``PROXY_HOOKS`` and the resolver so we control exactly + what gets registered. Verifies that the resulting instances land in + ``proxy_logging.proxy_hook_mapping`` keyed by hook name. + """ + hook_keys = ["cache_control_check", "max_budget_limiter"] + registered: List[Any] = [] + + from litellm.proxy import utils as utils_mod + + def fake_get_proxy_hook(hook_name): + class _Stub: + __name__ = hook_name + + def __init__(self, **kwargs): + self.hook_name = hook_name + + return _Stub + + monkeypatch.setattr(utils_mod, "PROXY_HOOKS", hook_keys) + monkeypatch.setattr(utils_mod, "get_proxy_hook", fake_get_proxy_hook) + monkeypatch.setattr( + litellm.logging_callback_manager, + "add_litellm_callback", + lambda cb: registered.append(cb), + ) + + with patch("litellm.proxy.proxy_server.prisma_client", None): + proxy_logging._add_proxy_hooks(llm_router=None) + + keys = list(proxy_logging.proxy_hook_mapping.keys()) + snapshot = { + "mapping_keys": keys, + "registered_count": len(registered), + "registered_hook_names": [getattr(r, "hook_name", None) for r in registered], + } + assert snapshot == { + "mapping_keys": hook_keys, + "registered_count": len(hook_keys), + "registered_hook_names": hook_keys, + } + + +def test_add_proxy_hooks_unknown_hook_raises(proxy_logging, monkeypatch): + from litellm.proxy import utils as utils_mod + + monkeypatch.setattr(utils_mod, "PROXY_HOOKS", ["bogus_hook"]) + + def bad_resolver(name): + raise KeyError(name) + + monkeypatch.setattr(utils_mod, "get_proxy_hook", bad_resolver) + with pytest.raises(KeyError): + proxy_logging._add_proxy_hooks(llm_router=None) + + +# --------------------------------------------------------------------------- +# get_proxy_hook +# --------------------------------------------------------------------------- + + +def test_get_proxy_hook_returns_registered_instance(proxy_logging): + s_cache = MagicMock() + s_budget = MagicMock() + s_parallel = MagicMock() + proxy_logging.proxy_hook_mapping = { + "cache_control_check": s_cache, + "max_budget_limiter": s_budget, + "max_parallel_request_limiter": s_parallel, + } + snapshot = { + "cache_control_check": proxy_logging.get_proxy_hook("cache_control_check") is s_cache, + "max_budget_limiter": proxy_logging.get_proxy_hook("max_budget_limiter") is s_budget, + "max_parallel_request_limiter": proxy_logging.get_proxy_hook("max_parallel_request_limiter") is s_parallel, + "unknown_returns_none": proxy_logging.get_proxy_hook("unknown") is None, + } + assert snapshot == { + "cache_control_check": True, + "max_budget_limiter": True, + "max_parallel_request_limiter": True, + "unknown_returns_none": True, + } + + +def test_get_proxy_hook_unknown_returns_none(proxy_logging): + proxy_logging.proxy_hook_mapping = {} + assert proxy_logging.get_proxy_hook("does-not-exist") is None + + +def test_get_proxy_hook_non_string_key_raises(proxy_logging): + # ``dict.get`` doesn't raise on unhashable types — but ``None`` returns None. + # The pin: passing an unhashable key blows up like dict access does. + proxy_logging.proxy_hook_mapping = {"k": object()} + with pytest.raises(TypeError): + proxy_logging.get_proxy_hook({"unhashable": True}) # type: ignore[arg-type] + + +# --------------------------------------------------------------------------- +# _init_litellm_callbacks +# --------------------------------------------------------------------------- + + +def test_init_litellm_callbacks_replaces_string_with_instance(proxy_logging, monkeypatch): + from litellm.proxy import utils as utils_mod + + sentinel_instance = MagicMock(spec=litellm.integrations.custom_logger.CustomLogger) + sentinel_instance.__class__ = litellm.integrations.custom_logger.CustomLogger + + monkeypatch.setattr(litellm, "callbacks", ["some-string-logger"]) + monkeypatch.setattr( + litellm.litellm_core_utils.litellm_logging, + "_init_custom_logger_compatible_class", + lambda *a, **kw: sentinel_instance, + ) + + monkeypatch.setattr(utils_mod, "PROXY_HOOKS", []) + proxy_logging._init_litellm_callbacks(llm_router=None) + snapshot = { + "replaced_first_item": litellm.callbacks[0] is sentinel_instance, + "callbacks_grew_with_service": len(litellm.callbacks) >= 2, + "service_logging_appended": any( + "ServiceLogging" in type(c).__name__ for c in litellm.callbacks + ), + } + assert snapshot == { + "replaced_first_item": True, + "callbacks_grew_with_service": True, + "service_logging_appended": True, + } + + +def test_init_litellm_callbacks_string_resolution_failure_keeps_string(proxy_logging, monkeypatch): + from litellm.proxy import utils as utils_mod + + monkeypatch.setattr(litellm, "callbacks", ["unknown-logger"]) + monkeypatch.setattr( + litellm.litellm_core_utils.litellm_logging, + "_init_custom_logger_compatible_class", + lambda *a, **kw: None, + ) + monkeypatch.setattr(utils_mod, "PROXY_HOOKS", []) + proxy_logging._init_litellm_callbacks(llm_router=None) + # Resolver returned None — original string remains in place at idx 0. + assert litellm.callbacks[0] == "unknown-logger" + + +def test_init_litellm_callbacks_propagates_resolver_error_raises(proxy_logging, monkeypatch): + from litellm.proxy import utils as utils_mod + + monkeypatch.setattr(litellm, "callbacks", ["raises-on-init"]) + monkeypatch.setattr( + litellm.litellm_core_utils.litellm_logging, + "_init_custom_logger_compatible_class", + MagicMock(side_effect=RuntimeError("bad init")), + ) + monkeypatch.setattr(utils_mod, "PROXY_HOOKS", []) + with pytest.raises(RuntimeError): + proxy_logging._init_litellm_callbacks(llm_router=None) + + +# --------------------------------------------------------------------------- +# update_request_status +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_update_request_status_when_alerting_set_writes_cache(proxy_logging): + proxy_logging.alerting = ["slack"] + proxy_logging.alerting_threshold = 5.0 + captured: Dict[str, Any] = {} + + async def fake_set_cache(**kwargs): + captured.update(kwargs) + + proxy_logging.internal_usage_cache.async_set_cache = fake_set_cache # type: ignore[assignment] + await proxy_logging.update_request_status(litellm_call_id="call-1", status="success") + snapshot = { + "key": captured["key"], + "value": captured["value"], + "local_only": captured["local_only"], + "ttl": captured["ttl"], + } + assert snapshot == { + "key": "request_status:call-1", + "value": "success", + "local_only": True, + "ttl": 105.0, + } + + +@pytest.mark.asyncio +async def test_update_request_status_no_alerting_skips_cache(proxy_logging): + proxy_logging.alerting = None + proxy_logging.internal_usage_cache.async_set_cache = AsyncMock() + await proxy_logging.update_request_status(litellm_call_id="call-1", status="success") + proxy_logging.internal_usage_cache.async_set_cache.assert_not_called() + + +@pytest.mark.asyncio +async def test_update_request_status_cache_error_raises(proxy_logging): + proxy_logging.alerting = ["slack"] + proxy_logging.internal_usage_cache.async_set_cache = AsyncMock(side_effect=ConnectionError("redis")) + with pytest.raises(ConnectionError): + await proxy_logging.update_request_status(litellm_call_id="x", status="fail") + + +# --------------------------------------------------------------------------- +# _convert_user_api_key_auth_to_dict +# --------------------------------------------------------------------------- + + +def test_convert_user_api_key_auth_to_dict_pydantic_uses_model_dump(proxy_logging, make_user_api_key_auth): + auth = make_user_api_key_auth(user_id="u-1", team_id="t-1") + result = proxy_logging._convert_user_api_key_auth_to_dict(auth) + snapshot = { + "user_id": result["user_id"], + "team_id": result["team_id"], + "is_dict": isinstance(result, dict), + } + assert snapshot == {"user_id": "u-1", "team_id": "t-1", "is_dict": True} + + +def test_convert_user_api_key_auth_to_dict_plain_object_uses_dict(proxy_logging): + class Obj: + pass + + obj = Obj() + obj.a = 1 + obj.b = 2 + obj.c = 3 + result = proxy_logging._convert_user_api_key_auth_to_dict(obj) + assert result == {"a": 1, "b": 2, "c": 3} + + +def test_convert_user_api_key_auth_to_dict_none_returns_empty_dict(proxy_logging): + assert proxy_logging._convert_user_api_key_auth_to_dict(None) == {} + + +def test_convert_user_api_key_auth_to_dict_unconvertible_object_returns_empty(proxy_logging): + class NoDict: + __slots__ = () + + assert proxy_logging._convert_user_api_key_auth_to_dict(NoDict()) == {} + + +def test_convert_user_api_key_auth_to_dict_pydantic_error_raises(proxy_logging): + """A ``model_dump`` that raises propagates.""" + + class _Boom: + def model_dump(self): + raise RuntimeError("model_dump failure") + + with pytest.raises(RuntimeError): + proxy_logging._convert_user_api_key_auth_to_dict(_Boom()) diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_mcp_bridging.py b/tests/test_litellm/proxy/utils/proxy_logging/test_mcp_bridging.py new file mode 100644 index 00000000000..9defb309863 --- /dev/null +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_mcp_bridging.py @@ -0,0 +1,426 @@ +"""Pin ProxyLogging's MCP-LLM bridging helpers. + +Covers: +- ``_convert_mcp_to_llm_format`` +- ``_convert_llm_result_to_mcp_response`` +- ``_extract_modified_arguments_from_content`` +- ``_parse_arguments_manually`` +- ``_convert_llm_result_to_mcp_during_response`` +- ``_parse_pre_mcp_call_hook_response`` +- ``_create_mcp_request_object_from_kwargs`` +- ``_convert_mcp_hook_response_to_kwargs`` +""" + +from __future__ import annotations + +import pytest + +from litellm.types.mcp import ( + MCPDuringCallResponseObject, + MCPPreCallRequestObject, + MCPPreCallResponseObject, +) + + +# --------------------------------------------------------------------------- +# _convert_mcp_to_llm_format +# --------------------------------------------------------------------------- + + +def test_convert_mcp_to_llm_format_returns_synthetic_data(proxy_logging, make_mcp_request_obj): + req = make_mcp_request_obj(tool_name="search", arguments={"q": "hello"}) + out = proxy_logging._convert_mcp_to_llm_format( + request_obj=req, + kwargs={ + "model": "gpt-4o-mini", + "user_api_key_user_id": "u-1", + "user_api_key_team_id": "t-1", + "user_api_key_end_user_id": "eu-1", + "user_api_key_hash": "hash", + "user_api_key_request_route": "/mcp", + "incoming_bearer_token": "tok", + }, + ) + snapshot = { + "model": out["model"], + "user_id": out["user_api_key_user_id"], + "mcp_tool_name": out["mcp_tool_name"], + "mcp_arguments": out["mcp_arguments"], + "incoming_bearer_token": out["incoming_bearer_token"], + "message_role": out["messages"][0]["role"], + } + assert snapshot == { + "model": "gpt-4o-mini", + "user_id": "u-1", + "mcp_tool_name": "search", + "mcp_arguments": {"q": "hello"}, + "incoming_bearer_token": "tok", + "message_role": "user", + } + + +def test_convert_mcp_to_llm_format_defaults_model(proxy_logging, make_mcp_request_obj): + req = make_mcp_request_obj() + out = proxy_logging._convert_mcp_to_llm_format(request_obj=req, kwargs={}) + snapshot = { + "model": out["model"], + "mcp_tool_name": out["mcp_tool_name"], + "incoming_bearer_token": out["incoming_bearer_token"], + "user_id": out["user_api_key_user_id"], + } + assert snapshot == { + "model": "mcp-tool-call", + "mcp_tool_name": "calculator", + "incoming_bearer_token": None, + "user_id": None, + } + + +def test_convert_mcp_to_llm_format_missing_request_obj_raises(proxy_logging): + with pytest.raises(AttributeError): + proxy_logging._convert_mcp_to_llm_format(request_obj=None, kwargs={}) + + +# --------------------------------------------------------------------------- +# _convert_llm_result_to_mcp_response +# --------------------------------------------------------------------------- + + +def test_convert_llm_result_to_mcp_response_exception_blocks(proxy_logging, make_mcp_request_obj): + req = make_mcp_request_obj() + result = proxy_logging._convert_llm_result_to_mcp_response( + llm_result=ValueError("boom"), + request_obj=req, + ) + assert isinstance(result, MCPPreCallResponseObject) + snapshot = { + "should_proceed": result.should_proceed, + "error_message": result.error_message, + "modified_arguments": result.modified_arguments, + } + assert snapshot == {"should_proceed": False, "error_message": "boom", "modified_arguments": None} + + +def test_convert_llm_result_to_mcp_response_blocked_content(proxy_logging, make_mcp_request_obj): + req = make_mcp_request_obj(tool_name="t", arguments={"a": 1}) + llm_result = {"messages": [{"content": "this is blocked"}]} + result = proxy_logging._convert_llm_result_to_mcp_response(llm_result=llm_result, request_obj=req) + assert isinstance(result, MCPPreCallResponseObject) + assert result.should_proceed is False + assert "blocked" in (result.error_message or "").lower() + + +def test_convert_llm_result_to_mcp_response_modified_content_redacted(proxy_logging, make_mcp_request_obj): + req = make_mcp_request_obj(tool_name="search", arguments={"q": "ssn 123"}) + llm_result = {"messages": [{"content": "Tool: search\nArguments: {\"q\": \"[REDACTED]\"}"}]} + result = proxy_logging._convert_llm_result_to_mcp_response(llm_result=llm_result, request_obj=req) + assert isinstance(result, MCPPreCallResponseObject) + snapshot = { + "should_proceed": result.should_proceed, + "modified_q": (result.modified_arguments or {}).get("q"), + "error": result.error_message, + } + assert snapshot == {"should_proceed": True, "modified_q": "[REDACTED]", "error": None} + + +def test_convert_llm_result_to_mcp_response_string_blocks(proxy_logging, make_mcp_request_obj): + req = make_mcp_request_obj() + result = proxy_logging._convert_llm_result_to_mcp_response(llm_result="bad input", request_obj=req) + assert isinstance(result, MCPPreCallResponseObject) + snapshot = { + "should_proceed": result.should_proceed, + "error_message": result.error_message, + "modified_arguments": result.modified_arguments, + } + assert snapshot == {"should_proceed": False, "error_message": "bad input", "modified_arguments": None} + + +def test_convert_llm_result_to_mcp_response_unmodified_returns_none(proxy_logging, make_mcp_request_obj): + req = make_mcp_request_obj(tool_name="x", arguments={"a": 1}) + same_content = "Tool: x\nArguments: {'a': 1}" + result = proxy_logging._convert_llm_result_to_mcp_response( + llm_result={"messages": [{"content": same_content}]}, + request_obj=req, + ) + assert result is None + + +def test_convert_llm_result_to_mcp_response_no_request_obj_raises(proxy_logging): + with pytest.raises(AttributeError): + proxy_logging._convert_llm_result_to_mcp_response(llm_result={"messages": [{"content": "x"}]}, request_obj=None) + + +# --------------------------------------------------------------------------- +# _extract_modified_arguments_from_content +# --------------------------------------------------------------------------- + + +def test_extract_modified_arguments_from_content_parses_json(proxy_logging, make_mcp_request_obj): + req = make_mcp_request_obj() + out = proxy_logging._extract_modified_arguments_from_content( + masked_content="Tool: x\nArguments: {\"a\": 1, \"b\": 2, \"c\": 3}", + request_obj=req, + ) + assert out == {"a": 1, "b": 2, "c": 3} + + +def test_extract_modified_arguments_from_content_no_arguments_line_returns_none(proxy_logging, make_mcp_request_obj): + out = proxy_logging._extract_modified_arguments_from_content( + masked_content="random content with no arguments", + request_obj=make_mcp_request_obj(), + ) + assert out is None + + +def test_extract_modified_arguments_from_content_empty_string_returns_none(proxy_logging, make_mcp_request_obj): + out = proxy_logging._extract_modified_arguments_from_content( + masked_content="", + request_obj=make_mcp_request_obj(), + ) + assert out is None + + +def test_extract_modified_arguments_from_content_invalid_json_falls_back(proxy_logging, make_mcp_request_obj): + req = make_mcp_request_obj(arguments={"name": "alice"}) + out = proxy_logging._extract_modified_arguments_from_content( + masked_content="Tool: x\nArguments: {name: REDACTED}", + request_obj=req, + ) + assert isinstance(out, dict) + assert "name" in out + + +def test_extract_modified_arguments_from_content_error_swallowed_returns_none(proxy_logging): + """Internal try/except swallows any unexpected error and returns None.""" + out = proxy_logging._extract_modified_arguments_from_content(masked_content=None, request_obj=None) + assert out is None + + +# --------------------------------------------------------------------------- +# _parse_arguments_manually +# --------------------------------------------------------------------------- + + +def test_parse_arguments_manually_applies_overrides(proxy_logging): + original = {"name": "alice", "ssn": "123-45-6789"} + out = proxy_logging._parse_arguments_manually( + args_text='"name": "[REDACTED]", "ssn": "[REDACTED]"', + original_args=original, + ) + snapshot = {"name": out["name"], "ssn": out["ssn"], "original_unchanged": original["name"]} + assert snapshot == {"name": "[REDACTED]", "ssn": "[REDACTED]", "original_unchanged": "alice"} + + +def test_parse_arguments_manually_returns_original_if_no_match(proxy_logging): + original = {"foo": "bar"} + out = proxy_logging._parse_arguments_manually(args_text="nothing here", original_args=original) + assert out == {"foo": "bar"} + + +def test_parse_arguments_manually_error_swallowed_returns_none(proxy_logging): + # Defensive: function catches any exception internally and returns None. + assert proxy_logging._parse_arguments_manually(args_text="x", original_args=None) is None # type: ignore[arg-type] + + +# --------------------------------------------------------------------------- +# _convert_llm_result_to_mcp_during_response +# --------------------------------------------------------------------------- + + +def test_convert_llm_result_to_mcp_during_response_exception(proxy_logging, make_mcp_request_obj): + req = make_mcp_request_obj() + result = proxy_logging._convert_llm_result_to_mcp_during_response( + llm_result=ValueError("during boom"), request_obj=req + ) + assert isinstance(result, MCPDuringCallResponseObject) + snapshot = { + "should_continue": result.should_continue, + "error_message": result.error_message, + "type": type(result).__name__, + } + assert snapshot == { + "should_continue": False, + "error_message": "during boom", + "type": "MCPDuringCallResponseObject", + } + + +def test_convert_llm_result_to_mcp_during_response_blocked_content(proxy_logging, make_mcp_request_obj): + req = make_mcp_request_obj(tool_name="t", arguments={"a": 1}) + result = proxy_logging._convert_llm_result_to_mcp_during_response( + llm_result={"messages": [{"content": "blocked content"}]}, + request_obj=req, + ) + assert isinstance(result, MCPDuringCallResponseObject) + assert result.should_continue is False + assert "blocked" in (result.error_message or "").lower() + + +def test_convert_llm_result_to_mcp_during_response_modified_stops(proxy_logging, make_mcp_request_obj): + req = make_mcp_request_obj(tool_name="t", arguments={"a": 1}) + result = proxy_logging._convert_llm_result_to_mcp_during_response( + llm_result={"messages": [{"content": "Tool: t\nArguments: {\"a\": \"[REDACTED]\"}"}]}, + request_obj=req, + ) + assert isinstance(result, MCPDuringCallResponseObject) + assert result.should_continue is False + assert "modified" in (result.error_message or "").lower() + + +def test_convert_llm_result_to_mcp_during_response_string_blocks(proxy_logging, make_mcp_request_obj): + req = make_mcp_request_obj() + result = proxy_logging._convert_llm_result_to_mcp_during_response( + llm_result="kill switch", request_obj=req + ) + assert isinstance(result, MCPDuringCallResponseObject) + snapshot = {"should_continue": result.should_continue, "error_message": result.error_message} + assert snapshot == {"should_continue": False, "error_message": "kill switch"} + + +def test_convert_llm_result_to_mcp_during_response_unmodified_returns_none(proxy_logging, make_mcp_request_obj): + req = make_mcp_request_obj(tool_name="t", arguments={"a": 1}) + same = "Tool: t\nArguments: {'a': 1}" + assert ( + proxy_logging._convert_llm_result_to_mcp_during_response( + llm_result={"messages": [{"content": same}]}, + request_obj=req, + ) + is None + ) + + +def test_convert_llm_result_to_mcp_during_response_no_request_obj_raises(proxy_logging): + with pytest.raises(AttributeError): + proxy_logging._convert_llm_result_to_mcp_during_response( + llm_result={"messages": [{"content": "x"}]}, request_obj=None + ) + + +# --------------------------------------------------------------------------- +# _parse_pre_mcp_call_hook_response +# --------------------------------------------------------------------------- + + +def test_parse_pre_mcp_call_hook_response_with_modified_args(proxy_logging, make_mcp_request_obj): + req = make_mcp_request_obj(arguments={"a": 1}) + resp = MCPPreCallResponseObject( + should_proceed=True, + modified_arguments={"a": "x", "b": "y"}, + error_message=None, + ) + out = proxy_logging._parse_pre_mcp_call_hook_response(response=resp, original_request=req) + snapshot = { + "should_proceed": out["should_proceed"], + "modified_arguments": out["modified_arguments"], + "error_message": out["error_message"], + "hidden_params_type": type(out["hidden_params"]).__name__, + } + assert snapshot == { + "should_proceed": True, + "modified_arguments": {"a": "x", "b": "y"}, + "error_message": None, + "hidden_params_type": "HiddenParams", + } + + +def test_parse_pre_mcp_call_hook_response_no_modifications_uses_original(proxy_logging, make_mcp_request_obj): + req = make_mcp_request_obj(arguments={"original": True}) + resp = MCPPreCallResponseObject( + should_proceed=True, modified_arguments=None, error_message=None + ) + out = proxy_logging._parse_pre_mcp_call_hook_response(response=resp, original_request=req) + assert out["modified_arguments"] == {"original": True} + + +def test_parse_pre_mcp_call_hook_response_invalid_response_raises(proxy_logging, make_mcp_request_obj): + with pytest.raises(AttributeError): + proxy_logging._parse_pre_mcp_call_hook_response( + response=None, original_request=make_mcp_request_obj() + ) + + +# --------------------------------------------------------------------------- +# _create_mcp_request_object_from_kwargs +# --------------------------------------------------------------------------- + + +def test_create_mcp_request_object_from_kwargs_full(proxy_logging, make_user_api_key_auth): + auth = make_user_api_key_auth(user_id="u-1") + obj = proxy_logging._create_mcp_request_object_from_kwargs( + kwargs={ + "name": "calc", + "arguments": {"x": 1}, + "server_name": "math", + "user_api_key_auth": auth, + } + ) + assert isinstance(obj, MCPPreCallRequestObject) + snapshot = { + "tool_name": obj.tool_name, + "arguments": obj.arguments, + "server_name": obj.server_name, + "auth_user_id": obj.user_api_key_auth.get("user_id"), + } + assert snapshot == {"tool_name": "calc", "arguments": {"x": 1}, "server_name": "math", "auth_user_id": "u-1"} + + +def test_create_mcp_request_object_from_kwargs_empty(proxy_logging): + obj = proxy_logging._create_mcp_request_object_from_kwargs(kwargs={}) + snapshot = { + "tool_name": obj.tool_name, + "arguments": obj.arguments, + "server_name": obj.server_name, + } + assert snapshot == {"tool_name": "", "arguments": {}, "server_name": None} + + +def test_create_mcp_request_object_from_kwargs_non_dict_raises(proxy_logging): + with pytest.raises(AttributeError): + proxy_logging._create_mcp_request_object_from_kwargs(kwargs=None) # type: ignore[arg-type] + + +# --------------------------------------------------------------------------- +# _convert_mcp_hook_response_to_kwargs +# --------------------------------------------------------------------------- + + +def test_convert_mcp_hook_response_to_kwargs_applies_modified_args(proxy_logging): + original = {"arguments": {"a": 1}, "name": "old"} + out = proxy_logging._convert_mcp_hook_response_to_kwargs( + response_data={"modified_arguments": {"a": 2}, "extra_headers": {"H": "1"}}, + original_kwargs=original, + ) + snapshot = { + "arguments": out["arguments"], + "extra_headers": out["extra_headers"], + "name": out["name"], + "original_unmodified": original["arguments"], + } + assert snapshot == { + "arguments": {"a": 2}, + "extra_headers": {"H": "1"}, + "name": "old", + "original_unmodified": {"a": 1}, + } + + +def test_convert_mcp_hook_response_to_kwargs_merges_headers(proxy_logging): + original = {"extra_headers": {"keep": "yes", "overwrite": "old"}} + out = proxy_logging._convert_mcp_hook_response_to_kwargs( + response_data={"extra_headers": {"overwrite": "new", "added": "1"}}, + original_kwargs=original, + ) + assert out["extra_headers"] == {"keep": "yes", "overwrite": "new", "added": "1"} + + +def test_convert_mcp_hook_response_to_kwargs_no_response_data_returns_original(proxy_logging): + original = {"a": 1} + out = proxy_logging._convert_mcp_hook_response_to_kwargs(response_data=None, original_kwargs=original) + assert out is original + + +def test_convert_mcp_hook_response_to_kwargs_invalid_original_raises(proxy_logging): + with pytest.raises(AttributeError): + proxy_logging._convert_mcp_hook_response_to_kwargs( + response_data={"modified_arguments": {"a": 1}}, original_kwargs=None # type: ignore[arg-type] + ) diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_module_helpers.py b/tests/test_litellm/proxy/utils/proxy_logging/test_module_helpers.py new file mode 100644 index 00000000000..c491f16f2e4 --- /dev/null +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_module_helpers.py @@ -0,0 +1,353 @@ +"""Pin behavior of top-of-file and bottom-of-region helpers. + +Covers ``print_verbose``, ``_get_email_logger_class``, +``_accepts_litellm_call_info``, ``_enrich_http_exception_with_guardrail_context``, +``on_backoff``, ``jsonify_object``, ``_lookup_deprecated_key``. +""" + +from __future__ import annotations + +from datetime import datetime, timedelta, timezone +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi import HTTPException + +import litellm +from litellm.proxy import utils as utils_mod +from litellm.proxy.utils import ( + _accepts_litellm_call_info, + _enrich_http_exception_with_guardrail_context, + _get_email_logger_class, + _lookup_deprecated_key, + jsonify_object, + on_backoff, + print_verbose, +) + + +# --------------------------------------------------------------------------- +# print_verbose +# --------------------------------------------------------------------------- + + +def test_print_verbose_when_set_verbose_true_prints_redacted(monkeypatch, capsys): + monkeypatch.setattr(litellm, "set_verbose", True) + print_verbose("hello world") + captured = capsys.readouterr() + snapshot = { + "out_has_prefix": "LiteLLM Proxy:" in captured.out, + "out_has_payload": "hello world" in captured.out, + "no_stderr": captured.err == "", + } + assert snapshot == {"out_has_prefix": True, "out_has_payload": True, "no_stderr": True} + + +def test_print_verbose_when_set_verbose_false_no_stdout(monkeypatch, capsys): + monkeypatch.setattr(litellm, "set_verbose", False) + print_verbose("quiet") + captured = capsys.readouterr() + assert captured.out == "" + + +def test_print_verbose_handles_unprintable_object_raises(monkeypatch): + monkeypatch.setattr(litellm, "set_verbose", True) + + class Bomb: + def __str__(self): + raise RuntimeError("bad str") + + with pytest.raises(RuntimeError): + print_verbose(Bomb()) + + +# --------------------------------------------------------------------------- +# _get_email_logger_class +# --------------------------------------------------------------------------- + + +def test_get_email_logger_class_priority_matrix(monkeypatch): + """Truth table for ``_get_email_logger_class`` priority: SendGrid > + Resend > SMTP > Base.""" + sg = object() + rs = object() + smtp = object() + base = object() + monkeypatch.setattr(utils_mod, "BaseEmailLogger", base) + monkeypatch.setattr(utils_mod, "SendGridEmailLogger", sg) + monkeypatch.setattr(utils_mod, "ResendEmailLogger", rs) + monkeypatch.setattr(utils_mod, "SMTPEmailLogger", smtp) + for k in ("SENDGRID_API_KEY", "RESEND_API_KEY", "SMTP_HOST"): + monkeypatch.delenv(k, raising=False) + + fallback = _get_email_logger_class() is base + monkeypatch.setenv("SMTP_HOST", "smtp.example") + smtp_choice = _get_email_logger_class() is smtp + monkeypatch.setenv("RESEND_API_KEY", "rs-x") + resend_choice = _get_email_logger_class() is rs + monkeypatch.setenv("SENDGRID_API_KEY", "sg-x") + sendgrid_choice = _get_email_logger_class() is sg + snapshot = { + "fallback_to_base": fallback, + "smtp_when_smtp_only": smtp_choice, + "resend_beats_smtp": resend_choice, + "sendgrid_wins": sendgrid_choice, + } + assert snapshot == { + "fallback_to_base": True, + "smtp_when_smtp_only": True, + "resend_beats_smtp": True, + "sendgrid_wins": True, + } + + +def test_get_email_logger_class_error_when_no_enterprise_module(monkeypatch): + monkeypatch.setattr(utils_mod, "BaseEmailLogger", None) + # Returns ``None`` rather than raising; this is the documented failure + # mode when the optional enterprise package is missing. + assert _get_email_logger_class() is None + # Sentinel: monkey-patch SendGrid env but keep BaseEmailLogger None; + # function still must return None and not blow up on the optional path. + monkeypatch.setenv("SENDGRID_API_KEY", "sg-x") + assert _get_email_logger_class() is None + + +# --------------------------------------------------------------------------- +# _accepts_litellm_call_info +# --------------------------------------------------------------------------- + + +class _CbAcceptsInfo: + async def async_post_call_response_headers_hook(self, *, litellm_call_info=None): + return None + + +class _CbRejectsInfo: + async def async_post_call_response_headers_hook(self, *, response): + return None + + +def test_accepts_litellm_call_info_matrix(monkeypatch): + monkeypatch.setattr(utils_mod, "_CALLBACK_ACCEPTS_CALL_INFO", {}) + cache = {id(_CbAcceptsInfo): True} + monkeypatch.setattr(utils_mod, "_CALLBACK_ACCEPTS_CALL_INFO", cache) + snapshot = { + "cache_hit_returns_true": _accepts_litellm_call_info(_CbAcceptsInfo()), + "cache_size_after_hit": len(cache), + "cache_keyed_by_type_id": id(_CbAcceptsInfo) in cache, + } + assert snapshot == { + "cache_hit_returns_true": True, + "cache_size_after_hit": 1, + "cache_keyed_by_type_id": True, + } + + +def test_accepts_litellm_call_info_signature_inspection(monkeypatch): + monkeypatch.setattr(utils_mod, "_CALLBACK_ACCEPTS_CALL_INFO", {}) + snapshot = { + "accepts_param_true": _accepts_litellm_call_info(_CbAcceptsInfo()), + "rejects_param_false": _accepts_litellm_call_info(_CbRejectsInfo()), + "cache_populated": len(utils_mod._CALLBACK_ACCEPTS_CALL_INFO) == 2, + } + assert snapshot == { + "accepts_param_true": True, + "rejects_param_false": False, + "cache_populated": True, + } + + +def test_accepts_litellm_call_info_error_on_callback_without_hook_raises(monkeypatch): + monkeypatch.setattr(utils_mod, "_CALLBACK_ACCEPTS_CALL_INFO", {}) + + class _Bad: + pass + + with pytest.raises(AttributeError): + _accepts_litellm_call_info(_Bad()) # type: ignore[arg-type] + + +# --------------------------------------------------------------------------- +# _enrich_http_exception_with_guardrail_context +# --------------------------------------------------------------------------- + + +def test_enrich_http_exception_adds_guardrail_name_and_mode(): + detail = {"error": "blocked"} + exc = HTTPException(status_code=400, detail=detail) + cb = MagicMock() + cb.guardrail_name = "presidio" + cb.event_hook = "pre_call" + + _enrich_http_exception_with_guardrail_context(exc, cb) + snapshot = { + "error": detail["error"], + "guardrail_name": detail["guardrail_name"], + "guardrail_mode": detail["guardrail_mode"], + } + assert snapshot == { + "error": "blocked", + "guardrail_name": "presidio", + "guardrail_mode": "pre_call", + } + + +def test_enrich_http_exception_does_not_overwrite_existing_keys(): + detail = {"error": "blocked", "guardrail_name": "explicit", "guardrail_mode": "during_call"} + exc = HTTPException(status_code=400, detail=detail) + cb = MagicMock() + cb.guardrail_name = "should-not-overwrite" + cb.event_hook = "should-not-overwrite" + _enrich_http_exception_with_guardrail_context(exc, cb) + assert detail == {"error": "blocked", "guardrail_name": "explicit", "guardrail_mode": "during_call"} + + +def test_enrich_http_exception_no_op_for_non_http_exception(): + other = ValueError("not http") + _enrich_http_exception_with_guardrail_context(other, MagicMock(guardrail_name="g")) + + +def test_enrich_http_exception_no_op_for_non_dict_detail(): + exc = HTTPException(status_code=400, detail="just a string") + _enrich_http_exception_with_guardrail_context(exc, MagicMock(guardrail_name="g")) + assert exc.detail == "just a string" + + +def test_enrich_http_exception_error_handling_does_not_raise(): + """``_enrich_http_exception_with_guardrail_context`` swallows mismatched + inputs (non-HTTPException, non-dict detail, no guardrail_name) and never + raises — verified by passing each pathological input in turn.""" + # Bare exception with no detail at all should not blow up. + bare = Exception("bare") + _enrich_http_exception_with_guardrail_context(bare, MagicMock(guardrail_name=None)) + # HTTPException with non-dict detail. + s = HTTPException(status_code=500, detail="str-detail") + _enrich_http_exception_with_guardrail_context(s, MagicMock(guardrail_name="g")) + assert s.detail == "str-detail" + + +def test_enrich_http_exception_with_falsy_attrs_does_not_set(): + detail = {"error": "blocked"} + exc = HTTPException(status_code=400, detail=detail) + cb = MagicMock() + cb.guardrail_name = None + cb.event_hook = None + _enrich_http_exception_with_guardrail_context(exc, cb) + assert detail == {"error": "blocked"} + + +# --------------------------------------------------------------------------- +# on_backoff +# --------------------------------------------------------------------------- + + +def test_on_backoff_invokes_print_verbose(monkeypatch): + captured = [] + monkeypatch.setattr(utils_mod, "print_verbose", lambda s: captured.append(s)) + on_backoff({"tries": 3}) + snapshot = {"len": len(captured), "first_has_attempt": "attempt" in captured[0], "first_has_3": "3" in captured[0]} + assert snapshot == {"len": 1, "first_has_attempt": True, "first_has_3": True} + + +def test_on_backoff_missing_tries_key_raises(): + with pytest.raises(KeyError): + on_backoff({}) + + +# --------------------------------------------------------------------------- +# jsonify_object +# --------------------------------------------------------------------------- + + +def test_jsonify_object_serializes_nested_dicts(): + src = {"plain": "x", "nested": {"a": 1, "b": 2}, "n": 42} + out = jsonify_object(src) + expected = {"plain": "x", "nested": '{"a": 1, "b": 2}', "n": 42} + assert out == expected + # Source is not mutated. + assert src == {"plain": "x", "nested": {"a": 1, "b": 2}, "n": 42} + + +def test_jsonify_object_failed_serialization_marks_value(monkeypatch): + class Unserialiseable: + pass + + src = {"name": "x", "bad": {"obj": Unserialiseable()}, "count": 1} + out = jsonify_object(src) + assert out == {"name": "x", "bad": "failed-to-serialize-json", "count": 1} + + +def test_jsonify_object_non_dict_input_raises(): + with pytest.raises(AttributeError): + jsonify_object("not a dict") # type: ignore[arg-type] + + +# --------------------------------------------------------------------------- +# _lookup_deprecated_key +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_lookup_deprecated_key_returns_active_token_id_and_caches(monkeypatch): + from litellm.caching.dual_cache import LimitedSizeOrderedDict + + fresh = LimitedSizeOrderedDict(max_size=1000) + monkeypatch.setattr(utils_mod, "_deprecated_key_cache", fresh) + + future = datetime.now(timezone.utc) + timedelta(hours=1) + deprecated_row = MagicMock() + deprecated_row.active_token_id = "active-123" + deprecated_row.revoke_at = future + + db = MagicMock() + db.litellm_deprecatedverificationtoken.find_first = AsyncMock(return_value=deprecated_row) + + result = await _lookup_deprecated_key(db=db, hashed_token="hash-abc") + cached_value = fresh.get("hash-abc") + snapshot = { + "result": result, + "cache_active_token_id": cached_value[0], + "cache_has_3_tuple": isinstance(cached_value, tuple) and len(cached_value) == 3, + } + assert snapshot == { + "result": "active-123", + "cache_active_token_id": "active-123", + "cache_has_3_tuple": True, + } + + +@pytest.mark.asyncio +async def test_lookup_deprecated_key_returns_none_when_not_found(monkeypatch): + from litellm.caching.dual_cache import LimitedSizeOrderedDict + + monkeypatch.setattr(utils_mod, "_deprecated_key_cache", LimitedSizeOrderedDict(max_size=10)) + db = MagicMock() + db.litellm_deprecatedverificationtoken.find_first = AsyncMock(return_value=None) + assert await _lookup_deprecated_key(db=db, hashed_token="missing") is None + + +@pytest.mark.asyncio +async def test_lookup_deprecated_key_db_error_returns_none(monkeypatch): + from litellm.caching.dual_cache import LimitedSizeOrderedDict + + monkeypatch.setattr(utils_mod, "_deprecated_key_cache", LimitedSizeOrderedDict(max_size=10)) + db = MagicMock() + db.litellm_deprecatedverificationtoken.find_first = AsyncMock(side_effect=RuntimeError("db down")) + result = await _lookup_deprecated_key(db=db, hashed_token="x") + assert result is None + + +@pytest.mark.asyncio +async def test_lookup_deprecated_key_uses_cache_within_ttl(monkeypatch): + from litellm.caching.dual_cache import LimitedSizeOrderedDict + + cache = LimitedSizeOrderedDict(max_size=10) + now_ts = datetime.now(timezone.utc).timestamp() + cache["hashY"] = ("active-from-cache", now_ts + 100, now_ts + 1000) + monkeypatch.setattr(utils_mod, "_deprecated_key_cache", cache) + + db = MagicMock() + db.litellm_deprecatedverificationtoken.find_first = AsyncMock(return_value=None) + result = await _lookup_deprecated_key(db=db, hashed_token="hashY") + assert result == "active-from-cache" + db.litellm_deprecatedverificationtoken.find_first.assert_not_called() diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_post_call_failure_hook.py b/tests/test_litellm/proxy/utils/proxy_logging/test_post_call_failure_hook.py new file mode 100644 index 00000000000..a2a57931d26 --- /dev/null +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_post_call_failure_hook.py @@ -0,0 +1,269 @@ +"""Pin ``ProxyLogging.post_call_failure_hook``, ``_is_proxy_only_llm_api_error``, +and ``_handle_logging_proxy_only_error``.""" + +from __future__ import annotations + +import asyncio +from typing import Any +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi import HTTPException + +import litellm +from litellm.integrations.custom_logger import CustomLogger +from litellm.proxy._types import AlertType, ProxyErrorTypes +from litellm.proxy.utils import ProxyLogging + + +@pytest.fixture(autouse=True) +def _clear_caps_cache(): + ProxyLogging._callback_capabilities_cache.clear() + yield + ProxyLogging._callback_capabilities_cache.clear() + + +# --------------------------------------------------------------------------- +# _is_proxy_only_llm_api_error +# --------------------------------------------------------------------------- + + +def test_is_proxy_only_llm_api_truth_table(proxy_logging): + """Pin the truth table of ``_is_proxy_only_llm_api_error`` in a single + snapshot. Covers no-route, non-LLM route, HTTPException on LLM route, + and auth-error short-circuit.""" + snapshot = { + "no_route": proxy_logging._is_proxy_only_llm_api_error( + original_exception=Exception(), route=None + ), + "non_llm_route": proxy_logging._is_proxy_only_llm_api_error( + original_exception=HTTPException(status_code=429, detail="rate"), + route="/random/path", + ), + "http_on_llm_route": proxy_logging._is_proxy_only_llm_api_error( + original_exception=HTTPException(status_code=429, detail="rate"), + route="/chat/completions", + ), + "auth_short_circuit": proxy_logging._is_proxy_only_llm_api_error( + original_exception=Exception("auth"), + error_type=ProxyErrorTypes.auth_error, + route="/chat/completions", + ), + } + assert snapshot == { + "no_route": False, + "non_llm_route": False, + "http_on_llm_route": True, + "auth_short_circuit": True, + } + + +def test_is_proxy_only_llm_api_missing_exception_raises(proxy_logging): + """Passing nothing should TypeError on the missing positional kwarg.""" + with pytest.raises(TypeError): + proxy_logging._is_proxy_only_llm_api_error() # type: ignore[call-arg] + + +# --------------------------------------------------------------------------- +# post_call_failure_hook +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_post_call_failure_hook_no_callbacks_returns_none( + proxy_logging, make_user_api_key_auth, mock_callbacks_disabled +): + proxy_logging.alert_types = [] + request_data = {"litellm_call_id": "abc", "model": "m", "messages": []} + out = await proxy_logging.post_call_failure_hook( + request_data=request_data, + original_exception=ValueError("oops"), + user_api_key_dict=make_user_api_key_auth(), + ) + snapshot = { + "out_is_none": out is None, + "litellm_logging_obj_popped": "litellm_logging_obj" not in request_data, + "call_id_preserved": request_data["litellm_call_id"] == "abc", + "first_api_call_start_time_present": "first_api_call_start_time" in request_data, + } + assert snapshot == { + "out_is_none": True, + "litellm_logging_obj_popped": True, + "call_id_preserved": True, + "first_api_call_start_time_present": False, + } + + +@pytest.mark.asyncio +async def test_post_call_failure_hook_callback_returns_http_exception( + proxy_logging, make_user_api_key_auth, monkeypatch +): + transformed = HTTPException(status_code=418, detail="teapot") + + class _Cb(CustomLogger): + async def async_post_call_failure_hook(self, **kwargs): # type: ignore[override] + return transformed + + monkeypatch.setattr(litellm, "callbacks", [_Cb()]) + proxy_logging.alert_types = [] + out = await proxy_logging.post_call_failure_hook( + request_data={"litellm_call_id": "abc"}, + original_exception=ValueError("oops"), + user_api_key_dict=make_user_api_key_auth(), + ) + assert out is transformed + + +@pytest.mark.asyncio +async def test_post_call_failure_hook_callback_raises_http_exception_first_wins( + proxy_logging, make_user_api_key_auth, monkeypatch +): + err = HTTPException(status_code=418, detail="raised teapot") + + class _Cb(CustomLogger): + async def async_post_call_failure_hook(self, **kwargs): # type: ignore[override] + raise err + + monkeypatch.setattr(litellm, "callbacks", [_Cb()]) + proxy_logging.alert_types = [] + out = await proxy_logging.post_call_failure_hook( + request_data={"litellm_call_id": "abc"}, + original_exception=ValueError("oops"), + user_api_key_dict=make_user_api_key_auth(), + ) + assert out is err + + +@pytest.mark.asyncio +async def test_post_call_failure_hook_non_http_exception_in_callback_swallowed( + proxy_logging, make_user_api_key_auth, monkeypatch +): + class _Cb(CustomLogger): + async def async_post_call_failure_hook(self, **kwargs): # type: ignore[override] + raise RuntimeError("non-http inside cb") + + monkeypatch.setattr(litellm, "callbacks", [_Cb()]) + proxy_logging.alert_types = [] + out = await proxy_logging.post_call_failure_hook( + request_data={"litellm_call_id": "abc"}, + original_exception=ValueError("oops"), + user_api_key_dict=make_user_api_key_auth(), + ) + assert out is None + + +# --------------------------------------------------------------------------- +# _handle_logging_proxy_only_error +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_handle_logging_proxy_only_path_uses_existing_logging_obj( + proxy_logging, make_user_api_key_auth +): + logging_obj = MagicMock() + logging_obj.call_type = "acompletion" + logging_obj.model_call_details = {} + logging_obj.async_failure_handler = AsyncMock() + + request_data = { + "litellm_logging_obj": logging_obj, + "messages": [{"role": "user", "content": "x"}], + "model": "m", + "metadata": {}, + } + await proxy_logging._handle_logging_proxy_only_error( + request_data=request_data, + user_api_key_dict=make_user_api_key_auth(), + route="/chat/completions", + original_exception=HTTPException(status_code=429, detail="rate"), + ) + from litellm.constants import LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL + + snapshot = { + "input_logged": "messages" in logging_obj.model_call_details, + "call_type_normalized": logging_obj.call_type, + "marker_present": logging_obj.model_call_details.get( + LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL + ) + is True, + "async_failure_called": logging_obj.async_failure_handler.called, + } + assert snapshot == { + "input_logged": True, + "call_type_normalized": "acompletion", + "marker_present": True, + "async_failure_called": True, + } + + +@pytest.mark.asyncio +async def test_handle_logging_proxy_only_path_skips_for_pass_through( + proxy_logging, make_user_api_key_auth +): + from litellm.types.utils import CallTypes + + logging_obj = MagicMock() + logging_obj.call_type = CallTypes.pass_through.value + logging_obj.model_call_details = {} + logging_obj.async_failure_handler = AsyncMock() + logging_obj.pre_call = MagicMock() + request_data = { + "litellm_logging_obj": logging_obj, + "messages": [{"role": "user"}], + "model": "m", + } + await proxy_logging._handle_logging_proxy_only_error( + request_data=request_data, + user_api_key_dict=make_user_api_key_auth(), + route="/chat/completions", + original_exception=HTTPException(status_code=429, detail="rate"), + ) + logging_obj.pre_call.assert_not_called() + logging_obj.async_failure_handler.assert_not_called() + + +@pytest.mark.asyncio +async def test_handle_logging_proxy_only_path_no_logging_obj_creates_one( + proxy_logging, make_user_api_key_auth, monkeypatch +): + fake_logging_obj = MagicMock() + fake_logging_obj.call_type = "acompletion" + fake_logging_obj.model_call_details = {} + fake_logging_obj.async_failure_handler = AsyncMock() + + def fake_function_setup(**kwargs): + return fake_logging_obj, {} + + monkeypatch.setattr(litellm.utils, "function_setup", fake_function_setup) + request_data = {"messages": [{"role": "user"}], "model": "m"} + await proxy_logging._handle_logging_proxy_only_error( + request_data=request_data, + user_api_key_dict=make_user_api_key_auth(), + route="/chat/completions", + original_exception=HTTPException(status_code=429, detail="rate"), + ) + assert "litellm_call_id" in request_data + fake_logging_obj.async_failure_handler.assert_called_once() + + +@pytest.mark.asyncio +async def test_handle_logging_proxy_only_path_propagates_async_failure_raises( + proxy_logging, make_user_api_key_auth +): + logging_obj = MagicMock() + logging_obj.call_type = "acompletion" + logging_obj.model_call_details = {} + logging_obj.async_failure_handler = AsyncMock(side_effect=RuntimeError("boom")) + request_data = { + "litellm_logging_obj": logging_obj, + "messages": [{"role": "user"}], + "model": "m", + } + with pytest.raises(RuntimeError): + await proxy_logging._handle_logging_proxy_only_error( + request_data=request_data, + user_api_key_dict=make_user_api_key_auth(), + route="/chat/completions", + original_exception=Exception("x"), + ) diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_post_call_success_hook.py b/tests/test_litellm/proxy/utils/proxy_logging/test_post_call_success_hook.py new file mode 100644 index 00000000000..6a339b37a80 --- /dev/null +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_post_call_success_hook.py @@ -0,0 +1,97 @@ +"""Pin ``ProxyLogging.post_call_success_hook``.""" + +from __future__ import annotations + +from typing import Any +from unittest.mock import AsyncMock, MagicMock + +import pytest + +import litellm +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.integrations.custom_logger import CustomLogger +from litellm.proxy.utils import ProxyLogging +from litellm.types.guardrails import GuardrailEventHooks + + +@pytest.fixture(autouse=True) +def _clear_caps_cache(): + ProxyLogging._callback_capabilities_cache.clear() + yield + ProxyLogging._callback_capabilities_cache.clear() + + +def _make_guardrail(name="g", should_run=True, override=None): + cb = MagicMock(spec=CustomGuardrail) + cb.__class__ = CustomGuardrail + cb.guardrail_name = name + cb.event_hook = GuardrailEventHooks.post_call + cb.should_run_guardrail = MagicMock(return_value=should_run) + cb.async_post_call_success_hook = AsyncMock(return_value=override) + return cb + + +@pytest.mark.asyncio +async def test_post_call_success_hook_returns_response_when_no_callbacks(proxy_logging, make_user_api_key_auth, mock_callbacks_disabled): + response = {"original": True, "model": "m", "choices": []} + out = await proxy_logging.post_call_success_hook( + data={}, response=response, user_api_key_dict=make_user_api_key_auth() + ) + assert out == {"original": True, "model": "m", "choices": []} + + +@pytest.mark.asyncio +async def test_post_call_success_hook_runs_other_callback_and_replaces_response( + proxy_logging, make_user_api_key_auth, monkeypatch +): + new_response = {"modified": True, "kept": "yes", "final": "v"} + + class _CL(CustomLogger): + async def async_post_call_success_hook(self, **kwargs): # type: ignore[override] + return new_response + + monkeypatch.setattr(litellm, "callbacks", [_CL()]) + out = await proxy_logging.post_call_success_hook( + data={}, response={"original": True}, user_api_key_dict=make_user_api_key_auth() + ) + assert out == new_response + + +@pytest.mark.asyncio +async def test_post_call_success_hook_guardrail_should_not_run_skipped( + proxy_logging, make_user_api_key_auth, monkeypatch +): + g = _make_guardrail(should_run=False) + monkeypatch.setattr(litellm, "callbacks", [g]) + response = MagicMock() + out = await proxy_logging.post_call_success_hook( + data={}, response=response, user_api_key_dict=make_user_api_key_auth() + ) + g.async_post_call_success_hook.assert_not_called() + assert out is response + + +@pytest.mark.asyncio +async def test_post_call_success_hook_guardrail_error_raises( + proxy_logging, make_user_api_key_auth, monkeypatch +): + g = _make_guardrail() + g.async_post_call_success_hook = AsyncMock(side_effect=RuntimeError("blocked")) + monkeypatch.setattr(litellm, "callbacks", [g]) + with pytest.raises(RuntimeError): + await proxy_logging.post_call_success_hook( + data={}, response=MagicMock(), user_api_key_dict=make_user_api_key_auth() + ) + + +@pytest.mark.asyncio +async def test_post_call_success_hook_guardrail_returns_modified_response( + proxy_logging, make_user_api_key_auth, monkeypatch +): + modified = {"a": 1, "b": 2, "c": 3} + g = _make_guardrail(override=modified) + monkeypatch.setattr(litellm, "callbacks", [g]) + out = await proxy_logging.post_call_success_hook( + data={}, response={"orig": True}, user_api_key_dict=make_user_api_key_auth() + ) + assert out == modified diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_pre_call_hook.py b/tests/test_litellm/proxy/utils/proxy_logging/test_pre_call_hook.py new file mode 100644 index 00000000000..05005dae797 --- /dev/null +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_pre_call_hook.py @@ -0,0 +1,168 @@ +"""Pin ``ProxyLogging.pre_call_hook`` and ``process_pre_call_hook_response``.""" + +from __future__ import annotations + +from typing import Any, Dict +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi import HTTPException + +import litellm +from litellm.exceptions import RejectedRequestError +from litellm.integrations.custom_logger import CustomLogger +from litellm.proxy.utils import ProxyLogging + + +@pytest.fixture(autouse=True) +def _clear_caps_cache(): + ProxyLogging._callback_capabilities_cache.clear() + yield + ProxyLogging._callback_capabilities_cache.clear() + + +# --------------------------------------------------------------------------- +# process_pre_call_hook_response +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_process_pre_call_hook_response_dict_returns_response(proxy_logging): + out = await proxy_logging.process_pre_call_hook_response( + response={"messages": [{"x": 1}], "model": "m", "temperature": 0.5}, + data={"original": True}, + call_type="completion", + ) + assert out == {"messages": [{"x": 1}], "model": "m", "temperature": 0.5} + + +@pytest.mark.asyncio +async def test_process_pre_call_hook_response_string_completion_raises_rejected(proxy_logging): + with pytest.raises(RejectedRequestError): + await proxy_logging.process_pre_call_hook_response( + response="rejected", + data={"model": "m"}, + call_type="completion", + ) + + +@pytest.mark.asyncio +async def test_process_pre_call_hook_response_string_other_call_type_raises_http(proxy_logging): + with pytest.raises(HTTPException) as info: + await proxy_logging.process_pre_call_hook_response( + response="bad", + data={}, + call_type="embeddings", + ) + assert info.value.status_code == 400 + + +@pytest.mark.asyncio +async def test_process_pre_call_hook_response_exception_reraises(proxy_logging): + err = RuntimeError("hook said no") + with pytest.raises(RuntimeError, match="hook said no"): + await proxy_logging.process_pre_call_hook_response( + response=err, data={}, call_type="completion" + ) + + +@pytest.mark.asyncio +async def test_process_pre_call_hook_response_other_type_returns_data(proxy_logging): + out = await proxy_logging.process_pre_call_hook_response( + response=12345, data={"a": 1, "b": 2, "c": 3}, call_type="completion" + ) + assert out == {"a": 1, "b": 2, "c": 3} + + +# --------------------------------------------------------------------------- +# pre_call_hook +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_pre_call_hook_returns_data_when_no_callbacks(proxy_logging, make_user_api_key_auth, mock_callbacks_disabled): + data = {"messages": [{"role": "user", "content": "hi"}], "model": "m", "temperature": 0.7} + proxy_logging.slack_alerting_instance = MagicMock(alerting=None) + out = await proxy_logging.pre_call_hook( + user_api_key_dict=make_user_api_key_auth(), + data=data, + call_type="completion", + ) + assert out is data + + +@pytest.mark.asyncio +async def test_pre_call_hook_returns_none_for_none_data(proxy_logging, make_user_api_key_auth, mock_callbacks_disabled): + proxy_logging.slack_alerting_instance = MagicMock(alerting=None) + out = await proxy_logging.pre_call_hook( + user_api_key_dict=make_user_api_key_auth(), + data=None, + call_type="completion", + ) + assert out is None + + +@pytest.mark.asyncio +async def test_pre_call_hook_invokes_pre_call_override(proxy_logging, make_user_api_key_auth, monkeypatch): + captured: Dict[str, Any] = {} + + class _Cb(CustomLogger): + async def async_pre_call_hook(self, **kwargs): # type: ignore[override] + captured.update(kwargs) + return {"messages": [{"x": "modified"}], "model": "m", "temperature": 0.1} + + monkeypatch.setattr(litellm, "callbacks", [_Cb()]) + proxy_logging.slack_alerting_instance = MagicMock(alerting=None) + out = await proxy_logging.pre_call_hook( + user_api_key_dict=make_user_api_key_auth(), + data={"messages": [{"x": "input"}], "model": "m", "temperature": 0.1}, + call_type="completion", + ) + snapshot = { + "out_messages": out["messages"], + "out_model": out["model"], + "out_temp": out["temperature"], + "cb_received_call_type": captured.get("call_type"), + } + assert snapshot == { + "out_messages": [{"x": "modified"}], + "out_model": "m", + "out_temp": 0.1, + "cb_received_call_type": "completion", + } + + +@pytest.mark.asyncio +async def test_pre_call_hook_propagates_callback_error_raises(proxy_logging, make_user_api_key_auth, monkeypatch): + class _BadCb(CustomLogger): + async def async_pre_call_hook(self, **kwargs): # type: ignore[override] + raise RuntimeError("rejected") + + monkeypatch.setattr(litellm, "callbacks", [_BadCb()]) + proxy_logging.slack_alerting_instance = MagicMock(alerting=None) + with pytest.raises(RuntimeError, match="rejected"): + await proxy_logging.pre_call_hook( + user_api_key_dict=make_user_api_key_auth(), + data={"model": "m"}, + call_type="completion", + ) + + +@pytest.mark.asyncio +async def test_pre_call_hook_processes_guardrail_metadata_when_no_overrides(proxy_logging, make_user_api_key_auth, mock_callbacks_disabled): + """Even when no callback overrides exist, ``_process_guardrail_metadata`` runs.""" + data = {"messages": [{"role": "user"}], "model": "m", "metadata": {"guardrails": ["g1"]}} + proxy_logging.slack_alerting_instance = MagicMock(alerting=None) + invoked = {} + + def fake_process(d): + invoked["data"] = d + + proxy_logging._process_guardrail_metadata = fake_process # type: ignore[assignment] + out = await proxy_logging.pre_call_hook( + user_api_key_dict=make_user_api_key_auth(), + data=data, + call_type="completion", + ) + assert out is data + assert invoked["data"] is data diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py b/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py new file mode 100644 index 00000000000..65d3c3c8079 --- /dev/null +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py @@ -0,0 +1,432 @@ +"""Pin ProxyLogging streaming + response-headers helpers. + +Covers ``_wrap_streaming_iterator_with_enrichment``, +``async_post_call_streaming_hook``, +``async_post_call_streaming_iterator_hook``, ``_fire_deferred_stream_logging``, +``is_a2a_streaming_response``, ``_init_response_taking_too_long_task``, +``post_call_response_headers_hook``, ``_build_litellm_call_info``. +""" + +from __future__ import annotations + +import asyncio +from typing import Any, Dict, List +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi import HTTPException + +import litellm +from litellm.integrations.custom_logger import CustomLogger +from litellm.proxy.utils import ProxyLogging + + +@pytest.fixture(autouse=True) +def _clear_caps_cache(): + ProxyLogging._callback_capabilities_cache.clear() + yield + ProxyLogging._callback_capabilities_cache.clear() + + +# --------------------------------------------------------------------------- +# is_a2a_streaming_response +# --------------------------------------------------------------------------- + + +def test_is_a2a_streaming_response_truth_matrix(proxy_logging): + snapshot = { + "all_three_keys_present": proxy_logging.is_a2a_streaming_response( + {"jsonrpc": "2.0", "id": "1", "result": {"x": 1}, "extra": "y"} + ), + "missing_result": proxy_logging.is_a2a_streaming_response( + {"jsonrpc": "2.0", "id": "1"} + ), + "missing_jsonrpc": proxy_logging.is_a2a_streaming_response( + {"id": "1", "result": {}} + ), + "empty_dict": proxy_logging.is_a2a_streaming_response({}), + } + assert snapshot == { + "all_three_keys_present": True, + "missing_result": False, + "missing_jsonrpc": False, + "empty_dict": False, + } + + +def test_is_a2a_streaming_response_invalid_input_raises(proxy_logging): + with pytest.raises(TypeError): + proxy_logging.is_a2a_streaming_response(None) # type: ignore[arg-type] + + +# --------------------------------------------------------------------------- +# _build_litellm_call_info +# --------------------------------------------------------------------------- + + +def test_build_litellm_call_info_pulls_from_hidden_params_and_metadata(proxy_logging): + response = MagicMock() + response._hidden_params = { + "custom_llm_provider": "openai", + "api_base": "https://api.openai.com", + "model_id": "model-1", + } + info = proxy_logging._build_litellm_call_info( + data={"metadata": {"model_info": {"name": "gpt-4o-mini"}}}, + response=response, + ) + assert info == { + "custom_llm_provider": "openai", + "model_info": {"name": "gpt-4o-mini"}, + "api_base": "https://api.openai.com", + "model_id": "model-1", + } + + +def test_build_litellm_call_info_fallbacks_to_litellm_metadata(proxy_logging): + response = MagicMock() + response._hidden_params = {"custom_llm_provider": "azure"} + info = proxy_logging._build_litellm_call_info( + data={"litellm_metadata": {"model_info": {"alias": "azure-gpt"}}}, + response=response, + ) + snapshot = { + "custom_llm_provider": info["custom_llm_provider"], + "model_info": info["model_info"], + "api_base": info["api_base"], + "model_id": info["model_id"], + } + assert snapshot == { + "custom_llm_provider": "azure", + "model_info": {"alias": "azure-gpt"}, + "api_base": None, + "model_id": None, + } + + +def test_build_litellm_call_info_invalid_data_raises(proxy_logging): + with pytest.raises(AttributeError): + proxy_logging._build_litellm_call_info(data=None, response=MagicMock()) # type: ignore[arg-type] + + +# --------------------------------------------------------------------------- +# _init_response_taking_too_long_task +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_init_response_taking_too_long_task_runs_when_alerting(proxy_logging): + proxy_logging.slack_alerting_instance = MagicMock() + proxy_logging.slack_alerting_instance.alerting = ["slack"] + captured: Dict[str, Any] = {} + + async def fake_resp_too_long(request_data): + captured["request_data"] = request_data + + proxy_logging.slack_alerting_instance.response_taking_too_long = fake_resp_too_long + payload = {"req": "y", "litellm_call_id": "c1", "model": "m"} + proxy_logging._init_response_taking_too_long_task(data=payload) + await asyncio.sleep(0) + snapshot = { + "received_payload": captured["request_data"], + "fired_once": len(captured) == 1, + "alerting_was_truthy": bool(proxy_logging.slack_alerting_instance.alerting), + } + assert snapshot == { + "received_payload": payload, + "fired_once": True, + "alerting_was_truthy": True, + } + + +@pytest.mark.asyncio +async def test_init_response_taking_too_long_task_no_op_when_alerting_off(proxy_logging): + proxy_logging.slack_alerting_instance = MagicMock() + proxy_logging.slack_alerting_instance.alerting = None + proxy_logging.slack_alerting_instance.response_taking_too_long = AsyncMock() + proxy_logging._init_response_taking_too_long_task(data=None) + await asyncio.sleep(0) + proxy_logging.slack_alerting_instance.response_taking_too_long.assert_not_called() + + +def test_init_response_taking_too_long_task_no_slack_instance_no_error_raises(proxy_logging): + proxy_logging.slack_alerting_instance = None + proxy_logging._init_response_taking_too_long_task(data=None) + + +# --------------------------------------------------------------------------- +# _wrap_streaming_iterator_with_enrichment +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_wrap_streaming_iterator_with_enrichment_passes_through_chunks(proxy_logging): + async def gen(): + for ch in ("a", "b", "c"): + yield ch + + cb = MagicMock(guardrail_name="g", event_hook="pre_call") + wrapped = proxy_logging._wrap_streaming_iterator_with_enrichment(callback=cb, gen=gen()) + out = [ch async for ch in wrapped] + snapshot = { + "chunks": out, + "count": len(out), + "first": out[0], + "last": out[-1], + } + assert snapshot == { + "chunks": ["a", "b", "c"], + "count": 3, + "first": "a", + "last": "c", + } + + +@pytest.mark.asyncio +async def test_wrap_streaming_iterator_with_enrichment_enriches_http_exception_raises(proxy_logging): + detail = {"error": "blocked"} + + async def boom_gen(): + if False: + yield # pragma: no cover + raise HTTPException(status_code=400, detail=detail) + + cb = MagicMock(guardrail_name="presidio", event_hook="post_call") + wrapped = proxy_logging._wrap_streaming_iterator_with_enrichment(callback=cb, gen=boom_gen()) + with pytest.raises(HTTPException): + async for _ in wrapped: + pass + assert detail["guardrail_name"] == "presidio" + assert detail["guardrail_mode"] == "post_call" + + +# --------------------------------------------------------------------------- +# async_post_call_streaming_hook +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_async_post_call_streaming_hook_fast_path_returns_response(proxy_logging, mock_callbacks_disabled, make_user_api_key_auth): + resp = "chunk-1" + out = await proxy_logging.async_post_call_streaming_hook( + data={}, response=resp, user_api_key_dict=make_user_api_key_auth() + ) + snapshot = { + "out_is_input": out is resp, + "out_value": out, + "type": type(out).__name__, + "callbacks_empty": len(litellm.callbacks) == 0, + } + assert snapshot == { + "out_is_input": True, + "out_value": "chunk-1", + "type": "str", + "callbacks_empty": True, + } + + +@pytest.mark.asyncio +async def test_async_post_call_streaming_hook_invokes_per_chunk_callback(proxy_logging, make_user_api_key_auth, monkeypatch): + class _Per(CustomLogger): + async def async_post_call_streaming_hook(self, **kwargs): # type: ignore[override] + return "modified-" + str(kwargs.get("response", "")) + + cb = _Per() + monkeypatch.setattr(litellm, "callbacks", [cb]) + + from litellm import ModelResponse + + fake_resp = ModelResponse( + id="rid", + choices=[{"index": 0, "delta": {"role": "assistant", "content": "hi"}, "finish_reason": None}], + created=0, + model="gpt-4o-mini", + object="chat.completion.chunk", + ) + out = await proxy_logging.async_post_call_streaming_hook( + data={}, + response=fake_resp, + user_api_key_dict=make_user_api_key_auth(), + ) + assert isinstance(out, str) + assert out.startswith("modified-") + + +@pytest.mark.asyncio +async def test_async_post_call_streaming_hook_callback_error_raises(proxy_logging, make_user_api_key_auth, monkeypatch): + class _Per(CustomLogger): + async def async_post_call_streaming_hook(self, **kwargs): # type: ignore[override] + raise RuntimeError("hook-fail") + + monkeypatch.setattr(litellm, "callbacks", [_Per()]) + + from litellm import ModelResponse + + fake_resp = ModelResponse( + id="rid", + choices=[{"index": 0, "delta": {"role": "assistant", "content": "hi"}, "finish_reason": None}], + created=0, + model="gpt-4o-mini", + object="chat.completion.chunk", + ) + with pytest.raises(RuntimeError): + await proxy_logging.async_post_call_streaming_hook( + data={}, + response=fake_resp, + user_api_key_dict=make_user_api_key_auth(), + ) + + +# --------------------------------------------------------------------------- +# async_post_call_streaming_iterator_hook +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_async_post_call_streaming_iterator_hook_no_overrides_passes_through(proxy_logging, make_user_api_key_auth, mock_callbacks_disabled): + async def gen(): + for ch in ("a", "b"): + yield ch + + chunks = [] + async for ch in proxy_logging.async_post_call_streaming_iterator_hook( + response=gen(), + user_api_key_dict=make_user_api_key_auth(), + request_data={}, + ): + chunks.append(ch) + snapshot = { + "chunks": chunks, + "count": len(chunks), + "passthrough_preserved_order": chunks == ["a", "b"], + } + assert snapshot == { + "chunks": ["a", "b"], + "count": 2, + "passthrough_preserved_order": True, + } + + +@pytest.mark.asyncio +async def test_async_post_call_streaming_iterator_hook_with_override_chains_callback(proxy_logging, make_user_api_key_auth, monkeypatch): + class _IterOverride(CustomLogger): + async def async_post_call_streaming_iterator_hook(self, **kwargs): # type: ignore[override] + async for ch in kwargs["response"]: + yield ch + "*" + + monkeypatch.setattr(litellm, "callbacks", [_IterOverride()]) + + async def gen(): + for ch in ("a", "b"): + yield ch + + out: List[str] = [] + async for ch in proxy_logging.async_post_call_streaming_iterator_hook( + response=gen(), + user_api_key_dict=make_user_api_key_auth(), + request_data={}, + ): + out.append(ch) + assert out == ["a*", "b*"] + + +@pytest.mark.asyncio +async def test_async_post_call_streaming_iterator_hook_upstream_error_raises(proxy_logging, make_user_api_key_auth, mock_callbacks_disabled): + async def gen(): + if False: + yield # pragma: no cover + raise RuntimeError("upstream") + + with pytest.raises(RuntimeError): + async for _ in proxy_logging.async_post_call_streaming_iterator_hook( + response=gen(), + user_api_key_dict=make_user_api_key_auth(), + request_data={}, + ): + pass + + +# --------------------------------------------------------------------------- +# _fire_deferred_stream_logging +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_fire_deferred_stream_logging_fires_callback(): + logging_obj = MagicMock() + captured: Dict[str, Any] = {} + + async def deferred(arg): + captured["arg"] = arg + + logging_obj._on_deferred_stream_complete = deferred + logging_obj._deferred_stream_complete_args = ("payload",) + + ProxyLogging._fire_deferred_stream_logging(request_data={"litellm_logging_obj": logging_obj}) + await asyncio.sleep(0) + snapshot = { + "arg": captured["arg"], + "callback_cleared": logging_obj._on_deferred_stream_complete is None, + "args_cleared": logging_obj._deferred_stream_complete_args is None, + } + assert snapshot == {"arg": "payload", "callback_cleared": True, "args_cleared": True} + + +def test_fire_deferred_stream_logging_no_logging_obj_no_error(): + ProxyLogging._fire_deferred_stream_logging(request_data={}) + + +def test_fire_deferred_stream_logging_missing_obj_raises_on_invalid_dict(): + with pytest.raises(AttributeError): + ProxyLogging._fire_deferred_stream_logging(request_data=None) # type: ignore[arg-type] + + +# --------------------------------------------------------------------------- +# post_call_response_headers_hook +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_post_call_response_headers_hook_returns_empty_when_no_callbacks( + proxy_logging, mock_callbacks_disabled, make_user_api_key_auth +): + out = await proxy_logging.post_call_response_headers_hook( + data={}, user_api_key_dict=make_user_api_key_auth(), response=MagicMock(_hidden_params={}) + ) + assert out == {} + + +@pytest.mark.asyncio +async def test_post_call_response_headers_hook_merges_callback_headers(proxy_logging, make_user_api_key_auth, monkeypatch): + class _Cb(CustomLogger): + async def async_post_call_response_headers_hook(self, **kwargs): # type: ignore[override] + return {"X-One": "1", "X-Two": "2", "X-Common": "first"} + + class _Cb2(CustomLogger): + async def async_post_call_response_headers_hook(self, **kwargs): # type: ignore[override] + return {"X-Common": "second", "X-Three": "3"} + + monkeypatch.setattr(litellm, "callbacks", [_Cb(), _Cb2()]) + response = MagicMock() + response._hidden_params = {} + out = await proxy_logging.post_call_response_headers_hook( + data={}, user_api_key_dict=make_user_api_key_auth(), response=response + ) + assert out == {"X-One": "1", "X-Two": "2", "X-Common": "second", "X-Three": "3"} + + +@pytest.mark.asyncio +async def test_post_call_response_headers_hook_swallows_callback_error(proxy_logging, make_user_api_key_auth, monkeypatch): + """Errors inside the hook are caught — function returns merged so-far.""" + + class _Cb(CustomLogger): + async def async_post_call_response_headers_hook(self, **kwargs): # type: ignore[override] + raise RuntimeError("bad header") + + monkeypatch.setattr(litellm, "callbacks", [_Cb()]) + response = MagicMock() + response._hidden_params = {} + out = await proxy_logging.post_call_response_headers_hook( + data={}, user_api_key_dict=make_user_api_key_auth(), response=response + ) + assert out == {} diff --git a/tests/test_litellm/test_bedrock_anthropic_1hr_cache_pricing.py b/tests/test_litellm/test_bedrock_anthropic_1hr_cache_pricing.py index 69af35dfeae..983f60b0339 100644 --- a/tests/test_litellm/test_bedrock_anthropic_1hr_cache_pricing.py +++ b/tests/test_litellm/test_bedrock_anthropic_1hr_cache_pricing.py @@ -72,9 +72,40 @@ US_EXPECTED = [ ("us.anthropic.claude-haiku-4-5-20251001-v1:0", 2.2e-06, None), ] +# EU/AU/JP cross-region inference profiles carry the same +10% regional +# premium as US (per AWS Bedrock pricing). Coverage list filters to entries +# that actually exist in the pricing JSON - e.g. Opus 4.6 has no JP profile. +REGIONAL_EXPECTED = [ + # Opus 4.6 - $11.00 / MTok (eu/au only; no jp profile) + ("eu.anthropic.claude-opus-4-6-v1", 1.1e-05, None), + ("au.anthropic.claude-opus-4-6-v1", 1.1e-05, None), + # Opus 4.7 - $11.00 / MTok (eu/au; jp is added in #28567) + ("eu.anthropic.claude-opus-4-7", 1.1e-05, None), + ("au.anthropic.claude-opus-4-7", 1.1e-05, None), + # Sonnet 4.6 - $6.60 / MTok + ("eu.anthropic.claude-sonnet-4-6", 6.6e-06, None), + ("au.anthropic.claude-sonnet-4-6", 6.6e-06, None), + ("jp.anthropic.claude-sonnet-4-6", 6.6e-06, None), + # Sonnet 4.5 - $6.60 / MTok with $13.20 / MTok long-context tier + ("eu.anthropic.claude-sonnet-4-5-20250929-v1:0", 6.6e-06, 1.32e-05), + ("au.anthropic.claude-sonnet-4-5-20250929-v1:0", 6.6e-06, 1.32e-05), + ("jp.anthropic.claude-sonnet-4-5-20250929-v1:0", 6.6e-06, 1.32e-05), + # Haiku 4.5 - $2.20 / MTok + ("eu.anthropic.claude-haiku-4-5-20251001-v1:0", 2.2e-06, None), + ("au.anthropic.claude-haiku-4-5-20251001-v1:0", 2.2e-06, None), + ("jp.anthropic.claude-haiku-4-5-20251001-v1:0", 2.2e-06, None), + # Note: eu.anthropic.claude-opus-4-5-20251101-v1:0 is intentionally NOT + # in this list. The existing entry carries base/global 5m rates + # (5e-06 / 6.25e-06) instead of the +10% regional premium (5.5e-06 / + # 6.875e-06), which would make the 1.6x 5m-to-1h invariant fail. + # Fixing the EU 5m rates first is left to a follow-up so this PR + # stays scoped to the 1-hour cache tier addition. +] + @pytest.mark.parametrize( - "model_key, expected_1hr, expected_1hr_lc", GLOBAL_EXPECTED + US_EXPECTED + "model_key, expected_1hr, expected_1hr_lc", + GLOBAL_EXPECTED + US_EXPECTED + REGIONAL_EXPECTED, ) def test_bedrock_anthropic_1hr_cache_write_pricing( model_data, model_key, expected_1hr, expected_1hr_lc diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index d3935523a32..92d2de55116 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -2165,11 +2165,11 @@ def test_gemini_3_1_flash_lite_pricing(): ): model_info = litellm.model_cost.get(model_name) assert model_info is not None, f"Missing model pricing entry: {model_name}" - assert model_info["input_cost_per_token"] == 4.5e-07 - assert model_info["input_cost_per_audio_token"] == 9e-07 - assert model_info["output_cost_per_token"] == 2.7e-06 - assert model_info["output_cost_per_reasoning_token"] == 2.7e-06 - assert model_info["cache_read_input_token_cost"] == 4.5e-08 + assert model_info["input_cost_per_token"] == 2.5e-07 + assert model_info["input_cost_per_audio_token"] == 5e-07 + assert model_info["output_cost_per_token"] == 1.5e-06 + assert model_info["output_cost_per_reasoning_token"] == 1.5e-06 + assert model_info["cache_read_input_token_cost"] == 2.5e-08 assert model_info["max_input_tokens"] == 1048576 diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 5e636b86ed6..e9287a95438 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -2376,6 +2376,74 @@ def test_get_deployment_model_info_base_model_flow(): # Should return None when no model info is found assert result is None + # Test Case 6: custom_model_info present but litellm_model_name_model_info is None + # (model has custom pricing in config but is not in built-in model_prices_and_context_window.json) + mock_custom_pricing_only = { + "input_cost_per_token": 1.74e-06, + "output_cost_per_token": 3.48e-06, + "cache_read_input_token_cost": 1.45e-08, + "mode": "chat", + } + + with patch.object( + litellm, + "model_cost", + {"custom-model-id": mock_custom_pricing_only}, + ): + with patch.object(litellm, "get_model_info") as mock_get_model_info: + # Model NOT in built-in cost map — raise exception + mock_get_model_info.side_effect = Exception("Model not in cost map") + + result = router.get_deployment_model_info( + model_id="custom-model-id", model_name="unknown-model" + ) + + # Should return custom_model_info even when litellm_model_name_model_info is None + assert result is not None + assert result["input_cost_per_token"] == 1.74e-06 + assert result["output_cost_per_token"] == 3.48e-06 + assert result["cache_read_input_token_cost"] == 1.45e-08 + assert result["mode"] == "chat" + + # Test Case 7: custom_model_info with base_model but litellm_model_name_model_info None + mock_custom_with_base = { + "base_model": "some-base-model", + "input_cost_per_token": 0.01, + "output_cost_per_token": 0.02, + } + mock_base_info = { + "key": "some-base-model", + "max_tokens": 8192, + "mode": "chat", + "litellm_provider": "openai", + } + + with patch.object( + litellm, + "model_cost", + {"custom-with-base": mock_custom_with_base}, + ): + with patch.object(litellm, "get_model_info") as mock_get_model_info: + + def get_info_side_effect(model): + if model == "some-base-model": + return mock_base_info + raise Exception("Model not in cost map") + + mock_get_model_info.side_effect = get_info_side_effect + + result = router.get_deployment_model_info( + model_id="custom-with-base", model_name="unknown-model" + ) + + # Should return custom_model_info merged with base model info + assert result is not None + assert ( + result["input_cost_per_token"] == 0.01 + ) # From custom (overrides base) + assert result["max_tokens"] == 8192 # From base model + assert result["litellm_provider"] == "openai" # From base model + print("✓ All base model flow test cases passed!") diff --git a/tests/test_litellm/test_service_logger.py b/tests/test_litellm/test_service_logger.py index 34fe73382b3..de46403b64d 100644 --- a/tests/test_litellm/test_service_logger.py +++ b/tests/test_litellm/test_service_logger.py @@ -99,6 +99,33 @@ async def test_async_log_success_event_should_handle_float_duration(): assert call_kwargs.kwargs["duration"] == 1.5 +@pytest.mark.asyncio +async def test_async_log_success_event_forwards_start_and_end_time(): + """The LITELLM service span must carry its real execution window, so + ``async_log_success_event`` forwards ``start_time``/``end_time`` to the service + hook. Without forwarding, the span emits with a synthetic now() boundary + instead of the call's actual timing.""" + service_logger = ServiceLogging(mock_testing=True) + + start_time = datetime(2026, 2, 13, 22, 35, 0) + end_time = datetime(2026, 2, 13, 22, 35, 1) + + with patch.object( + service_logger, "async_service_success_hook", new_callable=AsyncMock + ) as mock_hook: + await service_logger.async_log_success_event( + kwargs={"call_type": "completion"}, + response_obj=None, + start_time=start_time, + end_time=end_time, + ) + + mock_hook.assert_called_once() + forwarded = mock_hook.call_args.kwargs + assert forwarded["start_time"] == start_time + assert forwarded["end_time"] == end_time + + # --------------------------------------------------------------------------- # # V2 OpenTelemetry service-span dispatch (regression: service spans were always # dropped because the dispatch only recognized the legacy OpenTelemetry class). diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 6a78653ec99..2d75671f1cb 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -928,6 +928,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): }, }, "supports_native_streaming": {"type": "boolean"}, + "supports_image_size": {"type": "boolean"}, "supports_native_structured_output": {"type": "boolean"}, "tiered_pricing": { "type": "array", diff --git a/tests/test_litellm/test_vcr_safe_body_matcher.py b/tests/test_litellm/test_vcr_safe_body_matcher.py index 77b5416a15c..712ecf09911 100644 --- a/tests/test_litellm/test_vcr_safe_body_matcher.py +++ b/tests/test_litellm/test_vcr_safe_body_matcher.py @@ -14,6 +14,7 @@ from tests._vcr_conftest_common import ( # noqa: E402 KEY_FINGERPRINT_HEADER, KEY_FINGERPRINT_MATCHER_NAME, SAFE_BODY_MATCHER_NAME, + TOLERANT_PATH_MATCHER_NAME, TOLERANT_QUERY_MATCHER_NAME, _before_record_request, _is_credential_exchange_request, @@ -21,6 +22,7 @@ from tests._vcr_conftest_common import ( # noqa: E402 _key_fingerprint_matcher, _normalize_volatile_tokens, _safe_body_matcher, + _tolerant_path_matcher, _tolerant_query_matcher, vcr_config_dict, ) @@ -198,6 +200,20 @@ def test_normalize_volatile_tokens_collapses_uuid_and_timestamps(): assert _normalize_volatile_tokens(e) == _normalize_volatile_tokens(f) +def test_normalize_volatile_tokens_collapses_bedrock_batch_job_names(): + a = ( + b'{"jobName":"litellm-batch-aaaaaaaa",' + b'"outputDataConfig":{"s3OutputDataConfig":' + b'{"s3Uri":"s3://bucket/litellm-batch-outputs/litellm-batch-aaaaaaaa/"}}}' + ) + b = ( + b'{"jobName":"litellm-batch-bbbbbbbb",' + b'"outputDataConfig":{"s3OutputDataConfig":' + b'{"s3Uri":"s3://bucket/litellm-batch-outputs/litellm-batch-bbbbbbbb/"}}}' + ) + assert _normalize_volatile_tokens(a) == _normalize_volatile_tokens(b) + + def test_normalize_volatile_tokens_leaves_deterministic_bodies_unchanged(): body = b'{"model":"claude-haiku-4-5-20251001","temperature":0.0,"n":2}' assert _normalize_volatile_tokens(body) == body @@ -239,6 +255,83 @@ def test_match_on_uses_tolerant_query_not_builtin(): assert "query" not in cfg["match_on"] +def test_match_on_uses_tolerant_path_not_builtin(): + cfg = vcr_config_dict() + assert TOLERANT_PATH_MATCHER_NAME in cfg["match_on"] + assert "path" not in cfg["match_on"] + + +def test_tolerant_path_normalizes_bedrock_managed_s3_file_uuid(): + from vcr.request import Request + + a = Request( + method="PUT", + uri=( + "https://s3.us-west-2.amazonaws.com/litellm-proxy-test/" + "litellm-bedrock-files/us.anthropic.claude-haiku-4-5-20251001-v1-0-" + "123e4567-e89b-12d3-a456-426614174000.jsonl" + ), + body=b"", + headers={}, + ) + b = Request( + method="PUT", + uri=( + "https://s3.us-west-2.amazonaws.com/litellm-proxy-test/" + "litellm-bedrock-files/us.anthropic.claude-haiku-4-5-20251001-v1-0-" + "abcdefab-1234-5678-9abc-def012345678.jsonl" + ), + body=b"", + headers={}, + ) + _tolerant_path_matcher(a, b) + + +def test_tolerant_path_normalizes_bedrock_batch_s3_file_uuid(): + from vcr.request import Request + + a = Request( + method="PUT", + uri=( + "https://s3.us-west-2.amazonaws.com/litellm-proxy-test/" + "litellm-bedrock-files-us.anthropic.claude-haiku-4-5-20251001-v1-0-" + "a48e9ec2-5594-45e3-bdbb-44f5d71c06f3.jsonl" + ), + body=b"", + headers={}, + ) + b = Request( + method="PUT", + uri=( + "https://s3.us-west-2.amazonaws.com/litellm-proxy-test/" + "litellm-bedrock-files-us.anthropic.claude-haiku-4-5-20251001-v1-0-" + "123e4567-e89b-12d3-a456-426614174000.jsonl" + ), + body=b"", + headers={}, + ) + _tolerant_path_matcher(a, b) + + +def test_tolerant_path_still_rejects_different_regular_paths(): + from vcr.request import Request + + a = Request( + method="GET", + uri="https://api.openai.com/v1/files/file-a/content", + body=b"", + headers={}, + ) + b = Request( + method="GET", + uri="https://api.openai.com/v1/files/file-b/content", + body=b"", + headers={}, + ) + with pytest.raises(AssertionError): + _tolerant_path_matcher(a, b) + + def test_telemetry_request_detection(): assert _is_telemetry_request( _req(b"x", uri="https://us.cloud.langfuse.com/api/public/ingestion") diff --git a/ui/litellm-dashboard/public/assets/logos/langflow.svg b/ui/litellm-dashboard/public/assets/logos/langflow.svg new file mode 100644 index 00000000000..1c7b36c4dd6 --- /dev/null +++ b/ui/litellm-dashboard/public/assets/logos/langflow.svg @@ -0,0 +1,5 @@ + + + + + diff --git a/ui/litellm-dashboard/src/app/layout.tsx b/ui/litellm-dashboard/src/app/layout.tsx index f79b7eb7028..a73921ce35b 100644 --- a/ui/litellm-dashboard/src/app/layout.tsx +++ b/ui/litellm-dashboard/src/app/layout.tsx @@ -11,7 +11,7 @@ const inter = Inter({ subsets: ["latin"] }); export const metadata: Metadata = { title: "LiteLLM Dashboard", description: "LiteLLM Proxy Admin UI", - icons: { icon: "./favicon.ico" }, + icons: { icon: "/get_favicon" }, }; export default function RootLayout({ diff --git a/ui/litellm-dashboard/src/components/agents/add_agent_form.tsx b/ui/litellm-dashboard/src/components/agents/add_agent_form.tsx index 3929a9e1832..eee19171fb6 100644 --- a/ui/litellm-dashboard/src/components/agents/add_agent_form.tsx +++ b/ui/litellm-dashboard/src/components/agents/add_agent_form.tsx @@ -197,7 +197,7 @@ const AddAgentForm: React.FC = ({ const handleNext = async () => { try { if (currentStep === 0) { - await form.validateFields(["agent_name"]); + await form.validateFields(); const agentName = form.getFieldValue("agent_name"); if (agentName && !newKeyName) { setNewKeyName(`${agentName}-key`); diff --git a/ui/litellm-dashboard/src/components/agents/agent_config.ts b/ui/litellm-dashboard/src/components/agents/agent_config.ts index e87b191b19c..14b6729bf93 100644 --- a/ui/litellm-dashboard/src/components/agents/agent_config.ts +++ b/ui/litellm-dashboard/src/components/agents/agent_config.ts @@ -212,14 +212,14 @@ export const SKILL_FIELD_CONFIG = { }, tags: { name: "tags", - label: "Tags (comma-separated)", + label: "Tags", required: true, - placeholder: "e.g., hello world, greeting", + placeholder: "Type a tag and press Enter", }, examples: { name: "examples", - label: "Examples (comma-separated)", - placeholder: "e.g., hi, hello world", + label: "Examples", + placeholder: "Type an example and press Enter", }, }; diff --git a/ui/litellm-dashboard/src/components/agents/agent_form_fields.tsx b/ui/litellm-dashboard/src/components/agents/agent_form_fields.tsx index 42e55b8c56f..27b93838ce6 100644 --- a/ui/litellm-dashboard/src/components/agents/agent_form_fields.tsx +++ b/ui/litellm-dashboard/src/components/agents/agent_form_fields.tsx @@ -56,7 +56,7 @@ const AgentFormFields: React.FC = ({ showAgentName = true, {/* Skills */} {shouldShow(AGENT_FORM_CONFIG.skills.key) && ( - + {(fields, { add, remove }) => ( <> @@ -94,20 +94,26 @@ const AgentFormFields: React.FC = ({ showAgentName = true, label={SKILL_FIELD_CONFIG.tags.label} name={[field.name, 'tags']} rules={[{ required: SKILL_FIELD_CONFIG.tags.required, message: 'Required' }]} - getValueFromEvent={(e) => e.target.value.split(',').map((s: string) => s.trim())} - getValueProps={(value) => ({ value: Array.isArray(value) ? value.join(', ') : value })} > - + + + + {children} + + ); + }; + return render( + + + , + ); + }; + + it("shows only the oauth2 PKCE-delegation toggle for oauth2 servers", async () => { + renderWithInitialValues({ allow_all_keys: false, auth_type: "oauth2" }); + await expandPanel(); + expect( + screen.getByText("Delegate auth to upstream (PKCE passthrough)"), + ).toBeInTheDocument(); + // The non-oauth2 pass-through toggle must NOT appear for oauth2 servers. + expect(screen.queryByText("OAuth pass-through")).not.toBeInTheDocument(); + }); + + it("shows only the OAuth pass-through toggle for none-auth servers forwarding Authorization", async () => { + renderWithInitialValues({ + allow_all_keys: false, + auth_type: "none", + extra_headers: ["Authorization"], + }); + await expandPanel(); + expect(screen.getByText("OAuth pass-through")).toBeInTheDocument(); + // The oauth2-only PKCE delegation toggle must NOT appear here. + expect( + screen.queryByText("Delegate auth to upstream (PKCE passthrough)"), + ).not.toBeInTheDocument(); + }); + + it("hides both upstream-auth toggles for none-auth servers without an Authorization header", async () => { + renderWithInitialValues({ + allow_all_keys: false, + auth_type: "none", + extra_headers: ["x-api-key"], + }); + await expandPanel(); + expect(screen.queryByText("OAuth pass-through")).not.toBeInTheDocument(); + expect( + screen.queryByText("Delegate auth to upstream (PKCE passthrough)"), + ).not.toBeInTheDocument(); + }); + it("should reflect allow_all_keys when editing an existing server", async () => { renderWithForm({ mcpServer: { diff --git a/ui/litellm-dashboard/src/components/mcp_tools/MCPPermissionManagement.tsx b/ui/litellm-dashboard/src/components/mcp_tools/MCPPermissionManagement.tsx index 58848df39a0..b5f0fa2e7eb 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/MCPPermissionManagement.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/MCPPermissionManagement.tsx @@ -25,6 +25,21 @@ const MCPPermissionManagement: React.FC = ({ const form = Form.useFormInstance(); const watchedAuthType = Form.useWatch("auth_type", form); const isOAuth2 = watchedAuthType === AUTH_TYPE.OAUTH2; + const isNoneAuth = watchedAuthType === AUTH_TYPE.NONE || watchedAuthType == null; + const watchedExtraHeaders = Form.useWatch("extra_headers", form); + const hasAuthorizationHeader = Array.isArray(watchedExtraHeaders) + && watchedExtraHeaders.some( + (h) => typeof h === "string" && h.toLowerCase() === "authorization", + ); + // Two distinct, independent opt-ins: + // - delegate_auth_to_upstream: oauth2 servers only (PKCE passthrough — + // bypass LiteLLM admission). + // - oauth_passthrough: auth_type=none + Authorization in extra_headers + // (OAuth pass-through: proxy upstream oauth-protected-resource, emit 401 + // challenges, propagate upstream 401/403). + // Kept as separate flags so neither silently implies the other and existing + // oauth2 servers can't regress into pass-through behavior. + const canEnableOAuthPassthrough = isNoneAuth && hasAuthorizationHeader; const watchedDelegateAuth = Form.useWatch("delegate_auth_to_upstream", form); const watchedPublicInternet = Form.useWatch("available_on_public_internet", form); const showInternalDelegatePkceWarning = @@ -51,22 +66,34 @@ const MCPPermissionManagement: React.FC = ({ if (typeof mcpServer.delegate_auth_to_upstream === "boolean") { form.setFieldValue("delegate_auth_to_upstream", mcpServer.delegate_auth_to_upstream); } + if (typeof mcpServer.oauth_passthrough === "boolean") { + form.setFieldValue("oauth_passthrough", mcpServer.oauth_passthrough); + } } else { form.setFieldValue("allow_all_keys", false); form.setFieldValue("available_on_public_internet", true); form.setFieldValue("delegate_auth_to_upstream", false); + form.setFieldValue("oauth_passthrough", false); } }, [mcpServer, form]); - // delegate_auth_to_upstream is only honored server-side when auth_type=oauth2. + // delegate_auth_to_upstream is only honored server-side for oauth2 servers. // Force it back to false whenever the user switches away from oauth2 so a - // stale toggle value doesn't get persisted with another auth type. + // stale toggle value doesn't get persisted unexpectedly. useEffect(() => { if (!isOAuth2) { form.setFieldValue("delegate_auth_to_upstream", false); } }, [isOAuth2, form]); + // oauth_passthrough is only honored for auth_type=none servers that forward + // Authorization upstream. Force it back to false otherwise. + useEffect(() => { + if (!canEnableOAuthPassthrough) { + form.setFieldValue("oauth_passthrough", false); + } + }, [canEnableOAuthPassthrough, form]); + return ( = ({ )} + {canEnableOAuthPassthrough && ( +
+
+ + OAuth pass-through + + + + +

+ Forward upstream OAuth discovery and 401 challenges so clients negotiate OAuth directly with the upstream MCP server. +

+
+ + + +
+ )} + {showInternalDelegatePkceWarning && ( = ({ allow_all_keys: allowAllKeysRaw, available_on_public_internet: availableOnPublicInternetRaw, delegate_auth_to_upstream: delegateAuthToUpstreamRaw, + oauth_passthrough: oauthPassthroughRaw, token_validation_json: rawTokenValidationJson, ...restValues } = values; @@ -399,6 +400,7 @@ const CreateMCPServer: React.FC = ({ allow_all_keys: Boolean(allowAllKeysRaw), available_on_public_internet: Boolean(availableOnPublicInternetRaw), delegate_auth_to_upstream: Boolean(delegateAuthToUpstreamRaw), + oauth_passthrough: Boolean(oauthPassthroughRaw), static_headers: staticHeaders, ...(tokenValidation !== null && { token_validation: tokenValidation }), }; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx index ed3b22a569e..4070a5dd1af 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx @@ -243,6 +243,41 @@ describe("MCPServerEdit (delegate auth)", () => { expect(payload.auth_type).toBe("none"); expect(payload.delegate_auth_to_upstream).toBe(false); }); + + it("does not enable oauth_passthrough for an oauth2 server", async () => { + vi.mocked(networking.updateMCPServer).mockResolvedValue({ + ...interactiveOAuthServer, + oauth_passthrough: false, + }); + + render( + , + ); + + const saveButtons = screen.getAllByRole("button", { name: "Save Changes" }); + await act(async () => { + fireEvent.click(saveButtons[0]); + }); + + await waitFor(() => { + expect(networking.updateMCPServer).toHaveBeenCalledTimes(1); + }); + + const [, payload] = vi.mocked(networking.updateMCPServer).mock.calls[0]; + expect(payload.auth_type).toBe("oauth2"); + // oauth_passthrough is non-oauth2 only — must be forced false here. + expect(payload.oauth_passthrough).toBe(false); + }); }); describe("MCPServerEdit (tool allowlist)", () => { diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx index ab9c9ed6689..222de54f3cc 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx @@ -396,6 +396,7 @@ const MCPServerEdit: React.FC = ({ allow_all_keys: allowAllKeysRaw, available_on_public_internet: availableOnPublicInternetRaw, delegate_auth_to_upstream: delegateAuthToUpstreamRaw, + oauth_passthrough: oauthPassthroughRaw, token_validation_json: rawTokenValidationJson, ...restValues } = values; @@ -573,14 +574,34 @@ const MCPServerEdit: React.FC = ({ allow_all_keys: Boolean(allowAllKeysRaw ?? mcpServer.allow_all_keys), available_on_public_internet: Boolean(availableOnPublicInternetRaw ?? mcpServer.available_on_public_internet), // ``delegate_auth_to_upstream`` is only honored server-side for - // ``auth_type=oauth2``. The Form.Item is conditionally rendered so the - // value drops out of the form on auth_type change; force false for any - // non-oauth2 server to avoid persisting a stale ``true`` that would - // silently re-activate if auth_type is later switched back to oauth2. - delegate_auth_to_upstream: - restValues.auth_type === AUTH_TYPE.OAUTH2 + // ``auth_type=oauth2`` (PKCE passthrough). The Form.Item is + // conditionally rendered so the value drops out of the form on + // auth_type change; force false for any other configuration to avoid + // persisting a stale ``true`` that would silently re-activate if the + // configuration is later switched back. + delegate_auth_to_upstream: (() => { + const isOauth2 = restValues.auth_type === AUTH_TYPE.OAUTH2; + return isOauth2 ? Boolean(delegateAuthToUpstreamRaw ?? mcpServer.delegate_auth_to_upstream) - : false, + : false; + })(), + // ``oauth_passthrough`` is the dedicated, non-oauth2 opt-in. It is only + // honored for ``auth_type=none`` servers that forward ``Authorization`` + // upstream. Kept separate from ``delegate_auth_to_upstream`` so enabling + // pass-through never regresses oauth2 servers. Force false otherwise. + oauth_passthrough: (() => { + const isNoneAuth = + restValues.auth_type === AUTH_TYPE.NONE || restValues.auth_type == null; + const extraHeaders = Array.isArray(restValues.extra_headers) + ? restValues.extra_headers + : []; + const hasAuthorizationHeader = extraHeaders.some( + (h: unknown) => typeof h === "string" && h.toLowerCase() === "authorization", + ); + return isNoneAuth && hasAuthorizationHeader + ? Boolean(oauthPassthroughRaw ?? mcpServer.oauth_passthrough) + : false; + })(), // Include token_validation when it is set (non-null) or when clearing an existing value ...(tokenValidation !== null || mcpServer.token_validation ? { token_validation: tokenValidation } : {}), }; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx index 5a8035d4e0b..87e8b77837c 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx @@ -290,6 +290,27 @@ export const MCPServerView: React.FC = ({ )} + {handleAuth(mcpServer.auth_type) !== "oauth2" && + Array.isArray(mcpServer.extra_headers) && + mcpServer.extra_headers.some( + (h) => typeof h === "string" && h.toLowerCase() === "authorization", + ) && ( +
+ OAuth Pass-through +
+ {mcpServer.oauth_passthrough ? ( + + + Enabled + + ) : ( + + Disabled + + )} +
+
+ )}
Access Groups
diff --git a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx index 511f20baef5..3fa27afe1b2 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx @@ -211,6 +211,7 @@ export interface MCPServer { allow_all_keys?: boolean; available_on_public_internet?: boolean; delegate_auth_to_upstream?: boolean; + oauth_passthrough?: boolean; /** Stdio-only fields (present when transport === 'stdio') */ command?: string | null;