feat(mcp): hand listed-tool metadata to pre-call hooks with per-caller catalog identity

Track the tools each MCP server listed per caller identity so pre_mcp_call and during_mcp_call hooks receive the tool description and input schema the client saw. Servers with no caller-dependent inputs share one slot; user identity, forwarded headers, stdio env, relayed bearers, and server-specific auth get their own. Local registry and OpenAPI paths pass the registered metadata and admin description overrides. The Agent 365 guardrail reads the new fields into its evaluate payload.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-15 00:42:16 +00:00
parent 90873c46de
commit 3c13e3b457
11 changed files with 876 additions and 59 deletions

View file

@ -250,6 +250,21 @@ _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]]
_NO_LISTED_TOOLS: Final[_ListedToolsByCaller] = MappingProxyType({})
_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
@ -1128,6 +1143,25 @@ 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 _format_byok_openapi_auth_header(mcp_server: MCPServer, mcp_auth_header: str) -> str:
"""Format a raw BYOK credential for OpenAPI tool ``Authorization`` injection.
@ -1931,6 +1965,7 @@ class MCPServerManager:
"gmail_send_email": "zapier_mcp_server",
}
"""
self._listed_tools_by_server_id: dict[str, _ListedToolsByCaller] = {} # mutable-ok: refreshed per tools/list
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
@ -2629,7 +2664,7 @@ class MCPServerManager:
self._assign_unique_short_prefix(new_server)
_warn_legacy_delegate_auth_if_applicable(new_server, source="config")
_warn_config_id_jag_server_outruns_sso(new_server)
self._invalidate_discovery_lists(server_id)
self._invalidate_server_definition_caches(server_id)
self.config_mcp_servers[server_id] = new_server
self._set_oauth_discovery_deferred(
server_id,
@ -2833,7 +2868,7 @@ class MCPServerManager:
global_mcp_tool_registry,
)
self._invalidate_discovery_lists(server.server_id)
self._invalidate_server_definition_caches(server.server_id)
prefix_root: Final = normalize_server_name(get_server_prefix(server))
if server.spec_path and prefix_root:
openapi_key_prefix: Final = prefix_root + MCP_TOOL_PREFIX_SEPARATOR
@ -3240,7 +3275,7 @@ class MCPServerManager:
# env_vars_are_encrypted=False.
new_server: Final = await self.build_mcp_server_from_table(mcp_server, env_vars_are_encrypted=False)
self._assign_unique_short_prefix(new_server)
self._invalidate_discovery_lists(mcp_server.server_id)
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)
self.prime_oauth_metadata_discovery(new_server)
@ -3277,7 +3312,7 @@ class MCPServerManager:
previous_server=self.registry[mcp_server.server_id],
)
self._assign_unique_short_prefix(new_server)
self._invalidate_discovery_lists(mcp_server.server_id)
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)
self.prime_oauth_metadata_discovery(new_server)
@ -3732,19 +3767,7 @@ 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(
@ -3830,7 +3853,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."""
@ -4379,6 +4402,12 @@ 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=mcp_auth_header,
raw_headers=raw_headers,
oauth2_headers=oauth2_headers,
)
try:
# Tool *listing* must not be blocked by missing per-user env vars —
@ -4457,29 +4486,25 @@ class MCPServerManager:
if server.spec_path:
# OpenAPI tools were stored in the registry under the prefix
# active at registration time — fetch by that same prefix.
_tools: Final = global_mcp_tool_registry.list_tools(tool_prefix=get_server_prefix(server))
registry_prefix: Final = normalize_server_name(get_server_prefix(server)) + MCP_TOOL_PREFIX_SEPARATOR
_tools: Final = global_mcp_tool_registry.list_tools(tool_prefix=registry_prefix)
tools = global_mcp_tool_registry.convert_tools_to_mcp_sdk_tool_type(_tools)
# OpenAPI tools are stored in the registry with their prefix already
# 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".
if not add_prefix:
prefix: Final = get_server_prefix(server)
sep: Final = MCP_TOOL_PREFIX_SEPARATOR
tools = [
(
t.model_copy(update={"name": t.name[len(prefix) + len(sep) :]})
if t.name.startswith(f"{prefix}{sep}")
else t
)
for t in tools
]
return tools
unprefixed_tools: Final = [ # mutable-ok: returned through the list[MCPTool] listing contract
t.model_copy(update=MappingProxyType({"name": t.name[len(registry_prefix) :]})) for t in tools
]
self._record_listed_tools(server, unprefixed_tools, listed_caller)
return tools if add_prefix else unprefixed_tools
else:
tools = await self._fetch_tools_with_timeout(client, server.name)
self._remember_upstream_initialize_instructions(server, client)
prefixed_or_original_tools: Final = self._create_prefixed_tools(tools, server, add_prefix=add_prefix)
prefixed_or_original_tools: Final = self._create_prefixed_tools(
tools, server, add_prefix=add_prefix, caller=listed_caller
)
return prefixed_or_original_tools
@ -4523,6 +4548,86 @@ class MCPServerManager:
self._resource_discovery_cache.invalidate(server_id)
self._template_discovery_cache.invalidate(server_id)
def _invalidate_server_definition_caches(self, server_id: str) -> None:
self._invalidate_discovery_lists(server_id)
self._listed_tools_by_server_id.pop(server_id, None)
def _discovers_per_caller(self, server: MCPServer) -> bool:
return (
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)
or self._signs_caller_identity_upstream(server)
)
@staticmethod
def _signs_caller_identity_upstream(server: MCPServer) -> bool:
"""Whether MCPJWTSigner mints a per-caller ``Authorization`` for ``server``, so the upstream may
tailor its catalog to the caller even though the server itself is configured as shared."""
from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import ( # noqa: PLC0415 # lazy: guardrail package imports the proxy server
get_mcp_jwt_signer,
)
if get_mcp_jwt_signer() is None:
return False
return not any(k.lower() == "authorization" for k in (server.static_headers or {}))
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 upstream catalog.
Forwarded headers, header-driven stdio env, a relayed caller bearer, and the
server-specific auth header all reach upstream, so two callers differing in any of
them may be shown different tools. Shared servers with none of those stay on the
shared (``None``) slot. OpenAPI servers list from the process-wide registry.
"""
if server.spec_path or caller is None:
return None
auth: Final = caller.user_api_key_auth
identity: Final = (
(auth.user_id, auth.api_key) if auth is not None and self._discovers_per_caller(server) else None
)
forwarded: Final = self._forwarded_header_values(server, caller.raw_headers)
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
relayed_bearer: Final = (
self._extract_subject_token(caller.oauth2_headers, caller.raw_headers, auth)
if server.is_client_forwarded_token
else None
)
inputs: Final = (identity, caller.mcp_auth_header, forwarded, stdio_env, relayed_bearer)
if not any(inputs):
return None
material: Final = json.dumps(inputs, 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
) -> None:
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, _NO_LISTED_TOOLS)
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,
@ -4533,12 +4638,7 @@ class MCPServerManager:
subject_token: str | None,
credential_fingerprint: str | None = None,
) -> _DiscoveryKey:
per_user: Final = (
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)
)
per_user: Final = self._discovers_per_caller(server)
if not (per_user or mcp_auth_header or extra_headers or stdio_env or subject_token):
return server.server_id, None
identity: Final = (
@ -5368,7 +5468,13 @@ class MCPServerManager:
"attempts; the 3-character prefix space is too crowded."
)
def _create_prefixed_tools(self, tools: list[MCPTool], server: MCPServer, add_prefix: bool = True) -> list[MCPTool]:
def _create_prefixed_tools(
self,
tools: list[MCPTool],
server: MCPServer,
add_prefix: bool = True,
caller: ListedToolsCaller | None = None,
) -> list[MCPTool]:
"""
Create prefixed tools and update tool mapping.
@ -5393,9 +5499,21 @@ class MCPServerManager:
for spelling in iter_known_tool_name_spellings(original_name, server):
self.tool_name_to_mcp_server_name_mapping[spelling] = prefix
self._record_listed_tools(server, tools, caller)
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, _NO_LISTED_TOOLS).get(identity)
if not listed:
return None
tool: Final = listed.get(name) or listed.get(strip_known_server_prefix(name, server))
if tool is None:
return None
description: Final = (server.tool_name_to_description or {}).get(tool.name)
return tool if description is None else tool.model_copy(update={"description": description})
def _create_prefixed_prompts(
self, prompts: Sequence[Prompt], server: MCPServer, add_prefix: bool = True
) -> list[Prompt]:
@ -5629,6 +5747,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.
@ -5642,6 +5761,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
@ -5696,6 +5818,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
@ -5751,6 +5875,7 @@ class MCPServerManager:
start_time: datetime.datetime,
litellm_logging_obj: "LiteLLMLoggingObj | None" = None,
guardrail_context: Mapping[str, object] | None = None,
tool: MCPTool | None = None,
):
"""Create and return a during hook task for MCP tool calls.
@ -5765,6 +5890,8 @@ class MCPServerManager:
tool_name=name,
arguments=arguments,
server_name=server_name_from_prefix,
tool_description=tool.description if tool is not None else None,
tool_input_schema=tool.input_schema if tool is not None else None,
start_time=start_time.timestamp() if start_time else None,
hidden_params=HiddenParams(),
)
@ -5911,21 +6038,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
@ -6376,6 +6489,12 @@ class MCPServerManager:
user_api_key_auth,
mcp_auth_header,
)
listed_caller: Final = ListedToolsCaller(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=_server_auth_header_for(mcp_server, mcp_server_auth_headers, mcp_auth_header),
raw_headers=raw_headers,
oauth2_headers=oauth2_headers,
)
#########################################################
# Pre MCP Tool Call Hook
@ -6392,6 +6511,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 "arguments" in hook_result:
arguments = hook_result["arguments"]
@ -6408,6 +6528,7 @@ class MCPServerManager:
start_time=start_time,
litellm_logging_obj=litellm_logging_obj,
guardrail_context=guardrail_context,
tool=self.get_listed_tool(mcp_server, name, listed_caller),
)
tasks.append(during_hook_task)
@ -6669,7 +6790,7 @@ class MCPServerManager:
for server_id in previous_registry.keys() | registered_registry.keys():
if previous_registry.get(server_id) != registered_registry.get(server_id):
self._invalidate_discovery_lists(server_id)
self._invalidate_server_definition_caches(server_id)
self.registry = registered_registry
_warn_on_shared_identifier_prefixes(registered_registry.values())
# A discovery task may have published into ``previous_registry`` while

View file

@ -139,6 +139,7 @@ from litellm.types.mcp import (
without_header,
)
from litellm.types.mcp_server.mcp_server_manager import MCPInfo, MCPServer
from litellm.types.mcp_server.tool_registry import MCPTool as RegisteredTool
from litellm.types.utils import CallTypes, StandardLoggingMCPToolCall
from litellm.utils import Rules, client, function_setup
@ -1610,6 +1611,12 @@ async def _list_mcp_resource_templates(
return managed_resource_templates
def _registered_tool_metadata(name: str, registered: RegisteredTool, server: MCPServer) -> MCPTool:
overrides: Final = server.tool_name_to_description
description: Final = overrides.get(name, registered.description) if overrides else registered.description
return MCPTool(name=name, description=description, input_schema=registered.input_schema)
def _resolve_display_name_to_original(
name: str,
allowed_mcp_servers: list[MCPServer],
@ -2079,6 +2086,7 @@ async def _execute_mcp_tool(
raw_headers=raw_headers,
litellm_logging_obj=litellm_logging_obj,
guardrail_context=guardrail_context,
tool=_registered_tool_metadata(original_tool_name, local_tool, mcp_server),
)
# `pre_call_tool_check` may return guardrail-modified
# arguments; honor them on the local path too.
@ -2147,7 +2155,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
@ -2189,6 +2198,7 @@ async def _execute_mcp_tool(
raw_headers=raw_headers,
litellm_logging_obj=litellm_logging_obj,
guardrail_context=guardrail_context,
tool=_registered_tool_metadata(original_tool_name, registered_local_tool, prefix_server),
)
if "arguments" in hook_result:
arguments = hook_result["arguments"] # pyright: ignore[reportAny] # hook returns untyped args

View file

@ -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
@ -60,6 +60,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
@ -81,6 +82,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``."""
@ -99,6 +107,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]
@ -397,8 +413,14 @@ class Agent365Guardrail(CustomGuardrail):
arguments: Final = data.get("mcp_arguments")
server_name: Final = str(data.get("mcp_server_name") or "litellm")
agent_id: Final = self.agent_id or 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_tool_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),
}

View file

@ -1498,6 +1498,8 @@ class ProxyLogging:
"user_api_key_request_route": kwargs.get("user_api_key_request_route"),
"mcp_tool_name": request_obj.tool_name, # Keep original for reference
"mcp_arguments": request_obj.arguments, # Keep original for reference
"mcp_tool_description": request_obj.tool_description,
"mcp_tool_input_schema": request_obj.tool_input_schema,
# Surface the per-MCP-server rate-limit identity so the
# ParallelRequestLimiterV3 hook can apply mcp_rpm_limit on the
# synthetic call_mcp_tool payload (otherwise a key with
@ -1728,6 +1730,8 @@ class ProxyLogging:
tool_name=kwargs.get("name", ""),
arguments=kwargs.get("arguments", {}),
server_name=kwargs.get("server_name"),
tool_description=kwargs.get("tool_description"),
tool_input_schema=kwargs.get("tool_input_schema"),
user_api_key_auth=user_api_key_auth_dict,
hidden_params=HiddenParams(),
)

View file

@ -422,6 +422,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()
@ -445,6 +447,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()

View file

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

View file

@ -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
@ -6881,7 +6882,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
@ -7993,9 +7993,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(
@ -8046,6 +8049,141 @@ 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_registered_tool_metadata_to_pre_call_hooks():
"""OpenAPI-generated tools dispatch through the local registry, so the pre-call hooks must get the
registered description and input schema on that path too, even when no tools/list ran first."""
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": {"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)
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-list_pets",
arguments={"limit": 10},
allowed_mcp_servers=[petstore],
start_time=datetime.now(),
user_api_key_auth=UserAPIKeyAuth(api_key="sk-user", user_id="alice"),
)
finally:
mcp_module.global_mcp_tool_registry.unregister_tools_with_prefix("petstore-")
handed_tool = pre_call_tool_check.call_args.kwargs["tool"]
assert (handed_tool.name, handed_tool.description, handed_tool.input_schema) == (
"list_pets",
"List the pets",
schema,
)
@pytest.mark.asyncio
async def test_execute_mcp_tool_hands_openapi_hooks_the_admin_description_clients_saw():
"""tools/list shows the admin's tool_name_to_description wording, so the local-registry call path
must hand the pre-call hooks that same wording rather than the generated 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": "ADMIN DESC"},
)
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=schema, handler=lambda petId: "ok"
)
manager = mcp_module.global_mcp_server_manager
manager._listed_tools_by_server_id.pop(petstore.server_id, None)
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=UserAPIKeyAuth(api_key="sk-user", user_id="alice"),
)
finally:
mcp_module.global_mcp_tool_registry.unregister_tools_with_prefix("petstore-")
handed_tool = pre_call_tool_check.call_args.kwargs["tool"]
assert (handed_tool.description, handed_tool.input_schema) == ("ADMIN DESC", schema)
@pytest.mark.asyncio
async def test_execute_mcp_tool_hands_hooks_the_metadata_of_the_operation_it_runs_when_names_collide():
"""An OpenAPI operation whose name starts with its own server prefix must not be reported to the
pre-call hooks with the metadata of the shorter operation, since that is not the one that runs."""
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-")
handed_tool = pre_call_tool_check.call_args.kwargs["tool"]
assert (handed_tool.description, handed_tool.input_schema) == (
"long",
{"type": "object", "properties": {"petId": {"type": "integer"}}},
)
assert result.content[0].text == "long"
@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.

View file

@ -41,6 +41,7 @@ from pydantic import AnyUrl, TypeAdapter
from litellm.constants import MCP_METADATA_TIMEOUT
from litellm.proxy._experimental.mcp_server.tool_outcome import TextResult
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
ListedToolsCaller,
MCPServerManager,
_deserialize_json_dict,
_flow_endpoints_missing,
@ -6772,6 +6773,465 @@ class TestMCPServerManager:
# Verify the MCP client call was awaited exactly once
assert mock_client.call_tool.await_count == 1
@staticmethod
def _manager_ready_for_call_tool(listed_tools: list[MCPTool]) -> tuple[MCPServerManager, MagicMock]:
from mcp.types import CallToolResult
manager = MCPServerManager()
server = MCPServer(
server_id="test-server",
name="test-server",
transport=MCPTransport.http,
url="http://test-server.com",
)
manager.registry = {"test-server": server}
manager.tool_name_to_mcp_server_name_mapping["test_tool"] = "test-server"
manager.tool_name_to_mcp_server_name_mapping["test-server-test_tool"] = "test-server"
manager._create_prefixed_tools(listed_tools, server)
mock_client = AsyncMock()
mock_client.call_tool.return_value = MagicMock(spec=CallToolResult, content=[], isError=False)
manager._create_mcp_client = AsyncMock(return_value=mock_client)
proxy_logging_obj = MagicMock()
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={})
proxy_logging_obj.pre_call_hook = AsyncMock(return_value={})
proxy_logging_obj.during_call_hook = AsyncMock(return_value=None)
return manager, proxy_logging_obj
@staticmethod
def _unrestricted_auth() -> MagicMock:
user_api_key_auth = MagicMock()
user_api_key_auth.object_permission = None
user_api_key_auth.object_permission_id = None
return user_api_key_auth
@pytest.mark.asyncio
async def test_call_tool_hands_listed_tool_description_and_schema_to_pre_call_hooks(self):
schema = {"type": "object", "properties": {"param": {"type": "string"}}, "required": ["param"]}
listed = [MCPTool(name="test_tool", description="Runs the test tool", inputSchema=schema)]
manager, proxy_logging_obj = self._manager_ready_for_call_tool(listed)
await manager.call_tool(
server_name="test-server",
name="test_tool",
arguments={"param": "value"},
user_api_key_auth=self._unrestricted_auth(),
proxy_logging_obj=proxy_logging_obj,
)
hook_kwargs = proxy_logging_obj._create_mcp_request_object_from_kwargs.call_args.args[0]
assert (hook_kwargs["tool_description"], hook_kwargs["tool_input_schema"]) == ("Runs the test tool", schema)
@pytest.mark.asyncio
async def test_call_tool_hands_listed_tool_metadata_to_during_call_hooks_through_real_conversion(self):
schema = {"type": "object", "properties": {"param": {"type": "string"}}}
listed = [MCPTool(name="test_tool", description="Runs the test tool", inputSchema=schema)]
manager, _ = self._manager_ready_for_call_tool(listed)
proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging_obj.pre_call_hook = AsyncMock(return_value={})
proxy_logging_obj.during_call_hook = AsyncMock(return_value=None)
await manager.call_tool(
server_name="test-server",
name="test_tool",
arguments={"param": "value"},
user_api_key_auth=UserAPIKeyAuth(api_key="sk-test"),
proxy_logging_obj=proxy_logging_obj,
)
during_data = proxy_logging_obj.during_call_hook.call_args.kwargs["data"]
assert (during_data["mcp_tool_description"], during_data["mcp_tool_input_schema"]) == (
"Runs the test tool",
schema,
)
@pytest.mark.asyncio
async def test_call_tool_passes_no_tool_metadata_when_tool_was_never_listed(self):
manager, proxy_logging_obj = self._manager_ready_for_call_tool(
[MCPTool(name="other_tool", description="Unrelated", inputSchema={"type": "object"})]
)
await manager.call_tool(
server_name="test-server",
name="test_tool",
arguments={"param": "value"},
user_api_key_auth=self._unrestricted_auth(),
proxy_logging_obj=proxy_logging_obj,
)
hook_kwargs = proxy_logging_obj._create_mcp_request_object_from_kwargs.call_args.args[0]
assert (hook_kwargs["tool_description"], hook_kwargs["tool_input_schema"]) == (None, None)
def test_get_listed_tool_resolves_prefixed_name_and_latest_listing(self):
manager = MCPServerManager()
server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv")
manager._create_prefixed_tools([MCPTool(name="echo", description="v1", inputSchema={})], server)
manager._create_prefixed_tools([MCPTool(name="echo", description="v2", inputSchema={})], server)
by_prefixed_name = manager.get_listed_tool(server, "srv-echo")
assert by_prefixed_name is not None and by_prefixed_name.description == "v2"
assert manager.get_listed_tool(server, "missing") is None
def test_get_listed_tool_uses_admin_description_override_clients_saw(self):
manager = MCPServerManager()
server = MCPServer(
server_id="srv",
name="srv",
transport=MCPTransport.http,
url="http://srv",
tool_name_to_description={"echo": "Admin wording"},
)
schema = {"type": "object", "properties": {"text": {"type": "string"}}}
manager._create_prefixed_tools(
[
MCPTool(name="echo", description="Upstream wording", inputSchema=schema),
MCPTool(name="ping", description="Untouched", inputSchema={}),
],
server,
)
overridden = manager.get_listed_tool(server, "srv-echo")
assert overridden is not None
assert (overridden.name, overridden.description, overridden.input_schema) == ("echo", "Admin wording", schema)
untouched = manager.get_listed_tool(server, "ping")
assert untouched is not None and untouched.description == "Untouched"
def test_server_definition_change_drops_listed_tools(self):
manager = MCPServerManager()
server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv")
other = MCPServer(server_id="other", name="other", transport=MCPTransport.http, url="http://other")
manager._create_prefixed_tools([MCPTool(name="echo", description="old", inputSchema={})], server)
manager._create_prefixed_tools([MCPTool(name="ping", description="kept", inputSchema={})], other)
manager._invalidate_server_definition_caches(server.server_id)
assert manager.get_listed_tool(server, "echo") is None
kept = manager.get_listed_tool(other, "ping")
assert kept is not None and kept.description == "kept"
@pytest.mark.asyncio
async def test_user_oauth_refresh_keeps_listed_tools(self):
"""Tool definitions are server-wide, so one user's re-auth must not blank the metadata other
callers' tool calls hand to pre-call guardrails."""
manager = MCPServerManager()
server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv")
manager._create_prefixed_tools([MCPTool(name="echo", description="shared", inputSchema={})], server)
await manager.invalidate_user_oauth_token_cache("alice", server.server_id)
listed = manager.get_listed_tool(server, "echo")
assert listed is not None and listed.description == "shared"
def test_per_caller_server_keeps_listed_tools_per_identity(self):
manager = MCPServerManager()
server = MCPServer(
server_id="srv",
name="srv",
transport=MCPTransport.http,
url="http://srv",
auth_type=MCPAuth.oauth2_token_exchange,
)
alice = UserAPIKeyAuth(user_id="alice", api_key="hashed-alice")
bob = UserAPIKeyAuth(user_id="bob", api_key="hashed-bob")
alice_schema = {"type": "object", "properties": {"path": {"type": "string"}}}
bob_schema = {"type": "object", "properties": {"path": {"type": "string"}, "site": {"type": "string"}}}
manager._create_prefixed_tools(
[MCPTool(name="read", description="alice view", inputSchema=alice_schema)],
server,
caller=ListedToolsCaller(user_api_key_auth=alice),
)
manager._create_prefixed_tools(
[MCPTool(name="read", description="bob view", inputSchema=bob_schema)],
server,
caller=ListedToolsCaller(user_api_key_auth=bob),
)
alice_tool = manager.get_listed_tool(server, "srv-read", ListedToolsCaller(user_api_key_auth=alice))
bob_tool = manager.get_listed_tool(server, "srv-read", ListedToolsCaller(user_api_key_auth=bob))
assert alice_tool is not None and (alice_tool.description, alice_tool.input_schema) == (
"alice view",
alice_schema,
)
assert bob_tool is not None and (bob_tool.description, bob_tool.input_schema) == ("bob view", bob_schema)
carol = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="carol", api_key="k"))
assert manager.get_listed_tool(server, "srv-read", carol) is None
shared = MCPServer(server_id="shared", name="shared", transport=MCPTransport.http, url="http://shared")
manager._create_prefixed_tools(
[MCPTool(name="echo", description="everyone", inputSchema={})],
shared,
caller=ListedToolsCaller(user_api_key_auth=alice),
)
for_bob = manager.get_listed_tool(shared, "echo", ListedToolsCaller(user_api_key_auth=bob))
assert for_bob is not None and for_bob.description == "everyone"
@pytest.mark.parametrize(
("server_kwargs", "caller_a", "caller_b"),
[
pytest.param(
{"extra_headers": ["X-Workspace"]},
ListedToolsCaller(raw_headers={"x-workspace": "A"}),
ListedToolsCaller(raw_headers={"X-Workspace": "B"}),
id="forwarded-header",
),
pytest.param(
{"auth_type": MCPAuth.true_passthrough},
ListedToolsCaller(raw_headers={"authorization": "Bearer upstream-a"}),
ListedToolsCaller(raw_headers={"authorization": "Bearer upstream-b"}),
id="anonymous-passthrough-bearer",
),
pytest.param(
{"auth_type": MCPAuth.bearer_token},
ListedToolsCaller(mcp_auth_header="byok-a"),
ListedToolsCaller(mcp_auth_header="byok-b"),
id="per-server-auth-header",
),
pytest.param(
{"transport": MCPTransport.stdio, "command": "srv", "env": {"WS": "${X-WS}"}},
ListedToolsCaller(raw_headers={"X-WS": "A"}),
ListedToolsCaller(raw_headers={"X-WS": "B"}),
id="header-driven-stdio-env",
),
],
)
def test_upstream_identity_inputs_keep_listed_tools_apart(self, server_kwargs, caller_a, caller_b):
"""Whatever reaches upstream and can change its catalog must also split the listed-tool cache."""
manager = MCPServerManager()
server = MCPServer(
**{"server_id": "srv", "name": "srv", "transport": MCPTransport.http, "url": "http://srv", **server_kwargs}
)
manager._create_prefixed_tools(
[MCPTool(name="turn", description="Catalog A", inputSchema={})], server, caller=caller_a
)
manager._create_prefixed_tools(
[MCPTool(name="turn", description="Catalog B", inputSchema={})], server, caller=caller_b
)
for_a = manager.get_listed_tool(server, "srv-turn", caller_a)
for_b = manager.get_listed_tool(server, "srv-turn", caller_b)
assert for_a is not None and for_a.description == "Catalog A"
assert for_b is not None and for_b.description == "Catalog B"
assert manager.get_listed_tool(server, "srv-turn", ListedToolsCaller()) is None
def test_shared_server_ignores_headers_it_never_forwards(self):
manager = MCPServerManager()
server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv")
manager._create_prefixed_tools(
[MCPTool(name="turn", description="everyone", inputSchema={})],
server,
caller=ListedToolsCaller(raw_headers={"authorization": "Bearer sk-litellm", "x-workspace": "A"}),
)
other = ListedToolsCaller(raw_headers={"authorization": "Bearer sk-other", "x-workspace": "B"})
listed = manager.get_listed_tool(server, "turn", other)
assert listed is not None and listed.description == "everyone"
@pytest.mark.parametrize(
("signer", "static_headers", "shared"),
[
pytest.param(MagicMock(), None, False, id="signer-mints-per-caller-authorization"),
pytest.param(MagicMock(), {"Authorization": "Bearer admin-token"}, True, id="static-authorization-wins"),
pytest.param(None, None, True, id="no-signer-stays-shared"),
],
)
def test_jwt_signer_makes_a_shared_server_list_per_caller(self, signer, static_headers, shared):
"""MCPJWTSigner hands upstream a JWT naming the caller on an otherwise shared ``auth_type: none``
server, so the upstream may tailor the catalog and the cache must not hand one caller another's."""
manager = MCPServerManager()
server = MCPServer(
server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv", static_headers=static_headers
)
alice = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="alice", api_key="hashed-alice"))
bob = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="bob", api_key="hashed-bob"))
with patch( # test-quality-ok: the signer is a process-wide singleton the manager reads, no injection seam
"litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer",
return_value=signer,
):
manager._create_prefixed_tools(
[MCPTool(name="turn", description="alice view", inputSchema={})], server, caller=alice
)
for_bob = manager.get_listed_tool(server, "srv-turn", bob)
if shared:
assert for_bob is not None and for_bob.description == "alice view"
else:
assert for_bob is None
@pytest.mark.asyncio
async def test_call_tool_hands_hooks_the_catalog_the_same_forwarded_headers_listed(self):
"""Interleaved callers on a forwarded-header server: the hook must see the caller's own catalog."""
manager = MCPServerManager()
server = MCPServer(
server_id="catalog",
name="catalog",
transport=MCPTransport.http,
url="http://catalog",
extra_headers=["X-Workspace"],
)
manager.registry = {"catalog": server}
catalogs = {
"A": [
MCPTool(
name="turn", description="Catalog A", inputSchema={"properties": {"turn": {"description": "A"}}}
)
],
"B": [
MCPTool(
name="turn", description="Catalog B", inputSchema={"properties": {"turn": {"description": "B"}}}
)
],
}
mock_client = AsyncMock()
mock_client.call_tool.return_value = MagicMock(spec=CallToolResult, content=[], isError=False)
manager._create_mcp_client = AsyncMock(return_value=mock_client)
manager._fetch_tools_with_timeout = AsyncMock(side_effect=lambda client, name: catalogs[client.workspace])
for workspace in ("A", "B"):
manager._create_mcp_client.return_value.workspace = workspace
await manager._get_tools_from_server(
server=server,
extra_headers={"X-Workspace": workspace},
raw_headers={"x-workspace": workspace, "authorization": "Bearer sk-litellm"},
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm", user_id="shared-key"),
)
proxy_logging_obj = MagicMock()
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={})
proxy_logging_obj.pre_call_hook = AsyncMock(return_value={})
proxy_logging_obj.during_call_hook = AsyncMock(return_value=None)
await manager.call_tool(
server_name="catalog",
name="catalog-turn",
arguments={"turn": "A-1"},
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm", user_id="shared-key"),
proxy_logging_obj=proxy_logging_obj,
raw_headers={"x-workspace": "A", "authorization": "Bearer sk-litellm"},
)
hook_kwargs = proxy_logging_obj._create_mcp_request_object_from_kwargs.call_args.args[0]
assert (hook_kwargs["tool_description"], hook_kwargs["tool_input_schema"]) == (
"Catalog A",
{"properties": {"turn": {"description": "A"}}},
)
def test_per_caller_listed_tools_evict_oldest_caller_and_keep_shared(self):
from litellm.proxy._experimental.mcp_server.mcp_server_manager import _LISTED_TOOLS_CALLERS_PER_SERVER
manager = MCPServerManager()
server = MCPServer(
server_id="srv",
name="srv",
transport=MCPTransport.http,
url="http://srv",
auth_type=MCPAuth.oauth2_token_exchange,
)
manager._create_prefixed_tools([MCPTool(name="read", description="shared", inputSchema={})], server)
callers = [
ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id=f"u{i}", api_key=f"k{i}"))
for i in range(_LISTED_TOOLS_CALLERS_PER_SERVER + 1)
]
for caller in callers:
manager._create_prefixed_tools(
[MCPTool(name="read", description=caller.user_api_key_auth.user_id, inputSchema={})],
server,
caller=caller,
)
manager._create_prefixed_tools(
[MCPTool(name="read", description="u1 again", inputSchema={})], server, caller=callers[1]
)
assert manager.get_listed_tool(server, "srv-read", callers[0]) is None
second = manager.get_listed_tool(server, "srv-read", callers[1])
assert second is not None and second.description == "u1 again"
newest = manager.get_listed_tool(server, "srv-read", callers[-1])
assert newest is not None and newest.description == callers[-1].user_api_key_auth.user_id
assert len(manager._listed_tools_by_server_id[server.server_id]) == _LISTED_TOOLS_CALLERS_PER_SERVER + 1
shared = manager.get_listed_tool(server, "srv-read")
assert shared is not None and shared.description == "shared"
@pytest.mark.asyncio
@pytest.mark.parametrize("add_prefix", [True, False])
async def test_openapi_listing_records_listed_tools(self, add_prefix):
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
server = MCPServer(
server_id="petstore-id",
name="petstore",
alias="petstore",
transport=MCPTransport.http,
url=None,
spec_path="/spec.yaml",
)
manager = MCPServerManager()
manager._create_mcp_client = AsyncMock(return_value=AsyncMock())
async def _handler(**kwargs):
return None
global_mcp_tool_registry.unregister_tools_with_prefix("petstore-")
global_mcp_tool_registry.register_tool(
name="petstore-list_pets",
description="List pets",
input_schema={"type": "object", "properties": {"limit": {"type": "integer"}}},
handler=_handler,
)
try:
listed = await manager._get_tools_from_server(server=server, add_prefix=add_prefix)
finally:
global_mcp_tool_registry.unregister_tools_with_prefix("petstore-")
assert [t.name for t in listed] == ["petstore-list_pets" if add_prefix else "list_pets"]
for name in ("list_pets", "petstore-list_pets"):
tool = manager.get_listed_tool(server, name)
assert tool is not None and tool.description == "List pets"
assert tool.input_schema["properties"] == {"limit": {"type": "integer"}}
@pytest.mark.asyncio
async def test_openapi_listing_ignores_overlapping_server_prefix(self):
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
server = MCPServer(
server_id="pet-id",
name="pet",
alias="pet",
transport=MCPTransport.http,
url=None,
spec_path="/spec.yaml",
)
manager = MCPServerManager()
manager._create_mcp_client = AsyncMock(return_value=AsyncMock())
async def _handler(**kwargs):
return None
for prefix in ("pet-", "petstore-"):
global_mcp_tool_registry.unregister_tools_with_prefix(prefix)
global_mcp_tool_registry.register_tool(
name="pet-petstore-list",
description="Local pet tool",
input_schema={"type": "object", "properties": {"limit": {"type": "integer"}}},
handler=_handler,
)
global_mcp_tool_registry.register_tool(
name="petstore-list",
description="Foreign petstore tool",
input_schema={"type": "object", "properties": {"status": {"type": "string"}}},
handler=_handler,
)
try:
listed = await manager._get_tools_from_server(server=server, add_prefix=True)
finally:
for prefix in ("pet-", "petstore-"):
global_mcp_tool_registry.unregister_tools_with_prefix(prefix)
assert [t.name for t in listed] == ["pet-petstore-list"]
tool = manager.get_listed_tool(server, "petstore-list")
assert tool is not None and tool.description == "Local pet tool"
assert tool.input_schema["properties"] == {"limit": {"type": "integer"}}
@pytest.mark.asyncio
async def test_get_allowed_mcp_servers_with_user_api_key_auth(self):
"""

View file

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

View file

@ -300,6 +300,29 @@ class TestAllowFlow:
assert evaluate_call.json["conversationId"] == "sess-123"
assert evaluate_call.json["agentId"] == "agent-007"
@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_tool_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_tool_input_schema=schema))
assert handler.calls[1].json["tool"] == {"name": "send_email"}
@pytest.mark.asyncio
async def test_agent_id_falls_back_to_key_alias(self):
handler: Final = FakeHandler([_token_response(), _allow_response()])

View file

@ -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_tool_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 (out["mcp_tool_description"], out["mcp_tool_input_schema"]) == (None, None)
def test_create_mcp_request_object_from_kwargs_empty(proxy_logging):
obj = proxy_logging._create_mcp_request_object_from_kwargs(kwargs={})
snapshot = {