refactor: isolate database migrations and model defaults
This commit is contained in:
@@ -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":
|
||||||
|
|||||||
@@ -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
113
assistant_bot/migrations.py
Normal 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")
|
||||||
@@ -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,
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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__":
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
Reference in New Issue
Block a user