commit 6667e0d253d53afb29026b253681daa77ab5cb43 Author: kandrusyak Date: Sat Jul 25 15:07:53 2026 +0300 init diff --git a/.air/settings.json b/.air/settings.json new file mode 100644 index 0000000..a7858d1 --- /dev/null +++ b/.air/settings.json @@ -0,0 +1,3 @@ +{ + "editor.guides": [] +} \ No newline at end of file diff --git a/.dockerignore b/.dockerignore new file mode 100644 index 0000000..ea941b2 --- /dev/null +++ b/.dockerignore @@ -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 diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..a905d6d --- /dev/null +++ b/.gitignore @@ -0,0 +1,5 @@ +.env +.venv/ +__pycache__/ +*.py[cod] +assistant_data.sqlite3* diff --git a/.idea/.gitignore b/.idea/.gitignore new file mode 100644 index 0000000..b58b603 --- /dev/null +++ b/.idea/.gitignore @@ -0,0 +1,5 @@ +# Default ignored files +/shelf/ +/workspace.xml +# Editor-based HTTP Client requests +/httpRequests/ diff --git a/.idea/KAndrusyak_Bot.iml b/.idea/KAndrusyak_Bot.iml new file mode 100644 index 0000000..364598c --- /dev/null +++ b/.idea/KAndrusyak_Bot.iml @@ -0,0 +1,10 @@ + + + + + + + + + + \ No newline at end of file diff --git a/.idea/inspectionProfiles/profiles_settings.xml b/.idea/inspectionProfiles/profiles_settings.xml new file mode 100644 index 0000000..105ce2d --- /dev/null +++ b/.idea/inspectionProfiles/profiles_settings.xml @@ -0,0 +1,6 @@ + + + + \ No newline at end of file diff --git a/.idea/misc.xml b/.idea/misc.xml new file mode 100644 index 0000000..16ff1e7 --- /dev/null +++ b/.idea/misc.xml @@ -0,0 +1,7 @@ + + + + + + \ No newline at end of file diff --git a/.idea/modules.xml b/.idea/modules.xml new file mode 100644 index 0000000..4117856 --- /dev/null +++ b/.idea/modules.xml @@ -0,0 +1,8 @@ + + + + + + + + \ No newline at end of file diff --git a/.idea/vcs.xml b/.idea/vcs.xml new file mode 100644 index 0000000..35eb1dd --- /dev/null +++ b/.idea/vcs.xml @@ -0,0 +1,6 @@ + + + + + + \ No newline at end of file diff --git a/AGENTS.md b/AGENTS.md new file mode 100644 index 0000000..1b66c41 --- /dev/null +++ b/AGENTS.md @@ -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. diff --git a/Dockerfile b/Dockerfile new file mode 100644 index 0000000..53af450 --- /dev/null +++ b/Dockerfile @@ -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"] diff --git a/README.md b/README.md new file mode 100644 index 0000000..8c72c7c --- /dev/null +++ b/README.md @@ -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:///`. + +## 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 +``` diff --git a/assistant_bot/__init__.py b/assistant_bot/__init__.py new file mode 100644 index 0000000..3fd671f --- /dev/null +++ b/assistant_bot/__init__.py @@ -0,0 +1,2 @@ +"""Personal Telegram assistant package.""" + diff --git a/assistant_bot/__main__.py b/assistant_bot/__main__.py new file mode 100644 index 0000000..0c6c1c3 --- /dev/null +++ b/assistant_bot/__main__.py @@ -0,0 +1,5 @@ +from .application import run + + +if __name__ == "__main__": + run() diff --git a/assistant_bot/agent.py b/assistant_bot/agent.py new file mode 100644 index 0000000..4b70d52 --- /dev/null +++ b/assistant_bot/agent.py @@ -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) diff --git a/assistant_bot/ai.py b/assistant_bot/ai.py new file mode 100644 index 0000000..b0555ae --- /dev/null +++ b/assistant_bot/ai.py @@ -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: + ... diff --git a/assistant_bot/application.py b/assistant_bot/application.py new file mode 100644 index 0000000..f5cfc68 --- /dev/null +++ b/assistant_bot/application.py @@ -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) diff --git a/assistant_bot/config.py b/assistant_bot/config.py new file mode 100644 index 0000000..8c0b2f8 --- /dev/null +++ b/assistant_bot/config.py @@ -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 diff --git a/assistant_bot/handlers.py b/assistant_bot/handlers.py new file mode 100644 index 0000000..3ba4b4b --- /dev/null +++ b/assistant_bot/handlers.py @@ -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") diff --git a/assistant_bot/jobs.py b/assistant_bot/jobs.py new file mode 100644 index 0000000..93c0a3e --- /dev/null +++ b/assistant_bot/jobs.py @@ -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)) + diff --git a/assistant_bot/models.py b/assistant_bot/models.py new file mode 100644 index 0000000..ac56ebe --- /dev/null +++ b/assistant_bot/models.py @@ -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 diff --git a/assistant_bot/ollama.py b/assistant_bot/ollama.py new file mode 100644 index 0000000..159aa95 --- /dev/null +++ b/assistant_bot/ollama.py @@ -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 diff --git a/assistant_bot/prompts.py b/assistant_bot/prompts.py new file mode 100644 index 0000000..9f4e9b6 --- /dev/null +++ b/assistant_bot/prompts.py @@ -0,0 +1,7 @@ +ASSISTANT_SYSTEM_PROMPT = ( + "Ты персональный AI ассистент в Telegram. Отвечай кратко, по делу и на русском, " + "если пользователь не попросил другой язык. Используй долговременный контекст " + "только как вспомогательную информацию, не выдумывай факты и явно говори, " + "когда данных недостаточно." +) + diff --git a/assistant_bot/reminders.py b/assistant_bot/reminders.py new file mode 100644 index 0000000..7830920 --- /dev/null +++ b/assistant_bot/reminders.py @@ -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 + diff --git a/assistant_bot/speech.py b/assistant_bot/speech.py new file mode 100644 index 0000000..4a9e58a --- /dev/null +++ b/assistant_bot/speech.py @@ -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) diff --git a/assistant_bot/storage.py b/assistant_bot/storage.py new file mode 100644 index 0000000..edb9232 --- /dev/null +++ b/assistant_bot/storage.py @@ -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 diff --git a/assistant_bot/telegram_utils.py b/assistant_bot/telegram_utils.py new file mode 100644 index 0000000..31e9351 --- /dev/null +++ b/assistant_bot/telegram_utils.py @@ -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}", 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 diff --git a/assistant_bot/time_utils.py b/assistant_bot/time_utils.py new file mode 100644 index 0000000..c1d555a --- /dev/null +++ b/assistant_bot/time_utils.py @@ -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") + diff --git a/assistant_bot/yandex_ai.py b/assistant_bot/yandex_ai.py new file mode 100644 index 0000000..6d96e1d --- /dev/null +++ b/assistant_bot/yandex_ai.py @@ -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 "" diff --git a/main.py b/main.py new file mode 100644 index 0000000..29c5dd6 --- /dev/null +++ b/main.py @@ -0,0 +1,5 @@ +from assistant_bot.application import run + + +if __name__ == "__main__": + run() diff --git a/requirements-common.txt b/requirements-common.txt new file mode 100644 index 0000000..acbc4c7 --- /dev/null +++ b/requirements-common.txt @@ -0,0 +1 @@ +python-telegram-bot==22.8 diff --git a/requirements-local.txt b/requirements-local.txt new file mode 100644 index 0000000..306f57f --- /dev/null +++ b/requirements-local.txt @@ -0,0 +1,4 @@ +-r requirements-common.txt +faster-whisper==1.2.1 +ctranslate2==4.6.0 +setuptools==80.10.2 diff --git a/requirements-yandex.txt b/requirements-yandex.txt new file mode 100644 index 0000000..1c0b851 --- /dev/null +++ b/requirements-yandex.txt @@ -0,0 +1,2 @@ +-r requirements-common.txt +yandex-ai-studio-sdk==0.22.1 diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..9d0256d --- /dev/null +++ b/requirements.txt @@ -0,0 +1,3 @@ +-r requirements-local.txt +-r requirements-yandex.txt +tzdata>=2024.1; platform_system == "Windows" diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 0000000..db78820 --- /dev/null +++ b/tests/__init__.py @@ -0,0 +1 @@ +"""Project tests.""" diff --git a/tests/test_config.py b/tests/test_config.py new file mode 100644 index 0000000..7b37f00 --- /dev/null +++ b/tests/test_config.py @@ -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() diff --git a/tests/test_core.py b/tests/test_core.py new file mode 100644 index 0000000..4c0d91e --- /dev/null +++ b/tests/test_core.py @@ -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() diff --git a/tests/test_speech.py b/tests/test_speech.py new file mode 100644 index 0000000..ed66b1c --- /dev/null +++ b/tests/test_speech.py @@ -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() diff --git a/tests/test_voice_handler.py b/tests/test_voice_handler.py new file mode 100644 index 0000000..1b9d4e4 --- /dev/null +++ b/tests/test_voice_handler.py @@ -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() diff --git a/tests/test_yandex_ai.py b/tests/test_yandex_ai.py new file mode 100644 index 0000000..1779209 --- /dev/null +++ b/tests/test_yandex_ai.py @@ -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()