diff --git a/crud.py b/crud.py index 400b5d1..08d29e7 100644 --- a/crud.py +++ b/crud.py @@ -15,11 +15,28 @@ from .models import ( GetNWC, 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") 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( pubkey=data.pubkey, wallet=data.wallet, @@ -43,6 +60,12 @@ async def create_nwc(data: CreateNWCKey) -> NWCKey: async def delete_nwc(data: DeleteNWC) -> None: + + # hardening # + assert_valid_pubkey(data.pubkey) + assert_valid_wallet_id(data.wallet) + # ## # + await db.execute( "DELETE FROM nwcprovider.keys WHERE pubkey = :pubkey AND wallet = :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]: + 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( """ SELECT * FROM nwcprovider.keys @@ -57,15 +87,23 @@ async def get_wallet_nwcs(data: GetWalletNWC) -> List[NWCKey]: """, { "wallet": data.wallet, - "expires": int(time.time()) if not data.include_expired else -1, + "expires": expires, }, model=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 if data.wallet: + assert_valid_wallet_id(data.wallet) row = await db.fetchone( """ SELECT * FROM nwcprovider.keys @@ -75,7 +113,7 @@ async def get_nwc(data: GetNWC) -> Optional[NWCKey]: { "pubkey": data.pubkey, "wallet": data.wallet, - "expires": int(time.time()) if not data.include_expired else -1, + "expires": expires, }, NWCKey, ) @@ -87,7 +125,7 @@ async def get_nwc(data: GetNWC) -> Optional[NWCKey]: """, { "pubkey": data.pubkey, - "expires": int(time.time()) if not data.include_expired else -1, + "expires": expires, }, NWCKey, ) @@ -105,6 +143,11 @@ async def get_nwc(data: GetNWC) -> Optional[NWCKey]: async def get_budgets_nwc(data: GetBudgetsNWC) -> Optional[NWCBudget]: + + # hardening # + assert_valid_pubkey(data.pubkey) + # ## # + rows = await db.fetchall( "SELECT * FROM nwcprovider.budgets WHERE pubkey = :pubkey", {"pubkey": data.pubkey}, @@ -113,6 +156,12 @@ async def get_budgets_nwc(data: GetBudgetsNWC) -> Optional[NWCBudget]: if data.calculate_spent: for budget in budgets: 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( """ 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 + + # hardening # + assert_valid_msats(tot_spent_in_range_msats) + # ## # + budget.used_budget_msats = tot_spent_in_range_msats return budgets async def tracked_spend_nwc(data: TrackedSpendNWC, action): async def r(): + + # hardening # + assert_valid_pubkey(data.pubkey) + assert_valid_msats(data.amount_msats) + # ## # + created_at = int(time.time()) budgets = await get_budgets_nwc(GetBudgetsNWC( pubkey=data.pubkey @@ -139,6 +199,12 @@ async def tracked_spend_nwc(data: TrackedSpendNWC, action): in_budget = True for budget in budgets: 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 = ( next(iter((await db.fetchone( """ @@ -153,6 +219,12 @@ async def tracked_spend_nwc(data: TrackedSpendNWC, action): }, )).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: in_budget = False break @@ -185,6 +257,12 @@ async def get_config_nwc(key: str): async def set_config_nwc(key: str, value: str): + + # hardening # + assert_sane_string(key) + assert_sane_string(value) + # ## # + await db.execute( """ INSERT OR REPLACE INTO nwcprovider.config (key, value) diff --git a/paranoia.py b/paranoia.py new file mode 100644 index 0000000..4e56e93 --- /dev/null +++ b/paranoia.py @@ -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) + diff --git a/tasks.py b/tasks.py index 9ca9bf7..4b89913 100644 --- a/tasks.py +++ b/tasks.py @@ -22,6 +22,7 @@ from .execution_queue import execution_queue from .models import NWCKey, TrackedSpendNWC, GetNWC from .nwcp import NWCServiceProvider 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]: @@ -55,6 +56,16 @@ async def _process_invoice( amount_msats: int, 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 : payment = await pay_invoice( wallet_id=wallet_id, @@ -114,6 +125,11 @@ async def _on_pay_invoice( pubkey: str, payload: Dict ) -> List[Tuple[Optional[Dict], Optional[Dict], List]]: + + # hardening # + assert_valid_pubkey(pubkey) + # ## # + nwc = await get_nwc(GetNWC( pubkey=pubkey, refresh_last_used=True @@ -130,6 +146,12 @@ async def _on_pay_invoice( raise Exception("Missing invoice") invoice_data = bolt11_decode(invoice) amount_msats = int(invoice_data.amount_msat or 0) + + # hardening # + assert_valid_bolt11(invoice) + assert_valid_msats(amount_msats) + # ## # + res = await _process_invoice( nwc.wallet, pubkey, invoice, amount_msats, invoice_data.description ) @@ -149,6 +171,11 @@ async def _on_multi_pay_invoice( pubkey: str, payload: Dict ) -> List[Tuple[Optional[Dict], Optional[Dict], List]]: + + # hardening # + assert_valid_pubkey(pubkey) + # ## # + nwc = await get_nwc(GetNWC(pubkey=pubkey, refresh_last_used=True)) error = await _check(nwc, "multi_pay_invoice") if error: @@ -171,6 +198,12 @@ async def _on_multi_pay_invoice( invoice = i.get("invoice", None) invoice_data = bolt11_decode(invoice) amount_msats = int(invoice_data.amount_msat or 0) + + # hardening # + assert_valid_bolt11(invoice) + assert_valid_msats(amount_msats) + # ## # + res = await _process_invoice( nwc.wallet, pubkey, invoice, amount_msats, invoice_data.description ) @@ -197,6 +230,11 @@ async def _on_make_invoice( pubkey: str, payload: Dict ) -> List[Tuple[Optional[Dict], Optional[Dict], List]]: + + # hardening # + assert_valid_pubkey(pubkey) + # ## # + nwc = await get_nwc(GetNWC(pubkey=pubkey, refresh_last_used=True)) error = await _check(nwc, "make_invoice") if error: @@ -211,6 +249,17 @@ async def _on_make_invoice( description = params.get("description", "") description_hash = params.get("description_hash", 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( wallet_id=nwc.wallet, amount=int(amount_msats / 1000), @@ -253,6 +302,11 @@ async def _on_lookup_invoice( pubkey: str, payload: Dict ) -> List[Tuple[Optional[Dict], Optional[Dict], List]]: + + # hardening # + assert_valid_pubkey(pubkey) + # ## # + nwc = await get_nwc(GetNWC(pubkey=pubkey, refresh_last_used=True)) error = await _check(nwc, "lookup_invoice") if error: @@ -269,6 +323,12 @@ async def _on_lookup_invoice( if not payment_hash: invoice_data = bolt11_decode(invoice) payment_hash = invoice_data.payment_hash + + # hardening # + assert_valid_sha256(payment_hash) + assert_valid_bolt11(invoice) + # ## # + # Get payment data payment = await get_wallet_payment(nwc.wallet, payment_hash) if not payment: @@ -304,6 +364,10 @@ async def _on_list_transactions( pubkey: str, payload: Dict ) -> List[Tuple[Optional[Dict], Optional[Dict], List]]: + # hardening # + assert_valid_pubkey(pubkey) + # ## # + nwc = await get_nwc(GetNWC(pubkey=pubkey, refresh_last_used=True)) error = await _check(nwc, "list_transactions") if error: @@ -316,6 +380,15 @@ async def _on_list_transactions( offset = payload.get("offset", 0) unpaid = payload.get("unpaid", False) 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 = [] filters: Filters = Filters() filters.where(["time <= ?"]) @@ -363,6 +436,11 @@ async def _on_get_balance( pubkey: str, payload: Dict ) -> List[Tuple[Optional[Dict], Optional[Dict], List]]: + + # hardening # + assert_valid_pubkey(pubkey) + # ## # + nwc = await get_nwc(GetNWC(pubkey=pubkey, refresh_last_used= True)) error = await _check(nwc, "get_balance") if error: @@ -383,6 +461,11 @@ async def _on_get_info( pubkey: str, payload: Dict ) -> List[Tuple[Optional[Dict], Optional[Dict], List]]: + + # hardening # + assert_valid_pubkey(pubkey) + # ## # + nwc = await get_nwc(GetNWC(pubkey=pubkey, refresh_last_used=True)) error = await _check(nwc, "get_info") if error: diff --git a/views_api.py b/views_api.py index 2bd68fb..fcd755e 100644 --- a/views_api.py +++ b/views_api.py @@ -28,6 +28,7 @@ from .models import ( GetBudgetsNWC ) from .permission import nwc_permissions +from .paranoia import assert_valid_wallet_id, assert_valid_pubkey, assert_sane_string nwcprovider_api_router = APIRouter() @@ -53,6 +54,11 @@ async def api_get_nwcs( wallet: WalletTypeInfo = Depends(require_admin_key), ): wallet_id = wallet.wallet.id + + # hardening # + assert_valid_wallet_id(wallet_id) + # ## # + wallet_nwcs = GetWalletNWC( wallet=wallet_id, include_expired=include_expired @@ -82,6 +88,12 @@ async def api_get_nwc( wallet: WalletTypeInfo = Depends(require_admin_key) ) -> NWCGetResponse: 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)) if not nwc: raise Exception("Pubkey has no associated wallet") @@ -103,6 +115,11 @@ async def api_get_pairing_url( req: Request, secret: str ) -> str: + + # hardening # + assert_sane_string(secret) + # ## # + pprivkey: Optional[str] = await get_config_nwc("provider_key") if not pprivkey: raise Exception("Extension is not configured") @@ -147,6 +164,12 @@ async def api_register_nwc( wallet: WalletTypeInfo = Depends(require_admin_key), ): wallet_id = wallet.wallet.id + + # hardening # + assert_valid_pubkey(pubkey) + assert_valid_wallet_id(wallet_id) + # ## # + nwc = await create_nwc( CreateNWCKey( pubkey=pubkey, @@ -176,6 +199,12 @@ async def api_delete_nwc( wallet: WalletTypeInfo = Depends(require_admin_key) ): wallet_id = wallet.wallet.id + + # hardening # + assert_valid_pubkey(pubkey) + assert_valid_wallet_id(wallet_id) + # ## # + await delete_nwc(DeleteNWC( pubkey=pubkey, wallet=wallet_id @@ -217,6 +246,13 @@ async def api_get_config_nwc(key: str): ) async def api_set_config_nwc(req: Request): 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(): await set_config_nwc(key, value) return await api_get_all_config_nwc()