| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358 |
- import sqlite3
- from concurrent.futures import ThreadPoolExecutor
- from threading import Barrier
- import pytest
- 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] == "'[]'"
- @pytest.mark.parametrize(
- ("persisted_config", "requested_mode", "persisted_mode"),
- [
- ({"tool_invocation_mode": "chat_agent_tools"}, "dual_agent", "chat_agent_tools"),
- ({}, "chat_agent_tools", "dual_agent"),
- ],
- )
- def test_sqlite_store_rejects_session_mode_mismatch_without_mutation(
- tmp_path,
- persisted_config,
- requested_mode,
- persisted_mode,
- ):
- store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
- session_id = store.create_session(
- title="mode locked",
- config=persisted_config,
- session_id="mode-locked",
- )
- store.start_turn(session_id, turn_index=1, user_message="past user")
- store.append_message(
- session_id,
- turn_index=1,
- message=ChatMessage(role="user", content="past user"),
- )
- before_session = store.get_session(session_id)
- before_messages = store.list_messages(session_id)
- with pytest.raises(
- ValueError,
- match=(
- "session tool invocation mode mismatch: "
- f"persisted={persisted_mode} requested={requested_mode}"
- ),
- ):
- store.ensure_session(
- session_id,
- title="ignored",
- config={"tool_invocation_mode": requested_mode, "changed": True},
- )
- assert store.get_session(session_id) == before_session
- assert store.list_messages(session_id) == before_messages
- def test_sqlite_store_append_messages_rolls_back_the_whole_protocol_group(tmp_path):
- database_path = tmp_path / "agent_lab.sqlite3"
- store = SQLiteSessionStore(database_path)
- session_id = store.create_session(title="atomic group", config={})
- store.start_turn(session_id, turn_index=1, user_message="use tool")
- connection = sqlite3.connect(database_path)
- connection.execute(
- """
- CREATE TRIGGER reject_tool_message
- BEFORE INSERT ON messages
- WHEN NEW.role = 'tool'
- BEGIN
- SELECT RAISE(ABORT, 'tool blocked');
- END
- """
- )
- connection.commit()
- connection.close()
- tool_call = ToolCallEvent(
- id="atomic-call",
- name="mock_search",
- arguments={"query": "atomic"},
- raw_arguments='{"query":"atomic"}',
- )
- with pytest.raises(sqlite3.IntegrityError, match="tool blocked"):
- store.append_messages(
- session_id,
- turn_index=1,
- messages=[
- ChatMessage(role="assistant", content="", tool_calls=[tool_call]),
- ChatMessage(
- role="tool",
- content='{"ok":true}',
- name="mock_search",
- tool_call_id="atomic-call",
- ),
- ],
- )
- assert store.list_messages(session_id) == []
- def test_sqlite_store_migration_is_safe_across_concurrent_initialization(
- tmp_path,
- monkeypatch,
- ):
- database_path = tmp_path / "concurrent-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', 'visible', NULL, NULL, 'now')
- """
- )
- connection.commit()
- connection.close()
- worker_count = 8
- barrier = Barrier(worker_count)
- real_connect = sqlite3.connect
- class CoordinatedConnection:
- def __init__(self, connection):
- self._connection = connection
- @property
- def row_factory(self):
- return self._connection.row_factory
- @row_factory.setter
- def row_factory(self, value):
- self._connection.row_factory = value
- def execute(self, sql, *args):
- cursor = self._connection.execute(sql, *args)
- if (
- sql.strip() == "PRAGMA table_info(messages)"
- and not self._connection.in_transaction
- ):
- rows = cursor.fetchall()
- barrier.wait(timeout=5)
- return rows
- return cursor
- def __getattr__(self, name):
- return getattr(self._connection, name)
- def coordinated_connect(*args, **kwargs):
- return CoordinatedConnection(real_connect(*args, **kwargs))
- monkeypatch.setattr(
- "agent_lab.infrastructure.sqlite_store.sqlite3.connect",
- coordinated_connect,
- )
- def initialize_store() -> list[dict]:
- store = SQLiteSessionStore(database_path)
- barrier.wait()
- return store.list_messages("legacy")
- with ThreadPoolExecutor(max_workers=worker_count) as executor:
- results = list(executor.map(lambda _: initialize_store(), range(worker_count)))
- assert all(messages[0]["tool_calls"] == [] for messages in results)
- columns = real_connect(database_path).execute(
- "PRAGMA table_info(messages)"
- ).fetchall()
- assert [column[1] for column in columns].count("tool_calls_json") == 1
|