fix(gateway): stop typing loops on session interrupt
This commit is contained in:
@@ -1,13 +1,18 @@
|
||||
"""Tests for the pending_event None guard in recursive _run_agent calls.
|
||||
"""Tests for pending follow-up extraction in recursive _run_agent calls.
|
||||
|
||||
When pending_event is None (Path B: pending comes from interrupt_message),
|
||||
accessing pending_event.channel_prompt previously raised AttributeError.
|
||||
This verifies the fix: channel_prompt is captured inside the
|
||||
`if pending_event is not None:` block and falls back to None otherwise.
|
||||
|
||||
Also verifies that internal control interrupt reasons like "Stop requested"
|
||||
do not get recycled into the pending-user-message follow-up path.
|
||||
"""
|
||||
|
||||
from types import SimpleNamespace
|
||||
|
||||
from gateway.run import _is_control_interrupt_message
|
||||
|
||||
|
||||
def _extract_channel_prompt(pending_event):
|
||||
"""Reproduce the fixed logic from gateway/run.py.
|
||||
@@ -21,6 +26,15 @@ def _extract_channel_prompt(pending_event):
|
||||
return next_channel_prompt
|
||||
|
||||
|
||||
def _extract_pending_text(interrupted, pending_event, interrupt_message):
|
||||
"""Reproduce the fixed pending-text selection from gateway/run.py."""
|
||||
if interrupted and pending_event is None and interrupt_message:
|
||||
if _is_control_interrupt_message(interrupt_message):
|
||||
return None
|
||||
return interrupt_message
|
||||
return None
|
||||
|
||||
|
||||
class TestPendingEventNoneChannelPrompt:
|
||||
"""Guard against AttributeError when pending_event is None."""
|
||||
|
||||
@@ -40,3 +54,19 @@ class TestPendingEventNoneChannelPrompt:
|
||||
event = SimpleNamespace()
|
||||
result = _extract_channel_prompt(event)
|
||||
assert result is None
|
||||
|
||||
|
||||
class TestControlInterruptMessages:
|
||||
"""Control interrupt reasons must not become follow-up user input."""
|
||||
|
||||
def test_stop_requested_is_not_treated_as_pending_user_message(self):
|
||||
result = _extract_pending_text(True, None, "Stop requested")
|
||||
assert result is None
|
||||
|
||||
def test_session_reset_requested_is_not_treated_as_pending_user_message(self):
|
||||
result = _extract_pending_text(True, None, "Session reset requested")
|
||||
assert result is None
|
||||
|
||||
def test_real_user_interrupt_message_still_requeues(self):
|
||||
result = _extract_pending_text(True, None, "actually use postgres instead")
|
||||
assert result == "actually use postgres instead"
|
||||
|
||||
@@ -51,6 +51,9 @@ class ProgressCaptureAdapter(BasePlatformAdapter):
|
||||
async def send_typing(self, chat_id, metadata=None) -> None:
|
||||
self.typing.append({"chat_id": chat_id, "metadata": metadata})
|
||||
|
||||
async def stop_typing(self, chat_id) -> None:
|
||||
self.typing.append({"chat_id": chat_id, "metadata": {"stopped": True}})
|
||||
|
||||
async def get_chat_info(self, chat_id: str):
|
||||
return {"id": chat_id}
|
||||
|
||||
@@ -90,6 +93,40 @@ class LongPreviewAgent:
|
||||
}
|
||||
|
||||
|
||||
class DelayedProgressAgent:
|
||||
def __init__(self, **kwargs):
|
||||
self.tool_progress_callback = kwargs.get("tool_progress_callback")
|
||||
self.tools = []
|
||||
|
||||
def run_conversation(self, message, conversation_history=None, task_id=None):
|
||||
self.tool_progress_callback("tool.started", "terminal", "first command", {})
|
||||
time.sleep(0.45)
|
||||
self.tool_progress_callback("tool.started", "terminal", "second command", {})
|
||||
time.sleep(0.1)
|
||||
return {
|
||||
"final_response": "done",
|
||||
"messages": [],
|
||||
"api_calls": 1,
|
||||
}
|
||||
|
||||
|
||||
class DelayedInterimAgent:
|
||||
def __init__(self, **kwargs):
|
||||
self.interim_assistant_callback = kwargs.get("interim_assistant_callback")
|
||||
self.tools = []
|
||||
|
||||
def run_conversation(self, message, conversation_history=None, task_id=None):
|
||||
self.interim_assistant_callback("first interim")
|
||||
time.sleep(0.45)
|
||||
self.interim_assistant_callback("second interim")
|
||||
time.sleep(0.1)
|
||||
return {
|
||||
"final_response": "done",
|
||||
"messages": [],
|
||||
"api_calls": 1,
|
||||
}
|
||||
|
||||
|
||||
def _make_runner(adapter):
|
||||
gateway_run = importlib.import_module("gateway.run")
|
||||
GatewayRunner = gateway_run.GatewayRunner
|
||||
@@ -104,6 +141,7 @@ def _make_runner(adapter):
|
||||
runner._fallback_model = None
|
||||
runner._session_db = None
|
||||
runner._running_agents = {}
|
||||
runner._session_run_generation = {}
|
||||
runner.hooks = SimpleNamespace(loaded_hooks=False)
|
||||
runner.config = SimpleNamespace(
|
||||
thread_sessions_per_user=False,
|
||||
@@ -744,6 +782,154 @@ async def test_base_processing_releases_post_delivery_callback_after_main_send()
|
||||
assert released == [True]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_agent_drops_tool_progress_after_generation_invalidation(monkeypatch, tmp_path):
|
||||
import yaml
|
||||
|
||||
(tmp_path / "config.yaml").write_text(
|
||||
yaml.dump({"display": {"tool_progress": "all"}}),
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
fake_dotenv = types.ModuleType("dotenv")
|
||||
fake_dotenv.load_dotenv = lambda *args, **kwargs: None
|
||||
monkeypatch.setitem(sys.modules, "dotenv", fake_dotenv)
|
||||
|
||||
fake_run_agent = types.ModuleType("run_agent")
|
||||
fake_run_agent.AIAgent = DelayedProgressAgent
|
||||
monkeypatch.setitem(sys.modules, "run_agent", fake_run_agent)
|
||||
import tools.terminal_tool # noqa: F401 - register terminal tool metadata
|
||||
|
||||
adapter = ProgressCaptureAdapter(platform=Platform.DISCORD)
|
||||
runner = _make_runner(adapter)
|
||||
gateway_run = importlib.import_module("gateway.run")
|
||||
monkeypatch.setattr(gateway_run, "_hermes_home", tmp_path)
|
||||
monkeypatch.setattr(gateway_run, "_resolve_runtime_agent_kwargs", lambda: {"api_key": "***"})
|
||||
|
||||
source = SessionSource(
|
||||
platform=Platform.DISCORD,
|
||||
chat_id="dm-1",
|
||||
chat_type="dm",
|
||||
thread_id=None,
|
||||
)
|
||||
session_key = "agent:main:discord:dm:dm-1"
|
||||
runner._session_run_generation[session_key] = 1
|
||||
|
||||
original_send = adapter.send
|
||||
invalidated = {"done": False}
|
||||
|
||||
async def send_and_invalidate(chat_id, content, reply_to=None, metadata=None):
|
||||
result = await original_send(chat_id, content, reply_to=reply_to, metadata=metadata)
|
||||
if "first command" in content and not invalidated["done"]:
|
||||
invalidated["done"] = True
|
||||
runner._invalidate_session_run_generation(session_key, reason="test_stop")
|
||||
return result
|
||||
|
||||
adapter.send = send_and_invalidate
|
||||
|
||||
result = await runner._run_agent(
|
||||
message="hello",
|
||||
context_prompt="",
|
||||
history=[],
|
||||
source=source,
|
||||
session_id="sess-progress-stop",
|
||||
session_key=session_key,
|
||||
run_generation=1,
|
||||
)
|
||||
|
||||
all_progress_text = " ".join(call["content"] for call in adapter.sent)
|
||||
all_progress_text += " ".join(call["content"] for call in adapter.edits)
|
||||
assert result["final_response"] == "done"
|
||||
assert 'first command' in all_progress_text
|
||||
assert 'second command' not in all_progress_text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_agent_drops_interim_commentary_after_generation_invalidation(monkeypatch, tmp_path):
|
||||
import yaml
|
||||
|
||||
(tmp_path / "config.yaml").write_text(
|
||||
yaml.dump({"display": {"tool_progress": "off", "interim_assistant_messages": True}}),
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
fake_dotenv = types.ModuleType("dotenv")
|
||||
fake_dotenv.load_dotenv = lambda *args, **kwargs: None
|
||||
monkeypatch.setitem(sys.modules, "dotenv", fake_dotenv)
|
||||
|
||||
fake_run_agent = types.ModuleType("run_agent")
|
||||
fake_run_agent.AIAgent = DelayedInterimAgent
|
||||
monkeypatch.setitem(sys.modules, "run_agent", fake_run_agent)
|
||||
|
||||
adapter = ProgressCaptureAdapter(platform=Platform.DISCORD)
|
||||
runner = _make_runner(adapter)
|
||||
gateway_run = importlib.import_module("gateway.run")
|
||||
monkeypatch.setattr(gateway_run, "_hermes_home", tmp_path)
|
||||
monkeypatch.setattr(gateway_run, "_resolve_runtime_agent_kwargs", lambda: {"api_key": "***"})
|
||||
|
||||
source = SessionSource(
|
||||
platform=Platform.DISCORD,
|
||||
chat_id="dm-2",
|
||||
chat_type="dm",
|
||||
thread_id=None,
|
||||
)
|
||||
session_key = "agent:main:discord:dm:dm-2"
|
||||
runner._session_run_generation[session_key] = 1
|
||||
|
||||
original_send = adapter.send
|
||||
invalidated = {"done": False}
|
||||
|
||||
async def send_and_invalidate(chat_id, content, reply_to=None, metadata=None):
|
||||
result = await original_send(chat_id, content, reply_to=reply_to, metadata=metadata)
|
||||
if content == "first interim" and not invalidated["done"]:
|
||||
invalidated["done"] = True
|
||||
runner._invalidate_session_run_generation(session_key, reason="test_stop")
|
||||
return result
|
||||
|
||||
adapter.send = send_and_invalidate
|
||||
|
||||
result = await runner._run_agent(
|
||||
message="hello",
|
||||
context_prompt="",
|
||||
history=[],
|
||||
source=source,
|
||||
session_id="sess-commentary-stop",
|
||||
session_key=session_key,
|
||||
run_generation=1,
|
||||
)
|
||||
|
||||
sent_texts = [call["content"] for call in adapter.sent]
|
||||
assert result["final_response"] == "done"
|
||||
assert "first interim" in sent_texts
|
||||
assert "second interim" not in sent_texts
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_keep_typing_stops_immediately_when_interrupt_event_is_set():
|
||||
adapter = ProgressCaptureAdapter(platform=Platform.DISCORD)
|
||||
stop_event = asyncio.Event()
|
||||
|
||||
task = asyncio.create_task(
|
||||
adapter._keep_typing(
|
||||
"dm-typing-stop",
|
||||
interval=30.0,
|
||||
stop_event=stop_event,
|
||||
)
|
||||
)
|
||||
await asyncio.sleep(0.05)
|
||||
stop_event.set()
|
||||
await asyncio.wait_for(task, timeout=0.5)
|
||||
|
||||
normal_typing_calls = [
|
||||
call for call in adapter.typing if call.get("metadata") != {"stopped": True}
|
||||
]
|
||||
stopped_calls = [
|
||||
call for call in adapter.typing if call.get("metadata") == {"stopped": True}
|
||||
]
|
||||
assert len(normal_typing_calls) == 1
|
||||
assert len(stopped_calls) == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_verbose_mode_does_not_truncate_args_by_default(monkeypatch, tmp_path):
|
||||
"""Verbose mode with default tool_preview_length (0) should NOT truncate args.
|
||||
|
||||
@@ -24,10 +24,18 @@ class _FakeAdapter:
|
||||
|
||||
def __init__(self):
|
||||
self._pending_messages = {}
|
||||
self._active_sessions = {}
|
||||
self.interrupted_sessions = []
|
||||
|
||||
async def send(self, chat_id, text, **kwargs):
|
||||
pass
|
||||
|
||||
async def interrupt_session_activity(self, session_key, chat_id):
|
||||
self.interrupted_sessions.append((session_key, chat_id))
|
||||
event = self._active_sessions.get(session_key)
|
||||
if event is not None:
|
||||
event.set()
|
||||
|
||||
|
||||
def _make_runner():
|
||||
runner = object.__new__(GatewayRunner)
|
||||
@@ -37,6 +45,7 @@ def _make_runner():
|
||||
runner.adapters = {Platform.TELEGRAM: _FakeAdapter()}
|
||||
runner._running_agents = {}
|
||||
runner._running_agents_ts = {}
|
||||
runner._session_run_generation = {}
|
||||
runner._pending_messages = {}
|
||||
runner._pending_approvals = {}
|
||||
runner._voice_mode = {}
|
||||
@@ -81,7 +90,7 @@ async def test_sentinel_placed_before_agent_setup():
|
||||
# Patch _handle_message_with_agent to capture state at entry
|
||||
sentinel_was_set = False
|
||||
|
||||
async def mock_inner(self_inner, ev, src, qk):
|
||||
async def mock_inner(self_inner, ev, src, qk, generation):
|
||||
nonlocal sentinel_was_set
|
||||
sentinel_was_set = runner._running_agents.get(qk) is _AGENT_PENDING_SENTINEL
|
||||
return "ok"
|
||||
@@ -105,7 +114,7 @@ async def test_sentinel_cleaned_up_after_handler_returns():
|
||||
event = _make_event()
|
||||
session_key = build_session_key(event.source)
|
||||
|
||||
async def mock_inner(self_inner, ev, src, qk):
|
||||
async def mock_inner(self_inner, ev, src, qk, generation):
|
||||
return "ok"
|
||||
|
||||
with patch.object(GatewayRunner, "_handle_message_with_agent", mock_inner):
|
||||
@@ -127,7 +136,7 @@ async def test_sentinel_cleaned_up_on_exception():
|
||||
event = _make_event()
|
||||
session_key = build_session_key(event.source)
|
||||
|
||||
async def mock_inner(self_inner, ev, src, qk):
|
||||
async def mock_inner(self_inner, ev, src, qk, generation):
|
||||
raise RuntimeError("boom")
|
||||
|
||||
with patch.object(GatewayRunner, "_handle_message_with_agent", mock_inner):
|
||||
@@ -154,7 +163,7 @@ async def test_second_message_during_sentinel_queued_not_duplicate():
|
||||
|
||||
barrier = asyncio.Event()
|
||||
|
||||
async def slow_inner(self_inner, ev, src, qk):
|
||||
async def slow_inner(self_inner, ev, src, qk, generation):
|
||||
# Simulate slow setup — wait until test tells us to proceed
|
||||
await barrier.wait()
|
||||
return "ok"
|
||||
@@ -333,7 +342,7 @@ async def test_stop_during_sentinel_force_cleans_session():
|
||||
|
||||
barrier = asyncio.Event()
|
||||
|
||||
async def slow_inner(self_inner, ev, src, qk):
|
||||
async def slow_inner(self_inner, ev, src, qk, generation):
|
||||
await barrier.wait()
|
||||
return "ok"
|
||||
|
||||
@@ -381,6 +390,7 @@ async def test_stop_hard_kills_running_agent():
|
||||
fake_agent = MagicMock()
|
||||
fake_agent.get_activity_summary.return_value = {"seconds_since_activity": 0}
|
||||
runner._running_agents[session_key] = fake_agent
|
||||
runner.adapters[Platform.TELEGRAM]._active_sessions[session_key] = asyncio.Event()
|
||||
|
||||
# Send /stop
|
||||
stop_event = _make_event(text="/stop")
|
||||
@@ -393,6 +403,10 @@ async def test_stop_hard_kills_running_agent():
|
||||
assert session_key not in runner._running_agents, (
|
||||
"/stop must remove the agent from _running_agents so the session is unlocked"
|
||||
)
|
||||
assert runner.adapters[Platform.TELEGRAM].interrupted_sessions == [
|
||||
(session_key, "12345")
|
||||
]
|
||||
assert runner.adapters[Platform.TELEGRAM]._active_sessions[session_key].is_set()
|
||||
|
||||
# Must return a confirmation
|
||||
assert result is not None
|
||||
|
||||
@@ -50,6 +50,7 @@ def _make_runner(session_entry: SessionEntry):
|
||||
runner.session_store.rewrite_transcript = MagicMock()
|
||||
runner.session_store.update_session = MagicMock()
|
||||
runner._running_agents = {}
|
||||
runner._session_run_generation = {}
|
||||
runner._pending_messages = {}
|
||||
runner._pending_approvals = {}
|
||||
runner._session_db = MagicMock()
|
||||
@@ -223,6 +224,52 @@ async def test_handle_message_persists_agent_token_counts(monkeypatch):
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handle_message_discards_stale_result_after_session_invalidation(monkeypatch):
|
||||
import gateway.run as gateway_run
|
||||
|
||||
session_entry = SessionEntry(
|
||||
session_key=build_session_key(_make_source()),
|
||||
session_id="sess-1",
|
||||
created_at=datetime.now(),
|
||||
updated_at=datetime.now(),
|
||||
platform=Platform.TELEGRAM,
|
||||
chat_type="dm",
|
||||
)
|
||||
runner = _make_runner(session_entry)
|
||||
runner.session_store.load_transcript.return_value = [{"role": "user", "content": "earlier"}]
|
||||
session_key = session_entry.session_key
|
||||
runner.adapters[Platform.TELEGRAM]._post_delivery_callbacks = {session_key: object()}
|
||||
|
||||
async def _stale_result(**kwargs):
|
||||
runner._invalidate_session_run_generation(kwargs["session_key"], reason="test_stale_result")
|
||||
return {
|
||||
"final_response": "late reply",
|
||||
"messages": [],
|
||||
"tools": [],
|
||||
"history_offset": 0,
|
||||
"last_prompt_tokens": 80,
|
||||
"input_tokens": 120,
|
||||
"output_tokens": 45,
|
||||
"model": "openai/test-model",
|
||||
}
|
||||
|
||||
runner._run_agent = AsyncMock(side_effect=_stale_result)
|
||||
|
||||
monkeypatch.setattr(gateway_run, "_resolve_runtime_agent_kwargs", lambda: {"api_key": "***"})
|
||||
monkeypatch.setattr(
|
||||
"agent.model_metadata.get_model_context_length",
|
||||
lambda *_args, **_kwargs: 100000,
|
||||
)
|
||||
|
||||
result = await runner._handle_message(_make_event("hello"))
|
||||
|
||||
assert result is None
|
||||
runner.session_store.append_to_transcript.assert_not_called()
|
||||
runner.session_store.update_session.assert_not_called()
|
||||
assert session_key not in runner.adapters[Platform.TELEGRAM]._post_delivery_callbacks
|
||||
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_status_command_bypasses_active_session_guard():
|
||||
|
||||
Reference in New Issue
Block a user