diff --git a/nwcp.py b/nwcp.py index 20c3195..bb3afa6 100644 --- a/nwcp.py +++ b/nwcp.py @@ -131,6 +131,9 @@ class NWCServiceProvider: # Periodic info event resend loop self.info_event_task = None + # Requests are handled independently from the relay receive loop. + self.request_tasks: set[asyncio.Task[list[dict]]] = set() + # Subscription self.sub: MainSubscription | None = None self.rate_limit: dict[str, RateLimit] = {} @@ -431,6 +434,19 @@ class NWCServiceProvider: sent_events.append(res) return sent_events + def _log_request_task_exception(self, task: asyncio.Future[list[dict]]) -> None: + if task.cancelled(): + return + exception = task.exception() + if exception: + logger.error("Error handling request: " + str(exception)) + + def _dispatch_request(self, event: dict) -> None: + task = asyncio.create_task(self._handle_request(event)) + self.request_tasks.add(task) + task.add_done_callback(self.request_tasks.discard) + task.add_done_callback(self._log_request_task_exception) + def _extract_expiration_from_tags(self, tags: list) -> int: expiration = -1 for tag in tags: @@ -470,7 +486,7 @@ class NWCServiceProvider: # already handled or stale, all stale requests will be handled # later when eose is received if self.sub.requests_eose and self.sub.responses_eose: - await self._handle_request(event) + self._dispatch_request(event) elif event["kind"] == 23195 and sub_id == self.sub.responses_sub_id: # Ensure the response is from this service provider if event["pubkey"] != self.public_key_hex: @@ -498,7 +514,7 @@ class NWCServiceProvider: if self.sub.requests_eose and self.sub.responses_eose: stales = self.sub.get_stale() for stale in stales: - await self._handle_request(stale) + self._dispatch_request(stale) async def _on_closed_message(self, msg): if not self.sub: @@ -656,6 +672,12 @@ class NWCServiceProvider: self.info_event_task.cancel() except Exception as e: logger.warning("Error closing info event loop: " + str(e)) + request_tasks = list(self.request_tasks) + for task in request_tasks: + task.cancel() + if request_tasks: + await asyncio.gather(*request_tasks, return_exceptions=True) + self.request_tasks.clear() # close the websocket try: if self.ws: diff --git a/tests/unit/test_nwcp.py b/tests/unit/test_nwcp.py index 6cd4d13..a631037 100644 --- a/tests/unit/test_nwcp.py +++ b/tests/unit/test_nwcp.py @@ -181,6 +181,72 @@ async def test_handle_rejects_same_event_replay( assert calls == 1 +@pytest.mark.asyncio +async def test_relay_dispatches_requests_without_waiting_for_previous_request( + nwc_service_provider, monkeypatch +): + sub = nwc_service_provider._create_subscription() + sub.requests_sub_id = "requests" + sub.requests_eose = True + sub.responses_eose = True + monkeypatch.setattr(nwc_service_provider, "_verify_event", lambda event: True) + + first_request_finished = asyncio.Event() + second_request_finished = asyncio.Event() + + async def _handle_request(event): + if event["id"] == "first": + await first_request_finished.wait() + else: + second_request_finished.set() + return [] + + monkeypatch.setattr(nwc_service_provider, "_handle_request", _handle_request) + + def request(event_id): + return json.dumps( + [ + "EVENT", + sub.requests_sub_id, + { + "id": event_id, + "kind": 23194, + "pubkey": "a" * 64, + "content": "", + "tags": [["p", nwc_service_provider.public_key_hex]], + "created_at": int(time.time()), + }, + ] + ) + + await nwc_service_provider._on_message(None, request("first")) + await nwc_service_provider._on_message(None, request("second")) + + await asyncio.wait_for(second_request_finished.wait(), timeout=1) + assert not first_request_finished.is_set() + + first_request_finished.set() + await asyncio.gather(*list(nwc_service_provider.request_tasks)) + + +@pytest.mark.asyncio +async def test_cleanup_cancels_pending_request_tasks(nwc_service_provider, monkeypatch): + request_started = asyncio.Event() + + async def _handle_request(event): + request_started.set() + await asyncio.Event().wait() + return [] + + monkeypatch.setattr(nwc_service_provider, "_handle_request", _handle_request) + nwc_service_provider._dispatch_request({"id": "pending"}) + await request_started.wait() + + await nwc_service_provider.cleanup() + + assert not nwc_service_provider.request_tasks + + @pytest.mark.asyncio async def test_send_info_event(nwc_service_provider): """_send_info_event should publish a signed kind-13194 event."""