simplify(compat): terminal/file/environments — drop 42 re-exports/aliases, repoint 20 callers + 41 test files

tools/terminal_tool.py: drop 30 pure re-export names (lifecycle/config/backends/
sudo/guards/result/interrupt/utils/_DockerEnvironment/is_managed_tool_gateway_ready)
and the noqa-F401 comments on the 25 names the facade itself uses. Sibling modules
(terminal_tool_backends/_result/_sudo/_lifecycle, environments/base, process_registry)
that read removed names through the facade now import from the defining module.
tools/environments/base.py: drop 11 re-exports (base_output/base_session_env/
path_utils) and the BaseEnvironment.stop() compat alias (no in-tree caller; the
lifecycle hasattr(env, 'stop') fallback stays for third-party envs).
tools/environments/docker.py: drop 1 re-export + the re-export comment.
Callers/tests repointed to tools.terminal_tool_{lifecycle,backends,sudo,config,
guards,result}, tools.interrupt, tools.environments.{base_output,base_session_env,
path_utils}.
This commit is contained in:
Teknium
2026-09-03 13:29:55 -07:00
parent fcbe4acbef
commit f6938b37f3
75 changed files with 683 additions and 645 deletions
+1 -1
View File
@@ -33,7 +33,7 @@ def test_main_skips_configured_mcp_discovery_when_requested(monkeypatch):
monkeypatch.setattr(entry, "_load_env", lambda: None)
monkeypatch.setenv("HERMES_ACP_SKIP_CONFIGURED_MCP", "1")
monkeypatch.setattr(
"tools.mcp_tool.discover_mcp_tools",
"tools.mcp_tool_discovery.discover_mcp_tools",
lambda: discovery_calls.append(True),
)
monkeypatch.setattr(acp, "run_agent", fake_run_agent)
+5 -5
View File
@@ -82,7 +82,7 @@ class TestMcpRegistrationE2E:
{"function": {"name": "terminal"}},
]
with patch("tools.mcp_tool.register_mcp_servers", side_effect=mock_register), \
with patch("tools.mcp_tool_discovery.register_mcp_servers", side_effect=mock_register), \
patch("model_tools.get_tool_definitions", return_value=fake_tools):
resp = await acp_agent.new_session(cwd="/tmp", mcp_servers=servers)
@@ -217,7 +217,7 @@ class TestMcpSanitizationE2E:
fake_tools = [{"function": {"name": "mcp_ai_exa_exa_search"}}]
with patch("tools.mcp_tool.register_mcp_servers", side_effect=mock_register), \
with patch("tools.mcp_tool_discovery.register_mcp_servers", side_effect=mock_register), \
patch("model_tools.get_tool_definitions", return_value=fake_tools):
resp = await acp_agent.new_session(cwd="/tmp", mcp_servers=servers)
@@ -254,7 +254,7 @@ class TestSessionLifecycleMcpE2E:
state.agent.tools = []
state.agent.valid_tool_names = set()
with patch("tools.mcp_tool.register_mcp_servers", side_effect=mock_register), \
with patch("tools.mcp_tool_discovery.register_mcp_servers", side_effect=mock_register), \
patch("model_tools.get_tool_definitions", return_value=[]):
await acp_agent.load_session(cwd="/tmp", session_id=sid, mcp_servers=servers)
@@ -281,7 +281,7 @@ class TestSessionLifecycleMcpE2E:
state.agent.tools = []
state.agent.valid_tool_names = set()
with patch("tools.mcp_tool.register_mcp_servers", side_effect=mock_register), \
with patch("tools.mcp_tool_discovery.register_mcp_servers", side_effect=mock_register), \
patch("model_tools.get_tool_definitions", return_value=[]):
await acp_agent.resume_session(cwd="/tmp", session_id=sid, mcp_servers=servers)
@@ -303,7 +303,7 @@ class TestSessionLifecycleMcpE2E:
return []
# Need to set up the forked session's agent too
with patch("tools.mcp_tool.register_mcp_servers", side_effect=mock_register), \
with patch("tools.mcp_tool_discovery.register_mcp_servers", side_effect=mock_register), \
patch("model_tools.get_tool_definitions", return_value=[]):
fork_resp = await acp_agent.fork_session(
cwd="/tmp", session_id=sid, mcp_servers=servers
+3 -3
View File
@@ -641,7 +641,7 @@ class TestRegisterSessionMcpServers:
registered_config.update(config_map)
return ["mcp_test_server_tool1"]
with patch("tools.mcp_tool.register_mcp_servers", side_effect=capture_register), \
with patch("tools.mcp_tool_discovery.register_mcp_servers", side_effect=capture_register), \
patch("model_tools.get_tool_definitions", return_value=[]):
await agent._register_session_mcp_servers(state, [server])
@@ -682,7 +682,7 @@ class TestRegisterSessionMcpServers:
{"function": {"name": "terminal"}},
]
with patch("tools.mcp_tool.register_mcp_servers", return_value=["mcp_srv_search"]), \
with patch("tools.mcp_tool_discovery.register_mcp_servers", return_value=["mcp_srv_search"]), \
patch("model_tools.get_tool_definitions", return_value=fake_tools) as mock_defs:
await agent._register_session_mcp_servers(state, [server])
@@ -723,6 +723,6 @@ class TestRegisterSessionMcpServers:
env=[],
)
with patch("tools.mcp_tool.register_mcp_servers", side_effect=RuntimeError("boom")):
with patch("tools.mcp_tool_discovery.register_mcp_servers", side_effect=RuntimeError("boom")):
# Should not raise
await agent._register_session_mcp_servers(state, [server])
+14 -14
View File
@@ -101,8 +101,8 @@ def test_acp_background_discovery_does_not_block_startup(monkeypatch):
)
monkeypatch.setitem(
sys.modules,
"tools.mcp_tool",
_mod("tools.mcp_tool", discover_mcp_tools=_blocking_discover),
"tools.mcp_tool_discovery",
_mod("tools.mcp_tool_discovery", discover_mcp_tools=_blocking_discover),
)
start = time.monotonic()
@@ -149,8 +149,8 @@ def test_acp_late_refresh_adds_tools_when_discovery_lands_after_build(monkeypatc
)
monkeypatch.setitem(
sys.modules,
"tools.mcp_tool",
_mod("tools.mcp_tool", discover_mcp_tools=_slow_discover),
"tools.mcp_tool_discovery",
_mod("tools.mcp_tool_discovery", discover_mcp_tools=_slow_discover),
)
mcp_startup.start_background_mcp_discovery(
@@ -178,8 +178,8 @@ def test_acp_late_refresh_adds_tools_when_discovery_lands_after_build(monkeypatc
monkeypatch.setitem(
sys.modules,
"tools.mcp_tool",
_mod("tools.mcp_tool", refresh_agent_mcp_tools=_fake_refresh),
"tools.mcp_tool_agent",
_mod("tools.mcp_tool_agent", refresh_agent_mcp_tools=_fake_refresh),
)
# Trigger late-refresh.
@@ -226,8 +226,8 @@ def test_acp_late_refresh_skips_after_first_turn(monkeypatch):
)
monkeypatch.setitem(
sys.modules,
"tools.mcp_tool",
_mod("tools.mcp_tool", discover_mcp_tools=_slow_discover),
"tools.mcp_tool_discovery",
_mod("tools.mcp_tool_discovery", discover_mcp_tools=_slow_discover),
)
mcp_startup.start_background_mcp_discovery(
@@ -249,8 +249,8 @@ def test_acp_late_refresh_skips_after_first_turn(monkeypatch):
monkeypatch.setitem(
sys.modules,
"tools.mcp_tool",
_mod("tools.mcp_tool", refresh_agent_mcp_tools=_fake_refresh),
"tools.mcp_tool_agent",
_mod("tools.mcp_tool_agent", refresh_agent_mcp_tools=_fake_refresh),
)
acp_agent._schedule_mcp_late_refresh(state)
@@ -293,8 +293,8 @@ def test_acp_late_refresh_skips_while_turn_running(monkeypatch):
)
monkeypatch.setitem(
sys.modules,
"tools.mcp_tool",
_mod("tools.mcp_tool", discover_mcp_tools=_slow_discover),
"tools.mcp_tool_discovery",
_mod("tools.mcp_tool_discovery", discover_mcp_tools=_slow_discover),
)
mcp_startup.start_background_mcp_discovery(
@@ -316,8 +316,8 @@ def test_acp_late_refresh_skips_while_turn_running(monkeypatch):
monkeypatch.setitem(
sys.modules,
"tools.mcp_tool",
_mod("tools.mcp_tool", refresh_agent_mcp_tools=_fake_refresh),
"tools.mcp_tool_agent",
_mod("tools.mcp_tool_agent", refresh_agent_mcp_tools=_fake_refresh),
)
acp_agent._schedule_mcp_late_refresh(state)
+1 -1
View File
@@ -446,7 +446,7 @@ def test_between_turns_refresh_adds_late_tool_when_servers_registered():
new_def = {"type": "function", "function": {"name": "mcp_x_tool", "description": "", "parameters": {}}}
import model_tools
with patch("tools.mcp_tool.has_registered_mcp_tools", return_value=True), \
with patch("tools.mcp_tool_discovery.has_registered_mcp_tools", return_value=True), \
patch.object(model_tools, "get_tool_definitions", return_value=[new_def]):
_build(agent)
+1 -1
View File
@@ -132,7 +132,7 @@ class TestRunCleanupWiring(unittest.TestCase):
patch.object(
cli_mod, "_cleanup_all_browsers", patches["_cleanup_all_browsers"]
),
patch("tools.mcp_tool.shutdown_mcp_servers", lambda *a, **k: None),
patch("tools.mcp_tool_lifecycle.shutdown_mcp_servers", lambda *a, **k: None),
patch(
"agent.auxiliary_client.shutdown_cached_clients",
lambda *a, **k: None,
+6 -6
View File
@@ -53,11 +53,11 @@ def _tick(job, tmp_path, current_provider, deliveries):
return None
with patch("cron.scheduler._hermes_home", tmp_path), \
patch("cron.scheduler._resolve_origin", return_value=None), \
patch("cron.scheduler_delivery._resolve_origin", return_value=None), \
patch("hermes_cli.env_loader.load_hermes_dotenv"), \
patch("hermes_cli.env_loader.reset_secret_source_cache"), \
patch("hermes_state.get_shared_session_db", return_value=fake_db), \
patch("tools.mcp_tool.discover_mcp_tools", return_value=[]), \
patch("hermes_state_registry.acquire", return_value=fake_db), \
patch("tools.mcp_tool_discovery.discover_mcp_tools", return_value=[]), \
patch("hermes_cli.runtime_provider.resolve_runtime_provider",
return_value={
"api_key": "test-key",
@@ -138,11 +138,11 @@ class TestDriftAlertOnce:
cron_jobs.save_jobs([job])
fresh = [j for j in cron_jobs.load_jobs() if j["id"] == job["id"]][0]
with patch("cron.scheduler._hermes_home", tmp_path), \
patch("cron.scheduler._resolve_origin", return_value=None), \
patch("cron.scheduler_delivery._resolve_origin", return_value=None), \
patch("hermes_cli.env_loader.load_hermes_dotenv"), \
patch("hermes_cli.env_loader.reset_secret_source_cache"), \
patch("hermes_state.get_shared_session_db", return_value=fake_db), \
patch("tools.mcp_tool.discover_mcp_tools", return_value=[]), \
patch("hermes_state_registry.acquire", return_value=fake_db), \
patch("tools.mcp_tool_discovery.discover_mcp_tools", return_value=[]), \
patch("hermes_cli.runtime_provider.resolve_runtime_provider",
return_value={
"api_key": "test-key",
+3 -3
View File
@@ -52,11 +52,11 @@ def _tick_failing(job, tmp_path, deliveries, error="boom unrelated"):
with cron_jobs.use_cron_store(tmp_path), \
patch("cron.scheduler._hermes_home", tmp_path), \
patch("cron.scheduler._resolve_origin", return_value=None), \
patch("cron.scheduler_delivery._resolve_origin", return_value=None), \
patch("hermes_cli.env_loader.load_hermes_dotenv"), \
patch("hermes_cli.env_loader.reset_secret_source_cache"), \
patch("hermes_state.get_shared_session_db", return_value=fake_db), \
patch("tools.mcp_tool.discover_mcp_tools", return_value=[]), \
patch("hermes_state_registry.acquire", return_value=fake_db), \
patch("tools.mcp_tool_discovery.discover_mcp_tools", return_value=[]), \
patch("hermes_cli.runtime_provider.resolve_runtime_provider",
return_value={
"api_key": "test-key",
+1 -1
View File
@@ -75,7 +75,7 @@ def _orphan_claimed_row(executions, job_id: str) -> str:
def _run_tick():
with (
patch.object(scheduler_mod, "get_due_jobs", return_value=[]),
patch("tools.mcp_tool._kill_orphaned_mcp_children", lambda: None),
patch("tools.mcp_tool_lifecycle._kill_orphaned_mcp_children", lambda: None),
):
return scheduler_mod.tick(verbose=False)
+1 -1
View File
@@ -29,7 +29,7 @@ def _run_idle_tick(**kwargs):
patch.object(scheduler_mod, "get_due_jobs", return_value=[]),
patch.object(scheduler_mod, "load_config", side_effect=_fake_load_config),
patch(
"tools.mcp_tool._kill_orphaned_mcp_children",
"tools.mcp_tool_lifecycle._kill_orphaned_mcp_children",
side_effect=_fake_sweep,
),
):
+9 -9
View File
@@ -70,11 +70,11 @@ def _run_job_patched(job, tmp_path, *, resolve=None, skill_view=None):
fake_db = MagicMock()
patches = [
patch("cron.scheduler._hermes_home", tmp_path),
patch("cron.scheduler._resolve_origin", return_value=None),
patch("cron.scheduler_delivery._resolve_origin", return_value=None),
patch("hermes_cli.env_loader.load_hermes_dotenv"),
patch("hermes_cli.env_loader.reset_secret_source_cache"),
patch("hermes_state.get_shared_session_db", return_value=fake_db),
patch("tools.mcp_tool.discover_mcp_tools", return_value=[]),
patch("hermes_state_registry.acquire", return_value=fake_db),
patch("tools.mcp_tool_discovery.discover_mcp_tools", return_value=[]),
]
if resolve is None:
patches.append(
@@ -139,11 +139,11 @@ class TestMissingProviderKeyBlocks:
for _tick in range(2):
fresh = [j for j in cron_jobs.load_jobs() if j["id"] == job["id"]][0]
with patch("cron.scheduler._hermes_home", tmp_path), \
patch("cron.scheduler._resolve_origin", return_value=None), \
patch("cron.scheduler_delivery._resolve_origin", return_value=None), \
patch("hermes_cli.env_loader.load_hermes_dotenv"), \
patch("hermes_cli.env_loader.reset_secret_source_cache"), \
patch("hermes_state.get_shared_session_db", return_value=fake_db), \
patch("tools.mcp_tool.discover_mcp_tools", return_value=[]), \
patch("hermes_state_registry.acquire", return_value=fake_db), \
patch("tools.mcp_tool_discovery.discover_mcp_tools", return_value=[]), \
patch("hermes_cli.runtime_provider.resolve_runtime_provider",
side_effect=_AuthErrorFactory()), \
patch.object(sched, "_deliver_result", side_effect=fake_deliver), \
@@ -243,11 +243,11 @@ class TestOptOut:
for _tick in range(2):
fresh = [j for j in cron_jobs.load_jobs() if j["id"] == job["id"]][0]
with patch("cron.scheduler._hermes_home", tmp_path), \
patch("cron.scheduler._resolve_origin", return_value=None), \
patch("cron.scheduler_delivery._resolve_origin", return_value=None), \
patch("hermes_cli.env_loader.load_hermes_dotenv"), \
patch("hermes_cli.env_loader.reset_secret_source_cache"), \
patch("hermes_state.get_shared_session_db", return_value=fake_db), \
patch("tools.mcp_tool.discover_mcp_tools", return_value=[]), \
patch("hermes_state_registry.acquire", return_value=fake_db), \
patch("tools.mcp_tool_discovery.discover_mcp_tools", return_value=[]), \
patch("hermes_cli.runtime_provider.resolve_runtime_provider",
side_effect=_AuthErrorFactory()), \
patch.object(sched, "_deliver_result", side_effect=fake_deliver), \
+8 -8
View File
@@ -1042,14 +1042,14 @@ class TestRunJobConfigLogging:
# / hit the network and have caused this test to time out on CI
# (>30s wall clock) under load. See PR #33661 follow-up.
with patch("cron.scheduler._hermes_home", tmp_path), \
patch("cron.scheduler._resolve_origin", return_value=None), \
patch("cron.scheduler_delivery._resolve_origin", return_value=None), \
patch("hermes_cli.env_loader.load_hermes_dotenv"), \
patch("hermes_cli.env_loader.reset_secret_source_cache"), \
patch("hermes_cli.runtime_provider.resolve_runtime_provider",
return_value={"provider": "openrouter", "api_key": "x",
"base_url": "https://example.invalid",
"api_mode": "chat_completions"}), \
patch("tools.mcp_tool.discover_mcp_tools", return_value=[]), \
patch("tools.mcp_tool_discovery.discover_mcp_tools", return_value=[]), \
patch("run_agent.AIAgent") as mock_agent_cls:
mock_agent = MagicMock()
mock_agent.run_conversation.return_value = {"final_response": "ok"}
@@ -1145,13 +1145,13 @@ class TestRunJobConfigEnvVarExpansion:
return {**self._RUNTIME, "provider": "xai", "api_mode": "chat_completions"}
with patch("cron.scheduler._hermes_home", tmp_path), \
patch("cron.scheduler._resolve_origin", return_value=None), \
patch("cron.scheduler_delivery._resolve_origin", return_value=None), \
patch("hermes_cli.env_loader.load_hermes_dotenv"), \
patch("hermes_cli.env_loader.reset_secret_source_cache"), \
patch("hermes_state.get_shared_session_db", return_value=fake_db), \
patch("hermes_state_registry.acquire", return_value=fake_db), \
patch("hermes_cli.runtime_provider.resolve_runtime_provider",
side_effect=resolve_runtime), \
patch("tools.mcp_tool.discover_mcp_tools", return_value=[]), \
patch("tools.mcp_tool_discovery.discover_mcp_tools", return_value=[]), \
patch("run_agent.AIAgent") as mock_agent_cls:
mock_agent = MagicMock()
mock_agent.run_conversation.return_value = {"final_response": "ok"}
@@ -1201,13 +1201,13 @@ class TestRunJobConfigEnvVarExpansion:
return {**self._RUNTIME, "provider": "openrouter"}
with patch("cron.scheduler._hermes_home", tmp_path), \
patch("cron.scheduler._resolve_origin", return_value=None), \
patch("cron.scheduler_delivery._resolve_origin", return_value=None), \
patch("hermes_cli.env_loader.load_hermes_dotenv"), \
patch("hermes_cli.env_loader.reset_secret_source_cache"), \
patch("hermes_state.get_shared_session_db", return_value=fake_db), \
patch("hermes_state_registry.acquire", return_value=fake_db), \
patch("hermes_cli.runtime_provider.resolve_runtime_provider",
side_effect=resolve_runtime), \
patch("tools.mcp_tool.discover_mcp_tools", return_value=[]), \
patch("tools.mcp_tool_discovery.discover_mcp_tools", return_value=[]), \
patch("run_agent.AIAgent") as mock_agent_cls:
mock_agent = MagicMock()
mock_agent.run_conversation.return_value = {"final_response": "ok"}
@@ -97,7 +97,7 @@ def _run_booked_job(monkeypatch, tmp_path):
monkeypatch.setattr(
"hermes_cli.runtime_provider.resolve_runtime_provider", _fake_runtime
)
monkeypatch.setattr("tools.mcp_tool.discover_mcp_tools", lambda: [])
monkeypatch.setattr("tools.mcp_tool_discovery.discover_mcp_tools", lambda: [])
monkeypatch.setattr(cron_scheduler, "_get_hermes_home", lambda: tmp_path)
monkeypatch.setattr(cron_scheduler, "get_fallback_chain", lambda _cfg: [])
monkeypatch.setattr(
@@ -110,7 +110,7 @@ def test_run_job_cron_execute_code_deny_does_not_pollute_later_gateway_execute_c
"args": None,
},
)
monkeypatch.setattr("tools.mcp_tool.discover_mcp_tools", lambda: [])
monkeypatch.setattr("tools.mcp_tool_discovery.discover_mcp_tools", lambda: [])
monkeypatch.setattr(cron_scheduler, "_get_hermes_home", lambda: tmp_path)
monkeypatch.setattr(cron_scheduler, "get_fallback_chain", lambda _cfg: [])
monkeypatch.setattr(
+2 -2
View File
@@ -43,8 +43,8 @@ def test_no_agent_cron_job_does_not_initialize_mcp():
# _run_job_script returns (ok, output); make it fail cleanly so we
# don't need a real script file.
with patch("tools.mcp_tool.discover_mcp_tools", side_effect=fake_discover), \
patch("cron.scheduler._run_job_script", return_value=(False, "no such file")):
with patch("tools.mcp_tool_discovery.discover_mcp_tools", side_effect=fake_discover), \
patch("cron.scheduler_script._run_job_script", return_value=(False, "no such file")):
scheduler.run_job(job)
assert not discover_called, (
+2 -2
View File
@@ -421,7 +421,7 @@ async def test_shutdown_mcp_servers_nonblocking_keeps_loop_responsive():
hb = asyncio.create_task(heartbeat())
try:
with patch("tools.mcp_tool.shutdown_mcp_servers", wedged_shutdown):
with patch("tools.mcp_tool_lifecycle.shutdown_mcp_servers", wedged_shutdown):
done = await asyncio.wait_for(
gateway_run._shutdown_mcp_servers_nonblocking(timeout=0.5),
timeout=5,
@@ -438,7 +438,7 @@ async def test_shutdown_mcp_servers_nonblocking_keeps_loop_responsive():
@pytest.mark.asyncio
async def test_shutdown_mcp_servers_nonblocking_completes_fast_path():
calls = []
with patch("tools.mcp_tool.shutdown_mcp_servers", lambda: calls.append(1)):
with patch("tools.mcp_tool_lifecycle.shutdown_mcp_servers", lambda: calls.append(1)):
done = await gateway_run._shutdown_mcp_servers_nonblocking(timeout=5)
assert done is True
assert calls == [1]
@@ -106,8 +106,8 @@ async def test_reload_mcp_refreshes_cached_agent_tools():
]
with (
patch("tools.mcp_tool.shutdown_mcp_servers"),
patch("tools.mcp_tool.discover_mcp_tools", return_value=["HassTurnOn", "HassTurnOff"]),
patch("tools.mcp_tool_lifecycle.shutdown_mcp_servers"),
patch("tools.mcp_tool_discovery.discover_mcp_tools", return_value=["HassTurnOn", "HassTurnOff"]),
patch.dict("tools.mcp_tool._servers", {"homeassistant": object()}, clear=True),
patch("model_tools.get_tool_definitions", return_value=fresh_tool_defs),
):
@@ -136,8 +136,8 @@ async def test_reload_mcp_handles_empty_agent_cache():
assert len(runner._agent_cache) == 0
with (
patch("tools.mcp_tool.shutdown_mcp_servers"),
patch("tools.mcp_tool.discover_mcp_tools", return_value=[]),
patch("tools.mcp_tool_lifecycle.shutdown_mcp_servers"),
patch("tools.mcp_tool_discovery.discover_mcp_tools", return_value=[]),
patch.dict("tools.mcp_tool._servers", {}, clear=True),
patch("model_tools.get_tool_definitions", return_value=[]),
):
@@ -164,8 +164,8 @@ async def test_reload_mcp_preserves_per_agent_toolset_overrides():
return [{"type": "function", "function": {"name": "refreshed"}}]
with (
patch("tools.mcp_tool.shutdown_mcp_servers"),
patch("tools.mcp_tool.discover_mcp_tools", return_value=["refreshed"]),
patch("tools.mcp_tool_lifecycle.shutdown_mcp_servers"),
patch("tools.mcp_tool_discovery.discover_mcp_tools", return_value=["refreshed"]),
patch.dict("tools.mcp_tool._servers", {"homeassistant": object()}, clear=True),
patch("model_tools.get_tool_definitions", side_effect=_capture_get_tool_definitions),
):
@@ -44,7 +44,7 @@ class TestMcpInterpolationUsesScope:
"""MCP config ${VAR} interpolation resolves through the secret scope."""
def test_interpolation_reads_scope(self, monkeypatch):
from tools.mcp_tool import _interpolate_env_vars
from tools.mcp_tool_config import _interpolate_env_vars
monkeypatch.setenv("MY_MCP_TOKEN", "global-token")
ss.set_multiplex_active(True)
tok = ss.set_secret_scope({"MY_MCP_TOKEN": "profile-token"})
@@ -20,7 +20,7 @@ async def test_gateway_boot_discovers_mcp_under_every_profile_home(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
import gateway.run as gateway_run
from tools import mcp_tool
from tools import mcp_tool_discovery as _mcp_discovery
homes = [("default", tmp_path / "default"), ("worker", tmp_path / "worker")]
for _name, home in homes:
@@ -35,7 +35,7 @@ async def test_gateway_boot_discovers_mcp_under_every_profile_home(
"hermes_cli.profiles.profiles_to_serve",
lambda multiplex, profile_allowlist=None: homes,
)
monkeypatch.setattr(mcp_tool, "discover_mcp_tools", fake_discover)
monkeypatch.setattr(_mcp_discovery, "discover_mcp_tools", fake_discover)
await gateway_run._discover_gateway_mcp_tools(GatewayConfig(multiplex_profiles=True))
@@ -50,6 +50,8 @@ async def test_reload_mcp_only_touches_requesting_profile(
) -> None:
from gateway.run import GatewayRunner
from tools import mcp_tool
from tools import mcp_tool_discovery as _mcp_discovery
from tools import mcp_tool_lifecycle as _mcp_lifecycle
worker_home = tmp_path / "profiles" / "worker"
worker_home.mkdir(parents=True)
@@ -78,8 +80,8 @@ async def test_reload_mcp_only_touches_requesting_profile(
seen.append(("discover", get_hermes_home()))
return []
monkeypatch.setattr(mcp_tool, "shutdown_mcp_servers", fake_shutdown)
monkeypatch.setattr(mcp_tool, "discover_mcp_tools", fake_discover)
monkeypatch.setattr(_mcp_lifecycle, "shutdown_mcp_servers", fake_shutdown)
monkeypatch.setattr(_mcp_discovery, "discover_mcp_tools", fake_discover)
event = MessageEvent(
text="/reload-mcp", message_id="m1",
+3 -2
View File
@@ -193,6 +193,7 @@ class TestBuildSessionContextPrompt:
from unittest.mock import patch
from gateway.session import _slack_tools_loaded
import tools.mcp_tool as _mcp_tool_mod
from tools import mcp_tool_registration as _mcp_registration
# No native slack toolset / token configured.
with patch.dict(_os.environ, {}, clear=False):
@@ -202,14 +203,14 @@ class TestBuildSessionContextPrompt:
# registered a real tool, via the actual tracking function used
# by the live registration path (tools/mcp_tool.py:_track_mcp_tool_server),
# not a mock of the capability check.
_mcp_tool_mod._track_mcp_tool_server("mcp-company-slack_post_message", "company-slack")
_mcp_registration._track_mcp_tool_server("mcp-company-slack_post_message", "company-slack")
try:
assert _slack_tools_loaded() is True, (
"A connected MCP server with 'slack' in its name and "
"registered tools must be detected as Slack capability"
)
finally:
_mcp_tool_mod._forget_mcp_tool_server("mcp-company-slack_post_message")
_mcp_registration._forget_mcp_tool_server("mcp-company-slack_post_message")
def test_shared_slack_prompt_warns_against_guessed_self_mentions(self):
+1 -1
View File
@@ -208,7 +208,7 @@ async def test_start_gateway_does_not_start_cron_after_aborted_startup(tmp_path,
monkeypatch.setattr("hermes_logging.setup_logging", lambda hermes_home, mode: None)
monkeypatch.setattr("gateway.run.GatewayRunner", AbortedStartupRunner)
monkeypatch.setattr("gateway.run._start_cron_ticker", fail_if_cron_starts)
monkeypatch.setattr("tools.mcp_tool.shutdown_mcp_servers", lambda: None)
monkeypatch.setattr("tools.mcp_tool_lifecycle.shutdown_mcp_servers", lambda: None)
with pytest.raises(SystemExit) as exc:
await gateway_run.start_gateway(config=GatewayConfig(), replace=False, verbosity=None)
+4 -3
View File
@@ -6,7 +6,7 @@ from rich.console import Console
import hermes_cli.banner as banner
import model_tools
import tools.mcp_tool
import tools.mcp_tool_discovery
def test_cprint_falls_back_to_plain_print_when_prompt_toolkit_has_no_console(capsys):
@@ -32,6 +32,7 @@ def test_build_welcome_banner_title_falls_back_when_no_tag():
import hermes_cli.banner as _banner
import model_tools as _mt
import tools.mcp_tool as _mcp
from tools import mcp_tool_discovery as _mcp_discovery
_banner._latest_release_cache = None
buf = io.StringIO()
@@ -39,7 +40,7 @@ def test_build_welcome_banner_title_falls_back_when_no_tag():
_patch.object(_mt, "check_tool_availability", return_value=(["web"], [])),
_patch.object(_banner, "get_available_skills", return_value={}),
_patch.object(_banner, "get_update_result", return_value=None),
_patch.object(_mcp, "get_mcp_status", return_value=[]),
_patch.object(_mcp_discovery, "get_mcp_status", return_value=[]),
_patch.object(_banner, "get_latest_release_tag", return_value=None),
):
console = Console(file=buf, force_terminal=True, color_system="truecolor", width=160)
@@ -68,7 +69,7 @@ def test_build_welcome_banner_non_moa_unchanged(tmp_path, monkeypatch):
patch.object(model_tools, "check_tool_availability", return_value=([], [])),
patch.object(banner, "get_available_skills", return_value={}),
patch.object(banner, "get_update_result", return_value=None),
patch.object(tools.mcp_tool, "get_mcp_status", return_value=[]),
patch.object(tools.mcp_tool_discovery, "get_mcp_status", return_value=[]),
):
console = Console(record=True, force_terminal=False, color_system=None, width=160)
banner.build_welcome_banner(
+2 -2
View File
@@ -7,7 +7,7 @@ from rich.console import Console
import hermes_cli.banner as banner
import model_tools
import tools.mcp_tool
import tools.mcp_tool_discovery
def _build_banner_with_skills(skills_by_category, term_width=160):
@@ -20,7 +20,7 @@ def _build_banner_with_skills(skills_by_category, term_width=160):
),
patch.object(banner, "get_available_skills", return_value=skills_by_category),
patch.object(banner, "get_update_result", return_value=None),
patch.object(tools.mcp_tool, "get_mcp_status", return_value=[]),
patch.object(tools.mcp_tool_discovery, "get_mcp_status", return_value=[]),
patch("shutil.get_terminal_size", return_value=os.terminal_size((term_width, 50))),
):
console = Console(
+1 -1
View File
@@ -792,7 +792,7 @@ class TestToolsConfigIncludeMode:
import hermes_cli.tools_config as tc
# Mock the probe to return three tools
monkeypatch.setattr(
"tools.mcp_tool.probe_mcp_server_tools",
"tools.mcp_tool_discovery.probe_mcp_server_tools",
lambda: {"demo": [("a", "desc"), ("b", "desc"), ("c", "desc")]},
)
# Mock the checklist to return just the first tool
+25 -20
View File
@@ -10,6 +10,7 @@ import os
from pathlib import Path
import pytest
from tools import mcp_tool_config as _mcp_config
def _set_interactive_stdin(monkeypatch, *, is_tty: bool = True) -> None:
@@ -307,6 +308,9 @@ class TestMcpTest:
import asyncio
from hermes_cli import mcp_config
import tools.mcp_tool as mcp_tool
from tools import mcp_tool_discovery as _mcp_discovery
from tools import mcp_tool_lifecycle as _mcp_lifecycle
from tools import mcp_tool_loop as _mcp_loop
captured = {}
@@ -327,10 +331,10 @@ class TestMcpTest:
captured["inner_timeout"] = timeout
return await awaitable
monkeypatch.setattr(mcp_tool, "_ensure_mcp_loop", lambda: None)
monkeypatch.setattr(mcp_tool, "_stop_mcp_loop_if_idle", lambda: None)
monkeypatch.setattr(mcp_tool, "_connect_server", fake_connect)
monkeypatch.setattr(mcp_tool, "_run_on_mcp_loop", fake_run_on_mcp_loop)
monkeypatch.setattr(_mcp_loop, "_ensure_mcp_loop", lambda: None)
monkeypatch.setattr(_mcp_lifecycle, "_stop_mcp_loop_if_idle", lambda: None)
monkeypatch.setattr(_mcp_discovery, "_connect_server", fake_connect)
monkeypatch.setattr(_mcp_loop, "_run_on_mcp_loop", fake_run_on_mcp_loop)
monkeypatch.setattr(mcp_config.asyncio, "wait_for", fake_wait_for)
assert mcp_config._probe_single_server(
@@ -351,13 +355,13 @@ class TestEnvVarInterpolation:
def test_interpolate_cursor_env_prefix(self, monkeypatch):
"""Cursor-style ${env:VAR} resolves the same secret as ${VAR}."""
monkeypatch.setenv("MY_KEY", "secret123")
from tools.mcp_tool import _interpolate_env_vars
from tools.mcp_tool_config import _interpolate_env_vars
assert _interpolate_env_vars("Bearer ${env:MY_KEY}") == "Bearer secret123"
def test_env_ref_name_strips_prefix(self):
from tools.mcp_tool import _env_ref_name
from tools.mcp_tool_common import _env_ref_name
assert _env_ref_name("env:API_KEY") == "API_KEY"
assert _env_ref_name("API_KEY") == "API_KEY"
@@ -371,14 +375,14 @@ class TestContextVarInterpolation:
def test_user_home(self):
import os
from tools.mcp_tool import _interpolate_env_vars
from tools.mcp_tool_config import _interpolate_env_vars
assert _interpolate_env_vars("${userHome}") == os.path.expanduser("~")
def test_path_separator_and_slash_shorthand(self):
import os
from tools.mcp_tool import _interpolate_env_vars
from tools.mcp_tool_config import _interpolate_env_vars
assert _interpolate_env_vars("${pathSeparator}") == os.sep
assert _interpolate_env_vars("${/}") == os.sep
@@ -387,12 +391,12 @@ class TestContextVarInterpolation:
import tools.mcp_tool as mcp_tool
monkeypatch.setattr(
mcp_tool, "_workspace_folder", lambda: "/srv/projects/myapp"
_mcp_config, "_workspace_folder", lambda: "/srv/projects/myapp"
)
assert mcp_tool._interpolate_env_vars("${workspaceFolder}") == (
assert _mcp_config._interpolate_env_vars("${workspaceFolder}") == (
"/srv/projects/myapp"
)
assert mcp_tool._interpolate_env_vars(
assert _mcp_config._interpolate_env_vars(
"${workspaceFolderBasename}"
) == "myapp"
@@ -400,7 +404,7 @@ class TestContextVarInterpolation:
import os
import tools.file_tools_paths as file_tools_paths
from tools.mcp_tool import _workspace_folder
from tools.mcp_tool_config import _workspace_folder
monkeypatch.setattr(
file_tools_paths, "_authoritative_workspace_root", lambda task_id="default": None
@@ -413,8 +417,8 @@ class TestContextVarInterpolation:
import tools.mcp_tool as mcp_tool
monkeypatch.setenv("MY_TOKEN", "tok-1")
monkeypatch.setattr(mcp_tool, "_workspace_folder", lambda: "/ws/app")
result = mcp_tool._interpolate_env_vars(
monkeypatch.setattr(_mcp_config, "_workspace_folder", lambda: "/ws/app")
result = _mcp_config._interpolate_env_vars(
"${userHome}${/}.cache${/}${workspaceFolderBasename}-${MY_TOKEN}"
)
home = os.path.expanduser("~")
@@ -424,13 +428,13 @@ class TestContextVarInterpolation:
"""${USERHOME} is NOT a context var — it keeps env-var semantics
(literal placeholder when unset)."""
monkeypatch.delenv("USERHOME", raising=False)
from tools.mcp_tool import _interpolate_env_vars
from tools.mcp_tool_config import _interpolate_env_vars
assert _interpolate_env_vars("${USERHOME}") == "${USERHOME}"
def test_unknown_ref_keeps_literal_placeholder(self, monkeypatch):
monkeypatch.delenv("NOT_A_REAL_VAR_XYZ", raising=False)
from tools.mcp_tool import _interpolate_env_vars
from tools.mcp_tool_config import _interpolate_env_vars
assert _interpolate_env_vars("${NOT_A_REAL_VAR_XYZ}") == (
"${NOT_A_REAL_VAR_XYZ}"
@@ -440,8 +444,9 @@ class TestContextVarInterpolation:
import os
import tools.mcp_tool as mcp_tool
from tools import mcp_tool_config as _mcp_config
monkeypatch.setattr(mcp_tool, "_workspace_folder", lambda: "/ws/app")
monkeypatch.setattr(_mcp_config, "_workspace_folder", lambda: "/ws/app")
cfg = {
"command": "npx",
"args": ["-y", "server-fs", "${workspaceFolder}"],
@@ -449,7 +454,7 @@ class TestContextVarInterpolation:
"env": {"CACHE": "${userHome}${/}.cache"},
"headers": {"X-Ws": "${workspaceFolderBasename}"},
}
out = mcp_tool._interpolate_env_vars(cfg)
out = _mcp_config._interpolate_env_vars(cfg)
home = os.path.expanduser("~")
assert out["args"][2] == "/ws/app"
assert out["cwd"] == "/ws/app"
@@ -518,7 +523,7 @@ class TestProbeEnvResolution:
seen["config"] = config
return _FakeServer()
monkeypatch.setattr("tools.mcp_tool._connect_server", _fake_connect)
monkeypatch.setattr("tools.mcp_tool_discovery._connect_server", _fake_connect)
tools = mc._probe_single_server("n8n", {
"url": "http://localhost:5678/mcp-server/http",
@@ -591,7 +596,7 @@ class TestProbeCapabilityGating:
async def _fake_connect(name, cfg):
return self._make_server(called, caps)
monkeypatch.setattr("tools.mcp_tool._connect_server", _fake_connect)
monkeypatch.setattr("tools.mcp_tool_discovery._connect_server", _fake_connect)
details: dict = {}
mc._probe_single_server("srv", config, details=details)
return called, details
@@ -106,7 +106,7 @@ def _stub_mcp_modules(monkeypatch):
)
monkeypatch.setitem(
sys.modules,
"tools.mcp_tool",
"tools.mcp_tool_discovery",
types.SimpleNamespace(
discover_mcp_tools=lambda: None,
get_mcp_status=lambda: [{"connected": True}],
+6 -4
View File
@@ -82,9 +82,11 @@ def test_validator_flags_ssh_key_persistence_payload():
def test_explicit_registration_skips_dangerous_entry_before_connect(monkeypatch):
import tools.mcp_tool as mcp_tool
from tools import mcp_tool_discovery as _mcp_discovery
from tools import mcp_tool_loop as _mcp_loop
monkeypatch.setattr(mcp_tool, "_MCP_AVAILABLE", True)
monkeypatch.setattr(mcp_tool, "_ensure_mcp_loop", lambda: None)
monkeypatch.setattr(_mcp_loop, "_ensure_mcp_loop", lambda: None)
connected = []
@@ -99,8 +101,8 @@ def test_explicit_registration_skips_dangerous_entry_before_connect(monkeypatch)
assert inspect.iscoroutine(coro)
return asyncio.run(coro)
monkeypatch.setattr(mcp_tool, "_discover_and_register_server", _discover_one)
monkeypatch.setattr(mcp_tool, "_run_on_mcp_loop", _run_on_loop)
monkeypatch.setattr(_mcp_discovery, "_discover_and_register_server", _discover_one)
monkeypatch.setattr(_mcp_loop, "_run_on_mcp_loop", _run_on_loop)
with mcp_tool._lock:
saved_servers = dict(mcp_tool._servers)
@@ -111,7 +113,7 @@ def test_explicit_registration_skips_dangerous_entry_before_connect(monkeypatch)
mcp_tool._server_connect_errors.clear()
try:
mcp_tool.register_mcp_servers({
_mcp_discovery.register_mcp_servers({
"evil": _dangerous_entry(),
"clean": {"command": "npx", "args": ["-y", "clean-mcp"]},
})
+15 -12
View File
@@ -83,7 +83,7 @@ def test_prepare_agent_startup_backgrounds_blocking_mcp_for_chat(monkeypatch):
)
monkeypatch.setitem(
sys.modules,
"tools.mcp_tool",
"tools.mcp_tool_discovery",
types.SimpleNamespace(discover_mcp_tools=_blocking_discover),
)
@@ -139,7 +139,7 @@ def test_prepare_agent_startup_skips_discovery_when_chat_resolves_to_tui(
)
monkeypatch.setitem(
sys.modules,
"tools.mcp_tool",
"tools.mcp_tool_discovery",
types.SimpleNamespace(
discover_mcp_tools=lambda: calls.__setitem__("inline", calls["inline"] + 1),
),
@@ -181,7 +181,7 @@ def test_prepare_agent_startup_keeps_discovery_for_non_chat_commands(
)
monkeypatch.setitem(
sys.modules,
"tools.mcp_tool",
"tools.mcp_tool_discovery",
types.SimpleNamespace(
discover_mcp_tools=lambda: calls.__setitem__("inline", calls["inline"] + 1),
),
@@ -221,7 +221,7 @@ def test_background_mcp_discovery_suppresses_interactive_oauth(monkeypatch):
)
monkeypatch.setitem(
sys.modules,
"tools.mcp_tool",
"tools.mcp_tool_discovery",
types.SimpleNamespace(discover_mcp_tools=_discover),
)
@@ -281,7 +281,7 @@ def _install_retry_stubs(monkeypatch, *, connected: bool, calls: dict):
)
monkeypatch.setitem(
sys.modules,
"tools.mcp_tool",
"tools.mcp_tool_discovery",
types.SimpleNamespace(
discover_mcp_tools=lambda: calls.__setitem__("mcp", calls["mcp"] + 1),
get_mcp_status=lambda: [{"connected": connected}],
@@ -323,6 +323,9 @@ def test_discover_mcp_tools_spawns_only_allowed_servers(monkeypatch):
"""The filter must narrow the spawn set before any server is connected;
built-in toolset names in the list are ignored."""
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 import mcp_tool_loop as _mcp_loop
servers = {
"code-mcp": {"command": "true"},
@@ -336,33 +339,33 @@ def test_discover_mcp_tools_spawns_only_allowed_servers(monkeypatch):
sdk_probes["n"] += 1
return True
monkeypatch.setattr(mcp_tool, "_load_mcp_config", lambda: dict(servers))
monkeypatch.setattr(_mcp_config, "_load_mcp_config", lambda: dict(servers))
monkeypatch.setattr(mcp_tool, "_ensure_mcp_sdk", _fake_ensure_sdk)
monkeypatch.setattr(mcp_tool, "_try_acquire_mcp_discovery_lock", lambda: mcp_tool._LOCK_UNAVAILABLE)
monkeypatch.setattr(_mcp_loop, "_try_acquire_mcp_discovery_lock", lambda: mcp_tool._LOCK_UNAVAILABLE)
monkeypatch.setattr(mcp_tool, "_release_mcp_discovery_lock", lambda *_a, **_k: None, raising=False)
def _fake_register(cfgs):
seen.update(cfgs)
return []
monkeypatch.setattr(mcp_tool, "register_mcp_servers", _fake_register)
monkeypatch.setattr(_mcp_discovery, "register_mcp_servers", _fake_register)
monkeypatch.setattr(mcp_tool, "_servers", {})
monkeypatch.setattr(mcp_tool, "_server_connecting", set())
# Everything (no filter) — both would be registered.
mcp_tool.discover_mcp_tools()
_mcp_discovery.discover_mcp_tools()
assert set(seen) == {"code-mcp", "docs-mcp"}
# `-t terminal,code-mcp` — only the matching server; "terminal" is a no-op.
seen.clear()
mcp_tool.discover_mcp_tools(allowed_mcp_names=["terminal", "code-mcp"])
_mcp_discovery.discover_mcp_tools(allowed_mcp_names=["terminal", "code-mcp"])
assert set(seen) == {"code-mcp"}
# `-t terminal` — no MCP server in the filter: skip the whole MCP load,
# including the ~260ms `mcp` SDK import.
seen.clear()
sdk_probes["n"] = 0
assert mcp_tool.discover_mcp_tools(allowed_mcp_names=["terminal"]) == []
assert _mcp_discovery.discover_mcp_tools(allowed_mcp_names=["terminal"]) == []
assert seen == {}
assert sdk_probes["n"] == 0
@@ -371,7 +374,7 @@ def test_background_discovery_honors_server_filter(monkeypatch, _reset_mcp_serve
calls: list = []
monkeypatch.setitem(
sys.modules,
"tools.mcp_tool",
"tools.mcp_tool_discovery",
types.SimpleNamespace(discover_mcp_tools=lambda allowed_mcp_names=None: calls.append(allowed_mcp_names)),
)
monkeypatch.setitem(
+1 -1
View File
@@ -5,7 +5,7 @@ from unittest.mock import patch
from hermes_cli.tools_config import _configure_mcp_tools_interactive
# Patch targets: imports happen inside the function body, so patch at source
_PROBE = "tools.mcp_tool.probe_mcp_server_tools"
_PROBE = "tools.mcp_tool_discovery.probe_mcp_server_tools"
_CHECKLIST = "hermes_cli.curses_ui.curses_checklist"
_SAVE = "hermes_cli.tools_config.save_config"
+3 -2
View File
@@ -31,8 +31,9 @@ def _patch_config(monkeypatch, entries: dict) -> None:
def _patch_handler(monkeypatch, response: str, captured: dict | None = None):
"""Replace tools.mcp_tool._make_tool_handler with a transport mock."""
"""Replace tools.mcp_tool_handlers._make_tool_handler with a transport mock."""
import tools.mcp_tool as mcp_mod
from tools import mcp_tool_handlers as _mcp_handlers
def _fake_make_handler(server_name, tool_name, tool_timeout):
if captured is not None:
@@ -47,7 +48,7 @@ def _patch_handler(monkeypatch, response: str, captured: dict | None = None):
return _handler
monkeypatch.setattr(mcp_mod, "_make_tool_handler", _fake_make_handler)
monkeypatch.setattr(_mcp_handlers, "_make_tool_handler", _fake_make_handler)
# ---------------------------------------------------------------------------
+3 -3
View File
@@ -38,11 +38,11 @@ class TestContentAwareRefresh(unittest.TestCase):
return agent
def _refresh(self, agent, new_desc, **kw):
from tools.mcp_tool import refresh_agent_mcp_tools
from tools.mcp_tool_agent import refresh_agent_mcp_tools
with patch("model_tools.get_tool_definitions",
return_value=_defs(new_desc)), \
patch("tools.mcp_tool._reinject_post_build_tools",
patch("tools.mcp_tool_agent._reinject_post_build_tools",
return_value=set()):
return refresh_agent_mcp_tools(agent, **kw)
@@ -71,7 +71,7 @@ class TestCompactionWiring(unittest.TestCase):
from agent.conversation_compression import _refresh_agent_tool_definitions
agent = _Agent()
with patch("tools.mcp_tool.refresh_agent_mcp_tools",
with patch("tools.mcp_tool_agent.refresh_agent_mcp_tools",
return_value={"newly_added"}) as m:
changed = _refresh_agent_tool_definitions(agent)
self.assertTrue(changed)
+2 -2
View File
@@ -12151,9 +12151,9 @@ def test_session_info_includes_mcp_servers(monkeypatch):
{"name": "filesystem", "transport": "stdio", "tools": 4, "connected": True},
{"name": "broken", "transport": "stdio", "tools": 0, "connected": False},
]
fake_mod = types.ModuleType("tools.mcp_tool")
fake_mod = types.ModuleType("tools.mcp_tool_discovery")
fake_mod.get_mcp_status = lambda: fake_status
monkeypatch.setitem(sys.modules, "tools.mcp_tool", fake_mod)
monkeypatch.setitem(sys.modules, "tools.mcp_tool_discovery", fake_mod)
info = server._session_info(types.SimpleNamespace(tools=[], model="", provider="openai-codex"))
+22 -20
View File
@@ -17,6 +17,8 @@ from unittest.mock import patch
import pytest
import tools.mcp_tool as mcp_mod
from tools import mcp_tool_discovery as _mcp_discovery
from tools import mcp_tool_lifecycle as _mcp_lifecycle
@pytest.fixture(autouse=True)
@@ -49,9 +51,9 @@ class TestConnectCooldownHelpers:
def test_failure_arms_exponential_backoff(self):
now = 1000.0
with patch("tools.mcp_tool.time.monotonic", return_value=now):
mcp_mod._record_connect_failure("bad")
_mcp_discovery._record_connect_failure("bad")
d1 = mcp_mod._server_connect_retry_after["bad"]
mcp_mod._record_connect_failure("bad")
_mcp_discovery._record_connect_failure("bad")
d2 = mcp_mod._server_connect_retry_after["bad"]
assert d1 == now + mcp_mod._CONNECT_RETRY_BASE_BACKOFF_SEC
assert d2 == now + mcp_mod._CONNECT_RETRY_BASE_BACKOFF_SEC * 2
@@ -59,7 +61,7 @@ class TestConnectCooldownHelpers:
def test_unknown_server_not_in_cooldown(self):
assert mcp_mod._connect_cooldown_active("never-seen") is False
assert _mcp_discovery._connect_cooldown_active("never-seen") is False
@pytest.mark.skipif(not mcp_mod._MCP_AVAILABLE, reason="mcp SDK not installed")
@@ -76,7 +78,7 @@ class TestRegisterMcpServersIsolation:
server._tools = []
return server
return patch("tools.mcp_tool._connect_server", side_effect=fake_connect)
return patch("tools.mcp_tool_discovery._connect_server", side_effect=fake_connect)
def test_failing_server_skipped_on_second_pass(self):
attempts = []
@@ -85,16 +87,16 @@ class TestRegisterMcpServersIsolation:
"bad": {"command": "bad-cmd"},
}
with self._run_with_mocked_connect(attempts), \
patch("tools.mcp_tool._register_server_tools", return_value=[]), \
patch("tools.mcp_tool._filter_suspicious_mcp_servers", side_effect=lambda x: x):
mcp_mod.register_mcp_servers(cfg)
patch("tools.mcp_tool_registration._register_server_tools", return_value=[]), \
patch("tools.mcp_tool_config._filter_suspicious_mcp_servers", side_effect=lambda x: x):
_mcp_discovery.register_mcp_servers(cfg)
assert "good" in mcp_mod._servers
assert "bad" not in mcp_mod._servers
assert mcp_mod._connect_cooldown_active("bad") is True
assert _mcp_discovery._connect_cooldown_active("bad") is True
assert "bad" in attempts
attempts.clear()
mcp_mod.register_mcp_servers(cfg)
_mcp_discovery.register_mcp_servers(cfg)
assert "bad" not in attempts, (
"failing server was re-spawned despite active cooldown -- "
"restart storm not isolated (#50394)"
@@ -104,14 +106,14 @@ class TestRegisterMcpServersIsolation:
attempts = []
cfg = {"bad": {"command": "bad-cmd"}}
with self._run_with_mocked_connect(attempts), \
patch("tools.mcp_tool._register_server_tools", return_value=[]), \
patch("tools.mcp_tool._filter_suspicious_mcp_servers", side_effect=lambda x: x):
mcp_mod.register_mcp_servers(cfg)
assert mcp_mod._connect_cooldown_active("bad") is True
patch("tools.mcp_tool_registration._register_server_tools", return_value=[]), \
patch("tools.mcp_tool_config._filter_suspicious_mcp_servers", side_effect=lambda x: x):
_mcp_discovery.register_mcp_servers(cfg)
assert _mcp_discovery._connect_cooldown_active("bad") is True
mcp_mod._server_connect_retry_after["bad"] = mcp_mod.time.monotonic() - 1
attempts.clear()
mcp_mod.register_mcp_servers(cfg)
_mcp_discovery.register_mcp_servers(cfg)
assert "bad" in attempts, "elapsed cooldown should permit a retry"
@@ -126,18 +128,18 @@ class TestShutdownClearsCooldownState:
"""
def test_fast_path_clears_cooldown_state(self):
mcp_mod._record_connect_failure("bad")
_mcp_discovery._record_connect_failure("bad")
assert mcp_mod._server_connect_retry_after
assert not mcp_mod._servers # precondition: fast path taken
with patch("tools.mcp_tool._stop_mcp_loop"):
mcp_mod.shutdown_mcp_servers()
with patch("tools.mcp_tool_loop._stop_mcp_loop"):
_mcp_lifecycle.shutdown_mcp_servers()
assert mcp_mod._server_connect_retry_after == {}
assert mcp_mod._server_connect_failures == {}
def test_loop_not_running_path_clears_cooldown_state(self):
mcp_mod._record_connect_failure("bad")
_mcp_discovery._record_connect_failure("bad")
class _DeadServer:
name = "dead"
@@ -148,8 +150,8 @@ class TestShutdownClearsCooldownState:
mcp_mod._servers["dead"] = _DeadServer() # type: ignore[assignment]
# _mcp_loop is None in this test process, so the async _shutdown
# coroutine is never scheduled; only the final sweep can clear.
with patch("tools.mcp_tool._stop_mcp_loop"):
mcp_mod.shutdown_mcp_servers()
with patch("tools.mcp_tool_loop._stop_mcp_loop"):
_mcp_lifecycle.shutdown_mcp_servers()
assert mcp_mod._server_connect_retry_after == {}
assert mcp_mod._server_connect_failures == {}
+2 -2
View File
@@ -201,12 +201,12 @@ class TestMethodNotFoundDetection:
"""``_is_method_not_found_error`` underpins the ping→list_tools fallback."""
def test_structural_code_match(self):
from tools.mcp_tool import _is_method_not_found_error
from tools.mcp_tool_errors import _is_method_not_found_error
assert _is_method_not_found_error(_mcp_error(-32601)) is True
def test_unrelated_exception_is_not_match(self):
from tools.mcp_tool import _is_method_not_found_error
from tools.mcp_tool_errors import _is_method_not_found_error
assert _is_method_not_found_error(TimeoutError()) is False
assert _is_method_not_found_error(Exception("session terminated")) is False
+12 -9
View File
@@ -18,6 +18,7 @@ import pytest
pytest.importorskip("mcp.client.auth.oauth2")
from tools import mcp_tool_loop as _mcp_loop # noqa: E402
# ---------------------------------------------------------------------------
@@ -109,7 +110,7 @@ def test_circuit_breaker_half_opens_after_cooldown(monkeypatch, tmp_path):
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
from tools import mcp_tool
from tools.mcp_tool import _make_tool_handler
from tools.mcp_tool_handlers import _make_tool_handler
call_count = {"n": 0}
@@ -124,7 +125,7 @@ def test_circuit_breaker_half_opens_after_cooldown(monkeypatch, tmp_path):
return result
_install_stub_server(mcp_tool, "srv", _call_tool_success)
mcp_tool._ensure_mcp_loop()
_mcp_loop._ensure_mcp_loop()
try:
# Trip the breaker by setting the count at/above threshold and
@@ -177,7 +178,7 @@ def test_circuit_breaker_reopens_on_probe_failure(monkeypatch, tmp_path):
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
from tools import mcp_tool
from tools.mcp_tool import _make_tool_handler
from tools.mcp_tool_handlers import _make_tool_handler
call_count = {"n": 0}
@@ -186,7 +187,7 @@ def test_circuit_breaker_reopens_on_probe_failure(monkeypatch, tmp_path):
raise RuntimeError("still broken")
_install_stub_server(mcp_tool, "srv", _call_tool_fails)
mcp_tool._ensure_mcp_loop()
_mcp_loop._ensure_mcp_loop()
try:
mcp_tool._server_error_counts["srv"] = mcp_tool._CIRCUIT_BREAKER_THRESHOLD
@@ -234,7 +235,7 @@ def test_half_open_probe_on_dead_session_requests_reconnect(monkeypatch, tmp_pat
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
from tools import mcp_tool
from tools.mcp_tool import _make_tool_handler
from tools.mcp_tool_handlers import _make_tool_handler
server = _install_stub_server(mcp_tool, "srv", None)
# Simulate a dead/parked transport: no live session.
@@ -275,7 +276,7 @@ def test_half_open_dead_session_recovers_after_reconnect(monkeypatch, tmp_path):
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
from tools import mcp_tool
from tools.mcp_tool import _make_tool_handler
from tools.mcp_tool_handlers import _make_tool_handler
async def _call_tool_success(*a, **kw):
result = MagicMock()
@@ -289,7 +290,7 @@ def test_half_open_dead_session_recovers_after_reconnect(monkeypatch, tmp_path):
server = _install_stub_server(mcp_tool, "srv", _call_tool_success)
server.session = None # transport down at first
monkeypatch.setattr(mcp_tool, "_mcp_loop", None)
mcp_tool._ensure_mcp_loop()
_mcp_loop._ensure_mcp_loop()
try:
mcp_tool._server_error_counts["srv"] = mcp_tool._CIRCUIT_BREAKER_THRESHOLD
@@ -336,6 +337,8 @@ def test_circuit_breaker_cleared_on_reconnect(monkeypatch, tmp_path):
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
from tools import mcp_tool
from tools import mcp_tool_handlers as _mcp_handlers
from tools import mcp_tool_loop as _mcp_loop
from tools.mcp_oauth_manager import get_manager, reset_manager_for_tests
from mcp.client.auth import OAuthFlowError
@@ -345,7 +348,7 @@ def test_circuit_breaker_cleared_on_reconnect(monkeypatch, tmp_path):
raise AssertionError("session.call_tool should not be reached in this test")
_install_stub_server(mcp_tool, "srv", _call_tool_unused)
mcp_tool._ensure_mcp_loop()
_mcp_loop._ensure_mcp_loop()
# Open the breaker well above threshold, with a recent open-time so
# it would short-circuit everything without a reset.
@@ -371,7 +374,7 @@ def test_circuit_breaker_cleared_on_reconnect(monkeypatch, tmp_path):
def _retry_call():
raise OAuthFlowError("still failing post-reconnect")
result = mcp_tool._handle_auth_error_and_retry(
result = _mcp_handlers._handle_auth_error_and_retry(
"srv",
OAuthFlowError("initial"),
_retry_call,
+4 -4
View File
@@ -41,13 +41,13 @@ def _patch_sdk_async_client(dummy):
class TestResolveClientCert:
def test_returns_none_when_unset(self):
from tools.mcp_tool import _resolve_client_cert
from tools.mcp_tool_errors import _resolve_client_cert
assert _resolve_client_cert("srv", {}) is None
assert _resolve_client_cert("srv", {"url": "https://x"}) is None
def test_string_form_single_pem(self, tmp_path):
from tools.mcp_tool import _resolve_client_cert
from tools.mcp_tool_errors import _resolve_client_cert
pem = tmp_path / "combined.pem"
pem.write_text("dummy")
@@ -57,7 +57,7 @@ class TestResolveClientCert:
def test_list_form_two_elements(self, tmp_path):
from tools.mcp_tool import _resolve_client_cert
from tools.mcp_tool_errors import _resolve_client_cert
cert = tmp_path / "client.crt"
key = tmp_path / "client.key"
@@ -71,7 +71,7 @@ class TestResolveClientCert:
def test_password_must_be_string(self, tmp_path):
from tools.mcp_tool import _resolve_client_cert
from tools.mcp_tool_errors import _resolve_client_cert
cert = tmp_path / "client.crt"
key = tmp_path / "client.key"
@@ -10,14 +10,15 @@ import logging
import pytest
from tools import mcp_tool
from tools.mcp_tool import _warn_hidden_whitespace
from tools import mcp_tool_config as _mcp_config
from tools.mcp_tool_config import _warn_hidden_whitespace
@pytest.fixture(autouse=True)
def _reset_dedupe():
mcp_tool._whitespace_warned.clear()
_mcp_config._whitespace_warned.clear()
yield
mcp_tool._whitespace_warned.clear()
_mcp_config._whitespace_warned.clear()
def test_clean_config_no_warnings(caplog):
@@ -116,7 +117,7 @@ def test_load_mcp_config_emits_warning(tmp_path, monkeypatch, caplog):
with mock_patch("hermes_cli.config.load_config",
return_value={"mcp_servers": servers}), \
caplog.at_level(logging.WARNING, logger="tools.mcp_tool"):
result = mcp_tool._load_mcp_config()
result = _mcp_config._load_mcp_config()
assert "pasted" in result
# Value passes through unmutated.
+12 -11
View File
@@ -20,6 +20,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from tools import mcp_death_supervisor, mcp_tool
from tools import mcp_tool_lifecycle as _mcp_lifecycle
pytestmark = pytest.mark.skipif(
os.name != "posix", reason="the supervisor is POSIX-only (process groups)"
@@ -617,10 +618,10 @@ def _stdio_connection(child_pid, fake_supervisor):
# First call is the pids_before baseline; the second reports our child
# as the newly spawned server.
patch(
"tools.mcp_tool._snapshot_child_pids",
"tools.mcp_tool_lifecycle._snapshot_child_pids",
side_effect=[set(), {child_pid}],
),
patch("tools.mcp_tool._write_stderr_log_header"),
patch("tools.mcp_tool_config._write_stderr_log_header"),
patch("tools.mcp_tool._get_mcp_stderr_log", return_value=None),
patch(
"tools.mcp_tool._spawn_death_supervisor",
@@ -713,19 +714,19 @@ def test_scoped_teardown_of_one_owner_keeps_the_other_owner_supervised(monkeypat
"""
fake = _FakeSupervisor()
monkeypatch.setattr(mcp_tool, "_spawn_death_supervisor", lambda: fake)
monkeypatch.setattr(mcp_tool.time, "sleep", lambda _s: None) # skip the SIGTERM grace wait
monkeypatch.setattr(_mcp_lifecycle.time, "sleep", lambda _s: None) # skip the SIGTERM grace wait
a = subprocess.Popen(_VICTIM, start_new_session=True)
b = subprocess.Popen(_VICTIM, start_new_session=True)
try:
pg_a, pg_b = os.getpgid(a.pid), os.getpgid(b.pid)
with mcp_tool._lock:
mcp_tool._stdio_pids[a.pid] = "profile-a"
mcp_tool._stdio_pids[b.pid] = "profile-b"
mcp_tool._stdio_pgids[a.pid] = pg_a
mcp_tool._stdio_pgids[b.pid] = pg_b
_mcp_lifecycle._stdio_pids[a.pid] = "profile-a"
_mcp_lifecycle._stdio_pids[b.pid] = "profile-b"
_mcp_lifecycle._stdio_pgids[a.pid] = pg_a
_mcp_lifecycle._stdio_pgids[b.pid] = pg_b
mcp_tool._update_death_supervisor("register", [pg_a, pg_b])
mcp_tool._kill_orphaned_mcp_children(include_active=True, server_name="profile-a")
_mcp_lifecycle._kill_orphaned_mcp_children(include_active=True, server_name="profile-a")
a.wait(timeout=10)
assert b.poll() is None, "scoped teardown of profile-a killed profile-b's server"
@@ -734,7 +735,7 @@ def test_scoped_teardown_of_one_owner_keeps_the_other_owner_supervised(monkeypat
"scoped teardown released the OTHER owner's group from the supervisor"
)
assert mcp_tool._supervised_pgids == {pg_b}
assert b.pid in mcp_tool._stdio_pids and b.pid in mcp_tool._stdio_pgids
assert b.pid in _mcp_lifecycle._stdio_pids and b.pid in _mcp_lifecycle._stdio_pgids
finally:
for p in (a, b):
_kill(p.pid)
@@ -744,8 +745,8 @@ def test_scoped_teardown_of_one_owner_keeps_the_other_owner_supervised(monkeypat
pass
with mcp_tool._lock:
for p in (a, b):
mcp_tool._stdio_pids.pop(p.pid, None)
mcp_tool._stdio_pgids.pop(p.pid, None)
_mcp_lifecycle._stdio_pids.pop(p.pid, None)
_mcp_lifecycle._stdio_pgids.pop(p.pid, None)
@pytest.mark.live_system_guard_bypass
@@ -53,6 +53,8 @@ def test_two_processes_each_complete_local_mcp_discovery(tmp_path):
sys.path.insert(0, repo_root)
import tools.mcp_tool as mcp_tool
from tools import mcp_tool_config as _mcp_config
from tools import mcp_tool_discovery as _mcp_discovery
ready = Path(ready_arg)
release = Path(release_arg)
@@ -71,7 +73,7 @@ def test_two_processes_each_complete_local_mcp_discovery(tmp_path):
"enabled": True,
}
}
mcp_tool._load_mcp_config = lambda: config
_mcp_config._load_mcp_config = lambda: config
def fake_register_mcp_servers(servers):
tool_name = "mcp__test_srv__ping"
@@ -89,10 +91,10 @@ def test_two_processes_each_complete_local_mcp_discovery(tmp_path):
return [tool_name]
mcp_tool.register_mcp_servers = fake_register_mcp_servers
_mcp_discovery.register_mcp_servers = fake_register_mcp_servers
started.write_text("1", encoding="utf-8")
result = mcp_tool.discover_mcp_tools()
result = _mcp_discovery.discover_mcp_tools()
server = mcp_tool._servers.get("test_srv")
output.write_text(
json.dumps(
+2 -1
View File
@@ -6,7 +6,8 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from tools.mcp_tool import MCPServerTask, _register_server_tools
from tools.mcp_tool import MCPServerTask
from tools.mcp_tool_registration import _register_server_tools
from tools.registry import ToolRegistry
+2 -5
View File
@@ -1,4 +1,4 @@
"""Tests for the MCP elicitation handler in tools.mcp_tool.
"""Tests for the MCP elicitation handler in tools.mcp_tool_sampling.
These tests exercise ElicitationHandler in isolation -- the underlying
approval system and the MCP transport layer are mocked, so no real MCP
@@ -18,10 +18,7 @@ pytest.importorskip("mcp.types")
from mcp.types import ElicitResult # noqa: E402 -- after importorskip
from tools.mcp_tool import ( # noqa: E402
ElicitationHandler,
_format_elicitation_schema_summary,
)
from tools.mcp_tool_sampling import ElicitationHandler, _format_elicitation_schema_summary # noqa: E402
def _form_params(message="please confirm", schema=None):
+1 -1
View File
@@ -8,7 +8,7 @@ Fix: ``_exc_str()`` falls back to ``repr(exc)`` when ``str(exc)`` is empty.
"""
from tools.mcp_tool import _exc_str, _sanitize_error
from tools.mcp_tool_common import _exc_str, _sanitize_error
# ---------------------------------------------------------------------------
@@ -13,13 +13,9 @@ import logging
import pytest
from tools.mcp_tool import (
InvalidMcpUrlError,
MCPServerTask,
NonMcpEndpointError,
_classify_mcp_failure,
_unwrap_exception_group,
)
from tools.mcp_tool_errors import (
InvalidMcpUrlError, NonMcpEndpointError, _classify_mcp_failure, _unwrap_exception_group)
from tools.mcp_tool import MCPServerTask
def _group(*excs, msg="unhandled errors in a TaskGroup") -> BaseExceptionGroup:
+8 -8
View File
@@ -43,13 +43,13 @@ import pytest
class TestResolveIdentityHeader:
def test_returns_none_when_unset(self):
from tools.mcp_tool import _resolve_identity_header
from tools.mcp_tool_errors import _resolve_identity_header
assert _resolve_identity_header("srv", {}) is None
assert _resolve_identity_header("srv", {"url": "https://x"}) is None
def test_static_mode_returns_name_value(self):
from tools.mcp_tool import _resolve_identity_header
from tools.mcp_tool_errors import _resolve_identity_header
result = _resolve_identity_header("srv", {
"identity_header": {
@@ -61,7 +61,7 @@ class TestResolveIdentityHeader:
assert result == ("X-User-Id", "alice")
def test_static_is_default_value_from(self):
from tools.mcp_tool import _resolve_identity_header
from tools.mcp_tool_errors import _resolve_identity_header
result = _resolve_identity_header("srv", {
"identity_header": {"name": "X-User-Id", "value": "bob"},
@@ -69,7 +69,7 @@ class TestResolveIdentityHeader:
assert result == ("X-User-Id", "bob")
def test_profile_mode_uses_active_profile_name(self):
from tools.mcp_tool import _resolve_identity_header
from tools.mcp_tool_errors import _resolve_identity_header
with patch(
"hermes_cli.profiles.get_active_profile_name",
@@ -84,7 +84,7 @@ class TestResolveIdentityHeader:
assert result == ("X-Hermes-Profile", "workbot")
def test_missing_name_warns_and_returns_none(self, caplog):
from tools.mcp_tool import _resolve_identity_header
from tools.mcp_tool_errors import _resolve_identity_header
with caplog.at_level(logging.WARNING):
result = _resolve_identity_header("srv", {
@@ -94,7 +94,7 @@ class TestResolveIdentityHeader:
assert any("identity_header" in r.message for r in caplog.records)
def test_static_missing_value_warns_and_returns_none(self, caplog):
from tools.mcp_tool import _resolve_identity_header
from tools.mcp_tool_errors import _resolve_identity_header
with caplog.at_level(logging.WARNING):
result = _resolve_identity_header("srv", {
@@ -104,7 +104,7 @@ class TestResolveIdentityHeader:
assert any("identity_header" in r.message for r in caplog.records)
def test_unknown_value_from_warns_and_returns_none(self, caplog):
from tools.mcp_tool import _resolve_identity_header
from tools.mcp_tool_errors import _resolve_identity_header
with caplog.at_level(logging.WARNING):
result = _resolve_identity_header("srv", {
@@ -118,7 +118,7 @@ class TestResolveIdentityHeader:
assert any("identity_header" in r.message for r in caplog.records)
def test_non_dict_config_warns_and_returns_none(self, caplog):
from tools.mcp_tool import _resolve_identity_header
from tools.mcp_tool_errors import _resolve_identity_header
with caplog.at_level(logging.WARNING):
result = _resolve_identity_header("srv", {
+7 -7
View File
@@ -35,7 +35,7 @@ def _png_bytes():
class TestMimeExtension:
def test_maps_jpeg_variants_to_jpg(self):
from tools.mcp_tool import _mcp_image_extension_for_mime_type
from tools.mcp_tool_content import _mcp_image_extension_for_mime_type
assert _mcp_image_extension_for_mime_type("image/jpeg") == ".jpg"
assert _mcp_image_extension_for_mime_type("image/jpg") == ".jpg"
assert _mcp_image_extension_for_mime_type("IMAGE/JPEG") == ".jpg"
@@ -43,7 +43,7 @@ class TestMimeExtension:
def test_unknown_defaults_to_png(self):
from tools.mcp_tool import _mcp_image_extension_for_mime_type
from tools.mcp_tool_content import _mcp_image_extension_for_mime_type
assert _mcp_image_extension_for_mime_type("") == ".png"
assert _mcp_image_extension_for_mime_type("image/unheard-of-format") == ".png"
@@ -53,7 +53,7 @@ class TestCacheMcpImageBlock:
"""A well-formed ImageContent block with valid PNG bytes caches
to the image dir and the helper returns a ``MEDIA:<path>`` tag."""
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
from tools.mcp_tool import _cache_mcp_image_block
from tools.mcp_tool_content import _cache_mcp_image_block
block = SimpleNamespace(
data=base64.b64encode(_png_bytes()).decode("ascii"),
@@ -76,7 +76,7 @@ class TestCacheMcpImageBlock:
def test_returns_empty_when_block_is_not_an_image(self, tmp_path, monkeypatch):
"""Non-image MIME types shouldn't trigger caching."""
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
from tools.mcp_tool import _cache_mcp_image_block
from tools.mcp_tool_content import _cache_mcp_image_block
block = SimpleNamespace(
data=base64.b64encode(b"some bytes").decode("ascii"),
@@ -86,7 +86,7 @@ class TestCacheMcpImageBlock:
def test_returns_empty_when_block_has_no_data(self, tmp_path, monkeypatch):
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
from tools.mcp_tool import _cache_mcp_image_block
from tools.mcp_tool_content import _cache_mcp_image_block
block = SimpleNamespace(data=None, mimeType="image/png")
assert _cache_mcp_image_block(block) == ""
@@ -97,7 +97,7 @@ class TestCacheMcpImageBlock:
``image/png`` is actually an HTML error page, the cache raises and
we log + drop rather than propagate."""
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
from tools.mcp_tool import _cache_mcp_image_block
from tools.mcp_tool_content import _cache_mcp_image_block
block = SimpleNamespace(
data=base64.b64encode(b"<html>error</html>").decode("ascii"),
@@ -108,7 +108,7 @@ class TestCacheMcpImageBlock:
def test_handles_jpeg(self, tmp_path, monkeypatch):
"""JPEG signature should also be accepted."""
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
from tools.mcp_tool import _cache_mcp_image_block
from tools.mcp_tool_content import _cache_mcp_image_block
# minimal JPEG SOI marker + filler
jpeg = b"\xff\xd8\xff\xe0" + b"\x00" * 100 + b"\xff\xd9"
@@ -6,10 +6,14 @@ import threading
from types import SimpleNamespace
import pytest
from tools import mcp_tool_discovery as _mcp_discovery
from tools import mcp_tool_lifecycle as _mcp_lifecycle
from tools import mcp_tool_loop as _mcp_loop
def _reset_mcp_state(mcp_tool) -> None:
mcp_tool.shutdown_mcp_servers()
from tools.mcp_tool_lifecycle import shutdown_mcp_servers
shutdown_mcp_servers()
with mcp_tool._lock:
mcp_tool._servers.clear()
mcp_tool._server_connecting.clear()
@@ -17,14 +21,16 @@ def _reset_mcp_state(mcp_tool) -> None:
def _cleanup_mcp_state(mcp_tool, extra_servers=()) -> None:
from tools.mcp_tool_lifecycle import shutdown_mcp_servers
from tools.mcp_tool_loop import _run_on_mcp_loop
with mcp_tool._lock:
loop = mcp_tool._mcp_loop
if loop is not None and loop.is_running():
for server in extra_servers:
task = getattr(server, "_task", None)
if task is not None and not task.done():
mcp_tool._run_on_mcp_loop(server.shutdown, timeout=5)
mcp_tool.shutdown_mcp_servers()
_run_on_mcp_loop(server.shutdown, timeout=5)
shutdown_mcp_servers()
with mcp_tool._lock:
mcp_tool._servers.clear()
mcp_tool._server_connecting.clear()
@@ -53,7 +59,7 @@ def test_initial_connect_failure_is_registry_owned_and_reaped(monkeypatch, tmp_p
monkeypatch.setattr(mcp_tool, "_MAX_INITIAL_CONNECT_RETRIES", 0)
monkeypatch.setattr(mcp_tool, "_PARKED_RETRY_INTERVAL", 3600)
real_stop = mcp_tool._stop_mcp_loop
real_stop = _mcp_loop._stop_mcp_loop
pending_at_stop = []
async def _pending_tasks():
@@ -66,14 +72,14 @@ def test_initial_connect_failure_is_registry_owned_and_reaped(monkeypatch, tmp_p
def _observed_stop(*, only_if_idle=False):
pending_at_stop.extend(
mcp_tool._run_on_mcp_loop(_pending_tasks, timeout=5)
_mcp_loop._run_on_mcp_loop(_pending_tasks, timeout=5)
)
return real_stop(only_if_idle=only_if_idle)
monkeypatch.setattr(mcp_tool, "_stop_mcp_loop", _observed_stop)
monkeypatch.setattr(_mcp_loop, "_stop_mcp_loop", _observed_stop)
try:
assert mcp_tool.register_mcp_servers({
assert _mcp_discovery.register_mcp_servers({
"initial-failure": {"command": "unused", "connect_timeout": 5}
}) == []
@@ -87,7 +93,7 @@ def test_initial_connect_failure_is_registry_owned_and_reaped(monkeypatch, tmp_p
assert server._task is not None
assert not server._task.done(), "recoverable initial failure was not parked"
mcp_tool.shutdown_mcp_servers()
_mcp_lifecycle.shutdown_mcp_servers()
assert pending_at_stop == [], (
"shutdown left MCP tasks pending at loop stop: "
@@ -98,7 +104,7 @@ def test_initial_connect_failure_is_registry_owned_and_reaped(monkeypatch, tmp_p
assert mcp_tool._mcp_loop is None
assert mcp_tool._mcp_thread is None
finally:
monkeypatch.setattr(mcp_tool, "_stop_mcp_loop", real_stop)
monkeypatch.setattr(_mcp_loop, "_stop_mcp_loop", real_stop)
_cleanup_mcp_state(mcp_tool, created)
@@ -164,7 +170,7 @@ def test_initial_connect_failure_revives_same_registered_server(monkeypatch, tmp
}
try:
assert mcp_tool.register_mcp_servers(config) == []
assert _mcp_discovery.register_mcp_servers(config) == []
assert len(created) == 1
server = created[0]
with mcp_tool._lock:
@@ -175,7 +181,7 @@ def test_initial_connect_failure_revives_same_registered_server(monkeypatch, tmp
assert not server._task.done()
backend_up.set()
mcp_tool.register_mcp_servers(config)
_mcp_discovery.register_mcp_servers(config)
assert revived.wait(timeout=5), "cached parked server did not revive"
assert len(created) == 1, "revival created a duplicate server task"
@@ -208,6 +214,8 @@ def test_initial_auth_failure_is_retained_and_reaped(monkeypatch, tmp_path):
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
from tools import mcp_tool
from tools import mcp_tool_errors as _mcp_errors
from tools import mcp_tool_lifecycle as _mcp_lifecycle
_reset_mcp_state(mcp_tool)
created = []
@@ -223,10 +231,10 @@ def test_initial_auth_failure_is_retained_and_reaped(monkeypatch, tmp_path):
monkeypatch.setattr(mcp_tool, "MCPServerTask", _AuthFailingServerTask)
monkeypatch.setattr(mcp_tool, "_MCP_AVAILABLE", True)
monkeypatch.setattr(mcp_tool, "_PARKED_RETRY_INTERVAL", 3600)
monkeypatch.setattr(mcp_tool, "_is_auth_error", lambda exc: True)
monkeypatch.setattr(_mcp_errors, "_is_auth_error", lambda exc: True)
try:
assert mcp_tool.register_mcp_servers({
assert _mcp_discovery.register_mcp_servers({
"auth-failure": {"command": "unused", "connect_timeout": 5}
}) == []
assert len(created) == 1
@@ -240,7 +248,7 @@ def test_initial_auth_failure_is_retained_and_reaped(monkeypatch, tmp_path):
mcp_tool._server_connect_errors["auth-failure"]
)
mcp_tool.shutdown_mcp_servers()
_mcp_lifecycle.shutdown_mcp_servers()
assert server._task.done()
finally:
_cleanup_mcp_state(mcp_tool, created)
@@ -251,6 +259,8 @@ def test_standalone_failed_connect_is_reaped_without_global_owner(monkeypatch, t
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
from tools import mcp_tool
from tools import mcp_tool_discovery as _mcp_discovery
from tools import mcp_tool_loop as _mcp_loop
_reset_mcp_state(mcp_tool)
created = []
@@ -266,12 +276,12 @@ def test_standalone_failed_connect_is_reaped_without_global_owner(monkeypatch, t
monkeypatch.setattr(mcp_tool, "MCPServerTask", _ProbeServerTask)
monkeypatch.setattr(mcp_tool, "_MAX_INITIAL_CONNECT_RETRIES", 0)
monkeypatch.setattr(mcp_tool, "_PARKED_RETRY_INTERVAL", 3600)
mcp_tool._ensure_mcp_loop()
_mcp_loop._ensure_mcp_loop()
try:
with pytest.raises(ConnectionError, match="probe target unavailable"):
mcp_tool._run_on_mcp_loop(
lambda: mcp_tool._connect_server(
_mcp_loop._run_on_mcp_loop(
lambda: _mcp_discovery._connect_server(
"probe-only", {"command": "unused"}
),
timeout=5,
+1 -4
View File
@@ -18,10 +18,7 @@ from __future__ import annotations
import pytest
from tools.mcp_tool import (
InvalidMcpUrlError,
_validate_remote_mcp_url,
)
from tools.mcp_tool_errors import InvalidMcpUrlError, _validate_remote_mcp_url
class TestValidUrlsAccepted:
+49 -44
View File
@@ -13,6 +13,11 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
import tools.mcp_tool as mcp
from tools import mcp_tool_discovery as _mcp_discovery
from tools import mcp_tool_handlers as _mcp_handlers
from tools import mcp_tool_loop as _mcp_loop
from tools import mcp_tool_registration as _mcp_registration
from tools import mcp_tool_schema as _mcp_schema
@pytest.fixture(autouse=True)
@@ -66,14 +71,14 @@ class TestLazyMcpRegistration:
patch("tools.mcp_schema_cache.config_fingerprint", return_value="abc"), \
patch("tools.mcp_schema_cache.get_cached_entry", return_value=_fake_cache_entry()), \
patch(
"tools.mcp_tool._register_from_cache_sync",
"tools.mcp_tool_registration._register_from_cache_sync",
return_value=["mcp_playwright_browser_navigate"],
) as mock_register, \
patch("tools.mcp_tool._discover_and_register_server", new_callable=AsyncMock) as mock_discover, \
patch("tools.mcp_tool._ensure_mcp_loop") as mock_loop, \
patch("tools.mcp_tool._run_on_mcp_loop") as mock_run:
patch("tools.mcp_tool_discovery._discover_and_register_server", new_callable=AsyncMock) as mock_discover, \
patch("tools.mcp_tool_loop._ensure_mcp_loop") as mock_loop, \
patch("tools.mcp_tool_loop._run_on_mcp_loop") as mock_run:
mcp.register_mcp_servers(config)
_mcp_discovery.register_mcp_servers(config)
mock_register.assert_called_once()
mock_discover.assert_not_called()
@@ -85,10 +90,10 @@ class TestLazyMcpRegistration:
with patch("tools.mcp_tool._MCP_AVAILABLE", True), \
patch("tools.mcp_schema_cache.config_fingerprint", return_value="abc"), \
patch("tools.mcp_schema_cache.get_cached_entry", return_value=None), \
patch("tools.mcp_tool._ensure_mcp_loop"), \
patch("tools.mcp_tool._run_on_mcp_loop") as mock_run:
patch("tools.mcp_tool_loop._ensure_mcp_loop"), \
patch("tools.mcp_tool_loop._run_on_mcp_loop") as mock_run:
mcp.register_mcp_servers(config)
_mcp_discovery.register_mcp_servers(config)
mock_run.assert_called_once()
@@ -96,10 +101,10 @@ class TestLazyMcpRegistration:
config = {"playwright": {"command": "npx", "args": []}}
with patch("tools.mcp_tool._MCP_AVAILABLE", True), \
patch("tools.mcp_schema_cache.get_cached_entry") as mock_get, \
patch("tools.mcp_tool._ensure_mcp_loop"), \
patch("tools.mcp_tool._run_on_mcp_loop") as mock_run:
patch("tools.mcp_tool_loop._ensure_mcp_loop"), \
patch("tools.mcp_tool_loop._run_on_mcp_loop") as mock_run:
mcp.register_mcp_servers(config)
_mcp_discovery.register_mcp_servers(config)
mock_get.assert_not_called()
mock_run.assert_called_once()
@@ -109,10 +114,10 @@ class TestLazyMcpRegistration:
mcp._lazy_server_configs["playwright"] = dict(config["playwright"])
mcp._lazy_server_tool_names["playwright"] = ["mcp_playwright_browser_navigate"]
with patch("tools.mcp_tool._MCP_AVAILABLE", True), \
patch("tools.mcp_tool._register_from_cache_sync") as mock_register, \
patch("tools.mcp_tool._run_on_mcp_loop") as mock_run:
patch("tools.mcp_tool_registration._register_from_cache_sync") as mock_register, \
patch("tools.mcp_tool_loop._run_on_mcp_loop") as mock_run:
names = mcp.register_mcp_servers(config)
names = _mcp_discovery.register_mcp_servers(config)
mock_register.assert_not_called()
mock_run.assert_not_called()
@@ -156,9 +161,9 @@ class TestLazyFirstUseConnect:
mcp._servers["playwright"] = connected
return True
with patch.object(mcp, "_ensure_lazy_server_connected", side_effect=_connect) as mock_connect, \
patch.object(mcp, "_run_on_mcp_loop", side_effect=self._run_on_loop):
handler = mcp._make_tool_handler("playwright", "browser_navigate", 5)
with patch.object(_mcp_discovery, "_ensure_lazy_server_connected", side_effect=_connect) as mock_connect, \
patch.object(_mcp_loop, "_run_on_mcp_loop", side_effect=self._run_on_loop):
handler = _mcp_handlers._make_tool_handler("playwright", "browser_navigate", 5)
out = handler({}, task_id="t1")
mock_connect.assert_called_once_with("playwright")
@@ -183,10 +188,10 @@ class TestLazyFirstUseConnect:
async def _fake_paginate(list_method, items_attr, server_name):
return [SimpleNamespace(uri="file:///a", name="a", description="", mimeType="")]
with patch.object(mcp, "_ensure_lazy_server_connected", side_effect=_connect) as mock_connect, \
with patch.object(_mcp_discovery, "_ensure_lazy_server_connected", side_effect=_connect) as mock_connect, \
patch.object(mcp, "_paginate_full_list", side_effect=_fake_paginate), \
patch.object(mcp, "_run_on_mcp_loop", side_effect=self._run_on_loop):
handler = mcp._make_list_resources_handler("playwright", 5)
patch.object(_mcp_loop, "_run_on_mcp_loop", side_effect=self._run_on_loop):
handler = _mcp_handlers._make_list_resources_handler("playwright", 5)
out = handler({})
mock_connect.assert_called_once_with("playwright")
@@ -207,9 +212,9 @@ class TestLazyFirstUseConnect:
mcp._servers["playwright"] = connected
return True
with patch.object(mcp, "_ensure_lazy_server_connected", side_effect=_connect) as mock_connect, \
patch.object(mcp, "_run_on_mcp_loop", side_effect=self._run_on_loop):
handler = mcp._make_get_prompt_handler("playwright", 5)
with patch.object(_mcp_discovery, "_ensure_lazy_server_connected", side_effect=_connect) as mock_connect, \
patch.object(_mcp_loop, "_run_on_mcp_loop", side_effect=self._run_on_loop):
handler = _mcp_handlers._make_get_prompt_handler("playwright", 5)
out = handler({"name": "greeting"})
mock_connect.assert_called_once_with("playwright")
@@ -219,16 +224,16 @@ class TestLazyFirstUseConnect:
def test_check_fn_passes_for_lazy_registered_server(self):
mcp._lazy_server_configs["playwright"] = {"lazy": True}
mcp._lazy_server_fingerprints["playwright"] = "abc"
assert mcp._make_check_fn("playwright")() is True
assert _mcp_handlers._make_check_fn("playwright")() is True
def test_check_fn_fails_for_unknown_server(self):
assert mcp._make_check_fn("nope")() is False
assert _mcp_handlers._make_check_fn("nope")() is False
def test_lazy_connect_respects_connect_cooldown(self):
mcp._lazy_server_configs["playwright"] = {"command": "npx", "lazy": True}
with patch.object(mcp, "_connect_cooldown_active", return_value=True), \
patch.object(mcp, "_run_on_mcp_loop") as mock_run:
assert mcp._ensure_lazy_server_connected("playwright") is False
with patch.object(_mcp_discovery, "_connect_cooldown_active", return_value=True), \
patch.object(_mcp_loop, "_run_on_mcp_loop") as mock_run:
assert _mcp_discovery._ensure_lazy_server_connected("playwright") is False
mock_run.assert_not_called()
def test_lazy_connect_success_clears_lazy_state(self):
@@ -248,9 +253,9 @@ class TestLazyFirstUseConnect:
coro.close()
return ["mcp_playwright_browser_navigate"]
with patch.object(mcp, "_ensure_mcp_loop"), \
patch.object(mcp, "_run_on_mcp_loop", side_effect=_fake_run):
assert mcp._ensure_lazy_server_connected("playwright") is True
with patch.object(_mcp_loop, "_ensure_mcp_loop"), \
patch.object(_mcp_loop, "_run_on_mcp_loop", side_effect=_fake_run):
assert _mcp_discovery._ensure_lazy_server_connected("playwright") is True
assert "playwright" not in mcp._lazy_server_configs
assert "playwright" not in mcp._lazy_server_fingerprints
@@ -280,10 +285,10 @@ class TestLazyFirstUseConnect:
coro.close()
return ["mcp_playwright_tool_y"]
with patch.object(mcp, "_ensure_mcp_loop"), \
patch.object(mcp, "_run_on_mcp_loop", side_effect=_fake_run), \
with patch.object(_mcp_loop, "_ensure_mcp_loop"), \
patch.object(_mcp_loop, "_run_on_mcp_loop", side_effect=_fake_run), \
patch.object(registry, "deregister") as mock_dereg:
assert mcp._ensure_lazy_server_connected("playwright") is True
assert _mcp_discovery._ensure_lazy_server_connected("playwright") is True
mock_dereg.assert_called_once_with("mcp_playwright_tool_x", scope=None)
@@ -295,10 +300,10 @@ class TestLazyFirstUseConnect:
coro.close()
raise RuntimeError("spawn failed")
with patch.object(mcp, "_ensure_mcp_loop"), \
patch.object(mcp, "_run_on_mcp_loop", side_effect=_fake_run), \
patch.object(mcp, "_record_connect_failure") as mock_record:
assert mcp._ensure_lazy_server_connected("playwright") is False
with patch.object(_mcp_loop, "_ensure_mcp_loop"), \
patch.object(_mcp_loop, "_run_on_mcp_loop", side_effect=_fake_run), \
patch.object(_mcp_discovery, "_record_connect_failure") as mock_record:
assert _mcp_discovery._ensure_lazy_server_connected("playwright") is False
mock_record.assert_called_once_with("playwright")
# Config retained so a later call can retry after cooldown.
@@ -312,20 +317,20 @@ class TestCacheLoadDescriptionScan:
# eager discovery.
entry = _fake_cache_entry()
config = {"command": "npx", "args": [], "lazy": True}
with patch.object(mcp, "_scan_mcp_description", return_value=[]) as mock_scan, \
patch.object(mcp, "_convert_mcp_schema", side_effect=RuntimeError("stop")), \
with patch.object(_mcp_schema, "_scan_mcp_description", return_value=[]) as mock_scan, \
patch.object(_mcp_schema, "_convert_mcp_schema", side_effect=RuntimeError("stop")), \
pytest.raises(RuntimeError):
mcp._register_from_cache_sync("playwright", config, entry)
_mcp_registration._register_from_cache_sync("playwright", config, entry)
mock_scan.assert_called_once_with("playwright", "browser_navigate", "Navigate")
class TestResolveServerLazy:
def test_default_off(self):
assert mcp._resolve_server_lazy("s", {"command": "npx"}) is False
assert _mcp_discovery._resolve_server_lazy("s", {"command": "npx"}) is False
def test_explicit_true(self):
assert mcp._resolve_server_lazy("s", {"command": "npx", "lazy": True}) is True
assert _mcp_discovery._resolve_server_lazy("s", {"command": "npx", "lazy": True}) is True
def test_explicit_false(self):
assert mcp._resolve_server_lazy("s", {"command": "npx", "lazy": False}) is False
assert _mcp_discovery._resolve_server_lazy("s", {"command": "npx", "lazy": False}) is False
@@ -6,7 +6,7 @@ thread's. A per-request profile scope (dashboard ?profile= endpoints, e.g.
the MCP "Test server" probe) would silently vanish for anything resolving
get_hermes_home() inside the coroutine, most visibly OAuth token-store
paths. _run_on_mcp_loop now wraps scheduled coroutines with the caller's
override (mcp_tool._wrap_with_home_override).
override (mcp_tool_loop._wrap_with_home_override).
"""
import os
@@ -15,11 +15,11 @@ import pytest
@pytest.fixture
def mcp_loop():
import tools.mcp_tool as mcp_tool
from tools import mcp_tool_loop as _mcp_loop
mcp_tool._ensure_mcp_loop()
yield mcp_tool
mcp_tool._stop_mcp_loop()
_mcp_loop._ensure_mcp_loop()
yield _mcp_loop
_mcp_loop._stop_mcp_loop()
def test_override_propagates_to_mcp_loop(tmp_path, monkeypatch, mcp_loop):
+1 -1
View File
@@ -185,7 +185,7 @@ def test_windows_selects_launchers_never_the_sh_script():
`os.name`, which breaks path handling process-wide (it took pytest's own
traceback formatting down when I tried).
"""
from tools.mcp_tool import _npx_bin_candidates
from tools.mcp_tool_config import _npx_bin_candidates
win = _npx_bin_candidates("/c/bin", "mcp-linear", windows=True)
assert win == ["/c/bin/mcp-linear.cmd", "/c/bin/mcp-linear.exe"]
+2 -1
View File
@@ -17,6 +17,7 @@ import pytest
def test_revival_discovery_registers_tools_while_ready_is_cleared(monkeypatch):
"""A managed server revival must publish tools before readiness is reset."""
from tools import mcp_tool
from tools import mcp_tool_registration as _mcp_registration
from tools.mcp_tool import MCPServerTask
server = MCPServerTask("srv")
@@ -33,7 +34,7 @@ def test_revival_discovery_registers_tools_while_ready_is_cleared(monkeypatch):
monkeypatch.setitem(mcp_tool._servers, server.name, server)
register = MagicMock(return_value=["srv__send_message"])
monkeypatch.setattr(mcp_tool, "_register_server_tools", register)
monkeypatch.setattr(_mcp_registration, "_register_server_tools", register)
asyncio.run(server._discover_tools())
@@ -25,11 +25,11 @@ import pytest
@pytest.fixture
def mcp_loop():
import tools.mcp_tool as mcp_tool
from tools import mcp_tool_loop as _mcp_loop
mcp_tool._ensure_mcp_loop()
yield mcp_tool
mcp_tool._stop_mcp_loop()
_mcp_loop._ensure_mcp_loop()
yield _mcp_loop
_mcp_loop._stop_mcp_loop()
def test_inner_wait_for_timeout_surfaces_promptly_without_spinning(mcp_loop):
@@ -29,7 +29,8 @@ from contextlib import contextmanager
import pytest
from tools.mcp_tool import MCPServerTask, NonMcpEndpointError
from tools.mcp_tool import MCPServerTask
from tools.mcp_tool_errors import NonMcpEndpointError
def _make_task(name: str = "probe_srv") -> MCPServerTask:
@@ -201,6 +202,7 @@ def test_run_skips_preflight_for_oauth(monkeypatch):
(``.well-known/oauth-protected-resource``), not by a GET content-type check.
"""
import tools.mcp_tool as _mcp
from tools import mcp_tool_errors as _mcp_errors
preflight_calls: list[str] = []
@@ -215,7 +217,7 @@ def test_run_skips_preflight_for_oauth(monkeypatch):
raise asyncio.CancelledError()
# Bypass URL validation so the test doesn't need a live network.
monkeypatch.setattr(_mcp, "_validate_remote_mcp_url", lambda n, u: None)
monkeypatch.setattr(_mcp_errors, "_validate_remote_mcp_url", lambda n, u: None)
monkeypatch.setattr(_mcp.MCPServerTask, "_preflight_content_type", _fake_preflight)
monkeypatch.setattr(_mcp.MCPServerTask, "_run_http", _fake_run_http)
@@ -239,6 +241,7 @@ def test_run_skips_preflight_when_skip_preflight_set(monkeypatch):
non-OAuth auth schemes the probe headers don't satisfy).
"""
import tools.mcp_tool as _mcp
from tools import mcp_tool_errors as _mcp_errors
preflight_calls: list[str] = []
@@ -249,7 +252,7 @@ def test_run_skips_preflight_when_skip_preflight_set(monkeypatch):
async def _fake_run_http(self, config):
raise asyncio.CancelledError()
monkeypatch.setattr(_mcp, "_validate_remote_mcp_url", lambda n, u: None)
monkeypatch.setattr(_mcp_errors, "_validate_remote_mcp_url", lambda n, u: None)
monkeypatch.setattr(_mcp.MCPServerTask, "_preflight_content_type", _fake_preflight)
monkeypatch.setattr(_mcp.MCPServerTask, "_run_http", _fake_run_http)
+13 -13
View File
@@ -26,7 +26,7 @@ class TestProbeMcpServerTools:
def test_returns_empty_when_mcp_not_available(self):
with patch("tools.mcp_tool._MCP_AVAILABLE", False):
from tools.mcp_tool import probe_mcp_server_tools
from tools.mcp_tool_discovery import probe_mcp_server_tools
result = probe_mcp_server_tools()
assert result == {}
@@ -48,11 +48,11 @@ class TestProbeMcpServerTools:
return mock_server
with patch("tools.mcp_tool._MCP_AVAILABLE", True), \
patch("tools.mcp_tool._load_mcp_config", return_value=config), \
patch("tools.mcp_tool._connect_server", side_effect=fake_connect), \
patch("tools.mcp_tool._ensure_mcp_loop"), \
patch("tools.mcp_tool._run_on_mcp_loop") as mock_run, \
patch("tools.mcp_tool._stop_mcp_loop"):
patch("tools.mcp_tool_config._load_mcp_config", return_value=config), \
patch("tools.mcp_tool_discovery._connect_server", side_effect=fake_connect), \
patch("tools.mcp_tool_loop._ensure_mcp_loop"), \
patch("tools.mcp_tool_loop._run_on_mcp_loop") as mock_run, \
patch("tools.mcp_tool_loop._stop_mcp_loop"):
def run_coro(coro_or_factory, timeout=120):
coro = coro_or_factory() if callable(coro_or_factory) else coro_or_factory
@@ -64,7 +64,7 @@ class TestProbeMcpServerTools:
mock_run.side_effect = run_coro
from tools.mcp_tool import probe_mcp_server_tools
from tools.mcp_tool_discovery import probe_mcp_server_tools
result = probe_mcp_server_tools()
assert "github" in result
@@ -89,11 +89,11 @@ class TestProbeMcpServerTools:
return mock_server
with patch("tools.mcp_tool._MCP_AVAILABLE", True), \
patch("tools.mcp_tool._load_mcp_config", return_value=config), \
patch("tools.mcp_tool._connect_server", side_effect=fake_connect), \
patch("tools.mcp_tool._ensure_mcp_loop"), \
patch("tools.mcp_tool._run_on_mcp_loop") as mock_run, \
patch("tools.mcp_tool._stop_mcp_loop"):
patch("tools.mcp_tool_config._load_mcp_config", return_value=config), \
patch("tools.mcp_tool_discovery._connect_server", side_effect=fake_connect), \
patch("tools.mcp_tool_loop._ensure_mcp_loop"), \
patch("tools.mcp_tool_loop._run_on_mcp_loop") as mock_run, \
patch("tools.mcp_tool_loop._stop_mcp_loop"):
def run_coro(coro_or_factory, timeout=120):
coro = coro_or_factory() if callable(coro_or_factory) else coro_or_factory
@@ -105,7 +105,7 @@ class TestProbeMcpServerTools:
mock_run.side_effect = run_coro
from tools.mcp_tool import probe_mcp_server_tools
from tools.mcp_tool_discovery import probe_mcp_server_tools
result = probe_mcp_server_tools()
assert "github" in result
+2 -5
View File
@@ -10,11 +10,8 @@ import asyncio
import pytest
from tools.mcp_tool import (
MCPServerTask,
_handshake_rejected_as_modern,
_JSONRPC_UNSUPPORTED_PROTOCOL_VERSION,
)
from tools.mcp_tool import MCPServerTask
from tools.mcp_tool_errors import _handshake_rejected_as_modern, _JSONRPC_UNSUPPORTED_PROTOCOL_VERSION
class _Err(Exception):
@@ -16,7 +16,8 @@ import logging
import pytest
from tools.mcp_tool import MCPServerTask, _jittered
from tools.mcp_tool import MCPServerTask
from tools.mcp_tool_common import _jittered
# ── Jitter ───────────────────────────────────────────────────────────────────
+2 -1
View File
@@ -15,6 +15,7 @@ def test_register_wakes_stale_cached_server(monkeypatch, tmp_path):
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
from tools import mcp_tool
from tools import mcp_tool_discovery as _mcp_discovery
woken: list[str] = []
@@ -48,7 +49,7 @@ def test_register_wakes_stale_cached_server(monkeypatch, tmp_path):
monkeypatch.setitem(mcp_tool._servers, "healthy-srv", alive)
try:
result = mcp_tool.register_mcp_servers({
result = _mcp_discovery.register_mcp_servers({
"parked-srv": {"url": "http://127.0.0.1:9/mcp"},
"healthy-srv": {"url": "http://127.0.0.1:9/mcp"},
})
+16 -18
View File
@@ -42,7 +42,7 @@ def doc_cache(tmp_path, monkeypatch):
class TestRenderResourceBlock:
def test_embedded_pdf_blob_is_materialized(self, doc_cache):
from tools.mcp_tool import _render_mcp_resource_block
from tools.mcp_tool_content import _render_mcp_resource_block
out = _render_mcp_resource_block(_embedded(_blob_resource(PDF_BYTES)), "slack")
assert "saved to" in out
@@ -55,19 +55,19 @@ class TestRenderResourceBlock:
def test_malformed_base64_fails_explicitly(self):
from tools.mcp_tool import _render_mcp_resource_block
from tools.mcp_tool_content import _render_mcp_resource_block
res = SimpleNamespace(uri="x://y", mimeType="application/pdf", blob="!!!not-base64!!!", text=None)
out = _render_mcp_resource_block(_embedded(res), "srv")
assert "could not be decoded" in out
def test_non_resource_block_returns_empty(self):
from tools.mcp_tool import _render_mcp_resource_block
from tools.mcp_tool_content import _render_mcp_resource_block
assert _render_mcp_resource_block(SimpleNamespace(type="text", text="hi"), "srv") == ""
def test_path_traversal_uri_is_neutralized(self, doc_cache):
from tools.mcp_tool import _render_mcp_resource_block
from tools.mcp_tool_content import _render_mcp_resource_block
res = _blob_resource(PDF_BYTES, uri="evil://host/../../etc/passwd")
out = _render_mcp_resource_block(_embedded(res), "srv")
@@ -79,13 +79,13 @@ class TestRenderResourceBlock:
class TestResourceFilename:
def test_uri_last_segment_used(self):
from tools.mcp_tool import _mcp_resource_filename
from tools.mcp_tool_content import _mcp_resource_filename
assert _mcp_resource_filename("slack://f/ABC/quarterly.pdf", "application/pdf") == "quarterly.pdf"
def test_long_filename_capped_preserving_extension(self):
from tools.mcp_tool import _mcp_resource_filename
from tools.mcp_tool_content import _mcp_resource_filename
name = _mcp_resource_filename("x://h/" + "a" * 500 + ".pdf", "application/pdf")
assert len(name) <= 150
@@ -95,29 +95,30 @@ class TestResourceFilename:
class TestPreDecodeSizeCap:
def test_oversized_b64_rejected_before_decode(self, monkeypatch):
import tools.mcp_tool as m
from tools import mcp_tool_content as _mcp_content
monkeypatch.setattr(m, "_MCP_RESOURCE_MAX_B64_CHARS", 16)
monkeypatch.setattr(_mcp_content, "_MCP_RESOURCE_MAX_B64_CHARS", 16)
res = SimpleNamespace(
uri="x://y/big.pdf", mimeType="application/pdf",
blob="A" * 100, text=None,
)
called = []
monkeypatch.setattr(base64, "b64decode", lambda *a, **k: called.append(1))
out = m._render_mcp_resource_block(_embedded(res), "srv")
out = _mcp_content._render_mcp_resource_block(_embedded(res), "srv")
assert "too large" in out
assert not called
class TestAudioBlock:
def test_non_audio_returns_empty(self):
from tools.mcp_tool import _cache_mcp_audio_block
from tools.mcp_tool_content import _cache_mcp_audio_block
block = SimpleNamespace(data=base64.b64encode(b"x").decode(), mimeType="application/pdf")
assert _cache_mcp_audio_block(block) == ""
def test_audio_block_cached_as_media(self, tmp_path, monkeypatch):
import gateway.platforms.base as base
from tools.mcp_tool import _cache_mcp_audio_block
from tools.mcp_tool_content import _cache_mcp_audio_block
monkeypatch.setattr(base, "AUDIO_CACHE_DIR", tmp_path)
block = SimpleNamespace(
@@ -131,11 +132,7 @@ class TestAudioBlock:
class TestToolResultLoopOrdering:
def test_mixed_blocks_preserve_order(self, doc_cache):
"""Simulate the tool-result block loop with text + pdf resource."""
from tools.mcp_tool import (
_cache_mcp_image_block,
_cache_mcp_audio_block,
_render_mcp_resource_block,
)
from tools.mcp_tool_content import _cache_mcp_image_block, _cache_mcp_audio_block, _render_mcp_resource_block
blocks = [
SimpleNamespace(type="text", text="File ID: F123\nMIME Type: application/pdf"),
@@ -158,7 +155,7 @@ class TestToolResultLoopOrdering:
assert "saved to" in parts[1]
def test_existing_image_behavior_unchanged(self):
from tools.mcp_tool import _cache_mcp_image_block
from tools.mcp_tool_content import _cache_mcp_image_block
block = SimpleNamespace(
data=base64.b64encode(b"some bytes").decode("ascii"),
@@ -176,6 +173,7 @@ class TestErrorPathResourceText:
from unittest.mock import AsyncMock, MagicMock, patch as mock_patch
from tools import mcp_tool
from tools import mcp_tool_handlers as _mcp_handlers
fake_session = MagicMock()
fake_server = SimpleNamespace(session=fake_session, _rpc_lock=None)
@@ -197,9 +195,9 @@ class TestErrorPathResourceText:
mcp_tool._reset_server_error("test-server")
try:
with mock_patch.dict(mcp_tool._servers, {"test-server": fake_server}), \
mock_patch("tools.mcp_tool._run_on_mcp_loop", side_effect=_fake_run_on_mcp_loop):
mock_patch("tools.mcp_tool_loop._run_on_mcp_loop", side_effect=_fake_run_on_mcp_loop):
fake_session.call_tool = AsyncMock()
yield fake_session, mcp_tool._make_tool_handler("test-server", "my-tool", 30.0)
yield fake_session, _mcp_handlers._make_tool_handler("test-server", "my-tool", 30.0)
finally:
mcp_tool._reset_server_error("test-server")
+1 -1
View File
@@ -13,7 +13,7 @@ see _MCP_HARD_RESULT_CAP_CHARS in tools/mcp_tool.py.
from __future__ import annotations
from tools.mcp_tool import _MCP_HARD_RESULT_CAP_CHARS, _truncate_mcp_text_result
from tools.mcp_tool_content import _MCP_HARD_RESULT_CAP_CHARS, _truncate_mcp_text_result
class TestTruncateMcpTextResult:
+4 -2
View File
@@ -5,6 +5,7 @@ fingerprint keying, read/write round-trip, and invalidation behavior.
"""
import tools.mcp_schema_cache as msc
from tools import mcp_tool_registration as _mcp_registration
class TestConfigFingerprint:
@@ -155,7 +156,7 @@ class TestWriteThroughPreservesSchema:
server.session = MagicMock()
with patch("tools.registry.registry", ToolRegistry()):
registered = mt._register_server_tools("probe_srv", server, {})
registered = _mcp_registration._register_server_tools("probe_srv", server, {})
assert registered, "tool was not registered; write-through never fired"
entry = json.loads((tmp_path / "cache.json").read_text(encoding="utf-8"))["probe_srv"]
return entry
@@ -176,12 +177,13 @@ class TestWriteThroughPreservesSchema:
from unittest.mock import patch
import tools.mcp_tool as mt
from tools import mcp_tool_registration as _mcp_registration
from tools.registry import ToolRegistry
entry = self._cache_write_through(tmp_path, monkeypatch)
lazy_reg = ToolRegistry()
with patch("tools.registry.registry", lazy_reg):
names = mt._register_from_cache_sync("probe_srv", {}, entry)
names = _mcp_registration._register_from_cache_sync("probe_srv", {}, entry)
assert names, "lazy registration produced no tools"
schema = lazy_reg.get_schema("mcp__probe_srv__zhida")
assert schema is not None, "lazy path did not register the tool"
+41 -70
View File
@@ -6,6 +6,8 @@ import signal
from unittest.mock import patch, MagicMock
import pytest
from tools import mcp_tool_lifecycle as _mcp_lifecycle
from tools import mcp_tool_loop as _mcp_loop
# ---------------------------------------------------------------------------
@@ -16,7 +18,7 @@ class TestMCPLoopExceptionHandler:
"""_mcp_loop_exception_handler suppresses benign 'Event loop is closed'."""
def test_suppresses_event_loop_closed(self):
from tools.mcp_tool import _mcp_loop_exception_handler
from tools.mcp_tool_loop import _mcp_loop_exception_handler
loop = MagicMock()
context = {"exception": RuntimeError("Event loop is closed")}
# Should NOT call default handler
@@ -32,12 +34,12 @@ class TestMCPLoopExceptionHandler:
mcp_mod._servers.clear()
mcp_mod._server_connecting.clear()
try:
mcp_mod._ensure_mcp_loop()
_mcp_loop._ensure_mcp_loop()
with mcp_mod._lock:
loop = mcp_mod._mcp_loop
mcp_mod._servers["live"] = MagicMock(session=object())
assert mcp_mod._stop_mcp_loop_if_idle() is False
assert _mcp_lifecycle._stop_mcp_loop_if_idle() is False
with mcp_mod._lock:
assert mcp_mod._mcp_loop is loop
@@ -48,7 +50,7 @@ class TestMCPLoopExceptionHandler:
with mcp_mod._lock:
mcp_mod._servers.clear()
mcp_mod._server_connecting.clear()
mcp_mod._stop_mcp_loop()
_mcp_loop._stop_mcp_loop()
# ---------------------------------------------------------------------------
@@ -59,7 +61,7 @@ class TestStdioPidTracking:
"""_snapshot_child_pids and _stdio_pids track subprocess PIDs."""
def test_snapshot_returns_set(self):
from tools.mcp_tool import _snapshot_child_pids
from tools.mcp_tool_lifecycle import _snapshot_child_pids
result = _snapshot_child_pids()
assert isinstance(result, set)
# All elements should be ints
@@ -75,7 +77,7 @@ class TestStdioPidTracking:
import sys as _sys
import threading
from tools.mcp_tool import _snapshot_child_pids
from tools.mcp_tool_lifecycle import _snapshot_child_pids
procs = []
started = threading.Event()
@@ -99,12 +101,9 @@ class TestStdioPidTracking:
def test_kill_orphaned_handles_dead_pids(self):
"""_kill_orphaned_mcp_children gracefully handles already-dead PIDs."""
from tools.mcp_tool import (
_kill_orphaned_mcp_children,
_orphan_stdio_pid_servers,
_orphan_stdio_pids,
_lock,
)
from tools.mcp_tool_lifecycle import (
_kill_orphaned_mcp_children, _orphan_stdio_pid_servers, _orphan_stdio_pids)
from tools.mcp_tool import _lock
# Use a PID that definitely doesn't exist
fake_pid = 999999999
@@ -122,14 +121,9 @@ class TestStdioPidTracking:
def test_run_stdio_reaps_orphans_before_spawn(self):
"""_run_stdio kills orphaned PIDs from prior failed attempts (#57355)."""
from tools.mcp_tool import (
_kill_orphaned_mcp_children,
_orphan_stdio_pids,
_stdio_pids,
_stdio_pgids,
_lock,
MCPServerTask,
)
from tools.mcp_tool_lifecycle import (
_kill_orphaned_mcp_children, _orphan_stdio_pids, _stdio_pids, _stdio_pgids)
from tools.mcp_tool import _lock, MCPServerTask
from unittest.mock import patch, MagicMock, AsyncMock
# Seed an orphan PID that belongs to a prior failed connection.
@@ -156,10 +150,10 @@ class TestStdioPidTracking:
# stdio_client spawn. Patch the OSV check (local import)
# and stdio_client so no real subprocess is spawned.
with patch("tools.mcp_tool._MCP_AVAILABLE", True), \
patch("tools.mcp_tool._build_safe_env", return_value={}), \
patch("tools.mcp_tool._resolve_stdio_command",
patch("tools.mcp_tool_config._build_safe_env", return_value={}), \
patch("tools.mcp_tool_config._resolve_stdio_command",
return_value=("echo", {})), \
patch("tools.mcp_tool._write_stderr_log_header"), \
patch("tools.mcp_tool_config._write_stderr_log_header"), \
patch("tools.mcp_tool._get_mcp_stderr_log",
return_value=None), \
patch("tools.mcp_tool.check_package_for_malware",
@@ -185,12 +179,9 @@ class TestStdioPidTracking:
def test_kill_orphaned_can_filter_by_server_name(self):
"""Reconnect cleanup reaps only the orphan owned by that MCP server."""
from tools.mcp_tool import (
_kill_orphaned_mcp_children,
_orphan_stdio_pid_servers,
_orphan_stdio_pids,
_lock,
)
from tools.mcp_tool_lifecycle import (
_kill_orphaned_mcp_children, _orphan_stdio_pid_servers, _orphan_stdio_pids)
from tools.mcp_tool import _lock
target_pid = 454545
other_pid = 464646
@@ -229,13 +220,8 @@ class TestStdioPgroupReaping:
"""_kill_orphaned_mcp_children reaps via killpg when a pgid is tracked."""
def _reset_state(self):
from tools.mcp_tool import (
_orphan_stdio_pid_servers,
_orphan_stdio_pids,
_stdio_pgids,
_stdio_pids,
_lock,
)
from tools.mcp_tool_lifecycle import _orphan_stdio_pid_servers, _orphan_stdio_pids, _stdio_pgids, _stdio_pids
from tools.mcp_tool import _lock
with _lock:
_stdio_pids.clear()
_orphan_stdio_pids.clear()
@@ -244,12 +230,8 @@ class TestStdioPgroupReaping:
def test_killpg_used_when_pgid_tracked(self, monkeypatch):
"""SIGTERM and SIGKILL route through killpg when pgid is known."""
from tools.mcp_tool import (
_kill_orphaned_mcp_children,
_orphan_stdio_pids,
_stdio_pgids,
_lock,
)
from tools.mcp_tool_lifecycle import _kill_orphaned_mcp_children, _orphan_stdio_pids, _stdio_pgids
from tools.mcp_tool import _lock
self._reset_state()
fake_pid = 525252
@@ -287,12 +269,8 @@ class TestStdioPgroupReaping:
group, killpg(pgid) would signal the gateway itself and crash it.
The guard must skip killpg for that pgid and fall through to per-pid
os.kill instead."""
from tools.mcp_tool import (
_kill_orphaned_mcp_children,
_orphan_stdio_pids,
_stdio_pgids,
_lock,
)
from tools.mcp_tool_lifecycle import _kill_orphaned_mcp_children, _orphan_stdio_pids, _stdio_pgids
from tools.mcp_tool import _lock
if not hasattr(os, "killpg") or not hasattr(os, "getpgrp"):
pytest.skip("os.killpg/os.getpgrp not available on this platform")
@@ -335,12 +313,8 @@ class TestStdioPgroupReaping:
def test_no_pgid_uses_per_pid_kill(self, monkeypatch):
"""When no pgid is recorded (e.g. Windows), fall back to os.kill."""
from tools.mcp_tool import (
_kill_orphaned_mcp_children,
_orphan_stdio_pids,
_stdio_pgids,
_lock,
)
from tools.mcp_tool_lifecycle import _kill_orphaned_mcp_children, _orphan_stdio_pids, _stdio_pgids
from tools.mcp_tool import _lock
self._reset_state()
fake_pid = 747474
@@ -427,14 +401,9 @@ class TestStdioPgroupReaping:
assert os.getpgid(grandchild_pid) == parent_pgid
# Drive the reaper: register the parent pid + pgid as an orphan.
from tools.mcp_tool import (
_kill_orphaned_mcp_children,
_orphan_stdio_pid_servers,
_orphan_stdio_pids,
_stdio_pgids,
_stdio_pids,
_lock,
)
from tools.mcp_tool_lifecycle import (
_kill_orphaned_mcp_children, _orphan_stdio_pid_servers, _orphan_stdio_pids, _stdio_pgids, _stdio_pids)
from tools.mcp_tool import _lock
with _lock:
_stdio_pids.clear()
_orphan_stdio_pids.clear()
@@ -530,7 +499,7 @@ class TestMCPInitialConnectionRetry:
raise AssertionError("Should not attempt after shutdown")
with patch.object(MCPServerTask, '_run_stdio', fake_run_stdio), \
patch('tools.mcp_tool._jittered', lambda s: 0.01):
patch('tools.mcp_tool_common._jittered', lambda s: 0.01):
task = asyncio.ensure_future(server.run({"command": "fake"}))
# Give the first attempt time to fail, then set shutdown
@@ -592,7 +561,7 @@ class TestMCPLoopDrainOnStop:
mcp_mod._servers.clear()
mcp_mod._server_connecting.clear()
try:
mcp_mod._ensure_mcp_loop()
_mcp_loop._ensure_mcp_loop()
with mcp_mod._lock:
loop = mcp_mod._mcp_loop
assert loop is not None
@@ -631,7 +600,7 @@ class TestMCPLoopDrainOnStop:
),
),
):
mcp_mod._stop_mcp_loop()
_mcp_loop._stop_mcp_loop()
assert state["task"] is not None
assert state["task"].done(), "task left pending when the loop closed"
@@ -644,11 +613,12 @@ class TestMCPLoopDrainOnStop:
with mcp_mod._lock:
mcp_mod._servers.clear()
mcp_mod._server_connecting.clear()
mcp_mod._stop_mcp_loop()
_mcp_loop._stop_mcp_loop()
def test_drain_is_bounded_when_task_ignores_cancellation(self, caplog):
"""A cancellation-resistant task must not hang final MCP shutdown."""
import tools.mcp_tool as mcp_mod
from tools import mcp_tool_lifecycle as _mcp_lifecycle
async def _run():
release = asyncio.Event()
@@ -664,7 +634,7 @@ class TestMCPLoopDrainOnStop:
try:
with caplog.at_level("WARNING", logger=mcp_mod.logger.name):
async with asyncio.timeout(0.5):
await mcp_mod._drain_mcp_loop_tasks(timeout=0.01)
await _mcp_lifecycle._drain_mcp_loop_tasks(timeout=0.01)
assert not task.done(), "drain waited indefinitely for resistant task"
finally:
release.set()
@@ -685,6 +655,7 @@ class TestMCPLoopDrainOnStop:
"""A blocked loop must drain after it resumes, not stop ahead of the drain."""
import threading
import tools.mcp_tool as mcp_mod
from tools import mcp_tool_loop as _mcp_loop
parked_started = threading.Event()
cleanup_ran = threading.Event()
@@ -705,7 +676,7 @@ class TestMCPLoopDrainOnStop:
with mcp_mod._lock:
mcp_mod._servers.clear()
mcp_mod._server_connecting.clear()
mcp_mod._ensure_mcp_loop()
_mcp_loop._ensure_mcp_loop()
with mcp_mod._lock:
loop = mcp_mod._mcp_loop
assert loop is not None
@@ -720,7 +691,7 @@ class TestMCPLoopDrainOnStop:
release_timer.start()
try:
with caplog.at_level("WARNING", logger=mcp_mod.logger.name):
mcp_mod._stop_mcp_loop()
_mcp_loop._stop_mcp_loop()
assert cleanup_ran.is_set(), "drain was overtaken by loop.stop"
assert future.done(), "parked task remained pending after loop resumed"
@@ -732,7 +703,7 @@ class TestMCPLoopDrainOnStop:
with mcp_mod._lock:
mcp_mod._servers.clear()
mcp_mod._server_connecting.clear()
mcp_mod._stop_mcp_loop()
_mcp_loop._stop_mcp_loop()
assert any(
"Timed out waiting for MCP loop drain" in record.getMessage()
@@ -44,8 +44,8 @@ class TestStdioEncodingErrorHandler:
patch("tools.mcp_tool.StdioServerParameters") as mock_params,
patch("tools.mcp_tool.stdio_client", return_value=mock_stdio_cm),
patch("tools.mcp_tool.ClientSession", return_value=mock_session_cm),
patch("tools.mcp_tool._snapshot_child_pids", return_value=set()),
patch("tools.mcp_tool._write_stderr_log_header"),
patch("tools.mcp_tool_lifecycle._snapshot_child_pids", return_value=set()),
patch("tools.mcp_tool_config._write_stderr_log_header"),
patch("tools.mcp_tool._get_mcp_stderr_log", return_value=None),
):
server = MCPServerTask("test-encoding")
@@ -83,8 +83,8 @@ class TestStdioEncodingErrorHandler:
patch("tools.mcp_tool.StdioServerParameters") as mock_params,
patch("tools.mcp_tool.stdio_client", return_value=mock_stdio_cm),
patch("tools.mcp_tool.ClientSession", return_value=mock_session_cm),
patch("tools.mcp_tool._snapshot_child_pids", return_value=set()),
patch("tools.mcp_tool._write_stderr_log_header"),
patch("tools.mcp_tool_lifecycle._snapshot_child_pids", return_value=set()),
patch("tools.mcp_tool_config._write_stderr_log_header"),
patch("tools.mcp_tool._get_mcp_stderr_log", return_value=None),
):
server = MCPServerTask("test-encoding")
@@ -28,6 +28,7 @@ from unittest.mock import MagicMock
import pytest
pytest.importorskip("mcp")
from tools import mcp_tool_loop as _mcp_loop # noqa: E402
def _success_result():
@@ -94,7 +95,7 @@ def test_precall_dead_children_respawn_and_retry(monkeypatch, tmp_path):
retry once, and hand the model a normal result — no error at all."""
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
from tools import mcp_tool
from tools.mcp_tool import _make_tool_handler
from tools.mcp_tool_handlers import _make_tool_handler
called = {"n": 0}
alive = {"v": False}
@@ -117,7 +118,7 @@ def test_precall_dead_children_respawn_and_retry(monkeypatch, tmp_path):
children_dead=lambda: not alive["v"],
on_reconnect=_respawn,
)
mcp_tool._ensure_mcp_loop()
_mcp_loop._ensure_mcp_loop()
try:
handler = _make_tool_handler("srv-dead", "tool1", 10.0)
parsed = json.loads(handler({}))
@@ -135,7 +136,7 @@ def test_midcall_child_exit_respawn_and_retry(monkeypatch, tmp_path):
so the caller still gets its result."""
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
from tools import mcp_tool
from tools.mcp_tool import _make_tool_handler
from tools.mcp_tool_handlers import _make_tool_handler
alive = {"v": True}
@@ -163,7 +164,7 @@ def test_midcall_child_exit_respawn_and_retry(monkeypatch, tmp_path):
on_reconnect=_respawn,
)
server._watch_stdio_children = _watch_children
mcp_tool._ensure_mcp_loop()
_mcp_loop._ensure_mcp_loop()
try:
handler = _make_tool_handler("srv-midcall", "tool1", 10.0)
# The child dies once the RPC is in flight.
@@ -185,7 +186,7 @@ def test_dead_child_never_returning_is_not_reported_as_a_timeout(
investigation into a healthy remote backend)."""
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
from tools import mcp_tool
from tools.mcp_tool import _make_tool_handler
from tools.mcp_tool_handlers import _make_tool_handler
monkeypatch.setattr(mcp_tool, "_STDIO_RESPAWN_WAIT_SEC", 1.0)
called = {"n": 0}
@@ -197,7 +198,7 @@ def test_dead_child_never_returning_is_not_reported_as_a_timeout(
server = _install_stub_server(
mcp_tool, "srv-gone", _call_tool, children_dead=lambda: True,
)
mcp_tool._ensure_mcp_loop()
_mcp_loop._ensure_mcp_loop()
try:
handler = _make_tool_handler("srv-gone", "tool1", 300.0)
parsed = json.loads(handler({}))
@@ -221,7 +222,8 @@ def test_child_dying_again_after_respawn_does_not_hot_cycle(
is what parks it, and this path must not fight that."""
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
from tools import mcp_tool
from tools.mcp_tool import _make_tool_handler
from tools import mcp_tool_loop as _mcp_loop
from tools.mcp_tool_handlers import _make_tool_handler
monkeypatch.setattr(mcp_tool, "_STDIO_RESPAWN_WAIT_SEC", 1.0)
called = {"n": 0}
@@ -243,7 +245,7 @@ def test_child_dying_again_after_respawn_does_not_hot_cycle(
children_dead=lambda: True,
on_reconnect=_respawn_then_die,
)
mcp_tool._ensure_mcp_loop()
_mcp_loop._ensure_mcp_loop()
try:
handler = _make_tool_handler("srv-flap", "tool1", 10.0)
parsed = json.loads(handler({}))
+3 -2
View File
@@ -65,6 +65,7 @@ class TestStdioInitializeTimeout:
"""A stdio server that hangs at ``initialize`` must fail within
``connect_timeout`` — not block ``_run_stdio`` forever (#59349)."""
from tools import mcp_tool
from tools import mcp_tool_config as _mcp_config
server = mcp_tool.MCPServerTask("leak-guard")
config = {"command": "fake-mcp", "args": [], "connect_timeout": 0.2}
@@ -72,8 +73,8 @@ class TestStdioInitializeTimeout:
async def drive():
with patch.object(mcp_tool, "stdio_client", _fake_stdio_client), \
patch.object(mcp_tool, "ClientSession", _fake_client_session), \
patch.object(mcp_tool, "_resolve_stdio_command", lambda c, e: (c, e)), \
patch.object(mcp_tool, "_write_stderr_log_header", lambda *_a, **_k: None), \
patch.object(_mcp_config, "_resolve_stdio_command", lambda c, e: (c, e)), \
patch.object(_mcp_config, "_write_stderr_log_header", lambda *_a, **_k: None), \
patch.object(mcp_tool, "_get_mcp_stderr_log", lambda: None), \
patch("tools.osv_check.check_package_for_malware",
lambda *_a, **_k: None):
+26 -24
View File
@@ -8,6 +8,8 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from tools import mcp_tool
from tools import mcp_tool_content as _mcp_content
from tools import mcp_tool_handlers as _mcp_handlers
class _FakeContentBlock:
@@ -60,7 +62,7 @@ def _patch_mcp_server():
# fresh loop that _fake_run_on_mcp_loop spins up, not at fixture import.
fake_server = SimpleNamespace(session=fake_session, _rpc_lock=None)
with patch.dict(mcp_tool._servers, {"test-server": fake_server}), \
patch("tools.mcp_tool._run_on_mcp_loop", side_effect=_fake_run_on_mcp_loop):
patch("tools.mcp_tool_loop._run_on_mcp_loop", side_effect=_fake_run_on_mcp_loop):
yield fake_session
@@ -75,7 +77,7 @@ class TestStructuredContentPreservation:
content=[_FakeContentBlock("hello")],
)
)
handler = mcp_tool._make_tool_handler("test-server", "my-tool", 30.0)
handler = _mcp_handlers._make_tool_handler("test-server", "my-tool", 30.0)
raw = handler({})
data = json.loads(raw)
assert data == {"result": "hello"}
@@ -90,7 +92,7 @@ class TestStructuredContentPreservation:
structuredContent=None,
)
)
handler = mcp_tool._make_tool_handler("test-server", "my-tool", 30.0)
handler = _mcp_handlers._make_tool_handler("test-server", "my-tool", 30.0)
raw = handler({})
data = json.loads(raw)
assert data == {"result": "done"}
@@ -105,7 +107,7 @@ class TestStructuredContentPreservation:
structuredContent=payload,
)
)
handler = mcp_tool._make_tool_handler("test-server", "my-tool", 30.0)
handler = _mcp_handlers._make_tool_handler("test-server", "my-tool", 30.0)
raw = handler({})
data = json.loads(raw)
assert data["result"] == payload
@@ -125,7 +127,7 @@ class TestMetaPassthrough:
meta={"com.example/handoff": {"url": "https://x"}},
)
)
handler = mcp_tool._make_tool_handler("test-server", "my-tool", 30.0)
handler = _mcp_handlers._make_tool_handler("test-server", "my-tool", 30.0)
data = json.loads(handler({}))
assert data["result"] == "done"
assert data["_meta"] == {"com.example/handoff": {"url": "https://x"}}
@@ -143,7 +145,7 @@ class TestMetaPassthrough:
},
)
)
handler = mcp_tool._make_tool_handler("test-server", "my-tool", 30.0)
handler = _mcp_handlers._make_tool_handler("test-server", "my-tool", 30.0)
data = json.loads(handler({}))
assert data["_meta"] == {
"com.example.mcp/vendor": "keep",
@@ -158,7 +160,7 @@ class TestMetaPassthrough:
meta={"mcp.io/internal": True},
)
)
handler = mcp_tool._make_tool_handler("test-server", "my-tool", 30.0)
handler = _mcp_handlers._make_tool_handler("test-server", "my-tool", 30.0)
data = json.loads(handler({}))
assert data == {"result": "done"}
@@ -172,7 +174,7 @@ class TestMetaPassthrough:
meta={"com.example/k": "v"},
)
)
handler = mcp_tool._make_tool_handler("test-server", "my-tool", 30.0)
handler = _mcp_handlers._make_tool_handler("test-server", "my-tool", 30.0)
data = json.loads(handler({}))
assert data == {
"result": "txt",
@@ -187,7 +189,7 @@ class TestMetaPassthrough:
meta={"com.example/obj": object()},
)
)
handler = mcp_tool._make_tool_handler("test-server", "my-tool", 30.0)
handler = _mcp_handlers._make_tool_handler("test-server", "my-tool", 30.0)
data = json.loads(handler({}))
assert data == {"result": "done"}
@@ -199,22 +201,22 @@ class TestMetaPassthrough:
meta="not-a-dict",
)
)
handler = mcp_tool._make_tool_handler("test-server", "my-tool", 30.0)
handler = _mcp_handlers._make_tool_handler("test-server", "my-tool", 30.0)
data = json.loads(handler({}))
assert data == {"result": "done"}
class TestReservedMetaKeyPredicate:
def test_reserved_prefixes(self):
assert mcp_tool._is_reserved_mcp_meta_key("modelcontextprotocol.io/x")
assert mcp_tool._is_reserved_mcp_meta_key("mcp.dev/x")
assert mcp_tool._is_reserved_mcp_meta_key("tools.mcp.com/x")
assert _mcp_content._is_reserved_mcp_meta_key("modelcontextprotocol.io/x")
assert _mcp_content._is_reserved_mcp_meta_key("mcp.dev/x")
assert _mcp_content._is_reserved_mcp_meta_key("tools.mcp.com/x")
def test_vendor_and_unprefixed_not_reserved(self):
assert not mcp_tool._is_reserved_mcp_meta_key("com.example.mcp/x") # trailing label
assert not mcp_tool._is_reserved_mcp_meta_key("com.example/x")
assert not mcp_tool._is_reserved_mcp_meta_key("plain-key")
assert not mcp_tool._is_reserved_mcp_meta_key("/leading-slash")
assert not _mcp_content._is_reserved_mcp_meta_key("com.example.mcp/x") # trailing label
assert not _mcp_content._is_reserved_mcp_meta_key("com.example/x")
assert not _mcp_content._is_reserved_mcp_meta_key("plain-key")
assert not _mcp_content._is_reserved_mcp_meta_key("/leading-slash")
class TestContentStructuredArbitration:
@@ -235,7 +237,7 @@ class TestContentStructuredArbitration:
structuredContent=payload,
)
)
handler = mcp_tool._make_tool_handler("test-server", "my-tool", 30.0)
handler = _mcp_handlers._make_tool_handler("test-server", "my-tool", 30.0)
data = json.loads(handler({}))
assert data == {"result": json.dumps(payload)}
@@ -248,7 +250,7 @@ class TestContentStructuredArbitration:
structuredContent={"items": [1, 2, 3]},
)
)
handler = mcp_tool._make_tool_handler("test-server", "my-tool", 30.0)
handler = _mcp_handlers._make_tool_handler("test-server", "my-tool", 30.0)
data = json.loads(handler({}))
assert data == {"result": "3 item(s) found"}
@@ -262,7 +264,7 @@ class TestContentStructuredArbitration:
structuredContent=payload,
)
)
handler = mcp_tool._make_tool_handler("test-server", "my-tool", 30.0)
handler = _mcp_handlers._make_tool_handler("test-server", "my-tool", 30.0)
data = json.loads(handler({}))
assert data["structuredContent"] == payload
@@ -276,7 +278,7 @@ class TestContentStructuredArbitration:
structuredContent=payload,
)
)
handler = mcp_tool._make_tool_handler("test-server", "my-tool", 30.0)
handler = _mcp_handlers._make_tool_handler("test-server", "my-tool", 30.0)
data = json.loads(handler({}))
assert data["result"] == payload
@@ -299,7 +301,7 @@ class TestDroppedBlockNotice:
session.call_tool = AsyncMock(
return_value=_FakeCallToolResult(content=[weird])
)
handler = mcp_tool._make_tool_handler("test-server", "my-tool", 30.0)
handler = _mcp_handlers._make_tool_handler("test-server", "my-tool", 30.0)
data = json.loads(handler({}))
assert "[MCP content dropped: unsupported block" in data["result"]
assert "type=hologram" in data["result"]
@@ -316,13 +318,13 @@ class TestDroppedBlockNotice:
content=[weird], structuredContent=payload,
)
)
handler = mcp_tool._make_tool_handler("test-server", "my-tool", 30.0)
handler = _mcp_handlers._make_tool_handler("test-server", "my-tool", 30.0)
data = json.loads(handler({}))
assert data["structuredContent"] == payload
assert "[MCP content dropped" in data["result"]
def test_notice_helper_minimal_block(self):
notice = mcp_tool._render_mcp_dropped_block_notice(
notice = _mcp_content._render_mcp_dropped_block_notice(
SimpleNamespace(), "mystery"
)
assert notice == "[MCP content dropped: unsupported block (type=mystery)]"
+2 -1
View File
@@ -9,7 +9,8 @@ from __future__ import annotations
import pytest
from tools.mcp_tool import _DEFAULT_TOOL_TIMEOUT, _resolve_tool_timeout
from tools.mcp_tool import _DEFAULT_TOOL_TIMEOUT
from tools.mcp_tool_common import _resolve_tool_timeout
class TestMcpToolTimeoutResolution:
+154 -140
View File
@@ -68,6 +68,7 @@ class TestFilterMCPChildren:
import sys
import tools.mcp_tool as mcp_tool
from tools import mcp_tool_lifecycle as _mcp_lifecycle
cmdlines = {
101: [
@@ -99,7 +100,7 @@ class TestFilterMCPChildren:
)
monkeypatch.setitem(sys.modules, "psutil", fake_psutil)
assert mcp_tool._filter_mcp_children({101, 102, 103}) == {103}
assert _mcp_lifecycle._filter_mcp_children({101, 102, 103}) == {103}
# ---------------------------------------------------------------------------
@@ -118,7 +119,7 @@ class TestLoadMCPConfig:
}
}
with patch("hermes_cli.config.load_config", return_value={"mcp_servers": servers}):
from tools.mcp_tool import _load_mcp_config
from tools.mcp_tool_config import _load_mcp_config
result = _load_mcp_config()
assert "filesystem" in result
assert result["filesystem"]["command"] == "npx"
@@ -126,7 +127,7 @@ class TestLoadMCPConfig:
def test_mcp_servers_not_dict_returns_empty(self):
"""mcp_servers set to non-dict value -> empty dict."""
with patch("hermes_cli.config.load_config", return_value={"mcp_servers": "invalid"}):
from tools.mcp_tool import _load_mcp_config
from tools.mcp_tool_config import _load_mcp_config
result = _load_mcp_config()
assert result == {}
@@ -146,7 +147,7 @@ class TestLoadMCPConfig:
patch("hermes_cli.plugins.get_plugin_manager", return_value=manager),
patch.dict(os.environ, {"PORT": "3000"}),
):
from tools.mcp_tool import _load_mcp_config
from tools.mcp_tool_config import _load_mcp_config
result = _load_mcp_config()
@@ -187,7 +188,7 @@ class TestLoadMCPConfig:
monkeypatch.setenv("HERMES_BUNDLED_PLUGINS", str(bundled))
monkeypatch.setattr(plugins_mod, "_plugin_manager", None)
from tools.mcp_tool import _load_mcp_config
from tools.mcp_tool_config import _load_mcp_config
result = _load_mcp_config()
@@ -216,9 +217,9 @@ class TestMCPParallelSafetyProvenance:
try:
monkeypatch.setattr(mcp_tool, "_MCP_AVAILABLE", True)
monkeypatch.setattr(
mcp_tool, "_filter_suspicious_mcp_servers", lambda servers: servers
_mcp_config, "_filter_suspicious_mcp_servers", lambda servers: servers
)
mcp_tool.register_mcp_servers(
_mcp_discovery.register_mcp_servers(
{
"foo-bar": {"supports_parallel_tool_calls": True},
"foo_bar": {"supports_parallel_tool_calls": False},
@@ -236,6 +237,7 @@ class TestMCPParallelSafetyProvenance:
def test_tool_provenance_keeps_exact_raw_server_names(self):
import tools.mcp_tool as mcp_tool
from tools import mcp_tool_registration as _mcp_registration
first_tool = "mcp__foo_bar__first"
second_tool = "mcp__foo_bar__second"
@@ -247,12 +249,12 @@ class TestMCPParallelSafetyProvenance:
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")
_mcp_registration._track_mcp_tool_server(first_tool, "foo-bar")
_mcp_registration._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() == {
assert _mcp_discovery.is_mcp_tool_parallel_safe(first_tool) is True
assert _mcp_discovery.is_mcp_tool_parallel_safe(second_tool) is False
assert _mcp_discovery.get_registered_mcp_server_names() == {
"foo-bar",
"foo_bar",
}
@@ -268,10 +270,11 @@ class TestMCPStatus:
self, monkeypatch
):
import tools.mcp_tool as mcp_tool
from tools import mcp_tool_config as _mcp_config
from tools import mcp_tool_discovery as _mcp_discovery
monkeypatch.setattr(
mcp_tool,
"_load_mcp_config",
_mcp_config, "_load_mcp_config",
lambda: {
"configured": {"command": "docker", "args": ["mcp", "gateway", "run"]},
"connecting": {"command": "slow-mcp"},
@@ -292,7 +295,7 @@ class TestMCPStatus:
try:
statuses = {
entry["name"]: entry
for entry in mcp_tool.get_mcp_status()
for entry in _mcp_discovery.get_mcp_status()
}
finally:
with mcp_tool._lock:
@@ -315,7 +318,7 @@ class TestMCPStatus:
class TestLifecycleConfig:
def test_get_lifecycle_seconds_accepts_top_level_and_nested_values(self):
from tools.mcp_tool import _get_lifecycle_seconds
from tools.mcp_tool_common import _get_lifecycle_seconds
assert (
_get_lifecycle_seconds(
@@ -330,7 +333,7 @@ class TestLifecycleConfig:
) == 42.0
def test_get_lifecycle_seconds_ignores_invalid_values(self, caplog):
from tools.mcp_tool import _get_lifecycle_seconds
from tools.mcp_tool_common import _get_lifecycle_seconds
assert (
_get_lifecycle_seconds(
@@ -358,7 +361,7 @@ class TestLifecycleConfig:
class TestSchemaConversion:
def test_converts_mcp_tool_to_hermes_schema(self):
from tools.mcp_tool import _convert_mcp_schema
from tools.mcp_tool_schema import _convert_mcp_schema
mcp_tool = _make_mcp_tool(name="read_file", description="Read a file")
schema = _convert_mcp_schema("filesystem", mcp_tool)
@@ -379,7 +382,7 @@ class TestSchemaConversion:
CI/pipelines MCP tool whose ``definitions`` parameter is an array of
pipeline-definition IDs.
"""
from tools.mcp_tool import _convert_mcp_schema
from tools.mcp_tool_schema import _convert_mcp_schema
mcp_tool = _make_mcp_tool(
name="pipelines_build",
@@ -409,7 +412,7 @@ class TestSchemaConversion:
def test_optional_nullable_field_is_collapsed_to_non_null_schema(self):
"""Anthropic rejects MCP/Pydantic anyOf-null optional parameter schemas."""
from tools.mcp_tool import _normalize_mcp_input_schema
from tools.mcp_tool_schema import _normalize_mcp_input_schema
schema = _normalize_mcp_input_schema({
"type": "object",
@@ -434,7 +437,7 @@ class TestSchemaConversion:
def test_hyphens_sanitized_to_underscores(self):
"""Hyphens in tool/server names are replaced with underscores for LLM compat."""
from tools.mcp_tool import _convert_mcp_schema
from tools.mcp_tool_schema import _convert_mcp_schema
mcp_tool = _make_mcp_tool(name="get-sum")
schema = _convert_mcp_schema("my-server", mcp_tool)
@@ -449,7 +452,8 @@ class TestSchemaConversion:
class TestCheckFunction:
def test_disconnected_returns_false(self):
from tools.mcp_tool import _make_check_fn, _servers
from tools.mcp_tool_handlers import _make_check_fn
from tools.mcp_tool import _servers
_servers.pop("test_server", None)
check = _make_check_fn("test_server")
@@ -457,7 +461,8 @@ class TestCheckFunction:
def test_recycled_stdio_server_remains_available_for_lazy_reconnect(self):
from tools.mcp_tool import _make_check_fn, _servers
from tools.mcp_tool_handlers import _make_check_fn
from tools.mcp_tool import _servers
server = _make_mock_server("test_server", session=None)
server._config = {"command": "npx"}
@@ -501,7 +506,7 @@ class TestRunOnMcpLoop:
side_effect=RuntimeError("scheduler down"),
):
with pytest.raises(RuntimeError):
mcp._run_on_mcp_loop(factory)
_mcp_loop._run_on_mcp_loop(factory)
gc.collect()
assert created["coro"] is not None
@@ -528,7 +533,7 @@ class TestRunOnMcpLoop:
with warnings.catch_warnings(record=True) as caught:
warnings.simplefilter("always")
with pytest.raises(RuntimeError, match="not running"):
mcp._run_on_mcp_loop(coro)
_mcp_loop._run_on_mcp_loop(coro)
gc.collect()
assert coro.cr_frame is None
@@ -554,11 +559,12 @@ class TestToolHandler:
coro = coro_or_factory() if callable(coro_or_factory) else coro_or_factory
return asyncio.run(coro)
if coro_side_effect:
return patch("tools.mcp_tool._run_on_mcp_loop", side_effect=coro_side_effect)
return patch("tools.mcp_tool._run_on_mcp_loop", side_effect=fake_run)
return patch("tools.mcp_tool_loop._run_on_mcp_loop", side_effect=coro_side_effect)
return patch("tools.mcp_tool_loop._run_on_mcp_loop", side_effect=fake_run)
def test_successful_call(self):
from tools.mcp_tool import _make_tool_handler, _servers
from tools.mcp_tool_handlers import _make_tool_handler
from tools.mcp_tool import _servers
mock_session = MagicMock()
mock_session.call_tool = AsyncMock(
@@ -578,7 +584,8 @@ class TestToolHandler:
def test_recycled_stdio_server_reconnects_lazily_on_tool_call(self):
from tools.mcp_tool import _make_tool_handler, _servers
from tools.mcp_tool_handlers import _make_tool_handler
from tools.mcp_tool import _servers
mock_session = MagicMock()
mock_session.call_tool = AsyncMock(
@@ -598,7 +605,7 @@ class TestToolHandler:
try:
handler = _make_tool_handler("test_srv", "greet", 120)
with patch("tools.mcp_tool._request_lazy_reconnect", side_effect=fake_lazy_reconnect) as reconnect, \
with patch("tools.mcp_tool_discovery._request_lazy_reconnect", side_effect=fake_lazy_reconnect) as reconnect, \
self._patch_mcp_loop():
result = json.loads(handler({"name": "world"}))
assert result["result"] == "reconnected"
@@ -660,7 +667,7 @@ class TestRunOnMCPLoopInterrupts:
try:
with pytest.raises(InterruptedError, match="User sent a new message"):
mcp_mod._run_on_mcp_loop(_slow_call(), timeout=10)
_mcp_loop._run_on_mcp_loop(_slow_call(), timeout=10)
deadline = time.time() + 2
while time.time() < deadline and not cancelled.is_set():
@@ -699,7 +706,7 @@ class TestRunOnMCPLoopInterrupts:
try:
# 0.1s is the floor the MCP loop clamps short timeouts to.
with pytest.raises(TimeoutError, match=r"MCP call timed out after .*configured timeout: 0.1s"):
mcp_mod._run_on_mcp_loop(_slow_call(), timeout=0.1)
_mcp_loop._run_on_mcp_loop(_slow_call(), timeout=0.1)
deadline = time.time() + 2
while time.time() < deadline and not cancelled.is_set():
@@ -720,7 +727,8 @@ class TestDiscoverAndRegister:
def test_tools_registered_in_registry(self):
"""_discover_and_register_server registers tools with correct names."""
from tools.registry import ToolRegistry
from tools.mcp_tool import _discover_and_register_server, _servers, MCPServerTask
from tools.mcp_tool_discovery import _discover_and_register_server
from tools.mcp_tool import _servers, MCPServerTask
mock_registry = ToolRegistry()
mock_tools = [
@@ -735,7 +743,7 @@ class TestDiscoverAndRegister:
server._tools = mock_tools
return server
with patch("tools.mcp_tool._connect_server", side_effect=fake_connect), \
with patch("tools.mcp_tool_discovery._connect_server", side_effect=fake_connect), \
patch("tools.registry.registry", mock_registry):
registered = asyncio.run(
_discover_and_register_server("fs", {"command": "npx", "args": []})
@@ -750,7 +758,7 @@ class TestDiscoverAndRegister:
def test_same_server_normalization_collision_skips_all_ambiguous_tools(self, caplog):
from tools.mcp_tool import _register_server_tools
from tools.mcp_tool_registration import _register_server_tools
from tools.registry import ToolRegistry
registry = ToolRegistry()
@@ -766,7 +774,7 @@ class TestDiscoverAndRegister:
config = {"tools": {"resources": False, "prompts": False}}
with patch("tools.registry.registry", registry), \
patch("tools.mcp_tool._track_mcp_tool_server"), \
patch("tools.mcp_tool_registration._track_mcp_tool_server"), \
caplog.at_level(logging.ERROR, logger="tools.mcp_tool"):
registered = _register_server_tools("srv", server, config)
@@ -789,7 +797,7 @@ class TestDiscoverAndRegister:
generated utility is only sugar for servers that lack such a tool, so
the native tool wins and the utility is dropped.
"""
from tools.mcp_tool import _register_server_tools
from tools.mcp_tool_registration import _register_server_tools
from tools.registry import ToolRegistry
registry = ToolRegistry()
@@ -806,7 +814,7 @@ class TestDiscoverAndRegister:
config = {"tools": {"prompts": False}}
with patch("tools.registry.registry", registry), \
patch("tools.mcp_tool._track_mcp_tool_server"), \
patch("tools.mcp_tool_registration._track_mcp_tool_server"), \
caplog.at_level(logging.INFO, logger="tools.mcp_tool"):
registered = _register_server_tools("srv", server, config)
@@ -949,10 +957,10 @@ class TestToolsetInjection:
with patch("tools.mcp_tool._MCP_AVAILABLE", True), \
patch("tools.mcp_tool._servers", fresh_servers), \
patch("tools.mcp_tool._load_mcp_config", return_value=fake_config), \
patch("tools.mcp_tool._connect_server", side_effect=fake_connect), \
patch("tools.mcp_tool_config._load_mcp_config", return_value=fake_config), \
patch("tools.mcp_tool_discovery._connect_server", side_effect=fake_connect), \
patch("tools.registry.registry", mock_registry):
from tools.mcp_tool import discover_mcp_tools
from tools.mcp_tool_discovery import discover_mcp_tools
result = discover_mcp_tools()
assert "mcp__fs__list_files" in result
@@ -993,10 +1001,10 @@ class TestToolsetInjection:
with patch("tools.mcp_tool._MCP_AVAILABLE", True), \
patch("tools.mcp_tool._servers", fresh_servers), \
patch("tools.mcp_tool._load_mcp_config", return_value=fake_config), \
patch("tools.mcp_tool._connect_server", side_effect=flaky_connect), \
patch("tools.mcp_tool_config._load_mcp_config", return_value=fake_config), \
patch("tools.mcp_tool_discovery._connect_server", side_effect=flaky_connect), \
patch("toolsets.TOOLSETS", fake_toolsets):
from tools.mcp_tool import discover_mcp_tools
from tools.mcp_tool_discovery import discover_mcp_tools
# First call: good connects, broken fails
result1 = discover_mcp_tools()
@@ -1030,7 +1038,7 @@ class TestGracefulFallback:
def test_mcp_unavailable_returns_empty(self):
"""When _MCP_AVAILABLE is False, discover_mcp_tools is a no-op."""
with patch("tools.mcp_tool._MCP_AVAILABLE", False):
from tools.mcp_tool import discover_mcp_tools
from tools.mcp_tool_discovery import discover_mcp_tools
result = discover_mcp_tools()
assert result == []
@@ -1049,7 +1057,8 @@ class TestShutdown:
must still cancel and drain that waiter before closing the loop.
"""
import tools.mcp_tool as mcp_mod
from tools.mcp_tool import MCPServerTask, shutdown_mcp_servers
from tools.mcp_tool import MCPServerTask
from tools.mcp_tool_lifecycle import shutdown_mcp_servers
shutdown_started = threading.Event()
parked_task_done = threading.Event()
@@ -1065,7 +1074,7 @@ class TestShutdown:
with mcp_mod._lock:
mcp_mod._servers.clear()
mcp_mod._server_connecting.clear()
mcp_mod._ensure_mcp_loop()
_mcp_loop._ensure_mcp_loop()
with mcp_mod._lock:
loop = mcp_mod._mcp_loop
assert loop is not None
@@ -1118,12 +1127,13 @@ class TestShutdown:
with mcp_mod._lock:
mcp_mod._servers.clear()
mcp_mod._server_connecting.clear()
mcp_mod._stop_mcp_loop()
_mcp_loop._stop_mcp_loop()
def test_shutdown_deregisters_registered_tools(self):
"""shutdown_mcp_servers removes MCP tools and their raw alias."""
import tools.mcp_tool as mcp_mod
from tools.mcp_tool import MCPServerTask, shutdown_mcp_servers, _servers
from tools.mcp_tool import MCPServerTask, _servers
from tools.mcp_tool_lifecycle import shutdown_mcp_servers
from tools.registry import registry
from toolsets import resolve_toolset, validate_toolset
@@ -1144,7 +1154,7 @@ class TestShutdown:
server._registered_tool_names = ["mcp__test__ping"]
_servers["test"] = server
mcp_mod._ensure_mcp_loop()
_mcp_loop._ensure_mcp_loop()
try:
assert validate_toolset("test") is True
assert "mcp__test__ping" in resolve_toolset("test")
@@ -1159,7 +1169,9 @@ class TestShutdown:
def test_shutdown_is_parallel(self):
"""Multiple servers are shut down in parallel via asyncio.gather."""
import tools.mcp_tool as mcp_mod
from tools.mcp_tool import shutdown_mcp_servers, _servers
from tools import mcp_tool_loop as _mcp_loop
from tools.mcp_tool_lifecycle import shutdown_mcp_servers
from tools.mcp_tool import _servers
import time
_servers.clear()
@@ -1174,7 +1186,7 @@ class TestShutdown:
mock_server.shutdown = slow_shutdown
_servers[f"srv_{i}"] = mock_server
mcp_mod._ensure_mcp_loop()
_mcp_loop._ensure_mcp_loop()
try:
start = time.monotonic()
shutdown_mcp_servers()
@@ -1200,7 +1212,7 @@ class TestBuildSafeEnv:
def test_only_safe_vars_passed(self):
"""Only safe baseline vars and XDG_* from os.environ are included."""
from tools.mcp_tool import _build_safe_env
from tools.mcp_tool_config import _build_safe_env
fake_env = {
"PATH": "/usr/bin",
@@ -1230,7 +1242,7 @@ class TestBuildSafeEnv:
def test_secret_vars_excluded(self):
"""Sensitive env vars from os.environ are NOT passed through."""
from tools.mcp_tool import _build_safe_env
from tools.mcp_tool_config import _build_safe_env
fake_env = {
"PATH": "/usr/bin",
@@ -1254,7 +1266,7 @@ class TestBuildSafeEnv:
"""Vars tagged by an external secret source (Bitwarden/1Password) are
deliberately allowed for MCP stdio servers."""
from hermes_cli import env_loader
from tools.mcp_tool import _build_safe_env
from tools.mcp_tool_config import _build_safe_env
monkeypatch.setitem(env_loader._SECRET_SOURCES, "ALPACA_API_KEY", "bitwarden")
monkeypatch.setitem(env_loader._SECRET_SOURCES, "NOTION_TOKEN", "onepassword")
@@ -1274,7 +1286,7 @@ class TestBuildSafeEnv:
def test_windows_location_vars_passed_without_secrets(self):
"""Windows launcher tools need location vars, but secrets stay filtered."""
from tools.mcp_tool import _build_safe_env
from tools.mcp_tool_config import _build_safe_env
fake_env = {
"PATH": r"C:\Windows\System32",
@@ -1308,7 +1320,7 @@ class TestSanitizeError:
"""Tests for _sanitize_error() credential stripping."""
def test_strips_credentials(self):
from tools.mcp_tool import _sanitize_error
from tools.mcp_tool_common import _sanitize_error
for text, expected in (
("Error with ghp_abc123def456", "Error with [REDACTED]"),
@@ -1324,7 +1336,7 @@ class TestSanitizeError:
assert multi.count("[REDACTED]") == 3
def test_no_credentials_unchanged(self):
from tools.mcp_tool import _sanitize_error
from tools.mcp_tool_common import _sanitize_error
result = _sanitize_error("normal error message")
assert result == "normal error message"
@@ -1570,7 +1582,8 @@ class TestConfigurableTimeouts:
def test_timeout_passed_to_handler(self):
"""The tool handler uses the server's configured timeout."""
from tools.mcp_tool import _make_tool_handler, _servers
from tools.mcp_tool_handlers import _make_tool_handler
from tools.mcp_tool import _servers
mock_session = MagicMock()
mock_session.call_tool = AsyncMock(
@@ -1582,7 +1595,7 @@ class TestConfigurableTimeouts:
try:
handler = _make_tool_handler("test_srv", "my_tool", 180)
with patch("tools.mcp_tool._run_on_mcp_loop") as mock_run:
with patch("tools.mcp_tool_loop._run_on_mcp_loop") as mock_run:
def fake_run(coro, timeout=30):
coro.close()
return json.dumps({"result": "ok"})
@@ -1606,7 +1619,7 @@ class TestUtilitySchemas:
"""Tests for _build_utility_schemas() and the schema format of utility tools."""
def test_builds_four_utility_schemas(self):
from tools.mcp_tool import _build_utility_schemas
from tools.mcp_tool_schema import _build_utility_schemas
schemas = _build_utility_schemas("myserver")
assert len(schemas) == 4
@@ -1617,7 +1630,7 @@ class TestUtilitySchemas:
assert "mcp__myserver__get_prompt" in names
def test_read_resource_schema_requires_uri(self):
from tools.mcp_tool import _build_utility_schemas
from tools.mcp_tool_schema import _build_utility_schemas
schemas = _build_utility_schemas("srv")
rr = next(s for s in schemas if s["handler_key"] == "read_resource")
@@ -1638,12 +1651,13 @@ class TestUtilityHandlers:
def fake_run(coro_or_factory, timeout=30):
coro = coro_or_factory() if callable(coro_or_factory) else coro_or_factory
return asyncio.run(coro)
return patch("tools.mcp_tool._run_on_mcp_loop", side_effect=fake_run)
return patch("tools.mcp_tool_loop._run_on_mcp_loop", side_effect=fake_run)
# -- list_resources --
def test_list_resources_success(self):
from tools.mcp_tool import _make_list_resources_handler, _servers
from tools.mcp_tool_handlers import _make_list_resources_handler
from tools.mcp_tool import _servers
mock_resource = SimpleNamespace(
uri="file:///tmp/test.txt", name="test.txt",
@@ -1676,7 +1690,8 @@ class TestUtilityHandlers:
# -- get_prompt --
def test_get_prompt_success(self):
from tools.mcp_tool import _make_get_prompt_handler, _servers
from tools.mcp_tool_handlers import _make_get_prompt_handler
from tools.mcp_tool import _servers
mock_msg = SimpleNamespace(
role="assistant",
@@ -1713,7 +1728,8 @@ class TestUtilityToolRegistration:
def test_utility_tools_registered(self):
"""_discover_and_register_server registers all 4 utility tools."""
from tools.registry import ToolRegistry
from tools.mcp_tool import _discover_and_register_server, _servers, MCPServerTask
from tools.mcp_tool_discovery import _discover_and_register_server
from tools.mcp_tool import _servers, MCPServerTask
mock_registry = ToolRegistry()
mock_tools = [_make_mcp_tool("read_file", "Read a file")]
@@ -1725,7 +1741,7 @@ class TestUtilityToolRegistration:
server._tools = mock_tools
return server
with patch("tools.mcp_tool._connect_server", side_effect=fake_connect), \
with patch("tools.mcp_tool_discovery._connect_server", side_effect=fake_connect), \
patch("tools.registry.registry", mock_registry):
registered = asyncio.run(
_discover_and_register_server("fs", {"command": "npx", "args": []})
@@ -1784,13 +1800,11 @@ try:
except ImportError:
ToolUseContent = _CompatType
from tools.mcp_tool import (
CreateMessageResultWithTools,
SamplingHandler,
SamplingToolsCapability,
ToolUseContent,
_safe_numeric,
)
from tools.mcp_tool import CreateMessageResultWithTools, SamplingHandler, SamplingToolsCapability, ToolUseContent
from tools.mcp_tool_common import _safe_numeric
from tools import mcp_tool_config as _mcp_config
from tools import mcp_tool_discovery as _mcp_discovery
from tools import mcp_tool_loop as _mcp_loop
# ---------------------------------------------------------------------------
@@ -2333,7 +2347,9 @@ class TestDiscoveryFailedCount:
def test_failed_server_increments_failed_count(self):
"""When _discover_and_register_server raises, failed_count increments."""
from tools.mcp_tool import discover_mcp_tools, _servers, _ensure_mcp_loop
from tools.mcp_tool_discovery import discover_mcp_tools
from tools.mcp_tool import _servers
from tools.mcp_tool_loop import _ensure_mcp_loop
fake_config = {
"good_server": {"command": "npx", "args": ["good"]},
@@ -2351,10 +2367,10 @@ class TestDiscoveryFailedCount:
_servers[name] = server
return [f"mcp__{name}__tool_a"]
with patch("tools.mcp_tool._load_mcp_config", return_value=fake_config), \
patch("tools.mcp_tool._discover_and_register_server", side_effect=fake_register), \
with patch("tools.mcp_tool_config._load_mcp_config", return_value=fake_config), \
patch("tools.mcp_tool_discovery._discover_and_register_server", side_effect=fake_register), \
patch("tools.mcp_tool._MCP_AVAILABLE", True), \
patch("tools.mcp_tool._existing_tool_names", return_value=["mcp__good_server__tool_a"]):
patch("tools.mcp_tool_registration._existing_tool_names", return_value=["mcp__good_server__tool_a"]):
_ensure_mcp_loop()
# Capture the logger to verify failed_count in summary
@@ -2377,7 +2393,9 @@ class TestDiscoveryFailedCount:
def test_ok_servers_excludes_failures(self):
"""ok_servers count correctly excludes failed servers."""
from tools.mcp_tool import discover_mcp_tools, _servers, _ensure_mcp_loop
from tools.mcp_tool_discovery import discover_mcp_tools
from tools.mcp_tool import _servers
from tools.mcp_tool_loop import _ensure_mcp_loop
fake_config = {
"ok1": {"command": "npx", "args": ["ok1"]},
@@ -2395,10 +2413,10 @@ class TestDiscoveryFailedCount:
_servers[name] = server
return [f"mcp__{name}__t"]
with patch("tools.mcp_tool._load_mcp_config", return_value=fake_config), \
patch("tools.mcp_tool._discover_and_register_server", side_effect=selective_register), \
with patch("tools.mcp_tool_config._load_mcp_config", return_value=fake_config), \
patch("tools.mcp_tool_discovery._discover_and_register_server", side_effect=selective_register), \
patch("tools.mcp_tool._MCP_AVAILABLE", True), \
patch("tools.mcp_tool._existing_tool_names", return_value=["mcp__ok1__t", "mcp__ok2__t"]):
patch("tools.mcp_tool_registration._existing_tool_names", return_value=["mcp__ok1__t", "mcp__ok2__t"]):
_ensure_mcp_loop()
with patch("tools.mcp_tool_discovery.logger") as mock_logger:
@@ -2431,7 +2449,8 @@ class TestMCPSelectiveToolLoading:
def _run_discover(self, name, tool_names, config, session=None):
from tools.registry import ToolRegistry
from tools.mcp_tool import _discover_and_register_server, _servers
from tools.mcp_tool_discovery import _discover_and_register_server
from tools.mcp_tool import _servers
mock_registry = ToolRegistry()
server = self._make_server(name, tool_names, session=session)
@@ -2440,7 +2459,7 @@ class TestMCPSelectiveToolLoading:
return server
async def run():
with patch("tools.mcp_tool._connect_server", side_effect=fake_connect), \
with patch("tools.mcp_tool_discovery._connect_server", side_effect=fake_connect), \
patch("tools.registry.registry", mock_registry), \
patch("toolsets.create_custom_toolset"):
return await _discover_and_register_server(name, config)
@@ -2487,7 +2506,7 @@ class TestMCPSelectiveToolLoading:
assert registered == []
def test_enabled_false_skips_connection_attempt(self):
from tools.mcp_tool import discover_mcp_tools
from tools.mcp_tool_discovery import discover_mcp_tools
connect_called = []
@@ -2507,8 +2526,8 @@ class TestMCPSelectiveToolLoading:
with patch("tools.mcp_tool._MCP_AVAILABLE", True), \
patch("tools.mcp_tool._servers", {}), \
patch("tools.mcp_tool._load_mcp_config", return_value=fake_config), \
patch("tools.mcp_tool._connect_server", side_effect=fake_connect), \
patch("tools.mcp_tool_config._load_mcp_config", return_value=fake_config), \
patch("tools.mcp_tool_discovery._connect_server", side_effect=fake_connect), \
patch("toolsets.TOOLSETS", fake_toolsets):
result = discover_mcp_tools()
@@ -2548,7 +2567,8 @@ class TestMCPBuiltinCollisionGuard:
def test_mcp_tool_skipped_when_builtin_exists(self):
"""An MCP tool whose prefixed name collides with a built-in is skipped."""
from tools.registry import ToolRegistry
from tools.mcp_tool import _discover_and_register_server, _servers, MCPServerTask
from tools.mcp_tool_discovery import _discover_and_register_server
from tools.mcp_tool import _servers, MCPServerTask
mock_registry = ToolRegistry()
@@ -2573,7 +2593,7 @@ class TestMCPBuiltinCollisionGuard:
server._tools = mock_tools
return server
with patch("tools.mcp_tool._connect_server", side_effect=fake_connect), \
with patch("tools.mcp_tool_discovery._connect_server", side_effect=fake_connect), \
patch("tools.registry.registry", mock_registry):
registered = asyncio.run(
_discover_and_register_server("abc", {"command": "test", "args": []})
@@ -2588,7 +2608,8 @@ class TestMCPBuiltinCollisionGuard:
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
from tools.mcp_tool_discovery import _discover_and_register_server
from tools.mcp_tool import _servers, MCPServerTask
mock_registry = ToolRegistry()
@@ -2612,7 +2633,7 @@ class TestMCPBuiltinCollisionGuard:
server._tools = mock_tools
return server
with patch("tools.mcp_tool._connect_server", side_effect=fake_connect), \
with patch("tools.mcp_tool_discovery._connect_server", side_effect=fake_connect), \
patch("tools.registry.registry", mock_registry):
registered = asyncio.run(
_discover_and_register_server("srv", {"command": "test", "args": []})
@@ -2638,7 +2659,7 @@ class TestSanitizeMcpNameComponent:
"""Verify sanitize_mcp_name_component handles all edge cases."""
def test_hyphens_replaced(self):
from tools.mcp_tool import sanitize_mcp_name_component
from tools.mcp_tool_schema import sanitize_mcp_name_component
assert sanitize_mcp_name_component("my-server") == "my_server"
@@ -2670,7 +2691,7 @@ class TestRegisterMcpServers:
"""Verify the new register_mcp_servers() public API."""
def test_mcp_not_available_returns_empty(self):
from tools.mcp_tool import register_mcp_servers
from tools.mcp_tool_discovery import register_mcp_servers
with patch("tools.mcp_tool._MCP_AVAILABLE", False):
result = register_mcp_servers({"srv": {"command": "test"}})
@@ -2678,7 +2699,9 @@ class TestRegisterMcpServers:
def test_connects_new_servers(self):
from tools.mcp_tool import register_mcp_servers, _servers, _ensure_mcp_loop
from tools.mcp_tool_discovery import register_mcp_servers
from tools.mcp_tool import _servers
from tools.mcp_tool_loop import _ensure_mcp_loop
fake_config = {"my_server": {"command": "npx", "args": ["test"]}}
@@ -2689,8 +2712,8 @@ class TestRegisterMcpServers:
return ["mcp__my_server__tool1"]
with patch("tools.mcp_tool._MCP_AVAILABLE", True), \
patch("tools.mcp_tool._discover_and_register_server", side_effect=fake_register), \
patch("tools.mcp_tool._existing_tool_names", return_value=["mcp__my_server__tool1"]):
patch("tools.mcp_tool_discovery._discover_and_register_server", side_effect=fake_register), \
patch("tools.mcp_tool_registration._existing_tool_names", return_value=["mcp__my_server__tool1"]):
_ensure_mcp_loop()
result = register_mcp_servers(fake_config)
@@ -2699,9 +2722,9 @@ class TestRegisterMcpServers:
def test_skips_servers_already_connecting(self):
"""Servers in _server_connecting must not be spawned again (#58862)."""
from tools.mcp_tool import (
register_mcp_servers, _servers, _server_connecting, _ensure_mcp_loop,
)
from tools.mcp_tool_discovery import register_mcp_servers
from tools.mcp_tool import _servers, _server_connecting
from tools.mcp_tool_loop import _ensure_mcp_loop
fake_config = {"my_srv": {"command": "npx", "args": ["test"]}}
@@ -2718,9 +2741,9 @@ class TestRegisterMcpServers:
try:
with patch("tools.mcp_tool._MCP_AVAILABLE", True), \
patch("tools.mcp_tool._discover_and_register_server", side_effect=fake_register), \
patch("tools.mcp_tool._existing_tool_names", return_value=[]), \
patch("tools.mcp_tool._connect_cooldown_active", return_value=False):
patch("tools.mcp_tool_discovery._discover_and_register_server", side_effect=fake_register), \
patch("tools.mcp_tool_registration._existing_tool_names", return_value=[]), \
patch("tools.mcp_tool_discovery._connect_cooldown_active", return_value=False):
_ensure_mcp_loop()
result = register_mcp_servers(fake_config)
@@ -2736,10 +2759,9 @@ class TestRegisterMcpServers:
def test_clears_stale_connecting_on_timeout(self):
"""Stale entries in _server_connecting are cleaned up after timeout (#58862)."""
from tools.mcp_tool import (
register_mcp_servers, _servers, _server_connecting,
_server_connect_errors, _ensure_mcp_loop,
)
from tools.mcp_tool_discovery import register_mcp_servers
from tools.mcp_tool import _servers, _server_connecting, _server_connect_errors
from tools.mcp_tool_loop import _ensure_mcp_loop
fake_config = {
"srv_a": {"command": "npx", "args": ["a"]},
@@ -2750,9 +2772,9 @@ class TestRegisterMcpServers:
_server_connecting.add("srv_a")
with patch("tools.mcp_tool._MCP_AVAILABLE", True), \
patch("tools.mcp_tool._run_on_mcp_loop", side_effect=TimeoutError("timed out")), \
patch("tools.mcp_tool._existing_tool_names", return_value=[]), \
patch("tools.mcp_tool._connect_cooldown_active", return_value=False):
patch("tools.mcp_tool_loop._run_on_mcp_loop", side_effect=TimeoutError("timed out")), \
patch("tools.mcp_tool_registration._existing_tool_names", return_value=[]), \
patch("tools.mcp_tool_discovery._connect_cooldown_active", return_value=False):
_ensure_mcp_loop()
with pytest.raises(TimeoutError):
@@ -2779,10 +2801,8 @@ class TestMcpParallelToolCalls:
def test_is_mcp_tool_parallel_safe_with_flag(self):
"""MCP tool from a parallel-safe server returns True."""
from tools.mcp_tool import (
is_mcp_tool_parallel_safe, _mcp_tool_server_names,
_parallel_safe_servers, _lock,
)
from tools.mcp_tool_discovery import is_mcp_tool_parallel_safe
from tools.mcp_tool import _mcp_tool_server_names, _parallel_safe_servers, _lock
with _lock:
_parallel_safe_servers.add("docs")
_mcp_tool_server_names["mcp__docs__search"] = "docs"
@@ -2803,10 +2823,9 @@ class TestMcpParallelToolCalls:
def test_register_mcp_servers_tracks_parallel_flag(self):
"""register_mcp_servers populates _parallel_safe_servers from config."""
from tools.mcp_tool import (
register_mcp_servers, _parallel_safe_servers, _lock,
sanitize_mcp_name_component,
)
from tools.mcp_tool_discovery import register_mcp_servers
from tools.mcp_tool import _parallel_safe_servers, _lock
from tools.mcp_tool_schema import sanitize_mcp_name_component
fake_config = {
"parallel_srv": {
"command": "echo",
@@ -2822,9 +2841,9 @@ class TestMcpParallelToolCalls:
},
}
with patch("tools.mcp_tool._MCP_AVAILABLE", True), \
patch("tools.mcp_tool._ensure_mcp_loop"), \
patch("tools.mcp_tool._run_on_mcp_loop"), \
patch("tools.mcp_tool._existing_tool_names", return_value=[]):
patch("tools.mcp_tool_loop._ensure_mcp_loop"), \
patch("tools.mcp_tool_loop._run_on_mcp_loop"), \
patch("tools.mcp_tool_registration._existing_tool_names", return_value=[]):
register_mcp_servers(fake_config)
with _lock:
@@ -2872,10 +2891,8 @@ class TestMCPDiscoveryCrossProcessLock:
def test_lock_acquired_path(self, tmp_path):
"""Lock acquired -> discovery runs normally, lock released at end."""
from tools.mcp_tool import (
_LockCookie,
discover_mcp_tools,
)
from tools.mcp_tool_loop import _LockCookie
from tools.mcp_tool_discovery import discover_mcp_tools
lock_file = tmp_path / ".mcp-discovery.lock"
fh = open(lock_file, "w", encoding="utf-8")
@@ -2886,29 +2903,26 @@ class TestMCPDiscoveryCrossProcessLock:
mock_config = {"test_srv": {"command": "echo", "enabled": True}}
with patch.object(cookie, "release", wraps=cookie.release) as release_spy:
with patch("tools.mcp_tool._try_acquire_mcp_discovery_lock", mock_acquire), \
with patch("tools.mcp_tool_loop._try_acquire_mcp_discovery_lock", mock_acquire), \
patch("tools.mcp_tool._MCP_AVAILABLE", True), \
patch("tools.mcp_tool._load_mcp_config", return_value=mock_config), \
patch("tools.mcp_tool.register_mcp_servers", return_value=["mcp__test_srv__ping"]) as reg_spy:
patch("tools.mcp_tool_config._load_mcp_config", return_value=mock_config), \
patch("tools.mcp_tool_discovery.register_mcp_servers", return_value=["mcp__test_srv__ping"]) as reg_spy:
result = discover_mcp_tools()
assert result == ["mcp__test_srv__ping"]
release_spy.assert_called_once()
def test_lock_held_retries_exhausted_fallback(self):
"""All retry attempts see lock held -> runs discovery unguarded."""
from tools.mcp_tool import (
_LOCK_UNAVAILABLE,
discover_mcp_tools,
_MCP_DISCOVERY_LOCK_MAX_RETRIES,
)
from tools.mcp_tool import _LOCK_UNAVAILABLE, _MCP_DISCOVERY_LOCK_MAX_RETRIES
from tools.mcp_tool_discovery import discover_mcp_tools
mock_config = {"test_srv": {"command": "echo", "enabled": True}}
# Every attempt returns None (lock held)
with patch("tools.mcp_tool._try_acquire_mcp_discovery_lock", return_value=None), \
with patch("tools.mcp_tool_loop._try_acquire_mcp_discovery_lock", return_value=None), \
patch("tools.mcp_tool._MCP_AVAILABLE", True), \
patch("tools.mcp_tool._load_mcp_config", return_value=mock_config), \
patch("tools.mcp_tool.register_mcp_servers") as reg_spy, \
patch("tools.mcp_tool._existing_tool_names", return_value=[]):
patch("tools.mcp_tool_config._load_mcp_config", return_value=mock_config), \
patch("tools.mcp_tool_discovery.register_mcp_servers") as reg_spy, \
patch("tools.mcp_tool_registration._existing_tool_names", return_value=[]):
result = discover_mcp_tools()
# Must still run local discovery
reg_spy.assert_called_once_with(mock_config)
@@ -2930,7 +2944,7 @@ class TestMCPDiscoveryCrossProcessLock:
fh = open(lock_path, "w", encoding="utf-8")
with patch.dict("sys.modules", {"fcntl": mock_fcntl}), \
patch("tools.mcp_tool.os.name", "posix"):
from tools.mcp_tool import _acquire_lock_on_fh
from tools.mcp_tool_loop import _acquire_lock_on_fh
result = _acquire_lock_on_fh(fh)
assert result is True
mock_fcntl.flock.assert_called_once_with(
@@ -2962,7 +2976,7 @@ class TestRedirectHeaderStripper:
def test_default_strips_only_authorization(self):
import httpx
from tools.mcp_tool import _make_redirect_header_stripper
from tools.mcp_tool_errors import _make_redirect_header_stripper
hook = _make_redirect_header_stripper(
httpx.URL("https://origin.example.test/mcp")
@@ -2977,7 +2991,7 @@ class TestRedirectHeaderStripper:
def test_strict_strips_configured_headers_cross_origin(self):
import httpx
from tools.mcp_tool import _make_redirect_header_stripper
from tools.mcp_tool_errors import _make_redirect_header_stripper
hook = _make_redirect_header_stripper(
httpx.URL("https://origin.example.test/mcp"),
@@ -2996,7 +3010,7 @@ class TestRedirectHeaderStripper:
def test_same_origin_redirect_keeps_headers(self):
import httpx
from tools.mcp_tool import _make_redirect_header_stripper
from tools.mcp_tool_errors import _make_redirect_header_stripper
hook = _make_redirect_header_stripper(
httpx.URL("https://origin.example.test/mcp"),
+7 -5
View File
@@ -14,10 +14,11 @@ import pytest
pytest.importorskip("mcp.client.auth.oauth2")
from tools import mcp_tool_loop as _mcp_loop # noqa: E402
def test_is_auth_error_detects_oauth_flow_error():
from tools.mcp_tool import _is_auth_error
from tools.mcp_tool_errors import _is_auth_error
from mcp.client.auth import OAuthFlowError
assert _is_auth_error(OAuthFlowError("expired")) is True
@@ -28,7 +29,7 @@ def test_call_tool_handler_returns_needs_reauth_on_unrecoverable_401(monkeypatch
handler returns a structured needs_reauth error (not a generic failure)."""
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
from tools.mcp_tool import _make_tool_handler
from tools.mcp_tool_handlers import _make_tool_handler
from tools.mcp_oauth_manager import get_manager, reset_manager_for_tests
from mcp.client.auth import OAuthFlowError
@@ -53,7 +54,7 @@ def test_call_tool_handler_returns_needs_reauth_on_unrecoverable_401(monkeypatch
mcp_tool._server_error_counts.pop("srv", None)
# Ensure the MCP loop exists (run_on_mcp_loop needs it)
mcp_tool._ensure_mcp_loop()
_mcp_loop._ensure_mcp_loop()
# Force handle_401 to return False (no recovery available)
mgr = get_manager()
@@ -78,7 +79,7 @@ def test_call_tool_handler_returns_needs_reauth_on_unrecoverable_401(monkeypatch
def test_call_tool_handler_non_auth_error_still_generic(monkeypatch, tmp_path):
"""Non-auth exceptions still surface via the generic error path, not needs_reauth."""
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
from tools.mcp_tool import _make_tool_handler
from tools.mcp_tool_handlers import _make_tool_handler
server = MagicMock()
server.name = "srv"
@@ -91,9 +92,10 @@ def test_call_tool_handler_non_auth_error_still_generic(monkeypatch, tmp_path):
server.session = session
from tools import mcp_tool
from tools import mcp_tool_loop as _mcp_loop
mcp_tool._servers["srv"] = server
mcp_tool._server_error_counts.pop("srv", None)
mcp_tool._ensure_mcp_loop()
_mcp_loop._ensure_mcp_loop()
try:
handler = _make_tool_handler("srv", "tool1", 10.0)
+5 -3
View File
@@ -5,7 +5,9 @@ from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
from tools.mcp_tool import MCPServerTask, _format_connect_error, _resolve_stdio_command, _MCP_AVAILABLE
from tools.mcp_tool import MCPServerTask, _MCP_AVAILABLE
from tools.mcp_tool_errors import _format_connect_error
from tools.mcp_tool_config import _resolve_stdio_command
# Ensure the mcp module symbols exist for patching even when the SDK isn't installed
if not _MCP_AVAILABLE:
@@ -25,7 +27,7 @@ def test_resolve_stdio_command_falls_back_to_hermes_node_bin(tmp_path):
npx_path.write_text("#!/bin/sh\nexit 0\n", encoding="utf-8")
npx_path.chmod(0o755)
with patch("tools.mcp_tool.shutil.which", return_value=None), \
with patch("tools.mcp_tool_config.shutil.which", return_value=None), \
patch.dict("os.environ", {"HERMES_HOME": str(tmp_path)}, clear=False):
command, env = _resolve_stdio_command("npx", {"PATH": "/usr/bin"})
@@ -55,7 +57,7 @@ def test_resolve_stdio_command_falls_back_to_usr_local_bin():
def _fake_access(path, _mode):
return path == target
with patch("tools.mcp_tool.shutil.which", return_value=None), \
with patch("tools.mcp_tool_config.shutil.which", return_value=None), \
patch("tools.mcp_tool.os.path.isfile", side_effect=_fake_isfile), \
patch("tools.mcp_tool.os.access", side_effect=_fake_access):
command, env = _resolve_stdio_command("npx", {"PATH": "/opt/data/bin:/usr/bin:/bin"})
+18 -13
View File
@@ -17,6 +17,7 @@ import time
from unittest.mock import MagicMock
import pytest
from tools import mcp_tool_loop as _mcp_loop
# ---------------------------------------------------------------------------
@@ -26,7 +27,7 @@ import pytest
def test_is_session_expired_detects_invalid_or_expired_session():
"""Reporter's exact wpcom-mcp error message (#13383)."""
from tools.mcp_tool import _is_session_expired_error
from tools.mcp_tool_errors import _is_session_expired_error
exc = RuntimeError("Invalid params: Invalid or expired session")
assert _is_session_expired_error(exc) is True
@@ -34,7 +35,7 @@ def test_is_session_expired_detects_invalid_or_expired_session():
def test_is_session_expired_detects_expired_session_variant():
"""Generic ``session expired`` / ``expired session`` phrasings used
by other SDK servers."""
from tools.mcp_tool import _is_session_expired_error
from tools.mcp_tool_errors import _is_session_expired_error
assert _is_session_expired_error(RuntimeError("Session expired")) is True
assert _is_session_expired_error(RuntimeError("expired session: abc")) is True
@@ -42,7 +43,7 @@ def test_is_session_expired_detects_expired_session_variant():
def test_is_session_expired_detects_session_not_found():
"""Server-side GC produces ``session not found`` / ``unknown session``
on some implementations."""
from tools.mcp_tool import _is_session_expired_error
from tools.mcp_tool_errors import _is_session_expired_error
assert _is_session_expired_error(RuntimeError("session not found")) is True
assert _is_session_expired_error(RuntimeError("Unknown session: abc123")) is True
@@ -50,10 +51,11 @@ def test_is_session_expired_detects_session_not_found():
def test_is_session_expired_traversal_is_budget_bounded():
"""Pathologically long chains stop at the node budget without spinning."""
import tools.mcp_tool as mcp_mod
from tools.mcp_tool import _is_session_expired_error
from tools import mcp_tool_errors as _mcp_errors
from tools.mcp_tool_errors import _is_session_expired_error
exc: BaseException = RuntimeError("leaf")
for i in range(mcp_mod._EXC_TRAVERSAL_MAX_NODES * 2):
for i in range(_mcp_errors._EXC_TRAVERSAL_MAX_NODES * 2):
wrapper = RuntimeError(f"layer {i}")
wrapper.__cause__ = exc
exc = wrapper
@@ -75,7 +77,7 @@ def _install_stub_server(name: str = "wpcom"):
the event fires."""
from tools import mcp_tool
mcp_tool._ensure_mcp_loop()
_mcp_loop._ensure_mcp_loop()
server = MagicMock()
server.name = name
@@ -141,10 +143,11 @@ def test_call_tool_handler_rebuilds_configured_server_transport(
"""The real server run loop selects and rebuilds its configured transport."""
from anyio import ClosedResourceError
from tools import mcp_tool
from tools.mcp_tool import MCPServerTask, _make_tool_handler
from tools.mcp_tool import MCPServerTask
from tools.mcp_tool_handlers import _make_tool_handler
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
mcp_tool._ensure_mcp_loop()
_mcp_loop._ensure_mcp_loop()
transport_ready = threading.Event()
routes = []
configs = []
@@ -219,9 +222,10 @@ def test_session_expired_retry_waits_for_new_session(monkeypatch, tmp_path):
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
from tools import mcp_tool
from tools.mcp_tool import _make_tool_handler
from tools import mcp_tool_loop as _mcp_loop
from tools.mcp_tool_handlers import _make_tool_handler
mcp_tool._ensure_mcp_loop()
_mcp_loop._ensure_mcp_loop()
server = MagicMock()
server.name = "hindsight"
ready_flag = threading.Event()
@@ -290,7 +294,7 @@ def test_session_expired_handler_returns_none_without_loop(monkeypatch):
race), the handler must fall through cleanly instead of hanging
or raising."""
from tools import mcp_tool
from tools.mcp_tool import _handle_session_expired_and_retry
from tools.mcp_tool_handlers import _handle_session_expired_and_retry
# Install a server stub but make the event loop unavailable.
server = MagicMock()
@@ -320,7 +324,7 @@ def test_session_expired_handler_returns_none_without_loop(monkeypatch):
def test_session_expired_handler_returns_none_without_server_record():
"""If the server has been torn down / isn't in _servers, fall
through cleanly — nothing to reconnect to."""
from tools.mcp_tool import _handle_session_expired_and_retry
from tools.mcp_tool_handlers import _handle_session_expired_and_retry
out = _handle_session_expired_and_retry(
"does-not-exist",
RuntimeError("Invalid or expired session"),
@@ -378,7 +382,8 @@ def test_non_tool_handlers_also_reconnect_on_session_expired(
setattr(server.session, session_method, _sequence)
factory = getattr(mcp_tool, handler_factory)
from tools import mcp_tool_handlers as _mcp_handlers
factory = getattr(_mcp_handlers, handler_factory)
# list_resources / list_prompts take (server_name, timeout).
# read_resource / get_prompt take the same signature.
try:
+28 -26
View File
@@ -25,6 +25,8 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from tools import mcp_tool
from tools import mcp_tool_handlers as _mcp_handlers
from tools import mcp_tool_registration as _mcp_registration
class _FakeContentBlock:
@@ -63,7 +65,7 @@ def fake_session():
)
server = SimpleNamespace(session=session, _rpc_lock=None)
with patch.dict(mcp_tool._servers, {"srv": server}), \
patch("tools.mcp_tool._run_on_mcp_loop",
patch("tools.mcp_tool_loop._run_on_mcp_loop",
side_effect=_fake_run_on_mcp_loop), \
patch.dict(mcp_tool._server_error_counts, {}, clear=True):
yield session
@@ -94,9 +96,9 @@ class TestTrustGateAtCallTime:
"""Approval consulted; 'accept' lets the RPC through."""
_set_trust("srv", "untrusted")
# No readOnlyHint recorded for delete_repo → write-capable.
handler = mcp_tool._make_tool_handler("srv", "delete_repo", 30.0)
handler = _mcp_handlers._make_tool_handler("srv", "delete_repo", 30.0)
with patch(
"tools.approval.request_elicitation_consent",
"tools.approval_prompt.request_elicitation_consent",
return_value="accept",
) as consent:
raw = handler({"repo": "x"})
@@ -107,9 +109,9 @@ class TestTrustGateAtCallTime:
def test_denied_approval_blocks_rpc(self, fake_session):
"""'decline' blocks the call — the RPC must never fire."""
_set_trust("srv", "untrusted")
handler = mcp_tool._make_tool_handler("srv", "delete_repo", 30.0)
handler = _mcp_handlers._make_tool_handler("srv", "delete_repo", 30.0)
with patch(
"tools.approval.request_elicitation_consent",
"tools.approval_prompt.request_elicitation_consent",
return_value="decline",
):
raw = handler({"repo": "x"})
@@ -123,9 +125,9 @@ class TestTrustGateAtCallTime:
"""readOnlyHint=True tools pass without consulting approval."""
_set_trust("srv", "untrusted")
_set_read_only("srv", "list_repos", True)
handler = mcp_tool._make_tool_handler("srv", "list_repos", 30.0)
handler = _mcp_handlers._make_tool_handler("srv", "list_repos", 30.0)
with patch(
"tools.approval.request_elicitation_consent"
"tools.approval_prompt.request_elicitation_consent"
) as consent:
raw = handler({})
consent.assert_not_called()
@@ -136,9 +138,9 @@ class TestTrustGateAtCallTime:
):
"""trust: full (and the default) never consults approval."""
_set_trust("srv", "full")
handler = mcp_tool._make_tool_handler("srv", "delete_repo", 30.0)
handler = _mcp_handlers._make_tool_handler("srv", "delete_repo", 30.0)
with patch(
"tools.approval.request_elicitation_consent"
"tools.approval_prompt.request_elicitation_consent"
) as consent:
raw = handler({"repo": "x"})
consent.assert_not_called()
@@ -146,9 +148,9 @@ class TestTrustGateAtCallTime:
def test_unconfigured_server_defaults_to_full_trust(self, fake_session):
"""Backward compat: servers with no trust key behave as before."""
handler = mcp_tool._make_tool_handler("srv", "delete_repo", 30.0)
handler = _mcp_handlers._make_tool_handler("srv", "delete_repo", 30.0)
with patch(
"tools.approval.request_elicitation_consent"
"tools.approval_prompt.request_elicitation_consent"
) as consent:
raw = handler({"repo": "x"})
consent.assert_not_called()
@@ -158,9 +160,9 @@ class TestTrustGateAtCallTime:
"""An explicit readOnlyHint=False is write-capable."""
_set_trust("srv", "untrusted")
_set_read_only("srv", "write_file", False)
handler = mcp_tool._make_tool_handler("srv", "write_file", 30.0)
handler = _mcp_handlers._make_tool_handler("srv", "write_file", 30.0)
with patch(
"tools.approval.request_elicitation_consent",
"tools.approval_prompt.request_elicitation_consent",
return_value="decline",
) as consent:
handler({"path": "/etc/passwd"})
@@ -170,9 +172,9 @@ class TestTrustGateAtCallTime:
def test_approval_exception_fails_closed(self, fake_session):
"""Any exception in the consent path blocks the call."""
_set_trust("srv", "untrusted")
handler = mcp_tool._make_tool_handler("srv", "delete_repo", 30.0)
handler = _mcp_handlers._make_tool_handler("srv", "delete_repo", 30.0)
with patch(
"tools.approval.request_elicitation_consent",
"tools.approval_prompt.request_elicitation_consent",
side_effect=RuntimeError("approval backend down"),
):
raw = handler({"repo": "x"})
@@ -183,14 +185,14 @@ class TestTrustGateAtCallTime:
class TestTrustNormalization:
def test_unknown_trust_value_treated_as_untrusted(self):
"""Garbage trust strings fail closed to untrusted."""
assert mcp_tool._normalize_server_trust("banana") == "untrusted"
assert _mcp_registration._normalize_server_trust("banana") == "untrusted"
def test_known_values(self):
assert mcp_tool._normalize_server_trust("full") == "full"
assert mcp_tool._normalize_server_trust("UNTRUSTED") == "untrusted"
assert mcp_tool._normalize_server_trust(" Full ") == "full"
assert _mcp_registration._normalize_server_trust("full") == "full"
assert _mcp_registration._normalize_server_trust("UNTRUSTED") == "untrusted"
assert _mcp_registration._normalize_server_trust(" Full ") == "full"
# Missing key → default full (backward compatible; documented).
assert mcp_tool._normalize_server_trust(None) == "full"
assert _mcp_registration._normalize_server_trust(None) == "full"
class TestAnnotationCaptureAtDiscovery:
@@ -221,8 +223,8 @@ class TestAnnotationCaptureAtDiscovery:
"tools": {"resources": False, "prompts": False},
}
with patch("tools.registry.registry", ToolRegistry()), \
patch("tools.mcp_tool._track_mcp_tool_server"):
mcp_tool._register_server_tools("srv", server, config)
patch("tools.mcp_tool_registration._track_mcp_tool_server"):
_mcp_registration._register_server_tools("srv", server, config)
assert mcp_tool._server_trust_levels["srv"] == "untrusted"
hints = mcp_tool._tool_read_only_hints["srv"]
@@ -233,15 +235,15 @@ class TestAnnotationCaptureAtDiscovery:
def test_dict_annotations_supported(self):
"""Cached/JSON annotations arrive as plain dicts."""
assert mcp_tool._annotation_read_only_hint(
assert _mcp_registration._annotation_read_only_hint(
SimpleNamespace(annotations={"readOnlyHint": True})
) is True
assert mcp_tool._annotation_read_only_hint(
assert _mcp_registration._annotation_read_only_hint(
SimpleNamespace(annotations={"readOnlyHint": "yes"})
) is False # non-bool truthy → NOT read-only (hint must be True)
assert mcp_tool._annotation_read_only_hint(
assert _mcp_registration._annotation_read_only_hint(
SimpleNamespace(annotations=None)
) is False
assert mcp_tool._annotation_read_only_hint(
assert _mcp_registration._annotation_read_only_hint(
SimpleNamespace()
) is False
@@ -73,7 +73,7 @@ class TestCapabilityGatedRegistration:
"""Context7-shaped server (tools only, no prompts / resources) should
get zero utility stubs registered — this is the exact scenario
from the #18051 bug report."""
from tools.mcp_tool import _select_utility_schemas
from tools.mcp_tool_registration import _select_utility_schemas
server = _make_fake_server(
initialize_result=_make_init_result(resources=False, prompts=False)
@@ -85,7 +85,7 @@ class TestCapabilityGatedRegistration:
)
def test_resources_only_server_gets_resource_stubs_only(self):
from tools.mcp_tool import _select_utility_schemas
from tools.mcp_tool_registration import _select_utility_schemas
server = _make_fake_server(
initialize_result=_make_init_result(resources=True, prompts=False)
@@ -95,7 +95,7 @@ class TestCapabilityGatedRegistration:
def test_fully_capable_server_gets_all_four_stubs(self):
from tools.mcp_tool import _select_utility_schemas
from tools.mcp_tool_registration import _select_utility_schemas
server = _make_fake_server(
initialize_result=_make_init_result(resources=True, prompts=True)
@@ -111,7 +111,7 @@ class TestConfigFilterStillApplies:
must continue to override even when the server DOES advertise the capability."""
def test_config_disables_resources_even_when_advertised(self):
from tools.mcp_tool import _select_utility_schemas
from tools.mcp_tool_registration import _select_utility_schemas
server = _make_fake_server(
initialize_result=_make_init_result(resources=True, prompts=True)
@@ -124,7 +124,7 @@ class TestConfigFilterStillApplies:
assert _handler_keys(selected) == {"list_prompts", "get_prompt"}
def test_config_disables_prompts_even_when_advertised(self):
from tools.mcp_tool import _select_utility_schemas
from tools.mcp_tool_registration import _select_utility_schemas
server = _make_fake_server(
initialize_result=_make_init_result(resources=True, prompts=True)
@@ -143,7 +143,7 @@ class TestLegacyFallback:
check so pre-existing tests and servers keep working."""
def test_no_initialize_result_falls_back_to_hasattr_check(self):
from tools.mcp_tool import _select_utility_schemas
from tools.mcp_tool_registration import _select_utility_schemas
server = _make_fake_server(initialize_result=None)
# With the legacy fallback, session.spec includes all four methods,
@@ -156,7 +156,7 @@ class TestLegacyFallback:
def test_no_initialize_result_respects_session_spec(self):
"""Legacy fallback still filters by ``hasattr(session, method)``, so
a session whose spec lacks a method is correctly skipped."""
from tools.mcp_tool import _select_utility_schemas
from tools.mcp_tool_registration import _select_utility_schemas
server = _make_fake_server(initialize_result=None)
# Override session to a spec that only has list_resources
+11 -10
View File
@@ -12,6 +12,7 @@ import threading
import types
from tools import mcp_tool
from tools import mcp_tool_agent as _mcp_agent
def _tool(name):
@@ -37,7 +38,7 @@ def test_refresh_adds_late_landing_tools(monkeypatch):
import model_tools
monkeypatch.setattr(model_tools, "get_tool_definitions", lambda **kw: new_defs)
added = mcp_tool.refresh_agent_mcp_tools(agent)
added = _mcp_agent.refresh_agent_mcp_tools(agent)
assert added == {"mcp_granola_get_account_info"}
assert "mcp_granola_get_account_info" in agent.valid_tool_names
@@ -77,7 +78,7 @@ def test_refresh_preserves_memory_provider_and_context_engine_tools(monkeypatch)
lambda **kw: [_tool("read_file"), _tool("mcp_new_server_tool")],
)
added = mcp_tool.refresh_agent_mcp_tools(agent)
added = _mcp_agent.refresh_agent_mcp_tools(agent)
# The new MCP tool landed AND the injected families survived.
assert "mcp_new_server_tool" in agent.valid_tool_names
@@ -106,7 +107,7 @@ def test_refresh_does_not_reinject_disabled_memory_provider_tools(monkeypatch):
lambda **kw: [_tool("read_file")],
)
mcp_tool.refresh_agent_mcp_tools(agent)
_mcp_agent.refresh_agent_mcp_tools(agent)
assert "memory_search" not in agent.valid_tool_names
assert all(t["function"]["name"] != "memory_search" for t in agent.tools)
@@ -128,7 +129,7 @@ def test_refresh_respects_context_engine_toolset_gate(monkeypatch):
lambda **kw: [_tool("read_file"), _tool("mcp_new_tool")],
)
mcp_tool.refresh_agent_mcp_tools(agent)
_mcp_agent.refresh_agent_mcp_tools(agent)
assert "mcp_new_tool" in agent.valid_tool_names # MCP tool still lands
assert "lcm_grep" not in agent.valid_tool_names # gated out (#5544)
@@ -148,7 +149,7 @@ def test_refreshed_tool_is_callable_through_valid_tool_names_guard(monkeypatch):
# Before refresh the run loop would reject the call ("Tool does not exist").
assert "mcp_granola_list_meetings" not in agent.valid_tool_names
mcp_tool.refresh_agent_mcp_tools(agent)
_mcp_agent.refresh_agent_mcp_tools(agent)
# After refresh the same guard accepts it AND it's in the tools= payload.
assert "mcp_granola_list_meetings" in agent.valid_tool_names
@@ -184,7 +185,7 @@ def test_refresh_is_thread_safe_under_concurrent_calls(monkeypatch):
def _worker():
try:
for _ in range(50):
mcp_tool.refresh_agent_mcp_tools(agent)
_mcp_agent.refresh_agent_mcp_tools(agent)
# Coherence invariant: the name set must match the tool list
# at every observation, never a torn cross-attribute state.
names = {t["function"]["name"] for t in agent.tools}
@@ -261,7 +262,7 @@ def test_preserve_prefix_carries_a_flapping_tool_forward(monkeypatch):
_serve(monkeypatch, [_tool("read_file"), _tool("terminal")])
_registered(monkeypatch, ["read_file", "browser_navigate", "terminal"])
added = mcp_tool.refresh_agent_mcp_tools(agent, preserve_prefix=True)
added = _mcp_agent.refresh_agent_mcp_tools(agent, preserve_prefix=True)
assert added == set()
assert agent.tools == before
@@ -280,7 +281,7 @@ def test_preserve_prefix_appends_late_arrivals_at_the_tail(monkeypatch):
_serve(monkeypatch, [_tool("aaa_mcp_late"), _tool("read_file"), _tool("terminal")])
_registered(monkeypatch, ["aaa_mcp_late", "read_file", "terminal"])
added = mcp_tool.refresh_agent_mcp_tools(agent, preserve_prefix=True)
added = _mcp_agent.refresh_agent_mcp_tools(agent, preserve_prefix=True)
assert added == {"aaa_mcp_late"}
assert [t["function"]["name"] for t in agent.tools] == [
@@ -309,7 +310,7 @@ def test_eviction_rebuild_restores_the_sessions_saved_tool_order(monkeypatch):
monkeypatch.setattr(registry_mod.registry, "get_entry", lambda name, **kw: entries.get(name), raising=False)
rebuilt = _agent(["read_file", "terminal"]) # probe flipped: browser_navigate gone
changed = mcp_tool.restore_agent_tool_prefix(rebuilt, saved)
changed = _mcp_agent.restore_agent_tool_prefix(rebuilt, saved)
assert changed is True
assert [t["function"]["name"] for t in rebuilt.tools] == saved
@@ -333,7 +334,7 @@ def test_reprobe_tool_availability_drops_cached_check_fn_verdicts(monkeypatch):
with model_tools._tool_defs_cache_lock:
model_tools._tool_defs_cache[("sentinel",)] = []
mcp_tool.reprobe_tool_availability()
_mcp_agent.reprobe_tool_availability()
assert registry_mod._check_fn_cached(probe) is True
assert ("sentinel",) not in model_tools._tool_defs_cache
+2 -2
View File
@@ -463,7 +463,7 @@ def test_null_plus_const_union_ordering_with_nullable_strip():
the remaining null branch by collapsing consts and keeping nullability as
a hint.
"""
from tools.mcp_tool import _normalize_mcp_input_schema
from tools.mcp_tool_schema import _normalize_mcp_input_schema
schema = {
"type": "object",
@@ -486,7 +486,7 @@ def test_null_plus_const_union_ordering_with_nullable_strip():
def test_normalize_mcp_input_schema_collapses_const_unions():
from tools.mcp_tool import _normalize_mcp_input_schema
from tools.mcp_tool_schema import _normalize_mcp_input_schema
schema = {
"type": "object",
+6 -4
View File
@@ -24,6 +24,8 @@ import threading
import pytest
import tools.mcp_tool as mcp_tool
from tools import mcp_tool_discovery as _mcp_discovery
from tools import mcp_tool_lifecycle as _mcp_lifecycle
import tui_gateway.server as srv
@@ -33,8 +35,8 @@ def reload_env(monkeypatch):
calls = {"discover": 0, "shutdown": 0}
rev_box = {"rev": "rev-a"}
monkeypatch.setattr(mcp_tool, "shutdown_mcp_servers", lambda: calls.__setitem__("shutdown", calls["shutdown"] + 1))
monkeypatch.setattr(mcp_tool, "discover_mcp_tools", lambda: calls.__setitem__("discover", calls["discover"] + 1))
monkeypatch.setattr(_mcp_lifecycle, "shutdown_mcp_servers", lambda: calls.__setitem__("shutdown", calls["shutdown"] + 1))
monkeypatch.setattr(_mcp_discovery, "discover_mcp_tools", lambda: calls.__setitem__("discover", calls["discover"] + 1))
monkeypatch.setattr(srv, "_compute_mcp_rev", lambda: rev_box["rev"])
saved = (srv._mcp_reload_gen, srv._mcp_reload_loaded_rev)
@@ -72,7 +74,7 @@ def test_failed_reload_is_an_error_and_no_generation_advance(reload_env, monkeyp
def _boom():
raise RuntimeError("flapping server")
monkeypatch.setattr(mcp_tool, "discover_mcp_tools", _boom)
monkeypatch.setattr(_mcp_discovery, "discover_mcp_tools", _boom)
envelope = _reload(rev="rev-b")
@@ -146,7 +148,7 @@ def _run_leader_follower(reload_env, monkeypatch, follower_rev):
leader_in_discovery.set()
assert release_leader.wait(timeout=10)
monkeypatch.setattr(mcp_tool, "discover_mcp_tools", _slow_discover)
monkeypatch.setattr(_mcp_discovery, "discover_mcp_tools", _slow_discover)
results: dict = {}