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 237e38fa283..a93ffaeac9f 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,7 +280,6 @@ 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 @@ -299,14 +298,13 @@ 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}"{scope_attr}'}, + headers={"WWW-Authenticate": f'Bearer {error_attr}resource_metadata="{resource_metadata_url}"'}, ) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index d42c1c6b879..6b1f55dfeac 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -5,7 +5,7 @@ import secrets import time from collections.abc import Callable, Mapping from datetime import datetime, timezone -from typing import TYPE_CHECKING, Any, Final, Literal, Optional +from typing import TYPE_CHECKING, Annotated, Any, Final, Literal, Optional from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse import httpx @@ -2085,6 +2085,7 @@ async def authorize_complete( delivery: str | None = Form(None), team_id: str | None = Form(None), decision: str | None = Form(None), + selected_servers: Annotated[list[str] | None, Form(max_length=100)] = None, ) -> Response: """Finish an aggregate connect flow: mint the gateway authorization code for the signed-in user and hand it back to the DCR client, by 303 redirect (default) or, for @@ -2102,6 +2103,7 @@ async def authorize_complete( delivery=delivery, team_id=team_id, decision=decision, + selected_servers=tuple(selected_servers or ()), lookup_vendor_credential=_vendor_credential_state, lookup_server_reachability=_user_can_reach_mcp_server, ) diff --git a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py index f4ec41cb6e0..405815e6162 100644 --- a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py +++ b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py @@ -108,7 +108,6 @@ 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 @@ -287,13 +286,6 @@ 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`` @@ -309,7 +301,6 @@ 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 @@ -511,17 +502,6 @@ 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, @@ -557,30 +537,6 @@ 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, @@ -590,7 +546,6 @@ 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) @@ -750,7 +705,6 @@ 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( @@ -762,7 +716,6 @@ 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, ) @@ -820,15 +773,14 @@ def _open_flow_for( async def _flow_target( flow: _ConnectFlow, lookup_server_reachability: LookupServerReachability ) -> tuple[Literal["unscoped", "interactive", "m2m", "stale"], MCPServer | None]: - target_id: Final = flow.required_upstream_server_id or flow.resource_server_id - if target_id is None: + if flow.resource_server_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(target_id) + server: Final = global_mcp_server_manager.get_mcp_server_by_id(flow.resource_server_id) if ( server is None or not (server.is_gateway_managed_oauth2 or server.advertises_gateway_authorization_server) @@ -899,6 +851,35 @@ async def describe_connect_flow( ) +async def _selected_connections_refusal( + flow: _ConnectFlow, + selected_servers: tuple[str, ...], + lookup_vendor_credential: LookupVendorCredential, + lookup_server_reachability: LookupServerReachability, +) -> Response | None: + if not selected_servers: + return _oauth_error(400, "invalid_request", "select and connect an MCP server before finishing") + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # proxy import cycle + global_mcp_server_manager, + ) + + for server in (global_mcp_server_manager.get_mcp_server_by_name(name) for name in dict.fromkeys(selected_servers)): + if server is None or not await lookup_server_reachability(flow.user_id, server.server_id): + return _oauth_error(400, "invalid_request", "a selected MCP server is no longer available") + if ( + server.is_gateway_managed_oauth2 + and global_mcp_server_manager.effective_oauth2_flow(server) != "client_credentials" + ): + match await lookup_vendor_credential(flow.user_id, server.server_id): + case "unavailable": + return _oauth_error(503, "temporarily_unavailable", _DB_UNAVAILABLE_DESCRIPTION) + case "absent": + return _oauth_error(400, "invalid_request", "authorize the selected MCP servers before finishing") + case "present": + pass + return None + + async def complete_connect_flow( request: Request, flow_handle: str, @@ -909,6 +890,7 @@ async def complete_connect_flow( decision: str | None = None, lookup_vendor_credential: LookupVendorCredential = _unavailable_vendor_credential, lookup_server_reachability: LookupServerReachability = _unreachable_server, + selected_servers: tuple[str, ...] = (), ) -> Response: """Mint the code only after a deliberate POST by the sealed user. @@ -924,6 +906,12 @@ async def complete_connect_flow( opened: Final = _open_flow_for(request, flow_handle, session_user_id, now) if isinstance(opened, Response): return opened + if decision != "deny" and opened.resource_server_id is None and opened.audience is None: + refusal: Final = await _selected_connections_refusal( + opened, selected_servers, lookup_vendor_credential, lookup_server_reachability + ) + if refusal is not None: + return refusal if decision != "deny": described: Final = await _describe_opened_flow(opened, lookup_vendor_credential, lookup_server_reachability) if isinstance(described, Response): diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index f8842fa02f6..ea4799d9712 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -48,7 +48,6 @@ 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, @@ -1729,13 +1728,7 @@ 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, - 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 - ), + StarletteRequest(scope), get_route_relative_request_path(scope), None, invalid_token=False ) first: Final = results[0] if isinstance(first, HTTPException): diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index db57aa4f046..72afbcc34ae 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -31398,6 +31398,21 @@ "title": "Flow", "type": "string" }, + "selected_servers": { + "anyOf": [ + { + "items": { + "type": "string" + }, + "maxItems": 100, + "type": "array" + }, + { + "type": "null" + } + ], + "title": "Selected Servers" + }, "team_id": { "anyOf": [ { 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 40ec75329a6..7093e8f58c2 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 @@ -74,6 +74,13 @@ CODE_CHALLENGE = urlsafe_b64encode(hashlib.sha256(CODE_VERIFIER.encode("ascii")) @pytest.fixture(autouse=True) def _salt_key(monkeypatch): monkeypatch.setenv("LITELLM_SALT_KEY", MASTER_KEY) + from unittest.mock import patch + + with patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.get_mcp_server_by_name", + return_value=_scoped_mcp_server("public", auth_type="none"), + ): + yield def _request(path="/authorize", query="", cookies=None, method="GET"): @@ -339,6 +346,8 @@ async def test_full_walk_register_authorize_complete_token_and_replay(redirect_u handle, cookies = _flow_cookie_from(authorize_response) denied = await complete_connect_flow( + selected_servers=("public",), + lookup_server_reachability=_ServerReachability(), request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, session_user_id="attacker", @@ -347,6 +356,8 @@ async def test_full_walk_register_authorize_complete_token_and_replay(redirect_u assert denied.status_code == 403 anonymous = await complete_connect_flow( + selected_servers=("public",), + lookup_server_reachability=_ServerReachability(), request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, session_user_id=None, @@ -355,6 +366,8 @@ async def test_full_walk_register_authorize_complete_token_and_replay(redirect_u assert anonymous.status_code == 401 completed = await complete_connect_flow( + selected_servers=("public",), + lookup_server_reachability=_ServerReachability(), request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, session_user_id="u1", @@ -429,6 +442,8 @@ async def test_full_walk_register_authorize_complete_token_and_replay(redirect_u @pytest.mark.asyncio async def test_complete_rejects_missing_tampered_and_expired_flows(): missing = await complete_connect_flow( + selected_servers=("public",), + lookup_server_reachability=_ServerReachability(), request=_request("/authorize/complete", method="POST"), flow_handle="nope", session_user_id="u1", @@ -437,6 +452,8 @@ async def test_complete_rejects_missing_tampered_and_expired_flows(): assert missing.status_code == 400 tampered = await complete_connect_flow( + selected_servers=("public",), + lookup_server_reachability=_ServerReachability(), request=_request("/authorize/complete", cookies={f"{CONNECT_FLOW_COOKIE_PREFIX}h1": "garbage"}, method="POST"), flow_handle="h1", session_user_id="u1", @@ -518,6 +535,8 @@ async def test_token_gates_on_live_user_revalidation(failure, expected_status, e authorize_response = _authorize(client_id, session_user_id="deactivated-user") handle, cookies = _flow_cookie_from(authorize_response) completed = await complete_connect_flow( + selected_servers=("public",), + lookup_server_reachability=_ServerReachability(), request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, session_user_id="deactivated-user", @@ -572,6 +591,8 @@ async def test_flow_is_single_use_shared_cache_rejects_second_complete(): handle, cookies = _flow_cookie_from(_authorize(client_id, session_user_id="u1")) first = await complete_connect_flow( + selected_servers=("public",), + lookup_server_reachability=_ServerReachability(), request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, session_user_id="u1", @@ -579,6 +600,8 @@ async def test_flow_is_single_use_shared_cache_rejects_second_complete(): ) assert first.status_code == 303 second = await complete_connect_flow( + selected_servers=("public",), + lookup_server_reachability=_ServerReachability(), request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, session_user_id="u1", @@ -720,6 +743,8 @@ async def _complete(redirect_uri: str, delivery, cookies=None, handle=None, sess if cookies is None: handle, cookies = _flow_cookie_from(_authorize(client_id, session_user_id="u1", redirect_uri=redirect_uri)) response = await complete_connect_flow( + selected_servers=("public",), + lookup_server_reachability=_ServerReachability(), request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, session_user_id=session_user_id, @@ -827,6 +852,8 @@ async def test_unknown_delivery_value_is_rejected_before_the_flow_is_consumed(): handle, cookies = _flow_cookie_from(_authorize(client_id, session_user_id="u1", redirect_uri=LOOPBACK_REDIRECT_URI)) rejected = await complete_connect_flow( + selected_servers=("public",), + lookup_server_reachability=_ServerReachability(), request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, session_user_id="u1", @@ -837,6 +864,8 @@ async def test_unknown_delivery_value_is_rejected_before_the_flow_is_consumed(): assert json.loads(rejected.body)["error"] == "invalid_request" retried = await complete_connect_flow( + selected_servers=("public",), + lookup_server_reachability=_ServerReachability(), request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, session_user_id="u1", @@ -998,6 +1027,7 @@ async def _complete_page(response, scoped_server=None, vendor=None, reachable=No handle, cookies = _flow_cookie_from(response) with patch(_MANAGER_PATCH) as manager: manager.get_mcp_server_by_id.return_value = scoped_server + manager.get_mcp_server_by_name.return_value = scoped_server or _scoped_mcp_server("public", auth_type="none") return await complete_connect_flow( request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, @@ -1005,7 +1035,7 @@ async def _complete_page(response, scoped_server=None, vendor=None, reachable=No cache=cache or DualCache(), lookup_vendor_credential=vendor or _VendorCredential(), lookup_server_reachability=reachable or _ServerReachability(), - **overrides, + **{"selected_servers": ("public",), **overrides}, ) @@ -1814,6 +1844,8 @@ async def test_mcp_wire_formats_carry_no_native_client_fields(): assert "audience" not in flow_wire assert "team_id" not in flow_wire completed = await complete_connect_flow( + selected_servers=("public",), + lookup_server_reachability=_ServerReachability(), request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, session_user_id="u1", @@ -2357,110 +2389,98 @@ async def test_token_exchange_relays_a_mint_refusal(failure, status, 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" +async def test_unified_completion_requires_an_upstream_selection(): 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, - } + response = _authorize(client_id, session_user_id="u1") + completed = await _complete_page(response, selected_servers=()) + assert completed.status_code == 400 + assert "location" not in completed.headers + assert "select" in json.loads(completed.body)["error_description"] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "credential,reachable,status", + [("absent", True, 400), ("unavailable", True, 503), ("present", False, 400), ("present", True, 303)], +) +async def test_unified_completion_checks_selected_upstream_and_permissions(credential, reachable, status): + client_id = (await _register([REDIRECT_URI]))["client_id"] + response = _authorize(client_id, session_user_id="u1") 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") + server = _scoped_mcp_server(oauth2_flow="authorization_code") + vendor = _VendorCredential(credential) 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, + response, + scoped_server=server, + vendor=vendor, + reachable=_ServerReachability(reachable), + cache=cache, + selected_servers=("github",), ) - if condition == "cancelled": - assert completed.status_code == 303 - assert parse_qs(urlparse(completed.headers["location"]).query) == {"error": ["access_denied"], "state": ["client-state"]} + assert completed.status_code == status + if not reachable: assert vendor.calls == [] - else: - assert completed.status_code == (503 if condition == "vault_unavailable" else 400) + if status != 303: assert "location" not in completed.headers - assert vendor.calls == ([("u1", "github-id")] if condition == "vault_unavailable" else []) + retried = await _complete_page(response, scoped_server=server, cache=cache, selected_servers=("github",)) + assert retried.status_code == 303 + else: + code = parse_qs(urlparse(completed.headers["location"]).query)["code"][0] + token = await _redeem(code, client_id) + assert _opened_principal(json.loads(token.body)).resource_server_id is None + + +@pytest.mark.asyncio +async def test_unified_cancel_does_not_require_selected_servers(): + client_id = (await _register([REDIRECT_URI]))["client_id"] + response = _authorize(client_id, session_user_id="u1") + completed = await _complete_page(response, selected_servers=(), decision="deny") + assert completed.status_code == 303 + assert parse_qs(urlparse(completed.headers["location"]).query)["error"] == ["access_denied"] + + +@pytest.mark.asyncio +async def test_unified_completion_checks_every_selected_server_and_preserves_cancellation(): + import asyncio + from unittest.mock import patch + + client_id = (await _register([REDIRECT_URI]))["client_id"] + response = _authorize(client_id, session_user_id="u1") + handle, cookies = _flow_cookie_from(response) + servers = {name: _scoped_mcp_server(name, oauth2_flow="authorization_code") for name in ("github", "slack")} + cache = DualCache() + + async def credential(user_id, server_id): + if server_id == "slack-id": + return "absent" + return "present" + + async def cancelled(user_id, server_id): + raise asyncio.CancelledError() + + with patch(_MANAGER_PATCH) as manager: + manager.get_mcp_server_by_name.side_effect = servers.get + manager.get_mcp_server_by_id.side_effect = lambda server_id: next( + (server for server in servers.values() if server.server_id == server_id), None + ) + arguments = { + "request": _request("/authorize/complete", cookies=cookies, method="POST"), + "flow_handle": handle, + "session_user_id": "u1", + "cache": cache, + "lookup_server_reachability": _ServerReachability(), + } + missing = await complete_connect_flow( + **arguments, selected_servers=("missing",), lookup_vendor_credential=credential + ) + assert missing.status_code == 400 + unfinished = await complete_connect_flow( + **arguments, selected_servers=("github", "slack"), lookup_vendor_credential=credential + ) + assert unfinished.status_code == 400 + with pytest.raises(asyncio.CancelledError): + await complete_connect_flow(**arguments, selected_servers=("github",), lookup_vendor_credential=cancelled) + completed = await complete_connect_flow( + **arguments, selected_servers=("github", "slack"), lookup_vendor_credential=_VendorCredential() + ) + assert completed.status_code == 303 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 785bcc8ab70..a7c7e1853a0 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,7 +867,6 @@ 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, @@ -917,8 +916,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"].startswith( - 'Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp", scope="litellm:mcp:connect:' + assert response.headers["www-authenticate"] == ( + 'Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp"' ) 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 5777caf63de..0ba8b3f2a40 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,7 +10714,6 @@ 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)) diff --git a/ui/litellm-dashboard/src/components/chat/ConnectFlowBanner.test.tsx b/ui/litellm-dashboard/src/components/chat/ConnectFlowBanner.test.tsx index 5caf15d1fce..4c726b4a657 100644 --- a/ui/litellm-dashboard/src/components/chat/ConnectFlowBanner.test.tsx +++ b/ui/litellm-dashboard/src/components/chat/ConnectFlowBanner.test.tsx @@ -23,7 +23,7 @@ const unscoped = (client_origin: string): ConnectFlowStatus => ({ connected: null, }); -const renderBanner = (clientOrigin: string) => +const renderBanner = (clientOrigin: string, selectedServers: string[] = ["github"]) => render( accessToken="tok" onConnected={vi.fn()} failed={false} + selectedServers={selectedServers} />, ); describe("ConnectFlowBanner", () => { - it("posts only the flow handle to the proxy /authorize/complete as a full-page form", () => { + it("posts the flow handle and selected servers to the proxy /authorize/complete as a full-page form", () => { const { container } = renderBanner("https://claude.ai"); const form = container.querySelector("form")!; @@ -44,7 +45,14 @@ describe("ConnectFlowBanner", () => { expect(screen.getByDisplayValue("flow-handle-123")).toHaveAttribute("name", "flow"); expect(form.innerHTML).not.toContain("token"); expect(screen.getByRole("button", { name: /finish connecting/i })).toBeInTheDocument(); - expect(screen.queryByRole("button", { name: "Cancel" })).not.toBeInTheDocument(); + expect(screen.getByRole("button", { name: "Cancel" })).toBeInTheDocument(); + expect(new FormData(form).getAll("selected_servers")).toEqual(["github"]); + }); + + it("requires a selection before offering Finish but still allows cancellation", () => { + renderBanner("https://claude.ai", []); + expect(screen.queryByRole("button", { name: /finish connecting/i })).not.toBeInTheDocument(); + expect(screen.getByRole("button", { name: "Cancel" })).toBeInTheDocument(); }); it("offers manual delivery only for a loopback client, posted only when checked", () => { diff --git a/ui/litellm-dashboard/src/components/chat/ConnectFlowBanner.tsx b/ui/litellm-dashboard/src/components/chat/ConnectFlowBanner.tsx index 0d6e708f734..085e59884ce 100644 --- a/ui/litellm-dashboard/src/components/chat/ConnectFlowBanner.tsx +++ b/ui/litellm-dashboard/src/components/chat/ConnectFlowBanner.tsx @@ -11,6 +11,7 @@ interface Props { accessToken: string; onConnected: () => void; failed: boolean; + selectedServers?: readonly string[]; } /** Finish remains an explicit POST because a cross-site navigation must never mint a code. */ @@ -51,11 +52,17 @@ const copyFor = (flow: ConnectFlowStatus | undefined, failed: boolean): readonly ]; }; -const ConnectFlowBanner: React.FC = ({ flowHandle, flow, accessToken, onConnected, failed }) => { +const ConnectFlowBanner: React.FC = ({ + flowHandle, + flow, + accessToken, + onConnected, + failed, + selectedServers = [], +}) => { const action = `${getProxyBaseUrl()}/authorize/complete`; const state = failed || flow === undefined ? "stale" : flow.state; - const canFinish = state === "unscoped" || (state !== "stale" && flow?.connected === true); - const canCancel = state !== "unscoped"; + const canFinish = state === "unscoped" ? selectedServers.length > 0 : state !== "stale" && flow?.connected === true; const loopbackClient = isLoopbackOrigin(flow?.client_origin ?? null); const vendorServer = state === "interactive" && flow?.connected === false && flow.server_id !== null @@ -85,6 +92,10 @@ const ConnectFlowBanner: React.FC = ({ flowHandle, flow, accessToken, onC )}
+ {state === "unscoped" && + selectedServers.map((server) => ( + + ))} {canFinish && ( )} - {canCancel && ( - - )} + {loopbackClient && (