feat(telegram,core,backends,storage): telegram frontend with inbox, outbox, drafts and question buttons
This commit is contained in:
@@ -0,0 +1,235 @@
|
||||
"""Outbox (§3.8): a reply is a ``deliveries`` row first, a message second.
|
||||
|
||||
Rows are sent oldest first, retried with backoff on network errors and
|
||||
flood limits, resent as plain text when Telegram rejects our HTML, and
|
||||
given up only when Telegram says the window is 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.core.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))
|
||||
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) -> 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,
|
||||
)
|
||||
|
||||
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()))
|
||||
Reference in New Issue
Block a user