Persist regenerated answer versions

This commit is contained in:
Hamza Ayed
2026-10-01 12:51:48 +03:00
parent e1f981d29e
commit 7a042d036c
9 changed files with 445 additions and 27 deletions
@@ -0,0 +1,125 @@
import sqlite3
import unittest
from pathlib import Path
from uuid import uuid4
from fastapi.testclient import TestClient
from app import database
from app.main import app
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 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
);
INSERT INTO users(id) VALUES ('00000000-0000-4000-8000-000000000001');
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');
"""
)
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()
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"
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 TestClient(app) as client:
saved = client.put(
f"/v1/conversations/{conversation_id}",
headers={"X-User-ID": user_id},
json=payload,
)
self.assertEqual(saved.status_code, 200, saved.text)
loaded = client.get(
f"/v1/conversations/{conversation_id}",
headers={"X-User-ID": user_id},
)
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()