mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
feat(claude_code_gateway): issue rotating refresh tokens and a revocation endpoint (#44635)
* feat(claude_code_gateway): issue rotating refresh tokens and a revocation endpoint The Claude Code gateway's device-code grant now returns a refresh token alongside the session JWT, so a Claude Code session renews itself before the JWT expires instead of forcing the user back through the browser sign-in. grant_type=refresh_token re-mints the JWT from the live user row, rotates the refresh token, and refuses a replayed, foreign, or identity-only token with invalid_grant. The discovery document now advertises an RFC 7009 revocation endpoint, which Claude Code's /logout calls with both tokens, so sign-out burns the refresh token. The refresh token is the MCP gateway's sealed session refresh token bound to the fixed client id claude_code, so both front doors share one single-use record. * fix(claude_code_gateway): mint the whole credential before claiming the device code A refresh token that failed to mint answered 500 after the device code was already claimed and the login deleted, so the client could not redeem the completed sign-in again. The response is now built first, and the code is claimed only when it is a 200. * fix(claude-code-gateway): keep device sign-in working when session signing is unusable * feat(claude-code-gateway): end the whole refresh chain on a replay or a revocation * fix(mcp-gateway): fail the single-use peek closed on a Redis fault * fix(mcp-gateway): read the single-use marker under the cache namespace * fix(mcp-gateway): refuse a replayed refresh token without ending its chain A refresh token presented a second time is refused as already used and nothing else happens to the chain it was rotated from. Claude Code renews from its in-memory copy of the credential, so a second terminal on the same machine presents the token the first terminal already rotated and then recovers from the shared credential file; ending the chain there would sign both terminals out at every expiry. Only a revocation ends a chain, so the 503-on-unrecorded-chain-ending path and its tests go away. * fix(mcp-gateway): end the refresh chain before burning the revoked token --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
parent
e98bbc2f8e
commit
62dee3d730
7 changed files with 952 additions and 132 deletions
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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://<proxy-host>/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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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}"'
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue