fix(mcp): close the burn-before-check gate for both grants and validate master_key first

Follow-up to the pre-exchange identity gate, which I had only added to the
authorization_code branch and which left the master_key check inside the mint
(after the upstream exchange) - so the very burn-then-fail pattern it was meant to
prevent still applied to refresh_token grants and to a misconfigured gateway.

- Hoist a single pre-exchange gate above the upstream call that covers BOTH grant
  types: it fails closed (invalid_request) on an unresolvable litellm identity and
  500s on an unset master_key BEFORE the single-use code or refresh token is
  exchanged/rotated, so a bad key or a misconfigured gateway never burns the
  upstream credential.
- Report expires_in from the envelope JWT's own second-truncated exp (rounding the
  elapsed portion up) instead of the raw expires_at - now delta, so the client is
  never told the bearer is valid past the ~1s point admission already expires it.

Regression tests assert the upstream exchange is never called on the no-identity
refresh grant and the master_key-unset path, and that the reported expires_in does
not overstate the JWT exp.
This commit is contained in:
Tin Chi Lo 2026-07-11 16:58:14 -07:00
parent 7a63e51625
commit e16ad044c3
2 changed files with 115 additions and 10 deletions

View file

@ -1,6 +1,7 @@
import asyncio
import html as _html
import json
import math
import secrets
import time
from datetime import datetime, timezone
@ -820,7 +821,10 @@ async def _mint_bridge_delegate_token_response(
status_code=502, detail="Upstream token is too large to seal into a gateway-bound credential"
)
expires_in = max(1, int((sealed.expires_at - now).total_seconds()))
# The JWT exp is int(expires_at.timestamp()) (second-truncated), and admission expires the envelope
# against that exp. Report expires_in from the same truncated exp, rounding the elapsed portion up,
# so the client is never told the bearer lives past the point admission already rejects it.
expires_in = max(1, int(sealed.expires_at.timestamp()) - math.ceil(now.timestamp()))
body = {"access_token": sealed.token.get_secret_value(), "token_type": "Bearer", "expires_in": expires_in}
return JSONResponse(body, headers=TOKEN_NO_CACHE_HEADERS)
@ -898,15 +902,19 @@ async def exchange_token_with_server(
if code_verifier:
token_data["code_verifier"] = code_verifier
# For a bridge oauth_delegate mint, resolve the litellm identity BEFORE exchanging the
# single-use upstream code. A missing or transiently-unresolvable identity then fails closed
# with invalid_request without consuming the code, so the client can retry the same code
# instead of being forced back through the full interactive authorize. The mint below
# re-resolves authoritatively; get_key_object is cache-first, so that second call is a cache
# hit and this adds no extra database round-trip.
if mcp_server.is_oauth_delegate and mcp_server.is_dcr_bridge:
if not await _extract_active_key_hash_from_request(request):
return _bridge_invalid_request_response()
# A bridge oauth_delegate mint must fail closed BEFORE the upstream exchange consumes or rotates the
# single-use code (or refresh token): confirm the gateway can mint at all (master_key set) and that
# the request carries a resolvable litellm identity. Applies to both grant types, so an invalid key
# or a misconfigured gateway never burns the upstream credential. The mint below re-checks
# authoritatively; get_key_object is cache-first, so the identity re-resolution is a cache hit and
# adds no extra database round-trip.
if mcp_server.is_oauth_delegate and mcp_server.is_dcr_bridge:
from litellm.proxy.proxy_server import master_key as _bridge_master_key # noqa: PLC0415
if not _bridge_master_key:
raise HTTPException(status_code=500, detail="Server misconfigured: master_key is not set")
if not await _extract_active_key_hash_from_request(request):
return _bridge_invalid_request_response()
async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
response = await async_client.post(

View file

@ -4503,6 +4503,103 @@ async def test_bridge_envelope_does_not_seal_upstream_refresh_token():
assert opened.grant.refresh_token is None
@pytest.mark.asyncio
async def test_bridge_refresh_grant_fails_closed_before_upstream_when_no_identity():
"""The pre-exchange identity gate covers the refresh_token grant, not just authorization_code: an
unresolvable litellm identity fails closed with invalid_request BEFORE the upstream refresh is
exchanged, so the client's refresh token is not rotated/consumed on a rejected request."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import exchange_token_with_server
from litellm.types.mcp import MCPAuth
server = _bridge_server(auth_type=MCPAuth.oauth_delegate)
fake_http_client = MagicMock()
fake_http_client.post = AsyncMock()
with (
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client",
return_value=fake_http_client,
),
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._extract_active_key_hash_from_request",
new=AsyncMock(return_value=None),
),
patch("litellm.proxy.proxy_server.master_key", _BRIDGE_MASTER_KEY),
):
response = await exchange_token_with_server(
request=_bridge_mock_request(),
mcp_server=server,
grant_type="refresh_token",
code=None,
redirect_uri=None,
client_id="dcr-client-123",
client_secret=None,
code_verifier=None,
refresh_token="client-refresh-token",
)
assert response.status_code == 400
assert json.loads(response.body)["error"] == "invalid_request"
fake_http_client.post.assert_not_called()
@pytest.mark.asyncio
async def test_bridge_mint_fails_closed_before_upstream_when_master_key_unset():
"""master_key is validated BEFORE the upstream exchange, so a misconfigured gateway 500s without
consuming the single-use code, avoiding the burn-then-fail the pre-exchange gate exists to prevent."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import exchange_token_with_server
from litellm.types.mcp import MCPAuth
server = _bridge_server(auth_type=MCPAuth.oauth_delegate)
fake_http_client = MagicMock()
fake_http_client.post = AsyncMock()
with (
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client",
return_value=fake_http_client,
),
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._extract_active_key_hash_from_request",
new=AsyncMock(return_value="hashed-litellm-key-77"),
),
patch("litellm.proxy.proxy_server.master_key", None),
):
with pytest.raises(HTTPException) as exc:
await exchange_token_with_server(
request=_bridge_mock_request(),
mcp_server=server,
grant_type="authorization_code",
code="auth-code",
redirect_uri="https://claude.ai/api/mcp/auth_callback",
client_id="dcr-client-123",
client_secret=None,
code_verifier="verifier",
)
assert exc.value.status_code == 500
fake_http_client.post.assert_not_called()
@pytest.mark.asyncio
async def test_bridge_reported_expires_in_does_not_overstate_jwt_exp():
"""The reported expires_in is derived from the envelope JWT's second-truncated exp (rounding the
elapsed portion up), so the client is never told the bearer lives past the point admission expires
it. Regression for the sub-second overstatement of the raw (expires_at - now) delta."""
import time
import jwt as _jwt
from litellm.types.mcp import MCPAuth
server = _bridge_server(auth_type=MCPAuth.oauth_delegate)
upstream = {"access_token": "UP", "token_type": "Bearer", "expires_in": 300}
before = int(time.time())
response = await _exchange_for_bridge_server(server, upstream, key_hash="hashed-litellm-key-77")
body = json.loads(response.body)
claims = _jwt.decode(body["access_token"].removeprefix("llm_env_"), options={"verify_signature": False})
# projecting the reported lifetime from a time no later than the mint must not exceed the JWT exp
assert before + body["expires_in"] <= claims["exp"]
def test_bridge_grant_coerces_numeric_expires_in():
"""expires_in from an IdP may be an int, a float (3600.0), or a numeric string ("3600"); coerce
it to a positive int so the envelope TTL honors the real lifetime instead of dropping a non-int