fix(mcp): validate selected upstream connections before completing OAuth

This commit is contained in:
Joshua Valluru 2026-09-28 14:48:39 -07:00
parent 3a8ed9773b
commit 8c7a4372ee
13 changed files with 231 additions and 187 deletions

View file

@ -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}"'},
)

View file

@ -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,
)

View file

@ -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):

View file

@ -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):

View file

@ -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": [
{

View file

@ -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

View file

@ -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:

View file

@ -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))

View file

@ -23,7 +23,7 @@ const unscoped = (client_origin: string): ConnectFlowStatus => ({
connected: null,
});
const renderBanner = (clientOrigin: string) =>
const renderBanner = (clientOrigin: string, selectedServers: string[] = ["github"]) =>
render(
<ConnectFlowBanner
flowHandle="flow-handle-123"
@ -31,11 +31,12 @@ const renderBanner = (clientOrigin: string) =>
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", () => {

View file

@ -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<Props> = ({ flowHandle, flow, accessToken, onConnected, failed }) => {
const ConnectFlowBanner: React.FC<Props> = ({
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<Props> = ({ flowHandle, flow, accessToken, onC
)}
<form method="POST" action={action}>
<input type="hidden" name="flow" value={flowHandle} />
{state === "unscoped" &&
selectedServers.map((server) => (
<input key={server} type="hidden" name="selected_servers" value={server} />
))}
{canFinish && (
<button
type="submit"
@ -93,16 +104,14 @@ const ConnectFlowBanner: React.FC<Props> = ({ flowHandle, flow, accessToken, onC
Finish connecting
</button>
)}
{canCancel && (
<button
type="submit"
name="decision"
value="deny"
className="ml-2 h-[38px] rounded-md border px-4 text-sm font-semibold text-foreground hover:bg-accent/40"
>
Cancel
</button>
)}
<button
type="submit"
name="decision"
value="deny"
className="ml-2 h-[38px] rounded-md border px-4 text-sm font-semibold text-foreground hover:bg-accent/40"
>
Cancel
</button>
{loopbackClient && (
<label className="mt-2 flex items-center gap-2 text-[13px] text-muted-foreground">
<input type="checkbox" name="delivery" value="manual" />

View file

@ -40,10 +40,10 @@ const flow = (state: "unscoped" | "interactive" | "m2m" | "stale", connected: bo
connected,
});
const renderSurface = () =>
const renderSurface = (selectedServers: string[] = []) =>
render(
<QueryClientProvider client={new QueryClient({ defaultOptions: { queries: { retry: false } } })}>
<ConnectFlowSurface accessToken="token-123" selectedServers={[]} onChange={vi.fn()} />
<ConnectFlowSurface accessToken="token-123" selectedServers={selectedServers} onChange={vi.fn()} />
</QueryClientProvider>,
);
@ -57,7 +57,7 @@ afterEach(() => {
describe("ConnectFlowSurface", () => {
it.each([
{ result: flow("unscoped"), grid: true, finish: true, cancel: false, oauthStarts: 0 },
{ result: flow("unscoped"), grid: true, finish: false, cancel: true, oauthStarts: 0 },
{ result: flow("interactive", false), grid: false, finish: false, cancel: true, oauthStarts: 1 },
{ result: flow("interactive", true), grid: false, finish: true, cancel: true, oauthStarts: 0 },
{ result: flow("m2m", true), grid: false, finish: true, cancel: true, oauthStarts: 0 },
@ -77,6 +77,16 @@ describe("ConnectFlowSurface", () => {
},
);
it("submits the selected upstreams with the protected flow", async () => {
state.connectFlow = "flow-handle-123";
vi.mocked(fetchConnectFlow).mockResolvedValue(flow("unscoped"));
renderSurface(["github", "slack"]);
const finish = await screen.findByRole("button", { name: /finish connecting/i });
const submitted = new FormData((finish as HTMLButtonElement).form!);
expect(submitted.get("flow")).toBe("flow-handle-123");
expect(submitted.getAll("selected_servers")).toEqual(["github", "slack"]);
});
it("keeps the grid and Finish hidden until the gateway accepts a handle", () => {
state.connectFlow = "flow-handle-123";
vi.mocked(fetchConnectFlow).mockReturnValue(new Promise(() => {}));

View file

@ -48,6 +48,7 @@ const ConnectFlowSurface: React.FC<Props> = ({ accessToken, selectedServers, onC
accessToken={accessToken}
onConnected={refetch}
failed={isError}
selectedServers={selectedServers}
/>
{flow?.state === "unscoped" && (
<MCPAppsPanel accessToken={accessToken} selectedServers={selectedServers} onChange={onChange} connectMode />

View file

@ -25696,6 +25696,8 @@ export interface components {
delivery?: string | null;
/** Flow */
flow: string;
/** Selected Servers */
selected_servers?: string[] | null;
/** Team Id */
team_id?: string | null;
};