382 lines
14 KiB
Python
382 lines
14 KiB
Python
"""SQLite storage for local users, conversations, and messages."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import os
|
|
import hashlib
|
|
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,
|
|
password_hash TEXT,
|
|
created_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
|
UNIQUE(provider, provider_subject)
|
|
);
|
|
|
|
CREATE TABLE IF NOT EXISTS auth_sessions (
|
|
token_hash TEXT PRIMARY KEY,
|
|
user_id TEXT NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
|
expires_at INTEGER NOT NULL,
|
|
created_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP
|
|
);
|
|
CREATE INDEX IF NOT EXISTS idx_auth_sessions_user
|
|
ON auth_sessions(user_id, expires_at);
|
|
|
|
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)
|
|
);
|
|
|
|
CREATE TABLE IF NOT EXISTS answer_feedback (
|
|
user_id TEXT NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
|
conversation_id TEXT NOT NULL REFERENCES conversations(id) ON DELETE CASCADE,
|
|
message_index INTEGER NOT NULL CHECK(message_index >= 0),
|
|
version_index INTEGER NOT NULL CHECK(version_index >= 0),
|
|
answer_hash TEXT NOT NULL,
|
|
rating INTEGER NOT NULL CHECK(rating IN (-1, 1)),
|
|
updated_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
|
PRIMARY KEY(user_id, conversation_id, message_index, version_index, answer_hash)
|
|
);
|
|
CREATE INDEX IF NOT EXISTS idx_answer_feedback_user
|
|
ON answer_feedback(user_id, updated_at DESC);
|
|
|
|
CREATE TABLE IF NOT EXISTS 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
|
|
);
|
|
CREATE INDEX IF NOT EXISTS idx_agent_audit_created
|
|
ON agent_audit_events(created_at DESC);
|
|
"""
|
|
)
|
|
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"
|
|
)
|
|
identity_columns = {
|
|
row["name"] for row in connection.execute("PRAGMA table_info(user_identities)")
|
|
}
|
|
if "password_hash" not in identity_columns:
|
|
connection.execute("ALTER TABLE user_identities ADD COLUMN password_hash TEXT")
|
|
|
|
|
|
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"]
|
|
)
|
|
feedback_by_answer: dict[tuple[int, int, str], int] = {}
|
|
with _connect() as connection:
|
|
feedback_rows = connection.execute(
|
|
"""
|
|
SELECT message_index, version_index, answer_hash, rating
|
|
FROM answer_feedback
|
|
WHERE user_id = ? AND conversation_id = ?
|
|
""",
|
|
(user_id, conversation_id),
|
|
).fetchall()
|
|
for row in feedback_rows:
|
|
feedback_by_answer[(row["message_index"], row["version_index"], row["answer_hash"])] = row["rating"]
|
|
saved_messages = []
|
|
for message_index, message in enumerate(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,
|
|
"feedback_versions": (
|
|
[
|
|
feedback_by_answer.get(
|
|
(
|
|
message_index,
|
|
version_index,
|
|
hashlib.sha256(content.encode("utf-8")).hexdigest(),
|
|
)
|
|
)
|
|
for version_index, content in enumerate(message_versions)
|
|
]
|
|
if message["role"] == "assistant"
|
|
else []
|
|
),
|
|
"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
|
|
|
|
|
|
def save_answer_feedback(
|
|
user_id: str,
|
|
conversation_id: str,
|
|
message_index: int,
|
|
version_index: int,
|
|
rating: int,
|
|
timestamp: str,
|
|
) -> None:
|
|
"""Store a user's rating for the exact assistant answer version."""
|
|
with _connect() as connection:
|
|
messages = connection.execute(
|
|
"SELECT id, role, content FROM messages WHERE conversation_id = ? ORDER BY id",
|
|
(conversation_id,),
|
|
).fetchall()
|
|
if message_index >= len(messages):
|
|
raise ValueError("Message index does not exist.")
|
|
message = messages[message_index]
|
|
if message["role"] != "assistant":
|
|
raise ValueError("Feedback can only be attached to assistant messages.")
|
|
version = connection.execute(
|
|
"SELECT content FROM message_versions WHERE message_id = ? AND version_index = ?",
|
|
(message["id"], version_index),
|
|
).fetchone()
|
|
if version is None:
|
|
if version_index != 0:
|
|
raise ValueError("Answer version does not exist.")
|
|
answer = message["content"]
|
|
else:
|
|
answer = version["content"]
|
|
answer_hash = hashlib.sha256(answer.encode("utf-8")).hexdigest()
|
|
connection.execute(
|
|
"""
|
|
INSERT INTO answer_feedback(
|
|
user_id, conversation_id, message_index, version_index,
|
|
answer_hash, rating, updated_at
|
|
) VALUES (?, ?, ?, ?, ?, ?, ?)
|
|
ON CONFLICT(user_id, conversation_id, message_index, version_index, answer_hash)
|
|
DO UPDATE SET rating = excluded.rating, updated_at = excluded.updated_at
|
|
""",
|
|
(
|
|
user_id, conversation_id, message_index, version_index,
|
|
answer_hash, rating, timestamp,
|
|
),
|
|
)
|
|
|
|
|
|
def record_agent_audit_event(
|
|
event_id: str,
|
|
tool: str,
|
|
method: str,
|
|
status_code: int,
|
|
duration_ms: int,
|
|
) -> None:
|
|
"""Record agent route metadata only; never persist prompts or file contents."""
|
|
with _connect() as connection:
|
|
connection.execute(
|
|
"""
|
|
INSERT INTO agent_audit_events(id, tool, method, status_code, duration_ms)
|
|
VALUES (?, ?, ?, ?, ?)
|
|
""",
|
|
(event_id, tool, method, status_code, duration_ms),
|
|
)
|
|
|
|
|
|
def list_agent_audit_events(limit: int = 50) -> list[dict[str, Any]]:
|
|
bounded_limit = max(1, min(limit, 200))
|
|
with _connect() as connection:
|
|
rows = connection.execute(
|
|
"""
|
|
SELECT id, tool, method, status_code, duration_ms, created_at
|
|
FROM agent_audit_events
|
|
ORDER BY created_at DESC, rowid DESC
|
|
LIMIT ?
|
|
""",
|
|
(bounded_limit,),
|
|
).fetchall()
|
|
return [dict(row) for row in rows]
|
|
|
|
|
|
initialize_database()
|