perf(state): batch compression-tip row fetch in list_sessions_rich

list_sessions_rich()'s compression-root projection called
_get_session_rich_row() once per root — a separate single-row query per
compression root on every session-list render. Resolve every tip id
first, then fetch all tip rows in one WHERE id IN (...) query via the
new _get_session_rich_rows_batch().

_get_session_rich_row() is now a thin wrapper over the batch method, so
the enriched SELECT (preview + last_active) lives in exactly one place —
future column changes (e.g. #42196's include_system_prompt) only touch
one query.

get_compression_tip()'s chain walk is untouched; it's a genuine
per-session graph walk with branch/delegate-exclusion and race handling,
and batching it safely is out of scope here.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
jasoisjaso
2026-07-06 20:58:58 +10:00
committed by kshitij
parent 5017deb3db
commit adcdf9dc63
3 changed files with 135 additions and 14 deletions
+21 -6
View File
@@ -5751,16 +5751,31 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
# as the live conversation. Keep the root's started_at to preserve
# chronological ordering by original conversation start.
if project_compression_tips and not include_children:
projected = []
# get_compression_tip() walks each root's chain individually (it's
# a per-session graph walk, not batchable in one query), but the
# tip *row* fetch afterward was previously one _get_session_rich_row()
# call per compression root. Batch that half instead: resolve
# every tip id first, then fetch all tip rows in a single query.
tip_ids_by_root: Dict[str, str] = {}
for s in sessions:
if s.get("end_reason") != "compression":
projected.append(s)
continue
tip_id = self.get_compression_tip(s["id"])
if tip_id == s["id"]:
projected.append(s)
continue
tip_row = self._get_session_rich_row(tip_id, compact_rows=compact_rows)
if tip_id != s["id"]:
tip_ids_by_root[s["id"]] = tip_id
tip_rows = (
self._get_session_rich_rows_batch(
set(tip_ids_by_root.values()), compact_rows=compact_rows
)
if tip_ids_by_root
else {}
)
projected = []
for s in sessions:
tip_id = tip_ids_by_root.get(s["id"])
tip_row = tip_rows.get(tip_id) if tip_id else None
if not tip_row:
projected.append(s)
continue
+32 -8
View File
@@ -133,10 +133,33 @@ class SessionPortabilityMixin:
Pass ``compact_rows=True`` to omit the ``system_prompt`` blob (see
``list_sessions_rich`` for details).
Thin wrapper over ``_get_session_rich_rows_batch`` so the enriched
SELECT lives in exactly one place.
"""
return self._get_session_rich_rows_batch(
[session_id], compact_rows=compact_rows
).get(session_id)
def _get_session_rich_rows_batch(
self, session_ids, compact_rows: bool = False
) -> Dict[str, Dict[str, Any]]:
"""Fetch multiple sessions with the same enriched columns as
``_get_session_rich_row``, in a single query.
Used by ``list_sessions_rich``'s compression-tip projection to resolve
every tip row for a page in one round trip instead of one query per
compression-root row. Returns a dict keyed by session id; ids that
don't exist are simply absent from the result (same as
``_get_session_rich_row`` returning ``None`` for them).
"""
ids = [sid for sid in session_ids if sid]
if not ids:
return {}
# Same read-your-writes guarantee as list_sessions_rich.
self.flush_token_counts()
_sel = self._compact_session_cols() if compact_rows else "s.*"
placeholders = ",".join("?" for _ in ids)
query = f"""
SELECT {_sel},
COALESCE(
@@ -148,16 +171,17 @@ class SessionPortabilityMixin:
) AS _preview_raw,
{_sql_session_last_active("s")} AS last_active
FROM sessions s
WHERE s.id = ?
WHERE s.id IN ({placeholders})
"""
with self._lock:
cursor = self._conn.execute(query, (session_id,))
row = cursor.fetchone()
if not row:
return None
s = dict(row)
s["preview"] = _shape_preview(s.pop("_preview_raw", ""))
return s
cursor = self._conn.execute(query, ids)
rows = cursor.fetchall()
result: Dict[str, Dict[str, Any]] = {}
for row in rows:
s = dict(row)
s["preview"] = _shape_preview(s.pop("_preview_raw", ""))
result[s["id"]] = s
return result
def get_session_rich_row(self, session_id: str, compact_rows: bool = False) -> Optional[Dict[str, Any]]:
"""Public wrapper for :meth:`_get_session_rich_row`.
+82
View File
@@ -1851,6 +1851,88 @@ class TestCompressionChainProjection:
assert tip_row["ended_at"] is None # tip is still live
assert tip_row["end_reason"] is None
def test_list_projects_multiple_independent_chains_in_one_call(self, db):
"""Two unrelated compression chains in the same page must each
resolve to their own tip, not get cross-mixed by the batched tip-row
fetch (regression test for the single-query batch in
_get_session_rich_rows_batch — a wrong id->row mapping there would
silently swap one chain's data onto the other)."""
import time as _time
t0 = _time.time() - 7200
self._build_compression_chain(db, t0)
# Second, independent chain — same shape, different ids/content.
db.create_session("root2", "cli")
db._conn.execute("UPDATE sessions SET started_at=? WHERE id=?", (t0 + 100, "root2"))
db.append_message("root2", "user", "second conversation start")
db._conn.execute(
"UPDATE sessions SET ended_at=?, end_reason=? WHERE id=?",
(t0 + 200, "compression", "root2"),
)
db.create_session("tip2", "cli", parent_session_id="root2")
db._conn.execute("UPDATE sessions SET started_at=? WHERE id=?", (t0 + 201, "tip2"))
db.append_message("tip2", "user", "second conversation continuation")
db.update_session_cwd("tip2", "/tmp/workspaces/second")
db._conn.commit()
sessions = db.list_sessions_rich(source="cli", limit=20)
ids = [s["id"] for s in sessions]
assert "root1" not in ids and "root2" not in ids
assert "tip1" in ids and "tip2" in ids
tip1_row = next(s for s in sessions if s["id"] == "tip1")
tip2_row = next(s for s in sessions if s["id"] == "tip2")
assert tip1_row["_lineage_root_id"] == "root1"
assert tip1_row["preview"].startswith("latest message")
assert tip2_row["_lineage_root_id"] == "root2"
assert tip2_row["preview"].startswith("second conversation continuation")
assert tip2_row["cwd"] == "/tmp/workspaces/second"
def test_list_batches_tip_row_fetch_into_one_query(self, db, monkeypatch):
"""Projection must resolve tip rows for a whole page in one batched
query, not one _get_session_rich_row() call per compression root."""
import time as _time
t0 = _time.time() - 7200
self._build_compression_chain(db, t0)
db.create_session("root2", "cli")
db._conn.execute("UPDATE sessions SET started_at=? WHERE id=?", (t0 + 100, "root2"))
db.append_message("root2", "user", "second conversation start")
db._conn.execute(
"UPDATE sessions SET ended_at=?, end_reason=? WHERE id=?",
(t0 + 200, "compression", "root2"),
)
db.create_session("tip2", "cli", parent_session_id="root2")
db._conn.execute("UPDATE sessions SET started_at=? WHERE id=?", (t0 + 201, "tip2"))
db.append_message("tip2", "user", "second continuation")
db._conn.commit()
batch_calls = []
single_calls = []
original_batch = db._get_session_rich_rows_batch
original_single = db._get_session_rich_row
def counting_batch(session_ids, **kwargs):
batch_calls.append(list(session_ids))
return original_batch(session_ids, **kwargs)
def counting_single(session_id, **kwargs):
single_calls.append(session_id)
return original_single(session_id, **kwargs)
monkeypatch.setattr(db, "_get_session_rich_rows_batch", counting_batch)
monkeypatch.setattr(db, "_get_session_rich_row", counting_single)
sessions = db.list_sessions_rich(source="cli", limit=20)
assert len(sessions) >= 2 # sanity: both chains actually surfaced
# Two compression roots resolved with exactly one batched call, and
# zero single-row calls — not one single-row call per root.
assert len(batch_calls) == 1
assert set(batch_calls[0]) == {"tip1", "tip2"}
assert single_calls == []