This commit is contained in:
devin-ai-integration[bot] 2026-09-30 00:39:06 +00:00 • committed by GitHub
commit f2a3079aae
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
13 changed files with 1567 additions and 269 deletions

View file

@ -257,6 +257,20 @@ _user_env_vars_cache: Final[dict[tuple[str, str], tuple[dict[str, str], float]]]
_USER_ENV_VARS_CACHE_TTL: Final = 60 # seconds
_USER_ENV_VARS_CACHE_MAX_SIZE: Final = 4096 # cap to prevent unbounded growth
_ListedToolsByCaller: TypeAlias = Mapping[str | None, Mapping[str, MCPTool]]
_LISTED_TOOLS_CALLERS_PER_SERVER: Final = 256
@dataclass(frozen=True, slots=True)
class ListedToolsCaller:
"""Request inputs that select which upstream catalog a caller was shown by tools/list."""
user_api_key_auth: UserAPIKeyAuth | None = None
mcp_auth_header: str | dict[str, str] | None = None
raw_headers: Mapping[str, str] | None = None
oauth2_headers: Mapping[str, str] | None = None
# Auth types whose upstream OAuth endpoints (protected-resource + authorization-server metadata) the
# gateway discovers from the upstream itself: interactive oauth2 and the two client-forwarded modes.
# OBO/M2M endpoint discovery is decided separately via _obo_needs_endpoint_discovery. Shared by the
@ -1170,6 +1184,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.
@ -1295,6 +1328,21 @@ async def _resolve_byok_mcp_auth_header(
return mcp_auth_header
async def _byok_listing_auth_header(
mcp_server: MCPServer,
user_api_key_auth: UserAPIKeyAuth | None,
mcp_auth_header: str | dict[str, str] | None,
) -> str | dict[str, str] | None:
"""The credential a tools/list may use: a supplied header forwards unchanged, and a missing one
falls back to the stored credential without the tool-call path's byok_auth_required raise."""
if not mcp_server.is_byok or mcp_auth_header is not None:
return mcp_auth_header
from litellm.proxy._experimental.mcp_server.operations import _get_byok_credential
return await _get_byok_credential(mcp_server, user_api_key_auth)
def _client_forwarded_authorization_headers(
mcp_server: MCPServer,
oauth2_headers: dict[str, str] | None,
@ -1973,6 +2021,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
@ -3777,19 +3826,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(
@ -3875,7 +3912,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."""
@ -4406,7 +4443,7 @@ class MCPServerManager:
oauth2_headers: dict[str, str] | None = None,
client_ip: str | None = None,
proxy_logging_obj: ProxyLogging | None = None,
) -> list[MCPTool]:
) -> Sequence[MCPTool]:
"""
Helper method to get tools from a single MCP server with prefixed names.
@ -4425,6 +4462,13 @@ class MCPServerManager:
verbose_logger.info("_get_tools_from_server for %s...", server.name)
client = None
resolved_mcp_auth_header: Final = await _byok_listing_auth_header(server, user_api_key_auth, mcp_auth_header)
listed_caller: Final = ListedToolsCaller(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=resolved_mcp_auth_header,
raw_headers=raw_headers,
oauth2_headers=oauth2_headers,
)
try:
# Tool *listing* must not be blocked by missing per-user env vars —
@ -4468,7 +4512,7 @@ class MCPServerManager:
if (
get_mcp_jwt_signer() is not None
and not has_static_authorization
and not mcp_auth_header
and not resolved_mcp_auth_header
and not has_extra_authorization
):
extra_headers = await inject_mcp_jwt_headers_for_upstream(
@ -4491,7 +4535,7 @@ class MCPServerManager:
client = await self._create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
mcp_auth_header=resolved_mcp_auth_header,
extra_headers=extra_headers,
stdio_env=stdio_env,
subject_token=subject_token,
@ -4522,8 +4566,10 @@ class MCPServerManager:
# applied (e.g. "test_petstore-getinventory"). Do NOT pass them
# through _create_prefixed_tools — that would add the prefix a second
# time producing "test_petstore-test_petstore-getinventory".
unprefixed_tools: Final = guarded_openapi
self._record_listed_tools(server, unprefixed_tools, listed_caller)
if not add_prefix:
return list(guarded_openapi)
return unprefixed_tools
return [t.model_copy(update={"name": registered_names[t.name]}) for t in guarded_openapi]
else:
tools = await self._fetch_tools_with_timeout(client, server.name)
@ -4537,7 +4583,7 @@ class MCPServerManager:
raw_headers=raw_headers,
)
prefixed_or_original_tools: Final = self._create_prefixed_tools(
list(guarded_tools), server, add_prefix=add_prefix
guarded_tools, server, add_prefix=add_prefix, caller=listed_caller
)
return prefixed_or_original_tools
@ -4588,8 +4634,77 @@ class MCPServerManager:
)
self._invalidate_discovery_lists(server_id)
self._listed_tools_by_server_id.pop(server_id, None)
invalidate_oauth_metadata_cache(server_id)
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, the caller bearer (forwarded as-is or
exchanged as the OBO subject), the server-specific auth header, and the per-caller JWT
MCPJWTSigner mints for tools/list 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
forwarded: Final = self._forwarded_header_values(server, caller.raw_headers) or None
header_env: Final = self._build_stdio_env(server, caller.raw_headers)
stdio_env: Final = None if header_env == self._build_stdio_env(server) else header_env
caller_bearer: Final = (
self._extract_subject_token(caller.oauth2_headers, caller.raw_headers, auth)
if server.is_client_forwarded_token or server.auth_type == MCPAuth.oauth2_token_exchange
else None
)
_, digest = self._discovery_key(
server,
auth,
caller.mcp_auth_header,
forwarded,
stdio_env,
caller_bearer,
per_caller=self._signs_caller_identity_upstream(server),
)
return digest
@staticmethod
def _signs_caller_identity_upstream(server: MCPServer) -> bool:
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 server.static_headers is None or not any(k.lower() == "authorization" for k in server.static_headers)
@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, MappingProxyType({}))
shared: Final = existing.get(None)
callers: Final = tuple((key, value) for key, value in existing.items() if key not in (None, identity))
evicted: Final = 0 if identity is None else max(len(callers) + 1 - _LISTED_TOOLS_CALLERS_PER_SERVER, 0)
entries: Final = (
*(() if shared is None else ((None, shared),)),
*callers[evicted:],
(identity, listing),
)
self._listed_tools_by_server_id[server.server_id] = MappingProxyType(dict(entries))
def _discovery_key(
self,
server: MCPServer,
@ -4599,9 +4714,11 @@ class MCPServerManager:
stdio_env: dict[str, str] | None,
subject_token: str | None,
credential_fingerprint: str | None = None,
per_caller: bool = False,
) -> _DiscoveryKey:
per_user: Final = (
server.requires_per_user_auth
per_caller
or server.requires_per_user_auth
or self._references_per_user_env_var(server)
or server.delegate_auth_to_upstream
or server.auth_type in (MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag)
@ -5490,7 +5607,13 @@ class MCPServerManager:
{seen: kept for seen, kept in self._catalog_alert_signatures.items() if seen != key}
)
def _create_prefixed_tools(self, tools: list[MCPTool], server: MCPServer, add_prefix: bool = True) -> list[MCPTool]:
def _create_prefixed_tools(
self,
tools: Sequence[MCPTool],
server: MCPServer,
add_prefix: bool = True,
caller: ListedToolsCaller | None = None,
) -> list[MCPTool]:
"""
Create prefixed tools and update tool mapping.
@ -5515,9 +5638,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, MappingProxyType({})).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.get(tool.name) if server.tool_name_to_description else None
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]:
@ -5751,6 +5886,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.
@ -5764,6 +5900,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
@ -5827,6 +5966,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
@ -5882,6 +6023,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.
@ -5896,6 +6038,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(),
)
@ -6042,21 +6186,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
@ -6507,6 +6637,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
@ -6523,6 +6659,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"]
@ -6539,6 +6676,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)

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
@ -1606,6 +1607,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],
@ -2075,6 +2082,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.
@ -2143,7 +2151,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
@ -2185,6 +2194,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
@ -61,6 +61,7 @@ _GATEWAY_OWNED_TOKEN_ERRORS: Final = frozenset(
_INVALID_ASSERTION_AADSTS_PREFIX: Final = "50027"
_AADSTS_CODES_ADAPTER: Final = TypeAdapter(tuple[int, ...])
_MCP_CALL_TYPES: Final[tuple[str, ...]] = ("mcp_call", "call_mcp_tool")
_TOOL_INPUT_SCHEMA_ADAPTER: Final = TypeAdapter(dict[str, object])
_OBO_CACHE_MAX_ENTRIES: Final = 1000
_DEFAULT_TOKEN_TTL_SECONDS: Final = 3599.0
_TOKEN_EXPIRY_SLACK_SECONDS: Final = 60.0
@ -82,6 +83,13 @@ def _parse_aadsts_codes(raw: object) -> tuple[int, ...]:
return ()
def _parse_tool_input_schema(raw: object) -> Mapping[str, object] | None:
try:
return _TOOL_INPUT_SCHEMA_ADAPTER.validate_python(raw)
except ValidationError:
return None
def entra_assertion(value: object) -> str | None:
"""``value`` when it is a compact JWS, the only bearer shape the OBO exchange accepts as its assertion.
A LiteLLM virtual key, session bearer, or opaque upstream token in ``Authorization`` yields ``None``."""
@ -100,6 +108,14 @@ class _EvaluateResponse(TypedDict, total=False):
correlationId: ReadOnly[str]
class _ToolReference(BaseModel):
model_config = ConfigDict(frozen=True)
name: str
description: str | None = None
input_schema: Mapping[str, object] | None = Field(default=None, serialization_alias="inputSchema")
class _UnavailableDetail(TypedDict):
error: ReadOnly[str]
message: ReadOnly[str]
@ -392,8 +408,14 @@ class Agent365Guardrail(CustomGuardrail):
arguments: Final = data.get("mcp_arguments")
server_name: Final = str(data.get("mcp_server_name") or "litellm")
agent_id: Final = user_api_key_dict.key_alias
description: Final = data.get("mcp_tool_description")
tool_reference: Final = _ToolReference(
name=tool_name,
description=description if isinstance(description, str) and description else None,
input_schema=_parse_tool_input_schema(data.get("mcp_input_schema")),
)
payload: Final[dict[str, object]] = { # mutable-ok: JSON body with optional fields added below
"tool": {"name": tool_name},
"tool": tool_reference.model_dump(by_alias=True, exclude_none=True),
"serverName": server_name,
"conversationId": self._resolve_conversation_id(data),
}

View file

@ -1476,8 +1476,12 @@ class ProxyLogging:
TypeAdapter(dict[str, object]).validate_python(guardrail_context.get("metadata") or MappingProxyType({}))
)
mcp_tool_description: Final = kwargs.get("mcp_tool_description")
mcp_input_schema: Final = kwargs.get("mcp_input_schema")
mcp_tool_description: Final = request_obj.tool_description or kwargs.get("mcp_tool_description")
mcp_input_schema: Final = (
request_obj.tool_input_schema
if request_obj.tool_input_schema is not None
else kwargs.get("mcp_input_schema")
)
description_line: Final = f"\nDescription: {mcp_tool_description}" if mcp_tool_description else ""
tool_call_content: Final = (
f"Tool: {request_obj.tool_name}{description_line}\nArguments: {request_obj.arguments}"
@ -1736,6 +1740,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

@ -192,3 +192,26 @@ def test_byok_server_uses_the_calling_users_stored_credential_and_fails_closed_w
assert removed.status_code in (200, 204), removed.text
eventually(lambda: call_tool(gateway, owner_key, identity, name, ADD), lambda value: value.status_code == 401)
assert tool_calls(peer.drain()) == ()
def test_deprecated_string_x_mcp_auth_lists_a_byok_server_for_a_key_without_a_user(gateway: Gateway) -> None:
with mcp_peer() as peer, gateway.scenario() as scenario:
alias: Final = "byok" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias, auth_type="bearer_token", is_byok=True)
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
peer.drain()
response: Final = gateway.client.get(
"/mcp-rest/tools/list",
params={"server_id": identity},
headers={"x-litellm-api-key": key, "x-mcp-auth": "Bearer hdr"},
)
assert response.status_code == 200, response.text
names: Final = {tool["name"] for tool in response.json()["tools"]}
assert "add" in names, names
listings: Final = tuple(
item
for item in peer.drain()
if isinstance(item.get("body"), dict) and item["body"].get("method") == "tools/list"
)
assert len(listings) == 1, listings
assert listings[0]["headers"].get(b"authorization") == b"Bearer hdr"

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
@ -6877,7 +6878,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
@ -7989,9 +7989,12 @@ async def test_execute_mcp_tool_sets_model_in_model_call_details():
fake_server.server_name = "openapi-petstore"
fake_server.alias = None
fake_server.short_prefix = None
fake_server.tool_name_to_description = None
fake_tool = MagicMock()
fake_tool.name = "list_pets"
fake_tool.description = "test tool"
fake_tool.input_schema = {"type": "object"}
start_time = datetime.now(timezone.utc)
litellm_logging_obj, _ = function_setup(
@ -8042,6 +8045,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

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

@ -3300,7 +3300,8 @@ async def test_request_selected_tool_specific_guardrail_applies_to_virtual_execu
default_on=False,
custom_code="def apply_guardrail(inputs, request_data, input_type):\n"
' if inputs.get("tools", [{}])[0].get("function", {}).get("name") == "execute":\n'
f' return {{"action": "{action}", "reason": "resolved tool blocked", "texts": ["redacted"]}}\n'
' texts = [t.replace("confidential", "redacted") for t in inputs.get("texts", [])]\n'
f' return {{"action": "{action}", "reason": "resolved tool blocked", "texts": texts}}\n'
" return allow()\n",
)
manager: Final = mcp_server_manager.MCPServerManager()

View file

@ -346,6 +346,29 @@ class TestAllowFlow:
assert evaluate_call.json["conversationId"] == "sess-123"
assert evaluate_call.json["agentId"] == "my-agent-key"
@pytest.mark.asyncio
async def test_evaluate_payload_includes_listed_tool_metadata(self):
handler: Final = FakeHandler([_token_response(), _allow_response()])
guardrail: Final = _make_guardrail(handler)
schema: Final = {"type": "object", "properties": {"to": {"type": "string"}}, "required": ["to"]}
await _run(guardrail, _mcp_data(mcp_tool_description="Send an email", mcp_input_schema=schema))
assert handler.calls[1].json["tool"] == {
"name": "send_email",
"description": "Send an email",
"inputSchema": schema,
}
@pytest.mark.asyncio
@pytest.mark.parametrize(
("description", "schema"),
[(None, None), ("", None), (None, ["not", "a", "schema"]), (42, "type: object")],
)
async def test_evaluate_payload_omits_missing_or_malformed_tool_metadata(self, description, schema):
handler: Final = FakeHandler([_token_response(), _allow_response()])
guardrail: Final = _make_guardrail(handler)
await _run(guardrail, _mcp_data(mcp_tool_description=description, mcp_input_schema=schema))
assert handler.calls[1].json["tool"] == {"name": "send_email"}
@pytest.mark.asyncio
async def test_non_mcp_call_type_skipped(self):
handler: Final = FakeHandler([])

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_input_schema"]) == ("Adds numbers", schema)
def test_mcp_tool_metadata_absent_when_tool_was_never_listed(proxy_logging):
obj = proxy_logging._create_mcp_request_object_from_kwargs(kwargs={"name": "calc", "arguments": {}})
out = proxy_logging._convert_mcp_to_llm_format(request_obj=obj, kwargs={})
assert "mcp_tool_description" not in out and "mcp_input_schema" not in out
def test_create_mcp_request_object_from_kwargs_empty(proxy_logging):
obj = proxy_logging._create_mcp_request_object_from_kwargs(kwargs={})
snapshot = {