refactor(state): fold small get_messages/reaction/publish shapes; hug trailing closers (AST-identical)
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user