feat(mcp): add portable catalog pagination (#44446)
Some checks are pending
CI Coverage / assert-ci-coverage (push) Waiting to run
CodeQL / Analyze (actions) (push) Waiting to run
CodeQL / Analyze (javascript-typescript) (push) Waiting to run
CodeQL / Analyze (python) (push) Waiting to run
CodSpeed Benchmarks / benchmarks (push) Waiting to run
Helm unit test / unit-test (push) Waiting to run
Lens Worker Image / lens-worker-image (push) Waiting to run
Publish basedpyright base counts / publish (push) Waiting to run
Scorecard supply-chain security / Scorecard analysis (push) Waiting to run
Code Quality Checks / code-quality (push) Waiting to run
Code Quality Checks / python-310-import-smoke (push) Waiting to run
UI Unit Tests / ui-unit-tests (push) Waiting to run
LiteLLM Rust / rust-wheel (push) Waiting to run
LiteLLM Rust / rust-lint (push) Waiting to run
LiteLLM Rust / rust-test (push) Waiting to run
Terraform Provider / gofmt, vet, build, test (push) Waiting to run
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Waiting to run
Unit Tests: Documentation Validation / documentation (push) Waiting to run
Unit Tests: Proxy DB Operations / Lens Python 3.10 (push) Waiting to run
Unit Tests: Proxy DB Operations / assert-shard-coverage (push) Waiting to run
Unit Tests: Proxy DB Operations / auth-checks (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / budgets (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / custom-logging (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / db-and-spend (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / endpoints-and-responses (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / guardrails-hooks (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / jwt-and-keys (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / key-generation (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / logging-misc (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-runtime (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-server-core (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-utils (push) Blocked by required conditions
Unit Tests / Vertex AI (push) Blocked by required conditions
Unit Tests / misc (push) Blocked by required conditions
Unit Tests / Build the Rust bridge (push) Waiting to run
Unit Tests / caching-local (push) Blocked by required conditions
Unit Tests / core-utils (push) Blocked by required conditions
Unit Tests / enterprise-managed-files (push) Blocked by required conditions
Unit Tests / enterprise-package (push) Blocked by required conditions
Unit Tests / enterprise-routing (push) Blocked by required conditions
Unit Tests / integrations (push) Blocked by required conditions
Unit Tests / OpenAI and Meta Providers (push) Blocked by required conditions
Unit Tests / All Other Providers (push) Blocked by required conditions
Unit Tests / misc-dirs (push) Blocked by required conditions
Unit Tests / proxy-auth (push) Blocked by required conditions
Unit Tests / proxy-endpoints (push) Blocked by required conditions
Unit Tests / unit (push) Blocked by required conditions
GitHub Actions Security Analysis / zizmor (push) Waiting to run
Unit Tests / proxy-extras (push) Blocked by required conditions
Unit Tests / proxy-feature-endpoints (push) Blocked by required conditions
Unit Tests / proxy-hooks-client (push) Blocked by required conditions
Unit Tests / proxy-infra (push) Blocked by required conditions
Unit Tests / proxy-infra-root (push) Blocked by required conditions
Unit Tests / proxy-server (push) Blocked by required conditions
Unit Tests / responses-caching-types (push) Blocked by required conditions

* feat(mcp): add portable authorized catalog pagination

* fix(mcp): preserve routing and isolation across catalog pages

* fix(mcp): preserve bare calls after complete catalog pages

* test(mcp): verify bare calls with cached database catalog revisions

* test(mcp): cover refreshed authority and failed catalog continuations

* refactor(mcp): narrow pagination interfaces and preserve listed metadata

* fix(mcp): publish listed metadata only after successful pagination

* fix(mcp): defer bare routes until aggregate listing succeeds

* fix(mcp): refresh key access groups on catalog pages

* fix(mcp): accept read-only optional catalog headers

* fix(mcp): preserve consolidated listing metadata after restack

* fix(mcp): propagate unexpected catalog continuation failures

* fix(mcp): preserve key grants when a user record is absent

---------

Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
This commit is contained in:
joshua-berri 2026-10-06 18:54:50 -07:00 • committed by GitHub
parent 4909bd9e8c
commit d1cfe17518
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
26 changed files with 2861 additions and 345 deletions

View file

@ -12,7 +12,7 @@ from contextlib import AbstractAsyncContextManager
from functools import partial
from importlib.metadata import version
from types import MappingProxyType
from typing import Final, TypeAlias, TypeVar, cast
from typing import TYPE_CHECKING, Final, TypeAlias, TypeVar, cast
import anyio
import httpx2
@ -47,9 +47,13 @@ from mcp.types import (
InitializeRequestParams,
InitializeResult,
InputRequiredResult,
ListPromptsRequest,
ListPromptsResult,
ListResourcesRequest,
ListResourcesResult,
ListResourceTemplatesResult,
ListToolsRequest,
ListToolsResult,
PaginatedRequestParams,
PaginatedResult,
Prompt,
@ -89,6 +93,9 @@ from litellm.types.mcp import (
without_header,
)
if TYPE_CHECKING:
from litellm.proxy._experimental.mcp_server.contracts import CatalogListRequest, CatalogListResult
def to_basic_auth(auth_value: str) -> str:
"""Convert auth value to Basic Auth format."""
@ -830,6 +837,51 @@ class MCPClient:
return factory
async def list_page(self, request: "CatalogListRequest") -> "CatalogListResult":
from mcp.types import INVALID_PARAMS
params: Final = request.params or PaginatedRequestParams()
if isinstance(request, ListToolsRequest):
return await self.list_tools_page(params)
async def fetch(session: ClientSession) -> "CatalogListResult":
capabilities: Final = session.server_capabilities
empty: Final = (
ListPromptsResult(prompts=[])
if isinstance(request, ListPromptsRequest)
else ListResourcesResult(resources=[])
if isinstance(request, ListResourcesRequest)
else ListResourceTemplatesResult(resource_templates=[])
)
supported: Final = capabilities is None or (
capabilities.prompts is not None
if isinstance(request, ListPromptsRequest)
else capabilities.resources is not None
)
if not supported:
if params.cursor is not None:
raise MCPError(
code=INVALID_PARAMS, message="Upstream catalog became unavailable; start a fresh listing"
)
return empty
try:
if isinstance(request, ListPromptsRequest):
return await session.list_prompts(params=params)
if isinstance(request, ListResourcesRequest):
return await session.list_resources(params=params)
return await session.list_resource_templates(params=params)
except MCPError as error:
if error.error.code == METHOD_NOT_FOUND and params.cursor is None:
return empty
raise
with anyio.fail_after(max(self.timeout, MCP_TOOL_LISTING_TIMEOUT)):
return await self.run_with_session(fetch, quiet_on_error=True)
async def list_tools_page(self, params: PaginatedRequestParams) -> ListToolsResult:
with anyio.fail_after(max(self.timeout, MCP_TOOL_LISTING_TIMEOUT)):
return await self.run_with_session(lambda session: session.list_tools(params=params), quiet_on_error=True)
async def list_tools(self, raise_on_error: bool = False) -> list[MCPTool]:
"""List available tools from the server.
@ -846,7 +898,7 @@ class MCPClient:
# A per-server timeout above the global default extends the whole-walk deadline
listing_deadline: Final = max(self.timeout, MCP_TOOL_LISTING_TIMEOUT)
tools: Final = await self.run_with_session(
partial(list_tools_with_pagination, listing_deadline=listing_deadline),
partial(list_tools_with_pagination, listing_deadline=listing_deadline, require_complete=raise_on_error),
quiet_on_error=raise_on_error,
)
tool_count: Final = len(tools)

View file

@ -2,10 +2,10 @@ import json
from typing import Final, Literal
import anyio
from mcp import ClientSession
from mcp import ClientSession, MCPError
from mcp.types import INTERNAL_ERROR, PaginatedRequestParams
from mcp.types import CallToolRequestParams as MCPCallToolRequestParams
from mcp.types import CallToolResult as MCPCallToolResult
from mcp.types import PaginatedRequestParams
from mcp.types import Tool as MCPTool
from openai.types.chat import ChatCompletionToolParam
from openai.types.responses.function_tool_param import FunctionToolParam
@ -99,7 +99,7 @@ def transform_mcp_tool_to_anthropic_tool(mcp_tool: MCPTool) -> AnthropicMessages
async def list_tools_with_pagination(
session: ClientSession, listing_deadline: float | None = None
session: ClientSession, listing_deadline: float | None = None, *, require_complete: bool = False
) -> list[MCPTool]: # mutable-ok: list return contract
"""Collect tools from every tools/list page by following nextCursor.
@ -137,6 +137,10 @@ async def list_tools_with_pagination(
"MCP server repeated a tools/list cursor while listing tools; returning %s tools collected so far",
len(tools),
)
if require_complete:
raise MCPError(
code=INTERNAL_ERROR, message="Upstream tool discovery is incomplete: repeated cursor"
)
return tools
seen_cursors.add(next_cursor)
cursor = next_cursor
@ -146,6 +150,8 @@ async def list_tools_with_pagination(
MCP_TOOL_LISTING_MAX_PAGES,
len(tools),
)
if require_complete:
raise MCPError(code=INTERNAL_ERROR, message="Upstream tool discovery is incomplete: pagination limit")
return tools
verbose_logger.warning(
@ -153,6 +159,8 @@ async def list_tools_with_pagination(
effective_deadline,
len(tools),
)
if require_complete:
raise MCPError(code=INTERNAL_ERROR, message="Upstream tool discovery is incomplete: listing deadline")
return tools

View file

@ -0,0 +1,15 @@
# Gateway catalog pagination
Tools, prompts, resources, and resource templates retain upstream page boundaries. Follow the gateway's opaque `nextCursor` until it is absent. Each request checks the caller's current access and tool permissions. A cursor can be replayed and can be sent to another replica serving the same configuration
Set the same nonempty `LITELLM_SALT_KEY` on every replica. Pagination state uses authenticated encryption with purpose-separated HKDF keys derived only from this value. The master key is never a fallback. Complete single-page lists and direct tool calls work without a salt key; a listing that needs continuation returns an actionable configuration error instead of truncated results. Some SDKs automatically list tools when validating a tool-call response, so those SDK calls also require a salt when that listing is paginated
Cursors expire ten minutes after the first page. Continuations do not extend that deadline. Rotating the salt invalidates all outstanding cursors; clients must start a new listing. During a rolling key change, replicas with different keys cannot accept each other's cursors. Coordinate the change across the fleet
A changed registry, caller scope, or available upstream revision requires a fresh listing. When an upstream exposes a string or integer `_meta.revision`, subsequent pages must retain it. Otherwise consistency follows that upstream's own cursor guarantees. Repeated upstream cursors and the existing upstream page limit stop traversal with an explicit error
Listing failures retain per-server outcome metadata. An incomplete upstream catalog cannot establish a bare tool-name route. Complete initial pages retain legacy bare-name routing. Use the server-prefixed names returned by the gateway for paginated catalogs
For a scoped rollout, keep the previous source build serving the control pool and send only selected clients to a separate candidate pool. All candidate replicas must share configuration and salt. Verify page one on one candidate replica and continuation on another, plus a fresh listing and tool call on the control pool. Do not mirror tool calls between pools
For rollback, stop sending new requests to the candidate pool, drain its in-flight operations, and return selected clients to the control pool. Clients must discard candidate cursors and start a fresh listing when crossing versions; older gateways do not validate these cursors. Keep the registry-revision migration installed when rolling back pagination. Verify a fresh listing and a tool call after switching pools

View file

@ -73,6 +73,7 @@ from litellm.repositories.table_repositories import (
MCPServerRepository,
)
from litellm.types.mcp_server.mcp_server_manager import MCPServer
from litellm.types.proxy.auth.auth_checks import UserNotFoundError
if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient
@ -1201,6 +1202,31 @@ class MCPRequestHandler:
verbose_logger.warning("Failed to resolve per-team MCP rpm limits for admitted subject: %s", e)
return None
@staticmethod
async def refresh_catalog_authority(auth: UserAPIKeyAuth | None) -> UserAPIKeyAuth | None:
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
if auth is None:
return None
refreshed = auth.model_copy()
if auth.mcp_admitted_user_subject and auth.user_id:
current = await MCPRequestHandler.reload_admitted_user(auth.user_id, requires_fresh_policy=True)
refreshed.object_permission = current.object_permission
refreshed.object_permission_id = current.object_permission_id
refreshed.user_role = current.user_role
refreshed.org_id = current.org_id
elif auth.via_virtual_key and auth.api_key and auth.api_key != LITELLM_PROXY_MASTER_KEY_ALIAS:
current = await MCPRequestHandler._reload_admitted_key(auth.api_key, check_db_only=True)
refreshed.object_permission = current.object_permission
refreshed.object_permission_id = current.object_permission_id
refreshed.team_id = current.team_id
refreshed.org_id = current.org_id
refreshed.project_id = current.project_id
refreshed.user_id = current.user_id
refreshed.access_group_ids = current.access_group_ids
refreshed.requires_fresh_policy = True
return refreshed
@staticmethod
async def _reload_admitted_key(key_hash: str, *, check_db_only: bool = False) -> UserAPIKeyAuth:
"""Reload the live key record an admitted envelope references and re-check live policy.
@ -1244,6 +1270,8 @@ class MCPRequestHandler:
if not MCPRequestHandler._admitted_key_is_active(key_object):
raise HTTPException(status_code=401, detail="Invalid or expired credential")
await MCPRequestHandler._reject_if_admitted_owner_scim_deactivated(key_object)
key_object.api_key = key_hash
key_object.via_virtual_key = True
return key_object
@staticmethod
@ -3154,6 +3182,8 @@ class MCPRequestHandler:
ttl=get_management_object_ttl(user_api_key_cache),
)
return object_permission_id
except UserNotFoundError:
return None
except Exception as e: # noqa: BLE001 # Legacy callers retain their existing optional user-ceiling behavior
if check_db_only:
raise HTTPException(503, "User policy is unavailable") from e

View file

@ -3,29 +3,38 @@
from __future__ import annotations
import asyncio
import base64
import hashlib
import json
from collections import UserDict
from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, MutableMapping, Sequence
from contextlib import asynccontextmanager
from contextlib import ExitStack, asynccontextmanager
from contextvars import ContextVar
from dataclasses import dataclass, replace
from functools import wraps
from functools import partial, wraps
from itertools import chain
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, ParamSpec, TypeVar
from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter
from litellm._logging import verbose_logger
if TYPE_CHECKING:
from pydantic import BaseModel
from mcp.types import ListToolsResult, PaginatedRequestParams, PaginatedResult
from litellm.proxy._experimental.mcp_server.contracts import CatalogListRequest, CatalogListResult, OperationContext
from litellm.proxy._experimental.mcp_server.db import OAuthCredentialPayload
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing, ServerOutcome
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
from litellm.types.mcp_server.mcp_server_manager import MCPServer
from litellm.types.mcp_server.tool_registry import MCPTool
_P = ParamSpec("_P")
_R = TypeVar("_R")
_Page = TypeVar("_Page", bound="PaginatedResult")
_JSON_VALUE: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue)
_OUTCOME_VALUES: Final[TypeAdapter[dict[str, JsonValue]]] = TypeAdapter(dict[str, JsonValue])
class _OperationRoutes(UserDict[str, str]):
@ -95,7 +104,7 @@ def _snapshot(manager: MCPServerManager, database_identity: str) -> CatalogSnaps
)
class TargetCatalog:
class CatalogSnapshots:
def __init__(self, manager: MCPServerManager) -> None:
self.manager = manager
self._refresh_lock = asyncio.Lock()
@ -576,3 +585,632 @@ def global_manager() -> MCPServerManager:
def public_catalog_operation(function: Callable[_P, Awaitable[_R]]) -> Callable[_P, Awaitable[_R]]:
return catalog_operation(global_manager)(function)
async def paginate_catalog(
*,
method: str,
cursor: str | None,
caller_scope: str,
snapshot: str,
server_ids: tuple[str, ...],
fetch: Callable[[str, str | None], Awaitable[_Page]],
now: int,
) -> tuple[tuple[_Page, ...], str | None, Mapping[str, JsonValue]]:
from mcp.shared.exceptions import MCPError
from mcp.types import INVALID_PARAMS
from pydantic import ValidationError
from litellm.constants import MCP_TOOL_LISTING_MAX_PAGES
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error
from litellm.proxy._experimental.mcp_server.state_tokens import open_state, seal_state
ordered: Final = tuple(sorted(server_ids))
purpose: Final = "mcp.catalog.list.v1:" + method
if cursor is None:
state = _ListingState(
caller_scope=caller_scope,
snapshot=snapshot,
expires_at=now + 600,
positions=tuple(_UpstreamPosition(server_id=server_id) for server_id in ordered),
)
else:
opened: Final = open_state(cursor, purpose=purpose, now=now)
if isinstance(opened, Error):
raise MCPError(code=INVALID_PARAMS, message=opened.error.value)
try:
state = _ListingState.model_validate_json(json.dumps(opened.ok))
except ValidationError as error:
raise MCPError(code=INVALID_PARAMS, message="Invalid pagination state; start a fresh listing") from error
if (
state.caller_scope != caller_scope
or state.snapshot != snapshot
or tuple(position.server_id for position in state.positions) != ordered
or state.expires_at <= now
):
raise MCPError(code=INVALID_PARAMS, message="Pagination scope or snapshot changed; start a fresh listing")
async def advance(position: _UpstreamPosition) -> tuple[_UpstreamPosition, _Page | None]:
if position.complete:
return position, None
result: Final = await fetch(position.server_id, position.cursor)
revision: Final = (result.meta or {}).get("revision")
available_revision: Final = revision if isinstance(revision, (str, int)) else None
if position.revision is not None and position.revision != available_revision:
raise MCPError(code=INVALID_PARAMS, message="Upstream snapshot changed; start a fresh listing")
following: Final = result.next_cursor or None
fingerprint: Final = (
base64.urlsafe_b64encode(hashlib.sha256(following.encode()).digest()).decode("ascii").rstrip("=")
if following is not None
else None
)
if fingerprint in position.seen:
raise MCPError(code=INVALID_PARAMS, message="Upstream repeated a pagination cursor; start a fresh listing")
if following is not None and len(position.seen) + 1 >= MCP_TOOL_LISTING_MAX_PAGES:
raise MCPError(code=INVALID_PARAMS, message="Upstream pagination limit reached; start a fresh listing")
return (
_UpstreamPosition(
server_id=position.server_id,
cursor=following,
complete=following is None,
revision=available_revision,
seen=position.seen + ((fingerprint,) if fingerprint is not None else ()),
),
result,
)
tasks: Final = tuple(asyncio.create_task(advance(position)) for position in state.positions)
try:
results: Final = await asyncio.gather(*tasks)
finally:
for task in tasks:
task.cancel()
await asyncio.gather(*tasks, return_exceptions=True)
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import SERVER_OUTCOMES_META_KEY
pages: Final = tuple(result for _, result in results if result is not None)
page_outcomes: Final = (
_OUTCOME_VALUES.validate_python((result.meta or {}).get(SERVER_OUTCOMES_META_KEY, {})) for result in pages
)
outcomes: Final = dict(state.failures) | dict(chain.from_iterable(outcome.items() for outcome in page_outcomes))
failures: Final = {
key: value for key, value in outcomes.items() if isinstance(value, dict) and value.get("tag") != "ok"
}
following_state: Final = state.model_copy(
update={
"positions": tuple(position for position, _ in results),
"failures": failures,
}
)
next_cursor: str | None = None
if any(not position.complete for position in following_state.positions):
sealed: Final = seal_state(
_JSON_VALUE.validate_json(following_state.model_dump_json()),
purpose=purpose,
expires_at=following_state.expires_at,
now=now,
)
if isinstance(sealed, Error):
raise MCPError(code=INVALID_PARAMS, message=sealed.error.value)
next_cursor = sealed.ok
return pages, next_cursor, MappingProxyType(outcomes)
async def list_tools_page(
*,
cursor: str | None,
caller_scope: str,
snapshot: str,
server_ids: tuple[str, ...],
fetch: Callable[[str, str | None], Awaitable[ListToolsResult]],
now: int,
) -> ListToolsResult:
from mcp.types import ListToolsResult
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import SERVER_OUTCOMES_META_KEY
pages, next_cursor, outcomes = await paginate_catalog(
method="tools/list",
cursor=cursor,
caller_scope=caller_scope,
snapshot=snapshot,
server_ids=server_ids,
fetch=fetch,
now=now,
)
return ListToolsResult(
tools=list(chain.from_iterable(page.tools for page in pages)),
next_cursor=next_cursor,
_meta={SERVER_OUTCOMES_META_KEY: dict(outcomes)} if outcomes else None,
)
class _UpstreamPosition(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid", strict=True)
server_id: str
cursor: str | None = None
complete: bool = False
revision: str | int | None = None
seen: tuple[str, ...] = ()
class _ListingState(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid", strict=True)
caller_scope: str
snapshot: str
expires_at: int
positions: tuple[_UpstreamPosition, ...]
failures: Mapping[str, JsonValue] = {}
async def get_filtered_server_tools(
server: MCPServer,
*,
context: OperationContext,
allowed_mcp_servers: Sequence[MCPServer],
prefetched_oauth_creds: Mapping[str, OAuthCredentialPayload],
params: PaginatedRequestParams | None = None,
record_listing: bool = False,
listing_updates: ExitStack | None = None,
) -> tuple[ListToolsResult, ServerOutcome]:
from mcp.types import ListToolsResult
from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ServerListOk, classify_list_exception
from litellm.proxy._experimental.mcp_server.operations import (
_get_byok_credential,
_get_user_oauth_extra_headers_from_db,
_prepare_mcp_server_headers,
apply_display_name_overrides,
filter_tools_by_allowed_tools,
filter_tools_by_key_team_permissions,
global_mcp_server_manager,
)
from litellm.types.mcp import MCPAuth
user_api_key_auth, mcp_auth_header, _, mcp_server_auth_headers, oauth2_headers, raw_headers, client_ip = (
context.legacy_auth()
)
mcp_proxy_mode: Final = context.mcp_proxy_mode
if server is None:
return ListToolsResult(tools=[]), ServerListOk(tool_count=0)
server_auth_header, extra_headers = _prepare_mcp_server_headers(
server=server,
mcp_server_auth_headers=mcp_server_auth_headers,
mcp_auth_header=mcp_auth_header,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
scope_servers=list(allowed_mcp_servers),
)
# Prefer server-stored per-user OAuth when configured, so a stale
# Authorization header from the MCP client cannot override Redis/DB
# (same issue as call_tool in mcp_server_manager: VS Code caches tokens).
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415
to_server_spec,
)
# A server migrated to the v2 resolver gets its token from the resolver at connect
# time; building it here would double-resolve and be shadowed by the v2 graft. The
# preemptive 401 already challenged a missing token, so one exists for the connect.
migrated_to_v2: Final = to_server_spec(server) is not None
if (
not migrated_to_v2
and server.auth_type == MCPAuth.oauth2
and getattr(server, "needs_user_oauth_token", False)
and user_api_key_auth is not None
):
db_headers: Final = await _get_user_oauth_extra_headers_from_db(
server,
user_api_key_auth,
prefetched_creds=prefetched_oauth_creds,
)
if db_headers:
extra_headers = db_headers
# If still no OAuth2 token, fall back to pre-fetched creds (non-stale-client path)
elif not migrated_to_v2 and extra_headers is None and server.auth_type == MCPAuth.oauth2:
extra_headers = await _get_user_oauth_extra_headers_from_db(
server,
user_api_key_auth,
prefetched_creds=prefetched_oauth_creds,
)
catalog_auth_header: Final = server_auth_header
if server.is_byok and server.auth_type != MCPAuth.oauth2 and server_auth_header is None:
server_auth_header = await _get_byok_credential(server, user_api_key_auth)
try:
from litellm.proxy.proxy_server import proxy_logging_obj
listed_generation: Final = global_mcp_server_manager.listed_tools_generation(server.server_id)
if params is None:
page = ListToolsResult(
tools=await global_mcp_server_manager._get_tools_from_server(
server=server,
mcp_auth_header=server_auth_header,
extra_headers=extra_headers,
add_prefix=True,
raw_headers=raw_headers,
client_ip=client_ip,
user_api_key_auth=user_api_key_auth,
oauth2_headers=oauth2_headers,
proxy_logging_obj=proxy_logging_obj,
catalog_auth_header=catalog_auth_header,
record_listing=False,
)
)
else:
page = await global_mcp_server_manager.get_tools_page(
server=server,
mcp_auth_header=server_auth_header,
extra_headers=extra_headers,
add_prefix=True,
raw_headers=raw_headers,
client_ip=client_ip,
user_api_key_auth=user_api_key_auth,
oauth2_headers=oauth2_headers,
proxy_logging_obj=proxy_logging_obj,
params=params,
catalog_auth_header=catalog_auth_header,
record_listing=False,
listing_updates=listing_updates,
)
tools: Final = page.tools
filtered_tools = filter_tools_by_allowed_tools(tools, server)
filtered_tools = await filter_tools_by_key_team_permissions(
tools=filtered_tools,
server_id=server.server_id,
user_api_key_auth=user_api_key_auth,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller
from litellm.proxy._experimental.mcp_server.utils import strip_known_server_prefix
record: Final = partial(
global_mcp_server_manager.record_listed_tools,
server,
[tool.model_copy(update={"name": strip_known_server_prefix(tool.name, server)}) for tool in filtered_tools],
ListedToolsCaller(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=catalog_auth_header,
raw_headers=raw_headers,
oauth2_headers=oauth2_headers,
),
listed_generation,
record_listing=record_listing,
continuation=params is not None and params.cursor is not None,
)
if listing_updates is None:
record()
else:
listing_updates.callback(record)
if mcp_proxy_mode:
from litellm.proxy._experimental.mcp_server.tool_search import with_mcp_proxy_identity
filtered_tools = [with_mcp_proxy_identity(tool, server.server_id) for tool in filtered_tools]
else:
filtered_tools = apply_display_name_overrides(filtered_tools, server)
verbose_logger.debug(
"Successfully fetched %s tools from server %s, %s after filtering",
len(tools),
server.name,
len(filtered_tools),
)
return page.model_copy(update={"tools": filtered_tools}), ServerListOk(tool_count=len(filtered_tools))
except MCPUpstreamAuthError as e:
# Absorb so one unauthenticated server does not empty every other server's
# tools. Surfacing the upstream 401 to the client as a re-auth challenge is
# intentionally not done here: raising from this list handler cannot produce a
# 401 + WWW-Authenticate (the MCP session manager serializes it as a JSON-RPC
# error). Single-server routes surface it via the request-scope preemptive
# check in _raise_preemptive_401_for_unauthenticated_servers instead.
verbose_logger.debug("MCP list_tools: omitting %s; it needs upstream auth", server.name)
return ListToolsResult(tools=[]), classify_list_exception(e)
except Exception as e:
verbose_logger.exception("Error getting tools from server %s: %s", server.name, e)
return ListToolsResult(tools=[]), classify_list_exception(e)
def _caller_scope(context: OperationContext, servers: Sequence[MCPServer]) -> str:
from litellm.proxy._experimental.mcp_server.utils import upstream_credential_headers
caller: Final = context.user_api_key_auth
headers: Final = context.raw_headers or {}
credential_names: Final = (
upstream_credential_headers(headers)
| frozenset({"authorization"})
| frozenset(map(str.lower, chain.from_iterable(server.extra_headers or () for server in servers)))
)
material: Final = (
caller.model_dump(include={"api_key", "user_id", "team_id", "org_id", "end_user_id", "user_role"}, mode="json")
if caller is not None
else None,
tuple(sorted(context.mcp_servers)) if context.mcp_servers is not None else None,
context.mcp_auth_header,
{key: dict(value) for key, value in (context.mcp_server_auth_headers or {}).items()},
dict(context.oauth2_headers or {}),
{key.lower(): value for key, value in headers.items() if key.lower() in credential_names},
context.client_ip,
context.protocol_version,
context.mcp_proxy_mode,
)
return hashlib.sha256(json.dumps(material, sort_keys=True, separators=(",", ":")).encode()).hexdigest()
async def aggregate_gateway_tools(
context: OperationContext,
params: PaginatedRequestParams,
allowed: Sequence[MCPServer],
prefetched: Mapping[str, OAuthCredentialPayload],
*,
record_listing: bool = False,
) -> AggregateToolListing:
import time
from mcp.types import PaginatedRequestParams
from pydantic import TypeAdapter
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
SERVER_OUTCOMES_META_KEY,
AggregateToolListing,
ServerOutcome,
)
from litellm.proxy._experimental.mcp_server.operations import _aggregate_server_key, global_mcp_server_manager
async with global_mcp_server_manager.catalog.operation() as snapshot:
servers: Final = {server.server_id: server for server in allowed}
listing_updates: Final = ExitStack()
async def fetch(server_id: str, cursor: str | None) -> ListToolsResult:
result, outcome = await get_filtered_server_tools(
servers[server_id],
context=context,
allowed_mcp_servers=allowed,
prefetched_oauth_creds=prefetched,
params=PaginatedRequestParams(cursor=cursor),
record_listing=record_listing,
listing_updates=listing_updates,
)
if cursor is not None and outcome.tag != "ok":
from mcp.shared.exceptions import MCPError
from mcp.types import INVALID_PARAMS
raise MCPError(code=INVALID_PARAMS, message="Upstream continuation failed; start a fresh listing")
return result.model_copy(
update={
"meta": {
**(result.meta or {}),
SERVER_OUTCOMES_META_KEY: {
_aggregate_server_key(servers[server_id]): outcome.model_dump(mode="json")
},
}
}
)
result: Final = await list_tools_page(
cursor=params.cursor,
caller_scope=_caller_scope(context, allowed),
snapshot=snapshot.identity,
server_ids=tuple(servers),
fetch=fetch,
now=int(time.time()),
)
listing_updates.close()
return AggregateToolListing(
tools=result.tools,
outcomes=TypeAdapter(dict[str, ServerOutcome]).validate_python(
(result.meta or {}).get(SERVER_OUTCOMES_META_KEY, {})
),
next_cursor=result.next_cursor,
)
async def list_gateway_tools(
context: OperationContext, params: PaginatedRequestParams, *, log_list_tools_to_spendlogs: bool = True
) -> ListToolsResult:
from mcp.types import ListToolsResult
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import SERVER_OUTCOMES_META_KEY, outcome_wire_value
from litellm.proxy._experimental.mcp_server.operations import _list_mcp_tools
caller, auth, servers, server_headers, oauth_headers, headers, client_ip = context.legacy_auth()
listing: Final = await _list_mcp_tools(
user_api_key_auth=caller,
mcp_auth_header=auth,
mcp_servers=servers,
mcp_server_auth_headers=server_headers,
oauth2_headers=oauth_headers,
raw_headers=headers,
client_ip=client_ip,
params=params,
protocol_version=context.protocol_version,
log_list_tools_to_spendlogs=log_list_tools_to_spendlogs,
record_listing=True,
list_tools_log_source="mcp_protocol",
)
return ListToolsResult(
tools=listing.tools,
next_cursor=listing.next_cursor,
_meta={
SERVER_OUTCOMES_META_KEY: {key: outcome_wire_value(outcome) for key, outcome in listing.outcomes.items()}
}
if listing.outcomes
else None,
)
async def list_gateway_catalog(
context: OperationContext, request: CatalogListRequest, *, log_list_tools_to_spendlogs: bool = True
) -> CatalogListResult:
import time
from mcp.types import (
ListToolsRequest,
PaginatedRequestParams,
)
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
from litellm.proxy._experimental.mcp_server.operations import (
_get_allowed_mcp_servers,
global_mcp_server_manager,
raise_denied_scoped_mcp_access,
)
context = replace(context, _caller=await MCPRequestHandler.refresh_catalog_authority(context.user_api_key_auth))
params: Final = request.params or PaginatedRequestParams()
if isinstance(request, ListToolsRequest):
return await list_gateway_tools(context, params, log_list_tools_to_spendlogs=log_list_tools_to_spendlogs)
caller: Final = context.user_api_key_auth
scope: Final = context.mcp_servers
client_ip: Final = context.client_ip
async with global_mcp_server_manager.catalog.operation() as snapshot:
allowed: Final = await _get_allowed_mcp_servers(
user_api_key_auth=caller, mcp_servers=scope, client_ip=client_ip
)
if scope and not allowed:
await raise_denied_scoped_mcp_access(
requested_names=list(scope), user_api_key_auth=caller, client_ip=client_ip
)
servers: Final = {server.server_id: server for server in allowed}
async def fetch(server_id: str, cursor: str | None) -> CatalogListResult:
server: Final = servers[server_id]
from mcp.shared.exceptions import MCPError
from mcp.types import INVALID_PARAMS
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
SERVER_OUTCOMES_META_KEY,
classify_list_exception,
)
from litellm.proxy._experimental.mcp_server.operations import _aggregate_server_key
try:
page: Final = await fetch_optional_catalog_page(context, request, server, allowed, cursor)
return page.model_copy(
update={
"meta": {
key: value for key, value in (page.meta or {}).items() if key != SERVER_OUTCOMES_META_KEY
}
}
)
except Exception as error:
if cursor is not None:
raise MCPError(
code=INVALID_PARAMS, message="Upstream continuation failed; start a fresh listing"
) from error
return combine_optional_catalog(
request,
(),
None,
{
"litellm.ai/server_outcomes": {
_aggregate_server_key(server): classify_list_exception(error).model_dump(mode="json")
}
},
)
pages, next_cursor, outcomes = await paginate_catalog(
method=request.method,
cursor=params.cursor,
caller_scope=_caller_scope(context, allowed),
snapshot=snapshot.identity,
server_ids=tuple(servers),
fetch=fetch,
now=int(time.time()),
)
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
SERVER_OUTCOMES_META_KEY,
ServerOutcome,
outcome_wire_value,
)
typed_outcomes: Final = TypeAdapter(dict[str, ServerOutcome]).validate_python(outcomes)
return combine_optional_catalog(
request,
pages,
next_cursor,
_OUTCOME_VALUES.validate_python(
{SERVER_OUTCOMES_META_KEY: {key: outcome_wire_value(value) for key, value in typed_outcomes.items()}}
)
if outcomes
else None,
)
async def fetch_optional_catalog_page(
context: OperationContext,
request: CatalogListRequest,
server: MCPServer,
allowed: Sequence[MCPServer],
cursor: str | None,
) -> CatalogListResult:
from mcp.types import PaginatedRequestParams
from litellm.proxy._experimental.mcp_server.operations import _prepare_mcp_server_headers, global_mcp_server_manager
caller, auth, _, server_headers, oauth_headers, raw_headers, client_ip = context.legacy_auth()
auth_header, extra_headers = _prepare_mcp_server_headers(
server=server,
mcp_server_auth_headers=server_headers,
mcp_auth_header=auth,
oauth2_headers=oauth_headers,
raw_headers=raw_headers,
user_api_key_auth=caller,
scope_servers=list(allowed),
)
return await global_mcp_server_manager.get_optional_catalog_page(
server,
request.model_copy(update={"params": PaginatedRequestParams(cursor=cursor)}),
caller,
mcp_auth_header=auth_header,
extra_headers=extra_headers,
raw_headers=raw_headers,
client_ip=client_ip,
)
def combine_optional_catalog(
request: CatalogListRequest,
pages: Sequence[CatalogListResult],
next_cursor: str | None,
meta: Mapping[str, JsonValue] | None,
) -> CatalogListResult:
from mcp.types import (
ListPromptsRequest,
ListPromptsResult,
ListResourcesRequest,
ListResourcesResult,
ListResourceTemplatesResult,
)
if isinstance(request, ListPromptsRequest):
return ListPromptsResult(
prompts=list(chain.from_iterable(page.prompts for page in pages if isinstance(page, ListPromptsResult))),
next_cursor=next_cursor,
_meta=dict(meta) if meta is not None else None,
)
if isinstance(request, ListResourcesRequest):
return ListResourcesResult(
resources=list(
chain.from_iterable(page.resources for page in pages if isinstance(page, ListResourcesResult))
),
next_cursor=next_cursor,
_meta=dict(meta) if meta is not None else None,
)
return ListResourceTemplatesResult(
resource_templates=list(
chain.from_iterable(
page.resource_templates for page in pages if isinstance(page, ListResourceTemplatesResult)
)
),
next_cursor=next_cursor,
_meta=dict(meta) if meta is not None else None,
)

View file

@ -1,19 +1,41 @@
from __future__ import annotations
from collections.abc import Mapping
from copy import deepcopy
from dataclasses import dataclass, field
from datetime import datetime
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, Literal, Protocol
from typing import TYPE_CHECKING, Final, Literal, Protocol, TypeAlias
from litellm.proxy._experimental.mcp_server.tool_outcome import WireCompat
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
if TYPE_CHECKING:
from mcp.types import (
ListPromptsRequest,
ListPromptsResult,
ListResourcesRequest,
ListResourcesResult,
ListResourceTemplatesRequest,
ListResourceTemplatesResult,
ListToolsRequest,
ListToolsResult,
)
CatalogListRequest: TypeAlias = (
ListToolsRequest | ListPromptsRequest | ListResourcesRequest | ListResourceTemplatesRequest
)
CatalogListResult: TypeAlias = (
ListToolsResult | ListPromptsResult | ListResourcesResult | ListResourceTemplatesResult
)
from litellm.proxy._experimental.mcp_server.server_resolution import ResolvedMCPServer
class TargetCatalog(Protocol):
async def list(self, context: OperationContext, request: CatalogListRequest) -> CatalogListResult: ...
async def resolve(
self,
server_id: str,
@ -23,7 +45,7 @@ class TargetCatalog(Protocol):
not_found_detail: Mapping[str, str],
forbidden_detail: Mapping[str, str],
non_admin_missing: Literal["not_found", "forbidden"],
) -> "ResolvedMCPServer": ...
) -> ResolvedMCPServer: ...
def copy_caller(auth: UserAPIKeyAuth | None) -> UserAPIKeyAuth | None:

View file

@ -63,6 +63,7 @@ domain per the MCP spec's ``_meta`` key format so it cannot collide with spec-re
class AggregateToolListing(NamedTuple):
tools: list[MCPTool]
outcomes: dict[str, ServerOutcome]
next_cursor: str | None = None
def _iter_upstream_responses(exc: BaseException) -> Iterator[httpx.Response | httpx2.Response]:

View file

@ -24,7 +24,7 @@ from collections.abc import (
MutableMapping,
Sequence,
)
from contextlib import asynccontextmanager
from contextlib import ExitStack, asynccontextmanager
from dataclasses import dataclass, replace
from functools import lru_cache
from itertools import chain, groupby
@ -44,6 +44,11 @@ from mcp.types import (
GetPromptRequestParams,
GetPromptResult,
InputRequiredResult,
ListPromptsRequest,
ListResourcesRequest,
ListResourceTemplatesRequest,
ListToolsResult,
PaginatedRequestParams,
Prompt,
ResourceTemplate,
)
@ -215,6 +220,7 @@ from litellm.types.utils import CallTypes
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._experimental.mcp_server.contracts import CatalogListRequest, CatalogListResult
from litellm.types.mcp_server.mcp_toolset import MCPToolset
try:
@ -1970,9 +1976,9 @@ class MCPServerManager:
self._template_discovery_cache = _DiscoveryCache[ResourceTemplate](
discovery_ttl, discovery_clock, TypeAdapter(tuple[ResourceTemplate, ...])
)
from litellm.proxy._experimental.mcp_server.catalog import TargetCatalog
from litellm.proxy._experimental.mcp_server.catalog import CatalogSnapshots
self.catalog = TargetCatalog(self)
self.catalog = CatalogSnapshots(self)
self.registry: dict[str, MCPServer] = {}
self._openapi_health_probes: Callable[[str], _OpenAPIHealthProbe] = lru_cache(maxsize=128)(_OpenAPIHealthProbe)
self.config_mcp_servers: dict[str, MCPServer] = {}
@ -4267,27 +4273,47 @@ class MCPServerManager:
*,
catalog_auth_header: str | dict[str, str] | None | EllipsisType = ...,
record_listing: bool = False,
) -> Sequence[MCPTool]:
"""
Helper method to get tools from a single MCP server with prefixed names.
) -> list[MCPTool]:
result: Final = await self.get_tools_page(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=extra_headers,
add_prefix=add_prefix,
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
oauth2_headers=oauth2_headers,
client_ip=client_ip,
proxy_logging_obj=proxy_logging_obj,
catalog_auth_header=catalog_auth_header,
record_listing=record_listing,
)
return result.tools
Args:
server (MCPServer): The server to query tools from
mcp_auth_header: Optional auth header for MCP server
catalog_auth_header: The header the client supplied, keying the caller's catalog slot;
defaults to ``mcp_auth_header``
record_listing: Record the served catalog into the caller's listed-tools slot; only a
listing actually served to the caller sets it
Returns:
List[MCPTool]: List of tools available on the server with prefixed names
"""
async def get_tools_page(
self,
server: MCPServer,
mcp_auth_header: str | dict[str, str] | None = None,
extra_headers: dict[str, str] | None = None,
add_prefix: bool = True,
raw_headers: dict[str, str] | None = None,
user_api_key_auth: UserAPIKeyAuth | None = None,
oauth2_headers: dict[str, str] | None = None,
client_ip: str | None = None,
proxy_logging_obj: ProxyLogging | None = None,
*,
params: PaginatedRequestParams | None = None,
catalog_auth_header: str | dict[str, str] | None | EllipsisType = ...,
record_listing: bool = False,
listing_updates: ExitStack | None = None,
) -> ListToolsResult:
from litellm.proxy._experimental.mcp_server.tool_registry import (
global_mcp_tool_registry,
)
if self._skip_blocked_stdio_listing(server, "tool"):
return []
if params is not None and params.cursor is not None:
raise RuntimeError("Upstream catalog is unavailable")
return ListToolsResult(tools=[])
verbose_logger.debug("Connecting to url: %s", server.url)
verbose_logger.info("_get_tools_from_server for %s...", server.name)
@ -4401,10 +4427,17 @@ class MCPServerManager:
server, guarded_openapi, listed_caller, listed_generation, record_listing=record_listing
)
if not add_prefix:
return guarded_openapi
return [t.model_copy(update={"name": registered_names[t.name]}) for t in guarded_openapi]
return ListToolsResult(tools=list(guarded_openapi))
return ListToolsResult(
tools=[t.model_copy(update={"name": registered_names[t.name]}) for t in guarded_openapi]
)
else:
tools = await self._fetch_tools_with_timeout(client, server.name)
page: Final = (
await client.list_tools_page(params)
if params is not None
else ListToolsResult(tools=await self._fetch_tools_with_timeout(client, server.name))
)
tools = page.tools
self._remember_upstream_initialize_instructions(server, client)
guarded_tools: Final = await self._guard_tool_catalog(
@ -4415,13 +4448,17 @@ class MCPServerManager:
raw_headers=raw_headers,
)
prefixed_or_original_tools: Final = self._create_prefixed_tools(
guarded_tools, server, add_prefix=add_prefix
guarded_tools,
server,
add_prefix=add_prefix,
register_bare_names=params is None or (params.cursor is None and not page.next_cursor),
listing_updates=listing_updates,
)
self.record_listed_tools(
server, guarded_tools, listed_caller, listed_generation, record_listing=record_listing
)
return prefixed_or_original_tools
return page.model_copy(update={"tools": prefixed_or_original_tools})
except MCPUpstreamAuthError as upstream_auth_error:
# Pass-through 401 must surface to single-server routes so the
@ -4537,6 +4574,7 @@ class MCPServerManager:
generation: int | None = None,
*,
record_listing: bool = True,
continuation: bool = False,
) -> None:
"""Store the catalog served to ``caller``. ``generation`` is the server's listed-tools generation
read before the listing's upstream fetch; the record is skipped when it no longer matches."""
@ -4545,8 +4583,9 @@ class MCPServerManager:
if generation is not None and generation != self._listed_tools_generations.get(server.server_id, 0):
return
identity: Final = self._listed_tools_identity(server, caller)
listing: Final = MappingProxyType({tool.name: tool for tool in tools})
existing: Final = self._listed_tools_by_server_id.get(server.server_id, MappingProxyType({}))
prior: Final = existing.get(identity, {}) if continuation else {}
listing: Final = MappingProxyType({**prior, **{tool.name: tool for tool in tools}})
shared: Final = existing.get(None)
callers: Final = tuple((key, value) for key, value in existing.items() if key not in (None, identity))
evicted: Final = 0 if identity is None else max(len(callers) + 1 - _LISTED_TOOLS_CALLERS_PER_SERVER, 0)
@ -4602,6 +4641,67 @@ class MCPServerManager:
)
return True
async def get_optional_catalog_page(
self,
server: MCPServer,
request: "CatalogListRequest",
user_api_key_auth: UserAPIKeyAuth | None,
*,
mcp_auth_header: str | dict[str, str] | None = None,
extra_headers: Mapping[str, str] | None = None,
raw_headers: Mapping[str, str] | None = None,
client_ip: str | None = None,
) -> "CatalogListResult":
from mcp.types import ListPromptsResult, ListResourcesResult, ListResourceTemplatesResult
if self._skip_blocked_stdio_listing(server, "catalog"):
if request.params is not None and request.params.cursor is not None:
raise RuntimeError("Upstream catalog is unavailable")
match request:
case ListPromptsRequest():
return ListPromptsResult(prompts=[])
case ListResourcesRequest():
return ListResourcesResult(resources=[])
case ListResourceTemplatesRequest():
return ListResourceTemplatesResult(resource_templates=[])
case _:
raise RuntimeError("Unexpected catalog request type")
headers: Final = (
dict(
chain(
extra_headers.items() if extra_headers else (),
server.static_headers.items() if server.static_headers else (),
)
)
or None
)
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=headers,
stdio_env=stdio_env,
subject_token=subject_token,
user_api_key_auth=user_api_key_auth,
raw_headers=raw_headers,
client_ip=client_ip,
)
page: Final = await client.list_page(request)
match page:
case ListPromptsResult():
return page.model_copy(update={"prompts": self._create_prefixed_prompts(page.prompts, server)})
case ListResourcesResult():
return page.model_copy(update={"resources": self._create_prefixed_resources(page.resources, server)})
case ListResourceTemplatesResult():
return page.model_copy(
update={
"resource_templates": self._create_prefixed_resource_templates(page.resource_templates, server)
}
)
case _:
raise RuntimeError("Unexpected catalog result type")
async def get_prompts_from_server(
self,
server: MCPServer,
@ -5487,6 +5587,9 @@ class MCPServerManager:
tools: Sequence[MCPTool],
server: MCPServer,
add_prefix: bool = True,
*,
register_bare_names: bool = True,
listing_updates: ExitStack | None = None,
) -> list[MCPTool]:
"""
Create prefixed tools and update tool mapping.
@ -5500,6 +5603,10 @@ class MCPServerManager:
"""
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
if register_bare_names and listing_updates is not None:
listing_updates.callback(self._create_prefixed_tools, tools, server, add_prefix=add_prefix)
register_bare_names = False
prefixed_tools: Final = []
prefix: Final = get_server_prefix(server)
@ -5517,7 +5624,8 @@ class MCPServerManager:
continue
if namespace_owner is None and global_mcp_tool_registry.get_tool(spelling) is not None:
continue
self.tool_name_to_mcp_server_name_mapping[spelling] = prefix
if register_bare_names or spelling != original_name:
self.tool_name_to_mcp_server_name_mapping[spelling] = prefix
verbose_logger.info("Successfully fetched %s tools from server %s", len(prefixed_tools), server.name)
return prefixed_tools

View file

@ -6,6 +6,7 @@ import types
import uuid
from collections.abc import Mapping, Sequence
from datetime import datetime
from functools import partial
from typing import Any, Final, NoReturn, TypeAlias, overload
from fastapi import HTTPException
@ -74,15 +75,12 @@ from litellm.proxy._experimental.mcp_server.exceptions import (
MCPUpstreamAuthError,
)
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
SERVER_OUTCOMES_META_KEY,
AggregateToolListing,
ServerListOk,
ServerOutcome,
classify_list_exception,
outcome_wire_value,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
ListedToolsCaller,
MCPServerManager,
_caller_authorization_fans_out,
_client_forwarded_authorization_headers,
@ -959,6 +957,8 @@ async def _get_tools_from_mcp_servers(
request_tags: list[str] | None = None,
client_ip: str | None = None,
mcp_proxy_mode: bool = False,
params: PaginatedRequestParams | None = None,
protocol_version: str | None = None,
*,
record_listing: bool = False,
) -> AggregateToolListing:
@ -1068,149 +1068,54 @@ async def _get_tools_from_mcp_servers(
await _prefetch_oauth_creds_for_user(user_api_key_auth) if _has_oauth2_server else {}
)
async def _fetch_and_filter_server_tools(
server: MCPServer,
) -> "tuple[list[MCPTool], ServerOutcome]":
"""Fetch and filter tools from a single server, classifying any failure into that
server's outcome so the aggregate can keep serving the healthy subset without a
broken server masquerading as an empty one."""
if server is None:
return [], ServerListOk(tool_count=0)
from litellm.proxy._experimental.mcp_server.catalog import get_filtered_server_tools
server_auth_header, extra_headers = _prepare_mcp_server_headers(
server=server,
mcp_server_auth_headers=mcp_server_auth_headers,
mcp_auth_header=mcp_auth_header,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
scope_servers=allowed_mcp_servers,
context: Final = OperationContext(
_caller=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=tuple(mcp_servers) if mcp_servers is not None else None,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
client_ip=client_ip,
mcp_proxy_mode=mcp_proxy_mode,
protocol_version=protocol_version,
)
async def _fetch_and_filter_server_tools(server: MCPServer) -> tuple[Sequence[MCPTool], ServerOutcome]:
page, outcome = await get_filtered_server_tools(
server,
context=context,
allowed_mcp_servers=allowed_mcp_servers,
prefetched_oauth_creds=_prefetched_oauth_creds,
record_listing=record_listing,
)
return page.tools, outcome
# Prefer server-stored per-user OAuth when configured, so a stale
# Authorization header from the MCP client cannot override Redis/DB
# (same issue as call_tool in mcp_server_manager: VS Code caches tokens).
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415
to_server_spec,
if params is None:
results: Final = await asyncio.gather(
*(_fetch_and_filter_server_tools(server) for server in allowed_mcp_servers)
)
aggregated = AggregateToolListing(
tools=[tool for tools, _ in results for tool in tools],
outcomes={
_aggregate_server_key(server): outcome for server, (_, outcome) in zip(allowed_mcp_servers, results)
},
)
else:
from litellm.proxy._experimental.mcp_server.catalog import aggregate_gateway_tools
# A server migrated to the v2 resolver gets its token from the resolver at connect
# time; building it here would double-resolve and be shadowed by the v2 graft. The
# preemptive 401 already challenged a missing token, so one exists for the connect.
migrated_to_v2: Final = to_server_spec(server) is not None
if (
not migrated_to_v2
and server.auth_type == MCPAuth.oauth2
and getattr(server, "needs_user_oauth_token", False)
and user_api_key_auth is not None
):
db_headers: Final = await _get_user_oauth_extra_headers_from_db(
server,
user_api_key_auth,
prefetched_creds=_prefetched_oauth_creds,
)
if db_headers:
extra_headers = db_headers
# If still no OAuth2 token, fall back to pre-fetched creds (non-stale-client path)
elif not migrated_to_v2 and extra_headers is None and server.auth_type == MCPAuth.oauth2:
extra_headers = await _get_user_oauth_extra_headers_from_db(
server,
user_api_key_auth,
prefetched_creds=_prefetched_oauth_creds,
)
catalog_auth_header: Final = server_auth_header
if server.is_byok and server.auth_type != MCPAuth.oauth2 and server_auth_header is None:
server_auth_header = await _get_byok_credential(server, user_api_key_auth)
try:
from litellm.proxy.proxy_server import proxy_logging_obj
listed_generation: Final = global_mcp_server_manager.listed_tools_generation(server.server_id)
tools: Final = list(
await global_mcp_server_manager._get_tools_from_server(
server=server,
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,
user_api_key_auth=user_api_key_auth,
oauth2_headers=oauth2_headers,
proxy_logging_obj=proxy_logging_obj,
catalog_auth_header=catalog_auth_header,
record_listing=False,
)
)
filtered_tools = filter_tools_by_allowed_tools(tools, server)
filtered_tools = await filter_tools_by_key_team_permissions(
tools=filtered_tools,
server_id=server.server_id,
user_api_key_auth=user_api_key_auth,
)
global_mcp_server_manager.record_listed_tools(
server,
[
tool.model_copy(update={"name": strip_known_server_prefix(tool.name, server)})
for tool in filtered_tools
],
ListedToolsCaller(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=catalog_auth_header,
raw_headers=raw_headers,
oauth2_headers=oauth2_headers,
),
listed_generation,
record_listing=record_listing,
)
if mcp_proxy_mode:
from litellm.proxy._experimental.mcp_server.tool_search import with_mcp_proxy_identity
filtered_tools = [with_mcp_proxy_identity(tool, server.server_id) for tool in filtered_tools]
else:
filtered_tools = apply_display_name_overrides(filtered_tools, server)
verbose_logger.debug(
"Successfully fetched %s tools from server %s, %s after filtering",
len(tools),
server.name,
len(filtered_tools),
)
return filtered_tools, ServerListOk(tool_count=len(filtered_tools))
except MCPUpstreamAuthError as e:
# Absorb so one unauthenticated server does not empty every other server's
# tools. Surfacing the upstream 401 to the client as a re-auth challenge is
# intentionally not done here: raising from this list handler cannot produce a
# 401 + WWW-Authenticate (the MCP session manager serializes it as a JSON-RPC
# error). Single-server routes surface it via the request-scope preemptive
# check in _raise_preemptive_401_for_unauthenticated_servers instead.
verbose_logger.debug("MCP list_tools: omitting %s; it needs upstream auth", server.name)
return [], classify_list_exception(e)
except Exception as e:
verbose_logger.exception("Error getting tools from server %s: %s", server.name, e)
return [], classify_list_exception(e)
# Fetch tools from all servers in parallel
tasks: Final = [_fetch_and_filter_server_tools(server) for server in allowed_mcp_servers]
results: Final = await asyncio.gather(*tasks)
# Flatten results into single list
all_tools: Final[list[MCPTool]] = [tool for tools, _ in results for tool in tools]
server_outcomes: Final[dict[str, ServerOutcome]] = {
_aggregate_server_key(server): outcome
for server, (_, outcome) in zip(allowed_mcp_servers, results)
if server is not None
}
aggregated = await aggregate_gateway_tools(
context, params, allowed_mcp_servers, _prefetched_oauth_creds, record_listing=record_listing
)
all_tools: Final = aggregated.tools
server_outcomes: Final = aggregated.outcomes
# If logging is enabled, enrich spend_logs_metadata with counts
if litellm_logging_obj:
per_server_tool_counts: Final[dict[str, int]] = {
_aggregate_server_key(server): len(server_tools)
for server, (server_tools, _) in zip(allowed_mcp_servers, results)
if server is not None
per_server_tool_counts: Final = {
key: outcome.tool_count if isinstance(outcome, ServerListOk) else 0
for key, outcome in server_outcomes.items()
}
metadata_dict: Final = litellm_logging_obj.model_call_details.get("metadata")
@ -1243,7 +1148,7 @@ async def _get_tools_from_mcp_servers(
verbose_logger.info("Successfully fetched %s tools total from all MCP servers", len(all_tools))
return AggregateToolListing(tools=all_tools, outcomes=server_outcomes)
return aggregated
except Exception as e:
# Only fire failure hook if logging was requested for this list-tools execution
if log_list_tools_to_spendlogs and user_api_key_auth is not None:
@ -1489,6 +1394,8 @@ async def _list_mcp_tools(
list_tools_log_source: str | None = None,
client_ip: str | None = None,
mcp_proxy_mode: bool = False,
params: PaginatedRequestParams | None = None,
protocol_version: str | None = None,
*,
record_listing: bool = False,
) -> AggregateToolListing:
@ -1509,6 +1416,8 @@ async def _list_mcp_tools(
classified listing outcome
"""
from mcp.shared.exceptions import MCPError
try:
listing: Final = await _get_tools_from_mcp_servers(
user_api_key_auth=user_api_key_auth,
@ -1522,12 +1431,16 @@ async def _list_mcp_tools(
client_ip=client_ip,
mcp_proxy_mode=mcp_proxy_mode,
record_listing=record_listing,
params=params,
protocol_version=protocol_version,
)
verbose_logger.debug("Successfully fetched %s tools from managed MCP servers", len(listing.tools))
return listing
except HTTPException:
except (HTTPException, MCPError):
raise
except Exception as e:
if params is not None and params.cursor is not None:
raise
verbose_logger.exception("Error getting tools from managed MCP servers: %s", e)
# Continue with an empty listing instead of failing completely
return AggregateToolListing(tools=[], outcomes={})
@ -2759,16 +2672,13 @@ async def _execute_handle_list_tools(
*,
log_list_tools_to_spendlogs: bool = True,
) -> ListToolsResult:
from mcp.shared.exceptions import MCPError
from mcp.types import INVALID_PARAMS
try:
(
user_api_key_auth,
mcp_auth_header,
mcp_servers,
mcp_server_auth_headers,
oauth2_headers,
raw_headers,
_client_ip,
) = context.legacy_auth()
user_api_key_auth: Final = context.user_api_key_auth
mcp_servers: Final = context.mcp_servers
mcp_server_auth_headers: Final = context.mcp_server_auth_headers
verbose_logger.debug("MCP list_tools - User API Key Auth from context: %s", user_api_key_auth)
verbose_logger.debug("MCP list_tools - MCP servers from context: %s", mcp_servers)
verbose_logger.debug(
@ -2783,41 +2693,39 @@ async def _execute_handle_list_tools(
)
if context.mcp_proxy_mode:
if params.cursor is not None:
raise MCPError(code=INVALID_PARAMS, message="Invalid pagination cursor; start a fresh listing")
return ListToolsResult(tools=[Tool.model_validate(d) for d in get_mcp_proxy_tool_definitions()])
if getattr(
getattr(user_api_key_auth, "object_permission", None),
"mcp_tool_search_enabled",
False,
):
if params.cursor is not None:
raise MCPError(code=INVALID_PARAMS, message="Invalid pagination cursor; start a fresh listing")
return ListToolsResult(tools=[Tool.model_validate(d) for d in get_virtual_tool_definitions()])
# Get mcp_servers from context variable
verbose_logger.debug("MCP list_tools - Calling _list_mcp_tools")
listing: Final = await _list_mcp_tools(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=mcp_servers,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
log_list_tools_to_spendlogs=log_list_tools_to_spendlogs,
list_tools_log_source="mcp_protocol",
client_ip=_client_ip,
record_listing=True,
from litellm.proxy._experimental.mcp_server.catalog import list_gateway_catalog
from litellm.proxy._experimental.mcp_server.server_resolution import MCPServerTargetCatalog
catalog: Final = MCPServerTargetCatalog(
global_mcp_server_manager,
listing=partial(list_gateway_catalog, log_list_tools_to_spendlogs=log_list_tools_to_spendlogs),
)
verbose_logger.info("MCP list_tools - Successfully returned %s tools", len(listing.tools))
if not listing.outcomes:
return ListToolsResult(tools=listing.tools)
outcome_meta: Final = {
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})
result: Final = await catalog.list(context, ListToolsRequest(params=params))
assert isinstance(result, ListToolsResult)
return result
except MCPError:
raise
except HTTPException as e:
from mcp.shared.exceptions import MCPError
from mcp.types import INVALID_REQUEST
raise MCPError(code=INVALID_REQUEST, message=_http_detail_message(e.detail)) from e
except Exception as e:
if params.cursor is not None:
raise MCPError(
code=INVALID_PARAMS, message="Catalog continuation unavailable; start a fresh listing"
) from e
verbose_logger.exception("Error in list_tools endpoint: %s", e)
# Return empty list instead of failing completely
# This prevents the HTTP stream from failing and allows the client to get a response
@ -2975,41 +2883,28 @@ async def _execute_mcp_server_tool_call(
async def _execute_list_prompts(
context: OperationContext, params: PaginatedRequestParams, host_progress_callback: ProgressCallback | None = None
) -> ListPromptsResult:
from mcp.shared.exceptions import MCPError
from mcp.types import INVALID_PARAMS, INVALID_REQUEST
from litellm.proxy._experimental.mcp_server.catalog import list_gateway_catalog
from litellm.proxy._experimental.mcp_server.server_resolution import MCPServerTargetCatalog
if context.mcp_proxy_mode:
_reject_mcp_proxy_operation()
try:
(
user_api_key_auth,
mcp_auth_header,
mcp_servers,
mcp_server_auth_headers,
oauth2_headers,
raw_headers,
_client_ip,
) = context.legacy_auth()
verbose_logger.debug("MCP list_prompts - User API Key Auth from context: %s", user_api_key_auth)
verbose_logger.debug("MCP list_prompts - MCP servers from context: %s", mcp_servers)
verbose_logger.debug(
"MCP list_prompts - MCP server auth headers: %s",
list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None,
)
# Get mcp_servers from context variable
verbose_logger.debug("MCP list_prompts - Calling _list_prompts")
prompts: Final = await _list_mcp_prompts(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=mcp_servers,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
client_ip=_client_ip,
)
verbose_logger.info("MCP list_prompts - Successfully returned %s prompts", len(prompts))
return ListPromptsResult(prompts=prompts)
except Exception as e:
verbose_logger.exception("Error in list_prompts endpoint: %s", e)
# Return empty list instead of failing completely
# This prevents the HTTP stream from failing and allows the client to get a response
catalog: Final = MCPServerTargetCatalog(global_mcp_server_manager, listing=list_gateway_catalog)
result: Final = await catalog.list(context, ListPromptsRequest(params=params))
assert isinstance(result, ListPromptsResult)
return result
except MCPError:
raise
except HTTPException as error:
raise MCPError(code=INVALID_REQUEST, message=_http_detail_message(error.detail)) from error
except Exception as error:
if params.cursor is not None:
raise MCPError(
code=INVALID_PARAMS, message="Catalog continuation unavailable; start a fresh listing"
) from error
return ListPromptsResult(prompts=[])
@ -3045,78 +2940,56 @@ async def _execute_get_prompt(
async def _execute_list_resources(
context: OperationContext, params: PaginatedRequestParams, host_progress_callback: ProgressCallback | None = None
) -> ListResourcesResult:
from mcp.shared.exceptions import MCPError
from mcp.types import INVALID_PARAMS, INVALID_REQUEST
from litellm.proxy._experimental.mcp_server.catalog import list_gateway_catalog
from litellm.proxy._experimental.mcp_server.server_resolution import MCPServerTargetCatalog
if context.mcp_proxy_mode:
_reject_mcp_proxy_operation()
try:
(
user_api_key_auth,
mcp_auth_header,
mcp_servers,
mcp_server_auth_headers,
oauth2_headers,
raw_headers,
_client_ip,
) = context.legacy_auth()
verbose_logger.debug("MCP list_resources - User API Key Auth from context: %s", user_api_key_auth)
verbose_logger.debug("MCP list_resources - MCP servers from context: %s", mcp_servers)
verbose_logger.debug(
"MCP list_resources - MCP server auth headers: %s",
list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None,
)
resources: Final = await _list_mcp_resources(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=mcp_servers,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
client_ip=_client_ip,
)
verbose_logger.info("MCP list_resources - Successfully returned %s resources", len(resources))
return ListResourcesResult(resources=resources)
except Exception as e:
verbose_logger.exception("Error in list_resources endpoint: %s", e)
catalog: Final = MCPServerTargetCatalog(global_mcp_server_manager, listing=list_gateway_catalog)
result: Final = await catalog.list(context, ListResourcesRequest(params=params))
assert isinstance(result, ListResourcesResult)
return result
except MCPError:
raise
except HTTPException as error:
raise MCPError(code=INVALID_REQUEST, message=_http_detail_message(error.detail)) from error
except Exception as error:
if params.cursor is not None:
raise MCPError(
code=INVALID_PARAMS, message="Catalog continuation unavailable; start a fresh listing"
) from error
return ListResourcesResult(resources=[])
async def _execute_list_resource_templates(
context: OperationContext, params: PaginatedRequestParams, host_progress_callback: ProgressCallback | None = None
) -> ListResourceTemplatesResult:
from mcp.shared.exceptions import MCPError
from mcp.types import INVALID_PARAMS, INVALID_REQUEST
from litellm.proxy._experimental.mcp_server.catalog import list_gateway_catalog
from litellm.proxy._experimental.mcp_server.server_resolution import MCPServerTargetCatalog
if context.mcp_proxy_mode:
_reject_mcp_proxy_operation()
try:
(
user_api_key_auth,
mcp_auth_header,
mcp_servers,
mcp_server_auth_headers,
oauth2_headers,
raw_headers,
_client_ip,
) = context.legacy_auth()
verbose_logger.debug("MCP list_resource_templates - User API Key Auth from context: %s", user_api_key_auth)
verbose_logger.debug("MCP list_resource_templates - MCP servers from context: %s", mcp_servers)
verbose_logger.debug(
"MCP list_resource_templates - MCP server auth headers: %s",
list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None,
)
resource_templates: Final = await _list_mcp_resource_templates(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=mcp_servers,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
client_ip=_client_ip,
)
verbose_logger.info(
"MCP list_resource_templates - Successfully returned %s resource templates", len(resource_templates)
)
return ListResourceTemplatesResult(resource_templates=resource_templates)
except Exception as e:
verbose_logger.exception("Error in list_resource_templates endpoint: %s", e)
catalog: Final = MCPServerTargetCatalog(global_mcp_server_manager, listing=list_gateway_catalog)
result: Final = await catalog.list(context, ListResourceTemplatesRequest(params=params))
assert isinstance(result, ListResourceTemplatesResult)
return result
except MCPError:
raise
except HTTPException as error:
raise MCPError(code=INVALID_REQUEST, message=_http_detail_message(error.detail)) from error
except Exception as error:
if params.cursor is not None:
raise MCPError(
code=INVALID_PARAMS, message="Catalog continuation unavailable; start a fresh listing"
) from error
return ListResourceTemplatesResult(resource_templates=[])

View file

@ -913,6 +913,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=[])
@ -933,6 +935,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=[])
@ -947,6 +951,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=[])

View file

@ -2,7 +2,7 @@ from __future__ import annotations
from collections.abc import Awaitable, Callable, Mapping
from dataclasses import dataclass
from typing import Final, Literal, Protocol
from typing import TYPE_CHECKING, Final, Literal, Protocol
from fastapi import HTTPException, status
@ -10,6 +10,9 @@ from litellm.proxy._experimental.mcp_server.ui_session_utils import can_access_m
from litellm.proxy._types import LiteLLM_MCPServerTable, UserAPIKeyAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
if TYPE_CHECKING:
from litellm.proxy._experimental.mcp_server.contracts import CatalogListRequest, CatalogListResult, OperationContext
class MCPServerRegistry(Protocol):
def get_mcp_server_by_id(self, server_id: str) -> MCPServer | None: ...
@ -130,6 +133,12 @@ class MCPServerTargetCatalog:
id_client_ip: str | None = None
name_client_ip: str | None = None
match_name: bool = False
listing: Callable[[OperationContext, CatalogListRequest], Awaitable[CatalogListResult]] | None = None
async def list(self, context: OperationContext, request: CatalogListRequest) -> CatalogListResult:
if self.listing is None:
raise RuntimeError("Catalog listing dependency is not configured")
return await self.listing(context, request)
async def resolve(
self,

View file

@ -0,0 +1,87 @@
from __future__ import annotations
import base64
import binascii
import os
from enum import Enum
from typing import Final
from cryptography.exceptions import InvalidTag
from cryptography.hazmat.primitives import hashes
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
from cryptography.hazmat.primitives.kdf.hkdf import HKDF
from pydantic import BaseModel, ConfigDict, JsonValue, ValidationError
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error, Ok, Result
_PREFIX: Final = "mcp_state_v1."
_MAX_TOKEN_LENGTH: Final = 65536
_NONCE_BYTES: Final = 12
class StateTokenError(str, Enum):
MISSING_KEY = "Set the same LITELLM_SALT_KEY on every replica to enable pagination"
INVALID = "Invalid pagination state; start a fresh listing"
EXPIRED = "Pagination state expired; start a fresh listing"
TOO_LARGE = "Pagination state exceeds the supported size"
class _Envelope(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid", strict=True)
expires_at: int
value: JsonValue
def _cipher(purpose: str) -> Result[AESGCM, StateTokenError]:
salt_key: Final = os.getenv("LITELLM_SALT_KEY")
if not salt_key:
return Error(StateTokenError.MISSING_KEY)
if not purpose:
return Error(StateTokenError.INVALID)
key: Final = HKDF(
algorithm=hashes.SHA256(),
length=32,
salt=b"litellm:mcp:state:v1",
info=purpose.encode("utf-8"),
).derive(salt_key.encode("utf-8"))
return Ok(AESGCM(key))
def seal_state(value: JsonValue, *, purpose: str, expires_at: int, now: int) -> Result[str, StateTokenError]:
cipher: Final = _cipher(purpose)
if isinstance(cipher, Error):
return cipher
if expires_at <= now:
return Error(StateTokenError.EXPIRED)
plaintext: Final = _Envelope(expires_at=expires_at, value=value).model_dump_json().encode("utf-8")
if len(plaintext) > _MAX_TOKEN_LENGTH:
return Error(StateTokenError.TOO_LARGE)
nonce: Final = os.urandom(_NONCE_BYTES)
ciphertext: Final = cipher.ok.encrypt(nonce, plaintext, (_PREFIX + purpose).encode("utf-8"))
token: Final = _PREFIX + base64.urlsafe_b64encode(nonce + ciphertext).decode("ascii").rstrip("=")
return Error(StateTokenError.TOO_LARGE) if len(token) > _MAX_TOKEN_LENGTH else Ok(token)
def open_state(token: str, *, purpose: str, now: int) -> Result[JsonValue, StateTokenError]:
cipher: Final = _cipher(purpose)
if isinstance(cipher, Error):
return cipher
if len(token) > _MAX_TOKEN_LENGTH:
return Error(StateTokenError.TOO_LARGE)
if not token.startswith(_PREFIX):
return Error(StateTokenError.INVALID)
encoded: Final = token[len(_PREFIX) :]
try:
sealed: Final = base64.b64decode(encoded + "=" * (-len(encoded) % 4), altchars=b"-_", validate=True)
if len(sealed) < _NONCE_BYTES + 16 or base64.urlsafe_b64encode(sealed).decode("ascii").rstrip("=") != encoded:
return Error(StateTokenError.INVALID)
plaintext: Final = cipher.ok.decrypt(
sealed[:_NONCE_BYTES], sealed[_NONCE_BYTES:], (_PREFIX + purpose).encode("utf-8")
)
envelope: Final = _Envelope.model_validate_json(plaintext)
except (binascii.Error, ValueError, InvalidTag, ValidationError):
return Error(StateTokenError.INVALID)
if envelope.expires_at <= now:
return Error(StateTokenError.EXPIRED)
return Ok(envelope.value)

View file

@ -22,7 +22,7 @@ from mcp.server.mcpserver import Context, MCPServer
from mcp.server.transport_security import TransportSecuritySettings
from mcp.types import SamplingMessage, TextContent
from mcp_tests.mcp_e2e_upstream_server import add, multiply
from pydantic import BaseModel
from pydantic import BaseModel, JsonValue
from sse_starlette.sse import AppStatus
from starlette.requests import Request as StarletteRequest
from starlette.responses import Response
@ -622,3 +622,114 @@ def tool_calls(observed: tuple[dict[str, object], ...]) -> tuple[dict[str, objec
return tuple(
item for item in observed if isinstance(item.get("body"), dict) and item["body"].get("method") == "tools/call"
)
@contextmanager
def paginated_mcp_peer(
*,
page_size: int = 1,
repeat_cursor: bool = False,
fail_listing: bool = False,
fail_continuation: bool = False,
metadata: dict[str, JsonValue] | None = None,
) -> Iterator[McpPeer]:
from contextlib import asynccontextmanager
from mcp.server.lowlevel.server import Server
from mcp.server.streamable_http_manager import StreamableHTTPSessionManager
from mcp.types import (
CallToolResult,
ListPromptsResult,
ListResourcesResult,
ListResourceTemplatesResult,
ListToolsResult,
Prompt,
Resource,
ResourceTemplate,
Tool,
)
from starlette.applications import Starlette
from starlette.routing import Mount
def window(params):
if fail_listing or (fail_continuation and params is not None and params.cursor):
from mcp import MCPError
from mcp.types import INTERNAL_ERROR
raise MCPError(code=INTERNAL_ERROR, message="untrusted upstream message")
start = int(params.cursor) if params is not None and params.cursor else 0
end = min(start + page_size, 3)
return range(start, end), "1" if repeat_cursor else str(end) if end < 3 else None
async def tools(context, params):
indexes, cursor = window(params)
return ListToolsResult(
tools=[
Tool(
name=f"add{index}",
input_schema={
"type": "object",
"properties": {"a": {"type": "integer"}, "b": {"type": "integer"}},
"required": ["a", "b"],
},
)
for index in indexes
],
next_cursor=cursor,
meta={"revision": "stable", **(metadata or {})},
)
async def prompts(context, params):
indexes, cursor = window(params)
return ListPromptsResult(
prompts=[Prompt(name=f"prompt{index}") for index in indexes], next_cursor=cursor, meta=metadata
)
async def resources(context, params):
indexes, cursor = window(params)
return ListResourcesResult(
resources=[Resource(name=f"resource{index}", uri=f"status://item{index}") for index in indexes],
next_cursor=cursor,
meta=metadata,
)
async def templates(context, params):
indexes, cursor = window(params)
return ListResourceTemplatesResult(
resource_templates=[
ResourceTemplate(name=f"template{index}", uri_template=f"status{index}://{{item}}") for index in indexes
],
next_cursor=cursor,
meta=metadata,
)
async def call(context, params):
assert params.name in ("add0", "add1", "add2")
return CallToolResult(
content=[TextContent(type="text", text=str(params.arguments["a"] + params.arguments["b"]))]
)
service = Server(
"paginated-catalog",
on_list_tools=tools,
on_list_prompts=prompts,
on_list_resources=resources,
on_list_resource_templates=templates,
on_call_tool=call,
)
manager = StreamableHTTPSessionManager(
service,
stateless=True,
json_response=True,
security_settings=TransportSecuritySettings(enable_dns_rebinding_protection=False),
)
@asynccontextmanager
async def lifespan(app):
async with manager.run():
yield
app = Starlette(routes=[Mount("/mcp", app=manager.handle_request)], lifespan=lifespan)
observed = queue.Queue()
with asgi_server(_capturing(app, observed)) as url:
yield McpPeer(url + "/mcp/", observed)

View file

@ -0,0 +1,563 @@
import asyncio
from contextlib import asynccontextmanager
from pathlib import Path
from typing import Literal
import httpx
import pytest
import yaml
from mcp import ClientSession, MCPError
from mcp.client.streamable_http import streamable_http_client
from mcp.types import CallToolRequest, CallToolRequestParams, CallToolResult, PaginatedRequestParams
from integration._support.client import Gateway
from integration._support.mcp import paginated_mcp_peer
from integration._support.process import owned_proxy
from litellm.experimental_mcp_client.client import MCPClient
from litellm.types.mcp import MCPTransport
@asynccontextmanager
async def catalog_session(gateway):
async with httpx.AsyncClient(headers={"Authorization": "Bearer " + gateway.key}) as client:
async with streamable_http_client(
str(gateway.client.base_url).rstrip("/") + "/mcp/", http_client=client
) as streams:
async with ClientSession(streams[0], streams[1]) as session:
await session.initialize()
yield session
def config_file(directory: Path, upstream: str) -> Path:
config = directory / "proxy.yaml"
config.write_text(
yaml.safe_dump(
{
"model_list": [],
"mcp_servers": {"pages": {"url": upstream, "transport": "http"}},
"general_settings": {"master_key": "sk-pagination-test", "store_model_in_db": False},
}
)
)
return config
REMOVE_DATABASE = ("DATABASE_URL", "DATABASE_URL_READ_REPLICA", "LITELLM_LICENSE", "LITELLM_LICENSE_PATH")
def test_catalog_pages_are_portable_repeatable_and_complete(tmp_path: Path):
async def exercise(first_replica, second_replica):
for method, field in (
("list_tools", "tools"),
("list_prompts", "prompts"),
("list_resources", "resources"),
("list_resource_templates", "resource_templates"),
):
async with catalog_session(first_replica) as first_session:
first = await getattr(first_session, method)()
assert first.next_cursor
async with catalog_session(second_replica) as second_session:
second = await getattr(second_session, method)(params=PaginatedRequestParams(cursor=first.next_cursor))
replay = await getattr(second_session, method)(params=PaginatedRequestParams(cursor=first.next_cursor))
assert getattr(second, field) == getattr(replay, field)
assert second.next_cursor
third = await getattr(second_session, method)(params=PaginatedRequestParams(cursor=second.next_cursor))
names = [item.name for page in (first, second, third) for item in getattr(page, field)]
assert len(names) == len(set(names)) == 3
assert third.next_cursor is None
async with catalog_session(second_replica) as session:
called = await session.call_tool("pages-add2", {"a": 3, "b": 4})
assert not called.is_error
assert called.content[0].text == "7"
with paginated_mcp_peer() as peer, httpx.Client() as client:
seed = Gateway(client, "sk-pagination-test", peer.url)
config = config_file(tmp_path, peer.url)
options = dict(config=config, database_setup=(), remove_environment=REMOVE_DATABASE)
environment = {
"STORE_MODEL_IN_DB": "False",
"DISABLE_SCHEMA_UPDATE": "true",
"LITELLM_SALT_KEY": "shared-pagination-test",
}
with (
owned_proxy(seed, tmp_path / "a", environment, **options) as a,
owned_proxy(seed, tmp_path / "b", environment, **options) as b,
):
asyncio.run(exercise(a, b))
@pytest.mark.parametrize("page_size", [1, 3])
def test_no_salt_preserves_complete_lists_and_calls_but_rejects_continuations(tmp_path: Path, page_size: int):
async def exercise(gateway):
async with catalog_session(gateway) as session:
for method, field in (
("list_tools", "tools"),
("list_prompts", "prompts"),
("list_resources", "resources"),
("list_resource_templates", "resource_templates"),
):
if page_size == 1:
with pytest.raises(MCPError, match="LITELLM_SALT_KEY"):
await getattr(session, method)()
else:
result = await getattr(session, method)()
assert len(getattr(result, field)) == 3
assert result.next_cursor is None
called = await session.send_request(
CallToolRequest(params=CallToolRequestParams(name="pages-add2", arguments={"a": 3, "b": 4})),
CallToolResult,
)
assert not called.is_error
assert called.content[0].text == "7"
if page_size == 3:
bare = await session.send_request(
CallToolRequest(params=CallToolRequestParams(name="add2", arguments={"a": 3, "b": 4})),
CallToolResult,
)
assert not bare.is_error
assert bare.content[0].text == "7"
with paginated_mcp_peer(page_size=page_size) as peer, httpx.Client() as client:
seed = Gateway(client, "sk-pagination-test", peer.url)
config = config_file(tmp_path, peer.url)
environment = {"STORE_MODEL_IN_DB": "False", "DISABLE_SCHEMA_UPDATE": "true", "LITELLM_SALT_KEY": ""}
with owned_proxy(
seed, tmp_path / "proxy", environment, config=config, database_setup=(), remove_environment=REMOVE_DATABASE
) as gateway:
asyncio.run(exercise(gateway))
@pytest.mark.parametrize("grant", ["direct", "access_group"])
def test_continuations_reauthorize_and_reject_registry_changes(tmp_path: Path, monkeypatch, grant: str):
import os
import subprocess
import sys
from integration._support.database import scratch_database
from integration._support.mcp import register_mcp
assert os.environ.get("DATABASE_URL"), "This integration case requires disposable-database access"
async def exercise(a, b, peer, identity, owner, stranger, policy):
owner_a = Gateway(a.client, owner, peer.url)
owner_b = Gateway(b.client, owner, peer.url)
stranger_b = Gateway(b.client, stranger, peer.url)
first_pages = {}
for method in ("list_tools", "list_prompts", "list_resources", "list_resource_templates"):
async with catalog_session(owner_a) as session:
first_pages[method] = await getattr(session, method)()
assert first_pages[method].next_cursor
async with catalog_session(owner_b) as session:
continued = await getattr(session, method)(
params=PaginatedRequestParams(cursor=first_pages[method].next_cursor)
)
assert continued.next_cursor
async with catalog_session(stranger_b) as session:
peer.drain()
with pytest.raises(MCPError, match="fresh listing"):
await getattr(session, method)(
params=PaginatedRequestParams(cursor=first_pages[method].next_cursor)
)
assert not any(call["body"].get("method", "").endswith("/list") for call in peer.drain())
a.post(
"/key/update",
{
"key": owner,
**(
{"access_group_ids": []}
if grant == "access_group"
else {"object_permission": {"mcp_servers": ["no-mcp-servers"]}}
),
},
)
for method, first in first_pages.items():
async with catalog_session(owner_b) as session:
peer.drain()
with pytest.raises(MCPError, match="fresh listing"):
await getattr(session, method)(params=PaginatedRequestParams(cursor=first.next_cursor))
assert not any(call["body"].get("method", "").endswith("/list") for call in peer.drain())
a.post("/key/update", {"key": owner, **policy})
changed = a.request("PUT", "/v1/mcp/server", {"server_id": identity, "description": "new catalog generation"})
assert changed.status_code == 202, changed.text
for method, first in first_pages.items():
async with catalog_session(owner_b) as session:
peer.drain()
with pytest.raises(MCPError, match="fresh listing"):
await getattr(session, method)(params=PaginatedRequestParams(cursor=first.next_cursor))
assert not any(call["body"].get("method", "").endswith("/list") for call in peer.drain())
fresh = await getattr(session, method)()
assert fresh.next_cursor
with scratch_database() as database_url:
monkeypatch.setenv("DATABASE_URL", database_url)
subprocess.run(
[
sys.executable,
"-I",
"-m",
"prisma",
"db",
"push",
"--schema",
"litellm/proxy/schema.prisma",
"--skip-generate",
],
check=True,
capture_output=True,
text=True,
)
with paginated_mcp_peer() as peer, httpx.Client() as client:
seed = Gateway(client, "sk-pagination-test", peer.url)
config = tmp_path / "database-proxy.yaml"
config.write_text(
yaml.safe_dump(
{"model_list": [], "general_settings": {"master_key": seed.key, "store_model_in_db": True}}
)
)
environment = {
"DATABASE_URL": database_url,
"DISABLE_SCHEMA_UPDATE": "true",
"LITELLM_SALT_KEY": "shared-pagination-test",
}
options = dict(
config=config,
database_setup=(),
remove_environment=("DATABASE_URL_READ_REPLICA", "LITELLM_LICENSE", "LITELLM_LICENSE_PATH"),
)
with (
owned_proxy(seed, tmp_path / "a", environment, **options) as a,
owned_proxy(seed, tmp_path / "b", environment, **options) as b,
a.scenario() as scenario,
):
identity = register_mcp(scenario, peer, "pages")
group = a.request(
"POST",
"/v1/access_group",
{
"access_group_name": "pagination-grant",
"access_mcp_server_ids": [identity],
},
)
assert group.status_code == 201, group.text
policy = (
{"access_group_ids": [group.json()["access_group_id"]]}
if grant == "access_group"
else {"object_permission": {"mcp_servers": [identity]}}
)
owner = scenario.key(**policy)
stranger = scenario.key(object_permission={"mcp_servers": [identity]})
assert owner != stranger
asyncio.run(exercise(a, b, peer, identity, owner, stranger, policy))
@pytest.mark.parametrize("changed", ["key", "snapshot"])
def test_cursor_rejects_changed_replica_configuration_before_dispatch(tmp_path: Path, changed: str):
async def exercise(a, b, peer):
for method in ("list_tools", "list_prompts", "list_resources", "list_resource_templates"):
async with catalog_session(a) as session:
first = await getattr(session, method)()
assert first.next_cursor
async with catalog_session(b) as session:
peer.drain()
with pytest.raises(MCPError, match="fresh listing"):
await getattr(session, method)(params=PaginatedRequestParams(cursor=first.next_cursor))
assert not any(call["body"].get("method", "").endswith("/list") for call in peer.drain())
fresh = await getattr(session, method)()
assert fresh.next_cursor
with paginated_mcp_peer() as peer, httpx.Client() as client:
seed = Gateway(client, "sk-pagination-test", peer.url)
config = config_file(tmp_path, peer.url)
second_config = tmp_path / "second-proxy.yaml"
second_values = yaml.safe_load(config.read_text())
if changed == "snapshot":
second_values["mcp_servers"]["pages"]["description"] = "changed registry definition"
second_config.write_text(yaml.safe_dump(second_values))
environment = {
"STORE_MODEL_IN_DB": "False",
"DISABLE_SCHEMA_UPDATE": "true",
"LITELLM_SALT_KEY": "original-key",
}
second_environment = {**environment, "LITELLM_SALT_KEY": "rotated-key" if changed == "key" else "original-key"}
options = dict(database_setup=(), remove_environment=REMOVE_DATABASE)
with (
owned_proxy(seed, tmp_path / "a", environment, config=config, **options) as a,
owned_proxy(seed, tmp_path / "b", second_environment, config=second_config, **options) as b,
):
asyncio.run(exercise(a, b, peer))
def test_pages_preserve_supported_protocol_versions(tmp_path: Path):
from mcp.types import ListPromptsRequest, ListResourcesRequest, ListResourceTemplatesRequest, ListToolsRequest
from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS
async def exercise(gateway):
for revision in HANDSHAKE_PROTOCOL_VERSIONS:
client = MCPClient(
server_url=str(gateway.client.base_url).rstrip("/") + "/mcp",
transport_type=MCPTransport.http,
protocol_version=revision,
extra_headers={"Authorization": "Bearer " + gateway.key},
timeout=15,
)
for request, field in (
(ListToolsRequest, "tools"),
(ListPromptsRequest, "prompts"),
(ListResourcesRequest, "resources"),
(ListResourceTemplatesRequest, "resource_templates"),
):
first = await client.list_page(request())
assert first.next_cursor, revision
second = await client.list_page(request(params=PaginatedRequestParams(cursor=first.next_cursor)))
assert second.next_cursor, revision
third = await client.list_page(request(params=PaginatedRequestParams(cursor=second.next_cursor)))
assert third.next_cursor is None, revision
names = [item.name for page in (first, second, third) for item in getattr(page, field)]
assert len(names) == len(set(names)) == 3, revision
with paginated_mcp_peer() as peer, httpx.Client() as client:
seed = Gateway(client, "sk-pagination-test", peer.url)
config = config_file(tmp_path, peer.url)
environment = {"STORE_MODEL_IN_DB": "False", "DISABLE_SCHEMA_UPDATE": "true", "LITELLM_SALT_KEY": "shared-key"}
with owned_proxy(
seed, tmp_path / "proxy", environment, config=config, database_setup=(), remove_environment=REMOVE_DATABASE
) as gateway:
asyncio.run(exercise(gateway))
def test_incomplete_discovery_never_establishes_a_bare_tool_route(tmp_path: Path):
from integration._support.mcp import tool_calls
async def exercise(gateway, peer):
async with catalog_session(gateway) as session:
first = await session.list_tools()
assert [tool.name for tool in first.tools] == ["pages-add0"]
assert first.next_cursor
with pytest.raises(MCPError, match="repeated"):
await session.list_tools(params=PaginatedRequestParams(cursor=first.next_cursor))
peer.drain()
result = await session.send_request(
CallToolRequest(params=CallToolRequestParams(name="add0", arguments={"a": 3, "b": 4})),
CallToolResult,
)
assert result.is_error
assert tool_calls(peer.drain()) == ()
prefixed = await session.send_request(
CallToolRequest(params=CallToolRequestParams(name="pages-add0", arguments={"a": 3, "b": 4})),
CallToolResult,
)
assert not prefixed.is_error
assert prefixed.content[0].text == "7"
assert len(tool_calls(peer.drain())) == 1
with paginated_mcp_peer(repeat_cursor=True) as peer, httpx.Client() as client:
seed = Gateway(client, "sk-pagination-test", peer.url)
config = config_file(tmp_path, peer.url)
environment = {"STORE_MODEL_IN_DB": "False", "DISABLE_SCHEMA_UPDATE": "true", "LITELLM_SALT_KEY": "shared-key"}
with owned_proxy(
seed, tmp_path / "proxy", environment, config=config, database_setup=(), remove_environment=REMOVE_DATABASE
) as gateway:
asyncio.run(exercise(gateway, peer))
def test_partial_catalog_keeps_sanitized_failure_metadata_on_following_pages(tmp_path: Path):
async def exercise(gateway):
for method, field in (
("list_tools", "tools"),
("list_prompts", "prompts"),
("list_resources", "resources"),
("list_resource_templates", "resource_templates"),
):
async with catalog_session(gateway) as session:
first = await getattr(session, method)()
assert first.next_cursor
assert len(getattr(first, field)) == 1
fault = first.meta["litellm.ai/server_outcomes"]["broken"]
assert fault["status"] != "ok"
assert "tag" not in fault
assert "untrusted upstream message" not in first.model_dump_json()
second = await getattr(session, method)(params=PaginatedRequestParams(cursor=first.next_cursor))
assert len(getattr(second, field)) == 1
assert second.meta["litellm.ai/server_outcomes"]["broken"] == fault
with paginated_mcp_peer() as healthy, paginated_mcp_peer(fail_listing=True) as broken, httpx.Client() as client:
seed = Gateway(client, "sk-pagination-test", healthy.url)
config = config_file(tmp_path, healthy.url)
values = yaml.safe_load(config.read_text())
values["mcp_servers"]["broken"] = {"url": broken.url, "transport": "http"}
config.write_text(yaml.safe_dump(values))
environment = {"STORE_MODEL_IN_DB": "False", "DISABLE_SCHEMA_UPDATE": "true", "LITELLM_SALT_KEY": "shared-key"}
with owned_proxy(
seed, tmp_path / "proxy", environment, config=config, database_setup=(), remove_environment=REMOVE_DATABASE
) as gateway:
asyncio.run(exercise(gateway))
def test_failed_upstream_continuation_requires_restart_for_every_catalog(tmp_path: Path):
async def exercise(gateway):
async with catalog_session(gateway) as session:
for method in ("list_tools", "list_prompts", "list_resources", "list_resource_templates"):
first = await getattr(session, method)()
assert first.next_cursor
with pytest.raises(MCPError, match="fresh listing") as caught:
await getattr(session, method)(params=PaginatedRequestParams(cursor=first.next_cursor))
assert "untrusted upstream message" not in str(caught.value)
with paginated_mcp_peer(fail_continuation=True) as peer, httpx.Client() as client:
seed = Gateway(client, "sk-pagination-test", peer.url)
config = config_file(tmp_path, peer.url)
environment = {"STORE_MODEL_IN_DB": "False", "DISABLE_SCHEMA_UPDATE": "true", "LITELLM_SALT_KEY": "shared-key"}
with owned_proxy(
seed, tmp_path / "proxy", environment, config=config, database_setup=(), remove_environment=REMOVE_DATABASE
) as gateway:
asyncio.run(exercise(gateway))
@pytest.mark.parametrize("foreign", [{"foreign": {"status": "error", "message": "foreign outcome"}}, "malformed"])
def test_optional_catalog_ignores_upstream_gateway_outcomes(tmp_path: Path, foreign):
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import SERVER_OUTCOMES_META_KEY
async def exercise(gateway):
async with catalog_session(gateway) as session:
for method, field in (
("list_prompts", "prompts"),
("list_resources", "resources"),
("list_resource_templates", "resource_templates"),
):
first = await getattr(session, method)()
assert len(getattr(first, field)) == 1
assert first.next_cursor
second = await getattr(session, method)(params=PaginatedRequestParams(cursor=first.next_cursor))
assert len(getattr(second, field)) == 1
assert "foreign" not in (second.meta or {}).get(SERVER_OUTCOMES_META_KEY, {})
with paginated_mcp_peer(metadata={SERVER_OUTCOMES_META_KEY: foreign}) as peer, httpx.Client() as client:
seed = Gateway(client, "sk-pagination-test", peer.url)
config = config_file(tmp_path, peer.url)
environment = {"STORE_MODEL_IN_DB": "False", "DISABLE_SCHEMA_UPDATE": "true", "LITELLM_SALT_KEY": "shared-key"}
with owned_proxy(
seed, tmp_path / "proxy", environment, config=config, database_setup=(), remove_environment=REMOVE_DATABASE
) as gateway:
asyncio.run(exercise(gateway))
def test_complete_initial_page_keeps_bare_routes_with_a_cached_database_revision(tmp_path: Path, monkeypatch):
import os
import subprocess
import sys
from integration._support.database import read_rows, scratch_database, write_rows
from integration._support.mcp import register_mcp, tool_calls
assert os.environ.get("DATABASE_URL"), "This integration case requires disposable-database access"
async def exercise(gateway, first, second):
async with catalog_session(gateway) as session:
listing = await session.list_tools()
assert len(listing.tools) == 3 and listing.next_cursor is None
first.drain()
second.drain()
called = await session.send_request(
CallToolRequest(params=CallToolRequestParams(name="add2", arguments={"a": 3, "b": 4})),
CallToolResult,
)
assert not called.is_error
assert called.content[0].text == "7"
assert tool_calls(first.drain()) == ()
assert len(tool_calls(second.drain())) == 1
with scratch_database() as database_url:
monkeypatch.setenv("DATABASE_URL", database_url)
subprocess.run(
[sys.executable, "-I", "-m", "prisma", "db", "push", "--schema",
"litellm/proxy/schema.prisma", "--skip-generate"],
check=True, capture_output=True, text=True,
)
subprocess.run(
[sys.executable, "-I", "-m", "prisma", "db", "execute", "--schema",
"litellm/proxy/schema.prisma", "--file",
"litellm-proxy-extras/litellm_proxy_extras/migrations/20260923000000_add_mcp_catalog_revision_trigger/migration.sql"],
check=True, capture_output=True, text=True,
)
with paginated_mcp_peer(page_size=3) as first, paginated_mcp_peer(page_size=3) as second, httpx.Client() as client:
seed = Gateway(client, "sk-pagination-test", first.url)
config = tmp_path / "database-proxy.yaml"
config.write_text(yaml.safe_dump({
"model_list": [], "general_settings": {"master_key": seed.key, "store_model_in_db": True},
}))
environment = {"DATABASE_URL": database_url, "DISABLE_SCHEMA_UPDATE": "true", "LITELLM_SALT_KEY": "shared-pagination-test"}
with owned_proxy(
seed, tmp_path / "proxy", environment, config=config, database_setup=(),
remove_environment=("DATABASE_URL_READ_REPLICA", "LITELLM_LICENSE", "LITELLM_LICENSE_PATH"),
) as gateway, gateway.scenario() as scenario:
first_id = register_mcp(scenario, first, "first")
second_id = register_mcp(scenario, second, "second")
owner = scenario.key(object_permission={"mcp_servers": [second_id]})
warm = gateway.client.get("/mcp-rest/tools/list", headers={"Authorization": "Bearer " + gateway.key}, params={"server_id": first_id})
assert warm.status_code == 200, warm.text
# A no-op SQL writer bumps the trigger revision without changing either server.
write_rows('UPDATE "LiteLLM_MCPServerTable" SET "alias" = "alias" WHERE "server_id" = %s', (first_id,))
warm = gateway.client.get("/mcp-rest/tools/list", headers={"Authorization": "Bearer " + gateway.key}, params={"server_id": first_id})
assert warm.status_code == 200, warm.text
query = 'SELECT "reload_revision" FROM "LiteLLM_Config" WHERE "param_name" = %s'
revision = read_rows(query, ("mcp_catalog",))
assert revision and revision[0]["reload_revision"] > 0
asyncio.run(exercise(Gateway(gateway.client, owner, second.url), first, second))
assert read_rows(query, ("mcp_catalog",)) == revision
@pytest.mark.parametrize("entry", ["mcp", "server_mcp"])
def test_missing_user_keeps_explicit_key_and_team_grants(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch, entry: Literal["mcp", "server_mcp"]
):
import subprocess
import sys
import uuid
from integration._support.database import read_rows, scratch_database
from integration._support.mcp import McpCaller, register_mcp, tool_calls
with scratch_database() as database_url:
monkeypatch.setenv("DATABASE_URL", database_url)
subprocess.run(
[sys.executable, "-I", "-m", "prisma", "db", "push", "--schema",
"litellm/proxy/schema.prisma", "--skip-generate"],
check=True, capture_output=True, text=True,
)
with paginated_mcp_peer(page_size=3) as allowed, paginated_mcp_peer(page_size=3) as private, httpx.Client() as client:
seed = Gateway(client, "sk-pagination-test", allowed.url)
config = tmp_path / "database-proxy.yaml"
config.write_text(yaml.safe_dump({
"model_list": [], "general_settings": {"master_key": seed.key, "store_model_in_db": True},
}))
environment = {"DATABASE_URL": database_url, "DISABLE_SCHEMA_UPDATE": "true", "LITELLM_SALT_KEY": "shared-pagination-test"}
options = dict(config=config, database_setup=(), remove_environment=("DATABASE_URL_READ_REPLICA", "LITELLM_LICENSE", "LITELLM_LICENSE_PATH"))
with (
owned_proxy(seed, tmp_path / "a", environment, **options) as a,
owned_proxy(seed, tmp_path / "b", environment, **options) as b,
a.scenario() as scenario,
):
identity = register_mcp(scenario, allowed, "allowed")
register_mcp(scenario, private, "private")
team = scenario.team(object_permission={"mcp_servers": [identity]})
for policy in ({"object_permission": {"mcp_servers": [identity]}}, {"team_id": team}):
user = "missing-" + uuid.uuid4().hex
key = scenario.key(user_id=user, **policy)
assert read_rows('SELECT user_id FROM "LiteLLM_UserTable" WHERE user_id = %s', (user,)) == []
for replica in (a, b):
caller = McpCaller(replica, key, entry, alias="allowed")
allowed.drain()
private.drain()
listing = caller.list_tools()
assert listing.ok, listing
assert len(listing.tools) == 3, listing
name = next(name for name in listing.tools if name.endswith("add2"))
called = caller.call(name, {"a": 3, "b": 4})
assert called.ok and called.text == "7", called
assert len(tool_calls(allowed.drain())) == 1
assert private.drain() == ()
forbidden = caller.call("private-add2", {"a": 3, "b": 4})
assert not forbidden.ok, forbidden
assert tool_calls(allowed.drain()) == ()
assert private.drain() == ()

View file

@ -235,11 +235,11 @@ def _clear_proxy_database_env() -> typing.Iterator[None]:
async def _initialize_proxy(config_path: str) -> None:
from litellm.proxy._experimental.mcp_server.catalog import TargetCatalog
from litellm.proxy._experimental.mcp_server.catalog import CatalogSnapshots
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
cleanup_router_config_variables()
global_mcp_server_manager.catalog = TargetCatalog(global_mcp_server_manager)
global_mcp_server_manager.catalog = CatalogSnapshots(global_mcp_server_manager)
await initialize(config=config_path, debug=True)
for server_id, upstream in tuple(global_mcp_server_manager.registry.items()):
if upstream.server_name != "math_restricted":

View file

@ -3417,3 +3417,106 @@ async def test_cancelled_modern_catalog_load_prevents_tool_execution() -> None:
def test_modern_upstream_rejects_legacy_sse_transport() -> None:
with pytest.raises(ValueError, match="transport"):
MCPClient(protocol_version="2026-07-28", transport_type=MCPTransport.sse)
@pytest.mark.asyncio
@pytest.mark.parametrize("method, field, item", [
("tools/list", "tools", {"name": "second", "inputSchema": {"type": "object"}}),
("prompts/list", "prompts", {"name": "second"}),
("resources/list", "resources", {"name": "second", "uri": "status://second"}),
("resources/templates/list", "resourceTemplates", {"name": "second", "uriTemplate": "status://{name}"}),
])
async def test_single_catalog_page_preserves_cursor_metadata_and_request_cursor(method, field, item):
from mcp.types import ListToolsRequest, ListPromptsRequest, ListResourcesRequest, ListResourceTemplatesRequest, PaginatedRequestParams
def respond(request: httpx2.Request) -> httpx2.Response:
if request.method != "POST":
return httpx2.Response(405)
payload = _JSONRPC_MESSAGE_ADAPTER.validate_json(request.content)
if not isinstance(payload, JSONRPCRequest):
return httpx2.Response(202)
if payload.method == "initialize":
return httpx2.Response(
200,
json={
"jsonrpc": "2.0",
"id": payload.id,
"result": {
"protocolVersion": LATEST_HANDSHAKE_VERSION,
"capabilities": {"tools": {}, "prompts": {}, "resources": {}},
"serverInfo": {"name": "pagination", "version": "1"},
},
},
)
assert payload.method == method
assert payload.params is not None
assert payload.params["cursor"] == "upstream-position"
return httpx2.Response(
200,
json={
"jsonrpc": "2.0",
"id": payload.id,
"result": {
field: [item],
"nextCursor": "upstream-next",
"_meta": {"revision": "revision-two"},
},
},
)
client = _MockTransportClient(respond, server_url="https://upstream.example.com/mcp")
request_type = {
"tools/list": ListToolsRequest, "prompts/list": ListPromptsRequest,
"resources/list": ListResourcesRequest, "resources/templates/list": ListResourceTemplatesRequest,
}[method]
result = await client.list_page(request_type(params=PaginatedRequestParams(cursor="upstream-position")))
assert result.model_dump(by_alias=True)[field][0]["name"] == "second"
assert result.next_cursor == "upstream-next"
assert result.meta == {"revision": "revision-two"}
@pytest.mark.asyncio
@pytest.mark.parametrize("method", ["prompts/list", "resources/list", "resources/templates/list"])
@pytest.mark.parametrize("failure", ["unadvertised", "method_missing", "upstream_error"])
@pytest.mark.parametrize("cursor", [None, "continuation"])
async def test_optional_catalog_distinguishes_absent_capability_from_failed_continuation(method, failure, cursor):
from mcp.types import (
ListPromptsRequest, ListResourcesRequest, ListResourceTemplatesRequest, PaginatedRequestParams,
)
methods = []
def respond(request: httpx2.Request) -> httpx2.Response:
if request.method != "POST":
return httpx2.Response(405)
payload = _JSONRPC_MESSAGE_ADAPTER.validate_json(request.content)
if not isinstance(payload, JSONRPCRequest):
return httpx2.Response(202)
methods.append(payload.method)
if payload.method == "initialize":
return httpx2.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": {
"protocolVersion": LATEST_HANDSHAKE_VERSION,
"capabilities": {} if failure == "unadvertised" else {"prompts": {}, "resources": {}},
"serverInfo": {"name": "optional", "version": "1"},
}})
assert payload.method == method
return httpx2.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "error": {
"code": -32601 if failure == "method_missing" else -32603, "message": "Upstream unavailable",
}})
client = _MockTransportClient(respond, server_url="https://upstream.example.com/mcp")
request_type = {
"prompts/list": ListPromptsRequest, "resources/list": ListResourcesRequest,
"resources/templates/list": ListResourceTemplatesRequest,
}[method]
request = request_type(params=PaginatedRequestParams(cursor=cursor))
if cursor is not None or failure == "upstream_error":
with pytest.raises(MCPError):
await client.list_page(request)
else:
result = await client.list_page(request)
collection = {"prompts/list": "prompts", "resources/list": "resources", "resources/templates/list": "resource_templates"}[method]
assert getattr(result, collection) == []
assert result.next_cursor is None
if failure == "unadvertised":
assert method not in methods

View file

@ -129,7 +129,8 @@ async def test_load_mcp_tools_follows_pagination(mock_session):
@pytest.mark.asyncio()
async def test_pagination_walk_stops_at_page_cap(mock_session, monkeypatch):
@pytest.mark.parametrize("require_complete", [False, True])
async def test_pagination_walk_stops_at_page_cap(mock_session, monkeypatch, require_complete):
monkeypatch.setattr("litellm.experimental_mcp_client.tools.MCP_TOOL_LISTING_MAX_PAGES", 2)
mock_session.list_tools.side_effect = [
ListToolsResult(
@ -142,6 +143,12 @@ async def test_pagination_walk_stops_at_page_cap(mock_session, monkeypatch):
),
ListToolsResult(tools=[MCPTool(name="tool_2", description="2", inputSchema={})]),
]
if require_complete:
from mcp import MCPError
with pytest.raises(MCPError, match="incomplete"):
await list_tools_with_pagination(mock_session, require_complete=True)
return
result = await list_tools_with_pagination(mock_session)
assert [tool.name for tool in result] == ["tool_0", "tool_1"]
assert mock_session.list_tools.call_count == 2
@ -178,7 +185,8 @@ async def test_pagination_walk_treats_empty_cursor_as_terminal(mock_session):
@pytest.mark.asyncio()
async def test_pagination_walk_stops_at_whole_walk_deadline(mock_session, monkeypatch):
@pytest.mark.parametrize("require_complete", [False, True])
async def test_pagination_walk_stops_at_whole_walk_deadline(mock_session, monkeypatch, require_complete):
import anyio
from litellm.experimental_mcp_client.tools import list_tools_with_pagination
@ -195,6 +203,12 @@ async def test_pagination_walk_stops_at_whole_walk_deadline(mock_session, monkey
)
mock_session.list_tools = slow_page
if require_complete:
from mcp import MCPError
with pytest.raises(MCPError, match="incomplete"):
await list_tools_with_pagination(mock_session, require_complete=True)
return
result = await list_tools_with_pagination(mock_session)
assert [tool.name for tool in result] == ["tool_0"]
@ -466,3 +480,21 @@ def test_transform_mcp_tool_to_anthropic_tool_strips_keys_anthropic_rejects():
assert "oneOf" not in schema_keys
assert anthropic_tool["input_schema"]["properties"] == {"q": {"type": "string"}}
assert anthropic_tool["input_schema"]["required"] == ["q"]
@pytest.mark.asyncio
async def test_incomplete_discovery_cannot_satisfy_a_complete_listing(mock_session):
from mcp.shared.exceptions import MCPError
mock_session.list_tools.return_value = ListToolsResult(
tools=[MCPTool(name="partial", input_schema={"type": "object"})], next_cursor="repeat"
)
with pytest.raises(MCPError, match="incomplete"):
await list_tools_with_pagination(mock_session, require_complete=True)
@pytest.mark.asyncio
async def test_complete_discovery_preserves_tools_when_completeness_is_required(mock_session):
tool = MCPTool(name="complete", input_schema={"type": "object"})
mock_session.list_tools.return_value = ListToolsResult(tools=[tool])
assert await list_tools_with_pagination(mock_session, require_complete=True) == [tool]

View file

@ -6789,6 +6789,16 @@ class TestMCPDcrBridgeDelegateAdmission:
assert exc_info.value.status_code == 503
async def test_key_envelope_retains_verified_key_identity_for_catalog_reauthorization(self):
from litellm.proxy._types import hash_token
key_hash = hash_token("sk-owned-envelope-key")
record = UserAPIKeyAuth(token=key_hash)
with self._patch_key_reload(return_value=record):
admitted = await MCPRequestHandler._reload_admitted_key(key_hash)
assert admitted.api_key == key_hash
assert admitted.via_virtual_key is True
async def test_reload_admitted_key_returns_admin_for_master_key_hash(self):
"""An envelope sealed under the master key has no DB row to reload; the reload resolves it
to the PROXY_ADMIN auth context (api_key is the alias, never the hash) rather than failing.
@ -9798,13 +9808,16 @@ class TestGetUserObjectPermission:
mock_get_perm.assert_not_awaited()
prisma_client.db.litellm_usertable.find_unique.assert_awaited_once()
async def test_missing_user_row_places_no_ceiling(self):
@pytest.mark.parametrize("fresh_policy", [False, True])
async def test_missing_user_row_places_no_ceiling(self, fresh_policy):
"""Whether this human is entitled at all is unknown when their row is absent, which is the
state before the level existed, so it must not deny."""
from litellm.caching.dual_cache import DualCache
prisma_client = self._prisma_with_user(None)
auth = UserAPIKeyAuth(api_key="sk-test", user_id="ghost")
auth.requires_fresh_policy = fresh_policy
prisma_client.writer_db = prisma_client.db
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma_client),
@ -10167,7 +10180,8 @@ class TestScopedSessionAdmission:
@pytest.mark.asyncio
async def test_fresh_mcp_user_permission_link_ignores_cached_and_replica_grants(monkeypatch):
@pytest.mark.parametrize("failure", [RuntimeError("unavailable"), ValueError("User doesn't exist in db")])
async def test_fresh_mcp_user_permission_link_ignores_cached_and_replica_grants(monkeypatch, failure):
from litellm.caching.dual_cache import DualCache
from litellm.proxy import proxy_server
from litellm.proxy._types import LiteLLM_UserTable
@ -10182,7 +10196,7 @@ async def test_fresh_mcp_user_permission_link_ignores_cached_and_replica_grants(
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
assert await MCPRequestHandler._user_object_permission_id("fresh-human", database, check_db_only=True) == "current"
database.db.litellm_usertable.find_unique.assert_not_awaited()
database.writer_db.litellm_usertable.find_unique.side_effect = RuntimeError("unavailable")
database.writer_db.litellm_usertable.find_unique.side_effect = failure
with pytest.raises(HTTPException) as denied:
await MCPRequestHandler._user_object_permission_id("fresh-human", database, check_db_only=True)
assert denied.value.status_code == 503
@ -10226,3 +10240,74 @@ async def test_unreadable_empty_key_scope_cannot_gain_additive_grants(monkeypatc
access = await MCPRequestHandler.get_mcp_server_access(auth)
assert access.server_ids == ()
assert access.scope == "scoped"
@pytest.mark.asyncio
@pytest.mark.parametrize("kind", ["anonymous", "master", "custom"])
async def test_catalog_refresh_preserves_non_database_admission_and_resource_scope(kind):
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
if kind == "anonymous":
assert await MCPRequestHandler.refresh_catalog_authority(None) is None
return
caller = UserAPIKeyAuth(api_key=LITELLM_PROXY_MASTER_KEY_ALIAS if kind == "master" else "custom-subject")
caller.via_virtual_key = kind == "master"
caller.authenticated_by_custom_auth = kind == "custom"
caller.mcp_session_resource_server_id = "only-this-server"
caller.mcp_toolset_id = "only-this-toolset"
refreshed = await MCPRequestHandler.refresh_catalog_authority(caller)
assert refreshed is not caller
assert refreshed.api_key == caller.api_key
assert refreshed.authenticated_by_custom_auth == caller.authenticated_by_custom_auth
assert refreshed.mcp_session_resource_server_id == "only-this-server"
assert refreshed.mcp_toolset_id == "only-this-toolset"
assert refreshed.requires_fresh_policy is True
assert caller.requires_fresh_policy is False
@pytest.mark.asyncio
async def test_catalog_refresh_reads_current_user_org_without_losing_resource_scope(monkeypatch):
from types import SimpleNamespace
from litellm.caching.dual_cache import DualCache
from litellm.proxy import proxy_server
from litellm.proxy._types import LiteLLM_UserTable
current = LiteLLM_UserTable(user_id="catalog-user", organization_id="current-org", user_role="internal_user", teams=[])
table = SimpleNamespace(find_unique=AsyncMock(return_value=current))
database = SimpleNamespace(writer_db=SimpleNamespace(litellm_usertable=table))
monkeypatch.setattr(proxy_server, "prisma_client", database)
monkeypatch.setattr(proxy_server, "user_api_key_cache", DualCache())
caller = UserAPIKeyAuth(user_id="catalog-user", org_id="previous-org", user_role="proxy_admin")
caller.mcp_admitted_user_subject = True
caller.mcp_session_resource_server_id = "scoped-server"
refreshed = await MCPRequestHandler.refresh_catalog_authority(caller)
assert refreshed.org_id == "current-org"
assert refreshed.user_role == "internal_user"
assert refreshed.mcp_session_resource_server_id == "scoped-server"
assert refreshed.mcp_admitted_user_subject is True
assert caller.org_id == "previous-org"
@pytest.mark.asyncio
@pytest.mark.parametrize("current_groups", [[], ["replacement-group"]])
async def test_catalog_refresh_uses_current_virtual_key_policy_and_keeps_session_scope(monkeypatch, current_groups):
permission = LiteLLM_ObjectPermissionTable(object_permission_id="current-policy", mcp_servers=["current-server"])
current = UserAPIKeyAuth(object_permission=permission, object_permission_id="current-policy", team_id="new-team", org_id="new-org", project_id="new-project", user_id="new-owner", access_group_ids=current_groups)
reload_key = AsyncMock(return_value=current)
monkeypatch.setattr(MCPRequestHandler, "_reload_admitted_key", reload_key)
caller = UserAPIKeyAuth(api_key="owned-key-hash", team_id="old-team", org_id="old-org", project_id="old-project", user_id="old-owner", access_group_ids=["original-group"])
caller.via_virtual_key = True
caller.mcp_session_resource_server_id = "session-server"
caller.mcp_toolset_id = "session-toolset"
refreshed = await MCPRequestHandler.refresh_catalog_authority(caller)
reload_key.assert_awaited_once_with("owned-key-hash", check_db_only=True)
assert refreshed.object_permission == permission
assert refreshed.object_permission_id == "current-policy"
assert (refreshed.team_id, refreshed.org_id, refreshed.project_id, refreshed.user_id) == ("new-team", "new-org", "new-project", "new-owner")
assert refreshed.mcp_session_resource_server_id == "session-server"
assert refreshed.mcp_toolset_id == "session-toolset"
assert refreshed.via_virtual_key and refreshed.requires_fresh_policy
assert caller.team_id == "old-team" and not caller.requires_fresh_policy
assert refreshed.access_group_ids == current_groups
assert caller.access_group_ids == ["original-group"]

View file

@ -81,7 +81,7 @@ def config_only_mcp_manager_factory():
@pytest.fixture(autouse=True)
def _hermetic_mcp_server_registry():
from litellm.proxy._experimental.mcp_server.catalog import TargetCatalog
from litellm.proxy._experimental.mcp_server.catalog import CatalogSnapshots
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
@ -92,7 +92,7 @@ def _hermetic_mcp_server_registry():
saved_tools = global_mcp_tool_registry.published_tools
global_mcp_tool_registry.published_tools = {}
saved_catalog = global_mcp_server_manager.catalog
global_mcp_server_manager.catalog = TargetCatalog(global_mcp_server_manager)
global_mcp_server_manager.catalog = CatalogSnapshots(global_mcp_server_manager)
saved_registry = dict(global_mcp_server_manager.registry)
saved_config_servers = dict(global_mcp_server_manager.config_mcp_servers)
saved_tool_mapping = dict(global_mcp_server_manager.tool_name_to_mcp_server_name_mapping)

View file

@ -0,0 +1,376 @@
import asyncio
from collections.abc import Sequence
import pytest
from mcp.shared.exceptions import MCPError
from mcp.types import ListToolsResult, Tool
from litellm.proxy._experimental.mcp_server import catalog
def page(name: str, cursor: str | None = None, revision: str = "stable") -> ListToolsResult:
return ListToolsResult(
tools=[Tool(name=name, input_schema={"type": "object"})],
next_cursor=cursor,
meta={"revision": revision},
)
async def listing(
fetch,
cursor: str | None = None,
*,
caller_scope: str = "caller-and-scope",
snapshot: str = "registry-generation",
servers: Sequence[str] = ("a", "b"),
now: int = 100,
) -> ListToolsResult:
return await catalog.list_tools_page(
cursor=cursor,
caller_scope=caller_scope,
snapshot=snapshot,
server_ids=tuple(servers),
fetch=fetch,
now=now,
)
@pytest.mark.asyncio
async def test_listing_continues_on_another_replica_and_cursor_is_reusable(monkeypatch):
monkeypatch.setenv("LITELLM_SALT_KEY", "shared-test-key")
async def replica_a(server_id: str, cursor: str | None) -> ListToolsResult:
assert cursor is None
return page(server_id + "1", "page-two" if server_id == "a" else None)
async def replica_b(server_id: str, cursor: str | None) -> ListToolsResult:
assert (server_id, cursor) == ("a", "page-two")
return page("a2")
first = await listing(replica_a, servers=("b", "a"))
assert [tool.name for tool in first.tools] == ["a1", "b1"]
assert first.next_cursor
for _ in range(2):
second = await listing(replica_b, first.next_cursor)
assert [tool.name for tool in second.tools] == ["a2"]
assert second.next_cursor is None
@pytest.mark.asyncio
@pytest.mark.parametrize(
"changes",
[
{"caller_scope": "other-caller"},
{"snapshot": "new-registry-generation"},
{"servers": ("b",)},
{"now": 1000},
],
)
async def test_continuation_rejects_changed_binding_before_upstream_dispatch(monkeypatch, changes):
monkeypatch.setenv("LITELLM_SALT_KEY", "shared-test-key")
async def first_page(server_id: str, cursor: str | None) -> ListToolsResult:
return page(server_id, "next")
async def forbidden_dispatch(server_id: str, cursor: str | None) -> ListToolsResult:
pytest.fail("Rejected continuation must not contact upstream")
first = await listing(first_page)
with pytest.raises(MCPError, match=r"fresh listing|expired"):
await listing(forbidden_dispatch, first.next_cursor, **changes)
@pytest.mark.asyncio
async def test_complete_single_page_needs_no_key_but_continuation_does(monkeypatch):
monkeypatch.delenv("LITELLM_SALT_KEY", raising=False)
async def complete(server_id: str, cursor: str | None) -> ListToolsResult:
return page(server_id)
async def incomplete(server_id: str, cursor: str | None) -> ListToolsResult:
return page(server_id, "next")
result = await listing(complete)
assert len(result.tools) == 2
assert result.next_cursor is None
with pytest.raises(MCPError, match="LITELLM_SALT_KEY"):
await listing(incomplete)
@pytest.mark.asyncio
async def test_repeated_upstream_cursor_requires_restart(monkeypatch):
monkeypatch.setenv("LITELLM_SALT_KEY", "shared-test-key")
async def repeated(server_id: str, cursor: str | None) -> ListToolsResult:
return page(server_id, "repeated")
first = await listing(repeated)
with pytest.raises(MCPError, match=r"repeated.*cursor"):
await listing(repeated, first.next_cursor)
@pytest.mark.asyncio
async def test_changed_upstream_revision_requires_restart(monkeypatch):
monkeypatch.setenv("LITELLM_SALT_KEY", "shared-test-key")
async def original(server_id: str, cursor: str | None) -> ListToolsResult:
return page(server_id, "next")
async def revised(server_id: str, cursor: str | None) -> ListToolsResult:
return page(server_id, revision="changed")
first = await listing(original)
with pytest.raises(MCPError, match="fresh listing"):
await listing(revised, first.next_cursor)
@pytest.mark.asyncio
async def test_failed_upstream_remains_visible_when_other_sources_continue(monkeypatch):
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import SERVER_OUTCOMES_META_KEY
monkeypatch.setenv("LITELLM_SALT_KEY", "shared-test-key")
async def fetch(server_id: str, cursor: str | None) -> ListToolsResult:
if server_id == "a":
return ListToolsResult(tools=[], meta={SERVER_OUTCOMES_META_KEY: {"a": {"tag": "timeout"}}})
return page("b2" if cursor else "b1", None if cursor else "next")
first = await listing(fetch)
assert [tool.name for tool in first.tools] == ["b1"]
assert first.meta[SERVER_OUTCOMES_META_KEY]["a"] == {"tag": "timeout"}
second = await listing(fetch, first.next_cursor)
assert [tool.name for tool in second.tools] == ["b2"]
assert second.meta[SERVER_OUTCOMES_META_KEY]["a"] == {"tag": "timeout"}
@pytest.mark.asyncio
async def test_cursor_cannot_be_reused_for_a_different_catalog_kind(monkeypatch):
monkeypatch.setenv("LITELLM_SALT_KEY", "shared-test-key")
async def first_page(server_id: str, cursor: str | None) -> ListToolsResult:
return page(server_id, "next")
async def forbidden_dispatch(server_id: str, cursor: str | None) -> ListToolsResult:
pytest.fail("Wrong-kind cursor must not contact upstream")
first = await listing(first_page)
with pytest.raises(MCPError, match="Invalid pagination state"):
await catalog.paginate_catalog(
method="resources/list",
cursor=first.next_cursor,
caller_scope="caller-and-scope",
snapshot="registry-generation",
server_ids=("a", "b"),
fetch=forbidden_dispatch,
now=100,
)
@pytest.mark.asyncio
async def test_authenticated_state_with_invalid_catalog_schema_is_rejected(monkeypatch):
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok
from litellm.proxy._experimental.mcp_server.state_tokens import seal_state
monkeypatch.setenv("LITELLM_SALT_KEY", "shared-test-key")
token = seal_state({"not": "catalog state"}, purpose="mcp.catalog.list.v1:tools/list", expires_at=200, now=100)
assert isinstance(token, Ok)
async def forbidden_dispatch(server_id, cursor):
pytest.fail("Malformed catalog state must not dispatch")
with pytest.raises(MCPError, match="Invalid pagination state"):
await listing(forbidden_dispatch, token.ok)
@pytest.mark.asyncio
async def test_following_page_does_not_extend_original_expiry(monkeypatch):
monkeypatch.setenv("LITELLM_SALT_KEY", "shared-test-key")
async def fetch(server_id, cursor):
return page(server_id, "second" if cursor is None else "third")
first = await listing(fetch, now=100)
second = await listing(fetch, first.next_cursor, now=699)
assert second.next_cursor
with pytest.raises(MCPError, match="expired"):
await listing(fetch, second.next_cursor, now=700)
@pytest.mark.asyncio
async def test_page_limit_rejects_an_unending_upstream(monkeypatch):
from litellm import constants
monkeypatch.setenv("LITELLM_SALT_KEY", "shared-test-key")
monkeypatch.setattr(constants, "MCP_TOOL_LISTING_MAX_PAGES", 2)
async def fetch(server_id, cursor):
return page(server_id, "second" if cursor is None else "third")
first = await listing(fetch)
with pytest.raises(MCPError, match="limit"):
await listing(fetch, first.next_cursor)
@pytest.mark.parametrize(
"change",
[
{"_caller": None},
{"mcp_auth_header": "another upstream credential"},
{"mcp_servers": ("another-scope",)},
{"client_ip": "192.0.2.2"},
{"raw_headers": {"Authorization": "another bearer"}},
{"oauth2_headers": {"Authorization": "another upstream bearer"}},
{"mcp_server_auth_headers": {"a": {"Authorization": "another per-server bearer"}}},
{"protocol_version": "2024-11-05"},
{"mcp_proxy_mode": True},
],
)
def test_caller_binding_covers_identity_scope_and_forwarded_credentials(change):
from dataclasses import replace
from litellm.proxy._experimental.mcp_server.contracts import OperationContext
from litellm.proxy._types import UserAPIKeyAuth
original = OperationContext(_caller=UserAPIKeyAuth(api_key="synthetic-key", user_id="owner"))
assert catalog._caller_scope(original, ()) != catalog._caller_scope(replace(original, **change), ())
def test_caller_binding_ignores_transport_headers_and_header_case():
from dataclasses import replace
from litellm.proxy._experimental.mcp_server.contracts import OperationContext
from litellm.proxy._types import UserAPIKeyAuth
original = OperationContext(
_caller=UserAPIKeyAuth(user_id="owner"), raw_headers={"Authorization": "Bearer synthetic"}
)
retry = replace(original, raw_headers={"authorization": "Bearer synthetic", "mcp-session-id": "replica-b-session"})
assert catalog._caller_scope(original, ()) == catalog._caller_scope(retry, ())
@pytest.mark.asyncio
async def test_upstream_pages_overlap_and_keep_deterministic_order(monkeypatch):
monkeypatch.setenv("LITELLM_SALT_KEY", "shared-test-key")
both_started = asyncio.Event()
started = set()
async def fetch(server_id, cursor):
started.add(server_id)
if len(started) == 2:
both_started.set()
await asyncio.wait_for(both_started.wait(), timeout=1)
return page(server_id)
result = await listing(fetch, servers=("b", "a"))
assert [tool.name for tool in result.tools] == ["a", "b"]
@pytest.mark.asyncio
async def test_cursor_history_allows_the_full_supported_page_count(monkeypatch):
from litellm.constants import MCP_TOOL_LISTING_MAX_PAGES
monkeypatch.setenv("LITELLM_SALT_KEY", "shared-test-key")
async def fetch(server_id, cursor):
index = int(cursor or "0")
return page(str(index), str(index + 1) if index + 1 < MCP_TOOL_LISTING_MAX_PAGES else None)
cursor = None
for index in range(MCP_TOOL_LISTING_MAX_PAGES):
result = await listing(fetch, cursor, servers=("a",))
assert result.tools[0].name == str(index)
cursor = result.next_cursor
assert bool(cursor) == (index + 1 < MCP_TOOL_LISTING_MAX_PAGES)
@pytest.mark.asyncio
async def test_failed_page_cancels_and_drains_other_upstream_requests(monkeypatch):
monkeypatch.setenv("LITELLM_SALT_KEY", "shared-test-key")
pending_started = asyncio.Event()
pending_closed = asyncio.Event()
async def fetch(server_id, cursor):
if server_id == "a":
await asyncio.wait_for(pending_started.wait(), timeout=1)
raise ValueError("upstream unavailable")
pending_started.set()
try:
await asyncio.Event().wait()
finally:
pending_closed.set()
return page(server_id)
with pytest.raises(ValueError, match="upstream unavailable"):
await listing(fetch)
assert pending_closed.is_set()
@pytest.mark.asyncio
async def test_gateway_tools_continuation_rejects_an_upstream_failure(monkeypatch):
from unittest.mock import AsyncMock
from mcp.types import PaginatedRequestParams
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server import operations
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ServerListOk, classify_list_exception
from litellm.types.mcp import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
monkeypatch.setenv("LITELLM_SALT_KEY", "catalog-failure-test")
monkeypatch.setattr(proxy_server, "prisma_client", None)
server = MCPServer(server_id="pages", name="pages", transport=MCPTransport.http)
monkeypatch.setattr(operations.global_mcp_server_manager, "registry", {server.server_id: server})
fetch = AsyncMock(side_effect=[
(page("first", "next"), ServerListOk(tool_count=1)),
(ListToolsResult(tools=[]), classify_list_exception(TimeoutError("upstream secret"))),
])
monkeypatch.setattr(catalog, "get_filtered_server_tools", fetch)
context = operations.prepare_context()
first = await catalog.aggregate_gateway_tools(context, PaginatedRequestParams(), [server], {})
assert first.next_cursor and [tool.name for tool in first.tools] == ["first"]
with pytest.raises(MCPError, match="Upstream continuation failed; start a fresh listing") as denied:
await catalog.aggregate_gateway_tools(context, PaginatedRequestParams(cursor=first.next_cursor), [server], {})
assert "upstream secret" not in str(denied.value)
assert fetch.await_count == 2
assert fetch.await_args.kwargs["params"].cursor == "next"
@pytest.mark.asyncio
@pytest.mark.parametrize("request_name,result_name,field", [
("ListPromptsRequest", "ListPromptsResult", "prompts"),
("ListResourcesRequest", "ListResourcesResult", "resources"),
("ListResourceTemplatesRequest", "ListResourceTemplatesResult", "resource_templates"),
])
async def test_optional_gateway_catalog_reports_initial_failure_and_rejects_failed_continuation(monkeypatch, request_name, result_name, field):
from unittest.mock import AsyncMock
from mcp import types
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server import operations
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import SERVER_OUTCOMES_META_KEY
from litellm.types.mcp import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
monkeypatch.setenv("LITELLM_SALT_KEY", "catalog-failure-test")
monkeypatch.setattr(proxy_server, "prisma_client", None)
server = MCPServer(server_id="pages", name="pages", transport=MCPTransport.http)
monkeypatch.setattr(operations.global_mcp_server_manager, "registry", {server.server_id: server})
monkeypatch.setattr(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server]))
fetch = AsyncMock(side_effect=TimeoutError("upstream secret"))
monkeypatch.setattr(catalog, "fetch_optional_catalog_page", fetch)
context = operations.prepare_context()
request = getattr(types, request_name)
failed = await catalog.list_gateway_catalog(context, request())
assert getattr(failed, field) == [] and failed.next_cursor is None
assert next(iter(failed.meta[SERVER_OUTCOMES_META_KEY].values()))["status"] == "timeout"
assert "upstream secret" not in failed.model_dump_json()
fetch.side_effect = [getattr(types, result_name)(**{field: [], "next_cursor": "next"}), TimeoutError("upstream secret")]
first = await catalog.list_gateway_catalog(context, request())
assert first.next_cursor
with pytest.raises(MCPError, match="Upstream continuation failed; start a fresh listing") as denied:
await catalog.list_gateway_catalog(context, request(params=types.PaginatedRequestParams(cursor=first.next_cursor)))
assert "upstream secret" not in str(denied.value)
assert fetch.await_count == 3
assert fetch.await_args.args[-1] == "next"

View file

@ -101,8 +101,8 @@ def isolate_global_mcp_registry(monkeypatch):
"""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
from litellm.proxy._experimental.mcp_server.catalog import TargetCatalog
monkeypatch.setattr(global_mcp_server_manager, "catalog", TargetCatalog(global_mcp_server_manager))
from litellm.proxy._experimental.mcp_server.catalog import CatalogSnapshots
monkeypatch.setattr(global_mcp_server_manager, "catalog", CatalogSnapshots(global_mcp_server_manager))
snapshot = dict(global_mcp_server_manager.registry)
yield
global_mcp_server_manager.registry.clear()

View file

@ -18804,3 +18804,189 @@ async def test_catalog_rejects_a_changed_anchored_issuer_during_discovery(monkey
await manager.ensure_oauth_metadata_discovered(server)
assert rejected.value.status_code == 503
discovery.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize(
"request_cursor,next_cursor,expected_owner",
[(None, None, "notes"), (None, "next", "other"), ("last", None, "other")],
)
async def test_catalog_page_registers_bare_routes_only_for_complete_initial_discovery(
request_cursor, next_cursor, expected_owner
):
from types import SimpleNamespace
from mcp.types import ListToolsResult, PaginatedRequestParams
manager = _catalog_manager(LIST_NOTES)
server = _notes_server()
other = MCPServer(server_id="other", name="other", transport=MCPTransport.http)
manager.registry = {server.server_id: server, other.server_id: other}
manager._create_prefixed_tools([LIST_NOTES], other)
manager._create_mcp_client.return_value = SimpleNamespace(
list_tools_page=AsyncMock(return_value=ListToolsResult(tools=[LIST_NOTES], next_cursor=next_cursor))
)
result = await manager.get_tools_page(server, params=PaginatedRequestParams(cursor=request_cursor))
assert [tool.name for tool in result.tools] == ["notes-list_notes"]
assert result.next_cursor == next_cursor
assert manager._get_mcp_server_from_tool_name("notes-list_notes").server_id == "notes"
assert manager._get_mcp_server_from_tool_name("list_notes").server_id == expected_owner
@pytest.mark.asyncio
async def test_paginated_listing_keeps_earlier_tool_metadata_and_caller_isolation(monkeypatch):
from mcp.types import ListToolsResult, PaginatedRequestParams
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server import catalog, operations
monkeypatch.setenv("LITELLM_SALT_KEY", "listed-tools-test")
monkeypatch.setattr(proxy_server, "prisma_client", None)
first = MCPTool(name="first", description="First page", input_schema={"type": "object"})
second = MCPTool(name="second", description="Second page", input_schema={"type": "object"})
server = MCPServer(server_id="pages", name="pages", transport=MCPTransport.http)
manager = MCPServerManager()
manager.registry = {server.server_id: server}
client = AsyncMock()
client._last_initialize_instructions = None
client.list_tools_page.side_effect = [
ListToolsResult(tools=[first], next_cursor="next"),
ListToolsResult(tools=[second]),
ListToolsResult(tools=[second]),
]
manager._create_mcp_client = AsyncMock(return_value=client)
monkeypatch.setattr(operations, "global_mcp_server_manager", manager)
caller = UserAPIKeyAuth(api_key="owned-caller", user_id="alice")
context = operations.prepare_context(caller)
first_page = await catalog.aggregate_gateway_tools(
context, PaginatedRequestParams(), [server], {}, record_listing=True
)
assert first_page.next_cursor
await catalog.aggregate_gateway_tools(
context, PaginatedRequestParams(cursor=first_page.next_cursor), [server], {}, record_listing=True
)
identity = ListedToolsCaller(user_api_key_auth=caller)
assert manager.get_listed_tool(server, "first", identity) == first
assert manager.get_listed_tool(server, "second", identity) == second
other = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="other-caller", user_id="bob"))
assert manager.get_listed_tool(server, "first", other) is None
await catalog.aggregate_gateway_tools(context, PaginatedRequestParams(), [server], {}, record_listing=True)
assert manager.get_listed_tool(server, "first", identity) is None
assert manager.get_listed_tool(server, "second", identity) == second
@pytest.mark.asyncio
@pytest.mark.parametrize("request_name,result_field", [("ListPromptsRequest", "prompts"), ("ListResourcesRequest", "resources"), ("ListResourceTemplatesRequest", "resource_templates")])
@pytest.mark.parametrize("cursor", [None, "next"])
async def test_disabled_stdio_catalog_is_empty_initially_and_rejects_continuation(monkeypatch, request_name, result_field, cursor):
from mcp import types
monkeypatch.delenv("LITELLM_ENABLE_MCP_STDIO", raising=False)
server = MCPServer(server_id="disabled", name="disabled", transport=MCPTransport.stdio, command="blocked-executable", args=[])
request = getattr(types, request_name)(params=types.PaginatedRequestParams(cursor=cursor))
manager = MCPServerManager()
if cursor is not None:
with pytest.raises(RuntimeError, match="Upstream catalog is unavailable"):
await manager.get_optional_catalog_page(server, request, None)
else:
page = await manager.get_optional_catalog_page(server, request, None)
assert getattr(page, result_field) == []
assert page.next_cursor is None
@pytest.mark.asyncio
@pytest.mark.parametrize("cursor", [None, "next"])
async def test_disabled_stdio_tools_are_empty_initially_and_reject_continuation(monkeypatch, cursor):
from mcp.types import PaginatedRequestParams
monkeypatch.delenv("LITELLM_ENABLE_MCP_STDIO", raising=False)
server = MCPServer(server_id="disabled", name="disabled", transport=MCPTransport.stdio, command="blocked", args=[])
manager = MCPServerManager()
if cursor is not None:
with pytest.raises(RuntimeError, match="Upstream catalog is unavailable"):
await manager.get_tools_page(server, params=PaginatedRequestParams(cursor=cursor))
else:
page = await manager.get_tools_page(server, params=PaginatedRequestParams())
assert page.tools == []
assert page.next_cursor is None
@pytest.mark.asyncio
async def test_failed_aggregate_continuation_preserves_only_delivered_tool_metadata(monkeypatch):
from mcp.shared.exceptions import MCPError
from mcp.types import ListToolsResult, PaginatedRequestParams
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server import catalog, operations
monkeypatch.setenv("LITELLM_SALT_KEY", "failed-listing-test")
monkeypatch.setattr(proxy_server, "prisma_client", None)
first = MCPTool(name="first", description="Delivered", input_schema={"type": "object"})
unseen = MCPTool(name="unseen", description="Never delivered", input_schema={"type": "object"})
servers = [MCPServer(server_id=name, name=name, transport=MCPTransport.http) for name in ("alpha", "beta")]
manager = MCPServerManager()
manager.registry = {server.server_id: server for server in servers}
clients = {server.server_id: AsyncMock() for server in servers}
for client in clients.values():
client._last_initialize_instructions = None
client.list_tools_page.side_effect = [
ListToolsResult(tools=[first], next_cursor="next"),
ListToolsResult(tools=[unseen], next_cursor="next" if client is clients["beta"] else None),
]
async def create_client(server, **kwargs):
return clients[server.server_id]
manager._create_mcp_client = create_client
monkeypatch.setattr(operations, "global_mcp_server_manager", manager)
caller = UserAPIKeyAuth(api_key="owned-caller", user_id="alice")
context = operations.prepare_context(caller)
initial = await catalog.aggregate_gateway_tools(context, PaginatedRequestParams(), servers, {}, record_listing=True)
assert initial.next_cursor
with pytest.raises(MCPError, match="repeated a pagination cursor"):
await catalog.aggregate_gateway_tools(
context, PaginatedRequestParams(cursor=initial.next_cursor), servers, {}, record_listing=True
)
identity = ListedToolsCaller(user_api_key_auth=caller)
for server in servers:
assert manager.get_listed_tool(server, "first", identity) == first
assert manager.get_listed_tool(server, "unseen", identity) is None
@pytest.mark.asyncio
@pytest.mark.parametrize("can_seal", [False, True])
async def test_aggregate_publishes_complete_bare_routes_only_after_delivering_a_page(monkeypatch, can_seal):
from mcp.shared.exceptions import MCPError
from mcp.types import ListToolsResult, PaginatedRequestParams
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server import catalog, operations
monkeypatch.delenv("LITELLM_SALT_KEY", raising=False)
if can_seal:
monkeypatch.setenv("LITELLM_SALT_KEY", "delivered-page-test")
monkeypatch.setattr(proxy_server, "prisma_client", None)
tool = MCPTool(name="first", input_schema={"type": "object"})
servers = [MCPServer(server_id=name, name=name, transport=MCPTransport.http) for name in ("alpha", "beta")]
manager = MCPServerManager()
manager.registry = {server.server_id: server for server in servers}
clients = {server.server_id: AsyncMock() for server in servers}
for name, client in clients.items():
client._last_initialize_instructions = None
client.list_tools_page.return_value = ListToolsResult(tools=[tool], next_cursor="next" if name == "beta" else None)
async def create_client(server, **kwargs):
return clients[server.server_id]
manager._create_mcp_client = create_client
monkeypatch.setattr(operations, "global_mcp_server_manager", manager)
context = operations.prepare_context(UserAPIKeyAuth(api_key="owned-caller", user_id="alice"))
listing = catalog.aggregate_gateway_tools(context, PaginatedRequestParams(), servers, {}, record_listing=True)
if can_seal:
assert (await listing).next_cursor
assert manager._get_mcp_server_from_tool_name("first").server_id == "alpha"
else:
with pytest.raises(MCPError, match="LITELLM_SALT_KEY"):
await listing
assert manager._get_mcp_server_from_tool_name("first") is None

View file

@ -23,6 +23,7 @@ from mcp.types import (
ResourceTemplate,
TextContent,
TextResourceContents,
Tool,
)
from mcp.types import Tool as MCPTool
from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS, LATEST_HANDSHAKE_VERSION, MODERN_PROTOCOL_VERSIONS
@ -1292,7 +1293,7 @@ async def test_get_tools_from_mcp_servers_continues_when_one_server_fails():
):
if server.name == "working_server":
# Working server returns tools
tool1 = MagicMock()
tool1 = Tool(name="placeholder", inputSchema={})
tool1.name = "working_tool_1"
tool1.description = "Working tool 1"
tool1.input_schema = {}
@ -1308,12 +1309,12 @@ async def test_get_tools_from_mcp_servers_continues_when_one_server_fails():
mock_manager,
):
with patch(
"litellm.proxy._experimental.mcp_server.operations.verbose_logger",
) as mock_logger:
"litellm.proxy._experimental.mcp_server.catalog.verbose_logger",
) as mock_logger, patch("litellm.proxy._experimental.mcp_server.operations.verbose_logger", mock_logger):
# Test with server-specific auth headers
mcp_server_auth_headers = {
"working": "Bearer working-token",
"failing": "Bearer failing-token",
"working": {"Authorization": "Bearer working-token"},
"failing": {"Authorization": "Bearer failing-token"},
}
result = await _get_tools_from_mcp_servers(
@ -1404,12 +1405,12 @@ async def test_get_tools_from_mcp_servers_handles_all_servers_failing():
mock_manager,
):
with patch(
"litellm.proxy._experimental.mcp_server.operations.verbose_logger",
) as mock_logger:
"litellm.proxy._experimental.mcp_server.catalog.verbose_logger",
) as mock_logger, patch("litellm.proxy._experimental.mcp_server.operations.verbose_logger", mock_logger):
# Test with server-specific auth headers
mcp_server_auth_headers = {
"failing1": "Bearer failing1-token",
"failing2": "Bearer failing2-token",
"failing1": {"Authorization": "Bearer failing1-token"},
"failing2": {"Authorization": "Bearer failing2-token"},
}
result = await _get_tools_from_mcp_servers(
@ -4381,7 +4382,7 @@ async def test_list_tools_single_server_unprefixed_names():
raw_headers=None,
**kwargs,
):
tool = MagicMock()
tool = Tool(name="placeholder", inputSchema={})
tool.name = f"{server.alias}-toolA" if add_prefix else "toolA"
tool.description = "desc"
tool.input_schema = {}
@ -4459,7 +4460,7 @@ async def test_list_tools_multiple_servers_prefixed_names():
raw_headers=None,
**kwargs,
):
tool = MagicMock()
tool = Tool(name="placeholder", inputSchema={})
# When multiple servers, add_prefix should be True -> prefixed names
tool.name = f"{server.alias}-toolA" if add_prefix else "toolA"
tool.description = "desc"
@ -4873,22 +4874,22 @@ async def test_list_tools_filters_by_key_team_permissions():
**kwargs,
):
# Return 4 tools, but only 2 should be allowed
tool1 = MagicMock()
tool1 = Tool(name="placeholder", inputSchema={})
tool1.name = "tool1"
tool1.description = "Tool 1"
tool1.input_schema = {}
tool2 = MagicMock()
tool2 = Tool(name="placeholder", inputSchema={})
tool2.name = "tool2"
tool2.description = "Tool 2"
tool2.input_schema = {}
tool3 = MagicMock()
tool3 = Tool(name="placeholder", inputSchema={})
tool3.name = "tool3"
tool3.description = "Tool 3 - not allowed"
tool3.input_schema = {}
tool4 = MagicMock()
tool4 = Tool(name="placeholder", inputSchema={})
tool4.name = "tool4"
tool4.description = "Tool 4 - not allowed"
tool4.input_schema = {}
@ -4984,22 +4985,22 @@ async def test_list_tools_with_team_tool_permissions_inheritance():
**kwargs,
):
# Return 4 tools
tool1 = MagicMock()
tool1 = Tool(name="placeholder", inputSchema={})
tool1.name = "tool1"
tool1.description = "Tool 1"
tool1.input_schema = {}
tool2 = MagicMock()
tool2 = Tool(name="placeholder", inputSchema={})
tool2.name = "tool2"
tool2.description = "Tool 2"
tool2.input_schema = {}
tool3 = MagicMock()
tool3 = Tool(name="placeholder", inputSchema={})
tool3.name = "tool3"
tool3.description = "Tool 3"
tool3.input_schema = {}
tool4 = MagicMock()
tool4 = Tool(name="placeholder", inputSchema={})
tool4.name = "tool4"
tool4.description = "Tool 4"
tool4.input_schema = {}
@ -5081,17 +5082,17 @@ async def test_list_tools_with_no_tool_permissions_shows_all():
**kwargs,
):
# Return 3 tools
tool1 = MagicMock()
tool1 = Tool(name="placeholder", inputSchema={})
tool1.name = "tool1"
tool1.description = "Tool 1"
tool1.input_schema = {}
tool2 = MagicMock()
tool2 = Tool(name="placeholder", inputSchema={})
tool2.name = "tool2"
tool2.description = "Tool 2"
tool2.input_schema = {}
tool3 = MagicMock()
tool3 = Tool(name="placeholder", inputSchema={})
tool3.name = "tool3"
tool3.description = "Tool 3"
tool3.input_schema = {}
@ -5182,22 +5183,22 @@ async def test_list_tools_strips_prefix_when_matching_permissions():
**kwargs,
):
# Return tools WITH prefix (as they come from MCP server)
tool1 = MagicMock()
tool1 = Tool(name="placeholder", inputSchema={})
tool1.name = "GITMCP-fetch_litellm_documentation" # Prefixed
tool1.description = "Fetch docs"
tool1.input_schema = {}
tool2 = MagicMock()
tool2 = Tool(name="placeholder", inputSchema={})
tool2.name = "GITMCP-search_litellm_documentation" # Prefixed, not in allowed list
tool2.description = "Search docs"
tool2.input_schema = {}
tool3 = MagicMock()
tool3 = Tool(name="placeholder", inputSchema={})
tool3.name = "GITMCP-search_litellm_code" # Prefixed
tool3.description = "Search code"
tool3.input_schema = {}
tool4 = MagicMock()
tool4 = Tool(name="placeholder", inputSchema={})
tool4.name = "GITMCP-fetch_generic_url_content" # Prefixed, not in allowed list
tool4.description = "Fetch URL"
tool4.input_schema = {}
@ -5793,7 +5794,7 @@ async def test_get_tools_from_mcp_servers_returns_tools_when_success_logging_fai
server_a.auth_type = None
server_a.extra_headers = None
tool_1 = MagicMock()
tool_1 = Tool(name="placeholder", inputSchema={})
tool_1.name = "server_a-tool_1"
dummy_logging_obj = MagicMock()
@ -6119,7 +6120,7 @@ async def test_get_tools_from_mcp_servers_injects_stored_oauth2_token():
# Simulate the DB returning a valid credential for this user+server
prefetched_creds = {SERVER_ID: {"access_token": STORED_TOKEN, "server_id": SERVER_ID}}
tool_1 = MagicMock()
tool_1 = Tool(name="placeholder", inputSchema={})
tool_1.name = "atlassian_test-search"
with (
@ -6744,7 +6745,7 @@ async def test_list_tools_with_legacy_db_m2m_server_resolves_oauth2_flow():
)
)
tool_1 = MagicMock()
tool_1 = Tool(name="placeholder", inputSchema={})
tool_1.name = "legacy_m2m-tool"
captured_extra_headers = None
@ -9839,7 +9840,7 @@ async def test_aggregate_listing_reports_per_server_outcomes():
async def mock_get_tools_from_server(server, **kwargs):
if server.name == "working_server":
tool1 = MagicMock()
tool1 = Tool(name="placeholder", inputSchema={})
tool1.name = "working_tool_1"
tool1.description = "Working tool 1"
tool1.input_schema = {}
@ -10685,7 +10686,7 @@ async def test_list_tools_injects_byok_credential_for_non_oauth2_auth_types(auth
async def mock_get_tools_from_server(server, mcp_auth_header=None, add_prefix=False, **kwargs):
seen_auth_headers.append(mcp_auth_header)
tool = MagicMock()
tool = Tool(name="placeholder", inputSchema={})
tool.name = f"{server.alias}-toolA" if add_prefix else "toolA"
tool.description = "desc"
tool.input_schema = {}

View file

@ -289,7 +289,7 @@ def _catalog_case(method):
"method", ["prompts/list", "prompts/get", "resources/list", "resources/templates/list", "resources/read"]
)
@pytest.mark.parametrize("state", ["success", "denied", "upstream_failure", "scope_failure"])
async def test_native_catalog_operations_preserve_context_results_and_failure_policy(method, state):
async def test_catalog_helpers_preserve_context_results_and_failure_policy(method, state):
from types import SimpleNamespace
from fastapi import HTTPException
@ -325,10 +325,15 @@ async def test_native_catalog_operations_preserve_context_results_and_failure_po
with pytest.raises(expected_error):
await getattr(server, handler_name)(ctx, operation.params)
else:
result = await getattr(server, handler_name)(ctx, operation.params or PaginatedRequestParams())
if collection:
assert getattr(result, collection) == (payload if state == "success" else [])
helper = getattr(operations, "_list_mcp_" + collection)
result = await helper(
user_api_key_auth=caller, mcp_auth_header=None, mcp_servers=["catalog"],
mcp_server_auth_headers=None, oauth2_headers=None, raw_headers=headers, client_ip="192.0.2.41",
)
assert result == (payload if state == "success" else [])
else:
result = await getattr(server, handler_name)(ctx, operation.params or PaginatedRequestParams())
assert result == payload
assert allowed.await_args.kwargs == {
"user_api_key_auth": caller,
@ -948,3 +953,105 @@ async def test_local_handler_rejects_an_owner_absent_from_the_catalog(monkeypatc
await operations._handle_local_mcp_tool("private-export", {})
assert denied.value.status_code == 503
handler.assert_not_awaited()
@pytest.mark.asyncio
async def test_gateway_listing_rejects_unrecognized_continuation() -> None:
from mcp.shared.exceptions import MCPError
from mcp.types import ListToolsRequest, PaginatedRequestParams
with pytest.raises(MCPError, match=r"cursor|pagination"):
await GatewayOperations().execute(
ListToolsRequest(params=PaginatedRequestParams(cursor="forged-pagination-state")),
prepare_context(mcp_proxy_mode=True),
)
@pytest.mark.asyncio
@pytest.mark.parametrize("method", ["prompts/list", "resources/list", "resources/templates/list"])
async def test_continuation_preserves_current_authority_unavailable_error(monkeypatch, method):
from mcp import MCPError
from mcp.types import ListPromptsRequest, ListResourcesRequest, ListResourceTemplatesRequest, PaginatedRequestParams
from litellm.proxy import proxy_server
monkeypatch.setattr(proxy_server, "prisma_client", None)
caller = UserAPIKeyAuth(api_key="sk-owned-key-without-database")
caller.via_virtual_key = True
request = {"prompts/list": ListPromptsRequest, "resources/list": ListResourcesRequest, "resources/templates/list": ListResourceTemplatesRequest}[method]
with pytest.raises(MCPError, match="Server misconfigured: no database connection"):
await GatewayOperations().execute(request(params=PaginatedRequestParams(cursor="existing-state")), prepare_context(caller))
@pytest.mark.asyncio
async def test_virtual_tool_catalog_rejects_a_cursor_and_preserves_its_complete_listing():
from mcp import MCPError
from mcp.types import ListToolsRequest, PaginatedRequestParams
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
caller = UserAPIKeyAuth(object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="search", mcp_tool_search_enabled=True))
context = prepare_context(caller)
result = await GatewayOperations().execute(ListToolsRequest(), context)
assert result.tools
assert result.next_cursor is None
with pytest.raises(MCPError, match="fresh listing"):
await GatewayOperations().execute(ListToolsRequest(params=PaginatedRequestParams(cursor="existing-state")), context)
@pytest.mark.asyncio
@pytest.mark.parametrize("request_name", ["ListPromptsRequest", "ListResourcesRequest", "ListResourceTemplatesRequest"])
@pytest.mark.parametrize("cursor", [None, "existing-state"])
async def test_optional_catalog_preserves_revoked_user_error(monkeypatch, request_name, cursor):
from types import SimpleNamespace
from mcp import MCPError, types
from litellm.caching.dual_cache import DualCache
from litellm.proxy import proxy_server
monkeypatch.setattr(proxy_server, "general_settings", {"supported_db_objects": []})
table = SimpleNamespace(find_unique=AsyncMock(return_value=None))
monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(writer_db=SimpleNamespace(litellm_usertable=table)))
monkeypatch.setattr(proxy_server, "user_api_key_cache", DualCache())
caller = UserAPIKeyAuth(user_id="revoked-catalog-user")
caller.mcp_admitted_user_subject = True
request = getattr(types, request_name)(params=types.PaginatedRequestParams(cursor=cursor))
with pytest.raises(MCPError, match="Invalid or expired credential"):
await GatewayOperations().execute(request, prepare_context(caller))
@pytest.mark.asyncio
@pytest.mark.parametrize("cursor", ["continuation-state", ""])
async def test_tool_continuation_failure_requires_a_fresh_listing(monkeypatch: pytest.MonkeyPatch, cursor: str) -> None:
from mcp import MCPError
from mcp.types import INVALID_PARAMS, ListToolsRequest, PaginatedRequestParams
failure: Final = RuntimeError("catalog temporarily unavailable")
fetch: Final = AsyncMock(side_effect=failure)
monkeypatch.setattr(operations, "_get_tools_from_mcp_servers", fetch)
with pytest.raises(MCPError, match="start a fresh listing") as raised:
await GatewayOperations().execute(
ListToolsRequest(params=PaginatedRequestParams(cursor=cursor)), prepare_context()
)
assert raised.value.error.code == INVALID_PARAMS
assert raised.value.__cause__ is failure
assert fetch.await_count == 1
assert fetch.await_args.kwargs["params"].cursor == cursor
@pytest.mark.asyncio
@pytest.mark.parametrize("gateway", [False, True])
async def test_initial_tool_listing_preserves_legacy_error_fallback(monkeypatch: pytest.MonkeyPatch, gateway: bool) -> None:
from mcp.types import ListToolsRequest
fetch: Final = AsyncMock(side_effect=RuntimeError("catalog temporarily unavailable"))
monkeypatch.setattr(operations, "_get_tools_from_mcp_servers", fetch)
if gateway:
result: Final = await GatewayOperations().execute(ListToolsRequest(), prepare_context())
assert result.tools == []
assert result.next_cursor is None
else:
listing: Final = await operations._list_mcp_tools()
assert listing.tools == []
assert listing.next_cursor is None
fetch.assert_awaited_once()

View file

@ -553,3 +553,16 @@ async def test_target_catalog_does_not_reuse_admin_authorization_for_another_cal
)
assert (error.value.status_code, error.value.detail) == (403, {"error": "denied"})
manager.allowed_servers_spy.assert_called_once_with(_auth())
@pytest.mark.asyncio
async def test_catalog_without_listing_dependency_fails_explicitly():
from mcp.types import ListToolsRequest
from litellm.proxy._experimental.mcp_server.contracts import OperationContext
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
from litellm.proxy._experimental.mcp_server.server_resolution import MCPServerTargetCatalog
catalog = MCPServerTargetCatalog(MCPServerManager())
with pytest.raises(RuntimeError, match="listing dependency"):
await catalog.list(OperationContext(_caller=None), ListToolsRequest())

View file

@ -0,0 +1,90 @@
import base64
from typing import Final
import pytest
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error, Ok
from litellm.proxy._experimental.mcp_server.state_tokens import StateTokenError, open_state, seal_state
@pytest.fixture(autouse=True)
def state_key(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LITELLM_SALT_KEY", "synthetic-shared-pagination-key")
def test_state_is_portable_repeatable_and_preserves_json() -> None:
value: Final = {"caller": "user-a", "upstream": {"cursor": "opaque+/=", "offset": 3}, "revision": "r1"}
sealed: Final = seal_state(value, purpose="pagination", expires_at=200, now=100)
assert isinstance(sealed, Ok)
assert open_state(sealed.ok, purpose="pagination", now=150) == Ok(value)
assert open_state(sealed.ok, purpose="pagination", now=199) == Ok(value)
assert "user-a" not in sealed.ok
def test_sealing_the_same_state_uses_distinct_nonces() -> None:
first: Final = seal_state("same", purpose="pagination", expires_at=200, now=100)
second: Final = seal_state("same", purpose="pagination", expires_at=200, now=100)
assert isinstance(first, Ok) and isinstance(second, Ok)
assert first.ok != second.ok
assert open_state(first.ok, purpose="pagination", now=101) == Ok("same")
assert open_state(second.ok, purpose="pagination", now=101) == Ok("same")
@pytest.mark.parametrize("purpose", ("continuation", "pagination-other", ""))
def test_state_cannot_be_opened_for_another_purpose(purpose: str) -> None:
sealed: Final = seal_state("private", purpose="pagination", expires_at=200, now=100)
assert isinstance(sealed, Ok)
assert open_state(sealed.ok, purpose=purpose, now=100) == Error(StateTokenError.INVALID)
def test_rotating_the_key_invalidates_existing_state(monkeypatch: pytest.MonkeyPatch) -> None:
sealed: Final = seal_state("private", purpose="pagination", expires_at=200, now=100)
assert isinstance(sealed, Ok)
monkeypatch.setenv("LITELLM_SALT_KEY", "different-synthetic-key")
assert open_state(sealed.ok, purpose="pagination", now=100) == Error(StateTokenError.INVALID)
@pytest.mark.parametrize("missing", (True, False))
def test_master_key_cannot_replace_missing_or_empty_salt(monkeypatch: pytest.MonkeyPatch, missing: bool) -> None:
monkeypatch.setenv("LITELLM_MASTER_KEY", "synthetic-master-key")
if missing:
monkeypatch.delenv("LITELLM_SALT_KEY")
else:
monkeypatch.setenv("LITELLM_SALT_KEY", "")
assert seal_state("private", purpose="pagination", expires_at=200, now=100) == Error(StateTokenError.MISSING_KEY)
assert open_state("forged", purpose="pagination", now=100) == Error(StateTokenError.MISSING_KEY)
@pytest.mark.parametrize("now", (200, 201))
def test_state_expires_at_the_deadline(now: int) -> None:
sealed: Final = seal_state("private", purpose="pagination", expires_at=200, now=100)
assert isinstance(sealed, Ok)
assert open_state(sealed.ok, purpose="pagination", now=now) == Error(StateTokenError.EXPIRED)
assert seal_state("private", purpose="pagination", expires_at=200, now=now) == Error(StateTokenError.EXPIRED)
def test_altered_ciphertext_is_rejected() -> None:
sealed: Final = seal_state("private", purpose="pagination", expires_at=200, now=100)
assert isinstance(sealed, Ok)
prefix, encoded = sealed.ok.split(".", 1)
raw: Final = base64.urlsafe_b64decode(encoded + "=" * (-len(encoded) % 4))
altered: Final = bytes((raw[0] ^ 1,)) + raw[1:]
token: Final = prefix + "." + base64.urlsafe_b64encode(altered).decode("ascii").rstrip("=")
assert open_state(token, purpose="pagination", now=100) == Error(StateTokenError.INVALID)
@pytest.mark.parametrize(
"token", ("", "forged", "mcp_state_v2.abc", "mcp_state_v1.!", "mcp_state_v1.YQ", "mcp_state_v1.é")
)
def test_malformed_state_is_rejected(token: str) -> None:
assert open_state(token, purpose="pagination", now=100) == Error(StateTokenError.INVALID)
def test_excessive_state_is_rejected() -> None:
assert seal_state("x" * 65536, purpose="pagination", expires_at=200, now=100) == Error(StateTokenError.TOO_LARGE)
assert seal_state("x" * 50000, purpose="pagination", expires_at=200, now=100) == Error(StateTokenError.TOO_LARGE)
assert open_state("x" * 65537, purpose="pagination", now=100) == Error(StateTokenError.TOO_LARGE)
def test_empty_purpose_cannot_mint_state() -> None:
assert seal_state("private", purpose="", expires_at=200, now=100) == Error(StateTokenError.INVALID)