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 .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
View File
@@ -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
View File
@@ -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