fix(mcp): preserve safe OAuth retries and recovery challenges

This commit is contained in:
Joshua Valluru 2026-09-21 17:04:33 -07:00
parent 8136a67410
commit a7024590f8
6 changed files with 159 additions and 46 deletions

View file

@ -462,7 +462,8 @@ class MCPRequestHandler:
request_route: Final = get_request_route(request)
# Only OAuth metadata routes registered under /.well-known/ are public.
if request_route.startswith("/.well-known/"):
is_public_metadata: Final = request_route.startswith("/.well-known/")
if is_public_metadata:
validated_user_api_key_auth = UserAPIKeyAuth()
elif has_explicit_litellm_key:
# An explicit x-litellm-api-key is always a LiteLLM credential, even
@ -552,7 +553,7 @@ class MCPRequestHandler:
scope.pop(CONNECTION_SCOPE_KEY, None)
connection_header: Final = headers.get("authorization")
if is_connection_credential(connection_header):
if is_connection_credential(connection_header) and not is_public_metadata:
scope[CONNECTION_SCOPE_KEY] = await MCPRequestHandler._admit_connection_credential(
request=request,
request_route=request_route,
@ -627,10 +628,14 @@ class MCPRequestHandler:
resource=f"{get_request_base_url(request)}/mcp",
)
connection: Final = open_connection_credential(connection_header)
if connection is None:
if connection is None or connection.binding != expected_binding:
raise HTTPException(
status_code=401,
detail="Invalid or expired MCP connection credential",
detail=(
"Invalid or expired MCP connection credential"
if connection is None
else "Connection credential belongs to a different key or resource"
),
headers=MappingProxyType(
{
"www-authenticate": connection_challenge(request, expected_binding),
@ -638,8 +643,6 @@ class MCPRequestHandler:
}
),
)
if connection.binding != expected_binding:
raise HTTPException(status_code=401, detail="Connection credential belongs to a different key or resource")
return connection
@staticmethod
@ -1216,7 +1219,13 @@ class MCPRequestHandler:
raise HTTPException(status_code=401, detail="Invalid or expired credential")
@staticmethod
async def _enforce_admitted_live_policy(admitted: UserAPIKeyAuth, request: Request, route: str) -> None:
async def _enforce_admitted_live_policy(
admitted: UserAPIKeyAuth,
request: Request,
route: str,
*,
request_data: dict[str, object] | None = None,
) -> None:
"""Run the standard pipeline's authorization checks over the admitted identity.
Mirrors the ``user_api_key_auth`` wrapper between the builder and its return: clear the
@ -1246,7 +1255,7 @@ class MCPRequestHandler:
await _run_centralized_common_checks(
user_api_key_auth_obj=admitted,
request=request,
request_data=await _read_request_body(request=request),
request_data=await _read_request_body(request=request) if request_data is None else request_data,
route=route,
)
except (HTTPException, ProxyException):

View file

@ -1041,6 +1041,7 @@ async def exchange_token_with_server(
scope: str | None = None,
client_token_endpoint_auth_method: MCPTokenEndpointAuthMethod | None = None,
connection_binding: ConnectionBinding | None = None,
connection_claim: tuple[str, int] | None = None,
):
_raise_if_not_oauth2(mcp_server)
if grant_type not in ("authorization_code", "refresh_token"):
@ -1197,6 +1198,10 @@ async def exchange_token_with_server(
)
async_client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
if connection_claim is not None:
claimed: Final = await claim_connection_once(*connection_claim)
if claimed is not None:
return claimed
try:
response: Final = await async_client.post(
token_url,
@ -3092,11 +3097,8 @@ async def complete_connection(request: Request, flow: str = Form(...), decision:
if opened is None or opened.jti != flow or decision not in ("approve", "deny"):
return _oauth_error(400, "invalid_request", "Invalid or expired consent")
server: Final = await validate_connection_binding(request, opened.binding)
claim: Final = await claim_connection_once(f"consent:{opened.jti}", opened.exp)
if claim is not None:
return claim
if decision == "deny":
denied: Final = RedirectResponse(
response: Final = (
RedirectResponse(
append_connection_query(
opened.redirect_uri,
(
@ -3106,21 +3108,25 @@ async def complete_connection(request: Request, flow: str = Form(...), decision:
),
status_code=302,
)
cookie_path, _ = _cookie_path_and_secure(request)
denied.delete_cookie(cookie_name, path=cookie_path)
return denied
response: Final = await authorize_with_server(
request,
server,
opened.client_id,
opened.redirect_uri,
opened.state,
opened.code_challenge,
"S256",
"code",
opened.scope,
connection=opened,
if decision == "deny"
else await authorize_with_server(
request,
server,
opened.client_id,
opened.redirect_uri,
opened.state,
opened.code_challenge,
"S256",
"code",
opened.scope,
connection=opened,
)
)
if response.status_code >= 400:
return response
claim: Final = await claim_connection_once(f"consent:{opened.jti}", opened.exp)
if claim is not None:
return claim
cookie_path, _ = _cookie_path_and_secure(request)
response.delete_cookie(cookie_name, path=cookie_path)
return response
@ -3153,9 +3159,6 @@ async def exchange_connection_token(
or not _pkce_verifier_matches(code_verifier, opened.authorization.code_challenge)
):
return _oauth_error(400, "invalid_grant", "Invalid authorization code or PKCE verifier")
claim: Final = await claim_connection_once(f"code:{opened.jti}", opened.exp)
if claim is not None:
return claim
return await exchange_token_with_server(
request,
server,
@ -3167,6 +3170,7 @@ async def exchange_connection_token(
code_verifier,
scope=opened.authorization.scope,
connection_binding=binding,
connection_claim=(f"code:{opened.jti}", opened.exp),
)
if grant_type == "refresh_token":
refreshed: Final = open_connection_credential(refresh_token or "", refresh=True)
@ -3174,9 +3178,6 @@ async def exchange_connection_token(
return _oauth_error(400, "invalid_grant", "Invalid refresh credential")
if scope and not frozenset(scope.split()).issubset((refreshed.scope or "").split()):
return _oauth_error(400, "invalid_scope", "Refresh cannot expand the granted scopes")
claimed: Final = await claim_connection_once(f"refresh:{refreshed.jti}", refreshed.exp)
if claimed is not None:
return claimed
return await exchange_token_with_server(
request,
server,
@ -3189,5 +3190,6 @@ async def exchange_connection_token(
refresh_token=refreshed.token.get_secret_value(),
scope=scope or refreshed.scope,
connection_binding=binding,
connection_claim=(f"refresh:{refreshed.jti}", refreshed.exp),
)
return _oauth_error(400, "unsupported_grant_type", "Unsupported connection grant type")

View file

@ -1632,7 +1632,9 @@ async def validate_connection_binding(request: Request, binding: ConnectionBindi
if binding.resource != f"{get_request_base_url(request)}/mcp":
raise HTTPException(status_code=400, detail="Invalid connection resource")
key: Final = await MCPRequestHandler._reload_admitted_key(binding.key_hash) # pyright: ignore[reportPrivateUsage] # reuse key revocation and SCIM checks
await MCPRequestHandler._enforce_admitted_live_policy(key.model_copy(), request, "/mcp") # pyright: ignore[reportPrivateUsage] # enforce the same MCP route and budget policy
await MCPRequestHandler._enforce_admitted_live_policy( # pyright: ignore[reportPrivateUsage] # enforce the same MCP route and budget policy
key.model_copy(), request, "/mcp", request_data={}
)
allowed: Final = await MCPRequestHandler.get_allowed_mcp_servers(key)
server: Final = global_mcp_server_manager.get_mcp_server_by_id(
binding.server_id, client_ip=IPAddressUtils.get_mcp_client_ip(request)

View file

@ -953,7 +953,9 @@ async def test_get_tools_from_mcp_servers():
client_ip=None,
user_api_key_auth=None,
oauth2_headers=None,
connection_credential=None,
):
assert connection_credential is None
if server.server_id == "server1_id":
return [mock_tool_1]
return [mock_tool_2]

View file

@ -1817,7 +1817,8 @@ class TestMCPPublicRouteGuard:
await MCPRequestHandler.process_mcp_request(scope)
assert exc_info.value.status_code == 401
async def test_legitimate_well_known_path_still_bypasses_auth(self):
@pytest.mark.parametrize("bearer", [None, "llm_caccess_stale", "llm_crefresh_stale"])
async def test_legitimate_well_known_path_still_bypasses_auth(self, bearer):
"""
Real OAuth discovery routes registered under /.well-known/ must remain
public so unauthenticated clients can fetch them per RFC 8414/9728.
@ -1826,16 +1827,21 @@ class TestMCPPublicRouteGuard:
"type": "http",
"method": "GET",
"path": "/.well-known/oauth-protected-resource",
"headers": [],
"headers": [(b"authorization", f"Bearer {bearer}".encode())] if bearer else [],
}
# No mock needed — public path should not call user_api_key_auth at all
with patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
) as mock_auth:
auth_result, *_rest = await MCPRequestHandler.process_mcp_request(scope)
auth_result, mcp_header, _, server_headers, oauth_headers, raw_headers = (
await MCPRequestHandler.process_mcp_request(scope)
)
mock_auth.assert_not_called()
assert isinstance(auth_result, UserAPIKeyAuth)
assert "litellm.mcp.connection_grant" not in scope
assert not mcp_header and not server_headers and not oauth_headers
assert "authorization" not in raw_headers
@pytest.mark.asyncio
@ -9449,7 +9455,7 @@ class TestScopedSessionAdmission:
(None, "connection-target", 401),
],
)
@pytest.mark.parametrize("grant_state", ["valid", "expired", "refresh", "oversized"])
@pytest.mark.parametrize("grant_state", ["valid", "expired", "refresh", "oversized", "other-resource"])
@pytest.mark.parametrize("server_mode", ["managed", "delegated"])
async def test_connection_credential_requires_exact_key_and_server(
monkeypatch, key, selector, expected, grant_state, server_mode
@ -9485,7 +9491,9 @@ async def test_connection_credential_requires_exact_key_and_server(
)
monkeypatch.setattr(admission, "user_api_key_auth", AsyncMock(return_value=auth))
binding = ConnectionBinding(
key_hash=hash_token("sk-original"), server_id=server.server_id, resource="https://gateway.example/mcp"
key_hash=hash_token("sk-original"),
server_id=server.server_id,
resource="https://other.example/mcp" if grant_state == "other-resource" else "https://gateway.example/mcp",
)
issued = json.loads(
flow.mint_connection_tokens(
@ -9528,11 +9536,18 @@ async def test_connection_credential_requires_exact_key_and_server(
assert exc.value.status_code == expected_status
if (
server_mode == "managed"
and key == "sk-original"
and key is not None
and selector == "connection-target"
and grant_state != "valid"
):
assert "resource_metadata=" in exc.value.headers["www-authenticate"]
from urllib.parse import parse_qs, urlparse
challenge = exc.value.headers["www-authenticate"]
metadata_url = challenge.split('resource_metadata="', 1)[1].split('"', 1)[0]
bootstrap = parse_qs(urlparse(metadata_url).query)["connection"][0]
assert flow.open_connection_bootstrap(bootstrap) == ConnectionBinding(
key_hash=hash_token(key), server_id=server.server_id, resource="https://gateway.example/mcp"
)
assert exc.value.headers["Cache-Control"] == "no-store"
assert flow.CONNECTION_SCOPE_KEY not in scope
return
_, _, _, server_headers, oauth_headers, raw_headers = await MCPRequestHandler.process_mcp_request(scope)

View file

@ -12574,6 +12574,79 @@ def test_keyed_connection_headerless_exchange_and_rotating_refresh(keyed_oauth_c
harness.vault.assert_not_called()
def test_keyed_connection_consent_can_retry_failed_preparation(keyed_oauth_client, monkeypatch):
from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints
harness = keyed_oauth_client
_, _, handle = _start_keyed_oauth(harness)
prepare = AsyncMock(side_effect=[HTTPException(status_code=503, detail="discovery unavailable"), harness.server])
monkeypatch.setattr(endpoints, "_server_with_oauth_endpoints", prepare)
payload = {"flow": handle, "decision": "approve"}
failed = harness.client.post("/authorize/connection/complete", data=payload)
assert failed.status_code == 503
assert "location" not in failed.headers
assert not failed.headers.get("set-cookie")
retried = harness.client.post("/authorize/connection/complete", data=payload)
assert retried.status_code == 303
assert retried.headers["location"].startswith("https://provider.example/authorize?")
harness.upstream.post.assert_not_awaited()
@pytest.mark.parametrize("grant_type", ["authorization_code", "refresh_token"])
def test_keyed_connection_exchange_can_retry_failed_preparation(keyed_oauth_client, monkeypatch, grant_type):
from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints
harness = keyed_oauth_client
code_payload = _complete_keyed_oauth(harness)
tokens = harness.client.post("/token", data=code_payload).json() if grant_type == "refresh_token" else None
payload = (
{"grant_type": grant_type, "client_id": code_payload["client_id"], "refresh_token": tokens["refresh_token"]}
if tokens is not None else code_payload
)
harness.upstream.post.reset_mock()
prepare = AsyncMock(side_effect=[HTTPException(status_code=503, detail="discovery unavailable"), harness.server])
monkeypatch.setattr(endpoints, "_server_with_oauth_endpoints", prepare)
failed = harness.client.post("/token", data=payload)
assert failed.status_code == 503
harness.upstream.post.assert_not_awaited()
retried = harness.client.post("/token", data=payload)
assert retried.status_code == 200, retried.text
assert retried.json()["access_token"].startswith("llm_caccess_")
harness.upstream.post.assert_awaited_once()
@pytest.mark.parametrize("grant_type", ["authorization_code", "refresh_token"])
@pytest.mark.parametrize("failure", ["response", "timeout"])
def test_keyed_connection_dispatched_provider_failure_cannot_redeem_twice(keyed_oauth_client, grant_type, failure):
import httpx
harness = keyed_oauth_client
code_payload = _complete_keyed_oauth(harness)
tokens = harness.client.post("/token", data=code_payload).json() if grant_type == "refresh_token" else None
payload = (
{"grant_type": grant_type, "client_id": code_payload["client_id"], "refresh_token": tokens["refresh_token"]}
if tokens is not None else code_payload
)
harness.upstream.post.reset_mock()
harness.upstream.post.return_value = httpx.Response(
503, json={"error": "temporarily_unavailable"}, request=httpx.Request("POST", "https://provider.example/token")
)
if failure == "timeout":
harness.upstream.post.side_effect = httpx.ReadTimeout("provider response unavailable")
with pytest.raises(httpx.ReadTimeout):
harness.client.post("/token", data=payload)
else:
failed = harness.client.post("/token", data=payload)
assert failed.status_code >= 500
replayed = harness.client.post("/token", data=payload)
assert replayed.status_code == 400
assert replayed.json()["error"] == "invalid_grant"
harness.upstream.post.assert_awaited_once()
@pytest.mark.parametrize("change", ["verifier", "redirect", "resource", "tamper"])
def test_keyed_connection_rejects_bad_exchange_before_upstream(keyed_oauth_client, change):
harness = keyed_oauth_client
@ -12672,7 +12745,7 @@ def test_keyed_connection_client_cannot_enter_session_login_flow(keyed_oauth_cli
@pytest.mark.parametrize("stage", ["authorize", "exchange", "refresh"])
@pytest.mark.parametrize("policy", ["blocked", "expired", "denied_server", "denied_route", "allowed"])
@pytest.mark.parametrize("policy", ["blocked", "expired", "denied_server", "denied_route", "over_budget", "allowed"])
def test_keyed_connection_reloads_live_key_before_provider(keyed_oauth_client, monkeypatch, stage, policy):
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server import gateway_dcr_flow as flow
@ -12702,7 +12775,12 @@ def test_keyed_connection_reloads_live_key_before_provider(keyed_oauth_client, m
lookup = AsyncMock(return_value=key)
monkeypatch.setattr(auth_checks, "get_key_object", lookup)
monkeypatch.setattr(proxy_server, "prisma_client", MagicMock())
monkeypatch.setattr(admission, "_run_centralized_common_checks", AsyncMock())
import litellm
common_checks = AsyncMock(
side_effect=litellm.BudgetExceededError(current_cost=2, max_budget=1) if policy == "over_budget" else None
)
monkeypatch.setattr(admission, "_run_centralized_common_checks", common_checks)
monkeypatch.setitem(global_mcp_server_manager.registry, harness.server.server_id, harness.server)
if stage == "authorize":
challenge = urlsafe_b64encode(hashlib.sha256(payload["code_verifier"].encode()).digest()).rstrip(b"=").decode()
@ -12718,25 +12796,30 @@ def test_keyed_connection_reloads_live_key_before_provider(keyed_oauth_client, m
},
)
elif stage == "exchange":
response = harness.client.post("/token", data=payload)
response = harness.client.post("/token", data={**payload, "model": "untrusted\nforged log entry"})
else:
assert token is not None and token.status_code == 200
response = harness.client.post(
"/token",
data={
"model": "untrusted\nforged log entry",
"grant_type": "refresh_token",
"client_id": payload["client_id"],
"refresh_token": token.json()["refresh_token"],
"resource": harness.binding.resource,
},
)
assert response.status_code == (200 if policy == "allowed" else 401 if policy in ("blocked", "expired") else 403), (
assert response.status_code == (200 if policy == "allowed" else 401 if policy in ("blocked", "expired") else 422 if policy == "over_budget" else 403), (
response.text
)
lookup.assert_awaited_once()
assert lookup.call_args.kwargs["hashed_token"] == harness.binding.key_hash
assert harness.upstream.post.call_count == int(policy == "allowed" and stage != "authorize")
harness.vault.assert_not_called()
if policy not in ("blocked", "expired", "denied_route"):
common_checks.assert_awaited_once()
assert common_checks.call_args.kwargs["request_data"] == {}
assert common_checks.call_args.kwargs["route"] == "/mcp"
@pytest.mark.parametrize(