"""Outbox: a reply is a ``deliveries`` row first, a message second. Rows are sent oldest first, retried with backoff on network errors, resent as plain text when Telegram rejects our HTML, and given up once it's gone. """ from __future__ import annotations import asyncio import contextlib import logging from datetime import UTC, datetime, timedelta from typing import TYPE_CHECKING from aiogram.exceptions import ( TelegramBadRequest, TelegramForbiddenError, TelegramNetworkError, TelegramNotFound, TelegramRetryAfter, TelegramServerError, ) from aiogram.types import LinkPreviewOptions from sqlalchemy import func from sqlalchemy.exc import IntegrityError from sqlmodel import col, select from beaver_gateway.frontends.telegram.render import chunks, to_html from beaver_gateway.storage.models import Delivery if TYPE_CHECKING: from aiogram import Bot from beaver_gateway.events.bus import EventBus from beaver_gateway.storage.db import Database __all__ = ["Outbox"] _log = logging.getLogger("beaver_gateway.frontends.telegram.outbox") _MAX_BACKOFF = 300.0 _GONE = ("thread not found", "chat not found", "topic_deleted", "topic_closed") _NO_PREVIEW = LinkPreviewOptions(is_disabled=True) class Outbox: def __init__( self, db: Database, bot: Bot, *, bus: EventBus, backoff: float = 2.0 ) -> None: self._db = db self._bot = bot self._bus = bus self._backoff = backoff self._wake = asyncio.Event() async def enqueue( self, *, chat_id: int, thread_id: int | None, text: str, conversation_id: int | None = None, turn_id: str | None = None, dedupe_key: str | None = None, ) -> list[Delivery]: rows: list[Delivery] = [] for n, part in enumerate(chunks(text)): row = Delivery( conversation_id=conversation_id, chat_id=chat_id, thread_id=thread_id, text=part, turn_id=turn_id, dedupe_key=f"{dedupe_key}:{n}" if dedupe_key else None, ) async with self._db.session() as session: session.add(row) try: await session.commit() except IntegrityError: await session.rollback() continue await session.refresh(row) rows.append(row) if rows: self._wake.set() return rows async def run(self) -> None: while True: rows = await self._due() if not rows: self._wake.clear() with contextlib.suppress(TimeoutError): await asyncio.wait_for( self._wake.wait(), timeout=await self._wait_for_next() ) continue for row in rows: await self._send(row) async def _wait_for_next(self, cap: float = 5.0) -> float: async with self._db.session() as session: earliest = ( await session.exec( select(func.min(col(Delivery.next_attempt_at))).where( Delivery.status == "queued" ) ) ).one() if earliest is None: return cap now = datetime.now(UTC).replace(tzinfo=None) return max(0.05, min(cap, (earliest - now).total_seconds())) async def _due(self, limit: int = 50) -> list[Delivery]: now = datetime.now(UTC).replace(tzinfo=None) async with self._db.session() as session: result = await session.exec( select(Delivery) .where( Delivery.status == "queued", col(Delivery.next_attempt_at) <= now ) .order_by(col(Delivery.id)) .limit(limit) ) return list(result.all()) async def _send(self, row: Delivery) -> None: try: message = await self._bot.send_message( row.chat_id, row.text if row.plain else to_html(row.text), message_thread_id=row.thread_id, parse_mode=None if row.plain else "HTML", link_preview_options=_NO_PREVIEW, ) except TelegramRetryAfter as exc: await self._retry(row, str(exc), delay=float(exc.retry_after)) except TelegramBadRequest as exc: text = str(exc).lower() if "parse" in text and not row.plain: await self._retry(row, str(exc), delay=0.0, plain=True) elif any(marker in text for marker in _GONE): await self._fail(row, str(exc), gone=True) else: await self._fail(row, str(exc)) except (TelegramNotFound, TelegramForbiddenError) as exc: await self._fail(row, str(exc)) except (TelegramNetworkError, TelegramServerError, OSError) as exc: await self._retry( row, str(exc), delay=min(self._backoff ** (row.attempts + 1), _MAX_BACKOFF), ) else: await self._mark(row, status="sent", message_id=message.message_id) self._bus.publish( "delivery.sent", delivery=row.id, conversation_row=row.conversation_id, chat_id=row.chat_id, thread_id=row.thread_id, message_id=message.message_id, turn_id=row.turn_id, ) async def _retry( self, row: Delivery, error: str, *, delay: float, plain: bool = False ) -> None: _log.warning( "delivery #%s attempt %d failed: %s (retry in %.0fs)", row.id, row.attempts + 1, error, delay, ) await self._mark(row, status="queued", error=error, delay=delay, plain=plain) if delay < 5.0: self._wake.set() async def _fail(self, row: Delivery, error: str, *, gone: bool = False) -> None: _log.error( "delivery #%s to %s/%s given up: %s", row.id, row.chat_id, row.thread_id, error, ) await self._mark(row, status="failed", error=error) self._bus.publish( "delivery.failed", delivery=row.id, conversation_row=row.conversation_id, chat_id=row.chat_id, thread_id=row.thread_id, error=error, gone=gone, ) async def _mark( self, row: Delivery, *, status: str, error: str | None = None, delay: float = 0.0, plain: bool = False, message_id: int | None = None, ) -> None: async with self._db.session() as session: stored = await session.get(Delivery, row.id) if stored is None: return stored.status = status stored.attempts += 1 stored.last_error = error[:500] if error else None stored.next_attempt_at = ( datetime.now(UTC) + timedelta(seconds=delay) ).replace(tzinfo=None) if plain: stored.plain = True if message_id is not None: stored.message_id = message_id if status == "sent": stored.sent_at = datetime.now(UTC) session.add(stored) await session.commit() async def pending(self) -> int: async with self._db.session() as session: result = await session.exec( select(Delivery).where(Delivery.status == "queued") ) return len(list(result.all()))