perf(state): scan only the carried rows when every carried dict has a row id

The common micro-compaction pass carries dicts that all hold a matching
_row_id, yet _resolve_carried_row_ids re-read and re-derived identity for
every active row in the session inside the write transaction, O(active
rows) per turn on long sessions. Narrow the identity query to the carried
ids in that case and keep the full active-row scan for the fallbacks
(missing ids, or a stale id whose stored identity no longer matches), so
resolution semantics are unchanged.
This commit is contained in:
teknium1
2026-09-22 21:52:47 -07:00
committed by Teknium
parent 91df54184d
commit 8fc679b4e8

View File

@@ -674,37 +674,49 @@ class SessionMessagesMixin:
"""
if not carried_messages:
return []
rows = conn.execute(
"SELECT id, role, content, tool_call_id, tool_calls, timestamp FROM messages "
"WHERE session_id = ? AND active = 1 ORDER BY id",
(session_id,),
).fetchall()
by_id: Dict[int, Tuple[Any, ...]] = {}
by_key: Dict[Tuple[Any, ...], List[int]] = {}
for row in rows:
row_id = int(row["id"])
identity = self._row_identity(
row["role"], self._decode_content(row["content"]), row["tool_call_id"],
_parse_tool_calls(row["tool_calls"]))
by_id[row_id] = identity
timestamp = coerce_epoch(row["timestamp"], field="message timestamp")
if timestamp is not None:
by_key.setdefault((*identity, timestamp), []).append(row_id)
resolved: List[int] = []
carried: List[Tuple[Tuple[Any, ...], Any, Any]] = []
for message in carried_messages:
if not isinstance(message, dict):
continue
role = message.get("role", "unknown")
identity = self._row_identity(
role, message.get("content"), message.get("tool_call_id"),
message.get("role", "unknown"), message.get("content"), message.get("tool_call_id"),
_parse_tool_calls(message.get("tool_calls")))
row_id = message.get("_row_id")
if (isinstance(row_id, int) and not isinstance(row_id, bool)
and row_id > 0 and by_id.get(row_id) == identity):
if not (isinstance(row_id, int) and not isinstance(row_id, bool) and row_id > 0):
row_id = None
carried.append((identity, row_id, message.get("timestamp")))
def _index(ids: Optional[List[int]]):
by_id: Dict[int, Tuple[Any, ...]] = {}
by_key: Dict[Tuple[Any, ...], List[int]] = {}
narrow = f" AND id IN ({_placeholders(ids)})" if ids else ""
for row in conn.execute(
"SELECT id, role, content, tool_call_id, tool_calls, timestamp FROM messages "
f"WHERE session_id = ? AND active = 1{narrow} ORDER BY id",
(session_id, *(ids or ())),
).fetchall():
rid = int(row["id"])
by_id[rid] = self._row_identity(
row["role"], self._decode_content(row["content"]), row["tool_call_id"],
_parse_tool_calls(row["tool_calls"]))
ts = coerce_epoch(row["timestamp"], field="message timestamp")
if ts is not None:
by_key.setdefault((*by_id[rid], ts), []).append(rid)
return by_id, by_key
# The common micro pass carries dicts that all hold a matching _row_id, so the identity
# check only needs those rows; a full active-row scan is reserved for the fallbacks.
row_ids = [row_id for _, row_id, _ in carried if row_id is not None]
by_id, by_key = _index(row_ids if len(row_ids) == len(carried) else None)
if len(row_ids) == len(carried) and any(by_id.get(rid) != ident for ident, rid, _ in carried):
by_id, by_key = _index(None)
resolved: List[int] = []
for identity, row_id, raw_timestamp in carried:
if row_id is not None and by_id.get(row_id) == identity:
resolved.append(row_id)
continue
timestamp = coerce_epoch(message.get("timestamp"), field="message timestamp")
timestamp = coerce_epoch(raw_timestamp, field="message timestamp")
if timestamp is None:
continue
matches = by_key.get((*identity, timestamp), [])