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

View File

@@ -31,8 +31,7 @@ def _cooldown_row(exists: bool, cooldown_until, error) -> Dict[str, Any]:
return {
"session_exists": exists,
"cooldown_until": float(cooldown_until) if cooldown_until is not None else None,
"error": error,
}
"error": error}
def _claim_lease_row(conn, table: str, key_col: str, key: str, holder: str, now: float, expires_at: float,
@@ -47,8 +46,7 @@ def _claim_lease_row(conn, table: str, key_col: str, key: str, holder: str, now:
reclaimed_holder = row["holder"]
conn.execute(
f"INSERT OR IGNORE INTO {table} ({key_col}, holder, acquired_at, expires_at) VALUES (?, ?, ?, ?)",
(key, holder, now, expires_at),
)
(key, holder, now, expires_at))
owner = conn.execute(f"SELECT holder FROM {table} WHERE {key_col} = ?", (key,)).fetchone()
return owner is not None and owner["holder"] == holder, reclaimed_holder
@@ -122,15 +120,13 @@ class SessionCompressionMixin:
deleted = conn.execute(
"DELETE FROM compression_locks "
"WHERE session_id = ? AND holder = ? AND expires_at = ?",
(session_id, lock_row["holder"], expires_at),
)
(session_id, lock_row["holder"], expires_at))
if deleted.rowcount != 1:
return False
updated = conn.execute(
"UPDATE sessions SET ended_at = NULL, end_reason = NULL "
"WHERE id = ? AND ended_at IS NOT NULL AND end_reason = 'compression'",
(session_id,),
)
(session_id,))
# rowcount==1 is guaranteed by the parent SELECT in this same txn. A False
# return added past this point must raise instead: the lease DELETE above
# commits unless _do raises.
@@ -159,8 +155,7 @@ class SessionCompressionMixin:
parent["git_repo_root"],
profile_name or parent["profile_name"] or self._own_profile_name(),
parent["user_id"], parent["session_key"], parent["chat_id"], parent["chat_type"],
parent["thread_id"], parent["display_name"], parent["origin_json"], time.time(),
),
parent["thread_id"], parent["display_name"], parent["origin_json"], time.time()),
)
def publish_compression_child(
@@ -193,8 +188,7 @@ class SessionCompressionMixin:
conn.execute(
"UPDATE compression_locks SET expires_at = ? "
"WHERE session_id = ? AND holder = ?",
(time.time() + lease_ttl_seconds, parent_session_id, compression_lock_holder),
)
(time.time() + lease_ttl_seconds, parent_session_id, compression_lock_holder))
lock_row = conn.execute(_LOCK_ROW_SQL, (parent_session_id,)).fetchone()
if require_compression_lease and (
lock_row is None
@@ -223,8 +217,7 @@ class SessionCompressionMixin:
if is_automatic_end_reason(parent["end_reason"]):
conn.execute(
"UPDATE sessions SET ended_at = NULL, end_reason = NULL WHERE id = ?",
(parent_session_id,),
)
(parent_session_id,))
else:
raise RuntimeError(f"Compression parent already ended: {parent_session_id}")
if not messages:
@@ -232,21 +225,17 @@ class SessionCompressionMixin:
self._publish_child_session_row(
conn, parent, parent_session_id=parent_session_id, child_session_id=child_session_id,
source=source, model=model, model_config=model_config, system_prompt=system_prompt,
cwd=cwd, profile_name=profile_name,
)
cwd=cwd, profile_name=profile_name)
total_messages, total_tool_calls = self._insert_message_rows(conn, child_session_id, messages)
if watermark is not None:
# Clone the parent's concurrent tail into the child after the handoff;
# originals stay in the closed parent for lineage recovery.
_ceiling_clause = ""
_params: list = [parent_session_id, int(watermark)]
if watermark_ceiling is not None:
_ceiling_clause = " AND id <= ?"
_params.append(int(watermark_ceiling))
bounded = watermark_ceiling 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 > ?"
f"{_ceiling_clause} ORDER BY id", _params,
f"{' AND id <= ?' if bounded else ''} ORDER BY id",
[parent_session_id, int(watermark), *([int(watermark_ceiling)] if bounded else [])],
)
if tail_ids:
self._clone_message_rows(conn, tail_ids, session_id=child_session_id)
@@ -254,12 +243,10 @@ class SessionCompressionMixin:
total_tool_calls += tail_tool_calls
conn.execute(
"UPDATE sessions SET message_count = ?, tool_call_count = ? WHERE id = ?",
(total_messages, total_tool_calls, child_session_id),
)
(total_messages, total_tool_calls, child_session_id))
updated = conn.execute(
"UPDATE sessions SET ended_at = ?, end_reason = 'compression' "
"WHERE id = ? AND ended_at IS NULL", (time.time(), parent_session_id),
)
"WHERE id = ? AND ended_at IS NULL", (time.time(), parent_session_id))
if updated.rowcount != 1:
raise RuntimeError(f"Compression parent changed during publication: {parent_session_id}")
@@ -285,8 +272,7 @@ class SessionCompressionMixin:
" AND compression_failure_cooldown_until > ? "
"THEN compression_failure_cooldown_until ELSE ? END, "
"compression_failure_error = ? WHERE id = ?",
(cooldown_until, cooldown_until, error, session_id),
)
(cooldown_until, cooldown_until, error, session_id))
def get_compression_failure_cooldown(self, session_id: str) -> Optional[Dict[str, Any]]:
"""Return the active (unexpired) compression-failure cooldown, or None."""
@@ -320,8 +306,7 @@ class SessionCompressionMixin:
def _do(conn):
cursor = conn.execute(
"UPDATE sessions SET compression_failure_cooldown_until = ?, "
"compression_failure_error = ? WHERE id = ?", (deadline, error, session_id),
)
"compression_failure_error = ? WHERE id = ?", (deadline, error, session_id))
if cursor.rowcount != 1:
raise RuntimeError(f"compression cooldown rollback session missing: {session_id}")
self._execute_write(_do)
@@ -340,8 +325,7 @@ class SessionCompressionMixin:
self._write_sql_logged(
"clear_compression_failure_cooldown", session_id,
"UPDATE sessions SET compression_failure_cooldown_until = NULL, "
"compression_failure_error = NULL WHERE id = ?", (session_id,),
)
"compression_failure_error = NULL WHERE id = ?", (session_id,))
def _read_session_number(self, column: str, session_id: str, cast: type, zero: Any) -> Any:
"""Read one numeric ``sessions`` column clamped at ``zero``; a missing session,
@@ -365,8 +349,7 @@ class SessionCompressionMixin:
if session_id:
self._write_sql(
"UPDATE sessions SET compression_fallback_streak = ? WHERE id = ?",
(max(0, int(streak)), session_id),
)
(max(0, int(streak)), session_id))
def get_compression_ineffective_count(self, session_id: str) -> int:
"""Persisted ineffective-compaction strike count — the durable half of the
@@ -379,8 +362,7 @@ class SessionCompressionMixin:
if session_id:
self._write_sql(
"UPDATE sessions SET compression_ineffective_count = ? WHERE id = ?",
(max(0, int(count)), session_id),
)
(max(0, int(count)), session_id))
def get_compression_recovery_deadline(self, session_id: str) -> float:
"""Persisted anti-thrash recovery deadline (epoch; ``0.0`` = not armed). Durable
@@ -397,8 +379,7 @@ class SessionCompressionMixin:
normalized = 0.0
self._write_sql(
"UPDATE sessions SET compression_recovery_deadline = ? WHERE id = ?",
(normalized if normalized > 0.0 else None, session_id),
)
(normalized or None, session_id))
def refresh_compression_lock(self, session_id: str, holder: str, ttl_seconds: float = 300.0) -> bool:
"""Extend the compression lock lease if ``holder`` still owns it.
@@ -438,15 +419,12 @@ class SessionCompressionMixin:
def _do(conn):
return _claim_lease_row(
conn, "compression_locks", "session_id", session_id, holder, now, expires_at,
lambda h, e: e < now or _compression_lock_holder_process_is_dead(h),
)
lambda h, e: e < now or _compression_lock_holder_process_is_dead(h))
try:
acquired, reclaimed_holder = self._execute_write(_do)
if reclaimed_holder:
logger.warning(
"Reclaimed stale compression lock for session=%s (holder=%s)", session_id, reclaimed_holder,
)
logger.warning("Reclaimed stale compression lock for session=%s (holder=%s)", session_id, reclaimed_holder)
return bool(acquired)
except sqlite3.Error as exc:
# False makes the caller skip compression — safe when the lock subsystem is broken.
@@ -460,8 +438,7 @@ class SessionCompressionMixin:
self._write_sql_logged(
"release_compression_lock", session_id,
"DELETE FROM compression_locks WHERE session_id = ? AND holder = ?",
(session_id, holder),
)
(session_id, holder))
def _session_turn_lease_key_on_conn(self, conn, session_id: str) -> str:
"""Walk compression parents on ``conn`` to the conversation lease key.
@@ -474,8 +451,7 @@ class SessionCompressionMixin:
def _row(sid: str):
row = conn.execute(
"SELECT id, parent_session_id, source, model_config, end_reason "
"FROM sessions WHERE id = ?", (sid,),
"SELECT id, parent_session_id, source, model_config, end_reason FROM sessions WHERE id = ?", (sid,),
).fetchone()
return dict(row) if row else None
@@ -546,8 +522,7 @@ class SessionCompressionMixin:
logger.debug("session turn lease should_abort callback failed", exc_info=True)
try:
if self.try_acquire_session_turn_lease(
session_id, holder, ttl_seconds=ttl_seconds, patience_s=acquire_patience_s,
):
session_id, holder, ttl_seconds=ttl_seconds, patience_s=acquire_patience_s):
return True
except sqlite3.Error as exc:
# Long holder transactions can exhaust one write-patience budget; keep
@@ -594,8 +569,7 @@ class SessionCompressionMixin:
conversation_id = self._session_turn_lease_key_on_conn(conn, session_id)
conn.execute(
"DELETE FROM session_turn_leases WHERE conversation_id = ? AND holder = ?",
(conversation_id, holder),
)
(conversation_id, holder))
self._execute_write(_do)
@@ -605,8 +579,7 @@ class SessionCompressionMixin:
return None
row = self._read_one(
"SELECT holder FROM compression_locks WHERE session_id = ? AND expires_at >= ?",
(session_id, time.time()),
)
(session_id, time.time()))
return None if row is None else row[0]
def finalize_orphaned_compression_sessions(self) -> int:

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:

View File

@@ -174,8 +174,7 @@ class SessionTitlesMixin:
row = self._read_one(
"SELECT s.*, COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved "
"FROM sessions s LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash "
"WHERE s.title = ?", (title,),
)
"WHERE s.title = ?", (title,))
return self._session_row_dict(row) if row else None
def resolve_session_by_title(self, title: str) -> Optional[str]:
@@ -186,8 +185,7 @@ class SessionTitlesMixin:
numbered = self._read_all(
"SELECT id, title, started_at FROM sessions "
"WHERE title LIKE ? ESCAPE '\\' ORDER BY started_at DESC",
(f"{_escape_like(title)} #%",),
)
(f"{_escape_like(title)} #%",))
if numbered:
return numbered[0]["id"]
return exact["id"] if exact else None
@@ -199,8 +197,7 @@ class SessionTitlesMixin:
base = match.group(1) if match else base_title
rows = self._read_all(
"SELECT title FROM sessions WHERE title = ? OR title LIKE ? ESCAPE '\\'",
(base, f"{_escape_like(base)} #%"),
)
(base, f"{_escape_like(base)} #%"))
if not rows:
return base
# The unnumbered original counts as #1.

View File

@@ -83,8 +83,7 @@ _MODEL_USAGE_UPSERT_SQL = """INSERT INTO session_model_usage (
_MODEL_USAGE_FIELDS = frozenset((
"model", "billing_provider", "billing_base_url", "billing_mode", "input_tokens", "output_tokens",
"cache_read_tokens", "cache_write_tokens", "reasoning_tokens", "estimated_cost_usd",
"actual_cost_usd", "cost_status", "cost_source", "api_call_count",
))
"actual_cost_usd", "cost_status", "cost_source", "api_call_count"))
class SessionUsageMixin:
@@ -128,8 +127,7 @@ class SessionUsageMixin:
# leftovers. ``not is_alive()`` (not ``is None``) respawns a writer
# that died from an unexpected escape.
thread = threading.Thread(
target=self._token_writer_loop, name="session-db-token-writer", daemon=True,
)
target=self._token_writer_loop, name="session-db-token-writer", daemon=True)
self._token_writer_thread = thread
thread.start()
if self._token_atexit_hook is None:
@@ -255,8 +253,7 @@ class SessionUsageMixin:
# Writer stuck mid-apply: leave deltas unapplied rather than race it.
logger.warning(
"async token accounting: writer did not stop within %.0fs; "
"%d queued delta(s) not persisted", join_timeout, len(self._token_queue),
)
"%d queued delta(s) not persisted", join_timeout, len(self._token_queue))
return
# Writer gone: apply leftovers synchronously under the same busy protocol. Wait
# out a flush caller-drain that already claimed busy — close() nulls the
@@ -269,8 +266,7 @@ class SessionUsageMixin:
logger.warning(
"async token accounting: concurrent drain did not "
"finish within %.0fs; %d queued delta(s) not persisted",
join_timeout, len(self._token_queue),
)
join_timeout, len(self._token_queue))
return
self._token_queue_cond.wait(remaining)
# busy BEFORE clearing the queue (same ordering as the writer loop).
@@ -315,8 +311,7 @@ class SessionUsageMixin:
billing_provider if has_accounted_usage else None,
billing_base_url if has_accounted_usage else None,
billing_mode if has_accounted_usage else None, model if has_accounted_usage else None,
api_call_count, session_id,
)
api_call_count, session_id)
# Per-model attribution: the sessions row keeps one (model, provider) pair, so a
# mid-session /model switch would attribute every token to the initial model.
# Only the incremental path records here — absolute cumulative updates cannot be
@@ -378,9 +373,7 @@ class SessionUsageMixin:
billing_mode or sess.get("billing_mode") or "", task or "",
api_call_count or 0, *counts,
float(estimated_cost_usd or 0.0), float(actual_cost_usd or 0.0),
cost_status, cost_source, now, now,
),
)
cost_status, cost_source, now, now))
def record_auxiliary_usage(
self, session_id: str, task: str, *, model: Optional[str] = None,