mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
feat(mcp): source the delegated user's IdP token for token_exchange OBO
A delegated (agent-on-behalf-of-user) MCP tool call to a token_exchange (RFC 8693) server exchanged the agent's own inbound bearer and ignored the delegated user, silently minting agent-scoped upstream access. The authorization_code arm already resolved by subject_id so it honored delegation; token_exchange was the gap. Subject gains a delegated_user_id marker (set in to_subject); resolver _token_exchange and invalidate_credentials both route through one shared _obo_subject_token helper. A direct request exchanges the caller's own inbound token, unchanged; a delegated one sources the delegated user's live IdP token from the injected IdpSubjectTokenSource (Path B: read the stored IdP grant, refresh it single-flighted when expired) and fails closed with a 401 when the user has no grant. The agent's token is never a fallback, so the escalation is closed even before the consent capture lands. The IdP grant is stored in the existing per-user MCP credential table as a distinct type "idp_grant" payload keyed by IdP, so there is no new table and no migration and the oauth2-only server listings never surface it. Cross-user isolation falls out of the per-user subject token: the exchange cache is keyed on the subject_token, now each user's own IdP token, so N users behind one agent never share a minted token. Part of LIT-4448 (build item 3).
This commit is contained in:
parent
490028b0ef
commit
8ee8e2b097
12 changed files with 969 additions and 19 deletions
|
|
@ -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)."""
|
||||
|
|
|
|||
|
|
@ -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] = {}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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")]
|
||||
|
|
@ -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 == []
|
||||
|
|
@ -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"))
|
||||
|
|
|
|||
|
|
@ -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 ────────────────────────────────────────────
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue