diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 0c3932d2ba1..ce4928ff83d 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -3638,30 +3638,48 @@ class MCPServerManager: user_api_key_auth: UserAPIKeyAuth | None, raw_headers: Mapping[str, str] | None = None, ) -> None: - """Run the OBO exchange for a caller-supplied subject at the transport edge. + """Mint an exchange-backed server's upstream credential at the transport edge. Single-server routes call this before the MCP session opens, where an HTTP status and ``WWW-Authenticate`` still reach the client. A rejected subject raises the RFC 9728 challenge and any other ``CredError`` maps onto its public HTTP status, so an exchange failure surfaces as a failure instead of the session continuing into an empty tool list. A successful exchange is cached by the exchanger, so the session's list/call reuses it. + + Each mode pre-flights only where it would resolve the subject the session goes on to use, + which is what keeps the pre-flight from reaching a verdict the session would contradict. + ``oauth2_token_exchange`` mints from the caller's inbound bearer, so without one there is + nothing to exchange and the missing-subject case stays the preemptive challenge's job. + ``oauth2_id_jag`` is the mirror image: tool listing resolves it from the identity assertion + captured for this user at SSO login and never from the inbound bearer, so the pre-flight is + faithful exactly when no identity bearer was sent (a LiteLLM key in ``Authorization`` is not one), + and a caller that did send one is passed through + untouched rather than judged against a subject the listing will not use. That store-sourced + case is the one whose missing-assertion 412 and store-outage 503 the session cannot report. + Only OBO has a discovery challenge to raise; ID-JAG's failures are plain statuses whose body + already names what the user has to do, so they map through ``raise_public`` as at egress. """ - if server.auth_type != MCPAuth.oauth2_token_exchange: - return - if not self._extract_bearer_token(oauth2_headers, None): - return - resolved_server: Final = await self.ensure_oauth_metadata_discovered(server) - spec: Final = to_server_spec(resolved_server) - if spec is None or not isinstance(spec.config, TokenExchangeConfig): - return subject_token: Final = self._extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth) - if subject_token is None: + match server.auth_type: + case MCPAuth.oauth2_token_exchange: + if not self._extract_bearer_token(oauth2_headers, None): + return + case MCPAuth.oauth2_id_jag: + if subject_token is not None: + return + case _: + return + resolved_server: Final = await self.ensure_oauth_metadata_discovered(server) + spec: Final = _to_server_spec_fail_closed(resolved_server) + if spec is None or not isinstance(spec.config, (TokenExchangeConfig, IdJagConfig)): + return + if subject_token is None and isinstance(spec.config, TokenExchangeConfig): raise_token_exchange_challenge(resolved_server, root_path=get_server_root_path()) match await self._cred_provider.resolve_credentials(to_subject(user_api_key_auth, subject_token), spec): case Ok(_): return case Error(err): - if err.tag == "unauthorized": + if err.tag == "unauthorized" and isinstance(spec.config, TokenExchangeConfig): raise_token_exchange_challenge( resolved_server, root_path=get_server_root_path(), diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index bb075220530..3d7c947a913 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -3851,15 +3851,15 @@ if MCP_AVAILABLE: raise_token_exchange_challenge(server, root_path=get_server_root_path()) - # token_exchange (OBO) with a subject present: run the exchange here at the transport - # edge, so a rejected subject raises the RFC 9728 challenge (and a gateway fault its - # public status) instead of the session opening and list_tools masking the failure as - # an empty tool list. Gated to single-server routes; the multi-server aggregate keeps - # absorbing per-server auth failures so one bad server cannot 401 the whole connect. + # Exchange-backed modes (token_exchange's OBO mint, id_jag's stored-assertion mint): run + # the exchange here at the transport edge, so a rejected subject raises the RFC 9728 + # challenge and any other failure its public status, instead of the session opening and + # list_tools masking it as an empty tool list. The manager owns which modes pre-flight + # and what each mints from. Gated to single-server routes the key may reach; the + # multi-server aggregate keeps absorbing per-server auth failures so one bad server + # cannot 401 the whole connect. if ( server - and server.auth_type == MCPAuth.oauth2_token_exchange - and oauth2_headers and len(mcp_servers or []) == 1 and server.server_id in frozenset( 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 0fd35e674b7..9a6815a61e5 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 @@ -8201,6 +8201,129 @@ class TestPreemptive401ModeAware: await self._run(delegate, self.LITELLM_KEY_HEADERS, has_stored_token=False) +class TestSingleServerPreflightReachesIdJag: + """The connect-time preflight is what turns a credential failure into an HTTP status the client + can read. An oauth2_id_jag server has to reach it: its subject comes from the assertion stored at + SSO login, so the failure is decided before any IdP call and there is nothing later in the session + that can report it (tools/list degrades to an empty list, tools/call to 'tool not found').""" + + def _id_jag_server(self) -> MCPServer: + return MCPServer( + server_id="id-idjag", + name="idjag", + alias="idjag", + server_name="idjag", + url="https://idjag.test/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2_id_jag, + client_id="gateway-client", + client_secret="gateway-secret", + token_exchange_endpoint="https://org-idp.test/oauth2/token", + id_jag_resource_token_endpoint="https://resource-as.test/oauth2/token", + mcp_info={"server_name": "idjag"}, + ) + + async def _run(self, server: MCPServer, mcp_servers: list[str], preflight: AsyncMock) -> 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 + server_module.global_mcp_server_manager, + "get_mcp_server_by_name", + return_value=server, + ), + patch.object( # test-quality-ok: route wiring must invoke the manager preflight + server_module.global_mcp_server_manager, + "preflight_token_exchange", + preflight, + ), + patch.object( # test-quality-ok: allowed-set resolution needs the DB; the test controls its answer + server_module, "_get_allowed_mcp_servers", AsyncMock(return_value=[server]) + ), + ): + await server_module._raise_preemptive_401_for_unauthenticated_servers( + scope={"type": "http", "method": "POST", "path": "/mcp/idjag", "headers": []}, + mcp_servers=mcp_servers, + oauth2_headers={"Authorization": "Bearer sk-litellm-virtual-key"}, + mcp_server_auth_headers=None, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), + client_ip=None, + ) + + @pytest.mark.asyncio + async def test_id_jag_single_server_route_surfaces_the_preflight_status(self): + """The 412 the preflight raises must propagate out of connect, not be swallowed.""" + server = self._id_jag_server() + preflight = AsyncMock(side_effect=HTTPException(status_code=412, detail="no stored assertion")) + + with pytest.raises(HTTPException) as exc: + await self._run(server, ["idjag"], preflight) + + assert exc.value.status_code == 412 + assert preflight.await_args.kwargs["server"] is server + + @pytest.mark.asyncio + async def test_token_exchange_without_a_bearer_still_challenges_and_never_pre_flights(self): + """The already-shipped OBO path must be untouched by the call site dropping its mode test. + A token_exchange server with no inbound bearer has nothing to exchange, so it still gets the + RFC 9728 discovery challenge from the block above and the preflight is never reached; pushing + a subject-less exchange through the resolver would turn that challenge into some other status + and strand a client that only had to SSO and retry.""" + from litellm.proxy._experimental.mcp_server import server as server_module + + token_exchange = MCPServer( + server_id="id-obo", + name="obo", + alias="obo", + server_name="obo", + url="https://obo.test/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2_token_exchange, + token_exchange_endpoint="https://idp.test/oauth2/token", + client_id="cid", + client_secret="csec", + mcp_info={"server_name": "obo"}, + ) + preflight = AsyncMock() + + with ( + patch.object( # test-quality-ok: route wiring must use the manager's configured server + server_module.global_mcp_server_manager, + "get_mcp_server_by_name", + return_value=token_exchange, + ), + patch.object( # test-quality-ok: route wiring must invoke the manager preflight + server_module.global_mcp_server_manager, + "preflight_token_exchange", + preflight, + ), + pytest.raises(HTTPException) as exc, + ): + await server_module._raise_preemptive_401_for_unauthenticated_servers( + scope={"type": "http", "method": "POST", "path": "/mcp/obo", "headers": []}, + mcp_servers=["obo"], + oauth2_headers=None, + mcp_server_auth_headers=None, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), + client_ip=None, + ) + + assert exc.value.status_code == 401 + headers = exc.value.headers or {} + assert "resource_metadata" in (headers.get("WWW-Authenticate") or headers.get("www-authenticate") or "") + preflight.assert_not_awaited() + + @pytest.mark.asyncio + async def test_id_jag_multi_server_route_still_absorbs_the_failure(self): + """The aggregate contract is unchanged: with more than one target the preflight does not run, + so one server with no stored assertion cannot fail the whole connect.""" + preflight = AsyncMock(side_effect=HTTPException(status_code=412, detail="no stored assertion")) + + await self._run(self._id_jag_server(), ["idjag", "other"], preflight) + + preflight.assert_not_awaited() + + def _make_obo_server(alias: str) -> MCPServer: return MCPServer( server_id=f"id-{alias}", 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 482f779bcb8..02b1a19081a 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 @@ -2616,6 +2616,165 @@ class TestMCPServerManager: await manager.preflight_token_exchange(server=server, oauth2_headers=None, user_api_key_auth=None) assert resolved == ["good-subject"] + def _id_jag_server(self, server_id: str) -> "MCPServer": + return MCPServer( + server_id=server_id, + name=f"{server_id}-server", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2_id_jag, + client_id="gateway-client", + client_secret="gateway-secret", + token_exchange_endpoint="https://org-idp.example/oauth2/token", + id_jag_resource_token_endpoint="https://resource-as.example/oauth2/token", + ) + + @pytest.mark.asyncio + async def test_preflight_id_jag_surfaces_missing_assertion_as_a_plain_412(self): + """ID-JAG's missing/expired-assertion precondition must reach the client as a 412 whose body + names the fix, at the transport edge. Without the preflight the session opens and the caller + gets a 200 with an empty tool list and then 'tool not found', which is not what happened. + 412 is a precondition, not an RFC 9728 discovery challenge, so it carries no + WWW-Authenticate: there is nothing for the client to discover and retry against.""" + from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error + from litellm.proxy._experimental.mcp_server.outbound_credentials.types import CredError + + summary = ( + "ID-JAG requires an IdP identity assertion for this user and none is stored. " + "Sign in through LiteLLM SSO so the gateway captures one." + ) + + class _FakeProvider: + async def resolve_credentials(self, subject, server): + return Error(CredError.of_precondition_required(summary)) + + manager = MCPServerManager(cred_provider=_FakeProvider()) + + with pytest.raises(HTTPException) as exc_info: + await manager.preflight_token_exchange( + server=self._id_jag_server("id-jag-preflight-412"), + oauth2_headers=None, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), + ) + assert exc_info.value.status_code == 412 + assert summary in exc_info.value.detail + assert not (exc_info.value.headers or {}) + + @pytest.mark.asyncio + async def test_preflight_id_jag_surfaces_assertion_store_outage_as_503(self): + """A store outage is the other failure the session would swallow, and it is a different + answer than 412: the user has nothing to fix by signing in again.""" + from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error + from litellm.proxy._experimental.mcp_server.outbound_credentials.types import CredError + + class _FakeProvider: + async def resolve_credentials(self, subject, server): + return Error(CredError.of_upstream_unavailable("assertion store unreachable")) + + manager = MCPServerManager(cred_provider=_FakeProvider()) + + with pytest.raises(HTTPException) as exc_info: + await manager.preflight_token_exchange( + server=self._id_jag_server("id-jag-preflight-503"), + oauth2_headers=None, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), + ) + assert exc_info.value.status_code == 503 + + @pytest.mark.asyncio + async def test_preflight_id_jag_preflights_litellm_key_and_skips_identity_bearer(self): + """ID-JAG preflights when Authorization carries a LiteLLM key, but skips a caller identity + bearer that the session passes through unchanged.""" + from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import ( + StaticHeaderAuth, + ) + from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok + + subjects = [] + + class _FakeProvider: + async def resolve_credentials(self, subject, server): + subjects.append( + ( + subject.subject_id, + subject.inbound_token.get_secret_value() if subject.inbound_token else None, + ) + ) + return Ok(StaticHeaderAuth("Bearer minted-id-jag", header_name="Authorization")) + + manager = MCPServerManager(cred_provider=_FakeProvider()) + caller = UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1") + + await manager.preflight_token_exchange( + server=self._id_jag_server("id-jag-preflight-key"), + oauth2_headers={"Authorization": "Bearer sk-litellm-virtual-key"}, + raw_headers={"authorization": "Bearer sk-litellm-virtual-key"}, + user_api_key_auth=caller, + ) + assert subjects == [("u-1", None)] + + await manager.preflight_token_exchange( + server=self._id_jag_server("id-jag-preflight-identity"), + oauth2_headers={"Authorization": "Bearer caller-idp-id-token"}, + raw_headers={ + "x-litellm-api-key": "Bearer sk-admission-key", + "authorization": "Bearer caller-idp-id-token", + }, + user_api_key_auth=UserAPIKeyAuth(api_key="hashed-key", user_id="u-1"), + ) + assert subjects == [("u-1", None)] + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "server_fields", + [ + {"auth_type": MCPAuth.none}, + {"auth_type": MCPAuth.api_key, "authentication_token": "static-upstream-key"}, + {"auth_type": MCPAuth.bearer_token, "authentication_token": "static-upstream-key"}, + { + "auth_type": MCPAuth.oauth2, + "oauth2_flow": "client_credentials", + "client_id": "cid", + "client_secret": "csec", + "token_url": "https://idp.example.com/token", + }, + {"auth_type": MCPAuth.true_passthrough}, + ], + ) + async def test_preflight_resolves_nothing_for_a_mode_that_does_not_pre_flight(self, server_fields): + """The manager is the only thing deciding which modes pre-flight, so it has to reject every + other mode itself. The single-server call site no longer tests the mode before calling, so a + mode that falls through here would start resolving its credential a second time, at connect, + for flows that never had a connect-time resolution at all.""" + from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error + from litellm.proxy._experimental.mcp_server.outbound_credentials.types import CredError + + calls = [] + + class _FakeProvider: + async def resolve_credentials(self, subject, server): + calls.append(server.server_id) + return Error(CredError.of_misconfigured("the preflight must never get here")) + + manager = MCPServerManager(cred_provider=_FakeProvider()) + server = MCPServer( + server_id="not-pre-flighted", + name="not-pre-flighted-server", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + **server_fields, + ) + + assert ( + await manager.preflight_token_exchange( + server=server, + oauth2_headers={"Authorization": "Bearer sk-litellm-virtual-key"}, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), + ) + is None + ) + assert calls == [] + @pytest.mark.asyncio @pytest.mark.parametrize( "authorization", @@ -2638,7 +2797,9 @@ class TestMCPServerManager: resolved: Final[list[str | None]] = [] class _FakeProvider: - async def resolve_credentials(self, subject: Subject, server: ServerSpec) -> Ok[StaticHeaderAuth, CredError]: + async def resolve_credentials( + self, subject: Subject, server: ServerSpec + ) -> Ok[StaticHeaderAuth, CredError]: resolved.append(subject.inbound_token.get_secret_value() if subject.inbound_token else None) return Ok(StaticHeaderAuth("Bearer MINTED", header_name="Authorization")) @@ -2664,7 +2825,9 @@ class TestMCPServerManager: resolved: Final[list[str | None]] = [] class _FakeProvider: - async def resolve_credentials(self, subject: Subject, server: ServerSpec) -> Ok[StaticHeaderAuth, CredError]: + async def resolve_credentials( + self, subject: Subject, server: ServerSpec + ) -> Ok[StaticHeaderAuth, CredError]: resolved.append(subject.inbound_token.get_secret_value() if subject.inbound_token else None) return Ok(StaticHeaderAuth("Bearer MINTED", header_name="Authorization"))