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:
parent
9674110d04
commit
8f7aa33d4f
2 changed files with 87 additions and 7 deletions
39
nwcp.py
39
nwcp.py
|
|
@ -27,6 +27,7 @@ class MainSubscription:
|
|||
self.responses_eose = False
|
||||
self.events: dict[str, dict] = {}
|
||||
self.responses: list[str] = []
|
||||
self.seen_requests: dict[str, int] = {}
|
||||
|
||||
def get_stale(self) -> list[dict]:
|
||||
"""
|
||||
|
|
@ -49,6 +50,7 @@ class MainSubscription:
|
|||
"""
|
||||
Garbage collection, remove all the events that have a response older
|
||||
than expire seconds (defaults to 1 hour if 0 or None)
|
||||
and all seen requests that are expired
|
||||
"""
|
||||
expire = expire or 1 * 60 * 60
|
||||
now = int(time.time())
|
||||
|
|
@ -65,6 +67,11 @@ class MainSubscription:
|
|||
if len(deleted_ids) > 0:
|
||||
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:
|
||||
arbitrary_types_allowed = True
|
||||
|
||||
|
|
@ -125,7 +132,7 @@ class NWCServiceProvider:
|
|||
self.info_event_task = None
|
||||
|
||||
# Subscription
|
||||
self.sub = None
|
||||
self.sub: MainSubscription | None = None
|
||||
self.rate_limit: dict[str, RateLimit] = {}
|
||||
|
||||
# websocket connection
|
||||
|
|
@ -142,6 +149,8 @@ class NWCServiceProvider:
|
|||
# (handles reboots)
|
||||
self.handle_missed_events = handle_missed_events
|
||||
|
||||
self.event_max_age = self.handle_missed_events or 5 * 60
|
||||
|
||||
logger.info(
|
||||
"NWC Service is ready. relay: "
|
||||
+ str(self.relay)
|
||||
|
|
@ -270,11 +279,16 @@ class NWCServiceProvider:
|
|||
await asyncio.sleep(limit.backoff)
|
||||
limit.last_attempt_time = int(time.time())
|
||||
|
||||
def _create_subscription(self) -> MainSubscription:
|
||||
sub = MainSubscription()
|
||||
self.sub = sub
|
||||
return sub
|
||||
|
||||
async def _subscribe(self):
|
||||
"""
|
||||
[Re]Subscribe to receive nip 47 requests and responses from the relay
|
||||
"""
|
||||
self.sub = MainSubscription()
|
||||
sub = self._create_subscription()
|
||||
# Create requests subscription
|
||||
req_filter = {
|
||||
"kinds": [23194],
|
||||
|
|
@ -282,17 +296,17 @@ class NWCServiceProvider:
|
|||
# Since the last handle_missed_events seconds (handles reboots)
|
||||
"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)
|
||||
res_filter = {
|
||||
"kinds": [23195],
|
||||
"authors": [self.public_key_hex],
|
||||
"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
|
||||
await self._send(["REQ", self.sub.requests_sub_id, req_filter])
|
||||
await self._send(["REQ", self.sub.responses_sub_id, res_filter])
|
||||
await self._send(["REQ", sub.requests_sub_id, req_filter])
|
||||
await self._send(["REQ", sub.responses_sub_id, res_filter])
|
||||
|
||||
async def _on_connection(self, _):
|
||||
"""
|
||||
|
|
@ -335,6 +349,19 @@ class NWCServiceProvider:
|
|||
"""
|
||||
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"]
|
||||
content = event["content"]
|
||||
# Decrypt the content
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ import asyncio
|
|||
import json
|
||||
import random
|
||||
import string
|
||||
import time
|
||||
|
||||
import pytest
|
||||
from loguru import logger
|
||||
|
|
@ -75,6 +76,18 @@ def test_signverify(nwc_service_provider, nwc_service_provider2):
|
|||
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
|
||||
async def test_handle(nwc_service_provider, nwc_service_provider2):
|
||||
content = nwc_service_provider._json_dumps(
|
||||
|
|
@ -87,7 +100,7 @@ async def test_handle(nwc_service_provider, nwc_service_provider2):
|
|||
"kind": 23194,
|
||||
"content": content,
|
||||
"tags": [["p", nwc_service_provider2.public_key_hex]],
|
||||
"created_at": 1234567890,
|
||||
"created_at": int(time.time()),
|
||||
}
|
||||
signed = nwc_service_provider._sign_event(event)
|
||||
|
||||
|
|
@ -101,6 +114,7 @@ async def test_handle(nwc_service_provider, nwc_service_provider2):
|
|||
pass
|
||||
|
||||
nwc_service_provider2._send = _send_pass
|
||||
nwc_service_provider2._create_subscription()
|
||||
nwc_service_provider2.add_request_listener("pay_invoice", _handle_pay_invoice)
|
||||
sent_events = await nwc_service_provider2._handle_request(signed)
|
||||
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
|
||||
|
||||
|
||||
@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
|
||||
async def test_send_info_event(nwc_service_provider):
|
||||
"""_send_info_event should publish a signed kind-13194 event."""
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue