1590 lines
81 KiB
Python
1590 lines
81 KiB
Python
"""Transcript persistence for SessionDB.
|
|
|
|
Mixin bound onto ``SessionDB`` via the MRO, built on its ``_read_ctx`` /
|
|
``_execute_write`` / ``_write_sql`` / ``_read_one`` / ``_read_all`` primitives.
|
|
Covers message append / replace / rewind, reactions, resume-conversation
|
|
assembly and replayed-user-message duplicate detection.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import logging
|
|
import time
|
|
from typing import Any, Dict, List, Optional, Tuple
|
|
|
|
from agent.context_compressor import _DB_PERSISTED_MARKER as _DB_PERSISTED_MARKER_KEY
|
|
from agent.memory_manager import sanitize_context
|
|
from agent.message_sanitization import _sanitize_surrogates
|
|
from hermes_state_common import _RESET_END_REASONS, _RESET_END_REASONS_SQL, _legacy_reset_child_sql
|
|
|
|
# Log-record parity with the origin module (caplog tests pin "hermes_state").
|
|
logger = logging.getLogger("hermes_state")
|
|
|
|
# One INSERT shape for every message writer (append, batch, replace, compact, import).
|
|
_INSERT_MESSAGE_SQL = """INSERT INTO messages (session_id, role, content, tool_call_id,
|
|
tool_calls, tool_name, effect_disposition, timestamp, token_count, finish_reason,
|
|
reasoning, reasoning_content, reasoning_details, codex_reasoning_items,
|
|
codex_message_items, platform_message_id, observed, _compressed_summary, active, api_content, display_kind, display_metadata)
|
|
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)"""
|
|
|
|
_ENDED_BY_COMPRESSION_SQL = "SELECT ended_at, end_reason FROM sessions WHERE id = ?"
|
|
_COMPRESSION_LOCK_ROW_SQL = "SELECT holder, expires_at FROM compression_locks WHERE session_id = ?"
|
|
_TURN_LEASE_ROW_SQL = "SELECT holder, expires_at FROM session_turn_leases WHERE conversation_id = ?"
|
|
_DISPLAY_ACTIVE_CLAUSE = " AND (active = 1 OR compacted = 1)"
|
|
_DELETE_COMPRESSION_LOCK_SQL = "DELETE FROM compression_locks WHERE session_id = ? AND holder = ?"
|
|
|
|
|
|
def _placeholders(items) -> str:
|
|
return ",".join("?" for _ in items)
|
|
|
|
|
|
def _json_or(raw: Any, fallback: Any, warning: str) -> Any:
|
|
"""``json.loads(raw)``; on failure log *warning* and return *fallback*."""
|
|
try:
|
|
return json.loads(raw)
|
|
except (json.JSONDecodeError, TypeError):
|
|
logger.warning(warning)
|
|
return fallback
|
|
|
|
|
|
def _tool_calls_len(raw: Any, scalar: int = 0) -> int:
|
|
"""Tool-call count of a stored ``tool_calls`` column: list length, else *scalar*
|
|
for a truthy non-list value, 0 for empty/undecodable."""
|
|
if not raw:
|
|
return 0
|
|
try:
|
|
parsed = json.loads(raw) if isinstance(raw, str) else raw
|
|
except (TypeError, ValueError):
|
|
return 0
|
|
if isinstance(parsed, list):
|
|
return len(parsed)
|
|
return scalar if parsed else 0
|
|
|
|
|
|
def _coerce_timestamp(value: Any, default: float) -> float:
|
|
"""Explicit message timestamp (datetime or number) or *default* when invalid."""
|
|
if value is None:
|
|
return default
|
|
try:
|
|
return float(value.timestamp()) if hasattr(value, "timestamp") else float(value)
|
|
except (TypeError, ValueError):
|
|
logger.debug("Ignoring invalid explicit message timestamp: %r", value)
|
|
return default
|
|
|
|
|
|
def _parse_tool_calls(tool_calls: Any) -> Any:
|
|
"""tool_calls may be a list (live agent) or a JSON string (import/export); parse
|
|
first so json.dumps never double-encodes."""
|
|
if isinstance(tool_calls, str):
|
|
try:
|
|
return json.loads(tool_calls)
|
|
except (json.JSONDecodeError, TypeError):
|
|
return []
|
|
return tool_calls
|
|
|
|
|
|
def _tool_calls_count(tool_calls: Any) -> int:
|
|
if tool_calls is None:
|
|
return 0
|
|
return len(tool_calls) if isinstance(tool_calls, list) else 1
|
|
|
|
|
|
def _ended_by_compression(row) -> bool:
|
|
return row is not None and row["ended_at"] is not None and row["end_reason"] == "compression"
|
|
|
|
|
|
# _rows_to_conversation copies these row columns verbatim when truthy, in this order.
|
|
# ``api_content`` is returned VERBATIM (no sanitize/strip): the replay path substitutes
|
|
# it to keep the provider prompt cache byte-stable.
|
|
_VERBATIM_COLS = ("api_content", "display_kind")
|
|
_META_COLS = ("timestamp", "tool_call_id", "tool_name", "effect_disposition")
|
|
_ASSISTANT_JSON_COLS = ("reasoning_details", "codex_reasoning_items", "codex_message_items")
|
|
|
|
|
|
def _stale_holder(row, now: float) -> bool:
|
|
"""A lock/lease row whose holder is expired or a provably dead local process."""
|
|
from hermes_state import _compression_lock_holder_process_is_dead
|
|
return float(row["expires_at"]) <= now or _compression_lock_holder_process_is_dead(row["holder"])
|
|
|
|
|
|
class SessionMessagesMixin:
|
|
"""Message append/replace/rewind, reactions, resume conversations, replay dedupe."""
|
|
|
|
def _bump_conversation_generation(self, conn, session_id: str, end_reason: str) -> None:
|
|
"""Advance this peer's conversation generation past a boundary, inside the
|
|
transaction that writes the boundary.
|
|
|
|
Only ``_RESET_END_REASONS`` count (``compression`` continues one conversation).
|
|
The counter never reads the session rows: an aggregate over them could re-emit a
|
|
pair once ``delete_session()``/pruning removes an ended row and hand a new
|
|
conversation a retired affinity identity. It only ever increments.
|
|
"""
|
|
if end_reason not in _RESET_END_REASONS:
|
|
return
|
|
row = conn.execute("SELECT source, session_key FROM sessions WHERE id = ?", (session_id,)).fetchone()
|
|
if row is None:
|
|
return
|
|
source = str(row["source"] or "").strip()
|
|
session_key = str(row["session_key"] or "").strip()
|
|
if not source or not session_key:
|
|
return
|
|
conn.execute(
|
|
"""
|
|
INSERT INTO conversation_generations (source, session_key, generation)
|
|
VALUES (?, ?, 1)
|
|
ON CONFLICT(source, session_key) DO UPDATE
|
|
SET generation = conversation_generations.generation + 1
|
|
""",
|
|
(source, session_key),
|
|
)
|
|
|
|
@classmethod
|
|
def _encode_content(cls, content: Any) -> Any:
|
|
"""Serialize list/dict content (multimodal parts) as a sentinel-prefixed JSON
|
|
string; sqlite3 can only bind str/bytes/int/float/None. Lone surrogates are
|
|
scrubbed from text so persistence never fails. Paired with :meth:`_decode_content`.
|
|
"""
|
|
if isinstance(content, str):
|
|
return _sanitize_surrogates(content)
|
|
if content is None or isinstance(content, (bytes, int, float)):
|
|
return content
|
|
try:
|
|
# ensure_ascii=True escapes surrogates as \\udXXX — safe to bind.
|
|
return cls._CONTENT_JSON_PREFIX + json.dumps(content)
|
|
except (TypeError, ValueError):
|
|
return _sanitize_surrogates(str(content))
|
|
|
|
@classmethod
|
|
def _decode_content(cls, content: Any) -> Any:
|
|
"""Reverse :meth:`_encode_content`; returns scalars unchanged."""
|
|
if isinstance(content, str) and content.startswith(cls._CONTENT_JSON_PREFIX):
|
|
return _json_or(
|
|
content[len(cls._CONTENT_JSON_PREFIX):], content,
|
|
"Failed to decode JSON-encoded message content; returning raw string",
|
|
)
|
|
return content
|
|
|
|
@staticmethod
|
|
def _encode_display_metadata(display_metadata: Any) -> Optional[str]:
|
|
"""Serialize ``display_metadata`` for its TEXT column without double-encoding an
|
|
already-serialized JSON string (import/replace paths hand those in)."""
|
|
if not display_metadata:
|
|
return None
|
|
if isinstance(display_metadata, str):
|
|
try:
|
|
display_metadata = json.loads(display_metadata)
|
|
except (json.JSONDecodeError, TypeError):
|
|
logger.warning("Ignoring non-JSON display metadata on write")
|
|
return None
|
|
if not isinstance(display_metadata, dict):
|
|
logger.warning("Ignoring non-object display metadata on write")
|
|
return None
|
|
elif not isinstance(display_metadata, dict):
|
|
logger.warning(
|
|
"Ignoring unexpected display metadata type on write: %s", type(display_metadata).__name__,
|
|
)
|
|
return None
|
|
return json.dumps(display_metadata)
|
|
|
|
@staticmethod
|
|
def _decode_display_metadata(raw: Any) -> Optional[Dict[str, Any]]:
|
|
"""Decode a ``display_metadata`` column into a dict (never the raw TEXT — the
|
|
desktop does ``'task_count' in meta``). Pre-guard rows are double-encoded, so a
|
|
second string layer is unwrapped."""
|
|
if raw is None:
|
|
return None
|
|
try:
|
|
meta = json.loads(raw) if isinstance(raw, str) else raw
|
|
if isinstance(meta, str):
|
|
meta = json.loads(meta)
|
|
except (json.JSONDecodeError, TypeError):
|
|
logger.warning("Ignoring invalid display metadata on message row")
|
|
return None
|
|
if not isinstance(meta, dict):
|
|
logger.warning("Ignoring non-object display metadata on message row")
|
|
return None
|
|
return meta
|
|
|
|
@staticmethod
|
|
def _reasoning_json_text(value: Any) -> Optional[str]:
|
|
"""Serialize a structured reasoning field for its TEXT column. Strings are
|
|
stored as-is: round-tripping callers (get_messages -> replace_messages, e.g. the
|
|
fork handler) hand back the raw TEXT, and re-dumping would double-encode it so
|
|
every reasoning-replay consumer (``isinstance(..., list)``) drops it."""
|
|
if not value:
|
|
return None
|
|
return value if isinstance(value, str) else json.dumps(value)
|
|
|
|
def _check_transcript_write_guards(
|
|
self, conn, session_id: str, compression_lock_holder: Optional[str],
|
|
turn_lease_holder: Optional[str] = None, turn_lease_ttl_seconds: float = 300.0,
|
|
reject_active_turn_lease: bool = False, reject_active_compression_lock: bool = False,
|
|
allow_closed_compression_parent: bool = False,
|
|
) -> None:
|
|
"""Transcript-write admission checks, run INSIDE the write txn.
|
|
|
|
Shared by every transcript writer so they cannot diverge. Ordinary appends do
|
|
NOT check compression_locks: the lock only stops two COMPRESSIONS colliding, and
|
|
archive_and_compact() commits against a watermark and clones later rows, so
|
|
concurrent appends are safe by construction (blocking them killed turns while a
|
|
slow summary held the lease). Destructive user mutations opt in via
|
|
``reject_active_compression_lock`` / ``reject_active_turn_lease`` so a compressor
|
|
that captured its watermark cannot resurrect the removed turn.
|
|
"""
|
|
from hermes_state import CompressionSessionClosedError, SessionCompressionInProgressError, SessionTurnLeaseLostError
|
|
if reject_active_compression_lock:
|
|
active_lock = conn.execute(_COMPRESSION_LOCK_ROW_SQL, (session_id,)).fetchone()
|
|
if active_lock is not None:
|
|
if _stale_holder(active_lock, time.time()):
|
|
conn.execute(_DELETE_COMPRESSION_LOCK_SQL, (session_id, active_lock["holder"]))
|
|
elif active_lock["holder"] != compression_lock_holder:
|
|
raise SessionCompressionInProgressError(
|
|
f"Session {session_id!r} is being compressed by another writer"
|
|
)
|
|
if turn_lease_holder or reject_active_turn_lease:
|
|
conversation_id = self._session_turn_lease_key_on_conn(conn, session_id)
|
|
lease = conn.execute(_TURN_LEASE_ROW_SQL, (conversation_id,)).fetchone()
|
|
now = time.time()
|
|
if turn_lease_holder:
|
|
if lease is None or lease["holder"] != turn_lease_holder:
|
|
raise SessionTurnLeaseLostError(
|
|
f"Session turn lease lost; refusing transcript write for {session_id!r}"
|
|
)
|
|
if float(lease["expires_at"]) <= now:
|
|
# Expiry makes the row reclaimable, it does not prove a takeover;
|
|
# BEGIN IMMEDIATE serializes this renewal with acquisition, so a
|
|
# still-matching owner recovers from a starved refresher.
|
|
conn.execute(
|
|
"UPDATE session_turn_leases SET expires_at = ? "
|
|
"WHERE conversation_id = ? AND holder = ?",
|
|
(now + max(0.1, float(turn_lease_ttl_seconds)), conversation_id, turn_lease_holder),
|
|
)
|
|
elif lease is not None:
|
|
if not _stale_holder(lease, now):
|
|
raise SessionTurnLeaseLostError(
|
|
f"Session has an active turn lease; refusing transcript mutation for {session_id!r}"
|
|
)
|
|
# Same reclaim rule as acquisition (expired or provably dead owner);
|
|
# deleting here also fences a stale late flush after the mutation.
|
|
conn.execute(
|
|
"DELETE FROM session_turn_leases WHERE conversation_id = ? AND holder = ?",
|
|
(conversation_id, lease["holder"]),
|
|
)
|
|
session = conn.execute(_ENDED_BY_COMPRESSION_SQL, (session_id,)).fetchone()
|
|
if _ended_by_compression(session) and not allow_closed_compression_parent:
|
|
raise CompressionSessionClosedError(session_id)
|
|
|
|
def _message_row_params(
|
|
self, session_id: str, role: str, msg: Dict[str, Any], tool_calls: Any,
|
|
message_timestamp: float, *, keep_reasoning: bool,
|
|
) -> tuple:
|
|
"""Bind values for ``_INSERT_MESSAGE_SQL`` from one message dict.
|
|
|
|
*tool_calls* is the already-parsed value (see ``_parse_tool_calls``).
|
|
*keep_reasoning* False stores NULL for every reasoning/codex column.
|
|
"""
|
|
from hermes_state import _scrub_surrogates
|
|
|
|
def _str_or_none(value):
|
|
return _scrub_surrogates(value) if isinstance(value, str) else None
|
|
|
|
def _reasoning(key):
|
|
return msg.get(key) if keep_reasoning else None
|
|
|
|
return (
|
|
session_id,
|
|
role,
|
|
self._encode_content(msg.get("content")),
|
|
msg.get("tool_call_id"),
|
|
json.dumps(tool_calls) if tool_calls else None,
|
|
_scrub_surrogates(msg.get("tool_name")),
|
|
msg.get("effect_disposition"),
|
|
message_timestamp,
|
|
msg.get("token_count"),
|
|
msg.get("finish_reason"),
|
|
_scrub_surrogates(_reasoning("reasoning")),
|
|
_scrub_surrogates(_reasoning("reasoning_content")),
|
|
self._reasoning_json_text(_reasoning("reasoning_details")),
|
|
self._reasoning_json_text(_reasoning("codex_reasoning_items")),
|
|
self._reasoning_json_text(_reasoning("codex_message_items")),
|
|
# `message_id` is yuanbao's existing convention on message dicts.
|
|
msg.get("platform_message_id") or msg.get("message_id"),
|
|
1 if msg.get("observed") else 0,
|
|
1 if msg.get("_compressed_summary") else 0,
|
|
1,
|
|
_str_or_none(msg.get("api_content")),
|
|
_str_or_none(msg.get("display_kind")),
|
|
self._encode_display_metadata(msg.get("display_metadata")),
|
|
)
|
|
|
|
def append_message(
|
|
self, session_id: str, role: str, content: str = None, tool_name: str = None,
|
|
tool_calls: Any = None, tool_call_id: str = None, token_count: int = None,
|
|
finish_reason: str = None, reasoning: str = None, reasoning_content: str = None,
|
|
reasoning_details: Any = None, codex_reasoning_items: Any = None,
|
|
codex_message_items: Any = None, platform_message_id: str = None, observed: bool = False,
|
|
effect_disposition: Optional[str] = None, _compressed_summary: bool = False,
|
|
timestamp: Any = None, api_content: Optional[str] = None,
|
|
display_kind: Optional[str] = None, display_metadata: Optional[Dict[str, Any]] = None,
|
|
compression_lock_holder: Optional[str] = None, turn_lease_holder: Optional[str] = None,
|
|
turn_lease_ttl_seconds: float = 300.0,
|
|
) -> int:
|
|
"""Append one message; returns the row id. Bumps ``message_count`` (and
|
|
``tool_call_count`` when tool_calls are present).
|
|
|
|
``platform_message_id`` is the platform's own id (Telegram update_id, Yuanbao
|
|
msg_id) used by recall-style flows. ``api_content`` is the byte-fidelity
|
|
sidecar — the exact string sent to the API when it differed from ``content`` —
|
|
stored as sent except lone surrogates (which the loop scrubs anyway).
|
|
"""
|
|
msg = {
|
|
"content": content, "tool_name": tool_name, "tool_call_id": tool_call_id,
|
|
"token_count": token_count, "finish_reason": finish_reason, "reasoning": reasoning,
|
|
"reasoning_content": reasoning_content, "reasoning_details": reasoning_details,
|
|
"codex_reasoning_items": codex_reasoning_items, "codex_message_items": codex_message_items,
|
|
"platform_message_id": platform_message_id, "message_id": platform_message_id,
|
|
"observed": observed, "effect_disposition": effect_disposition,
|
|
"_compressed_summary": _compressed_summary, "api_content": api_content,
|
|
"display_kind": display_kind, "display_metadata": display_metadata,
|
|
}
|
|
# Encode outside the write txn (display metadata first: log-order parity).
|
|
display_metadata_json = self._encode_display_metadata(display_metadata)
|
|
msg["display_metadata"] = display_metadata_json
|
|
tool_calls = _parse_tool_calls(tool_calls)
|
|
message_timestamp = _coerce_timestamp(timestamp, time.time())
|
|
num_tool_calls = _tool_calls_count(tool_calls)
|
|
params = self._message_row_params(
|
|
session_id, role, msg, tool_calls, message_timestamp, keep_reasoning=True,
|
|
)
|
|
|
|
def _do(conn):
|
|
self._check_transcript_write_guards(
|
|
conn, session_id, compression_lock_holder,
|
|
turn_lease_holder=turn_lease_holder, turn_lease_ttl_seconds=turn_lease_ttl_seconds,
|
|
)
|
|
msg_id = conn.execute(_INSERT_MESSAGE_SQL, params).lastrowid
|
|
if num_tool_calls > 0:
|
|
conn.execute(
|
|
"""UPDATE sessions SET message_count = message_count + 1,
|
|
tool_call_count = tool_call_count + ? WHERE id = ?""",
|
|
(num_tool_calls, session_id),
|
|
)
|
|
else:
|
|
conn.execute(
|
|
"UPDATE sessions SET message_count = message_count + 1 WHERE id = ?", (session_id,),
|
|
)
|
|
return msg_id
|
|
|
|
# THE critical write (its failure aborts the turn): long patience so a sibling
|
|
# legitimately holding the lock for seconds (VACUUM, checkpoint) can't kill it.
|
|
return self._execute_write(_do, patience_s=self._TRANSCRIPT_WRITE_PATIENCE_S)
|
|
|
|
def append_messages_batch(
|
|
self, session_id: str, messages: List[Dict[str, Any]],
|
|
compression_lock_holder: Optional[str] = None, turn_lease_holder: Optional[str] = None,
|
|
chunk_rows: Optional[int] = None, turn_lease_ttl_seconds: float = 300.0,
|
|
) -> int:
|
|
"""Append *messages* (``_insert_message_rows`` dict shape) in ONE write txn.
|
|
|
|
All rows land or none do; the admission guards run once for the batch.
|
|
``chunk_rows`` bounds transaction size for LARGE copies (branch seeds: FTS
|
|
triggers run per row, 10k rows ≈ 2.4s under one lock) — commits in chunks with
|
|
the old per-row-loop recovery semantics. Returns the inserted row count.
|
|
"""
|
|
if not messages:
|
|
return 0
|
|
if chunk_rows is not None and len(messages) > chunk_rows:
|
|
return sum(
|
|
self.append_messages_batch(
|
|
session_id, messages[start:start + chunk_rows],
|
|
compression_lock_holder=compression_lock_holder,
|
|
turn_lease_holder=turn_lease_holder,
|
|
turn_lease_ttl_seconds=turn_lease_ttl_seconds,
|
|
)
|
|
for start in range(0, len(messages), chunk_rows)
|
|
)
|
|
|
|
def _do(conn):
|
|
self._check_transcript_write_guards(
|
|
conn, session_id, compression_lock_holder,
|
|
turn_lease_holder=turn_lease_holder, turn_lease_ttl_seconds=turn_lease_ttl_seconds,
|
|
)
|
|
from agent.transcript_repair import resolve_and_repair_transcript_batch
|
|
|
|
inserted_rows = resolve_and_repair_transcript_batch(
|
|
conn, session_id, messages,
|
|
encode_content_fn=self._encode_content, decode_content_fn=self._decode_content,
|
|
)
|
|
inserted = tool_calls_total = 0
|
|
if inserted_rows:
|
|
inserted, tool_calls_total = self._insert_message_rows(conn, session_id, inserted_rows)
|
|
if tool_calls_total > 0:
|
|
conn.execute(
|
|
"""UPDATE sessions SET message_count = message_count + ?,
|
|
tool_call_count = tool_call_count + ? WHERE id = ?""",
|
|
(inserted, tool_calls_total, session_id),
|
|
)
|
|
elif inserted > 0:
|
|
conn.execute(
|
|
"UPDATE sessions SET message_count = message_count + ? WHERE id = ?",
|
|
(inserted, session_id),
|
|
)
|
|
return inserted
|
|
|
|
return self._execute_write(_do, patience_s=self._TRANSCRIPT_WRITE_PATIENCE_S)
|
|
|
|
def set_latest_matching_message_display_kind(
|
|
self, session_id: str, *, role: str, content: str, display_kind: str,
|
|
display_metadata: Optional[Dict[str, Any]] = None,
|
|
) -> bool:
|
|
"""Stamp presentation metadata on this turn's freshly persisted row; the model
|
|
still receives ``role``/``content`` unchanged."""
|
|
from hermes_state import _scrub_surrogates
|
|
if not session_id or not content or not display_kind:
|
|
return False
|
|
|
|
def _do(conn):
|
|
row = conn.execute(
|
|
"SELECT id FROM messages WHERE session_id = ? AND role = ? "
|
|
"AND content = ? AND active = 1 ORDER BY id DESC LIMIT 1",
|
|
(session_id, role, self._encode_content(content)),
|
|
).fetchone()
|
|
if row is None:
|
|
return False
|
|
conn.execute(
|
|
"UPDATE messages SET display_kind = ?, display_metadata = ? WHERE id = ?",
|
|
(_scrub_surrogates(display_kind), self._encode_display_metadata(display_metadata), row[0]),
|
|
)
|
|
return True
|
|
|
|
return bool(self._execute_write(_do))
|
|
|
|
def _reaction_list(self, meta: Optional[Dict[str, Any]]) -> List[Dict[str, Any]]:
|
|
"""Well-formed (dict) reactions stored under ``REACTIONS_METADATA_KEY``."""
|
|
reactions = (meta or {}).get(self.REACTIONS_METADATA_KEY)
|
|
return [r for r in reactions if isinstance(r, dict)] if isinstance(reactions, list) else []
|
|
|
|
def set_message_reaction(
|
|
self, session_id: str, message_row_id: int, emoji: Optional[str], *, author: str = "user",
|
|
) -> Optional[List[Dict[str, Any]]]:
|
|
"""Set (or with ``emoji=None`` clear) *author*'s reaction on one message.
|
|
|
|
Tapback semantics: one reaction per author per message; the same emoji again
|
|
clears it, a different one replaces it. Returns the message's reaction list
|
|
after the write, or ``None`` when the row isn't part of *session_id*.
|
|
"""
|
|
from hermes_state import _scrub_surrogates
|
|
if not session_id or message_row_id is None:
|
|
return None
|
|
|
|
def _do(conn):
|
|
row = conn.execute(
|
|
"SELECT display_metadata FROM messages WHERE id = ? AND session_id = ?",
|
|
(message_row_id, session_id),
|
|
).fetchone()
|
|
if row is None:
|
|
return None
|
|
meta = self._decode_display_metadata(row[0]) or {}
|
|
existing = self._reaction_list(meta)
|
|
reactions = [r for r in existing if r.get("author") != author]
|
|
previous = next((r for r in existing if r.get("author") == author), None)
|
|
toggling_off = emoji is not None and previous is not None and previous.get("emoji") == emoji
|
|
if emoji and not toggling_off:
|
|
reactions.append({"emoji": _scrub_surrogates(emoji), "author": author, "at": time.time()})
|
|
if reactions:
|
|
meta[self.REACTIONS_METADATA_KEY] = reactions
|
|
else:
|
|
meta.pop(self.REACTIONS_METADATA_KEY, None)
|
|
conn.execute(
|
|
"UPDATE messages SET display_metadata = ? WHERE id = ?",
|
|
(self._encode_display_metadata(meta) if meta else None, message_row_id),
|
|
)
|
|
return reactions
|
|
|
|
return self._execute_write(_do)
|
|
|
|
def get_message_reactions(self, session_id: str, message_row_id: int) -> List[Dict[str, Any]]:
|
|
"""Reaction list persisted on one message row (never ``None``)."""
|
|
if not session_id or message_row_id is None:
|
|
return []
|
|
row = self._read_one(
|
|
"SELECT display_metadata FROM messages WHERE id = ? AND session_id = ?",
|
|
(message_row_id, session_id),
|
|
)
|
|
return self._reaction_list(self._decode_display_metadata(row[0])) if row is not None else []
|
|
|
|
def take_unseen_reactions(self, session_id: str, *, author: str = "user") -> List[Dict[str, Any]]:
|
|
"""Return *author*'s not-yet-surfaced reactions and mark them seen.
|
|
|
|
Reactions are announced on the NEXT user turn (never by rewriting the reacted
|
|
message — cache-safe); the ``seen`` stamp makes each announcement exactly once.
|
|
"""
|
|
if not session_id:
|
|
return []
|
|
|
|
def _do(conn):
|
|
rows = conn.execute(
|
|
"SELECT id, role, content, display_metadata FROM messages "
|
|
"WHERE session_id = ? AND active = 1 AND display_metadata IS NOT NULL ORDER BY id",
|
|
(session_id,),
|
|
).fetchall()
|
|
pending = []
|
|
for row in rows:
|
|
meta = self._decode_display_metadata(row["display_metadata"])
|
|
if not meta:
|
|
continue
|
|
reactions = meta.get(self.REACTIONS_METADATA_KEY)
|
|
if not isinstance(reactions, list):
|
|
continue
|
|
changed = False
|
|
for reaction in reactions:
|
|
if not isinstance(reaction, dict) or reaction.get("author") != author or reaction.get("seen"):
|
|
continue
|
|
reaction["seen"] = True
|
|
changed = True
|
|
content = self._decode_content(row["content"])
|
|
pending.append({
|
|
"row_id": row["id"],
|
|
"role": row["role"],
|
|
"emoji": reaction.get("emoji") or "",
|
|
"text": content if isinstance(content, str) else "",
|
|
})
|
|
if changed:
|
|
conn.execute(
|
|
"UPDATE messages SET display_metadata = ? WHERE id = ?",
|
|
(self._encode_display_metadata(meta), row["id"]),
|
|
)
|
|
return pending
|
|
|
|
return self._execute_write(_do) or []
|
|
|
|
def latest_message_row_id(
|
|
self, session_id: str, *, role: str = "user", offset: int = 0, require_text: bool = True
|
|
) -> Optional[int]:
|
|
"""Row id of the most recent active message with *role*, or ``None``.
|
|
|
|
``offset`` steps to earlier turns (1 = the one before the latest). ``require_text``
|
|
skips rows without plain-text content (tool-call-only turns, attachment stubs)
|
|
so "the latest message" never resolves to an invisible bubble.
|
|
"""
|
|
if not session_id or role not in {"user", "assistant"} or offset < 0:
|
|
return None
|
|
text_filter = "AND content IS NOT NULL AND TRIM(content) != '' " if require_text else ""
|
|
row = self._read_one(
|
|
"SELECT id FROM messages WHERE session_id = ? AND role = ? "
|
|
f"AND active = 1 {text_filter}ORDER BY id DESC LIMIT 1 OFFSET ?",
|
|
(session_id, role, int(offset)),
|
|
)
|
|
return row[0] if row else None
|
|
|
|
def latest_user_message_row_id(self, session_id: str) -> Optional[int]:
|
|
"""Row id of the most recent active user message ("the message that triggered
|
|
me"), or ``None``."""
|
|
return self.latest_message_row_id(session_id, role="user")
|
|
|
|
def get_message_role(self, session_id: str, row_id: int) -> Optional[str]:
|
|
"""Role of the active message at *row_id* in *session_id*, or ``None``."""
|
|
if not session_id:
|
|
return None
|
|
row = self._read_one(
|
|
"SELECT role FROM messages WHERE id = ? AND session_id = ? AND active = 1",
|
|
(int(row_id), session_id),
|
|
)
|
|
return row[0] if row else None
|
|
|
|
def _insert_message_rows(self, conn, session_id: str, messages: List[Dict[str, Any]]) -> tuple[int, int]:
|
|
"""Insert *messages* as fresh active rows inside the caller's write txn.
|
|
|
|
Returns ``(inserted_count, tool_call_count)``; does NOT touch sessions.*
|
|
counters (callers reconcile them differently). Reasoning columns are kept for
|
|
assistant rows only. Stamps ``msg["_row_id"]``.
|
|
"""
|
|
now_ts = time.time()
|
|
inserted = tool_calls_total = 0
|
|
for msg in messages:
|
|
role = msg.get("role", "unknown")
|
|
tool_calls = _parse_tool_calls(msg.get("tool_calls"))
|
|
message_timestamp = _coerce_timestamp(msg.get("timestamp"), now_ts)
|
|
cur = conn.execute(
|
|
_INSERT_MESSAGE_SQL,
|
|
self._message_row_params(
|
|
session_id, role, msg, tool_calls, message_timestamp, keep_reasoning=role == "assistant",
|
|
),
|
|
)
|
|
if isinstance(msg, dict) and cur.lastrowid is not None:
|
|
msg["_row_id"] = cur.lastrowid
|
|
inserted += 1
|
|
tool_calls_total += _tool_calls_count(tool_calls)
|
|
now_ts = max(now_ts + 1e-6, message_timestamp + 1e-6)
|
|
return inserted, tool_calls_total
|
|
|
|
def replace_messages(
|
|
self, session_id: str, messages: List[Dict[str, Any]], active_only: bool = False,
|
|
archive_dropped: bool = False, reject_active_turn_lease: bool = False,
|
|
) -> None:
|
|
"""Atomically replace the stored messages for a session (/retry, /undo, /compress).
|
|
|
|
DESTRUCTIVE by default: every row is DELETEd (and leaves the FTS index).
|
|
``active_only=True`` replaces only ``active = 1`` rows, leaving soft-archived
|
|
rows (compacted turns, rewind rows) untouched — required when sharing a session
|
|
id with an agent doing in-place compaction. ``archive_dropped=True`` SOFT-archives
|
|
the live rows (``active = 0, compacted = 0``, rewind-style) instead of deleting:
|
|
the mode rewind/edit/regenerate must use, since a DELETE leaves nothing to
|
|
recover from; it implies active-only handling. ``reject_active_turn_lease=True``
|
|
runs the lease check in the same write txn for user-initiated rewrites that do
|
|
not own the cross-process lease.
|
|
"""
|
|
from hermes_state import CompressionSessionClosedError
|
|
active_clause = " AND active = 1" if active_only else ""
|
|
|
|
def _do(conn):
|
|
if reject_active_turn_lease:
|
|
self._check_transcript_write_guards(
|
|
conn, session_id, None, reject_active_turn_lease=True, reject_active_compression_lock=True,
|
|
)
|
|
elif _ended_by_compression(conn.execute(_ENDED_BY_COMPRESSION_SQL, (session_id,)).fetchone()):
|
|
raise CompressionSessionClosedError(session_id)
|
|
if archive_dropped:
|
|
# Content-preserving UPDATE: FTS triggers don't fire on `active`, so the
|
|
# replaced turns stay searchable/readable with include_inactive=True.
|
|
conn.execute(
|
|
"UPDATE messages SET active = 0 WHERE session_id = ? AND active = 1", (session_id,),
|
|
)
|
|
else:
|
|
conn.execute(f"DELETE FROM messages WHERE session_id = ?{active_clause}", (session_id,))
|
|
conn.execute(
|
|
"UPDATE sessions SET message_count = 0, tool_call_count = 0 WHERE id = ?", (session_id,),
|
|
)
|
|
total_messages, total_tool_calls = self._insert_message_rows(conn, session_id, messages)
|
|
conn.execute(
|
|
"UPDATE sessions SET message_count = ?, tool_call_count = ? WHERE id = ?",
|
|
(total_messages, total_tool_calls, session_id),
|
|
)
|
|
|
|
self._execute_write(_do)
|
|
|
|
def has_archived_messages(self, session_id: str) -> bool:
|
|
"""True if the session has any soft-archived (``active = 0``) rows. Cheap probe;
|
|
production rewrite paths no longer branch on it (kept for tests/diagnostics)."""
|
|
return self._read_one(
|
|
"SELECT 1 FROM messages WHERE session_id = ? AND active = 0 LIMIT 1", (session_id,),
|
|
) is not None
|
|
|
|
def get_active_message_watermark(self, session_id: str) -> int:
|
|
"""MAX(id) of the session's active rows — captured at compression START; every
|
|
active row above it arrived concurrently and must survive compaction verbatim.
|
|
0 for an empty/unknown session."""
|
|
if not session_id:
|
|
return 0
|
|
row = self._read_one(
|
|
"SELECT COALESCE(MAX(id), 0) FROM messages WHERE session_id = ? AND active = 1", (session_id,),
|
|
)
|
|
return int(row[0]) if row else 0
|
|
|
|
def _tail_rows_after_watermark(self, conn, sql: str, params) -> Tuple[List[int], int]:
|
|
"""``(ids, tool_call_count)`` of the concurrent-tail rows selected by *sql*
|
|
(``SELECT id, tool_calls ...``)."""
|
|
rows = conn.execute(sql, params).fetchall()
|
|
return [int(r["id"]) for r in rows], sum(_tool_calls_len(r["tool_calls"]) for r in rows)
|
|
|
|
def _clone_message_rows(self, conn, tail_ids: List[int], *, session_id: Optional[str] = None) -> None:
|
|
"""Pure-SQL column clone of *tail_ids* as fresh live rows (new id, active=1,
|
|
compacted=0, everything else byte-exact; FTS triggers index the clones). With
|
|
*session_id* the clones land in that session instead of the originals'."""
|
|
skip = ("id", "active", "compacted") + (("session_id",) if session_id is not None else ())
|
|
col_list = ", ".join(c for c in self._message_column_names(conn) if c not in skip)
|
|
placeholders = _placeholders(tail_ids)
|
|
if session_id is None:
|
|
conn.execute(
|
|
f"INSERT INTO messages ({col_list}, active, compacted) "
|
|
f"SELECT {col_list}, 1, 0 FROM messages "
|
|
f"WHERE id IN ({placeholders}) ORDER BY id", tail_ids,
|
|
)
|
|
else:
|
|
conn.execute(
|
|
f"INSERT INTO messages ({col_list}, session_id, active, compacted) "
|
|
f"SELECT {col_list}, ?, 1, 0 FROM messages "
|
|
f"WHERE id IN ({placeholders}) ORDER BY id", [session_id, *tail_ids],
|
|
)
|
|
|
|
def archive_and_compact(
|
|
self, session_id: str, compacted_messages: List[Dict[str, Any]],
|
|
model_config_patch: Optional[Dict[str, Any]] = None, watermark: Optional[int] = None,
|
|
lock_holder: Optional[str] = None, tail_count: int = 0,
|
|
) -> int:
|
|
"""Non-destructive in-place compaction under ONE durable session id.
|
|
|
|
Soft-archives the active rows (``active=0, compacted=1`` — "summarized away",
|
|
still found by search_messages and readable with include_inactive) and inserts
|
|
*compacted_messages* as fresh active rows, atomically. Live-context loads
|
|
filter ``active = 1`` so the model reloads only the compacted set.
|
|
|
|
*watermark* (``get_active_message_watermark`` at compression START): rows with
|
|
``id > watermark`` arrived during the slow summary and are re-sequenced after
|
|
the compacted set by a pure-SQL column clone (fresh ids — consumers re-resolve
|
|
by content); ``None`` archives everything. *lock_holder*: the commit verifies
|
|
inside the txn that the compression lock is still held and unexpired, so a
|
|
reclaimed lease fails instead of clobbering the winner. *tail_count*: the LAST
|
|
N rows of *compacted_messages* are the verbatim carried-forward tail; their
|
|
originals (at/below the watermark) and the watermark clones' originals are
|
|
superseded duplicates and get rewind-style flags (``active=0, compacted=0``) so
|
|
session_search doesn't return each carried message once per compaction.
|
|
|
|
``message_count`` becomes the ACTIVE count; ``model_config_patch`` merges into
|
|
the session JSON in the same txn (``None`` value removes a key). Returns the new
|
|
active count.
|
|
"""
|
|
from hermes_state import SessionCompressionInProgressError
|
|
|
|
def _do(conn):
|
|
if lock_holder is not None:
|
|
lock_row = conn.execute(_COMPRESSION_LOCK_ROW_SQL, (session_id,)).fetchone()
|
|
if (
|
|
lock_row is None
|
|
or lock_row["holder"] != lock_holder
|
|
or float(lock_row["expires_at"]) <= time.time()
|
|
):
|
|
raise SessionCompressionInProgressError(
|
|
f"Compression lease for {session_id!r} lost before "
|
|
"commit; refusing to publish a stale compaction"
|
|
)
|
|
patched_model_config = None
|
|
if model_config_patch is not None:
|
|
# on_missing="raise": never commit against a vanished session row (the
|
|
# compressor's caller turns the error into a keep-the-original no-op).
|
|
patched_model_config = self._merge_model_config_json(
|
|
conn, session_id, model_config_patch, on_missing="raise"
|
|
)
|
|
tail_ids: list[int] = []
|
|
tail_tool_calls = 0
|
|
if watermark is not None:
|
|
tail_ids, tail_tool_calls = self._tail_rows_after_watermark(
|
|
conn, "SELECT id, tool_calls FROM messages "
|
|
"WHERE session_id = ? AND active = 1 AND id > ? ORDER BY id",
|
|
(session_id, int(watermark)),
|
|
)
|
|
# Rewind targets sit AT/BELOW the watermark (the compressor only saw rows up
|
|
# to it); without the bound a concurrent append would steal a LIMIT slot.
|
|
rewind_ids: list[int] = []
|
|
if tail_count > 0:
|
|
if watermark is not None:
|
|
tail_rows = conn.execute(
|
|
"SELECT id FROM messages WHERE session_id = ? AND active = 1 AND id <= ? "
|
|
"ORDER BY id DESC LIMIT ?", (session_id, int(watermark), int(tail_count)),
|
|
).fetchall()
|
|
else:
|
|
tail_rows = conn.execute(
|
|
"SELECT id FROM messages "
|
|
"WHERE session_id = ? AND active = 1 ORDER BY id DESC LIMIT ?",
|
|
(session_id, int(tail_count)),
|
|
).fetchall()
|
|
rewind_ids = [int(row["id"]) for row in tail_rows]
|
|
rewind_ids += tail_ids
|
|
if rewind_ids:
|
|
placeholders = _placeholders(rewind_ids)
|
|
conn.execute(
|
|
"UPDATE messages SET active = 0, compacted = 0 "
|
|
f"WHERE session_id = ? AND id IN ({placeholders})", [session_id, *rewind_ids],
|
|
)
|
|
conn.execute(
|
|
"UPDATE messages SET active = 0, compacted = 1 "
|
|
"WHERE session_id = ? AND active = 1 "
|
|
f"AND id NOT IN ({placeholders})", [session_id, *rewind_ids],
|
|
)
|
|
else:
|
|
conn.execute(
|
|
"UPDATE messages SET active = 0, compacted = 1 "
|
|
"WHERE session_id = ? AND active = 1", (session_id,),
|
|
)
|
|
inserted, tool_calls_total = self._insert_message_rows(conn, session_id, compacted_messages)
|
|
if tail_ids:
|
|
self._clone_message_rows(conn, tail_ids)
|
|
inserted += len(tail_ids)
|
|
tool_calls_total += tail_tool_calls
|
|
if model_config_patch is None:
|
|
conn.execute(
|
|
"UPDATE sessions SET message_count = ?, tool_call_count = ? WHERE id = ?",
|
|
(inserted, tool_calls_total, session_id),
|
|
)
|
|
else:
|
|
conn.execute(
|
|
"UPDATE sessions SET message_count = ?, tool_call_count = ?, "
|
|
"model_config = ? WHERE id = ?",
|
|
(inserted, tool_calls_total, patched_model_config, session_id),
|
|
)
|
|
return inserted
|
|
|
|
return self._execute_write(_do)
|
|
|
|
def _message_column_names(self, conn) -> List[str]:
|
|
"""Column names of the messages table, cached per-connection era."""
|
|
cached = getattr(self, "_message_columns_cache", None)
|
|
if cached:
|
|
return cached
|
|
cols = [r[1] for r in conn.execute("PRAGMA table_info(messages)").fetchall()]
|
|
self._message_columns_cache = cols
|
|
return cols
|
|
|
|
def set_latest_user_api_content(self, session_id: str, content: Any, api_content: str) -> int:
|
|
"""Backfill the ``api_content`` sidecar onto the newest ACTIVE user row.
|
|
|
|
In-place preflight compaction inserts the current user row BEFORE the turn
|
|
prologue composes the sidecar, and the later persist identity-skips compacted
|
|
dicts; without this the reload would replay clean content and reopen the
|
|
prompt-cache divergence. The ``content`` match guards against a racing rewrite.
|
|
Returns rows updated (0 or 1).
|
|
"""
|
|
from hermes_state import _scrub_surrogates
|
|
return self._write_rowcount(
|
|
"UPDATE messages SET api_content = ? WHERE id = (SELECT id FROM messages "
|
|
"WHERE session_id = ? AND role = 'user' AND active = 1 ORDER BY id DESC LIMIT 1"
|
|
") AND content IS ?",
|
|
(_scrub_surrogates(api_content), session_id, self._encode_content(content)),
|
|
)
|
|
|
|
def _dedupe_display_generations(self, rows):
|
|
"""Collapse compaction generations so each logical message appears once.
|
|
|
|
Compaction copies the protected tail into each generation (same
|
|
role/content/timestamp, different ``active``/id); prefer the live row, then the
|
|
newest generation. The ONE definition shared by every display projection
|
|
(get_messages, get_resume_conversations, get_ancestor_display_prefix,
|
|
get_messages_as_conversation). *rows* must be ordered by ``id``; order is kept.
|
|
"""
|
|
seen: Dict[Tuple[Any, ...], Any] = {}
|
|
for row in rows:
|
|
dedupe_content = row["content"]
|
|
if row["role"] == "user":
|
|
from agent.context_compressor import split_user_originated_turn
|
|
|
|
handoff, live_view = split_user_originated_turn({
|
|
"role": "user",
|
|
"content": self._decode_content(row["content"]),
|
|
"display_kind": row["display_kind"],
|
|
"display_metadata": self._decode_display_metadata(row["display_metadata"]),
|
|
})
|
|
if handoff is not None and live_view is not None:
|
|
dedupe_content = self._encode_content(live_view.get("content"))
|
|
# Tool fields are part of the key: identical tool messages across generations
|
|
# collapse, distinct tool calls sharing role/content/timestamp never merge.
|
|
key = (
|
|
row["role"], dedupe_content, row["timestamp"],
|
|
row["tool_call_id"], row["tool_calls"], row["tool_name"],
|
|
)
|
|
cur = seen.get(key)
|
|
if cur is None or (row["active"], row["id"]) > (cur["active"], cur["id"]):
|
|
seen[key] = row
|
|
return sorted(seen.values(), key=lambda r: r["id"])
|
|
|
|
def _row_to_message_dict(self, row, *, warn_context: str, summary_flag: bool) -> Dict[str, Any]:
|
|
"""``dict(row)`` with content/tool_calls/display_metadata decoded. *summary_flag*
|
|
pops ``_compressed_summary`` and keeps it only as ``True``."""
|
|
msg = dict(row)
|
|
if summary_flag and msg.pop("_compressed_summary", 0):
|
|
msg["_compressed_summary"] = True
|
|
if "content" in msg:
|
|
msg["content"] = self._decode_content(msg["content"])
|
|
if msg.get("tool_calls"):
|
|
msg["tool_calls"] = _json_or(
|
|
msg["tool_calls"], [],
|
|
f"Failed to deserialize tool_calls in {warn_context}, falling back to []",
|
|
)
|
|
if msg.get("display_metadata") is not None:
|
|
msg["display_metadata"] = self._decode_display_metadata(msg["display_metadata"])
|
|
return msg
|
|
|
|
@staticmethod
|
|
def _active_clause(include_inactive: bool, include_compacted: bool) -> str:
|
|
"""Audit reads: every row; display reads: active plus compaction-archived
|
|
(never Undo/Rewind rows); default: live only."""
|
|
if include_inactive:
|
|
return ""
|
|
return _DISPLAY_ACTIVE_CLAUSE if include_compacted else " AND active = 1"
|
|
|
|
def get_messages(
|
|
self, session_id: str, include_inactive: bool = False, include_compacted: bool = False,
|
|
limit: Optional[int] = None, offset: int = 0, latest: bool = False,
|
|
after_id: Optional[int] = None,
|
|
) -> List[Dict[str, Any]]:
|
|
"""Load messages for a session in insertion order (AUTOINCREMENT id, never
|
|
timestamp — clocks regress on WSL2/NTP steps).
|
|
|
|
``include_inactive`` loads soft-deleted rewind rows; ``include_compacted`` adds
|
|
rows preserved by in-place compaction (durable display history a transcript
|
|
read must not drop) but not rewind rows. ``limit``/``offset`` page; ``latest``
|
|
measures the offset back from the newest row and still returns chronological
|
|
order. ``after_id`` is keyset paging (``id > after_id``), ascending only.
|
|
"""
|
|
if after_id is not None and (latest or offset):
|
|
raise ValueError("after_id is incompatible with latest/offset paging")
|
|
if after_id is not None and include_compacted:
|
|
raise ValueError("after_id is incompatible with include_compacted (deduped display reads use offset paging)")
|
|
active_clause = self._active_clause(include_inactive, include_compacted)
|
|
if include_compacted:
|
|
# Read the full display set (the UI-level row cap lives in the endpoint),
|
|
# dedupe generations, then page.
|
|
rows = self._dedupe_display_generations(self._read_all(
|
|
"SELECT * FROM messages WHERE session_id = ?" + active_clause + " ORDER BY id ASC",
|
|
[session_id],
|
|
))
|
|
if latest:
|
|
rows = rows[::-1]
|
|
rows = rows[offset:]
|
|
if limit is not None:
|
|
rows = rows[:limit]
|
|
if latest:
|
|
rows = rows[::-1]
|
|
else:
|
|
keyset_clause = " AND id > ?" if after_id is not None else ""
|
|
sql = (
|
|
"SELECT * FROM messages WHERE session_id = ?"
|
|
f"{active_clause}{keyset_clause} ORDER BY id {'DESC' if latest else 'ASC'}"
|
|
)
|
|
params: list = [session_id]
|
|
if after_id is not None:
|
|
params.append(after_id)
|
|
if limit is not None or offset:
|
|
# SQLite's OFFSET requires LIMIT; -1 means "no limit".
|
|
sql += " LIMIT ? OFFSET ?"
|
|
params.extend([-1 if limit is None else limit, offset])
|
|
rows = self._read_all(sql, params)
|
|
if latest:
|
|
rows.reverse()
|
|
return [self._row_to_message_dict(row, warn_context="get_messages", summary_flag=True) for row in rows]
|
|
|
|
def find_pr_url_messages(self, session_ids: List[str]) -> List[Dict[str, Any]]:
|
|
"""Tool results in these sessions containing ``/pull/`` — a deliberately loose
|
|
candidate scan, oldest-first per session so the caller can take the last match."""
|
|
found: List[Dict[str, Any]] = []
|
|
ids = [s for s in session_ids if s]
|
|
for start in range(0, len(ids), 900): # SQLite's bound-variable ceiling.
|
|
chunk = ids[start : start + 900]
|
|
rows = self._read_all(
|
|
f"""SELECT session_id, content FROM messages
|
|
WHERE session_id IN ({_placeholders(chunk)})
|
|
AND role = 'tool' AND content LIKE '%/pull/%'
|
|
ORDER BY id ASC""",
|
|
chunk,
|
|
)
|
|
found.extend({"session_id": row[0], "content": row[1]} for row in rows)
|
|
return found
|
|
|
|
def get_messages_around(self, session_id: str, around_message_id: int, window: int = 5) -> Dict[str, Any]:
|
|
"""Window of up to *window* messages either side of an anchor id (id ascending).
|
|
|
|
``messages_before``/``messages_after`` count rows in the returned slice strictly
|
|
before/after the anchor; less than *window* means a session boundary. Empty
|
|
window when the anchor is not a row of *session_id*.
|
|
"""
|
|
window = max(window, 0)
|
|
with self._read_ctx() as conn:
|
|
anchor_exists = conn.execute(
|
|
"SELECT 1 FROM messages WHERE id = ? AND session_id = ? LIMIT 1",
|
|
(around_message_id, session_id),
|
|
).fetchone()
|
|
if not anchor_exists:
|
|
return {"window": [], "messages_before": 0, "messages_after": 0}
|
|
before_rows = conn.execute(
|
|
"SELECT * FROM messages WHERE session_id = ? AND id <= ? ORDER BY id DESC LIMIT ?",
|
|
(session_id, around_message_id, window + 1),
|
|
).fetchall()
|
|
after_rows = conn.execute(
|
|
"SELECT * FROM messages WHERE session_id = ? AND id > ? ORDER BY id ASC LIMIT ?",
|
|
(session_id, around_message_id, window),
|
|
).fetchall()
|
|
rows = list(reversed(before_rows)) + list(after_rows)
|
|
window_msgs = [self._row_to_message_dict(r, warn_context="get_messages_around", summary_flag=False) for r in rows]
|
|
# before_rows includes the anchor itself.
|
|
return {"window": window_msgs, "messages_before": max(0, len(before_rows) - 1), "messages_after": len(after_rows)}
|
|
|
|
def resolve_resume_session_id(self, session_id: str) -> str:
|
|
"""Redirect a resume target to the descendant session that holds the messages.
|
|
|
|
Follows the compression chain to the live tip first (``get_compression_tip`` is
|
|
lineage-aware: only children of compression-ended parents, so delegation/branch
|
|
children never hijack the resume), then walks ``parent_session_id`` forward,
|
|
returning the deepest node with messages — never short-circuiting on the start
|
|
node, since a continuation may hold the newer turns. Branch, delegate, reset
|
|
and tool children are skipped (they carry ``parent_session_id`` too). Returns
|
|
*session_id* unchanged when nothing has messages. Depth cap 32.
|
|
"""
|
|
if not session_id:
|
|
return session_id
|
|
try:
|
|
tip = self.get_compression_tip(session_id)
|
|
except Exception:
|
|
tip = session_id
|
|
if tip and tip != session_id:
|
|
session_id = tip
|
|
with self._read_ctx() as conn:
|
|
current = session_id
|
|
seen = {current}
|
|
best = None # deepest node with messages
|
|
for _ in range(32):
|
|
try:
|
|
if conn.execute(
|
|
"SELECT 1 FROM messages WHERE session_id = ? LIMIT 1", (current,),
|
|
).fetchone() is not None:
|
|
best = current
|
|
child_row = conn.execute(
|
|
"SELECT id FROM sessions AS child WHERE child.parent_session_id = ? "
|
|
" AND json_extract(COALESCE(child.model_config, '{}'), '$._branched_from') IS NULL "
|
|
" AND json_extract(COALESCE(child.model_config, '{}'), '$._delegate_from') IS NULL "
|
|
" AND json_extract(COALESCE(child.model_config, '{}'), '$._reset_from') IS NULL "
|
|
f" AND NOT {_legacy_reset_child_sql('child', _RESET_END_REASONS_SQL)} "
|
|
" AND COALESCE(child.source, '') != 'tool' "
|
|
"ORDER BY child.started_at DESC, child.id DESC LIMIT 1", (current,),
|
|
).fetchone()
|
|
except Exception:
|
|
return session_id
|
|
if child_row is None:
|
|
break
|
|
child_id = child_row["id"] if hasattr(child_row, "keys") else child_row[0]
|
|
if not child_id or child_id in seen:
|
|
break
|
|
seen.add(child_id)
|
|
current = child_id
|
|
return best if best is not None else session_id
|
|
|
|
def _fetch_conversation_rows(self, session_ids: List[str], active_clause: str, *, with_session_id: bool):
|
|
"""``_CONVERSATION_ROW_COLUMNS`` rows for *session_ids*, ORDER BY id (insertion
|
|
order — timestamps are not monotonic and would break tool-call adjacency)."""
|
|
prefix = "SELECT session_id, " if with_session_id else "SELECT "
|
|
with self._read_ctx() as conn:
|
|
return conn.execute(
|
|
f"{prefix}{self._CONVERSATION_ROW_COLUMNS} "
|
|
f"FROM messages WHERE session_id IN ({_placeholders(session_ids)})"
|
|
f"{active_clause} ORDER BY id", tuple(session_ids),
|
|
).fetchall()
|
|
|
|
def get_messages_as_conversation(
|
|
self, session_id: str, include_ancestors: bool = False, include_inactive: bool = False,
|
|
repair_alternation: bool = False, include_row_ids: bool = False,
|
|
include_compacted: bool = False,
|
|
) -> List[Dict[str, Any]]:
|
|
"""Load messages in OpenAI conversation format (gateway history restore).
|
|
|
|
``include_compacted`` adds compaction-archived rows deduped by
|
|
:meth:`_dedupe_display_generations` — DISPLAY reads only; the model-fed restore
|
|
must not pass it or a resume regrows the history compaction summarized away.
|
|
``repair_alternation`` runs ``repair_message_sequence`` on the loaded list
|
|
(LIVE REPLAY callers) so a durable ``user;user`` pair doesn't re-trigger the
|
|
per-request repair forever; the stored transcript is never mutated.
|
|
"""
|
|
session_ids = [session_id]
|
|
if include_ancestors and not self._is_explicit_branch_session(session_id):
|
|
session_ids = self._session_lineage_root_to_tip(session_id)
|
|
rows = self._fetch_conversation_rows(
|
|
session_ids, self._active_clause(include_inactive, include_compacted), with_session_id=False,
|
|
)
|
|
if include_compacted:
|
|
rows = self._dedupe_display_generations(rows)
|
|
return self._rows_to_conversation(
|
|
rows, session_id=session_id, include_ancestors=include_ancestors,
|
|
repair_alternation=repair_alternation, include_row_ids=include_row_ids,
|
|
)
|
|
|
|
def _dedupe_replayed_user(self, messages, msg, exact_user_clones) -> Tuple[bool, Any]:
|
|
"""Ancestor-lineage dedupe for one decoded user *msg*.
|
|
|
|
Returns ``(skip, exact_clone_key)``. Watermark rotation column-clones the
|
|
concurrent tail into the child after the summary, so the copies need not be
|
|
adjacent: an exact ``(timestamp, canonical content)`` clone index is checked
|
|
first, then the adjacent-duplicate heuristic. A rotated child carrier wins over
|
|
the simpler ancestor copy (it owns the durable row id and the summary scaffold).
|
|
"""
|
|
canonical_content, _is_composite = self._canonical_replayed_user_content(msg)
|
|
exact_clone_key = self._exact_replayed_user_clone_key(msg.get("timestamp"), canonical_content)
|
|
previous_exact = exact_user_clones.get(exact_clone_key) if exact_clone_key is not None else None
|
|
duplicate = None
|
|
if previous_exact is not None:
|
|
previous_index = next(
|
|
(index for index, candidate in enumerate(messages) if candidate is previous_exact), None,
|
|
)
|
|
if previous_index is not None:
|
|
duplicate = (previous_index, True)
|
|
if duplicate is None:
|
|
duplicate = self._find_duplicate_replayed_user_message(messages, msg)
|
|
if duplicate is not None:
|
|
duplicate_index, prefer_current = duplicate
|
|
if not prefer_current:
|
|
return True, exact_clone_key
|
|
messages.pop(duplicate_index)
|
|
return False, exact_clone_key
|
|
|
|
def _rows_to_conversation(
|
|
self, rows, *, session_id: str, include_ancestors: bool, repair_alternation: bool,
|
|
include_row_ids: bool = False, include_summary_markers: bool = False,
|
|
) -> List[Dict[str, Any]]:
|
|
"""Decode fetched message rows (ordered by id, pre-filtered) into OpenAI format.
|
|
|
|
Every dict is stamped ``_DB_PERSISTED_MARKER_KEY`` at the source (born durable)
|
|
so an identity-losing handoff never re-appends the whole transcript on flush.
|
|
``_row_id`` is opt-in (gateway reactions). Reasoning fields are restored on
|
|
assistant rows only. Key order of each dict is stable (see the column tables).
|
|
"""
|
|
from hermes_state import _strip_background_review_harness, _strip_stale_tool_call_markers
|
|
messages = []
|
|
exact_user_clones: Dict[Tuple[Any, str], Dict[str, Any]] = {}
|
|
for row in rows:
|
|
content = self._decode_content(row["content"])
|
|
if row["role"] in {"user", "assistant"} and isinstance(content, str):
|
|
content = sanitize_context(content).strip()
|
|
msg = {"role": row["role"], "content": content, _DB_PERSISTED_MARKER_KEY: True}
|
|
if include_row_ids and row["id"] is not None:
|
|
msg["_row_id"] = row["id"]
|
|
msg.update((col, row[col]) for col in _VERBATIM_COLS if row[col])
|
|
if row["display_metadata"]:
|
|
decoded = self._decode_display_metadata(row["display_metadata"])
|
|
if decoded is not None:
|
|
msg["display_metadata"] = decoded
|
|
if include_summary_markers and row["_compressed_summary"]:
|
|
msg["_compressed_summary"] = True
|
|
msg.update((col, row[col]) for col in _META_COLS if row[col])
|
|
if row["tool_calls"]:
|
|
msg["tool_calls"] = _json_or(
|
|
row["tool_calls"], [],
|
|
"Failed to deserialize tool_calls in conversation replay, falling back to []",
|
|
)
|
|
# Platform-side id exposed as ``message_id`` (JSONL transcript compat).
|
|
if row["platform_message_id"]:
|
|
msg["message_id"] = row["platform_message_id"]
|
|
if row["observed"]:
|
|
msg["observed"] = True
|
|
if row["role"] == "assistant":
|
|
msg.update((col, row[col]) for col in ("finish_reason", "reasoning") if row[col])
|
|
if row["reasoning_content"] is not None:
|
|
msg["reasoning_content"] = row["reasoning_content"]
|
|
msg.update(
|
|
(col, _json_or(row[col], None, f"Failed to deserialize {col}, falling back to None"))
|
|
for col in _ASSISTANT_JSON_COLS if row[col]
|
|
)
|
|
exact_clone_key = None
|
|
if include_ancestors:
|
|
skip, exact_clone_key = self._dedupe_replayed_user(messages, msg, exact_user_clones)
|
|
if skip:
|
|
continue
|
|
messages.append(msg)
|
|
if include_ancestors and exact_clone_key is not None:
|
|
exact_user_clones[exact_clone_key] = msg
|
|
# Defense-in-depth: strip a background-review harness turn (older builds shared
|
|
# the parent's session_id) plus its curator reply, and bare tool-call marker
|
|
# content ("[memory]") persisted as an answer before the loop fix.
|
|
messages = _strip_background_review_harness(messages)
|
|
messages = _strip_stale_tool_call_markers(messages)
|
|
if repair_alternation and messages:
|
|
from agent.agent_runtime_helpers import repair_message_sequence
|
|
|
|
repaired = repair_message_sequence(None, messages)
|
|
if repaired:
|
|
logger.info(
|
|
"Repaired %d message-alternation violation(s) while "
|
|
"restoring session %s — durable transcript kept them, "
|
|
"see repair_message_sequence", repaired, session_id,
|
|
)
|
|
return messages
|
|
|
|
def get_resume_conversations(self, session_id: str) -> Tuple[List[Dict[str, Any]], List[Dict[str, Any]]]:
|
|
"""``(model_history, display_history)`` for a session resume from ONE SELECT.
|
|
|
|
``model_history``: the tip's active rows, alternation-repaired, with the summary
|
|
marker kept for pre-compress checkpointing. ``display_history``: the full
|
|
compression lineage (``/branch`` sessions are their own lineage) verbatim, with
|
|
compaction-archived rows included and deduped, plus replayed-user dedup. Byte-
|
|
identical to the separate reads (test_get_resume_conversations_matches_separate_reads).
|
|
"""
|
|
session_ids = self._resume_lineage_ids(session_id)
|
|
rows = self._fetch_conversation_rows(session_ids, _DISPLAY_ACTIVE_CLAUSE, with_session_id=True)
|
|
# The model projection stays active-only: it is the compressed working context.
|
|
tip_rows = [r for r in rows if r["session_id"] == session_id and r["active"]]
|
|
model_history = self._rows_to_conversation(
|
|
tip_rows, session_id=session_id, include_ancestors=False, repair_alternation=True,
|
|
include_row_ids=True, include_summary_markers=True,
|
|
)
|
|
display_history = self._rows_to_conversation(
|
|
self._dedupe_display_generations(rows), session_id=session_id,
|
|
include_ancestors=True, repair_alternation=False, include_row_ids=True,
|
|
)
|
|
return model_history, display_history
|
|
|
|
def _resume_lineage_ids(self, session_id: str) -> List[str]:
|
|
"""Session ids a full (display) resume materializes: the compression lineage,
|
|
or the session alone for an explicit ``/branch`` copy. Shared by the resume
|
|
readers and the resume guard so the guard counts exactly what a resume loads."""
|
|
if self._is_explicit_branch_session(session_id):
|
|
return [session_id]
|
|
return self._session_lineage_root_to_tip(session_id)
|
|
|
|
def _resume_count_scope(self, session_id: str, tip_only: bool) -> Tuple[List[str], str]:
|
|
"""``tip_only``: the tip's ACTIVE rows (model restore); else the full-lineage
|
|
DISPLAY set (active + compaction-archived) that get_resume_conversations loads."""
|
|
if tip_only:
|
|
return [session_id], "active = 1"
|
|
return self._resume_lineage_ids(session_id), "(active = 1 OR compacted = 1)"
|
|
|
|
def get_resume_message_count(self, session_id: str, *, tip_only: bool = False) -> int:
|
|
"""Count the rows a resume would materialize (see ``_resume_count_scope``)."""
|
|
session_ids, active_clause = self._resume_count_scope(session_id, tip_only)
|
|
row = self._read_one(
|
|
f"SELECT COUNT(*) FROM messages WHERE session_id IN ({_placeholders(session_ids)}) AND {active_clause}",
|
|
tuple(session_ids),
|
|
)
|
|
return int(row[0] if row else 0)
|
|
|
|
def assert_resume_safe(self, session_id: str, max_messages: Optional[int] = None, *, tip_only: bool = False) -> int:
|
|
"""Return the resume row count or raise ``SessionResumeTooLargeError``.
|
|
|
|
``max_messages=None`` reads ``sessions.max_resume_messages``; 0 disables the
|
|
guard and returns 0 without counting. ``tip_only`` bounds only the tip's active
|
|
rows, for callers that never materialize the lineage in memory — a heavily
|
|
compressed conversation (~29k lineage rows behind a ~700-row tip) is exactly
|
|
what compression should produce and must not be rejected.
|
|
"""
|
|
from hermes_state import SessionResumeTooLargeError, resolved_max_resume_messages
|
|
if max_messages is None:
|
|
max_messages = resolved_max_resume_messages()
|
|
if max_messages < 0:
|
|
raise ValueError("max_messages must be non-negative")
|
|
if max_messages == 0:
|
|
return 0
|
|
session_ids, active_clause = self._resume_count_scope(session_id, tip_only)
|
|
row = self._read_one(
|
|
"SELECT COUNT(*) FROM ("
|
|
f"SELECT 1 FROM messages WHERE session_id IN ({_placeholders(session_ids)}) "
|
|
f"AND {active_clause} LIMIT ?"
|
|
")", (*session_ids, max_messages + 1),
|
|
)
|
|
message_count = int(row[0] if row else 0)
|
|
if message_count > max_messages:
|
|
raise SessionResumeTooLargeError(
|
|
message_count, max_messages,
|
|
scope="in its tip segment" if tip_only else "across its lineage",
|
|
)
|
|
return message_count
|
|
|
|
def get_ancestor_display_prefix(self, session_id: str) -> List[Dict[str, Any]]:
|
|
"""Ancestor-only display messages of a lineage (rows with ``session_id !=`` tip).
|
|
|
|
``session.resume`` prepends this to the live model history. Identifying
|
|
ancestors by row origin (not ``display[:len(display) - len(model)]``) avoids
|
|
overcounting when alternation repair removes tip messages from the middle.
|
|
"""
|
|
session_ids = self._resume_lineage_ids(session_id)
|
|
if len(session_ids) <= 1:
|
|
return []
|
|
rows = self._dedupe_display_generations(
|
|
self._fetch_conversation_rows(session_ids, _DISPLAY_ACTIVE_CLAUSE, with_session_id=True)
|
|
)
|
|
ancestor_ids = {int(row["id"]) for row in rows if row["session_id"] != session_id and row["id"] is not None}
|
|
if not ancestor_ids:
|
|
return []
|
|
lineage = self._rows_to_conversation(
|
|
rows, session_id=session_id, include_ancestors=True, repair_alternation=False, include_row_ids=True,
|
|
)
|
|
return [
|
|
{k: v for k, v in message.items() if k != "_row_id"}
|
|
for message in lineage if message.get("_row_id") in ancestor_ids
|
|
]
|
|
|
|
def get_conversation_root(self, session_id: str) -> str:
|
|
"""ROOT id of *session_id*'s lineage — the stable conversation id across
|
|
compression segments and delegate subagents (Nous Portal usage tagging).
|
|
Unchanged when there is no recorded parent."""
|
|
chain = self._session_lineage_root_to_tip(session_id)
|
|
return chain[0] if chain and chain[0] else session_id
|
|
|
|
@staticmethod
|
|
def _canonical_replayed_user_content(msg: Dict[str, Any]) -> Tuple[Any, bool]:
|
|
"""Return canonical live content and whether *msg* is composite."""
|
|
if msg.get("role") != "user":
|
|
return None, False
|
|
from agent.context_compressor import split_user_originated_turn
|
|
|
|
handoff, live_view = split_user_originated_turn(msg)
|
|
is_composite = handoff is not None and live_view is not None
|
|
return (
|
|
live_view.get("content") if is_composite and live_view is not None else msg.get("content"),
|
|
is_composite,
|
|
)
|
|
|
|
@staticmethod
|
|
def _exact_replayed_user_clone_key(timestamp: Any, content: Any) -> Optional[Tuple[Any, str]]:
|
|
"""Return a hashable key for a column-exact rotation clone."""
|
|
if timestamp is None or content in (None, "", []):
|
|
return None
|
|
try:
|
|
encoded = json.dumps(content, ensure_ascii=False, sort_keys=True, separators=(",", ":"))
|
|
except (TypeError, ValueError):
|
|
return None
|
|
return timestamp, encoded
|
|
|
|
@staticmethod
|
|
def _find_duplicate_replayed_user_message(
|
|
messages: List[Dict[str, Any]], msg: Dict[str, Any]
|
|
) -> Optional[Tuple[int, bool]]:
|
|
"""Return an adjacent replay duplicate and whether *msg* must win.
|
|
|
|
Rotation may persist the current ask once in the parent and again inside a
|
|
composite child carrier; compare the canonical live payload for carriers while
|
|
keeping the exact-string dedupe for ordinary replayed users. The child carrier
|
|
wins (it owns the durable row id and the retained scaffold).
|
|
"""
|
|
from hermes_state import SessionDB
|
|
if msg.get("role") != "user":
|
|
return None
|
|
content, prefer_current = SessionDB._canonical_replayed_user_content(msg)
|
|
if content in (None, "", []):
|
|
return None
|
|
for index in range(len(messages) - 1, -1, -1):
|
|
prev = messages[index]
|
|
if prev.get("role") == "user":
|
|
prev_content, prev_is_composite = SessionDB._canonical_replayed_user_content(prev)
|
|
if prev_content == content and (prefer_current or prev_is_composite or isinstance(content, str)):
|
|
return index, prefer_current
|
|
if prev.get("role") == "assistant" and (prev.get("content") or prev.get("tool_calls")):
|
|
return None
|
|
return None
|
|
|
|
def get_active_message_ids(self, session_id: str) -> List[int]:
|
|
"""Ordered physical active ids pinned by rewind CAS checks (includes legacy
|
|
harness rows that conversation projections omit)."""
|
|
rows = self._read_all(
|
|
"SELECT id FROM messages WHERE session_id = ? AND active = 1 ORDER BY id", (session_id,),
|
|
)
|
|
return [int(row[0]) for row in rows]
|
|
|
|
@staticmethod
|
|
def _active_transcript_counts(conn, session_id: str) -> tuple[int, int]:
|
|
"""Return active message/tool-call counts inside the caller's txn."""
|
|
rows = conn.execute(
|
|
"SELECT tool_calls FROM messages WHERE session_id = ? AND active = 1", (session_id,),
|
|
).fetchall()
|
|
return len(rows), sum(_tool_calls_len(row[0], scalar=1) for row in rows)
|
|
|
|
def _split_rewind_target(self, target_row: Dict[str, Any], expected_target_content: Any, preserve_compaction_handoff: bool):
|
|
"""Validate an active rewind target and return its handoff scaffold (or None).
|
|
|
|
Raises ``ValueError`` for an inactive / non-user-originated target or a missing
|
|
composite carrier, ``RuntimeError`` when the canonical live payload no longer
|
|
matches *expected_target_content*.
|
|
"""
|
|
if not target_row.get("active"):
|
|
raise ValueError("rewind target is not active")
|
|
from agent.context_compressor import split_user_originated_turn
|
|
|
|
split_target = target_row.copy()
|
|
split_target["content"] = self._decode_content(split_target.get("content"))
|
|
split_target["display_metadata"] = self._decode_display_metadata(split_target.get("display_metadata"))
|
|
handoff, live_view = split_user_originated_turn(split_target)
|
|
if live_view is None:
|
|
raise ValueError("rewind target is not a user-originated turn")
|
|
live_content = live_view.get("content")
|
|
if isinstance(live_content, str):
|
|
live_content = sanitize_context(live_content).strip()
|
|
if expected_target_content is not None and live_content != expected_target_content:
|
|
raise RuntimeError("rewind target changed before it could be persisted")
|
|
if preserve_compaction_handoff and handoff is None:
|
|
raise ValueError("preserve_compaction_handoff requires an active composite carrier")
|
|
return handoff if preserve_compaction_handoff else None
|
|
|
|
def rewind_to_message(
|
|
self, session_id: str, target_message_id: int, *, preserve_compaction_handoff: bool = False,
|
|
expected_active_ids: Optional[List[int]] = None, expected_target_content: Any = None,
|
|
) -> Dict[str, Any]:
|
|
"""Soft-delete (``active=0``) every message with id >= *target_message_id*.
|
|
|
|
The target itself goes inactive so the caller can pre-fill it as the next
|
|
prompt. Returns ``{"rewound_count", "target_message", "new_head_id"}`` (plus
|
|
``replacement_message_id`` with ``preserve_compaction_handoff``, which archives a
|
|
composite summary carrier and inserts its hidden handoff scaffold as the new
|
|
head in the same txn). Raises ``ValueError`` when the target is missing or not a
|
|
``user`` row.
|
|
|
|
``expected_active_ids`` / ``expected_target_content`` pin the active row set and
|
|
the canonical live payload inside the txn before any mutation (presentation-only
|
|
metadata changes do not invalidate a rewind). A live cross-process turn lease
|
|
refuses the rewind; expired/dead holders are reclaimed. ``rewind_count`` always
|
|
increments, even when the target was already inactive.
|
|
"""
|
|
|
|
def _do(conn):
|
|
self._check_transcript_write_guards(
|
|
conn, session_id, None, reject_active_turn_lease=True, reject_active_compression_lock=True,
|
|
)
|
|
if expected_active_ids is not None:
|
|
active_rows = conn.execute(
|
|
"SELECT id FROM messages WHERE session_id = ? AND active = 1 ORDER BY id", (session_id,),
|
|
).fetchall()
|
|
if [int(r[0]) for r in active_rows] != expected_active_ids:
|
|
raise RuntimeError("active transcript changed before the rewind could be persisted")
|
|
row = conn.execute(
|
|
"SELECT * FROM messages WHERE id = ? AND session_id = ?", (target_message_id, session_id),
|
|
).fetchone()
|
|
if row is None:
|
|
raise ValueError(f"message {target_message_id} not found in session {session_id}")
|
|
target_row = dict(row)
|
|
if target_row.get("role") != "user":
|
|
raise ValueError(
|
|
f"rewind target must be a 'user' message (got role="
|
|
f"{target_row.get('role')!r}, id={target_message_id})"
|
|
)
|
|
replacement_message_id = replacement = None
|
|
if preserve_compaction_handoff or expected_target_content is not None:
|
|
replacement = self._split_rewind_target(target_row, expected_target_content, preserve_compaction_handoff)
|
|
ids = [r[0] for r in conn.execute(
|
|
"SELECT id FROM messages WHERE session_id = ? AND id >= ? AND active = 1",
|
|
(session_id, target_message_id),
|
|
).fetchall()]
|
|
if ids:
|
|
conn.execute(f"UPDATE messages SET active = 0 WHERE id IN ({_placeholders(ids)})", ids)
|
|
if replacement is not None:
|
|
self._insert_message_rows(conn, session_id, [replacement])
|
|
replacement_message_id = int(conn.execute("SELECT last_insert_rowid()").fetchone()[0])
|
|
conn.execute(
|
|
"UPDATE sessions SET rewind_count = COALESCE(rewind_count, 0) + 1 WHERE id = ?", (session_id,),
|
|
)
|
|
message_count, tool_call_count = self._active_transcript_counts(conn, session_id)
|
|
conn.execute(
|
|
"UPDATE sessions SET message_count = ?, tool_call_count = ? WHERE id = ?",
|
|
(message_count, tool_call_count, session_id),
|
|
)
|
|
head_row = conn.execute(
|
|
"SELECT MAX(id) FROM messages WHERE session_id = ? AND active = 1", (session_id,),
|
|
).fetchone()
|
|
return target_row, ids, head_row[0] if head_row else None, replacement_message_id
|
|
|
|
target_row, rewound, new_head_id, replacement_message_id = self._execute_write(_do)
|
|
# Decode for the prompt-buffer prefill without a second fallible DB operation.
|
|
target_row["content"] = self._decode_content(target_row.get("content"))
|
|
result = {"rewound_count": len(rewound), "target_message": target_row, "new_head_id": new_head_id}
|
|
if preserve_compaction_handoff:
|
|
result["replacement_message_id"] = replacement_message_id
|
|
return result
|
|
|
|
def message_count(self, session_id: str = None) -> int:
|
|
"""Count messages, optionally for a specific session."""
|
|
if session_id:
|
|
return self._read_one("SELECT COUNT(*) FROM messages WHERE session_id = ?", (session_id,))[0]
|
|
return self._read_one("SELECT COUNT(*) FROM messages")[0]
|
|
|
|
def has_platform_message_id(self, session_id: str, platform_message_id: str) -> bool:
|
|
"""True when a message with *platform_message_id* exists (partial index lookup;
|
|
the gateway's transient-failure dedupe guard)."""
|
|
return self._read_one(
|
|
"SELECT 1 FROM messages WHERE session_id = ? AND platform_message_id = ? LIMIT 1",
|
|
(session_id, platform_message_id),
|
|
) is not None
|
|
|
|
def _is_explicit_fork_child_row(self, session: Dict[str, Any]) -> bool:
|
|
"""True when *session* is a branch, delegate, or tool child of its parent.
|
|
|
|
Markers only count when they point at ``parent_session_id``: compression copies
|
|
``model_config`` onto the continuation, so a delegate's continuation carries
|
|
``_delegate_from=<the delegate's own parent>`` and presence-only matching would
|
|
misclassify it (same binding as ``_NON_CONTINUATION_CHILD_FILTER_SQL``).
|
|
"""
|
|
if session.get("source") == "tool":
|
|
return True
|
|
raw = session.get("model_config")
|
|
if not raw:
|
|
return False
|
|
try:
|
|
cfg = json.loads(raw) if isinstance(raw, str) else raw
|
|
except (TypeError, json.JSONDecodeError):
|
|
return False
|
|
if not isinstance(cfg, dict):
|
|
return False
|
|
markers = (cfg.get("_branched_from"), cfg.get("_delegate_from"))
|
|
parent_id = session.get("parent_session_id")
|
|
return parent_id in markers if parent_id else any(m is not None for m in markers)
|
|
|
|
def is_explicit_fork_child(self, session_id: str) -> bool:
|
|
"""Public read-only view of :meth:`_is_explicit_fork_child_row`; a missing row
|
|
is not a fork."""
|
|
session = self.get_session(session_id)
|
|
return bool(session and self._is_explicit_fork_child_row(session))
|
|
|
|
def latest_conversation_boundary(self, session_key: str, source: str) -> Optional[int]:
|
|
"""How many conversation boundaries (``_RESET_END_REASONS`` ends) this routing
|
|
peer has crossed, or ``None`` when never reset.
|
|
|
|
The peer is ``(session_key, source)`` — the identity recovery uses — never the
|
|
key alone (an API caller may legally reuse a Telegram row's key). Read from
|
|
``conversation_generations`` (advanced inside each boundary's txn), not an
|
|
aggregate over session rows: deletes/prunes would let an aggregate re-emit a
|
|
retired pair. Rows are never garbage-collected, by design (dropping one would
|
|
re-issue generation 1 — the ABA this counter prevents).
|
|
"""
|
|
if not session_key or not source:
|
|
return None
|
|
row = self._read_one(
|
|
"SELECT generation FROM conversation_generations WHERE source = ? AND session_key = ?",
|
|
(source, session_key),
|
|
)
|
|
if row is None or row["generation"] is None or int(row["generation"]) <= 0:
|
|
return None
|
|
return int(row["generation"])
|
|
|
|
def clear_messages(self, session_id: str) -> None:
|
|
"""Delete all messages for a session and reset its counters."""
|
|
def _do(conn):
|
|
conn.execute("DELETE FROM messages WHERE session_id = ?", (session_id,))
|
|
conn.execute(
|
|
"UPDATE sessions SET message_count = 0, tool_call_count = 0 WHERE id = ?", (session_id,),
|
|
)
|
|
self._execute_write(_do)
|
|
|
|
def purge_stale_tool_call_markers(self, *, dry_run: bool = False, backup: bool = True) -> Dict[str, Any]:
|
|
"""Permanently clear bare tool-call marker content (e.g. "[memory]") left by
|
|
pre-fix sessions. ``_rows_to_conversation`` repairs this in memory on every load,
|
|
so this is optional; it just stops the re-scan and removes the bytes.
|
|
|
|
Only ``content`` is touched (tool_call pairing unaffected). With ``backup`` a
|
|
``VACUUM INTO`` snapshot (safe against a live connection) is taken first; none
|
|
when nothing changes. ``dry_run`` reports without writing or backing up.
|
|
Returns ``{"dry_run", "rows_affected", "row_ids", "backup_path"}``.
|
|
"""
|
|
from hermes_state import _STALE_TOOL_CALL_MARKER_RE
|
|
|
|
def _find_affected(conn) -> List[int]:
|
|
cursor = conn.execute(
|
|
"SELECT id, content FROM messages "
|
|
"WHERE role = 'assistant' AND tool_calls IS NOT NULL AND tool_calls != ''"
|
|
)
|
|
return [
|
|
row["id"] for row in cursor.fetchall()
|
|
if isinstance(row["content"], str) and _STALE_TOOL_CALL_MARKER_RE.fullmatch(row["content"].strip())
|
|
]
|
|
|
|
def _result(affected, backup_path=None):
|
|
return {"dry_run": dry_run, "rows_affected": len(affected), "row_ids": affected, "backup_path": backup_path}
|
|
|
|
with self._read_ctx() as conn:
|
|
affected_ids = _find_affected(conn)
|
|
if dry_run or not affected_ids:
|
|
return _result(affected_ids)
|
|
backup_path: Optional[str] = None
|
|
if backup:
|
|
import datetime
|
|
|
|
stamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
|
|
dest = self.db_path.with_name(f"{self.db_path.name}.pre-clean-markers-backup-{stamp}")
|
|
with self._lock:
|
|
self._conn.execute("VACUUM INTO ?", (str(dest),))
|
|
backup_path = str(dest)
|
|
logger.info("Backed up state.db to %s before clean-markers write", backup_path)
|
|
|
|
def _do(conn):
|
|
ids = _find_affected(conn)
|
|
if ids:
|
|
conn.execute(f"UPDATE messages SET content = '' WHERE id IN ({_placeholders(ids)})", ids)
|
|
return ids
|
|
|
|
affected_ids = self._execute_write(_do)
|
|
if affected_ids:
|
|
logger.info(
|
|
"Permanently cleared %d stale tool-call marker row(s) in state.db (#78148)", len(affected_ids),
|
|
)
|
|
return _result(affected_ids, backup_path)
|