This commit is contained in:
2026-07-15 14:11:59 +07:00
commit 3d171206ca
10 changed files with 700 additions and 0 deletions

432
bot.py Normal file
View File

@@ -0,0 +1,432 @@
import asyncio
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 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
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"),
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()
client = AsyncOpenAI(
api_key=settings.api_token,
base_url=settings.api_base_url,
timeout=settings.api_timeout_seconds,
max_retries=settings.api_max_retries,
)
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))
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 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
@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 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
async with user_lock(message.from_user.id):
pending_user_messages.setdefault(message.from_user.id, []).append(message.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)
)
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())