diff --git a/litellm/proxy/gateway/mcp/outbound_credentials/resolver.py b/litellm/proxy/gateway/mcp/outbound_credentials/resolver.py index 57802a7fdbe..a88a0a72c1c 100644 --- a/litellm/proxy/gateway/mcp/outbound_credentials/resolver.py +++ b/litellm/proxy/gateway/mcp/outbound_credentials/resolver.py @@ -13,8 +13,9 @@ Implemented: `none`, `passthrough`, `api_key` (shared from config, or per-user / from the injected `CredentialStore`), `authorization_code` (per-user token read from the injected `TokenStore`, refreshed proactively via the `TokenRefresher`), and `client_credentials` (shared service-account token cached in the `ServiceTokenStore`, minted by the -`ClientCredentialsFetcher`). The `token_exchange` and `aws_sigv4` modes are typed stubs that -fail closed until their collaborators are injected. +`ClientCredentialsFetcher`), and `token_exchange` (RFC 8693 OBO: the caller's inbound token is +swapped for an upstream-audience token via the `TokenExchanger`, cached per-user). The +`aws_sigv4` mode is a typed stub that fails closed until its signer is injected. """ from __future__ import annotations @@ -30,6 +31,7 @@ from .clock import Clock from .credential_store import CredentialKey, CredentialStore from .httpx_auth import NoOpAuth, StaticHeaderAuth from .service_token_store import ServiceTokenKey, ServiceTokenStore +from .token_exchanger import TokenExchanger from .token_refresher import TokenRefresher from .token_store import StoredToken, TokenKey, TokenStore from .types import ( @@ -64,6 +66,7 @@ class UpstreamCredentialProvider: clock: Clock, service_token_store: ServiceTokenStore, client_credentials_fetcher: ClientCredentialsFetcher, + token_exchanger: TokenExchanger, ) -> None: self._credential_store = credential_store self._token_store = token_store @@ -71,6 +74,7 @@ class UpstreamCredentialProvider: self._clock = clock self._service_token_store = service_token_store self._client_credentials_fetcher = client_credentials_fetcher + self._token_exchanger = token_exchanger async def resolve( self, subject: Subject, server: ServerSpec @@ -223,12 +227,44 @@ class UpstreamCredentialProvider: ) # best-effort; token is valid anyway return Ok(_bearer(minted)) - # --- arms awaiting their collaborators (typed stubs, fail closed) ---------------------- async def _token_exchange( self, subject: Subject, server: ServerSpec, config: TokenExchangeConfig ) -> Result[httpx.Auth, CredError]: - return _todo(AuthSpecKind.token_exchange) + # OBO: swap the caller's live inbound token for a token bound to server.resource at the + # IdP's exchange endpoint. The inbound token is sent ONLY to the exchanger, never to the + # upstream; the upstream gets the exchanged token. Per-user cache, best-effort like M2M + # (read failure re-exchanges, write failure is ignored) since it is re-exchangeable. + if subject.inbound_token is None: + return Error( + CredError.of_unauthorized( + "token_exchange: no inbound token to exchange; authenticate to the gateway" + ) + ) + key = TokenKey( + tenant_id=subject.tenant_id, + subject_id=subject.subject_id, + server_id=server.server_id, + resource=server.resource, + ) + cached = await self._token_store.get(key) + if isinstance(cached, Ok): + token = cached.ok + if token is not None and not self._is_near_expiry(token): + return Ok(_bearer(token)) + exchanged = await self._token_exchanger.exchange( + config, subject.inbound_token, server.resource + ) + if isinstance(exchanged, Error): + return Error( + exchanged.error + ) # subject_token bad -> 401, config -> 500, down -> 503 + minted = exchanged.ok + await self._token_store.put( + key, minted + ) # best-effort; token is re-exchangeable + return Ok(_bearer(minted)) + # --- arms awaiting their collaborators (typed stubs, fail closed) ---------------------- async def _aws_sigv4( self, subject: Subject, server: ServerSpec, config: AwsSigV4Config ) -> Result[httpx.Auth, CredError]: diff --git a/litellm/proxy/gateway/mcp/outbound_credentials/token_exchanger.py b/litellm/proxy/gateway/mcp/outbound_credentials/token_exchanger.py new file mode 100644 index 00000000000..af0867dfdee --- /dev/null +++ b/litellm/proxy/gateway/mcp/outbound_credentials/token_exchanger.py @@ -0,0 +1,32 @@ +"""The RFC 8693 token-exchange (OBO) grant, isolated behind a port. + +`resolve()` owns the orchestration (cache, expiry, fail closed, and sending the inbound token +ONLY here, never to the upstream); the exchange itself is delegated to this port. The +production body is a hand-rolled RFC 8693 POST over httpx (the SDK ships only the deprecated +RFC 7523 jwt-bearer grant), naming the gateway as `act` and requesting a token bound to +`resource` (RFC 8707) at the discovered token endpoint. Faked in tests. +""" + +from __future__ import annotations + +from typing import Protocol + +from pydantic import SecretStr + +from ..result import Result +from .token_store import StoredToken +from .types import CredError, TokenExchangeConfig + + +class TokenExchanger(Protocol): + """Swaps the caller's `subject_token` for a token bound to `resource`, or fails closed. + + Returns `unauthorized` when the subject_token is invalid/expired (the user must + re-authenticate to the gateway), `misconfigured` when the exchange config is wrong (the + endpoint or IdP does not support exchange), and `upstream_unavailable` when the endpoint is + unreachable. + """ + + async def exchange( + self, config: TokenExchangeConfig, subject_token: SecretStr, resource: str + ) -> Result[StoredToken, CredError]: ... diff --git a/litellm/proxy/gateway/mcp/outbound_credentials/types.py b/litellm/proxy/gateway/mcp/outbound_credentials/types.py index dc4a7ef6c79..21cf6d1597f 100644 --- a/litellm/proxy/gateway/mcp/outbound_credentials/types.py +++ b/litellm/proxy/gateway/mcp/outbound_credentials/types.py @@ -168,13 +168,15 @@ class ClientCredentialsConfig(BaseModel): class TokenExchangeConfig(BaseModel): - """RFC 8693 OBO; swap the caller's live subject_token for an upstream-audience token.""" + """RFC 8693 OBO; swap the caller's live subject_token for a token bound to the upstream's + audience (`server.resource`). The exchange runs at the IdP's token endpoint, discovered + (RFC 8414); only the inbound token type and an optional manual endpoint override are config. + """ model_config = ConfigDict(frozen=True) kind: Literal[AuthSpecKind.token_exchange] = AuthSpecKind.token_exchange - token_exchange_endpoint: str - audience: str subject_token_type: str = "urn:ietf:params:oauth:token-type:access_token" + token_exchange_endpoint: str | None = None class SharedKey(BaseModel): diff --git a/tests/mcp_tests/gateway/test_resolver.py b/tests/mcp_tests/gateway/test_resolver.py index eecff64103d..5ae04953ce3 100644 --- a/tests/mcp_tests/gateway/test_resolver.py +++ b/tests/mcp_tests/gateway/test_resolver.py @@ -32,6 +32,9 @@ from litellm.proxy.gateway.mcp.outbound_credentials.service_token_store import ( ServiceTokenKey, ServiceTokenStore, ) +from litellm.proxy.gateway.mcp.outbound_credentials.token_exchanger import ( + TokenExchanger, +) from litellm.proxy.gateway.mcp.outbound_credentials.token_refresher import ( TokenRefresher, ) @@ -54,6 +57,7 @@ from litellm.proxy.gateway.mcp.outbound_credentials.types import ( ServerSpec, SharedKey, Subject, + TokenExchangeConfig, ) from litellm.proxy.gateway.mcp.result import Error, Ok, Result @@ -124,6 +128,22 @@ class FlakyServiceTokenStore: return Ok(None) +class FakeExchanger: + def __init__(self, result: Result[StoredToken, CredError]) -> None: + self._result = result + self.calls = 0 + self.received_subject_token: str | None = None + self.received_resource: str | None = None + + async def exchange( + self, config: TokenExchangeConfig, subject_token: SecretStr, resource: str + ) -> Result[StoredToken, CredError]: + self.calls += 1 + self.received_subject_token = subject_token.get_secret_value() + self.received_resource = resource + return self._result + + def _provider( *, credential_store: CredentialStore | None = None, @@ -132,6 +152,7 @@ def _provider( clock: Clock | None = None, service_token_store: ServiceTokenStore | None = None, fetcher: ClientCredentialsFetcher | None = None, + token_exchanger: TokenExchanger | None = None, ) -> UpstreamCredentialProvider: return UpstreamCredentialProvider( credential_store=credential_store or InMemoryCredentialStore(), @@ -142,6 +163,8 @@ def _provider( service_token_store=service_token_store or InMemoryServiceTokenStore(), client_credentials_fetcher=fetcher or FakeFetcher(Error(CredError.of_upstream_unavailable("unused"))), + token_exchanger=token_exchanger + or FakeExchanger(Error(CredError.of_upstream_unavailable("unused"))), ) @@ -175,6 +198,10 @@ def _svc_key() -> ServiceTokenKey: return ServiceTokenKey(server_id="s1", resource=RESOURCE) +def _tx_cfg() -> TokenExchangeConfig: + return TokenExchangeConfig() + + def _applied_headers(auth: httpx.Auth) -> httpx.Headers: request = httpx.Request("POST", "https://up.example/mcp") return next(auth.auth_flow(request)).headers @@ -312,11 +339,6 @@ async def test_self_contained_arms_never_read_the_inbound_token(): @pytest.mark.parametrize( "config", [ - { - "kind": "token_exchange", - "token_exchange_endpoint": "https://idp/token", - "audience": "https://up.example", - }, {"kind": "aws_sigv4", "region": "us-east-1"}, ], ) @@ -560,6 +582,125 @@ async def test_client_credentials_cache_write_failure_is_best_effort(): assert _applied_headers(result.ok)["Authorization"] == "Bearer fresh" +async def test_token_exchange_swaps_inbound_for_an_upstream_token(): + # Core invariant: the inbound token goes to the exchanger; the upstream gets the EXCHANGED + # token bound to server.resource, never the inbound one. + subject = Subject(tenant_id="t1", subject_id="u1", inbound_token="user-idp-jwt") + exchanger = FakeExchanger( + Ok(StoredToken(access_token="exchanged", expires_at=NOW + timedelta(hours=1))) + ) + result = await _provider(token_exchanger=exchanger).resolve( + subject, _spec(_tx_cfg()) + ) + assert isinstance(result, Ok) + header = _applied_headers(result.ok)["Authorization"] + assert header == "Bearer exchanged" + assert "user-idp-jwt" not in header # inbound is never forwarded to the upstream + assert ( + exchanger.received_subject_token == "user-idp-jwt" + ) # it went to the exchanger + assert exchanger.received_resource == RESOURCE # bound to server.resource + + +async def test_token_exchange_without_inbound_token_fails_closed(): + result = await _provider().resolve(SUBJECT, _spec(_tx_cfg())) + assert isinstance(result, Error) + assert result.error.tag == "unauthorized" + + +async def test_token_exchange_uses_cached_fresh_token(): + store = InMemoryTokenStore( + { + _token_key(): StoredToken( + access_token="cached", expires_at=NOW + timedelta(hours=1) + ) + } + ) + exchanger = FakeExchanger( + Error(CredError.of_upstream_unavailable("must not be called")) + ) + subject = Subject(tenant_id="t1", subject_id="u1", inbound_token="jwt") + result = await _provider(token_store=store, token_exchanger=exchanger).resolve( + subject, _spec(_tx_cfg()) + ) + assert isinstance(result, Ok) + assert _applied_headers(result.ok)["Authorization"] == "Bearer cached" + assert exchanger.calls == 0 + + +async def test_token_exchange_remints_near_expiry(): + store = InMemoryTokenStore( + { + _token_key(): StoredToken( + access_token="old", expires_at=NOW + timedelta(seconds=30) + ) + } + ) + exchanger = FakeExchanger( + Ok(StoredToken(access_token="new", expires_at=NOW + timedelta(hours=1))) + ) + subject = Subject(tenant_id="t1", subject_id="u1", inbound_token="jwt") + result = await _provider(token_store=store, token_exchanger=exchanger).resolve( + subject, _spec(_tx_cfg()) + ) + assert isinstance(result, Ok) + assert _applied_headers(result.ok)["Authorization"] == "Bearer new" + assert exchanger.calls == 1 + + +async def test_token_exchange_subject_token_invalid_is_unauthorized(): + exchanger = FakeExchanger(Error(CredError.of_unauthorized("invalid subject_token"))) + subject = Subject(tenant_id="t1", subject_id="u1", inbound_token="stale") + result = await _provider(token_exchanger=exchanger).resolve( + subject, _spec(_tx_cfg()) + ) + assert isinstance(result, Error) + assert result.error.tag == "unauthorized" + + +async def test_token_exchange_config_error_is_misconfigured(): + exchanger = FakeExchanger( + Error(CredError.of_misconfigured("IdP lacks token exchange")) + ) + subject = Subject(tenant_id="t1", subject_id="u1", inbound_token="jwt") + result = await _provider(token_exchanger=exchanger).resolve( + subject, _spec(_tx_cfg()) + ) + assert isinstance(result, Error) + assert result.error.tag == "misconfigured" + + +async def test_token_exchange_endpoint_down_is_upstream_unavailable(): + exchanger = FakeExchanger( + Error(CredError.of_upstream_unavailable("exchange endpoint timeout")) + ) + subject = Subject(tenant_id="t1", subject_id="u1", inbound_token="jwt") + result = await _provider(token_exchanger=exchanger).resolve( + subject, _spec(_tx_cfg()) + ) + assert isinstance(result, Error) + assert result.error.tag == "upstream_unavailable" + + +async def test_token_exchange_cache_failure_degrades_to_exchange(): + # Same TokenStore as authorization_code, but token_exchange treats it as an optimization: + # a read/write outage degrades to a fresh exchange rather than failing (token re-exchangeable). + exchanger = FakeExchanger( + Ok(StoredToken(access_token="fresh", expires_at=NOW + timedelta(hours=1))) + ) + subject = Subject(tenant_id="t1", subject_id="u1", inbound_token="jwt") + result = await _provider( + token_store=FailingTokenStore(), token_exchanger=exchanger + ).resolve(subject, _spec(_tx_cfg())) + assert isinstance(result, Ok) + assert _applied_headers(result.ok)["Authorization"] == "Bearer fresh" + + +def test_token_exchange_config_allows_discovery_defaults(): + spec = _spec({"kind": "token_exchange"}) + assert isinstance(spec.config, TokenExchangeConfig) + + async def test_api_key_per_user_store_error_is_upstream_unavailable(): # A store/DB outage on read is distinct from a miss: 503, not the 401/412 of "not found". provider = _provider(credential_store=FailingCredentialStore())