mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
refactor(mcp): extract upstream preparation and support modern clients (#44232)
* refactor(mcp): extract upstream preparation and support modern clients * fix(mcp): reject incompatible upstream transport before saving * fix(mcp): serialize protocol validation with server updates --------- Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
This commit is contained in:
parent
485ad76635
commit
04bc354525
22 changed files with 1115 additions and 381 deletions
|
|
@ -36,6 +36,7 @@ from mcp.types import (
|
|||
METHOD_NOT_FOUND,
|
||||
REQUEST_TIMEOUT,
|
||||
ClientCapabilities,
|
||||
DiscoverResult,
|
||||
ElicitationCapability,
|
||||
FormElicitationCapability,
|
||||
GetPromptRequestParams,
|
||||
|
|
@ -84,6 +85,7 @@ from litellm.types.mcp import (
|
|||
MCPUpstreamProtocol,
|
||||
credential_redirect_hook,
|
||||
has_header,
|
||||
validate_mcp_protocol_transport,
|
||||
without_header,
|
||||
)
|
||||
|
||||
|
|
@ -401,7 +403,10 @@ class MCPClient:
|
|||
logging_callback: Callable | None = None,
|
||||
protocol_version: MCPUpstreamProtocol = "auto",
|
||||
):
|
||||
self.protocol_version: MCPUpstreamProtocol = TypeAdapter(MCPUpstreamProtocol).validate_python(protocol_version)
|
||||
self.protocol_version: MCPUpstreamProtocol = TypeAdapter[MCPUpstreamProtocol](
|
||||
MCPUpstreamProtocol
|
||||
).validate_python(protocol_version)
|
||||
validate_mcp_protocol_transport(self.protocol_version, transport_type)
|
||||
self.server_url: str = server_url
|
||||
self.transport_type: MCPTransport = transport_type
|
||||
self.auth_type: MCPAuthType = auth_type
|
||||
|
|
@ -540,6 +545,17 @@ class MCPClient:
|
|||
|
||||
return safe_env
|
||||
|
||||
async def _prepare_session(self, session: ClientSession) -> InitializeResult | DiscoverResult:
|
||||
if self.protocol_version != "2026-07-28":
|
||||
return await self._initialize_session(session)
|
||||
discovery: Final = DiscoverResult.model_validate(await session.send_discover(self.protocol_version))
|
||||
if self.protocol_version not in discovery.supported_versions:
|
||||
raise MCPError(code=-32022, message="Upstream did not accept the configured MCP protocol version")
|
||||
session.adopt(discovery)
|
||||
if session.protocol_version != self.protocol_version:
|
||||
raise MCPError(code=-32022, message="Upstream selected an unsupported MCP protocol version")
|
||||
return discovery
|
||||
|
||||
async def _initialize_session(self, session: ClientSession) -> InitializeResult:
|
||||
if self.protocol_version == "auto":
|
||||
automatic: Final = await session.initialize()
|
||||
|
|
@ -623,7 +639,7 @@ class MCPClient:
|
|||
)
|
||||
session: Final = await session_ctx.__aenter__()
|
||||
try:
|
||||
init_result: Final = await self._initialize_session(session)
|
||||
init_result: Final = await self._prepare_session(session)
|
||||
instructions: Final = getattr(init_result, "instructions", None)
|
||||
self._last_initialize_instructions = (
|
||||
instructions.strip() or None if isinstance(instructions, str) else None
|
||||
|
|
|
|||
|
|
@ -3,12 +3,28 @@ from copy import deepcopy
|
|||
from dataclasses import dataclass, field
|
||||
from datetime import datetime
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Protocol
|
||||
from typing import TYPE_CHECKING, Final, Literal, Protocol
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.tool_outcome import WireCompat
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._experimental.mcp_server.server_resolution import ResolvedMCPServer
|
||||
|
||||
|
||||
class TargetCatalog(Protocol):
|
||||
async def resolve(
|
||||
self,
|
||||
server_id: str,
|
||||
caller: UserAPIKeyAuth,
|
||||
*,
|
||||
is_admin_view: bool,
|
||||
not_found_detail: Mapping[str, str],
|
||||
forbidden_detail: Mapping[str, str],
|
||||
non_admin_missing: Literal["not_found", "forbidden"],
|
||||
) -> "ResolvedMCPServer": ...
|
||||
|
||||
|
||||
def copy_caller(auth: UserAPIKeyAuth | None) -> UserAPIKeyAuth | None:
|
||||
if auth is None:
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ from datetime import datetime, timedelta, timezone
|
|||
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypedDict, cast
|
||||
|
||||
from fastapi import HTTPException
|
||||
from pydantic import TypeAdapter
|
||||
from typing_extensions import ReadOnly
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -52,8 +53,8 @@ from litellm.repositories.verification_token_repository import (
|
|||
VerificationTokenRepository,
|
||||
)
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
from litellm.types.mcp import MCPCredentials
|
||||
from litellm.types.mcp_server.mcp_server_manager import PinnedMCPTool
|
||||
from litellm.types.mcp import MCPCredentials, MCPTransportType, MCPUpstreamProtocol, validate_mcp_protocol_transport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPInfo, PinnedMCPTool
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma import models as prisma_db_models
|
||||
|
|
@ -600,6 +601,7 @@ async def _mcp_server_write_if_identifier_free(
|
|||
alias: str | None,
|
||||
exclude_server_id: str | None,
|
||||
write: "Callable[[TableActions[prisma_db_models.LiteLLM_MCPServerTable]], Awaitable[prisma_db_models.LiteLLM_MCPServerTable | None]]",
|
||||
lock_server_id: str | None = None,
|
||||
) -> "prisma_db_models.LiteLLM_MCPServerTable | McpIdentifierConflict | None":
|
||||
"""Run ``write`` only when no other live row owns ``server_name``/``alias``.
|
||||
|
||||
|
|
@ -619,6 +621,10 @@ async def _mcp_server_write_if_identifier_free(
|
|||
)
|
||||
if conflict is not None:
|
||||
return conflict
|
||||
if lock_server_id is not None:
|
||||
await tx.execute_raw(
|
||||
'SELECT server_id FROM "LiteLLM_MCPServerTable" WHERE server_id=$1 FOR UPDATE', lock_server_id
|
||||
)
|
||||
return await write(tx.litellm_mcpservertable)
|
||||
|
||||
|
||||
|
|
@ -1100,6 +1106,27 @@ async def get_draft_mcp_server(
|
|||
return table
|
||||
|
||||
|
||||
def _validate_mcp_protocol_write(
|
||||
stored: "prisma_db_models.LiteLLM_MCPServerTable", data_dict: Mapping[str, object]
|
||||
) -> None:
|
||||
raw_info: Final = data_dict.get("mcp_info", stored.mcp_info)
|
||||
adapter: Final = TypeAdapter[MCPInfo | None](MCPInfo | None)
|
||||
try:
|
||||
info: Final = (
|
||||
adapter.validate_json(raw_info) if isinstance(raw_info, str) else adapter.validate_python(raw_info)
|
||||
)
|
||||
validate_mcp_protocol_transport(
|
||||
TypeAdapter[MCPUpstreamProtocol](MCPUpstreamProtocol).validate_python(
|
||||
(info or {}).get("protocol_version", "auto")
|
||||
),
|
||||
TypeAdapter[MCPTransportType](MCPTransportType).validate_python(
|
||||
data_dict.get("transport", stored.transport)
|
||||
),
|
||||
)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
|
||||
|
||||
async def _update_mcp_server_row(
|
||||
prisma_client: PrismaClient,
|
||||
*,
|
||||
|
|
@ -1107,29 +1134,36 @@ async def _update_mcp_server_row(
|
|||
data_dict: Mapping[str, object],
|
||||
) -> "prisma_db_models.LiteLLM_MCPServerTable | McpIdentifierConflict | None":
|
||||
identifier_write: Final = any(field in data_dict for field in ("server_name", "alias"))
|
||||
protocol_write: Final = bool({"transport", "mcp_info"}.intersection(data_dict))
|
||||
|
||||
async def _update(
|
||||
table: "TableActions[prisma_db_models.LiteLLM_MCPServerTable]",
|
||||
) -> "prisma_db_models.LiteLLM_MCPServerTable | None":
|
||||
if protocol_write:
|
||||
stored: Final = await table.find_unique(where={"server_id": server_id})
|
||||
if stored is None:
|
||||
return None
|
||||
_validate_mcp_protocol_write(stored, data_dict)
|
||||
return await table.update(
|
||||
where={"server_id": server_id},
|
||||
data=data_dict,
|
||||
)
|
||||
|
||||
if not identifier_write:
|
||||
if not identifier_write and not protocol_write:
|
||||
return await _update(_mcp_server_table_actions(prisma_client))
|
||||
if "alias" in data_dict and not data_dict["alias"] and "server_name" not in data_dict:
|
||||
# Clearing the alias drops the prefix to the stored server_name, which
|
||||
# may already belong to another row, so that name needs the check too.
|
||||
existing: Final = await _db_find_mcp_server_row(prisma_client, server_id)
|
||||
if existing is None:
|
||||
return await _update(_mcp_server_table_actions(prisma_client))
|
||||
return None
|
||||
return await _mcp_server_write_if_identifier_free(
|
||||
prisma_client,
|
||||
server_name=existing.server_name,
|
||||
alias=None,
|
||||
exclude_server_id=server_id,
|
||||
write=_update,
|
||||
lock_server_id=server_id if protocol_write else None,
|
||||
)
|
||||
return await _mcp_server_write_if_identifier_free(
|
||||
prisma_client,
|
||||
|
|
@ -1137,6 +1171,7 @@ async def _update_mcp_server_row(
|
|||
alias=_identifier_field(data_dict, "alias"),
|
||||
exclude_server_id=server_id,
|
||||
write=_update,
|
||||
lock_server_id=server_id if protocol_write else None,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -58,12 +58,10 @@ from litellm.constants import (
|
|||
MCP_CLIENT_TIMEOUT,
|
||||
MCP_HEALTH_CHECK_TIMEOUT,
|
||||
MCP_METADATA_TIMEOUT,
|
||||
MCP_NPM_CACHE_DIR,
|
||||
MCP_STDIO_ALLOWED_COMMANDS,
|
||||
MCP_TOOL_LISTING_TIMEOUT,
|
||||
)
|
||||
from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException
|
||||
from litellm.experimental_mcp_client.client import MCPClient, MCPSigV4Auth, strip_auth_scheme, to_basic_credentials
|
||||
from litellm.experimental_mcp_client.client import MCPClient, strip_auth_scheme, to_basic_credentials
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
_sync_guardrail_info_to_logging_obj, # pyright: ignore[reportPrivateUsage] - the same bridge @log_guardrail_information uses; reimplementing it here would fork the metadata-key logic
|
||||
)
|
||||
|
|
@ -91,8 +89,6 @@ from litellm.proxy._experimental.mcp_server.mcp_debug import describe_upstream_h
|
|||
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import (
|
||||
MCPPerUserTokenCache,
|
||||
mcp_per_user_token_cache,
|
||||
resolve_mcp_auth,
|
||||
resolved_token_header,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import (
|
||||
_redact_mcp_resource_url,
|
||||
|
|
@ -105,10 +101,8 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials import (
|
|||
UpstreamCredentialProvider,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
|
||||
prepare_mcp_client,
|
||||
raise_public,
|
||||
raise_token_exchange_challenge,
|
||||
raise_user_oauth_challenge,
|
||||
to_server_spec,
|
||||
to_subject,
|
||||
)
|
||||
|
|
@ -118,19 +112,14 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_sto
|
|||
from litellm.proxy._experimental.mcp_server.outbound_credentials.per_user_oauth_store import (
|
||||
LazyPerUserOAuthTokenStore,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.resolver import resolve_credentials_with_source
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_provider import (
|
||||
build_token_exchanger,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
||||
DEFAULT_CREDENTIAL_HEADER,
|
||||
AuthorizationCodeConfig,
|
||||
AuthResolution,
|
||||
ClientCredentialsConfig,
|
||||
CredError,
|
||||
IdJagConfig,
|
||||
PassthroughConfig,
|
||||
ServerSpec,
|
||||
TokenExchangeConfig,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.result_conversion import (
|
||||
|
|
@ -145,7 +134,6 @@ from litellm.proxy._experimental.mcp_server.sampling_handler import (
|
|||
from litellm.proxy._experimental.mcp_server.stdio_gate import (
|
||||
MCP_STDIO_DISABLED_MESSAGE,
|
||||
is_mcp_stdio_blocked,
|
||||
is_mcp_stdio_enabled,
|
||||
warn_if_mcp_stdio_blocked,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.tool_catalog_guard import (
|
||||
|
|
@ -154,6 +142,13 @@ from litellm.proxy._experimental.mcp_server.tool_catalog_guard import (
|
|||
pin_tool_catalog,
|
||||
scan_tool_descriptions,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.upstream import (
|
||||
passthrough_token_from_mcp_auth_header,
|
||||
prepare_upstream_client,
|
||||
resolve_upstream_auth,
|
||||
take_forwarded_authorization,
|
||||
to_server_spec_fail_closed,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.utils import (
|
||||
MCP_TOOL_PREFIX_SEPARATOR,
|
||||
MCPMissingUserEnvVarsError,
|
||||
|
|
@ -204,10 +199,8 @@ from litellm.types.llms.custom_http import httpxSpecialProvider
|
|||
from litellm.types.mcp import (
|
||||
DEFAULT_SUBJECT_TOKEN_TYPE,
|
||||
MCPAuth,
|
||||
MCPStdioConfig,
|
||||
MCPTokenEndpointAuthMethod,
|
||||
MCPUpstreamProtocol,
|
||||
has_header,
|
||||
without_header,
|
||||
)
|
||||
from litellm.types.mcp_server.mcp_server_manager import (
|
||||
|
|
@ -1234,7 +1227,7 @@ def _resolve_openapi_tool_auth(
|
|||
|
||||
Returns the ``Authorization`` value to inject, the extra headers to forward, and the credential to
|
||||
hand ``resolve_openapi_upstream_auth``, whose passthrough arm reads it via
|
||||
``_passthrough_token_from_mcp_auth_header``. The per-server Authorization travels only in the
|
||||
``passthrough_token_from_mcp_auth_header``. The per-server Authorization travels only in the
|
||||
credential, never also in the forwarded headers, because the resolver pops Authorization out of
|
||||
those and would otherwise have two sources to reconcile.
|
||||
"""
|
||||
|
|
@ -1325,35 +1318,6 @@ def _client_forwarded_authorization_headers(
|
|||
return extra_headers
|
||||
|
||||
|
||||
def _take_forwarded_authorization(
|
||||
headers: dict[str, str] | None,
|
||||
) -> tuple[str | None, dict[str, str] | None]:
|
||||
"""Pop the ``Authorization`` value out of ``headers`` (case-insensitive), returning it with the
|
||||
remaining headers, so the passthrough resolver arm is the single Authorization source rather than
|
||||
the header also riding in ``extra_headers`` (which the resolved auth would then defer to)."""
|
||||
if not headers:
|
||||
return None, headers
|
||||
value: Final = next((v for k, v in headers.items() if k.lower() == "authorization"), None)
|
||||
return value, without_header(headers, DEFAULT_CREDENTIAL_HEADER)
|
||||
|
||||
|
||||
def _passthrough_token_from_mcp_auth_header(
|
||||
mcp_auth_header: str | dict[str, str] | None,
|
||||
) -> str | None:
|
||||
"""The caller's per-server upstream credential for a passthrough-mode server, or None.
|
||||
|
||||
Sourced from ``x-mcp-{alias}-authorization`` (string or per-header dict form) or the deprecated
|
||||
global ``x-mcp-auth`` fallback. Per-server headers are the multi-server shape: they bind one
|
||||
token to one server, so an aggregate scope with several passthrough-mode servers never replays
|
||||
a single credential across upstreams. The value is forwarded verbatim, so it must be the full
|
||||
header value (e.g. ``Bearer <upstream-token>``)."""
|
||||
if isinstance(mcp_auth_header, str):
|
||||
return mcp_auth_header or None
|
||||
if isinstance(mcp_auth_header, dict):
|
||||
return next((v for k, v in mcp_auth_header.items() if k.lower() == "authorization"), None)
|
||||
return None
|
||||
|
||||
|
||||
async def _materialize_auth_headers(auth: httpx2.Auth | None) -> dict[str, str] | None:
|
||||
"""Extract the header a resolved ``httpx2.Auth`` would set, as a plain dict, or None.
|
||||
|
||||
|
|
@ -1418,25 +1382,6 @@ def _redacted_registry_dump(servers: dict[str, MCPServer]) -> dict[str, dict[str
|
|||
}
|
||||
|
||||
|
||||
def _to_server_spec_fail_closed(server: MCPServer) -> ServerSpec | None:
|
||||
"""`to_server_spec`, except a half-configured `oauth2_id_jag` server refuses instead of deferring.
|
||||
|
||||
ID-JAG has no v1 arm, so deferring to v1 would let `resolve_mcp_auth` honor a caller x-mcp-*
|
||||
override or fall through to the static `authentication_token`, both of which bypass the per-user
|
||||
identity assertion the mode promises. That is an operator misconfiguration, not a fallback.
|
||||
"""
|
||||
spec: Final = to_server_spec(server)
|
||||
if spec is None and server.auth_type == MCPAuth.oauth2_id_jag:
|
||||
raise_public(
|
||||
CredError.of_misconfigured(
|
||||
"oauth2_id_jag requires token_exchange_endpoint, id_jag_resource_token_endpoint, "
|
||||
"client_id, and a client_secret or client_private_key; refusing to fall back to "
|
||||
"a static credential."
|
||||
)
|
||||
)
|
||||
return spec
|
||||
|
||||
|
||||
def _caller_authorization_fans_out(
|
||||
server: MCPServer,
|
||||
scope_servers: list[MCPServer] | None,
|
||||
|
|
@ -4088,75 +4033,6 @@ class MCPServerManager:
|
|||
_write_user_env_vars_cache(user_id, server.server_id, values)
|
||||
return values
|
||||
|
||||
async def _resolve_v2_auth(
|
||||
self,
|
||||
*,
|
||||
server: MCPServer,
|
||||
spec: ServerSpec,
|
||||
provider: UpstreamCredentialProvider,
|
||||
subject_token: str | None,
|
||||
user_api_key_auth: UserAPIKeyAuth | None,
|
||||
extra_headers: dict[str, str] | None,
|
||||
) -> tuple[httpx2.Auth | None, dict[str, str] | None]:
|
||||
"""Resolve a v2-owned server's upstream credential into ``(resolved_auth, extra_headers)``.
|
||||
|
||||
On a missing/rejected per-user credential this raises the mode's discovery challenge
|
||||
(authorization_code's browser-OAuth 401, token_exchange's RFC 9728 challenge) or maps any
|
||||
other ``CredError`` onto its public HTTP status; it never returns an error as a value.
|
||||
"""
|
||||
match await resolve_credentials_with_source(provider, to_subject(user_api_key_auth, subject_token), spec):
|
||||
case Ok(credential):
|
||||
auth: Final = credential.auth
|
||||
# NoOpAuth has no header_name and so never conflicts.
|
||||
header_name: Final[str | None] = getattr(auth, "header_name", None)
|
||||
if header_name is None or not extra_headers:
|
||||
source: Final = (
|
||||
AuthResolution.extra_headers
|
||||
if credential.source == AuthResolution.no_auth and extra_headers
|
||||
else credential.source
|
||||
)
|
||||
record_auth_resolution(server.server_id, source)
|
||||
return auth, extra_headers
|
||||
if not has_header(extra_headers, header_name):
|
||||
record_auth_resolution(server.server_id, credential.source)
|
||||
return auth, extra_headers
|
||||
if isinstance(
|
||||
spec.config,
|
||||
(TokenExchangeConfig, AuthorizationCodeConfig, IdJagConfig, ClientCredentialsConfig),
|
||||
):
|
||||
# The resolver owns the credential here (token_exchange's exchanged token,
|
||||
# authorization_code's stored token, id_jag's minted assertion,
|
||||
# client_credentials' gateway-minted M2M token). It is authoritative: a
|
||||
# guardrail such as MCPJWTSigner, static_headers, or any other injected
|
||||
# Authorization must NOT shadow it (otherwise the upstream gets e.g. the
|
||||
# signer's JWT instead of the minted token and rejects it, and for M2M the
|
||||
# one-shot 401 refetch is lost with it). Drop only the header the resolved
|
||||
# credential is about to occupy, so a static credential the operator aimed at a
|
||||
# DIFFERENT header still reaches upstream.
|
||||
record_auth_resolution(server.server_id, credential.source)
|
||||
return auth, without_header(extra_headers, header_name)
|
||||
# Other modes: an Authorization already supplied via extra_headers (a forwarded caller
|
||||
# header or static_headers) is intentional and wins; v1 applies those last.
|
||||
record_auth_resolution(server.server_id, AuthResolution.extra_headers)
|
||||
return None, extra_headers
|
||||
case Error(err):
|
||||
record_auth_resolution(server.server_id, AuthResolution.failed)
|
||||
if err.tag == "unauthorized" and isinstance(spec.config, AuthorizationCodeConfig):
|
||||
# authorization_code's missing per-user token -> the per-server browser-OAuth
|
||||
# challenge, built here where the full MCPServer is in hand.
|
||||
raise_user_oauth_challenge(server, root_path=get_request_root_path())
|
||||
if err.tag == "unauthorized" and isinstance(spec.config, TokenExchangeConfig):
|
||||
# token_exchange (OBO): a missing/rejected subject token -> the RFC 9728 challenge
|
||||
# pointing at the IdP the client must SSO with to obtain one, rather than an opaque
|
||||
# 401. No gateway-side browser flow. An IdP step-up rejection (Entra Conditional
|
||||
# Access) threads its claims blob into the challenge for the client to satisfy.
|
||||
raise_token_exchange_challenge(
|
||||
server,
|
||||
root_path=get_request_root_path(),
|
||||
claims=err.unauthorized.claims,
|
||||
)
|
||||
raise_public(err)
|
||||
|
||||
async def preflight_token_exchange(
|
||||
self,
|
||||
server: MCPServer,
|
||||
|
|
@ -4196,7 +4072,7 @@ class MCPServerManager:
|
|||
case _:
|
||||
return
|
||||
resolved_server: Final = await self.ensure_oauth_metadata_discovered(server)
|
||||
spec: Final = _to_server_spec_fail_closed(resolved_server)
|
||||
spec: Final = to_server_spec_fail_closed(resolved_server)
|
||||
if spec is None or not isinstance(spec.config, (TokenExchangeConfig, IdJagConfig)):
|
||||
return
|
||||
if subject_token is None and isinstance(spec.config, TokenExchangeConfig):
|
||||
|
|
@ -4226,202 +4102,33 @@ class MCPServerManager:
|
|||
client_ip: str | None = None,
|
||||
protocol_version_override: MCPUpstreamProtocol | None = None,
|
||||
) -> MCPClient:
|
||||
"""
|
||||
Create an MCPClient instance for the given server.
|
||||
|
||||
Auth resolution (single place for all auth logic):
|
||||
1. ``mcp_auth_header`` — per-request/per-user override
|
||||
2. OAuth2 Token Exchange (OBO) — exchange user token for scoped token
|
||||
3. OAuth2 client_credentials token — auto-fetched and cached
|
||||
4. ``server.authentication_token`` — static token from config/DB
|
||||
|
||||
Args:
|
||||
server: The server configuration.
|
||||
mcp_auth_header: Optional per-request auth override.
|
||||
extra_headers: Additional headers to forward.
|
||||
stdio_env: Environment variables for stdio transport.
|
||||
subject_token: Optional user JWT for token exchange (OBO) flow.
|
||||
user_api_key_auth: Optional auth context for sampling callbacks.
|
||||
|
||||
Returns:
|
||||
Configured MCP client instance.
|
||||
"""
|
||||
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
|
||||
# A caller-supplied per-request override (mcp_auth_header / x-mcp-*) defers to the v1 path
|
||||
# so it wins - except for the modes the v2 resolver owns per-caller (authorization_code's
|
||||
# stored token, token_exchange's RFC 8693 minted token, id_jag's minted assertion, and the
|
||||
# passthrough modes' forwarded caller token). A caller must not be able to substitute another
|
||||
# user's stored credential, nor silently disable the OBO / ID-JAG exchange and forward an
|
||||
# arbitrary bearer upstream, so we keep the v2 spec and ignore the override for these; the
|
||||
# REST tools preview supplies its not-yet-persisted token through the resolver
|
||||
# (cred_provider), never this path.
|
||||
if (
|
||||
spec is not None
|
||||
and mcp_auth_header
|
||||
and not isinstance(
|
||||
spec.config,
|
||||
(AuthorizationCodeConfig, IdJagConfig, PassthroughConfig, TokenExchangeConfig),
|
||||
)
|
||||
):
|
||||
spec = None
|
||||
auth_value: Final = await resolve_mcp_auth(resolved_server, mcp_auth_header) if spec is None else None
|
||||
auth_header_name: Final = resolved_token_header(resolved_server, mcp_auth_header) if spec is None else None
|
||||
|
||||
# Create sampling and elicitation callbacks for this client
|
||||
sampling_cb = (
|
||||
_create_sampling_callback(
|
||||
operation_context=OperationContext(
|
||||
_caller=user_api_key_auth, raw_headers=raw_headers, client_ip=client_ip
|
||||
)
|
||||
)
|
||||
if resolved_server.allow_sampling
|
||||
else None
|
||||
)
|
||||
elicitation_cb: Final = _create_elicitation_callback() if resolved_server.allow_elicitation else None
|
||||
|
||||
# Handle stdio transport
|
||||
if transport == MCPTransport.stdio:
|
||||
if not is_mcp_stdio_enabled():
|
||||
raise HTTPException(status_code=403, detail=MCP_STDIO_DISABLED_MESSAGE)
|
||||
resolved_env: Final = (
|
||||
stdio_env
|
||||
if stdio_env is not None
|
||||
else (dict(resolved_server.env) if resolved_server.env is not None else None)
|
||||
)
|
||||
|
||||
# Ensure npm-based STDIO MCP servers have a writable cache dir.
|
||||
# In containers the default (~/.npm or /app/.npm) may not exist
|
||||
# or be read-only, causing npx to fail with ENOENT.
|
||||
if resolved_env is not None and "NPM_CONFIG_CACHE" not in resolved_env:
|
||||
resolved_env["NPM_CONFIG_CACHE"] = MCP_NPM_CACHE_DIR
|
||||
# Defense-in-depth: block commands not in the allowlist.
|
||||
# The Pydantic validator blocks new servers; this catches legacy
|
||||
# config/DB records predating the allowlist.
|
||||
if resolved_server.command:
|
||||
base_command: Final = os.path.basename(resolved_server.command)
|
||||
# Strip .exe/.cmd/.bat/.com suffix for Windows compatibility
|
||||
base_command_no_ext = base_command.lower()
|
||||
for ext in [".exe", ".cmd", ".bat", ".com"]:
|
||||
if base_command.lower().endswith(ext):
|
||||
base_command_no_ext = base_command[: -len(ext)].lower()
|
||||
break
|
||||
if (
|
||||
base_command.lower() not in MCP_STDIO_ALLOWED_COMMANDS
|
||||
and base_command_no_ext not in MCP_STDIO_ALLOWED_COMMANDS
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=f"MCP stdio command '{resolved_server.command}' is not in the allowlist ({sorted(MCP_STDIO_ALLOWED_COMMANDS)}). "
|
||||
f"Add it to LITELLM_MCP_STDIO_EXTRA_COMMANDS to allow this command.",
|
||||
return await prepare_upstream_client(
|
||||
resolved_server,
|
||||
provider=cred_provider or self._cred_provider,
|
||||
root_path=get_request_root_path(),
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
stdio_env=stdio_env,
|
||||
subject_token=subject_token,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
protocol_version=(
|
||||
protocol_version_override if protocol_version_override is not None else resolved_server.protocol_version
|
||||
),
|
||||
sampling_callback=(
|
||||
_create_sampling_callback(
|
||||
operation_context=OperationContext(
|
||||
_caller=user_api_key_auth,
|
||||
raw_headers=raw_headers,
|
||||
client_ip=client_ip,
|
||||
)
|
||||
|
||||
stdio_config: MCPStdioConfig | None = None
|
||||
if resolved_server.command and resolved_server.args is not None:
|
||||
stdio_config = MCPStdioConfig(
|
||||
command=resolved_server.command,
|
||||
args=resolved_server.args,
|
||||
env=resolved_env,
|
||||
)
|
||||
|
||||
record_auth_resolution(server.server_id, AuthResolution.not_applicable)
|
||||
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),
|
||||
stdio_config=stdio_config,
|
||||
extra_headers=extra_headers,
|
||||
sampling_callback=sampling_cb,
|
||||
elicitation_callback=elicitation_cb,
|
||||
)
|
||||
else:
|
||||
# For HTTP/SSE transports
|
||||
server_url: Final = resolved_server.url or ""
|
||||
|
||||
if spec is not None:
|
||||
inbound_token = subject_token
|
||||
if isinstance(spec.config, PassthroughConfig):
|
||||
inbound_token, extra_headers = _take_forwarded_authorization(extra_headers)
|
||||
per_server_token: Final = _passthrough_token_from_mcp_auth_header(mcp_auth_header)
|
||||
if per_server_token is not None:
|
||||
inbound_token = per_server_token
|
||||
resolved_auth, extra_headers = await self._resolve_v2_auth(
|
||||
server=resolved_server,
|
||||
spec=spec,
|
||||
provider=provider,
|
||||
subject_token=inbound_token,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
extra_headers=extra_headers,
|
||||
)
|
||||
return await prepare_mcp_client(
|
||||
resolved_server,
|
||||
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
|
||||
),
|
||||
extra_headers=extra_headers,
|
||||
resolved_auth=resolved_auth,
|
||||
sampling_callback=sampling_cb,
|
||||
elicitation_callback=elicitation_cb,
|
||||
),
|
||||
)
|
||||
|
||||
# Create SigV4 auth if configured
|
||||
aws_auth = None
|
||||
if resolved_server.auth_type == MCPAuth.aws_sigv4:
|
||||
aws_auth = MCPSigV4Auth(
|
||||
aws_access_key_id=resolved_server.aws_access_key_id,
|
||||
aws_secret_access_key=resolved_server.aws_secret_access_key,
|
||||
aws_session_token=resolved_server.aws_session_token,
|
||||
aws_region_name=resolved_server.aws_region_name,
|
||||
aws_service_name=resolved_server.aws_service_name,
|
||||
aws_role_name=resolved_server.aws_role_name,
|
||||
aws_session_name=resolved_server.aws_session_name,
|
||||
)
|
||||
|
||||
legacy_source: Final = (
|
||||
AuthResolution.aws_sigv4
|
||||
if aws_auth is not None
|
||||
else AuthResolution.extra_headers
|
||||
if extra_headers and has_header(extra_headers, auth_header_name or "Authorization")
|
||||
else AuthResolution.per_request_header
|
||||
if mcp_auth_header
|
||||
else AuthResolution.static_token
|
||||
if auth_value
|
||||
else AuthResolution.extra_headers
|
||||
if extra_headers
|
||||
else AuthResolution.no_auth
|
||||
)
|
||||
record_auth_resolution(server.server_id, legacy_source)
|
||||
return await prepare_mcp_client(
|
||||
resolved_server,
|
||||
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,
|
||||
timeout=(resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT),
|
||||
extra_headers=extra_headers,
|
||||
aws_auth=aws_auth,
|
||||
sampling_callback=sampling_cb,
|
||||
elicitation_callback=elicitation_cb,
|
||||
),
|
||||
)
|
||||
if resolved_server.allow_sampling
|
||||
else None
|
||||
),
|
||||
elicitation_callback=(_create_elicitation_callback() if resolved_server.allow_elicitation else None),
|
||||
)
|
||||
|
||||
async def _get_tools_from_server(
|
||||
self,
|
||||
|
|
@ -6457,7 +6164,7 @@ class MCPServerManager:
|
|||
token, token_exchange's exchanged token, passthrough's forwarded caller token) must be
|
||||
materialized into headers here. Returns ``(resolved_auth_headers, forwarded_headers)``:
|
||||
the resolved headers are authoritative over every other Authorization source (the same
|
||||
rule ``_resolve_v2_auth`` applies on the MCPClient path) and ``forwarded_headers`` comes
|
||||
rule ``resolve_upstream_auth`` applies on the MCPClient path) and ``forwarded_headers`` comes
|
||||
back with any header the resolver claimed already dropped. Unmigrated (v1) servers resolve
|
||||
through the stored-token lookup instead, and a missing per-user credential raises the same
|
||||
discovery challenge the MCPClient path serves, rather than egressing unauthenticated.
|
||||
|
|
@ -6480,11 +6187,12 @@ class MCPServerManager:
|
|||
if isinstance(spec.config, (TokenExchangeConfig, IdJagConfig)):
|
||||
subject_token = self._extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth)
|
||||
elif isinstance(spec.config, PassthroughConfig):
|
||||
inbound_token, forwarded_headers = _take_forwarded_authorization(forwarded_headers)
|
||||
per_server_token: Final = _passthrough_token_from_mcp_auth_header(mcp_auth_header)
|
||||
inbound_token, forwarded_headers = take_forwarded_authorization(forwarded_headers)
|
||||
per_server_token: Final = passthrough_token_from_mcp_auth_header(mcp_auth_header)
|
||||
subject_token = per_server_token if per_server_token is not None else inbound_token
|
||||
|
||||
resolved_auth, forwarded_headers = await self._resolve_v2_auth(
|
||||
resolved_auth, forwarded_headers = await resolve_upstream_auth(
|
||||
root_path=get_request_root_path(),
|
||||
server=mcp_server,
|
||||
spec=spec,
|
||||
provider=self._cred_provider,
|
||||
|
|
|
|||
|
|
@ -120,3 +120,42 @@ async def authorize_mcp_server(
|
|||
)
|
||||
|
||||
return resolved
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class MCPServerTargetCatalog:
|
||||
manager: MCPServerRegistry
|
||||
db_lookup: Callable[[str], Awaitable[LiteLLM_MCPServerTable | None]] | None = None
|
||||
temp_lookup: Callable[[str], Awaitable[MCPServer | None]] | None = None
|
||||
id_client_ip: str | None = None
|
||||
name_client_ip: str | None = None
|
||||
match_name: bool = False
|
||||
|
||||
async def resolve(
|
||||
self,
|
||||
server_id: str,
|
||||
caller: UserAPIKeyAuth,
|
||||
*,
|
||||
is_admin_view: bool,
|
||||
not_found_detail: Mapping[str, str],
|
||||
forbidden_detail: Mapping[str, str],
|
||||
non_admin_missing: Literal["not_found", "forbidden"],
|
||||
) -> ResolvedMCPServer:
|
||||
resolved: Final = await resolve_mcp_server(
|
||||
server_id,
|
||||
manager=self.manager,
|
||||
db_lookup=self.db_lookup,
|
||||
temp_lookup=self.temp_lookup,
|
||||
id_client_ip=self.id_client_ip,
|
||||
name_client_ip=self.name_client_ip,
|
||||
match_name=self.match_name,
|
||||
)
|
||||
return await authorize_mcp_server(
|
||||
resolved,
|
||||
caller,
|
||||
manager=self.manager,
|
||||
is_admin_view=is_admin_view,
|
||||
not_found_detail=not_found_detail,
|
||||
forbidden_detail=forbidden_detail,
|
||||
non_admin_missing=non_admin_missing,
|
||||
)
|
||||
|
|
|
|||
315
litellm/proxy/_experimental/mcp_server/upstream.py
Normal file
315
litellm/proxy/_experimental/mcp_server/upstream.py
Normal file
|
|
@ -0,0 +1,315 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from typing import Final
|
||||
|
||||
import httpx2
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.constants import MCP_CLIENT_TIMEOUT, MCP_NPM_CACHE_DIR, MCP_STDIO_ALLOWED_COMMANDS
|
||||
from litellm.experimental_mcp_client.client import MCPClient, MCPSigV4Auth
|
||||
from litellm.proxy._experimental.mcp_server.legacy_callbacks import ElicitationCallback, SamplingCallback
|
||||
from litellm.proxy._experimental.mcp_server.mcp_debug import record_auth_resolution
|
||||
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import resolve_mcp_auth, resolved_token_header
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials import Error, Ok, UpstreamCredentialProvider
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
|
||||
prepare_mcp_client,
|
||||
raise_public,
|
||||
raise_token_exchange_challenge,
|
||||
raise_user_oauth_challenge,
|
||||
to_server_spec,
|
||||
to_subject,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.resolver import resolve_credentials_with_source
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
||||
DEFAULT_CREDENTIAL_HEADER,
|
||||
AuthorizationCodeConfig,
|
||||
AuthResolution,
|
||||
ClientCredentialsConfig,
|
||||
CredError,
|
||||
IdJagConfig,
|
||||
PassthroughConfig,
|
||||
ServerSpec,
|
||||
TokenExchangeConfig,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.stdio_gate import MCP_STDIO_DISABLED_MESSAGE, is_mcp_stdio_enabled
|
||||
from litellm.proxy._types import MCPTransport, UserAPIKeyAuth
|
||||
from litellm.types.mcp import (
|
||||
MCPAuth,
|
||||
MCPStdioConfig,
|
||||
MCPTransportType,
|
||||
MCPUpstreamProtocol,
|
||||
has_header,
|
||||
without_header,
|
||||
)
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
|
||||
def take_forwarded_authorization(
|
||||
headers: dict[str, str] | None,
|
||||
) -> tuple[str | None, dict[str, str] | None]:
|
||||
"""Pop the ``Authorization`` value out of ``headers`` (case-insensitive), returning it with the
|
||||
remaining headers, so the passthrough resolver arm is the single Authorization source rather than
|
||||
the header also riding in ``extra_headers`` (which the resolved auth would then defer to)."""
|
||||
if not headers:
|
||||
return None, headers
|
||||
value: Final = next((v for k, v in headers.items() if k.lower() == "authorization"), None)
|
||||
return value, without_header(headers, DEFAULT_CREDENTIAL_HEADER)
|
||||
|
||||
|
||||
def passthrough_token_from_mcp_auth_header(
|
||||
mcp_auth_header: str | dict[str, str] | None,
|
||||
) -> str | None:
|
||||
"""The caller's per-server upstream credential for a passthrough-mode server, or None.
|
||||
|
||||
Sourced from ``x-mcp-{alias}-authorization`` (string or per-header dict form) or the deprecated
|
||||
global ``x-mcp-auth`` fallback. Per-server headers are the multi-server shape: they bind one
|
||||
token to one server, so an aggregate scope with several passthrough-mode servers never replays
|
||||
a single credential across upstreams. The value is forwarded verbatim, so it must be the full
|
||||
header value (e.g. ``Bearer <upstream-token>``)."""
|
||||
if isinstance(mcp_auth_header, str):
|
||||
return mcp_auth_header or None
|
||||
if isinstance(mcp_auth_header, dict):
|
||||
return next((v for k, v in mcp_auth_header.items() if k.lower() == "authorization"), None)
|
||||
return None
|
||||
|
||||
|
||||
def to_server_spec_fail_closed(server: MCPServer) -> ServerSpec | None:
|
||||
"""`to_server_spec`, except a half-configured `oauth2_id_jag` server refuses instead of deferring.
|
||||
|
||||
ID-JAG has no v1 arm, so deferring to v1 would let `resolve_mcp_auth` honor a caller x-mcp-*
|
||||
override or fall through to the static `authentication_token`, both of which bypass the per-user
|
||||
identity assertion the mode promises. That is an operator misconfiguration, not a fallback.
|
||||
"""
|
||||
spec: Final = to_server_spec(server)
|
||||
if spec is None and server.auth_type == MCPAuth.oauth2_id_jag:
|
||||
raise_public(
|
||||
CredError.of_misconfigured(
|
||||
"oauth2_id_jag requires token_exchange_endpoint, id_jag_resource_token_endpoint, "
|
||||
"client_id, and a client_secret or client_private_key; refusing to fall back to "
|
||||
"a static credential."
|
||||
)
|
||||
)
|
||||
return spec
|
||||
|
||||
|
||||
async def resolve_upstream_auth(
|
||||
*,
|
||||
server: MCPServer,
|
||||
spec: ServerSpec,
|
||||
root_path: str,
|
||||
provider: UpstreamCredentialProvider,
|
||||
subject_token: str | None,
|
||||
user_api_key_auth: UserAPIKeyAuth | None,
|
||||
extra_headers: dict[str, str] | None,
|
||||
) -> tuple[httpx2.Auth | None, dict[str, str] | None]:
|
||||
"""Resolve a v2-owned server's upstream credential into ``(resolved_auth, extra_headers)``.
|
||||
|
||||
On a missing/rejected per-user credential this raises the mode's discovery challenge
|
||||
(authorization_code's browser-OAuth 401, token_exchange's RFC 9728 challenge) or maps any
|
||||
other ``CredError`` onto its public HTTP status; it never returns an error as a value.
|
||||
"""
|
||||
match await resolve_credentials_with_source(provider, to_subject(user_api_key_auth, subject_token), spec):
|
||||
case Ok(credential):
|
||||
auth: Final = credential.auth
|
||||
# NoOpAuth has no header_name and so never conflicts.
|
||||
header_name: Final[str | None] = getattr(auth, "header_name", None)
|
||||
if header_name is None or not extra_headers:
|
||||
source: Final = (
|
||||
AuthResolution.extra_headers
|
||||
if credential.source == AuthResolution.no_auth and extra_headers
|
||||
else credential.source
|
||||
)
|
||||
record_auth_resolution(server.server_id, source)
|
||||
return auth, extra_headers
|
||||
if not has_header(extra_headers, header_name):
|
||||
record_auth_resolution(server.server_id, credential.source)
|
||||
return auth, extra_headers
|
||||
if isinstance(
|
||||
spec.config,
|
||||
(TokenExchangeConfig, AuthorizationCodeConfig, IdJagConfig, ClientCredentialsConfig),
|
||||
):
|
||||
# The resolver owns the credential here (token_exchange's exchanged token,
|
||||
# authorization_code's stored token, id_jag's minted assertion,
|
||||
# client_credentials' gateway-minted M2M token). It is authoritative: a
|
||||
# guardrail such as MCPJWTSigner, static_headers, or any other injected
|
||||
# Authorization must NOT shadow it (otherwise the upstream gets e.g. the
|
||||
# signer's JWT instead of the minted token and rejects it, and for M2M the
|
||||
# one-shot 401 refetch is lost with it). Drop only the header the resolved
|
||||
# credential is about to occupy, so a static credential the operator aimed at a
|
||||
# DIFFERENT header still reaches upstream.
|
||||
record_auth_resolution(server.server_id, credential.source)
|
||||
return auth, without_header(extra_headers, header_name)
|
||||
# Other modes: an Authorization already supplied via extra_headers (a forwarded caller
|
||||
# header or static_headers) is intentional and wins; v1 applies those last.
|
||||
record_auth_resolution(server.server_id, AuthResolution.extra_headers)
|
||||
return None, extra_headers
|
||||
case Error(err):
|
||||
record_auth_resolution(server.server_id, AuthResolution.failed)
|
||||
if err.tag == "unauthorized" and isinstance(spec.config, AuthorizationCodeConfig):
|
||||
# authorization_code's missing per-user token -> the per-server browser-OAuth
|
||||
# challenge, built here where the full MCPServer is in hand.
|
||||
raise_user_oauth_challenge(server, root_path=root_path)
|
||||
if err.tag == "unauthorized" and isinstance(spec.config, TokenExchangeConfig):
|
||||
# token_exchange (OBO): a missing/rejected subject token -> the RFC 9728 challenge
|
||||
# pointing at the IdP the client must SSO with to obtain one, rather than an opaque
|
||||
# 401. No gateway-side browser flow. An IdP step-up rejection (Entra Conditional
|
||||
# Access) threads its claims blob into the challenge for the client to satisfy.
|
||||
raise_token_exchange_challenge(
|
||||
server,
|
||||
root_path=root_path,
|
||||
claims=err.unauthorized.claims,
|
||||
)
|
||||
raise_public(err)
|
||||
|
||||
|
||||
def _stdio_config(server: MCPServer, stdio_env: dict[str, str] | None) -> MCPStdioConfig | None:
|
||||
if not is_mcp_stdio_enabled():
|
||||
raise HTTPException(status_code=403, detail=MCP_STDIO_DISABLED_MESSAGE)
|
||||
if server.command:
|
||||
command: Final = os.path.basename(server.command)
|
||||
lowercase: Final = command.lower()
|
||||
normalized: Final = next(
|
||||
(
|
||||
lowercase.removesuffix(suffix)
|
||||
for suffix in (".exe", ".cmd", ".bat", ".com")
|
||||
if lowercase.endswith(suffix)
|
||||
),
|
||||
lowercase,
|
||||
)
|
||||
if command not in MCP_STDIO_ALLOWED_COMMANDS and normalized not in MCP_STDIO_ALLOWED_COMMANDS:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=f"MCP stdio command '{server.command}' is not in the allowlist ({sorted(MCP_STDIO_ALLOWED_COMMANDS)}). "
|
||||
"Add it to LITELLM_MCP_STDIO_EXTRA_COMMANDS to allow this command.",
|
||||
)
|
||||
if not server.command or server.args is None:
|
||||
return None
|
||||
environment: Final = stdio_env if stdio_env is not None else server.env
|
||||
return MCPStdioConfig(
|
||||
command=server.command,
|
||||
args=server.args,
|
||||
env={"NPM_CONFIG_CACHE": MCP_NPM_CACHE_DIR, **environment} if environment is not None else None,
|
||||
)
|
||||
|
||||
|
||||
async def prepare_upstream_client(
|
||||
server: MCPServer,
|
||||
*,
|
||||
provider: UpstreamCredentialProvider,
|
||||
root_path: str,
|
||||
mcp_auth_header: str | dict[str, str] | None = None,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
stdio_env: dict[str, str] | None = None,
|
||||
subject_token: str | None = None,
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
protocol_version: MCPUpstreamProtocol,
|
||||
sampling_callback: SamplingCallback | None = None,
|
||||
elicitation_callback: ElicitationCallback | None = None,
|
||||
) -> MCPClient:
|
||||
transport: Final[MCPTransportType] = server.transport or MCPTransport.sse
|
||||
server_spec: Final = None if transport == MCPTransport.stdio else to_server_spec_fail_closed(server)
|
||||
spec: Final = (
|
||||
None
|
||||
if server_spec is not None
|
||||
and mcp_auth_header
|
||||
and not isinstance(
|
||||
server_spec.config, (AuthorizationCodeConfig, IdJagConfig, PassthroughConfig, TokenExchangeConfig)
|
||||
)
|
||||
else server_spec
|
||||
)
|
||||
auth_value: Final = await resolve_mcp_auth(server, mcp_auth_header) if spec is None else None
|
||||
auth_header_name: Final = resolved_token_header(server, mcp_auth_header) if spec is None else None
|
||||
timeout: Final = server.timeout if server.timeout is not None else MCP_CLIENT_TIMEOUT
|
||||
if transport == MCPTransport.stdio:
|
||||
config: Final = _stdio_config(server, stdio_env)
|
||||
record_auth_resolution(server.server_id, AuthResolution.not_applicable)
|
||||
return MCPClient(
|
||||
server_url="",
|
||||
transport_type=transport,
|
||||
protocol_version=protocol_version,
|
||||
auth_type=server.auth_type,
|
||||
auth_value=auth_value,
|
||||
timeout=timeout,
|
||||
stdio_config=config,
|
||||
extra_headers=extra_headers,
|
||||
sampling_callback=sampling_callback,
|
||||
elicitation_callback=elicitation_callback,
|
||||
)
|
||||
if spec is not None:
|
||||
inbound_token, forwarded_headers = (
|
||||
take_forwarded_authorization(extra_headers)
|
||||
if isinstance(spec.config, PassthroughConfig)
|
||||
else (subject_token, extra_headers)
|
||||
)
|
||||
per_server_token: Final = (
|
||||
passthrough_token_from_mcp_auth_header(mcp_auth_header)
|
||||
if isinstance(spec.config, PassthroughConfig)
|
||||
else None
|
||||
)
|
||||
resolved_auth, resolved_headers = await resolve_upstream_auth(
|
||||
server=server,
|
||||
spec=spec,
|
||||
provider=provider,
|
||||
root_path=root_path,
|
||||
subject_token=per_server_token if per_server_token is not None else inbound_token,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
extra_headers=forwarded_headers,
|
||||
)
|
||||
return await prepare_mcp_client(
|
||||
server,
|
||||
MCPClient(
|
||||
server_url=server.url or "",
|
||||
transport_type=transport,
|
||||
protocol_version=protocol_version,
|
||||
auth_type=server.auth_type,
|
||||
timeout=timeout,
|
||||
extra_headers=resolved_headers,
|
||||
resolved_auth=resolved_auth,
|
||||
sampling_callback=sampling_callback,
|
||||
elicitation_callback=elicitation_callback,
|
||||
),
|
||||
)
|
||||
aws_auth: Final = (
|
||||
MCPSigV4Auth(
|
||||
aws_access_key_id=server.aws_access_key_id,
|
||||
aws_secret_access_key=server.aws_secret_access_key,
|
||||
aws_session_token=server.aws_session_token,
|
||||
aws_region_name=server.aws_region_name,
|
||||
aws_service_name=server.aws_service_name,
|
||||
aws_role_name=server.aws_role_name,
|
||||
aws_session_name=server.aws_session_name,
|
||||
)
|
||||
if server.auth_type == MCPAuth.aws_sigv4
|
||||
else None
|
||||
)
|
||||
source: Final = (
|
||||
AuthResolution.aws_sigv4
|
||||
if aws_auth is not None
|
||||
else AuthResolution.extra_headers
|
||||
if extra_headers and has_header(extra_headers, auth_header_name or "Authorization")
|
||||
else AuthResolution.per_request_header
|
||||
if mcp_auth_header
|
||||
else AuthResolution.static_token
|
||||
if auth_value
|
||||
else AuthResolution.extra_headers
|
||||
if extra_headers
|
||||
else AuthResolution.no_auth
|
||||
)
|
||||
record_auth_resolution(server.server_id, source)
|
||||
return await prepare_mcp_client(
|
||||
server,
|
||||
MCPClient(
|
||||
server_url=server.url or "",
|
||||
transport_type=transport,
|
||||
protocol_version=protocol_version,
|
||||
auth_type=server.auth_type,
|
||||
auth_value=auth_value,
|
||||
auth_header_name=auth_header_name,
|
||||
timeout=timeout,
|
||||
extra_headers=extra_headers,
|
||||
aws_auth=aws_auth,
|
||||
sampling_callback=sampling_callback,
|
||||
elicitation_callback=elicitation_callback,
|
||||
),
|
||||
)
|
||||
|
|
@ -16,6 +16,7 @@ from pydantic import (
|
|||
JsonValue,
|
||||
PositiveInt,
|
||||
PrivateAttr,
|
||||
TypeAdapter,
|
||||
field_validator,
|
||||
model_validator,
|
||||
)
|
||||
|
|
@ -46,6 +47,8 @@ from litellm.types.mcp import (
|
|||
MCPCredentials,
|
||||
MCPTransport,
|
||||
MCPTransportType,
|
||||
MCPUpstreamProtocol,
|
||||
validate_mcp_protocol_transport,
|
||||
)
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPInfo
|
||||
from litellm.types.proxy.agent_identity import ManagedAgentContext
|
||||
|
|
@ -1673,6 +1676,16 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase):
|
|||
description="Server-managed: set by the endpoint; caller values are overridden.",
|
||||
)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_protocol_transport(self) -> "NewMCPServerRequest":
|
||||
validate_mcp_protocol_transport(
|
||||
TypeAdapter[MCPUpstreamProtocol](MCPUpstreamProtocol).validate_python(
|
||||
(self.mcp_info or {}).get("protocol_version", "auto")
|
||||
),
|
||||
self.transport,
|
||||
)
|
||||
return self
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def validate_transport_fields(cls, values):
|
||||
|
|
@ -1756,6 +1769,18 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase):
|
|||
timeout: float | None = None
|
||||
max_concurrent_requests: int | None = None
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_protocol_transport(self) -> "UpdateMCPServerRequest":
|
||||
if not {"transport", "mcp_info"}.issubset(self.model_fields_set):
|
||||
return self
|
||||
validate_mcp_protocol_transport(
|
||||
TypeAdapter[MCPUpstreamProtocol](MCPUpstreamProtocol).validate_python(
|
||||
(self.mcp_info or {}).get("protocol_version", "auto")
|
||||
),
|
||||
self.transport,
|
||||
)
|
||||
return self
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def validate_transport_fields(cls, values):
|
||||
|
|
|
|||
|
|
@ -44,6 +44,7 @@ from fastapi import (
|
|||
status,
|
||||
)
|
||||
from fastapi.responses import JSONResponse
|
||||
from pydantic import TypeAdapter
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
try:
|
||||
|
|
@ -137,6 +138,7 @@ if MCP_AVAILABLE:
|
|||
def validate_tool_name(name: str) -> _ToolNameValidationResult:
|
||||
return _ToolNameValidationResult()
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.contracts import TargetCatalog
|
||||
from litellm.proxy._experimental.mcp_server.db import (
|
||||
McpIdentifierConflict,
|
||||
approve_mcp_server,
|
||||
|
|
@ -178,6 +180,7 @@ if MCP_AVAILABLE:
|
|||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.server_resolution import (
|
||||
MCPServerTargetCatalog,
|
||||
authorize_mcp_server,
|
||||
resolve_mcp_server,
|
||||
)
|
||||
|
|
@ -237,7 +240,9 @@ if MCP_AVAILABLE:
|
|||
MCPCredentials,
|
||||
MCPGatewaySessionsResponse,
|
||||
MCPGatewaySessionsTerminateResponse,
|
||||
MCPUpstreamProtocol,
|
||||
normalize_upstream_header_name,
|
||||
validate_mcp_protocol_transport,
|
||||
)
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer, PinnedMCPTool
|
||||
|
||||
|
|
@ -2185,18 +2190,16 @@ if MCP_AVAILABLE:
|
|||
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
|
||||
|
||||
client_ip: Final = IPAddressUtils.get_mcp_client_ip(request) if request is not None else None
|
||||
resolved: Final = await resolve_mcp_server(
|
||||
server_id,
|
||||
catalog: Final[TargetCatalog] = MCPServerTargetCatalog(
|
||||
manager=global_mcp_server_manager,
|
||||
temp_lookup=get_cached_temporary_mcp_server,
|
||||
id_client_ip=None,
|
||||
name_client_ip=client_ip,
|
||||
match_name=True,
|
||||
)
|
||||
authorized: Final = await authorize_mcp_server(
|
||||
resolved,
|
||||
authorized: Final = await catalog.resolve(
|
||||
server_id,
|
||||
user_api_key_dict,
|
||||
manager=global_mcp_server_manager,
|
||||
is_admin_view=_user_has_admin_view(user_api_key_dict),
|
||||
not_found_detail={"error": f"MCP server {server_id} not found"},
|
||||
forbidden_detail={"error": f"Access denied to MCP server {server_id}"},
|
||||
|
|
@ -2739,15 +2742,13 @@ if MCP_AVAILABLE:
|
|||
404, so server ids can't be enumerated), using the same allowed-server
|
||||
resolution the MCP gateway enforces on tool calls.
|
||||
"""
|
||||
resolved: Final = await resolve_mcp_server(
|
||||
server_id,
|
||||
catalog: Final[TargetCatalog] = MCPServerTargetCatalog(
|
||||
manager=global_mcp_server_manager,
|
||||
db_lookup=lambda sid: get_mcp_server(prisma_client, sid),
|
||||
)
|
||||
authorized: Final = await authorize_mcp_server(
|
||||
resolved,
|
||||
authorized: Final = await catalog.resolve(
|
||||
server_id,
|
||||
user_api_key_dict,
|
||||
manager=global_mcp_server_manager,
|
||||
is_admin_view=_user_has_admin_view(user_api_key_dict),
|
||||
not_found_detail={"error": f"MCP Server {server_id} not found"},
|
||||
forbidden_detail={
|
||||
|
|
@ -2938,6 +2939,32 @@ if MCP_AVAILABLE:
|
|||
statuses.append(status_obj)
|
||||
return statuses
|
||||
|
||||
def _validate_mcp_protocol_update(
|
||||
payload: UpdateMCPServerRequest,
|
||||
fields_set: set[str],
|
||||
stored: LiteLLM_MCPServerTable | None,
|
||||
read_failed: bool,
|
||||
) -> None:
|
||||
if not {"transport", "mcp_info"}.intersection(fields_set):
|
||||
return
|
||||
if read_failed:
|
||||
raise HTTPException(
|
||||
status_code=503, detail="Cannot validate MCP configuration while stored state is unavailable"
|
||||
)
|
||||
if stored is None:
|
||||
return
|
||||
effective_transport: Final = payload.transport if "transport" in fields_set else stored.transport
|
||||
effective_info: Final = payload.mcp_info if "mcp_info" in fields_set else stored.mcp_info
|
||||
try:
|
||||
validate_mcp_protocol_transport(
|
||||
TypeAdapter[MCPUpstreamProtocol](MCPUpstreamProtocol).validate_python(
|
||||
(effective_info or {}).get("protocol_version", "auto")
|
||||
),
|
||||
effective_transport,
|
||||
)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
|
||||
@router.put(
|
||||
"/server",
|
||||
description="Allows deleting mcp serves in the db",
|
||||
|
|
@ -2986,9 +3013,9 @@ if MCP_AVAILABLE:
|
|||
},
|
||||
)
|
||||
|
||||
# Snapshot the pre-update identity so we can detect a mint-relevant change below. The read is
|
||||
# advisory (it only feeds the stale-token purge decision), so a failure skips the purge with a
|
||||
# warning instead of failing the edit, whose primary job is the update itself.
|
||||
# Snapshot stored configuration for protocol validation and mint-relevant changes below.
|
||||
# Protocol or transport edits require this read; other edits may continue on read failure
|
||||
# while skipping the best-effort stale-token purge.
|
||||
try:
|
||||
old_server_record = await get_mcp_server(prisma_client, payload.server_id)
|
||||
old_server_record_read_failed = False
|
||||
|
|
@ -3001,6 +3028,8 @@ if MCP_AVAILABLE:
|
|||
old_server_record = None
|
||||
old_server_record_read_failed = True
|
||||
|
||||
_validate_mcp_protocol_update(payload, payload_fields_set, old_server_record, old_server_record_read_failed)
|
||||
|
||||
if payload.per_server_oauth_discovery and (old_server_record is not None or old_server_record_read_failed):
|
||||
relay_eligible: Final = old_server_record is not None and is_per_server_oauth_discovery_eligible(
|
||||
payload.auth_type if "auth_type" in payload_fields_set else old_server_record.auth_type,
|
||||
|
|
|
|||
|
|
@ -63,7 +63,14 @@ DEFAULT_SUBJECT_TOKEN_TYPE: Final = "urn:ietf:params:oauth:token-type:access_tok
|
|||
MCPTransportType = Literal[MCPTransport.sse, MCPTransport.http, MCPTransport.stdio]
|
||||
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"]
|
||||
MCPUpstreamProtocol = MCPLegacyVersion | Literal["auto", "2026-07-28"]
|
||||
|
||||
|
||||
def validate_mcp_protocol_transport(protocol_version: MCPUpstreamProtocol, transport: MCPTransportType) -> None:
|
||||
if protocol_version == "2026-07-28" and transport == MCPTransport.sse:
|
||||
raise ValueError("Modern MCP requires HTTP or stdio transport")
|
||||
|
||||
|
||||
MCPAdvertisedVersions = Annotated[tuple[MCPLegacyVersion, ...], Field(min_length=1)]
|
||||
MCPSpecVersionType = Literal[
|
||||
MCPSpecVersion.nov_2024,
|
||||
|
|
|
|||
|
|
@ -13,13 +13,14 @@ from litellm.types.mcp import (
|
|||
MCPTransportType,
|
||||
MCPUpstreamProtocol,
|
||||
normalize_upstream_header_name,
|
||||
validate_mcp_protocol_transport,
|
||||
)
|
||||
|
||||
|
||||
# MCPInfo now allows arbitrary additional fields for custom metadata
|
||||
def _validate_mcp_protocol_metadata(value: dict[str, object]) -> dict[str, object]:
|
||||
if "protocol_version" in value:
|
||||
TypeAdapter(MCPUpstreamProtocol).validate_python(value["protocol_version"])
|
||||
TypeAdapter[MCPUpstreamProtocol](MCPUpstreamProtocol).validate_python(value["protocol_version"])
|
||||
return value
|
||||
|
||||
|
||||
|
|
@ -277,9 +278,10 @@ class MCPServer(BaseModel):
|
|||
@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.protocol_version = TypeAdapter[MCPUpstreamProtocol](MCPUpstreamProtocol).validate_python(
|
||||
self.mcp_info.get("protocol_version", "auto")
|
||||
)
|
||||
validate_mcp_protocol_transport(self.protocol_version, self.transport)
|
||||
return self
|
||||
|
||||
@model_validator(mode="after")
|
||||
|
|
|
|||
|
|
@ -390,3 +390,97 @@ def test_ui_session_lists_and_fetches_team_granted_config_server(
|
|||
assert detail.status_code == 200, f"Team-granted server detail access should succeed: {detail.text}"
|
||||
assert detail.json()["server_id"] == server_id, detail.text
|
||||
assert detail.json()["alias"] == alias, detail.text
|
||||
|
||||
|
||||
@pytest.mark.parametrize("explicit_transport", [False, True])
|
||||
def test_modern_sse_registration_rejected_without_saving(gateway: Gateway, explicit_transport: bool) -> None:
|
||||
identity: Final = str(uuid.uuid4())
|
||||
with mcp_peer() as peer:
|
||||
response: Final = gateway.request("POST", "/v1/mcp/server", {
|
||||
"server_id": identity, "server_name": "invalid" + uuid.uuid4().hex[:8],
|
||||
"url": peer.url, "mcp_info": {"protocol_version": "2026-07-28"},
|
||||
**({"transport": "sse"} if explicit_transport else {}),
|
||||
})
|
||||
try:
|
||||
assert response.status_code == 422, response.text
|
||||
assert "Modern MCP requires HTTP or stdio" in response.text
|
||||
assert identity not in _servers(gateway)
|
||||
assert peer.drain() == (), "Rejected configuration reached upstream"
|
||||
finally:
|
||||
if response.status_code == 201:
|
||||
delete_mcp(gateway, identity)
|
||||
|
||||
|
||||
def test_protocol_transport_updates_validate_effective_configuration(gateway: Gateway) -> None:
|
||||
from integration._support.database import read_rows
|
||||
|
||||
with mcp_peer() as peer, gateway.scenario() as scenario:
|
||||
identity: Final = register_mcp(scenario, peer, "protocol" + uuid.uuid4().hex[:8])
|
||||
modern: Final = gateway.request("PUT", "/v1/mcp/server", {
|
||||
"server_id": identity, "mcp_info": {"protocol_version": "2026-07-28"},
|
||||
})
|
||||
assert modern.status_code == 202, modern.text
|
||||
assert modern.json()["transport"] == "http"
|
||||
changed: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "description": "renamed"})
|
||||
assert changed.status_code == 202, changed.text
|
||||
assert changed.json()["mcp_info"]["protocol_version"] == "2026-07-28"
|
||||
snapshot: Final = read_rows('SELECT transport, mcp_info, updated_at::text FROM "LiteLLM_MCPServerTable" WHERE server_id=%s', (identity,))
|
||||
peer.drain()
|
||||
rejected: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "transport": "sse", "url": peer.url})
|
||||
assert rejected.status_code == 400, rejected.text
|
||||
assert "Modern MCP requires HTTP or stdio" in rejected.text
|
||||
assert read_rows('SELECT transport, mcp_info, updated_at::text FROM "LiteLLM_MCPServerTable" WHERE server_id=%s', (identity,)) == snapshot
|
||||
assert tool_calls(peer.drain()) == ()
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
assert call_tool(gateway, key, identity, "add", ADD).status_code == 200
|
||||
|
||||
legacy: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "transport": "sse", "url": peer.url, "mcp_info": {}})
|
||||
assert legacy.status_code == 202, legacy.text
|
||||
legacy_snapshot: Final = read_rows('SELECT transport, mcp_info, updated_at::text FROM "LiteLLM_MCPServerTable" WHERE server_id=%s', (identity,))
|
||||
rejected_protocol: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "mcp_info": {"protocol_version": "2026-07-28"}})
|
||||
assert rejected_protocol.status_code == 400, rejected_protocol.text
|
||||
assert read_rows('SELECT transport, mcp_info, updated_at::text FROM "LiteLLM_MCPServerTable" WHERE server_id=%s', (identity,)) == legacy_snapshot
|
||||
repaired: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "transport": "http", "url": peer.url, "mcp_info": {"protocol_version": "2026-07-28"}})
|
||||
assert repaired.status_code == 202, repaired.text
|
||||
assert call_tool(gateway, key, identity, "add", ADD).status_code == 200
|
||||
|
||||
|
||||
@pytest.mark.parametrize("rename", [False, True])
|
||||
def test_concurrent_protocol_transport_edits_cannot_save_incompatible_configuration(
|
||||
gateway: Gateway, peer: Gateway, rename: bool
|
||||
) -> None:
|
||||
import os
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
import psycopg
|
||||
from integration._support.database import read_rows
|
||||
|
||||
with mcp_peer() as upstream, gateway.scenario() as scenario:
|
||||
identity: Final = register_mcp(scenario, upstream, "race" + uuid.uuid4().hex[:8])
|
||||
with ThreadPoolExecutor(max_workers=2) as pool:
|
||||
# Hold the row so both workers read the old configuration before either can write.
|
||||
with psycopg.connect(os.environ["DATABASE_URL"]) as blocker:
|
||||
blocker.execute('SELECT server_id FROM "LiteLLM_MCPServerTable" WHERE server_id=%s FOR UPDATE', (identity,))
|
||||
protocol = pool.submit(gateway.request, "PUT", "/v1/mcp/server", {
|
||||
"server_id": identity, "mcp_info": {"protocol_version": "2026-07-28"},
|
||||
**({"alias": "renamed" + uuid.uuid4().hex[:8]} if rename else {}),
|
||||
})
|
||||
transport = pool.submit(peer.request, "PUT", "/v1/mcp/server", {
|
||||
"server_id": identity, "transport": "sse", "url": upstream.url,
|
||||
})
|
||||
eventually(
|
||||
lambda: read_rows("SELECT pid FROM pg_stat_activity WHERE datname=current_database() AND wait_event_type='Lock' AND query LIKE %s", ('%LiteLLM_MCPServerTable%',)),
|
||||
lambda rows: len(rows) >= 2,
|
||||
seconds=3,
|
||||
)
|
||||
responses: Final = [protocol.result(), transport.result()]
|
||||
assert sorted(response.status_code for response in responses) == [202, 400], [r.text for r in responses]
|
||||
saved: Final = read_rows('SELECT transport, mcp_info FROM "LiteLLM_MCPServerTable" WHERE server_id=%s', (identity,))[0]
|
||||
assert saved["transport"] == "http" or saved["mcp_info"] != {"protocol_version": "2026-07-28"}
|
||||
repaired: Final = gateway.request("PUT", "/v1/mcp/server", {
|
||||
"server_id": identity, "transport": "http", "url": upstream.url,
|
||||
"mcp_info": {"protocol_version": "2026-07-28"},
|
||||
})
|
||||
assert repaired.status_code == 202, repaired.text
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
assert call_tool(gateway, key, identity, "add", ADD).status_code == 200
|
||||
|
|
|
|||
|
|
@ -195,3 +195,52 @@ def test_pinned_revision_pairs_list_and_call_through_gateway(
|
|||
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
|
||||
|
||||
|
||||
@pytest.mark.parametrize("downstream", ("2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"))
|
||||
@pytest.mark.parametrize("peer_kind", ("http", "stdio"))
|
||||
def test_legacy_gateway_calls_modern_upstream_without_initialize(
|
||||
gateway: Gateway,
|
||||
downstream: str,
|
||||
peer_kind: PeerKind,
|
||||
) -> 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 = "modern" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias, mcp_info={"protocol_version": "2026-07-28"})
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
client: Final = MCPClient(
|
||||
server_url=str(gateway.client.base_url).rstrip("/") + "/mcp",
|
||||
transport_type=MCPTransport.http,
|
||||
protocol_version=downstream,
|
||||
extra_headers={"Authorization": f"Bearer {key}"},
|
||||
)
|
||||
|
||||
async def exercise() -> None:
|
||||
listed: Final = await client.list_tools(raise_on_error=True)
|
||||
assert f"{alias}-add" in tuple(tool.name for tool in listed)
|
||||
called: Final = await client.call_tool(
|
||||
CallToolRequestParams(name=f"{alias}-add", arguments={"a": 2, "b": 3}),
|
||||
raise_on_error=True,
|
||||
)
|
||||
assert called.is_error is False
|
||||
assert called.content[0].text == "5"
|
||||
|
||||
peer.drain()
|
||||
asyncio.run(exercise())
|
||||
observed: Final = peer.drain()
|
||||
assert len(tool_calls(observed)) == 1
|
||||
assert all(row["body"].get("method") not in ("initialize", "notifications/initialized") for row in observed)
|
||||
requests: Final = tuple(row for row in observed if "id" in row["body"])
|
||||
assert requests
|
||||
for row in requests:
|
||||
metadata: Final = row["body"]["params"]["_meta"]
|
||||
assert metadata["io.modelcontextprotocol/protocolVersion"] == "2026-07-28"
|
||||
assert "io.modelcontextprotocol/clientCapabilities" in metadata
|
||||
assert b"mcp-session-id" not in row.get("headers", {})
|
||||
|
|
|
|||
|
|
@ -3028,9 +3028,101 @@ async def test_configured_upstream_revision_is_offered_and_checked(revision, acc
|
|||
await client.list_tools(raise_on_error=True)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("revision", ["2026-07-28", "unknown", "", None])
|
||||
@pytest.mark.parametrize("revision", ["unknown", "", None])
|
||||
def test_upstream_protocol_configuration_rejects_unavailable_modes(revision):
|
||||
from pydantic import ValidationError
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
MCPClient(protocol_version=revision)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("accepted", [True, False])
|
||||
async def test_modern_upstream_requests_are_self_contained_without_initialization(accepted: bool) -> None:
|
||||
from queue import SimpleQueue
|
||||
|
||||
from mcp.types import DiscoverResult, ToolsCapability
|
||||
|
||||
methods: Final[SimpleQueue[str]] = SimpleQueue()
|
||||
|
||||
def respond(request: httpx2.Request) -> httpx2.Response:
|
||||
if request.method != "POST":
|
||||
return httpx2.Response(405)
|
||||
payload: Final = _JSONRPC_MESSAGE_ADAPTER.validate_json(request.content)
|
||||
assert isinstance(payload, JSONRPCRequest)
|
||||
methods.put(payload.method)
|
||||
assert payload.method != "initialize", "Modern operations must not establish a legacy session"
|
||||
assert "mcp-session-id" not in request.headers
|
||||
assert request.headers["mcp-protocol-version"] == "2026-07-28"
|
||||
assert request.headers["authorization"] == "Bearer upstream-credential"
|
||||
assert payload.params is not None
|
||||
metadata: Final = payload.params["_meta"]
|
||||
assert metadata["io.modelcontextprotocol/protocolVersion"] == "2026-07-28"
|
||||
assert "io.modelcontextprotocol/clientCapabilities" in metadata
|
||||
if payload.method == "server/discover":
|
||||
discovery: Final = DiscoverResult(
|
||||
supported_versions=["2026-07-28"] if accepted else ["2025-11-25"],
|
||||
capabilities=ServerCapabilities(tools=ToolsCapability()),
|
||||
instructions="modern instructions",
|
||||
)
|
||||
return httpx2.Response(
|
||||
200,
|
||||
json={
|
||||
"jsonrpc": "2.0",
|
||||
"id": payload.id,
|
||||
"result": discovery.model_dump(by_alias=True, exclude_none=True),
|
||||
},
|
||||
)
|
||||
assert accepted, "Rejected negotiation must prevent upstream execution"
|
||||
if payload.method == "tools/list":
|
||||
return httpx2.Response(
|
||||
200,
|
||||
json={
|
||||
"jsonrpc": "2.0",
|
||||
"id": payload.id,
|
||||
"result": {
|
||||
"resultType": "complete",
|
||||
"cacheScope": "private",
|
||||
"ttlMs": 0,
|
||||
"tools": [{"name": "add", "inputSchema": {"type": "object"}}],
|
||||
},
|
||||
},
|
||||
)
|
||||
assert payload.method == "tools/call"
|
||||
assert payload.params["arguments"] == {"a": 2, "b": 3}
|
||||
return httpx2.Response(
|
||||
200,
|
||||
json={
|
||||
"jsonrpc": "2.0",
|
||||
"id": payload.id,
|
||||
"result": {"resultType": "complete", "content": [{"type": "text", "text": "5"}], "isError": False},
|
||||
},
|
||||
)
|
||||
|
||||
client: Final = _MockTransportClient(
|
||||
respond,
|
||||
server_url="https://example.com/mcp",
|
||||
protocol_version="2026-07-28",
|
||||
auth_type=MCPAuth.bearer_token,
|
||||
auth_value="upstream-credential",
|
||||
)
|
||||
params: Final = CallToolRequestParams(name="add", arguments={"a": 2, "b": 3})
|
||||
if accepted:
|
||||
result: Final = await client.call_tool(params, raise_on_error=True)
|
||||
assert result.content[0].text == "5"
|
||||
assert not result.is_error
|
||||
assert client._last_initialize_instructions == "modern instructions"
|
||||
assert tuple(methods.get_nowait() for _ in range(methods.qsize())) == (
|
||||
"server/discover",
|
||||
"tools/call",
|
||||
"tools/list",
|
||||
)
|
||||
else:
|
||||
with pytest.raises((MCPError, RuntimeError), match="protocol version"):
|
||||
await client.call_tool(params, raise_on_error=True)
|
||||
assert tuple(methods.get_nowait() for _ in range(methods.qsize())) == ("server/discover",)
|
||||
|
||||
|
||||
def test_modern_upstream_rejects_legacy_sse_transport() -> None:
|
||||
with pytest.raises(ValueError, match="transport"):
|
||||
MCPClient(protocol_version="2026-07-28", transport_type=MCPTransport.sse)
|
||||
|
|
|
|||
|
|
@ -111,7 +111,7 @@ async def test_mcp_cost_tracking():
|
|||
local_mcp_server_manager = MCPServerManager()
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
|
||||
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
|
||||
mock_client_constructor,
|
||||
):
|
||||
# Load the server config
|
||||
|
|
@ -244,7 +244,7 @@ async def test_mcp_cost_tracking_per_tool():
|
|||
local_mcp_server_manager = MCPServerManager()
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
|
||||
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
|
||||
mock_client_constructor,
|
||||
):
|
||||
# Load the server config with per-tool costs
|
||||
|
|
@ -417,7 +417,7 @@ async def test_mcp_tool_call_hook():
|
|||
local_mcp_server_manager = MCPServerManager()
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
|
||||
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
|
||||
mock_client_constructor,
|
||||
):
|
||||
# Load the server config
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ from unittest.mock import AsyncMock, MagicMock
|
|||
|
||||
import pytest
|
||||
from prisma import Json, models
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.db import (
|
||||
create_mcp_server,
|
||||
|
|
@ -35,7 +36,7 @@ def _mock_prisma():
|
|||
mock_prisma.db.litellm_mcpservertable.update = AsyncMock(return_value=row)
|
||||
mock_prisma.db.litellm_mcpservertable.create = AsyncMock(return_value=row)
|
||||
mock_prisma.db.litellm_mcpservertable.find_first = AsyncMock(return_value=None)
|
||||
mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=None)
|
||||
mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=row)
|
||||
tx_client = MagicMock()
|
||||
tx_client.execute_raw = AsyncMock()
|
||||
tx_client.litellm_mcpservertable = mock_prisma.db.litellm_mcpservertable
|
||||
|
|
@ -1141,6 +1142,42 @@ async def test_set_mcp_server_pinned_tools_writes_the_snapshot_and_null_clears_i
|
|||
@pytest.mark.asyncio
|
||||
async def test_set_mcp_server_pinned_tools_on_a_missing_server_writes_nothing():
|
||||
mock_prisma = _mock_prisma()
|
||||
mock_prisma.db.litellm_mcpservertable.find_unique.return_value = None
|
||||
|
||||
assert await set_mcp_server_pinned_tools(mock_prisma, "ghost", None, "admin") is None
|
||||
mock_prisma.db.litellm_mcpservertable.update.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("protocol_only", [False, True])
|
||||
async def test_protocol_update_revalidates_current_stored_configuration_before_writing(protocol_only: bool):
|
||||
prisma = _mock_prisma()
|
||||
table = prisma.db.litellm_mcpservertable
|
||||
table.find_unique.return_value = models.LiteLLM_MCPServerTable.model_construct(
|
||||
server_id="test-server", transport="sse" if protocol_only else "http",
|
||||
mcp_info={} if protocol_only else {"protocol_version": "2026-07-28"}, env={}, env_vars=[],
|
||||
)
|
||||
payload = UpdateMCPServerRequest.model_validate({
|
||||
"server_id": "test-server",
|
||||
**({"mcp_info": {"protocol_version": "2026-07-28"}} if protocol_only else {"transport": "sse", "url": "https://upstream.example/sse"}),
|
||||
})
|
||||
with pytest.raises(HTTPException) as error:
|
||||
await update_mcp_server(prisma, payload, "admin")
|
||||
assert error.value.status_code == 400
|
||||
assert "Modern MCP requires HTTP or stdio" in str(error.value.detail)
|
||||
table.update.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("clear_alias", [False, True])
|
||||
async def test_protocol_update_preserves_missing_server_without_writing(clear_alias: bool):
|
||||
prisma = _mock_prisma()
|
||||
table = prisma.db.litellm_mcpservertable
|
||||
table.find_unique.return_value = None
|
||||
payload = UpdateMCPServerRequest.model_validate({
|
||||
"server_id": "missing", "mcp_info": {"protocol_version": "2026-07-28"},
|
||||
**({"alias": None} if clear_alias else {}),
|
||||
})
|
||||
result = await update_mcp_server(prisma, payload, "admin")
|
||||
assert result is None
|
||||
table.update.assert_not_awaited()
|
||||
|
|
|
|||
|
|
@ -71,7 +71,7 @@ async def test_mcp_server_manager_https_server():
|
|||
return mock_client
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
|
||||
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
|
||||
mock_client_constructor,
|
||||
):
|
||||
await mcp_server_manager.load_servers_from_config(
|
||||
|
|
@ -179,7 +179,7 @@ async def test_mcp_http_transport_list_tools_mock():
|
|||
return mock_client
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
|
||||
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
|
||||
mock_client_constructor,
|
||||
):
|
||||
# Load server config with HTTP transport
|
||||
|
|
@ -256,7 +256,7 @@ async def test_mcp_http_transport_call_tool_mock():
|
|||
return mock_client
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
|
||||
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
|
||||
mock_client_constructor,
|
||||
):
|
||||
# Load server config with HTTP transport
|
||||
|
|
@ -322,7 +322,7 @@ async def test_mcp_http_transport_call_tool_error_mock():
|
|||
return mock_client
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
|
||||
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
|
||||
mock_client_constructor,
|
||||
):
|
||||
# Load server config with HTTP transport
|
||||
|
|
@ -1093,7 +1093,7 @@ async def test_list_tools_only_returns_allowed_servers(monkeypatch):
|
|||
return mock_client
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
|
||||
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
|
||||
mock_client_constructor,
|
||||
):
|
||||
# Call list_tools
|
||||
|
|
@ -1390,7 +1390,7 @@ async def test_mcp_server_manager_alias_tool_prefixing():
|
|||
return mock_client
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
|
||||
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
|
||||
mock_client_constructor,
|
||||
):
|
||||
# Get tools from server
|
||||
|
|
@ -1450,7 +1450,7 @@ async def test_mcp_server_manager_server_name_tool_prefixing():
|
|||
return mock_client
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
|
||||
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
|
||||
mock_client_constructor,
|
||||
):
|
||||
# Get tools from server
|
||||
|
|
@ -1510,7 +1510,7 @@ async def test_mcp_server_manager_server_id_tool_prefixing():
|
|||
return mock_client
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
|
||||
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
|
||||
mock_client_constructor,
|
||||
):
|
||||
# Get tools from server
|
||||
|
|
@ -2506,7 +2506,7 @@ async def test_filter_tools_by_allowed_tools_integration():
|
|||
|
||||
# Mock the MCPClient constructor
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
|
||||
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
|
||||
mock_client_constructor,
|
||||
):
|
||||
# Call _get_tools_from_mcp_servers which should apply the filtering
|
||||
|
|
@ -2620,7 +2620,7 @@ async def test_filter_tools_by_disallowed_tools_integration():
|
|||
|
||||
# Mock the MCPClient constructor
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
|
||||
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
|
||||
mock_client_constructor,
|
||||
):
|
||||
# Call _get_tools_from_mcp_servers which should apply the filtering
|
||||
|
|
@ -2722,7 +2722,7 @@ async def test_filter_tools_no_restrictions_integration():
|
|||
|
||||
# Mock the MCPClient constructor
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
|
||||
"litellm.proxy._experimental.mcp_server.upstream.MCPClient",
|
||||
mock_client_constructor,
|
||||
):
|
||||
# Call _get_tools_from_mcp_servers which should apply the filtering
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
from litellm.proxy._experimental.mcp_server.upstream import resolve_upstream_auth
|
||||
import importlib
|
||||
import asyncio
|
||||
import functools
|
||||
|
|
@ -94,7 +95,7 @@ async def test_manager_sampling_preserves_explicit_headers_without_ambient_conte
|
|||
client.call_tool = AsyncMock(return_value=CallToolResult(content=[]))
|
||||
assert legacy_server.get_active_auth_context() is None
|
||||
with (
|
||||
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient", return_value=client) as factory,
|
||||
patch("litellm.proxy._experimental.mcp_server.upstream.MCPClient", return_value=client) as factory,
|
||||
patch("litellm.proxy._experimental.mcp_server.sampling_handler.handle_sampling_create_message", sampling),
|
||||
):
|
||||
await MCPServerManager()._call_regular_mcp_tool(
|
||||
|
|
@ -1343,7 +1344,7 @@ class TestMCPServerManager:
|
|||
"ensure_oauth_metadata_discovered",
|
||||
new=ensure_oauth_metadata_discovered,
|
||||
),
|
||||
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient"),
|
||||
patch("litellm.proxy._experimental.mcp_server.upstream.MCPClient"),
|
||||
):
|
||||
await manager._create_mcp_client(server)
|
||||
|
||||
|
|
@ -3900,10 +3901,10 @@ class TestMCPServerManager:
|
|||
)
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.resolve_mcp_auth",
|
||||
"litellm.proxy._experimental.mcp_server.upstream.resolve_mcp_auth",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_resolve,
|
||||
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient") as mock_client_cls,
|
||||
patch("litellm.proxy._experimental.mcp_server.upstream.MCPClient") as mock_client_cls,
|
||||
):
|
||||
await manager._create_mcp_client(server=server, extra_headers={"Authorization": "Bearer upstream-token"})
|
||||
mock_resolve.assert_not_awaited()
|
||||
|
|
@ -3953,10 +3954,10 @@ class TestMCPServerManager:
|
|||
)
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.resolve_mcp_auth",
|
||||
"litellm.proxy._experimental.mcp_server.upstream.resolve_mcp_auth",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_resolve,
|
||||
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient") as mock_client_cls,
|
||||
patch("litellm.proxy._experimental.mcp_server.upstream.MCPClient") as mock_client_cls,
|
||||
):
|
||||
await manager._create_mcp_client(
|
||||
server=server,
|
||||
|
|
@ -3996,10 +3997,10 @@ class TestMCPServerManager:
|
|||
)
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.resolve_mcp_auth",
|
||||
"litellm.proxy._experimental.mcp_server.upstream.resolve_mcp_auth",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_resolve,
|
||||
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient") as mock_client_cls,
|
||||
patch("litellm.proxy._experimental.mcp_server.upstream.MCPClient") as mock_client_cls,
|
||||
):
|
||||
await manager._create_mcp_client(
|
||||
server=server,
|
||||
|
|
@ -13707,7 +13708,8 @@ async def test_debug_resolution_matches_final_header_conflict_winner(_mcp_reques
|
|||
"none": NoneConfig(),
|
||||
}[config]
|
||||
try:
|
||||
auth, remaining = await MCPServerManager()._resolve_v2_auth(
|
||||
auth, remaining = await resolve_upstream_auth(
|
||||
root_path="",
|
||||
server=MCPServer(
|
||||
server_id="s",
|
||||
name="s",
|
||||
|
|
@ -15085,7 +15087,7 @@ async def test_client_sampling_does_not_fill_explicit_context_from_another_ambie
|
|||
try:
|
||||
legacy_server.set_auth_context(UserAPIKeyAuth(user_id="unrelated"), raw_headers={"authorization": "unrelated-credential"}, client_ip="192.0.2.99")
|
||||
with (
|
||||
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient") as factory,
|
||||
patch("litellm.proxy._experimental.mcp_server.upstream.MCPClient") as factory,
|
||||
patch("litellm.proxy._experimental.mcp_server.sampling_handler.handle_sampling_create_message", sampling),
|
||||
):
|
||||
if legacy_factory:
|
||||
|
|
@ -15800,3 +15802,57 @@ class TestToolCatalogGuard:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
server=server,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("command,args", [(None, []), ("python", None), ("blocked-executable", [])])
|
||||
async def test_upstream_preparation_rejects_blocked_or_preserves_incomplete_stdio_config(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
command: str | None,
|
||||
args: list[str] | None,
|
||||
) -> None:
|
||||
monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", "true")
|
||||
server: Final = MCPServer(server_id="stdio", name="stdio", transport=MCPTransport.stdio, command=command, args=args)
|
||||
if command == "blocked-executable":
|
||||
with pytest.raises(HTTPException) as error:
|
||||
await MCPServerManager()._create_mcp_client(server)
|
||||
assert error.value.status_code == 403
|
||||
assert "not in the allowlist" in error.value.detail
|
||||
else:
|
||||
client: Final = await MCPServerManager()._create_mcp_client(server)
|
||||
assert client.stdio_config is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upstream_preparation_preserves_windows_command_and_caller_environment(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from litellm.constants import MCP_NPM_CACHE_DIR
|
||||
|
||||
monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", "true")
|
||||
environment: Final = {"PEER_USER": "alice"}
|
||||
server: Final = MCPServer(
|
||||
server_id="stdio", name="stdio", transport=MCPTransport.stdio, command="python.exe", args=[]
|
||||
)
|
||||
client: Final = await MCPServerManager()._create_mcp_client(server, stdio_env=environment)
|
||||
assert client.stdio_config == {
|
||||
"command": "python.exe",
|
||||
"args": [],
|
||||
"env": {"PEER_USER": "alice", "NPM_CONFIG_CACHE": MCP_NPM_CACHE_DIR},
|
||||
}
|
||||
assert environment == {"PEER_USER": "alice"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upstream_preparation_honors_case_sensitive_extra_command(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
from litellm.proxy._experimental.mcp_server import upstream
|
||||
|
||||
monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", "true")
|
||||
monkeypatch.setattr(upstream, "MCP_STDIO_ALLOWED_COMMANDS", frozenset({"CustomRunner"}))
|
||||
server: Final = MCPServer(
|
||||
server_id="custom-stdio", name="custom-stdio", transport=MCPTransport.stdio,
|
||||
command="/opt/tools/CustomRunner", args=[],
|
||||
)
|
||||
client: Final = await MCPServerManager()._create_mcp_client(server)
|
||||
assert client.stdio_config is not None
|
||||
assert client.stdio_config["command"] == "/opt/tools/CustomRunner"
|
||||
|
|
|
|||
|
|
@ -851,3 +851,28 @@ def test_the_openapi_arm_keeps_the_shared_client_when_no_guard_is_needed(resolve
|
|||
assert not client.client.event_hooks.get("request")
|
||||
finally:
|
||||
_request_resolved_auth_headers.reset(token)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("mode", [MCPAuth.true_passthrough, MCPAuth.oauth_delegate])
|
||||
@pytest.mark.parametrize("per_server", [None, "Bearer per-server"])
|
||||
async def test_openapi_passthrough_preparation_preserves_credential_precedence(
|
||||
mode: MCPAuth, per_server: str | None,
|
||||
) -> None:
|
||||
from typing import Final
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
|
||||
server: Final = MCPServer(
|
||||
server_id="openapi-passthrough", name="openapi-passthrough", transport=MCPTransport.http,
|
||||
url="https://upstream.example", spec_path="https://upstream.example/openapi.json", auth_type=mode,
|
||||
)
|
||||
headers: Final = {"authorization": "Bearer forwarded", "X-Trace": "trace"}
|
||||
resolved, remaining = await MCPServerManager().resolve_openapi_upstream_auth(
|
||||
mcp_server=server, oauth2_headers=None, raw_headers=None,
|
||||
mcp_auth_header=per_server, user_api_key_auth=UserAPIKeyAuth(user_id="alice"),
|
||||
forwarded_headers=headers,
|
||||
)
|
||||
assert resolved == {"Authorization": per_server or "Bearer forwarded"}
|
||||
assert remaining == {"X-Trace": "trace"}
|
||||
assert headers == {"authorization": "Bearer forwarded", "X-Trace": "trace"}
|
||||
|
|
|
|||
|
|
@ -136,7 +136,7 @@ async def test_prompt_sampling_receives_explicit_operation_caller_headers_and_ip
|
|||
sampling = AsyncMock()
|
||||
with (
|
||||
patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[upstream])),
|
||||
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient", return_value=client) as factory,
|
||||
patch("litellm.proxy._experimental.mcp_server.upstream.MCPClient", return_value=client) as factory,
|
||||
patch("litellm.proxy._experimental.mcp_server.sampling_handler.handle_sampling_create_message", sampling),
|
||||
):
|
||||
result = await GatewayOperations().execute(
|
||||
|
|
|
|||
|
|
@ -10,7 +10,9 @@ from unittest.mock import Mock
|
|||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.contracts import TargetCatalog
|
||||
from litellm.proxy._experimental.mcp_server.server_resolution import (
|
||||
MCPServerTargetCatalog,
|
||||
ResolutionSource,
|
||||
ResolvedMCPServer,
|
||||
authorize_mcp_server,
|
||||
|
|
@ -460,3 +462,94 @@ async def test_missing_alias_does_not_produce_a_resolution() -> None:
|
|||
manager: Final = _manager()
|
||||
assert await resolve_mcp_server("missing", manager=manager, match_name=True) is None
|
||||
manager.name_lookup_spy.assert_called_once_with("missing", None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("source", ["db", "registry", "temp"])
|
||||
@pytest.mark.parametrize("allowed", [False, True])
|
||||
async def test_target_catalog_authorizes_canonical_identity(source: ResolutionSource, allowed: bool) -> None:
|
||||
server: Final = _runtime_server()
|
||||
manager: Final = _manager(
|
||||
servers_by_name={"requested-alias": server},
|
||||
allowed_server_ids=(server.server_id,) if allowed else ("requested-alias",),
|
||||
)
|
||||
|
||||
async def database_lookup(server_id: str) -> LiteLLM_MCPServerTable | None:
|
||||
return _table_server(server.server_id)
|
||||
|
||||
async def temporary_lookup(server_id: str) -> MCPServer | None:
|
||||
return server
|
||||
|
||||
catalog: Final[TargetCatalog] = MCPServerTargetCatalog(
|
||||
manager=manager,
|
||||
db_lookup=database_lookup if source == "db" else None,
|
||||
temp_lookup=temporary_lookup if source == "temp" else None,
|
||||
match_name=True,
|
||||
)
|
||||
operation: Final = catalog.resolve(
|
||||
"requested-alias",
|
||||
_auth(),
|
||||
is_admin_view=False,
|
||||
not_found_detail={"error": "missing"},
|
||||
forbidden_detail={"error": "denied"},
|
||||
non_admin_missing="forbidden",
|
||||
)
|
||||
if allowed and source != "temp":
|
||||
result: Final = await operation
|
||||
assert result.table.server_id == server.server_id
|
||||
assert result.source == source
|
||||
assert result.runtime is (None if source == "db" else server)
|
||||
else:
|
||||
with pytest.raises(HTTPException) as error:
|
||||
await operation
|
||||
assert (error.value.status_code, error.value.detail) == (403, {"error": "denied"})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"admin,missing,status_code", [(True, "forbidden", 404), (False, "forbidden", 403), (False, "not_found", 404)]
|
||||
)
|
||||
async def test_target_catalog_preserves_missing_target_policy(
|
||||
admin: bool,
|
||||
missing: Literal["forbidden", "not_found"],
|
||||
status_code: int,
|
||||
) -> None:
|
||||
catalog: Final[TargetCatalog] = MCPServerTargetCatalog(manager=_manager())
|
||||
with pytest.raises(HTTPException) as error:
|
||||
await catalog.resolve(
|
||||
"missing",
|
||||
_auth(),
|
||||
is_admin_view=admin,
|
||||
not_found_detail={"error": "missing"},
|
||||
forbidden_detail={"error": "denied"},
|
||||
non_admin_missing=missing,
|
||||
)
|
||||
assert error.value.status_code == status_code
|
||||
assert error.value.detail == {"error": "missing" if status_code == 404 else "denied"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_target_catalog_does_not_reuse_admin_authorization_for_another_caller() -> None:
|
||||
server: Final = _runtime_server()
|
||||
manager: Final = _manager(servers_by_id={server.server_id: server})
|
||||
catalog: Final[TargetCatalog] = MCPServerTargetCatalog(manager=manager)
|
||||
admin: Final = await catalog.resolve(
|
||||
server.server_id,
|
||||
UserAPIKeyAuth(user_id="admin"),
|
||||
is_admin_view=True,
|
||||
not_found_detail={"error": "missing"},
|
||||
forbidden_detail={"error": "denied"},
|
||||
non_admin_missing="forbidden",
|
||||
)
|
||||
assert admin.runtime is server
|
||||
with pytest.raises(HTTPException) as error:
|
||||
await catalog.resolve(
|
||||
server.server_id,
|
||||
_auth(),
|
||||
is_admin_view=False,
|
||||
not_found_detail={"error": "missing"},
|
||||
forbidden_detail={"error": "denied"},
|
||||
non_admin_missing="forbidden",
|
||||
)
|
||||
assert (error.value.status_code, error.value.detail) == (403, {"error": "denied"})
|
||||
manager.allowed_servers_spy.assert_called_once_with(_auth())
|
||||
|
|
|
|||
|
|
@ -11059,3 +11059,89 @@ class TestMCPServerResolutionCharacterization:
|
|||
health_check.assert_not_awaited()
|
||||
effects.assert_no_writes()
|
||||
assert httpx_mock.calls.call_count == 0
|
||||
|
||||
|
||||
@pytest.mark.parametrize("explicit_transport", [False, True])
|
||||
def test_modern_sse_create_is_rejected_before_persistence(explicit_transport: bool) -> None:
|
||||
with pytest.raises(ValidationError, match="Modern MCP requires HTTP or stdio"):
|
||||
NewMCPServerRequest.model_validate({
|
||||
"url": "https://upstream.example/sse",
|
||||
"mcp_info": {"protocol_version": "2026-07-28"},
|
||||
**({"transport": "sse"} if explicit_transport else {}),
|
||||
})
|
||||
|
||||
|
||||
@pytest.mark.parametrize("metadata", [False, True])
|
||||
def test_modern_sse_runtime_configuration_is_rejected(metadata: bool) -> None:
|
||||
with pytest.raises(ValidationError, match="Modern MCP requires HTTP or stdio"):
|
||||
MCPServer.model_validate({
|
||||
"server_id": "modern", "name": "modern", "transport": "sse",
|
||||
**({"mcp_info": {"protocol_version": "2026-07-28"}} if metadata else {"protocol_version": "2026-07-28"}),
|
||||
})
|
||||
|
||||
|
||||
@pytest.mark.parametrize("transport,version", [("http", "2026-07-28"), ("stdio", "2026-07-28"), ("sse", "2025-11-25"), ("sse", "auto")])
|
||||
def test_supported_protocol_transport_configurations_remain_valid(transport: str, version: str, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", "true")
|
||||
payload: Final = NewMCPServerRequest.model_validate({
|
||||
"transport": transport, "url": "https://upstream.example/mcp", "command": "python", "args": ["peer.py"],
|
||||
"mcp_info": {"protocol_version": version},
|
||||
})
|
||||
assert payload.transport == transport
|
||||
assert payload.mcp_info == {"protocol_version": version}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("protocol_only", [False, True])
|
||||
async def test_modern_sse_partial_update_rejected_without_writes(protocol_only: bool) -> None:
|
||||
old_record: Final = LiteLLM_MCPServerTable(
|
||||
server_id="srv-1", transport="sse" if protocol_only else "http",
|
||||
mcp_info={"protocol_version": "auto" if protocol_only else "2026-07-28"},
|
||||
)
|
||||
payload: Final = UpdateMCPServerRequest.model_validate({
|
||||
"server_id": "srv-1",
|
||||
**({"mcp_info": {"protocol_version": "2026-07-28"}} if protocol_only else {"transport": "sse", "url": "https://upstream.example/sse"}),
|
||||
})
|
||||
update_mock: Final = AsyncMock(side_effect=HTTPException(status_code=418, detail="Unexpected persistence"))
|
||||
p1, p2, p3, p4, p5 = _edit_endpoint_patches(old_record, update_mock)
|
||||
with p1, p2, p3, p4, p5:
|
||||
with pytest.raises(HTTPException) as error:
|
||||
await mgmt_endpoints.edit_mcp_server(payload=payload, user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN))
|
||||
assert error.value.status_code == 400
|
||||
assert "Modern MCP requires HTTP or stdio" in str(error.value.detail)
|
||||
update_mock.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_protocol_partial_update_fails_closed_when_stored_configuration_is_unreadable() -> None:
|
||||
update_mock: Final = AsyncMock(side_effect=HTTPException(status_code=418, detail="Unexpected persistence"))
|
||||
p1, p2, p3, p4, p5 = _edit_endpoint_patches(RuntimeError("db unavailable"), update_mock)
|
||||
with p1, p2, p3, p4, p5:
|
||||
with pytest.raises(HTTPException) as error:
|
||||
await mgmt_endpoints.edit_mcp_server(
|
||||
payload=UpdateMCPServerRequest(server_id="srv-1", mcp_info={"protocol_version": "2026-07-28"}),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
|
||||
)
|
||||
assert error.value.status_code == 503
|
||||
update_mock.assert_not_awaited()
|
||||
|
||||
|
||||
def test_modern_sse_complete_update_is_rejected() -> None:
|
||||
with pytest.raises(ValidationError, match="Modern MCP requires HTTP or stdio"):
|
||||
UpdateMCPServerRequest(
|
||||
server_id="server", transport=MCPTransport.sse, url="https://upstream.example/sse",
|
||||
mcp_info={"protocol_version": "2026-07-28"},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_protocol_update_on_missing_server_preserves_not_found() -> None:
|
||||
update_mock: Final = AsyncMock(return_value=None)
|
||||
p1, p2, p3, p4, p5 = _edit_endpoint_patches(None, update_mock)
|
||||
with p1, p2, p3, p4, p5:
|
||||
with pytest.raises(HTTPException) as error:
|
||||
await mgmt_endpoints.edit_mcp_server(
|
||||
payload=UpdateMCPServerRequest(server_id="missing", mcp_info={"protocol_version": "2026-07-28"}),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
|
||||
)
|
||||
assert error.value.status_code == 404
|
||||
|
|
|
|||
|
|
@ -397,7 +397,7 @@ def test_mcp_advertised_versions_reject_unavailable_revisions(versions):
|
|||
ConfigGeneralSettings(mcp_advertised_versions=versions)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("revision", ["2026-07-28", "unknown", None])
|
||||
@pytest.mark.parametrize("revision", ["unknown", None])
|
||||
def test_mcp_metadata_rejects_unavailable_upstream_protocol(revision):
|
||||
from litellm.proxy._types import NewMCPServerRequest, UpdateMCPServerRequest
|
||||
|
||||
|
|
@ -466,3 +466,13 @@ def test_an_http_mcp_server_is_unaffected_by_the_stdio_flag(monkeypatch, request
|
|||
def test_a_non_mapping_mcp_server_payload_gets_a_validation_error(request_model):
|
||||
with pytest.raises(ValidationError, match="valid dictionary"):
|
||||
request_model.model_validate("not-a-server")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("request_model", MCP_SERVER_REQUESTS)
|
||||
def test_modern_http_upstream_protocol_is_available(request_model):
|
||||
parsed = request_model.model_validate({
|
||||
"server_id": "modern", "transport": "http", "url": "https://example.com/mcp",
|
||||
"mcp_info": {"protocol_version": "2026-07-28"},
|
||||
})
|
||||
assert parsed.mcp_info["protocol_version"] == "2026-07-28"
|
||||
assert parsed.transport == "http"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue