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 []