From 97a662f2b223f08e9208b9a11f2303fb92d1a6d6 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Mon, 28 Sep 2026 11:41:54 -0700 Subject: [PATCH] fix(mcp): preserve upstream authentication challenges on protocol requests --- .../mcp_server/faults/list_outcomes.py | 45 ++- .../mcp_server/mcp_server_manager.py | 129 ++++--- .../_experimental/mcp_server/operations.py | 212 +++++------ .../proxy/_experimental/mcp_server/server.py | 172 +++++++-- .../mcp_server/faults/test_list_outcomes.py | 54 +++ .../test_mcp_oauth_passthrough_tools.py | 356 ++++++++++++++++++ .../mcp_server/test_mcp_server.py | 3 +- .../mcp_server/test_mcp_server_manager.py | 14 +- 8 files changed, 747 insertions(+), 238 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/faults/list_outcomes.py b/litellm/proxy/_experimental/mcp_server/faults/list_outcomes.py index b96a7a74e4a..0059498e60a 100644 --- a/litellm/proxy/_experimental/mcp_server/faults/list_outcomes.py +++ b/litellm/proxy/_experimental/mcp_server/faults/list_outcomes.py @@ -10,13 +10,14 @@ becomes an outcome, never a second failure. from __future__ import annotations -from collections.abc import Iterator +from collections.abc import Iterator, Mapping from typing import Final, Literal, NamedTuple, NoReturn, TypeAlias import httpx import httpx2 +from fastapi import HTTPException from mcp.types import Tool as MCPTool -from pydantic import BaseModel, ConfigDict +from pydantic import BaseModel, ConfigDict, Field from typing_extensions import assert_never from litellm.proxy._experimental.mcp_server.exceptions import ( @@ -50,6 +51,8 @@ class ServerListFault(BaseModel): model_config = ConfigDict(frozen=True) tag: ListFaultCategory status_code: int | None = None + www_authenticate: str | None = Field(default=None, exclude=True, repr=False) + server_name: str | None = Field(default=None, exclude=True, repr=False) ServerOutcome: TypeAlias = ServerListOk | ServerListFault @@ -64,6 +67,20 @@ class AggregateToolListing(NamedTuple): outcomes: dict[str, ServerOutcome] +def listing_auth_error(outcomes: Mapping[str, ServerOutcome]) -> MCPUpstreamAuthError | None: + blocked: Final = tuple( + (name, outcome) + for name, outcome in outcomes.items() + if isinstance(outcome, ServerListFault) and outcome.tag in ("auth_required", "forbidden") + ) + if not blocked or len(blocked) != len(outcomes): + return None + name, outcome = next((entry for entry in blocked if entry[1].tag == "auth_required"), blocked[0]) + return MCPUpstreamAuthError( + 401 if outcome.tag == "auth_required" else 403, outcome.www_authenticate, outcome.server_name or name + ) + + def _iter_upstream_responses(exc: BaseException) -> Iterator[httpx.Response | httpx2.Response]: """Yield every upstream ``httpx``/``httpx2`` ``Response`` in the exception tree, in the shared traversal's deliberate order (explicit causes first, ExceptionGroup members in raise order, the incidental @@ -87,9 +104,20 @@ def upstream_auth_challenge(exc: BaseException) -> tuple[int, str | None] | None rides with it can never come from two different responses in the tree. Non-auth responses do not end the scan: a causal 401 behind an unrelated 5xx must still be found, or the client never receives the challenge it needs to re-authenticate.""" - for response in _iter_upstream_responses(exc): - if response.status_code in (401, 403): - return response.status_code, response.headers.get("www-authenticate") + return next( + (challenge for current in iter_exception_tree(exc) if (challenge := _auth_challenge(current)) is not None), None + ) + + +def _auth_challenge(exc: BaseException) -> tuple[int, str | None] | None: + if isinstance(exc, MCPUpstreamAuthError): + return exc.status_code, exc.www_authenticate + if isinstance(exc, HTTPException) and exc.status_code in (401, 403): + headers: Final = exc.headers or {} + return exc.status_code, headers.get("WWW-Authenticate") or headers.get("www-authenticate") + response: Final = getattr(exc, "response", None) + if isinstance(response, (httpx.Response, httpx2.Response)) and response.status_code in (401, 403): + return response.status_code, response.headers.get("www-authenticate") return None @@ -122,17 +150,20 @@ def classify_list_exception(exc: BaseException) -> ServerListFault: return exc.fault if isinstance(exc, MCPUpstreamAuthError): tag: Final = "forbidden" if exc.status_code == 403 else "auth_required" - return ServerListFault(tag=tag, status_code=exc.status_code) + return ServerListFault( + tag=tag, status_code=exc.status_code, www_authenticate=exc.www_authenticate, server_name=exc.server_name + ) if isinstance(exc, TimeoutError): return ServerListFault(tag="timeout") if isinstance(exc, ConnectionError): return ServerListFault(tag="unreachable") auth: Final = upstream_auth_challenge(exc) if auth is not None: - status_code, _ = auth + status_code, challenge = auth return ServerListFault( tag="forbidden" if status_code == 403 else "auth_required", status_code=status_code, + www_authenticate=challenge, ) response: Final = _find_upstream_response(exc) if response is not None: diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index ea685c431bd..375cf9fc200 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -4632,6 +4632,8 @@ class MCPServerManager: return self._create_prefixed_prompts(items, server, add_prefix=add_prefix) except Exception as error: verbose_logger.warning("Failed to get prompts from server %s: %s", server.name, error) + if upstream_auth_challenge(error) is not None: + raise_classified_list_failure(error, server.name, suppress_challenge=server.is_dcr_bridge) return [] async def get_resources_from_server( @@ -4678,6 +4680,8 @@ class MCPServerManager: return self._create_prefixed_resources(items, server, add_prefix=add_prefix) except Exception as error: verbose_logger.warning("Failed to get resources from server %s: %s", server.name, error) + if upstream_auth_challenge(error) is not None: + raise_classified_list_failure(error, server.name, suppress_challenge=server.is_dcr_bridge) return [] async def get_resource_templates_from_server( @@ -4724,6 +4728,8 @@ class MCPServerManager: return self._create_prefixed_resource_templates(items, server, add_prefix=add_prefix) except Exception as error: verbose_logger.warning("Failed to get resource_templates from server %s: %s", server.name, error) + if upstream_auth_challenge(error) is not None: + raise_classified_list_failure(error, server.name, suppress_challenge=server.is_dcr_bridge) return [] async def read_resource_from_server( @@ -4738,29 +4744,38 @@ class MCPServerManager: ) -> ReadResourceResult: """Read resource contents from a specific MCP server.""" - verbose_logger.debug("Connecting to url: %s", server.url) - verbose_logger.info("read_resource_from_server for %s...", server.name) + try: + verbose_logger.debug("Connecting to url: %s", server.url) + verbose_logger.info("read_resource_from_server for %s...", server.name) - if server.static_headers: - if extra_headers is None: - extra_headers = {} - extra_headers.update(server.static_headers) + if server.static_headers: + if extra_headers is None: + extra_headers = {} + extra_headers.update(server.static_headers) - stdio_env: Final = self._build_stdio_env(server, raw_headers) - subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth) + stdio_env: Final = self._build_stdio_env(server, raw_headers) + subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth) - client: Final = await self._create_mcp_client( - server=server, - mcp_auth_header=mcp_auth_header, - extra_headers=extra_headers, - stdio_env=stdio_env, - subject_token=subject_token, - raw_headers=raw_headers, - client_ip=client_ip, - user_api_key_auth=user_api_key_auth, - ) + client: Final = await self._create_mcp_client( + server=server, + mcp_auth_header=mcp_auth_header, + extra_headers=extra_headers, + stdio_env=stdio_env, + subject_token=subject_token, + raw_headers=raw_headers, + client_ip=client_ip, + user_api_key_auth=user_api_key_auth, + ) - return await client.read_resource(url) + return await client.read_resource(url) + except Exception as exc: + auth_failure: Final = upstream_auth_challenge(exc) + if auth_failure is not None: + status_code, challenge = auth_failure + raise MCPUpstreamAuthError( + status_code, None if server.is_dcr_bridge else challenge, server.name + ) from exc + raise async def get_prompt_from_server( self, @@ -4775,33 +4790,42 @@ class MCPServerManager: ) -> GetPromptResult: """Fetch a specific prompt definition from a single MCP server.""" - verbose_logger.debug("Connecting to url: %s", server.url) - verbose_logger.info("get_prompt_from_server for %s...", server.name) + try: + verbose_logger.debug("Connecting to url: %s", server.url) + verbose_logger.info("get_prompt_from_server for %s...", server.name) - if server.static_headers: - if extra_headers is None: - extra_headers = {} - extra_headers.update(server.static_headers) + if server.static_headers: + if extra_headers is None: + extra_headers = {} + extra_headers.update(server.static_headers) - stdio_env: Final = self._build_stdio_env(server, raw_headers) - subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth) + stdio_env: Final = self._build_stdio_env(server, raw_headers) + subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth) - client: Final = await self._create_mcp_client( - server=server, - mcp_auth_header=mcp_auth_header, - extra_headers=extra_headers, - stdio_env=stdio_env, - subject_token=subject_token, - raw_headers=raw_headers, - client_ip=client_ip, - user_api_key_auth=user_api_key_auth, - ) + client: Final = await self._create_mcp_client( + server=server, + mcp_auth_header=mcp_auth_header, + extra_headers=extra_headers, + stdio_env=stdio_env, + subject_token=subject_token, + raw_headers=raw_headers, + client_ip=client_ip, + user_api_key_auth=user_api_key_auth, + ) - get_prompt_request_params: Final = GetPromptRequestParams( - name=prompt_name, - arguments=arguments, - ) - return await client.get_prompt(get_prompt_request_params) + get_prompt_request_params: Final = GetPromptRequestParams( + name=prompt_name, + arguments=arguments, + ) + return await client.get_prompt(get_prompt_request_params) + except Exception as exc: + auth_failure: Final = upstream_auth_challenge(exc) + if auth_failure is not None: + status_code, challenge = auth_failure + raise MCPUpstreamAuthError( + status_code, None if server.is_dcr_bridge else challenge, server.name + ) from exc + raise @staticmethod def _is_same_authority_metadata_url(url: str, server_url: str) -> bool: @@ -6094,29 +6118,10 @@ class MCPServerManager: tool_call_coro = _obo_call_tool_limited() else: - # Scoped to the two client-forwarded token modes this stack introduced; legacy - # oauth2 + delegate_auth_to_upstream (is_oauth_passthrough) is being removed, so it is not - # added here even though the list path still relays for it. - relays_upstream_auth: Final = mcp_server.is_client_forwarded_token server_label: Final = mcp_server.name or mcp_server.server_name or mcp_server.alias or "" async def _call_tool_via_client(client, params): async with self._limit_outbound_concurrency(mcp_server): - if not relays_upstream_auth: - return await client.call_tool( - params, - host_progress_callback=host_progress_callback, - allow_input_required=allow_input_required, - ) - # The client-forwarded modes carry the caller's own upstream token, so an upstream - # 401 (expired/invalid token) is the caller's to resolve: relay it as - # MCPUpstreamAuthError so single-server REST callers turn it into a 401 + - # WWW-Authenticate and re-run the upstream OAuth flow. Only 401 is a re-auth signal - # (mirrors the list path and MCPUpstreamAuthError's contract); a 403 is a genuine - # authorization failure that re-auth won't fix, so it takes the non-auth branch and - # stays a visible warning. raise_on_error only re-raises transport failures - # (tool-level isError results are still returned normally); a non-auth failure keeps - # the same isError degradation the default path produces. try: return await client.call_tool( params, @@ -6142,7 +6147,7 @@ class MCPServerManager: _, www_authenticate = auth_info raise MCPUpstreamAuthError( status_code=401, - www_authenticate=www_authenticate, + www_authenticate=None if mcp_server.is_dcr_bridge else www_authenticate, server_name=server_label, ) from e diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index a19246b6e90..0c037ebf893 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -4,9 +4,10 @@ import asyncio import traceback import types import uuid -from collections.abc import Mapping, Sequence +from collections.abc import Awaitable, Callable, Mapping, Sequence from datetime import datetime -from typing import Any, Final, NoReturn, TypeAlias, overload +from itertools import chain +from typing import Any, Final, NoReturn, TypeAlias, TypeVar, overload from fastapi import HTTPException from mcp import ReadResourceResult, Resource @@ -78,6 +79,7 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ( ServerListOk, ServerOutcome, classify_list_exception, + listing_auth_error, outcome_wire_value, ) from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( @@ -1242,6 +1244,26 @@ async def _get_tools_from_mcp_servers( raise +_ListingItem = TypeVar("_ListingItem", Prompt, Resource, ResourceTemplate) + + +async def _collect_mcp_listing( + servers: Sequence[MCPServer], fetch: Callable[[MCPServer], Awaitable[list[_ListingItem]]] +) -> list[_ListingItem]: + async def fetch_one(server: MCPServer) -> tuple[list[_ListingItem], ServerOutcome]: + try: + items: Final = await fetch(server) + return items, ServerListOk(tool_count=len(items)) + except Exception as exc: + return [], classify_list_exception(exc) + + results: Final = await asyncio.gather(*(fetch_one(server) for server in servers)) + failure: Final = listing_auth_error({server.name: result[1] for server, result in zip(servers, results)}) + if failure is not None: + raise failure + return list(chain.from_iterable(items for items, _ in results)) + + async def _get_prompts_from_mcp_servers( user_api_key_auth: UserAPIKeyAuth | None, mcp_auth_header: str | None, @@ -1251,32 +1273,13 @@ async def _get_prompts_from_mcp_servers( raw_headers: dict[str, str] | None = None, client_ip: str | None = None, ) -> list[Prompt]: - """ - Helper method to fetch prompt from MCP servers based on server filtering criteria. - - Args: - user_api_key_auth: User authentication info for access control - mcp_auth_header: Optional auth header for MCP server (deprecated) - mcp_servers: Optional list of server names/aliases to filter by - mcp_server_auth_headers: Optional dict of server-specific auth headers - oauth2_headers: Optional dict of oauth2 headers - - Returns: - List[Prompt]: Combined list of prompts from filtered servers - """ - - allowed_mcp_servers: Final = await _get_allowed_mcp_servers( + allowed: Final = await _get_allowed_mcp_servers( user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip, ) - # Get prompts from each allowed server - all_prompts: Final = [] - for server in allowed_mcp_servers: - if server is None: - continue - + async def fetch(server: MCPServer) -> list[Prompt]: server_auth_header, extra_headers = _prepare_mcp_server_headers( server=server, mcp_server_auth_headers=mcp_server_auth_headers, @@ -1284,30 +1287,19 @@ async def _get_prompts_from_mcp_servers( oauth2_headers=oauth2_headers, raw_headers=raw_headers, user_api_key_auth=user_api_key_auth, - scope_servers=allowed_mcp_servers, + scope_servers=allowed, + ) + return await global_mcp_server_manager.get_prompts_from_server( + server=server, + user_api_key_auth=user_api_key_auth, + mcp_auth_header=server_auth_header, + extra_headers=extra_headers, + add_prefix=True, + raw_headers=raw_headers, + client_ip=client_ip, ) - try: - prompts = await global_mcp_server_manager.get_prompts_from_server( - server=server, - user_api_key_auth=user_api_key_auth, - mcp_auth_header=server_auth_header, - extra_headers=extra_headers, - add_prefix=True, # Always add server prefix - raw_headers=raw_headers, - client_ip=client_ip, - ) - - all_prompts.extend(prompts) - - verbose_logger.debug("Successfully fetched %s prompts from server %s", len(prompts), server.name) - except Exception as e: - verbose_logger.exception("Error getting prompts from server %s: %s", server.name, e) - # Continue with other servers instead of failing completely - - verbose_logger.info("Successfully fetched %s prompts total from all MCP servers", len(all_prompts)) - - return all_prompts + return await _collect_mcp_listing(tuple(server for server in allowed if server is not None), fetch) async def _get_resources_from_mcp_servers( @@ -1319,19 +1311,13 @@ async def _get_resources_from_mcp_servers( raw_headers: dict[str, str] | None = None, client_ip: str | None = None, ) -> list[Resource]: - """Fetch resources from allowed MCP servers.""" - - allowed_mcp_servers: Final = await _get_allowed_mcp_servers( + allowed: Final = await _get_allowed_mcp_servers( user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip, ) - all_resources: Final[list[Resource]] = [] - for server in allowed_mcp_servers: - if server is None: - continue - + async def fetch(server: MCPServer) -> list[Resource]: server_auth_header, extra_headers = _prepare_mcp_server_headers( server=server, mcp_server_auth_headers=mcp_server_auth_headers, @@ -1339,28 +1325,19 @@ async def _get_resources_from_mcp_servers( oauth2_headers=oauth2_headers, raw_headers=raw_headers, user_api_key_auth=user_api_key_auth, - scope_servers=allowed_mcp_servers, + scope_servers=allowed, + ) + return await global_mcp_server_manager.get_resources_from_server( + server=server, + user_api_key_auth=user_api_key_auth, + mcp_auth_header=server_auth_header, + extra_headers=extra_headers, + add_prefix=True, + raw_headers=raw_headers, + client_ip=client_ip, ) - try: - resources = await global_mcp_server_manager.get_resources_from_server( - server=server, - user_api_key_auth=user_api_key_auth, - mcp_auth_header=server_auth_header, - extra_headers=extra_headers, - add_prefix=True, # Always add server prefix - raw_headers=raw_headers, - client_ip=client_ip, - ) - all_resources.extend(resources) - - verbose_logger.debug("Successfully fetched %s resources from server %s", len(resources), server.name) - except Exception as e: - verbose_logger.exception("Error getting resources from server %s: %s", server.name, e) - - verbose_logger.info("Successfully fetched %s resources total from all MCP servers", len(all_resources)) - - return all_resources + return await _collect_mcp_listing(tuple(server for server in allowed if server is not None), fetch) async def _get_resource_templates_from_mcp_servers( @@ -1372,19 +1349,13 @@ async def _get_resource_templates_from_mcp_servers( raw_headers: dict[str, str] | None = None, client_ip: str | None = None, ) -> list[ResourceTemplate]: - """Fetch resource templates from allowed MCP servers.""" - - allowed_mcp_servers: Final = await _get_allowed_mcp_servers( + allowed: Final = await _get_allowed_mcp_servers( user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip, ) - all_resource_templates: Final[list[ResourceTemplate]] = [] - for server in allowed_mcp_servers: - if server is None: - continue - + async def fetch(server: MCPServer) -> list[ResourceTemplate]: server_auth_header, extra_headers = _prepare_mcp_server_headers( server=server, mcp_server_auth_headers=mcp_server_auth_headers, @@ -1392,38 +1363,19 @@ async def _get_resource_templates_from_mcp_servers( oauth2_headers=oauth2_headers, raw_headers=raw_headers, user_api_key_auth=user_api_key_auth, - scope_servers=allowed_mcp_servers, + scope_servers=allowed, + ) + return await global_mcp_server_manager.get_resource_templates_from_server( + server=server, + user_api_key_auth=user_api_key_auth, + mcp_auth_header=server_auth_header, + extra_headers=extra_headers, + add_prefix=True, + raw_headers=raw_headers, + client_ip=client_ip, ) - try: - resource_templates = await global_mcp_server_manager.get_resource_templates_from_server( - server=server, - user_api_key_auth=user_api_key_auth, - mcp_auth_header=server_auth_header, - extra_headers=extra_headers, - add_prefix=True, # Always add server prefix - raw_headers=raw_headers, - client_ip=client_ip, - ) - all_resource_templates.extend(resource_templates) - verbose_logger.debug( - "Successfully fetched %s resource templates from server %s", - len(resource_templates), - server.name, - ) - except Exception as e: - verbose_logger.exception( - "Error getting resource templates from server %s: %s", - server.name, - str(e), - ) - - verbose_logger.info( - "Successfully fetched %s resource templates total from all MCP servers", - len(all_resource_templates), - ) - - return all_resource_templates + return await _collect_mcp_listing(tuple(server for server in allowed if server is not None), fetch) async def filter_tools_by_key_team_permissions( @@ -1539,6 +1491,8 @@ async def _list_mcp_prompts( client_ip=client_ip, ) verbose_logger.debug("Successfully fetched %s prompts from managed MCP servers", len(managed_prompts)) + except MCPUpstreamAuthError: + raise except Exception as e: verbose_logger.exception("Error getting tools from managed MCP servers: %s", e) # Continue with empty managed tools list instead of failing completely @@ -1569,6 +1523,8 @@ async def _list_mcp_resources( client_ip=client_ip, ) verbose_logger.debug("Successfully fetched %s resources from managed MCP servers", len(managed_resources)) + except MCPUpstreamAuthError: + raise except Exception as e: verbose_logger.exception("Error getting resources from managed MCP servers: %s", e) @@ -1601,6 +1557,8 @@ async def _list_mcp_resource_templates( "Successfully fetched %s resource templates from managed MCP servers", len(managed_resource_templates), ) + except MCPUpstreamAuthError: + raise except Exception as e: verbose_logger.exception( "Error getting resource templates from managed MCP servers: %s", @@ -2721,6 +2679,9 @@ async def _execute_handle_list_tools( list_tools_log_source="mcp_protocol", client_ip=_client_ip, ) + auth_failure: Final = listing_auth_error(listing.outcomes) + if auth_failure is not None: + raise auth_failure verbose_logger.info("MCP list_tools - Successfully returned %s tools", len(listing.tools)) if not listing.outcomes: return ListToolsResult(tools=listing.tools) @@ -2728,6 +2689,8 @@ async def _execute_handle_list_tools( SERVER_OUTCOMES_META_KEY: {key: outcome_wire_value(outcome) for key, outcome in listing.outcomes.items()} } return ListToolsResult.model_validate({"tools": listing.tools, "_meta": outcome_meta}) + except MCPUpstreamAuthError: + raise except HTTPException as e: from mcp.shared.exceptions import MCPError from mcp.types import INVALID_REQUEST @@ -2857,27 +2820,16 @@ async def _execute_mcp_server_tool_call( is_error=True, ) except HTTPException as e: + if e.status_code == 401 and e.headers and any(name.lower() == "www-authenticate" for name in e.headers): + raise verbose_logger.error("HTTPException in MCP tool call: %s", e) return CallToolResult( content=[TextContent(text=f"Error: {_http_detail_message(e.detail)}", type="text")], is_error=True, ) - except MCPUpstreamAuthError as e: - # The MCP session manager serializes handler exceptions as JSON-RPC errors, so a - # mid-session tool call cannot emit a raw 401 + WWW-Authenticate the way the REST - # call path and the connect-time preemptive check do. Return an explicit isError - # naming the upstream status (at info level, not a traceback) so the client still - # learns it must re-authenticate upstream and expected pass-through 401s don't spam. - verbose_logger.info("Upstream auth failure calling MCP tool: HTTP %s", e.status_code) - return CallToolResult( - content=[ - TextContent( - text=f"Error: upstream authentication required (HTTP {e.status_code})", - type="text", - ) - ], - is_error=True, - ) + except MCPUpstreamAuthError as exc: + verbose_logger.info("Upstream auth failure calling MCP tool: HTTP %s", exc.status_code) + raise except Exception as e: verbose_logger.exception("MCP mcp_server_tool_call - error: %s", e) return CallToolResult( @@ -2922,6 +2874,8 @@ async def _execute_list_prompts( ) verbose_logger.info("MCP list_prompts - Successfully returned %s prompts", len(prompts)) return ListPromptsResult(prompts=prompts) + except MCPUpstreamAuthError: + raise except Exception as e: verbose_logger.exception("Error in list_prompts endpoint: %s", e) # Return empty list instead of failing completely @@ -2991,6 +2945,8 @@ async def _execute_list_resources( ) verbose_logger.info("MCP list_resources - Successfully returned %s resources", len(resources)) return ListResourcesResult(resources=resources) + except MCPUpstreamAuthError: + raise except Exception as e: verbose_logger.exception("Error in list_resources endpoint: %s", e) return ListResourcesResult(resources=[]) # mutable-ok: MCP result payload @@ -3031,6 +2987,8 @@ async def _execute_list_resource_templates( "MCP list_resource_templates - Successfully returned %s resource templates", len(resource_templates) ) return ListResourceTemplatesResult(resource_templates=resource_templates) + except MCPUpstreamAuthError: + raise except Exception as e: verbose_logger.exception("Error in list_resource_templates endpoint: %s", e) return ListResourceTemplatesResult(resource_templates=[]) # mutable-ok: MCP result payload diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 555aebc7434..926aead9b7a 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -92,6 +92,62 @@ if TYPE_CHECKING: from mcp.server.session import ServerSession as _McpServerSession +_MCP_AUTH_RESPONSE_SCOPE_KEY: Final = "litellm_mcp_auth_response" + + +class MCPAuthResponse: + """Defer HTTP success until the SDK produces data, preserving late auth challenges.""" + + def __init__(self, send: Send) -> None: + self._send = send + self._start: Message | None = None + self._preamble: tuple[Message, ...] = () + self._committed = False + self._replaced = False + self._sse = False + self.challenge: HTTPException | None = None + + async def send(self, message: Message) -> None: + if self._replaced: + return + if message["type"] == "http.response.start" and message["status"] == 200: + self._start = message + self._sse = any( + name.lower() == b"content-type" and b"text/event-stream" in value + for name, value in message.get("headers", ()) + ) + return + if self._start is None or self._committed: + await self._send(message) + return + body: Final = message.get("body", b"") + if ( + self._sse + and message.get("more_body", False) + and not any(line.startswith(b"data:") and line[5:].strip() for line in body.splitlines()) + ): + if not body.startswith(b":"): + self._preamble = (*self._preamble, message) + return + self._committed = True + if self.challenge is not None: + self._replaced = True + response: Final = JSONResponse( + {"detail": self.challenge.detail}, + status_code=self.challenge.status_code, + headers=self.challenge.headers, + ) + await self._send( + {"type": "http.response.start", "status": response.status_code, "headers": response.raw_headers} + ) + await self._send({"type": "http.response.body", "body": response.body, "more_body": False}) + return + await self._send(self._start) + for preamble in self._preamble: + await self._send(preamble) + await self._send(message) + + _STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS: Final = 30 * 60 # Upper bound on concurrent stateful sessions a single caller may hold. Each # `initialize` creates a session that survives until the idle timeout, so @@ -530,6 +586,7 @@ if MCP_AVAILABLE: ListToolsResult, PaginatedRequestParams, ReadResourceRequestParams, + TextContent, ) from litellm.proxy._experimental.mcp_server.auth.litellm_auth_handler import ( @@ -806,38 +863,60 @@ if MCP_AVAILABLE: @contextlib.asynccontextmanager async def _legacy_operation_context(ctx: ServerRequestContext, *, trace: bool) -> AsyncGenerator[OperationContext]: - with contextlib.ExitStack() as cleanup: - cleanup.callback(active_mcp_request_ctx_var.reset, active_mcp_request_ctx_var.set(ctx)) - cleanup.callback(active_mcp_session_var.reset, active_mcp_session_var.set(ctx.session)) - if trace: - cleanup.callback( - _otel_reset_mcp_trace_carrier, _otel_set_mcp_trace_carrier(_mcp_meta_trace_carrier(ctx)) + try: + with contextlib.ExitStack() as cleanup: + cleanup.callback(active_mcp_request_ctx_var.reset, active_mcp_request_ctx_var.set(ctx)) + cleanup.callback(active_mcp_session_var.reset, active_mcp_session_var.set(ctx.session)) + if trace: + cleanup.callback( + _otel_reset_mcp_trace_carrier, _otel_set_mcp_trace_carrier(_mcp_meta_trace_carrier(ctx)) + ) + cleanup.callback( + _otel_reset_mcp_transport_span, + _otel_set_mcp_transport_span(_otel_transport_span_from_message(ctx)), + ) + cleanup.callback(_otel_reset_mcp_request_destinations, _otel_set_mcp_request_destinations(ctx)) + ( + auth, + token, + servers, + server_headers, + oauth_headers, + headers, + client_ip, + ) = await get_or_extract_auth_context() + yield operations.prepare_context( + auth, + token, + servers, + server_headers, + oauth_headers, + headers, + client_ip, + _mcp_proxy_mode.get(), + wire_compat_for(ctx.protocol_version), + ctx.protocol_version, ) - cleanup.callback( - _otel_reset_mcp_transport_span, _otel_set_mcp_transport_span(_otel_transport_span_from_message(ctx)) - ) - cleanup.callback(_otel_reset_mcp_request_destinations, _otel_set_mcp_request_destinations(ctx)) - ( - auth, - token, - servers, - server_headers, - oauth_headers, - headers, - client_ip, - ) = await get_or_extract_auth_context() - yield operations.prepare_context( - auth, - token, - servers, - server_headers, - oauth_headers, - headers, - client_ip, - _mcp_proxy_mode.get(), - wire_compat_for(ctx.protocol_version), - ctx.protocol_version, - ) + except (MCPUpstreamAuthError, HTTPException) as exc: + if isinstance(exc, HTTPException) and ( + exc.status_code != 401 + or not exc.headers + or not any(name.lower() == "www-authenticate" for name in exc.headers) + ): + raise + if isinstance(ctx.request, StarletteRequest): + response: Final = ctx.request.scope.get(_MCP_AUTH_RESPONSE_SCOPE_KEY) + if isinstance(response, MCPAuthResponse): + response.challenge = ( + exc.to_http_exception( + base_url=get_request_base_url(ctx.request), request_path=ctx.request.url.path + ) + if isinstance(exc, MCPUpstreamAuthError) + else exc + ) + raise MCPError( + code=INVALID_REQUEST, message=f"Upstream authorization failed (HTTP {exc.status_code})" + ) from exc async def handle_list_tools(ctx: ServerRequestContext, params: PaginatedRequestParams) -> ListToolsResult: try: @@ -896,9 +975,21 @@ if MCP_AVAILABLE: async def mcp_server_tool_call( ctx: ServerRequestContext, params: CallToolRequestParams ) -> CallToolResult | InputRequiredResult: - async with _legacy_operation_context(ctx, trace=True) as context: - return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute( - CallToolRequest(params=params), context + try: + async with _legacy_operation_context(ctx, trace=True) as context: + return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute( + CallToolRequest(params=params), context + ) + except MCPError as exc: + if not isinstance(exc.__cause__, (MCPUpstreamAuthError, HTTPException)): + raise + return CallToolResult( + content=[ + TextContent( + type="text", text=f"Error: upstream authentication required (HTTP {exc.__cause__.status_code})" + ) + ], + is_error=True, ) async def list_prompts(ctx: ServerRequestContext, params: PaginatedRequestParams) -> ListPromptsResult: @@ -909,6 +1000,8 @@ if MCP_AVAILABLE: return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute( ListPromptsRequest(params=params), context ) + except MCPError: + raise except Exception as exc: # noqa: BLE001 # preserve native listing fallback for ingress failures verbose_logger.exception("Error in list_prompts endpoint: %s", exc) return ListPromptsResult(prompts=[]) @@ -929,6 +1022,8 @@ if MCP_AVAILABLE: return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute( ListResourcesRequest(params=params), context ) + except MCPError: + raise except Exception as exc: # noqa: BLE001 # preserve native listing fallback for ingress failures verbose_logger.exception("Error in list_resources endpoint: %s", exc) return ListResourcesResult(resources=[]) @@ -943,6 +1038,8 @@ if MCP_AVAILABLE: return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute( ListResourceTemplatesRequest(params=params), context ) + except MCPError: + raise except Exception as exc: # noqa: BLE001 # preserve native listing fallback for ingress failures verbose_logger.exception("Error in list_resource_templates endpoint: %s", exc) return ListResourceTemplatesResult(resource_templates=[]) @@ -2248,7 +2345,12 @@ if MCP_AVAILABLE: scoped_server_endpoint=scoped_server_endpoint, is_initialize=is_initialize, ): - await target_manager.handle_request(scope, receive, local_send) + if request_method == "POST" and body and not is_initialize: + auth_response: Final = MCPAuthResponse(local_send) + scope[_MCP_AUTH_RESPONSE_SCOPE_KEY] = auth_response + await target_manager.handle_request(scope, receive, auth_response.send) + else: + await target_manager.handle_request(scope, receive, local_send) if use_stateful and session_id and scope.get("method") == "DELETE": _remove_stateful_session_tracking(session_id) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_list_outcomes.py b/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_list_outcomes.py index f951499e18f..fc02fed0945 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_list_outcomes.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_list_outcomes.py @@ -255,3 +255,57 @@ def test_pure_non_auth_response_still_classifies_upstream_error(): fault = classify_list_exception(exc) assert fault.tag == "upstream_error" assert fault.status_code == 502 + + +@pytest.mark.parametrize( + "outcomes,expected_status", + ( + ({}, None), + ({"empty": ServerListOk(tool_count=0)}, None), + ({"auth": ServerListFault(tag="auth_required", status_code=401)}, 401), + ({"denied": ServerListFault(tag="forbidden", status_code=403)}, 403), + ({"denied": ServerListFault(tag="forbidden", status_code=403), "auth": ServerListFault(tag="auth_required", status_code=401)}, 401), + ({"auth": ServerListFault(tag="auth_required", status_code=401), "empty": ServerListOk(tool_count=0)}, None), + ({"auth": ServerListFault(tag="auth_required", status_code=401), "healthy": ServerListOk(tool_count=2)}, None), + ({"auth": ServerListFault(tag="auth_required", status_code=401), "timeout": ServerListFault(tag="timeout")}, None), + ), +) +def test_listing_auth_failure_requires_every_server_to_be_blocked( + outcomes: dict[str, ServerListOk | ServerListFault], expected_status: int | None +) -> None: + from litellm.proxy._experimental.mcp_server.faults.list_outcomes import listing_auth_error + from typing import Final + failure: Final = listing_auth_error(outcomes) + assert (failure.status_code if failure else None) == expected_status + + +def test_classified_auth_challenge_is_preserved_without_serializing_it() -> None: + from typing import Final + from fastapi import HTTPException + from litellm.proxy._experimental.mcp_server.faults.list_outcomes import listing_auth_error + challenge: Final = 'Bearer resource_metadata="https://gateway/.well-known/oauth-protected-resource/mcp"' + fault: Final = classify_list_exception(HTTPException(401, headers={"WWW-Authenticate": challenge})) + failure: Final = listing_auth_error({"upstream": fault}) + assert failure is not None + assert failure.www_authenticate == challenge + assert failure.server_name == "upstream" + assert "www_authenticate" not in fault.model_dump() + assert challenge not in repr(fault) + assert outcome_wire_value(fault) == {"status": "auth_required", "http_status": 401} + + +def test_auth_recovery_uses_routable_server_name_without_exposing_it_in_outcomes() -> None: + from typing import Final + from litellm.proxy._experimental.mcp_server.faults.list_outcomes import listing_auth_error + fault: Final = classify_list_exception(MCPUpstreamAuthError(401, None, "upstream-internal")) + failure: Final = listing_auth_error({"short-prefix": fault}) + assert failure is not None + assert failure.server_name == "upstream-internal" + assert "server_name" not in fault.model_dump() + assert "upstream-internal" not in repr(fault) + assert outcome_wire_value(fault) == {"status": "auth_required", "http_status": 401} + + +def test_auth_challenge_traversal_preserves_typed_carrier() -> None: + from litellm.proxy._experimental.mcp_server.faults.list_outcomes import upstream_auth_challenge + assert upstream_auth_challenge(MCPUpstreamAuthError(401, "Bearer", "upstream")) == (401, "Bearer") diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py index 1909e3306a2..15aa089637a 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py @@ -3,6 +3,7 @@ from litellm.proxy._experimental.mcp_server import operations as mcp_operations import logging import sys +from typing import Final from unittest.mock import AsyncMock, MagicMock import httpx @@ -498,3 +499,358 @@ def test_passthrough_admission_recognizes_only_matching_authorization(oauth_head server = MCPServer(server_id="catalog", name="catalog", alias="catalog", transport=MCPTransport.http) assert _client_has_passthrough_authorization(server, oauth_headers, server_headers) is authorized + + +@pytest.mark.asyncio +async def test_listing_transport_preserves_auth_challenge_before_sse_success() -> None: + from starlette.exceptions import HTTPException + from litellm.proxy._experimental.mcp_server.server import MCPAuthResponse + + send: Final = AsyncMock() + response: Final = MCPAuthResponse(send) + await response.send( + {"type": "http.response.start", "status": 200, "headers": [(b"content-type", b"text/event-stream")]} + ) + await response.send({"type": "http.response.body", "body": b": ping\r\n\r\n", "more_body": True}) + response.challenge = HTTPException( + 401, "Unauthorized", headers={"WWW-Authenticate": 'Bearer resource_metadata="http://localhost/mcp-metadata"'} + ) + await response.send( + { + "type": "http.response.body", + "body": b'event: message\r\ndata: {"jsonrpc":"2.0","id":1,"result":{"tools":[]}}\r\n\r\n', + "more_body": True, + } + ) + await response.send({"type": "http.response.body", "body": b"", "more_body": False}) + sent: Final = tuple(call.args[0] for call in send.await_args_list) + assert [m["status"] for m in sent if m["type"] == "http.response.start"] == [401] + assert dict(sent[0]["headers"])[b"www-authenticate"] == b'Bearer resource_metadata="http://localhost/mcp-metadata"' + assert b'"tools"' not in b"".join(m.get("body", b"") for m in sent) + assert sent[-1].get("more_body", False) is False + + +@pytest.mark.asyncio +@pytest.mark.parametrize("status", [200, 400, 403]) +async def test_listing_transport_preserves_non_auth_responses(status: int) -> None: + from litellm.proxy._experimental.mcp_server.server import MCPAuthResponse + + send: Final = AsyncMock() + response: Final = MCPAuthResponse(send) + start: Final = { + "type": "http.response.start", + "status": status, + "headers": [(b"content-type", b"application/json")], + } + body: Final = { + "type": "http.response.body", + "body": b'{"jsonrpc":"2.0","id":1,"result":{"tools":[]}}', + "more_body": False, + } + await response.send(start) + await response.send(body) + assert tuple(call.args[0] for call in send.await_args_list) == (start, body) + + +@pytest.mark.asyncio +async def test_protocol_listing_does_not_report_success_when_every_server_requires_auth() -> None: + from unittest.mock import patch + from mcp.types import ListToolsRequest + from litellm.proxy._experimental.mcp_server import operations + from litellm.proxy._experimental.mcp_server.contracts import OperationContext + from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing, ServerListFault + from litellm.proxy._types import UserAPIKeyAuth + + context: Final = OperationContext(_caller=UserAPIKeyAuth(user_id="reader")) + listing: Final = AggregateToolListing([], {"github": ServerListFault(tag="auth_required", status_code=401)}) + with patch.object(operations, "_list_mcp_tools", AsyncMock(return_value=listing)): + with pytest.raises(MCPUpstreamAuthError) as caught: + await operations.GatewayOperations().execute(ListToolsRequest(), context) + assert caught.value.status_code == 401 + assert caught.value.server_name == "github" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("method", ("tools/list", "prompts/list", "resources/list", "resources/templates/list")) +@pytest.mark.parametrize("protocol", ("2025-06-18", "2025-11-25")) +@pytest.mark.parametrize("json_response", (False, True)) +@pytest.mark.parametrize("stateful", (False, True)) +@pytest.mark.parametrize("path", ("/mcp", "/github/mcp")) +async def test_streamable_http_listing_returns_late_oauth_challenge( + monkeypatch: pytest.MonkeyPatch, stateful: bool, path: str, json_response: bool, protocol: str, method: str +) -> None: + from litellm.proxy._experimental.mcp_server import server + from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing, ServerListFault + from litellm.proxy._types import UserAPIKeyAuth + + auth: Final = UserAPIKeyAuth(api_key="test-owner", user_id="test-user") + monkeypatch.setattr( + server, "extract_mcp_auth_context", AsyncMock(return_value=(auth, None, None, None, None, None)) + ) + monkeypatch.setattr(server, "_raise_preemptive_401_for_unauthenticated_servers", AsyncMock()) + monkeypatch.setattr(server, "_check_passthrough_upstream_auth", AsyncMock()) + listing: Final = AsyncMock( + return_value=AggregateToolListing( + [], + { + "github": ServerListFault( + tag="auth_required", + status_code=401, + www_authenticate='Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp/github"', + ) + }, + ) + ) + if method != "tools/list": + listing.side_effect = MCPUpstreamAuthError( + 401, 'Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp/github"', "github" + ) + list_function: Final = { + "tools/list": "_list_mcp_tools", + "prompts/list": "_list_mcp_prompts", + "resources/list": "_list_mcp_resources", + "resources/templates/list": "_list_mcp_resource_templates", + }[method] + monkeypatch.setattr(server.operations, list_function, listing) + monkeypatch.setattr(server.operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[])) + monkeypatch.setattr(server.operations, "_raise_if_initialize_grants_no_mcp_servers", AsyncMock()) + from mcp.server.streamable_http_manager import StreamableHTTPSessionManager + + monkeypatch.setattr( + server, + "session_manager_stateless", + StreamableHTTPSessionManager(app=server.server, stateless=True, json_response=json_response), + ) + monkeypatch.setattr( + server, + "session_manager_stateful", + StreamableHTTPSessionManager(app=server.server, stateless=False, json_response=json_response), + ) + await server.initialize_session_managers() + try: + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=server.app), base_url="http://gateway" + ) as client: + headers: Final = {"accept": "application/json, text/event-stream", "mcp-protocol-version": protocol} + if stateful: + initialized: Final = await client.post( + path, + headers=headers, + json={ + "jsonrpc": "2.0", + "id": 0, + "method": "initialize", + "params": { + "protocolVersion": protocol, + "capabilities": {}, + "clientInfo": {"name": "test", "version": "1"}, + }, + }, + ) + assert initialized.status_code == 200, initialized.text + client.headers["mcp-session-id"] = initialized.headers["mcp-session-id"] + notification: Final = await client.post( + path, headers=headers, json={"jsonrpc": "2.0", "method": "notifications/initialized"} + ) + assert notification.status_code == 202 + malformed: Final = await client.post( + path, headers={**headers, "content-type": "application/json"}, content=b"{" + ) + assert malformed.status_code == 400 + listing.assert_not_awaited() + response: Final = await client.post( + path, headers=headers, json={"jsonrpc": "2.0", "id": 1, "method": method} + ) + assert response.status_code == 401, response.text + assert ( + response.headers["www-authenticate"] + == 'Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp/github"' + ) + listing.assert_awaited_once() + finally: + await server.shutdown_session_managers() + + +@pytest.mark.asyncio +async def test_late_auth_failure_does_not_rewrite_committed_stream() -> None: + from fastapi import HTTPException + from litellm.proxy._experimental.mcp_server.server import MCPAuthResponse + + send: Final = AsyncMock() + response: Final = MCPAuthResponse(send) + start: Final = {"type": "http.response.start", "status": 200, "headers": [(b"content-type", b"text/event-stream")]} + progress: Final = { + "type": "http.response.body", + "body": b'data: {"method":"notifications/progress"}\n\n', + "more_body": True, + } + error: Final = { + "type": "http.response.body", + "body": b'data: {"error":{"code":-32600,"message":"Upstream authorization failed (HTTP 401)"}}\n\n', + "more_body": False, + } + await response.send(start) + await response.send(progress) + response.challenge = HTTPException(401, headers={"WWW-Authenticate": "Bearer"}) + await response.send(error) + assert tuple(call.args[0] for call in send.await_args_list) == (start, progress, error) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("kind", ("prompts", "resources", "resource_templates")) +async def test_optional_listing_propagates_auth_without_discarding_healthy_servers( + kind: str, monkeypatch: pytest.MonkeyPatch +) -> None: + from litellm.proxy._experimental.mcp_server import operations + + blocked: Final = _http_server("blocked", "blocked") + healthy: Final = _http_server("healthy", "healthy") + fetch: Final = AsyncMock(side_effect=MCPUpstreamAuthError(401, "Bearer", "blocked")) + manager: Final = MagicMock() + setattr(manager, f"get_{kind}_from_server", fetch) + monkeypatch.setattr(operations, "global_mcp_server_manager", manager) + monkeypatch.setattr(operations, "_prepare_mcp_server_headers", MagicMock(return_value=(None, None))) + monkeypatch.setattr(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[blocked])) + listing: Final = getattr(operations, f"_list_mcp_{kind}") + with pytest.raises(MCPUpstreamAuthError) as failure: + await listing() + assert failure.value.www_authenticate == "Bearer" + fetch.assert_awaited_once() + fetch.reset_mock(side_effect=True) + fetch.side_effect = [MCPUpstreamAuthError(401, "Bearer", "blocked"), []] + monkeypatch.setattr(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[blocked, healthy])) + assert await listing() == [] + assert fetch.await_count == 2 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "operation", + ( + "get_prompts_from_server", + "get_resources_from_server", + "get_resource_templates_from_server", + "get_prompt_from_server", + "read_resource_from_server", + ), +) +@pytest.mark.parametrize("status", (401, 403)) +async def test_manager_preserves_auth_failures_for_prompts_and_resources( + operation: str, status: int, monkeypatch: pytest.MonkeyPatch +) -> None: + from fastapi import HTTPException + from pydantic import AnyUrl + + manager: Final = MCPServerManager() + upstream: Final = _http_server("upstream", "upstream") + create: Final = AsyncMock(side_effect=HTTPException(status, headers={"WWW-Authenticate": "Bearer"})) + monkeypatch.setattr(manager, "_create_mcp_client", create) + kwargs: Final = ( + {"prompt_name": "example"} + if operation == "get_prompt_from_server" + else {"url": AnyUrl("https://example.com/resource")} + if operation == "read_resource_from_server" + else {} + ) + with pytest.raises(MCPUpstreamAuthError) as failure: + await getattr(manager, operation)(server=upstream, user_api_key_auth=None, **kwargs) + assert failure.value.status_code == status + assert failure.value.www_authenticate == "Bearer" + assert failure.value.server_name == "upstream" + create.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_optional_listing_preserves_cancellation() -> None: + import asyncio + from litellm.proxy._experimental.mcp_server.operations import _collect_mcp_listing + + fetch: Final = AsyncMock(side_effect=asyncio.CancelledError()) + with pytest.raises(asyncio.CancelledError): + await _collect_mcp_listing((_http_server("upstream", "upstream"),), fetch) + fetch.assert_awaited_once() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ("get_prompt_from_server", "read_resource_from_server")) +async def test_prompt_and_resource_calls_preserve_static_headers_and_non_auth_failures( + operation: str, monkeypatch: pytest.MonkeyPatch +) -> None: + from pydantic import AnyUrl + + manager: Final = MCPServerManager() + upstream: Final = _http_server("upstream", "upstream", static_headers={"x-upstream": "configured"}) + failure: Final = RuntimeError("Upstream unavailable") + create: Final = AsyncMock(side_effect=failure) + monkeypatch.setattr(manager, "_create_mcp_client", create) + kwargs: Final = ( + {"prompt_name": "example"} + if operation == "get_prompt_from_server" + else {"url": AnyUrl("https://example.com/resource")} + ) + with pytest.raises(RuntimeError) as caught: + await getattr(manager, operation)(server=upstream, user_api_key_auth=None, **kwargs) + assert caught.value is failure + assert create.await_args.kwargs["extra_headers"] == {"x-upstream": "configured"} + create.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_transport_preserves_sse_priming_event_on_success() -> None: + from litellm.proxy._experimental.mcp_server.server import MCPAuthResponse + + send: Final = AsyncMock() + response: Final = MCPAuthResponse(send) + start: Final = {"type": "http.response.start", "status": 200, "headers": [(b"content-type", b"text/event-stream")]} + priming: Final = {"type": "http.response.body", "body": b"id: resume-token\r\ndata: \r\n\r\n", "more_body": True} + tools: Final = { + "type": "http.response.body", + "body": b'data: {"jsonrpc":"2.0","id":1,"result":{"tools":[]}}\n\n', + "more_body": True, + } + await response.send(start) + await response.send(priming) + send.assert_not_awaited() + await response.send(tools) + assert tuple(call.args[0] for call in send.await_args_list) == (start, priming, tools) + + +@pytest.mark.asyncio +async def test_tool_call_preserves_resolver_http_challenge(monkeypatch: pytest.MonkeyPatch) -> None: + from fastapi import HTTPException + from mcp.types import CallToolRequest, CallToolRequestParams + from litellm.proxy._experimental.mcp_server import operations + from litellm.proxy._experimental.mcp_server.contracts import OperationContext + + challenge: Final = HTTPException(401, "Unauthorized", headers={"WWW-Authenticate": "Bearer"}) + call: Final = AsyncMock(side_effect=challenge) + monkeypatch.setattr(operations, "call_mcp_tool", call) + with pytest.raises(HTTPException) as failure: + await operations.GatewayOperations().execute( + CallToolRequest(params=CallToolRequestParams(name="upstream-tool", arguments={})), + OperationContext(_caller=None), + ) + assert failure.value is challenge + call.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_tool_handler_preserves_unrelated_protocol_errors( + _mcp_request_ctx, monkeypatch: pytest.MonkeyPatch +) -> None: + from mcp import MCPError + from mcp.types import CallToolRequestParams + from litellm.proxy._experimental.mcp_server import server + + failure: Final = MCPError(code=-32602, message="Invalid tool parameters") + execute: Final = AsyncMock(side_effect=failure) + gateway: Final = MagicMock() + gateway.execute = execute + monkeypatch.setattr(server.operations, "GatewayOperations", MagicMock(return_value=gateway)) + monkeypatch.setattr( + server, "get_or_extract_auth_context", AsyncMock(return_value=(None, None, None, None, None, None, None)) + ) + with pytest.raises(MCPError) as caught: + await server.mcp_server_tool_call(_mcp_request_ctx(), CallToolRequestParams(name="example", arguments={})) + assert caught.value is failure + execute.assert_awaited_once() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index ab00ec4da1e..da11b290a1d 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -1969,7 +1969,8 @@ async def test_mcp_routing_initialize_to_stateful_no_session_to_stateless(_mcp_r assert stateful_handle.await_count == (1 if stateful else 0) assert stateless_handle.await_count == (0 if stateful else 1) - observe_start.assert_awaited_once_with(0 if debug and method == "POST" else 1) + deferred: Final = method == "POST" and (debug or (bool(request_body) and not stateful)) + observe_start.assert_awaited_once_with(0 if deferred else 1) assert send.await_count == 2 assert send.call_args_list[0].args[0]["status"] == 200 assert send.call_args_list[1].args[0] == body diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index bb4e9e0e0f0..9f713c4175b 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -2047,9 +2047,8 @@ class TestMCPServerManager: assert mock_log.warning.called @pytest.mark.asyncio - async def test_call_non_passthrough_does_not_opt_into_raise_on_error(self): - """Non-client-forwarded auth types keep the default call_tool masking (raise_on_error stays - off), so this relay is scoped to the pass-through modes and cannot regress api_key/OBO calls.""" + async def test_call_static_auth_preserves_success_with_transport_errors_enabled(self): + """A static credential uses the same transport-auth error channel while preserving tool results.""" server = MCPServer( server_id="ak-call", name="ak-call-server", @@ -2076,7 +2075,7 @@ class TestMCPServerManager: ) assert result.is_error is False - assert mock_client.call_tool.call_args.kwargs.get("raise_on_error") is not True + assert mock_client.call_tool.call_args.kwargs.get("raise_on_error") is True def _token_exchange_server(self, server_id: str) -> "MCPServer": return MCPServer( @@ -6909,7 +6908,7 @@ class TestMCPServerManager: # Create mock client that tracks call_tool usage mock_client = AsyncMock() - async def mock_call_tool(params, host_progress_callback=None, allow_input_required=False): + async def mock_call_tool(params, host_progress_callback=None, allow_input_required=False, raise_on_error=False): # Return a mock CallToolResult result = MagicMock(spec=CallToolResult) result.content = [{"type": "text", "text": "Tool executed successfully"}] @@ -14141,6 +14140,7 @@ async def test_discovery_cache_bounds_detached_fetches_without_dropping_results( @pytest.mark.asyncio async def test_discovery_cache_tracks_resolved_credentials_across_workers() -> None: + from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError import respx from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import StaticHeaderAuth from litellm.proxy._experimental.mcp_server.outbound_credentials.resolver import UpstreamCredentialProvider @@ -14198,7 +14198,9 @@ async def test_discovery_cache_tracks_resolved_credentials_across_workers() -> N assert upstream.initializes == 4 source.token = None for manager in managers: - assert await manager.get_prompts_from_server(server, user) == [] + with pytest.raises(MCPUpstreamAuthError) as failure: + await manager.get_prompts_from_server(server, user) + assert failure.value.status_code == 401 assert upstream.initializes == 4