feat: scroll aware sticky prompt
This commit is contained in:
@@ -380,7 +380,7 @@ class TestStubSchemaDrift(unittest.TestCase):
|
||||
# Parameters that are internal (injected by the handler, not user-facing)
|
||||
_INTERNAL_PARAMS = {"task_id", "user_task"}
|
||||
# Parameters intentionally blocked in the sandbox
|
||||
_BLOCKED_TERMINAL_PARAMS = {"background", "pty", "notify_on_complete"}
|
||||
_BLOCKED_TERMINAL_PARAMS = {"background", "pty", "notify_on_complete", "watch_patterns"}
|
||||
|
||||
def test_stubs_cover_all_schema_params(self):
|
||||
"""Every user-facing parameter in the real schema must appear in the
|
||||
|
||||
@@ -29,8 +29,11 @@ class TestInterruptModule:
|
||||
|
||||
def test_thread_safety(self):
|
||||
"""Set from one thread targeting another thread's ident."""
|
||||
from tools.interrupt import set_interrupt, is_interrupted
|
||||
from tools.interrupt import set_interrupt, is_interrupted, _interrupted_threads, _lock
|
||||
set_interrupt(False)
|
||||
# Clear any stale thread idents left by prior tests in this worker.
|
||||
with _lock:
|
||||
_interrupted_threads.clear()
|
||||
|
||||
seen = {"value": False}
|
||||
|
||||
|
||||
@@ -6,6 +6,8 @@ All tests use mocks -- no real MCP servers or subprocesses are started.
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import threading
|
||||
import time
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
@@ -255,6 +257,77 @@ class TestToolHandler:
|
||||
finally:
|
||||
_servers.pop("test_srv", None)
|
||||
|
||||
def test_interrupted_call_returns_interrupted_error(self):
|
||||
from tools.mcp_tool import _make_tool_handler, _servers
|
||||
|
||||
mock_session = MagicMock()
|
||||
server = _make_mock_server("test_srv", session=mock_session)
|
||||
_servers["test_srv"] = server
|
||||
|
||||
try:
|
||||
handler = _make_tool_handler("test_srv", "greet", 120)
|
||||
def _interrupting_run(coro, timeout=30):
|
||||
coro.close()
|
||||
raise InterruptedError("User sent a new message")
|
||||
with patch(
|
||||
"tools.mcp_tool._run_on_mcp_loop",
|
||||
side_effect=_interrupting_run,
|
||||
):
|
||||
result = json.loads(handler({}))
|
||||
assert result == {"error": "MCP call interrupted: user sent a new message"}
|
||||
finally:
|
||||
_servers.pop("test_srv", None)
|
||||
|
||||
|
||||
class TestRunOnMCPLoopInterrupts:
|
||||
def test_interrupt_cancels_waiting_mcp_call(self):
|
||||
import tools.mcp_tool as mcp_mod
|
||||
from tools.interrupt import set_interrupt
|
||||
|
||||
loop = asyncio.new_event_loop()
|
||||
thread = threading.Thread(target=loop.run_forever, daemon=True)
|
||||
thread.start()
|
||||
|
||||
cancelled = threading.Event()
|
||||
|
||||
async def _slow_call():
|
||||
try:
|
||||
await asyncio.sleep(5)
|
||||
return "done"
|
||||
except asyncio.CancelledError:
|
||||
cancelled.set()
|
||||
raise
|
||||
|
||||
old_loop = mcp_mod._mcp_loop
|
||||
old_thread = mcp_mod._mcp_thread
|
||||
mcp_mod._mcp_loop = loop
|
||||
mcp_mod._mcp_thread = thread
|
||||
|
||||
waiter_tid = threading.current_thread().ident
|
||||
|
||||
def _interrupt_soon():
|
||||
time.sleep(0.2)
|
||||
set_interrupt(True, waiter_tid)
|
||||
|
||||
interrupter = threading.Thread(target=_interrupt_soon, daemon=True)
|
||||
interrupter.start()
|
||||
|
||||
try:
|
||||
with pytest.raises(InterruptedError, match="User sent a new message"):
|
||||
mcp_mod._run_on_mcp_loop(_slow_call(), timeout=2)
|
||||
|
||||
deadline = time.time() + 2
|
||||
while time.time() < deadline and not cancelled.is_set():
|
||||
time.sleep(0.05)
|
||||
assert cancelled.is_set()
|
||||
finally:
|
||||
set_interrupt(False, waiter_tid)
|
||||
loop.call_soon_threadsafe(loop.stop)
|
||||
thread.join(timeout=2)
|
||||
loop.close()
|
||||
mcp_mod._mcp_loop = old_loop
|
||||
mcp_mod._mcp_thread = old_thread
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tool registration (discovery + register)
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
"""Tests for the central tool registry."""
|
||||
|
||||
import json
|
||||
import threading
|
||||
|
||||
from tools.registry import ToolRegistry
|
||||
|
||||
@@ -167,6 +168,32 @@ class TestToolsetAvailability:
|
||||
)
|
||||
assert reg.get_all_tool_names() == ["a_tool", "z_tool"]
|
||||
|
||||
def test_get_registered_toolset_names(self):
|
||||
reg = ToolRegistry()
|
||||
reg.register(
|
||||
name="first", toolset="zeta", schema=_make_schema(), handler=_dummy_handler
|
||||
)
|
||||
reg.register(
|
||||
name="second", toolset="alpha", schema=_make_schema(), handler=_dummy_handler
|
||||
)
|
||||
reg.register(
|
||||
name="third", toolset="alpha", schema=_make_schema(), handler=_dummy_handler
|
||||
)
|
||||
assert reg.get_registered_toolset_names() == ["alpha", "zeta"]
|
||||
|
||||
def test_get_tool_names_for_toolset(self):
|
||||
reg = ToolRegistry()
|
||||
reg.register(
|
||||
name="z_tool", toolset="grouped", schema=_make_schema(), handler=_dummy_handler
|
||||
)
|
||||
reg.register(
|
||||
name="a_tool", toolset="grouped", schema=_make_schema(), handler=_dummy_handler
|
||||
)
|
||||
reg.register(
|
||||
name="other_tool", toolset="other", schema=_make_schema(), handler=_dummy_handler
|
||||
)
|
||||
assert reg.get_tool_names_for_toolset("grouped") == ["a_tool", "z_tool"]
|
||||
|
||||
def test_handler_exception_returns_error(self):
|
||||
reg = ToolRegistry()
|
||||
|
||||
@@ -301,6 +328,22 @@ class TestEmojiMetadata:
|
||||
assert reg.get_emoji("t") == "⚡"
|
||||
|
||||
|
||||
class TestEntryLookup:
|
||||
def test_get_entry_returns_registered_entry(self):
|
||||
reg = ToolRegistry()
|
||||
reg.register(
|
||||
name="alpha", toolset="core", schema=_make_schema("alpha"), handler=_dummy_handler
|
||||
)
|
||||
entry = reg.get_entry("alpha")
|
||||
assert entry is not None
|
||||
assert entry.name == "alpha"
|
||||
assert entry.toolset == "core"
|
||||
|
||||
def test_get_entry_returns_none_for_unknown_tool(self):
|
||||
reg = ToolRegistry()
|
||||
assert reg.get_entry("missing") is None
|
||||
|
||||
|
||||
class TestSecretCaptureResultContract:
|
||||
def test_secret_request_result_does_not_include_secret_value(self):
|
||||
result = {
|
||||
@@ -309,3 +352,141 @@ class TestSecretCaptureResultContract:
|
||||
"validated": False,
|
||||
}
|
||||
assert "secret" not in json.dumps(result).lower()
|
||||
|
||||
|
||||
class TestThreadSafety:
|
||||
def test_get_available_toolsets_uses_coherent_snapshot(self, monkeypatch):
|
||||
reg = ToolRegistry()
|
||||
reg.register(
|
||||
name="alpha",
|
||||
toolset="gated",
|
||||
schema=_make_schema("alpha"),
|
||||
handler=_dummy_handler,
|
||||
check_fn=lambda: False,
|
||||
)
|
||||
|
||||
entries, toolset_checks = reg._snapshot_state()
|
||||
|
||||
def snapshot_then_mutate():
|
||||
reg.deregister("alpha")
|
||||
return entries, toolset_checks
|
||||
|
||||
monkeypatch.setattr(reg, "_snapshot_state", snapshot_then_mutate)
|
||||
|
||||
toolsets = reg.get_available_toolsets()
|
||||
assert toolsets["gated"]["available"] is False
|
||||
assert toolsets["gated"]["tools"] == ["alpha"]
|
||||
|
||||
def test_check_tool_availability_tolerates_concurrent_register(self):
|
||||
reg = ToolRegistry()
|
||||
check_started = threading.Event()
|
||||
writer_done = threading.Event()
|
||||
errors = []
|
||||
result_holder = {}
|
||||
writer_completed_during_check = {}
|
||||
|
||||
def blocking_check():
|
||||
check_started.set()
|
||||
writer_completed_during_check["value"] = writer_done.wait(timeout=1)
|
||||
return True
|
||||
|
||||
reg.register(
|
||||
name="alpha",
|
||||
toolset="gated",
|
||||
schema=_make_schema("alpha"),
|
||||
handler=_dummy_handler,
|
||||
check_fn=blocking_check,
|
||||
)
|
||||
reg.register(
|
||||
name="beta",
|
||||
toolset="plain",
|
||||
schema=_make_schema("beta"),
|
||||
handler=_dummy_handler,
|
||||
)
|
||||
|
||||
def reader():
|
||||
try:
|
||||
result_holder["value"] = reg.check_tool_availability()
|
||||
except Exception as exc: # pragma: no cover - exercised on failure only
|
||||
errors.append(exc)
|
||||
|
||||
def writer():
|
||||
assert check_started.wait(timeout=1)
|
||||
reg.register(
|
||||
name="gamma",
|
||||
toolset="new",
|
||||
schema=_make_schema("gamma"),
|
||||
handler=_dummy_handler,
|
||||
)
|
||||
writer_done.set()
|
||||
|
||||
reader_thread = threading.Thread(target=reader)
|
||||
writer_thread = threading.Thread(target=writer)
|
||||
reader_thread.start()
|
||||
writer_thread.start()
|
||||
reader_thread.join(timeout=2)
|
||||
writer_thread.join(timeout=2)
|
||||
|
||||
assert not reader_thread.is_alive()
|
||||
assert not writer_thread.is_alive()
|
||||
assert writer_completed_during_check["value"] is True
|
||||
assert errors == []
|
||||
|
||||
available, unavailable = result_holder["value"]
|
||||
assert "gated" in available
|
||||
assert "plain" in available
|
||||
assert unavailable == []
|
||||
|
||||
def test_get_available_toolsets_tolerates_concurrent_deregister(self):
|
||||
reg = ToolRegistry()
|
||||
check_started = threading.Event()
|
||||
writer_done = threading.Event()
|
||||
errors = []
|
||||
result_holder = {}
|
||||
writer_completed_during_check = {}
|
||||
|
||||
def blocking_check():
|
||||
check_started.set()
|
||||
writer_completed_during_check["value"] = writer_done.wait(timeout=1)
|
||||
return True
|
||||
|
||||
reg.register(
|
||||
name="alpha",
|
||||
toolset="gated",
|
||||
schema=_make_schema("alpha"),
|
||||
handler=_dummy_handler,
|
||||
check_fn=blocking_check,
|
||||
)
|
||||
reg.register(
|
||||
name="beta",
|
||||
toolset="plain",
|
||||
schema=_make_schema("beta"),
|
||||
handler=_dummy_handler,
|
||||
)
|
||||
|
||||
def reader():
|
||||
try:
|
||||
result_holder["value"] = reg.get_available_toolsets()
|
||||
except Exception as exc: # pragma: no cover - exercised on failure only
|
||||
errors.append(exc)
|
||||
|
||||
def writer():
|
||||
assert check_started.wait(timeout=1)
|
||||
reg.deregister("beta")
|
||||
writer_done.set()
|
||||
|
||||
reader_thread = threading.Thread(target=reader)
|
||||
writer_thread = threading.Thread(target=writer)
|
||||
reader_thread.start()
|
||||
writer_thread.start()
|
||||
reader_thread.join(timeout=2)
|
||||
writer_thread.join(timeout=2)
|
||||
|
||||
assert not reader_thread.is_alive()
|
||||
assert not writer_thread.is_alive()
|
||||
assert writer_completed_during_check["value"] is True
|
||||
assert errors == []
|
||||
|
||||
toolsets = result_holder["value"]
|
||||
assert "gated" in toolsets
|
||||
assert toolsets["gated"]["available"] is True
|
||||
|
||||
Reference in New Issue
Block a user