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>
This commit is contained in:
yucheng 2026-09-30 09:08:00 +00:00
parent 2e1c6bfc70
commit 7dee160b9f
4 changed files with 143 additions and 31 deletions

View file

@ -1203,6 +1203,23 @@ def _server_auth_header_for(
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 _format_byok_openapi_auth_header(mcp_server: MCPServer, mcp_auth_header: str) -> str:
"""Format a raw BYOK credential for OpenAPI tool ``Authorization`` injection.
@ -4638,15 +4655,15 @@ class MCPServerManager:
invalidate_oauth_metadata_cache(server_id)
def _listed_tools_identity(self, server: MCPServer, caller: ListedToolsCaller | None) -> str | None:
"""Key the listed-tool cache by every request input that can change the upstream catalog.
"""Key the listed-tool cache by every request input that can change the served catalog.
Forwarded headers, header-driven stdio env, the caller bearer (forwarded as-is or
exchanged as the OBO subject), the server-specific auth header, and the per-caller JWT
MCPJWTSigner mints for tools/list all reach upstream, so two callers differing in any of
them may be shown different tools. Shared servers with none of those stay on the shared
(``None``) slot. OpenAPI servers list from the process-wide registry.
The catalog is guardrail-shaped for the caller's own key (default-on guardrails, key or team
selections and opt-outs), so every keyed caller gets its own slot, on OpenAPI servers too.
Forwarded headers, header-driven stdio env, the caller bearer (forwarded as-is or exchanged 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 server.spec_path or caller is None:
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
@ -4664,20 +4681,10 @@ class MCPServerManager:
forwarded,
stdio_env,
caller_bearer,
per_caller=self._signs_caller_identity_upstream(server),
per_caller=auth is not None,
)
return digest
@staticmethod
def _signs_caller_identity_upstream(server: MCPServer) -> bool:
from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import ( # noqa: PLC0415 # lazy: guardrail package imports the proxy server
get_mcp_jwt_signer,
)
if get_mcp_jwt_signer() is None:
return False
return server.static_headers is None or not any(k.lower() == "authorization" for k in server.static_headers)
@staticmethod
def _forwarded_header_values(
server: MCPServer, raw_headers: Mapping[str, str] | None
@ -6633,11 +6640,8 @@ class MCPServerManager:
user_api_key_auth,
mcp_auth_header,
)
listed_caller: Final = ListedToolsCaller(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=_server_auth_header_for(mcp_server, mcp_server_auth_headers, mcp_auth_header),
raw_headers=raw_headers,
oauth2_headers=oauth2_headers,
listed_caller: Final = listed_tools_caller_for(
mcp_server, user_api_key_auth, mcp_auth_header, mcp_server_auth_headers, raw_headers, oauth2_headers
)
#########################################################

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,
@ -1607,10 +1609,12 @@ async def _list_mcp_resource_templates(
return managed_resource_templates
def _registered_tool_metadata(name: str, registered: RegisteredTool, server: MCPServer) -> MCPTool:
"""The tool as ``tools/list`` served it (pinned, overridden, guardrail-masked) when a listing was
recorded for ``server``, else the registry entry with the admin description override applied."""
listed: Final = global_mcp_server_manager.get_listed_tool(server, name)
def _registered_tool_metadata(
name: str, registered: RegisteredTool, server: MCPServer, caller: ListedToolsCaller
) -> MCPTool:
"""The tool as ``tools/list`` served it to this caller (pinned, overridden, guardrail-masked) when a
listing was recorded, else the registry entry with the admin description override applied."""
listed: Final = global_mcp_server_manager.get_listed_tool(server, name, caller)
if listed is not None:
return listed
overrides: Final = server.tool_name_to_description
@ -2087,7 +2091,14 @@ async def _execute_mcp_tool(
raw_headers=raw_headers,
litellm_logging_obj=litellm_logging_obj,
guardrail_context=guardrail_context,
tool=_registered_tool_metadata(original_tool_name, local_tool, mcp_server),
tool=_registered_tool_metadata(
original_tool_name,
local_tool,
mcp_server,
listed_tools_caller_for(
mcp_server, user_api_key_auth, mcp_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.
@ -2199,7 +2210,19 @@ async def _execute_mcp_tool(
raw_headers=raw_headers,
litellm_logging_obj=litellm_logging_obj,
guardrail_context=guardrail_context,
tool=_registered_tool_metadata(original_tool_name, registered_local_tool, prefix_server),
tool=_registered_tool_metadata(
original_tool_name,
registered_local_tool,
prefix_server,
listed_tools_caller_for(
prefix_server,
user_api_key_auth,
mcp_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

View file

@ -167,3 +167,28 @@ def test_openapi_call_is_evaluated_against_the_masked_override_the_listing_serve
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"
)

View file

@ -30,6 +30,7 @@ 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,
@ -8179,8 +8180,11 @@ async def test_execute_mcp_tool_hands_openapi_hooks_the_guarded_catalog_entry_cl
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)], None
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={})
@ -8194,7 +8198,7 @@ async def test_execute_mcp_tool_hands_openapi_hooks_the_guarded_catalog_entry_cl
arguments={"petId": 1},
allowed_mcp_servers=[petstore],
start_time=datetime.now(),
user_api_key_auth=UserAPIKeyAuth(api_key="sk-user", user_id="alice"),
user_api_key_auth=alice,
)
finally:
mcp_module.global_mcp_tool_registry.unregister_tools_with_prefix("petstore-")
@ -8206,6 +8210,62 @@ async def test_execute_mcp_tool_hands_openapi_hooks_the_guarded_catalog_entry_cl
)
@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_hands_hooks_the_metadata_of_the_operation_it_runs_when_names_collide():
"""An OpenAPI operation whose name starts with its own server prefix must not be reported to the