225 lines
6.1 KiB
Python
225 lines
6.1 KiB
Python
import logging
|
|
import os
|
|
from pathlib import Path
|
|
from zoneinfo import ZoneInfo, ZoneInfoNotFoundError
|
|
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
PROJECT_ROOT = Path(__file__).resolve().parent.parent
|
|
ENV_FILE = PROJECT_ROOT / ".env"
|
|
|
|
TOKEN_ENV_NAME = "BOT_TOKEN"
|
|
PASSWORD_ENV_NAME = "ASSISTANT_PASSWORD"
|
|
DB_ENV_NAME = "ASSISTANT_DB"
|
|
MODE_ENV_NAME = "ASSISTANT_MODE"
|
|
OLLAMA_BASE_URL_ENV_NAME = "OLLAMA_BASE_URL"
|
|
OLLAMA_MODEL_ENV_NAME = "OLLAMA_MODEL"
|
|
YANDEX_CLOUD_FOLDER_ENV_NAME = "YANDEX_CLOUD_FOLDER"
|
|
YANDEX_CLOUD_MODEL_ENV_NAME = "YANDEX_CLOUD_MODEL"
|
|
YANDEX_STT_MODEL_ENV_NAME = "YANDEX_STT_MODEL"
|
|
YANDEX_STT_LANGUAGE_ENV_NAME = "YANDEX_STT_LANGUAGE"
|
|
TIMEZONE_ENV_NAME = "ASSISTANT_TIMEZONE"
|
|
WHISPER_MODEL_ENV_NAME = "WHISPER_MODEL"
|
|
WHISPER_DEVICE_ENV_NAME = "WHISPER_DEVICE"
|
|
WHISPER_COMPUTE_TYPE_ENV_NAME = "WHISPER_COMPUTE_TYPE"
|
|
WHISPER_LANGUAGE_ENV_NAME = "WHISPER_LANGUAGE"
|
|
VOICE_MAX_DURATION_ENV_NAME = "VOICE_MAX_DURATION_SECONDS"
|
|
|
|
DEFAULT_DB_FILE = PROJECT_ROOT / "assistant_data.sqlite3"
|
|
DEFAULT_MODE = "local"
|
|
DEFAULT_OLLAMA_BASE_URL = "http://localhost:11434"
|
|
DEFAULT_OLLAMA_MODEL = "qwen3.5:9b"
|
|
DEFAULT_YANDEX_CLOUD_MODEL = "yandexgpt/latest"
|
|
DEFAULT_YANDEX_STT_MODEL = "general"
|
|
DEFAULT_YANDEX_STT_LANGUAGE = "ru-RU"
|
|
DEFAULT_TIMEZONE = "Europe/Moscow"
|
|
DEFAULT_WHISPER_MODEL = "large-v3"
|
|
DEFAULT_WHISPER_DEVICE = "cuda"
|
|
DEFAULT_WHISPER_COMPUTE_TYPE = "int8"
|
|
DEFAULT_WHISPER_LANGUAGE = "ru"
|
|
DEFAULT_VOICE_MAX_DURATION_SECONDS = 120
|
|
|
|
REMINDER_POLL_SECONDS = 30
|
|
MAX_TELEGRAM_MESSAGE_LENGTH = 3900
|
|
MAX_AGENT_STEPS = 5
|
|
CONVERSATION_HISTORY_LIMIT = 12
|
|
|
|
HTML_FORMATS = (
|
|
("Bold", "b"),
|
|
("Italic", "i"),
|
|
)
|
|
|
|
|
|
def load_env_file(path: Path = ENV_FILE) -> None:
|
|
if not path.exists():
|
|
return
|
|
|
|
for raw_line in path.read_text(encoding="utf-8").splitlines():
|
|
line = raw_line.strip()
|
|
if not line or line.startswith("#") or "=" not in line:
|
|
continue
|
|
|
|
key, value = line.split("=", 1)
|
|
key = key.strip()
|
|
value = value.strip()
|
|
|
|
if not key or key in os.environ:
|
|
continue
|
|
|
|
if len(value) >= 2 and value[0] == value[-1] and value[0] in {'"', "'"}:
|
|
value = value[1:-1]
|
|
|
|
os.environ[key] = value
|
|
|
|
|
|
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
|