add dbversion model
This commit is contained in:
+2
-2
@@ -17,7 +17,7 @@ from slowapi.util import get_remote_address
|
|||||||
from starlette.middleware.sessions import SessionMiddleware
|
from starlette.middleware.sessions import SessionMiddleware
|
||||||
|
|
||||||
from lnbits.core.crud import (
|
from lnbits.core.crud import (
|
||||||
get_dbversions,
|
get_db_version,
|
||||||
get_installed_extensions,
|
get_installed_extensions,
|
||||||
update_installed_extension_state,
|
update_installed_extension_state,
|
||||||
)
|
)
|
||||||
@@ -313,7 +313,7 @@ async def restore_installed_extension(app: FastAPI, ext: InstallableExtension):
|
|||||||
extension = Extension.from_installable_ext(ext)
|
extension = Extension.from_installable_ext(ext)
|
||||||
register_ext_routes(app, extension)
|
register_ext_routes(app, extension)
|
||||||
|
|
||||||
current_version = (await get_dbversions()).get(ext.id, 0)
|
current_version = await get_db_version(ext.id)
|
||||||
await migrate_extension_database(ext, current_version)
|
await migrate_extension_database(ext, current_version)
|
||||||
|
|
||||||
# mount routes for the new version
|
# mount routes for the new version
|
||||||
|
|||||||
+2
-2
@@ -17,7 +17,7 @@ from lnbits.core.crud import (
|
|||||||
delete_unused_wallets,
|
delete_unused_wallets,
|
||||||
delete_wallet_by_id,
|
delete_wallet_by_id,
|
||||||
delete_wallet_payment,
|
delete_wallet_payment,
|
||||||
get_dbversions,
|
get_db_versions,
|
||||||
get_installed_extension,
|
get_installed_extension,
|
||||||
get_installed_extensions,
|
get_installed_extensions,
|
||||||
get_payments,
|
get_payments,
|
||||||
@@ -122,7 +122,7 @@ def database_migrate():
|
|||||||
async def db_versions():
|
async def db_versions():
|
||||||
"""Show current database versions"""
|
"""Show current database versions"""
|
||||||
async with core_db.connect() as conn:
|
async with core_db.connect() as conn:
|
||||||
click.echo(await get_dbversions(conn))
|
click.echo(await get_db_versions(conn))
|
||||||
|
|
||||||
|
|
||||||
@db.command("cleanup-wallets")
|
@db.command("cleanup-wallets")
|
||||||
|
|||||||
+25
-16
@@ -25,6 +25,7 @@ from .models import (
|
|||||||
AccountFilters,
|
AccountFilters,
|
||||||
AccountOverview,
|
AccountOverview,
|
||||||
CreatePayment,
|
CreatePayment,
|
||||||
|
DbVersion,
|
||||||
Payment,
|
Payment,
|
||||||
PaymentFilters,
|
PaymentFilters,
|
||||||
PaymentHistoryPoint,
|
PaymentHistoryPoint,
|
||||||
@@ -842,35 +843,37 @@ async def check_internal(
|
|||||||
Returns the checking_id of the internal payment if it exists,
|
Returns the checking_id of the internal payment if it exists,
|
||||||
otherwise None
|
otherwise None
|
||||||
"""
|
"""
|
||||||
row: dict = await (conn or db).fetchone(
|
payment = await (conn or db).fetchone(
|
||||||
f"""
|
f"""
|
||||||
SELECT checking_id FROM apipayments
|
SELECT * FROM apipayments
|
||||||
WHERE payment_hash = :hash AND status = '{PaymentState.PENDING}' AND amount > 0
|
WHERE payment_hash = :hash AND status = '{PaymentState.PENDING}' AND amount > 0
|
||||||
""",
|
""",
|
||||||
{"hash": payment_hash},
|
{"hash": payment_hash},
|
||||||
|
Payment,
|
||||||
)
|
)
|
||||||
if not row:
|
if not payment:
|
||||||
return None
|
return None
|
||||||
else:
|
else:
|
||||||
return row["checking_id"]
|
return payment.checking_id
|
||||||
|
|
||||||
|
|
||||||
async def check_internal_status(
|
async def is_internal_status_success(
|
||||||
payment_hash: str, conn: Optional[Connection] = None
|
payment_hash: str, conn: Optional[Connection] = None
|
||||||
) -> bool:
|
) -> bool:
|
||||||
"""
|
"""
|
||||||
Returns True if the internal payment was successful
|
Returns True if the internal payment was found and has the given status,
|
||||||
"""
|
"""
|
||||||
row: dict = await (conn or db).fetchone(
|
payment = await (conn or db).fetchone(
|
||||||
"""
|
"""
|
||||||
SELECT status FROM apipayments
|
SELECT * FROM apipayments
|
||||||
WHERE payment_hash = :payment_hash AND amount > 0
|
WHERE payment_hash = :payment_hash AND amount > 0
|
||||||
""",
|
""",
|
||||||
{"payment_hash": payment_hash},
|
{"payment_hash": payment_hash},
|
||||||
|
Payment,
|
||||||
)
|
)
|
||||||
if not row:
|
if not payment:
|
||||||
return True
|
return True
|
||||||
return row["status"] == PaymentState.SUCCESS.value
|
return payment.status == PaymentState.SUCCESS.value
|
||||||
|
|
||||||
|
|
||||||
async def mark_webhook_sent(payment_hash: str, status: int) -> None:
|
async def mark_webhook_sent(payment_hash: str, status: int) -> None:
|
||||||
@@ -950,12 +953,18 @@ async def create_admin_settings(super_user: str, new_settings: dict):
|
|||||||
|
|
||||||
# db versions
|
# db versions
|
||||||
# --------------
|
# --------------
|
||||||
async def get_dbversions(conn: Optional[Connection] = None) -> dict:
|
async def get_db_version(
|
||||||
result = await (conn or db).execute("SELECT db, version FROM dbversions")
|
ext_id: str, conn: Optional[Connection] = None
|
||||||
_dict = {}
|
) -> Optional[DbVersion]:
|
||||||
for row in result.mappings().all():
|
return await (conn or db).fetchone(
|
||||||
_dict[row["db"]] = row["version"]
|
"SELECT * FROM dbversions WHERE db = :ext_id",
|
||||||
return _dict
|
{"ext_id": ext_id},
|
||||||
|
model=DbVersion,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def get_db_versions(conn: Optional[Connection] = None) -> list[DbVersion]:
|
||||||
|
return await (conn or db).fetchall("SELECT * FROM dbversions", model=DbVersion)
|
||||||
|
|
||||||
|
|
||||||
async def update_migration_version(conn, db_name, version):
|
async def update_migration_version(conn, db_name, version):
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ from lnbits.core import core_app_extra
|
|||||||
from lnbits.core.crud import (
|
from lnbits.core.crud import (
|
||||||
create_installed_extension,
|
create_installed_extension,
|
||||||
delete_installed_extension,
|
delete_installed_extension,
|
||||||
get_dbversions,
|
get_db_version,
|
||||||
get_installed_extension,
|
get_installed_extension,
|
||||||
update_installed_extension_state,
|
update_installed_extension_state,
|
||||||
)
|
)
|
||||||
@@ -28,7 +28,7 @@ async def install_extension(ext_info: InstallableExtension) -> Extension:
|
|||||||
|
|
||||||
ext_info.extract_archive()
|
ext_info.extract_archive()
|
||||||
|
|
||||||
db_version = (await get_dbversions()).get(ext_id, 0)
|
db_version = await get_db_version(ext_id)
|
||||||
await migrate_extension_database(ext_info, db_version)
|
await migrate_extension_database(ext_info, db_version)
|
||||||
|
|
||||||
await create_installed_extension(ext_info)
|
await create_installed_extension(ext_info)
|
||||||
|
|||||||
+20
-8
@@ -1,6 +1,6 @@
|
|||||||
import importlib
|
import importlib
|
||||||
import re
|
import re
|
||||||
from typing import Any
|
from typing import Any, Optional
|
||||||
from urllib.parse import urlparse
|
from urllib.parse import urlparse
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
@@ -8,17 +8,20 @@ from loguru import logger
|
|||||||
|
|
||||||
from lnbits.core import migrations as core_migrations
|
from lnbits.core import migrations as core_migrations
|
||||||
from lnbits.core.crud import (
|
from lnbits.core.crud import (
|
||||||
get_dbversions,
|
get_db_versions,
|
||||||
get_installed_extensions,
|
get_installed_extensions,
|
||||||
update_migration_version,
|
update_migration_version,
|
||||||
)
|
)
|
||||||
from lnbits.core.db import db as core_db
|
from lnbits.core.db import db as core_db
|
||||||
from lnbits.core.extensions.models import InstallableExtension
|
from lnbits.core.extensions.models import InstallableExtension
|
||||||
|
from lnbits.core.models import DbVersion
|
||||||
from lnbits.db import COCKROACH, POSTGRES, SQLITE, Connection
|
from lnbits.db import COCKROACH, POSTGRES, SQLITE, Connection
|
||||||
from lnbits.settings import settings
|
from lnbits.settings import settings
|
||||||
|
|
||||||
|
|
||||||
async def migrate_extension_database(ext: InstallableExtension, current_version: int):
|
async def migrate_extension_database(
|
||||||
|
ext: InstallableExtension, current_version: Optional[DbVersion] = None
|
||||||
|
):
|
||||||
|
|
||||||
try:
|
try:
|
||||||
ext_migrations = importlib.import_module(f"{ext.module_name}.migrations")
|
ext_migrations = importlib.import_module(f"{ext.module_name}.migrations")
|
||||||
@@ -32,14 +35,17 @@ async def migrate_extension_database(ext: InstallableExtension, current_version:
|
|||||||
|
|
||||||
|
|
||||||
async def run_migration(
|
async def run_migration(
|
||||||
db: Connection, migrations_module: Any, db_name: str, current_version: int
|
db: Connection,
|
||||||
|
migrations_module: Any,
|
||||||
|
db_name: str,
|
||||||
|
current_version: Optional[DbVersion] = None,
|
||||||
):
|
):
|
||||||
matcher = re.compile(r"^m(\d\d\d)_")
|
matcher = re.compile(r"^m(\d\d\d)_")
|
||||||
for key, migrate in migrations_module.__dict__.items():
|
for key, migrate in migrations_module.__dict__.items():
|
||||||
match = matcher.match(key)
|
match = matcher.match(key)
|
||||||
if match:
|
if match:
|
||||||
version = int(match.group(1))
|
version = int(match.group(1))
|
||||||
if version > current_version:
|
if not current_version or version > current_version.version:
|
||||||
logger.debug(f"running migration {db_name}.{version}")
|
logger.debug(f"running migration {db_name}.{version}")
|
||||||
print(f"running migration {db_name}.{version}")
|
print(f"running migration {db_name}.{version}")
|
||||||
await migrate(db)
|
await migrate(db)
|
||||||
@@ -71,7 +77,7 @@ async def load_disabled_extension_list() -> None:
|
|||||||
async def migrate_databases():
|
async def migrate_databases():
|
||||||
"""Creates the necessary databases if they don't exist already; or migrates them."""
|
"""Creates the necessary databases if they don't exist already; or migrates them."""
|
||||||
|
|
||||||
current_versions = await get_dbversions()
|
current_versions = await get_db_versions()
|
||||||
async with core_db.connect() as conn:
|
async with core_db.connect() as conn:
|
||||||
exists = False
|
exists = False
|
||||||
if conn.type == SQLITE:
|
if conn.type == SQLITE:
|
||||||
@@ -87,7 +93,10 @@ async def migrate_databases():
|
|||||||
if not exists:
|
if not exists:
|
||||||
await core_migrations.m000_create_migrations_table(conn)
|
await core_migrations.m000_create_migrations_table(conn)
|
||||||
|
|
||||||
core_version = current_versions.get("core", 0)
|
core_version = next(
|
||||||
|
(v for v in current_versions if v.db == "core"),
|
||||||
|
DbVersion(db="core", version=0),
|
||||||
|
)
|
||||||
await run_migration(conn, core_migrations, "core", core_version)
|
await run_migration(conn, core_migrations, "core", core_version)
|
||||||
|
|
||||||
# here is the first place we can be sure that the
|
# here is the first place we can be sure that the
|
||||||
@@ -95,7 +104,10 @@ async def migrate_databases():
|
|||||||
await load_disabled_extension_list()
|
await load_disabled_extension_list()
|
||||||
|
|
||||||
for ext in await get_installed_extensions():
|
for ext in await get_installed_extensions():
|
||||||
current_version = current_versions.get(ext.id)
|
current_version = next(
|
||||||
|
(v for v in current_versions if v.db == ext.id),
|
||||||
|
DbVersion(db=ext.id, version=0),
|
||||||
|
)
|
||||||
if current_version is None:
|
if current_version is None:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
f"Extension {ext.id} has no migration version. This should not happen."
|
f"Extension {ext.id} has no migration version. This should not happen."
|
||||||
|
|||||||
@@ -469,3 +469,8 @@ class BalanceDelta(BaseModel):
|
|||||||
class SimpleStatus(BaseModel):
|
class SimpleStatus(BaseModel):
|
||||||
success: bool
|
success: bool
|
||||||
message: str
|
message: str
|
||||||
|
|
||||||
|
|
||||||
|
class DbVersion(BaseModel):
|
||||||
|
db: str
|
||||||
|
version: int
|
||||||
|
|||||||
@@ -46,7 +46,6 @@ from lnbits.wallets.base import (
|
|||||||
|
|
||||||
from .crud import (
|
from .crud import (
|
||||||
check_internal,
|
check_internal,
|
||||||
check_internal_status,
|
|
||||||
create_account,
|
create_account,
|
||||||
create_admin_settings,
|
create_admin_settings,
|
||||||
create_payment,
|
create_payment,
|
||||||
@@ -62,6 +61,7 @@ from .crud import (
|
|||||||
get_user,
|
get_user,
|
||||||
get_wallet,
|
get_wallet,
|
||||||
get_wallet_payment,
|
get_wallet_payment,
|
||||||
|
is_internal_status_success,
|
||||||
update_admin_settings,
|
update_admin_settings,
|
||||||
update_payment_details,
|
update_payment_details,
|
||||||
update_payment_status,
|
update_payment_status,
|
||||||
@@ -247,9 +247,7 @@ async def pay_invoice(
|
|||||||
extra=extra,
|
extra=extra,
|
||||||
)
|
)
|
||||||
|
|
||||||
# we check if an internal invoice exists that has already been paid
|
if await is_internal_status_success(invoice.payment_hash, conn=conn):
|
||||||
# (not pending anymore)
|
|
||||||
if await check_internal_status(invoice.payment_hash, conn=conn):
|
|
||||||
raise PaymentError("Internal invoice already paid.", status="failed")
|
raise PaymentError("Internal invoice already paid.", status="failed")
|
||||||
|
|
||||||
# check_internal() returns the checking_id of the invoice we're waiting for
|
# check_internal() returns the checking_id of the invoice we're waiting for
|
||||||
|
|||||||
@@ -40,7 +40,7 @@ from ..crud import (
|
|||||||
create_user_extension,
|
create_user_extension,
|
||||||
delete_dbversion,
|
delete_dbversion,
|
||||||
drop_extension_db,
|
drop_extension_db,
|
||||||
get_dbversions,
|
get_db_version,
|
||||||
get_installed_extension,
|
get_installed_extension,
|
||||||
get_installed_extensions,
|
get_installed_extensions,
|
||||||
get_user_extension,
|
get_user_extension,
|
||||||
@@ -474,7 +474,7 @@ async def get_extension_release(org: str, repo: str, tag_name: str):
|
|||||||
)
|
)
|
||||||
async def delete_extension_db(ext_id: str):
|
async def delete_extension_db(ext_id: str):
|
||||||
try:
|
try:
|
||||||
db_version = (await get_dbversions()).get(ext_id, None)
|
db_version = await get_db_version(ext_id)
|
||||||
if not db_version:
|
if not db_version:
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
status_code=HTTPStatus.BAD_REQUEST,
|
status_code=HTTPStatus.BAD_REQUEST,
|
||||||
|
|||||||
@@ -23,7 +23,7 @@ from lnbits.wallets import get_funding_source
|
|||||||
from ...utils.exchange_rates import allowed_currencies, currencies
|
from ...utils.exchange_rates import allowed_currencies, currencies
|
||||||
from ..crud import (
|
from ..crud import (
|
||||||
create_wallet,
|
create_wallet,
|
||||||
get_dbversions,
|
get_db_versions,
|
||||||
get_installed_extensions,
|
get_installed_extensions,
|
||||||
get_user_by_id,
|
get_user_by_id,
|
||||||
get_wallet,
|
get_wallet,
|
||||||
@@ -104,7 +104,8 @@ async def extensions(request: Request, user: User = Depends(check_user_exists)):
|
|||||||
|
|
||||||
all_ext_ids = [ext.code for ext in Extension.get_valid_extensions()]
|
all_ext_ids = [ext.code for ext in Extension.get_valid_extensions()]
|
||||||
inactive_extensions = [e.id for e in await get_installed_extensions(active=False)]
|
inactive_extensions = [e.id for e in await get_installed_extensions(active=False)]
|
||||||
db_version = await get_dbversions()
|
db_versions = await get_db_versions()
|
||||||
|
|
||||||
extensions = [
|
extensions = [
|
||||||
{
|
{
|
||||||
"id": ext.id,
|
"id": ext.id,
|
||||||
@@ -115,7 +116,9 @@ async def extensions(request: Request, user: User = Depends(check_user_exists)):
|
|||||||
"isFeatured": ext.meta.featured if ext.meta else False,
|
"isFeatured": ext.meta.featured if ext.meta else False,
|
||||||
"dependencies": ext.meta.dependencies if ext.meta else "",
|
"dependencies": ext.meta.dependencies if ext.meta else "",
|
||||||
"isInstalled": ext.id in installed_exts_ids,
|
"isInstalled": ext.id in installed_exts_ids,
|
||||||
"hasDatabaseTables": ext.id in db_version,
|
"hasDatabaseTables": next(
|
||||||
|
(True for version in db_versions if version.db == ext.id), False
|
||||||
|
),
|
||||||
"isAvailable": ext.id in all_ext_ids,
|
"isAvailable": ext.id in all_ext_ids,
|
||||||
"isAdminOnly": ext.id in settings.lnbits_admin_extensions,
|
"isAdminOnly": ext.id in settings.lnbits_admin_extensions,
|
||||||
"isActive": ext.id not in inactive_extensions,
|
"isActive": ext.id not in inactive_extensions,
|
||||||
|
|||||||
Reference in New Issue
Block a user