update sqlalchemy to 1.4

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