refactor(state): fold small get_messages/reaction/publish shapes; hug trailing closers (AST-identical)

This commit is contained in:
Teknium
2026-09-02 17:20:16 -07:00
parent 1991695511
commit 831f2e542c
4 changed files with 83 additions and 171 deletions
+48 -99
View File
@@ -156,8 +156,7 @@ class SessionMessagesMixin:
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",
)
"Failed to decode JSON-encoded message content; returning raw string")
return content
@staticmethod
@@ -260,8 +259,7 @@ class SessionMessagesMixin:
# 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"]),
)
(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)
@@ -326,14 +324,12 @@ class SessionMessagesMixin:
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,
)
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,
)
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(
@@ -368,25 +364,21 @@ class SessionMessagesMixin:
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,
)
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,
)
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,
encode_content_fn=self._encode_content, decode_content_fn=self._decode_content)
inserted, tool_calls_total = (
self._insert_message_rows(conn, session_id, inserted_rows) if inserted_rows else (0, 0)
)
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 + ?,
@@ -396,8 +388,7 @@ class SessionMessagesMixin:
elif inserted > 0:
conn.execute(
"UPDATE sessions SET message_count = message_count + ? WHERE id = ?",
(inserted, session_id),
)
(inserted, session_id))
return inserted
return self._execute_write(_do, patience_s=self._TRANSCRIPT_WRITE_PATIENCE_S)
@@ -454,8 +445,7 @@ class SessionMessagesMixin:
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:
if emoji and not (previous is not None and previous.get("emoji") == emoji):
reactions.append({"emoji": _scrub_surrogates(emoji), "author": author, "at": time.time()})
if reactions:
meta[self.REACTIONS_METADATA_KEY] = reactions
@@ -463,8 +453,7 @@ class SessionMessagesMixin:
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),
)
(self._encode_display_metadata(meta) if meta else None, message_row_id))
return reactions
return self._execute_write(_do)
@@ -475,8 +464,7 @@ class SessionMessagesMixin:
return []
row = self._read_one(
"SELECT display_metadata FROM messages WHERE id = ? AND session_id = ?",
(message_row_id, 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]]:
@@ -510,13 +498,11 @@ class SessionMessagesMixin:
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 "",
})
"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"]),
)
(self._encode_display_metadata(meta), row["id"]))
return pending
return self._execute_write(_do) or []
@@ -533,8 +519,7 @@ class SessionMessagesMixin:
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)),
)
(session_id, role, int(offset)))
return row[0] if row else None
def latest_user_message_row_id(self, session_id: str) -> Optional[int]:
@@ -548,8 +533,7 @@ class SessionMessagesMixin:
return None
row = self._read_one(
"SELECT role FROM messages WHERE id = ? AND session_id = ? AND active = 1",
(int(row_id), session_id),
)
(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]:
@@ -566,8 +550,7 @@ class SessionMessagesMixin:
_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
@@ -609,8 +592,7 @@ class SessionMessagesMixin:
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),
)
(total_messages, total_tool_calls, session_id))
self._execute_write(_do)
@@ -647,8 +629,7 @@ class SessionMessagesMixin:
f"INSERT INTO messages ({col_list}, {'session_id, ' if retarget else ''}active, compacted) "
f"SELECT {col_list}, {'?, ' if retarget else ''}1, 0 FROM messages "
f"WHERE id IN ({_placeholders(tail_ids)}) ORDER BY id",
[session_id, *tail_ids] if retarget else tail_ids,
)
[session_id, *tail_ids] if retarget else tail_ids)
def archive_and_compact(
self, session_id: str, compacted_messages: List[Dict[str, Any]],
@@ -692,8 +673,7 @@ class SessionMessagesMixin:
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)),
)
(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] = []
@@ -715,18 +695,15 @@ class SessionMessagesMixin:
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],
)
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],
)
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,),
)
"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)
@@ -735,14 +712,12 @@ class SessionMessagesMixin:
if model_config_patch is None:
conn.execute(
"UPDATE sessions SET message_count = ?, tool_call_count = ? WHERE id = ?",
(inserted, tool_calls_total, session_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),
)
(inserted, tool_calls_total, patched_model_config, session_id))
return inserted
return self._execute_write(_do)
@@ -766,8 +741,7 @@ class SessionMessagesMixin:
"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)),
)
(_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: the
@@ -779,21 +753,18 @@ class SessionMessagesMixin:
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"]),
})
"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"],
)
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
@@ -809,8 +780,7 @@ class SessionMessagesMixin:
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 []",
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"])
@@ -843,8 +813,7 @@ class SessionMessagesMixin:
# 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],
))
[session_id]))
if latest:
rows = rows[::-1]
rows = rows[offset:]
@@ -858,9 +827,7 @@ class SessionMessagesMixin:
"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)
params: list = [session_id] if after_id is None else [session_id, after_id]
if limit is not None or offset:
# SQLite's OFFSET requires LIMIT; -1 means "no limit".
sql += " LIMIT ? OFFSET ?"
@@ -877,14 +844,13 @@ class SessionMessagesMixin:
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(
found.extend({"session_id": row[0], "content": row[1]} for row in 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]:
@@ -893,11 +859,9 @@ class SessionMessagesMixin:
*window* = session boundary). Empty when the anchor is not in *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:
if not conn.execute(
"SELECT 1 FROM messages WHERE id = ? AND session_id = ? LIMIT 1", (around_message_id, session_id),
).fetchone():
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 ?",
@@ -987,8 +951,7 @@ class SessionMessagesMixin:
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,
)
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* -> ``(skip, exact_clone_key)``.
@@ -1045,8 +1008,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 []",
)
"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"]
@@ -1075,14 +1037,12 @@ class SessionMessagesMixin:
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,
)
"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]]]:
@@ -1096,12 +1056,10 @@ class SessionMessagesMixin:
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,
)
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,
)
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]:
@@ -1122,8 +1080,7 @@ class SessionMessagesMixin:
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),
)
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:
@@ -1143,14 +1100,12 @@ class SessionMessagesMixin:
"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),
)
")", (*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",
)
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]]:
@@ -1187,13 +1142,11 @@ class SessionMessagesMixin:
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,
)
is_composite)
@staticmethod
def _exact_replayed_user_clone_key(timestamp: Any, content: Any) -> Optional[Tuple[Any, str]]:
@@ -1253,7 +1206,6 @@ class SessionMessagesMixin:
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"))
@@ -1323,8 +1275,7 @@ class SessionMessagesMixin:
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),
)
(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()
@@ -1389,8 +1340,7 @@ class SessionMessagesMixin:
return None
row = self._read_one(
"SELECT generation FROM conversation_generations WHERE source = ? AND session_key = ?",
(source, 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"])
@@ -1432,7 +1382,6 @@ class SessionMessagesMixin:
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: