mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
90873c46de
commit
3c13e3b457
11 changed files with 876 additions and 59 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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()])
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue