diff --git a/tests/gateway/test_multiplex_mcp_discovery.py b/tests/gateway/test_multiplex_mcp_discovery.py index 2e39405c69..28905a5144 100644 --- a/tests/gateway/test_multiplex_mcp_discovery.py +++ b/tests/gateway/test_multiplex_mcp_discovery.py @@ -179,6 +179,109 @@ def test_scope_visibility_rejects_a_foreign_server_route( assert mcp_tool._server_tool_scopes["shared"] == {launch_scope} +def test_shared_server_tools_are_callable_and_removed_on_non_owner_reload( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + from agent.secret_scope import set_multiplex_active + from hermes_constants import ( + hermes_home_key, + reset_hermes_home_override, + set_hermes_home_override, + ) + from tools import mcp_tool + from tools import mcp_tool_config as _mcp_config + from tools import mcp_tool_discovery as _mcp_discovery + from tools.registry import registry + + worker_home = tmp_path / "profiles" / "worker" + launch_home = tmp_path / "default" + worker_home.mkdir(parents=True) + launch_home.mkdir() + worker_token = set_hermes_home_override(worker_home) + previous_multiplex = set_multiplex_active(True) + worker_scope = hermes_home_key() + launch_scope = hermes_home_key(launch_home) + tool = SimpleNamespace( + name="echo", + description="Echo a value", + inputSchema={"type": "object", "properties": {}}, + annotations=None, + ) + server = SimpleNamespace( + name="shared", + session=object(), + _tools=[tool], + tool_timeout=30, + _registered_tool_names=[], + _config={}, + initialize_result=None, + ) + owner_tool_name = "mcp__shared__echo" + registry.register( + owner_tool_name, + "mcp-shared", + {"name": owner_tool_name, "description": "Echo a value", "type": "object"}, + lambda **_kwargs: None, + scope=launch_scope, + ) + registry.register_toolset_alias("shared", "mcp-shared") + server._registered_tool_names = [owner_tool_name] + with mcp_tool._lock: + saved = { + "_servers": dict(mcp_tool._servers), + "_server_scope_keys": dict(mcp_tool._server_scope_keys), + "_server_tool_scopes": dict(mcp_tool._server_tool_scopes), + "_mcp_tool_server_names": dict(mcp_tool._mcp_tool_server_names), + } + mcp_tool._servers.clear() + mcp_tool._server_scope_keys.clear() + mcp_tool._server_tool_scopes.clear() + mcp_tool._mcp_tool_server_names.clear() + mcp_tool._servers["shared"] = server + mcp_tool._server_scope_keys["shared"] = launch_scope + mcp_tool._server_tool_scopes["shared"] = {launch_scope} + + try: + monkeypatch.setattr(mcp_tool, "_ensure_mcp_sdk", lambda: True) + monkeypatch.setattr(_mcp_config, "_filter_suspicious_mcp_servers", lambda servers: servers) + assert _mcp_discovery.register_mcp_servers({"shared": {}}) + tool_names = registry.get_tool_names_for_toolset("mcp-shared") + assert tool_names + assert callable(registry.get_entry(tool_names[0]).handler) + + # Changing the worker route removes only the worker overlay; the shared + # live connection and launch owner remain intact. + assert _mcp_discovery.register_mcp_servers( + {"shared": {"url": "https://worker.example/mcp"}} + ) == [] + assert registry.get_tool_names_for_toolset("mcp-shared") == [] + with mcp_tool._lock: + assert mcp_tool._server_scope_keys["shared"] == launch_scope + assert mcp_tool._server_tool_scopes["shared"] == {launch_scope} + assert mcp_tool._servers["shared"] is server + assert registry.snapshot_registration(owner_tool_name, scope=launch_scope) is not None + assert registry.get_toolset_alias_target("shared") == "mcp-shared" + + # Removing the server from the worker config has the same scoped cleanup. + assert _mcp_discovery.register_mcp_servers({}) == [] + assert registry.get_tool_names_for_toolset("mcp-shared") == [] + with mcp_tool._lock: + assert mcp_tool._server_scope_keys["shared"] == launch_scope + assert mcp_tool._server_tool_scopes["shared"] == {launch_scope} + assert mcp_tool._servers["shared"] is server + finally: + for tool_name in list(registry.get_tool_names_for_toolset("mcp-shared")): + registry.deregister(tool_name, scope=worker_scope) + registry.deregister(owner_tool_name, scope=launch_scope) + with mcp_tool._lock: + for name, value in saved.items(): + target = getattr(mcp_tool, name) + target.clear() + target.update(value) + set_multiplex_active(previous_multiplex) + reset_hermes_home_override(worker_token) + + def test_deregister_scope_kwarg_targets_overlay_and_keeps_plugin_confinement() -> None: from tools.registry import ToolRegistry diff --git a/tools/mcp_tool_discovery.py b/tools/mcp_tool_discovery.py index a782d8e78c..686ead1c1d 100644 --- a/tools/mcp_tool_discovery.py +++ b/tools/mcp_tool_discovery.py @@ -168,10 +168,8 @@ def _ensure_lazy_server_connected(server_name: str) -> bool: # The cached manifest may advertise tools the live server no longer serves. phantom_names = [n for n in cached_names if n not in live_names] if phantom_names: - from tools.registry import registry for tool_name in phantom_names: - registry.deregister(tool_name, scope=_core._server_registry_scope(server_name)) - _registration._forget_mcp_tool_server(tool_name) + _registration._deregister_mcp_tool_all_scopes(server_name, tool_name) logger.info("MCP server '%s': deregistered %d phantom cached tool(s) not served live (stale schema-cache " "fingerprint %s): %s", server_name, len(phantom_names), stale_fingerprint, ", ".join(phantom_names)) return server is not None and server.session is not None @@ -375,10 +373,13 @@ def register_mcp_servers(servers: Dict[str, dict]) -> List[str]: logger.debug("MCP SDK not available -- skipping explicit MCP registration") return [] servers = _config._filter_suspicious_mcp_servers(servers) + scoped_healed = _registration.register_connected_into_current_scope(servers) if not servers: logger.debug("No explicit MCP servers provided") - return [] + return _registration._existing_tool_names() if _core._mcp_registry_scope() is not None else [] new_servers = _select_new_servers(servers) + if scoped_healed: + logger.info("MCP: registered %d already-connected server(s) into this profile scope", scoped_healed) if not new_servers: return _registration._existing_tool_names() new_servers, lazy_registered, lazy_server_count = _register_lazy_from_cache(new_servers) diff --git a/tools/mcp_tool_health.py b/tools/mcp_tool_health.py index 1b328cfeda..6db839bb5c 100644 --- a/tools/mcp_tool_health.py +++ b/tools/mcp_tool_health.py @@ -9,7 +9,6 @@ import time from typing import Iterable, Optional from tools.mcp_tool_errors import _is_method_not_found_error, _unwrap_exception_group from tools.mcp_tool_schema import mcp_prefixed_tool_name -from tools.mcp_tool_registration import _forget_mcp_tool_server from tools.mcp_tool_common import _core from tools import mcp_tool_registration as _registration @@ -128,8 +127,7 @@ class MCPServerHealthMixin: from tools.registry import registry for tool_name in tool_names: if registry.get_toolset_for_tool(tool_name) == f"mcp-{self.name}": - registry.deregister(tool_name, scope=_core._server_registry_scope(self.name)) - _forget_mcp_tool_server(tool_name) + _registration._deregister_mcp_tool_all_scopes(self.name, tool_name) async def _refresh_tools(self): """Re-fetch tools on ``tools/list_changed`` and update the registry. The lock serializes rapid-fire diff --git a/tools/mcp_tool_registration.py b/tools/mcp_tool_registration.py index 10b65b5238..98cef3064a 100644 --- a/tools/mcp_tool_registration.py +++ b/tools/mcp_tool_registration.py @@ -4,6 +4,7 @@ resolution and the schema-cache write-through. Both entry points (``_register_se live, ``_register_from_cache_sync`` lazy) build ``_Candidate`` records for ``_register_candidates``.""" import logging +import threading from dataclasses import dataclass from types import SimpleNamespace from typing import TYPE_CHECKING, Any, Callable, Dict, Iterable, List, Optional @@ -20,6 +21,7 @@ if TYPE_CHECKING: # pragma: no cover from tools.mcp_tool import MCPServerTask logger = logging.getLogger("tools.mcp_tool") +_SCOPE_REFRESH_LOCKS = tuple(threading.RLock() for _ in range(16)) _UTILITY_ORIGIN_PREFIX = "generated utility " # Utility tool key -> handler factory; each takes (server_name, tool_timeout). @@ -68,6 +70,51 @@ def _forget_mcp_tool_server(tool_name: str) -> None: _core._mcp_tool_server_names.pop(tool_name, None) +def _deregister_mcp_tool_all_scopes(server_name: str, tool_name: str) -> None: + """Deregister one server tool from every profile overlay that owns it.""" + from tools.registry import registry + + with _core._lock: + scopes = set(_core._server_tool_scopes.get(server_name, ())) + if not scopes: + scopes = {_core._server_registry_scope(server_name)} + for scope in scopes: + registry.deregister(tool_name, scope=scope) + _forget_mcp_tool_server(tool_name) + _restore_server_toolset_alias(server_name) + + +def _restore_server_toolset_alias(server_name: str) -> None: + """Keep the process-global alias while any profile still owns this server's tools.""" + from tools.registry import registry + + with _core._lock: + server = _core._servers.get(server_name) + scopes = set(_core._server_tool_scopes.get(server_name, ())) + tool_names = list(getattr(server, "_registered_tool_names", ()) if server is not None else ()) + if any( + registry.snapshot_registration(tool_name, scope=scope) is not None + for scope in scopes for tool_name in tool_names + ): + registry.register_toolset_alias(server_name, f"mcp-{server_name}") + + +def _remove_server_scope(server_name: str, scope: str) -> None: + """Remove one profile's MCP overlay for a shared live connection.""" + from tools.registry import registry + + for tool_name in registry.get_tool_names_for_toolset(f"mcp-{server_name}"): + registry.deregister(tool_name, scope=scope) + with _core._lock: + scopes = set(_core._server_tool_scopes.get(server_name, ())) + scopes.discard(scope) + if scopes: + _core._server_tool_scopes[server_name] = scopes + else: + _core._server_tool_scopes.pop(server_name, None) + _restore_server_toolset_alias(server_name) + + def _select_utility_schemas(server_name: str, server: "MCPServerTask", config: dict) -> List[dict]: """Utility schemas allowed by config (``tools.resources``/``tools.prompts``) and advertised capabilities. ``initialize_result.capabilities`` is the truth (sub-object non-None iff the @@ -99,6 +146,25 @@ def _select_utility_schemas(server_name: str, server: "MCPServerTask", config: d def _existing_tool_names() -> List[str]: """Tool names for all connected servers plus lazy (cache-registered) servers, whose tools live only in the registry.""" + scope = _core._mcp_registry_scope() + if scope is not None: + from tools.registry import registry + + with _core._lock: + server_names = [ + name for name in _core._servers + if _core._server_visible_in_scope(name, scope) + ] + server_names.extend( + name for name in _core._lazy_server_tool_names + if name not in _core._servers + ) + return sorted({ + tool_name + for server_name in server_names + for tool_name in registry.get_tool_names_for_toolset(f"mcp-{server_name}") + }) + names: List[str] = [] for server in _core._servers.values(): names.extend(server._registered_tool_names if hasattr(server, "_registered_tool_names") @@ -228,6 +294,7 @@ def _register_candidates(name: str, candidates: List[_Candidate], *, check_fn: C from tools.registry import registry toolset_name = f"mcp-{name}" registered: List[str] = [] + scope_value = scope() for c in candidates: existing_toolset = registry.get_toolset_for_tool(c.registry_name) if existing_toolset and existing_toolset != toolset_name: # foreign owner: skip, preserve it @@ -244,9 +311,12 @@ def _register_candidates(name: str, candidates: List[_Candidate], *, check_fn: C continue registry.register( name=c.registry_name, toolset=toolset_name, schema=c.schema, handler=c.handler, check_fn=check_fn, - is_async=False, description=c.schema.get("description") or "", scope=scope()) + is_async=False, description=c.schema.get("description") or "", scope=scope_value) if registry.get_toolset_for_tool(c.registry_name) == toolset_name: _track_mcp_tool_server(c.registry_name, name) + if scope_value is not None: + with _core._lock: + _core._server_tool_scopes.setdefault(name, set()).add(scope_value) registered.append(c.registry_name) elif not lazy: logger.error("MCP server '%s': registration of %s as '%s' was rejected by the registry; " @@ -302,6 +372,79 @@ def _register_server_tools(name: str, server: "MCPServerTask", config: dict) -> return registered +def _server_enabled(config: dict) -> bool: + return _parse_boolish(config.get("enabled", True), default=True) + + +def _same_server_route(server: Any, config: dict) -> bool: + from tools.mcp_schema_cache import config_fingerprint + + return config_fingerprint(getattr(server, "_config", {}) or {}) == config_fingerprint(config) + + +def register_connected_into_current_scope(servers: dict) -> int: + """Serialize shared-scope reconciliation and registration for one discovery pass.""" + scope = _core._mcp_registry_scope() + if scope is None: + return 0 + with _SCOPE_REFRESH_LOCKS[hash(scope) % len(_SCOPE_REFRESH_LOCKS)]: + return _register_connected_into_current_scope(servers) + + +def _register_connected_into_current_scope(servers: dict) -> int: + """Heal the current profile's MCP overlay from already-connected shared servers. + + A shared live connection remains owned by the profile that opened it, but a profile with the + same route must still receive its callable tool entries. The current profile's config is the + allowlist, and route fingerprints prevent borrowing a differently-authenticated connection. + Missing or changed config entries remove only this profile's overlay. + """ + from tools.registry import registry + + scope = _core._mcp_registry_scope() + if scope is None: + return 0 + + from tools.mcp_schema_cache import config_fingerprint + + with _core._lock: + stale = [] + for name, scopes in _core._server_tool_scopes.items(): + if scope not in scopes: + continue + server = _core._servers.get(name) + config = servers.get(name) + if (config is None or not _server_enabled(config) or server is None + or getattr(server, "session", None) is None + or config_fingerprint(getattr(server, "_config", {}) or {}) != config_fingerprint(config)): + stale.append(name) + for name in stale: + _remove_server_scope(name, scope) + + registered_servers = 0 + for name, config in servers.items(): + if not _server_enabled(config): + continue + with _core._lock: + server = _core._servers.get(name) + if server is None or getattr(server, "session", None) is None or not _same_server_route(server, config): + continue + if registry.get_tool_names_for_toolset(f"mcp-{name}"): + continue + candidates = _tool_candidates(name, server._tools, _make_tool_filter(name, config), server.tool_timeout) + candidates += _utility_candidates( + name, _select_utility_schemas(name, server, config), server.tool_timeout) + names = _register_candidates( + name, _resolve_name_collisions(name, candidates), + check_fn=_make_check_fn(name), scope=lambda: scope, lazy=False) + if names: + registered_servers += 1 + with _core._lock: + server._registered_tool_names = sorted( + set(getattr(server, "_registered_tool_names", []) or ()) | set(names)) + return registered_servers + + def _register_from_cache_sync(name: str, config: dict, entry: dict) -> List[str]: """Lazy startup: register from a cached manifest with no child process (first real call goes through ``_ensure_lazy_server_connected``). Trust metadata is recorded first so the diff --git a/tools/mcp_tool_server_run.py b/tools/mcp_tool_server_run.py index 20ff8e770a..4daded5437 100644 --- a/tools/mcp_tool_server_run.py +++ b/tools/mcp_tool_server_run.py @@ -413,8 +413,6 @@ class MCPServerRunMixin: def _deregister_tools(self) -> None: """Drop this server's tools from the registry (idempotent); on shutdown AND budget exhaustion, so a dead server never leaves phantom tools in the prompt.""" - from tools.registry import registry for tool_name in list(getattr(self, "_registered_tool_names", [])): - registry.deregister(tool_name, scope=_core._server_registry_scope(self.name)) - _registration._forget_mcp_tool_server(tool_name) + _registration._deregister_mcp_tool_all_scopes(self.name, tool_name) self._registered_tool_names = []