diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 1bb3cedbd02..30e9d7f833e 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -5971,8 +5971,14 @@ class MCPServerManager: if proxy_logging_obj is None: return hook_result - # Admission credentials are never handed to guardrails as the caller's assertion. - incoming_bearer_token: Final = self._extract_subject_token(None, raw_headers, user_api_key_auth) + inbound_authorization: Final = next( + (v for k, v in (raw_headers or {}).items() if isinstance(k, str) and k.lower() == "authorization"), + "", + ) + incoming_bearer_token: Final = ( + inbound_authorization[len("bearer ") :] if inbound_authorization.lower().startswith("bearer ") else None + ) + incoming_subject_token: Final = self._extract_subject_token(None, raw_headers, user_api_key_auth) pre_hook_kwargs: Final = { "guardrail_context": guardrail_context, @@ -5988,6 +5994,7 @@ class MCPServerManager: ), "user_api_key_hash": (getattr(user_api_key_auth, "api_key_hash", None) if user_api_key_auth else None), "incoming_bearer_token": incoming_bearer_token, + "incoming_subject_token": incoming_subject_token, "headers": logging_safe_mcp_headers(raw_headers), "tool_description": tool.description if tool is not None else None, "tool_input_schema": tool.input_schema if tool is not None else None, @@ -7192,17 +7199,16 @@ class MCPServerManager: return None def get_mcp_server_answering_to(self, name: str, client_ip: str | None = None) -> MCPServer | None: - """The server a scoped ``/mcp/{name}`` connect resolves to, matched the way the router matches - it: case-insensitive over server_id, name and every published prefix form, then the exact - name lookup as the fallback.""" - return next( + """The server a scoped ``/mcp/{name}`` connect resolves to: the alias-first exact lookup, then + the router's case-insensitive prefix match.""" + return self.get_mcp_server_by_name(name, client_ip=client_ip) or next( ( server for server in self.get_filtered_registry(client_ip).values() if server_answers_to_name(server, name) ), None, - ) or self.get_mcp_server_by_name(name, client_ip=client_ip) + ) def get_filtered_registry(self, client_ip: str | None = None) -> dict[str, MCPServer]: """ diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index b73d183db86..b9dcb79dafe 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1715,8 +1715,8 @@ if MCP_AVAILABLE: continue # Caller sign-in: challenge at connect because a tool-call-time 401 is wrapped into a - # JSON-RPC error and the WWW-Authenticate header is lost. Non-OBO gates fire only on a - # single-server connect the key's grant admits. + # JSON-RPC error and the WWW-Authenticate header is lost. OBO keeps its connect gate; + # guardrail-only gates fire only on a single-server connect the key's grant admits. sign_in = caller_sign_in_for(server, user_api_key_auth) if server is not None else None if ( server @@ -1726,7 +1726,7 @@ if MCP_AVAILABLE: ) is None and ( - server.auth_type == MCPAuth.oauth2_token_exchange + (server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers) or await _key_granted_single_server(server, mcp_servers, user_api_key_auth, client_ip) ) ): diff --git a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py index cc2df181574..b33c61951d0 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py +++ b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py @@ -195,7 +195,7 @@ class Agent365Guardrail(CustomGuardrail): return data tool_name: Final = str(data.get("mcp_tool_name") or "") - assertion: Final = entra_assertion(data.get("incoming_bearer_token")) + assertion: Final = entra_assertion(data.get("incoming_subject_token")) if assertion is None: self._handle_caller_fault( data=data, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 8d2642a66f1..48d43519b6a 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -10270,6 +10270,50 @@ class TestOboPreflightScopedToAllowedServers: ) +class TestOboChallengeGateKeepsBaseConnectRules: + """An OBO connect carrying a bearer in oauth2_headers is challenged by the exchange path, not the + preemptive gate, so a multi-server connect or a single-server connect with any bearer at all must + not be refused before the session opens.""" + + LITELLM_KEY_BEARER = {"Authorization": "Bearer sk-1234"} + + async def _run(self, servers: list[MCPServer], mcp_servers: list[str]) -> None: + from litellm.proxy._experimental.mcp_server import server as server_module + + with ( + patch.object( # test-quality-ok: route wiring must use the manager's configured server + mcp_operations.global_mcp_server_manager, + "get_mcp_server_answering_to", + return_value=servers[0], + ), + patch.object( # test-quality-ok: allowed-set resolution needs the DB; the test controls its answer + mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=servers) + ), + ): + await server_module._raise_preemptive_401_for_unauthenticated_servers( + scope={"type": "http", "method": "POST", "path": "/mcp/obo", "headers": []}, + mcp_servers=mcp_servers, + oauth2_headers=self.LITELLM_KEY_BEARER, + mcp_server_auth_headers=None, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-1234", user_id="u-1"), + client_ip=None, + raw_headers={"x-litellm-api-key": "sk-1234", "authorization": "Bearer sk-1234"}, + ) + + @pytest.mark.asyncio + async def test_multi_server_connect_with_any_bearer_is_not_preemptively_challenged(self): + obo = _make_obo_server("obo") + catalog = MCPServer(server_id="id-catalog", name="catalog", alias="catalog", transport=MCPTransport.http) + + await self._run([obo, catalog], ["obo", "catalog"]) + + @pytest.mark.asyncio + async def test_single_obo_connect_with_litellm_key_bearer_still_challenges(self): + with pytest.raises(HTTPException) as exc: + await self._run([_make_obo_server("obo")], ["obo"]) + assert exc.value.status_code == 401 + + @pytest.mark.asyncio async def test_post_mcp_call_guardrails_return_the_rewritten_result(): """The result a post_mcp_call guardrail rewrote must be what the caller sends back.""" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 97b455bfd39..8194d7ab5c5 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -6186,6 +6186,24 @@ class TestMCPServerManager: assert manager._resolve_mcp_server_for_tool_call("zapier", "create_zap") is zapier assert manager._resolve_mcp_server_for_tool_call("other", "create_zap") is other + @pytest.mark.parametrize("gh_first", [True, False], ids=["gh-listed-first", "gh-public-listed-first"]) + def test_answering_to_prefers_alias_over_earlier_prefix_match(self, gh_first): + manager = MCPServerManager() + gh = MCPServer( + server_id="gh-id", name="gh", server_name="gh", transport=MCPTransport.http, auth_type=MCPAuth.oauth2 + ) + gh_public = MCPServer( + server_id="gh-public-id", name="gh_public", server_name="gh_public", alias="gh", transport=MCPTransport.http + ) + manager.registry = ( + {"gh-id": gh, "gh-public-id": gh_public} if gh_first else {"gh-public-id": gh_public, "gh-id": gh} + ) + + assert manager.get_mcp_server_answering_to("gh") is gh_public + assert manager.get_mcp_server_answering_to("gh-public-id") is gh_public + assert manager.get_mcp_server_answering_to("GH_PUBLIC") is gh_public + assert manager.get_mcp_server_answering_to("gh-id") is gh + def test_remove_server_drops_only_its_own_tool_mapping_rows(self): manager = self._manager_with_deepwiki_and_huggingface() @@ -6813,6 +6831,59 @@ class TestMCPServerManager: assert exc_info.value.status_code == 403 + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("raw_headers", "api_key", "expected_bearer", "expected_subject"), + [ + pytest.param( + {"authorization": "Bearer eyJ.x.y"}, + "eyJ.x.y", + "eyJ.x.y", + None, + id="idp-token-as-admission-stays-raw-bearer", + ), + pytest.param( + {"authorization": "Bearer sk-1234"}, + "sk-1234", + "sk-1234", + None, + id="litellm-key-as-bearer-is-not-a-subject", + ), + pytest.param( + {"x-litellm-api-key": "sk-1234", "authorization": "Bearer eyJ.x.y"}, + "sk-1234", + "eyJ.x.y", + "eyJ.x.y", + id="key-admission-plus-idp-bearer-subject", + ), + ], + ) + async def test_pre_call_tool_check_separates_raw_bearer_from_subject( + self, raw_headers, api_key, expected_bearer, expected_subject + ): + manager = MCPServerManager() + server = MCPServer( + server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv", allowed_tools=None + ) + proxy_logging = MagicMock() + proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) + proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging.pre_call_hook = AsyncMock(return_value=None) + + await manager.pre_call_tool_check( + server_name="srv", + name="turn", + arguments={}, + user_api_key_auth=UserAPIKeyAuth(api_key=api_key, user_id="u"), + proxy_logging_obj=proxy_logging, + server=server, + raw_headers=raw_headers, + ) + + kwargs: Final = proxy_logging._create_mcp_request_object_from_kwargs.call_args.args[0] + assert kwargs["incoming_bearer_token"] == expected_bearer + assert kwargs["incoming_subject_token"] == expected_subject + @pytest.mark.asyncio async def test_check_tool_permission_for_key_team_allows_permitted_tool(self): """ diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py index cfd6745308d..47f1d9f5f0c 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py @@ -210,7 +210,7 @@ def _mcp_data(**overrides: Any) -> dict: "mcp_tool_name": "send_email", "mcp_arguments": {"to": "user@example.com", "body": "hello"}, "mcp_server_name": "outlook_mcp", - "incoming_bearer_token": FAKE_ASSERTION, + "incoming_subject_token": FAKE_ASSERTION, "metadata": {"headers": {"mcp-session-id": "sess-123"}}, } data.update(overrides) @@ -699,17 +699,27 @@ class TestUnreachableFallback: handler: Final = FakeHandler([]) guardrail: Final = _make_guardrail(handler) with pytest.raises(HTTPException) as exc_info: - await _run(guardrail, _mcp_data(incoming_bearer_token=None)) + await _run(guardrail, _mcp_data(incoming_subject_token=None)) assert exc_info.value.status_code == 401 assert handler.calls == [] + @pytest.mark.asyncio + async def test_raw_bearer_without_subject_token_is_no_bearer(self): + exchanger: Final = StubTokenExchanger() + handler: Final = FakeHandler([]) + guardrail: Final = _make_guardrail(handler, exchanger=exchanger) + with pytest.raises(HTTPException) as exc_info: + await _run(guardrail, _mcp_data(incoming_subject_token=None, incoming_bearer_token=FAKE_ASSERTION)) + assert exc_info.value.status_code == 401 + assert exchanger.calls == [] + @pytest.mark.asyncio async def test_non_jwt_bearer_token_fail_closed(self): exchanger: Final = StubTokenExchanger() handler: Final = FakeHandler([]) guardrail: Final = _make_guardrail(handler, exchanger=exchanger) with pytest.raises(HTTPException) as exc_info: - await _run(guardrail, _mcp_data(incoming_bearer_token="sk-litellm-virtual-key")) + await _run(guardrail, _mcp_data(incoming_subject_token="sk-litellm-virtual-key")) assert exc_info.value.status_code == 401 assert exchanger.calls == [] @@ -717,7 +727,7 @@ class TestUnreachableFallback: async def test_missing_bearer_token_blocks_even_fail_open(self): handler: Final = FakeHandler([]) guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_open") - data: Final = _mcp_data(incoming_bearer_token=None) + data: Final = _mcp_data(incoming_subject_token=None) with pytest.raises(HTTPException) as exc_info: await _run(guardrail, data) assert exc_info.value.status_code == 401 @@ -871,7 +881,7 @@ class TestOboTokenCache: handler: Final = FakeHandler([_allow_response(), _allow_response()]) guardrail: Final = _make_guardrail(handler, exchanger=exchanger) await _run(guardrail, _mcp_data()) - await _run(guardrail, _mcp_data(incoming_bearer_token=other_assertion)) + await _run(guardrail, _mcp_data(incoming_subject_token=other_assertion)) assert len(exchanger.calls) == 2 assert handler.calls[1].headers["Authorization"] == "Bearer token-b"