173 lines
5.6 KiB
Python
173 lines
5.6 KiB
Python
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
import asyncio
|
||
|
|
import logging
|
||
|
|
from collections import defaultdict
|
||
|
|
from typing import Any
|
||
|
|
|
||
|
|
from yuxi.channel.extensions.qqbot.types import QQBotChatType, QQBotEventType, QueuedMessage
|
||
|
|
|
||
|
|
logger = logging.getLogger(__name__)
|
||
|
|
|
||
|
|
|
||
|
|
class QQBotMessageQueue:
|
||
|
|
GLOBAL_QUEUE_SIZE = 1000
|
||
|
|
PER_PEER_QUEUE_SIZE = 20
|
||
|
|
GROUP_QUEUE_SIZE = 50
|
||
|
|
MAX_CONCURRENT_USERS = 10
|
||
|
|
|
||
|
|
def __init__(self):
|
||
|
|
self._global_queue: asyncio.Queue[QueuedMessage] = asyncio.Queue(maxsize=self.GLOBAL_QUEUE_SIZE)
|
||
|
|
self._peer_queues: dict[str, list[QueuedMessage]] = defaultdict(list)
|
||
|
|
self._active_peers: set[str] = set()
|
||
|
|
self._semaphore = asyncio.Semaphore(self.MAX_CONCURRENT_USERS)
|
||
|
|
self._consumer_task: asyncio.Task | None = None
|
||
|
|
self._handler: Any | None = None
|
||
|
|
self._cancel_event = asyncio.Event()
|
||
|
|
|
||
|
|
def set_handler(self, handler: Any) -> None:
|
||
|
|
self._handler = handler
|
||
|
|
|
||
|
|
async def enqueue(self, msg: QueuedMessage) -> None:
|
||
|
|
peer_id = self._get_peer_id(msg)
|
||
|
|
peer_queue = self._peer_queues[peer_id]
|
||
|
|
|
||
|
|
max_size = self.GROUP_QUEUE_SIZE if msg.chat_type == QQBotChatType.GROUP else self.PER_PEER_QUEUE_SIZE
|
||
|
|
|
||
|
|
if len(peer_queue) >= max_size:
|
||
|
|
if msg.chat_type == QQBotChatType.GROUP:
|
||
|
|
bot_msgs = [i for i, m in enumerate(peer_queue) if m.sender_is_bot]
|
||
|
|
if bot_msgs:
|
||
|
|
peer_queue.pop(bot_msgs[0])
|
||
|
|
else:
|
||
|
|
logger.warning("Group queue full for peer=%s, dropping", peer_id)
|
||
|
|
return
|
||
|
|
else:
|
||
|
|
peer_queue.pop(0)
|
||
|
|
|
||
|
|
peer_queue.append(msg)
|
||
|
|
|
||
|
|
async def start_consumer(self) -> None:
|
||
|
|
self._cancel_event.clear()
|
||
|
|
self._consumer_task = asyncio.create_task(self._consume_loop(), name="qqbot-msg-queue-consumer")
|
||
|
|
|
||
|
|
async def stop_consumer(self) -> None:
|
||
|
|
self._cancel_event.set()
|
||
|
|
if self._consumer_task and not self._consumer_task.done():
|
||
|
|
self._consumer_task.cancel()
|
||
|
|
try:
|
||
|
|
await self._consumer_task
|
||
|
|
except asyncio.CancelledError:
|
||
|
|
pass
|
||
|
|
self._consumer_task = None
|
||
|
|
|
||
|
|
async def _consume_loop(self) -> None:
|
||
|
|
while not self._cancel_event.is_set():
|
||
|
|
try:
|
||
|
|
await self._drain_next_peer()
|
||
|
|
except Exception:
|
||
|
|
logger.exception("Message queue consumer error")
|
||
|
|
|
||
|
|
async def _drain_next_peer(self) -> None:
|
||
|
|
ready_peers = [p for p in self._peer_queues if p not in self._active_peers and self._peer_queues[p]]
|
||
|
|
|
||
|
|
if not ready_peers:
|
||
|
|
await asyncio.sleep(0.1)
|
||
|
|
return
|
||
|
|
|
||
|
|
for peer_id in ready_peers:
|
||
|
|
if peer_id in self._active_peers:
|
||
|
|
continue
|
||
|
|
|
||
|
|
async with self._semaphore:
|
||
|
|
self._active_peers.add(peer_id)
|
||
|
|
try:
|
||
|
|
await self._drain_peer(peer_id)
|
||
|
|
finally:
|
||
|
|
self._active_peers.discard(peer_id)
|
||
|
|
return
|
||
|
|
|
||
|
|
await asyncio.sleep(0.1)
|
||
|
|
|
||
|
|
async def _drain_peer(self, peer_id: str) -> None:
|
||
|
|
queue = self._peer_queues.get(peer_id, [])
|
||
|
|
if not queue:
|
||
|
|
return
|
||
|
|
|
||
|
|
messages: list[QueuedMessage] = []
|
||
|
|
while queue:
|
||
|
|
msg = queue.pop(0)
|
||
|
|
messages.append(msg)
|
||
|
|
|
||
|
|
if not messages:
|
||
|
|
return
|
||
|
|
|
||
|
|
if len(messages) > 1 and any(m.chat_type == QQBotChatType.GROUP for m in messages):
|
||
|
|
messages = self._merge_group_messages(messages)
|
||
|
|
|
||
|
|
for msg in messages:
|
||
|
|
if self._handler:
|
||
|
|
try:
|
||
|
|
await self._handler(msg)
|
||
|
|
except Exception:
|
||
|
|
logger.exception("Message handler error for msg_id=%s", msg.msg_id)
|
||
|
|
|
||
|
|
def _merge_group_messages(self, messages: list[QueuedMessage]) -> list[QueuedMessage]:
|
||
|
|
control_msgs = []
|
||
|
|
normal_msgs = []
|
||
|
|
bot_msgs = []
|
||
|
|
|
||
|
|
for msg in messages:
|
||
|
|
if msg.content.startswith("/"):
|
||
|
|
control_msgs.append(msg)
|
||
|
|
elif msg.sender_is_bot:
|
||
|
|
bot_msgs.append(msg)
|
||
|
|
else:
|
||
|
|
normal_msgs.append(msg)
|
||
|
|
|
||
|
|
result = control_msgs.copy()
|
||
|
|
|
||
|
|
if normal_msgs:
|
||
|
|
merged = self._merge_messages(normal_msgs)
|
||
|
|
result.append(merged)
|
||
|
|
|
||
|
|
if bot_msgs:
|
||
|
|
bot_merged = self._merge_messages(bot_msgs)
|
||
|
|
result.append(bot_merged)
|
||
|
|
|
||
|
|
return result
|
||
|
|
|
||
|
|
def _merge_messages(self, messages: list[QueuedMessage]) -> QueuedMessage:
|
||
|
|
if len(messages) == 1:
|
||
|
|
return messages[0]
|
||
|
|
|
||
|
|
base = messages[-1]
|
||
|
|
contents = []
|
||
|
|
all_mentions: list[str] = []
|
||
|
|
all_attachments = []
|
||
|
|
|
||
|
|
for m in messages:
|
||
|
|
name = m.sender_name or m.sender_id
|
||
|
|
contents.append(f"[{name}]: {m.content}")
|
||
|
|
all_mentions.extend(m.mentions)
|
||
|
|
all_attachments.extend(m.attachments)
|
||
|
|
|
||
|
|
base.content = "\n".join(contents)
|
||
|
|
base.mentions = list(dict.fromkeys(all_mentions))
|
||
|
|
base.attachments = all_attachments
|
||
|
|
base.sender_is_bot = all(m.sender_is_bot for m in messages)
|
||
|
|
|
||
|
|
is_at = any(m.event_type == QQBotEventType.GROUP_AT_MESSAGE_CREATE for m in messages)
|
||
|
|
if is_at:
|
||
|
|
base.event_type = QQBotEventType.GROUP_AT_MESSAGE_CREATE
|
||
|
|
|
||
|
|
base.merge = {"count": len(messages)}
|
||
|
|
|
||
|
|
return base
|
||
|
|
|
||
|
|
def _get_peer_id(self, msg: QueuedMessage) -> str:
|
||
|
|
if msg.chat_type == QQBotChatType.GUILD:
|
||
|
|
return f"guild:{msg.channel_id}"
|
||
|
|
if msg.chat_type == QQBotChatType.GROUP:
|
||
|
|
return f"group:{msg.group_openid}"
|
||
|
|
return f"dm:{msg.sender_id}"
|