Compare commits

...

10 Commits

Author SHA1 Message Date
kandrusyak
701cdd35d2 refactor: unify Telegram delivery and add quality checks
Some checks failed
quality / test (3.10) (push) Has been cancelled
quality / test (3.12) (push) Has been cancelled
2026-07-27 17:54:25 +03:00
kandrusyak
c20d893255 refactor: isolate database migrations and model defaults 2026-07-27 17:51:16 +03:00
kandrusyak
ab029a4c67 refactor: centralize settings and application services 2026-07-27 17:47:39 +03:00
kandrusyak
248e90faf5 refactor: centralize agent tool registry 2026-07-27 17:44:37 +03:00
kandrusyak
79ed741a3b refactor: remove unreachable legacy command handlers 2026-07-27 17:41:33 +03:00
kandrusyak
3cc60c69e1 test: lock down application and agent contracts 2026-07-27 17:37:45 +03:00
kandrusyak
4b7a50d0cc add typing simulation 2026-07-27 17:20:32 +03:00
kandrusyak
275eca8dc5 add formating 2026-07-27 17:14:07 +03:00
kandrusyak
7b7ff51d0f add skill 2026-07-27 17:04:36 +03:00
kandrusyak
053b124a9a add password authentication for bot users 2026-07-25 15:57:30 +03:00
33 changed files with 2039 additions and 996 deletions

View File

@@ -0,0 +1,7 @@
---
name: grill-me
description: A relentless interview to sharpen a plan or design.
disable-model-invocation: true
---
Run a `/grilling` session.

View File

@@ -0,0 +1,5 @@
interface:
display_name: "Grill Me"
short_description: "Sharpen a plan through interview"
policy:
allow_implicit_invocation: false

View File

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

29
.github/workflows/quality.yml vendored Normal file
View File

@@ -0,0 +1,29 @@
name: quality
on:
push:
pull_request:
jobs:
test:
runs-on: ubuntu-latest
strategy:
matrix:
python-version: ["3.10", "3.12"]
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: ${{ matrix.python-version }}
cache: pip
- name: Install test dependencies
run: |
python -m pip install --upgrade pip
python -m pip install -r requirements-common.txt -r requirements-dev.txt
- name: Lint
run: python -m ruff check .
- name: Test with coverage
run: |
python -m coverage run -m unittest discover -v
python -m coverage report

5
.gitignore vendored
View File

@@ -1,5 +1,10 @@
.env .env
.venv/ .venv/
.idea/
__pycache__/ __pycache__/
*.py[cod] *.py[cod]
assistant_data.sqlite3* assistant_data.sqlite3*
.coverage
.mypy_cache/
.pytest_cache/
.ruff_cache/

5
.idea/.gitignore generated vendored
View File

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

View File

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

View File

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

7
.idea/misc.xml generated
View File

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

8
.idea/modules.xml generated
View File

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

6
.idea/vcs.xml generated
View File

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

View File

@@ -20,9 +20,12 @@ Telegram-ассистент с двумя режимами AI: локальны
- `assistant_bot/ollama.py` — HTTP-клиент Ollama; - `assistant_bot/ollama.py` — HTTP-клиент Ollama;
- `assistant_bot/yandex_ai.py` — адаптер Yandex AI Studio SDK; - `assistant_bot/yandex_ai.py` — адаптер Yandex AI Studio SDK;
- `assistant_bot/agent.py` — агентный цикл и выполнение внутренних tools; - `assistant_bot/agent.py` — агентный цикл и выполнение внутренних tools;
- `assistant_bot/agent_tools.py` — единый реестр и исполнители agent tools;
- `assistant_bot/services.py` — типизированный контейнер сервисов приложения;
- `assistant_bot/handlers.py` — Telegram-команды и сообщения; - `assistant_bot/handlers.py` — Telegram-команды и сообщения;
- `assistant_bot/reminders.py` — разбор времени напоминаний; - `assistant_bot/reminders.py` — разбор времени напоминаний;
- `assistant_bot/jobs.py` — фоновые задачи; - `assistant_bot/jobs.py` — фоновые задачи;
- `assistant_bot/migrations.py` — версионируемые миграции SQLite;
- `tests/` — модульные тесты ядра. - `tests/` — модульные тесты ядра.
## Запуск ## Запуск
@@ -39,9 +42,14 @@ python main.py
```dotenv ```dotenv
BOT_TOKEN=... BOT_TOKEN=...
ASSISTANT_PASSWORD=...
ASSISTANT_MODE=local ASSISTANT_MODE=local
``` ```
`ASSISTANT_PASSWORD` — общий пароль доступа к боту. Новый пользователь должен
один раз отправить его боту в личном чате; после успешной проверки авторизация
сохраняется в SQLite, а сам пароль в базу данных не записывается.
`ASSISTANT_MODE` принимает `local` (значение по умолчанию) или `yandex`. `ASSISTANT_MODE` принимает `local` (значение по умолчанию) или `yandex`.
Общие дополнительные настройки: `ASSISTANT_DB`, `ASSISTANT_TIMEZONE` и Общие дополнительные настройки: `ASSISTANT_DB`, `ASSISTANT_TIMEZONE` и
`VOICE_MAX_DURATION_SECONDS`. `VOICE_MAX_DURATION_SECONDS`.
@@ -146,3 +154,12 @@ Ollama должна быть доступна контейнеру по адре
conda activate kandrusyak_bot conda activate kandrusyak_bot
python -m unittest discover -v python -m unittest discover -v
``` ```
Для локальных проверок качества установите dev-зависимости:
```powershell
python -m pip install -r requirements-common.txt -r requirements-dev.txt
python -m ruff check .
python -m coverage run -m unittest discover -v
python -m coverage report
```

View File

@@ -1,29 +1,29 @@
import json import json
import logging import logging
import re import re
from datetime import datetime, timezone from datetime import datetime
from typing import Any from typing import Any
from zoneinfo import ZoneInfo from zoneinfo import ZoneInfo
from telegram import Update from telegram import Update
from telegram.constants import ChatAction
from telegram.ext import ContextTypes from telegram.ext import ContextTypes
from .agent_tools import execute_agent_tool, render_agent_tool_catalog
from .ai import AIClientError from .ai import AIClientError
from .config import CONVERSATION_HISTORY_LIMIT, MAX_AGENT_STEPS from .config import CONVERSATION_HISTORY_LIMIT, MAX_AGENT_STEPS
from .models import AgentDecision, ReminderParseResult from .models import AgentDecision
from .prompts import ASSISTANT_SYSTEM_PROMPT from .prompts import ASSISTANT_SYSTEM_PROMPT
from .reminders import parse_reminder
from .storage import AssistantStorage from .storage import AssistantStorage
from .telegram_utils import ( from .telegram_utils import (
edit_or_reply, escape_markdown_text,
get_ai_client, get_ai_client,
get_storage, get_storage,
get_tz, get_tz,
parse_positive_int, reply_markdown,
require_user_id, require_user_id,
typing_action,
) )
from .time_utils import format_local_dt, to_utc_iso, utc_now from .time_utils import format_local_dt
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -80,32 +80,18 @@ def build_agent_tool_prompt(tz: ZoneInfo) -> str:
"Ты работаешь как агент с внутренними tools. Пользователь пишет свободным текстом, " "Ты работаешь как агент с внутренними tools. Пользователь пишет свободным текстом, "
"а Telegram-команды ему не нужны.\n" "а Telegram-команды ему не нужны.\n"
f"Текущее локальное время: {now_local}. Таймзона: {tz.key}.\n\n" f"Текущее локальное время: {now_local}. Таймзона: {tz.key}.\n\n"
"Отвечай СТРОГО одним JSON-объектом без Markdown и без текста вокруг.\n" "Отвечай СТРОГО одним JSON-объектом без Markdown-блока и без текста вокруг.\n"
"Если нужно выполнить действие, верни tool_calls. Если действие уже выполнено " "Если нужно выполнить действие, верни tool_calls. Если действие уже выполнено "
"или tool не нужен, верни final.\n\n" "или tool не нужен, верни final.\n\n"
"Форматы ответа:\n" "Форматы ответа:\n"
'{"tool_calls":[{"name":"create_note","arguments":{"text":"..."}}],"reset_context":false}\n' '{"tool_calls":[{"name":"create_note","arguments":{"text":"..."}}],"reset_context":false}\n'
'{"final":"Короткий ответ пользователю","reset_context":false}\n\n' '{"final":"Короткий ответ пользователю","reset_context":false}\n\n'
"Значение final оформляй обычным Markdown, не MarkdownV2. Умеренно используй "
"жирный и курсивный текст, списки, ссылки и блоки кода, когда они улучшают "
"читаемость. Для короткого простого ответа разметка не обязательна. "
"Не добавляй декоративные эмодзи чаще одного раза на сообщение.\n\n"
"Доступные tools:\n" "Доступные tools:\n"
"- get_current_datetime {} - узнать текущее локальное и UTC-время.\n" f"{render_agent_tool_catalog()}\n\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" "Правила:\n"
"- Для просьб 'запомни', 'сохрани как факт', 'учти на будущее' используй remember.\n" "- Для просьб 'запомни', 'сохрани как факт', 'учти на будущее' используй remember.\n"
"- Для заметок используй create_note, для напоминаний create_reminder, " "- Для заметок используй create_note, для напоминаний create_reminder, "
@@ -200,185 +186,6 @@ def save_conversation_exchange(
storage.add_conversation_exchange(user_id, chat_id, prompt, answer) 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( async def run_agent_prompt(
update: Update, update: Update,
context: ContextTypes.DEFAULT_TYPE, context: ContextTypes.DEFAULT_TYPE,
@@ -389,7 +196,10 @@ async def run_agent_prompt(
user_id = require_user_id(update) user_id = require_user_id(update)
if user_id is None: if user_id is None:
await update.effective_message.reply_text("Не могу определить пользователя.") await reply_markdown(
update.effective_message,
"⚠️ **Не могу определить пользователя.**",
)
return return
storage = get_storage(context) storage = get_storage(context)
@@ -424,17 +234,12 @@ async def run_agent_prompt(
{"role": "user", "content": prompt}, {"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: try:
last_tool_results: list[dict[str, Any]] = [] last_tool_results: list[dict[str, Any]] = []
reset_context = False reset_context = False
final_answer: str | None = None
async with typing_action(context.bot, chat_id):
for _step in range(MAX_AGENT_STEPS): for _step in range(MAX_AGENT_STEPS):
raw_answer = await ai_client.chat(model, messages, json_mode=True) raw_answer = await ai_client.chat(model, messages, json_mode=True)
decision = parse_agent_decision(raw_answer) decision = parse_agent_decision(raw_answer)
@@ -463,7 +268,9 @@ async def run_agent_prompt(
messages.append( messages.append(
{ {
"role": "assistant", "role": "assistant",
"content": json_dumps({"tool_calls": decision.tool_calls}), "content": json_dumps(
{"tool_calls": decision.tool_calls}
),
} }
) )
messages.append( messages.append(
@@ -480,19 +287,12 @@ async def run_agent_prompt(
continue continue
if decision.final: if decision.final:
await edit_or_reply(placeholder, decision.final) final_answer = decision.final
save_conversation_exchange( break
storage,
user_id,
chat_id,
prompt,
decision.final,
reset_context,
)
return
messages.append({"role": "assistant", "content": raw_answer}) messages.append({"role": "assistant", "content": raw_answer})
if final_answer is None:
final_answer = await ai_client.chat( final_answer = await ai_client.chat(
model, model,
[ [
@@ -502,12 +302,13 @@ async def run_agent_prompt(
"content": ( "content": (
"Лимит tool-шагов исчерпан. Больше не вызывай tools. " "Лимит tool-шагов исчерпан. Больше не вызывай tools. "
f"Последние результаты tools: {json_dumps(last_tool_results)}. " f"Последние результаты tools: {json_dumps(last_tool_results)}. "
"Сформулируй короткий финальный ответ пользователю обычным текстом." "Сформулируй короткий финальный ответ пользователю обычным Markdown."
), ),
}, },
], ],
) )
await edit_or_reply(placeholder, final_answer)
await reply_markdown(update.effective_message, final_answer)
save_conversation_exchange( save_conversation_exchange(
storage, storage,
user_id, user_id,
@@ -523,18 +324,17 @@ async def run_agent_prompt(
hint = ( hint = (
"Проверь YANDEX_CLOUD_FOLDER, YC_API_KEY и доступ к выбранной модели." "Проверь YANDEX_CLOUD_FOLDER, YC_API_KEY и доступ к выбранной модели."
) )
await placeholder.edit_text( await reply_markdown(
f"Не получилось вызвать {ai_client.display_name}.\n{exc}\n\n{hint}" update.effective_message,
f"⚠️ **Не получилось вызвать "
f"{escape_markdown_text(ai_client.display_name)}.**\n\n"
f"{escape_markdown_text(exc)}\n\n"
f"{escape_markdown_text(hint)}",
) )
return return
except Exception: except Exception:
logger.exception("Failed to run AI prompt") logger.exception("Failed to run AI prompt")
await placeholder.edit_text("Произошла внутренняя ошибка при запросе к модели.") await reply_markdown(
update.effective_message,
"⚠️ **Произошла внутренняя ошибка** при запросе к модели.",
async def run_ai_prompt( )
update: Update,
context: ContextTypes.DEFAULT_TYPE,
prompt: str,
) -> None:
await run_agent_prompt(update, context, prompt)

View File

@@ -0,0 +1,417 @@
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Any, Callable
from zoneinfo import ZoneInfo
from .models import ReminderParseResult
from .reminders import parse_reminder
from .storage import AssistantStorage
from .telegram_utils import parse_positive_int
from .time_utils import format_local_dt, to_utc_iso, utc_now
ToolResult = dict[str, Any]
ToolExecutor = Callable[[dict[str, Any], "ToolContext"], ToolResult]
@dataclass(frozen=True)
class ToolContext:
storage: AssistantStorage
user_id: int
chat_id: int
tz: ZoneInfo
@dataclass(frozen=True)
class AgentTool:
name: str
usage: str
description: str
execute: ToolExecutor
paragraph_before: bool = False
def prompt_line(self) -> str:
suffix = f" {self.usage}" if self.usage else " {}"
prefix = "\n" if self.paragraph_before else ""
return f"{prefix}- {self.name}{suffix} - {self.description}"
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 get_current_datetime(
_arguments: dict[str, Any],
context: ToolContext,
) -> ToolResult:
now = utc_now()
return {
"ok": True,
"local": now.astimezone(context.tz).strftime("%Y-%m-%d %H:%M:%S"),
"utc": to_utc_iso(now),
"timezone": context.tz.key,
}
def remember(arguments: dict[str, Any], context: ToolContext) -> ToolResult:
text = coerce_required_text(arguments, "text")
if not text:
return {"ok": False, "error": "text is required"}
memory_id = context.storage.add_memory(context.user_id, text)
return {"ok": True, "id": memory_id, "text": text}
def list_memory(arguments: dict[str, Any], context: ToolContext) -> ToolResult:
rows = context.storage.list_memories(
context.user_id,
coerce_positive_int(arguments.get("limit"), 20),
)
return {"ok": True, "items": rows}
def delete_memory(arguments: dict[str, Any], context: ToolContext) -> ToolResult:
memory_id = parse_positive_int(str(arguments.get("id", "")))
if memory_id is None:
return {"ok": False, "error": "valid id is required"}
return {
"ok": context.storage.delete_memory(context.user_id, memory_id),
"id": memory_id,
}
def create_note(arguments: dict[str, Any], context: ToolContext) -> ToolResult:
text = coerce_required_text(arguments, "text")
if not text:
return {"ok": False, "error": "text is required"}
note_id = context.storage.add_note(context.user_id, text)
return {"ok": True, "id": note_id, "text": text}
def list_notes(arguments: dict[str, Any], context: ToolContext) -> ToolResult:
rows = context.storage.list_notes(
context.user_id,
coerce_positive_int(arguments.get("limit"), 20),
)
return {"ok": True, "items": rows}
def delete_note(arguments: dict[str, Any], context: ToolContext) -> ToolResult:
note_id = parse_positive_int(str(arguments.get("id", "")))
if note_id is None:
return {"ok": False, "error": "valid id is required"}
return {
"ok": context.storage.delete_note(context.user_id, note_id),
"id": note_id,
}
def create_reminder(
arguments: dict[str, Any],
context: ToolContext,
) -> ToolResult:
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, context.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 = context.storage.add_reminder(
context.user_id,
context.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, context.tz),
}
def list_reminders(
arguments: dict[str, Any],
context: ToolContext,
) -> ToolResult:
rows = context.storage.list_reminders(
context.user_id,
coerce_positive_int(arguments.get("limit"), 20),
)
for row in rows:
row["remind_at_local"] = format_local_dt(
row["remind_at"],
context.tz,
)
return {"ok": True, "items": rows}
def cancel_reminder(
arguments: dict[str, Any],
context: ToolContext,
) -> ToolResult:
reminder_id = parse_positive_int(str(arguments.get("id", "")))
if reminder_id is None:
return {"ok": False, "error": "valid id is required"}
return {
"ok": context.storage.cancel_reminder(
context.user_id,
reminder_id,
),
"id": reminder_id,
}
def create_status(arguments: dict[str, Any], context: ToolContext) -> ToolResult:
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 = context.storage.add_tracked_item(
context.user_id,
title,
status,
)
return {
"ok": True,
"id": item_id,
"title": title,
"status": status,
}
def list_statuses(
arguments: dict[str, Any],
context: ToolContext,
) -> ToolResult:
rows = context.storage.list_tracked_items(
context.user_id,
coerce_positive_int(arguments.get("limit"), 30),
)
return {"ok": True, "items": rows}
def update_status(arguments: dict[str, Any], context: ToolContext) -> ToolResult:
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": context.storage.set_tracked_status(
context.user_id,
item_id,
status,
),
"id": item_id,
"status": status,
}
def delete_status(arguments: dict[str, Any], context: ToolContext) -> ToolResult:
item_id = parse_positive_int(str(arguments.get("id", "")))
if item_id is None:
return {"ok": False, "error": "valid id is required"}
return {
"ok": context.storage.delete_tracked_item(
context.user_id,
item_id,
),
"id": item_id,
}
def search_conversation(
arguments: dict[str, Any],
context: ToolContext,
) -> ToolResult:
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 = context.storage.search_conversation_messages(
context.user_id,
context.chat_id,
query,
limit,
)
discussions = [
{
"matched_message_id": hit["id"],
"score": hit["score"],
"messages": context.storage.conversation_message_window(
context.user_id,
context.chat_id,
int(hit["id"]),
),
}
for hit in hits
]
return {
"ok": True,
"query": query,
"discussions": discussions,
}
AGENT_TOOLS = (
AgentTool(
"get_current_datetime",
"",
"узнать текущее локальное и UTC-время.",
get_current_datetime,
),
AgentTool(
"remember",
'{"text": string}',
"сохранить важный долгосрочный факт о пользователе.",
remember,
),
AgentTool(
"list_memory",
'{"limit": number}',
"показать сохраненную память.",
list_memory,
),
AgentTool(
"delete_memory",
'{"id": number}',
"удалить запись памяти.",
delete_memory,
),
AgentTool(
"create_note",
'{"text": string}',
"сохранить заметку.",
create_note,
),
AgentTool(
"list_notes",
'{"limit": number}',
"показать заметки.",
list_notes,
),
AgentTool(
"delete_note",
'{"id": number}',
"удалить заметку.",
delete_note,
),
AgentTool(
"create_reminder",
'{"when": string, "text": string}',
(
"поставить напоминание. when можно указывать как '30m', "
"'через 2 часа', '18:30', '2026-07-16 18:30'. Если пользователь "
"говорит 'завтра/послезавтра/через неделю', сам рассчитай дату от "
"текущего локального времени и передай 'YYYY-MM-DD HH:MM'."
),
create_reminder,
),
AgentTool(
"list_reminders",
'{"limit": number}',
"показать активные напоминания.",
list_reminders,
),
AgentTool(
"cancel_reminder",
'{"id": number}',
"отменить напоминание.",
cancel_reminder,
),
AgentTool(
"create_status",
'{"title": string, "status": string}',
"начать отслеживать статус.",
create_status,
),
AgentTool(
"list_statuses",
'{"limit": number}',
"показать отслеживаемые статусы.",
list_statuses,
),
AgentTool(
"update_status",
'{"id": number, "status": string}',
"обновить статус.",
update_status,
),
AgentTool(
"delete_status",
'{"id": number}',
"удалить отслеживаемый объект.",
delete_status,
),
AgentTool(
"search_conversation",
'{"query": string, "limit": number}',
(
"найти старое обсуждение во всей сохраненной переписке по "
"содержательным ключевым словам."
),
search_conversation,
paragraph_before=True,
),
)
AGENT_TOOL_BY_NAME = {tool.name: tool for tool in AGENT_TOOLS}
def render_agent_tool_catalog() -> str:
return "\n".join(tool.prompt_line() for tool in AGENT_TOOLS)
def execute_agent_tool(
name: str,
arguments: dict[str, Any],
storage: AssistantStorage,
user_id: int,
chat_id: int,
tz: ZoneInfo,
) -> ToolResult:
tool = AGENT_TOOL_BY_NAME.get(name.strip().lower())
if tool is None:
return {"ok": False, "error": f"unknown tool: {name}"}
return tool.execute(
arguments,
ToolContext(
storage=storage,
user_id=user_id,
chat_id=chat_id,
tz=tz,
),
)

View File

@@ -6,25 +6,13 @@ from telegram.ext import (
CommandHandler, CommandHandler,
InlineQueryHandler, InlineQueryHandler,
MessageHandler, MessageHandler,
TypeHandler,
filters, filters,
) )
from .ai import AIClient from .ai import AIClient
from .config import ( from .auth import authentication_guard
get_assistant_mode, from .config import AppSettings, load_app_settings
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 ( from .handlers import (
ask_command, ask_command,
help_command, help_command,
@@ -39,6 +27,7 @@ from .handlers import (
) )
from .jobs import post_init, post_shutdown from .jobs import post_init, post_shutdown
from .ollama import OllamaClient from .ollama import OllamaClient
from .services import SERVICES_KEY, ApplicationServices
from .storage import AssistantStorage from .storage import AssistantStorage
from .speech import SpeechRecognizer, SpeechTranscriber from .speech import SpeechRecognizer, SpeechTranscriber
from .yandex_ai import YandexAIClient, YandexSpeechRecognizer from .yandex_ai import YandexAIClient, YandexSpeechRecognizer
@@ -52,42 +41,56 @@ def configure_logging() -> None:
logging.getLogger("httpx").setLevel(logging.WARNING) logging.getLogger("httpx").setLevel(logging.WARNING)
def create_application() -> Application: def create_services(settings: AppSettings) -> ApplicationServices:
storage = AssistantStorage(get_db_path()) storage = AssistantStorage(
mode = get_assistant_mode() settings.db_path,
default_models=settings.default_models(),
)
ai_client: AIClient
speech_recognizer: SpeechTranscriber
if settings.mode == "local":
ai_client = OllamaClient(settings.ollama_base_url)
speech_recognizer = SpeechRecognizer(
model_name=settings.whisper_model,
device=settings.whisper_device,
compute_type=settings.whisper_compute_type,
language=settings.whisper_language,
)
else:
folder_id = settings.yandex_cloud_folder
if folder_id is None:
raise RuntimeError(
"YANDEX_CLOUD_FOLDER is required in yandex mode."
)
ai_client = YandexAIClient(folder_id=folder_id)
speech_recognizer = YandexSpeechRecognizer(
folder_id=folder_id,
language=settings.yandex_stt_language,
model=settings.yandex_stt_model,
)
return ApplicationServices(
storage=storage,
ai_client=ai_client,
timezone=settings.timezone,
speech_recognizer=speech_recognizer,
voice_max_duration_seconds=settings.voice_max_duration_seconds,
assistant_password=settings.assistant_password,
)
def create_application(settings: AppSettings | None = None) -> Application:
settings = settings or load_app_settings()
application = ( application = (
Application.builder() Application.builder()
.token(get_bot_token()) .token(settings.bot_token)
.post_init(post_init) .post_init(post_init)
.post_shutdown(post_shutdown) .post_shutdown(post_shutdown)
.build() .build()
) )
application.bot_data["storage"] = storage application.bot_data[SERVICES_KEY] = create_services(settings)
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(TypeHandler(Update, authentication_guard), group=-1)
application.add_handler(CommandHandler("start", start)) application.add_handler(CommandHandler("start", start))
application.add_handler(CommandHandler("help", help_command)) application.add_handler(CommandHandler("help", help_command))
application.add_handler(CommandHandler("ask", ask_command)) application.add_handler(CommandHandler("ask", ask_command))

68
assistant_bot/auth.py Normal file
View File

@@ -0,0 +1,68 @@
from hmac import compare_digest
from telegram import Update
from telegram.constants import ChatType
from telegram.ext import ApplicationHandlerStop, ContextTypes
from .telegram_utils import (
get_services,
get_storage,
reply_markdown,
require_user_id,
)
def passwords_match(candidate: str, expected: str) -> bool:
return compare_digest(
candidate.encode("utf-8"),
expected.encode("utf-8"),
)
async def authentication_guard(
update: Update,
context: ContextTypes.DEFAULT_TYPE,
) -> None:
user_id = require_user_id(update)
if user_id is None:
raise ApplicationHandlerStop
storage = get_storage(context)
if storage.is_user_authorized(user_id):
return
message = update.effective_message
chat = update.effective_chat
if chat is not None and chat.type != ChatType.PRIVATE:
if message:
await reply_markdown(
message,
"🔐 **Сначала авторизуйся в личном чате с ботом.**",
)
raise ApplicationHandlerStop
candidate = message.text.strip() if message and message.text else ""
password = get_services(context).assistant_password
if candidate and passwords_match(candidate, password):
storage.authorize_user(user_id)
await reply_markdown(
message,
"✅ **Пароль принят.** Доступ открыт — повторно вводить его не нужно.",
)
raise ApplicationHandlerStop
if candidate and not candidate.startswith("/"):
await reply_markdown(
message,
"🔐 **Неверный пароль.** Попробуй ещё раз.",
)
raise ApplicationHandlerStop
if message:
await reply_markdown(
message,
"🔐 **Для доступа к боту введи пароль** одним текстовым сообщением.",
)
elif update.inline_query:
await update.inline_query.answer([], cache_time=0)
raise ApplicationHandlerStop

View File

@@ -1,5 +1,6 @@
import logging import logging
import os import os
from dataclasses import dataclass
from pathlib import Path from pathlib import Path
from zoneinfo import ZoneInfo, ZoneInfoNotFoundError from zoneinfo import ZoneInfo, ZoneInfoNotFoundError
@@ -10,6 +11,7 @@ PROJECT_ROOT = Path(__file__).resolve().parent.parent
ENV_FILE = PROJECT_ROOT / ".env" ENV_FILE = PROJECT_ROOT / ".env"
TOKEN_ENV_NAME = "BOT_TOKEN" TOKEN_ENV_NAME = "BOT_TOKEN"
PASSWORD_ENV_NAME = "ASSISTANT_PASSWORD"
DB_ENV_NAME = "ASSISTANT_DB" DB_ENV_NAME = "ASSISTANT_DB"
MODE_ENV_NAME = "ASSISTANT_MODE" MODE_ENV_NAME = "ASSISTANT_MODE"
OLLAMA_BASE_URL_ENV_NAME = "OLLAMA_BASE_URL" OLLAMA_BASE_URL_ENV_NAME = "OLLAMA_BASE_URL"
@@ -44,10 +46,34 @@ MAX_TELEGRAM_MESSAGE_LENGTH = 3900
MAX_AGENT_STEPS = 5 MAX_AGENT_STEPS = 5
CONVERSATION_HISTORY_LIMIT = 12 CONVERSATION_HISTORY_LIMIT = 12
HTML_FORMATS = (
("Bold", "b"), @dataclass(frozen=True)
("Italic", "i"), class AppSettings:
bot_token: str
assistant_password: str
db_path: Path
mode: str
timezone: ZoneInfo
voice_max_duration_seconds: int
ollama_base_url: str
ollama_model: str
yandex_cloud_folder: str | None
yandex_cloud_model: str
yandex_stt_model: str
yandex_stt_language: str
whisper_model: str
whisper_device: str
whisper_compute_type: str
whisper_language: str | None
def default_models(self) -> dict[str, str]:
models = {"local": self.ollama_model}
if self.yandex_cloud_folder:
models["yandex"] = (
f"gpt://{self.yandex_cloud_folder}/"
f"{self.yandex_cloud_model}"
) )
return models
def load_env_file(path: Path = ENV_FILE) -> None: def load_env_file(path: Path = ENV_FILE) -> None:
@@ -72,142 +98,123 @@ def load_env_file(path: Path = ENV_FILE) -> None:
os.environ[key] = value os.environ[key] = value
def get_bot_token() -> str: def load_app_settings() -> AppSettings:
"""Load and validate the complete startup configuration once."""
load_env_file() load_env_file()
token = os.getenv(TOKEN_ENV_NAME) token = os.getenv(TOKEN_ENV_NAME)
if not token: if not token:
raise RuntimeError(f"Set {TOKEN_ENV_NAME} in environment or .env file.") raise RuntimeError(
return token f"Set {TOKEN_ENV_NAME} in environment or .env file."
)
password = os.getenv(PASSWORD_ENV_NAME, "").strip()
if not password:
raise RuntimeError(
f"Set {PASSWORD_ENV_NAME} in environment or .env file."
)
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() mode = os.getenv(MODE_ENV_NAME, DEFAULT_MODE).strip().lower()
if mode not in {"local", "yandex"}: if mode not in {"local", "yandex"}:
raise RuntimeError( raise RuntimeError(
f"{MODE_ENV_NAME} must be either 'local' or 'yandex', got {mode!r}." f"{MODE_ENV_NAME} must be either 'local' or 'yandex', got {mode!r}."
) )
return mode
raw_db_path = os.getenv(DB_ENV_NAME)
if raw_db_path:
db_path = Path(raw_db_path).expanduser()
if not db_path.is_absolute():
db_path = PROJECT_ROOT / db_path
else:
db_path = DEFAULT_DB_FILE
def get_ollama_base_url() -> str:
load_env_file()
return os.getenv(OLLAMA_BASE_URL_ENV_NAME, DEFAULT_OLLAMA_BASE_URL).rstrip("/")
def get_default_ollama_model() -> str:
load_env_file()
return os.getenv(OLLAMA_MODEL_ENV_NAME, DEFAULT_OLLAMA_MODEL)
def get_yandex_cloud_folder() -> str:
load_env_file()
folder = os.getenv(YANDEX_CLOUD_FOLDER_ENV_NAME, "").strip()
if not folder or folder == "":
raise RuntimeError(
f"Set {YANDEX_CLOUD_FOLDER_ENV_NAME} in environment or .env file."
)
return folder.strip("/")
def get_yandex_cloud_model() -> str:
load_env_file()
model = os.getenv(
YANDEX_CLOUD_MODEL_ENV_NAME,
DEFAULT_YANDEX_CLOUD_MODEL,
).strip()
return model.strip("/") or DEFAULT_YANDEX_CLOUD_MODEL
def get_default_yandex_model() -> str:
return f"gpt://{get_yandex_cloud_folder()}/{get_yandex_cloud_model()}"
def get_default_model(provider: str) -> str:
if provider == "local":
return get_default_ollama_model()
if provider == "yandex":
return get_default_yandex_model()
raise ValueError(f"Unknown AI provider: {provider}")
def get_yandex_stt_model() -> str:
load_env_file()
return os.getenv(YANDEX_STT_MODEL_ENV_NAME, DEFAULT_YANDEX_STT_MODEL).strip()
def get_yandex_stt_language() -> str:
load_env_file()
return os.getenv(
YANDEX_STT_LANGUAGE_ENV_NAME,
DEFAULT_YANDEX_STT_LANGUAGE,
).strip()
def get_local_timezone() -> ZoneInfo:
load_env_file()
timezone_name = os.getenv(TIMEZONE_ENV_NAME, DEFAULT_TIMEZONE) timezone_name = os.getenv(TIMEZONE_ENV_NAME, DEFAULT_TIMEZONE)
try: try:
return ZoneInfo(timezone_name) local_timezone = ZoneInfo(timezone_name)
except ZoneInfoNotFoundError: except ZoneInfoNotFoundError:
logger.warning("Unknown timezone %s, falling back to UTC", timezone_name) logger.warning(
return ZoneInfo("UTC") "Unknown timezone %s, falling back to UTC",
timezone_name,
)
local_timezone = ZoneInfo("UTC")
raw_voice_duration = os.getenv(
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, VOICE_MAX_DURATION_ENV_NAME,
str(DEFAULT_VOICE_MAX_DURATION_SECONDS), str(DEFAULT_VOICE_MAX_DURATION_SECONDS),
) )
try: try:
value = int(raw_value) voice_duration = int(raw_voice_duration)
except ValueError: except ValueError:
value = 0 voice_duration = 0
if value <= 0: if voice_duration <= 0:
logger.warning( logger.warning(
"%s must be a positive integer, using %s", "%s must be a positive integer, using %s",
VOICE_MAX_DURATION_ENV_NAME, VOICE_MAX_DURATION_ENV_NAME,
DEFAULT_VOICE_MAX_DURATION_SECONDS, DEFAULT_VOICE_MAX_DURATION_SECONDS,
) )
return DEFAULT_VOICE_MAX_DURATION_SECONDS voice_duration = DEFAULT_VOICE_MAX_DURATION_SECONDS
return value
yandex_folder = os.getenv(
YANDEX_CLOUD_FOLDER_ENV_NAME,
"",
).strip().strip("/")
if mode == "yandex" and not yandex_folder:
raise RuntimeError(
f"Set {YANDEX_CLOUD_FOLDER_ENV_NAME} in environment or .env file."
)
yandex_model = os.getenv(
YANDEX_CLOUD_MODEL_ENV_NAME,
DEFAULT_YANDEX_CLOUD_MODEL,
).strip().strip("/")
if not yandex_model:
yandex_model = DEFAULT_YANDEX_CLOUD_MODEL
whisper_language = os.getenv(
WHISPER_LANGUAGE_ENV_NAME,
DEFAULT_WHISPER_LANGUAGE,
).strip()
return AppSettings(
bot_token=token,
assistant_password=password,
db_path=db_path,
mode=mode,
timezone=local_timezone,
voice_max_duration_seconds=voice_duration,
ollama_base_url=os.getenv(
OLLAMA_BASE_URL_ENV_NAME,
DEFAULT_OLLAMA_BASE_URL,
).rstrip("/"),
ollama_model=os.getenv(
OLLAMA_MODEL_ENV_NAME,
DEFAULT_OLLAMA_MODEL,
),
yandex_cloud_folder=yandex_folder or None,
yandex_cloud_model=yandex_model,
yandex_stt_model=os.getenv(
YANDEX_STT_MODEL_ENV_NAME,
DEFAULT_YANDEX_STT_MODEL,
).strip(),
yandex_stt_language=os.getenv(
YANDEX_STT_LANGUAGE_ENV_NAME,
DEFAULT_YANDEX_STT_LANGUAGE,
).strip(),
whisper_model=os.getenv(
WHISPER_MODEL_ENV_NAME,
DEFAULT_WHISPER_MODEL,
).strip(),
whisper_device=os.getenv(
WHISPER_DEVICE_ENV_NAME,
DEFAULT_WHISPER_DEVICE,
).strip(),
whisper_compute_type=os.getenv(
WHISPER_COMPUTE_TYPE_ENV_NAME,
DEFAULT_WHISPER_COMPUTE_TYPE,
).strip(),
whisper_language=(
None
if whisper_language.lower() == "auto"
else whisper_language or None
),
)

View File

@@ -3,7 +3,6 @@ import logging
import os import os
import tempfile import tempfile
from pathlib import Path from pathlib import Path
from typing import Any
from telegram import Update from telegram import Update
from telegram.constants import ChatAction from telegram.constants import ChatAction
@@ -12,17 +11,18 @@ from telegram.ext import ContextTypes
from .ai import AIClientError from .ai import AIClientError
from .agent import run_agent_prompt from .agent import run_agent_prompt
from .reminders import parse_reminder
from .telegram_utils import ( from .telegram_utils import (
build_inline_results, build_inline_results,
command_text, command_text,
edit_markdown,
escape_markdown_text,
get_ai_client, get_ai_client,
get_speech_recognizer, get_speech_recognizer,
get_storage, get_storage,
get_tz, get_tz,
get_voice_max_duration, get_voice_max_duration,
parse_positive_int, markdown_code,
reply_long, reply_markdown,
require_user_id, require_user_id,
) )
from .speech import SpeechRecognitionError from .speech import SpeechRecognitionError
@@ -32,42 +32,37 @@ from .time_utils import format_local_dt
logger = logging.getLogger(__name__) 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: async def start(update: Update, _context: ContextTypes.DEFAULT_TYPE) -> None:
if update.message: if update.message:
await update.message.reply_text( await reply_markdown(
"Я готов как персональный ассистент.\n" update.message,
"Пиши обычным текстом или отправляй голосовые сообщения: " "**👋 Персональный ассистент готов**\n\n"
"'запомни, что...', 'напомни завтра в 10:00...', " "Пиши обычным текстом или отправляй голосовые сообщения. Например:\n\n"
"'сохрани заметку...', 'покажи мои статусы'.\n" "- Запомни, что я предпочитаю короткие ответы\n"
"Я помню последние реплики диалога; команда /new очищает текущий контекст.\n" "- Напомни завтра в 10:00 проверить почту\n"
"Я сам решу, когда нужно вызвать внутренний tool." "- Сохрани заметку с идеей проекта\n"
"- Покажи мои статусы\n\n"
"Я помню последние реплики диалога. Команда `/new` начинает новую тему, "
"не удаляя архив переписки.",
) )
async def help_command(update: Update, _context: ContextTypes.DEFAULT_TYPE) -> None: async def help_command(update: Update, _context: ContextTypes.DEFAULT_TYPE) -> None:
if update.message: if update.message:
await update.message.reply_text( await reply_markdown(
"Основной режим - свободный диалог текстом или голосовыми сообщениями.\n\n" update.message,
"Примеры:\n" "**🧭 Как пользоваться ботом**\n\n"
"Запомни, что я предпочитаю короткие ответы.\n" "Основной режим — свободный диалог текстом или голосовыми сообщениями.\n\n"
"Сохрани заметку: идея для проекта.\n" "**Примеры запросов**\n\n"
"Напомни через 30 минут проверить сборку.\n" "- Запомни, что я предпочитаю короткие ответы\n"
"Отслеживай паспорт, статус: жду ответа.\n" "- Сохрани заметку: идея для проекта\n"
"Покажи активные напоминания.\n\n" "- Напомни через 30 минут проверить сборку\n"
"Вся переписка сохраняется. /history тема ищет старое обсуждение, " "- Отслеживай паспорт, статус: жду ответа\n"
"/new начинает новый контекст без удаления архива.\n\n" "- Покажи активные напоминания\n\n"
"Служебные команды оставлены для настройки и отладки: /models, /model, /ask.\n" "Вся переписка сохраняется. `/history тема` ищет старое обсуждение, "
"Режим выбирается через ASSISTANT_MODE=local или ASSISTANT_MODE=yandex." "а `/new` начинает новый контекст без удаления архива.\n\n"
"**Служебные команды:** `/models`, `/model`, `/ask`\n"
"**Режим:** `ASSISTANT_MODE=local` или `ASSISTANT_MODE=yandex`",
) )
@@ -76,10 +71,13 @@ async def new_conversation_command(update: Update, context: ContextTypes.DEFAULT
return return
user_id = require_user_id(update) user_id = require_user_id(update)
if user_id is None: if user_id is None:
await update.message.reply_text("Не могу определить пользователя.") await reply_markdown(update.message, "⚠️ **Не могу определить пользователя.**")
return return
get_storage(context).start_new_conversation(user_id, int(update.message.chat_id)) get_storage(context).start_new_conversation(user_id, int(update.message.chat_id))
await update.message.reply_text("Начинаем новую тему. Предыдущая переписка сохранена в архиве.") await reply_markdown(
update.message,
"🆕 **Начинаем новую тему.** Предыдущая переписка сохранена в архиве.",
)
async def history_command(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None: async def history_command(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
@@ -87,12 +85,15 @@ async def history_command(update: Update, context: ContextTypes.DEFAULT_TYPE) ->
return return
user_id = require_user_id(update) user_id = require_user_id(update)
if user_id is None: if user_id is None:
await update.message.reply_text("Не могу определить пользователя.") await reply_markdown(update.message, "⚠️ **Не могу определить пользователя.**")
return return
query = command_text(context) query = command_text(context)
if not query: if not query:
await update.message.reply_text("Укажи тему или ключевые слова: /history отпуск") await reply_markdown(
update.message,
"🔎 Укажи тему или ключевые слова: `/history отпуск`",
)
return return
rows = get_storage(context).search_conversation_messages( rows = get_storage(context).search_conversation_messages(
@@ -102,24 +103,31 @@ async def history_command(update: Update, context: ContextTypes.DEFAULT_TYPE) ->
limit=10, limit=10,
) )
if not rows: if not rows:
await update.message.reply_text("В сохраненной переписке ничего не найдено.") await reply_markdown(
update.message,
"🔎 **В сохранённой переписке ничего не найдено.**",
)
return return
role_names = {"user": "Вы", "assistant": "Бот"} role_names = {"user": "Вы", "assistant": "Бот"}
tz = get_tz(context) tz = get_tz(context)
text = "Найденные сообщения:\n\n" + "\n\n".join( text = "**🔎 Найденные сообщения**\n\n" + "\n\n".join(
f"{role_names.get(str(row['role']), row['role'])} · " f"**{escape_markdown_text(role_names.get(str(row['role']), row['role']))}** · "
f"{format_local_dt(row['created_at'], tz)}\n{row['content']}" f"{markdown_code(format_local_dt(row['created_at'], tz))}\n"
f"{str(row['content']) if row['role'] == 'assistant' else escape_markdown_text(row['content'])}"
for row in rows for row in rows
) )
await reply_long(update.message, text) await reply_markdown(update.message, text)
async def ask_command(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None: async def ask_command(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None:
prompt = command_text(context) prompt = command_text(context)
if not prompt: if not prompt:
if update.message: if update.message:
await update.message.reply_text("Напиши текст после /ask или просто отправь сообщение в личный чат.") await reply_markdown(
update.message,
"💬 Напиши текст после `/ask` или просто отправь сообщение в личный чат.",
)
return return
await run_agent_prompt(update, context, prompt) await run_agent_prompt(update, context, prompt)
@@ -137,8 +145,10 @@ async def private_voice(update: Update, context: ContextTypes.DEFAULT_TYPE) -> N
voice = update.message.voice voice = update.message.voice
max_duration = get_voice_max_duration(context) max_duration = get_voice_max_duration(context)
if voice.duration > max_duration: if voice.duration > max_duration:
await update.message.reply_text( await reply_markdown(
f"Голосовое сообщение слишком длинное. Максимум: {max_duration} сек." update.message,
"🎙️ **Голосовое сообщение слишком длинное.** "
f"Максимум: {markdown_code(max_duration)} сек.",
) )
return return
@@ -146,7 +156,10 @@ async def private_voice(update: Update, context: ContextTypes.DEFAULT_TYPE) -> N
chat_id=int(update.message.chat_id), chat_id=int(update.message.chat_id),
action=ChatAction.TYPING, action=ChatAction.TYPING,
) )
status_message = await update.message.reply_text("Распознаю голосовое сообщение...") status_message = await reply_markdown(
update.message,
"🎙️ *Распознаю голосовое сообщение…*",
)
temp_path: Path | None = None temp_path: Path | None = None
transcript = "" transcript = ""
@@ -163,14 +176,17 @@ async def private_voice(update: Update, context: ContextTypes.DEFAULT_TYPE) -> N
) )
except (SpeechRecognitionError, TelegramError): except (SpeechRecognitionError, TelegramError):
logger.exception("Failed to transcribe Telegram voice message") logger.exception("Failed to transcribe Telegram voice message")
await status_message.edit_text( await edit_markdown(
"Не получилось распознать голосовое сообщение. Попробуй еще раз позже." status_message,
"⚠️ **Не получилось распознать голосовое сообщение.** "
"Попробуй ещё раз позже.",
) )
return return
except Exception: except Exception:
logger.exception("Unexpected error while processing Telegram voice message") logger.exception("Unexpected error while processing Telegram voice message")
await status_message.edit_text( await edit_markdown(
"Произошла внутренняя ошибка при обработке голосового сообщения." status_message,
"⚠️ **Произошла внутренняя ошибка** при обработке голосового сообщения.",
) )
return return
finally: finally:
@@ -181,11 +197,17 @@ async def private_voice(update: Update, context: ContextTypes.DEFAULT_TYPE) -> N
logger.warning("Could not delete temporary voice file %s", temp_path) logger.warning("Could not delete temporary voice file %s", temp_path)
if not transcript: if not transcript:
await status_message.edit_text("Не удалось расслышать речь в сообщении.") await edit_markdown(
status_message,
"⚠️ **Не удалось расслышать речь в сообщении.**",
)
return return
preview = transcript if len(transcript) <= 500 else f"{transcript[:497]}..." preview = transcript if len(transcript) <= 500 else f"{transcript[:497]}..."
await status_message.edit_text(f"Распознано: {preview}") await edit_markdown(
status_message,
f"🎙️ **Распознано:** {escape_markdown_text(preview)}",
)
await run_agent_prompt(update, context, transcript) await run_agent_prompt(update, context, transcript)
@@ -197,20 +219,26 @@ async def models_command(update: Update, context: ContextTypes.DEFAULT_TYPE) ->
try: try:
models = await ai_client.list_models() models = await ai_client.list_models()
except AIClientError as exc: except AIClientError as exc:
await update.message.reply_text( await reply_markdown(
f"Не получилось получить список моделей {ai_client.display_name}.\n{exc}" update.message,
"⚠️ **Не получилось получить список моделей "
f"{escape_markdown_text(ai_client.display_name)}.**\n\n"
f"{escape_markdown_text(exc)}",
) )
return return
if not models: if not models:
await update.message.reply_text( await reply_markdown(
f"{ai_client.display_name} доступна, но моделей не найдено." update.message,
f"🤖 **{escape_markdown_text(ai_client.display_name)} доступна, "
"но моделей не найдено.**",
) )
return return
await update.message.reply_text( await reply_markdown(
f"Модели {ai_client.display_name}:\n" update.message,
+ "\n".join(f"- {model}" for model in models) f"**🤖 Модели {escape_markdown_text(ai_client.display_name)}**\n\n"
+ "\n".join(f"- {markdown_code(model)}" for model in models),
) )
@@ -220,276 +248,29 @@ async def model_command(update: Update, context: ContextTypes.DEFAULT_TYPE) -> N
user_id = require_user_id(update) user_id = require_user_id(update)
if user_id is None: if user_id is None:
await update.message.reply_text("Не могу определить пользователя.") await reply_markdown(update.message, "⚠️ **Не могу определить пользователя.**")
return return
storage = get_storage(context) storage = get_storage(context)
ai_client = get_ai_client(context) ai_client = get_ai_client(context)
model = command_text(context) model = command_text(context)
if not model: if not model:
await update.message.reply_text( await reply_markdown(
f"Текущая модель {ai_client.display_name}: " update.message,
f"{ai_client.normalize_model(storage.get_user_model(user_id, ai_client.provider))}" f"🤖 **Текущая модель {escape_markdown_text(ai_client.display_name)}:** "
f"{markdown_code(ai_client.normalize_model(storage.get_user_model(user_id, ai_client.provider)))}",
) )
return return
model = ai_client.normalize_model(model) model = ai_client.normalize_model(model)
storage.set_user_model(user_id, model, ai_client.provider) storage.set_user_model(user_id, model, ai_client.provider)
await update.message.reply_text( await reply_markdown(
f"Модель {ai_client.display_name} сохранена: {model}" update.message,
f"✅ **Модель {escape_markdown_text(ai_client.display_name)} сохранена:** "
f"{markdown_code(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.message.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: async def inline_query(update: Update, _context: ContextTypes.DEFAULT_TYPE) -> None:
if not update.inline_query or not update.inline_query.query: if not update.inline_query or not update.inline_query.query:
return return

View File

@@ -6,7 +6,8 @@ from zoneinfo import ZoneInfo
from telegram.ext import Application from telegram.ext import Application
from .config import REMINDER_POLL_SECONDS from .config import REMINDER_POLL_SECONDS
from .storage import AssistantStorage from .services import SERVICES_KEY, ApplicationServices
from .telegram_utils import escape_markdown_text, markdown_code, send_markdown
from .time_utils import format_local_dt, utc_now from .time_utils import format_local_dt, utc_now
@@ -14,19 +15,23 @@ logger = logging.getLogger(__name__)
async def reminder_loop(application: Application) -> None: async def reminder_loop(application: Application) -> None:
storage: AssistantStorage = application.bot_data["storage"] services: ApplicationServices = application.bot_data[SERVICES_KEY]
tz: ZoneInfo = application.bot_data["timezone"] storage = services.storage
tz: ZoneInfo = services.timezone
while True: while True:
try: try:
due_reminders = storage.due_reminders(utc_now()) due_reminders = storage.due_reminders(utc_now())
for reminder in due_reminders: for reminder in due_reminders:
await application.bot.send_message( await send_markdown(
application.bot,
chat_id=reminder["chat_id"], chat_id=reminder["chat_id"],
text=( markdown=(
f"Напоминание #{reminder['id']} " f"**⏰ Напоминание "
f"({format_local_dt(reminder['remind_at'], tz)}):\n" f"{markdown_code('#' + str(reminder['id']))}**\n\n"
f"{reminder['text']}" f"**Время:** "
f"{markdown_code(format_local_dt(reminder['remind_at'], tz))}\n"
f"**Задача:** {escape_markdown_text(reminder['text'])}"
), ),
) )
storage.mark_reminder_sent(int(reminder["id"])) storage.mark_reminder_sent(int(reminder["id"]))

113
assistant_bot/migrations.py Normal file
View File

@@ -0,0 +1,113 @@
import sqlite3
LATEST_SCHEMA_VERSION = 2
INITIAL_SCHEMA = """
CREATE TABLE IF NOT EXISTS user_settings (
user_id INTEGER PRIMARY KEY,
ollama_model TEXT,
updated_at TEXT NOT NULL
);
CREATE TABLE IF NOT EXISTS authorized_users (
user_id INTEGER PRIMARY KEY,
authorized_at TEXT NOT NULL
);
CREATE TABLE IF NOT EXISTS memories (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL,
text TEXT NOT NULL,
created_at TEXT NOT NULL
);
CREATE TABLE IF NOT EXISTS notes (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL,
text TEXT NOT NULL,
created_at TEXT NOT NULL
);
CREATE TABLE IF NOT EXISTS reminders (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL,
chat_id INTEGER NOT NULL,
text TEXT NOT NULL,
remind_at TEXT NOT NULL,
status TEXT NOT NULL DEFAULT 'pending',
created_at TEXT NOT NULL,
sent_at TEXT
);
CREATE TABLE IF NOT EXISTS tracked_items (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL,
title TEXT NOT NULL,
status TEXT NOT NULL,
created_at TEXT NOT NULL,
updated_at TEXT NOT NULL
);
CREATE TABLE IF NOT EXISTS conversation_messages (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL,
chat_id INTEGER NOT NULL,
role TEXT NOT NULL CHECK(role IN ('user', 'assistant')),
content TEXT NOT NULL,
created_at TEXT NOT NULL
);
CREATE TABLE IF NOT EXISTS conversation_contexts (
user_id INTEGER NOT NULL,
chat_id INTEGER NOT NULL,
started_after_id INTEGER NOT NULL DEFAULT 0,
updated_at TEXT NOT NULL,
PRIMARY KEY (user_id, chat_id)
);
CREATE INDEX IF NOT EXISTS idx_memories_user_created
ON memories(user_id, created_at DESC);
CREATE INDEX IF NOT EXISTS idx_notes_user_created
ON notes(user_id, created_at DESC);
CREATE INDEX IF NOT EXISTS idx_reminders_due
ON reminders(status, remind_at);
CREATE INDEX IF NOT EXISTS idx_tracked_user_updated
ON tracked_items(user_id, updated_at DESC);
CREATE INDEX IF NOT EXISTS idx_conversation_chat_id
ON conversation_messages(user_id, chat_id, id DESC);
"""
def _schema_version(connection: sqlite3.Connection) -> int:
row = connection.execute("PRAGMA user_version").fetchone()
return int(row[0])
def _has_column(
connection: sqlite3.Connection,
table: str,
column: str,
) -> bool:
return any(
str(row["name"]) == column
for row in connection.execute(f"PRAGMA table_info({table})")
)
def migrate_database(connection: sqlite3.Connection) -> None:
"""Bring a new or existing database to the latest known schema."""
connection.execute("PRAGMA journal_mode=WAL")
version = _schema_version(connection)
if version < 1:
connection.executescript(INITIAL_SCHEMA)
connection.execute("PRAGMA user_version = 1")
version = 1
if version < 2:
if not _has_column(connection, "user_settings", "yandex_model"):
connection.execute(
"ALTER TABLE user_settings ADD COLUMN yandex_model TEXT"
)
connection.execute("PRAGMA user_version = 2")

19
assistant_bot/services.py Normal file
View File

@@ -0,0 +1,19 @@
from dataclasses import dataclass
from zoneinfo import ZoneInfo
from .ai import AIClient
from .speech import SpeechTranscriber
from .storage import AssistantStorage
SERVICES_KEY = "services"
@dataclass(frozen=True)
class ApplicationServices:
storage: AssistantStorage
ai_client: AIClient
timezone: ZoneInfo
speech_recognizer: SpeechTranscriber
voice_max_duration_seconds: int
assistant_password: str

View File

@@ -3,9 +3,9 @@ import sqlite3
from contextlib import contextmanager from contextlib import contextmanager
from datetime import datetime from datetime import datetime
from pathlib import Path from pathlib import Path
from typing import Any from typing import Any, Mapping
from .config import get_default_model from .migrations import migrate_database
from .time_utils import to_utc_iso, utc_now from .time_utils import to_utc_iso, utc_now
@@ -21,8 +21,13 @@ def required_lastrowid(cursor: sqlite3.Cursor) -> int:
class AssistantStorage: class AssistantStorage:
def __init__(self, path: Path) -> None: def __init__(
self,
path: Path,
default_models: Mapping[str, str],
) -> None:
self.path = path self.path = path
self._default_models = dict(default_models)
self.path.parent.mkdir(parents=True, exist_ok=True) self.path.parent.mkdir(parents=True, exist_ok=True)
self._init_db() self._init_db()
@@ -45,86 +50,25 @@ class AssistantStorage:
def _init_db(self) -> None: def _init_db(self) -> None:
with self._connection() as connection: with self._connection() as connection:
connection.execute("PRAGMA journal_mode=WAL") migrate_database(connection)
connection.executescript(
"""
CREATE TABLE IF NOT EXISTS user_settings (
user_id INTEGER PRIMARY KEY,
ollama_model TEXT,
yandex_model TEXT,
updated_at TEXT NOT NULL
);
CREATE TABLE IF NOT EXISTS memories ( def is_user_authorized(self, user_id: int) -> bool:
id INTEGER PRIMARY KEY AUTOINCREMENT, with self._connection() as connection:
user_id INTEGER NOT NULL, row = connection.execute(
text TEXT NOT NULL, "SELECT 1 FROM authorized_users WHERE user_id = ?",
created_at TEXT NOT NULL (user_id,),
); ).fetchone()
return row is not None
CREATE TABLE IF NOT EXISTS notes ( def authorize_user(self, user_id: int) -> None:
id INTEGER PRIMARY KEY AUTOINCREMENT, with self._connection() as connection:
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( connection.execute(
"ALTER TABLE user_settings ADD COLUMN yandex_model TEXT" """
INSERT INTO authorized_users(user_id, authorized_at)
VALUES (?, ?)
ON CONFLICT(user_id) DO NOTHING
""",
(user_id, to_utc_iso(utc_now())),
) )
def add_conversation_exchange( def add_conversation_exchange(
@@ -288,7 +232,12 @@ class AssistantStorage:
).fetchone() ).fetchone()
if row and row["model"]: if row and row["model"]:
return str(row["model"]) return str(row["model"])
return get_default_model(provider) try:
return self._default_models[provider]
except KeyError as exc:
raise ValueError(
f"No default model configured for provider: {provider}"
) from exc
def set_user_model( def set_user_model(
self, self,

View File

@@ -1,59 +1,300 @@
from html import escape import asyncio
import logging
import re
from collections.abc import AsyncIterator, Awaitable, Callable
from contextlib import asynccontextmanager, suppress
from uuid import uuid4 from uuid import uuid4
from zoneinfo import ZoneInfo from zoneinfo import ZoneInfo
from telegram import InlineQueryResultArticle, InputTextMessageContent, Message, Update from telegram import Bot, InlineQueryResultArticle, InputTextMessageContent, Message, Update
from telegram.constants import ParseMode from telegram.constants import ChatAction, ParseMode
from telegram.error import BadRequest
from telegram.ext import ContextTypes from telegram.ext import ContextTypes
from telegramify_markdown import (
convert,
entities_to_markdownv2,
split_entities,
utf16_len,
)
from telegramify_markdown.config import RenderConfig
from .ai import AIClient from .ai import AIClient
from .config import HTML_FORMATS, MAX_TELEGRAM_MESSAGE_LENGTH from .config import MAX_TELEGRAM_MESSAGE_LENGTH
from .services import SERVICES_KEY, ApplicationServices
from .speech import SpeechTranscriber from .speech import SpeechTranscriber
from .storage import AssistantStorage from .storage import AssistantStorage
logger = logging.getLogger(__name__)
FORMATTING_FALLBACK_NOTICE = "⚠️ Не удалось применить форматирование."
TYPING_REFRESH_INTERVAL_SECONDS = 4.0
_MARKDOWN_SPECIAL_CHARS = re.compile(r"([\\`*_[\]{}()#+\-.!|>~])")
_RENDER_CONFIG = RenderConfig()
_RENDER_CONFIG.markdown_symbol.heading_level_1 = ""
_RENDER_CONFIG.markdown_symbol.heading_level_2 = ""
_RENDER_CONFIG.markdown_symbol.heading_level_3 = ""
_RENDER_CONFIG.markdown_symbol.heading_level_4 = ""
async def _refresh_typing(
bot: Bot,
chat_id: int,
interval: float,
) -> None:
while True:
await asyncio.sleep(interval)
try:
await bot.send_chat_action(
chat_id=chat_id,
action=ChatAction.TYPING,
)
except asyncio.CancelledError:
raise
except Exception:
logger.warning("Failed to refresh Telegram typing action", exc_info=True)
@asynccontextmanager
async def typing_action(
bot: Bot,
chat_id: int,
interval: float = TYPING_REFRESH_INTERVAL_SECONDS,
) -> AsyncIterator[None]:
"""Keep Telegram's typing indicator active while work is in progress."""
try:
await bot.send_chat_action(
chat_id=chat_id,
action=ChatAction.TYPING,
)
except Exception:
logger.warning("Failed to send Telegram typing action", exc_info=True)
refresh_task = asyncio.create_task(
_refresh_typing(bot, chat_id, interval),
name=f"telegram-typing-{chat_id}",
)
try:
yield
finally:
refresh_task.cancel()
with suppress(asyncio.CancelledError):
await refresh_task
def escape_markdown_text(value: object) -> str:
"""Escape dynamic text before embedding it into ordinary Markdown."""
return _MARKDOWN_SPECIAL_CHARS.sub(r"\\\1", str(value))
def markdown_code(value: object) -> str:
"""Render a dynamic value as a safe CommonMark inline-code span."""
text = str(value)
longest_run = max((len(run) for run in re.findall(r"`+", text)), default=0)
delimiter = "`" * (longest_run + 1)
padding = " " if text.startswith(("`", " ")) or text.endswith(("`", " ")) else ""
return f"{delimiter}{padding}{text}{padding}{delimiter}"
def render_markdown_chunks(markdown: str) -> list[tuple[str, str]]:
"""Return (MarkdownV2, plain text) pairs that fit Telegram's limit."""
plain_text, entities = convert(
markdown,
latex_escape=False,
config=_RENDER_CONFIG,
)
pending = list(
split_entities(
plain_text,
entities,
max_utf16_len=MAX_TELEGRAM_MESSAGE_LENGTH,
)
)
chunks: list[tuple[str, str]] = []
while pending:
chunk_text, chunk_entities = pending.pop(0)
markdown_v2 = entities_to_markdownv2(chunk_text, chunk_entities)
if utf16_len(markdown_v2) <= MAX_TELEGRAM_MESSAGE_LENGTH:
chunks.append((markdown_v2, chunk_text))
continue
plain_length = utf16_len(chunk_text)
if plain_length <= 1:
raise ValueError("Unable to split MarkdownV2 within Telegram's limit")
sub_limit = max(1, plain_length // 2)
sub_chunks = list(
split_entities(
chunk_text,
chunk_entities,
max_utf16_len=sub_limit,
)
)
if len(sub_chunks) == 1:
sub_chunks = list(
split_entities(
chunk_text,
chunk_entities,
max_utf16_len=max(1, plain_length - 1),
)
)
if len(sub_chunks) == 1:
raise ValueError("Unable to split MarkdownV2 within Telegram's limit")
pending = sub_chunks + pending
return chunks
def _plain_chunks(text: str) -> list[str]:
chunks = split_entities(
text,
[],
max_utf16_len=MAX_TELEGRAM_MESSAGE_LENGTH,
)
return [chunk_text for chunk_text, _entities in chunks] or [text]
def _prepare_chunks(markdown: str) -> tuple[list[tuple[str, str]], bool]:
try:
chunks = render_markdown_chunks(markdown)
except Exception:
logger.exception("Failed to convert Markdown to Telegram MarkdownV2")
return [(chunk, chunk) for chunk in _plain_chunks(markdown)], True
return chunks or [("", "")], False
def _with_fallback_notice(text: str) -> str:
return f"{text}\n\n{FORMATTING_FALLBACK_NOTICE}"
def _with_markdown_notice(markdown_v2: str) -> str:
notice = entities_to_markdownv2(FORMATTING_FALLBACK_NOTICE, [])
return f"{markdown_v2}\n\n{notice}"
async def reply_markdown(message: Message, markdown: str) -> Message:
async def send_formatted(text: str) -> Message:
return await message.reply_text(
text,
parse_mode=ParseMode.MARKDOWN_V2,
)
async def send_plain(text: str) -> Message:
return await message.reply_text(text)
return await _send_markdown_chunks(
markdown,
send_formatted,
send_plain,
)
MarkdownSender = Callable[[str], Awaitable[Message]]
async def _send_markdown_chunks(
markdown: str,
send_formatted: MarkdownSender,
send_plain: MarkdownSender,
) -> Message:
chunks, formatting_failed = _prepare_chunks(markdown)
first_message: Message | None = None
for index, (markdown_v2, plain_text) in enumerate(chunks):
is_last = index == len(chunks) - 1
formatted_text = (
_with_markdown_notice(markdown_v2)
if is_last and formatting_failed
else markdown_v2
)
try:
sent = await send_formatted(formatted_text)
except BadRequest:
logger.warning("Telegram rejected MarkdownV2; retrying as plain text")
formatting_failed = True
fallback_text = (
_with_fallback_notice(plain_text) if is_last else plain_text
)
sent = await send_plain(fallback_text)
if first_message is None:
first_message = sent
assert first_message is not None
return first_message
async def edit_markdown(message: Message, markdown: str) -> None:
await _edit_or_reply(message, markdown)
async def send_markdown(bot: Bot, chat_id: int, markdown: str) -> Message:
async def send_formatted(text: str) -> Message:
return await bot.send_message(
chat_id=chat_id,
text=text,
parse_mode=ParseMode.MARKDOWN_V2,
)
async def send_plain(text: str) -> Message:
return await bot.send_message(chat_id=chat_id, text=text)
return await _send_markdown_chunks(
markdown,
send_formatted,
send_plain,
)
def make_article( def make_article(
title: str, title: str,
text: str, markdown: str,
parse_mode: str | None = None,
) -> InlineQueryResultArticle: ) -> InlineQueryResultArticle:
markdown_v2 = render_markdown_chunks(markdown)[0][0]
return InlineQueryResultArticle( return InlineQueryResultArticle(
id=f"inline_{uuid4()}", id=f"inline_{uuid4()}",
title=title, title=title,
input_message_content=InputTextMessageContent(text, parse_mode=parse_mode), input_message_content=InputTextMessageContent(
markdown_v2,
parse_mode=ParseMode.MARKDOWN_V2,
),
) )
def build_inline_results(query: str) -> list[InlineQueryResultArticle]: def build_inline_results(query: str) -> list[InlineQueryResultArticle]:
escaped_query = escape(query) escaped_query = escape_markdown_text(query)
return [ return [
make_article("Caps", query.upper()), make_article("Caps", escape_markdown_text(query.upper())),
*[ make_article("Bold", f"**{escaped_query}**"),
make_article(title, f"<{tag}>{escaped_query}</{tag}>", ParseMode.HTML) make_article("Italic", f"*{escaped_query}*"),
for title, tag in HTML_FORMATS
],
] ]
def get_services(
context: ContextTypes.DEFAULT_TYPE,
) -> ApplicationServices:
return context.application.bot_data[SERVICES_KEY]
def get_storage(context: ContextTypes.DEFAULT_TYPE) -> AssistantStorage: def get_storage(context: ContextTypes.DEFAULT_TYPE) -> AssistantStorage:
return context.application.bot_data["storage"] return get_services(context).storage
def get_ai_client(context: ContextTypes.DEFAULT_TYPE) -> AIClient: def get_ai_client(context: ContextTypes.DEFAULT_TYPE) -> AIClient:
return context.application.bot_data["ai_client"] return get_services(context).ai_client
def get_tz(context: ContextTypes.DEFAULT_TYPE) -> ZoneInfo: def get_tz(context: ContextTypes.DEFAULT_TYPE) -> ZoneInfo:
return context.application.bot_data["timezone"] return get_services(context).timezone
def get_speech_recognizer(context: ContextTypes.DEFAULT_TYPE) -> SpeechTranscriber: def get_speech_recognizer(context: ContextTypes.DEFAULT_TYPE) -> SpeechTranscriber:
return context.application.bot_data["speech_recognizer"] return get_services(context).speech_recognizer
def get_voice_max_duration(context: ContextTypes.DEFAULT_TYPE) -> int: def get_voice_max_duration(context: ContextTypes.DEFAULT_TYPE) -> int:
return int(context.application.bot_data["voice_max_duration"]) return get_services(context).voice_max_duration_seconds
def command_text(context: ContextTypes.DEFAULT_TYPE) -> str: def command_text(context: ContextTypes.DEFAULT_TYPE) -> str:
@@ -67,25 +308,38 @@ def require_user_id(update: Update) -> int | None:
return None return None
async def reply_long(message: Message, text: str) -> None: async def _edit_or_reply(message: Message, text: str) -> None:
chunks = [ chunks, formatting_failed = _prepare_chunks(text)
text[index : index + MAX_TELEGRAM_MESSAGE_LENGTH]
for index in range(0, len(text), MAX_TELEGRAM_MESSAGE_LENGTH)
] or [""]
for chunk in chunks: for index, (markdown_v2, plain_text) in enumerate(chunks):
await message.reply_text(chunk) is_first = index == 0
is_last = index == len(chunks) - 1
formatted_text = (
async def edit_or_reply(message: Message, text: str) -> None: _with_markdown_notice(markdown_v2)
chunks = [ if is_last and formatting_failed
text[index : index + MAX_TELEGRAM_MESSAGE_LENGTH] else markdown_v2
for index in range(0, len(text), MAX_TELEGRAM_MESSAGE_LENGTH) )
] or [""] try:
if is_first:
await message.edit_text(chunks[0]) await message.edit_text(
for chunk in chunks[1:]: formatted_text,
await message.reply_text(chunk) parse_mode=ParseMode.MARKDOWN_V2,
)
else:
await message.reply_text(
formatted_text,
parse_mode=ParseMode.MARKDOWN_V2,
)
except BadRequest:
logger.warning("Telegram rejected MarkdownV2; retrying as plain text")
formatting_failed = True
fallback_text = (
_with_fallback_notice(plain_text) if is_last else plain_text
)
if is_first:
await message.edit_text(fallback_text)
else:
await message.reply_text(fallback_text)
def parse_positive_int(value: str) -> int | None: def parse_positive_int(value: str) -> int | None:

16
pyproject.toml Normal file
View File

@@ -0,0 +1,16 @@
[tool.ruff]
target-version = "py310"
line-length = 100
extend-exclude = [".agents"]
[tool.ruff.lint]
select = ["E4", "E7", "E9", "F"]
[tool.coverage.run]
branch = true
source = ["assistant_bot"]
[tool.coverage.report]
fail_under = 60
show_missing = true
skip_covered = true

View File

@@ -1 +1,2 @@
python-telegram-bot==22.8 python-telegram-bot==22.8
telegramify-markdown==1.2.0

2
requirements-dev.txt Normal file
View File

@@ -0,0 +1,2 @@
coverage[toml]>=7.6,<8
ruff>=0.9,<1

11
skills-lock.json Normal file
View File

@@ -0,0 +1,11 @@
{
"version": 1,
"skills": {
"grill-me": {
"source": "mattpocock/skills",
"sourceType": "github",
"skillPath": "skills/productivity/grill-me/SKILL.md",
"computedHash": "7edd436132b43ea2973710f0c28a9879c9b0910651eb89a991a7e98cb6dbbacb"
}
}
}

68
tests/test_application.py Normal file
View File

@@ -0,0 +1,68 @@
import tempfile
import unittest
from pathlib import Path
from zoneinfo import ZoneInfo
from assistant_bot import application as application_module
from assistant_bot.auth import authentication_guard
from assistant_bot.config import AppSettings
from assistant_bot.handlers import (
ask_command,
help_command,
history_command,
inline_query,
model_command,
models_command,
new_conversation_command,
private_text,
private_voice,
start,
)
class ApplicationRegistrationTests(unittest.TestCase):
def test_registers_the_current_public_handlers_in_order(self) -> None:
with tempfile.TemporaryDirectory() as directory:
settings = AppSettings(
bot_token="123456:TEST",
assistant_password="secret",
db_path=Path(directory) / "assistant.sqlite3",
mode="local",
timezone=ZoneInfo("UTC"),
voice_max_duration_seconds=120,
ollama_base_url="http://localhost:11434",
ollama_model="qwen3.5:9b",
yandex_cloud_folder=None,
yandex_cloud_model="yandexgpt/latest",
yandex_stt_model="general",
yandex_stt_language="ru-RU",
whisper_model="small",
whisper_device="cpu",
whisper_compute_type="int8",
whisper_language="ru",
)
application = application_module.create_application(settings)
self.assertEqual(
[handler.callback for handler in application.handlers[-1]],
[authentication_guard],
)
self.assertEqual(
[handler.callback for handler in application.handlers[0]],
[
start,
help_command,
ask_command,
models_command,
model_command,
new_conversation_command,
history_command,
inline_query,
private_voice,
private_text,
],
)
if __name__ == "__main__":
unittest.main()

93
tests/test_auth.py Normal file
View File

@@ -0,0 +1,93 @@
import tempfile
import unittest
from pathlib import Path
from types import SimpleNamespace
from typing import Any
from unittest.mock import AsyncMock
from zoneinfo import ZoneInfo
from telegram.constants import ChatType, ParseMode
from telegram.ext import ApplicationHandlerStop
from assistant_bot.auth import authentication_guard
from assistant_bot.services import SERVICES_KEY, ApplicationServices
from assistant_bot.storage import AssistantStorage
from assistant_bot.telegram_utils import render_markdown_chunks
class AuthenticationGuardTests(unittest.IsolatedAsyncioTestCase):
def setUp(self) -> None:
self.temporary_directory = tempfile.TemporaryDirectory()
self.storage = AssistantStorage(
Path(self.temporary_directory.name) / "assistant.sqlite3",
default_models={"local": "qwen3.5:9b"},
)
def tearDown(self) -> None:
self.temporary_directory.cleanup()
def make_update_and_context(self, text: str) -> tuple[Any, Any]:
message = SimpleNamespace(text=text, reply_text=AsyncMock())
update = SimpleNamespace(
effective_user=SimpleNamespace(id=42),
effective_chat=SimpleNamespace(type=ChatType.PRIVATE),
effective_message=message,
inline_query=None,
)
context = SimpleNamespace(
application=SimpleNamespace(
bot_data={
SERVICES_KEY: ApplicationServices(
storage=self.storage,
ai_client=SimpleNamespace(),
timezone=ZoneInfo("UTC"),
speech_recognizer=SimpleNamespace(),
voice_max_duration_seconds=120,
assistant_password="секрет",
),
}
)
)
return update, context
async def test_correct_password_authorizes_user(self) -> None:
update, context = self.make_update_and_context("секрет")
with self.assertRaises(ApplicationHandlerStop):
await authentication_guard(update, context)
self.assertTrue(self.storage.is_user_authorized(42))
rendered = render_markdown_chunks(
"✅ **Пароль принят.** Доступ открыт — повторно вводить его не нужно."
)[0][0]
update.effective_message.reply_text.assert_awaited_once_with(
rendered,
parse_mode=ParseMode.MARKDOWN_V2,
)
async def test_wrong_password_does_not_authorize_user(self) -> None:
update, context = self.make_update_and_context("неверно")
with self.assertRaises(ApplicationHandlerStop):
await authentication_guard(update, context)
self.assertFalse(self.storage.is_user_authorized(42))
rendered = render_markdown_chunks(
"🔐 **Неверный пароль.** Попробуй ещё раз."
)[0][0]
update.effective_message.reply_text.assert_awaited_once_with(
rendered,
parse_mode=ParseMode.MARKDOWN_V2,
)
async def test_authorized_user_passes_guard(self) -> None:
self.storage.authorize_user(42)
update, context = self.make_update_and_context("обычное сообщение")
await authentication_guard(update, context)
update.effective_message.reply_text.assert_not_awaited()
if __name__ == "__main__":
unittest.main()

View File

@@ -5,67 +5,102 @@ from unittest.mock import patch
from assistant_bot import config from assistant_bot import config
class WhisperConfigTests(unittest.TestCase): class AppSettingsTests(unittest.TestCase):
@staticmethod
def load_settings(**overrides):
environment = {
config.TOKEN_ENV_NAME: "token",
config.PASSWORD_ENV_NAME: "secret",
**overrides,
}
with patch("assistant_bot.config.load_env_file"), patch.dict(
os.environ,
environment,
clear=True,
):
return config.load_app_settings()
def test_loads_complete_startup_configuration_once(self) -> None:
environment = {
config.TOKEN_ENV_NAME: "token",
config.PASSWORD_ENV_NAME: "secret",
config.MODE_ENV_NAME: "local",
config.DB_ENV_NAME: "data/test.sqlite3",
config.TIMEZONE_ENV_NAME: "UTC",
config.WHISPER_LANGUAGE_ENV_NAME: "auto",
}
with patch(
"assistant_bot.config.load_env_file"
) as load_env_file, patch.dict(
os.environ,
environment,
clear=True,
):
settings = config.load_app_settings()
load_env_file.assert_called_once_with()
self.assertEqual(settings.bot_token, "token")
self.assertEqual(settings.assistant_password, "secret")
self.assertEqual(settings.mode, "local")
self.assertEqual(settings.db_path, config.PROJECT_ROOT / "data/test.sqlite3")
self.assertEqual(settings.timezone.key, "UTC")
self.assertIsNone(settings.whisper_language)
self.assertIsNone(settings.yandex_cloud_folder)
def test_gpu_int8_defaults(self) -> None: def test_gpu_int8_defaults(self) -> None:
variable_names = ( settings = self.load_settings()
config.WHISPER_MODEL_ENV_NAME,
config.WHISPER_DEVICE_ENV_NAME,
config.WHISPER_COMPUTE_TYPE_ENV_NAME,
)
with patch("assistant_bot.config.load_env_file"), patch.dict(
os.environ,
{},
clear=False,
):
for variable_name in variable_names:
os.environ.pop(variable_name, None)
self.assertEqual(config.get_whisper_model(), "large-v3") self.assertEqual(settings.whisper_model, "large-v3")
self.assertEqual(config.get_whisper_device(), "cuda") self.assertEqual(settings.whisper_device, "cuda")
self.assertEqual(config.get_whisper_compute_type(), "int8") self.assertEqual(settings.whisper_compute_type, "int8")
class AssistantModeConfigTests(unittest.TestCase):
def test_local_mode_is_default(self) -> None: def test_local_mode_is_default(self) -> None:
with patch("assistant_bot.config.load_env_file"), patch.dict( self.assertEqual(self.load_settings().mode, "local")
os.environ,
{},
clear=False,
):
os.environ.pop(config.MODE_ENV_NAME, None)
self.assertEqual(config.get_assistant_mode(), "local")
def test_yandex_mode_is_supported(self) -> None: def test_yandex_mode_is_supported(self) -> None:
with patch("assistant_bot.config.load_env_file"), patch.dict( settings = self.load_settings(
os.environ, **{
{config.MODE_ENV_NAME: "YANDEX"}, config.MODE_ENV_NAME: "YANDEX",
clear=False, config.YANDEX_CLOUD_FOLDER_ENV_NAME: "folder-id",
): }
self.assertEqual(config.get_assistant_mode(), "yandex") )
self.assertEqual(settings.mode, "yandex")
def test_unknown_mode_is_rejected(self) -> None: def test_unknown_mode_is_rejected(self) -> None:
with patch("assistant_bot.config.load_env_file"), patch.dict(
os.environ,
{config.MODE_ENV_NAME: "cloud"},
clear=False,
):
with self.assertRaises(RuntimeError): with self.assertRaises(RuntimeError):
config.get_assistant_mode() self.load_settings(**{config.MODE_ENV_NAME: "cloud"})
def test_yandex_model_is_full_gpt_uri(self) -> None: def test_yandex_model_is_full_gpt_uri(self) -> None:
with patch("assistant_bot.config.load_env_file"), patch.dict( settings = self.load_settings(
os.environ, **{
{ config.MODE_ENV_NAME: "yandex",
config.YANDEX_CLOUD_FOLDER_ENV_NAME: "folder-id", config.YANDEX_CLOUD_FOLDER_ENV_NAME: "folder-id",
config.YANDEX_CLOUD_MODEL_ENV_NAME: "yandexgpt/latest", config.YANDEX_CLOUD_MODEL_ENV_NAME: "yandexgpt/latest",
}, }
clear=False, )
):
self.assertEqual( self.assertEqual(
config.get_default_yandex_model(), settings.default_models()["yandex"],
"gpt://folder-id/yandexgpt/latest", "gpt://folder-id/yandexgpt/latest",
) )
def test_password_is_required(self) -> None:
with patch("assistant_bot.config.load_env_file"), patch.dict(
os.environ,
{config.TOKEN_ENV_NAME: "token"},
clear=True,
):
with self.assertRaises(RuntimeError):
config.load_app_settings()
def test_password_is_read_from_environment(self) -> None:
settings = self.load_settings(
**{config.PASSWORD_ENV_NAME: " test-password "}
)
self.assertEqual(settings.assistant_password, "test-password")
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()

View File

@@ -6,11 +6,27 @@ from datetime import datetime, timezone
from pathlib import Path from pathlib import Path
from zoneinfo import ZoneInfo from zoneinfo import ZoneInfo
from assistant_bot.agent import execute_agent_tool, parse_agent_decision from assistant_bot.agent import (
build_agent_tool_prompt,
execute_agent_tool,
parse_agent_decision,
)
from assistant_bot.agent_tools import AGENT_TOOLS
from assistant_bot.migrations import LATEST_SCHEMA_VERSION
from assistant_bot.reminders import parse_reminder, unit_to_timedelta from assistant_bot.reminders import parse_reminder, unit_to_timedelta
from assistant_bot.storage import AssistantStorage from assistant_bot.storage import AssistantStorage
DEFAULT_MODELS = {
"local": "qwen3.5:9b",
"yandex": "gpt://folder/yandexgpt/latest",
}
def create_storage(path: Path) -> AssistantStorage:
return AssistantStorage(path, default_models=DEFAULT_MODELS)
class ReminderParserTests(unittest.TestCase): class ReminderParserTests(unittest.TestCase):
def test_supported_relative_unit(self) -> None: def test_supported_relative_unit(self) -> None:
delta = unit_to_timedelta(2, "часа") delta = unit_to_timedelta(2, "часа")
@@ -52,7 +68,173 @@ class AgentDecisionTests(unittest.TestCase):
self.assertTrue(decision.reset_context) self.assertTrue(decision.reset_context)
class AgentToolContractTests(unittest.TestCase):
TOOL_NAMES = (
"get_current_datetime",
"remember",
"list_memory",
"delete_memory",
"create_note",
"list_notes",
"delete_note",
"create_reminder",
"list_reminders",
"cancel_reminder",
"create_status",
"list_statuses",
"update_status",
"delete_status",
"search_conversation",
)
def test_prompt_exposes_all_supported_tool_names(self) -> None:
prompt = build_agent_tool_prompt(ZoneInfo("UTC"))
self.assertEqual(
tuple(tool.name for tool in AGENT_TOOLS),
self.TOOL_NAMES,
)
for tool_name in self.TOOL_NAMES:
with self.subTest(tool_name=tool_name):
self.assertIn(tool_name, prompt)
def test_crud_tool_result_shapes_remain_stable(self) -> None:
with tempfile.TemporaryDirectory() as directory:
storage = create_storage(Path(directory) / "assistant.sqlite3")
common = {
"storage": storage,
"user_id": 42,
"chat_id": 100,
"tz": ZoneInfo("UTC"),
}
memory = execute_agent_tool(
"remember",
{"text": "короткие ответы"},
**common,
)
self.assertEqual(
memory,
{"ok": True, "id": memory["id"], "text": "короткие ответы"},
)
self.assertEqual(
execute_agent_tool("list_memory", {"limit": 10}, **common)[
"items"
][0]["id"],
memory["id"],
)
self.assertEqual(
execute_agent_tool(
"delete_memory",
{"id": memory["id"]},
**common,
),
{"ok": True, "id": memory["id"]},
)
note = execute_agent_tool(
"create_note",
{"text": "идея"},
**common,
)
self.assertEqual(
note,
{"ok": True, "id": note["id"], "text": "идея"},
)
self.assertEqual(
execute_agent_tool("list_notes", {}, **common)["items"][0][
"id"
],
note["id"],
)
self.assertEqual(
execute_agent_tool(
"delete_note",
{"id": note["id"]},
**common,
),
{"ok": True, "id": note["id"]},
)
status = execute_agent_tool(
"create_status",
{"title": "паспорт", "status": "ожидание"},
**common,
)
self.assertEqual(
status,
{
"ok": True,
"id": status["id"],
"title": "паспорт",
"status": "ожидание",
},
)
self.assertEqual(
execute_agent_tool(
"update_status",
{"id": status["id"], "status": "готово"},
**common,
),
{"ok": True, "id": status["id"], "status": "готово"},
)
self.assertEqual(
execute_agent_tool("list_statuses", {}, **common)["items"][0][
"status"
],
"готово",
)
self.assertEqual(
execute_agent_tool(
"delete_status",
{"id": status["id"]},
**common,
),
{"ok": True, "id": status["id"]},
)
def test_unknown_tool_error_remains_stable(self) -> None:
with tempfile.TemporaryDirectory() as directory:
result = execute_agent_tool(
"missing",
{},
create_storage(Path(directory) / "assistant.sqlite3"),
user_id=42,
chat_id=100,
tz=ZoneInfo("UTC"),
)
self.assertEqual(
result,
{"ok": False, "error": "unknown tool: missing"},
)
class StorageTests(unittest.TestCase): class StorageTests(unittest.TestCase):
def test_new_database_uses_latest_schema_version(self) -> None:
with tempfile.TemporaryDirectory() as directory:
database_path = Path(directory) / "assistant.sqlite3"
create_storage(database_path)
with closing(sqlite3.connect(database_path)) as connection:
version = connection.execute(
"PRAGMA user_version"
).fetchone()[0]
self.assertEqual(version, LATEST_SCHEMA_VERSION)
def test_user_authorization_is_persisted(self) -> None:
with tempfile.TemporaryDirectory() as directory:
database_path = Path(directory) / "assistant.sqlite3"
storage = create_storage(database_path)
self.assertFalse(storage.is_user_authorized(42))
storage.authorize_user(42)
reopened_storage = create_storage(database_path)
self.assertTrue(reopened_storage.is_user_authorized(42))
self.assertFalse(reopened_storage.is_user_authorized(43))
def test_legacy_user_settings_gets_yandex_model_column(self) -> None: def test_legacy_user_settings_gets_yandex_model_column(self) -> None:
with tempfile.TemporaryDirectory() as directory: with tempfile.TemporaryDirectory() as directory:
database_path = Path(directory) / "assistant.sqlite3" database_path = Path(directory) / "assistant.sqlite3"
@@ -68,14 +250,32 @@ class StorageTests(unittest.TestCase):
) )
connection.commit() connection.commit()
storage = AssistantStorage(database_path) storage = create_storage(database_path)
storage.set_user_model(42, "yandexgpt", "yandex") storage.set_user_model(42, "yandexgpt", "yandex")
self.assertEqual(storage.get_user_model(42, "yandex"), "yandexgpt") self.assertEqual(storage.get_user_model(42, "yandex"), "yandexgpt")
with closing(sqlite3.connect(database_path)) as connection:
version = connection.execute(
"PRAGMA user_version"
).fetchone()[0]
self.assertEqual(version, LATEST_SCHEMA_VERSION)
def test_default_models_are_injected(self) -> None:
with tempfile.TemporaryDirectory() as directory:
storage = create_storage(Path(directory) / "assistant.sqlite3")
self.assertEqual(
storage.get_user_model(42, "local"),
DEFAULT_MODELS["local"],
)
self.assertEqual(
storage.get_user_model(42, "yandex"),
DEFAULT_MODELS["yandex"],
)
def test_provider_models_are_stored_separately(self) -> None: def test_provider_models_are_stored_separately(self) -> None:
with tempfile.TemporaryDirectory() as directory: with tempfile.TemporaryDirectory() as directory:
storage = AssistantStorage(Path(directory) / "assistant.sqlite3") storage = create_storage(Path(directory) / "assistant.sqlite3")
storage.set_user_model(42, "qwen3.5:9b", "local") storage.set_user_model(42, "qwen3.5:9b", "local")
storage.set_user_model(42, "yandexgpt", "yandex") storage.set_user_model(42, "yandexgpt", "yandex")
@@ -85,7 +285,7 @@ class StorageTests(unittest.TestCase):
def test_memory_crud(self) -> None: def test_memory_crud(self) -> None:
with tempfile.TemporaryDirectory() as directory: with tempfile.TemporaryDirectory() as directory:
storage = AssistantStorage(Path(directory) / "assistant.sqlite3") storage = create_storage(Path(directory) / "assistant.sqlite3")
memory_id = storage.add_memory(42, "короткие ответы") memory_id = storage.add_memory(42, "короткие ответы")
self.assertEqual(storage.list_memories(42)[0]["id"], memory_id) self.assertEqual(storage.list_memories(42)[0]["id"], memory_id)
@@ -94,7 +294,7 @@ class StorageTests(unittest.TestCase):
def test_new_context_preserves_searchable_archive(self) -> None: def test_new_context_preserves_searchable_archive(self) -> None:
with tempfile.TemporaryDirectory() as directory: with tempfile.TemporaryDirectory() as directory:
storage = AssistantStorage(Path(directory) / "assistant.sqlite3") storage = create_storage(Path(directory) / "assistant.sqlite3")
storage.add_conversation_exchange( storage.add_conversation_exchange(
42, 42,
100, 100,

View File

@@ -0,0 +1,162 @@
import asyncio
import unittest
from types import SimpleNamespace
from unittest.mock import AsyncMock
from telegram.constants import ChatAction, ParseMode
from telegram.error import BadRequest
from assistant_bot.telegram_utils import (
FORMATTING_FALLBACK_NOTICE,
build_inline_results,
escape_markdown_text,
render_markdown_chunks,
reply_markdown,
send_markdown,
typing_action,
)
class MarkdownRenderingTests(unittest.TestCase):
def test_converts_common_markdown_to_markdown_v2(self) -> None:
chunks = render_markdown_chunks(
"**📝 Заметки**\n\n- `#3` Проект\\_v2\\.0"
)
self.assertEqual(len(chunks), 1)
markdown_v2, plain_text = chunks[0]
self.assertIn("*📝 Заметки*", markdown_v2)
self.assertIn("Проект\\_v2\\.0", markdown_v2)
self.assertNotIn("**", markdown_v2)
self.assertIn("Проект_v2.0", plain_text)
def test_splits_long_formatted_text_without_losing_format(self) -> None:
chunks = render_markdown_chunks(f"**{'слово ' * 999}слово**")
self.assertGreater(len(chunks), 1)
for markdown_v2, plain_text in chunks:
self.assertLessEqual(len(plain_text.encode("utf-16-le")) // 2, 3900)
self.assertLessEqual(len(markdown_v2.encode("utf-16-le")) // 2, 3900)
self.assertTrue(markdown_v2.startswith("*"))
self.assertTrue(markdown_v2.rstrip().endswith("*"))
def test_splits_by_rendered_markdown_v2_length(self) -> None:
chunks = render_markdown_chunks(escape_markdown_text("_" * 5000))
self.assertGreater(len(chunks), 2)
for markdown_v2, _plain_text in chunks:
self.assertLessEqual(len(markdown_v2.encode("utf-16-le")) // 2, 3900)
def test_inline_results_use_markdown_v2(self) -> None:
results = build_inline_results("проект_v2")
bold_content = results[1].input_message_content
self.assertEqual(bold_content.parse_mode, ParseMode.MARKDOWN_V2)
self.assertEqual(bold_content.message_text, "*проект\\_v2*")
class MarkdownDeliveryTests(unittest.IsolatedAsyncioTestCase):
async def test_send_markdown_uses_the_shared_formatted_delivery(self) -> None:
sent_message = SimpleNamespace()
bot = SimpleNamespace(
send_message=AsyncMock(return_value=sent_message)
)
result = await send_markdown(bot, 42, "**Готово.**")
self.assertIs(result, sent_message)
bot.send_message.assert_awaited_once_with(
chat_id=42,
text="*Готово\\.*",
parse_mode=ParseMode.MARKDOWN_V2,
)
async def test_retries_plain_text_with_notice_when_telegram_rejects_markup(
self,
) -> None:
fallback_message = SimpleNamespace()
message = SimpleNamespace(
reply_text=AsyncMock(
side_effect=[
BadRequest("Can't parse entities"),
fallback_message,
]
)
)
sent = await reply_markdown(message, "**Готово.**")
self.assertIs(sent, fallback_message)
self.assertEqual(message.reply_text.await_count, 2)
formatted_call, fallback_call = message.reply_text.await_args_list
self.assertEqual(
formatted_call.kwargs["parse_mode"],
ParseMode.MARKDOWN_V2,
)
self.assertNotIn("parse_mode", fallback_call.kwargs)
self.assertEqual(
fallback_call.args[0],
f"Готово.\n\n{FORMATTING_FALLBACK_NOTICE}",
)
async def test_places_fallback_notice_only_in_last_chunk(self) -> None:
first_message = SimpleNamespace()
markdown = f"**{'слово ' * 999}слово**"
expected_chunks = render_markdown_chunks(markdown)
call_number = 0
async def send(*_args, **_kwargs):
nonlocal call_number
call_number += 1
if call_number == 1:
raise BadRequest("Can't parse entities")
return first_message if call_number == 2 else SimpleNamespace()
message = SimpleNamespace(reply_text=AsyncMock(side_effect=send))
await reply_markdown(message, markdown)
self.assertEqual(message.reply_text.await_count, len(expected_chunks) + 1)
first_plain_call = message.reply_text.await_args_list[1]
last_formatted_call = message.reply_text.await_args_list[-1]
self.assertNotIn(FORMATTING_FALLBACK_NOTICE, first_plain_call.args[0])
self.assertIn(
"Не удалось применить форматирование",
last_formatted_call.args[0],
)
self.assertEqual(
last_formatted_call.kwargs["parse_mode"],
ParseMode.MARKDOWN_V2,
)
class TypingActionTests(unittest.IsolatedAsyncioTestCase):
async def test_refreshes_typing_until_context_exits(self) -> None:
refreshed = asyncio.Event()
call_count = 0
async def send_chat_action(**_kwargs) -> None:
nonlocal call_count
call_count += 1
if call_count >= 2:
refreshed.set()
bot = SimpleNamespace(
send_chat_action=AsyncMock(side_effect=send_chat_action)
)
async with typing_action(bot, chat_id=42, interval=0.001):
await asyncio.wait_for(refreshed.wait(), timeout=0.5)
calls_after_exit = call_count
await asyncio.sleep(0.01)
self.assertGreaterEqual(calls_after_exit, 2)
self.assertEqual(call_count, calls_after_exit)
for call in bot.send_chat_action.await_args_list:
self.assertEqual(call.kwargs["chat_id"], 42)
self.assertEqual(call.kwargs["action"], ChatAction.TYPING)
if __name__ == "__main__":
unittest.main()

View File

@@ -3,8 +3,13 @@ from pathlib import Path
from types import SimpleNamespace from types import SimpleNamespace
from typing import Any from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch from unittest.mock import AsyncMock, MagicMock, patch
from zoneinfo import ZoneInfo
from telegram.constants import ParseMode
from assistant_bot.handlers import private_voice from assistant_bot.handlers import private_voice
from assistant_bot.services import SERVICES_KEY, ApplicationServices
from assistant_bot.telegram_utils import render_markdown_chunks
class VoiceHandlerTests(unittest.IsolatedAsyncioTestCase): class VoiceHandlerTests(unittest.IsolatedAsyncioTestCase):
@@ -28,8 +33,14 @@ class VoiceHandlerTests(unittest.IsolatedAsyncioTestCase):
bot=bot, bot=bot,
application=SimpleNamespace( application=SimpleNamespace(
bot_data={ bot_data={
"speech_recognizer": recognizer, SERVICES_KEY: ApplicationServices(
"voice_max_duration": 120, storage=SimpleNamespace(),
ai_client=SimpleNamespace(),
timezone=ZoneInfo("UTC"),
speech_recognizer=recognizer,
voice_max_duration_seconds=120,
assistant_password="secret",
),
} }
), ),
) )
@@ -42,8 +53,12 @@ class VoiceHandlerTests(unittest.IsolatedAsyncioTestCase):
await private_voice(update, context) await private_voice(update, context)
rendered = render_markdown_chunks(
"🎙️ **Голосовое сообщение слишком длинное.** Максимум: `120` сек."
)[0][0]
update.message.reply_text.assert_awaited_once_with( update.message.reply_text.assert_awaited_once_with(
"Голосовое сообщение слишком длинное. Максимум: 120 сек." rendered,
parse_mode=ParseMode.MARKDOWN_V2,
) )
recognizer.transcribe.assert_not_called() recognizer.transcribe.assert_not_called()
@@ -64,7 +79,13 @@ class VoiceHandlerTests(unittest.IsolatedAsyncioTestCase):
telegram_file.download_to_drive.assert_awaited_once_with( telegram_file.download_to_drive.assert_awaited_once_with(
custom_path=temporary_path custom_path=temporary_path
) )
status.edit_text.assert_awaited_once_with("Распознано: Напомни позвонить") rendered = render_markdown_chunks(
"🎙️ **Распознано:** Напомни позвонить"
)[0][0]
status.edit_text.assert_awaited_once_with(
rendered,
parse_mode=ParseMode.MARKDOWN_V2,
)
run_agent.assert_awaited_once_with( run_agent.assert_awaited_once_with(
update, update,
context, context,