diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index de5281998ce..f7638846a81 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -28,7 +28,7 @@ from contextlib import asynccontextmanager from dataclasses import dataclass, replace from functools import lru_cache from itertools import chain, groupby -from types import MappingProxyType +from types import EllipsisType, MappingProxyType from typing import TYPE_CHECKING, Any, Final, Generic, Literal, TypeAlias, TypedDict, TypeVar, cast from urllib.parse import ParseResult, urlparse @@ -1346,18 +1346,14 @@ async def _resolve_byok_mcp_auth_header( return mcp_auth_header -async def _byok_catalog_auth_header( - mcp_server: MCPServer, - user_api_key_auth: UserAPIKeyAuth | None, +def _catalog_auth_header( mcp_auth_header: str | dict[str, str] | None, + catalog_auth_header: str | dict[str, str] | None | EllipsisType, ) -> str | dict[str, str] | None: - """Keys the caller's catalog slot the way tools/call will look it up; never sent upstream.""" - if not mcp_server.is_byok or mcp_auth_header is not None: - return mcp_auth_header - - from litellm.proxy._experimental.mcp_server.operations import _get_byok_credential - - return await _get_byok_credential(mcp_server, user_api_key_auth) + """The header the client supplied, which keys the caller's catalog slot on both tools/list and + tools/call. A caller that already swapped a stored BYOK credential into ``mcp_auth_header`` passes + the client's value explicitly, since the stored credential must never be read to find the slot.""" + return mcp_auth_header if catalog_auth_header is ... else catalog_auth_header def _client_forwarded_authorization_headers( @@ -4479,6 +4475,8 @@ class MCPServerManager: oauth2_headers: dict[str, str] | None = None, client_ip: str | None = None, proxy_logging_obj: ProxyLogging | None = None, + *, + catalog_auth_header: str | dict[str, str] | None | EllipsisType = ..., ) -> Sequence[MCPTool]: """ Helper method to get tools from a single MCP server with prefixed names. @@ -4486,6 +4484,8 @@ class MCPServerManager: Args: server (MCPServer): The server to query tools from mcp_auth_header: Optional auth header for MCP server + catalog_auth_header: The header the client supplied, keying the caller's catalog slot; + defaults to ``mcp_auth_header`` Returns: List[MCPTool]: List of tools available on the server with prefixed names @@ -4500,7 +4500,7 @@ class MCPServerManager: client = None listed_caller: Final = ListedToolsCaller( user_api_key_auth=user_api_key_auth, - mcp_auth_header=await _byok_catalog_auth_header(server, user_api_key_auth, mcp_auth_header), + mcp_auth_header=_catalog_auth_header(mcp_auth_header, catalog_auth_header), raw_headers=raw_headers, oauth2_headers=oauth2_headers, ) @@ -6627,6 +6627,8 @@ class MCPServerManager: guardrail_context: Mapping[str, object] | None = None, client_ip: str | None = None, wire_compat: WireCompat = WireCompat.LEGACY, + *, + catalog_auth_header: str | None | EllipsisType = ..., ) -> CallToolResult | InputRequiredResult: """ Call a tool with the given name and arguments @@ -6638,6 +6640,8 @@ class MCPServerManager: user_api_key_auth: User authentication mcp_auth_header: MCP auth header (deprecated) mcp_server_auth_headers: Optional dict of server-specific auth headers {server_alias: auth_value} + catalog_auth_header: The header the client supplied, keying the caller's catalog slot; + defaults to ``mcp_auth_header`` as received, before BYOK resolution proxy_logging_obj: Optional ProxyLogging object for hook integration litellm_logging_obj: Optional request logger the guardrail hooks record their evaluations onto, so MCP guardrail activity reaches the @@ -6649,6 +6653,7 @@ class MCPServerManager: """ start_time: Final = datetime.datetime.now() mcp_server: Final = self._resolve_mcp_server_for_tool_call(server_name, name) + client_auth_header: Final = _catalog_auth_header(mcp_auth_header, catalog_auth_header) # Resolved before any hook runs so a missing BYOK credential (401) never # leaves during-hook side effects (audit logging, rate-limit bookkeeping) @@ -6659,7 +6664,7 @@ class MCPServerManager: mcp_auth_header, ) listed_caller: Final = listed_tools_caller_for( - mcp_server, user_api_key_auth, mcp_auth_header, mcp_server_auth_headers, raw_headers, oauth2_headers + mcp_server, user_api_key_auth, client_auth_header, mcp_server_auth_headers, raw_headers, oauth2_headers ) ######################################################### diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 33b61ad310e..be2ca076de8 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -1115,6 +1115,7 @@ async def _get_tools_from_mcp_servers( prefetched_creds=_prefetched_oauth_creds, ) + catalog_auth_header: Final = server_auth_header if server.is_byok and server.auth_type != MCPAuth.oauth2 and server_auth_header is None: server_auth_header = await _get_byok_credential(server, user_api_key_auth) @@ -1131,6 +1132,7 @@ async def _get_tools_from_mcp_servers( user_api_key_auth=user_api_key_auth, oauth2_headers=oauth2_headers, proxy_logging_obj=proxy_logging_obj, + catalog_auth_header=catalog_auth_header, ) filtered_tools = filter_tools_by_allowed_tools(tools, server) @@ -2013,6 +2015,7 @@ async def _execute_mcp_tool( if mcp_server is None: mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name) + client_auth_header: Final = mcp_auth_header if mcp_server: standard_logging_mcp_tool_call["mcp_server_cost_info"] = (mcp_server.mcp_info or {}).get("mcp_server_cost_info") if litellm_logging_obj: @@ -2087,7 +2090,12 @@ async def _execute_mcp_tool( original_tool_name, mcp_server, listed_tools_caller_for( - mcp_server, user_api_key_auth, mcp_auth_header, mcp_server_auth_headers, raw_headers, oauth2_headers + mcp_server, + user_api_key_auth, + client_auth_header, + mcp_server_auth_headers, + raw_headers, + oauth2_headers, ), ), ) @@ -2139,6 +2147,7 @@ async def _execute_mcp_tool( arguments=arguments, user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, + catalog_auth_header=client_auth_header, mcp_server_auth_headers=mcp_server_auth_headers, oauth2_headers=oauth2_headers, raw_headers=raw_headers, @@ -2207,7 +2216,7 @@ async def _execute_mcp_tool( listed_tools_caller_for( prefix_server, user_api_key_auth, - mcp_auth_header, + client_auth_header, mcp_server_auth_headers, raw_headers, oauth2_headers, @@ -2615,8 +2624,12 @@ async def _handle_managed_mcp_tool( guardrail_context: Mapping[str, object] | None = None, client_ip: str | None = None, wire_compat: WireCompat = WireCompat.LEGACY, + *, + catalog_auth_header: str | None, ) -> CallToolResult | InputRequiredResult: - """Handle tool execution for managed server tools""" + """Handle tool execution for managed server tools. ``catalog_auth_header`` is the header the client + supplied, which keys the caller's catalog slot; ``mcp_auth_header`` may already be the resolved + BYOK credential.""" # Import here to avoid circular import from litellm.proxy.proxy_server import proxy_logging_obj @@ -2626,6 +2639,7 @@ async def _handle_managed_mcp_tool( arguments=arguments, user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, + catalog_auth_header=catalog_auth_header, mcp_server_auth_headers=mcp_server_auth_headers, oauth2_headers=oauth2_headers, raw_headers=raw_headers, diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py index f8bf72428aa..00ea7c0c78c 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py @@ -963,6 +963,7 @@ async def test_get_tools_from_mcp_servers(): user_api_key_auth=None, oauth2_headers=None, proxy_logging_obj=None, + catalog_auth_header=None, ): if server.server_id == "server1_id": return [mock_tool_1] diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 8fd931c8b7c..8756ec421ae 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -53,6 +53,7 @@ from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( _obo_retry_applies, _resolve_openapi_tool_auth, _should_strip_caller_authorization, + listed_tools_caller_for, ) from litellm.proxy._types import ( LiteLLM_MCPServerTable, @@ -7322,7 +7323,54 @@ class TestMCPServerManager: assert listed is not None and listed.description == "everyone" @pytest.mark.asyncio - async def test_byok_stored_credential_lists_into_the_slot_tools_call_reads(self): + async def test_byok_listing_never_reads_the_credential_store(self): + """tools/list keys the caller's catalog slot by what the client supplied plus the caller's key. + Resolving the stored BYOK credential for that would fail every REST listing while the DB is + down and would seed a per-worker cache the next tools/call trusts over the store.""" + manager = MCPServerManager() + server = MCPServer( + server_id="byok-cold", + name="byok_cold", + transport=MCPTransport.http, + url="http://byok-cold", + is_byok=True, + auth_type=MCPAuth.api_key, + ) + user = UserAPIKeyAuth(api_key="sk-litellm", user_id="byok-cold-user") + manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager._fetch_tools_with_timeout = AsyncMock( + return_value=[MCPTool(name="turn", description="listed while db down", inputSchema={})] + ) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch( + "litellm.proxy._experimental.mcp_server.db.get_user_credential", + AsyncMock(side_effect=RuntimeError("DB DOWN")), + ), + ): + await manager._get_tools_from_server(server=server, user_api_key_auth=user) + + listed = manager.get_listed_tool(server, "turn", listed_tools_caller_for(server, user, None, None, None, None)) + assert listed is not None and listed.description == "listed while db down" + + @pytest.mark.parametrize( + ("list_header", "call_kwargs"), + [ + pytest.param( + None, {"mcp_auth_header": "stored-secret", "catalog_auth_header": None}, id="execute-mcp-tool" + ), + pytest.param(None, {"mcp_auth_header": None}, id="responses-api"), + pytest.param("Bearer hdr", {"mcp_auth_header": "Bearer hdr"}, id="client-supplied-header"), + ], + ) + @pytest.mark.asyncio + async def test_byok_tools_call_reads_the_slot_the_clients_own_header_listed( + self, list_header: str | None, call_kwargs: dict[str, str | None] + ): + """A REST listing records under the header the client sent (none here). tools/call then swaps the + stored credential in, either before reaching ``call_tool`` (``execute_mcp_tool``) or inside it (the + Responses API), and must still read that slot rather than one keyed by the credential.""" from litellm.proxy._experimental.mcp_server.byok_credential_cache import ( byok_credential_cache_key, cache_byok_credential, @@ -7337,20 +7385,40 @@ class TestMCPServerManager: url="http://byok-catalog", is_byok=True, ) + manager.registry = {"byok-catalog": server} user = UserAPIKeyAuth(api_key="sk-litellm", user_id="byok-user") - manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + mock_client = AsyncMock() + mock_client.call_tool.return_value = MagicMock(spec=CallToolResult, content=[], isError=False) + manager._create_mcp_client = AsyncMock(return_value=mock_client) manager._fetch_tools_with_timeout = AsyncMock( return_value=[MCPTool(name="turn", description="stored cred catalog", inputSchema={})] ) + proxy_logging_obj = MagicMock() + proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) + proxy_logging_obj.during_call_hook = AsyncMock(return_value=None) cache_byok_credential("byok-user", "byok-catalog", "stored-secret") try: - await manager._get_tools_from_server(server=server, user_api_key_auth=user) + await manager._get_tools_from_server(server=server, mcp_auth_header=list_header, user_api_key_auth=user) + listed = manager.get_listed_tool( + server, "turn", listed_tools_caller_for(server, user, list_header, None, None, None) + ) + assert listed is not None and listed.description == "stored cred catalog" + + await manager.call_tool( + server_name="byok_catalog", + name="turn", + arguments={}, + user_api_key_auth=user, + proxy_logging_obj=proxy_logging_obj, + **call_kwargs, + ) finally: byok_credential_cache.delete_cache(byok_credential_cache_key("byok-user", "byok-catalog")) - call_side = ListedToolsCaller(user_api_key_auth=user, mcp_auth_header="stored-secret") - listed = manager.get_listed_tool(server, "turn", call_side) - assert listed is not None and listed.description == "stored cred catalog" + hook_kwargs = proxy_logging_obj._create_mcp_request_object_from_kwargs.call_args.args[0] + assert hook_kwargs["tool_description"] == "stored cred catalog" @pytest.mark.asyncio async def test_byok_supplied_header_lists_without_credential_validation(self): @@ -7401,12 +7469,13 @@ class TestMCPServerManager: ], ) @pytest.mark.asyncio - async def test_byok_listing_keys_the_catalog_by_the_stored_secret_but_never_sends_it_upstream( + async def test_byok_listing_keys_the_catalog_by_the_caller_and_never_touches_the_stored_secret( self, server_auth: dict[str, object] ): - """The stored BYOK secret keys the catalog slot tools/call reads, but tools/list sends upstream - exactly what the caller supplied (nothing here), so the static token, the M2M mint and - MCPJWTSigner all behave as they did before the catalog existed, whatever the auth_type.""" + """The caller's key plus what the caller supplied (nothing here) keys the catalog slot tools/call + reads, even with the stored BYOK secret at hand in the cache, and tools/list sends upstream exactly + what the caller supplied, so the static token, the M2M mint and MCPJWTSigner all behave as they + did before the catalog existed, whatever the auth_type.""" from litellm.proxy._experimental.mcp_server.byok_credential_cache import ( byok_credential_cache_key, cache_byok_credential, @@ -7448,8 +7517,7 @@ class TestMCPServerManager: assert client_kwargs["mcp_auth_header"] is None, client_kwargs assert client_kwargs["extra_headers"] == {"Authorization": "Bearer signed-jwt"} signer_headers.assert_awaited_once() - call_side = ListedToolsCaller(user_api_key_auth=alice, mcp_auth_header="BYOK-ALICE-SECRET") - listed = manager.get_listed_tool(server, "echo", call_side) + listed = manager.get_listed_tool(server, "echo", ListedToolsCaller(user_api_key_auth=alice)) assert listed is not None and listed.description == "listed catalog" @pytest.mark.parametrize(