import asyncio import hashlib from cachetools import TTLCache from yuxi.channels.models import ChannelMessage from yuxi.utils.logging_config import logger class DedupPolicy: def __init__(self, ttl: int = 300, maxsize: int = 10000, redis_url: str | None = None): self._seen: TTLCache = TTLCache(maxsize=maxsize, ttl=ttl) self._lock = asyncio.Lock() self._redis_url = redis_url self._redis = None async def _ensure_redis(self): if self._redis is not None: return self._redis if not self._redis_url: return None try: import redis.asyncio as aioredis self._redis = aioredis.from_url(self._redis_url, decode_responses=False) await self._redis.ping() logger.info(f"DedupPolicy: Redis connected ({self._redis_url})") return self._redis except Exception as e: logger.warning(f"DedupPolicy: Redis unavailable, falling back to memory: {e}") if self._redis is not None: try: await self._redis.aclose() except Exception: pass self._redis = None return None async def check_and_mark(self, key: str, ttl: int | None = None) -> bool: effective_ttl = ttl if ttl is not None else self._seen.ttl redis = await self._ensure_redis() if redis is not None: try: acquired = await redis.set(key, "1", nx=True, ex=effective_ttl) if acquired is None: return True async with self._lock: self._seen[key] = True return False except Exception as e: logger.debug(f"DedupPolicy: Redis error, falling back to memory: {e}") async with self._lock: if key in self._seen: return True self._seen[key] = True return False async def is_duplicate(self, message: ChannelMessage) -> bool: msg_id = message.identity.channel_message_id if not msg_id: return False key = f"{message.identity.channel_id}:{msg_id}" return await self.check_and_mark(key) async def check_and_remember(self, message: ChannelMessage) -> bool: async with self._lock: if self._check_identity_key(message): return True if self._check_fingerprint(message): return True self._record(message) return False def _check_identity_key(self, message: ChannelMessage) -> bool: msg_id = message.identity.channel_message_id if not msg_id: return False key = f"{message.identity.channel_id}:{msg_id}" if key in self._seen: logger.debug(f"Duplicate message filtered (identity): {key}") return True return False def _check_fingerprint(self, message: ChannelMessage) -> bool: fp = self._build_fingerprint(message) if not fp: return False key = f"fp:{message.identity.channel_id}:{fp}" if key in self._seen: logger.debug(f"Duplicate message filtered (fingerprint): {key}") return True return False def _record(self, message: ChannelMessage) -> None: msg_id = message.identity.channel_message_id if msg_id: self._seen[f"{message.identity.channel_id}:{msg_id}"] = True fp = self._build_fingerprint(message) if fp: self._seen[f"fp:{message.identity.channel_id}:{fp}"] = True @staticmethod def _build_fingerprint(message: ChannelMessage) -> str | None: if not message.content: return None raw = "|".join( [ message.identity.channel_id, message.identity.channel_chat_id, message.content, ] ) return hashlib.sha256(raw.encode()).hexdigest()[:32] async def clear(self) -> None: async with self._lock: self._seen.clear() async def close(self) -> None: if self._redis is not None: try: await self._redis.aclose() except Exception: pass self._redis = None