from __future__ import annotations import asyncio import logging import time from collections.abc import Awaitable, Callable from dataclasses import dataclass, field logger = logging.getLogger(__name__) @dataclass class RetryTask: task_id: str chat_id: str payload: dict attempt: int = 0 max_attempts: int = 5 created_at: float = field(default_factory=time.time) next_retry_at: float = field(default_factory=time.time) last_error: str = "" base_delay: float = 1.0 max_delay: float = 120.0 class MessageRetryQueue: def __init__(self, max_concurrent: int = 3, poll_interval: float = 1.0): self._queue: list[RetryTask] = [] self._dead_letter: list[RetryTask] = [] self._lock = asyncio.Lock() self._max_concurrent = max_concurrent self._poll_interval = poll_interval self._running = False self._worker_task: asyncio.Task | None = None self._send_cb: Callable[[str, dict], Awaitable[bool]] | None = None self._active_tasks: set[str] = set() @property def pending_count(self) -> int: return len(self._queue) @property def dead_count(self) -> int: return len(self._dead_letter) def set_send_callback(self, callback: Callable[[str, dict], Awaitable[bool]]) -> None: self._send_cb = callback async def enqueue(self, task: RetryTask) -> None: async with self._lock: if task.task_id in self._active_tasks: return self._queue.append(task) self._queue.sort(key=lambda t: t.next_retry_at) async def start(self) -> None: if self._running: return self._running = True self._worker_task = asyncio.create_task(self._worker_loop()) logger.info("MessageRetryQueue: worker started") async def stop(self) -> None: self._running = False if self._worker_task: self._worker_task.cancel() try: await self._worker_task except asyncio.CancelledError: pass self._worker_task = None logger.info("MessageRetryQueue: worker stopped, pending=%d dead=%d", len(self._queue), len(self._dead_letter)) async def _worker_loop(self) -> None: while self._running: task = None async with self._lock: for t in self._queue: if t.task_id in self._active_tasks: continue if time.time() >= t.next_retry_at: task = t self._active_tasks.add(t.task_id) break if task is None: await asyncio.sleep(self._poll_interval) continue if len(self._active_tasks) >= self._max_concurrent: async with self._lock: self._active_tasks.discard(task.task_id) await asyncio.sleep(self._poll_interval) continue try: success = await self._process_task(task) async with self._lock: self._active_tasks.discard(task.task_id) if success: self._queue.remove(task) elif task.attempt >= task.max_attempts: self._queue.remove(task) self._dead_letter.append(task) logger.warning( "MessageRetryQueue: task %s exhausted retries (chat=%s)", task.task_id, task.chat_id, ) except Exception: async with self._lock: self._active_tasks.discard(task.task_id) await asyncio.sleep(self._poll_interval) async def _process_task(self, task: RetryTask) -> bool: if self._send_cb is None: return False try: task.attempt += 1 success = await self._send_cb(task.chat_id, task.payload) if success: logger.debug("MessageRetryQueue: task %s succeeded on attempt %d", task.task_id, task.attempt) return True task.last_error = "send_failed" except Exception as e: task.last_error = str(e) delay = min(task.base_delay * (2 ** (task.attempt - 1)), task.max_delay) task.next_retry_at = time.time() + delay logger.info( "MessageRetryQueue: task %s retry %d/%d, next in %.1fs", task.task_id, task.attempt, task.max_attempts, delay, ) return False def get_dead_letter_tasks(self) -> list[RetryTask]: return list(self._dead_letter) def clear_dead_letter(self) -> int: count = len(self._dead_letter) self._dead_letter.clear() return count