"""Observable per-account serial queue for delayed automatic replies.""" from __future__ import annotations import asyncio import logging import time import uuid from collections import deque from copy import deepcopy from dataclasses import dataclass, field from datetime import datetime, timezone from typing import Any, Awaitable, Callable, Iterable, Optional logger = logging.getLogger("douyin_im.reply_queue") ReplyCallback = Callable[[], Awaitable[Any]] ErrorCallback = Callable[[str, BaseException], None] DetailsMerger = Callable[[dict[str, Any]], dict[str, Any]] @dataclass class _QueueItem: job_id: str due_at: float slot_seconds: float callback: ReplyCallback description: str queued_at: float merge_keys: frozenset[str] = field(default_factory=frozenset) details: dict[str, Any] = field(default_factory=dict) expedited: bool = False class AccountReplyQueue: """Run and expose delayed reply jobs for one hosted account. A single consumer is the only code path allowed to invoke callbacks. Jobs selected for immediate delivery are moved to an urgent FIFO, so they can never overlap an already active send. Removing a scheduled job also moves every job behind it forward by the removed job's reserved slot. """ def __init__( self, account_id: int, on_error: Optional[ErrorCallback] = None, ) -> None: self.account_id = account_id self._on_error = on_error self._waiting: list[_QueueItem] = [] self._urgent: deque[_QueueItem] = deque() self._active_item: Optional[_QueueItem] = None self._task: Optional[asyncio.Task] = None self._running = False self._state_lock = asyncio.Lock() self._wake = asyncio.Event() self._tail_due_at = 0.0 @property def pending_count(self) -> int: """Approximate active + urgent + waiting count for lightweight badges.""" return ( len(self._waiting) + len(self._urgent) + (1 if self._active_item is not None else 0) ) async def start(self) -> None: async with self._state_lock: if self._task and not self._task.done(): return self._running = True self._tail_due_at = 0.0 self._wake.clear() self._task = asyncio.create_task( self._run(), name=f"account-reply-queue-{self.account_id}", ) async def enqueue( self, delay_seconds: float, callback: ReplyCallback, description: str = "", details: Optional[dict[str, Any]] = None, merge_key: str = "", merge_keys: Optional[Iterable[str]] = None, ) -> int: """Append one reply job and return its current 1-based queue position.""" interval = max(0.0, float(delay_seconds or 0)) loop = asyncio.get_running_loop() async with self._state_lock: if not self._running or not self._task or self._task.done(): raise RuntimeError("reply queue is not running") due_at = max(loop.time(), self._tail_due_at) + interval self._tail_due_at = due_at self._waiting.append( _QueueItem( job_id=uuid.uuid4().hex, due_at=due_at, slot_seconds=interval, callback=callback, description=description, queued_at=time.time(), merge_keys=self._normalize_merge_keys( merge_keys if merge_keys is not None else merge_key ), details=deepcopy(details or {}), ) ) position = self.pending_count self._wake.set() return position @staticmethod def _normalize_merge_keys(value: str | Iterable[str]) -> frozenset[str]: values = [value] if isinstance(value, str) else list(value or []) return frozenset(str(item or "").strip() for item in values if str(item or "").strip()) @staticmethod def _merge_keys_match( existing: frozenset[str], incoming: frozenset[str], ) -> bool: existing_conversations = {key for key in existing if key.startswith("conv:")} incoming_conversations = {key for key in incoming if key.startswith("conv:")} if existing_conversations & incoming_conversations: return True # Two explicit, different conversation IDs must never merge just because # their partial source data happens to expose the same peer identifier. if existing_conversations and incoming_conversations: return False existing_peers = {key for key in existing if key.startswith("peer:")} incoming_peers = {key for key in incoming if key.startswith("peer:")} return bool(existing_peers & incoming_peers) async def merge_pending( self, merge_key: str | Iterable[str], details_merger: DetailsMerger, ) -> dict[str, Any]: """Merge details into one queued conversation without changing its slot. Only waiting and urgent jobs are mutable. Once the consumer marks a job active, its callback is sealed and a later message must follow the normal new-message path. """ normalized_keys = self._normalize_merge_keys(merge_key) if not normalized_keys: return {"status": "not_found"} async with self._state_lock: if not self._running or not self._task or self._task.done(): return {"status": "not_running"} active_offset = 1 if self._active_item is not None else 0 matches: list[tuple[_QueueItem, str, int]] = [] for index, candidate in enumerate(self._urgent): if self._merge_keys_match(candidate.merge_keys, normalized_keys): matches.append((candidate, "ready", active_offset + index + 1)) waiting_offset = active_offset + len(self._urgent) for index, candidate in enumerate(self._waiting): if self._merge_keys_match(candidate.merge_keys, normalized_keys): matches.append((candidate, "waiting", waiting_offset + index + 1)) if not matches: return {"status": "not_found"} incoming_conversations = { key for key in normalized_keys if key.startswith("conv:") } if not incoming_conversations: matched_conversations = { key for candidate, _, _ in matches for key in candidate.merge_keys if key.startswith("conv:") } if len(matched_conversations) > 1: return {"status": "not_found", "reason": "ambiguous_peer"} item, item_status, position = matches[0] merged_details = details_merger(deepcopy(item.details)) if not isinstance(merged_details, dict): raise TypeError("reply queue details merger must return a dict") item.details = deepcopy(merged_details) item.merge_keys = frozenset(item.merge_keys | normalized_keys) return { "status": "merged", "job_id": item.job_id, "position": position, "queue_status": item_status, "message_count": int(item.details.get("message_count") or 1), } async def snapshot(self) -> list[dict[str, Any]]: """Return a callback-free management snapshot ordered by execution.""" loop = asyncio.get_running_loop() now_mono = loop.time() now_epoch = time.time() async with self._state_lock: ordered: list[tuple[_QueueItem, str]] = [] if self._active_item is not None: ordered.append((self._active_item, "sending")) ordered.extend((item, "ready") for item in self._urgent) ordered.extend((item, "waiting") for item in self._waiting) result = [] for position, (item, status) in enumerate(ordered, start=1): remaining = 0.0 if status != "waiting" else max(0.0, item.due_at - now_mono) scheduled_epoch = now_epoch + max(0.0, item.due_at - now_mono) payload = { "job_id": item.job_id, "account_id": self.account_id, "position": position, "status": status, "expedited": bool(item.expedited), "description": item.description, "interval_seconds": int(round(item.slot_seconds)), "enqueued_at": datetime.fromtimestamp( item.queued_at, tz=timezone.utc ).isoformat(), "scheduled_at": datetime.fromtimestamp( scheduled_epoch, tz=timezone.utc ).isoformat(), "remaining_seconds": int(max(0, round(remaining))), } # Details are controlled by DouyinImService and never contain callbacks/session data. payload.update(deepcopy(item.details)) result.append(payload) return result async def send_now(self, job_id: str) -> dict[str, Any]: """Move one waiting job to the urgent FIFO and free its future slot.""" job_id = str(job_id or "").strip() async with self._state_lock: if not self._running or not self._task or self._task.done(): return {"status": "not_running", "job_id": job_id} if self._active_item and self._active_item.job_id == job_id: return {"status": "already_sending", "job_id": job_id} if any(item.job_id == job_id for item in self._urgent): return {"status": "already_requested", "job_id": job_id} selected_index = next( (index for index, item in enumerate(self._waiting) if item.job_id == job_id), None, ) if selected_index is None: return {"status": "not_found", "job_id": job_id} item = self._waiting.pop(selected_index) shift_seconds = max(0.0, item.slot_seconds) shifted_count = 0 for later in self._waiting[selected_index:]: later.due_at -= shift_seconds shifted_count += 1 item.due_at = asyncio.get_running_loop().time() item.expedited = True self._urgent.append(item) self._recalculate_tail_due_at() self._wake.set() return { "status": "accepted", "job_id": job_id, "shifted_count": shifted_count, } async def stop(self) -> None: """Cancel the active wait/send and discard all remaining jobs.""" async with self._state_lock: self._running = False self._wake.set() task = self._task self._task = None if task: task.cancel() try: await task except asyncio.CancelledError: pass async with self._state_lock: self._waiting.clear() self._urgent.clear() self._active_item = None self._tail_due_at = 0.0 self._wake.clear() def _recalculate_tail_due_at(self) -> None: scheduled = [item.due_at for item in self._waiting] self._tail_due_at = max(scheduled, default=0.0) async def _run(self) -> None: while True: item: Optional[_QueueItem] = None wait_seconds: Optional[float] = None async with self._state_lock: if not self._running: return if self._urgent: item = self._urgent.popleft() elif self._waiting: candidate = self._waiting[0] remaining = candidate.due_at - asyncio.get_running_loop().time() if remaining <= 0: item = self._waiting.pop(0) else: wait_seconds = remaining if item is not None: self._active_item = item self._wake.clear() if item is None: try: if wait_seconds is None: await self._wake.wait() else: await asyncio.wait_for(self._wake.wait(), timeout=wait_seconds) except asyncio.TimeoutError: pass continue try: await item.callback() except asyncio.CancelledError: raise except Exception as exc: logger.exception( "Account %s queued reply failed (%s)", self.account_id, item.description, ) if self._on_error: try: self._on_error(item.description, exc) except Exception: logger.debug("Reply queue error callback failed", exc_info=True) finally: async with self._state_lock: if self._active_item is item: self._active_item = None if not self._waiting: self._tail_due_at = 0.0 self._wake.set()