From d1cfe17518f69d5fbf56d6d26d392bba83ab52e4 Mon Sep 17 00:00:00 2001 From: joshua-berri Date: Tue, 6 Oct 2026 18:54:50 -0700 Subject: [PATCH] feat(mcp): add portable catalog pagination (#44446) * 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> --- litellm/experimental_mcp_client/client.py | 56 +- litellm/experimental_mcp_client/tools.py | 14 +- .../proxy/_experimental/mcp_server/README.md | 15 + .../mcp_server/auth/user_api_key_auth_mcp.py | 30 + .../proxy/_experimental/mcp_server/catalog.py | 646 +++++++++++++++++- .../_experimental/mcp_server/contracts.py | 26 +- .../mcp_server/faults/list_outcomes.py | 1 + .../mcp_server/mcp_server_manager.py | 158 ++++- .../_experimental/mcp_server/operations.py | 397 ++++------- .../proxy/_experimental/mcp_server/server.py | 6 + .../mcp_server/server_resolution.py | 11 +- .../_experimental/mcp_server/state_tokens.py | 87 +++ tests/integration/_support/mcp.py | 113 ++- tests/integration/mcp/test_pagination.py | 563 +++++++++++++++ tests/mcp_tests/test_proxy_mcp_e2e.py | 4 +- .../test_mcp_client.py | 103 +++ .../experimental_mcp_client/test_tools.py | 36 +- .../auth/test_user_api_key_auth_mcp.py | 91 ++- .../_experimental/mcp_server/conftest.py | 4 +- .../_experimental/mcp_server/test_catalog.py | 376 ++++++++++ .../mcp_server/test_discoverable_endpoints.py | 4 +- .../mcp_server/test_mcp_server_manager.py | 186 +++++ .../test_mcp_server_tool_calls_and_headers.py | 63 +- .../mcp_server/test_operations.py | 113 ++- .../mcp_server/test_server_resolution.py | 13 + .../mcp_server/test_state_tokens.py | 90 +++ 26 files changed, 2861 insertions(+), 345 deletions(-) create mode 100644 litellm/proxy/_experimental/mcp_server/README.md create mode 100644 litellm/proxy/_experimental/mcp_server/state_tokens.py create mode 100644 tests/integration/mcp/test_pagination.py create mode 100644 tests/unit/proxy/_experimental/mcp_server/test_catalog.py create mode 100644 tests/unit/proxy/_experimental/mcp_server/test_state_tokens.py diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index c3e4881d689..8f78253eaca 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -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) diff --git a/litellm/experimental_mcp_client/tools.py b/litellm/experimental_mcp_client/tools.py index a73a12b9e03..8015b065899 100644 --- a/litellm/experimental_mcp_client/tools.py +++ b/litellm/experimental_mcp_client/tools.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/README.md b/litellm/proxy/_experimental/mcp_server/README.md new file mode 100644 index 00000000000..1123a9779ac --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/README.md @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 19e7f716f81..418bdcd2941 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/catalog.py b/litellm/proxy/_experimental/mcp_server/catalog.py index 1f3cd77cdfa..a957b352705 100644 --- a/litellm/proxy/_experimental/mcp_server/catalog.py +++ b/litellm/proxy/_experimental/mcp_server/catalog.py @@ -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, + ) diff --git a/litellm/proxy/_experimental/mcp_server/contracts.py b/litellm/proxy/_experimental/mcp_server/contracts.py index 82b1861c1e8..f3800c93a77 100644 --- a/litellm/proxy/_experimental/mcp_server/contracts.py +++ b/litellm/proxy/_experimental/mcp_server/contracts.py @@ -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: diff --git a/litellm/proxy/_experimental/mcp_server/faults/list_outcomes.py b/litellm/proxy/_experimental/mcp_server/faults/list_outcomes.py index 06bf56f6f96..0c1f7599718 100644 --- a/litellm/proxy/_experimental/mcp_server/faults/list_outcomes.py +++ b/litellm/proxy/_experimental/mcp_server/faults/list_outcomes.py @@ -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]: diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index e301aa53e50..cb0eed8823c 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index fdea39bf5bc..2164dd332ac 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -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=[]) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 42c2e55a1ee..c046fa93d54 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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=[]) diff --git a/litellm/proxy/_experimental/mcp_server/server_resolution.py b/litellm/proxy/_experimental/mcp_server/server_resolution.py index 54fd17fa280..e828d8d8ee6 100644 --- a/litellm/proxy/_experimental/mcp_server/server_resolution.py +++ b/litellm/proxy/_experimental/mcp_server/server_resolution.py @@ -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, diff --git a/litellm/proxy/_experimental/mcp_server/state_tokens.py b/litellm/proxy/_experimental/mcp_server/state_tokens.py new file mode 100644 index 00000000000..88f98bef6b8 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/state_tokens.py @@ -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) diff --git a/tests/integration/_support/mcp.py b/tests/integration/_support/mcp.py index d39d7e4029e..5991c35c140 100644 --- a/tests/integration/_support/mcp.py +++ b/tests/integration/_support/mcp.py @@ -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) diff --git a/tests/integration/mcp/test_pagination.py b/tests/integration/mcp/test_pagination.py new file mode 100644 index 00000000000..84689232038 --- /dev/null +++ b/tests/integration/mcp/test_pagination.py @@ -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() == () diff --git a/tests/mcp_tests/test_proxy_mcp_e2e.py b/tests/mcp_tests/test_proxy_mcp_e2e.py index 21d6633f9f2..f56b68922c5 100644 --- a/tests/mcp_tests/test_proxy_mcp_e2e.py +++ b/tests/mcp_tests/test_proxy_mcp_e2e.py @@ -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": diff --git a/tests/unit/experimental_mcp_client/test_mcp_client.py b/tests/unit/experimental_mcp_client/test_mcp_client.py index 83b669e68f0..6c488d770c3 100644 --- a/tests/unit/experimental_mcp_client/test_mcp_client.py +++ b/tests/unit/experimental_mcp_client/test_mcp_client.py @@ -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 diff --git a/tests/unit/experimental_mcp_client/test_tools.py b/tests/unit/experimental_mcp_client/test_tools.py index 6645b06664d..ac7afe4c74b 100644 --- a/tests/unit/experimental_mcp_client/test_tools.py +++ b/tests/unit/experimental_mcp_client/test_tools.py @@ -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] diff --git a/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 1a4a01a0141..8e439cc822f 100644 --- a/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -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"] diff --git a/tests/unit/proxy/_experimental/mcp_server/conftest.py b/tests/unit/proxy/_experimental/mcp_server/conftest.py index 90bb783ad4e..de7ed4d45e8 100644 --- a/tests/unit/proxy/_experimental/mcp_server/conftest.py +++ b/tests/unit/proxy/_experimental/mcp_server/conftest.py @@ -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) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_catalog.py b/tests/unit/proxy/_experimental/mcp_server/test_catalog.py new file mode 100644 index 00000000000..5365bda516a --- /dev/null +++ b/tests/unit/proxy/_experimental/mcp_server/test_catalog.py @@ -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" diff --git a/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 195c1ae2c74..8a0352e82a1 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -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() diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 4359762ba33..417c6cad3b1 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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 diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py index 189424a9f90..59e21404301 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py @@ -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 = {} diff --git a/tests/unit/proxy/_experimental/mcp_server/test_operations.py b/tests/unit/proxy/_experimental/mcp_server/test_operations.py index e3dad9c513e..8274b01ceeb 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_operations.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_operations.py @@ -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() diff --git a/tests/unit/proxy/_experimental/mcp_server/test_server_resolution.py b/tests/unit/proxy/_experimental/mcp_server/test_server_resolution.py index 853118b8dc2..51b97f7c8c5 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_server_resolution.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_server_resolution.py @@ -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()) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_state_tokens.py b/tests/unit/proxy/_experimental/mcp_server/test_state_tokens.py new file mode 100644 index 00000000000..b3fc5a89526 --- /dev/null +++ b/tests/unit/proxy/_experimental/mcp_server/test_state_tokens.py @@ -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)