From c1bd0511cfc2c310274fd95b515addbb06a2ec99 Mon Sep 17 00:00:00 2001 From: Kyzcreig <9063726+Kyzcreig@users.noreply.github.com> Date: Sun, 30 Aug 2026 06:16:33 -0700 Subject: [PATCH] fix(gateway): persist the platform message id on every user turn --- agent/conversation_loop.py | 6 + agent/turn_context.py | 9 + gateway/run.py | 7 + run_agent.py | 20 +- .../test_restart_drain_recovery_dedup.py | 206 ++++++++++++++++++ tests/gateway/test_run_progress_topics.py | 44 ++-- tests/gateway/test_tts_media_routing.py | 2 +- 7 files changed, 270 insertions(+), 24 deletions(-) create mode 100644 tests/gateway/test_restart_drain_recovery_dedup.py diff --git a/agent/conversation_loop.py b/agent/conversation_loop.py index f38c406ee7..4158e97304 100644 --- a/agent/conversation_loop.py +++ b/agent/conversation_loop.py @@ -1918,6 +1918,7 @@ def run_conversation( persist_user_timestamp: Optional[float] = None, persist_user_display_kind: Optional[str] = None, persist_user_display_metadata: Optional[Dict[str, Any]] = None, + persist_user_platform_id: Optional[str] = None, moa_config: Optional[dict[str, Any]] = None, ) -> Dict[str, Any]: """ @@ -1943,6 +1944,10 @@ def run_conversation( the message unchanged. persist_user_display_metadata: Optional payload for that event (e.g. a delegation's task count). + persist_user_platform_id: Optional platform-side message id (e.g. the + Discord/Telegram message id) to store as metadata on that + persisted user message, so restart drain-window recovery can + dedup an interrupted turn against the transcript. or queuing follow-up prefetch work. Returns: @@ -1995,6 +2000,7 @@ def run_conversation( persist_user_timestamp, persist_user_display_kind=persist_user_display_kind, persist_user_display_metadata=persist_user_display_metadata, + persist_user_platform_id=persist_user_platform_id, restore_or_build_system_prompt=_restore_or_build_system_prompt, install_safe_stdio=_install_safe_stdio, sanitize_surrogates=_sanitize_surrogates, diff --git a/agent/turn_context.py b/agent/turn_context.py index d6882e9129..8f333870cf 100644 --- a/agent/turn_context.py +++ b/agent/turn_context.py @@ -557,6 +557,7 @@ def build_turn_context( stream_callback, persist_user_message: Optional[Any], persist_user_timestamp: Optional[float] = None, + persist_user_platform_id: Optional[str] = None, *, persist_user_display_kind: Optional[str] = None, persist_user_display_metadata: Optional[Dict[str, Any]] = None, @@ -670,6 +671,7 @@ def build_turn_context( agent._persist_user_message_idx = None agent._persist_user_message_override = persist_user_message agent._persist_user_message_timestamp = persist_user_timestamp + agent._persist_user_message_platform_id = persist_user_platform_id # Generate unique task_id if not provided to isolate VMs between tasks. effective_task_id = task_id or str(uuid.uuid4()) agent._current_task_id = effective_task_id @@ -807,6 +809,13 @@ def build_turn_context( if persist_user_display_metadata: user_msg["display_metadata"] = persist_user_display_metadata + # Stamp the platform-side message id (e.g. the Discord/Telegram message id) + # as metadata on the user turn so it survives the early crash-resilience + # persist below (the turn-start flush). Load-bearing for restart + # drain-window recovery: a recovery pass dedups via + # ``has_platform_message_id`` against this row. + if persist_user_platform_id is not None: + user_msg["platform_message_id"] = persist_user_platform_id append_message(messages, user_msg) current_turn_user_idx = len(messages) - 1 agent._persist_user_message_idx = current_turn_user_idx diff --git a/gateway/run.py b/gateway/run.py index 410aa14a14..74be532104 100644 --- a/gateway/run.py +++ b/gateway/run.py @@ -6963,6 +6963,13 @@ class TurnRunner: _conversation_kwargs["moa_config"] = ctx.moa_config if _persist_user_timestamp_override is not None: _conversation_kwargs["persist_user_timestamp"] = _persist_user_timestamp_override + # Thread the platform-side inbound message id onto the persisted + # user turn so a turn interrupted by a gateway restart is durably + # recorded WITH its id — restart drain-window recovery dedups + # against has_platform_message_id, and without this the + # interrupted turn is invisible to that check. + if ctx.event_message_id is not None: + _conversation_kwargs["persist_user_platform_id"] = str(ctx.event_message_id) result = agent.run_conversation(_api_run_message, **_conversation_kwargs) finally: unregister_gateway_notify(_approval_session_key) diff --git a/run_agent.py b/run_agent.py index 4d5ea6f5ab..aaca141527 100644 --- a/run_agent.py +++ b/run_agent.py @@ -2073,7 +2073,10 @@ class AIAgent: idx = getattr(self, "_persist_user_message_idx", None) override = getattr(self, "_persist_user_message_override", None) timestamp = getattr(self, "_persist_user_message_timestamp", None) - if idx is None or (override is None and timestamp is None): + platform_id = getattr(self, "_persist_user_message_platform_id", None) + if idx is None or ( + override is None and timestamp is None and platform_id is None + ): return if 0 <= idx < len(messages): msg = messages[idx] @@ -2105,6 +2108,14 @@ class AIAgent: msg["content"] = override if timestamp is not None: msg["timestamp"] = timestamp + # Platform-side message id (e.g. the Discord/Telegram message + # id) — metadata, load-bearing for restart drain-window + # recovery dedup: it lets a recovery pass ask + # ``has_platform_message_id`` whether an interrupted turn + # already reached the transcript. Stamped here in addition to + # ``build_turn_context`` so it survives the override path. + if platform_id is not None: + msg["platform_message_id"] = platform_id def _persist_session(self, messages: List[Dict], conversation_history: List[Dict] = None): """Save session state to both JSON log and SQLite on any exit path. @@ -2485,6 +2496,11 @@ class AIAgent: else msg.get("display_kind") ), "display_metadata": msg.get("display_metadata"), + # Platform-side message id (e.g. the Discord/Telegram + # message id). _insert_message_rows reads it off the row + # dict; load-bearing for restart drain-window recovery + # dedup via has_platform_message_id. + "platform_message_id": msg.get("platform_message_id"), } if isinstance(msg.get("_row_id"), int): _row["_row_id"] = msg["_row_id"] @@ -8769,6 +8785,7 @@ class AIAgent: persist_user_timestamp: Optional[float] = None, persist_user_display_kind: Optional[str] = None, persist_user_display_metadata: Optional[Dict[str, Any]] = None, + persist_user_platform_id: Optional[str] = None, moa_config: Optional[dict[str, Any]] = None, ) -> Dict[str, Any]: """Forwarder — see ``agent.conversation_loop.run_conversation``.""" @@ -9143,6 +9160,7 @@ class AIAgent: persist_user_timestamp=persist_user_timestamp, persist_user_display_kind=persist_user_display_kind, persist_user_display_metadata=persist_user_display_metadata, + persist_user_platform_id=persist_user_platform_id, moa_config=moa_config, ) finally: diff --git a/tests/gateway/test_restart_drain_recovery_dedup.py b/tests/gateway/test_restart_drain_recovery_dedup.py new file mode 100644 index 0000000000..15c2b54e1e --- /dev/null +++ b/tests/gateway/test_restart_drain_recovery_dedup.py @@ -0,0 +1,206 @@ +"""Restart drain-window recovery must be able to dedup an interrupted turn. + +The Discord missed-message backfill (``_run_missed_message_backfill``) exists +to recover messages the bot never saw while it was down. A gateway RESTART +produces a harder case: the message WAS received and a turn WAS started, then +the drain window force-interrupted it. The transcript is the only durable +record of that, and the transcript row for the user turn is written WITHOUT +the platform-side message id — so nothing downstream can ask "did this +Discord message already reach the transcript?" and the recovery pass has no +authority to dedup against. + +``SessionDB`` already carries a ``platform_message_id`` column, a partial +unique index over ``(session_id, platform_message_id)``, and a +``has_platform_message_id`` lookup — the storage and the query exist. What is +missing is the WRITE on the normal agent-persisted turn path: the id is only +attached on the gateway-side transient-failure fallback +(``_handle_message_with_agent``), never on the path the agent itself flushes. +""" + +from __future__ import annotations + +import types + +import pytest + +from hermes_state import SessionDB + + +def _make_db(tmp_path) -> SessionDB: + return SessionDB(db_path=tmp_path / "state.db") + + +class _MinimalAgent: + """The narrow slice of AIAgent that ``_apply_persist_user_message_override`` + and ``_flush_messages_to_session_db`` read.""" + + def __init__(self, db: SessionDB, session_id: str): + self._session_db = db + self._session_db_created = True + self.session_id = session_id + self._last_flushed_db_idx = 0 + self._flushed_db_message_ids = set() + self._flushed_db_message_session_id = session_id + self._persist_user_message_idx = None + self._persist_user_message_override = None + self._persist_user_message_timestamp = None + self._persist_disabled = False + + def _ensure_db_session(self): # pragma: no cover - already created + return None + + +def test_build_turn_context_stamps_the_platform_message_id_on_the_user_turn(): + """The turn prologue must carry the platform id onto the user turn dict. + + This is the row the early crash-resilience persist writes, so it is the + only place a drain-interrupted turn can pick the id up. + """ + from agent.turn_context import build_turn_context + + agent = types.SimpleNamespace() + ctx = _build_turn_context_for_test( + build_turn_context, agent, persist_user_platform_id="discord-991" + ) + + user_msgs = [m for m in ctx.messages if m.get("role") == "user"] + assert user_msgs, "no user turn in the built context" + assert user_msgs[-1].get("platform_message_id") == "discord-991", ( + "the user turn reached persistence without its platform message id — a " + "drain-interrupted turn is then unrecoverable/undedupable by " + "has_platform_message_id" + ) + + +def test_persisted_interrupted_turn_is_findable_by_platform_message_id(tmp_path): + """E2E: flush a turn the way the agent does, then ask the dedup authority. + + This is the exact question the restart drain-window recovery pass asks + before re-dispatching a message. On main the answer is False even though + the turn IS in the transcript, so recovery would re-run a turn that + already ran (duplicate work, duplicate spend, duplicate reply). + """ + from run_agent import AIAgent + + db = _make_db(tmp_path) + session_id = db.create_session("sess-drain-window", "gateway") + + agent = _MinimalAgent(db, session_id) + agent._persist_user_message_idx = 0 + agent._persist_user_message_platform_id = "discord-4242" + + messages = [{"role": "user", "content": "please do the thing"}] + + AIAgent._apply_persist_user_message_override(agent, messages) + AIAgent._flush_messages_to_session_db_unlocked( + agent, messages, conversation_history=None + ) + + assert db.has_platform_message_id(session_id, "discord-4242"), ( + "the interrupted turn is in the transcript but carries no " + "platform_message_id, so restart drain-window recovery cannot tell it " + "already ran and will re-dispatch it" + ) + + +def test_platform_message_id_survives_a_persist_content_override(tmp_path): + """The id must not be lost on the override path. + + Group-chat / observed-context turns route through + ``_persist_user_message_override``; the id has to survive that rewrite or + the dedup authority is blind for exactly the busy channels that need it. + """ + from run_agent import AIAgent + + db = _make_db(tmp_path) + session_id = db.create_session("sess-override", "gateway") + + agent = _MinimalAgent(db, session_id) + agent._persist_user_message_idx = 0 + agent._persist_user_message_override = "clean transcript text" + agent._persist_user_message_platform_id = "discord-7777" + + messages = [{"role": "user", "content": "api-facing text with context"}] + + AIAgent._apply_persist_user_message_override(agent, messages) + AIAgent._flush_messages_to_session_db_unlocked( + agent, messages, conversation_history=None + ) + + assert db.has_platform_message_id(session_id, "discord-7777") + + +def _build_turn_context_for_test(build_turn_context, agent, **overrides): + """Construct a minimal build_turn_context call. + + Mirrors ``tests/agent/test_turn_context.py::_build`` but is kept local so + this file stays self-contained. + """ + from tests.agent.test_turn_context import _FakeAgent, _stub_runtime_main + + fake = _FakeAgent() + kwargs = dict( + agent=fake, + user_message="hello", + system_message=None, + conversation_history=None, + task_id=None, + stream_callback=None, + persist_user_message=None, + restore_or_build_system_prompt=lambda *a, **k: None, + install_safe_stdio=lambda: None, + sanitize_surrogates=lambda s: s, + summarize_user_message_for_log=lambda s: s, + set_session_context=lambda _sid: None, + set_current_write_origin=lambda _o: None, + ra=lambda: types.SimpleNamespace(_set_interrupt=lambda *a, **k: None), + ) + kwargs.update(overrides) + return build_turn_context(**kwargs) + + +def test_gateway_run_agent_threads_the_event_message_id_into_the_turn(): + """AST proof that the gateway call site passes the id down. + + The unit tests above prove the persistence layer STORES the id once it is + given one. This pins the wiring: without the gateway forwarding + ``event_message_id`` as ``persist_user_platform_id``, the whole path is + dead code and every real inbound turn still persists without its id. + """ + import ast + import inspect + + import gateway.run as gateway_run + + source = inspect.getsource(gateway_run) + tree = ast.parse(source) + + forwards = [ + node + for node in ast.walk(tree) + if isinstance(node, ast.Subscript) + and isinstance(node.slice, ast.Constant) + and node.slice.value == "persist_user_platform_id" + ] + assert forwards, ( + "gateway/run.py never forwards persist_user_platform_id — the inbound " + "platform message id never reaches the persisted user turn, so a " + "drain-interrupted turn stays undedupable" + ) + + +def test_run_conversation_accepts_persist_user_platform_id(): + """The public forwarder must expose the kwarg the gateway passes.""" + import inspect + + from agent.conversation_loop import run_conversation + from run_agent import AIAgent + + assert ( + "persist_user_platform_id" + in inspect.signature(run_conversation).parameters + ) + assert ( + "persist_user_platform_id" + in inspect.signature(AIAgent.run_conversation).parameters + ) diff --git a/tests/gateway/test_run_progress_topics.py b/tests/gateway/test_run_progress_topics.py index c955605070..583c48d52d 100644 --- a/tests/gateway/test_run_progress_topics.py +++ b/tests/gateway/test_run_progress_topics.py @@ -215,7 +215,7 @@ class FakeAgent: self.tool_progress_callback = kwargs.get("tool_progress_callback") self.tools = [] - def run_conversation(self, message, conversation_history=None, task_id=None): + def run_conversation(self, message, conversation_history=None, task_id=None, **kwargs): cb = self.tool_progress_callback if cb is not None: cb("tool.started", "terminal", "pwd", {}) @@ -290,7 +290,7 @@ class DuplicateNativeToolsAgent: self.tool_complete_callback = kwargs.get("tool_complete_callback") self.tools = [] - def run_conversation(self, message, conversation_history=None, task_id=None): + def run_conversation(self, message, conversation_history=None, task_id=None, **kwargs): self.tool_start_callback("call-a", "web_search", {"query": "alpha"}) time.sleep(0.15) self.tool_start_callback("call-b", "web_search", {"query": "beta"}) @@ -319,7 +319,7 @@ class ThinkingAgent: self.tool_progress_callback = kwargs.get("tool_progress_callback") self.tools = [] - def run_conversation(self, message, conversation_history=None, task_id=None): + def run_conversation(self, message, conversation_history=None, task_id=None, **kwargs): cb = self.tool_progress_callback if cb is not None: cb("_thinking", "weighing the options here") @@ -339,7 +339,7 @@ class LongPreviewAgent: self.tool_progress_callback = kwargs.get("tool_progress_callback") self.tools = [] - def run_conversation(self, message, conversation_history=None, task_id=None): + def run_conversation(self, message, conversation_history=None, task_id=None, **kwargs): self.tool_progress_callback("tool.started", "terminal", self.LONG_CMD, {}) time.sleep(0.35) return { @@ -356,7 +356,7 @@ class UrlPreviewAgent: self.tool_progress_callback = kwargs.get("tool_progress_callback") self.tools = [] - def run_conversation(self, message, conversation_history=None, task_id=None): + def run_conversation(self, message, conversation_history=None, task_id=None, **kwargs): self.tool_progress_callback( "tool.started", "web_extract", @@ -376,7 +376,7 @@ class DelayedProgressAgent: self.tool_progress_callback = kwargs.get("tool_progress_callback") self.tools = [] - def run_conversation(self, message, conversation_history=None, task_id=None): + def run_conversation(self, message, conversation_history=None, task_id=None, **kwargs): self.tool_progress_callback("tool.started", "terminal", "first command", {}) time.sleep(0.45) self.tool_progress_callback("tool.started", "terminal", "second command", {}) @@ -395,7 +395,7 @@ class RetryableEditProgressAgent: self.tool_progress_callback = kwargs.get("tool_progress_callback") self.tools = [] - def run_conversation(self, message, conversation_history=None, task_id=None): + def run_conversation(self, message, conversation_history=None, task_id=None, **kwargs): callback = self.tool_progress_callback assert callback is not None callback("tool.started", "terminal", "first command", {}) @@ -420,7 +420,7 @@ class ManyProgressLinesAgent: self.tool_progress_callback = kwargs.get("tool_progress_callback") self.tools = [] - def run_conversation(self, message, conversation_history=None, task_id=None): + def run_conversation(self, message, conversation_history=None, task_id=None, **kwargs): cb = self.tool_progress_callback assert cb is not None cb("tool.started", "terminal", "first-short", {}) @@ -443,7 +443,7 @@ class DelayedInterimAgent: self.interim_assistant_callback = kwargs.get("interim_assistant_callback") self.tools = [] - def run_conversation(self, message, conversation_history=None, task_id=None): + def run_conversation(self, message, conversation_history=None, task_id=None, **kwargs): self.interim_assistant_callback("first interim") time.sleep(0.45) self.interim_assistant_callback("second interim") @@ -802,7 +802,7 @@ class CommentaryAgent: self.stream_delta_callback = kwargs.get("stream_delta_callback") self.tools = [] - def run_conversation(self, message, conversation_history=None, task_id=None): + def run_conversation(self, message, conversation_history=None, task_id=None, **kwargs): if self.interim_assistant_callback: self.interim_assistant_callback("I'll inspect the repo first.", already_streamed=False) time.sleep(0.1) @@ -820,7 +820,7 @@ class PreviewedResponseAgent: self.interim_assistant_callback = kwargs.get("interim_assistant_callback") self.tools = [] - def run_conversation(self, message, conversation_history=None, task_id=None): + def run_conversation(self, message, conversation_history=None, task_id=None, **kwargs): if self.interim_assistant_callback: self.interim_assistant_callback("You're welcome.", already_streamed=False) return { @@ -837,7 +837,7 @@ class PreviewedSplitAfterCommentaryAgent: self.session_id = kwargs.get("session_id") self.tools = [] - def run_conversation(self, message, conversation_history=None, task_id=None): + def run_conversation(self, message, conversation_history=None, task_id=None, **kwargs): if self.interim_assistant_callback: self.interim_assistant_callback("I'll inspect the repo first.", already_streamed=False) self.session_id = f"{self.session_id}-child" @@ -854,7 +854,7 @@ class StreamingRefineAgent: self.stream_delta_callback = kwargs.get("stream_delta_callback") self.tools = [] - def run_conversation(self, message, conversation_history=None, task_id=None): + def run_conversation(self, message, conversation_history=None, task_id=None, **kwargs): if self.stream_delta_callback: self.stream_delta_callback("Continuing to refine:") time.sleep(0.1) @@ -875,7 +875,7 @@ class QueuedCommentaryAgent: self.interim_assistant_callback = kwargs.get("interim_assistant_callback") self.tools = [] - def run_conversation(self, message, conversation_history=None, task_id=None): + def run_conversation(self, message, conversation_history=None, task_id=None, **kwargs): type(self).calls += 1 if type(self).calls == 1 and self.interim_assistant_callback: self.interim_assistant_callback("I'll inspect the repo first.", already_streamed=False) @@ -896,7 +896,7 @@ class QueuedMediaAgent: self.stream_delta_callback = kwargs.get("stream_delta_callback") self.tools = [] - def run_conversation(self, message, conversation_history=None, task_id=None): + def run_conversation(self, message, conversation_history=None, task_id=None, **kwargs): type(self).calls += 1 if type(self).calls == 1: final_response = f"first response\nMEDIA:{type(self).media_path}" @@ -921,7 +921,7 @@ class QueuedSilenceAgent: def __init__(self, **kwargs): self.tools = [] - def run_conversation(self, message, conversation_history=None, task_id=None): + def run_conversation(self, message, conversation_history=None, task_id=None, **kwargs): type(self).calls += 1 return { "final_response": "NO_REPLY" if type(self).calls == 1 else "follow-up processed", @@ -938,7 +938,7 @@ class QueuedFailedEmptyAgent: def __init__(self, **kwargs): self.tools = [] - def run_conversation(self, message, conversation_history=None, task_id=None): + def run_conversation(self, message, conversation_history=None, task_id=None, **kwargs): type(self).calls += 1 if type(self).calls == 1: return { @@ -960,7 +960,7 @@ class BackgroundReviewAgent: self.background_review_callback = kwargs.get("background_review_callback") self.tools = [] - def run_conversation(self, message, conversation_history=None, task_id=None): + def run_conversation(self, message, conversation_history=None, task_id=None, **kwargs): if self.background_review_callback: self.background_review_callback("💾 Skill 'prospect-scanner' created.") return { @@ -978,7 +978,7 @@ class VerboseAgent: self.tool_progress_callback = kwargs.get("tool_progress_callback") self.tools = [] - def run_conversation(self, message, conversation_history=None, task_id=None): + def run_conversation(self, message, conversation_history=None, task_id=None, **kwargs): self.tool_progress_callback( "tool.started", "execute_code", None, {"code": self.LONG_CODE}, @@ -1194,7 +1194,7 @@ class TransformedStreamAgent: self.stream_delta_callback = kwargs.get("stream_delta_callback") self.tools = [] - def run_conversation(self, message, conversation_history=None, task_id=None): + def run_conversation(self, message, conversation_history=None, task_id=None, **kwargs): if self.stream_delta_callback: self.stream_delta_callback("original answer") return { @@ -1686,7 +1686,7 @@ class TerminalCommandAgent: self.tool_progress_callback = kwargs.get("tool_progress_callback") self.tools = [] - def run_conversation(self, message, conversation_history=None, task_id=None): + def run_conversation(self, message, conversation_history=None, task_id=None, **kwargs): self.tool_progress_callback( "tool.started", "terminal", self.CMD, {"command": self.CMD} ) @@ -1849,7 +1849,7 @@ class MultiTerminalCommandAgent: self.tool_progress_callback = kwargs.get("tool_progress_callback") self.tools = [] - def run_conversation(self, message, conversation_history=None, task_id=None): + def run_conversation(self, message, conversation_history=None, task_id=None, **kwargs): cb = self.tool_progress_callback cb("tool.started", "terminal", "echo one", {"command": "echo one"}) cb("tool.started", "terminal", "echo two", {"command": "echo two"}) diff --git a/tests/gateway/test_tts_media_routing.py b/tests/gateway/test_tts_media_routing.py index 006d654a91..fd3ed9ee44 100644 --- a/tests/gateway/test_tts_media_routing.py +++ b/tests/gateway/test_tts_media_routing.py @@ -493,7 +493,7 @@ class _QueuedMediaAgent: def __init__(self, **kwargs): self.tools = [] - def run_conversation(self, message, conversation_history=None, task_id=None): + def run_conversation(self, message, conversation_history=None, task_id=None, **kwargs): type(self).calls += 1 if type(self).calls == 1: return {