diff --git a/README.md b/README.md index 70593a8..dee86a1 100644 --- a/README.md +++ b/README.md @@ -50,6 +50,7 @@ flowchart LR 3. **Fan-Out** - Subscription requests are sent to all configured relays 4. **Aggregation** - Events from all relays are collected and deduplicated 5. **Response** - Events are sent back to the client with the original subscription ID +6. **Publish Acknowledgement** - Every `EVENT` a client publishes gets exactly one `["OK", , , ]` reply (NIP-01): `true` as soon as any relay accepts it, `false` once every relay has rejected it, the wait times out (10 s), or no relay is connected ## Configuration diff --git a/config.json b/config.json index 1f58e7b..b9e4956 100644 --- a/config.json +++ b/config.json @@ -1,7 +1,7 @@ { "name": "Nostr Client", "short_description": "Nostr relay multiplexer", - "version": "1.1.0", + "version": "1.2.0-aio.2", "tile": "/nostrclient/static/images/nostr-bitcoin.png", "contributors": ["calle", "motorina0", "dni"], "min_lnbits_version": "1.4.0", diff --git a/nostr/client/client.py b/nostr/client/client.py index d6fb5c8..0412c1b 100644 --- a/nostr/client/client.py +++ b/nostr/client/client.py @@ -39,11 +39,13 @@ class NostrClient: callback_events_func=None, callback_notices_func=None, callback_eosenotices_func=None, + callback_command_results_func=None, ): while self.running: self._check_events(callback_events_func) self._check_notices(callback_notices_func) self._check_eos_notices(callback_eosenotices_func) + self._check_command_results(callback_command_results_func) await asyncio.sleep(0.2) @@ -73,3 +75,12 @@ class NostrClient: callback_eosenotices_func(event_msg) except Exception as e: logger.debug(e) + + def _check_command_results(self, callback_command_results_func=None): + try: + while self.relay_manager.message_pool.has_command_results(): + result_msg = self.relay_manager.message_pool.get_command_result() + if callback_command_results_func: + callback_command_results_func(result_msg) + except Exception as e: + logger.debug(e) diff --git a/nostr/message_pool.py b/nostr/message_pool.py index a3e6c5f..30c1469 100644 --- a/nostr/message_pool.py +++ b/nostr/message_pool.py @@ -27,11 +27,22 @@ class EndOfStoredEventsMessage: self.url = url +class CommandResultMessage: + """An `["OK", , , ]` reply from one relay.""" + + def __init__(self, event_id: str, accepted: bool, message: str, url: str) -> None: + self.event_id = event_id + self.accepted = accepted + self.message = message + self.url = url + + class MessagePool: def __init__(self) -> None: self.events: Queue[EventMessage] = Queue() self.notices: Queue[NoticeMessage] = Queue() self.eose_notices: Queue[EndOfStoredEventsMessage] = Queue() + self.command_results: Queue[CommandResultMessage] = Queue() self._unique_events: set = set() self.lock: Lock = Lock() @@ -47,6 +58,9 @@ class MessagePool: def get_eose_notice(self): return self.eose_notices.get() + def get_command_result(self): + return self.command_results.get() + def has_events(self): return self.events.qsize() > 0 @@ -56,6 +70,9 @@ class MessagePool: def has_eose_notices(self): return self.eose_notices.qsize() > 0 + def has_command_results(self): + return self.command_results.qsize() > 0 + def _process_message(self, message: str, url: str): message_json = json.loads(message) message_type = message_json[0] @@ -75,6 +92,13 @@ class MessagePool: self.notices.put(NoticeMessage(message_json[1], url)) elif message_type == RelayMessageType.END_OF_STORED_EVENTS: self.eose_notices.put(EndOfStoredEventsMessage(message_json[1], url)) + elif message_type == RelayMessageType.COMMAND_RESULT: + message = message_json[3] if len(message_json) > 3 else "" + self.command_results.put( + CommandResultMessage( + message_json[1], bool(message_json[2]), str(message), url + ) + ) def _accept_event(self, event_message: EventMessage): """ diff --git a/router.py b/router.py index a7054e9..53ce387 100644 --- a/router.py +++ b/router.py @@ -1,5 +1,6 @@ import asyncio import json +import time from typing import ClassVar from fastapi import WebSocket, WebSocketDisconnect @@ -9,11 +10,30 @@ from loguru import logger from .nostr.client.client import NostrClient # from . import nostr_client -from .nostr.message_pool import EndOfStoredEventsMessage, EventMessage, NoticeMessage +from .nostr.message_pool import ( + CommandResultMessage, + EndOfStoredEventsMessage, + EventMessage, + NoticeMessage, +) nostr_client: NostrClient = NostrClient() all_routers: list["NostrRouter"] = [] +# How long to wait for relays to answer an EVENT before replying `OK false`. +PUBLISH_TIMEOUT_SECONDS = 10 + + +class PendingPublish: + """An EVENT a client sent that still awaits its `OK` reply.""" + + def __init__(self, event_id: str, relay_count: int) -> None: + self.event_id = event_id + # relays connected at publish time; once this many answered we stop + # waiting for an acceptance + self.relay_count = relay_count + self.sent_at = time.time() + class NostrRouter: received_subscription_events: ClassVar[dict[str, list[EventMessage]]] = {} @@ -21,12 +41,14 @@ class NostrRouter: received_subscription_eosenotices: ClassVar[dict[str, EndOfStoredEventsMessage]] = ( {} ) + received_command_results: ClassVar[dict[str, list[CommandResultMessage]]] = {} def __init__(self, websocket: WebSocket): self.connected: bool = True self.websocket: WebSocket = websocket self.tasks: list[asyncio.Task] = [] self.original_subscription_ids: dict[str, str] = {} + self.pending_publishes: dict[str, PendingPublish] = {} @property def subscriptions(self) -> list[str]: @@ -40,6 +62,9 @@ class NostrRouter: async def stop(self): nostr_client.relay_manager.close_subscriptions(self.subscriptions) self.connected = False + for event_id in self.pending_publishes: + NostrRouter.received_command_results.pop(event_id, None) + self.pending_publishes.clear() for t in self.tasks: try: @@ -74,6 +99,7 @@ class NostrRouter: while self.connected: try: await self._handle_subscriptions() + await self._handle_command_results() self._handle_notices() except Exception as e: logger.debug(f"Failed to handle response for client: '{e!s}'.") @@ -120,6 +146,42 @@ class NostrRouter: f"[NOSTRCLIENT] Error in _handle_received_subscription_events: {e}" ) + async def _handle_command_results(self): + """ + Reply exactly one `["OK", , , ]` per EVENT a + client published (NIP-01). The EVENT was fanned out to every relay, so + several OKs can come back for one id: the client gets `true` as soon + as any relay accepts, and `false` once every relay has rejected it or + the wait times out. + """ + for event_id in list(self.pending_publishes.keys()): + pending = self.pending_publishes[event_id] + results = NostrRouter.received_command_results.get(event_id, []) + accepted = next((r for r in results if r.accepted), None) + if accepted: + await self._send_ok(event_id, True, accepted.message) + elif len(results) >= pending.relay_count: + await self._send_ok(event_id, False, results[-1].message) + elif time.time() - pending.sent_at > PUBLISH_TIMEOUT_SECONDS: + message = ( + results[-1].message + if results + else "error: timed out waiting for relays" + ) + await self._send_ok(event_id, False, message) + else: + continue + self.pending_publishes.pop(event_id, None) + NostrRouter.received_command_results.pop(event_id, None) + + async def _send_ok(self, event_id: str, accepted: bool, message: str): + try: + await self.websocket.send_text( + json.dumps(["OK", event_id, accepted, message]) + ) + except Exception as e: + logger.debug(f"Failed to send OK for '{event_id}': {e}") + def _handle_notices(self): while len(NostrRouter.received_subscription_notices): my_event = NostrRouter.received_subscription_notices.pop(0) @@ -141,9 +203,26 @@ class NostrRouter: return if json_data[0] == "EVENT": - nostr_client.relay_manager.publish_message(json_str) + await self._handle_client_event(json_data, json_str) return + async def _handle_client_event(self, json_data, json_str): + event = json_data[1] if len(json_data) > 1 else None + event_id = event.get("id") if isinstance(event, dict) else None + if not event_id: + logger.debug("Ignoring EVENT without an id.") + return + + connected = [ + r for r in nostr_client.relay_manager.relays.values() if r.connected + ] + if not connected: + await self._send_ok(event_id, False, "error: no relays connected") + return + + self.pending_publishes[event_id] = PendingPublish(event_id, len(connected)) + nostr_client.relay_manager.publish_message(json_str) + def _handle_client_req(self, json_data): subscription_id = json_data[1] logger.info(f"New subscription: '{subscription_id}'") diff --git a/tasks.py b/tasks.py index 2a76765..e2abdcb 100644 --- a/tasks.py +++ b/tasks.py @@ -4,8 +4,13 @@ import threading from loguru import logger from .crud import get_relays -from .nostr.message_pool import EndOfStoredEventsMessage, EventMessage, NoticeMessage -from .router import NostrRouter, nostr_client +from .nostr.message_pool import ( + CommandResultMessage, + EndOfStoredEventsMessage, + EventMessage, + NoticeMessage, +) +from .router import NostrRouter, all_routers, nostr_client async def init_relays(): @@ -55,12 +60,24 @@ async def subscribe_events(): NostrRouter.received_subscription_eosenotices[sub_id] = event_message + def callback_command_results(result_message: CommandResultMessage): + event_id = result_message.event_id + # Only keep OKs some client is still waiting on; the rest would + # accumulate forever (events published by other relay users, or + # results arriving after the client already got its reply). + if not any(event_id in r.pending_publishes for r in all_routers): + return + NostrRouter.received_command_results.setdefault(event_id, []).append( + result_message + ) + def wrap_async_subscribe(): asyncio.run( nostr_client.subscribe( callback_events, callback_notices, callback_eose_notices, + callback_command_results, ) ) diff --git a/tests/test_router_ok.py b/tests/test_router_ok.py new file mode 100644 index 0000000..44863a1 --- /dev/null +++ b/tests/test_router_ok.py @@ -0,0 +1,141 @@ +import json +import time + +import pytest + +from .. import router as router_module +from ..nostr.message_pool import CommandResultMessage +from ..router import PUBLISH_TIMEOUT_SECONDS, NostrRouter + + +class FakeWebSocket: + def __init__(self): + self.sent: list[list] = [] + + async def send_text(self, text: str): + self.sent.append(json.loads(text)) + + +class FakeRelay: + def __init__(self, connected: bool): + self.connected = connected + + +class FakeRelayManager: + def __init__(self, relays: dict[str, FakeRelay]): + self.relays = relays + self.published: list[str] = [] + + def publish_message(self, message: str): + self.published.append(message) + + def close_subscriptions(self, subscriptions): + pass + + +EVENT_ID = "ab" * 32 +EVENT_MSG = json.dumps(["EVENT", {"id": EVENT_ID, "kind": 1, "content": "hi"}]) + + +def _router(monkeypatch, relays: dict[str, FakeRelay]): + manager = FakeRelayManager(relays) + monkeypatch.setattr(router_module.nostr_client, "relay_manager", manager) + NostrRouter.received_command_results.clear() + ws = FakeWebSocket() + return NostrRouter(ws), ws, manager # type: ignore[arg-type] + + +def _ok_from(url: str, accepted: bool, message: str = ""): + NostrRouter.received_command_results.setdefault(EVENT_ID, []).append( + CommandResultMessage(EVENT_ID, accepted, message, url) + ) + + +@pytest.mark.asyncio +async def test_no_connected_relays_replies_ok_false_immediately(monkeypatch): + router, ws, manager = _router(monkeypatch, {"wss://a": FakeRelay(False)}) + + await router._handle_client_to_nostr(EVENT_MSG) + + assert manager.published == [] + assert ws.sent == [["OK", EVENT_ID, False, "error: no relays connected"]] + assert router.pending_publishes == {} + + +@pytest.mark.asyncio +async def test_event_without_id_is_ignored(monkeypatch): + router, ws, manager = _router(monkeypatch, {"wss://a": FakeRelay(True)}) + + await router._handle_client_to_nostr(json.dumps(["EVENT", {"kind": 1}])) + + assert manager.published == [] + assert ws.sent == [] + + +@pytest.mark.asyncio +async def test_any_accepting_relay_yields_ok_true(monkeypatch): + relays = {"wss://a": FakeRelay(True), "wss://b": FakeRelay(True)} + router, ws, manager = _router(monkeypatch, relays) + + await router._handle_client_to_nostr(EVENT_MSG) + assert manager.published == [EVENT_MSG] + assert EVENT_ID in router.pending_publishes + + # nothing answered yet: no OK + await router._handle_command_results() + assert ws.sent == [] + + _ok_from("wss://a", False, "blocked: kind not allowed") + await router._handle_command_results() + assert ws.sent == [] # one rejection out of two relays: keep waiting + + _ok_from("wss://b", True, "") + await router._handle_command_results() + assert ws.sent == [["OK", EVENT_ID, True, ""]] + assert router.pending_publishes == {} + assert EVENT_ID not in NostrRouter.received_command_results + + # a late OK must not produce a second reply + _ok_from("wss://a", True, "") + await router._handle_command_results() + assert len(ws.sent) == 1 + + +@pytest.mark.asyncio +async def test_all_relays_rejecting_yields_ok_false(monkeypatch): + relays = {"wss://a": FakeRelay(True), "wss://b": FakeRelay(True)} + router, ws, _ = _router(monkeypatch, relays) + + await router._handle_client_to_nostr(EVENT_MSG) + _ok_from("wss://a", False, "invalid: bad sig") + _ok_from("wss://b", False, "blocked: kind not allowed") + await router._handle_command_results() + + assert ws.sent == [["OK", EVENT_ID, False, "blocked: kind not allowed"]] + assert router.pending_publishes == {} + + +@pytest.mark.asyncio +async def test_timeout_yields_ok_false(monkeypatch): + router, ws, _ = _router(monkeypatch, {"wss://a": FakeRelay(True)}) + + await router._handle_client_to_nostr(EVENT_MSG) + router.pending_publishes[EVENT_ID].sent_at = ( + time.time() - PUBLISH_TIMEOUT_SECONDS - 1 + ) + await router._handle_command_results() + + assert ws.sent == [["OK", EVENT_ID, False, "error: timed out waiting for relays"]] + assert router.pending_publishes == {} + + +@pytest.mark.asyncio +async def test_stop_drops_pending_publishes(monkeypatch): + router, _, _ = _router(monkeypatch, {"wss://a": FakeRelay(True)}) + + await router._handle_client_to_nostr(EVENT_MSG) + _ok_from("wss://a", True) + await router.stop() + + assert router.pending_publishes == {} + assert NostrRouter.received_command_results == {}