fix(agent): freeze turn confirmation expiry and verify wire parity

This commit is contained in:
joaomarcos
2026-09-08 00:03:40 -03:00
committed by kshitij
parent e5ca5207de
commit 401fef6e6f
3 changed files with 357 additions and 23 deletions

View File

@@ -151,6 +151,10 @@ def canonicalize_replay_history(
return strip_stale_dangerous_confirmations(cleaned, now=now)
# Backward-compatible alias for the send-path name (2026-09-07 code).
canonicalize_history_for_send = canonicalize_replay_history
# --- Stale dangerous-confirmation text expiry ---
# Short on purpose: a dangerous confirmation must not survive any restart or resume gap.
@@ -199,12 +203,21 @@ def strip_stale_dangerous_confirmations(
cleaned: List[Dict[str, Any]] = []
for msg in agent_history:
ts = msg.get("timestamp") if isinstance(msg, dict) and msg.get("role") == "user" else None
if ts is None or not is_dangerous_confirmation(msg.get("content", "")) or (now - float(ts)) <= expiry_seconds:
try:
is_stale = (
ts is not None
and is_dangerous_confirmation(msg.get("content", ""))
and (float(now) - float(ts)) > expiry_seconds
)
except (ValueError, TypeError):
is_stale = False
if not is_stale:
cleaned.append(msg)
continue
logger.debug(
"Redacting stale dangerous-confirmation text in user message (age=%.1fs, expiry=%.1fs): %r",
now - float(ts), expiry_seconds, (msg.get("content") or "")[:80],
float(now) - float(ts), expiry_seconds, (msg.get("content") or "")[:80],
)
redacted = dict(msg)
redacted["content"] = _EXPIRED_CONFIRMATION_SENTINEL

View File

@@ -466,6 +466,7 @@ def _bind_turn_identity(
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
agent._current_turn_timestamp = persist_user_timestamp
# Unique task_id when not provided isolates VMs between tasks.
effective_task_id = task_id or str(uuid.uuid4())
agent._current_task_id = effective_task_id
@@ -494,7 +495,7 @@ _PER_TURN_RESET_STATE: Tuple[Tuple[str, Any], ...] = (
("_tool_guardrail_halt_decision", None), ("_vision_supported", True),
("_iteration_budget_warning_injected", False),
("_run_budget_wrapup_injected", False), ("_verification_stop_nudges", 0),
("_pre_verify_nudges", 0),
("_pre_verify_nudges", 0), ("_current_turn_timestamp", None),
)
@@ -504,8 +505,11 @@ def _reset_per_turn_agent_state(agent: Any) -> None:
setattr(agent, name, value)
agent._turn_failed_file_mutations = {}
agent._turn_file_mutation_paths = set()
agent._tool_guardrails.reset_for_turn()
_reset_consol = getattr(agent._memory_store, "reset_consolidation_failures", None)
_guardrails = getattr(agent, "_tool_guardrails", None)
if _guardrails is not None and hasattr(_guardrails, "reset_for_turn"):
_guardrails.reset_for_turn()
_mem = getattr(agent, "_memory_store", None)
_reset_consol = getattr(_mem, "reset_consolidation_failures", None) if _mem is not None else None
if callable(_reset_consol):
_reset_consol()
@@ -563,6 +567,8 @@ def _stage_turn_user_message(
# CLI input is stamped when staged; gateway input may carry the platform event
# time. Preserve either value and cover any legacy unstamped handoff.
stamp_message_timestamp(user_msg, timestamp=persist_user_timestamp)
if agent is not None and getattr(agent, "_current_turn_timestamp", None) is None:
agent._current_turn_timestamp = user_msg.get("timestamp")
# Synthesized turns stamp their transcript type so the crash persist writes a typed
# row; the model still receives role/content unchanged (api_messages strips both).
@@ -1034,6 +1040,7 @@ def _sanitize_model_for(agent: Any, moa_config: Any) -> Any:
def build_api_messages(
agent: Any, messages: List[Dict[str, Any]], *, current_turn_user_idx: Any,
ext_prefetch_cache: Any, plugin_user_context: Any, moa_config: Any, active_system_prompt: Any,
now: Optional[float] = None,
) -> Tuple[List[Dict[str, Any]], str]:
"""Build the wire copy of ``messages`` for one API call plus the effective system
message. Returns ``(api_messages, effective_system)``.
@@ -1055,10 +1062,33 @@ def build_api_messages(
and 0 <= current_turn_user_idx < len(messages)
else None
)
turn_now = now
if turn_now is None and agent is not None:
_agent_ts = getattr(agent, "_current_turn_timestamp", None)
if isinstance(_agent_ts, (int, float)):
turn_now = float(_agent_ts)
if turn_now is None and isinstance(current_turn_message, dict):
_msg_ts = current_turn_message.get("timestamp")
if isinstance(_msg_ts, (int, float)):
turn_now = float(_msg_ts)
elif isinstance(_msg_ts, str):
try:
turn_now = float(_msg_ts)
except ValueError:
pass
if turn_now is None:
turn_now = time.time()
if agent is not None:
with suppress(Exception):
agent._current_turn_timestamp = turn_now
# Replay consumers rewrite interrupted blocks, dangling tails, and expired
# confirmations on read. Apply the exact same transform to this request-only
# copy before sidecars are substituted; the durable transcript remains intact.
canonical_messages = canonicalize_replay_history(messages)
# The expiry evaluation is frozen for the active turn so tool-loop iterations
# cannot rewrite the prefix or withdraw confirmation mid-turn.
canonical_messages = canonicalize_replay_history(messages, now=turn_now)
api_messages = []
for idx, msg in enumerate(canonical_messages):

View File

@@ -7,8 +7,13 @@ because the dangling tool-call tail was replayed on every resume).
"""
import copy
import json
from pathlib import Path
import tempfile
import time
from agent.replay_cleanup import (
canonicalize_history_for_send,
canonicalize_replay_history,
is_interrupted_tool_result,
strip_dangling_tool_call_tail,
@@ -16,6 +21,36 @@ from agent.replay_cleanup import (
strip_stale_dangerous_confirmations,
sanitize_replay_history,
)
from agent.transports.chat_completions import ChatCompletionsTransport
from agent.turn_context import build_api_messages
from hermes_state import SessionDB
def _wire(messages):
return ChatCompletionsTransport().convert_messages(list(messages))
def _canon(objs):
return json.dumps(objs, sort_keys=True, separators=(",", ":"))
class _Agent:
api_mode = "chat_completions"
ephemeral_system_prompt = None
_compression_warning = None
max_iterations = 10
@staticmethod
def _copy_reasoning_content_for_api(_source, _target):
return None
@staticmethod
def _should_sanitize_tool_calls():
return False
@staticmethod
def _sanitize_tool_calls_for_strict_api(*_args, **_kwargs):
return None
def _user(text):
@@ -117,27 +152,17 @@ def test_canonicalize_replay_history_matches_all_resume_transforms():
assert actual == expected
def test_canonicalize_history_for_send_alias():
assert canonicalize_history_for_send is canonicalize_replay_history
def test_send_builder_uses_canonical_history_without_mutating_source():
"""The request copy must match replay cleanup while durable history stays intact."""
from agent.turn_context import build_api_messages
class _Agent:
api_mode = "chat_completions"
ephemeral_system_prompt = None
@staticmethod
def _copy_reasoning_content_for_api(_source, _target):
return None
@staticmethod
def _should_sanitize_tool_calls():
return False
@staticmethod
def _sanitize_tool_calls_for_strict_api(*_args, **_kwargs):
return None
now = 10_000.0
agent = _Agent()
agent._current_turn_timestamp = now
history = [
_user("before"),
{"role": "assistant", "content": "ack"},
@@ -149,7 +174,7 @@ def test_send_builder_uses_canonical_history_without_mutating_source():
original = copy.deepcopy(history)
request, _ = build_api_messages(
_Agent(), history, current_turn_user_idx=len(history) - 1, ext_prefetch_cache="",
agent, history, current_turn_user_idx=len(history) - 1, ext_prefetch_cache="",
plugin_user_context="", moa_config=None, active_system_prompt="",
)
@@ -159,3 +184,269 @@ def test_send_builder_uses_canonical_history_without_mutating_source():
]
assert "EXPIRED" in request[2]["content"]
assert request[-1]["content"] == "current-wire"
# Wire representation comparison: send wire matches replay wire with sidecar applied
wire_request = _wire(request)
expected_replay = canonicalize_replay_history(copy.deepcopy(history[:-1]), now=now) + [
{"role": "user", "content": "current-wire"}
]
assert _canon(wire_request) == _canon(_wire(expected_replay))
def test_send_byte_identity_with_tui_replay_interrupted_block():
"""Interrupted read-only assistant->tool block: send path wire bytes match TUI replay."""
now = 10_000.0
agent = _Agent()
agent._current_turn_timestamp = now
history = [
_user("u1"),
{"role": "assistant", "content": "a1"},
_assistant_tc("read_file"), _tool("[command interrupted]"),
_user("u2"),
]
send_request, _ = build_api_messages(
agent, history, current_turn_user_idx=len(history) - 1, ext_prefetch_cache="",
plugin_user_context="", moa_config=None, active_system_prompt="",
)
wire_send = _wire(send_request)
wire_replay = _wire(sanitize_replay_history(copy.deepcopy(history)))
assert _canon(wire_send) == _canon(wire_replay)
def test_send_byte_identity_with_dangling_tool_call_tail():
"""Trailing unanswered assistant(tool_calls): send path wire bytes match TUI replay."""
now = 10_000.0
history = [
_user("u1"),
{"role": "assistant", "content": "a1"},
{"role": "assistant", "content": "", "tool_calls": [{"id": "c9", "type": "function", "function": {"name": "read_file", "arguments": "{}"}}]},
]
wire_replay = _wire(sanitize_replay_history(copy.deepcopy(history)))
wire_canon = _wire(canonicalize_replay_history(copy.deepcopy(history), now=now))
assert _canon(wire_canon) == _canon(wire_replay)
def test_send_byte_identity_with_stale_dangerous_confirmation():
"""Stale confirmation (>60s): send path wire bytes match gateway replay (redacted to sentinel)."""
now = 10_000.0
agent = _Agent()
agent._current_turn_timestamp = now
history = [
{"role": "user", "content": "confirm forced restart", "timestamp": now - 120.0},
{"role": "assistant", "content": "a1"},
{"role": "user", "content": "u2", "timestamp": now},
]
send_request, _ = build_api_messages(
agent, history, current_turn_user_idx=len(history) - 1, ext_prefetch_cache="",
plugin_user_context="", moa_config=None, active_system_prompt="",
)
wire_send = _wire(send_request)
wire_replay = _wire(strip_stale_dangerous_confirmations(copy.deepcopy(history), now=now))
assert _canon(wire_send) == _canon(wire_replay)
assert any("EXPIRED" in (m.get("content") or "") for m in wire_send)
def test_send_byte_identity_with_fresh_confirmation():
"""Fresh confirmation (<60s): send path wire bytes match gateway replay (preserved verbatim)."""
now = 10_000.0
agent = _Agent()
agent._current_turn_timestamp = now
history = [
{"role": "user", "content": "confirm reboot", "timestamp": now - 30.0},
{"role": "assistant", "content": "a1"},
{"role": "user", "content": "u2", "timestamp": now},
]
send_request, _ = build_api_messages(
agent, history, current_turn_user_idx=len(history) - 1, ext_prefetch_cache="",
plugin_user_context="", moa_config=None, active_system_prompt="",
)
wire_send = _wire(send_request)
wire_replay = _wire(strip_stale_dangerous_confirmations(copy.deepcopy(history), now=now))
assert _canon(wire_send) == _canon(wire_replay)
assert wire_send[0]["content"] == "confirm reboot"
def test_send_byte_identity_clean_history():
"""Clean history (no interrupted blocks, no stale confirmations): send matches replay."""
now = 10_000.0
agent = _Agent()
agent._current_turn_timestamp = now
history = [
{"role": "user", "content": "u1"},
{"role": "assistant", "content": "a1"},
{"role": "user", "content": "u2"},
]
send_request, _ = build_api_messages(
agent, history, current_turn_user_idx=len(history) - 1, ext_prefetch_cache="",
plugin_user_context="", moa_config=None, active_system_prompt="",
)
wire_send = _wire(send_request)
wire_replay = _wire(canonicalize_replay_history(copy.deepcopy(history), now=now))
assert _canon(wire_send) == _canon(wire_replay) == _canon(_wire(history))
def test_send_byte_identity_with_sidecar():
"""Historical user turn with api_content sidecar: wire representation reproduces sidecar bytes."""
now = 10_000.0
agent = _Agent()
agent._current_turn_timestamp = now
history = [
{"role": "user", "content": "hello", "api_content": "hello [with memory]", "timestamp": now - 100},
{"role": "assistant", "content": "hi"},
{"role": "user", "content": "current", "timestamp": now},
]
send_request, _ = build_api_messages(
agent, history, current_turn_user_idx=len(history) - 1, ext_prefetch_cache="",
plugin_user_context="", moa_config=None, active_system_prompt="",
)
wire_send = _wire(send_request)
assert wire_send[0]["content"] == "hello [with memory]"
# Replay with sidecar intact matches
replay = canonicalize_replay_history(copy.deepcopy(history), now=now)
assert replay[0]["api_content"] == "hello [with memory]"
def test_active_turn_expiry_decision_frozen_across_tool_iterations(monkeypatch):
"""Deterministic wire-level probe (ehz0ah blocking defect):
Confirmation timestamp: 9941.0.
Active turn starts: 10000.0 (age = 59s <= 60s, fresh).
Request 1 assembled at 10000.0 sends 'confirm reboot' on wire.
Tool executes; request 2 assembled at 10002.0 (age = 61s > 60s).
Because both requests belong to the same active turn, the expiry decision
is frozen: request 2 does NOT rewrite the confirmation to EXPIRED, preserving
the prompt cache prefix across tool iterations.
Subsequent turn N+1 at 10070.0 DOES expire the confirmation and matches replay.
"""
from agent.turn_context import _reset_per_turn_agent_state
agent = _Agent()
history = [
{"role": "user", "content": "confirm reboot", "timestamp": 9941.0},
{"role": "assistant", "content": "Preparing reboot..."},
{"role": "user", "content": "proceed now", "timestamp": 10000.0},
]
turn_user_idx = 2
# Request 1 (first iteration of turn at t=10000.0):
monkeypatch.setattr(time, "time", lambda: 10000.0)
req1, _ = build_api_messages(
agent, history, current_turn_user_idx=turn_user_idx, ext_prefetch_cache="",
plugin_user_context="", moa_config=None, active_system_prompt="",
)
wire1 = _wire(req1)
assert wire1[0]["content"] == "confirm reboot"
# Tool executes during this turn; model response + tool result appended to history:
history.append({
"role": "assistant", "content": "",
"tool_calls": [{"id": "c1", "type": "function", "function": {"name": "reboot_check", "arguments": "{}"}}],
})
history.append({"role": "tool", "tool_call_id": "c1", "content": "ready"})
# Request 2 (second iteration of the SAME turn at t=10002.0 > 60s expiry threshold):
monkeypatch.setattr(time, "time", lambda: 10002.0)
req2, _ = build_api_messages(
agent, history, current_turn_user_idx=turn_user_idx, ext_prefetch_cache="",
plugin_user_context="", moa_config=None, active_system_prompt="",
)
wire2 = _wire(req2)
# The confirmation row MUST NOT have mutated to EXPIRED mid-turn:
assert wire2[0]["content"] == "confirm reboot"
# The prefix (all messages prior to the new tool call) remains byte-identical:
assert _canon(wire1) == _canon(wire2[:len(wire1)])
# Turn N+1: user sends a new message at t=10070.0 (well past expiry):
_reset_per_turn_agent_state(agent)
history.append({"role": "user", "content": "system status", "timestamp": 10070.0})
monkeypatch.setattr(time, "time", lambda: 10070.0)
req3, _ = build_api_messages(
agent, history, current_turn_user_idx=len(history) - 1, ext_prefetch_cache="",
plugin_user_context="", moa_config=None, active_system_prompt="",
)
wire3 = _wire(req3)
# In the new turn, the confirmation row IS expired on wire:
assert "EXPIRED" in wire3[0]["content"]
# Wire matches replay canonicalization at this turn boundary:
wire_replay = _wire(canonicalize_replay_history(copy.deepcopy(history[:-1]), now=10070.0) + [history[-1]])
assert _canon(wire3) == _canon(wire_replay)
def test_canonicalize_is_idempotent_and_non_mutating():
"""canonicalize_replay_history is idempotent and does not mutate source."""
now = 10_000.0
history = [
_user("u1"),
{"role": "assistant", "content": "a1"},
_assistant_tc("read_file"), _tool("[command interrupted]"),
{"role": "user", "content": "confirm forced restart", "timestamp": now - 120},
{"role": "assistant", "content": "a2"},
_user("u3"),
]
original = copy.deepcopy(history)
out1 = canonicalize_replay_history(copy.deepcopy(history), now=now)
out2 = canonicalize_replay_history(copy.deepcopy(out1), now=now)
assert _canon(_wire(out1)) == _canon(_wire(out2))
assert history == original
def test_db_roundtrip_byte_identity():
"""SessionDB round-trip: stored messages read back and canonicalized match wire."""
messages = [
{"role": "user", "content": "check system"},
{"role": "assistant", "content": "all ok"},
{"role": "user", "content": "proceed"},
]
with tempfile.TemporaryDirectory() as tmp:
db = SessionDB(db_path=Path(tmp) / "t.db")
try:
db.create_session(session_id="s1", source="cli")
for m in messages:
db.append_message("s1", role=m["role"], content=m["content"])
conv = db.get_messages_as_conversation("s1")
read_back = [
{"role": m["role"], "content": m["content"]}
for m in conv
if m.get("content") is not None
]
canon_read = canonicalize_replay_history(read_back)
assert _canon(_wire(canon_read)) == _canon(_wire(messages))
finally:
db.close()
def test_canonicalize_replay_history_handles_malformed_timestamps():
"""Malformed or non-numeric timestamps must not raise exceptions and be safely preserved."""
now = 10_000.0
history = [
{"role": "user", "content": "confirm reboot", "timestamp": "2026-09-08T00:00:00Z"},
{"role": "user", "content": "confirm reboot", "timestamp": "not_a_number"},
{"role": "user", "content": "confirm reboot", "timestamp": None},
{"role": "user", "content": "confirm reboot", "timestamp": now - 120.0},
]
agent = _Agent()
# Should not raise TypeError or ValueError
canon = canonicalize_replay_history(history, now=now)
assert len(canon) == 4
# String / None timestamps are left untouched (not expired)
assert canon[0]["content"] == "confirm reboot"
assert canon[1]["content"] == "confirm reboot"
assert canon[2]["content"] == "confirm reboot"
# The valid numeric timestamp older than 60s is expired cleanly
assert "EXPIRED" in canon[3]["content"]
# Also verify build_api_messages with string timestamp on current_turn_message
res, _ = build_api_messages(
agent,
[{"role": "user", "content": "test", "timestamp": "2026-09-08T00:00:00Z"}],
current_turn_user_idx=0,
ext_prefetch_cache="",
plugin_user_context="",
moa_config=None,
active_system_prompt="",
)
assert len(res) == 1
assert res[0]["content"] == "test"