refactor(agent/session_persistence): shared _persist_lock, _existing_log_is_larger phase helper, comprehension durable-content
This commit is contained in:
@@ -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,
|
||||||
|
|||||||
Reference in New Issue
Block a user