fix(mcp): pre-flight the ID-JAG credential at the transport edge (#35392)

This commit is contained in:
Yassin Kortam 2026-09-03 18:39:37 -07:00 • committed by GitHub
parent 7ae352e5cf
commit b7f53ce9a9
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 324 additions and 20 deletions

View file

@ -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(),

View file

@ -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(

View file

@ -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}",

View file

@ -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"))