test_sqlite_store.py 23 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735
  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_deduplicates_metric_keys_before_creating_unique_index(
  300. tmp_path,
  301. ):
  302. database_path = tmp_path / "duplicate-metric-keys.sqlite3"
  303. connection = sqlite3.connect(database_path)
  304. connection.execute(
  305. """
  306. CREATE TABLE usage_stats (
  307. id INTEGER PRIMARY KEY AUTOINCREMENT,
  308. session_id TEXT NOT NULL,
  309. turn_index INTEGER NOT NULL,
  310. round_index INTEGER NOT NULL,
  311. prompt_tokens INTEGER NOT NULL,
  312. completion_tokens INTEGER NOT NULL,
  313. total_tokens INTEGER NOT NULL,
  314. cached_tokens INTEGER NOT NULL,
  315. ttft_ms INTEGER,
  316. elapsed_ms INTEGER NOT NULL,
  317. metric_key TEXT,
  318. created_at TEXT NOT NULL
  319. )
  320. """
  321. )
  322. connection.executemany(
  323. """
  324. INSERT INTO usage_stats (
  325. id, session_id, turn_index, round_index, prompt_tokens,
  326. completion_tokens, total_tokens, cached_tokens, ttft_ms,
  327. elapsed_ms, metric_key, created_at
  328. ) VALUES (?, 'legacy', 1, 1, 0, 0, 0, 0, NULL, 1, ?, 'now')
  329. """,
  330. [
  331. (7, "chat:1:1"),
  332. (11, "chat:1:1"),
  333. (13, None),
  334. (17, None),
  335. (19, None),
  336. ],
  337. )
  338. connection.commit()
  339. connection.close()
  340. SQLiteSessionStore(database_path).list_sessions()
  341. migrated = sqlite3.connect(database_path)
  342. rows = migrated.execute(
  343. """
  344. SELECT id, session_id, metric_key
  345. FROM usage_stats
  346. ORDER BY id
  347. """
  348. ).fetchall()
  349. indexes = [row[1] for row in migrated.execute("PRAGMA index_list(usage_stats)")]
  350. migrated.close()
  351. assert rows == [
  352. (7, "legacy", "chat:1:1"),
  353. (13, "legacy", None),
  354. (17, "legacy", None),
  355. (19, "legacy", None),
  356. ]
  357. assert indexes.count("usage_stats_metric_key_uq") == 1
  358. def test_sqlite_store_session_list_totals_are_not_multiplied_by_turns(tmp_path):
  359. store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
  360. session_id = store.create_session(title="two turns", config={})
  361. for turn_index, tokens in [(1, 2), (2, 4)]:
  362. store.start_turn(
  363. session_id,
  364. turn_index=turn_index,
  365. user_message=f"turn {turn_index}",
  366. )
  367. store.append_usage(
  368. session_id,
  369. turn_index=turn_index,
  370. round_index=1,
  371. usage=TokenUsage(total_tokens=tokens),
  372. ttft_ms=None,
  373. elapsed_ms=10,
  374. )
  375. assert store.list_sessions()[0]["turn_count"] == 2
  376. assert store.list_sessions()[0]["total_tokens"] == 6
  377. def test_sqlite_store_roundtrips_assistant_tool_calls_and_tool_replies(tmp_path):
  378. store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
  379. session_id = store.create_session(title="direct tools", config={})
  380. store.start_turn(session_id, turn_index=1, user_message="search")
  381. tool_calls = [
  382. ToolCallEvent(
  383. id="call-1",
  384. name="mock_search",
  385. arguments={"query": "agent lab"},
  386. raw_arguments='{"query":"agent lab"}',
  387. ),
  388. ToolCallEvent(
  389. id="call-2",
  390. name="mock_ticket",
  391. arguments={"title": "follow up"},
  392. raw_arguments='{"title":"follow up"}',
  393. ),
  394. ]
  395. store.append_message(
  396. session_id,
  397. turn_index=1,
  398. message=ChatMessage(
  399. role="assistant",
  400. content="",
  401. tool_calls=tool_calls,
  402. ),
  403. )
  404. store.append_message(
  405. session_id,
  406. turn_index=1,
  407. message=ChatMessage(
  408. role="tool",
  409. content='{"ok":true}',
  410. name="mock_search",
  411. tool_call_id="call-1",
  412. ),
  413. )
  414. messages = store.list_messages(session_id)
  415. assert messages[0]["tool_calls"] == [call.model_dump() for call in tool_calls]
  416. assert messages[1]["tool_calls"] == []
  417. assert messages[1]["tool_call_id"] == "call-1"
  418. def test_sqlite_store_migrates_legacy_messages_with_empty_tool_calls(tmp_path):
  419. database_path = tmp_path / "legacy.sqlite3"
  420. connection = sqlite3.connect(database_path)
  421. connection.execute(
  422. """
  423. CREATE TABLE messages (
  424. id INTEGER PRIMARY KEY AUTOINCREMENT,
  425. session_id TEXT NOT NULL,
  426. turn_index INTEGER NOT NULL,
  427. role TEXT NOT NULL,
  428. content TEXT NOT NULL,
  429. name TEXT,
  430. tool_call_id TEXT,
  431. created_at TEXT NOT NULL
  432. )
  433. """
  434. )
  435. connection.execute(
  436. """
  437. INSERT INTO messages
  438. (session_id, turn_index, role, content, name, tool_call_id, created_at)
  439. VALUES ('legacy', 1, 'assistant', 'old answer', NULL, NULL, 'now')
  440. """
  441. )
  442. connection.commit()
  443. connection.close()
  444. store = SQLiteSessionStore(database_path)
  445. assert store.list_messages("legacy")[0]["tool_calls"] == []
  446. migrated = sqlite3.connect(database_path).execute(
  447. "PRAGMA table_info(messages)"
  448. ).fetchall()
  449. tool_calls_column = next(column for column in migrated if column[1] == "tool_calls_json")
  450. assert tool_calls_column[3] == 1
  451. assert tool_calls_column[4] == "'[]'"
  452. @pytest.mark.parametrize(
  453. ("persisted_config", "requested_mode", "persisted_mode"),
  454. [
  455. ({"tool_invocation_mode": "chat_agent_tools"}, "dual_agent", "chat_agent_tools"),
  456. ({}, "chat_agent_tools", "dual_agent"),
  457. ],
  458. )
  459. def test_sqlite_store_rejects_session_mode_mismatch_without_mutation(
  460. tmp_path,
  461. persisted_config,
  462. requested_mode,
  463. persisted_mode,
  464. ):
  465. store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
  466. session_id = store.create_session(
  467. title="mode locked",
  468. config=persisted_config,
  469. session_id="mode-locked",
  470. )
  471. store.start_turn(session_id, turn_index=1, user_message="past user")
  472. store.append_message(
  473. session_id,
  474. turn_index=1,
  475. message=ChatMessage(role="user", content="past user"),
  476. )
  477. before_session = store.get_session(session_id)
  478. before_messages = store.list_messages(session_id)
  479. with pytest.raises(
  480. ValueError,
  481. match=(
  482. "session tool invocation mode mismatch: "
  483. f"persisted={persisted_mode} requested={requested_mode}"
  484. ),
  485. ):
  486. store.ensure_session(
  487. session_id,
  488. title="ignored",
  489. config={"tool_invocation_mode": requested_mode, "changed": True},
  490. )
  491. assert store.get_session(session_id) == before_session
  492. assert store.list_messages(session_id) == before_messages
  493. def test_sqlite_store_append_messages_rolls_back_the_whole_protocol_group(tmp_path):
  494. database_path = tmp_path / "agent_lab.sqlite3"
  495. store = SQLiteSessionStore(database_path)
  496. session_id = store.create_session(title="atomic group", config={})
  497. store.start_turn(session_id, turn_index=1, user_message="use tool")
  498. connection = sqlite3.connect(database_path)
  499. connection.execute(
  500. """
  501. CREATE TRIGGER reject_tool_message
  502. BEFORE INSERT ON messages
  503. WHEN NEW.role = 'tool'
  504. BEGIN
  505. SELECT RAISE(ABORT, 'tool blocked');
  506. END
  507. """
  508. )
  509. connection.commit()
  510. connection.close()
  511. tool_call = ToolCallEvent(
  512. id="atomic-call",
  513. name="mock_search",
  514. arguments={"query": "atomic"},
  515. raw_arguments='{"query":"atomic"}',
  516. )
  517. with pytest.raises(sqlite3.IntegrityError, match="tool blocked"):
  518. store.append_messages(
  519. session_id,
  520. turn_index=1,
  521. messages=[
  522. ChatMessage(role="assistant", content="", tool_calls=[tool_call]),
  523. ChatMessage(
  524. role="tool",
  525. content='{"ok":true}',
  526. name="mock_search",
  527. tool_call_id="atomic-call",
  528. ),
  529. ],
  530. )
  531. assert store.list_messages(session_id) == []
  532. def test_sqlite_store_migration_is_safe_across_concurrent_initialization(
  533. tmp_path,
  534. monkeypatch,
  535. ):
  536. database_path = tmp_path / "concurrent-legacy.sqlite3"
  537. connection = sqlite3.connect(database_path)
  538. connection.execute(
  539. """
  540. CREATE TABLE messages (
  541. id INTEGER PRIMARY KEY AUTOINCREMENT,
  542. session_id TEXT NOT NULL,
  543. turn_index INTEGER NOT NULL,
  544. role TEXT NOT NULL,
  545. content TEXT NOT NULL,
  546. name TEXT,
  547. tool_call_id TEXT,
  548. created_at TEXT NOT NULL
  549. )
  550. """
  551. )
  552. connection.execute(
  553. """
  554. CREATE TABLE turns (
  555. session_id TEXT NOT NULL,
  556. turn_index INTEGER NOT NULL,
  557. user_message TEXT NOT NULL,
  558. started_at TEXT NOT NULL,
  559. completed_at TEXT,
  560. PRIMARY KEY (session_id, turn_index)
  561. )
  562. """
  563. )
  564. connection.execute(
  565. """
  566. CREATE TABLE usage_stats (
  567. id INTEGER PRIMARY KEY AUTOINCREMENT,
  568. session_id TEXT NOT NULL,
  569. turn_index INTEGER NOT NULL,
  570. round_index INTEGER NOT NULL,
  571. prompt_tokens INTEGER NOT NULL,
  572. completion_tokens INTEGER NOT NULL,
  573. total_tokens INTEGER NOT NULL,
  574. cached_tokens INTEGER NOT NULL,
  575. ttft_ms INTEGER,
  576. elapsed_ms INTEGER NOT NULL,
  577. metric_key TEXT,
  578. created_at TEXT NOT NULL
  579. )
  580. """
  581. )
  582. connection.executemany(
  583. """
  584. INSERT INTO usage_stats (
  585. id, session_id, turn_index, round_index, prompt_tokens,
  586. completion_tokens, total_tokens, cached_tokens, ttft_ms,
  587. elapsed_ms, metric_key, created_at
  588. ) VALUES (?, 'legacy', 1, 1, 0, 0, 0, 0, NULL, 1, ?, 'now')
  589. """,
  590. [
  591. (7, "chat:1:1"),
  592. (11, "chat:1:1"),
  593. (13, None),
  594. (17, None),
  595. ],
  596. )
  597. connection.execute(
  598. """
  599. INSERT INTO messages
  600. (session_id, turn_index, role, content, name, tool_call_id, created_at)
  601. VALUES ('legacy', 1, 'assistant', 'visible', NULL, NULL, 'now')
  602. """
  603. )
  604. connection.commit()
  605. connection.close()
  606. worker_count = 8
  607. barrier = Barrier(worker_count)
  608. real_connect = sqlite3.connect
  609. class CoordinatedConnection:
  610. def __init__(self, connection):
  611. self._connection = connection
  612. @property
  613. def row_factory(self):
  614. return self._connection.row_factory
  615. @row_factory.setter
  616. def row_factory(self, value):
  617. self._connection.row_factory = value
  618. def execute(self, sql, *args):
  619. cursor = self._connection.execute(sql, *args)
  620. if (
  621. sql.strip() == "PRAGMA table_info(messages)"
  622. and not self._connection.in_transaction
  623. ):
  624. rows = cursor.fetchall()
  625. barrier.wait(timeout=5)
  626. return rows
  627. return cursor
  628. def __getattr__(self, name):
  629. return getattr(self._connection, name)
  630. def coordinated_connect(*args, **kwargs):
  631. return CoordinatedConnection(real_connect(*args, **kwargs))
  632. monkeypatch.setattr(
  633. "agent_lab.infrastructure.sqlite_store.sqlite3.connect",
  634. coordinated_connect,
  635. )
  636. def initialize_store() -> list[dict]:
  637. store = SQLiteSessionStore(database_path)
  638. barrier.wait()
  639. return store.list_messages("legacy")
  640. with ThreadPoolExecutor(max_workers=worker_count) as executor:
  641. results = list(executor.map(lambda _: initialize_store(), range(worker_count)))
  642. assert all(messages[0]["tool_calls"] == [] for messages in results)
  643. columns = real_connect(database_path).execute(
  644. "PRAGMA table_info(messages)"
  645. ).fetchall()
  646. assert [column[1] for column in columns].count("tool_calls_json") == 1
  647. migrated = real_connect(database_path)
  648. turn_columns = [row[1] for row in migrated.execute("PRAGMA table_info(turns)")]
  649. usage_columns = [
  650. row[1] for row in migrated.execute("PRAGMA table_info(usage_stats)")
  651. ]
  652. metric_indexes = [
  653. row[1] for row in migrated.execute("PRAGMA index_list(usage_stats)")
  654. ]
  655. usage_rows = migrated.execute(
  656. """
  657. SELECT id, session_id, metric_key
  658. FROM usage_stats
  659. ORDER BY id
  660. """
  661. ).fetchall()
  662. migrated.close()
  663. assert turn_columns.count("wall_time_ms") == 1
  664. for column_name in (
  665. "agent",
  666. "call_kind",
  667. "event_id",
  668. "event_name",
  669. "used_fallback",
  670. "tool_latency_ms",
  671. "metric_key",
  672. ):
  673. assert usage_columns.count(column_name) == 1
  674. assert metric_indexes.count("usage_stats_metric_key_uq") == 1
  675. assert usage_rows == [
  676. (7, "legacy", "chat:1:1"),
  677. (13, "legacy", None),
  678. (17, "legacy", None),
  679. ]