update sqlalchemy to 1.4
This commit is contained in:
+1
-2
@@ -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
@@ -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
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user