From 383f6ec7219ea0a93ea0bbaa1f3907ebb5bcb30f Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?dni=20=E2=9A=A1?= Date: Thu, 26 Sep 2024 15:48:55 +0200 Subject: [PATCH] fix part1 --- lnbits/core/crud.py | 114 ++++++++------------- lnbits/core/migrations.py | 9 ++ lnbits/core/models.py | 59 +++++------ lnbits/db.py | 18 +++- lnbits/settings.py | 6 -- tests/regtest/test_real_invoice.py | 8 -- tests/regtest/test_services_pay_invoice.py | 3 +- 7 files changed, 94 insertions(+), 123 deletions(-) diff --git a/lnbits/core/crud.py b/lnbits/core/crud.py index 42e31b9ff..8bb93a027 100644 --- a/lnbits/core/crud.py +++ b/lnbits/core/crud.py @@ -11,7 +11,6 @@ from lnbits.core.extensions.models import ( InstallableExtension, PayToEnableInfo, UserExtension, - UserExtensionInfo, ) from lnbits.core.models import PaymentState from lnbits.db import DB_TYPE, SQLITE, Connection, Database, Filters, Page @@ -164,17 +163,14 @@ async def get_accounts( async def get_account( user_id: str, conn: Optional[Connection] = None ) -> Optional[User]: - row = await (conn or db).fetchone( + user = await (conn or db).fetchone( """ SELECT id, email, username, pubkey, created_at, updated_at, extra FROM accounts WHERE id = :id """, {"id": user_id}, + User, ) - - user = User(**row) if row else None - if user and row["extra"]: - user.config = UserConfig(**json.loads(row["extra"])) return user @@ -463,93 +459,77 @@ async def drop_extension_db(*, ext_id: str, conn: Optional[Connection] = None) - async def get_installed_extension( ext_id: str, conn: Optional[Connection] = None ) -> Optional[InstallableExtension]: - row = await (conn or db).fetchone( + extension = await (conn or db).fetchone( "SELECT * FROM installed_extensions WHERE id = :id", {"id": ext_id}, + InstallableExtension, ) - - return InstallableExtension.from_row(row) if row else None + return extension async def get_installed_extensions( active: Optional[bool] = None, conn: Optional[Connection] = None, ) -> list[InstallableExtension]: - rows = await (conn or db).fetchall( - "SELECT * FROM installed_extensions", + where = "WHERE active = :active" if active else "" + values = {"active": active} if active else {} + all_extensions = await (conn or db).fetchall( + f"SELECT * FROM installed_extensions {where}", + values, + model=InstallableExtension, ) - all_extensions = [InstallableExtension.from_row(row) for row in rows] - if active is None: - return all_extensions - - return [e for e in all_extensions if e.active == active] + return all_extensions async def get_user_extension( user_id: str, extension: str, conn: Optional[Connection] = None ) -> Optional[UserExtension]: - row = await (conn or db).fetchone( + return await (conn or db).fetchone( """ - SELECT extension, active, extra as _extra FROM extensions + SELECT * FROM extensions WHERE "user" = :user AND extension = :ext """, {"user": user_id, "ext": extension}, + model=UserExtension, ) - return UserExtension.from_row(row) if row else None async def get_user_extensions( user_id: str, conn: Optional[Connection] = None ) -> list[UserExtension]: - rows = await (conn or db).fetchall( - """ - SELECT extension, active, extra as _extra FROM extensions - WHERE "user" = :user - """, + return await (conn or db).fetchall( + """SELECT * FROM extensions WHERE "user" = :user""", {"user": user_id}, + model=UserExtension, ) - return [UserExtension.from_row(row) for row in rows] async def update_user_extension( - *, user_id: str, extension: str, active: bool, conn: Optional[Connection] = None + user_extension: UserExtension, conn: Optional[Connection] = None ) -> None: - await (conn or db).execute( - """ - INSERT INTO extensions ("user", extension, active) VALUES (:user, :ext, :active) - ON CONFLICT ("user", extension) DO UPDATE SET active = :active - """, - {"user": user_id, "ext": extension, "active": active}, - ) + where = """extension = :extension AND "user" = :user""" + await (conn or db).update("extensions", user_extension, where) + # await (conn or db).execute( + # """ + # INSERT INTO extensions ("user", extension, active) + # VALUES (:user, :ext, :active) + # ON CONFLICT ("user", extension) DO UPDATE SET active = :active + # """, + # {"user": user_id, "ext": extension, "active": active}, + # ) async def get_user_active_extensions_ids( user_id: str, conn: Optional[Connection] = None ) -> list[str]: - rows = await (conn or db).fetchall( + exts = await (conn or db).fetchall( """ - SELECT extension FROM extensions WHERE "user" = :user AND active + SELECT * FROM extensions WHERE "user" = :user AND active """, {"user": user_id}, + UserExtension, ) - return [e.get("extension", "") for e in rows] - - -async def update_user_extension_extra( - user_id: str, - extension: str, - extra: UserExtensionInfo, - conn: Optional[Connection] = None, -) -> None: - extra_json = json.dumps(dict(extra)) - await (conn or db).execute( - """ - INSERT INTO extensions ("user", extension, extra) VALUES - (:user, :ext, :extra) - ON CONFLICT ("user", extension) DO UPDATE SET extra = :extra - """, - {"user": user_id, "ext": extension, "extra": extra_json}, - ) + return [ext.extension for ext in exts] # wallets @@ -754,31 +734,29 @@ async def get_standalone_payment( row = await (conn or db).fetchone( f""" - SELECT * - FROM apipayments + SELECT * FROM apipayments WHERE {clause} - ORDER BY amount - LIMIT 1 + ORDER BY amount LIMIT 1 """, values, + Payment, ) - - return Payment.from_row(row) if row else None + return row async def get_wallet_payment( wallet_id: str, payment_hash: str, conn: Optional[Connection] = None ) -> Optional[Payment]: - row = await (conn or db).fetchone( + payment = await (conn or db).fetchone( """ SELECT * FROM apipayments WHERE wallet = :wallet AND hash = :hash """, {"wallet": wallet_id, "hash": payment_hash}, + Payment, ) - - return Payment.from_row(row) if row else None + return payment async def get_latest_payments_by_extension(ext_name: str, ext_id: str, limit: int = 5): @@ -1305,24 +1283,22 @@ async def delete_tinyurl(tinyurl_id: str): async def get_webpush_subscription( endpoint: str, user: str ) -> Optional[WebPushSubscription]: - row = await db.fetchone( + return await db.fetchone( """ SELECT * FROM webpush_subscriptions WHERE endpoint = :endpoint AND "user" = :user """, {"endpoint": endpoint, "user": user}, + WebPushSubscription, ) - return WebPushSubscription(**dict(row)) if row else None -async def get_webpush_subscriptions_for_user( - user: str, -) -> list[WebPushSubscription]: - rows = await db.fetchall( +async def get_webpush_subscriptions_for_user(user: str) -> list[WebPushSubscription]: + return await db.fetchall( """SELECT * FROM webpush_subscriptions WHERE "user" = :user""", {"user": user}, + WebPushSubscription, ) - return [WebPushSubscription(**dict(row)) for row in rows] async def create_webpush_subscription( diff --git a/lnbits/core/migrations.py b/lnbits/core/migrations.py index 24d09056a..d79e072e2 100644 --- a/lnbits/core/migrations.py +++ b/lnbits/core/migrations.py @@ -553,3 +553,12 @@ async def m022_add_pubkey_to_accounts(db): await db.execute("ALTER TABLE accounts ADD COLUMN pubkey TEXT") except OperationalError: pass + + +async def m023_add_column_column_to_apipayments(db): + """ + renames hash to payment_hash and drops unused index + """ + await db.execute("DROP INDEX by_hash") + await db.execute("ALTER TABLE apipayments RENAME COLUMN hash TO payment_hash") + await db.execute("ALTER TABLE apipayments RENAME COLUMN wallet TO wallet_id") diff --git a/lnbits/core/models.py b/lnbits/core/models.py index 916109554..227a627fc 100644 --- a/lnbits/core/models.py +++ b/lnbits/core/models.py @@ -3,7 +3,6 @@ from __future__ import annotations import datetime import hashlib import hmac -import json import time from dataclasses import dataclass from enum import Enum @@ -90,7 +89,7 @@ class WalletTypeInfo: wallet: Wallet -class UserConfig(BaseModel): +class UserExtra(BaseModel): email_verified: Optional[bool] = False first_name: Optional[str] = None last_name: Optional[str] = None @@ -144,7 +143,7 @@ class User(BaseModel): admin: bool = False super_user: bool = False has_password: bool = False - config: Optional[UserConfig] = None + extra: Optional[UserExtra] = None created_at: Optional[int] = None updated_at: Optional[int] = None @@ -178,7 +177,7 @@ class UpdateUser(BaseModel): user_id: str email: Optional[str] = Query(default=None) username: Optional[str] = Query(default=..., min_length=2, max_length=20) - config: Optional[UserConfig] = None + extra: Optional[UserExtra] = None class UpdateUserPassword(BaseModel): @@ -244,23 +243,38 @@ class CreatePayment(BaseModel): fee: int = 0 +# class Extra(BaseModel): +# _raw_json: str + +# @property +# def _json(self): +# return json.loads(self._raw_json) + +# class PaymentExtra(Extra): +# tag: Optional[str] = None +# def __getitem__(self, key): +# return self[key] or self._raw_json[key] + + class Payment(BaseModel): status: str - # TODO should be removed in the future, backward compatibility - pending: bool checking_id: str + payment_hash: str + wallet_id: str amount: int fee: int memo: Optional[str] time: int bolt11: str - preimage: str - payment_hash: str expiry: Optional[float] extra: Optional[dict] - wallet_id: str webhook: Optional[str] - webhook_status: Optional[int] + webhook_status: Optional[int] = None + preimage: Optional[str] = "0" * 64 + + @property + def pending(self) -> bool: + return self.status == PaymentState.PENDING.value @property def success(self) -> bool: @@ -270,27 +284,6 @@ class Payment(BaseModel): def failed(self) -> bool: return self.status == PaymentState.FAILED.value - @classmethod - def from_row(cls, row: dict): - return cls( - checking_id=row["checking_id"], - payment_hash=row["hash"] or "0" * 64, - bolt11=row["bolt11"] or "", - preimage=row["preimage"] or "0" * 64, - extra=json.loads(row["extra"] or "{}"), - status=row["status"], - # TODO should be removed in the future, backward compatibility - pending=row["status"] == PaymentState.PENDING.value, - amount=row["amount"], - fee=row["fee"], - memo=row["memo"], - time=row["time"], - expiry=row["expiry"], - wallet_id=row["wallet"], - webhook=row["webhook"], - webhook_status=row["webhook_status"], - ) - @property def tag(self) -> Optional[str]: if self.extra is None: @@ -377,10 +370,6 @@ class TinyURL(BaseModel): wallet: str time: float - @classmethod - def from_row(cls, row: dict): - return cls(**dict(row)) - class ConversionData(BaseModel): from_: str = "sat" diff --git a/lnbits/db.py b/lnbits/db.py index fe20b6c8b..eb7a2078f 100644 --- a/lnbits/db.py +++ b/lnbits/db.py @@ -328,13 +328,23 @@ class Database(Compat): finally: self.lock.release() - async def fetchall(self, query: str, values: Optional[dict] = None) -> list[dict]: + async def fetchall( + self, + query: str, + values: Optional[dict] = None, + model: Optional[type[TModel]] = None, + ) -> list[TModel]: async with self.connect() as conn: - return await conn.fetchall(query, values) + return await conn.fetchall(query, values, model) - async def fetchone(self, query: str, values: Optional[dict] = None) -> dict: + async def fetchone( + self, + query: str, + values: Optional[dict] = None, + model: Optional[type[TModel]] = None, + ) -> TModel: async with self.connect() as conn: - return await conn.fetchone(query, values) + return await conn.fetchone(query, values, model) async def insert(self, table_name: str, model: BaseModel) -> None: async with self.connect() as conn: diff --git a/lnbits/settings.py b/lnbits/settings.py index f4aeb07cf..34d2e2230 100644 --- a/lnbits/settings.py +++ b/lnbits/settings.py @@ -7,7 +7,6 @@ import json from enum import Enum from hashlib import sha256 from os import path -from sqlite3 import Row from time import time from typing import Any, Optional @@ -635,11 +634,6 @@ class ReadOnlySettings( class Settings(EditableSettings, ReadOnlySettings, TransientSettings, BaseSettings): - @classmethod - def from_row(cls, row: Row) -> Settings: - data = dict(row) - return cls(**data) - class Config: env_file = ".env" env_file_encoding = "utf-8" diff --git a/tests/regtest/test_real_invoice.py b/tests/regtest/test_real_invoice.py index 6ecb9bae2..9daa54964 100644 --- a/tests/regtest/test_real_invoice.py +++ b/tests/regtest/test_real_invoice.py @@ -143,7 +143,6 @@ async def test_pay_real_invoice_set_pending_and_check_state( payment = await get_standalone_payment(invoice["payment_hash"]) assert payment assert payment.success - assert payment.pending is False @pytest.mark.asyncio @@ -167,7 +166,6 @@ async def test_pay_hold_invoice_check_pending( payment_db = await get_standalone_payment(invoice_obj.payment_hash) assert payment_db - assert payment_db.pending is True settle_invoice(preimage) @@ -181,7 +179,6 @@ async def test_pay_hold_invoice_check_pending( payment_db_after_settlement = await get_standalone_payment(invoice_obj.payment_hash) assert payment_db_after_settlement - assert payment_db_after_settlement.pending is False @pytest.mark.asyncio @@ -205,7 +202,6 @@ async def test_pay_hold_invoice_check_pending_and_fail( payment_db = await get_standalone_payment(invoice_obj.payment_hash) assert payment_db - assert payment_db.pending is True preimage_hash = hashlib.sha256(bytes.fromhex(preimage)).hexdigest() @@ -221,7 +217,6 @@ async def test_pay_hold_invoice_check_pending_and_fail( # payment should be in database as failed payment_db_after_settlement = await get_standalone_payment(invoice_obj.payment_hash) assert payment_db_after_settlement - assert payment_db_after_settlement.pending is False assert payment_db_after_settlement.failed is True @@ -246,7 +241,6 @@ async def test_pay_hold_invoice_check_pending_and_fail_cancel_payment_task_in_me payment_db = await get_standalone_payment(invoice_obj.payment_hash) assert payment_db - assert payment_db.pending is True # cancel payment task, this simulates the client dropping the connection task.cancel() @@ -307,7 +301,6 @@ async def test_receive_real_invoice_set_pending_and_check_state( assert payment_status["paid"] assert payment - assert payment.pending is False # set the incoming invoice to pending await update_payment_details(payment.checking_id, status=PaymentState.PENDING) @@ -316,7 +309,6 @@ async def test_receive_real_invoice_set_pending_and_check_state( invoice["payment_hash"], incoming=True ) assert payment_pending - assert payment_pending.pending is True assert payment_pending.success is False assert payment_pending.failed is False diff --git a/tests/regtest/test_services_pay_invoice.py b/tests/regtest/test_services_pay_invoice.py index 849d87872..d24d01653 100644 --- a/tests/regtest/test_services_pay_invoice.py +++ b/tests/regtest/test_services_pay_invoice.py @@ -5,6 +5,7 @@ from lnbits.core.crud import ( ) from lnbits.core.services import ( PaymentError, + PaymentState, pay_invoice, ) @@ -21,7 +22,7 @@ async def test_services_pay_invoice(to_wallet, real_invoice): assert payment_hash payment = await get_standalone_payment(payment_hash) assert payment - assert not payment.pending + assert not payment.status == PaymentState.SUCCESS assert payment.memo == description