From 96ed00e1840669accc54c1fbeb19639ce32023a2 Mon Sep 17 00:00:00 2001 From: Milan Date: Tue, 14 Apr 2026 14:19:31 +0300 Subject: [PATCH] feat(mcp): gateway InitializeResult.instructions from upstream or YAML - Add optional instructions on MCPServer (config/DB/types) and Prisma migration. - MCPClient: fetch_upstream_initialize_instructions() for one-shot initialize. - Gateway merges per-request instructions: YAML/API overrides; otherwise fetch upstream initialize instructions (skip spec_path/OpenAPI-only servers). - Pass auth headers into instruction merge; ContextVar for gateway Server. - REST: wire instructions on connection-test MCPServer payloads. Made-with: Cursor --- .../migration.sql | 2 + .../litellm_proxy_extras/schema.prisma | 1 + litellm/experimental_mcp_client/client.py | 46 +++++ .../_experimental/mcp_server/mcp_context.py | 5 + .../mcp_server/mcp_server_manager.py | 4 + .../mcp_server/rest_endpoints.py | 1 + .../proxy/_experimental/mcp_server/server.py | 178 ++++++++++++++++-- litellm/proxy/_types.py | 4 + litellm/proxy/schema.prisma | 1 + .../types/mcp_server/mcp_server_manager.py | 2 + schema.prisma | 1 + 11 files changed, 229 insertions(+), 16 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260414140000_add_mcp_server_instructions/migration.sql 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 fce95465b55..a728d912715 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? // MCP InitializeResult.instructions (optional) url String? spec_path String? transport String @default("sse") diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 1423617cac0..fe56a418e7b 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -329,6 +329,52 @@ class MCPClient: except BaseException as e: verbose_logger.debug(f"Error during http_client cleanup: {e}") + async def fetch_upstream_initialize_instructions(self) -> Optional[str]: + """Open a transport, run ``initialize`` once, return upstream ``instructions``.""" + http_client: Optional[httpx.AsyncClient] = None + try: + transport_ctx, http_client = self._create_transport_context() + transport = await transport_ctx.__aenter__() + try: + read_stream, write_stream = transport[0], transport[1] + session_ctx = ClientSession(read_stream, write_stream) + session = await session_ctx.__aenter__() + try: + init = await session.initialize() + return init.instructions + finally: + try: + await session_ctx.__aexit__(None, None, None) + except BaseException as e: + verbose_logger.debug( + "Error during session context exit (instructions fetch): %s", + e, + ) + finally: + try: + await transport_ctx.__aexit__(None, None, None) + except BaseException as e: + verbose_logger.debug( + "Error during transport context exit (instructions fetch): %s", + e, + ) + except Exception as e: + verbose_logger.debug( + "fetch_upstream_initialize_instructions failed for %s: %s", + self.server_url or "stdio", + e, + ) + return None + finally: + if http_client is not None: + try: + await http_client.aclose() + except BaseException as e: + verbose_logger.debug( + "Error during http_client cleanup (instructions fetch): %s", + e, + ) + def update_auth_value(self, mcp_auth_value: Union[str, Dict[str, str]]): """ Set the authentication header for the MCP client. 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 8d3831e75fb..dd7f092cb53 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -351,6 +351,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 +694,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 @@ -2946,6 +2948,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( @@ -3041,6 +3044,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..1402b385a8e 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -6,6 +6,7 @@ LiteLLM MCP Server Routes import asyncio import contextlib +import contextvars import time import traceback import uuid @@ -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,10 +206,27 @@ if MCP_AVAILABLE: ) return normalized + class _LitellmMcpGatewayServer(Server): + """Gateway server that injects per-request ``InitializeResult.instructions``.""" + + def create_initialization_options( # type: ignore[override] + self, + notification_options: Optional[NotificationOptions] = None, + experimental_capabilities: Optional[Dict[str, Dict[str, Any]]] = None, + ) -> InitializationOptions: + opts = super().create_initialization_options( + 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 ################# ######################################################## - server: Server = Server( + server: Server = _LitellmMcpGatewayServer( name=LITELLM_MCP_SERVER_NAME, version=LITELLM_MCP_SERVER_VERSION, ) @@ -814,10 +837,7 @@ if MCP_AVAILABLE: return tools def _get_client_ip_from_context() -> Optional[str]: - """ - Extract client_ip from auth context. - Returns None if context not set (caller should handle this as "no IP filtering"). - """ + """Return ``client_ip`` from MCP auth context (set by HTTP/SSE handlers), or None.""" try: auth_user = auth_context_var.get() if auth_user and isinstance(auth_user, MCPAuthenticatedUser): @@ -836,19 +856,15 @@ if MCP_AVAILABLE: Args: user_api_key_auth: The authenticated user's API key info. mcp_servers: Optional list of server names to filter to. - client_ip: Client IP for IP-based access control. If None, falls back to - auth context. Pass explicitly from request handlers for safety. - Note: If client_ip is None and auth context is not set, IP filtering is skipped. - This is intentional for internal callers but may indicate a bug if called - from a request handler without proper context setup. + client_ip: Client IP for IP-based access control. MCP HTTP/SSE handlers set auth context (including ``client_ip``) before MCP work; when this is + ``None``, ``client_ip`` is taken from that context. Callers may still + pass ``client_ip`` explicitly when already computed. """ - # Use explicit client_ip if provided, otherwise try auth context if client_ip is None: client_ip = _get_client_ip_from_context() if client_ip is None: verbose_logger.debug( - "MCP _get_allowed_mcp_servers called without client_ip and no auth context. " - "IP filtering will be skipped. This is expected for internal calls." + "MCP _get_allowed_mcp_servers: client IP unknown; skipping public-internet IP filter." ) allowed_mcp_server_ids = ( @@ -1103,6 +1119,112 @@ if MCP_AVAILABLE: return server_auth_header, extra_headers + async def _merge_gateway_initialize_instructions( + allowed_mcp_servers: List[MCPServer], + user_api_key_auth: Optional[UserAPIKeyAuth], + mcp_auth_header: Optional[str], + mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]], + oauth2_headers: Optional[Dict[str, str]], + raw_headers: Optional[Dict[str, str]], + ) -> Optional[str]: + """Merge ``instructions`` for gateway ``initialize``: YAML/API overrides upstream.""" + if not allowed_mcp_servers: + return None + + _has_oauth2_server = any( + getattr(s, "auth_type", None) == MCPAuth.oauth2 + for s in allowed_mcp_servers + ) + _prefetched_oauth_creds = ( + await _prefetch_oauth_creds_for_user(user_api_key_auth) + if _has_oauth2_server + else {} + ) + + async def _one(server: MCPServer) -> Optional[Tuple[str, str]]: + label = ( + server.alias + or server.server_name + or server.name + or server.server_id + or "mcp" + ) + if server.instructions and server.instructions.strip(): + return (label, server.instructions.strip()) + if server.spec_path: + return None + + server_auth_header, extra_headers = _prepare_mcp_server_headers( + server=server, + mcp_server_auth_headers=mcp_server_auth_headers, + mcp_auth_header=mcp_auth_header, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + ) + if extra_headers is None and server.auth_type == MCPAuth.oauth2: + extra_headers = await _get_user_oauth_extra_headers_from_db( + server, + user_api_key_auth, + prefetched_creds=_prefetched_oauth_creds, + ) + try: + if server.static_headers: + if extra_headers is None: + extra_headers = {} + extra_headers.update(server.static_headers) + stdio_env = global_mcp_server_manager._build_stdio_env( + server, raw_headers + ) + client = await global_mcp_server_manager._create_mcp_client( + server=server, + mcp_auth_header=server_auth_header, + extra_headers=extra_headers, + stdio_env=stdio_env, + ) + text = await client.fetch_upstream_initialize_instructions() + if text and text.strip(): + return (label, text.strip()) + except Exception as e: + verbose_logger.debug( + "MCP gateway: upstream instructions fetch failed for %s: %s", + server.name, + e, + ) + return None + + pairs = await asyncio.gather(*(_one(s) for s in allowed_mcp_servers)) + texts = [p for p in pairs if p is not None] + 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) + + async def _set_mcp_gateway_initialize_instructions_token( + user_api_key_auth: Optional[UserAPIKeyAuth], + mcp_servers: Optional[List[str]], + client_ip: Optional[str], + mcp_auth_header: Optional[str], + mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]], + oauth2_headers: Optional[Dict[str, str]], + raw_headers: Optional[Dict[str, str]], + ) -> contextvars.Token[Optional[str]]: + """Resolve merged gateway ``instructions``; return ContextVar token to reset.""" + allowed = await _get_allowed_mcp_servers( + user_api_key_auth=user_api_key_auth, + mcp_servers=mcp_servers, + client_ip=client_ip, + ) + merged = await _merge_gateway_initialize_instructions( + allowed_mcp_servers=allowed, + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + ) + return _mcp_gateway_initialize_instructions.set(merged) + async def _get_tools_from_mcp_servers( # noqa: PLR0915 user_api_key_auth: Optional[UserAPIKeyAuth], mcp_auth_header: Optional[str], @@ -2670,7 +2792,19 @@ if MCP_AVAILABLE: # Request was fully handled (e.g., DELETE on non-existent session) return - await session_manager.handle_request(scope, receive, send) + _instr_tok = await _set_mcp_gateway_initialize_instructions_token( + user_api_key_auth, + mcp_servers, + _client_ip, + mcp_auth_header, + mcp_server_auth_headers, + oauth2_headers, + raw_headers, + ) + try: + await session_manager.handle_request(scope, receive, send) + finally: + _mcp_gateway_initialize_instructions.reset(_instr_tok) except HTTPException: # Re-raise HTTP exceptions to preserve status codes and details raise @@ -2729,7 +2863,19 @@ if MCP_AVAILABLE: await initialize_session_managers() await asyncio.sleep(0.1) - await sse_session_manager.handle_request(scope, receive, send) + _sse_instr_tok = await _set_mcp_gateway_initialize_instructions_token( + user_api_key_auth, + mcp_servers, + _sse_client_ip, + mcp_auth_header, + mcp_server_auth_headers, + oauth2_headers, + raw_headers, + ) + try: + await sse_session_manager.handle_request(scope, receive, send) + finally: + _mcp_gateway_initialize_instructions.reset(_sse_instr_tok) 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..6ba8d0b68a3 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1137,6 +1137,8 @@ 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 + # Shown to MCP clients in InitializeResult.instructions (optional) + instructions: Optional[str] = None # Stdio-specific fields command: Optional[str] = None args: List[str] = Field(default_factory=list) @@ -1219,6 +1221,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 +1273,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 diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index fce95465b55..a728d912715 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? // MCP InitializeResult.instructions (optional) 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..805494b1854 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -27,6 +27,8 @@ class MCPServer(BaseModel): spec_path: Optional[str] = None auth_type: Optional[MCPAuthType] = None authentication_token: Optional[str] = None + # Optional text returned on MCP `initialize` (InitializeResult.instructions) + instructions: Optional[str] = None mcp_info: Optional[MCPInfo] = None extra_headers: Optional[ List[str] diff --git a/schema.prisma b/schema.prisma index fce95465b55..a728d912715 100644 --- a/schema.prisma +++ b/schema.prisma @@ -289,6 +289,7 @@ model LiteLLM_MCPServerTable { server_name String? alias String? description String? + instructions String? // MCP InitializeResult.instructions (optional) url String? spec_path String? transport String @default("sse")