diff --git a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py index 2ea1bd062ce..5b45c0dd859 100644 --- a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py +++ b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py @@ -126,6 +126,9 @@ _CLAIM_TTL_BUFFER_SECONDS: Final = 60 _USED_CODE_CACHE_PREFIX: Final = "mcp_gateway_dcr_code_used:" _USED_FLOW_CACHE_PREFIX: Final = "mcp_gateway_dcr_flow_used:" _USED_REFRESH_CACHE_PREFIX: Final = "mcp_gateway_dcr_refresh_used:" +_REFRESH_REPLAYED_DESCRIPTION: Final = "the refresh token was already used" +_REVOKED_REFRESH_FAMILY_CACHE_PREFIX: Final = "mcp_gateway_dcr_refresh_family_revoked:" +_REFRESH_FAMILY_REVOKED_DESCRIPTION: Final = "the refresh token was revoked" MAX_REDIRECT_URIS: Final = 4 MAX_REDIRECT_URI_LENGTH: Final = 256 @@ -344,6 +347,41 @@ def _oauth_error(status_code: int, error: str, description: str) -> JSONResponse ) +class SessionSigning(LiteLLMBaseModel): + """The key material a token verb mints and opens session tokens under, resolved once per + call together with the instant it serves, so every token the call issues or checks agrees + on ``now``.""" + + model_config = ConfigDict(frozen=True) + keys: SessionSigningKeys + now: datetime + + +def resolve_session_signing(master_key: str | None, caller: str) -> SessionSigning | Response: + """A token verb's precondition: the proxy has a master key and the session signing + configuration validates. Either defect is an operator fault the verb answers as a 500 + (never a 4xx the client would treat as its own), logged under ``caller``.""" + if master_key is None: + verbose_logger.error("%s rejected: no master_key configured", caller) + return _oauth_error(500, "server_error", "the gateway has no master key configured") + keys: Final = active_session_signing_keys(master_key) + if isinstance(keys, SessionSigningConfigError): + verbose_logger.error("%s rejected: %s", caller, keys.detail) + return _oauth_error(500, "server_error", "the gateway session signing configuration is invalid") + return SessionSigning(keys=keys, now=datetime.now(timezone.utc)) + + +def _refresh_claim_key(jti: str) -> str: + return f"{_USED_REFRESH_CACHE_PREFIX}{jti}" + + +def _refresh_family_key(family: str) -> str: + return f"{_REVOKED_REFRESH_FAMILY_CACHE_PREFIX}{family}" + + +_REFRESH_CLAIM_TTL_SECONDS: Final = SESSION_REFRESH_TTL_SECONDS + _CLAIM_TTL_BUFFER_SECONDS + + def _seal(prefix: str, payload: BaseModel) -> str: """Serialized ``exclude_none`` for the same reason session JWTs are minted that way: an optional claim that is unset never reaches the wire, so during a rolling deploy a blob @@ -1056,13 +1094,17 @@ class _SingleUseGuard: async def peek(self, key: str) -> Literal["unclaimed", "claimed", "unavailable"]: """Read-only view of a single-use marker, resolved against the same shared authority as :meth:`claim` so introspection observes exactly the record redemption and revocation wrote. - A backend fault is ``"unavailable"`` (fail closed) rather than a guess either way.""" + A backend fault is ``"unavailable"`` (fail closed) rather than a guess either way, which is why + the read goes to the Redis client itself: the cache wrapper's ``async_get_cache`` turns a fault + into ``None``, and ``None`` here would pass as unclaimed. The key is namespaced exactly as the + wrapper namespaces every write, or a ``namespace`` configured under ``coordination_redis`` or + ``cache_params`` would hide the marker.""" from litellm.proxy.proxy_server import redis_usage_cache # noqa: PLC0415 # circular import at module load redis_cache: Final = redis_usage_cache or getattr(self._cache, "redis_cache", None) if redis_cache is not None: try: - value = await redis_cache.async_get_cache(key) + value = await redis_cache.init_async_client().get(redis_cache.check_and_fix_namespace(key=key)) except Exception as e: # noqa: BLE001 # ANY Redis fault fails the read closed verbose_logger.warning("mcp gateway single-use peek: shared cache backend unavailable: %s", e) return "unavailable" @@ -1071,9 +1113,28 @@ class _SingleUseGuard: return "unclaimed" if local is None else "claimed" -def _session_token_pair(principal: SessionPrincipal, keys: SessionSigningKeys, now: datetime) -> Response: +async def _family_refusal(guard: _SingleUseGuard, family: str | None) -> Response | None: + """A refresh token whose rotation chain a revocation ended is refused before anything is + minted; a backend that cannot say answers 503.""" + if family is None: + return None + peeked: Final = await guard.peek(_refresh_family_key(family)) + match peeked: + case "unclaimed": + return None + case "claimed": + return _oauth_error(400, "invalid_grant", _REFRESH_FAMILY_REVOKED_DESCRIPTION) + case "unavailable": + return _oauth_error(503, "temporarily_unavailable", _CLAIM_UNAVAILABLE_DESCRIPTION) + case _: + assert_never(peeked) + + +def _session_token_pair( + principal: SessionPrincipal, keys: SessionSigningKeys, now: datetime, family: str | None +) -> Response: access: Final = mint_session_token(principal, keys, now) - refresh: Final = mint_session_refresh_token(principal, keys, now) + refresh: Final = mint_session_refresh_token(principal, keys, now, family=family) if not isinstance(access, MintedSessionToken) or not isinstance(refresh, MintedSessionToken): return _oauth_error(500, "server_error", "failed to mint the session credential") return JSONResponse( @@ -1098,20 +1159,21 @@ class _ProxyCredentialTokenResponse(TypedDict): issued_token_type: NotRequired[ReadOnly[_IssuedTokenType]] -def _proxy_credential_response( +def proxy_credential_response( minted: MintedProxyCredential, principal: SessionPrincipal, - keys: SessionSigningKeys, - now: datetime, + signing: SessionSigning, issued_token_type: _IssuedTokenType | None = None, + family: str | None = None, ) -> Response: """The proxy-API token response: the access token is the very credential ``lite login`` stores (accepted on every proxy route with user and team attribution), and the refresh token is a gateway-sealed rotating token bound to the team the credential - was minted for, so a renewal keeps the team the user consented to. A token exchange - also states ``issued_token_type``, which RFC 8693 section 2.2.1 requires.""" + was minted for, so a renewal keeps the team the user consented to, and to the rotation + chain (``family``) the presented token belonged to. A token exchange also states + ``issued_token_type``, which RFC 8693 section 2.2.1 requires.""" bound_principal: Final = principal.model_copy(update=MappingProxyType({"team_id": minted.team_id})) - refresh: Final = mint_session_refresh_token(bound_principal, keys, now) + refresh: Final = mint_session_refresh_token(bound_principal, signing.keys, signing.now, family=family) if not isinstance(refresh, MintedSessionToken): return _oauth_error(500, "server_error", "failed to mint the session credential") credential: Final[_ProxyCredentialTokenResponse] = { @@ -1162,7 +1224,7 @@ def _mint_failure_response(failure: ProxyCredentialMintFailure) -> Response: ) case "team_required": return _oauth_error( - 400, "invalid_grant", "this user belongs to a team; sign in again and pick the team for this credential" + 400, "invalid_grant", "this user belongs to a team; sign in again to get a credential issued for it" ) case "unavailable" | "faulted" | "unresolvable" | "no_active_key": return _reload_failure_response(failure) @@ -1207,19 +1269,13 @@ async def aggregate_token( with that audience, and the RFC 8693 token exchange that turns an IdP token straight into the proxy-API credential. Every path re-validates the litellm user live before minting, so a deactivated user cannot obtain or renew a session.""" - if master_key is None: - verbose_logger.error("mcp_gateway_dcr token grant rejected: no master_key configured") - return _oauth_error(500, "server_error", "the gateway has no master key configured") - keys: Final = active_session_signing_keys(master_key) - if isinstance(keys, SessionSigningConfigError): - verbose_logger.error("mcp_gateway_dcr token grant rejected: %s", keys.detail) - return _oauth_error(500, "server_error", "the gateway session signing configuration is invalid") - now: Final = datetime.now(timezone.utc) + signing: Final = resolve_session_signing(master_key, "mcp_gateway_dcr token grant") + if isinstance(signing, Response): + return signing issue: Final = _GrantIssuer( request=request, resource=resource, - keys=keys, - now=now, + signing=signing, reload_user=reload_user, mint_proxy_credential=mint_proxy_credential, guard=_SingleUseGuard(cache), @@ -1232,7 +1288,7 @@ async def aggregate_token( client_id=client_id, code_verifier=code_verifier, resource=resource, - now=now, + now=signing.now, issue=issue, ) if grant_type == "refresh_token": @@ -1241,8 +1297,7 @@ async def aggregate_token( refresh_token=refresh_token, client_id=client_id, resource=resource, - keys=keys, - now=now, + signing=signing, issue=issue, ) if grant_type == TOKEN_EXCHANGE_GRANT_TYPE: @@ -1261,6 +1316,44 @@ async def aggregate_token( ) +class _ProxyCredentialIssuer: + """The proxy-API tail every grant for that audience shares once its own proof has checked + out: refuse a rotation chain that was already ended, mint the live credential (which + revalidates the user and the team membership against the database), then claim the + single-use marker, ending the whole chain when the marker was already claimed. The claim + comes AFTER minting so a transient DB 503 or a membership refusal never burns a + still-valid grant, and it fails closed when it cannot be recorded.""" + + def __init__( + self, signing: SessionSigning, mint_proxy_credential: MintProxyCredential, guard: _SingleUseGuard + ) -> None: + self._signing: Final = signing + self._mint_proxy_credential: Final = mint_proxy_credential + self._guard: Final = guard + + async def __call__( + self, + principal: SessionPrincipal, + claim_key: str, + claim_ttl_seconds: int, + replayed: str, + family: str | None = None, + ) -> Response: + ended: Final = await _family_refusal(self._guard, family) + if ended is not None: + return ended + minted: Final = await self._mint_proxy_credential(principal.user_id, principal.team_id) + if not isinstance(minted, MintedProxyCredential): + return _mint_failure_response(minted) + refusal: Final = _claim_refusal( + await self._guard.claim(claim_key, claim_ttl_seconds), + replayed=_oauth_error(400, "invalid_grant", replayed), + ) + if refusal is not None: + return refusal + return proxy_credential_response(minted, principal, self._signing, family=family) + + class _GrantIssuer: """The tail every grant shares once its own proof (code + PKCE, or a refresh token) has checked out: revalidate the user live, claim the single-use marker, mint. The @@ -1271,55 +1364,59 @@ class _GrantIssuer: self, request: Request, resource: str | None, - keys: SessionSigningKeys, - now: datetime, + signing: SessionSigning, reload_user: ReloadUser, mint_proxy_credential: MintProxyCredential, guard: _SingleUseGuard, ) -> None: self._request: Final = request self._resource: Final = resource - self._keys: Final = keys - self._now: Final = now + self._signing: Final = signing self._reload_user: Final = reload_user self._mint_proxy_credential: Final = mint_proxy_credential self._guard: Final = guard + self._proxy_credential: Final = _ProxyCredentialIssuer(signing, mint_proxy_credential, guard) async def __call__( - self, principal: SessionPrincipal, claim_key: str, claim_ttl_seconds: int, replayed: str + self, + principal: SessionPrincipal, + claim_key: str, + claim_ttl_seconds: int, + replayed: str, + family: str | None = None, ) -> Response: match principal.audience: case None: - return await self._issue_session_pair(principal, claim_key, claim_ttl_seconds, replayed) + return await self._issue_session_pair(principal, claim_key, claim_ttl_seconds, replayed, family) case "proxy_api": - return await self._issue_proxy_credential(principal, claim_key, claim_ttl_seconds, replayed) + return await self._issue_proxy_credential(principal, claim_key, claim_ttl_seconds, replayed, family) case _: assert_never(principal.audience) async def _issue_session_pair( - self, principal: SessionPrincipal, claim_key: str, claim_ttl_seconds: int, replayed: str + self, principal: SessionPrincipal, claim_key: str, claim_ttl_seconds: int, replayed: str, family: str | None ) -> Response: + ended: Final = await _family_refusal(self._guard, family) + if ended is not None: + return ended failure: Final = await self._reload_user(principal.user_id) if failure is not None: return _reload_failure_response(failure) - refusal: Final = await self._claim_refusal(claim_key, claim_ttl_seconds, replayed) + refusal: Final = _claim_refusal( + await self._guard.claim(claim_key, claim_ttl_seconds), + replayed=_oauth_error(400, "invalid_grant", replayed), + ) if refusal is not None: return refusal - return _session_token_pair(principal, self._keys, self._now) + return _session_token_pair(principal, self._signing.keys, self._signing.now, family) async def _issue_proxy_credential( - self, principal: SessionPrincipal, claim_key: str, claim_ttl_seconds: int, replayed: str + self, principal: SessionPrincipal, claim_key: str, claim_ttl_seconds: int, replayed: str, family: str | None ) -> Response: target_refusal: Final = self._proxy_api_target_refusal() if target_refusal is not None: return target_refusal - minted: Final = await self._mint_proxy_credential(principal.user_id, principal.team_id) - if not isinstance(minted, MintedProxyCredential): - return _mint_failure_response(minted) - refusal: Final = await self._claim_refusal(claim_key, claim_ttl_seconds, replayed) - if refusal is not None: - return refusal - return _proxy_credential_response(minted, principal, self._keys, self._now) + return await self._proxy_credential(principal, claim_key, claim_ttl_seconds, replayed, family) async def exchange( self, subject_token: str, client_id: str, exchange_subject_token: ExchangeSubjectToken @@ -1339,20 +1436,13 @@ class _GrantIssuer: minted: Final = await self._mint_proxy_credential(principal.user_id, principal.team_id) if not isinstance(minted, MintedProxyCredential): return _mint_failure_response(minted) - return _proxy_credential_response( - minted, principal, self._keys, self._now, issued_token_type=ACCESS_TOKEN_TOKEN_TYPE - ) + return proxy_credential_response(minted, principal, self._signing, issued_token_type=ACCESS_TOKEN_TOKEN_TYPE) def _proxy_api_target_refusal(self) -> Response | None: if self._resource is None or is_proxy_api_resource(self._request, self._resource): return None return _oauth_error(400, "invalid_target", "resource does not match the proxy API this grant was issued for") - async def _claim_refusal(self, claim_key: str, claim_ttl_seconds: int, replayed: str) -> Response | None: - return _claim_refusal( - await self._guard.claim(claim_key, claim_ttl_seconds), replayed=_oauth_error(400, "invalid_grant", replayed) - ) - async def _authorization_code_grant( request: Request, @@ -1395,30 +1485,73 @@ async def _authorization_code_grant( ) +def _open_presented_refresh_token( + refresh_token: str | None, client_id: str, signing: SessionSigning +) -> SessionRefreshOpened | Response: + if not refresh_token: + return _oauth_error(400, "invalid_request", "refresh_token is required") + opened: Final = open_session_refresh_bearer(refresh_token, signing.keys, signing.now, expected_client_id=client_id) + if not isinstance(opened, SessionRefreshOpened): + return _oauth_error(400, "invalid_grant", "the refresh token is invalid for this client") + return opened + + async def _refresh_token_grant( request: Request, refresh_token: str | None, client_id: str, resource: str | None, - keys: SessionSigningKeys, - now: datetime, + signing: SessionSigning, issue: _GrantIssuer, ) -> Response: - if not refresh_token: - return _oauth_error(400, "invalid_request", "refresh_token is required") - opened: Final = open_session_refresh_bearer(refresh_token, keys, now, expected_client_id=client_id) - if not isinstance(opened, SessionRefreshOpened): - return _oauth_error(400, "invalid_grant", "the refresh token is invalid for this client") + opened: Final = _open_presented_refresh_token(refresh_token, client_id, signing) + if isinstance(opened, Response): + return opened if _resource_conflicts_with_scope(request, resource, opened.principal.resource_server_id): return _oauth_error(400, "invalid_target", "resource does not match the scope this token was issued for") # Refresh-token rotation (OAuth 2.0 Security BCP section 4.13): the presented refresh token is # single-use, so a captured or replayed refresh token cannot mint a second pair after the - # legitimate holder rotated. + # legitimate holder rotated. A replay refuses only itself, never the chain (section 4.13.2's + # chain ending): Claude Code renews from its in-memory copy, so a second terminal on the same + # machine presents the token the first one already rotated, then recovers from the shared + # credential file, and the chain it recovers into has to be alive. Only a revocation ends it. return await issue( opened.principal, - claim_key=f"{_USED_REFRESH_CACHE_PREFIX}{opened.jti}", - claim_ttl_seconds=SESSION_REFRESH_TTL_SECONDS + _CLAIM_TTL_BUFFER_SECONDS, - replayed="the refresh token was already used", + claim_key=_refresh_claim_key(opened.jti), + claim_ttl_seconds=_REFRESH_CLAIM_TTL_SECONDS, + replayed=_REFRESH_REPLAYED_DESCRIPTION, + family=opened.family, + ) + + +async def refresh_proxy_credential( + refresh_token: str | None, + client_id: str, + master_key: str | None, + cache: DualCache, + mint_proxy_credential: MintProxyCredential, +) -> Response: + """The refresh_token grant for a fixed public client the gateway never registered (the + Claude Code CLI, which presents no client_id of its own): the token must have been minted + for ``client_id`` with the proxy-API audience, the credential is re-minted live, and the + presented token rotates under the same single-use record the DCR flow burns, so one + revocation covers both front doors. The identity-only session pair is never issued here: + a token of that audience answers invalid_grant even when its client binding matches.""" + signing: Final = resolve_session_signing(master_key, "mcp_gateway refresh grant") + if isinstance(signing, Response): + return signing + opened: Final = _open_presented_refresh_token(refresh_token, client_id, signing) + if isinstance(opened, Response): + return opened + if opened.principal.audience != PROXY_API_AUDIENCE: + return _oauth_error(400, "invalid_grant", "the refresh token is invalid for this client") + issue: Final = _ProxyCredentialIssuer(signing, mint_proxy_credential, _SingleUseGuard(cache)) + return await issue( + opened.principal, + claim_key=_refresh_claim_key(opened.jti), + claim_ttl_seconds=_REFRESH_CLAIM_TTL_SECONDS, + replayed=_REFRESH_REPLAYED_DESCRIPTION, + family=opened.family, ) @@ -1450,7 +1583,8 @@ async def _token_exchange_grant( async def revoke_refresh_token(token: str, client_id: str, master_key: str | None, cache: DualCache) -> Response: """RFC 7009 revocation for the gateway's refresh tokens: burn the presented token's - ``jti`` so neither the holder nor a thief can rotate it again. Access tokens are + ``jti`` and the rotation chain it belongs to, so no rotation of it, whoever holds one, + mints again (section 2.1: the grant is revoked, not one token). Access tokens are stateless and expire on their own (the proxy-API credential within ``CLI_JWT_EXPIRATION_HOURS``), so per RFC 7009 section 2.2 an unrecognized or already dead token still answers 200; only an unknown client is refused. A live token whose @@ -1459,20 +1593,28 @@ async def revoke_refresh_token(token: str, client_id: str, master_key: str | Non never happened.""" if not is_gateway_dcr_client_id(client_id) or open_gateway_dcr_client(client_id) is None: return _oauth_error(401, "invalid_client", "unknown or malformed client_id") - if master_key is None: - verbose_logger.error("mcp_gateway_dcr revoke rejected: no master_key configured") - return _oauth_error(500, "server_error", "the gateway has no master key configured") - keys: Final = active_session_signing_keys(master_key) - if isinstance(keys, SessionSigningConfigError): - verbose_logger.error("mcp_gateway_dcr revoke rejected: %s", keys.detail) - return _oauth_error(500, "server_error", "the gateway session signing configuration is invalid") - now: Final = datetime.now(timezone.utc) - opened: Final = open_session_refresh_bearer(token, keys, now, expected_client_id=client_id) + return await revoke_session_refresh_token(token, client_id, master_key, cache) + + +async def revoke_session_refresh_token( + token: str, client_id: str, master_key: str | None, cache: DualCache +) -> Response: + """The client-agnostic half of RFC 7009 revocation, shared with the fixed public clients + the gateway never registered: burn the presented refresh token's ``jti`` and its rotation + chain when it was issued to ``client_id``, answer 200 for anything else (an access token, + a dead or foreign token, garbage), and 503 when a burn could not be recorded in the shared + backend. The chain marker is written first: a fault between the two writes then leaves + every token of the chain refused rather than a descendant still renewable, and the retry + the 503 asks for only has the presented ``jti`` left to burn.""" + signing: Final = resolve_session_signing(master_key, "mcp_gateway revoke") + if isinstance(signing, Response): + return signing + opened: Final = open_session_refresh_bearer(token, signing.keys, signing.now, expected_client_id=client_id) if isinstance(opened, SessionRefreshOpened): - burned: Final = await _SingleUseGuard(cache).claim( - f"{_USED_REFRESH_CACHE_PREFIX}{opened.jti}", SESSION_REFRESH_TTL_SECONDS + _CLAIM_TTL_BUFFER_SECONDS - ) - if burned == "unavailable": + guard: Final = _SingleUseGuard(cache) + ended: Final = await guard.claim(_refresh_family_key(opened.family), _REFRESH_CLAIM_TTL_SECONDS) + burned: Final = await guard.claim(_refresh_claim_key(opened.jti), _REFRESH_CLAIM_TTL_SECONDS) + if "unavailable" in (burned, ended): return _oauth_error(503, "temporarily_unavailable", _CLAIM_UNAVAILABLE_DESCRIPTION) return Response(content="{}", media_type="application/json", headers=TOKEN_NO_CACHE_HEADERS) @@ -1544,10 +1686,12 @@ async def introspect_gateway_token( if not isinstance(opened, OpenedSessionToken): return _inactive_introspection_response() if opened.kind == "session_refresh": - peeked: Final = await _SingleUseGuard(cache).peek(f"{_USED_REFRESH_CACHE_PREFIX}{opened.jti}") - if peeked == "unavailable": + guard: Final = _SingleUseGuard(cache) + used: Final = await guard.peek(_refresh_claim_key(opened.jti)) + ended: Final = await guard.peek(_refresh_family_key(opened.family)) + if "unavailable" in (used, ended): return _oauth_error(503, "temporarily_unavailable", _CLAIM_UNAVAILABLE_DESCRIPTION) - if peeked == "claimed": + if "claimed" in (used, ended): return _inactive_introspection_response() failure: Final = await reload_user(opened.principal.user_id) if failure == "unavailable" or failure == "faulted": diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_credentials.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_credentials.py index b8eda703fa8..b216f463824 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_credentials.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_credentials.py @@ -241,12 +241,13 @@ def resolve_session_bearer( class SessionRefreshOpened(LiteLLMBaseModel): """A valid session refresh token presented to the token endpoint: the principal to - re-validate and renew under.""" + re-validate and renew under, and the rotation chain the renewal continues.""" model_config = ConfigDict(frozen=True) tag: Literal["opened"] = "opened" principal: SessionPrincipal jti: str + family: str class SessionRefreshInvalid(LiteLLMBaseModel): @@ -285,4 +286,4 @@ def open_session_refresh_bearer( return SessionRefreshInvalid() if opened.principal.client_id != expected_client_id: return SessionRefreshInvalid() - return SessionRefreshOpened(principal=opened.principal, jti=opened.jti) + return SessionRefreshOpened(principal=opened.principal, jti=opened.jti, family=opened.family) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py index da91faccadd..319ce81d619 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py @@ -14,8 +14,10 @@ signing approach as :mod:`.envelope`), or RS256 under an operator-provided RSA p key (:class:`AsymmetricSessionKeys`) so downstream validators hold only the public half. Claims are ``iss``/``iat``/``exp`` plus ``jti`` (per-mint uniqueness, so two tokens minted in the same second never -collide and a future revocation list has a stable handle), ``kind``, ``user_id``, and -``client_id``; ``client_id`` binds the refresh token +collide and the single-use record has a stable handle), ``kind``, ``user_id``, +``client_id``, and on a rotated refresh token ``family`` (the ``jti`` of the refresh token +its chain started from, so a revocation can end every rotation that descends +from one sign-in); ``client_id`` binds the refresh token to the DCR client it was issued to (RFC 6749 section 6) and is carried on the access token for parity and audit. There is no encrypted payload: nothing in a session token is secret beyond the signature, and reprs never print the signed value because minted @@ -224,13 +226,16 @@ class MintedSessionToken(LiteLLMBaseModel): class OpenedSessionToken(LiteLLMBaseModel): """A validated session token of either kind: the principal it was minted for, the - ``jti`` so the token endpoint can enforce single-use rotation on a refresh token, and + ``jti`` so the token endpoint can enforce single-use rotation on a refresh token, the + ``family`` every rotation of one refresh token shares (the root token's ``jti``; a + token minted before families were stamped, or an access token, is its own root), and the signed ``kind``/``iat``/``exp`` so an introspection response can report the token's metadata without re-decoding.""" model_config = ConfigDict(frozen=True) principal: SessionPrincipal jti: str + family: str kind: SessionTokenKind iat: int exp: int @@ -304,6 +309,7 @@ class _SessionClaims(LiteLLMBaseModel): resource_server_id: str | None = None audience: SessionAudience | None = None team_id: str | None = None + family: str | None = Field(default=None, min_length=1) def is_session_token(candidate: str) -> bool: @@ -342,12 +348,14 @@ def mint_session_refresh_token( principal: SessionPrincipal, keys: SessionSigningKeys, now: datetime, + family: str | None = None, ) -> MintedSessionToken | SessionTokenMintError: """Mint the long-lived session REFRESH token for ``principal``. ``exp`` is ``SESSION_REFRESH_TTL_SECONDS`` from ``now``. Minting a distinct ``kind="session_refresh"`` claim is what keeps a refresh token from ever opening as an - access credential at the MCP edge. + access credential at the MCP edge. ``family`` is the rotation chain the token continues + (the opened predecessor's ``family``); ``None`` starts a chain rooted at this token. """ return _mint( kind="session_refresh", @@ -356,6 +364,7 @@ def mint_session_refresh_token( expires_at=now + timedelta(seconds=SESSION_REFRESH_TTL_SECONDS), keys=keys, now=now, + family=family, ) @@ -393,6 +402,7 @@ def _mint( expires_at: datetime, keys: SessionSigningKeys, now: datetime, + family: str | None = None, ) -> MintedSessionToken | SessionTokenTooLarge: """Sign the claims for either token kind and enforce the size cap. Shared by both mints so the JWT shape, issuer, and size guard cannot drift between access and refresh.""" @@ -407,6 +417,7 @@ def _mint( resource_server_id=principal.resource_server_id, audience=principal.audience, team_id=principal.team_id, + family=family, ) token: Final = prefix + _sign_claims(claims, keys) size_bytes: Final = len(token.encode("utf-8")) @@ -465,6 +476,7 @@ def _open( team_id=claims.team_id, ), jti=claims.jti, + family=claims.family or claims.jti, kind=claims.kind, iat=claims.iat, exp=claims.exp, diff --git a/litellm/proxy/anthropic_endpoints/gateway_endpoints.py b/litellm/proxy/anthropic_endpoints/gateway_endpoints.py index 48acfa67a28..42ed6c425ab 100644 --- a/litellm/proxy/anthropic_endpoints/gateway_endpoints.py +++ b/litellm/proxy/anthropic_endpoints/gateway_endpoints.py @@ -2,8 +2,9 @@ Claude Code gateway protocol. Implements the wire contract the Claude Code CLI uses to talk to a gateway: -OAuth 2.0 device-authorization sign-in (RFC 8414 / RFC 8628), inference via the -Anthropic Messages API, managed settings, and OTLP telemetry ingestion. See +OAuth 2.0 device-authorization sign-in (RFC 8414 / RFC 8628) with a rotating +refresh token (RFC 6749 section 6) and RFC 7009 revocation for sign-out, inference +via the Anthropic Messages API, managed settings, and OTLP telemetry ingestion. See https://code.claude.com/docs/en/claude-apps-gateway. Everything lives under the ``/claude_code_gateway`` base so operators point @@ -11,7 +12,9 @@ Claude Code at ``https:///claude_code_gateway`` via ``/login``. The device flow reuses the proxy's existing SSO login machinery: the browser leg is served by ``/sso/key/generate`` and the shared ``cli_sso_session_cache`` flow, so the bearer token minted here is the same session JWT the LiteLLM CLI uses and -is accepted by every bearer-authenticated proxy route. +is accepted by every bearer-authenticated proxy route. The refresh token is the +MCP gateway's sealed session refresh token bound to a fixed client id, so a renewal +re-mints that JWT from the live user row and shares the DCR flow's single-use record. """ import hashlib @@ -20,9 +23,9 @@ import secrets from collections.abc import Mapping from dataclasses import dataclass from types import MappingProxyType -from typing import Final +from typing import Annotated, Final -from fastapi import APIRouter, Depends, Request, Response +from fastapi import APIRouter, Depends, Request, Response, status from fastapi.responses import JSONResponse from pydantic import Field, TypeAdapter, ValidationError @@ -34,6 +37,18 @@ from litellm.constants import ( CLI_SSO_SESSION_TTL_SECONDS, LITELLM_CLI_SOURCE_IDENTIFIER, ) +from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import ( + PROXY_API_AUDIENCE, + MintedProxyCredential, + MintProxyCredential, + SessionSigning, + proxy_credential_response, + refresh_proxy_credential, + resolve_session_signing, + revoke_session_refresh_token, +) +from litellm.proxy._experimental.mcp_server.oauth_utils import TOKEN_NO_CACHE_HEADERS +from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token import SessionPrincipal from litellm.proxy._types import LiteLLM_UserTable, LitellmUserRoles from litellm.proxy.anthropic_endpoints.endpoints import anthropic_response, count_tokens from litellm.proxy.auth.user_api_key_auth import user_api_key_auth @@ -43,6 +58,7 @@ from litellm.proxy.management_endpoints.ui_sso import CliSsoTeamDetail from litellm.types.llms.base import LiteLLMBaseModel GATEWAY_PREFIX: Final = "/claude_code_gateway" +CLAUDE_CODE_CLIENT_ID: Final = "claude_code" _DEVICE_CODE_GRANT: Final = "urn:ietf:params:oauth:grant-type:device_code" _REFRESH_TOKEN_GRANT: Final = "refresh_token" _DEVICE_CODE_SEPARATOR: Final = "." @@ -77,6 +93,7 @@ class _AuthorizationServerMetadata(LiteLLMBaseModel): issuer: str device_authorization_endpoint: str token_endpoint: str + revocation_endpoint: str grant_types_supported: tuple[str, ...] @@ -89,12 +106,6 @@ class _DeviceAuthorizationBody(LiteLLMBaseModel): interval: int -class _AccessTokenBody(LiteLLMBaseModel): - access_token: str - expires_in: int - token_type: str = "Bearer" - - class _ManagedSettingsBody(LiteLLMBaseModel): uuid: str checksum: str @@ -174,6 +185,7 @@ async def oauth_authorization_server(request: Request) -> JSONResponse: request_base_url=request_base_url, route="claude_code_gateway/oauth/device_authorization" ), token_endpoint=get_custom_url(request_base_url=request_base_url, route="claude_code_gateway/oauth/token"), + revocation_endpoint=get_custom_url(request_base_url=request_base_url, route="claude_code_gateway/oauth/revoke"), grant_types_supported=(_DEVICE_CODE_GRANT, _REFRESH_TOKEN_GRANT), ) return JSONResponse(content=metadata.model_dump()) @@ -274,6 +286,34 @@ def _mint_access_token(login: _GatewayLogin) -> str: ) +def _session_only_credential(login: _GatewayLogin) -> Response: + return JSONResponse( + status_code=status.HTTP_200_OK, + content={ + "access_token": _mint_access_token(login), + "token_type": "Bearer", + "expires_in": CLI_JWT_EXPIRATION_HOURS * _SECONDS_PER_HOUR, + }, + headers=TOKEN_NO_CACHE_HEADERS, + ) + + +def _renewable_credential(login: _GatewayLogin, signing: SessionSigning) -> Response: + user_id: Final = login.user_info.user_id + return proxy_credential_response( + MintedProxyCredential( + key=_mint_access_token(login), + expires_in=CLI_JWT_EXPIRATION_HOURS * _SECONDS_PER_HOUR, + user_id=user_id, + team_id=login.team_id, + ), + SessionPrincipal( + user_id=user_id, client_id=CLAUDE_CODE_CLIENT_ID, audience=PROXY_API_AUDIENCE, team_id=login.team_id + ), + signing, + ) + + @with_service_target(CLI_SSO_SESSIONS_TARGET) async def _claim_device_code(login_id: str, cache: DualCache) -> bool: from litellm.proxy.management_endpoints.ui_sso import ( @@ -288,8 +328,14 @@ async def _claim_device_code(login_id: str, cache: DualCache) -> bool: return claims == 1 +def _proxy_credential_minter() -> MintProxyCredential: + from litellm.proxy._experimental.mcp_server.proxy_api_credentials import mint_proxy_credential + + return mint_proxy_credential + + @with_service_target(CLI_SSO_SESSIONS_TARGET) -async def _handle_device_code_grant(device_code: str | None) -> JSONResponse: +async def _handle_device_code_grant(device_code: str | None) -> Response: from fastapi import HTTPException from litellm.proxy.management_endpoints.ui_sso import ( @@ -297,7 +343,7 @@ async def _handle_device_code_grant(device_code: str | None) -> JSONResponse: _get_cli_sso_flow_or_raise, # pyright: ignore[reportPrivateUsage] # shared device-flow helper _verify_cli_sso_poll_secret, # pyright: ignore[reportPrivateUsage] # shared device-flow helper ) - from litellm.proxy.proxy_server import cli_sso_session_cache + from litellm.proxy.proxy_server import cli_sso_session_cache, master_key if not device_code: return _oauth_error_response( @@ -320,17 +366,23 @@ async def _handle_device_code_grant(device_code: str | None) -> JSONResponse: if isinstance(login, _OAuthError): return _oauth_error_response(login) - access_token: Final = _mint_access_token(login) + signing: Final = resolve_session_signing(master_key, f"{CLAUDE_CODE_CLIENT_ID} token grant") + credential: Final = ( + _session_only_credential(login) if isinstance(signing, Response) else _renewable_credential(login, signing) + ) + if credential.status_code != status.HTTP_200_OK: + return credential if not await _claim_device_code(login_id, cli_sso_session_cache): return _oauth_error_response(_OAuthError(status_code=400, error="expired_token")) await cli_sso_session_cache.async_delete_cache(key=_get_cli_sso_flow_cache_key(login_id)) - body: Final = _AccessTokenBody(access_token=access_token, expires_in=CLI_JWT_EXPIRATION_HOURS * _SECONDS_PER_HOUR) - return JSONResponse(content=body.model_dump()) + return credential @router.post("/oauth/token", include_in_schema=False) -async def oauth_token(request: Request) -> JSONResponse: +async def oauth_token( + request: Request, mint_proxy_credential: Annotated[MintProxyCredential, Depends(_proxy_credential_minter)] +) -> Response: if not _is_gateway_enabled(): return _oauth_error_response(_OAuthError(status_code=404, error="not_found")) @@ -342,12 +394,15 @@ async def oauth_token(request: Request) -> JSONResponse: return await _handle_device_code_grant(device_code if isinstance(device_code, str) else None) if grant_type == _REFRESH_TOKEN_GRANT: - return _oauth_error_response( - _OAuthError( - status_code=401, - error="invalid_grant", - description="This gateway does not issue refresh tokens; sign in again", - ) + from litellm.proxy.proxy_server import master_key, user_api_key_cache + + refresh_token: Final = form.get("refresh_token") + return await refresh_proxy_credential( + refresh_token=refresh_token if isinstance(refresh_token, str) else None, + client_id=CLAUDE_CODE_CLIENT_ID, + master_key=master_key, + cache=user_api_key_cache, + mint_proxy_credential=mint_proxy_credential, ) return _oauth_error_response( @@ -357,6 +412,24 @@ async def oauth_token(request: Request) -> JSONResponse: ) +@router.post("/oauth/revoke", include_in_schema=False) +async def oauth_revoke(request: Request) -> Response: + if not _is_gateway_enabled(): + return _oauth_error_response(_OAuthError(status_code=404, error="not_found")) + + from litellm.proxy.proxy_server import master_key, user_api_key_cache + + form: Final = await request.form() + token: Final = form.get("token") + if not isinstance(token, str) or not token: + return _oauth_error_response( + _OAuthError(status_code=400, error="invalid_request", description="token is required") + ) + return await revoke_session_refresh_token( + token=token, client_id=CLAUDE_CODE_CLIENT_ID, master_key=master_key, cache=user_api_key_cache + ) + + @router.get("/managed/settings", include_in_schema=False, dependencies=_AUTHENTICATED) async def managed_settings(request: Request) -> Response: ensure_gateway_enabled() diff --git a/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_session_token.py b/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_session_token.py index bc7ec1e1112..bcde8a99566 100644 --- a/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_session_token.py +++ b/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_session_token.py @@ -189,6 +189,7 @@ def test_alg_none_token_is_rejected(): _valid_claims(iat="evil"), _valid_claims(kind="access"), _valid_claims(user_id=""), + _valid_claims(family=""), _valid_claims(nbf=0), {k: v for k, v in _valid_claims().items() if k != "client_id"}, {k: v for k, v in _valid_claims().items() if k != "exp"}, @@ -275,6 +276,31 @@ def test_proxy_api_audience_and_team_round_trip_through_the_refresh_token(): assert opened.principal == principal +def test_refresh_token_family_is_rooted_at_the_first_jti_and_carried_through_rotation(): + root = mint_session_refresh_token(PRINCIPAL, KEYS, NOW) + assert isinstance(root, MintedSessionToken) + root_token = root.token.get_secret_value() + assert "family" not in _decoded_claims(root_token, SESSION_REFRESH_PREFIX) + opened_root = open_session_refresh_token(root_token, KEYS, NOW) + assert isinstance(opened_root, OpenedSessionToken) + assert opened_root.family == opened_root.jti + + rotated = mint_session_refresh_token(PRINCIPAL, KEYS, NOW, family=opened_root.family) + assert isinstance(rotated, MintedSessionToken) + rotated_token = rotated.token.get_secret_value() + assert _decoded_claims(rotated_token, SESSION_REFRESH_PREFIX)["family"] == opened_root.jti + opened_rotated = open_session_refresh_token(rotated_token, KEYS, NOW) + assert isinstance(opened_rotated, OpenedSessionToken) + assert opened_rotated.jti != opened_root.jti + assert opened_rotated.family == opened_root.jti + + twice_rotated = mint_session_refresh_token(PRINCIPAL, KEYS, NOW, family=opened_rotated.family) + assert isinstance(twice_rotated, MintedSessionToken) + opened_twice = open_session_refresh_token(twice_rotated.token.get_secret_value(), KEYS, NOW) + assert isinstance(opened_twice, OpenedSessionToken) + assert opened_twice.family == opened_root.jti + + def test_signed_claims_with_an_unknown_audience_are_rejected(): token = _sign_claims(_valid_claims(audience="bogus")) assert isinstance(open_session_token(token, KEYS, NOW), SessionMalformed) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py b/tests/unit/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py index 2943ff4b74a..76d071b1852 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py @@ -712,6 +712,62 @@ async def test_single_use_guard_fails_closed_when_redis_errors(): assert await guard.claim("jti-fault", 60) == "unavailable" # fail closed, not a fallback count of 1 +@pytest.mark.asyncio +async def test_single_use_guard_peek_fails_closed_when_redis_errors(): + """The peek that gates a refresh on its chain's revocation marker must read the Redis client + itself: the cache wrapper's get swallows a fault into ``None``, which would make a revoked chain + look unclaimed for exactly as long as Redis is down.""" + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import _SingleUseGuard + + cache = DualCache() + cache.redis_cache = MagicMock() + cache.redis_cache.async_get_cache = AsyncMock(return_value=None) + cache.redis_cache.check_and_fix_namespace = MagicMock(side_effect=lambda key: key) + cache.redis_cache.init_async_client.return_value.get = AsyncMock(side_effect=ConnectionError("redis down")) + + guard = _SingleUseGuard(cache) + assert await guard.peek("family-fault") == "unavailable" + + cache.redis_cache.init_async_client.return_value.get = AsyncMock(return_value=b"1") + assert await guard.peek("family-fault") == "claimed" + + +@pytest.mark.asyncio +async def test_single_use_guard_peek_reads_the_key_under_the_namespace_claim_wrote(): + """A configured ``redis_namespace`` prefixes every key the cache wrapper writes, so the raw client + read behind peek must ask for the same prefixed key: a revocation written under ``ns:...`` and + peeked at the bare key would never be seen, and the revoked chain would keep renewing.""" + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import _SingleUseGuard + + stored: dict[str, int] = {} # mutable-ok: the fake Redis store the test inspects + + def namespaced(key: str) -> str: + return f"ns:{key}" + + async def increment(key: str, value: int, ttl: int) -> int: + stored[namespaced(key)] = stored.get(namespaced(key), 0) + value + return stored[namespaced(key)] + + async def raw_get(key: str) -> bytes | None: + return None if key not in stored else str(stored[key]).encode() + + cache = DualCache() + cache.redis_cache = MagicMock() + cache.redis_cache.check_and_fix_namespace = MagicMock(side_effect=namespaced) + cache.redis_cache.async_increment = AsyncMock(side_effect=increment) + cache.redis_cache.init_async_client.return_value.get = AsyncMock(side_effect=raw_get) + + guard = _SingleUseGuard(cache) + assert await guard.peek("family-ns") == "unclaimed" + assert await guard.claim("family-ns", 60) == "first" + assert set(stored) == {"ns:family-ns"} + assert await guard.peek("family-ns") == "claimed" + + LOOPBACK_REDIRECT_URI = "http://localhost:3118/callback" @@ -1850,12 +1906,15 @@ async def test_revoke_burns_the_refresh_token_and_answers_200_for_dead_or_unknow assert garbage.status_code == 200 -def _redis_that(async_increment): +def _redis_that(async_increment, get=None): from unittest.mock import AsyncMock, MagicMock cache = DualCache() cache.redis_cache = MagicMock() cache.redis_cache.async_increment = async_increment + cache.redis_cache.check_and_fix_namespace = MagicMock(side_effect=lambda key: key) + cache.redis_cache.init_async_client.return_value.get = get or AsyncMock(return_value=None) + cache.redis_cache.async_get_cache = AsyncMock(side_effect=AssertionError("peek must read the client, not the wrapper")) cache.async_increment_cache = AsyncMock(side_effect=AssertionError("must not fall back to in-memory")) return cache @@ -1903,6 +1962,64 @@ async def test_revoke_answers_503_while_the_shared_record_cannot_be_written_then assert already_burned.status_code == 200 +class _RedisThatFaultsOnTheSecondWrite: + """A shared Redis double whose first increment lands and whose second raises, the fault between + revoke's two writes; its raw client reads back exactly the keys that landed, and later writes land.""" + + def __init__(self) -> None: + self.landed: tuple[str, ...] = () + self.writes = 0 + + def check_and_fix_namespace(self, key: str) -> str: + return key + + def init_async_client(self) -> "_RedisThatFaultsOnTheSecondWrite": + return self + + async def async_increment(self, key: str, value: int, **kwargs: object) -> int: + self.writes += 1 + if self.writes == 2: + raise ConnectionError("redis down") + self.landed = (*self.landed, key) + return self.landed.count(key) + + async def get(self, key: str) -> int | None: + return self.landed.count(key) or None + + +@pytest.mark.asyncio +async def test_revoke_ends_the_chain_first_so_a_fault_before_the_jti_burn_leaves_no_descendant_renewable(): + """A revocation whose chain marker landed but whose ``jti`` burn faulted answers 503, and the + descendant a copy already rotated must be refused from that moment, not renewable until the + client retries; the retry then lands with only the ``jti`` left to burn.""" + client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"] + issued = DualCache() + payload = json.loads( + (await _redeem_native(await _native_code(client_id, cache=issued), client_id, _Minter(), cache=issued)).body + ) + copied_and_rotated = json.loads( + (await _refresh_native(payload["refresh_token"], client_id, _Minter(), issued)).body + ) + redis: Final = _RedisThatFaultsOnTheSecondWrite() + shared: Final = DualCache(redis_cache=redis) # pyright: ignore[reportArgumentType] # duck-typed Redis double + + half_written = await revoke_refresh_token( + token=payload["refresh_token"], client_id=client_id, master_key=MASTER_KEY, cache=shared + ) + assert half_written.status_code == 503 + + refused = await _refresh_native(copied_and_rotated["refresh_token"], client_id, _Minter(), shared) + assert refused.status_code == 400 + assert json.loads(refused.body)["error"] == "invalid_grant" + assert "revoked" in json.loads(refused.body)["error_description"] + + retried = await revoke_refresh_token( + token=payload["refresh_token"], client_id=client_id, master_key=MASTER_KEY, cache=shared + ) + assert retried.status_code == 200 + assert json.loads(retried.body) == {} + + @pytest.mark.asyncio async def test_refresh_answers_503_without_burning_the_token_while_redis_is_down(): """Fail closed, but say why: a refresh the shared backend could not record is refused with 503 @@ -1939,6 +2056,36 @@ async def test_refresh_answers_503_without_burning_the_token_while_redis_is_down assert json.loads(replayed.body)["error"] == "invalid_grant" +@pytest.mark.asyncio +async def test_refresh_of_a_rotated_token_answers_503_before_minting_while_redis_cannot_be_read(): + """A rotated token's chain marker is read before anything is minted or claimed; a Redis read fault + is a 503, never a pass: the token stays unburned and unminted until Redis answers.""" + from unittest.mock import AsyncMock + + client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"] + issued = DualCache() + payload = json.loads( + (await _redeem_native(await _native_code(client_id, cache=issued), client_id, _Minter(), cache=issued)).body + ) + rotated = json.loads( + (await _refresh_native(payload["refresh_token"], client_id, _Minter(), _redis_that(AsyncMock(return_value=1)))).body + )["refresh_token"] + + minter = _Minter() + claim = AsyncMock(return_value=1) + unreadable = await _refresh_native( + rotated, client_id, minter, _redis_that(claim, get=AsyncMock(side_effect=ConnectionError("redis down"))) + ) + assert unreadable.status_code == 503 + assert json.loads(unreadable.body)["error"] == "temporarily_unavailable" + assert minter.calls == [] + assert claim.await_count == 0 + + readable = await _refresh_native(rotated, client_id, _Minter(), _redis_that(AsyncMock(return_value=1))) + assert readable.status_code == 200 + assert json.loads(readable.body)["refresh_token"] != rotated + + @pytest.mark.asyncio async def test_revoke_from_another_client_leaves_the_token_usable(): client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"] @@ -1955,6 +2102,24 @@ async def test_revoke_from_another_client_leaves_the_token_usable(): assert refreshed.status_code == 200 +@pytest.mark.asyncio +async def test_revoke_of_an_already_rotated_token_ends_its_chain(): + client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"] + cache = DualCache() + payload = json.loads( + (await _redeem_native(await _native_code(client_id, cache=cache), client_id, _Minter(), cache=cache)).body + ) + rotated = json.loads((await _refresh_native(payload["refresh_token"], client_id, _Minter(), cache)).body) + revoked = await revoke_refresh_token( + token=payload["refresh_token"], client_id=client_id, master_key=MASTER_KEY, cache=cache + ) + assert revoked.status_code == 200 + refused = await _refresh_native(rotated["refresh_token"], client_id, _Minter(), cache) + assert refused.status_code == 400 + assert json.loads(refused.body)["error"] == "invalid_grant" + assert "revoked" in json.loads(refused.body)["error_description"] + + @pytest.mark.asyncio async def test_revoke_refuses_unknown_clients_and_a_missing_master_key(): client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"] @@ -2075,6 +2240,50 @@ async def test_introspect_refresh_token_goes_inactive_once_rotated(): assert (status, body) == (200, {"active": False}) +@pytest.mark.asyncio +async def test_refresh_replay_refuses_only_itself_and_leaves_the_session_pair_chain_renewing(): + """A second terminal of the same sign-in presents the refresh token the first one already + rotated (Claude Code renews from its in-memory copy, then recovers from the shared credential + file), so a replay refuses only itself: the live descendant keeps renewing, introspection keeps + reporting it active, and a token the gateway minted before chains were stamped (no ``family`` + claim) roots its own chain and rotates too.""" + keys, now, principal = _introspection_fixtures() + client_id = (await _register([REDIRECT_URI]))["client_id"] + cache = DualCache() + root = mint_session_refresh_token(SessionPrincipal(user_id="u1", client_id=client_id), keys, now) + + async def _refresh(token): + return await aggregate_token( + request=_request("/token", method="POST"), + grant_type="refresh_token", + code=None, + redirect_uri=None, + client_id=client_id, + code_verifier=None, + refresh_token=token, + master_key=MASTER_KEY, + reload_user=_reload_user_active, + cache=cache, + ) + + root_token = root.token.get_secret_value() + rotated = json.loads((await _refresh(root_token)).body)["refresh_token"] + twice_rotated = json.loads((await _refresh(rotated)).body)["refresh_token"] + replayed = await _refresh(root_token) + assert replayed.status_code == 400 + assert json.loads(replayed.body)["error"] == "invalid_grant" + assert "already used" in json.loads(replayed.body)["error_description"] + + renewed = await _refresh(twice_rotated) + assert renewed.status_code == 200 + thrice_rotated = json.loads(renewed.body)["refresh_token"] + status, body = await _introspect(thrice_rotated, cache=cache) + assert (status, body["active"]) == (200, True) + + unrelated = mint_session_refresh_token(principal.model_copy(update={"client_id": client_id}), keys, now) + assert (await _refresh(unrelated.token.get_secret_value())).status_code == 200 + + @pytest.mark.asyncio async def test_introspect_accepts_rs256_signed_tokens_under_configured_signing(monkeypatch): from cryptography.hazmat.primitives import serialization @@ -2137,6 +2346,21 @@ async def test_introspect_fails_closed_on_dead_user_and_503s_on_outage(): assert (status, body["error"]) == (500, "server_error") +@pytest.mark.asyncio +async def test_introspect_503s_while_the_single_use_record_cannot_be_read(): + """A token whose chain may have been revoked is never reported active on a Redis read fault.""" + from unittest.mock import AsyncMock + + keys, now, principal = _introspection_fixtures() + minted = mint_session_refresh_token(principal, keys, now) + unreadable = _redis_that(AsyncMock(return_value=1), get=AsyncMock(side_effect=ConnectionError("redis down"))) + status, body = await _introspect(minted.token.get_secret_value(), cache=unreadable) + assert (status, body["error"]) == (503, "temporarily_unavailable") + + status, body = await _introspect(minted.token.get_secret_value(), cache=_redis_that(AsyncMock(return_value=1))) + assert (status, body["active"]) == (200, True) + + @pytest.mark.asyncio @pytest.mark.parametrize( "auth_type", [None, "none", "api_key", "bearer_token", "basic", "authorization", "token", "aws_sigv4"] diff --git a/tests/unit/proxy/anthropic_endpoints/test_gateway_endpoints.py b/tests/unit/proxy/anthropic_endpoints/test_gateway_endpoints.py index 7c3e8f56a21..8ac0e2a6f6d 100644 --- a/tests/unit/proxy/anthropic_endpoints/test_gateway_endpoints.py +++ b/tests/unit/proxy/anthropic_endpoints/test_gateway_endpoints.py @@ -2,12 +2,14 @@ Tests for the Claude Code gateway protocol (anthropic_endpoints/gateway_endpoints.py). Covers the OAuth device-flow surface (RFC 8414 discovery, RFC 8628 device -authorization + token), managed settings, OTLP ingestion, and the enable flag. +authorization + token, the rotating refresh grant, RFC 7009 revocation), managed +settings, OTLP ingestion, and the enable flag. """ import asyncio from collections.abc import Iterator, Mapping from contextlib import ExitStack, contextmanager +from datetime import datetime, timezone from types import MappingProxyType from typing import Final from unittest.mock import AsyncMock, MagicMock, patch @@ -18,6 +20,21 @@ from fastapi import FastAPI from fastapi.testclient import TestClient from litellm.caching.dual_cache import DualCache +from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import ( + MintedProxyCredential, + ProxyCredentialMintFailure, +) +from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credentials import ( + SessionRefreshOpened, + open_session_refresh_bearer, + session_keys_from_master_key, +) +from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token import ( + SESSION_REFRESH_PREFIX, + MintedSessionToken, + SessionPrincipal, + mint_session_refresh_token, +) from litellm.proxy._types import ProxyException from litellm.proxy.anthropic_endpoints import gateway_endpoints from litellm.proxy.management_endpoints.ui_sso import ( @@ -29,6 +46,8 @@ from litellm.proxy.middleware.prometheus_auth_middleware import PrometheusAuthMi _DEVICE_CODE_GRANT: Final = "urn:ietf:params:oauth:grant-type:device_code" _MASTER_KEY: Final = "sk-master-key" +_TOKEN_URL: Final = "/claude_code_gateway/oauth/token" +_REVOKE_URL: Final = "/claude_code_gateway/oauth/revoke" _SHARED_LOGIN_ID: Final = "cli-shared-login-code" _SHARED_POLL_SECRET: Final = "shared-poll-secret" _SHARED_DEVICE_CODE: Final = f"{_SHARED_LOGIN_ID}.{_SHARED_POLL_SECRET}" @@ -53,41 +72,75 @@ _COMPLETED_SESSION: Final = MappingProxyType( class _SharedRedisFake: + """A Redis double with a configured namespace, prefixing every key its wrapper methods touch the + way ``RedisCache`` does, while its raw client answers only the exact key it is asked for.""" + + namespace: Final = "qa" + def __init__(self) -> None: self.values: Mapping[str, object] = MappingProxyType({}) self.counters: Mapping[str, float] = MappingProxyType({}) + def check_and_fix_namespace(self, key: str) -> str: + return key if key.startswith(f"{self.namespace}:") else f"{self.namespace}:{key}" + def set_cache(self, key: str, value: object, **kwargs: object) -> None: - self.values = MappingProxyType({**self.values, key: value}) + self.values = MappingProxyType({**self.values, self.check_and_fix_namespace(key): value}) def get_cache(self, key: str, **kwargs: object) -> object: - return self.values.get(key) + return self.values.get(self.check_and_fix_namespace(key)) def delete_cache(self, key: str) -> None: - self.values = MappingProxyType({name: value for name, value in self.values.items() if name != key}) + gone: Final = self.check_and_fix_namespace(key) + self.values = MappingProxyType({name: value for name, value in self.values.items() if name != gone}) async def async_delete_cache(self, key: str) -> None: self.delete_cache(key) async def async_increment(self, key: str, value: float, **kwargs: object) -> float: - incremented: Final = self.counters.get(key, 0) + value - self.counters = MappingProxyType({**self.counters, key: incremented}) + namespaced: Final = self.check_and_fix_namespace(key) + incremented: Final = self.counters.get(namespaced, 0) + value + self.counters = MappingProxyType({**self.counters, namespaced: incremented}) return incremented + async def async_get_cache(self, key: str, **kwargs: object) -> object: + return self.values.get(self.check_and_fix_namespace(key)) + + def init_async_client(self) -> "_SharedRedisClientFake": + return _SharedRedisClientFake(self) + + +class _SharedRedisClientFake: + def __init__(self, store: _SharedRedisFake) -> None: + self._store = store + + async def get(self, key: str) -> object: + return self._store.counters.get(key, self._store.values.get(key)) + def _replica(redis: _SharedRedisFake) -> DualCache: return DualCache(redis_cache=redis, default_in_memory_ttl=600) # pyright: ignore[reportArgumentType] # duck-typed Redis double +class _Minter: + def __init__(self, failure: ProxyCredentialMintFailure | None = None) -> None: + self.failure: Final = failure + self.calls: tuple[tuple[str, str | None], ...] = () + + async def __call__(self, user_id: str, team_id: str | None) -> MintedProxyCredential | ProxyCredentialMintFailure: + self.calls = (*self.calls, (user_id, team_id)) + if self.failure is not None: + return self.failure + return MintedProxyCredential(key=f"sk-cli-{user_id}", expires_in=7200, user_id=user_id, team_id=team_id) + + def _real_auth_proxy_attrs() -> Mapping[str, object]: proxy_logging_obj: Final = MagicMock() proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock() proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) return MappingProxyType( { - "master_key": _MASTER_KEY, "prisma_client": None, - "user_api_key_cache": DualCache(), "proxy_logging_obj": proxy_logging_obj, "llm_router": None, "llm_model_list": [], @@ -106,6 +159,8 @@ def _gateway_env( enabled: bool = True, managed_settings: Mapping[str, object] | None = None, cache: DualCache | None = None, + claim_cache: DualCache | None = None, + minter: _Minter | None = None, real_auth: bool = False, extra_settings: Mapping[str, object] = MappingProxyType({}), ) -> Iterator[tuple[TestClient, DualCache]]: @@ -119,6 +174,7 @@ def _gateway_env( app: Final = FastAPI() app.add_middleware(PrometheusAuthMiddleware) app.include_router(gateway_endpoints.router) + app.dependency_overrides[gateway_endpoints._proxy_credential_minter] = lambda: minter or _Minter() async def _fake_auth() -> object: return object() @@ -134,6 +190,21 @@ def _gateway_env( "litellm.proxy.proxy_server.cli_sso_session_cache", session_cache ) ) + stack.enter_context( + patch( # test-quality-ok: the refresh token is signed under this proxy_server module global + "litellm.proxy.proxy_server.master_key", _MASTER_KEY + ) + ) + stack.enter_context( + patch( # test-quality-ok: the single-use record lives in this proxy_server module global + "litellm.proxy.proxy_server.user_api_key_cache", claim_cache or DualCache() + ) + ) + stack.enter_context( + patch( # test-quality-ok: the single-use guard prefers this proxy_server module global when set + "litellm.proxy.proxy_server.redis_usage_cache", None + ) + ) if real_auth: for name, value in _real_auth_proxy_attrs().items(): stack.enter_context(patch(f"litellm.proxy.proxy_server.{name}", value)) @@ -148,10 +219,43 @@ def _start_device_flow(client: TestClient) -> str: def _request_token(client: TestClient, device_code: str) -> httpx.Response: - return client.post( - "/claude_code_gateway/oauth/token", - data={"grant_type": _DEVICE_CODE_GRANT, "device_code": device_code}, + return client.post(_TOKEN_URL, data={"grant_type": _DEVICE_CODE_GRANT, "device_code": device_code}) + + +def _refresh(client: TestClient, refresh_token: str) -> httpx.Response: + return client.post(_TOKEN_URL, data={"grant_type": "refresh_token", "refresh_token": refresh_token}) + + +def _revoke(client: TestClient, token: str, hint: str | None = None) -> httpx.Response: + return client.post(_REVOKE_URL, data={"token": token, **({} if hint is None else {"token_type_hint": hint})}) + + +def _signed_in(client: TestClient, cache: DualCache) -> Mapping[str, object]: + device_code: Final = _start_device_flow(client) + _complete_flow(cache, device_code) + with patch(_MINT, return_value="sk-litellm-session-token"): + resp: Final = _request_token(client, device_code) + assert resp.status_code == 200 + return MappingProxyType(resp.json()) + + +def _opened_refresh(refresh_token: str, client_id: str = gateway_endpoints.CLAUDE_CODE_CLIENT_ID) -> SessionPrincipal: + opened: Final = open_session_refresh_bearer( + refresh_token, + session_keys_from_master_key(_MASTER_KEY), + datetime.now(timezone.utc), + expected_client_id=client_id, ) + assert isinstance(opened, SessionRefreshOpened) + return opened.principal + + +def _foreign_refresh_token(principal: SessionPrincipal) -> str: + minted: Final = mint_session_refresh_token( + principal, session_keys_from_master_key(_MASTER_KEY), datetime.now(timezone.utc) + ) + assert isinstance(minted, MintedSessionToken) + return minted.token.get_secret_value() def _completed_flow(session_data: Mapping[str, object] = _COMPLETED_SESSION) -> dict[str, object]: @@ -168,9 +272,7 @@ def _login_id(device_code: str) -> str: return device_code.partition(".")[0] -def _complete_flow( - cache: DualCache, device_code: str, session_data: Mapping[str, object] = _COMPLETED_SESSION -) -> None: +def _complete_flow(cache: DualCache, device_code: str, session_data: Mapping[str, object] = _COMPLETED_SESSION) -> None: key: Final = _get_cli_sso_flow_cache_key(_login_id(device_code)) flow: Final = cache.get_cache(key=key) assert isinstance(flow, dict) @@ -185,15 +287,17 @@ def test_discovery_shape(): body = resp.json() assert body["device_authorization_endpoint"].endswith("/claude_code_gateway/oauth/device_authorization") assert body["token_endpoint"].endswith("/claude_code_gateway/oauth/token") + assert body["revocation_endpoint"].endswith("/claude_code_gateway/oauth/revoke") assert body["grant_types_supported"] == [ "urn:ietf:params:oauth:grant-type:device_code", "refresh_token", ] # authorization_endpoint is intentionally absent (device flow only). assert "authorization_endpoint" not in body - # Both endpoints must be same-origin with the issuer. + # Every endpoint must be same-origin with the issuer, or Claude Code ignores it. assert body["device_authorization_endpoint"].startswith(body["issuer"]) assert body["token_endpoint"].startswith(body["issuer"]) + assert body["revocation_endpoint"].startswith(body["issuer"]) def test_discovery_404_when_disabled(): @@ -270,10 +374,16 @@ def test_token_success_mints_bearer_and_is_single_use(): with patch(_MINT, return_value="sk-litellm-session-token") as mint: resp = _request_token(client, device_code) assert resp.status_code == 200 + assert resp.headers["cache-control"] == "no-store" body = resp.json() assert body["access_token"] == "sk-litellm-session-token" assert body["token_type"] == "Bearer" assert body["expires_in"] > 0 + assert body["refresh_token"].startswith(SESSION_REFRESH_PREFIX) + principal = _opened_refresh(body["refresh_token"]) + assert principal.user_id == "user-123" + assert principal.team_id == "team-a" + assert principal.audience == "proxy_api" called_user = mint.call_args.kwargs["user_info"] assert called_user.user_id == "user-123" @@ -332,6 +442,22 @@ def test_token_mint_failure_leaves_the_login_unconsumed(): assert retry.json()["access_token"] == "sk-session" +def test_token_refresh_mint_failure_leaves_the_login_unconsumed(): + oversized_user: Final = {**_COMPLETED_SESSION, "user_id": "u" * 5000} + with _gateway_env() as (client, cache): + device_code = _start_device_flow(client) + _complete_flow(cache, device_code, session_data=oversized_user) + with patch(_MINT, return_value="sk-session"): + failed = _request_token(client, device_code) + _complete_flow(cache, device_code) + with patch(_MINT, return_value="sk-session"): + retry = _request_token(client, device_code) + assert failed.status_code == 500 + assert failed.json()["error"] == "server_error" + assert retry.status_code == 200 + assert retry.json()["refresh_token"].startswith(SESSION_REFRESH_PREFIX) + + def test_token_unknown_team_grants_is_invalid_grant(): with _gateway_env() as (client, cache): device_code = _start_device_flow(client) @@ -374,14 +500,230 @@ def test_token_unknown_device_code_is_expired_token(): assert resp.json()["error"] == "expired_token" -def test_refresh_grant_forces_relogin(): - with _gateway_env() as (client, _): - resp = client.post( - "/claude_code_gateway/oauth/token", - data={"grant_type": "refresh_token", "refresh_token": "whatever"}, +def test_token_signs_in_without_a_refresh_token_when_session_signing_is_unusable(): + """An ``mcp_session_token_signing`` block whose key reference does not resolve must not turn the device + sign-in that worked before refresh tokens existed into a 500: the login still answers the session token + and consumes the device code, only the refresh token is left out, and the refresh grant stays fail-closed.""" + unusable_signing: Final = {"algorithm": "RS256", "kid": "rotated", "private_key": "os.environ/LITELLM_TEST_NO_PEM"} + with _gateway_env(extra_settings={"mcp_session_token_signing": unusable_signing}) as (client, cache): + device_code = _start_device_flow(client) + _complete_flow(cache, device_code) + with patch(_MINT, return_value="sk-litellm-session-token"): + resp = _request_token(client, device_code) + assert resp.status_code == 200 + assert resp.headers["cache-control"] == "no-store" + body = resp.json() + assert body["access_token"] == "sk-litellm-session-token" + assert body["token_type"] == "Bearer" + assert body["expires_in"] > 0 + assert "refresh_token" not in body + replay = _request_token(client, device_code) + assert replay.status_code == 400 + assert replay.json()["error"] == "expired_token" + principal = SessionPrincipal( + user_id="user-123", client_id=gateway_endpoints.CLAUDE_CODE_CLIENT_ID, audience="proxy_api", team_id="team-a" ) - assert resp.status_code == 401 + refused = _refresh(client, _foreign_refresh_token(principal)) + assert refused.status_code == 500 + assert refused.json()["error"] == "server_error" + + +def test_refresh_grant_mints_a_new_bearer_and_rotates_the_refresh_token(): + """Claude Code refreshes with grant_type and refresh_token alone, no client_id: the gateway + re-mints the session JWT from the live user row for the team the login picked, hands back a + fresh refresh token, and the presented one is dead from then on (rotation). The mint runs + before the single-use claim on every presentation, a replay included, so a transient mint + failure never burns a still-valid token; the chain itself keeps rotating until a replay.""" + minter: Final = _Minter() + with _gateway_env(minter=minter) as (client, cache): + signed_in = _signed_in(client, cache) + refreshed = _refresh(client, str(signed_in["refresh_token"])) + assert refreshed.status_code == 200 + assert refreshed.headers["cache-control"] == "no-store" + rotated = refreshed.json() + assert minter.calls == (("user-123", "team-a"),) + assert rotated["access_token"] == "sk-cli-user-123" + assert rotated["token_type"] == "Bearer" + assert rotated["expires_in"] == 7200 + assert rotated["refresh_token"] != signed_in["refresh_token"] + assert _opened_refresh(rotated["refresh_token"]).team_id == "team-a" + + chained = _refresh(client, rotated["refresh_token"]) + assert chained.status_code == 200 + assert chained.json()["refresh_token"] != rotated["refresh_token"] + + replayed = _refresh(client, str(signed_in["refresh_token"])) + assert replayed.status_code == 400 + assert replayed.json()["error"] == "invalid_grant" + assert "already used" in replayed.json()["error_description"] + assert "access_token" not in replayed.json() + assert minter.calls == (("user-123", "team-a"),) * 3 + + +def test_refresh_grant_replay_is_refused_on_a_replica_that_did_not_serve_the_rotation(): + redis: Final = _SharedRedisFake() + minter: Final = _Minter() + with _gateway_env(claim_cache=_replica(redis), minter=minter) as (client, cache): + refresh_token = str(_signed_in(client, cache)["refresh_token"]) + assert _refresh(client, refresh_token).status_code == 200 + with _gateway_env(claim_cache=_replica(redis), minter=minter) as (client, _): + replayed = _refresh(client, refresh_token) + assert replayed.status_code == 400 + assert replayed.json()["error"] == "invalid_grant" + assert len(minter.calls) == 2 + + +def test_refresh_grant_replay_refuses_only_itself_so_a_second_terminal_recovers_into_a_live_chain(): + """Claude Code renews from its in-memory copy of the credential, so a second terminal on the + same machine presents the refresh token the first one already rotated and then recovers from + the shared credential file. The replay is refused as already used and nothing else happens: + the rotation the first terminal saved keeps renewing, as does its own rotation after it, so + neither terminal is signed out. Only a sign-out ends the chain.""" + minter: Final = _Minter() + with _gateway_env(minter=minter) as (client, cache): + first = str(_signed_in(client, cache)["refresh_token"]) + rotated = str(_refresh(client, first).json()["refresh_token"]) + replayed = _refresh(client, first) + assert replayed.status_code == 400 + assert replayed.json()["error"] == "invalid_grant" + assert "already used" in replayed.json()["error_description"] + + renewed = _refresh(client, rotated) + assert renewed.status_code == 200 + assert "access_token" in renewed.json() + renewed_again = _refresh(client, str(renewed.json()["refresh_token"])) + assert renewed_again.status_code == 200 + assert len(minter.calls) == 4 + + +@pytest.mark.parametrize( + "principal", + [ + SessionPrincipal(user_id="user-123", client_id="llm_dcrc_other", audience="proxy_api", team_id="team-a"), + SessionPrincipal(user_id="user-123", client_id=gateway_endpoints.CLAUDE_CODE_CLIENT_ID, team_id="team-a"), + ], + ids=["dcr_client", "identity_only_audience"], +) +def test_refresh_grant_refuses_a_token_minted_for_another_client_or_audience(principal: SessionPrincipal): + minter: Final = _Minter() + with _gateway_env(minter=minter) as (client, _): + resp = _refresh(client, _foreign_refresh_token(principal)) + assert resp.status_code == 400 assert resp.json()["error"] == "invalid_grant" + assert minter.calls == () + + +@pytest.mark.parametrize("refresh_token", ["", "not-a-refresh-token"]) +def test_refresh_grant_without_a_usable_token_mints_nothing(refresh_token: str): + minter: Final = _Minter() + with _gateway_env(minter=minter) as (client, _): + resp = _refresh(client, refresh_token) + assert resp.status_code == 400 + assert resp.json()["error"] == ("invalid_request" if refresh_token == "" else "invalid_grant") + assert minter.calls == () + + +@pytest.mark.parametrize( + "failure, status, error", + [ + ("no_active_key", 400, "invalid_grant"), + ("not_a_member", 400, "invalid_grant"), + ("unavailable", 503, "temporarily_unavailable"), + ], +) +def test_refresh_grant_mint_refusal_leaves_the_refresh_token_usable( + failure: ProxyCredentialMintFailure, status: int, error: str +): + """A deactivated user, a lost team membership, or a DB outage refuses the renewal without + burning the presented token, so a transient outage never forces a re-login.""" + claim_cache: Final = DualCache() + with _gateway_env(claim_cache=claim_cache, minter=_Minter(failure)) as (client, cache): + refresh_token = str(_signed_in(client, cache)["refresh_token"]) + refused = _refresh(client, refresh_token) + assert refused.status_code == status + assert refused.json()["error"] == error + assert "refresh_token" not in refused.json() + with _gateway_env(claim_cache=claim_cache, minter=_Minter()) as (client, _): + recovered = _refresh(client, refresh_token) + assert recovered.status_code == 200 + + +def test_revoke_burns_the_refresh_token_and_answers_200_for_every_other_token(): + """Claude Code's /logout posts the session JWT and then the refresh token with a hint, each as + a form, and tolerates nothing but a 2xx: the refresh token is dead afterwards, while the + stateless JWT, an already-dead token, and garbage all answer 200 (RFC 7009 section 2.2).""" + minter: Final = _Minter() + with _gateway_env(minter=minter) as (client, cache): + signed_in = _signed_in(client, cache) + jwt_revoked = _revoke(client, str(signed_in["access_token"])) + assert jwt_revoked.status_code == 200 + assert jwt_revoked.json() == {} + assert jwt_revoked.headers["cache-control"] == "no-store" + assert _refresh(client, str(signed_in["refresh_token"])).status_code == 200 + + rotated = _refresh(client, str(signed_in["refresh_token"])) + assert rotated.status_code == 400 + fresh = str(_signed_in(client, cache)["refresh_token"]) + revoked = _revoke(client, fresh, hint="refresh_token") + assert revoked.status_code == 200 + refused = _refresh(client, fresh) + assert refused.status_code == 400 + assert refused.json()["error"] == "invalid_grant" + assert _revoke(client, fresh, hint="refresh_token").status_code == 200 + assert _revoke(client, "nonsense").status_code == 200 + missing = client.post(_REVOKE_URL, data={"token_type_hint": "refresh_token"}) + assert missing.status_code == 400 + assert missing.json()["error"] == "invalid_request" + assert len(minter.calls) == 2 + + +def test_revoke_from_another_client_leaves_the_refresh_token_usable(): + foreign: Final = SessionPrincipal( + user_id="user-123", client_id="llm_dcrc_other", audience="proxy_api", team_id="team-a" + ) + with _gateway_env() as (client, cache): + own = str(_signed_in(client, cache)["refresh_token"]) + assert _revoke(client, _foreign_refresh_token(foreign)).status_code == 200 + assert _refresh(client, own).status_code == 200 + + +def test_revoke_ends_the_rotation_chain_of_the_presented_token(): + """``/logout`` posts the refresh token Claude Code holds. When a copy of that token already + rotated it, the copy's live descendant must die with the sign-out (RFC 7009 section 2.1 revokes + the grant, not one token), instead of the single-use record only confirming what the copy did.""" + with _gateway_env(minter=_Minter()) as (client, cache): + stored = str(_signed_in(client, cache)["refresh_token"]) + copied_and_rotated = str(_refresh(client, stored).json()["refresh_token"]) + assert _revoke(client, stored).status_code == 200 + refused = _refresh(client, copied_and_rotated) + assert refused.status_code == 400 + assert refused.json()["error"] == "invalid_grant" + assert "revoked" in refused.json()["error_description"] + + +def test_revoke_on_one_replica_ends_the_chain_on_another_through_a_namespaced_redis(): + """The sign-out's revocation marker is written through the cache wrapper, which prefixes the key + with the configured ``redis_namespace``; the replica serving the next renewal must look for the + marker under that same prefix, or the revoked chain keeps renewing everywhere but where it signed + out.""" + redis: Final = _SharedRedisFake() + minter: Final = _Minter() + with _gateway_env(claim_cache=_replica(redis), minter=minter) as (client, cache): + stored = str(_signed_in(client, cache)["refresh_token"]) + copied_and_rotated = str(_refresh(client, stored).json()["refresh_token"]) + assert _revoke(client, stored).status_code == 200 + with _gateway_env(claim_cache=_replica(redis), minter=minter) as (client, _): + refused = _refresh(client, copied_and_rotated) + assert refused.status_code == 400 + assert refused.json()["error"] == "invalid_grant" + assert "revoked" in refused.json()["error_description"] + assert len(minter.calls) == 1 + + +def test_revoke_404_when_gateway_disabled(): + with _gateway_env(enabled=False) as (client, _): + resp = _revoke(client, "anything") + assert resp.status_code == 404 def test_unsupported_grant_type(): @@ -409,9 +751,7 @@ def test_managed_settings_returns_client_envelope_and_304_on_cached_checksum(): assert body["uuid"] == checksum assert resp.headers["ETag"] == f'"{checksum}"' - not_modified = client.get( - "/claude_code_gateway/managed/settings", headers={"If-None-Match": f'"{checksum}"'} - ) + not_modified = client.get("/claude_code_gateway/managed/settings", headers={"If-None-Match": f'"{checksum}"'}) assert not_modified.status_code == 304 assert not_modified.headers["ETag"] == f'"{checksum}"'