diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index d55eb3ac014..a44ca342d82 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -1169,6 +1169,83 @@ async def list_user_oauth_credentials( return results +# The user's IdP (e.g. Okta) grant for delegated on-behalf-of exchange is a THIRD credential kind in +# this table, distinct from a BYOK secret and a per-upstream-server oauth2 token. It is tagged with its +# own payload ``type`` so the oauth2-only readers (``get_user_oauth_credential`` / +# ``list_user_oauth_credentials``, which gate on ``type == "oauth2"``) never surface it as a connected +# server, and it is keyed by the IdP (in ``server_id``) rather than an upstream server, so one grant +# serves every token_exchange upstream that IdP fronts. +_IDP_GRANT_TYPE = "idp_grant" + + +def _decode_idp_grant_payload(stored: str) -> dict[str, Any] | None: + """Return the decoded payload iff ``stored`` holds an IdP grant (``type == "idp_grant"``).""" + decoded = _decode_user_credential(stored) + if decoded is None: + return None + try: + parsed = json.loads(decoded) + except (ValueError, TypeError): + return None + if isinstance(parsed, dict) and parsed.get("type") == _IDP_GRANT_TYPE: + return parsed + return None + + +async def store_user_idp_grant( + prisma_client: PrismaClient, + user_id: str, + idp_key: str, + access_token: str, + refresh_token: str | None = None, + expires_in: int | None = None, + scopes: list[str] | None = None, +) -> None: + """Persist a user's IdP grant (for delegated OBO exchange), keyed by IdP in ``server_id``. + + Stored as a ``type: "idp_grant"`` payload so the oauth2-only readers never mistake it for a + connected upstream server. The refresh path overwrites the same row with the rotated token. + """ + expires_at: str | None = None + if expires_in is not None: + expires_at = (datetime.now(timezone.utc) + timedelta(seconds=expires_in)).isoformat() + + payload: dict[str, Any] = { + "type": _IDP_GRANT_TYPE, + "access_token": access_token, + "connected_at": datetime.now(timezone.utc).isoformat(), + } + if refresh_token: + payload["refresh_token"] = refresh_token + if expires_at: + payload["expires_at"] = expires_at + if scopes: + payload["scopes"] = scopes + + encoded = encrypt_value_helper(json.dumps(payload)) + await MCPUserCredentialsRepository(prisma_client).table.upsert( + where={"user_id_server_id": {"user_id": user_id, "server_id": idp_key}}, + data={ + "create": {"user_id": user_id, "server_id": idp_key, "credential_b64": encoded}, + "update": {"credential_b64": encoded}, + }, + ) + + +async def get_user_idp_grant( + prisma_client: PrismaClient, + user_id: str, + idp_key: str, +) -> dict[str, Any] | None: + """Return the decoded IdP-grant payload for a user+IdP pair, or None.""" + row = await MCPUserCredentialsRepository(prisma_client).table.find_unique( + where={"user_id_server_id": {"user_id": user_id, "server_id": idp_key}} + ) + if row is None: + return None + return _decode_idp_grant_payload(row.credential_b64) + + def _decrypted_credential_field(creds: Dict[str, object], field: str) -> object: """Return one credential field decrypted with the global salt key; non-string and legacy plaintext values come back unchanged (decrypt_value_helper returns the original on failure).""" diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 115ff2e492c..41b985eb48b 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -74,6 +74,9 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import to_server_spec, to_subject, ) +from litellm.proxy._experimental.mcp_server.outbound_credentials.idp_subject_provider import ( + build_idp_subject_source, +) from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import ( InvalidatableOAuthTokenStore, ) @@ -1009,6 +1012,7 @@ class MCPServerManager: self._cred_provider = cred_provider or UpstreamCredentialProvider( oauth_token_store=self._per_user_oauth_token_store, token_exchanger=build_token_exchanger(), + idp_subject_source=build_idp_subject_source(), ) self.registry: dict[str, MCPServer] = {} self.config_mcp_servers: dict[str, MCPServer] = {} diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py index d244455b09b..5cb3ad6bec5 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py @@ -43,7 +43,10 @@ def to_subject(user_api_key_auth: Optional[UserAPIKeyAuth], subject_token: Optio an empty subject rather than share one credential slot across callers. A validated delegation assertion (UserAPIKeyAuth.delegated_user_id, stamped at MCP admission after the consent check) replaces the credential subject so per-user upstream credentials resolve as the delegated - user; admission, permissions, and attribution stay on the calling key. + user; admission, permissions, and attribution stay on the calling key. The same value is + surfaced on ``Subject.delegated_user_id`` as the delegation marker: it tells the token_exchange + arm that ``inbound_token`` is the agent's admission credential (not the subject's own token), so + the arm sources the delegated user's IdP grant instead of exchanging the agent's bearer. """ inbound = SecretStr(subject_token) if subject_token else None if user_api_key_auth is None: @@ -52,6 +55,7 @@ def to_subject(user_api_key_auth: Optional[UserAPIKeyAuth], subject_token: Optio tenant_id=user_api_key_auth.org_id or user_api_key_auth.team_id or "", subject_id=user_api_key_auth.delegated_user_id or user_api_key_auth.user_id or "", inbound_token=inbound, + delegated_user_id=user_api_key_auth.delegated_user_id, ) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/idp_subject_provider.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/idp_subject_provider.py new file mode 100644 index 00000000000..5898b74ce32 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/idp_subject_provider.py @@ -0,0 +1,151 @@ +"""Composition root for the delegated-OBO IdP subject-token source (Path B). + +Wires the pure ``StoredIdpGrantSource`` to its runtime edges: the DB read/persist of the user's stored +IdP grant (a distinct ``idp_grant``-typed credential in the per-user MCP credential store, keyed by IdP +rather than upstream server, so the oauth2-only server listings never surface it) and the httpx +refresh_token POST against the IdP. The token-endpoint POST is the same OAuth2 helper the +authorization_code refresh uses, reused here rather than re-implementing the untyped-httpx boundary. +``store_user_idp_grant`` is the store-back path the first-time consent flow calls to persist the +captured grant; the consent UI itself is out of scope here. Nothing reads a runtime global at build +time (prisma/httpx are acquired per call), so no lazy wrapper is needed, mirroring +``build_token_exchanger``. +""" + +from __future__ import annotations + +from litellm._logging import verbose_logger +from litellm.proxy._experimental.mcp_server.outbound_credentials.idp_subject_source import ( + IdpGrantRefresher, + PersistIdpGrant, + StoredIdpGrantSource, + TokenEndpointPost, + idp_grant_key, +) +from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import ( + CachedOAuthTokenStore, + OAuthToken, + TokenStoreUnavailable, +) +from litellm.proxy._experimental.mcp_server.outbound_credentials.v2_token_store import ( + CredentialReader, + V2PerUserTokenStore, +) + +# A grant with no declared expiry is cached for this long; one with an expiry is cached until then. +# Mirrors the authorization_code store's default so a warm grant avoids a DB read per delegated call. +_GRANT_READ_CACHE_TTL_SECONDS = 300.0 + + +async def _read_credential(user_id: str, idp_key: str) -> dict[str, object] | None: + from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 # lazy: avoids import cycle + get_user_idp_grant, + ) + from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # lazy: runtime global + + if prisma_client is None: + raise TokenStoreUnavailable("Database not connected") + return await get_user_idp_grant(prisma_client, user_id, idp_key) + + +async def _persist_credential( + user_id: str, + idp_key: str, + access_token: str, + refresh_token: str | None, + expires_in: int | None, + scopes: tuple[str, ...] | None, +) -> None: + from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 # lazy: avoids import cycle + store_user_idp_grant, + ) + from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # lazy: runtime global + + if prisma_client is None: + # Symmetric with _read_credential (which raises TokenStoreUnavailable): make the skipped + # persist observable, since a refresh that rotated the IdP refresh_token then failed to save + # it strands the user until re-consent. + verbose_logger.warning("MCP IdP grant persist skipped: database not connected; a rotated grant may be lost") + return + await store_user_idp_grant( + prisma_client=prisma_client, + user_id=user_id, + idp_key=idp_key, + access_token=access_token, + refresh_token=refresh_token, + expires_in=expires_in, + scopes=list(scopes) if scopes else None, + ) + + +async def _post_token_endpoint(url: str, form: dict[str, str], headers: dict[str, str]) -> dict[str, object] | None: + # get_async_httpx_client's signature carries an untyped param, so the imported symbol reads as + # partially unknown; the httpx.Response JSON body is Any and the refresher validates each field, so + # the untyped boundary is contained here. A failed refresh is a miss, not a 500 (matching the + # authorization_code refresher), so any error becomes None. + from litellm.llms.custom_httpx.http_handler import ( # noqa: PLC0415 # lazy import; avoids cycle + get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # httpx handler untyped + ) + from litellm.types.llms.custom_http import httpxSpecialProvider # noqa: PLC0415 # lazy import + + request_headers = {"Accept": "application/json", **headers} + try: + client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check) + response = await client.post(url, headers=request_headers, data=form) # pyright: ignore[reportUnknownMemberType] # AsyncHTTPHandler.post signature is untyped + if response is None: + return None + response.raise_for_status() + body: dict[str, object] = response.json() # pyright: ignore[reportAny] # untyped JSON body, validated by the refresher + except Exception as exc: # noqa: BLE001 # any IdP/transport error is a refresh miss, not a 500 + verbose_logger.warning("MCP IdP grant refresh request failed: %s", exc) + return None + else: + return body + + +def build_idp_subject_source( + *, + read_credential: CredentialReader = _read_credential, + token_endpoint: TokenEndpointPost = _post_token_endpoint, + persist: PersistIdpGrant = _persist_credential, +) -> StoredIdpGrantSource: + # An expiry-aware cache in front of the DB read, mirroring the authorization_code chain + # (Cached(V2PerUserTokenStore)), so a warm grant is served without a DB round-trip on every + # delegated call; the config-aware refresh stays in the source. Built once so the cache is shared + # across requests. Collaborators are injectable for tests. + cached_read = CachedOAuthTokenStore( + V2PerUserTokenStore(read_credential), default_ttl_seconds=_GRANT_READ_CACHE_TTL_SECONDS + ) + + async def read_grant(user_id: str, idp_key: str) -> OAuthToken | None: + try: + return await cached_read.fetch(user_id, idp_key) + except TokenStoreUnavailable: + return None + + refresher = IdpGrantRefresher(token_endpoint, persist) + return StoredIdpGrantSource(read_grant, refresher.refresh) + + +async def capture_user_idp_grant( + user_id: str, + token_exchange_endpoint: str, + access_token: str, + *, + refresh_token: str | None = None, + expires_in: int | None = None, + scopes: tuple[str, ...] | None = None, +) -> None: + """Persist a user's IdP grant for on-behalf-of exchange, keyed by the IdP (the AS token endpoint). + + The store-back path the first-time consent flow calls once it has captured the user's IdP grant + (``authorization_code + offline_access`` against the IdP). One grant serves every ``token_exchange`` + upstream that IdP fronts, since it is keyed by the IdP endpoint rather than an upstream server. + """ + await _persist_credential( + user_id, + idp_grant_key(token_exchange_endpoint), + access_token, + refresh_token, + expires_in, + scopes, + ) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/idp_subject_source.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/idp_subject_source.py new file mode 100644 index 00000000000..340162d8c68 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/idp_subject_source.py @@ -0,0 +1,218 @@ +"""Source the delegated user's live IdP subject token for an OBO (token_exchange) mint. + +Path B of the agent-delegation flow: at consent the user's IdP (e.g. Okta) grant is captured and +stored; at runtime the ``token_exchange`` arm needs the user's *own* subject token, not the agent's +admission bearer, to feed into the RFC 8693 exchange. ``StoredIdpGrantSource`` reads the stored grant +keyed by ``(user, idp)``, refreshes it to a live access token when it has expired, and returns that +token. It returns ``None`` when the user has no usable grant (never consented, or an expired grant with +no refresh_token) so the arm fails closed with a 401 that tells the client to authenticate the *user* +with the IdP; the agent's own token is never a fallback (the escalation the flow forbids). + +The IdP is identified by its token endpoint, so one grant serves every ``token_exchange`` upstream +fronted by the same authorization server (keyed by IdP, not by upstream server). The DB read and the +IdP refresh POST are injected as edges (this module stays v1-free and deals only in ``OAuthToken``); +the refresh is single-flighted per ``(user, idp)`` so concurrent callers collapse to one refresh, +which also avoids invalidating each other's grant under IdP refresh-token rotation. +""" + +from __future__ import annotations + +import time +from collections.abc import Awaitable, Callable +from typing import Protocol + +from litellm._logging import verbose_logger +from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import ( + TokenEndpointAuthConfigError, + build_token_endpoint_client_auth, +) +from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import ( + InProcessRefreshCoordinator, + OAuthToken, + RefreshCoordinator, +) +from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( + TokenExchangeConfig, +) + +# Prefix the (user, idp) credential-store key so an IdP grant row is not mistaken for a +# per-upstream-server credential row that shares the same table, and does not collide with a normal +# server_id (a UUID or hostname) absent a deliberately pathological "idp::"-prefixed one. The row is +# additionally tagged with a distinct payload type, so the oauth2 readers ignore it regardless. +_IDP_GRANT_KEY_PREFIX = "idp::" + +_EXPIRY_SKEW_SECONDS = 60.0 + +# Reads the stored IdP grant for (user_id, idp_key) as an OAuthToken, or None when the user has none. +ReadIdpGrant = Callable[[str, str], Awaitable["OAuthToken | None"]] +# Refreshes the grant against the IdP (refresh_token grant) and persists the result, returning the +# fresh token or None when it cannot be refreshed. Threads the config so the endpoint and client +# credentials come from the calling server's token_exchange config. +RefreshIdpGrant = Callable[[str, str, TokenExchangeConfig, "OAuthToken"], Awaitable["OAuthToken | None"]] +# POSTs an OAuth form to a token endpoint and returns the parsed JSON body, or None on any failure. +TokenEndpointPost = Callable[[str, "dict[str, str]", "dict[str, str]"], Awaitable["dict[str, object] | None"]] + + +class PersistIdpGrant(Protocol): + async def __call__( + self, + user_id: str, + idp_key: str, + access_token: str, + refresh_token: str | None, + expires_in: int | None, + scopes: tuple[str, ...] | None, + ) -> None: ... + + +def _parse_expires_in(raw: object) -> int | None: + if isinstance(raw, bool): + return None + if isinstance(raw, (int, float)): + return int(raw) + if isinstance(raw, str): + try: + return int(float(raw)) + except ValueError: + return None + return None + + +def idp_grant_key(token_exchange_endpoint: str) -> str: + """The (user, idp) storage key's idp component, derived from the AS token endpoint. + + All ``token_exchange`` servers fronted by the same authorization server share one endpoint, so + they share one grant. Normalized (trailing slash stripped) so trivially different spellings of the + same endpoint resolve to one key. + """ + return f"{_IDP_GRANT_KEY_PREFIX}{token_exchange_endpoint.rstrip('/')}" + + +class IdpSubjectTokenSource(Protocol): + """Produces the delegated user's live IdP subject token for a ``token_exchange`` mint, or None.""" + + async def subject_token(self, user_id: str, config: TokenExchangeConfig) -> str | None: ... + + +class StoredIdpGrantSource: + """Reads the user's stored IdP grant and refreshes it to a live subject token when expired. + + The DB read and the refresh POST are injected so the pure logic (expiry, single-flight, fail-closed + on miss) is testable without I/O. The clock is injected for deterministic expiry in tests. + """ + + def __init__( + self, + read_grant: ReadIdpGrant, + refresh_grant: RefreshIdpGrant, + *, + coordinator: RefreshCoordinator | None = None, + clock: Callable[[], float] = time.time, + expiry_skew_seconds: float = _EXPIRY_SKEW_SECONDS, + ) -> None: + self._read_grant = read_grant + self._refresh_grant = refresh_grant + self._coordinator: RefreshCoordinator = coordinator or InProcessRefreshCoordinator() + self._clock = clock + self._expiry_skew_seconds = expiry_skew_seconds + + async def subject_token(self, user_id: str, config: TokenExchangeConfig) -> str | None: + endpoint = config.token_exchange_endpoint + if not endpoint: + # No IdP endpoint to source or refresh against; the exchanger separately fails closed (412) + # for a non-delegated call, so here we simply have no subject material for the user. + return None + idp_key = idp_grant_key(endpoint) + grant = await self._read_grant(user_id, idp_key) + if grant is None: + return None + if not self._is_expired(grant): + return grant.access_token + refreshed = await self._coordinator.run( + user_id, + idp_key, + refresh=lambda: self._refresh_if_still_expired(user_id, idp_key, config), + reread=lambda: self._read_grant(user_id, idp_key), + ) + if refreshed is None or self._is_expired(refreshed): + return None + return refreshed.access_token + + async def _refresh_if_still_expired( + self, user_id: str, idp_key: str, config: TokenExchangeConfig + ) -> OAuthToken | None: + # Re-read under the single-flight so a token another caller just refreshed is reused rather + # than refreshed again (and its refresh_token spent again under IdP rotation). + latest = await self._read_grant(user_id, idp_key) + if latest is None: + return None + if not self._is_expired(latest): + return latest + if latest.refresh_token is None: + return None + return await self._refresh_grant(user_id, idp_key, config, latest) + + def _is_expired(self, token: OAuthToken) -> bool: + return token.expires_at is not None and self._clock() >= token.expires_at - self._expiry_skew_seconds + + +class IdpGrantRefresher: + """Refreshes a stored IdP grant via the RFC 6749 refresh_token grant, then persists the rotation. + + Unlike the authorization_code refresher, the token endpoint and client credentials come from the + calling server's ``TokenExchangeConfig`` (the IdP is the token-exchange authorization server), not + a per-server lookup, since one IdP grant is shared across every ``token_exchange`` upstream it + fronts. The HTTP POST and the persist are injected so the form-building and (untyped) response + parsing stay testable without a live IdP or DB. Returns ``None`` (the source then challenges) when + there is no refresh_token, the config lacks the client credentials to authenticate to the endpoint, + or the grant fails; never a stale or partial token. A rotated refresh_token replaces the old one; + an omitted one is carried forward. + """ + + def __init__( + self, + token_endpoint: TokenEndpointPost, + persist: PersistIdpGrant, + *, + clock: Callable[[], float] = time.time, + ) -> None: + self._token_endpoint = token_endpoint + self._persist = persist + self._clock = clock + + async def refresh( + self, user_id: str, idp_key: str, config: TokenExchangeConfig, token: OAuthToken + ) -> OAuthToken | None: + endpoint = config.token_exchange_endpoint + if token.refresh_token is None or not endpoint or not config.client_id or config.client_secret is None: + return None + try: + client_auth = build_token_endpoint_client_auth( + auth_method=config.token_endpoint_auth_method, + client_id=config.client_id, + client_secret=config.client_secret.get_secret_value(), + ) + except TokenEndpointAuthConfigError as exc: + verbose_logger.warning("MCP IdP grant refresh misconfigured for %s: %s", idp_key, exc) + return None + form = { + "grant_type": "refresh_token", + "refresh_token": token.refresh_token, + **client_auth.body, + } + body = await self._token_endpoint(endpoint, form, client_auth.headers) + if body is None: + return None + access_token = body.get("access_token") + if not isinstance(access_token, str) or not access_token: + return None + rotated = body.get("refresh_token") + new_refresh = rotated if isinstance(rotated, str) and rotated else token.refresh_token + expires_in = _parse_expires_in(body.get("expires_in")) + await self._persist(user_id, idp_key, access_token, new_refresh, expires_in, token.scopes or None) + return OAuthToken( + access_token=access_token, + expires_at=self._clock() + expires_in if expires_in is not None else None, + refresh_token=new_refresh, + scopes=token.scopes, + ) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py index ecfd471190c..961f1ccd5eb 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py @@ -23,6 +23,9 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth impo NoOpAuth, StaticHeaderAuth, ) +from litellm.proxy._experimental.mcp_server.outbound_credentials.idp_subject_source import ( + IdpSubjectTokenSource, +) from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import ( OAuthToken, OAuthTokenStore, @@ -74,6 +77,14 @@ class _NullTokenExchanger: return None +class _NullIdpSubjectTokenSource: + """Fail-closed default: with no source wired, a delegated OBO request has no subject token, so + the token_exchange arm challenges rather than falling back to the agent's admission credential.""" + + async def subject_token(self, user_id: str, config: TokenExchangeConfig) -> str | None: + return None + + class UpstreamCredentialProvider: """Produces the one `httpx.Auth` for a `(subject, upstream)` pair, per declared mode. @@ -87,9 +98,11 @@ class UpstreamCredentialProvider: self, oauth_token_store: OAuthTokenStore | None = None, token_exchanger: TokenExchanger | None = None, + idp_subject_source: IdpSubjectTokenSource | None = None, ) -> None: self._oauth_token_store: OAuthTokenStore = oauth_token_store or _NullOAuthTokenStore() self._token_exchanger: TokenExchanger = token_exchanger or _NullTokenExchanger() + self._idp_subject_source: IdpSubjectTokenSource = idp_subject_source or _NullIdpSubjectTokenSource() async def resolve_credentials(self, subject: Subject, server: ServerSpec) -> Result[httpx.Auth, CredError]: match server.config: @@ -150,39 +163,76 @@ class UpstreamCredentialProvider: async def _token_exchange( self, subject: Subject, server: ServerSpec, config: TokenExchangeConfig ) -> Result[StaticHeaderAuth, CredError]: - """RFC 8693 OBO: exchange the caller's inbound token for an upstream-bound bearer. + """RFC 8693 OBO: exchange the subject's token for an upstream-bound bearer. - No inbound token means there is nothing to exchange, so the arm fails closed with a 401 rather + The subject token is the caller's own inbound token for a direct request, or the delegated + user's IdP token (Path B) for an agent-delegated one; it is never the agent's admission + credential (see ``_obo_subject_token``). A missing subject token fails closed with a 401 rather than falling through to a weaker source (§1.5); the exchanger handles the IdP round-trip and caching and returns the upstream token or a typed error. """ - inbound = subject.inbound_token - if inbound is None: - return Error( - CredError.of_unauthorized( - "Token exchange requires a caller token to exchange (OBO).", - www_authenticate='Bearer error="invalid_request"', - ) - ) - match await self._token_exchanger.exchange( - inbound.get_secret_value(), server, config, tenant_id=subject.tenant_id - ): + subject_token = await self._obo_subject_token(subject, config) + if isinstance(subject_token, Error): + return Error(subject_token.error) + match await self._token_exchanger.exchange(subject_token.ok, server, config, tenant_id=subject.tenant_id): case Ok(token): return Ok(StaticHeaderAuth(f"Bearer {token.access_token}", header_name="Authorization")) case Error(err): return Error(err) + async def _obo_subject_token(self, subject: Subject, config: TokenExchangeConfig) -> Result[str, CredError]: + """The subject_token to exchange (OBO): the caller's own token, or the delegated user's IdP + token, never the agent's admission credential. + + For a direct (non-delegated) request the subject presents its own token inline, so the inbound + token IS the subject's proof and is exchanged as before; a missing one is the arm's 401. For an + agent-delegated request (``delegated_user_id`` set) the inbound token is the agent's admission + credential and must not be exchanged (exchanging it would mint agent-scoped access, the exact + escalation this design prevents); the delegated user's IdP grant is sourced instead (Path B), + and its absence fails closed with a 401 telling the client to authenticate the user with the + IdP. The agent's token is never a fallback. + """ + if subject.delegated_user_id is None: + inbound = subject.inbound_token + if inbound is None: + return Error( + CredError.of_unauthorized( + "Token exchange requires a caller token to exchange (OBO).", + www_authenticate='Bearer error="invalid_request"', + ) + ) + return Ok(inbound.get_secret_value()) + subject_token = await self._idp_subject_source.subject_token(subject.delegated_user_id, config) + if subject_token is None: + return Error( + CredError.of_unauthorized( + "Delegated user has not authorized the IdP for on-behalf-of access; complete consent and retry.", + www_authenticate='Bearer error="invalid_token"', + ) + ) + return Ok(subject_token) + async def invalidate_credentials(self, subject: Subject, server: ServerSpec) -> None: """Drop any cached credential the resolver owns for this `(subject, server)`. Used after an upstream rejects the injected credential, so the next resolve re-mints rather than serving the same rejected token until TTL. Only `token_exchange` holds a re-mintable - cached credential here; other modes are a no-op. + cached credential here; other modes are a no-op. The exchange cache is keyed on the + subject_token that minted it, so invalidation recomputes the key from the same source + `_token_exchange` uses (the delegated user's IdP token, or the caller's own inbound token), + never the agent's admission credential. This is best-effort: if the delegated user's IdP token + has since rotated the recomputed key no longer matches the rejected entry, which then lapses on + its own TTL rather than being dropped early. """ - if isinstance(server.config, TokenExchangeConfig) and subject.inbound_token is not None: - await self._token_exchanger.invalidate( - subject.inbound_token.get_secret_value(), server, server.config, tenant_id=subject.tenant_id - ) + if not isinstance(server.config, TokenExchangeConfig): + return + match await self._obo_subject_token(subject, server.config): + case Error(_): + return + case Ok(subject_token): + await self._token_exchanger.invalidate( + subject_token, server, server.config, tenant_id=subject.tenant_id + ) async def _authz_token(self, subject: Subject, server: ServerSpec) -> OAuthToken | None: """The user's authorization_code token, or None when absent or the store is unreachable. diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/types.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/types.py index 7e04be4f045..69d2cc49816 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/types.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/types.py @@ -340,6 +340,13 @@ class Subject(BaseModel): subject_id: str # Opaque, already-validated inbound identity. Only `token_exchange` / `passthrough` read it. inbound_token: SecretStr | None = None + # Set only for an agent-delegated (on-behalf-of) request; equals subject_id when present. Its + # presence means inbound_token is the AGENT's admission credential, NOT this subject's own token, + # so an arm that would otherwise consume inbound_token as the subject's proof (token_exchange) + # must instead source the subject's own material. subject_id stays the principal whose upstream + # credential is resolved either way, so arms keyed purely on subject_id (authorization_code) are + # already correct and read nothing here. + delegated_user_id: str | None = None class ServerSpec(BaseModel): diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_adapter.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_adapter.py index 349a81e73ed..648aa9641b2 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_adapter.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_adapter.py @@ -272,6 +272,9 @@ def test_to_subject_maps_principal_fields(): assert subject.tenant_id == "org1" assert subject.subject_id == "user1" assert subject.inbound_token is None + # A non-delegated request carries no delegation marker: the token_exchange arm then treats the + # inbound token as the subject's own, unchanged from before delegation existed. + assert subject.delegated_user_id is None def test_to_subject_delegated_user_replaces_credential_subject(): @@ -281,6 +284,9 @@ def test_to_subject_delegated_user_replaces_credential_subject(): subject = to_subject(principal, None) assert subject.subject_id == "delegated-user" assert subject.tenant_id == "team1" + # The delegation marker is surfaced so the token_exchange arm knows the inbound token is the + # agent's admission credential and sources the delegated user's IdP grant instead. + assert subject.delegated_user_id == "delegated-user" @pytest.mark.parametrize( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_idp_subject_provider.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_idp_subject_provider.py new file mode 100644 index 00000000000..a44586c197b --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_idp_subject_provider.py @@ -0,0 +1,56 @@ +"""Tests for the IdP subject-source composition root. + +The composition wires an expiry-aware cache in front of the DB read (mirroring the authorization_code +token store chain), so a warm IdP grant is served without a DB round-trip on every delegated call. +""" + +import pytest +from pydantic import SecretStr + +from litellm.proxy._experimental.mcp_server.outbound_credentials.idp_subject_provider import ( + build_idp_subject_source, +) +from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( + TokenExchangeConfig, +) + +_CONFIG = TokenExchangeConfig( + token_exchange_endpoint="https://idp.example.com/token", + client_id="gateway-client", + client_secret=SecretStr("gateway-secret"), +) + + +@pytest.mark.asyncio +async def test_grant_read_is_cached_across_delegated_calls(): + reads: list[tuple[str, str]] = [] + + async def counting_read_credential(user_id, idp_key): + reads.append((user_id, idp_key)) + # A far-future expiry so the grant is fresh -> cacheable, no refresh. + return {"type": "idp_grant", "access_token": "live-at", "expires_at": "2999-01-01T00:00:00+00:00"} + + source = build_idp_subject_source(read_credential=counting_read_credential) + + first = await source.subject_token("alice", _CONFIG) + second = await source.subject_token("alice", _CONFIG) + + assert first == "live-at" + assert second == "live-at" + # The second delegated call is served from the shared cache: exactly one DB read, not one per call. + assert len(reads) == 1 + + +@pytest.mark.asyncio +async def test_distinct_users_each_read_once(): + reads: list[tuple[str, str]] = [] + + async def counting_read_credential(user_id, idp_key): + reads.append((user_id, idp_key)) + return {"type": "idp_grant", "access_token": f"at-{user_id}", "expires_at": "2999-01-01T00:00:00+00:00"} + + source = build_idp_subject_source(read_credential=counting_read_credential) + assert await source.subject_token("alice", _CONFIG) == "at-alice" + assert await source.subject_token("bob", _CONFIG) == "at-bob" + # Cache is keyed per user, so the two users do not share an entry (and neither re-reads). + assert sorted(reads) == [("alice", "idp::https://idp.example.com/token"), ("bob", "idp::https://idp.example.com/token")] diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_idp_subject_source.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_idp_subject_source.py new file mode 100644 index 00000000000..a5310af0ab8 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_idp_subject_source.py @@ -0,0 +1,200 @@ +"""Tests for the delegated-OBO IdP subject-token source (Path B). + +``StoredIdpGrantSource`` reads a stored IdP grant and refreshes it to a live subject token, failing +closed (None) whenever the user has no usable grant. ``IdpGrantRefresher`` runs the refresh_token +grant against the config's IdP endpoint and persists the rotation. The DB read, refresh POST, persist, +and clock are injected, so every case is exercised without a live IdP or DB. +""" + +import asyncio + +import pytest +from pydantic import SecretStr + +from litellm.proxy._experimental.mcp_server.outbound_credentials.idp_subject_source import ( + IdpGrantRefresher, + StoredIdpGrantSource, + idp_grant_key, +) +from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import ( + OAuthToken, +) +from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( + TokenExchangeConfig, +) + +_CONFIG = TokenExchangeConfig( + token_exchange_endpoint="https://idp.example.com/oauth2/v1/token", + client_id="gateway-client", + client_secret=SecretStr("gateway-secret"), +) +_IDP_KEY = idp_grant_key("https://idp.example.com/oauth2/v1/token") + + +def _reader(grants): + async def read(user_id, idp_key): + return grants.get((user_id, idp_key)) + + return read + + +def _never_refresh(): + async def refresh(user_id, idp_key, config, token): + raise AssertionError("refresh must not run for a fresh grant") + + return refresh + + +def test_idp_grant_key_is_namespaced_and_normalizes_trailing_slash(): + assert idp_grant_key("https://idp.example.com/token/") == "idp::https://idp.example.com/token" + assert idp_grant_key("https://idp.example.com/token") == "idp::https://idp.example.com/token" + assert idp_grant_key("https://a/token").startswith("idp::") + + +@pytest.mark.asyncio +async def test_fresh_grant_is_returned_without_refreshing(): + grants = {("alice", _IDP_KEY): OAuthToken(access_token="live-at", expires_at=10_000.0)} + source = StoredIdpGrantSource(_reader(grants), _never_refresh(), clock=lambda: 100.0) + assert await source.subject_token("alice", _CONFIG) == "live-at" + + +@pytest.mark.asyncio +async def test_no_grant_returns_none(): + source = StoredIdpGrantSource(_reader({}), _never_refresh(), clock=lambda: 100.0) + assert await source.subject_token("alice", _CONFIG) is None + + +@pytest.mark.asyncio +async def test_no_endpoint_returns_none(): + grants = {("alice", _IDP_KEY): OAuthToken(access_token="live-at")} + source = StoredIdpGrantSource(_reader(grants), _never_refresh(), clock=lambda: 100.0) + no_endpoint = TokenExchangeConfig(client_id="c", client_secret=SecretStr("s")) + assert await source.subject_token("alice", no_endpoint) is None + + +@pytest.mark.asyncio +async def test_expired_grant_with_no_refresh_token_returns_none(): + grants = {("alice", _IDP_KEY): OAuthToken(access_token="stale", expires_at=50.0, refresh_token=None)} + source = StoredIdpGrantSource(_reader(grants), _never_refresh(), clock=lambda: 100.0) + assert await source.subject_token("alice", _CONFIG) is None + + +@pytest.mark.asyncio +async def test_expired_grant_is_refreshed_to_a_live_token(): + grants = {("alice", _IDP_KEY): OAuthToken(access_token="stale", expires_at=50.0, refresh_token="rt")} + refreshes = [] + + async def refresh(user_id, idp_key, config, token): + refreshes.append((user_id, idp_key, token.refresh_token)) + return OAuthToken(access_token="fresh-at", expires_at=10_000.0, refresh_token="rt2") + + source = StoredIdpGrantSource(_reader(grants), refresh, clock=lambda: 100.0) + assert await source.subject_token("alice", _CONFIG) == "fresh-at" + assert refreshes == [("alice", _IDP_KEY, "rt")] + + +@pytest.mark.asyncio +async def test_a_refresh_that_still_yields_an_expired_token_returns_none(): + grants = {("alice", _IDP_KEY): OAuthToken(access_token="stale", expires_at=50.0, refresh_token="rt")} + + async def refresh(user_id, idp_key, config, token): + return OAuthToken(access_token="still-stale", expires_at=60.0, refresh_token="rt") + + source = StoredIdpGrantSource(_reader(grants), refresh, clock=lambda: 100.0) + assert await source.subject_token("alice", _CONFIG) is None + + +@pytest.mark.asyncio +async def test_concurrent_expired_reads_refresh_once(): + """Refresh is single-flighted per (user, idp) so IdP refresh-token rotation isn't raced: two + concurrent callers for the same expired grant trigger exactly one refresh, not two (a second + would spend an already-rotated refresh_token).""" + grants = {("alice", _IDP_KEY): OAuthToken(access_token="stale", expires_at=50.0, refresh_token="rt")} + refresh_count = 0 + gate = asyncio.Event() + + async def refresh(user_id, idp_key, config, token): + nonlocal refresh_count + refresh_count += 1 + await gate.wait() + return OAuthToken(access_token="fresh-at", expires_at=10_000.0, refresh_token="rt2") + + source = StoredIdpGrantSource(_reader(grants), refresh, clock=lambda: 100.0) + task_a = asyncio.ensure_future(source.subject_token("alice", _CONFIG)) + task_b = asyncio.ensure_future(source.subject_token("alice", _CONFIG)) + await asyncio.sleep(0) + gate.set() + results = await asyncio.gather(task_a, task_b) + assert results == ["fresh-at", "fresh-at"] + assert refresh_count == 1 + + +class _RecordingPost: + def __init__(self, body): + self._body = body + self.calls = [] + + async def __call__(self, url, form, headers): + self.calls.append((url, form, headers)) + return self._body + + +class _RecordingPersist: + def __init__(self): + self.calls = [] + + async def __call__(self, user_id, idp_key, access_token, refresh_token, expires_in, scopes): + self.calls.append((user_id, idp_key, access_token, refresh_token, expires_in, scopes)) + + +@pytest.mark.asyncio +async def test_refresher_runs_refresh_token_grant_persists_and_returns_the_new_token(): + post = _RecordingPost({"access_token": "new-at", "refresh_token": "rt2", "expires_in": 3600}) + persist = _RecordingPersist() + refresher = IdpGrantRefresher(post, persist, clock=lambda: 1000.0) + stale = OAuthToken(access_token="old", expires_at=1.0, refresh_token="rt1", scopes=("read",)) + + result = await refresher.refresh("alice", _IDP_KEY, _CONFIG, stale) + + assert result is not None + assert result.access_token == "new-at" + assert result.refresh_token == "rt2" + assert result.expires_at == 1000.0 + 3600 + url, form, _headers = post.calls[0] + assert url == "https://idp.example.com/oauth2/v1/token" + assert form["grant_type"] == "refresh_token" + assert form["refresh_token"] == "rt1" + # The rotated triple is persisted so later reads see the new token without refreshing again. + assert persist.calls == [("alice", _IDP_KEY, "new-at", "rt2", 3600, ("read",))] + + +@pytest.mark.asyncio +async def test_refresher_carries_forward_an_omitted_refresh_token(): + post = _RecordingPost({"access_token": "new-at", "expires_in": 3600}) # no rotated refresh_token + persist = _RecordingPersist() + refresher = IdpGrantRefresher(post, persist, clock=lambda: 0.0) + stale = OAuthToken(access_token="old", expires_at=1.0, refresh_token="rt1") + result = await refresher.refresh("alice", _IDP_KEY, _CONFIG, stale) + assert result is not None and result.refresh_token == "rt1" + + +@pytest.mark.asyncio +async def test_refresher_returns_none_when_the_idp_call_fails(): + post = _RecordingPost(None) # transport / HTTP failure surfaces as a miss + persist = _RecordingPersist() + refresher = IdpGrantRefresher(post, persist) + stale = OAuthToken(access_token="old", expires_at=1.0, refresh_token="rt1") + assert await refresher.refresh("alice", _IDP_KEY, _CONFIG, stale) is None + assert persist.calls == [] + + +@pytest.mark.asyncio +async def test_refresher_returns_none_without_client_credentials(): + # A token_exchange config missing client credentials cannot authenticate to the endpoint, so the + # refresh cannot run and the caller fails closed rather than posting an unauthenticated grant. + post = _RecordingPost({"access_token": "new-at"}) + refresher = IdpGrantRefresher(post, _RecordingPersist()) + no_creds = TokenExchangeConfig(token_exchange_endpoint="https://idp.example.com/token") + stale = OAuthToken(access_token="old", expires_at=1.0, refresh_token="rt1") + assert await refresher.refresh("alice", _IDP_KEY, no_creds, stale) is None + assert post.calls == [] diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py index c88027abcd4..250ba6c2162 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py @@ -34,6 +34,9 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_sto OAuthToken, TokenStoreUnavailable, ) +from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import ( + OboTokenExchanger, +) _SUBJECT = Subject(tenant_id="", subject_id="") @@ -276,6 +279,122 @@ async def test_token_exchange_without_an_exchanger_fails_closed(): assert result.error.tag == "misconfigured" +class _FakeIdpSubjectSource: + """An IdpSubjectTokenSource returning a canned per-user IdP subject token (None == no grant).""" + + def __init__(self, by_user: dict) -> None: + self._by_user = by_user + self.calls: list[str] = [] + + async def subject_token(self, user_id, config): + self.calls.append(user_id) + return self._by_user.get(user_id) + + +def _delegated(subject_id, *, agent_token="agent-jwt", tenant="acme"): + # An agent-delegated request: the delegated user is the credential subject, but the inbound token + # on the wire is the AGENT's admission credential, never the user's own token. + return Subject( + tenant_id=tenant, + subject_id=subject_id, + inbound_token=SecretStr(agent_token), + delegated_user_id=subject_id, + ) + + +@pytest.mark.asyncio +async def test_delegated_token_exchange_exchanges_the_delegated_users_idp_token_not_the_agent_bearer(): + exchanger = _FakeExchanger(Ok(OAuthToken(access_token="upstream-alice"))) + idp = _FakeIdpSubjectSource({"alice": "alice-idp-token"}) + provider = UpstreamCredentialProvider(token_exchanger=exchanger, idp_subject_source=idp) + result = await provider.resolve_credentials(_delegated("alice"), _spec(_OBO)) + assert isinstance(result, Ok) + assert _emitted(result.ok)["Authorization"] == "Bearer upstream-alice" + # The subject token handed to the exchanger is the delegated user's IdP token, sourced for that + # user -- NOT the agent's inbound bearer (which would mint agent-scoped access, the escalation + # this arm must prevent). + assert exchanger.calls == [("alice-idp-token", "acme", "s")] + assert idp.calls == ["alice"] + + +@pytest.mark.asyncio +async def test_delegated_token_exchange_without_a_grant_is_unauthorized_and_never_uses_the_agent_bearer(): + exchanger = _FakeExchanger(Ok(OAuthToken(access_token="never"))) + idp = _FakeIdpSubjectSource({}) # the delegated user has not consented / has no stored grant + provider = UpstreamCredentialProvider(token_exchanger=exchanger, idp_subject_source=idp) + result = await provider.resolve_credentials(_delegated("alice"), _spec(_OBO)) + assert isinstance(result, Error) + assert result.error.tag == "unauthorized" + assert result.error.unauthorized.www_authenticate == 'Bearer error="invalid_token"' + # Fail closed: the agent's admission bearer is never a fallback subject, so the IdP is not hit. + assert exchanger.calls == [] + + +@pytest.mark.asyncio +async def test_delegated_token_exchange_without_an_idp_source_wired_fails_closed(): + # The fail-closed default (no IdP source wired) must not fall back to the agent's bearer. + exchanger = _FakeExchanger(Ok(OAuthToken(access_token="never"))) + provider = UpstreamCredentialProvider(token_exchanger=exchanger) + result = await provider.resolve_credentials(_delegated("alice"), _spec(_OBO)) + assert isinstance(result, Error) + assert result.error.tag == "unauthorized" + assert exchanger.calls == [] + + +@pytest.mark.asyncio +async def test_non_delegated_token_exchange_ignores_the_idp_source(): + # A direct (non-delegated) request presents its own token inline, so the arm exchanges the inbound + # token exactly as before delegation existed and never consults the IdP source. + exchanger = _FakeExchanger(Ok(OAuthToken(access_token="exchanged"))) + idp = _FakeIdpSubjectSource({"alice": "should-not-be-used"}) + provider = UpstreamCredentialProvider(token_exchanger=exchanger, idp_subject_source=idp) + subject = Subject(tenant_id="acme", subject_id="alice", inbound_token=SecretStr("caller-jwt")) + result = await provider.resolve_credentials(subject, _spec(_OBO)) + assert isinstance(result, Ok) + assert exchanger.calls == [("caller-jwt", "acme", "s")] + assert idp.calls == [] + + +@pytest.mark.asyncio +async def test_delegated_token_exchange_isolates_users_behind_one_agent(): + """Two users delegating through ONE agent (same inbound agent bearer) must not share a minted + token. The subject token is each user's own IdP token, so the exchange cache -- keyed on the + subject token -- isolates them. If the arm exchanged the shared agent bearer instead, both would + hash to one cache entry and the second user would be served the first user's upstream token.""" + + class _EchoPost: + def __init__(self): + self.count = 0 + + async def __call__(self, url, form, headers): + self.count += 1 + return {"access_token": f"upstream-for-{form['subject_token']}", "token_type": "Bearer"} + + post = _EchoPost() + exchanger = OboTokenExchanger(post) + idp = _FakeIdpSubjectSource({"alice": "alice-idp-token", "bob": "bob-idp-token"}) + provider = UpstreamCredentialProvider(token_exchanger=exchanger, idp_subject_source=idp) + + alice = await provider.resolve_credentials(_delegated("alice"), _spec(_OBO)) + bob = await provider.resolve_credentials(_delegated("bob"), _spec(_OBO)) + + assert isinstance(alice, Ok) and _emitted(alice.ok)["Authorization"] == "Bearer upstream-for-alice-idp-token" + assert isinstance(bob, Ok) and _emitted(bob.ok)["Authorization"] == "Bearer upstream-for-bob-idp-token" + # Two distinct real exchanges: no cache collision on the shared agent bearer. + assert post.count == 2 + + +@pytest.mark.asyncio +async def test_delegated_invalidate_credentials_drops_the_delegated_users_entry_not_the_agent_bearer(): + exchanger = _FakeExchanger(Ok(OAuthToken(access_token="exchanged"))) + idp = _FakeIdpSubjectSource({"alice": "alice-idp-token"}) + provider = UpstreamCredentialProvider(token_exchanger=exchanger, idp_subject_source=idp) + await provider.invalidate_credentials(_delegated("alice"), _spec(_OBO)) + # Invalidation must recompute the cache key from the same subject token the exchange used (the + # user's IdP token), or it would drop the wrong entry and leave the rejected token cached. + assert exchanger.invalidations == [("alice-idp-token", "acme", "s")] + + @pytest.mark.asyncio async def test_passthrough_forwards_the_inbound_token_verbatim(): subject = Subject(tenant_id="", subject_id="", inbound_token=SecretStr("Bearer upstream-xyz")) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py index 7269774442b..c795b5f31bd 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py @@ -20,6 +20,7 @@ from litellm.proxy._experimental.mcp_server.db import ( _decode_user_credential, _prepare_mcp_server_data, get_user_credential, + get_user_idp_grant, get_user_oauth_credential, is_oauth_credential_expired, list_user_oauth_credentials, @@ -27,6 +28,7 @@ from litellm.proxy._experimental.mcp_server.db import ( rotate_mcp_user_credentials_master_key, rotate_mcp_user_env_vars_master_key, store_user_credential, + store_user_idp_grant, store_user_oauth_credential, ) from litellm.proxy._types import NewMCPServerRequest, UpdateMCPServerRequest @@ -537,6 +539,62 @@ async def test_list_oauth_credentials_filters_byok_and_returns_payloads(): assert tokens == {"tok-enc", "tok-legacy"} +# ── IdP grant (delegated OBO subject material) ──────────────────────────────── + + +@pytest.mark.asyncio +async def test_idp_grant_round_trips(): + prisma = _make_prisma_with_existing(row=None) + await store_user_idp_grant( + prisma, "alice", "idp::https://idp.example.com/token", "idp-at", refresh_token="idp-rt", scopes=["read"] + ) + + stored = _stored_value(prisma) + row = MagicMock() + row.credential_b64 = stored + prisma.db.litellm_mcpusercredentials.find_unique = AsyncMock(return_value=row) + + result = await get_user_idp_grant(prisma, "alice", "idp::https://idp.example.com/token") + assert result is not None + assert result["type"] == "idp_grant" + assert result["access_token"] == "idp-at" + assert result["refresh_token"] == "idp-rt" + assert result["scopes"] == ["read"] + + +@pytest.mark.asyncio +async def test_idp_grant_does_not_persist_plaintext(): + prisma = _make_prisma_with_existing(row=None) + await store_user_idp_grant(prisma, "alice", "idp::https://idp/token", "idp-secret-at", refresh_token="idp-secret-rt") + stored = _stored_value(prisma) + try: + decoded_bytes = base64.urlsafe_b64decode(stored) + except Exception: + decoded_bytes = b"" + assert b"idp-secret-at" not in decoded_bytes + assert b"idp-secret-rt" not in decoded_bytes + + +@pytest.mark.asyncio +async def test_idp_grant_is_not_surfaced_as_an_oauth2_credential(): + # An IdP grant is a THIRD credential kind sharing this table; the oauth2-only readers gate on + # type == "oauth2", so an idp_grant row must be invisible to get_user_oauth_credential and to the + # connected-servers listing -- otherwise it would show up as a phantom server the user "connected". + prisma = _make_prisma_with_existing(row=None) + await store_user_idp_grant(prisma, "alice", "idp::https://idp/token", "idp-at") + stored = _stored_value(prisma) + idp_row = MagicMock() + idp_row.credential_b64 = stored + idp_row.server_id = "idp::https://idp/token" + + prisma.db.litellm_mcpusercredentials.find_unique = AsyncMock(return_value=idp_row) + assert await get_user_oauth_credential(prisma, "alice", "idp::https://idp/token") is None + assert await get_user_idp_grant(prisma, "alice", "idp::https://idp/token") is not None + + prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[idp_row]) + assert await list_user_oauth_credentials(prisma, "alice") == [] + + # ── _decode_user_credential helper ────────────────────────────────────────────