diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index 21b82db44b7..bd672f4c160 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -344,8 +344,15 @@ jobs: - shard: mcp-elicitation artifact-name: mcp-elicitation test-path: >- + --override-ini=pythonpath=tests + tests/unit/experimental_mcp_client/test_mcp_client.py + tests/unit/proxy/_experimental/mcp_server/test_capabilities.py + tests/unit/proxy/_experimental/mcp_server/test_interactions.py tests/unit/proxy/_experimental/mcp_server/test_mcp_elicitation_handler.py tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py + tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py + tests/unit/proxy/_experimental/mcp_server/test_operations.py + tests/integration/mcp/test_interactions.py workers: 2 reruns: 0 timeout-minutes: 20 diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 960b9dde53f..729a6acb20d 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -20,6 +20,7 @@ import httpx2 from httpx2._client import UseClientDefault from httpx2._types import AuthTypes from mcp import ClientSession, MCPError, ReadResourceResult, Resource, StdioServerParameters +from mcp.client._input_required import run_input_required_driver from mcp.client.sse import sse_client from mcp.client.stdio import stdio_client from mcp.client.streamable_http import streamable_http_client @@ -49,6 +50,7 @@ from mcp.types import ( InitializeRequestParams, InitializeResult, InputRequiredResult, + InputResponses, ListPromptsRequest, ListPromptsResult, ListResourcesRequest, @@ -100,6 +102,7 @@ from litellm.types.mcp import ( if TYPE_CHECKING: from litellm.proxy._experimental.mcp_server.contracts import CatalogListRequest, CatalogListResult + from litellm.proxy._experimental.mcp_server.legacy_callbacks import ElicitationCallback def to_basic_auth(auth_value: str) -> str: @@ -411,7 +414,7 @@ class MCPClient: aws_auth: httpx2.Auth | None = None, resolved_auth: httpx2.Auth | None = None, sampling_callback: Callable | None = None, - elicitation_callback: Callable | None = None, + elicitation_callback: "ElicitationCallback | None" = None, logging_callback: Callable | None = None, protocol_version: MCPUpstreamProtocol = "auto", ): @@ -438,7 +441,7 @@ class MCPClient: self._resolved_auth: httpx2.Auth | None = resolved_auth self._last_initialize_instructions: str | None = None self._sampling_callback: Callable | None = sampling_callback - self._elicitation_callback: Callable | None = elicitation_callback + self._elicitation_callback: ElicitationCallback | None = elicitation_callback self._logging_callback: Callable | None = logging_callback # handle the basic auth value if provided if auth_value: @@ -631,15 +634,6 @@ class MCPClient: # The SDK closes pending requests when its message handler raises. raise RuntimeError("MCP response stream failed") - session_kwargs: Final = { - name: callback - for name, callback in ( - ("sampling_callback", self._sampling_callback), - ("elicitation_callback", self._elicitation_callback), - ("logging_callback", self._logging_callback), - ) - if callback is not None - } # The SDK drops a response stream that ends without a JSON-RPC reply, so nothing else # ever fails the request. session_ctx: Final = ClientSession( @@ -647,7 +641,9 @@ class MCPClient: write_stream, read_timeout_seconds=self.timeout, message_handler=receive_message, - **session_kwargs, + sampling_callback=self._sampling_callback, + elicitation_callback=self._elicitation_callback, + logging_callback=self._logging_callback, ) session: Final = await session_ctx.__aenter__() try: @@ -948,6 +944,28 @@ class MCPClient: """The error result ``call_tool`` returns when it swallows a failure (no re-execution).""" return error_text_result(exc) + async def _request_with_interaction( + self, + session: ClientSession, + request: Callable[[InputResponses | None, str | None], Awaitable[TSessionResult | InputRequiredResult]], + input_responses: InputResponses | None, + request_state: str | None, + allow_input_required: bool, + ) -> TSessionResult | InputRequiredResult: + from litellm.proxy._experimental.mcp_server.contracts import ClientInteraction + from litellm.proxy._experimental.mcp_server.interactions import LegacyClientInteraction, ModernClientInteraction + + with anyio.fail_after(self.timeout): + first: Final = await request(input_responses, request_state) + if allow_input_required: + return await ModernClientInteraction( + session, allow_elicitation=self._elicitation_callback is not None + ).complete(first, request) + if not isinstance(first, InputRequiredResult): + return first + interaction: Final[ClientInteraction] = LegacyClientInteraction(session) + return await run_input_required_driver(first, dispatch=interaction.request, retry=request) + async def call_tool( self, call_tool_request_params: MCPCallToolRequestParams, @@ -990,11 +1008,25 @@ class MCPClient: ) if not any(tool.name == call_tool_request_params.name for tool in tools): raise MCPError(code=-32603, message="Tool schema is unavailable from the bounded upstream catalog") - return await session.call_tool( - name=call_tool_request_params.name, - arguments=call_tool_request_params.arguments, - progress_callback=on_progress, - allow_input_required=allow_input_required, + + async def request( + responses: InputResponses | None, state: str | None + ) -> MCPCallToolResult | InputRequiredResult: + return await session.call_tool( + name=call_tool_request_params.name, + arguments=call_tool_request_params.arguments, + input_responses=responses, + request_state=state, + progress_callback=on_progress, + allow_input_required=True, + ) + + return await self._request_with_interaction( + session, + request, + call_tool_request_params.input_responses, + call_tool_request_params.request_state, + allow_input_required, ) try: @@ -1129,15 +1161,32 @@ class MCPClient: # Return empty list instead of raising to allow graceful degradation return ListPromptsResult(prompts=[]) - async def get_prompt(self, get_prompt_request_params: GetPromptRequestParams) -> GetPromptResult: + async def get_prompt( + self, get_prompt_request_params: GetPromptRequestParams, *, allow_input_required: bool = False + ) -> GetPromptResult | InputRequiredResult: """Fetch a prompt definition from the MCP server.""" verbose_logger.info("MCP client fetching prompt '%s'", get_prompt_request_params.name) async def _get_prompt_operation(session: ClientSession): verbose_logger.debug("MCP client sending get_prompt request to session") - return await session.get_prompt( - name=get_prompt_request_params.name, - arguments=get_prompt_request_params.arguments, + + async def request( + responses: InputResponses | None, state: str | None + ) -> GetPromptResult | InputRequiredResult: + return await session.get_prompt( + name=get_prompt_request_params.name, + arguments=get_prompt_request_params.arguments, + input_responses=responses, + request_state=state, + allow_input_required=True, + ) + + return await self._request_with_interaction( + session, + request, + get_prompt_request_params.input_responses, + get_prompt_request_params.request_state, + allow_input_required, ) try: @@ -1285,13 +1334,37 @@ class MCPClient: # Return empty list instead of raising to allow graceful degradation return ListResourceTemplatesResult(resource_templates=[]) - async def read_resource(self, url: AnyUrl) -> ReadResourceResult: + async def read_resource( + self, + url: AnyUrl, + *, + input_responses: InputResponses | None = None, + request_state: str | None = None, + allow_input_required: bool = False, + ) -> ReadResourceResult | InputRequiredResult: """Fetch resource contents from the MCP server.""" verbose_logger.info("MCP client fetching resource '%s'", url) async def _read_resource_operation(session: ClientSession): verbose_logger.debug("MCP client sending read_resource request to session") - return await session.read_resource(str(url)) + + async def request( + responses: InputResponses | None, state: str | None + ) -> ReadResourceResult | InputRequiredResult: + return await session.read_resource( + str(url), + input_responses=responses, + request_state=state, + allow_input_required=True, + ) + + return await self._request_with_interaction( + session, + request, + input_responses, + request_state, + allow_input_required, + ) try: read_resource_result: Final = await self.run_with_session(_read_resource_operation) diff --git a/litellm/proxy/_experimental/mcp_server/capabilities.py b/litellm/proxy/_experimental/mcp_server/capabilities.py index bfd00327eb4..1ee263633be 100644 --- a/litellm/proxy/_experimental/mcp_server/capabilities.py +++ b/litellm/proxy/_experimental/mcp_server/capabilities.py @@ -11,7 +11,13 @@ from mcp_types.methods import CLIENT_REQUESTS from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS, LATEST_HANDSHAKE_VERSION from pydantic import TypeAdapter -from litellm.types.mcp import MCP_LEGACY_VERSIONS, MCPAdvertisedVersions, MCPLegacyVersion, MCPSpecVersion, MCPTransport +from litellm.types.mcp import ( + MCP_LEGACY_VERSIONS, + MCPAdvertisedVersion, + MCPAdvertisedVersions, + MCPSpecVersion, + MCPTransport, +) GATEWAY_OPERATIONS: Final = frozenset( { @@ -46,14 +52,14 @@ REVISION_SUPPORT: Final[Mapping[str, RevisionSupport]] = MappingProxyType( if version.value in HANDSHAKE_PROTOCOL_VERSIONS else frozenset({"complete", "input_required"}), extensions=frozenset(), - completed=version.value in HANDSHAKE_PROTOCOL_VERSIONS, + completed=version.value in HANDSHAKE_PROTOCOL_VERSIONS or version.value == "2026-07-28", ) for version in MCPSpecVersion } ) _COMPLETED_REVISIONS: Final = tuple(version for version, support in REVISION_SUPPORT.items() if support.completed) TRANSLATION_PAIRS: Final = frozenset(product(_COMPLETED_REVISIONS, repeat=2)) -_ADVERTISED_VERSIONS: Final[TypeAdapter[tuple[MCPLegacyVersion, ...]]] = TypeAdapter(MCPAdvertisedVersions) +_ADVERTISED_VERSIONS: Final[TypeAdapter[tuple[MCPAdvertisedVersion, ...]]] = TypeAdapter(MCPAdvertisedVersions) def configured_versions() -> tuple[str, ...]: diff --git a/litellm/proxy/_experimental/mcp_server/contracts.py b/litellm/proxy/_experimental/mcp_server/contracts.py index f3800c93a77..bf0fae1a210 100644 --- a/litellm/proxy/_experimental/mcp_server/contracts.py +++ b/litellm/proxy/_experimental/mcp_server/contracts.py @@ -7,6 +7,8 @@ from datetime import datetime from types import MappingProxyType from typing import TYPE_CHECKING, Final, Literal, Protocol, TypeAlias +from mcp.types import ErrorData, InputRequest, InputResponse, InputResponses + 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 @@ -121,6 +123,10 @@ class ProgressCallback(Protocol): async def __call__(self, progress: float, total: float | None, /) -> None: ... +class ClientInteraction(Protocol): + async def request(self, key: str, request: InputRequest) -> InputResponse | ErrorData: ... + + @dataclass(frozen=True, slots=True) class AuthorizedToolCall: name: str @@ -130,3 +136,5 @@ class AuthorizedToolCall: host_progress_callback: ProgressCallback | None guardrail_context: Mapping[str, object] | None logging_data: Mapping[str, object] + input_responses: InputResponses | None = None + request_state: str | None = None diff --git a/litellm/proxy/_experimental/mcp_server/interactions.py b/litellm/proxy/_experimental/mcp_server/interactions.py new file mode 100644 index 00000000000..42addcfb65b --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/interactions.py @@ -0,0 +1,262 @@ +import asyncio +import hashlib +import json +import secrets +from collections.abc import Awaitable, Callable +from dataclasses import dataclass +from typing import Final, Literal, TypeAlias, TypeVar +from uuid import uuid4 + +from mcp import MCPError +from mcp.client._input_required import run_input_required_driver +from mcp.client.session import ClientRequestContext, ClientSession +from mcp.types import ( + CallToolRequest, + ElicitRequest, + ElicitRequestURLParams, + ErrorData, + GetPromptRequest, + InputRequest, + InputRequests, + InputRequiredResult, + InputResponse, + InputResponses, + ReadResourceRequest, +) +from pydantic import BaseModel, ConfigDict, Field, JsonValue, TypeAdapter, ValidationError + +from litellm.proxy._experimental.mcp_server.contracts import OperationContext +from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error +from litellm.proxy._experimental.mcp_server.state_tokens import StateTokenError, open_state, seal_state +from litellm.types.mcp_server.mcp_server_manager import MCPServer + + +@dataclass(frozen=True, slots=True) +class LegacyClientInteraction: + session: ClientSession + + async def request(self, key: str, request: InputRequest) -> InputResponse | ErrorData: + context: Final = ClientRequestContext( + session=self.session, request_id=key, meta=request.params.meta if request.params else None + ) + legacy_request: Final = ( + request.model_copy(update={"params": request.params.model_copy(update={"elicitation_id": str(uuid4())})}) + if isinstance(request, ElicitRequest) + and isinstance(request.params, ElicitRequestURLParams) + and request.params.elicitation_id is None + else request + ) + return await self.session.dispatch_input_request(context, legacy_request) + + +InteractionOperation: TypeAlias = CallToolRequest | GetPromptRequest | ReadResourceRequest +_JSON: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) +_PURPOSE: Final = "mcp:interaction:repeatable:v1" + + +class BoundInputRequiredResult(InputRequiredResult): + target_id: str | None = Field(default=None, exclude=True) + target_digest: str | None = Field(default=None, exclude=True) + gateway_responses: InputResponses | None = Field(default=None, exclude=True) + + +class ContinuationState(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid") + + principal: str + operation: str + target_id: str + target_digest: str + upstream_state: str | None + gateway_responses: InputResponses | None = None + policy: Literal["repeatable"] = "repeatable" + expires_at: int + nonce: str + + +def _digest(value: JsonValue) -> str: + return hashlib.sha256( + json.dumps(value, sort_keys=True, separators=(",", ":"), allow_nan=False).encode() + ).hexdigest() + + +def target_digest(server: MCPServer) -> str: + return _digest( + _JSON.validate_python( + { + "id": server.server_id, + "url": server.url, + "transport": server.transport, + "protocol": server.protocol_version, + "command": server.command, + "args": server.args, + } + ) + ) + + +def bind_target(result: InputRequiredResult, server: MCPServer) -> BoundInputRequiredResult: + bound: Final = ( + result + if isinstance(result, BoundInputRequiredResult) + else BoundInputRequiredResult.model_validate(result.model_dump()) + ) + return bound.model_copy(update={"target_id": server.server_id, "target_digest": target_digest(server)}) + + +def _principal(context: OperationContext) -> str: + caller: Final = context.user_api_key_auth + if caller is None or not caller.user_id: + raise MCPError(code=-32602, message="MCP continuations require an authenticated caller identity") + return _digest( + _JSON.validate_python( + { + "user": caller.user_id, + "team": caller.team_id, + "org": caller.org_id, + "end_user": caller.end_user_id, + "servers": sorted(context.mcp_servers) if context.mcp_servers is not None else None, + } + ) + ) + + +def _operation(operation: InteractionOperation) -> str: + return _digest( + _JSON.validate_python( + { + "method": operation.method, + "params": operation.params.model_dump( + mode="json", by_alias=True, exclude={"meta", "input_responses", "request_state"} + ), + } + ) + ) + + +def _state_error(error: StateTokenError) -> MCPError: + return MCPError( + code=-32602, + message=( + "Set the same LITELLM_SALT_KEY on every replica to enable MCP continuations" + if error is StateTokenError.MISSING_KEY + else "Invalid or expired MCP continuation; start a fresh request" + ), + ) + + +def open_continuation( + operation: InteractionOperation, context: OperationContext, *, now: int +) -> ContinuationState | None: + token: Final = operation.params.request_state + if token is None: + if operation.params.input_responses: + raise MCPError(code=-32602, message="Input responses require a gateway continuation") + return None + opened: Final = open_state(token, purpose=_PURPOSE, now=now) + if isinstance(opened, Error): + raise _state_error(opened.error) + try: + state: Final = ContinuationState.model_validate(opened.ok) + except ValidationError as error: + raise _state_error(StateTokenError.INVALID) from error + if state.principal != _principal(context) or state.operation != _operation(operation) or state.expires_at <= now: + raise _state_error(StateTokenError.INVALID) + return state + + +def seal_continuation( + result: BoundInputRequiredResult, + operation: InteractionOperation, + context: OperationContext, + *, + now: int, + previous: ContinuationState | None = None, +) -> InputRequiredResult: + if result.target_id is None or result.target_digest is None: + raise MCPError(code=-32602, message="MCP continuation target is unavailable") + state: Final = ContinuationState( + principal=_principal(context), + operation=_operation(operation), + target_id=result.target_id, + target_digest=result.target_digest, + upstream_state=result.request_state, + gateway_responses=result.gateway_responses, + expires_at=previous.expires_at if previous is not None else now + 600, + nonce=previous.nonce if previous is not None else secrets.token_urlsafe(24), + ) + sealed: Final = seal_state( + _JSON.validate_json(state.model_dump_json()), purpose=_PURPOSE, expires_at=state.expires_at, now=now + ) + if isinstance(sealed, Error): + raise _state_error(sealed.error) + return InputRequiredResult(input_requests=result.input_requests, request_state=sealed.ok, _meta=result.meta) + + +_Terminal: Final = TypeVar("_Terminal") + + +@dataclass(frozen=True, slots=True) +class _DeferredInteraction: + result: BoundInputRequiredResult + + +@dataclass(frozen=True, slots=True) +class ModernClientInteraction: + session: ClientSession + allow_elicitation: bool + + async def request(self, key: str, request: InputRequest) -> InputResponse | ErrorData: + if isinstance(request, ElicitRequest): + return ErrorData(code=-32602, message="Modern elicitation requires a continuation") + return await LegacyClientInteraction(self.session).request(key, request) + + async def prepare( + self, result: _Terminal | InputRequiredResult + ) -> _Terminal | InputRequiredResult | _DeferredInteraction: + if not isinstance(result, InputRequiredResult): + return result + pending: Final[InputRequests] = { + key: request for key, request in (result.input_requests or {}).items() if isinstance(request, ElicitRequest) + } + if pending and not self.allow_elicitation: + raise MCPError(code=-32602, message="Elicitation is disabled for this MCP server") + if result.input_requests and not pending: + return result + local: Final = tuple( + (key, request) for key, request in (result.input_requests or {}).items() if key not in pending + ) + responses: Final = await asyncio.gather(*(self.request(key, request) for key, request in local)) + for response in responses: + if isinstance(response, ErrorData): + raise MCPError(code=response.code, message=response.message) + return _DeferredInteraction( + BoundInputRequiredResult( + input_requests=pending or None, + request_state=result.request_state, + _meta=result.meta, + gateway_responses={ + key: response for (key, _), response in zip(local, responses) if not isinstance(response, ErrorData) + } + or None, + ) + ) + + async def complete( + self, + first: _Terminal | InputRequiredResult, + retry: Callable[[InputResponses | None, str | None], Awaitable[_Terminal | InputRequiredResult]], + ) -> _Terminal | InputRequiredResult: + prepared: Final = await self.prepare(first) + if isinstance(prepared, _DeferredInteraction): + return prepared.result + if not isinstance(prepared, InputRequiredResult): + return prepared + + async def resume( + responses: InputResponses | None, state: str | None + ) -> _Terminal | InputRequiredResult | _DeferredInteraction: + return await self.prepare(await retry(responses, state)) + + completed: Final = await run_input_required_driver(prepared, dispatch=self.request, retry=resume) + return completed.result if isinstance(completed, _DeferredInteraction) else completed diff --git a/litellm/proxy/_experimental/mcp_server/legacy_callbacks.py b/litellm/proxy/_experimental/mcp_server/legacy_callbacks.py index 13c3d265429..17bb84e066a 100644 --- a/litellm/proxy/_experimental/mcp_server/legacy_callbacks.py +++ b/litellm/proxy/_experimental/mcp_server/legacy_callbacks.py @@ -24,7 +24,7 @@ class SamplingCallback(Protocol): class ElicitationCallback(Protocol): - async def __call__(self, context: object, params: ElicitRequestParams, /) -> ElicitResult | ErrorData: ... + async def __call__(self, context: object, params: ElicitRequestParams) -> ElicitResult | ErrorData: ... def create_sampling_callback( @@ -79,6 +79,11 @@ def create_elicitation_callback(timeout: float | None = None) -> ElicitationCall relay_timeout: Final = timeout if timeout is not None else MCP_CLIENT_TIMEOUT async def callback(context: object, params: ElicitRequestParams) -> ElicitResult | ErrorData: + if request is not None and request.protocol_version == "2026-07-28": + return ErrorData( + code=-32602, + message="A legacy upstream cannot resume input for a modern client; the operation may have partially completed", + ) from litellm.proxy._experimental.mcp_server.elicitation_handler import handle_elicitation_request return await handle_elicitation_request( diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 217e373030d..aa31799f1aa 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -44,6 +44,7 @@ from mcp.types import ( GetPromptRequestParams, GetPromptResult, InputRequiredResult, + InputResponses, ListPromptsRequest, ListPromptsResult, ListResourcesRequest, @@ -96,6 +97,7 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ( raise_classified_list_failure, upstream_auth_challenge, ) +from litellm.proxy._experimental.mcp_server.interactions import bind_target from litellm.proxy._experimental.mcp_server.mcp_debug import describe_upstream_http_failure, record_auth_resolution from litellm.proxy._experimental.mcp_server.oauth2_token_cache import ( MCPPerUserTokenCache, @@ -4888,7 +4890,10 @@ class MCPServerManager: extra_headers: dict[str, str] | None = None, raw_headers: dict[str, str] | None = None, client_ip: str | None = None, - ) -> ReadResourceResult: + input_responses: InputResponses | None = None, + request_state: str | None = None, + allow_input_required: bool = False, + ) -> ReadResourceResult | InputRequiredResult: """Read resource contents from a specific MCP server.""" verbose_logger.debug("Connecting to url: %s", server.url) @@ -4913,7 +4918,10 @@ class MCPServerManager: user_api_key_auth=user_api_key_auth, ) - return await client.read_resource(url) + result: Final = await client.read_resource( + url, input_responses=input_responses, request_state=request_state, allow_input_required=allow_input_required + ) + return bind_target(result, server) if isinstance(result, InputRequiredResult) else result async def get_prompt_from_server( self, @@ -4925,7 +4933,10 @@ class MCPServerManager: extra_headers: dict[str, str] | None = None, raw_headers: dict[str, str] | None = None, client_ip: str | None = None, - ) -> GetPromptResult: + input_responses: InputResponses | None = None, + request_state: str | None = None, + allow_input_required: bool = False, + ) -> GetPromptResult | InputRequiredResult: """Fetch a specific prompt definition from a single MCP server.""" verbose_logger.debug("Connecting to url: %s", server.url) @@ -4953,8 +4964,11 @@ class MCPServerManager: get_prompt_request_params: Final = GetPromptRequestParams( name=prompt_name, arguments=arguments, + input_responses=input_responses, + request_state=request_state, ) - return await client.get_prompt(get_prompt_request_params) + result: Final = await client.get_prompt(get_prompt_request_params, allow_input_required=allow_input_required) + return bind_target(result, server) if isinstance(result, InputRequiredResult) else result @staticmethod def _is_same_authority_metadata_url(url: str, server_url: str) -> bool: @@ -6178,6 +6192,8 @@ class MCPServerManager: user_api_key_auth: UserAPIKeyAuth | None = None, client_ip: str | None = None, allow_input_required: bool = False, + input_responses: InputResponses | None = None, + request_state: str | None = None, ) -> CallToolResult | InputRequiredResult: """ Call a regular MCP tool using the MCP client. @@ -6319,6 +6335,8 @@ class MCPServerManager: call_tool_params: Final = MCPCallToolRequestParams( name=original_tool_name, arguments=arguments, + input_responses=input_responses, + request_state=request_state, ) if _obo_retry_applies(mcp_server, subject_token): @@ -6425,7 +6443,11 @@ class MCPServerManager: result: Final = mcp_responses[result_index] self._remember_upstream_initialize_instructions(mcp_server, client) - return cast("CallToolResult | InputRequiredResult", result) + return ( + bind_target(result, mcp_server) + if isinstance(result, InputRequiredResult) + else cast("CallToolResult", result) + ) def _resolve_mcp_server_for_tool_call( self, @@ -6637,6 +6659,8 @@ class MCPServerManager: guardrail_context: Mapping[str, object] | None = None, client_ip: str | None = None, wire_compat: WireCompat = WireCompat.LEGACY, + input_responses: InputResponses | None = None, + request_state: str | None = None, *, catalog_auth_header: str | None | EllipsisType = ..., listed_tool: MCPTool | None | EllipsisType = ..., @@ -6779,6 +6803,8 @@ class MCPServerManager: hook_extra_headers=hook_result.get("extra_headers"), user_api_key_auth=user_api_key_auth, allow_input_required=wire_compat is WireCompat.MODERN, + input_responses=input_responses, + request_state=request_state, ) return await self._gather_openapi_tool_tasks(tasks, proxy_logging_obj) diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 6ec5bdf6af7..039fbfdb916 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -1,18 +1,19 @@ """Shared MCP operation policy and dispatch.""" import asyncio +import time import traceback import types import uuid from collections.abc import Mapping, Sequence from contextvars import ContextVar -from dataclasses import dataclass +from dataclasses import dataclass, replace from datetime import datetime from functools import partial from typing import Any, Final, NoReturn, TypeAlias, overload from fastapi import HTTPException -from mcp import ReadResourceResult, Resource +from mcp import MCPError, ReadResourceResult, Resource from mcp.types import ( CallToolRequest, CallToolRequestParams, @@ -23,6 +24,7 @@ from mcp.types import ( GetPromptRequestParams, GetPromptResult, InputRequiredResult, + InputResponses, ListPromptsRequest, ListPromptsResult, ListResourcesRequest, @@ -83,6 +85,14 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ( classify_list_exception, outcome_wire_value, ) +from litellm.proxy._experimental.mcp_server.interactions import ( + BoundInputRequiredResult, + ContinuationState, + InteractionOperation, + open_continuation, + seal_continuation, + target_digest, +) from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: F401 # legacy module exports MCPServerManager, _caller_authorization_fans_out, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export @@ -1862,6 +1872,8 @@ async def execute_mcp_tool( guardrail_context: Mapping[str, object] | None = None, client_ip: str | None = None, wire_compat: WireCompat = WireCompat.LEGACY, + input_responses: InputResponses | None = None, + request_state: str | None = None, **kwargs: object, # kwargs-ok: preserves the existing REST and decorated logging call contract ) -> CallToolResult | InputRequiredResult: context: Final = prepare_context( @@ -1881,6 +1893,8 @@ async def execute_mcp_tool( host_progress_callback=host_progress_callback, guardrail_context=guardrail_context, logging_data=types.MappingProxyType(kwargs), + input_responses=input_responses, + request_state=request_state, ) return await GatewayOperations().execute(operation, context) @@ -1899,6 +1913,8 @@ async def _execute_mcp_tool( guardrail_context: Mapping[str, object] | None = None, client_ip: str | None = None, wire_compat: WireCompat = WireCompat.LEGACY, + input_responses: InputResponses | None = None, + request_state: str | None = None, **kwargs: Any, ) -> CallToolResult | InputRequiredResult: """ @@ -2194,6 +2210,8 @@ async def _execute_mcp_tool( guardrail_context=guardrail_context, host_progress_callback=host_progress_callback, wire_compat=wire_compat, + input_responses=input_responses, + request_state=request_state, ) # Fall back to local tool registry with original name (legacy support) @@ -2435,6 +2453,8 @@ async def call_mcp_tool( raw_headers: dict[str, str] | None = None, client_ip: str | None = None, wire_compat: WireCompat = WireCompat.LEGACY, + input_responses: InputResponses | None = None, + request_state: str | None = None, **kwargs: Any, ) -> CallToolResult | InputRequiredResult: """ @@ -2493,6 +2513,8 @@ async def call_mcp_tool( raw_headers=raw_headers, client_ip=client_ip, wire_compat=wire_compat, + input_responses=input_responses, + request_state=request_state, **kwargs, ) except Exception as e: @@ -2525,7 +2547,10 @@ async def mcp_get_prompt( oauth2_headers: dict[str, str] | None = None, raw_headers: dict[str, str] | None = None, client_ip: str | None = None, -) -> GetPromptResult: + input_responses: InputResponses | None = None, + request_state: str | None = None, + allow_input_required: bool = False, +) -> GetPromptResult | InputRequiredResult: """ Fetch a specific MCP prompt, handling both prefixed and unprefixed names. """ @@ -2570,6 +2595,9 @@ async def mcp_get_prompt( extra_headers=extra_headers, raw_headers=raw_headers, client_ip=client_ip, + input_responses=input_responses, + request_state=request_state, + allow_input_required=allow_input_required, ) @@ -2582,7 +2610,10 @@ async def mcp_read_resource( oauth2_headers: dict[str, str] | None = None, raw_headers: dict[str, str] | None = None, client_ip: str | None = None, -) -> ReadResourceResult: + input_responses: InputResponses | None = None, + request_state: str | None = None, + allow_input_required: bool = False, +) -> ReadResourceResult | InputRequiredResult: """Read resource contents from upstream MCP servers.""" allowed_mcp_servers: Final = await _get_allowed_mcp_servers( @@ -2623,6 +2654,9 @@ async def mcp_read_resource( extra_headers=extra_headers, raw_headers=raw_headers, client_ip=client_ip, + input_responses=input_responses, + request_state=request_state, + allow_input_required=allow_input_required, ) @@ -2671,6 +2705,8 @@ async def _handle_managed_mcp_tool( guardrail_context: Mapping[str, object] | None = None, client_ip: str | None = None, wire_compat: WireCompat = WireCompat.LEGACY, + input_responses: InputResponses | None = None, + request_state: str | None = None, *, catalog_auth_header: str | None, ) -> CallToolResult | InputRequiredResult: @@ -2696,6 +2732,8 @@ async def _handle_managed_mcp_tool( litellm_logging_obj=litellm_logging_obj, guardrail_context=guardrail_context, wire_compat=wire_compat, + input_responses=input_responses, + request_state=request_state, ) verbose_logger.debug("CALL TOOL RESULT: %s", call_tool_result) return call_tool_result @@ -2917,6 +2955,8 @@ async def _execute_mcp_server_tool_call( client_ip=_client_ip, host_progress_callback=host_progress_callback, wire_compat=context.wire_compat, + input_responses=params.input_responses, + request_state=params.request_state, **data, # for logging ) except MCPMissingUserEnvVarsError as e: @@ -3008,7 +3048,7 @@ async def _execute_list_prompts( async def _execute_get_prompt( context: OperationContext, params: GetPromptRequestParams, host_progress_callback: ProgressCallback | None = None -) -> GetPromptResult: +) -> GetPromptResult | InputRequiredResult: if context.mcp_proxy_mode: _reject_mcp_proxy_operation() ( @@ -3032,6 +3072,9 @@ async def _execute_get_prompt( oauth2_headers=oauth2_headers, raw_headers=raw_headers, client_ip=_client_ip, + input_responses=params.input_responses, + request_state=params.request_state, + allow_input_required=context.wire_compat is WireCompat.MODERN, ) @@ -3093,7 +3136,7 @@ async def _execute_list_resource_templates( async def _execute_read_resource( context: OperationContext, params: ReadResourceRequestParams, host_progress_callback: ProgressCallback | None = None -) -> ReadResourceResult: +) -> ReadResourceResult | InputRequiredResult: if context.mcp_proxy_mode: _reject_mcp_proxy_operation() ( @@ -3115,6 +3158,9 @@ async def _execute_read_resource( oauth2_headers=oauth2_headers, raw_headers=raw_headers, client_ip=_client_ip, + input_responses=params.input_responses, + request_state=params.request_state, + allow_input_required=context.wire_compat is WireCompat.MODERN, ) return read_resource_result @@ -3177,6 +3223,17 @@ GatewayResult: TypeAlias = ( ) +def validate_continuation(operation: InteractionOperation, context: OperationContext) -> ContinuationState | None: + state: Final = open_continuation(operation, context, now=int(time.time())) + if state is not None: + target: Final = global_mcp_server_manager.get_mcp_server_by_id(state.target_id) + if target is None or target_digest(target) != state.target_digest: + raise MCPError(code=-32602, message="MCP continuation target changed; start a fresh request") + if set(operation.params.input_responses or {}) & set(state.gateway_responses or {}): + raise MCPError(code=-32602, message="Cannot replace gateway input responses") + return state + + class GatewayOperations: def __init__(self, host_progress_callback: ProgressCallback | None = None) -> None: self._host_progress_callback = host_progress_callback @@ -3201,7 +3258,9 @@ class GatewayOperations: async def execute(self, operation: ListPromptsRequest, context: OperationContext) -> ListPromptsResult: ... @overload - async def execute(self, operation: GetPromptRequest, context: OperationContext) -> GetPromptResult: ... + async def execute( + self, operation: GetPromptRequest, context: OperationContext + ) -> GetPromptResult | InputRequiredResult: ... @overload async def execute(self, operation: ListResourcesRequest, context: OperationContext) -> ListResourcesResult: ... @@ -3212,10 +3271,43 @@ class GatewayOperations: ) -> ListResourceTemplatesResult: ... @overload - async def execute(self, operation: ReadResourceRequest, context: OperationContext) -> ReadResourceResult: ... + async def execute( + self, operation: ReadResourceRequest, context: OperationContext + ) -> ReadResourceResult | InputRequiredResult: ... @catalog_operation(lambda: global_mcp_server_manager) async def execute(self, operation: GatewayOperation, context: OperationContext) -> GatewayResult: + if not isinstance(operation, (CallToolRequest, GetPromptRequest, ReadResourceRequest)): + return await self._execute(operation, context) + if context.wire_compat is not WireCompat.MODERN: + if operation.params.request_state is not None or operation.params.input_responses: + raise MCPError(code=-32602, message="Continuations require the modern MCP protocol") + return await self._execute(operation, context) + state: Final = validate_continuation(operation, context) + upstream: Final = operation.model_copy( + update={ + "params": operation.params.model_copy( + update={ + "request_state": state.upstream_state if state is not None else None, + "input_responses": { + **(state.gateway_responses or {}), + **(operation.params.input_responses or {}), + } + if state is not None + else None, + } + ) + } + ) + dispatch_context: Final = replace(context, mcp_servers=(state.target_id,)) if state is not None else context + result: Final = await self._execute(upstream, dispatch_context) + if isinstance(result, BoundInputRequiredResult): + return seal_continuation(result, operation, context, now=int(time.time()), previous=state) + if isinstance(result, InputRequiredResult): + raise MCPError(code=-32602, message="MCP continuation target is unavailable") + return result + + async def _execute(self, operation: GatewayOperation, context: OperationContext) -> GatewayResult: match operation: case DiscoverRequest(): listings: Final = ( @@ -3282,6 +3374,8 @@ class GatewayOperations: host_progress_callback=operation.host_progress_callback, guardrail_context=operation.guardrail_context, wire_compat=context.wire_compat, + input_responses=operation.input_responses, + request_state=operation.request_state, **operation.logging_data, ) case ListToolsRequest(params=params): diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index d1ac574dc91..dce792cfee1 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -131,12 +131,7 @@ def reject_disallowed_mcp_origin(request: StarletteRequest) -> None: def unsupported_protocol_version(scope: Scope) -> str | None: - """Return the unsupported ``MCP-Protocol-Version`` header value, if any. - - SDK 2's ``StreamableHTTPSessionManager`` routes any version outside - ``HANDSHAKE_PROTOCOL_VERSIONS`` to the modern single-exchange path, which - bypasses litellm's session/auth model, so the ASGI entry rejects it. - """ + """Admit configured HTTP revisions while keeping SSE on the legacy protocol.""" from litellm.proxy._experimental.mcp_server.capabilities import configured_versions headers: Final[Iterable[tuple[bytes, bytes]]] = scope.get("headers") or () @@ -144,6 +139,8 @@ def unsupported_protocol_version(scope: Scope) -> str | None: raw.decode("latin-1").strip() for key, raw in headers if key.lower() == _MCP_PROTOCOL_VERSION_HEADER ) for value in values: + if value == "2026-07-28" and scope.get("path", "").rstrip("/").endswith("/sse"): + return value if value and value not in configured_versions(): return value return None @@ -507,9 +504,11 @@ if MCP_AVAILABLE: from mcp.server.lowlevel.server import NotificationOptions from mcp.server.models import InitializationOptions from mcp.shared.exceptions import MCPError + from mcp.shared.inbound import InboundLadderRejection, classify_inbound_request, find_duplicated_routing_header from mcp.types import ( CallToolRequest, GetPromptRequest, + JSONRPCRequest, ListPromptsRequest, ListResourcesRequest, ListResourceTemplatesRequest, @@ -519,6 +518,7 @@ if MCP_AVAILABLE: from litellm.proxy._experimental.mcp_server import operations from litellm.proxy._experimental.mcp_server.contracts import OperationContext + from litellm.proxy._experimental.mcp_server.interactions import InteractionOperation from litellm.proxy._experimental.mcp_server.operations import ( _invalidate_byok_cred_cache, _mcp_session_id_from_headers, @@ -926,7 +926,9 @@ if MCP_AVAILABLE: verbose_logger.exception("Error in list_prompts endpoint: %s", exc) return ListPromptsResult(prompts=[]) - async def get_prompt(ctx: ServerRequestContext, params: GetPromptRequestParams) -> GetPromptResult: + async def get_prompt( + ctx: ServerRequestContext, params: GetPromptRequestParams + ) -> GetPromptResult | InputRequiredResult: if _mcp_proxy_mode.get(): _reject_mcp_proxy_operation() async with _legacy_operation_context(ctx, trace=False) as context: @@ -964,7 +966,9 @@ if MCP_AVAILABLE: verbose_logger.exception("Error in list_resource_templates endpoint: %s", exc) return ListResourceTemplatesResult(resource_templates=[]) - async def read_resource(ctx: ServerRequestContext, params: ReadResourceRequestParams) -> ReadResourceResult: + async def read_resource( + ctx: ServerRequestContext, params: ReadResourceRequestParams + ) -> ReadResourceResult | InputRequiredResult: if _mcp_proxy_mode.get(): _reject_mcp_proxy_operation() async with _legacy_operation_context(ctx, trace=False) as context: @@ -1389,6 +1393,8 @@ if MCP_AVAILABLE: async def _read_request_body_for_routing( receive: Receive, + *, + full_body: bool = False, ) -> tuple[list[Message], bytes]: """ Read just enough of the request body to decide whether this is a @@ -1401,7 +1407,8 @@ if MCP_AVAILABLE: The remainder of an oversized body is streamed lazily through ``wrapped_receive`` in the caller — so an authenticated client cannot force the proxy to buffer an arbitrarily large payload just to make a - routing decision. + routing decision. Modern interaction preflight requests the full body, + which the SDK's single-exchange transport also requires. """ consumed_messages: Final[list[Message]] = [] body_chunks: Final[list[bytes]] = [] @@ -1409,6 +1416,12 @@ if MCP_AVAILABLE: while True: message = await receive() + if ( + full_body + and peeked_bytes + len(message.get("body", b"") or b"") + > session_manager_stateless.max_request_body_size + ): + raise HTTPException(status_code=413, detail="Request body too large") consumed_messages.append(message) if message.get("type") != "http.request": @@ -1422,7 +1435,7 @@ if MCP_AVAILABLE: # handler via ``consumed_messages``, but ``body_chunks`` is # purely for the JSON-RPC method check — there is no reason # to copy a large body frame into a second buffer. - remaining = _MCP_ROUTING_PEEK_MAX_BYTES - peeked_bytes + remaining = len(body) if full_body else _MCP_ROUTING_PEEK_MAX_BYTES - peeked_bytes if remaining > 0: body_chunks.append(body[:remaining]) peeked_bytes += min(len(body), remaining) @@ -1430,7 +1443,7 @@ if MCP_AVAILABLE: if not message.get("more_body", False): break - if peeked_bytes >= _MCP_ROUTING_PEEK_MAX_BYTES: + if not full_body and peeked_bytes >= _MCP_ROUTING_PEEK_MAX_BYTES: # Stop draining; downstream replay will pull remaining chunks # directly from the original `receive` via wrapped_receive. break @@ -1997,6 +2010,59 @@ if MCP_AVAILABLE: detail="Forbidden", ) + _INTERACTION_REQUEST: Final[TypeAdapter[InteractionOperation]] = TypeAdapter(InteractionOperation) + + @catalog_operation(lambda: operations.global_mcp_server_manager) + async def _preflight_modern_interaction( + scope: Scope, body: bytes, context: OperationContext + ) -> JSONResponse | None: + try: + envelope: Final = JSONRPCRequest.model_validate_json(body) + operation: Final = _INTERACTION_REQUEST.validate_json(body) + except ValidationError: + return None + headers: Final = StarletteRequest(scope).headers + if find_duplicated_routing_header(headers.items()) is not None or isinstance( + classify_inbound_request(envelope.model_dump(by_alias=True), headers=dict(headers)), InboundLadderRejection + ): + return None + try: + state: Final = operations.validate_continuation(operation, context) + except MCPError as error: + return JSONResponse( + status_code=400, + content={"jsonrpc": "2.0", "id": envelope.id, "error": error.error.model_dump(exclude_none=True)}, + ) + if state is None and context.mcp_servers is None: + return None + targets: Final = ( + [state.target_id] + if state is not None + else list(context.mcp_servers) + if context.mcp_servers is not None + else None + ) + allowed: Final = await operations._get_allowed_mcp_servers( + user_api_key_auth=context.user_api_key_auth, mcp_servers=targets, client_ip=context.client_ip + ) + if state is not None and not any(target.server_id == state.target_id for target in allowed): + raise HTTPException(status_code=403, detail="MCP continuation target is no longer authorized") + authorized_names: Final = [target.alias or target.name for target in allowed] + await _raise_preemptive_401_for_unauthenticated_servers( + scope=scope, + mcp_servers=list(context.mcp_servers) if context.mcp_servers is not None else authorized_names, + oauth2_headers=dict(context.oauth2_headers) if context.oauth2_headers is not None else None, + mcp_server_auth_headers={key: dict(value) for key, value in context.mcp_server_auth_headers.items()} + if context.mcp_server_auth_headers is not None + else None, + user_api_key_auth=context.user_api_key_auth, + client_ip=context.client_ip, + allowed_server_ids={target.server_id for target in allowed}, + raw_headers=context.raw_headers, + ) + await _check_passthrough_upstream_auth(scope, context.user_api_key_auth, authorized_names, context.client_ip) + return None + async def handle_streamable_http_mcp(scope: Scope, receive: Receive, send: Send) -> None: """Handle MCP requests through StreamableHTTP.""" try: @@ -2027,6 +2093,10 @@ if MCP_AVAILABLE: ) = await extract_mcp_auth_context(scope, path) reject_disallowed_mcp_client(StarletteRequest(scope).headers, user_api_key_auth) scoped_server_endpoint: Final = len(_get_mcp_servers_in_path(path) or []) == 1 + request_headers: Final = StarletteRequest(scope).headers + defer_upstream_probes: Final = request_headers.get( + "mcp-protocol-version" + ) == "2026-07-28" and request_headers.get("mcp-method") in {"tools/call", "prompts/get", "resources/read"} # Extract client IP for MCP access control _client_ip: Final = IPAddressUtils.get_mcp_client_ip(StarletteRequest(scope)) @@ -2055,21 +2125,22 @@ if MCP_AVAILABLE: # from the fully-authorized server set: a passthrough server that # the active toolset excludes should not trigger an OAuth flow # for a server the caller will be 403'd on after authentication. - await _raise_preemptive_401_for_unauthenticated_servers( - scope=scope, - mcp_servers=mcp_servers, - oauth2_headers=oauth2_headers, - mcp_server_auth_headers=mcp_server_auth_headers, - user_api_key_auth=user_api_key_auth, - client_ip=_client_ip, - allowed_server_ids=toolset_allowed_server_ids, - raw_headers=raw_headers, - ) + if not defer_upstream_probes: + await _raise_preemptive_401_for_unauthenticated_servers( + scope=scope, + mcp_servers=mcp_servers, + oauth2_headers=oauth2_headers, + mcp_server_auth_headers=mcp_server_auth_headers, + user_api_key_auth=user_api_key_auth, + client_ip=_client_ip, + allowed_server_ids=toolset_allowed_server_ids, + raw_headers=raw_headers, + ) - # Pre-flight auth check for pass-through servers. Must run after - # toolset scoping so the probe list is derived from the fully-authorized - # server set, not the raw user-supplied names. - await _check_passthrough_upstream_auth(scope, user_api_key_auth, mcp_servers, _client_ip) + # Pre-flight auth check for pass-through servers. Must run after + # toolset scoping so the probe list is derived from the fully-authorized + # server set, not the raw user-supplied names. + await _check_passthrough_upstream_auth(scope, user_api_key_auth, mcp_servers, _client_ip) # Inject masked debug headers when client sends x-litellm-mcp-debug: true _debug_headers: Final = MCPDebug.maybe_build_debug_headers( @@ -2138,7 +2209,24 @@ if MCP_AVAILABLE: body = b"" if scope.get("method") == "POST": - consumed_messages, body = await _read_request_body_for_routing(receive) + consumed_messages, body = await _read_request_body_for_routing(receive, full_body=defer_upstream_probes) + if defer_upstream_probes: + rejection: Final = await _preflight_modern_interaction( + scope, + body, + 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, + ), + ) + if rejection is not None: + await rejection(scope, receive, send) + return is_initialize = _is_initialize_request(body) use_stateful: Final = bool(session_id or is_initialize) @@ -2279,12 +2367,16 @@ if MCP_AVAILABLE: client_info=_extract_initialize_client_info(body), ) - async with _gateway_initialize_instructions_request_scope( - user_api_key_auth, - mcp_servers, - _client_ip, - scoped_server_endpoint=scoped_server_endpoint, - is_initialize=is_initialize, + async with ( + contextlib.nullcontext() + if defer_upstream_probes + else _gateway_initialize_instructions_request_scope( + user_api_key_auth, + mcp_servers, + _client_ip, + scoped_server_endpoint=scoped_server_endpoint, + is_initialize=is_initialize, + ) ): await target_manager.handle_request(scope, receive, local_send) if use_stateful and session_id and scope.get("method") == "DELETE": diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index bed2503d211..d46dd3ab3e9 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -3206,7 +3206,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): mcp_advertised_versions: MCPAdvertisedVersions | None = Field( None, description="MCP revisions enabled by the gateway. Defaults to all completed legacy revisions. " - "Modern protocol serving and Apps/Tasks remain disabled.", + "Modern protocol serving requires explicit opt-in. Apps/Tasks remain disabled.", ) mcp_allowed_clients: list[MCPAllowedClient] | None = Field( None, diff --git a/litellm/types/mcp.py b/litellm/types/mcp.py index 556cb6712ae..e104f74faf7 100644 --- a/litellm/types/mcp.py +++ b/litellm/types/mcp.py @@ -4,7 +4,7 @@ import enum import re from collections.abc import Awaitable, Callable, Mapping from types import MappingProxyType -from typing import TYPE_CHECKING, Annotated, Any, Final, Literal +from typing import TYPE_CHECKING, Annotated, Any, Final, Literal, TypeAlias from urllib.parse import urlsplit import httpx @@ -71,7 +71,8 @@ def validate_mcp_protocol_transport(protocol_version: MCPUpstreamProtocol, trans raise ValueError("Modern MCP requires HTTP or stdio transport") -MCPAdvertisedVersions = Annotated[tuple[MCPLegacyVersion, ...], Field(min_length=1)] +MCPAdvertisedVersion: TypeAlias = MCPLegacyVersion | Literal["2026-07-28"] +MCPAdvertisedVersions: TypeAlias = Annotated[tuple[MCPAdvertisedVersion, ...], Field(min_length=1)] MCPSpecVersionType = Literal[ MCPSpecVersion.nov_2024, MCPSpecVersion.mar_2025, diff --git a/tests/integration/mcp/test_interactions.py b/tests/integration/mcp/test_interactions.py new file mode 100644 index 00000000000..42af8cf5dba --- /dev/null +++ b/tests/integration/mcp/test_interactions.py @@ -0,0 +1,395 @@ +import json +import os +from pathlib import Path +from typing import Final + +import httpx +import pytest +import yaml +from pydantic import JsonValue + +from integration._support.client import Gateway +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, wire_server + + +_KEY: Final = "sk-interaction-test" +_PROXY_PYTHONPATH: Final = os.pathsep.join( + (str(Path(__file__).resolve().parents[3]), str(Path(__file__).resolve().parents[2])) +) +_META: Final = { + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientInfo": {"name": "continuation-test", "version": "1"}, + "io.modelcontextprotocol/clientCapabilities": {"elicitation": {"form": {}, "url": {}}}, +} + + +def interaction_peer(request: Request) -> Reply: + if request.method != "POST": + return Reply(status=405) + body: Final = json.loads(request.body) + method: Final = body["method"] + params: Final = body.get("params", {}) + if method == "server/discover": + result = { + "supportedVersions": ["2026-07-28"], + "capabilities": {"tools": {}, "prompts": {}, "resources": {}}, + "cacheScope": "private", + "ttlMs": 0, + } + elif method == "tools/list": + result = { + "tools": [{"name": "confirm", "inputSchema": {"type": "object"}}], + "cacheScope": "private", + "ttlMs": 0, + } + elif method == "prompts/list": + result = {"prompts": [{"name": "confirm"}], "cacheScope": "private", "ttlMs": 0} + elif method == "resources/list": + result = {"resources": [{"name": "confirm", "uri": "test://confirm"}], "cacheScope": "private", "ttlMs": 0} + elif method == "resources/templates/list": + result = {"resourceTemplates": [], "cacheScope": "private", "ttlMs": 0} + elif not params.get("requestState"): + return Reply( + body=json.dumps( + { + "jsonrpc": "2.0", + "id": body["id"], + "result": { + "resultType": "input_required", + "requestState": "opaque:" + method, + "inputRequests": { + "consent": { + "method": "elicitation/create", + "params": { + "mode": "form", + "message": "Confirm", + "requestedSchema": {"type": "object", "properties": {}}, + }, + } + }, + }, + } + ).encode() + ) + else: + assert params["requestState"] == "opaque:" + method + assert params["inputResponses"] == {"consent": {"action": "accept"}} + if method == "tools/call": + result = {"content": [{"type": "text", "text": "confirmed"}], "isError": False} + elif method == "prompts/get": + result = {"messages": [{"role": "user", "content": {"type": "text", "text": "confirmed"}}]} + else: + assert method == "resources/read" + result = {"contents": [{"uri": "test://confirm", "text": "confirmed"}]} + return Reply( + body=json.dumps({"jsonrpc": "2.0", "id": body["id"], "result": {"resultType": "complete", **result}}).encode() + ) + + +def rpc(gateway: Gateway, method: str, params: dict[str, JsonValue], *, status: int = 200) -> dict: + response: Final = gateway.client.post( + "/mcp/", + headers={ + "Authorization": "Bearer " + gateway.key, + "MCP-Protocol-Version": "2026-07-28", + "Mcp-Method": method, + "Mcp-Name": str(params.get("name", params.get("uri", ""))), + "Accept": "application/json, text/event-stream", + }, + json={"jsonrpc": "2.0", "id": 1, "method": method, "params": {**params, "_meta": _META}}, + ) + assert response.status_code == status, response.text + return response.json() + + +@pytest.mark.parametrize("changed_target", [False, True]) +def test_continuations_resume_on_another_replica_and_reject_changed_operations( + tmp_path: Path, changed_target: bool +) -> None: + with wire_server(interaction_peer) as peer, httpx.Client() as client: + config: Final = tmp_path / "proxy.yaml" + config.write_text( + yaml.safe_dump( + { + "model_list": [], + "mcp_servers": { + "mrtr": { + "url": peer.url, + "transport": "http", + "protocol_version": "2026-07-28", + "allow_elicitation": True, + } + }, + "general_settings": { + "master_key": _KEY, + "store_model_in_db": False, + "mcp_advertised_versions": ["2025-11-25", "2026-07-28"], + }, + } + ) + ) + other_config: Final = tmp_path / "other-proxy.yaml" + other_config.write_text( + config.read_text().replace(peer.url, peer.url + "/changed-target") if changed_target else config.read_text() + ) + seed: Final = Gateway(client, _KEY, peer.url) + environment: Final = { + "PYTHONPATH": _PROXY_PYTHONPATH, + "STORE_MODEL_IN_DB": "False", + "DISABLE_SCHEMA_UPDATE": "true", + "LITELLM_SALT_KEY": "shared-interaction-test", + } + options: Final = { + "database_setup": (), + "remove_environment": ( + "DATABASE_URL", + "DATABASE_URL_READ_REPLICA", + "LITELLM_LICENSE", + "LITELLM_LICENSE_PATH", + "REDIS_URL", + "REDIS_HOST", + ), + } + with ( + owned_proxy(seed, tmp_path / "a", environment, config=config, **options) as first, + owned_proxy(seed, tmp_path / "b", environment, config=other_config, **options) as second, + ): + for method, params, terminal_field in ( + ("tools/call", {"name": "mrtr-confirm", "arguments": {}}, "content"), + ("prompts/get", {"name": "mrtr-confirm", "arguments": {}}, "messages"), + ("resources/read", {"uri": "test://confirm"}, "contents"), + ): + initial: Final = rpc(first, method, params) + assert initial["result"]["resultType"] == "input_required", initial + state: Final = initial["result"]["requestState"] + assert state.startswith("mcp_state_v1."), initial + retry: Final = {**params, "requestState": state, "inputResponses": {"consent": {"action": "accept"}}} + if changed_target: + peer.drain() + refused: Final = rpc(second, method, retry, status=400) + assert refused["error"]["code"] == -32602, refused + assert peer.drain() == (), "Changed target must reject before upstream dispatch" + continue + completed: Final = rpc(second, method, retry) + assert "confirmed" in json.dumps(completed["result"][terminal_field]), completed + assert rpc(first, method, retry) == completed + peer.drain() + changed: Final = { + **retry, + **({"uri": "test://other"} if method == "resources/read" else {"name": "mrtr-other"}), + } + rejected: Final = rpc(second, method, changed, status=400) + assert rejected["error"]["code"] == -32602, rejected + assert peer.drain() == (), "Rejected continuation must not contact the upstream" + + +def test_continuation_reauthenticates_caller_and_rechecks_revoked_permissions(tmp_path: Path) -> None: + auth_module: Final = tmp_path / "interaction_auth.py" + auth_module.write_text( + "from pathlib import Path\n" + "from fastapi import HTTPException, Request\n" + "from litellm.proxy._types import UserAPIKeyAuth\n" + "async def authenticate(request: Request, api_key: str) -> UserAPIKeyAuth:\n" + " if api_key != 'sk-interaction-test':\n" + " raise HTTPException(status_code=401, detail='Unknown test caller')\n" + " return UserAPIKeyAuth.model_validate_json(Path(__file__).with_suffix('.json').read_text())\n" + ) + identity: Final = { + "user_id": "alice", + "team_id": "team-a", + "user_role": "internal_user", + "object_permission": {"object_permission_id": "test-permission", "mcp_servers": ["interaction-server"]}, + } + auth_state: Final = auth_module.with_suffix(".json") + auth_state.write_text(json.dumps(identity)) + with wire_server(interaction_peer) as peer, wire_server(interaction_peer) as other_peer, httpx.Client() as client: + config: Final = tmp_path / "proxy.yaml" + config.write_text( + yaml.safe_dump( + { + "model_list": [], + "mcp_servers": { + "mrtr": { + "server_id": "interaction-server", + "url": peer.url, + "transport": "http", + "protocol_version": "2026-07-28", + "allow_elicitation": True, + }, + "other": { + "server_id": "other-server", + "url": other_peer.url, + "transport": "http", + "protocol_version": "2026-07-28", + "allow_elicitation": True, + }, + }, + "general_settings": { + "master_key": _KEY, + "custom_auth": "interaction_auth.authenticate", + "store_model_in_db": False, + "mcp_advertised_versions": ["2025-11-25", "2026-07-28"], + }, + } + ) + ) + with owned_proxy( + Gateway(client, _KEY, peer.url), + tmp_path / "proxy", + { + "PYTHONPATH": _PROXY_PYTHONPATH, + "STORE_MODEL_IN_DB": "False", + "DISABLE_SCHEMA_UPDATE": "true", + "LITELLM_SALT_KEY": "caller-test-salt", + }, + config=config, + database_setup=(), + remove_environment=("DATABASE_URL", "DATABASE_URL_READ_REPLICA", "REDIS_URL", "REDIS_HOST"), + ) as gateway: + params: Final = {"name": "mrtr-confirm", "arguments": {}} + initial: Final = rpc(gateway, "tools/call", params) + assert initial["result"]["resultType"] == "input_required", initial + retry: Final = { + **params, + "requestState": initial["result"]["requestState"], + "inputResponses": {"consent": {"action": "accept"}}, + } + for changed in ({"user_id": "bob"}, {"team_id": "team-b"}): + auth_state.write_text(json.dumps({**identity, **changed})) + peer.drain() + rejected: Final = rpc(gateway, "tools/call", retry, status=400) + assert rejected["error"]["code"] == -32602, rejected + assert peer.drain() == (), "Caller-bound state must reject before contacting upstream" + auth_state.write_text(json.dumps(identity)) + resumed: Final = rpc(gateway, "tools/call", retry) + assert resumed["result"]["content"] == [{"type": "text", "text": "confirmed"}], resumed + auth_state.write_text( + json.dumps( + { + **identity, + "object_permission": { + "object_permission_id": "test-permission", + "mcp_servers": ["denied-server"], + }, + } + ) + ) + peer.drain() + revoked: Final = rpc(gateway, "tools/call", retry, status=403) + assert revoked["detail"] == "MCP continuation target is no longer authorized", revoked + assert peer.drain() == (), "Revoked access must reject a valid continuation before upstream dispatch" + + auth_state.write_text(json.dumps(identity)) + resource: Final = rpc(gateway, "resources/read", {"uri": "test://confirm"}) + assert resource["result"]["resultType"] == "input_required", resource + auth_state.write_text( + json.dumps( + { + **identity, + "object_permission": { + "object_permission_id": "test-permission", + "mcp_servers": ["other-server"], + }, + } + ) + ) + peer.drain() + other_peer.drain() + response: Final = gateway.client.post( + "/mcp/", + headers={ + "Authorization": "Bearer " + gateway.key, + "MCP-Protocol-Version": "2026-07-28", + "Mcp-Method": "resources/read", + "Mcp-Name": "test://confirm", + "Accept": "application/json, text/event-stream", + }, + json={ + "jsonrpc": "2.0", + "id": 1, + "method": "resources/read", + "params": { + "uri": "test://confirm", + "_meta": _META, + "requestState": resource["result"]["requestState"], + "inputResponses": {"consent": {"action": "accept"}}, + }, + }, + ) + other_requests: Final = tuple( + (json.loads(request.body)["method"], json.loads(request.body).get("params", {}).get("requestState")) + for request in other_peer.drain() + ) + assert other_requests == (), "Continuation must never send upstream state to another authorized server" + assert peer.drain() == (), "Revoked original target must not receive a retry" + assert response.status_code >= 400 or "error" in response.json(), response.text + + auth_state.write_text( + json.dumps( + { + **identity, + "object_permission": { + "object_permission_id": "test-permission", + "mcp_servers": ["interaction-server", "other-server"], + }, + } + ) + ) + completed: Final = rpc( + gateway, + "resources/read", + { + "uri": "test://confirm", + "requestState": resource["result"]["requestState"], + "inputResponses": {"consent": {"action": "accept"}}, + }, + ) + assert completed["result"]["contents"] == [{"uri": "test://confirm", "text": "confirmed"}], completed + assert other_peer.drain() == (), "Expanded access must keep the continuation on its original target" + + +def test_missing_continuation_key_reports_configuration_for_each_carrier(tmp_path: Path) -> None: + with wire_server(interaction_peer) as peer, httpx.Client() as client: + config: Final = tmp_path / "proxy.yaml" + config.write_text( + yaml.safe_dump( + { + "model_list": [], + "mcp_servers": { + "mrtr": { + "url": peer.url, + "transport": "http", + "protocol_version": "2026-07-28", + "allow_elicitation": True, + } + }, + "general_settings": { + "master_key": _KEY, + "store_model_in_db": False, + "mcp_advertised_versions": ["2025-11-25", "2026-07-28"], + }, + } + ) + ) + with owned_proxy( + Gateway(client, _KEY, peer.url), + tmp_path / "proxy", + { + "PYTHONPATH": _PROXY_PYTHONPATH, + "STORE_MODEL_IN_DB": "False", + "DISABLE_SCHEMA_UPDATE": "true", + "LITELLM_SALT_KEY": "", + }, + config=config, + database_setup=(), + remove_environment=("DATABASE_URL", "DATABASE_URL_READ_REPLICA", "REDIS_URL", "REDIS_HOST"), + ) as gateway: + for method, params in ( + ("tools/call", {"name": "mrtr-confirm", "arguments": {}}), + ("prompts/get", {"name": "mrtr-confirm", "arguments": {}}), + ("resources/read", {"uri": "test://confirm"}), + ): + response: Final = rpc(gateway, method, params, status=400) + assert response["error"]["code"] == -32602, response + assert "LITELLM_SALT_KEY" in response["error"]["message"], response diff --git a/tests/integration/run.py b/tests/integration/run.py index f7b2dead197..aa5ead066f2 100644 --- a/tests/integration/run.py +++ b/tests/integration/run.py @@ -23,7 +23,9 @@ GROUPS: Final = MappingProxyType( "security": ("security",), } ) -GITHUB_FILES: Final = frozenset({"tests/integration/database/test_roi_observed.py"}) +GITHUB_FILES: Final = frozenset( + {"tests/integration/database/test_roi_observed.py", "tests/integration/mcp/test_interactions.py"} +) @dataclass(frozen=True, slots=True) diff --git a/tests/unit/experimental_mcp_client/test_mcp_client.py b/tests/unit/experimental_mcp_client/test_mcp_client.py index 9a30e1990b1..ff192c9279c 100644 --- a/tests/unit/experimental_mcp_client/test_mcp_client.py +++ b/tests/unit/experimental_mcp_client/test_mcp_client.py @@ -61,6 +61,199 @@ from litellm.types.mcp_server.mcp_server_manager import MCPServer _JSONRPC_MESSAGE_ADAPTER: Final = TypeAdapter(JSONRPCMessage) +@pytest.mark.asyncio +async def test_tool_continuation_reaches_modern_upstream() -> None: + from mcp.types import ElicitResult, TextContent + + def respond(request: httpx2.Request) -> httpx2.Response: + payload: Final = _JSONRPC_MESSAGE_ADAPTER.validate_json(request.content) + assert isinstance(payload, JSONRPCRequest) + if payload.method == "server/discover": + return httpx2.Response( + 200, + json={ + "jsonrpc": "2.0", + "id": payload.id, + "result": { + "supportedVersions": ["2026-07-28"], + "capabilities": {"tools": {}}, + "resultType": "complete", + "cacheScope": "private", + "ttlMs": 0, + }, + }, + ) + if payload.method == "tools/list": + return httpx2.Response( + 200, + json={ + "jsonrpc": "2.0", + "id": payload.id, + "result": { + "tools": [{"name": "confirm", "inputSchema": {"type": "object"}}], + "resultType": "complete", + "cacheScope": "private", + "ttlMs": 0, + }, + }, + ) + assert payload.method == "tools/call" + assert payload.params is not None + assert payload.params.get("requestState") == "opaque-upstream-state" + assert payload.params.get("inputResponses") == {"confirmation": {"action": "accept"}} + return httpx2.Response( + 200, + json={ + "jsonrpc": "2.0", + "id": payload.id, + "result": { + "resultType": "complete", + "content": [{"type": "text", "text": "confirmed"}], + "isError": False, + }, + }, + ) + + client: Final = _MockTransportClient(respond, server_url="https://example.com/mcp", protocol_version="2026-07-28") + result: Final = await client.call_tool( + CallToolRequestParams( + name="confirm", + request_state="opaque-upstream-state", + input_responses={"confirmation": ElicitResult(action="accept")}, + ), + raise_on_error=True, + ) + assert isinstance(result, CallToolResult) + assert result.content == [TextContent(type="text", text="confirmed")] + assert result.is_error is False + + +@pytest.mark.asyncio +@pytest.mark.parametrize("modern_caller", [False, True]) +@pytest.mark.parametrize( + "elicitation_mode,sampling", [("form", False), ("form", True), ("url", False), ("url", True), ("none", True)] +) +async def test_modern_input_request_uses_existing_elicitation_callback( + modern_caller: bool, sampling: bool, elicitation_mode: str +) -> None: + from queue import SimpleQueue + from mcp.types import ( + ElicitResult, + ElicitRequestParams, + TextContent, + CreateMessageRequestParams, + CreateMessageResult, + ) + from litellm.proxy._experimental.mcp_server.interactions import BoundInputRequiredResult + + observed: Final[SimpleQueue[str]] = SimpleQueue() + sampled: Final = CreateMessageResult( + role="assistant", content=TextContent(type="text", text="sampled"), model="test" + ) + expected_responses: Final = { + **({"consent": {"action": "accept"}} if elicitation_mode != "none" else {}), + **({"sample": sampled.model_dump(by_alias=True, exclude_none=True)} if sampling else {}), + } + + async def elicit(context: object, params: ElicitRequestParams) -> ElicitResult: + if params.mode == "url": + assert params.elicitation_id, "Legacy URL input must carry its required elicitation ID" + observed.put(params.message) + return ElicitResult(action="accept") + + async def sample(context: object, params: CreateMessageRequestParams) -> CreateMessageResult: + assert params.max_tokens == 10 + observed.put("sample") + return sampled + + def respond(request: httpx2.Request) -> httpx2.Response: + payload: Final = _JSONRPC_MESSAGE_ADAPTER.validate_json(request.content) + assert isinstance(payload, JSONRPCRequest) + if payload.method == "server/discover": + result = { + "supportedVersions": ["2026-07-28"], + "capabilities": {"tools": {}}, + "resultType": "complete", + "cacheScope": "private", + "ttlMs": 0, + } + elif payload.method == "tools/list": + result = { + "tools": [{"name": "confirm", "inputSchema": {"type": "object"}}], + "resultType": "complete", + "cacheScope": "private", + "ttlMs": 0, + } + elif not (payload.params or {}).get("requestState"): + result = { + "resultType": "input_required", + "requestState": "pending", + "inputRequests": { + **( + { + "consent": { + "method": "elicitation/create", + "params": { + "message": "Confirm operation", + **( + {"mode": "form", "requestedSchema": {"type": "object", "properties": {}}} + if elicitation_mode == "form" + else {"mode": "url", "url": "https://example.com/confirm"} + ), + }, + } + } + if elicitation_mode != "none" + else {} + ), + **( + {"sample": {"method": "sampling/createMessage", "params": {"messages": [], "maxTokens": 10}}} + if sampling + else {} + ), + }, + } + else: + assert (payload.params or {}).get("requestState") == "pending" + assert (payload.params or {}).get("inputResponses") == expected_responses + result = {"resultType": "complete", "content": [{"type": "text", "text": "confirmed"}], "isError": False} + return httpx2.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": result}) + + client: Final = _MockTransportClient( + respond, + server_url="https://example.com/mcp", + protocol_version="2026-07-28", + elicitation_callback=elicit, + sampling_callback=sample, + ) + result: Final = await client.call_tool( + CallToolRequestParams(name="confirm", arguments={}), raise_on_error=True, allow_input_required=modern_caller + ) + if modern_caller and elicitation_mode != "none": + assert isinstance(result, BoundInputRequiredResult) + assert tuple(result.input_requests or {}) == ("consent",) + assert result.gateway_responses == ({"sample": sampled} if sampling else None) + resumed: Final = await client.call_tool( + CallToolRequestParams( + name="confirm", + arguments={}, + request_state=result.request_state, + input_responses={"consent": ElicitResult(action="accept"), **(result.gateway_responses or {})}, + ), + raise_on_error=True, + allow_input_required=True, + ) + assert isinstance(resumed, CallToolResult) + assert resumed.content == [TextContent(type="text", text="confirmed")] + else: + assert isinstance(result, CallToolResult) + assert result.content == [TextContent(type="text", text="confirmed")] + assert sorted(observed.get_nowait() for _ in range(observed.qsize())) == sorted( + ([] if modern_caller or elicitation_mode == "none" else ["Confirm operation"]) + + (["sample"] if sampling else []) + ) + + def _initialized(instructions: str | None = None) -> InitializeResult: return InitializeResult( protocol_version=LATEST_HANDSHAKE_VERSION, @@ -3615,3 +3808,79 @@ async def test_optional_discovery_retains_freshness_across_pages( assert result.next_cursor is None assert result.ttl_ms == max(0, ttl - cleanup_seconds * 1000) assert len(await getattr(client, "list_" + kind)(raise_on_error=True)) == 2 + + +def test_prompt_continuation_polling_respects_the_original_deadline() -> None: + from mcp.types import GetPromptRequestParams + + def respond(request: httpx2.Request) -> httpx2.Response: + payload: Final = _JSONRPC_MESSAGE_ADAPTER.validate_json(request.content) + assert isinstance(payload, JSONRPCRequest) + result: Final = ( + { + "supportedVersions": ["2026-07-28"], + "capabilities": {"prompts": {}}, + "resultType": "complete", + "cacheScope": "private", + "ttlMs": 0, + } + if payload.method == "server/discover" + else {"resultType": "input_required", "requestState": "pending"} + ) + return httpx2.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": result}) + + loop: Final = _AutojumpClockLoop() + client: Final = _MockTransportClient( + respond, server_url="https://example.com/mcp", protocol_version="2026-07-28", timeout=0.12 + ) + try: + with pytest.raises(TimeoutError): + loop.run_until_complete(client.get_prompt(GetPromptRequestParams(name="pending"))) + assert loop.time() == pytest.approx(0.12) + finally: + loop.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", ("form", "url")) +@pytest.mark.parametrize("enabled", (False, True)) +async def test_modern_elicitation_honors_server_permission(mode: str, enabled: bool) -> None: + import anyio + from mcp import ClientSession, MCPError + from mcp.shared.message import SessionMessage + from mcp.types import ( + ElicitRequest, + ElicitRequestFormParams, + ElicitRequestParams, + ElicitRequestURLParams, + ElicitResult, + InputRequiredResult, + ) + + async def elicit(context: object, params: ElicitRequestParams) -> ElicitResult: + return ElicitResult(action="accept") + + client: Final = MCPClient(server_url="https://example.com/mcp", elicitation_callback=elicit if enabled else None) + request: Final = AsyncMock( + return_value=InputRequiredResult( + request_state="pending", + input_requests={ + "consent": ElicitRequest( + params=ElicitRequestFormParams(message="Confirm", requested_schema={"type": "object"}) + if mode == "form" + else ElicitRequestURLParams(message="Confirm", url="https://example.com/confirm") + ) + }, + ) + ) + send, receive = anyio.create_memory_object_stream[SessionMessage](1) + async with send, receive: + session: Final = ClientSession(receive, send) + if enabled: + result: Final = await client._request_with_interaction(session, request, None, None, True) + assert isinstance(result, InputRequiredResult) + assert result.input_requests["consent"].params.mode == mode + else: + with pytest.raises(MCPError, match="Elicitation is disabled"): + await client._request_with_interaction(session, request, None, None, True) + request.assert_awaited_once_with(None, None) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_capabilities.py b/tests/unit/proxy/_experimental/mcp_server/test_capabilities.py index f104e655704..f8e3655b2f2 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_capabilities.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_capabilities.py @@ -33,7 +33,7 @@ def test_discovery_only_exposes_authorized_completed_support(revision, transport client_extensions=frozenset({"io.modelcontextprotocol/ui"}), upstream_extensions=frozenset({"io.modelcontextprotocol/ui"}), ) - assert result.supported_versions == [revision] + assert result.supported_versions == ([revision] if transport is MCPTransport.sse else [revision, "2026-07-28"]) assert result.capabilities.tools is not None assert result.capabilities.prompts is None assert result.capabilities.resources is None @@ -43,7 +43,7 @@ def test_discovery_only_exposes_authorized_completed_support(revision, transport assert result.ttl_ms == 0 -@pytest.mark.parametrize("upstream", [frozenset(), frozenset({"unknown"}), frozenset({"2026-07-28"})]) +@pytest.mark.parametrize("upstream", [frozenset(), frozenset({"unknown"})]) def test_unproven_translation_never_advertises_operations(upstream): result = build_discovery( configured=HANDSHAKE_PROTOCOL_VERSIONS, @@ -82,12 +82,12 @@ def test_discovery_results_do_not_share_mutable_capabilities(): assert capabilities.tools.list_changed is not True -def test_modern_candidates_do_not_enable_public_serving(): - modern = REVISION_SUPPORT["2026-07-28"] - assert modern.completed is False +def test_modern_support_excludes_legacy_sse() -> None: + modern: Final = REVISION_SUPPORT["2026-07-28"] + assert modern.completed is True assert "input_required" in modern.results assert MCPTransport.sse not in modern.transports - assert not any("2026-07-28" in pair for pair in TRANSLATION_PAIRS) + assert ("2026-07-28", "2026-07-28") in TRANSLATION_PAIRS @pytest.mark.asyncio @@ -104,3 +104,39 @@ async def test_version_policy_gates_the_actual_sdk_handshake(versions, accepted) with pytest.RaisesGroup(pytest.RaisesExc(MCPError, match="Unsupported MCP protocol version"), flatten_subgroups=True): async with Client(server, mode="legacy"): pytest.fail("The excluded revision must not initialize") + + +def test_modern_discovery_requires_opt_in_and_keeps_unsupported_features_disabled() -> None: + from pydantic import TypeAdapter + from litellm.types.mcp import MCPAdvertisedVersions + + configured: Final = TypeAdapter(MCPAdvertisedVersions).validate_python(["2026-07-28"]) + result: Final = build_discovery( + configured=configured, + revision="2026-07-28", + transport=MCPTransport.http, + authorized_operations=GATEWAY_OPERATIONS, + upstream_versions=frozenset({"2026-07-28"}), + capabilities=ServerCapabilities( + tools=ToolsCapability(), prompts=PromptsCapability(), resources=ResourcesCapability() + ), + ) + assert result.supported_versions == ["2026-07-28"] + assert result.capabilities.tools is not None + assert result.capabilities.prompts is not None + assert result.capabilities.resources is not None + assert result.capabilities.tasks is None + assert result.capabilities.extensions is None + + +@pytest.mark.parametrize("path", ["/mcp/sse", "/mcp/example/sse/"]) +def test_modern_protocol_is_rejected_on_legacy_sse_paths(path: str, monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy import proxy_server + from litellm.proxy._experimental.mcp_server.server import unsupported_protocol_version + + monkeypatch.setitem(proxy_server.general_settings, "mcp_advertised_versions", ["2025-11-25", "2026-07-28"]) + assert unsupported_protocol_version({"path": "/mcp", "headers": [(b"mcp-protocol-version", b"2026-07-28")]}) is None + assert ( + unsupported_protocol_version({"path": path, "headers": [(b"mcp-protocol-version", b"2026-07-28")]}) + == "2026-07-28" + ) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_interactions.py b/tests/unit/proxy/_experimental/mcp_server/test_interactions.py new file mode 100644 index 00000000000..8286269525a --- /dev/null +++ b/tests/unit/proxy/_experimental/mcp_server/test_interactions.py @@ -0,0 +1,253 @@ +from typing import Final, Literal + +import pytest +from mcp import MCPError +from mcp.types import CallToolRequest, CallToolRequestParams, InputRequiredResult + +from litellm.proxy._experimental.mcp_server.contracts import OperationContext +from litellm.proxy._experimental.mcp_server.interactions import bind_target, open_continuation, seal_continuation +from litellm.proxy._types import UserAPIKeyAuth +from litellm.types.mcp_server.mcp_server_manager import MCPServer + + +def test_continuation_is_repeatable_and_bound_to_caller_operation_and_expiry(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", "local-continuation-test") + context: Final = OperationContext( + _caller=UserAPIKeyAuth(user_id="alice", team_id="team"), mcp_servers=("upstream",) + ) + operation: Final = CallToolRequest(params=CallToolRequestParams(name="confirm", arguments={"amount": 3})) + server: Final = MCPServer(server_id="upstream", name="upstream", url="https://example.com/mcp", transport="http") + sealed: Final = seal_continuation( + bind_target(InputRequiredResult(request_state="opaque"), server), operation, context, now=100 + ) + retry: Final = operation.model_copy( + update={"params": operation.params.model_copy(update={"request_state": sealed.request_state})} + ) + state: Final = open_continuation(retry, context, now=101) + assert state is not None + assert state.upstream_state == "opaque" + assert open_continuation(retry, context, now=102) == state + with pytest.raises(MCPError, match="Invalid or expired"): + open_continuation(retry, OperationContext(_caller=UserAPIKeyAuth(user_id="bob", team_id="team")), now=101) + with pytest.raises(MCPError, match="Invalid or expired"): + open_continuation( + retry.model_copy(update={"params": retry.params.model_copy(update={"arguments": {"amount": 4}})}), + context, + now=101, + ) + with pytest.raises(MCPError, match="Invalid or expired"): + open_continuation(retry, context, now=700) + monkeypatch.setenv("LITELLM_SALT_KEY", "rotated-test-salt") + with pytest.raises(MCPError, match="Invalid or expired"): + open_continuation(retry, context, now=101) + + +def test_continuation_missing_salt_does_not_reject_initial_request(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("LITELLM_SALT_KEY", raising=False) + context: Final = OperationContext(_caller=UserAPIKeyAuth(user_id="alice")) + operation: Final = CallToolRequest(params=CallToolRequestParams(name="confirm", arguments={})) + assert open_continuation(operation, context, now=100) is None + server: Final = MCPServer(server_id="upstream", name="upstream", url="https://example.com/mcp", transport="http") + with pytest.raises(MCPError, match="LITELLM_SALT_KEY"): + seal_continuation(bind_target(InputRequiredResult(request_state="opaque"), server), operation, context, now=100) + + +@pytest.mark.parametrize("identity", [None, UserAPIKeyAuth(), UserAPIKeyAuth(team_id="team")]) +def test_continuation_requires_stable_authenticated_principal( + identity: UserAPIKeyAuth | None, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", "test-salt") + operation: Final = CallToolRequest(params=CallToolRequestParams(name="confirm", arguments={})) + server: Final = MCPServer(server_id="server", name="server", url="https://example.com/mcp", transport="http") + with pytest.raises(MCPError, match="authenticated caller identity"): + seal_continuation( + bind_target(InputRequiredResult(request_state="opaque"), server), + operation, + OperationContext(_caller=identity), + now=100, + ) + + +def test_continuation_keeps_original_expiry_and_survives_credential_rotation(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", "test-salt") + before: Final = OperationContext(_caller=UserAPIKeyAuth(user_id="alice", api_key="old-test-key")) + after: Final = OperationContext(_caller=UserAPIKeyAuth(user_id="alice", api_key="new-test-key")) + operation: Final = CallToolRequest(params=CallToolRequestParams(name="confirm", arguments={"a": 1, "b": 2})) + server: Final = MCPServer(server_id="server", name="server", url="https://example.com/mcp", transport="http") + bound: Final = bind_target(InputRequiredResult(request_state="opaque"), server) + first: Final = seal_continuation(bound, operation, before, now=100) + retry: Final = operation.model_copy( + update={ + "params": operation.params.model_copy( + update={"request_state": first.request_state, "arguments": {"b": 2, "a": 1}} + ) + } + ) + state: Final = open_continuation(retry, after, now=200) + assert state is not None + second: Final = seal_continuation(bound, retry, after, now=699, previous=state) + final: Final = retry.model_copy( + update={"params": retry.params.model_copy(update={"request_state": second.request_state})} + ) + assert open_continuation(final, after, now=699) == state + with pytest.raises(MCPError, match="Invalid or expired"): + open_continuation(final, after, now=700) + + +@pytest.mark.parametrize("change", ["tamper", "team", "org", "principal", "method"]) +def test_altered_continuations_are_rejected(change: str, monkeypatch: pytest.MonkeyPatch) -> None: + from mcp.types import GetPromptRequest, GetPromptRequestParams + + monkeypatch.setenv("LITELLM_SALT_KEY", "test-salt") + context: Final = OperationContext(_caller=UserAPIKeyAuth(user_id="alice", team_id="one", org_id="org")) + operation: Final = CallToolRequest(params=CallToolRequestParams(name="confirm", arguments={})) + server: Final = MCPServer(server_id="server", name="server", url="https://example.com/mcp", transport="http") + first: Final = seal_continuation( + bind_target(InputRequiredResult(request_state="opaque"), server), operation, context, now=100 + ) + token: Final = first.request_state + assert token is not None + retry: Final = ( + GetPromptRequest(params=GetPromptRequestParams(name="confirm", arguments={}, request_state=token)) + if change == "method" + else operation.model_copy( + update={ + "params": operation.params.model_copy( + update={"request_state": token + "a" if change == "tamper" else token} + ) + } + ) + ) + caller: Final = UserAPIKeyAuth( + user_id="bob" if change == "principal" else "alice", + team_id="two" if change == "team" else "one", + org_id="other" if change == "org" else "org", + ) + with pytest.raises(MCPError) as error: + open_continuation(retry, OperationContext(_caller=caller), now=101) + assert error.value.error.code == -32602 + + +def test_input_responses_without_state_and_unbound_results_are_rejected(monkeypatch: pytest.MonkeyPatch) -> None: + from mcp.types import ElicitResult + from litellm.proxy._experimental.mcp_server.interactions import BoundInputRequiredResult + + monkeypatch.setenv("LITELLM_SALT_KEY", "test-salt") + context: Final = OperationContext(_caller=UserAPIKeyAuth(user_id="alice")) + operation: Final = CallToolRequest( + params=CallToolRequestParams(name="confirm", input_responses={"consent": ElicitResult(action="accept")}) + ) + with pytest.raises(MCPError, match="require a gateway continuation"): + open_continuation(operation, context, now=100) + with pytest.raises(MCPError, match="target is unavailable"): + seal_continuation(BoundInputRequiredResult(request_state="opaque"), operation, context, now=100) + + +def test_authenticated_but_invalid_state_payload_is_rejected(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy._experimental.mcp_server.state_tokens import seal_state + from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok + + monkeypatch.setenv("LITELLM_SALT_KEY", "test-salt") + sealed: Final = seal_state( + {"unsupported": "state-schema"}, purpose="mcp:interaction:repeatable:v1", expires_at=200, now=100 + ) + assert isinstance(sealed, Ok) + operation: Final = CallToolRequest(params=CallToolRequestParams(name="confirm", request_state=sealed.ok)) + with pytest.raises(MCPError, match="Invalid or expired"): + open_continuation(operation, OperationContext(_caller=UserAPIKeyAuth(user_id="alice")), now=101) + + +@pytest.mark.asyncio +async def test_modern_sampling_errors_abort_without_relaying_form_input() -> None: + import anyio + from mcp import ClientSession + from mcp.shared.message import SessionMessage + from mcp.types import ( + CreateMessageRequest, + CreateMessageRequestParams, + ElicitRequest, + ElicitRequestFormParams, + ErrorData, + ) + from litellm.proxy._experimental.mcp_server.interactions import ModernClientInteraction + + async def sampling(context: object, params: CreateMessageRequestParams) -> ErrorData: + return ErrorData(code=-32603, message="Sampling refused by gateway policy") + + send, receive = anyio.create_memory_object_stream[SessionMessage](1) + try: + interaction: Final = ModernClientInteraction(ClientSession(receive, send, sampling_callback=sampling), allow_elicitation=True) + form: Final = ElicitRequest( + params=ElicitRequestFormParams(message="Confirm", requested_schema={"type": "object", "properties": {}}) + ) + refused: Final = await interaction.request("consent", form) + assert isinstance(refused, ErrorData) + assert refused.code == -32602 + with pytest.raises(MCPError, match="Sampling refused by gateway policy"): + await interaction.prepare( + InputRequiredResult( + request_state="opaque", + input_requests={ + "consent": form, + "sample": CreateMessageRequest(params=CreateMessageRequestParams(messages=[], max_tokens=1)), + }, + ) + ) + finally: + await send.aclose() + await receive.aclose() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("invalid_response", ["legacy_state", "gateway_override", "unbound_result"]) +async def test_gateway_rejects_invalid_interaction_boundaries( + invalid_response: Literal["legacy_state", "gateway_override", "unbound_result"], monkeypatch: pytest.MonkeyPatch +) -> None: + from unittest.mock import AsyncMock, patch + from mcp.types import CreateMessageResult, ElicitResult, TextContent + from litellm.proxy._experimental.mcp_server import operations + from litellm.proxy._experimental.mcp_server.contracts import WireCompat + from litellm.proxy._experimental.mcp_server.interactions import BoundInputRequiredResult + + monkeypatch.setenv("LITELLM_SALT_KEY", "test-salt") + context: Final = OperationContext( + _caller=UserAPIKeyAuth(user_id="alice"), + wire_compat=WireCompat.LEGACY if invalid_response == "legacy_state" else WireCompat.MODERN, + ) + operation: Final = CallToolRequest(params=CallToolRequestParams(name="confirm", arguments={})) + server: Final = MCPServer(server_id="server", name="server", url="https://example.com/mcp", transport="http") + if invalid_response == "legacy_state": + operation.params.request_state = "untrusted-state" + elif invalid_response == "gateway_override": + bound: Final = bind_target( + BoundInputRequiredResult( + request_state="opaque", + gateway_responses={ + "sample": CreateMessageResult( + role="assistant", content=TextContent(type="text", text="gateway-owned"), model="test-model" + ) + }, + ), + server, + ) + sealed: Final = seal_continuation(bound, operation, context, now=100) + operation.params.request_state = sealed.request_state + operation.params.input_responses = {"sample": ElicitResult(action="accept")} + message: Final = { + "legacy_state": "Continuations require the modern MCP protocol", + "gateway_override": "Cannot replace gateway input responses", + "unbound_result": "MCP continuation target is unavailable", + }[invalid_response] + dispatch: Final = AsyncMock(return_value=InputRequiredResult(request_state="unbound-upstream-state")) + with ( + patch.object(operations, "_execute_mcp_server_tool_call", dispatch), + patch.object(operations.global_mcp_server_manager, "get_mcp_server_by_id", return_value=server), + patch.object(operations.time, "time", return_value=101), + ): + with pytest.raises(MCPError, match=message) as rejected: + await operations.GatewayOperations().execute(operation, context) + assert rejected.value.error.code == -32602 + if invalid_response == "unbound_result": + dispatch.assert_awaited_once() + else: + dispatch.assert_not_awaited() diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py index c7b477a6cc6..b67e4d24f47 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -164,7 +164,9 @@ async def test_elicitation_callback_keeps_initiating_session(): request = AsyncMock(return_value=accepted) initiating = SimpleNamespace(client_params=SimpleNamespace(capabilities=capabilities), elicit_form=request) token = legacy_server.active_mcp_session_var.set(initiating) - request_token = active_mcp_request_ctx_var.set(SimpleNamespace(session=initiating, request_id="initiating-call")) + request_token = active_mcp_request_ctx_var.set( + SimpleNamespace(session=initiating, request_id="initiating-call", protocol_version="2025-11-25") + ) try: callback = _create_elicitation_callback() legacy_server.active_mcp_session_var.set(SimpleNamespace()) @@ -4323,7 +4325,9 @@ class TestMCPServerManager: mock_create_client.assert_called_once() called_kwargs = mock_create_client.call_args.kwargs assert called_kwargs["extra_headers"] == {"X-Test": "1", "X-Static": "1"} - mock_client.read_resource.assert_awaited_once_with("https://example.com/resource") + mock_client.read_resource.assert_awaited_once_with( + "https://example.com/resource", input_responses=None, request_state=None, allow_input_required=False + ) assert result is read_result @pytest.mark.asyncio @@ -19275,3 +19279,32 @@ async def test_repeated_stale_discovery_uses_current_callers_endpoint(endpoint: original, needed_endpoint=lambda server: getattr(server, endpoint), retry_stale=False, ) assert resolved is replacement + + +@pytest.mark.asyncio +async def test_legacy_upstream_elicitation_rejects_modern_downstream_without_consent() -> None: + from types import SimpleNamespace + from mcp.types import ElicitRequestFormParams, ErrorData + from litellm.proxy._experimental.mcp_server import server as legacy_server + from litellm.proxy._experimental.mcp_server.legacy_callbacks import create_elicitation_callback + from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var + + session: Final = SimpleNamespace(client_params=None) + session_token: Final = legacy_server.active_mcp_session_var.set(session) + request_token: Final = active_mcp_request_ctx_var.set( + SimpleNamespace(session=session, request_id="modern-call", protocol_version="2026-07-28") + ) + relay: Final = AsyncMock() + try: + callback: Final = create_elicitation_callback() + with patch("litellm.proxy._experimental.mcp_server.elicitation_handler.handle_elicitation_request", relay): + result: Final = await callback( + None, ElicitRequestFormParams(message="Confirm", requested_schema={"type": "object"}) + ) + assert isinstance(result, ErrorData) + assert result.code == -32602 + assert "may have partially completed" in result.message + relay.assert_not_awaited() + finally: + active_mcp_request_ctx_var.reset(request_token) + legacy_server.active_mcp_session_var.reset(session_token) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py index 97c6e2703ad..49d21924d9c 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py @@ -6,7 +6,7 @@ import json import os from datetime import datetime, timedelta from types import SimpleNamespace -from typing import Final +from typing import Final, Literal from unittest.mock import AsyncMock, MagicMock, patch import httpx @@ -27,7 +27,7 @@ from mcp.types import ( ) from mcp.types import Tool as MCPTool from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS, LATEST_HANDSHAKE_VERSION, MODERN_PROTOCOL_VERSIONS -from pydantic import TypeAdapter +from pydantic import JsonValue, TypeAdapter from starlette.types import Message, Receive, Scope, Send from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var @@ -1086,6 +1086,9 @@ async def test_mcp_get_prompt_success(): extra_headers={"X-Test": "1"}, raw_headers=None, client_ip=None, + input_responses=None, + request_state=None, + allow_input_required=False, ) assert result is prompt_result @@ -1149,6 +1152,9 @@ async def test_mcp_read_resource_success(): extra_headers={"X-Test": "1"}, raw_headers=None, client_ip=None, + input_responses=None, + request_state=None, + allow_input_required=False, ) assert result is read_result @@ -11149,3 +11155,285 @@ async def test_discovery_adapter_preserves_authenticated_context(_mcp_request_ct context = dispatched.await_args.args[1] assert context.user_api_key_auth.user_id == "discover-caller" assert context.mcp_servers == ("allowed",) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "method,params", + ( + ("tools/call", {"name": "confirm", "arguments": {}}), + ("prompts/get", {"name": "confirm"}), + ("resources/read", {"uri": "test://confirm"}), + ), +) +@pytest.mark.parametrize("continuation", ("initial", "valid", "tampered", "revoked")) +@pytest.mark.parametrize("large_body", (False, True)) +@pytest.mark.parametrize("passthrough", (False, True)) +@pytest.mark.parametrize("route_name", ("interactive", "friendly")) +async def test_modern_oauth_challenge_follows_continuation_authorization( + method: str, + params: dict[str, JsonValue], + continuation: Literal["initial", "valid", "tampered", "revoked"], + large_body: bool, + passthrough: bool, + route_name: str, + monkeypatch: pytest.MonkeyPatch, +) -> None: + import time + from pydantic import TypeAdapter + from litellm.proxy._experimental.mcp_server import server + from litellm.proxy._experimental.mcp_server.contracts import OperationContext + from litellm.proxy._experimental.mcp_server.interactions import InteractionOperation, bind_target, seal_continuation + from mcp.types import InputRequiredResult + + monkeypatch.setenv("LITELLM_SALT_KEY", "challenge-test-salt") + monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {"mcp_advertised_versions": ["2026-07-28"]}) + caller: Final = UserAPIKeyAuth(user_id="alice") + target: Final = _make_oauth2_server("interactive", oauth2_flow="authorization_code").model_copy( + update={ + "alias": "friendly", + **( + {"auth_type": MCPAuth.none, "extra_headers": ["Authorization"], "oauth_passthrough": True} + if passthrough + else {} + ), + } + ) + operation: Final = TypeAdapter(InteractionOperation).validate_python({"method": method, "params": params}) + context: Final = OperationContext(_caller=caller, mcp_servers=(route_name,)) + sealed: Final = seal_continuation( + bind_target(InputRequiredResult(request_state="upstream"), target), operation, context, now=int(time.time()) + ) + request_params: Final = { + **params, + **( + {"requestState": "invalid" if continuation == "tampered" else sealed.request_state} + if continuation != "initial" + else {} + ), + } + body: Final = json.dumps( + { + "jsonrpc": "2.0", + "id": 7, + "method": method, + "params": { + **request_params, + "_meta": { + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {}, + "padding": "x" * (server._MCP_ROUTING_PEEK_MAX_BYTES * 2 if large_body else 0), + }, + }, + } + ).encode() + route_value: Final = params["uri" if method == "resources/read" else "name"] + assert isinstance(route_value, str) + scope: Final[Scope] = { + "type": "http", + "method": "POST", + "path": f"/mcp/{route_name}", + "scheme": "http", + "server": ("localhost", 8000), + "query_string": b"", + "root_path": "", + "headers": [ + (b"content-type", b"application/json"), + (b"x-litellm-api-key", b"test-gateway-key"), + (b"authorization", b"Bearer expired-upstream-token"), + (b"mcp-protocol-version", b"2026-07-28"), + (b"mcp-method", method.encode()), + (b"mcp-name", route_value.encode()), + (b"accept", b"application/json, text/event-stream"), + ], + } + receive: Final = AsyncMock( + side_effect=[ + {"type": "http.request", "body": body[: len(body) // 2], "more_body": True}, + {"type": "http.request", "body": body[len(body) // 2 :], "more_body": False}, + ] + ) + send: Final = AsyncMock() + with ( + patch.object( + server, + "extract_mcp_auth_context", + AsyncMock( + return_value=( + caller, + None, + [route_name], + None, + {"Authorization": "Bearer expired-upstream-token"}, + None, + ) + ), + ), + patch.object(server, "_SESSION_MANAGERS_INITIALIZED", True), + patch.object(server, "_probe_upstream_auth", AsyncMock(return_value=(401, None))) as probe, + patch.object(server.session_manager_stateless, "handle_request", AsyncMock()) as dispatch, + patch.object( + mcp_operations, + "_get_allowed_mcp_servers", + AsyncMock(return_value=[] if continuation == "revoked" else [target]), + ), + patch.object(mcp_operations.global_mcp_server_manager, "get_mcp_server_by_id", return_value=target), + patch.object(mcp_operations.global_mcp_server_manager, "get_mcp_server_by_name", return_value=target), + patch.object( + mcp_operations.global_mcp_server_manager, "ensure_oauth_metadata_discovered", AsyncMock(return_value=target) + ) as discovery, + patch.object( + mcp_operations.global_mcp_server_manager, "has_user_oauth_token", AsyncMock(return_value=False) + ) as token, + ): + if continuation in ("initial", "valid"): + with pytest.raises(HTTPException) as rejected: + await server.handle_streamable_http_mcp(scope, receive, send) + assert rejected.value.status_code == 401 + if passthrough: + assert "resource_metadata=" in rejected.value.headers["www-authenticate"] + assert "invalid_token" in rejected.value.headers["www-authenticate"] + probe.assert_awaited_once_with(target.url, "Bearer expired-upstream-token") + token.assert_not_awaited() + else: + assert ( + f'Bearer authorization_uri="http://localhost:8000/.well-known/oauth-authorization-server/mcp/{route_name}"' + == rejected.value.headers["www-authenticate"] + ) + token.assert_awaited_once() + probe.assert_not_awaited() + elif continuation == "revoked": + with pytest.raises(HTTPException) as rejected: + await server.handle_streamable_http_mcp(scope, receive, send) + assert rejected.value.status_code == 403 + probe.assert_not_awaited() + discovery.assert_not_awaited() + token.assert_not_awaited() + else: + await server.handle_streamable_http_mcp(scope, receive, send) + response_body: Final = json.loads( + next( + call.args[0]["body"] + for call in send.await_args_list + if call.args[0]["type"] == "http.response.body" + ) + ) + assert response_body["error"]["code"] == -32602 + probe.assert_not_awaited() + discovery.assert_not_awaited() + token.assert_not_awaited() + dispatch.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("method", ("tools/call", "prompts/get", "resources/read")) +@pytest.mark.parametrize("size_delta", (-1, 0, 1)) +@pytest.mark.parametrize("chunked", (False, True)) +async def test_modern_preflight_enforces_sdk_body_limit( + method: str, size_delta: int, chunked: bool, monkeypatch: pytest.MonkeyPatch +) -> None: + from litellm.proxy._experimental.mcp_server import server + + monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {"mcp_advertised_versions": ["2026-07-28"]}) + limit: Final = server.session_manager_stateless.max_request_body_size + params: Final = {"uri": "test://confirm"} if method == "resources/read" else {"name": "confirm"} + payload: Final = json.dumps( + { + "jsonrpc": "2.0", + "id": 7, + "method": method, + "params": { + **params, + "_meta": { + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {}, + "padding": "", + }, + }, + } + ).encode() + body: Final = payload.replace( + b'"padding": ""', b'"padding": "' + b"x" * (limit + size_delta - len(payload)) + b'"', 1 + ) + assert len(body) == limit + size_delta + chunks: Final = (body[: limit - 1], body[limit - 1 :]) if chunked else (body,) + messages: Final[list[Message]] = [ + {"type": "http.request", "body": chunk, "more_body": index < len(chunks) - 1 or size_delta > 0} + for index, chunk in enumerate(chunks) + ] + receive: Final = AsyncMock(side_effect=[*messages, {"type": "http.request", "body": b"", "more_body": False}]) + scope: Final[Scope] = { + "type": "http", + "method": "POST", + "path": "/mcp", + "scheme": "http", + "server": ("localhost", 8000), + "query_string": b"", + "root_path": "", + "headers": [ + (b"content-type", b"application/json"), + (b"authorization", b"Bearer test-key"), + (b"mcp-protocol-version", b"2026-07-28"), + (b"mcp-method", method.encode()), + (b"mcp-name", next(iter(params.values())).encode()), + *([] if chunked else [(b"content-length", str(len(body)).encode())]), + ], + } + with ( + patch.object( + server, "extract_mcp_auth_context", AsyncMock(return_value=(UserAPIKeyAuth(), None, None, None, None, None)) + ), + patch.object(server, "_SESSION_MANAGERS_INITIALIZED", True), + patch.object(server.session_manager_stateless, "handle_request", AsyncMock()) as dispatch, + patch.object(server.operations, "_get_allowed_mcp_servers", AsyncMock()) as authorization, + ): + if size_delta > 0: + with pytest.raises(HTTPException) as rejected: + await server.handle_streamable_http_mcp(scope, receive, AsyncMock()) + assert rejected.value.status_code == 413 + assert rejected.value.detail == "Request body too large" + dispatch.assert_not_awaited() + else: + await server.handle_streamable_http_mcp(scope, receive, AsyncMock()) + dispatch.assert_awaited_once() + assert receive.await_count == len(chunks) + authorization.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("invalid_request", ("malformed", "other_method", "header_mismatch", "duplicate_header")) +async def test_modern_preflight_leaves_invalid_envelopes_to_sdk_without_upstream_work( + invalid_request: Literal["malformed", "other_method", "header_mismatch", "duplicate_header"], +) -> None: + from litellm.proxy._experimental.mcp_server import server + from litellm.proxy._experimental.mcp_server.contracts import OperationContext + + envelope: Final = { + "jsonrpc": "2.0", + "id": 7, + "method": "tools/list" if invalid_request == "other_method" else "tools/call", + "params": { + "name": "confirm", + "_meta": { + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {}, + }, + }, + } + headers: Final = [ + (b"mcp-protocol-version", b"2026-07-28"), + (b"mcp-method", b"tools/call"), + (b"mcp-name", b"wrong" if invalid_request == "header_mismatch" else b"confirm"), + ] + scope: Final[Scope] = { + "type": "http", + "headers": [*headers, *([(b"mcp-method", b"tools/call")] if invalid_request == "duplicate_header" else [])], + } + with patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock()) as resolve: + result: Final = await server._preflight_modern_interaction( + scope, + b"{" if invalid_request == "malformed" else json.dumps(envelope).encode(), + OperationContext(_caller=UserAPIKeyAuth(user_id="alice"), mcp_servers=("interactive",)), + ) + assert result is None + resolve.assert_not_awaited() diff --git a/tests/unit/proxy/proxy_server/test_proxy_config.py b/tests/unit/proxy/proxy_server/test_proxy_config.py index d1b80627b3a..912404f91b8 100644 --- a/tests/unit/proxy/proxy_server/test_proxy_config.py +++ b/tests/unit/proxy/proxy_server/test_proxy_config.py @@ -5449,14 +5449,24 @@ async def test_model_refresh_updates_availability_catalog_and_retains_it_on_db_f @pytest.mark.asyncio -@pytest.mark.parametrize("versions", [None, ["2024-11-05"], [], ["2026-07-28"], ["unknown"]]) -async def test_proxy_config_validates_advertised_mcp_versions_at_load(tmp_path, monkeypatch, versions): +@pytest.mark.parametrize( + ("versions", "valid"), + [ + (None, True), + (["2024-11-05"], True), + (["2026-07-28"], True), + (["2025-11-25", "2026-07-28"], True), + ([], False), + (["unknown"], False), + ], +) +async def test_proxy_config_validates_advertised_mcp_versions_at_load(tmp_path, monkeypatch, versions, valid): config = tmp_path / "mcp-versions.yaml" config.write_text(json.dumps({"model_list": [], "general_settings": {"mcp_advertised_versions": versions}})) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False) monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False) - if versions is None or versions == ["2024-11-05"]: + if valid: _, _, settings = await ProxyConfig().load_config(router=None, config_file_path=str(config)) assert settings["mcp_advertised_versions"] == versions return diff --git a/tests/unit/proxy/test__types.py b/tests/unit/proxy/test__types.py index dd7bd86982e..50a4c3a6ca7 100644 --- a/tests/unit/proxy/test__types.py +++ b/tests/unit/proxy/test__types.py @@ -387,7 +387,7 @@ def test_change_password_request_passwords_hidden_from_repr(): for rendered in (repr(request), str(request)): assert "hunter2hunter2" not in rendered assert "NewP@ssw0rd-2026" not in rendered -@pytest.mark.parametrize("versions", [[], ["2099-01-01"], ["2026-07-28"]]) +@pytest.mark.parametrize("versions", [[], ["2099-01-01"]]) def test_mcp_advertised_versions_reject_unavailable_revisions(versions): from pydantic import ValidationError diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index fa3c22c0713..52f56ddc92e 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -30304,9 +30304,9 @@ export interface components { maximum_spend_logs_retention_period?: string | null; /** * Mcp Advertised Versions - * @description MCP revisions enabled by the gateway. Defaults to all completed legacy revisions. Modern protocol serving and Apps/Tasks remain disabled. + * @description MCP revisions enabled by the gateway. Defaults to all completed legacy revisions. Modern protocol serving requires explicit opt-in. Apps/Tasks remain disabled. */ - mcp_advertised_versions?: ("2024-11-05" | "2025-03-26" | "2025-06-18" | "2025-11-25")[] | null; + mcp_advertised_versions?: (("2024-11-05" | "2025-03-26" | "2025-06-18" | "2025-11-25") | "2026-07-28")[] | null; /** * Mcp Allowed Clients * @description MCP client applications admitted by the gateway, each an {alias, value} pair where alias is the name shown in the dashboard and logs and value is the identity that must match exactly. When set, every MCP request must carry a client identity equal to one of the values: a JWT caller is identified by the claim named in litellm_jwtauth.mcp_client_id_jwt_field, any other caller by the header named in mcp_client_id_header. A request with no resolvable identity, or an unlisted one, is rejected with 403. Unset means every client is admitted.