更新
This commit is contained in:
@@ -0,0 +1,352 @@
|
||||
"""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()
|
||||
Reference in New Issue
Block a user