53 lines
1.8 KiB
Python
53 lines
1.8 KiB
Python
|
|
import asyncio
|
||
|
|
import logging
|
||
|
|
import time
|
||
|
|
|
||
|
|
from yuxi.channel.message.models import UnifiedMessage
|
||
|
|
|
||
|
|
logger = logging.getLogger(__name__)
|
||
|
|
|
||
|
|
_DEFAULT_FENCE_TTL = 300.0
|
||
|
|
|
||
|
|
|
||
|
|
class ConversationFence:
|
||
|
|
def __init__(self, ttl: float = _DEFAULT_FENCE_TTL):
|
||
|
|
self._ttl = ttl
|
||
|
|
self._versions: dict[str, int] = {}
|
||
|
|
self._version_ts: dict[str, float] = {}
|
||
|
|
self._locks: dict[str, asyncio.Lock] = {}
|
||
|
|
self._abort_events: dict[str, asyncio.Event] = {}
|
||
|
|
|
||
|
|
@staticmethod
|
||
|
|
def key_for(msg: UnifiedMessage) -> str:
|
||
|
|
if msg.group and msg.group.id:
|
||
|
|
return f"{msg.channel_type}:{msg.account_id}:group:{msg.group.id}"
|
||
|
|
return f"{msg.channel_type}:{msg.account_id}:dm:{msg.sender.id}"
|
||
|
|
|
||
|
|
def enter(self, key: str) -> tuple[int, asyncio.Event]:
|
||
|
|
now = time.monotonic()
|
||
|
|
self._gc(now)
|
||
|
|
self._versions[key] = self._versions.get(key, 0) + 1
|
||
|
|
self._version_ts[key] = now
|
||
|
|
|
||
|
|
old_event = self._abort_events.get(key)
|
||
|
|
if old_event is not None and not old_event.is_set():
|
||
|
|
old_event.set()
|
||
|
|
logger.debug("Foreground fence: aborting previous run for %s", key)
|
||
|
|
|
||
|
|
new_event = asyncio.Event()
|
||
|
|
self._abort_events[key] = new_event
|
||
|
|
return self._versions[key], new_event
|
||
|
|
|
||
|
|
def lock_for(self, key: str) -> asyncio.Lock:
|
||
|
|
return self._locks.setdefault(key, asyncio.Lock())
|
||
|
|
|
||
|
|
def current_version(self, key: str) -> int:
|
||
|
|
return self._versions.get(key, 0)
|
||
|
|
|
||
|
|
def _gc(self, now: float) -> None:
|
||
|
|
expired = [k for k, ts in self._version_ts.items() if now - ts > self._ttl]
|
||
|
|
for k in expired:
|
||
|
|
self._versions.pop(k, None)
|
||
|
|
self._version_ts.pop(k, None)
|
||
|
|
self._locks.pop(k, None)
|
||
|
|
self._abort_events.pop(k, None)
|