mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
2e1c6bfc70
commit
7dee160b9f
4 changed files with 143 additions and 31 deletions
|
|
@ -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
|
||||
)
|
||||
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue