refactor: isolate database migrations and model defaults

This commit is contained in:
kandrusyak
2026-07-27 17:51:16 +03:00
parent ab029a4c67
commit c20d893255
7 changed files with 238 additions and 311 deletions

View File

@@ -42,7 +42,10 @@ def configure_logging() -> None:
def create_services(settings: AppSettings) -> ApplicationServices: def create_services(settings: AppSettings) -> ApplicationServices:
storage = AssistantStorage(settings.db_path) storage = AssistantStorage(
settings.db_path,
default_models=settings.default_models(),
)
ai_client: AIClient ai_client: AIClient
speech_recognizer: SpeechTranscriber speech_recognizer: SpeechTranscriber
if settings.mode == "local": if settings.mode == "local":

View File

@@ -66,6 +66,15 @@ class AppSettings:
whisper_compute_type: str whisper_compute_type: str
whisper_language: str | None whisper_language: str | None
def default_models(self) -> dict[str, str]:
models = {"local": self.ollama_model}
if self.yandex_cloud_folder:
models["yandex"] = (
f"gpt://{self.yandex_cloud_folder}/"
f"{self.yandex_cloud_model}"
)
return models
def load_env_file(path: Path = ENV_FILE) -> None: def load_env_file(path: Path = ENV_FILE) -> None:
if not path.exists(): if not path.exists():
@@ -209,154 +218,3 @@ def load_app_settings() -> AppSettings:
else whisper_language or None else whisper_language or None
), ),
) )
def get_bot_token() -> str:
load_env_file()
token = os.getenv(TOKEN_ENV_NAME)
if not token:
raise RuntimeError(f"Set {TOKEN_ENV_NAME} in environment or .env file.")
return token
def get_assistant_password() -> str:
load_env_file()
password = os.getenv(PASSWORD_ENV_NAME, "").strip()
if not password:
raise RuntimeError(
f"Set {PASSWORD_ENV_NAME} in environment or .env file."
)
return password
def get_db_path() -> Path:
load_env_file()
raw_path = os.getenv(DB_ENV_NAME)
if not raw_path:
return DEFAULT_DB_FILE
path = Path(raw_path).expanduser()
if path.is_absolute():
return path
return PROJECT_ROOT / path
def get_assistant_mode() -> str:
load_env_file()
mode = os.getenv(MODE_ENV_NAME, DEFAULT_MODE).strip().lower()
if mode not in {"local", "yandex"}:
raise RuntimeError(
f"{MODE_ENV_NAME} must be either 'local' or 'yandex', got {mode!r}."
)
return mode
def get_ollama_base_url() -> str:
load_env_file()
return os.getenv(OLLAMA_BASE_URL_ENV_NAME, DEFAULT_OLLAMA_BASE_URL).rstrip("/")
def get_default_ollama_model() -> str:
load_env_file()
return os.getenv(OLLAMA_MODEL_ENV_NAME, DEFAULT_OLLAMA_MODEL)
def get_yandex_cloud_folder() -> str:
load_env_file()
folder = os.getenv(YANDEX_CLOUD_FOLDER_ENV_NAME, "").strip()
if not folder or folder == "":
raise RuntimeError(
f"Set {YANDEX_CLOUD_FOLDER_ENV_NAME} in environment or .env file."
)
return folder.strip("/")
def get_yandex_cloud_model() -> str:
load_env_file()
model = os.getenv(
YANDEX_CLOUD_MODEL_ENV_NAME,
DEFAULT_YANDEX_CLOUD_MODEL,
).strip()
return model.strip("/") or DEFAULT_YANDEX_CLOUD_MODEL
def get_default_yandex_model() -> str:
return f"gpt://{get_yandex_cloud_folder()}/{get_yandex_cloud_model()}"
def get_default_model(provider: str) -> str:
if provider == "local":
return get_default_ollama_model()
if provider == "yandex":
return get_default_yandex_model()
raise ValueError(f"Unknown AI provider: {provider}")
def get_yandex_stt_model() -> str:
load_env_file()
return os.getenv(YANDEX_STT_MODEL_ENV_NAME, DEFAULT_YANDEX_STT_MODEL).strip()
def get_yandex_stt_language() -> str:
load_env_file()
return os.getenv(
YANDEX_STT_LANGUAGE_ENV_NAME,
DEFAULT_YANDEX_STT_LANGUAGE,
).strip()
def get_local_timezone() -> ZoneInfo:
load_env_file()
timezone_name = os.getenv(TIMEZONE_ENV_NAME, DEFAULT_TIMEZONE)
try:
return ZoneInfo(timezone_name)
except ZoneInfoNotFoundError:
logger.warning("Unknown timezone %s, falling back to UTC", timezone_name)
return ZoneInfo("UTC")
def get_whisper_model() -> str:
load_env_file()
return os.getenv(WHISPER_MODEL_ENV_NAME, DEFAULT_WHISPER_MODEL).strip()
def get_whisper_device() -> str:
load_env_file()
return os.getenv(WHISPER_DEVICE_ENV_NAME, DEFAULT_WHISPER_DEVICE).strip()
def get_whisper_compute_type() -> str:
load_env_file()
return os.getenv(
WHISPER_COMPUTE_TYPE_ENV_NAME,
DEFAULT_WHISPER_COMPUTE_TYPE,
).strip()
def get_whisper_language() -> str | None:
load_env_file()
language = os.getenv(
WHISPER_LANGUAGE_ENV_NAME,
DEFAULT_WHISPER_LANGUAGE,
).strip()
return None if language.lower() == "auto" else language or None
def get_voice_max_duration_seconds() -> int:
load_env_file()
raw_value = os.getenv(
VOICE_MAX_DURATION_ENV_NAME,
str(DEFAULT_VOICE_MAX_DURATION_SECONDS),
)
try:
value = int(raw_value)
except ValueError:
value = 0
if value <= 0:
logger.warning(
"%s must be a positive integer, using %s",
VOICE_MAX_DURATION_ENV_NAME,
DEFAULT_VOICE_MAX_DURATION_SECONDS,
)
return DEFAULT_VOICE_MAX_DURATION_SECONDS
return value

113
assistant_bot/migrations.py Normal file
View File

@@ -0,0 +1,113 @@
import sqlite3
LATEST_SCHEMA_VERSION = 2
INITIAL_SCHEMA = """
CREATE TABLE IF NOT EXISTS user_settings (
user_id INTEGER PRIMARY KEY,
ollama_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);
"""
def _schema_version(connection: sqlite3.Connection) -> int:
row = connection.execute("PRAGMA user_version").fetchone()
return int(row[0])
def _has_column(
connection: sqlite3.Connection,
table: str,
column: str,
) -> bool:
return any(
str(row["name"]) == column
for row in connection.execute(f"PRAGMA table_info({table})")
)
def migrate_database(connection: sqlite3.Connection) -> None:
"""Bring a new or existing database to the latest known schema."""
connection.execute("PRAGMA journal_mode=WAL")
version = _schema_version(connection)
if version < 1:
connection.executescript(INITIAL_SCHEMA)
connection.execute("PRAGMA user_version = 1")
version = 1
if version < 2:
if not _has_column(connection, "user_settings", "yandex_model"):
connection.execute(
"ALTER TABLE user_settings ADD COLUMN yandex_model TEXT"
)
connection.execute("PRAGMA user_version = 2")

View File

@@ -3,9 +3,9 @@ import sqlite3
from contextlib import contextmanager from contextlib import contextmanager
from datetime import datetime from datetime import datetime
from pathlib import Path 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 from .time_utils import to_utc_iso, utc_now
@@ -21,8 +21,13 @@ def required_lastrowid(cursor: sqlite3.Cursor) -> int:
class AssistantStorage: class AssistantStorage:
def __init__(self, path: Path) -> None: def __init__(
self,
path: Path,
default_models: Mapping[str, str],
) -> None:
self.path = path self.path = path
self._default_models = dict(default_models)
self.path.parent.mkdir(parents=True, exist_ok=True) self.path.parent.mkdir(parents=True, exist_ok=True)
self._init_db() self._init_db()
@@ -45,92 +50,7 @@ class AssistantStorage:
def _init_db(self) -> None: def _init_db(self) -> None:
with self._connection() as connection: with self._connection() as connection:
connection.execute("PRAGMA journal_mode=WAL") migrate_database(connection)
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"
)
def is_user_authorized(self, user_id: int) -> bool: def is_user_authorized(self, user_id: int) -> bool:
with self._connection() as connection: with self._connection() as connection:
@@ -312,7 +232,12 @@ class AssistantStorage:
).fetchone() ).fetchone()
if row and row["model"]: if row and row["model"]:
return str(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( def set_user_model(
self, self,

View File

@@ -19,7 +19,8 @@ class AuthenticationGuardTests(unittest.IsolatedAsyncioTestCase):
def setUp(self) -> None: def setUp(self) -> None:
self.temporary_directory = tempfile.TemporaryDirectory() self.temporary_directory = tempfile.TemporaryDirectory()
self.storage = AssistantStorage( self.storage = AssistantStorage(
Path(self.temporary_directory.name) / "assistant.sqlite3" Path(self.temporary_directory.name) / "assistant.sqlite3",
default_models={"local": "qwen3.5:9b"},
) )
def tearDown(self) -> None: def tearDown(self) -> None:

View File

@@ -6,6 +6,20 @@ from assistant_bot import config
class AppSettingsTests(unittest.TestCase): class AppSettingsTests(unittest.TestCase):
@staticmethod
def load_settings(**overrides):
environment = {
config.TOKEN_ENV_NAME: "token",
config.PASSWORD_ENV_NAME: "secret",
**overrides,
}
with patch("assistant_bot.config.load_env_file"), patch.dict(
os.environ,
environment,
clear=True,
):
return config.load_app_settings()
def test_loads_complete_startup_configuration_once(self) -> None: def test_loads_complete_startup_configuration_once(self) -> None:
environment = { environment = {
config.TOKEN_ENV_NAME: "token", config.TOKEN_ENV_NAME: "token",
@@ -33,87 +47,59 @@ class AppSettingsTests(unittest.TestCase):
self.assertIsNone(settings.whisper_language) self.assertIsNone(settings.whisper_language)
self.assertIsNone(settings.yandex_cloud_folder) self.assertIsNone(settings.yandex_cloud_folder)
class WhisperConfigTests(unittest.TestCase):
def test_gpu_int8_defaults(self) -> None: def test_gpu_int8_defaults(self) -> None:
variable_names = ( settings = self.load_settings()
config.WHISPER_MODEL_ENV_NAME,
config.WHISPER_DEVICE_ENV_NAME,
config.WHISPER_COMPUTE_TYPE_ENV_NAME,
)
with patch("assistant_bot.config.load_env_file"), patch.dict(
os.environ,
{},
clear=False,
):
for variable_name in variable_names:
os.environ.pop(variable_name, None)
self.assertEqual(config.get_whisper_model(), "large-v3") self.assertEqual(settings.whisper_model, "large-v3")
self.assertEqual(config.get_whisper_device(), "cuda") self.assertEqual(settings.whisper_device, "cuda")
self.assertEqual(config.get_whisper_compute_type(), "int8") self.assertEqual(settings.whisper_compute_type, "int8")
class AssistantModeConfigTests(unittest.TestCase):
def test_local_mode_is_default(self) -> None: def test_local_mode_is_default(self) -> None:
with patch("assistant_bot.config.load_env_file"), patch.dict( self.assertEqual(self.load_settings().mode, "local")
os.environ,
{},
clear=False,
):
os.environ.pop(config.MODE_ENV_NAME, None)
self.assertEqual(config.get_assistant_mode(), "local")
def test_yandex_mode_is_supported(self) -> None: def test_yandex_mode_is_supported(self) -> None:
with patch("assistant_bot.config.load_env_file"), patch.dict( settings = self.load_settings(
os.environ, **{
{config.MODE_ENV_NAME: "YANDEX"}, config.MODE_ENV_NAME: "YANDEX",
clear=False, config.YANDEX_CLOUD_FOLDER_ENV_NAME: "folder-id",
): }
self.assertEqual(config.get_assistant_mode(), "yandex") )
self.assertEqual(settings.mode, "yandex")
def test_unknown_mode_is_rejected(self) -> None: def test_unknown_mode_is_rejected(self) -> None:
with patch("assistant_bot.config.load_env_file"), patch.dict(
os.environ,
{config.MODE_ENV_NAME: "cloud"},
clear=False,
):
with self.assertRaises(RuntimeError): with self.assertRaises(RuntimeError):
config.get_assistant_mode() self.load_settings(**{config.MODE_ENV_NAME: "cloud"})
def test_yandex_model_is_full_gpt_uri(self) -> None: def test_yandex_model_is_full_gpt_uri(self) -> None:
with patch("assistant_bot.config.load_env_file"), patch.dict( settings = self.load_settings(
os.environ, **{
{ config.MODE_ENV_NAME: "yandex",
config.YANDEX_CLOUD_FOLDER_ENV_NAME: "folder-id", config.YANDEX_CLOUD_FOLDER_ENV_NAME: "folder-id",
config.YANDEX_CLOUD_MODEL_ENV_NAME: "yandexgpt/latest", config.YANDEX_CLOUD_MODEL_ENV_NAME: "yandexgpt/latest",
}, }
clear=False, )
):
self.assertEqual( self.assertEqual(
config.get_default_yandex_model(), settings.default_models()["yandex"],
"gpt://folder-id/yandexgpt/latest", "gpt://folder-id/yandexgpt/latest",
) )
class PasswordConfigTests(unittest.TestCase):
def test_password_is_required(self) -> None: def test_password_is_required(self) -> None:
with patch("assistant_bot.config.load_env_file"), patch.dict( with patch("assistant_bot.config.load_env_file"), patch.dict(
os.environ, os.environ,
{}, {config.TOKEN_ENV_NAME: "token"},
clear=False, clear=True,
): ):
os.environ.pop(config.PASSWORD_ENV_NAME, None)
with self.assertRaises(RuntimeError): with self.assertRaises(RuntimeError):
config.get_assistant_password() config.load_app_settings()
def test_password_is_read_from_environment(self) -> None: def test_password_is_read_from_environment(self) -> None:
with patch("assistant_bot.config.load_env_file"), patch.dict( settings = self.load_settings(
os.environ, **{config.PASSWORD_ENV_NAME: " test-password "}
{config.PASSWORD_ENV_NAME: " test-password "}, )
clear=False,
): self.assertEqual(settings.assistant_password, "test-password")
self.assertEqual(config.get_assistant_password(), "test-password")
if __name__ == "__main__": if __name__ == "__main__":

View File

@@ -12,10 +12,21 @@ from assistant_bot.agent import (
parse_agent_decision, parse_agent_decision,
) )
from assistant_bot.agent_tools import AGENT_TOOLS from assistant_bot.agent_tools import AGENT_TOOLS
from assistant_bot.migrations import LATEST_SCHEMA_VERSION
from assistant_bot.reminders import parse_reminder, unit_to_timedelta from assistant_bot.reminders import parse_reminder, unit_to_timedelta
from assistant_bot.storage import AssistantStorage from assistant_bot.storage import AssistantStorage
DEFAULT_MODELS = {
"local": "qwen3.5:9b",
"yandex": "gpt://folder/yandexgpt/latest",
}
def create_storage(path: Path) -> AssistantStorage:
return AssistantStorage(path, default_models=DEFAULT_MODELS)
class ReminderParserTests(unittest.TestCase): class ReminderParserTests(unittest.TestCase):
def test_supported_relative_unit(self) -> None: def test_supported_relative_unit(self) -> None:
delta = unit_to_timedelta(2, "часа") delta = unit_to_timedelta(2, "часа")
@@ -89,7 +100,7 @@ class AgentToolContractTests(unittest.TestCase):
def test_crud_tool_result_shapes_remain_stable(self) -> None: def test_crud_tool_result_shapes_remain_stable(self) -> None:
with tempfile.TemporaryDirectory() as directory: with tempfile.TemporaryDirectory() as directory:
storage = AssistantStorage(Path(directory) / "assistant.sqlite3") storage = create_storage(Path(directory) / "assistant.sqlite3")
common = { common = {
"storage": storage, "storage": storage,
"user_id": 42, "user_id": 42,
@@ -187,7 +198,7 @@ class AgentToolContractTests(unittest.TestCase):
result = execute_agent_tool( result = execute_agent_tool(
"missing", "missing",
{}, {},
AssistantStorage(Path(directory) / "assistant.sqlite3"), create_storage(Path(directory) / "assistant.sqlite3"),
user_id=42, user_id=42,
chat_id=100, chat_id=100,
tz=ZoneInfo("UTC"), tz=ZoneInfo("UTC"),
@@ -200,15 +211,27 @@ class AgentToolContractTests(unittest.TestCase):
class StorageTests(unittest.TestCase): class StorageTests(unittest.TestCase):
def test_new_database_uses_latest_schema_version(self) -> None:
with tempfile.TemporaryDirectory() as directory:
database_path = Path(directory) / "assistant.sqlite3"
create_storage(database_path)
with closing(sqlite3.connect(database_path)) as connection:
version = connection.execute(
"PRAGMA user_version"
).fetchone()[0]
self.assertEqual(version, LATEST_SCHEMA_VERSION)
def test_user_authorization_is_persisted(self) -> None: def test_user_authorization_is_persisted(self) -> None:
with tempfile.TemporaryDirectory() as directory: with tempfile.TemporaryDirectory() as directory:
database_path = Path(directory) / "assistant.sqlite3" database_path = Path(directory) / "assistant.sqlite3"
storage = AssistantStorage(database_path) storage = create_storage(database_path)
self.assertFalse(storage.is_user_authorized(42)) self.assertFalse(storage.is_user_authorized(42))
storage.authorize_user(42) storage.authorize_user(42)
reopened_storage = AssistantStorage(database_path) reopened_storage = create_storage(database_path)
self.assertTrue(reopened_storage.is_user_authorized(42)) self.assertTrue(reopened_storage.is_user_authorized(42))
self.assertFalse(reopened_storage.is_user_authorized(43)) self.assertFalse(reopened_storage.is_user_authorized(43))
@@ -227,14 +250,32 @@ class StorageTests(unittest.TestCase):
) )
connection.commit() connection.commit()
storage = AssistantStorage(database_path) storage = create_storage(database_path)
storage.set_user_model(42, "yandexgpt", "yandex") storage.set_user_model(42, "yandexgpt", "yandex")
self.assertEqual(storage.get_user_model(42, "yandex"), "yandexgpt") self.assertEqual(storage.get_user_model(42, "yandex"), "yandexgpt")
with closing(sqlite3.connect(database_path)) as connection:
version = connection.execute(
"PRAGMA user_version"
).fetchone()[0]
self.assertEqual(version, LATEST_SCHEMA_VERSION)
def test_default_models_are_injected(self) -> None:
with tempfile.TemporaryDirectory() as directory:
storage = create_storage(Path(directory) / "assistant.sqlite3")
self.assertEqual(
storage.get_user_model(42, "local"),
DEFAULT_MODELS["local"],
)
self.assertEqual(
storage.get_user_model(42, "yandex"),
DEFAULT_MODELS["yandex"],
)
def test_provider_models_are_stored_separately(self) -> None: def test_provider_models_are_stored_separately(self) -> None:
with tempfile.TemporaryDirectory() as directory: with tempfile.TemporaryDirectory() as directory:
storage = AssistantStorage(Path(directory) / "assistant.sqlite3") storage = create_storage(Path(directory) / "assistant.sqlite3")
storage.set_user_model(42, "qwen3.5:9b", "local") storage.set_user_model(42, "qwen3.5:9b", "local")
storage.set_user_model(42, "yandexgpt", "yandex") storage.set_user_model(42, "yandexgpt", "yandex")
@@ -244,7 +285,7 @@ class StorageTests(unittest.TestCase):
def test_memory_crud(self) -> None: def test_memory_crud(self) -> None:
with tempfile.TemporaryDirectory() as directory: with tempfile.TemporaryDirectory() as directory:
storage = AssistantStorage(Path(directory) / "assistant.sqlite3") storage = create_storage(Path(directory) / "assistant.sqlite3")
memory_id = storage.add_memory(42, "короткие ответы") memory_id = storage.add_memory(42, "короткие ответы")
self.assertEqual(storage.list_memories(42)[0]["id"], memory_id) self.assertEqual(storage.list_memories(42)[0]["id"], memory_id)
@@ -253,7 +294,7 @@ class StorageTests(unittest.TestCase):
def test_new_context_preserves_searchable_archive(self) -> None: def test_new_context_preserves_searchable_archive(self) -> None:
with tempfile.TemporaryDirectory() as directory: with tempfile.TemporaryDirectory() as directory:
storage = AssistantStorage(Path(directory) / "assistant.sqlite3") storage = create_storage(Path(directory) / "assistant.sqlite3")
storage.add_conversation_exchange( storage.add_conversation_exchange(
42, 42,
100, 100,