70 lines
2.1 KiB
Python
70 lines
2.1 KiB
Python
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
import time
|
||
|
|
from collections import defaultdict
|
||
|
|
from dataclasses import dataclass, field
|
||
|
|
|
||
|
|
|
||
|
|
@dataclass
|
||
|
|
class GroupMessage:
|
||
|
|
msg_id: str
|
||
|
|
author_id: str
|
||
|
|
author_name: str
|
||
|
|
content: str
|
||
|
|
timestamp: float
|
||
|
|
mentions_bot: bool = False
|
||
|
|
|
||
|
|
|
||
|
|
@dataclass
|
||
|
|
class GroupSession:
|
||
|
|
group_id: str
|
||
|
|
messages: list[GroupMessage] = field(default_factory=list)
|
||
|
|
last_active: float = 0.0
|
||
|
|
buffer_limit: int = 50
|
||
|
|
ttl_seconds: float = 3600.0
|
||
|
|
|
||
|
|
def add(self, msg: GroupMessage) -> None:
|
||
|
|
self.messages.append(msg)
|
||
|
|
self.last_active = time.time()
|
||
|
|
if len(self.messages) > self.buffer_limit:
|
||
|
|
self.messages = self.messages[-self.buffer_limit :]
|
||
|
|
|
||
|
|
def is_expired(self, now: float | None = None) -> bool:
|
||
|
|
if now is None:
|
||
|
|
now = time.time()
|
||
|
|
return now - self.last_active > self.ttl_seconds
|
||
|
|
|
||
|
|
def recent_context(self, count: int = 10) -> list[GroupMessage]:
|
||
|
|
return self.messages[-count:]
|
||
|
|
|
||
|
|
|
||
|
|
class GroupHistoryBuffer:
|
||
|
|
def __init__(self, buffer_limit: int = 50, ttl_seconds: float = 3600.0) -> None:
|
||
|
|
self._sessions: dict[str, GroupSession] = defaultdict(GroupSession)
|
||
|
|
self._buffer_limit = buffer_limit
|
||
|
|
self._ttl_seconds = ttl_seconds
|
||
|
|
|
||
|
|
def record(self, group_id: str, msg: GroupMessage) -> None:
|
||
|
|
session = self._sessions[group_id]
|
||
|
|
session.buffer_limit = self._buffer_limit
|
||
|
|
session.ttl_seconds = self._ttl_seconds
|
||
|
|
if not session.group_id:
|
||
|
|
session.group_id = group_id
|
||
|
|
session.add(msg)
|
||
|
|
|
||
|
|
def recent_context(self, group_id: str, count: int = 10) -> list[GroupMessage]:
|
||
|
|
session = self._sessions.get(group_id)
|
||
|
|
if session is None:
|
||
|
|
return []
|
||
|
|
if session.is_expired():
|
||
|
|
del self._sessions[group_id]
|
||
|
|
return []
|
||
|
|
return session.recent_context(count)
|
||
|
|
|
||
|
|
def gc(self) -> int:
|
||
|
|
now = time.time()
|
||
|
|
expired = [gid for gid, s in self._sessions.items() if s.is_expired(now)]
|
||
|
|
for gid in expired:
|
||
|
|
del self._sessions[gid]
|
||
|
|
return len(expired)
|