perf(state): batch the turn flush into one SQLite transaction
Re-derivation of #23254 (@devsart95) on today's flush loop. The turn flush in _flush_messages_to_session_db wrote one BEGIN IMMEDIATE transaction per message row; a typical agent turn (user + assistant + tool results) paid 3-8 transactions -- and, off WAL (the default on macOS while the WAL-reset guard is active), 3-8 fsyncs -- per turn. Adds SessionDB.append_messages_batch: same row shape as append_message (shared _prepare_message_row serializer + _MESSAGE_INSERT_SQL column list, so the two writers cannot drift), same compression-lock and compression-closed guards, one aggregated session-counter UPDATE, one transaction for the whole batch. Row serialization stays outside the write lock. The flush loop now collects the turn's new rows and writes them in one call. All-or-nothing pairs exactly with the persisted-marker stamping: on failure no rows landed and no markers were stamped, so the next flush re-writes the whole tail (same recovery contract as before, minus the partial-prefix case that could double-count). Measured (same harness, 5-message turn, journal_mode=DELETE, synchronous=FULL): 2.32ms -> 0.83ms median per turn flush (64% faster, 5 fsyncs -> 1). On WAL the win is smaller but the atomicity fix holds.
This commit is contained in:
+240
-72
@@ -6001,6 +6001,117 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
)
|
||||
return None
|
||||
|
||||
# INSERT column list shared by append_message and append_messages_batch so
|
||||
# the row shape can never drift between the single and batched writers.
|
||||
_MESSAGE_INSERT_SQL = (
|
||||
"INSERT INTO messages (session_id, role, content, tool_call_id, "
|
||||
"tool_calls, tool_name, effect_disposition, timestamp, token_count, finish_reason, "
|
||||
"reasoning, reasoning_content, reasoning_details, codex_reasoning_items, "
|
||||
"codex_message_items, platform_message_id, observed, active, api_content, display_kind, display_metadata) "
|
||||
"VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)"
|
||||
)
|
||||
|
||||
def _prepare_message_row(
|
||||
self,
|
||||
*,
|
||||
session_id: str,
|
||||
role: str,
|
||||
content: Any = None,
|
||||
tool_name: Optional[str] = None,
|
||||
tool_calls: Any = None,
|
||||
tool_call_id: Optional[str] = None,
|
||||
token_count: Optional[int] = None,
|
||||
finish_reason: Optional[str] = None,
|
||||
reasoning: Optional[str] = None,
|
||||
reasoning_content: Optional[str] = None,
|
||||
reasoning_details: Any = None,
|
||||
codex_reasoning_items: Any = None,
|
||||
codex_message_items: Any = None,
|
||||
platform_message_id: Optional[str] = None,
|
||||
observed: bool = False,
|
||||
effect_disposition: Optional[str] = None,
|
||||
timestamp: Any = None,
|
||||
api_content: Optional[str] = None,
|
||||
display_kind: Optional[str] = None,
|
||||
display_metadata: Optional[Dict[str, Any]] = None,
|
||||
) -> Tuple[tuple, int]:
|
||||
"""Serialize one message into ``_MESSAGE_INSERT_SQL`` bind params.
|
||||
|
||||
Runs entirely OUTSIDE the write transaction (JSON encoding, surrogate
|
||||
scrubbing, timestamp coercion), so batched flushes keep the BEGIN
|
||||
IMMEDIATE window as short as possible. Returns ``(params,
|
||||
num_tool_calls)`` where ``num_tool_calls`` feeds the session counter
|
||||
update.
|
||||
"""
|
||||
# Display metadata is presentation-only and never changes the model
|
||||
# context role/content replayed to providers.
|
||||
display_metadata_json = self._encode_display_metadata(display_metadata)
|
||||
# Serialize structured fields to JSON before entering the write txn
|
||||
reasoning_details_json = (
|
||||
json.dumps(reasoning_details)
|
||||
if reasoning_details else None
|
||||
)
|
||||
codex_items_json = (
|
||||
json.dumps(codex_reasoning_items)
|
||||
if codex_reasoning_items else None
|
||||
)
|
||||
codex_message_items_json = (
|
||||
json.dumps(codex_message_items)
|
||||
if codex_message_items else None
|
||||
)
|
||||
# tool_calls may arrive as a Python list (from the live agent) or
|
||||
# as a JSON string (from import/export). Parse first to avoid
|
||||
# double-encoding.
|
||||
if isinstance(tool_calls, str):
|
||||
try:
|
||||
tool_calls = json.loads(tool_calls)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
tool_calls = []
|
||||
tool_calls_json = json.dumps(tool_calls) if tool_calls else None
|
||||
# Multimodal content (list of parts) must be JSON-encoded: sqlite3
|
||||
# cannot bind list/dict parameters directly.
|
||||
stored_content = self._encode_content(content)
|
||||
|
||||
message_timestamp = time.time()
|
||||
if timestamp is not None:
|
||||
try:
|
||||
if hasattr(timestamp, "timestamp"):
|
||||
message_timestamp = float(timestamp.timestamp())
|
||||
else:
|
||||
message_timestamp = float(timestamp)
|
||||
except (TypeError, ValueError):
|
||||
logger.debug("Ignoring invalid explicit message timestamp: %r", timestamp)
|
||||
|
||||
# Pre-compute tool call count
|
||||
num_tool_calls = 0
|
||||
if tool_calls is not None:
|
||||
num_tool_calls = len(tool_calls) if isinstance(tool_calls, list) else 1
|
||||
|
||||
params = (
|
||||
session_id,
|
||||
role,
|
||||
stored_content,
|
||||
tool_call_id,
|
||||
tool_calls_json,
|
||||
_scrub_surrogates(tool_name),
|
||||
effect_disposition,
|
||||
message_timestamp,
|
||||
token_count,
|
||||
finish_reason,
|
||||
_scrub_surrogates(reasoning),
|
||||
_scrub_surrogates(reasoning_content),
|
||||
reasoning_details_json,
|
||||
codex_items_json,
|
||||
codex_message_items_json,
|
||||
platform_message_id,
|
||||
1 if observed else 0,
|
||||
1,
|
||||
_scrub_surrogates(api_content) if isinstance(api_content, str) else None,
|
||||
_scrub_surrogates(display_kind) if isinstance(display_kind, str) else None,
|
||||
display_metadata_json,
|
||||
)
|
||||
return params, num_tool_calls
|
||||
|
||||
@staticmethod
|
||||
def _decode_display_metadata(raw: Any) -> Optional[Dict[str, Any]]:
|
||||
"""Decode a ``display_metadata`` column into the dict every reader expects.
|
||||
@@ -6068,49 +6179,28 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
from every outgoing payload anyway, so the scrubbed form IS the
|
||||
wire bytes).
|
||||
"""
|
||||
# Display metadata is presentation-only and never changes the model
|
||||
# context role/content replayed to providers.
|
||||
display_metadata_json = self._encode_display_metadata(display_metadata)
|
||||
# Serialize structured fields to JSON before entering the write txn
|
||||
reasoning_details_json = (
|
||||
json.dumps(reasoning_details)
|
||||
if reasoning_details else None
|
||||
row_params, num_tool_calls = self._prepare_message_row(
|
||||
session_id=session_id,
|
||||
role=role,
|
||||
content=content,
|
||||
tool_name=tool_name,
|
||||
tool_calls=tool_calls,
|
||||
tool_call_id=tool_call_id,
|
||||
token_count=token_count,
|
||||
finish_reason=finish_reason,
|
||||
reasoning=reasoning,
|
||||
reasoning_content=reasoning_content,
|
||||
reasoning_details=reasoning_details,
|
||||
codex_reasoning_items=codex_reasoning_items,
|
||||
codex_message_items=codex_message_items,
|
||||
platform_message_id=platform_message_id,
|
||||
observed=observed,
|
||||
effect_disposition=effect_disposition,
|
||||
timestamp=timestamp,
|
||||
api_content=api_content,
|
||||
display_kind=display_kind,
|
||||
display_metadata=display_metadata,
|
||||
)
|
||||
codex_items_json = (
|
||||
json.dumps(codex_reasoning_items)
|
||||
if codex_reasoning_items else None
|
||||
)
|
||||
codex_message_items_json = (
|
||||
json.dumps(codex_message_items)
|
||||
if codex_message_items else None
|
||||
)
|
||||
# tool_calls may arrive as a Python list (from the live agent) or
|
||||
# as a JSON string (from import/export). Parse first to avoid
|
||||
# double-encoding.
|
||||
if isinstance(tool_calls, str):
|
||||
try:
|
||||
tool_calls = json.loads(tool_calls)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
tool_calls = []
|
||||
tool_calls_json = json.dumps(tool_calls) if tool_calls else None
|
||||
# Multimodal content (list of parts) must be JSON-encoded: sqlite3
|
||||
# cannot bind list/dict parameters directly.
|
||||
stored_content = self._encode_content(content)
|
||||
|
||||
message_timestamp = time.time()
|
||||
if timestamp is not None:
|
||||
try:
|
||||
if hasattr(timestamp, "timestamp"):
|
||||
message_timestamp = float(timestamp.timestamp())
|
||||
else:
|
||||
message_timestamp = float(timestamp)
|
||||
except (TypeError, ValueError):
|
||||
logger.debug("Ignoring invalid explicit message timestamp: %r", timestamp)
|
||||
|
||||
# Pre-compute tool call count
|
||||
num_tool_calls = 0
|
||||
if tool_calls is not None:
|
||||
num_tool_calls = len(tool_calls) if isinstance(tool_calls, list) else 1
|
||||
|
||||
def _do(conn):
|
||||
active_lock = conn.execute(
|
||||
@@ -6135,36 +6225,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
and session["end_reason"] == "compression"
|
||||
):
|
||||
raise CompressionSessionClosedError(session_id)
|
||||
cursor = conn.execute(
|
||||
"""INSERT INTO messages (session_id, role, content, tool_call_id,
|
||||
tool_calls, tool_name, effect_disposition, timestamp, token_count, finish_reason,
|
||||
reasoning, reasoning_content, reasoning_details, codex_reasoning_items,
|
||||
codex_message_items, platform_message_id, observed, active, api_content, display_kind, display_metadata)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""",
|
||||
(
|
||||
session_id,
|
||||
role,
|
||||
stored_content,
|
||||
tool_call_id,
|
||||
tool_calls_json,
|
||||
_scrub_surrogates(tool_name),
|
||||
effect_disposition,
|
||||
message_timestamp,
|
||||
token_count,
|
||||
finish_reason,
|
||||
_scrub_surrogates(reasoning),
|
||||
_scrub_surrogates(reasoning_content),
|
||||
reasoning_details_json,
|
||||
codex_items_json,
|
||||
codex_message_items_json,
|
||||
platform_message_id,
|
||||
1 if observed else 0,
|
||||
1,
|
||||
_scrub_surrogates(api_content) if isinstance(api_content, str) else None,
|
||||
_scrub_surrogates(display_kind) if isinstance(display_kind, str) else None,
|
||||
display_metadata_json,
|
||||
),
|
||||
)
|
||||
cursor = conn.execute(self._MESSAGE_INSERT_SQL, row_params)
|
||||
msg_id = cursor.lastrowid
|
||||
|
||||
# Update counters
|
||||
@@ -6190,6 +6251,113 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
_do, patience_s=self._TRANSCRIPT_WRITE_PATIENCE_S
|
||||
)
|
||||
|
||||
def append_messages_batch(
|
||||
self,
|
||||
session_id: str,
|
||||
messages: List[Dict[str, Any]],
|
||||
compression_lock_holder: Optional[str] = None,
|
||||
) -> List[int]:
|
||||
"""Append multiple messages atomically in ONE write transaction.
|
||||
|
||||
``messages`` is a list of dicts whose keys mirror
|
||||
:meth:`append_message`'s keyword arguments (role, content, tool_name,
|
||||
tool_calls, tool_call_id, finish_reason, reasoning, reasoning_content,
|
||||
reasoning_details, codex_reasoning_items, codex_message_items,
|
||||
timestamp, api_content, display_kind, display_metadata, ...).
|
||||
|
||||
A turn-boundary flush writes the whole turn (user + assistant + tool
|
||||
rows, typically 3-8 messages) as one BEGIN IMMEDIATE / commit pair
|
||||
instead of one transaction (and, off WAL, one fsync) per row. Row
|
||||
serialization happens OUTSIDE the transaction via
|
||||
``_prepare_message_row`` — the same helper ``append_message`` uses,
|
||||
so the row shape cannot drift between the two writers.
|
||||
|
||||
Atomicity contract: all rows land or none do (the caller re-flushes
|
||||
unstamped messages on the next attempt). The compression-lock and
|
||||
compression-closed guards from ``append_message`` run once for the
|
||||
batch — same session, same instant. Returns the inserted row IDs in
|
||||
input order.
|
||||
"""
|
||||
if not messages:
|
||||
return []
|
||||
|
||||
prepared: List[tuple] = []
|
||||
total_tool_calls = 0
|
||||
for msg in messages:
|
||||
role = msg.get("role", "unknown")
|
||||
params, num_tc = self._prepare_message_row(
|
||||
session_id=session_id,
|
||||
role=role,
|
||||
content=msg.get("content"),
|
||||
tool_name=msg.get("tool_name"),
|
||||
tool_calls=msg.get("tool_calls"),
|
||||
tool_call_id=msg.get("tool_call_id"),
|
||||
token_count=msg.get("token_count"),
|
||||
finish_reason=msg.get("finish_reason"),
|
||||
reasoning=msg.get("reasoning") if role == "assistant" else None,
|
||||
reasoning_content=msg.get("reasoning_content") if role == "assistant" else None,
|
||||
reasoning_details=msg.get("reasoning_details") if role == "assistant" else None,
|
||||
codex_reasoning_items=msg.get("codex_reasoning_items") if role == "assistant" else None,
|
||||
codex_message_items=msg.get("codex_message_items") if role == "assistant" else None,
|
||||
platform_message_id=msg.get("platform_message_id"),
|
||||
observed=bool(msg.get("observed")),
|
||||
effect_disposition=msg.get("effect_disposition"),
|
||||
timestamp=msg.get("timestamp"),
|
||||
api_content=msg.get("api_content"),
|
||||
display_kind=msg.get("display_kind"),
|
||||
display_metadata=msg.get("display_metadata"),
|
||||
)
|
||||
prepared.append(params)
|
||||
total_tool_calls += num_tc
|
||||
|
||||
def _do(conn):
|
||||
active_lock = conn.execute(
|
||||
"SELECT holder FROM compression_locks "
|
||||
"WHERE session_id = ? AND expires_at > ?",
|
||||
(session_id, time.time()),
|
||||
).fetchone()
|
||||
if (
|
||||
active_lock is not None
|
||||
and active_lock["holder"] != compression_lock_holder
|
||||
):
|
||||
raise SessionCompressionInProgressError(
|
||||
f"Session {session_id!r} is being compressed by another writer"
|
||||
)
|
||||
session = conn.execute(
|
||||
"SELECT ended_at, end_reason FROM sessions WHERE id = ?",
|
||||
(session_id,),
|
||||
).fetchone()
|
||||
if (
|
||||
session is not None
|
||||
and session["ended_at"] is not None
|
||||
and session["end_reason"] == "compression"
|
||||
):
|
||||
raise CompressionSessionClosedError(session_id)
|
||||
|
||||
row_ids: List[int] = []
|
||||
for params in prepared:
|
||||
cursor = conn.execute(self._MESSAGE_INSERT_SQL, params)
|
||||
row_ids.append(cursor.lastrowid)
|
||||
|
||||
# One aggregated counter update for the whole batch.
|
||||
if total_tool_calls > 0:
|
||||
conn.execute(
|
||||
"""UPDATE sessions SET message_count = message_count + ?,
|
||||
tool_call_count = tool_call_count + ? WHERE id = ?""",
|
||||
(len(prepared), total_tool_calls, session_id),
|
||||
)
|
||||
else:
|
||||
conn.execute(
|
||||
"UPDATE sessions SET message_count = message_count + ? WHERE id = ?",
|
||||
(len(prepared), session_id),
|
||||
)
|
||||
return row_ids
|
||||
|
||||
# Same criticality as append_message: this IS the turn's transcript.
|
||||
return self._execute_write(
|
||||
_do, patience_s=self._TRANSCRIPT_WRITE_PATIENCE_S
|
||||
)
|
||||
|
||||
def set_latest_matching_message_display_kind(
|
||||
self, session_id: str, *, role: str, content: str, display_kind: str,
|
||||
display_metadata: Optional[Dict[str, Any]] = None,
|
||||
|
||||
+35
-18
@@ -2096,6 +2096,10 @@ class AIAgent:
|
||||
):
|
||||
_scan_start += 1
|
||||
|
||||
# Collect this flush's new rows and write them in ONE transaction
|
||||
# at the end of the scan (see append_messages_batch).
|
||||
_batch_rows: List[Dict[str, Any]] = []
|
||||
_batch_msgs: List[Dict] = []
|
||||
for _msg_idx in range(_scan_start, len(messages)):
|
||||
msg = messages[_msg_idx]
|
||||
if not isinstance(msg, dict):
|
||||
@@ -2214,33 +2218,46 @@ class AIAgent:
|
||||
]
|
||||
elif isinstance(msg.get("tool_calls"), list):
|
||||
tool_calls_data = msg["tool_calls"]
|
||||
self._session_db.append_message(
|
||||
session_id=self.session_id,
|
||||
role=role,
|
||||
content=content,
|
||||
tool_name=msg.get("tool_name"),
|
||||
tool_calls=tool_calls_data,
|
||||
tool_call_id=msg.get("tool_call_id"),
|
||||
finish_reason=msg.get("finish_reason"),
|
||||
reasoning=msg.get("reasoning") if role == "assistant" else None,
|
||||
reasoning_content=msg.get("reasoning_content") if role == "assistant" else None,
|
||||
reasoning_details=msg.get("reasoning_details") if role == "assistant" else None,
|
||||
codex_reasoning_items=msg.get("codex_reasoning_items") if role == "assistant" else None,
|
||||
codex_message_items=msg.get("codex_message_items") if role == "assistant" else None,
|
||||
timestamp=_row_timestamp,
|
||||
api_content=_row_api_content,
|
||||
display_kind=(
|
||||
_batch_rows.append({
|
||||
"role": role,
|
||||
"content": content,
|
||||
"tool_name": msg.get("tool_name"),
|
||||
"tool_calls": tool_calls_data,
|
||||
"tool_call_id": msg.get("tool_call_id"),
|
||||
"finish_reason": msg.get("finish_reason"),
|
||||
"reasoning": msg.get("reasoning") if role == "assistant" else None,
|
||||
"reasoning_content": msg.get("reasoning_content") if role == "assistant" else None,
|
||||
"reasoning_details": msg.get("reasoning_details") if role == "assistant" else None,
|
||||
"codex_reasoning_items": msg.get("codex_reasoning_items") if role == "assistant" else None,
|
||||
"codex_message_items": msg.get("codex_message_items") if role == "assistant" else None,
|
||||
"timestamp": _row_timestamp,
|
||||
"api_content": _row_api_content,
|
||||
"display_kind": (
|
||||
"hidden"
|
||||
if msg.get(COMPRESSED_SUMMARY_METADATA_KEY)
|
||||
and not msg.get("_compressed_summary_has_user_turn")
|
||||
else msg.get("display_kind")
|
||||
),
|
||||
display_metadata=msg.get("display_metadata"),
|
||||
"display_metadata": msg.get("display_metadata"),
|
||||
})
|
||||
_batch_msgs.append(msg)
|
||||
# One transaction for the whole turn's new rows (typically 3-8
|
||||
# messages): one BEGIN IMMEDIATE / commit — and, off WAL, one
|
||||
# fsync — instead of one per row. All-or-nothing pairs exactly
|
||||
# with the marker stamping below: on failure NO rows landed and
|
||||
# NO markers were stamped, so the next flush re-scans and
|
||||
# re-writes the whole tail (same recovery contract as before,
|
||||
# minus the partial-prefix case that could double-pay counters).
|
||||
if _batch_rows:
|
||||
self._session_db.append_messages_batch(
|
||||
session_id=self.session_id,
|
||||
messages=_batch_rows,
|
||||
compression_lock_holder=getattr(
|
||||
self, "_active_compression_lock_holder", None
|
||||
),
|
||||
)
|
||||
msg[_DB_PERSISTED_MARKER] = True
|
||||
for _written in _batch_msgs:
|
||||
_written[_DB_PERSISTED_MARKER] = True
|
||||
# The intrinsic markers are now the sole source of truth. Reset the
|
||||
# one-shot seed so no id() outlives this flush to alias a message
|
||||
# allocated next turn at a recycled address.
|
||||
|
||||
@@ -146,6 +146,12 @@ def test_parity_persist_bounded_scan():
|
||||
self.rows = []
|
||||
def append_message(self, **kw):
|
||||
self.rows.append({k: copy.deepcopy(v) for k, v in kw.items()})
|
||||
def append_messages_batch(self, session_id, messages, **kw):
|
||||
for m in messages:
|
||||
row = {k: copy.deepcopy(v) for k, v in m.items()}
|
||||
row["session_id"] = session_id
|
||||
self.rows.append(row)
|
||||
return list(range(1, len(messages) + 1))
|
||||
|
||||
def make_agent(bounded):
|
||||
a = ra.AIAgent.__new__(ra.AIAgent)
|
||||
|
||||
@@ -0,0 +1,180 @@
|
||||
"""Tests for SessionDB.append_messages_batch (#23254 salvage).
|
||||
|
||||
The batch writer must be row-shape-identical to append_message (shared
|
||||
_prepare_message_row + _MESSAGE_INSERT_SQL), atomic (all rows or none),
|
||||
and must aggregate the session counters in one UPDATE.
|
||||
"""
|
||||
|
||||
import json
|
||||
import sqlite3
|
||||
|
||||
import pytest
|
||||
|
||||
from hermes_state import (
|
||||
CompressionSessionClosedError,
|
||||
SessionDB,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def db(tmp_path):
|
||||
d = SessionDB(db_path=tmp_path / "state.db")
|
||||
d.create_session("sess-batch", source="cli")
|
||||
yield d
|
||||
d.close()
|
||||
|
||||
|
||||
def _turn_messages():
|
||||
return [
|
||||
{"role": "user", "content": "question"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "let me check",
|
||||
"tool_calls": [{"name": "terminal", "arguments": "{}"}],
|
||||
"reasoning_content": "thinking...",
|
||||
"finish_reason": "tool_calls",
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"content": "tool output",
|
||||
"tool_name": "terminal",
|
||||
"tool_call_id": "call_1",
|
||||
},
|
||||
{"role": "assistant", "content": "answer", "finish_reason": "stop"},
|
||||
]
|
||||
|
||||
|
||||
class TestAppendMessagesBatch:
|
||||
def test_batch_rows_identical_to_single_appends(self, db, tmp_path):
|
||||
"""The batch writer stores the same bytes append_message would."""
|
||||
db2 = SessionDB(db_path=tmp_path / "state2.db")
|
||||
db2.create_session("sess-batch", source="cli")
|
||||
try:
|
||||
msgs = _turn_messages()
|
||||
db.append_messages_batch("sess-batch", msgs)
|
||||
for m in msgs:
|
||||
role = m["role"]
|
||||
db2.append_message(
|
||||
session_id="sess-batch",
|
||||
role=role,
|
||||
content=m.get("content"),
|
||||
tool_name=m.get("tool_name"),
|
||||
tool_calls=m.get("tool_calls"),
|
||||
tool_call_id=m.get("tool_call_id"),
|
||||
finish_reason=m.get("finish_reason"),
|
||||
reasoning_content=(
|
||||
m.get("reasoning_content") if role == "assistant" else None
|
||||
),
|
||||
)
|
||||
cols = (
|
||||
"role, content, tool_call_id, tool_calls, tool_name, "
|
||||
"finish_reason, reasoning_content, observed, active"
|
||||
)
|
||||
rows_a = db._conn.execute(
|
||||
f"SELECT {cols} FROM messages ORDER BY id"
|
||||
).fetchall()
|
||||
rows_b = db2._conn.execute(
|
||||
f"SELECT {cols} FROM messages ORDER BY id"
|
||||
).fetchall()
|
||||
assert [tuple(r) for r in rows_a] == [tuple(r) for r in rows_b]
|
||||
finally:
|
||||
db2.close()
|
||||
|
||||
def test_counters_aggregate_once(self, db):
|
||||
db.append_messages_batch("sess-batch", _turn_messages())
|
||||
row = db._conn.execute(
|
||||
"SELECT message_count, tool_call_count FROM sessions WHERE id = ?",
|
||||
("sess-batch",),
|
||||
).fetchone()
|
||||
assert row["message_count"] == 4
|
||||
assert row["tool_call_count"] == 1
|
||||
|
||||
def test_returns_row_ids_in_input_order(self, db):
|
||||
ids = db.append_messages_batch("sess-batch", _turn_messages())
|
||||
assert ids == sorted(ids)
|
||||
assert len(ids) == 4
|
||||
|
||||
def test_empty_batch_is_noop(self, db):
|
||||
assert db.append_messages_batch("sess-batch", []) == []
|
||||
row = db._conn.execute(
|
||||
"SELECT message_count FROM sessions WHERE id = ?", ("sess-batch",)
|
||||
).fetchone()
|
||||
assert row["message_count"] == 0
|
||||
|
||||
def test_atomicity_all_or_nothing(self, db, monkeypatch):
|
||||
"""A failure mid-batch leaves ZERO rows and untouched counters."""
|
||||
real_execute_write = db._execute_write
|
||||
original_insert = SessionDB._MESSAGE_INSERT_SQL
|
||||
|
||||
calls = {"n": 0}
|
||||
|
||||
def _do_wrapper(fn, **kwargs):
|
||||
def failing(conn):
|
||||
real_conn_execute = conn.execute
|
||||
|
||||
def exec_counting(sql, *args):
|
||||
if sql == original_insert:
|
||||
calls["n"] += 1
|
||||
if calls["n"] == 3:
|
||||
raise sqlite3.OperationalError("boom mid-batch")
|
||||
return real_conn_execute(sql, *args)
|
||||
|
||||
conn.execute = exec_counting
|
||||
try:
|
||||
return fn(conn)
|
||||
finally:
|
||||
conn.execute = real_conn_execute
|
||||
|
||||
return real_execute_write(failing, **kwargs)
|
||||
|
||||
monkeypatch.setattr(db, "_execute_write", _do_wrapper)
|
||||
with pytest.raises(sqlite3.OperationalError):
|
||||
db.append_messages_batch("sess-batch", _turn_messages())
|
||||
monkeypatch.undo()
|
||||
|
||||
count = db._conn.execute("SELECT COUNT(*) FROM messages").fetchone()[0]
|
||||
assert count == 0
|
||||
row = db._conn.execute(
|
||||
"SELECT message_count, tool_call_count FROM sessions WHERE id = ?",
|
||||
("sess-batch",),
|
||||
).fetchone()
|
||||
assert row["message_count"] == 0
|
||||
assert row["tool_call_count"] == 0
|
||||
|
||||
def test_compression_closed_session_rejected(self, db):
|
||||
db._conn.execute(
|
||||
"UPDATE sessions SET ended_at = 1.0, end_reason = 'compression' "
|
||||
"WHERE id = ?",
|
||||
("sess-batch",),
|
||||
)
|
||||
db._conn.commit()
|
||||
with pytest.raises(CompressionSessionClosedError):
|
||||
db.append_messages_batch("sess-batch", _turn_messages())
|
||||
|
||||
def test_multimodal_content_encoded(self, db):
|
||||
msgs = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "look"},
|
||||
{"type": "image_url", "image_url": {"url": "data:x"}},
|
||||
],
|
||||
}
|
||||
]
|
||||
db.append_messages_batch("sess-batch", msgs)
|
||||
raw = db._conn.execute("SELECT content FROM messages").fetchone()[0]
|
||||
# encoded via _encode_content — same sentinel prefix as append_message
|
||||
loaded = db.get_messages("sess-batch")
|
||||
assert loaded, raw
|
||||
|
||||
def test_tool_calls_json_string_not_double_encoded(self, db):
|
||||
msgs = [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "x",
|
||||
"tool_calls": json.dumps([{"name": "t", "arguments": "{}"}]),
|
||||
}
|
||||
]
|
||||
db.append_messages_batch("sess-batch", msgs)
|
||||
raw = db._conn.execute("SELECT tool_calls FROM messages").fetchone()[0]
|
||||
assert json.loads(raw) == [{"name": "t", "arguments": "{}"}]
|
||||
@@ -13,6 +13,12 @@ class _CapturingSessionDB:
|
||||
self.rows.append({"role": role, "content": content})
|
||||
return len(self.rows)
|
||||
|
||||
def append_messages_batch(self, session_id, messages, **kwargs):
|
||||
# Mirror the real batch writer: same rows, one call.
|
||||
for m in messages:
|
||||
self.rows.append({"role": m.get("role"), "content": m.get("content")})
|
||||
return list(range(len(self.rows) - len(messages) + 1, len(self.rows) + 1))
|
||||
|
||||
|
||||
def _agent_with_capturing_db():
|
||||
agent = AIAgent.__new__(AIAgent)
|
||||
|
||||
@@ -227,6 +227,11 @@ def test_flush_guard_clamps_overshooting_cursor():
|
||||
def append_message(self, **kw):
|
||||
self.rows.append(kw)
|
||||
|
||||
def append_messages_batch(self, session_id, messages, **kw):
|
||||
for m in messages:
|
||||
self.rows.append(dict(m, session_id=session_id))
|
||||
return list(range(1, len(messages) + 1))
|
||||
|
||||
agent = _bare_agent()
|
||||
agent._session_db = _DB()
|
||||
agent._session_db_created = True
|
||||
|
||||
@@ -27,7 +27,8 @@ def _make_agent(session_db):
|
||||
|
||||
def test_tool_name_persisted_to_session_db():
|
||||
"""tool_name set by make_tool_result_message must be passed through to
|
||||
append_message so the column is populated on first flush to the session DB."""
|
||||
the batched flush so the column is populated on first write to the
|
||||
session DB."""
|
||||
session_db = MagicMock()
|
||||
agent = _make_agent(session_db)
|
||||
|
||||
@@ -37,9 +38,8 @@ def test_tool_name_persisted_to_session_db():
|
||||
]
|
||||
agent._flush_messages_to_session_db(messages)
|
||||
|
||||
tool_appends = [
|
||||
c for c in session_db.append_message.call_args_list
|
||||
if c.kwargs.get("role") == "tool"
|
||||
]
|
||||
assert len(tool_appends) == 1
|
||||
assert tool_appends[0].kwargs["tool_name"] == "terminal"
|
||||
assert session_db.append_messages_batch.call_count == 1
|
||||
batch = session_db.append_messages_batch.call_args.kwargs["messages"]
|
||||
tool_rows = [m for m in batch if m.get("role") == "tool"]
|
||||
assert len(tool_rows) == 1
|
||||
assert tool_rows[0]["tool_name"] == "terminal"
|
||||
|
||||
Reference in New Issue
Block a user