fix(mcp): key the listed-tool slot by the caller's admission identity and forwarded bearer

The slot a tools/list records for a later tools/call was keyed by (user_id, api_key)
only, so every team-only JWT caller shared one slot and one JWT user acting in two
teams shared a slot; a tools/call then handed pre_mcp_call hooks a description another
caller was served. The slot is now keyed by the hashed key, user, team and organization,
plus the admission credential of a caller admitted with neither a key nor a user.

The caller bearer split the slot only on client-forwarded-token and token-exchange
servers; a legacy delegated oauth2 server (delegate_auth_to_upstream without client
credentials) also forwards it upstream and served a different catalog per bearer into
one slot. The bearer now splits the slot on every server whose egress forwards it
(_consumes_caller_authorization) or exchanges it.
This commit is contained in:
Yucheng He 2026-10-01 16:16:56 -07:00
parent 0219521559
commit ddb672155f
2 changed files with 149 additions and 15 deletions

View file

@ -1221,6 +1221,21 @@ def listed_tools_caller_for(
)
def _admission_identity(
auth: UserAPIKeyAuth, raw_headers: Mapping[str, str] | None
) -> tuple[str | None, str | None, str | None, str | None, str | None]:
"""The admission identity the served catalog is shaped for: the hashed key, user, team and
organization, plus the admission credential (``x-litellm-api-key``, else ``Authorization``) of a
caller admitted with neither a key nor a user."""
keyless: Final = auth.api_key is None and auth.user_id is None
credential: Final = (
_raw_header_value(raw_headers, "x-litellm-api-key") or _raw_header_value(raw_headers, "authorization")
if keyless
else None
)
return auth.api_key, auth.user_id, auth.team_id, auth.org_id, credential
def _format_byok_openapi_auth_header(mcp_server: MCPServer, mcp_auth_header: str) -> str:
"""Format a raw BYOK credential for OpenAPI tool ``Authorization`` injection.
@ -4692,11 +4707,14 @@ class MCPServerManager:
def _listed_tools_identity(self, server: MCPServer, caller: ListedToolsCaller | None) -> str | None:
"""Key the listed-tool cache by every request input that can change the served catalog.
The catalog is guardrail-shaped for the caller's own key (default-on guardrails, key or team
selections and opt-outs), so every keyed caller gets its own slot, on OpenAPI servers too.
Forwarded headers, header-driven stdio env, the caller bearer (forwarded as-is or exchanged as
the OBO subject) and the server-specific auth header also reach upstream and split the slot
further. Only unkeyed listings with none of those share the ``None`` slot.
The catalog is guardrail-shaped for the caller's admission identity (default-on guardrails,
key or team selections and opt-outs), so every admitted caller gets its own slot, keyed by
``_admission_identity``: the hashed key, user, team and organization, plus the admission
credential of a caller admitted with neither a key nor a user (a team-only JWT). Forwarded
headers, header-driven stdio env, the caller bearer on every server whose egress forwards it
(``_consumes_caller_authorization``) or exchanges it as the OBO subject, and the
server-specific auth header also reach upstream and split the slot further. Only unkeyed
listings with none of those share the ``None`` slot.
"""
if caller is None:
return None
@ -4706,19 +4724,18 @@ class MCPServerManager:
stdio_env: Final = None if header_env == self._build_stdio_env(server) else header_env
caller_bearer: Final = (
self._extract_subject_token(caller.oauth2_headers, caller.raw_headers, auth)
if server.is_client_forwarded_token or server.auth_type == MCPAuth.oauth2_token_exchange
if _consumes_caller_authorization(server) or server.auth_type == MCPAuth.oauth2_token_exchange
else None
)
_, digest = self._discovery_key(
server,
auth,
caller.mcp_auth_header,
forwarded,
stdio_env,
caller_bearer,
per_caller=auth is not None,
identity: Final = None if auth is None else _admission_identity(auth, caller.raw_headers)
if not (identity or caller.mcp_auth_header or forwarded or stdio_env or caller_bearer):
return None
material: Final = json.dumps(
(identity, caller.mcp_auth_header, forwarded, stdio_env, caller_bearer),
sort_keys=True,
separators=(",", ":"),
)
return digest
return hashlib.sha256(material.encode()).hexdigest()
@staticmethod
def _forwarded_header_values(

View file

@ -7722,6 +7722,123 @@ class TestMCPServerManager:
assert listed is not None and listed.description == "slot a"
def test_listed_tools_slot_is_split_per_team_for_keyless_callers(self):
"""A team-only JWT admits a caller with neither a key nor a user, so the team keys the slot."""
manager: Final = MCPServerManager()
server: Final = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv")
team_one: Final = ListedToolsCaller(
user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id=None, team_id="team-one")
)
team_two: Final = ListedToolsCaller(
user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id=None, team_id="team-two")
)
manager._record_listed_tools(
server, [MCPTool(name="foo", description="Fetch rows FLAGWORD", inputSchema={})], team_one
)
assert manager.get_listed_tool(server, "foo", team_two) is None
listed: Final = manager.get_listed_tool(server, "foo", team_one)
assert listed is not None and listed.description == "Fetch rows FLAGWORD"
def test_listed_tools_slot_is_split_per_team_for_the_same_keyless_user(self):
"""One JWT user acting in two teams is served two team-shaped catalogs, so each team is a slot."""
manager: Final = MCPServerManager()
server: Final = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv")
alice_in_one: Final = ListedToolsCaller(
user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id="alice", team_id="team-one")
)
alice_in_two: Final = ListedToolsCaller(
user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id="alice", team_id="team-two")
)
manager._record_listed_tools(
server, [MCPTool(name="foo", description="Fetch rows FLAGWORD", inputSchema={})], alice_in_one
)
assert manager.get_listed_tool(server, "foo", alice_in_two) is None
listed: Final = manager.get_listed_tool(server, "foo", alice_in_one)
assert listed is not None and listed.description == "Fetch rows FLAGWORD"
def test_listed_tools_slot_is_split_by_the_admission_bearer_of_keyless_callers_without_a_user(self):
"""Two team-only JWT callers of one team differ only in the JWT they were admitted with, so that
credential keys the slot, on a server that never forwards it."""
manager: Final = MCPServerManager()
server: Final = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv")
alice: Final = ListedToolsCaller(
user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id=None, team_id="team-one"),
raw_headers={"authorization": "Bearer jwt-alice"},
)
bob: Final = ListedToolsCaller(
user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id=None, team_id="team-one"),
raw_headers={"authorization": "Bearer jwt-bob"},
)
manager._record_listed_tools(server, [MCPTool(name="foo", description="alice view", inputSchema={})], alice)
assert manager.get_listed_tool(server, "foo", bob) is None
listed: Final = manager.get_listed_tool(server, "foo", alice)
assert listed is not None and listed.description == "alice view"
@pytest.mark.parametrize(
("server_kwargs", "forwards_bearer"),
[
pytest.param(
{"auth_type": MCPAuth.oauth2, "delegate_auth_to_upstream": True, "oauth2_flow": "authorization_code"},
True,
id="oauth2-delegated-to-upstream",
),
pytest.param({"auth_type": MCPAuth.oauth_delegate}, True, id="oauth-delegate"),
pytest.param({"auth_type": MCPAuth.true_passthrough}, True, id="true-passthrough"),
pytest.param({"auth_type": MCPAuth.oauth2_token_exchange}, True, id="token-exchange"),
pytest.param(
{"auth_type": MCPAuth.none, "extra_headers": ["Authorization"], "oauth_passthrough": True},
True,
id="oauth-passthrough",
),
pytest.param({}, False, id="plain"),
pytest.param(
{"auth_type": MCPAuth.oauth2, "oauth2_flow": "authorization_code"},
False,
id="oauth2-gateway-managed",
),
pytest.param(
{
"auth_type": MCPAuth.oauth2,
"oauth2_flow": "client_credentials",
"delegate_auth_to_upstream": True,
"client_id": "gateway",
"client_secret": "secret",
"token_url": "http://idp/token",
},
False,
id="oauth2-client-credentials",
),
],
)
def test_listed_tools_slot_is_split_by_the_forwarded_bearer_on_servers_that_forward_it(
self, server_kwargs: dict[str, object], forwards_bearer: bool
):
"""Two callers sharing one key but carrying different upstream bearers are served two upstream
catalogs exactly on the servers whose egress forwards or exchanges that bearer."""
manager: Final = MCPServerManager()
server: Final = MCPServer(
**{"server_id": "dg", "name": "dg", "transport": MCPTransport.http, "url": "http://dg", **server_kwargs}
)
caller_a: Final = ListedToolsCaller(
user_api_key_auth=UserAPIKeyAuth(api_key="sk-master"),
raw_headers={"x-litellm-api-key": "Bearer sk-master", "authorization": "Bearer UP-A"},
)
caller_b: Final = ListedToolsCaller(
user_api_key_auth=UserAPIKeyAuth(api_key="sk-master"),
raw_headers={"x-litellm-api-key": "Bearer sk-master", "authorization": "Bearer UP-B"},
)
manager._record_listed_tools(
server, [MCPTool(name="lookup", description="Workspace A lookup FLAGWORD", inputSchema={})], caller_a
)
for_b: Final = manager.get_listed_tool(server, "lookup", caller_b)
assert (for_b is None) is forwards_bearer
for_a: Final = manager.get_listed_tool(server, "lookup", caller_a)
assert for_a is not None and for_a.description == "Workspace A lookup FLAGWORD"
@pytest.mark.asyncio
async def test_call_tool_hands_hooks_the_catalog_the_same_forwarded_headers_listed(self):
manager = MCPServerManager()