Files
sovereign_ai/SovereignAI-Starter/app/database.py
T

237 lines
8.4 KiB
Python

"""SQLite storage for local users, conversations, and messages."""
from __future__ import annotations
import os
import sqlite3
from uuid import UUID
from pathlib import Path
from contextlib import contextmanager
from typing import Any, Iterator
def _database_path() -> Path:
configured = os.getenv("SOVEREIGNAI_DATA_DIR")
if configured:
data_dir = Path(configured).expanduser()
else:
local_app_data = os.getenv("LOCALAPPDATA")
base_dir = Path(local_app_data) if local_app_data else Path.home() / ".local" / "share"
data_dir = base_dir / "SovereignAI" / "data"
data_dir.mkdir(parents=True, exist_ok=True)
return data_dir / "sovereign_ai.sqlite3"
DATABASE_PATH = _database_path()
LOCAL_USER_ID = str(
UUID(os.getenv("SOVEREIGNAI_LOCAL_USER_ID", "00000000-0000-4000-8000-000000000001"))
)
@contextmanager
def _connect() -> Iterator[sqlite3.Connection]:
connection = sqlite3.connect(DATABASE_PATH, timeout=10)
connection.row_factory = sqlite3.Row
connection.execute("PRAGMA foreign_keys = ON")
try:
with connection:
yield connection
finally:
connection.close()
def initialize_database() -> None:
with _connect() as connection:
connection.execute("PRAGMA journal_mode = WAL")
connection.executescript(
"""
CREATE TABLE IF NOT EXISTS users (
id TEXT PRIMARY KEY,
created_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE TABLE IF NOT EXISTS 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 IF NOT EXISTS 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 INDEX IF NOT EXISTS idx_conversations_user_updated
ON conversations(user_id, updated_at DESC);
CREATE TABLE IF NOT EXISTS 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,
selected_version INTEGER NOT NULL DEFAULT 0,
created_at TEXT NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_messages_conversation
ON messages(conversation_id, id);
CREATE TABLE IF NOT EXISTS message_versions (
message_id INTEGER NOT NULL REFERENCES messages(id) ON DELETE CASCADE,
version_index INTEGER NOT NULL CHECK(version_index >= 0),
content TEXT NOT NULL,
PRIMARY KEY(message_id, version_index)
);
"""
)
message_columns = {
row["name"] for row in connection.execute("PRAGMA table_info(messages)")
}
if "selected_version" not in message_columns:
connection.execute(
"ALTER TABLE messages ADD COLUMN selected_version INTEGER NOT NULL DEFAULT 0"
)
def ensure_user(user_id: str) -> None:
with _connect() as connection:
connection.execute("INSERT OR IGNORE INTO users(id) VALUES (?)", (user_id,))
def list_conversations(user_id: str) -> list[dict[str, Any]]:
with _connect() as connection:
rows = connection.execute(
"""
SELECT c.id, c.title, c.created_at, c.updated_at,
COUNT(m.id) AS message_count
FROM conversations AS c
LEFT JOIN messages AS m ON m.conversation_id = c.id
WHERE c.user_id = ?
GROUP BY c.id
ORDER BY c.updated_at DESC
""",
(user_id,),
).fetchall()
return [dict(row) for row in rows]
def get_conversation(user_id: str, conversation_id: str) -> dict[str, Any] | None:
with _connect() as connection:
conversation = connection.execute(
"""
SELECT id, title, created_at, updated_at
FROM conversations WHERE id = ? AND user_id = ?
""",
(conversation_id, user_id),
).fetchone()
if conversation is None:
return None
messages = connection.execute(
"SELECT id, role, content, selected_version, created_at FROM messages "
"WHERE conversation_id = ? ORDER BY id",
(conversation_id,),
).fetchall()
versions = connection.execute(
"""
SELECT mv.message_id, mv.content
FROM message_versions AS mv
JOIN messages AS m ON m.id = mv.message_id
WHERE m.conversation_id = ?
ORDER BY mv.message_id, mv.version_index
""",
(conversation_id,),
).fetchall()
versions_by_message: dict[int, list[str]] = {}
for version in versions:
versions_by_message.setdefault(version["message_id"], []).append(
version["content"]
)
saved_messages = []
for message in messages:
message_versions = versions_by_message.get(message["id"], [])
if message["role"] == "assistant" and not message_versions:
message_versions = [message["content"]]
saved_messages.append(
{
"role": message["role"],
"content": message["content"],
"selected_version": message["selected_version"],
"versions": message_versions,
"created_at": message["created_at"],
}
)
return {**dict(conversation), "messages": saved_messages}
def save_conversation(
user_id: str,
conversation_id: str,
title: str,
messages: list[dict[str, Any]],
timestamp: str,
) -> None:
with _connect() as connection:
connection.execute("INSERT OR IGNORE INTO users(id) VALUES (?)", (user_id,))
owner = connection.execute(
"SELECT user_id FROM conversations WHERE id = ?", (conversation_id,)
).fetchone()
if owner is not None and owner["user_id"] != user_id:
raise PermissionError("Conversation does not belong to this user.")
connection.execute(
"""
INSERT INTO conversations(id, user_id, title, created_at, updated_at)
VALUES (?, ?, ?, ?, ?)
ON CONFLICT(id) DO UPDATE SET title = excluded.title,
updated_at = excluded.updated_at
""",
(conversation_id, user_id, title, timestamp, timestamp),
)
connection.execute("DELETE FROM messages WHERE conversation_id = ?", (conversation_id,))
for message in messages:
cursor = connection.execute(
"""
INSERT INTO messages(
conversation_id, role, content, selected_version, created_at
) VALUES (?, ?, ?, ?, ?)
""",
(
conversation_id,
message["role"],
message["content"],
message.get("selected_version", 0),
timestamp,
),
)
if message["role"] == "assistant":
versions = message.get("versions") or [message["content"]]
connection.executemany(
"""
INSERT INTO message_versions(message_id, version_index, content)
VALUES (?, ?, ?)
""",
[
(cursor.lastrowid, index, content)
for index, content in enumerate(versions)
],
)
def delete_conversation(user_id: str, conversation_id: str) -> bool:
with _connect() as connection:
cursor = connection.execute(
"DELETE FROM conversations WHERE id = ? AND user_id = ?",
(conversation_id, user_id),
)
return cursor.rowcount > 0
initialize_database()