mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
e9b8c0fd5c
commit
4b0b326dc2
4 changed files with 116 additions and 28 deletions
|
|
@ -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
|
||||
)
|
||||
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue