feat(mcp): translate input requests and bind continuations (#45464)
Some checks are pending
CI Coverage / assert-ci-coverage (push) Waiting to run
CodSpeed Benchmarks / benchmarks (push) Waiting to run
Helm unit test / unit-test (push) Waiting to run
Lens Worker Image / lens-worker-image (amd64, ubuntu-latest) (push) Waiting to run
Lens Worker Image / lens-worker-image (arm64, ubuntu-24.04-arm) (push) Waiting to run
Lens Worker Image / Publish Lens development index (push) Blocked by required conditions
Publish lint base counts / publish (basedpyright, scripts/type_check_gate.py) (push) Waiting to run
Publish lint base counts / publish (ruff-strict, scripts/ruff_strict_gate.py) (push) Waiting to run
Publish lint base counts / publish (test-quality, scripts/test_quality_gate.py) (push) Waiting to run
Publish lint base counts / publish (type-discipline, scripts/type_discipline_gate.py) (push) Waiting to run
Scorecard supply-chain security / Scorecard analysis (push) Waiting to run
Code Quality Checks / code-quality (push) Waiting to run
Code Quality Checks / python-310-import-smoke (push) Waiting to run
UI Unit Tests / ui-unit-tests (push) Waiting to run
LiteLLM Rust / rust-lint (push) Waiting to run
LiteLLM Rust / rust-test (push) Waiting to run
LiteLLM Rust / rust-wheel (push) Waiting to run
Unit Tests: Documentation Validation / documentation (push) Waiting to run
Unit Tests / proxy-runtime (push) Blocked by required conditions
Unit Tests / proxy-server-core (push) Blocked by required conditions
Unit Tests / proxy-utils (push) Blocked by required conditions
Unit Tests / unit passed (push) Blocked by required conditions
Unit Tests / mcp-oauth (push) Blocked by required conditions
Unit Tests / proxy-feature-endpoints (push) Blocked by required conditions
Unit Tests / enterprise-routing (push) Blocked by required conditions
Unit Tests / proxy-endpoints (push) Blocked by required conditions
Unit Tests / Build the Rust bridge (push) Waiting to run
Unit Tests / caching-local (push) Blocked by required conditions
Unit Tests / core-utils (push) Blocked by required conditions
Unit Tests / enterprise-managed-files (push) Blocked by required conditions
Unit Tests / enterprise-package (push) Blocked by required conditions
Unit Tests / integrations (push) Blocked by required conditions
Unit Tests / OpenAI and Meta Providers (push) Blocked by required conditions
Unit Tests / All Other Providers (push) Blocked by required conditions
Unit Tests / Vertex AI (push) Blocked by required conditions
Unit Tests / proxy-extras (push) Blocked by required conditions
Unit Tests / mcp-elicitation (push) Blocked by required conditions
Unit Tests / misc (push) Blocked by required conditions
Unit Tests / misc-dirs (push) Blocked by required conditions
Unit Tests / proxy-auth (push) Blocked by required conditions
Unit Tests / proxy-hooks-client (push) Blocked by required conditions
Unit Tests / proxy-infra (push) Blocked by required conditions
Unit Tests / proxy-infra-root (push) Blocked by required conditions
Unit Tests / proxy-server (push) Blocked by required conditions
Unit Tests / unit (push) Blocked by required conditions
Unit Tests / responses-caching-types (push) Blocked by required conditions
Unit Tests / Lens Python 3.10 (push) Waiting to run
Unit Tests / assert-shard-coverage (push) Waiting to run
Unit Tests / auth-checks (push) Blocked by required conditions
Unit Tests / budgets (push) Blocked by required conditions
Unit Tests / custom-logging (push) Blocked by required conditions
Unit Tests / db-and-spend (push) Blocked by required conditions
Unit Tests / endpoints-and-responses (push) Blocked by required conditions
Unit Tests / guardrails-hooks (push) Blocked by required conditions
Unit Tests / jwt-and-keys (push) Blocked by required conditions
Unit Tests / key-generation (push) Blocked by required conditions
Unit Tests / logging-misc (push) Blocked by required conditions
GitHub Actions Security Analysis / zizmor (push) Waiting to run

* feat(mcp): translate input requests and bind resumable continuations

* fix(mcp): keep continuations bound to their original upstream

* fix(mcp): sync advertised protocol API schema

* test(mcp): run interaction regressions in GitHub CI

* test(mcp): set source paths for interaction proxy processes

* test(mcp): verify salt guidance across interaction carriers

* fix(mcp): preserve interaction authentication and elicitation policy

* fix(mcp): bound preflight bodies and preserve challenge routes

* test(mcp): type interaction regression fixtures

---------

Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
This commit is contained in:
joshua-berri 2026-10-08 20:10:10 -07:00 • committed by GitHub
parent 0ae62c5d02
commit 04f0ea8c33
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
21 changed files with 1952 additions and 92 deletions

View file

@ -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

View file

@ -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)

View file

@ -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, ...]:

View file

@ -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

View file

@ -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

View file

@ -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(

View file

@ -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)

View file

@ -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):

View file

@ -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":

View file

@ -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,

View file

@ -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,

View file

@ -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

View file

@ -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)

View file

@ -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)

View file

@ -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"
)

View file

@ -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()

View file

@ -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)

View file

@ -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()

View file

@ -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

View file

@ -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

View file

@ -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.