fix(gateway): register shared MCP tools per profile

This commit is contained in:
joaomarcos
2026-09-08 02:16:09 -03:00
committed by kshitij
parent 7107185f58
commit d02edc2cbc
5 changed files with 254 additions and 11 deletions

View File

@@ -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

View File

@@ -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)

View File

@@ -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

View File

@@ -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

View File

@@ -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 = []