From 3a8ed9773b94f4acbbd878e364955bbe473a5c3a Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Mon, 28 Sep 2026 12:57:00 -0700 Subject: [PATCH] fix(mcp): require challenged upstream consent before completing OAuth --- .../mcp_server/auth/user_api_key_auth_mcp.py | 4 +- .../mcp_server/gateway_dcr_flow.py | 52 ++++++++- .../proxy/_experimental/mcp_server/server.py | 9 +- .../mcp_server/test_gateway_dcr_flow.py | 110 ++++++++++++++++++ .../test_mcp_oauth_passthrough_tools.py | 5 +- .../mcp_server/test_mcp_server.py | 1 + 6 files changed, 175 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index a93ffaeac9f..237e38fa283 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -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}'}, ) diff --git a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py index e66504af47a..f4ec41cb6e0 100644 --- a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py +++ b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py @@ -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) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index ea4799d9712..f8842fa02f6 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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): 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 2943ff4b74a..40ec75329a6 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 @@ -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 []) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py index a7c7e1853a0..785bcc8ab70 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py @@ -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: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 0ba8b3f2a40..5777caf63de 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -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))