mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
feat(mcp): add portable catalog pagination (#44446)
Some checks are pending
CI Coverage / assert-ci-coverage (push) Waiting to run
CodeQL / Analyze (actions) (push) Waiting to run
CodeQL / Analyze (javascript-typescript) (push) Waiting to run
CodeQL / Analyze (python) (push) Waiting to run
CodSpeed Benchmarks / benchmarks (push) Waiting to run
Helm unit test / unit-test (push) Waiting to run
Lens Worker Image / lens-worker-image (push) Waiting to run
Publish basedpyright base counts / publish (push) Waiting to run
Scorecard supply-chain security / Scorecard analysis (push) Waiting to run
Code Quality Checks / code-quality (push) Waiting to run
Code Quality Checks / python-310-import-smoke (push) Waiting to run
UI Unit Tests / ui-unit-tests (push) Waiting to run
LiteLLM Rust / rust-wheel (push) Waiting to run
LiteLLM Rust / rust-lint (push) Waiting to run
LiteLLM Rust / rust-test (push) Waiting to run
Terraform Provider / gofmt, vet, build, test (push) Waiting to run
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Waiting to run
Unit Tests: Documentation Validation / documentation (push) Waiting to run
Unit Tests: Proxy DB Operations / Lens Python 3.10 (push) Waiting to run
Unit Tests: Proxy DB Operations / assert-shard-coverage (push) Waiting to run
Unit Tests: Proxy DB Operations / auth-checks (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / budgets (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / custom-logging (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / db-and-spend (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / endpoints-and-responses (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / guardrails-hooks (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / jwt-and-keys (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / key-generation (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / logging-misc (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-runtime (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-server-core (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-utils (push) Blocked by required conditions
Unit Tests / Vertex AI (push) Blocked by required conditions
Unit Tests / misc (push) Blocked by required conditions
Unit Tests / Build the Rust bridge (push) Waiting to run
Unit Tests / caching-local (push) Blocked by required conditions
Unit Tests / core-utils (push) Blocked by required conditions
Unit Tests / enterprise-managed-files (push) Blocked by required conditions
Unit Tests / enterprise-package (push) Blocked by required conditions
Unit Tests / enterprise-routing (push) Blocked by required conditions
Unit Tests / integrations (push) Blocked by required conditions
Unit Tests / OpenAI and Meta Providers (push) Blocked by required conditions
Unit Tests / All Other Providers (push) Blocked by required conditions
Unit Tests / misc-dirs (push) Blocked by required conditions
Unit Tests / proxy-auth (push) Blocked by required conditions
Unit Tests / proxy-endpoints (push) Blocked by required conditions
Unit Tests / unit (push) Blocked by required conditions
GitHub Actions Security Analysis / zizmor (push) Waiting to run
Unit Tests / proxy-extras (push) Blocked by required conditions
Unit Tests / proxy-feature-endpoints (push) Blocked by required conditions
Unit Tests / proxy-hooks-client (push) Blocked by required conditions
Unit Tests / proxy-infra (push) Blocked by required conditions
Unit Tests / proxy-infra-root (push) Blocked by required conditions
Unit Tests / proxy-server (push) Blocked by required conditions
Unit Tests / responses-caching-types (push) Blocked by required conditions
Some checks are pending
CI Coverage / assert-ci-coverage (push) Waiting to run
CodeQL / Analyze (actions) (push) Waiting to run
CodeQL / Analyze (javascript-typescript) (push) Waiting to run
CodeQL / Analyze (python) (push) Waiting to run
CodSpeed Benchmarks / benchmarks (push) Waiting to run
Helm unit test / unit-test (push) Waiting to run
Lens Worker Image / lens-worker-image (push) Waiting to run
Publish basedpyright base counts / publish (push) Waiting to run
Scorecard supply-chain security / Scorecard analysis (push) Waiting to run
Code Quality Checks / code-quality (push) Waiting to run
Code Quality Checks / python-310-import-smoke (push) Waiting to run
UI Unit Tests / ui-unit-tests (push) Waiting to run
LiteLLM Rust / rust-wheel (push) Waiting to run
LiteLLM Rust / rust-lint (push) Waiting to run
LiteLLM Rust / rust-test (push) Waiting to run
Terraform Provider / gofmt, vet, build, test (push) Waiting to run
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Waiting to run
Unit Tests: Documentation Validation / documentation (push) Waiting to run
Unit Tests: Proxy DB Operations / Lens Python 3.10 (push) Waiting to run
Unit Tests: Proxy DB Operations / assert-shard-coverage (push) Waiting to run
Unit Tests: Proxy DB Operations / auth-checks (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / budgets (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / custom-logging (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / db-and-spend (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / endpoints-and-responses (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / guardrails-hooks (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / jwt-and-keys (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / key-generation (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / logging-misc (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-runtime (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-server-core (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-utils (push) Blocked by required conditions
Unit Tests / Vertex AI (push) Blocked by required conditions
Unit Tests / misc (push) Blocked by required conditions
Unit Tests / Build the Rust bridge (push) Waiting to run
Unit Tests / caching-local (push) Blocked by required conditions
Unit Tests / core-utils (push) Blocked by required conditions
Unit Tests / enterprise-managed-files (push) Blocked by required conditions
Unit Tests / enterprise-package (push) Blocked by required conditions
Unit Tests / enterprise-routing (push) Blocked by required conditions
Unit Tests / integrations (push) Blocked by required conditions
Unit Tests / OpenAI and Meta Providers (push) Blocked by required conditions
Unit Tests / All Other Providers (push) Blocked by required conditions
Unit Tests / misc-dirs (push) Blocked by required conditions
Unit Tests / proxy-auth (push) Blocked by required conditions
Unit Tests / proxy-endpoints (push) Blocked by required conditions
Unit Tests / unit (push) Blocked by required conditions
GitHub Actions Security Analysis / zizmor (push) Waiting to run
Unit Tests / proxy-extras (push) Blocked by required conditions
Unit Tests / proxy-feature-endpoints (push) Blocked by required conditions
Unit Tests / proxy-hooks-client (push) Blocked by required conditions
Unit Tests / proxy-infra (push) Blocked by required conditions
Unit Tests / proxy-infra-root (push) Blocked by required conditions
Unit Tests / proxy-server (push) Blocked by required conditions
Unit Tests / responses-caching-types (push) Blocked by required conditions
* feat(mcp): add portable authorized catalog pagination * fix(mcp): preserve routing and isolation across catalog pages * fix(mcp): preserve bare calls after complete catalog pages * test(mcp): verify bare calls with cached database catalog revisions * test(mcp): cover refreshed authority and failed catalog continuations * refactor(mcp): narrow pagination interfaces and preserve listed metadata * fix(mcp): publish listed metadata only after successful pagination * fix(mcp): defer bare routes until aggregate listing succeeds * fix(mcp): refresh key access groups on catalog pages * fix(mcp): accept read-only optional catalog headers * fix(mcp): preserve consolidated listing metadata after restack * fix(mcp): propagate unexpected catalog continuation failures * fix(mcp): preserve key grants when a user record is absent --------- Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
This commit is contained in:
parent
4909bd9e8c
commit
d1cfe17518
26 changed files with 2861 additions and 345 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
15
litellm/proxy/_experimental/mcp_server/README.md
Normal file
15
litellm/proxy/_experimental/mcp_server/README.md
Normal file
|
|
@ -0,0 +1,15 @@
|
|||
# Gateway catalog pagination
|
||||
|
||||
Tools, prompts, resources, and resource templates retain upstream page boundaries. Follow the gateway's opaque `nextCursor` until it is absent. Each request checks the caller's current access and tool permissions. A cursor can be replayed and can be sent to another replica serving the same configuration
|
||||
|
||||
Set the same nonempty `LITELLM_SALT_KEY` on every replica. Pagination state uses authenticated encryption with purpose-separated HKDF keys derived only from this value. The master key is never a fallback. Complete single-page lists and direct tool calls work without a salt key; a listing that needs continuation returns an actionable configuration error instead of truncated results. Some SDKs automatically list tools when validating a tool-call response, so those SDK calls also require a salt when that listing is paginated
|
||||
|
||||
Cursors expire ten minutes after the first page. Continuations do not extend that deadline. Rotating the salt invalidates all outstanding cursors; clients must start a new listing. During a rolling key change, replicas with different keys cannot accept each other's cursors. Coordinate the change across the fleet
|
||||
|
||||
A changed registry, caller scope, or available upstream revision requires a fresh listing. When an upstream exposes a string or integer `_meta.revision`, subsequent pages must retain it. Otherwise consistency follows that upstream's own cursor guarantees. Repeated upstream cursors and the existing upstream page limit stop traversal with an explicit error
|
||||
|
||||
Listing failures retain per-server outcome metadata. An incomplete upstream catalog cannot establish a bare tool-name route. Complete initial pages retain legacy bare-name routing. Use the server-prefixed names returned by the gateway for paginated catalogs
|
||||
|
||||
For a scoped rollout, keep the previous source build serving the control pool and send only selected clients to a separate candidate pool. All candidate replicas must share configuration and salt. Verify page one on one candidate replica and continuation on another, plus a fresh listing and tool call on the control pool. Do not mirror tool calls between pools
|
||||
|
||||
For rollback, stop sending new requests to the candidate pool, drain its in-flight operations, and return selected clients to the control pool. Clients must discard candidate cursors and start a fresh listing when crossing versions; older gateways do not validate these cursors. Keep the registry-revision migration installed when rolling back pagination. Verify a fresh listing and a tool call after switching pools
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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=[])
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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=[])
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
87
litellm/proxy/_experimental/mcp_server/state_tokens.py
Normal file
87
litellm/proxy/_experimental/mcp_server/state_tokens.py
Normal file
|
|
@ -0,0 +1,87 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import binascii
|
||||
import os
|
||||
from enum import Enum
|
||||
from typing import Final
|
||||
|
||||
from cryptography.exceptions import InvalidTag
|
||||
from cryptography.hazmat.primitives import hashes
|
||||
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
|
||||
from cryptography.hazmat.primitives.kdf.hkdf import HKDF
|
||||
from pydantic import BaseModel, ConfigDict, JsonValue, ValidationError
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error, Ok, Result
|
||||
|
||||
_PREFIX: Final = "mcp_state_v1."
|
||||
_MAX_TOKEN_LENGTH: Final = 65536
|
||||
_NONCE_BYTES: Final = 12
|
||||
|
||||
|
||||
class StateTokenError(str, Enum):
|
||||
MISSING_KEY = "Set the same LITELLM_SALT_KEY on every replica to enable pagination"
|
||||
INVALID = "Invalid pagination state; start a fresh listing"
|
||||
EXPIRED = "Pagination state expired; start a fresh listing"
|
||||
TOO_LARGE = "Pagination state exceeds the supported size"
|
||||
|
||||
|
||||
class _Envelope(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="forbid", strict=True)
|
||||
|
||||
expires_at: int
|
||||
value: JsonValue
|
||||
|
||||
|
||||
def _cipher(purpose: str) -> Result[AESGCM, StateTokenError]:
|
||||
salt_key: Final = os.getenv("LITELLM_SALT_KEY")
|
||||
if not salt_key:
|
||||
return Error(StateTokenError.MISSING_KEY)
|
||||
if not purpose:
|
||||
return Error(StateTokenError.INVALID)
|
||||
key: Final = HKDF(
|
||||
algorithm=hashes.SHA256(),
|
||||
length=32,
|
||||
salt=b"litellm:mcp:state:v1",
|
||||
info=purpose.encode("utf-8"),
|
||||
).derive(salt_key.encode("utf-8"))
|
||||
return Ok(AESGCM(key))
|
||||
|
||||
|
||||
def seal_state(value: JsonValue, *, purpose: str, expires_at: int, now: int) -> Result[str, StateTokenError]:
|
||||
cipher: Final = _cipher(purpose)
|
||||
if isinstance(cipher, Error):
|
||||
return cipher
|
||||
if expires_at <= now:
|
||||
return Error(StateTokenError.EXPIRED)
|
||||
plaintext: Final = _Envelope(expires_at=expires_at, value=value).model_dump_json().encode("utf-8")
|
||||
if len(plaintext) > _MAX_TOKEN_LENGTH:
|
||||
return Error(StateTokenError.TOO_LARGE)
|
||||
nonce: Final = os.urandom(_NONCE_BYTES)
|
||||
ciphertext: Final = cipher.ok.encrypt(nonce, plaintext, (_PREFIX + purpose).encode("utf-8"))
|
||||
token: Final = _PREFIX + base64.urlsafe_b64encode(nonce + ciphertext).decode("ascii").rstrip("=")
|
||||
return Error(StateTokenError.TOO_LARGE) if len(token) > _MAX_TOKEN_LENGTH else Ok(token)
|
||||
|
||||
|
||||
def open_state(token: str, *, purpose: str, now: int) -> Result[JsonValue, StateTokenError]:
|
||||
cipher: Final = _cipher(purpose)
|
||||
if isinstance(cipher, Error):
|
||||
return cipher
|
||||
if len(token) > _MAX_TOKEN_LENGTH:
|
||||
return Error(StateTokenError.TOO_LARGE)
|
||||
if not token.startswith(_PREFIX):
|
||||
return Error(StateTokenError.INVALID)
|
||||
encoded: Final = token[len(_PREFIX) :]
|
||||
try:
|
||||
sealed: Final = base64.b64decode(encoded + "=" * (-len(encoded) % 4), altchars=b"-_", validate=True)
|
||||
if len(sealed) < _NONCE_BYTES + 16 or base64.urlsafe_b64encode(sealed).decode("ascii").rstrip("=") != encoded:
|
||||
return Error(StateTokenError.INVALID)
|
||||
plaintext: Final = cipher.ok.decrypt(
|
||||
sealed[:_NONCE_BYTES], sealed[_NONCE_BYTES:], (_PREFIX + purpose).encode("utf-8")
|
||||
)
|
||||
envelope: Final = _Envelope.model_validate_json(plaintext)
|
||||
except (binascii.Error, ValueError, InvalidTag, ValidationError):
|
||||
return Error(StateTokenError.INVALID)
|
||||
if envelope.expires_at <= now:
|
||||
return Error(StateTokenError.EXPIRED)
|
||||
return Ok(envelope.value)
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
563
tests/integration/mcp/test_pagination.py
Normal file
563
tests/integration/mcp/test_pagination.py
Normal file
|
|
@ -0,0 +1,563 @@
|
|||
import asyncio
|
||||
from contextlib import asynccontextmanager
|
||||
from pathlib import Path
|
||||
from typing import Literal
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import yaml
|
||||
from mcp import ClientSession, MCPError
|
||||
from mcp.client.streamable_http import streamable_http_client
|
||||
from mcp.types import CallToolRequest, CallToolRequestParams, CallToolResult, PaginatedRequestParams
|
||||
|
||||
from integration._support.client import Gateway
|
||||
from integration._support.mcp import paginated_mcp_peer
|
||||
from integration._support.process import owned_proxy
|
||||
from litellm.experimental_mcp_client.client import MCPClient
|
||||
from litellm.types.mcp import MCPTransport
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def catalog_session(gateway):
|
||||
async with httpx.AsyncClient(headers={"Authorization": "Bearer " + gateway.key}) as client:
|
||||
async with streamable_http_client(
|
||||
str(gateway.client.base_url).rstrip("/") + "/mcp/", http_client=client
|
||||
) as streams:
|
||||
async with ClientSession(streams[0], streams[1]) as session:
|
||||
await session.initialize()
|
||||
yield session
|
||||
|
||||
|
||||
def config_file(directory: Path, upstream: str) -> Path:
|
||||
config = directory / "proxy.yaml"
|
||||
config.write_text(
|
||||
yaml.safe_dump(
|
||||
{
|
||||
"model_list": [],
|
||||
"mcp_servers": {"pages": {"url": upstream, "transport": "http"}},
|
||||
"general_settings": {"master_key": "sk-pagination-test", "store_model_in_db": False},
|
||||
}
|
||||
)
|
||||
)
|
||||
return config
|
||||
|
||||
|
||||
REMOVE_DATABASE = ("DATABASE_URL", "DATABASE_URL_READ_REPLICA", "LITELLM_LICENSE", "LITELLM_LICENSE_PATH")
|
||||
|
||||
|
||||
def test_catalog_pages_are_portable_repeatable_and_complete(tmp_path: Path):
|
||||
async def exercise(first_replica, second_replica):
|
||||
for method, field in (
|
||||
("list_tools", "tools"),
|
||||
("list_prompts", "prompts"),
|
||||
("list_resources", "resources"),
|
||||
("list_resource_templates", "resource_templates"),
|
||||
):
|
||||
async with catalog_session(first_replica) as first_session:
|
||||
first = await getattr(first_session, method)()
|
||||
assert first.next_cursor
|
||||
async with catalog_session(second_replica) as second_session:
|
||||
second = await getattr(second_session, method)(params=PaginatedRequestParams(cursor=first.next_cursor))
|
||||
replay = await getattr(second_session, method)(params=PaginatedRequestParams(cursor=first.next_cursor))
|
||||
assert getattr(second, field) == getattr(replay, field)
|
||||
assert second.next_cursor
|
||||
third = await getattr(second_session, method)(params=PaginatedRequestParams(cursor=second.next_cursor))
|
||||
names = [item.name for page in (first, second, third) for item in getattr(page, field)]
|
||||
assert len(names) == len(set(names)) == 3
|
||||
assert third.next_cursor is None
|
||||
async with catalog_session(second_replica) as session:
|
||||
called = await session.call_tool("pages-add2", {"a": 3, "b": 4})
|
||||
assert not called.is_error
|
||||
assert called.content[0].text == "7"
|
||||
|
||||
with paginated_mcp_peer() as peer, httpx.Client() as client:
|
||||
seed = Gateway(client, "sk-pagination-test", peer.url)
|
||||
config = config_file(tmp_path, peer.url)
|
||||
options = dict(config=config, database_setup=(), remove_environment=REMOVE_DATABASE)
|
||||
environment = {
|
||||
"STORE_MODEL_IN_DB": "False",
|
||||
"DISABLE_SCHEMA_UPDATE": "true",
|
||||
"LITELLM_SALT_KEY": "shared-pagination-test",
|
||||
}
|
||||
with (
|
||||
owned_proxy(seed, tmp_path / "a", environment, **options) as a,
|
||||
owned_proxy(seed, tmp_path / "b", environment, **options) as b,
|
||||
):
|
||||
asyncio.run(exercise(a, b))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("page_size", [1, 3])
|
||||
def test_no_salt_preserves_complete_lists_and_calls_but_rejects_continuations(tmp_path: Path, page_size: int):
|
||||
async def exercise(gateway):
|
||||
async with catalog_session(gateway) as session:
|
||||
for method, field in (
|
||||
("list_tools", "tools"),
|
||||
("list_prompts", "prompts"),
|
||||
("list_resources", "resources"),
|
||||
("list_resource_templates", "resource_templates"),
|
||||
):
|
||||
if page_size == 1:
|
||||
with pytest.raises(MCPError, match="LITELLM_SALT_KEY"):
|
||||
await getattr(session, method)()
|
||||
else:
|
||||
result = await getattr(session, method)()
|
||||
assert len(getattr(result, field)) == 3
|
||||
assert result.next_cursor is None
|
||||
called = await session.send_request(
|
||||
CallToolRequest(params=CallToolRequestParams(name="pages-add2", arguments={"a": 3, "b": 4})),
|
||||
CallToolResult,
|
||||
)
|
||||
assert not called.is_error
|
||||
assert called.content[0].text == "7"
|
||||
if page_size == 3:
|
||||
bare = await session.send_request(
|
||||
CallToolRequest(params=CallToolRequestParams(name="add2", arguments={"a": 3, "b": 4})),
|
||||
CallToolResult,
|
||||
)
|
||||
assert not bare.is_error
|
||||
assert bare.content[0].text == "7"
|
||||
|
||||
with paginated_mcp_peer(page_size=page_size) as peer, httpx.Client() as client:
|
||||
seed = Gateway(client, "sk-pagination-test", peer.url)
|
||||
config = config_file(tmp_path, peer.url)
|
||||
environment = {"STORE_MODEL_IN_DB": "False", "DISABLE_SCHEMA_UPDATE": "true", "LITELLM_SALT_KEY": ""}
|
||||
with owned_proxy(
|
||||
seed, tmp_path / "proxy", environment, config=config, database_setup=(), remove_environment=REMOVE_DATABASE
|
||||
) as gateway:
|
||||
asyncio.run(exercise(gateway))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("grant", ["direct", "access_group"])
|
||||
def test_continuations_reauthorize_and_reject_registry_changes(tmp_path: Path, monkeypatch, grant: str):
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
from integration._support.database import scratch_database
|
||||
from integration._support.mcp import register_mcp
|
||||
|
||||
assert os.environ.get("DATABASE_URL"), "This integration case requires disposable-database access"
|
||||
|
||||
async def exercise(a, b, peer, identity, owner, stranger, policy):
|
||||
owner_a = Gateway(a.client, owner, peer.url)
|
||||
owner_b = Gateway(b.client, owner, peer.url)
|
||||
stranger_b = Gateway(b.client, stranger, peer.url)
|
||||
first_pages = {}
|
||||
for method in ("list_tools", "list_prompts", "list_resources", "list_resource_templates"):
|
||||
async with catalog_session(owner_a) as session:
|
||||
first_pages[method] = await getattr(session, method)()
|
||||
assert first_pages[method].next_cursor
|
||||
async with catalog_session(owner_b) as session:
|
||||
continued = await getattr(session, method)(
|
||||
params=PaginatedRequestParams(cursor=first_pages[method].next_cursor)
|
||||
)
|
||||
assert continued.next_cursor
|
||||
async with catalog_session(stranger_b) as session:
|
||||
peer.drain()
|
||||
with pytest.raises(MCPError, match="fresh listing"):
|
||||
await getattr(session, method)(
|
||||
params=PaginatedRequestParams(cursor=first_pages[method].next_cursor)
|
||||
)
|
||||
assert not any(call["body"].get("method", "").endswith("/list") for call in peer.drain())
|
||||
|
||||
a.post(
|
||||
"/key/update",
|
||||
{
|
||||
"key": owner,
|
||||
**(
|
||||
{"access_group_ids": []}
|
||||
if grant == "access_group"
|
||||
else {"object_permission": {"mcp_servers": ["no-mcp-servers"]}}
|
||||
),
|
||||
},
|
||||
)
|
||||
for method, first in first_pages.items():
|
||||
async with catalog_session(owner_b) as session:
|
||||
peer.drain()
|
||||
with pytest.raises(MCPError, match="fresh listing"):
|
||||
await getattr(session, method)(params=PaginatedRequestParams(cursor=first.next_cursor))
|
||||
assert not any(call["body"].get("method", "").endswith("/list") for call in peer.drain())
|
||||
a.post("/key/update", {"key": owner, **policy})
|
||||
changed = a.request("PUT", "/v1/mcp/server", {"server_id": identity, "description": "new catalog generation"})
|
||||
assert changed.status_code == 202, changed.text
|
||||
for method, first in first_pages.items():
|
||||
async with catalog_session(owner_b) as session:
|
||||
peer.drain()
|
||||
with pytest.raises(MCPError, match="fresh listing"):
|
||||
await getattr(session, method)(params=PaginatedRequestParams(cursor=first.next_cursor))
|
||||
assert not any(call["body"].get("method", "").endswith("/list") for call in peer.drain())
|
||||
fresh = await getattr(session, method)()
|
||||
assert fresh.next_cursor
|
||||
|
||||
with scratch_database() as database_url:
|
||||
monkeypatch.setenv("DATABASE_URL", database_url)
|
||||
subprocess.run(
|
||||
[
|
||||
sys.executable,
|
||||
"-I",
|
||||
"-m",
|
||||
"prisma",
|
||||
"db",
|
||||
"push",
|
||||
"--schema",
|
||||
"litellm/proxy/schema.prisma",
|
||||
"--skip-generate",
|
||||
],
|
||||
check=True,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
with paginated_mcp_peer() as peer, httpx.Client() as client:
|
||||
seed = Gateway(client, "sk-pagination-test", peer.url)
|
||||
config = tmp_path / "database-proxy.yaml"
|
||||
config.write_text(
|
||||
yaml.safe_dump(
|
||||
{"model_list": [], "general_settings": {"master_key": seed.key, "store_model_in_db": True}}
|
||||
)
|
||||
)
|
||||
environment = {
|
||||
"DATABASE_URL": database_url,
|
||||
"DISABLE_SCHEMA_UPDATE": "true",
|
||||
"LITELLM_SALT_KEY": "shared-pagination-test",
|
||||
}
|
||||
options = dict(
|
||||
config=config,
|
||||
database_setup=(),
|
||||
remove_environment=("DATABASE_URL_READ_REPLICA", "LITELLM_LICENSE", "LITELLM_LICENSE_PATH"),
|
||||
)
|
||||
with (
|
||||
owned_proxy(seed, tmp_path / "a", environment, **options) as a,
|
||||
owned_proxy(seed, tmp_path / "b", environment, **options) as b,
|
||||
a.scenario() as scenario,
|
||||
):
|
||||
identity = register_mcp(scenario, peer, "pages")
|
||||
group = a.request(
|
||||
"POST",
|
||||
"/v1/access_group",
|
||||
{
|
||||
"access_group_name": "pagination-grant",
|
||||
"access_mcp_server_ids": [identity],
|
||||
},
|
||||
)
|
||||
assert group.status_code == 201, group.text
|
||||
policy = (
|
||||
{"access_group_ids": [group.json()["access_group_id"]]}
|
||||
if grant == "access_group"
|
||||
else {"object_permission": {"mcp_servers": [identity]}}
|
||||
)
|
||||
owner = scenario.key(**policy)
|
||||
stranger = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
assert owner != stranger
|
||||
asyncio.run(exercise(a, b, peer, identity, owner, stranger, policy))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("changed", ["key", "snapshot"])
|
||||
def test_cursor_rejects_changed_replica_configuration_before_dispatch(tmp_path: Path, changed: str):
|
||||
async def exercise(a, b, peer):
|
||||
for method in ("list_tools", "list_prompts", "list_resources", "list_resource_templates"):
|
||||
async with catalog_session(a) as session:
|
||||
first = await getattr(session, method)()
|
||||
assert first.next_cursor
|
||||
async with catalog_session(b) as session:
|
||||
peer.drain()
|
||||
with pytest.raises(MCPError, match="fresh listing"):
|
||||
await getattr(session, method)(params=PaginatedRequestParams(cursor=first.next_cursor))
|
||||
assert not any(call["body"].get("method", "").endswith("/list") for call in peer.drain())
|
||||
fresh = await getattr(session, method)()
|
||||
assert fresh.next_cursor
|
||||
|
||||
with paginated_mcp_peer() as peer, httpx.Client() as client:
|
||||
seed = Gateway(client, "sk-pagination-test", peer.url)
|
||||
config = config_file(tmp_path, peer.url)
|
||||
second_config = tmp_path / "second-proxy.yaml"
|
||||
second_values = yaml.safe_load(config.read_text())
|
||||
if changed == "snapshot":
|
||||
second_values["mcp_servers"]["pages"]["description"] = "changed registry definition"
|
||||
second_config.write_text(yaml.safe_dump(second_values))
|
||||
environment = {
|
||||
"STORE_MODEL_IN_DB": "False",
|
||||
"DISABLE_SCHEMA_UPDATE": "true",
|
||||
"LITELLM_SALT_KEY": "original-key",
|
||||
}
|
||||
second_environment = {**environment, "LITELLM_SALT_KEY": "rotated-key" if changed == "key" else "original-key"}
|
||||
options = dict(database_setup=(), remove_environment=REMOVE_DATABASE)
|
||||
with (
|
||||
owned_proxy(seed, tmp_path / "a", environment, config=config, **options) as a,
|
||||
owned_proxy(seed, tmp_path / "b", second_environment, config=second_config, **options) as b,
|
||||
):
|
||||
asyncio.run(exercise(a, b, peer))
|
||||
|
||||
|
||||
def test_pages_preserve_supported_protocol_versions(tmp_path: Path):
|
||||
from mcp.types import ListPromptsRequest, ListResourcesRequest, ListResourceTemplatesRequest, ListToolsRequest
|
||||
from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS
|
||||
|
||||
async def exercise(gateway):
|
||||
for revision in HANDSHAKE_PROTOCOL_VERSIONS:
|
||||
client = MCPClient(
|
||||
server_url=str(gateway.client.base_url).rstrip("/") + "/mcp",
|
||||
transport_type=MCPTransport.http,
|
||||
protocol_version=revision,
|
||||
extra_headers={"Authorization": "Bearer " + gateway.key},
|
||||
timeout=15,
|
||||
)
|
||||
for request, field in (
|
||||
(ListToolsRequest, "tools"),
|
||||
(ListPromptsRequest, "prompts"),
|
||||
(ListResourcesRequest, "resources"),
|
||||
(ListResourceTemplatesRequest, "resource_templates"),
|
||||
):
|
||||
first = await client.list_page(request())
|
||||
assert first.next_cursor, revision
|
||||
second = await client.list_page(request(params=PaginatedRequestParams(cursor=first.next_cursor)))
|
||||
assert second.next_cursor, revision
|
||||
third = await client.list_page(request(params=PaginatedRequestParams(cursor=second.next_cursor)))
|
||||
assert third.next_cursor is None, revision
|
||||
names = [item.name for page in (first, second, third) for item in getattr(page, field)]
|
||||
assert len(names) == len(set(names)) == 3, revision
|
||||
|
||||
with paginated_mcp_peer() as peer, httpx.Client() as client:
|
||||
seed = Gateway(client, "sk-pagination-test", peer.url)
|
||||
config = config_file(tmp_path, peer.url)
|
||||
environment = {"STORE_MODEL_IN_DB": "False", "DISABLE_SCHEMA_UPDATE": "true", "LITELLM_SALT_KEY": "shared-key"}
|
||||
with owned_proxy(
|
||||
seed, tmp_path / "proxy", environment, config=config, database_setup=(), remove_environment=REMOVE_DATABASE
|
||||
) as gateway:
|
||||
asyncio.run(exercise(gateway))
|
||||
|
||||
|
||||
def test_incomplete_discovery_never_establishes_a_bare_tool_route(tmp_path: Path):
|
||||
from integration._support.mcp import tool_calls
|
||||
|
||||
async def exercise(gateway, peer):
|
||||
async with catalog_session(gateway) as session:
|
||||
first = await session.list_tools()
|
||||
assert [tool.name for tool in first.tools] == ["pages-add0"]
|
||||
assert first.next_cursor
|
||||
with pytest.raises(MCPError, match="repeated"):
|
||||
await session.list_tools(params=PaginatedRequestParams(cursor=first.next_cursor))
|
||||
peer.drain()
|
||||
result = await session.send_request(
|
||||
CallToolRequest(params=CallToolRequestParams(name="add0", arguments={"a": 3, "b": 4})),
|
||||
CallToolResult,
|
||||
)
|
||||
assert result.is_error
|
||||
assert tool_calls(peer.drain()) == ()
|
||||
prefixed = await session.send_request(
|
||||
CallToolRequest(params=CallToolRequestParams(name="pages-add0", arguments={"a": 3, "b": 4})),
|
||||
CallToolResult,
|
||||
)
|
||||
assert not prefixed.is_error
|
||||
assert prefixed.content[0].text == "7"
|
||||
assert len(tool_calls(peer.drain())) == 1
|
||||
|
||||
with paginated_mcp_peer(repeat_cursor=True) as peer, httpx.Client() as client:
|
||||
seed = Gateway(client, "sk-pagination-test", peer.url)
|
||||
config = config_file(tmp_path, peer.url)
|
||||
environment = {"STORE_MODEL_IN_DB": "False", "DISABLE_SCHEMA_UPDATE": "true", "LITELLM_SALT_KEY": "shared-key"}
|
||||
with owned_proxy(
|
||||
seed, tmp_path / "proxy", environment, config=config, database_setup=(), remove_environment=REMOVE_DATABASE
|
||||
) as gateway:
|
||||
asyncio.run(exercise(gateway, peer))
|
||||
|
||||
|
||||
def test_partial_catalog_keeps_sanitized_failure_metadata_on_following_pages(tmp_path: Path):
|
||||
async def exercise(gateway):
|
||||
for method, field in (
|
||||
("list_tools", "tools"),
|
||||
("list_prompts", "prompts"),
|
||||
("list_resources", "resources"),
|
||||
("list_resource_templates", "resource_templates"),
|
||||
):
|
||||
async with catalog_session(gateway) as session:
|
||||
first = await getattr(session, method)()
|
||||
assert first.next_cursor
|
||||
assert len(getattr(first, field)) == 1
|
||||
fault = first.meta["litellm.ai/server_outcomes"]["broken"]
|
||||
assert fault["status"] != "ok"
|
||||
assert "tag" not in fault
|
||||
assert "untrusted upstream message" not in first.model_dump_json()
|
||||
second = await getattr(session, method)(params=PaginatedRequestParams(cursor=first.next_cursor))
|
||||
assert len(getattr(second, field)) == 1
|
||||
assert second.meta["litellm.ai/server_outcomes"]["broken"] == fault
|
||||
|
||||
with paginated_mcp_peer() as healthy, paginated_mcp_peer(fail_listing=True) as broken, httpx.Client() as client:
|
||||
seed = Gateway(client, "sk-pagination-test", healthy.url)
|
||||
config = config_file(tmp_path, healthy.url)
|
||||
values = yaml.safe_load(config.read_text())
|
||||
values["mcp_servers"]["broken"] = {"url": broken.url, "transport": "http"}
|
||||
config.write_text(yaml.safe_dump(values))
|
||||
environment = {"STORE_MODEL_IN_DB": "False", "DISABLE_SCHEMA_UPDATE": "true", "LITELLM_SALT_KEY": "shared-key"}
|
||||
with owned_proxy(
|
||||
seed, tmp_path / "proxy", environment, config=config, database_setup=(), remove_environment=REMOVE_DATABASE
|
||||
) as gateway:
|
||||
asyncio.run(exercise(gateway))
|
||||
|
||||
|
||||
def test_failed_upstream_continuation_requires_restart_for_every_catalog(tmp_path: Path):
|
||||
async def exercise(gateway):
|
||||
async with catalog_session(gateway) as session:
|
||||
for method in ("list_tools", "list_prompts", "list_resources", "list_resource_templates"):
|
||||
first = await getattr(session, method)()
|
||||
assert first.next_cursor
|
||||
with pytest.raises(MCPError, match="fresh listing") as caught:
|
||||
await getattr(session, method)(params=PaginatedRequestParams(cursor=first.next_cursor))
|
||||
assert "untrusted upstream message" not in str(caught.value)
|
||||
|
||||
with paginated_mcp_peer(fail_continuation=True) as peer, httpx.Client() as client:
|
||||
seed = Gateway(client, "sk-pagination-test", peer.url)
|
||||
config = config_file(tmp_path, peer.url)
|
||||
environment = {"STORE_MODEL_IN_DB": "False", "DISABLE_SCHEMA_UPDATE": "true", "LITELLM_SALT_KEY": "shared-key"}
|
||||
with owned_proxy(
|
||||
seed, tmp_path / "proxy", environment, config=config, database_setup=(), remove_environment=REMOVE_DATABASE
|
||||
) as gateway:
|
||||
asyncio.run(exercise(gateway))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("foreign", [{"foreign": {"status": "error", "message": "foreign outcome"}}, "malformed"])
|
||||
def test_optional_catalog_ignores_upstream_gateway_outcomes(tmp_path: Path, foreign):
|
||||
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import SERVER_OUTCOMES_META_KEY
|
||||
|
||||
async def exercise(gateway):
|
||||
async with catalog_session(gateway) as session:
|
||||
for method, field in (
|
||||
("list_prompts", "prompts"),
|
||||
("list_resources", "resources"),
|
||||
("list_resource_templates", "resource_templates"),
|
||||
):
|
||||
first = await getattr(session, method)()
|
||||
assert len(getattr(first, field)) == 1
|
||||
assert first.next_cursor
|
||||
second = await getattr(session, method)(params=PaginatedRequestParams(cursor=first.next_cursor))
|
||||
assert len(getattr(second, field)) == 1
|
||||
assert "foreign" not in (second.meta or {}).get(SERVER_OUTCOMES_META_KEY, {})
|
||||
|
||||
with paginated_mcp_peer(metadata={SERVER_OUTCOMES_META_KEY: foreign}) as peer, httpx.Client() as client:
|
||||
seed = Gateway(client, "sk-pagination-test", peer.url)
|
||||
config = config_file(tmp_path, peer.url)
|
||||
environment = {"STORE_MODEL_IN_DB": "False", "DISABLE_SCHEMA_UPDATE": "true", "LITELLM_SALT_KEY": "shared-key"}
|
||||
with owned_proxy(
|
||||
seed, tmp_path / "proxy", environment, config=config, database_setup=(), remove_environment=REMOVE_DATABASE
|
||||
) as gateway:
|
||||
asyncio.run(exercise(gateway))
|
||||
|
||||
|
||||
def test_complete_initial_page_keeps_bare_routes_with_a_cached_database_revision(tmp_path: Path, monkeypatch):
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
from integration._support.database import read_rows, scratch_database, write_rows
|
||||
from integration._support.mcp import register_mcp, tool_calls
|
||||
|
||||
assert os.environ.get("DATABASE_URL"), "This integration case requires disposable-database access"
|
||||
|
||||
async def exercise(gateway, first, second):
|
||||
async with catalog_session(gateway) as session:
|
||||
listing = await session.list_tools()
|
||||
assert len(listing.tools) == 3 and listing.next_cursor is None
|
||||
first.drain()
|
||||
second.drain()
|
||||
called = await session.send_request(
|
||||
CallToolRequest(params=CallToolRequestParams(name="add2", arguments={"a": 3, "b": 4})),
|
||||
CallToolResult,
|
||||
)
|
||||
assert not called.is_error
|
||||
assert called.content[0].text == "7"
|
||||
assert tool_calls(first.drain()) == ()
|
||||
assert len(tool_calls(second.drain())) == 1
|
||||
|
||||
with scratch_database() as database_url:
|
||||
monkeypatch.setenv("DATABASE_URL", database_url)
|
||||
subprocess.run(
|
||||
[sys.executable, "-I", "-m", "prisma", "db", "push", "--schema",
|
||||
"litellm/proxy/schema.prisma", "--skip-generate"],
|
||||
check=True, capture_output=True, text=True,
|
||||
)
|
||||
subprocess.run(
|
||||
[sys.executable, "-I", "-m", "prisma", "db", "execute", "--schema",
|
||||
"litellm/proxy/schema.prisma", "--file",
|
||||
"litellm-proxy-extras/litellm_proxy_extras/migrations/20260923000000_add_mcp_catalog_revision_trigger/migration.sql"],
|
||||
check=True, capture_output=True, text=True,
|
||||
)
|
||||
with paginated_mcp_peer(page_size=3) as first, paginated_mcp_peer(page_size=3) as second, httpx.Client() as client:
|
||||
seed = Gateway(client, "sk-pagination-test", first.url)
|
||||
config = tmp_path / "database-proxy.yaml"
|
||||
config.write_text(yaml.safe_dump({
|
||||
"model_list": [], "general_settings": {"master_key": seed.key, "store_model_in_db": True},
|
||||
}))
|
||||
environment = {"DATABASE_URL": database_url, "DISABLE_SCHEMA_UPDATE": "true", "LITELLM_SALT_KEY": "shared-pagination-test"}
|
||||
with owned_proxy(
|
||||
seed, tmp_path / "proxy", environment, config=config, database_setup=(),
|
||||
remove_environment=("DATABASE_URL_READ_REPLICA", "LITELLM_LICENSE", "LITELLM_LICENSE_PATH"),
|
||||
) as gateway, gateway.scenario() as scenario:
|
||||
first_id = register_mcp(scenario, first, "first")
|
||||
second_id = register_mcp(scenario, second, "second")
|
||||
owner = scenario.key(object_permission={"mcp_servers": [second_id]})
|
||||
warm = gateway.client.get("/mcp-rest/tools/list", headers={"Authorization": "Bearer " + gateway.key}, params={"server_id": first_id})
|
||||
assert warm.status_code == 200, warm.text
|
||||
# A no-op SQL writer bumps the trigger revision without changing either server.
|
||||
write_rows('UPDATE "LiteLLM_MCPServerTable" SET "alias" = "alias" WHERE "server_id" = %s', (first_id,))
|
||||
warm = gateway.client.get("/mcp-rest/tools/list", headers={"Authorization": "Bearer " + gateway.key}, params={"server_id": first_id})
|
||||
assert warm.status_code == 200, warm.text
|
||||
query = 'SELECT "reload_revision" FROM "LiteLLM_Config" WHERE "param_name" = %s'
|
||||
revision = read_rows(query, ("mcp_catalog",))
|
||||
assert revision and revision[0]["reload_revision"] > 0
|
||||
asyncio.run(exercise(Gateway(gateway.client, owner, second.url), first, second))
|
||||
assert read_rows(query, ("mcp_catalog",)) == revision
|
||||
|
||||
|
||||
@pytest.mark.parametrize("entry", ["mcp", "server_mcp"])
|
||||
def test_missing_user_keeps_explicit_key_and_team_grants(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch, entry: Literal["mcp", "server_mcp"]
|
||||
):
|
||||
import subprocess
|
||||
import sys
|
||||
import uuid
|
||||
|
||||
from integration._support.database import read_rows, scratch_database
|
||||
from integration._support.mcp import McpCaller, register_mcp, tool_calls
|
||||
|
||||
with scratch_database() as database_url:
|
||||
monkeypatch.setenv("DATABASE_URL", database_url)
|
||||
subprocess.run(
|
||||
[sys.executable, "-I", "-m", "prisma", "db", "push", "--schema",
|
||||
"litellm/proxy/schema.prisma", "--skip-generate"],
|
||||
check=True, capture_output=True, text=True,
|
||||
)
|
||||
with paginated_mcp_peer(page_size=3) as allowed, paginated_mcp_peer(page_size=3) as private, httpx.Client() as client:
|
||||
seed = Gateway(client, "sk-pagination-test", allowed.url)
|
||||
config = tmp_path / "database-proxy.yaml"
|
||||
config.write_text(yaml.safe_dump({
|
||||
"model_list": [], "general_settings": {"master_key": seed.key, "store_model_in_db": True},
|
||||
}))
|
||||
environment = {"DATABASE_URL": database_url, "DISABLE_SCHEMA_UPDATE": "true", "LITELLM_SALT_KEY": "shared-pagination-test"}
|
||||
options = dict(config=config, database_setup=(), remove_environment=("DATABASE_URL_READ_REPLICA", "LITELLM_LICENSE", "LITELLM_LICENSE_PATH"))
|
||||
with (
|
||||
owned_proxy(seed, tmp_path / "a", environment, **options) as a,
|
||||
owned_proxy(seed, tmp_path / "b", environment, **options) as b,
|
||||
a.scenario() as scenario,
|
||||
):
|
||||
identity = register_mcp(scenario, allowed, "allowed")
|
||||
register_mcp(scenario, private, "private")
|
||||
team = scenario.team(object_permission={"mcp_servers": [identity]})
|
||||
for policy in ({"object_permission": {"mcp_servers": [identity]}}, {"team_id": team}):
|
||||
user = "missing-" + uuid.uuid4().hex
|
||||
key = scenario.key(user_id=user, **policy)
|
||||
assert read_rows('SELECT user_id FROM "LiteLLM_UserTable" WHERE user_id = %s', (user,)) == []
|
||||
for replica in (a, b):
|
||||
caller = McpCaller(replica, key, entry, alias="allowed")
|
||||
allowed.drain()
|
||||
private.drain()
|
||||
listing = caller.list_tools()
|
||||
assert listing.ok, listing
|
||||
assert len(listing.tools) == 3, listing
|
||||
name = next(name for name in listing.tools if name.endswith("add2"))
|
||||
called = caller.call(name, {"a": 3, "b": 4})
|
||||
assert called.ok and called.text == "7", called
|
||||
assert len(tool_calls(allowed.drain())) == 1
|
||||
assert private.drain() == ()
|
||||
forbidden = caller.call("private-add2", {"a": 3, "b": 4})
|
||||
assert not forbidden.ok, forbidden
|
||||
assert tool_calls(allowed.drain()) == ()
|
||||
assert private.drain() == ()
|
||||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
376
tests/unit/proxy/_experimental/mcp_server/test_catalog.py
Normal file
376
tests/unit/proxy/_experimental/mcp_server/test_catalog.py
Normal file
|
|
@ -0,0 +1,376 @@
|
|||
import asyncio
|
||||
from collections.abc import Sequence
|
||||
|
||||
import pytest
|
||||
from mcp.shared.exceptions import MCPError
|
||||
from mcp.types import ListToolsResult, Tool
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import catalog
|
||||
|
||||
|
||||
def page(name: str, cursor: str | None = None, revision: str = "stable") -> ListToolsResult:
|
||||
return ListToolsResult(
|
||||
tools=[Tool(name=name, input_schema={"type": "object"})],
|
||||
next_cursor=cursor,
|
||||
meta={"revision": revision},
|
||||
)
|
||||
|
||||
|
||||
async def listing(
|
||||
fetch,
|
||||
cursor: str | None = None,
|
||||
*,
|
||||
caller_scope: str = "caller-and-scope",
|
||||
snapshot: str = "registry-generation",
|
||||
servers: Sequence[str] = ("a", "b"),
|
||||
now: int = 100,
|
||||
) -> ListToolsResult:
|
||||
return await catalog.list_tools_page(
|
||||
cursor=cursor,
|
||||
caller_scope=caller_scope,
|
||||
snapshot=snapshot,
|
||||
server_ids=tuple(servers),
|
||||
fetch=fetch,
|
||||
now=now,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_listing_continues_on_another_replica_and_cursor_is_reusable(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "shared-test-key")
|
||||
|
||||
async def replica_a(server_id: str, cursor: str | None) -> ListToolsResult:
|
||||
assert cursor is None
|
||||
return page(server_id + "1", "page-two" if server_id == "a" else None)
|
||||
|
||||
async def replica_b(server_id: str, cursor: str | None) -> ListToolsResult:
|
||||
assert (server_id, cursor) == ("a", "page-two")
|
||||
return page("a2")
|
||||
|
||||
first = await listing(replica_a, servers=("b", "a"))
|
||||
assert [tool.name for tool in first.tools] == ["a1", "b1"]
|
||||
assert first.next_cursor
|
||||
for _ in range(2):
|
||||
second = await listing(replica_b, first.next_cursor)
|
||||
assert [tool.name for tool in second.tools] == ["a2"]
|
||||
assert second.next_cursor is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"changes",
|
||||
[
|
||||
{"caller_scope": "other-caller"},
|
||||
{"snapshot": "new-registry-generation"},
|
||||
{"servers": ("b",)},
|
||||
{"now": 1000},
|
||||
],
|
||||
)
|
||||
async def test_continuation_rejects_changed_binding_before_upstream_dispatch(monkeypatch, changes):
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "shared-test-key")
|
||||
|
||||
async def first_page(server_id: str, cursor: str | None) -> ListToolsResult:
|
||||
return page(server_id, "next")
|
||||
|
||||
async def forbidden_dispatch(server_id: str, cursor: str | None) -> ListToolsResult:
|
||||
pytest.fail("Rejected continuation must not contact upstream")
|
||||
|
||||
first = await listing(first_page)
|
||||
with pytest.raises(MCPError, match=r"fresh listing|expired"):
|
||||
await listing(forbidden_dispatch, first.next_cursor, **changes)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_complete_single_page_needs_no_key_but_continuation_does(monkeypatch):
|
||||
monkeypatch.delenv("LITELLM_SALT_KEY", raising=False)
|
||||
|
||||
async def complete(server_id: str, cursor: str | None) -> ListToolsResult:
|
||||
return page(server_id)
|
||||
|
||||
async def incomplete(server_id: str, cursor: str | None) -> ListToolsResult:
|
||||
return page(server_id, "next")
|
||||
|
||||
result = await listing(complete)
|
||||
assert len(result.tools) == 2
|
||||
assert result.next_cursor is None
|
||||
with pytest.raises(MCPError, match="LITELLM_SALT_KEY"):
|
||||
await listing(incomplete)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_repeated_upstream_cursor_requires_restart(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "shared-test-key")
|
||||
|
||||
async def repeated(server_id: str, cursor: str | None) -> ListToolsResult:
|
||||
return page(server_id, "repeated")
|
||||
|
||||
first = await listing(repeated)
|
||||
with pytest.raises(MCPError, match=r"repeated.*cursor"):
|
||||
await listing(repeated, first.next_cursor)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_changed_upstream_revision_requires_restart(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "shared-test-key")
|
||||
|
||||
async def original(server_id: str, cursor: str | None) -> ListToolsResult:
|
||||
return page(server_id, "next")
|
||||
|
||||
async def revised(server_id: str, cursor: str | None) -> ListToolsResult:
|
||||
return page(server_id, revision="changed")
|
||||
|
||||
first = await listing(original)
|
||||
with pytest.raises(MCPError, match="fresh listing"):
|
||||
await listing(revised, first.next_cursor)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_failed_upstream_remains_visible_when_other_sources_continue(monkeypatch):
|
||||
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import SERVER_OUTCOMES_META_KEY
|
||||
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "shared-test-key")
|
||||
|
||||
async def fetch(server_id: str, cursor: str | None) -> ListToolsResult:
|
||||
if server_id == "a":
|
||||
return ListToolsResult(tools=[], meta={SERVER_OUTCOMES_META_KEY: {"a": {"tag": "timeout"}}})
|
||||
return page("b2" if cursor else "b1", None if cursor else "next")
|
||||
|
||||
first = await listing(fetch)
|
||||
assert [tool.name for tool in first.tools] == ["b1"]
|
||||
assert first.meta[SERVER_OUTCOMES_META_KEY]["a"] == {"tag": "timeout"}
|
||||
second = await listing(fetch, first.next_cursor)
|
||||
assert [tool.name for tool in second.tools] == ["b2"]
|
||||
assert second.meta[SERVER_OUTCOMES_META_KEY]["a"] == {"tag": "timeout"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cursor_cannot_be_reused_for_a_different_catalog_kind(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "shared-test-key")
|
||||
|
||||
async def first_page(server_id: str, cursor: str | None) -> ListToolsResult:
|
||||
return page(server_id, "next")
|
||||
|
||||
async def forbidden_dispatch(server_id: str, cursor: str | None) -> ListToolsResult:
|
||||
pytest.fail("Wrong-kind cursor must not contact upstream")
|
||||
|
||||
first = await listing(first_page)
|
||||
with pytest.raises(MCPError, match="Invalid pagination state"):
|
||||
await catalog.paginate_catalog(
|
||||
method="resources/list",
|
||||
cursor=first.next_cursor,
|
||||
caller_scope="caller-and-scope",
|
||||
snapshot="registry-generation",
|
||||
server_ids=("a", "b"),
|
||||
fetch=forbidden_dispatch,
|
||||
now=100,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticated_state_with_invalid_catalog_schema_is_rejected(monkeypatch):
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok
|
||||
from litellm.proxy._experimental.mcp_server.state_tokens import seal_state
|
||||
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "shared-test-key")
|
||||
token = seal_state({"not": "catalog state"}, purpose="mcp.catalog.list.v1:tools/list", expires_at=200, now=100)
|
||||
assert isinstance(token, Ok)
|
||||
|
||||
async def forbidden_dispatch(server_id, cursor):
|
||||
pytest.fail("Malformed catalog state must not dispatch")
|
||||
|
||||
with pytest.raises(MCPError, match="Invalid pagination state"):
|
||||
await listing(forbidden_dispatch, token.ok)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_following_page_does_not_extend_original_expiry(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "shared-test-key")
|
||||
|
||||
async def fetch(server_id, cursor):
|
||||
return page(server_id, "second" if cursor is None else "third")
|
||||
|
||||
first = await listing(fetch, now=100)
|
||||
second = await listing(fetch, first.next_cursor, now=699)
|
||||
assert second.next_cursor
|
||||
with pytest.raises(MCPError, match="expired"):
|
||||
await listing(fetch, second.next_cursor, now=700)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_page_limit_rejects_an_unending_upstream(monkeypatch):
|
||||
from litellm import constants
|
||||
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "shared-test-key")
|
||||
monkeypatch.setattr(constants, "MCP_TOOL_LISTING_MAX_PAGES", 2)
|
||||
|
||||
async def fetch(server_id, cursor):
|
||||
return page(server_id, "second" if cursor is None else "third")
|
||||
|
||||
first = await listing(fetch)
|
||||
with pytest.raises(MCPError, match="limit"):
|
||||
await listing(fetch, first.next_cursor)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"change",
|
||||
[
|
||||
{"_caller": None},
|
||||
{"mcp_auth_header": "another upstream credential"},
|
||||
{"mcp_servers": ("another-scope",)},
|
||||
{"client_ip": "192.0.2.2"},
|
||||
{"raw_headers": {"Authorization": "another bearer"}},
|
||||
{"oauth2_headers": {"Authorization": "another upstream bearer"}},
|
||||
{"mcp_server_auth_headers": {"a": {"Authorization": "another per-server bearer"}}},
|
||||
{"protocol_version": "2024-11-05"},
|
||||
{"mcp_proxy_mode": True},
|
||||
],
|
||||
)
|
||||
def test_caller_binding_covers_identity_scope_and_forwarded_credentials(change):
|
||||
from dataclasses import replace
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.contracts import OperationContext
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
original = OperationContext(_caller=UserAPIKeyAuth(api_key="synthetic-key", user_id="owner"))
|
||||
assert catalog._caller_scope(original, ()) != catalog._caller_scope(replace(original, **change), ())
|
||||
|
||||
|
||||
def test_caller_binding_ignores_transport_headers_and_header_case():
|
||||
from dataclasses import replace
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.contracts import OperationContext
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
original = OperationContext(
|
||||
_caller=UserAPIKeyAuth(user_id="owner"), raw_headers={"Authorization": "Bearer synthetic"}
|
||||
)
|
||||
retry = replace(original, raw_headers={"authorization": "Bearer synthetic", "mcp-session-id": "replica-b-session"})
|
||||
assert catalog._caller_scope(original, ()) == catalog._caller_scope(retry, ())
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upstream_pages_overlap_and_keep_deterministic_order(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "shared-test-key")
|
||||
both_started = asyncio.Event()
|
||||
started = set()
|
||||
|
||||
async def fetch(server_id, cursor):
|
||||
started.add(server_id)
|
||||
if len(started) == 2:
|
||||
both_started.set()
|
||||
await asyncio.wait_for(both_started.wait(), timeout=1)
|
||||
return page(server_id)
|
||||
|
||||
result = await listing(fetch, servers=("b", "a"))
|
||||
assert [tool.name for tool in result.tools] == ["a", "b"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cursor_history_allows_the_full_supported_page_count(monkeypatch):
|
||||
from litellm.constants import MCP_TOOL_LISTING_MAX_PAGES
|
||||
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "shared-test-key")
|
||||
|
||||
async def fetch(server_id, cursor):
|
||||
index = int(cursor or "0")
|
||||
return page(str(index), str(index + 1) if index + 1 < MCP_TOOL_LISTING_MAX_PAGES else None)
|
||||
|
||||
cursor = None
|
||||
for index in range(MCP_TOOL_LISTING_MAX_PAGES):
|
||||
result = await listing(fetch, cursor, servers=("a",))
|
||||
assert result.tools[0].name == str(index)
|
||||
cursor = result.next_cursor
|
||||
assert bool(cursor) == (index + 1 < MCP_TOOL_LISTING_MAX_PAGES)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_failed_page_cancels_and_drains_other_upstream_requests(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "shared-test-key")
|
||||
pending_started = asyncio.Event()
|
||||
pending_closed = asyncio.Event()
|
||||
|
||||
async def fetch(server_id, cursor):
|
||||
if server_id == "a":
|
||||
await asyncio.wait_for(pending_started.wait(), timeout=1)
|
||||
raise ValueError("upstream unavailable")
|
||||
pending_started.set()
|
||||
try:
|
||||
await asyncio.Event().wait()
|
||||
finally:
|
||||
pending_closed.set()
|
||||
return page(server_id)
|
||||
|
||||
with pytest.raises(ValueError, match="upstream unavailable"):
|
||||
await listing(fetch)
|
||||
assert pending_closed.is_set()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gateway_tools_continuation_rejects_an_upstream_failure(monkeypatch):
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
from mcp.types import PaginatedRequestParams
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._experimental.mcp_server import operations
|
||||
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ServerListOk, classify_list_exception
|
||||
from litellm.types.mcp import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "catalog-failure-test")
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
server = MCPServer(server_id="pages", name="pages", transport=MCPTransport.http)
|
||||
monkeypatch.setattr(operations.global_mcp_server_manager, "registry", {server.server_id: server})
|
||||
fetch = AsyncMock(side_effect=[
|
||||
(page("first", "next"), ServerListOk(tool_count=1)),
|
||||
(ListToolsResult(tools=[]), classify_list_exception(TimeoutError("upstream secret"))),
|
||||
])
|
||||
monkeypatch.setattr(catalog, "get_filtered_server_tools", fetch)
|
||||
context = operations.prepare_context()
|
||||
first = await catalog.aggregate_gateway_tools(context, PaginatedRequestParams(), [server], {})
|
||||
assert first.next_cursor and [tool.name for tool in first.tools] == ["first"]
|
||||
with pytest.raises(MCPError, match="Upstream continuation failed; start a fresh listing") as denied:
|
||||
await catalog.aggregate_gateway_tools(context, PaginatedRequestParams(cursor=first.next_cursor), [server], {})
|
||||
assert "upstream secret" not in str(denied.value)
|
||||
assert fetch.await_count == 2
|
||||
assert fetch.await_args.kwargs["params"].cursor == "next"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("request_name,result_name,field", [
|
||||
("ListPromptsRequest", "ListPromptsResult", "prompts"),
|
||||
("ListResourcesRequest", "ListResourcesResult", "resources"),
|
||||
("ListResourceTemplatesRequest", "ListResourceTemplatesResult", "resource_templates"),
|
||||
])
|
||||
async def test_optional_gateway_catalog_reports_initial_failure_and_rejects_failed_continuation(monkeypatch, request_name, result_name, field):
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
from mcp import types
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._experimental.mcp_server import operations
|
||||
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import SERVER_OUTCOMES_META_KEY
|
||||
from litellm.types.mcp import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "catalog-failure-test")
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
server = MCPServer(server_id="pages", name="pages", transport=MCPTransport.http)
|
||||
monkeypatch.setattr(operations.global_mcp_server_manager, "registry", {server.server_id: server})
|
||||
monkeypatch.setattr(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server]))
|
||||
fetch = AsyncMock(side_effect=TimeoutError("upstream secret"))
|
||||
monkeypatch.setattr(catalog, "fetch_optional_catalog_page", fetch)
|
||||
context = operations.prepare_context()
|
||||
request = getattr(types, request_name)
|
||||
failed = await catalog.list_gateway_catalog(context, request())
|
||||
assert getattr(failed, field) == [] and failed.next_cursor is None
|
||||
assert next(iter(failed.meta[SERVER_OUTCOMES_META_KEY].values()))["status"] == "timeout"
|
||||
assert "upstream secret" not in failed.model_dump_json()
|
||||
fetch.side_effect = [getattr(types, result_name)(**{field: [], "next_cursor": "next"}), TimeoutError("upstream secret")]
|
||||
first = await catalog.list_gateway_catalog(context, request())
|
||||
assert first.next_cursor
|
||||
with pytest.raises(MCPError, match="Upstream continuation failed; start a fresh listing") as denied:
|
||||
await catalog.list_gateway_catalog(context, request(params=types.PaginatedRequestParams(cursor=first.next_cursor)))
|
||||
assert "upstream secret" not in str(denied.value)
|
||||
assert fetch.await_count == 3
|
||||
assert fetch.await_args.args[-1] == "next"
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 = {}
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
Loading…
Add table
Reference in a new issue