diff --git a/lnbits/app.py b/lnbits/app.py index 3854c173c..6ea992c74 100644 --- a/lnbits/app.py +++ b/lnbits/app.py @@ -62,7 +62,6 @@ from .middleware import ( ) from .requestvars import g from .tasks import ( - check_pending_payments, create_task, internal_invoice_listener, invoice_listener, @@ -410,7 +409,7 @@ def register_async_tasks(app: FastAPI): if not settings.lnbits_extensions_deactivate_all: create_task(check_and_register_extensions(app)) - create_permanent_task(check_pending_payments) + # create_permanent_task(check_pending_payments) create_permanent_task(invoice_listener) create_permanent_task(internal_invoice_listener) create_permanent_task(cache.invalidate_forever) diff --git a/lnbits/core/crud.py b/lnbits/core/crud.py index 39fbad93b..b09fa0cdf 100644 --- a/lnbits/core/crud.py +++ b/lnbits/core/crud.py @@ -1,7 +1,7 @@ import datetime import json from time import time -from typing import Any, Dict, List, Literal, Optional +from typing import Dict, List, Literal, Optional from uuid import uuid4 import shortuuid @@ -54,11 +54,18 @@ async def create_account( extra = json.dumps(dict(user_config)) if user_config else "{}" now = int(time()) await (conn or db).execute( - f""" + """ INSERT INTO accounts (id, username, pass, email, extra, created_at, updated_at) - VALUES (?, ?, ?, ?, ?, {db.timestamp_placeholder}, {db.timestamp_placeholder}) + VALUES (:user, :username, :password, :email, :extra, :now, :now) """, - (user_id, username, password, email, extra, now, now), + { + "user": user_id, + "username": username, + "password": password, + "email": email, + "extra": extra, + "now": now, + }, ) new_account = await get_account(user_id=user_id, conn=conn) @@ -92,18 +99,18 @@ async def update_account( now = int(time()) await db.execute( - f""" + """ UPDATE accounts SET (username, email, extra, updated_at) = - (?, ?, ?, {db.timestamp_placeholder}) - WHERE id = ? + (:username, :email, :extra, :now) + WHERE id = :user """, - ( - username, - email, - json.dumps(dict(extra)) if extra else "{}", - now, - user_id, - ), + { + "username": username, + "email": email, + "extra": json.dumps(dict(extra)) if extra else "{}", + "now": now, + "user": user_id, + }, ) user = await get_user(user_id) @@ -113,8 +120,8 @@ async def update_account( async def delete_account(user_id: str, conn: Optional[Connection] = None) -> None: await (conn or db).execute( - "DELETE from accounts WHERE id = ?", - (user_id,), + "DELETE from accounts WHERE id = :user", + {"user": user_id}, ) @@ -144,7 +151,7 @@ async def get_accounts( FROM accounts LEFT JOIN wallets ON accounts.id = wallets.user """, [], - [], + {}, filters=filters, model=Account, group_by=["accounts.id"], @@ -157,9 +164,9 @@ async def get_account( row = await (conn or db).fetchone( """ SELECT id, email, username, created_at, updated_at, extra - FROM accounts WHERE id = ? + FROM accounts WHERE id = :id """, - (user_id,), + {"id": user_id}, ) user = User(**row) if row else None @@ -174,26 +181,23 @@ async def delete_accounts_no_wallets( ) -> None: delta = int(time()) - time_delta await (conn or db).execute( - f""" + """ DELETE FROM accounts WHERE NOT EXISTS ( SELECT wallets.id FROM wallets WHERE wallets.user = accounts.id ) AND ( - (updated_at is null AND created_at < {db.timestamp_placeholder}) - OR updated_at < {db.timestamp_placeholder} + (updated_at is null AND created_at < :delta) + OR updated_at < :delta ) """, - ( - delta, - delta, - ), + {"delta": delta}, ) async def get_user_password(user_id: str) -> Optional[str]: row = await db.fetchone( - "SELECT pass FROM accounts WHERE id = ?", - (user_id,), + "SELECT pass FROM accounts WHERE id = :user", + {"user": user_id}, ) if not row: return None @@ -201,6 +205,7 @@ async def get_user_password(user_id: str) -> Optional[str]: return row[0] +# TODO: refactor not a crud function async def verify_user_password(user_id: str, password: str) -> bool: existing_password = await get_user_password(user_id) if not existing_password: @@ -210,7 +215,7 @@ async def verify_user_password(user_id: str, password: str) -> bool: return pwd_context.verify(password, existing_password) -# todo: , conn: Optional[Connection] = None ?? +# TODO: , conn: Optional[Connection] = None ??, maybe also not a crud function async def update_user_password(data: UpdateUserPassword) -> Optional[User]: assert data.password == data.password_repeat, "Passwords do not match." @@ -224,15 +229,15 @@ async def update_user_password(data: UpdateUserPassword) -> Optional[User]: now = int(time()) await db.execute( - f""" - UPDATE accounts SET pass = ?, updated_at = {db.timestamp_placeholder} - WHERE id = ? + """ + UPDATE accounts SET pass = :pass, updated_at = :now + WHERE id = :user """, - ( - pwd_context.hash(data.password), - now, - data.user_id, - ), + { + "pass": pwd_context.hash(data.password), + "now": now, + "user": data.user_id, + }, ) user = await get_user(data.user_id) @@ -246,9 +251,9 @@ async def get_account_by_username( row = await (conn or db).fetchone( """ SELECT id, username, email, created_at, updated_at - FROM accounts WHERE username = ? + FROM accounts WHERE username = :username """, - (username,), + {"username": username}, ) return User(**row) if row else None @@ -260,9 +265,9 @@ async def get_account_by_email( row = await (conn or db).fetchone( """ SELECT id, username, email, created_at, updated_at - FROM accounts WHERE email = ? + FROM accounts WHERE email = :email """, - (email,), + {"email": email}, ) return User(**row) if row else None @@ -281,9 +286,9 @@ async def get_user(user_id: str, conn: Optional[Connection] = None) -> Optional[ user = await (conn or db).fetchone( """ SELECT id, email, username, pass, extra, created_at, updated_at - FROM accounts WHERE id = ? + FROM accounts WHERE id = :id """, - (user_id,), + {"id": user_id}, ) if user: @@ -294,9 +299,9 @@ async def get_user(user_id: str, conn: Optional[Connection] = None) -> Optional[ SELECT balance FROM balances WHERE wallet = wallets.id ), 0) AS balance_msat FROM wallets - WHERE "user" = ? and wallets.deleted = false + WHERE "user" = :user and wallets.deleted = false """, - (user_id,), + {"user": user_id}, ) else: return None @@ -340,27 +345,22 @@ async def add_installed_extension( """ INSERT INTO installed_extensions (id, version, name, active, short_description, icon, stars, meta) - VALUES (?, ?, ?, ?, ?, ?, ?, ?) ON CONFLICT (id) DO UPDATE SET + VALUES + (:ext, :version, :name, :active, :short_description, :icon, :stars, :meta) + ON CONFLICT (id) DO UPDATE SET (version, name, active, short_description, icon, stars, meta) = - (?, ?, ?, ?, ?, ?, ?) + (:version, :name, :active, :short_description, :icon, :stars, :meta) """, - ( - ext.id, - version, - ext.name, - ext.active, - ext.short_description, - ext.icon, - ext.stars, - json.dumps(meta), - version, - ext.name, - ext.active, - ext.short_description, - ext.icon, - ext.stars, - json.dumps(meta), - ), + { + "ext": ext.id, + "version": version, + "name": ext.name, + "active": ext.active, + "short_description": ext.short_description, + "icon": ext.icon, + "stars": ext.stars, + "meta": json.dumps(meta), + }, ) @@ -369,9 +369,9 @@ async def update_installed_extension_state( ) -> None: await (conn or db).execute( """ - UPDATE installed_extensions SET active = ? WHERE id = ? + UPDATE installed_extensions SET active = :active WHERE id = :ext """, - (active, ext_id), + {"ext": ext_id, "active": active}, ) @@ -391,15 +391,16 @@ async def delete_installed_extension( ) -> None: await (conn or db).execute( """ - DELETE from installed_extensions WHERE id = ? + DELETE from installed_extensions WHERE id = :ext """, - (ext_id,), + {"ext": ext_id}, ) async def drop_extension_db(*, ext_id: str, conn: Optional[Connection] = None) -> None: db_version = await (conn or db).fetchone( - "SELECT * FROM dbversions WHERE db = ?", (ext_id,) + "SELECT * FROM dbversions WHERE db = :id", + {"id": ext_id}, ) # Check that 'ext_id' is a valid extension id and not a malicious string assert db_version, f"Extension '{ext_id}' db version cannot be found" @@ -412,7 +413,6 @@ async def drop_extension_db(*, ext_id: str, conn: Optional[Connection] = None) - # The `ext_id` value is verified above. await (conn or db).execute( f"DROP SCHEMA IF EXISTS {ext_id} CASCADE", - (), ) @@ -420,8 +420,8 @@ async def get_installed_extension( ext_id: str, conn: Optional[Connection] = None ) -> Optional[InstallableExtension]: row = await (conn or db).fetchone( - "SELECT * FROM installed_extensions WHERE id = ?", - (ext_id,), + "SELECT * FROM installed_extensions WHERE id = :id", + {"id": ext_id}, ) return InstallableExtension.from_row(row) if row else None @@ -433,7 +433,6 @@ async def get_installed_extensions( ) -> List["InstallableExtension"]: rows = await (conn or db).fetchall( "SELECT * FROM installed_extensions", - (), ) all_extensions = [InstallableExtension.from_row(row) for row in rows] if active is None: @@ -448,9 +447,9 @@ async def get_user_extension( row = await (conn or db).fetchone( """ SELECT extension, active, extra as _extra FROM extensions - WHERE "user" = ? AND extension = ? + WHERE "user" = :user AND extension = :ext """, - (user_id, extension), + {"user": user_id, "ext": extension}, ) return UserExtension.from_row(row) if row else None @@ -461,9 +460,9 @@ async def get_user_extensions( rows = await (conn or db).fetchall( """ SELECT extension, active, extra as _extra FROM extensions - WHERE "user" = ? + WHERE "user" = :user """, - (user_id,), + {"user": user_id}, ) return [UserExtension.from_row(row) for row in rows] @@ -473,10 +472,10 @@ async def update_user_extension( ) -> None: await (conn or db).execute( """ - INSERT INTO extensions ("user", extension, active) VALUES (?, ?, ?) - ON CONFLICT ("user", extension) DO UPDATE SET active = ? + INSERT INTO extensions ("user", extension, active) VALUES (:user, :ext, :active) + ON CONFLICT ("user", extension) DO UPDATE SET active = :active """, - (user_id, extension, active, active), + {"user": user_id, "ext": extension, "active": active}, ) @@ -484,8 +483,10 @@ async def get_user_active_extensions_ids( user_id: str, conn: Optional[Connection] = None ) -> List[str]: rows = await (conn or db).fetchall( - """SELECT extension FROM extensions WHERE "user" = ? AND active""", - (user_id,), + """ + SELECT extension FROM extensions WHERE "user" = :user AND active + """, + {"user": user_id}, ) return [e[0] for e in rows] @@ -499,10 +500,11 @@ async def update_user_extension_extra( extra_json = json.dumps(dict(extra)) await (conn or db).execute( """ - INSERT INTO extensions ("user", extension, extra) VALUES (?, ?, ?) - ON CONFLICT ("user", extension) DO UPDATE SET extra = ? + INSERT INTO extensions ("user", extension, extra) VALUES + (:user, :ext, :extra) + ON CONFLICT ("user", extension) DO UPDATE SET extra = :extra """, - (user_id, extension, extra_json, extra_json), + {"user": user_id, "ext": extension, "extra": extra_json}, ) @@ -519,19 +521,18 @@ async def create_wallet( wallet_id = uuid4().hex now = int(time()) await (conn or db).execute( - f""" + """ INSERT INTO wallets (id, name, "user", adminkey, inkey, created_at, updated_at) - VALUES (?, ?, ?, ?, ?, {db.timestamp_placeholder}, {db.timestamp_placeholder}) + VALUES (:wallet, :name, :user, :adminkey, :inkey, :now, :now) """, - ( - wallet_id, - wallet_name or settings.lnbits_default_wallet_name, - user_id, - uuid4().hex, - uuid4().hex, - now, - now, - ), + { + "wallet": wallet_id, + "name": wallet_name or settings.lnbits_default_wallet_name, + "user": user_id, + "adminkey": uuid4().hex, + "inkey": uuid4().hex, + "now": now, + }, ) new_wallet = await get_wallet(wallet_id=wallet_id, conn=conn) @@ -547,22 +548,22 @@ async def update_wallet( conn: Optional[Connection] = None, ) -> Optional[Wallet]: set_clause = [] - values: list = [] - set_clause.append(f"updated_at = {db.timestamp_placeholder}") - now = int(time()) - values.append(now) + set_clause.append("updated_at = :now") + values: dict = { + "wallet": wallet_id, + "now": int(time()), + } if name: - set_clause.append("name = ?") - values.append(name) + set_clause.append("name = :name") + values["name"] = name if currency is not None: - set_clause.append("currency = ?") - values.append(currency) - values.append(wallet_id) + set_clause.append("currency = :currency") + values["currency"] = currency await (conn or db).execute( f""" - UPDATE wallets SET {', '.join(set_clause)} WHERE id = ? + UPDATE wallets SET {', '.join(set_clause)} WHERE id = :wallet """, - tuple(values), + values, ) wallet = await get_wallet(wallet_id=wallet_id, conn=conn) assert wallet, "updated created wallet couldn't be retrieved" @@ -578,12 +579,12 @@ async def delete_wallet( ) -> None: now = int(time()) await (conn or db).execute( - f""" + """ UPDATE wallets - SET deleted = ?, updated_at = {db.timestamp_placeholder} - WHERE id = ? AND "user" = ? + SET deleted = :deleted, updated_at = :now + WHERE id = :wallet AND "user" = :user """, - (deleted, now, wallet_id, user_id), + {"wallet": wallet_id, "user": user_id, "deleted": deleted, "now": now}, ) @@ -591,8 +592,8 @@ async def force_delete_wallet( wallet_id: str, conn: Optional[Connection] = None ) -> None: await (conn or db).execute( - "DELETE FROM wallets WHERE id = ?", - (wallet_id,), + "DELETE FROM wallets WHERE id = :wallet", + {"wallet": wallet_id}, ) @@ -601,12 +602,12 @@ async def delete_wallet_by_id( ) -> Optional[int]: now = int(time()) result = await (conn or db).execute( - f""" + """ UPDATE wallets - SET deleted = true, updated_at = {db.timestamp_placeholder} - WHERE id = ? + SET deleted = true, updated_at = :now + WHERE id = :wallet """, - (now, wallet_id), + {"wallet": wallet_id, "now": now}, ) return result.rowcount @@ -621,19 +622,16 @@ async def delete_unused_wallets( ) -> None: delta = int(time()) - time_delta await (conn or db).execute( - f""" + """ DELETE FROM wallets WHERE ( SELECT COUNT(*) FROM apipayments WHERE wallet = wallets.id ) = 0 AND ( - (updated_at is null AND created_at < {db.timestamp_placeholder}) - OR updated_at < {db.timestamp_placeholder} + (updated_at is null AND created_at < :delta) + OR updated_at < :delta ) """, - ( - delta, - delta, - ), + {"delta": delta}, ) @@ -643,9 +641,9 @@ async def get_wallet( row = await (conn or db).fetchone( """ SELECT *, COALESCE((SELECT balance FROM balances WHERE wallet = wallets.id), 0) - AS balance_msat FROM wallets WHERE id = ? + AS balance_msat FROM wallets WHERE id = :wallet """, - (wallet_id,), + {"wallet": wallet_id}, ) return Wallet(**row) if row else None @@ -655,9 +653,9 @@ async def get_wallets(user_id: str, conn: Optional[Connection] = None) -> List[W rows = await (conn or db).fetchall( """ SELECT *, COALESCE((SELECT balance FROM balances WHERE wallet = wallets.id), 0) - AS balance_msat FROM wallets WHERE "user" = ? + AS balance_msat FROM wallets WHERE "user" = :user """, - (user_id,), + {"user": user_id}, ) return [Wallet(**row) for row in rows] @@ -671,9 +669,9 @@ async def get_wallet_for_key( """ SELECT *, COALESCE((SELECT balance FROM balances WHERE wallet = wallets.id), 0) AS balance_msat FROM wallets - WHERE (adminkey = ? OR inkey = ?) AND deleted = false + WHERE (adminkey = :key OR inkey = :key) AND deleted = false """, - (key, key), + {"key": key}, ) if not row: @@ -697,14 +695,17 @@ async def get_standalone_payment( incoming: Optional[bool] = False, wallet_id: Optional[str] = None, ) -> Optional[Payment]: - clause: str = "checking_id = ? OR hash = ?" - values = [checking_id_or_hash, checking_id_or_hash] + clause: str = "checking_id = :checking_id OR hash = :hash" + values = { + "wallet": wallet_id, + "checking_id": checking_id_or_hash, + "hash": checking_id_or_hash, + } if incoming: clause = f"({clause}) AND amount > 0" if wallet_id: - clause = f"({clause}) AND wallet = ?" - values.append(wallet_id) + clause = f"({clause}) AND wallet = :wallet" row = await (conn or db).fetchone( f""" @@ -714,7 +715,7 @@ async def get_standalone_payment( ORDER BY amount LIMIT 1 """, - tuple(values), + values, ) return Payment.from_row(row) if row else None @@ -727,9 +728,9 @@ async def get_wallet_payment( """ SELECT * FROM apipayments - WHERE wallet = ? AND hash = ? + WHERE wallet = :wallet AND hash = :hash """, - (wallet_id, payment_hash), + {"wallet": wallet_id, "hash": payment_hash}, ) return Payment.from_row(row) if row else None @@ -740,14 +741,11 @@ async def get_latest_payments_by_extension(ext_name: str, ext_id: str, limit: in f""" SELECT * FROM apipayments WHERE status = '{PaymentState.SUCCESS}' - AND extra LIKE ? - AND extra LIKE ? + AND extra LIKE :ext_name + AND extra LIKE :ext_id ORDER BY time DESC LIMIT {limit} """, - ( - f"%{ext_name}%", - f"%{ext_id}%", - ), + {"ext_name": f"%{ext_name}%", "ext_id": f"%{ext_id}%"}, ) return rows @@ -769,16 +767,17 @@ async def get_payments_paginated( Filters payments to be returned by complete | pending | outgoing | incoming. """ - values: List[Any] = [] + values: dict = { + "wallet": wallet_id, + "time": since, + } clause: List[str] = [] if since is not None: - clause.append(f"time > {db.timestamp_placeholder}") - values.append(since) + clause.append("time > :time") if wallet_id: - clause.append("wallet = ?") - values.append(wallet_id) + clause.append("wallet = :wallet") if complete and pending: pass @@ -857,20 +856,23 @@ async def delete_expired_invoices( conn: Optional[Connection] = None, ) -> None: # first we delete all invoices older than one month + await (conn or db).execute( f""" DELETE FROM apipayments WHERE status = '{PaymentState.PENDING}' AND amount > 0 - AND time < {db.timestamp_now} - {db.interval_seconds(2592000)} - """ + AND time < :delta + """, + {"delta": int(time() - 2592000)}, ) # then we delete all invoices whose expiry date is in the past await (conn or db).execute( f""" DELETE FROM apipayments WHERE status = '{PaymentState.PENDING}' AND amount > 0 - AND expiry < {db.timestamp_now} - """ + AND expiry < :now + """, + {"now": int(time())}, ) @@ -904,27 +906,28 @@ async def create_payment( INSERT INTO apipayments (wallet, checking_id, bolt11, hash, preimage, amount, status, memo, fee, extra, webhook, expiry, pending) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + VALUES (:wallet, :checking_id, :bolt11, :hash, :preimage, + :amount, :status, :memo, :fee, :extra, :webhook, :expiry, :pending) """, - ( - wallet_id, - checking_id, - payment_request, - payment_hash, - preimage, - amount, - status.value, - memo, - fee, - ( + { + "wallet": wallet_id, + "checking_id": checking_id, + "bolt11": payment_request, + "hash": payment_hash, + "preimage": preimage, + "amount": amount, + "status": status.value, + "memo": memo, + "fee": fee, + "extra": ( json.dumps(extra) if extra and extra != {} and isinstance(extra, dict) else None ), - webhook, - db.datetime_to_timestamp(expiry) if expiry else None, - False, # TODO: remove this in next release - ), + "webhook": webhook, + "expiry": db.datetime_to_timestamp(expiry) if expiry else None, + "pending": False, # TODO: remove this in next release + }, ) new_payment = await get_wallet_payment(wallet_id, payment_hash, conn=conn) @@ -937,8 +940,8 @@ async def update_payment_status( checking_id: str, status: PaymentState, conn: Optional[Connection] = None ) -> None: await (conn or db).execute( - "UPDATE apipayments SET status = ? WHERE checking_id = ?", - (status.value, checking_id), + "UPDATE apipayments SET status = :status WHERE checking_id = :checking_id", + {"status": status.value, "checking_id": checking_id}, ) @@ -950,27 +953,30 @@ async def update_payment_details( new_checking_id: Optional[str] = None, conn: Optional[Connection] = None, ) -> None: + set_variables: dict = { + "checking_id": checking_id, + "new_checking_id": new_checking_id, + "status": status.value if status else None, + "fee": fee, + "preimage": preimage, + } + set_clause: List[str] = [] - set_variables: List[Any] = [] - if new_checking_id is not None: - set_clause.append("checking_id = ?") - set_variables.append(new_checking_id) + set_clause.append("checking_id = :checking_id") if status is not None: - set_clause.append("status = ?") - set_variables.append(status.value) + set_clause.append("status = :status") if fee is not None: - set_clause.append("fee = ?") - set_variables.append(fee) + set_clause.append("fee = :fee") if preimage is not None: - set_clause.append("preimage = ?") - set_variables.append(preimage) - - set_variables.append(checking_id) + set_clause.append("preimage = :preimage") await (conn or db).execute( - f"UPDATE apipayments SET {', '.join(set_clause)} WHERE checking_id = ?", - tuple(set_variables), + f""" + UPDATE apipayments SET {', '.join(set_clause)} + WHERE checking_id = :checking_id + """, + set_variables, ) @@ -989,8 +995,8 @@ async def update_payment_extra( amount_clause = "AND amount < 0" if outgoing else "AND amount > 0" row = await (conn or db).fetchone( - f"SELECT hash, extra from apipayments WHERE hash = ? {amount_clause}", - (payment_hash,), + f"SELECT hash, extra from apipayments WHERE hash = :hash {amount_clause}", + {"hash": payment_hash}, ) if not row: return @@ -998,8 +1004,8 @@ async def update_payment_extra( db_extra.update(extra) await (conn or db).execute( - f"UPDATE apipayments SET extra = ? WHERE hash = ? {amount_clause} ", - (json.dumps(db_extra), payment_hash), + f"UPDATE apipayments SET extra = :extra WHERE hash = :hash {amount_clause} ", + {"extra": json.dumps(db_extra), "hash": payment_hash}, ) @@ -1019,10 +1025,11 @@ async def get_payments_history( if not filters: filters = Filters() where = [f"(status = '{PaymentState.SUCCESS}' OR amount < 0)"] - values = [] + values: dict = { + "wallet": wallet_id, + } if wallet_id: - where.append("wallet = ?") - values.append(wallet_id) + where.append("wallet = :wallet") if DB_TYPE == SQLITE and group in sqlite_formats: date_trunc = f"strftime('{sqlite_formats[group]}', time, 'unixepoch')" @@ -1070,8 +1077,8 @@ async def delete_wallet_payment( checking_id: str, wallet_id: str, conn: Optional[Connection] = None ) -> None: await (conn or db).execute( - "DELETE FROM apipayments WHERE checking_id = ? AND wallet = ?", - (checking_id, wallet_id), + "DELETE FROM apipayments WHERE checking_id = :checking_id AND wallet = :wallet", + {"checking_id": checking_id, "wallet": wallet_id}, ) @@ -1085,9 +1092,9 @@ async def check_internal( row = await (conn or db).fetchone( f""" SELECT checking_id FROM apipayments - WHERE hash = ? AND status = '{PaymentState.PENDING}' AND amount > 0 + WHERE hash = :hash AND status = '{PaymentState.PENDING}' AND amount > 0 """, - (payment_hash,), + {"hash": payment_hash}, ) if not row: return None @@ -1105,9 +1112,9 @@ async def check_internal_pending( row = await (conn or db).fetchone( """ SELECT status FROM apipayments - WHERE hash = ? AND amount > 0 + WHERE hash = :hash AND amount > 0 """, - (payment_hash,), + {"hash": payment_hash}, ) if not row: return True @@ -1117,10 +1124,10 @@ async def check_internal_pending( async def mark_webhook_sent(payment_hash: str, status: int) -> None: await db.execute( """ - UPDATE apipayments SET webhook_status = ? - WHERE hash = ? + UPDATE apipayments SET webhook_status = :status + WHERE hash = :hash """, - (status, payment_hash), + {"status": status, "hash": payment_hash}, ) @@ -1161,20 +1168,29 @@ async def update_admin_settings(data: EditableSettings) -> None: editable_settings = json.loads(row["editable_settings"]) if row else {} editable_settings.update(data.dict(exclude_unset=True)) await db.execute( - "UPDATE settings SET editable_settings = ?", (json.dumps(editable_settings),) + "UPDATE settings SET editable_settings = :settings", + {"settings": json.dumps(editable_settings)}, ) async def update_super_user(super_user: str) -> SuperSettings: - await db.execute("UPDATE settings SET super_user = ?", (super_user,)) + await db.execute( + "UPDATE settings SET super_user = :user", + {"user": super_user}, + ) settings = await get_super_settings() assert settings, "updated super_user settings could not be retrieved" return settings async def create_admin_settings(super_user: str, new_settings: dict): - sql = "INSERT INTO settings (super_user, editable_settings) VALUES (?, ?)" - await db.execute(sql, (super_user, json.dumps(new_settings))) + await db.execute( + """ + INSERT INTO settings (super_user, editable_settings) + VALUES (:user, :settings) + """, + {"user": super_user, "settings": json.dumps(new_settings)}, + ) settings = await get_super_settings() assert settings, "created admin settings could not be retrieved" return settings @@ -1190,19 +1206,19 @@ async def get_dbversions(conn: Optional[Connection] = None): async def update_migration_version(conn, db_name, version): await (conn or db).execute( """ - INSERT INTO dbversions (db, version) VALUES (?, ?) - ON CONFLICT (db) DO UPDATE SET version = ? + INSERT INTO dbversions (db, version) VALUES (:db, :version) + ON CONFLICT (db) DO UPDATE SET version = :version """, - (db_name, version, version), + {"db": db_name, "version": version}, ) async def delete_dbversion(*, ext_id: str, conn: Optional[Connection] = None) -> None: await (conn or db).execute( """ - DELETE FROM dbversions WHERE db = ? + DELETE FROM dbversions WHERE db = :ext """, - (ext_id,), + {"ext": ext_id}, ) @@ -1213,37 +1229,35 @@ async def delete_dbversion(*, ext_id: str, conn: Optional[Connection] = None) -> async def create_tinyurl(domain: str, endless: bool, wallet: str): tinyurl_id = shortuuid.uuid()[:8] await db.execute( - "INSERT INTO tiny_url (id, url, endless, wallet) VALUES (?, ?, ?, ?)", - ( - tinyurl_id, - domain, - endless, - wallet, - ), + """ + INSERT INTO tiny_url (id, url, endless, wallet) + VALUES (:tinyurl, :domain, :endless, :wallet) + """, + {"tinyurl": tinyurl_id, "domain": domain, "endless": endless, "wallet": wallet}, ) return await get_tinyurl(tinyurl_id) async def get_tinyurl(tinyurl_id: str) -> Optional[TinyURL]: row = await db.fetchone( - "SELECT * FROM tiny_url WHERE id = ?", - (tinyurl_id,), + "SELECT * FROM tiny_url WHERE id = :tinyurl", + {"tinyurl": tinyurl_id}, ) return TinyURL.from_row(row) if row else None async def get_tinyurl_by_url(url: str) -> List[TinyURL]: rows = await db.fetchall( - "SELECT * FROM tiny_url WHERE url = ?", - (url,), + "SELECT * FROM tiny_url WHERE url = :url", + {"url": url}, ) return [TinyURL.from_row(row) for row in rows] async def delete_tinyurl(tinyurl_id: str): await db.execute( - "DELETE FROM tiny_url WHERE id = ?", - (tinyurl_id,), + "DELETE FROM tiny_url WHERE id = :tinyurl", + {"tinyurl": tinyurl_id}, ) @@ -1261,8 +1275,10 @@ async def get_webpush_settings() -> Optional[WebPushSettings]: async def create_webpush_settings(webpush_settings: dict): await db.execute( - "INSERT INTO webpush_settings (vapid_keypair) VALUES (?)", - (json.dumps(webpush_settings),), + "INSERT INTO webpush_settings (vapid_keypair) VALUES (:vapid_keypair)", + { + "vapid_keypair": json.dumps(webpush_settings), + }, ) return await get_webpush_settings() @@ -1271,11 +1287,11 @@ async def get_webpush_subscription( endpoint: str, user: str ) -> Optional[WebPushSubscription]: row = await db.fetchone( - """SELECT * FROM webpush_subscriptions WHERE endpoint = ? AND "user" = ?""", - ( - endpoint, - user, - ), + """ + SELECT * FROM webpush_subscriptions + WHERE endpoint = :endpoint AND "user" = :user + """, + {"endpoint": endpoint, "user": user}, ) return WebPushSubscription(**dict(row)) if row else None @@ -1284,8 +1300,8 @@ async def get_webpush_subscriptions_for_user( user: str, ) -> List[WebPushSubscription]: rows = await db.fetchall( - """SELECT * FROM webpush_subscriptions WHERE "user" = ?""", - (user,), + """SELECT * FROM webpush_subscriptions WHERE "user" = :user""", + {"user": user}, ) return [WebPushSubscription(**dict(row)) for row in rows] @@ -1296,14 +1312,9 @@ async def create_webpush_subscription( await db.execute( """ INSERT INTO webpush_subscriptions (endpoint, "user", data, host) - VALUES (?, ?, ?, ?) + VALUES (:endpoint, :user, :data, :host) """, - ( - endpoint, - user, - data, - host, - ), + {"endpoint": endpoint, "user": user, "data": data, "host": host}, ) subscription = await get_webpush_subscription(endpoint, user) assert subscription, "Newly created webpush subscription couldn't be retrieved" @@ -1312,17 +1323,17 @@ async def create_webpush_subscription( async def delete_webpush_subscription(endpoint: str, user: str) -> int: resp = await db.execute( - """DELETE FROM webpush_subscriptions WHERE endpoint = ? AND "user" = ?""", - ( - endpoint, - user, - ), + """ + DELETE FROM webpush_subscriptions WHERE endpoint = :endpoint AND "user" = :user + """, + {"endpoint": endpoint, "user": user}, ) return resp.rowcount async def delete_webpush_subscriptions(endpoint: str) -> int: resp = await db.execute( - "DELETE FROM webpush_subscriptions WHERE endpoint = ?", (endpoint,) + "DELETE FROM webpush_subscriptions WHERE endpoint = :endpoint", + {"endpoint": endpoint}, ) return resp.rowcount diff --git a/lnbits/db.py b/lnbits/db.py index e71bc0b83..d2baca4ed 100644 --- a/lnbits/db.py +++ b/lnbits/db.py @@ -144,36 +144,31 @@ class Connection(Compat): query = query.replace("?", "%s") return query - def rewrite_values(self, values): + def rewrite_values(self, values: dict) -> dict: # strip html clean_regex = re.compile("<.*?>|&([a-z0-9]+|#[0-9]{1,6}|#x[0-9a-f]{1,6});") - - # tuple to list and back to tuple - raw_values = [values] if isinstance(values, str) else list(values) - values = [] - for raw_value in raw_values: + clean_values: dict = {} + for key, raw_value in values.items(): if isinstance(raw_value, str): - values.append(re.sub(clean_regex, "", raw_value)) + clean_values[key] = re.sub(clean_regex, "", raw_value) elif isinstance(raw_value, datetime.datetime): ts = raw_value.timestamp() if self.type == SQLITE: - values.append(int(ts)) + clean_values[key] = int(ts) else: - values.append(ts) + clean_values[key] = ts else: - values.append(raw_value) - return tuple(values) + clean_values[key] = raw_value + return clean_values - async def fetchall(self, query: str, values: tuple = ()) -> list: - result = await self.conn.execute( - text(self.rewrite_query(query)), self.rewrite_values(values) - ) + async def fetchall(self, query: str, values: Optional[dict] = None) -> list: + params = self.rewrite_values(values) if values else {} + result = await self.conn.execute(text(self.rewrite_query(query)), params) return result.fetchall() - async def fetchone(self, query: str, values: tuple = ()): - result = await self.conn.execute( - text(self.rewrite_query(query)), self.rewrite_values(values) - ) + async def fetchone(self, query: str, values: Optional[dict] = None): + params = self.rewrite_values(values) if values else {} + result = await self.conn.execute(text(self.rewrite_query(query)), params) row = result.fetchone() result.close() return row @@ -182,7 +177,7 @@ class Connection(Compat): self, query: str, where: Optional[list[str]] = None, - values: Optional[list[str]] = None, + values: Optional[dict] = None, filters: Optional[Filters] = None, model: Optional[type[TRowModel]] = None, group_by: Optional[list[str]] = None, @@ -235,10 +230,9 @@ class Connection(Compat): total=count, ) - async def execute(self, query: str, values: tuple = ()): - return await self.conn.execute( - text(self.rewrite_query(query)), self.rewrite_values(values) - ) + async def execute(self, query: str, values: Optional[dict] = None): + params = self.rewrite_values(values) if values else {} + return await self.conn.execute(text(self.rewrite_query(query)), params) class Database(Compat): @@ -280,21 +274,23 @@ class Database(Compat): if self.schema: if self.type in {POSTGRES, COCKROACH}: await wconn.execute( - f"CREATE SCHEMA IF NOT EXISTS {self.schema}" + f"CREATE SCHEMA IF NOT EXISTS {self.schema}", {} ) elif self.type == SQLITE: - await wconn.execute(f"ATTACH '{self.path}' AS {self.schema}") + await wconn.execute( + f"ATTACH '{self.path}' AS {self.schema}", {} + ) yield wconn finally: self.lock.release() - async def fetchall(self, query: str, values: tuple = ()) -> list: + async def fetchall(self, query: str, values: Optional[dict] = None) -> list: async with self.connect() as conn: result = await conn.execute(query, values) return result.fetchall() - async def fetchone(self, query: str, values: tuple = ()): + async def fetchone(self, query: str, values: Optional[dict] = None): async with self.connect() as conn: result = await conn.execute(query, values) row = result.fetchone() @@ -305,7 +301,7 @@ class Database(Compat): self, query: str, where: Optional[list[str]] = None, - values: Optional[list[str]] = None, + values: Optional[dict] = None, filters: Optional[Filters] = None, model: Optional[type[TRowModel]] = None, group_by: Optional[list[str]] = None, @@ -313,7 +309,7 @@ class Database(Compat): async with self.connect() as conn: return await conn.fetch_page(query, where, values, filters, model, group_by) - async def execute(self, query: str, values: tuple = ()): + async def execute(self, query: str, values: Optional[dict] = None): async with self.connect() as conn: return await conn.execute(query, values) @@ -394,9 +390,8 @@ class Page(BaseModel, Generic[T]): class Filter(BaseModel, Generic[TFilterModel]): field: str op: Operator = Operator.EQ - values: list[Any] - model: Optional[type[TFilterModel]] + values: Optional[dict] = None @classmethod def parse_query(cls, key: str, raw_values: list[Any], model: type[TFilterModel]): @@ -415,27 +410,25 @@ class Filter(BaseModel, Generic[TFilterModel]): if field in model.__fields__: compare_field = model.__fields__[field] - values = [] + values: dict = {} for raw_value in raw_values: validated, errors = compare_field.validate(raw_value, {}, loc="none") if errors: raise ValidationError(errors=[errors], model=model) - values.append(validated) + values[field](validated) else: raise ValueError("Unknown filter field") return cls(field=field, op=op, values=values, model=model) - @property - def statement(self): - assert self.model, "Model is required for statement generation" - placeholder = get_placeholder(self.model, self.field) - if self.op in (Operator.INCLUDE, Operator.EXCLUDE): - placeholders = ", ".join([placeholder] * len(self.values)) - stmt = [f"{self.field} {self.op.as_sql} ({placeholders})"] - else: - stmt = [f"{self.field} {self.op.as_sql} {placeholder}"] * len(self.values) - return " OR ".join(stmt) + # @property + # def statement(self): + # if self.op in (Operator.INCLUDE, Operator.EXCLUDE): + # placeholders = ", ".join([placeholder] * len(self.values)) + # stmt = [f"{self.field} {self.op.as_sql} ({placeholders})"] + # else: + # stmt = [f"{self.field} {self.op.as_sql} {placeholder}"] * len(self.values) + # return " OR ".join(stmt) class Filters(BaseModel, Generic[TFilterModel]): @@ -481,18 +474,15 @@ class Filters(BaseModel, Generic[TFilterModel]): def where(self, where_stmts: Optional[list[str]] = None) -> str: if not where_stmts: where_stmts = [] - if self.filters: - for page_filter in self.filters: - where_stmts.append(page_filter.statement) + # if self.filters: + # for page_filter in self.filters: + # where_stmts.append(page_filter.statement) if self.search and self.model: + fields = self.model.__search_fields__ if DB_TYPE == POSTGRES: - where_stmts.append( - f"lower(concat({', '.join(self.model.__search_fields__)})) LIKE ?" - ) + where_stmts.append(f"lower(concat({', '.join(fields)})) LIKE :search") elif DB_TYPE == SQLITE: - where_stmts.append( - f"lower({'||'.join(self.model.__search_fields__)}) LIKE ?" - ) + where_stmts.append(f"lower({'||'.join(fields)}) LIKE :search") if where_stmts: return "WHERE " + " AND ".join(where_stmts) return "" @@ -502,12 +492,14 @@ class Filters(BaseModel, Generic[TFilterModel]): return f"ORDER BY {self.sortby} {self.direction or 'asc'}" return "" - def values(self, values: Optional[list[str]] = None) -> tuple: + def values(self, values: Optional[dict] = None) -> dict: if not values: - values = [] + values = {} if self.filters: for page_filter in self.filters: - values.extend(page_filter.values) + if page_filter.values: + for key, value in page_filter.values.items(): + values[key] = value if self.search and self.model: - values.append(f"%{self.search}%") - return tuple(values) + values["search"] = f"%{self.search}%" + return values