mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
7a63e51625
commit
e16ad044c3
2 changed files with 115 additions and 10 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue