mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
a20cfc5f0e
commit
238609ea5d
2 changed files with 143 additions and 21 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue