feat(mcp): hand listed-tool description and input schema to pre-call hooks per caller (#41162)

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

* refactor(mcp): drop the listed-tools empty sentinel and routine test docstrings

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): mark the listed-tools cache digest as a non-security hash

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): key the listed-tools cache by the OBO subject token

token_exchange servers list upstream with the caller's own Entra bearer, so two callers on one
LiteLLM key with different subjects were sharing a catalog slot

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): resolve the BYOK credential before keying the listed-tools slot

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): drop the OAuth discovery cache when a server definition changes

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(mcp): drop a diff-narrating comment

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): never validate a supplied header on the tools/list BYOK path

The pre-listing resolver ran the tool-call byok_auth_required check even
when the caller already supplied x-mcp-auth, and it ran outside the
per-server error boundary, so a single deprecated-header caller dropped
the server from the aggregate list. Listing now returns a supplied header
unchanged and falls back to the stored credential without raising

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(mcp): assert the BYOK listing lands in the caller's listed-tool slot

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(mcp): cover the deprecated string x-mcp-auth header on a BYOK tools/list

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): key the per-caller listed-tool slot by the hashed token

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): key discovery cache by the hashed token

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): key discovery caches per caller correctly and drop stale caches on server updates

Discovery-list cache identity now uses the hashed token instead of the raw
api_key and treats MCPJWTSigner-signed servers as per caller. Server
definition changes also drop the cached upstream OAuth metadata. OpenAPI
listings look tools up under the normalized registry prefix with the
separator, so an overlapping sibling prefix no longer leaks into the list.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): keep the discovery cache digest call unchanged so CodeQL matches the existing alert

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(mcp): derive the listed-tool caller identity from the discovery cache key

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): guard OAuth metadata cache writes with a per-server generation and drop unproven per-caller discovery keys

An upstream metadata fetch that started before a server edit could store its stale reply after
invalidate_oauth_metadata_cache ran. Invalidation now bumps a per-server generation and the fetch
only stores when the generation it captured before I/O is unchanged.

The MCPJWTSigner-based per-caller discovery classification and the api_key to token key change had no
reproduction (the signer only injects on tools/list, and UserAPIKeyAuth hashes api_key in place), so
both go back to the merge-base behavior.

Integration coverage under tests/integration/mcp: overlapping OpenAPI aliases, a config-declared
server name with a space, OAuth metadata refetch after a save, and the in-flight stale-write race

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): keep OAuth metadata generations only while a fetch is in flight

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): count queued OAuth metadata fetchers so invalidation survives lock handoff

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): keep a held OAuth metadata lock registered even when no fetcher slot claims it

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(mcp): prove a peer worker drops stale upstream OAuth metadata after a save elsewhere

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(mcp): return one masked text per scanned string in the selected-guardrail REST test

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): fold the signed caller into the discovery digest instead of a second key hash

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): satisfy type discipline gate on listed-tool identity

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): hand tools/call hooks the exact catalog entry tools/list served

get_listed_tool re-applied the admin description override on top of the cached listing, so a
guardrail-masked description was restored to its original wording at call time, and the OpenAPI /
local-registry call path built its metadata from the registry instead of the guarded caller catalog.
Both paths now return the cached entry as served, falling back to the registry only when no listing
was recorded

Adds tests/integration/mcp/test_mcp_listed_tool_metadata.py (red on the prior head for the two
regressions, red on the merge base for the feature, green on this head)

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): key OpenAPI listed-tool entries per caller so tools/call reads its own guarded listing

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(mcp): align listed-tool slot tests with per-caller keying

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): keep oauth2 listing on the minted or signed credential, not the stored BYOK secret

The listing helper that keys the per-caller catalog by the stored BYOK credential also handed that
credential to the upstream client, which on an oauth2 server short-circuited the client_credentials
mint and the MCPJWTSigner gate. Split the two: the catalog identity keeps the stored credential so
tools/call finds the caller's slot, while an oauth2 server's tools/list sends only the per-request
header, letting the M2M mint or signed JWT proceed as on main

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(mcp): oauth2 BYOK listing sends the minted token, not the stored secret, through the real proxy

Integration cell for the listing fix: a client_credentials BYOK server with a stored user credential,
one tools/list as that user, the peer must see a live minted bearer and one /token mint. Red at the
pre-fix tip (zero mints, stored secret upstream), green at the fixed head and at the merge base

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): keep the stored BYOK credential for catalog identity only on tools/list

Listing used the resolved stored credential both to key the caller's catalog slot and as the
upstream transport header, so REST api_key and bearer_token listings sent the user's secret
instead of the server's static token and the MCPJWTSigner gate went quiet. The upstream client
and the signer gate now read the caller-supplied mcp_auth_header for every auth type, exactly
as before the catalog existed, and the stored credential only names the slot tools/call reads

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): type the listed-tool metadata read from pre-call kwargs for the basedpyright gate

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): hand never-listed tools/call hooks name and arguments only

The local-registry call path fell back to the registry entry with the admin description override
when no tools/list had been recorded for the caller, so a pre_mcp_call guardrail scanned a
description the caller was never served and blocked OpenAPI calls that passed before, and base's
own selected-guardrail REST test failed on the two-text redaction. _registered_tool_metadata now
returns the listed entry or None, so a tools/call with no prior listing sends name and arguments
only as promised, and that REST test double goes back to its base shape

* fix(mcp): keep during_mcp_call hooks on name and arguments only

call_tool handed the caller's listed entry to the during-hook task as well, so during_mcp_call
guardrails scanned the description line and schema leaves of any listed tool after the upstream
call had already run, blocking calls that passed before whenever the policy matched the
description, returned a fixed-length texts list, or hit the depth guard on a deep schema. The
listed entry is only disclosed for pre_mcp_call, so the during task no longer receives it and its
request object carries no description or schema, as before

* fix(mcp): key the BYOK catalog slot by the client's header, not the stored credential

tools/list resolved the stored BYOK credential to pick the caller's catalog slot, which read the
credential store before the classified try block. With Postgres down and a cold per-worker cache
that made every REST tools/list on an is_byok server fail with tools=[] and no upstream call, and
the read seeded the per-worker cache (including a negative entry), so a tools/call on another
worker after a store, rotate or revoke on this one kept using the stale value.

The slot is now keyed by what the client supplied plus the caller's hashed key, on both sides.
_get_tools_from_server and call_tool take a keyword-only catalog_auth_header that defaults to
mcp_auth_header as received (the default is the builtin Ellipsis so it survives a module reload).
The /mcp fan-out and execute_mcp_tool, which swap the resolved credential into mcp_auth_header,
pass the client's value explicitly. What goes upstream is unchanged. _byok_catalog_auth_header is
gone.

* fix(mcp): drop a listed catalog recorded across a server save

_record_listed_tools ran after the awaited upstream fetch, so a PUT /v1/mcp/server that landed
mid-fetch had its invalidation undone when the fetch completed: hooks then saw the pre-save
description next to the post-save definition until the next listing, instead of name and
arguments only.

The manager now keeps a per-server listed-tools generation, bumped by
_invalidate_server_definition_caches. _get_tools_from_server reads it before the fetch and
_record_listed_tools skips the write when it moved; the next listing records normally.

* fix(mcp): drop the catalog again once a saved OpenAPI server's registry is rebuilt

add_server and update_server publish the saved definition before the OpenAPI registry entries are
rebuilt from the spec, so a listing recorded during that fetch held the pre-save entries under the
new generation. The generation is bumped a second time after the registry refresh.

The during-hook task no longer accepts a listed entry, the one-line wrapper over get_listed_tool is
inlined at its two call sites, and the per-server generation map is a plain dict.

* fix(mcp): keep discovery and OAuth metadata caches across an OpenAPI spec re-read

add_server and update_server ran the full server-definition invalidation a
second time after the awaited OpenAPI spec fetch, which also dropped the
prompts/resources/templates discovery entries and the OAuth protected-resource
metadata filled under the already-published definition, so the next request
went upstream again. Only the listed-tool catalog recorded during the fetch
holds pre-save entries, so the post-fetch pass now drops just that catalog
and bumps its generation via the new _drop_listed_tools helper, which the
full invalidation also calls.

* fix(mcp): look a called tool up in the listed catalog by its bare name only

get_listed_tool stripped the server prefix a second time when the exact
name was absent from the caller's listing, so a never-listed upstream tool
whose bare name starts with the server prefix resolved to the listed
sibling and that sibling's description and input schema reached the
pre-call hooks for a call to a different tool. Every caller already passes
the once-stripped bare name, so the lookup is now exact.

Tests that looked the catalog up by a prefixed name now use the bare name
the callers pass; two new tests pin the never-listed sibling case at the
manager and at the tools/call path.

* fix(mcp): record a listed-tool catalog only for a listing the caller is served

_get_tools_from_server now records the catalog into the caller's
listed-tools slot only when asked (record_listing=True), which the
served listings pass: the /mcp and Responses API tools/list handlers via
_get_tools_from_mcp_servers, MCPServerManager.list_tools, and the REST
listing via _list_server_tools. Four internal listings stop recording,
so a later tools/call hands pre_mcp_call hooks name and arguments only,
as on main:

- _list_tools_before_first_call, the implicit listing inside tools/call
  when this worker does not yet expose the tool
- fetch_pinnable_tool_catalog, the admin pin snapshot listed without the
  catalog guard and without description overrides
- _initialize_tool_name_to_mcp_server_name_mapping, the startup fill
- get_tools_for_server, used by the semantic tool filter

_create_prefixed_tools returns to its tool-name mapping job only; the
record follows it in _get_tools_from_server.

* fix(mcp): opt every listing out of catalog recording unless it is served

The aggregate listing and _list_mcp_tools now default to record_listing=False,
so a catalog fetched inside a tools/call no longer fills the caller's
listed-tools slot. The /mcp/proxy meta-tools (call_tool, search_tools,
get_tool_schema) and the tool-search virtual tool stop recording: /mcp/proxy
serves only the meta-tools and the search serves only its hits, so a later
pre_mcp_call hook was reading a description the caller never listed.

The tools/list handler, the Responses MCP handler and the /v1/mcp/tools
management listing opt in with record_listing=True, since each serves the
catalog to the caller.

* fix(mcp): key the listed-tool slot by the caller's admission identity and forwarded bearer

The slot a tools/list records for a later tools/call was keyed by (user_id, api_key)
only, so every team-only JWT caller shared one slot and one JWT user acting in two
teams shared a slot; a tools/call then handed pre_mcp_call hooks a description another
caller was served. The slot is now keyed by the hashed key, user, team and organization,
plus the admission credential of a caller admitted with neither a key nor a user.

The caller bearer split the slot only on client-forwarded-token and token-exchange
servers; a legacy delegated oauth2 server (delegate_auth_to_upstream without client
credentials) also forwards it upstream and served a different catalog per bearer into
one slot. The bearer now splits the slot on every server whose egress forwards it
(_consumes_caller_authorization) or exchanges it.

* fix(mcp): record only tools served by the bridge

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): keep bridge tool metadata request-local

* refactor(mcp): centralize listed catalog recording guard

* fix(mcp): preserve base TPM reservations for listed tool calls

* fix(mcp): preserve project token reservations for listed calls

* fix(mcp): record served catalogs and preserve call message bytes

* test(mcp): align listing expectation with deferred recording

* test(mcp): audit listed metadata across callers and bridge lifecycles

* test(mcp): preserve guardrail fixture worker affinity

* refactor(mcp): expose listed catalog recording API

* feat(mcp): pass served_tools through the anthropic messages bridge

/v1/messages auto-execution now hands this request's resolved tool
definitions to _execute_tool_calls, matching the Responses and chat
completions bridges: the pre_mcp_call hook receives the description
and input schema the model was shown for that call. Request-local
only; the shared listed-tools catalog is untouched.

* test(mcp): pin served_tools handoff on the anthropic messages bridge

Mirrors the credentials-forwarding test: the request's resolved tool
definitions must reach _execute_tool_calls under served_tools so
pre_mcp_call hooks judge the call on the description and input schema
the model was shown. Fails without the previous commit's one-liner.

* style(mcp): sort the local import block ruff flagged

* fix(mcp): keep the admin include_disabled_tools view off the listed-tools catalog

GET /mcp-rest/tools/list?include_disabled_tools=true is the admin-only
configuration view: apply_tool_filters is False, so it serves the full
server catalog. Recording that response into the caller's listed-tools
slot warmed tools/call metadata no runtime listing ever served,
breaking the only-a-served-listing-records invariant (Bugbot).

The record is now gated on apply_tool_filters; disabled tools stay
unreachable (the call-time allowlist 403 fires before hooks), so the
observable fix is the slot no longer warming from a settings view.
Verified live: the new test fails on the unfixed head and passes here,
and the rest of the listed-tool-metadata suite is unchanged.

---------

Co-authored-by: yucheng <yucheng@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-04 00:53:58 -07:00 • committed by GitHub
parent 9a5e828310
commit 0b74ae9c5c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
36 changed files with 5585 additions and 406 deletions

View file

@ -147,6 +147,7 @@ async def anthropic_messages_with_mcp(
tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
tool_server_map=tool_server_map,
served_tools=deduplicated_mcp_tools,
tool_calls=list(tool_use_blocks),
user_api_key_auth=context.user_api_key_auth,
mcp_auth_header=context.mcp_auth_header,

View file

@ -28,7 +28,7 @@ from contextlib import asynccontextmanager
from dataclasses import dataclass, replace
from functools import lru_cache
from itertools import chain, groupby
from types import MappingProxyType
from types import EllipsisType, MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Generic, Literal, TypeAlias, TypedDict, TypeVar, cast
from urllib.parse import ParseResult, urlparse
@ -259,6 +259,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
@ -1172,6 +1186,57 @@ def _authorization_is_litellm_admission_credential(
return bool(user_api_key_auth and user_api_key_auth.api_key and not admission_header)
def _server_auth_header_for(
server: MCPServer,
mcp_server_auth_headers: Mapping[str, str | dict[str, str]] | None,
mcp_auth_header: str | dict[str, str] | None,
) -> str | dict[str, str] | None:
"""Server-specific ``x-mcp-<alias>-authorization`` header, else the deprecated global one."""
server_specific: Final = (
lookup_mcp_server_auth_in_headers(
mcp_server_auth_headers,
alias=server.alias,
server_name=server.server_name,
access_groups=server.access_groups,
)
if mcp_server_auth_headers
else None
)
return mcp_auth_header if server_specific is None else server_specific
def listed_tools_caller_for(
server: MCPServer,
user_api_key_auth: UserAPIKeyAuth | None,
mcp_auth_header: str | dict[str, str] | None,
mcp_server_auth_headers: Mapping[str, str | dict[str, str]] | None,
raw_headers: Mapping[str, str] | None,
oauth2_headers: Mapping[str, str] | None,
) -> ListedToolsCaller:
"""The caller a tools/call must look its listed entry up under: the same inputs tools/list keyed by."""
return ListedToolsCaller(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=_server_auth_header_for(server, mcp_server_auth_headers, mcp_auth_header),
raw_headers=raw_headers,
oauth2_headers=oauth2_headers,
)
def _admission_identity(
auth: UserAPIKeyAuth, raw_headers: Mapping[str, str] | None
) -> tuple[str | None, str | None, str | None, str | None, str | None]:
"""The admission identity the served catalog is shaped for: the hashed key, user, team and
organization, plus the admission credential (``x-litellm-api-key``, else ``Authorization``) of a
caller admitted with neither a key nor a user."""
keyless: Final = auth.api_key is None and auth.user_id is None
credential: Final = (
_raw_header_value(raw_headers, "x-litellm-api-key") or _raw_header_value(raw_headers, "authorization")
if keyless
else None
)
return auth.api_key, auth.user_id, auth.team_id, auth.org_id, credential
def _format_byok_openapi_auth_header(mcp_server: MCPServer, mcp_auth_header: str) -> str:
"""Format a raw BYOK credential for OpenAPI tool ``Authorization`` injection.
@ -1297,6 +1362,16 @@ async def _resolve_byok_mcp_auth_header(
return mcp_auth_header
def _catalog_auth_header(
mcp_auth_header: str | dict[str, str] | None,
catalog_auth_header: str | dict[str, str] | None | EllipsisType,
) -> str | dict[str, str] | None:
"""The header the client supplied, which keys the caller's catalog slot on both tools/list and
tools/call. A caller that already swapped a stored BYOK credential into ``mcp_auth_header`` passes
the client's value explicitly, since the stored credential must never be read to find the slot."""
return mcp_auth_header if catalog_auth_header is ... else catalog_auth_header
def _client_forwarded_authorization_headers(
mcp_server: MCPServer,
oauth2_headers: dict[str, str] | None,
@ -1925,6 +2000,8 @@ class MCPServerManager:
"gmail_send_email": "zapier_mcp_server",
}
"""
self._listed_tools_by_server_id: dict[str, _ListedToolsByCaller] = {} # mutable-ok: refreshed per tools/list
self._listed_tools_generations: dict[str, int] = {} # mutable-ok: bumped per server save
self._upstream_initialize_instructions_by_server_id: dict[str, str] = {}
# Per-server monotonic timestamp of last upstream prefetch attempt (success,
# empty result, or failure). Used to throttle re-probes for servers that do
@ -3242,6 +3319,8 @@ class MCPServerManager:
self._invalidate_server_definition_caches(mcp_server.server_id)
self.registry[mcp_server.server_id] = new_server
await self._maybe_register_openapi_tools(new_server)
if new_server.spec_path:
self._drop_listed_tools(mcp_server.server_id)
self.prime_oauth_metadata_discovery(new_server)
verbose_logger.debug("Added MCP Server: %s", new_server.name)
@ -3279,6 +3358,8 @@ class MCPServerManager:
self._invalidate_server_definition_caches(mcp_server.server_id)
self.registry[mcp_server.server_id] = new_server
await self._maybe_register_openapi_tools(new_server)
if new_server.spec_path:
self._drop_listed_tools(mcp_server.server_id)
self.prime_oauth_metadata_discovery(new_server)
verbose_logger.debug("Updated MCP Server: %s", new_server.name)
@ -3754,25 +3835,14 @@ class MCPServerManager:
verbose_logger.warning("MCP Server %s not found", server_id)
return []
# Get server-specific auth header if available
server_auth_header: str | dict[str, str] | None = None
if mcp_server_auth_headers:
server_auth_header = lookup_mcp_server_auth_in_headers(
mcp_server_auth_headers,
alias=server.alias,
server_name=server.server_name,
access_groups=server.access_groups,
)
# Fall back to deprecated mcp_auth_header if no server-specific header found
if server_auth_header is None:
server_auth_header = mcp_auth_header
server_auth_header: Final = _server_auth_header_for(server, mcp_server_auth_headers, mcp_auth_header)
try:
tools: Final = await self._get_tools_from_server(
server=server,
mcp_auth_header=server_auth_header,
user_api_key_auth=user_api_key_auth,
record_listing=True,
)
return tools
except Exception as e:
@ -3852,7 +3922,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."""
@ -4147,13 +4217,20 @@ class MCPServerManager:
oauth2_headers: dict[str, str] | None = None,
client_ip: str | None = None,
proxy_logging_obj: ProxyLogging | None = None,
) -> list[MCPTool]:
*,
catalog_auth_header: str | dict[str, str] | None | EllipsisType = ...,
record_listing: bool = False,
) -> Sequence[MCPTool]:
"""
Helper method to get tools from a single MCP server with prefixed names.
Args:
server (MCPServer): The server to query tools from
mcp_auth_header: Optional auth header for MCP server
catalog_auth_header: The header the client supplied, keying the caller's catalog slot;
defaults to ``mcp_auth_header``
record_listing: Record the served catalog into the caller's listed-tools slot; only a
listing actually served to the caller sets it
Returns:
List[MCPTool]: List of tools available on the server with prefixed names
@ -4169,6 +4246,13 @@ class MCPServerManager:
verbose_logger.info("_get_tools_from_server for %s...", server.name)
client = None
listed_caller: Final = ListedToolsCaller(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=_catalog_auth_header(mcp_auth_header, catalog_auth_header),
raw_headers=raw_headers,
oauth2_headers=oauth2_headers,
)
listed_generation: Final = self._listed_tools_generations.get(server.server_id, 0)
try:
# Tool *listing* must not be blocked by missing per-user env vars —
@ -4266,8 +4350,12 @@ class MCPServerManager:
# applied (e.g. "test_petstore-getinventory"). Do NOT pass them
# through _create_prefixed_tools — that would add the prefix a second
# time producing "test_petstore-test_petstore-getinventory".
unprefixed_tools: Final = guarded_openapi
self._record_listed_tools(
server, unprefixed_tools, listed_caller, listed_generation, record_listing=record_listing
)
if not add_prefix:
return list(guarded_openapi)
return unprefixed_tools
return [t.model_copy(update={"name": registered_names[t.name]}) for t in guarded_openapi]
else:
tools = await self._fetch_tools_with_timeout(client, server.name)
@ -4281,7 +4369,10 @@ class MCPServerManager:
raw_headers=raw_headers,
)
prefixed_or_original_tools: Final = self._create_prefixed_tools(
list(guarded_tools), server, add_prefix=add_prefix
guarded_tools, server, add_prefix=add_prefix
)
self._record_listed_tools(
server, guarded_tools, listed_caller, listed_generation, record_listing=record_listing
)
return prefixed_or_original_tools
@ -4332,8 +4423,99 @@ class MCPServerManager:
)
self._invalidate_discovery_lists(server_id)
self._drop_listed_tools(server_id)
invalidate_oauth_metadata_cache(server_id)
def _drop_listed_tools(self, server_id: str) -> None:
self._listed_tools_by_server_id.pop(server_id, None)
self._listed_tools_generations[server_id] = self._listed_tools_generations.get(server_id, 0) + 1
def _listed_tools_identity(self, server: MCPServer, caller: ListedToolsCaller | None) -> str | None:
"""Key the listed-tool cache by every request input that can change the served catalog.
The catalog is guardrail-shaped for the caller's admission identity (default-on guardrails,
key or team selections and opt-outs), so every admitted caller gets its own slot, keyed by
``_admission_identity``: the hashed key, user, team and organization, plus the admission
credential of a caller admitted with neither a key nor a user (a team-only JWT). Forwarded
headers, header-driven stdio env, the caller bearer on every server whose egress forwards it
(``_consumes_caller_authorization``) or exchanges it as the OBO subject, and the
server-specific auth header also reach upstream and split the slot further. Only unkeyed
listings with none of those share the ``None`` slot.
"""
if caller is None:
return None
auth: Final = caller.user_api_key_auth
forwarded: Final = self._forwarded_header_values(server, caller.raw_headers) or None
header_env: Final = self._build_stdio_env(server, caller.raw_headers)
stdio_env: Final = None if header_env == self._build_stdio_env(server) else header_env
caller_bearer: Final = (
self._extract_subject_token(caller.oauth2_headers, caller.raw_headers, auth)
if _consumes_caller_authorization(server) or server.auth_type == MCPAuth.oauth2_token_exchange
else None
)
identity: Final = None if auth is None else _admission_identity(auth, caller.raw_headers)
if not (identity or caller.mcp_auth_header or forwarded or stdio_env or caller_bearer):
return None
material: Final = json.dumps(
(identity, caller.mcp_auth_header, forwarded, stdio_env, caller_bearer),
sort_keys=True,
separators=(",", ":"),
)
return hashlib.sha256(material.encode()).hexdigest()
@staticmethod
def _forwarded_header_values(
server: MCPServer, raw_headers: Mapping[str, str] | None
) -> tuple[tuple[str, str], ...]:
if not raw_headers or not server.extra_headers:
return ()
forwarded_names: Final = frozenset(name.lower() for name in server.extra_headers)
return tuple(
sorted((name.lower(), value) for name, value in raw_headers.items() if name.lower() in forwarded_names)
)
def listed_tools_generation(self, server_id: str) -> int:
return self._listed_tools_generations.get(server_id, 0)
def record_listed_tools(
self,
server: MCPServer,
tools: Sequence[MCPTool],
caller: ListedToolsCaller | None,
generation: int,
*,
record_listing: bool = True,
) -> None:
self._record_listed_tools(server, tools, caller, generation, record_listing=record_listing)
def _record_listed_tools(
self,
server: MCPServer,
tools: Sequence[MCPTool],
caller: ListedToolsCaller | None,
generation: int | None = None,
*,
record_listing: bool = True,
) -> None:
"""Store the catalog served to ``caller``. ``generation`` is the server's listed-tools generation
read before the listing's upstream fetch; the record is skipped when it no longer matches."""
if not record_listing:
return
if generation is not None and generation != self._listed_tools_generations.get(server.server_id, 0):
return
identity: Final = self._listed_tools_identity(server, caller)
listing: Final = MappingProxyType({tool.name: tool for tool in tools})
existing: Final = self._listed_tools_by_server_id.get(server.server_id, MappingProxyType({}))
shared: Final = existing.get(None)
callers: Final = tuple((key, value) for key, value in existing.items() if key not in (None, identity))
evicted: Final = 0 if identity is None else max(len(callers) + 1 - _LISTED_TOOLS_CALLERS_PER_SERVER, 0)
entries: Final = (
*(() if shared is None else ((None, shared),)),
*callers[evicted:],
(identity, listing),
)
self._listed_tools_by_server_id[server.server_id] = MappingProxyType(dict(entries))
def _discovery_key(
self,
server: MCPServer,
@ -4343,9 +4525,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)
@ -5253,7 +5437,12 @@ class MCPServerManager:
{seen: kept for seen, kept in self._catalog_alert_signatures.items() if seen != key}
)
def _create_prefixed_tools(self, tools: list[MCPTool], server: MCPServer, add_prefix: bool = True) -> list[MCPTool]:
def _create_prefixed_tools(
self,
tools: Sequence[MCPTool],
server: MCPServer,
add_prefix: bool = True,
) -> list[MCPTool]:
"""
Create prefixed tools and update tool mapping.
@ -5281,6 +5470,13 @@ class MCPServerManager:
verbose_logger.info("Successfully fetched %s tools from server %s", len(prefixed_tools), server.name)
return prefixed_tools
def get_listed_tool(self, server: MCPServer, name: str, caller: ListedToolsCaller | None = None) -> MCPTool | None:
identity: Final = self._listed_tools_identity(server, caller)
listed: Final = self._listed_tools_by_server_id.get(server.server_id, MappingProxyType({})).get(identity)
if not listed:
return None
return listed.get(name)
def _create_prefixed_prompts(
self, prompts: Sequence[Prompt], server: MCPServer, add_prefix: bool = True
) -> list[Prompt]:
@ -5514,6 +5710,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.
@ -5527,6 +5724,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
@ -5590,6 +5790,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
@ -5805,21 +6007,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
@ -6242,6 +6430,9 @@ class MCPServerManager:
guardrail_context: Mapping[str, object] | None = None,
client_ip: str | None = None,
wire_compat: WireCompat = WireCompat.LEGACY,
*,
catalog_auth_header: str | None | EllipsisType = ...,
listed_tool: MCPTool | None | EllipsisType = ...,
) -> CallToolResult | InputRequiredResult:
"""
Call a tool with the given name and arguments
@ -6253,6 +6444,8 @@ class MCPServerManager:
user_api_key_auth: User authentication
mcp_auth_header: MCP auth header (deprecated)
mcp_server_auth_headers: Optional dict of server-specific auth headers {server_alias: auth_value}
catalog_auth_header: The header the client supplied, keying the caller's catalog slot;
defaults to ``mcp_auth_header`` as received, before BYOK resolution
proxy_logging_obj: Optional ProxyLogging object for hook integration
litellm_logging_obj: Optional request logger the guardrail hooks record
their evaluations onto, so MCP guardrail activity reaches the
@ -6264,6 +6457,7 @@ class MCPServerManager:
"""
start_time: Final = datetime.datetime.now()
mcp_server: Final = self._resolve_mcp_server_for_tool_call(server_name, name)
client_auth_header: Final = _catalog_auth_header(mcp_auth_header, catalog_auth_header)
# Resolved before any hook runs so a missing BYOK credential (401) never
# leaves during-hook side effects (audit logging, rate-limit bookkeeping)
@ -6273,6 +6467,9 @@ class MCPServerManager:
user_api_key_auth,
mcp_auth_header,
)
listed_caller: Final = listed_tools_caller_for(
mcp_server, user_api_key_auth, client_auth_header, mcp_server_auth_headers, raw_headers, oauth2_headers
)
#########################################################
# Pre MCP Tool Call Hook
@ -6289,6 +6486,7 @@ class MCPServerManager:
raw_headers=raw_headers,
litellm_logging_obj=litellm_logging_obj,
guardrail_context=guardrail_context,
tool=self.get_listed_tool(mcp_server, name, listed_caller) if listed_tool is ... else listed_tool,
)
if "arguments" in hook_result:
arguments = hook_result["arguments"]

View file

@ -81,12 +81,14 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
outcome_wire_value,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
ListedToolsCaller,
MCPServerManager,
_caller_authorization_fans_out,
_client_forwarded_authorization_headers,
_resolve_openapi_tool_auth,
_should_strip_caller_authorization,
global_mcp_server_manager,
listed_tools_caller_for,
)
from litellm.proxy._experimental.mcp_server.oauth_utils import (
_redact_mcp_resource_url,
@ -954,6 +956,8 @@ async def _get_tools_from_mcp_servers(
request_tags: list[str] | None = None,
client_ip: str | None = None,
mcp_proxy_mode: bool = False,
*,
record_listing: bool = False,
) -> AggregateToolListing:
"""
Helper method to fetch tools from MCP servers based on server filtering criteria.
@ -964,6 +968,8 @@ async def _get_tools_from_mcp_servers(
mcp_servers: Optional list of server names/aliases to filter by
mcp_server_auth_headers: Optional dict of server-specific auth headers
oauth2_headers: Optional dict of oauth2 headers
record_listing: Record each served catalog into the caller's listed-tools slot; only a
listing actually served to the caller sets it
Returns:
AggregateToolListing: Combined tools from filtered servers plus each server's
@ -1111,12 +1117,14 @@ async def _get_tools_from_mcp_servers(
prefetched_creds=_prefetched_oauth_creds,
)
catalog_auth_header: Final = server_auth_header
if server.is_byok and server.auth_type != MCPAuth.oauth2 and server_auth_header is None:
server_auth_header = await _get_byok_credential(server, user_api_key_auth)
try:
from litellm.proxy.proxy_server import proxy_logging_obj
listed_generation: Final = global_mcp_server_manager.listed_tools_generation(server.server_id)
tools: Final = await global_mcp_server_manager._get_tools_from_server(
server=server,
mcp_auth_header=server_auth_header,
@ -1127,6 +1135,8 @@ async def _get_tools_from_mcp_servers(
user_api_key_auth=user_api_key_auth,
oauth2_headers=oauth2_headers,
proxy_logging_obj=proxy_logging_obj,
catalog_auth_header=catalog_auth_header,
record_listing=False,
)
filtered_tools = filter_tools_by_allowed_tools(tools, server)
@ -1135,6 +1145,21 @@ async def _get_tools_from_mcp_servers(
server_id=server.server_id,
user_api_key_auth=user_api_key_auth,
)
global_mcp_server_manager.record_listed_tools(
server,
[
tool.model_copy(update={"name": strip_known_server_prefix(tool.name, server)})
for tool in filtered_tools
],
ListedToolsCaller(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=catalog_auth_header,
raw_headers=raw_headers,
oauth2_headers=oauth2_headers,
),
listed_generation,
record_listing=record_listing,
)
if mcp_proxy_mode:
from litellm.proxy._experimental.mcp_server.tool_search import with_mcp_proxy_identity
@ -1458,6 +1483,8 @@ async def _list_mcp_tools(
list_tools_log_source: str | None = None,
client_ip: str | None = None,
mcp_proxy_mode: bool = False,
*,
record_listing: bool = False,
) -> AggregateToolListing:
"""
List all available MCP tools.
@ -1468,6 +1495,8 @@ async def _list_mcp_tools(
mcp_servers: Optional list of server names/aliases to filter by
mcp_server_auth_headers: Optional dict of server-specific auth headers {server_alias: auth_value}
client_ip: Client IP for IP-based server access control
record_listing: Record each served catalog into the caller's listed-tools slot; only a
listing actually served to the caller sets it
Returns:
AggregateToolListing: Combined tools from all accessible servers plus each server's
@ -1486,6 +1515,7 @@ async def _list_mcp_tools(
list_tools_log_source=list_tools_log_source,
client_ip=client_ip,
mcp_proxy_mode=mcp_proxy_mode,
record_listing=record_listing,
)
verbose_logger.debug("Successfully fetched %s tools from managed MCP servers", len(listing.tools))
return listing
@ -1798,6 +1828,7 @@ async def _list_tools_before_first_call(
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
client_ip=client_ip,
record_listing=False,
)
except Exception as e: # noqa: BLE001 # best effort: resolution below answers as it did before
verbose_logger.debug("MCP tools/call: listing %s before its first call failed: %s", server.name, e)
@ -2001,6 +2032,7 @@ async def _execute_mcp_tool(
if mcp_server is None:
mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name)
client_auth_header: Final = mcp_auth_header
if mcp_server:
standard_logging_mcp_tool_call["mcp_server_cost_info"] = (mcp_server.mcp_info or {}).get("mcp_server_cost_info")
if litellm_logging_obj:
@ -2071,6 +2103,18 @@ async def _execute_mcp_tool(
raw_headers=raw_headers,
litellm_logging_obj=litellm_logging_obj,
guardrail_context=guardrail_context,
tool=global_mcp_server_manager.get_listed_tool(
mcp_server,
original_tool_name,
listed_tools_caller_for(
mcp_server,
user_api_key_auth,
client_auth_header,
mcp_server_auth_headers,
raw_headers,
oauth2_headers,
),
),
)
# `pre_call_tool_check` may return guardrail-modified
# arguments; honor them on the local path too.
@ -2120,6 +2164,7 @@ async def _execute_mcp_tool(
arguments=arguments,
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
catalog_auth_header=client_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
@ -2139,7 +2184,8 @@ async def _execute_mcp_tool(
# not in the registry either, `_handle_local_mcp_tool` below reports
# 404 and nothing runs, so demanding a server here would turn every
# unknown tool name into a misleading 503.
if global_mcp_tool_registry.get_tool(original_tool_name) is not None:
registered_local_tool: Final = global_mcp_tool_registry.get_tool(original_tool_name)
if registered_local_tool is not None:
# `mcp_server` is None here because the tool name is not in the
# tool -> server mapping, but the name still carries a prefix
# that the server-level check above compared against the
@ -2181,6 +2227,18 @@ async def _execute_mcp_tool(
raw_headers=raw_headers,
litellm_logging_obj=litellm_logging_obj,
guardrail_context=guardrail_context,
tool=global_mcp_server_manager.get_listed_tool(
prefix_server,
original_tool_name,
listed_tools_caller_for(
prefix_server,
user_api_key_auth,
client_auth_header,
mcp_server_auth_headers,
raw_headers,
oauth2_headers,
),
),
)
if "arguments" in hook_result:
arguments = hook_result["arguments"] # pyright: ignore[reportAny] # hook returns untyped args
@ -2583,8 +2641,12 @@ async def _handle_managed_mcp_tool(
guardrail_context: Mapping[str, object] | None = None,
client_ip: str | None = None,
wire_compat: WireCompat = WireCompat.LEGACY,
*,
catalog_auth_header: str | None,
) -> CallToolResult | InputRequiredResult:
"""Handle tool execution for managed server tools"""
"""Handle tool execution for managed server tools. ``catalog_auth_header`` is the header the client
supplied, which keys the caller's catalog slot; ``mcp_auth_header`` may already be the resolved
BYOK credential."""
# Import here to avoid circular import
from litellm.proxy.proxy_server import proxy_logging_obj
@ -2594,6 +2656,7 @@ async def _handle_managed_mcp_tool(
arguments=arguments,
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
catalog_auth_header=catalog_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
@ -2712,6 +2775,7 @@ async def _execute_handle_list_tools(
log_list_tools_to_spendlogs=log_list_tools_to_spendlogs,
list_tools_log_source="mcp_protocol",
client_ip=_client_ip,
record_listing=True,
)
verbose_logger.info("MCP list_tools - Successfully returned %s tools", len(listing.tools))
if not listing.outcomes:

View file

@ -223,6 +223,7 @@ if MCP_AVAILABLE:
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
_UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES,
ListedToolsCaller,
global_mcp_server_manager,
)
from litellm.proxy._experimental.mcp_server.oauth_utils import (
@ -701,6 +702,8 @@ if MCP_AVAILABLE:
extra_headers: dict[str, str] | None,
client_ip: str | None,
proxy_logging_obj: "ProxyLogging | None",
*,
record_listing: bool,
) -> list[MCPTool]:
return await global_mcp_server_manager._get_tools_from_server(
server=server,
@ -711,11 +714,12 @@ if MCP_AVAILABLE:
client_ip=client_ip,
user_api_key_auth=user_api_key_auth,
proxy_logging_obj=proxy_logging_obj,
record_listing=record_listing,
)
async def _get_tools_for_single_server(
server,
server_auth_header,
server: MCPServer,
server_auth_header: dict[str, str] | str | None,
raw_headers: dict[str, str] | None = None,
user_api_key_auth: UserAPIKeyAuth | None = None,
extra_headers: dict[str, str] | None = None,
@ -731,31 +735,44 @@ if MCP_AVAILABLE:
"""
from litellm.proxy.proxy_server import proxy_logging_obj
tools = await _list_server_tools(
server, server_auth_header, raw_headers, user_api_key_auth, extra_headers, client_ip, proxy_logging_obj
listed_generation: Final = global_mcp_server_manager.listed_tools_generation(server.server_id)
tools: Final = await _list_server_tools(
server,
server_auth_header,
raw_headers,
user_api_key_auth,
extra_headers,
client_ip,
proxy_logging_obj,
record_listing=False,
)
if not apply_tool_filters:
return _create_tool_response_objects(tools, server)
# Always apply allowed_tools/disallowed_tools so the blacklist is
# enforced even when no allowlist is set (matches the SSE/HTTP path).
tools = filter_tools_by_allowed_tools(tools, server)
# Filter by the key's effective tool permissions through the same
# function the MCP protocol path uses (direct grants, toolset grants,
# and team/agent/org ceilings), so REST listing cannot drift from it.
# Entries here are tool names on one server, written bare by every
# writer, and dispatch compares them bare; matching a wider set of
# spellings would advertise a tool that tools/call then refuses
if user_api_key_auth:
tools = await filter_tools_by_key_team_permissions(
tools=tools,
server_filtered: Final = filter_tools_by_allowed_tools(tools, server) if apply_tool_filters else tools
served_tools: Final = (
await filter_tools_by_key_team_permissions(
tools=server_filtered,
server_id=server.server_id,
user_api_key_auth=user_api_key_auth,
)
if apply_tool_filters and user_api_key_auth
else server_filtered
)
if apply_tool_filters:
# Only a listing shaped for the caller's runtime view may set their
# listed-tools slot; the admin-only unfiltered configuration view
# must not warm it.
global_mcp_server_manager.record_listed_tools(
server,
served_tools,
ListedToolsCaller(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=server_auth_header,
raw_headers=raw_headers,
),
listed_generation,
)
return _create_tool_response_objects(tools, server)
return _create_tool_response_objects(served_tools, server)
async def fetch_pinnable_tool_catalog(
server: MCPServer, request: Request, user_api_key_dict: UserAPIKeyAuth
@ -774,6 +791,7 @@ if MCP_AVAILABLE:
await _get_user_oauth_extra_headers(server, user_api_key_dict),
IPAddressUtils.get_mcp_client_ip(request),
None,
record_listing=False,
)
scan: Final = await scan_tool_descriptions(
apply_description_overrides(upstream, server), server, proxy_logging_obj, user_api_key_dict, raw_headers

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

@ -1018,6 +1018,7 @@ if MCP_AVAILABLE:
mcp_auth_header=None,
mcp_servers=None,
mcp_server_auth_headers=None,
record_listing=True,
)
tools: Final = listing.tools
dumped_tools: Final = [tool.model_dump(by_alias=True) for tool in tools]

View file

@ -1151,6 +1151,8 @@ def _overrides_moderation_hook(callback: CustomLogger) -> bool:
_LISTED_MODEL_NAMES: Final = TypeAdapter(tuple[str, ...])
_MCP_TOOL_DESCRIPTION: Final[TypeAdapter[str | None]] = TypeAdapter(str | None)
_MCP_TOOL_INPUT_SCHEMA: Final[TypeAdapter[Mapping[str, object] | None]] = TypeAdapter(Mapping[str, object] | None)
@dataclass(frozen=True, slots=True)
@ -1475,9 +1477,14 @@ class ProxyLogging:
TypeAdapter(dict[str, object]).validate_python(guardrail_context.get("metadata") or MappingProxyType({}))
)
mcp_tool_description: Final = kwargs.get("mcp_tool_description")
mcp_input_schema: Final = kwargs.get("mcp_input_schema")
description_line: Final = f"\nDescription: {mcp_tool_description}" if mcp_tool_description else ""
mcp_tool_description: Final = request_obj.tool_description or kwargs.get("mcp_tool_description")
mcp_input_schema: Final = (
request_obj.tool_input_schema
if request_obj.tool_input_schema is not None
else kwargs.get("mcp_input_schema")
)
listing_description: Final = kwargs.get("mcp_tool_description")
description_line: Final = f"\nDescription: {listing_description}" if listing_description else ""
tool_call_content: Final = (
f"Tool: {request_obj.tool_name}{description_line}\nArguments: {request_obj.arguments}"
)
@ -1735,6 +1742,8 @@ class ProxyLogging:
tool_name=kwargs.get("name", ""),
arguments=kwargs.get("arguments", {}),
server_name=kwargs.get("server_name"),
tool_description=_MCP_TOOL_DESCRIPTION.validate_python(kwargs.get("tool_description")),
tool_input_schema=_MCP_TOOL_INPUT_SCHEMA.validate_python(kwargs.get("tool_input_schema")),
user_api_key_auth=user_api_key_auth_dict,
hidden_params=HiddenParams(),
)

View file

@ -286,6 +286,7 @@ async def aresponses_api_with_mcp(
call_params=call_params,
previous_response_id=previous_response_id,
tool_server_map=tool_server_map,
served_tools=original_mcp_tools,
**kwargs,
)
await mcp_streaming_response._create_initial_response_iterator()
@ -339,6 +340,7 @@ async def aresponses_api_with_mcp(
tool_results: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
tool_server_map=tool_server_map,
served_tools=original_mcp_tools,
tool_calls=tool_calls,
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
@ -395,6 +397,7 @@ async def aresponses_api_with_mcp(
final_response = MCPEnhancedStreamingIterator(
tool_server_map=tool_server_map,
served_tools=original_mcp_tools,
base_iterator=final_response,
mcp_events=tool_execution_events,
user_api_key_auth=user_api_key_auth,

View file

@ -435,6 +435,7 @@ async def acompletion_with_mcp(
# Execute tool calls
self.tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
tool_server_map=self.tool_server_map,
served_tools=deduplicated_mcp_tools,
tool_calls=self.tool_calls,
user_api_key_auth=self.user_api_key_auth,
mcp_auth_header=self.mcp_auth_header,
@ -609,6 +610,7 @@ async def acompletion_with_mcp(
tool_results: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
tool_server_map=tool_server_map,
tool_calls=tool_calls,
served_tools=deduplicated_mcp_tools,
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,

View file

@ -696,6 +696,7 @@ class LiteLLM_Proxy_MCP_Handler:
litellm_trace_id: str | None = None,
request_tags: list[str] | None = None,
guardrail_context: Mapping[str, object] | None = None,
served_tools: Sequence[MCPTool] | None = None,
) -> list[MCPToolResult]:
"""Execute tool calls and return results."""
from fastapi import HTTPException
@ -860,6 +861,11 @@ class LiteLLM_Proxy_MCP_Handler:
proxy_logging_obj=proxy_logging_obj,
litellm_logging_obj=litellm_logging_obj,
guardrail_context=guardrail_context,
listed_tool=(
next((tool for tool in served_tools if tool.name == tool_name), None)
if served_tools is not None
else ...
),
)
if proxy_logging_obj:
@ -1152,6 +1158,7 @@ class LiteLLM_Proxy_MCP_Handler:
call_params: Mapping[str, object],
previous_response_id: str | None,
tool_server_map: dict[str, str],
served_tools: Sequence[MCPTool] | None = None,
**kwargs,
) -> Any:
"""
@ -1181,6 +1188,7 @@ class LiteLLM_Proxy_MCP_Handler:
base_iterator=None, # Will be created internally
mcp_events=mcp_discovery_events, # Pre-generated MCP discovery events
tool_server_map=tool_server_map,
served_tools=served_tools,
mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy,
user_api_key_auth=kwargs.get("user_api_key_auth")
or kwargs.get("litellm_metadata", {}).get("user_api_key_auth"),

View file

@ -281,6 +281,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
mcp_tools_with_litellm_proxy: Sequence[Mapping[str, object]] | None = None,
user_api_key_auth: "UserAPIKeyAuth | None" = None,
original_request_params: dict[str, Any] | None = None,
served_tools: Sequence[MCPTool] | None = None,
):
# MCP setup
self.mcp_tools_with_litellm_proxy = mcp_tools_with_litellm_proxy or []
@ -300,6 +301,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
self.mcp_discovery_generated = True # Events are already generated
self.mcp_events = mcp_events # Store the initial MCP events for backward compatibility
self.tool_server_map = tool_server_map
self.served_tools = tuple(served_tools) if served_tools is not None else None
# Iterator references
self.base_iterator: BaseResponsesAPIStreamingIterator | ResponsesAPIResponse | None = (
@ -796,6 +798,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
# Execute the tools
tool_results: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
tool_server_map=self.tool_server_map,
served_tools=self.served_tools,
tool_calls=tool_calls,
user_api_key_auth=self.user_api_key_auth,
mcp_auth_header=self.mcp_auth_header,

View file

@ -429,6 +429,8 @@ class MCPPreCallRequestObject(BaseModel):
tool_name: str
arguments: dict[str, Any]
server_name: str | None = None
tool_description: str | None = None
tool_input_schema: Mapping[str, object] | None = None
user_api_key_auth: dict[str, Any] | None = None
hidden_params: HiddenParams = HiddenParams()
@ -452,6 +454,8 @@ class MCPDuringCallRequestObject(BaseModel):
tool_name: str
arguments: dict[str, Any]
server_name: str | None = None
tool_description: str | None = None
tool_input_schema: Mapping[str, object] | None = None
start_time: float | None = None
hidden_params: HiddenParams = HiddenParams()

View file

@ -216,6 +216,16 @@ JsonRpc = Mapping[str, object]
class ScriptedTool:
name: str
respond: Callable[[JsonRpc], Reply | JsonRpc]
description: str | Callable[[Mapping[str, str]], str] | None = None
input_schema: JsonRpc = field(default_factory=lambda: {"type": "object"})
def listing(self, headers: Mapping[str, str]) -> JsonRpc:
described: Final = self.description(headers) if callable(self.description) else self.description
return {
"name": self.name,
"inputSchema": self.input_schema,
**({} if described is None else {"description": described}),
}
def jsonrpc_reply(identity: object, result: JsonRpc) -> Reply:
@ -253,9 +263,7 @@ def scripted_peer(*tools: ScriptedTool) -> Iterator[McpPeer]:
},
)
if method == "tools/list":
return jsonrpc_reply(
identity, {"tools": [{"name": name, "inputSchema": {"type": "object"}} for name in by_name]}
)
return jsonrpc_reply(identity, {"tools": [tool.listing(request.headers) for tool in by_name.values()]})
if method != "tools/call":
return jsonrpc_error(identity, -32601, f"unsupported method {method}")
tool: Final = by_name.get(body["params"]["name"])

View file

@ -133,6 +133,7 @@ def wire_server(
class OwnedHTTPServer(ThreadingHTTPServer):
daemon_threads = False
request_queue_size = 128
def server_bind(self) -> None:
super().server_bind()

View file

@ -1,20 +1,29 @@
import re
import textwrap
import uuid
from collections.abc import Iterator
from pathlib import Path
from typing import Final
import pytest
from integration._support.client import Gateway
from integration._support.client import Gateway, gateway_from_environment
from integration._support.database import read_rows
from integration._support.mcp import (
ENTRY_POINTS,
EntryPoint,
McpCaller,
McpPeer,
Outcome,
PeerKind,
ScriptedTool,
peer_of,
register_mcp,
scripted_peer,
text_result,
tool_calls,
)
from integration._support.mcp_grants import SUBJECTS, Subject, grant
from integration._support.process import owned_proxy
CALLABLE: Final = {"add": {"a": 1, "b": 2}, "multiply": {"a": 2, "b": 3}}
RESULTS: Final = {"add": "3", "multiply": "6"}
@ -135,3 +144,87 @@ def test_same_tool_name_on_two_servers_routes_by_prefix(gateway: Gateway) -> Non
assert outcome.ok and outcome.text == "10", outcome.raw
assert tool_calls(first.drain()) == ()
assert [call["body"]["params"]["name"] for call in tool_calls(second.drain())] == ["add"]
_PROBE: Final = "catalog-probe"
_ECHO: Final = "catalog-echo"
_UNLISTED: Final = ""
_GUARDRAIL_CODE: Final = (
"def apply_guardrail(inputs, request_data, input_type):\n"
f' if "{_PROBE}" not in list(inputs.get("texts") or []):\n'
" return allow()\n"
' function = inputs.get("tools", [{}])[0].get("function", {})\n'
f' return block("{_ECHO}[" + function.get("description") + "]")\n'
)
_ECHO_GUARDRAIL_YAML: Final = (
"guardrails:\n"
" - guardrail_name: catalog-echo\n"
" litellm_params:\n"
" guardrail: custom_code\n"
" mode: pre_mcp_call\n"
" default_on: true\n"
" custom_code: |\n" + textwrap.indent(_GUARDRAIL_CODE, 8 * " ")
)
@pytest.fixture(scope="module")
def echo_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]:
directory: Final = tmp_path_factory.mktemp("catalog-echo")
path: Final = directory / "catalog_echo.yaml"
path.write_text((Path(__file__).resolve().parents[1] / "proxy_config.yaml").read_text() + _ECHO_GUARDRAIL_YAML)
with gateway_from_environment() as gateway, owned_proxy(gateway, directory, {}, config=path, workers=2) as rig:
yield rig
def _echoed_description(outcome: Outcome) -> str:
found: Final = re.search(rf"{_ECHO}\[(.*?)\]", outcome.raw)
assert found is not None, outcome.raw
return found.group(1)
@pytest.mark.parametrize("subject", SUBJECTS)
def test_each_subjects_call_is_evaluated_only_against_the_catalog_its_own_listing_served(
echo_rig: Gateway, subject: Subject
) -> None:
described: Final = "Adds under grant " + uuid.uuid4().hex[:8]
tool: Final = ScriptedTool("add", lambda _: text_result("3"), description=described)
with scripted_peer(tool) as peer, echo_rig.scenario() as scenario:
group: Final = "grp" + uuid.uuid4().hex[:8]
alias: Final = "cat" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias, mcp_access_groups=[group])
caller: Final = grant(
scenario, subject, (identity,), (identity,), access_group=group, allowed_tools={identity: ("add",)}
)
reach: Final = McpCaller(echo_rig, caller.key, "mcp", alias, caller.headers)
assert reach.initialize().ok
cold: Final = _echoed_description(reach.call(f"{alias}-add", {"probe": _PROBE}))
listed: Final = reach.list_tools()
assert listed.ok and f"{alias}-add" in listed.tools, listed.raw
warm: Final = _echoed_description(reach.call(f"{alias}-add", {"probe": _PROBE}))
assert (cold, warm) == (_UNLISTED, described), (cold, warm)
assert tool_calls(peer.drain()) == (), "a blocked probe reached the peer"
def test_end_users_of_one_key_share_its_catalog_slot_because_the_identity_excludes_the_end_user(
echo_rig: Gateway,
) -> None:
described: Final = "Adds for end users " + uuid.uuid4().hex[:8]
tool: Final = ScriptedTool("add", lambda _: text_result("3"), description=described)
with scripted_peer(tool) as peer, echo_rig.scenario() as scenario:
alias: Final = "eu" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias)
granted: Final = grant(scenario, "end_user", (identity,), (identity,))
first: Final = McpCaller(echo_rig, granted.key, "mcp", alias, granted.headers)
second: Final = McpCaller(
echo_rig, granted.key, "mcp", alias, {"x-litellm-end-user-id": "integration-" + uuid.uuid4().hex[:10]}
)
assert second.initialize().ok
assert _echoed_description(second.call(f"{alias}-add", {"probe": _PROBE})) == _UNLISTED
listed: Final = first.list_tools()
assert listed.ok and f"{alias}-add" in listed.tools, listed.raw
assert _echoed_description(second.call(f"{alias}-add", {"probe": _PROBE})) == described, (
"the end-user header is intentionally not part of the catalog identity: one key, one slot"
)
assert tool_calls(peer.drain()) == (), "a blocked probe reached the peer"

View file

@ -1,11 +1,23 @@
import json
import uuid
from collections.abc import Iterator
from collections.abc import Generator, Iterator, Mapping
from contextlib import contextmanager
from dataclasses import dataclass
from hashlib import sha256
from pathlib import Path
from typing import Final
import httpx
import pytest
from integration._support.client import Gateway, JsonValue, Scenario, eventually
import yaml
from integration._support.client import (
JSON_OBJECT,
Gateway,
Scenario,
eventually,
gateway_from_environment,
object_value,
)
from integration._support.database import read_rows
from integration._support.mcp import (
ENTRY_POINTS,
@ -13,10 +25,16 @@ from integration._support.mcp import (
McpCaller,
McpPeer,
Outcome,
ScriptedTool,
mcp_peer,
register_mcp,
scripted_peer,
text_result,
tool_calls,
)
from integration._support.process import owned_proxy
from integration._support.wire import Reply, Request, Wire, wire_server
from pydantic import JsonValue
DEFAULT_COST: Final = 0.25
ADD_COST: Final = 0.5
@ -191,3 +209,472 @@ def test_guardrail_removal_stops_blocking_without_restart(gateway: Gateway) -> N
lambda calls: len(calls) >= 1,
seconds=40,
)
MASK_ME: Final = "mask-integration-secret"
MASKED: Final = "[MASKED]"
COUNT_MISMATCH: Final = "count-mismatch-marker"
LOOKUP_DESCRIPTION: Final = "Look up one record"
LOOKUP_SCHEMA: Final = {
"type": "object",
"properties": {"record": {"type": "string", "description": "record identifier"}},
}
SPEND_ROW: Final = 'SELECT status, metadata FROM "LiteLLM_SpendLogs" WHERE request_id = %s'
_RECORDER_CODE: Final = """\
import os
import httpx
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_logger import CustomLogger
SINK = "{sink}/native"
def _record(stage, data, call_type):
logging_obj = data.get("litellm_logging_obj")
return {{
"stage": stage,
"pid": os.getpid(),
"call_type": call_type,
"litellm_call_id": None if logging_obj is None else logging_obj.litellm_call_id,
"messages": data.get("messages"),
"mcp_tool_name": data.get("mcp_tool_name"),
"mcp_arguments": data.get("mcp_arguments"),
"mcp_tool_description": data.get("mcp_tool_description"),
"mcp_input_schema": data.get("mcp_input_schema"),
}}
async def _post(record):
async with httpx.AsyncClient(timeout=5) as client:
await client.post(SINK, json=record)
class HookRecorder(CustomLogger):
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
await _post(_record("pre", data, call_type))
async def async_moderation_hook(self, data, user_api_key_dict, call_type):
await _post(_record("during", data, call_type))
async def async_post_mcp_tool_call_hook(self, kwargs, response_obj, start_time, end_time):
await _post(
{{
"stage": "post",
"pid": os.getpid(),
"litellm_call_id": kwargs.get("litellm_call_id"),
"tool": kwargs.get("mcp_tool_call_metadata"),
"content": [item.model_dump() for item in response_obj.mcp_tool_call_response],
}}
)
class SinkGuardrail(CustomGuardrail):
def __init__(self, api_base, **kwargs):
super().__init__(**kwargs)
self.api_base = api_base
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
payload = {{
"pid": os.getpid(),
"input_type": input_type,
"litellm_call_id": None if logging_obj is None else logging_obj.litellm_call_id,
"texts": inputs.get("texts"),
"tools": inputs.get("tools"),
"structured_messages": inputs.get("structured_messages"),
"mcp_tool_name": request_data.get("mcp_tool_name"),
}}
async with httpx.AsyncClient(timeout=5) as client:
verdict = (await client.post(self.api_base, json=payload)).json()
return {{**inputs, "texts": verdict["texts"]}}
recorder = HookRecorder()
"""
@dataclass(frozen=True, slots=True)
class Sunk:
target: str
body: dict[str, JsonValue]
@dataclass(frozen=True, slots=True)
class HooksRig:
gateway: Gateway
sibling: Gateway
sink: Wire
guardrail: str
def sunk(self) -> tuple[Sunk, ...]:
return tuple(Sunk(request.target, JSON_OBJECT.validate_json(request.body)) for request in self.sink.drain())
def _guardrail_sink(request: Request) -> Reply:
if not request.target.startswith("/guardrail"):
return Reply()
texts: Final = JSON_OBJECT.validate_json(request.body).get("texts")
assert isinstance(texts, list), texts
masked: Final = [str(text).replace(MASK_ME, MASKED) for text in texts]
extra: Final = ["extra"] if any(COUNT_MISMATCH in text for text in masked) else []
return Reply(body=json.dumps({"texts": [*masked, *extra]}).encode())
@pytest.fixture(scope="module")
def hooks_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[HooksRig]:
directory: Final = tmp_path_factory.mktemp("guardrail-payloads")
guardrail: Final = "sink" + uuid.uuid4().hex[:8]
with wire_server(_guardrail_sink) as sink:
(directory / "hook_recorder.py").write_text(_RECORDER_CODE.format(sink=sink.url))
config: Final = JSON_OBJECT.validate_python(
yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
)
config["guardrails"] = [
{
"guardrail_name": guardrail,
"litellm_params": {
"guardrail": "hook_recorder.SinkGuardrail",
"mode": ["pre_mcp_call", "post_mcp_call"],
"default_on": True,
"api_base": f"{sink.url}/guardrail",
},
}
]
config["litellm_settings"] = {
**object_value(config["litellm_settings"]),
"callbacks": ["hook_recorder.recorder"],
}
path: Final = directory / "config.yaml"
path.write_text(yaml.safe_dump(config))
with (
gateway_from_environment() as gateway,
owned_proxy(gateway, directory, {"KEEPALIVE_TIMEOUT": "120"}, config=path, workers=2) as candidate,
owned_proxy(gateway, directory, {}, config=path) as sibling,
):
yield HooksRig(candidate, sibling, sink, guardrail)
def _worker(gateway: Gateway) -> int:
response: Final = gateway.client.get("/debug/memory/summary", headers={"x-litellm-api-key": gateway.key})
assert response.status_code == 200, response.text
worker: Final = JSON_OBJECT.validate_json(response.content)["worker_pid"]
assert isinstance(worker, int), response.text
return worker
@contextmanager
def _pinned(gateway: Gateway) -> Generator[tuple[Gateway, int], None, None]:
limits: Final = httpx.Limits(max_connections=1, max_keepalive_connections=1, keepalive_expiry=120)
with httpx.Client(base_url=gateway.client.base_url, timeout=15, trust_env=False, limits=limits) as client:
pinned: Final = Gateway(client, gateway.key, gateway.upstream_url)
yield pinned, _worker(pinned)
def _lookup_tool(result: str = "found") -> ScriptedTool:
return ScriptedTool(
"lookup", lambda _: text_result(result), description=LOOKUP_DESCRIPTION, input_schema=LOOKUP_SCHEMA
)
def _generic(sunk: tuple[Sunk, ...], input_type: str) -> tuple[dict[str, JsonValue], ...]:
return tuple(item.body for item in sunk if item.target == "/guardrail" and item.body["input_type"] == input_type)
def _native(sunk: tuple[Sunk, ...], stage: str) -> tuple[dict[str, JsonValue], ...]:
return tuple(item.body for item in sunk if item.target == "/native" and item.body["stage"] == stage)
def _only(records: tuple[dict[str, JsonValue], ...]) -> dict[str, JsonValue]:
assert len(records) == 1, records
return records[0]
def _scan(sunk: tuple[Sunk, ...], call_id: JsonValue) -> dict[str, JsonValue]:
return _only(tuple(record for record in _generic(sunk, "request") if record["litellm_call_id"] == call_id))
def _texts(content: JsonValue) -> list[JsonValue]:
assert isinstance(content, list), content
return [object_value(item)["text"] for item in content]
def _has_lookup(listing: Outcome) -> bool:
return any(tool.endswith("lookup") for tool in listing.tools)
def _spend_row(call_id: JsonValue) -> dict[str, JsonValue]:
assert isinstance(call_id, str), call_id
rows: Final = eventually(lambda: read_rows(SPEND_ROW, (call_id,)), lambda found: len(found) == 1, seconds=70)
return rows[0]
def _synthetic_message(name: str, arguments: Mapping[str, str]) -> list[dict[str, str]]:
return [{"role": "user", "content": f"Tool: {name}\nArguments: {dict(arguments)}"}]
def test_generic_sink_and_native_hooks_receive_listed_metadata_on_typed_keys_with_the_message_bytes_unchanged(
hooks_rig: HooksRig,
) -> None:
with (
scripted_peer(_lookup_tool()) as peer,
_pinned(hooks_rig.gateway) as (pinned, worker),
pinned.scenario() as scenario,
):
alias: Final = "payload" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias)
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
caller: Final = McpCaller(pinned, key, "mcp", headers={"x-mcp-servers": alias})
assert eventually(caller.initialize, lambda outcome: outcome.ok, seconds=30).ok
hooks_rig.sunk()
listed: Final = eventually(caller.list_tools, _has_lookup, seconds=30)
name: Final = next(tool for tool in listed.tools if tool.endswith("lookup"))
scans: Final = _generic(hooks_rig.sunk(), "request")
assert scans and all(scan["texts"] == [LOOKUP_DESCRIPTION, "record identifier"] for scan in scans), scans
arguments: Final = {"record": "r-1"}
outcome: Final = caller.call(name, arguments)
assert outcome.text == "found", outcome.raw
assert _worker(pinned) == worker
sunk: Final = hooks_rig.sunk()
pre: Final = _only(_native(sunk, "pre"))
call_id: Final = pre["litellm_call_id"]
generic: Final = _scan(sunk, call_id)
assert generic["texts"] == [LOOKUP_DESCRIPTION, "record identifier", "r-1"], generic
assert generic["tools"] == [
{
"type": "function",
"function": {
"name": "lookup",
"description": LOOKUP_DESCRIPTION,
"parameters": {**LOOKUP_SCHEMA, "additionalProperties": False},
"strict": False,
},
}
], generic
response: Final = _only(_generic(sunk, "response"))
assert (response["texts"], response["pid"]) == (["found"], worker), response
during: Final = _only(_native(sunk, "during"))
post: Final = _only(_native(sunk, "post"))
assert all(record["pid"] == worker for record in (pre, during, post)), sunk
assert all(record["litellm_call_id"] == call_id for record in (pre, during, post)), sunk
assert pre["call_type"] == "call_mcp_tool" and during["call_type"] == "call_mcp_tool", sunk
assert pre["messages"] == _synthetic_message("lookup", arguments), pre
assert during["messages"] == _synthetic_message("lookup", arguments), during
assert (pre["mcp_tool_name"], pre["mcp_arguments"]) == ("lookup", arguments), pre
assert pre["mcp_tool_description"] == LOOKUP_DESCRIPTION, pre
assert pre["mcp_input_schema"] == LOOKUP_SCHEMA, pre
assert (during["mcp_tool_description"], during["mcp_input_schema"]) == (None, None), during
assert _texts(post["content"]) == ["found"], post
row: Final = _spend_row(call_id)
assert row["status"] == "success", row
assert _tool_metadata(row)["name"] == "lookup", row
metadata: Final = row["metadata"]
assert isinstance(metadata, dict) and metadata["applied_guardrails"] == [hooks_rig.guardrail], metadata
def test_pre_call_mask_reaches_the_peer_and_post_call_mask_reaches_the_caller_on_one_call_id(
hooks_rig: HooksRig,
) -> None:
with (
scripted_peer(_lookup_tool(f"found {MASK_ME}")) as peer,
_pinned(hooks_rig.gateway) as (pinned, worker),
pinned.scenario() as scenario,
):
alias: Final = "mask" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias)
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
caller: Final = McpCaller(pinned, key, "mcp", headers={"x-mcp-servers": alias})
assert eventually(caller.initialize, lambda outcome: outcome.ok, seconds=30).ok
name: Final = next(
tool for tool in eventually(caller.list_tools, _has_lookup, seconds=30).tools if "lookup" in tool
)
peer.drain()
hooks_rig.sunk()
outcome: Final = caller.call(name, {"record": MASK_ME})
assert outcome.text == f"found {MASKED}", outcome.raw
assert _worker(pinned) == worker
reached: Final = tool_calls(peer.drain())
assert len(reached) == 1, reached
params: Final = object_value(JSON_OBJECT.validate_python(reached[0]["body"])["params"])
assert params["arguments"] == {"record": MASKED}, params
sunk: Final = hooks_rig.sunk()
generic: Final = _scan(sunk, _only(_native(sunk, "pre"))["litellm_call_id"])
scanned: Final = generic["texts"]
assert isinstance(scanned, list) and scanned[-1] == MASK_ME and MASKED not in scanned, generic
assert _only(_generic(sunk, "response"))["texts"] == [f"found {MASK_ME}"], sunk
during: Final = _only(_native(sunk, "during"))
assert during["mcp_arguments"] == {"record": MASKED}, during
assert during["messages"] == _synthetic_message("lookup", {"record": MASKED}), during
assert _only(_native(sunk, "post"))["litellm_call_id"] == generic["litellm_call_id"], sunk
assert _spend_row(generic["litellm_call_id"])["status"] == "success"
def test_call_time_description_and_schema_come_from_the_catalog_of_the_worker_that_served_the_listing(
hooks_rig: HooksRig,
) -> None:
with (
scripted_peer(_lookup_tool()) as peer,
hooks_rig.gateway.scenario() as scenario,
_pinned(hooks_rig.sibling) as (second, second_worker),
):
alias: Final = "local" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias)
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
other: Final = McpCaller(second, key, "mcp", headers={"x-mcp-servers": alias})
assert eventually(other.initialize, lambda outcome: outcome.ok, seconds=60).ok
with _pinned(hooks_rig.gateway) as (first, first_worker):
assert first_worker != second_worker
lister: Final = McpCaller(first, key, "mcp", headers={"x-mcp-servers": alias})
assert eventually(lister.initialize, lambda outcome: outcome.ok, seconds=30).ok
name: Final = next(
tool for tool in eventually(lister.list_tools, _has_lookup, seconds=30).tools if "lookup" in tool
)
hooks_rig.sunk()
assert other.call(name, {"record": "r-2"}).text == "found"
assert _worker(second) == second_worker
elsewhere: Final = hooks_rig.sunk()
unlisted: Final = _only(_native(elsewhere, "pre"))
assert unlisted["pid"] == second_worker, unlisted
assert (unlisted["mcp_tool_description"], unlisted["mcp_input_schema"]) == (None, None), unlisted
assert _scan(elsewhere, unlisted["litellm_call_id"])["texts"] == ["r-2"], elsewhere
assert lister.call(name, {"record": "r-3"}).text == "found"
assert _worker(first) == first_worker
at_lister: Final = hooks_rig.sunk()
listed: Final = _only(_native(at_lister, "pre"))
assert listed["pid"] == first_worker, listed
assert (listed["mcp_tool_description"], listed["mcp_input_schema"]) == (LOOKUP_DESCRIPTION, LOOKUP_SCHEMA)
assert _scan(at_lister, listed["litellm_call_id"])["texts"] == [
LOOKUP_DESCRIPTION,
"record identifier",
"r-3",
]
assert _has_lookup(other.list_tools())
hooks_rig.sunk()
assert other.call(name, {"record": "r-4"}).text == "found"
assert _worker(second) == second_worker
populated: Final = _only(_native(hooks_rig.sunk(), "pre"))
assert populated["pid"] == second_worker, populated
assert populated["mcp_tool_description"] == LOOKUP_DESCRIPTION, populated
@pytest.mark.parametrize("entry", ("mcp", "rest"))
@pytest.mark.parametrize("listed", (False, True))
@pytest.mark.parametrize("record", ("r-1", MASK_ME))
def test_long_descriptions_do_not_refuse_small_tpm_calls_or_change_masked_message_bytes(
hooks_rig: HooksRig, entry: EntryPoint, listed: bool, record: str
) -> None:
description: Final = "Gateway tool metadata. " * 300
tool: Final = ScriptedTool(
"lookup", lambda _: text_result("found"), description=description, input_schema=LOOKUP_SCHEMA
)
with (
scripted_peer(tool) as peer,
_pinned(hooks_rig.gateway) as (pinned, worker),
pinned.scenario() as scenario,
):
alias: Final = "quota" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias)
key: Final = scenario.key(object_permission={"mcp_servers": [identity]}, tpm_limit=64)
caller: Final = McpCaller(pinned, key, entry, headers={"x-mcp-servers": alias})
assert eventually(caller.initialize, lambda outcome: outcome.ok, seconds=30).ok
if listed:
listing: Final = eventually(lambda: caller.list_tools(identity), _has_lookup, seconds=30)
assert _has_lookup(listing), listing.raw
hooks_rig.sunk()
peer.drain()
arguments: Final = {"record": record}
outcome: Final = caller.call(f"{alias}-lookup", arguments, identity)
assert outcome.ok and outcome.text == "found", outcome.raw
assert _worker(pinned) == worker
reached: Final = tool_calls(peer.drain())
assert len(reached) == 1, reached
params: Final = object_value(JSON_OBJECT.validate_python(reached[0]["body"])["params"])
masked_arguments: Final = {"record": record.replace(MASK_ME, MASKED)}
assert params["arguments"] == masked_arguments, params
sunk: Final = hooks_rig.sunk()
pre: Final = _only(tuple(hook for hook in _native(sunk, "pre") if hook["mcp_tool_name"] == "lookup"))
during: Final = _only(tuple(hook for hook in _native(sunk, "during") if hook["mcp_tool_name"] == "lookup"))
assert pre["messages"] == _synthetic_message("lookup", arguments), pre
assert during["messages"] == _synthetic_message("lookup", masked_arguments), during
assert _spend_row(pre["litellm_call_id"])["status"] == "success"
def _nested_schema(levels: int) -> dict[str, object]:
if levels == 0:
return {"type": "string", "description": "deepest leaf"}
return {"type": "object", "properties": {"a": _nested_schema(levels - 1)}}
def test_a_schema_past_the_scan_depth_is_not_published_while_one_at_the_limit_is_scanned_on_the_call(
hooks_rig: HooksRig,
) -> None:
shallow: Final = ScriptedTool("shallow", lambda _: text_result("found"), input_schema=_nested_schema(49))
deep: Final = ScriptedTool("deep", lambda _: text_result("found"), input_schema=_nested_schema(50))
with (
scripted_peer(shallow, deep) as peer,
_pinned(hooks_rig.gateway) as (pinned, worker),
pinned.scenario() as scenario,
):
alias: Final = "depth" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias)
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
caller: Final = McpCaller(pinned, key, "mcp", headers={"x-mcp-servers": alias})
assert eventually(caller.initialize, lambda outcome: outcome.ok, seconds=30).ok
listing: Final = eventually(caller.list_tools, lambda outcome: len(outcome.tools) > 0, seconds=30)
assert listing.tools == (f"{alias}-shallow",), listing.raw
assert _worker(pinned) == worker
hooks_rig.sunk()
assert caller.call(f"{alias}-shallow", {"record": "r-1"}).text == "found"
scanned: Final = _only(_generic(hooks_rig.sunk(), "request"))
assert scanned["texts"] == ["deepest leaf", "r-1"], scanned
unpublished: Final = caller.call(f"{alias}-deep", {"record": "r-2"})
assert _worker(pinned) == worker
assert unpublished.error is not None and "Tool 'deep' not found" in unpublished.raw, unpublished.raw
reached: Final = tool_calls(peer.drain())
assert [object_value(JSON_OBJECT.validate_python(call["body"])["params"])["name"] for call in reached] == [
"shallow"
]
relisted: Final = _generic(hooks_rig.sunk(), "request")
assert all((scan["mcp_tool_name"], scan["litellm_call_id"]) == ("shallow", None) for scan in relisted), relisted
rows: Final = _rows(key, 3)
assert [(row["call_type"], row["status"]) for row in rows] == [
("list_mcp_tools", "success"),
("call_mcp_tool", "success"),
("call_mcp_tool", "failure"),
], rows
assert [_tool_metadata(row)["name"] for row in rows[1:]] == ["shallow", "deep"], rows
failed: Final = object_value(JSON_OBJECT.validate_python(rows[2]["metadata"])["error_information"])
assert failed["error_message"] == "404: Tool 'deep' not found", failed
def test_an_adapter_returning_the_wrong_number_of_texts_fails_closed_before_the_peer(hooks_rig: HooksRig) -> None:
with (
scripted_peer(_lookup_tool()) as peer,
_pinned(hooks_rig.gateway) as (pinned, worker),
pinned.scenario() as scenario,
):
alias: Final = "count" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias)
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
caller: Final = McpCaller(pinned, key, "mcp", headers={"x-mcp-servers": alias})
assert eventually(caller.initialize, lambda outcome: outcome.ok, seconds=30).ok
name: Final = next(
tool for tool in eventually(caller.list_tools, _has_lookup, seconds=30).tools if "lookup" in tool
)
assert _worker(pinned) == worker
hooks_rig.sunk()
blocked: Final = caller.call(name, {"record": COUNT_MISMATCH})
assert _worker(pinned) == worker
assert blocked.error is not None, blocked.raw
assert (
"guardrail returned 4 texts for 3 MCP tool strings, so the redaction cannot be mapped back" in blocked.raw
)
assert tool_calls(peer.drain()) == (), "the blocked call reached the peer"
sunk: Final = hooks_rig.sunk()
assert _only(_generic(sunk, "request"))["texts"] == [LOOKUP_DESCRIPTION, "record identifier", COUNT_MISMATCH]
assert _native(sunk, "post") == (), sunk
rows: Final = _rows(key, 2)
assert [(row["call_type"], row["status"]) for row in rows] == [
("list_mcp_tools", "success"),
("call_mcp_tool", "failure"),
], rows
assert _tool_metadata(rows[1])["arguments"] == {"record": COUNT_MISMATCH}, rows[1]

View file

@ -1,21 +1,32 @@
import base64
import re
import textwrap
import uuid
from collections.abc import Iterator
from pathlib import Path
from typing import Final
import pytest
from integration._support.client import Gateway, eventually
from integration._support.client import Gateway, Scenario, eventually, gateway_from_environment
from integration._support.database import read_rows
from integration._support.mcp import (
ENTRY_POINTS,
EntryPoint,
McpCaller,
McpPeer,
Outcome,
ScriptedTool,
call_tool,
mcp_peer,
register_mcp,
scripted_peer,
text_result,
tool_calls,
tool_names,
)
from integration._support.oauth_server import oauth_server
from integration._support.process import owned_proxy
from pydantic import TypeAdapter
ADD: Final = {"a": 2, "b": 3}
STATIC_MODES: Final = (
@ -192,3 +203,288 @@ def test_byok_server_uses_the_calling_users_stored_credential_and_fails_closed_w
assert removed.status_code in (200, 204), removed.text
eventually(lambda: call_tool(gateway, owner_key, identity, name, ADD), lambda value: value.status_code == 401)
assert tool_calls(peer.drain()) == ()
def _listings(peer: McpPeer) -> tuple[dict[str, object], ...]:
return tuple(
item
for item in peer.drain()
if isinstance(item.get("body"), dict) and item["body"].get("method") == "tools/list"
)
def test_oauth2_byok_listing_sends_the_minted_token_not_the_users_stored_secret(gateway: Gateway) -> None:
with mcp_peer() as peer, oauth_server() as auth, gateway.scenario() as scenario:
alias: Final = "cc" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(
scenario,
peer,
alias,
auth_type="oauth2",
oauth2_flow="client_credentials",
is_byok=True,
token_url=auth.issuer + "/token",
credentials={"client_id": "cc-client", "client_secret": "cc-secret-" + uuid.uuid4().hex},
)
owner: Final = scenario.user()
owner_key: Final = scenario.key(user_id=owner, object_permission={"mcp_servers": [identity]})
secret: Final = "byok-" + uuid.uuid4().hex
stored: Final = gateway.client.post(
f"/v1/mcp/server/{identity}/user-credential",
json={"credential": secret},
headers={"x-litellm-api-key": owner_key},
)
assert stored.status_code in (200, 201), stored.text
scenario.cleanups.callback(
gateway.client.delete,
f"/v1/mcp/server/{identity}/user-credential",
headers={"x-litellm-api-key": owner_key},
)
peer.drain()
auth.drain()
response: Final = gateway.client.get(
"/mcp-rest/tools/list", params={"server_id": identity}, headers={"x-litellm-api-key": owner_key}
)
assert response.status_code == 200, response.text
assert "add" in {tool["name"] for tool in response.json()["tools"]}, response.text
assert [request["grant_type"] for request in auth.token_requests()] == ["client_credentials"]
listings: Final = _listings(peer)
assert len(listings) == 1, listings
sent: Final = _header(listings[0], b"authorization")
assert sent is not None and auth.is_live(sent.decode().removeprefix("Bearer ")), sent
assert secret.encode() not in sent, "stored BYOK secret replaced the minted token on tools/list"
@pytest.mark.parametrize(("auth_type", "header", "shape"), STATIC_MODES[:2])
def test_byok_rest_listing_sends_the_servers_static_credential_not_the_users_stored_secret(
gateway: Gateway, auth_type: str, header: bytes, shape: str
) -> None:
with mcp_peer() as peer, gateway.scenario() as scenario:
alias: Final = "byok" + uuid.uuid4().hex[:8]
static: Final = "static-" + uuid.uuid4().hex
identity: Final = register_mcp(
scenario, peer, alias, auth_type=auth_type, is_byok=True, credentials={"auth_value": static}
)
owner: Final = scenario.user()
owner_key: Final = scenario.key(user_id=owner, object_permission={"mcp_servers": [identity]})
secret: Final = "byok-" + uuid.uuid4().hex
stored: Final = gateway.client.post(
f"/v1/mcp/server/{identity}/user-credential",
json={"credential": secret},
headers={"x-litellm-api-key": owner_key},
)
assert stored.status_code in (200, 201), stored.text
scenario.cleanups.callback(
gateway.client.delete,
f"/v1/mcp/server/{identity}/user-credential",
headers={"x-litellm-api-key": owner_key},
)
peer.drain()
response: Final = gateway.client.get(
"/mcp-rest/tools/list", params={"server_id": identity}, headers={"x-litellm-api-key": owner_key}
)
assert response.status_code == 200, response.text
assert "add" in {tool["name"] for tool in response.json()["tools"]}, response.text
listings: Final = _listings(peer)
assert len(listings) == 1, listings
assert _header(listings[0], header) == shape.format(secret=static, basic="").encode(), listings[0]["headers"]
peer.drain()
called: Final = call_tool(gateway, owner_key, identity, f"{alias}-add", ADD)
assert called.status_code == 200, called.text
assert _header(_one_call(peer), header) == shape.format(secret=secret, basic="").encode()
def test_deprecated_string_x_mcp_auth_lists_a_byok_server_for_a_key_without_a_user(gateway: Gateway) -> None:
with mcp_peer() as peer, gateway.scenario() as scenario:
alias: Final = "byok" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias, auth_type="bearer_token", is_byok=True)
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
peer.drain()
response: Final = gateway.client.get(
"/mcp-rest/tools/list",
params={"server_id": identity},
headers={"x-litellm-api-key": key, "x-mcp-auth": "Bearer hdr"},
)
assert response.status_code == 200, response.text
names: Final = {tool["name"] for tool in response.json()["tools"]}
assert "add" in names, names
listings: Final = _listings(peer)
assert len(listings) == 1, listings
assert listings[0]["headers"].get(b"authorization") == b"Bearer hdr"
_PROBE: Final = "catalog-probe"
_ECHO: Final = "catalog-echo"
_UNLISTED: Final = ""
_GUARDRAIL_CODE: Final = (
"def apply_guardrail(inputs, request_data, input_type):\n"
f' if "{_PROBE}" not in list(inputs.get("texts") or []):\n'
" return allow()\n"
' function = inputs.get("tools", [{}])[0].get("function", {})\n'
f' return block("{_ECHO}[" + function.get("description") + "]")\n'
)
_ECHO_GUARDRAIL_YAML: Final = (
"guardrails:\n"
" - guardrail_name: catalog-echo\n"
" litellm_params:\n"
" guardrail: custom_code\n"
" mode: pre_mcp_call\n"
" default_on: true\n"
" custom_code: |\n" + textwrap.indent(_GUARDRAIL_CODE, 8 * " ")
)
@pytest.fixture(scope="module")
def echo_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]:
directory: Final = tmp_path_factory.mktemp("catalog-echo")
path: Final = directory / "catalog_echo.yaml"
path.write_text((Path(__file__).resolve().parents[1] / "proxy_config.yaml").read_text() + _ECHO_GUARDRAIL_YAML)
with gateway_from_environment() as gateway, owned_proxy(gateway, directory, {}, config=path, workers=2) as rig:
yield rig
def _echoed_description(outcome: Outcome) -> str:
found: Final = re.search(rf"{_ECHO}\[(.*?)\]", outcome.raw)
assert found is not None, outcome.raw
return found.group(1)
_PROBE_ARGUMENTS: Final = {"probe": _PROBE}
_HEADERS: Final = TypeAdapter(dict[str, str])
def _store_byok_credential(scenario: Scenario, identity: str, key: str, secret: str) -> None:
stored: Final = scenario.gateway.client.post(
f"/v1/mcp/server/{identity}/user-credential", json={"credential": secret}, headers={"x-litellm-api-key": key}
)
assert stored.status_code in (200, 201), stored.text
scenario.cleanups.callback(
scenario.gateway.client.delete, f"/v1/mcp/server/{identity}/user-credential", headers={"x-litellm-api-key": key}
)
def test_rotating_the_credential_drops_the_callers_listing_until_it_lists_again(echo_rig: Gateway) -> None:
with mcp_peer() as peer, echo_rig.scenario() as scenario:
first: Final = "cred-" + uuid.uuid4().hex
second: Final = "cred-" + uuid.uuid4().hex
alias: Final = "rot" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(
scenario, peer, alias, auth_type="bearer_token", credentials={"auth_value": first}
)
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
caller: Final = McpCaller(echo_rig, key, "mcp", alias)
assert caller.list_tools().ok
assert _echoed_description(caller.call(f"{alias}-add", _PROBE_ARGUMENTS)) == "Add two integers"
rotated: Final = echo_rig.request(
"PUT", "/v1/mcp/server", {"server_id": identity, "credentials": {"auth_value": second}}
)
assert rotated.status_code == 202, rotated.text
eventually(
lambda: _echoed_description(caller.call(f"{alias}-add", _PROBE_ARGUMENTS)), lambda seen: seen == _UNLISTED
)
peer.drain()
assert caller.list_tools().ok
assert _echoed_description(caller.call(f"{alias}-add", _PROBE_ARGUMENTS)) == "Add two integers"
relisted: Final = _listings(peer)
assert len(relisted) == 1, relisted
assert _header(relisted[0], b"authorization") == f"Bearer {second}".encode(), relisted[0]["headers"]
def test_byok_callers_are_evaluated_against_their_own_listing_and_the_stored_secret_never_keys_the_slot(
echo_rig: Gateway,
) -> None:
with mcp_peer() as peer, echo_rig.scenario() as scenario:
alias: Final = "byok" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias, auth_type="api_key", is_byok=True)
owner_key: Final = scenario.key(user_id=scenario.user(), object_permission={"mcp_servers": [identity]})
stranger_key: Final = scenario.key(user_id=scenario.user(), object_permission={"mcp_servers": [identity]})
owner_secret: Final = "byok-" + uuid.uuid4().hex
replacement: Final = "byok-" + uuid.uuid4().hex
_store_byok_credential(scenario, identity, owner_key, owner_secret)
_store_byok_credential(scenario, identity, stranger_key, "byok-" + uuid.uuid4().hex)
owner: Final = McpCaller(echo_rig, owner_key, "mcp", alias)
stranger: Final = McpCaller(echo_rig, stranger_key, "mcp", alias)
peer.drain()
assert owner.list_tools().ok
listings: Final = _listings(peer)
assert [_header(item, b"x-api-key") for item in listings] == [owner_secret.encode()], listings
own: Final = _echoed_description(owner.call(f"{alias}-add", _PROBE_ARGUMENTS))
other: Final = _echoed_description(stranger.call(f"{alias}-add", _PROBE_ARGUMENTS))
assert (own, other) == ("Add two integers", _UNLISTED), (own, other)
_store_byok_credential(scenario, identity, owner_key, replacement)
assert _echoed_description(owner.call(f"{alias}-add", _PROBE_ARGUMENTS)) == "Add two integers", (
"the slot is keyed by the client-supplied header, never by the stored credential"
)
sent: Final = eventually(
lambda: (owner.call(f"{alias}-add", ADD).ok, tool_calls(peer.drain())),
lambda value: any(_header(call, b"x-api-key") == replacement.encode() for call in value[1]),
)
assert sent[0], sent
def test_callers_with_different_server_scoped_auth_headers_are_evaluated_against_their_own_listings(
echo_rig: Gateway,
) -> None:
tool: Final = ScriptedTool(
"add",
lambda _: text_result("3"),
description=lambda headers: "Adds for " + headers.get("authorization", "nobody"),
)
with scripted_peer(tool) as peer, echo_rig.scenario() as scenario:
alias: Final = "scoped" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias)
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
acme_token: Final = "acme-" + uuid.uuid4().hex
globex_token: Final = "globex-" + uuid.uuid4().hex
acme: Final = McpCaller(echo_rig, key, "mcp", alias, {f"x-mcp-{alias}-authorization": f"Bearer {acme_token}"})
globex: Final = McpCaller(
echo_rig, key, "mcp", alias, {f"x-mcp-{alias}-authorization": f"Bearer {globex_token}"}
)
assert acme.list_tools().ok and globex.list_tools().ok
seen: Final = (
_echoed_description(acme.call(f"{alias}-add", _PROBE_ARGUMENTS)),
_echoed_description(globex.call(f"{alias}-add", _PROBE_ARGUMENTS)),
)
assert seen == (f"Adds for Bearer {acme_token}", f"Adds for Bearer {globex_token}"), seen
assert tool_calls(peer.drain()) == (), "a blocked probe reached the peer"
def test_deprecated_string_x_mcp_auth_callers_on_a_user_less_key_own_separate_listings(echo_rig: Gateway) -> None:
tool: Final = ScriptedTool(
"add",
lambda _: text_result("3"),
description=lambda headers: "Adds for " + headers.get("authorization", "nobody"),
)
with scripted_peer(tool) as peer, echo_rig.scenario() as scenario:
alias: Final = "legacy" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias, auth_type="bearer_token")
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
first_token: Final = "first-" + uuid.uuid4().hex
second_token: Final = "second-" + uuid.uuid4().hex
first: Final = McpCaller(echo_rig, key, "mcp", alias, {"x-mcp-auth": f"Bearer {first_token}"})
second: Final = McpCaller(echo_rig, key, "mcp", alias, {"x-mcp-auth": f"Bearer {second_token}"})
assert _echoed_description(first.call(f"{alias}-add", _PROBE_ARGUMENTS)) == _UNLISTED
peer.drain()
assert first.list_tools().ok
listings: Final = _listings(peer)
assert len(listings) == 1, listings
listed_with: Final = _HEADERS.validate_python(listings[0]["headers"])
assert listed_with.get("authorization") == f"Bearer {first_token}", listed_with
warm: Final = eventually(
lambda: _echoed_description(first.call(f"{alias}-add", _PROBE_ARGUMENTS)),
lambda seen: seen != _UNLISTED,
)
assert warm == f"Adds for Bearer {first_token}", warm
assert _echoed_description(second.call(f"{alias}-add", _PROBE_ARGUMENTS)) == _UNLISTED
assert second.list_tools().ok
seen: Final = eventually(
lambda: (
_echoed_description(first.call(f"{alias}-add", _PROBE_ARGUMENTS)),
_echoed_description(second.call(f"{alias}-add", _PROBE_ARGUMENTS)),
),
lambda pair: _UNLISTED not in pair,
)
assert seen == (f"Adds for Bearer {first_token}", f"Adds for Bearer {second_token}"), seen
assert tool_calls(peer.drain()) == (), "a blocked probe reached the peer"

View file

@ -1,7 +1,9 @@
import functools
import json
import uuid
from collections.abc import Mapping
from contextlib import ExitStack
from hashlib import sha256
from pathlib import Path
from typing import Final
@ -14,15 +16,27 @@ from integration._support.client import Gateway, eventually
from integration._support.database import read_rows
from integration._support.generation import LIFECYCLE_SETTINGS, bounded_http_requests
from integration._support.mcp import (
JsonRpc,
McpCaller,
Outcome,
ScriptedTool,
call_tool,
mcp_peer,
register_mcp,
scripted_peer,
text_result,
tool_calls,
tool_names,
)
from integration._support.process import owned_proxy
from pydantic import TypeAdapter
_SPEND_NONCES: Final = (
"SELECT status, metadata->'mcp_tool_call_metadata'->'arguments'->>'nonce' AS nonce"
' FROM "LiteLLM_SpendLogs" WHERE api_key = %s AND call_type = %s'
)
_OBJECTS: Final = TypeAdapter(Mapping[str, object])
_STRINGS: Final = TypeAdapter(Mapping[str, str])
@pytest.mark.covers("mcp.call_tool.saved_headers.reach_actual_transport")
@ -463,9 +477,7 @@ def _update_tool_permissions(
assert updated.status_code == 200, updated.text
def _listing_on_both(
gateway: Gateway, peer: Gateway, key: str, expected: set[str]
) -> None:
def _listing_on_both(gateway: Gateway, peer: Gateway, key: str, expected: set[str]) -> None:
for worker in (gateway, peer):
listing: Final = eventually(
functools.partial(_granted_view, worker, key),
@ -525,3 +537,56 @@ def test_key_update_tool_permission_widen_narrow_and_clear_apply_on_both_workers
_listing_on_both(gateway, peer, key, all_tools)
nulled: Final = _multiply_outcome_on_both(gateway, peer, key, alias)
assert [call.text for call in nulled] == ["6", "6"], [call.raw for call in nulled]
def _nonce_echo(params: JsonRpc) -> JsonRpc:
return text_result(_STRINGS.validate_python(params["arguments"])["nonce"])
def _call_params(call: Mapping[str, object]) -> Mapping[str, object]:
return _OBJECTS.validate_python(_OBJECTS.validate_python(call["body"])["params"])
def _listed(caller: McpCaller, name: str) -> None:
listing: Final = eventually(caller.list_tools, lambda outcome: name in outcome.tools, seconds=45)
assert listing.error is None, (caller.gateway.client.base_url, listing.raw)
def test_tool_calls_on_both_workers_stay_base_compatible_after_each_worker_lists(
gateway: Gateway, peer: Gateway
) -> None:
schema: Final = {"type": "object", "properties": {"nonce": {"type": "string"}}}
tool: Final = ScriptedTool("echo", _nonce_echo, description="Echo the nonce back", input_schema=schema)
with scripted_peer(tool) as upstream, gateway.scenario() as scenario:
alias: Final = "compat" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, upstream, alias)
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
name: Final = f"{alias}-echo"
first: Final = McpCaller(gateway, key, "mcp", alias)
second: Final = McpCaller(peer, key, "mcp", alias)
_listed(first, name)
_listed(second, name)
upstream.drain()
nonces: Final = (uuid.uuid4().hex, uuid.uuid4().hex)
outcomes: Final = (
first.call(name, {"nonce": nonces[0]}),
second.call(name, {"nonce": nonces[1]}),
first.call(name, {"nonce": nonces[0]}),
)
assert [outcome.text for outcome in outcomes] == [nonces[0], nonces[1], nonces[0]], [o.raw for o in outcomes]
params: Final = [_call_params(call) for call in tool_calls(upstream.drain())]
assert [set(entry) - {"_meta"} for entry in params] == [{"name", "arguments"}] * 3, params
assert [entry["name"] for entry in params] == ["echo"] * 3, params
assert [entry["arguments"] for entry in params] == [
{"nonce": nonces[0]},
{"nonce": nonces[1]},
{"nonce": nonces[0]},
], params
rows: Final = eventually(
lambda: read_rows(_SPEND_NONCES, (sha256(key.encode()).hexdigest(), "call_mcp_tool")),
lambda found: len(found) >= 3,
seconds=70,
)
assert sorted((str(row["status"]), str(row["nonce"])) for row in rows) == sorted(
("success", nonce) for nonce in (nonces[0], nonces[1], nonces[0])
), rows

View file

@ -0,0 +1,531 @@
"""pre_mcp_call guardrails are handed the tool entry ``tools/list`` served to the caller.
One owned proxy carries a default-on ``custom_code`` pre_mcp_call guardrail. At listing time it masks
``SECRET`` out of every scanned text. At call time, when an argument carries the probe marker, it
blocks and echoes the description and parameters it was handed, which is the only way to observe from outside
what metadata the gateway attached to the hook
"""
import json
import threading
import uuid
from collections.abc import Callable, Generator, Iterator, Mapping
from concurrent.futures import ThreadPoolExecutor
from contextlib import ExitStack, contextmanager
from pathlib import Path
from typing import Final
import httpx
import pytest
import yaml
from integration._support.client import Gateway, eventually, gateway_from_environment
from integration._support.mcp import (
EntryPoint,
JsonRpc,
McpCaller,
ScriptedTool,
listed_tools,
openapi_peer,
register_mcp,
scripted_peer,
text_result,
tool_calls,
)
from integration._support.process import owned_proxy
from pydantic import TypeAdapter
_ECHO: Final = "catalog-echo:"
_PROBE: Final = "catalog-probe"
_CALLERS_PER_SERVER: Final = 256
_PIN_SECONDS: Final = 120
_COLD: Final[tuple[str, JsonRpc]] = ("", {"type": "object", "properties": {}, "additionalProperties": False})
_PID: Final = TypeAdapter(int)
_HEADERS: Final = TypeAdapter(dict[bytes, bytes])
_SCHEMA: Final = TypeAdapter(dict[str, object])
_LOOKUP_SCHEMA: Final = {
"type": "object",
"properties": {"probe": {"type": "string", "description": "a probe marker"}},
"additionalProperties": False,
}
_GUARDRAIL_CODE: Final = (
"def apply_guardrail(inputs, request_data, input_type):\n"
' texts = list(inputs.get("texts") or [])\n'
' function = inputs.get("tools", [{}])[0].get("function", {})\n'
f' if "{_PROBE}" in texts:\n'
f' return block("{_ECHO}" + json_stringify('
'{"description": function.get("description"), "parameters": function.get("parameters")}))\n'
' masked = [text.replace("SECRET", "[MASKED]") for text in texts]\n'
" if masked != texts:\n"
" return modify(texts=masked)\n"
" return allow()\n"
)
@pytest.fixture(scope="module")
def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]:
directory: Final = tmp_path_factory.mktemp("listed-tool-metadata")
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["guardrails"] = [
{
"guardrail_name": "catalog-echo-" + uuid.uuid4().hex[:8],
"litellm_params": {
"guardrail": "custom_code",
"mode": "pre_mcp_call",
"default_on": True,
"custom_code": _GUARDRAIL_CODE,
},
}
]
config["general_settings"] = {**config["general_settings"], "proxy_config_reload_interval_seconds": 1}
path: Final = directory / "config.yaml"
path.write_text(yaml.safe_dump(config))
with (
gateway_from_environment() as gateway,
owned_proxy(gateway, directory, {"KEEPALIVE_TIMEOUT": "120"}, config=path, workers=2) as candidate,
ExitStack() as stack,
):
_two_workers(stack, candidate)
yield candidate
def _strings(value: object) -> Iterator[str]:
if isinstance(value, str):
yield value
return
children: Final = value.values() if isinstance(value, Mapping) else value if isinstance(value, list) else ()
for child in children:
yield from _strings(child)
def _decoded(raw: str) -> object:
data: Final = tuple(line[5:].strip() for line in raw.splitlines() if line.startswith("data:"))
return json.loads(data[-1] if data else raw)
def _echoed(raw: str) -> tuple[str | None, Mapping[str, object] | None]:
"""The (description, parameters) the guardrail was handed, recovered from its block reason."""
carrier: Final = next((text for text in _strings(_decoded(raw)) if _ECHO in text), None)
assert carrier is not None, raw
echoed, _ = json.JSONDecoder().raw_decode(carrier.split(_ECHO, 1)[1])
assert isinstance(echoed, dict), carrier
return echoed.get("description"), echoed.get("parameters")
def _probe(caller: McpCaller, name: str, server_id: str) -> tuple[str | None, Mapping[str, object] | None]:
outcome: Final = caller.call(name, {"probe": _PROBE}, server_id=server_id)
assert outcome.error is not None, outcome.raw
return _echoed(outcome.raw)
def _worker(gateway: Gateway) -> int:
response: Final = gateway.request("GET", "/debug/memory/summary")
assert response.status_code == 200, response.text
return _PID.validate_python(response.json()["worker_pid"])
@contextmanager
def _pinned(rig: Gateway) -> Generator[Gateway, None, None]:
"""A single keep-alive connection, so every request on it is served by the worker that accepted it."""
limits: Final = httpx.Limits(max_connections=1, max_keepalive_connections=1)
with httpx.Client(base_url=rig.client.base_url, timeout=15, trust_env=False, limits=limits) as client:
yield Gateway(client, rig.key, rig.upstream_url)
def _connection(stack: ExitStack, rig: Gateway, wanted: Callable[[int], bool]) -> tuple[Gateway, int]:
"""A pinned connection to a worker ``wanted`` accepts; a worker still starting up accepts nothing yet."""
def attempt() -> tuple[Gateway, int] | None:
with ExitStack() as candidate:
gateway: Final = candidate.enter_context(_pinned(rig))
pid: Final = _worker(gateway)
if not wanted(pid):
return None
stack.enter_context(candidate.pop_all())
return gateway, pid
found: Final = eventually(attempt, lambda pair: pair is not None, seconds=_PIN_SECONDS)
assert found is not None
return found
def _two_workers(stack: ExitStack, rig: Gateway) -> tuple[tuple[Gateway, int], tuple[Gateway, int]]:
"""Connects while the first worker is busy answering, so the idle worker wins the accept race."""
first: Final = _connection(stack, rig, lambda _: True)
stop: Final = threading.Event()
def keep_busy() -> None:
while not stop.is_set():
_worker(first[0])
with ThreadPoolExecutor(max_workers=1) as pool:
busy: Final = pool.submit(keep_busy)
try:
other: Final = _connection(stack, rig, lambda pid: pid != first[1])
finally:
stop.set()
busy.result()
return first, other
def _served_name(gateway: Gateway, key: str, identity: str, tool: str) -> str:
"""The prefixed name this worker lists for ``tool`` once its registry reload carries the server."""
listing: Final = eventually(
lambda: McpCaller(gateway, key, "rest").list_tools(server_id=identity),
lambda value: value.ok and any(full.endswith(tool) for full in value.tools),
)
return next(full for full in listing.tools if full.endswith(tool))
def _settled_probe(
gateway: Gateway, key: str, name: str, identity: str
) -> tuple[str | None, Mapping[str, object] | None]:
"""The hook echo for a direct call, once this worker's registry reload carries the server."""
outcome: Final = eventually(
lambda: McpCaller(gateway, key, "rest").call(name, {"probe": _PROBE}, server_id=identity),
lambda value: value.error is not None and _ECHO in value.raw,
)
return _echoed(outcome.raw)
def _forwarded_tenants(observed: tuple[dict[str, object], ...]) -> frozenset[bytes]:
return frozenset(_HEADERS.validate_python(item["headers"]).get(b"x-tenant", b"") for item in observed) - {b""}
def _lookup_tool(description: str | Callable[[Mapping[str, str]], str] = "Look up one record") -> ScriptedTool:
return ScriptedTool("lookup", lambda _: text_result("found"), description=description, input_schema=_LOOKUP_SCHEMA)
@pytest.mark.parametrize("entry", ["rest", "mcp"])
def test_pre_call_hook_receives_the_description_and_input_schema_the_caller_was_listed(
rig: Gateway, entry: EntryPoint
) -> None:
schema: Final = {"type": "object", "properties": {"probe": {"type": "string", "description": "a probe marker"}}}
tool: Final = ScriptedTool(
"lookup", lambda _: text_result("found"), description="Look up one record", input_schema=schema
)
with scripted_peer(tool) as peer, rig.scenario() as scenario:
alias: Final = "meta" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias)
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
caller: Final = McpCaller(rig, key, entry, headers={"x-mcp-servers": alias})
assert caller.initialize().ok
listed: Final = caller.list_tools(server_id=identity)
assert listed.ok, listed.raw
name: Final = next(full for full in listed.tools if full.endswith("lookup"))
description, parameters = _probe(caller, name, identity)
assert description == "Look up one record", (description, parameters)
assert parameters is not None and parameters.get("properties") == schema["properties"], parameters
def test_each_caller_is_evaluated_against_the_catalog_its_own_forwarded_headers_produced(rig: Gateway) -> None:
tool: Final = ScriptedTool(
"report",
lambda _: text_result("ok"),
description=lambda headers: f"Report for tenant {headers.get('x-tenant', 'nobody')}",
)
with scripted_peer(tool) as peer, rig.scenario() as scenario:
alias: Final = "tenant" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias, extra_headers=["x-tenant"])
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
acme: Final = McpCaller(rig, key, "server_mcp", alias, headers={"x-tenant": "acme"})
globex: Final = McpCaller(rig, key, "server_mcp", alias, headers={"x-tenant": "globex"})
acme_listing: Final = acme.list_tools()
globex_listing: Final = globex.list_tools()
assert acme_listing.ok and globex_listing.ok, (acme_listing.raw, globex_listing.raw)
name: Final = next(full for full in acme_listing.tools if full.endswith("report"))
acme_seen, _ = _probe(acme, name, identity)
globex_seen, _ = _probe(globex, name, identity)
assert (acme_seen, globex_seen) == ("Report for tenant acme", "Report for tenant globex"), (
"each caller's tools/call must be evaluated against the catalog its own headers listed"
)
def test_call_is_evaluated_against_the_masked_description_the_listing_served(rig: Gateway) -> None:
tool: Final = ScriptedTool("read_note", lambda _: text_result("note"), description="Read a note")
with scripted_peer(tool) as peer, rig.scenario() as scenario:
alias: Final = "note" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(
scenario, peer, alias, tool_name_to_description={"read_note": "Read a SECRET note"}
)
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
served: Final = listed_tools(rig, key, identity)
name: Final = next(full for full in served if full.endswith("read_note"))
assert served[name]["description"] == "Read a [MASKED] note", served[name]
seen, _ = _probe(McpCaller(rig, key, "rest"), name, identity)
assert seen == "Read a [MASKED] note", "the admin override must not restore wording the listing masked"
def test_openapi_call_is_evaluated_against_the_masked_override_the_listing_served(rig: Gateway) -> None:
with openapi_peer() as peer, rig.scenario() as scenario:
alias: Final = "pets" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(
scenario, peer, alias, tool_name_to_description={"getpet": "Fetch one SECRET pet"}
)
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
served: Final = listed_tools(rig, key, identity)
name: Final = next(full for full in served if full.endswith("getpet"))
assert served[name]["description"] == "Fetch one [MASKED] pet", served[name]
seen, parameters = _probe(McpCaller(rig, key, "rest"), name, identity)
assert seen == "Fetch one [MASKED] pet", "the OpenAPI call path must hand hooks the entry the listing served"
assert parameters is not None and "petId" in parameters.get("properties", {}), parameters
assert not [call for call in peer.drain() if call["path"].startswith("/pets")], "blocked before upstream"
def test_openapi_call_is_evaluated_against_the_entry_this_key_was_listed_not_the_last_listing(
rig: Gateway,
) -> None:
with openapi_peer() as peer, rig.scenario() as scenario:
alias: Final = "pets" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(
scenario, peer, alias, tool_name_to_description={"getpet": "Fetch one SECRET pet"}
)
guarded: Final = scenario.key(object_permission={"mcp_servers": [identity]})
opted_out: Final = scenario.key(
object_permission={"mcp_servers": [identity]}, metadata={"disable_global_guardrails": True}
)
guarded_served: Final = listed_tools(rig, guarded, identity)
opted_out_served: Final = listed_tools(rig, opted_out, identity)
name: Final = next(full for full in guarded_served if full.endswith("getpet"))
assert (guarded_served[name]["description"], opted_out_served[name]["description"]) == (
"Fetch one [MASKED] pet",
"Fetch one SECRET pet",
), (guarded_served[name], opted_out_served[name])
seen, _ = _probe(McpCaller(rig, guarded, "rest"), name, identity)
assert seen == "Fetch one [MASKED] pet", (
"the guarded key must be evaluated against its own listing, not the opted-out key's later one"
)
def test_direct_call_without_a_listing_hands_the_hook_no_metadata_on_either_worker(rig: Gateway) -> None:
with scripted_peer(_lookup_tool()) as peer, ExitStack() as stack:
(first, first_pid), (other, other_pid) = _two_workers(stack, rig)
with first.scenario() as scenario:
identity: Final = register_mcp(scenario, peer, "cold" + uuid.uuid4().hex[:8])
observer: Final = scenario.key(object_permission={"mcp_servers": [identity]})
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
name: Final = _served_name(first, observer, identity, "lookup")
seen: Final = tuple(_settled_probe(gateway, key, name, identity) for gateway in (first, other))
assert seen == (_COLD, _COLD), (seen, first_pid, other_pid)
assert (_worker(first), _worker(other)) == (first_pid, other_pid)
assert tool_calls(peer.drain()) == (), "the probe is blocked at the hook, before the upstream"
def test_warm_metadata_is_local_to_the_worker_that_served_the_listing(rig: Gateway) -> None:
with scripted_peer(_lookup_tool()) as peer, ExitStack() as stack:
(first, first_pid), (other, other_pid) = _two_workers(stack, rig)
with first.scenario() as scenario:
identity: Final = register_mcp(scenario, peer, "local" + uuid.uuid4().hex[:8])
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
name: Final = _served_name(first, key, identity, "lookup")
warm: Final = ("Look up one record", _LOOKUP_SCHEMA)
assert _probe(McpCaller(first, key, "rest"), name, identity) == warm, first_pid
assert _settled_probe(other, key, name, identity) == _COLD, (
"a worker that never served this caller a listing has no catalog for it",
other_pid,
)
assert _served_name(other, key, identity, "lookup") == name
assert _probe(McpCaller(other, key, "rest"), name, identity) == warm, other_pid
assert (_worker(first), _worker(other)) == (first_pid, other_pid)
assert tool_calls(peer.drain()) == ()
def test_listed_tools_without_a_description_still_hand_the_hook_the_schema_the_listing_served(rig: Gateway) -> None:
open_schema: Final[JsonRpc] = {"type": "object", "properties": {}, "additionalProperties": True}
undescribed: Final = ScriptedTool("undescribed", lambda _: text_result("ok"), input_schema=open_schema)
blank: Final = ScriptedTool("blank", lambda _: text_result("ok"), description="", input_schema=open_schema)
with scripted_peer(undescribed, blank) as peer, _pinned(rig) as worker, worker.scenario() as scenario:
pid: Final = _worker(worker)
identity: Final = register_mcp(scenario, peer, "bare" + uuid.uuid4().hex[:8])
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
served: Final = listed_tools(worker, key, identity)
names: Final = tuple(next(full for full in served if full.endswith(tool)) for tool in ("undescribed", "blank"))
assert tuple((served[name].get("description") or "", served[name]["inputSchema"]) for name in names) == (
("", open_schema),
("", open_schema),
), served
seen: Final = tuple(_probe(McpCaller(worker, key, "rest"), name, identity) for name in names)
assert seen == (("", open_schema), ("", open_schema)), (seen, _COLD)
assert _worker(worker) == pid
assert tool_calls(peer.drain()) == ()
def test_hook_receives_the_nested_schema_with_the_leaves_the_listing_masked(rig: Gateway) -> None:
schema: Final = {
"type": "object",
"required": ["filter"],
"additionalProperties": False,
"properties": {
"filter": {
"type": "object",
"properties": {
"path": {"type": "string", "description": "SECRET path"},
"tags": {"type": "array", "items": {"type": "string", "description": "one SECRET tag"}},
},
}
},
}
masked: Final = _SCHEMA.validate_python(json.loads(json.dumps(schema).replace("SECRET", "[MASKED]")))
tool: Final = ScriptedTool(
"search", lambda _: text_result("hit"), description="Search SECRET records", input_schema=schema
)
with scripted_peer(tool) as peer, _pinned(rig) as worker, worker.scenario() as scenario:
pid: Final = _worker(worker)
identity: Final = register_mcp(scenario, peer, "nested" + uuid.uuid4().hex[:8])
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
served: Final = listed_tools(worker, key, identity)
name: Final = next(full for full in served if full.endswith("search"))
assert (served[name]["description"], served[name]["inputSchema"]) == ("Search [MASKED] records", masked), served
assert _probe(McpCaller(worker, key, "rest"), name, identity) == ("Search [MASKED] records", masked), (
"the hook must be handed the nested schema exactly as the listing served it"
)
assert _worker(worker) == pid
assert tool_calls(peer.drain()) == ()
def test_a_server_definition_update_drops_the_catalog_on_every_worker_until_the_caller_lists_again(
rig: Gateway,
) -> None:
with scripted_peer(_lookup_tool()) as peer, ExitStack() as stack:
(first, first_pid), (other, other_pid) = _two_workers(stack, rig)
with first.scenario() as scenario:
identity: Final = register_mcp(scenario, peer, "upd" + uuid.uuid4().hex[:8])
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
name: Final = _served_name(first, key, identity, "lookup")
assert _served_name(other, key, identity, "lookup") == name
before: Final = tuple(_probe(McpCaller(gateway, key, "rest"), name, identity) for gateway in (first, other))
assert before == (("Look up one record", _LOOKUP_SCHEMA),) * 2, before
updated: Final = first.request(
"PUT",
"/v1/mcp/server",
{"server_id": identity, "tool_name_to_description": {"lookup": "Audited lookup"}},
)
assert updated.status_code == 202, updated.text
assert _probe(McpCaller(first, key, "rest"), name, identity) == _COLD, (
"the worker that applied the update must drop its catalog at once",
first_pid,
)
assert (
eventually(lambda: _probe(McpCaller(other, key, "rest"), name, identity), lambda seen: seen == _COLD)
== _COLD
), other_pid
relisted: Final = tuple(
listed_tools(gateway, key, identity)[name]["description"] for gateway in (first, other)
)
assert relisted == ("Audited lookup", "Audited lookup"), relisted
after: Final = tuple(_probe(McpCaller(gateway, key, "rest"), name, identity) for gateway in (first, other))
assert after == (("Audited lookup", _LOOKUP_SCHEMA),) * 2, after
assert (_worker(first), _worker(other)) == (first_pid, other_pid)
assert tool_calls(peer.drain()) == ()
def test_a_listing_in_flight_across_a_server_update_does_not_resurrect_the_old_catalog(rig: Gateway) -> None:
started: Final = threading.Event()
release: Final = threading.Event()
def describe(_: Mapping[str, str]) -> str:
started.set()
assert release.wait(20), "the listing was never released"
return "Look up one record"
with (
scripted_peer(_lookup_tool(describe)) as peer,
ExitStack() as stack,
ThreadPoolExecutor(max_workers=1) as pool,
):
worker, pid = _connection(stack, rig, lambda _: True)
sibling, _ = _connection(stack, rig, lambda candidate: candidate == pid)
with worker.scenario() as scenario:
identity: Final = register_mcp(scenario, peer, "stale" + uuid.uuid4().hex[:8])
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
release.set()
name: Final = _served_name(worker, key, identity, "lookup")
assert _probe(McpCaller(worker, key, "rest"), name, identity) == ("Look up one record", _LOOKUP_SCHEMA)
started.clear()
release.clear()
pending: Final = pool.submit(McpCaller(worker, key, "rest").list_tools, identity)
assert started.wait(10), "the upstream never saw the in-flight listing"
updated: Final = sibling.request(
"PUT",
"/v1/mcp/server",
{"server_id": identity, "tool_name_to_description": {"lookup": "Audited lookup"}},
)
assert updated.status_code == 202, updated.text
release.set()
stale: Final = pending.result(timeout=20)
assert stale.ok, stale.raw
assert _probe(McpCaller(worker, key, "rest"), name, identity) == _COLD, (
"a listing fetched before the update must not be recorded after it"
)
assert listed_tools(worker, key, identity)[name]["description"] == "Audited lookup"
assert _probe(McpCaller(worker, key, "rest"), name, identity) == ("Audited lookup", _LOOKUP_SCHEMA)
assert (_worker(worker), _worker(sibling)) == (pid, pid)
assert tool_calls(peer.drain()) == ()
def test_a_server_keeps_the_newest_256_caller_catalogs_and_evicts_the_oldest(rig: Gateway) -> None:
tool: Final = _lookup_tool(lambda headers: f"Lookup for tenant {headers.get('x-tenant', 'nobody')}")
with scripted_peer(tool) as peer, _pinned(rig) as worker, worker.scenario() as scenario:
pid: Final = _worker(worker)
alias: Final = "cap" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias, extra_headers=["x-tenant"])
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
tenants: Final = tuple(f"t{index}" for index in range(_CALLERS_PER_SERVER + 1))
callers: Final = {
tenant: McpCaller(worker, key, "server_mcp", alias, headers={"x-tenant": tenant}) for tenant in tenants
}
listings: Final = tuple(callers[tenant].list_tools() for tenant in tenants)
assert all(listing.ok for listing in listings), [listing.raw for listing in listings if not listing.ok]
name: Final = next(full for full in listings[0].tools if full.endswith("lookup"))
def seen(tenant: str) -> str | None:
return _probe(callers[tenant], name, identity)[0]
assert (seen("t0"), seen("t1"), seen(tenants[-1])) == (
"",
"Lookup for tenant t1",
f"Lookup for tenant {tenants[-1]}",
), "the oldest of 257 callers is evicted, the newest 256 keep their own catalog"
assert callers["t0"].list_tools().ok
assert (seen("t0"), seen("t1")) == ("Lookup for tenant t0", ""), "relisting makes t0 newest and evicts t1"
assert _worker(worker) == pid
observed: Final = peer.drain()
assert _forwarded_tenants(observed) == frozenset(tenant.encode() for tenant in tenants), len(observed)
assert tool_calls(observed) == ()
def test_an_admin_include_disabled_tools_listing_does_not_warm_the_runtime_catalog(rig: Gateway) -> None:
"""``include_disabled_tools=true`` is the admin-only configuration view, not a listing the caller
runs against: recording it would warm tools/call metadata no runtime listing ever served."""
tool: Final = _lookup_tool()
with scripted_peer(tool) as peer, rig.scenario() as scenario:
alias: Final = "adminview" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias)
admin: Final = rig.key
caller: Final = McpCaller(rig, admin, "rest")
name: Final = eventually(
lambda: caller.call(f"{alias}-lookup", {"probe": _PROBE}, server_id=identity),
lambda value: value.error is not None and _ECHO in value.raw,
)
assert _echoed(name.raw) == _COLD, "before any listing the call is cold"
view: Final = eventually(
lambda: rig.client.get(
"/mcp-rest/tools/list",
params={"server_id": identity, "include_disabled_tools": "true"},
headers={"x-litellm-api-key": admin},
),
lambda response: response.status_code == 200
and any(entry["name"].endswith("lookup") for entry in response.json()["tools"]),
)
assert view.status_code == 200, view.text
after_view: Final = _probe(caller, f"{alias}-lookup", identity)
assert after_view == _COLD, (
"the admin-only include_disabled_tools view must not record the caller's listed-tools slot"
)
runtime: Final = caller.list_tools(server_id=identity)
assert runtime.ok, runtime.raw
assert _probe(caller, f"{alias}-lookup", identity)[0] == "Look up one record", (
"a genuine runtime listing still warms the slot"
)

View file

@ -1,16 +1,35 @@
import asyncio
import json
import uuid
from collections.abc import Callable, Iterator, Mapping, Sequence
from collections.abc import Callable, Generator, Iterator, Mapping, Sequence
from concurrent.futures import ThreadPoolExecutor
from contextlib import contextmanager
from dataclasses import dataclass
from typing import Final, Literal
from hashlib import sha256
from pathlib import Path
from typing import Final, Literal, TypeVar
import httpx
import pytest
from integration._support.client import Gateway, Scenario
from integration._support.mcp import McpPeer, mcp_peer, register_mcp, tool_calls
import yaml
from integration._support.client import Gateway, Scenario, eventually, gateway_from_environment, object_value
from integration._support.database import read_rows
from integration._support.mcp import (
McpPeer,
ScriptedTool,
mcp_peer,
register_mcp,
scripted_peer,
text_result,
tool_calls,
)
from integration._support.mcp_grants import create_toolset
from integration._support.process import owned_proxy
from integration._support.wire import Reply, Request, Wire, wire_server
from openai import AsyncOpenAI, OpenAI
from openai.types.chat import ChatCompletionMessageParam
from openai.types.responses.tool_param import Mcp
from pydantic import JsonValue, TypeAdapter
Surface = Literal["chat", "responses", "messages", "messages_bridge"]
SURFACES: Final[tuple[Surface, ...]] = ("chat", "responses", "messages", "messages_bridge")
@ -18,6 +37,9 @@ ADD: Final = {"a": 2, "b": 3}
ANSWER: Final = "the sum is 5"
GATEWAY_REF: Final = {"type": "mcp", "server_url": "litellm_proxy", "server_label": "litellm"}
AUTO: Final = {**GATEWAY_REF, "require_approval": "never"}
OUTAGE: Final = "bridge-outage"
JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
JSON_VALUE: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue)
def _json(body: Mapping[str, object]) -> Reply:
@ -44,25 +66,94 @@ def _has_tool_result(body: Mapping[str, object]) -> bool:
return False
def _model_double(tool: str) -> Callable[[Request], Reply]:
arguments: Final = json.dumps(ADD)
@dataclass(frozen=True, slots=True)
class Turn:
tool: str
arguments: str
answer: str
def _fixed_turn(tool: str) -> Callable[[Mapping[str, JsonValue]], Turn]:
return lambda _: Turn(tool, json.dumps(ADD), ANSWER)
def _echoing_turn(body: Mapping[str, JsonValue]) -> Turn:
names: Final = _tool_names(body)
return Turn(names[0] if names else "", json.dumps({"query": _prompt(body)}), _tool_result_text(body) or "")
def _prompt(body: Mapping[str, JsonValue]) -> str:
inputs: Final = body.get("input")
if isinstance(inputs, str):
return inputs
items: Final = inputs if isinstance(inputs, list) else body.get("messages")
first: Final = items[0] if isinstance(items, list) and items else None
return str(first["content"]) if isinstance(first, dict) and isinstance(first.get("content"), str) else ""
def _tool_result_text(body: Mapping[str, JsonValue]) -> str | None:
inputs: Final = body.get("input")
messages: Final = body.get("messages")
items: Final = inputs if isinstance(inputs, list) else messages if isinstance(messages, list) else ()
for item in items:
if not isinstance(item, dict):
continue
if item.get("type") == "function_call_output":
return str(item["output"])
if item.get("role") == "tool":
return str(item["content"])
if (block := _tool_result_block(item)) is not None:
return block
return None
def _tool_result_block(item: Mapping[str, JsonValue]) -> str | None:
content: Final = item.get("content")
for block in content if isinstance(content, list) else ():
if isinstance(block, dict) and block.get("type") == "tool_result":
return str(block["content"])
return None
def _responses_stream(response: Mapping[str, JsonValue], item: Mapping[str, JsonValue]) -> Reply:
events: Final = (
{"type": "response.created", "sequence_number": 0, "response": {**response, "status": "in_progress"}},
{"type": "response.in_progress", "sequence_number": 1, "response": {**response, "status": "in_progress"}},
{"type": "response.output_item.added", "sequence_number": 2, "output_index": 0, "item": item},
{"type": "response.output_item.done", "sequence_number": 3, "output_index": 0, "item": item},
{"type": "response.completed", "sequence_number": 4, "response": response},
)
return Reply(
content_type="text/event-stream",
chunks=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events),
)
def _model_double(plan: Callable[[Mapping[str, JsonValue]], Turn]) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
if request.method == "GET" and request.target.endswith("/models"):
return _json({"object": "list", "data": []})
body: Final = json.loads(request.body)
assert isinstance(body, dict), request.body
body: Final = JSON_OBJECT.validate_json(request.body)
if OUTAGE in _prompt(body):
outage: Final = {"error": {"message": "scripted provider outage", "type": "server_error", "code": None}}
return Reply(status=500, body=json.dumps(outage).encode())
turn: Final = plan(body)
done: Final = _has_tool_result(body)
identity: Final = uuid.uuid4().hex[:12]
usage: Final = {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}
if request.target.endswith("/chat/completions"):
message: Final = (
{"role": "assistant", "content": ANSWER}
{"role": "assistant", "content": turn.answer}
if done
else {
"role": "assistant",
"content": None,
"tool_calls": [
{"id": "call_1", "type": "function", "function": {"name": tool, "arguments": arguments}}
{
"id": "call_1",
"type": "function",
"function": {"name": turn.tool, "arguments": turn.arguments},
}
],
}
)
@ -74,7 +165,7 @@ def _model_double(tool: str) -> Callable[[Request], Reply]:
else message
)
chunk: Final = {
"id": "chatcmpl-1",
"id": f"chatcmpl-{identity}",
"object": "chat.completion.chunk",
"created": 1,
"model": "gpt-4o-mini",
@ -89,7 +180,7 @@ def _model_double(tool: str) -> Callable[[Request], Reply]:
)
return _json(
{
"id": "chatcmpl-1",
"id": f"chatcmpl-{identity}",
"object": "chat.completion",
"created": 1,
"model": "gpt-4o-mini",
@ -99,13 +190,13 @@ def _model_double(tool: str) -> Callable[[Request], Reply]:
)
if request.target.endswith("/messages"):
content: Final = (
[{"type": "text", "text": ANSWER}]
[{"type": "text", "text": turn.answer}]
if done
else [{"type": "tool_use", "id": "toolu_1", "name": tool, "input": ADD}]
else [{"type": "tool_use", "id": "toolu_1", "name": turn.tool, "input": json.loads(turn.arguments)}]
)
return _json(
{
"id": "msg_1",
"id": f"msg_{identity}",
"type": "message",
"role": "assistant",
"model": "claude",
@ -116,39 +207,34 @@ def _model_double(tool: str) -> Callable[[Request], Reply]:
}
)
assert request.target.endswith("/responses"), request.target
output: Final = (
[
{
"type": "message",
"id": "msg_1",
"role": "assistant",
"status": "completed",
"content": [{"type": "output_text", "text": ANSWER, "annotations": []}],
}
]
if done
else [
{
"type": "function_call",
"id": "fc_1",
"call_id": "call_1",
"name": tool,
"arguments": arguments,
"status": "completed",
}
]
)
return _json(
item: Final[dict[str, JsonValue]] = (
{
"id": "resp_1",
"object": "response",
"created_at": 1,
"type": "message",
"id": f"msg_{identity}",
"role": "assistant",
"status": "completed",
"content": [{"type": "output_text", "text": turn.answer, "annotations": []}],
}
if done
else {
"type": "function_call",
"id": f"fc_{identity}",
"call_id": f"call_{identity}",
"name": turn.tool,
"arguments": turn.arguments,
"status": "completed",
"model": "gpt-4o-mini",
"output": output,
"usage": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15},
}
)
response: Final[dict[str, JsonValue]] = {
"id": f"resp_{identity}",
"object": "response",
"created_at": 1,
"status": "completed",
"model": "gpt-4o-mini",
"output": [item],
"usage": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15},
}
return _responses_stream(response, item) if body.get("stream") is True else _json(response)
return respond
@ -240,7 +326,7 @@ def _rig(gateway: Gateway, surface: Surface) -> Iterator[Rig]:
alias: Final = "llm" + uuid.uuid4().hex[:8]
with (
mcp_peer() as peer,
wire_server(_model_double(f"{alias}-add")) as wire,
wire_server(_model_double(_fixed_turn(f"{alias}-add"))) as wire,
gateway.scenario() as scenario,
):
server_id: Final = register_mcp(scenario, peer, alias)
@ -403,3 +489,460 @@ def test_toolset_gateway_url_gives_a_key_of_an_ungranted_team_no_tools_and_never
assert _peer_add_calls(rig.peer) == (), "denied caller reached the peer"
assert all(rig.tool not in names for names in rig.upstream_tools()), rig.upstream_tools()
assert response.status_code in (200, 400, 401, 403), response.text
Bridge = Literal["chat", "responses", "messages"]
BRIDGES: Final[tuple[Bridge, ...]] = ("chat", "responses")
Client = Literal["sync", "async"]
HOOK_ECHO: Final = "bridge-echo:"
HOOK_PROBE: Final = "bridge-probe"
HOOK_CODE: Final = (
"def apply_guardrail(inputs, request_data, input_type):\n"
' texts = list(inputs.get("texts") or [])\n'
' function = inputs.get("tools", [{}])[0].get("function", {})\n'
" for text in texts:\n"
f' if "{HOOK_PROBE}" in text:\n'
f' return block("{HOOK_ECHO}" + json_stringify('
'{"description": function.get("description"), "parameters": function.get("parameters")}))\n'
" return allow()\n"
)
LOOKUP: Final[tuple[str, dict[str, JsonValue]]] = (
"Look up one record",
{"type": "object", "properties": {"query": {"type": "string"}}},
)
REPORT: Final[tuple[str, dict[str, JsonValue]]] = (
"Write one report",
{"type": "object", "properties": {"query": {"type": "string"}, "format": {}}},
)
COLD: Final[tuple[str, dict[str, JsonValue]]] = (
"",
{"type": "object", "properties": {}, "additionalProperties": False},
)
RELOAD_FAST: Final = {"PROXY_CONFIG_RELOAD_INTERVAL_SECONDS": "5"}
CALL_ID: Final = "x-litellm-call-id"
T = TypeVar("T")
Definition = tuple[str, str, JsonValue]
Echo = tuple[JsonValue, JsonValue]
def _served(listing: tuple[str, Mapping[str, JsonValue]]) -> Echo:
return listing[0], {**listing[1], "additionalProperties": False}
@dataclass(frozen=True, slots=True)
class Hooked:
proxy: Gateway
sink: Wire
@pytest.fixture(scope="module")
def hooked(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Hooked]:
directory: Final = tmp_path_factory.mktemp("bridge-hooks")
base: Final = JSON_OBJECT.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()))
with wire_server(lambda _: _json({"flagged": False, "session_id": "scripted"})) as sink:
echo: Final = {"guardrail": "custom_code", "mode": "pre_mcp_call", "default_on": True, "custom_code": HOOK_CODE}
pillar: Final = {
"guardrail": "pillar",
"mode": ["pre_call", "pre_mcp_call"],
"default_on": True,
"api_key": "sk-pillar-" + uuid.uuid4().hex,
"api_base": sink.url,
"on_flagged_action": "monitor",
}
guardrails: Final = [
{"guardrail_name": "bridge-echo-" + uuid.uuid4().hex[:8], "litellm_params": echo},
{"guardrail_name": "bridge-sink-" + uuid.uuid4().hex[:8], "litellm_params": pillar},
]
path: Final = directory / "config.yaml"
path.write_text(yaml.safe_dump({**base, "guardrails": guardrails}))
with (
gateway_from_environment() as gateway,
owned_proxy(gateway, directory, RELOAD_FAST, config=path, workers=2) as proxy,
):
yield Hooked(proxy, sink)
@dataclass(frozen=True, slots=True)
class BridgeRig:
hooked: Hooked
scenario: Scenario
peer: McpPeer
wire: Wire
alias: str
server_id: str
model: str
bridge: Bridge
def tool(self, name: str) -> str:
return f"{self.alias}-{name}"
def names(self) -> frozenset[str]:
return frozenset(("lookup", "report"))
def mcp(self, name: str) -> Mcp:
return {**AUTO_MCP, "allowed_tools": [self.tool(name)]}
def url(self, path: str) -> str:
return str(self.hooked.proxy.client.base_url).rstrip("/") + path
def post(self, key: str, prompt: str, tools: Sequence[Mcp], **extra: object) -> httpx.Response:
headers: Final = {"Authorization": f"Bearer {key}"}
if self.bridge == "chat":
body: Final = {"model": self.model, "messages": [{"role": "user", "content": prompt}], "tools": list(tools)}
return httpx.post(self.url("/v1/chat/completions"), headers=headers, json={**body, **extra}, timeout=90)
if self.bridge == "responses":
body_r: Final = {"model": self.model, "input": prompt, "tools": list(tools), **extra}
return httpx.post(self.url("/v1/responses"), headers=headers, json=body_r, timeout=90)
body_m: Final = {
"model": self.model,
"max_tokens": 64,
"messages": [{"role": "user", "content": prompt}],
"tools": list(tools),
**extra,
}
return httpx.post(self.url("/v1/messages"), headers=headers, json=body_m, timeout=90)
def upstream_by_prompt(self) -> Mapping[str, tuple[tuple[Definition, ...], ...]]:
bodies: Final = tuple(
JSON_OBJECT.validate_json(request.body) for request in self.wire.drain() if request.method == "POST"
)
prompts: Final = frozenset(_prompt(body) for body in bodies)
return {prompt: tuple(_definitions(body) for body in bodies if _prompt(body) == prompt) for prompt in prompts}
def peer_calls(self) -> tuple[tuple[str, JsonValue], ...]:
return tuple(_peer_call(call) for call in tool_calls(self.peer.drain()))
def hook_messages(self, marker: str) -> tuple[str, ...]:
posted: Final = tuple(
JSON_OBJECT.validate_json(request.body) for request in self.hooked.sink.drain() if request.method == "POST"
)
contents: Final = tuple(content for payload in posted for content in _contents(payload))
return tuple(content for content in contents if _synthetic(content, marker))
AUTO_MCP: Final[Mcp] = {
"type": "mcp",
"server_label": "litellm",
"server_url": "litellm_proxy",
"require_approval": "never",
}
def _peer_call(call: Mapping[str, object]) -> tuple[str, JsonValue]:
params: Final = object_value(object_value(JSON_VALUE.validate_python(call["body"]))["params"])
return str(params["name"]), params["arguments"]
def _contents(payload: Mapping[str, JsonValue]) -> Iterator[str]:
messages: Final = payload.get("messages")
for message in messages if isinstance(messages, list) else ():
if isinstance(message, dict):
yield str(message.get("content"))
def _function(tool: JsonValue) -> dict[str, JsonValue] | None:
if not isinstance(tool, dict):
return None
function: Final = tool.get("function", tool)
return function if isinstance(function, dict) else None
def _definitions(body: Mapping[str, JsonValue]) -> tuple[Definition, ...]:
tools: Final = body.get("tools")
functions: Final = tuple(_function(tool) for tool in tools) if isinstance(tools, list) else ()
return tuple(
(str(function["name"]), str(function.get("description", "")), function.get("parameters"))
for function in functions
if function is not None
)
def _uniform(upstream: Mapping[str, tuple[tuple[Definition, ...], ...]]) -> Mapping[str, tuple[Definition, ...] | None]:
return {
prompt: rounds[0] if all(definitions == rounds[0] for definitions in rounds) else None
for prompt, rounds in upstream.items()
}
def _strings(value: JsonValue) -> Iterator[str]:
if isinstance(value, str):
yield value
return
children: Final = value.values() if isinstance(value, dict) else value if isinstance(value, list) else ()
for child in children:
yield from _strings(child)
def _echoed(value: JsonValue) -> Echo:
carrier: Final = next((text for text in _strings(value) if HOOK_ECHO in text), None)
assert carrier is not None, value
payload: Final = carrier.split(HOOK_ECHO, 1)[1]
end: Final = json.JSONDecoder().raw_decode(payload)[1]
echoed: Final = JSON_OBJECT.validate_json(payload[:end])
return echoed.get("description"), echoed.get("parameters")
@contextmanager
def _bridge_rig(hooked: Hooked, bridge: Bridge) -> Generator[BridgeRig, None, None]:
alias: Final = "brg" + uuid.uuid4().hex[:8]
lookup: Final = ScriptedTool(
"lookup",
lambda params: text_result("found:" + json.dumps(JSON_OBJECT.validate_python(params)["arguments"])),
description=LOOKUP[0],
input_schema=LOOKUP[1],
)
report: Final = ScriptedTool(
"report", lambda _: text_result("reported"), description=REPORT[0], input_schema=REPORT[1]
)
with (
scripted_peer(lookup, report) as peer,
wire_server(_model_double(_echoing_turn)) as wire,
hooked.proxy.scenario() as scenario,
):
server_id: Final = register_mcp(scenario, peer, alias)
model: Final = scenario.model(model=_upstream_model(bridge), api_base=wire.url + "/v1")
rig: Final = BridgeRig(hooked, scenario, peer, wire, alias, server_id, model, bridge)
eventually(
lambda: tuple(_on_worker(hooked.proxy, lambda client: _master_listing(rig, client)) for _ in range(6)),
lambda seen: len({pid for pid, _ in seen}) >= 2 and all(names == rig.names() for _, names in seen),
seconds=45,
)
peer.drain()
hooked.sink.drain()
yield rig
def _bridge_key(rig: BridgeRig) -> str:
return rig.scenario.key(object_permission={"mcp_servers": [rig.server_id]})
@dataclass(frozen=True, slots=True)
class Seen:
response_id: str
call_id: str
text: str
def _ask_sync(rig: BridgeRig, key: str, prompt: str, tools: Sequence[Mcp], stream: bool) -> Seen:
sdk: Final = OpenAI(base_url=rig.url("/v1"), api_key=key, max_retries=0, timeout=90)
if rig.bridge == "responses":
if stream:
raw_events: Final = sdk.responses.with_raw_response.create(
model=rig.model, input=prompt, tools=tools, stream=True
)
completed: Final = next(
event.response for event in raw_events.parse() if event.type == "response.completed"
)
return Seen(completed.id, raw_events.headers[CALL_ID], completed.output_text)
raw_response: Final = sdk.responses.with_raw_response.create(model=rig.model, input=prompt, tools=tools)
response: Final = raw_response.parse()
return Seen(response.id, raw_response.headers[CALL_ID], response.output_text)
messages: Final[list[ChatCompletionMessageParam]] = [{"role": "user", "content": prompt}]
extra: Final = {"tools": list(tools)}
if stream:
raw_chunks: Final = sdk.chat.completions.with_raw_response.create(
model=rig.model, messages=messages, stream=True, extra_body=extra
)
parts: Final = tuple(
(chunk.id, chunk.choices[0].delta.content or "") for chunk in raw_chunks.parse() if chunk.choices
)
return Seen(parts[0][0], raw_chunks.headers[CALL_ID], "".join(text for _, text in parts))
raw_completion: Final = sdk.chat.completions.with_raw_response.create(
model=rig.model, messages=messages, extra_body=extra
)
completion: Final = raw_completion.parse()
return Seen(completion.id, raw_completion.headers[CALL_ID], completion.choices[0].message.content or "")
async def _ask_async(rig: BridgeRig, key: str, prompt: str, tools: Sequence[Mcp], stream: bool) -> Seen:
sdk: Final = AsyncOpenAI(base_url=rig.url("/v1"), api_key=key, max_retries=0, timeout=90)
if rig.bridge == "responses":
if stream:
raw_events: Final = await sdk.responses.with_raw_response.create(
model=rig.model, input=prompt, tools=tools, stream=True
)
completed: Final = [
event.response async for event in raw_events.parse() if event.type == "response.completed"
]
return Seen(completed[0].id, raw_events.headers[CALL_ID], completed[0].output_text)
raw_response: Final = await sdk.responses.with_raw_response.create(model=rig.model, input=prompt, tools=tools)
response: Final = raw_response.parse()
return Seen(response.id, raw_response.headers[CALL_ID], response.output_text)
messages: Final[list[ChatCompletionMessageParam]] = [{"role": "user", "content": prompt}]
extra: Final = {"tools": list(tools)}
if stream:
raw_chunks: Final = await sdk.chat.completions.with_raw_response.create(
model=rig.model, messages=messages, stream=True, extra_body=extra
)
parts: Final = [
(chunk.id, chunk.choices[0].delta.content or "") async for chunk in raw_chunks.parse() if chunk.choices
]
return Seen(parts[0][0], raw_chunks.headers[CALL_ID], "".join(text for _, text in parts))
raw_completion: Final = await sdk.chat.completions.with_raw_response.create(
model=rig.model, messages=messages, extra_body=extra
)
completion: Final = raw_completion.parse()
return Seen(completion.id, raw_completion.headers[CALL_ID], completion.choices[0].message.content or "")
def _ask(rig: BridgeRig, key: str, prompt: str, tools: Sequence[Mcp], stream: bool, client: Client) -> Seen:
if client == "async":
return asyncio.run(_ask_async(rig, key, prompt, tools, stream))
return _ask_sync(rig, key, prompt, tools, stream)
def _synthetic(content: str, marker: str) -> bool:
return content.startswith("Tool: lookup\n") and marker in content
def _on_worker(gateway: Gateway, act: Callable[[httpx.Client], T]) -> tuple[int, T]:
with httpx.Client(base_url=str(gateway.client.base_url), timeout=30) as client:
summary: Final = client.get("/debug/memory/summary", headers={"Authorization": f"Bearer {gateway.key}"})
assert summary.status_code == 200, summary.text
pid: Final = JSON_OBJECT.validate_json(summary.content)["worker_pid"]
assert isinstance(pid, int), summary.text
return pid, act(client)
def _both_workers(gateway: Gateway) -> frozenset[int]:
return eventually(
lambda: frozenset(_on_worker(gateway, lambda _: None)[0] for _ in range(6)), lambda pids: len(pids) >= 2
)
def _master_listing(rig: BridgeRig, client: httpx.Client) -> frozenset[str]:
headers: Final = {"x-litellm-api-key": rig.hooked.proxy.key}
response: Final = client.get("/mcp-rest/tools/list", headers=headers, params={"server_id": rig.server_id})
tools: Final = JSON_OBJECT.validate_json(response.content).get("tools") if response.status_code == 200 else None
return frozenset(str(object_value(tool)["name"]) for tool in tools) if isinstance(tools, list) else frozenset()
def _direct_probe(rig: BridgeRig, key: str, name: str, client: httpx.Client) -> Echo:
body: Final = {"server_id": rig.server_id, "name": rig.tool(name), "arguments": {"query": HOOK_PROBE}}
response: Final = client.post("/mcp-rest/tools/call", headers={"x-litellm-api-key": key}, json=body)
return _echoed(JSON_VALUE.validate_json(response.content))
def _spend_row(key: str, call_id: str) -> dict[str, JsonValue]:
rows: Final = eventually(
lambda: read_rows(
'SELECT api_key, call_type, status, cache_hit FROM "LiteLLM_SpendLogs" WHERE litellm_call_id=%s', (call_id,)
),
lambda found: len(found) >= 1,
seconds=70,
)
assert rows[0]["api_key"] == sha256(key.encode()).hexdigest(), rows
return rows[0]
@pytest.mark.parametrize("client", ("sync", "async"))
@pytest.mark.parametrize("stream", (False, True), ids=("plain", "stream"))
@pytest.mark.parametrize("bridge", BRIDGES)
def test_bridge_hook_sees_the_definition_of_the_tool_filtered_for_that_request(
hooked: Hooked, bridge: Bridge, stream: bool, client: Client
) -> None:
with _bridge_rig(hooked, bridge) as rig:
key: Final = _bridge_key(rig)
marker: Final = "m" + uuid.uuid4().hex
probe: Final = f"{marker} {HOOK_PROBE}"
found: Final = _ask(rig, key, marker, [rig.mcp("lookup")], stream, client)
assert found.text == "found:" + json.dumps({"query": marker}), found
blocked: Final = _ask(rig, key, probe, [rig.mcp("lookup")], stream, client)
assert _echoed(blocked.text) == _served(LOOKUP), blocked
expected: Final = ((rig.tool("lookup"), *_served(LOOKUP)),)
upstream: Final = rig.upstream_by_prompt()
assert set(upstream) == {marker, probe} and all(
definitions == expected for definitions in upstream[marker] + upstream[probe]
), upstream
assert rig.peer_calls() == (("lookup", {"query": marker}),)
assert rig.hook_messages(marker) == (f"Tool: lookup\nArguments: {dict(query=marker)}",)
assert _spend_row(key, found.call_id)["status"] == "success"
@pytest.mark.parametrize("bridge", BRIDGES)
def test_direct_call_after_bridge_only_discovery_stays_cold(hooked: Hooked, bridge: Bridge) -> None:
with _bridge_rig(hooked, bridge) as rig:
key: Final = _bridge_key(rig)
workers: Final = _both_workers(rig.hooked.proxy)
bridged: Final = _ask(rig, key, HOOK_PROBE, [rig.mcp("lookup")], False, "sync")
assert _echoed(bridged.text) == _served(LOOKUP), bridged
direct: Final = eventually(
lambda: tuple(
_on_worker(rig.hooked.proxy, lambda client: _direct_probe(rig, key, "lookup", client)) for _ in range(6)
),
lambda seen: frozenset(pid for pid, _ in seen) == workers,
)
assert all(echo == COLD for _, echo in direct) and frozenset(pid for pid, _ in direct) == workers, (
direct,
workers,
)
assert rig.peer_calls() == ()
@pytest.mark.parametrize("bridge", BRIDGES)
def test_concurrent_requests_of_one_key_with_different_allowed_tools_each_see_their_own_definition(
hooked: Hooked, bridge: Bridge
) -> None:
with _bridge_rig(hooked, bridge) as rig:
key: Final = _bridge_key(rig)
prompts: Final = {name: f"{name} {uuid.uuid4().hex} {HOOK_PROBE}" for name in ("lookup", "report")}
def ask(name: str) -> Seen:
return _ask(rig, key, prompts[name], [rig.mcp(name)], False, "sync")
with ThreadPoolExecutor(2) as pool:
lookup, report = pool.map(ask, ("lookup", "report"))
assert (_echoed(lookup.text), _echoed(report.text)) == (_served(LOOKUP), _served(REPORT)), (lookup, report)
upstream: Final = rig.upstream_by_prompt()
assert _uniform(upstream) == {
prompts["lookup"]: ((rig.tool("lookup"), *_served(LOOKUP)),),
prompts["report"]: ((rig.tool("report"), *_served(REPORT)),),
}, upstream
assert rig.peer_calls() == ()
@pytest.mark.parametrize("bridge", BRIDGES)
def test_provider_outage_reaches_the_caller_and_never_the_peer_or_the_hooks(hooked: Hooked, bridge: Bridge) -> None:
with _bridge_rig(hooked, bridge) as rig:
key: Final = _bridge_key(rig)
prompt: Final = f"{OUTAGE} {uuid.uuid4().hex}"
response: Final = rig.post(key, prompt, [rig.mcp("lookup")])
assert response.status_code == 500, response.text
upstream: Final = rig.upstream_by_prompt()
assert set(upstream) == {prompt} and all(
definitions == ((rig.tool("lookup"), *_served(LOOKUP)),) for definitions in upstream[prompt]
), upstream
assert rig.peer_calls() == ()
assert rig.hook_messages(prompt) == ()
@pytest.mark.parametrize("bridge", BRIDGES)
def test_identical_nonstream_repeat_is_a_cache_hit_without_new_model_peer_or_hook_traffic(
hooked: Hooked, bridge: Bridge
) -> None:
with _bridge_rig(hooked, bridge) as rig:
key: Final = _bridge_key(rig)
marker: Final = "m" + uuid.uuid4().hex
first: Final = _ask(rig, key, marker, [rig.mcp("lookup")], False, "sync")
assert first.text == "found:" + json.dumps({"query": marker}), first
assert set(rig.upstream_by_prompt()) == {marker} and rig.peer_calls() == (("lookup", {"query": marker}),)
assert rig.hook_messages(marker) == (f"Tool: lookup\nArguments: {dict(query=marker)}",)
assert _spend_row(key, first.call_id)["cache_hit"] != "True"
repeat: Final = _ask(rig, key, marker, [rig.mcp("lookup")], False, "sync")
assert repeat.text == first.text, (first, repeat)
assert (rig.upstream_by_prompt(), rig.peer_calls(), rig.hook_messages(marker)) == ({}, (), ()), repeat
assert _spend_row(key, repeat.call_id)["cache_hit"] == "True"
def test_messages_bridge_hook_keeps_the_base_shape_without_request_local_metadata(hooked: Hooked) -> None:
with _bridge_rig(hooked, "messages") as rig:
key: Final = _bridge_key(rig)
marker: Final = "m" + uuid.uuid4().hex
probe: Final = f"{marker} {HOOK_PROBE}"
found: Final = rig.post(key, marker, [rig.mcp("lookup")])
assert found.status_code == 200, found.text
blocked: Final = rig.post(key, probe, [rig.mcp("lookup")])
assert blocked.status_code == 200, blocked.text
assert _echoed(JSON_VALUE.validate_json(blocked.content)) == COLD, blocked.text
assert rig.peer_calls() == (("lookup", {"query": marker}),)
assert rig.hook_messages(marker) == (f"Tool: lookup\nArguments: {dict(query=marker)}",)

View file

@ -1,16 +1,20 @@
import base64
import hashlib
import re
import secrets
import textwrap
import time
import uuid
from collections.abc import Iterator
from dataclasses import dataclass
from pathlib import Path
from typing import Final
from urllib.parse import parse_qs, urlsplit
import httpx
import jwt
import pytest
from integration._support.client import Gateway, eventually
from integration._support.client import Gateway, eventually, gateway_from_environment
from integration._support.database import read_rows
from integration._support.mcp import (
ENTRY_POINTS,
@ -27,6 +31,7 @@ from integration._support.mcp import (
)
from integration._support.mcp_grants import create_toolset
from integration._support.oauth_server import AuthorizationServer, oauth_server
from integration._support.process import owned_proxy
ADD: Final = {"a": 2, "b": 3}
CLIENT_REDIRECT: Final = "http://127.0.0.1:9/cb"
@ -503,3 +508,70 @@ def test_resource_scoped_session_bearer_opens_a_team_toolset_inside_its_server_a
refused: Final = _toolset_rpc(gateway, bearer, outside_name, "tools/list", {})
assert refused.status == 403, refused.raw
assert tool_calls(peer.drain()) == ()
_PROBE: Final = "catalog-probe"
_ECHO: Final = "catalog-echo"
_UNLISTED: Final = ""
_GUARDRAIL_CODE: Final = (
"def apply_guardrail(inputs, request_data, input_type):\n"
f' if "{_PROBE}" not in list(inputs.get("texts") or []):\n'
" return allow()\n"
' function = inputs.get("tools", [{}])[0].get("function", {})\n'
f' return block("{_ECHO}[" + function.get("description") + "]")\n'
)
_ECHO_GUARDRAIL_YAML: Final = (
"guardrails:\n"
" - guardrail_name: catalog-echo\n"
" litellm_params:\n"
" guardrail: custom_code\n"
" mode: pre_mcp_call\n"
" default_on: true\n"
" custom_code: |\n" + textwrap.indent(_GUARDRAIL_CODE, 8 * " ")
)
@pytest.fixture(scope="module")
def echo_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]:
directory: Final = tmp_path_factory.mktemp("catalog-echo")
path: Final = directory / "catalog_echo.yaml"
path.write_text((Path(__file__).resolve().parents[1] / "proxy_config.yaml").read_text() + _ECHO_GUARDRAIL_YAML)
with gateway_from_environment() as gateway, owned_proxy(gateway, directory, {}, config=path, workers=2) as rig:
yield rig
def _echoed_description(outcome: Outcome) -> str:
found: Final = re.search(rf"{_ECHO}\[(.*?)\]", outcome.raw)
assert found is not None, outcome.raw
return found.group(1)
def test_token_exchange_callers_with_different_subject_tokens_own_separate_listings(echo_rig: Gateway) -> None:
with mcp_peer() as peer, oauth_server() as auth, echo_rig.scenario() as scenario:
alias: Final = "te" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(
scenario,
peer,
alias,
auth_type="oauth2_token_exchange",
token_exchange_endpoint=auth.issuer + "/token",
credentials={"client_id": "te-client", "client_secret": "te-secret"},
)
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
first_subject: Final = "subject-" + uuid.uuid4().hex
first: Final = McpCaller(echo_rig, key, "mcp", alias, {"Authorization": f"Bearer {first_subject}"})
second: Final = McpCaller(echo_rig, key, "mcp", alias, {"Authorization": "Bearer subject-" + uuid.uuid4().hex})
auth.drain()
assert first.list_tools().ok
assert [request["subject_token"] for request in auth.token_requests()] == [first_subject]
probe: Final = {"probe": _PROBE}
own: Final = _echoed_description(first.call(f"{alias}-add", probe))
other: Final = _echoed_description(second.call(f"{alias}-add", probe))
assert (own, other) == ("Add two integers", _UNLISTED), (
"the caller bearer is part of the identity on a token-exchange server: one subject, one slot"
)
assert second.list_tools().ok
assert _echoed_description(second.call(f"{alias}-add", probe)) == "Add two integers"
assert tool_calls(peer.drain()) == (), "a blocked probe reached the peer"

View file

@ -1,13 +1,23 @@
import itertools
import uuid
from collections.abc import Mapping
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass
from hashlib import sha256
from pathlib import Path
from typing import Final
import pytest
import yaml
from integration._support.client import Gateway, eventually
from integration._support.database import read_rows
from integration._support.mcp import (
ENTRY_POINTS,
EntryPoint,
JsonRpc,
McpCaller,
Outcome,
ScriptedTool,
disconnecting_tool,
echo_tool,
listed_tools,
@ -15,8 +25,67 @@ from integration._support.mcp import (
register_mcp,
scripted_peer,
slow_tool,
text_result,
tool_calls,
)
from integration._support.process import owned_proxy_process
from integration._support.wire import Reply
from pydantic import BaseModel, TypeAdapter
_ECHO: Final = "catalog-echo:"
_PROBE: Final = "catalog-probe"
_DESCRIPTION: Final = "Look up one record"
_GUARDRAIL_CODE: Final = (
"def apply_guardrail(inputs, request_data, input_type):\n"
' texts = list(inputs.get("texts") or [])\n'
' function = inputs.get("tools", [{}])[0].get("function", {})\n'
f' if "{_PROBE}" in texts:\n'
f' return block("{_ECHO}" + json_stringify({{"description": function.get("description")}}))\n'
" return allow()\n"
)
_BURST: Final = 20
_OUTAGE: Final = 6
_SPEND_NONCES: Final = (
"SELECT status, metadata->'mcp_tool_call_metadata'->'arguments'->>'nonce' AS nonce"
' FROM "LiteLLM_SpendLogs" WHERE api_key = %s AND call_type = %s'
)
_OBJECTS: Final = TypeAdapter(Mapping[str, object])
_STRINGS: Final = TypeAdapter(Mapping[str, str])
_ECHOED: Final = TypeAdapter(Mapping[str, str | None])
class _Content(BaseModel):
text: str
class _Result(BaseModel):
content: tuple[_Content, ...]
class _RpcReply(BaseModel):
id: int
result: _Result
class _SessionsReport(BaseModel):
worker_pid: int
@pytest.fixture(scope="module")
def echo_config(tmp_path_factory: pytest.TempPathFactory) -> Path:
base: Final = _OBJECTS.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()))
guardrail: Final = {
"guardrail_name": "catalog-echo-" + uuid.uuid4().hex[:8],
"litellm_params": {
"guardrail": "custom_code",
"mode": "pre_mcp_call",
"default_on": True,
"custom_code": _GUARDRAIL_CODE,
},
}
path: Final = tmp_path_factory.mktemp("failure-recovery") / "config.yaml"
path.write_text(yaml.safe_dump({**base, "guardrails": [guardrail]}))
return path
def _call(caller: McpCaller, name: str, arguments: dict[str, object], entry: EntryPoint, identity: str) -> Outcome:
@ -134,3 +203,151 @@ def test_peer_restart_on_the_same_url_is_picked_up_without_gateway_restart(gatew
)
assert back.text == '{"a": 1, "b": 1}', back.raw
assert len(tool_calls(replacement.drain())) >= 1
def _rpc_reply(raw: str) -> _RpcReply:
data: Final = tuple(line[5:].strip() for line in raw.splitlines() if line.startswith("data:"))
return _RpcReply.model_validate_json(data[-1] if data else raw)
def _call_params(call: Mapping[str, object]) -> Mapping[str, object]:
return _OBJECTS.validate_python(_OBJECTS.validate_python(call["body"])["params"])
def _call_nonce(call: Mapping[str, object]) -> str:
return _STRINGS.validate_python(_call_params(call)["arguments"])["nonce"]
def _listed(caller: McpCaller, name: str) -> None:
listing: Final = eventually(caller.list_tools, lambda outcome: name in outcome.tools, seconds=45)
assert listing.error is None, (caller.gateway.client.base_url, listing.raw)
@dataclass(frozen=True, slots=True)
class _Worker:
caller: McpCaller
pid: int
def _worker(proxy: Gateway, key: str, alias: str) -> _Worker:
sessions: Final = proxy.client.get("/v1/mcp/sessions", headers={"x-litellm-api-key": proxy.key})
assert sessions.status_code == 200, sessions.text
return _Worker(McpCaller(proxy, key, "mcp", alias), _SessionsReport.model_validate_json(sessions.text).worker_pid)
def _served(worker: _Worker, name: str, nonce: str) -> None:
served: Final = worker.caller.call(name, {"nonce": nonce})
assert served.text == "found", (worker.pid, served.raw)
def _probed_description(worker: _Worker, name: str) -> str | None:
"""The description the pre_mcp_call guardrail on that worker was handed, recovered from its block reason."""
blocked: Final = worker.caller.call(name, {"nonce": _PROBE})
assert blocked.error is not None, (worker.pid, blocked.raw)
carrier: Final = next((item.text for item in _rpc_reply(blocked.raw).result.content if _ECHO in item.text), None)
assert carrier is not None, (worker.pid, blocked.raw)
return _ECHOED.validate_json(carrier.split(_ECHO, 1)[1])["description"]
def _catalog_is_cold(worker: _Worker, name: str) -> bool:
return not _probed_description(worker, name)
@pytest.mark.timeout(600)
def test_worker_restart_cools_its_listed_catalog_while_the_sibling_worker_keeps_serving(
gateway: Gateway, echo_config: Path, tmp_path: Path
) -> None:
tool: Final = ScriptedTool("lookup", lambda _: text_result("found"), description=_DESCRIPTION)
with scripted_peer(tool) as peer, gateway.scenario() as scenario:
alias: Final = "cold" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, peer, alias)
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
name: Final = f"{alias}-lookup"
with owned_proxy_process(gateway, tmp_path / "sibling", {}, config=echo_config) as sibling_proxy:
sibling: Final = _worker(sibling_proxy.gateway, key, alias)
with owned_proxy_process(gateway, tmp_path / "first", {}, config=echo_config) as first_proxy:
first: Final = _worker(first_proxy.gateway, key, alias)
assert first.pid != sibling.pid
_served(first, name, "first-unlisted")
_served(sibling, name, "sibling-unlisted")
assert _catalog_is_cold(first, name) and _catalog_is_cold(sibling, name)
_listed(first.caller, name)
assert _probed_description(first, name) == _DESCRIPTION
assert _catalog_is_cold(sibling, name), "a listing on one worker warmed its sibling"
_listed(sibling.caller, name)
assert _probed_description(sibling, name) == _DESCRIPTION
_served(sibling, name, "sibling-alone")
with owned_proxy_process(gateway, tmp_path / "restarted", {}, config=echo_config) as restarted_proxy:
restarted: Final = _worker(restarted_proxy.gateway, key, alias)
assert restarted.pid not in (first.pid, sibling.pid)
assert _catalog_is_cold(restarted, name), "a restarted worker kept the old process's catalog"
assert _probed_description(sibling, name) == _DESCRIPTION
_served(restarted, name, "restarted-unlisted")
_listed(restarted.caller, name)
assert _probed_description(restarted, name) == _DESCRIPTION
calls: Final = tool_calls(peer.drain())
assert [_call_nonce(call) for call in calls] == [
"first-unlisted",
"sibling-unlisted",
"sibling-alone",
"restarted-unlisted",
], calls
assert all(set(_call_params(call)) - {"_meta"} == {"name", "arguments"} for call in calls), calls
def _outage_echo(name: str, failures: int) -> ScriptedTool:
attempts: Final = itertools.count(1)
def respond(params: JsonRpc) -> Reply | JsonRpc:
if next(attempts) <= failures:
return Reply(status=503, body=b'{"error": "scripted outage"}')
return text_result(_STRINGS.validate_python(params["arguments"])["nonce"])
return ScriptedTool(name, respond)
def _echo_call(caller: McpCaller, name: str, nonce: str) -> Outcome:
return caller.call(name, {"nonce": nonce})
def _logged_nonces(key: str, count: int) -> tuple[tuple[str, str], ...]:
rows: Final = eventually(
lambda: read_rows(_SPEND_NONCES, (sha256(key.encode()).hexdigest(), "call_mcp_tool")),
lambda found: len(found) >= count,
seconds=70,
)
return tuple(sorted((str(row["status"]), str(row["nonce"])) for row in rows))
def test_peer_outage_during_a_bounded_burst_fails_exactly_the_outage_calls_and_lands_each_call_once(
gateway: Gateway, peer: Gateway
) -> None:
with scripted_peer(_outage_echo("echo", _OUTAGE)) as upstream, gateway.scenario() as scenario:
alias: Final = "burst" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, upstream, alias)
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
name: Final = f"{alias}-echo"
callers: Final = (McpCaller(gateway, key, "mcp", alias), McpCaller(peer, key, "mcp", alias))
for caller in callers:
_listed(caller, name)
upstream.drain()
nonces: Final = tuple(uuid.uuid4().hex for _ in range(_BURST))
with ThreadPoolExecutor(max_workers=_BURST) as pool:
outcomes: Final = tuple(pool.map(_echo_call, itertools.cycle(callers), itertools.repeat(name), nonces))
raws: Final = [outcome.raw for outcome in outcomes]
failed: Final = tuple(nonce for nonce, outcome in zip(nonces, outcomes) if outcome.error is not None)
assert len(failed) == _OUTAGE, raws
assert all(outcome.error is not None or outcome.text == nonce for nonce, outcome in zip(nonces, outcomes)), raws
assert all(_rpc_reply(outcome.raw).id == 1 for outcome in outcomes), raws
burst_calls: Final = tool_calls(upstream.drain())
assert sorted(_call_nonce(call) for call in burst_calls) == sorted(nonces), burst_calls
assert all(_call_params(call)["arguments"] == {"nonce": _call_nonce(call)} for call in burst_calls), burst_calls
assert all(set(_call_params(call)) - {"_meta"} == {"name", "arguments"} for call in burst_calls), burst_calls
recovered: Final = tuple(
_echo_call(caller, name, nonce) for caller, nonce in zip(itertools.cycle(callers), failed)
)
assert [outcome.text for outcome in recovered] == list(failed), [outcome.raw for outcome in recovered]
assert sorted(_call_nonce(call) for call in tool_calls(upstream.drain())) == sorted(failed)
assert _logged_nonces(key, _BURST + len(failed)) == tuple(
sorted([("success", nonce) for nonce in nonces] + [("failure", nonce) for nonce in failed])
)

View file

@ -1,5 +1,8 @@
import re
import secrets
import textwrap
import uuid
from collections.abc import Iterator
from datetime import datetime, timedelta, timezone
from pathlib import Path
from typing import Final
@ -7,14 +10,17 @@ from typing import Final
import httpx
import pytest
import yaml
from integration._support.client import Gateway, Scenario, object_value
from integration._support.client import Gateway, Scenario, gateway_from_environment, object_value
from integration._support.mcp import (
INITIALIZE,
Outcome,
ScriptedTool,
_outcome_from_rest,
_outcome_from_rpc,
mcp_peer,
register_mcp,
scripted_peer,
text_result,
tool_calls,
)
from integration._support.mcp_grants import create_toolset
@ -507,3 +513,63 @@ def test_a_member_of_two_teams_sees_the_union_and_each_route_stays_narrowed_to_i
crossed: Final = _route_call(gateway, headers, first_name, f"{alias}-multiply")
assert not crossed.ok, crossed.raw
assert tool_calls(peer.drain()) == ()
_PROBE: Final = "catalog-probe"
_ECHO: Final = "catalog-echo"
_UNLISTED: Final = ""
_GUARDRAIL_CODE: Final = (
"def apply_guardrail(inputs, request_data, input_type):\n"
f' if "{_PROBE}" not in list(inputs.get("texts") or []):\n'
" return allow()\n"
' function = inputs.get("tools", [{}])[0].get("function", {})\n'
f' return block("{_ECHO}[" + function.get("description") + "]")\n'
)
_ECHO_GUARDRAIL_YAML: Final = (
"guardrails:\n"
" - guardrail_name: catalog-echo\n"
" litellm_params:\n"
" guardrail: custom_code\n"
" mode: pre_mcp_call\n"
" default_on: true\n"
" custom_code: |\n" + textwrap.indent(_GUARDRAIL_CODE, 8 * " ")
)
@pytest.fixture(scope="module")
def echo_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]:
directory: Final = tmp_path_factory.mktemp("catalog-echo")
path: Final = directory / "catalog_echo.yaml"
path.write_text((Path(__file__).resolve().parents[1] / "proxy_config.yaml").read_text() + _ECHO_GUARDRAIL_YAML)
with gateway_from_environment() as gateway, owned_proxy(gateway, directory, {}, config=path, workers=2) as rig:
yield rig
def _echoed_description(outcome: Outcome) -> str:
found: Final = re.search(rf"{_ECHO}\[(.*?)\]", outcome.raw)
assert found is not None, outcome.raw
return found.group(1)
def test_a_team_keys_toolset_route_listing_feeds_its_own_calls_but_not_a_team_mates(echo_rig: Gateway) -> None:
described: Final = "Adds for the team " + uuid.uuid4().hex[:8]
tool: Final = ScriptedTool("add", lambda _: text_result("9"), description=described)
with scripted_peer(tool) as peer, echo_rig.scenario() as scenario:
alias: Final = "lit6029echo" + uuid.uuid4().hex[:6]
server_id: Final = register_mcp(scenario, peer, alias)
granted_id, granted_name = _toolset(scenario, server_id, "add")
team_id: Final = scenario.team(object_permission={"mcp_toolsets": [granted_id]})
key: Final = scenario.key(team_id=team_id)
team_mate: Final = scenario.key(team_id=team_id)
_assert_team_grants_only(echo_rig, team_id, key, granted_id)
listed: Final = _toolset_rpc(echo_rig, _bearer(key), granted_name, "tools/list", {})
assert listed.ok and listed.tools == (f"{alias}-add",), listed.raw
probe: Final[dict[str, object]] = {"name": f"{alias}-add", "arguments": {"probe": _PROBE}}
own: Final = _echoed_description(_toolset_rpc(echo_rig, _bearer(key), granted_name, "tools/call", probe))
mate: Final = _echoed_description(_toolset_rpc(echo_rig, _bearer(team_mate), granted_name, "tools/call", probe))
assert (own, mate) == (described, _UNLISTED), (
"the slot is keyed by the hashed key, so a team-mate that never listed is handed nothing"
)
assert tool_calls(peer.drain()) == (), "a blocked probe reached the peer"

View file

@ -226,6 +226,76 @@ async def test_anthropic_messages_with_mcp_forwards_the_callers_mcp_credentials(
assert execution["guardrail_context"] == {"metadata": {"guardrails": ("block-all",)}}
@pytest.mark.asyncio
async def test_anthropic_messages_with_mcp_hands_execution_the_requests_served_tools():
"""
Regression test: /v1/messages auto-execution must carry this request's
resolved tool definitions into execution, matching the Responses and chat
completions bridges.
Given: A request whose MCP reference resolves to a definition carrying a
description and input schema
When: The model asks for that tool and the gateway executes it
Then: _execute_tool_calls receives the definition under served_tools, so
pre_mcp_call hooks can judge the call on what the model was shown
Dropping it does not fail loudly; the call still runs, but the hook sees
only the name and arguments, leaving /v1/messages permanently colder than
the other two bridges even though all three resolve the same definitions.
"""
from mcp.types import Tool
from litellm.llms.anthropic.pass_through.messages import mcp_handler
from litellm.responses.mcp.request_context import MCPRequestContext
served = [
Tool(
name="read_wiki_structure",
description="Read the structure of a wiki",
inputSchema={"type": "object", "properties": {"repoName": {"type": "string"}}},
)
]
process = AsyncMock(return_value=(served, {"read_wiki_structure": "deepwiki"}))
execute = AsyncMock(
return_value=[{"tool_call_id": "toolu_1", "result": "ok", "name": "read_wiki_structure"}]
)
responses = [
{
"stop_reason": "tool_use",
"content": [{"type": "tool_use", "id": "toolu_1", "name": "read_wiki_structure", "input": {}}],
},
{"stop_reason": "end_turn", "content": [{"type": "text", "text": "done"}]},
]
with (
patch.object(MCPRequestContext, "resolve", return_value=MCPRequestContext(user_api_key_auth="auth")),
patch.object(
import_module("litellm.responses.mcp.litellm_proxy_mcp_handler").LiteLLM_Proxy_MCP_Handler,
"_process_mcp_tools_without_openai_transform",
new=process,
),
patch.object(
import_module("litellm.responses.mcp.litellm_proxy_mcp_handler").LiteLLM_Proxy_MCP_Handler,
"_execute_tool_calls",
new=execute,
),
patch("litellm.anthropic_messages", new=AsyncMock(side_effect=responses)),
):
await mcp_handler.anthropic_messages_with_mcp(
max_tokens=100,
messages=[{"role": "user", "content": "hi"}],
model="claude-sonnet-4-5",
tools=[MCP_REFERENCE],
)
execution = execute.call_args.kwargs
assert execution.get("served_tools") == served, (
"The request's resolved tool definitions must reach execution so pre_mcp_call "
"hooks see the listed description and input schema"
)
@pytest.mark.asyncio
async def test_anthropic_messages_with_mcp_stops_when_every_tool_call_is_skipped():
"""

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

@ -1,17 +1,29 @@
from litellm.proxy._experimental.mcp_server import operations as mcp_operations
import asyncio
import json
from datetime import datetime
from unittest.mock import AsyncMock, patch
import pytest
from fastapi import HTTPException
from mcp.shared.exceptions import MCPError
from mcp.types import CallToolResult, TextContent
from mcp.types import Tool as MCPTool
from pydantic import AnyUrl
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._experimental.mcp_server import server
from litellm.proxy._experimental.mcp_server.mcp_context import _mcp_proxy_mode
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller
from litellm.proxy._experimental.mcp_server.tool_search import (
handle_mcp_proxy_tool,
mcp_proxy_tool_id,
with_mcp_proxy_identity,
)
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth
from litellm.types.mcp import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
AUTH = UserAPIKeyAuth(api_key="key")
@ -130,3 +142,46 @@ async def test_proxy_scope_exception_emits_failure_log(monkeypatch: pytest.Monke
assert hook_payload["arguments"] == arguments
assert "raw_headers" not in hook_payload
assert "raw-scope-secret" not in recorder.events[1][1]
@pytest.mark.asyncio
async def test_proxy_call_tool_on_a_never_listed_tool_hands_the_pre_hook_no_listed_tool() -> None:
"""/mcp/proxy tools/list serves only the meta-tools, so the catalog call_tool reads to resolve its
tool_id was never served: it must not fill the caller's listed-tools slot, and the pre-call hook
must see no listed tool for the call."""
manager = mcp_operations.global_mcp_server_manager
server = MCPServer(server_id="proxy-meta", name="proxy-meta", transport=MCPTransport.http, url="http://meta")
auth = UserAPIKeyAuth(api_key="sk-proxy-meta", user_id="proxy-caller")
upstream = [MCPTool(name="echo", description="Echo text back", inputSchema={"type": "object"})]
served_as = with_mcp_proxy_identity(MCPTool(name="proxy-meta-echo", inputSchema={}), server.server_id)
pre_call_tool_check = AsyncMock(return_value={})
async def call_regular_mcp_tool(*, tasks: list[asyncio.Task[object]], **_: object) -> CallToolResult:
await asyncio.gather(*tasks)
return CallToolResult(content=[TextContent(type="text", text="echoed")])
with (
patch.dict(manager.registry, {server.server_id: server}),
patch.dict(manager.tool_name_to_mcp_server_name_mapping),
patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])),
patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())),
patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)),
patch.object(manager, "pre_call_tool_check", pre_call_tool_check),
patch.object(manager, "_call_regular_mcp_tool", call_regular_mcp_tool),
):
try:
result = await handle_mcp_proxy_tool(
name="call_tool",
arguments={"tool_id": mcp_proxy_tool_id(served_as), "arguments": {}},
user_api_key_dict=auth,
)
listed = manager.get_listed_tool(server, "echo", ListedToolsCaller(user_api_key_auth=auth))
finally:
manager._drop_listed_tools(server.server_id)
assert result.is_error is False
assert result.content[0].text == "echoed"
pre_call_tool_check.assert_awaited_once()
assert pre_call_tool_check.await_args.kwargs["name"] == "echo"
assert pre_call_tool_check.await_args.kwargs["tool"] is None
assert listed is None

View file

@ -963,6 +963,8 @@ async def test_get_tools_from_mcp_servers():
user_api_key_auth=None,
oauth2_headers=None,
proxy_logging_obj=None,
catalog_auth_header=None,
record_listing=True,
):
if server.server_id == "server1_id":
return [mock_tool_1]
@ -1998,6 +2000,7 @@ async def test_get_tools_for_single_server():
client_ip=None,
user_api_key_auth=None,
proxy_logging_obj=ANY,
record_listing=False,
)
# Verify the result

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
@ -23,18 +24,20 @@ from mcp.types import (
TextContent,
TextResourceContents,
)
from mcp.types import Tool as MCPTool
from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS, LATEST_HANDSHAKE_VERSION, MODERN_PROTOCOL_VERSIONS
from pydantic import TypeAdapter
from starlette.types import Message, Receive, Scope, Send
from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller
from litellm.proxy._types import (
LiteLLM_MCPServerTable,
MCPTransport,
UserAPIKeyAuth,
)
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer
from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer, PinnedMCPTool
def test_mcp_available_on_sdk2():
@ -85,9 +88,6 @@ def cleanup_mcp_global_state():
yield
def _call_tool_params(name, arguments=None):
from mcp.types import CallToolRequestParams
@ -99,6 +99,7 @@ def _paged_params():
return PaginatedRequestParams()
@pytest.mark.asyncio
async def test_mcp_server_tool_call_body_contains_request_data(_mcp_request_ctx):
"""Test that proxy_server_request body contains name and arguments"""
@ -295,7 +296,9 @@ async def test_mcp_server_tool_call_relays_upstream_auth_error_as_iserror(_mcp_r
):
with patch("litellm.proxy.proxy_server.proxy_config", MagicMock()):
with patch("litellm.proxy._experimental.mcp_server.operations.verbose_logger", mock_logger):
result = await mcp_server_tool_call(_mcp_request_ctx(), _call_tool_params("test_tool", {"param": "value"}))
result = await mcp_server_tool_call(
_mcp_request_ctx(), _call_tool_params("test_tool", {"param": "value"})
)
assert result.is_error is True
# The dedicated MCPUpstreamAuthError branch (not the generic Exception fallthrough) produces this
@ -1167,20 +1170,32 @@ async def test_read_resource_preserves_content_metadata(_mcp_request_ctx, kind,
else BlobResourceContents(uri=uri, blob="aGVsbG8=", mimeType="image/png", meta=metadata)
)
with (
patch.object(server, "get_or_extract_auth_context", AsyncMock(return_value=(caller, None, ["catalog"], None, None, None, None))),
patch.object(
server,
"get_or_extract_auth_context",
AsyncMock(return_value=(caller, None, ["catalog"], None, None, None, None)),
),
patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[upstream_server])),
patch.object(operations.global_mcp_server_manager, "read_resource_from_server", AsyncMock(return_value=ReadResourceResult(contents=[content]))),
patch.object(
operations.global_mcp_server_manager,
"read_resource_from_server",
AsyncMock(return_value=ReadResourceResult(contents=[content])),
),
):
result: Final = await server.read_resource(_mcp_request_ctx(), ReadResourceRequestParams(uri=uri))
assert result.model_dump(mode="json", by_alias=True, exclude_none=True) == {
"cacheScope": "private", "resultType": "complete", "ttlMs": 0,
"contents": [{
"uri": uri,
"mimeType": "text/plain" if kind == "text" else "image/png",
"text" if kind == "text" else "blob": "hello world" if kind == "text" else "aGVsbG8=",
**({"_meta": metadata} if metadata is not None else {}),
}],
"cacheScope": "private",
"resultType": "complete",
"ttlMs": 0,
"contents": [
{
"uri": uri,
"mimeType": "text/plain" if kind == "text" else "image/png",
"text" if kind == "text" else "blob": "hello world" if kind == "text" else "aGVsbG8=",
**({"_meta": metadata} if metadata is not None else {}),
}
],
}
@ -1674,7 +1689,9 @@ async def test_handle_list_tools_converts_permission_httpexception_to_mcp_error(
with (
patch( # test-quality-ok: the protocol handler reads auth from module context; no injection seam
"litellm.proxy._experimental.mcp_server.server.get_or_extract_auth_context",
new=AsyncMock(return_value=(None, None, None, None, None, None, None), side_effect=denial if denial_at_auth else None),
new=AsyncMock(
return_value=(None, None, None, None, None, None, None), side_effect=denial if denial_at_auth else None
),
),
patch( # test-quality-ok: the listing helper is the handler's only collaborator; the suite's seam
"litellm.proxy._experimental.mcp_server.operations._list_mcp_tools",
@ -1913,8 +1930,8 @@ async def test_streamable_http_session_manager_is_stateless():
("DELETE", b"", False),
),
)
async def test_mcp_routing_initialize_to_stateful_no_session_to_stateless(_mcp_request_ctx,
debug: bool, method: str, request_body: bytes, stateful: bool
async def test_mcp_routing_initialize_to_stateful_no_session_to_stateless(
_mcp_request_ctx, debug: bool, method: str, request_body: bytes, stateful: bool
) -> None:
from starlette.requests import Request
from starlette.types import Message, Receive, Scope, Send
@ -4056,7 +4073,8 @@ async def test_truncated_jsonrpc_response_with_nested_method_skips_lock(
# parsed, with a nested "method" key in the first bytes to trip a flat
# substring heuristic.
response_prefix: Final = (
'{"jsonrpc":"2.0","id":99,"' + response_field
'{"jsonrpc":"2.0","id":99,"'
+ response_field
+ '":{"code":-32000,"message":"test","data":{"method":"GET","payload":"'
).encode()
response_body: Final = (
@ -6564,8 +6582,12 @@ class TestGatewayCreateInitializationOptions:
yield (None, None)
async def record_request(
serving_server: object, read_stream: object, write_stream: object,
*, lifespan_state: object, init_options: InitializationOptions,
serving_server: object,
read_stream: object,
write_stream: object,
*,
lifespan_state: object,
init_options: InitializationOptions,
) -> None:
captured["server_name"] = init_options.server_name
@ -6877,7 +6899,6 @@ async def test_probe_upstream_auth_surfaces_httpx_status_error():
returning the response. The probe must catch that specifically (before the
fail-open `except Exception`) so the auth check is not silently defeated.
"""
import httpx
from litellm.proxy._experimental.mcp_server.server import _probe_upstream_auth
@ -7412,7 +7433,8 @@ async def test_execute_mcp_tool_rest_server_id_authoritative_for_unprefixed_tool
return_value=oauth_server,
),
patch.object(
mcp_operations, "_handle_managed_mcp_tool",
mcp_operations,
"_handle_managed_mcp_tool",
new=fake_handle_managed_mcp_tool,
),
patch.object(
@ -7658,7 +7680,8 @@ async def test_execute_mcp_tool_strips_a_prefix_that_contains_the_separator():
return_value=alias_less_server,
),
patch.object(
mcp_operations, "_handle_managed_mcp_tool",
mcp_operations,
"_handle_managed_mcp_tool",
new=fake_handle_managed_mcp_tool,
),
patch.object(
@ -7933,7 +7956,8 @@ async def test_execute_mcp_tool_rest_hyphenated_upstream_tool_name_routes_to_req
return_value=None,
),
patch.object(
mcp_operations, "_handle_managed_mcp_tool",
mcp_operations,
"_handle_managed_mcp_tool",
new=fake_handle_managed_mcp_tool,
),
patch.object(
@ -7989,9 +8013,12 @@ async def test_execute_mcp_tool_sets_model_in_model_call_details():
fake_server.server_name = "openapi-petstore"
fake_server.alias = None
fake_server.short_prefix = None
fake_server.tool_name_to_description = None
fake_tool = MagicMock()
fake_tool.name = "list_pets"
fake_tool.description = "test tool"
fake_tool.input_schema = {"type": "object"}
start_time = datetime.now(timezone.utc)
litellm_logging_obj, _ = function_setup(
@ -8042,6 +8069,348 @@ async def test_execute_mcp_tool_sets_model_in_model_call_details():
assert litellm_logging_obj.model == "MCP: list_pets"
@pytest.mark.asyncio
async def test_execute_mcp_tool_hands_openapi_hooks_the_listed_entry_and_nothing_before_a_listing():
"""A local-registry tools/call with no prior tools/list hands the pre-call hooks name and arguments
only, as before this metadata existed, so a pre_mcp_call policy never scans a description the caller was
not served. Once the caller has listed, the same call hands the entry that listing served."""
from litellm.caching.caching import DualCache
from litellm.proxy._experimental.mcp_server import operations as mcp_module
from litellm.proxy.utils import ProxyLogging
petstore = MCPServer(
server_id="petstore-id",
name="petstore",
server_name="petstore",
transport=MCPTransport.http,
url=None,
spec_path="https://example.com/petstore.yaml",
tool_name_to_description={"list_pets": "ADMIN DESC"},
)
schema = {"type": "object", "properties": {"limit": {"type": "integer"}}}
mcp_module.global_mcp_tool_registry.register_tool(
name="petstore-list_pets", description="List the pets", input_schema=schema, handler=lambda limit: "ok"
)
manager = mcp_module.global_mcp_server_manager
manager._listed_tools_by_server_id.pop(petstore.server_id, None)
alice = UserAPIKeyAuth(api_key="sk-user", user_id="alice")
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging.pre_call_hook = AsyncMock(return_value={})
pre_call_tool_check = AsyncMock(wraps=manager.pre_call_tool_check)
async def call() -> tuple[MCPTool | None, dict]:
await mcp_module.execute_mcp_tool(
name="petstore-list_pets",
arguments={"limit": 10},
allowed_mcp_servers=[petstore],
start_time=datetime.now(),
user_api_key_auth=alice,
)
return pre_call_tool_check.call_args.kwargs["tool"], proxy_logging.pre_call_hook.call_args.kwargs["data"]
try:
with (
patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore),
patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check),
patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging),
):
never_listed_tool, never_listed_data = await call()
manager._record_listed_tools(
petstore,
[MCPTool(name="list_pets", description="ADMIN DESC", inputSchema=schema)],
ListedToolsCaller(user_api_key_auth=alice),
)
listed_tool, listed_data = await call()
finally:
mcp_module.global_mcp_tool_registry.unregister_tools_with_prefix("petstore-")
manager._listed_tools_by_server_id.pop(petstore.server_id, None)
assert never_listed_tool is None
assert (never_listed_data.get("mcp_tool_description"), never_listed_data.get("mcp_input_schema")) == (None, None)
assert listed_tool is not None and (listed_tool.description, listed_tool.input_schema) == ("ADMIN DESC", schema)
assert (listed_data["mcp_tool_description"], listed_data["mcp_input_schema"]) == ("ADMIN DESC", schema)
@pytest.mark.asyncio
async def test_execute_mcp_tool_hands_openapi_hooks_the_guarded_catalog_entry_clients_saw():
"""When tools/list pinned the schema and masked the description of an OpenAPI tool, the local-registry
call path must hand the pre-call hooks that served entry, not the raw registry one."""
from litellm.proxy._experimental.mcp_server import operations as mcp_module
petstore = MCPServer(
server_id="petstore-id",
name="petstore",
server_name="petstore",
transport=MCPTransport.http,
url=None,
spec_path="https://example.com/petstore.yaml",
tool_name_to_description={"getpetbyid": "Find a SECRET pet"},
)
registry_schema = {"type": "object", "properties": {"petId": {"type": "integer"}, "dump_all": {"type": "boolean"}}}
pinned_schema = {"type": "object", "properties": {"petId": {"type": "integer"}}}
mcp_module.global_mcp_tool_registry.register_tool(
name="petstore-getpetbyid",
description="Find pet by ID",
input_schema=registry_schema,
handler=lambda petId: "ok",
)
manager = mcp_module.global_mcp_server_manager
alice = UserAPIKeyAuth(api_key="sk-user", user_id="alice")
manager._record_listed_tools(
petstore,
[MCPTool(name="getpetbyid", description="Find a [MASKED] pet", inputSchema=pinned_schema)],
ListedToolsCaller(user_api_key_auth=alice),
)
pre_call_tool_check = AsyncMock(return_value={})
try:
with (
patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore),
patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check),
):
await mcp_module.execute_mcp_tool(
name="petstore-getpetbyid",
arguments={"petId": 1},
allowed_mcp_servers=[petstore],
start_time=datetime.now(),
user_api_key_auth=alice,
)
finally:
mcp_module.global_mcp_tool_registry.unregister_tools_with_prefix("petstore-")
manager._listed_tools_by_server_id.pop(petstore.server_id, None)
handed_tool = pre_call_tool_check.call_args.kwargs["tool"]
assert (handed_tool.description, handed_tool.input_schema) == ("Find a [MASKED] pet", pinned_schema), (
"the pre-call policy must evaluate the entry tools/list served, not the raw registry entry"
)
@pytest.mark.asyncio
async def test_execute_mcp_tool_hands_openapi_hooks_each_callers_own_listed_entry():
"""Two keys can be shown differently guarded OpenAPI catalogs. The call path must evaluate each key
against the entry its own tools/list served, not the entry the most recent listing left behind."""
from litellm.proxy._experimental.mcp_server import operations as mcp_module
petstore = MCPServer(
server_id="petstore-id",
name="petstore",
server_name="petstore",
transport=MCPTransport.http,
url=None,
spec_path="https://example.com/petstore.yaml",
)
schema = {"type": "object", "properties": {"petId": {"type": "integer"}}}
mcp_module.global_mcp_tool_registry.register_tool(
name="petstore-getpetbyid", description="Find a SECRET pet", input_schema=schema, handler=lambda petId: "ok"
)
manager = mcp_module.global_mcp_server_manager
guarded = UserAPIKeyAuth(api_key="sk-guarded", user_id="alice")
opted_out = UserAPIKeyAuth(api_key="sk-opted-out", user_id="bob")
manager._record_listed_tools(
petstore,
[MCPTool(name="getpetbyid", description="Find a [MASKED] pet", inputSchema=schema)],
ListedToolsCaller(user_api_key_auth=guarded),
)
manager._record_listed_tools(
petstore,
[MCPTool(name="getpetbyid", description="Find a SECRET pet", inputSchema=schema)],
ListedToolsCaller(user_api_key_auth=opted_out),
)
pre_call_tool_check = AsyncMock(return_value={})
try:
with (
patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore),
patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check),
):
for caller in (guarded, opted_out):
await mcp_module.execute_mcp_tool(
name="petstore-getpetbyid",
arguments={"petId": 1},
allowed_mcp_servers=[petstore],
start_time=datetime.now(),
user_api_key_auth=caller,
)
finally:
mcp_module.global_mcp_tool_registry.unregister_tools_with_prefix("petstore-")
manager._listed_tools_by_server_id.pop(petstore.server_id, None)
handed = [call.kwargs["tool"].description for call in pre_call_tool_check.call_args_list]
assert handed == ["Find a [MASKED] pet", "Find a SECRET pet"], (
"each key's tools/call must be evaluated against the OpenAPI entry its own listing served"
)
@pytest.mark.asyncio
async def test_execute_mcp_tool_runs_the_longer_colliding_operation_and_hands_hooks_no_registry_metadata():
"""An OpenAPI operation whose name starts with its own server prefix runs instead of the shorter one, and
with no prior listing the pre-call hooks get name and arguments only, never either registry entry."""
from litellm.proxy._experimental.mcp_server import operations as mcp_module
petstore = MCPServer(
server_id="petstore-id",
name="petstore",
server_name="petstore",
transport=MCPTransport.http,
url=None,
spec_path="https://example.com/petstore.yaml",
)
registry = mcp_module.global_mcp_tool_registry
registry.register_tool(name="petstore-get_pet", description="short", input_schema={}, handler=lambda: "short")
registry.register_tool(
name="petstore-petstore-get_pet",
description="long",
input_schema={"type": "object", "properties": {"petId": {"type": "integer"}}},
handler=lambda: "long",
)
manager = mcp_module.global_mcp_server_manager
pre_call_tool_check = AsyncMock(return_value={})
try:
with (
patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore),
patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check),
):
result = await mcp_module.execute_mcp_tool(
name="petstore-petstore-get_pet",
arguments={},
allowed_mcp_servers=[petstore],
start_time=datetime.now(),
user_api_key_auth=UserAPIKeyAuth(api_key="sk-user", user_id="alice"),
)
finally:
registry.unregister_tools_with_prefix("petstore-")
assert pre_call_tool_check.call_args.kwargs["tool"] is None
assert result.content[0].text == "long"
@pytest.mark.asyncio
async def test_execute_mcp_tool_hands_hooks_nothing_for_a_never_listed_operation_named_after_a_listed_one():
"""After the caller listed ``get_pet``, a call to the never-listed ``petstore-get_pet`` operation hands the
pre-call hooks name and arguments only, not the listed sibling's description and schema."""
from litellm.proxy._experimental.mcp_server import operations as mcp_module
petstore = MCPServer(
server_id="petstore-id",
name="petstore",
server_name="petstore",
transport=MCPTransport.http,
url=None,
spec_path="https://example.com/petstore.yaml",
)
registry = mcp_module.global_mcp_tool_registry
registry.register_tool(
name="petstore-petstore-get_pet", description="long", input_schema={}, handler=lambda: "long"
)
manager = mcp_module.global_mcp_server_manager
alice = UserAPIKeyAuth(api_key="sk-user", user_id="alice")
manager._record_listed_tools(
petstore,
[MCPTool(name="get_pet", description="Fetches pet records. FLAGWORD", inputSchema={"type": "object"})],
ListedToolsCaller(user_api_key_auth=alice),
)
pre_call_tool_check = AsyncMock(return_value={})
try:
with (
patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore),
patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check),
):
result = await mcp_module.execute_mcp_tool(
name="petstore-petstore-get_pet",
arguments={},
allowed_mcp_servers=[petstore],
start_time=datetime.now(),
user_api_key_auth=alice,
)
finally:
registry.unregister_tools_with_prefix("petstore-")
manager._listed_tools_by_server_id.pop(petstore.server_id, None)
assert pre_call_tool_check.call_args.kwargs["tool"] is None
assert result.content[0].text == "long"
@pytest.mark.asyncio
async def test_execute_mcp_tool_implicit_listing_before_the_first_call_hands_hooks_no_description():
"""The listing tools/call runs on its own when this worker does not yet expose the tool is never served
to the caller, so it leaves the caller's listed slot empty and the pre-call hooks still get name and
arguments only, as on main."""
manager = mcp_operations.global_mcp_server_manager
server = _never_listed_passthrough_server()
manager.registry[server.server_id] = server
manager._listed_tools_by_server_id.pop(server.server_id, None)
upstream = AsyncMock()
upstream.call_tool.return_value = CallToolResult(content=[TextContent(type="text", text="ok")], isError=False)
proxy_logging = _mock_mcp_proxy_logging()
proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={})
proxy_logging.pre_call_hook = AsyncMock(return_value={})
proxy_logging.during_call_hook = AsyncMock(return_value=None)
fetch_tools = AsyncMock(
return_value=[MCPTool(name="add", description="Adds. FLAGWORD", inputSchema={"type": "object"})]
)
with (
patch.object(manager, "_create_mcp_client", new=AsyncMock(return_value=upstream)),
patch.object(manager, "_fetch_tools_with_timeout", new=fetch_tools),
patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging),
):
result = await mcp_operations.execute_mcp_tool(
name="lazy_map-add",
arguments={"a": 1, "b": 2},
allowed_mcp_servers=[server],
start_time=datetime.now(),
mcp_auth_header="Bearer caller-token",
raw_headers={"authorization": "Bearer caller-token"},
)
assert fetch_tools.await_count == 1
assert upstream.call_tool.await_count == 1
assert result.content[0].text == "ok"
hook_kwargs = proxy_logging._create_mcp_request_object_from_kwargs.call_args.args[0]
assert (hook_kwargs["tool_description"], hook_kwargs["tool_input_schema"]) == (None, None)
assert server.server_id not in manager._listed_tools_by_server_id
@pytest.mark.asyncio
async def test_fetch_pinnable_tool_catalog_records_no_listed_catalog_for_the_admin():
"""The pin snapshot lists the raw upstream catalog, without the catalog guard or the admin's description
overrides, so it must not become what the admin's own later tools/call is evaluated against."""
from litellm.caching.caching import DualCache
from litellm.proxy._experimental.mcp_server.rest_endpoints import fetch_pinnable_tool_catalog
from litellm.proxy.utils import ProxyLogging
manager = mcp_operations.global_mcp_server_manager
server = MCPServer(
server_id="pin-srv",
name="pin_srv",
transport=MCPTransport.http,
url="https://up.example.com/mcp",
tool_name_to_description={"add": "Admin wording"},
)
manager._listed_tools_by_server_id.pop(server.server_id, None)
admin = UserAPIKeyAuth(api_key="sk-admin", user_id="admin")
request = MagicMock()
request.client.host = "10.1.2.3"
request.headers = {"x-litellm-api-key": "sk-admin"}
fetch_tools = AsyncMock(
return_value=[MCPTool(name="add", description="Upstream wording", inputSchema={"type": "object"})]
)
with (
patch.object(manager, "_create_mcp_client", new=AsyncMock(return_value=MagicMock())),
patch.object(manager, "_fetch_tools_with_timeout", new=fetch_tools),
patch("litellm.proxy.proxy_server.proxy_logging_obj", ProxyLogging(user_api_key_cache=DualCache())),
):
snapshot = await fetch_pinnable_tool_catalog(server, request, admin)
assert snapshot == {"add": PinnedMCPTool(description="Upstream wording", input_schema={"type": "object"})}
assert server.server_id not in manager._listed_tools_by_server_id
assert manager.get_listed_tool(server, "add", ListedToolsCaller(user_api_key_auth=admin)) is None
@pytest.mark.asyncio
async def test_execute_mcp_tool_rest_unresolved_prefixed_name_routes_to_requested_server():
"""A prefixed REST name that resolves to no tool must still dispatch to the server_id.
@ -8098,7 +8467,8 @@ async def test_execute_mcp_tool_rest_unresolved_prefixed_name_routes_to_requeste
return_value=None,
),
patch.object(
mcp_operations, "_handle_managed_mcp_tool",
mcp_operations,
"_handle_managed_mcp_tool",
new=fake_handle_managed_mcp_tool,
),
patch.object(
@ -8597,7 +8967,9 @@ class TestMCPMetaTraceCarrier:
assert _mcp_meta_trace_carrier(None) is None
assert _mcp_meta_trace_carrier(SimpleNamespace(meta=None)) is None
only_progress = CallToolRequestParams.model_validate({"name": "t", "_meta": {"progressToken": "p1"}}, by_name=False).meta
only_progress = CallToolRequestParams.model_validate(
{"name": "t", "_meta": {"progressToken": "p1"}}, by_name=False
).meta
assert _mcp_meta_trace_carrier(SimpleNamespace(meta=only_progress)) is None
@ -10394,7 +10766,9 @@ async def test_mcp_origin_admission_precedes_authentication(
patch("litellm.proxy.proxy_server.origins", allowed_origins),
patch.object(server, "extract_mcp_auth_context", authenticate),
):
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=server.app), base_url="http://gateway") as client:
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app=server.app), base_url="http://gateway"
) as client:
response: Final = await client.request(method, path, headers=(*session_headers, *origin_headers))
assert response.status_code == expected_status
@ -10477,12 +10851,15 @@ async def test_streamable_http_rejects_modern_protocol_version(
@pytest.mark.asyncio
@pytest.mark.parametrize("handler_name,field", [
("handle_list_tools", "tools"),
("list_prompts", "prompts"),
("list_resources", "resources"),
("list_resource_templates", "resource_templates"),
])
@pytest.mark.parametrize(
"handler_name,field",
[
("handle_list_tools", "tools"),
("list_prompts", "prompts"),
("list_resources", "resources"),
("list_resource_templates", "resource_templates"),
],
)
async def test_native_listing_preserves_empty_result_on_auth_failure(_mcp_request_ctx, handler_name, field):
from litellm.proxy._experimental.mcp_server import server
@ -10500,7 +10877,9 @@ async def test_tool_listing_preserves_permission_denial_when_failure_logging_fai
auth = UserAPIKeyAuth(user_id="denied-caller")
denial = HTTPException(status_code=403, detail="scope denied")
logger = MagicMock()
logger.post_call_failure_hook = AsyncMock(side_effect=RuntimeError("log unavailable") if failure_hook_raises else None)
logger.post_call_failure_hook = AsyncMock(
side_effect=RuntimeError("log unavailable") if failure_hook_raises else None
)
upstream = AsyncMock()
with (
patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(side_effect=denial)),
@ -10509,7 +10888,9 @@ async def test_tool_listing_preserves_permission_denial_when_failure_logging_fai
patch.object(operations.global_mcp_server_manager, "_get_tools_from_server", upstream),
):
with pytest.raises(HTTPException) as rejected:
await operations._get_tools_from_mcp_servers(user_api_key_auth=auth, mcp_auth_header=None, mcp_servers=["catalog"], log_list_tools_to_spendlogs=True)
await operations._get_tools_from_mcp_servers(
user_api_key_auth=auth, mcp_auth_header=None, mcp_servers=["catalog"], log_list_tools_to_spendlogs=True
)
assert rejected.value is denial
upstream.assert_not_awaited()
logger.post_call_failure_hook.assert_awaited_once()
@ -10521,7 +10902,9 @@ async def test_tool_listing_preserves_permission_denial_when_failure_logging_fai
@pytest.mark.parametrize("prefix,suffix", (("", ""), ("/gateway", "/")))
@pytest.mark.parametrize("opening_protocol", (None, *MODERN_PROTOCOL_VERSIONS))
async def test_legacy_sse_mount_emits_message_endpoint(
prefix: str, suffix: str, opening_protocol: str | None,
prefix: str,
suffix: str,
opening_protocol: str | None,
) -> None:
from starlette.applications import Starlette
from starlette.routing import Mount
@ -10586,16 +10969,20 @@ async def test_legacy_sse_mount_emits_message_endpoint(
return (await messages.get())["status"]
if opening_protocol is not None:
discover: Final = json.dumps({
"jsonrpc": "2.0",
"id": 0,
"method": "server/discover",
"params": {"_meta": {
"io.modelcontextprotocol/protocolVersion": opening_protocol,
"io.modelcontextprotocol/clientInfo": {"name": "modern-client", "version": "1"},
"io.modelcontextprotocol/clientCapabilities": {},
}},
}).encode()
discover: Final = json.dumps(
{
"jsonrpc": "2.0",
"id": 0,
"method": "server/discover",
"params": {
"_meta": {
"io.modelcontextprotocol/protocolVersion": opening_protocol,
"io.modelcontextprotocol/clientInfo": {"name": "modern-client", "version": "1"},
"io.modelcontextprotocol/clientCapabilities": {},
}
},
}
).encode()
assert await post(discover) == 202
discovered_frame: Final = (await asyncio.wait_for(outgoing.get(), 2))["body"].decode()
discovered: Final = json.loads(discovered_frame.split("data: ", 1)[1].splitlines()[0])
@ -10628,7 +11015,16 @@ async def test_legacy_sse_mount_emits_message_endpoint(
patch.object(
mcp_server,
"extract_mcp_auth_context",
AsyncMock(return_value=(post_auth, None, [marker], {marker: {"Authorization": marker}}, {"Authorization": marker}, {"x-request-marker": marker})),
AsyncMock(
return_value=(
post_auth,
None,
[marker],
{marker: {"Authorization": marker}},
{"Authorization": marker},
{"x-request-marker": marker},
)
),
),
patch.object(mcp_server.operations, "_get_tools_from_mcp_servers", listing),
):
@ -10677,7 +11073,11 @@ async def test_discovery_adapter_preserves_authenticated_context(_mcp_request_ct
dispatched = AsyncMock(return_value=expected)
auth = UserAPIKeyAuth(user_id="discover-caller")
with (
patch.object(server, "get_or_extract_auth_context", AsyncMock(return_value=(auth, None, ["allowed"], None, None, None, None))),
patch.object(
server,
"get_or_extract_auth_context",
AsyncMock(return_value=(auth, None, ["allowed"], None, None, None, None)),
),
patch.object(server.operations.GatewayOperations, "execute", dispatched),
):
result = await server.discover(_mcp_request_ctx(), RequestParams())

View file

@ -23,6 +23,7 @@ from mcp.types import Tool
import litellm
from litellm.models.object_permission import LiteLLM_ObjectPermissionTable
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller
from litellm.proxy._experimental.mcp_server.tool_search import (
AGENT_SEARCH_TOOL_NAME,
MCP_TOOL_CALL_TOOL_NAME,
@ -32,12 +33,14 @@ from litellm.proxy._experimental.mcp_server.tool_search import (
ToolSearchResult,
coerce_top_k,
get_virtual_tool_definitions,
handle_mcp_tool_search,
search_mcp_tools,
search_tools,
)
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.common_utils.semantic_text_index import EmbeddingFailed, SemanticTextIndex, Vector
from litellm.types.mcp import MCPToolSearchSettings
from litellm.types.mcp import MCPToolSearchSettings, MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
def _make_tools(specs: list[tuple[str, str]]) -> tuple[Tool, ...]:
@ -1353,3 +1356,33 @@ async def test_handle_mcp_tool_call_scoped_denial_names_the_binding_agent() -> N
assert exc_info.value.status_code == 403
assert "MCP server 'github'" in exc_info.value.detail["error"]
assert "agent 'agent-123'" in exc_info.value.detail["error"]
@pytest.mark.asyncio
async def test_mcp_tool_search_leaves_the_listed_tools_slot_empty(monkeypatch: pytest.MonkeyPatch) -> None:
"""The search lists the whole catalog but serves only its hits, so the listing must not fill the
caller's listed-tools slot: a later call to a tool the search never returned is not a listed tool."""
monkeypatch.setattr(litellm, "mcp_tool_search", None)
manager = mcp_operations.global_mcp_server_manager
server = MCPServer(server_id="search-slot", name="search-slot", transport=MCPTransport.http, url="http://slot")
user = UserAPIKeyAuth(api_key="sk-search-slot", user_id="searcher")
upstream = [
Tool(name="echo", description="Echo text back", inputSchema={"type": "object"}),
Tool(name="delete_note", description="Delete a note", inputSchema={"type": "object"}),
]
with (
patch.dict(manager.tool_name_to_mcp_server_name_mapping),
patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])),
patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())),
patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)),
):
try:
result = await handle_mcp_tool_search(query="echo", top_k=1, user_api_key_dict=user)
caller = ListedToolsCaller(user_api_key_auth=user)
listed = [manager.get_listed_tool(server, tool.name, caller) for tool in upstream]
finally:
manager._drop_listed_tools(server.server_id)
assert result.is_error is False
assert [hit["name"] for hit in json.loads(result.content[0].text)] == ["search-slot-echo"]
assert listed == [None, None]

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

@ -1,15 +1,101 @@
import asyncio
from typing import Final
from unittest.mock import AsyncMock, patch
import pytest
from mcp.types import GetPromptRequest, GetPromptRequestParams, GetPromptResult
from mcp.types import Tool as MCPTool
import litellm
from litellm.caching.dual_cache import DualCache
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._experimental.mcp_server import operations
from litellm.proxy._experimental.mcp_server import rest_endpoints
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller
from litellm.proxy._experimental.mcp_server.operations import GatewayOperations, prepare_context
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.utils import ProxyLogging
from litellm.types.mcp import MCPAuth, MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
class _CatalogHookCapture(CustomLogger):
data: dict[str, object] | None = None
async def async_pre_call_hook(
self, user_api_key_dict: UserAPIKeyAuth, cache: DualCache, data: dict[str, object], call_type: str
) -> None:
if call_type == "call_mcp_tool":
self.data = data.copy()
async def _served_catalog_tool() -> str:
return "ok"
@pytest.mark.asyncio
@pytest.mark.parametrize("surface", ["mcp", "rest"])
@pytest.mark.parametrize("restriction", ["key", "server"])
async def test_listing_records_only_tools_the_caller_received(
monkeypatch: pytest.MonkeyPatch, surface: str, restriction: str
) -> None:
manager: Final = operations.global_mcp_server_manager
server: Final = MCPServer(
server_id="served-catalog", name="served-catalog", transport=MCPTransport.http,
spec_path="/catalog.yaml", allow_all_keys=True,
allowed_tools=["echo"] if restriction == "server" else None,
)
auth: Final = UserAPIKeyAuth(
api_key="sk-served-catalog", user_id="lister",
object_permission={
"object_permission_id": "served-permission",
"mcp_servers": [server.server_id],
"mcp_tool_permissions": {server.server_id: ["echo"]} if restriction == "key" else None,
},
)
monkeypatch.setitem(manager.registry, server.server_id, server)
monkeypatch.setitem(manager.tool_name_to_mcp_server_name_mapping, "status", server.server_id)
monkeypatch.setitem(manager.tool_name_to_mcp_server_name_mapping, "served-catalog-status", server.server_id)
capture: Final = _CatalogHookCapture()
monkeypatch.setattr(litellm, "callbacks", [capture])
for name in ("echo", "status"):
global_mcp_tool_registry.register_tool(
name=f"served-catalog-{name}", description=f"{name} description",
input_schema={"type": "object"}, handler=_served_catalog_tool,
)
try:
if surface == "mcp":
listing: Final = await operations._list_mcp_tools(
user_api_key_auth=auth, mcp_servers=[server.server_id], record_listing=True,
)
assert [tool.name for tool in listing.tools] == ["served-catalog-echo"]
else:
rest_listing: Final = await rest_endpoints._get_tools_for_single_server(
server, None, user_api_key_auth=auth,
)
assert [tool.name for tool in rest_listing] == ["echo"]
granted: Final = auth.model_copy(update={"object_permission": None})
caller: Final = ListedToolsCaller(user_api_key_auth=granted)
assert manager.get_listed_tool(server, "status", caller) is None
served: Final = manager.get_listed_tool(server, "echo", caller)
assert served is not None
assert (served.description, served.input_schema) == ("echo description", {"type": "object"})
server.allowed_tools = None
result: Final = await manager.call_tool(
server_name=server.server_id, name="status", arguments={}, user_api_key_auth=granted,
proxy_logging_obj=ProxyLogging(user_api_key_cache=UserApiKeyCache()),
)
assert result.is_error is False
assert capture.data is not None
assert capture.data["messages"] == [{"role": "user", "content": "Tool: status\nArguments: {}"}]
assert (capture.data.get("mcp_tool_description"), capture.data.get("mcp_input_schema")) == (None, None)
finally:
manager._drop_listed_tools(server.server_id)
global_mcp_tool_registry.unregister_tools_with_prefix("served-catalog-")
@pytest.mark.asyncio
async def test_oauth_prefetch_failure_does_not_log_caller_or_exception_text(caplog):
from litellm.proxy._experimental.mcp_server.operations import _prefetch_oauth_creds_for_user
@ -665,3 +751,30 @@ async def test_tools_listing_preserves_explicit_spend_log_policy(log_enabled):
)
assert result.tools == []
assert listing.await_args.kwargs["log_list_tools_to_spendlogs"] is log_enabled
assert listing.await_args.kwargs["record_listing"] is True
@pytest.mark.asyncio
@pytest.mark.parametrize(("listing_kwargs", "recorded"), [({}, False), ({"record_listing": True}, True)])
async def test_list_mcp_tools_records_the_catalog_only_when_asked(
listing_kwargs: dict[str, bool], recorded: bool
) -> None:
"""The aggregate listing fills the caller's listed-tools slot only when asked: a listing an internal
caller never serves must not hand a later tools/call a description the caller never saw."""
manager = operations.global_mcp_server_manager
server = MCPServer(server_id="listing-slot", name="listing-slot", transport=MCPTransport.http, url="http://slot")
user = UserAPIKeyAuth(api_key="sk-listing-slot", user_id="lister")
upstream = [MCPTool(name="echo", description="Echo text back", inputSchema={"type": "object"})]
with (
patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])),
patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())),
patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)),
patch.dict(manager.tool_name_to_mcp_server_name_mapping),
):
try:
listing = await operations._list_mcp_tools(user_api_key_auth=user, **listing_kwargs)
listed = manager.get_listed_tool(server, "echo", ListedToolsCaller(user_api_key_auth=user))
finally:
manager._drop_listed_tools(server.server_id)
assert [tool.name for tool in listing.tools] == ["listing-slot-echo"]
assert (listed is not None) is recorded

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

@ -13,6 +13,7 @@ from typing import Any, Dict, Final, List, Optional
import pytest
from fastapi import HTTPException
from pydantic import TypeAdapter
import litellm
from litellm import Router
@ -21,6 +22,7 @@ from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
PARALLEL_REQUEST_SLOT_TTL_SECONDS,
ParallelSlotAcquisition,
@ -39,6 +41,7 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import (
from litellm.proxy.utils import InternalUsageCache, ProxyLogging, hash_token
from litellm.types.caching import RedisPipelineIncrementOperation
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.mcp import MCPPreCallRequestObject
from litellm.types.utils import (
EmbeddingResponse,
ModelResponse,
@ -108,6 +111,159 @@ def test_api_key_descriptor_applies_budget_throttle(
assert api_key_descriptor["rate_limit"]["tokens_per_unit"] == expected_tpm
@pytest.mark.asyncio
@pytest.mark.parametrize(
"description", [None, "Gateway metadata, not caller input. " * 100], ids=["unlisted", "listed"]
)
@pytest.mark.parametrize("arguments_rewritten", [False, True])
async def test_mcp_description_does_not_change_admission_or_reserved_tokens(
description: str | None, arguments_rewritten: bool, monkeypatch: pytest.MonkeyPatch
) -> None:
cache: Final = DualCache()
handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache))
logger: Final = ProxyLogging(user_api_key_cache=UserApiKeyCache())
schema: Final = {"type": "object", "properties": {"q": {"type": "string", "description": "Schema text " * 100}}}
request: Final = MCPPreCallRequestObject(
tool_name="echo", arguments={"q": "hello"}, tool_description=description, tool_input_schema=schema
)
data: Final = TypeAdapter(dict[str, object]).validate_python(logger._convert_mcp_to_llm_format(request, {}))
messages: Final = data["messages"]
caller: Final = UserAPIKeyAuth(api_key=hash_token("sk-mcp-description-reservation"), tpm_limit=64)
if arguments_rewritten:
data["mcp_arguments"] = {"q": "Transformed arguments " * 100}
monkeypatch.setattr(litellm, "callbacks", [handler])
await logger.pre_call_hook(user_api_key_dict=caller, data=data, call_type="call_mcp_tool")
stash: Final = get_request_stash()
assert stash is not None
assert stash.reserved_tokens == 25
assert (
await cache.async_get_cache(
key=handler.create_rate_limit_keys("api_key", caller.api_key, "tokens"), local_only=True
)
== 25
)
assert data["messages"] is messages
assert data.get("mcp_tool_description") == description
assert data["mcp_input_schema"] == schema
assert messages == [
{
"role": "user",
"content": "Tool: echo\nArguments: {'q': 'hello'}",
}
]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"description", [None, "Gateway metadata, not caller input. " * 100], ids=["unlisted", "listed"]
)
@pytest.mark.parametrize("itpm_limit,otpm_limit", [(64, 4096), (4096, 64), (4096, 4096)])
@pytest.mark.parametrize("arguments_rewritten", [False, True])
async def test_mcp_description_preserves_project_input_and_output_reservations(
description: str | None, itpm_limit: int, otpm_limit: int,
arguments_rewritten: bool, monkeypatch: pytest.MonkeyPatch
) -> None:
cache: Final = DualCache()
handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache))
logger: Final = ProxyLogging(user_api_key_cache=UserApiKeyCache())
schema: Final = {"type": "object", "properties": {"q": {"type": "string", "description": "Schema text " * 100}}}
request: Final = MCPPreCallRequestObject(
tool_name="echo", arguments={"q": "hello"}, tool_description=description, tool_input_schema=schema
)
data: Final = TypeAdapter(dict[str, object]).validate_python(logger._convert_mcp_to_llm_format(request, {}))
messages: Final = data["messages"]
base_data: Final[dict[str, object]] = {
"messages": [{"role": "user", "content": "Tool: echo\nArguments: {'q': 'hello'}"}]
}
expected_input: Final = handler._estimate_precise_input_tokens(base_data, "mcp-tool-call", "call_mcp_tool")
expected_output: Final = handler.no_max_tokens_output_floor(otpm_limit)
expected_combined: Final = handler._estimate_tokens_for_request(
base_data, min_configured_tpm_limit=4096, call_type="call_mcp_tool"
)
caller: Final = UserAPIKeyAuth(
api_key=hash_token("sk-mcp-project-reservation"),
tpm_limit=4096,
project_id="mcp-project-reservation",
project_metadata={
"model_itpm_limit": {"mcp-tool-call": itpm_limit},
"model_otpm_limit": {"mcp-tool-call": otpm_limit},
},
)
if arguments_rewritten:
data["mcp_arguments"] = {"q": "Transformed arguments " * 100}
monkeypatch.setattr(litellm, "callbacks", [handler])
await logger.pre_call_hook(user_api_key_dict=caller, data=data, call_type="call_mcp_tool")
stash: Final = get_request_stash()
assert stash is not None
assert (stash.reserved_tokens, stash.itpm_reserved_tokens, stash.otpm_reserved_tokens) == (
expected_combined,
expected_input,
expected_output,
)
assert (
await cache.async_get_cache(
key=handler.create_rate_limit_keys(
"model_per_project_itpm", f"{caller.project_id}:mcp-tool-call", "tokens"
),
local_only=True,
)
== expected_input
)
assert (
await cache.async_get_cache(
key=handler.create_rate_limit_keys(
"model_per_project_otpm", f"{caller.project_id}:mcp-tool-call", "tokens"
),
local_only=True,
)
== expected_output
)
assert data["messages"] is messages
assert data.get("mcp_tool_description") == description
assert data["mcp_input_schema"] == schema
assert messages == [
{
"role": "user",
"content": "Tool: echo\nArguments: {'q': 'hello'}",
}
]
def test_llm_tpm_estimation_still_counts_messages_with_mcp_metadata() -> None:
handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache()))
data: Final[dict[str, object]] = {
"messages": [{"role": "user", "content": "x" * 400}],
"max_tokens": 1,
"mcp_tool_name": "echo",
"mcp_arguments": {},
}
assert handler._estimate_tokens_for_request(data, call_type="acompletion") == 101
@pytest.mark.asyncio
async def test_unconverted_mcp_request_keeps_its_reservation() -> None:
cache: Final = DualCache()
handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache))
caller: Final = UserAPIKeyAuth(api_key=hash_token("sk-raw-mcp-request"), tpm_limit=64)
data: Final[dict[str, object]] = {"name": "echo", "arguments": {"q": "hello"}, "server_id": "fixture"}
await handler.async_pre_call_hook(user_api_key_dict=caller, cache=cache, data=data, call_type="call_mcp_tool")
stash: Final = get_request_stash()
assert stash is not None
assert stash.reserved_tokens == 16
assert (
await cache.async_get_cache(
key=handler.create_rate_limit_keys("api_key", caller.api_key, "tokens"), local_only=True
)
== 16
)
@pytest.mark.flaky(reruns=3)
@pytest.mark.asyncio
async def test_sliding_window_rate_limit_v3(monkeypatch, time_controller):

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 = {

View file

@ -1,10 +1,11 @@
import asyncio
import importlib
import subprocess
import sys
import textwrap
import types
from typing import Any, Final, cast
from unittest.mock import AsyncMock, MagicMock
from typing import Any, Final, Literal, cast
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
@ -12,15 +13,26 @@ from mcp.types import CallToolResult, TextContent
from mcp.types import Tool as MCPTool
from openai.types.responses.tool_param import Mcp
import litellm
from litellm.caching.caching import DualCache
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.proxy._experimental.mcp_server import operations as mcp_operations
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller, MCPServerManager
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.utils import ProxyLogging
from litellm.responses import main as responses_main
from litellm.responses.mcp import litellm_proxy_mcp_handler as mcp_handler_module
from litellm.responses.mcp.litellm_proxy_mcp_handler import (
LiteLLM_Proxy_MCP_Handler,
)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.mcp import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
from litellm.types.responses.main import OutputFunctionToolCall
from litellm.types.utils import ModelResponse
from litellm.types.utils import GenericGuardrailAPIInputs, ModelResponse
class _DummyMCPResult:
@ -1310,6 +1322,167 @@ async def test_responses_discovery_logs_sanitized_caller_headers(monkeypatch: py
assert headers["x-mcp-deepwiki-authorization"] == "upstream-sentinel"
@pytest.mark.asyncio
@pytest.mark.parametrize("real_listing", [False, True])
@pytest.mark.parametrize(
("allowed_tools", "expected_names"),
[
([], ["responses_slot-echo", "responses_slot-status"]),
(["echo"], ["responses_slot-echo"]),
(["responses_slot-echo"], ["responses_slot-echo"]),
(["absent"], []),
],
)
async def test_bridge_listing_leaves_the_callers_catalog_unchanged(
monkeypatch: pytest.MonkeyPatch, allowed_tools: list[str], expected_names: list[str], real_listing: bool
) -> None:
manager: Final = mcp_operations.global_mcp_server_manager
server: Final = MCPServer(
server_id="responses-slot", name="responses_slot", alias="responses_slot", transport=MCPTransport.http
)
user: Final = UserAPIKeyAuth(api_key="sk-responses-slot", user_id="responder")
upstream: Final = [
MCPTool(name="echo", description="Echo text back", inputSchema={"type": "object"}),
MCPTool(name="status", description="Report status", inputSchema={"type": "object"}),
MCPTool(name="echo", description="Duplicate echo", inputSchema={"type": "object", "properties": {}}),
]
fake_manager: Final = types.SimpleNamespace(
get_registry=MagicMock(return_value={}),
get_allowed_mcp_servers=AsyncMock(return_value=[]),
get_mcp_servers_from_ids=MagicMock(return_value=[]),
get_mcp_server_by_name=MagicMock(return_value=None),
)
monkeypatch.setattr(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
fake_manager,
)
with (
patch.dict(manager.tool_name_to_mcp_server_name_mapping),
patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])),
patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())),
patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)),
):
try:
if real_listing:
await manager._get_tools_from_server(server, user_api_key_auth=user, record_listing=True)
caller: Final = ListedToolsCaller(user_api_key_auth=user)
before: Final = {
tool.name: (listed.description, listed.input_schema)
for tool in upstream
if (listed := manager.get_listed_tool(server, tool.name, caller)) is not None
}
assert bool(before) is real_listing
tools, _server_names = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform(
user_api_key_auth=user,
mcp_tools_with_litellm_proxy=[
{
"type": "mcp",
"server_url": "litellm_proxy/mcp/responses-slot",
"allowed_tools": allowed_tools,
}
],
)
recorded: Final = {
tool.name: (listed.description, listed.input_schema)
for tool in upstream
if (listed := manager.get_listed_tool(server, tool.name, caller)) is not None
}
assert recorded == before
assert (
manager.get_listed_tool(
server, "echo", ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-other-caller"))
)
is None
)
finally:
manager._drop_listed_tools(server.server_id)
assert [tool.name for tool in tools] == expected_names
class _BridgeMetadataGuardrail(CustomGuardrail):
def __init__(self) -> None:
super().__init__(guardrail_name="bridge-metadata", event_hook=GuardrailEventHooks.pre_mcp_call, default_on=True)
self.calls: tuple[tuple[object, object], ...] = ()
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict[str, object],
input_type: Literal["request", "response"],
logging_obj: Logging | None = None,
) -> GenericGuardrailAPIInputs:
if request_data.get("mcp_arguments") == {"probe": "bridge"}:
self.calls += ((request_data.get("mcp_tool_description"), request_data.get("mcp_input_schema")),)
return inputs
@pytest.mark.asyncio
async def test_concurrent_bridge_calls_use_their_own_served_metadata(monkeypatch: pytest.MonkeyPatch) -> None:
manager: Final = MCPServerManager()
server: Final = MCPServer(server_id="bridge", name="bridge", transport=MCPTransport.http, url="http://upstream")
manager.registry = {server.server_id: server}
user: Final = UserAPIKeyAuth(api_key="sk-bridge", user_id="bridge-user")
upstream: Final = [
MCPTool(
name="echo",
description="Echo text",
inputSchema={"type": "object", "properties": {"text": {"type": "string"}}},
),
MCPTool(name="status", description="Read status", inputSchema={"type": "object"}),
]
client: Final = AsyncMock()
client.call_tool.return_value = CallToolResult(content=[TextContent(type="text", text="ok")])
manager._create_mcp_client = AsyncMock(return_value=client)
manager._fetch_tools_with_timeout = AsyncMock(return_value=upstream)
guardrail: Final = _BridgeMetadataGuardrail()
logger: Final = ProxyLogging(user_api_key_cache=DualCache())
monkeypatch.setattr(litellm, "callbacks", [guardrail])
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", logger)
monkeypatch.setattr(mcp_operations, "global_mcp_server_manager", manager)
monkeypatch.setattr(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server]))
monkeypatch.setattr("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", manager)
first_listed: Final = asyncio.Event()
second_listed: Final = asyncio.Event()
async def bridge(name: str, first: bool) -> None:
if not first:
await first_listed.wait()
tools, server_map = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform(
user_api_key_auth=user,
mcp_tools_with_litellm_proxy=[
{"type": "mcp", "server_url": "litellm_proxy/mcp/bridge", "allowed_tools": [name]}
],
)
(first_listed if first else second_listed).set()
await second_listed.wait()
result: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
tool_server_map=server_map,
tool_calls=[
{"type": "function_call", "name": f"bridge-{name}", "arguments": '{"probe":"bridge"}', "call_id": name}
],
user_api_key_auth=user,
served_tools=tools,
)
assert [entry["result"] for entry in result] == ["ok"]
try:
await asyncio.gather(bridge("echo", True), bridge("status", False))
assert sorted(guardrail.calls, key=str) == sorted(
((tool.description, tool.input_schema) for tool in upstream), key=str
)
await manager.call_tool("bridge", "echo", {"probe": "bridge"}, user_api_key_auth=user, proxy_logging_obj=logger)
assert guardrail.calls[-1] == (None, None)
await manager._get_tools_from_server(
server, user_api_key_auth=user, proxy_logging_obj=logger, record_listing=True
)
await manager.call_tool("bridge", "echo", {"probe": "bridge"}, user_api_key_auth=user, proxy_logging_obj=logger)
assert guardrail.calls[-1] == (upstream[0].description, upstream[0].input_schema)
finally:
manager._drop_listed_tools(server.server_id)
ProxyLogging._callback_capabilities_cache.clear()
def _toolset_gateway_manager(toolset_id: str, server_id: str) -> types.SimpleNamespace:
return types.SimpleNamespace(
get_registry=MagicMock(return_value={}),