diff --git a/tests/tools/test_mcp_tool.py b/tests/tools/test_mcp_tool.py index 0f5cfc43ba3e..a2cf61e04a56 100644 --- a/tests/tools/test_mcp_tool.py +++ b/tests/tools/test_mcp_tool.py @@ -6,6 +6,7 @@ All tests use mocks -- no real MCP servers or subprocesses are started. import asyncio import concurrent.futures import json +import logging import os import sys import threading @@ -125,6 +126,71 @@ class TestLoadMCPConfig: assert result == {} + +class TestMCPParallelSafetyProvenance: + def test_parallel_safe_servers_keep_exact_raw_names(self, monkeypatch): + import tools.mcp_tool as mcp_tool + + first = SimpleNamespace(session=object(), _registered_tool_names=[]) + second = SimpleNamespace(session=object(), _registered_tool_names=[]) + + with mcp_tool._lock: + saved_servers = dict(mcp_tool._servers) + saved_parallel = set(mcp_tool._parallel_safe_servers) + mcp_tool._servers.clear() + mcp_tool._servers.update({"foo-bar": first, "foo_bar": second}) + mcp_tool._parallel_safe_servers.clear() + + try: + monkeypatch.setattr(mcp_tool, "_MCP_AVAILABLE", True) + monkeypatch.setattr( + mcp_tool, "_filter_suspicious_mcp_servers", lambda servers: servers + ) + mcp_tool.register_mcp_servers( + { + "foo-bar": {"supports_parallel_tool_calls": True}, + "foo_bar": {"supports_parallel_tool_calls": False}, + } + ) + with mcp_tool._lock: + assert "foo-bar" in mcp_tool._parallel_safe_servers + assert "foo_bar" not in mcp_tool._parallel_safe_servers + finally: + with mcp_tool._lock: + mcp_tool._servers.clear() + mcp_tool._servers.update(saved_servers) + mcp_tool._parallel_safe_servers.clear() + mcp_tool._parallel_safe_servers.update(saved_parallel) + + def test_tool_provenance_keeps_exact_raw_server_names(self): + import tools.mcp_tool as mcp_tool + + first_tool = "mcp__foo_bar__first" + second_tool = "mcp__foo_bar__second" + with mcp_tool._lock: + saved_map = dict(mcp_tool._mcp_tool_server_names) + saved_parallel = set(mcp_tool._parallel_safe_servers) + mcp_tool._mcp_tool_server_names.clear() + mcp_tool._parallel_safe_servers.clear() + mcp_tool._parallel_safe_servers.add("foo-bar") + + try: + mcp_tool._track_mcp_tool_server(first_tool, "foo-bar") + mcp_tool._track_mcp_tool_server(second_tool, "foo_bar") + + assert mcp_tool.is_mcp_tool_parallel_safe(first_tool) is True + assert mcp_tool.is_mcp_tool_parallel_safe(second_tool) is False + assert mcp_tool.get_registered_mcp_server_names() == { + "foo-bar", + "foo_bar", + } + finally: + with mcp_tool._lock: + mcp_tool._mcp_tool_server_names.clear() + mcp_tool._mcp_tool_server_names.update(saved_map) + mcp_tool._parallel_safe_servers.clear() + mcp_tool._parallel_safe_servers.update(saved_parallel) + class TestMCPStatus: def test_status_distinguishes_configured_connecting_failed_and_disabled( self, monkeypatch @@ -1087,6 +1153,103 @@ class TestDiscoverAndRegister: _servers.pop("srv", None) + def test_same_server_normalization_collision_skips_all_ambiguous_tools(self, caplog): + from tools.mcp_tool import _register_server_tools + from tools.registry import ToolRegistry + + registry = ToolRegistry() + server = _make_mock_server( + "srv", + session=MagicMock(), + tools=[ + _make_mcp_tool("read-file"), + _make_mcp_tool("read_file"), + _make_mcp_tool("safe_tool"), + ], + ) + config = {"tools": {"resources": False, "prompts": False}} + + with patch("tools.registry.registry", registry), \ + patch("tools.mcp_tool._track_mcp_tool_server"), \ + caplog.at_level(logging.ERROR, logger="tools.mcp_tool"): + registered = _register_server_tools("srv", server, config) + + assert registered == ["mcp__srv__safe_tool"] + assert registry.get_entry("mcp__srv__read_file") is None + assert registry.get_entry("mcp__srv__safe_tool") is not None + assert any( + "name normalization collision" in record.message + and "tool 'read-file'" in record.message + and "tool 'read_file'" in record.message + for record in caplog.records + ) + + def test_cross_server_normalization_collision_preserves_first_owner(self, caplog): + from tools.mcp_tool import _register_server_tools + from tools.registry import ToolRegistry + + registry = ToolRegistry() + first = _make_mock_server( + "foo-bar", + session=MagicMock(), + tools=[_make_mcp_tool("search")], + ) + second = _make_mock_server( + "foo_bar", + session=MagicMock(), + tools=[_make_mcp_tool("search")], + ) + config = {"tools": {"resources": False, "prompts": False}} + + with patch("tools.registry.registry", registry), \ + patch("tools.mcp_tool._track_mcp_tool_server"), \ + caplog.at_level(logging.ERROR, logger="tools.mcp_tool"): + first_registered = _register_server_tools("foo-bar", first, config) + second_registered = _register_server_tools("foo_bar", second, config) + + assert first_registered == ["mcp__foo_bar__search"] + assert second_registered == [] + entry = registry.get_entry("mcp__foo_bar__search") + assert entry is not None + assert entry.toolset == "mcp-foo-bar" + assert any( + "already owned by MCP toolset 'mcp-foo-bar'" in record.message + for record in caplog.records + ) + + def test_raw_tool_collision_with_generated_utility_skips_both(self, caplog): + from tools.mcp_tool import _register_server_tools + from tools.registry import ToolRegistry + + registry = ToolRegistry() + server = _make_mock_server( + "srv", + session=MagicMock(), + tools=[ + _make_mcp_tool("list_resources"), + _make_mcp_tool("safe_tool"), + ], + ) + server.initialize_result = SimpleNamespace( + capabilities=SimpleNamespace(resources=object(), prompts=None) + ) + config = {"tools": {"prompts": False}} + + with patch("tools.registry.registry", registry), \ + patch("tools.mcp_tool._track_mcp_tool_server"), \ + caplog.at_level(logging.ERROR, logger="tools.mcp_tool"): + registered = _register_server_tools("srv", server, config) + + assert "mcp__srv__list_resources" not in registered + assert registry.get_entry("mcp__srv__list_resources") is None + assert "mcp__srv__safe_tool" in registered + assert "mcp__srv__read_resource" in registered + assert any( + "tool 'list_resources'" in record.message + and "generated utility 'list_resources'" in record.message + for record in caplog.records + ) + # --------------------------------------------------------------------------- # MCPServerTask (run / start / shutdown) # --------------------------------------------------------------------------- @@ -1195,6 +1358,49 @@ class TestMCPServerTask: "mcp__srv__get_prompt", } + def test_refresh_removes_old_tool_when_new_list_becomes_ambiguous(self, caplog): + """A newly ambiguous list must not leave the old handler callable.""" + from tools.mcp_tool import MCPServerTask + from tools.registry import ToolRegistry + + registry = ToolRegistry() + server = MCPServerTask("srv") + server._config = {"tools": {"resources": False, "prompts": False}} + server._tools = [_make_mcp_tool("read_file")] + server._registered_tool_names = ["mcp__srv__read_file"] + server.session = MagicMock() + server.session.list_tools = AsyncMock( + return_value=SimpleNamespace( + tools=[ + _make_mcp_tool("read_file"), + _make_mcp_tool("read-file"), + ] + ) + ) + registry.register( + name="mcp__srv__read_file", + toolset="mcp-srv", + schema={ + "name": "mcp__srv__read_file", + "description": "Old", + "parameters": {"type": "object", "properties": {}}, + }, + handler=lambda *_args, **_kwargs: "{}", + ) + + with patch("tools.registry.registry", registry), \ + patch("tools.mcp_tool._track_mcp_tool_server"), \ + patch("tools.mcp_tool._forget_mcp_tool_server"), \ + caplog.at_level(logging.ERROR, logger="tools.mcp_tool"): + asyncio.run(server._refresh_tools()) + + assert registry.get_entry("mcp__srv__read_file") is None + assert server._registered_tool_names == [] + assert any( + "name normalization collision" in record.message + for record in caplog.records + ) + def test_schedule_tools_refresh_keeps_task_until_done(self): """Background refresh tasks are strongly referenced and then discarded.""" from tools.mcp_tool import MCPServerTask @@ -4197,8 +4403,8 @@ class TestMCPBuiltinCollisionGuard: _servers.pop("minimax", None) - def test_mcp_tool_allowed_when_collision_is_another_mcp(self): - """Collision between two MCP toolsets is allowed (last wins).""" + def test_mcp_tool_rejected_when_collision_is_another_mcp(self): + """Cross-server MCP collisions preserve the existing owner.""" from tools.registry import ToolRegistry from tools.mcp_tool import _discover_and_register_server, _servers, MCPServerTask @@ -4230,9 +4436,13 @@ class TestMCPBuiltinCollisionGuard: _discover_and_register_server("srv", {"command": "test", "args": []}) ) - # MCP-to-MCP collision is allowed — the new server wins. - assert "mcp__srv__do_thing" in registered - assert mock_registry.get_toolset_for_tool("mcp__srv__do_thing") == "mcp-srv" + # Cross-server MCP collisions fail closed: the existing owner stays active. + assert "mcp__srv__do_thing" not in registered + entry = mock_registry.get_entry("mcp__srv__do_thing") + assert entry is not None + assert entry.toolset == "mcp-old" + assert entry.schema["description"] == "From another MCP server" + assert mock_registry.get_toolset_for_tool("mcp__srv__do_thing") == "mcp-old" _servers.pop("srv", None) diff --git a/tests/tools/test_registry.py b/tests/tools/test_registry.py index 97415c991e29..b355f22e41bb 100644 --- a/tests/tools/test_registry.py +++ b/tests/tools/test_registry.py @@ -1,6 +1,7 @@ """Tests for the central tool registry.""" import json +import logging import threading from pathlib import Path from unittest.mock import patch @@ -112,6 +113,57 @@ class TestRegisterAndDispatch: assert result["result_type"] == "NoneType" + def test_cross_mcp_toolsets_do_not_overwrite_atomically(self, caplog): + """Parallel MCP registrations with one name leave exactly one owner.""" + reg = ToolRegistry() + barrier = threading.Barrier(3) + errors = [] + + def _register(toolset, owner): + try: + barrier.wait(timeout=5) + + def _handler(args, **kwargs): + return json.dumps({"owner": owner}) + + reg.register( + name="mcp__foo_bar__search", + toolset=toolset, + schema=_make_schema("mcp__foo_bar__search"), + handler=_handler, + ) + except BaseException as exc: # pragma: no cover - asserted below + errors.append(exc) + + threads = [ + threading.Thread(target=_register, args=("mcp-foo-bar", "dash")), + threading.Thread(target=_register, args=("mcp-foo_bar", "underscore")), + ] + + with caplog.at_level(logging.ERROR, logger="tools.registry"): + for thread in threads: + thread.start() + barrier.wait(timeout=5) + for thread in threads: + thread.join(timeout=10) + + assert all(not thread.is_alive() for thread in threads) + assert errors == [] + assert reg._generation == 1 + + entry = reg.get_entry("mcp__foo_bar__search") + assert entry is not None + assert entry.toolset in {"mcp-foo-bar", "mcp-foo_bar"} + assert json.loads(reg.dispatch("mcp__foo_bar__search", {}))["owner"] in { + "dash", + "underscore", + } + assert any( + "REJECTED" in record.message + and "mcp__foo_bar__search" in record.message + for record in caplog.records + ) + class TestGetDefinitions: def test_returns_openai_format(self): reg = ToolRegistry() diff --git a/tools/mcp_tool.py b/tools/mcp_tool.py index 5398700cc4ac..0ae26b4d3b4b 100644 --- a/tools/mcp_tool.py +++ b/tools/mcp_tool.py @@ -2105,21 +2105,39 @@ class MCPServerTask: # "tool not connected" / stale-handler races during startup # notifications. Tools absent from the fresh list are no longer # callable, so remove only those stale registry entries first. + toolset_name = f"mcp-{self.name}" stale_tool_names = old_tool_names - { mcp_prefixed_tool_name(self.name, tool.name) for tool in new_mcp_tools } for tool_name in stale_tool_names: + # Never let one server's refresh remove a colliding name that + # is currently owned by another server. + if registry.get_toolset_for_tool(tool_name) != toolset_name: + continue registry.deregister(tool_name) _forget_mcp_tool_server(tool_name) - # 3. Re-register with fresh tool list + # 3. Re-register with the fresh list. The helper may skip names that + # are ambiguous after normalization. self._tools = new_mcp_tools - self._registered_tool_names = _register_server_tools( + registered_names = _register_server_tools( self.name, self, self._config ) - # 5. Log what changed (user-visible notification) + # A previously unique raw name can become ambiguous without changing + # its normalized registry name. In that case the pre-pass above does + # not consider it stale, so remove any old entry that the final, + # collision-checked registration set no longer owns. + registered_name_set = set(registered_names) + for tool_name in old_tool_names - registered_name_set: + if registry.get_toolset_for_tool(tool_name) != toolset_name: + continue + registry.deregister(tool_name) + _forget_mcp_tool_server(tool_name) + self._registered_tool_names = registered_names + + # 4. Log what changed (user-visible notification) new_tool_names = set(self._registered_tool_names) added = new_tool_names - old_tool_names removed = old_tool_names - new_tool_names @@ -4075,17 +4093,15 @@ def _handle_session_expired_and_retry( return None -# Sanitized server names whose ``supports_parallel_tool_calls`` config is True. -# Populated during ``register_mcp_servers()`` and queried by -# ``is_mcp_tool_parallel_safe()`` for the parallel-execution check in run_agent. +# Exact raw server names whose ``supports_parallel_tool_calls`` config is True. +# Raw identity matters: distinct names such as ``foo-bar`` and ``foo_bar`` both +# sanitize to ``foo_bar`` but must not share policy. _parallel_safe_servers: set = set() -# Exact MCP tool-name provenance. MCP tool names are formatted as -# ``mcp_{sanitized_server}_{sanitized_tool}``, which is ambiguous when server -# names contain underscores (``mcp_a_b_tool`` could be server ``a`` + tool -# ``b_tool`` or server ``a_b`` + tool ``tool``). Keep the server component -# captured at registration time so parallel safety never relies on prefix -# guessing. +# Exact MCP tool-name provenance. The generated registry name is lossy because +# provider-safe normalization maps punctuation to ``_``. Keep the raw server +# name captured at registration time so policy and capability checks never rely +# on parsing or re-sanitizing the generated name. _mcp_tool_server_names: Dict[str, str] = {} # Dedicated event loop running in a background daemon thread. @@ -5527,10 +5543,9 @@ _UTILITY_CAPABILITY_ATTRS = { def _track_mcp_tool_server(tool_name: str, server_name: str) -> None: - """Remember the exact MCP server that registered *tool_name*.""" - safe_server_name = sanitize_mcp_name_component(server_name) + """Remember the exact raw MCP server that registered *tool_name*.""" with _lock: - _mcp_tool_server_names[tool_name] = safe_server_name + _mcp_tool_server_names[tool_name] = server_name def _forget_mcp_tool_server(tool_name: str) -> None: @@ -5616,6 +5631,11 @@ def _register_server_tools(name: str, server: MCPServerTask, config: dict) -> Li for ``mcp-{server}`` and raw server-name aliases is derived from the live registry, rather than mutating ``toolsets.TOOLSETS`` at runtime. + Lossy provider-safe name normalization can map distinct raw names to the + same registry name (for example ``read-file`` and ``read_file``). Such + collisions fail closed: every ambiguous entry is skipped rather than + selecting an arbitrary handler. + Used by both initial discovery and dynamic refresh (list_changed). Returns: @@ -5634,8 +5654,12 @@ def _register_server_tools(name: str, server: MCPServerTask, config: dict) -> Li # include takes precedence over exclude # Neither set → register all tools (backward-compatible default) tools_filter = config.get("tools") or {} - include_set = _normalize_name_filter(tools_filter.get("include"), f"mcp_servers.{name}.tools.include") - exclude_set = _normalize_name_filter(tools_filter.get("exclude"), f"mcp_servers.{name}.tools.exclude") + include_set = _normalize_name_filter( + tools_filter.get("include"), f"mcp_servers.{name}.tools.include" + ) + exclude_set = _normalize_name_filter( + tools_filter.get("exclude"), f"mcp_servers.{name}.tools.exclude" + ) def _should_register(tool_name: str) -> bool: if include_set: @@ -5644,82 +5668,149 @@ def _register_server_tools(name: str, server: MCPServerTask, config: dict) -> Li return not matches_name_filter(tool_name, exclude_set) return True + check_fn = _make_check_fn(name) + candidates: List[dict] = [] + for mcp_tool in server._tools: if not _should_register(mcp_tool.name): - logger.debug("MCP server '%s': skipping tool '%s' (filtered by config)", name, mcp_tool.name) - continue - - # Scan tool description for prompt injection patterns - _scan_mcp_description(name, mcp_tool.name, mcp_tool.description or "") - - schema = _convert_mcp_schema(name, mcp_tool) - tool_name_prefixed = schema["name"] - - # Guard against collisions with built-in (non-MCP) tools. - existing_toolset = registry.get_toolset_for_tool(tool_name_prefixed) - if existing_toolset and not existing_toolset.startswith("mcp-"): - logger.warning( - "MCP server '%s': tool '%s' (→ '%s') collides with built-in " - "tool in toolset '%s' — skipping to preserve built-in", - name, mcp_tool.name, tool_name_prefixed, existing_toolset, + logger.debug( + "MCP server '%s': skipping tool '%s' (filtered by config)", + name, + mcp_tool.name, ) continue - registry.register( - name=tool_name_prefixed, - toolset=toolset_name, - schema=schema, - handler=_make_tool_handler(name, mcp_tool.name, server.tool_timeout), - check_fn=_make_check_fn(name), - is_async=False, - description=schema["description"], + _scan_mcp_description(name, mcp_tool.name, mcp_tool.description or "") + schema = _convert_mcp_schema(name, mcp_tool) + candidates.append( + { + "registry_name": schema["name"], + "origin": f"tool {mcp_tool.name!r}", + "schema": schema, + "handler": _make_tool_handler( + name, mcp_tool.name, server.tool_timeout + ), + "check_fn": check_fn, + } ) - _track_mcp_tool_server(tool_name_prefixed, name) - registered_names.append(tool_name_prefixed) - # Register MCP Resources & Prompts utility tools, filtered by config and - # only when the server actually supports the corresponding capability. - _handler_factories = { + # Generated resource/prompt utility tools share the same namespace as raw + # MCP tools, so they must participate in the same collision preflight. + handler_factories = { "list_resources": _make_list_resources_handler, "read_resource": _make_read_resource_handler, "list_prompts": _make_list_prompts_handler, "get_prompt": _make_get_prompt_handler, } - check_fn = _make_check_fn(name) for entry in _select_utility_schemas(name, server, config): schema = entry["schema"] handler_key = entry["handler_key"] - handler = _handler_factories[handler_key](name, server.tool_timeout) - util_name = schema["name"] + candidates.append( + { + "registry_name": schema["name"], + "origin": f"generated utility {handler_key!r}", + "schema": schema, + "handler": handler_factories[handler_key]( + name, server.tool_timeout + ), + "check_fn": check_fn, + } + ) - # Same collision guard for utility tools. - existing_toolset = registry.get_toolset_for_tool(util_name) - if existing_toolset and not existing_toolset.startswith("mcp-"): - logger.warning( - "MCP server '%s': utility tool '%s' collides with built-in " - "tool in toolset '%s' — skipping to preserve built-in", - name, util_name, existing_toolset, + # Exact duplicate rows from a server are harmless but should not inflate + # counts. Distinct origins that collapse to one normalized name are unsafe. + unique_candidates: List[dict] = [] + seen_candidates: set[tuple[str, str]] = set() + origins_by_name: Dict[str, set[str]] = {} + for candidate in candidates: + key = (candidate["registry_name"], candidate["origin"]) + if key in seen_candidates: + logger.debug( + "MCP server '%s': duplicate registration candidate %s for '%s'; " + "keeping one", + name, + candidate["origin"], + candidate["registry_name"], ) continue + seen_candidates.add(key) + unique_candidates.append(candidate) + origins_by_name.setdefault(candidate["registry_name"], set()).add( + candidate["origin"] + ) + + ambiguous_names = { + registry_name: sorted(origins) + for registry_name, origins in origins_by_name.items() + if len(origins) > 1 + } + for registry_name, origins in sorted(ambiguous_names.items()): + logger.error( + "MCP server '%s': name normalization collision for '%s' from %s; " + "skipping every colliding entry instead of choosing an arbitrary " + "handler", + name, + registry_name, + ", ".join(origins), + ) + + for candidate in unique_candidates: + registry_name = candidate["registry_name"] + if registry_name in ambiguous_names: + continue + + existing_toolset = registry.get_toolset_for_tool(registry_name) + if existing_toolset and existing_toolset != toolset_name: + if existing_toolset.startswith("mcp-"): + logger.error( + "MCP server '%s': %s normalizes to '%s', already owned by " + "MCP toolset '%s' — skipping to preserve the existing owner", + name, + candidate["origin"], + registry_name, + existing_toolset, + ) + else: + logger.warning( + "MCP server '%s': %s (→ '%s') collides with built-in tool " + "in toolset '%s' — skipping to preserve built-in", + name, + candidate["origin"], + registry_name, + existing_toolset, + ) + continue registry.register( - name=util_name, + name=registry_name, toolset=toolset_name, - schema=schema, - handler=handler, - check_fn=check_fn, + schema=candidate["schema"], + handler=candidate["handler"], + check_fn=candidate["check_fn"], is_async=False, - description=schema["description"], + description=candidate["schema"]["description"], ) - _track_mcp_tool_server(util_name, name) - registered_names.append(util_name) + + # The pre-check above is advisory only. Multiple servers connect in + # parallel, so ToolRegistry.register() is the atomic ownership gate. + if registry.get_toolset_for_tool(registry_name) != toolset_name: + logger.error( + "MCP server '%s': registration of %s as '%s' was rejected by " + "the registry; skipping provenance/count updates", + name, + candidate["origin"], + registry_name, + ) + continue + + _track_mcp_tool_server(registry_name, name) + registered_names.append(registry_name) if registered_names: registry.register_toolset_alias(name, toolset_name) return registered_names - async def _discover_and_register_server(name: str, config: dict) -> List[str]: """Connect to a single MCP server, discover tools, and register them. @@ -5804,9 +5895,9 @@ def register_mcp_servers(servers: Dict[str, dict]) -> List[str]: # Track which servers opt-in to parallel tool calls (idempotent). for srv_name, srv_cfg in servers.items(): if _parse_boolish(srv_cfg.get("supports_parallel_tool_calls", False), default=False): - _parallel_safe_servers.add(sanitize_mcp_name_component(srv_name)) + _parallel_safe_servers.add(srv_name) else: - _parallel_safe_servers.discard(sanitize_mcp_name_component(srv_name)) + _parallel_safe_servers.discard(srv_name) for srv in stale_cached: _signal_reconnect(srv) diff --git a/tools/registry.py b/tools/registry.py index 354da7123fd7..ede7c487dd9d 100644 --- a/tools/registry.py +++ b/tools/registry.py @@ -388,18 +388,7 @@ class ToolRegistry: with self._lock: existing = self._tools.get(name) if existing and existing.toolset != toolset: - # Allow MCP-to-MCP overwrites (legitimate: server refresh, - # or two MCP servers with overlapping tool names). - both_mcp = ( - existing.toolset.startswith("mcp-") - and toolset.startswith("mcp-") - ) - if both_mcp: - logger.debug( - "Tool '%s': MCP toolset '%s' overwriting MCP toolset '%s'", - name, toolset, existing.toolset, - ) - elif override: + if override: _owner = self._plugin_owner_of(handler) if _owner is not None and not self._plugin_override_policy.get(_owner, False): logger.error( @@ -423,8 +412,9 @@ class ToolRegistry: name, toolset, existing.toolset, ) else: - # Reject shadowing — prevent plugins/MCP from overwriting - # built-in tools or vice versa. + # Reject every cross-toolset shadow, including MCP-to-MCP + # collisions. Legitimate MCP reconnect/refresh re-registers + # within the same canonical toolset and remains allowed. logger.error( "Tool registration REJECTED: '%s' (toolset '%s') would " "shadow existing tool from toolset '%s'. Pass "