diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index a247ce64bf5..6341f7972f0 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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 ) ######################################################### diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 5a63fedce73..891c37f4fe7 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -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 diff --git a/tests/integration/mcp/test_mcp_listed_tool_metadata.py b/tests/integration/mcp/test_mcp_listed_tool_metadata.py index 0381e8abb00..92cf56e9899 100644 --- a/tests/integration/mcp/test_mcp_listed_tool_metadata.py +++ b/tests/integration/mcp/test_mcp_listed_tool_metadata.py @@ -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" + ) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 57268902614..070d87eb012 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -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