edges hardening

This commit is contained in:
Riccardo Balbo 2025-02-22 12:16:21 +00:00
commit e3957b5135
4 changed files with 322 additions and 3 deletions

84
crud.py
View file

@ -15,11 +15,28 @@ from .models import (
GetNWC, GetNWC,
NWCNewBudget NWCNewBudget
) )
from .paranoia import assert_valid_wallet_id, assert_valid_pubkey, assert_sane_string, assert_valid_timestamp_seconds, assert_valid_msats, assert_valid_positive_int
db = Database("ext_nwcprovider") db = Database("ext_nwcprovider")
async def create_nwc(data: CreateNWCKey) -> NWCKey: async def create_nwc(data: CreateNWCKey) -> NWCKey:
# hardening #
assert_valid_pubkey(data.pubkey)
assert_valid_wallet_id(data.wallet)
assert_sane_string(data.description)
assert_valid_timestamp_seconds(data.expires_at)
for permission in data.permissions:
assert_sane_string(permission)
if data.budgets:
for budget in data.budgets:
assert_valid_msats(budget.budget_msats)
assert_valid_positive_int(budget.refresh_window)
assert_valid_timestamp_seconds(budget.created_at)
# ## #
nwckey_entry = NWCKey( nwckey_entry = NWCKey(
pubkey=data.pubkey, pubkey=data.pubkey,
wallet=data.wallet, wallet=data.wallet,
@ -43,6 +60,12 @@ async def create_nwc(data: CreateNWCKey) -> NWCKey:
async def delete_nwc(data: DeleteNWC) -> None: async def delete_nwc(data: DeleteNWC) -> None:
# hardening #
assert_valid_pubkey(data.pubkey)
assert_valid_wallet_id(data.wallet)
# ## #
await db.execute( await db.execute(
"DELETE FROM nwcprovider.keys WHERE pubkey = :pubkey AND wallet = :wallet", "DELETE FROM nwcprovider.keys WHERE pubkey = :pubkey AND wallet = :wallet",
{"pubkey": data.pubkey, "wallet": data.wallet}, {"pubkey": data.pubkey, "wallet": data.wallet},
@ -50,6 +73,13 @@ async def delete_nwc(data: DeleteNWC) -> None:
async def get_wallet_nwcs(data: GetWalletNWC) -> List[NWCKey]: async def get_wallet_nwcs(data: GetWalletNWC) -> List[NWCKey]:
expires = int(time.time()) if not data.include_expired else -1
# hardening #
assert_valid_wallet_id(data.wallet)
assert_valid_timestamp_seconds(expires)
# ## #
return await db.fetchall( return await db.fetchall(
""" """
SELECT * FROM nwcprovider.keys SELECT * FROM nwcprovider.keys
@ -57,15 +87,23 @@ async def get_wallet_nwcs(data: GetWalletNWC) -> List[NWCKey]:
""", """,
{ {
"wallet": data.wallet, "wallet": data.wallet,
"expires": int(time.time()) if not data.include_expired else -1, "expires": expires,
}, },
model=NWCKey, model=NWCKey,
) )
async def get_nwc(data: GetNWC) -> Optional[NWCKey]: async def get_nwc(data: GetNWC) -> Optional[NWCKey]:
expires = int(time.time()) if not data.include_expired else -1
# hardening #
assert_valid_pubkey(data.pubkey)
assert_valid_timestamp_seconds(expires)
# ## #
# expires_at = 0 means it never expires # expires_at = 0 means it never expires
if data.wallet: if data.wallet:
assert_valid_wallet_id(data.wallet)
row = await db.fetchone( row = await db.fetchone(
""" """
SELECT * FROM nwcprovider.keys SELECT * FROM nwcprovider.keys
@ -75,7 +113,7 @@ async def get_nwc(data: GetNWC) -> Optional[NWCKey]:
{ {
"pubkey": data.pubkey, "pubkey": data.pubkey,
"wallet": data.wallet, "wallet": data.wallet,
"expires": int(time.time()) if not data.include_expired else -1, "expires": expires,
}, },
NWCKey, NWCKey,
) )
@ -87,7 +125,7 @@ async def get_nwc(data: GetNWC) -> Optional[NWCKey]:
""", """,
{ {
"pubkey": data.pubkey, "pubkey": data.pubkey,
"expires": int(time.time()) if not data.include_expired else -1, "expires": expires,
}, },
NWCKey, NWCKey,
) )
@ -105,6 +143,11 @@ async def get_nwc(data: GetNWC) -> Optional[NWCKey]:
async def get_budgets_nwc(data: GetBudgetsNWC) -> Optional[NWCBudget]: async def get_budgets_nwc(data: GetBudgetsNWC) -> Optional[NWCBudget]:
# hardening #
assert_valid_pubkey(data.pubkey)
# ## #
rows = await db.fetchall( rows = await db.fetchall(
"SELECT * FROM nwcprovider.budgets WHERE pubkey = :pubkey", "SELECT * FROM nwcprovider.budgets WHERE pubkey = :pubkey",
{"pubkey": data.pubkey}, {"pubkey": data.pubkey},
@ -113,6 +156,12 @@ async def get_budgets_nwc(data: GetBudgetsNWC) -> Optional[NWCBudget]:
if data.calculate_spent: if data.calculate_spent:
for budget in budgets: for budget in budgets:
last_cycle, next_cycle = budget.get_timestamp_range() last_cycle, next_cycle = budget.get_timestamp_range()
# hardening #
assert_valid_timestamp_seconds(last_cycle)
assert_valid_timestamp_seconds(next_cycle)
# ## #
tot_spent_in_range_msats = await db.fetchone( tot_spent_in_range_msats = await db.fetchone(
""" """
SELECT SUM(amount_msats) FROM nwcprovider.spent SELECT SUM(amount_msats) FROM nwcprovider.spent
@ -126,12 +175,23 @@ async def get_budgets_nwc(data: GetBudgetsNWC) -> Optional[NWCBudget]:
}, },
) )
tot_spent_in_range_msats = next(iter(tot_spent_in_range_msats.values())) or 0 tot_spent_in_range_msats = next(iter(tot_spent_in_range_msats.values())) or 0
# hardening #
assert_valid_msats(tot_spent_in_range_msats)
# ## #
budget.used_budget_msats = tot_spent_in_range_msats budget.used_budget_msats = tot_spent_in_range_msats
return budgets return budgets
async def tracked_spend_nwc(data: TrackedSpendNWC, action): async def tracked_spend_nwc(data: TrackedSpendNWC, action):
async def r(): async def r():
# hardening #
assert_valid_pubkey(data.pubkey)
assert_valid_msats(data.amount_msats)
# ## #
created_at = int(time.time()) created_at = int(time.time())
budgets = await get_budgets_nwc(GetBudgetsNWC( budgets = await get_budgets_nwc(GetBudgetsNWC(
pubkey=data.pubkey pubkey=data.pubkey
@ -139,6 +199,12 @@ async def tracked_spend_nwc(data: TrackedSpendNWC, action):
in_budget = True in_budget = True
for budget in budgets: for budget in budgets:
last_cycle, next_cycle = budget.get_timestamp_range() last_cycle, next_cycle = budget.get_timestamp_range()
# hardening #
assert_valid_timestamp_seconds(last_cycle)
assert_valid_timestamp_seconds(next_cycle)
# ## #
tot_spent_in_range_msats = ( tot_spent_in_range_msats = (
next(iter((await db.fetchone( next(iter((await db.fetchone(
""" """
@ -153,6 +219,12 @@ async def tracked_spend_nwc(data: TrackedSpendNWC, action):
}, },
)).values())) or 0 )).values())) or 0
) )
# hardening #
assert_valid_msats(tot_spent_in_range_msats)
assert_valid_msats(budget.budget_msats)
# ## #
if tot_spent_in_range_msats + data.amount_msats > budget.budget_msats: if tot_spent_in_range_msats + data.amount_msats > budget.budget_msats:
in_budget = False in_budget = False
break break
@ -185,6 +257,12 @@ async def get_config_nwc(key: str):
async def set_config_nwc(key: str, value: str): async def set_config_nwc(key: str, value: str):
# hardening #
assert_sane_string(key)
assert_sane_string(value)
# ## #
await db.execute( await db.execute(
""" """
INSERT OR REPLACE INTO nwcprovider.config (key, value) INSERT OR REPLACE INTO nwcprovider.config (key, value)

122
paranoia.py Normal file
View file

@ -0,0 +1,122 @@
# Run-time hardening to detect unexpected inputs at the edges
from loguru import logger
ENABLE_HARDENING = True
def panic(reason:str):
if not ENABLE_HARDENING:
return
logger.error(f"hardening error: {reason}")
raise ValueError(f"hardening error: {reason}")
# Throw if string contains any non-printable characters
def assert_printable(v:str):
if not ENABLE_HARDENING:
return
if not isinstance(v, str):
panic("not a string "+str(v))
if not v.isprintable():
panic("string contains non-printable characters")
# Check if number is valid int and not NaN
def assert_valid_int(v:int):
if not ENABLE_HARDENING:
return
if not isinstance(v, int):
panic("number is not a valid int")
# Check if number is valid positive int
def assert_valid_positive_int(v:int):
if not ENABLE_HARDENING:
return
assert_valid_int(v)
if v < 0:
panic("number is not positive")
# Check if number is a valid sats amount
def assert_valid_sats(v:int):
if not ENABLE_HARDENING:
return
assert_valid_positive_int(v)
max_sats_value = 10_000_000
if v >= max_sats_value:
panic("sats amount looks too high")
# Check if number is a valid msats amount
def assert_valid_msats(v:int):
if not ENABLE_HARDENING:
return
assert_valid_positive_int(v)
max_msats_value = 10_000_000 * 1000
if v >= max_msats_value:
panic("msats amount looks too high")
# Check if string is a valid sha256 hash
def assert_valid_sha256(v:str):
if not ENABLE_HARDENING:
return
assert_printable(v)
if len(v) != 64 or not all(c in "0123456789abcdef" for c in v):
panic("string is not a valid sha256 hash")
# Check if valid nostr pubkey
def assert_valid_pubkey(v:str):
if not ENABLE_HARDENING:
return
assert_valid_sha256(v)
# Check if valid wallet id
def assert_valid_wallet_id(v:str):
if not ENABLE_HARDENING:
return
assert_printable(v)
if not v.isalnum():
panic("string is not a valid wallet id")
# Check if valid timestamp in seconds
def assert_valid_timestamp_seconds(v:int):
if not ENABLE_HARDENING:
return
assert_valid_int(v)
if v < 0:
panic("timestamp is negative")
if v > 2**31:
panic("timestamp is too high")
# Check if string is within sane parameters
def assert_sane_string(v:str):
if not ENABLE_HARDENING:
return
assert_printable(v)
if len(v) > 1024:
panic("string is too long")
# Check if string is a non-empty string
def assert_non_empty_string(v:str):
if not ENABLE_HARDENING:
return
assert_printable(v)
if len(v.strip()) == 0:
panic("string is empty")
# Assert valid json
def assert_valid_json(v:str):
if not ENABLE_HARDENING:
return
assert_non_empty_string(v)
try:
import json
json.loads(v)
except:
panic("string is not valid json")
# Check if string is a valid bolt11 invoice
def assert_valid_bolt11(invoice:str):
if not ENABLE_HARDENING:
return
assert_printable(invoice)

View file

@ -22,6 +22,7 @@ from .execution_queue import execution_queue
from .models import NWCKey, TrackedSpendNWC, GetNWC from .models import NWCKey, TrackedSpendNWC, GetNWC
from .nwcp import NWCServiceProvider from .nwcp import NWCServiceProvider
from .permission import nwc_permissions from .permission import nwc_permissions
from .paranoia import assert_valid_wallet_id, assert_valid_pubkey, assert_sane_string, assert_valid_timestamp_seconds, assert_valid_msats, assert_valid_positive_int, assert_valid_bolt11, assert_valid_sha256
async def _check(nwc: Optional[NWCKey], method: str) -> Optional[Dict]: async def _check(nwc: Optional[NWCKey], method: str) -> Optional[Dict]:
@ -55,6 +56,16 @@ async def _process_invoice(
amount_msats: int, amount_msats: int,
description: Optional[str] = None, description: Optional[str] = None,
): ):
# hardening #
assert_valid_wallet_id(wallet_id)
assert_valid_pubkey(pubkey)
assert_valid_bolt11(invoice)
assert_valid_msats(amount_msats)
if description:
assert_sane_string(description)
# ## #
async def execute_payment() -> str : async def execute_payment() -> str :
payment = await pay_invoice( payment = await pay_invoice(
wallet_id=wallet_id, wallet_id=wallet_id,
@ -114,6 +125,11 @@ async def _on_pay_invoice(
pubkey: str, pubkey: str,
payload: Dict payload: Dict
) -> List[Tuple[Optional[Dict], Optional[Dict], List]]: ) -> List[Tuple[Optional[Dict], Optional[Dict], List]]:
# hardening #
assert_valid_pubkey(pubkey)
# ## #
nwc = await get_nwc(GetNWC( nwc = await get_nwc(GetNWC(
pubkey=pubkey, pubkey=pubkey,
refresh_last_used=True refresh_last_used=True
@ -130,6 +146,12 @@ async def _on_pay_invoice(
raise Exception("Missing invoice") raise Exception("Missing invoice")
invoice_data = bolt11_decode(invoice) invoice_data = bolt11_decode(invoice)
amount_msats = int(invoice_data.amount_msat or 0) amount_msats = int(invoice_data.amount_msat or 0)
# hardening #
assert_valid_bolt11(invoice)
assert_valid_msats(amount_msats)
# ## #
res = await _process_invoice( res = await _process_invoice(
nwc.wallet, pubkey, invoice, amount_msats, invoice_data.description nwc.wallet, pubkey, invoice, amount_msats, invoice_data.description
) )
@ -149,6 +171,11 @@ async def _on_multi_pay_invoice(
pubkey: str, pubkey: str,
payload: Dict payload: Dict
) -> List[Tuple[Optional[Dict], Optional[Dict], List]]: ) -> List[Tuple[Optional[Dict], Optional[Dict], List]]:
# hardening #
assert_valid_pubkey(pubkey)
# ## #
nwc = await get_nwc(GetNWC(pubkey=pubkey, refresh_last_used=True)) nwc = await get_nwc(GetNWC(pubkey=pubkey, refresh_last_used=True))
error = await _check(nwc, "multi_pay_invoice") error = await _check(nwc, "multi_pay_invoice")
if error: if error:
@ -171,6 +198,12 @@ async def _on_multi_pay_invoice(
invoice = i.get("invoice", None) invoice = i.get("invoice", None)
invoice_data = bolt11_decode(invoice) invoice_data = bolt11_decode(invoice)
amount_msats = int(invoice_data.amount_msat or 0) amount_msats = int(invoice_data.amount_msat or 0)
# hardening #
assert_valid_bolt11(invoice)
assert_valid_msats(amount_msats)
# ## #
res = await _process_invoice( res = await _process_invoice(
nwc.wallet, pubkey, invoice, amount_msats, invoice_data.description nwc.wallet, pubkey, invoice, amount_msats, invoice_data.description
) )
@ -197,6 +230,11 @@ async def _on_make_invoice(
pubkey: str, pubkey: str,
payload: Dict payload: Dict
) -> List[Tuple[Optional[Dict], Optional[Dict], List]]: ) -> List[Tuple[Optional[Dict], Optional[Dict], List]]:
# hardening #
assert_valid_pubkey(pubkey)
# ## #
nwc = await get_nwc(GetNWC(pubkey=pubkey, refresh_last_used=True)) nwc = await get_nwc(GetNWC(pubkey=pubkey, refresh_last_used=True))
error = await _check(nwc, "make_invoice") error = await _check(nwc, "make_invoice")
if error: if error:
@ -211,6 +249,17 @@ async def _on_make_invoice(
description = params.get("description", "") description = params.get("description", "")
description_hash = params.get("description_hash", None) description_hash = params.get("description_hash", None)
expiry = params.get("expiry", None) expiry = params.get("expiry", None)
# hardening #
assert_valid_msats(amount_msats)
if description:
assert_sane_string(description)
if description_hash:
assert_valid_sha256(description_hash)
if expiry:
assert_valid_timestamp_seconds(expiry)
# ## #
payment = await create_invoice( payment = await create_invoice(
wallet_id=nwc.wallet, wallet_id=nwc.wallet,
amount=int(amount_msats / 1000), amount=int(amount_msats / 1000),
@ -253,6 +302,11 @@ async def _on_lookup_invoice(
pubkey: str, pubkey: str,
payload: Dict payload: Dict
) -> List[Tuple[Optional[Dict], Optional[Dict], List]]: ) -> List[Tuple[Optional[Dict], Optional[Dict], List]]:
# hardening #
assert_valid_pubkey(pubkey)
# ## #
nwc = await get_nwc(GetNWC(pubkey=pubkey, refresh_last_used=True)) nwc = await get_nwc(GetNWC(pubkey=pubkey, refresh_last_used=True))
error = await _check(nwc, "lookup_invoice") error = await _check(nwc, "lookup_invoice")
if error: if error:
@ -269,6 +323,12 @@ async def _on_lookup_invoice(
if not payment_hash: if not payment_hash:
invoice_data = bolt11_decode(invoice) invoice_data = bolt11_decode(invoice)
payment_hash = invoice_data.payment_hash payment_hash = invoice_data.payment_hash
# hardening #
assert_valid_sha256(payment_hash)
assert_valid_bolt11(invoice)
# ## #
# Get payment data # Get payment data
payment = await get_wallet_payment(nwc.wallet, payment_hash) payment = await get_wallet_payment(nwc.wallet, payment_hash)
if not payment: if not payment:
@ -304,6 +364,10 @@ async def _on_list_transactions(
pubkey: str, pubkey: str,
payload: Dict payload: Dict
) -> List[Tuple[Optional[Dict], Optional[Dict], List]]: ) -> List[Tuple[Optional[Dict], Optional[Dict], List]]:
# hardening #
assert_valid_pubkey(pubkey)
# ## #
nwc = await get_nwc(GetNWC(pubkey=pubkey, refresh_last_used=True)) nwc = await get_nwc(GetNWC(pubkey=pubkey, refresh_last_used=True))
error = await _check(nwc, "list_transactions") error = await _check(nwc, "list_transactions")
if error: if error:
@ -316,6 +380,15 @@ async def _on_list_transactions(
offset = payload.get("offset", 0) offset = payload.get("offset", 0)
unpaid = payload.get("unpaid", False) unpaid = payload.get("unpaid", False)
tx_type = payload.get("type", None) tx_type = payload.get("type", None)
# hardening #
assert_valid_positive_int(tfrom)
assert_valid_positive_int(tto)
assert_valid_positive_int(limit)
assert_valid_positive_int(offset)
assert_sane_string(tx_type)
# ## #
values = [] values = []
filters: Filters = Filters() filters: Filters = Filters()
filters.where(["time <= ?"]) filters.where(["time <= ?"])
@ -363,6 +436,11 @@ async def _on_get_balance(
pubkey: str, pubkey: str,
payload: Dict payload: Dict
) -> List[Tuple[Optional[Dict], Optional[Dict], List]]: ) -> List[Tuple[Optional[Dict], Optional[Dict], List]]:
# hardening #
assert_valid_pubkey(pubkey)
# ## #
nwc = await get_nwc(GetNWC(pubkey=pubkey, refresh_last_used= True)) nwc = await get_nwc(GetNWC(pubkey=pubkey, refresh_last_used= True))
error = await _check(nwc, "get_balance") error = await _check(nwc, "get_balance")
if error: if error:
@ -383,6 +461,11 @@ async def _on_get_info(
pubkey: str, pubkey: str,
payload: Dict payload: Dict
) -> List[Tuple[Optional[Dict], Optional[Dict], List]]: ) -> List[Tuple[Optional[Dict], Optional[Dict], List]]:
# hardening #
assert_valid_pubkey(pubkey)
# ## #
nwc = await get_nwc(GetNWC(pubkey=pubkey, refresh_last_used=True)) nwc = await get_nwc(GetNWC(pubkey=pubkey, refresh_last_used=True))
error = await _check(nwc, "get_info") error = await _check(nwc, "get_info")
if error: if error:

View file

@ -28,6 +28,7 @@ from .models import (
GetBudgetsNWC GetBudgetsNWC
) )
from .permission import nwc_permissions from .permission import nwc_permissions
from .paranoia import assert_valid_wallet_id, assert_valid_pubkey, assert_sane_string
nwcprovider_api_router = APIRouter() nwcprovider_api_router = APIRouter()
@ -53,6 +54,11 @@ async def api_get_nwcs(
wallet: WalletTypeInfo = Depends(require_admin_key), wallet: WalletTypeInfo = Depends(require_admin_key),
): ):
wallet_id = wallet.wallet.id wallet_id = wallet.wallet.id
# hardening #
assert_valid_wallet_id(wallet_id)
# ## #
wallet_nwcs = GetWalletNWC( wallet_nwcs = GetWalletNWC(
wallet=wallet_id, wallet=wallet_id,
include_expired=include_expired include_expired=include_expired
@ -82,6 +88,12 @@ async def api_get_nwc(
wallet: WalletTypeInfo = Depends(require_admin_key) wallet: WalletTypeInfo = Depends(require_admin_key)
) -> NWCGetResponse: ) -> NWCGetResponse:
wallet_id = wallet.wallet.id wallet_id = wallet.wallet.id
# hardening #
assert_valid_pubkey(pubkey)
assert_valid_wallet_id(wallet_id)
# ## #
nwc = await get_nwc(GetNWC(pubkey=pubkey, wallet=wallet_id, include_expired=include_expired)) nwc = await get_nwc(GetNWC(pubkey=pubkey, wallet=wallet_id, include_expired=include_expired))
if not nwc: if not nwc:
raise Exception("Pubkey has no associated wallet") raise Exception("Pubkey has no associated wallet")
@ -103,6 +115,11 @@ async def api_get_pairing_url(
req: Request, req: Request,
secret: str secret: str
) -> str: ) -> str:
# hardening #
assert_sane_string(secret)
# ## #
pprivkey: Optional[str] = await get_config_nwc("provider_key") pprivkey: Optional[str] = await get_config_nwc("provider_key")
if not pprivkey: if not pprivkey:
raise Exception("Extension is not configured") raise Exception("Extension is not configured")
@ -147,6 +164,12 @@ async def api_register_nwc(
wallet: WalletTypeInfo = Depends(require_admin_key), wallet: WalletTypeInfo = Depends(require_admin_key),
): ):
wallet_id = wallet.wallet.id wallet_id = wallet.wallet.id
# hardening #
assert_valid_pubkey(pubkey)
assert_valid_wallet_id(wallet_id)
# ## #
nwc = await create_nwc( nwc = await create_nwc(
CreateNWCKey( CreateNWCKey(
pubkey=pubkey, pubkey=pubkey,
@ -176,6 +199,12 @@ async def api_delete_nwc(
wallet: WalletTypeInfo = Depends(require_admin_key) wallet: WalletTypeInfo = Depends(require_admin_key)
): ):
wallet_id = wallet.wallet.id wallet_id = wallet.wallet.id
# hardening #
assert_valid_pubkey(pubkey)
assert_valid_wallet_id(wallet_id)
# ## #
await delete_nwc(DeleteNWC( await delete_nwc(DeleteNWC(
pubkey=pubkey, pubkey=pubkey,
wallet=wallet_id wallet=wallet_id
@ -217,6 +246,13 @@ async def api_get_config_nwc(key: str):
) )
async def api_set_config_nwc(req: Request): async def api_set_config_nwc(req: Request):
data = await req.json() data = await req.json()
# hardening #
for key, value in data.items():
assert_sane_string(key)
assert_sane_string(value)
# ## #
for key, value in data.items(): for key, value in data.items():
await set_config_nwc(key, value) await set_config_nwc(key, value)
return await api_get_all_config_nwc() return await api_get_all_config_nwc()