refactor(agent/session_persistence): tighten seed/scan/row helpers, hoist transcript_repair import

This commit is contained in:
Teknium
2026-09-02 19:03:31 -07:00
parent 649c34dcb3
commit c47195762e
+14 -36
View File
@@ -22,6 +22,7 @@ from agent.memory_manager import sanitize_context
from agent.redact import redact_sensitive_text
from agent.tool_dispatch_helpers import _is_multimodal_tool_result, _multimodal_text_summary
from agent.trajectory import convert_scratchpad_to_think, save_trajectory as _save_trajectory_to_file
from agent.transcript_repair import sync_flushed_message_markers
from utils import atomic_json_write
# Same logger name as the origin module so log records / caplog filters are unchanged.
@@ -134,15 +135,11 @@ def _db_flush_seed_ids(agent) -> set:
"""One-shot ``_flushed_db_message_ids`` seed, honoured only for the same session after a
non-empty flush; translated to markers by the scan and cleared afterwards."""
current_session_id = getattr(agent, "session_id", None)
flushed_session_id = getattr(agent, "_flushed_db_message_session_id", None)
if flushed_session_id != current_session_id or agent._last_flushed_db_idx == 0:
seed_ids = set()
else:
seed_ids = None
if getattr(agent, "_flushed_db_message_session_id", None) == current_session_id and agent._last_flushed_db_idx != 0:
seed_ids = getattr(agent, "_flushed_db_message_ids", None)
if not isinstance(seed_ids, set):
seed_ids = set()
agent._flushed_db_message_session_id = current_session_id
return seed_ids
return seed_ids if isinstance(seed_ids, set) else set()
def _db_flush_scan_start(agent, messages: List[Dict]) -> int:
@@ -151,12 +148,9 @@ def _db_flush_scan_start(agent, messages: List[Dict]) -> int:
scan_start = 0
prev_prefix = getattr(agent, "_db_flush_scan_prefix", None)
if isinstance(prev_prefix, list):
limit = min(len(prev_prefix), len(messages))
while (
scan_start < limit
and messages[scan_start] is prev_prefix[scan_start]
and bool(messages[scan_start].get(_DB_PERSISTED_MARKER))
):
for prev, cur in zip(prev_prefix, messages):
if cur is not prev or not cur.get(_DB_PERSISTED_MARKER):
break
scan_start += 1
return scan_start
@@ -167,9 +161,7 @@ def _db_flush_row(agent, msg: Dict, is_current_turn_user: bool) -> Dict[str, Any
content = msg.get("content")
# api_content sidecar: exact bytes sent to the API when they differ from clean content, so
# replay reproduces the sent prefix byte-for-byte.
api_content = msg.get("api_content")
if not isinstance(api_content, str):
api_content = None
api_content = msg.get("api_content") if isinstance(msg.get("api_content"), str) else None
timestamp = msg.get("timestamp")
if is_current_turn_user and msg.get("role") == "user":
override = getattr(agent, "_persist_user_message_override", None)
@@ -186,14 +178,9 @@ def _db_flush_row(agent, msg: Dict, is_current_turn_user: bool) -> Dict[str, Any
api_content = None
# get_messages_as_conversation replays rows through sanitize_context().strip(); capture the
# sent bytes when they would differ (compared in wire form).
if (
api_content is None
and role in ("user", "assistant")
and isinstance(content, str)
and content
and sanitize_context(content).strip() != content.strip()
):
api_content = content
if api_content is None and role in ("user", "assistant") and isinstance(content, str) and content:
if sanitize_context(content).strip() != content.strip():
api_content = content
# Key order is the divert-JSONL wire order (divert_session_transcript_jsonl).
row = {
"role": role,
@@ -255,8 +242,6 @@ def _db_flush_write(agent, batch_rows: List[Dict[str, Any]], batch_msgs: List[Di
turn_lease_holder=getattr(agent, "_active_session_turn_lease_holder", None),
turn_lease_ttl_seconds=getattr(agent, "_active_session_turn_lease_ttl_seconds", 300.0) or 300.0,
)
from agent.transcript_repair import sync_flushed_message_markers
sync_flushed_message_markers(batch_msgs, batch_rows)
@@ -267,11 +252,11 @@ def _db_flush_adopt_compression_tip(agent) -> bool:
is missing or already ended is not adopted either.
"""
old_id = agent.session_id
tip = None
try:
tip = agent._session_db.get_compression_tip(old_id)
except Exception as tip_exc:
logger.warning("compression tip lookup failed for %s: %s", old_id, tip_exc)
return False
if not tip or tip == old_id:
return False
try:
@@ -339,9 +324,7 @@ class SessionPersistenceMixin:
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 not 0 <= idx < len(messages):
return
msg = messages[idx]
msg = messages[idx] if 0 <= idx < len(messages) else None
if not (isinstance(msg, dict) and msg.get("role") == "user"):
return
if _override_replaces_content(msg, msg.get("content"), override):
@@ -391,12 +374,7 @@ class SessionPersistenceMixin:
while messages and isinstance(messages[-1], dict) and messages[-1].get("role") == "tool":
messages.pop()
# Providers reject a dangling assistant(tool_calls) whose results were just popped.
if (
messages
and isinstance(messages[-1], dict)
and messages[-1].get("role") == "assistant"
and messages[-1].get("tool_calls")
):
if messages and isinstance(messages[-1], dict) and messages[-1].get("role") == "assistant" and messages[-1].get("tool_calls"):
messages.pop()
_repair_message_sequence = _forward("agent.agent_runtime_helpers", "repair_message_sequence")