74 lines
2.6 KiB
Python
74 lines
2.6 KiB
Python
|
|
import asyncio
|
||
|
|
from collections import defaultdict
|
||
|
|
|
||
|
|
from yuxi.channel.message.models import UnifiedMessage
|
||
|
|
|
||
|
|
|
||
|
|
class DebounceManager:
|
||
|
|
def __init__(self, window_ms: int = 2500, coalesce_dms: bool = False):
|
||
|
|
self.window_ms = window_ms
|
||
|
|
self.coalesce_dms = coalesce_dms
|
||
|
|
self._pending: dict[str, list[UnifiedMessage]] = defaultdict(list)
|
||
|
|
self._timers: dict[str, asyncio.Task] = {}
|
||
|
|
self._dm_buffers: dict[str, tuple[float, list[UnifiedMessage]]] = {}
|
||
|
|
self._lock = asyncio.Lock()
|
||
|
|
|
||
|
|
def _dm_key(self, msg: UnifiedMessage) -> str:
|
||
|
|
return f"dm:{msg.sender.id}"
|
||
|
|
|
||
|
|
def _group_key(self, msg: UnifiedMessage) -> str:
|
||
|
|
gid = msg.group.id if msg.group else "unknown"
|
||
|
|
return f"group:{gid}:{msg.sender.id}"
|
||
|
|
|
||
|
|
def _debounce_key(self, msg: UnifiedMessage) -> str:
|
||
|
|
if self.coalesce_dms and not msg.group:
|
||
|
|
return self._dm_key(msg)
|
||
|
|
return self._group_key(msg)
|
||
|
|
|
||
|
|
async def enqueue(self, msg: UnifiedMessage, on_dispatch: callable):
|
||
|
|
key = self._debounce_key(msg)
|
||
|
|
async with self._lock:
|
||
|
|
self._pending[key].append(msg)
|
||
|
|
if key in self._timers and not self._timers[key].done():
|
||
|
|
self._timers[key].cancel()
|
||
|
|
self._timers[key] = asyncio.create_task(self._flush_after_delay(key, on_dispatch))
|
||
|
|
|
||
|
|
async def _flush_after_delay(self, key: str, on_dispatch: callable):
|
||
|
|
await asyncio.sleep(self.window_ms / 1000.0)
|
||
|
|
await self.flush_key(key, on_dispatch)
|
||
|
|
|
||
|
|
async def flush_key(self, key: str, on_dispatch: callable):
|
||
|
|
async with self._lock:
|
||
|
|
messages = self._pending.pop(key, [])
|
||
|
|
if key in self._timers:
|
||
|
|
self._timers.pop(key, None)
|
||
|
|
|
||
|
|
if not messages:
|
||
|
|
return
|
||
|
|
|
||
|
|
if self.coalesce_dms and key.startswith("dm:"):
|
||
|
|
merged = self._coalesce_dm_messages(messages)
|
||
|
|
await on_dispatch(merged)
|
||
|
|
else:
|
||
|
|
last = messages[-1]
|
||
|
|
await on_dispatch(last)
|
||
|
|
|
||
|
|
def _coalesce_dm_messages(self, messages: list[UnifiedMessage]) -> UnifiedMessage:
|
||
|
|
if len(messages) <= 1:
|
||
|
|
return messages[0]
|
||
|
|
|
||
|
|
combined = " ".join(m.content for m in messages if m.content)
|
||
|
|
all_urls = []
|
||
|
|
for m in messages:
|
||
|
|
all_urls.extend(m.media_urls)
|
||
|
|
|
||
|
|
result = messages[-1]
|
||
|
|
result.content = combined
|
||
|
|
result.media_urls = all_urls
|
||
|
|
return result
|
||
|
|
|
||
|
|
async def flush_all(self, on_dispatch: callable):
|
||
|
|
keys = list(self._pending.keys())
|
||
|
|
for key in keys:
|
||
|
|
await self.flush_key(key, on_dispatch)
|