test_sqlite_store.py 20 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643
  1. import sqlite3
  2. from concurrent.futures import ThreadPoolExecutor
  3. from threading import Barrier
  4. import pytest
  5. from agent_lab.domain.events import ToolCallEvent
  6. from agent_lab.domain.messages import ChatMessage, TokenUsage
  7. from agent_lab.infrastructure.sqlite_store import SQLiteSessionStore
  8. def test_sqlite_store_persists_session_messages_audit_and_usage(tmp_path):
  9. store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
  10. session_id = store.create_session(
  11. title="Debug first turn",
  12. config={
  13. "tool_invocation_mode": "chat_agent_tools",
  14. "chat_agent": {"model": "chat-model"},
  15. },
  16. )
  17. store.start_turn(session_id, turn_index=1, user_message="debug this")
  18. store.append_message(
  19. session_id,
  20. turn_index=1,
  21. message=ChatMessage(role="user", content="debug this"),
  22. )
  23. store.append_message(
  24. session_id,
  25. turn_index=1,
  26. message=ChatMessage(role="assistant", content="answer"),
  27. )
  28. store.append_audit(
  29. session_id,
  30. event="chat_agent_request",
  31. details={"model": "chat-model"},
  32. turn_index=1,
  33. round_index=1,
  34. )
  35. store.append_usage(
  36. session_id,
  37. turn_index=1,
  38. round_index=1,
  39. usage=TokenUsage(
  40. prompt_tokens=3,
  41. completion_tokens=7,
  42. total_tokens=10,
  43. cached_tokens=2,
  44. ),
  45. ttft_ms=120,
  46. elapsed_ms=450,
  47. )
  48. store.complete_turn(session_id, turn_index=1, wall_time_ms=425)
  49. sessions = store.list_sessions()
  50. assert sessions == [
  51. {
  52. "id": session_id,
  53. "title": "Debug first turn",
  54. "created_at": sessions[0]["created_at"],
  55. "updated_at": sessions[0]["updated_at"],
  56. "turn_count": 1,
  57. "total_tokens": 10,
  58. }
  59. ]
  60. assert store.get_session(session_id)["config"]["chat_agent"]["model"] == "chat-model"
  61. assert [message["content"] for message in store.list_messages(session_id)] == [
  62. "debug this",
  63. "answer",
  64. ]
  65. assert store.list_audit_logs(session_id)[0]["event"] == "chat_agent_request"
  66. usage = store.usage_summary(session_id)
  67. assert usage["calls"][0] | {"created_at": "ignored"} == {
  68. "id": usage["calls"][0]["id"],
  69. "turn_index": 1,
  70. "round_index": 1,
  71. "prompt_tokens": 3,
  72. "completion_tokens": 7,
  73. "total_tokens": 10,
  74. "cached_tokens": 2,
  75. "ttft_ms": 120,
  76. "elapsed_ms": 450,
  77. "mode": "chat_agent_tools",
  78. "agent": "chat_agent",
  79. "call_kind": "chat_completion",
  80. "event_id": None,
  81. "event_name": None,
  82. "used_fallback": False,
  83. "tool_latency_ms": None,
  84. "metric_key": None,
  85. "created_at": "ignored",
  86. }
  87. assert usage["turns"] == [
  88. {
  89. "turn_index": 1,
  90. "prompt_tokens": 3,
  91. "completion_tokens": 7,
  92. "total_tokens": 10,
  93. "cached_tokens": 2,
  94. "elapsed_ms": 450,
  95. "call_count": 1,
  96. "fallback_count": 0,
  97. "tool_count": 0,
  98. "turn_wall_time_ms": 425,
  99. }
  100. ]
  101. assert usage["session"] == {
  102. "prompt_tokens": 3,
  103. "completion_tokens": 7,
  104. "total_tokens": 10,
  105. "cached_tokens": 2,
  106. "elapsed_ms": 450,
  107. "call_count": 1,
  108. "fallback_count": 0,
  109. "tool_count": 0,
  110. "turn_wall_time_ms": 425,
  111. }
  112. def test_sqlite_store_persists_typed_calls_without_double_counting_metrics(tmp_path):
  113. store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
  114. session_id = store.create_session(
  115. title="comparable metrics",
  116. config={"tool_invocation_mode": "dual_agent"},
  117. )
  118. store.start_turn(session_id, turn_index=1, user_message="search")
  119. store.append_usage(
  120. session_id,
  121. turn_index=1,
  122. round_index=1,
  123. usage=TokenUsage(prompt_tokens=2, completion_tokens=3, total_tokens=5),
  124. ttft_ms=10,
  125. elapsed_ms=30,
  126. metadata={"metric_key": "chat:1:1"},
  127. )
  128. fallback_metadata = {
  129. "agent": "event_agent",
  130. "call_kind": "argument_fallback",
  131. "event_id": "event-1",
  132. "event_name": "mock_search",
  133. "used_fallback": True,
  134. "metric_key": "fallback:1:1:event-1",
  135. }
  136. for _ in range(2):
  137. store.append_usage(
  138. session_id,
  139. turn_index=1,
  140. round_index=1,
  141. usage=TokenUsage(prompt_tokens=4, completion_tokens=6, total_tokens=10),
  142. ttft_ms=5,
  143. elapsed_ms=20,
  144. metadata=fallback_metadata,
  145. )
  146. store.append_usage(
  147. session_id,
  148. turn_index=1,
  149. round_index=1,
  150. usage=TokenUsage(prompt_tokens=100, completion_tokens=100, total_tokens=200),
  151. ttft_ms=None,
  152. elapsed_ms=80,
  153. metadata={
  154. "agent": "event_kernel",
  155. "call_kind": "tool_execution",
  156. "event_id": "event-1",
  157. "event_name": "mock_search",
  158. "used_fallback": True,
  159. "tool_latency_ms": 80,
  160. "metric_key": "tool:1:1:event-1",
  161. },
  162. )
  163. store.complete_turn(session_id, turn_index=1, wall_time_ms=95)
  164. usage = store.usage_summary(session_id)
  165. assert len(usage["calls"]) == 3
  166. assert [call["mode"] for call in usage["calls"]] == ["dual_agent"] * 3
  167. assert [call["agent"] for call in usage["calls"]] == [
  168. "chat_agent",
  169. "event_agent",
  170. "event_kernel",
  171. ]
  172. assert [call["call_kind"] for call in usage["calls"]] == [
  173. "chat_completion",
  174. "argument_fallback",
  175. "tool_execution",
  176. ]
  177. assert usage["calls"][1]["event_id"] == "event-1"
  178. assert usage["calls"][1]["event_name"] == "mock_search"
  179. assert usage["calls"][1]["used_fallback"] is True
  180. assert usage["calls"][2]["tool_latency_ms"] == 80
  181. assert usage["calls"][2]["total_tokens"] == 0
  182. assert usage["turns"] == [
  183. {
  184. "turn_index": 1,
  185. "prompt_tokens": 6,
  186. "completion_tokens": 9,
  187. "total_tokens": 15,
  188. "cached_tokens": 0,
  189. "elapsed_ms": 130,
  190. "call_count": 3,
  191. "fallback_count": 1,
  192. "tool_count": 1,
  193. "turn_wall_time_ms": 95,
  194. }
  195. ]
  196. assert usage["session"] == {
  197. "prompt_tokens": 6,
  198. "completion_tokens": 9,
  199. "total_tokens": 15,
  200. "cached_tokens": 0,
  201. "elapsed_ms": 130,
  202. "call_count": 3,
  203. "fallback_count": 1,
  204. "tool_count": 1,
  205. "turn_wall_time_ms": 95,
  206. }
  207. assert store.list_sessions()[0]["total_tokens"] == 15
  208. def test_sqlite_store_metric_keys_are_unique_within_each_session(tmp_path):
  209. store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
  210. for session_id in ("session-a", "session-b"):
  211. store.create_session(title=session_id, config={}, session_id=session_id)
  212. store.start_turn(session_id, turn_index=1, user_message="test")
  213. for _ in range(2):
  214. store.append_usage(
  215. session_id,
  216. turn_index=1,
  217. round_index=1,
  218. usage=TokenUsage(total_tokens=1),
  219. ttft_ms=None,
  220. elapsed_ms=1,
  221. metadata={"metric_key": "chat:1:1"},
  222. )
  223. assert len(store.usage_summary("session-a")["calls"]) == 1
  224. assert len(store.usage_summary("session-b")["calls"]) == 1
  225. def test_sqlite_store_migrates_legacy_usage_and_turn_metrics_additively(tmp_path):
  226. database_path = tmp_path / "legacy-usage.sqlite3"
  227. connection = sqlite3.connect(database_path)
  228. connection.executescript(
  229. """
  230. CREATE TABLE sessions (
  231. id TEXT PRIMARY KEY,
  232. title TEXT NOT NULL,
  233. created_at TEXT NOT NULL,
  234. updated_at TEXT NOT NULL,
  235. config_json TEXT NOT NULL
  236. );
  237. CREATE TABLE turns (
  238. session_id TEXT NOT NULL,
  239. turn_index INTEGER NOT NULL,
  240. user_message TEXT NOT NULL,
  241. started_at TEXT NOT NULL,
  242. completed_at TEXT,
  243. PRIMARY KEY (session_id, turn_index)
  244. );
  245. CREATE TABLE usage_stats (
  246. id INTEGER PRIMARY KEY AUTOINCREMENT,
  247. session_id TEXT NOT NULL,
  248. turn_index INTEGER NOT NULL,
  249. round_index INTEGER NOT NULL,
  250. prompt_tokens INTEGER NOT NULL,
  251. completion_tokens INTEGER NOT NULL,
  252. total_tokens INTEGER NOT NULL,
  253. cached_tokens INTEGER NOT NULL,
  254. ttft_ms INTEGER,
  255. elapsed_ms INTEGER NOT NULL,
  256. created_at TEXT NOT NULL
  257. );
  258. INSERT INTO sessions VALUES (
  259. 'legacy', 'Legacy', 'now', 'now',
  260. '{"tool_invocation_mode":"chat_agent_tools"}'
  261. );
  262. INSERT INTO turns VALUES ('legacy', 1, 'old', 'now', 'now');
  263. INSERT INTO usage_stats (
  264. session_id, turn_index, round_index, prompt_tokens,
  265. completion_tokens, total_tokens, cached_tokens,
  266. ttft_ms, elapsed_ms, created_at
  267. ) VALUES ('legacy', 1, 1, 1, 2, 3, 0, 4, 5, 'now');
  268. """
  269. )
  270. connection.commit()
  271. connection.close()
  272. store = SQLiteSessionStore(database_path)
  273. usage = store.usage_summary("legacy")
  274. assert usage["calls"][0]["agent"] == "chat_agent"
  275. assert usage["calls"][0]["call_kind"] == "chat_completion"
  276. assert usage["calls"][0]["mode"] == "chat_agent_tools"
  277. assert usage["calls"][0]["used_fallback"] is False
  278. assert usage["turns"][0]["turn_wall_time_ms"] is None
  279. migrated = sqlite3.connect(database_path)
  280. usage_columns = {row[1] for row in migrated.execute("PRAGMA table_info(usage_stats)")}
  281. turn_columns = {row[1] for row in migrated.execute("PRAGMA table_info(turns)")}
  282. indexes = {row[1] for row in migrated.execute("PRAGMA index_list(usage_stats)")}
  283. metric_index_sql = migrated.execute(
  284. "SELECT sql FROM sqlite_master WHERE name = 'usage_stats_metric_key_uq'"
  285. ).fetchone()[0]
  286. migrated.close()
  287. assert {
  288. "agent",
  289. "call_kind",
  290. "event_id",
  291. "event_name",
  292. "used_fallback",
  293. "tool_latency_ms",
  294. "metric_key",
  295. } <= usage_columns
  296. assert "wall_time_ms" in turn_columns
  297. assert "usage_stats_metric_key_uq" in indexes
  298. assert "session_id, metric_key" in metric_index_sql
  299. def test_sqlite_store_session_list_totals_are_not_multiplied_by_turns(tmp_path):
  300. store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
  301. session_id = store.create_session(title="two turns", config={})
  302. for turn_index, tokens in [(1, 2), (2, 4)]:
  303. store.start_turn(
  304. session_id,
  305. turn_index=turn_index,
  306. user_message=f"turn {turn_index}",
  307. )
  308. store.append_usage(
  309. session_id,
  310. turn_index=turn_index,
  311. round_index=1,
  312. usage=TokenUsage(total_tokens=tokens),
  313. ttft_ms=None,
  314. elapsed_ms=10,
  315. )
  316. assert store.list_sessions()[0]["turn_count"] == 2
  317. assert store.list_sessions()[0]["total_tokens"] == 6
  318. def test_sqlite_store_roundtrips_assistant_tool_calls_and_tool_replies(tmp_path):
  319. store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
  320. session_id = store.create_session(title="direct tools", config={})
  321. store.start_turn(session_id, turn_index=1, user_message="search")
  322. tool_calls = [
  323. ToolCallEvent(
  324. id="call-1",
  325. name="mock_search",
  326. arguments={"query": "agent lab"},
  327. raw_arguments='{"query":"agent lab"}',
  328. ),
  329. ToolCallEvent(
  330. id="call-2",
  331. name="mock_ticket",
  332. arguments={"title": "follow up"},
  333. raw_arguments='{"title":"follow up"}',
  334. ),
  335. ]
  336. store.append_message(
  337. session_id,
  338. turn_index=1,
  339. message=ChatMessage(
  340. role="assistant",
  341. content="",
  342. tool_calls=tool_calls,
  343. ),
  344. )
  345. store.append_message(
  346. session_id,
  347. turn_index=1,
  348. message=ChatMessage(
  349. role="tool",
  350. content='{"ok":true}',
  351. name="mock_search",
  352. tool_call_id="call-1",
  353. ),
  354. )
  355. messages = store.list_messages(session_id)
  356. assert messages[0]["tool_calls"] == [call.model_dump() for call in tool_calls]
  357. assert messages[1]["tool_calls"] == []
  358. assert messages[1]["tool_call_id"] == "call-1"
  359. def test_sqlite_store_migrates_legacy_messages_with_empty_tool_calls(tmp_path):
  360. database_path = tmp_path / "legacy.sqlite3"
  361. connection = sqlite3.connect(database_path)
  362. connection.execute(
  363. """
  364. CREATE TABLE messages (
  365. id INTEGER PRIMARY KEY AUTOINCREMENT,
  366. session_id TEXT NOT NULL,
  367. turn_index INTEGER NOT NULL,
  368. role TEXT NOT NULL,
  369. content TEXT NOT NULL,
  370. name TEXT,
  371. tool_call_id TEXT,
  372. created_at TEXT NOT NULL
  373. )
  374. """
  375. )
  376. connection.execute(
  377. """
  378. INSERT INTO messages
  379. (session_id, turn_index, role, content, name, tool_call_id, created_at)
  380. VALUES ('legacy', 1, 'assistant', 'old answer', NULL, NULL, 'now')
  381. """
  382. )
  383. connection.commit()
  384. connection.close()
  385. store = SQLiteSessionStore(database_path)
  386. assert store.list_messages("legacy")[0]["tool_calls"] == []
  387. migrated = sqlite3.connect(database_path).execute(
  388. "PRAGMA table_info(messages)"
  389. ).fetchall()
  390. tool_calls_column = next(column for column in migrated if column[1] == "tool_calls_json")
  391. assert tool_calls_column[3] == 1
  392. assert tool_calls_column[4] == "'[]'"
  393. @pytest.mark.parametrize(
  394. ("persisted_config", "requested_mode", "persisted_mode"),
  395. [
  396. ({"tool_invocation_mode": "chat_agent_tools"}, "dual_agent", "chat_agent_tools"),
  397. ({}, "chat_agent_tools", "dual_agent"),
  398. ],
  399. )
  400. def test_sqlite_store_rejects_session_mode_mismatch_without_mutation(
  401. tmp_path,
  402. persisted_config,
  403. requested_mode,
  404. persisted_mode,
  405. ):
  406. store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
  407. session_id = store.create_session(
  408. title="mode locked",
  409. config=persisted_config,
  410. session_id="mode-locked",
  411. )
  412. store.start_turn(session_id, turn_index=1, user_message="past user")
  413. store.append_message(
  414. session_id,
  415. turn_index=1,
  416. message=ChatMessage(role="user", content="past user"),
  417. )
  418. before_session = store.get_session(session_id)
  419. before_messages = store.list_messages(session_id)
  420. with pytest.raises(
  421. ValueError,
  422. match=(
  423. "session tool invocation mode mismatch: "
  424. f"persisted={persisted_mode} requested={requested_mode}"
  425. ),
  426. ):
  427. store.ensure_session(
  428. session_id,
  429. title="ignored",
  430. config={"tool_invocation_mode": requested_mode, "changed": True},
  431. )
  432. assert store.get_session(session_id) == before_session
  433. assert store.list_messages(session_id) == before_messages
  434. def test_sqlite_store_append_messages_rolls_back_the_whole_protocol_group(tmp_path):
  435. database_path = tmp_path / "agent_lab.sqlite3"
  436. store = SQLiteSessionStore(database_path)
  437. session_id = store.create_session(title="atomic group", config={})
  438. store.start_turn(session_id, turn_index=1, user_message="use tool")
  439. connection = sqlite3.connect(database_path)
  440. connection.execute(
  441. """
  442. CREATE TRIGGER reject_tool_message
  443. BEFORE INSERT ON messages
  444. WHEN NEW.role = 'tool'
  445. BEGIN
  446. SELECT RAISE(ABORT, 'tool blocked');
  447. END
  448. """
  449. )
  450. connection.commit()
  451. connection.close()
  452. tool_call = ToolCallEvent(
  453. id="atomic-call",
  454. name="mock_search",
  455. arguments={"query": "atomic"},
  456. raw_arguments='{"query":"atomic"}',
  457. )
  458. with pytest.raises(sqlite3.IntegrityError, match="tool blocked"):
  459. store.append_messages(
  460. session_id,
  461. turn_index=1,
  462. messages=[
  463. ChatMessage(role="assistant", content="", tool_calls=[tool_call]),
  464. ChatMessage(
  465. role="tool",
  466. content='{"ok":true}',
  467. name="mock_search",
  468. tool_call_id="atomic-call",
  469. ),
  470. ],
  471. )
  472. assert store.list_messages(session_id) == []
  473. def test_sqlite_store_migration_is_safe_across_concurrent_initialization(
  474. tmp_path,
  475. monkeypatch,
  476. ):
  477. database_path = tmp_path / "concurrent-legacy.sqlite3"
  478. connection = sqlite3.connect(database_path)
  479. connection.execute(
  480. """
  481. CREATE TABLE messages (
  482. id INTEGER PRIMARY KEY AUTOINCREMENT,
  483. session_id TEXT NOT NULL,
  484. turn_index INTEGER NOT NULL,
  485. role TEXT NOT NULL,
  486. content TEXT NOT NULL,
  487. name TEXT,
  488. tool_call_id TEXT,
  489. created_at TEXT NOT NULL
  490. )
  491. """
  492. )
  493. connection.execute(
  494. """
  495. CREATE TABLE turns (
  496. session_id TEXT NOT NULL,
  497. turn_index INTEGER NOT NULL,
  498. user_message TEXT NOT NULL,
  499. started_at TEXT NOT NULL,
  500. completed_at TEXT,
  501. PRIMARY KEY (session_id, turn_index)
  502. )
  503. """
  504. )
  505. connection.execute(
  506. """
  507. CREATE TABLE usage_stats (
  508. id INTEGER PRIMARY KEY AUTOINCREMENT,
  509. session_id TEXT NOT NULL,
  510. turn_index INTEGER NOT NULL,
  511. round_index INTEGER NOT NULL,
  512. prompt_tokens INTEGER NOT NULL,
  513. completion_tokens INTEGER NOT NULL,
  514. total_tokens INTEGER NOT NULL,
  515. cached_tokens INTEGER NOT NULL,
  516. ttft_ms INTEGER,
  517. elapsed_ms INTEGER NOT NULL,
  518. created_at TEXT NOT NULL
  519. )
  520. """
  521. )
  522. connection.execute(
  523. """
  524. INSERT INTO messages
  525. (session_id, turn_index, role, content, name, tool_call_id, created_at)
  526. VALUES ('legacy', 1, 'assistant', 'visible', NULL, NULL, 'now')
  527. """
  528. )
  529. connection.commit()
  530. connection.close()
  531. worker_count = 8
  532. barrier = Barrier(worker_count)
  533. real_connect = sqlite3.connect
  534. class CoordinatedConnection:
  535. def __init__(self, connection):
  536. self._connection = connection
  537. @property
  538. def row_factory(self):
  539. return self._connection.row_factory
  540. @row_factory.setter
  541. def row_factory(self, value):
  542. self._connection.row_factory = value
  543. def execute(self, sql, *args):
  544. cursor = self._connection.execute(sql, *args)
  545. if (
  546. sql.strip() == "PRAGMA table_info(messages)"
  547. and not self._connection.in_transaction
  548. ):
  549. rows = cursor.fetchall()
  550. barrier.wait(timeout=5)
  551. return rows
  552. return cursor
  553. def __getattr__(self, name):
  554. return getattr(self._connection, name)
  555. def coordinated_connect(*args, **kwargs):
  556. return CoordinatedConnection(real_connect(*args, **kwargs))
  557. monkeypatch.setattr(
  558. "agent_lab.infrastructure.sqlite_store.sqlite3.connect",
  559. coordinated_connect,
  560. )
  561. def initialize_store() -> list[dict]:
  562. store = SQLiteSessionStore(database_path)
  563. barrier.wait()
  564. return store.list_messages("legacy")
  565. with ThreadPoolExecutor(max_workers=worker_count) as executor:
  566. results = list(executor.map(lambda _: initialize_store(), range(worker_count)))
  567. assert all(messages[0]["tool_calls"] == [] for messages in results)
  568. columns = real_connect(database_path).execute(
  569. "PRAGMA table_info(messages)"
  570. ).fetchall()
  571. assert [column[1] for column in columns].count("tool_calls_json") == 1
  572. migrated = real_connect(database_path)
  573. turn_columns = [row[1] for row in migrated.execute("PRAGMA table_info(turns)")]
  574. usage_columns = [
  575. row[1] for row in migrated.execute("PRAGMA table_info(usage_stats)")
  576. ]
  577. metric_indexes = [
  578. row[1] for row in migrated.execute("PRAGMA index_list(usage_stats)")
  579. ]
  580. migrated.close()
  581. assert turn_columns.count("wall_time_ms") == 1
  582. for column_name in (
  583. "agent",
  584. "call_kind",
  585. "event_id",
  586. "event_name",
  587. "used_fallback",
  588. "tool_latency_ms",
  589. "metric_key",
  590. ):
  591. assert usage_columns.count(column_name) == 1
  592. assert metric_indexes.count("usage_stats_metric_key_uq") == 1