mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-26 01:12:21 +00:00
feat(mcp): configure protocol versions and capability discovery (#43169)
* feat(mcp): configure protocol versions and capability discovery * test(mcp): return SDK initialization result in REST pagination fixture * test(mcp): arm cancellation deadlines after TCP calls start * fix(mcp): avoid serialized discovery and listing spend logs * test(mcp): isolate default protocol header policy * fix(mcp): honor preview protocol pins and refresh migrated tests * fix(mcp): preserve edited protocol pins in saved OAuth previews * fix(mcp): retain saved protocol pins when previews omit versions --------- Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
This commit is contained in:
parent
cf491d1df9
commit
6b7688869e
25 changed files with 939 additions and 42 deletions
|
|
@ -10,6 +10,7 @@ import os
|
|||
from collections.abc import AsyncIterator, Awaitable, Callable, Generator, Sequence
|
||||
from contextlib import AbstractAsyncContextManager
|
||||
from functools import partial
|
||||
from importlib.metadata import version
|
||||
from types import MappingProxyType
|
||||
from typing import Final, TypeAlias, TypeVar, cast
|
||||
|
||||
|
|
@ -34,8 +35,16 @@ _TransportContext: TypeAlias = AbstractAsyncContextManager[_TransportStreams]
|
|||
from mcp.types import (
|
||||
METHOD_NOT_FOUND,
|
||||
REQUEST_TIMEOUT,
|
||||
ClientCapabilities,
|
||||
ElicitationCapability,
|
||||
FormElicitationCapability,
|
||||
GetPromptRequestParams,
|
||||
GetPromptResult,
|
||||
Implementation,
|
||||
InitializedNotification,
|
||||
InitializeRequest,
|
||||
InitializeRequestParams,
|
||||
InitializeResult,
|
||||
InputRequiredResult,
|
||||
ListPromptsResult,
|
||||
ListResourcesResult,
|
||||
|
|
@ -44,12 +53,14 @@ from mcp.types import (
|
|||
PaginatedResult,
|
||||
Prompt,
|
||||
ResourceTemplate,
|
||||
SamplingCapability,
|
||||
ServerNotification,
|
||||
UrlElicitationCapability,
|
||||
)
|
||||
from mcp.types import CallToolRequestParams as MCPCallToolRequestParams
|
||||
from mcp.types import CallToolResult as MCPCallToolResult
|
||||
from mcp.types import Tool as MCPTool
|
||||
from pydantic import AnyUrl
|
||||
from pydantic import AnyUrl, TypeAdapter
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import (
|
||||
|
|
@ -64,11 +75,13 @@ from litellm.proxy._experimental.mcp_server.mcp_debug import capture_upstream_er
|
|||
from litellm.proxy._experimental.mcp_server.result_conversion import error_text_result
|
||||
from litellm.types.llms.custom_http import VerifyTypes
|
||||
from litellm.types.mcp import (
|
||||
MCP_LEGACY_VERSIONS,
|
||||
MCPAuth,
|
||||
MCPAuthType,
|
||||
MCPStdioConfig,
|
||||
MCPTransport,
|
||||
MCPTransportType,
|
||||
MCPUpstreamProtocol,
|
||||
credential_redirect_hook,
|
||||
has_header,
|
||||
without_header,
|
||||
|
|
@ -386,7 +399,9 @@ class MCPClient:
|
|||
sampling_callback: Callable | None = None,
|
||||
elicitation_callback: Callable | None = None,
|
||||
logging_callback: Callable | None = None,
|
||||
protocol_version: MCPUpstreamProtocol = "auto",
|
||||
):
|
||||
self.protocol_version: MCPUpstreamProtocol = TypeAdapter(MCPUpstreamProtocol).validate_python(protocol_version)
|
||||
self.server_url: str = server_url
|
||||
self.transport_type: MCPTransport = transport_type
|
||||
self.auth_type: MCPAuthType = auth_type
|
||||
|
|
@ -525,6 +540,35 @@ class MCPClient:
|
|||
|
||||
return safe_env
|
||||
|
||||
async def _initialize_session(self, session: ClientSession) -> InitializeResult:
|
||||
if self.protocol_version == "auto":
|
||||
automatic: Final = await session.initialize()
|
||||
if automatic.protocol_version not in MCP_LEGACY_VERSIONS:
|
||||
raise MCPError(code=-32022, message="Upstream selected an unsupported MCP protocol version")
|
||||
return automatic
|
||||
result: Final = await session.send_request(
|
||||
InitializeRequest(
|
||||
params=InitializeRequestParams(
|
||||
protocol_version=self.protocol_version,
|
||||
client_info=Implementation(name="litellm", version=version("litellm")),
|
||||
capabilities=ClientCapabilities(
|
||||
sampling=SamplingCapability() if self._sampling_callback is not None else None,
|
||||
elicitation=ElicitationCapability(
|
||||
form=FormElicitationCapability(), url=UrlElicitationCapability()
|
||||
)
|
||||
if self._elicitation_callback is not None
|
||||
else None,
|
||||
),
|
||||
)
|
||||
),
|
||||
InitializeResult,
|
||||
)
|
||||
if result.protocol_version != self.protocol_version:
|
||||
raise MCPError(code=-32022, message="Upstream did not accept the configured MCP protocol version")
|
||||
session.adopt(result)
|
||||
await session.send_notification(InitializedNotification())
|
||||
return result
|
||||
|
||||
async def _execute_session_operation(
|
||||
self,
|
||||
transport_ctx: _TransportContext,
|
||||
|
|
@ -579,7 +623,7 @@ class MCPClient:
|
|||
)
|
||||
session: Final = await session_ctx.__aenter__()
|
||||
try:
|
||||
init_result: Final = await session.initialize()
|
||||
init_result: Final = await self._initialize_session(session)
|
||||
instructions: Final = getattr(init_result, "instructions", None)
|
||||
self._last_initialize_instructions = (
|
||||
instructions.strip() or None if isinstance(instructions, str) else None
|
||||
|
|
|
|||
145
litellm/proxy/_experimental/mcp_server/capabilities.py
Normal file
145
litellm/proxy/_experimental/mcp_server/capabilities.py
Normal file
|
|
@ -0,0 +1,145 @@
|
|||
from collections.abc import Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from itertools import product
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal
|
||||
|
||||
from mcp.server.context import CallNext, HandlerResult, ServerRequestContext
|
||||
from mcp.shared.exceptions import MCPError
|
||||
from mcp.types import DiscoverResult, InitializeRequestParams, InitializeResult, ServerCapabilities
|
||||
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
|
||||
|
||||
GATEWAY_OPERATIONS: Final = frozenset(
|
||||
{
|
||||
"tools/list",
|
||||
"tools/call",
|
||||
"prompts/list",
|
||||
"prompts/get",
|
||||
"resources/list",
|
||||
"resources/read",
|
||||
"resources/templates/list",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RevisionSupport:
|
||||
transports: frozenset[MCPTransport]
|
||||
operations: frozenset[str]
|
||||
results: frozenset[Literal["complete", "input_required"]]
|
||||
extensions: frozenset[str]
|
||||
completed: bool
|
||||
|
||||
|
||||
REVISION_SUPPORT: Final[Mapping[str, RevisionSupport]] = MappingProxyType(
|
||||
{
|
||||
version.value: RevisionSupport(
|
||||
transports=frozenset(MCPTransport)
|
||||
if version.value in HANDSHAKE_PROTOCOL_VERSIONS
|
||||
else frozenset({MCPTransport.http, MCPTransport.stdio}),
|
||||
operations=frozenset(method for method in GATEWAY_OPERATIONS if (method, version.value) in CLIENT_REQUESTS),
|
||||
results=frozenset({"complete"})
|
||||
if version.value in HANDSHAKE_PROTOCOL_VERSIONS
|
||||
else frozenset({"complete", "input_required"}),
|
||||
extensions=frozenset(),
|
||||
completed=version.value in HANDSHAKE_PROTOCOL_VERSIONS,
|
||||
)
|
||||
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)
|
||||
|
||||
|
||||
def configured_versions() -> tuple[str, ...]:
|
||||
from litellm.proxy.proxy_server import general_settings_view
|
||||
|
||||
configured: Final = general_settings_view().get("mcp_advertised_versions")
|
||||
return _ADVERTISED_VERSIONS.validate_python(MCP_LEGACY_VERSIONS if configured is None else configured)
|
||||
|
||||
|
||||
def build_discovery(
|
||||
*,
|
||||
configured: tuple[str, ...],
|
||||
revision: str,
|
||||
transport: MCPTransport,
|
||||
authorized_operations: frozenset[str],
|
||||
upstream_versions: frozenset[str],
|
||||
capabilities: ServerCapabilities,
|
||||
client_extensions: frozenset[str] = frozenset(),
|
||||
upstream_extensions: frozenset[str] = frozenset(),
|
||||
instructions: str | None = None,
|
||||
) -> DiscoverResult:
|
||||
supported: Final = tuple(
|
||||
version
|
||||
for version, support in REVISION_SUPPORT.items()
|
||||
if version in configured and support.completed and transport in support.transports
|
||||
)
|
||||
revision_support: Final = REVISION_SUPPORT.get(revision)
|
||||
operations: Final[frozenset[str]] = (
|
||||
authorized_operations & revision_support.operations
|
||||
if revision in supported
|
||||
and revision_support is not None
|
||||
and any((revision, upstream) in TRANSLATION_PAIRS for upstream in upstream_versions)
|
||||
else frozenset()
|
||||
)
|
||||
extensions: Final[frozenset[str]] = (
|
||||
revision_support.extensions & client_extensions & upstream_extensions
|
||||
if operations and revision_support is not None
|
||||
else frozenset()
|
||||
)
|
||||
caller_capabilities: Final = capabilities.model_copy(deep=True)
|
||||
return DiscoverResult(
|
||||
supported_versions=list(supported),
|
||||
capabilities=ServerCapabilities(
|
||||
tools=caller_capabilities.tools if {"tools/list", "tools/call"} <= operations else None,
|
||||
prompts=caller_capabilities.prompts if {"prompts/list", "prompts/get"} <= operations else None,
|
||||
resources=caller_capabilities.resources if {"resources/list", "resources/read"} <= operations else None,
|
||||
extensions={
|
||||
key: value for key, value in (caller_capabilities.extensions or {}).items() if key in extensions
|
||||
}
|
||||
or None,
|
||||
),
|
||||
instructions=instructions,
|
||||
cache_scope="private",
|
||||
ttl_ms=0,
|
||||
)
|
||||
|
||||
|
||||
class GatewayVersionPolicy:
|
||||
def __init__(self, versions: Callable[[], tuple[str, ...]] = configured_versions) -> None:
|
||||
self._versions = versions
|
||||
|
||||
async def __call__(self, ctx: ServerRequestContext[object, object], call_next: CallNext) -> HandlerResult:
|
||||
versions: Final = self._versions()
|
||||
requested: Final = (
|
||||
InitializeRequestParams.model_validate(ctx.params or {}).protocol_version
|
||||
if ctx.method == "initialize"
|
||||
else ctx.protocol_version
|
||||
)
|
||||
negotiated: Final = (
|
||||
(requested if requested in HANDSHAKE_PROTOCOL_VERSIONS else LATEST_HANDSHAKE_VERSION)
|
||||
if ctx.method == "initialize"
|
||||
else requested
|
||||
)
|
||||
if negotiated not in versions:
|
||||
raise MCPError(code=-32022, message="Unsupported MCP protocol version", data={"supported": list(versions)})
|
||||
result: Final = await call_next(ctx)
|
||||
if ctx.method != "initialize":
|
||||
return result
|
||||
initialized: Final = InitializeResult.model_validate(result)
|
||||
discovery: Final = build_discovery(
|
||||
configured=versions,
|
||||
revision=initialized.protocol_version,
|
||||
transport=MCPTransport.http,
|
||||
authorized_operations=GATEWAY_OPERATIONS,
|
||||
upstream_versions=frozenset(HANDSHAKE_PROTOCOL_VERSIONS),
|
||||
capabilities=initialized.capabilities,
|
||||
instructions=initialized.instructions,
|
||||
)
|
||||
return initialized.model_copy(update={"capabilities": discovery.capabilities})
|
||||
|
|
@ -28,6 +28,7 @@ class OperationContext:
|
|||
client_ip: str | None = None
|
||||
mcp_proxy_mode: bool = False
|
||||
wire_compat: WireCompat = WireCompat.LEGACY
|
||||
protocol_version: str | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
object.__setattr__(self, "_caller", copy_caller(self._caller))
|
||||
|
|
|
|||
|
|
@ -193,6 +193,7 @@ from litellm.types.mcp import (
|
|||
MCPAuth,
|
||||
MCPStdioConfig,
|
||||
MCPTokenEndpointAuthMethod,
|
||||
MCPUpstreamProtocol,
|
||||
has_header,
|
||||
without_header,
|
||||
)
|
||||
|
|
@ -340,6 +341,7 @@ class MCPServerConfig(TypedDict, total=False):
|
|||
whatever the admin wrote, and each read applies its own default."""
|
||||
|
||||
server_id: ReadOnly[str]
|
||||
protocol_version: ReadOnly[MCPUpstreamProtocol]
|
||||
alias: str
|
||||
description: str
|
||||
mcp_info: MCPInfo
|
||||
|
|
@ -2549,6 +2551,9 @@ class MCPServerManager:
|
|||
new_server = MCPServer(
|
||||
server_id=server_id,
|
||||
name=name_for_prefix,
|
||||
protocol_version=TypeAdapter(MCPUpstreamProtocol).validate_python(
|
||||
server_config.get("protocol_version", mcp_info.get("protocol_version", "auto"))
|
||||
),
|
||||
alias=alias,
|
||||
server_name=server_name,
|
||||
spec_path=server_config.get("spec_path", None),
|
||||
|
|
@ -3109,6 +3114,9 @@ class MCPServerManager:
|
|||
new_server: Final = MCPServer(
|
||||
server_id=mcp_server.server_id,
|
||||
name=name_for_prefix,
|
||||
protocol_version=TypeAdapter(MCPUpstreamProtocol).validate_python(
|
||||
_mcp_info.get("protocol_version", "auto")
|
||||
),
|
||||
alias=getattr(mcp_server, "alias", None),
|
||||
server_name=getattr(mcp_server, "server_name", None),
|
||||
url=mcp_server.url,
|
||||
|
|
@ -4145,6 +4153,7 @@ class MCPServerManager:
|
|||
cred_provider: UpstreamCredentialProvider | None = None,
|
||||
raw_headers: Mapping[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
protocol_version_override: MCPUpstreamProtocol | None = None,
|
||||
) -> MCPClient:
|
||||
"""
|
||||
Create an MCPClient instance for the given server.
|
||||
|
|
@ -4168,6 +4177,9 @@ class MCPServerManager:
|
|||
"""
|
||||
record_auth_resolution(server.server_id, AuthResolution.unresolved)
|
||||
resolved_server: Final = await self.ensure_oauth_metadata_discovered(server)
|
||||
protocol_version: Final = (
|
||||
protocol_version_override if protocol_version_override is not None else resolved_server.protocol_version
|
||||
)
|
||||
transport: Final = resolved_server.transport or MCPTransport.sse
|
||||
spec = None if transport == MCPTransport.stdio else _to_server_spec_fail_closed(resolved_server)
|
||||
provider: Final = cred_provider or self._cred_provider
|
||||
|
|
@ -4249,6 +4261,7 @@ class MCPServerManager:
|
|||
return MCPClient(
|
||||
server_url="", # Not used for stdio
|
||||
transport_type=transport,
|
||||
protocol_version=protocol_version,
|
||||
auth_type=resolved_server.auth_type,
|
||||
auth_value=auth_value,
|
||||
timeout=(resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT),
|
||||
|
|
@ -4281,6 +4294,7 @@ class MCPServerManager:
|
|||
MCPClient(
|
||||
server_url=server_url,
|
||||
transport_type=transport,
|
||||
protocol_version=protocol_version,
|
||||
auth_type=resolved_server.auth_type,
|
||||
timeout=(
|
||||
resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT
|
||||
|
|
@ -4324,6 +4338,7 @@ class MCPServerManager:
|
|||
MCPClient(
|
||||
server_url=server_url,
|
||||
transport_type=transport,
|
||||
protocol_version=protocol_version,
|
||||
auth_type=resolved_server.auth_type,
|
||||
auth_value=auth_value,
|
||||
auth_header_name=auth_header_name,
|
||||
|
|
|
|||
|
|
@ -14,6 +14,8 @@ from mcp.types import (
|
|||
CallToolRequest,
|
||||
CallToolRequestParams,
|
||||
CallToolResult,
|
||||
DiscoverRequest,
|
||||
DiscoverResult,
|
||||
GetPromptRequest,
|
||||
GetPromptRequestParams,
|
||||
GetPromptResult,
|
||||
|
|
@ -28,10 +30,14 @@ from mcp.types import (
|
|||
ListToolsResult,
|
||||
PaginatedRequestParams,
|
||||
Prompt,
|
||||
PromptsCapability,
|
||||
ReadResourceRequest,
|
||||
ReadResourceRequestParams,
|
||||
ResourcesCapability,
|
||||
ResourceTemplate,
|
||||
ServerCapabilities,
|
||||
TextContent,
|
||||
ToolsCapability,
|
||||
)
|
||||
from mcp.types import Tool as MCPTool
|
||||
from pydantic import AnyUrl, ConfigDict, Field, TypeAdapter
|
||||
|
|
@ -51,6 +57,11 @@ from litellm.proxy._experimental.mcp_server.byok_credential_cache import (
|
|||
cache_byok_credential,
|
||||
get_cached_byok_credential,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.capabilities import (
|
||||
GATEWAY_OPERATIONS,
|
||||
build_discovery,
|
||||
configured_versions,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.contracts import (
|
||||
AuthorizedToolCall,
|
||||
OperationContext,
|
||||
|
|
@ -122,7 +133,9 @@ from litellm.proxy.litellm_pre_call_utils import (
|
|||
)
|
||||
from litellm.types.mcp import (
|
||||
DEFAULT_CREDENTIAL_HEADER,
|
||||
MCP_LEGACY_VERSIONS,
|
||||
MCPAuth,
|
||||
MCPTransport,
|
||||
without_header,
|
||||
)
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPInfo, MCPServer
|
||||
|
|
@ -2657,7 +2670,11 @@ class _McpDeniedDetail(TypedDict):
|
|||
|
||||
|
||||
async def _execute_handle_list_tools(
|
||||
context: OperationContext, params: PaginatedRequestParams, host_progress_callback: ProgressCallback | None = None
|
||||
context: OperationContext,
|
||||
params: PaginatedRequestParams,
|
||||
host_progress_callback: ProgressCallback | None = None,
|
||||
*,
|
||||
log_list_tools_to_spendlogs: bool = True,
|
||||
) -> ListToolsResult:
|
||||
try:
|
||||
(
|
||||
|
|
@ -2700,7 +2717,7 @@ async def _execute_handle_list_tools(
|
|||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
log_list_tools_to_spendlogs=True,
|
||||
log_list_tools_to_spendlogs=log_list_tools_to_spendlogs,
|
||||
list_tools_log_source="mcp_protocol",
|
||||
client_ip=_client_ip,
|
||||
)
|
||||
|
|
@ -3065,6 +3082,7 @@ def prepare_context(
|
|||
client_ip: str | None = None,
|
||||
mcp_proxy_mode: bool = False,
|
||||
wire_compat: WireCompat = WireCompat.LEGACY,
|
||||
protocol_version: str | None = None,
|
||||
) -> OperationContext:
|
||||
return OperationContext(
|
||||
_caller=user_api_key_auth,
|
||||
|
|
@ -3076,11 +3094,13 @@ def prepare_context(
|
|||
client_ip=client_ip,
|
||||
mcp_proxy_mode=mcp_proxy_mode,
|
||||
wire_compat=wire_compat,
|
||||
protocol_version=protocol_version,
|
||||
)
|
||||
|
||||
|
||||
GatewayOperation: TypeAlias = (
|
||||
AuthorizedToolCall
|
||||
| DiscoverRequest
|
||||
| ListToolsRequest
|
||||
| CallToolRequest
|
||||
| ListPromptsRequest
|
||||
|
|
@ -3090,7 +3110,8 @@ GatewayOperation: TypeAlias = (
|
|||
| ReadResourceRequest
|
||||
)
|
||||
GatewayResult: TypeAlias = (
|
||||
ListToolsResult
|
||||
DiscoverResult
|
||||
| ListToolsResult
|
||||
| CallToolResult
|
||||
| InputRequiredResult
|
||||
| ListPromptsResult
|
||||
|
|
@ -3105,6 +3126,9 @@ class GatewayOperations:
|
|||
def __init__(self, host_progress_callback: ProgressCallback | None = None) -> None:
|
||||
self._host_progress_callback = host_progress_callback
|
||||
|
||||
@overload
|
||||
async def execute(self, operation: DiscoverRequest, context: OperationContext) -> DiscoverResult: ...
|
||||
|
||||
@overload
|
||||
async def execute(
|
||||
self, operation: AuthorizedToolCall, context: OperationContext
|
||||
|
|
@ -3137,6 +3161,51 @@ class GatewayOperations:
|
|||
|
||||
async def execute(self, operation: GatewayOperation, context: OperationContext) -> GatewayResult:
|
||||
match operation:
|
||||
case DiscoverRequest():
|
||||
listings: Final = (
|
||||
()
|
||||
if context.mcp_proxy_mode
|
||||
else (ListPromptsRequest(), ListResourcesRequest(), ListResourceTemplatesRequest())
|
||||
)
|
||||
tasks: Final = (
|
||||
asyncio.create_task(
|
||||
_execute_handle_list_tools(
|
||||
context,
|
||||
PaginatedRequestParams(),
|
||||
self._host_progress_callback,
|
||||
log_list_tools_to_spendlogs=False,
|
||||
)
|
||||
),
|
||||
*(asyncio.create_task(self.execute(listing, context)) for listing in listings),
|
||||
)
|
||||
try:
|
||||
results: Final = await asyncio.gather(*tasks)
|
||||
finally:
|
||||
for task in tasks:
|
||||
task.cancel()
|
||||
await asyncio.gather(*tasks, return_exceptions=True)
|
||||
return build_discovery(
|
||||
configured=configured_versions(),
|
||||
revision=context.protocol_version or "2025-11-25",
|
||||
transport=MCPTransport.http,
|
||||
authorized_operations=GATEWAY_OPERATIONS,
|
||||
upstream_versions=frozenset(MCP_LEGACY_VERSIONS),
|
||||
capabilities=ServerCapabilities(
|
||||
tools=ToolsCapability()
|
||||
if any(isinstance(result, ListToolsResult) and result.tools for result in results)
|
||||
else None,
|
||||
prompts=PromptsCapability()
|
||||
if any(isinstance(result, ListPromptsResult) and result.prompts for result in results)
|
||||
else None,
|
||||
resources=ResourcesCapability()
|
||||
if any(
|
||||
(isinstance(result, ListResourcesResult) and result.resources)
|
||||
or (isinstance(result, ListResourceTemplatesResult) and result.resource_templates)
|
||||
for result in results
|
||||
)
|
||||
else None,
|
||||
),
|
||||
)
|
||||
case AuthorizedToolCall():
|
||||
auth, token, _servers, server_headers, oauth_headers, headers, _client_ip = context.legacy_auth()
|
||||
return await _execute_mcp_tool(
|
||||
|
|
|
|||
|
|
@ -1375,7 +1375,16 @@ if MCP_AVAILABLE:
|
|||
and headers.get(MCPRequestHandler.LITELLM_API_KEY_HEADER_NAME_PRIMARY)
|
||||
else None
|
||||
)
|
||||
return _StagedServerTest(request=request, mcp_auth_header=mcp_auth_header, oauth2_headers=oauth2_headers)
|
||||
preview_request: Final = (
|
||||
request.model_copy(
|
||||
update={"mcp_info": {**(request.mcp_info or {}), "protocol_version": saved_server.protocol_version}}
|
||||
)
|
||||
if saved_server is not None and "protocol_version" not in (request.mcp_info or {})
|
||||
else request
|
||||
)
|
||||
return _StagedServerTest(
|
||||
request=preview_request, mcp_auth_header=mcp_auth_header, oauth2_headers=oauth2_headers
|
||||
)
|
||||
|
||||
async def _list_tools_within(client: MCPClient, deadline: float) -> list[MCPTool] | None:
|
||||
with anyio.move_on_after(deadline):
|
||||
|
|
@ -1512,6 +1521,7 @@ if MCP_AVAILABLE:
|
|||
extra_headers=merged_headers,
|
||||
stdio_env=stdio_env,
|
||||
cred_provider=preview_cred_provider,
|
||||
protocol_version_override=server_model.protocol_version,
|
||||
)
|
||||
|
||||
return await operation(client)
|
||||
|
|
|
|||
|
|
@ -125,12 +125,14 @@ def unsupported_protocol_version(scope: Scope) -> str | None:
|
|||
``HANDSHAKE_PROTOCOL_VERSIONS`` to the modern single-exchange path, which
|
||||
bypasses litellm's session/auth model, so the ASGI entry rejects it.
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.capabilities import configured_versions
|
||||
|
||||
headers: Final[Iterable[tuple[bytes, bytes]]] = scope.get("headers") or ()
|
||||
values: Final = tuple(
|
||||
raw.decode("latin-1").strip() for key, raw in headers if key.lower() == _MCP_PROTOCOL_VERSION_HEADER
|
||||
)
|
||||
for value in values:
|
||||
if value and value not in HANDSHAKE_PROTOCOL_VERSIONS:
|
||||
if value and value not in configured_versions():
|
||||
return value
|
||||
return None
|
||||
|
||||
|
|
@ -149,7 +151,10 @@ try:
|
|||
from mcp.server.session import ServerSession as _McpServerSession
|
||||
from mcp.types import (
|
||||
BlobResourceContents,
|
||||
DiscoverRequest,
|
||||
DiscoverResult,
|
||||
GetPromptResult,
|
||||
RequestParams,
|
||||
ResourceTemplate,
|
||||
TextResourceContents,
|
||||
)
|
||||
|
|
@ -526,11 +531,11 @@ if MCP_AVAILABLE:
|
|||
PaginatedRequestParams,
|
||||
ReadResourceRequestParams,
|
||||
)
|
||||
from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.auth.litellm_auth_handler import (
|
||||
MCPAuthenticatedUser,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.capabilities import GatewayVersionPolicy, configured_versions
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
MCPServerManager,
|
||||
global_mcp_server_manager,
|
||||
|
|
@ -585,6 +590,7 @@ if MCP_AVAILABLE:
|
|||
name=LITELLM_MCP_SERVER_NAME,
|
||||
version=LITELLM_MCP_SERVER_VERSION,
|
||||
)
|
||||
server.middleware.append(GatewayVersionPolicy())
|
||||
server.create_initialization_options = types.MethodType(_gateway_create_initialization_options, server)
|
||||
sse: Final[SseServerTransport] = SseServerTransport("/sse/messages")
|
||||
|
||||
|
|
@ -830,6 +836,7 @@ if MCP_AVAILABLE:
|
|||
client_ip,
|
||||
_mcp_proxy_mode.get(),
|
||||
wire_compat_for(ctx.protocol_version),
|
||||
ctx.protocol_version,
|
||||
)
|
||||
|
||||
async def handle_list_tools(ctx: ServerRequestContext, params: PaginatedRequestParams) -> ListToolsResult:
|
||||
|
|
@ -948,6 +955,11 @@ if MCP_AVAILABLE:
|
|||
ReadResourceRequest(params=params), context
|
||||
)
|
||||
|
||||
async def discover(ctx: ServerRequestContext, params: RequestParams) -> DiscoverResult:
|
||||
async with _legacy_operation_context(ctx, trace=False) as context:
|
||||
return await operations.GatewayOperations().execute(DiscoverRequest(params=params), context)
|
||||
|
||||
server.add_request_handler("server/discover", RequestParams, discover)
|
||||
server.add_request_handler("tools/list", PaginatedRequestParams, handle_list_tools)
|
||||
server.add_request_handler("tools/call", CallToolRequestParams, mcp_server_tool_call)
|
||||
server.add_request_handler("prompts/list", PaginatedRequestParams, list_prompts)
|
||||
|
|
@ -1954,7 +1966,7 @@ if MCP_AVAILABLE:
|
|||
reject_disallowed_mcp_origin(StarletteRequest(scope))
|
||||
bad_version: Final = unsupported_protocol_version(scope)
|
||||
if bad_version is not None:
|
||||
supported: Final = ", ".join(sorted(HANDSHAKE_PROTOCOL_VERSIONS))
|
||||
supported: Final = ", ".join(configured_versions())
|
||||
await JSONResponse(
|
||||
status_code=400,
|
||||
content={ # mutable-ok: JSON-RPC error payload
|
||||
|
|
@ -2299,7 +2311,7 @@ if MCP_AVAILABLE:
|
|||
reject_disallowed_mcp_origin(StarletteRequest(scope))
|
||||
bad_version: Final = unsupported_protocol_version(scope)
|
||||
if bad_version is not None:
|
||||
supported: Final = ", ".join(sorted(HANDSHAKE_PROTOCOL_VERSIONS))
|
||||
supported: Final = ", ".join(configured_versions())
|
||||
await JSONResponse(
|
||||
status_code=400,
|
||||
content={ # mutable-ok: JSON-RPC error payload
|
||||
|
|
|
|||
|
|
@ -37,6 +37,7 @@ from litellm.types.llms.openai import (
|
|||
ResponsesAPIResponse,
|
||||
)
|
||||
from litellm.types.mcp import (
|
||||
MCPAdvertisedVersions,
|
||||
MCPAllowedClient,
|
||||
MCPAuth,
|
||||
MCPAuthType,
|
||||
|
|
@ -2998,6 +2999,11 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
|
|||
None,
|
||||
description="Custom CIDR ranges that define internal/private networks for MCP access control. When set, only these ranges are treated as internal. Defaults to RFC 1918 private ranges (10.0.0.0/8, 172.16.0.0/12, 192.168.0.0/16, 127.0.0.0/8).",
|
||||
)
|
||||
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.",
|
||||
)
|
||||
mcp_allowed_clients: list[MCPAllowedClient] | None = Field(
|
||||
None,
|
||||
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.",
|
||||
|
|
|
|||
|
|
@ -6342,6 +6342,11 @@ class ProxyConfig:
|
|||
if general_settings is None:
|
||||
general_settings = {}
|
||||
|
||||
if general_settings.get("mcp_advertised_versions") is not None:
|
||||
from litellm.types.mcp import MCPAdvertisedVersions
|
||||
|
||||
TypeAdapter(MCPAdvertisedVersions).validate_python(general_settings["mcp_advertised_versions"])
|
||||
|
||||
if os.getenv("NUM_WORKERS", "1") != "1" and redis_usage_cache is None:
|
||||
warn_login_counters_are_per_worker(os.getenv("NUM_WORKERS", "1"))
|
||||
if declared_proxy_ranges(general_settings) is 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, Any, Final, Literal
|
||||
from typing import TYPE_CHECKING, Annotated, Any, Final, Literal
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
import httpx
|
||||
|
|
@ -34,6 +34,8 @@ class MCPSpecVersion(str, enum.Enum):
|
|||
nov_2024 = "2024-11-05"
|
||||
mar_2025 = "2025-03-26"
|
||||
jun_2025 = "2025-06-18"
|
||||
nov_2025 = "2025-11-25"
|
||||
jul_2026 = "2026-07-28"
|
||||
|
||||
|
||||
class MCPAuth(str, enum.Enum):
|
||||
|
|
@ -59,7 +61,17 @@ DEFAULT_SUBJECT_TOKEN_TYPE: Final = "urn:ietf:params:oauth:token-type:access_tok
|
|||
|
||||
# MCP Literals
|
||||
MCPTransportType = Literal[MCPTransport.sse, MCPTransport.http, MCPTransport.stdio]
|
||||
MCPSpecVersionType = Literal[MCPSpecVersion.nov_2024, MCPSpecVersion.mar_2025, MCPSpecVersion.jun_2025]
|
||||
MCPLegacyVersion = Literal["2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"]
|
||||
MCP_LEGACY_VERSIONS: Final[tuple[MCPLegacyVersion, ...]] = ("2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25")
|
||||
MCPUpstreamProtocol = MCPLegacyVersion | Literal["auto"]
|
||||
MCPAdvertisedVersions = Annotated[tuple[MCPLegacyVersion, ...], Field(min_length=1)]
|
||||
MCPSpecVersionType = Literal[
|
||||
MCPSpecVersion.nov_2024,
|
||||
MCPSpecVersion.mar_2025,
|
||||
MCPSpecVersion.jun_2025,
|
||||
MCPSpecVersion.nov_2025,
|
||||
MCPSpecVersion.jul_2026,
|
||||
]
|
||||
MCPAuthType = (
|
||||
Literal[
|
||||
MCPAuth.none,
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
from datetime import datetime
|
||||
from typing import Any, Final, Literal
|
||||
from typing import Annotated, Any, Final, Literal
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
|
||||
from pydantic import AfterValidator, BaseModel, ConfigDict, Field, TypeAdapter, field_validator, model_validator
|
||||
from typing_extensions import Self
|
||||
|
||||
from litellm.types.mcp import (
|
||||
|
|
@ -10,11 +10,19 @@ from litellm.types.mcp import (
|
|||
MCPAuthType,
|
||||
MCPTokenEndpointAuthMethod,
|
||||
MCPTransportType,
|
||||
MCPUpstreamProtocol,
|
||||
normalize_upstream_header_name,
|
||||
)
|
||||
|
||||
|
||||
# MCPInfo now allows arbitrary additional fields for custom metadata
|
||||
MCPInfo = dict[str, Any]
|
||||
def _validate_mcp_protocol_metadata(value: dict[str, object]) -> dict[str, object]:
|
||||
if "protocol_version" in value:
|
||||
TypeAdapter(MCPUpstreamProtocol).validate_python(value["protocol_version"])
|
||||
return value
|
||||
|
||||
|
||||
MCPInfo = Annotated[dict[str, Any], AfterValidator(_validate_mcp_protocol_metadata)]
|
||||
|
||||
|
||||
class MCPOAuthMetadata(BaseModel):
|
||||
|
|
@ -66,6 +74,7 @@ class MCPServer(BaseModel):
|
|||
server_name: str | None = None
|
||||
url: str | None = None
|
||||
transport: MCPTransportType
|
||||
protocol_version: MCPUpstreamProtocol = "auto"
|
||||
spec_path: str | None = None
|
||||
auth_type: MCPAuthType | None = None
|
||||
authentication_token: str | None = None
|
||||
|
|
@ -246,6 +255,14 @@ class MCPServer(BaseModel):
|
|||
"""
|
||||
return self.oauth2_flow == "client_credentials"
|
||||
|
||||
@model_validator(mode="after")
|
||||
def resolve_protocol_version(self) -> Self:
|
||||
if "protocol_version" not in self.model_fields_set and self.mcp_info is not None:
|
||||
self.protocol_version = TypeAdapter(MCPUpstreamProtocol).validate_python(
|
||||
self.mcp_info.get("protocol_version", "auto")
|
||||
)
|
||||
return self
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_identity_binding_mode(self) -> Self:
|
||||
binding: Final = self.oauth_identity_binding
|
||||
|
|
|
|||
|
|
@ -84,3 +84,42 @@ def test_jsonrpc_error_and_malformed_tool_result_remain_errors(gateway: Gateway)
|
|||
control: Final = call_tool(gateway, key, identity, names["add"], {"a": 3, "b": 5})
|
||||
assert control.status_code == 200 and control.json()["isError"] is False, control.text
|
||||
assert control.json()["content"][0]["text"] == "8"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("ingress", ("http", "sse"))
|
||||
def test_configured_revision_blocks_unadvertised_handshake_and_keeps_allowed_control(
|
||||
gateway: Gateway, tmp_path, ingress: str
|
||||
) -> None:
|
||||
import asyncio
|
||||
from pathlib import Path
|
||||
|
||||
import yaml
|
||||
from integration._support.mcp import mcp_peer
|
||||
from integration._support.process import owned_proxy
|
||||
from litellm.experimental_mcp_client.client import MCPClient
|
||||
from litellm.types.mcp import MCPTransport
|
||||
from mcp import MCPError
|
||||
from mcp.types import CallToolRequestParams
|
||||
|
||||
with mcp_peer() as upstream, gateway.scenario() as scenario:
|
||||
alias: Final = "restricted" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, upstream, alias)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["general_settings"]["mcp_advertised_versions"] = ["2024-11-05"]
|
||||
config_path: Final = tmp_path / "restricted.yaml"
|
||||
config_path.write_text(yaml.safe_dump(config))
|
||||
with owned_proxy(gateway, tmp_path, {"DISABLE_SCHEMA_UPDATE": "true"}, config=config_path) as restricted:
|
||||
endpoint: Final = str(restricted.client.base_url).rstrip("/") + ("/mcp/sse" if ingress == "sse" else "/mcp")
|
||||
headers: Final = {"Authorization": f"Bearer {key}", "x-mcp-servers": identity}
|
||||
denied: Final = MCPClient(server_url=endpoint, transport_type=MCPTransport(ingress), protocol_version="2025-11-25", extra_headers=headers)
|
||||
allowed: Final = MCPClient(server_url=endpoint, transport_type=MCPTransport(ingress), protocol_version="2024-11-05", extra_headers=headers)
|
||||
|
||||
async def exercise() -> None:
|
||||
with pytest.raises(MCPError, match="Unsupported MCP protocol version"):
|
||||
await denied.list_tools(raise_on_error=True)
|
||||
assert f"{alias}-add" in tuple(tool.name for tool in await allowed.list_tools(raise_on_error=True))
|
||||
result: Final = await allowed.call_tool(CallToolRequestParams(name=f"{alias}-add", arguments={"a": 2, "b": 5}))
|
||||
assert result.is_error is False and result.content[0].text == "7"
|
||||
|
||||
asyncio.run(exercise())
|
||||
|
|
|
|||
|
|
@ -155,3 +155,43 @@ def test_server_initiated_sampling_and_elicitation_surface_as_errors_not_success
|
|||
assert outcome.error is not None, outcome.raw
|
||||
assert outcome.text is None or not outcome.text.startswith(("sampled:", "elicited:")), outcome.raw
|
||||
assert len(tool_calls(peer.drain())) == 1
|
||||
|
||||
|
||||
@pytest.mark.parametrize("downstream", ("2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"))
|
||||
@pytest.mark.parametrize("upstream", ("2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"))
|
||||
@pytest.mark.parametrize("peer_kind", ("http", "sse", "stdio"))
|
||||
@pytest.mark.parametrize("ingress", ("http", "sse"))
|
||||
def test_pinned_revision_pairs_list_and_call_through_gateway(
|
||||
gateway: Gateway, downstream: str, upstream: str, peer_kind: PeerKind, ingress: str
|
||||
) -> None:
|
||||
import asyncio
|
||||
|
||||
from mcp.types import CallToolRequestParams
|
||||
|
||||
from litellm.experimental_mcp_client.client import MCPClient
|
||||
from litellm.types.mcp import MCPTransport
|
||||
|
||||
with peer_of(peer_kind) as peer, gateway.scenario() as scenario:
|
||||
alias: Final = "versions" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias, mcp_info={"protocol_version": upstream})
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
endpoint: Final = str(gateway.client.base_url).rstrip("/") + ("/mcp/sse" if ingress == "sse" else "/mcp")
|
||||
client: Final = MCPClient(
|
||||
server_url=endpoint, transport_type=MCPTransport(ingress), protocol_version=downstream,
|
||||
extra_headers={"Authorization": f"Bearer {key}", "x-mcp-servers": identity}, timeout=15,
|
||||
)
|
||||
|
||||
async def exercise() -> None:
|
||||
tools: Final = await client.list_tools(raise_on_error=True)
|
||||
assert f"{alias}-add" in tuple(tool.name for tool in tools)
|
||||
result: Final = await client.call_tool(CallToolRequestParams(name=f"{alias}-add", arguments={"a": 3, "b": 4}))
|
||||
assert result.is_error is False
|
||||
assert result.content[0].text == "7"
|
||||
|
||||
peer.drain()
|
||||
asyncio.run(exercise())
|
||||
observed: Final = peer.drain()
|
||||
negotiations: Final = tuple(item["body"] for item in observed if item["body"].get("method") == "initialize")
|
||||
assert negotiations, "The operation must reach the upstream negotiation"
|
||||
assert all(request["params"]["protocolVersion"] == upstream for request in negotiations), negotiations
|
||||
assert len(tool_calls(observed)) == 1
|
||||
|
|
|
|||
|
|
@ -0,0 +1,106 @@
|
|||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from mcp import Client
|
||||
from mcp.server import Server
|
||||
from mcp.shared.exceptions import MCPError
|
||||
from mcp.types import PromptsCapability, ResourcesCapability, ServerCapabilities, ToolsCapability
|
||||
from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.capabilities import (
|
||||
GATEWAY_OPERATIONS,
|
||||
REVISION_SUPPORT,
|
||||
TRANSLATION_PAIRS,
|
||||
GatewayVersionPolicy,
|
||||
build_discovery,
|
||||
)
|
||||
from litellm.types.mcp import MCPTransport
|
||||
|
||||
|
||||
@pytest.mark.parametrize("revision", HANDSHAKE_PROTOCOL_VERSIONS)
|
||||
@pytest.mark.parametrize("transport", tuple(MCPTransport))
|
||||
def test_discovery_only_exposes_authorized_completed_support(revision, transport):
|
||||
result = build_discovery(
|
||||
configured=(revision, "2026-07-28", "unknown"),
|
||||
revision=revision,
|
||||
transport=transport,
|
||||
authorized_operations=frozenset({"tools/list", "tools/call"}),
|
||||
upstream_versions=frozenset(HANDSHAKE_PROTOCOL_VERSIONS),
|
||||
capabilities=ServerCapabilities(
|
||||
tools=ToolsCapability(), prompts=PromptsCapability(), resources=ResourcesCapability(),
|
||||
extensions={"io.modelcontextprotocol/ui": {}},
|
||||
),
|
||||
client_extensions=frozenset({"io.modelcontextprotocol/ui"}),
|
||||
upstream_extensions=frozenset({"io.modelcontextprotocol/ui"}),
|
||||
)
|
||||
assert result.supported_versions == [revision]
|
||||
assert result.capabilities.tools is not None
|
||||
assert result.capabilities.prompts is None
|
||||
assert result.capabilities.resources is None
|
||||
assert result.capabilities.extensions is None
|
||||
assert result.capabilities.tasks is None
|
||||
assert result.cache_scope == "private"
|
||||
assert result.ttl_ms == 0
|
||||
|
||||
|
||||
@pytest.mark.parametrize("upstream", [frozenset(), frozenset({"unknown"}), frozenset({"2026-07-28"})])
|
||||
def test_unproven_translation_never_advertises_operations(upstream):
|
||||
result = build_discovery(
|
||||
configured=HANDSHAKE_PROTOCOL_VERSIONS,
|
||||
revision="2025-11-25",
|
||||
transport=MCPTransport.http,
|
||||
authorized_operations=GATEWAY_OPERATIONS,
|
||||
upstream_versions=upstream,
|
||||
capabilities=ServerCapabilities(tools=ToolsCapability()),
|
||||
)
|
||||
assert result.capabilities.tools is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize("revision", ["2026-07-28", "unknown", "2024-11-05"])
|
||||
def test_unadvertised_revision_never_gains_capabilities(revision):
|
||||
result = build_discovery(
|
||||
configured=("2025-11-25",), revision=revision, transport=MCPTransport.http,
|
||||
authorized_operations=GATEWAY_OPERATIONS, upstream_versions=frozenset(HANDSHAKE_PROTOCOL_VERSIONS),
|
||||
capabilities=ServerCapabilities(tools=ToolsCapability()),
|
||||
)
|
||||
assert result.capabilities.tools is None
|
||||
|
||||
|
||||
def test_discovery_results_do_not_share_mutable_capabilities():
|
||||
capabilities = ServerCapabilities(tools=ToolsCapability(), prompts=PromptsCapability(), resources=ResourcesCapability())
|
||||
args = dict(
|
||||
configured=HANDSHAKE_PROTOCOL_VERSIONS, revision="2025-11-25", transport=MCPTransport.http,
|
||||
upstream_versions=frozenset(HANDSHAKE_PROTOCOL_VERSIONS), capabilities=capabilities,
|
||||
)
|
||||
allowed = build_discovery(**args, authorized_operations=GATEWAY_OPERATIONS)
|
||||
denied = build_discovery(**args, authorized_operations=frozenset())
|
||||
assert allowed.capabilities.prompts is not None
|
||||
assert allowed.capabilities.resources is not None
|
||||
assert denied.capabilities.model_dump(exclude_none=True) == {}
|
||||
assert allowed.capabilities.tools is not None
|
||||
allowed.capabilities.tools.list_changed = True
|
||||
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
|
||||
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)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("versions,accepted", [(("2025-11-25",), True), (("2025-06-18",), False)])
|
||||
async def test_version_policy_gates_the_actual_sdk_handshake(versions, accepted):
|
||||
server: Final = Server("test-gateway", version="1")
|
||||
server.middleware.append(GatewayVersionPolicy(lambda: versions))
|
||||
if accepted:
|
||||
async with Client(server, mode="legacy") as client:
|
||||
assert client.protocol_version == "2025-11-25"
|
||||
result = await client.session.send_ping()
|
||||
assert result is not None
|
||||
else:
|
||||
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")
|
||||
|
|
@ -10444,11 +10444,12 @@ async def test_active_request_ctx_var_feeds_auth_resolution_recording(_mcp_reque
|
|||
)
|
||||
@pytest.mark.parametrize("handler", ("handle_streamable_http_mcp", "handle_sse_mcp"))
|
||||
async def test_streamable_http_rejects_modern_protocol_version(
|
||||
header_value: str, expected_rejected: bool, handler: str
|
||||
header_value: str, expected_rejected: bool, handler: str, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
from litellm.proxy._experimental.mcp_server import server as mcp_module
|
||||
from litellm.proxy._experimental.mcp_server.server import unsupported_protocol_version
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
|
||||
scope: Scope = {
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
|
|
@ -10659,3 +10660,32 @@ async def test_legacy_sse_mount_emits_message_endpoint(
|
|||
await incoming.put({"type": "http.disconnect"})
|
||||
await asyncio.wait_for(task, 2)
|
||||
assert await post(initialization) == 404
|
||||
|
||||
|
||||
@pytest.mark.parametrize("revision,rejected", [("2024-11-05", False), ("2025-11-25", True), ("2026-07-28", True)])
|
||||
def test_protocol_header_respects_configured_advertisement(revision, rejected):
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._experimental.mcp_server.server import unsupported_protocol_version
|
||||
|
||||
with patch.dict(proxy_server.general_settings, {"mcp_advertised_versions": ["2024-11-05"]}):
|
||||
result = unsupported_protocol_version({"headers": [(b"mcp-protocol-version", revision.encode())]})
|
||||
assert result == (revision if rejected else None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_discovery_adapter_preserves_authenticated_context(_mcp_request_ctx):
|
||||
from mcp.types import DiscoverResult, RequestParams, ServerCapabilities
|
||||
from litellm.proxy._experimental.mcp_server import server
|
||||
|
||||
expected = DiscoverResult(supported_versions=["2025-11-25"], capabilities=ServerCapabilities())
|
||||
dispatched = AsyncMock(return_value=expected)
|
||||
auth = UserAPIKeyAuth(user_id="discover-caller")
|
||||
with (
|
||||
patch.object(server, "get_or_extract_auth_context", AsyncMock(return_value=(auth, None, ["allowed"], None, None, None, None))),
|
||||
patch.object(server.operations.GatewayOperations, "execute", dispatched),
|
||||
):
|
||||
result = await server.discover(_mcp_request_ctx(), RequestParams())
|
||||
assert result is expected
|
||||
context = dispatched.await_args.args[1]
|
||||
assert context.user_api_key_auth.user_id == "discover-caller"
|
||||
assert context.mcp_servers == ("allowed",)
|
||||
|
|
|
|||
|
|
@ -63,7 +63,7 @@ from litellm.proxy._types import (
|
|||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
from litellm.types.mcp import MCPAuth, MCPAuthType
|
||||
from litellm.types.mcp import MCPAuth, MCPAuthType, MCPUpstreamProtocol
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.caching.llm_caching_handler import LLMClientCache
|
||||
|
|
@ -14637,3 +14637,27 @@ class TestSharedIdentifierPrefixWarning:
|
|||
assert "srv-b" in shared_warnings[0]
|
||||
assert "srv-c" not in shared_warnings[0]
|
||||
assert "'shared'" in shared_warnings[0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("revision", ["auto", "2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"])
|
||||
async def test_configured_protocol_reaches_the_upstream_client(config_only_mcp_manager_factory, revision):
|
||||
manager = config_only_mcp_manager_factory()
|
||||
await manager.load_servers_from_config({"versions": {"url": "http://127.0.0.1:9/mcp", "transport": "http", "protocol_version": revision}})
|
||||
server = next(iter(manager.config_mcp_servers.values()))
|
||||
client = await manager._create_mcp_client(server)
|
||||
assert server.protocol_version == revision
|
||||
assert client.protocol_version == revision
|
||||
|
||||
|
||||
@pytest.mark.parametrize("revision", ("auto", "2024-11-05", "2025-06-18"))
|
||||
@pytest.mark.parametrize("explicit", (None, "auto", "2025-11-25"))
|
||||
def test_runtime_protocol_metadata_preserves_explicit_precedence(
|
||||
revision: MCPUpstreamProtocol, explicit: MCPUpstreamProtocol | None
|
||||
) -> None:
|
||||
server: Final = MCPServer.model_validate({
|
||||
"server_id": "preview", "name": "preview", "transport": "http",
|
||||
"mcp_info": {"protocol_version": revision},
|
||||
**({"protocol_version": explicit} if explicit is not None else {}),
|
||||
})
|
||||
assert server.protocol_version == (explicit if explicit is not None else revision)
|
||||
|
|
|
|||
|
|
@ -542,3 +542,126 @@ async def test_local_tool_json_array_is_converted_once_for_the_caller_revision(c
|
|||
|
||||
assert [block.text for block in result.content] == [body]
|
||||
assert result.structured_content == (["a", "b"] if compat == "modern" else None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_discovery_preserves_caller_scope_and_proxy_restrictions():
|
||||
from mcp.types import DiscoverRequest, ListToolsResult, Tool
|
||||
|
||||
listed = AsyncMock(return_value=ListToolsResult(tools=[Tool(name="allowed", input_schema={"type": "object"})]))
|
||||
context = prepare_context(UserAPIKeyAuth(user_id="scoped"), mcp_servers=["only-this"], mcp_proxy_mode=True, protocol_version="2025-06-18")
|
||||
with patch("litellm.proxy._experimental.mcp_server.operations._execute_handle_list_tools", listed):
|
||||
result = await GatewayOperations().execute(DiscoverRequest(), context)
|
||||
assert result.capabilities.tools is not None
|
||||
assert result.capabilities.resources is None
|
||||
assert result.capabilities.prompts is None
|
||||
assert listed.await_args.args[0] is context
|
||||
assert listed.await_args.args[0].user_api_key_auth.user_id == "scoped"
|
||||
assert listed.await_args.args[0].mcp_servers == ("only-this",)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_discovery_denial_cannot_advertise_tools():
|
||||
from mcp.types import DiscoverRequest
|
||||
from fastapi import HTTPException
|
||||
|
||||
denied = AsyncMock(side_effect=HTTPException(status_code=403, detail="Forbidden"))
|
||||
with patch("litellm.proxy._experimental.mcp_server.operations._execute_handle_list_tools", denied):
|
||||
with pytest.raises(HTTPException) as error:
|
||||
await GatewayOperations().execute(DiscoverRequest(), prepare_context(UserAPIKeyAuth(user_id="denied")))
|
||||
assert error.value.status_code == 403
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("available", ["none", "resources", "templates", "prompts"])
|
||||
async def test_discovery_lists_each_capability_with_the_same_caller(available):
|
||||
from mcp.types import (
|
||||
DiscoverRequest, ListToolsResult, ListPromptsResult, ListResourcesResult,
|
||||
ListResourceTemplatesResult, Prompt, Resource, ResourceTemplate,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server import operations
|
||||
|
||||
context = prepare_context(UserAPIKeyAuth(user_id="scoped"), mcp_servers=["authorized"])
|
||||
tools = AsyncMock(return_value=ListToolsResult(tools=[]))
|
||||
prompts = AsyncMock(return_value=ListPromptsResult(prompts=[Prompt(name="allowed")] if available == "prompts" else []))
|
||||
resources = AsyncMock(return_value=ListResourcesResult(resources=[Resource(name="allowed", uri="test://allowed")] if available == "resources" else []))
|
||||
templates = AsyncMock(return_value=ListResourceTemplatesResult(resource_templates=[ResourceTemplate(name="allowed", uri_template="test://{id}")] if available == "templates" else []))
|
||||
with (
|
||||
patch.object(operations, "_execute_handle_list_tools", tools),
|
||||
patch.object(operations, "_execute_list_prompts", prompts),
|
||||
patch.object(operations, "_execute_list_resources", resources),
|
||||
patch.object(operations, "_execute_list_resource_templates", templates),
|
||||
):
|
||||
result = await GatewayOperations().execute(DiscoverRequest(), context)
|
||||
assert result.capabilities.tools is None
|
||||
assert (result.capabilities.prompts is not None) == (available == "prompts")
|
||||
assert (result.capabilities.resources is not None) == (available in {"resources", "templates"})
|
||||
for listing in (tools, prompts, resources, templates):
|
||||
assert listing.await_args.args[0] is context
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("outcome", ["success", "failure", "cancel"])
|
||||
async def test_discovery_concurrent_listings_drain_on_failure_and_cancellation(outcome):
|
||||
from mcp.types import DiscoverRequest, ListToolsResult, ListPromptsResult, ListResourcesResult, ListResourceTemplatesResult
|
||||
from litellm.proxy._experimental.mcp_server import operations
|
||||
|
||||
ready = [asyncio.Event() for _ in range(4)]
|
||||
closed = [asyncio.Event() for _ in range(4)]
|
||||
release = asyncio.Event()
|
||||
responses = (ListToolsResult(tools=[]), ListPromptsResult(prompts=[]), ListResourcesResult(resources=[]), ListResourceTemplatesResult(resource_templates=[]))
|
||||
|
||||
def listing(index):
|
||||
async def run(*args, **kwargs):
|
||||
ready[index].set()
|
||||
try:
|
||||
await release.wait()
|
||||
if index == 0 and outcome == "failure":
|
||||
raise ValueError("discovery failed")
|
||||
if outcome != "success":
|
||||
await asyncio.Event().wait()
|
||||
return responses[index]
|
||||
finally:
|
||||
closed[index].set()
|
||||
return run
|
||||
|
||||
with (
|
||||
patch.object(operations, "_execute_handle_list_tools", side_effect=listing(0)) as tools,
|
||||
patch.object(operations, "_execute_list_prompts", side_effect=listing(1)),
|
||||
patch.object(operations, "_execute_list_resources", side_effect=listing(2)),
|
||||
patch.object(operations, "_execute_list_resource_templates", side_effect=listing(3)),
|
||||
):
|
||||
task = asyncio.create_task(GatewayOperations().execute(DiscoverRequest(), prepare_context(UserAPIKeyAuth(user_id="scoped"))))
|
||||
try:
|
||||
await asyncio.wait_for(asyncio.gather(*(event.wait() for event in ready)), 1)
|
||||
if outcome == "cancel":
|
||||
task.cancel()
|
||||
else:
|
||||
release.set()
|
||||
if outcome == "success":
|
||||
result = await asyncio.wait_for(task, 1)
|
||||
assert result.capabilities.model_dump(exclude_none=True) == {}
|
||||
else:
|
||||
with pytest.raises(asyncio.CancelledError if outcome == "cancel" else ValueError):
|
||||
await asyncio.wait_for(task, 1)
|
||||
assert all(event.is_set() for event in closed)
|
||||
assert tools.call_args.kwargs["log_list_tools_to_spendlogs"] is False
|
||||
finally:
|
||||
task.cancel()
|
||||
await asyncio.gather(task, return_exceptions=True)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("log_enabled", [False, True])
|
||||
async def test_tools_listing_preserves_explicit_spend_log_policy(log_enabled):
|
||||
from mcp.types import PaginatedRequestParams
|
||||
from litellm.proxy._experimental.mcp_server import operations
|
||||
|
||||
listing = AsyncMock(return_value=operations.AggregateToolListing(tools=[], outcomes={}))
|
||||
with patch.object(operations, "_list_mcp_tools", listing):
|
||||
result = await operations._execute_handle_list_tools(
|
||||
prepare_context(UserAPIKeyAuth(user_id="caller")), PaginatedRequestParams(),
|
||||
log_list_tools_to_spendlogs=log_enabled,
|
||||
)
|
||||
assert result.tools == []
|
||||
assert listing.await_args.kwargs["log_list_tools_to_spendlogs"] is log_enabled
|
||||
|
|
|
|||
|
|
@ -28,7 +28,7 @@ from litellm.proxy._types import (
|
|||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.types.mcp import MCPAuth, MCPTransport
|
||||
from litellm.types.mcp import MCPAuth, MCPTransport, MCPUpstreamProtocol
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
_OK_TOOL_RESULT: Final = CallToolResult(content=[TextContent(type="text", text='{"result": "ok"}')], is_error=False)
|
||||
|
|
@ -1476,7 +1476,7 @@ class TestListToolsRestAPI:
|
|||
monkeypatch,
|
||||
):
|
||||
"""The REST tools/list path should include tools beyond the upstream first page."""
|
||||
from mcp.types import ListToolsResult, PaginatedRequestParams
|
||||
from mcp.types import Implementation, InitializeResult, ListToolsResult, PaginatedRequestParams, ServerCapabilities
|
||||
from mcp.types import Tool as MCPTool
|
||||
|
||||
import litellm.experimental_mcp_client.client as mcp_client_module
|
||||
|
|
@ -1512,7 +1512,11 @@ class TestListToolsRestAPI:
|
|||
|
||||
mock_session_ctx = AsyncMock()
|
||||
mock_session_instance = AsyncMock()
|
||||
mock_session_instance.initialize = AsyncMock(return_value=None)
|
||||
mock_session_instance.initialize = AsyncMock(return_value=InitializeResult(
|
||||
protocol_version="2025-11-25",
|
||||
capabilities=ServerCapabilities(),
|
||||
server_info=Implementation(name="stub", version="1"),
|
||||
))
|
||||
mock_session_instance.list_tools.side_effect = [
|
||||
ListToolsResult(
|
||||
tools=[
|
||||
|
|
@ -4628,3 +4632,74 @@ class TestClientAllowlistOnRestRoutes:
|
|||
assert denied.value.detail["error"] == "Forbidden"
|
||||
assert "'claude-code'" in denied.value.detail["details"]
|
||||
acting.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("revision", ("auto", "2024-11-05", "2025-06-18"))
|
||||
async def test_preview_client_honors_protocol_metadata(revision: MCPUpstreamProtocol) -> None:
|
||||
from litellm.experimental_mcp_client.client import MCPClient
|
||||
|
||||
payload: Final = NewMCPServerRequest(
|
||||
server_name="preview", url="http://127.0.0.1:9/mcp", transport="http",
|
||||
auth_type=MCPAuth.none, mcp_info={"protocol_version": revision},
|
||||
)
|
||||
|
||||
async def inspect_client(client: MCPClient) -> dict[str, str]:
|
||||
return {"protocol_version": client.protocol_version}
|
||||
|
||||
result: Final = await rest_endpoints._execute_with_mcp_client(payload, inspect_client)
|
||||
assert result == {"protocol_version": revision}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type", (MCPAuth.none, MCPAuth.bearer_token, MCPAuth.oauth2))
|
||||
@pytest.mark.parametrize(
|
||||
("metadata", "expected"),
|
||||
(
|
||||
(None, "2025-11-25"),
|
||||
({}, "2025-11-25"),
|
||||
({"description": "edited"}, "2025-11-25"),
|
||||
({"protocol_version": "auto"}, "auto"),
|
||||
({"protocol_version": "2024-11-05"}, "2024-11-05"),
|
||||
({"protocol_version": "2025-06-18"}, "2025-06-18"),
|
||||
),
|
||||
)
|
||||
async def test_saved_preview_protocol_omission_and_explicit_edits(
|
||||
monkeypatch: pytest.MonkeyPatch, auth_type: MCPAuth,
|
||||
metadata: dict[str, str] | None, expected: MCPUpstreamProtocol,
|
||||
) -> None:
|
||||
from starlette.datastructures import Headers
|
||||
|
||||
from litellm.experimental_mcp_client.client import MCPClient
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
from litellm.proxy.management_endpoints import mcp_management_endpoints
|
||||
|
||||
saved: Final = MCPServer(
|
||||
server_id="saved-preview", name="preview", url="https://example.com/mcp",
|
||||
transport="http", auth_type=auth_type, protocol_version="2025-11-25",
|
||||
authentication_token="stored-token",
|
||||
authorization_url="https://example.com/authorize", token_url="https://example.com/token",
|
||||
)
|
||||
manager: Final = MCPServerManager()
|
||||
manager.registry = {saved.server_id: saved}
|
||||
monkeypatch.setattr(rest_endpoints, "global_mcp_server_manager", manager)
|
||||
monkeypatch.setattr(mcp_management_endpoints, "global_mcp_server_manager", manager)
|
||||
payload: Final = NewMCPServerRequest(
|
||||
server_id=saved.server_id, server_name=saved.name, url=saved.url, transport="http",
|
||||
auth_type=auth_type, mcp_info=metadata,
|
||||
authorization_url=saved.authorization_url, token_url=saved.token_url,
|
||||
)
|
||||
staged: Final = rest_endpoints._stage_server_test(
|
||||
payload, Headers({"x-litellm-api-key": "sk-admin", "authorization": "Bearer preview-token"})
|
||||
)
|
||||
|
||||
async def inspect_client(client: MCPClient) -> dict[str, str]:
|
||||
return {"protocol_version": client.protocol_version}
|
||||
|
||||
result: Final = await rest_endpoints._execute_with_mcp_client(
|
||||
staged.request, inspect_client,
|
||||
mcp_auth_header=staged.mcp_auth_header, oauth2_headers=staged.oauth2_headers,
|
||||
)
|
||||
assert result == {"protocol_version": expected}
|
||||
assert saved.protocol_version == "2025-11-25"
|
||||
assert payload.mcp_info == metadata
|
||||
|
|
|
|||
|
|
@ -4903,3 +4903,19 @@ async def test_model_refresh_updates_availability_catalog_and_retains_it_on_db_f
|
|||
assert await pc._get_models_from_db(client) == []
|
||||
assert pc.auto_router_db_catalog == ()
|
||||
assert find_many.await_count == 3
|
||||
|
||||
|
||||
@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):
|
||||
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"]:
|
||||
_, _, settings = await ProxyConfig().load_config(router=None, config_file_path=str(config))
|
||||
assert settings["mcp_advertised_versions"] == versions
|
||||
return
|
||||
with pytest.raises(ValidationError):
|
||||
await ProxyConfig().load_config(router=None, config_file_path=str(config))
|
||||
|
|
|
|||
|
|
@ -377,3 +377,21 @@ 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"]])
|
||||
def test_mcp_advertised_versions_reject_unavailable_revisions(versions):
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm.proxy._types import ConfigGeneralSettings
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
ConfigGeneralSettings(mcp_advertised_versions=versions)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("revision", ["2026-07-28", "unknown", None])
|
||||
def test_mcp_metadata_rejects_unavailable_upstream_protocol(revision):
|
||||
from litellm.proxy._types import NewMCPServerRequest, UpdateMCPServerRequest
|
||||
|
||||
payload = {"server_id": "test", "transport": "http", "url": "https://example.com/mcp", "mcp_info": {"protocol_version": revision}}
|
||||
for model in (NewMCPServerRequest, UpdateMCPServerRequest):
|
||||
with pytest.raises(ValidationError):
|
||||
model.model_validate(payload)
|
||||
|
|
|
|||
|
|
@ -58,6 +58,15 @@ from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
|||
_JSONRPC_MESSAGE_ADAPTER: Final = TypeAdapter(JSONRPCMessage)
|
||||
|
||||
|
||||
def _initialized(instructions: str | None = None) -> InitializeResult:
|
||||
return InitializeResult(
|
||||
protocol_version=LATEST_HANDSHAKE_VERSION,
|
||||
capabilities=ServerCapabilities(),
|
||||
server_info=Implementation(name="test", version="1"),
|
||||
instructions=instructions,
|
||||
)
|
||||
|
||||
|
||||
class _MockTransportClient(MCPClient):
|
||||
"""An MCPClient whose streamable-HTTP transport runs on an httpx2 MockTransport."""
|
||||
|
||||
|
|
@ -125,7 +134,7 @@ class TestMCPClient:
|
|||
mock_stdio_client.return_value = mock_stdio_ctx
|
||||
|
||||
mock_session_instance = AsyncMock()
|
||||
mock_session_instance.initialize = AsyncMock()
|
||||
mock_session_instance.initialize = AsyncMock(return_value=_initialized())
|
||||
mock_session_ctx = AsyncMock()
|
||||
mock_session_ctx.__aenter__.return_value = mock_session_instance
|
||||
mock_session_ctx.__aexit__.return_value = None
|
||||
|
|
@ -168,7 +177,7 @@ class TestMCPClient:
|
|||
# Mock the session
|
||||
with patch("litellm.experimental_mcp_client.client.ClientSession") as mock_session:
|
||||
mock_session_instance = AsyncMock()
|
||||
mock_session_instance.initialize = AsyncMock()
|
||||
mock_session_instance.initialize = AsyncMock(return_value=_initialized())
|
||||
mock_session_ctx = AsyncMock()
|
||||
mock_session_ctx.__aenter__.return_value = mock_session_instance
|
||||
mock_session_ctx.__aexit__.return_value = None
|
||||
|
|
@ -214,7 +223,7 @@ class TestMCPClient:
|
|||
# Mock the session
|
||||
with patch("litellm.experimental_mcp_client.client.ClientSession") as mock_session:
|
||||
mock_session_instance = AsyncMock()
|
||||
mock_session_instance.initialize = AsyncMock()
|
||||
mock_session_instance.initialize = AsyncMock(return_value=_initialized())
|
||||
mock_session_ctx = AsyncMock()
|
||||
mock_session_ctx.__aenter__.return_value = mock_session_instance
|
||||
mock_session_ctx.__aexit__.return_value = None
|
||||
|
|
@ -266,7 +275,7 @@ class TestMCPClient:
|
|||
# Mock the session
|
||||
with patch("litellm.experimental_mcp_client.client.ClientSession") as mock_session:
|
||||
mock_session_instance = AsyncMock()
|
||||
mock_session_instance.initialize = AsyncMock()
|
||||
mock_session_instance.initialize = AsyncMock(return_value=_initialized())
|
||||
mock_session_ctx = AsyncMock()
|
||||
mock_session_ctx.__aenter__.return_value = mock_session_instance
|
||||
mock_session_ctx.__aexit__.return_value = None
|
||||
|
|
@ -413,8 +422,7 @@ class TestMCPClientInstructionsCapture:
|
|||
)
|
||||
|
||||
mock_session = AsyncMock()
|
||||
init_result = MagicMock()
|
||||
init_result.instructions = " upstream says hello "
|
||||
init_result = _initialized(" upstream says hello ")
|
||||
mock_session.initialize = AsyncMock(return_value=init_result)
|
||||
|
||||
session_ctx = MagicMock()
|
||||
|
|
@ -442,8 +450,7 @@ class TestMCPClientInstructionsCapture:
|
|||
)
|
||||
|
||||
mock_session = AsyncMock()
|
||||
init_result = MagicMock()
|
||||
init_result.instructions = None
|
||||
init_result = _initialized()
|
||||
mock_session.initialize = AsyncMock(return_value=init_result)
|
||||
|
||||
session_ctx = MagicMock()
|
||||
|
|
@ -600,8 +607,7 @@ class TestExecuteSessionOperationSurfacesTransportError:
|
|||
@patch("litellm.experimental_mcp_client.client.ClientSession")
|
||||
async def test_cleanup_error_after_success_is_swallowed(self, mock_session_cls):
|
||||
client = MCPClient(server_url="http://example.com/mcp", transport_type="http")
|
||||
init_result = MagicMock()
|
||||
init_result.instructions = None
|
||||
init_result = _initialized()
|
||||
self._make_session(mock_session_cls, AsyncMock(return_value=init_result))
|
||||
transport_ctx = self._make_transport(_FakeExceptionGroup("late", [httpx2.ConnectError("late cleanup error")]))
|
||||
|
||||
|
|
@ -634,7 +640,7 @@ class TestExecuteSessionOperationSurfacesTransportError:
|
|||
@pytest.mark.parametrize("original_error", (False, True))
|
||||
@patch("litellm.experimental_mcp_client.client.ClientSession")
|
||||
async def test_session_exit_cancellation_preserves_original_failure(self, session_class, original_error):
|
||||
self._make_session(session_class, AsyncMock(return_value=None))
|
||||
self._make_session(session_class, AsyncMock(return_value=_initialized()))
|
||||
cancelled: Final = asyncio.CancelledError("cancelled while closing session")
|
||||
session_class.return_value.__aexit__ = AsyncMock(side_effect=cancelled)
|
||||
original: Final = RuntimeError("operation failed")
|
||||
|
|
@ -656,7 +662,7 @@ class TestExecuteSessionOperationSurfacesTransportError:
|
|||
@pytest.mark.parametrize("phase", ("session", "transport"))
|
||||
@patch("litellm.experimental_mcp_client.client.ClientSession")
|
||||
async def test_cleanup_preserves_process_exit(self, session_class, phase, signal_type):
|
||||
self._make_session(session_class, AsyncMock(return_value=None))
|
||||
self._make_session(session_class, AsyncMock(return_value=_initialized()))
|
||||
signal: Final = signal_type("process stopping")
|
||||
if phase == "session":
|
||||
session_class.return_value.__aexit__ = AsyncMock(side_effect=signal)
|
||||
|
|
@ -670,7 +676,7 @@ class TestExecuteSessionOperationSurfacesTransportError:
|
|||
@pytest.mark.asyncio
|
||||
@patch("litellm.experimental_mcp_client.client.ClientSession")
|
||||
async def test_session_and_termination_share_one_cleanup_deadline(self, session_class):
|
||||
self._make_session(session_class, AsyncMock(return_value=None))
|
||||
self._make_session(session_class, AsyncMock(return_value=_initialized()))
|
||||
deleting: Final = asyncio.Event()
|
||||
|
||||
async def close_session(*args):
|
||||
|
|
@ -1883,16 +1889,17 @@ async def test_sse_read_failure_is_preserved() -> None:
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("protocol_version", ["auto", "2025-06-18"])
|
||||
@pytest.mark.parametrize("transport", [MCPTransport.sse, MCPTransport.stdio])
|
||||
@pytest.mark.parametrize("mode", ["ok", "closed", "silent"])
|
||||
async def test_transport_completion_and_normal_messages(transport: MCPTransport, mode: str) -> None:
|
||||
async def test_transport_completion_and_normal_messages(transport: MCPTransport, mode: str, protocol_version: str) -> None:
|
||||
from mcp import ClientSession
|
||||
from litellm.proxy._experimental.mcp_server.rest_endpoints import _connection_error_message
|
||||
|
||||
logging_callback: Final = AsyncMock()
|
||||
read_timeout: Final = 0.2 if mode == "silent" else 30
|
||||
client: Final = MCPClient(
|
||||
server_url="https://example.com/sse", transport_type=transport, timeout=read_timeout, logging_callback=logging_callback
|
||||
server_url="https://example.com/sse", transport_type=transport, timeout=read_timeout, logging_callback=logging_callback, protocol_version=protocol_version
|
||||
)
|
||||
|
||||
async def operation(session: ClientSession) -> CallToolResult:
|
||||
|
|
@ -2754,12 +2761,13 @@ async def test_http_close_cancellation_cannot_turn_into_success(original_error:
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("protocol_version", ("auto", "2025-06-18"))
|
||||
@pytest.mark.parametrize("cancel_mode", ("scope", "task", "wait_for", "read_timeout"))
|
||||
@pytest.mark.parametrize("concurrency", (1, 5))
|
||||
@pytest.mark.parametrize("termination", ("ok", "hang", "hang_body"))
|
||||
@pytest.mark.parametrize("raise_on_error", (False, True))
|
||||
async def test_cancellation_delivers_termination_over_tcp(
|
||||
cancel_mode: str, concurrency: int, termination: str, raise_on_error: bool
|
||||
cancel_mode: str, concurrency: int, termination: str, raise_on_error: bool, protocol_version: str
|
||||
) -> None:
|
||||
started: Final = asyncio.Event()
|
||||
scope_ready: Final[asyncio.Future[anyio.CancelScope]] = asyncio.get_running_loop().create_future()
|
||||
|
|
@ -2813,6 +2821,8 @@ async def test_cancellation_delivers_termination_over_tcp(
|
|||
await stop.wait()
|
||||
return
|
||||
if payload["method"] == "initialize":
|
||||
if cancel_mode != "read_timeout":
|
||||
await asyncio.sleep(0.75)
|
||||
response: Final = json.dumps(
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
|
|
@ -2839,7 +2849,7 @@ async def test_cancellation_delivers_termination_over_tcp(
|
|||
listener: Final = await asyncio.start_server(handle_connection, "127.0.0.1", 0)
|
||||
port: Final = listener.sockets[0].getsockname()[1]
|
||||
client: Final = MCPClient(
|
||||
server_url=f"http://127.0.0.1:{port}/mcp", timeout=2 if cancel_mode == "read_timeout" else 0.5 if termination != "ok" else 30
|
||||
server_url=f"http://127.0.0.1:{port}/mcp", protocol_version=protocol_version, timeout=2 if cancel_mode == "read_timeout" else 30
|
||||
)
|
||||
|
||||
async def calls():
|
||||
|
|
@ -2866,7 +2876,7 @@ async def test_cancellation_delivers_termination_over_tcp(
|
|||
|
||||
try:
|
||||
task: Final = asyncio.create_task(invoke())
|
||||
await asyncio.wait_for(started.wait(), 3)
|
||||
await asyncio.wait_for(started.wait(), 30)
|
||||
if cancel_mode == "scope":
|
||||
(await scope_ready).deadline = anyio.current_time() + 0.2
|
||||
if cancel_mode == "task":
|
||||
|
|
@ -2901,3 +2911,52 @@ async def test_cancellation_delivers_termination_over_tcp(
|
|||
closed: Final = await asyncio.wait_for(asyncio.gather(*connections, return_exceptions=True), 2)
|
||||
assert all(result is None or isinstance(result, asyncio.CancelledError) for result in closed), closed
|
||||
await asyncio.wait_for(listener.wait_closed(), 2)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("revision", ["2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25", "auto"])
|
||||
@pytest.mark.parametrize("accepted", [True, False])
|
||||
@pytest.mark.parametrize("callbacks", [False, True])
|
||||
async def test_configured_upstream_revision_is_offered_and_checked(revision, accepted, callbacks):
|
||||
from mcp.types import JSONRPCRequest
|
||||
from mcp_types.version import LATEST_HANDSHAKE_VERSION
|
||||
|
||||
offered = LATEST_HANDSHAKE_VERSION if revision == "auto" else revision
|
||||
|
||||
def respond(request):
|
||||
if request.method == "DELETE":
|
||||
return httpx2.Response(200)
|
||||
payload = _JSONRPC_MESSAGE_ADAPTER.validate_json(request.content)
|
||||
if not isinstance(payload, JSONRPCRequest):
|
||||
return httpx2.Response(202)
|
||||
if payload.method == "initialize":
|
||||
assert payload.params["protocolVersion"] == offered
|
||||
assert ("sampling" in payload.params["capabilities"]) == callbacks
|
||||
assert ("elicitation" in payload.params["capabilities"]) == callbacks
|
||||
return httpx2.Response(200, json={
|
||||
"jsonrpc": "2.0", "id": payload.id,
|
||||
"result": {"protocolVersion": offered if accepted else "unsupported",
|
||||
"capabilities": {"tools": {}}, "serverInfo": {"name": "upstream", "version": "1"}},
|
||||
})
|
||||
assert accepted, "No operation may execute after failed version negotiation"
|
||||
return httpx2.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": {"tools": [{"name": "echo", "inputSchema": {"type": "object"}}]}})
|
||||
|
||||
client = _MockTransportClient(
|
||||
respond, server_url="https://example.com/mcp", protocol_version=revision,
|
||||
sampling_callback=AsyncMock() if callbacks else None,
|
||||
elicitation_callback=AsyncMock() if callbacks else None,
|
||||
)
|
||||
if accepted:
|
||||
result = await client.list_tools(raise_on_error=True)
|
||||
assert [tool.name for tool in result] == ["echo"]
|
||||
else:
|
||||
with pytest.raises((MCPError, RuntimeError), match="protocol version"):
|
||||
await client.list_tools(raise_on_error=True)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("revision", ["2026-07-28", "unknown", "", None])
|
||||
def test_upstream_protocol_configuration_rejects_unavailable_modes(revision):
|
||||
from pydantic import ValidationError
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
MCPClient(protocol_version=revision)
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@ import litellm.experimental_mcp_client.client as mcp_client_module
|
|||
from litellm.experimental_mcp_client.client import MCPClient
|
||||
from litellm.types.mcp import MCPAuth, MCPTransport
|
||||
from mcp.types import CallToolResult as MCPCallToolResult
|
||||
from mcp.types import ListToolsResult, PaginatedRequestParams
|
||||
from mcp.types import Implementation, InitializeResult, ListToolsResult, PaginatedRequestParams, ServerCapabilities
|
||||
from mcp.types import Tool as MCPTool
|
||||
|
||||
|
||||
|
|
@ -128,6 +128,11 @@ class TestMCPClientUnitTests:
|
|||
mock_session_ctx = AsyncMock()
|
||||
mock_session_class.return_value = mock_session_ctx
|
||||
mock_session_instance = AsyncMock()
|
||||
mock_session_instance.initialize.return_value = InitializeResult(
|
||||
protocol_version="2025-11-25",
|
||||
capabilities=ServerCapabilities(),
|
||||
server_info=Implementation(name="test-peer", version="1"),
|
||||
)
|
||||
mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session_instance)
|
||||
|
||||
client = MCPClient(
|
||||
|
|
@ -163,6 +168,11 @@ class TestMCPClientUnitTests:
|
|||
mock_session_ctx = AsyncMock()
|
||||
mock_session_class.return_value = mock_session_ctx
|
||||
mock_session_instance = AsyncMock()
|
||||
mock_session_instance.initialize.return_value = InitializeResult(
|
||||
protocol_version="2025-11-25",
|
||||
capabilities=ServerCapabilities(),
|
||||
server_info=Implementation(name="test-peer", version="1"),
|
||||
)
|
||||
mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session_instance)
|
||||
|
||||
mock_tools = [
|
||||
|
|
@ -204,6 +214,11 @@ class TestMCPClientUnitTests:
|
|||
mock_session_ctx = AsyncMock()
|
||||
mock_session_class.return_value = mock_session_ctx
|
||||
mock_session_instance = AsyncMock()
|
||||
mock_session_instance.initialize.return_value = InitializeResult(
|
||||
protocol_version="2025-11-25",
|
||||
capabilities=ServerCapabilities(),
|
||||
server_info=Implementation(name="test-peer", version="1"),
|
||||
)
|
||||
mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session_instance)
|
||||
|
||||
first_page_tools = [
|
||||
|
|
@ -245,6 +260,11 @@ class TestMCPClientUnitTests:
|
|||
mock_session_ctx = AsyncMock()
|
||||
mock_session_class.return_value = mock_session_ctx
|
||||
mock_session_instance = AsyncMock()
|
||||
mock_session_instance.initialize.return_value = InitializeResult(
|
||||
protocol_version="2025-11-25",
|
||||
capabilities=ServerCapabilities(),
|
||||
server_info=Implementation(name="test-peer", version="1"),
|
||||
)
|
||||
mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session_instance)
|
||||
|
||||
mock_session_instance.list_tools.side_effect = [
|
||||
|
|
@ -277,6 +297,11 @@ class TestMCPClientUnitTests:
|
|||
mock_session_ctx = AsyncMock()
|
||||
mock_session_class.return_value = mock_session_ctx
|
||||
mock_session_instance = AsyncMock()
|
||||
mock_session_instance.initialize.return_value = InitializeResult(
|
||||
protocol_version="2025-11-25",
|
||||
capabilities=ServerCapabilities(),
|
||||
server_info=Implementation(name="test-peer", version="1"),
|
||||
)
|
||||
mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session_instance)
|
||||
|
||||
mock_result = MCPCallToolResult(content=[])
|
||||
|
|
@ -289,7 +314,7 @@ class TestMCPClientUnitTests:
|
|||
assert result == mock_result
|
||||
mock_session_instance.initialize.assert_called_once()
|
||||
mock_session_instance.call_tool.assert_called_once_with(
|
||||
name="test_tool", arguments={"arg1": "value1"}, progress_callback=ANY
|
||||
name="test_tool", arguments={"arg1": "value1"}, progress_callback=ANY, allow_input_required=False
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -489,7 +489,7 @@ async def test_sse_mcp_handler_mock():
|
|||
)
|
||||
|
||||
with (
|
||||
patch("litellm.proxy._experimental.mcp_server.server.server.run", run),
|
||||
patch("litellm.proxy._experimental.mcp_server.server.serve_loop", run),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED",
|
||||
True,
|
||||
|
|
@ -511,7 +511,7 @@ async def test_sse_mcp_handler_mock():
|
|||
# Call the handler
|
||||
await handle_sse_mcp(mock_scope, mock_receive, mock_send)
|
||||
|
||||
assert run.await_args.args[:2] == (read_stream, write_stream)
|
||||
assert run.await_args.args[1:3] == (read_stream, write_stream)
|
||||
assert mock_sse.connect_sse.call_args.args[0]["path"] == "/mcp/sse"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -38,6 +38,7 @@ def _run_shard(tmp_path: Path, test_path: str, workers: str) -> subprocess.Compl
|
|||
"PATH": f"{shim_dir}{os.pathsep}{os.environ['PATH']}",
|
||||
"GITHUB_OUTPUT": str(tmp_path / "github_output"),
|
||||
"TEST_PATH": test_path,
|
||||
"UNIT_FLAG": "",
|
||||
"WORKERS": workers,
|
||||
"UNIT_FLAG": "",
|
||||
},
|
||||
|
|
|
|||
5
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
5
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -28350,6 +28350,11 @@ export interface components {
|
|||
* @description Maximum retention period for spend logs (e.g., '7d' for 7 days). Logs older than this will be deleted.
|
||||
*/
|
||||
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.
|
||||
*/
|
||||
mcp_advertised_versions?: ("2024-11-05" | "2025-03-26" | "2025-06-18" | "2025-11-25")[] | 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