diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 7a0f59c3c2b..cd5fe6c58ac 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -6,7 +6,7 @@ import time from collections.abc import AsyncIterator, Callable, Mapping from contextlib import asynccontextmanager 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 @@ -2133,6 +2133,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 @@ -2150,6 +2151,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/faults/list_outcomes.py b/litellm/proxy/_experimental/mcp_server/faults/list_outcomes.py index b96a7a74e4a..a71e0d10ec9 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 any(isinstance(outcome, ServerListOk) for outcome in outcomes.values()): + 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 @@ -104,17 +132,22 @@ def raise_classified_list_failure( a classified fault. Every fetch site delegates here so the two channels cannot drift apart per call site. ``suppress_challenge`` is for dcr_bridge servers, whose upstream challenge points clients at the wrong protected-resource metadata and must never relay.""" - auth: Final = upstream_auth_challenge(exc) + auth: Final = upstream_auth_error(exc, server_name, suppress_challenge=suppress_challenge) if auth is not None: - status_code, challenge = auth - raise MCPUpstreamAuthError( - status_code=status_code, - www_authenticate=None if suppress_challenge else challenge, - server_name=server_name, - ) from exc + raise auth from exc raise MCPServerListError(classify_list_exception(exc), server_name) from exc +def upstream_auth_error( + exc: BaseException, server_name: str, *, suppress_challenge: bool = False +) -> MCPUpstreamAuthError | None: + auth: Final = upstream_auth_challenge(exc) + if auth is None: + return None + status_code, challenge = auth + return MCPUpstreamAuthError(status_code, None if suppress_challenge else challenge, server_name) + + def classify_list_exception(exc: BaseException) -> ServerListFault: """Classify a per-server listing failure into exactly one outcome. Total: an exception this function cannot recognize is the gateway's own fault (``internal``), never a re-raise.""" @@ -122,17 +155,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/gateway_dcr_flow.py b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py index e66504af47a..bbc40433516 100644 --- a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py +++ b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py @@ -851,6 +851,37 @@ 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_id(server_id) for server_id 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, @@ -861,6 +892,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. @@ -876,6 +908,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/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index c0792c32de2..3fdbad3177e 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -86,6 +86,7 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ( ServerListFault, raise_classified_list_failure, upstream_auth_challenge, + upstream_auth_error, ) from litellm.proxy._experimental.mcp_server.mcp_debug import describe_upstream_http_failure, record_auth_resolution from litellm.proxy._experimental.mcp_server.oauth2_token_cache import ( @@ -4664,7 +4665,7 @@ 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) - return [] + raise_classified_list_failure(error, server.name, suppress_challenge=server.is_dcr_bridge) async def get_resources_from_server( self, @@ -4710,7 +4711,7 @@ 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) - return [] + raise_classified_list_failure(error, server.name, suppress_challenge=server.is_dcr_bridge) async def get_resource_templates_from_server( self, @@ -4756,7 +4757,7 @@ 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) - return [] + raise_classified_list_failure(error, server.name, suppress_challenge=server.is_dcr_bridge) async def read_resource_from_server( self, @@ -4770,29 +4771,35 @@ 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_error(exc, server.name, suppress_challenge=server.is_dcr_bridge) + if auth_failure is not None: + raise auth_failure from exc + raise async def get_prompt_from_server( self, @@ -4807,33 +4814,39 @@ 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_error(exc, server.name, suppress_challenge=server.is_dcr_bridge) + if auth_failure is not None: + raise auth_failure from exc + raise @staticmethod def _is_same_authority_metadata_url(url: str, server_url: str) -> bool: @@ -6190,29 +6203,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, @@ -6238,7 +6232,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 8bb2772c760..944e88646c5 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 ( @@ -1238,6 +1240,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.server_id: 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, @@ -1247,32 +1269,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, @@ -1280,30 +1283,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( @@ -1315,19 +1307,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, @@ -1335,28 +1321,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( @@ -1368,19 +1345,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, @@ -1388,38 +1359,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( @@ -1535,6 +1487,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 @@ -1565,6 +1519,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) @@ -1597,6 +1553,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", @@ -2717,6 +2675,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) @@ -2724,6 +2685,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 @@ -2853,27 +2816,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( @@ -2918,6 +2870,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 @@ -2987,6 +2941,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 @@ -3027,6 +2983,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 2412e83b9d9..053c934db89 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 ( @@ -92,6 +93,61 @@ 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()) + ): + 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=[]) @@ -1580,6 +1677,151 @@ if MCP_AVAILABLE: ) return user_api_key_auth.model_copy(update={"object_permission": updated_op, "mcp_toolset_id": toolset_id}) + async def _server_auth_challenge( + configured_server: MCPServer, + server_name: str, + scope: Scope, + mcp_servers: list[str], + oauth2_headers: dict[str, str] | None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None, + user_api_key_auth: UserAPIKeyAuth | None, + client_ip: str | None, + raw_headers: Mapping[str, str] | None, + ) -> HTTPException | None: + try: + if configured_server.auth_type == MCPAuth.oauth2 and configured_server.oauth2_flow == "client_credentials": + return None + server: Final = await operations.global_mcp_server_manager.ensure_oauth_metadata_discovered( + configured_server + ) + if server.auth_type == MCPAuth.oauth2: + if MCPServerManager.effective_oauth2_flow(server) == "client_credentials": + return None + + if getattr(server, "delegate_auth_to_upstream", False) is not True: + if await operations.global_mcp_server_manager.has_user_oauth_token(server, user_api_key_auth): + return None + + if _is_mcp_admitted_user_subject(user_api_key_auth): + return HTTPException( + status_code=401, + detail="Unauthorized", + headers={ + "www-authenticate": get_passthrough_www_authenticate( + scope=scope, + server_name=server_name, + ) + }, + ) + + request: Final = StarletteRequest(scope) + base_url: Final = get_request_base_url(request) + _path: Final = get_route_relative_request_path(scope) + + as_metadata_root: Final = ( + f"{base_url}/.well-known/oauth-authorization-server{well_known_root_suffix()}" + ) + as_url: Final = ( + f"{as_metadata_root}/mcp/{server_name}" + if _path.startswith(f"/mcp/{server_name}") + else f"{as_metadata_root}/{server_name}" + ) + authorization_uri: Final = f'Bearer authorization_uri="{as_url}"' + + return HTTPException( + status_code=401, + detail="Unauthorized", + headers={"www-authenticate": authorization_uri}, + ) + + if not oauth2_headers: + return HTTPException( + status_code=401, + detail="Unauthorized", + headers={ + "www-authenticate": get_passthrough_www_authenticate(scope=scope, server_name=server_name) + }, + ) + return None + + if server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers: + from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 # lazy: adapter pulls MCP subgraph + raise_token_exchange_challenge, + ) + from litellm.proxy.middleware.per_request_root_path_middleware import ( # noqa: PLC0415 # lazy: middleware imports proxy utils + get_request_root_path, + ) + + raise_token_exchange_challenge(server, root_path=get_request_root_path()) + + if len(mcp_servers) == 1 and server.server_id in frozenset( + allowed.server_id + for allowed in await operations._get_allowed_mcp_servers( + user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip + ) + ): + await operations.global_mcp_server_manager.preflight_token_exchange( + server=server, + oauth2_headers=oauth2_headers, + user_api_key_auth=user_api_key_auth, + raw_headers=raw_headers, + ) + + if server.is_oauth_passthrough and not operations._client_has_passthrough_authorization( + server, oauth2_headers, mcp_server_auth_headers + ): + return HTTPException( + status_code=401, + detail="Unauthorized", + headers={ + "www-authenticate": get_passthrough_www_authenticate(scope=scope, server_name=server_name) + }, + ) + + if ( + server.is_oauth_delegate + and len(mcp_servers) == 1 + and _get_forwarded_auth_from_scope(scope) is None + and not operations._client_has_per_server_auth_header(server, mcp_server_auth_headers) + ): + return HTTPException( + status_code=401, + detail="Unauthorized", + headers={ + "www-authenticate": get_passthrough_www_authenticate(scope=scope, server_name=server_name) + }, + ) + + if ( + server.is_true_passthrough + and len(mcp_servers) == 1 + and not _scope_has_authorization_header(scope) + and not operations._client_has_per_server_auth_header(server, mcp_server_auth_headers) + ): + if server.is_dcr_bridge: + return HTTPException( + status_code=401, + detail="Unauthorized", + headers={ + "www-authenticate": get_passthrough_www_authenticate( + scope=scope, + server_name=server_name, + ) + }, + ) + upstream_status, upstream_www_authenticate = await _probe_upstream_auth(server.url or "", "") + if upstream_status == 401 and upstream_www_authenticate: + return HTTPException( + status_code=401, + detail="Unauthorized", + headers={"www-authenticate": upstream_www_authenticate}, + ) + except HTTPException as exc: + if exc.status_code != 401: + raise + return exc + return None + async def _raise_preemptive_401_for_unauthenticated_servers( scope: Scope, mcp_servers: list[str] | None, @@ -1601,208 +1843,71 @@ 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 []: - 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 - # preemptive challenge and let downstream authorization - # return 403. - continue - if server is not None and server.auth_type == MCPAuth.oauth2 and server.oauth2_flow == "client_credentials": - # Stamped M2M: the challenge decision below never reads discovered - # metadata, so deferred-discovery failures must not 503 this loop. - # Unstamped rows stay on the discover-first path because filling - # authorization_url/token_url can change their inferred flow. - continue - if server is not None: - server = await operations.global_mcp_server_manager.ensure_oauth_metadata_discovered(server) - if server and server.auth_type == MCPAuth.oauth2: - # The challenge decision is per oauth2 sub-mode, not per header: - # gateway-managed modes (M2M and interactive authorization_code) - # never receive a client-supplied upstream token, so a bearer in - # Authorization is a LiteLLM key (surfaced here as oauth2_headers) - # and must not suppress the challenge. Only the delegate mode - # treats a present bearer as the upstream token. The sub-mode is - # resolved the same way egress resolves it, via - # effective_oauth2_flow: an unstamped (null oauth2_flow) row with - # the M2M shape resolves to client_credentials, so the bare - # has_client_credentials column is never trusted here. - if MCPServerManager.effective_oauth2_flow(server) == "client_credentials": - # M2M: the gateway mints its own token at egress from the - # stored client credentials, so there is nothing to challenge. - continue - - if getattr(server, "delegate_auth_to_upstream", False) is not True: - # Gateway-managed interactive (authorization_code): the only - # thing that authorizes egress is a stored per-user token, so - # challenge whenever one is absent, regardless of any bearer. - # The v2 resolver owns the existence check, so every - # authorization_code resolution (egress and this discovery - # challenge) runs through it. A keyless admitted subject is - # challenged with the per-server resource_metadata (whose - # authorization server is the gateway itself, vaulting via the - # authorize interlude); the per-server relay advertised below - # cannot vault without a litellm key on its token request. - if await operations.global_mcp_server_manager.has_user_oauth_token(server, user_api_key_auth): - continue - - if _is_mcp_admitted_user_subject(user_api_key_auth): - raise HTTPException( - status_code=401, - detail="Unauthorized", - headers={ - "www-authenticate": get_passthrough_www_authenticate( - scope=scope, - server_name=server_name, - ) - }, - ) - - request = StarletteRequest(scope) - base_url = get_request_base_url(request) - _path = get_route_relative_request_path(scope) - - # Pick the well-known AS-metadata form that matches the inbound route - # so strict RFC 9728 §3.2 clients can resolve it correctly. - as_metadata_root = f"{base_url}/.well-known/oauth-authorization-server{well_known_root_suffix()}" - if _path.startswith(f"/mcp/{server_name}"): - _as_url = f"{as_metadata_root}/mcp/{server_name}" - else: - _as_url = f"{as_metadata_root}/{server_name}" - authorization_uri = f'Bearer authorization_uri="{_as_url}"' - - raise HTTPException( - status_code=401, - detail="Unauthorized", - headers={"www-authenticate": authorization_uri}, - ) - - if not oauth2_headers: - # Delegate-auth servers run upstream PKCE: a present bearer is - # the upstream token, so only challenge when it is absent, with - # the proxied resource_metadata (RFC 9728), not the gateway - # authorization_uri above which would authorize against the - # gateway instead of the upstream IdP. - www_authenticate = get_passthrough_www_authenticate( + 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( + *( + _server_auth_challenge( + configured_server=server, + server_name=server.alias or server.server_name or server.name, scope=scope, - server_name=server_name, + 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, + raw_headers=raw_headers, ) - raise HTTPException( - status_code=401, - detail="Unauthorized", - headers={"www-authenticate": www_authenticate}, + for server in eligible + ), + return_exceptions=True, + ) + failures: Final = tuple( + (server, result) for server, result in zip(eligible, results) if isinstance(result, BaseException) + ) + for server, failure in failures: + if not isinstance(failure, Exception): + raise failure + if not isinstance(failure, HTTPException) or failure.status_code != 401: + verbose_logger.warning( + "MCP authentication preflight failed for %s (%s)", server.name, type(failure).__name__ ) - # Delegate server with a bearer present: it is the upstream token, - # so admit the session and move to the next target. Every oauth2 - # sub-mode is terminal here (continue or raise) so no oauth2 server - # reaches the token_exchange / pass-through blocks below. - continue - - # token_exchange (OBO): the caller supplied no subject token. Challenge at connect - # (transport level, where WWW-Authenticate survives) with the RFC 9728 resource_metadata - # so the client discovers the IdP, SSOs, and retries with a subject token, which LiteLLM - # then exchanges. A tool-call-time 401 would be wrapped into a JSON-RPC error and the - # header lost, so the discovery flow needs this pre-emptive challenge. - if server and server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers: - from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 # lazy: adapter pulls MCP subgraph - raise_token_exchange_challenge, + if not failures or len(failures) != len(results): + return + for _, failure in failures: + if not isinstance(failure, HTTPException) or failure.status_code != 401: + raise failure + 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 ) - from litellm.proxy.middleware.per_request_root_path_middleware import ( # noqa: PLC0415 # lazy: middleware imports proxy utils - get_request_root_path, - ) - - raise_token_exchange_challenge(server, root_path=get_request_root_path()) - - # Exchange-backed modes (token_exchange's OBO mint, id_jag's stored-assertion mint): run - # the exchange here at the transport edge, so a rejected subject raises the RFC 9728 - # challenge and any other failure its public status, instead of the session opening and - # list_tools masking it as an empty tool list. The manager owns which modes pre-flight - # and what each mints from. Gated to single-server routes the key may reach; the - # multi-server aggregate keeps absorbing per-server auth failures so one bad server - # cannot 401 the whole connect. + raise failures[0][1] + for server_name in mcp_servers: if ( - server - and len(mcp_servers or []) == 1 - and server.server_id - in frozenset( - allowed.server_id - for allowed in await operations._get_allowed_mcp_servers( - user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip - ) - ) - ): - await operations.global_mcp_server_manager.preflight_token_exchange( - server=server, + server := operations.global_mcp_server_manager.get_mcp_server_by_name(server_name, client_ip=client_ip) + ) is None: + continue + if allowed_server_ids is not None and server.server_id not in allowed_server_ids: + continue + if ( + challenge := await _server_auth_challenge( + configured_server=server, + server_name=server_name, + 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, raw_headers=raw_headers, ) - - # Pass-through OAuth: when the admin has opted a server into - # forwarding the client's bearer token (is_oauth_passthrough) and - # the client hasn't supplied one, fail fast with 401 and point - # them at the gateway's oauth-protected-resource well-known URL. - # That endpoint proxies the upstream's metadata so the client - # kicks off OAuth against the real upstream IdP, not the gateway. - if ( - server - and server.is_oauth_passthrough - and not operations._client_has_passthrough_authorization( - server, oauth2_headers, mcp_server_auth_headers - ) - ): - www_authenticate = get_passthrough_www_authenticate( - scope=scope, - server_name=server_name, - ) - raise HTTPException( - status_code=401, - detail="Unauthorized", - headers={"www-authenticate": www_authenticate}, - ) - - if ( - server - and server.is_oauth_delegate - and len(mcp_servers or []) == 1 - and _get_forwarded_auth_from_scope(scope) is None - and not operations._client_has_per_server_auth_header(server, mcp_server_auth_headers) - ): - www_authenticate = get_passthrough_www_authenticate( - scope=scope, - server_name=server_name, - ) - raise HTTPException( - status_code=401, - detail="Unauthorized", - headers={"www-authenticate": www_authenticate}, - ) - - if ( - server - and server.is_true_passthrough - and len(mcp_servers or []) == 1 - and not _scope_has_authorization_header(scope) - and not operations._client_has_per_server_auth_header(server, mcp_server_auth_headers) - ): - if server.is_dcr_bridge: - raise HTTPException( - status_code=401, - detail="Unauthorized", - headers={ - "www-authenticate": get_passthrough_www_authenticate( - scope=scope, - server_name=server_name, - ) - }, - ) - upstream_status, upstream_www_authenticate = await _probe_upstream_auth(server.url or "", "") - if upstream_status == 401 and upstream_www_authenticate: - raise HTTPException( - status_code=401, - detail="Unauthorized", - headers={"www-authenticate": upstream_www_authenticate}, - ) + ) is not None: + raise challenge def _get_authorization_header_from_scope(scope: Scope) -> str | None: """First ``Authorization`` header value in the ASGI scope, or None.""" @@ -2017,16 +2122,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 @@ -2103,6 +2209,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 @@ -2248,7 +2366,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/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 78b7b729375..99639d30a05 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -31506,6 +31506,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/faults/test_list_outcomes.py b/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_list_outcomes.py index f951499e18f..33c4c2a6c19 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,59 @@ 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), + ({"timeout": ServerListFault(tag="timeout")}, None), + ({"denied": ServerListFault(tag="forbidden", status_code=403), "timeout": ServerListFault(tag="timeout")}, 403), + ({"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")}, 401), + ), +) +def test_listing_auth_failure_requires_no_successful_server( + 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_gateway_dcr_flow.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py index 2943ff4b74a..9a752ab14a4 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_id", + 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-id",), + 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-id",), + 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-id",), + 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-id",), + 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-id",), + 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-id",), + 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-id",), + 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-id",), + 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-id",), + 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-id",), + 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-id",), + lookup_server_reachability=_ServerReachability(), request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, session_user_id="u1", @@ -997,7 +1026,9 @@ 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_id.side_effect = lambda server_id: ( + scoped_server or _scoped_mcp_server("public", auth_type="none") + ) if server_id == "public-id" else scoped_server return await complete_connect_flow( request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, @@ -1005,7 +1036,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-id",), **overrides}, ) @@ -1814,6 +1845,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-id",), + lookup_server_reachability=_ServerReachability(), request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, session_user_id="u1", @@ -2354,3 +2387,123 @@ 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_completion_requires_an_upstream_selection(): + client_id = (await _register([REDIRECT_URI]))["client_id"] + 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() + server = _scoped_mcp_server(oauth2_flow="authorization_code") + vendor = _VendorCredential(credential) + completed = await _complete_page( + response, + scoped_server=server, + vendor=vendor, + reachable=_ServerReachability(reachable), + cache=cache, + selected_servers=("github-id",), + ) + assert completed.status_code == status + if not reachable: + assert vendor.calls == [] + if status != 303: + assert "location" not in completed.headers + retried = await _complete_page(response, scoped_server=server, cache=cache, selected_servers=("github-id",)) + 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-id", "slack-id"), lookup_vendor_credential=credential + ) + assert unfinished.status_code == 400 + with pytest.raises(asyncio.CancelledError): + await complete_connect_flow(**arguments, selected_servers=("github-id",), lookup_vendor_credential=cancelled) + completed = await complete_connect_flow( + **arguments, selected_servers=("github-id", "slack-id"), lookup_vendor_credential=_VendorCredential() + ) + assert completed.status_code == 303 + + +@pytest.mark.asyncio +async def test_unified_completion_validates_selected_id_despite_alias_collision(): + 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) + selected = _scoped_mcp_server("github", oauth2_flow="authorization_code") + other = _scoped_mcp_server("other", auth_type="none").model_copy(update={"alias": selected.server_id}) + cache = DualCache() + with patch(_MANAGER_PATCH) as manager: + manager.get_mcp_server_by_name.return_value = other + manager.get_mcp_server_by_id.side_effect = {selected.server_id: selected, other.server_id: other}.get + vendor = _VendorCredential("absent") + arguments = dict(request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, session_user_id="u1", cache=cache, selected_servers=(selected.server_id,), lookup_server_reachability=_ServerReachability()) + refused = await complete_connect_flow(**arguments, lookup_vendor_credential=vendor) + assert refused.status_code == 400 + assert vendor.calls == [("u1", selected.server_id)] + completed = await complete_connect_flow(**arguments, lookup_vendor_credential=_VendorCredential()) + assert completed.status_code == 303 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_max_concurrent_requests.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_max_concurrent_requests.py index 0dc7ac5ecd9..1f426bddc3b 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_max_concurrent_requests.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_max_concurrent_requests.py @@ -47,7 +47,7 @@ def _make_server(server_id: str, max_concurrent_requests: Optional[int]) -> MCPS def _patch_client_with_tracker(manager: MCPServerManager, tracker: _ConcurrencyTracker): async def fake_create_mcp_client(server, **kwargs): class _ProbeClient: - async def call_tool(self, params, host_progress_callback=None, allow_input_required=False): + async def call_tool(self, params, host_progress_callback=None, allow_input_required=False, raise_on_error: bool = False): tracker.enter(server.server_id) try: await asyncio.sleep(HOLD_SECONDS) 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..c625b80e1fb 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,469 @@ 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)) +@pytest.mark.parametrize("dcr_bridge", (False, True)) +async def test_manager_preserves_auth_failures_for_prompts_and_resources( + operation: str, status: int, dcr_bridge: bool, monkeypatch: pytest.MonkeyPatch +) -> None: + from fastapi import HTTPException + from pydantic import AnyUrl + + manager: Final = MCPServerManager() + upstream: Final = _http_server("upstream", "upstream", auth_type=MCPAuth.oauth_delegate, dcr_bridge=dcr_bridge) + 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 == (None if dcr_bridge else "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")) +@pytest.mark.parametrize("extra_headers", (None, {"x-forwarded": "caller"})) +async def test_prompt_and_resource_calls_preserve_static_headers_and_non_auth_failures( + operation: str, extra_headers: dict[str, str] | None, 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, extra_headers=extra_headers, **kwargs) + assert caught.value is failure + assert create.await_args.kwargs["extra_headers"] == {**(extra_headers or {}), "x-upstream": "configured"} + create.assert_awaited_once() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("preamble", (b"id: resume-token\r\ndata: \r\n\r\n", b": ping\r\n\r\n")) +async def test_transport_preserves_sse_priming_event_on_success(preamble: bytes) -> 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": preamble, "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() + + +@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() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("kind", ("prompts", "resources", "resource_templates")) +async def test_optional_listing_challenges_auth_when_other_server_times_out( + kind: str, monkeypatch: pytest.MonkeyPatch +) -> None: + from litellm.proxy._experimental.mcp_server import operations + + blocked: Final = _http_server("blocked", "blocked") + unavailable: Final = _http_server("unavailable", "unavailable") + manager: Final = MCPServerManager() + create: Final = AsyncMock(side_effect=[MCPUpstreamAuthError(401, "Bearer", "blocked"), TimeoutError()]) + monkeypatch.setattr(manager, "_create_mcp_client", create) + 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, unavailable])) + with pytest.raises(MCPUpstreamAuthError) as caught: + await getattr(operations, f"_list_mcp_{kind}")() + assert caught.value.status_code == 401 + assert caught.value.www_authenticate == "Bearer" + assert caught.value.server_name == "blocked" + assert create.await_count == 2 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("kind", ("prompts", "resources", "resource_templates")) +@pytest.mark.parametrize("healthy_first", (False, True)) +async def test_optional_listing_preserves_healthy_duplicate_names(kind: str, healthy_first: bool, monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy._experimental.mcp_server import operations + healthy: Final = _http_server("healthy-id", "duplicate") + blocked: Final = _http_server("blocked-id", "duplicate") + manager: Final = MagicMock() + fetch: Final = AsyncMock(side_effect=[[], MCPUpstreamAuthError(401, "Bearer", "duplicate")] if healthy_first else [MCPUpstreamAuthError(401, "Bearer", "duplicate"), []]) + 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=[healthy, blocked] if healthy_first else [blocked, healthy])) + assert await getattr(operations, f"_list_mcp_{kind}")() == [] + assert fetch.await_count == 2 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 a7e56f3f84a..441552692e4 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 @@ -2804,7 +2805,7 @@ async def test_initialize_request_tracks_active_session_after_response_header(): patch( # test-quality-ok: registry is empty in unit tests; key owns one server "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new_callable=AsyncMock, - return_value=[MagicMock()], + return_value=[MCPServer(server_id="available", name="available", transport=MCPTransport.http)], ), patch( "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", @@ -2957,7 +2958,7 @@ async def test_initialize_request_records_client_name_in_gateway_sessions_report patch( # test-quality-ok: registry is empty in unit tests; key owns one server "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new_callable=AsyncMock, - return_value=[MagicMock()], + return_value=[MCPServer(server_id="available", name="available", transport=MCPTransport.http)], ), patch( # test-quality-ok: session manager init is a module-level flag; the suite's only seam "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", @@ -10685,3 +10686,254 @@ 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"), RuntimeError("discovery failed"), 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, + ) + with pytest.raises(type(failure)) as caught: + await request + if not isinstance(failure, asyncio.CancelledError): + assert caught.value is failure + 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"' + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("unified", (False, True)) +@pytest.mark.parametrize( + "auth_type,bridge,upstream_status,expected_challenge", + ( + (MCPAuth.none, False, 401, "gateway"), + (MCPAuth.oauth_delegate, False, 401, "gateway"), + (MCPAuth.true_passthrough, True, 401, "gateway"), + (MCPAuth.true_passthrough, False, 401, "upstream"), + (MCPAuth.true_passthrough, False, 200, None), + ), +) +async def test_preflight_preserves_client_forwarded_auth_challenges( + monkeypatch: pytest.MonkeyPatch, + unified: bool, + auth_type: MCPAuth, + bridge: bool, + upstream_status: int, + expected_challenge: str | None, +) -> None: + from litellm.proxy._experimental.mcp_server import server as server_module + + upstream: Final = _client_forwarded_mode_server("forwarded", auth_type).model_copy( + update={ + "dcr_bridge": bridge, + "oauth_passthrough": auth_type == MCPAuth.none, + "extra_headers": ["Authorization"] if auth_type == MCPAuth.none else None, + } + ) + 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) + probe: Final = AsyncMock(return_value=(upstream_status, "Bearer realm=upstream")) + monkeypatch.setattr(server_module, "_probe_upstream_auth", probe) + request: Final = server_module._raise_preemptive_401_for_unauthenticated_servers( + scope={"type": "http", "path": "/mcp", "method": "POST", "headers": [(b"host", b"gateway")]}, + mcp_servers=None if unified else [upstream.name], + oauth2_headers=None, + mcp_server_auth_headers=None, + user_api_key_auth=UserAPIKeyAuth(user_id="reader"), + client_ip=None, + ) + if expected_challenge is None: + await request + else: + with pytest.raises(HTTPException) as caught: + await request + assert caught.value.status_code == 401 + expected: Final = ( + 'Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp/forwarded"' + if expected_challenge == "gateway" + else "Bearer realm=upstream" + ) + assert (caught.value.headers or {})["www-authenticate"] == expected + assert probe.await_count == int(auth_type == MCPAuth.true_passthrough and not bridge) + + +@pytest.mark.asyncio +async def test_preflight_gateway_subject_uses_server_resource_metadata(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy._experimental.mcp_server import server as server_module + + upstream: Final = _make_oauth2_server("managed") + auth: Final = UserAPIKeyAuth(user_id="reader") + auth.mcp_admitted_user_subject = True + manager: Final = mcp_operations.global_mcp_server_manager + monkeypatch.setattr(manager, "get_mcp_server_by_name", lambda *args, **kwargs: upstream) + monkeypatch.setattr(manager, "has_user_oauth_token", AsyncMock(return_value=False)) + with pytest.raises(HTTPException) as caught: + await server_module._raise_preemptive_401_for_unauthenticated_servers( + scope={"type": "http", "path": "/mcp/managed", "method": "POST", "headers": [(b"host", b"gateway")]}, + mcp_servers=["managed"], + oauth2_headers=None, + mcp_server_auth_headers=None, + user_api_key_auth=auth, + 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/managed"' + ) + + +@pytest.mark.asyncio +async def test_preflight_does_not_request_oauth_for_excluded_server(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy._experimental.mcp_server import server as server_module + + upstream: Final = _make_oauth2_server("excluded") + manager: Final = mcp_operations.global_mcp_server_manager + monkeypatch.setattr(manager, "get_mcp_server_by_name", lambda *args, **kwargs: upstream) + discovery: Final = AsyncMock(return_value=upstream) + tokens: Final = AsyncMock(return_value=False) + monkeypatch.setattr(manager, "ensure_oauth_metadata_discovered", discovery) + monkeypatch.setattr(manager, "has_user_oauth_token", tokens) + await server_module._raise_preemptive_401_for_unauthenticated_servers( + scope={"type": "http", "path": "/mcp/excluded", "method": "POST", "headers": []}, + mcp_servers=["excluded"], + oauth2_headers=None, + mcp_server_auth_headers=None, + user_api_key_auth=UserAPIKeyAuth(user_id="reader"), + client_ip=None, + allowed_server_ids={"other"}, + ) + discovery.assert_not_awaited() + tokens.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("healthy_peer", (False, True)) +@pytest.mark.parametrize("failure_first", (False, True)) +@pytest.mark.parametrize("failure", (HTTPException(503, "unavailable"), RuntimeError("discovery failed"), asyncio.CancelledError())) +async def test_unified_preflight_preserves_usable_peer_during_discovery_failure( + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, + healthy_peer: bool, + failure_first: bool, + failure: BaseException, +) -> None: + from litellm.proxy._experimental.mcp_server import server as server_module + + unavailable: Final = _make_oauth2_server("unavailable") + peer: Final = _make_oauth2_server("peer") + servers: Final = [unavailable, peer] if failure_first else [peer, unavailable] + monkeypatch.setattr(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=servers)) + manager: Final = mcp_operations.global_mcp_server_manager + + async def discover(server: MCPServer) -> MCPServer: + if server.server_id == unavailable.server_id: + raise failure + return server + + monkeypatch.setattr(manager, "ensure_oauth_metadata_discovered", discover) + tokens: Final = AsyncMock(return_value=healthy_peer) + monkeypatch.setattr(manager, "has_user_oauth_token", tokens) + request: Final = 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, + ) + if healthy_peer and not isinstance(failure, asyncio.CancelledError): + with caplog.at_level("WARNING", logger="LiteLLM"): + await request + assert any("unavailable" in record.getMessage() for record in caplog.records) + assert str(failure) not in caplog.text + else: + with pytest.raises(type(failure)) as caught: + await request + if not isinstance(failure, asyncio.CancelledError): + assert caught.value is failure + tokens.assert_awaited_once() + assert tokens.await_args.args[0].server_id == peer.server_id 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 a1dc0e779da..88fb0f87b59 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 @@ -1844,7 +1844,11 @@ class TestMCPServerManager: server, mcp_auth_header, extra_headers, stdio_env, subject_token=None, **kwargs ): # pragma: no cover - helper captured["subject_token"] = subject_token - return AsyncMock() + return AsyncMock( + discovery_auth_fingerprint=AsyncMock(return_value="test-credential-hash"), + list_prompts=AsyncMock(return_value=[]), + list_resources=AsyncMock(return_value=[]), + ) manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) manager._fetch_tools_with_timeout = AsyncMock(return_value=[]) @@ -2049,9 +2053,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", @@ -2078,7 +2081,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( @@ -2633,7 +2636,11 @@ class TestMCPServerManager: server, mcp_auth_header, extra_headers, stdio_env, subject_token=None, **kwargs ): # pragma: no cover - helper captured["subject_token"] = subject_token - return AsyncMock() + return AsyncMock( + discovery_auth_fingerprint=AsyncMock(return_value="test-credential-hash"), + list_prompts=AsyncMock(return_value=[]), + list_resources=AsyncMock(return_value=[]), + ) manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) await call(manager) @@ -6911,7 +6918,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"}] @@ -13146,6 +13153,7 @@ class TestLitellmAdmissionKeyIsNeverTheSubjectToken: client: Final = AsyncMock() client.call_tool = AsyncMock(return_value=CallToolResult(content=[], isError=False)) client.list_prompts = AsyncMock(return_value=[]) + client.discovery_auth_fingerprint = AsyncMock(return_value="test-credential-hash") client.read_resource = AsyncMock(return_value=ReadResourceResult(contents=[])) manager._create_mcp_client = AsyncMock(return_value=client) return manager @@ -13925,8 +13933,12 @@ async def test_discovery_cache_empty_results_and_failures(kind: str, outcome: st "templates": manager.get_resource_templates_from_server, }[kind] with _mcp_upstream(upstream.respond): - assert await operation(_discovery_server(), None) == [] - assert await operation(_discovery_server(), None) == [] + for _ in range(2): + if outcome == "failure": + with pytest.raises(MCPServerListError, match="discovery"): + await operation(_discovery_server(), None) + else: + assert await operation(_discovery_server(), None) == [] assert upstream.initializes == (2 if outcome == "failure" else 1) if outcome == "failure": upstream.outcome = "supported" @@ -13946,7 +13958,8 @@ async def test_discovery_cache_retries_failed_pagination_before_caching_complete "templates": manager.get_resource_templates_from_server, }[kind] with _mcp_upstream(upstream.respond): - assert await operation(_discovery_server(), None) == [] + with pytest.raises(MCPServerListError, match="discovery"): + await operation(_discovery_server(), None) assert upstream.initializes == 1 upstream.outcome = "paged" recovered: Final = await operation(_discovery_server(), None) @@ -14197,6 +14210,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 @@ -14254,7 +14268,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 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 && (