This commit is contained in:
kandrusyak
2026-07-25 15:07:53 +03:00
commit 6667e0d253
40 changed files with 3258 additions and 0 deletions

3
.air/settings.json Normal file
View File

@@ -0,0 +1,3 @@
{
"editor.guides": []
}

17
.dockerignore Normal file
View File

@@ -0,0 +1,17 @@
.git
.github
.idea
.air
.agents
.env
.venv
venv
__pycache__
*.py[cod]
*.sqlite3
*.sqlite3-*
.coverage
.pytest_cache
.mypy_cache
.ruff_cache
tests

5
.gitignore vendored Normal file
View File

@@ -0,0 +1,5 @@
.env
.venv/
__pycache__/
*.py[cod]
assistant_data.sqlite3*

5
.idea/.gitignore generated vendored Normal file
View File

@@ -0,0 +1,5 @@
# Default ignored files
/shelf/
/workspace.xml
# Editor-based HTTP Client requests
/httpRequests/

10
.idea/KAndrusyak_Bot.iml generated Normal file
View File

@@ -0,0 +1,10 @@
<?xml version="1.0" encoding="UTF-8"?>
<module type="PYTHON_MODULE" version="4">
<component name="NewModuleRootManager">
<content url="file://$MODULE_DIR$">
<excludeFolder url="file://$MODULE_DIR$/.venv" />
</content>
<orderEntry type="jdk" jdkName="kandrusyak_bot" jdkType="Python SDK" />
<orderEntry type="sourceFolder" forTests="false" />
</component>
</module>

View File

@@ -0,0 +1,6 @@
<component name="InspectionProjectProfileManager">
<settings>
<option name="USE_PROJECT_PROFILE" value="false" />
<version value="1.0" />
</settings>
</component>

7
.idea/misc.xml generated Normal file
View File

@@ -0,0 +1,7 @@
<?xml version="1.0" encoding="UTF-8"?>
<project version="4">
<component name="Black">
<option name="sdkName" value="kandrusyak_bot" />
</component>
<component name="ProjectRootManager" version="2" project-jdk-name="kandrusyak_bot" project-jdk-type="Python SDK" />
</project>

8
.idea/modules.xml generated Normal file
View File

@@ -0,0 +1,8 @@
<?xml version="1.0" encoding="UTF-8"?>
<project version="4">
<component name="ProjectModuleManager">
<modules>
<module fileurl="file://$PROJECT_DIR$/.idea/KAndrusyak_Bot.iml" filepath="$PROJECT_DIR$/.idea/KAndrusyak_Bot.iml" />
</modules>
</component>
</project>

6
.idea/vcs.xml generated Normal file
View File

@@ -0,0 +1,6 @@
<?xml version="1.0" encoding="UTF-8"?>
<project version="4">
<component name="VcsDirectoryMappings">
<mapping directory="" vcs="Git" />
</component>
</project>

70
AGENTS.md Normal file
View File

@@ -0,0 +1,70 @@
# AGENTS.md
This repository contains a Python Telegram assistant bot with local Ollama integration and SQLite-backed storage.
## Working scope
- Keep changes targeted to this repository.
- Prefer small, reviewable edits over broad refactors.
- Do not commit secrets from [.env](air-file://v3c0f26v8ert02idn5b5/C:/Users/solda/PycharmProjects/KAndrusyak_Bot/.env?type=file&root=C%3A).
- Treat [assistant_data.sqlite3](air-file://v3c0f26v8ert02idn5b5/C:/Users/solda/PycharmProjects/KAndrusyak_Bot/assistant_data.sqlite3?type=file&root=C%3A) as local runtime data, not source.
## Project structure
- [main.py](air-file://v3c0f26v8ert02idn5b5/C:/Users/solda/PycharmProjects/KAndrusyak_Bot/main.py?type=file&root=C%3A) is the compatible entry point.
- [assistant_bot](air-file://v3c0f26v8ert02idn5b5/C:/Users/solda/PycharmProjects/KAndrusyak_Bot/assistant_bot?type=file&root=C%3A) contains the application package.
- [tests](air-file://v3c0f26v8ert02idn5b5/C:/Users/solda/PycharmProjects/KAndrusyak_Bot/tests?type=file&root=C%3A) contains unittest-based tests.
Key modules:
- [application.py](air-file://v3c0f26v8ert02idn5b5/C:/Users/solda/PycharmProjects/KAndrusyak_Bot/assistant_bot/application.py?type=file&root=C%3A): Telegram app construction and handler registration.
- [handlers.py](air-file://v3c0f26v8ert02idn5b5/C:/Users/solda/PycharmProjects/KAndrusyak_Bot/assistant_bot/handlers.py?type=file&root=C%3A): bot commands and message handling.
- [agent.py](air-file://v3c0f26v8ert02idn5b5/C:/Users/solda/PycharmProjects/KAndrusyak_Bot/assistant_bot/agent.py?type=file&root=C%3A): agent loop and tool decision handling.
- [storage.py](air-file://v3c0f26v8ert02idn5b5/C:/Users/solda/PycharmProjects/KAndrusyak_Bot/assistant_bot/storage.py?type=file&root=C%3A): SQLite persistence.
- [ollama.py](air-file://v3c0f26v8ert02idn5b5/C:/Users/solda/PycharmProjects/KAndrusyak_Bot/assistant_bot/ollama.py?type=file&root=C%3A): Ollama client integration.
- [reminders.py](air-file://v3c0f26v8ert02idn5b5/C:/Users/solda/PycharmProjects/KAndrusyak_Bot/assistant_bot/reminders.py?type=file&root=C%3A): reminder parsing and scheduling logic.
## Environment and dependencies
- Python project with dependencies from [requirements.txt](air-file://v3c0f26v8ert02idn5b5/C:/Users/solda/PycharmProjects/KAndrusyak_Bot/requirements.txt?type=file&root=C%3A).
- Required environment variable: `BOT_TOKEN`.
- Supported optional variables include `OLLAMA_BASE_URL`, `OLLAMA_MODEL`, `ASSISTANT_DB`, and `ASSISTANT_TIMEZONE`.
## Run and test
Install dependencies:
```powershell
python -m pip install -r requirements.txt
```
Run the bot:
```powershell
python main.py
```
Alternative package entry point:
```powershell
python -m assistant_bot
```
Run tests:
```powershell
python -m unittest discover -v
```
## Change guidance
- Preserve the existing stdlib `unittest` style unless there is a clear reason to introduce a new test runner.
- Keep Telegram handlers, storage logic, and Ollama client concerns separated by module.
- Prefer configuration through environment variables instead of hardcoded values.
- When changing reminder parsing, storage behavior, or agent decision parsing, update or add tests in [test_core.py](air-file://v3c0f26v8ert02idn5b5/C:/Users/solda/PycharmProjects/KAndrusyak_Bot/tests/test_core.py?type=file&root=C%3A).
- Avoid editing generated caches under `__pycache__`.
## Operational cautions
- Bot startup depends on a valid Telegram token and reachable Ollama endpoint when model-backed features are used.
- SQLite files may be environment-specific; avoid relying on checked-in runtime state for tests.

64
Dockerfile Normal file
View File

@@ -0,0 +1,64 @@
# Build the local GPU image with: docker build --target local -t kandrusyak-bot:local .
FROM nvidia/cuda:12.3.2-cudnn9-runtime-ubuntu22.04 AS local
ENV PYTHONDONTWRITEBYTECODE=1 \
PYTHONUNBUFFERED=1 \
PIP_DISABLE_PIP_VERSION_CHECK=1 \
ASSISTANT_MODE=local \
ASSISTANT_DB=/data/assistant_data.sqlite3 \
HF_HOME=/data/huggingface
WORKDIR /app
RUN apt-get update \
&& apt-get install --no-install-recommends --yes python3 python3-pip tzdata \
&& rm -rf /var/lib/apt/lists/*
COPY requirements-common.txt requirements-local.txt ./
RUN python3 -m pip install --no-cache-dir --requirement requirements-local.txt
RUN groupadd --system bot \
&& useradd --system --gid bot --create-home bot \
&& mkdir /data \
&& chown bot:bot /data
COPY --chown=bot:bot assistant_bot ./assistant_bot
COPY --chown=bot:bot main.py ./main.py
USER bot
VOLUME ["/data"]
STOPSIGNAL SIGTERM
CMD ["python3", "-m", "assistant_bot"]
# Build the cloud image with: docker build --target yandex -t kandrusyak-bot:yandex .
# This is the default final target when --target is omitted.
FROM python:3.12.13-slim-bookworm AS yandex
ENV PYTHONDONTWRITEBYTECODE=1 \
PYTHONUNBUFFERED=1 \
PIP_DISABLE_PIP_VERSION_CHECK=1 \
ASSISTANT_MODE=yandex \
ASSISTANT_DB=/data/assistant_data.sqlite3
WORKDIR /app
RUN apt-get update \
&& apt-get install --no-install-recommends --yes tzdata \
&& rm -rf /var/lib/apt/lists/*
COPY requirements-common.txt requirements-yandex.txt ./
RUN python -m pip install --no-cache-dir --requirement requirements-yandex.txt
RUN groupadd --system bot \
&& useradd --system --gid bot --create-home bot \
&& mkdir /data \
&& chown bot:bot /data
COPY --chown=bot:bot assistant_bot ./assistant_bot
COPY --chown=bot:bot main.py ./main.py
USER bot
VOLUME ["/data"]
STOPSIGNAL SIGTERM
CMD ["python", "-m", "assistant_bot"]

141
README.md Normal file
View File

@@ -0,0 +1,141 @@
# KAndrusyak Bot
Telegram-ассистент с двумя режимами AI: локальными Ollama + faster-whisper или
облачными YandexGPT + SpeechKit через Yandex AI Studio,
контекстом текущего диалога,
долговременной памятью, заметками, напоминаниями и отслеживанием статусов.
Вся переписка сохраняется в SQLite отдельно для каждого чата. В текущий контекст попадают
последние реплики активной темы, причем более свежие сообщения имеют больший приоритет.
При явной смене темы модель начинает новый контекст, не удаляя архив. К старому обсуждению
можно вернуться обычной просьбой к боту или найти сообщения командой `/history ключевые слова`.
Команда `/new` вручную начинает новую тему и также сохраняет предыдущую переписку.
## Структура
- `main.py` — совместимая точка входа;
- `assistant_bot/application.py` — сборка Telegram-приложения;
- `assistant_bot/config.py` — переменные окружения и настройки;
- `assistant_bot/storage.py` — SQLite-хранилище;
- `assistant_bot/ollama.py` — HTTP-клиент Ollama;
- `assistant_bot/yandex_ai.py` — адаптер Yandex AI Studio SDK;
- `assistant_bot/agent.py` — агентный цикл и выполнение внутренних tools;
- `assistant_bot/handlers.py` — Telegram-команды и сообщения;
- `assistant_bot/reminders.py` — разбор времени напоминаний;
- `assistant_bot/jobs.py` — фоновые задачи;
- `tests/` — модульные тесты ядра.
## Запуск
```powershell
conda activate kandrusyak_bot
python -m pip install -r requirements.txt
python main.py
```
Также пакет можно запустить командой `python -m assistant_bot`.
Обязательная настройка в `.env`:
```dotenv
BOT_TOKEN=...
ASSISTANT_MODE=local
```
`ASSISTANT_MODE` принимает `local` (значение по умолчанию) или `yandex`.
Общие дополнительные настройки: `ASSISTANT_DB`, `ASSISTANT_TIMEZONE` и
`VOICE_MAX_DURATION_SECONDS`.
### Локальный режим
В режиме `local` текст обрабатывает Ollama, а голосовые сообщения распознаются
локально через faster-whisper. Настройки:
```dotenv
ASSISTANT_MODE=local
OLLAMA_BASE_URL=http://localhost:11434
OLLAMA_MODEL=qwen3.5:9b
```
При первом голосовом сообщении модель faster-whisper будет загружена
автоматически. Настройки распознавания:
```dotenv
WHISPER_MODEL=large-v3
WHISPER_DEVICE=cuda
WHISPER_COMPUTE_TYPE=int8
WHISPER_LANGUAGE=ru
VOICE_MAX_DURATION_SECONDS=120
```
По умолчанию используется полная модель `large-v3` на NVIDIA GPU с вычислениями INT8
и beam size 5. Крупная модель сохраняет приоритет качества, а INT8 уменьшает расход
видеопамяти. Для этого режима необходимы CUDA 12, cuBLAS для CUDA 12 и cuDNN 9.
Чтобы определять язык автоматически, укажите `WHISPER_LANGUAGE=auto`.
На Windows используется зафиксированная версия CTranslate2 4.6.0: более новые
сборки 4.7.x могут аварийно завершаться при инициализации модели. Проект настроен
для Conda-окружения `kandrusyak_bot`. Полный набор DLL cuDNN 9 должен находиться
в системном `PATH` или рядом с библиотекой CTranslate2.
### Режим Yandex AI Studio
В режиме `yandex` бот использует официальный Yandex AI Studio SDK: YandexGPT для
текста и SpeechKit для Telegram-аудио в формате OGG Opus. Создайте API-ключ с
необходимыми ролями и укажите стандартные переменные авторизации SDK:
```dotenv
ASSISTANT_MODE=yandex
YANDEX_CLOUD_FOLDER=...
YC_API_KEY=...
YANDEX_CLOUD_MODEL=yandexgpt/latest
YANDEX_STT_MODEL=general
YANDEX_STT_LANGUAGE=ru-RU
VOICE_MAX_DURATION_SECONDS=120
```
Значения моделей можно посмотреть командой `/models`, а пользовательский выбор
сохранить командой `/model имя-модели`. Выбор хранится отдельно для локального и
облачного режимов. Короткое имя облачной модели автоматически преобразуется в URI
`gpt://<YANDEX_CLOUD_FOLDER>/<YANDEX_CLOUD_MODEL>`.
## Docker
Один `Dockerfile` содержит два независимых target. Облачный Yandex-образ является
target по умолчанию и не содержит CUDA, faster-whisper и CTranslate2:
```bash
docker build --target yandex -t kandrusyak-bot:yandex .
docker run -d \
--name kandrusyak-bot \
--restart unless-stopped \
--env-file .env \
-e ASSISTANT_MODE=yandex \
-v kandrusyak-data:/data \
kandrusyak-bot:yandex
```
Локальный образ содержит CUDA 12, cuDNN 9 и faster-whisper. Для него требуется
NVIDIA Container Toolkit:
```bash
docker build --target local -t kandrusyak-bot:local .
docker run -d \
--name kandrusyak-bot \
--restart unless-stopped \
--gpus all \
--env-file .env \
-e ASSISTANT_MODE=local \
-v kandrusyak-data:/data \
kandrusyak-bot:local
```
Ollama должна быть доступна контейнеру по адресу из `OLLAMA_BASE_URL`; адрес
`localhost` внутри контейнера указывает на сам контейнер, а не на Linux-хост.
## Проверка
```powershell
conda activate kandrusyak_bot
python -m unittest discover -v
```

View File

@@ -0,0 +1,2 @@
"""Personal Telegram assistant package."""

View File

@@ -0,0 +1,5 @@
from .application import run
if __name__ == "__main__":
run()

516
assistant_bot/agent.py Normal file
View File

@@ -0,0 +1,516 @@
import json
import logging
import re
from datetime import datetime, timezone
from typing import Any
from zoneinfo import ZoneInfo
from telegram import Update
from telegram.constants import ChatAction
from telegram.ext import ContextTypes
from .ai import AIClientError
from .config import CONVERSATION_HISTORY_LIMIT, MAX_AGENT_STEPS
from .models import AgentDecision, ReminderParseResult
from .prompts import ASSISTANT_SYSTEM_PROMPT
from .reminders import parse_reminder
from .storage import AssistantStorage
from .telegram_utils import (
edit_or_reply,
get_ai_client,
get_storage,
get_tz,
parse_positive_int,
require_user_id,
)
from .time_utils import format_local_dt, to_utc_iso, utc_now
logger = logging.getLogger(__name__)
def build_assistant_context(storage: AssistantStorage, user_id: int, tz: ZoneInfo) -> str:
memories = storage.list_memories(user_id, limit=20)
notes = storage.list_notes(user_id, limit=10)
reminders = storage.list_reminders(user_id, limit=10)
tracked_items = storage.list_tracked_items(user_id, limit=20)
sections: list[str] = []
if memories:
sections.append(
"Долговременная память:\n"
+ "\n".join(f"- #{row['id']}: {row['text']}" for row in memories)
)
if notes:
sections.append(
"Последние заметки:\n"
+ "\n".join(f"- #{row['id']}: {row['text']}" for row in notes)
)
if reminders:
sections.append(
"Активные напоминания:\n"
+ "\n".join(
f"- #{row['id']} {format_local_dt(row['remind_at'], tz)}: {row['text']}"
for row in reminders
)
)
if tracked_items:
sections.append(
"Отслеживаемые статусы:\n"
+ "\n".join(
f"- #{row['id']} {row['title']}: {row['status']}"
for row in tracked_items
)
)
if not sections:
return "Долговременный контекст пока пуст."
return "\n\n".join(sections)
def json_dumps(data: Any) -> str:
return json.dumps(data, ensure_ascii=False, indent=2)
def build_agent_tool_prompt(tz: ZoneInfo) -> str:
now_local = datetime.now(tz).strftime("%Y-%m-%d %H:%M")
return (
"Ты работаешь как агент с внутренними tools. Пользователь пишет свободным текстом, "
"а Telegram-команды ему не нужны.\n"
f"Текущее локальное время: {now_local}. Таймзона: {tz.key}.\n\n"
"Отвечай СТРОГО одним JSON-объектом без Markdown и без текста вокруг.\n"
"Если нужно выполнить действие, верни tool_calls. Если действие уже выполнено "
"или tool не нужен, верни final.\n\n"
"Форматы ответа:\n"
"{\"tool_calls\":[{\"name\":\"create_note\",\"arguments\":{\"text\":\"...\"}}],\"reset_context\":false}\n"
"{\"final\":\"Короткий ответ пользователю\",\"reset_context\":false}\n\n"
"Доступные tools:\n"
"- get_current_datetime {} - узнать текущее локальное и UTC-время.\n"
"- remember {\"text\": string} - сохранить важный долгосрочный факт о пользователе.\n"
"- list_memory {\"limit\": number} - показать сохраненную память.\n"
"- delete_memory {\"id\": number} - удалить запись памяти.\n"
"- create_note {\"text\": string} - сохранить заметку.\n"
"- list_notes {\"limit\": number} - показать заметки.\n"
"- delete_note {\"id\": number} - удалить заметку.\n"
"- create_reminder {\"when\": string, \"text\": string} - поставить напоминание. "
"when можно указывать как '30m', 'через 2 часа', '18:30', '2026-07-16 18:30'. "
"Если пользователь говорит 'завтра/послезавтра/через неделю', сам рассчитай "
"дату от текущего локального времени и передай 'YYYY-MM-DD HH:MM'.\n"
"- list_reminders {\"limit\": number} - показать активные напоминания.\n"
"- cancel_reminder {\"id\": number} - отменить напоминание.\n"
"- create_status {\"title\": string, \"status\": string} - начать отслеживать статус.\n"
"- list_statuses {\"limit\": number} - показать отслеживаемые статусы.\n"
"- update_status {\"id\": number, \"status\": string} - обновить статус.\n"
"- delete_status {\"id\": number} - удалить отслеживаемый объект.\n\n"
"- search_conversation {\"query\": string, \"limit\": number} - найти старое обсуждение "
"во всей сохраненной переписке по содержательным ключевым словам.\n\n"
"Правила:\n"
"- Для просьб 'запомни', 'сохрани как факт', 'учти на будущее' используй remember.\n"
"- Для заметок используй create_note, для напоминаний create_reminder, "
"для контроля дел/заявок/ожиданий create_status или update_status.\n"
"- Если для действия не хватает данных, не вызывай tool, а задай уточняющий вопрос через final.\n"
"- Учитывай историю диалога для коротких ответов на уточняющие вопросы. "
"Если новый запрос явно начинает другую, не связанную с историей тему, не опирайся на старую тему "
"и верни reset_context=true. Для продолжения темы и сомнительных случаев верни false.\n"
"- Если пользователь просит найти, вспомнить или продолжить старое обсуждение, вызови "
"search_conversation. Передавай в query только ключевые слова темы, без общих слов. "
"Результаты поиска содержат соседние реплики; более свежие совпадения при прочих равных важнее.\n"
"- После результата tool верни final с кратким подтверждением или следующим tool_calls."
)
def extract_json_object(raw_text: str) -> dict[str, Any] | None:
candidates = [raw_text.strip()]
code_block = re.search(r"```(?:json)?\s*(.*?)```", raw_text, re.IGNORECASE | re.DOTALL)
if code_block:
candidates.insert(0, code_block.group(1).strip())
first_brace = raw_text.find("{")
last_brace = raw_text.rfind("}")
if 0 <= first_brace < last_brace:
candidates.append(raw_text[first_brace : last_brace + 1])
for candidate in candidates:
if not candidate:
continue
try:
parsed = json.loads(candidate)
except json.JSONDecodeError:
continue
if isinstance(parsed, dict):
return parsed
return None
def parse_agent_decision(raw_text: str) -> AgentDecision:
parsed = extract_json_object(raw_text)
if parsed is None:
return AgentDecision(final=raw_text.strip(), tool_calls=[], reset_context=False)
raw_calls = parsed.get("tool_calls")
if raw_calls is None and parsed.get("tool"):
raw_calls = [
{
"name": parsed.get("tool"),
"arguments": parsed.get("arguments", {}),
}
]
tool_calls: list[dict[str, Any]] = []
if isinstance(raw_calls, list):
for raw_call in raw_calls:
if not isinstance(raw_call, dict):
continue
name = raw_call.get("name") or raw_call.get("tool")
arguments = raw_call.get("arguments", {})
if isinstance(name, str):
tool_calls.append(
{
"name": name.strip(),
"arguments": arguments if isinstance(arguments, dict) else {},
}
)
final = parsed.get("final") or parsed.get("answer")
reset_context = parsed.get("reset_context") is True
if isinstance(final, str) and final.strip():
return AgentDecision(
final=final.strip(),
tool_calls=tool_calls,
reset_context=reset_context,
)
return AgentDecision(final=None, tool_calls=tool_calls, reset_context=reset_context)
def save_conversation_exchange(
storage: AssistantStorage,
user_id: int,
chat_id: int,
prompt: str,
answer: str,
reset_context: bool,
) -> None:
if reset_context:
storage.start_new_conversation(user_id, chat_id)
storage.add_conversation_exchange(user_id, chat_id, prompt, answer)
def coerce_positive_int(value: Any, default: int, maximum: int = 50) -> int:
try:
parsed = int(value)
except (TypeError, ValueError):
return default
if parsed <= 0:
return default
return min(parsed, maximum)
def coerce_required_text(arguments: dict[str, Any], key: str) -> str:
value = arguments.get(key, "")
return str(value).strip()
def parse_agent_reminder(when: str, text: str, tz: ZoneInfo) -> ReminderParseResult | None:
parsed = parse_reminder(f"{when} {text}".strip(), tz)
if parsed:
return parsed
normalized_when = when.strip().replace("T", " ")
try:
parsed_dt = datetime.fromisoformat(normalized_when)
except ValueError:
return None
if parsed_dt.tzinfo is None:
parsed_dt = parsed_dt.replace(tzinfo=tz)
return ReminderParseResult(parsed_dt.astimezone(timezone.utc), text)
def execute_agent_tool(
name: str,
arguments: dict[str, Any],
storage: AssistantStorage,
user_id: int,
chat_id: int,
tz: ZoneInfo,
) -> dict[str, Any]:
tool_name = name.strip().lower()
if tool_name == "get_current_datetime":
now = utc_now()
return {
"ok": True,
"local": now.astimezone(tz).strftime("%Y-%m-%d %H:%M:%S"),
"utc": to_utc_iso(now),
"timezone": tz.key,
}
if tool_name == "remember":
text = coerce_required_text(arguments, "text")
if not text:
return {"ok": False, "error": "text is required"}
memory_id = storage.add_memory(user_id, text)
return {"ok": True, "id": memory_id, "text": text}
if tool_name == "list_memory":
rows = storage.list_memories(user_id, coerce_positive_int(arguments.get("limit"), 20))
return {"ok": True, "items": rows}
if tool_name == "delete_memory":
memory_id = parse_positive_int(str(arguments.get("id", "")))
if memory_id is None:
return {"ok": False, "error": "valid id is required"}
return {"ok": storage.delete_memory(user_id, memory_id), "id": memory_id}
if tool_name == "create_note":
text = coerce_required_text(arguments, "text")
if not text:
return {"ok": False, "error": "text is required"}
note_id = storage.add_note(user_id, text)
return {"ok": True, "id": note_id, "text": text}
if tool_name == "list_notes":
rows = storage.list_notes(user_id, coerce_positive_int(arguments.get("limit"), 20))
return {"ok": True, "items": rows}
if tool_name == "delete_note":
note_id = parse_positive_int(str(arguments.get("id", "")))
if note_id is None:
return {"ok": False, "error": "valid id is required"}
return {"ok": storage.delete_note(user_id, note_id), "id": note_id}
if tool_name == "create_reminder":
when = coerce_required_text(arguments, "when")
text = coerce_required_text(arguments, "text")
if not when or not text:
return {"ok": False, "error": "when and text are required"}
parsed = parse_agent_reminder(when, text, tz)
if not parsed:
return {"ok": False, "error": "could not parse reminder time"}
if parsed.remind_at_utc <= utc_now():
return {"ok": False, "error": "reminder time must be in the future"}
reminder_id = storage.add_reminder(user_id, chat_id, parsed.text, parsed.remind_at_utc)
return {
"ok": True,
"id": reminder_id,
"text": parsed.text,
"remind_at": format_local_dt(parsed.remind_at_utc, tz),
}
if tool_name == "list_reminders":
rows = storage.list_reminders(user_id, coerce_positive_int(arguments.get("limit"), 20))
for row in rows:
row["remind_at_local"] = format_local_dt(row["remind_at"], tz)
return {"ok": True, "items": rows}
if tool_name == "cancel_reminder":
reminder_id = parse_positive_int(str(arguments.get("id", "")))
if reminder_id is None:
return {"ok": False, "error": "valid id is required"}
return {"ok": storage.cancel_reminder(user_id, reminder_id), "id": reminder_id}
if tool_name == "create_status":
title = coerce_required_text(arguments, "title")
status = coerce_required_text(arguments, "status") or "open"
if not title:
return {"ok": False, "error": "title is required"}
item_id = storage.add_tracked_item(user_id, title, status)
return {"ok": True, "id": item_id, "title": title, "status": status}
if tool_name == "list_statuses":
rows = storage.list_tracked_items(user_id, coerce_positive_int(arguments.get("limit"), 30))
return {"ok": True, "items": rows}
if tool_name == "update_status":
item_id = parse_positive_int(str(arguments.get("id", "")))
status = coerce_required_text(arguments, "status")
if item_id is None or not status:
return {"ok": False, "error": "valid id and status are required"}
return {"ok": storage.set_tracked_status(user_id, item_id, status), "id": item_id, "status": status}
if tool_name == "delete_status":
item_id = parse_positive_int(str(arguments.get("id", "")))
if item_id is None:
return {"ok": False, "error": "valid id is required"}
return {"ok": storage.delete_tracked_item(user_id, item_id), "id": item_id}
if tool_name == "search_conversation":
query = coerce_required_text(arguments, "query")
if not query:
return {"ok": False, "error": "query is required"}
limit = coerce_positive_int(arguments.get("limit"), 5, maximum=10)
hits = storage.search_conversation_messages(user_id, chat_id, query, limit)
discussions = [
{
"matched_message_id": hit["id"],
"score": hit["score"],
"messages": storage.conversation_message_window(
user_id,
chat_id,
int(hit["id"]),
),
}
for hit in hits
]
return {"ok": True, "query": query, "discussions": discussions}
return {"ok": False, "error": f"unknown tool: {name}"}
async def run_agent_prompt(
update: Update,
context: ContextTypes.DEFAULT_TYPE,
prompt: str,
) -> None:
if not update.effective_message:
return
user_id = require_user_id(update)
if user_id is None:
await update.effective_message.reply_text("Не могу определить пользователя.")
return
storage = get_storage(context)
ai_client = get_ai_client(context)
tz = get_tz(context)
model = ai_client.normalize_model(
storage.get_user_model(user_id, ai_client.provider)
)
chat_id = int(update.effective_message.chat_id)
assistant_context = build_assistant_context(storage, user_id, tz)
conversation_history = storage.list_conversation_messages(
user_id,
chat_id,
limit=CONVERSATION_HISTORY_LIMIT,
)
messages = [
{"role": "system", "content": ASSISTANT_SYSTEM_PROMPT},
{"role": "system", "content": build_agent_tool_prompt(tz)},
{"role": "system", "content": assistant_context},
{
"role": "system",
"content": (
"Далее идет недавняя история текущей темы от старых реплик к новым. "
"Чем старше реплика, тем меньше ее приоритет; при противоречии опирайся "
"на более новые сообщения. Архив других тем доступен через search_conversation."
),
},
*(
{"role": str(row["role"]), "content": str(row["content"])}
for row in conversation_history
),
{"role": "user", "content": prompt},
]
await context.bot.send_chat_action(
chat_id=chat_id,
action=ChatAction.TYPING,
)
placeholder = await update.effective_message.reply_text(
f"Думаю через {ai_client.display_name} / {model}..."
)
try:
last_tool_results: list[dict[str, Any]] = []
reset_context = False
for _step in range(MAX_AGENT_STEPS):
raw_answer = await ai_client.chat(model, messages, json_mode=True)
decision = parse_agent_decision(raw_answer)
reset_context = reset_context or decision.reset_context
if decision.tool_calls:
tool_results = []
for call in decision.tool_calls:
result = execute_agent_tool(
name=str(call.get("name", "")),
arguments=call.get("arguments", {}),
storage=storage,
user_id=user_id,
chat_id=chat_id,
tz=tz,
)
tool_results.append(
{
"name": call.get("name"),
"arguments": call.get("arguments", {}),
"result": result,
}
)
last_tool_results = tool_results
messages.append({"role": "assistant", "content": json_dumps({"tool_calls": decision.tool_calls})})
messages.append(
{
"role": "user",
"content": (
"Результаты tools:\n"
f"{json_dumps({'tool_results': tool_results})}\n"
"Продолжи. Верни либо следующий tool_calls, либо final. "
"Ответ снова строго JSON."
),
}
)
continue
if decision.final:
await edit_or_reply(placeholder, decision.final)
save_conversation_exchange(
storage,
user_id,
chat_id,
prompt,
decision.final,
reset_context,
)
return
messages.append({"role": "assistant", "content": raw_answer})
final_answer = await ai_client.chat(
model,
[
*messages,
{
"role": "user",
"content": (
"Лимит tool-шагов исчерпан. Больше не вызывай tools. "
f"Последние результаты tools: {json_dumps(last_tool_results)}. "
"Сформулируй короткий финальный ответ пользователю обычным текстом."
),
},
],
)
await edit_or_reply(placeholder, final_answer)
save_conversation_exchange(
storage,
user_id,
chat_id,
prompt,
final_answer,
reset_context,
)
except AIClientError as exc:
if ai_client.provider == "local":
hint = f"Проверь, что Ollama запущена и модель установлена: ollama pull {model}"
else:
hint = (
"Проверь YANDEX_CLOUD_FOLDER, YC_API_KEY "
"и доступ к выбранной модели."
)
await placeholder.edit_text(
f"Не получилось вызвать {ai_client.display_name}.\n{exc}\n\n{hint}"
)
return
except Exception:
logger.exception("Failed to run AI prompt")
await placeholder.edit_text("Произошла внутренняя ошибка при запросе к модели.")
async def run_ai_prompt(
update: Update,
context: ContextTypes.DEFAULT_TYPE,
prompt: str,
) -> None:
await run_agent_prompt(update, context, prompt)

24
assistant_bot/ai.py Normal file
View File

@@ -0,0 +1,24 @@
from typing import Protocol
class AIClientError(RuntimeError):
"""Raised when the configured AI provider cannot complete a request."""
class AIClient(Protocol):
provider: str
display_name: str
async def list_models(self) -> list[str]:
...
def normalize_model(self, model: str) -> str:
...
async def chat(
self,
model: str,
messages: list[dict[str, str]],
json_mode: bool = False,
) -> str:
...

View File

@@ -0,0 +1,112 @@
import logging
from telegram import Update
from telegram.ext import (
Application,
CommandHandler,
InlineQueryHandler,
MessageHandler,
filters,
)
from .ai import AIClient
from .config import (
get_assistant_mode,
get_bot_token,
get_db_path,
get_local_timezone,
get_ollama_base_url,
get_voice_max_duration_seconds,
get_whisper_compute_type,
get_whisper_device,
get_whisper_language,
get_whisper_model,
get_yandex_cloud_folder,
get_yandex_stt_language,
get_yandex_stt_model,
)
from .handlers import (
ask_command,
help_command,
history_command,
inline_query,
model_command,
models_command,
new_conversation_command,
private_text,
private_voice,
start,
)
from .jobs import post_init
from .ollama import OllamaClient
from .storage import AssistantStorage
from .speech import SpeechRecognizer, SpeechTranscriber
from .yandex_ai import YandexAIClient, YandexSpeechRecognizer
def configure_logging() -> None:
logging.basicConfig(
format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
level=logging.INFO,
)
logging.getLogger("httpx").setLevel(logging.WARNING)
def create_application() -> Application:
storage = AssistantStorage(get_db_path())
mode = get_assistant_mode()
application = (
Application.builder()
.token(get_bot_token())
.post_init(post_init)
.build()
)
application.bot_data["storage"] = storage
ai_client: AIClient
speech_recognizer: SpeechTranscriber
if mode == "local":
ai_client = OllamaClient(get_ollama_base_url())
speech_recognizer = SpeechRecognizer(
model_name=get_whisper_model(),
device=get_whisper_device(),
compute_type=get_whisper_compute_type(),
language=get_whisper_language(),
)
else:
folder_id = get_yandex_cloud_folder()
ai_client = YandexAIClient(folder_id=folder_id)
speech_recognizer = YandexSpeechRecognizer(
folder_id=folder_id,
language=get_yandex_stt_language(),
model=get_yandex_stt_model(),
)
application.bot_data["ai_client"] = ai_client
application.bot_data["assistant_mode"] = mode
application.bot_data["timezone"] = get_local_timezone()
application.bot_data["speech_recognizer"] = speech_recognizer
application.bot_data["voice_max_duration"] = get_voice_max_duration_seconds()
application.add_handler(CommandHandler("start", start))
application.add_handler(CommandHandler("help", help_command))
application.add_handler(CommandHandler("ask", ask_command))
application.add_handler(CommandHandler("models", models_command))
application.add_handler(CommandHandler("model", model_command))
application.add_handler(CommandHandler("new", new_conversation_command))
application.add_handler(CommandHandler("history", history_command))
application.add_handler(InlineQueryHandler(inline_query))
application.add_handler(
MessageHandler(filters.VOICE & filters.ChatType.PRIVATE, private_voice)
)
application.add_handler(
MessageHandler(
filters.TEXT & filters.ChatType.PRIVATE & ~filters.COMMAND,
private_text,
)
)
return application
def run() -> None:
configure_logging()
create_application().run_polling(allowed_updates=Update.ALL_TYPES)

213
assistant_bot/config.py Normal file
View File

@@ -0,0 +1,213 @@
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"
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_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:
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

498
assistant_bot/handlers.py Normal file
View File

@@ -0,0 +1,498 @@
import asyncio
import logging
import os
import tempfile
from pathlib import Path
from typing import Any
from telegram import Update
from telegram.constants import ChatAction
from telegram.error import TelegramError
from telegram.ext import ContextTypes
from .ai import AIClientError
from .agent import run_agent_prompt
from .reminders import parse_reminder
from .telegram_utils import (
build_inline_results,
command_text,
get_ai_client,
get_speech_recognizer,
get_storage,
get_tz,
get_voice_max_duration,
parse_positive_int,
reply_long,
require_user_id,
)
from .speech import SpeechRecognitionError
from .time_utils import format_local_dt
logger = logging.getLogger(__name__)
def format_numbered_rows(
rows: list[dict[str, Any]],
formatter,
empty_text: str,
) -> str:
if not rows:
return empty_text
return "\n".join(formatter(row) for row in rows)
async def start(update: Update, _context: ContextTypes.DEFAULT_TYPE) -> None:
if update.message:
await update.message.reply_text(
"Я готов как персональный ассистент.\n"
"Пиши обычным текстом или отправляй голосовые сообщения: "
"'запомни, что...', 'напомни завтра в 10:00...', "
"'сохрани заметку...', 'покажи мои статусы'.\n"
"Я помню последние реплики диалога; команда /new очищает текущий контекст.\n"
"Я сам решу, когда нужно вызвать внутренний tool."
)
async def help_command(update: Update, _context: ContextTypes.DEFAULT_TYPE) -> None:
if update.message:
await update.message.reply_text(
"Основной режим - свободный диалог текстом или голосовыми сообщениями.\n\n"
"Примеры:\n"
"Запомни, что я предпочитаю короткие ответы.\n"
"Сохрани заметку: идея для проекта.\n"
"Напомни через 30 минут проверить сборку.\n"
"Отслеживай паспорт, статус: жду ответа.\n"
"Покажи активные напоминания.\n\n"
"Вся переписка сохраняется. /history тема ищет старое обсуждение, "
"/new начинает новый контекст без удаления архива.\n\n"
"Служебные команды оставлены для настройки и отладки: /models, /model, /ask.\n"
"Режим выбирается через ASSISTANT_MODE=local или ASSISTANT_MODE=yandex."
)
async def new_conversation_command(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
if not update.message:
return
user_id = require_user_id(update)
if user_id is None:
await update.message.reply_text("Не могу определить пользователя.")
return
get_storage(context).start_new_conversation(user_id, int(update.message.chat_id))
await update.message.reply_text("Начинаем новую тему. Предыдущая переписка сохранена в архиве.")
async def history_command(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
if not update.message:
return
user_id = require_user_id(update)
if user_id is None:
await update.message.reply_text("Не могу определить пользователя.")
return
query = command_text(context)
if not query:
await update.message.reply_text("Укажи тему или ключевые слова: /history отпуск")
return
rows = get_storage(context).search_conversation_messages(
user_id,
int(update.message.chat_id),
query,
limit=10,
)
if not rows:
await update.message.reply_text("В сохраненной переписке ничего не найдено.")
return
role_names = {"user": "Вы", "assistant": "Бот"}
tz = get_tz(context)
text = "Найденные сообщения:\n\n" + "\n\n".join(
f"{role_names.get(str(row['role']), row['role'])} · "
f"{format_local_dt(row['created_at'], tz)}\n{row['content']}"
for row in rows
)
await reply_long(update.message, text)
async def ask_command(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
prompt = command_text(context)
if not prompt:
if update.message:
await update.message.reply_text("Напиши текст после /ask или просто отправь сообщение в личный чат.")
return
await run_agent_prompt(update, context, prompt)
async def private_text(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
if not update.message or not update.message.text:
return
await run_agent_prompt(update, context, update.message.text.strip())
async def private_voice(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
if not update.message or not update.message.voice:
return
voice = update.message.voice
max_duration = get_voice_max_duration(context)
if voice.duration > max_duration:
await update.message.reply_text(
f"Голосовое сообщение слишком длинное. Максимум: {max_duration} сек."
)
return
await context.bot.send_chat_action(
chat_id=int(update.message.chat_id),
action=ChatAction.TYPING,
)
status_message = await update.message.reply_text("Распознаю голосовое сообщение...")
temp_path: Path | None = None
try:
telegram_file = await context.bot.get_file(voice.file_id)
file_descriptor, raw_path = tempfile.mkstemp(suffix=".ogg")
os.close(file_descriptor)
temp_path = Path(raw_path)
await telegram_file.download_to_drive(custom_path=temp_path)
transcript = await asyncio.to_thread(
get_speech_recognizer(context).transcribe,
temp_path,
)
except (SpeechRecognitionError, TelegramError):
logger.exception("Failed to transcribe Telegram voice message")
await status_message.edit_text(
"Не получилось распознать голосовое сообщение. Попробуй еще раз позже."
)
return
except Exception:
logger.exception("Unexpected error while processing Telegram voice message")
await status_message.edit_text(
"Произошла внутренняя ошибка при обработке голосового сообщения."
)
return
finally:
if temp_path is not None:
try:
temp_path.unlink(missing_ok=True)
except OSError:
logger.warning("Could not delete temporary voice file %s", temp_path)
if not transcript:
await status_message.edit_text("Не удалось расслышать речь в сообщении.")
return
preview = transcript if len(transcript) <= 500 else f"{transcript[:497]}..."
await status_message.edit_text(f"Распознано: {preview}")
await run_agent_prompt(update, context, transcript)
async def models_command(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
if not update.message:
return
ai_client = get_ai_client(context)
try:
models = await ai_client.list_models()
except AIClientError as exc:
await update.message.reply_text(
f"Не получилось получить список моделей {ai_client.display_name}.\n{exc}"
)
return
if not models:
await update.message.reply_text(
f"{ai_client.display_name} доступна, но моделей не найдено."
)
return
await update.message.reply_text(
f"Модели {ai_client.display_name}:\n"
+ "\n".join(f"- {model}" for model in models)
)
async def model_command(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
if not update.message:
return
user_id = require_user_id(update)
if user_id is None:
await update.message.reply_text("Не могу определить пользователя.")
return
storage = get_storage(context)
ai_client = get_ai_client(context)
model = command_text(context)
if not model:
await update.message.reply_text(
f"Текущая модель {ai_client.display_name}: "
f"{ai_client.normalize_model(storage.get_user_model(user_id, ai_client.provider))}"
)
return
model = ai_client.normalize_model(model)
storage.set_user_model(user_id, model, ai_client.provider)
await update.message.reply_text(
f"Модель {ai_client.display_name} сохранена: {model}"
)
async def remember_command(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
if not update.message:
return
user_id = require_user_id(update)
text = command_text(context)
if user_id is None:
await update.message.reply_text("Не могу определить пользователя.")
return
if not text:
await update.message.reply_text("Напиши факт после команды: /remember я предпочитаю короткие ответы")
return
memory_id = get_storage(context).add_memory(user_id, text)
await update.message.reply_text(f"Запомнил. ID памяти: {memory_id}")
async def memory_command(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
if not update.message:
return
user_id = require_user_id(update)
if user_id is None:
await update.message.reply_text("Не могу определить пользователя.")
return
rows = get_storage(context).list_memories(user_id)
text = format_numbered_rows(
rows,
lambda row: f"#{row['id']} - {row['text']}",
"Долговременная память пока пустая.",
)
await reply_long(update.message, text)
async def forget_memory_command(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
if not update.message:
return
user_id = require_user_id(update)
memory_id = parse_positive_int(command_text(context))
if user_id is None:
await update.message.reply_text("Не могу определить пользователя.")
return
if memory_id is None:
await update.message.reply_text("Укажи ID: /forget_memory 3")
return
deleted = get_storage(context).delete_memory(user_id, memory_id)
await update.message.reply_text("Удалил." if deleted else "Не нашел такую запись памяти.")
async def note_command(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
if not update.message:
return
user_id = require_user_id(update)
text = command_text(context)
if user_id is None:
await update.message.reply_text("Не могу определить пользователя.")
return
if not text:
await update.message.reply_text("Напиши текст заметки: /note идея для проекта")
return
note_id = get_storage(context).add_note(user_id, text)
await update.message.reply_text(f"Заметка сохранена. ID: {note_id}")
async def notes_command(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
if not update.message:
return
user_id = require_user_id(update)
if user_id is None:
await update.message.reply_text("Не могу определить пользователя.")
return
rows = get_storage(context).list_notes(user_id)
text = format_numbered_rows(
rows,
lambda row: f"#{row['id']} - {row['text']}",
"Заметок пока нет.",
)
await reply_long(update.message, text)
async def forget_note_command(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
if not update.message:
return
user_id = require_user_id(update)
note_id = parse_positive_int(command_text(context))
if user_id is None:
await update.message.reply_text("Не могу определить пользователя.")
return
if note_id is None:
await update.message.reply_text("Укажи ID: /forget 3")
return
deleted = get_storage(context).delete_note(user_id, note_id)
await update.message.reply_text("Удалил." if deleted else "Не нашел такую заметку.")
async def remind_command(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
if not update.message or not update.effective_chat:
return
user_id = require_user_id(update)
if user_id is None:
await update.message.reply_text("Не могу определить пользователя.")
return
parsed = parse_reminder(command_text(context), get_tz(context))
if not parsed:
await update.message.reply_text(
"Формат напоминания:\n"
"/remind 30m позвонить\n"
"/remind через 2 часа проверить статус\n"
"/remind 18:30 отправить отчет\n"
"/remind 2026-07-16 18:30 отправить отчет"
)
return
reminder_id = get_storage(context).add_reminder(
user_id=user_id,
chat_id=int(update.effective_chat.id),
text=parsed.text,
remind_at_utc=parsed.remind_at_utc,
)
await update.message.reply_text(
f"Напоминание #{reminder_id} поставлено на "
f"{format_local_dt(parsed.remind_at_utc, get_tz(context))}."
)
async def reminders_command(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
if not update.message:
return
user_id = require_user_id(update)
if user_id is None:
await update.message.reply_text("Не могу определить пользователя.")
return
tz = get_tz(context)
rows = get_storage(context).list_reminders(user_id)
text = format_numbered_rows(
rows,
lambda row: f"#{row['id']} - {format_local_dt(row['remind_at'], tz)} - {row['text']}",
"Активных напоминаний нет.",
)
await reply_long(update.message, text)
async def cancel_reminder_command(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
if not update.message:
return
user_id = require_user_id(update)
reminder_id = parse_positive_int(command_text(context))
if user_id is None:
await update.message.reply_text("Не могу определить пользователя.")
return
if reminder_id is None:
await update.message.reply_text("Укажи ID: /cancelreminder 3")
return
cancelled = get_storage(context).cancel_reminder(user_id, reminder_id)
await update.message.reply_text("Отменил." if cancelled else "Не нашел активное напоминание.")
def split_tracking_text(raw: str) -> tuple[str, str]:
if "|" not in raw:
return raw.strip(), "open"
title, status = raw.split("|", 1)
return title.strip(), status.strip() or "open"
async def track_command(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
if not update.message:
return
user_id = require_user_id(update)
raw_text = command_text(context)
if user_id is None:
await update.message.reply_text("Не могу определить пользователя.")
return
if not raw_text:
await update.message.reply_text("Формат: /track паспорт | жду ответа")
return
title, status = split_tracking_text(raw_text)
if not title:
await update.message.reply_text("Название не должно быть пустым.")
return
item_id = get_storage(context).add_tracked_item(user_id, title, status)
await update.message.reply_text(f"Добавил отслеживание #{item_id}: {title} - {status}")
async def status_command(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
if not update.message:
return
user_id = require_user_id(update)
if user_id is None:
await update.message.reply_text("Не могу определить пользователя.")
return
raw_text = command_text(context)
storage = get_storage(context)
if not raw_text:
rows = storage.list_tracked_items(user_id)
text = format_numbered_rows(
rows,
lambda row: f"#{row['id']} - {row['title']}: {row['status']}",
"Нет отслеживаемых статусов.",
)
await reply_long(update.message, text)
return
parts = raw_text.split(maxsplit=1)
item_id = parse_positive_int(parts[0])
new_status = parts[1].strip() if len(parts) > 1 else ""
if item_id is None or not new_status:
await update.message.reply_text("Формат: /status 3 новый статус")
return
updated = storage.set_tracked_status(user_id, item_id, new_status)
await update.message.reply_text("Статус обновлен." if updated else "Не нашел такой объект.")
async def untrack_command(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
if not update.message:
return
user_id = require_user_id(update)
item_id = parse_positive_int(command_text(context))
if user_id is None:
await update.message.reply_text("Не могу определить пользователя.")
return
if item_id is None:
await update.message.reply_text("Укажи ID: /untrack 3")
return
deleted = get_storage(context).delete_tracked_item(user_id, item_id)
await update.message.reply_text("Удалил." if deleted else "Не нашел такой объект.")
async def inline_query(update: Update, _context: ContextTypes.DEFAULT_TYPE) -> None:
if not update.inline_query or not update.inline_query.query:
return
try:
await update.inline_query.answer(build_inline_results(update.inline_query.query))
except Exception:
logger.exception("Failed to answer inline query")

42
assistant_bot/jobs.py Normal file
View File

@@ -0,0 +1,42 @@
import asyncio
import logging
from zoneinfo import ZoneInfo
from telegram.ext import Application
from .config import REMINDER_POLL_SECONDS
from .storage import AssistantStorage
from .time_utils import format_local_dt, utc_now
logger = logging.getLogger(__name__)
async def reminder_loop(application: Application) -> None:
storage: AssistantStorage = application.bot_data["storage"]
tz: ZoneInfo = application.bot_data["timezone"]
while True:
try:
due_reminders = storage.due_reminders(utc_now())
for reminder in due_reminders:
await application.bot.send_message(
chat_id=reminder["chat_id"],
text=(
f"Напоминание #{reminder['id']} "
f"({format_local_dt(reminder['remind_at'], tz)}):\n"
f"{reminder['text']}"
),
)
storage.mark_reminder_sent(int(reminder["id"]))
except asyncio.CancelledError:
raise
except Exception:
logger.exception("Reminder loop failed")
await asyncio.sleep(REMINDER_POLL_SECONDS)
async def post_init(application: Application) -> None:
application.create_task(reminder_loop(application))

16
assistant_bot/models.py Normal file
View File

@@ -0,0 +1,16 @@
from dataclasses import dataclass
from datetime import datetime
from typing import Any
@dataclass(frozen=True)
class ReminderParseResult:
remind_at_utc: datetime
text: str
@dataclass(frozen=True)
class AgentDecision:
final: str | None
tool_calls: list[dict[str, Any]]
reset_context: bool = False

99
assistant_bot/ollama.py Normal file
View File

@@ -0,0 +1,99 @@
import asyncio
import json
import urllib.error
import urllib.request
from typing import Any
from .ai import AIClientError
class OllamaError(AIClientError):
pass
class OllamaClient:
provider = "local"
display_name = "Ollama"
def __init__(self, base_url: str) -> None:
self.base_url = base_url.rstrip("/")
def normalize_model(self, model: str) -> str:
return model.strip()
async def list_models(self) -> list[str]:
return await asyncio.to_thread(self._list_models)
async def chat(
self,
model: str,
messages: list[dict[str, str]],
json_mode: bool = False,
) -> str:
return await asyncio.to_thread(self._chat, model, messages, json_mode)
def _list_models(self) -> list[str]:
data = self._request_json("GET", "/api/tags", None, timeout=15)
models = data.get("models", [])
return sorted(model["name"] for model in models if model.get("name"))
def _chat(
self,
model: str,
messages: list[dict[str, str]],
json_mode: bool,
) -> str:
payload = {
"model": model,
"messages": messages,
"stream": False,
}
if json_mode:
payload["format"] = "json"
data = self._request_json("POST", "/api/chat", payload, timeout=180)
content = data.get("message", {}).get("content")
if isinstance(content, str) and content.strip():
return content.strip()
fallback = data.get("response")
if isinstance(fallback, str) and fallback.strip():
return fallback.strip()
raise OllamaError("Ollama returned an empty response.")
def _request_json(
self,
method: str,
path: str,
payload: dict[str, Any] | None,
timeout: int,
) -> dict[str, Any]:
url = f"{self.base_url}{path}"
data = json.dumps(payload).encode("utf-8") if payload is not None else None
request = urllib.request.Request(
url,
data=data,
method=method,
headers={"Content-Type": "application/json"},
)
try:
with urllib.request.urlopen(request, timeout=timeout) as response:
raw = response.read().decode("utf-8")
except urllib.error.HTTPError as exc:
body = exc.read().decode("utf-8", errors="replace")[:400]
raise OllamaError(f"Ollama HTTP {exc.code}: {body}") from exc
except urllib.error.URLError as exc:
raise OllamaError(f"Ollama is unavailable at {self.base_url}: {exc.reason}") from exc
except TimeoutError as exc:
raise OllamaError("Ollama request timed out.") from exc
try:
parsed = json.loads(raw)
except json.JSONDecodeError as exc:
raise OllamaError("Ollama returned invalid JSON.") from exc
if not isinstance(parsed, dict):
raise OllamaError("Ollama returned an unexpected response.")
return parsed

7
assistant_bot/prompts.py Normal file
View File

@@ -0,0 +1,7 @@
ASSISTANT_SYSTEM_PROMPT = (
"Ты персональный AI ассистент в Telegram. Отвечай кратко, по делу и на русском, "
"если пользователь не попросил другой язык. Используй долговременный контекст "
"только как вспомогательную информацию, не выдумывай факты и явно говори, "
"когда данных недостаточно."
)

View File

@@ -0,0 +1,64 @@
import re
from datetime import datetime, timedelta, timezone
from zoneinfo import ZoneInfo
from .models import ReminderParseResult
from .time_utils import utc_now
def unit_to_timedelta(amount: int, unit: str) -> timedelta | None:
normalized = unit.lower().strip(".")
if normalized in {"s", "sec", "secs", "second", "seconds", "с", "сек", "секунд"}:
return timedelta(seconds=amount)
if normalized in {"m", "min", "mins", "minute", "minutes", "м"} or normalized.startswith("мин"):
return timedelta(minutes=amount)
if normalized in {"h", "hr", "hrs", "hour", "hours", "ч"} or normalized.startswith("час"):
return timedelta(hours=amount)
if normalized in {"d", "day", "days", "д"} or normalized.startswith(("дн", "ден")):
return timedelta(days=amount)
return None
def parse_reminder(raw_text: str, tz: ZoneInfo) -> ReminderParseResult | None:
text = raw_text.strip()
if not text:
return None
relative_match = re.match(r"^(?:через\s+)?(\d+)\s*([a-zа-я.]+)\s+(.+)$", text, re.IGNORECASE)
if relative_match:
amount = int(relative_match.group(1))
delta = unit_to_timedelta(amount, relative_match.group(2))
reminder_text = relative_match.group(3).strip()
if delta and reminder_text:
return ReminderParseResult(utc_now() + delta, reminder_text)
now_local = datetime.now(tz)
absolute_patterns = (
(r"^(\d{4}-\d{2}-\d{2})\s+(\d{1,2}:\d{2})\s+(.+)$", "%Y-%m-%d %H:%M"),
(r"^(\d{1,2}\.\d{1,2}\.\d{4})\s+(\d{1,2}:\d{2})\s+(.+)$", "%d.%m.%Y %H:%M"),
)
for pattern, date_format in absolute_patterns:
match = re.match(pattern, text)
if not match:
continue
naive = datetime.strptime(f"{match.group(1)} {match.group(2)}", date_format)
remind_at = naive.replace(tzinfo=tz)
reminder_text = match.group(3).strip()
if reminder_text:
return ReminderParseResult(remind_at.astimezone(timezone.utc), reminder_text)
time_match = re.match(r"^(\d{1,2}:\d{2})\s+(.+)$", text)
if time_match:
hour, minute = (int(part) for part in time_match.group(1).split(":", 1))
if 0 <= hour <= 23 and 0 <= minute <= 59:
remind_at = now_local.replace(hour=hour, minute=minute, second=0, microsecond=0)
if remind_at <= now_local:
remind_at += timedelta(days=1)
reminder_text = time_match.group(2).strip()
if reminder_text:
return ReminderParseResult(remind_at.astimezone(timezone.utc), reminder_text)
return None

85
assistant_bot/speech.py Normal file
View File

@@ -0,0 +1,85 @@
import sys
import threading
import warnings
from pathlib import Path
from typing import Any, Callable, Protocol
class SpeechRecognitionError(RuntimeError):
"""Raised when a voice message cannot be transcribed."""
class SpeechTranscriber(Protocol):
def transcribe(self, audio_path: str | Path) -> str:
...
def load_whisper_model_class() -> Any:
# CTranslate2 imports torch for optional model converters. Speech inference
# does not need it, so do not let an unrelated torch installation prevent
# faster-whisper from loading.
missing = object()
loaded_torch = sys.modules.get("torch", missing)
if loaded_torch is missing:
sys.modules["torch"] = None
try:
with warnings.catch_warnings():
warnings.filterwarnings(
"ignore",
message="pkg_resources is deprecated as an API.*",
category=UserWarning,
)
from faster_whisper import WhisperModel
finally:
if loaded_torch is missing:
sys.modules.pop("torch", None)
return WhisperModel
class SpeechRecognizer:
def __init__(
self,
model_name: str,
device: str,
compute_type: str,
language: str | None,
model_factory: Callable[..., Any] | None = None,
) -> None:
self.model_name = model_name
self.device = device
self.compute_type = compute_type
self.language = language
self._model_factory = model_factory
self._model: Any | None = None
self._lock = threading.Lock()
def _create_model(self) -> Any:
factory = self._model_factory
if factory is None:
factory = load_whisper_model_class()
return factory(
self.model_name,
device=self.device,
compute_type=self.compute_type,
)
def transcribe(self, audio_path: str | Path) -> str:
try:
with self._lock:
if self._model is None:
self._model = self._create_model()
segments, _info = self._model.transcribe(
str(audio_path),
language=self.language,
beam_size=5,
vad_filter=True,
)
parts = [
segment.text.strip()
for segment in segments
if segment.text.strip()
]
except Exception as exc:
raise SpeechRecognitionError("Voice transcription failed") from exc
return " ".join(parts)

490
assistant_bot/storage.py Normal file
View File

@@ -0,0 +1,490 @@
import re
import sqlite3
from contextlib import contextmanager
from datetime import datetime
from pathlib import Path
from typing import Any
from .config import get_default_model
from .time_utils import to_utc_iso, utc_now
class AssistantStorage:
def __init__(self, path: Path) -> None:
self.path = path
self.path.parent.mkdir(parents=True, exist_ok=True)
self._init_db()
def _connect(self) -> sqlite3.Connection:
connection = sqlite3.connect(self.path)
connection.row_factory = sqlite3.Row
return connection
@contextmanager
def _connection(self):
connection = self._connect()
try:
yield connection
connection.commit()
except Exception:
connection.rollback()
raise
finally:
connection.close()
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 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 add_conversation_exchange(
self,
user_id: int,
chat_id: int,
user_content: str,
assistant_content: str,
) -> None:
now = to_utc_iso(utc_now())
with self._connection() as connection:
connection.executemany(
"""
INSERT INTO conversation_messages(user_id, chat_id, role, content, created_at)
VALUES (?, ?, ?, ?, ?)
""",
(
(user_id, chat_id, "user", user_content, now),
(user_id, chat_id, "assistant", assistant_content, now),
),
)
def list_conversation_messages(
self,
user_id: int,
chat_id: int,
limit: int = 12,
) -> list[dict[str, Any]]:
with self._connection() as connection:
context_row = connection.execute(
"""
SELECT started_after_id
FROM conversation_contexts
WHERE user_id = ? AND chat_id = ?
""",
(user_id, chat_id),
).fetchone()
started_after_id = int(context_row["started_after_id"]) if context_row else 0
rows = connection.execute(
"""
SELECT id, role, content, created_at
FROM conversation_messages
WHERE user_id = ? AND chat_id = ? AND id > ?
ORDER BY id DESC
LIMIT ?
""",
(user_id, chat_id, started_after_id, limit),
).fetchall()
return [dict(row) for row in reversed(rows)]
def start_new_conversation(self, user_id: int, chat_id: int) -> None:
now = to_utc_iso(utc_now())
with self._connection() as connection:
row = connection.execute(
"""
SELECT COALESCE(MAX(id), 0) AS last_id
FROM conversation_messages
WHERE user_id = ? AND chat_id = ?
""",
(user_id, chat_id),
).fetchone()
last_id = int(row["last_id"])
connection.execute(
"""
INSERT INTO conversation_contexts(user_id, chat_id, started_after_id, updated_at)
VALUES (?, ?, ?, ?)
ON CONFLICT(user_id, chat_id) DO UPDATE SET
started_after_id = excluded.started_after_id,
updated_at = excluded.updated_at
""",
(user_id, chat_id, last_id, now),
)
def search_conversation_messages(
self,
user_id: int,
chat_id: int,
query: str,
limit: int = 10,
) -> list[dict[str, Any]]:
normalized_query = query.casefold().strip()
terms = list(
dict.fromkeys(
part for part in re.findall(r"\w+", normalized_query) if len(part) > 1
)
)
if not terms:
return []
with self._connection() as connection:
rows = connection.execute(
"""
SELECT id, role, content, created_at
FROM conversation_messages
WHERE user_id = ? AND chat_id = ?
ORDER BY id DESC
""",
(user_id, chat_id),
).fetchall()
if not rows:
return []
newest_id = int(rows[0]["id"])
matches: list[tuple[float, dict[str, Any]]] = []
for row in rows:
content = str(row["content"]).casefold()
matched_terms = sum(term in content for term in terms)
if not matched_terms:
continue
coverage = matched_terms / len(terms)
phrase_bonus = 1.0 if len(normalized_query) > 2 and normalized_query in content else 0.0
age = max(0, newest_id - int(row["id"]))
recency_weight = 1.0 / (1.0 + age / 50.0)
score = (coverage + phrase_bonus) * (0.5 + 0.5 * recency_weight)
item = dict(row)
item["score"] = round(score, 4)
matches.append((score, item))
matches.sort(key=lambda item: (item[0], int(item[1]["id"])), reverse=True)
return [item for _score, item in matches[: max(1, limit)]]
def conversation_message_window(
self,
user_id: int,
chat_id: int,
message_id: int,
radius: int = 2,
) -> list[dict[str, Any]]:
with self._connection() as connection:
before = connection.execute(
"""
SELECT id, role, content, created_at
FROM conversation_messages
WHERE user_id = ? AND chat_id = ? AND id <= ?
ORDER BY id DESC
LIMIT ?
""",
(user_id, chat_id, message_id, radius + 1),
).fetchall()
after = connection.execute(
"""
SELECT id, role, content, created_at
FROM conversation_messages
WHERE user_id = ? AND chat_id = ? AND id > ?
ORDER BY id ASC
LIMIT ?
""",
(user_id, chat_id, message_id, radius),
).fetchall()
rows = [*reversed(before), *after]
return [dict(row) for row in rows]
def get_user_model(self, user_id: int, provider: str = "local") -> str:
column = self._model_column(provider)
with self._connection() as connection:
row = connection.execute(
f"SELECT {column} AS model FROM user_settings WHERE user_id = ?",
(user_id,),
).fetchone()
if row and row["model"]:
return str(row["model"])
return get_default_model(provider)
def set_user_model(
self,
user_id: int,
model: str,
provider: str = "local",
) -> None:
column = self._model_column(provider)
now = to_utc_iso(utc_now())
with self._connection() as connection:
connection.execute(
f"""
INSERT INTO user_settings(user_id, {column}, updated_at)
VALUES (?, ?, ?)
ON CONFLICT(user_id) DO UPDATE SET
{column} = excluded.{column},
updated_at = excluded.updated_at
""",
(user_id, model, now),
)
@staticmethod
def _model_column(provider: str) -> str:
try:
return {
"local": "ollama_model",
"yandex": "yandex_model",
}[provider]
except KeyError as exc:
raise ValueError(f"Unknown AI provider: {provider}") from exc
def add_memory(self, user_id: int, text: str) -> int:
with self._connection() as connection:
cursor = connection.execute(
"INSERT INTO memories(user_id, text, created_at) VALUES (?, ?, ?)",
(user_id, text, to_utc_iso(utc_now())),
)
return int(cursor.lastrowid)
def list_memories(self, user_id: int, limit: int = 20) -> list[dict[str, Any]]:
with self._connection() as connection:
rows = connection.execute(
"""
SELECT id, text, created_at
FROM memories
WHERE user_id = ?
ORDER BY created_at DESC, id DESC
LIMIT ?
""",
(user_id, limit),
).fetchall()
return [dict(row) for row in rows]
def delete_memory(self, user_id: int, memory_id: int) -> bool:
with self._connection() as connection:
cursor = connection.execute(
"DELETE FROM memories WHERE user_id = ? AND id = ?",
(user_id, memory_id),
)
return cursor.rowcount > 0
def add_note(self, user_id: int, text: str) -> int:
with self._connection() as connection:
cursor = connection.execute(
"INSERT INTO notes(user_id, text, created_at) VALUES (?, ?, ?)",
(user_id, text, to_utc_iso(utc_now())),
)
return int(cursor.lastrowid)
def list_notes(self, user_id: int, limit: int = 20) -> list[dict[str, Any]]:
with self._connection() as connection:
rows = connection.execute(
"""
SELECT id, text, created_at
FROM notes
WHERE user_id = ?
ORDER BY created_at DESC, id DESC
LIMIT ?
""",
(user_id, limit),
).fetchall()
return [dict(row) for row in rows]
def delete_note(self, user_id: int, note_id: int) -> bool:
with self._connection() as connection:
cursor = connection.execute(
"DELETE FROM notes WHERE user_id = ? AND id = ?",
(user_id, note_id),
)
return cursor.rowcount > 0
def add_reminder(
self,
user_id: int,
chat_id: int,
text: str,
remind_at_utc: datetime,
) -> int:
with self._connection() as connection:
cursor = connection.execute(
"""
INSERT INTO reminders(user_id, chat_id, text, remind_at, created_at)
VALUES (?, ?, ?, ?, ?)
""",
(
user_id,
chat_id,
text,
to_utc_iso(remind_at_utc),
to_utc_iso(utc_now()),
),
)
return int(cursor.lastrowid)
def list_reminders(self, user_id: int, limit: int = 20) -> list[dict[str, Any]]:
with self._connection() as connection:
rows = connection.execute(
"""
SELECT id, text, remind_at, status, sent_at
FROM reminders
WHERE user_id = ? AND status = 'pending'
ORDER BY remind_at ASC, id ASC
LIMIT ?
""",
(user_id, limit),
).fetchall()
return [dict(row) for row in rows]
def cancel_reminder(self, user_id: int, reminder_id: int) -> bool:
with self._connection() as connection:
cursor = connection.execute(
"""
UPDATE reminders
SET status = 'cancelled'
WHERE user_id = ? AND id = ? AND status = 'pending'
""",
(user_id, reminder_id),
)
return cursor.rowcount > 0
def due_reminders(self, now_utc: datetime, limit: int = 20) -> list[dict[str, Any]]:
with self._connection() as connection:
rows = connection.execute(
"""
SELECT id, user_id, chat_id, text, remind_at
FROM reminders
WHERE status = 'pending' AND remind_at <= ?
ORDER BY remind_at ASC, id ASC
LIMIT ?
""",
(to_utc_iso(now_utc), limit),
).fetchall()
return [dict(row) for row in rows]
def mark_reminder_sent(self, reminder_id: int) -> None:
with self._connection() as connection:
connection.execute(
"""
UPDATE reminders
SET status = 'sent', sent_at = ?
WHERE id = ? AND status = 'pending'
""",
(to_utc_iso(utc_now()), reminder_id),
)
def add_tracked_item(self, user_id: int, title: str, status: str) -> int:
now = to_utc_iso(utc_now())
with self._connection() as connection:
cursor = connection.execute(
"""
INSERT INTO tracked_items(user_id, title, status, created_at, updated_at)
VALUES (?, ?, ?, ?, ?)
""",
(user_id, title, status, now, now),
)
return int(cursor.lastrowid)
def list_tracked_items(self, user_id: int, limit: int = 30) -> list[dict[str, Any]]:
with self._connection() as connection:
rows = connection.execute(
"""
SELECT id, title, status, created_at, updated_at
FROM tracked_items
WHERE user_id = ?
ORDER BY updated_at DESC, id DESC
LIMIT ?
""",
(user_id, limit),
).fetchall()
return [dict(row) for row in rows]
def set_tracked_status(self, user_id: int, item_id: int, status: str) -> bool:
with self._connection() as connection:
cursor = connection.execute(
"""
UPDATE tracked_items
SET status = ?, updated_at = ?
WHERE user_id = ? AND id = ?
""",
(status, to_utc_iso(utc_now()), user_id, item_id),
)
return cursor.rowcount > 0
def delete_tracked_item(self, user_id: int, item_id: int) -> bool:
with self._connection() as connection:
cursor = connection.execute(
"DELETE FROM tracked_items WHERE user_id = ? AND id = ?",
(user_id, item_id),
)
return cursor.rowcount > 0

View File

@@ -0,0 +1,95 @@
from html import escape
from uuid import uuid4
from zoneinfo import ZoneInfo
from telegram import InlineQueryResultArticle, InputTextMessageContent, Message, Update
from telegram.constants import ParseMode
from telegram.ext import ContextTypes
from .ai import AIClient
from .config import HTML_FORMATS, MAX_TELEGRAM_MESSAGE_LENGTH
from .speech import SpeechTranscriber
from .storage import AssistantStorage
def make_article(
title: str,
text: str,
parse_mode: str | None = None,
) -> InlineQueryResultArticle:
return InlineQueryResultArticle(
id=f"inline_{uuid4()}",
title=title,
input_message_content=InputTextMessageContent(text, parse_mode=parse_mode),
)
def build_inline_results(query: str) -> list[InlineQueryResultArticle]:
escaped_query = escape(query)
return [
make_article("Caps", query.upper()),
*[
make_article(title, f"<{tag}>{escaped_query}</{tag}>", ParseMode.HTML)
for title, tag in HTML_FORMATS
],
]
def get_storage(context: ContextTypes.DEFAULT_TYPE) -> AssistantStorage:
return context.application.bot_data["storage"]
def get_ai_client(context: ContextTypes.DEFAULT_TYPE) -> AIClient:
return context.application.bot_data["ai_client"]
def get_tz(context: ContextTypes.DEFAULT_TYPE) -> ZoneInfo:
return context.application.bot_data["timezone"]
def get_speech_recognizer(context: ContextTypes.DEFAULT_TYPE) -> SpeechTranscriber:
return context.application.bot_data["speech_recognizer"]
def get_voice_max_duration(context: ContextTypes.DEFAULT_TYPE) -> int:
return int(context.application.bot_data["voice_max_duration"])
def command_text(context: ContextTypes.DEFAULT_TYPE) -> str:
return " ".join(context.args).strip()
def require_user_id(update: Update) -> int | None:
if update.effective_user:
return int(update.effective_user.id)
return None
async def reply_long(message: Message, text: str) -> None:
chunks = [
text[index : index + MAX_TELEGRAM_MESSAGE_LENGTH]
for index in range(0, len(text), MAX_TELEGRAM_MESSAGE_LENGTH)
] or [""]
for chunk in chunks:
await message.reply_text(chunk)
async def edit_or_reply(message: Message, text: str) -> None:
chunks = [
text[index : index + MAX_TELEGRAM_MESSAGE_LENGTH]
for index in range(0, len(text), MAX_TELEGRAM_MESSAGE_LENGTH)
] or [""]
await message.edit_text(chunks[0])
for chunk in chunks[1:]:
await message.reply_text(chunk)
def parse_positive_int(value: str) -> int | None:
try:
parsed = int(value)
except ValueError:
return None
return parsed if parsed > 0 else None

View File

@@ -0,0 +1,23 @@
from datetime import datetime, timezone
from zoneinfo import ZoneInfo
def utc_now() -> datetime:
return datetime.now(timezone.utc)
def to_utc_iso(value: datetime) -> str:
return value.astimezone(timezone.utc).isoformat(timespec="seconds")
def parse_utc_iso(value: str) -> datetime:
parsed = datetime.fromisoformat(value)
if parsed.tzinfo is None:
return parsed.replace(tzinfo=timezone.utc)
return parsed.astimezone(timezone.utc)
def format_local_dt(value: str | datetime, tz: ZoneInfo) -> str:
parsed = parse_utc_iso(value) if isinstance(value, str) else value
return parsed.astimezone(tz).strftime("%Y-%m-%d %H:%M")

115
assistant_bot/yandex_ai.py Normal file
View File

@@ -0,0 +1,115 @@
from pathlib import Path
from typing import Any
from .ai import AIClientError
from .speech import SpeechRecognitionError
class YandexAIError(AIClientError):
pass
def create_async_sdk(folder_id: str) -> Any:
try:
from yandex_ai_studio_sdk import AsyncAIStudio
except ImportError as exc:
raise RuntimeError(
"Yandex AI Studio SDK is not installed. "
"Run: python -m pip install -r requirements.txt"
) from exc
return AsyncAIStudio(folder_id=folder_id)
def create_sync_sdk(folder_id: str) -> Any:
try:
from yandex_ai_studio_sdk import AIStudio
except ImportError as exc:
raise RuntimeError(
"Yandex AI Studio SDK is not installed. "
"Run: python -m pip install -r requirements.txt"
) from exc
return AIStudio(folder_id=folder_id)
class YandexAIClient:
provider = "yandex"
display_name = "Yandex AI Studio"
def __init__(self, folder_id: str, sdk: Any | None = None) -> None:
self.folder_id = folder_id.strip("/")
self._sdk = sdk if sdk is not None else create_async_sdk(self.folder_id)
def normalize_model(self, model: str) -> str:
normalized = model.strip()
if normalized.startswith("gpt://"):
return normalized
return f"gpt://{self.folder_id}/{normalized.lstrip('/')}"
async def list_models(self) -> list[str]:
try:
models = await self._sdk.chat.completions.list()
except Exception as exc:
raise YandexAIError(f"Не удалось получить список моделей: {exc}") from exc
model_names = {
str(uri)
for model in models
if (uri := getattr(model, "uri", None))
}
return sorted(model_names)
async def chat(
self,
model: str,
messages: list[dict[str, str]],
json_mode: bool = False,
) -> str:
sdk_messages = [
{
"role": message.get("role", "user"),
"text": message.get("content", ""),
}
for message in messages
]
try:
completion = self._sdk.models.completions(self.normalize_model(model))
if json_mode:
completion = completion.configure(response_format="json")
result = await completion.run(sdk_messages, timeout=180)
content = getattr(result[0], "text", None)
except Exception as exc:
raise YandexAIError(f"Запрос к модели завершился ошибкой: {exc}") from exc
if isinstance(content, str) and content.strip():
return content.strip()
raise YandexAIError("Yandex AI Studio вернула пустой ответ.")
class YandexSpeechRecognizer:
def __init__(
self,
folder_id: str,
language: str,
model: str,
sdk: Any | None = None,
) -> None:
self.folder_id = folder_id.strip("/")
self.language = language
self.model = model
self._sdk = sdk if sdk is not None else create_sync_sdk(self.folder_id)
def transcribe(self, audio_path: str | Path) -> str:
try:
audio = Path(audio_path).read_bytes()
recognizer = self._sdk.speechkit.speech_to_text(
audio_format=self._sdk.speechkit.AudioFormat.OGG_OPUS,
language_codes=self.language,
model=self.model,
)
result = recognizer.run(audio, timeout=180)
text = getattr(result, "text", None)
except Exception as exc:
raise SpeechRecognitionError("Yandex SpeechKit transcription failed") from exc
return text.strip() if isinstance(text, str) else ""

5
main.py Normal file
View File

@@ -0,0 +1,5 @@
from assistant_bot.application import run
if __name__ == "__main__":
run()

1
requirements-common.txt Normal file
View File

@@ -0,0 +1 @@
python-telegram-bot==22.8

4
requirements-local.txt Normal file
View File

@@ -0,0 +1,4 @@
-r requirements-common.txt
faster-whisper==1.2.1
ctranslate2==4.6.0
setuptools==80.10.2

2
requirements-yandex.txt Normal file
View File

@@ -0,0 +1,2 @@
-r requirements-common.txt
yandex-ai-studio-sdk==0.22.1

3
requirements.txt Normal file
View File

@@ -0,0 +1,3 @@
-r requirements-local.txt
-r requirements-yandex.txt
tzdata>=2024.1; platform_system == "Windows"

1
tests/__init__.py Normal file
View File

@@ -0,0 +1 @@
"""Project tests."""

71
tests/test_config.py Normal file
View File

@@ -0,0 +1,71 @@
import os
import unittest
from unittest.mock import patch
from assistant_bot import config
class WhisperConfigTests(unittest.TestCase):
def test_gpu_int8_defaults(self) -> None:
variable_names = (
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(config.get_whisper_device(), "cuda")
self.assertEqual(config.get_whisper_compute_type(), "int8")
class AssistantModeConfigTests(unittest.TestCase):
def test_local_mode_is_default(self) -> None:
with patch("assistant_bot.config.load_env_file"), patch.dict(
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:
with patch("assistant_bot.config.load_env_file"), patch.dict(
os.environ,
{config.MODE_ENV_NAME: "YANDEX"},
clear=False,
):
self.assertEqual(config.get_assistant_mode(), "yandex")
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):
config.get_assistant_mode()
def test_yandex_model_is_full_gpt_uri(self) -> None:
with patch("assistant_bot.config.load_env_file"), patch.dict(
os.environ,
{
config.YANDEX_CLOUD_FOLDER_ENV_NAME: "folder-id",
config.YANDEX_CLOUD_MODEL_ENV_NAME: "yandexgpt/latest",
},
clear=False,
):
self.assertEqual(
config.get_default_yandex_model(),
"gpt://folder-id/yandexgpt/latest",
)
if __name__ == "__main__":
unittest.main()

148
tests/test_core.py Normal file
View File

@@ -0,0 +1,148 @@
import sqlite3
import tempfile
import unittest
from contextlib import closing
from datetime import datetime, timezone
from pathlib import Path
from zoneinfo import ZoneInfo
from assistant_bot.agent import execute_agent_tool, parse_agent_decision
from assistant_bot.reminders import parse_reminder, unit_to_timedelta
from assistant_bot.storage import AssistantStorage
class ReminderParserTests(unittest.TestCase):
def test_supported_relative_unit(self) -> None:
self.assertEqual(unit_to_timedelta(2, "часа").total_seconds(), 7200)
def test_absolute_datetime(self) -> None:
result = parse_reminder(
"2030-01-02 10:30 проверить отчет",
ZoneInfo("UTC"),
)
self.assertIsNotNone(result)
assert result is not None
self.assertEqual(
result.remind_at_utc,
datetime(2030, 1, 2, 10, 30, tzinfo=timezone.utc),
)
self.assertEqual(result.text, "проверить отчет")
class AgentDecisionTests(unittest.TestCase):
def test_tool_call_json(self) -> None:
decision = parse_agent_decision(
'{"tool_calls":[{"name":"create_note","arguments":{"text":"идея"}}]}'
)
self.assertIsNone(decision.final)
self.assertEqual(decision.tool_calls[0]["name"], "create_note")
self.assertEqual(decision.tool_calls[0]["arguments"]["text"], "идея")
def test_context_reset_flag(self) -> None:
decision = parse_agent_decision(
'{"final":"Перейдем к новой теме","reset_context":true}'
)
self.assertEqual(decision.final, "Перейдем к новой теме")
self.assertTrue(decision.reset_context)
class StorageTests(unittest.TestCase):
def test_legacy_user_settings_gets_yandex_model_column(self) -> None:
with tempfile.TemporaryDirectory() as directory:
database_path = Path(directory) / "assistant.sqlite3"
with closing(sqlite3.connect(database_path)) as connection:
connection.execute(
"""
CREATE TABLE user_settings (
user_id INTEGER PRIMARY KEY,
ollama_model TEXT,
updated_at TEXT NOT NULL
)
"""
)
connection.commit()
storage = AssistantStorage(database_path)
storage.set_user_model(42, "yandexgpt", "yandex")
self.assertEqual(storage.get_user_model(42, "yandex"), "yandexgpt")
def test_provider_models_are_stored_separately(self) -> None:
with tempfile.TemporaryDirectory() as directory:
storage = AssistantStorage(Path(directory) / "assistant.sqlite3")
storage.set_user_model(42, "qwen3.5:9b", "local")
storage.set_user_model(42, "yandexgpt", "yandex")
self.assertEqual(storage.get_user_model(42, "local"), "qwen3.5:9b")
self.assertEqual(storage.get_user_model(42, "yandex"), "yandexgpt")
def test_memory_crud(self) -> None:
with tempfile.TemporaryDirectory() as directory:
storage = AssistantStorage(Path(directory) / "assistant.sqlite3")
memory_id = storage.add_memory(42, "короткие ответы")
self.assertEqual(storage.list_memories(42)[0]["id"], memory_id)
self.assertTrue(storage.delete_memory(42, memory_id))
self.assertEqual(storage.list_memories(42), [])
def test_new_context_preserves_searchable_archive(self) -> None:
with tempfile.TemporaryDirectory() as directory:
storage = AssistantStorage(Path(directory) / "assistant.sqlite3")
storage.add_conversation_exchange(
42,
100,
"Обсудим проект Альфа",
"Какой аспект проекта интересует?",
)
storage.start_new_conversation(42, 100)
storage.add_conversation_exchange(
42,
100,
"Снова обсуждаем проект Альфа",
"Продолжаем обсуждение проекта.",
)
storage.add_conversation_exchange(42, 200, "Другой чат", "Другой ответ")
messages = storage.list_conversation_messages(42, 100)
matches = storage.search_conversation_messages(42, 100, "проект Альфа")
self.assertEqual(
[(row["role"], row["content"]) for row in messages],
[
("user", "Снова обсуждаем проект Альфа"),
("assistant", "Продолжаем обсуждение проекта."),
],
)
self.assertEqual(matches[0]["content"], "Снова обсуждаем проект Альфа")
self.assertIn("Обсудим проект Альфа", [row["content"] for row in matches])
self.assertEqual(len(storage.list_conversation_messages(42, 200)), 2)
archive_window = storage.conversation_message_window(
42,
100,
int(matches[-1]["id"]),
)
self.assertIn(
"Какой аспект проекта интересует?",
[row["content"] for row in archive_window],
)
tool_result = execute_agent_tool(
"search_conversation",
{"query": "проект Альфа", "limit": 2},
storage,
user_id=42,
chat_id=100,
tz=ZoneInfo("UTC"),
)
self.assertTrue(tool_result["ok"])
self.assertTrue(tool_result["discussions"])
self.assertTrue(tool_result["discussions"][0]["messages"])
if __name__ == "__main__":
unittest.main()

68
tests/test_speech.py Normal file
View File

@@ -0,0 +1,68 @@
import tempfile
import unittest
from pathlib import Path
from assistant_bot.speech import SpeechRecognitionError, SpeechRecognizer
class FakeSegment:
def __init__(self, text: str) -> None:
self.text = text
class SpeechRecognizerTests(unittest.TestCase):
def test_transcribe_joins_non_empty_segments(self) -> None:
factory_calls = []
transcribe_calls = []
class FakeModel:
def transcribe(self, audio_path, **kwargs):
transcribe_calls.append((audio_path, kwargs))
return iter(
[FakeSegment(" Привет "), FakeSegment(""), FakeSegment("мир")]
), None
def model_factory(*args, **kwargs):
factory_calls.append((args, kwargs))
return FakeModel()
recognizer = SpeechRecognizer(
model_name="small",
device="cpu",
compute_type="int8",
language="ru",
model_factory=model_factory,
)
with tempfile.TemporaryDirectory() as directory:
audio_path = Path(directory) / "voice.ogg"
result = recognizer.transcribe(audio_path)
self.assertEqual(result, "Привет мир")
self.assertEqual(
factory_calls[0],
(("small",), {"device": "cpu", "compute_type": "int8"}),
)
self.assertEqual(
transcribe_calls[0][1],
{"language": "ru", "beam_size": 5, "vad_filter": True},
)
def test_transcribe_wraps_model_errors(self) -> None:
def failing_factory(*_args, **_kwargs):
raise RuntimeError("model is unavailable")
recognizer = SpeechRecognizer(
model_name="small",
device="cpu",
compute_type="int8",
language=None,
model_factory=failing_factory,
)
with self.assertRaises(SpeechRecognitionError):
recognizer.transcribe("voice.ogg")
if __name__ == "__main__":
unittest.main()

View File

@@ -0,0 +1,75 @@
import unittest
from pathlib import Path
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
from assistant_bot.handlers import private_voice
class VoiceHandlerTests(unittest.IsolatedAsyncioTestCase):
@staticmethod
def make_update_and_context(duration: int, transcript: str = "Напомни позвонить"):
status_message = SimpleNamespace(edit_text=AsyncMock())
message = SimpleNamespace(
chat_id=100,
voice=SimpleNamespace(duration=duration, file_id="voice-file-id"),
reply_text=AsyncMock(return_value=status_message),
)
update = SimpleNamespace(message=message)
telegram_file = SimpleNamespace(download_to_drive=AsyncMock())
bot = SimpleNamespace(
get_file=AsyncMock(return_value=telegram_file),
send_chat_action=AsyncMock(),
)
recognizer = MagicMock()
recognizer.transcribe.return_value = transcript
context = SimpleNamespace(
bot=bot,
application=SimpleNamespace(
bot_data={
"speech_recognizer": recognizer,
"voice_max_duration": 120,
}
),
)
return update, context, status_message, recognizer, telegram_file
async def test_rejects_voice_message_over_duration_limit(self) -> None:
update, context, _status, recognizer, _telegram_file = (
self.make_update_and_context(duration=121)
)
await private_voice(update, context)
update.message.reply_text.assert_awaited_once_with(
"Голосовое сообщение слишком длинное. Максимум: 120 сек."
)
recognizer.transcribe.assert_not_called()
async def test_transcribes_voice_and_passes_text_to_agent(self) -> None:
update, context, status, recognizer, telegram_file = (
self.make_update_and_context(duration=10)
)
with patch(
"assistant_bot.handlers.run_agent_prompt",
new_callable=AsyncMock,
) as run_agent:
await private_voice(update, context)
recognizer.transcribe.assert_called_once()
temporary_path = Path(recognizer.transcribe.call_args.args[0])
self.assertFalse(temporary_path.exists())
telegram_file.download_to_drive.assert_awaited_once_with(
custom_path=temporary_path
)
status.edit_text.assert_awaited_once_with("Распознано: Напомни позвонить")
run_agent.assert_awaited_once_with(
update,
context,
"Напомни позвонить",
)
if __name__ == "__main__":
unittest.main()

132
tests/test_yandex_ai.py Normal file
View File

@@ -0,0 +1,132 @@
import unittest
from pathlib import Path
from types import SimpleNamespace
from unittest.mock import patch
from assistant_bot.yandex_ai import YandexAIClient, YandexSpeechRecognizer
class FakeCompletion:
def __init__(self) -> None:
self.configuration = None
self.messages = None
self.timeout = None
def configure(self, **kwargs):
self.configuration = kwargs
return self
async def run(self, messages, timeout):
self.messages = messages
self.timeout = timeout
return [SimpleNamespace(text=" готово ")]
class YandexAIClientTests(unittest.IsolatedAsyncioTestCase):
async def test_chat_uses_sdk_messages_and_json_mode(self) -> None:
completion = FakeCompletion()
requested_models = []
def completions(model):
requested_models.append(model)
return completion
sdk = SimpleNamespace(
models=SimpleNamespace(completions=completions),
chat=SimpleNamespace(
completions=SimpleNamespace(
list=self._list_models,
)
),
)
client = YandexAIClient(folder_id="folder-id", sdk=sdk)
answer = await client.chat(
"yandexgpt",
[
{"role": "system", "content": "Отвечай кратко"},
{"role": "user", "content": "Привет"},
],
json_mode=True,
)
self.assertEqual(answer, "готово")
self.assertEqual(requested_models, ["gpt://folder-id/yandexgpt"])
self.assertEqual(completion.configuration, {"response_format": "json"})
self.assertEqual(
completion.messages,
[
{"role": "system", "text": "Отвечай кратко"},
{"role": "user", "text": "Привет"},
],
)
self.assertEqual(completion.timeout, 180)
async def test_list_models_returns_sorted_uris(self) -> None:
sdk = SimpleNamespace(
chat=SimpleNamespace(
completions=SimpleNamespace(
list=self._list_models,
)
)
)
client = YandexAIClient(folder_id="folder-id", sdk=sdk)
self.assertEqual(
await client.list_models(),
["gpt://folder/alice-ai/latest", "gpt://folder/yandexgpt/latest"],
)
@staticmethod
async def _list_models():
return [
SimpleNamespace(uri="gpt://folder/yandexgpt/latest"),
SimpleNamespace(uri="gpt://folder/alice-ai/latest"),
]
class YandexSpeechRecognizerTests(unittest.TestCase):
def test_transcribe_uses_speechkit_ogg_opus(self) -> None:
calls = []
class FakeRecognizer:
def run(self, audio, timeout):
calls.append((audio, timeout))
return SimpleNamespace(text=" Привет, мир ")
format_value = object()
def speech_to_text(**kwargs):
calls.append(kwargs)
return FakeRecognizer()
sdk = SimpleNamespace(
speechkit=SimpleNamespace(
AudioFormat=SimpleNamespace(OGG_OPUS=format_value),
speech_to_text=speech_to_text,
)
)
recognizer = YandexSpeechRecognizer(
folder_id="folder-id",
language="ru-RU",
model="general",
sdk=sdk,
)
with patch.object(Path, "read_bytes", return_value=b"ogg-data"):
result = recognizer.transcribe("voice.ogg")
self.assertEqual(result, "Привет, мир")
self.assertEqual(
calls[0],
{
"audio_format": format_value,
"language_codes": "ru-RU",
"model": "general",
},
)
self.assertEqual(calls[1], (b"ogg-data", 180))
if __name__ == "__main__":
unittest.main()