mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(mcp): validate selected upstream connections before completing OAuth
This commit is contained in:
parent
3a8ed9773b
commit
8c7a4372ee
13 changed files with 231 additions and 187 deletions
|
|
@ -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}"'},
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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": [
|
||||
{
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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", () => {
|
||||
|
|
|
|||
|
|
@ -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" />
|
||||
|
|
|
|||
|
|
@ -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(() => {}));
|
||||
|
|
|
|||
|
|
@ -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 />
|
||||
|
|
|
|||
2
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
2
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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;
|
||||
};
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue