Files
beaver-gateway/src/beaver_gateway/frontends/telegram/outbox.py
T

236 lines
7.7 KiB
Python

"""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()))