This commit is contained in:
joshua-berri 2026-09-30 10:28:42 -07:00 • committed by GitHub
commit 920e80d44d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
20 changed files with 1721 additions and 500 deletions

View file

@ -6,7 +6,7 @@ import time
from collections.abc import AsyncIterator, Callable, Mapping
from contextlib import asynccontextmanager
from datetime import datetime, timezone
from typing import TYPE_CHECKING, Any, Final, Literal, Optional
from typing import TYPE_CHECKING, Annotated, Any, Final, Literal, Optional
from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
import httpx
@ -2133,6 +2133,7 @@ async def authorize_complete(
delivery: str | None = Form(None),
team_id: str | None = Form(None),
decision: str | None = Form(None),
selected_servers: Annotated[list[str] | None, Form(max_length=100)] = None,
) -> Response:
"""Finish an aggregate connect flow: mint the gateway authorization code for the
signed-in user and hand it back to the DCR client, by 303 redirect (default) or, for
@ -2150,6 +2151,7 @@ async def authorize_complete(
delivery=delivery,
team_id=team_id,
decision=decision,
selected_servers=tuple(selected_servers or ()),
lookup_vendor_credential=_vendor_credential_state,
lookup_server_reachability=_user_can_reach_mcp_server,
)

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 any(isinstance(outcome, ServerListOk) for outcome in outcomes.values()):
return None
name, outcome = next((entry for entry in blocked if entry[1].tag == "auth_required"), blocked[0])
return MCPUpstreamAuthError(
401 if outcome.tag == "auth_required" else 403, outcome.www_authenticate, outcome.server_name or name
)
def _iter_upstream_responses(exc: BaseException) -> Iterator[httpx.Response | httpx2.Response]:
"""Yield every upstream ``httpx``/``httpx2`` ``Response`` in the exception tree, in the shared traversal's deliberate
order (explicit causes first, ExceptionGroup members in raise order, the incidental
@ -87,9 +104,20 @@ def upstream_auth_challenge(exc: BaseException) -> tuple[int, str | None] | None
rides with it can never come from two different responses in the tree. Non-auth responses do not
end the scan: a causal 401 behind an unrelated 5xx must still be found, or the client never
receives the challenge it needs to re-authenticate."""
for response in _iter_upstream_responses(exc):
if response.status_code in (401, 403):
return response.status_code, response.headers.get("www-authenticate")
return next(
(challenge for current in iter_exception_tree(exc) if (challenge := _auth_challenge(current)) is not None), None
)
def _auth_challenge(exc: BaseException) -> tuple[int, str | None] | None:
if isinstance(exc, MCPUpstreamAuthError):
return exc.status_code, exc.www_authenticate
if isinstance(exc, HTTPException) and exc.status_code in (401, 403):
headers: Final = exc.headers or {}
return exc.status_code, headers.get("WWW-Authenticate") or headers.get("www-authenticate")
response: Final = getattr(exc, "response", None)
if isinstance(response, (httpx.Response, httpx2.Response)) and response.status_code in (401, 403):
return response.status_code, response.headers.get("www-authenticate")
return None
@ -104,17 +132,22 @@ def raise_classified_list_failure(
a classified fault. Every fetch site delegates here so the two channels cannot drift apart per
call site. ``suppress_challenge`` is for dcr_bridge servers, whose upstream challenge points
clients at the wrong protected-resource metadata and must never relay."""
auth: Final = upstream_auth_challenge(exc)
auth: Final = upstream_auth_error(exc, server_name, suppress_challenge=suppress_challenge)
if auth is not None:
status_code, challenge = auth
raise MCPUpstreamAuthError(
status_code=status_code,
www_authenticate=None if suppress_challenge else challenge,
server_name=server_name,
) from exc
raise auth from exc
raise MCPServerListError(classify_list_exception(exc), server_name) from exc
def upstream_auth_error(
exc: BaseException, server_name: str, *, suppress_challenge: bool = False
) -> MCPUpstreamAuthError | None:
auth: Final = upstream_auth_challenge(exc)
if auth is None:
return None
status_code, challenge = auth
return MCPUpstreamAuthError(status_code, None if suppress_challenge else challenge, server_name)
def classify_list_exception(exc: BaseException) -> ServerListFault:
"""Classify a per-server listing failure into exactly one outcome. Total: an exception this
function cannot recognize is the gateway's own fault (``internal``), never a re-raise."""
@ -122,17 +155,20 @@ def classify_list_exception(exc: BaseException) -> ServerListFault:
return exc.fault
if isinstance(exc, MCPUpstreamAuthError):
tag: Final = "forbidden" if exc.status_code == 403 else "auth_required"
return ServerListFault(tag=tag, status_code=exc.status_code)
return ServerListFault(
tag=tag, status_code=exc.status_code, www_authenticate=exc.www_authenticate, server_name=exc.server_name
)
if isinstance(exc, TimeoutError):
return ServerListFault(tag="timeout")
if isinstance(exc, ConnectionError):
return ServerListFault(tag="unreachable")
auth: Final = upstream_auth_challenge(exc)
if auth is not None:
status_code, _ = auth
status_code, challenge = auth
return ServerListFault(
tag="forbidden" if status_code == 403 else "auth_required",
status_code=status_code,
www_authenticate=challenge,
)
response: Final = _find_upstream_response(exc)
if response is not None:

View file

@ -851,6 +851,37 @@ async def describe_connect_flow(
)
async def _selected_connections_refusal(
flow: _ConnectFlow,
selected_servers: tuple[str, ...],
lookup_vendor_credential: LookupVendorCredential,
lookup_server_reachability: LookupServerReachability,
) -> Response | None:
if not selected_servers:
return _oauth_error(400, "invalid_request", "select and connect an MCP server before finishing")
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # proxy import cycle
global_mcp_server_manager,
)
for server in (
global_mcp_server_manager.get_mcp_server_by_id(server_id) for server_id in dict.fromkeys(selected_servers)
):
if server is None or not await lookup_server_reachability(flow.user_id, server.server_id):
return _oauth_error(400, "invalid_request", "a selected MCP server is no longer available")
if (
server.is_gateway_managed_oauth2
and global_mcp_server_manager.effective_oauth2_flow(server) != "client_credentials"
):
match await lookup_vendor_credential(flow.user_id, server.server_id):
case "unavailable":
return _oauth_error(503, "temporarily_unavailable", _DB_UNAVAILABLE_DESCRIPTION)
case "absent":
return _oauth_error(400, "invalid_request", "authorize the selected MCP servers before finishing")
case "present":
pass
return None
async def complete_connect_flow(
request: Request,
flow_handle: str,
@ -861,6 +892,7 @@ async def complete_connect_flow(
decision: str | None = None,
lookup_vendor_credential: LookupVendorCredential = _unavailable_vendor_credential,
lookup_server_reachability: LookupServerReachability = _unreachable_server,
selected_servers: tuple[str, ...] = (),
) -> Response:
"""Mint the code only after a deliberate POST by the sealed user.
@ -876,6 +908,12 @@ async def complete_connect_flow(
opened: Final = _open_flow_for(request, flow_handle, session_user_id, now)
if isinstance(opened, Response):
return opened
if decision != "deny" and opened.resource_server_id is None and opened.audience is None:
refusal: Final = await _selected_connections_refusal(
opened, selected_servers, lookup_vendor_credential, lookup_server_reachability
)
if refusal is not None:
return refusal
if decision != "deny":
described: Final = await _describe_opened_flow(opened, lookup_vendor_credential, lookup_server_reachability)
if isinstance(described, Response):

View file

@ -86,6 +86,7 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
ServerListFault,
raise_classified_list_failure,
upstream_auth_challenge,
upstream_auth_error,
)
from litellm.proxy._experimental.mcp_server.mcp_debug import describe_upstream_http_failure, record_auth_resolution
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import (
@ -4664,7 +4665,7 @@ class MCPServerManager:
return self._create_prefixed_prompts(items, server, add_prefix=add_prefix)
except Exception as error:
verbose_logger.warning("Failed to get prompts from server %s: %s", server.name, error)
return []
raise_classified_list_failure(error, server.name, suppress_challenge=server.is_dcr_bridge)
async def get_resources_from_server(
self,
@ -4710,7 +4711,7 @@ class MCPServerManager:
return self._create_prefixed_resources(items, server, add_prefix=add_prefix)
except Exception as error:
verbose_logger.warning("Failed to get resources from server %s: %s", server.name, error)
return []
raise_classified_list_failure(error, server.name, suppress_challenge=server.is_dcr_bridge)
async def get_resource_templates_from_server(
self,
@ -4756,7 +4757,7 @@ class MCPServerManager:
return self._create_prefixed_resource_templates(items, server, add_prefix=add_prefix)
except Exception as error:
verbose_logger.warning("Failed to get resource_templates from server %s: %s", server.name, error)
return []
raise_classified_list_failure(error, server.name, suppress_challenge=server.is_dcr_bridge)
async def read_resource_from_server(
self,
@ -4770,29 +4771,35 @@ class MCPServerManager:
) -> ReadResourceResult:
"""Read resource contents from a specific MCP server."""
verbose_logger.debug("Connecting to url: %s", server.url)
verbose_logger.info("read_resource_from_server for %s...", server.name)
try:
verbose_logger.debug("Connecting to url: %s", server.url)
verbose_logger.info("read_resource_from_server for %s...", server.name)
if server.static_headers:
if extra_headers is None:
extra_headers = {}
extra_headers.update(server.static_headers)
if server.static_headers:
if extra_headers is None:
extra_headers = {}
extra_headers.update(server.static_headers)
stdio_env: Final = self._build_stdio_env(server, raw_headers)
subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth)
stdio_env: Final = self._build_stdio_env(server, raw_headers)
subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth)
client: Final = await self._create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=extra_headers,
stdio_env=stdio_env,
subject_token=subject_token,
raw_headers=raw_headers,
client_ip=client_ip,
user_api_key_auth=user_api_key_auth,
)
client: Final = await self._create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=extra_headers,
stdio_env=stdio_env,
subject_token=subject_token,
raw_headers=raw_headers,
client_ip=client_ip,
user_api_key_auth=user_api_key_auth,
)
return await client.read_resource(url)
return await client.read_resource(url)
except Exception as exc:
auth_failure: Final = upstream_auth_error(exc, server.name, suppress_challenge=server.is_dcr_bridge)
if auth_failure is not None:
raise auth_failure from exc
raise
async def get_prompt_from_server(
self,
@ -4807,33 +4814,39 @@ class MCPServerManager:
) -> GetPromptResult:
"""Fetch a specific prompt definition from a single MCP server."""
verbose_logger.debug("Connecting to url: %s", server.url)
verbose_logger.info("get_prompt_from_server for %s...", server.name)
try:
verbose_logger.debug("Connecting to url: %s", server.url)
verbose_logger.info("get_prompt_from_server for %s...", server.name)
if server.static_headers:
if extra_headers is None:
extra_headers = {}
extra_headers.update(server.static_headers)
if server.static_headers:
if extra_headers is None:
extra_headers = {}
extra_headers.update(server.static_headers)
stdio_env: Final = self._build_stdio_env(server, raw_headers)
subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth)
stdio_env: Final = self._build_stdio_env(server, raw_headers)
subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth)
client: Final = await self._create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=extra_headers,
stdio_env=stdio_env,
subject_token=subject_token,
raw_headers=raw_headers,
client_ip=client_ip,
user_api_key_auth=user_api_key_auth,
)
client: Final = await self._create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=extra_headers,
stdio_env=stdio_env,
subject_token=subject_token,
raw_headers=raw_headers,
client_ip=client_ip,
user_api_key_auth=user_api_key_auth,
)
get_prompt_request_params: Final = GetPromptRequestParams(
name=prompt_name,
arguments=arguments,
)
return await client.get_prompt(get_prompt_request_params)
get_prompt_request_params: Final = GetPromptRequestParams(
name=prompt_name,
arguments=arguments,
)
return await client.get_prompt(get_prompt_request_params)
except Exception as exc:
auth_failure: Final = upstream_auth_error(exc, server.name, suppress_challenge=server.is_dcr_bridge)
if auth_failure is not None:
raise auth_failure from exc
raise
@staticmethod
def _is_same_authority_metadata_url(url: str, server_url: str) -> bool:
@ -6190,29 +6203,10 @@ class MCPServerManager:
tool_call_coro = _obo_call_tool_limited()
else:
# Scoped to the two client-forwarded token modes this stack introduced; legacy
# oauth2 + delegate_auth_to_upstream (is_oauth_passthrough) is being removed, so it is not
# added here even though the list path still relays for it.
relays_upstream_auth: Final = mcp_server.is_client_forwarded_token
server_label: Final = mcp_server.name or mcp_server.server_name or mcp_server.alias or ""
async def _call_tool_via_client(client, params):
async with self._limit_outbound_concurrency(mcp_server):
if not relays_upstream_auth:
return await client.call_tool(
params,
host_progress_callback=host_progress_callback,
allow_input_required=allow_input_required,
)
# The client-forwarded modes carry the caller's own upstream token, so an upstream
# 401 (expired/invalid token) is the caller's to resolve: relay it as
# MCPUpstreamAuthError so single-server REST callers turn it into a 401 +
# WWW-Authenticate and re-run the upstream OAuth flow. Only 401 is a re-auth signal
# (mirrors the list path and MCPUpstreamAuthError's contract); a 403 is a genuine
# authorization failure that re-auth won't fix, so it takes the non-auth branch and
# stays a visible warning. raise_on_error only re-raises transport failures
# (tool-level isError results are still returned normally); a non-auth failure keeps
# the same isError degradation the default path produces.
try:
return await client.call_tool(
params,
@ -6238,7 +6232,7 @@ class MCPServerManager:
_, www_authenticate = auth_info
raise MCPUpstreamAuthError(
status_code=401,
www_authenticate=www_authenticate,
www_authenticate=None if mcp_server.is_dcr_bridge else www_authenticate,
server_name=server_label,
) from e

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 (
@ -1238,6 +1240,26 @@ async def _get_tools_from_mcp_servers(
raise
_ListingItem = TypeVar("_ListingItem", Prompt, Resource, ResourceTemplate)
async def _collect_mcp_listing(
servers: Sequence[MCPServer], fetch: Callable[[MCPServer], Awaitable[list[_ListingItem]]]
) -> list[_ListingItem]:
async def fetch_one(server: MCPServer) -> tuple[list[_ListingItem], ServerOutcome]:
try:
items: Final = await fetch(server)
return items, ServerListOk(tool_count=len(items))
except Exception as exc:
return [], classify_list_exception(exc)
results: Final = await asyncio.gather(*(fetch_one(server) for server in servers))
failure: Final = listing_auth_error({server.server_id: result[1] for server, result in zip(servers, results)})
if failure is not None:
raise failure
return list(chain.from_iterable(items for items, _ in results))
async def _get_prompts_from_mcp_servers(
user_api_key_auth: UserAPIKeyAuth | None,
mcp_auth_header: str | None,
@ -1247,32 +1269,13 @@ async def _get_prompts_from_mcp_servers(
raw_headers: dict[str, str] | None = None,
client_ip: str | None = None,
) -> list[Prompt]:
"""
Helper method to fetch prompt from MCP servers based on server filtering criteria.
Args:
user_api_key_auth: User authentication info for access control
mcp_auth_header: Optional auth header for MCP server (deprecated)
mcp_servers: Optional list of server names/aliases to filter by
mcp_server_auth_headers: Optional dict of server-specific auth headers
oauth2_headers: Optional dict of oauth2 headers
Returns:
List[Prompt]: Combined list of prompts from filtered servers
"""
allowed_mcp_servers: Final = await _get_allowed_mcp_servers(
allowed: Final = await _get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_servers=mcp_servers,
client_ip=client_ip,
)
# Get prompts from each allowed server
all_prompts: Final = []
for server in allowed_mcp_servers:
if server is None:
continue
async def fetch(server: MCPServer) -> list[Prompt]:
server_auth_header, extra_headers = _prepare_mcp_server_headers(
server=server,
mcp_server_auth_headers=mcp_server_auth_headers,
@ -1280,30 +1283,19 @@ async def _get_prompts_from_mcp_servers(
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
scope_servers=allowed_mcp_servers,
scope_servers=allowed,
)
return await global_mcp_server_manager.get_prompts_from_server(
server=server,
user_api_key_auth=user_api_key_auth,
mcp_auth_header=server_auth_header,
extra_headers=extra_headers,
add_prefix=True,
raw_headers=raw_headers,
client_ip=client_ip,
)
try:
prompts = await global_mcp_server_manager.get_prompts_from_server(
server=server,
user_api_key_auth=user_api_key_auth,
mcp_auth_header=server_auth_header,
extra_headers=extra_headers,
add_prefix=True, # Always add server prefix
raw_headers=raw_headers,
client_ip=client_ip,
)
all_prompts.extend(prompts)
verbose_logger.debug("Successfully fetched %s prompts from server %s", len(prompts), server.name)
except Exception as e:
verbose_logger.exception("Error getting prompts from server %s: %s", server.name, e)
# Continue with other servers instead of failing completely
verbose_logger.info("Successfully fetched %s prompts total from all MCP servers", len(all_prompts))
return all_prompts
return await _collect_mcp_listing(tuple(server for server in allowed if server is not None), fetch)
async def _get_resources_from_mcp_servers(
@ -1315,19 +1307,13 @@ async def _get_resources_from_mcp_servers(
raw_headers: dict[str, str] | None = None,
client_ip: str | None = None,
) -> list[Resource]:
"""Fetch resources from allowed MCP servers."""
allowed_mcp_servers: Final = await _get_allowed_mcp_servers(
allowed: Final = await _get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_servers=mcp_servers,
client_ip=client_ip,
)
all_resources: Final[list[Resource]] = []
for server in allowed_mcp_servers:
if server is None:
continue
async def fetch(server: MCPServer) -> list[Resource]:
server_auth_header, extra_headers = _prepare_mcp_server_headers(
server=server,
mcp_server_auth_headers=mcp_server_auth_headers,
@ -1335,28 +1321,19 @@ async def _get_resources_from_mcp_servers(
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
scope_servers=allowed_mcp_servers,
scope_servers=allowed,
)
return await global_mcp_server_manager.get_resources_from_server(
server=server,
user_api_key_auth=user_api_key_auth,
mcp_auth_header=server_auth_header,
extra_headers=extra_headers,
add_prefix=True,
raw_headers=raw_headers,
client_ip=client_ip,
)
try:
resources = await global_mcp_server_manager.get_resources_from_server(
server=server,
user_api_key_auth=user_api_key_auth,
mcp_auth_header=server_auth_header,
extra_headers=extra_headers,
add_prefix=True, # Always add server prefix
raw_headers=raw_headers,
client_ip=client_ip,
)
all_resources.extend(resources)
verbose_logger.debug("Successfully fetched %s resources from server %s", len(resources), server.name)
except Exception as e:
verbose_logger.exception("Error getting resources from server %s: %s", server.name, e)
verbose_logger.info("Successfully fetched %s resources total from all MCP servers", len(all_resources))
return all_resources
return await _collect_mcp_listing(tuple(server for server in allowed if server is not None), fetch)
async def _get_resource_templates_from_mcp_servers(
@ -1368,19 +1345,13 @@ async def _get_resource_templates_from_mcp_servers(
raw_headers: dict[str, str] | None = None,
client_ip: str | None = None,
) -> list[ResourceTemplate]:
"""Fetch resource templates from allowed MCP servers."""
allowed_mcp_servers: Final = await _get_allowed_mcp_servers(
allowed: Final = await _get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_servers=mcp_servers,
client_ip=client_ip,
)
all_resource_templates: Final[list[ResourceTemplate]] = []
for server in allowed_mcp_servers:
if server is None:
continue
async def fetch(server: MCPServer) -> list[ResourceTemplate]:
server_auth_header, extra_headers = _prepare_mcp_server_headers(
server=server,
mcp_server_auth_headers=mcp_server_auth_headers,
@ -1388,38 +1359,19 @@ async def _get_resource_templates_from_mcp_servers(
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
scope_servers=allowed_mcp_servers,
scope_servers=allowed,
)
return await global_mcp_server_manager.get_resource_templates_from_server(
server=server,
user_api_key_auth=user_api_key_auth,
mcp_auth_header=server_auth_header,
extra_headers=extra_headers,
add_prefix=True,
raw_headers=raw_headers,
client_ip=client_ip,
)
try:
resource_templates = await global_mcp_server_manager.get_resource_templates_from_server(
server=server,
user_api_key_auth=user_api_key_auth,
mcp_auth_header=server_auth_header,
extra_headers=extra_headers,
add_prefix=True, # Always add server prefix
raw_headers=raw_headers,
client_ip=client_ip,
)
all_resource_templates.extend(resource_templates)
verbose_logger.debug(
"Successfully fetched %s resource templates from server %s",
len(resource_templates),
server.name,
)
except Exception as e:
verbose_logger.exception(
"Error getting resource templates from server %s: %s",
server.name,
str(e),
)
verbose_logger.info(
"Successfully fetched %s resource templates total from all MCP servers",
len(all_resource_templates),
)
return all_resource_templates
return await _collect_mcp_listing(tuple(server for server in allowed if server is not None), fetch)
async def filter_tools_by_key_team_permissions(
@ -1535,6 +1487,8 @@ async def _list_mcp_prompts(
client_ip=client_ip,
)
verbose_logger.debug("Successfully fetched %s prompts from managed MCP servers", len(managed_prompts))
except MCPUpstreamAuthError:
raise
except Exception as e:
verbose_logger.exception("Error getting tools from managed MCP servers: %s", e)
# Continue with empty managed tools list instead of failing completely
@ -1565,6 +1519,8 @@ async def _list_mcp_resources(
client_ip=client_ip,
)
verbose_logger.debug("Successfully fetched %s resources from managed MCP servers", len(managed_resources))
except MCPUpstreamAuthError:
raise
except Exception as e:
verbose_logger.exception("Error getting resources from managed MCP servers: %s", e)
@ -1597,6 +1553,8 @@ async def _list_mcp_resource_templates(
"Successfully fetched %s resource templates from managed MCP servers",
len(managed_resource_templates),
)
except MCPUpstreamAuthError:
raise
except Exception as e:
verbose_logger.exception(
"Error getting resource templates from managed MCP servers: %s",
@ -2717,6 +2675,9 @@ async def _execute_handle_list_tools(
list_tools_log_source="mcp_protocol",
client_ip=_client_ip,
)
auth_failure: Final = listing_auth_error(listing.outcomes)
if auth_failure is not None:
raise auth_failure
verbose_logger.info("MCP list_tools - Successfully returned %s tools", len(listing.tools))
if not listing.outcomes:
return ListToolsResult(tools=listing.tools)
@ -2724,6 +2685,8 @@ async def _execute_handle_list_tools(
SERVER_OUTCOMES_META_KEY: {key: outcome_wire_value(outcome) for key, outcome in listing.outcomes.items()}
}
return ListToolsResult.model_validate({"tools": listing.tools, "_meta": outcome_meta})
except MCPUpstreamAuthError:
raise
except HTTPException as e:
from mcp.shared.exceptions import MCPError
from mcp.types import INVALID_REQUEST
@ -2853,27 +2816,16 @@ async def _execute_mcp_server_tool_call(
is_error=True,
)
except HTTPException as e:
if e.status_code == 401 and e.headers and any(name.lower() == "www-authenticate" for name in e.headers):
raise
verbose_logger.error("HTTPException in MCP tool call: %s", e)
return CallToolResult(
content=[TextContent(text=f"Error: {_http_detail_message(e.detail)}", type="text")],
is_error=True,
)
except MCPUpstreamAuthError as e:
# The MCP session manager serializes handler exceptions as JSON-RPC errors, so a
# mid-session tool call cannot emit a raw 401 + WWW-Authenticate the way the REST
# call path and the connect-time preemptive check do. Return an explicit isError
# naming the upstream status (at info level, not a traceback) so the client still
# learns it must re-authenticate upstream and expected pass-through 401s don't spam.
verbose_logger.info("Upstream auth failure calling MCP tool: HTTP %s", e.status_code)
return CallToolResult(
content=[
TextContent(
text=f"Error: upstream authentication required (HTTP {e.status_code})",
type="text",
)
],
is_error=True,
)
except MCPUpstreamAuthError as exc:
verbose_logger.info("Upstream auth failure calling MCP tool: HTTP %s", exc.status_code)
raise
except Exception as e:
verbose_logger.exception("MCP mcp_server_tool_call - error: %s", e)
return CallToolResult(
@ -2918,6 +2870,8 @@ async def _execute_list_prompts(
)
verbose_logger.info("MCP list_prompts - Successfully returned %s prompts", len(prompts))
return ListPromptsResult(prompts=prompts)
except MCPUpstreamAuthError:
raise
except Exception as e:
verbose_logger.exception("Error in list_prompts endpoint: %s", e)
# Return empty list instead of failing completely
@ -2987,6 +2941,8 @@ async def _execute_list_resources(
)
verbose_logger.info("MCP list_resources - Successfully returned %s resources", len(resources))
return ListResourcesResult(resources=resources)
except MCPUpstreamAuthError:
raise
except Exception as e:
verbose_logger.exception("Error in list_resources endpoint: %s", e)
return ListResourcesResult(resources=[]) # mutable-ok: MCP result payload
@ -3027,6 +2983,8 @@ async def _execute_list_resource_templates(
"MCP list_resource_templates - Successfully returned %s resource templates", len(resource_templates)
)
return ListResourceTemplatesResult(resource_templates=resource_templates)
except MCPUpstreamAuthError:
raise
except Exception as e:
verbose_logger.exception("Error in list_resource_templates endpoint: %s", e)
return ListResourceTemplatesResult(resource_templates=[]) # mutable-ok: MCP result payload

View file

@ -34,6 +34,7 @@ from litellm.llms.custom_httpx.http_handler import (
)
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
_gateway_dcr_challenge,
_is_mcp_admitted_user_subject,
)
from litellm.proxy._experimental.mcp_server.client_allowlist import (
@ -92,6 +93,61 @@ if TYPE_CHECKING:
from mcp.server.session import ServerSession as _McpServerSession
_MCP_AUTH_RESPONSE_SCOPE_KEY: Final = "litellm_mcp_auth_response"
class MCPAuthResponse:
"""Defer HTTP success until the SDK produces data, preserving late auth challenges."""
def __init__(self, send: Send) -> None:
self._send = send
self._start: Message | None = None
self._preamble: tuple[Message, ...] = ()
self._committed = False
self._replaced = False
self._sse = False
self.challenge: HTTPException | None = None
async def send(self, message: Message) -> None:
if self._replaced:
return
if message["type"] == "http.response.start" and message["status"] == 200:
self._start = message
self._sse = any(
name.lower() == b"content-type" and b"text/event-stream" in value
for name, value in message.get("headers", ())
)
return
if self._start is None or self._committed:
await self._send(message)
return
body: Final = message.get("body", b"")
if (
self._sse
and message.get("more_body", False)
and not any(line.startswith(b"data:") and line[5:].strip() for line in body.splitlines())
):
self._preamble = (*self._preamble, message)
return
self._committed = True
if self.challenge is not None:
self._replaced = True
response: Final = JSONResponse(
{"detail": self.challenge.detail},
status_code=self.challenge.status_code,
headers=self.challenge.headers,
)
await self._send(
{"type": "http.response.start", "status": response.status_code, "headers": response.raw_headers}
)
await self._send({"type": "http.response.body", "body": response.body, "more_body": False})
return
await self._send(self._start)
for preamble in self._preamble:
await self._send(preamble)
await self._send(message)
_STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS: Final = 30 * 60
# Upper bound on concurrent stateful sessions a single caller may hold. Each
# `initialize` creates a session that survives until the idle timeout, so
@ -530,6 +586,7 @@ if MCP_AVAILABLE:
ListToolsResult,
PaginatedRequestParams,
ReadResourceRequestParams,
TextContent,
)
from litellm.proxy._experimental.mcp_server.auth.litellm_auth_handler import (
@ -806,38 +863,60 @@ if MCP_AVAILABLE:
@contextlib.asynccontextmanager
async def _legacy_operation_context(ctx: ServerRequestContext, *, trace: bool) -> AsyncGenerator[OperationContext]:
with contextlib.ExitStack() as cleanup:
cleanup.callback(active_mcp_request_ctx_var.reset, active_mcp_request_ctx_var.set(ctx))
cleanup.callback(active_mcp_session_var.reset, active_mcp_session_var.set(ctx.session))
if trace:
cleanup.callback(
_otel_reset_mcp_trace_carrier, _otel_set_mcp_trace_carrier(_mcp_meta_trace_carrier(ctx))
try:
with contextlib.ExitStack() as cleanup:
cleanup.callback(active_mcp_request_ctx_var.reset, active_mcp_request_ctx_var.set(ctx))
cleanup.callback(active_mcp_session_var.reset, active_mcp_session_var.set(ctx.session))
if trace:
cleanup.callback(
_otel_reset_mcp_trace_carrier, _otel_set_mcp_trace_carrier(_mcp_meta_trace_carrier(ctx))
)
cleanup.callback(
_otel_reset_mcp_transport_span,
_otel_set_mcp_transport_span(_otel_transport_span_from_message(ctx)),
)
cleanup.callback(_otel_reset_mcp_request_destinations, _otel_set_mcp_request_destinations(ctx))
(
auth,
token,
servers,
server_headers,
oauth_headers,
headers,
client_ip,
) = await get_or_extract_auth_context()
yield operations.prepare_context(
auth,
token,
servers,
server_headers,
oauth_headers,
headers,
client_ip,
_mcp_proxy_mode.get(),
wire_compat_for(ctx.protocol_version),
ctx.protocol_version,
)
cleanup.callback(
_otel_reset_mcp_transport_span, _otel_set_mcp_transport_span(_otel_transport_span_from_message(ctx))
)
cleanup.callback(_otel_reset_mcp_request_destinations, _otel_set_mcp_request_destinations(ctx))
(
auth,
token,
servers,
server_headers,
oauth_headers,
headers,
client_ip,
) = await get_or_extract_auth_context()
yield operations.prepare_context(
auth,
token,
servers,
server_headers,
oauth_headers,
headers,
client_ip,
_mcp_proxy_mode.get(),
wire_compat_for(ctx.protocol_version),
ctx.protocol_version,
)
except (MCPUpstreamAuthError, HTTPException) as exc:
if isinstance(exc, HTTPException) and (
exc.status_code != 401
or not exc.headers
or not any(name.lower() == "www-authenticate" for name in exc.headers)
):
raise
if isinstance(ctx.request, StarletteRequest):
response: Final = ctx.request.scope.get(_MCP_AUTH_RESPONSE_SCOPE_KEY)
if isinstance(response, MCPAuthResponse):
response.challenge = (
exc.to_http_exception(
base_url=get_request_base_url(ctx.request), request_path=ctx.request.url.path
)
if isinstance(exc, MCPUpstreamAuthError)
else exc
)
raise MCPError(
code=INVALID_REQUEST, message=f"Upstream authorization failed (HTTP {exc.status_code})"
) from exc
async def handle_list_tools(ctx: ServerRequestContext, params: PaginatedRequestParams) -> ListToolsResult:
try:
@ -896,9 +975,21 @@ if MCP_AVAILABLE:
async def mcp_server_tool_call(
ctx: ServerRequestContext, params: CallToolRequestParams
) -> CallToolResult | InputRequiredResult:
async with _legacy_operation_context(ctx, trace=True) as context:
return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute(
CallToolRequest(params=params), context
try:
async with _legacy_operation_context(ctx, trace=True) as context:
return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute(
CallToolRequest(params=params), context
)
except MCPError as exc:
if not isinstance(exc.__cause__, (MCPUpstreamAuthError, HTTPException)):
raise
return CallToolResult(
content=[
TextContent(
type="text", text=f"Error: upstream authentication required (HTTP {exc.__cause__.status_code})"
)
],
is_error=True,
)
async def list_prompts(ctx: ServerRequestContext, params: PaginatedRequestParams) -> ListPromptsResult:
@ -909,6 +1000,8 @@ if MCP_AVAILABLE:
return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute(
ListPromptsRequest(params=params), context
)
except MCPError:
raise
except Exception as exc: # noqa: BLE001 # preserve native listing fallback for ingress failures
verbose_logger.exception("Error in list_prompts endpoint: %s", exc)
return ListPromptsResult(prompts=[])
@ -929,6 +1022,8 @@ if MCP_AVAILABLE:
return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute(
ListResourcesRequest(params=params), context
)
except MCPError:
raise
except Exception as exc: # noqa: BLE001 # preserve native listing fallback for ingress failures
verbose_logger.exception("Error in list_resources endpoint: %s", exc)
return ListResourcesResult(resources=[])
@ -943,6 +1038,8 @@ if MCP_AVAILABLE:
return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute(
ListResourceTemplatesRequest(params=params), context
)
except MCPError:
raise
except Exception as exc: # noqa: BLE001 # preserve native listing fallback for ingress failures
verbose_logger.exception("Error in list_resource_templates endpoint: %s", exc)
return ListResourceTemplatesResult(resource_templates=[])
@ -1580,6 +1677,151 @@ if MCP_AVAILABLE:
)
return user_api_key_auth.model_copy(update={"object_permission": updated_op, "mcp_toolset_id": toolset_id})
async def _server_auth_challenge(
configured_server: MCPServer,
server_name: str,
scope: Scope,
mcp_servers: list[str],
oauth2_headers: dict[str, str] | None,
mcp_server_auth_headers: dict[str, dict[str, str]] | None,
user_api_key_auth: UserAPIKeyAuth | None,
client_ip: str | None,
raw_headers: Mapping[str, str] | None,
) -> HTTPException | None:
try:
if configured_server.auth_type == MCPAuth.oauth2 and configured_server.oauth2_flow == "client_credentials":
return None
server: Final = await operations.global_mcp_server_manager.ensure_oauth_metadata_discovered(
configured_server
)
if server.auth_type == MCPAuth.oauth2:
if MCPServerManager.effective_oauth2_flow(server) == "client_credentials":
return None
if getattr(server, "delegate_auth_to_upstream", False) is not True:
if await operations.global_mcp_server_manager.has_user_oauth_token(server, user_api_key_auth):
return None
if _is_mcp_admitted_user_subject(user_api_key_auth):
return HTTPException(
status_code=401,
detail="Unauthorized",
headers={
"www-authenticate": get_passthrough_www_authenticate(
scope=scope,
server_name=server_name,
)
},
)
request: Final = StarletteRequest(scope)
base_url: Final = get_request_base_url(request)
_path: Final = get_route_relative_request_path(scope)
as_metadata_root: Final = (
f"{base_url}/.well-known/oauth-authorization-server{well_known_root_suffix()}"
)
as_url: Final = (
f"{as_metadata_root}/mcp/{server_name}"
if _path.startswith(f"/mcp/{server_name}")
else f"{as_metadata_root}/{server_name}"
)
authorization_uri: Final = f'Bearer authorization_uri="{as_url}"'
return HTTPException(
status_code=401,
detail="Unauthorized",
headers={"www-authenticate": authorization_uri},
)
if not oauth2_headers:
return HTTPException(
status_code=401,
detail="Unauthorized",
headers={
"www-authenticate": get_passthrough_www_authenticate(scope=scope, server_name=server_name)
},
)
return None
if server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers:
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 # lazy: adapter pulls MCP subgraph
raise_token_exchange_challenge,
)
from litellm.proxy.middleware.per_request_root_path_middleware import ( # noqa: PLC0415 # lazy: middleware imports proxy utils
get_request_root_path,
)
raise_token_exchange_challenge(server, root_path=get_request_root_path())
if len(mcp_servers) == 1 and server.server_id in frozenset(
allowed.server_id
for allowed in await operations._get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip
)
):
await operations.global_mcp_server_manager.preflight_token_exchange(
server=server,
oauth2_headers=oauth2_headers,
user_api_key_auth=user_api_key_auth,
raw_headers=raw_headers,
)
if server.is_oauth_passthrough and not operations._client_has_passthrough_authorization(
server, oauth2_headers, mcp_server_auth_headers
):
return HTTPException(
status_code=401,
detail="Unauthorized",
headers={
"www-authenticate": get_passthrough_www_authenticate(scope=scope, server_name=server_name)
},
)
if (
server.is_oauth_delegate
and len(mcp_servers) == 1
and _get_forwarded_auth_from_scope(scope) is None
and not operations._client_has_per_server_auth_header(server, mcp_server_auth_headers)
):
return HTTPException(
status_code=401,
detail="Unauthorized",
headers={
"www-authenticate": get_passthrough_www_authenticate(scope=scope, server_name=server_name)
},
)
if (
server.is_true_passthrough
and len(mcp_servers) == 1
and not _scope_has_authorization_header(scope)
and not operations._client_has_per_server_auth_header(server, mcp_server_auth_headers)
):
if server.is_dcr_bridge:
return HTTPException(
status_code=401,
detail="Unauthorized",
headers={
"www-authenticate": get_passthrough_www_authenticate(
scope=scope,
server_name=server_name,
)
},
)
upstream_status, upstream_www_authenticate = await _probe_upstream_auth(server.url or "", "")
if upstream_status == 401 and upstream_www_authenticate:
return HTTPException(
status_code=401,
detail="Unauthorized",
headers={"www-authenticate": upstream_www_authenticate},
)
except HTTPException as exc:
if exc.status_code != 401:
raise
return exc
return None
async def _raise_preemptive_401_for_unauthenticated_servers(
scope: Scope,
mcp_servers: list[str] | None,
@ -1601,208 +1843,71 @@ if MCP_AVAILABLE:
excludes a passthrough server is not pushed into an OAuth flow for
a server it will be 403'd on immediately after authentication.
"""
for server_name in mcp_servers or []:
server = operations.global_mcp_server_manager.get_mcp_server_by_name(server_name, client_ip=client_ip)
if server is not None and allowed_server_ids is not None and server.server_id not in allowed_server_ids:
# Caller's narrowed scope excludes this server — skip the
# preemptive challenge and let downstream authorization
# return 403.
continue
if server is not None and server.auth_type == MCPAuth.oauth2 and server.oauth2_flow == "client_credentials":
# Stamped M2M: the challenge decision below never reads discovered
# metadata, so deferred-discovery failures must not 503 this loop.
# Unstamped rows stay on the discover-first path because filling
# authorization_url/token_url can change their inferred flow.
continue
if server is not None:
server = await operations.global_mcp_server_manager.ensure_oauth_metadata_discovered(server)
if server and server.auth_type == MCPAuth.oauth2:
# The challenge decision is per oauth2 sub-mode, not per header:
# gateway-managed modes (M2M and interactive authorization_code)
# never receive a client-supplied upstream token, so a bearer in
# Authorization is a LiteLLM key (surfaced here as oauth2_headers)
# and must not suppress the challenge. Only the delegate mode
# treats a present bearer as the upstream token. The sub-mode is
# resolved the same way egress resolves it, via
# effective_oauth2_flow: an unstamped (null oauth2_flow) row with
# the M2M shape resolves to client_credentials, so the bare
# has_client_credentials column is never trusted here.
if MCPServerManager.effective_oauth2_flow(server) == "client_credentials":
# M2M: the gateway mints its own token at egress from the
# stored client credentials, so there is nothing to challenge.
continue
if getattr(server, "delegate_auth_to_upstream", False) is not True:
# Gateway-managed interactive (authorization_code): the only
# thing that authorizes egress is a stored per-user token, so
# challenge whenever one is absent, regardless of any bearer.
# The v2 resolver owns the existence check, so every
# authorization_code resolution (egress and this discovery
# challenge) runs through it. A keyless admitted subject is
# challenged with the per-server resource_metadata (whose
# authorization server is the gateway itself, vaulting via the
# authorize interlude); the per-server relay advertised below
# cannot vault without a litellm key on its token request.
if await operations.global_mcp_server_manager.has_user_oauth_token(server, user_api_key_auth):
continue
if _is_mcp_admitted_user_subject(user_api_key_auth):
raise HTTPException(
status_code=401,
detail="Unauthorized",
headers={
"www-authenticate": get_passthrough_www_authenticate(
scope=scope,
server_name=server_name,
)
},
)
request = StarletteRequest(scope)
base_url = get_request_base_url(request)
_path = get_route_relative_request_path(scope)
# Pick the well-known AS-metadata form that matches the inbound route
# so strict RFC 9728 §3.2 clients can resolve it correctly.
as_metadata_root = f"{base_url}/.well-known/oauth-authorization-server{well_known_root_suffix()}"
if _path.startswith(f"/mcp/{server_name}"):
_as_url = f"{as_metadata_root}/mcp/{server_name}"
else:
_as_url = f"{as_metadata_root}/{server_name}"
authorization_uri = f'Bearer authorization_uri="{_as_url}"'
raise HTTPException(
status_code=401,
detail="Unauthorized",
headers={"www-authenticate": authorization_uri},
)
if not oauth2_headers:
# Delegate-auth servers run upstream PKCE: a present bearer is
# the upstream token, so only challenge when it is absent, with
# the proxied resource_metadata (RFC 9728), not the gateway
# authorization_uri above which would authorize against the
# gateway instead of the upstream IdP.
www_authenticate = get_passthrough_www_authenticate(
if mcp_servers is None:
allowed: Final = await operations._get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth, mcp_servers=None, client_ip=client_ip
)
eligible: Final = tuple(
server for server in allowed if allowed_server_ids is None or server.server_id in allowed_server_ids
)
results: Final = await asyncio.gather(
*(
_server_auth_challenge(
configured_server=server,
server_name=server.alias or server.server_name or server.name,
scope=scope,
server_name=server_name,
mcp_servers=[server.alias or server.server_name or server.name],
oauth2_headers=oauth2_headers,
mcp_server_auth_headers=mcp_server_auth_headers,
user_api_key_auth=user_api_key_auth,
client_ip=client_ip,
raw_headers=raw_headers,
)
raise HTTPException(
status_code=401,
detail="Unauthorized",
headers={"www-authenticate": www_authenticate},
for server in eligible
),
return_exceptions=True,
)
failures: Final = tuple(
(server, result) for server, result in zip(eligible, results) if isinstance(result, BaseException)
)
for server, failure in failures:
if not isinstance(failure, Exception):
raise failure
if not isinstance(failure, HTTPException) or failure.status_code != 401:
verbose_logger.warning(
"MCP authentication preflight failed for %s (%s)", server.name, type(failure).__name__
)
# Delegate server with a bearer present: it is the upstream token,
# so admit the session and move to the next target. Every oauth2
# sub-mode is terminal here (continue or raise) so no oauth2 server
# reaches the token_exchange / pass-through blocks below.
continue
# token_exchange (OBO): the caller supplied no subject token. Challenge at connect
# (transport level, where WWW-Authenticate survives) with the RFC 9728 resource_metadata
# so the client discovers the IdP, SSOs, and retries with a subject token, which LiteLLM
# then exchanges. A tool-call-time 401 would be wrapped into a JSON-RPC error and the
# header lost, so the discovery flow needs this pre-emptive challenge.
if server and server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers:
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 # lazy: adapter pulls MCP subgraph
raise_token_exchange_challenge,
if not failures or len(failures) != len(results):
return
for _, failure in failures:
if not isinstance(failure, HTTPException) or failure.status_code != 401:
raise failure
if all(server.is_gateway_managed_oauth2 for server in eligible):
raise _gateway_dcr_challenge(
StarletteRequest(scope), get_route_relative_request_path(scope), None, invalid_token=False
)
from litellm.proxy.middleware.per_request_root_path_middleware import ( # noqa: PLC0415 # lazy: middleware imports proxy utils
get_request_root_path,
)
raise_token_exchange_challenge(server, root_path=get_request_root_path())
# Exchange-backed modes (token_exchange's OBO mint, id_jag's stored-assertion mint): run
# the exchange here at the transport edge, so a rejected subject raises the RFC 9728
# challenge and any other failure its public status, instead of the session opening and
# list_tools masking it as an empty tool list. The manager owns which modes pre-flight
# and what each mints from. Gated to single-server routes the key may reach; the
# multi-server aggregate keeps absorbing per-server auth failures so one bad server
# cannot 401 the whole connect.
raise failures[0][1]
for server_name in mcp_servers:
if (
server
and len(mcp_servers or []) == 1
and server.server_id
in frozenset(
allowed.server_id
for allowed in await operations._get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip
)
)
):
await operations.global_mcp_server_manager.preflight_token_exchange(
server=server,
server := operations.global_mcp_server_manager.get_mcp_server_by_name(server_name, client_ip=client_ip)
) is None:
continue
if allowed_server_ids is not None and server.server_id not in allowed_server_ids:
continue
if (
challenge := await _server_auth_challenge(
configured_server=server,
server_name=server_name,
scope=scope,
mcp_servers=mcp_servers,
oauth2_headers=oauth2_headers,
mcp_server_auth_headers=mcp_server_auth_headers,
user_api_key_auth=user_api_key_auth,
client_ip=client_ip,
raw_headers=raw_headers,
)
# Pass-through OAuth: when the admin has opted a server into
# forwarding the client's bearer token (is_oauth_passthrough) and
# the client hasn't supplied one, fail fast with 401 and point
# them at the gateway's oauth-protected-resource well-known URL.
# That endpoint proxies the upstream's metadata so the client
# kicks off OAuth against the real upstream IdP, not the gateway.
if (
server
and server.is_oauth_passthrough
and not operations._client_has_passthrough_authorization(
server, oauth2_headers, mcp_server_auth_headers
)
):
www_authenticate = get_passthrough_www_authenticate(
scope=scope,
server_name=server_name,
)
raise HTTPException(
status_code=401,
detail="Unauthorized",
headers={"www-authenticate": www_authenticate},
)
if (
server
and server.is_oauth_delegate
and len(mcp_servers or []) == 1
and _get_forwarded_auth_from_scope(scope) is None
and not operations._client_has_per_server_auth_header(server, mcp_server_auth_headers)
):
www_authenticate = get_passthrough_www_authenticate(
scope=scope,
server_name=server_name,
)
raise HTTPException(
status_code=401,
detail="Unauthorized",
headers={"www-authenticate": www_authenticate},
)
if (
server
and server.is_true_passthrough
and len(mcp_servers or []) == 1
and not _scope_has_authorization_header(scope)
and not operations._client_has_per_server_auth_header(server, mcp_server_auth_headers)
):
if server.is_dcr_bridge:
raise HTTPException(
status_code=401,
detail="Unauthorized",
headers={
"www-authenticate": get_passthrough_www_authenticate(
scope=scope,
server_name=server_name,
)
},
)
upstream_status, upstream_www_authenticate = await _probe_upstream_auth(server.url or "", "")
if upstream_status == 401 and upstream_www_authenticate:
raise HTTPException(
status_code=401,
detail="Unauthorized",
headers={"www-authenticate": upstream_www_authenticate},
)
) is not None:
raise challenge
def _get_authorization_header_from_scope(scope: Scope) -> str | None:
"""First ``Authorization`` header value in the ASGI scope, or None."""
@ -2017,16 +2122,17 @@ if MCP_AVAILABLE:
# from the fully-authorized server set: a passthrough server that
# the active toolset excludes should not trigger an OAuth flow
# for a server the caller will be 403'd on after authentication.
await _raise_preemptive_401_for_unauthenticated_servers(
scope=scope,
mcp_servers=mcp_servers,
oauth2_headers=oauth2_headers,
mcp_server_auth_headers=mcp_server_auth_headers,
user_api_key_auth=user_api_key_auth,
client_ip=_client_ip,
allowed_server_ids=toolset_allowed_server_ids,
raw_headers=raw_headers,
)
if mcp_servers is not None:
await _raise_preemptive_401_for_unauthenticated_servers(
scope=scope,
mcp_servers=mcp_servers,
oauth2_headers=oauth2_headers,
mcp_server_auth_headers=mcp_server_auth_headers,
user_api_key_auth=user_api_key_auth,
client_ip=_client_ip,
allowed_server_ids=toolset_allowed_server_ids,
raw_headers=raw_headers,
)
# Pre-flight auth check for pass-through servers. Must run after
# toolset scoping so the probe list is derived from the fully-authorized
@ -2103,6 +2209,18 @@ if MCP_AVAILABLE:
consumed_messages, body = await _read_request_body_for_routing(receive)
is_initialize = _is_initialize_request(body)
if is_initialize and mcp_servers is None:
await _raise_preemptive_401_for_unauthenticated_servers(
scope=scope,
mcp_servers=None,
oauth2_headers=oauth2_headers,
mcp_server_auth_headers=mcp_server_auth_headers,
user_api_key_auth=user_api_key_auth,
client_ip=_client_ip,
allowed_server_ids=toolset_allowed_server_ids,
raw_headers=raw_headers,
)
use_stateful: Final = bool(session_id or is_initialize)
target_manager: Final = session_manager_stateful if use_stateful else session_manager_stateless
@ -2248,7 +2366,12 @@ if MCP_AVAILABLE:
scoped_server_endpoint=scoped_server_endpoint,
is_initialize=is_initialize,
):
await target_manager.handle_request(scope, receive, local_send)
if request_method == "POST" and body and not is_initialize:
auth_response: Final = MCPAuthResponse(local_send)
scope[_MCP_AUTH_RESPONSE_SCOPE_KEY] = auth_response
await target_manager.handle_request(scope, receive, auth_response.send)
else:
await target_manager.handle_request(scope, receive, local_send)
if use_stateful and session_id and scope.get("method") == "DELETE":
_remove_stateful_session_tracking(session_id)

View file

@ -31506,6 +31506,21 @@
"title": "Flow",
"type": "string"
},
"selected_servers": {
"anyOf": [
{
"items": {
"type": "string"
},
"maxItems": 100,
"type": "array"
},
{
"type": "null"
}
],
"title": "Selected Servers"
},
"team_id": {
"anyOf": [
{

View file

@ -255,3 +255,59 @@ def test_pure_non_auth_response_still_classifies_upstream_error():
fault = classify_list_exception(exc)
assert fault.tag == "upstream_error"
assert fault.status_code == 502
@pytest.mark.parametrize(
"outcomes,expected_status",
(
({}, None),
({"timeout": ServerListFault(tag="timeout")}, None),
({"denied": ServerListFault(tag="forbidden", status_code=403), "timeout": ServerListFault(tag="timeout")}, 403),
({"empty": ServerListOk(tool_count=0)}, None),
({"auth": ServerListFault(tag="auth_required", status_code=401)}, 401),
({"denied": ServerListFault(tag="forbidden", status_code=403)}, 403),
({"denied": ServerListFault(tag="forbidden", status_code=403), "auth": ServerListFault(tag="auth_required", status_code=401)}, 401),
({"auth": ServerListFault(tag="auth_required", status_code=401), "empty": ServerListOk(tool_count=0)}, None),
({"auth": ServerListFault(tag="auth_required", status_code=401), "healthy": ServerListOk(tool_count=2)}, None),
({"auth": ServerListFault(tag="auth_required", status_code=401), "timeout": ServerListFault(tag="timeout")}, 401),
),
)
def test_listing_auth_failure_requires_no_successful_server(
outcomes: dict[str, ServerListOk | ServerListFault], expected_status: int | None
) -> None:
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import listing_auth_error
from typing import Final
failure: Final = listing_auth_error(outcomes)
assert (failure.status_code if failure else None) == expected_status
def test_classified_auth_challenge_is_preserved_without_serializing_it() -> None:
from typing import Final
from fastapi import HTTPException
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import listing_auth_error
challenge: Final = 'Bearer resource_metadata="https://gateway/.well-known/oauth-protected-resource/mcp"'
fault: Final = classify_list_exception(HTTPException(401, headers={"WWW-Authenticate": challenge}))
failure: Final = listing_auth_error({"upstream": fault})
assert failure is not None
assert failure.www_authenticate == challenge
assert failure.server_name == "upstream"
assert "www_authenticate" not in fault.model_dump()
assert challenge not in repr(fault)
assert outcome_wire_value(fault) == {"status": "auth_required", "http_status": 401}
def test_auth_recovery_uses_routable_server_name_without_exposing_it_in_outcomes() -> None:
from typing import Final
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import listing_auth_error
fault: Final = classify_list_exception(MCPUpstreamAuthError(401, None, "upstream-internal"))
failure: Final = listing_auth_error({"short-prefix": fault})
assert failure is not None
assert failure.server_name == "upstream-internal"
assert "server_name" not in fault.model_dump()
assert "upstream-internal" not in repr(fault)
assert outcome_wire_value(fault) == {"status": "auth_required", "http_status": 401}
def test_auth_challenge_traversal_preserves_typed_carrier() -> None:
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import upstream_auth_challenge
assert upstream_auth_challenge(MCPUpstreamAuthError(401, "Bearer", "upstream")) == (401, "Bearer")

View file

@ -74,6 +74,13 @@ CODE_CHALLENGE = urlsafe_b64encode(hashlib.sha256(CODE_VERIFIER.encode("ascii"))
@pytest.fixture(autouse=True)
def _salt_key(monkeypatch):
monkeypatch.setenv("LITELLM_SALT_KEY", MASTER_KEY)
from unittest.mock import patch
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.get_mcp_server_by_id",
return_value=_scoped_mcp_server("public", auth_type="none"),
):
yield
def _request(path="/authorize", query="", cookies=None, method="GET"):
@ -339,6 +346,8 @@ async def test_full_walk_register_authorize_complete_token_and_replay(redirect_u
handle, cookies = _flow_cookie_from(authorize_response)
denied = await complete_connect_flow(
selected_servers=("public-id",),
lookup_server_reachability=_ServerReachability(),
request=_request("/authorize/complete", cookies=cookies, method="POST"),
flow_handle=handle,
session_user_id="attacker",
@ -347,6 +356,8 @@ async def test_full_walk_register_authorize_complete_token_and_replay(redirect_u
assert denied.status_code == 403
anonymous = await complete_connect_flow(
selected_servers=("public-id",),
lookup_server_reachability=_ServerReachability(),
request=_request("/authorize/complete", cookies=cookies, method="POST"),
flow_handle=handle,
session_user_id=None,
@ -355,6 +366,8 @@ async def test_full_walk_register_authorize_complete_token_and_replay(redirect_u
assert anonymous.status_code == 401
completed = await complete_connect_flow(
selected_servers=("public-id",),
lookup_server_reachability=_ServerReachability(),
request=_request("/authorize/complete", cookies=cookies, method="POST"),
flow_handle=handle,
session_user_id="u1",
@ -429,6 +442,8 @@ async def test_full_walk_register_authorize_complete_token_and_replay(redirect_u
@pytest.mark.asyncio
async def test_complete_rejects_missing_tampered_and_expired_flows():
missing = await complete_connect_flow(
selected_servers=("public-id",),
lookup_server_reachability=_ServerReachability(),
request=_request("/authorize/complete", method="POST"),
flow_handle="nope",
session_user_id="u1",
@ -437,6 +452,8 @@ async def test_complete_rejects_missing_tampered_and_expired_flows():
assert missing.status_code == 400
tampered = await complete_connect_flow(
selected_servers=("public-id",),
lookup_server_reachability=_ServerReachability(),
request=_request("/authorize/complete", cookies={f"{CONNECT_FLOW_COOKIE_PREFIX}h1": "garbage"}, method="POST"),
flow_handle="h1",
session_user_id="u1",
@ -518,6 +535,8 @@ async def test_token_gates_on_live_user_revalidation(failure, expected_status, e
authorize_response = _authorize(client_id, session_user_id="deactivated-user")
handle, cookies = _flow_cookie_from(authorize_response)
completed = await complete_connect_flow(
selected_servers=("public-id",),
lookup_server_reachability=_ServerReachability(),
request=_request("/authorize/complete", cookies=cookies, method="POST"),
flow_handle=handle,
session_user_id="deactivated-user",
@ -572,6 +591,8 @@ async def test_flow_is_single_use_shared_cache_rejects_second_complete():
handle, cookies = _flow_cookie_from(_authorize(client_id, session_user_id="u1"))
first = await complete_connect_flow(
selected_servers=("public-id",),
lookup_server_reachability=_ServerReachability(),
request=_request("/authorize/complete", cookies=cookies, method="POST"),
flow_handle=handle,
session_user_id="u1",
@ -579,6 +600,8 @@ async def test_flow_is_single_use_shared_cache_rejects_second_complete():
)
assert first.status_code == 303
second = await complete_connect_flow(
selected_servers=("public-id",),
lookup_server_reachability=_ServerReachability(),
request=_request("/authorize/complete", cookies=cookies, method="POST"),
flow_handle=handle,
session_user_id="u1",
@ -720,6 +743,8 @@ async def _complete(redirect_uri: str, delivery, cookies=None, handle=None, sess
if cookies is None:
handle, cookies = _flow_cookie_from(_authorize(client_id, session_user_id="u1", redirect_uri=redirect_uri))
response = await complete_connect_flow(
selected_servers=("public-id",),
lookup_server_reachability=_ServerReachability(),
request=_request("/authorize/complete", cookies=cookies, method="POST"),
flow_handle=handle,
session_user_id=session_user_id,
@ -827,6 +852,8 @@ async def test_unknown_delivery_value_is_rejected_before_the_flow_is_consumed():
handle, cookies = _flow_cookie_from(_authorize(client_id, session_user_id="u1", redirect_uri=LOOPBACK_REDIRECT_URI))
rejected = await complete_connect_flow(
selected_servers=("public-id",),
lookup_server_reachability=_ServerReachability(),
request=_request("/authorize/complete", cookies=cookies, method="POST"),
flow_handle=handle,
session_user_id="u1",
@ -837,6 +864,8 @@ async def test_unknown_delivery_value_is_rejected_before_the_flow_is_consumed():
assert json.loads(rejected.body)["error"] == "invalid_request"
retried = await complete_connect_flow(
selected_servers=("public-id",),
lookup_server_reachability=_ServerReachability(),
request=_request("/authorize/complete", cookies=cookies, method="POST"),
flow_handle=handle,
session_user_id="u1",
@ -997,7 +1026,9 @@ async def _complete_page(response, scoped_server=None, vendor=None, reachable=No
handle, cookies = _flow_cookie_from(response)
with patch(_MANAGER_PATCH) as manager:
manager.get_mcp_server_by_id.return_value = scoped_server
manager.get_mcp_server_by_id.side_effect = lambda server_id: (
scoped_server or _scoped_mcp_server("public", auth_type="none")
) if server_id == "public-id" else scoped_server
return await complete_connect_flow(
request=_request("/authorize/complete", cookies=cookies, method="POST"),
flow_handle=handle,
@ -1005,7 +1036,7 @@ async def _complete_page(response, scoped_server=None, vendor=None, reachable=No
cache=cache or DualCache(),
lookup_vendor_credential=vendor or _VendorCredential(),
lookup_server_reachability=reachable or _ServerReachability(),
**overrides,
**{"selected_servers": ("public-id",), **overrides},
)
@ -1814,6 +1845,8 @@ async def test_mcp_wire_formats_carry_no_native_client_fields():
assert "audience" not in flow_wire
assert "team_id" not in flow_wire
completed = await complete_connect_flow(
selected_servers=("public-id",),
lookup_server_reachability=_ServerReachability(),
request=_request("/authorize/complete", cookies=cookies, method="POST"),
flow_handle=handle,
session_user_id="u1",
@ -2354,3 +2387,123 @@ async def test_token_exchange_relays_a_mint_refusal(failure, status, error):
response = await _exchange_native(client_id, _Minter(failure), _Exchanger())
assert response.status_code == status
assert json.loads(response.body)["error"] == error
@pytest.mark.asyncio
async def test_unified_completion_requires_an_upstream_selection():
client_id = (await _register([REDIRECT_URI]))["client_id"]
response = _authorize(client_id, session_user_id="u1")
completed = await _complete_page(response, selected_servers=())
assert completed.status_code == 400
assert "location" not in completed.headers
assert "select" in json.loads(completed.body)["error_description"]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"credential,reachable,status",
[("absent", True, 400), ("unavailable", True, 503), ("present", False, 400), ("present", True, 303)],
)
async def test_unified_completion_checks_selected_upstream_and_permissions(credential, reachable, status):
client_id = (await _register([REDIRECT_URI]))["client_id"]
response = _authorize(client_id, session_user_id="u1")
cache = DualCache()
server = _scoped_mcp_server(oauth2_flow="authorization_code")
vendor = _VendorCredential(credential)
completed = await _complete_page(
response,
scoped_server=server,
vendor=vendor,
reachable=_ServerReachability(reachable),
cache=cache,
selected_servers=("github-id",),
)
assert completed.status_code == status
if not reachable:
assert vendor.calls == []
if status != 303:
assert "location" not in completed.headers
retried = await _complete_page(response, scoped_server=server, cache=cache, selected_servers=("github-id",))
assert retried.status_code == 303
else:
code = parse_qs(urlparse(completed.headers["location"]).query)["code"][0]
token = await _redeem(code, client_id)
assert _opened_principal(json.loads(token.body)).resource_server_id is None
@pytest.mark.asyncio
async def test_unified_cancel_does_not_require_selected_servers():
client_id = (await _register([REDIRECT_URI]))["client_id"]
response = _authorize(client_id, session_user_id="u1")
completed = await _complete_page(response, selected_servers=(), decision="deny")
assert completed.status_code == 303
assert parse_qs(urlparse(completed.headers["location"]).query)["error"] == ["access_denied"]
@pytest.mark.asyncio
async def test_unified_completion_checks_every_selected_server_and_preserves_cancellation():
import asyncio
from unittest.mock import patch
client_id = (await _register([REDIRECT_URI]))["client_id"]
response = _authorize(client_id, session_user_id="u1")
handle, cookies = _flow_cookie_from(response)
servers = {name: _scoped_mcp_server(name, oauth2_flow="authorization_code") for name in ("github", "slack")}
cache = DualCache()
async def credential(user_id, server_id):
if server_id == "slack-id":
return "absent"
return "present"
async def cancelled(user_id, server_id):
raise asyncio.CancelledError()
with patch(_MANAGER_PATCH) as manager:
manager.get_mcp_server_by_name.side_effect = servers.get
manager.get_mcp_server_by_id.side_effect = lambda server_id: next(
(server for server in servers.values() if server.server_id == server_id), None
)
arguments = {
"request": _request("/authorize/complete", cookies=cookies, method="POST"),
"flow_handle": handle,
"session_user_id": "u1",
"cache": cache,
"lookup_server_reachability": _ServerReachability(),
}
missing = await complete_connect_flow(
**arguments, selected_servers=("missing",), lookup_vendor_credential=credential
)
assert missing.status_code == 400
unfinished = await complete_connect_flow(
**arguments, selected_servers=("github-id", "slack-id"), lookup_vendor_credential=credential
)
assert unfinished.status_code == 400
with pytest.raises(asyncio.CancelledError):
await complete_connect_flow(**arguments, selected_servers=("github-id",), lookup_vendor_credential=cancelled)
completed = await complete_connect_flow(
**arguments, selected_servers=("github-id", "slack-id"), lookup_vendor_credential=_VendorCredential()
)
assert completed.status_code == 303
@pytest.mark.asyncio
async def test_unified_completion_validates_selected_id_despite_alias_collision():
from unittest.mock import patch
client_id = (await _register([REDIRECT_URI]))["client_id"]
response = _authorize(client_id, session_user_id="u1")
handle, cookies = _flow_cookie_from(response)
selected = _scoped_mcp_server("github", oauth2_flow="authorization_code")
other = _scoped_mcp_server("other", auth_type="none").model_copy(update={"alias": selected.server_id})
cache = DualCache()
with patch(_MANAGER_PATCH) as manager:
manager.get_mcp_server_by_name.return_value = other
manager.get_mcp_server_by_id.side_effect = {selected.server_id: selected, other.server_id: other}.get
vendor = _VendorCredential("absent")
arguments = dict(request=_request("/authorize/complete", cookies=cookies, method="POST"), flow_handle=handle, session_user_id="u1", cache=cache, selected_servers=(selected.server_id,), lookup_server_reachability=_ServerReachability())
refused = await complete_connect_flow(**arguments, lookup_vendor_credential=vendor)
assert refused.status_code == 400
assert vendor.calls == [("u1", selected.server_id)]
completed = await complete_connect_flow(**arguments, lookup_vendor_credential=_VendorCredential())
assert completed.status_code == 303

View file

@ -47,7 +47,7 @@ def _make_server(server_id: str, max_concurrent_requests: Optional[int]) -> MCPS
def _patch_client_with_tracker(manager: MCPServerManager, tracker: _ConcurrencyTracker):
async def fake_create_mcp_client(server, **kwargs):
class _ProbeClient:
async def call_tool(self, params, host_progress_callback=None, allow_input_required=False):
async def call_tool(self, params, host_progress_callback=None, allow_input_required=False, raise_on_error: bool = False):
tracker.enter(server.server_id)
try:
await asyncio.sleep(HOLD_SECONDS)

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,469 @@ def test_passthrough_admission_recognizes_only_matching_authorization(oauth_head
server = MCPServer(server_id="catalog", name="catalog", alias="catalog", transport=MCPTransport.http)
assert _client_has_passthrough_authorization(server, oauth_headers, server_headers) is authorized
@pytest.mark.asyncio
async def test_listing_transport_preserves_auth_challenge_before_sse_success() -> None:
from starlette.exceptions import HTTPException
from litellm.proxy._experimental.mcp_server.server import MCPAuthResponse
send: Final = AsyncMock()
response: Final = MCPAuthResponse(send)
await response.send(
{"type": "http.response.start", "status": 200, "headers": [(b"content-type", b"text/event-stream")]}
)
await response.send({"type": "http.response.body", "body": b": ping\r\n\r\n", "more_body": True})
response.challenge = HTTPException(
401, "Unauthorized", headers={"WWW-Authenticate": 'Bearer resource_metadata="http://localhost/mcp-metadata"'}
)
await response.send(
{
"type": "http.response.body",
"body": b'event: message\r\ndata: {"jsonrpc":"2.0","id":1,"result":{"tools":[]}}\r\n\r\n',
"more_body": True,
}
)
await response.send({"type": "http.response.body", "body": b"", "more_body": False})
sent: Final = tuple(call.args[0] for call in send.await_args_list)
assert [m["status"] for m in sent if m["type"] == "http.response.start"] == [401]
assert dict(sent[0]["headers"])[b"www-authenticate"] == b'Bearer resource_metadata="http://localhost/mcp-metadata"'
assert b'"tools"' not in b"".join(m.get("body", b"") for m in sent)
assert sent[-1].get("more_body", False) is False
@pytest.mark.asyncio
@pytest.mark.parametrize("status", [200, 400, 403])
async def test_listing_transport_preserves_non_auth_responses(status: int) -> None:
from litellm.proxy._experimental.mcp_server.server import MCPAuthResponse
send: Final = AsyncMock()
response: Final = MCPAuthResponse(send)
start: Final = {
"type": "http.response.start",
"status": status,
"headers": [(b"content-type", b"application/json")],
}
body: Final = {
"type": "http.response.body",
"body": b'{"jsonrpc":"2.0","id":1,"result":{"tools":[]}}',
"more_body": False,
}
await response.send(start)
await response.send(body)
assert tuple(call.args[0] for call in send.await_args_list) == (start, body)
@pytest.mark.asyncio
async def test_protocol_listing_does_not_report_success_when_every_server_requires_auth() -> None:
from unittest.mock import patch
from mcp.types import ListToolsRequest
from litellm.proxy._experimental.mcp_server import operations
from litellm.proxy._experimental.mcp_server.contracts import OperationContext
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing, ServerListFault
from litellm.proxy._types import UserAPIKeyAuth
context: Final = OperationContext(_caller=UserAPIKeyAuth(user_id="reader"))
listing: Final = AggregateToolListing([], {"github": ServerListFault(tag="auth_required", status_code=401)})
with patch.object(operations, "_list_mcp_tools", AsyncMock(return_value=listing)):
with pytest.raises(MCPUpstreamAuthError) as caught:
await operations.GatewayOperations().execute(ListToolsRequest(), context)
assert caught.value.status_code == 401
assert caught.value.server_name == "github"
@pytest.mark.asyncio
@pytest.mark.parametrize("method", ("tools/list", "prompts/list", "resources/list", "resources/templates/list"))
@pytest.mark.parametrize("protocol", ("2025-06-18", "2025-11-25"))
@pytest.mark.parametrize("json_response", (False, True))
@pytest.mark.parametrize("stateful", (False, True))
@pytest.mark.parametrize("path", ("/mcp", "/github/mcp"))
async def test_streamable_http_listing_returns_late_oauth_challenge(
monkeypatch: pytest.MonkeyPatch, stateful: bool, path: str, json_response: bool, protocol: str, method: str
) -> None:
from litellm.proxy._experimental.mcp_server import server
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing, ServerListFault
from litellm.proxy._types import UserAPIKeyAuth
auth: Final = UserAPIKeyAuth(api_key="test-owner", user_id="test-user")
monkeypatch.setattr(
server, "extract_mcp_auth_context", AsyncMock(return_value=(auth, None, None, None, None, None))
)
monkeypatch.setattr(server, "_raise_preemptive_401_for_unauthenticated_servers", AsyncMock())
monkeypatch.setattr(server, "_check_passthrough_upstream_auth", AsyncMock())
listing: Final = AsyncMock(
return_value=AggregateToolListing(
[],
{
"github": ServerListFault(
tag="auth_required",
status_code=401,
www_authenticate='Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp/github"',
)
},
)
)
if method != "tools/list":
listing.side_effect = MCPUpstreamAuthError(
401, 'Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp/github"', "github"
)
list_function: Final = {
"tools/list": "_list_mcp_tools",
"prompts/list": "_list_mcp_prompts",
"resources/list": "_list_mcp_resources",
"resources/templates/list": "_list_mcp_resource_templates",
}[method]
monkeypatch.setattr(server.operations, list_function, listing)
monkeypatch.setattr(server.operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[]))
monkeypatch.setattr(server.operations, "_raise_if_initialize_grants_no_mcp_servers", AsyncMock())
from mcp.server.streamable_http_manager import StreamableHTTPSessionManager
monkeypatch.setattr(
server,
"session_manager_stateless",
StreamableHTTPSessionManager(app=server.server, stateless=True, json_response=json_response),
)
monkeypatch.setattr(
server,
"session_manager_stateful",
StreamableHTTPSessionManager(app=server.server, stateless=False, json_response=json_response),
)
await server.initialize_session_managers()
try:
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app=server.app), base_url="http://gateway"
) as client:
headers: Final = {"accept": "application/json, text/event-stream", "mcp-protocol-version": protocol}
if stateful:
initialized: Final = await client.post(
path,
headers=headers,
json={
"jsonrpc": "2.0",
"id": 0,
"method": "initialize",
"params": {
"protocolVersion": protocol,
"capabilities": {},
"clientInfo": {"name": "test", "version": "1"},
},
},
)
assert initialized.status_code == 200, initialized.text
client.headers["mcp-session-id"] = initialized.headers["mcp-session-id"]
notification: Final = await client.post(
path, headers=headers, json={"jsonrpc": "2.0", "method": "notifications/initialized"}
)
assert notification.status_code == 202
malformed: Final = await client.post(
path, headers={**headers, "content-type": "application/json"}, content=b"{"
)
assert malformed.status_code == 400
listing.assert_not_awaited()
response: Final = await client.post(
path, headers=headers, json={"jsonrpc": "2.0", "id": 1, "method": method}
)
assert response.status_code == 401, response.text
assert (
response.headers["www-authenticate"]
== 'Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp/github"'
)
listing.assert_awaited_once()
finally:
await server.shutdown_session_managers()
@pytest.mark.asyncio
async def test_late_auth_failure_does_not_rewrite_committed_stream() -> None:
from fastapi import HTTPException
from litellm.proxy._experimental.mcp_server.server import MCPAuthResponse
send: Final = AsyncMock()
response: Final = MCPAuthResponse(send)
start: Final = {"type": "http.response.start", "status": 200, "headers": [(b"content-type", b"text/event-stream")]}
progress: Final = {
"type": "http.response.body",
"body": b'data: {"method":"notifications/progress"}\n\n',
"more_body": True,
}
error: Final = {
"type": "http.response.body",
"body": b'data: {"error":{"code":-32600,"message":"Upstream authorization failed (HTTP 401)"}}\n\n',
"more_body": False,
}
await response.send(start)
await response.send(progress)
response.challenge = HTTPException(401, headers={"WWW-Authenticate": "Bearer"})
await response.send(error)
assert tuple(call.args[0] for call in send.await_args_list) == (start, progress, error)
@pytest.mark.asyncio
@pytest.mark.parametrize("kind", ("prompts", "resources", "resource_templates"))
async def test_optional_listing_propagates_auth_without_discarding_healthy_servers(
kind: str, monkeypatch: pytest.MonkeyPatch
) -> None:
from litellm.proxy._experimental.mcp_server import operations
blocked: Final = _http_server("blocked", "blocked")
healthy: Final = _http_server("healthy", "healthy")
fetch: Final = AsyncMock(side_effect=MCPUpstreamAuthError(401, "Bearer", "blocked"))
manager: Final = MagicMock()
setattr(manager, f"get_{kind}_from_server", fetch)
monkeypatch.setattr(operations, "global_mcp_server_manager", manager)
monkeypatch.setattr(operations, "_prepare_mcp_server_headers", MagicMock(return_value=(None, None)))
monkeypatch.setattr(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[blocked]))
listing: Final = getattr(operations, f"_list_mcp_{kind}")
with pytest.raises(MCPUpstreamAuthError) as failure:
await listing()
assert failure.value.www_authenticate == "Bearer"
fetch.assert_awaited_once()
fetch.reset_mock(side_effect=True)
fetch.side_effect = [MCPUpstreamAuthError(401, "Bearer", "blocked"), []]
monkeypatch.setattr(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[blocked, healthy]))
assert await listing() == []
assert fetch.await_count == 2
@pytest.mark.asyncio
@pytest.mark.parametrize(
"operation",
(
"get_prompts_from_server",
"get_resources_from_server",
"get_resource_templates_from_server",
"get_prompt_from_server",
"read_resource_from_server",
),
)
@pytest.mark.parametrize("status", (401, 403))
@pytest.mark.parametrize("dcr_bridge", (False, True))
async def test_manager_preserves_auth_failures_for_prompts_and_resources(
operation: str, status: int, dcr_bridge: bool, monkeypatch: pytest.MonkeyPatch
) -> None:
from fastapi import HTTPException
from pydantic import AnyUrl
manager: Final = MCPServerManager()
upstream: Final = _http_server("upstream", "upstream", auth_type=MCPAuth.oauth_delegate, dcr_bridge=dcr_bridge)
create: Final = AsyncMock(side_effect=HTTPException(status, headers={"WWW-Authenticate": "Bearer"}))
monkeypatch.setattr(manager, "_create_mcp_client", create)
kwargs: Final = (
{"prompt_name": "example"}
if operation == "get_prompt_from_server"
else {"url": AnyUrl("https://example.com/resource")}
if operation == "read_resource_from_server"
else {}
)
with pytest.raises(MCPUpstreamAuthError) as failure:
await getattr(manager, operation)(server=upstream, user_api_key_auth=None, **kwargs)
assert failure.value.status_code == status
assert failure.value.www_authenticate == (None if dcr_bridge else "Bearer")
assert failure.value.server_name == "upstream"
create.assert_awaited_once()
@pytest.mark.asyncio
async def test_optional_listing_preserves_cancellation() -> None:
import asyncio
from litellm.proxy._experimental.mcp_server.operations import _collect_mcp_listing
fetch: Final = AsyncMock(side_effect=asyncio.CancelledError())
with pytest.raises(asyncio.CancelledError):
await _collect_mcp_listing((_http_server("upstream", "upstream"),), fetch)
fetch.assert_awaited_once()
@pytest.mark.asyncio
@pytest.mark.parametrize("operation", ("get_prompt_from_server", "read_resource_from_server"))
@pytest.mark.parametrize("extra_headers", (None, {"x-forwarded": "caller"}))
async def test_prompt_and_resource_calls_preserve_static_headers_and_non_auth_failures(
operation: str, extra_headers: dict[str, str] | None, monkeypatch: pytest.MonkeyPatch
) -> None:
from pydantic import AnyUrl
manager: Final = MCPServerManager()
upstream: Final = _http_server("upstream", "upstream", static_headers={"x-upstream": "configured"})
failure: Final = RuntimeError("Upstream unavailable")
create: Final = AsyncMock(side_effect=failure)
monkeypatch.setattr(manager, "_create_mcp_client", create)
kwargs: Final = (
{"prompt_name": "example"}
if operation == "get_prompt_from_server"
else {"url": AnyUrl("https://example.com/resource")}
)
with pytest.raises(RuntimeError) as caught:
await getattr(manager, operation)(server=upstream, user_api_key_auth=None, extra_headers=extra_headers, **kwargs)
assert caught.value is failure
assert create.await_args.kwargs["extra_headers"] == {**(extra_headers or {}), "x-upstream": "configured"}
create.assert_awaited_once()
@pytest.mark.asyncio
@pytest.mark.parametrize("preamble", (b"id: resume-token\r\ndata: \r\n\r\n", b": ping\r\n\r\n"))
async def test_transport_preserves_sse_priming_event_on_success(preamble: bytes) -> None:
from litellm.proxy._experimental.mcp_server.server import MCPAuthResponse
send: Final = AsyncMock()
response: Final = MCPAuthResponse(send)
start: Final = {"type": "http.response.start", "status": 200, "headers": [(b"content-type", b"text/event-stream")]}
priming: Final = {"type": "http.response.body", "body": preamble, "more_body": True}
tools: Final = {
"type": "http.response.body",
"body": b'data: {"jsonrpc":"2.0","id":1,"result":{"tools":[]}}\n\n',
"more_body": True,
}
await response.send(start)
await response.send(priming)
send.assert_not_awaited()
await response.send(tools)
assert tuple(call.args[0] for call in send.await_args_list) == (start, priming, tools)
@pytest.mark.asyncio
async def test_tool_call_preserves_resolver_http_challenge(monkeypatch: pytest.MonkeyPatch) -> None:
from fastapi import HTTPException
from mcp.types import CallToolRequest, CallToolRequestParams
from litellm.proxy._experimental.mcp_server import operations
from litellm.proxy._experimental.mcp_server.contracts import OperationContext
challenge: Final = HTTPException(401, "Unauthorized", headers={"WWW-Authenticate": "Bearer"})
call: Final = AsyncMock(side_effect=challenge)
monkeypatch.setattr(operations, "call_mcp_tool", call)
with pytest.raises(HTTPException) as failure:
await operations.GatewayOperations().execute(
CallToolRequest(params=CallToolRequestParams(name="upstream-tool", arguments={})),
OperationContext(_caller=None),
)
assert failure.value is challenge
call.assert_awaited_once()
@pytest.mark.asyncio
async def test_tool_handler_preserves_unrelated_protocol_errors(
_mcp_request_ctx, monkeypatch: pytest.MonkeyPatch
) -> None:
from mcp import MCPError
from mcp.types import CallToolRequestParams
from litellm.proxy._experimental.mcp_server import server
failure: Final = MCPError(code=-32602, message="Invalid tool parameters")
execute: Final = AsyncMock(side_effect=failure)
gateway: Final = MagicMock()
gateway.execute = execute
monkeypatch.setattr(server.operations, "GatewayOperations", MagicMock(return_value=gateway))
monkeypatch.setattr(
server, "get_or_extract_auth_context", AsyncMock(return_value=(None, None, None, None, None, None, None))
)
with pytest.raises(MCPError) as caught:
await server.mcp_server_tool_call(_mcp_request_ctx(), CallToolRequestParams(name="example", arguments={}))
assert caught.value is failure
execute.assert_awaited_once()
@pytest.mark.asyncio
@pytest.mark.parametrize("path", ("/mcp", "/github/mcp"))
@pytest.mark.parametrize("has_token", (False, True))
@pytest.mark.parametrize("healthy_companion", (False, True))
async def test_initialize_challenges_missing_upstream_credentials_before_creating_session(
monkeypatch: pytest.MonkeyPatch, path: str, has_token: bool, healthy_companion: bool
) -> None:
from mcp.server.streamable_http_manager import StreamableHTTPSessionManager
from litellm.proxy._experimental.mcp_server import server
from litellm.proxy._types import UserAPIKeyAuth
github: Final = MCPServer(
server_id="github-id", name="github", alias="github", server_name="github",
url="https://github.example/mcp", transport=MCPTransport.http,
auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code",
authorization_url="https://github.example/authorize", token_url="https://github.example/token",
client_id="registered-client",
)
auth: Final = UserAPIKeyAuth(api_key="test-owner", user_id="test-user")
selected: Final = ["github"] if path != "/mcp" else None
monkeypatch.setattr(
server, "extract_mcp_auth_context", AsyncMock(return_value=(auth, None, selected, None, None, None))
)
manager: Final = server.operations.global_mcp_server_manager
public: Final = MCPServer(
server_id="public-id", name="public", alias="public", server_name="public",
url="https://public.example/mcp", transport=MCPTransport.http, auth_type=MCPAuth.none,
)
eligible: Final = [github, public] if healthy_companion and selected is None else [github]
monkeypatch.setattr(manager, "get_mcp_server_by_name", lambda name, **kwargs: next(s for s in eligible if s.alias == name))
monkeypatch.setattr(manager, "has_user_oauth_token", AsyncMock(return_value=has_token))
monkeypatch.setattr(manager, "_ensure_upstream_initialize_instructions_cached", AsyncMock())
monkeypatch.setattr(server.operations, "_get_allowed_mcp_servers", AsyncMock(return_value=eligible))
monkeypatch.setattr(server, "_check_passthrough_upstream_auth", AsyncMock())
monkeypatch.setattr(
server, "session_manager_stateful",
StreamableHTTPSessionManager(app=server.server, stateless=False, json_response=True),
)
monkeypatch.setattr(
server, "session_manager_stateless",
StreamableHTTPSessionManager(app=server.server, stateless=True, json_response=True),
)
await server.initialize_session_managers()
try:
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=server.app), base_url="http://gateway") as client:
response: Final = await client.post(
path, headers={"accept": "application/json, text/event-stream"},
json={"jsonrpc": "2.0", "id": 1, "method": "initialize", "params": {
"protocolVersion": "2025-06-18", "capabilities": {},
"clientInfo": {"name": "test-client", "version": "1"},
}},
)
can_initialize: Final = has_token or (healthy_companion and selected is None)
assert response.status_code == (200 if can_initialize else 401), response.text
if can_initialize:
assert response.json()["result"]["serverInfo"]["name"]
assert response.headers["mcp-session-id"]
else:
assert response.headers["www-authenticate"].startswith("Bearer ")
if path == "/mcp":
assert response.headers["www-authenticate"] == (
'Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp"'
)
assert "mcp-session-id" not in response.headers
finally:
await server.shutdown_session_managers()
@pytest.mark.asyncio
@pytest.mark.parametrize("kind", ("prompts", "resources", "resource_templates"))
async def test_optional_listing_challenges_auth_when_other_server_times_out(
kind: str, monkeypatch: pytest.MonkeyPatch
) -> None:
from litellm.proxy._experimental.mcp_server import operations
blocked: Final = _http_server("blocked", "blocked")
unavailable: Final = _http_server("unavailable", "unavailable")
manager: Final = MCPServerManager()
create: Final = AsyncMock(side_effect=[MCPUpstreamAuthError(401, "Bearer", "blocked"), TimeoutError()])
monkeypatch.setattr(manager, "_create_mcp_client", create)
monkeypatch.setattr(operations, "global_mcp_server_manager", manager)
monkeypatch.setattr(operations, "_prepare_mcp_server_headers", MagicMock(return_value=(None, None)))
monkeypatch.setattr(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[blocked, unavailable]))
with pytest.raises(MCPUpstreamAuthError) as caught:
await getattr(operations, f"_list_mcp_{kind}")()
assert caught.value.status_code == 401
assert caught.value.www_authenticate == "Bearer"
assert caught.value.server_name == "blocked"
assert create.await_count == 2
@pytest.mark.asyncio
@pytest.mark.parametrize("kind", ("prompts", "resources", "resource_templates"))
@pytest.mark.parametrize("healthy_first", (False, True))
async def test_optional_listing_preserves_healthy_duplicate_names(kind: str, healthy_first: bool, monkeypatch: pytest.MonkeyPatch) -> None:
from litellm.proxy._experimental.mcp_server import operations
healthy: Final = _http_server("healthy-id", "duplicate")
blocked: Final = _http_server("blocked-id", "duplicate")
manager: Final = MagicMock()
fetch: Final = AsyncMock(side_effect=[[], MCPUpstreamAuthError(401, "Bearer", "duplicate")] if healthy_first else [MCPUpstreamAuthError(401, "Bearer", "duplicate"), []])
setattr(manager, f"get_{kind}_from_server", fetch)
monkeypatch.setattr(operations, "global_mcp_server_manager", manager)
monkeypatch.setattr(operations, "_prepare_mcp_server_headers", MagicMock(return_value=(None, None)))
monkeypatch.setattr(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[healthy, blocked] if healthy_first else [blocked, healthy]))
assert await getattr(operations, f"_list_mcp_{kind}")() == []
assert fetch.await_count == 2

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
@ -2804,7 +2805,7 @@ async def test_initialize_request_tracks_active_session_after_response_header():
patch( # test-quality-ok: registry is empty in unit tests; key owns one server
"litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers",
new_callable=AsyncMock,
return_value=[MagicMock()],
return_value=[MCPServer(server_id="available", name="available", transport=MCPTransport.http)],
),
patch(
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED",
@ -2957,7 +2958,7 @@ async def test_initialize_request_records_client_name_in_gateway_sessions_report
patch( # test-quality-ok: registry is empty in unit tests; key owns one server
"litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers",
new_callable=AsyncMock,
return_value=[MagicMock()],
return_value=[MCPServer(server_id="available", name="available", transport=MCPTransport.http)],
),
patch( # test-quality-ok: session manager init is a module-level flag; the suite's only seam
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED",
@ -10685,3 +10686,254 @@ async def test_discovery_adapter_preserves_authenticated_context(_mcp_request_ct
context = dispatched.await_args.args[1]
assert context.user_api_key_auth.user_id == "discover-caller"
assert context.mcp_servers == ("allowed",)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"token_states, allowed_ids, expected_status",
(
((False,), None, 401),
((False, False), None, 401),
((False, True), None, None),
((True, False), None, None),
((False, False), {"server-1"}, 401),
((False, True), {"server-1"}, None),
((False,), set(), None),
((), None, None),
),
)
async def test_unified_preflight_challenges_only_when_all_authorized_servers_need_oauth(
monkeypatch: pytest.MonkeyPatch,
token_states: tuple[bool, ...],
allowed_ids: set[str] | None,
expected_status: int | None,
) -> None:
from litellm.proxy._experimental.mcp_server import server as server_module
servers: Final = tuple(
_make_oauth2_server(f"server-{index}").model_copy(update={"server_id": f"server-{index}"})
for index in range(len(token_states))
)
auth: Final = UserAPIKeyAuth(api_key="test-key", user_id="reader")
lookup: Final = AsyncMock(return_value=servers)
monkeypatch.setattr(mcp_operations, "_get_allowed_mcp_servers", lookup)
manager: Final = mcp_operations.global_mcp_server_manager
monkeypatch.setattr(manager, "get_mcp_server_by_name", lambda name, **kwargs: next(s for s in servers if s.alias == name))
tokens: Final = AsyncMock(side_effect=lambda s, user: token_states[int(s.server_id.rsplit("-", 1)[1])])
monkeypatch.setattr(manager, "has_user_oauth_token", tokens)
scope: Final[Scope] = {"type": "http", "method": "POST", "path": "/mcp", "headers": [(b"host", b"gateway")]}
request: Final = server_module._raise_preemptive_401_for_unauthenticated_servers(
scope=scope, mcp_servers=None, oauth2_headers=None, mcp_server_auth_headers=None,
user_api_key_auth=auth, client_ip="127.0.0.1", allowed_server_ids=allowed_ids,
)
if expected_status is not None:
with pytest.raises(HTTPException) as caught:
await request
assert caught.value.status_code == expected_status
assert "www-authenticate" in {key.lower() for key in (caught.value.headers or {})}
else:
await request
lookup.assert_awaited_once_with(user_api_key_auth=auth, mcp_servers=None, client_ip="127.0.0.1")
assert {call.args[0].server_id for call in tokens.await_args_list} == {
s.server_id for s in servers if allowed_ids is None or s.server_id in allowed_ids
}
@pytest.mark.asyncio
@pytest.mark.parametrize("failure", (HTTPException(status_code=503, detail="unavailable"), RuntimeError("discovery failed"), asyncio.CancelledError()))
async def test_unified_preflight_does_not_misclassify_discovery_failure_as_oauth(
monkeypatch: pytest.MonkeyPatch, failure: BaseException
) -> None:
from litellm.proxy._experimental.mcp_server import server as server_module
upstream: Final = _make_oauth2_server("unavailable")
monkeypatch.setattr(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[upstream]))
manager: Final = mcp_operations.global_mcp_server_manager
monkeypatch.setattr(manager, "get_mcp_server_by_name", lambda *args, **kwargs: upstream)
discovery: Final = AsyncMock(side_effect=failure)
monkeypatch.setattr(manager, "ensure_oauth_metadata_discovered", discovery)
request: Final = server_module._raise_preemptive_401_for_unauthenticated_servers(
scope={"type": "http", "path": "/mcp", "method": "POST", "headers": []},
mcp_servers=None, oauth2_headers=None, mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(user_id="reader"), client_ip=None,
)
with pytest.raises(type(failure)) as caught:
await request
if not isinstance(failure, asyncio.CancelledError):
assert caught.value is failure
discovery.assert_awaited_once_with(upstream)
@pytest.mark.asyncio
async def test_unified_preflight_preserves_delegated_oauth_challenge(monkeypatch: pytest.MonkeyPatch) -> None:
from litellm.proxy._experimental.mcp_server import server as server_module
upstream: Final = _make_oauth2_server("delegated", delegate_auth_to_upstream=True)
monkeypatch.setattr(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[upstream]))
monkeypatch.setattr(
mcp_operations.global_mcp_server_manager, "get_mcp_server_by_name", lambda *args, **kwargs: upstream
)
with pytest.raises(HTTPException) as caught:
await server_module._raise_preemptive_401_for_unauthenticated_servers(
scope={"type": "http", "path": "/mcp", "method": "POST", "headers": [(b"host", b"gateway")]},
mcp_servers=None, oauth2_headers=None, mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(user_id="reader"), client_ip=None,
)
assert caught.value.status_code == 401
assert (caught.value.headers or {})["www-authenticate"] == (
'Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp/delegated"'
)
@pytest.mark.asyncio
@pytest.mark.parametrize("unified", (False, True))
@pytest.mark.parametrize(
"auth_type,bridge,upstream_status,expected_challenge",
(
(MCPAuth.none, False, 401, "gateway"),
(MCPAuth.oauth_delegate, False, 401, "gateway"),
(MCPAuth.true_passthrough, True, 401, "gateway"),
(MCPAuth.true_passthrough, False, 401, "upstream"),
(MCPAuth.true_passthrough, False, 200, None),
),
)
async def test_preflight_preserves_client_forwarded_auth_challenges(
monkeypatch: pytest.MonkeyPatch,
unified: bool,
auth_type: MCPAuth,
bridge: bool,
upstream_status: int,
expected_challenge: str | None,
) -> None:
from litellm.proxy._experimental.mcp_server import server as server_module
upstream: Final = _client_forwarded_mode_server("forwarded", auth_type).model_copy(
update={
"dcr_bridge": bridge,
"oauth_passthrough": auth_type == MCPAuth.none,
"extra_headers": ["Authorization"] if auth_type == MCPAuth.none else None,
}
)
monkeypatch.setattr(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[upstream]))
monkeypatch.setattr(mcp_operations.global_mcp_server_manager, "get_mcp_server_by_name", lambda *args, **kwargs: upstream)
probe: Final = AsyncMock(return_value=(upstream_status, "Bearer realm=upstream"))
monkeypatch.setattr(server_module, "_probe_upstream_auth", probe)
request: Final = server_module._raise_preemptive_401_for_unauthenticated_servers(
scope={"type": "http", "path": "/mcp", "method": "POST", "headers": [(b"host", b"gateway")]},
mcp_servers=None if unified else [upstream.name],
oauth2_headers=None,
mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(user_id="reader"),
client_ip=None,
)
if expected_challenge is None:
await request
else:
with pytest.raises(HTTPException) as caught:
await request
assert caught.value.status_code == 401
expected: Final = (
'Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp/forwarded"'
if expected_challenge == "gateway"
else "Bearer realm=upstream"
)
assert (caught.value.headers or {})["www-authenticate"] == expected
assert probe.await_count == int(auth_type == MCPAuth.true_passthrough and not bridge)
@pytest.mark.asyncio
async def test_preflight_gateway_subject_uses_server_resource_metadata(monkeypatch: pytest.MonkeyPatch) -> None:
from litellm.proxy._experimental.mcp_server import server as server_module
upstream: Final = _make_oauth2_server("managed")
auth: Final = UserAPIKeyAuth(user_id="reader")
auth.mcp_admitted_user_subject = True
manager: Final = mcp_operations.global_mcp_server_manager
monkeypatch.setattr(manager, "get_mcp_server_by_name", lambda *args, **kwargs: upstream)
monkeypatch.setattr(manager, "has_user_oauth_token", AsyncMock(return_value=False))
with pytest.raises(HTTPException) as caught:
await server_module._raise_preemptive_401_for_unauthenticated_servers(
scope={"type": "http", "path": "/mcp/managed", "method": "POST", "headers": [(b"host", b"gateway")]},
mcp_servers=["managed"],
oauth2_headers=None,
mcp_server_auth_headers=None,
user_api_key_auth=auth,
client_ip=None,
)
assert caught.value.status_code == 401
assert (caught.value.headers or {})["www-authenticate"] == (
'Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp/managed"'
)
@pytest.mark.asyncio
async def test_preflight_does_not_request_oauth_for_excluded_server(monkeypatch: pytest.MonkeyPatch) -> None:
from litellm.proxy._experimental.mcp_server import server as server_module
upstream: Final = _make_oauth2_server("excluded")
manager: Final = mcp_operations.global_mcp_server_manager
monkeypatch.setattr(manager, "get_mcp_server_by_name", lambda *args, **kwargs: upstream)
discovery: Final = AsyncMock(return_value=upstream)
tokens: Final = AsyncMock(return_value=False)
monkeypatch.setattr(manager, "ensure_oauth_metadata_discovered", discovery)
monkeypatch.setattr(manager, "has_user_oauth_token", tokens)
await server_module._raise_preemptive_401_for_unauthenticated_servers(
scope={"type": "http", "path": "/mcp/excluded", "method": "POST", "headers": []},
mcp_servers=["excluded"],
oauth2_headers=None,
mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(user_id="reader"),
client_ip=None,
allowed_server_ids={"other"},
)
discovery.assert_not_awaited()
tokens.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("healthy_peer", (False, True))
@pytest.mark.parametrize("failure_first", (False, True))
@pytest.mark.parametrize("failure", (HTTPException(503, "unavailable"), RuntimeError("discovery failed"), asyncio.CancelledError()))
async def test_unified_preflight_preserves_usable_peer_during_discovery_failure(
monkeypatch: pytest.MonkeyPatch,
caplog: pytest.LogCaptureFixture,
healthy_peer: bool,
failure_first: bool,
failure: BaseException,
) -> None:
from litellm.proxy._experimental.mcp_server import server as server_module
unavailable: Final = _make_oauth2_server("unavailable")
peer: Final = _make_oauth2_server("peer")
servers: Final = [unavailable, peer] if failure_first else [peer, unavailable]
monkeypatch.setattr(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=servers))
manager: Final = mcp_operations.global_mcp_server_manager
async def discover(server: MCPServer) -> MCPServer:
if server.server_id == unavailable.server_id:
raise failure
return server
monkeypatch.setattr(manager, "ensure_oauth_metadata_discovered", discover)
tokens: Final = AsyncMock(return_value=healthy_peer)
monkeypatch.setattr(manager, "has_user_oauth_token", tokens)
request: Final = server_module._raise_preemptive_401_for_unauthenticated_servers(
scope={"type": "http", "path": "/mcp", "method": "POST", "headers": [(b"host", b"gateway")]},
mcp_servers=None,
oauth2_headers=None,
mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(user_id="reader"),
client_ip=None,
)
if healthy_peer and not isinstance(failure, asyncio.CancelledError):
with caplog.at_level("WARNING", logger="LiteLLM"):
await request
assert any("unavailable" in record.getMessage() for record in caplog.records)
assert str(failure) not in caplog.text
else:
with pytest.raises(type(failure)) as caught:
await request
if not isinstance(failure, asyncio.CancelledError):
assert caught.value is failure
tokens.assert_awaited_once()
assert tokens.await_args.args[0].server_id == peer.server_id

View file

@ -1844,7 +1844,11 @@ class TestMCPServerManager:
server, mcp_auth_header, extra_headers, stdio_env, subject_token=None, **kwargs
): # pragma: no cover - helper
captured["subject_token"] = subject_token
return AsyncMock()
return AsyncMock(
discovery_auth_fingerprint=AsyncMock(return_value="test-credential-hash"),
list_prompts=AsyncMock(return_value=[]),
list_resources=AsyncMock(return_value=[]),
)
manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client)
manager._fetch_tools_with_timeout = AsyncMock(return_value=[])
@ -2049,9 +2053,8 @@ class TestMCPServerManager:
assert mock_log.warning.called
@pytest.mark.asyncio
async def test_call_non_passthrough_does_not_opt_into_raise_on_error(self):
"""Non-client-forwarded auth types keep the default call_tool masking (raise_on_error stays
off), so this relay is scoped to the pass-through modes and cannot regress api_key/OBO calls."""
async def test_call_static_auth_preserves_success_with_transport_errors_enabled(self):
"""A static credential uses the same transport-auth error channel while preserving tool results."""
server = MCPServer(
server_id="ak-call",
name="ak-call-server",
@ -2078,7 +2081,7 @@ class TestMCPServerManager:
)
assert result.is_error is False
assert mock_client.call_tool.call_args.kwargs.get("raise_on_error") is not True
assert mock_client.call_tool.call_args.kwargs.get("raise_on_error") is True
def _token_exchange_server(self, server_id: str) -> "MCPServer":
return MCPServer(
@ -2633,7 +2636,11 @@ class TestMCPServerManager:
server, mcp_auth_header, extra_headers, stdio_env, subject_token=None, **kwargs
): # pragma: no cover - helper
captured["subject_token"] = subject_token
return AsyncMock()
return AsyncMock(
discovery_auth_fingerprint=AsyncMock(return_value="test-credential-hash"),
list_prompts=AsyncMock(return_value=[]),
list_resources=AsyncMock(return_value=[]),
)
manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client)
await call(manager)
@ -6911,7 +6918,7 @@ class TestMCPServerManager:
# Create mock client that tracks call_tool usage
mock_client = AsyncMock()
async def mock_call_tool(params, host_progress_callback=None, allow_input_required=False):
async def mock_call_tool(params, host_progress_callback=None, allow_input_required=False, raise_on_error=False):
# Return a mock CallToolResult
result = MagicMock(spec=CallToolResult)
result.content = [{"type": "text", "text": "Tool executed successfully"}]
@ -13146,6 +13153,7 @@ class TestLitellmAdmissionKeyIsNeverTheSubjectToken:
client: Final = AsyncMock()
client.call_tool = AsyncMock(return_value=CallToolResult(content=[], isError=False))
client.list_prompts = AsyncMock(return_value=[])
client.discovery_auth_fingerprint = AsyncMock(return_value="test-credential-hash")
client.read_resource = AsyncMock(return_value=ReadResourceResult(contents=[]))
manager._create_mcp_client = AsyncMock(return_value=client)
return manager
@ -13925,8 +13933,12 @@ async def test_discovery_cache_empty_results_and_failures(kind: str, outcome: st
"templates": manager.get_resource_templates_from_server,
}[kind]
with _mcp_upstream(upstream.respond):
assert await operation(_discovery_server(), None) == []
assert await operation(_discovery_server(), None) == []
for _ in range(2):
if outcome == "failure":
with pytest.raises(MCPServerListError, match="discovery"):
await operation(_discovery_server(), None)
else:
assert await operation(_discovery_server(), None) == []
assert upstream.initializes == (2 if outcome == "failure" else 1)
if outcome == "failure":
upstream.outcome = "supported"
@ -13946,7 +13958,8 @@ async def test_discovery_cache_retries_failed_pagination_before_caching_complete
"templates": manager.get_resource_templates_from_server,
}[kind]
with _mcp_upstream(upstream.respond):
assert await operation(_discovery_server(), None) == []
with pytest.raises(MCPServerListError, match="discovery"):
await operation(_discovery_server(), None)
assert upstream.initializes == 1
upstream.outcome = "paged"
recovered: Final = await operation(_discovery_server(), None)
@ -14197,6 +14210,7 @@ async def test_discovery_cache_bounds_detached_fetches_without_dropping_results(
@pytest.mark.asyncio
async def test_discovery_cache_tracks_resolved_credentials_across_workers() -> None:
from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError
import respx
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import StaticHeaderAuth
from litellm.proxy._experimental.mcp_server.outbound_credentials.resolver import UpstreamCredentialProvider
@ -14254,7 +14268,9 @@ async def test_discovery_cache_tracks_resolved_credentials_across_workers() -> N
assert upstream.initializes == 4
source.token = None
for manager in managers:
assert await manager.get_prompts_from_server(server, user) == []
with pytest.raises(MCPUpstreamAuthError) as failure:
await manager.get_prompts_from_server(server, user)
assert failure.value.status_code == 401
assert upstream.initializes == 4

View file

@ -23,7 +23,7 @@ const unscoped = (client_origin: string): ConnectFlowStatus => ({
connected: null,
});
const renderBanner = (clientOrigin: string) =>
const renderBanner = (clientOrigin: string, selectedServers: string[] = ["github"]) =>
render(
<ConnectFlowBanner
flowHandle="flow-handle-123"
@ -31,11 +31,12 @@ const renderBanner = (clientOrigin: string) =>
accessToken="tok"
onConnected={vi.fn()}
failed={false}
selectedServers={selectedServers}
/>,
);
describe("ConnectFlowBanner", () => {
it("posts only the flow handle to the proxy /authorize/complete as a full-page form", () => {
it("posts the flow handle and selected servers to the proxy /authorize/complete as a full-page form", () => {
const { container } = renderBanner("https://claude.ai");
const form = container.querySelector("form")!;
@ -44,7 +45,14 @@ describe("ConnectFlowBanner", () => {
expect(screen.getByDisplayValue("flow-handle-123")).toHaveAttribute("name", "flow");
expect(form.innerHTML).not.toContain("token");
expect(screen.getByRole("button", { name: /finish connecting/i })).toBeInTheDocument();
expect(screen.queryByRole("button", { name: "Cancel" })).not.toBeInTheDocument();
expect(screen.getByRole("button", { name: "Cancel" })).toBeInTheDocument();
expect(new FormData(form).getAll("selected_servers")).toEqual(["github"]);
});
it("requires a selection before offering Finish but still allows cancellation", () => {
renderBanner("https://claude.ai", []);
expect(screen.queryByRole("button", { name: /finish connecting/i })).not.toBeInTheDocument();
expect(screen.getByRole("button", { name: "Cancel" })).toBeInTheDocument();
});
it("offers manual delivery only for a loopback client, posted only when checked", () => {

View file

@ -11,6 +11,7 @@ interface Props {
accessToken: string;
onConnected: () => void;
failed: boolean;
selectedServers?: readonly string[];
}
/** Finish remains an explicit POST because a cross-site navigation must never mint a code. */
@ -51,11 +52,17 @@ const copyFor = (flow: ConnectFlowStatus | undefined, failed: boolean): readonly
];
};
const ConnectFlowBanner: React.FC<Props> = ({ flowHandle, flow, accessToken, onConnected, failed }) => {
const ConnectFlowBanner: React.FC<Props> = ({
flowHandle,
flow,
accessToken,
onConnected,
failed,
selectedServers = [],
}) => {
const action = `${getProxyBaseUrl()}/authorize/complete`;
const state = failed || flow === undefined ? "stale" : flow.state;
const canFinish = state === "unscoped" || (state !== "stale" && flow?.connected === true);
const canCancel = state !== "unscoped";
const canFinish = state === "unscoped" ? selectedServers.length > 0 : state !== "stale" && flow?.connected === true;
const loopbackClient = isLoopbackOrigin(flow?.client_origin ?? null);
const vendorServer =
state === "interactive" && flow?.connected === false && flow.server_id !== null
@ -85,6 +92,10 @@ const ConnectFlowBanner: React.FC<Props> = ({ flowHandle, flow, accessToken, onC
)}
<form method="POST" action={action}>
<input type="hidden" name="flow" value={flowHandle} />
{state === "unscoped" &&
selectedServers.map((server) => (
<input key={server} type="hidden" name="selected_servers" value={server} />
))}
{canFinish && (
<button
type="submit"
@ -93,16 +104,14 @@ const ConnectFlowBanner: React.FC<Props> = ({ flowHandle, flow, accessToken, onC
Finish connecting
</button>
)}
{canCancel && (
<button
type="submit"
name="decision"
value="deny"
className="ml-2 h-[38px] rounded-md border px-4 text-sm font-semibold text-foreground hover:bg-accent/40"
>
Cancel
</button>
)}
<button
type="submit"
name="decision"
value="deny"
className="ml-2 h-[38px] rounded-md border px-4 text-sm font-semibold text-foreground hover:bg-accent/40"
>
Cancel
</button>
{loopbackClient && (
<label className="mt-2 flex items-center gap-2 text-[13px] text-muted-foreground">
<input type="checkbox" name="delivery" value="manual" />

View file

@ -1,5 +1,5 @@
import { afterEach, describe, expect, it, vi } from "vitest";
import { act, render, screen, waitFor } from "@testing-library/react";
import { act, fireEvent, render, screen, waitFor } from "@testing-library/react";
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import ConnectFlowSurface from "./ConnectFlowSurface";
import { fetchConnectFlow } from "@/components/networking";
@ -23,7 +23,12 @@ vi.mock("@/components/networking", async (importOriginal) => ({
}));
vi.mock("@/components/chat/MCPAppsPanel", async (importOriginal) => ({
...(await importOriginal<typeof import("@/components/chat/MCPAppsPanel")>()),
default: () => <div data-testid="mcp-apps-panel" />,
default: ({ selectedServers, onChange }: { selectedServers: string[]; onChange: (ids: string[]) => void }) => (
<div data-testid="mcp-apps-panel">
<span data-testid="selected-servers">{selectedServers.join(",")}</span>
<button onClick={() => onChange(["github-id", "slack-id"])}>Select upstreams</button>
</div>
),
}));
vi.mock("@/hooks/useUserMcpOAuthFlow", () => ({
useUserMcpOAuthFlow: ({ onSuccess: success }: { onSuccess: () => void }) => {
@ -40,12 +45,16 @@ const flow = (state: "unscoped" | "interactive" | "m2m" | "stale", connected: bo
connected,
});
const renderSurface = () =>
render(
<QueryClientProvider client={new QueryClient({ defaultOptions: { queries: { retry: false } } })}>
<ConnectFlowSurface accessToken="token-123" selectedServers={[]} onChange={vi.fn()} />
</QueryClientProvider>,
const renderSurface = (selectedServers: string[] = [], onChange = vi.fn()) => {
const client = new QueryClient({ defaultOptions: { queries: { retry: false } } });
const surface = () => (
<QueryClientProvider client={client}>
<ConnectFlowSurface accessToken="token-123" selectedServers={selectedServers} onChange={onChange} />
</QueryClientProvider>
);
const result = render(surface());
return { ...result, refresh: () => result.rerender(surface()) };
};
afterEach(() => {
state.oauthReturn = null;
@ -57,7 +66,7 @@ afterEach(() => {
describe("ConnectFlowSurface", () => {
it.each([
{ result: flow("unscoped"), grid: true, finish: true, cancel: false, oauthStarts: 0 },
{ result: flow("unscoped"), grid: true, finish: false, cancel: true, oauthStarts: 0 },
{ result: flow("interactive", false), grid: false, finish: false, cancel: true, oauthStarts: 1 },
{ result: flow("interactive", true), grid: false, finish: true, cancel: true, oauthStarts: 0 },
{ result: flow("m2m", true), grid: false, finish: true, cancel: true, oauthStarts: 0 },
@ -77,6 +86,49 @@ describe("ConnectFlowSurface", () => {
},
);
it("submits the selected upstreams with the protected flow", async () => {
state.connectFlow = "flow-handle-123";
vi.mocked(fetchConnectFlow).mockResolvedValue(flow("unscoped"));
const chatSelectionChanged = vi.fn();
renderSurface(["github", "slack"], chatSelectionChanged);
await screen.findByRole("button", { name: "Select upstreams" });
expect(screen.queryByRole("button", { name: /finish connecting/i })).not.toBeInTheDocument();
expect(screen.getByTestId("selected-servers")).toBeEmptyDOMElement();
fireEvent.click(screen.getByRole("button", { name: "Select upstreams" }));
const finish = await screen.findByRole("button", { name: /finish connecting/i });
const submitted = new FormData((finish as HTMLButtonElement).form!);
expect(submitted.get("flow")).toBe("flow-handle-123");
expect(submitted.getAll("selected_servers")).toEqual(["github-id", "slack-id"]);
expect(chatSelectionChanged).not.toHaveBeenCalled();
});
it("isolates selections between flow handles and preserves ordinary chat selections", async () => {
const chatSelectionChanged = vi.fn();
vi.mocked(fetchConnectFlow).mockResolvedValue(flow("unscoped"));
const view = renderSurface(["github"], chatSelectionChanged);
expect(screen.getByTestId("selected-servers")).toHaveTextContent("github");
state.connectFlow = "first-flow";
view.refresh();
await screen.findByRole("button", { name: "Select upstreams" });
expect(screen.getByTestId("selected-servers")).toBeEmptyDOMElement();
fireEvent.click(screen.getByRole("button", { name: "Select upstreams" }));
await screen.findByRole("button", { name: /finish connecting/i });
state.connectFlow = "second-flow";
view.refresh();
await screen.findByRole("button", { name: "Select upstreams" });
expect(screen.getByTestId("selected-servers")).toBeEmptyDOMElement();
expect(screen.queryByRole("button", { name: /finish connecting/i })).not.toBeInTheDocument();
state.connectFlow = null;
view.refresh();
expect(screen.getByTestId("selected-servers")).toHaveTextContent("github");
expect(chatSelectionChanged).not.toHaveBeenCalled();
fireEvent.click(screen.getByRole("button", { name: "Select upstreams" }));
expect(chatSelectionChanged).toHaveBeenCalledWith(["github-id", "slack-id"]);
});
it("keeps the grid and Finish hidden until the gateway accepts a handle", () => {
state.connectFlow = "flow-handle-123";
vi.mocked(fetchConnectFlow).mockReturnValue(new Promise(() => {}));

View file

@ -1,6 +1,6 @@
"use client";
import React, { useEffect } from "react";
import React, { useEffect, useState } from "react";
import { useRouter, useSearchParams } from "next/navigation";
import { useQuery } from "@tanstack/react-query";
import MCPAppsPanel from "@/components/chat/MCPAppsPanel";
@ -19,6 +19,11 @@ const ConnectFlowSurface: React.FC<Props> = ({ accessToken, selectedServers, onC
const searchParams = useSearchParams();
const oauthReturn = searchParams.get("mcpOauthReturn");
const connectFlow = searchParams.get("connect_flow");
const [flowSelection, setFlowSelection] = useState<{ handle: string | null; serverIds: string[] }>({
handle: null,
serverIds: [],
});
const selectedServerIds = flowSelection.handle === connectFlow ? flowSelection.serverIds : [];
useEffect(() => {
if (oauthReturn) {
@ -48,9 +53,15 @@ const ConnectFlowSurface: React.FC<Props> = ({ accessToken, selectedServers, onC
accessToken={accessToken}
onConnected={refetch}
failed={isError}
selectedServers={selectedServerIds}
/>
{flow?.state === "unscoped" && (
<MCPAppsPanel accessToken={accessToken} selectedServers={selectedServers} onChange={onChange} connectMode />
<MCPAppsPanel
accessToken={accessToken}
selectedServers={selectedServerIds}
onChange={(serverIds) => setFlowSelection({ handle: connectFlow, serverIds })}
connectMode
/>
)}
</>
);

View file

@ -125,7 +125,7 @@ describe("MCPAppsPanel connected-app reachability (LIT-4861)", () => {
vi.mocked(fetchMCPServers).mockResolvedValue(connectServers);
vi.mocked(listMCPTools).mockResolvedValue({ tools: [] });
renderConnectPanel(true, ["reachable_srv", "unreachable_srv"]);
renderConnectPanel(true, ["s-reach", "s-unreach"]);
expect(await screen.findByText("reachable_srv")).toBeInTheDocument();
expect(vi.mocked(fetchMCPServers)).toHaveBeenCalledWith("tok", undefined, true);
@ -300,3 +300,26 @@ describe("MCPAppsPanel connected-app reachability (LIT-4861)", () => {
expect(screen.queryByText("revoked_srv")).not.toBeInTheDocument();
});
});
it.each([true, false])("selects an unambiguous upstream in connect mode=%s", async (connectMode) => {
const onChange = vi.fn();
vi.mocked(fetchMCPServers).mockResolvedValue([
{
server_id: "selected-id",
server_name: "github",
alias: "github-selected",
auth_type: "none",
connected_app_reachable: true,
},
{ server_id: "other-id", server_name: "other", alias: "github", auth_type: "none", connected_app_reachable: true },
] as MCPServer[]);
vi.mocked(listMCPTools).mockResolvedValue({ tools: [] });
render(
<QueryClientProvider client={new QueryClient({ defaultOptions: { queries: { retry: false } } })}>
<MCPAppsPanel accessToken="tok" selectedServers={[]} onChange={onChange} connectMode={connectMode} />
</QueryClientProvider>,
);
fireEvent.click(await screen.findByText("github"));
fireEvent.click(await screen.findByRole("button", { name: "Connect", exact: true }));
await waitFor(() => expect(onChange).toHaveBeenCalledWith([connectMode ? "selected-id" : "github"]));
});

View file

@ -140,6 +140,11 @@ const MCPAppsPanel: React.FC<Props> = ({ accessToken, selectedServers, onChange,
const nameOf = (s: MCPServer) => s.server_name ?? s.alias ?? s.server_id;
const selectionOf = useCallback(
(server: MCPServer) => (connectMode ? server.server_id : nameOf(server)),
[connectMode],
);
const detailServer = servers.find((s) => s.server_id === detailServerId);
const connectUnavailabilityLabel = useCallback(
@ -239,19 +244,20 @@ const MCPAppsPanel: React.FC<Props> = ({ accessToken, selectedServers, onChange,
.filter(
(s) =>
oauthConnected.has(s.server_id) &&
!selectedServersRef.current.includes(nameOf(s)) &&
!selectedServersRef.current.includes(selectionOf(s)) &&
connectUnavailabilityLabel(s) === null,
)
.map(nameOf);
.map(selectionOf);
if (namesToAdd.length > 0) {
onChangeRef.current([...selectedServersRef.current, ...namesToAdd]);
}
}, [oauthConnected, connectUnavailabilityLabel]);
}, [oauthConnected, connectUnavailabilityLabel, selectionOf]);
const handleToggle = async (server: MCPServer, checked: boolean) => {
const serverName = nameOf(server);
const selection = selectionOf(server);
if (!checked) {
onChange(selectedServers.filter((s) => s !== serverName));
onChange(selectedServers.filter((s) => s !== selection));
setOauthConnected((prev) => {
const next = new Set(prev);
next.delete(server.server_id);
@ -268,8 +274,8 @@ const MCPAppsPanel: React.FC<Props> = ({ accessToken, selectedServers, onChange,
return;
}
if (connectableNow(server.server_id) === undefined) return;
if (!selectedServersRef.current.includes(serverName)) {
onChange([...selectedServersRef.current, serverName]);
if (!selectedServersRef.current.includes(selection)) {
onChange([...selectedServersRef.current, selection]);
}
} catch {
toast.warning(`Could not load tools for ${serverName}`);
@ -308,7 +314,7 @@ const MCPAppsPanel: React.FC<Props> = ({ accessToken, selectedServers, onChange,
/>
);
}
if (selectedServers.includes(nameOf(server))) {
if (selectedServers.includes(selectionOf(server))) {
return <span className="w-[7px] h-[7px] rounded-full bg-success shrink-0" />;
}
return null;
@ -328,12 +334,12 @@ const MCPAppsPanel: React.FC<Props> = ({ accessToken, selectedServers, onChange,
name.toLowerCase().includes(query.toLowerCase()) ||
(s.description ?? "").toLowerCase().includes(query.toLowerCase());
const matchesTab =
activeTab === "all" || (selectedServers.includes(name) && connectUnavailabilityLabel(s) === null);
activeTab === "all" || (selectedServers.includes(selectionOf(s)) && connectUnavailabilityLabel(s) === null);
return matchesQuery && matchesTab;
});
const connectedCount = servers.filter(
(s) => selectedServers.includes(nameOf(s)) && connectUnavailabilityLabel(s) === null,
(s) => selectedServers.includes(selectionOf(s)) && connectUnavailabilityLabel(s) === null,
).length;
const emptyStateText = () => {
@ -348,7 +354,7 @@ const MCPAppsPanel: React.FC<Props> = ({ accessToken, selectedServers, onChange,
if (detailServer) {
const name = nameOf(detailServer);
const isConnected = selectedServers.includes(name);
const isConnected = selectedServers.includes(selectionOf(detailServer));
const isTogglingOn = togglingOn.has(name);
const color = getAvatarColor(name);
@ -388,7 +394,7 @@ const MCPAppsPanel: React.FC<Props> = ({ accessToken, selectedServers, onChange,
n.delete(detailServer.server_id);
return n;
});
onChangeRef.current(selectedServersRef.current.filter((s) => s !== name));
onChangeRef.current(selectedServersRef.current.filter((s) => s !== selectionOf(detailServer)));
}}
className="font-semibold h-[38px] min-w-[110px]"
>

View file

@ -25781,6 +25781,8 @@ export interface components {
delivery?: string | null;
/** Flow */
flow: string;
/** Selected Servers */
selected_servers?: string[] | null;
/** Team Id */
team_id?: string | null;
};