fix(mcp): answer 503 when a refresh-token burn cannot be recorded

revoke_refresh_token discarded the single-use claim result, so a revocation that arrived while Redis was unreachable answered 200 and left the refresh token live. The token endpoint reported the same outage as invalid_grant "already used". The guard now reports first, replayed, or unavailable, and both endpoints answer 503 temporarily_unavailable for an outage (RFC 7009 section 2.2.1, RFC 6749 section 5.2), which the CLI surfaces as a one-line warning while keeping the key it has
This commit is contained in:
mateo-berri 2026-08-20 07:56:38 -07:00
parent a20cfc5f0e
commit 238609ea5d
2 changed files with 143 additions and 21 deletions

View file

@ -723,10 +723,16 @@ async def complete_connect_flow(
return _oauth_error(401, "login_required", "sign in to LiteLLM to finish connecting")
if session_user_id != flow.user_id:
return _oauth_error(403, "access_denied", "the signed-in user does not match this connect flow")
if not await _SingleUseGuard(cache).claim(
f"{_USED_FLOW_CACHE_PREFIX}{flow.jti}", CONNECT_FLOW_TTL_SECONDS + _CLAIM_TTL_BUFFER_SECONDS
):
return _oauth_error(400, "invalid_request", "this connect flow was already completed; restart the connection")
flow_refusal: Final = _claim_refusal(
await _SingleUseGuard(cache).claim(
f"{_USED_FLOW_CACHE_PREFIX}{flow.jti}", CONNECT_FLOW_TTL_SECONDS + _CLAIM_TTL_BUFFER_SECONDS
),
replayed=_oauth_error(
400, "invalid_request", "this connect flow was already completed; restart the connection"
),
)
if flow_refusal is not None:
return flow_refusal
response: Final = (
_denied_flow_response(flow) if decision == "deny" else _approved_flow_response(flow, delivery, team_id, now)
)
@ -806,6 +812,27 @@ def _pkce_verifier_matches(code_verifier: str, code_challenge: str) -> bool:
return hmac.compare_digest(computed, code_challenge.encode("utf-8"))
ClaimOutcome = Literal["first", "replayed", "unavailable"]
_CLAIM_UNAVAILABLE_DESCRIPTION: Final = "the single-use record is unavailable right now; try again shortly"
def _claim_refusal(outcome: ClaimOutcome, replayed: Response) -> Response | None:
"""A claim that is not the first caller's is refused, but the two reasons must stay apart on the
wire: a replay is the grant's own 4xx, while a shared backend that could not record the claim is
a 503 (RFC 7009 section 2.2.1, RFC 6749 section 5.2 ``temporarily_unavailable``), so the client
keeps the still-valid token and retries instead of being told it was already used."""
match outcome:
case "first":
return None
case "replayed":
return replayed
case "unavailable":
return _oauth_error(503, "temporarily_unavailable", _CLAIM_UNAVAILABLE_DESCRIPTION)
case _:
assert_never(outcome)
class _SingleUseGuard:
"""Atomic single-use claim for a one-time id (an auth-code, connect-flow ``jti``, or refresh-token
``jti``) over the injected proxy cache.
@ -829,9 +856,10 @@ class _SingleUseGuard:
def __init__(self, cache: DualCache) -> None:
self._cache = cache
async def claim(self, key: str, ttl_seconds: int) -> bool:
"""Atomically claim ``key``. ``True`` iff this caller is the first (increment to 1); ``False``
on a replay (>1) or when the claim could not be recorded in the shared backend (fail closed)."""
async def claim(self, key: str, ttl_seconds: int) -> ClaimOutcome:
"""Atomically claim ``key``. ``"first"`` iff this caller is the first (increment to 1),
``"replayed"`` on a replay (>1), and ``"unavailable"`` when the claim could not be recorded in
the shared backend, which every caller treats as a refusal (fail closed)."""
from litellm.proxy.proxy_server import redis_usage_cache # noqa: PLC0415 # circular import at module load
# Resolve the shared authority HERE rather than trusting the injected cache: callers pass
@ -850,11 +878,11 @@ class _SingleUseGuard:
verbose_logger.warning(
"mcp gateway single-use claim: shared cache backend unavailable, failing closed: %s", e
)
return False
return count == 1
return "unavailable"
return "first" if count == 1 else "replayed"
# No shared backend configured (single-replica): the in-memory increment is authoritative.
count = await self._cache.async_increment_cache(key, 1, ttl=ttl_seconds, local_only=True)
return count == 1
return "first" if count == 1 else "replayed"
def _session_token_pair(principal: SessionPrincipal, keys: SessionKeys, now: datetime) -> Response:
@ -1046,8 +1074,9 @@ class _GrantIssuer:
failure: Final = await self._reload_user(principal.user_id)
if failure is not None:
return _reload_failure_response(failure)
if not await self._guard.claim(claim_key, claim_ttl_seconds):
return _oauth_error(400, "invalid_grant", replayed)
refusal: Final = await self._claim_refusal(claim_key, claim_ttl_seconds, replayed)
if refusal is not None:
return refusal
return _session_token_pair(principal, self._keys, self._now)
async def _issue_proxy_credential(
@ -1060,10 +1089,16 @@ 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)
if not await self._guard.claim(claim_key, claim_ttl_seconds):
return _oauth_error(400, "invalid_grant", replayed)
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)
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,
@ -1138,7 +1173,10 @@ async def revoke_refresh_token(token: str, client_id: str, master_key: str | Non
``jti`` so neither the holder nor a thief can rotate it again. 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."""
dead token still answers 200; only an unknown client is refused. A live token whose
burn could not be recorded in the shared backend answers 503 (section 2.2.1), so the
client knows the token still stands and retries instead of reporting a logout that
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:
@ -1148,7 +1186,9 @@ async def revoke_refresh_token(token: str, client_id: str, master_key: str | Non
now: Final = datetime.now(timezone.utc)
opened: Final = open_session_refresh_bearer(token, keys, now, expected_client_id=client_id)
if isinstance(opened, SessionRefreshOpened):
_ = await _SingleUseGuard(cache).claim(
burned: Final = await _SingleUseGuard(cache).claim(
f"{_USED_REFRESH_CACHE_PREFIX}{opened.jti}", SESSION_REFRESH_TTL_SECONDS + _CLAIM_TTL_BUFFER_SECONDS
)
if burned == "unavailable":
return _oauth_error(503, "temporarily_unavailable", _CLAIM_UNAVAILABLE_DESCRIPTION)
return Response(content="{}", media_type="application/json", headers=TOKEN_NO_CACHE_HEADERS)

View file

@ -560,8 +560,8 @@ async def test_single_use_guard_in_memory_is_single_use_within_process():
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import _SingleUseGuard
guard = _SingleUseGuard(DualCache()) # redis_cache is None
assert await guard.claim("jti-inmem", 60) is True
assert await guard.claim("jti-inmem", 60) is False # replay of the same id
assert await guard.claim("jti-inmem", 60) == "first"
assert await guard.claim("jti-inmem", 60) == "replayed"
@pytest.mark.asyncio
@ -579,9 +579,9 @@ async def test_single_use_guard_uses_redis_as_sole_authority_when_configured():
cache.async_increment_cache = AsyncMock(side_effect=AssertionError("must not fall back to in-memory"))
guard = _SingleUseGuard(cache)
assert await guard.claim("jti-redis", 60) is True
assert await guard.claim("jti-redis", 60) == "first"
cache.redis_cache.async_increment = AsyncMock(return_value=2)
assert await guard.claim("jti-redis", 60) is False # Redis says 2 → replay
assert await guard.claim("jti-redis", 60) == "replayed"
@pytest.mark.asyncio
@ -599,7 +599,7 @@ async def test_single_use_guard_fails_closed_when_redis_errors():
cache.async_increment_cache = AsyncMock(return_value=1) # would fail OPEN if the guard fell back
guard = _SingleUseGuard(cache)
assert await guard.claim("jti-fault", 60) is False # fail closed, not a fallback count of 1
assert await guard.claim("jti-fault", 60) == "unavailable" # fail closed, not a fallback count of 1
LOOPBACK_REDIRECT_URI = "http://localhost:3118/callback"
@ -1531,6 +1531,88 @@ 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):
from unittest.mock import AsyncMock, MagicMock
cache = DualCache()
cache.redis_cache = MagicMock()
cache.redis_cache.async_increment = async_increment
cache.async_increment_cache = AsyncMock(side_effect=AssertionError("must not fall back to in-memory"))
return cache
@pytest.mark.asyncio
async def test_revoke_answers_503_while_the_shared_record_cannot_be_written_then_burns_the_token():
"""A revocation whose single-use marker never reached Redis must not report success: the token is
still redeemable on every worker, so the client has to hear 503 (RFC 7009 2.2.1) and retry."""
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
)
redis_down = await revoke_refresh_token(
token=payload["refresh_token"],
client_id=client_id,
master_key=MASTER_KEY,
cache=_redis_that(AsyncMock(side_effect=ConnectionError("redis down"))),
)
assert redis_down.status_code == 503
assert json.loads(redis_down.body) == {
"error": "temporarily_unavailable",
"error_description": "the single-use record is unavailable right now; try again shortly",
}
assert redis_down.headers["cache-control"] == "no-store"
redis_back = await revoke_refresh_token(
token=payload["refresh_token"],
client_id=client_id,
master_key=MASTER_KEY,
cache=_redis_that(AsyncMock(return_value=1)),
)
assert redis_back.status_code == 200
assert json.loads(redis_back.body) == {}
already_burned = await revoke_refresh_token(
token=payload["refresh_token"],
client_id=client_id,
master_key=MASTER_KEY,
cache=_redis_that(AsyncMock(return_value=2)),
)
assert already_burned.status_code == 200
@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
rather than ``invalid_grant``, so the CLI keeps the key it has and retries instead of telling the
user the token was already used and sending them back through the browser."""
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
)
redis_down = await _refresh_native(
payload["refresh_token"], client_id, _Minter(), _redis_that(AsyncMock(side_effect=ConnectionError("redis down")))
)
assert redis_down.status_code == 503
assert json.loads(redis_down.body)["error"] == "temporarily_unavailable"
assert "refresh_token" not in json.loads(redis_down.body)
redis_back = await _refresh_native(payload["refresh_token"], client_id, _Minter(), _redis_that(AsyncMock(return_value=1)))
assert redis_back.status_code == 200
assert json.loads(redis_back.body)["refresh_token"] != payload["refresh_token"]
replayed = await _refresh_native(payload["refresh_token"], client_id, _Minter(), _redis_that(AsyncMock(return_value=2)))
assert replayed.status_code == 400
assert json.loads(replayed.body)["error"] == "invalid_grant"
@pytest.mark.asyncio
async def test_revoke_from_another_client_leaves_the_token_usable():
client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"]