fix(mcp): require challenged upstream consent before completing OAuth

This commit is contained in:
Joshua Valluru 2026-09-28 12:57:00 -07:00
parent cb6050b521
commit 3a8ed9773b
6 changed files with 175 additions and 6 deletions

View file

@ -280,6 +280,7 @@ def _gateway_dcr_challenge(
route: str,
mcp_servers: list[str] | None,
invalid_token: bool,
oauth_scope: str | None = None,
) -> HTTPException:
"""The RFC 9728 challenge pointing the client at the protected-resource metadata
matching the scope it requested: the per-server document (same URL spelling the
@ -298,13 +299,14 @@ def _gateway_dcr_challenge(
else f"{get_request_base_url(request)}/.well-known/oauth-protected-resource{well_known_root_suffix()}/mcp"
)
error_attr: Final = 'error="invalid_token", ' if invalid_token else ""
scope_attr: Final = f', scope="{oauth_scope}"' if oauth_scope else ""
return HTTPException(
status_code=401,
detail={
"error": "authentication_required",
"message": "Authenticate with the gateway to use the MCP endpoint.",
},
headers={"WWW-Authenticate": f'Bearer {error_attr}resource_metadata="{resource_metadata_url}"'},
headers={"WWW-Authenticate": f'Bearer {error_attr}resource_metadata="{resource_metadata_url}"{scope_attr}'},
)

View file

@ -108,6 +108,7 @@ handle carried in the connect-page URL (the same handle-plus-cookie pattern as t
``mcp_oauth_state_`` upstream relay, for the same reasons: replica-safe with no
server-side session store, and the sealed value never appears in a URL)."""
UPSTREAM_AUTHORIZATION_SCOPE_PREFIX: Final = "litellm:mcp:connect:"
CONNECT_FLOW_TTL_SECONDS: Final = 600
GATEWAY_AUTH_CODE_TTL_SECONDS: Final = 120
MANUAL_DELIVERY_AUTH_CODE_TTL_SECONDS: Final = 300
@ -286,6 +287,13 @@ class GatewayDcrClient(BaseModel):
iat: int
class _UpstreamAuthorizationRequirement(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
server_id: str = Field(min_length=1)
user_id: str | None = None
exp: int
class _ConnectFlow(BaseModel):
"""One in-flight authorize: the SSO user it belongs to and the client parameters
needed to mint the code at the finish step. Sealed into the per-flow cookie. ``jti``
@ -301,6 +309,7 @@ class _ConnectFlow(BaseModel):
jti: str = Field(min_length=1)
exp: int
resource_server_id: str | None = None
required_upstream_server_id: str | None = None
audience: SessionAudience | None = None
@ -502,6 +511,17 @@ def resolve_scoped_resource_server(request: Request, resource: str | None) -> MC
return server
def upstream_authorization_scope(server_id: str, user_id: str | None) -> str:
return _seal(
UPSTREAM_AUTHORIZATION_SCOPE_PREFIX,
_UpstreamAuthorizationRequirement(
server_id=server_id,
user_id=user_id,
exp=int(datetime.now(timezone.utc).timestamp()) + CONNECT_FLOW_TTL_SECONDS,
),
)
def aggregate_authorize(
request: Request,
client_id: str,
@ -537,6 +557,30 @@ def aggregate_authorize(
if session_user_id is None:
return _login_redirect(base_url, request)
scoped_server: Final = resolve_scoped_resource_server(request, resource)
requested: Final = tuple(
value
for value in request.query_params.get("scope", "").split()
if value.startswith(UPSTREAM_AUTHORIZATION_SCOPE_PREFIX)
)
if len(requested) > 1:
return _oauth_error(400, "invalid_scope", "only one upstream authorization requirement is supported")
requirement: Final = (
_open_sealed(
requested[0],
UPSTREAM_AUTHORIZATION_SCOPE_PREFIX,
_UpstreamAuthorizationRequirement,
"mcp_upstream_authorization",
)
if requested
else None
)
if requested and (requirement is None or datetime.now(timezone.utc).timestamp() >= requirement.exp):
return _oauth_error(400, "invalid_scope", "invalid or expired upstream authorization requirement; reconnect")
if requirement is not None:
if requirement.user_id is not None and requirement.user_id != session_user_id:
return _oauth_error(403, "access_denied", "sign in as the user that requested this MCP connection")
if scoped_server is not None and scoped_server.server_id != requirement.server_id:
return _oauth_error(400, "invalid_scope", "upstream authorization requirement does not match the resource")
handle: Final = secrets.token_urlsafe(24)
flow: Final = _new_connect_flow(
session_user_id=session_user_id,
@ -546,6 +590,7 @@ def aggregate_authorize(
code_challenge=code_challenge or "",
resource_server_id=scoped_server.server_id if scoped_server is not None else None,
audience=None,
required_upstream_server_id=requirement.server_id if requirement is not None else None,
)
connect_url: Final = _append_query_params(f"{base_url}/ui/connect", (("connect_flow", handle),))
response: Final = RedirectResponse(connect_url, status_code=303)
@ -705,6 +750,7 @@ def _new_connect_flow(
code_challenge: str,
resource_server_id: str | None,
audience: SessionAudience | None,
required_upstream_server_id: str | None = None,
) -> _ConnectFlow:
now: Final = datetime.now(timezone.utc)
return _ConnectFlow(
@ -716,6 +762,7 @@ def _new_connect_flow(
jti=secrets.token_urlsafe(24),
exp=int(now.timestamp()) + CONNECT_FLOW_TTL_SECONDS,
resource_server_id=resource_server_id,
required_upstream_server_id=required_upstream_server_id,
audience=audience,
)
@ -773,14 +820,15 @@ def _open_flow_for(
async def _flow_target(
flow: _ConnectFlow, lookup_server_reachability: LookupServerReachability
) -> tuple[Literal["unscoped", "interactive", "m2m", "stale"], MCPServer | None]:
if flow.resource_server_id is None:
target_id: Final = flow.required_upstream_server_id or flow.resource_server_id
if target_id is None:
return "unscoped", None
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # import cycle
MCPServerManager,
global_mcp_server_manager,
)
server: Final = global_mcp_server_manager.get_mcp_server_by_id(flow.resource_server_id)
server: Final = global_mcp_server_manager.get_mcp_server_by_id(target_id)
if (
server is None
or not (server.is_gateway_managed_oauth2 or server.advertises_gateway_authorization_server)

View file

@ -48,6 +48,7 @@ from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
from litellm.proxy._experimental.mcp_server.exceptions import (
MCPUpstreamAuthError,
)
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import upstream_authorization_scope
from litellm.proxy._experimental.mcp_server.mcp_context import (
_mcp_active_toolset_id,
_mcp_gateway_initialize_instructions,
@ -1728,7 +1729,13 @@ if MCP_AVAILABLE:
if results and all(isinstance(result, HTTPException) and result.status_code == 401 for result in results):
if all(server.is_gateway_managed_oauth2 for server in eligible):
raise _gateway_dcr_challenge(
StarletteRequest(scope), get_route_relative_request_path(scope), None, invalid_token=False
StarletteRequest(scope),
get_route_relative_request_path(scope),
None,
invalid_token=False,
oauth_scope=upstream_authorization_scope(
eligible[0].server_id, user_api_key_auth.user_id if user_api_key_auth is not None else None
),
)
first: Final = results[0]
if isinstance(first, HTTPException):

View file

@ -2354,3 +2354,113 @@ async def test_token_exchange_relays_a_mint_refusal(failure, status, error):
response = await _exchange_native(client_id, _Minter(failure), _Exchanger())
assert response.status_code == status
assert json.loads(response.body)["error"] == error
@pytest.mark.asyncio
async def test_unified_challenge_requires_upstream_consent_without_narrowing_gateway_token(monkeypatch) -> None:
from unittest.mock import AsyncMock
from urllib.parse import urlencode
from fastapi import HTTPException
from litellm.proxy._experimental.mcp_server import operations, server
from litellm.proxy._types import UserAPIKeyAuth
github = _scoped_mcp_server(oauth2_flow="authorization_code")
manager = operations.global_mcp_server_manager
monkeypatch.setattr(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[github]))
monkeypatch.setattr(manager, "get_mcp_server_by_name", lambda *args, **kwargs: github)
monkeypatch.setattr(manager, "ensure_oauth_metadata_discovered", AsyncMock(return_value=github))
monkeypatch.setattr(manager, "has_user_oauth_token", AsyncMock(return_value=False))
with pytest.raises(HTTPException) as challenged:
await server._raise_preemptive_401_for_unauthenticated_servers(
scope=_request("/mcp").scope, mcp_servers=None, oauth2_headers=None,
mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(user_id="u1", api_key="test-key"),
client_ip=None,
)
headers = {k.lower(): v for k, v in (challenged.value.headers or {}).items()}
requested = re.search(r'scope="([^"]+)"', headers["www-authenticate"])
assert requested is not None, "The challenge must carry the upstream authorization requirement"
client_id = (await _register([REDIRECT_URI]))["client_id"]
response = aggregate_authorize(
request=_request("/authorize/mcp-session", query=urlencode({"scope": requested.group(1)})),
client_id=client_id, redirect_uri=REDIRECT_URI, state="client-state", code_challenge=CODE_CHALLENGE,
code_challenge_method="S256", response_type="code", session_user_id="u1",
resource="https://llm.example.com/mcp",
)
assert response.status_code == 303
described = await _describe_page(response, scoped_server=github, vendor=_VendorCredential("absent"))
assert json.loads(described.body) == {
"state": "interactive", "client_origin": "https://claude.ai",
"server_id": "github-id", "server_name": "github", "connected": False,
}
cache = DualCache()
premature = await _complete_page(response, scoped_server=github, vendor=_VendorCredential("absent"), cache=cache)
assert premature.status_code == 400
assert "location" not in premature.headers
vendor = _VendorCredential("present")
completed = await _complete_page(response, scoped_server=github, vendor=vendor, cache=cache)
assert completed.status_code == 303
assert vendor.calls == [("u1", "github-id")]
code = parse_qs(urlparse(completed.headers["location"]).query)["code"][0]
tokens = await _redeem(code, client_id, resource="https://llm.example.com/mcp")
assert tokens.status_code == 200
assert _opened_principal(json.loads(tokens.body)).resource_server_id is None
@pytest.mark.asyncio
@pytest.mark.parametrize("invalid", ("tampered", "expired", "duplicate", "other_user", "other_resource"))
async def test_upstream_authorization_requirement_rejects_invalid_binding(invalid: str, monkeypatch) -> None:
from urllib.parse import urlencode
from litellm.proxy._experimental.mcp_server import gateway_dcr_flow as flow
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
hint = flow.upstream_authorization_scope("github-id", "u2" if invalid == "other_user" else "u1")
if invalid == "tampered":
hint = flow.UPSTREAM_AUTHORIZATION_SCOPE_PREFIX + "invalid-ciphertext"
elif invalid == "expired":
hint = flow._seal(flow.UPSTREAM_AUTHORIZATION_SCOPE_PREFIX, flow._UpstreamAuthorizationRequirement(
server_id="github-id", user_id="u1", exp=int(datetime.now(timezone.utc).timestamp()) - 1,
))
elif invalid == "duplicate":
hint = f"{hint} {hint}"
monkeypatch.setattr(global_mcp_server_manager, "get_mcp_server_by_name", lambda *args, **kwargs: _scoped_mcp_server("other"))
client_id = (await _register([REDIRECT_URI]))["client_id"]
response = aggregate_authorize(
request=_request(query=urlencode({"scope": hint})),
client_id=client_id, redirect_uri=REDIRECT_URI, state="client-state", code_challenge=CODE_CHALLENGE,
code_challenge_method="S256", response_type="code", session_user_id="u1",
resource="https://llm.example.com/mcp/other" if invalid == "other_resource" else "https://llm.example.com/mcp",
)
assert response.status_code == (403 if invalid == "other_user" else 400)
assert json.loads(response.body)["error"] == ("access_denied" if invalid == "other_user" else "invalid_scope")
assert "location" not in response.headers
assert "set-cookie" not in response.headers
@pytest.mark.asyncio
@pytest.mark.parametrize("condition", ("deleted", "revoked", "vault_unavailable", "cancelled"))
async def test_required_upstream_completion_preserves_failure_and_cancellation_guards(condition: str) -> None:
from urllib.parse import urlencode
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import upstream_authorization_scope
client_id = (await _register([REDIRECT_URI]))["client_id"]
hint = upstream_authorization_scope("github-id", "u1")
response = aggregate_authorize(
request=_request(query=urlencode({"scope": hint})), client_id=client_id, redirect_uri=REDIRECT_URI,
state="client-state", code_challenge=CODE_CHALLENGE, code_challenge_method="S256",
response_type="code", session_user_id="u1", resource="https://llm.example.com/mcp",
)
assert response.status_code == 303
vendor = _VendorCredential("unavailable" if condition == "vault_unavailable" else "absent")
completed = await _complete_page(
response, scoped_server=None if condition == "deleted" else _scoped_mcp_server(),
reachable=_ServerReachability(condition != "revoked"), vendor=vendor,
decision="deny" if condition == "cancelled" else None,
)
if condition == "cancelled":
assert completed.status_code == 303
assert parse_qs(urlparse(completed.headers["location"]).query) == {"error": ["access_denied"], "state": ["client-state"]}
assert vendor.calls == []
else:
assert completed.status_code == (503 if condition == "vault_unavailable" else 400)
assert "location" not in completed.headers
assert vendor.calls == ([("u1", "github-id")] if condition == "vault_unavailable" else [])

View file

@ -867,6 +867,7 @@ async def test_initialize_challenges_missing_upstream_credentials_before_creatin
from litellm.proxy._experimental.mcp_server import server
from litellm.proxy._types import UserAPIKeyAuth
monkeypatch.setenv("LITELLM_SALT_KEY", "test-mcp-oauth-signing")
github: Final = MCPServer(
server_id="github-id", name="github", alias="github", server_name="github",
url="https://github.example/mcp", transport=MCPTransport.http,
@ -916,8 +917,8 @@ async def test_initialize_challenges_missing_upstream_credentials_before_creatin
else:
assert response.headers["www-authenticate"].startswith("Bearer ")
if path == "/mcp":
assert response.headers["www-authenticate"] == (
'Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp"'
assert response.headers["www-authenticate"].startswith(
'Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp", scope="litellm:mcp:connect:'
)
assert "mcp-session-id" not in response.headers
finally:

View file

@ -10714,6 +10714,7 @@ async def test_unified_preflight_challenges_only_when_all_authorized_servers_nee
) -> None:
from litellm.proxy._experimental.mcp_server import server as server_module
monkeypatch.setenv("LITELLM_SALT_KEY", "test-mcp-oauth-signing")
servers: Final = tuple(
_make_oauth2_server(f"server-{index}").model_copy(update={"server_id": f"server-{index}"})
for index in range(len(token_states))