refactor: use pynostr/coincurve instead of secp256k1 (#18)

---------

Co-authored-by: Riccardo Balbo <os@rblb.it>
Co-authored-by: Vlad Stan <stan.v.vlad@gmail.com>
This commit is contained in:
dni ⚡ 2025-12-02 21:04:55 +01:00 • committed by GitHub
commit 667a4b0528
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
10 changed files with 2056 additions and 1237 deletions

View file

@ -10,7 +10,7 @@ if [ ! -d ./lnbits ] ; then
fi fi
cd lnbits cd lnbits
echo $PWD echo $PWD
git checkout v1.0.0-rc7 git checkout dev
poetry env use 3.12 poetry env use 3.12
POETRY_PYTHON_PATH=$(poetry env info -p)/bin/python POETRY_PYTHON_PATH=$(poetry env info -p)/bin/python
ln -sf $POETRY_PYTHON_PATH /home/vscode/python ln -sf $POETRY_PYTHON_PATH /home/vscode/python

View file

@ -1,4 +1,4 @@
import secp256k1 from coincurve import PrivateKey
async def m001_initial(db): async def m001_initial(db):
@ -72,14 +72,14 @@ async def m003_default_config(db):
ON CONFLICT (key) DO UPDATE SET value = EXCLUDED.value; ON CONFLICT (key) DO UPDATE SET value = EXCLUDED.value;
""" """
) )
new_private_key = bytes.hex(secp256k1._gen_private_key()) private_key = PrivateKey()
await db.execute( await db.execute(
""" """
INSERT INTO nwcprovider.config (key, value) INSERT INTO nwcprovider.config (key, value)
VALUES ('provider_key', :provider_key) VALUES ('provider_key', :provider_key)
ON CONFLICT (key) DO UPDATE SET value = EXCLUDED.value; ON CONFLICT (key) DO UPDATE SET value = EXCLUDED.value;
""", """,
{"provider_key": new_private_key}, {"provider_key": private_key.to_hex()},
) )

101
nwcp.py
View file

@ -1,5 +1,4 @@
import asyncio import asyncio
import base64
import hashlib import hashlib
import json import json
import random import random
@ -7,13 +6,11 @@ import time
from collections.abc import Awaitable, Callable from collections.abc import Awaitable, Callable
from typing import Any, Union from typing import Any, Union
import secp256k1 from coincurve import PublicKeyXOnly
from Cryptodome import Random
from Cryptodome.Cipher import AES
from Cryptodome.Util.Padding import pad, unpad
from lnbits.helpers import encrypt_internal_message from lnbits.helpers import encrypt_internal_message
from lnbits.settings import settings from lnbits.settings import settings
from loguru import logger from loguru import logger
from pynostr.key import PrivateKey
from websockets.legacy.client import connect from websockets.legacy.client import connect
@ -75,7 +72,7 @@ class MainSubscription:
class NWCServiceProvider: class NWCServiceProvider:
def __init__( def __init__(
self, self,
private_key: str | None = None, private_key_hex: str | None = None,
relay: str | None = None, relay: str | None = None,
handle_missed_events: int = 0, handle_missed_events: int = 0,
): ):
@ -90,15 +87,18 @@ class NWCServiceProvider:
) )
self.relay = relay self.relay = relay
if not private_key: # Create random key if not private_key_hex: # Create random key
private_key = bytes.hex(secp256k1._gen_private_key()) self.private_key = PrivateKey()
self.private_key_hex = self.private_key.hex()
else:
self.private_key = PrivateKey.from_hex(private_key_hex)
self.private_key_hex = private_key_hex
self.private_key = secp256k1.PrivateKey(bytes.fromhex(private_key)) self.public_key = self.private_key.public_key
self.private_key_hex = private_key
self.public_key = self.private_key.pubkey
if not self.public_key: if not self.public_key:
raise Exception("Invalid public key") raise Exception("Invalid public key")
self.public_key_hex = self.public_key.serialize().hex()[2:]
self.public_key_hex = self.public_key.hex()
# List of supported methods # List of supported methods
self.supported_methods: list[str] = [] self.supported_methods: list[str] = []
@ -314,7 +314,7 @@ class NWCServiceProvider:
nwc_pubkey = event["pubkey"] nwc_pubkey = event["pubkey"]
content = event["content"] content = event["content"]
# Decrypt the content # Decrypt the content
content = self._decrypt_content(content, nwc_pubkey) content = self.private_key.decrypt_message(content, nwc_pubkey)
# Deserialize content # Deserialize content
content = json.loads(content) content = json.loads(content)
# Handle request # Handle request
@ -364,7 +364,9 @@ class NWCServiceProvider:
# Reference user # Reference user
res["tags"].append(["p", nwc_pubkey]) res["tags"].append(["p", nwc_pubkey])
# Finalize response event # Finalize response event
res["content"] = self._encrypt_content(res["content"], nwc_pubkey) res["content"] = self.private_key.encrypt_message(
res["content"], nwc_pubkey
)
self._sign_event(res) self._sign_event(res)
# Register response for this request, so we knows it is not stale # Register response for this request, so we knows it is not stale
@ -513,65 +515,6 @@ class NWCServiceProvider:
logger.debug("Reconnecting to NWC relay...") logger.debug("Reconnecting to NWC relay...")
await self._ratelimit("connecting") await self._ratelimit("connecting")
def _encrypt_content(
self, content: str, pubkey_hex: str, iv_seed: int | None = None
) -> str:
"""
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)
shared = pubkey.tweak_mul(bytes.fromhex(self.private_key_hex)).serialize()[1:]
# random iv (16B)
if not iv_seed:
iv = Random.new().read(AES.block_size)
else:
iv = hashlib.sha256(iv_seed.to_bytes(32, byteorder="big")).digest()
iv = iv[: AES.block_size]
aes = AES.new(shared, AES.MODE_CBC, iv)
content_bytes = content.encode("utf-8")
# padding
content_bytes = pad(content_bytes, AES.block_size)
encrypted_b64 = base64.b64encode(aes.encrypt(content_bytes)).decode("ascii")
iv_b64 = base64.b64encode(iv).decode("ascii")
encrypted_content = encrypted_b64 + "?iv=" + iv_b64
return encrypted_content
def _decrypt_content(self, content: str, pubkey_hex: str) -> str:
"""
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.
"""
pubkey = secp256k1.PublicKey(bytes.fromhex("02" + pubkey_hex), True)
shared = pubkey.tweak_mul(bytes.fromhex(self.private_key_hex)).serialize()[1:]
# extract iv and content
(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_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: def _verify_event(self, event: dict) -> bool:
""" """
Verify the event signature Verify the event signature
@ -596,10 +539,8 @@ class NWCServiceProvider:
if event_id != event["id"]: # Invalid event id if event_id != event["id"]: # Invalid event id
return False return False
pubkey_hex = event["pubkey"] pubkey_hex = event["pubkey"]
pubkey = secp256k1.PublicKey(bytes.fromhex("02" + pubkey_hex), True) pubkey = PublicKeyXOnly(bytes.fromhex(pubkey_hex))
if not pubkey.schnorr_verify( if not pubkey.verify(bytes.fromhex(event["sig"]), bytes.fromhex(event_id)):
bytes.fromhex(event_id), bytes.fromhex(event["sig"]), None, raw=True
):
return False return False
return True return True
@ -628,10 +569,8 @@ class NWCServiceProvider:
event["id"] = event_id event["id"] = event_id
event["pubkey"] = self.public_key_hex event["pubkey"] = self.public_key_hex
signature = ( signature = self.private_key.sign(bytes.fromhex(event_id))
self.private_key.schnorr_sign(bytes.fromhex(event_id), None, raw=True) event["sig"] = signature.hex() # type: ignore
).hex()
event["sig"] = signature
return event return event
async def cleanup(self): async def cleanup(self):

3054
poetry.lock generated

File diff suppressed because it is too large Load diff

View file

@ -7,11 +7,8 @@ authors = [{ name = "Riccardo Balbo", email = "oc@rblb.it" }]
urls = { Homepage = "https://lnbits.com", Repository = "https://github.com/lnbits/nwcprovider" } urls = { Homepage = "https://lnbits.com", Repository = "https://github.com/lnbits/nwcprovider" }
dependencies = [ "lnbits>1" ] dependencies = [ "lnbits>1" ]
[tool.poetry] [dependency-groups]
package-mode = false dev = [
[tool.uv]
dev-dependencies = [
"black", "black",
"pytest-asyncio", "pytest-asyncio",
"pytest", "pytest",
@ -21,12 +18,15 @@ dev-dependencies = [
"pytest-md", "pytest-md",
] ]
[tool.poetry]
package-mode = false
[tool.mypy] [tool.mypy]
plugins = ["pydantic.mypy"] plugins = ["pydantic.mypy"]
[[tool.mypy.overrides]] [[tool.mypy.overrides]]
module = [ module = [
"secp256k1.*", "pynostr.*",
] ]
ignore_missing_imports = "True" ignore_missing_imports = "True"

View file

@ -19,8 +19,8 @@ fi
docker run --name=lnbits_nwcprovider_ext_nostr_test \ docker run --name=lnbits_nwcprovider_ext_nostr_test \
-d \ -d \
--rm \ --rm \
-v $PWD/strfry.conf:/etc/strfry.conf \ -v $PWD/strfry.conf:/etc/strfry.conf:Z \
-v $PWD/strfry-data:/app/strfry-db \ -v $PWD/strfry-data:/app/strfry-db:Z \
-p 7777:7777 \ -p 7777:7777 \
ghcr.io/hoytech/strfry:latest ghcr.io/hoytech/strfry:latest

View file

@ -47,7 +47,7 @@ relay {
port = 7777 port = 7777
# Set OS-limit on maximum number of open files/sockets (if 0, don't attempt to set) (restart required) # Set OS-limit on maximum number of open files/sockets (if 0, don't attempt to set) (restart required)
nofiles = 1000000 nofiles = 0
# HTTP header that contains the client's real IP, before reverse proxying (ie x-real-ip) (MUST be all lower-case) # HTTP header that contains the client's real IP, before reverse proxying (ie x-real-ip) (MUST be all lower-case)
realIpHeader = "" realIpHeader = ""

View file

@ -1,5 +1,4 @@
import asyncio import asyncio
import base64
import hashlib import hashlib
import json import json
import random import random
@ -9,11 +8,8 @@ from typing import Union
import bolt11 import bolt11
import httpx import httpx
import pytest import pytest
import secp256k1
from Cryptodome import Random
from Cryptodome.Cipher import AES
from Cryptodome.Util.Padding import pad, unpad
from loguru import logger from loguru import logger
from pynostr.key import PrivateKey
from websockets.legacy.client import connect from websockets.legacy.client import connect
wallets = { wallets = {
@ -95,12 +91,12 @@ async def refresh_wallet_balances():
def gen_keypair(): def gen_keypair():
private_key_hex = bytes.hex(secp256k1._gen_private_key()) private_key = PrivateKey()
private_key = secp256k1.PrivateKey(bytes.fromhex(private_key_hex)) private_key_hex = private_key.hex()
public_key = private_key.pubkey public_key = private_key.public_key
if not public_key: if not public_key:
raise Exception("Error generating pubkey") raise Exception("Error generating pubkey")
public_key_hex = public_key.serialize().hex()[2:] public_key_hex = public_key.hex()
return {"priv": private_key_hex, "pub": public_key_hex} return {"priv": private_key_hex, "pub": public_key_hex}
@ -164,12 +160,12 @@ class NWCWallet:
self.event_queue = [] self.event_queue = []
self.subscriptions_count = 0 self.subscriptions_count = 0
self.sub_id = "" self.sub_id = ""
self.private_key = secp256k1.PrivateKey(bytes.fromhex(self.secret)) self.private_key = PrivateKey.from_hex(self.secret)
self.private_key_hex = self.secret self.private_key_hex = self.secret
self.public_key = self.private_key.pubkey self.public_key = self.private_key.public_key
if not self.public_key: if not self.public_key:
raise Exception("Error generating pubkey") raise Exception("Error generating pubkey")
self.public_key_hex = self.public_key.serialize().hex()[2:] self.public_key_hex = self.public_key.hex()
self.task = None self.task = None
async def close(self): async def close(self):
@ -243,43 +239,13 @@ class NWCWallet:
else: else:
break break
def _encrypt_content(
self, content: str, pubkey_hex: str, iv_seed: int | None = None
) -> str:
pubkey = secp256k1.PublicKey(bytes.fromhex("02" + pubkey_hex), True)
shared = pubkey.tweak_mul(bytes.fromhex(self.private_key_hex)).serialize()[1:]
if not iv_seed:
iv = Random.new().read(AES.block_size)
else:
iv = hashlib.sha256(iv_seed.to_bytes(32, byteorder="big")).digest()
iv = iv[: AES.block_size]
aes = AES.new(shared, AES.MODE_CBC, iv)
content_bytes = content.encode("utf-8")
content_bytes = pad(content_bytes, AES.block_size)
encrypted_b64 = base64.b64encode(aes.encrypt(content_bytes)).decode("ascii")
iv_b64 = base64.b64encode(iv).decode("ascii")
encrypted_content = encrypted_b64 + "?iv=" + iv_b64
return encrypted_content
def _decrypt_content(self, content: str, pubkey_hex: str) -> str:
pubkey = secp256k1.PublicKey(bytes.fromhex("02" + pubkey_hex), True)
shared = pubkey.tweak_mul(bytes.fromhex(self.private_key_hex)).serialize()[1:]
(encrypted_content_b64, iv_b64) = content.split("?iv=")
encrypted_content = base64.b64decode(encrypted_content_b64.encode("ascii"))
iv = base64.b64decode(iv_b64.encode("ascii"))
aes = AES.new(shared, AES.MODE_CBC, iv)
decrypted_bytes = aes.decrypt(encrypted_content)
decrypted_bytes = unpad(decrypted_bytes, AES.block_size)
decrypted = decrypted_bytes.decode("utf-8")
return decrypted
async def _on_message(self, _, message: str): async def _on_message(self, _, message: str):
logger.debug("Received message: " + message) logger.debug("Received message: " + message)
msg = json.loads(message) msg = json.loads(message)
if msg[0] == "EVENT": # Event message if msg[0] == "EVENT": # Event message
event = msg[2] event = msg[2]
nwc_pubkey = event["pubkey"] nwc_pubkey = event["pubkey"]
content = self._decrypt_content(event["content"], nwc_pubkey) content = self.private_key.decrypt_message(event["content"], nwc_pubkey)
content = json.loads(content) content = json.loads(content)
self.event_queue.append( self.event_queue.append(
{ {
@ -312,10 +278,9 @@ class NWCWallet:
event_id = hashlib.sha256(signature_data.encode()).hexdigest() event_id = hashlib.sha256(signature_data.encode()).hexdigest()
event["id"] = event_id event["id"] = event_id
event["pubkey"] = self.public_key_hex event["pubkey"] = self.public_key_hex
signature = ( signature = self.private_key.sign(bytes.fromhex(event_id))
self.private_key.schnorr_sign(bytes.fromhex(event_id), None, raw=True) # type error? returns str but is bytes
).hex() event["sig"] = signature.hex() # type: ignore
event["sig"] = signature
return event return event
async def send_event(self, method, params): async def send_event(self, method, params):
@ -331,7 +296,7 @@ class NWCWallet:
"content": json.dumps({"method": method, "params": params}), "content": json.dumps({"method": method, "params": params}),
} }
logger.debug("Sending event: " + str(event)) logger.debug("Sending event: " + str(event))
event["content"] = self._encrypt_content( event["content"] = self.private_key.encrypt_message(
event["content"], self.provider_pub_hex event["content"], self.provider_pub_hex
) )
self._sign_event(event) self._sign_event(event)
@ -888,7 +853,7 @@ async def test_idor_vulnerability():
f"http://localhost:5002/nwcprovider/api/v1/nwc/{nwc_wallet1['pubkey']}", f"http://localhost:5002/nwcprovider/api/v1/nwc/{nwc_wallet1['pubkey']}",
headers={"X-Api-Key": wallets["wallet2"]["admin_key"]}, headers={"X-Api-Key": wallets["wallet2"]["admin_key"]},
) )
assert resp.status_code == 500 assert resp.status_code == 400
assert "Pubkey has no associated wallet" in resp.text assert "Pubkey has no associated wallet" in resp.text

View file

@ -33,25 +33,22 @@ def test_supported_methods(nwc_service_provider):
def test_encrytdecrypt(nwc_service_provider, nwc_service_provider2): def test_encrytdecrypt(nwc_service_provider, nwc_service_provider2):
content = "Hello World" content = "Hello World"
expected_enc = "qVurNVISSl/9CfREIhk5Lg==?iv=QpCo5dI9gUcoLsSMLA7o7Q==" enc_a = nwc_service_provider.private_key.encrypt_message(
enc_a = nwc_service_provider._encrypt_content( content, nwc_service_provider2.public_key_hex
content, nwc_service_provider2.public_key_hex, 21
) )
enc_b = nwc_service_provider2._encrypt_content( enc_b = nwc_service_provider2.private_key.encrypt_message(
content, nwc_service_provider.public_key_hex, 21 content, nwc_service_provider.public_key_hex
) )
dec_a = nwc_service_provider2._decrypt_content( dec_a = nwc_service_provider2.private_key.decrypt_message(
enc_a, nwc_service_provider.public_key_hex enc_a, nwc_service_provider.public_key_hex
) )
dec_b = nwc_service_provider._decrypt_content( dec_b = nwc_service_provider.private_key.decrypt_message(
enc_b, nwc_service_provider2.public_key_hex enc_b, nwc_service_provider2.public_key_hex
) )
assert dec_a == content assert dec_a == content
assert dec_b == content assert dec_b == content
assert enc_a == expected_enc
assert enc_b == expected_enc
def test_signverify(nwc_service_provider, nwc_service_provider2): def test_signverify(nwc_service_provider, nwc_service_provider2):
@ -82,8 +79,8 @@ async def test_handle(nwc_service_provider, nwc_service_provider2):
content = nwc_service_provider._json_dumps( content = nwc_service_provider._json_dumps(
{"method": "pay_invoice", "params": {"invoice": "abc"}} {"method": "pay_invoice", "params": {"invoice": "abc"}}
) )
content = nwc_service_provider._encrypt_content( content = nwc_service_provider.private_key.encrypt_message(
content, nwc_service_provider2.public_key_hex, 21 content, nwc_service_provider2.public_key_hex
) )
event = { event = {
"kind": 23194, "kind": 23194,
@ -108,7 +105,7 @@ async def test_handle(nwc_service_provider, nwc_service_provider2):
assert len(sent_events) == 1 assert len(sent_events) == 1
for revent in sent_events: for revent in sent_events:
assert nwc_service_provider2._verify_event(revent) assert nwc_service_provider2._verify_event(revent)
content = nwc_service_provider2._decrypt_content( content = nwc_service_provider2.private_key.decrypt_message(
revent["content"], nwc_service_provider.public_key_hex revent["content"], nwc_service_provider.public_key_hex
) )
logger.debug(event) logger.debug(event)

View file

@ -1,10 +1,10 @@
from http import HTTPStatus from http import HTTPStatus
import secp256k1
from fastapi import APIRouter, Depends, Request from fastapi import APIRouter, Depends, Request
from fastapi.responses import JSONResponse from fastapi.responses import JSONResponse
from lnbits.core.models import WalletTypeInfo from lnbits.core.models import WalletTypeInfo
from lnbits.decorators import check_admin, require_admin_key from lnbits.decorators import check_admin, require_admin_key
from pynostr.key import PrivateKey
from .crud import ( from .crud import (
create_nwc, create_nwc,
@ -88,11 +88,13 @@ async def api_get_nwc(
nwc = await get_nwc( nwc = await get_nwc(
GetNWC(pubkey=pubkey, wallet=wallet_id, include_expired=include_expired) GetNWC(pubkey=pubkey, wallet=wallet_id, include_expired=include_expired)
) )
if not nwc: if not nwc:
raise Exception("Pubkey has no associated wallet") raise ValueError("Pubkey has no associated wallet")
res = NWCGetResponse( res = NWCGetResponse(
data=nwc, budgets=await get_budgets_nwc(GetBudgetsNWC(pubkey=pubkey)) data=nwc, budgets=await get_budgets_nwc(GetBudgetsNWC(pubkey=pubkey))
) )
return res return res
@ -123,11 +125,11 @@ async def api_get_pairing_url(req: Request, secret: str) -> str:
scheme = "wss" scheme = "wss"
netloc += "/nostrclient/api/v1/relay" netloc += "/nostrclient/api/v1/relay"
relay = f"{scheme}://{netloc}" relay = f"{scheme}://{netloc}"
psk = secp256k1.PrivateKey(bytes.fromhex(pprivkey)) psk = PrivateKey.from_hex(pprivkey)
ppk = psk.pubkey ppk = psk.public_key
if not ppk: if not ppk:
raise Exception("Error generating pubkey") raise Exception("Error generating pubkey")
ppubkey = ppk.serialize().hex()[2:] ppubkey = ppk.hex()
url = "nostr+walletconnect://" url = "nostr+walletconnect://"
url += ppubkey url += ppubkey
url += "?relay=" + relay url += "?relay=" + relay