ForcePilot/backend/package/yuxi/channels/adapters/nostr/guard.py
Kris 1f78c44b03 refactor: 整理并清理项目中的冗余代码与格式问题
这是一个批量整理提交,包含以下主要改动:
1.  删除多处冗余的空行和未使用的导入
2.  修复文件末尾缺少换行符的问题
3.  调整部分模块的导入顺序与代码排版
4.  修复部分配置默认值与策略逻辑
5.  新增多个功能模块与辅助工具
6.  完善异常处理与日志记录
7.  修复速率限制、消息缓存、权限校验等逻辑bug
8.  废弃部分旧有API与配置项并添加警告提示
2026-05-12 14:51:53 +08:00

228 lines
8.2 KiB
Python

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, since_ts: int = 0) -> 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 in ("whitelist", "allowlist") 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 since_ts > 0 and created_at < since_ts:
self._record_audit(f"stale event: created_at={created_at}, since={since_ts}", raw_event)
return f"stale event: created_at={created_at}, since={since_ts}"
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)