新增 Nostr 渠道扩展,支持在 Yuxi 平台中集成 Nostr 去中心化社交协议。 包含以下功能模块: - bus: 事件总线与中继通信 - account: 账户管理 - key_utils: 密钥工具 - gateway: SSE/WebSocket 网关接入 - outbound: 外发消息管理 - gift_wrap: Gift Wrap 加密 - nip44: NIP-44 加密协议 - config_schema: 配置模式 - defaults: 默认配置 - profile_core: 用户资料核心 - profile_publisher: 资料发布 - event_utils: 事件工具 - deletion: 事件删除 - reactions: 表情反应 - metrics: 指标监控 - seen_tracker: 已读追踪 - session_route: 会话路由 - state_store: 状态存储
785 lines
27 KiB
Python
785 lines
27 KiB
Python
import asyncio
|
|
import base64
|
|
import hashlib
|
|
import json
|
|
import os
|
|
import time
|
|
from collections import defaultdict
|
|
from collections.abc import Callable
|
|
from dataclasses import dataclass
|
|
from enum import Enum
|
|
|
|
from cryptography.hazmat.primitives.ciphers import Cipher, algorithms, modes
|
|
from coincurve import PrivateKey, PublicKey
|
|
|
|
from .defaults import (
|
|
AUTH_KIND,
|
|
CIRCUIT_BREAKER_RESET_MS,
|
|
CIRCUIT_BREAKER_THRESHOLD,
|
|
DM_CHAT_KIND,
|
|
DM_FILE_KIND,
|
|
DM_KIND,
|
|
GIFT_WRAP_KIND,
|
|
HEALTH_WINDOW_MS,
|
|
MAX_CIPHERTEXT_BYTES,
|
|
MAX_PLAINTEXT_CHARS,
|
|
RATE_LIMIT_GLOBAL_PER_WINDOW,
|
|
RATE_LIMIT_GLOBAL_WINDOW_MS,
|
|
RATE_LIMIT_PER_SENDER_PER_WINDOW,
|
|
RATE_LIMIT_PER_SENDER_WINDOW_MS,
|
|
RECONNECT_BACKOFF_MULTIPLIER,
|
|
RECONNECT_INITIAL_DELAY_S,
|
|
RECONNECT_MAX_DELAY_S,
|
|
STATE_PERSIST_DEBOUNCE_MS,
|
|
TIMESTAMP_SKEW_SECONDS,
|
|
)
|
|
from .event_utils import sign_event, verify_event
|
|
from .gift_wrap import create_gift_wrap, unwrap_gift_wrap
|
|
from .key_utils import normalize_pubkey
|
|
from .metrics import MetricsCollector, MetricsSnapshot
|
|
from .seen_tracker import SeenTracker
|
|
from .state_store import NostrBusState, save_bus_state
|
|
|
|
|
|
class BreakerState(Enum):
|
|
CLOSED = "closed"
|
|
OPEN = "open"
|
|
HALF_OPEN = "half_open"
|
|
|
|
|
|
class CircuitBreaker:
|
|
THRESHOLD = CIRCUIT_BREAKER_THRESHOLD
|
|
RESET_MS = CIRCUIT_BREAKER_RESET_MS
|
|
|
|
def __init__(self):
|
|
self.state = BreakerState.CLOSED
|
|
self.failure_count = 0
|
|
self.last_failure_time = 0.0
|
|
|
|
def can_attempt(self) -> bool:
|
|
if self.state == BreakerState.CLOSED:
|
|
return True
|
|
if self.state == BreakerState.OPEN:
|
|
elapsed_ms = (time.monotonic() - self.last_failure_time) * 1000
|
|
if elapsed_ms >= self.RESET_MS:
|
|
self.state = BreakerState.HALF_OPEN
|
|
return True
|
|
return False
|
|
return self.state == BreakerState.HALF_OPEN
|
|
|
|
def record_success(self):
|
|
self.state = BreakerState.CLOSED
|
|
self.failure_count = 0
|
|
|
|
def record_failure(self):
|
|
self.failure_count += 1
|
|
self.last_failure_time = time.monotonic()
|
|
if self.state == BreakerState.HALF_OPEN:
|
|
self.state = BreakerState.OPEN
|
|
elif self.failure_count >= self.THRESHOLD:
|
|
self.state = BreakerState.OPEN
|
|
|
|
|
|
class RelayHealthTracker:
|
|
WINDOW_MS = HEALTH_WINDOW_MS
|
|
|
|
def __init__(self):
|
|
self._success: dict[str, int] = defaultdict(int)
|
|
self._failure: dict[str, int] = defaultdict(int)
|
|
self._last_success: dict[str, float] = {}
|
|
self._latencies: dict[str, list[tuple[float, float]]] = defaultdict(list)
|
|
|
|
def record_success(self, relay: str, latency_ms: float = 0):
|
|
now = time.monotonic()
|
|
self._success[relay] += 1
|
|
self._last_success[relay] = now
|
|
self._latencies[relay].append((now, latency_ms))
|
|
|
|
def record_failure(self, relay: str):
|
|
self._failure[relay] += 1
|
|
|
|
def score(self, relay: str) -> float:
|
|
now = time.monotonic()
|
|
total = self._success[relay] + self._failure[relay]
|
|
if total == 0:
|
|
return 0.5
|
|
success_rate = self._success[relay] / total
|
|
recency_bonus = 0.0
|
|
if relay in self._last_success:
|
|
elapsed = (now - self._last_success[relay]) * 1000
|
|
recency_bonus = max(0, 1 - elapsed / self.WINDOW_MS) * 0.2
|
|
latency_penalty = 0.0
|
|
recent_latencies = [
|
|
latency_ms for t, latency_ms in self._latencies[relay]
|
|
if (now - t) * 1000 < self.WINDOW_MS
|
|
]
|
|
if recent_latencies:
|
|
avg_latency = sum(recent_latencies) / len(recent_latencies)
|
|
latency_penalty = min(0.2, avg_latency / 10000)
|
|
return max(0.0, min(1.0, success_rate + recency_bonus - latency_penalty))
|
|
|
|
def get_sorted_relays(self, relays: list[str]) -> list[str]:
|
|
return sorted(relays, key=lambda r: self.score(r), reverse=True)
|
|
|
|
|
|
class FixedWindowRateLimiter:
|
|
def __init__(self, window_ms: int, max_per_window: int, max_keys: int = 2048):
|
|
self.window_ms = window_ms
|
|
self.max_per_window = max_per_window
|
|
self.max_keys = max_keys
|
|
self._windows: dict[str, tuple[int, int]] = {}
|
|
|
|
def check(self, key: str) -> bool:
|
|
now_ms = int(time.monotonic() * 1000)
|
|
if key in self._windows:
|
|
win_start, count = self._windows[key]
|
|
if now_ms - win_start >= self.window_ms:
|
|
self._windows[key] = (now_ms, 1)
|
|
return self.max_per_window > 0
|
|
if count >= self.max_per_window:
|
|
return False
|
|
self._windows[key] = (win_start, count + 1)
|
|
return True
|
|
if len(self._windows) >= self.max_keys:
|
|
self._cleanup(now_ms)
|
|
self._windows[key] = (now_ms, 1)
|
|
return self.max_per_window > 0
|
|
|
|
@property
|
|
def entry_count(self) -> int:
|
|
return len(self._windows)
|
|
|
|
def _cleanup(self, now_ms: int):
|
|
expired = [
|
|
k
|
|
for k, (ws, _) in self._windows.items()
|
|
if now_ms - ws >= self.window_ms
|
|
]
|
|
for k in expired:
|
|
del self._windows[k]
|
|
|
|
|
|
@dataclass
|
|
class WebSocketConnection:
|
|
relay: str
|
|
ws: object = None
|
|
task: asyncio.Task | None = None
|
|
connected: bool = False
|
|
authenticated: bool = False
|
|
|
|
|
|
class WebSocketPool:
|
|
def __init__(self):
|
|
self._connections: dict[str, WebSocketConnection] = {}
|
|
self._subscriptions: dict[str, list[str]] = defaultdict(list)
|
|
self._on_auth_challenge: Callable | None = None
|
|
self._pending_ok: dict[str, asyncio.Future] = {}
|
|
|
|
def set_auth_handler(self, handler: Callable):
|
|
self._on_auth_challenge = handler
|
|
|
|
async def subscribe(
|
|
self,
|
|
relays: list[str],
|
|
filters: list[dict],
|
|
on_event: Callable,
|
|
on_eose: Callable,
|
|
on_close: Callable,
|
|
on_error: Callable,
|
|
on_connect: Callable | None = None,
|
|
):
|
|
for relay in relays:
|
|
if relay in self._connections:
|
|
continue
|
|
task = asyncio.create_task(
|
|
self._connect_loop(relay, filters, on_event, on_eose, on_close, on_error, on_connect)
|
|
)
|
|
self._connections[relay] = WebSocketConnection(relay=relay, task=task)
|
|
|
|
async def _connect_loop(
|
|
self,
|
|
relay: str,
|
|
filters: list[dict],
|
|
on_event: Callable,
|
|
on_eose: Callable,
|
|
on_close: Callable,
|
|
on_error: Callable,
|
|
on_connect: Callable | None = None,
|
|
):
|
|
delay = RECONNECT_INITIAL_DELAY_S
|
|
import websockets
|
|
|
|
is_reconnect = False
|
|
while True:
|
|
try:
|
|
async with websockets.connect(relay, max_size=MAX_CIPHERTEXT_BYTES * 4) as ws:
|
|
conn = self._connections.get(relay)
|
|
if conn:
|
|
conn.ws = ws
|
|
conn.connected = True
|
|
conn.authenticated = False
|
|
|
|
if on_connect:
|
|
await on_connect(relay, is_reconnect)
|
|
is_reconnect = True
|
|
|
|
sub_id = hashlib.sha256(os.urandom(16)).hexdigest()[:16]
|
|
self._subscriptions[relay].append(sub_id)
|
|
req = json.dumps(["REQ", sub_id, *filters])
|
|
await ws.send(req)
|
|
|
|
try:
|
|
async for raw_msg in ws:
|
|
try:
|
|
msg = json.loads(raw_msg)
|
|
except json.JSONDecodeError:
|
|
continue
|
|
|
|
msg_type = msg[0] if isinstance(msg, list) else None
|
|
if msg_type == "EVENT":
|
|
event = msg[2] if len(msg) > 2 else {}
|
|
await on_event(event, relay)
|
|
elif msg_type == "EOSE":
|
|
await on_eose(relay, sub_id)
|
|
elif msg_type == "AUTH":
|
|
challenge = msg[1] if len(msg) > 1 else ""
|
|
if self._on_auth_challenge:
|
|
await self._on_auth_challenge(relay, challenge, ws)
|
|
elif msg_type == "OK":
|
|
event_id = msg[1] if len(msg) > 1 else ""
|
|
accepted = msg[2] if len(msg) > 2 else False
|
|
message = msg[3] if len(msg) > 3 else ""
|
|
fut = self._pending_ok.pop(event_id, None)
|
|
if fut is not None and not fut.done():
|
|
fut.set_result({"accepted": accepted, "message": message})
|
|
if not accepted and "auth-required" in str(message):
|
|
if self._on_auth_challenge:
|
|
await self._on_auth_challenge(relay, "", ws)
|
|
await on_error(relay, msg)
|
|
elif msg_type in ("CLOSED", "NOTICE"):
|
|
if msg_type == "CLOSED" and len(msg) > 2 and "auth-required" in str(msg[2]):
|
|
if self._on_auth_challenge:
|
|
await self._on_auth_challenge(relay, "", ws)
|
|
await on_error(relay, msg)
|
|
except asyncio.CancelledError:
|
|
break
|
|
except Exception:
|
|
pass
|
|
|
|
except asyncio.CancelledError:
|
|
break
|
|
except Exception as e:
|
|
await on_error(relay, str(e))
|
|
|
|
conn = self._connections.get(relay)
|
|
if conn:
|
|
conn.connected = False
|
|
conn.ws = None
|
|
conn.authenticated = False
|
|
await on_close(relay)
|
|
delay = min(delay * RECONNECT_BACKOFF_MULTIPLIER, RECONNECT_MAX_DELAY_S)
|
|
await asyncio.sleep(delay)
|
|
|
|
async def publish(self, relay: str, event: dict, wait_ok: bool = False, ok_timeout: float = 5.0) -> bool:
|
|
conn = self._connections.get(relay)
|
|
if not conn or not conn.ws:
|
|
return False
|
|
try:
|
|
import websockets
|
|
|
|
req = json.dumps(["EVENT", event])
|
|
await conn.ws.send(req)
|
|
if wait_ok:
|
|
event_id = event.get("id", "")
|
|
fut = asyncio.get_event_loop().create_future()
|
|
self._pending_ok[event_id] = fut
|
|
try:
|
|
result = await asyncio.wait_for(fut, timeout=ok_timeout)
|
|
return result.get("accepted", False)
|
|
except TimeoutError:
|
|
self._pending_ok.pop(event_id, None)
|
|
return False
|
|
return True
|
|
except websockets.exceptions.ConnectionClosed:
|
|
return False
|
|
except Exception:
|
|
return False
|
|
|
|
async def unsubscribe(self, relay: str, sub_id: str):
|
|
conn = self._connections.get(relay)
|
|
if conn and conn.ws:
|
|
req = json.dumps(["CLOSE", sub_id])
|
|
await conn.ws.send(req)
|
|
|
|
async def close_all(self):
|
|
for relay, sub_ids in self._subscriptions.items():
|
|
for sub_id in sub_ids:
|
|
await self.unsubscribe(relay, sub_id)
|
|
for conn in self._connections.values():
|
|
if conn.task:
|
|
conn.task.cancel()
|
|
for conn in self._connections.values():
|
|
if conn.task:
|
|
try:
|
|
await conn.task
|
|
except asyncio.CancelledError:
|
|
pass
|
|
self._connections.clear()
|
|
self._subscriptions.clear()
|
|
self._pending_ok.clear()
|
|
|
|
|
|
def _get_shared_secret(sk_bytes: bytes, recipient_pk_hex: str) -> bytes:
|
|
recipient_pk = PublicKey(bytes.fromhex("02" + recipient_pk_hex))
|
|
shared_point = recipient_pk.multiply(sk_bytes)
|
|
return shared_point.format(compressed=False)[1:33]
|
|
|
|
|
|
def _aes_cbc_encrypt(key: bytes, plaintext: str) -> str:
|
|
iv = os.urandom(16)
|
|
pad_length = 16 - (len(plaintext.encode()) % 16)
|
|
padded = plaintext.encode() + b"\x00" * pad_length
|
|
cipher = Cipher(algorithms.AES(key), modes.CBC(iv))
|
|
encryptor = cipher.encryptor()
|
|
ciphertext = encryptor.update(padded) + encryptor.finalize()
|
|
encrypted_b64 = base64.b64encode(ciphertext).decode()
|
|
iv_b64 = base64.b64encode(iv).decode()
|
|
return f"{encrypted_b64}?iv={iv_b64}"
|
|
|
|
|
|
def _aes_cbc_decrypt(key: bytes, ciphertext_str: str) -> str:
|
|
if "?iv=" in ciphertext_str:
|
|
encrypted_b64, iv_b64 = ciphertext_str.split("?iv=", 1)
|
|
ciphertext = base64.b64decode(encrypted_b64)
|
|
iv = base64.b64decode(iv_b64)
|
|
else:
|
|
raw = base64.b64decode(ciphertext_str)
|
|
iv, ciphertext = raw[:16], raw[16:]
|
|
cipher = Cipher(algorithms.AES(key), modes.CBC(iv))
|
|
decryptor = cipher.decryptor()
|
|
plaintext = decryptor.update(ciphertext) + decryptor.finalize()
|
|
return plaintext.rstrip(b"\x00").decode()
|
|
|
|
|
|
def nip04_encrypt(sk: bytes, recipient_pk_hex: str, plaintext: str) -> str:
|
|
shared_secret = _get_shared_secret(sk, recipient_pk_hex)
|
|
return _aes_cbc_encrypt(shared_secret, plaintext)
|
|
|
|
|
|
def nip04_decrypt(sk: bytes, sender_pk_hex: str, ciphertext_b64: str) -> str:
|
|
shared_secret = _get_shared_secret(sk, sender_pk_hex)
|
|
return _aes_cbc_decrypt(shared_secret, ciphertext_b64)
|
|
|
|
|
|
@dataclass
|
|
class BusOptions:
|
|
on_metric: Callable | None = None
|
|
state_dir: str = ""
|
|
seen_max_entries: int = 100_000
|
|
encryption: str = "nip17"
|
|
|
|
|
|
class NostrBus:
|
|
def __init__(
|
|
self,
|
|
sk: bytes,
|
|
pk: str,
|
|
relays: list[str],
|
|
account_id: str,
|
|
options: BusOptions,
|
|
):
|
|
self.sk = sk
|
|
self.pk = pk
|
|
self.relays = relays
|
|
self.account_id = account_id
|
|
self.pool = WebSocketPool()
|
|
self.seen = SeenTracker()
|
|
self.circuit_breakers: dict[str, CircuitBreaker] = {}
|
|
self.health_tracker = RelayHealthTracker()
|
|
self.global_rate_limiter = FixedWindowRateLimiter(
|
|
window_ms=RATE_LIMIT_GLOBAL_WINDOW_MS,
|
|
max_per_window=RATE_LIMIT_GLOBAL_PER_WINDOW,
|
|
)
|
|
self.per_sender_rate_limiter = FixedWindowRateLimiter(
|
|
window_ms=RATE_LIMIT_PER_SENDER_WINDOW_MS,
|
|
max_per_window=RATE_LIMIT_PER_SENDER_PER_WINDOW,
|
|
)
|
|
self.metrics = MetricsCollector(on_metric=options.on_metric)
|
|
self.options = options
|
|
self._inflight: set[str] = set()
|
|
self._abort = asyncio.Event()
|
|
self._subscriptions: list[asyncio.Task] = []
|
|
self._persist_timer: asyncio.Task | None = None
|
|
self._pending_persist = False
|
|
self._state: NostrBusState | None = None
|
|
self._auth_relays: set[str] = set()
|
|
|
|
def _sk_hex(self) -> str:
|
|
return self.sk.hex() if isinstance(self.sk, bytes) else self.sk
|
|
|
|
async def _handle_auth(self, relay: str, challenge: str, ws):
|
|
event = sign_event(
|
|
sk=self.sk,
|
|
pubkey=self.pk,
|
|
kind=AUTH_KIND,
|
|
content="",
|
|
tags=[["relay", relay], ["challenge", challenge]],
|
|
)
|
|
req = json.dumps(["AUTH", event])
|
|
await ws.send(req)
|
|
conn = self.pool._connections.get(relay)
|
|
if conn:
|
|
conn.authenticated = True
|
|
self._auth_relays.add(relay)
|
|
|
|
async def start(
|
|
self,
|
|
since: int,
|
|
on_message: Callable,
|
|
authorize_sender: Callable,
|
|
):
|
|
for relay in self.relays:
|
|
self.circuit_breakers[relay] = CircuitBreaker()
|
|
|
|
encryption = getattr(self.options, "encryption", "nip17")
|
|
filters = []
|
|
if encryption in ("nip04", "both"):
|
|
filters.append({"kinds": [DM_KIND], "since": since, "#p": [self.pk]})
|
|
if encryption in ("nip17", "both"):
|
|
filters.append({"kinds": [GIFT_WRAP_KIND], "since": since, "#p": [self.pk]})
|
|
if not filters:
|
|
filters = [{"kinds": [DM_KIND], "since": since, "#p": [self.pk]}]
|
|
|
|
self.pool.set_auth_handler(self._handle_auth)
|
|
|
|
async def _on_event(event: dict, relay: str):
|
|
await self._handle_event(event, on_message, authorize_sender, since)
|
|
|
|
async def _on_eose(relay: str, sub_id: str):
|
|
self.metrics.emit("relay.message.eose", relay)
|
|
|
|
async def _on_connect(relay: str, is_reconnect: bool):
|
|
self.metrics.emit("relay.reconnect" if is_reconnect else "relay.connect", relay)
|
|
|
|
async def _on_close(relay: str):
|
|
self.metrics.emit("relay.disconnect", relay)
|
|
|
|
async def _on_error(relay: str, error):
|
|
self.metrics.emit("relay.error", relay)
|
|
|
|
await self.pool.subscribe(
|
|
relays=self.relays,
|
|
filters=filters,
|
|
on_event=_on_event,
|
|
on_eose=_on_eose,
|
|
on_close=_on_close,
|
|
on_error=_on_error,
|
|
on_connect=_on_connect,
|
|
)
|
|
|
|
async def _handle_event(
|
|
self,
|
|
event: dict,
|
|
on_message: Callable,
|
|
authorize_sender: Callable,
|
|
since: int,
|
|
):
|
|
if not isinstance(event, dict):
|
|
return
|
|
|
|
event_id = event.get("id", "")
|
|
pubkey = event.get("pubkey", "")
|
|
kind = event.get("kind", 0)
|
|
created_at = event.get("created_at", 0)
|
|
|
|
self.metrics.emit("event.received")
|
|
|
|
if not event_id:
|
|
return
|
|
|
|
if self.seen.peek(event_id) or event_id in self._inflight:
|
|
self.metrics.emit("event.duplicate")
|
|
return
|
|
|
|
if pubkey == self.pk:
|
|
self.metrics.emit("event.rejected.self_message")
|
|
return
|
|
|
|
now = int(time.time())
|
|
if since > 0 and created_at < since:
|
|
self.metrics.emit("event.rejected.stale")
|
|
return
|
|
|
|
if created_at > now + TIMESTAMP_SKEW_SECONDS:
|
|
self.metrics.emit("event.rejected.future")
|
|
return
|
|
|
|
encryption = getattr(self.options, "encryption", "nip17")
|
|
|
|
if kind == GIFT_WRAP_KIND and encryption in ("nip17", "both"):
|
|
return await self._handle_gift_wrap(event, on_message, authorize_sender, since)
|
|
|
|
if kind == DM_KIND and encryption in ("nip04", "both"):
|
|
return await self._handle_nip04_dm(event, on_message, authorize_sender, since)
|
|
|
|
self.metrics.emit("event.rejected.wrong_kind")
|
|
|
|
async def _handle_nip04_dm(
|
|
self,
|
|
event: dict,
|
|
on_message: Callable,
|
|
authorize_sender: Callable,
|
|
since: int,
|
|
):
|
|
event_id = event.get("id", "")
|
|
pubkey = event.get("pubkey", "")
|
|
|
|
tags = event.get("tags", [])
|
|
has_p_tag = any(
|
|
isinstance(t, list) and len(t) >= 2 and t[0] == "p" and t[1] == self.pk
|
|
for t in tags
|
|
)
|
|
if not has_p_tag:
|
|
self.metrics.emit("event.rejected.wrong_kind")
|
|
return
|
|
|
|
content = event.get("content", "")
|
|
if len(content.encode()) > MAX_CIPHERTEXT_BYTES:
|
|
self.metrics.emit("event.rejected.oversized_ciphertext")
|
|
return
|
|
|
|
if not self.global_rate_limiter.check("global"):
|
|
self.metrics.emit("rate_limit.global")
|
|
self.metrics.emit("event.rejected.rate_limited")
|
|
return
|
|
|
|
if not verify_event(event):
|
|
self.metrics.emit("event.rejected.invalid_signature")
|
|
return
|
|
|
|
if not self.per_sender_rate_limiter.check(pubkey):
|
|
self.metrics.emit("rate_limit.per_sender")
|
|
self.metrics.emit("event.rejected.rate_limited")
|
|
return
|
|
|
|
async def reply_to_sender(reply_text: str):
|
|
await self.send_dm(pubkey, reply_text)
|
|
|
|
try:
|
|
authorized = await authorize_sender(pubkey, reply_to_sender)
|
|
except Exception:
|
|
authorized = False
|
|
if not authorized:
|
|
return
|
|
|
|
try:
|
|
sender_pk = normalize_pubkey(pubkey)
|
|
plaintext = nip04_decrypt(self.sk, sender_pk, content)
|
|
except Exception:
|
|
self.metrics.emit("decrypt.failure")
|
|
self.metrics.emit("event.rejected.decrypt_failed")
|
|
return
|
|
|
|
self.metrics.emit("decrypt.success")
|
|
|
|
if len(plaintext) > MAX_PLAINTEXT_CHARS:
|
|
self.metrics.emit("event.rejected.oversized_plaintext")
|
|
return
|
|
|
|
self._inflight.add(event_id)
|
|
try:
|
|
await on_message(pubkey, plaintext, reply_to_sender)
|
|
self.seen.has(event_id)
|
|
self.metrics.emit("event.processed")
|
|
self._schedule_persist(event_id)
|
|
finally:
|
|
self._inflight.discard(event_id)
|
|
|
|
async def _handle_gift_wrap(
|
|
self,
|
|
event: dict,
|
|
on_message: Callable,
|
|
authorize_sender: Callable,
|
|
since: int,
|
|
):
|
|
event_id = event.get("id", "")
|
|
|
|
unwrapped = unwrap_gift_wrap(self.sk, event)
|
|
if unwrapped is None:
|
|
self.metrics.emit("event.rejected.gift_wrap_invalid")
|
|
return
|
|
|
|
if unwrapped.kind not in (DM_CHAT_KIND, DM_FILE_KIND):
|
|
self.metrics.emit("event.rejected.wrong_inner_kind")
|
|
return
|
|
|
|
pubkey = unwrapped.pubkey
|
|
content = unwrapped.content
|
|
|
|
if len(content.encode()) > MAX_CIPHERTEXT_BYTES:
|
|
self.metrics.emit("event.rejected.oversized_ciphertext")
|
|
return
|
|
|
|
if not self.global_rate_limiter.check("global"):
|
|
self.metrics.emit("rate_limit.global")
|
|
self.metrics.emit("event.rejected.rate_limited")
|
|
return
|
|
|
|
if not self.per_sender_rate_limiter.check(pubkey):
|
|
self.metrics.emit("rate_limit.per_sender")
|
|
self.metrics.emit("event.rejected.rate_limited")
|
|
return
|
|
|
|
async def reply_to_sender(reply_text: str):
|
|
await self.send_dm(pubkey, reply_text)
|
|
|
|
try:
|
|
authorized = await authorize_sender(pubkey, reply_to_sender)
|
|
except Exception:
|
|
authorized = False
|
|
if not authorized:
|
|
return
|
|
|
|
self.metrics.emit("decrypt.success")
|
|
|
|
if len(content) > MAX_PLAINTEXT_CHARS:
|
|
self.metrics.emit("event.rejected.oversized_plaintext")
|
|
return
|
|
|
|
self._inflight.add(event_id)
|
|
try:
|
|
await on_message(pubkey, content, reply_to_sender)
|
|
self.seen.has(event_id)
|
|
self.metrics.emit("event.processed")
|
|
self._schedule_persist(event_id)
|
|
finally:
|
|
self._inflight.discard(event_id)
|
|
|
|
def _schedule_persist(self, event_id: str):
|
|
self._pending_persist = True
|
|
if self._persist_timer is None or self._persist_timer.done():
|
|
self._persist_timer = asyncio.create_task(self._debounced_persist())
|
|
|
|
async def _debounced_persist(self):
|
|
try:
|
|
await asyncio.sleep(STATE_PERSIST_DEBOUNCE_MS / 1000)
|
|
if self._pending_persist and self.options.state_dir:
|
|
await self._flush_state()
|
|
except asyncio.CancelledError:
|
|
pass
|
|
|
|
async def send_dm(
|
|
self, to_pubkey: str, text: str, *,
|
|
reply_to_event_id: str | None = None, kind: int = DM_CHAT_KIND,
|
|
) -> dict:
|
|
to_pk = normalize_pubkey(to_pubkey)
|
|
encryption = getattr(self.options, "encryption", "nip17")
|
|
|
|
if encryption == "nip04":
|
|
return await self._send_nip04_dm(to_pk, text, reply_to_event_id)
|
|
return await self._send_nip17_dm(to_pk, text, reply_to_event_id, kind)
|
|
|
|
async def _send_nip04_dm(self, to_pk: str, text: str, reply_to_event_id: str | None = None) -> dict:
|
|
ciphertext = nip04_encrypt(self.sk, to_pk, text)
|
|
tags = [["p", to_pk]]
|
|
if reply_to_event_id:
|
|
tags.append(["e", reply_to_event_id, "", "reply"])
|
|
event = sign_event(
|
|
sk=self.sk,
|
|
pubkey=self.pk,
|
|
kind=DM_KIND,
|
|
content=ciphertext,
|
|
tags=tags,
|
|
)
|
|
return await self._publish_event(event)
|
|
|
|
async def _send_nip17_dm(
|
|
self, to_pk: str, text: str,
|
|
reply_to_event_id: str | None = None, kind: int = DM_CHAT_KIND,
|
|
) -> dict:
|
|
tags = [["p", to_pk]]
|
|
if reply_to_event_id:
|
|
tags.append(["e", reply_to_event_id, "", "reply"])
|
|
result = create_gift_wrap(
|
|
sender_sk_hex=self._sk_hex(),
|
|
sender_pk=self.pk,
|
|
receiver_pk=to_pk,
|
|
content=text,
|
|
kind=kind,
|
|
tags=tags,
|
|
)
|
|
return await self._publish_event(result.event)
|
|
|
|
async def send_reaction(self, target_event_id: str, reaction: str = "+", target_pubkey: str | None = None) -> dict:
|
|
from .reactions import send_reaction as _send_reaction
|
|
return await _send_reaction(self, target_event_id, reaction, target_pubkey)
|
|
|
|
async def send_deletion(self, event_ids: list[str], reason: str = "") -> dict:
|
|
from .deletion import create_deletion_event
|
|
event = create_deletion_event(self.sk, self.pk, event_ids, reason)
|
|
return await self._publish_event(event)
|
|
|
|
async def _publish_event(self, event: dict) -> dict:
|
|
sorted_relays = self.health_tracker.get_sorted_relays(self.relays)
|
|
last_error = None
|
|
for relay in sorted_relays:
|
|
breaker = self.circuit_breakers.get(relay)
|
|
if breaker and not breaker.can_attempt():
|
|
continue
|
|
try:
|
|
success = await self.pool.publish(relay, event)
|
|
if success:
|
|
self.health_tracker.record_success(relay)
|
|
if breaker:
|
|
breaker.record_success()
|
|
self.metrics.emit("relay.message.ok", relay)
|
|
import uuid
|
|
|
|
return {"ok": True, "message_id": uuid.uuid4().hex}
|
|
else:
|
|
self.health_tracker.record_failure(relay)
|
|
if breaker:
|
|
breaker.record_failure()
|
|
last_error = f"publish to {relay} returned False"
|
|
except Exception as e:
|
|
self.health_tracker.record_failure(relay)
|
|
if breaker:
|
|
breaker.record_failure()
|
|
last_error = str(e)
|
|
return {"ok": False, "error": last_error or "No relays available"}
|
|
|
|
async def _flush_state(self):
|
|
if not self._state or not self.options.state_dir:
|
|
return
|
|
self._state.recent_event_ids = self.seen.get_recent_ids()
|
|
from pathlib import Path
|
|
|
|
save_bus_state(self.account_id, self._state, Path(self.options.state_dir))
|
|
self._pending_persist = False
|
|
|
|
async def close(self):
|
|
self._abort.set()
|
|
if self._persist_timer and not self._persist_timer.done():
|
|
self._persist_timer.cancel()
|
|
try:
|
|
await self._persist_timer
|
|
except asyncio.CancelledError:
|
|
pass
|
|
self._persist_timer = None
|
|
if self._pending_persist:
|
|
await self._flush_state()
|
|
await self.pool.close_all()
|
|
|
|
def _set_state(self, state: NostrBusState):
|
|
self._state = state
|
|
|
|
def _sign_event_hash(self, event_id: str) -> bytes:
|
|
pk = PrivateKey(self.sk)
|
|
return pk.sign_schnorr(bytes.fromhex(event_id))
|
|
|
|
def get_metrics(self) -> MetricsSnapshot:
|
|
return self.metrics.snapshot(
|
|
seen_tracker_size=self.seen.size,
|
|
rate_limiter_entries=self.global_rate_limiter.entry_count,
|
|
)
|