From 28897ca120e114dfd6de749befdb025cd6d5ebed Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?dni=20=E2=9A=A1?= Date: Thu, 3 Oct 2024 11:17:32 +0200 Subject: [PATCH] fixup! --- lnbits/core/crud.py | 15 +++++++-------- lnbits/core/models.py | 33 ++++++++++++++------------------- lnbits/core/services.py | 7 +++---- lnbits/core/views/api.py | 7 ++++--- lnbits/core/views/generic.py | 5 ++--- lnbits/core/views/user_api.py | 4 ++-- lnbits/core/views/wallet_api.py | 7 +++---- lnbits/db.py | 17 ++++++++--------- 8 files changed, 43 insertions(+), 52 deletions(-) diff --git a/lnbits/core/crud.py b/lnbits/core/crud.py index 6eb2725d2..371a9d100 100644 --- a/lnbits/core/crud.py +++ b/lnbits/core/crud.py @@ -31,7 +31,6 @@ from .models import ( TinyURL, User, Wallet, - WalletBalance, WebPushSubscription, ) @@ -340,7 +339,7 @@ async def create_wallet( adminkey=uuid4().hex, inkey=uuid4().hex, ) - await (conn or db).update("wallets", wallet) + await (conn or db).insert("wallets", wallet) return wallet @@ -420,7 +419,7 @@ async def delete_unused_wallets( async def get_wallet( wallet_id: str, deleted: Optional[bool] = None, conn: Optional[Connection] = None -) -> Optional[WalletBalance]: +) -> Optional[Wallet]: where = "AND deleted = :deleted" if deleted is not None else "" return await (conn or db).fetchone( f""" @@ -430,13 +429,13 @@ async def get_wallet( WHERE id = :wallet {where} """, {"wallet": wallet_id, "deleted": deleted}, - WalletBalance, + Wallet, ) async def get_wallets( user_id: str, deleted: Optional[bool] = None, conn: Optional[Connection] = None -) -> list[WalletBalance]: +) -> list[Wallet]: where = "AND deleted = :deleted" if deleted is not None else "" return await (conn or db).fetchall( f""" @@ -446,14 +445,14 @@ async def get_wallets( WHERE "user" = :user {where} """, {"user": user_id, "deleted": deleted}, - WalletBalance, + Wallet, ) async def get_wallet_for_key( key: str, conn: Optional[Connection] = None, -) -> Optional[WalletBalance]: +) -> Optional[Wallet]: return await (conn or db).fetchone( """ SELECT *, COALESCE(( @@ -463,7 +462,7 @@ async def get_wallet_for_key( WHERE (adminkey = :key OR inkey = :key) AND deleted = false """, {"key": key}, - WalletBalance, + Wallet, ) diff --git a/lnbits/core/models.py b/lnbits/core/models.py index 2093e6df0..9b31ece76 100644 --- a/lnbits/core/models.py +++ b/lnbits/core/models.py @@ -11,7 +11,7 @@ from typing import Callable, Optional from ecdsa import SECP256k1, SigningKey from fastapi import Query from passlib.context import CryptContext -from pydantic import BaseModel, validator +from pydantic import BaseModel, Field, validator from lnbits.db import FilterModel from lnbits.helpers import url_for @@ -45,6 +45,17 @@ class Wallet(BaseModel): created_at: datetime = datetime.now(timezone.utc) updated_at: datetime = datetime.now(timezone.utc) currency: Optional[str] = None + balance_msat: int = Field(default=0, no_database=True) + + @property + def balance(self) -> int: + return self.balance_msat // 1000 + + @property + def withdrawable_balance(self) -> int: + from .services import fee_reserve + + return self.balance_msat - fee_reserve(self.balance_msat) @property def lnurlwithdraw_full(self) -> str: @@ -63,22 +74,6 @@ class Wallet(BaseModel): ) -class WalletBalance(Wallet): - """Wallet with balance properties""" - - balance_msat: int = 0 - - @property - def balance(self) -> int: - return self.balance_msat // 1000 - - @property - def withdrawable_balance(self) -> int: - from .services import fee_reserve - - return self.balance_msat - fee_reserve(self.balance_msat) - - class KeyType(Enum): admin = 0 invoice = 1 @@ -92,7 +87,7 @@ class KeyType(Enum): @dataclass class WalletTypeInfo: key_type: KeyType - wallet: WalletBalance + wallet: Wallet class UserExtra(BaseModel): @@ -174,7 +169,7 @@ class User(BaseModel): username: Optional[str] = None pubkey: Optional[str] = None extensions: list[str] = [] - wallets: list[WalletBalance] = [] + wallets: list[Wallet] = [] admin: bool = False super_user: bool = False has_password: bool = False diff --git a/lnbits/core/services.py b/lnbits/core/services.py index c35a87ad4..a5bb3b54a 100644 --- a/lnbits/core/services.py +++ b/lnbits/core/services.py @@ -78,7 +78,7 @@ from .models import ( PaymentState, User, UserExtra, - WalletBalance, + Wallet, ) @@ -452,7 +452,7 @@ async def _create_external_payment( def _check_wallet_balance( - wallet: WalletBalance, + wallet: Wallet, fee_reserve_total_msat: int, internal_checking_id: Optional[str] = None, ): @@ -700,7 +700,7 @@ def fee_reserve_total(amount_msat: int, internal: bool = False) -> int: return fee_reserve(amount_msat, internal) + service_fee(amount_msat, internal) -async def send_payment_notification(wallet: WalletBalance, payment: Payment): +async def send_payment_notification(wallet: Wallet, payment: Payment): await websocket_updater( wallet.inkey, json.dumps( @@ -857,7 +857,6 @@ async def create_user_account( account.id = uuid4().hex account = await create_account(account) - await create_wallet( user_id=account.id, wallet_name=wallet_name or settings.lnbits_default_wallet_name, diff --git a/lnbits/core/views/api.py b/lnbits/core/views/api.py index 92b0c22f6..49e006d79 100644 --- a/lnbits/core/views/api.py +++ b/lnbits/core/views/api.py @@ -55,12 +55,13 @@ async def health() -> dict: "/api/v1/wallets", name="Wallets", description="Get basic info for all of user's wallets.", + response_model=list[BaseWallet], ) -async def api_wallets(user: User = Depends(check_user_exists)) -> list[BaseWallet]: - return [BaseWallet(**w.dict()) for w in user.wallets] +async def api_wallets(user: User = Depends(check_user_exists)) -> list[Wallet]: + return user.wallets -@api_router.post("/api/v1/account", response_model=Wallet) +@api_router.post("/api/v1/account") async def api_create_account(data: CreateWallet) -> Wallet: user = await create_user_account(wallet_name=data.name) return user.wallets[0] diff --git a/lnbits/core/views/generic.py b/lnbits/core/views/generic.py index 7db37af89..85694cac3 100644 --- a/lnbits/core/views/generic.py +++ b/lnbits/core/views/generic.py @@ -13,7 +13,7 @@ from pydantic.types import UUID4 from lnbits.core.extensions.models import Extension, ExtensionMeta, InstallableExtension from lnbits.core.helpers import to_valid_user_id -from lnbits.core.models import User, WalletBalance +from lnbits.core.models import User from lnbits.core.services import create_invoice, create_user_account from lnbits.decorators import check_admin, check_user_exists from lnbits.helpers import template_renderer @@ -167,8 +167,7 @@ async def wallet( if wal: wallet = await get_wallet(wal.hex) elif len(user.wallets) == 0: - _wallet = await create_wallet(user_id=user.id) - wallet = WalletBalance(**_wallet.dict()) + wallet = await create_wallet(user_id=user.id) user.wallets.append(wallet) elif lnbits_last_active_wallet and user.get_wallet(lnbits_last_active_wallet): wallet = await get_wallet(lnbits_last_active_wallet) diff --git a/lnbits/core/views/user_api.py b/lnbits/core/views/user_api.py index ccd527146..b4780ee48 100644 --- a/lnbits/core/views/user_api.py +++ b/lnbits/core/views/user_api.py @@ -21,7 +21,7 @@ from lnbits.core.models import ( AccountOverview, CreateTopup, User, - WalletBalance, + Wallet, ) from lnbits.core.services import update_wallet_balance from lnbits.db import Filters, Page @@ -103,7 +103,7 @@ async def api_users_toggle_admin(user_id: str) -> None: @users_router.get("/user/{user_id}/wallet") -async def api_users_get_user_wallet(user_id: str) -> List[WalletBalance]: +async def api_users_get_user_wallet(user_id: str) -> List[Wallet]: return await get_wallets(user_id) diff --git a/lnbits/core/views/wallet_api.py b/lnbits/core/views/wallet_api.py index 4ce1ee9b1..1a484a1b4 100644 --- a/lnbits/core/views/wallet_api.py +++ b/lnbits/core/views/wallet_api.py @@ -12,7 +12,6 @@ from lnbits.core.models import ( CreateWallet, KeyType, Wallet, - WalletBalance, ) from lnbits.decorators import ( WalletTypeInfo, @@ -62,7 +61,7 @@ async def api_update_wallet( name: Optional[str] = Body(None), currency: Optional[str] = Body(None), key_info: WalletTypeInfo = Depends(require_admin_key), -) -> WalletBalance: +) -> Wallet: wallet = await get_wallet(key_info.wallet.id) if not wallet: raise HTTPException(status_code=HTTPStatus.NOT_FOUND, detail="Wallet not found") @@ -85,6 +84,6 @@ async def api_delete_wallet( @wallet_router.post("") async def api_create_wallet( data: CreateWallet, - wallet: WalletTypeInfo = Depends(require_admin_key), + key_info: WalletTypeInfo = Depends(require_admin_key), ) -> Wallet: - return await create_wallet(user_id=wallet.wallet.user, wallet_name=data.name) + return await create_wallet(user_id=key_info.wallet.user, wallet_name=data.name) diff --git a/lnbits/db.py b/lnbits/db.py index 478db343a..54b4e5b28 100644 --- a/lnbits/db.py +++ b/lnbits/db.py @@ -570,10 +570,11 @@ def insert_query(table_name: str, model: BaseModel) -> str: :param model: Pydantic model """ placeholders = [] - for field in model.dict().keys(): + keys = model_to_dict(model).keys() + for field in keys: placeholders.append(get_placeholder(model, field)) # add quotes to keys to avoid SQL conflicts (e.g. `user` is a reserved keyword) - fields = ", ".join([f'"{key}"' for key in model.dict().keys()]) + fields = ", ".join([f'"{key}"' for key in keys]) values = ", ".join(placeholders) return f"INSERT INTO {table_name} ({fields}) VALUES ({values})" @@ -586,7 +587,7 @@ def update_query(table_name: str, model: BaseModel, where: str = "id = :id") -> :param where: Where string, default to `id = :id` """ fields = [] - for field in model.dict().keys(): + for field in model_to_dict(model).keys(): placeholder = get_placeholder(model, field) # add quotes to keys to avoid SQL conflicts (e.g. `user` is a reserved keyword) fields.append(f'"{field}" = {placeholder}') @@ -602,9 +603,9 @@ def model_to_dict(model: BaseModel) -> dict: """ _dict: dict = {} for key, value in model.dict().items(): - if key.startswith("_"): - continue type_ = model.__fields__[key].type_ + if model.__fields__[key].field_info.extra.get("no_database", False): + continue if isinstance(value, datetime.datetime): _dict[key] = value.timestamp() continue @@ -643,9 +644,6 @@ def dict_to_model(_row: dict, model: type[TModel]) -> TModel: logger.warning(f"Converting {key} to model `{model}`.") continue type_ = model.__fields__[key].type_ - # if issubclass(type_, datetime.datetime): - # _dict[key] = datetime.datetime.fromtimestamp(value) - # continue if issubclass(type_, bool): _dict[key] = bool(value) continue @@ -654,4 +652,5 @@ def dict_to_model(_row: dict, model: type[TModel]) -> TModel: continue _dict[key] = value continue - return model.construct(**_dict) + _model = model.construct(**_dict) + return _model