Files
sovereign_ai/SovereignAI-Starter/tests/test_database_versions.py
T

145 lines
6.1 KiB
Python

import sqlite3
import unittest
from pathlib import Path
from uuid import uuid4
from fastapi.testclient import TestClient
from app import auth, 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 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
);
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');
"""
)
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()
self.assertEqual(identity["provider_subject"], "subject-1")
self.assertEqual(identity["email"], "old@example.test")
self.assertIsNone(identity["password_hash"])
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 TestClient(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()