ForcePilot/backend/package/yuxi/channels/adapters/mattermost/cache.py
Kris 002d601d1b feat(mattermost): 实现完整的 Mattermost 适配器模块
新增 Mattermost 渠道完整实现,包含适配器核心、消息处理、交互回调、命令支持、安全校验、多账号管理等功能,支持机器人消息发送、交互按钮、命令注册、投票功能以及配置动态修改等特性。
2026-05-12 00:46:12 +08:00

190 lines
5.8 KiB
Python

from __future__ import annotations
import time
from collections import OrderedDict
from typing import Any
SENT_CACHE_MAX = 500
SENT_CACHE_TTL_S = 600
BOT_CACHE_TTL_S = 600
REACTION_BOT_CACHE_TTL_S = 600
USER_CACHE_TTL_S = 300
CHANNEL_CACHE_TTL_S = 300
DM_CHANNEL_CACHE_TTL_S = 300
DM_CHANNEL_CACHE_MAX = 200
class TTLCache:
"""通用 TTL 缓存基类。"""
def __init__(self, ttl_s: int = 300, max_size: int = 1000):
self._cache: dict[str, tuple[float, object]] = {}
self._ttl_s = ttl_s
self._max_size = max_size
def get(self, key: str) -> object | None:
entry = self._cache.get(key)
if entry is None:
return None
ts, value = entry
if time.monotonic() - ts > self._ttl_s:
del self._cache[key]
return None
return value
def set(self, key: str, value: object) -> None:
now = time.monotonic()
self._cache[key] = (now, value)
self._evict(now)
def delete(self, key: str) -> bool:
return self._cache.pop(key, None) is not None
def _evict(self, now: float) -> None:
cutoff = now - self._ttl_s
expired = [k for k, (ts, _) in self._cache.items() if ts < cutoff]
for k in expired:
del self._cache[k]
if len(self._cache) > self._max_size:
sorted_keys = sorted(self._cache, key=lambda k: self._cache[k][0])
for k in sorted_keys[: len(self._cache) - self._max_size]:
del self._cache[k]
def clear(self) -> None:
self._cache.clear()
class LRUCache:
"""LRU 缓存,适合 DM channel ID 缓存。"""
def __init__(self, max_size: int = DM_CHANNEL_CACHE_MAX, ttl_s: int = DM_CHANNEL_CACHE_TTL_S):
self._cache: OrderedDict[str, tuple[float, Any]] = OrderedDict()
self._max_size = max_size
self._ttl_s = ttl_s
def get(self, key: str) -> Any | None:
entry = self._cache.get(key)
if entry is None:
return None
ts, value = entry
if time.monotonic() - ts > self._ttl_s:
del self._cache[key]
return None
self._cache.move_to_end(key)
return value
def set(self, key: str, value: Any) -> None:
now = time.monotonic()
if key in self._cache:
self._cache.move_to_end(key)
self._cache[key] = (now, value)
self._evict_expired(now)
while len(self._cache) > self._max_size:
self._cache.popitem(last=False)
def _evict_expired(self, now: float) -> None:
cutoff = now - self._ttl_s
expired = [k for k, (ts, _) in self._cache.items() if ts < cutoff]
for k in expired:
del self._cache[k]
def clear(self) -> None:
self._cache.clear()
def size(self) -> int:
return len(self._cache)
class MattermostChannelCache:
"""Mattermost 渠道缓存集合。
- botUserCache: bot 用户信息缓存 (TTL 10min)
- userByNameCache: 用户名 → 用户ID 缓存
- channelByNameCache: 频道名 → 频道ID 缓存
- dmChannelCache: DM 频道 ID 缓存 (LRU)
"""
def __init__(self):
self.bot_user: TTLCache = TTLCache(ttl_s=BOT_CACHE_TTL_S, max_size=1)
self.user_by_name: TTLCache = TTLCache(ttl_s=USER_CACHE_TTL_S, max_size=500)
self.channel_by_name: TTLCache = TTLCache(ttl_s=CHANNEL_CACHE_TTL_S, max_size=500)
self.dm_channel: LRUCache = LRUCache(max_size=DM_CHANNEL_CACHE_MAX, ttl_s=DM_CHANNEL_CACHE_TTL_S)
def clear(self) -> None:
self.bot_user.clear()
self.user_by_name.clear()
self.channel_by_name.clear()
self.dm_channel.clear()
def stats(self) -> dict[str, int]:
return {
"bot_user": 1 if self.bot_user.get("_cached") else 0,
"user_by_name": len(self.user_by_name._cache),
"channel_by_name": len(self.channel_by_name._cache),
"dm_channel": self.dm_channel.size(),
}
class SentMessageCache:
"""已发送消息缓存 — 支持自动追加 thread_ts。"""
def __init__(self, max_size: int = SENT_CACHE_MAX, ttl_s: int = SENT_CACHE_TTL_S):
self._cache: dict[str, dict] = {}
self._max_size = max_size
self._ttl_s = ttl_s
def record(
self,
msg_id: str,
chat_id: str,
channel_id: str = "",
thread_id: str = "",
user_id: str = "",
) -> None:
now = time.monotonic()
self._cache[msg_id] = {
"msg_id": msg_id,
"chat_id": chat_id,
"channel_id": channel_id,
"thread_id": thread_id,
"user_id": user_id,
"recorded_at": now,
}
self._evict_if_needed(now)
def get(self, msg_id: str) -> dict | None:
entry = self._cache.get(msg_id)
if entry is None:
return None
if time.monotonic() - entry["recorded_at"] > self._ttl_s:
del self._cache[msg_id]
return None
return entry
def get_thread_id(self, msg_id: str) -> str | None:
entry = self.get(msg_id)
if entry:
return entry.get("thread_id") or None
return None
def clear(self) -> None:
self._cache.clear()
def size(self) -> int:
return len(self._cache)
def _evict_if_needed(self, now: float) -> None:
cutoff = now - self._ttl_s
expired = [k for k, v in self._cache.items() if v["recorded_at"] < cutoff]
for k in expired:
del self._cache[k]
if len(self._cache) > self._max_size:
sorted_entries = sorted(
self._cache.items(),
key=lambda item: item[1]["recorded_at"],
)
excess = len(self._cache) - self._max_size
for k, _ in sorted_entries[:excess]:
del self._cache[k]