diff --git a/docs/my-website/docs/proxy/config_settings.md b/docs/my-website/docs/proxy/config_settings.md index 18b27edfa71..3a95d22ea3f 100644 --- a/docs/my-website/docs/proxy/config_settings.md +++ b/docs/my-website/docs/proxy/config_settings.md @@ -487,6 +487,7 @@ router_settings: | AZURE_STORAGE_CLIENT_ID | The Application Client ID to use for Authentication to Azure Blob Storage logging | AZURE_STORAGE_CLIENT_SECRET | The Application Client Secret to use for Authentication to Azure Blob Storage logging | AZURE_VECTOR_STORE_COST_PER_GB_PER_DAY | Cost per GB per day for Azure Vector Store service +| BACKGROUND_HEALTH_CHECK_MAX_TOKENS | Optional global default for `max_tokens` on proxy background health checks when a model has no `health_check_max_tokens`. If unset, non-wildcard models default to 1. Applies to wildcard routes when set. Default is unset | BATCH_STATUS_POLL_INTERVAL_SECONDS | Interval in seconds for polling batch status. Default is 3600 (1 hour) | BATCH_STATUS_POLL_MAX_ATTEMPTS | Maximum number of attempts for polling batch status. Default is 24 (for 24 hours) | BEDROCK_MAX_POLICY_SIZE | Maximum size for Bedrock policy. Default is 75 diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260414140000_add_mcp_server_instructions/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260414140000_add_mcp_server_instructions/migration.sql new file mode 100644 index 00000000000..531024c519f --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260414140000_add_mcp_server_instructions/migration.sql @@ -0,0 +1,2 @@ +-- AlterTable +ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "instructions" TEXT; diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 1b3db52eaa7..ce3f5f131f7 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -289,6 +289,7 @@ model LiteLLM_MCPServerTable { server_name String? alias String? description String? + instructions String? url String? spec_path String? transport String @default("sse") diff --git a/litellm/constants.py b/litellm/constants.py index 1a547740ad0..348c405506b 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1,6 +1,6 @@ import os import sys -from typing import List, Literal +from typing import List, Literal, Optional from litellm.litellm_core_utils.env_utils import get_env_int @@ -1343,6 +1343,22 @@ BATCH_STATUS_POLL_MAX_ATTEMPTS = int( HEALTH_CHECK_TIMEOUT_SECONDS = int( os.getenv("HEALTH_CHECK_TIMEOUT_SECONDS", 60) ) # 60 seconds +_background_health_check_max_tokens_env = os.getenv( + "BACKGROUND_HEALTH_CHECK_MAX_TOKENS" +) +try: + _raw_background_health_check_max_tokens = ( + _background_health_check_max_tokens_env.strip() + if _background_health_check_max_tokens_env is not None + else "" + ) + BACKGROUND_HEALTH_CHECK_MAX_TOKENS: Optional[int] = ( + int(_raw_background_health_check_max_tokens) + if _raw_background_health_check_max_tokens + else None + ) +except (ValueError, TypeError): + BACKGROUND_HEALTH_CHECK_MAX_TOKENS = None LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME = "litellm-internal-health-check" LITTELM_CLI_SERVICE_ACCOUNT_NAME = "litellm-cli" LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME = "litellm_internal_jobs" diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 1423617cac0..e703a3956b9 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -221,6 +221,7 @@ class MCPClient: self.extra_headers: Optional[Dict[str, str]] = extra_headers self.ssl_verify: Optional[VerifyTypes] = ssl_verify self._aws_auth: Optional[httpx.Auth] = aws_auth + self._last_initialize_instructions: Optional[str] = None # handle the basic auth value if provided if auth_value: self.update_auth_value(auth_value) @@ -296,7 +297,12 @@ class MCPClient: session_ctx = ClientSession(read_stream, write_stream) session = await session_ctx.__aenter__() try: - await session.initialize() + init_result = await session.initialize() + self._last_initialize_instructions = None + if init_result is not None: + ins = getattr(init_result, "instructions", None) + if isinstance(ins, str) and ins.strip(): + self._last_initialize_instructions = ins.strip() return await operation(session) finally: try: @@ -315,6 +321,7 @@ class MCPClient: """Open a session, run the provided coroutine, and clean up.""" http_client: Optional[httpx.AsyncClient] = None try: + self._last_initialize_instructions = None transport_ctx, http_client = self._create_transport_context() return await self._execute_session_operation(transport_ctx, operation) except Exception: diff --git a/litellm/litellm_core_utils/model_param_helper.py b/litellm/litellm_core_utils/model_param_helper.py index 66b174feac4..b4fa5cb60aa 100644 --- a/litellm/litellm_core_utils/model_param_helper.py +++ b/litellm/litellm_core_utils/model_param_helper.py @@ -11,6 +11,10 @@ from openai.types.completion_create_params import ( CompletionCreateParamsStreaming as TextCompletionCreateParamsStreaming, ) from openai.types.embedding_create_params import EmbeddingCreateParams +from openai.types.responses.response_create_params import ( + ResponseCreateParamsNonStreaming, + ResponseCreateParamsStreaming, +) from litellm._logging import verbose_logger from litellm.types.rerank import RerankRequest @@ -65,6 +69,9 @@ class ModelParamHelper: ModelParamHelper._get_litellm_supported_transcription_kwargs() ) rerank_kwargs = ModelParamHelper._get_litellm_supported_rerank_kwargs() + responses_api_kwargs = ( + ModelParamHelper._get_litellm_supported_responses_api_kwargs() + ) exclude_kwargs = ModelParamHelper._get_exclude_kwargs() combined_kwargs = chat_completion_kwargs.union( @@ -72,6 +79,7 @@ class ModelParamHelper: embedding_kwargs, transcription_kwargs, rerank_kwargs, + responses_api_kwargs, ) combined_kwargs = combined_kwargs.difference(exclude_kwargs) return combined_kwargs @@ -93,9 +101,9 @@ class ModelParamHelper: streaming_params: Set[str] = set( getattr(CompletionCreateParamsStreaming, "__annotations__", {}).keys() ) - litellm_provider_specific_params: Set[ - str - ] = ModelParamHelper.get_litellm_provider_specific_params_for_chat_params() + litellm_provider_specific_params: Set[str] = ( + ModelParamHelper.get_litellm_provider_specific_params_for_chat_params() + ) all_chat_completion_kwargs: Set[str] = non_streaming_params.union( streaming_params ).union(litellm_provider_specific_params) @@ -167,6 +175,21 @@ class ModelParamHelper: verbose_logger.debug("Error getting transcription kwargs %s", str(e)) return set() + @staticmethod + def _get_litellm_supported_responses_api_kwargs() -> Set[str]: + """ + Get the litellm supported responses API kwargs + + This follows the OpenAI API Spec + """ + non_streaming_params: Set[str] = set( + getattr(ResponseCreateParamsNonStreaming, "__annotations__", {}).keys() + ) + streaming_params: Set[str] = set( + getattr(ResponseCreateParamsStreaming, "__annotations__", {}).keys() + ) + return non_streaming_params.union(streaming_params) + @staticmethod def _get_exclude_kwargs() -> Set[str]: """ diff --git a/litellm/llms/anthropic/common_utils.py b/litellm/llms/anthropic/common_utils.py index a0da14bcc2b..3be5a0c816f 100644 --- a/litellm/llms/anthropic/common_utils.py +++ b/litellm/llms/anthropic/common_utils.py @@ -2,6 +2,7 @@ This file contains common utils for anthropic calls. """ +import copy from typing import Any, Dict, List, Optional, Union import httpx @@ -736,6 +737,69 @@ def strip_advisor_blocks_from_messages( return messages +def is_anthropic_invalid_thinking_signature_error(error_text: str) -> bool: + """ + Detect Anthropic 400 when encrypted thinking signatures in history do not match + the current deployment (e.g. user rotated API key or switched model endpoint). + + Example API message: + messages.N.content.M: Invalid `signature` in `thinking` block + """ + if not error_text: + return False + lower = error_text.lower() + return ( + "invalid" in lower + and "signature" in lower + and "thinking" in lower + and "block" in lower + ) + + +def strip_thinking_blocks_from_anthropic_messages(messages: List[Any]) -> List[Any]: + """ + Return a new message list with thinking / redacted_thinking content blocks removed + from each message. Used to recover from invalid thinking signatures on retry. + + Messages whose content is a list and becomes empty after stripping are omitted, + since Anthropic rejects empty content arrays. + """ + out: List[Any] = [] + for m in messages: + if not isinstance(m, dict): + out.append(m) + continue + mm = copy.deepcopy(m) + content = mm.get("content") + if isinstance(content, list): + filtered = [ + b + for b in content + if not ( + isinstance(b, dict) + and b.get("type") in ("thinking", "redacted_thinking") + ) + ] + if not filtered: + continue + mm["content"] = filtered + out.append(mm) + return out + + +def strip_thinking_blocks_from_anthropic_messages_request_dict( + data: Dict[str, Any], +) -> None: + """ + Mutate an Anthropic Messages-style request dict: strip thinking blocks from + ``messages`` and remove the top-level ``thinking`` extended-thinking param. + """ + msgs = data.get("messages") + if isinstance(msgs, list): + data["messages"] = strip_thinking_blocks_from_anthropic_messages(msgs) + data.pop("thinking", None) + + def process_anthropic_headers(headers: Union[httpx.Headers, dict]) -> dict: openai_headers = {} if "anthropic-ratelimit-requests-limit" in headers: diff --git a/litellm/llms/azure/passthrough/transformation.py b/litellm/llms/azure/passthrough/transformation.py index 4e9de4b314f..9b1d95e5314 100644 --- a/litellm/llms/azure/passthrough/transformation.py +++ b/litellm/llms/azure/passthrough/transformation.py @@ -1,7 +1,9 @@ from typing import TYPE_CHECKING, List, Optional, Tuple import httpx +from httpx import Response +from litellm.litellm_core_utils.litellm_logging import Logging from litellm.llms.azure.common_utils import BaseAzureLLM from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig from litellm.secret_managers.main import get_secret_str @@ -11,6 +13,8 @@ from litellm.types.router import GenericLiteLLMParams if TYPE_CHECKING: from httpx import URL + from litellm.types.utils import CostResponseTypes + class AzurePassthroughConfig(BasePassthroughConfig): def is_streaming_request(self, endpoint: str, request_data: dict) -> bool: @@ -83,3 +87,36 @@ class AzurePassthroughConfig(BasePassthroughConfig): self, api_key: Optional[str] = None, api_base: Optional[str] = None ) -> List[str]: return super().get_models(api_key, api_base) + + def logging_non_streaming_response( + self, + model: str, + custom_llm_provider: str, + httpx_response: Response, + request_data: dict, + logging_obj: Logging, + endpoint: str, + ) -> Optional["CostResponseTypes"]: + from litellm import encoding + from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig + from litellm.types.utils import ModelResponse + + if "chat/completions" not in endpoint: + return None + + openai_chat_config = OpenAIGPTConfig() + + litellm_model_response: ModelResponse = openai_chat_config.transform_response( + model=model, + messages=[{"role": "user", "content": "no-message-pass-through-endpoint"}], + raw_response=httpx_response, + model_response=ModelResponse(), + logging_obj=logging_obj, + optional_params={}, + litellm_params={}, + api_key="", + request_data=request_data, + encoding=encoding, + ) + + return litellm_model_response diff --git a/litellm/llms/base_llm/anthropic_messages/transformation.py b/litellm/llms/base_llm/anthropic_messages/transformation.py index fdad1633e8f..49aa563781f 100644 --- a/litellm/llms/base_llm/anthropic_messages/transformation.py +++ b/litellm/llms/base_llm/anthropic_messages/transformation.py @@ -120,3 +120,46 @@ class BaseAnthropicMessagesConfig(ABC): return BaseLLMException( message=error_message, status_code=status_code, headers=headers ) + + @property + def max_retry_on_anthropic_messages_http_error(self) -> int: + """ + Max HTTP attempts for /v1/messages when the handler may mutate the body and + retry (e.g. strip invalid encrypted thinking signatures after a deployment or + credential change). + """ + return 2 + + def should_retry_anthropic_messages_on_http_error( + self, e: httpx.HTTPStatusError, litellm_params: dict + ) -> bool: + """ + When True, async_anthropic_messages_handler will transform the request body + and issue one more attempt (bounded by max_retry_on_anthropic_messages_http_error). + """ + from litellm.llms.anthropic.common_utils import ( + is_anthropic_invalid_thinking_signature_error, + ) + + return ( + e.response.status_code == 400 + and is_anthropic_invalid_thinking_signature_error(e.response.text) + ) + + def transform_anthropic_messages_request_on_http_error( + self, e: httpx.HTTPStatusError, request_data: dict + ) -> dict: + """ + Mutates request_data in place when retrying after a recoverable HTTP error. + """ + from litellm.llms.anthropic.common_utils import ( + is_anthropic_invalid_thinking_signature_error, + strip_thinking_blocks_from_anthropic_messages_request_dict, + ) + + if ( + e.response.status_code == 400 + and is_anthropic_invalid_thinking_signature_error(e.response.text) + ): + strip_thinking_blocks_from_anthropic_messages_request_dict(request_data) + return request_data diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 7a8820a8785..ba76e13b6a8 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -1816,6 +1816,73 @@ class BaseLLMHTTPHandler: logging_obj=logging_obj, ) + async def _async_post_anthropic_messages_with_http_error_retry( + self, + async_httpx_client: AsyncHTTPHandler, + request_url: str, + headers: dict, + signed_json_body: Optional[bytes], + request_body: dict, + stream: bool, + logging_obj: LiteLLMLoggingObj, + provider_config: BaseAnthropicMessagesConfig, + litellm_params: GenericLiteLLMParams, + api_key: Optional[str], + model: str, + ) -> httpx.Response: + max_attempts = max( + provider_config.max_retry_on_anthropic_messages_http_error, 1 + ) + litellm_params_dict = dict(litellm_params) + optional_params_dict = dict(litellm_params) + for attempt_idx in range(max_attempts): + try: + response = await async_httpx_client.post( + url=request_url, + headers=headers, + data=signed_json_body or json.dumps(request_body), + stream=stream or False, + logging_obj=logging_obj, + ) + response.raise_for_status() + return response + except httpx.HTTPStatusError as e: + hit_max_attempt = attempt_idx + 1 == max_attempts + should_retry = ( + provider_config.should_retry_anthropic_messages_on_http_error( + e=e, litellm_params=litellm_params_dict + ) + ) + if should_retry and not hit_max_attempt: + verbose_logger.debug( + "Anthropic /v1/messages: invalid thinking signature; " + "stripping thinking blocks and retrying (attempt %s/%s).", + attempt_idx + 2, + max_attempts, + ) + provider_config.transform_anthropic_messages_request_on_http_error( + e=e, request_data=request_body + ) + headers, signed_json_body = provider_config.sign_request( + headers=headers, + optional_params=optional_params_dict, + request_data=request_body, + api_base=request_url, + api_key=api_key, + stream=stream, + fake_stream=False, + model=model, + ) + logging_obj.model_call_details.update(request_body) + continue + raise self._handle_error(e=e, provider_config=provider_config) + except Exception as e: + raise self._handle_error(e=e, provider_config=provider_config) + + raise RuntimeError( + "unreachable: anthropic messages HTTP retry loop exited without return" + ) + async def async_anthropic_messages_handler( self, model: str, @@ -1955,19 +2022,19 @@ class BaseLLMHTTPHandler: }, ) - try: - response = await async_httpx_client.post( - url=request_url, - headers=headers, - data=signed_json_body or json.dumps(request_body), - stream=stream or False, - logging_obj=logging_obj, - ) - response.raise_for_status() - except Exception as e: - raise self._handle_error( - e=e, provider_config=anthropic_messages_provider_config - ) + response = await self._async_post_anthropic_messages_with_http_error_retry( + async_httpx_client=async_httpx_client, + request_url=request_url, + headers=headers, + signed_json_body=signed_json_body, + request_body=request_body, + stream=stream or False, + logging_obj=logging_obj, + provider_config=anthropic_messages_provider_config, + litellm_params=litellm_params, + api_key=api_key, + model=model, + ) # used for logging + cost tracking logging_obj.model_call_details["httpx_response"] = response @@ -4496,9 +4563,9 @@ class BaseLLMHTTPHandler: # Second: Execute agentic loop # Add custom_llm_provider to kwargs so the agentic loop can reconstruct the full model name kwargs_with_provider = kwargs.copy() if kwargs else {} - kwargs_with_provider[ - "custom_llm_provider" - ] = custom_llm_provider + kwargs_with_provider["custom_llm_provider"] = ( + custom_llm_provider + ) agentic_response = await callback.async_run_agentic_loop( tools=tool_calls, model=model, @@ -4614,9 +4681,9 @@ class BaseLLMHTTPHandler: # Second: Execute agentic loop # Add custom_llm_provider to kwargs so the agentic loop can reconstruct the full model name kwargs_with_provider = kwargs.copy() if kwargs else {} - kwargs_with_provider[ - "custom_llm_provider" - ] = custom_llm_provider + kwargs_with_provider["custom_llm_provider"] = ( + custom_llm_provider + ) agentic_response = ( await callback.async_run_chat_completion_agentic_loop( tools=tool_calls, @@ -5110,7 +5177,10 @@ class BaseLLMHTTPHandler: _is_async: bool = False, fake_stream: bool = False, litellm_metadata: Optional[Dict[str, Any]] = None, - ) -> Union[ImageResponse, Coroutine[Any, Any, ImageResponse],]: + ) -> Union[ + ImageResponse, + Coroutine[Any, Any, ImageResponse], + ]: """ Handles image edit requests. @@ -5322,7 +5392,10 @@ class BaseLLMHTTPHandler: fake_stream: bool = False, litellm_metadata: Optional[Dict[str, Any]] = None, api_key: Optional[str] = None, - ) -> Union[ImageResponse, Coroutine[Any, Any, ImageResponse],]: + ) -> Union[ + ImageResponse, + Coroutine[Any, Any, ImageResponse], + ]: """ Handles image generation requests. When _is_async=True, returns a coroutine instead of making the call directly. @@ -5562,7 +5635,10 @@ class BaseLLMHTTPHandler: fake_stream: bool = False, litellm_metadata: Optional[Dict[str, Any]] = None, api_key: Optional[str] = None, - ) -> Union[VideoObject, Coroutine[Any, Any, VideoObject],]: + ) -> Union[ + VideoObject, + Coroutine[Any, Any, VideoObject], + ]: """ Handles video generation requests. When _is_async=True, returns a coroutine instead of making the call directly. diff --git a/litellm/proxy/_experimental/mcp_server/mcp_context.py b/litellm/proxy/_experimental/mcp_server/mcp_context.py index 12830db1d6a..a60138dd340 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_context.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_context.py @@ -14,3 +14,8 @@ from typing import Optional _mcp_active_toolset_id: ContextVar[Optional[str]] = ContextVar( "_mcp_active_toolset_id", default=None ) + +# Per-request merged InitializeResult.instructions; set in MCP HTTP/SSE handlers. +_mcp_gateway_initialize_instructions: ContextVar[Optional[str]] = ContextVar( + "_mcp_gateway_initialize_instructions", default=None +) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index c50bfe3ab1a..8aec30bfc3c 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -184,6 +184,16 @@ class MCPServerManager: "gmail_send_email": "zapier_mcp_server", } """ + self._upstream_initialize_instructions_by_server_id: Dict[str, str] = {} + + def _remember_upstream_initialize_instructions( + self, server: MCPServer, client: MCPClient + ) -> None: + raw = getattr(client, "_last_initialize_instructions", None) + if raw and str(raw).strip(): + self._upstream_initialize_instructions_by_server_id[server.server_id] = str( + raw + ).strip() def get_registry(self) -> Dict[str, MCPServer]: """ @@ -204,6 +214,7 @@ class MCPServerManager: mcp_aliases: Optional dictionary mapping aliases to server names from litellm_settings """ verbose_logger.debug("Loading MCP Servers from config-----") + self._upstream_initialize_instructions_by_server_id.clear() # Track which aliases have been used to ensure only first occurrence is used used_aliases = set() @@ -351,6 +362,7 @@ class MCPServerManager: aws_service_name=server_config.get("aws_service_name", None), aws_role_name=server_config.get("aws_role_name", None), aws_session_name=server_config.get("aws_session_name", None), + instructions=server_config.get("instructions", None), ) self.config_mcp_servers[server_id] = new_server @@ -693,6 +705,7 @@ class MCPServerManager: aws_service_name=aws_creds.get("aws_service_name"), aws_role_name=aws_creds.get("aws_role_name"), aws_session_name=aws_creds.get("aws_session_name"), + instructions=mcp_server.instructions, ) return new_server @@ -1247,6 +1260,7 @@ class MCPServerManager: return tools else: tools = await self._fetch_tools_with_timeout(client, server.name) + self._remember_upstream_initialize_instructions(server, client) prefixed_or_original_tools = self._create_prefixed_tools( tools, server, add_prefix=add_prefix @@ -2383,6 +2397,7 @@ class MCPServerManager: # If proxy_logging_obj is not None, the tool call result is at index 1 (after the during hook task) result_index = 1 if proxy_logging_obj else 0 result = mcp_responses[result_index] + self._remember_upstream_initialize_instructions(mcp_server, client) return cast(CallToolResult, result) @@ -2627,6 +2642,7 @@ class MCPServerManager: ) verbose_logger.debug("Loading MCP servers from database into registry...") + self._upstream_initialize_instructions_by_server_id.clear() # perform authz check to filter the mcp servers user has access to prisma_client = get_prisma_client_or_throw( @@ -2910,6 +2926,7 @@ class MCPServerManager: await asyncio.wait_for( client.run_with_session(_noop), timeout=MCP_HEALTH_CHECK_TIMEOUT ) + self._remember_upstream_initialize_instructions(server, client) status = "healthy" except asyncio.TimeoutError: health_check_error = ( @@ -2951,6 +2968,7 @@ class MCPServerManager: token_url=server.token_url, registration_url=server.registration_url, allow_all_keys=server.allow_all_keys, + instructions=server.instructions, ) async def get_all_mcp_servers_with_health_and_teams( @@ -3046,6 +3064,7 @@ class MCPServerManager: is_byok=server.is_byok, byok_description=server.byok_description, byok_api_key_help_url=server.byok_api_key_help_url, + instructions=server.instructions, ) async def get_all_mcp_servers_unfiltered(self) -> List[LiteLLM_MCPServerTable]: diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 32560a2211d..8131c040136 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -933,6 +933,7 @@ if MCP_AVAILABLE: authorization_url=request.authorization_url, registration_url=request.registration_url, oauth2_flow=_oauth2_flow, + instructions=request.instructions, ) stdio_env = global_mcp_server_manager._build_stdio_env( diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 99578d006e1..cefe2ef763b 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -7,6 +7,7 @@ LiteLLM MCP Server Routes import asyncio import contextlib import time +import types import traceback import uuid from datetime import datetime @@ -37,7 +38,10 @@ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( get_request_base_url, ) -from litellm.proxy._experimental.mcp_server.mcp_context import _mcp_active_toolset_id +from litellm.proxy._experimental.mcp_server.mcp_context import ( + _mcp_active_toolset_id, + _mcp_gateway_initialize_instructions, +) from litellm.proxy._experimental.mcp_server.mcp_debug import MCPDebug from litellm.proxy._experimental.mcp_server.utils import ( LITELLM_MCP_SERVER_DESCRIPTION, @@ -122,6 +126,8 @@ _INITIALIZATION_LOCK = asyncio.Lock() if MCP_AVAILABLE: from mcp.server import Server + from mcp.server.lowlevel.server import NotificationOptions + from mcp.server.models import InitializationOptions # Import auth context variables and middleware from mcp.server.auth.middleware.auth_context import ( @@ -200,6 +206,21 @@ if MCP_AVAILABLE: ) return normalized + def _gateway_create_initialization_options( + self, + notification_options: Optional[NotificationOptions] = None, + experimental_capabilities: Optional[Dict[str, Dict[str, Any]]] = None, + ) -> InitializationOptions: + opts = Server.create_initialization_options( + self, + notification_options=notification_options, + experimental_capabilities=experimental_capabilities or {}, + ) + merged = _mcp_gateway_initialize_instructions.get() + if merged is not None: + return opts.model_copy(update={"instructions": merged}) + return opts + ######################################################## ############ Initialize the MCP Server ################# ######################################################## @@ -207,6 +228,9 @@ if MCP_AVAILABLE: name=LITELLM_MCP_SERVER_NAME, version=LITELLM_MCP_SERVER_VERSION, ) + server.create_initialization_options = types.MethodType( # type: ignore[method-assign] + _gateway_create_initialization_options, server + ) sse: SseServerTransport = SseServerTransport("/mcp/sse/messages") # Create session managers @@ -1021,7 +1045,9 @@ if MCP_AVAILABLE: except (ValueError, TypeError): pass ttl = _compute_per_user_token_ttl(server, raw_expires) - await mcp_per_user_token_cache.set(user_id, server_id, access_token, ttl) + await mcp_per_user_token_cache.set( + user_id, server_id, access_token, ttl + ) return {"Authorization": f"Bearer {access_token}"} except Exception as e: @@ -1103,6 +1129,57 @@ if MCP_AVAILABLE: return server_auth_header, extra_headers + def _merge_gateway_initialize_instructions( + allowed_mcp_servers: List[MCPServer], + ) -> Optional[str]: + """YAML/DB override, else in-memory upstream text from list_tools / health_check / call_tool.""" + if not allowed_mcp_servers: + return None + + texts: List[Tuple[str, str]] = [] + for server in allowed_mcp_servers: + label = ( + server.alias + or server.server_name + or server.name + or server.server_id + or "mcp" + ) + if server.instructions and server.instructions.strip(): + texts.append((label, server.instructions.strip())) + continue + if server.spec_path: + continue + cached = global_mcp_server_manager._upstream_initialize_instructions_by_server_id.get( + server.server_id + ) + if cached and cached.strip(): + texts.append((label, cached.strip())) + + if not texts: + return None + if len(texts) == 1: + return texts[0][1] + return "\n\n---\n\n".join(f"[{lbl}]\n{txt}" for lbl, txt in texts) + + @contextlib.asynccontextmanager + async def _gateway_initialize_instructions_request_scope( + user_api_key_auth: Optional[UserAPIKeyAuth], + mcp_servers: Optional[List[str]], + client_ip: Optional[str], + ) -> AsyncIterator[None]: + allowed = await _get_allowed_mcp_servers( + user_api_key_auth=user_api_key_auth, + mcp_servers=mcp_servers, + client_ip=client_ip, + ) + merged = _merge_gateway_initialize_instructions(allowed_mcp_servers=allowed) + tok = _mcp_gateway_initialize_instructions.set(merged) + try: + yield + finally: + _mcp_gateway_initialize_instructions.reset(tok) + async def _get_tools_from_mcp_servers( # noqa: PLR0915 user_api_key_auth: Optional[UserAPIKeyAuth], mcp_auth_header: Optional[str], @@ -2670,7 +2747,12 @@ if MCP_AVAILABLE: # Request was fully handled (e.g., DELETE on non-existent session) return - await session_manager.handle_request(scope, receive, send) + async with _gateway_initialize_instructions_request_scope( + user_api_key_auth, + mcp_servers, + _client_ip, + ): + await session_manager.handle_request(scope, receive, send) except HTTPException: # Re-raise HTTP exceptions to preserve status codes and details raise @@ -2729,7 +2811,12 @@ if MCP_AVAILABLE: await initialize_session_managers() await asyncio.sleep(0.1) - await sse_session_manager.handle_request(scope, receive, send) + async with _gateway_initialize_instructions_request_scope( + user_api_key_auth, + mcp_servers, + _sse_client_ip, + ): + await sse_session_manager.handle_request(scope, receive, send) except Exception as e: verbose_logger.exception(f"Error handling MCP request: {e}") # Instead of re-raising, try to send a graceful error response diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 0bbee56d5e0..7fff640c497 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -904,9 +904,9 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase): allowed_cache_controls: Optional[list] = [] config: Optional[dict] = {} permissions: Optional[dict] = {} - model_max_budget: Optional[ - dict - ] = {} # {"gpt-4": 5.0, "gpt-3.5-turbo": 5.0}, defaults to {} + model_max_budget: Optional[dict] = ( + {} + ) # {"gpt-4": 5.0, "gpt-3.5-turbo": 5.0}, defaults to {} model_config = ConfigDict(protected_namespaces=()) model_rpm_limit: Optional[dict] = None @@ -1048,9 +1048,9 @@ class RegenerateKeyRequest(GenerateKeyRequest): spend: Optional[float] = None metadata: Optional[dict] = None new_master_key: Optional[str] = None - grace_period: Optional[ - str - ] = None # Duration to keep old key valid (e.g. "24h", "2d"); None = immediate revoke + grace_period: Optional[str] = ( + None # Duration to keep old key valid (e.g. "24h", "2d"); None = immediate revoke + ) class ResetSpendRequest(LiteLLMPydanticObjectBase): @@ -1137,6 +1137,7 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase): tool_name_to_description: Optional[Dict[str, str]] = None extra_headers: Optional[List[str]] = None static_headers: Optional[Dict[str, str]] = None + instructions: Optional[str] = None # Stdio-specific fields command: Optional[str] = None args: List[str] = Field(default_factory=list) @@ -1219,6 +1220,7 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase): tool_name_to_description: Optional[Dict[str, str]] = None extra_headers: Optional[List[str]] = None static_headers: Optional[Dict[str, str]] = None + instructions: Optional[str] = None # Stdio-specific fields command: Optional[str] = None args: List[str] = Field(default_factory=list) @@ -1270,6 +1272,7 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase): transport: MCPTransportType auth_type: Optional[MCPAuthType] = None credentials: Optional[MCPCredentials] = None + instructions: Optional[str] = None created_at: Optional[datetime] = None created_by: Optional[str] = None updated_at: Optional[datetime] = None @@ -1574,12 +1577,12 @@ class NewCustomerRequest(BudgetNewRequest): blocked: bool = False # allow/disallow requests for this end-user budget_id: Optional[str] = None # give either a budget_id or max_budget spend: Optional[float] = None - allowed_model_region: Optional[ - AllowedModelRegion - ] = None # require all user requests to use models in this specific region - default_model: Optional[ - str - ] = None # if no equivalent model in allowed region - default all requests to this model + allowed_model_region: Optional[AllowedModelRegion] = ( + None # require all user requests to use models in this specific region + ) + default_model: Optional[str] = ( + None # if no equivalent model in allowed region - default all requests to this model + ) object_permission: Optional[LiteLLM_ObjectPermissionBase] = None @model_validator(mode="before") @@ -1602,12 +1605,12 @@ class UpdateCustomerRequest(LiteLLMPydanticObjectBase): blocked: bool = False # allow/disallow requests for this end-user max_budget: Optional[float] = None budget_id: Optional[str] = None # give either a budget_id or max_budget - allowed_model_region: Optional[ - AllowedModelRegion - ] = None # require all user requests to use models in this specific region - default_model: Optional[ - str - ] = None # if no equivalent model in allowed region - default all requests to this model + allowed_model_region: Optional[AllowedModelRegion] = ( + None # require all user requests to use models in this specific region + ) + default_model: Optional[str] = ( + None # if no equivalent model in allowed region - default all requests to this model + ) object_permission: Optional[LiteLLM_ObjectPermissionBase] = None @@ -1697,15 +1700,15 @@ 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 - team_member_budget: Optional[ - float - ] = None # allow user to set a budget for all team members - team_member_rpm_limit: Optional[ - int - ] = None # allow user to set RPM limit for all team members - team_member_tpm_limit: Optional[ - int - ] = None # allow user to set TPM limit for all team members + team_member_budget: Optional[float] = ( + None # allow user to set a budget for all team members + ) + team_member_rpm_limit: Optional[int] = ( + None # allow user to set RPM limit for all team members + ) + team_member_tpm_limit: Optional[int] = ( + None # allow user to set TPM limit for all team members + ) team_member_key_duration: Optional[str] = None # e.g. "1d", "1w", "1m" team_member_budget_duration: Optional[str] = None # e.g. "30d", "1mo" allowed_vector_store_indexes: Optional[List[AllowedVectorStoreIndexItem]] = None @@ -1802,9 +1805,9 @@ class BlockKeyRequest(LiteLLMPydanticObjectBase): class AddTeamCallback(LiteLLMPydanticObjectBase): callback_name: str - callback_type: Optional[ - Literal["success", "failure", "success_and_failure"] - ] = "success_and_failure" + callback_type: Optional[Literal["success", "failure", "success_and_failure"]] = ( + "success_and_failure" + ) callback_vars: Dict[str, str] @model_validator(mode="before") @@ -2146,9 +2149,9 @@ class ConfigList(LiteLLMPydanticObjectBase): stored_in_db: Optional[bool] field_default_value: Any premium_field: bool = False - nested_fields: Optional[ - List[FieldDetail] - ] = None # For nested dictionary or Pydantic fields + nested_fields: Optional[List[FieldDetail]] = ( + None # For nested dictionary or Pydantic fields + ) class UserHeaderMapping(LiteLLMPydanticObjectBase): @@ -2507,9 +2510,9 @@ class UserAPIKeyAuth( user_max_budget: Optional[float] = None request_route: Optional[str] = None user: Optional[Any] = None # Expanded user object when expand=user is used - created_by_user: Optional[ - Any - ] = None # Expanded created_by user when expand=user is used + created_by_user: Optional[Any] = ( + None # Expanded created_by user when expand=user is used + ) end_user_object_permission: Optional[LiteLLM_ObjectPermissionTable] = None # Decoded upstream IdP claims (groups, roles, etc.) propagated by JWT auth machinery # and forwarded into outbound tokens by guardrails such as MCPJWTSigner. @@ -2648,9 +2651,9 @@ class LiteLLM_OrganizationMembershipTable(LiteLLMPydanticObjectBase): budget_id: Optional[str] = None created_at: datetime updated_at: datetime - user: Optional[ - Any - ] = None # You might want to replace 'Any' with a more specific type if available + user: Optional[Any] = ( + None # You might want to replace 'Any' with a more specific type if available + ) litellm_budget_table: Optional[LiteLLM_BudgetTable] = None user_email: Optional[str] = None @@ -3805,9 +3808,9 @@ class TeamModelDeleteRequest(BaseModel): # Organization Member Requests class OrganizationMemberAddRequest(OrgMemberAddRequest): organization_id: str - max_budget_in_organization: Optional[ - float - ] = None # Users max budget within the organization + max_budget_in_organization: Optional[float] = ( + None # Users max budget within the organization + ) class OrganizationMemberDeleteRequest(MemberDeleteRequest): @@ -4062,9 +4065,9 @@ class ProviderBudgetResponse(LiteLLMPydanticObjectBase): Maps provider names to their budget configs. """ - providers: Dict[ - str, ProviderBudgetResponseObject - ] = {} # Dictionary mapping provider names to their budget configurations + providers: Dict[str, ProviderBudgetResponseObject] = ( + {} + ) # Dictionary mapping provider names to their budget configurations class ProxyStateVariables(TypedDict): @@ -4226,9 +4229,9 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase): enforce_rbac: bool = False roles_jwt_field: Optional[str] = None # v2 on role mappings role_mappings: Optional[List[RoleMapping]] = None - object_id_jwt_field: Optional[ - str - ] = None # can be either user / team, inferred from the role mapping + object_id_jwt_field: Optional[str] = ( + None # can be either user / team, inferred from the role mapping + ) scope_mappings: Optional[List[ScopeMapping]] = None enforce_scope_based_access: bool = False enforce_team_based_model_access: bool = False diff --git a/litellm/proxy/health_check.py b/litellm/proxy/health_check.py index 5d1bcf31f84..1518ed66ab1 100644 --- a/litellm/proxy/health_check.py +++ b/litellm/proxy/health_check.py @@ -11,7 +11,11 @@ from typing import List, Optional import litellm logger = logging.getLogger(__name__) -from litellm.constants import DEFAULT_HEALTH_CHECK_PROMPT, HEALTH_CHECK_TIMEOUT_SECONDS +from litellm.constants import ( + BACKGROUND_HEALTH_CHECK_MAX_TOKENS, + DEFAULT_HEALTH_CHECK_PROMPT, + HEALTH_CHECK_TIMEOUT_SECONDS, +) ILLEGAL_DISPLAY_PARAMS = [ "messages", @@ -242,7 +246,9 @@ async def _perform_health_check( cleaned["model_id"] = _model_id if isinstance(is_healthy, Exception): exceptions_by_model_id[_model_id] = is_healthy - cleaned["exception_status"] = getattr(is_healthy, "status_code", 500) + cleaned["exception_status"] = getattr( + is_healthy, "status_code", 500 + ) unhealthy_endpoints.append(cleaned) return healthy_endpoints, unhealthy_endpoints, exceptions_by_model_id @@ -301,6 +307,8 @@ def _update_litellm_params_for_health_check( _health_check_max_tokens = model_info.get("health_check_max_tokens", None) if _health_check_max_tokens is not None: litellm_params["max_tokens"] = _health_check_max_tokens + elif BACKGROUND_HEALTH_CHECK_MAX_TOKENS is not None: + litellm_params["max_tokens"] = BACKGROUND_HEALTH_CHECK_MAX_TOKENS elif "*" not in ( model_info.get("health_check_model") or litellm_params.get("model") or "" ): diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 1b3db52eaa7..ce3f5f131f7 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -289,6 +289,7 @@ model LiteLLM_MCPServerTable { server_name String? alias String? description String? + instructions String? url String? spec_path String? transport String @default("sse") diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index a7d0968c0ef..ace8c8a4188 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -27,20 +27,21 @@ class MCPServer(BaseModel): spec_path: Optional[str] = None auth_type: Optional[MCPAuthType] = None authentication_token: Optional[str] = None + instructions: Optional[str] = None mcp_info: Optional[MCPInfo] = None - extra_headers: Optional[ - List[str] - ] = None # allow admin to specify which headers to forward from client to the MCP server + extra_headers: Optional[List[str]] = ( + None # allow admin to specify which headers to forward from client to the MCP server + ) allowed_tools: Optional[List[str]] = None disallowed_tools: Optional[List[str]] = None tool_name_to_display_name: Optional[Dict[str, str]] = None tool_name_to_description: Optional[Dict[str, str]] = None - allowed_params: Optional[ - Dict[str, List[str]] - ] = None # map of tool names to allowed parameter lists - static_headers: Optional[ - Dict[str, str] - ] = None # static headers to forward to the MCP server + allowed_params: Optional[Dict[str, List[str]]] = ( + None # map of tool names to allowed parameter lists + ) + static_headers: Optional[Dict[str, str]] = ( + None # static headers to forward to the MCP server + ) # OAuth-specific fields client_id: Optional[str] = None client_secret: Optional[str] = None diff --git a/schema.prisma b/schema.prisma index 1b3db52eaa7..ce3f5f131f7 100644 --- a/schema.prisma +++ b/schema.prisma @@ -289,6 +289,7 @@ model LiteLLM_MCPServerTable { server_name String? alias String? description String? + instructions String? url String? spec_path String? transport String @default("sse") diff --git a/tests/local_testing/test_unit_test_caching.py b/tests/local_testing/test_unit_test_caching.py index fa5cf802546..e25b75e658f 100644 --- a/tests/local_testing/test_unit_test_caching.py +++ b/tests/local_testing/test_unit_test_caching.py @@ -131,6 +131,55 @@ def test_get_cache_key_text_completion(): assert cache_key_2 == cache_key_3 +def test_get_cache_key_responses_api(): + """ + Regression test: two /v1/responses calls that differ only in + `instructions` (or any Responses-API-only param) must produce + different cache keys. Mirrors the chat / embedding / text-completion + cache-key tests above. + """ + cache = Cache() + + base_kwargs = { + "model": "openai/gpt-4.1", + "input": [{"role": "user", "content": "what is the weather"}], + "temperature": 0.3, + } + + kwargs_a = { + **base_kwargs, + "instructions": "summarize the weather on 10th May", + } + kwargs_b = { + **base_kwargs, + "instructions": "summarize the weather on 7th May", + } + + key_a = cache.get_cache_key(**kwargs_a) + key_b = cache.get_cache_key(**kwargs_b) + + assert isinstance(key_a, str) and len(key_a) > 0 + assert key_a != key_b, "instructions must be part of the Responses API cache key" + + # Sanity: identical payloads must still collide (cache hits still work) + key_a_again = cache.get_cache_key(**kwargs_a) + assert key_a == key_a_again + + # Spot-check a handful of other Responses-only params individually. + for param, value_x, value_y in [ + ("previous_response_id", "resp_aaa", "resp_bbb"), + ("reasoning", {"effort": "low"}, {"effort": "high"}), + ("include", ["reasoning.encrypted_content"], []), + ("max_output_tokens", 100, 500), + ("background", True, False), + ]: + kx = {**base_kwargs, param: value_x} + ky = {**base_kwargs, param: value_y} + assert cache.get_cache_key(**kx) != cache.get_cache_key( + **ky + ), f"Responses-API param `{param}` is not part of the cache key" + + def test_get_hashed_cache_key(): cache = Cache() cache_key = "model:gpt-3.5-turbo,messages:Hello world" diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/mcp_tests/test_mcp_server.py index a4a28215e16..6af07585796 100644 --- a/tests/mcp_tests/test_mcp_server.py +++ b/tests/mcp_tests/test_mcp_server.py @@ -417,17 +417,22 @@ async def test_streamable_http_mcp_handler_mock(): # Mock extract_mcp_auth_context to bypass auth checks in the handler mock_auth_context = (None, None, None, {}, {}, {}) - with patch( - "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", - True, - ), patch( - "litellm.proxy._experimental.mcp_server.server.session_manager", - mock_session_manager, - ), patch( - "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", - AsyncMock(return_value=mock_auth_context), - ), patch( - "litellm.proxy._experimental.mcp_server.server.set_auth_context", + with ( + patch( + "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", + True, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.session_manager", + mock_session_manager, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + AsyncMock(return_value=mock_auth_context), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.set_auth_context", + ), ): from litellm.proxy._experimental.mcp_server.server import ( handle_streamable_http_mcp, @@ -471,17 +476,22 @@ async def test_sse_mcp_handler_mock(): [], ) - with patch( - "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", - True, - ), patch( - "litellm.proxy._experimental.mcp_server.server.sse_session_manager", - mock_sse_session_manager, - ), 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", + with ( + patch( + "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", + True, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.sse_session_manager", + mock_sse_session_manager, + ), + 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", + ), ): from litellm.proxy._experimental.mcp_server.server import handle_sse_mcp @@ -833,7 +843,9 @@ async def test_get_tools_from_mcp_servers(): mock_manager.get_allowed_mcp_servers = AsyncMock( return_value=["server1_id", "server2_id"] ) - mock_manager.get_mcp_server_by_id = lambda server_id: mock_server_1 if server_id == "server1_id" else mock_server_2 + mock_manager.get_mcp_server_by_id = lambda server_id: ( + mock_server_1 if server_id == "server1_id" else mock_server_2 + ) mock_manager._get_tools_from_server = AsyncMock(return_value=[mock_tool_1]) # Mock filter_server_ids_by_ip_with_info to return input unchanged (no IP filtering in test) mock_manager.filter_server_ids_by_ip_with_info = MagicMock( @@ -859,7 +871,10 @@ async def test_get_tools_from_mcp_servers(): mock_manager_2.get_allowed_mcp_servers = AsyncMock( return_value=["server1_id", "server2_id"] ) - mock_manager_2.get_mcp_server_by_id = lambda server_id: mock_server_1 if server_id == "server1_id" else mock_server_2 + mock_manager_2.get_mcp_server_by_id = lambda server_id: ( + mock_server_1 if server_id == "server1_id" else mock_server_2 + ) + async def mock_get_tools_side_effect( server, mcp_auth_header=None, @@ -900,7 +915,11 @@ async def test_get_tools_from_mcp_servers(): mock_manager.get_allowed_mcp_servers = AsyncMock( return_value=["server1_id", "server2_id", "server3_id"] ) - mock_manager.get_mcp_server_by_id = lambda server_id: mock_server_1 if server_id == "server1_id" else (mock_server_2 if server_id == "server2_id" else mock_server_3) + mock_manager.get_mcp_server_by_id = lambda server_id: ( + mock_server_1 + if server_id == "server1_id" + else (mock_server_2 if server_id == "server2_id" else mock_server_3) + ) mock_manager._get_tools_from_server = AsyncMock(return_value=[mock_tool_1]) # Mock filter_server_ids_by_ip_with_info to return input unchanged (no IP filtering in test) mock_manager.filter_server_ids_by_ip_with_info = MagicMock( @@ -1050,15 +1069,15 @@ async def test_mcp_server_manager_access_groups_from_config(): # Should find config_server for group-a, both for group-b, other_server for group-c import asyncio - server_ids_a = await MCPRequestHandler._get_mcp_servers_from_access_groups([ - "group-a" - ]) - server_ids_b = await MCPRequestHandler._get_mcp_servers_from_access_groups([ - "group-b" - ]) - server_ids_c = await MCPRequestHandler._get_mcp_servers_from_access_groups([ - "group-c" - ]) + server_ids_a = await MCPRequestHandler._get_mcp_servers_from_access_groups( + ["group-a"] + ) + server_ids_b = await MCPRequestHandler._get_mcp_servers_from_access_groups( + ["group-b"] + ) + server_ids_c = await MCPRequestHandler._get_mcp_servers_from_access_groups( + ["group-c"] + ) assert any(config_server.server_id == sid for sid in server_ids_a) assert set(server_ids_b) == set( [ @@ -1474,6 +1493,7 @@ async def test_add_update_server_with_alias(): mock_mcp_server.byok_api_key_help_url = None mock_mcp_server.created_at = None mock_mcp_server.updated_at = None + mock_mcp_server.instructions = None # Add server to manager await test_manager.add_server(mock_mcp_server) @@ -1530,6 +1550,7 @@ async def test_add_update_server_without_alias(): mock_mcp_server.byok_api_key_help_url = None mock_mcp_server.created_at = None mock_mcp_server.updated_at = None + mock_mcp_server.instructions = None # Add server to manager await test_manager.add_server(mock_mcp_server) @@ -1587,7 +1608,7 @@ async def test_add_update_server_fallback_to_server_id(): mock_mcp_server.byok_api_key_help_url = None mock_mcp_server.created_at = None mock_mcp_server.updated_at = None - + mock_mcp_server.instructions = None # Add server to manager await test_manager.add_server(mock_mcp_server) @@ -2151,8 +2172,12 @@ async def test_list_tool_rest_api_all_servers_with_auth(): for call_args in mock_get_tools.call_args_list } - assert server_auth_map.get(mock_zapier_server) == "Bearer zapier_token" - assert server_auth_map.get(mock_slack_server) == "Bearer slack_token" + assert ( + server_auth_map.get(mock_zapier_server) == "Bearer zapier_token" + ) + assert ( + server_auth_map.get(mock_slack_server) == "Bearer slack_token" + ) @pytest.mark.asyncio @@ -2690,26 +2715,33 @@ async def test_call_mcp_tool_uses_manager_permission_lookup(): expected_response = [TextContent(type="text", text="ok")] - with patch.object( - global_mcp_server_manager, - "get_allowed_mcp_servers", - new_callable=AsyncMock, - ) as mock_get_allowed, patch.object( - global_mcp_server_manager, - "get_mcp_server_by_id", - return_value=mock_server, - ), patch.object( - global_mcp_server_manager, - "_get_mcp_server_from_tool_name", - return_value=mock_server, - ) as mock_get_server, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_tool_registry" - ) as mock_tool_registry, patch( - "litellm.proxy._experimental.mcp_server.server._handle_managed_mcp_tool", - new_callable=AsyncMock, - ) as mock_handle_managed, patch( - "litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed", - return_value=True, + with ( + patch.object( + global_mcp_server_manager, + "get_allowed_mcp_servers", + new_callable=AsyncMock, + ) as mock_get_allowed, + patch.object( + global_mcp_server_manager, + "get_mcp_server_by_id", + return_value=mock_server, + ), + patch.object( + global_mcp_server_manager, + "_get_mcp_server_from_tool_name", + return_value=mock_server, + ) as mock_get_server, + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_tool_registry" + ) as mock_tool_registry, + patch( + "litellm.proxy._experimental.mcp_server.server._handle_managed_mcp_tool", + new_callable=AsyncMock, + ) as mock_handle_managed, + patch( + "litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed", + return_value=True, + ), ): mock_get_allowed.return_value = [mock_server.server_id] mock_tool_registry.get_tool.return_value = None @@ -2759,27 +2791,34 @@ async def test_call_mcp_tool_resolves_unprefixed_tool_name_and_checks_permission expected_response = [TextContent(type="text", text="ok")] - with patch.object( - global_mcp_server_manager, - "get_allowed_mcp_servers", - new_callable=AsyncMock, - ) as mock_get_allowed, patch.object( - global_mcp_server_manager, - "get_mcp_server_by_id", - return_value=mock_server, - ), patch.object( - global_mcp_server_manager, - "_get_mcp_server_from_tool_name", - return_value=mock_server, - ) as mock_get_server, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_tool_registry" - ) as mock_tool_registry, patch( - "litellm.proxy._experimental.mcp_server.server._handle_managed_mcp_tool", - new_callable=AsyncMock, - ) as mock_handle_managed, patch( - "litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed", - return_value=True, - ) as mock_is_allowed: + with ( + patch.object( + global_mcp_server_manager, + "get_allowed_mcp_servers", + new_callable=AsyncMock, + ) as mock_get_allowed, + patch.object( + global_mcp_server_manager, + "get_mcp_server_by_id", + return_value=mock_server, + ), + patch.object( + global_mcp_server_manager, + "_get_mcp_server_from_tool_name", + return_value=mock_server, + ) as mock_get_server, + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_tool_registry" + ) as mock_tool_registry, + patch( + "litellm.proxy._experimental.mcp_server.server._handle_managed_mcp_tool", + new_callable=AsyncMock, + ) as mock_handle_managed, + patch( + "litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed", + return_value=True, + ) as mock_is_allowed, + ): mock_get_allowed.return_value = [mock_server.server_id] mock_tool_registry.get_tool.return_value = None mock_handle_managed.return_value = expected_response diff --git a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py index 13a09f54e68..dee689708c3 100644 --- a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py +++ b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py @@ -40,6 +40,7 @@ class TestMCPClient: with pytest.raises( ValueError, match="stdio_config is required for stdio transport" ): + async def _noop(session): return None @@ -251,11 +252,11 @@ class TestMCPClient: server_url="http://example.com/sse", transport_type="sse", auth_type=MCPAuth.token, - auth_value="my-secret-token" + auth_value="my-secret-token", ) - + headers = client._get_auth_headers() - + assert "Authorization" in headers assert headers["Authorization"] == "token my-secret-token" @@ -266,27 +267,27 @@ class TestMCPClient: server_url="http://example.com/sse", transport_type="sse", auth_type=MCPAuth.bearer_token, - auth_value="bearer-token" + auth_value="bearer-token", ) headers = client._get_auth_headers() assert headers["Authorization"] == "Bearer bearer-token" - + # Test API key client = MCPClient( server_url="http://example.com/sse", transport_type="sse", auth_type=MCPAuth.api_key, - auth_value="api-key" + auth_value="api-key", ) headers = client._get_auth_headers() assert headers["X-API-Key"] == "api-key" - + # Test basic auth (gets base64 encoded) client = MCPClient( server_url="http://example.com/sse", transport_type="sse", auth_type=MCPAuth.basic, - auth_value="user:pass" + auth_value="user:pass", ) headers = client._get_auth_headers() assert headers["Authorization"].startswith("Basic ") @@ -298,11 +299,11 @@ class TestMCPClient: transport_type="sse", auth_type=MCPAuth.token, auth_value="my-token", - extra_headers={"X-Custom-Header": "custom-value"} + extra_headers={"X-Custom-Header": "custom-value"}, ) - + headers = client._get_auth_headers() - + assert headers["Authorization"] == "token my-token" assert headers["X-Custom-Header"] == "custom-value" @@ -312,5 +313,80 @@ class TestMCPClient: assert MCPAuth.token.value == "token" +# --------------------------------------------------------------------------- +# _last_initialize_instructions capture +# --------------------------------------------------------------------------- + + +class TestMCPClientInstructionsCapture: + """Tests for _last_initialize_instructions capture during session init.""" + + def test_initial_value_is_none(self): + """Fresh client has no cached instructions.""" + client = MCPClient( + server_url="http://example.com/mcp", + transport_type="http", + ) + assert client._last_initialize_instructions is None + + @pytest.mark.asyncio + @patch("litellm.experimental_mcp_client.client.ClientSession") + async def test_captures_instructions_from_initialize(self, mock_session_cls): + """Instructions from upstream initialize() are captured and stripped.""" + client = MCPClient( + server_url="http://example.com/mcp", + transport_type="http", + ) + + mock_session = AsyncMock() + init_result = MagicMock() + init_result.instructions = " upstream says hello " + mock_session.initialize = AsyncMock(return_value=init_result) + + session_ctx = MagicMock() + session_ctx.__aenter__ = AsyncMock(return_value=mock_session) + session_ctx.__aexit__ = AsyncMock(return_value=False) + mock_session_cls.return_value = session_ctx + + transport_ctx = MagicMock() + transport_ctx.__aenter__ = AsyncMock(return_value=(MagicMock(), MagicMock())) + transport_ctx.__aexit__ = AsyncMock(return_value=False) + + async def _op(session): + return "done" + + await client._execute_session_operation(transport_ctx, _op) + assert client._last_initialize_instructions == "upstream says hello" + + @pytest.mark.asyncio + @patch("litellm.experimental_mcp_client.client.ClientSession") + async def test_none_instructions_stays_none(self, mock_session_cls): + """When upstream returns no instructions the field stays None.""" + client = MCPClient( + server_url="http://example.com/mcp", + transport_type="http", + ) + + mock_session = AsyncMock() + init_result = MagicMock() + init_result.instructions = None + mock_session.initialize = AsyncMock(return_value=init_result) + + session_ctx = MagicMock() + session_ctx.__aenter__ = AsyncMock(return_value=mock_session) + session_ctx.__aexit__ = AsyncMock(return_value=False) + mock_session_cls.return_value = session_ctx + + transport_ctx = MagicMock() + transport_ctx.__aenter__ = AsyncMock(return_value=(MagicMock(), MagicMock())) + transport_ctx.__aexit__ = AsyncMock(return_value=False) + + async def _op(session): + return "done" + + await client._execute_session_operation(transport_ctx, _op) + assert client._last_initialize_instructions is None + + if __name__ == "__main__": pytest.main([__file__]) diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py index 22470b93540..2f57ce5d180 100644 --- a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py +++ b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py @@ -440,14 +440,18 @@ class TestProxyOAuthHeaderForwarding: (b"content-type", b"application/json"), ] ) - + # Should preserve OAuth even with flag=False - cleaned_without_flag = clean_headers(raw_headers, forward_llm_provider_auth_headers=False) + cleaned_without_flag = clean_headers( + raw_headers, forward_llm_provider_auth_headers=False + ) assert "authorization" in cleaned_without_flag assert cleaned_without_flag["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}" - + # Should also preserve OAuth with flag=True - cleaned_with_flag = clean_headers(raw_headers, forward_llm_provider_auth_headers=True) + cleaned_with_flag = clean_headers( + raw_headers, forward_llm_provider_auth_headers=True + ) assert "authorization" in cleaned_with_flag assert cleaned_with_flag["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}" @@ -867,8 +871,6 @@ class TestValidateEnvironmentAuthToken: assert "authorization" not in headers - - class TestGetAuthToken: """Tests for AnthropicModelInfo.get_auth_token() static method.""" @@ -1092,7 +1094,10 @@ class TestPassthroughAuthToken: config = AnthropicMessagesConfig() with mock_patch.dict( "os.environ", - {"ANTHROPIC_API_KEY": FAKE_REGULAR_KEY, "ANTHROPIC_AUTH_TOKEN": FAKE_AUTH_TOKEN}, + { + "ANTHROPIC_API_KEY": FAKE_REGULAR_KEY, + "ANTHROPIC_AUTH_TOKEN": FAKE_AUTH_TOKEN, + }, clear=True, ): updated_headers, _ = config.validate_anthropic_messages_environment( @@ -1131,3 +1136,147 @@ class TestPassthroughAuthToken: ) assert url == "https://custom.example.com/v1/messages" + + +class TestAnthropicThinkingSignatureSelfHeal: + """Helpers for retrying after invalid encrypted thinking signatures.""" + + def test_is_anthropic_invalid_thinking_signature_error_positive(self): + from litellm.llms.anthropic.common_utils import ( + is_anthropic_invalid_thinking_signature_error, + ) + + raw = ( + '{"type":"error","error":{"type":"invalid_request_error",' + '"message":"messages.3.content.3: Invalid `signature` in `thinking` block"},' + '"request_id":"req_011Ca2EtQDxp7x6RGUY2jVn9"}' + ) + assert is_anthropic_invalid_thinking_signature_error(raw) is True + + def test_is_anthropic_invalid_thinking_signature_error_negative(self): + from litellm.llms.anthropic.common_utils import ( + is_anthropic_invalid_thinking_signature_error, + ) + + assert is_anthropic_invalid_thinking_signature_error("") is False + assert ( + is_anthropic_invalid_thinking_signature_error("rate limit exceeded") + is False + ) + + def test_strip_thinking_blocks_from_anthropic_messages(self): + from litellm.llms.anthropic.common_utils import ( + strip_thinking_blocks_from_anthropic_messages, + ) + + messages = [ + {"role": "user", "content": "hi"}, + { + "role": "assistant", + "content": [ + {"type": "thinking", "thinking": "plan", "signature": "sig"}, + {"type": "text", "text": "hello"}, + ], + }, + ] + out = strip_thinking_blocks_from_anthropic_messages(messages) + assert len(out) == 2 + assert out[0] == messages[0] + assert len(out[1]["content"]) == 1 + assert out[1]["content"][0]["type"] == "text" + assert messages[1]["content"][0]["type"] == "thinking" + + def test_strip_thinking_blocks_drops_message_when_only_thinking_blocks(self): + from litellm.llms.anthropic.common_utils import ( + strip_thinking_blocks_from_anthropic_messages, + ) + + messages = [ + {"role": "user", "content": "hi"}, + { + "role": "assistant", + "content": [ + {"type": "thinking", "thinking": "plan", "signature": "sig"}, + ], + }, + ] + out = strip_thinking_blocks_from_anthropic_messages(messages) + assert len(out) == 1 + assert out[0]["role"] == "user" + + def test_strip_thinking_blocks_from_anthropic_messages_request_dict(self): + from litellm.llms.anthropic.common_utils import ( + strip_thinking_blocks_from_anthropic_messages_request_dict, + ) + + data = { + "model": "claude-sonnet-4-20250514", + "messages": [ + { + "role": "assistant", + "content": [ + { + "type": "thinking", + "thinking": "x", + "signature": "y", + }, + ], + } + ], + "thinking": {"type": "enabled", "budget_tokens": 1024}, + } + strip_thinking_blocks_from_anthropic_messages_request_dict(data) + assert "thinking" not in data + assert data["messages"] == [] + + def test_anthropic_messages_config_http_retry_helpers(self): + import httpx + + from litellm.llms.anthropic.experimental_pass_through.messages.transformation import ( + AnthropicMessagesConfig, + ) + + config = AnthropicMessagesConfig() + assert config.max_retry_on_anthropic_messages_http_error == 2 + + req = httpx.Request("POST", "https://api.anthropic.com/v1/messages") + err_text = ( + '{"type":"error","error":{"type":"invalid_request_error",' + '"message":"messages.3.content.3: Invalid `signature` in `thinking` block"},' + '"request_id":"req_011Ca2EtQDxp7x6RGUY2jVn9"}' + ) + resp = httpx.Response(400, request=req, text=err_text) + err = httpx.HTTPStatusError("bad", request=req, response=resp) + assert config.should_retry_anthropic_messages_on_http_error(err, {}) is True + + resp_bad = httpx.Response(400, request=req, text="rate limit exceeded") + err_bad = httpx.HTTPStatusError("bad", request=req, response=resp_bad) + assert ( + config.should_retry_anthropic_messages_on_http_error(err_bad, {}) is False + ) + + resp_500 = httpx.Response(500, request=req, text=err_text) + err_500 = httpx.HTTPStatusError("bad", request=req, response=resp_500) + assert ( + config.should_retry_anthropic_messages_on_http_error(err_500, {}) is False + ) + + data = { + "model": "claude-sonnet-4-20250514", + "messages": [ + { + "role": "assistant", + "content": [ + { + "type": "thinking", + "thinking": "x", + "signature": "y", + }, + ], + } + ], + "thinking": {"type": "enabled", "budget_tokens": 1024}, + } + config.transform_anthropic_messages_request_on_http_error(err, data) + assert "thinking" not in data + assert data["messages"] == [] diff --git a/tests/test_litellm/llms/azure/passthrough/test_azure_passthrough_transformation.py b/tests/test_litellm/llms/azure/passthrough/test_azure_passthrough_transformation.py new file mode 100644 index 00000000000..529a7453d74 --- /dev/null +++ b/tests/test_litellm/llms/azure/passthrough/test_azure_passthrough_transformation.py @@ -0,0 +1,97 @@ +import json +import os +import sys +from unittest.mock import MagicMock + +import httpx + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.llms.azure.passthrough.transformation import AzurePassthroughConfig +from litellm.types.utils import ModelResponse + + +def _azure_chat_completion_body(): + return { + "id": "chatcmpl-abc123", + "object": "chat.completion", + "created": 1700000000, + "model": "gpt-4.1-mini-2025-04-14", + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": "Hello! How can I assist you today?", + }, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 10, + "completion_tokens": 8, + "total_tokens": 18, + }, + } + + +def _make_httpx_response(body: dict) -> httpx.Response: + return httpx.Response( + status_code=200, + headers={"content-type": "application/json"}, + content=json.dumps(body).encode("utf-8"), + request=httpx.Request( + "POST", + "https://example.openai.azure.com/openai/deployments/gpt-4.1-mini/chat/completions", + ), + ) + + +def test_azure_passthrough_logging_non_streaming_response_chat_completions(): + """ + Returns a populated ModelResponse (with usage + content) for a chat/completions + endpoint. This is what _success_handler_helper_fn needs to build + standard_logging_object — without it, Datadog/cost-tracking/router-success all + raise on every Azure passthrough request. + """ + config = AzurePassthroughConfig() + logging_obj = MagicMock() + + result = config.logging_non_streaming_response( + model="gpt-4.1-mini", + custom_llm_provider="azure", + httpx_response=_make_httpx_response(_azure_chat_completion_body()), + request_data={ + "model": "gpt-4.1-mini", + "messages": [{"role": "user", "content": "hi"}], + }, + logging_obj=logging_obj, + endpoint="openai/deployments/gpt-4.1-mini/chat/completions", + ) + + assert isinstance(result, ModelResponse) + assert result.choices[0].message.content == "Hello! How can I assist you today?" + assert result.usage.prompt_tokens == 10 + assert result.usage.completion_tokens == 8 + assert result.usage.total_tokens == 18 + + +def test_azure_passthrough_logging_non_streaming_response_unknown_endpoint_returns_none(): + """ + Endpoints other than chat/completions (responses, messages, images) fall + through to None — matches base-class behavior and Bedrock's "unknown + endpoint" handling. Not a regression; just scoping. + """ + config = AzurePassthroughConfig() + logging_obj = MagicMock() + + result = config.logging_non_streaming_response( + model="gpt-4.1-mini", + custom_llm_provider="azure", + httpx_response=_make_httpx_response(_azure_chat_completion_body()), + request_data={}, + logging_obj=logging_obj, + endpoint="openai/responses", + ) + + assert result is None 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 384d428888f..9df6408b0d7 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 @@ -5,7 +5,12 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi import HTTPException from mcp import ReadResourceResult, Resource -from mcp.types import BlobResourceContents, Prompt, ResourceTemplate, TextResourceContents +from mcp.types import ( + BlobResourceContents, + Prompt, + ResourceTemplate, + TextResourceContents, +) from litellm.proxy._types import ( LiteLLM_MCPServerTable, @@ -157,15 +162,19 @@ async def test_get_prompts_from_mcp_servers_success(): server_b.auth_type = None server_b.extra_headers = None - with patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", - AsyncMock(return_value=[server_a, server_b]), - ) as mock_allowed, patch( - "litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers", - return_value=(None, None), - ) as mock_headers, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", - ) as mock_manager: + with ( + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + AsyncMock(return_value=[server_a, server_b]), + ) as mock_allowed, + patch( + "litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers", + return_value=(None, None), + ) as mock_headers, + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + ) as mock_manager, + ): mock_manager.get_prompts_from_server = AsyncMock( side_effect=[ [Prompt(name="hello", description="hi")], @@ -213,15 +222,19 @@ async def test_get_resources_from_mcp_servers_success(): server_b.auth_type = None server_b.extra_headers = None - with patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", - AsyncMock(return_value=[server_a, server_b]), - ) as mock_allowed, patch( - "litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers", - return_value=(None, None), - ) as mock_headers, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", - ) as mock_manager: + with ( + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + AsyncMock(return_value=[server_a, server_b]), + ) as mock_allowed, + patch( + "litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers", + return_value=(None, None), + ) as mock_headers, + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + ) as mock_manager, + ): mock_manager.get_resources_from_server = AsyncMock( side_effect=[ [ @@ -274,15 +287,19 @@ async def test_get_resource_templates_from_mcp_servers_success(): server.auth_type = None server.extra_headers = None - with patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", - AsyncMock(return_value=[server]), - ) as mock_allowed, patch( - "litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers", - return_value=(None, None), - ) as mock_headers, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", - ) as mock_manager: + with ( + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + AsyncMock(return_value=[server]), + ) as mock_allowed, + patch( + "litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers", + return_value=(None, None), + ) as mock_headers, + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + ) as mock_manager, + ): mock_manager.get_resource_templates_from_server = AsyncMock( return_value=[ ResourceTemplate( @@ -320,15 +337,19 @@ async def test_mcp_get_prompt_success(): prompt_result = MagicMock(name="prompt_result") - with patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", - AsyncMock(return_value=[server]), - ) as mock_allowed, patch( - "litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers", - return_value=({"Authorization": "token"}, {"X-Test": "1"}), - ) as mock_headers, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", - ) as mock_manager: + with ( + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + AsyncMock(return_value=[server]), + ) as mock_allowed, + patch( + "litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers", + return_value=({"Authorization": "token"}, {"X-Test": "1"}), + ) as mock_headers, + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + ) as mock_manager, + ): mock_manager.get_prompt_from_server = AsyncMock(return_value=prompt_result) result = await mcp_get_prompt( @@ -378,15 +399,19 @@ async def test_mcp_read_resource_success(): ] ) - with patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", - AsyncMock(return_value=[server]), - ) as mock_allowed, patch( - "litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers", - return_value=({"Authorization": "token"}, {"X-Test": "1"}), - ) as mock_headers, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", - ) as mock_manager: + with ( + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + AsyncMock(return_value=[server]), + ) as mock_allowed, + patch( + "litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers", + return_value=({"Authorization": "token"}, {"X-Test": "1"}), + ) as mock_headers, + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + ) as mock_manager, + ): mock_manager.read_resource_from_server = AsyncMock(return_value=read_result) result = await mcp_read_resource( @@ -591,7 +616,10 @@ async def test_get_tools_from_mcp_servers_continues_when_one_server_fails(): working_server if server_id == "working_server" else failing_server ) # Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering) - mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0) + mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: ( + server_ids, + 0, + ) async def mock_get_tools_from_server( server, @@ -693,7 +721,10 @@ async def test_get_tools_from_mcp_servers_handles_all_servers_failing(): failing_server1 if server_id == "failing_server1" else failing_server2 ) # Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering) - mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0) + mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: ( + server_ids, + 0, + ) async def mock_get_tools_from_server( server, @@ -830,12 +861,14 @@ async def test_concurrent_initialize_session_managers(): mcp_server._sse_session_manager_cm = None # Mock the session managers to avoid actual MCP initialization - with patch( - "litellm.proxy._experimental.mcp_server.server.session_manager" - ) as mock_session_manager, patch( - "litellm.proxy._experimental.mcp_server.server.sse_session_manager" - ) as mock_sse_session_manager, patch( - "litellm.proxy._experimental.mcp_server.server.verbose_logger" + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.session_manager" + ) as mock_session_manager, + patch( + "litellm.proxy._experimental.mcp_server.server.sse_session_manager" + ) as mock_sse_session_manager, + patch("litellm.proxy._experimental.mcp_server.server.verbose_logger"), ): # Mock the run() method to return a mock context manager mock_cm = AsyncMock() @@ -961,15 +994,19 @@ async def test_mcp_routing_with_conflicting_alias_and_group_name(): return_value=[specific_server.server_id, other_server.server_id] ) - with patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_allowed_mcp_servers", - mock_get_allowed, - ), patch( - "litellm.proxy._experimental.mcp_server.server.MCPRequestHandler._get_mcp_servers_from_access_groups", - mock_db_lookup, - ), patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager._get_tools_from_server", - mock_get_tools_spy, + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_allowed_mcp_servers", + mock_get_allowed, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.MCPRequestHandler._get_mcp_servers_from_access_groups", + mock_db_lookup, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager._get_tools_from_server", + mock_get_tools_spy, + ), ): mcp_servers_from_path = _get_mcp_servers_in_path(test_path) @@ -1062,17 +1099,21 @@ async def test_oauth2_headers_passed_to_mcp_client(): async def mock_fetch_tools_with_timeout(client, server_name): return [] # Return empty list of tools - with patch.object( - global_mcp_server_manager, - "_create_mcp_client", - side_effect=mock_create_mcp_client, - ) as mock_create_client, patch.object( - global_mcp_server_manager, - "_fetch_tools_with_timeout", - side_effect=mock_fetch_tools_with_timeout, - ), patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", - AsyncMock(return_value=[oauth2_server]), + with ( + patch.object( + global_mcp_server_manager, + "_create_mcp_client", + side_effect=mock_create_mcp_client, + ) as mock_create_client, + patch.object( + global_mcp_server_manager, + "_fetch_tools_with_timeout", + side_effect=mock_fetch_tools_with_timeout, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + AsyncMock(return_value=[oauth2_server]), + ), ): # Call _get_tools_from_mcp_servers which should eventually call _create_mcp_client await _get_tools_from_mcp_servers( @@ -1138,7 +1179,10 @@ async def test_list_tools_single_server_unprefixed_names(): mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1"]) mock_manager.get_mcp_server_by_id = MagicMock(return_value=server) # Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering) - mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0) + mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: ( + server_ids, + 0, + ) async def mock_get_tools_from_server( server, @@ -1216,7 +1260,10 @@ async def test_list_tools_multiple_servers_prefixed_names(): server1 if server_id == "server1" else server2 ) # Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering) - mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0) + mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: ( + server_ids, + 0, + ) async def mock_get_tools_from_server( server, @@ -1270,12 +1317,15 @@ async def test_mcp_manager_allows_public_servers_without_permissions(): ) manager.registry = {public_server.server_id: public_server} - with patch( - "litellm.proxy.management_endpoints.common_utils._user_has_admin_view", - return_value=False, - ), patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPRequestHandler.get_allowed_mcp_servers", - AsyncMock(return_value=[]), + with ( + patch( + "litellm.proxy.management_endpoints.common_utils._user_has_admin_view", + return_value=False, + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPRequestHandler.get_allowed_mcp_servers", + AsyncMock(return_value=[]), + ), ): allowed = await manager.get_allowed_mcp_servers(UserAPIKeyAuth()) @@ -1302,12 +1352,15 @@ async def test_mcp_manager_returns_public_when_permission_lookup_fails(): ) manager.registry = {public_server.server_id: public_server} - with patch( - "litellm.proxy.management_endpoints.common_utils._user_has_admin_view", - return_value=False, - ), patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPRequestHandler.get_allowed_mcp_servers", - AsyncMock(side_effect=Exception("boom")), + with ( + patch( + "litellm.proxy.management_endpoints.common_utils._user_has_admin_view", + return_value=False, + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPRequestHandler.get_allowed_mcp_servers", + AsyncMock(side_effect=Exception("boom")), + ), ): allowed = await manager.get_allowed_mcp_servers(UserAPIKeyAuth()) @@ -1342,12 +1395,15 @@ async def test_mcp_manager_merges_public_and_restricted_servers(): scoped_server.server_id: scoped_server, } - with patch( - "litellm.proxy.management_endpoints.common_utils._user_has_admin_view", - return_value=False, - ), patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPRequestHandler.get_allowed_mcp_servers", - AsyncMock(return_value=["restricted"]), + with ( + patch( + "litellm.proxy.management_endpoints.common_utils._user_has_admin_view", + return_value=False, + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPRequestHandler.get_allowed_mcp_servers", + AsyncMock(return_value=["restricted"]), + ), ): allowed = await manager.get_allowed_mcp_servers(UserAPIKeyAuth()) @@ -1399,12 +1455,15 @@ async def test_call_mcp_tool_user_unauthorized_access(): return another_server_obj return None - with patch( - "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_allowed_mcp_servers", - AsyncMock(return_value=["allowed_server", "another_server"]), - ), patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_id", - side_effect=mock_get_server_by_id, + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_allowed_mcp_servers", + AsyncMock(return_value=["allowed_server", "another_server"]), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_id", + side_effect=mock_get_server_by_id, + ), ): # Try to call a tool from "restricted_server" - should raise HTTPException with 403 status with pytest.raises(HTTPException) as exc_info: @@ -1467,7 +1526,10 @@ async def test_list_tools_filters_by_key_team_permissions(): mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1"]) mock_manager.get_mcp_server_by_id = lambda server_id: server # Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering) - mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0) + mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: ( + server_ids, + 0, + ) async def mock_get_tools_from_server( server, @@ -1573,7 +1635,10 @@ async def test_list_tools_with_team_tool_permissions_inheritance(): mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1"]) mock_manager.get_mcp_server_by_id = lambda server_id: server # Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering) - mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0) + mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: ( + server_ids, + 0, + ) async def mock_get_tools_from_server( server, @@ -1665,7 +1730,10 @@ async def test_list_tools_with_no_tool_permissions_shows_all(): mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1"]) mock_manager.get_mcp_server_by_id = lambda server_id: server # Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering) - mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0) + mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: ( + server_ids, + 0, + ) async def mock_get_tools_from_server( server, @@ -1760,7 +1828,10 @@ async def test_list_tools_strips_prefix_when_matching_permissions(): mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["gitmcp_server"]) mock_manager.get_mcp_server_by_id = MagicMock(return_value=server) # Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering) - mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0) + mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: ( + server_ids, + 0, + ) async def mock_get_tools_from_server( server, @@ -2002,12 +2073,15 @@ class TestMCPServerManagerReload: mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock( return_value=[db_row] ) - with patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", - return_value=mock_prisma, - ), patch.object( - manager, "build_mcp_server_from_table", AsyncMock() - ) as mock_build: + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=mock_prisma, + ), + patch.object( + manager, "build_mcp_server_from_table", AsyncMock() + ) as mock_build, + ): await manager.reload_servers_from_database() mock_build.assert_not_awaited() @@ -2045,14 +2119,17 @@ class TestMCPServerManagerReload: mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock( return_value=[db_row] ) - with patch( - "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", - return_value=mock_prisma, - ), patch.object( - manager, - "build_mcp_server_from_table", - AsyncMock(return_value=rebuilt_server), - ) as mock_build: + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=mock_prisma, + ), + patch.object( + manager, + "build_mcp_server_from_table", + AsyncMock(return_value=rebuilt_server), + ) as mock_build, + ): await manager.reload_servers_from_database() mock_build.assert_awaited_once_with(db_row) @@ -2090,26 +2167,32 @@ async def test_call_mcp_tool_logs_failure_via_post_call_failure_hook(): user_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user") - with patch.object( - global_mcp_server_manager, - "get_allowed_mcp_servers", - new_callable=AsyncMock, - return_value=[mock_server.server_id], - ), patch.object( - global_mcp_server_manager, - "get_mcp_server_by_id", - return_value=mock_server, - ), patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers_from_mcp_server_names", - new_callable=AsyncMock, - return_value=[mock_server], - ), patch( - "litellm.proxy._experimental.mcp_server.server.execute_mcp_tool", - new_callable=AsyncMock, - side_effect=Exception("boom"), - ), patch( - "litellm.proxy.proxy_server.proxy_logging_obj", - proxy_logging_mock, + with ( + patch.object( + global_mcp_server_manager, + "get_allowed_mcp_servers", + new_callable=AsyncMock, + return_value=[mock_server.server_id], + ), + patch.object( + global_mcp_server_manager, + "get_mcp_server_by_id", + return_value=mock_server, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers_from_mcp_server_names", + new_callable=AsyncMock, + return_value=[mock_server], + ), + patch( + "litellm.proxy._experimental.mcp_server.server.execute_mcp_tool", + new_callable=AsyncMock, + side_effect=Exception("boom"), + ), + patch( + "litellm.proxy.proxy_server.proxy_logging_obj", + proxy_logging_mock, + ), ): with pytest.raises(Exception): await call_mcp_tool( @@ -2157,23 +2240,30 @@ async def test_get_tools_from_mcp_servers_logs_list_tools_to_spendlogs_when_enab dummy_logging_obj.model_call_details = {"metadata": {"spend_logs_metadata": {}}} dummy_logging_obj.async_success_handler = AsyncMock() - with patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", - new=AsyncMock(return_value=[server_a]), - ), patch( - "litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers", - return_value=(None, None), - ), patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", - ) as mock_manager, patch( - "litellm.proxy._experimental.mcp_server.server.filter_tools_by_allowed_tools", - side_effect=lambda tools, _server: tools, - ), patch( - "litellm.proxy._experimental.mcp_server.server.filter_tools_by_key_team_permissions", - new=AsyncMock(side_effect=lambda tools, **_: tools), - ), patch( - "litellm.proxy._experimental.mcp_server.server.function_setup", - return_value=(dummy_logging_obj, None), + with ( + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new=AsyncMock(return_value=[server_a]), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers", + return_value=(None, None), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + ) as mock_manager, + patch( + "litellm.proxy._experimental.mcp_server.server.filter_tools_by_allowed_tools", + side_effect=lambda tools, _server: tools, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.filter_tools_by_key_team_permissions", + new=AsyncMock(side_effect=lambda tools, **_: tools), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.function_setup", + return_value=(dummy_logging_obj, None), + ), ): mock_manager._get_tools_from_server = AsyncMock(return_value=[tool_1]) @@ -2188,7 +2278,9 @@ async def test_get_tools_from_mcp_servers_logs_list_tools_to_spendlogs_when_enab assert tools == [tool_1] dummy_logging_obj.async_success_handler.assert_awaited_once() - assert dummy_logging_obj.async_success_handler.await_args.kwargs["result"] == [tool_1] + assert dummy_logging_obj.async_success_handler.await_args.kwargs["result"] == [ + tool_1 + ] spend_meta = dummy_logging_obj.model_call_details["metadata"]["spend_logs_metadata"] assert spend_meta["tool_count_total"] == 1 @@ -2381,26 +2473,34 @@ async def test_get_tools_from_mcp_servers_injects_stored_oauth2_token(): oauth2_server.extra_headers = None # Simulate the DB returning a valid credential for this user+server - prefetched_creds = {SERVER_ID: {"access_token": STORED_TOKEN, "server_id": SERVER_ID}} + prefetched_creds = { + SERVER_ID: {"access_token": STORED_TOKEN, "server_id": SERVER_ID} + } tool_1 = MagicMock() tool_1.name = "atlassian_test-search" - with patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", - new=AsyncMock(return_value=[oauth2_server]), - ), patch( - # Patch the bulk prefetch so no real DB connection is needed - "litellm.proxy._experimental.mcp_server.server._prefetch_oauth_creds_for_user", - new=AsyncMock(return_value=prefetched_creds), - ) as mock_prefetch, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", - ) as mock_manager, patch( - "litellm.proxy._experimental.mcp_server.server.filter_tools_by_allowed_tools", - side_effect=lambda tools, _server: tools, - ), patch( - "litellm.proxy._experimental.mcp_server.server.filter_tools_by_key_team_permissions", - new=AsyncMock(side_effect=lambda tools, **_: tools), + with ( + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new=AsyncMock(return_value=[oauth2_server]), + ), + patch( + # Patch the bulk prefetch so no real DB connection is needed + "litellm.proxy._experimental.mcp_server.server._prefetch_oauth_creds_for_user", + new=AsyncMock(return_value=prefetched_creds), + ) as mock_prefetch, + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + ) as mock_manager, + patch( + "litellm.proxy._experimental.mcp_server.server.filter_tools_by_allowed_tools", + side_effect=lambda tools, _server: tools, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.filter_tools_by_key_team_permissions", + new=AsyncMock(side_effect=lambda tools, **_: tools), + ), ): mock_manager._get_tools_from_server = AsyncMock(return_value=[tool_1]) @@ -2421,3 +2521,201 @@ async def test_get_tools_from_mcp_servers_injects_stored_oauth2_token(): assert call_kwargs["extra_headers"] == {"Authorization": f"Bearer {STORED_TOKEN}"} assert tools == [tool_1] + + +# --------------------------------------------------------------------------- +# _merge_gateway_initialize_instructions + ContextVar / InitializationOptions +# --------------------------------------------------------------------------- + + +def _make_instruction_server( + server_id="s1", + name="s1", + *, + alias=None, + server_name=None, + instructions=None, + spec_path=None, + url="https://example.com", +): + return MCPServer( + server_id=server_id, + name=name, + alias=alias, + server_name=server_name, + url=url, + transport=MCPTransport.http, + instructions=instructions, + spec_path=spec_path, + ) + + +class TestMergeGatewayInitializeInstructions: + """Tests for _merge_gateway_initialize_instructions.""" + + def _merge(self, servers): + try: + from litellm.proxy._experimental.mcp_server.server import ( + _merge_gateway_initialize_instructions, + ) + except ImportError: + pytest.skip("MCP server not available") + return _merge_gateway_initialize_instructions(servers) + + def test_empty_server_list_returns_none(self): + """No servers yields no instructions.""" + assert self._merge([]) is None + + def test_single_server_yaml_instructions(self): + """A single server with YAML instructions returns them verbatim.""" + s = _make_instruction_server(instructions="Use add() for sums.") + assert self._merge([s]) == "Use add() for sums." + + def test_yaml_instructions_strips_whitespace(self): + """Leading/trailing whitespace is stripped.""" + s = _make_instruction_server(instructions=" padded \n") + assert self._merge([s]) == "padded" + + def test_yaml_override_beats_upstream_cache(self): + """YAML/DB instructions take precedence over upstream cache.""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + global_mcp_server_manager._upstream_initialize_instructions_by_server_id[ + "s1" + ] = "upstream" + try: + s = _make_instruction_server(instructions="yaml wins") + assert self._merge([s]) == "yaml wins" + finally: + global_mcp_server_manager._upstream_initialize_instructions_by_server_id.pop( + "s1", None + ) + + def test_upstream_cache_used_when_no_yaml(self): + """Upstream cached instructions are used when no YAML override is set.""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + global_mcp_server_manager._upstream_initialize_instructions_by_server_id[ + "s1" + ] = "from upstream" + try: + s = _make_instruction_server(instructions=None) + assert self._merge([s]) == "from upstream" + finally: + global_mcp_server_manager._upstream_initialize_instructions_by_server_id.pop( + "s1", None + ) + + def test_spec_path_servers_skipped(self): + """OpenAPI (spec_path) servers do not contribute instructions.""" + s = _make_instruction_server(spec_path="/openapi.json", url=None) + assert self._merge([s]) is None + + def test_no_instructions_no_cache_returns_none(self): + """Server with no instructions and no cache yields None.""" + s = _make_instruction_server() + assert self._merge([s]) is None + + def test_multiple_servers_merged_with_labels(self): + """Multiple servers get label-prefixed and separator-joined.""" + s1 = _make_instruction_server( + server_id="a", name="a", alias="Alpha", instructions="instr A" + ) + s2 = _make_instruction_server( + server_id="b", name="b", alias="Beta", instructions="instr B" + ) + result = self._merge([s1, s2]) + assert result is not None + assert "[Alpha]" in result and "[Beta]" in result + assert "instr A" in result and "instr B" in result + assert "---" in result + + def test_single_server_no_label_wrapping(self): + """A single server's instructions are not wrapped with a label.""" + s = _make_instruction_server(alias="MyServer", instructions="single") + result = self._merge([s]) + assert result == "single" + assert "[MyServer]" not in result + + def test_mixed_yaml_cache_specpath(self): + """YAML, upstream-cache, and spec_path servers are handled correctly together.""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + global_mcp_server_manager._upstream_initialize_instructions_by_server_id[ + "c" + ] = "cached C" + try: + s_yaml = _make_instruction_server( + server_id="a", name="a", alias="A", instructions="yaml A" + ) + s_spec = _make_instruction_server( + server_id="b", name="b", alias="B", spec_path="/spec.json", url=None + ) + s_cached = _make_instruction_server(server_id="c", name="c", alias="C") + result = self._merge([s_yaml, s_spec, s_cached]) + assert "yaml A" in result + assert "cached C" in result + assert "[B]" not in result + finally: + global_mcp_server_manager._upstream_initialize_instructions_by_server_id.pop( + "c", None + ) + + +class TestGatewayCreateInitializationOptions: + """Tests for the patched server.create_initialization_options via ContextVar.""" + + def test_no_contextvar_returns_default_options(self): + """When ContextVar is None, instructions are absent.""" + try: + from litellm.proxy._experimental.mcp_server.mcp_context import ( + _mcp_gateway_initialize_instructions, + ) + from litellm.proxy._experimental.mcp_server.server import server + except ImportError: + pytest.skip("MCP server not available") + + tok = _mcp_gateway_initialize_instructions.set(None) + try: + opts = server.create_initialization_options() + assert getattr(opts, "instructions", None) is None + finally: + _mcp_gateway_initialize_instructions.reset(tok) + + def test_contextvar_set_injects_instructions(self): + """When ContextVar has a value, it appears in InitializationOptions.""" + try: + from litellm.proxy._experimental.mcp_server.mcp_context import ( + _mcp_gateway_initialize_instructions, + ) + from litellm.proxy._experimental.mcp_server.server import server + except ImportError: + pytest.skip("MCP server not available") + + tok = _mcp_gateway_initialize_instructions.set("hello from merge") + try: + opts = server.create_initialization_options() + assert opts.instructions == "hello from merge" + finally: + _mcp_gateway_initialize_instructions.reset(tok) + + def test_contextvar_reset_removes_instructions(self): + """After resetting the ContextVar, instructions disappear.""" + try: + from litellm.proxy._experimental.mcp_server.mcp_context import ( + _mcp_gateway_initialize_instructions, + ) + from litellm.proxy._experimental.mcp_server.server import server + except ImportError: + pytest.skip("MCP server not available") + + tok = _mcp_gateway_initialize_instructions.set("temporary") + _mcp_gateway_initialize_instructions.reset(tok) + opts = server.create_initialization_options() + assert getattr(opts, "instructions", None) is None 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 656a9c616e8..ac5349e7105 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 @@ -43,10 +43,10 @@ def _reload_mcp_manager_module(): # After reload, server.py still holds a stale reference to the old # global_mcp_server_manager. Update it so tests that exercise server.py # functions (e.g. _get_tools_from_mcp_servers) use the fresh instance. - server_module = sys.modules.get( - "litellm.proxy._experimental.mcp_server.server" - ) - if server_module is not None and hasattr(server_module, "global_mcp_server_manager"): + server_module = sys.modules.get("litellm.proxy._experimental.mcp_server.server") + if server_module is not None and hasattr( + server_module, "global_mcp_server_manager" + ): server_module.global_mcp_server_manager = reloaded.global_mcp_server_manager return reloaded @@ -223,9 +223,7 @@ class TestMCPServerManager: with caplog.at_level(logging.WARNING, logger="LiteLLM"): await manager.load_servers_from_config(config) - assert any( - "invalid alias 'bad/name'" in message for message in caplog.messages - ) + assert any("invalid alias 'bad/name'" in message for message in caplog.messages) @pytest.mark.asyncio async def test_load_servers_from_config_accepts_valid_alias(self, caplog): @@ -492,7 +490,12 @@ class TestMCPServerManager: mock_client = AsyncMock() mock_client.list_prompts = AsyncMock(return_value=[mock_prompt]) - with patch.object(manager, "_create_mcp_client", new_callable=AsyncMock, return_value=mock_client): + with patch.object( + manager, + "_create_mcp_client", + new_callable=AsyncMock, + return_value=mock_client, + ): prompts = await manager.get_prompts_from_server(server, add_prefix=True) mock_client.list_prompts.assert_awaited_once() @@ -520,7 +523,12 @@ class TestMCPServerManager: mock_client = AsyncMock() mock_client.get_prompt = AsyncMock(return_value=mock_result) - with patch.object(manager, "_create_mcp_client", new_callable=AsyncMock, return_value=mock_client): + with patch.object( + manager, + "_create_mcp_client", + new_callable=AsyncMock, + return_value=mock_client, + ): result = await manager.get_prompt_from_server( server=server, prompt_name="hello", @@ -551,13 +559,23 @@ class TestMCPServerManager: mock_client = AsyncMock() mock_resources = [Resource(name="file", uri="https://example.com/file")] mock_client.list_resources = AsyncMock(return_value=mock_resources) - prefixed_resources = [Resource(name="alias-server-file", uri="https://example.com/file")] + prefixed_resources = [ + Resource(name="alias-server-file", uri="https://example.com/file") + ] - with patch.object(manager, "_create_mcp_client", new_callable=AsyncMock, return_value=mock_client) as mock_create_client, patch.object( - manager, - "_create_prefixed_resources", - return_value=prefixed_resources, - ) as mock_prefix: + with ( + patch.object( + manager, + "_create_mcp_client", + new_callable=AsyncMock, + return_value=mock_client, + ) as mock_create_client, + patch.object( + manager, + "_create_prefixed_resources", + return_value=prefixed_resources, + ) as mock_prefix, + ): result = await manager.get_resources_from_server( server=server, mcp_auth_header="auth", @@ -602,11 +620,19 @@ class TestMCPServerManager: ) ] - with patch.object(manager, "_create_mcp_client", new_callable=AsyncMock, return_value=mock_client) as mock_create_client, patch.object( - manager, - "_create_prefixed_resource_templates", - return_value=prefixed_templates, - ) as mock_prefix: + with ( + patch.object( + manager, + "_create_mcp_client", + new_callable=AsyncMock, + return_value=mock_client, + ) as mock_create_client, + patch.object( + manager, + "_create_prefixed_resource_templates", + return_value=prefixed_templates, + ) as mock_prefix, + ): result = await manager.get_resource_templates_from_server( server=server, mcp_auth_header="auth", @@ -650,7 +676,12 @@ class TestMCPServerManager: ) mock_client.read_resource = AsyncMock(return_value=read_result) - with patch.object(manager, "_create_mcp_client", new_callable=AsyncMock, return_value=mock_client) as mock_create_client: + with patch.object( + manager, + "_create_mcp_client", + new_callable=AsyncMock, + return_value=mock_client, + ) as mock_create_client: result = await manager.read_resource_from_server( server=server, url="https://example.com/resource", @@ -661,7 +692,9 @@ class TestMCPServerManager: mock_create_client.assert_called_once() called_kwargs = mock_create_client.call_args.kwargs assert called_kwargs["extra_headers"] == {"X-Test": "1", "X-Static": "1"} - mock_client.read_resource.assert_awaited_once_with("https://example.com/resource") + mock_client.read_resource.assert_awaited_once_with( + "https://example.com/resource" + ) assert result is read_result @pytest.mark.asyncio @@ -724,22 +757,27 @@ class TestMCPServerManager: registration_url=None, ) - with patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client", - return_value=mock_client, - ), patch.object( - manager, - "_fetch_oauth_metadata_from_resource", - AsyncMock(return_value=([], None)), - ), patch.object( - manager, - "_attempt_well_known_discovery", - AsyncMock(return_value=([], None)), - ), patch.object( - manager, - "_fetch_authorization_server_metadata", - AsyncMock(return_value=mock_metadata), - ) as mock_fetch_auth: + with ( + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client", + return_value=mock_client, + ), + patch.object( + manager, + "_fetch_oauth_metadata_from_resource", + AsyncMock(return_value=([], None)), + ), + patch.object( + manager, + "_attempt_well_known_discovery", + AsyncMock(return_value=([], None)), + ), + patch.object( + manager, + "_fetch_authorization_server_metadata", + AsyncMock(return_value=mock_metadata), + ) as mock_fetch_auth, + ): result = await manager._descovery_metadata(server_url) mock_fetch_auth.assert_awaited_once_with(["https://example.com"]) @@ -779,9 +817,8 @@ class TestMCPServerManager: assert server.scopes == ["config"] # config overrides discovery assert server.authorization_url == "https://config.example.com/auth" assert server.token_url == "https://discovered.example.com/token" - assert ( - server.registration_url == "https://discovered.example.com/register" - ) + assert server.registration_url == "https://discovered.example.com/register" + @pytest.mark.asyncio async def test_config_oauth_initialize_tool_name_to_mcp_server_name_mapping(self): manager = MCPServerManager() @@ -801,7 +838,7 @@ class TestMCPServerManager: # Initialize the tool mapping await manager._initialize_tool_name_to_mcp_server_name_mapping() assert manager.tool_name_to_mcp_server_name_mapping == {} - + @pytest.mark.asyncio async def test_list_tools_handles_missing_server_alias(self): """Test that list_tools handles servers without alias gracefully""" @@ -1017,7 +1054,9 @@ class TestMCPServerManager: # Capture the extra_headers passed to _create_mcp_client captured_extra_headers = None - async def capture_create_mcp_client(server, mcp_auth_header, extra_headers, stdio_env): + async def capture_create_mcp_client( + server, mcp_auth_header, extra_headers, stdio_env + ): nonlocal captured_extra_headers captured_extra_headers = extra_headers return mock_client @@ -1314,15 +1353,19 @@ class TestMCPServerManager: return tool_func - with patch( - "litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.create_tool_function", - side_effect=fake_create_tool_function, - ), patch( - "litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.build_input_schema", - return_value={"type": "object", "properties": {}, "required": []}, - ), patch( - "litellm.proxy._experimental.mcp_server.tool_registry.global_mcp_tool_registry.register_tool", - return_value=None, + with ( + patch( + "litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.create_tool_function", + side_effect=fake_create_tool_function, + ), + patch( + "litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.build_input_schema", + return_value={"type": "object", "properties": {}, "required": []}, + ), + patch( + "litellm.proxy._experimental.mcp_server.tool_registry.global_mcp_tool_registry.register_tool", + return_value=None, + ), ): await manager._register_openapi_tools( spec_path=str(spec_path), @@ -2161,7 +2204,9 @@ class TestMCPServerManager: # Register the server and map a tool to it manager.registry = {"test-server": server} manager.tool_name_to_mcp_server_name_mapping["test_tool"] = "test-server" - manager.tool_name_to_mcp_server_name_mapping["test-server-test_tool"] = "test-server" + manager.tool_name_to_mcp_server_name_mapping["test-server-test_tool"] = ( + "test-server" + ) # Create mock client that tracks call_tool usage mock_client = AsyncMock() @@ -2252,11 +2297,16 @@ class TestMCPServerManager: # Verify MCPRequestHandler.get_allowed_mcp_servers was called with user_api_key_auth mock_get_allowed.assert_called_once() call_args = mock_get_allowed.call_args - assert call_args[0][0] is user_api_key_auth # First positional arg should be user_api_key_auth + assert ( + call_args[0][0] is user_api_key_auth + ) # First positional arg should be user_api_key_auth assert call_args[0][0].user_id == "user-123" assert call_args[0][0].object_permission_id == "perm_123" assert call_args[0][0].object_permission is not None - assert call_args[0][0].object_permission.mcp_servers == ["test_server_1", "test_server_2"] + assert call_args[0][0].object_permission.mcp_servers == [ + "test_server_1", + "test_server_2", + ] # Verify result contains the expected servers assert "test_server_1" in result @@ -2483,5 +2533,82 @@ class TestHasClientCredentialsOAuth2Flow: assert server.needs_user_oauth_token is False +# --------------------------------------------------------------------------- +# Upstream initialize-instructions cache +# --------------------------------------------------------------------------- + + +class TestMCPServerManagerUpstreamInstructionsCache: + """Tests for the upstream initialize-instructions cache.""" + + def test_get_returns_none_when_empty(self): + """Empty cache returns None for any key.""" + manager = MCPServerManager() + assert ( + manager._upstream_initialize_instructions_by_server_id.get("nonexistent") + is None + ) + + def test_remember_stores_stripped_value(self): + """_remember_upstream_initialize_instructions stores a stripped string.""" + manager = MCPServerManager() + fake_server = MagicMock(server_id="srv") + fake_client = MagicMock(_last_initialize_instructions=" hello \n") + manager._remember_upstream_initialize_instructions(fake_server, fake_client) + assert ( + manager._upstream_initialize_instructions_by_server_id.get("srv") == "hello" + ) + + def test_remember_ignores_empty_string(self): + """Whitespace-only instructions are not stored.""" + manager = MCPServerManager() + fake_server = MagicMock(server_id="srv") + fake_client = MagicMock(_last_initialize_instructions=" ") + manager._remember_upstream_initialize_instructions(fake_server, fake_client) + assert manager._upstream_initialize_instructions_by_server_id.get("srv") is None + + def test_remember_ignores_none(self): + """None instructions are not stored.""" + manager = MCPServerManager() + fake_server = MagicMock(server_id="srv") + fake_client = MagicMock(_last_initialize_instructions=None) + manager._remember_upstream_initialize_instructions(fake_server, fake_client) + assert manager._upstream_initialize_instructions_by_server_id.get("srv") is None + + @pytest.mark.asyncio + async def test_load_servers_from_config_clears_cache(self): + """Reloading config clears any previously cached upstream instructions.""" + manager = MCPServerManager() + manager._upstream_initialize_instructions_by_server_id["old"] = "stale" + await manager.load_servers_from_config( + mcp_servers_config={ + "fresh_srv": { + "url": "https://example.com", + "instructions": "from yaml", + } + } + ) + assert manager._upstream_initialize_instructions_by_server_id.get("old") is None + + @pytest.mark.asyncio + async def test_load_servers_reads_instructions_from_config(self): + """instructions field from YAML config is persisted on the MCPServer.""" + manager = MCPServerManager() + await manager.load_servers_from_config( + mcp_servers_config={ + "srv_a": { + "url": "https://a.example.com", + "instructions": "A instructions", + }, + "srv_b": { + "url": "https://b.example.com", + }, + } + ) + by_name = {s.server_name: s for s in manager.config_mcp_servers.values()} + assert "srv_a" in by_name and by_name["srv_a"].instructions == "A instructions" + assert "srv_b" in by_name and by_name["srv_b"].instructions is None + + if __name__ == "__main__": pytest.main([__file__]) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py index 7c142e3a771..32b988ddb22 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py @@ -399,7 +399,9 @@ class TestMCPServerManagerSigV4: server = next(iter(manager.config_mcp_servers.values())) assert server.auth_type == MCPAuth.aws_sigv4 assert server.aws_access_key_id == "AKIAIOSFODNN7EXAMPLE" - assert server.aws_secret_access_key == "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY" + assert ( + server.aws_secret_access_key == "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY" + ) assert server.aws_region_name == "us-east-1" assert server.aws_service_name == "bedrock-agentcore" @@ -529,7 +531,9 @@ class TestMCPServerManagerSigV4: "aws_session_name": "my-session", } - result = manager._extract_aws_credentials(creds, credentials_are_encrypted=False) + result = manager._extract_aws_credentials( + creds, credentials_are_encrypted=False + ) assert result["aws_role_name"] == "arn:aws:iam::123456789012:role/TestRole" assert result["aws_session_name"] == "my-session" @@ -615,12 +619,15 @@ class TestCredentialMergeOnUpdate: credentials={"aws_region_name": "eu-west-1"}, ) - with patch( - "litellm.proxy._experimental.mcp_server.db._get_salt_key", - return_value=None, - ), patch( - "litellm.proxy._experimental.mcp_server.db.encrypt_value_helper", - side_effect=lambda value, new_encryption_key: value, + with ( + patch( + "litellm.proxy._experimental.mcp_server.db._get_salt_key", + return_value=None, + ), + patch( + "litellm.proxy._experimental.mcp_server.db.encrypt_value_helper", + side_effect=lambda value, new_encryption_key: value, + ), ): await update_mcp_server(mock_prisma, data, "test-user") @@ -685,12 +692,15 @@ class TestCredentialMergeOnUpdate: credentials={"aws_region_name": "us-east-1"}, ) - with patch( - "litellm.proxy._experimental.mcp_server.db._get_salt_key", - return_value=None, - ), patch( - "litellm.proxy._experimental.mcp_server.db.encrypt_value_helper", - side_effect=lambda value, new_encryption_key: value, + with ( + patch( + "litellm.proxy._experimental.mcp_server.db._get_salt_key", + return_value=None, + ), + patch( + "litellm.proxy._experimental.mcp_server.db.encrypt_value_helper", + side_effect=lambda value, new_encryption_key: value, + ), ): await update_mcp_server(mock_prisma, data, "test-user") @@ -728,12 +738,15 @@ class TestCredentialMergeOnUpdate: credentials={"auth_value": "my-key"}, ) - with patch( - "litellm.proxy._experimental.mcp_server.db._get_salt_key", - return_value=None, - ), patch( - "litellm.proxy._experimental.mcp_server.db.encrypt_value_helper", - side_effect=lambda value, new_encryption_key: f"enc:{value}", + with ( + patch( + "litellm.proxy._experimental.mcp_server.db._get_salt_key", + return_value=None, + ), + patch( + "litellm.proxy._experimental.mcp_server.db.encrypt_value_helper", + side_effect=lambda value, new_encryption_key: f"enc:{value}", + ), ): await update_mcp_server(mock_prisma, data, "test-user") @@ -772,12 +785,15 @@ class TestCredentialMergeOnUpdate: credentials={"scopes": ["read", "write"]}, ) - with patch( - "litellm.proxy._experimental.mcp_server.db._get_salt_key", - return_value=None, - ), patch( - "litellm.proxy._experimental.mcp_server.db.encrypt_value_helper", - side_effect=lambda value, new_encryption_key: value, + with ( + patch( + "litellm.proxy._experimental.mcp_server.db._get_salt_key", + return_value=None, + ), + patch( + "litellm.proxy._experimental.mcp_server.db.encrypt_value_helper", + side_effect=lambda value, new_encryption_key: value, + ), ): await update_mcp_server(mock_prisma, data, "test-user") @@ -803,7 +819,9 @@ class TestSigV4BuildFromTable: table_record.server_name = "sigv4_server" table_record.alias = None table_record.description = None - table_record.url = "https://bedrock-agentcore.us-east-1.amazonaws.com/invocations" + table_record.url = ( + "https://bedrock-agentcore.us-east-1.amazonaws.com/invocations" + ) table_record.spec_path = None table_record.transport = "http" table_record.auth_type = "aws_sigv4" @@ -838,6 +856,7 @@ class TestSigV4BuildFromTable: table_record.tool_name_to_description = None table_record.byok_api_key_help_url = None table_record.oauth2_flow = None + table_record.instructions = None manager = MCPServerManager() @@ -895,6 +914,7 @@ class TestSigV4BuildFromTable: table_record.tool_name_to_description = None table_record.byok_api_key_help_url = None table_record.oauth2_flow = None + table_record.instructions = None manager = MCPServerManager() @@ -934,7 +954,9 @@ class TestDecryptCredentials: with patch( "litellm.proxy._experimental.mcp_server.db.decrypt_value_helper", - side_effect=lambda value, key, exception_type="error", return_original_value=False: value.replace("enc:", ""), + side_effect=lambda value, key, exception_type="error", return_original_value=False: value.replace( + "enc:", "" + ), ): result = decrypt_credentials(credentials=creds) @@ -956,7 +978,9 @@ class TestDecryptCredentials: with patch( "litellm.proxy._experimental.mcp_server.db.decrypt_value_helper", - side_effect=lambda value, key, exception_type="error", return_original_value=False: value.replace("enc:", ""), + side_effect=lambda value, key, exception_type="error", return_original_value=False: value.replace( + "enc:", "" + ), ): result = decrypt_credentials(credentials=creds) @@ -988,15 +1012,21 @@ class TestRotateCredentials: ) mock_prisma.db.litellm_mcpservertable.update = AsyncMock() - with patch( - "litellm.proxy._experimental.mcp_server.db._get_salt_key", - return_value="old-key", - ), patch( - "litellm.proxy._experimental.mcp_server.db.decrypt_value_helper", - side_effect=lambda value, key, exception_type="error", return_original_value=False: value.replace("enc_old:", ""), - ), patch( - "litellm.proxy._experimental.mcp_server.db.encrypt_value_helper", - side_effect=lambda value, new_encryption_key: f"enc_new:{value}", + with ( + patch( + "litellm.proxy._experimental.mcp_server.db._get_salt_key", + return_value="old-key", + ), + patch( + "litellm.proxy._experimental.mcp_server.db.decrypt_value_helper", + side_effect=lambda value, key, exception_type="error", return_original_value=False: value.replace( + "enc_old:", "" + ), + ), + patch( + "litellm.proxy._experimental.mcp_server.db.encrypt_value_helper", + side_effect=lambda value, new_encryption_key: f"enc_new:{value}", + ), ): await rotate_mcp_server_credentials_master_key( mock_prisma, "admin", "new-key" diff --git a/tests/test_litellm/proxy/test_health_check_max_tokens.py b/tests/test_litellm/proxy/test_health_check_max_tokens.py index e26f7fb9f20..4d417b40b59 100644 --- a/tests/test_litellm/proxy/test_health_check_max_tokens.py +++ b/tests/test_litellm/proxy/test_health_check_max_tokens.py @@ -1,7 +1,10 @@ +from unittest.mock import AsyncMock, MagicMock, patch + import pytest -from litellm.proxy.health_check import _update_litellm_params_for_health_check + from litellm.litellm_core_utils.health_check_helpers import HealthCheckHelpers -from unittest.mock import AsyncMock, patch, MagicMock +from litellm.proxy import health_check as hc_module +from litellm.proxy.health_check import _update_litellm_params_for_health_check @pytest.mark.asyncio @@ -50,10 +53,13 @@ async def test_ahealth_check_wildcard_models_respects_max_tokens(): Test that ahealth_check_wildcard_models respects max_tokens if passed, otherwise defaults to 10. """ - with patch( - "litellm.litellm_core_utils.llm_request_utils.pick_cheapest_chat_models_from_llm_provider", - return_value=["gpt-4o-mini"], - ), patch("litellm.acompletion", new_callable=AsyncMock): + with ( + patch( + "litellm.litellm_core_utils.llm_request_utils.pick_cheapest_chat_models_from_llm_provider", + return_value=["gpt-4o-mini"], + ), + patch("litellm.acompletion", new_callable=AsyncMock), + ): # Test Case 1: No max_tokens passed, should default to 10 model_params = {} await HealthCheckHelpers.ahealth_check_wildcard_models( @@ -73,3 +79,50 @@ async def test_ahealth_check_wildcard_models_respects_max_tokens(): litellm_logging_obj=MagicMock(), ) assert model_params["max_tokens"] == 3 + + +@pytest.mark.asyncio +async def test_background_health_check_max_tokens_env_var(monkeypatch): + """ + Test that BACKGROUND_HEALTH_CHECK_MAX_TOKENS env var is used as global default + for explicit (non-wildcard) models. + """ + monkeypatch.setattr(hc_module, "BACKGROUND_HEALTH_CHECK_MAX_TOKENS", 10) + + model_info = {} + litellm_params = {"model": "azure/gpt-4"} + + updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + + assert updated_params["max_tokens"] == 10 + + +@pytest.mark.asyncio +async def test_per_model_overrides_global_env_var(monkeypatch): + """ + Test that per-model health_check_max_tokens takes priority over + BACKGROUND_HEALTH_CHECK_MAX_TOKENS env var. + """ + monkeypatch.setattr(hc_module, "BACKGROUND_HEALTH_CHECK_MAX_TOKENS", 10) + + model_info = {"health_check_max_tokens": 5} + litellm_params = {"model": "azure/gpt-4"} + + updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + + assert updated_params["max_tokens"] == 5 + + +@pytest.mark.asyncio +async def test_global_env_var_applies_to_wildcard_models(monkeypatch): + """ + Test that BACKGROUND_HEALTH_CHECK_MAX_TOKENS env var also applies to wildcard models. + """ + monkeypatch.setattr(hc_module, "BACKGROUND_HEALTH_CHECK_MAX_TOKENS", 15) + + model_info = {} + litellm_params = {"model": "openai/*"} + + updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + + assert updated_params["max_tokens"] == 15 diff --git a/tests/test_litellm/test_model_param_helper.py b/tests/test_litellm/test_model_param_helper.py index c6e4b864a22..a62779aeab1 100644 --- a/tests/test_litellm/test_model_param_helper.py +++ b/tests/test_litellm/test_model_param_helper.py @@ -31,3 +31,32 @@ def test_get_standard_logging_model_parameters_excludes_prompt_content(): assert "prompt" not in result assert "input" not in result assert result == {"temperature": 0.5} + + +def test_get_all_llm_api_params_includes_responses_api(): + """ + Regression guard for the Responses API cache-key bug: + Responses-API-only kwargs must be present in the cache-key allow-list, + otherwise Cache.get_cache_key() silently drops them and two requests + that differ only in (e.g.) `instructions` collide on the same key. + """ + all_params = ModelParamHelper._get_all_llm_api_params() + responses_only_params = { + "instructions", + "previous_response_id", + "reasoning", + "include", + "store", + "background", + "max_output_tokens", + "max_tool_calls", + "prompt_cache_key", + "prompt_cache_retention", + "context_management", + "conversation", + "safety_identifier", + } + missing = responses_only_params - all_params + assert ( + missing == set() + ), f"Responses-API kwargs missing from cache-key allow-list: {sorted(missing)}"