fix(gateway): persist the platform message id on every user turn

This commit is contained in:
Kyzcreig
2026-08-30 06:16:33 -07:00
committed by Teknium
parent 5505042f40
commit c1bd0511cf
7 changed files with 270 additions and 24 deletions
+6
View File
@@ -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,
+9
View File
@@ -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
+7
View File
@@ -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)
+19 -1
View File
@@ -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:
@@ -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
)
+22 -22
View File
@@ -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"})
+1 -1
View File
@@ -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 {