mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(mcp): preserve upstream authentication challenges on protocol requests
This commit is contained in:
parent
6d4ccf7e97
commit
97a662f2b2
8 changed files with 747 additions and 238 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue