520 lines
18 KiB
Python
520 lines
18 KiB
Python
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"<break(?:\s+timeout=(?P<timeout>\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())
|