mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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
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:
parent
0ae62c5d02
commit
04f0ea8c33
21 changed files with 1952 additions and 92 deletions
7
.github/workflows/test-unit.yml
vendored
7
.github/workflows/test-unit.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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, ...]:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
262
litellm/proxy/_experimental/mcp_server/interactions.py
Normal file
262
litellm/proxy/_experimental/mcp_server/interactions.py
Normal 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
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
395
tests/integration/mcp/test_interactions.py
Normal file
395
tests/integration/mcp/test_interactions.py
Normal 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
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
253
tests/unit/proxy/_experimental/mcp_server/test_interactions.py
Normal file
253
tests/unit/proxy/_experimental/mcp_server/test_interactions.py
Normal 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()
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
4
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
4
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue