refactor: isolate database migrations and model defaults
This commit is contained in:
@@ -3,9 +3,9 @@ import sqlite3
|
||||
from contextlib import contextmanager
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Any, Mapping
|
||||
|
||||
from .config import get_default_model
|
||||
from .migrations import migrate_database
|
||||
from .time_utils import to_utc_iso, utc_now
|
||||
|
||||
|
||||
@@ -21,8 +21,13 @@ def required_lastrowid(cursor: sqlite3.Cursor) -> int:
|
||||
|
||||
|
||||
class AssistantStorage:
|
||||
def __init__(self, path: Path) -> None:
|
||||
def __init__(
|
||||
self,
|
||||
path: Path,
|
||||
default_models: Mapping[str, str],
|
||||
) -> None:
|
||||
self.path = path
|
||||
self._default_models = dict(default_models)
|
||||
self.path.parent.mkdir(parents=True, exist_ok=True)
|
||||
self._init_db()
|
||||
|
||||
@@ -45,92 +50,7 @@ class AssistantStorage:
|
||||
|
||||
def _init_db(self) -> None:
|
||||
with self._connection() as connection:
|
||||
connection.execute("PRAGMA journal_mode=WAL")
|
||||
connection.executescript(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS user_settings (
|
||||
user_id INTEGER PRIMARY KEY,
|
||||
ollama_model TEXT,
|
||||
yandex_model TEXT,
|
||||
updated_at TEXT NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS authorized_users (
|
||||
user_id INTEGER PRIMARY KEY,
|
||||
authorized_at TEXT NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS memories (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
user_id INTEGER NOT NULL,
|
||||
text TEXT NOT NULL,
|
||||
created_at TEXT NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS notes (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
user_id INTEGER NOT NULL,
|
||||
text TEXT NOT NULL,
|
||||
created_at TEXT NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS reminders (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
user_id INTEGER NOT NULL,
|
||||
chat_id INTEGER NOT NULL,
|
||||
text TEXT NOT NULL,
|
||||
remind_at TEXT NOT NULL,
|
||||
status TEXT NOT NULL DEFAULT 'pending',
|
||||
created_at TEXT NOT NULL,
|
||||
sent_at TEXT
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS tracked_items (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
user_id INTEGER NOT NULL,
|
||||
title TEXT NOT NULL,
|
||||
status TEXT NOT NULL,
|
||||
created_at TEXT NOT NULL,
|
||||
updated_at TEXT NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS conversation_messages (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
user_id INTEGER NOT NULL,
|
||||
chat_id INTEGER NOT NULL,
|
||||
role TEXT NOT NULL CHECK(role IN ('user', 'assistant')),
|
||||
content TEXT NOT NULL,
|
||||
created_at TEXT NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS conversation_contexts (
|
||||
user_id INTEGER NOT NULL,
|
||||
chat_id INTEGER NOT NULL,
|
||||
started_after_id INTEGER NOT NULL DEFAULT 0,
|
||||
updated_at TEXT NOT NULL,
|
||||
PRIMARY KEY (user_id, chat_id)
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_memories_user_created
|
||||
ON memories(user_id, created_at DESC);
|
||||
CREATE INDEX IF NOT EXISTS idx_notes_user_created
|
||||
ON notes(user_id, created_at DESC);
|
||||
CREATE INDEX IF NOT EXISTS idx_reminders_due
|
||||
ON reminders(status, remind_at);
|
||||
CREATE INDEX IF NOT EXISTS idx_tracked_user_updated
|
||||
ON tracked_items(user_id, updated_at DESC);
|
||||
CREATE INDEX IF NOT EXISTS idx_conversation_chat_id
|
||||
ON conversation_messages(user_id, chat_id, id DESC);
|
||||
"""
|
||||
)
|
||||
columns = {
|
||||
str(row["name"])
|
||||
for row in connection.execute("PRAGMA table_info(user_settings)")
|
||||
}
|
||||
if "yandex_model" not in columns:
|
||||
connection.execute(
|
||||
"ALTER TABLE user_settings ADD COLUMN yandex_model TEXT"
|
||||
)
|
||||
migrate_database(connection)
|
||||
|
||||
def is_user_authorized(self, user_id: int) -> bool:
|
||||
with self._connection() as connection:
|
||||
@@ -312,7 +232,12 @@ class AssistantStorage:
|
||||
).fetchone()
|
||||
if row and row["model"]:
|
||||
return str(row["model"])
|
||||
return get_default_model(provider)
|
||||
try:
|
||||
return self._default_models[provider]
|
||||
except KeyError as exc:
|
||||
raise ValueError(
|
||||
f"No default model configured for provider: {provider}"
|
||||
) from exc
|
||||
|
||||
def set_user_model(
|
||||
self,
|
||||
|
||||
Reference in New Issue
Block a user