refactor(state): pack messages mixin signatures, inline single-use tool-call count

This commit is contained in:
Teknium
2026-09-02 23:54:52 -07:00
parent a7fa1ad696
commit 0c98e684ed
+40 -61
View File
@@ -174,12 +174,10 @@ class SessionMessagesMixin:
reasoning-replay consumers (``isinstance(..., list)``) would drop it."""
return None if not value else (value if isinstance(value, str) else json.dumps(value))
def _check_transcript_write_guards(
self, conn, session_id: str, compression_lock_holder: Optional[str],
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:
allow_closed_compression_parent: bool = False) -> None:
"""Transcript-write admission checks, run INSIDE the write txn by every writer. Ordinary appends do
NOT check compression_locks: the lock only stops two COMPRESSIONS colliding and archive_and_compact()
commits against a watermark, so concurrent appends are safe (blocking them killed turns during slow
@@ -218,10 +216,8 @@ class SessionMessagesMixin:
if _ended_by_compression(conn.execute(_ENDED_ROW_SQL, (session_id,)).fetchone()) 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:
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* already parsed;
*keep_reasoning* False NULLs every reasoning column). ``platform_message_id`` falls back to
``message_id`` (yuanbao's message-dict convention)."""
@@ -259,8 +255,7 @@ class SessionMessagesMixin:
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:
turn_lease_holder: Optional[str] = None, turn_lease_ttl_seconds: float = 300.0) -> int:
"""Append one message; returns the row id and bumps the session counters. ``platform_message_id``:
the platform's own id. ``api_content``: byte-fidelity sidecar, the exact string sent to the API when
it differed from ``content``, stored as sent except lone surrogates."""
@@ -268,14 +263,13 @@ class SessionMessagesMixin:
# Encode outside the write txn (display metadata first: log-order parity).
msg["display_metadata"] = self._encode_display_metadata(display_metadata)
tool_calls = _parse_tool_calls(tool_calls)
num_tool_calls = _tool_calls_count(tool_calls)
params = self._message_row_params(
session_id, role, msg, tool_calls, _coerce_timestamp(timestamp, time.time()), 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
self._bump_session_counters(conn, session_id, 1, num_tool_calls, unit=True)
self._bump_session_counters(conn, session_id, 1, _tool_calls_count(tool_calls), unit=True)
return msg_id
# THE critical write (failure aborts the turn): long patience so a sibling legitimately
# holding the lock for seconds (VACUUM, checkpoint) can't kill it.
@@ -284,8 +278,7 @@ class SessionMessagesMixin:
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:
turn_lease_ttl_seconds: float = 300.0) -> int:
"""Append *messages* in ONE write txn (all rows land or none, guards run once); returns the inserted
count. ``chunk_rows`` bounds txn size for LARGE copies (branch seeds; FTS triggers run per row)."""
if not messages:
@@ -306,10 +299,9 @@ class SessionMessagesMixin:
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:
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 (newest active row by content,
right after the serial turn flushed); the model still sees ``role``/``content`` unchanged, so
producer provenance survives without classifying by content at render time."""
@@ -331,9 +323,8 @@ class SessionMessagesMixin:
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]]]:
def set_message_reaction(self, session_id: str, message_row_id: int, emoji: Optional[str], *,
author: str = "user") -> Optional[List[Dict[str, Any]]]:
"""Set (``emoji=None``: clear) *author*'s reaction. Tapback semantics: one per author per message;
the same emoji again clears, a different one replaces. Returns the list after the write, or
``None`` for a foreign row."""
@@ -393,9 +384,8 @@ class SessionMessagesMixin:
return pending
return self._execute_write(_do)
def latest_message_row_id(
self, session_id: str, *, role: str = "user", offset: int = 0, require_text: bool = True
) -> Optional[int]:
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 *role* message, or ``None``. ``offset`` steps back; ``require_text``
skips rows without plain-text content so "the latest message" never resolves to an invisible bubble."""
if not session_id or role not in {"user", "assistant"} or offset < 0:
@@ -431,10 +421,8 @@ class SessionMessagesMixin:
now_ts = max(now_ts, 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:
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 a session's messages (/retry, /undo, /compress). DESTRUCTIVE by default (rows
DELETEd, leave FTS). ``active_only`` spares soft-archived rows (needed with in-place compaction).
``archive_dropped`` SOFT-archives live rows rewind-style: what rewind/edit/regenerate must use, since
@@ -486,11 +474,9 @@ class SessionMessagesMixin:
f"WHERE id IN ({_placeholders(tail_ids)}) ORDER BY id",
[session_id, *tail_ids] if retarget else tail_ids)
def archive_and_compact(
self, session_id: str, compacted_messages: List[Dict[str, Any]],
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:
lock_holder: Optional[str] = None, tail_count: int = 0) -> int:
"""Non-destructive in-place compaction under ONE session id: soft-archive the active rows (``active=0,
compacted=1``: summarized away, still searchable) and insert *compacted_messages* as fresh active
rows, atomically; returns the new ACTIVE count (= ``message_count``). *watermark* (compression
@@ -599,10 +585,9 @@ class SessionMessagesMixin:
"""Audit: every row; display: active plus compaction-archived (never Undo/Rewind rows); default: live."""
return "" if include_inactive else (_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]]:
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 in insertion order (id, never timestamp: clocks regress). ``include_inactive``:
rewind rows too; ``include_compacted``: compaction-archived display history (not rewind rows).
``latest`` pages back from the newest but returns chronological order; ``after_id``: keyset paging."""
@@ -702,10 +687,10 @@ class SessionMessagesMixin:
f"FROM messages WHERE session_id IN ({_placeholders(session_ids)})"
f"{active_clause} ORDER BY id", tuple(session_ids))
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]]:
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 format. ``include_compacted`` (deduped display history) is for DISPLAY reads
only: the model-fed restore must not regrow what compaction summarized away. ``repair_alternation``
repairs the loaded list for LIVE REPLAY callers (a durable ``user;user`` pair would re-trigger the
@@ -740,10 +725,9 @@ class SessionMessagesMixin:
messages.pop(duplicate_index)
return not prefer_current, 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]]:
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 rows (ordered by id, pre-filtered) into OpenAI format, stable key order. Every dict is
stamped ``_DB_PERSISTED_MARKER_KEY`` (born durable) so an identity-losing handoff never re-appends the
transcript on flush. ``_row_id`` is opt-in (gateway reactions); reasoning restored on assistant rows
@@ -770,8 +754,7 @@ class SessionMessagesMixin:
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"]:
if row["platform_message_id"]: # platform-side id exposed as ``message_id`` (JSONL transcript compat)
msg["message_id"] = row["platform_message_id"]
if row["observed"]:
msg["observed"] = True
@@ -782,14 +765,13 @@ class SessionMessagesMixin:
msg.update(
(col, _json_or(row[col], None, f"Failed to deserialize {col}, falling back to None"))
for col in ("reasoning_details", "codex_reasoning_items", "codex_message_items") 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
if exact_clone_key is not None:
exact_user_clones[exact_clone_key] = msg
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.
messages = _strip_stale_tool_call_markers(_strip_background_review_harness(messages))
@@ -809,10 +791,9 @@ class SessionMessagesMixin:
rows = self._fetch_conversation_rows(
self._resume_lineage_ids(session_id), _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)
[r for r in rows if r["session_id"] == session_id and r["active"]], 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)
@@ -899,8 +880,8 @@ class SessionMessagesMixin:
return None
@staticmethod
def _find_duplicate_replayed_user_message(
messages: List[Dict[str, Any]], msg: Dict[str, Any]) -> Optional[Tuple[int, bool]]:
def _find_duplicate_replayed_user_message(messages: List[Dict[str, Any]],
msg: Dict[str, Any]) -> Optional[Tuple[int, bool]]:
"""Adjacent replay duplicate ``(index, prefer_current)`` or None. Rotation may persist the current ask
in the parent and again inside a composite child carrier: carriers compare by canonical live payload,
ordinary users by exact string. The child carrier wins (it owns the durable row id and scaffold)."""
@@ -949,10 +930,9 @@ class SessionMessagesMixin:
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]:
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*, target included (the caller
pre-fills it as the next prompt). Returns ``{"rewound_count", "target_message", "new_head_id"}``, plus
``replacement_message_id`` with ``preserve_compaction_handoff`` (archives a composite summary carrier,
@@ -1077,10 +1057,9 @@ class SessionMessagesMixin:
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}")
backup_path = str(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)
self._conn.execute("VACUUM INTO ?", (backup_path,))
logger.info("Backed up state.db to %s before clean-markers write", backup_path)
def _do(conn):
ids = _find_affected(conn)