161 lines
7.0 KiB
Python
161 lines
7.0 KiB
Python
import sqlite3
|
|
import unittest
|
|
from pathlib import Path
|
|
from uuid import uuid4
|
|
|
|
from app import auth, database
|
|
from app.main import app
|
|
from tests.api_client import authenticated_client
|
|
|
|
|
|
class ConversationVersionMigrationTests(unittest.TestCase):
|
|
def setUp(self) -> None:
|
|
self.temp_dir = Path(__file__).parent / f".tmp-db-{uuid4().hex}"
|
|
self.temp_dir.mkdir()
|
|
self.database_path = self.temp_dir / "legacy.sqlite3"
|
|
self.original_database_path = database.DATABASE_PATH
|
|
database.DATABASE_PATH = self.database_path
|
|
connection = sqlite3.connect(self.database_path)
|
|
try:
|
|
connection.executescript(
|
|
"""
|
|
CREATE TABLE users (
|
|
id TEXT PRIMARY KEY,
|
|
created_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP
|
|
);
|
|
CREATE TABLE user_identities (
|
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
|
user_id TEXT NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
|
provider TEXT NOT NULL,
|
|
provider_subject TEXT NOT NULL,
|
|
email TEXT,
|
|
created_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
|
UNIQUE(provider, provider_subject)
|
|
);
|
|
CREATE TABLE conversations (
|
|
id TEXT PRIMARY KEY,
|
|
user_id TEXT NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
|
title TEXT NOT NULL,
|
|
created_at TEXT NOT NULL,
|
|
updated_at TEXT NOT NULL
|
|
);
|
|
CREATE TABLE messages (
|
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
|
conversation_id TEXT NOT NULL REFERENCES conversations(id) ON DELETE CASCADE,
|
|
role TEXT NOT NULL CHECK(role IN ('user', 'assistant')),
|
|
content TEXT NOT NULL,
|
|
created_at TEXT NOT NULL
|
|
);
|
|
CREATE TABLE agent_audit_events (
|
|
id TEXT PRIMARY KEY,
|
|
tool TEXT NOT NULL,
|
|
method TEXT NOT NULL,
|
|
status_code INTEGER NOT NULL,
|
|
duration_ms INTEGER NOT NULL,
|
|
created_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP
|
|
);
|
|
INSERT INTO users(id) VALUES ('00000000-0000-4000-8000-000000000001');
|
|
INSERT INTO user_identities(user_id,provider,provider_subject,email)
|
|
VALUES ('00000000-0000-4000-8000-000000000001','google','subject-1','old@example.test');
|
|
INSERT INTO conversations(id, user_id, title, created_at, updated_at)
|
|
VALUES ('conversation-1', '00000000-0000-4000-8000-000000000001',
|
|
'قديم', '2026-01-01', '2026-01-01');
|
|
INSERT INTO messages(conversation_id, role, content, created_at)
|
|
VALUES ('conversation-1', 'assistant', 'جواب قديم', '2026-01-01');
|
|
INSERT INTO agent_audit_events(id,tool,method,status_code,duration_ms)
|
|
VALUES ('legacy-event','/v1/agent/run','POST',200,12);
|
|
"""
|
|
)
|
|
finally:
|
|
connection.close()
|
|
database.initialize_database()
|
|
|
|
def tearDown(self) -> None:
|
|
database.DATABASE_PATH = self.original_database_path
|
|
for path in self.temp_dir.iterdir():
|
|
path.unlink()
|
|
self.temp_dir.rmdir()
|
|
|
|
def test_old_history_migrates_and_answer_versions_round_trip(self) -> None:
|
|
database.initialize_database()
|
|
with database._connect() as connection:
|
|
identity = connection.execute(
|
|
"SELECT provider_subject,email,password_hash FROM user_identities WHERE provider='google'"
|
|
).fetchone()
|
|
legacy_audit = connection.execute(
|
|
"SELECT user_id FROM agent_audit_events WHERE id='legacy-event'"
|
|
).fetchone()
|
|
self.assertEqual(identity["provider_subject"], "subject-1")
|
|
self.assertEqual(identity["email"], "old@example.test")
|
|
self.assertIsNone(identity["password_hash"])
|
|
self.assertIsNone(legacy_audit["user_id"])
|
|
self.assertEqual(database.list_agent_audit_events(
|
|
"00000000-0000-4000-8000-000000000001"
|
|
), [])
|
|
old_conversation = database.get_conversation(
|
|
"00000000-0000-4000-8000-000000000001", "conversation-1"
|
|
)
|
|
self.assertIsNotNone(old_conversation)
|
|
self.assertEqual(old_conversation["messages"][0]["versions"], ["جواب قديم"])
|
|
self.assertEqual(old_conversation["messages"][0]["selected_version"], 0)
|
|
|
|
database.save_conversation(
|
|
"00000000-0000-4000-8000-000000000001",
|
|
"conversation-1",
|
|
"تجربة النسخ",
|
|
[
|
|
{"role": "user", "content": "السؤال"},
|
|
{
|
|
"role": "assistant",
|
|
"content": "الجواب الأول",
|
|
"versions": ["الجواب الأول", "الجواب الثاني"],
|
|
"selected_version": 0,
|
|
},
|
|
],
|
|
"2026-01-02T00:00:00+00:00",
|
|
)
|
|
conversation = database.get_conversation(
|
|
"00000000-0000-4000-8000-000000000001", "conversation-1"
|
|
)
|
|
self.assertEqual(
|
|
conversation["messages"][1]["versions"], ["الجواب الأول", "الجواب الثاني"]
|
|
)
|
|
self.assertEqual(conversation["messages"][1]["selected_version"], 0)
|
|
|
|
def test_api_saves_and_returns_selected_answer_version(self) -> None:
|
|
user_id = "00000000-0000-4000-8000-000000000001"
|
|
token, _ = auth.issue_session(user_id)
|
|
conversation_id = "00000000-0000-4000-8000-000000000099"
|
|
payload = {
|
|
"title": "API version test",
|
|
"messages": [
|
|
{"role": "user", "content": "question"},
|
|
{
|
|
"role": "assistant",
|
|
"content": "second answer",
|
|
"versions": ["first answer", "second answer"],
|
|
"selected_version": 1,
|
|
},
|
|
],
|
|
}
|
|
with authenticated_client(app) as client:
|
|
saved = client.put(
|
|
f"/v1/conversations/{conversation_id}",
|
|
headers={"Authorization": f"Bearer {token}"},
|
|
json=payload,
|
|
)
|
|
self.assertEqual(saved.status_code, 200, saved.text)
|
|
loaded = client.get(
|
|
f"/v1/conversations/{conversation_id}",
|
|
headers={"Authorization": f"Bearer {token}"},
|
|
)
|
|
self.assertEqual(loaded.status_code, 200, loaded.text)
|
|
assistant = loaded.json()["messages"][1]
|
|
self.assertEqual(assistant["content"], "second answer")
|
|
self.assertEqual(assistant["versions"], ["first answer", "second answer"])
|
|
self.assertEqual(assistant["selected_version"], 1)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|