refactor(agent/session_persistence): shared _persist_lock, _existing_log_is_larger phase helper, comprehension durable-content

This commit is contained in:
Teknium
2026-09-02 21:21:34 -07:00
parent f6fe5452de
commit b41e73d0e7
+44 -55
View File
@@ -62,12 +62,10 @@ def _safe_session_filename_component(session_id: str) -> str:
"""Path-safe filename component for a (possibly untrusted ``X-Hermes-Session-Id``) session ID: """Path-safe filename component for a (possibly untrusted ``X-Hermes-Session-Id``) session ID:
non ``[A-Za-z0-9_-]`` → ``_``, capped, plus a content hash when changed so distinct IDs cannot collide.""" non ``[A-Za-z0-9_-]`` → ``_``, capped, plus a content hash when changed so distinct IDs cannot collide."""
raw = str(session_id or "").strip() raw = str(session_id or "").strip()
sanitized = re.sub(r"[^\w-]", "_", raw).strip("._") sanitized = re.sub(r"[^\w-]", "_", raw).strip("._")[:96] or "session"
sanitized = sanitized[:96] or "session"
if raw and sanitized == raw: if raw and sanitized == raw:
return sanitized return sanitized
digest = hashlib.sha256(raw.encode("utf-8", errors="surrogatepass")).hexdigest()[:12] return f"{sanitized}_{hashlib.sha256(raw.encode('utf-8', errors='surrogatepass')).hexdigest()[:12]}"
return f"{sanitized}_{digest}"
def _override_replaces_content(msg: Dict, content: Any, override: Any) -> bool: def _override_replaces_content(msg: Dict, content: Any, override: Any) -> bool:
@@ -100,17 +98,14 @@ def _durable_content(content: Any) -> Any:
"""Text-only DB projection: multimodal envelopes → summary; part lists keep text, images → ``[screenshot]``.""" """Text-only DB projection: multimodal envelopes → summary; part lists keep text, images → ``[screenshot]``."""
if _is_multimodal_tool_result(content): if _is_multimodal_tool_result(content):
return _multimodal_text_summary(content) return _multimodal_text_summary(content)
if isinstance(content, list): if not isinstance(content, list):
txt = [] return content
for p in content: txt = [
if not isinstance(p, dict): str(p.get("text", "")) if p.get("type") == "text" else "[screenshot]"
continue for p in content
if p.get("type") == "text": if isinstance(p, dict) and (p.get("type") == "text" or p.get("type") in _IMAGE_PART_TYPES)
txt.append(str(p.get("text", ""))) ]
elif p.get("type") in _IMAGE_PART_TYPES: return "\n".join(txt) if txt else None
txt.append("[screenshot]")
return "\n".join(txt) if txt else None
return content
def _tool_calls_data(msg: Dict) -> Any: def _tool_calls_data(msg: Dict) -> Any:
@@ -121,6 +116,11 @@ def _tool_calls_data(msg: Dict) -> Any:
return None return None
def _persist_lock(agent):
"""Close and turn-start persistence can run on separate CLI threads: one critical section."""
return getattr(agent, "_session_persist_lock", None) or nullcontext()
# --- flush phases (module-level so the flush also works bound onto duck-typed agents) --- # --- flush phases (module-level so the flush also works bound onto duck-typed agents) ---
@@ -138,12 +138,10 @@ def _db_flush_seed_ids(agent) -> set:
def _db_flush_scan_start(agent, messages: List[Dict]) -> int: def _db_flush_scan_start(agent, messages: List[Dict]) -> int:
"""Skip the identity-matched, still-marked prefix of the previous flush's snapshot.""" """Skip the identity-matched, still-marked prefix of the previous flush's snapshot."""
scan_start = 0 scan_start = 0
prev_prefix = getattr(agent, "_db_flush_scan_prefix", None) for prev, cur in zip(getattr(agent, "_db_flush_scan_prefix", None) or (), messages):
if isinstance(prev_prefix, list): if cur is not prev or not cur.get(_DB_PERSISTED_MARKER):
for prev, cur in zip(prev_prefix, messages): break
if cur is not prev or not cur.get(_DB_PERSISTED_MARKER): scan_start += 1
break
scan_start += 1
return scan_start return scan_start
@@ -209,11 +207,9 @@ def _db_flush_collect(agent, messages: List[Dict], conversation_history: Optiona
batch_msgs: List[Dict] = [] batch_msgs: List[Dict] = []
for msg_idx in range(_db_flush_scan_start(agent, messages), len(messages)): for msg_idx in range(_db_flush_scan_start(agent, messages), len(messages)):
msg = messages[msg_idx] msg = messages[msg_idx]
if not isinstance(msg, dict):
continue
# The flush is append-only: a mid-turn persist of scaffolding could commit a synthetic # The flush is append-only: a mid-turn persist of scaffolding could commit a synthetic
# turn the end-of-turn drop cannot un-write. Skip regardless of position. # turn the end-of-turn drop cannot un-write. Skip regardless of position.
if _is_ephemeral_scaffolding(msg) or msg.get(_DB_PERSISTED_MARKER): if not isinstance(msg, dict) or _is_ephemeral_scaffolding(msg) or msg.get(_DB_PERSISTED_MARKER):
continue continue
# Already durable (history copy or caller-seeded): stamp so future flushes skip it. # Already durable (history copy or caller-seeded): stamp so future flushes skip it.
if id(msg) in history_ids or id(msg) in seed_ids: if id(msg) in history_ids or id(msg) in seed_ids:
@@ -309,6 +305,22 @@ def _session_log_entry(agent, msg: Dict[str, Any]) -> Dict[str, Any]:
return {**msg, "content": agent._redact_message_content(content)} return {**msg, "content": agent._redact_message_content(content)}
def _existing_log_is_larger(log_file, count: int) -> bool:
"""Never overwrite a larger log with fewer messages (resumed agent with partial history);
a corrupted existing file allows the overwrite."""
if not log_file.exists():
return False
try:
existing = json.loads(log_file.read_text(encoding="utf-8"))
existing_count = existing.get("message_count", len(existing.get("messages", [])))
except Exception:
return False
if existing_count > count:
logging.debug("Skipping session log overwrite: existing has %d messages, current has %d", existing_count, count)
return True
return False
class SessionPersistenceMixin: class SessionPersistenceMixin:
"""Session DB flush, session log and trajectory persistence (see module docstring).""" """Session DB flush, session log and trajectory persistence (see module docstring)."""
@@ -335,10 +347,9 @@ class SessionPersistenceMixin:
def _persist_session(self, messages: List[Dict], conversation_history: List[Dict] = None): 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. Trailing empty-response """Save session state to both JSON log and SQLite on any exit path. Trailing empty-response
scaffolding is dropped from the live list; the persist override is applied to the DB row only.""" scaffolding is dropped from the live list; the persist override is applied to the DB row only."""
# Close and turn-start persistence can run on separate CLI threads: one critical section.
from agent.agent_runtime_helpers import note_turn_persisted from agent.agent_runtime_helpers import note_turn_persisted
with getattr(self, "_session_persist_lock", None) or nullcontext(): with _persist_lock(self):
self._drop_trailing_empty_response_scaffolding(messages) self._drop_trailing_empty_response_scaffolding(messages)
self._session_messages = messages self._session_messages = messages
self._save_session_log(messages) self._save_session_log(messages)
@@ -374,14 +385,11 @@ class SessionPersistenceMixin:
def _flush_messages_to_session_db(self, messages: List[Dict], conversation_history: Optional[List[Dict]] = None): def _flush_messages_to_session_db(self, messages: List[Dict], conversation_history: Optional[List[Dict]] = None):
"""Serialize direct and turn-boundary session flushes per agent.""" """Serialize direct and turn-boundary session flushes per agent."""
with getattr(self, "_session_persist_lock", None) or nullcontext(): with _persist_lock(self):
return self._flush_messages_to_session_db_unlocked(messages, conversation_history) return self._flush_messages_to_session_db_unlocked(messages, conversation_history)
def _flush_messages_to_session_db_unlocked( def _flush_messages_to_session_db_unlocked(
self, self, messages: List[Dict], conversation_history: Optional[List[Dict]] = None, _adoption_budget: int = 1,
messages: List[Dict],
conversation_history: Optional[List[Dict]] = None,
_adoption_budget: int = 1,
): ):
"""Persist un-flushed messages to SQLite. Dedup is the intrinsic ``_DB_PERSISTED_MARKER`` on """Persist un-flushed messages to SQLite. Dedup is the intrinsic ``_DB_PERSISTED_MARKER`` on
each written dict — not positional slices (drift after sequence repair) nor an ``id(msg)`` set each written dict — not positional slices (drift after sequence repair) nor an ``id(msg)`` set
@@ -389,9 +397,7 @@ class SessionPersistenceMixin:
session adopts its live tip and retries exactly once.""" session adopts its live tip and retries exactly once."""
# Persistence-isolated agents (background review fork) share the parent's session_id for cache # Persistence-isolated agents (background review fork) share the parent's session_id for cache
# warmth; a write here would land the curator's turn in the user's real history. # warmth; a write here would land the curator's turn in the user's real history.
if getattr(self, "_persist_disabled", False): if getattr(self, "_persist_disabled", False) or not self._session_db:
return None
if not self._session_db:
return None return None
batch_rows: List[Dict[str, Any]] = [] batch_rows: List[Dict[str, Any]] = []
try: try:
@@ -420,7 +426,6 @@ class SessionPersistenceMixin:
return messages.copy() return messages.copy()
_format_tools_for_system_message = _forward("agent.system_prompt", "format_tools_for_system_message") _format_tools_for_system_message = _forward("agent.system_prompt", "format_tools_for_system_message")
_convert_to_trajectory_format = _forward("agent.agent_runtime_helpers", "convert_to_trajectory_format") _convert_to_trajectory_format = _forward("agent.agent_runtime_helpers", "convert_to_trajectory_format")
def _save_trajectory(self, messages: List[Dict[str, Any]], user_query: str, completed: bool): def _save_trajectory(self, messages: List[Dict[str, Any]], user_query: str, completed: bool):
@@ -431,7 +436,6 @@ class SessionPersistenceMixin:
_save_trajectory_to_file(trajectory, self.model, completed) _save_trajectory_to_file(trajectory, self.model, completed)
_extract_api_error_context = _forward_static("agent.agent_runtime_helpers", "extract_api_error_context") _extract_api_error_context = _forward_static("agent.agent_runtime_helpers", "extract_api_error_context")
_dump_api_request_debug = _forward("agent.agent_runtime_helpers", "dump_api_request_debug") _dump_api_request_debug = _forward("agent.agent_runtime_helpers", "dump_api_request_debug")
@staticmethod @staticmethod
@@ -459,8 +463,7 @@ class SessionPersistenceMixin:
def _save_session_log(self, messages: List[Dict[str, Any]] = None): def _save_session_log(self, messages: List[Dict[str, Any]] = None):
"""Optional per-session JSON snapshot (``sessions.write_json_snapshots``, default False) for """Optional per-session JSON snapshot (``sessions.write_json_snapshots``, default False) for
external tooling; state.db is canonical. Rewrites the full list after every persistence point, external tooling; state.db is canonical. Rewrites the full list after every persistence point."""
never overwriting a larger log with fewer messages (resumed agent with partial history)."""
if not getattr(self, "_session_json_enabled", False): if not getattr(self, "_session_json_enabled", False):
return return
messages = messages or self._session_messages messages = messages or self._session_messages
@@ -468,28 +471,14 @@ class SessionPersistenceMixin:
return return
# Re-derive the path each call so /branch and /compress land in the right file. # Re-derive the path each call so /branch and /compress land in the right file.
try: try:
safe_sid = _safe_session_filename_component(self.session_id) log_file = self.logs_dir / f"session_{_safe_session_filename_component(self.session_id)}.json"
log_file = self.logs_dir / f"session_{safe_sid}.json"
except Exception: except Exception:
return return
try: try:
# Mirror the SQLite flush: scaffolding is never durable transcript content. # Mirror the SQLite flush: scaffolding is never durable transcript content.
cleaned = [_session_log_entry(self, msg) for msg in messages if not _is_ephemeral_scaffolding(msg)] cleaned = [_session_log_entry(self, msg) for msg in messages if not _is_ephemeral_scaffolding(msg)]
if _existing_log_is_larger(log_file, len(cleaned)):
if log_file.exists(): return
try:
existing = json.loads(log_file.read_text(encoding="utf-8"))
existing_count = existing.get("message_count", len(existing.get("messages", [])))
if existing_count > len(cleaned):
logging.debug(
"Skipping session log overwrite: existing has %d messages, current has %d",
existing_count, len(cleaned),
)
return
except Exception:
pass # corrupted existing file — allow the overwrite
entry = { entry = {
"session_id": self.session_id, "session_id": self.session_id,
"model": self.model, "model": self.model,