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:
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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])
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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,
|
||||
),
|
||||
):
|
||||
|
||||
@@ -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), \
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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, (
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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}],
|
||||
|
||||
@@ -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"]},
|
||||
})
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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"
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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"))
|
||||
|
||||
|
||||
@@ -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 == {}
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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", {
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 ───────────────────────────────────────────────────────────────────
|
||||
|
||||
@@ -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"},
|
||||
})
|
||||
|
||||
@@ -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")
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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({}))
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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)]"
|
||||
|
||||
@@ -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
@@ -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"),
|
||||
|
||||
@@ -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,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"})
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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 = {}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user