mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge 671701e1ab into 025292e75b
This commit is contained in:
commit
920e80d44d
20 changed files with 1721 additions and 500 deletions
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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": [
|
||||
{
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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", () => {
|
||||
|
|
|
|||
|
|
@ -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" />
|
||||
|
|
|
|||
|
|
@ -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(() => {}));
|
||||
|
|
|
|||
|
|
@ -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
|
||||
/>
|
||||
)}
|
||||
</>
|
||||
);
|
||||
|
|
|
|||
|
|
@ -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"]));
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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]"
|
||||
>
|
||||
|
|
|
|||
2
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
2
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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;
|
||||
};
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue