ForcePilot/backend/package/yuxi/channels/adapters/yuanbao/send_cache.py

196 lines
6.3 KiB
Python
Raw Normal View History

from __future__ import annotations
import json
import time
from dataclasses import dataclass, field
from typing import Any
@dataclass
class SentMessageEntry:
msg_id: str
chat_id: str
content_preview: str = ""
sent_at: float = field(default_factory=time.time)
status: str = "sent"
def to_dict(self) -> dict[str, Any]:
return {
"msg_id": self.msg_id,
"chat_id": self.chat_id,
"content_preview": self.content_preview,
"sent_at": self.sent_at,
"status": self.status,
}
@classmethod
def from_dict(cls, data: dict[str, Any]) -> SentMessageEntry:
return cls(
msg_id=data.get("msg_id", ""),
chat_id=data.get("chat_id", ""),
content_preview=data.get("content_preview", ""),
sent_at=data.get("sent_at", 0.0),
status=data.get("status", "sent"),
)
class SendMessageCache:
MAX_ENTRIES = 10000
def __init__(self, max_entries: int | None = None):
self._cache: dict[str, SentMessageEntry] = {}
self._max_entries = max_entries or self.MAX_ENTRIES
self._redis: Any = None
self._redis_prefix = "yuanbao:sendcache:"
self._redis_ttl = 86400
def set_redis_backend(self, redis_client: Any, prefix: str = "", ttl: int = 86400) -> None:
self._redis = redis_client
if prefix:
self._redis_prefix = prefix
self._redis_ttl = ttl
def add(self, msg_id: str, chat_id: str, content_preview: str = "") -> None:
if len(self._cache) >= self._max_entries:
oldest = min(self._cache.values(), key=lambda e: e.sent_at)
self._cache.pop(oldest.msg_id, None)
entry = SentMessageEntry(
msg_id=msg_id,
chat_id=chat_id,
content_preview=content_preview[:200],
)
self._cache[msg_id] = entry
self._redis_save(entry)
def get(self, msg_id: str) -> SentMessageEntry | None:
entry = self._cache.get(msg_id)
if entry:
return entry
return self._redis_load(msg_id)
def get_by_chat(self, chat_id: str) -> list[SentMessageEntry]:
from_redis = self._redis_get_by_chat(chat_id)
memory = [e for e in self._cache.values() if e.chat_id == chat_id]
seen = {e.msg_id for e in memory}
for e in from_redis:
if e.msg_id not in seen:
memory.append(e)
return memory
def update_status(self, msg_id: str, status: str) -> bool:
entry = self._cache.get(msg_id)
if entry:
entry.status = status
self._redis_save(entry)
return True
entry = self._redis_load(msg_id)
if entry:
entry.status = status
self._cache[msg_id] = entry
self._redis_save(entry)
return True
return False
def remove(self, msg_id: str) -> bool:
self._redis_delete(msg_id)
return self._cache.pop(msg_id, None) is not None
def clear(self) -> None:
self._cache.clear()
@property
def size(self) -> int:
return len(self._cache)
def has_sent_in_chat(self, chat_id: str, since: float | None = None) -> bool:
for entry in self._cache.values():
if entry.chat_id == chat_id:
if since is None or entry.sent_at >= since:
return True
redis_entries = self._redis_get_by_chat(chat_id)
for entry in redis_entries:
if since is None or entry.sent_at >= since:
return True
return False
def mark_deleted(self, msg_id: str) -> bool:
return self.update_status(msg_id, "deleted")
def mark_edited(self, msg_id: str, new_content_preview: str = "") -> bool:
entry = self._cache.get(msg_id)
if entry:
entry.status = "edited"
if new_content_preview:
entry.content_preview = new_content_preview[:200]
self._redis_save(entry)
return True
entry = self._redis_load(msg_id)
if entry:
entry.status = "edited"
if new_content_preview:
entry.content_preview = new_content_preview[:200]
self._cache[msg_id] = entry
self._redis_save(entry)
return True
return False
def get_by_status(self, status: str) -> list[SentMessageEntry]:
return [e for e in self._cache.values() if e.status == status]
@property
def stats(self) -> dict[str, int]:
stats = {}
for entry in self._cache.values():
stats[entry.status] = stats.get(entry.status, 0) + 1
return stats
def _redis_save(self, entry: SentMessageEntry) -> None:
if not self._redis:
return
try:
key = f"{self._redis_prefix}{entry.msg_id}"
data = json.dumps(entry.to_dict(), ensure_ascii=False)
self._redis.setex(key, self._redis_ttl, data)
idx_key = f"{self._redis_prefix}chat:{entry.chat_id}"
self._redis.sadd(idx_key, entry.msg_id)
self._redis.expire(idx_key, self._redis_ttl)
except Exception:
pass
def _redis_load(self, msg_id: str) -> SentMessageEntry | None:
if not self._redis:
return None
try:
key = f"{self._redis_prefix}{msg_id}"
data = self._redis.get(key)
if data:
return SentMessageEntry.from_dict(json.loads(data))
except Exception:
pass
return None
def _redis_delete(self, msg_id: str) -> None:
if not self._redis:
return
try:
key = f"{self._redis_prefix}{msg_id}"
self._redis.delete(key)
except Exception:
pass
def _redis_get_by_chat(self, chat_id: str) -> list[SentMessageEntry]:
if not self._redis:
return []
try:
idx_key = f"{self._redis_prefix}chat:{chat_id}"
msg_ids = self._redis.smembers(idx_key) or []
entries = []
for msg_id in msg_ids:
entry = self._redis_load(msg_id.decode() if isinstance(msg_id, bytes) else msg_id)
if entry:
entries.append(entry)
return entries
except Exception:
return []