| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184 |
- import sqlite3
- from agent_lab.domain.events import ToolCallEvent
- from agent_lab.domain.messages import ChatMessage, TokenUsage
- from agent_lab.infrastructure.sqlite_store import SQLiteSessionStore
- def test_sqlite_store_persists_session_messages_audit_and_usage(tmp_path):
- store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
- session_id = store.create_session(
- title="Debug first turn",
- config={"chat_agent": {"model": "chat-model"}},
- )
- store.start_turn(session_id, turn_index=1, user_message="debug this")
- store.append_message(
- session_id,
- turn_index=1,
- message=ChatMessage(role="user", content="debug this"),
- )
- store.append_message(
- session_id,
- turn_index=1,
- message=ChatMessage(role="assistant", content="answer"),
- )
- store.append_audit(
- session_id,
- event="chat_agent_request",
- details={"model": "chat-model"},
- turn_index=1,
- round_index=1,
- )
- store.append_usage(
- session_id,
- turn_index=1,
- round_index=1,
- usage=TokenUsage(
- prompt_tokens=3,
- completion_tokens=7,
- total_tokens=10,
- cached_tokens=2,
- ),
- ttft_ms=120,
- elapsed_ms=450,
- )
- store.complete_turn(session_id, turn_index=1)
- sessions = store.list_sessions()
- assert sessions == [
- {
- "id": session_id,
- "title": "Debug first turn",
- "created_at": sessions[0]["created_at"],
- "updated_at": sessions[0]["updated_at"],
- "turn_count": 1,
- "total_tokens": 10,
- }
- ]
- assert store.get_session(session_id)["config"]["chat_agent"]["model"] == "chat-model"
- assert [message["content"] for message in store.list_messages(session_id)] == [
- "debug this",
- "answer",
- ]
- assert store.list_audit_logs(session_id)[0]["event"] == "chat_agent_request"
- usage = store.usage_summary(session_id)
- assert usage["calls"][0]["total_tokens"] == 10
- assert usage["turns"] == [
- {
- "turn_index": 1,
- "prompt_tokens": 3,
- "completion_tokens": 7,
- "total_tokens": 10,
- "cached_tokens": 2,
- "elapsed_ms": 450,
- }
- ]
- assert usage["session"]["total_tokens"] == 10
- def test_sqlite_store_session_list_totals_are_not_multiplied_by_turns(tmp_path):
- store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
- session_id = store.create_session(title="two turns", config={})
- for turn_index, tokens in [(1, 2), (2, 4)]:
- store.start_turn(
- session_id,
- turn_index=turn_index,
- user_message=f"turn {turn_index}",
- )
- store.append_usage(
- session_id,
- turn_index=turn_index,
- round_index=1,
- usage=TokenUsage(total_tokens=tokens),
- ttft_ms=None,
- elapsed_ms=10,
- )
- assert store.list_sessions()[0]["turn_count"] == 2
- assert store.list_sessions()[0]["total_tokens"] == 6
- def test_sqlite_store_roundtrips_assistant_tool_calls_and_tool_replies(tmp_path):
- store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
- session_id = store.create_session(title="direct tools", config={})
- store.start_turn(session_id, turn_index=1, user_message="search")
- tool_calls = [
- ToolCallEvent(
- id="call-1",
- name="mock_search",
- arguments={"query": "agent lab"},
- raw_arguments='{"query":"agent lab"}',
- ),
- ToolCallEvent(
- id="call-2",
- name="mock_ticket",
- arguments={"title": "follow up"},
- raw_arguments='{"title":"follow up"}',
- ),
- ]
- store.append_message(
- session_id,
- turn_index=1,
- message=ChatMessage(
- role="assistant",
- content="",
- tool_calls=tool_calls,
- ),
- )
- store.append_message(
- session_id,
- turn_index=1,
- message=ChatMessage(
- role="tool",
- content='{"ok":true}',
- name="mock_search",
- tool_call_id="call-1",
- ),
- )
- messages = store.list_messages(session_id)
- assert messages[0]["tool_calls"] == [call.model_dump() for call in tool_calls]
- assert messages[1]["tool_calls"] == []
- assert messages[1]["tool_call_id"] == "call-1"
- def test_sqlite_store_migrates_legacy_messages_with_empty_tool_calls(tmp_path):
- database_path = tmp_path / "legacy.sqlite3"
- connection = sqlite3.connect(database_path)
- connection.execute(
- """
- CREATE TABLE messages (
- id INTEGER PRIMARY KEY AUTOINCREMENT,
- session_id TEXT NOT NULL,
- turn_index INTEGER NOT NULL,
- role TEXT NOT NULL,
- content TEXT NOT NULL,
- name TEXT,
- tool_call_id TEXT,
- created_at TEXT NOT NULL
- )
- """
- )
- connection.execute(
- """
- INSERT INTO messages
- (session_id, turn_index, role, content, name, tool_call_id, created_at)
- VALUES ('legacy', 1, 'assistant', 'old answer', NULL, NULL, 'now')
- """
- )
- connection.commit()
- connection.close()
- store = SQLiteSessionStore(database_path)
- assert store.list_messages("legacy")[0]["tool_calls"] == []
- migrated = sqlite3.connect(database_path).execute(
- "PRAGMA table_info(messages)"
- ).fetchall()
- tool_calls_column = next(column for column in migrated if column[1] == "tool_calls_json")
- assert tool_calls_column[3] == 1
- assert tool_calls_column[4] == "'[]'"
|