fix(gateway): register shared MCP tools per profile
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 = []
|
||||
|
||||
Reference in New Issue
Block a user