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:
non ``[A-Za-z0-9_-]`` → ``_``, capped, plus a content hash when changed so distinct IDs cannot collide."""
raw = str(session_id or "").strip()
sanitized = re.sub(r"[^\w-]", "_", raw).strip("._")
sanitized = sanitized[:96] or "session"
sanitized = re.sub(r"[^\w-]", "_", raw).strip("._")[:96] or "session"
if raw and sanitized == raw:
return sanitized
digest = hashlib.sha256(raw.encode("utf-8", errors="surrogatepass")).hexdigest()[:12]
return f"{sanitized}_{digest}"
return f"{sanitized}_{hashlib.sha256(raw.encode('utf-8', errors='surrogatepass')).hexdigest()[:12]}"
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]``."""
if _is_multimodal_tool_result(content):
return _multimodal_text_summary(content)
if isinstance(content, list):
txt = []
for p in content:
if not isinstance(p, dict):
continue
if p.get("type") == "text":
txt.append(str(p.get("text", "")))
elif p.get("type") in _IMAGE_PART_TYPES:
txt.append("[screenshot]")
return "\n".join(txt) if txt else None
return content
if not isinstance(content, list):
return content
txt = [
str(p.get("text", "")) if p.get("type") == "text" else "[screenshot]"
for p in content
if isinstance(p, dict) and (p.get("type") == "text" or p.get("type") in _IMAGE_PART_TYPES)
]
return "\n".join(txt) if txt else None
def _tool_calls_data(msg: Dict) -> Any:
@@ -121,6 +116,11 @@ def _tool_calls_data(msg: Dict) -> Any:
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) ---
@@ -138,12 +138,10 @@ def _db_flush_seed_ids(agent) -> set:
def _db_flush_scan_start(agent, messages: List[Dict]) -> int:
"""Skip the identity-matched, still-marked prefix of the previous flush's snapshot."""
scan_start = 0
prev_prefix = getattr(agent, "_db_flush_scan_prefix", None)
if isinstance(prev_prefix, list):
for prev, cur in zip(prev_prefix, messages):
if cur is not prev or not cur.get(_DB_PERSISTED_MARKER):
break
scan_start += 1
for prev, cur in zip(getattr(agent, "_db_flush_scan_prefix", None) or (), messages):
if cur is not prev or not cur.get(_DB_PERSISTED_MARKER):
break
scan_start += 1
return scan_start
@@ -209,11 +207,9 @@ def _db_flush_collect(agent, messages: List[Dict], conversation_history: Optiona
batch_msgs: List[Dict] = []
for msg_idx in range(_db_flush_scan_start(agent, messages), len(messages)):
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
# 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
# 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:
@@ -309,6 +305,22 @@ def _session_log_entry(agent, msg: Dict[str, Any]) -> Dict[str, Any]:
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:
"""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):
"""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."""
# Close and turn-start persistence can run on separate CLI threads: one critical section.
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._session_messages = 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):
"""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)
def _flush_messages_to_session_db_unlocked(
self,
messages: List[Dict],
conversation_history: Optional[List[Dict]] = None,
_adoption_budget: int = 1,
self, 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
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."""
# 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.
if getattr(self, "_persist_disabled", False):
return None
if not self._session_db:
if getattr(self, "_persist_disabled", False) or not self._session_db:
return None
batch_rows: List[Dict[str, Any]] = []
try:
@@ -420,7 +426,6 @@ class SessionPersistenceMixin:
return messages.copy()
_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")
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)
_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")
@staticmethod
@@ -459,8 +463,7 @@ class SessionPersistenceMixin:
def _save_session_log(self, messages: List[Dict[str, Any]] = None):
"""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,
never overwriting a larger log with fewer messages (resumed agent with partial history)."""
external tooling; state.db is canonical. Rewrites the full list after every persistence point."""
if not getattr(self, "_session_json_enabled", False):
return
messages = messages or self._session_messages
@@ -468,28 +471,14 @@ class SessionPersistenceMixin:
return
# Re-derive the path each call so /branch and /compress land in the right file.
try:
safe_sid = _safe_session_filename_component(self.session_id)
log_file = self.logs_dir / f"session_{safe_sid}.json"
log_file = self.logs_dir / f"session_{_safe_session_filename_component(self.session_id)}.json"
except Exception:
return
try:
# 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)]
if log_file.exists():
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
if _existing_log_is_larger(log_file, len(cleaned)):
return
entry = {
"session_id": self.session_id,
"model": self.model,