fix part1

This commit is contained in:
dni ⚡
2024-10-10 09:07:44 +02:00
parent 2410c271a6
commit 383f6ec721
7 changed files with 94 additions and 123 deletions
+45 -69
View File
@@ -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(
+9
View File
@@ -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")
+24 -35
View File
@@ -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"
+14 -4
View File
@@ -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:
-6
View File
@@ -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"
-8
View File
@@ -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
+2 -1
View File
@@ -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