test_sqlite_store.py 5.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184
  1. import sqlite3
  2. from agent_lab.domain.events import ToolCallEvent
  3. from agent_lab.domain.messages import ChatMessage, TokenUsage
  4. from agent_lab.infrastructure.sqlite_store import SQLiteSessionStore
  5. def test_sqlite_store_persists_session_messages_audit_and_usage(tmp_path):
  6. store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
  7. session_id = store.create_session(
  8. title="Debug first turn",
  9. config={"chat_agent": {"model": "chat-model"}},
  10. )
  11. store.start_turn(session_id, turn_index=1, user_message="debug this")
  12. store.append_message(
  13. session_id,
  14. turn_index=1,
  15. message=ChatMessage(role="user", content="debug this"),
  16. )
  17. store.append_message(
  18. session_id,
  19. turn_index=1,
  20. message=ChatMessage(role="assistant", content="answer"),
  21. )
  22. store.append_audit(
  23. session_id,
  24. event="chat_agent_request",
  25. details={"model": "chat-model"},
  26. turn_index=1,
  27. round_index=1,
  28. )
  29. store.append_usage(
  30. session_id,
  31. turn_index=1,
  32. round_index=1,
  33. usage=TokenUsage(
  34. prompt_tokens=3,
  35. completion_tokens=7,
  36. total_tokens=10,
  37. cached_tokens=2,
  38. ),
  39. ttft_ms=120,
  40. elapsed_ms=450,
  41. )
  42. store.complete_turn(session_id, turn_index=1)
  43. sessions = store.list_sessions()
  44. assert sessions == [
  45. {
  46. "id": session_id,
  47. "title": "Debug first turn",
  48. "created_at": sessions[0]["created_at"],
  49. "updated_at": sessions[0]["updated_at"],
  50. "turn_count": 1,
  51. "total_tokens": 10,
  52. }
  53. ]
  54. assert store.get_session(session_id)["config"]["chat_agent"]["model"] == "chat-model"
  55. assert [message["content"] for message in store.list_messages(session_id)] == [
  56. "debug this",
  57. "answer",
  58. ]
  59. assert store.list_audit_logs(session_id)[0]["event"] == "chat_agent_request"
  60. usage = store.usage_summary(session_id)
  61. assert usage["calls"][0]["total_tokens"] == 10
  62. assert usage["turns"] == [
  63. {
  64. "turn_index": 1,
  65. "prompt_tokens": 3,
  66. "completion_tokens": 7,
  67. "total_tokens": 10,
  68. "cached_tokens": 2,
  69. "elapsed_ms": 450,
  70. }
  71. ]
  72. assert usage["session"]["total_tokens"] == 10
  73. def test_sqlite_store_session_list_totals_are_not_multiplied_by_turns(tmp_path):
  74. store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
  75. session_id = store.create_session(title="two turns", config={})
  76. for turn_index, tokens in [(1, 2), (2, 4)]:
  77. store.start_turn(
  78. session_id,
  79. turn_index=turn_index,
  80. user_message=f"turn {turn_index}",
  81. )
  82. store.append_usage(
  83. session_id,
  84. turn_index=turn_index,
  85. round_index=1,
  86. usage=TokenUsage(total_tokens=tokens),
  87. ttft_ms=None,
  88. elapsed_ms=10,
  89. )
  90. assert store.list_sessions()[0]["turn_count"] == 2
  91. assert store.list_sessions()[0]["total_tokens"] == 6
  92. def test_sqlite_store_roundtrips_assistant_tool_calls_and_tool_replies(tmp_path):
  93. store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
  94. session_id = store.create_session(title="direct tools", config={})
  95. store.start_turn(session_id, turn_index=1, user_message="search")
  96. tool_calls = [
  97. ToolCallEvent(
  98. id="call-1",
  99. name="mock_search",
  100. arguments={"query": "agent lab"},
  101. raw_arguments='{"query":"agent lab"}',
  102. ),
  103. ToolCallEvent(
  104. id="call-2",
  105. name="mock_ticket",
  106. arguments={"title": "follow up"},
  107. raw_arguments='{"title":"follow up"}',
  108. ),
  109. ]
  110. store.append_message(
  111. session_id,
  112. turn_index=1,
  113. message=ChatMessage(
  114. role="assistant",
  115. content="",
  116. tool_calls=tool_calls,
  117. ),
  118. )
  119. store.append_message(
  120. session_id,
  121. turn_index=1,
  122. message=ChatMessage(
  123. role="tool",
  124. content='{"ok":true}',
  125. name="mock_search",
  126. tool_call_id="call-1",
  127. ),
  128. )
  129. messages = store.list_messages(session_id)
  130. assert messages[0]["tool_calls"] == [call.model_dump() for call in tool_calls]
  131. assert messages[1]["tool_calls"] == []
  132. assert messages[1]["tool_call_id"] == "call-1"
  133. def test_sqlite_store_migrates_legacy_messages_with_empty_tool_calls(tmp_path):
  134. database_path = tmp_path / "legacy.sqlite3"
  135. connection = sqlite3.connect(database_path)
  136. connection.execute(
  137. """
  138. CREATE TABLE messages (
  139. id INTEGER PRIMARY KEY AUTOINCREMENT,
  140. session_id TEXT NOT NULL,
  141. turn_index INTEGER NOT NULL,
  142. role TEXT NOT NULL,
  143. content TEXT NOT NULL,
  144. name TEXT,
  145. tool_call_id TEXT,
  146. created_at TEXT NOT NULL
  147. )
  148. """
  149. )
  150. connection.execute(
  151. """
  152. INSERT INTO messages
  153. (session_id, turn_index, role, content, name, tool_call_id, created_at)
  154. VALUES ('legacy', 1, 'assistant', 'old answer', NULL, NULL, 'now')
  155. """
  156. )
  157. connection.commit()
  158. connection.close()
  159. store = SQLiteSessionStore(database_path)
  160. assert store.list_messages("legacy")[0]["tool_calls"] == []
  161. migrated = sqlite3.connect(database_path).execute(
  162. "PRAGMA table_info(messages)"
  163. ).fetchall()
  164. tool_calls_column = next(column for column in migrated if column[1] == "tool_calls_json")
  165. assert tool_calls_column[3] == 1
  166. assert tool_calls_column[4] == "'[]'"