diff --git a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py index bc1dea50a74..314c80adbc4 100644 --- a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py +++ b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py @@ -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) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py index 6c2b1274aa4..761f823076b 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py @@ -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"]