Persist regenerated answer versions
This commit is contained in:
@@ -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()
|
||||
Reference in New Issue
Block a user