mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge cbc12d1e1c into 4ece6c9fb8
This commit is contained in:
commit
9068c2441d
26 changed files with 3374 additions and 352 deletions
|
|
@ -28,7 +28,7 @@ from contextlib import asynccontextmanager
|
|||
from dataclasses import dataclass, replace
|
||||
from functools import lru_cache
|
||||
from itertools import chain, groupby
|
||||
from types import MappingProxyType
|
||||
from types import EllipsisType, MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Generic, Literal, TypeAlias, TypedDict, TypeVar, cast
|
||||
from urllib.parse import ParseResult, urlparse
|
||||
|
||||
|
|
@ -257,6 +257,20 @@ _user_env_vars_cache: Final[dict[tuple[str, str], tuple[dict[str, str], float]]]
|
|||
_USER_ENV_VARS_CACHE_TTL: Final = 60 # seconds
|
||||
_USER_ENV_VARS_CACHE_MAX_SIZE: Final = 4096 # cap to prevent unbounded growth
|
||||
|
||||
_ListedToolsByCaller: TypeAlias = Mapping[str | None, Mapping[str, MCPTool]]
|
||||
_LISTED_TOOLS_CALLERS_PER_SERVER: Final = 256
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ListedToolsCaller:
|
||||
"""Request inputs that select which upstream catalog a caller was shown by tools/list."""
|
||||
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None
|
||||
mcp_auth_header: str | dict[str, str] | None = None
|
||||
raw_headers: Mapping[str, str] | None = None
|
||||
oauth2_headers: Mapping[str, str] | None = None
|
||||
|
||||
|
||||
# Auth types whose upstream OAuth endpoints (protected-resource + authorization-server metadata) the
|
||||
# gateway discovers from the upstream itself: interactive oauth2 and the two client-forwarded modes.
|
||||
# OBO/M2M endpoint discovery is decided separately via _obo_needs_endpoint_discovery. Shared by the
|
||||
|
|
@ -1170,6 +1184,57 @@ def _authorization_is_litellm_admission_credential(
|
|||
return bool(user_api_key_auth and user_api_key_auth.api_key and not admission_header)
|
||||
|
||||
|
||||
def _server_auth_header_for(
|
||||
server: MCPServer,
|
||||
mcp_server_auth_headers: Mapping[str, str | dict[str, str]] | None,
|
||||
mcp_auth_header: str | dict[str, str] | None,
|
||||
) -> str | dict[str, str] | None:
|
||||
"""Server-specific ``x-mcp-<alias>-authorization`` header, else the deprecated global one."""
|
||||
server_specific: Final = (
|
||||
lookup_mcp_server_auth_in_headers(
|
||||
mcp_server_auth_headers,
|
||||
alias=server.alias,
|
||||
server_name=server.server_name,
|
||||
access_groups=server.access_groups,
|
||||
)
|
||||
if mcp_server_auth_headers
|
||||
else None
|
||||
)
|
||||
return mcp_auth_header if server_specific is None else server_specific
|
||||
|
||||
|
||||
def listed_tools_caller_for(
|
||||
server: MCPServer,
|
||||
user_api_key_auth: UserAPIKeyAuth | None,
|
||||
mcp_auth_header: str | dict[str, str] | None,
|
||||
mcp_server_auth_headers: Mapping[str, str | dict[str, str]] | None,
|
||||
raw_headers: Mapping[str, str] | None,
|
||||
oauth2_headers: Mapping[str, str] | None,
|
||||
) -> ListedToolsCaller:
|
||||
"""The caller a tools/call must look its listed entry up under: the same inputs tools/list keyed by."""
|
||||
return ListedToolsCaller(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=_server_auth_header_for(server, mcp_server_auth_headers, mcp_auth_header),
|
||||
raw_headers=raw_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
)
|
||||
|
||||
|
||||
def _admission_identity(
|
||||
auth: UserAPIKeyAuth, raw_headers: Mapping[str, str] | None
|
||||
) -> tuple[str | None, str | None, str | None, str | None, str | None]:
|
||||
"""The admission identity the served catalog is shaped for: the hashed key, user, team and
|
||||
organization, plus the admission credential (``x-litellm-api-key``, else ``Authorization``) of a
|
||||
caller admitted with neither a key nor a user."""
|
||||
keyless: Final = auth.api_key is None and auth.user_id is None
|
||||
credential: Final = (
|
||||
_raw_header_value(raw_headers, "x-litellm-api-key") or _raw_header_value(raw_headers, "authorization")
|
||||
if keyless
|
||||
else None
|
||||
)
|
||||
return auth.api_key, auth.user_id, auth.team_id, auth.org_id, credential
|
||||
|
||||
|
||||
def _format_byok_openapi_auth_header(mcp_server: MCPServer, mcp_auth_header: str) -> str:
|
||||
"""Format a raw BYOK credential for OpenAPI tool ``Authorization`` injection.
|
||||
|
||||
|
|
@ -1295,6 +1360,16 @@ async def _resolve_byok_mcp_auth_header(
|
|||
return mcp_auth_header
|
||||
|
||||
|
||||
def _catalog_auth_header(
|
||||
mcp_auth_header: str | dict[str, str] | None,
|
||||
catalog_auth_header: str | dict[str, str] | None | EllipsisType,
|
||||
) -> str | dict[str, str] | None:
|
||||
"""The header the client supplied, which keys the caller's catalog slot on both tools/list and
|
||||
tools/call. A caller that already swapped a stored BYOK credential into ``mcp_auth_header`` passes
|
||||
the client's value explicitly, since the stored credential must never be read to find the slot."""
|
||||
return mcp_auth_header if catalog_auth_header is ... else catalog_auth_header
|
||||
|
||||
|
||||
def _client_forwarded_authorization_headers(
|
||||
mcp_server: MCPServer,
|
||||
oauth2_headers: dict[str, str] | None,
|
||||
|
|
@ -1923,6 +1998,8 @@ class MCPServerManager:
|
|||
"gmail_send_email": "zapier_mcp_server",
|
||||
}
|
||||
"""
|
||||
self._listed_tools_by_server_id: dict[str, _ListedToolsByCaller] = {} # mutable-ok: refreshed per tools/list
|
||||
self._listed_tools_generations: dict[str, int] = {} # mutable-ok: bumped per server save
|
||||
self._upstream_initialize_instructions_by_server_id: dict[str, str] = {}
|
||||
# Per-server monotonic timestamp of last upstream prefetch attempt (success,
|
||||
# empty result, or failure). Used to throttle re-probes for servers that do
|
||||
|
|
@ -3240,6 +3317,8 @@ class MCPServerManager:
|
|||
self._invalidate_server_definition_caches(mcp_server.server_id)
|
||||
self.registry[mcp_server.server_id] = new_server
|
||||
await self._maybe_register_openapi_tools(new_server)
|
||||
if new_server.spec_path:
|
||||
self._drop_listed_tools(mcp_server.server_id)
|
||||
self.prime_oauth_metadata_discovery(new_server)
|
||||
verbose_logger.debug("Added MCP Server: %s", new_server.name)
|
||||
|
||||
|
|
@ -3277,6 +3356,8 @@ class MCPServerManager:
|
|||
self._invalidate_server_definition_caches(mcp_server.server_id)
|
||||
self.registry[mcp_server.server_id] = new_server
|
||||
await self._maybe_register_openapi_tools(new_server)
|
||||
if new_server.spec_path:
|
||||
self._drop_listed_tools(mcp_server.server_id)
|
||||
self.prime_oauth_metadata_discovery(new_server)
|
||||
verbose_logger.debug("Updated MCP Server: %s", new_server.name)
|
||||
|
||||
|
|
@ -3748,25 +3829,14 @@ class MCPServerManager:
|
|||
verbose_logger.warning("MCP Server %s not found", server_id)
|
||||
return []
|
||||
|
||||
# Get server-specific auth header if available
|
||||
server_auth_header: str | dict[str, str] | None = None
|
||||
if mcp_server_auth_headers:
|
||||
server_auth_header = lookup_mcp_server_auth_in_headers(
|
||||
mcp_server_auth_headers,
|
||||
alias=server.alias,
|
||||
server_name=server.server_name,
|
||||
access_groups=server.access_groups,
|
||||
)
|
||||
|
||||
# Fall back to deprecated mcp_auth_header if no server-specific header found
|
||||
if server_auth_header is None:
|
||||
server_auth_header = mcp_auth_header
|
||||
server_auth_header: Final = _server_auth_header_for(server, mcp_server_auth_headers, mcp_auth_header)
|
||||
|
||||
try:
|
||||
tools: Final = await self._get_tools_from_server(
|
||||
server=server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
record_listing=True,
|
||||
)
|
||||
return tools
|
||||
except Exception as e:
|
||||
|
|
@ -3846,7 +3916,7 @@ class MCPServerManager:
|
|||
def _build_stdio_env(
|
||||
self,
|
||||
server: MCPServer,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
raw_headers: Mapping[str, str] | None = None,
|
||||
) -> dict[str, str] | None:
|
||||
"""Resolve stdio env values, supporting header-driven placeholders."""
|
||||
|
||||
|
|
@ -4141,13 +4211,20 @@ class MCPServerManager:
|
|||
oauth2_headers: dict[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
proxy_logging_obj: ProxyLogging | None = None,
|
||||
) -> list[MCPTool]:
|
||||
*,
|
||||
catalog_auth_header: str | dict[str, str] | None | EllipsisType = ...,
|
||||
record_listing: bool = False,
|
||||
) -> Sequence[MCPTool]:
|
||||
"""
|
||||
Helper method to get tools from a single MCP server with prefixed names.
|
||||
|
||||
Args:
|
||||
server (MCPServer): The server to query tools from
|
||||
mcp_auth_header: Optional auth header for MCP server
|
||||
catalog_auth_header: The header the client supplied, keying the caller's catalog slot;
|
||||
defaults to ``mcp_auth_header``
|
||||
record_listing: Record the served catalog into the caller's listed-tools slot; only a
|
||||
listing actually served to the caller sets it
|
||||
|
||||
Returns:
|
||||
List[MCPTool]: List of tools available on the server with prefixed names
|
||||
|
|
@ -4163,6 +4240,13 @@ class MCPServerManager:
|
|||
verbose_logger.info("_get_tools_from_server for %s...", server.name)
|
||||
|
||||
client = None
|
||||
listed_caller: Final = ListedToolsCaller(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=_catalog_auth_header(mcp_auth_header, catalog_auth_header),
|
||||
raw_headers=raw_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
)
|
||||
listed_generation: Final = self._listed_tools_generations.get(server.server_id, 0)
|
||||
|
||||
try:
|
||||
# Tool *listing* must not be blocked by missing per-user env vars —
|
||||
|
|
@ -4260,8 +4344,12 @@ class MCPServerManager:
|
|||
# applied (e.g. "test_petstore-getinventory"). Do NOT pass them
|
||||
# through _create_prefixed_tools — that would add the prefix a second
|
||||
# time producing "test_petstore-test_petstore-getinventory".
|
||||
unprefixed_tools: Final = guarded_openapi
|
||||
self._record_listed_tools(
|
||||
server, unprefixed_tools, listed_caller, listed_generation, record_listing=record_listing
|
||||
)
|
||||
if not add_prefix:
|
||||
return list(guarded_openapi)
|
||||
return unprefixed_tools
|
||||
return [t.model_copy(update={"name": registered_names[t.name]}) for t in guarded_openapi]
|
||||
else:
|
||||
tools = await self._fetch_tools_with_timeout(client, server.name)
|
||||
|
|
@ -4275,7 +4363,10 @@ class MCPServerManager:
|
|||
raw_headers=raw_headers,
|
||||
)
|
||||
prefixed_or_original_tools: Final = self._create_prefixed_tools(
|
||||
list(guarded_tools), server, add_prefix=add_prefix
|
||||
guarded_tools, server, add_prefix=add_prefix
|
||||
)
|
||||
self._record_listed_tools(
|
||||
server, guarded_tools, listed_caller, listed_generation, record_listing=record_listing
|
||||
)
|
||||
|
||||
return prefixed_or_original_tools
|
||||
|
|
@ -4326,8 +4417,85 @@ class MCPServerManager:
|
|||
)
|
||||
|
||||
self._invalidate_discovery_lists(server_id)
|
||||
self._drop_listed_tools(server_id)
|
||||
invalidate_oauth_metadata_cache(server_id)
|
||||
|
||||
def _drop_listed_tools(self, server_id: str) -> None:
|
||||
self._listed_tools_by_server_id.pop(server_id, None)
|
||||
self._listed_tools_generations[server_id] = self._listed_tools_generations.get(server_id, 0) + 1
|
||||
|
||||
def _listed_tools_identity(self, server: MCPServer, caller: ListedToolsCaller | None) -> str | None:
|
||||
"""Key the listed-tool cache by every request input that can change the served catalog.
|
||||
|
||||
The catalog is guardrail-shaped for the caller's admission identity (default-on guardrails,
|
||||
key or team selections and opt-outs), so every admitted caller gets its own slot, keyed by
|
||||
``_admission_identity``: the hashed key, user, team and organization, plus the admission
|
||||
credential of a caller admitted with neither a key nor a user (a team-only JWT). Forwarded
|
||||
headers, header-driven stdio env, the caller bearer on every server whose egress forwards it
|
||||
(``_consumes_caller_authorization``) or exchanges it as the OBO subject, and the
|
||||
server-specific auth header also reach upstream and split the slot further. Only unkeyed
|
||||
listings with none of those share the ``None`` slot.
|
||||
"""
|
||||
if caller is None:
|
||||
return None
|
||||
auth: Final = caller.user_api_key_auth
|
||||
forwarded: Final = self._forwarded_header_values(server, caller.raw_headers) or None
|
||||
header_env: Final = self._build_stdio_env(server, caller.raw_headers)
|
||||
stdio_env: Final = None if header_env == self._build_stdio_env(server) else header_env
|
||||
caller_bearer: Final = (
|
||||
self._extract_subject_token(caller.oauth2_headers, caller.raw_headers, auth)
|
||||
if _consumes_caller_authorization(server) or server.auth_type == MCPAuth.oauth2_token_exchange
|
||||
else None
|
||||
)
|
||||
identity: Final = None if auth is None else _admission_identity(auth, caller.raw_headers)
|
||||
if not (identity or caller.mcp_auth_header or forwarded or stdio_env or caller_bearer):
|
||||
return None
|
||||
material: Final = json.dumps(
|
||||
(identity, caller.mcp_auth_header, forwarded, stdio_env, caller_bearer),
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
)
|
||||
return hashlib.sha256(material.encode()).hexdigest()
|
||||
|
||||
@staticmethod
|
||||
def _forwarded_header_values(
|
||||
server: MCPServer, raw_headers: Mapping[str, str] | None
|
||||
) -> tuple[tuple[str, str], ...]:
|
||||
if not raw_headers or not server.extra_headers:
|
||||
return ()
|
||||
forwarded_names: Final = frozenset(name.lower() for name in server.extra_headers)
|
||||
return tuple(
|
||||
sorted((name.lower(), value) for name, value in raw_headers.items() if name.lower() in forwarded_names)
|
||||
)
|
||||
|
||||
def _record_listed_tools(
|
||||
self,
|
||||
server: MCPServer,
|
||||
tools: Sequence[MCPTool],
|
||||
caller: ListedToolsCaller | None,
|
||||
generation: int | None = None,
|
||||
*,
|
||||
record_listing: bool = True,
|
||||
) -> None:
|
||||
"""Store the catalog served to ``caller``. ``generation`` is the server's listed-tools generation
|
||||
read before the listing's upstream fetch; the record is skipped when it no longer matches."""
|
||||
if not record_listing:
|
||||
return
|
||||
if generation is not None and generation != self._listed_tools_generations.get(server.server_id, 0):
|
||||
return
|
||||
identity: Final = self._listed_tools_identity(server, caller)
|
||||
listing: Final = MappingProxyType({tool.name: tool for tool in tools})
|
||||
existing: Final = self._listed_tools_by_server_id.get(server.server_id, MappingProxyType({}))
|
||||
shared: Final = existing.get(None)
|
||||
callers: Final = tuple((key, value) for key, value in existing.items() if key not in (None, identity))
|
||||
evicted: Final = 0 if identity is None else max(len(callers) + 1 - _LISTED_TOOLS_CALLERS_PER_SERVER, 0)
|
||||
entries: Final = (
|
||||
*(() if shared is None else ((None, shared),)),
|
||||
*callers[evicted:],
|
||||
(identity, listing),
|
||||
)
|
||||
self._listed_tools_by_server_id[server.server_id] = MappingProxyType(dict(entries))
|
||||
|
||||
def _discovery_key(
|
||||
self,
|
||||
server: MCPServer,
|
||||
|
|
@ -4337,9 +4505,11 @@ class MCPServerManager:
|
|||
stdio_env: dict[str, str] | None,
|
||||
subject_token: str | None,
|
||||
credential_fingerprint: str | None = None,
|
||||
per_caller: bool = False,
|
||||
) -> _DiscoveryKey:
|
||||
per_user: Final = (
|
||||
server.requires_per_user_auth
|
||||
per_caller
|
||||
or server.requires_per_user_auth
|
||||
or self._references_per_user_env_var(server)
|
||||
or server.delegate_auth_to_upstream
|
||||
or server.auth_type in (MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag)
|
||||
|
|
@ -5247,7 +5417,12 @@ class MCPServerManager:
|
|||
{seen: kept for seen, kept in self._catalog_alert_signatures.items() if seen != key}
|
||||
)
|
||||
|
||||
def _create_prefixed_tools(self, tools: list[MCPTool], server: MCPServer, add_prefix: bool = True) -> list[MCPTool]:
|
||||
def _create_prefixed_tools(
|
||||
self,
|
||||
tools: Sequence[MCPTool],
|
||||
server: MCPServer,
|
||||
add_prefix: bool = True,
|
||||
) -> list[MCPTool]:
|
||||
"""
|
||||
Create prefixed tools and update tool mapping.
|
||||
|
||||
|
|
@ -5275,6 +5450,13 @@ class MCPServerManager:
|
|||
verbose_logger.info("Successfully fetched %s tools from server %s", len(prefixed_tools), server.name)
|
||||
return prefixed_tools
|
||||
|
||||
def get_listed_tool(self, server: MCPServer, name: str, caller: ListedToolsCaller | None = None) -> MCPTool | None:
|
||||
identity: Final = self._listed_tools_identity(server, caller)
|
||||
listed: Final = self._listed_tools_by_server_id.get(server.server_id, MappingProxyType({})).get(identity)
|
||||
if not listed:
|
||||
return None
|
||||
return listed.get(name)
|
||||
|
||||
def _create_prefixed_prompts(
|
||||
self, prompts: Sequence[Prompt], server: MCPServer, add_prefix: bool = True
|
||||
) -> list[Prompt]:
|
||||
|
|
@ -5508,6 +5690,7 @@ class MCPServerManager:
|
|||
raw_headers: dict[str, str] | None = None,
|
||||
litellm_logging_obj: "LiteLLMLoggingObj | None" = None,
|
||||
guardrail_context: Mapping[str, object] | None = None,
|
||||
tool: MCPTool | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Run pre-call checks and guardrail hooks for an MCP tool call.
|
||||
|
|
@ -5521,6 +5704,9 @@ class MCPServerManager:
|
|||
``pre_mcp_call`` evaluation (or a block) on the spend-log row the Guardrails
|
||||
Monitor counts. It stays optional so callers that do no logging are unchanged.
|
||||
|
||||
``tool`` is the upstream tool definition when one was listed, so guardrails
|
||||
can see its description and input schema, not just the name and arguments.
|
||||
|
||||
Returns a dict that may contain:
|
||||
- "arguments": hook-modified tool arguments (only if changed)
|
||||
- "extra_headers": headers injected by pre_mcp_call guardrail hooks
|
||||
|
|
@ -5584,6 +5770,8 @@ class MCPServerManager:
|
|||
"user_api_key_hash": (getattr(user_api_key_auth, "api_key_hash", None) if user_api_key_auth else None),
|
||||
"incoming_bearer_token": incoming_bearer_token,
|
||||
"headers": logging_safe_mcp_headers(raw_headers),
|
||||
"tool_description": tool.description if tool is not None else None,
|
||||
"tool_input_schema": tool.input_schema if tool is not None else None,
|
||||
}
|
||||
|
||||
# Create MCP request object for processing
|
||||
|
|
@ -5799,21 +5987,7 @@ class MCPServerManager:
|
|||
GuardrailRaisedException: If guardrails block the call
|
||||
HTTPException: If an HTTP error occurs
|
||||
"""
|
||||
# Get server-specific auth header if available (case-insensitive)
|
||||
# FIX: Added case-insensitive matching to handle auth header keys that may not match
|
||||
# the exact case of server alias/name (e.g., '1litellmagcgateway' vs '1LiteLLMAGCGateway')
|
||||
server_auth_header: dict[str, str] | str | None = None
|
||||
if mcp_server_auth_headers:
|
||||
server_auth_header = lookup_mcp_server_auth_in_headers(
|
||||
mcp_server_auth_headers,
|
||||
alias=mcp_server.alias,
|
||||
server_name=mcp_server.server_name,
|
||||
access_groups=mcp_server.access_groups,
|
||||
)
|
||||
|
||||
# Fall back to deprecated mcp_auth_header if no server-specific header found
|
||||
if server_auth_header is None:
|
||||
server_auth_header = mcp_auth_header
|
||||
server_auth_header: Final = _server_auth_header_for(mcp_server, mcp_server_auth_headers, mcp_auth_header)
|
||||
|
||||
# Extract subject token for OAuth2 Token Exchange (OBO) and ID-JAG flows
|
||||
subject_token: str | None = None
|
||||
|
|
@ -6236,6 +6410,9 @@ class MCPServerManager:
|
|||
guardrail_context: Mapping[str, object] | None = None,
|
||||
client_ip: str | None = None,
|
||||
wire_compat: WireCompat = WireCompat.LEGACY,
|
||||
*,
|
||||
catalog_auth_header: str | None | EllipsisType = ...,
|
||||
listed_tool: MCPTool | None | EllipsisType = ...,
|
||||
) -> CallToolResult | InputRequiredResult:
|
||||
"""
|
||||
Call a tool with the given name and arguments
|
||||
|
|
@ -6247,6 +6424,8 @@ class MCPServerManager:
|
|||
user_api_key_auth: User authentication
|
||||
mcp_auth_header: MCP auth header (deprecated)
|
||||
mcp_server_auth_headers: Optional dict of server-specific auth headers {server_alias: auth_value}
|
||||
catalog_auth_header: The header the client supplied, keying the caller's catalog slot;
|
||||
defaults to ``mcp_auth_header`` as received, before BYOK resolution
|
||||
proxy_logging_obj: Optional ProxyLogging object for hook integration
|
||||
litellm_logging_obj: Optional request logger the guardrail hooks record
|
||||
their evaluations onto, so MCP guardrail activity reaches the
|
||||
|
|
@ -6258,6 +6437,7 @@ class MCPServerManager:
|
|||
"""
|
||||
start_time: Final = datetime.datetime.now()
|
||||
mcp_server: Final = self._resolve_mcp_server_for_tool_call(server_name, name)
|
||||
client_auth_header: Final = _catalog_auth_header(mcp_auth_header, catalog_auth_header)
|
||||
|
||||
# Resolved before any hook runs so a missing BYOK credential (401) never
|
||||
# leaves during-hook side effects (audit logging, rate-limit bookkeeping)
|
||||
|
|
@ -6267,6 +6447,9 @@ class MCPServerManager:
|
|||
user_api_key_auth,
|
||||
mcp_auth_header,
|
||||
)
|
||||
listed_caller: Final = listed_tools_caller_for(
|
||||
mcp_server, user_api_key_auth, client_auth_header, mcp_server_auth_headers, raw_headers, oauth2_headers
|
||||
)
|
||||
|
||||
#########################################################
|
||||
# Pre MCP Tool Call Hook
|
||||
|
|
@ -6283,6 +6466,7 @@ class MCPServerManager:
|
|||
raw_headers=raw_headers,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
guardrail_context=guardrail_context,
|
||||
tool=self.get_listed_tool(mcp_server, name, listed_caller) if listed_tool is ... else listed_tool,
|
||||
)
|
||||
if "arguments" in hook_result:
|
||||
arguments = hook_result["arguments"]
|
||||
|
|
|
|||
|
|
@ -81,12 +81,14 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
|
|||
outcome_wire_value,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
ListedToolsCaller,
|
||||
MCPServerManager,
|
||||
_caller_authorization_fans_out,
|
||||
_client_forwarded_authorization_headers,
|
||||
_resolve_openapi_tool_auth,
|
||||
_should_strip_caller_authorization,
|
||||
global_mcp_server_manager,
|
||||
listed_tools_caller_for,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import (
|
||||
_redact_mcp_resource_url,
|
||||
|
|
@ -954,6 +956,8 @@ async def _get_tools_from_mcp_servers(
|
|||
request_tags: list[str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
mcp_proxy_mode: bool = False,
|
||||
*,
|
||||
record_listing: bool = False,
|
||||
) -> AggregateToolListing:
|
||||
"""
|
||||
Helper method to fetch tools from MCP servers based on server filtering criteria.
|
||||
|
|
@ -964,6 +968,8 @@ async def _get_tools_from_mcp_servers(
|
|||
mcp_servers: Optional list of server names/aliases to filter by
|
||||
mcp_server_auth_headers: Optional dict of server-specific auth headers
|
||||
oauth2_headers: Optional dict of oauth2 headers
|
||||
record_listing: Record each served catalog into the caller's listed-tools slot; only a
|
||||
listing actually served to the caller sets it
|
||||
|
||||
Returns:
|
||||
AggregateToolListing: Combined tools from filtered servers plus each server's
|
||||
|
|
@ -1111,12 +1117,14 @@ async def _get_tools_from_mcp_servers(
|
|||
prefetched_creds=_prefetched_oauth_creds,
|
||||
)
|
||||
|
||||
catalog_auth_header: Final = server_auth_header
|
||||
if server.is_byok and server.auth_type != MCPAuth.oauth2 and server_auth_header is None:
|
||||
server_auth_header = await _get_byok_credential(server, user_api_key_auth)
|
||||
|
||||
try:
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
|
||||
listed_generation: Final = global_mcp_server_manager._listed_tools_generations.get(server.server_id, 0)
|
||||
tools: Final = await global_mcp_server_manager._get_tools_from_server(
|
||||
server=server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
|
|
@ -1127,6 +1135,8 @@ async def _get_tools_from_mcp_servers(
|
|||
user_api_key_auth=user_api_key_auth,
|
||||
oauth2_headers=oauth2_headers,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
catalog_auth_header=catalog_auth_header,
|
||||
record_listing=False,
|
||||
)
|
||||
filtered_tools = filter_tools_by_allowed_tools(tools, server)
|
||||
|
||||
|
|
@ -1135,6 +1145,21 @@ async def _get_tools_from_mcp_servers(
|
|||
server_id=server.server_id,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
global_mcp_server_manager._record_listed_tools(
|
||||
server,
|
||||
[
|
||||
tool.model_copy(update={"name": strip_known_server_prefix(tool.name, server)})
|
||||
for tool in filtered_tools
|
||||
],
|
||||
ListedToolsCaller(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=catalog_auth_header,
|
||||
raw_headers=raw_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
),
|
||||
listed_generation,
|
||||
record_listing=record_listing,
|
||||
)
|
||||
|
||||
if mcp_proxy_mode:
|
||||
from litellm.proxy._experimental.mcp_server.tool_search import with_mcp_proxy_identity
|
||||
|
|
@ -1458,6 +1483,8 @@ async def _list_mcp_tools(
|
|||
list_tools_log_source: str | None = None,
|
||||
client_ip: str | None = None,
|
||||
mcp_proxy_mode: bool = False,
|
||||
*,
|
||||
record_listing: bool = False,
|
||||
) -> AggregateToolListing:
|
||||
"""
|
||||
List all available MCP tools.
|
||||
|
|
@ -1468,6 +1495,8 @@ async def _list_mcp_tools(
|
|||
mcp_servers: Optional list of server names/aliases to filter by
|
||||
mcp_server_auth_headers: Optional dict of server-specific auth headers {server_alias: auth_value}
|
||||
client_ip: Client IP for IP-based server access control
|
||||
record_listing: Record each served catalog into the caller's listed-tools slot; only a
|
||||
listing actually served to the caller sets it
|
||||
|
||||
Returns:
|
||||
AggregateToolListing: Combined tools from all accessible servers plus each server's
|
||||
|
|
@ -1486,6 +1515,7 @@ async def _list_mcp_tools(
|
|||
list_tools_log_source=list_tools_log_source,
|
||||
client_ip=client_ip,
|
||||
mcp_proxy_mode=mcp_proxy_mode,
|
||||
record_listing=record_listing,
|
||||
)
|
||||
verbose_logger.debug("Successfully fetched %s tools from managed MCP servers", len(listing.tools))
|
||||
return listing
|
||||
|
|
@ -1798,6 +1828,7 @@ async def _list_tools_before_first_call(
|
|||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
client_ip=client_ip,
|
||||
record_listing=False,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # best effort: resolution below answers as it did before
|
||||
verbose_logger.debug("MCP tools/call: listing %s before its first call failed: %s", server.name, e)
|
||||
|
|
@ -2001,6 +2032,7 @@ async def _execute_mcp_tool(
|
|||
if mcp_server is None:
|
||||
mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name)
|
||||
|
||||
client_auth_header: Final = mcp_auth_header
|
||||
if mcp_server:
|
||||
standard_logging_mcp_tool_call["mcp_server_cost_info"] = (mcp_server.mcp_info or {}).get("mcp_server_cost_info")
|
||||
if litellm_logging_obj:
|
||||
|
|
@ -2071,6 +2103,18 @@ async def _execute_mcp_tool(
|
|||
raw_headers=raw_headers,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
guardrail_context=guardrail_context,
|
||||
tool=global_mcp_server_manager.get_listed_tool(
|
||||
mcp_server,
|
||||
original_tool_name,
|
||||
listed_tools_caller_for(
|
||||
mcp_server,
|
||||
user_api_key_auth,
|
||||
client_auth_header,
|
||||
mcp_server_auth_headers,
|
||||
raw_headers,
|
||||
oauth2_headers,
|
||||
),
|
||||
),
|
||||
)
|
||||
# `pre_call_tool_check` may return guardrail-modified
|
||||
# arguments; honor them on the local path too.
|
||||
|
|
@ -2120,6 +2164,7 @@ async def _execute_mcp_tool(
|
|||
arguments=arguments,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
catalog_auth_header=client_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
|
|
@ -2139,7 +2184,8 @@ async def _execute_mcp_tool(
|
|||
# not in the registry either, `_handle_local_mcp_tool` below reports
|
||||
# 404 and nothing runs, so demanding a server here would turn every
|
||||
# unknown tool name into a misleading 503.
|
||||
if global_mcp_tool_registry.get_tool(original_tool_name) is not None:
|
||||
registered_local_tool: Final = global_mcp_tool_registry.get_tool(original_tool_name)
|
||||
if registered_local_tool is not None:
|
||||
# `mcp_server` is None here because the tool name is not in the
|
||||
# tool -> server mapping, but the name still carries a prefix
|
||||
# that the server-level check above compared against the
|
||||
|
|
@ -2181,6 +2227,18 @@ async def _execute_mcp_tool(
|
|||
raw_headers=raw_headers,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
guardrail_context=guardrail_context,
|
||||
tool=global_mcp_server_manager.get_listed_tool(
|
||||
prefix_server,
|
||||
original_tool_name,
|
||||
listed_tools_caller_for(
|
||||
prefix_server,
|
||||
user_api_key_auth,
|
||||
client_auth_header,
|
||||
mcp_server_auth_headers,
|
||||
raw_headers,
|
||||
oauth2_headers,
|
||||
),
|
||||
),
|
||||
)
|
||||
if "arguments" in hook_result:
|
||||
arguments = hook_result["arguments"] # pyright: ignore[reportAny] # hook returns untyped args
|
||||
|
|
@ -2583,8 +2641,12 @@ async def _handle_managed_mcp_tool(
|
|||
guardrail_context: Mapping[str, object] | None = None,
|
||||
client_ip: str | None = None,
|
||||
wire_compat: WireCompat = WireCompat.LEGACY,
|
||||
*,
|
||||
catalog_auth_header: str | None,
|
||||
) -> CallToolResult | InputRequiredResult:
|
||||
"""Handle tool execution for managed server tools"""
|
||||
"""Handle tool execution for managed server tools. ``catalog_auth_header`` is the header the client
|
||||
supplied, which keys the caller's catalog slot; ``mcp_auth_header`` may already be the resolved
|
||||
BYOK credential."""
|
||||
# Import here to avoid circular import
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
|
||||
|
|
@ -2594,6 +2656,7 @@ async def _handle_managed_mcp_tool(
|
|||
arguments=arguments,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
catalog_auth_header=catalog_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
|
|
@ -2712,6 +2775,7 @@ async def _execute_handle_list_tools(
|
|||
log_list_tools_to_spendlogs=log_list_tools_to_spendlogs,
|
||||
list_tools_log_source="mcp_protocol",
|
||||
client_ip=_client_ip,
|
||||
record_listing=True,
|
||||
)
|
||||
verbose_logger.info("MCP list_tools - Successfully returned %s tools", len(listing.tools))
|
||||
if not listing.outcomes:
|
||||
|
|
|
|||
|
|
@ -223,6 +223,7 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
_UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES,
|
||||
ListedToolsCaller,
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import (
|
||||
|
|
@ -703,6 +704,8 @@ if MCP_AVAILABLE:
|
|||
extra_headers: dict[str, str] | None,
|
||||
client_ip: str | None,
|
||||
proxy_logging_obj: "ProxyLogging | None",
|
||||
*,
|
||||
record_listing: bool,
|
||||
) -> list[MCPTool]:
|
||||
return await global_mcp_server_manager._get_tools_from_server(
|
||||
server=server,
|
||||
|
|
@ -713,11 +716,12 @@ if MCP_AVAILABLE:
|
|||
client_ip=client_ip,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
record_listing=record_listing,
|
||||
)
|
||||
|
||||
async def _get_tools_for_single_server(
|
||||
server,
|
||||
server_auth_header,
|
||||
server: MCPServer,
|
||||
server_auth_header: dict[str, str] | str | None,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
|
|
@ -733,31 +737,40 @@ if MCP_AVAILABLE:
|
|||
"""
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
|
||||
tools = await _list_server_tools(
|
||||
server, server_auth_header, raw_headers, user_api_key_auth, extra_headers, client_ip, proxy_logging_obj
|
||||
listed_generation: Final = global_mcp_server_manager._listed_tools_generations.get(server.server_id, 0)
|
||||
tools: Final = await _list_server_tools(
|
||||
server,
|
||||
server_auth_header,
|
||||
raw_headers,
|
||||
user_api_key_auth,
|
||||
extra_headers,
|
||||
client_ip,
|
||||
proxy_logging_obj,
|
||||
record_listing=False,
|
||||
)
|
||||
|
||||
if not apply_tool_filters:
|
||||
return _create_tool_response_objects(tools, server)
|
||||
|
||||
# Always apply allowed_tools/disallowed_tools so the blacklist is
|
||||
# enforced even when no allowlist is set (matches the SSE/HTTP path).
|
||||
tools = filter_tools_by_allowed_tools(tools, server)
|
||||
|
||||
# Filter by the key's effective tool permissions through the same
|
||||
# function the MCP protocol path uses (direct grants, toolset grants,
|
||||
# and team/agent/org ceilings), so REST listing cannot drift from it.
|
||||
# Entries here are tool names on one server, written bare by every
|
||||
# writer, and dispatch compares them bare; matching a wider set of
|
||||
# spellings would advertise a tool that tools/call then refuses
|
||||
if user_api_key_auth:
|
||||
tools = await filter_tools_by_key_team_permissions(
|
||||
tools=tools,
|
||||
server_filtered: Final = filter_tools_by_allowed_tools(tools, server) if apply_tool_filters else tools
|
||||
served_tools: Final = (
|
||||
await filter_tools_by_key_team_permissions(
|
||||
tools=server_filtered,
|
||||
server_id=server.server_id,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
if apply_tool_filters and user_api_key_auth
|
||||
else server_filtered
|
||||
)
|
||||
global_mcp_server_manager._record_listed_tools(
|
||||
server,
|
||||
served_tools,
|
||||
ListedToolsCaller(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=server_auth_header,
|
||||
raw_headers=raw_headers,
|
||||
),
|
||||
listed_generation,
|
||||
)
|
||||
|
||||
return _create_tool_response_objects(tools, server)
|
||||
return _create_tool_response_objects(served_tools, server)
|
||||
|
||||
async def fetch_pinnable_tool_catalog(
|
||||
server: MCPServer, request: Request, user_api_key_dict: UserAPIKeyAuth
|
||||
|
|
@ -776,6 +789,7 @@ if MCP_AVAILABLE:
|
|||
await _get_user_oauth_extra_headers(server, user_api_key_dict),
|
||||
IPAddressUtils.get_mcp_client_ip(request),
|
||||
None,
|
||||
record_listing=False,
|
||||
)
|
||||
scan: Final = await scan_tool_descriptions(
|
||||
apply_description_overrides(upstream, server), server, proxy_logging_obj, user_api_key_dict, raw_headers
|
||||
|
|
|
|||
|
|
@ -19,7 +19,7 @@ from typing import TYPE_CHECKING, ClassVar, Final, Literal, NoReturn
|
|||
|
||||
import httpx
|
||||
from fastapi import HTTPException
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -61,6 +61,7 @@ _GATEWAY_OWNED_TOKEN_ERRORS: Final = frozenset(
|
|||
_INVALID_ASSERTION_AADSTS_PREFIX: Final = "50027"
|
||||
_AADSTS_CODES_ADAPTER: Final = TypeAdapter(tuple[int, ...])
|
||||
_MCP_CALL_TYPES: Final[tuple[str, ...]] = ("mcp_call", "call_mcp_tool")
|
||||
_TOOL_INPUT_SCHEMA_ADAPTER: Final = TypeAdapter(dict[str, object])
|
||||
_OBO_CACHE_MAX_ENTRIES: Final = 1000
|
||||
_DEFAULT_TOKEN_TTL_SECONDS: Final = 3599.0
|
||||
_TOKEN_EXPIRY_SLACK_SECONDS: Final = 60.0
|
||||
|
|
@ -82,6 +83,13 @@ def _parse_aadsts_codes(raw: object) -> tuple[int, ...]:
|
|||
return ()
|
||||
|
||||
|
||||
def _parse_tool_input_schema(raw: object) -> Mapping[str, object] | None:
|
||||
try:
|
||||
return _TOOL_INPUT_SCHEMA_ADAPTER.validate_python(raw)
|
||||
except ValidationError:
|
||||
return None
|
||||
|
||||
|
||||
def entra_assertion(value: object) -> str | None:
|
||||
"""``value`` when it is a compact JWS, the only bearer shape the OBO exchange accepts as its assertion.
|
||||
A LiteLLM virtual key, session bearer, or opaque upstream token in ``Authorization`` yields ``None``."""
|
||||
|
|
@ -100,6 +108,14 @@ class _EvaluateResponse(TypedDict, total=False):
|
|||
correlationId: ReadOnly[str]
|
||||
|
||||
|
||||
class _ToolReference(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
name: str
|
||||
description: str | None = None
|
||||
input_schema: Mapping[str, object] | None = Field(default=None, serialization_alias="inputSchema")
|
||||
|
||||
|
||||
class _UnavailableDetail(TypedDict):
|
||||
error: ReadOnly[str]
|
||||
message: ReadOnly[str]
|
||||
|
|
@ -392,8 +408,14 @@ class Agent365Guardrail(CustomGuardrail):
|
|||
arguments: Final = data.get("mcp_arguments")
|
||||
server_name: Final = str(data.get("mcp_server_name") or "litellm")
|
||||
agent_id: Final = user_api_key_dict.key_alias
|
||||
description: Final = data.get("mcp_tool_description")
|
||||
tool_reference: Final = _ToolReference(
|
||||
name=tool_name,
|
||||
description=description if isinstance(description, str) and description else None,
|
||||
input_schema=_parse_tool_input_schema(data.get("mcp_input_schema")),
|
||||
)
|
||||
payload: Final[dict[str, object]] = { # mutable-ok: JSON body with optional fields added below
|
||||
"tool": {"name": tool_name},
|
||||
"tool": tool_reference.model_dump(by_alias=True, exclude_none=True),
|
||||
"serverName": server_name,
|
||||
"conversationId": self._resolve_conversation_id(data),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1014,6 +1014,7 @@ if MCP_AVAILABLE:
|
|||
mcp_auth_header=None,
|
||||
mcp_servers=None,
|
||||
mcp_server_auth_headers=None,
|
||||
record_listing=True,
|
||||
)
|
||||
tools: Final = listing.tools
|
||||
dumped_tools: Final = [tool.model_dump(by_alias=True) for tool in tools]
|
||||
|
|
|
|||
|
|
@ -1148,6 +1148,8 @@ def _overrides_moderation_hook(callback: CustomLogger) -> bool:
|
|||
|
||||
|
||||
_LISTED_MODEL_NAMES: Final = TypeAdapter(tuple[str, ...])
|
||||
_MCP_TOOL_DESCRIPTION: Final[TypeAdapter[str | None]] = TypeAdapter(str | None)
|
||||
_MCP_TOOL_INPUT_SCHEMA: Final[TypeAdapter[Mapping[str, object] | None]] = TypeAdapter(Mapping[str, object] | None)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -1472,9 +1474,14 @@ class ProxyLogging:
|
|||
TypeAdapter(dict[str, object]).validate_python(guardrail_context.get("metadata") or MappingProxyType({}))
|
||||
)
|
||||
|
||||
mcp_tool_description: Final = kwargs.get("mcp_tool_description")
|
||||
mcp_input_schema: Final = kwargs.get("mcp_input_schema")
|
||||
description_line: Final = f"\nDescription: {mcp_tool_description}" if mcp_tool_description else ""
|
||||
mcp_tool_description: Final = request_obj.tool_description or kwargs.get("mcp_tool_description")
|
||||
mcp_input_schema: Final = (
|
||||
request_obj.tool_input_schema
|
||||
if request_obj.tool_input_schema is not None
|
||||
else kwargs.get("mcp_input_schema")
|
||||
)
|
||||
listing_description: Final = kwargs.get("mcp_tool_description")
|
||||
description_line: Final = f"\nDescription: {listing_description}" if listing_description else ""
|
||||
tool_call_content: Final = (
|
||||
f"Tool: {request_obj.tool_name}{description_line}\nArguments: {request_obj.arguments}"
|
||||
)
|
||||
|
|
@ -1732,6 +1739,8 @@ class ProxyLogging:
|
|||
tool_name=kwargs.get("name", ""),
|
||||
arguments=kwargs.get("arguments", {}),
|
||||
server_name=kwargs.get("server_name"),
|
||||
tool_description=_MCP_TOOL_DESCRIPTION.validate_python(kwargs.get("tool_description")),
|
||||
tool_input_schema=_MCP_TOOL_INPUT_SCHEMA.validate_python(kwargs.get("tool_input_schema")),
|
||||
user_api_key_auth=user_api_key_auth_dict,
|
||||
hidden_params=HiddenParams(),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -286,6 +286,7 @@ async def aresponses_api_with_mcp(
|
|||
call_params=call_params,
|
||||
previous_response_id=previous_response_id,
|
||||
tool_server_map=tool_server_map,
|
||||
served_tools=original_mcp_tools,
|
||||
**kwargs,
|
||||
)
|
||||
await mcp_streaming_response._create_initial_response_iterator()
|
||||
|
|
@ -339,6 +340,7 @@ async def aresponses_api_with_mcp(
|
|||
|
||||
tool_results: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
|
||||
tool_server_map=tool_server_map,
|
||||
served_tools=original_mcp_tools,
|
||||
tool_calls=tool_calls,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
|
|
@ -395,6 +397,7 @@ async def aresponses_api_with_mcp(
|
|||
|
||||
final_response = MCPEnhancedStreamingIterator(
|
||||
tool_server_map=tool_server_map,
|
||||
served_tools=original_mcp_tools,
|
||||
base_iterator=final_response,
|
||||
mcp_events=tool_execution_events,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
|
|
|
|||
|
|
@ -435,6 +435,7 @@ async def acompletion_with_mcp(
|
|||
# Execute tool calls
|
||||
self.tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
|
||||
tool_server_map=self.tool_server_map,
|
||||
served_tools=deduplicated_mcp_tools,
|
||||
tool_calls=self.tool_calls,
|
||||
user_api_key_auth=self.user_api_key_auth,
|
||||
mcp_auth_header=self.mcp_auth_header,
|
||||
|
|
@ -609,6 +610,7 @@ async def acompletion_with_mcp(
|
|||
tool_results: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
|
||||
tool_server_map=tool_server_map,
|
||||
tool_calls=tool_calls,
|
||||
served_tools=deduplicated_mcp_tools,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
|
|
|
|||
|
|
@ -696,6 +696,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
litellm_trace_id: str | None = None,
|
||||
request_tags: list[str] | None = None,
|
||||
guardrail_context: Mapping[str, object] | None = None,
|
||||
served_tools: Sequence[MCPTool] | None = None,
|
||||
) -> list[MCPToolResult]:
|
||||
"""Execute tool calls and return results."""
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -860,6 +861,11 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
guardrail_context=guardrail_context,
|
||||
listed_tool=(
|
||||
next((tool for tool in served_tools if tool.name == tool_name), None)
|
||||
if served_tools is not None
|
||||
else ...
|
||||
),
|
||||
)
|
||||
|
||||
if proxy_logging_obj:
|
||||
|
|
@ -1152,6 +1158,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
call_params: Mapping[str, object],
|
||||
previous_response_id: str | None,
|
||||
tool_server_map: dict[str, str],
|
||||
served_tools: Sequence[MCPTool] | None = None,
|
||||
**kwargs,
|
||||
) -> Any:
|
||||
"""
|
||||
|
|
@ -1181,6 +1188,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
base_iterator=None, # Will be created internally
|
||||
mcp_events=mcp_discovery_events, # Pre-generated MCP discovery events
|
||||
tool_server_map=tool_server_map,
|
||||
served_tools=served_tools,
|
||||
mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy,
|
||||
user_api_key_auth=kwargs.get("user_api_key_auth")
|
||||
or kwargs.get("litellm_metadata", {}).get("user_api_key_auth"),
|
||||
|
|
|
|||
|
|
@ -281,6 +281,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
mcp_tools_with_litellm_proxy: Sequence[Mapping[str, object]] | None = None,
|
||||
user_api_key_auth: "UserAPIKeyAuth | None" = None,
|
||||
original_request_params: dict[str, Any] | None = None,
|
||||
served_tools: Sequence[MCPTool] | None = None,
|
||||
):
|
||||
# MCP setup
|
||||
self.mcp_tools_with_litellm_proxy = mcp_tools_with_litellm_proxy or []
|
||||
|
|
@ -300,6 +301,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
self.mcp_discovery_generated = True # Events are already generated
|
||||
self.mcp_events = mcp_events # Store the initial MCP events for backward compatibility
|
||||
self.tool_server_map = tool_server_map
|
||||
self.served_tools = tuple(served_tools) if served_tools is not None else None
|
||||
|
||||
# Iterator references
|
||||
self.base_iterator: BaseResponsesAPIStreamingIterator | ResponsesAPIResponse | None = (
|
||||
|
|
@ -796,6 +798,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
# Execute the tools
|
||||
tool_results: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
|
||||
tool_server_map=self.tool_server_map,
|
||||
served_tools=self.served_tools,
|
||||
tool_calls=tool_calls,
|
||||
user_api_key_auth=self.user_api_key_auth,
|
||||
mcp_auth_header=self.mcp_auth_header,
|
||||
|
|
|
|||
|
|
@ -429,6 +429,8 @@ class MCPPreCallRequestObject(BaseModel):
|
|||
tool_name: str
|
||||
arguments: dict[str, Any]
|
||||
server_name: str | None = None
|
||||
tool_description: str | None = None
|
||||
tool_input_schema: Mapping[str, object] | None = None
|
||||
user_api_key_auth: dict[str, Any] | None = None
|
||||
hidden_params: HiddenParams = HiddenParams()
|
||||
|
||||
|
|
@ -452,6 +454,8 @@ class MCPDuringCallRequestObject(BaseModel):
|
|||
tool_name: str
|
||||
arguments: dict[str, Any]
|
||||
server_name: str | None = None
|
||||
tool_description: str | None = None
|
||||
tool_input_schema: Mapping[str, object] | None = None
|
||||
start_time: float | None = None
|
||||
hidden_params: HiddenParams = HiddenParams()
|
||||
|
||||
|
|
|
|||
|
|
@ -216,6 +216,16 @@ JsonRpc = Mapping[str, object]
|
|||
class ScriptedTool:
|
||||
name: str
|
||||
respond: Callable[[JsonRpc], Reply | JsonRpc]
|
||||
description: str | Callable[[Mapping[str, str]], str] | None = None
|
||||
input_schema: JsonRpc = field(default_factory=lambda: {"type": "object"})
|
||||
|
||||
def listing(self, headers: Mapping[str, str]) -> JsonRpc:
|
||||
described: Final = self.description(headers) if callable(self.description) else self.description
|
||||
return {
|
||||
"name": self.name,
|
||||
"inputSchema": self.input_schema,
|
||||
**({} if described is None else {"description": described}),
|
||||
}
|
||||
|
||||
|
||||
def jsonrpc_reply(identity: object, result: JsonRpc) -> Reply:
|
||||
|
|
@ -253,9 +263,7 @@ def scripted_peer(*tools: ScriptedTool) -> Iterator[McpPeer]:
|
|||
},
|
||||
)
|
||||
if method == "tools/list":
|
||||
return jsonrpc_reply(
|
||||
identity, {"tools": [{"name": name, "inputSchema": {"type": "object"}} for name in by_name]}
|
||||
)
|
||||
return jsonrpc_reply(identity, {"tools": [tool.listing(request.headers) for tool in by_name.values()]})
|
||||
if method != "tools/call":
|
||||
return jsonrpc_error(identity, -32601, f"unsupported method {method}")
|
||||
tool: Final = by_name.get(body["params"]["name"])
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ from integration._support.mcp import (
|
|||
tool_calls,
|
||||
tool_names,
|
||||
)
|
||||
from integration._support.oauth_server import oauth_server
|
||||
|
||||
ADD: Final = {"a": 2, "b": 3}
|
||||
STATIC_MODES: Final = (
|
||||
|
|
@ -192,3 +193,111 @@ def test_byok_server_uses_the_calling_users_stored_credential_and_fails_closed_w
|
|||
assert removed.status_code in (200, 204), removed.text
|
||||
eventually(lambda: call_tool(gateway, owner_key, identity, name, ADD), lambda value: value.status_code == 401)
|
||||
assert tool_calls(peer.drain()) == ()
|
||||
|
||||
|
||||
def _listings(peer: McpPeer) -> tuple[dict[str, object], ...]:
|
||||
return tuple(
|
||||
item
|
||||
for item in peer.drain()
|
||||
if isinstance(item.get("body"), dict) and item["body"].get("method") == "tools/list"
|
||||
)
|
||||
|
||||
|
||||
def test_oauth2_byok_listing_sends_the_minted_token_not_the_users_stored_secret(gateway: Gateway) -> None:
|
||||
with mcp_peer() as peer, oauth_server() as auth, gateway.scenario() as scenario:
|
||||
alias: Final = "cc" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(
|
||||
scenario,
|
||||
peer,
|
||||
alias,
|
||||
auth_type="oauth2",
|
||||
oauth2_flow="client_credentials",
|
||||
is_byok=True,
|
||||
token_url=auth.issuer + "/token",
|
||||
credentials={"client_id": "cc-client", "client_secret": "cc-secret-" + uuid.uuid4().hex},
|
||||
)
|
||||
owner: Final = scenario.user()
|
||||
owner_key: Final = scenario.key(user_id=owner, object_permission={"mcp_servers": [identity]})
|
||||
secret: Final = "byok-" + uuid.uuid4().hex
|
||||
stored: Final = gateway.client.post(
|
||||
f"/v1/mcp/server/{identity}/user-credential",
|
||||
json={"credential": secret},
|
||||
headers={"x-litellm-api-key": owner_key},
|
||||
)
|
||||
assert stored.status_code in (200, 201), stored.text
|
||||
scenario.cleanups.callback(
|
||||
gateway.client.delete,
|
||||
f"/v1/mcp/server/{identity}/user-credential",
|
||||
headers={"x-litellm-api-key": owner_key},
|
||||
)
|
||||
peer.drain()
|
||||
auth.drain()
|
||||
response: Final = gateway.client.get(
|
||||
"/mcp-rest/tools/list", params={"server_id": identity}, headers={"x-litellm-api-key": owner_key}
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert "add" in {tool["name"] for tool in response.json()["tools"]}, response.text
|
||||
assert [request["grant_type"] for request in auth.token_requests()] == ["client_credentials"]
|
||||
listings: Final = _listings(peer)
|
||||
assert len(listings) == 1, listings
|
||||
sent: Final = _header(listings[0], b"authorization")
|
||||
assert sent is not None and auth.is_live(sent.decode().removeprefix("Bearer ")), sent
|
||||
assert secret.encode() not in sent, "stored BYOK secret replaced the minted token on tools/list"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("auth_type", "header", "shape"), STATIC_MODES[:2])
|
||||
def test_byok_rest_listing_sends_the_servers_static_credential_not_the_users_stored_secret(
|
||||
gateway: Gateway, auth_type: str, header: bytes, shape: str
|
||||
) -> None:
|
||||
with mcp_peer() as peer, gateway.scenario() as scenario:
|
||||
alias: Final = "byok" + uuid.uuid4().hex[:8]
|
||||
static: Final = "static-" + uuid.uuid4().hex
|
||||
identity: Final = register_mcp(
|
||||
scenario, peer, alias, auth_type=auth_type, is_byok=True, credentials={"auth_value": static}
|
||||
)
|
||||
owner: Final = scenario.user()
|
||||
owner_key: Final = scenario.key(user_id=owner, object_permission={"mcp_servers": [identity]})
|
||||
secret: Final = "byok-" + uuid.uuid4().hex
|
||||
stored: Final = gateway.client.post(
|
||||
f"/v1/mcp/server/{identity}/user-credential",
|
||||
json={"credential": secret},
|
||||
headers={"x-litellm-api-key": owner_key},
|
||||
)
|
||||
assert stored.status_code in (200, 201), stored.text
|
||||
scenario.cleanups.callback(
|
||||
gateway.client.delete,
|
||||
f"/v1/mcp/server/{identity}/user-credential",
|
||||
headers={"x-litellm-api-key": owner_key},
|
||||
)
|
||||
peer.drain()
|
||||
response: Final = gateway.client.get(
|
||||
"/mcp-rest/tools/list", params={"server_id": identity}, headers={"x-litellm-api-key": owner_key}
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert "add" in {tool["name"] for tool in response.json()["tools"]}, response.text
|
||||
listings: Final = _listings(peer)
|
||||
assert len(listings) == 1, listings
|
||||
assert _header(listings[0], header) == shape.format(secret=static, basic="").encode(), listings[0]["headers"]
|
||||
peer.drain()
|
||||
called: Final = call_tool(gateway, owner_key, identity, f"{alias}-add", ADD)
|
||||
assert called.status_code == 200, called.text
|
||||
assert _header(_one_call(peer), header) == shape.format(secret=secret, basic="").encode()
|
||||
|
||||
|
||||
def test_deprecated_string_x_mcp_auth_lists_a_byok_server_for_a_key_without_a_user(gateway: Gateway) -> None:
|
||||
with mcp_peer() as peer, gateway.scenario() as scenario:
|
||||
alias: Final = "byok" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias, auth_type="bearer_token", is_byok=True)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
peer.drain()
|
||||
response: Final = gateway.client.get(
|
||||
"/mcp-rest/tools/list",
|
||||
params={"server_id": identity},
|
||||
headers={"x-litellm-api-key": key, "x-mcp-auth": "Bearer hdr"},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
names: Final = {tool["name"] for tool in response.json()["tools"]}
|
||||
assert "add" in names, names
|
||||
listings: Final = _listings(peer)
|
||||
assert len(listings) == 1, listings
|
||||
assert listings[0]["headers"].get(b"authorization") == b"Bearer hdr"
|
||||
|
|
|
|||
194
tests/integration/mcp/test_mcp_listed_tool_metadata.py
Normal file
194
tests/integration/mcp/test_mcp_listed_tool_metadata.py
Normal file
|
|
@ -0,0 +1,194 @@
|
|||
"""pre_mcp_call guardrails are handed the tool entry ``tools/list`` served to the caller.
|
||||
|
||||
One owned proxy carries a default-on ``custom_code`` pre_mcp_call guardrail. At listing time it masks
|
||||
``SECRET`` out of every scanned text. At call time, when an argument carries the probe marker, it
|
||||
blocks and echoes the description and parameters it was handed, which is the only way to observe from outside
|
||||
what metadata the gateway attached to the hook
|
||||
"""
|
||||
|
||||
import json
|
||||
import uuid
|
||||
from collections.abc import Iterator, Mapping
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
import yaml
|
||||
from integration._support.client import Gateway, gateway_from_environment
|
||||
from integration._support.mcp import (
|
||||
EntryPoint,
|
||||
McpCaller,
|
||||
ScriptedTool,
|
||||
listed_tools,
|
||||
openapi_peer,
|
||||
register_mcp,
|
||||
scripted_peer,
|
||||
text_result,
|
||||
)
|
||||
from integration._support.process import owned_proxy
|
||||
|
||||
_ECHO: Final = "catalog-echo:"
|
||||
_PROBE: Final = "catalog-probe"
|
||||
_GUARDRAIL_CODE: Final = (
|
||||
"def apply_guardrail(inputs, request_data, input_type):\n"
|
||||
' texts = list(inputs.get("texts") or [])\n'
|
||||
' function = inputs.get("tools", [{}])[0].get("function", {})\n'
|
||||
f' if "{_PROBE}" in texts:\n'
|
||||
f' return block("{_ECHO}" + json_stringify('
|
||||
'{"description": function.get("description"), "parameters": function.get("parameters")}))\n'
|
||||
' masked = [text.replace("SECRET", "[MASKED]") for text in texts]\n'
|
||||
" if masked != texts:\n"
|
||||
" return modify(texts=masked)\n"
|
||||
" return allow()\n"
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]:
|
||||
directory: Final = tmp_path_factory.mktemp("listed-tool-metadata")
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["guardrails"] = [
|
||||
{
|
||||
"guardrail_name": "catalog-echo-" + uuid.uuid4().hex[:8],
|
||||
"litellm_params": {
|
||||
"guardrail": "custom_code",
|
||||
"mode": "pre_mcp_call",
|
||||
"default_on": True,
|
||||
"custom_code": _GUARDRAIL_CODE,
|
||||
},
|
||||
}
|
||||
]
|
||||
path: Final = directory / "config.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
with gateway_from_environment() as gateway, owned_proxy(gateway, directory, {}, config=path) as candidate:
|
||||
yield candidate
|
||||
|
||||
|
||||
def _strings(value: object) -> Iterator[str]:
|
||||
if isinstance(value, str):
|
||||
yield value
|
||||
return
|
||||
children: Final = value.values() if isinstance(value, Mapping) else value if isinstance(value, list) else ()
|
||||
for child in children:
|
||||
yield from _strings(child)
|
||||
|
||||
|
||||
def _decoded(raw: str) -> object:
|
||||
data: Final = tuple(line[5:].strip() for line in raw.splitlines() if line.startswith("data:"))
|
||||
return json.loads(data[-1] if data else raw)
|
||||
|
||||
|
||||
def _echoed(raw: str) -> tuple[str | None, Mapping[str, object] | None]:
|
||||
"""The (description, parameters) the guardrail was handed, recovered from its block reason."""
|
||||
carrier: Final = next((text for text in _strings(_decoded(raw)) if _ECHO in text), None)
|
||||
assert carrier is not None, raw
|
||||
echoed, _ = json.JSONDecoder().raw_decode(carrier.split(_ECHO, 1)[1])
|
||||
assert isinstance(echoed, dict), carrier
|
||||
return echoed.get("description"), echoed.get("parameters")
|
||||
|
||||
|
||||
def _probe(caller: McpCaller, name: str, server_id: str) -> tuple[str | None, Mapping[str, object] | None]:
|
||||
outcome: Final = caller.call(name, {"probe": _PROBE}, server_id=server_id)
|
||||
assert outcome.error is not None, outcome.raw
|
||||
return _echoed(outcome.raw)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("entry", ["rest", "mcp"])
|
||||
def test_pre_call_hook_receives_the_description_and_input_schema_the_caller_was_listed(
|
||||
rig: Gateway, entry: EntryPoint
|
||||
) -> None:
|
||||
schema: Final = {"type": "object", "properties": {"probe": {"type": "string", "description": "a probe marker"}}}
|
||||
tool: Final = ScriptedTool(
|
||||
"lookup", lambda _: text_result("found"), description="Look up one record", input_schema=schema
|
||||
)
|
||||
with scripted_peer(tool) as peer, rig.scenario() as scenario:
|
||||
alias: Final = "meta" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
caller: Final = McpCaller(rig, key, entry, headers={"x-mcp-servers": alias})
|
||||
assert caller.initialize().ok
|
||||
listed: Final = caller.list_tools(server_id=identity)
|
||||
assert listed.ok, listed.raw
|
||||
name: Final = next(full for full in listed.tools if full.endswith("lookup"))
|
||||
description, parameters = _probe(caller, name, identity)
|
||||
assert description == "Look up one record", (description, parameters)
|
||||
assert parameters is not None and parameters.get("properties") == schema["properties"], parameters
|
||||
|
||||
|
||||
def test_each_caller_is_evaluated_against_the_catalog_its_own_forwarded_headers_produced(rig: Gateway) -> None:
|
||||
tool: Final = ScriptedTool(
|
||||
"report",
|
||||
lambda _: text_result("ok"),
|
||||
description=lambda headers: f"Report for tenant {headers.get('x-tenant', 'nobody')}",
|
||||
)
|
||||
with scripted_peer(tool) as peer, rig.scenario() as scenario:
|
||||
alias: Final = "tenant" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias, extra_headers=["x-tenant"])
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
acme: Final = McpCaller(rig, key, "server_mcp", alias, headers={"x-tenant": "acme"})
|
||||
globex: Final = McpCaller(rig, key, "server_mcp", alias, headers={"x-tenant": "globex"})
|
||||
acme_listing: Final = acme.list_tools()
|
||||
globex_listing: Final = globex.list_tools()
|
||||
assert acme_listing.ok and globex_listing.ok, (acme_listing.raw, globex_listing.raw)
|
||||
name: Final = next(full for full in acme_listing.tools if full.endswith("report"))
|
||||
acme_seen, _ = _probe(acme, name, identity)
|
||||
globex_seen, _ = _probe(globex, name, identity)
|
||||
assert (acme_seen, globex_seen) == ("Report for tenant acme", "Report for tenant globex"), (
|
||||
"each caller's tools/call must be evaluated against the catalog its own headers listed"
|
||||
)
|
||||
|
||||
|
||||
def test_call_is_evaluated_against_the_masked_description_the_listing_served(rig: Gateway) -> None:
|
||||
tool: Final = ScriptedTool("read_note", lambda _: text_result("note"), description="Read a note")
|
||||
with scripted_peer(tool) as peer, rig.scenario() as scenario:
|
||||
alias: Final = "note" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(
|
||||
scenario, peer, alias, tool_name_to_description={"read_note": "Read a SECRET note"}
|
||||
)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
served: Final = listed_tools(rig, key, identity)
|
||||
name: Final = next(full for full in served if full.endswith("read_note"))
|
||||
assert served[name]["description"] == "Read a [MASKED] note", served[name]
|
||||
seen, _ = _probe(McpCaller(rig, key, "rest"), name, identity)
|
||||
assert seen == "Read a [MASKED] note", "the admin override must not restore wording the listing masked"
|
||||
|
||||
|
||||
def test_openapi_call_is_evaluated_against_the_masked_override_the_listing_served(rig: Gateway) -> None:
|
||||
with openapi_peer() as peer, rig.scenario() as scenario:
|
||||
alias: Final = "pets" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(
|
||||
scenario, peer, alias, tool_name_to_description={"getpet": "Fetch one SECRET pet"}
|
||||
)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
served: Final = listed_tools(rig, key, identity)
|
||||
name: Final = next(full for full in served if full.endswith("getpet"))
|
||||
assert served[name]["description"] == "Fetch one [MASKED] pet", served[name]
|
||||
seen, parameters = _probe(McpCaller(rig, key, "rest"), name, identity)
|
||||
assert seen == "Fetch one [MASKED] pet", "the OpenAPI call path must hand hooks the entry the listing served"
|
||||
assert parameters is not None and "petId" in parameters.get("properties", {}), parameters
|
||||
assert not [call for call in peer.drain() if call["path"].startswith("/pets")], "blocked before upstream"
|
||||
|
||||
|
||||
def test_openapi_call_is_evaluated_against_the_entry_this_key_was_listed_not_the_last_listing(
|
||||
rig: Gateway,
|
||||
) -> None:
|
||||
with openapi_peer() as peer, rig.scenario() as scenario:
|
||||
alias: Final = "pets" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(
|
||||
scenario, peer, alias, tool_name_to_description={"getpet": "Fetch one SECRET pet"}
|
||||
)
|
||||
guarded: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
opted_out: Final = scenario.key(
|
||||
object_permission={"mcp_servers": [identity]}, metadata={"disable_global_guardrails": True}
|
||||
)
|
||||
guarded_served: Final = listed_tools(rig, guarded, identity)
|
||||
opted_out_served: Final = listed_tools(rig, opted_out, identity)
|
||||
name: Final = next(full for full in guarded_served if full.endswith("getpet"))
|
||||
assert (guarded_served[name]["description"], opted_out_served[name]["description"]) == (
|
||||
"Fetch one [MASKED] pet",
|
||||
"Fetch one SECRET pet",
|
||||
), (guarded_served[name], opted_out_served[name])
|
||||
seen, _ = _probe(McpCaller(rig, guarded, "rest"), name, identity)
|
||||
assert seen == "Fetch one [MASKED] pet", (
|
||||
"the guarded key must be evaluated against its own listing, not the opted-out key's later one"
|
||||
)
|
||||
|
|
@ -48,6 +48,7 @@ def _bare_manager() -> MOD.MCPServerManager:
|
|||
reaches the guardrail hooks; they have their own coverage elsewhere.
|
||||
"""
|
||||
mgr = MOD.MCPServerManager.__new__(MOD.MCPServerManager)
|
||||
mgr._listed_tools_by_server_id = {}
|
||||
mgr.check_allowed_or_banned_tools = lambda name, server: True
|
||||
mgr.validate_allowed_params = lambda tool_name, arguments, server: None
|
||||
|
||||
|
|
|
|||
|
|
@ -1,17 +1,29 @@
|
|||
from litellm.proxy._experimental.mcp_server import operations as mcp_operations
|
||||
import asyncio
|
||||
import json
|
||||
from datetime import datetime
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from mcp.shared.exceptions import MCPError
|
||||
from mcp.types import CallToolResult, TextContent
|
||||
from mcp.types import Tool as MCPTool
|
||||
from pydantic import AnyUrl
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy._experimental.mcp_server import server
|
||||
from litellm.proxy._experimental.mcp_server.mcp_context import _mcp_proxy_mode
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller
|
||||
from litellm.proxy._experimental.mcp_server.tool_search import (
|
||||
handle_mcp_proxy_tool,
|
||||
mcp_proxy_tool_id,
|
||||
with_mcp_proxy_identity,
|
||||
)
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth
|
||||
from litellm.types.mcp import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
AUTH = UserAPIKeyAuth(api_key="key")
|
||||
|
||||
|
|
@ -130,3 +142,46 @@ async def test_proxy_scope_exception_emits_failure_log(monkeypatch: pytest.Monke
|
|||
assert hook_payload["arguments"] == arguments
|
||||
assert "raw_headers" not in hook_payload
|
||||
assert "raw-scope-secret" not in recorder.events[1][1]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_proxy_call_tool_on_a_never_listed_tool_hands_the_pre_hook_no_listed_tool() -> None:
|
||||
"""/mcp/proxy tools/list serves only the meta-tools, so the catalog call_tool reads to resolve its
|
||||
tool_id was never served: it must not fill the caller's listed-tools slot, and the pre-call hook
|
||||
must see no listed tool for the call."""
|
||||
manager = mcp_operations.global_mcp_server_manager
|
||||
server = MCPServer(server_id="proxy-meta", name="proxy-meta", transport=MCPTransport.http, url="http://meta")
|
||||
auth = UserAPIKeyAuth(api_key="sk-proxy-meta", user_id="proxy-caller")
|
||||
upstream = [MCPTool(name="echo", description="Echo text back", inputSchema={"type": "object"})]
|
||||
served_as = with_mcp_proxy_identity(MCPTool(name="proxy-meta-echo", inputSchema={}), server.server_id)
|
||||
pre_call_tool_check = AsyncMock(return_value={})
|
||||
|
||||
async def call_regular_mcp_tool(*, tasks: list[asyncio.Task[object]], **_: object) -> CallToolResult:
|
||||
await asyncio.gather(*tasks)
|
||||
return CallToolResult(content=[TextContent(type="text", text="echoed")])
|
||||
|
||||
with (
|
||||
patch.dict(manager.registry, {server.server_id: server}),
|
||||
patch.dict(manager.tool_name_to_mcp_server_name_mapping),
|
||||
patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])),
|
||||
patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())),
|
||||
patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)),
|
||||
patch.object(manager, "pre_call_tool_check", pre_call_tool_check),
|
||||
patch.object(manager, "_call_regular_mcp_tool", call_regular_mcp_tool),
|
||||
):
|
||||
try:
|
||||
result = await handle_mcp_proxy_tool(
|
||||
name="call_tool",
|
||||
arguments={"tool_id": mcp_proxy_tool_id(served_as), "arguments": {}},
|
||||
user_api_key_dict=auth,
|
||||
)
|
||||
listed = manager.get_listed_tool(server, "echo", ListedToolsCaller(user_api_key_auth=auth))
|
||||
finally:
|
||||
manager._drop_listed_tools(server.server_id)
|
||||
|
||||
assert result.is_error is False
|
||||
assert result.content[0].text == "echoed"
|
||||
pre_call_tool_check.assert_awaited_once()
|
||||
assert pre_call_tool_check.await_args.kwargs["name"] == "echo"
|
||||
assert pre_call_tool_check.await_args.kwargs["tool"] is None
|
||||
assert listed is None
|
||||
|
|
|
|||
|
|
@ -963,6 +963,8 @@ async def test_get_tools_from_mcp_servers():
|
|||
user_api_key_auth=None,
|
||||
oauth2_headers=None,
|
||||
proxy_logging_obj=None,
|
||||
catalog_auth_header=None,
|
||||
record_listing=True,
|
||||
):
|
||||
if server.server_id == "server1_id":
|
||||
return [mock_tool_1]
|
||||
|
|
@ -1998,6 +2000,7 @@ async def test_get_tools_for_single_server():
|
|||
client_ip=None,
|
||||
user_api_key_auth=None,
|
||||
proxy_logging_obj=ANY,
|
||||
record_listing=False,
|
||||
)
|
||||
|
||||
# Verify the result
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -9,6 +9,7 @@ from types import SimpleNamespace
|
|||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from mcp import ReadResourceResult, Resource
|
||||
|
|
@ -23,18 +24,20 @@ from mcp.types import (
|
|||
TextContent,
|
||||
TextResourceContents,
|
||||
)
|
||||
from mcp.types import Tool as MCPTool
|
||||
from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS, LATEST_HANDSHAKE_VERSION, MODERN_PROTOCOL_VERSIONS
|
||||
from pydantic import TypeAdapter
|
||||
from starlette.types import Message, Receive, Scope, Send
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_MCPServerTable,
|
||||
MCPTransport,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer, PinnedMCPTool
|
||||
|
||||
|
||||
def test_mcp_available_on_sdk2():
|
||||
|
|
@ -85,9 +88,6 @@ def cleanup_mcp_global_state():
|
|||
yield
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def _call_tool_params(name, arguments=None):
|
||||
from mcp.types import CallToolRequestParams
|
||||
|
||||
|
|
@ -99,6 +99,7 @@ def _paged_params():
|
|||
|
||||
return PaginatedRequestParams()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_server_tool_call_body_contains_request_data(_mcp_request_ctx):
|
||||
"""Test that proxy_server_request body contains name and arguments"""
|
||||
|
|
@ -295,7 +296,9 @@ async def test_mcp_server_tool_call_relays_upstream_auth_error_as_iserror(_mcp_r
|
|||
):
|
||||
with patch("litellm.proxy.proxy_server.proxy_config", MagicMock()):
|
||||
with patch("litellm.proxy._experimental.mcp_server.operations.verbose_logger", mock_logger):
|
||||
result = await mcp_server_tool_call(_mcp_request_ctx(), _call_tool_params("test_tool", {"param": "value"}))
|
||||
result = await mcp_server_tool_call(
|
||||
_mcp_request_ctx(), _call_tool_params("test_tool", {"param": "value"})
|
||||
)
|
||||
|
||||
assert result.is_error is True
|
||||
# The dedicated MCPUpstreamAuthError branch (not the generic Exception fallthrough) produces this
|
||||
|
|
@ -1167,20 +1170,32 @@ async def test_read_resource_preserves_content_metadata(_mcp_request_ctx, kind,
|
|||
else BlobResourceContents(uri=uri, blob="aGVsbG8=", mimeType="image/png", meta=metadata)
|
||||
)
|
||||
with (
|
||||
patch.object(server, "get_or_extract_auth_context", AsyncMock(return_value=(caller, None, ["catalog"], None, None, None, None))),
|
||||
patch.object(
|
||||
server,
|
||||
"get_or_extract_auth_context",
|
||||
AsyncMock(return_value=(caller, None, ["catalog"], None, None, None, None)),
|
||||
),
|
||||
patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[upstream_server])),
|
||||
patch.object(operations.global_mcp_server_manager, "read_resource_from_server", AsyncMock(return_value=ReadResourceResult(contents=[content]))),
|
||||
patch.object(
|
||||
operations.global_mcp_server_manager,
|
||||
"read_resource_from_server",
|
||||
AsyncMock(return_value=ReadResourceResult(contents=[content])),
|
||||
),
|
||||
):
|
||||
result: Final = await server.read_resource(_mcp_request_ctx(), ReadResourceRequestParams(uri=uri))
|
||||
|
||||
assert result.model_dump(mode="json", by_alias=True, exclude_none=True) == {
|
||||
"cacheScope": "private", "resultType": "complete", "ttlMs": 0,
|
||||
"contents": [{
|
||||
"uri": uri,
|
||||
"mimeType": "text/plain" if kind == "text" else "image/png",
|
||||
"text" if kind == "text" else "blob": "hello world" if kind == "text" else "aGVsbG8=",
|
||||
**({"_meta": metadata} if metadata is not None else {}),
|
||||
}],
|
||||
"cacheScope": "private",
|
||||
"resultType": "complete",
|
||||
"ttlMs": 0,
|
||||
"contents": [
|
||||
{
|
||||
"uri": uri,
|
||||
"mimeType": "text/plain" if kind == "text" else "image/png",
|
||||
"text" if kind == "text" else "blob": "hello world" if kind == "text" else "aGVsbG8=",
|
||||
**({"_meta": metadata} if metadata is not None else {}),
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
|
|
@ -1674,7 +1689,9 @@ async def test_handle_list_tools_converts_permission_httpexception_to_mcp_error(
|
|||
with (
|
||||
patch( # test-quality-ok: the protocol handler reads auth from module context; no injection seam
|
||||
"litellm.proxy._experimental.mcp_server.server.get_or_extract_auth_context",
|
||||
new=AsyncMock(return_value=(None, None, None, None, None, None, None), side_effect=denial if denial_at_auth else None),
|
||||
new=AsyncMock(
|
||||
return_value=(None, None, None, None, None, None, None), side_effect=denial if denial_at_auth else None
|
||||
),
|
||||
),
|
||||
patch( # test-quality-ok: the listing helper is the handler's only collaborator; the suite's seam
|
||||
"litellm.proxy._experimental.mcp_server.operations._list_mcp_tools",
|
||||
|
|
@ -1913,8 +1930,8 @@ async def test_streamable_http_session_manager_is_stateless():
|
|||
("DELETE", b"", False),
|
||||
),
|
||||
)
|
||||
async def test_mcp_routing_initialize_to_stateful_no_session_to_stateless(_mcp_request_ctx,
|
||||
debug: bool, method: str, request_body: bytes, stateful: bool
|
||||
async def test_mcp_routing_initialize_to_stateful_no_session_to_stateless(
|
||||
_mcp_request_ctx, debug: bool, method: str, request_body: bytes, stateful: bool
|
||||
) -> None:
|
||||
from starlette.requests import Request
|
||||
from starlette.types import Message, Receive, Scope, Send
|
||||
|
|
@ -4056,7 +4073,8 @@ async def test_truncated_jsonrpc_response_with_nested_method_skips_lock(
|
|||
# parsed, with a nested "method" key in the first bytes to trip a flat
|
||||
# substring heuristic.
|
||||
response_prefix: Final = (
|
||||
'{"jsonrpc":"2.0","id":99,"' + response_field
|
||||
'{"jsonrpc":"2.0","id":99,"'
|
||||
+ response_field
|
||||
+ '":{"code":-32000,"message":"test","data":{"method":"GET","payload":"'
|
||||
).encode()
|
||||
response_body: Final = (
|
||||
|
|
@ -6564,8 +6582,12 @@ class TestGatewayCreateInitializationOptions:
|
|||
yield (None, None)
|
||||
|
||||
async def record_request(
|
||||
serving_server: object, read_stream: object, write_stream: object,
|
||||
*, lifespan_state: object, init_options: InitializationOptions,
|
||||
serving_server: object,
|
||||
read_stream: object,
|
||||
write_stream: object,
|
||||
*,
|
||||
lifespan_state: object,
|
||||
init_options: InitializationOptions,
|
||||
) -> None:
|
||||
captured["server_name"] = init_options.server_name
|
||||
|
||||
|
|
@ -6877,7 +6899,6 @@ async def test_probe_upstream_auth_surfaces_httpx_status_error():
|
|||
returning the response. The probe must catch that specifically (before the
|
||||
fail-open `except Exception`) so the auth check is not silently defeated.
|
||||
"""
|
||||
import httpx
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.server import _probe_upstream_auth
|
||||
|
||||
|
|
@ -7412,7 +7433,8 @@ async def test_execute_mcp_tool_rest_server_id_authoritative_for_unprefixed_tool
|
|||
return_value=oauth_server,
|
||||
),
|
||||
patch.object(
|
||||
mcp_operations, "_handle_managed_mcp_tool",
|
||||
mcp_operations,
|
||||
"_handle_managed_mcp_tool",
|
||||
new=fake_handle_managed_mcp_tool,
|
||||
),
|
||||
patch.object(
|
||||
|
|
@ -7658,7 +7680,8 @@ async def test_execute_mcp_tool_strips_a_prefix_that_contains_the_separator():
|
|||
return_value=alias_less_server,
|
||||
),
|
||||
patch.object(
|
||||
mcp_operations, "_handle_managed_mcp_tool",
|
||||
mcp_operations,
|
||||
"_handle_managed_mcp_tool",
|
||||
new=fake_handle_managed_mcp_tool,
|
||||
),
|
||||
patch.object(
|
||||
|
|
@ -7933,7 +7956,8 @@ async def test_execute_mcp_tool_rest_hyphenated_upstream_tool_name_routes_to_req
|
|||
return_value=None,
|
||||
),
|
||||
patch.object(
|
||||
mcp_operations, "_handle_managed_mcp_tool",
|
||||
mcp_operations,
|
||||
"_handle_managed_mcp_tool",
|
||||
new=fake_handle_managed_mcp_tool,
|
||||
),
|
||||
patch.object(
|
||||
|
|
@ -7989,9 +8013,12 @@ async def test_execute_mcp_tool_sets_model_in_model_call_details():
|
|||
fake_server.server_name = "openapi-petstore"
|
||||
fake_server.alias = None
|
||||
fake_server.short_prefix = None
|
||||
fake_server.tool_name_to_description = None
|
||||
|
||||
fake_tool = MagicMock()
|
||||
fake_tool.name = "list_pets"
|
||||
fake_tool.description = "test tool"
|
||||
fake_tool.input_schema = {"type": "object"}
|
||||
|
||||
start_time = datetime.now(timezone.utc)
|
||||
litellm_logging_obj, _ = function_setup(
|
||||
|
|
@ -8042,6 +8069,348 @@ async def test_execute_mcp_tool_sets_model_in_model_call_details():
|
|||
assert litellm_logging_obj.model == "MCP: list_pets"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_hands_openapi_hooks_the_listed_entry_and_nothing_before_a_listing():
|
||||
"""A local-registry tools/call with no prior tools/list hands the pre-call hooks name and arguments
|
||||
only, as before this metadata existed, so a pre_mcp_call policy never scans a description the caller was
|
||||
not served. Once the caller has listed, the same call hands the entry that listing served."""
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy._experimental.mcp_server import operations as mcp_module
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
||||
petstore = MCPServer(
|
||||
server_id="petstore-id",
|
||||
name="petstore",
|
||||
server_name="petstore",
|
||||
transport=MCPTransport.http,
|
||||
url=None,
|
||||
spec_path="https://example.com/petstore.yaml",
|
||||
tool_name_to_description={"list_pets": "ADMIN DESC"},
|
||||
)
|
||||
schema = {"type": "object", "properties": {"limit": {"type": "integer"}}}
|
||||
mcp_module.global_mcp_tool_registry.register_tool(
|
||||
name="petstore-list_pets", description="List the pets", input_schema=schema, handler=lambda limit: "ok"
|
||||
)
|
||||
manager = mcp_module.global_mcp_server_manager
|
||||
manager._listed_tools_by_server_id.pop(petstore.server_id, None)
|
||||
alice = UserAPIKeyAuth(api_key="sk-user", user_id="alice")
|
||||
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
|
||||
proxy_logging.pre_call_hook = AsyncMock(return_value={})
|
||||
pre_call_tool_check = AsyncMock(wraps=manager.pre_call_tool_check)
|
||||
|
||||
async def call() -> tuple[MCPTool | None, dict]:
|
||||
await mcp_module.execute_mcp_tool(
|
||||
name="petstore-list_pets",
|
||||
arguments={"limit": 10},
|
||||
allowed_mcp_servers=[petstore],
|
||||
start_time=datetime.now(),
|
||||
user_api_key_auth=alice,
|
||||
)
|
||||
return pre_call_tool_check.call_args.kwargs["tool"], proxy_logging.pre_call_hook.call_args.kwargs["data"]
|
||||
|
||||
try:
|
||||
with (
|
||||
patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore),
|
||||
patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check),
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging),
|
||||
):
|
||||
never_listed_tool, never_listed_data = await call()
|
||||
manager._record_listed_tools(
|
||||
petstore,
|
||||
[MCPTool(name="list_pets", description="ADMIN DESC", inputSchema=schema)],
|
||||
ListedToolsCaller(user_api_key_auth=alice),
|
||||
)
|
||||
listed_tool, listed_data = await call()
|
||||
finally:
|
||||
mcp_module.global_mcp_tool_registry.unregister_tools_with_prefix("petstore-")
|
||||
manager._listed_tools_by_server_id.pop(petstore.server_id, None)
|
||||
|
||||
assert never_listed_tool is None
|
||||
assert (never_listed_data.get("mcp_tool_description"), never_listed_data.get("mcp_input_schema")) == (None, None)
|
||||
assert listed_tool is not None and (listed_tool.description, listed_tool.input_schema) == ("ADMIN DESC", schema)
|
||||
assert (listed_data["mcp_tool_description"], listed_data["mcp_input_schema"]) == ("ADMIN DESC", schema)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_hands_openapi_hooks_the_guarded_catalog_entry_clients_saw():
|
||||
"""When tools/list pinned the schema and masked the description of an OpenAPI tool, the local-registry
|
||||
call path must hand the pre-call hooks that served entry, not the raw registry one."""
|
||||
from litellm.proxy._experimental.mcp_server import operations as mcp_module
|
||||
|
||||
petstore = MCPServer(
|
||||
server_id="petstore-id",
|
||||
name="petstore",
|
||||
server_name="petstore",
|
||||
transport=MCPTransport.http,
|
||||
url=None,
|
||||
spec_path="https://example.com/petstore.yaml",
|
||||
tool_name_to_description={"getpetbyid": "Find a SECRET pet"},
|
||||
)
|
||||
registry_schema = {"type": "object", "properties": {"petId": {"type": "integer"}, "dump_all": {"type": "boolean"}}}
|
||||
pinned_schema = {"type": "object", "properties": {"petId": {"type": "integer"}}}
|
||||
mcp_module.global_mcp_tool_registry.register_tool(
|
||||
name="petstore-getpetbyid",
|
||||
description="Find pet by ID",
|
||||
input_schema=registry_schema,
|
||||
handler=lambda petId: "ok",
|
||||
)
|
||||
manager = mcp_module.global_mcp_server_manager
|
||||
alice = UserAPIKeyAuth(api_key="sk-user", user_id="alice")
|
||||
manager._record_listed_tools(
|
||||
petstore,
|
||||
[MCPTool(name="getpetbyid", description="Find a [MASKED] pet", inputSchema=pinned_schema)],
|
||||
ListedToolsCaller(user_api_key_auth=alice),
|
||||
)
|
||||
pre_call_tool_check = AsyncMock(return_value={})
|
||||
|
||||
try:
|
||||
with (
|
||||
patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore),
|
||||
patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check),
|
||||
):
|
||||
await mcp_module.execute_mcp_tool(
|
||||
name="petstore-getpetbyid",
|
||||
arguments={"petId": 1},
|
||||
allowed_mcp_servers=[petstore],
|
||||
start_time=datetime.now(),
|
||||
user_api_key_auth=alice,
|
||||
)
|
||||
finally:
|
||||
mcp_module.global_mcp_tool_registry.unregister_tools_with_prefix("petstore-")
|
||||
manager._listed_tools_by_server_id.pop(petstore.server_id, None)
|
||||
|
||||
handed_tool = pre_call_tool_check.call_args.kwargs["tool"]
|
||||
assert (handed_tool.description, handed_tool.input_schema) == ("Find a [MASKED] pet", pinned_schema), (
|
||||
"the pre-call policy must evaluate the entry tools/list served, not the raw registry entry"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_hands_openapi_hooks_each_callers_own_listed_entry():
|
||||
"""Two keys can be shown differently guarded OpenAPI catalogs. The call path must evaluate each key
|
||||
against the entry its own tools/list served, not the entry the most recent listing left behind."""
|
||||
from litellm.proxy._experimental.mcp_server import operations as mcp_module
|
||||
|
||||
petstore = MCPServer(
|
||||
server_id="petstore-id",
|
||||
name="petstore",
|
||||
server_name="petstore",
|
||||
transport=MCPTransport.http,
|
||||
url=None,
|
||||
spec_path="https://example.com/petstore.yaml",
|
||||
)
|
||||
schema = {"type": "object", "properties": {"petId": {"type": "integer"}}}
|
||||
mcp_module.global_mcp_tool_registry.register_tool(
|
||||
name="petstore-getpetbyid", description="Find a SECRET pet", input_schema=schema, handler=lambda petId: "ok"
|
||||
)
|
||||
manager = mcp_module.global_mcp_server_manager
|
||||
guarded = UserAPIKeyAuth(api_key="sk-guarded", user_id="alice")
|
||||
opted_out = UserAPIKeyAuth(api_key="sk-opted-out", user_id="bob")
|
||||
manager._record_listed_tools(
|
||||
petstore,
|
||||
[MCPTool(name="getpetbyid", description="Find a [MASKED] pet", inputSchema=schema)],
|
||||
ListedToolsCaller(user_api_key_auth=guarded),
|
||||
)
|
||||
manager._record_listed_tools(
|
||||
petstore,
|
||||
[MCPTool(name="getpetbyid", description="Find a SECRET pet", inputSchema=schema)],
|
||||
ListedToolsCaller(user_api_key_auth=opted_out),
|
||||
)
|
||||
pre_call_tool_check = AsyncMock(return_value={})
|
||||
|
||||
try:
|
||||
with (
|
||||
patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore),
|
||||
patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check),
|
||||
):
|
||||
for caller in (guarded, opted_out):
|
||||
await mcp_module.execute_mcp_tool(
|
||||
name="petstore-getpetbyid",
|
||||
arguments={"petId": 1},
|
||||
allowed_mcp_servers=[petstore],
|
||||
start_time=datetime.now(),
|
||||
user_api_key_auth=caller,
|
||||
)
|
||||
finally:
|
||||
mcp_module.global_mcp_tool_registry.unregister_tools_with_prefix("petstore-")
|
||||
manager._listed_tools_by_server_id.pop(petstore.server_id, None)
|
||||
|
||||
handed = [call.kwargs["tool"].description for call in pre_call_tool_check.call_args_list]
|
||||
assert handed == ["Find a [MASKED] pet", "Find a SECRET pet"], (
|
||||
"each key's tools/call must be evaluated against the OpenAPI entry its own listing served"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_runs_the_longer_colliding_operation_and_hands_hooks_no_registry_metadata():
|
||||
"""An OpenAPI operation whose name starts with its own server prefix runs instead of the shorter one, and
|
||||
with no prior listing the pre-call hooks get name and arguments only, never either registry entry."""
|
||||
from litellm.proxy._experimental.mcp_server import operations as mcp_module
|
||||
|
||||
petstore = MCPServer(
|
||||
server_id="petstore-id",
|
||||
name="petstore",
|
||||
server_name="petstore",
|
||||
transport=MCPTransport.http,
|
||||
url=None,
|
||||
spec_path="https://example.com/petstore.yaml",
|
||||
)
|
||||
registry = mcp_module.global_mcp_tool_registry
|
||||
registry.register_tool(name="petstore-get_pet", description="short", input_schema={}, handler=lambda: "short")
|
||||
registry.register_tool(
|
||||
name="petstore-petstore-get_pet",
|
||||
description="long",
|
||||
input_schema={"type": "object", "properties": {"petId": {"type": "integer"}}},
|
||||
handler=lambda: "long",
|
||||
)
|
||||
manager = mcp_module.global_mcp_server_manager
|
||||
pre_call_tool_check = AsyncMock(return_value={})
|
||||
|
||||
try:
|
||||
with (
|
||||
patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore),
|
||||
patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check),
|
||||
):
|
||||
result = await mcp_module.execute_mcp_tool(
|
||||
name="petstore-petstore-get_pet",
|
||||
arguments={},
|
||||
allowed_mcp_servers=[petstore],
|
||||
start_time=datetime.now(),
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="sk-user", user_id="alice"),
|
||||
)
|
||||
finally:
|
||||
registry.unregister_tools_with_prefix("petstore-")
|
||||
|
||||
assert pre_call_tool_check.call_args.kwargs["tool"] is None
|
||||
assert result.content[0].text == "long"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_hands_hooks_nothing_for_a_never_listed_operation_named_after_a_listed_one():
|
||||
"""After the caller listed ``get_pet``, a call to the never-listed ``petstore-get_pet`` operation hands the
|
||||
pre-call hooks name and arguments only, not the listed sibling's description and schema."""
|
||||
from litellm.proxy._experimental.mcp_server import operations as mcp_module
|
||||
|
||||
petstore = MCPServer(
|
||||
server_id="petstore-id",
|
||||
name="petstore",
|
||||
server_name="petstore",
|
||||
transport=MCPTransport.http,
|
||||
url=None,
|
||||
spec_path="https://example.com/petstore.yaml",
|
||||
)
|
||||
registry = mcp_module.global_mcp_tool_registry
|
||||
registry.register_tool(
|
||||
name="petstore-petstore-get_pet", description="long", input_schema={}, handler=lambda: "long"
|
||||
)
|
||||
manager = mcp_module.global_mcp_server_manager
|
||||
alice = UserAPIKeyAuth(api_key="sk-user", user_id="alice")
|
||||
manager._record_listed_tools(
|
||||
petstore,
|
||||
[MCPTool(name="get_pet", description="Fetches pet records. FLAGWORD", inputSchema={"type": "object"})],
|
||||
ListedToolsCaller(user_api_key_auth=alice),
|
||||
)
|
||||
pre_call_tool_check = AsyncMock(return_value={})
|
||||
|
||||
try:
|
||||
with (
|
||||
patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore),
|
||||
patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check),
|
||||
):
|
||||
result = await mcp_module.execute_mcp_tool(
|
||||
name="petstore-petstore-get_pet",
|
||||
arguments={},
|
||||
allowed_mcp_servers=[petstore],
|
||||
start_time=datetime.now(),
|
||||
user_api_key_auth=alice,
|
||||
)
|
||||
finally:
|
||||
registry.unregister_tools_with_prefix("petstore-")
|
||||
manager._listed_tools_by_server_id.pop(petstore.server_id, None)
|
||||
|
||||
assert pre_call_tool_check.call_args.kwargs["tool"] is None
|
||||
assert result.content[0].text == "long"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_implicit_listing_before_the_first_call_hands_hooks_no_description():
|
||||
"""The listing tools/call runs on its own when this worker does not yet expose the tool is never served
|
||||
to the caller, so it leaves the caller's listed slot empty and the pre-call hooks still get name and
|
||||
arguments only, as on main."""
|
||||
manager = mcp_operations.global_mcp_server_manager
|
||||
server = _never_listed_passthrough_server()
|
||||
manager.registry[server.server_id] = server
|
||||
manager._listed_tools_by_server_id.pop(server.server_id, None)
|
||||
upstream = AsyncMock()
|
||||
upstream.call_tool.return_value = CallToolResult(content=[TextContent(type="text", text="ok")], isError=False)
|
||||
proxy_logging = _mock_mcp_proxy_logging()
|
||||
proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
|
||||
proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={})
|
||||
proxy_logging.pre_call_hook = AsyncMock(return_value={})
|
||||
proxy_logging.during_call_hook = AsyncMock(return_value=None)
|
||||
fetch_tools = AsyncMock(
|
||||
return_value=[MCPTool(name="add", description="Adds. FLAGWORD", inputSchema={"type": "object"})]
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(manager, "_create_mcp_client", new=AsyncMock(return_value=upstream)),
|
||||
patch.object(manager, "_fetch_tools_with_timeout", new=fetch_tools),
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging),
|
||||
):
|
||||
result = await mcp_operations.execute_mcp_tool(
|
||||
name="lazy_map-add",
|
||||
arguments={"a": 1, "b": 2},
|
||||
allowed_mcp_servers=[server],
|
||||
start_time=datetime.now(),
|
||||
mcp_auth_header="Bearer caller-token",
|
||||
raw_headers={"authorization": "Bearer caller-token"},
|
||||
)
|
||||
|
||||
assert fetch_tools.await_count == 1
|
||||
assert upstream.call_tool.await_count == 1
|
||||
assert result.content[0].text == "ok"
|
||||
hook_kwargs = proxy_logging._create_mcp_request_object_from_kwargs.call_args.args[0]
|
||||
assert (hook_kwargs["tool_description"], hook_kwargs["tool_input_schema"]) == (None, None)
|
||||
assert server.server_id not in manager._listed_tools_by_server_id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_pinnable_tool_catalog_records_no_listed_catalog_for_the_admin():
|
||||
"""The pin snapshot lists the raw upstream catalog, without the catalog guard or the admin's description
|
||||
overrides, so it must not become what the admin's own later tools/call is evaluated against."""
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy._experimental.mcp_server.rest_endpoints import fetch_pinnable_tool_catalog
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
||||
manager = mcp_operations.global_mcp_server_manager
|
||||
server = MCPServer(
|
||||
server_id="pin-srv",
|
||||
name="pin_srv",
|
||||
transport=MCPTransport.http,
|
||||
url="https://up.example.com/mcp",
|
||||
tool_name_to_description={"add": "Admin wording"},
|
||||
)
|
||||
manager._listed_tools_by_server_id.pop(server.server_id, None)
|
||||
admin = UserAPIKeyAuth(api_key="sk-admin", user_id="admin")
|
||||
request = MagicMock()
|
||||
request.client.host = "10.1.2.3"
|
||||
request.headers = {"x-litellm-api-key": "sk-admin"}
|
||||
fetch_tools = AsyncMock(
|
||||
return_value=[MCPTool(name="add", description="Upstream wording", inputSchema={"type": "object"})]
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(manager, "_create_mcp_client", new=AsyncMock(return_value=MagicMock())),
|
||||
patch.object(manager, "_fetch_tools_with_timeout", new=fetch_tools),
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", ProxyLogging(user_api_key_cache=DualCache())),
|
||||
):
|
||||
snapshot = await fetch_pinnable_tool_catalog(server, request, admin)
|
||||
|
||||
assert snapshot == {"add": PinnedMCPTool(description="Upstream wording", input_schema={"type": "object"})}
|
||||
assert server.server_id not in manager._listed_tools_by_server_id
|
||||
assert manager.get_listed_tool(server, "add", ListedToolsCaller(user_api_key_auth=admin)) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_rest_unresolved_prefixed_name_routes_to_requested_server():
|
||||
"""A prefixed REST name that resolves to no tool must still dispatch to the server_id.
|
||||
|
|
@ -8098,7 +8467,8 @@ async def test_execute_mcp_tool_rest_unresolved_prefixed_name_routes_to_requeste
|
|||
return_value=None,
|
||||
),
|
||||
patch.object(
|
||||
mcp_operations, "_handle_managed_mcp_tool",
|
||||
mcp_operations,
|
||||
"_handle_managed_mcp_tool",
|
||||
new=fake_handle_managed_mcp_tool,
|
||||
),
|
||||
patch.object(
|
||||
|
|
@ -8597,7 +8967,9 @@ class TestMCPMetaTraceCarrier:
|
|||
|
||||
assert _mcp_meta_trace_carrier(None) is None
|
||||
assert _mcp_meta_trace_carrier(SimpleNamespace(meta=None)) is None
|
||||
only_progress = CallToolRequestParams.model_validate({"name": "t", "_meta": {"progressToken": "p1"}}, by_name=False).meta
|
||||
only_progress = CallToolRequestParams.model_validate(
|
||||
{"name": "t", "_meta": {"progressToken": "p1"}}, by_name=False
|
||||
).meta
|
||||
assert _mcp_meta_trace_carrier(SimpleNamespace(meta=only_progress)) is None
|
||||
|
||||
|
||||
|
|
@ -10394,7 +10766,9 @@ async def test_mcp_origin_admission_precedes_authentication(
|
|||
patch("litellm.proxy.proxy_server.origins", allowed_origins),
|
||||
patch.object(server, "extract_mcp_auth_context", authenticate),
|
||||
):
|
||||
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=server.app), base_url="http://gateway") as client:
|
||||
async with httpx.AsyncClient(
|
||||
transport=httpx.ASGITransport(app=server.app), base_url="http://gateway"
|
||||
) as client:
|
||||
response: Final = await client.request(method, path, headers=(*session_headers, *origin_headers))
|
||||
|
||||
assert response.status_code == expected_status
|
||||
|
|
@ -10477,12 +10851,15 @@ async def test_streamable_http_rejects_modern_protocol_version(
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("handler_name,field", [
|
||||
("handle_list_tools", "tools"),
|
||||
("list_prompts", "prompts"),
|
||||
("list_resources", "resources"),
|
||||
("list_resource_templates", "resource_templates"),
|
||||
])
|
||||
@pytest.mark.parametrize(
|
||||
"handler_name,field",
|
||||
[
|
||||
("handle_list_tools", "tools"),
|
||||
("list_prompts", "prompts"),
|
||||
("list_resources", "resources"),
|
||||
("list_resource_templates", "resource_templates"),
|
||||
],
|
||||
)
|
||||
async def test_native_listing_preserves_empty_result_on_auth_failure(_mcp_request_ctx, handler_name, field):
|
||||
from litellm.proxy._experimental.mcp_server import server
|
||||
|
||||
|
|
@ -10500,7 +10877,9 @@ async def test_tool_listing_preserves_permission_denial_when_failure_logging_fai
|
|||
auth = UserAPIKeyAuth(user_id="denied-caller")
|
||||
denial = HTTPException(status_code=403, detail="scope denied")
|
||||
logger = MagicMock()
|
||||
logger.post_call_failure_hook = AsyncMock(side_effect=RuntimeError("log unavailable") if failure_hook_raises else None)
|
||||
logger.post_call_failure_hook = AsyncMock(
|
||||
side_effect=RuntimeError("log unavailable") if failure_hook_raises else None
|
||||
)
|
||||
upstream = AsyncMock()
|
||||
with (
|
||||
patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(side_effect=denial)),
|
||||
|
|
@ -10509,7 +10888,9 @@ async def test_tool_listing_preserves_permission_denial_when_failure_logging_fai
|
|||
patch.object(operations.global_mcp_server_manager, "_get_tools_from_server", upstream),
|
||||
):
|
||||
with pytest.raises(HTTPException) as rejected:
|
||||
await operations._get_tools_from_mcp_servers(user_api_key_auth=auth, mcp_auth_header=None, mcp_servers=["catalog"], log_list_tools_to_spendlogs=True)
|
||||
await operations._get_tools_from_mcp_servers(
|
||||
user_api_key_auth=auth, mcp_auth_header=None, mcp_servers=["catalog"], log_list_tools_to_spendlogs=True
|
||||
)
|
||||
assert rejected.value is denial
|
||||
upstream.assert_not_awaited()
|
||||
logger.post_call_failure_hook.assert_awaited_once()
|
||||
|
|
@ -10521,7 +10902,9 @@ async def test_tool_listing_preserves_permission_denial_when_failure_logging_fai
|
|||
@pytest.mark.parametrize("prefix,suffix", (("", ""), ("/gateway", "/")))
|
||||
@pytest.mark.parametrize("opening_protocol", (None, *MODERN_PROTOCOL_VERSIONS))
|
||||
async def test_legacy_sse_mount_emits_message_endpoint(
|
||||
prefix: str, suffix: str, opening_protocol: str | None,
|
||||
prefix: str,
|
||||
suffix: str,
|
||||
opening_protocol: str | None,
|
||||
) -> None:
|
||||
from starlette.applications import Starlette
|
||||
from starlette.routing import Mount
|
||||
|
|
@ -10586,16 +10969,20 @@ async def test_legacy_sse_mount_emits_message_endpoint(
|
|||
return (await messages.get())["status"]
|
||||
|
||||
if opening_protocol is not None:
|
||||
discover: Final = json.dumps({
|
||||
"jsonrpc": "2.0",
|
||||
"id": 0,
|
||||
"method": "server/discover",
|
||||
"params": {"_meta": {
|
||||
"io.modelcontextprotocol/protocolVersion": opening_protocol,
|
||||
"io.modelcontextprotocol/clientInfo": {"name": "modern-client", "version": "1"},
|
||||
"io.modelcontextprotocol/clientCapabilities": {},
|
||||
}},
|
||||
}).encode()
|
||||
discover: Final = json.dumps(
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"id": 0,
|
||||
"method": "server/discover",
|
||||
"params": {
|
||||
"_meta": {
|
||||
"io.modelcontextprotocol/protocolVersion": opening_protocol,
|
||||
"io.modelcontextprotocol/clientInfo": {"name": "modern-client", "version": "1"},
|
||||
"io.modelcontextprotocol/clientCapabilities": {},
|
||||
}
|
||||
},
|
||||
}
|
||||
).encode()
|
||||
assert await post(discover) == 202
|
||||
discovered_frame: Final = (await asyncio.wait_for(outgoing.get(), 2))["body"].decode()
|
||||
discovered: Final = json.loads(discovered_frame.split("data: ", 1)[1].splitlines()[0])
|
||||
|
|
@ -10628,7 +11015,16 @@ async def test_legacy_sse_mount_emits_message_endpoint(
|
|||
patch.object(
|
||||
mcp_server,
|
||||
"extract_mcp_auth_context",
|
||||
AsyncMock(return_value=(post_auth, None, [marker], {marker: {"Authorization": marker}}, {"Authorization": marker}, {"x-request-marker": marker})),
|
||||
AsyncMock(
|
||||
return_value=(
|
||||
post_auth,
|
||||
None,
|
||||
[marker],
|
||||
{marker: {"Authorization": marker}},
|
||||
{"Authorization": marker},
|
||||
{"x-request-marker": marker},
|
||||
)
|
||||
),
|
||||
),
|
||||
patch.object(mcp_server.operations, "_get_tools_from_mcp_servers", listing),
|
||||
):
|
||||
|
|
@ -10677,7 +11073,11 @@ async def test_discovery_adapter_preserves_authenticated_context(_mcp_request_ct
|
|||
dispatched = AsyncMock(return_value=expected)
|
||||
auth = UserAPIKeyAuth(user_id="discover-caller")
|
||||
with (
|
||||
patch.object(server, "get_or_extract_auth_context", AsyncMock(return_value=(auth, None, ["allowed"], None, None, None, None))),
|
||||
patch.object(
|
||||
server,
|
||||
"get_or_extract_auth_context",
|
||||
AsyncMock(return_value=(auth, None, ["allowed"], None, None, None, None)),
|
||||
),
|
||||
patch.object(server.operations.GatewayOperations, "execute", dispatched),
|
||||
):
|
||||
result = await server.discover(_mcp_request_ctx(), RequestParams())
|
||||
|
|
|
|||
|
|
@ -22,6 +22,7 @@ from mcp.types import Tool
|
|||
import litellm
|
||||
from litellm.models.object_permission import LiteLLM_ObjectPermissionTable
|
||||
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller
|
||||
from litellm.proxy._experimental.mcp_server.tool_search import (
|
||||
AGENT_SEARCH_TOOL_NAME,
|
||||
MCP_TOOL_CALL_TOOL_NAME,
|
||||
|
|
@ -31,12 +32,14 @@ from litellm.proxy._experimental.mcp_server.tool_search import (
|
|||
ToolSearchResult,
|
||||
coerce_top_k,
|
||||
get_virtual_tool_definitions,
|
||||
handle_mcp_tool_search,
|
||||
search_mcp_tools,
|
||||
search_tools,
|
||||
)
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.semantic_text_index import EmbeddingFailed, SemanticTextIndex, Vector
|
||||
from litellm.types.mcp import MCPToolSearchSettings
|
||||
from litellm.types.mcp import MCPToolSearchSettings, MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
|
||||
def _make_tools(specs: list[tuple[str, str]]) -> tuple[Tool, ...]:
|
||||
|
|
@ -1268,3 +1271,33 @@ async def test_handle_mcp_tool_call_scoped_denial_names_the_binding_agent() -> N
|
|||
assert exc_info.value.status_code == 403
|
||||
assert "MCP server 'github'" in exc_info.value.detail["error"]
|
||||
assert "agent 'agent-123'" in exc_info.value.detail["error"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_tool_search_leaves_the_listed_tools_slot_empty(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""The search lists the whole catalog but serves only its hits, so the listing must not fill the
|
||||
caller's listed-tools slot: a later call to a tool the search never returned is not a listed tool."""
|
||||
monkeypatch.setattr(litellm, "mcp_tool_search", None)
|
||||
manager = mcp_operations.global_mcp_server_manager
|
||||
server = MCPServer(server_id="search-slot", name="search-slot", transport=MCPTransport.http, url="http://slot")
|
||||
user = UserAPIKeyAuth(api_key="sk-search-slot", user_id="searcher")
|
||||
upstream = [
|
||||
Tool(name="echo", description="Echo text back", inputSchema={"type": "object"}),
|
||||
Tool(name="delete_note", description="Delete a note", inputSchema={"type": "object"}),
|
||||
]
|
||||
with (
|
||||
patch.dict(manager.tool_name_to_mcp_server_name_mapping),
|
||||
patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])),
|
||||
patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())),
|
||||
patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)),
|
||||
):
|
||||
try:
|
||||
result = await handle_mcp_tool_search(query="echo", top_k=1, user_api_key_dict=user)
|
||||
caller = ListedToolsCaller(user_api_key_auth=user)
|
||||
listed = [manager.get_listed_tool(server, tool.name, caller) for tool in upstream]
|
||||
finally:
|
||||
manager._drop_listed_tools(server.server_id)
|
||||
|
||||
assert result.is_error is False
|
||||
assert [hit["name"] for hit in json.loads(result.content[0].text)] == ["search-slot-echo"]
|
||||
assert listed == [None, None]
|
||||
|
|
|
|||
|
|
@ -42,9 +42,12 @@ async def test_openapi_local_tool_runs_pre_call_tool_check():
|
|||
fake_server.server_name = "openapi-petstore"
|
||||
fake_server.alias = None
|
||||
fake_server.short_prefix = None
|
||||
fake_server.tool_name_to_description = None
|
||||
|
||||
fake_tool = MagicMock()
|
||||
fake_tool.name = "list_pets"
|
||||
fake_tool.description = "test tool"
|
||||
fake_tool.input_schema = {"type": "object"}
|
||||
|
||||
pre_call = AsyncMock(return_value={})
|
||||
handle_local = AsyncMock(return_value=CallToolResult(content=[], is_error=False))
|
||||
|
|
@ -125,9 +128,12 @@ async def test_openapi_local_tool_blocked_when_pre_call_check_raises():
|
|||
fake_server.server_name = "openapi-petstore"
|
||||
fake_server.alias = None
|
||||
fake_server.short_prefix = None
|
||||
fake_server.tool_name_to_description = None
|
||||
|
||||
fake_tool = MagicMock()
|
||||
fake_tool.name = "delete_pet"
|
||||
fake_tool.description = "test tool"
|
||||
fake_tool.input_schema = {"type": "object"}
|
||||
|
||||
pre_call = AsyncMock(
|
||||
side_effect=HTTPException(status_code=403, detail="not allowed")
|
||||
|
|
@ -190,6 +196,8 @@ async def test_openapi_local_tool_denied_when_server_not_resolvable():
|
|||
|
||||
fake_tool = MagicMock()
|
||||
fake_tool.name = "list_pets"
|
||||
fake_tool.description = "test tool"
|
||||
fake_tool.input_schema = {"type": "object"}
|
||||
|
||||
pre_call = AsyncMock(return_value={})
|
||||
handle_local = AsyncMock(return_value=CallToolResult(content=[], is_error=False))
|
||||
|
|
@ -274,6 +282,8 @@ async def test_openapi_local_tool_injects_resolved_oauth_token():
|
|||
|
||||
fake_tool = MagicMock()
|
||||
fake_tool.name = "get_values"
|
||||
fake_tool.description = "test tool"
|
||||
fake_tool.input_schema = {"type": "object"}
|
||||
captured: dict = {}
|
||||
|
||||
async def handle_local(_name, _arguments, _wire_compat):
|
||||
|
|
@ -620,6 +630,8 @@ async def test_per_server_auth_header_reaches_both_openapi_dispatch_arms(dispatc
|
|||
if dispatch_arm == "local_registry":
|
||||
fake_tool = MagicMock()
|
||||
fake_tool.name = "list_reports"
|
||||
fake_tool.description = "test tool"
|
||||
fake_tool.input_schema = {"type": "object"}
|
||||
with (
|
||||
patch.object(manager, "_get_mcp_server_from_tool_name", return_value=server),
|
||||
patch.object(mcp_operations.global_mcp_tool_registry, "get_tool", return_value=fake_tool),
|
||||
|
|
@ -691,6 +703,8 @@ async def test_local_dispatch_reports_the_outcome_instead_of_success(failure: st
|
|||
|
||||
fake_tool = MagicMock()
|
||||
fake_tool.name = "list_reports"
|
||||
fake_tool.description = "test tool"
|
||||
fake_tool.input_schema = {"type": "object"}
|
||||
fake_tool.handler = raising_handler
|
||||
server = MCPServer(
|
||||
server_id="srv-openapi",
|
||||
|
|
|
|||
|
|
@ -1,15 +1,101 @@
|
|||
import asyncio
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
from mcp.types import GetPromptRequest, GetPromptRequestParams, GetPromptResult
|
||||
from mcp.types import Tool as MCPTool
|
||||
|
||||
import litellm
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy._experimental.mcp_server import operations
|
||||
from litellm.proxy._experimental.mcp_server import rest_endpoints
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller
|
||||
from litellm.proxy._experimental.mcp_server.operations import GatewayOperations, prepare_context
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
from litellm.types.mcp import MCPAuth, MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
|
||||
class _CatalogHookCapture(CustomLogger):
|
||||
data: dict[str, object] | None = None
|
||||
|
||||
async def async_pre_call_hook(
|
||||
self, user_api_key_dict: UserAPIKeyAuth, cache: DualCache, data: dict[str, object], call_type: str
|
||||
) -> None:
|
||||
if call_type == "call_mcp_tool":
|
||||
self.data = data.copy()
|
||||
|
||||
|
||||
async def _served_catalog_tool() -> str:
|
||||
return "ok"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("surface", ["mcp", "rest"])
|
||||
@pytest.mark.parametrize("restriction", ["key", "server"])
|
||||
async def test_listing_records_only_tools_the_caller_received(
|
||||
monkeypatch: pytest.MonkeyPatch, surface: str, restriction: str
|
||||
) -> None:
|
||||
manager: Final = operations.global_mcp_server_manager
|
||||
server: Final = MCPServer(
|
||||
server_id="served-catalog", name="served-catalog", transport=MCPTransport.http,
|
||||
spec_path="/catalog.yaml", allow_all_keys=True,
|
||||
allowed_tools=["echo"] if restriction == "server" else None,
|
||||
)
|
||||
auth: Final = UserAPIKeyAuth(
|
||||
api_key="sk-served-catalog", user_id="lister",
|
||||
object_permission={
|
||||
"object_permission_id": "served-permission",
|
||||
"mcp_servers": [server.server_id],
|
||||
"mcp_tool_permissions": {server.server_id: ["echo"]} if restriction == "key" else None,
|
||||
},
|
||||
)
|
||||
monkeypatch.setitem(manager.registry, server.server_id, server)
|
||||
monkeypatch.setitem(manager.tool_name_to_mcp_server_name_mapping, "status", server.server_id)
|
||||
monkeypatch.setitem(manager.tool_name_to_mcp_server_name_mapping, "served-catalog-status", server.server_id)
|
||||
capture: Final = _CatalogHookCapture()
|
||||
monkeypatch.setattr(litellm, "callbacks", [capture])
|
||||
for name in ("echo", "status"):
|
||||
global_mcp_tool_registry.register_tool(
|
||||
name=f"served-catalog-{name}", description=f"{name} description",
|
||||
input_schema={"type": "object"}, handler=_served_catalog_tool,
|
||||
)
|
||||
try:
|
||||
if surface == "mcp":
|
||||
listing: Final = await operations._list_mcp_tools(
|
||||
user_api_key_auth=auth, mcp_servers=[server.server_id], record_listing=True,
|
||||
)
|
||||
assert [tool.name for tool in listing.tools] == ["served-catalog-echo"]
|
||||
else:
|
||||
rest_listing: Final = await rest_endpoints._get_tools_for_single_server(
|
||||
server, None, user_api_key_auth=auth,
|
||||
)
|
||||
assert [tool.name for tool in rest_listing] == ["echo"]
|
||||
granted: Final = auth.model_copy(update={"object_permission": None})
|
||||
caller: Final = ListedToolsCaller(user_api_key_auth=granted)
|
||||
assert manager.get_listed_tool(server, "status", caller) is None
|
||||
served: Final = manager.get_listed_tool(server, "echo", caller)
|
||||
assert served is not None
|
||||
assert (served.description, served.input_schema) == ("echo description", {"type": "object"})
|
||||
server.allowed_tools = None
|
||||
result: Final = await manager.call_tool(
|
||||
server_name=server.server_id, name="status", arguments={}, user_api_key_auth=granted,
|
||||
proxy_logging_obj=ProxyLogging(user_api_key_cache=UserApiKeyCache()),
|
||||
)
|
||||
assert result.is_error is False
|
||||
assert capture.data is not None
|
||||
assert capture.data["messages"] == [{"role": "user", "content": "Tool: status\nArguments: {}"}]
|
||||
assert (capture.data.get("mcp_tool_description"), capture.data.get("mcp_input_schema")) == (None, None)
|
||||
finally:
|
||||
manager._drop_listed_tools(server.server_id)
|
||||
global_mcp_tool_registry.unregister_tools_with_prefix("served-catalog-")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_oauth_prefetch_failure_does_not_log_caller_or_exception_text(caplog):
|
||||
from litellm.proxy._experimental.mcp_server.operations import _prefetch_oauth_creds_for_user
|
||||
|
|
@ -665,3 +751,30 @@ async def test_tools_listing_preserves_explicit_spend_log_policy(log_enabled):
|
|||
)
|
||||
assert result.tools == []
|
||||
assert listing.await_args.kwargs["log_list_tools_to_spendlogs"] is log_enabled
|
||||
assert listing.await_args.kwargs["record_listing"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(("listing_kwargs", "recorded"), [({}, False), ({"record_listing": True}, True)])
|
||||
async def test_list_mcp_tools_records_the_catalog_only_when_asked(
|
||||
listing_kwargs: dict[str, bool], recorded: bool
|
||||
) -> None:
|
||||
"""The aggregate listing fills the caller's listed-tools slot only when asked: a listing an internal
|
||||
caller never serves must not hand a later tools/call a description the caller never saw."""
|
||||
manager = operations.global_mcp_server_manager
|
||||
server = MCPServer(server_id="listing-slot", name="listing-slot", transport=MCPTransport.http, url="http://slot")
|
||||
user = UserAPIKeyAuth(api_key="sk-listing-slot", user_id="lister")
|
||||
upstream = [MCPTool(name="echo", description="Echo text back", inputSchema={"type": "object"})]
|
||||
with (
|
||||
patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])),
|
||||
patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())),
|
||||
patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)),
|
||||
patch.dict(manager.tool_name_to_mcp_server_name_mapping),
|
||||
):
|
||||
try:
|
||||
listing = await operations._list_mcp_tools(user_api_key_auth=user, **listing_kwargs)
|
||||
listed = manager.get_listed_tool(server, "echo", ListedToolsCaller(user_api_key_auth=user))
|
||||
finally:
|
||||
manager._drop_listed_tools(server.server_id)
|
||||
assert [tool.name for tool in listing.tools] == ["listing-slot-echo"]
|
||||
assert (listed is not None) is recorded
|
||||
|
|
|
|||
|
|
@ -346,6 +346,29 @@ class TestAllowFlow:
|
|||
assert evaluate_call.json["conversationId"] == "sess-123"
|
||||
assert evaluate_call.json["agentId"] == "my-agent-key"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_evaluate_payload_includes_listed_tool_metadata(self):
|
||||
handler: Final = FakeHandler([_token_response(), _allow_response()])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
schema: Final = {"type": "object", "properties": {"to": {"type": "string"}}, "required": ["to"]}
|
||||
await _run(guardrail, _mcp_data(mcp_tool_description="Send an email", mcp_input_schema=schema))
|
||||
assert handler.calls[1].json["tool"] == {
|
||||
"name": "send_email",
|
||||
"description": "Send an email",
|
||||
"inputSchema": schema,
|
||||
}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("description", "schema"),
|
||||
[(None, None), ("", None), (None, ["not", "a", "schema"]), (42, "type: object")],
|
||||
)
|
||||
async def test_evaluate_payload_omits_missing_or_malformed_tool_metadata(self, description, schema):
|
||||
handler: Final = FakeHandler([_token_response(), _allow_response()])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
await _run(guardrail, _mcp_data(mcp_tool_description=description, mcp_input_schema=schema))
|
||||
assert handler.calls[1].json["tool"] == {"name": "send_email"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_mcp_call_type_skipped(self):
|
||||
handler: Final = FakeHandler([])
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ from typing import Any, Dict, Final, List, Optional
|
|||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
|
|
@ -21,6 +22,7 @@ from litellm.caching.in_memory_cache import InMemoryCache
|
|||
from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
|
||||
PARALLEL_REQUEST_SLOT_TTL_SECONDS,
|
||||
ParallelSlotAcquisition,
|
||||
|
|
@ -39,6 +41,7 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import (
|
|||
from litellm.proxy.utils import InternalUsageCache, ProxyLogging, hash_token
|
||||
from litellm.types.caching import RedisPipelineIncrementOperation
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
from litellm.types.mcp import MCPPreCallRequestObject
|
||||
from litellm.types.utils import (
|
||||
EmbeddingResponse,
|
||||
ModelResponse,
|
||||
|
|
@ -108,6 +111,159 @@ def test_api_key_descriptor_applies_budget_throttle(
|
|||
assert api_key_descriptor["rate_limit"]["tokens_per_unit"] == expected_tpm
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"description", [None, "Gateway metadata, not caller input. " * 100], ids=["unlisted", "listed"]
|
||||
)
|
||||
@pytest.mark.parametrize("arguments_rewritten", [False, True])
|
||||
async def test_mcp_description_does_not_change_admission_or_reserved_tokens(
|
||||
description: str | None, arguments_rewritten: bool, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
cache: Final = DualCache()
|
||||
handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache))
|
||||
logger: Final = ProxyLogging(user_api_key_cache=UserApiKeyCache())
|
||||
schema: Final = {"type": "object", "properties": {"q": {"type": "string", "description": "Schema text " * 100}}}
|
||||
request: Final = MCPPreCallRequestObject(
|
||||
tool_name="echo", arguments={"q": "hello"}, tool_description=description, tool_input_schema=schema
|
||||
)
|
||||
data: Final = TypeAdapter(dict[str, object]).validate_python(logger._convert_mcp_to_llm_format(request, {}))
|
||||
messages: Final = data["messages"]
|
||||
caller: Final = UserAPIKeyAuth(api_key=hash_token("sk-mcp-description-reservation"), tpm_limit=64)
|
||||
|
||||
if arguments_rewritten:
|
||||
data["mcp_arguments"] = {"q": "Transformed arguments " * 100}
|
||||
monkeypatch.setattr(litellm, "callbacks", [handler])
|
||||
await logger.pre_call_hook(user_api_key_dict=caller, data=data, call_type="call_mcp_tool")
|
||||
|
||||
stash: Final = get_request_stash()
|
||||
assert stash is not None
|
||||
assert stash.reserved_tokens == 25
|
||||
assert (
|
||||
await cache.async_get_cache(
|
||||
key=handler.create_rate_limit_keys("api_key", caller.api_key, "tokens"), local_only=True
|
||||
)
|
||||
== 25
|
||||
)
|
||||
assert data["messages"] is messages
|
||||
assert data.get("mcp_tool_description") == description
|
||||
assert data["mcp_input_schema"] == schema
|
||||
assert messages == [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Tool: echo\nArguments: {'q': 'hello'}",
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"description", [None, "Gateway metadata, not caller input. " * 100], ids=["unlisted", "listed"]
|
||||
)
|
||||
@pytest.mark.parametrize("itpm_limit,otpm_limit", [(64, 4096), (4096, 64), (4096, 4096)])
|
||||
@pytest.mark.parametrize("arguments_rewritten", [False, True])
|
||||
async def test_mcp_description_preserves_project_input_and_output_reservations(
|
||||
description: str | None, itpm_limit: int, otpm_limit: int,
|
||||
arguments_rewritten: bool, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
cache: Final = DualCache()
|
||||
handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache))
|
||||
logger: Final = ProxyLogging(user_api_key_cache=UserApiKeyCache())
|
||||
schema: Final = {"type": "object", "properties": {"q": {"type": "string", "description": "Schema text " * 100}}}
|
||||
request: Final = MCPPreCallRequestObject(
|
||||
tool_name="echo", arguments={"q": "hello"}, tool_description=description, tool_input_schema=schema
|
||||
)
|
||||
data: Final = TypeAdapter(dict[str, object]).validate_python(logger._convert_mcp_to_llm_format(request, {}))
|
||||
messages: Final = data["messages"]
|
||||
base_data: Final[dict[str, object]] = {
|
||||
"messages": [{"role": "user", "content": "Tool: echo\nArguments: {'q': 'hello'}"}]
|
||||
}
|
||||
expected_input: Final = handler._estimate_precise_input_tokens(base_data, "mcp-tool-call", "call_mcp_tool")
|
||||
expected_output: Final = handler.no_max_tokens_output_floor(otpm_limit)
|
||||
expected_combined: Final = handler._estimate_tokens_for_request(
|
||||
base_data, min_configured_tpm_limit=4096, call_type="call_mcp_tool"
|
||||
)
|
||||
caller: Final = UserAPIKeyAuth(
|
||||
api_key=hash_token("sk-mcp-project-reservation"),
|
||||
tpm_limit=4096,
|
||||
project_id="mcp-project-reservation",
|
||||
project_metadata={
|
||||
"model_itpm_limit": {"mcp-tool-call": itpm_limit},
|
||||
"model_otpm_limit": {"mcp-tool-call": otpm_limit},
|
||||
},
|
||||
)
|
||||
|
||||
if arguments_rewritten:
|
||||
data["mcp_arguments"] = {"q": "Transformed arguments " * 100}
|
||||
monkeypatch.setattr(litellm, "callbacks", [handler])
|
||||
await logger.pre_call_hook(user_api_key_dict=caller, data=data, call_type="call_mcp_tool")
|
||||
|
||||
stash: Final = get_request_stash()
|
||||
assert stash is not None
|
||||
assert (stash.reserved_tokens, stash.itpm_reserved_tokens, stash.otpm_reserved_tokens) == (
|
||||
expected_combined,
|
||||
expected_input,
|
||||
expected_output,
|
||||
)
|
||||
assert (
|
||||
await cache.async_get_cache(
|
||||
key=handler.create_rate_limit_keys(
|
||||
"model_per_project_itpm", f"{caller.project_id}:mcp-tool-call", "tokens"
|
||||
),
|
||||
local_only=True,
|
||||
)
|
||||
== expected_input
|
||||
)
|
||||
assert (
|
||||
await cache.async_get_cache(
|
||||
key=handler.create_rate_limit_keys(
|
||||
"model_per_project_otpm", f"{caller.project_id}:mcp-tool-call", "tokens"
|
||||
),
|
||||
local_only=True,
|
||||
)
|
||||
== expected_output
|
||||
)
|
||||
assert data["messages"] is messages
|
||||
assert data.get("mcp_tool_description") == description
|
||||
assert data["mcp_input_schema"] == schema
|
||||
assert messages == [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Tool: echo\nArguments: {'q': 'hello'}",
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
def test_llm_tpm_estimation_still_counts_messages_with_mcp_metadata() -> None:
|
||||
handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache()))
|
||||
data: Final[dict[str, object]] = {
|
||||
"messages": [{"role": "user", "content": "x" * 400}],
|
||||
"max_tokens": 1,
|
||||
"mcp_tool_name": "echo",
|
||||
"mcp_arguments": {},
|
||||
}
|
||||
assert handler._estimate_tokens_for_request(data, call_type="acompletion") == 101
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unconverted_mcp_request_keeps_its_reservation() -> None:
|
||||
cache: Final = DualCache()
|
||||
handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache))
|
||||
caller: Final = UserAPIKeyAuth(api_key=hash_token("sk-raw-mcp-request"), tpm_limit=64)
|
||||
data: Final[dict[str, object]] = {"name": "echo", "arguments": {"q": "hello"}, "server_id": "fixture"}
|
||||
|
||||
await handler.async_pre_call_hook(user_api_key_dict=caller, cache=cache, data=data, call_type="call_mcp_tool")
|
||||
|
||||
stash: Final = get_request_stash()
|
||||
assert stash is not None
|
||||
assert stash.reserved_tokens == 16
|
||||
assert (
|
||||
await cache.async_get_cache(
|
||||
key=handler.create_rate_limit_keys("api_key", caller.api_key, "tokens"), local_only=True
|
||||
)
|
||||
== 16
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.flaky(reruns=3)
|
||||
@pytest.mark.asyncio
|
||||
async def test_sliding_window_rate_limit_v3(monkeypatch, time_controller):
|
||||
|
|
|
|||
|
|
@ -403,6 +403,26 @@ def test_create_mcp_request_object_from_kwargs_full(proxy_logging, make_user_api
|
|||
assert snapshot == {"tool_name": "calc", "arguments": {"x": 1}, "server_name": "math", "auth_user_id": "u-1"}
|
||||
|
||||
|
||||
def test_mcp_tool_metadata_flows_from_kwargs_to_synthetic_data(proxy_logging):
|
||||
schema = {"type": "object", "properties": {"x": {"type": "integer"}}}
|
||||
obj = proxy_logging._create_mcp_request_object_from_kwargs(
|
||||
kwargs={
|
||||
"name": "calc",
|
||||
"arguments": {"x": 1},
|
||||
"tool_description": "Adds numbers",
|
||||
"tool_input_schema": schema,
|
||||
}
|
||||
)
|
||||
out = proxy_logging._convert_mcp_to_llm_format(request_obj=obj, kwargs={})
|
||||
assert (out["mcp_tool_description"], out["mcp_input_schema"]) == ("Adds numbers", schema)
|
||||
|
||||
|
||||
def test_mcp_tool_metadata_absent_when_tool_was_never_listed(proxy_logging):
|
||||
obj = proxy_logging._create_mcp_request_object_from_kwargs(kwargs={"name": "calc", "arguments": {}})
|
||||
out = proxy_logging._convert_mcp_to_llm_format(request_obj=obj, kwargs={})
|
||||
assert "mcp_tool_description" not in out and "mcp_input_schema" not in out
|
||||
|
||||
|
||||
def test_create_mcp_request_object_from_kwargs_empty(proxy_logging):
|
||||
obj = proxy_logging._create_mcp_request_object_from_kwargs(kwargs={})
|
||||
snapshot = {
|
||||
|
|
|
|||
|
|
@ -1,10 +1,11 @@
|
|||
import asyncio
|
||||
import importlib
|
||||
import subprocess
|
||||
import sys
|
||||
import textwrap
|
||||
import types
|
||||
from typing import Any, Final, cast
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
from typing import Any, Final, Literal, cast
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -12,15 +13,26 @@ from mcp.types import CallToolResult, TextContent
|
|||
from mcp.types import Tool as MCPTool
|
||||
from openai.types.responses.tool_param import Mcp
|
||||
|
||||
import litellm
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.proxy._experimental.mcp_server import operations as mcp_operations
|
||||
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller, MCPServerManager
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
from litellm.responses import main as responses_main
|
||||
from litellm.responses.mcp import litellm_proxy_mcp_handler as mcp_handler_module
|
||||
from litellm.responses.mcp.litellm_proxy_mcp_handler import (
|
||||
LiteLLM_Proxy_MCP_Handler,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
from litellm.types.mcp import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
from litellm.types.responses.main import OutputFunctionToolCall
|
||||
from litellm.types.utils import ModelResponse
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs, ModelResponse
|
||||
|
||||
|
||||
class _DummyMCPResult:
|
||||
|
|
@ -1310,6 +1322,167 @@ async def test_responses_discovery_logs_sanitized_caller_headers(monkeypatch: py
|
|||
assert headers["x-mcp-deepwiki-authorization"] == "upstream-sentinel"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("real_listing", [False, True])
|
||||
@pytest.mark.parametrize(
|
||||
("allowed_tools", "expected_names"),
|
||||
[
|
||||
([], ["responses_slot-echo", "responses_slot-status"]),
|
||||
(["echo"], ["responses_slot-echo"]),
|
||||
(["responses_slot-echo"], ["responses_slot-echo"]),
|
||||
(["absent"], []),
|
||||
],
|
||||
)
|
||||
async def test_bridge_listing_leaves_the_callers_catalog_unchanged(
|
||||
monkeypatch: pytest.MonkeyPatch, allowed_tools: list[str], expected_names: list[str], real_listing: bool
|
||||
) -> None:
|
||||
manager: Final = mcp_operations.global_mcp_server_manager
|
||||
server: Final = MCPServer(
|
||||
server_id="responses-slot", name="responses_slot", alias="responses_slot", transport=MCPTransport.http
|
||||
)
|
||||
user: Final = UserAPIKeyAuth(api_key="sk-responses-slot", user_id="responder")
|
||||
upstream: Final = [
|
||||
MCPTool(name="echo", description="Echo text back", inputSchema={"type": "object"}),
|
||||
MCPTool(name="status", description="Report status", inputSchema={"type": "object"}),
|
||||
MCPTool(name="echo", description="Duplicate echo", inputSchema={"type": "object", "properties": {}}),
|
||||
]
|
||||
fake_manager: Final = types.SimpleNamespace(
|
||||
get_registry=MagicMock(return_value={}),
|
||||
get_allowed_mcp_servers=AsyncMock(return_value=[]),
|
||||
get_mcp_servers_from_ids=MagicMock(return_value=[]),
|
||||
get_mcp_server_by_name=MagicMock(return_value=None),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
|
||||
fake_manager,
|
||||
)
|
||||
with (
|
||||
patch.dict(manager.tool_name_to_mcp_server_name_mapping),
|
||||
patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])),
|
||||
patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())),
|
||||
patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)),
|
||||
):
|
||||
try:
|
||||
if real_listing:
|
||||
await manager._get_tools_from_server(server, user_api_key_auth=user, record_listing=True)
|
||||
caller: Final = ListedToolsCaller(user_api_key_auth=user)
|
||||
before: Final = {
|
||||
tool.name: (listed.description, listed.input_schema)
|
||||
for tool in upstream
|
||||
if (listed := manager.get_listed_tool(server, tool.name, caller)) is not None
|
||||
}
|
||||
assert bool(before) is real_listing
|
||||
tools, _server_names = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform(
|
||||
user_api_key_auth=user,
|
||||
mcp_tools_with_litellm_proxy=[
|
||||
{
|
||||
"type": "mcp",
|
||||
"server_url": "litellm_proxy/mcp/responses-slot",
|
||||
"allowed_tools": allowed_tools,
|
||||
}
|
||||
],
|
||||
)
|
||||
recorded: Final = {
|
||||
tool.name: (listed.description, listed.input_schema)
|
||||
for tool in upstream
|
||||
if (listed := manager.get_listed_tool(server, tool.name, caller)) is not None
|
||||
}
|
||||
assert recorded == before
|
||||
assert (
|
||||
manager.get_listed_tool(
|
||||
server, "echo", ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-other-caller"))
|
||||
)
|
||||
is None
|
||||
)
|
||||
finally:
|
||||
manager._drop_listed_tools(server.server_id)
|
||||
|
||||
assert [tool.name for tool in tools] == expected_names
|
||||
|
||||
|
||||
class _BridgeMetadataGuardrail(CustomGuardrail):
|
||||
def __init__(self) -> None:
|
||||
super().__init__(guardrail_name="bridge-metadata", event_hook=GuardrailEventHooks.pre_mcp_call, default_on=True)
|
||||
self.calls: tuple[tuple[object, object], ...] = ()
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict[str, object],
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Logging | None = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
if request_data.get("mcp_arguments") == {"probe": "bridge"}:
|
||||
self.calls += ((request_data.get("mcp_tool_description"), request_data.get("mcp_input_schema")),)
|
||||
return inputs
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_bridge_calls_use_their_own_served_metadata(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
manager: Final = MCPServerManager()
|
||||
server: Final = MCPServer(server_id="bridge", name="bridge", transport=MCPTransport.http, url="http://upstream")
|
||||
manager.registry = {server.server_id: server}
|
||||
user: Final = UserAPIKeyAuth(api_key="sk-bridge", user_id="bridge-user")
|
||||
upstream: Final = [
|
||||
MCPTool(
|
||||
name="echo",
|
||||
description="Echo text",
|
||||
inputSchema={"type": "object", "properties": {"text": {"type": "string"}}},
|
||||
),
|
||||
MCPTool(name="status", description="Read status", inputSchema={"type": "object"}),
|
||||
]
|
||||
client: Final = AsyncMock()
|
||||
client.call_tool.return_value = CallToolResult(content=[TextContent(type="text", text="ok")])
|
||||
manager._create_mcp_client = AsyncMock(return_value=client)
|
||||
manager._fetch_tools_with_timeout = AsyncMock(return_value=upstream)
|
||||
guardrail: Final = _BridgeMetadataGuardrail()
|
||||
logger: Final = ProxyLogging(user_api_key_cache=DualCache())
|
||||
monkeypatch.setattr(litellm, "callbacks", [guardrail])
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", logger)
|
||||
monkeypatch.setattr(mcp_operations, "global_mcp_server_manager", manager)
|
||||
monkeypatch.setattr(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server]))
|
||||
monkeypatch.setattr("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", manager)
|
||||
first_listed: Final = asyncio.Event()
|
||||
second_listed: Final = asyncio.Event()
|
||||
|
||||
async def bridge(name: str, first: bool) -> None:
|
||||
if not first:
|
||||
await first_listed.wait()
|
||||
tools, server_map = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform(
|
||||
user_api_key_auth=user,
|
||||
mcp_tools_with_litellm_proxy=[
|
||||
{"type": "mcp", "server_url": "litellm_proxy/mcp/bridge", "allowed_tools": [name]}
|
||||
],
|
||||
)
|
||||
(first_listed if first else second_listed).set()
|
||||
await second_listed.wait()
|
||||
result: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
|
||||
tool_server_map=server_map,
|
||||
tool_calls=[
|
||||
{"type": "function_call", "name": f"bridge-{name}", "arguments": '{"probe":"bridge"}', "call_id": name}
|
||||
],
|
||||
user_api_key_auth=user,
|
||||
served_tools=tools,
|
||||
)
|
||||
assert [entry["result"] for entry in result] == ["ok"]
|
||||
|
||||
try:
|
||||
await asyncio.gather(bridge("echo", True), bridge("status", False))
|
||||
assert sorted(guardrail.calls, key=str) == sorted(
|
||||
((tool.description, tool.input_schema) for tool in upstream), key=str
|
||||
)
|
||||
await manager.call_tool("bridge", "echo", {"probe": "bridge"}, user_api_key_auth=user, proxy_logging_obj=logger)
|
||||
assert guardrail.calls[-1] == (None, None)
|
||||
await manager._get_tools_from_server(
|
||||
server, user_api_key_auth=user, proxy_logging_obj=logger, record_listing=True
|
||||
)
|
||||
await manager.call_tool("bridge", "echo", {"probe": "bridge"}, user_api_key_auth=user, proxy_logging_obj=logger)
|
||||
assert guardrail.calls[-1] == (upstream[0].description, upstream[0].input_schema)
|
||||
finally:
|
||||
manager._drop_listed_tools(server.server_id)
|
||||
ProxyLogging._callback_capabilities_cache.clear()
|
||||
|
||||
|
||||
def _toolset_gateway_manager(toolset_id: str, server_id: str) -> types.SimpleNamespace:
|
||||
return types.SimpleNamespace(
|
||||
get_registry=MagicMock(return_value={}),
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue