From 97a662f2b223f08e9208b9a11f2303fb92d1a6d6 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Mon, 28 Sep 2026 11:41:54 -0700 Subject: [PATCH 1/9] fix(mcp): preserve upstream authentication challenges on protocol requests --- .../mcp_server/faults/list_outcomes.py | 45 ++- .../mcp_server/mcp_server_manager.py | 129 ++++--- .../_experimental/mcp_server/operations.py | 212 +++++------ .../proxy/_experimental/mcp_server/server.py | 172 +++++++-- .../mcp_server/faults/test_list_outcomes.py | 54 +++ .../test_mcp_oauth_passthrough_tools.py | 356 ++++++++++++++++++ .../mcp_server/test_mcp_server.py | 3 +- .../mcp_server/test_mcp_server_manager.py | 14 +- 8 files changed, 747 insertions(+), 238 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/faults/list_outcomes.py b/litellm/proxy/_experimental/mcp_server/faults/list_outcomes.py index b96a7a74e4a..0059498e60a 100644 --- a/litellm/proxy/_experimental/mcp_server/faults/list_outcomes.py +++ b/litellm/proxy/_experimental/mcp_server/faults/list_outcomes.py @@ -10,13 +10,14 @@ becomes an outcome, never a second failure. from __future__ import annotations -from collections.abc import Iterator +from collections.abc import Iterator, Mapping from typing import Final, Literal, NamedTuple, NoReturn, TypeAlias import httpx import httpx2 +from fastapi import HTTPException from mcp.types import Tool as MCPTool -from pydantic import BaseModel, ConfigDict +from pydantic import BaseModel, ConfigDict, Field from typing_extensions import assert_never from litellm.proxy._experimental.mcp_server.exceptions import ( @@ -50,6 +51,8 @@ class ServerListFault(BaseModel): model_config = ConfigDict(frozen=True) tag: ListFaultCategory status_code: int | None = None + www_authenticate: str | None = Field(default=None, exclude=True, repr=False) + server_name: str | None = Field(default=None, exclude=True, repr=False) ServerOutcome: TypeAlias = ServerListOk | ServerListFault @@ -64,6 +67,20 @@ class AggregateToolListing(NamedTuple): outcomes: dict[str, ServerOutcome] +def listing_auth_error(outcomes: Mapping[str, ServerOutcome]) -> MCPUpstreamAuthError | None: + blocked: Final = tuple( + (name, outcome) + for name, outcome in outcomes.items() + if isinstance(outcome, ServerListFault) and outcome.tag in ("auth_required", "forbidden") + ) + if not blocked or len(blocked) != len(outcomes): + return None + name, outcome = next((entry for entry in blocked if entry[1].tag == "auth_required"), blocked[0]) + return MCPUpstreamAuthError( + 401 if outcome.tag == "auth_required" else 403, outcome.www_authenticate, outcome.server_name or name + ) + + def _iter_upstream_responses(exc: BaseException) -> Iterator[httpx.Response | httpx2.Response]: """Yield every upstream ``httpx``/``httpx2`` ``Response`` in the exception tree, in the shared traversal's deliberate order (explicit causes first, ExceptionGroup members in raise order, the incidental @@ -87,9 +104,20 @@ def upstream_auth_challenge(exc: BaseException) -> tuple[int, str | None] | None rides with it can never come from two different responses in the tree. Non-auth responses do not end the scan: a causal 401 behind an unrelated 5xx must still be found, or the client never receives the challenge it needs to re-authenticate.""" - for response in _iter_upstream_responses(exc): - if response.status_code in (401, 403): - return response.status_code, response.headers.get("www-authenticate") + return next( + (challenge for current in iter_exception_tree(exc) if (challenge := _auth_challenge(current)) is not None), None + ) + + +def _auth_challenge(exc: BaseException) -> tuple[int, str | None] | None: + if isinstance(exc, MCPUpstreamAuthError): + return exc.status_code, exc.www_authenticate + if isinstance(exc, HTTPException) and exc.status_code in (401, 403): + headers: Final = exc.headers or {} + return exc.status_code, headers.get("WWW-Authenticate") or headers.get("www-authenticate") + response: Final = getattr(exc, "response", None) + if isinstance(response, (httpx.Response, httpx2.Response)) and response.status_code in (401, 403): + return response.status_code, response.headers.get("www-authenticate") return None @@ -122,17 +150,20 @@ def classify_list_exception(exc: BaseException) -> ServerListFault: return exc.fault if isinstance(exc, MCPUpstreamAuthError): tag: Final = "forbidden" if exc.status_code == 403 else "auth_required" - return ServerListFault(tag=tag, status_code=exc.status_code) + return ServerListFault( + tag=tag, status_code=exc.status_code, www_authenticate=exc.www_authenticate, server_name=exc.server_name + ) if isinstance(exc, TimeoutError): return ServerListFault(tag="timeout") if isinstance(exc, ConnectionError): return ServerListFault(tag="unreachable") auth: Final = upstream_auth_challenge(exc) if auth is not None: - status_code, _ = auth + status_code, challenge = auth return ServerListFault( tag="forbidden" if status_code == 403 else "auth_required", status_code=status_code, + www_authenticate=challenge, ) response: Final = _find_upstream_response(exc) if response is not None: diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index ea685c431bd..375cf9fc200 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -4632,6 +4632,8 @@ class MCPServerManager: return self._create_prefixed_prompts(items, server, add_prefix=add_prefix) except Exception as error: verbose_logger.warning("Failed to get prompts from server %s: %s", server.name, error) + if upstream_auth_challenge(error) is not None: + raise_classified_list_failure(error, server.name, suppress_challenge=server.is_dcr_bridge) return [] async def get_resources_from_server( @@ -4678,6 +4680,8 @@ class MCPServerManager: return self._create_prefixed_resources(items, server, add_prefix=add_prefix) except Exception as error: verbose_logger.warning("Failed to get resources from server %s: %s", server.name, error) + if upstream_auth_challenge(error) is not None: + raise_classified_list_failure(error, server.name, suppress_challenge=server.is_dcr_bridge) return [] async def get_resource_templates_from_server( @@ -4724,6 +4728,8 @@ class MCPServerManager: return self._create_prefixed_resource_templates(items, server, add_prefix=add_prefix) except Exception as error: verbose_logger.warning("Failed to get resource_templates from server %s: %s", server.name, error) + if upstream_auth_challenge(error) is not None: + raise_classified_list_failure(error, server.name, suppress_challenge=server.is_dcr_bridge) return [] async def read_resource_from_server( @@ -4738,29 +4744,38 @@ class MCPServerManager: ) -> ReadResourceResult: """Read resource contents from a specific MCP server.""" - verbose_logger.debug("Connecting to url: %s", server.url) - verbose_logger.info("read_resource_from_server for %s...", server.name) + try: + verbose_logger.debug("Connecting to url: %s", server.url) + verbose_logger.info("read_resource_from_server for %s...", server.name) - if server.static_headers: - if extra_headers is None: - extra_headers = {} - extra_headers.update(server.static_headers) + if server.static_headers: + if extra_headers is None: + extra_headers = {} + extra_headers.update(server.static_headers) - stdio_env: Final = self._build_stdio_env(server, raw_headers) - subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth) + stdio_env: Final = self._build_stdio_env(server, raw_headers) + subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth) - client: Final = await self._create_mcp_client( - server=server, - mcp_auth_header=mcp_auth_header, - extra_headers=extra_headers, - stdio_env=stdio_env, - subject_token=subject_token, - raw_headers=raw_headers, - client_ip=client_ip, - user_api_key_auth=user_api_key_auth, - ) + client: Final = await self._create_mcp_client( + server=server, + mcp_auth_header=mcp_auth_header, + extra_headers=extra_headers, + stdio_env=stdio_env, + subject_token=subject_token, + raw_headers=raw_headers, + client_ip=client_ip, + user_api_key_auth=user_api_key_auth, + ) - return await client.read_resource(url) + return await client.read_resource(url) + except Exception as exc: + auth_failure: Final = upstream_auth_challenge(exc) + if auth_failure is not None: + status_code, challenge = auth_failure + raise MCPUpstreamAuthError( + status_code, None if server.is_dcr_bridge else challenge, server.name + ) from exc + raise async def get_prompt_from_server( self, @@ -4775,33 +4790,42 @@ class MCPServerManager: ) -> GetPromptResult: """Fetch a specific prompt definition from a single MCP server.""" - verbose_logger.debug("Connecting to url: %s", server.url) - verbose_logger.info("get_prompt_from_server for %s...", server.name) + try: + verbose_logger.debug("Connecting to url: %s", server.url) + verbose_logger.info("get_prompt_from_server for %s...", server.name) - if server.static_headers: - if extra_headers is None: - extra_headers = {} - extra_headers.update(server.static_headers) + if server.static_headers: + if extra_headers is None: + extra_headers = {} + extra_headers.update(server.static_headers) - stdio_env: Final = self._build_stdio_env(server, raw_headers) - subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth) + stdio_env: Final = self._build_stdio_env(server, raw_headers) + subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth) - client: Final = await self._create_mcp_client( - server=server, - mcp_auth_header=mcp_auth_header, - extra_headers=extra_headers, - stdio_env=stdio_env, - subject_token=subject_token, - raw_headers=raw_headers, - client_ip=client_ip, - user_api_key_auth=user_api_key_auth, - ) + client: Final = await self._create_mcp_client( + server=server, + mcp_auth_header=mcp_auth_header, + extra_headers=extra_headers, + stdio_env=stdio_env, + subject_token=subject_token, + raw_headers=raw_headers, + client_ip=client_ip, + user_api_key_auth=user_api_key_auth, + ) - get_prompt_request_params: Final = GetPromptRequestParams( - name=prompt_name, - arguments=arguments, - ) - return await client.get_prompt(get_prompt_request_params) + get_prompt_request_params: Final = GetPromptRequestParams( + name=prompt_name, + arguments=arguments, + ) + return await client.get_prompt(get_prompt_request_params) + except Exception as exc: + auth_failure: Final = upstream_auth_challenge(exc) + if auth_failure is not None: + status_code, challenge = auth_failure + raise MCPUpstreamAuthError( + status_code, None if server.is_dcr_bridge else challenge, server.name + ) from exc + raise @staticmethod def _is_same_authority_metadata_url(url: str, server_url: str) -> bool: @@ -6094,29 +6118,10 @@ class MCPServerManager: tool_call_coro = _obo_call_tool_limited() else: - # Scoped to the two client-forwarded token modes this stack introduced; legacy - # oauth2 + delegate_auth_to_upstream (is_oauth_passthrough) is being removed, so it is not - # added here even though the list path still relays for it. - relays_upstream_auth: Final = mcp_server.is_client_forwarded_token server_label: Final = mcp_server.name or mcp_server.server_name or mcp_server.alias or "" async def _call_tool_via_client(client, params): async with self._limit_outbound_concurrency(mcp_server): - if not relays_upstream_auth: - return await client.call_tool( - params, - host_progress_callback=host_progress_callback, - allow_input_required=allow_input_required, - ) - # The client-forwarded modes carry the caller's own upstream token, so an upstream - # 401 (expired/invalid token) is the caller's to resolve: relay it as - # MCPUpstreamAuthError so single-server REST callers turn it into a 401 + - # WWW-Authenticate and re-run the upstream OAuth flow. Only 401 is a re-auth signal - # (mirrors the list path and MCPUpstreamAuthError's contract); a 403 is a genuine - # authorization failure that re-auth won't fix, so it takes the non-auth branch and - # stays a visible warning. raise_on_error only re-raises transport failures - # (tool-level isError results are still returned normally); a non-auth failure keeps - # the same isError degradation the default path produces. try: return await client.call_tool( params, @@ -6142,7 +6147,7 @@ class MCPServerManager: _, www_authenticate = auth_info raise MCPUpstreamAuthError( status_code=401, - www_authenticate=www_authenticate, + www_authenticate=None if mcp_server.is_dcr_bridge else www_authenticate, server_name=server_label, ) from e diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index a19246b6e90..0c037ebf893 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -4,9 +4,10 @@ import asyncio import traceback import types import uuid -from collections.abc import Mapping, Sequence +from collections.abc import Awaitable, Callable, Mapping, Sequence from datetime import datetime -from typing import Any, Final, NoReturn, TypeAlias, overload +from itertools import chain +from typing import Any, Final, NoReturn, TypeAlias, TypeVar, overload from fastapi import HTTPException from mcp import ReadResourceResult, Resource @@ -78,6 +79,7 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ( ServerListOk, ServerOutcome, classify_list_exception, + listing_auth_error, outcome_wire_value, ) from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( @@ -1242,6 +1244,26 @@ async def _get_tools_from_mcp_servers( raise +_ListingItem = TypeVar("_ListingItem", Prompt, Resource, ResourceTemplate) + + +async def _collect_mcp_listing( + servers: Sequence[MCPServer], fetch: Callable[[MCPServer], Awaitable[list[_ListingItem]]] +) -> list[_ListingItem]: + async def fetch_one(server: MCPServer) -> tuple[list[_ListingItem], ServerOutcome]: + try: + items: Final = await fetch(server) + return items, ServerListOk(tool_count=len(items)) + except Exception as exc: + return [], classify_list_exception(exc) + + results: Final = await asyncio.gather(*(fetch_one(server) for server in servers)) + failure: Final = listing_auth_error({server.name: result[1] for server, result in zip(servers, results)}) + if failure is not None: + raise failure + return list(chain.from_iterable(items for items, _ in results)) + + async def _get_prompts_from_mcp_servers( user_api_key_auth: UserAPIKeyAuth | None, mcp_auth_header: str | None, @@ -1251,32 +1273,13 @@ async def _get_prompts_from_mcp_servers( raw_headers: dict[str, str] | None = None, client_ip: str | None = None, ) -> list[Prompt]: - """ - Helper method to fetch prompt from MCP servers based on server filtering criteria. - - Args: - user_api_key_auth: User authentication info for access control - mcp_auth_header: Optional auth header for MCP server (deprecated) - mcp_servers: Optional list of server names/aliases to filter by - mcp_server_auth_headers: Optional dict of server-specific auth headers - oauth2_headers: Optional dict of oauth2 headers - - Returns: - List[Prompt]: Combined list of prompts from filtered servers - """ - - allowed_mcp_servers: Final = await _get_allowed_mcp_servers( + allowed: Final = await _get_allowed_mcp_servers( user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip, ) - # Get prompts from each allowed server - all_prompts: Final = [] - for server in allowed_mcp_servers: - if server is None: - continue - + async def fetch(server: MCPServer) -> list[Prompt]: server_auth_header, extra_headers = _prepare_mcp_server_headers( server=server, mcp_server_auth_headers=mcp_server_auth_headers, @@ -1284,30 +1287,19 @@ async def _get_prompts_from_mcp_servers( oauth2_headers=oauth2_headers, raw_headers=raw_headers, user_api_key_auth=user_api_key_auth, - scope_servers=allowed_mcp_servers, + scope_servers=allowed, + ) + return await global_mcp_server_manager.get_prompts_from_server( + server=server, + user_api_key_auth=user_api_key_auth, + mcp_auth_header=server_auth_header, + extra_headers=extra_headers, + add_prefix=True, + raw_headers=raw_headers, + client_ip=client_ip, ) - try: - prompts = await global_mcp_server_manager.get_prompts_from_server( - server=server, - user_api_key_auth=user_api_key_auth, - mcp_auth_header=server_auth_header, - extra_headers=extra_headers, - add_prefix=True, # Always add server prefix - raw_headers=raw_headers, - client_ip=client_ip, - ) - - all_prompts.extend(prompts) - - verbose_logger.debug("Successfully fetched %s prompts from server %s", len(prompts), server.name) - except Exception as e: - verbose_logger.exception("Error getting prompts from server %s: %s", server.name, e) - # Continue with other servers instead of failing completely - - verbose_logger.info("Successfully fetched %s prompts total from all MCP servers", len(all_prompts)) - - return all_prompts + return await _collect_mcp_listing(tuple(server for server in allowed if server is not None), fetch) async def _get_resources_from_mcp_servers( @@ -1319,19 +1311,13 @@ async def _get_resources_from_mcp_servers( raw_headers: dict[str, str] | None = None, client_ip: str | None = None, ) -> list[Resource]: - """Fetch resources from allowed MCP servers.""" - - allowed_mcp_servers: Final = await _get_allowed_mcp_servers( + allowed: Final = await _get_allowed_mcp_servers( user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip, ) - all_resources: Final[list[Resource]] = [] - for server in allowed_mcp_servers: - if server is None: - continue - + async def fetch(server: MCPServer) -> list[Resource]: server_auth_header, extra_headers = _prepare_mcp_server_headers( server=server, mcp_server_auth_headers=mcp_server_auth_headers, @@ -1339,28 +1325,19 @@ async def _get_resources_from_mcp_servers( oauth2_headers=oauth2_headers, raw_headers=raw_headers, user_api_key_auth=user_api_key_auth, - scope_servers=allowed_mcp_servers, + scope_servers=allowed, + ) + return await global_mcp_server_manager.get_resources_from_server( + server=server, + user_api_key_auth=user_api_key_auth, + mcp_auth_header=server_auth_header, + extra_headers=extra_headers, + add_prefix=True, + raw_headers=raw_headers, + client_ip=client_ip, ) - try: - resources = await global_mcp_server_manager.get_resources_from_server( - server=server, - user_api_key_auth=user_api_key_auth, - mcp_auth_header=server_auth_header, - extra_headers=extra_headers, - add_prefix=True, # Always add server prefix - raw_headers=raw_headers, - client_ip=client_ip, - ) - all_resources.extend(resources) - - verbose_logger.debug("Successfully fetched %s resources from server %s", len(resources), server.name) - except Exception as e: - verbose_logger.exception("Error getting resources from server %s: %s", server.name, e) - - verbose_logger.info("Successfully fetched %s resources total from all MCP servers", len(all_resources)) - - return all_resources + return await _collect_mcp_listing(tuple(server for server in allowed if server is not None), fetch) async def _get_resource_templates_from_mcp_servers( @@ -1372,19 +1349,13 @@ async def _get_resource_templates_from_mcp_servers( raw_headers: dict[str, str] | None = None, client_ip: str | None = None, ) -> list[ResourceTemplate]: - """Fetch resource templates from allowed MCP servers.""" - - allowed_mcp_servers: Final = await _get_allowed_mcp_servers( + allowed: Final = await _get_allowed_mcp_servers( user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip, ) - all_resource_templates: Final[list[ResourceTemplate]] = [] - for server in allowed_mcp_servers: - if server is None: - continue - + async def fetch(server: MCPServer) -> list[ResourceTemplate]: server_auth_header, extra_headers = _prepare_mcp_server_headers( server=server, mcp_server_auth_headers=mcp_server_auth_headers, @@ -1392,38 +1363,19 @@ async def _get_resource_templates_from_mcp_servers( oauth2_headers=oauth2_headers, raw_headers=raw_headers, user_api_key_auth=user_api_key_auth, - scope_servers=allowed_mcp_servers, + scope_servers=allowed, + ) + return await global_mcp_server_manager.get_resource_templates_from_server( + server=server, + user_api_key_auth=user_api_key_auth, + mcp_auth_header=server_auth_header, + extra_headers=extra_headers, + add_prefix=True, + raw_headers=raw_headers, + client_ip=client_ip, ) - try: - resource_templates = await global_mcp_server_manager.get_resource_templates_from_server( - server=server, - user_api_key_auth=user_api_key_auth, - mcp_auth_header=server_auth_header, - extra_headers=extra_headers, - add_prefix=True, # Always add server prefix - raw_headers=raw_headers, - client_ip=client_ip, - ) - all_resource_templates.extend(resource_templates) - verbose_logger.debug( - "Successfully fetched %s resource templates from server %s", - len(resource_templates), - server.name, - ) - except Exception as e: - verbose_logger.exception( - "Error getting resource templates from server %s: %s", - server.name, - str(e), - ) - - verbose_logger.info( - "Successfully fetched %s resource templates total from all MCP servers", - len(all_resource_templates), - ) - - return all_resource_templates + return await _collect_mcp_listing(tuple(server for server in allowed if server is not None), fetch) async def filter_tools_by_key_team_permissions( @@ -1539,6 +1491,8 @@ async def _list_mcp_prompts( client_ip=client_ip, ) verbose_logger.debug("Successfully fetched %s prompts from managed MCP servers", len(managed_prompts)) + except MCPUpstreamAuthError: + raise except Exception as e: verbose_logger.exception("Error getting tools from managed MCP servers: %s", e) # Continue with empty managed tools list instead of failing completely @@ -1569,6 +1523,8 @@ async def _list_mcp_resources( client_ip=client_ip, ) verbose_logger.debug("Successfully fetched %s resources from managed MCP servers", len(managed_resources)) + except MCPUpstreamAuthError: + raise except Exception as e: verbose_logger.exception("Error getting resources from managed MCP servers: %s", e) @@ -1601,6 +1557,8 @@ async def _list_mcp_resource_templates( "Successfully fetched %s resource templates from managed MCP servers", len(managed_resource_templates), ) + except MCPUpstreamAuthError: + raise except Exception as e: verbose_logger.exception( "Error getting resource templates from managed MCP servers: %s", @@ -2721,6 +2679,9 @@ async def _execute_handle_list_tools( list_tools_log_source="mcp_protocol", client_ip=_client_ip, ) + auth_failure: Final = listing_auth_error(listing.outcomes) + if auth_failure is not None: + raise auth_failure verbose_logger.info("MCP list_tools - Successfully returned %s tools", len(listing.tools)) if not listing.outcomes: return ListToolsResult(tools=listing.tools) @@ -2728,6 +2689,8 @@ async def _execute_handle_list_tools( SERVER_OUTCOMES_META_KEY: {key: outcome_wire_value(outcome) for key, outcome in listing.outcomes.items()} } return ListToolsResult.model_validate({"tools": listing.tools, "_meta": outcome_meta}) + except MCPUpstreamAuthError: + raise except HTTPException as e: from mcp.shared.exceptions import MCPError from mcp.types import INVALID_REQUEST @@ -2857,27 +2820,16 @@ async def _execute_mcp_server_tool_call( is_error=True, ) except HTTPException as e: + if e.status_code == 401 and e.headers and any(name.lower() == "www-authenticate" for name in e.headers): + raise verbose_logger.error("HTTPException in MCP tool call: %s", e) return CallToolResult( content=[TextContent(text=f"Error: {_http_detail_message(e.detail)}", type="text")], is_error=True, ) - except MCPUpstreamAuthError as e: - # The MCP session manager serializes handler exceptions as JSON-RPC errors, so a - # mid-session tool call cannot emit a raw 401 + WWW-Authenticate the way the REST - # call path and the connect-time preemptive check do. Return an explicit isError - # naming the upstream status (at info level, not a traceback) so the client still - # learns it must re-authenticate upstream and expected pass-through 401s don't spam. - verbose_logger.info("Upstream auth failure calling MCP tool: HTTP %s", e.status_code) - return CallToolResult( - content=[ - TextContent( - text=f"Error: upstream authentication required (HTTP {e.status_code})", - type="text", - ) - ], - is_error=True, - ) + except MCPUpstreamAuthError as exc: + verbose_logger.info("Upstream auth failure calling MCP tool: HTTP %s", exc.status_code) + raise except Exception as e: verbose_logger.exception("MCP mcp_server_tool_call - error: %s", e) return CallToolResult( @@ -2922,6 +2874,8 @@ async def _execute_list_prompts( ) verbose_logger.info("MCP list_prompts - Successfully returned %s prompts", len(prompts)) return ListPromptsResult(prompts=prompts) + except MCPUpstreamAuthError: + raise except Exception as e: verbose_logger.exception("Error in list_prompts endpoint: %s", e) # Return empty list instead of failing completely @@ -2991,6 +2945,8 @@ async def _execute_list_resources( ) verbose_logger.info("MCP list_resources - Successfully returned %s resources", len(resources)) return ListResourcesResult(resources=resources) + except MCPUpstreamAuthError: + raise except Exception as e: verbose_logger.exception("Error in list_resources endpoint: %s", e) return ListResourcesResult(resources=[]) # mutable-ok: MCP result payload @@ -3031,6 +2987,8 @@ async def _execute_list_resource_templates( "MCP list_resource_templates - Successfully returned %s resource templates", len(resource_templates) ) return ListResourceTemplatesResult(resource_templates=resource_templates) + except MCPUpstreamAuthError: + raise except Exception as e: verbose_logger.exception("Error in list_resource_templates endpoint: %s", e) return ListResourceTemplatesResult(resource_templates=[]) # mutable-ok: MCP result payload diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 555aebc7434..926aead9b7a 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -92,6 +92,62 @@ if TYPE_CHECKING: from mcp.server.session import ServerSession as _McpServerSession +_MCP_AUTH_RESPONSE_SCOPE_KEY: Final = "litellm_mcp_auth_response" + + +class MCPAuthResponse: + """Defer HTTP success until the SDK produces data, preserving late auth challenges.""" + + def __init__(self, send: Send) -> None: + self._send = send + self._start: Message | None = None + self._preamble: tuple[Message, ...] = () + self._committed = False + self._replaced = False + self._sse = False + self.challenge: HTTPException | None = None + + async def send(self, message: Message) -> None: + if self._replaced: + return + if message["type"] == "http.response.start" and message["status"] == 200: + self._start = message + self._sse = any( + name.lower() == b"content-type" and b"text/event-stream" in value + for name, value in message.get("headers", ()) + ) + return + if self._start is None or self._committed: + await self._send(message) + return + body: Final = message.get("body", b"") + if ( + self._sse + and message.get("more_body", False) + and not any(line.startswith(b"data:") and line[5:].strip() for line in body.splitlines()) + ): + if not body.startswith(b":"): + self._preamble = (*self._preamble, message) + return + self._committed = True + if self.challenge is not None: + self._replaced = True + response: Final = JSONResponse( + {"detail": self.challenge.detail}, + status_code=self.challenge.status_code, + headers=self.challenge.headers, + ) + await self._send( + {"type": "http.response.start", "status": response.status_code, "headers": response.raw_headers} + ) + await self._send({"type": "http.response.body", "body": response.body, "more_body": False}) + return + await self._send(self._start) + for preamble in self._preamble: + await self._send(preamble) + await self._send(message) + + _STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS: Final = 30 * 60 # Upper bound on concurrent stateful sessions a single caller may hold. Each # `initialize` creates a session that survives until the idle timeout, so @@ -530,6 +586,7 @@ if MCP_AVAILABLE: ListToolsResult, PaginatedRequestParams, ReadResourceRequestParams, + TextContent, ) from litellm.proxy._experimental.mcp_server.auth.litellm_auth_handler import ( @@ -806,38 +863,60 @@ if MCP_AVAILABLE: @contextlib.asynccontextmanager async def _legacy_operation_context(ctx: ServerRequestContext, *, trace: bool) -> AsyncGenerator[OperationContext]: - with contextlib.ExitStack() as cleanup: - cleanup.callback(active_mcp_request_ctx_var.reset, active_mcp_request_ctx_var.set(ctx)) - cleanup.callback(active_mcp_session_var.reset, active_mcp_session_var.set(ctx.session)) - if trace: - cleanup.callback( - _otel_reset_mcp_trace_carrier, _otel_set_mcp_trace_carrier(_mcp_meta_trace_carrier(ctx)) + try: + with contextlib.ExitStack() as cleanup: + cleanup.callback(active_mcp_request_ctx_var.reset, active_mcp_request_ctx_var.set(ctx)) + cleanup.callback(active_mcp_session_var.reset, active_mcp_session_var.set(ctx.session)) + if trace: + cleanup.callback( + _otel_reset_mcp_trace_carrier, _otel_set_mcp_trace_carrier(_mcp_meta_trace_carrier(ctx)) + ) + cleanup.callback( + _otel_reset_mcp_transport_span, + _otel_set_mcp_transport_span(_otel_transport_span_from_message(ctx)), + ) + cleanup.callback(_otel_reset_mcp_request_destinations, _otel_set_mcp_request_destinations(ctx)) + ( + auth, + token, + servers, + server_headers, + oauth_headers, + headers, + client_ip, + ) = await get_or_extract_auth_context() + yield operations.prepare_context( + auth, + token, + servers, + server_headers, + oauth_headers, + headers, + client_ip, + _mcp_proxy_mode.get(), + wire_compat_for(ctx.protocol_version), + ctx.protocol_version, ) - cleanup.callback( - _otel_reset_mcp_transport_span, _otel_set_mcp_transport_span(_otel_transport_span_from_message(ctx)) - ) - cleanup.callback(_otel_reset_mcp_request_destinations, _otel_set_mcp_request_destinations(ctx)) - ( - auth, - token, - servers, - server_headers, - oauth_headers, - headers, - client_ip, - ) = await get_or_extract_auth_context() - yield operations.prepare_context( - auth, - token, - servers, - server_headers, - oauth_headers, - headers, - client_ip, - _mcp_proxy_mode.get(), - wire_compat_for(ctx.protocol_version), - ctx.protocol_version, - ) + except (MCPUpstreamAuthError, HTTPException) as exc: + if isinstance(exc, HTTPException) and ( + exc.status_code != 401 + or not exc.headers + or not any(name.lower() == "www-authenticate" for name in exc.headers) + ): + raise + if isinstance(ctx.request, StarletteRequest): + response: Final = ctx.request.scope.get(_MCP_AUTH_RESPONSE_SCOPE_KEY) + if isinstance(response, MCPAuthResponse): + response.challenge = ( + exc.to_http_exception( + base_url=get_request_base_url(ctx.request), request_path=ctx.request.url.path + ) + if isinstance(exc, MCPUpstreamAuthError) + else exc + ) + raise MCPError( + code=INVALID_REQUEST, message=f"Upstream authorization failed (HTTP {exc.status_code})" + ) from exc async def handle_list_tools(ctx: ServerRequestContext, params: PaginatedRequestParams) -> ListToolsResult: try: @@ -896,9 +975,21 @@ if MCP_AVAILABLE: async def mcp_server_tool_call( ctx: ServerRequestContext, params: CallToolRequestParams ) -> CallToolResult | InputRequiredResult: - async with _legacy_operation_context(ctx, trace=True) as context: - return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute( - CallToolRequest(params=params), context + try: + async with _legacy_operation_context(ctx, trace=True) as context: + return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute( + CallToolRequest(params=params), context + ) + except MCPError as exc: + if not isinstance(exc.__cause__, (MCPUpstreamAuthError, HTTPException)): + raise + return CallToolResult( + content=[ + TextContent( + type="text", text=f"Error: upstream authentication required (HTTP {exc.__cause__.status_code})" + ) + ], + is_error=True, ) async def list_prompts(ctx: ServerRequestContext, params: PaginatedRequestParams) -> ListPromptsResult: @@ -909,6 +1000,8 @@ if MCP_AVAILABLE: return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute( ListPromptsRequest(params=params), context ) + except MCPError: + raise except Exception as exc: # noqa: BLE001 # preserve native listing fallback for ingress failures verbose_logger.exception("Error in list_prompts endpoint: %s", exc) return ListPromptsResult(prompts=[]) @@ -929,6 +1022,8 @@ if MCP_AVAILABLE: return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute( ListResourcesRequest(params=params), context ) + except MCPError: + raise except Exception as exc: # noqa: BLE001 # preserve native listing fallback for ingress failures verbose_logger.exception("Error in list_resources endpoint: %s", exc) return ListResourcesResult(resources=[]) @@ -943,6 +1038,8 @@ if MCP_AVAILABLE: return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute( ListResourceTemplatesRequest(params=params), context ) + except MCPError: + raise except Exception as exc: # noqa: BLE001 # preserve native listing fallback for ingress failures verbose_logger.exception("Error in list_resource_templates endpoint: %s", exc) return ListResourceTemplatesResult(resource_templates=[]) @@ -2248,7 +2345,12 @@ if MCP_AVAILABLE: scoped_server_endpoint=scoped_server_endpoint, is_initialize=is_initialize, ): - await target_manager.handle_request(scope, receive, local_send) + if request_method == "POST" and body and not is_initialize: + auth_response: Final = MCPAuthResponse(local_send) + scope[_MCP_AUTH_RESPONSE_SCOPE_KEY] = auth_response + await target_manager.handle_request(scope, receive, auth_response.send) + else: + await target_manager.handle_request(scope, receive, local_send) if use_stateful and session_id and scope.get("method") == "DELETE": _remove_stateful_session_tracking(session_id) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_list_outcomes.py b/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_list_outcomes.py index f951499e18f..fc02fed0945 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_list_outcomes.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_list_outcomes.py @@ -255,3 +255,57 @@ def test_pure_non_auth_response_still_classifies_upstream_error(): fault = classify_list_exception(exc) assert fault.tag == "upstream_error" assert fault.status_code == 502 + + +@pytest.mark.parametrize( + "outcomes,expected_status", + ( + ({}, None), + ({"empty": ServerListOk(tool_count=0)}, None), + ({"auth": ServerListFault(tag="auth_required", status_code=401)}, 401), + ({"denied": ServerListFault(tag="forbidden", status_code=403)}, 403), + ({"denied": ServerListFault(tag="forbidden", status_code=403), "auth": ServerListFault(tag="auth_required", status_code=401)}, 401), + ({"auth": ServerListFault(tag="auth_required", status_code=401), "empty": ServerListOk(tool_count=0)}, None), + ({"auth": ServerListFault(tag="auth_required", status_code=401), "healthy": ServerListOk(tool_count=2)}, None), + ({"auth": ServerListFault(tag="auth_required", status_code=401), "timeout": ServerListFault(tag="timeout")}, None), + ), +) +def test_listing_auth_failure_requires_every_server_to_be_blocked( + outcomes: dict[str, ServerListOk | ServerListFault], expected_status: int | None +) -> None: + from litellm.proxy._experimental.mcp_server.faults.list_outcomes import listing_auth_error + from typing import Final + failure: Final = listing_auth_error(outcomes) + assert (failure.status_code if failure else None) == expected_status + + +def test_classified_auth_challenge_is_preserved_without_serializing_it() -> None: + from typing import Final + from fastapi import HTTPException + from litellm.proxy._experimental.mcp_server.faults.list_outcomes import listing_auth_error + challenge: Final = 'Bearer resource_metadata="https://gateway/.well-known/oauth-protected-resource/mcp"' + fault: Final = classify_list_exception(HTTPException(401, headers={"WWW-Authenticate": challenge})) + failure: Final = listing_auth_error({"upstream": fault}) + assert failure is not None + assert failure.www_authenticate == challenge + assert failure.server_name == "upstream" + assert "www_authenticate" not in fault.model_dump() + assert challenge not in repr(fault) + assert outcome_wire_value(fault) == {"status": "auth_required", "http_status": 401} + + +def test_auth_recovery_uses_routable_server_name_without_exposing_it_in_outcomes() -> None: + from typing import Final + from litellm.proxy._experimental.mcp_server.faults.list_outcomes import listing_auth_error + fault: Final = classify_list_exception(MCPUpstreamAuthError(401, None, "upstream-internal")) + failure: Final = listing_auth_error({"short-prefix": fault}) + assert failure is not None + assert failure.server_name == "upstream-internal" + assert "server_name" not in fault.model_dump() + assert "upstream-internal" not in repr(fault) + assert outcome_wire_value(fault) == {"status": "auth_required", "http_status": 401} + + +def test_auth_challenge_traversal_preserves_typed_carrier() -> None: + from litellm.proxy._experimental.mcp_server.faults.list_outcomes import upstream_auth_challenge + assert upstream_auth_challenge(MCPUpstreamAuthError(401, "Bearer", "upstream")) == (401, "Bearer") 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 1909e3306a2..15aa089637a 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 @@ -3,6 +3,7 @@ from litellm.proxy._experimental.mcp_server import operations as mcp_operations import logging import sys +from typing import Final from unittest.mock import AsyncMock, MagicMock import httpx @@ -498,3 +499,358 @@ def test_passthrough_admission_recognizes_only_matching_authorization(oauth_head server = MCPServer(server_id="catalog", name="catalog", alias="catalog", transport=MCPTransport.http) assert _client_has_passthrough_authorization(server, oauth_headers, server_headers) is authorized + + +@pytest.mark.asyncio +async def test_listing_transport_preserves_auth_challenge_before_sse_success() -> None: + from starlette.exceptions import HTTPException + from litellm.proxy._experimental.mcp_server.server import MCPAuthResponse + + send: Final = AsyncMock() + response: Final = MCPAuthResponse(send) + await response.send( + {"type": "http.response.start", "status": 200, "headers": [(b"content-type", b"text/event-stream")]} + ) + await response.send({"type": "http.response.body", "body": b": ping\r\n\r\n", "more_body": True}) + response.challenge = HTTPException( + 401, "Unauthorized", headers={"WWW-Authenticate": 'Bearer resource_metadata="http://localhost/mcp-metadata"'} + ) + await response.send( + { + "type": "http.response.body", + "body": b'event: message\r\ndata: {"jsonrpc":"2.0","id":1,"result":{"tools":[]}}\r\n\r\n', + "more_body": True, + } + ) + await response.send({"type": "http.response.body", "body": b"", "more_body": False}) + sent: Final = tuple(call.args[0] for call in send.await_args_list) + assert [m["status"] for m in sent if m["type"] == "http.response.start"] == [401] + assert dict(sent[0]["headers"])[b"www-authenticate"] == b'Bearer resource_metadata="http://localhost/mcp-metadata"' + assert b'"tools"' not in b"".join(m.get("body", b"") for m in sent) + assert sent[-1].get("more_body", False) is False + + +@pytest.mark.asyncio +@pytest.mark.parametrize("status", [200, 400, 403]) +async def test_listing_transport_preserves_non_auth_responses(status: int) -> None: + from litellm.proxy._experimental.mcp_server.server import MCPAuthResponse + + send: Final = AsyncMock() + response: Final = MCPAuthResponse(send) + start: Final = { + "type": "http.response.start", + "status": status, + "headers": [(b"content-type", b"application/json")], + } + body: Final = { + "type": "http.response.body", + "body": b'{"jsonrpc":"2.0","id":1,"result":{"tools":[]}}', + "more_body": False, + } + await response.send(start) + await response.send(body) + assert tuple(call.args[0] for call in send.await_args_list) == (start, body) + + +@pytest.mark.asyncio +async def test_protocol_listing_does_not_report_success_when_every_server_requires_auth() -> None: + from unittest.mock import patch + from mcp.types import ListToolsRequest + from litellm.proxy._experimental.mcp_server import operations + from litellm.proxy._experimental.mcp_server.contracts import OperationContext + from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing, ServerListFault + from litellm.proxy._types import UserAPIKeyAuth + + context: Final = OperationContext(_caller=UserAPIKeyAuth(user_id="reader")) + listing: Final = AggregateToolListing([], {"github": ServerListFault(tag="auth_required", status_code=401)}) + with patch.object(operations, "_list_mcp_tools", AsyncMock(return_value=listing)): + with pytest.raises(MCPUpstreamAuthError) as caught: + await operations.GatewayOperations().execute(ListToolsRequest(), context) + assert caught.value.status_code == 401 + assert caught.value.server_name == "github" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("method", ("tools/list", "prompts/list", "resources/list", "resources/templates/list")) +@pytest.mark.parametrize("protocol", ("2025-06-18", "2025-11-25")) +@pytest.mark.parametrize("json_response", (False, True)) +@pytest.mark.parametrize("stateful", (False, True)) +@pytest.mark.parametrize("path", ("/mcp", "/github/mcp")) +async def test_streamable_http_listing_returns_late_oauth_challenge( + monkeypatch: pytest.MonkeyPatch, stateful: bool, path: str, json_response: bool, protocol: str, method: str +) -> None: + from litellm.proxy._experimental.mcp_server import server + from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing, ServerListFault + from litellm.proxy._types import UserAPIKeyAuth + + auth: Final = UserAPIKeyAuth(api_key="test-owner", user_id="test-user") + monkeypatch.setattr( + server, "extract_mcp_auth_context", AsyncMock(return_value=(auth, None, None, None, None, None)) + ) + monkeypatch.setattr(server, "_raise_preemptive_401_for_unauthenticated_servers", AsyncMock()) + monkeypatch.setattr(server, "_check_passthrough_upstream_auth", AsyncMock()) + listing: Final = AsyncMock( + return_value=AggregateToolListing( + [], + { + "github": ServerListFault( + tag="auth_required", + status_code=401, + www_authenticate='Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp/github"', + ) + }, + ) + ) + if method != "tools/list": + listing.side_effect = MCPUpstreamAuthError( + 401, 'Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp/github"', "github" + ) + list_function: Final = { + "tools/list": "_list_mcp_tools", + "prompts/list": "_list_mcp_prompts", + "resources/list": "_list_mcp_resources", + "resources/templates/list": "_list_mcp_resource_templates", + }[method] + monkeypatch.setattr(server.operations, list_function, listing) + monkeypatch.setattr(server.operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[])) + monkeypatch.setattr(server.operations, "_raise_if_initialize_grants_no_mcp_servers", AsyncMock()) + from mcp.server.streamable_http_manager import StreamableHTTPSessionManager + + monkeypatch.setattr( + server, + "session_manager_stateless", + StreamableHTTPSessionManager(app=server.server, stateless=True, json_response=json_response), + ) + monkeypatch.setattr( + server, + "session_manager_stateful", + StreamableHTTPSessionManager(app=server.server, stateless=False, json_response=json_response), + ) + await server.initialize_session_managers() + try: + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=server.app), base_url="http://gateway" + ) as client: + headers: Final = {"accept": "application/json, text/event-stream", "mcp-protocol-version": protocol} + if stateful: + initialized: Final = await client.post( + path, + headers=headers, + json={ + "jsonrpc": "2.0", + "id": 0, + "method": "initialize", + "params": { + "protocolVersion": protocol, + "capabilities": {}, + "clientInfo": {"name": "test", "version": "1"}, + }, + }, + ) + assert initialized.status_code == 200, initialized.text + client.headers["mcp-session-id"] = initialized.headers["mcp-session-id"] + notification: Final = await client.post( + path, headers=headers, json={"jsonrpc": "2.0", "method": "notifications/initialized"} + ) + assert notification.status_code == 202 + malformed: Final = await client.post( + path, headers={**headers, "content-type": "application/json"}, content=b"{" + ) + assert malformed.status_code == 400 + listing.assert_not_awaited() + response: Final = await client.post( + path, headers=headers, json={"jsonrpc": "2.0", "id": 1, "method": method} + ) + assert response.status_code == 401, response.text + assert ( + response.headers["www-authenticate"] + == 'Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp/github"' + ) + listing.assert_awaited_once() + finally: + await server.shutdown_session_managers() + + +@pytest.mark.asyncio +async def test_late_auth_failure_does_not_rewrite_committed_stream() -> None: + from fastapi import HTTPException + from litellm.proxy._experimental.mcp_server.server import MCPAuthResponse + + send: Final = AsyncMock() + response: Final = MCPAuthResponse(send) + start: Final = {"type": "http.response.start", "status": 200, "headers": [(b"content-type", b"text/event-stream")]} + progress: Final = { + "type": "http.response.body", + "body": b'data: {"method":"notifications/progress"}\n\n', + "more_body": True, + } + error: Final = { + "type": "http.response.body", + "body": b'data: {"error":{"code":-32600,"message":"Upstream authorization failed (HTTP 401)"}}\n\n', + "more_body": False, + } + await response.send(start) + await response.send(progress) + response.challenge = HTTPException(401, headers={"WWW-Authenticate": "Bearer"}) + await response.send(error) + assert tuple(call.args[0] for call in send.await_args_list) == (start, progress, error) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("kind", ("prompts", "resources", "resource_templates")) +async def test_optional_listing_propagates_auth_without_discarding_healthy_servers( + kind: str, monkeypatch: pytest.MonkeyPatch +) -> None: + from litellm.proxy._experimental.mcp_server import operations + + blocked: Final = _http_server("blocked", "blocked") + healthy: Final = _http_server("healthy", "healthy") + fetch: Final = AsyncMock(side_effect=MCPUpstreamAuthError(401, "Bearer", "blocked")) + manager: Final = MagicMock() + setattr(manager, f"get_{kind}_from_server", fetch) + monkeypatch.setattr(operations, "global_mcp_server_manager", manager) + monkeypatch.setattr(operations, "_prepare_mcp_server_headers", MagicMock(return_value=(None, None))) + monkeypatch.setattr(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[blocked])) + listing: Final = getattr(operations, f"_list_mcp_{kind}") + with pytest.raises(MCPUpstreamAuthError) as failure: + await listing() + assert failure.value.www_authenticate == "Bearer" + fetch.assert_awaited_once() + fetch.reset_mock(side_effect=True) + fetch.side_effect = [MCPUpstreamAuthError(401, "Bearer", "blocked"), []] + monkeypatch.setattr(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[blocked, healthy])) + assert await listing() == [] + assert fetch.await_count == 2 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "operation", + ( + "get_prompts_from_server", + "get_resources_from_server", + "get_resource_templates_from_server", + "get_prompt_from_server", + "read_resource_from_server", + ), +) +@pytest.mark.parametrize("status", (401, 403)) +async def test_manager_preserves_auth_failures_for_prompts_and_resources( + operation: str, status: int, monkeypatch: pytest.MonkeyPatch +) -> None: + from fastapi import HTTPException + from pydantic import AnyUrl + + manager: Final = MCPServerManager() + upstream: Final = _http_server("upstream", "upstream") + create: Final = AsyncMock(side_effect=HTTPException(status, headers={"WWW-Authenticate": "Bearer"})) + monkeypatch.setattr(manager, "_create_mcp_client", create) + kwargs: Final = ( + {"prompt_name": "example"} + if operation == "get_prompt_from_server" + else {"url": AnyUrl("https://example.com/resource")} + if operation == "read_resource_from_server" + else {} + ) + with pytest.raises(MCPUpstreamAuthError) as failure: + await getattr(manager, operation)(server=upstream, user_api_key_auth=None, **kwargs) + assert failure.value.status_code == status + assert failure.value.www_authenticate == "Bearer" + assert failure.value.server_name == "upstream" + create.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_optional_listing_preserves_cancellation() -> None: + import asyncio + from litellm.proxy._experimental.mcp_server.operations import _collect_mcp_listing + + fetch: Final = AsyncMock(side_effect=asyncio.CancelledError()) + with pytest.raises(asyncio.CancelledError): + await _collect_mcp_listing((_http_server("upstream", "upstream"),), fetch) + fetch.assert_awaited_once() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ("get_prompt_from_server", "read_resource_from_server")) +async def test_prompt_and_resource_calls_preserve_static_headers_and_non_auth_failures( + operation: str, monkeypatch: pytest.MonkeyPatch +) -> None: + from pydantic import AnyUrl + + manager: Final = MCPServerManager() + upstream: Final = _http_server("upstream", "upstream", static_headers={"x-upstream": "configured"}) + failure: Final = RuntimeError("Upstream unavailable") + create: Final = AsyncMock(side_effect=failure) + monkeypatch.setattr(manager, "_create_mcp_client", create) + kwargs: Final = ( + {"prompt_name": "example"} + if operation == "get_prompt_from_server" + else {"url": AnyUrl("https://example.com/resource")} + ) + with pytest.raises(RuntimeError) as caught: + await getattr(manager, operation)(server=upstream, user_api_key_auth=None, **kwargs) + assert caught.value is failure + assert create.await_args.kwargs["extra_headers"] == {"x-upstream": "configured"} + create.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_transport_preserves_sse_priming_event_on_success() -> None: + from litellm.proxy._experimental.mcp_server.server import MCPAuthResponse + + send: Final = AsyncMock() + response: Final = MCPAuthResponse(send) + start: Final = {"type": "http.response.start", "status": 200, "headers": [(b"content-type", b"text/event-stream")]} + priming: Final = {"type": "http.response.body", "body": b"id: resume-token\r\ndata: \r\n\r\n", "more_body": True} + tools: Final = { + "type": "http.response.body", + "body": b'data: {"jsonrpc":"2.0","id":1,"result":{"tools":[]}}\n\n', + "more_body": True, + } + await response.send(start) + await response.send(priming) + send.assert_not_awaited() + await response.send(tools) + assert tuple(call.args[0] for call in send.await_args_list) == (start, priming, tools) + + +@pytest.mark.asyncio +async def test_tool_call_preserves_resolver_http_challenge(monkeypatch: pytest.MonkeyPatch) -> None: + from fastapi import HTTPException + from mcp.types import CallToolRequest, CallToolRequestParams + from litellm.proxy._experimental.mcp_server import operations + from litellm.proxy._experimental.mcp_server.contracts import OperationContext + + challenge: Final = HTTPException(401, "Unauthorized", headers={"WWW-Authenticate": "Bearer"}) + call: Final = AsyncMock(side_effect=challenge) + monkeypatch.setattr(operations, "call_mcp_tool", call) + with pytest.raises(HTTPException) as failure: + await operations.GatewayOperations().execute( + CallToolRequest(params=CallToolRequestParams(name="upstream-tool", arguments={})), + OperationContext(_caller=None), + ) + assert failure.value is challenge + call.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_tool_handler_preserves_unrelated_protocol_errors( + _mcp_request_ctx, monkeypatch: pytest.MonkeyPatch +) -> None: + from mcp import MCPError + from mcp.types import CallToolRequestParams + from litellm.proxy._experimental.mcp_server import server + + failure: Final = MCPError(code=-32602, message="Invalid tool parameters") + execute: Final = AsyncMock(side_effect=failure) + gateway: Final = MagicMock() + gateway.execute = execute + monkeypatch.setattr(server.operations, "GatewayOperations", MagicMock(return_value=gateway)) + monkeypatch.setattr( + server, "get_or_extract_auth_context", AsyncMock(return_value=(None, None, None, None, None, None, None)) + ) + with pytest.raises(MCPError) as caught: + await server.mcp_server_tool_call(_mcp_request_ctx(), CallToolRequestParams(name="example", arguments={})) + assert caught.value is failure + execute.assert_awaited_once() 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 ab00ec4da1e..da11b290a1d 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 @@ -1969,7 +1969,8 @@ async def test_mcp_routing_initialize_to_stateful_no_session_to_stateless(_mcp_r assert stateful_handle.await_count == (1 if stateful else 0) assert stateless_handle.await_count == (0 if stateful else 1) - observe_start.assert_awaited_once_with(0 if debug and method == "POST" else 1) + deferred: Final = method == "POST" and (debug or (bool(request_body) and not stateful)) + observe_start.assert_awaited_once_with(0 if deferred else 1) assert send.await_count == 2 assert send.call_args_list[0].args[0]["status"] == 200 assert send.call_args_list[1].args[0] == body diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index bb4e9e0e0f0..9f713c4175b 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -2047,9 +2047,8 @@ class TestMCPServerManager: assert mock_log.warning.called @pytest.mark.asyncio - async def test_call_non_passthrough_does_not_opt_into_raise_on_error(self): - """Non-client-forwarded auth types keep the default call_tool masking (raise_on_error stays - off), so this relay is scoped to the pass-through modes and cannot regress api_key/OBO calls.""" + async def test_call_static_auth_preserves_success_with_transport_errors_enabled(self): + """A static credential uses the same transport-auth error channel while preserving tool results.""" server = MCPServer( server_id="ak-call", name="ak-call-server", @@ -2076,7 +2075,7 @@ class TestMCPServerManager: ) assert result.is_error is False - assert mock_client.call_tool.call_args.kwargs.get("raise_on_error") is not True + assert mock_client.call_tool.call_args.kwargs.get("raise_on_error") is True def _token_exchange_server(self, server_id: str) -> "MCPServer": return MCPServer( @@ -6909,7 +6908,7 @@ class TestMCPServerManager: # Create mock client that tracks call_tool usage mock_client = AsyncMock() - async def mock_call_tool(params, host_progress_callback=None, allow_input_required=False): + async def mock_call_tool(params, host_progress_callback=None, allow_input_required=False, raise_on_error=False): # Return a mock CallToolResult result = MagicMock(spec=CallToolResult) result.content = [{"type": "text", "text": "Tool executed successfully"}] @@ -14141,6 +14140,7 @@ async def test_discovery_cache_bounds_detached_fetches_without_dropping_results( @pytest.mark.asyncio async def test_discovery_cache_tracks_resolved_credentials_across_workers() -> None: + from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError import respx from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import StaticHeaderAuth from litellm.proxy._experimental.mcp_server.outbound_credentials.resolver import UpstreamCredentialProvider @@ -14198,7 +14198,9 @@ async def test_discovery_cache_tracks_resolved_credentials_across_workers() -> N assert upstream.initializes == 4 source.token = None for manager in managers: - assert await manager.get_prompts_from_server(server, user) == [] + with pytest.raises(MCPUpstreamAuthError) as failure: + await manager.get_prompts_from_server(server, user) + assert failure.value.status_code == 401 assert upstream.initializes == 4 From cb6050b5215be82d2f0379d861d65da46f528ed1 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Mon, 28 Sep 2026 12:18:26 -0700 Subject: [PATCH 2/9] fix(mcp): challenge missing upstream OAuth during initialization --- .../proxy/_experimental/mcp_server/server.py | 71 +++++++++++--- .../test_mcp_oauth_passthrough_tools.py | 68 +++++++++++++ .../mcp_server/test_mcp_server.py | 98 +++++++++++++++++++ 3 files changed, 226 insertions(+), 11 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 926aead9b7a..ea4799d9712 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -34,6 +34,7 @@ from litellm.llms.custom_httpx.http_handler import ( ) from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, + _gateway_dcr_challenge, _is_mcp_admitted_user_subject, ) from litellm.proxy._experimental.mcp_server.client_allowlist import ( @@ -1698,7 +1699,42 @@ if MCP_AVAILABLE: excludes a passthrough server is not pushed into an OAuth flow for a server it will be 403'd on immediately after authentication. """ - for server_name in mcp_servers or []: + if mcp_servers is None: + allowed: Final = await operations._get_allowed_mcp_servers( + user_api_key_auth=user_api_key_auth, mcp_servers=None, client_ip=client_ip + ) + eligible: Final = tuple( + server for server in allowed if allowed_server_ids is None or server.server_id in allowed_server_ids + ) + results: Final = await asyncio.gather( + *( + _raise_preemptive_401_for_unauthenticated_servers( + scope=scope, + mcp_servers=[server.alias or server.server_name or server.name], + oauth2_headers=oauth2_headers, + mcp_server_auth_headers=mcp_server_auth_headers, + user_api_key_auth=user_api_key_auth, + client_ip=client_ip, + allowed_server_ids=allowed_server_ids, + raw_headers=raw_headers, + ) + for server in eligible + ), + return_exceptions=True, + ) + for result in results: + if isinstance(result, asyncio.CancelledError): + raise result + 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 + ) + first: Final = results[0] + if isinstance(first, HTTPException): + raise first + return + for server_name in mcp_servers: server = operations.global_mcp_server_manager.get_mcp_server_by_name(server_name, client_ip=client_ip) if server is not None and allowed_server_ids is not None and server.server_id not in allowed_server_ids: # Caller's narrowed scope excludes this server โ€” skip the @@ -2114,16 +2150,17 @@ if MCP_AVAILABLE: # from the fully-authorized server set: a passthrough server that # the active toolset excludes should not trigger an OAuth flow # for a server the caller will be 403'd on after authentication. - await _raise_preemptive_401_for_unauthenticated_servers( - scope=scope, - mcp_servers=mcp_servers, - oauth2_headers=oauth2_headers, - mcp_server_auth_headers=mcp_server_auth_headers, - user_api_key_auth=user_api_key_auth, - client_ip=_client_ip, - allowed_server_ids=toolset_allowed_server_ids, - raw_headers=raw_headers, - ) + if mcp_servers is not None: + await _raise_preemptive_401_for_unauthenticated_servers( + scope=scope, + mcp_servers=mcp_servers, + oauth2_headers=oauth2_headers, + mcp_server_auth_headers=mcp_server_auth_headers, + user_api_key_auth=user_api_key_auth, + client_ip=_client_ip, + allowed_server_ids=toolset_allowed_server_ids, + raw_headers=raw_headers, + ) # Pre-flight auth check for pass-through servers. Must run after # toolset scoping so the probe list is derived from the fully-authorized @@ -2200,6 +2237,18 @@ if MCP_AVAILABLE: consumed_messages, body = await _read_request_body_for_routing(receive) is_initialize = _is_initialize_request(body) + if is_initialize and mcp_servers is None: + await _raise_preemptive_401_for_unauthenticated_servers( + scope=scope, + mcp_servers=None, + oauth2_headers=oauth2_headers, + mcp_server_auth_headers=mcp_server_auth_headers, + user_api_key_auth=user_api_key_auth, + client_ip=_client_ip, + allowed_server_ids=toolset_allowed_server_ids, + raw_headers=raw_headers, + ) + use_stateful: Final = bool(session_id or is_initialize) target_manager: Final = session_manager_stateful if use_stateful else session_manager_stateless 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 15aa089637a..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 @@ -854,3 +854,71 @@ async def test_tool_handler_preserves_unrelated_protocol_errors( await server.mcp_server_tool_call(_mcp_request_ctx(), CallToolRequestParams(name="example", arguments={})) assert caught.value is failure execute.assert_awaited_once() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("path", ("/mcp", "/github/mcp")) +@pytest.mark.parametrize("has_token", (False, True)) +@pytest.mark.parametrize("healthy_companion", (False, True)) +async def test_initialize_challenges_missing_upstream_credentials_before_creating_session( + monkeypatch: pytest.MonkeyPatch, path: str, has_token: bool, healthy_companion: bool +) -> None: + from mcp.server.streamable_http_manager import StreamableHTTPSessionManager + from litellm.proxy._experimental.mcp_server import server + from litellm.proxy._types import UserAPIKeyAuth + + github: Final = MCPServer( + server_id="github-id", name="github", alias="github", server_name="github", + url="https://github.example/mcp", transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", + authorization_url="https://github.example/authorize", token_url="https://github.example/token", + client_id="registered-client", + ) + auth: Final = UserAPIKeyAuth(api_key="test-owner", user_id="test-user") + selected: Final = ["github"] if path != "/mcp" else None + monkeypatch.setattr( + server, "extract_mcp_auth_context", AsyncMock(return_value=(auth, None, selected, None, None, None)) + ) + manager: Final = server.operations.global_mcp_server_manager + public: Final = MCPServer( + server_id="public-id", name="public", alias="public", server_name="public", + url="https://public.example/mcp", transport=MCPTransport.http, auth_type=MCPAuth.none, + ) + eligible: Final = [github, public] if healthy_companion and selected is None else [github] + monkeypatch.setattr(manager, "get_mcp_server_by_name", lambda name, **kwargs: next(s for s in eligible if s.alias == name)) + monkeypatch.setattr(manager, "has_user_oauth_token", AsyncMock(return_value=has_token)) + monkeypatch.setattr(manager, "_ensure_upstream_initialize_instructions_cached", AsyncMock()) + monkeypatch.setattr(server.operations, "_get_allowed_mcp_servers", AsyncMock(return_value=eligible)) + monkeypatch.setattr(server, "_check_passthrough_upstream_auth", AsyncMock()) + monkeypatch.setattr( + server, "session_manager_stateful", + StreamableHTTPSessionManager(app=server.server, stateless=False, json_response=True), + ) + monkeypatch.setattr( + server, "session_manager_stateless", + StreamableHTTPSessionManager(app=server.server, stateless=True, json_response=True), + ) + await server.initialize_session_managers() + try: + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=server.app), base_url="http://gateway") as client: + response: Final = await client.post( + path, headers={"accept": "application/json, text/event-stream"}, + json={"jsonrpc": "2.0", "id": 1, "method": "initialize", "params": { + "protocolVersion": "2025-06-18", "capabilities": {}, + "clientInfo": {"name": "test-client", "version": "1"}, + }}, + ) + can_initialize: Final = has_token or (healthy_companion and selected is None) + assert response.status_code == (200 if can_initialize else 401), response.text + if can_initialize: + assert response.json()["result"]["serverInfo"]["name"] + assert response.headers["mcp-session-id"] + else: + assert response.headers["www-authenticate"].startswith("Bearer ") + if path == "/mcp": + assert response.headers["www-authenticate"] == ( + 'Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp"' + ) + assert "mcp-session-id" not in response.headers + finally: + await server.shutdown_session_managers() 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 da11b290a1d..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 @@ -10690,3 +10690,101 @@ async def test_discovery_adapter_preserves_authenticated_context(_mcp_request_ct context = dispatched.await_args.args[1] assert context.user_api_key_auth.user_id == "discover-caller" assert context.mcp_servers == ("allowed",) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "token_states, allowed_ids, expected_status", + ( + ((False,), None, 401), + ((False, False), None, 401), + ((False, True), None, None), + ((True, False), None, None), + ((False, False), {"server-1"}, 401), + ((False, True), {"server-1"}, None), + ((False,), set(), None), + ((), None, None), + ), +) +async def test_unified_preflight_challenges_only_when_all_authorized_servers_need_oauth( + monkeypatch: pytest.MonkeyPatch, + token_states: tuple[bool, ...], + allowed_ids: set[str] | None, + expected_status: int | None, +) -> None: + from litellm.proxy._experimental.mcp_server import server as server_module + + servers: Final = tuple( + _make_oauth2_server(f"server-{index}").model_copy(update={"server_id": f"server-{index}"}) + for index in range(len(token_states)) + ) + auth: Final = UserAPIKeyAuth(api_key="test-key", user_id="reader") + lookup: Final = AsyncMock(return_value=servers) + monkeypatch.setattr(mcp_operations, "_get_allowed_mcp_servers", lookup) + manager: Final = mcp_operations.global_mcp_server_manager + monkeypatch.setattr(manager, "get_mcp_server_by_name", lambda name, **kwargs: next(s for s in servers if s.alias == name)) + tokens: Final = AsyncMock(side_effect=lambda s, user: token_states[int(s.server_id.rsplit("-", 1)[1])]) + monkeypatch.setattr(manager, "has_user_oauth_token", tokens) + scope: Final[Scope] = {"type": "http", "method": "POST", "path": "/mcp", "headers": [(b"host", b"gateway")]} + request: Final = server_module._raise_preemptive_401_for_unauthenticated_servers( + scope=scope, mcp_servers=None, oauth2_headers=None, mcp_server_auth_headers=None, + user_api_key_auth=auth, client_ip="127.0.0.1", allowed_server_ids=allowed_ids, + ) + if expected_status is not None: + with pytest.raises(HTTPException) as caught: + await request + assert caught.value.status_code == expected_status + assert "www-authenticate" in {key.lower() for key in (caught.value.headers or {})} + else: + await request + lookup.assert_awaited_once_with(user_api_key_auth=auth, mcp_servers=None, client_ip="127.0.0.1") + assert {call.args[0].server_id for call in tokens.await_args_list} == { + s.server_id for s in servers if allowed_ids is None or s.server_id in allowed_ids + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize("failure", (HTTPException(status_code=503, detail="unavailable"), asyncio.CancelledError())) +async def test_unified_preflight_does_not_misclassify_discovery_failure_as_oauth( + monkeypatch: pytest.MonkeyPatch, failure: BaseException +) -> None: + from litellm.proxy._experimental.mcp_server import server as server_module + + upstream: Final = _make_oauth2_server("unavailable") + monkeypatch.setattr(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[upstream])) + manager: Final = mcp_operations.global_mcp_server_manager + monkeypatch.setattr(manager, "get_mcp_server_by_name", lambda *args, **kwargs: upstream) + discovery: Final = AsyncMock(side_effect=failure) + monkeypatch.setattr(manager, "ensure_oauth_metadata_discovered", discovery) + request: Final = server_module._raise_preemptive_401_for_unauthenticated_servers( + scope={"type": "http", "path": "/mcp", "method": "POST", "headers": []}, + mcp_servers=None, oauth2_headers=None, mcp_server_auth_headers=None, + user_api_key_auth=UserAPIKeyAuth(user_id="reader"), client_ip=None, + ) + if isinstance(failure, asyncio.CancelledError): + with pytest.raises(asyncio.CancelledError): + await request + else: + await request + discovery.assert_awaited_once_with(upstream) + + +@pytest.mark.asyncio +async def test_unified_preflight_preserves_delegated_oauth_challenge(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy._experimental.mcp_server import server as server_module + + upstream: Final = _make_oauth2_server("delegated", delegate_auth_to_upstream=True) + monkeypatch.setattr(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[upstream])) + monkeypatch.setattr( + mcp_operations.global_mcp_server_manager, "get_mcp_server_by_name", lambda *args, **kwargs: upstream + ) + with pytest.raises(HTTPException) as caught: + await server_module._raise_preemptive_401_for_unauthenticated_servers( + scope={"type": "http", "path": "/mcp", "method": "POST", "headers": [(b"host", b"gateway")]}, + mcp_servers=None, oauth2_headers=None, mcp_server_auth_headers=None, + user_api_key_auth=UserAPIKeyAuth(user_id="reader"), client_ip=None, + ) + assert caught.value.status_code == 401 + assert (caught.value.headers or {})["www-authenticate"] == ( + 'Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp/delegated"' + ) From 3a8ed9773b94f4acbbd878e364955bbe473a5c3a Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Mon, 28 Sep 2026 12:57:00 -0700 Subject: [PATCH 3/9] fix(mcp): require challenged upstream consent before completing OAuth --- .../mcp_server/auth/user_api_key_auth_mcp.py | 4 +- .../mcp_server/gateway_dcr_flow.py | 52 ++++++++- .../proxy/_experimental/mcp_server/server.py | 9 +- .../mcp_server/test_gateway_dcr_flow.py | 110 ++++++++++++++++++ .../test_mcp_oauth_passthrough_tools.py | 5 +- .../mcp_server/test_mcp_server.py | 1 + 6 files changed, 175 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index a93ffaeac9f..237e38fa283 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -280,6 +280,7 @@ def _gateway_dcr_challenge( route: str, mcp_servers: list[str] | None, invalid_token: bool, + oauth_scope: str | None = None, ) -> HTTPException: """The RFC 9728 challenge pointing the client at the protected-resource metadata matching the scope it requested: the per-server document (same URL spelling the @@ -298,13 +299,14 @@ def _gateway_dcr_challenge( else f"{get_request_base_url(request)}/.well-known/oauth-protected-resource{well_known_root_suffix()}/mcp" ) error_attr: Final = 'error="invalid_token", ' if invalid_token else "" + scope_attr: Final = f', scope="{oauth_scope}"' if oauth_scope else "" return HTTPException( status_code=401, detail={ "error": "authentication_required", "message": "Authenticate with the gateway to use the MCP endpoint.", }, - headers={"WWW-Authenticate": f'Bearer {error_attr}resource_metadata="{resource_metadata_url}"'}, + headers={"WWW-Authenticate": f'Bearer {error_attr}resource_metadata="{resource_metadata_url}"{scope_attr}'}, ) diff --git a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py index e66504af47a..f4ec41cb6e0 100644 --- a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py +++ b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py @@ -108,6 +108,7 @@ handle carried in the connect-page URL (the same handle-plus-cookie pattern as t ``mcp_oauth_state_`` upstream relay, for the same reasons: replica-safe with no server-side session store, and the sealed value never appears in a URL).""" +UPSTREAM_AUTHORIZATION_SCOPE_PREFIX: Final = "litellm:mcp:connect:" CONNECT_FLOW_TTL_SECONDS: Final = 600 GATEWAY_AUTH_CODE_TTL_SECONDS: Final = 120 MANUAL_DELIVERY_AUTH_CODE_TTL_SECONDS: Final = 300 @@ -286,6 +287,13 @@ class GatewayDcrClient(BaseModel): iat: int +class _UpstreamAuthorizationRequirement(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid") + server_id: str = Field(min_length=1) + user_id: str | None = None + exp: int + + class _ConnectFlow(BaseModel): """One in-flight authorize: the SSO user it belongs to and the client parameters needed to mint the code at the finish step. Sealed into the per-flow cookie. ``jti`` @@ -301,6 +309,7 @@ class _ConnectFlow(BaseModel): jti: str = Field(min_length=1) exp: int resource_server_id: str | None = None + required_upstream_server_id: str | None = None audience: SessionAudience | None = None @@ -502,6 +511,17 @@ def resolve_scoped_resource_server(request: Request, resource: str | None) -> MC return server +def upstream_authorization_scope(server_id: str, user_id: str | None) -> str: + return _seal( + UPSTREAM_AUTHORIZATION_SCOPE_PREFIX, + _UpstreamAuthorizationRequirement( + server_id=server_id, + user_id=user_id, + exp=int(datetime.now(timezone.utc).timestamp()) + CONNECT_FLOW_TTL_SECONDS, + ), + ) + + def aggregate_authorize( request: Request, client_id: str, @@ -537,6 +557,30 @@ def aggregate_authorize( if session_user_id is None: return _login_redirect(base_url, request) scoped_server: Final = resolve_scoped_resource_server(request, resource) + requested: Final = tuple( + value + for value in request.query_params.get("scope", "").split() + if value.startswith(UPSTREAM_AUTHORIZATION_SCOPE_PREFIX) + ) + if len(requested) > 1: + return _oauth_error(400, "invalid_scope", "only one upstream authorization requirement is supported") + requirement: Final = ( + _open_sealed( + requested[0], + UPSTREAM_AUTHORIZATION_SCOPE_PREFIX, + _UpstreamAuthorizationRequirement, + "mcp_upstream_authorization", + ) + if requested + else None + ) + if requested and (requirement is None or datetime.now(timezone.utc).timestamp() >= requirement.exp): + return _oauth_error(400, "invalid_scope", "invalid or expired upstream authorization requirement; reconnect") + if requirement is not None: + if requirement.user_id is not None and requirement.user_id != session_user_id: + return _oauth_error(403, "access_denied", "sign in as the user that requested this MCP connection") + if scoped_server is not None and scoped_server.server_id != requirement.server_id: + return _oauth_error(400, "invalid_scope", "upstream authorization requirement does not match the resource") handle: Final = secrets.token_urlsafe(24) flow: Final = _new_connect_flow( session_user_id=session_user_id, @@ -546,6 +590,7 @@ def aggregate_authorize( code_challenge=code_challenge or "", resource_server_id=scoped_server.server_id if scoped_server is not None else None, audience=None, + required_upstream_server_id=requirement.server_id if requirement is not None else None, ) connect_url: Final = _append_query_params(f"{base_url}/ui/connect", (("connect_flow", handle),)) response: Final = RedirectResponse(connect_url, status_code=303) @@ -705,6 +750,7 @@ def _new_connect_flow( code_challenge: str, resource_server_id: str | None, audience: SessionAudience | None, + required_upstream_server_id: str | None = None, ) -> _ConnectFlow: now: Final = datetime.now(timezone.utc) return _ConnectFlow( @@ -716,6 +762,7 @@ def _new_connect_flow( jti=secrets.token_urlsafe(24), exp=int(now.timestamp()) + CONNECT_FLOW_TTL_SECONDS, resource_server_id=resource_server_id, + required_upstream_server_id=required_upstream_server_id, audience=audience, ) @@ -773,14 +820,15 @@ def _open_flow_for( async def _flow_target( flow: _ConnectFlow, lookup_server_reachability: LookupServerReachability ) -> tuple[Literal["unscoped", "interactive", "m2m", "stale"], MCPServer | None]: - if flow.resource_server_id is None: + target_id: Final = flow.required_upstream_server_id or flow.resource_server_id + if target_id is None: return "unscoped", None from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # import cycle MCPServerManager, global_mcp_server_manager, ) - server: Final = global_mcp_server_manager.get_mcp_server_by_id(flow.resource_server_id) + server: Final = global_mcp_server_manager.get_mcp_server_by_id(target_id) if ( server is None or not (server.is_gateway_managed_oauth2 or server.advertises_gateway_authorization_server) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index ea4799d9712..f8842fa02f6 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -48,6 +48,7 @@ from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( from litellm.proxy._experimental.mcp_server.exceptions import ( MCPUpstreamAuthError, ) +from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import upstream_authorization_scope from litellm.proxy._experimental.mcp_server.mcp_context import ( _mcp_active_toolset_id, _mcp_gateway_initialize_instructions, @@ -1728,7 +1729,13 @@ if MCP_AVAILABLE: if results and all(isinstance(result, HTTPException) and result.status_code == 401 for result in results): if all(server.is_gateway_managed_oauth2 for server in eligible): raise _gateway_dcr_challenge( - StarletteRequest(scope), get_route_relative_request_path(scope), None, invalid_token=False + StarletteRequest(scope), + get_route_relative_request_path(scope), + None, + invalid_token=False, + oauth_scope=upstream_authorization_scope( + eligible[0].server_id, user_api_key_auth.user_id if user_api_key_auth is not None else None + ), ) first: Final = results[0] if isinstance(first, HTTPException): diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py index 2943ff4b74a..40ec75329a6 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py @@ -2354,3 +2354,113 @@ async def test_token_exchange_relays_a_mint_refusal(failure, status, error): response = await _exchange_native(client_id, _Minter(failure), _Exchanger()) assert response.status_code == status assert json.loads(response.body)["error"] == error + + +@pytest.mark.asyncio +async def test_unified_challenge_requires_upstream_consent_without_narrowing_gateway_token(monkeypatch) -> None: + from unittest.mock import AsyncMock + from urllib.parse import urlencode + from fastapi import HTTPException + from litellm.proxy._experimental.mcp_server import operations, server + from litellm.proxy._types import UserAPIKeyAuth + + github = _scoped_mcp_server(oauth2_flow="authorization_code") + manager = operations.global_mcp_server_manager + monkeypatch.setattr(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[github])) + monkeypatch.setattr(manager, "get_mcp_server_by_name", lambda *args, **kwargs: github) + monkeypatch.setattr(manager, "ensure_oauth_metadata_discovered", AsyncMock(return_value=github)) + monkeypatch.setattr(manager, "has_user_oauth_token", AsyncMock(return_value=False)) + with pytest.raises(HTTPException) as challenged: + await server._raise_preemptive_401_for_unauthenticated_servers( + scope=_request("/mcp").scope, mcp_servers=None, oauth2_headers=None, + mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(user_id="u1", api_key="test-key"), + client_ip=None, + ) + headers = {k.lower(): v for k, v in (challenged.value.headers or {}).items()} + requested = re.search(r'scope="([^"]+)"', headers["www-authenticate"]) + assert requested is not None, "The challenge must carry the upstream authorization requirement" + client_id = (await _register([REDIRECT_URI]))["client_id"] + response = aggregate_authorize( + request=_request("/authorize/mcp-session", query=urlencode({"scope": requested.group(1)})), + client_id=client_id, redirect_uri=REDIRECT_URI, state="client-state", code_challenge=CODE_CHALLENGE, + code_challenge_method="S256", response_type="code", session_user_id="u1", + resource="https://llm.example.com/mcp", + ) + assert response.status_code == 303 + described = await _describe_page(response, scoped_server=github, vendor=_VendorCredential("absent")) + assert json.loads(described.body) == { + "state": "interactive", "client_origin": "https://claude.ai", + "server_id": "github-id", "server_name": "github", "connected": False, + } + cache = DualCache() + premature = await _complete_page(response, scoped_server=github, vendor=_VendorCredential("absent"), cache=cache) + assert premature.status_code == 400 + assert "location" not in premature.headers + vendor = _VendorCredential("present") + completed = await _complete_page(response, scoped_server=github, vendor=vendor, cache=cache) + assert completed.status_code == 303 + assert vendor.calls == [("u1", "github-id")] + code = parse_qs(urlparse(completed.headers["location"]).query)["code"][0] + tokens = await _redeem(code, client_id, resource="https://llm.example.com/mcp") + assert tokens.status_code == 200 + assert _opened_principal(json.loads(tokens.body)).resource_server_id is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("invalid", ("tampered", "expired", "duplicate", "other_user", "other_resource")) +async def test_upstream_authorization_requirement_rejects_invalid_binding(invalid: str, monkeypatch) -> None: + from urllib.parse import urlencode + from litellm.proxy._experimental.mcp_server import gateway_dcr_flow as flow + from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager + + hint = flow.upstream_authorization_scope("github-id", "u2" if invalid == "other_user" else "u1") + if invalid == "tampered": + hint = flow.UPSTREAM_AUTHORIZATION_SCOPE_PREFIX + "invalid-ciphertext" + elif invalid == "expired": + hint = flow._seal(flow.UPSTREAM_AUTHORIZATION_SCOPE_PREFIX, flow._UpstreamAuthorizationRequirement( + server_id="github-id", user_id="u1", exp=int(datetime.now(timezone.utc).timestamp()) - 1, + )) + elif invalid == "duplicate": + hint = f"{hint} {hint}" + monkeypatch.setattr(global_mcp_server_manager, "get_mcp_server_by_name", lambda *args, **kwargs: _scoped_mcp_server("other")) + client_id = (await _register([REDIRECT_URI]))["client_id"] + response = aggregate_authorize( + request=_request(query=urlencode({"scope": hint})), + client_id=client_id, redirect_uri=REDIRECT_URI, state="client-state", code_challenge=CODE_CHALLENGE, + code_challenge_method="S256", response_type="code", session_user_id="u1", + resource="https://llm.example.com/mcp/other" if invalid == "other_resource" else "https://llm.example.com/mcp", + ) + assert response.status_code == (403 if invalid == "other_user" else 400) + assert json.loads(response.body)["error"] == ("access_denied" if invalid == "other_user" else "invalid_scope") + assert "location" not in response.headers + assert "set-cookie" not in response.headers + + +@pytest.mark.asyncio +@pytest.mark.parametrize("condition", ("deleted", "revoked", "vault_unavailable", "cancelled")) +async def test_required_upstream_completion_preserves_failure_and_cancellation_guards(condition: str) -> None: + from urllib.parse import urlencode + from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import upstream_authorization_scope + + client_id = (await _register([REDIRECT_URI]))["client_id"] + hint = upstream_authorization_scope("github-id", "u1") + response = aggregate_authorize( + request=_request(query=urlencode({"scope": hint})), client_id=client_id, redirect_uri=REDIRECT_URI, + state="client-state", code_challenge=CODE_CHALLENGE, code_challenge_method="S256", + response_type="code", session_user_id="u1", resource="https://llm.example.com/mcp", + ) + assert response.status_code == 303 + vendor = _VendorCredential("unavailable" if condition == "vault_unavailable" else "absent") + completed = await _complete_page( + response, scoped_server=None if condition == "deleted" else _scoped_mcp_server(), + reachable=_ServerReachability(condition != "revoked"), vendor=vendor, + decision="deny" if condition == "cancelled" else None, + ) + if condition == "cancelled": + assert completed.status_code == 303 + assert parse_qs(urlparse(completed.headers["location"]).query) == {"error": ["access_denied"], "state": ["client-state"]} + assert vendor.calls == [] + else: + assert completed.status_code == (503 if condition == "vault_unavailable" else 400) + assert "location" not in completed.headers + assert vendor.calls == ([("u1", "github-id")] if condition == "vault_unavailable" else []) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py index a7c7e1853a0..785bcc8ab70 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py @@ -867,6 +867,7 @@ async def test_initialize_challenges_missing_upstream_credentials_before_creatin from litellm.proxy._experimental.mcp_server import server from litellm.proxy._types import UserAPIKeyAuth + monkeypatch.setenv("LITELLM_SALT_KEY", "test-mcp-oauth-signing") github: Final = MCPServer( server_id="github-id", name="github", alias="github", server_name="github", url="https://github.example/mcp", transport=MCPTransport.http, @@ -916,8 +917,8 @@ async def test_initialize_challenges_missing_upstream_credentials_before_creatin else: assert response.headers["www-authenticate"].startswith("Bearer ") if path == "/mcp": - assert response.headers["www-authenticate"] == ( - 'Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp"' + assert response.headers["www-authenticate"].startswith( + 'Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp", scope="litellm:mcp:connect:' ) assert "mcp-session-id" not in response.headers finally: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 0ba8b3f2a40..5777caf63de 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -10714,6 +10714,7 @@ async def test_unified_preflight_challenges_only_when_all_authorized_servers_nee ) -> None: from litellm.proxy._experimental.mcp_server import server as server_module + monkeypatch.setenv("LITELLM_SALT_KEY", "test-mcp-oauth-signing") servers: Final = tuple( _make_oauth2_server(f"server-{index}").model_copy(update={"server_id": f"server-{index}"}) for index in range(len(token_states)) From 8c7a4372ee4ac4e5a37b3c16ab9e098f566bc784 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Mon, 28 Sep 2026 14:48:39 -0700 Subject: [PATCH 4/9] fix(mcp): validate selected upstream connections before completing OAuth --- .../mcp_server/auth/user_api_key_auth_mcp.py | 4 +- .../mcp_server/discoverable_endpoints.py | 4 +- .../mcp_server/gateway_dcr_flow.py | 88 +++---- .../proxy/_experimental/mcp_server/server.py | 9 +- litellm/proxy/_lazy_openapi_snapshot.json | 15 ++ .../mcp_server/test_gateway_dcr_flow.py | 224 ++++++++++-------- .../test_mcp_oauth_passthrough_tools.py | 5 +- .../mcp_server/test_mcp_server.py | 1 - .../chat/ConnectFlowBanner.test.tsx | 14 +- .../src/components/chat/ConnectFlowBanner.tsx | 35 ++- .../chat/ConnectFlowSurface.test.tsx | 16 +- .../components/chat/ConnectFlowSurface.tsx | 1 + ui/litellm-dashboard/src/lib/http/schema.d.ts | 2 + 13 files changed, 231 insertions(+), 187 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 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 && (