Files
hermes-agent/tests/agent/test_cursor_optimizations_parity.py
devsart95 06ae5b6faa 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.
2026-08-03 20:43:38 +05:30

276 lines
11 KiB
Python

"""Byte-parity + benchmark harness for the per-iteration cursor optimizations.
Drives the pure functions directly (no AIAgent):
1. sanitize_tool_call_arguments (with/without cursor)
2. estimate_messages_tokens_rough (memoized) vs a reference reimplementation
3. _flush_messages_to_session_db bounded scan — simulated via a stub agent
Run: HERMES worktree venv python parity_harness.py
"""
import copy
import json
import random
import statistics
import sys
import time
sys.path.insert(0, ".")
random.seed(1234)
UNI = "日本語テキスト🎉 café Ω ≈ 中文字符串"
def build_history(n):
"""Synthetic conversation: user/assistant/tool cycles, malformed args, unicode."""
msgs = []
i = 0
while len(msgs) < n:
msgs.append({"role": "user", "content": f"question {i} {UNI} " + "x" * random.randint(10, 400)})
if i % 3 == 0:
args = json.dumps({"q": f"val {i}", "u": UNI, "n": i})
if i % 9 == 0:
args = '{"broken": tru' # malformed
elif i % 6 == 0:
args = "" # empty
msgs.append({
"role": "assistant", "content": "",
"tool_calls": [{"id": f"call_{i}", "type": "function",
"function": {"name": "web_search", "arguments": args}}],
})
msgs.append({"role": "tool", "tool_call_id": f"call_{i}",
"name": "web_search", "content": f"result {i} {UNI}"})
else:
msgs.append({"role": "assistant", "content": f"answer {i} " + "y" * random.randint(10, 600),
"reasoning_content": f"thinking {i}"})
if i % 7 == 0 and msgs:
msgs[-1]["content"] = [{"type": "text", "text": f"part {i}"},
{"type": "image_url", "image_url": {"url": "data:image/png;base64,AAAA"}}]
i += 1
return msgs[:n]
# ---------- reference (pre-optimization) implementations ----------
from agent.model_metadata import (
estimate_messages_tokens_rough,
_estimate_message_tokens_without_images,
_count_image_tokens,
_MSG_TOKENS_CACHE,
)
def estimate_messages_tokens_rough_OLD(messages):
_IMAGE_TOKEN_COST = 1500
text_tokens = 0
image_tokens = 0
for msg in messages:
text_tokens += _estimate_message_tokens_without_images(msg)
image_tokens += _count_image_tokens(msg, _IMAGE_TOKEN_COST)
return text_tokens + image_tokens
from agent.agent_runtime_helpers import sanitize_tool_call_arguments
def simulate_compression(msgs):
"""Rewrite the middle of the history with fresh dict copies + a summary."""
head, mid, tail = msgs[:2], msgs[2:-6], msgs[-6:]
summary = {"role": "user", "content": "SUMMARY OF DROPPED CONTEXT " + UNI}
new = [dict(m) if isinstance(m, dict) else m for m in head]
new.append(summary)
new.extend(dict(m) if isinstance(m, dict) else m for m in tail)
msgs[:] = new
def test_parity_sanitize_cursor():
print("=== parity: sanitize_tool_call_arguments cursor ===")
for n in (50, 200, 500):
base = build_history(n)
old_list = copy.deepcopy(base)
new_list = copy.deepcopy(base)
cursor = {}
for iteration in range(3):
r_old = sanitize_tool_call_arguments(old_list)
r_new = sanitize_tool_call_arguments(new_list, cursor=cursor)
assert r_old == r_new, (n, iteration, r_old, r_new)
assert old_list == new_list, f"list mismatch n={n} it={iteration}"
# append a new exchange (one malformed) between iterations
for lst in (old_list, new_list):
lst.append({"role": "assistant", "content": "",
"tool_calls": [{"id": f"c{iteration}", "type": "function",
"function": {"name": "t", "arguments": '{"bad": '}}]})
lst.append({"role": "tool", "tool_call_id": f"c{iteration}", "content": "ok"})
if iteration == 1:
simulate_compression(old_list)
simulate_compression(new_list)
# after mutations, one more full compare
r_old = sanitize_tool_call_arguments(old_list)
r_new = sanitize_tool_call_arguments(new_list, cursor=cursor)
assert r_old == r_new and old_list == new_list
print(f" n={n}: OK (element-wise equal across 3 iterations + compression)")
def test_parity_token_memo():
print("=== parity: estimate_messages_tokens_rough memo ===")
for n in (50, 200, 500):
msgs = build_history(n)
_MSG_TOKENS_CACHE.clear()
for iteration in range(3):
# simulate api_messages copies each iteration (shallow copies)
api = [m.copy() for m in msgs]
old = estimate_messages_tokens_rough_OLD(api)
new = estimate_messages_tokens_rough(api)
assert old == new, (n, iteration, old, new)
msgs.append({"role": "user", "content": f"followup {iteration} {UNI}"})
if iteration == 1:
simulate_compression(msgs)
# mutate a string in place-ish: replace content of an existing dict
msgs[0]["content"] = "EDITED " + UNI
api = [m.copy() for m in msgs]
assert estimate_messages_tokens_rough_OLD(api) == estimate_messages_tokens_rough(api)
# odd types fall through the memo
weird = [{"role": "user", "content": {"_multimodal": True, "text_summary": "s"}},
{"role": "user", "content": None}, "not-a-dict",
{"role": "tool", "content": [{"type": "text", "text": UNI}, "raw"], "meta": (1, 2)}]
assert estimate_messages_tokens_rough_OLD(weird) == estimate_messages_tokens_rough(weird)
print(f" n={n}: OK (equal across 3 iterations + compression + in-place edit + odd types)")
def test_parity_persist_bounded_scan():
print("=== parity: _flush_messages_to_session_db bounded scan ===")
import run_agent as ra
class FakeDB:
def __init__(self):
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)
a.session_id = "s1"
a._session_db = FakeDB()
a._session_db_created = True
a._last_flushed_db_idx = 0
a._flushed_db_message_ids = set()
a._persist_disabled = False
a._session_persist_lock = None
if not bounded:
# neutralize the cursor: force full scan every time
a._db_flush_scan_prefix = None
return a
for n in (50, 200, 500):
base = build_history(n)
la, lb = copy.deepcopy(base), copy.deepcopy(base)
A, B = make_agent(False), make_agent(True)
for iteration in range(3):
A._db_flush_scan_prefix = None # baseline: always full scan
ra_ok = A._flush_messages_to_session_db_unlocked(la, None)
rb_ok = B._flush_messages_to_session_db_unlocked(lb, None)
assert ra_ok is True and rb_ok is True
assert A._session_db.rows == B._session_db.rows, f"rows diverge n={n} it={iteration}"
assert la == lb
for lst in (la, lb):
lst.append({"role": "user", "content": f"turn {iteration} {UNI}"})
lst.append({"role": "assistant", "content": f"reply {iteration}",
"_empty_recovery_synthetic": iteration == 0}) # scaffolding once
if iteration == 1:
# compression-style rewrite: fresh copies without markers
for lst in (la, lb):
head = [dict(m) for m in lst[:3]]
for m in head:
m.pop(ra._DB_PERSISTED_MARKER, None)
tail = [dict(m) for m in lst[-4:]]
for m in tail:
m.pop(ra._DB_PERSISTED_MARKER, None)
lst[:] = head + [{"role": "user", "content": "SUMMARY"}] + tail
A._db_flush_scan_prefix = None
A._flush_messages_to_session_db_unlocked(la, None)
B._flush_messages_to_session_db_unlocked(lb, None)
assert A._session_db.rows == B._session_db.rows and la == lb
print(f" n={n}: OK (identical DB rows + marker stamps across 3 flushes + compression rewrite)")
def bench():
print("=== benchmarks (median of 5, per call) ===")
def timeit(fn, reps=5):
ts = []
for _ in range(reps):
t0 = time.perf_counter()
fn()
ts.append(time.perf_counter() - t0)
return statistics.median(ts) * 1e3 # ms
for n in (50, 200, 500):
msgs = build_history(n)
sanitize_tool_call_arguments(msgs) # settle repairs first
# sanitize: old (no cursor) vs new (warm cursor)
old_ms = timeit(lambda: sanitize_tool_call_arguments(msgs))
cur = {}
sanitize_tool_call_arguments(msgs, cursor=cur)
new_ms = timeit(lambda: sanitize_tool_call_arguments(msgs, cursor=cur))
# tokens: old walk vs warm memo (on fresh shallow copies, like api_messages)
api = [m.copy() for m in msgs]
told = timeit(lambda: estimate_messages_tokens_rough_OLD([m.copy() for m in msgs]))
_MSG_TOKENS_CACHE.clear()
estimate_messages_tokens_rough([m.copy() for m in msgs]) # warm
tnew = timeit(lambda: estimate_messages_tokens_rough([m.copy() for m in msgs]))
# persist scan: fully-flushed list, old full walk vs bounded skip
import run_agent as ra
flushed = copy.deepcopy(msgs)
for m in flushed:
if isinstance(m, dict):
m[ra._DB_PERSISTED_MARKER] = True
def old_scan():
for _idx, m in enumerate(flushed):
if not isinstance(m, dict):
continue
if ra._is_ephemeral_scaffolding(m):
continue
if m.get(ra._DB_PERSISTED_MARKER):
continue
prefix = flushed[:]
def new_scan():
s = 0
lim = min(len(prefix), len(flushed))
while s < lim and flushed[s] is prefix[s]:
s += 1
for _idx in range(s, len(flushed)):
m = flushed[_idx]
if not isinstance(m, dict):
continue
if ra._is_ephemeral_scaffolding(m):
continue
if m.get(ra._DB_PERSISTED_MARKER):
continue
pold = timeit(old_scan)
pnew = timeit(new_scan)
print(f" n={n:3d}: sanitize {old_ms:.3f}ms -> {new_ms:.3f}ms | "
f"tokens {told:.3f}ms -> {tnew:.3f}ms | "
f"persist-scan {pold*1000:.1f}us -> {pnew*1000:.1f}us")
if __name__ == "__main__":
test_parity_sanitize_cursor()
test_parity_token_memo()
test_parity_persist_bounded_scan()
bench()
print("ALL PARITY CHECKS PASSED")