feat: track seen events until they expire (#47)

* yield event loop when processing multiple invoices
* track seen events until they expire
This commit is contained in:
Riccardo Balbo 2026-06-26 13:28:12 +02:00 • committed by GitHub
commit 8f7aa33d4f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 87 additions and 7 deletions

39
nwcp.py
View file

@ -27,6 +27,7 @@ class MainSubscription:
self.responses_eose = False self.responses_eose = False
self.events: dict[str, dict] = {} self.events: dict[str, dict] = {}
self.responses: list[str] = [] self.responses: list[str] = []
self.seen_requests: dict[str, int] = {}
def get_stale(self) -> list[dict]: def get_stale(self) -> list[dict]:
""" """
@ -49,6 +50,7 @@ class MainSubscription:
""" """
Garbage collection, remove all the events that have a response older Garbage collection, remove all the events that have a response older
than expire seconds (defaults to 1 hour if 0 or None) than expire seconds (defaults to 1 hour if 0 or None)
and all seen requests that are expired
""" """
expire = expire or 1 * 60 * 60 expire = expire or 1 * 60 * 60
now = int(time.time()) now = int(time.time())
@ -65,6 +67,11 @@ class MainSubscription:
if len(deleted_ids) > 0: if len(deleted_ids) > 0:
logger.debug("Garbage collected " + str(len(deleted_ids)) + " events") logger.debug("Garbage collected " + str(len(deleted_ids)) + " events")
# Clean seen requests
for event_id, expiry in list(self.seen_requests.items()):
if expiry < now:
del self.seen_requests[event_id]
class Config: class Config:
arbitrary_types_allowed = True arbitrary_types_allowed = True
@ -125,7 +132,7 @@ class NWCServiceProvider:
self.info_event_task = None self.info_event_task = None
# Subscription # Subscription
self.sub = None self.sub: MainSubscription | None = None
self.rate_limit: dict[str, RateLimit] = {} self.rate_limit: dict[str, RateLimit] = {}
# websocket connection # websocket connection
@ -142,6 +149,8 @@ class NWCServiceProvider:
# (handles reboots) # (handles reboots)
self.handle_missed_events = handle_missed_events self.handle_missed_events = handle_missed_events
self.event_max_age = self.handle_missed_events or 5 * 60
logger.info( logger.info(
"NWC Service is ready. relay: " "NWC Service is ready. relay: "
+ str(self.relay) + str(self.relay)
@ -270,11 +279,16 @@ class NWCServiceProvider:
await asyncio.sleep(limit.backoff) await asyncio.sleep(limit.backoff)
limit.last_attempt_time = int(time.time()) limit.last_attempt_time = int(time.time())
def _create_subscription(self) -> MainSubscription:
sub = MainSubscription()
self.sub = sub
return sub
async def _subscribe(self): async def _subscribe(self):
""" """
[Re]Subscribe to receive nip 47 requests and responses from the relay [Re]Subscribe to receive nip 47 requests and responses from the relay
""" """
self.sub = MainSubscription() sub = self._create_subscription()
# Create requests subscription # Create requests subscription
req_filter = { req_filter = {
"kinds": [23194], "kinds": [23194],
@ -282,17 +296,17 @@ class NWCServiceProvider:
# Since the last handle_missed_events seconds (handles reboots) # Since the last handle_missed_events seconds (handles reboots)
"since": int(time.time()) - self.handle_missed_events, "since": int(time.time()) - self.handle_missed_events,
} }
self.sub.requests_sub_id = self._get_new_subid() sub.requests_sub_id = self._get_new_subid()
# Create responses subscription (needed to track previosly responded requests) # Create responses subscription (needed to track previosly responded requests)
res_filter = { res_filter = {
"kinds": [23195], "kinds": [23195],
"authors": [self.public_key_hex], "authors": [self.public_key_hex],
"since": int(time.time()) - self.handle_missed_events, "since": int(time.time()) - self.handle_missed_events,
} }
self.sub.responses_sub_id = self._get_new_subid() sub.responses_sub_id = self._get_new_subid()
# Subscribe # Subscribe
await self._send(["REQ", self.sub.requests_sub_id, req_filter]) await self._send(["REQ", sub.requests_sub_id, req_filter])
await self._send(["REQ", self.sub.responses_sub_id, res_filter]) await self._send(["REQ", sub.responses_sub_id, res_filter])
async def _on_connection(self, _): async def _on_connection(self, _):
""" """
@ -335,6 +349,19 @@ class NWCServiceProvider:
""" """
Handle a nwc request Handle a nwc request
""" """
if not self.sub:
raise Exception("Subscription is not established")
sub = self.sub
expire = sub.seen_requests.get(event["id"])
if expire or event["created_at"] < int(time.time() - self.event_max_age):
raise Exception("Event is too old or already handled")
expiration = self._extract_expiration_from_tags(event["tags"])
if expiration <= 0:
expiration = int(time.time() + self.event_max_age)
sub.seen_requests[event["id"]] = expiration
nwc_pubkey = event["pubkey"] nwc_pubkey = event["pubkey"]
content = event["content"] content = event["content"]
# Decrypt the content # Decrypt the content

View file

@ -2,6 +2,7 @@ import asyncio
import json import json
import random import random
import string import string
import time
import pytest import pytest
from loguru import logger from loguru import logger
@ -75,6 +76,18 @@ def test_signverify(nwc_service_provider, nwc_service_provider2):
assert nwc_service_provider2._verify_event(signed) assert nwc_service_provider2._verify_event(signed)
def test_default_event_max_age(nwc_service_provider):
assert nwc_service_provider.event_max_age == 5 * 60
assert (
NWCServiceProvider(
"d7b5232fba0e02e32cfe26f20cdf2c803b27ecd81052c2dd5d17e5e1a333fe58",
"",
handle_missed_events=123,
).event_max_age
== 123
)
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_handle(nwc_service_provider, nwc_service_provider2): async def test_handle(nwc_service_provider, nwc_service_provider2):
content = nwc_service_provider._json_dumps( content = nwc_service_provider._json_dumps(
@ -87,7 +100,7 @@ async def test_handle(nwc_service_provider, nwc_service_provider2):
"kind": 23194, "kind": 23194,
"content": content, "content": content,
"tags": [["p", nwc_service_provider2.public_key_hex]], "tags": [["p", nwc_service_provider2.public_key_hex]],
"created_at": 1234567890, "created_at": int(time.time()),
} }
signed = nwc_service_provider._sign_event(event) signed = nwc_service_provider._sign_event(event)
@ -101,6 +114,7 @@ async def test_handle(nwc_service_provider, nwc_service_provider2):
pass pass
nwc_service_provider2._send = _send_pass nwc_service_provider2._send = _send_pass
nwc_service_provider2._create_subscription()
nwc_service_provider2.add_request_listener("pay_invoice", _handle_pay_invoice) nwc_service_provider2.add_request_listener("pay_invoice", _handle_pay_invoice)
sent_events = await nwc_service_provider2._handle_request(signed) sent_events = await nwc_service_provider2._handle_request(signed)
assert len(sent_events) == 1 assert len(sent_events) == 1
@ -128,6 +142,45 @@ async def test_handle(nwc_service_provider, nwc_service_provider2):
assert p_tag[0][1] == nwc_service_provider.public_key_hex assert p_tag[0][1] == nwc_service_provider.public_key_hex
@pytest.mark.asyncio
async def test_handle_rejects_same_event_replay(
nwc_service_provider, nwc_service_provider2
):
content = nwc_service_provider._json_dumps(
{"method": "pay_invoice", "params": {"invoice": "abc"}}
)
content = nwc_service_provider.private_key.encrypt_message(
content, nwc_service_provider2.public_key_hex
)
event = {
"kind": 23194,
"content": content,
"tags": [["p", nwc_service_provider2.public_key_hex]],
"created_at": int(time.time()),
}
signed = nwc_service_provider._sign_event(event)
calls = 0
async def _handle_pay_invoice(provider, pubkey, content):
nonlocal calls
calls += 1
return [({"preimage": "00000"}, None, [])]
async def _send_pass(obj):
pass
nwc_service_provider2._send = _send_pass
nwc_service_provider2._create_subscription()
nwc_service_provider2.add_request_listener("pay_invoice", _handle_pay_invoice)
await nwc_service_provider2._handle_request(signed)
with pytest.raises(Exception, match="already handled"):
await nwc_service_provider2._handle_request(signed)
assert calls == 1
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_send_info_event(nwc_service_provider): async def test_send_info_event(nwc_service_provider):
"""_send_info_event should publish a signed kind-13194 event.""" """_send_info_event should publish a signed kind-13194 event."""