from __future__ import annotations import time from collections import OrderedDict, defaultdict from dataclasses import dataclass, field from yuxi.channels.adapters.nostr.crypto import normalize_pubkey from yuxi.utils.logging_config import logger import json as _json @dataclass class AuditRecord: timestamp: float reason: str pubkey: str kind: int event_id: str @dataclass class GuardPolicy: allowed_kinds: set[int] = field(default_factory=lambda: {1, 4, 5, 7, 1059}) max_ciphertext_bytes: int = 50_000 max_plaintext_bytes: int = 10_000 max_future_skew_sec: int = 30 rate_limit_window_ms: int = 10_000 rate_limit_max_per_sender_per_window: int = 20 rate_limit_max_global_per_window: int = 200 dm_policy: str = "pairing" allow_from: list[str] = field(default_factory=list) class SeenTracker: def __init__(self, max_size: int = 100_000, ttl_sec: int = 3600): self._max_size = max_size self._ttl_sec = ttl_sec self._store: OrderedDict[str, float] = OrderedDict() self._eviction_counter: int = 0 def is_seen(self, event_id: str) -> bool: if event_id in self._store: ts = self._store[event_id] if time.monotonic() - ts < self._ttl_sec: self._store.move_to_end(event_id) return True del self._store[event_id] return False def mark_seen(self, event_id: str) -> None: if event_id in self._store: self._store.move_to_end(event_id) self._store[event_id] = time.monotonic() return self._store[event_id] = time.monotonic() self._eviction_counter += 1 self._evict() def _evict(self) -> None: while len(self._store) > self._max_size: oldest_key, _ = self._store.popitem(last=False) if self._eviction_counter % 1000 == 0: now = time.monotonic() expired_keys = [k for k, v in self._store.items() if now - v >= self._ttl_sec] for k in expired_keys: del self._store[k] def __len__(self) -> int: return len(self._store) class RateLimiter: def __init__(self, window_ms: int = 10_000, max_per_sender: int = 20, max_global: int = 200): self._window_ms = window_ms self._max_per_sender = max_per_sender self._max_global = max_global self._sender_buckets: dict[str, list[float]] = defaultdict(list) self._global_bucket: list[float] = [] def check_and_record(self, sender_pubkey: str) -> bool: now = time.monotonic() * 1000 window_start = now - self._window_ms sender_times = self._sender_buckets[sender_pubkey] sender_times[:] = [t for t in sender_times if t > window_start] self._global_bucket[:] = [t for t in self._global_bucket if t > window_start] if len(sender_times) >= self._max_per_sender: logger.debug(f"Rate limit hit for sender: {sender_pubkey[:8]}") return False if len(self._global_bucket) >= self._max_global: logger.debug("Global rate limit hit") return False sender_times.append(now) self._global_bucket.append(now) return True def clear(self) -> None: self._sender_buckets.clear() self._global_bucket.clear() class NostrGuard: def __init__( self, own_pubkey: str, policy: GuardPolicy | None = None, ): self._own_pubkey = own_pubkey self.policy = policy or GuardPolicy() self._seen = SeenTracker() self._rate_limiter = RateLimiter( window_ms=self.policy.rate_limit_window_ms, max_per_sender=self.policy.rate_limit_max_per_sender_per_window, max_global=self.policy.rate_limit_max_global_per_window, ) self._inflight: set[str] = set() self._allow_from_norm: set[str] = {normalize_pubkey(p) for p in self.policy.allow_from if normalize_pubkey(p)} self._audit_log: list[AuditRecord] = [] def _record_audit(self, reason: str, raw_event: dict) -> None: record = AuditRecord( timestamp=time.time(), reason=reason, pubkey=raw_event.get("pubkey", ""), kind=raw_event.get("kind", 0), event_id=raw_event.get("id", ""), ) self._audit_log.append(record) if len(self._audit_log) > 1000: self._audit_log = self._audit_log[-500:] logger.warning( "[NostrAudit] event rejected: reason=%s, pubkey=%s, kind=%s, event_id=%s", reason, raw_event.get("pubkey", "?")[:12], raw_event.get("kind", 0), raw_event.get("id", "")[:12], ) def get_audit_log(self) -> list[dict]: return [_json.loads(_json.dumps(r.__dict__)) for r in self._audit_log] def clear_audit_log(self) -> None: self._audit_log.clear() def check(self, raw_event: dict) -> str | None: event_id = raw_event.get("id", "") if not event_id: self._record_audit("missing event id", raw_event) return "missing event id" if self._seen.is_seen(event_id): self._record_audit("duplicate event", raw_event) return "duplicate event" if event_id in self._inflight: self._record_audit("event in-flight (concurrent processing)", raw_event) return "event in-flight (concurrent processing)" if self.policy.dm_policy == "disabled": self._record_audit("dm policy disabled — inbound rejected", raw_event) return "dm policy disabled — inbound rejected" kind = raw_event.get("kind", 0) if kind not in self.policy.allowed_kinds: self._record_audit(f"disallowed kind: {kind}", raw_event) return f"disallowed kind: {kind}" pubkey = raw_event.get("pubkey", "") pubkey_norm = normalize_pubkey(pubkey) if pubkey == self._own_pubkey or pubkey_norm == self._own_pubkey: self._record_audit("self-message (echo)", raw_event) return "self-message (echo)" if self._allow_from_norm and pubkey_norm not in self._allow_from_norm: self._record_audit(f"pubkey not in allow_from: {pubkey_norm[:12]}...", raw_event) return f"pubkey not in allow_from: {pubkey_norm[:12]}..." if self.policy.dm_policy == "whitelist" and not self._allow_from_norm: self._record_audit("whitelist mode without allow_from entries", raw_event) return "whitelist mode without allow_from entries — allowFrom 白名单为空" created_at = raw_event.get("created_at", 0) now = int(time.time()) if created_at > now + self.policy.max_future_skew_sec: self._record_audit(f"future event: created_at={created_at}, now={now}", raw_event) return f"future event: created_at={created_at}, now={now}" content = raw_event.get("content", "") if len(content) > self.policy.max_ciphertext_bytes: self._record_audit(f"ciphertext too large: {len(content)} > {self.policy.max_ciphertext_bytes}", raw_event) return f"ciphertext too large: {len(content)} > {self.policy.max_ciphertext_bytes}" if not self._rate_limiter.check_and_record(pubkey): self._record_audit("rate limited", raw_event) return "rate limited" self._seen.mark_seen(event_id) self._inflight.add(event_id) return None def check_plaintext_size(self, plaintext: str) -> str | None: if len(plaintext) > self.policy.max_plaintext_bytes: return f"plaintext too large: {len(plaintext)} > {self.policy.max_plaintext_bytes}" return None def done_processing(self, event_id: str) -> None: self._inflight.discard(event_id) @property def own_pubkey(self) -> str: return self._own_pubkey @property def inflight(self) -> frozenset[str]: return frozenset(self._inflight)