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