import asyncio import io import logging import os import re import sqlite3 from dataclasses import dataclass from pathlib import Path from aiogram import Bot, Dispatcher, F, Router from aiogram.client.default import DefaultBotProperties from aiogram.client.session.aiohttp import AiohttpSession from aiogram.enums import ChatAction, ParseMode from aiogram.exceptions import TelegramBadRequest from aiogram.filters import Command, CommandStart from aiogram.types import Message from dotenv import load_dotenv from httpx import AsyncClient as HttpxAsyncClient from openai import APIConnectionError, APIError, APITimeoutError, AsyncOpenAI, NotFoundError, RateLimitError load_dotenv() logger = logging.getLogger(__name__) router = Router() BREAK_RE = re.compile(r"\d{1,2}))?\s*/?>", re.IGNORECASE) @dataclass(frozen=True) class Settings: bot_token: str api_base_url: str api_token: str model: str transcription_model: str telegram_proxy: str | None whitelist_user_ids: frozenset[int] system_prompt: str api_timeout_seconds: float api_max_retries: int max_prompt_chars: int max_response_chars: int database_path: Path max_context_messages: int message_batch_delay_seconds: float def _required_env(name: str) -> str: value = os.getenv(name, "").strip() if not value: raise RuntimeError(f"{name} is required") return value def _int_env(name: str, default: int) -> int: value = os.getenv(name, "").strip() if not value: return default try: parsed = int(value) except ValueError as exc: raise RuntimeError(f"{name} must be an integer") from exc if parsed <= 0: raise RuntimeError(f"{name} must be greater than zero") return parsed def _float_env(name: str, default: float) -> float: value = os.getenv(name, "").strip() if not value: return default try: parsed = float(value) except ValueError as exc: raise RuntimeError(f"{name} must be a number") from exc if parsed <= 0: raise RuntimeError(f"{name} must be greater than zero") return parsed def _whitelist_env() -> frozenset[int]: raw = os.getenv("WHITELIST_USER_IDS", "").strip() if not raw: raise RuntimeError("WHITELIST_USER_IDS is required; do not run an open public proxy bot") ids: set[int] = set() for item in raw.split(","): value = item.strip() if not value: continue try: ids.add(int(value)) except ValueError as exc: raise RuntimeError("WHITELIST_USER_IDS must contain only numeric Telegram user IDs") from exc if not ids: raise RuntimeError("WHITELIST_USER_IDS must contain at least one Telegram user ID") return frozenset(ids) def _load_system_prompt() -> str: prompt_path = Path(os.getenv("SYSTEM_PROMPT_PATH", "config/system_prompt.txt")).expanduser() if not prompt_path.exists() or not prompt_path.is_file(): raise RuntimeError(f"System prompt file not found: {prompt_path}") prompt = prompt_path.read_text(encoding="utf-8").strip() if not prompt: raise RuntimeError("System prompt file is empty") return prompt def _api_base_url_env() -> str: base_url = _required_env("API_BASE_URL").rstrip("/") chat_completions_suffix = "/chat/completions" if base_url.endswith(chat_completions_suffix): normalized = base_url[: -len(chat_completions_suffix)] logger.warning("API_BASE_URL should be the API root, normalized to %s", normalized) return normalized return base_url def load_settings() -> Settings: return Settings( bot_token=_required_env("BOT_TOKEN"), api_base_url=_api_base_url_env(), api_token=_required_env("API_TOKEN"), model=_required_env("MODEL"), transcription_model=_required_env("TRANSCRIPTION_MODEL"), telegram_proxy=os.getenv("TELEGRAM_PROXY", "").strip() or None, whitelist_user_ids=_whitelist_env(), system_prompt=_load_system_prompt(), api_timeout_seconds=_float_env("API_TIMEOUT_SECONDS", 60), api_max_retries=_int_env("API_MAX_RETRIES", 2), max_prompt_chars=_int_env("MAX_PROMPT_CHARS", 12000), max_response_chars=_int_env("MAX_RESPONSE_CHARS", 3900), database_path=Path(os.getenv("DATABASE_PATH", "data/bot.sqlite3")).expanduser(), max_context_messages=_int_env("MAX_CONTEXT_MESSAGES", 20), message_batch_delay_seconds=_float_env("MESSAGE_BATCH_DELAY_SECONDS", 3), ) settings = load_settings() http_client = HttpxAsyncClient(proxy=settings.telegram_proxy) if settings.telegram_proxy else None client = AsyncOpenAI( api_key=settings.api_token, base_url=settings.api_base_url, timeout=settings.api_timeout_seconds, max_retries=settings.api_max_retries, http_client=http_client, ) user_locks: dict[int, asyncio.Lock] = {} pending_user_messages: dict[int, list[str]] = {} pending_tasks: dict[int, asyncio.Task[None]] = {} class ChatHistory: def __init__(self, database_path: Path) -> None: self.database_path = database_path async def init(self) -> None: await asyncio.to_thread(self._init_sync) def _connect(self) -> sqlite3.Connection: return sqlite3.connect(self.database_path) def _init_sync(self) -> None: self.database_path.parent.mkdir(parents=True, exist_ok=True) with self._connect() as connection: connection.execute("PRAGMA journal_mode=WAL") connection.execute( """ CREATE TABLE IF NOT EXISTS messages ( id INTEGER PRIMARY KEY AUTOINCREMENT, user_id INTEGER NOT NULL, role TEXT NOT NULL CHECK(role IN ('user', 'assistant')), content TEXT NOT NULL, created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP ) """ ) connection.execute( "CREATE INDEX IF NOT EXISTS idx_messages_user_id_id ON messages(user_id, id)" ) async def get_context(self, user_id: int) -> list[dict[str, str]]: return await asyncio.to_thread(self._get_context_sync, user_id) def _get_context_sync(self, user_id: int) -> list[dict[str, str]]: with self._connect() as connection: rows = connection.execute( """ SELECT role, content FROM messages WHERE user_id = ? ORDER BY id DESC LIMIT ? """, (user_id, settings.max_context_messages), ).fetchall() return [{"role": role, "content": content} for role, content in reversed(rows)] async def append_turn(self, user_id: int, user_text: str, assistant_text: str) -> None: await asyncio.to_thread(self._append_turn_sync, user_id, user_text, assistant_text) def _append_turn_sync(self, user_id: int, user_text: str, assistant_text: str) -> None: with self._connect() as connection: connection.executemany( "INSERT INTO messages(user_id, role, content) VALUES (?, ?, ?)", ( (user_id, "user", clamp_user_text(user_text)), (user_id, "assistant", assistant_text[: settings.max_response_chars]), ), ) async def reset(self, user_id: int) -> None: await asyncio.to_thread(self._reset_sync, user_id) def _reset_sync(self, user_id: int) -> None: with self._connect() as connection: connection.execute("DELETE FROM messages WHERE user_id = ?", (user_id,)) history = ChatHistory(settings.database_path) def is_allowed(message: Message) -> bool: return bool(message.from_user and message.from_user.id in settings.whitelist_user_ids) def user_lock(user_id: int) -> asyncio.Lock: lock = user_locks.get(user_id) if lock is None: lock = asyncio.Lock() user_locks[user_id] = lock return lock def combine_user_messages(messages: list[str]) -> str: if len(messages) == 1: return messages[0] return "\n".join(f"Message {index}: {text}" for index, text in enumerate(messages, start=1)) def voice_transcript_text(transcript: str) -> str: return f"[Voice message transcript]\n{transcript.strip()}" async def keep_typing(bot: Bot, chat_id: int, stop: asyncio.Event) -> None: while not stop.is_set(): try: await bot.send_chat_action(chat_id=chat_id, action=ChatAction.TYPING) except Exception: logger.exception("Failed to send typing action") try: await asyncio.wait_for(stop.wait(), timeout=4.0) except TimeoutError: continue def clamp_user_text(text: str) -> str: if len(text) <= settings.max_prompt_chars: return text return text[-settings.max_prompt_chars :] async def stream_model_reply(user_text: str, context: list[dict[str, str]]) -> str: chunks: list[str] = [] response = await client.chat.completions.create( model=settings.model, messages=[ {"role": "system", "content": settings.system_prompt}, *context, {"role": "user", "content": clamp_user_text(user_text)}, ], stream=True, ) async for chunk in response: delta = chunk.choices[0].delta.content if chunk.choices else None if not delta: continue chunks.append(delta) if sum(len(part) for part in chunks) >= settings.max_response_chars: chunks.append("\n\n[Response truncated]") break return "".join(chunks).strip() async def download_voice_message(bot: Bot, message: Message) -> tuple[bytes, str]: if not message.voice: raise ValueError("Message does not contain a voice attachment") file = await bot.get_file(message.voice.file_id) if not file.file_path: raise RuntimeError("Telegram did not return a file path for the voice message") buffer = io.BytesIO() await bot.download_file(file.file_path, destination=buffer) return buffer.getvalue(), Path(file.file_path).name async def transcribe_voice_message(bot: Bot, message: Message) -> str: voice_bytes, filename = await download_voice_message(bot, message) transcription = await client.audio.transcriptions.create( model=settings.transcription_model, file=(filename, voice_bytes, "audio/ogg"), ) transcript = transcription.text.strip() if not transcript: raise RuntimeError("Voice message transcription was empty") return transcript async def safe_answer(message: Message, text: str) -> None: try: await message.answer(text) except TelegramBadRequest: await message.answer(text[: settings.max_response_chars]) async def wait_between_reply_parts(bot: Bot, chat_id: int, timeout: int) -> None: if timeout <= 0: await bot.send_chat_action(chat_id=chat_id, action=ChatAction.TYPING) await asyncio.sleep(0.8) return stop_typing = asyncio.Event() typing_task = asyncio.create_task(keep_typing(bot, chat_id, stop_typing)) try: await asyncio.sleep(timeout) finally: stop_typing.set() await typing_task def split_reply(reply: str) -> list[tuple[str, int]]: parts: list[tuple[str, int]] = [] position = 0 for match in BREAK_RE.finditer(reply): text = reply[position : match.start()].strip() timeout = min(int(match.group("timeout") or 0), 30) if text: parts.append((text, timeout)) position = match.end() tail = reply[position:].strip() if tail: parts.append((tail, 0)) return parts or [(reply.strip(), 0)] async def send_reply(message: Message, bot: Bot, reply: str) -> None: parts = split_reply(reply or "The model returned an empty response.") for index, (text, timeout) in enumerate(parts): await safe_answer(message, text) if index < len(parts) - 1: await wait_between_reply_parts(bot, message.chat.id, timeout) async def process_pending_messages(message: Message, bot: Bot, user_id: int) -> None: try: await asyncio.sleep(settings.message_batch_delay_seconds) async with user_lock(user_id): messages = pending_user_messages.pop(user_id, []) pending_tasks.pop(user_id, None) if not messages: return user_text = combine_user_messages(messages) stop_typing = asyncio.Event() typing_task = asyncio.create_task(keep_typing(bot, message.chat.id, stop_typing)) should_store_reply = False try: context = await history.get_context(user_id) reply = await stream_model_reply(user_text, context) should_store_reply = bool(reply) except RateLimitError: logger.warning("API rate limit exceeded", exc_info=True) reply = "The model provider is rate-limiting requests. Please try again later." except (APITimeoutError, APIConnectionError): logger.warning("API connection problem", exc_info=True) reply = "The model provider is temporarily unreachable. Please try again later." except NotFoundError: logger.exception("API endpoint or model was not found") reply = "The model provider returned 404. Check API_BASE_URL and MODEL configuration." except APIError: logger.exception("API returned an error") reply = "The model provider returned an error. Please try again later." except Exception: logger.exception("Unexpected bot error") reply = "Unexpected error while processing the request. Please try again later." finally: stop_typing.set() await typing_task if should_store_reply: await history.append_turn(user_id, user_text, reply) await send_reply(message, bot, reply) except asyncio.CancelledError: raise async def queue_user_message(message: Message, bot: Bot, user_text: str) -> None: if not message.from_user: await message.answer("Cannot identify Telegram user.") return async with user_lock(message.from_user.id): pending_user_messages.setdefault(message.from_user.id, []).append(user_text) task = pending_tasks.get(message.from_user.id) if task: task.cancel() pending_tasks[message.from_user.id] = asyncio.create_task( process_pending_messages(message, bot, message.from_user.id) ) @router.message(CommandStart()) async def start(message: Message) -> None: if not is_allowed(message): await message.answer("Access denied.") return await message.answer("Send a text or voice message and I will forward it to the configured model.") @router.message(Command("new")) async def reset_context(message: Message) -> None: if not is_allowed(message): await message.answer("Access denied.") return if not message.from_user: await message.answer("Cannot identify Telegram user.") return async with user_lock(message.from_user.id): task = pending_tasks.pop(message.from_user.id, None) if task: task.cancel() pending_user_messages.pop(message.from_user.id, None) await history.reset(message.from_user.id) await message.answer("Context reset. The next message will start a new conversation.") @router.message(F.text) async def handle_text(message: Message, bot: Bot) -> None: if not is_allowed(message): await message.answer("Access denied.") return if not message.text: await message.answer("Only text messages are supported.") return await queue_user_message(message, bot, message.text) @router.message(F.voice) async def handle_voice(message: Message, bot: Bot) -> None: if not is_allowed(message): await message.answer("Access denied.") return if not message.voice: await message.answer("Only Telegram voice messages are supported.") return logger.info("Transcribing voice message for user_id=%s", message.from_user.id) try: transcript = await transcribe_voice_message(bot, message) except RateLimitError: logger.warning("Transcription API rate limit exceeded", exc_info=True) await message.answer("The transcription model provider is rate-limiting requests. Please try again later.") return except (APITimeoutError, APIConnectionError): logger.warning("Transcription API connection problem", exc_info=True) await message.answer("The transcription model provider is temporarily unreachable. Please try again later.") return except NotFoundError: logger.exception("Transcription endpoint or model was not found") await message.answer( "The transcription model provider returned 404. Check API_BASE_URL and TRANSCRIPTION_MODEL configuration." ) return except APIError: logger.exception("Transcription API returned an error") await message.answer("The transcription model provider returned an error. Please try again later.") return except Exception: logger.exception("Unexpected voice transcription error") await message.answer("Unexpected error while transcribing the voice message. Please try again later.") return transcript_message = voice_transcript_text(transcript) logger.info("Voice message transcribed for user_id=%s", message.from_user.id) # Telegram bots cannot mark a voice note as listened in the client UI. await queue_user_message(message, bot, transcript_message) async def main() -> None: logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(name)s: %(message)s") await history.init() session = AiohttpSession(proxy=settings.telegram_proxy) if settings.telegram_proxy else AiohttpSession() bot = Bot( token=settings.bot_token, session=session, default=DefaultBotProperties(parse_mode=ParseMode.HTML), ) dispatcher = Dispatcher() dispatcher.include_router(router) await dispatcher.start_polling(bot, allowed_updates=dispatcher.resolve_used_update_types()) if __name__ == "__main__": asyncio.run(main())