fix(mcp): key the BYOK catalog slot by the client's header, not the stored credential

tools/list resolved the stored BYOK credential to pick the caller's catalog slot, which read the
credential store before the classified try block. With Postgres down and a cold per-worker cache
that made every REST tools/list on an is_byok server fail with tools=[] and no upstream call, and
the read seeded the per-worker cache (including a negative entry), so a tools/call on another
worker after a store, rotate or revoke on this one kept using the stale value.

The slot is now keyed by what the client supplied plus the caller's hashed key, on both sides.
_get_tools_from_server and call_tool take a keyword-only catalog_auth_header that defaults to
mcp_auth_header as received (the default is the builtin Ellipsis so it survives a module reload).
The /mcp fan-out and execute_mcp_tool, which swap the resolved credential into mcp_auth_header,
pass the client's value explicitly. What goes upstream is unchanged. _byok_catalog_auth_header is
gone.
This commit is contained in:
Yucheng He 2026-10-01 02:37:50 -07:00
parent e9b8c0fd5c
commit 4b0b326dc2
4 changed files with 116 additions and 28 deletions

View file

@ -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
)
#########################################################

View file

@ -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,

View file

@ -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]

View file

@ -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(