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.helpers import encrypt_internal_message
|
||||
from urllib.parse import quote
|
||||
from Cryptodome.Util.Padding import pad, unpad
|
||||
|
||||
class MainSubscription:
|
||||
def __init__(self):
|
||||
|
|
@ -26,6 +27,9 @@ class MainSubscription:
|
|||
self.responses:List[str] = []
|
||||
|
||||
def getStale(self) -> List[Dict]:
|
||||
"""
|
||||
Get all the pending events that do not have a response yet.
|
||||
"""
|
||||
pending_events = []
|
||||
for [id, event] in self.events.items():
|
||||
if not id in self.responses:
|
||||
|
|
@ -33,6 +37,9 @@ class MainSubscription:
|
|||
return pending_events
|
||||
|
||||
def registerResponse(self, event_id:str):
|
||||
"""
|
||||
Register a response for a request event (not stale anymore)
|
||||
"""
|
||||
if not event_id in self.responses:
|
||||
self.responses.append(event_id)
|
||||
|
||||
|
|
@ -56,37 +63,54 @@ class NWCServiceProvider:
|
|||
self.public_key = self.private_key.pubkey
|
||||
self.public_key_hex = self.public_key.serialize().hex()[2:]
|
||||
|
||||
# List of supported methods
|
||||
self.supported_methods = []
|
||||
|
||||
# Keep track of the number of subscriptions (used for unique subid)
|
||||
self.subscriptions_count = 0
|
||||
|
||||
# Request listeners, listen to specific methods
|
||||
self.request_listeners = {}
|
||||
|
||||
# Reconnect task (if the connection is lost)
|
||||
self.reconnect_task = None
|
||||
|
||||
# Subscription
|
||||
self.sub = None
|
||||
|
||||
# websocket connection
|
||||
self.ws = None
|
||||
|
||||
# if True the websocket is connected
|
||||
self.connected = False
|
||||
# if True the wallet is shutting down
|
||||
|
||||
# if True the instance is shutting down
|
||||
self.shutdown = False
|
||||
|
||||
logger.info("NWC Service is ready. relay: "+str(self.relay)+" pubkey: " +
|
||||
self.public_key_hex)
|
||||
|
||||
def getSupportedMethods(self):
|
||||
"""
|
||||
Returns the list of supported methods by this service provider.
|
||||
"""
|
||||
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:
|
||||
self.supported_methods.append(method)
|
||||
self.request_listeners[method] = l
|
||||
|
||||
async def start(self):
|
||||
"""
|
||||
Starts the NWC service connection.
|
||||
Starts the NWC service provider.
|
||||
"""
|
||||
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}
|
||||
return json.dumps(data, separators=(',', ':'), ensure_ascii=False)
|
||||
|
||||
|
||||
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
|
||||
|
||||
|
||||
async def _send(self, data: Dict):
|
||||
"""
|
||||
Sends data to the NWC relay.
|
||||
Sends data to the relay.
|
||||
|
||||
Args:
|
||||
data (Dict): The data to be sent.
|
||||
|
|
@ -125,6 +151,7 @@ class NWCServiceProvider:
|
|||
tx = self._json_dumps(data)
|
||||
await self.ws.send(tx)
|
||||
|
||||
|
||||
def _get_new_subid(self) -> str:
|
||||
"""
|
||||
Generates a unique subscription id.
|
||||
|
|
@ -145,7 +172,7 @@ class NWCServiceProvider:
|
|||
|
||||
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:
|
||||
if self._is_shutting_down():
|
||||
|
|
@ -158,8 +185,7 @@ class NWCServiceProvider:
|
|||
"""
|
||||
[Re]Subscribe to receive nip 47 requests and responses from the relay
|
||||
"""
|
||||
self.sub = MainSubscription()
|
||||
|
||||
self.sub = MainSubscription()
|
||||
# Create requests subscription
|
||||
req_filter = {
|
||||
"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.responses_sub_id, res_filter])
|
||||
|
||||
|
||||
async def _on_connection(self,ws):
|
||||
"""
|
||||
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
|
||||
self.sub.registerResponse(event["id"])
|
||||
# Send response event
|
||||
|
||||
await self._send(["EVENT", res])
|
||||
|
||||
|
||||
|
||||
async def _on_message(self, ws, message: str):
|
||||
"""
|
||||
Handle incoming messages from the relay.
|
||||
|
|
@ -340,20 +365,21 @@ class NWCServiceProvider:
|
|||
except Exception as e:
|
||||
logger.error("Error parsing event: "+str(e))
|
||||
|
||||
|
||||
async def _connect_to_relay(self):
|
||||
"""
|
||||
Initiate websocket connection to the relay.
|
||||
"""
|
||||
await asyncio.sleep(1)
|
||||
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...')
|
||||
try:
|
||||
async with websockets.connect(self.relay) as ws:
|
||||
self.ws = ws
|
||||
self.connected = True
|
||||
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:
|
||||
reply = await ws.recv()
|
||||
await self._on_message(ws, reply)
|
||||
|
|
@ -372,39 +398,45 @@ class NWCServiceProvider:
|
|||
logger.debug("Reconnecting to NWC relay in 5 seconds...")
|
||||
await asyncio.sleep(5)
|
||||
|
||||
|
||||
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:
|
||||
content (str): The content to be encrypted.
|
||||
pubkey_hex (str): The public key in hex format.
|
||||
|
||||
Returns:
|
||||
str: The encrypted content.
|
||||
"""
|
||||
pubkey = secp256k1.PublicKey(
|
||||
bytes.fromhex("02" + pubkey_hex), True)
|
||||
|
||||
bytes.fromhex("02" + pubkey_hex), True)
|
||||
shared = pubkey.tweak_mul(bytes.fromhex(
|
||||
self.private_key_hex)).serialize()[1:]
|
||||
# random iv (16B)
|
||||
iv = Random.new().read(AES.block_size)
|
||||
aes = AES.new(shared, AES.MODE_CBC, iv)
|
||||
|
||||
content_bytes = content.encode("utf-8")
|
||||
|
||||
# padding
|
||||
def pad(s): return s + (16 - len(s) % 16) * chr(16 - len(s) % 16)
|
||||
content = pad(content).encode("utf-8")
|
||||
# Encrypt
|
||||
encryptedB64 = base64.b64encode(aes.encrypt(content)).decode("ascii")
|
||||
content_bytes = pad(content_bytes, AES.block_size)
|
||||
|
||||
encrypted_b64 = base64.b64encode(
|
||||
aes.encrypt(content_bytes)).decode("ascii")
|
||||
ivB64 = base64.b64encode(iv).decode("ascii")
|
||||
encryptedContent = encryptedB64 + "?iv=" + ivB64
|
||||
return encryptedContent
|
||||
encrypted_content = encrypted_b64 + "?iv=" + ivB64
|
||||
return encrypted_content
|
||||
|
||||
|
||||
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:
|
||||
content (str): The encrypted content.
|
||||
pubkey_hex (str): The public key in hex format.
|
||||
|
||||
Returns:
|
||||
str: The decrypted content.
|
||||
|
|
@ -415,25 +447,27 @@ class NWCServiceProvider:
|
|||
shared = pubkey.tweak_mul(bytes.fromhex(
|
||||
self.private_key_hex)).serialize()[1:]
|
||||
# extract iv and content
|
||||
(encryptedContentB64, ivB64) = content.split("?iv=")
|
||||
encryptedContent = base64.b64decode(
|
||||
encryptedContentB64.encode("ascii"))
|
||||
iv = base64.b64decode(ivB64.encode("ascii"))
|
||||
(encrypted_content_b64, iv_b64) = content.split("?iv=")
|
||||
encrypted_content = base64.b64decode(
|
||||
encrypted_content_b64.encode("ascii"))
|
||||
iv = base64.b64decode(iv_b64.encode("ascii"))
|
||||
# Decrypt
|
||||
aes = AES.new(shared, AES.MODE_CBC, iv)
|
||||
decrypted = aes.decrypt(encryptedContent).decode("utf-8")
|
||||
def unpad(s): return s[:-ord(s[len(s)-1:])]
|
||||
return unpad(decrypted)
|
||||
decrypted_bytes = aes.decrypt(encrypted_content)
|
||||
decrypted_bytes = unpad(decrypted_bytes, AES.block_size)
|
||||
decrypted = decrypted_bytes.decode("utf-8")
|
||||
return decrypted
|
||||
|
||||
|
||||
def _verify_event(self, event: Dict) -> bool:
|
||||
"""
|
||||
Signs the event (in place) with the service secret
|
||||
Verify the event signature
|
||||
|
||||
Args:
|
||||
event (Dict): The event to be signed.
|
||||
event (Dict): The event to verify.
|
||||
|
||||
Returns:
|
||||
Dict: The input event with the signature added.
|
||||
bool: True if the event signature is valid, False otherwise.
|
||||
"""
|
||||
signature_data = self._json_dumps([
|
||||
0,
|
||||
|
|
@ -452,9 +486,10 @@ class NWCServiceProvider:
|
|||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _sign_event(self, event: Dict) -> Dict:
|
||||
"""
|
||||
Signs the event (in place) with the service secret
|
||||
Signs the event (in place)
|
||||
|
||||
Args:
|
||||
event (Dict): The event to be signed.
|
||||
|
|
@ -479,6 +514,7 @@ class NWCServiceProvider:
|
|||
bytes.fromhex(event_id), None, raw=True)).hex()
|
||||
event["sig"] = signature
|
||||
return event
|
||||
|
||||
|
||||
async def cleanup(self):
|
||||
logger.debug("Closing NWC Service Provider connection")
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue