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] == "'[]'"