Fix padding and docs
This commit is contained in:
parent
88da075509
commit
3dd608ff47
1 changed files with 69 additions and 33 deletions
|
|
@ -15,6 +15,7 @@ from typing import Union, List, Callable, Tuple
|
||||||
from lnbits.app import settings
|
from lnbits.app import settings
|
||||||
from lnbits.helpers import encrypt_internal_message
|
from lnbits.helpers import encrypt_internal_message
|
||||||
from urllib.parse import quote
|
from urllib.parse import quote
|
||||||
|
from Cryptodome.Util.Padding import pad, unpad
|
||||||
|
|
||||||
class MainSubscription:
|
class MainSubscription:
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
|
|
@ -26,6 +27,9 @@ class MainSubscription:
|
||||||
self.responses:List[str] = []
|
self.responses:List[str] = []
|
||||||
|
|
||||||
def getStale(self) -> List[Dict]:
|
def getStale(self) -> List[Dict]:
|
||||||
|
"""
|
||||||
|
Get all the pending events that do not have a response yet.
|
||||||
|
"""
|
||||||
pending_events = []
|
pending_events = []
|
||||||
for [id, event] in self.events.items():
|
for [id, event] in self.events.items():
|
||||||
if not id in self.responses:
|
if not id in self.responses:
|
||||||
|
|
@ -33,6 +37,9 @@ class MainSubscription:
|
||||||
return pending_events
|
return pending_events
|
||||||
|
|
||||||
def registerResponse(self, event_id:str):
|
def registerResponse(self, event_id:str):
|
||||||
|
"""
|
||||||
|
Register a response for a request event (not stale anymore)
|
||||||
|
"""
|
||||||
if not event_id in self.responses:
|
if not event_id in self.responses:
|
||||||
self.responses.append(event_id)
|
self.responses.append(event_id)
|
||||||
|
|
||||||
|
|
@ -56,37 +63,54 @@ class NWCServiceProvider:
|
||||||
self.public_key = self.private_key.pubkey
|
self.public_key = self.private_key.pubkey
|
||||||
self.public_key_hex = self.public_key.serialize().hex()[2:]
|
self.public_key_hex = self.public_key.serialize().hex()[2:]
|
||||||
|
|
||||||
|
# List of supported methods
|
||||||
self.supported_methods = []
|
self.supported_methods = []
|
||||||
|
|
||||||
|
# Keep track of the number of subscriptions (used for unique subid)
|
||||||
self.subscriptions_count = 0
|
self.subscriptions_count = 0
|
||||||
|
|
||||||
|
# Request listeners, listen to specific methods
|
||||||
self.request_listeners = {}
|
self.request_listeners = {}
|
||||||
|
|
||||||
|
# Reconnect task (if the connection is lost)
|
||||||
self.reconnect_task = None
|
self.reconnect_task = None
|
||||||
|
|
||||||
|
# Subscription
|
||||||
self.sub = None
|
self.sub = None
|
||||||
|
|
||||||
# websocket connection
|
# websocket connection
|
||||||
self.ws = None
|
self.ws = None
|
||||||
|
|
||||||
# if True the websocket is connected
|
# if True the websocket is connected
|
||||||
self.connected = False
|
self.connected = False
|
||||||
# if True the wallet is shutting down
|
|
||||||
|
# if True the instance is shutting down
|
||||||
self.shutdown = False
|
self.shutdown = False
|
||||||
|
|
||||||
logger.info("NWC Service is ready. relay: "+str(self.relay)+" pubkey: " +
|
logger.info("NWC Service is ready. relay: "+str(self.relay)+" pubkey: " +
|
||||||
self.public_key_hex)
|
self.public_key_hex)
|
||||||
|
|
||||||
def getSupportedMethods(self):
|
def getSupportedMethods(self):
|
||||||
|
"""
|
||||||
|
Returns the list of supported methods by this service provider.
|
||||||
|
"""
|
||||||
return self.supported_methods
|
return self.supported_methods
|
||||||
|
|
||||||
def addRequestListener(self, method: str, l: Callable[["NWCConnector", str, Dict], List[Tuple[Dict, Dict]]]):
|
def addRequestListener(self, method: str, l: Callable[["NWCServiceProvider", str, Dict], List[Tuple[Dict, Dict]]]):
|
||||||
|
"""
|
||||||
|
Adds a request listener for a specific method.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
method (str): The method name.
|
||||||
|
l (Callable[["NWCServiceProvider", str, Dict], List[Tuple[Dict, Dict]]]): The listener function
|
||||||
|
"""
|
||||||
if not method in self.supported_methods:
|
if not method in self.supported_methods:
|
||||||
self.supported_methods.append(method)
|
self.supported_methods.append(method)
|
||||||
self.request_listeners[method] = l
|
self.request_listeners[method] = l
|
||||||
|
|
||||||
async def start(self):
|
async def start(self):
|
||||||
"""
|
"""
|
||||||
Starts the NWC service connection.
|
Starts the NWC service provider.
|
||||||
"""
|
"""
|
||||||
self.reconnect_task = asyncio.create_task(self._connect_to_relay())
|
self.reconnect_task = asyncio.create_task(self._connect_to_relay())
|
||||||
|
|
||||||
|
|
@ -105,15 +129,17 @@ class NWCServiceProvider:
|
||||||
data = {k: v for k, v in data.items() if v is not None}
|
data = {k: v for k, v in data.items() if v is not None}
|
||||||
return json.dumps(data, separators=(',', ':'), ensure_ascii=False)
|
return json.dumps(data, separators=(',', ':'), ensure_ascii=False)
|
||||||
|
|
||||||
|
|
||||||
def _is_shutting_down(self) -> bool:
|
def _is_shutting_down(self) -> bool:
|
||||||
"""
|
"""
|
||||||
Returns True if the wallet is shutting down.
|
Returns True if the instance is shutting down.
|
||||||
"""
|
"""
|
||||||
return self.shutdown or not settings.lnbits_running
|
return self.shutdown or not settings.lnbits_running
|
||||||
|
|
||||||
|
|
||||||
async def _send(self, data: Dict):
|
async def _send(self, data: Dict):
|
||||||
"""
|
"""
|
||||||
Sends data to the NWC relay.
|
Sends data to the relay.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
data (Dict): The data to be sent.
|
data (Dict): The data to be sent.
|
||||||
|
|
@ -125,6 +151,7 @@ class NWCServiceProvider:
|
||||||
tx = self._json_dumps(data)
|
tx = self._json_dumps(data)
|
||||||
await self.ws.send(tx)
|
await self.ws.send(tx)
|
||||||
|
|
||||||
|
|
||||||
def _get_new_subid(self) -> str:
|
def _get_new_subid(self) -> str:
|
||||||
"""
|
"""
|
||||||
Generates a unique subscription id.
|
Generates a unique subscription id.
|
||||||
|
|
@ -145,7 +172,7 @@ class NWCServiceProvider:
|
||||||
|
|
||||||
async def _wait_for_connection(self):
|
async def _wait_for_connection(self):
|
||||||
"""
|
"""
|
||||||
Waits until the wallet is connected to the relay.
|
Waits until the connection is established.
|
||||||
"""
|
"""
|
||||||
while not self.connected:
|
while not self.connected:
|
||||||
if self._is_shutting_down():
|
if self._is_shutting_down():
|
||||||
|
|
@ -159,7 +186,6 @@ class NWCServiceProvider:
|
||||||
[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()
|
self.sub = MainSubscription()
|
||||||
|
|
||||||
# Create requests subscription
|
# Create requests subscription
|
||||||
req_filter = {
|
req_filter = {
|
||||||
"kinds": [23194],
|
"kinds": [23194],
|
||||||
|
|
@ -179,6 +205,7 @@ class NWCServiceProvider:
|
||||||
await self._send(["REQ", self.sub.requests_sub_id, req_filter])
|
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", self.sub.responses_sub_id, res_filter])
|
||||||
|
|
||||||
|
|
||||||
async def _on_connection(self,ws):
|
async def _on_connection(self,ws):
|
||||||
"""
|
"""
|
||||||
On connection callback, announce the service provider methods and subscribe to nip67 events.
|
On connection callback, announce the service provider methods and subscribe to nip67 events.
|
||||||
|
|
@ -256,11 +283,9 @@ class NWCServiceProvider:
|
||||||
# Register response for this request, so we knows it is not stale
|
# Register response for this request, so we knows it is not stale
|
||||||
self.sub.registerResponse(event["id"])
|
self.sub.registerResponse(event["id"])
|
||||||
# Send response event
|
# Send response event
|
||||||
|
|
||||||
await self._send(["EVENT", res])
|
await self._send(["EVENT", res])
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
async def _on_message(self, ws, message: str):
|
async def _on_message(self, ws, message: str):
|
||||||
"""
|
"""
|
||||||
Handle incoming messages from the relay.
|
Handle incoming messages from the relay.
|
||||||
|
|
@ -340,20 +365,21 @@ class NWCServiceProvider:
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error("Error parsing event: "+str(e))
|
logger.error("Error parsing event: "+str(e))
|
||||||
|
|
||||||
|
|
||||||
async def _connect_to_relay(self):
|
async def _connect_to_relay(self):
|
||||||
"""
|
"""
|
||||||
Initiate websocket connection to the relay.
|
Initiate websocket connection to the relay.
|
||||||
"""
|
"""
|
||||||
await asyncio.sleep(1)
|
await asyncio.sleep(1)
|
||||||
logger.debug("Connecting to NWC relay "+self.relay)
|
logger.debug("Connecting to NWC relay "+self.relay)
|
||||||
while not self._is_shutting_down(): # Reconnect until the wallet is shutting down
|
while not self._is_shutting_down(): # Reconnect until the instance is shutting down
|
||||||
logger.debug('Creating new connection...')
|
logger.debug('Creating new connection...')
|
||||||
try:
|
try:
|
||||||
async with websockets.connect(self.relay) as ws:
|
async with websockets.connect(self.relay) as ws:
|
||||||
self.ws = ws
|
self.ws = ws
|
||||||
self.connected = True
|
self.connected = True
|
||||||
await self._on_connection(ws)
|
await self._on_connection(ws)
|
||||||
while not self._is_shutting_down(): # receive messages until the wallet is shutting down
|
while not self._is_shutting_down(): # receive messages until the instance is shutting down
|
||||||
try:
|
try:
|
||||||
reply = await ws.recv()
|
reply = await ws.recv()
|
||||||
await self._on_message(ws, reply)
|
await self._on_message(ws, reply)
|
||||||
|
|
@ -372,39 +398,45 @@ class NWCServiceProvider:
|
||||||
logger.debug("Reconnecting to NWC relay in 5 seconds...")
|
logger.debug("Reconnecting to NWC relay in 5 seconds...")
|
||||||
await asyncio.sleep(5)
|
await asyncio.sleep(5)
|
||||||
|
|
||||||
|
|
||||||
def _encrypt_content(self, content: str, pubkey_hex:str) -> str:
|
def _encrypt_content(self, content: str, pubkey_hex:str) -> str:
|
||||||
"""
|
"""
|
||||||
Encrypts the content to be sent to the service.
|
Encrypts the content for the given public key
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
content (str): The content to be encrypted.
|
content (str): The content to be encrypted.
|
||||||
|
pubkey_hex (str): The public key in hex format.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
str: The encrypted content.
|
str: The encrypted content.
|
||||||
"""
|
"""
|
||||||
pubkey = secp256k1.PublicKey(
|
pubkey = secp256k1.PublicKey(
|
||||||
bytes.fromhex("02" + pubkey_hex), True)
|
bytes.fromhex("02" + pubkey_hex), True)
|
||||||
|
|
||||||
shared = pubkey.tweak_mul(bytes.fromhex(
|
shared = pubkey.tweak_mul(bytes.fromhex(
|
||||||
self.private_key_hex)).serialize()[1:]
|
self.private_key_hex)).serialize()[1:]
|
||||||
# random iv (16B)
|
# random iv (16B)
|
||||||
iv = Random.new().read(AES.block_size)
|
iv = Random.new().read(AES.block_size)
|
||||||
aes = AES.new(shared, AES.MODE_CBC, iv)
|
aes = AES.new(shared, AES.MODE_CBC, iv)
|
||||||
|
|
||||||
|
content_bytes = content.encode("utf-8")
|
||||||
|
|
||||||
# padding
|
# padding
|
||||||
def pad(s): return s + (16 - len(s) % 16) * chr(16 - len(s) % 16)
|
content_bytes = pad(content_bytes, AES.block_size)
|
||||||
content = pad(content).encode("utf-8")
|
|
||||||
# Encrypt
|
encrypted_b64 = base64.b64encode(
|
||||||
encryptedB64 = base64.b64encode(aes.encrypt(content)).decode("ascii")
|
aes.encrypt(content_bytes)).decode("ascii")
|
||||||
ivB64 = base64.b64encode(iv).decode("ascii")
|
ivB64 = base64.b64encode(iv).decode("ascii")
|
||||||
encryptedContent = encryptedB64 + "?iv=" + ivB64
|
encrypted_content = encrypted_b64 + "?iv=" + ivB64
|
||||||
return encryptedContent
|
return encrypted_content
|
||||||
|
|
||||||
|
|
||||||
def _decrypt_content(self, content: str , pubkey_hex:str) -> str:
|
def _decrypt_content(self, content: str , pubkey_hex:str) -> str:
|
||||||
"""
|
"""
|
||||||
Decrypts the content coming from the service.
|
Decrypts the content for the given public key
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
content (str): The encrypted content.
|
content (str): The encrypted content.
|
||||||
|
pubkey_hex (str): The public key in hex format.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
str: The decrypted content.
|
str: The decrypted content.
|
||||||
|
|
@ -415,25 +447,27 @@ class NWCServiceProvider:
|
||||||
shared = pubkey.tweak_mul(bytes.fromhex(
|
shared = pubkey.tweak_mul(bytes.fromhex(
|
||||||
self.private_key_hex)).serialize()[1:]
|
self.private_key_hex)).serialize()[1:]
|
||||||
# extract iv and content
|
# extract iv and content
|
||||||
(encryptedContentB64, ivB64) = content.split("?iv=")
|
(encrypted_content_b64, iv_b64) = content.split("?iv=")
|
||||||
encryptedContent = base64.b64decode(
|
encrypted_content = base64.b64decode(
|
||||||
encryptedContentB64.encode("ascii"))
|
encrypted_content_b64.encode("ascii"))
|
||||||
iv = base64.b64decode(ivB64.encode("ascii"))
|
iv = base64.b64decode(iv_b64.encode("ascii"))
|
||||||
# Decrypt
|
# Decrypt
|
||||||
aes = AES.new(shared, AES.MODE_CBC, iv)
|
aes = AES.new(shared, AES.MODE_CBC, iv)
|
||||||
decrypted = aes.decrypt(encryptedContent).decode("utf-8")
|
decrypted_bytes = aes.decrypt(encrypted_content)
|
||||||
def unpad(s): return s[:-ord(s[len(s)-1:])]
|
decrypted_bytes = unpad(decrypted_bytes, AES.block_size)
|
||||||
return unpad(decrypted)
|
decrypted = decrypted_bytes.decode("utf-8")
|
||||||
|
return decrypted
|
||||||
|
|
||||||
|
|
||||||
def _verify_event(self, event: Dict) -> bool:
|
def _verify_event(self, event: Dict) -> bool:
|
||||||
"""
|
"""
|
||||||
Signs the event (in place) with the service secret
|
Verify the event signature
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
event (Dict): The event to be signed.
|
event (Dict): The event to verify.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Dict: The input event with the signature added.
|
bool: True if the event signature is valid, False otherwise.
|
||||||
"""
|
"""
|
||||||
signature_data = self._json_dumps([
|
signature_data = self._json_dumps([
|
||||||
0,
|
0,
|
||||||
|
|
@ -452,9 +486,10 @@ class NWCServiceProvider:
|
||||||
return False
|
return False
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
|
||||||
def _sign_event(self, event: Dict) -> Dict:
|
def _sign_event(self, event: Dict) -> Dict:
|
||||||
"""
|
"""
|
||||||
Signs the event (in place) with the service secret
|
Signs the event (in place)
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
event (Dict): The event to be signed.
|
event (Dict): The event to be signed.
|
||||||
|
|
@ -480,6 +515,7 @@ class NWCServiceProvider:
|
||||||
event["sig"] = signature
|
event["sig"] = signature
|
||||||
return event
|
return event
|
||||||
|
|
||||||
|
|
||||||
async def cleanup(self):
|
async def cleanup(self):
|
||||||
logger.debug("Closing NWC Service Provider connection")
|
logger.debug("Closing NWC Service Provider connection")
|
||||||
self.shutdown = True # Mark for shutdown
|
self.shutdown = True # Mark for shutdown
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue