fix(mcp): preserve upstream authentication challenges on protocol requests

This commit is contained in:
Joshua Valluru 2026-09-28 11:41:54 -07:00
parent 6d4ccf7e97
commit 97a662f2b2
8 changed files with 747 additions and 238 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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