mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(mcp): pre-flight the ID-JAG credential at the transport edge (#35392)
This commit is contained in:
parent
7ae352e5cf
commit
b7f53ce9a9
4 changed files with 324 additions and 20 deletions
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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}",
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue