diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 8ce7ef2b20c..2656b40ffdb 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -35,6 +35,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credenti ) from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( ConnectionBinding, + ConnectionCredential, EnvelopeIdentity, ) from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credentials import ( @@ -552,49 +553,15 @@ class MCPRequestHandler: scope.pop(CONNECTION_SCOPE_KEY, None) connection_header: Final = headers.get("authorization") if is_connection_credential(connection_header): - if not has_explicit_litellm_key or request_route != "/mcp": - raise HTTPException(status_code=401, detail="A connection credential requires the original MCP key") - targets: Final = MCPRequestHandler._resolve_target_server_names(request_route, mcp_servers) - from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager - - target: Final = ( - global_mcp_server_manager.get_mcp_server_by_name( - targets[0], client_ip=IPAddressUtils.get_mcp_client_ip(request) - ) - if len(targets) == 1 - else None + scope[CONNECTION_SCOPE_KEY] = await MCPRequestHandler._admit_connection_credential( + request=request, + request_route=request_route, + connection_header=connection_header or "", + litellm_api_key=litellm_api_key, + mcp_servers=mcp_servers, + validated_user_api_key_auth=validated_user_api_key_auth, + has_explicit_litellm_key=has_explicit_litellm_key, ) - allowed: Final = await MCPRequestHandler.get_allowed_mcp_servers(validated_user_api_key_auth) - if ( - target is None - or target.server_id not in allowed - or not target.is_gateway_managed_oauth2 - or not target.needs_user_oauth_token - or target.oauth_identity_binding is not None - ): - raise HTTPException(status_code=403, detail="Connection credential does not authorize this MCP server") - expected_binding: Final = ConnectionBinding( - key_hash=hash_token(_get_bearer_token_or_received_api_key(litellm_api_key)), - server_id=target.server_id, - resource=f"{get_request_base_url(request)}/mcp", - ) - connection: Final = open_connection_credential(connection_header or "") - if connection is None: - raise HTTPException( - status_code=401, - detail="Invalid or expired MCP connection credential", - headers=MappingProxyType( - { - "www-authenticate": connection_challenge(request, expected_binding), - "Cache-Control": "no-store", - } - ), - ) - if connection.binding != expected_binding: - raise HTTPException( - status_code=401, detail="Connection credential belongs to a different key or resource" - ) - scope[CONNECTION_SCOPE_KEY] = connection # Leak-defense (single chokepoint): a gateway admission credential (session bearer or bridge # envelope) is NEVER a valid upstream token. Scrub it from EVERY egress context so no @@ -623,6 +590,58 @@ class MCPRequestHandler: raw_headers, ) + @staticmethod + async def _admit_connection_credential( + request: Request, + request_route: str, + connection_header: str, + litellm_api_key: str, + mcp_servers: list[str] | None, + validated_user_api_key_auth: UserAPIKeyAuth, + has_explicit_litellm_key: bool, + ) -> ConnectionCredential: + if not has_explicit_litellm_key or request_route != "/mcp": + raise HTTPException(status_code=401, detail="A connection credential requires the original MCP key") + targets: Final = MCPRequestHandler._resolve_target_server_names(request_route, mcp_servers) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager + + target: Final = ( + global_mcp_server_manager.get_mcp_server_by_name( + targets[0], client_ip=IPAddressUtils.get_mcp_client_ip(request) + ) + if len(targets) == 1 + else None + ) + allowed: Final = await MCPRequestHandler.get_allowed_mcp_servers(validated_user_api_key_auth) + if ( + target is None + or target.server_id not in allowed + or not target.is_gateway_managed_oauth2 + or not target.needs_user_oauth_token + or target.oauth_identity_binding is not None + ): + raise HTTPException(status_code=403, detail="Connection credential does not authorize this MCP server") + expected_binding: Final = ConnectionBinding( + key_hash=hash_token(_get_bearer_token_or_received_api_key(litellm_api_key)), + server_id=target.server_id, + resource=f"{get_request_base_url(request)}/mcp", + ) + connection: Final = open_connection_credential(connection_header) + if connection is None: + raise HTTPException( + status_code=401, + detail="Invalid or expired MCP connection credential", + headers=MappingProxyType( + { + "www-authenticate": connection_challenge(request, expected_binding), + "Cache-Control": "no-store", + } + ), + ) + if connection.binding != expected_binding: + raise HTTPException(status_code=401, detail="Connection credential belongs to a different key or resource") + return connection + @staticmethod def _is_gateway_admission_credential(value: str | None) -> bool: """True when a header value is a gateway admission credential — a session bearer or bridge diff --git a/litellm/proxy/_experimental/mcp_server/contracts.py b/litellm/proxy/_experimental/mcp_server/contracts.py index c3129d171ad..35b6ebc4204 100644 --- a/litellm/proxy/_experimental/mcp_server/contracts.py +++ b/litellm/proxy/_experimental/mcp_server/contracts.py @@ -5,6 +5,7 @@ from datetime import datetime from types import MappingProxyType from typing import Final, Protocol +from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ConnectionCredential from litellm.proxy._types import UserAPIKeyAuth from litellm.types.mcp_server.mcp_server_manager import MCPServer @@ -26,6 +27,7 @@ class OperationContext: raw_headers: Mapping[str, str] | None = field(default=None, repr=False) client_ip: str | None = None mcp_proxy_mode: bool = False + connection_credential: ConnectionCredential | None = field(default=None, repr=False) def __post_init__(self) -> None: object.__setattr__(self, "_caller", copy_caller(self._caller)) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_context.py b/litellm/proxy/_experimental/mcp_server/mcp_context.py index c280bef337e..11325a9f127 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_context.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_context.py @@ -11,8 +11,6 @@ from typing import TYPE_CHECKING, Final if TYPE_CHECKING: from mcp.server.context import ServerRequestContext - from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ConnectionCredential - # The SDK 1.x ``mcp.server.lowlevel.server.request_ctx`` ContextVar was removed in # SDK 2, which hands each request handler a ``ServerRequestContext`` argument # instead. The handlers set this var so downstream helpers (session auth caching, @@ -42,19 +40,3 @@ _mcp_gateway_server_name: Final[ContextVar[str | None]] = ContextVar("_mcp_gatew # Set server-side by the /mcp/proxy route. Never populated from client-supplied headers. _mcp_proxy_mode: Final[ContextVar[bool]] = ContextVar("_mcp_proxy_mode", default=False) - - -def get_connection_credential(server_id: str) -> "ConnectionCredential | None": - - from starlette.requests import Request - - from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ConnectionCredential - - context: Final = get_active_mcp_request_ctx() - request: Final = context.request if context is not None else None - if not isinstance(request, Request): - return None - value: Final = request.scope.get("litellm.mcp.connection_grant") - if not isinstance(value, ConnectionCredential) or value.binding.server_id != server_id: - return None - return value diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 9dbc1588f0f..3db7f60e1b3 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -111,6 +111,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import to_server_spec, to_subject, ) +from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ConnectionCredential from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import ( InvalidatableOAuthTokenStore, ) @@ -4108,6 +4109,7 @@ class MCPServerManager: cred_provider: UpstreamCredentialProvider | None = None, raw_headers: Mapping[str, str] | None = None, client_ip: str | None = None, + connection_credential: ConnectionCredential | None = None, ) -> MCPClient: """ Create an MCPClient instance for the given server. @@ -4133,13 +4135,17 @@ class MCPServerManager: resolved_server: Final = await self.ensure_oauth_metadata_discovered(server) transport: Final = resolved_server.transport or MCPTransport.sse spec = None if transport == MCPTransport.stdio else _to_server_spec_fail_closed(resolved_server) - from litellm.proxy._experimental.mcp_server.mcp_context import get_connection_credential from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import OAuthToken from litellm.proxy._experimental.mcp_server.outbound_credentials.presented_token_store import ( PresentedOAuthTokenStore, ) - connection: Final = get_connection_credential(resolved_server.server_id) + connection: Final = ( + connection_credential + if connection_credential is not None + and connection_credential.binding.server_id == resolved_server.server_id + else None + ) if connection is not None and connection.exp <= int(datetime.datetime.now(datetime.timezone.utc).timestamp()): raise HTTPException(status_code=401, detail="MCP connection credential expired; reconnect") provider: Final = ( @@ -4173,7 +4179,9 @@ class MCPServerManager: sampling_cb = ( _create_sampling_callback( operation_context=OperationContext( - _caller=user_api_key_auth, raw_headers=raw_headers, client_ip=client_ip + _caller=user_api_key_auth, + raw_headers=raw_headers, + client_ip=client_ip, ) ) if resolved_server.allow_sampling @@ -4323,6 +4331,7 @@ class MCPServerManager: user_api_key_auth: UserAPIKeyAuth | None = None, oauth2_headers: dict[str, str] | None = None, client_ip: str | None = None, + connection_credential: ConnectionCredential | None = None, ) -> list[MCPTool]: """ Helper method to get tools from a single MCP server with prefixed names. @@ -4414,6 +4423,7 @@ class MCPServerManager: user_api_key_auth=user_api_key_auth, raw_headers=raw_headers, client_ip=client_ip, + connection_credential=connection_credential, ) ## HANDLE OPENAPI TOOLS @@ -4525,6 +4535,7 @@ class MCPServerManager: add_prefix: bool = True, raw_headers: dict[str, str] | None = None, client_ip: str | None = None, + connection_credential: ConnectionCredential | None = None, ) -> list[Prompt]: try: headers: Final = ( @@ -4547,6 +4558,7 @@ class MCPServerManager: user_api_key_auth=user_api_key_auth, raw_headers=raw_headers, client_ip=client_ip, + connection_credential=connection_credential, ) credential_fingerprint: Final = await client.discovery_auth_fingerprint() key: Final = self._discovery_key( @@ -4571,6 +4583,7 @@ class MCPServerManager: add_prefix: bool = True, raw_headers: dict[str, str] | None = None, client_ip: str | None = None, + connection_credential: ConnectionCredential | None = None, ) -> list[Resource]: try: headers: Final = ( @@ -4593,6 +4606,7 @@ class MCPServerManager: user_api_key_auth=user_api_key_auth, raw_headers=raw_headers, client_ip=client_ip, + connection_credential=connection_credential, ) credential_fingerprint: Final = await client.discovery_auth_fingerprint() key: Final = self._discovery_key( @@ -4617,6 +4631,7 @@ class MCPServerManager: add_prefix: bool = True, raw_headers: dict[str, str] | None = None, client_ip: str | None = None, + connection_credential: ConnectionCredential | None = None, ) -> list[ResourceTemplate]: try: headers: Final = ( @@ -4639,6 +4654,7 @@ class MCPServerManager: user_api_key_auth=user_api_key_auth, raw_headers=raw_headers, client_ip=client_ip, + connection_credential=connection_credential, ) credential_fingerprint: Final = await client.discovery_auth_fingerprint() key: Final = self._discovery_key( @@ -4663,6 +4679,7 @@ class MCPServerManager: extra_headers: dict[str, str] | None = None, raw_headers: dict[str, str] | None = None, client_ip: str | None = None, + connection_credential: ConnectionCredential | None = None, ) -> ReadResourceResult: """Read resource contents from a specific MCP server.""" @@ -4686,6 +4703,7 @@ class MCPServerManager: raw_headers=raw_headers, client_ip=client_ip, user_api_key_auth=user_api_key_auth, + connection_credential=connection_credential, ) return await client.read_resource(url) @@ -4700,6 +4718,7 @@ class MCPServerManager: extra_headers: dict[str, str] | None = None, raw_headers: dict[str, str] | None = None, client_ip: str | None = None, + connection_credential: ConnectionCredential | None = None, ) -> GetPromptResult: """Fetch a specific prompt definition from a single MCP server.""" @@ -4723,6 +4742,7 @@ class MCPServerManager: raw_headers=raw_headers, client_ip=client_ip, user_api_key_auth=user_api_key_auth, + connection_credential=connection_credential, ) get_prompt_request_params: Final = GetPromptRequestParams( @@ -5805,6 +5825,7 @@ class MCPServerManager: user_api_key_auth: UserAPIKeyAuth | None, raw_headers: Mapping[str, str] | None = None, client_ip: str | None = None, + connection_credential: ConnectionCredential | None = None, ) -> CallToolResult: """Call a token_exchange (OBO) tool; on an upstream 401/403 re-mint the token once and retry. @@ -5832,6 +5853,7 @@ class MCPServerManager: user_api_key_auth=user_api_key_auth, raw_headers=raw_headers, client_ip=client_ip, + connection_credential=connection_credential, ) return await retry_client.call_tool(call_tool_params, host_progress_callback=host_progress_callback) @@ -5850,6 +5872,7 @@ class MCPServerManager: hook_extra_headers: dict[str, str] | None = None, user_api_key_auth: UserAPIKeyAuth | None = None, client_ip: str | None = None, + connection_credential: ConnectionCredential | None = None, ) -> CallToolResult: """ Call a regular MCP tool using the MCP client. @@ -5996,6 +6019,7 @@ class MCPServerManager: user_api_key_auth=user_api_key_auth, raw_headers=raw_headers, client_ip=client_ip, + connection_credential=connection_credential, ) call_tool_params: Final = MCPCallToolRequestParams( @@ -6021,6 +6045,7 @@ class MCPServerManager: user_api_key_auth=user_api_key_auth, raw_headers=raw_headers, client_ip=client_ip, + connection_credential=connection_credential, ) tool_call_coro = _obo_call_tool_limited() @@ -6303,6 +6328,7 @@ class MCPServerManager: litellm_logging_obj: "LiteLLMLoggingObj | None" = None, guardrail_context: Mapping[str, object] | None = None, client_ip: str | None = None, + connection_credential: ConnectionCredential | None = None, ) -> CallToolResult: """ Call a tool with the given name and arguments @@ -6434,6 +6460,7 @@ class MCPServerManager: host_progress_callback=host_progress_callback, hook_extra_headers=hook_result.get("extra_headers"), user_api_key_auth=user_api_key_auth, + connection_credential=connection_credential, ) return await self._gather_openapi_tool_tasks(tasks, proxy_logging_obj) diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index fcee3483e15..506372b741a 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -85,6 +85,7 @@ from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( _request_extra_headers, _request_resolved_auth_headers, ) +from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ConnectionCredential from litellm.proxy._experimental.mcp_server.tool_registry import ( global_mcp_tool_registry, ) @@ -250,6 +251,7 @@ async def _dispatch_virtual_mcp_tool( oauth2_headers: dict[str, str] | None = None, raw_headers: dict[str, str] | None = None, mcp_proxy_mode: bool = False, + connection_credential: ConnectionCredential | None = None, ) -> CallToolResult | None: """Handle the mcp_tool_search / mcp_tool_call virtual tools. @@ -297,6 +299,7 @@ async def _dispatch_virtual_mcp_tool( ) try: proxy_result: Final = await handle_mcp_proxy_tool( + connection_credential=connection_credential, name=name, arguments=arguments or {}, # mutable-ok: proxy handler payload user_api_key_dict=user_api_key_auth, @@ -364,6 +367,7 @@ async def _dispatch_virtual_mcp_tool( args: Final = arguments or {} if name == MCP_TOOL_SEARCH_TOOL_NAME: return await handle_mcp_tool_search( + connection_credential=connection_credential, query=TypeAdapter(str).validate_python(args.get("query", "")), top_k=coerce_top_k(args.get("top_k", 5)), user_api_key_dict=user_api_key_auth, @@ -399,6 +403,7 @@ async def _dispatch_virtual_mcp_tool( types.MappingProxyType({"name": args.get("tool_name", ""), "arguments": args.get("arguments") or {}}) ) return await handle_mcp_tool_call( + connection_credential=connection_credential, tool_name=tool_request.name, arguments=tool_request.arguments or {}, user_api_key_dict=user_api_key_auth, @@ -943,6 +948,7 @@ async def _get_tools_from_mcp_servers( request_tags: list[str] | None = None, client_ip: str | None = None, mcp_proxy_mode: bool = False, + connection_credential: ConnectionCredential | None = None, ) -> AggregateToolListing: """ Helper method to fetch tools from MCP servers based on server filtering criteria. @@ -1105,6 +1111,7 @@ async def _get_tools_from_mcp_servers( try: tools: Final = await global_mcp_server_manager._get_tools_from_server( + connection_credential=connection_credential, server=server, mcp_auth_header=server_auth_header, extra_headers=extra_headers, @@ -1230,6 +1237,7 @@ async def _get_prompts_from_mcp_servers( oauth2_headers: dict[str, str] | None = None, raw_headers: dict[str, str] | None = None, client_ip: str | None = None, + connection_credential: ConnectionCredential | None = None, ) -> list[Prompt]: """ Helper method to fetch prompt from MCP servers based on server filtering criteria. @@ -1269,6 +1277,7 @@ async def _get_prompts_from_mcp_servers( try: prompts = await global_mcp_server_manager.get_prompts_from_server( + connection_credential=connection_credential, server=server, user_api_key_auth=user_api_key_auth, mcp_auth_header=server_auth_header, @@ -1298,6 +1307,7 @@ async def _get_resources_from_mcp_servers( oauth2_headers: dict[str, str] | None = None, raw_headers: dict[str, str] | None = None, client_ip: str | None = None, + connection_credential: ConnectionCredential | None = None, ) -> list[Resource]: """Fetch resources from allowed MCP servers.""" @@ -1324,6 +1334,7 @@ async def _get_resources_from_mcp_servers( try: resources = await global_mcp_server_manager.get_resources_from_server( + connection_credential=connection_credential, server=server, user_api_key_auth=user_api_key_auth, mcp_auth_header=server_auth_header, @@ -1351,6 +1362,7 @@ async def _get_resource_templates_from_mcp_servers( oauth2_headers: dict[str, str] | None = None, raw_headers: dict[str, str] | None = None, client_ip: str | None = None, + connection_credential: ConnectionCredential | None = None, ) -> list[ResourceTemplate]: """Fetch resource templates from allowed MCP servers.""" @@ -1377,6 +1389,7 @@ async def _get_resource_templates_from_mcp_servers( try: resource_templates = await global_mcp_server_manager.get_resource_templates_from_server( + connection_credential=connection_credential, server=server, user_api_key_auth=user_api_key_auth, mcp_auth_header=server_auth_header, @@ -1446,6 +1459,7 @@ async def _list_mcp_tools( list_tools_log_source: str | None = None, client_ip: str | None = None, mcp_proxy_mode: bool = False, + connection_credential: ConnectionCredential | None = None, ) -> AggregateToolListing: """ List all available MCP tools. @@ -1464,6 +1478,7 @@ async def _list_mcp_tools( try: listing: Final = await _get_tools_from_mcp_servers( + connection_credential=connection_credential, user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, mcp_servers=mcp_servers, @@ -1493,6 +1508,7 @@ async def _list_mcp_prompts( oauth2_headers: dict[str, str] | None = None, raw_headers: dict[str, str] | None = None, client_ip: str | None = None, + connection_credential: ConnectionCredential | None = None, ) -> list[Prompt]: """ List all available MCP prompts. @@ -1510,6 +1526,7 @@ async def _list_mcp_prompts( managed_prompts = [] try: managed_prompts = await _get_prompts_from_mcp_servers( + connection_credential=connection_credential, user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, mcp_servers=mcp_servers, @@ -1534,12 +1551,14 @@ async def _list_mcp_resources( oauth2_headers: dict[str, str] | None = None, raw_headers: dict[str, str] | None = None, client_ip: str | None = None, + connection_credential: ConnectionCredential | None = None, ) -> list[Resource]: """List all available MCP resources.""" managed_resources: list[Resource] = [] try: managed_resources = await _get_resources_from_mcp_servers( + connection_credential=connection_credential, user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, mcp_servers=mcp_servers, @@ -1563,12 +1582,14 @@ async def _list_mcp_resource_templates( oauth2_headers: dict[str, str] | None = None, raw_headers: dict[str, str] | None = None, client_ip: str | None = None, + connection_credential: ConnectionCredential | None = None, ) -> list[ResourceTemplate]: """List all available MCP resource templates.""" managed_resource_templates: list[ResourceTemplate] = [] try: managed_resource_templates = await _get_resource_templates_from_mcp_servers( + connection_credential=connection_credential, user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, mcp_servers=mcp_servers, @@ -1731,6 +1752,7 @@ async def _list_tools_before_first_call( oauth2_headers: dict[str, str] | None, raw_headers: dict[str, str] | None, client_ip: str | None = None, + connection_credential: ConnectionCredential | None = None, ) -> None: """List ``server`` with the caller's own credentials when it does not yet expose ``tool_name`` here. @@ -1746,6 +1768,7 @@ async def _list_tools_before_first_call( return try: await _get_tools_from_mcp_servers( + connection_credential=connection_credential, user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, mcp_servers=[server.server_id], @@ -1771,9 +1794,11 @@ async def execute_mcp_tool( host_progress_callback: ProgressCallback | None = None, guardrail_context: Mapping[str, object] | None = None, client_ip: str | None = None, + connection_credential: ConnectionCredential | None = None, **kwargs: object, # kwargs-ok: preserves the existing REST and decorated logging call contract ) -> CallToolResult: context: Final = prepare_context( + connection_credential=connection_credential, user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, mcp_server_auth_headers=mcp_server_auth_headers, @@ -1806,6 +1831,7 @@ async def _execute_mcp_tool( host_progress_callback: ProgressCallback | None = None, guardrail_context: Mapping[str, object] | None = None, client_ip: str | None = None, + connection_credential: ConnectionCredential | None = None, **kwargs: Any, ) -> CallToolResult: """ @@ -1865,6 +1891,7 @@ async def _execute_mcp_tool( else strip_known_server_prefix(name, first_call_target) ) await _list_tools_before_first_call( + connection_credential=connection_credential, server=first_call_target, tool_name=first_call_tool_name, allowed_mcp_servers=allowed_mcp_servers, @@ -2059,6 +2086,7 @@ async def _execute_mcp_tool( ######################################################### elif mcp_server: response = await _handle_managed_mcp_tool( + connection_credential=connection_credential, server_name=server_name, name=original_tool_name, arguments=arguments, @@ -2281,6 +2309,7 @@ async def call_mcp_tool( oauth2_headers: dict[str, str] | None = None, raw_headers: dict[str, str] | None = None, client_ip: str | None = None, + connection_credential: ConnectionCredential | None = None, **kwargs: Any, ) -> CallToolResult: """ @@ -2325,6 +2354,7 @@ async def call_mcp_tool( # Delegate to execute_mcp_tool for execution response = await execute_mcp_tool( + connection_credential=connection_credential, name=name, arguments=arguments, allowed_mcp_servers=allowed_mcp_servers, @@ -2363,6 +2393,7 @@ async def mcp_get_prompt( oauth2_headers: dict[str, str] | None = None, raw_headers: dict[str, str] | None = None, client_ip: str | None = None, + connection_credential: ConnectionCredential | None = None, ) -> GetPromptResult: """ Fetch a specific MCP prompt, handling both prefixed and unprefixed names. @@ -2399,6 +2430,7 @@ async def mcp_get_prompt( ) return await global_mcp_server_manager.get_prompt_from_server( + connection_credential=connection_credential, server=server, user_api_key_auth=user_api_key_auth, prompt_name=original_prompt_name, @@ -2419,6 +2451,7 @@ async def mcp_read_resource( oauth2_headers: dict[str, str] | None = None, raw_headers: dict[str, str] | None = None, client_ip: str | None = None, + connection_credential: ConnectionCredential | None = None, ) -> ReadResourceResult: """Read resource contents from upstream MCP servers.""" @@ -2452,6 +2485,7 @@ async def mcp_read_resource( ) return await global_mcp_server_manager.read_resource_from_server( + connection_credential=connection_credential, server=server, user_api_key_auth=user_api_key_auth, url=url, @@ -2506,12 +2540,14 @@ async def _handle_managed_mcp_tool( host_progress_callback: ProgressCallback | None = None, guardrail_context: Mapping[str, object] | None = None, client_ip: str | None = None, + connection_credential: ConnectionCredential | None = None, ) -> CallToolResult: """Handle tool execution for managed server tools""" # Import here to avoid circular import from litellm.proxy.proxy_server import proxy_logging_obj call_tool_result: Final = await global_mcp_server_manager.call_tool( + connection_credential=connection_credential, server_name=server_name, name=name, arguments=arguments, @@ -2576,6 +2612,7 @@ _MCP_CREDENTIAL_REQUEST_FIELDS: Final = frozenset( "mcp_server_auth_headers", "oauth2_headers", "user_api_key_auth", + "connection_credential", } ) @@ -2622,6 +2659,7 @@ async def _execute_handle_list_tools( # Get mcp_servers from context variable verbose_logger.debug("MCP list_tools - Calling _list_mcp_tools") listing: Final = await _list_mcp_tools( + connection_credential=context.connection_credential, user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, mcp_servers=mcp_servers, @@ -2681,6 +2719,7 @@ async def _execute_mcp_server_tool_call( # Inside this try so virtual-tool errors convert to isError # CallToolResult instead of raising out of the protocol handler. virtual_tool_result: Final = await _dispatch_virtual_mcp_tool( + connection_credential=context.connection_credential, name=params.name, arguments=params.arguments, user_api_key_auth=user_api_key_auth, @@ -2729,6 +2768,7 @@ async def _execute_mcp_server_tool_call( data = body_data response: Final = await call_mcp_tool( + connection_credential=context.connection_credential, user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, mcp_servers=mcp_servers, @@ -2822,6 +2862,7 @@ async def _execute_list_prompts( # Get mcp_servers from context variable verbose_logger.debug("MCP list_prompts - Calling _list_prompts") prompts: Final = await _list_mcp_prompts( + connection_credential=context.connection_credential, user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, mcp_servers=mcp_servers, @@ -2856,6 +2897,7 @@ async def _execute_get_prompt( verbose_logger.debug("MCP mcp_server_tool_call - User API Key Auth from context: %s", user_api_key_auth) return await mcp_get_prompt( + connection_credential=context.connection_credential, name=params.name, arguments=params.arguments, user_api_key_auth=user_api_key_auth, @@ -2891,6 +2933,7 @@ async def _execute_list_resources( ) resources: Final = await _list_mcp_resources( + connection_credential=context.connection_credential, user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, mcp_servers=mcp_servers, @@ -2929,6 +2972,7 @@ async def _execute_list_resource_templates( ) resource_templates: Final = await _list_mcp_resource_templates( + connection_credential=context.connection_credential, user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, mcp_servers=mcp_servers, @@ -2962,6 +3006,7 @@ async def _execute_read_resource( ) = context.legacy_auth() read_resource_result: Final = await mcp_read_resource( + connection_credential=context.connection_credential, url=params.uri, user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, @@ -2991,8 +3036,10 @@ def prepare_context( raw_headers: Mapping[str, str] | None = None, client_ip: str | None = None, mcp_proxy_mode: bool = False, + connection_credential: ConnectionCredential | None = None, ) -> OperationContext: return OperationContext( + connection_credential=connection_credential, _caller=user_api_key_auth, mcp_auth_header=mcp_auth_header, mcp_servers=tuple(mcp_servers) if mcp_servers is not None else None, @@ -3060,6 +3107,7 @@ class GatewayOperations: case AuthorizedToolCall(): auth, token, _servers, server_headers, oauth_headers, headers, _client_ip = context.legacy_auth() return await _execute_mcp_tool( + connection_credential=context.connection_credential, name=operation.name, arguments=dict(operation.arguments), # mutable-ok: existing tool hooks own mutable argument data allowed_mcp_servers=list( diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index ce4ad532217..4ea379e2ea7 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -828,8 +828,19 @@ if MCP_AVAILABLE: headers, client_ip, ) = await get_or_extract_auth_context() + credential: Final = ( + ctx.request.scope.get(CONNECTION_SCOPE_KEY) if isinstance(ctx.request, StarletteRequest) else None + ) yield operations.prepare_context( - auth, token, servers, server_headers, oauth_headers, headers, client_ip, _mcp_proxy_mode.get() + auth, + token, + servers, + server_headers, + oauth_headers, + headers, + client_ip, + _mcp_proxy_mode.get(), + connection_credential=credential if isinstance(credential, ConnectionCredential) else None, ) async def handle_list_tools(ctx: ServerRequestContext, params: PaginatedRequestParams) -> ListToolsResult: diff --git a/litellm/proxy/_experimental/mcp_server/tool_search.py b/litellm/proxy/_experimental/mcp_server/tool_search.py index 3650c722103..0004c62a05a 100644 --- a/litellm/proxy/_experimental/mcp_server/tool_search.py +++ b/litellm/proxy/_experimental/mcp_server/tool_search.py @@ -26,6 +26,7 @@ if TYPE_CHECKING: from mcp.types import CallToolResult, Tool from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ConnectionCredential from litellm.proxy._types import UserAPIKeyAuth MCP_TOOL_SEARCH_SETTINGS_KEY: Final[str] = "mcp_tool_search" @@ -462,6 +463,7 @@ async def handle_mcp_tool_search( mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, oauth2_headers: dict[str, str] | None = None, raw_headers: dict[str, str] | None = None, + connection_credential: ConnectionCredential | None = None, ) -> CallToolResult: from litellm.proxy._experimental.mcp_server.operations import ( _list_mcp_tools, @@ -488,6 +490,7 @@ async def handle_mcp_tool_search( else None ) mcp_listing: Final = await _list_mcp_tools( + connection_credential=connection_credential, user_api_key_auth=user_api_key_dict, mcp_servers=mcp_servers, client_ip=client_ip, @@ -513,6 +516,7 @@ async def handle_mcp_proxy_tool( oauth2_headers: dict[str, str] | None = None, # mutable-ok: preserve forwarded headers raw_headers: dict[str, str] | None = None, # mutable-ok: preserve request headers litellm_logging_obj: LiteLLMLoggingObj | None = None, + connection_credential: ConnectionCredential | None = None, ) -> CallToolResult: from fastapi import HTTPException from jsonschema import ValidationError as JsonSchemaValidationError @@ -524,6 +528,7 @@ async def handle_mcp_proxy_tool( ) listing: Final = await _list_mcp_tools( + connection_credential=connection_credential, user_api_key_auth=user_api_key_dict, mcp_servers=mcp_servers, client_ip=client_ip, @@ -579,6 +584,7 @@ async def handle_mcp_proxy_tool( return _text_tool_result(f"Invalid arguments: {exc.message}", is_error=True) return await handle_mcp_tool_call( + connection_credential=connection_credential, tool_name=_mcp_proxy_identity(tool)["tool_name"], arguments=tool_arguments, user_api_key_dict=user_api_key_dict, @@ -606,6 +612,7 @@ async def handle_mcp_tool_call( litellm_logging_obj: LiteLLMLoggingObj | None = None, requested_server_id: str | None = None, guardrail_context: Mapping[str, object] | None = None, + connection_credential: ConnectionCredential | None = None, ) -> CallToolResult: from litellm.proxy._experimental.mcp_server.operations import ( _get_allowed_mcp_servers, @@ -634,6 +641,7 @@ async def handle_mcp_tool_call( raise HTTPException(status_code=403, detail="User not allowed to call this tool.") return await execute_mcp_tool( + connection_credential=connection_credential, name=tool_name, arguments=arguments, allowed_mcp_servers=allowed_mcp_servers, 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 cbfec2bc1ba..977a1ac432c 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 @@ -1080,6 +1080,7 @@ async def test_mcp_get_prompt_success(): extra_headers={"X-Test": "1"}, raw_headers=None, client_ip=None, + connection_credential=None, ) assert result is prompt_result @@ -1143,6 +1144,7 @@ async def test_mcp_read_resource_success(): extra_headers={"X-Test": "1"}, raw_headers=None, client_ip=None, + connection_credential=None, ) assert result is read_result @@ -8986,6 +8988,7 @@ async def test_fire_mcp_tool_call_logging_strips_credentials_from_failure_hook() "mcp_auth_header": "upstream-secret", "mcp_server_auth_headers": {"srv": {"authorization": "Bearer srv-secret"}}, "oauth2_headers": {"authorization": "Bearer oauth-secret"}, + "connection_credential": "connection-secret", "user_api_key_auth": user_auth, } 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 24c4b7a5691..fd8e323bf0b 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 @@ -4070,6 +4070,7 @@ class TestMCPServerManager: user_api_key_auth=None, raw_headers=None, client_ip=None, + connection_credential=None, ) mock_client.list_resource_templates.assert_awaited_once() assert result == expected_templates @@ -14513,7 +14514,7 @@ async def test_connection_grants_follow_current_message_and_never_leak_to_anothe ) reset = active_mcp_request_ctx_var.set(SimpleNamespace(request=request)) try: - client = await manager._create_mcp_client(server) + client = await manager._create_mcp_client(server, connection_credential=credential) sent = await client.prepare_request_auth() assert sent.headers["authorization"] == f"Bearer {token}" store.fetch.assert_not_awaited() @@ -14522,10 +14523,14 @@ async def test_connection_grants_follow_current_message_and_never_leak_to_anothe reset = active_mcp_request_ctx_var.set(SimpleNamespace(request=request)) try: - other_client = await manager._create_mcp_client(other) + other_client = await manager._create_mcp_client(other, connection_credential=credential) other_sent = await other_client.prepare_request_auth() assert other_sent.headers["authorization"] == "Bearer saved-vault-token" assert store.fetch.call_args.args[1] == "other-target" + explicit_client = await manager._create_mcp_client( + server, user_api_key_auth=UserAPIKeyAuth(user_id="independent-caller") + ) + assert (await explicit_client.prepare_request_auth()).headers["authorization"] == "Bearer saved-vault-token" finally: active_mcp_request_ctx_var.reset(reset) saved_client = await manager._create_mcp_client(server) @@ -14570,7 +14575,7 @@ async def test_connection_expiry_between_admission_and_egress_never_uses_vault() reset = active_mcp_request_ctx_var.set(SimpleNamespace(request=request)) try: with pytest.raises(HTTPException) as exc: - await manager._create_mcp_client(server) + await manager._create_mcp_client(server, connection_credential=credential) assert exc.value.status_code == 401 assert "expired" in exc.value.detail store.fetch.assert_not_awaited() @@ -14609,3 +14614,96 @@ async def test_client_sampling_does_not_fill_explicit_context_from_another_ambie assert captured["client_ip"] is None finally: auth_context_var.reset(token) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("method", [ + "tools/list", "tools/call", "prompts/list", "prompts/get", "resources/list", + "resources/templates/list", "resources/read", "virtual/search", "virtual/call", + "proxy/search", "proxy/schema", "proxy/call", +]) +async def test_native_operations_send_only_current_connection_credential(method): + from types import SimpleNamespace + from datetime import timezone + from mcp import types + from mcp.server.context import ServerRequestContext + from pydantic import SecretStr + from starlette.requests import Request + from litellm.proxy._experimental.mcp_server import operations, server as ingress + from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import CONNECTION_SCOPE_KEY + from litellm.proxy._experimental.mcp_server.outbound_credentials import UpstreamCredentialProvider + from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ConnectionBinding, ConnectionCredential + from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import OAuthToken + + store = SimpleNamespace(fetch=AsyncMock(return_value=OAuthToken(access_token="saved-vault-token"))) + manager = MCPServerManager(cred_provider=UpstreamCredentialProvider(oauth_token_store=store)) + target = MCPServer( + server_id="catalog", name="catalog", server_name="catalog", url="https://catalog.example/mcp", + transport="http", auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", + ) + manager.registry = {target.server_id: target} + upstream = _DiscoveryUpstream() + tool = {"name": "example", "description": "Example tool", "inputSchema": {"type": "object", "properties": {}}} + + async def respond(request): + payload = _JSONRPC_ADAPTER.validate_json(request.content) if request.method == "POST" else None + if isinstance(payload, types.JSONRPCRequest) and payload.method in ("tools/list", "tools/call", "prompts/get", "resources/read"): + upstream.requests = (*upstream.requests, (payload.method, request.headers.get("authorization", ""))) + results = { + "tools/list": {"tools": [tool]}, + "tools/call": {"content": [{"type": "text", "text": "executed"}], "isError": False}, + "prompts/get": {"messages": []}, + "resources/read": {"contents": [{"uri": "test://example", "text": "resource body"}]}, + } + return httpx2.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": results[payload.method]}) + return await upstream.respond(request) + + requests = { + "tools/list": types.ListToolsRequest(), + "tools/call": types.CallToolRequest(params=types.CallToolRequestParams(name="catalog-example", arguments={})), + "prompts/list": types.ListPromptsRequest(), + "prompts/get": types.GetPromptRequest(params=types.GetPromptRequestParams(name="catalog-example")), + "resources/list": types.ListResourcesRequest(), + "resources/templates/list": types.ListResourceTemplatesRequest(), + "resources/read": types.ReadResourceRequest(params=types.ReadResourceRequestParams(uri="test://example")), + "virtual/search": types.CallToolRequest(params=types.CallToolRequestParams(name="mcp_tool_search", arguments={"query": "example"})), + "virtual/call": types.CallToolRequest(params=types.CallToolRequestParams(name="mcp_tool_call", arguments={"tool_name": "catalog-example", "arguments": {}})), + "proxy/search": types.CallToolRequest(params=types.CallToolRequestParams(name="search_tools", arguments={"query": "example"})), + "proxy/schema": types.CallToolRequest(params=types.CallToolRequestParams(name="get_tool_schema", arguments={"tool_id": "28a7a373ebe572627a98e19b5347b405"})), + "proxy/call": types.CallToolRequest(params=types.CallToolRequestParams(name="call_tool", arguments={"tool_id": "28a7a373ebe572627a98e19b5347b405", "arguments": {}})), + } + caller = UserAPIKeyAuth(object_permission={"object_permission_id": "test", "mcp_tool_search_enabled": method.startswith("virtual/")}) + auth = (caller, None, ["catalog"], None, None, {}, None) + proxy_reset = ingress._mcp_proxy_mode.set(method.startswith("proxy/")) + try: + with ( + _mcp_upstream(respond), + patch.object(operations, "global_mcp_server_manager", manager), + patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[target])), + patch.object(manager, "get_allowed_mcp_servers", AsyncMock(return_value=["catalog"])), + patch.object(ingress, "get_or_extract_auth_context", AsyncMock(return_value=auth)), + ): + for value in ("first", "second", "unvalidated"): + credential = ConnectionCredential( + kind="connection_access", binding=ConnectionBinding(key_hash="key", server_id="catalog", resource="https://gateway.example/mcp"), + client_id="client", token=SecretStr(value or "unused"), jti=value or "unused", + exp=int(datetime.now(timezone.utc).timestamp()) + 300, + ) if value in ("first", "second") else value + request = Request({"type": "http", "method": "POST", "path": "/mcp", "headers": [], CONNECTION_SCOPE_KEY: credential}) + ctx = ServerRequestContext(session=SimpleNamespace(), lifespan_context={}, protocol_version="2025-06-18", method=requests[method].method, request=request) + start = len(upstream.requests) + async with ingress._legacy_operation_context(ctx, trace=False) as context: + result = await operations.GatewayOperations().execute(requests[method], context) + sent = upstream.requests[start:] + assert sent, result + expected = f"Bearer {value}" if value in ("first", "second") else "Bearer saved-vault-token" + assert {authorization for _, authorization in sent} == {expected} + expected_method = "tools/call" if method in ("virtual/call", "proxy/call") else "tools/list" if method.startswith(("virtual/", "proxy/")) else method + assert expected_method in {name for name, _ in sent} + if isinstance(result, types.CallToolResult): + assert result.is_error is False, result + assert result.content + if value in ("first", "second"): + store.fetch.assert_not_awaited() + finally: + ingress._mcp_proxy_mode.reset(proxy_reset) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_operations.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_operations.py index abb925ddc77..5897782f7be 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_operations.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_operations.py @@ -67,7 +67,7 @@ async def test_legacy_adapter_cleans_context_after_cancelled_operation(): previous_session = server.active_mcp_session_var.get() previous_request = active_mcp_request_ctx_var.get() - request = SimpleNamespace(session=object()) + request = SimpleNamespace(session=object(), request=None) auth = (None, None, None, None, None, None, None) async def cancelled_operation(): @@ -93,7 +93,7 @@ async def test_legacy_adapter_cleans_context_when_trace_setup_fails(): previous_session = server.active_mcp_session_var.get() previous_request = active_mcp_request_ctx_var.get() - request = SimpleNamespace(session=object()) + request = SimpleNamespace(session=object(), request=None) async def enter_operation(): async with server._legacy_operation_context(request, trace=True):