This commit is contained in:
dni ⚡
2024-10-10 09:08:42 +02:00
parent e8d10abee3
commit 28897ca120
8 changed files with 43 additions and 52 deletions
+7 -8
View File
@@ -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,
)
+14 -19
View File
@@ -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
+3 -4
View File
@@ -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,
+4 -3
View File
@@ -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]
+2 -3
View File
@@ -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)
+2 -2
View File
@@ -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)
+3 -4
View File
@@ -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)
+8 -9
View File
@@ -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