diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 2d13dc0dcf3..49b2c2691c5 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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( 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 b0b40982ade..d73028a6751 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 @@ -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()