test_sqlite_store.py 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358
  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={"chat_agent": {"model": "chat-model"}},
  13. )
  14. store.start_turn(session_id, turn_index=1, user_message="debug this")
  15. store.append_message(
  16. session_id,
  17. turn_index=1,
  18. message=ChatMessage(role="user", content="debug this"),
  19. )
  20. store.append_message(
  21. session_id,
  22. turn_index=1,
  23. message=ChatMessage(role="assistant", content="answer"),
  24. )
  25. store.append_audit(
  26. session_id,
  27. event="chat_agent_request",
  28. details={"model": "chat-model"},
  29. turn_index=1,
  30. round_index=1,
  31. )
  32. store.append_usage(
  33. session_id,
  34. turn_index=1,
  35. round_index=1,
  36. usage=TokenUsage(
  37. prompt_tokens=3,
  38. completion_tokens=7,
  39. total_tokens=10,
  40. cached_tokens=2,
  41. ),
  42. ttft_ms=120,
  43. elapsed_ms=450,
  44. )
  45. store.complete_turn(session_id, turn_index=1)
  46. sessions = store.list_sessions()
  47. assert sessions == [
  48. {
  49. "id": session_id,
  50. "title": "Debug first turn",
  51. "created_at": sessions[0]["created_at"],
  52. "updated_at": sessions[0]["updated_at"],
  53. "turn_count": 1,
  54. "total_tokens": 10,
  55. }
  56. ]
  57. assert store.get_session(session_id)["config"]["chat_agent"]["model"] == "chat-model"
  58. assert [message["content"] for message in store.list_messages(session_id)] == [
  59. "debug this",
  60. "answer",
  61. ]
  62. assert store.list_audit_logs(session_id)[0]["event"] == "chat_agent_request"
  63. usage = store.usage_summary(session_id)
  64. assert usage["calls"][0]["total_tokens"] == 10
  65. assert usage["turns"] == [
  66. {
  67. "turn_index": 1,
  68. "prompt_tokens": 3,
  69. "completion_tokens": 7,
  70. "total_tokens": 10,
  71. "cached_tokens": 2,
  72. "elapsed_ms": 450,
  73. }
  74. ]
  75. assert usage["session"]["total_tokens"] == 10
  76. def test_sqlite_store_session_list_totals_are_not_multiplied_by_turns(tmp_path):
  77. store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
  78. session_id = store.create_session(title="two turns", config={})
  79. for turn_index, tokens in [(1, 2), (2, 4)]:
  80. store.start_turn(
  81. session_id,
  82. turn_index=turn_index,
  83. user_message=f"turn {turn_index}",
  84. )
  85. store.append_usage(
  86. session_id,
  87. turn_index=turn_index,
  88. round_index=1,
  89. usage=TokenUsage(total_tokens=tokens),
  90. ttft_ms=None,
  91. elapsed_ms=10,
  92. )
  93. assert store.list_sessions()[0]["turn_count"] == 2
  94. assert store.list_sessions()[0]["total_tokens"] == 6
  95. def test_sqlite_store_roundtrips_assistant_tool_calls_and_tool_replies(tmp_path):
  96. store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
  97. session_id = store.create_session(title="direct tools", config={})
  98. store.start_turn(session_id, turn_index=1, user_message="search")
  99. tool_calls = [
  100. ToolCallEvent(
  101. id="call-1",
  102. name="mock_search",
  103. arguments={"query": "agent lab"},
  104. raw_arguments='{"query":"agent lab"}',
  105. ),
  106. ToolCallEvent(
  107. id="call-2",
  108. name="mock_ticket",
  109. arguments={"title": "follow up"},
  110. raw_arguments='{"title":"follow up"}',
  111. ),
  112. ]
  113. store.append_message(
  114. session_id,
  115. turn_index=1,
  116. message=ChatMessage(
  117. role="assistant",
  118. content="",
  119. tool_calls=tool_calls,
  120. ),
  121. )
  122. store.append_message(
  123. session_id,
  124. turn_index=1,
  125. message=ChatMessage(
  126. role="tool",
  127. content='{"ok":true}',
  128. name="mock_search",
  129. tool_call_id="call-1",
  130. ),
  131. )
  132. messages = store.list_messages(session_id)
  133. assert messages[0]["tool_calls"] == [call.model_dump() for call in tool_calls]
  134. assert messages[1]["tool_calls"] == []
  135. assert messages[1]["tool_call_id"] == "call-1"
  136. def test_sqlite_store_migrates_legacy_messages_with_empty_tool_calls(tmp_path):
  137. database_path = tmp_path / "legacy.sqlite3"
  138. connection = sqlite3.connect(database_path)
  139. connection.execute(
  140. """
  141. CREATE TABLE messages (
  142. id INTEGER PRIMARY KEY AUTOINCREMENT,
  143. session_id TEXT NOT NULL,
  144. turn_index INTEGER NOT NULL,
  145. role TEXT NOT NULL,
  146. content TEXT NOT NULL,
  147. name TEXT,
  148. tool_call_id TEXT,
  149. created_at TEXT NOT NULL
  150. )
  151. """
  152. )
  153. connection.execute(
  154. """
  155. INSERT INTO messages
  156. (session_id, turn_index, role, content, name, tool_call_id, created_at)
  157. VALUES ('legacy', 1, 'assistant', 'old answer', NULL, NULL, 'now')
  158. """
  159. )
  160. connection.commit()
  161. connection.close()
  162. store = SQLiteSessionStore(database_path)
  163. assert store.list_messages("legacy")[0]["tool_calls"] == []
  164. migrated = sqlite3.connect(database_path).execute(
  165. "PRAGMA table_info(messages)"
  166. ).fetchall()
  167. tool_calls_column = next(column for column in migrated if column[1] == "tool_calls_json")
  168. assert tool_calls_column[3] == 1
  169. assert tool_calls_column[4] == "'[]'"
  170. @pytest.mark.parametrize(
  171. ("persisted_config", "requested_mode", "persisted_mode"),
  172. [
  173. ({"tool_invocation_mode": "chat_agent_tools"}, "dual_agent", "chat_agent_tools"),
  174. ({}, "chat_agent_tools", "dual_agent"),
  175. ],
  176. )
  177. def test_sqlite_store_rejects_session_mode_mismatch_without_mutation(
  178. tmp_path,
  179. persisted_config,
  180. requested_mode,
  181. persisted_mode,
  182. ):
  183. store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
  184. session_id = store.create_session(
  185. title="mode locked",
  186. config=persisted_config,
  187. session_id="mode-locked",
  188. )
  189. store.start_turn(session_id, turn_index=1, user_message="past user")
  190. store.append_message(
  191. session_id,
  192. turn_index=1,
  193. message=ChatMessage(role="user", content="past user"),
  194. )
  195. before_session = store.get_session(session_id)
  196. before_messages = store.list_messages(session_id)
  197. with pytest.raises(
  198. ValueError,
  199. match=(
  200. "session tool invocation mode mismatch: "
  201. f"persisted={persisted_mode} requested={requested_mode}"
  202. ),
  203. ):
  204. store.ensure_session(
  205. session_id,
  206. title="ignored",
  207. config={"tool_invocation_mode": requested_mode, "changed": True},
  208. )
  209. assert store.get_session(session_id) == before_session
  210. assert store.list_messages(session_id) == before_messages
  211. def test_sqlite_store_append_messages_rolls_back_the_whole_protocol_group(tmp_path):
  212. database_path = tmp_path / "agent_lab.sqlite3"
  213. store = SQLiteSessionStore(database_path)
  214. session_id = store.create_session(title="atomic group", config={})
  215. store.start_turn(session_id, turn_index=1, user_message="use tool")
  216. connection = sqlite3.connect(database_path)
  217. connection.execute(
  218. """
  219. CREATE TRIGGER reject_tool_message
  220. BEFORE INSERT ON messages
  221. WHEN NEW.role = 'tool'
  222. BEGIN
  223. SELECT RAISE(ABORT, 'tool blocked');
  224. END
  225. """
  226. )
  227. connection.commit()
  228. connection.close()
  229. tool_call = ToolCallEvent(
  230. id="atomic-call",
  231. name="mock_search",
  232. arguments={"query": "atomic"},
  233. raw_arguments='{"query":"atomic"}',
  234. )
  235. with pytest.raises(sqlite3.IntegrityError, match="tool blocked"):
  236. store.append_messages(
  237. session_id,
  238. turn_index=1,
  239. messages=[
  240. ChatMessage(role="assistant", content="", tool_calls=[tool_call]),
  241. ChatMessage(
  242. role="tool",
  243. content='{"ok":true}',
  244. name="mock_search",
  245. tool_call_id="atomic-call",
  246. ),
  247. ],
  248. )
  249. assert store.list_messages(session_id) == []
  250. def test_sqlite_store_migration_is_safe_across_concurrent_initialization(
  251. tmp_path,
  252. monkeypatch,
  253. ):
  254. database_path = tmp_path / "concurrent-legacy.sqlite3"
  255. connection = sqlite3.connect(database_path)
  256. connection.execute(
  257. """
  258. CREATE TABLE messages (
  259. id INTEGER PRIMARY KEY AUTOINCREMENT,
  260. session_id TEXT NOT NULL,
  261. turn_index INTEGER NOT NULL,
  262. role TEXT NOT NULL,
  263. content TEXT NOT NULL,
  264. name TEXT,
  265. tool_call_id TEXT,
  266. created_at TEXT NOT NULL
  267. )
  268. """
  269. )
  270. connection.execute(
  271. """
  272. INSERT INTO messages
  273. (session_id, turn_index, role, content, name, tool_call_id, created_at)
  274. VALUES ('legacy', 1, 'assistant', 'visible', NULL, NULL, 'now')
  275. """
  276. )
  277. connection.commit()
  278. connection.close()
  279. worker_count = 8
  280. barrier = Barrier(worker_count)
  281. real_connect = sqlite3.connect
  282. class CoordinatedConnection:
  283. def __init__(self, connection):
  284. self._connection = connection
  285. @property
  286. def row_factory(self):
  287. return self._connection.row_factory
  288. @row_factory.setter
  289. def row_factory(self, value):
  290. self._connection.row_factory = value
  291. def execute(self, sql, *args):
  292. cursor = self._connection.execute(sql, *args)
  293. if (
  294. sql.strip() == "PRAGMA table_info(messages)"
  295. and not self._connection.in_transaction
  296. ):
  297. rows = cursor.fetchall()
  298. barrier.wait(timeout=5)
  299. return rows
  300. return cursor
  301. def __getattr__(self, name):
  302. return getattr(self._connection, name)
  303. def coordinated_connect(*args, **kwargs):
  304. return CoordinatedConnection(real_connect(*args, **kwargs))
  305. monkeypatch.setattr(
  306. "agent_lab.infrastructure.sqlite_store.sqlite3.connect",
  307. coordinated_connect,
  308. )
  309. def initialize_store() -> list[dict]:
  310. store = SQLiteSessionStore(database_path)
  311. barrier.wait()
  312. return store.list_messages("legacy")
  313. with ThreadPoolExecutor(max_workers=worker_count) as executor:
  314. results = list(executor.map(lambda _: initialize_store(), range(worker_count)))
  315. assert all(messages[0]["tool_calls"] == [] for messages in results)
  316. columns = real_connect(database_path).execute(
  317. "PRAGMA table_info(messages)"
  318. ).fetchall()
  319. assert [column[1] for column in columns].count("tool_calls_json") == 1