perf: use check_account_exists decorator (#3600)

This commit is contained in:
Vlad Stan
2025-12-04 10:17:47 +02:00
committed by GitHub
parent 5213508dc1
commit b3efb4d378
12 changed files with 182 additions and 114 deletions
+21 -15
View File
@@ -23,6 +23,7 @@ from lnbits.core.models.extensions import (
UserExtension,
UserExtensionInfo,
)
from lnbits.core.models.users import AccountId
from lnbits.core.services import check_transaction_status, create_invoice
from lnbits.core.services.extensions import (
activate_extension,
@@ -33,8 +34,9 @@ from lnbits.core.services.extensions import (
uninstall_extension,
)
from lnbits.decorators import (
check_account_exists,
check_account_id_exists,
check_admin,
check_user_exists,
)
from lnbits.settings import settings
@@ -163,7 +165,7 @@ async def api_update_pay_to_enable(
@extension_router.put("/{ext_id}/enable")
async def api_enable_extension(
ext_id: str, user: User = Depends(check_user_exists)
ext_id: str, account_id: AccountId = Depends(check_account_id_exists)
) -> SimpleStatus:
if ext_id not in [e.code for e in await get_valid_extensions()]:
raise HTTPException(
@@ -177,12 +179,12 @@ async def api_enable_extension(
if not ext.active:
raise ValueError(f"Extension '{ext_id}' is not activated.")
user_ext = await get_user_extension(user.id, ext_id)
user_ext = await get_user_extension(account_id.id, ext_id)
if not user_ext:
user_ext = UserExtension(user=user.id, extension=ext_id, active=False)
user_ext = UserExtension(user=account_id.id, extension=ext_id, active=False)
await create_user_extension(user_ext)
if user.admin or not ext.requires_payment:
if account_id.is_admin_id or not ext.requires_payment:
user_ext.active = True
await update_user_extension(user_ext)
return SimpleStatus(success=True, message=f"Extension '{ext_id}' enabled.")
@@ -219,13 +221,13 @@ async def api_enable_extension(
@extension_router.put("/{ext_id}/disable")
async def api_disable_extension(
ext_id: str, user: User = Depends(check_user_exists)
ext_id: str, account_id: AccountId = Depends(check_account_id_exists)
) -> SimpleStatus:
if ext_id not in [e.code for e in await get_valid_extensions()]:
raise HTTPException(
HTTPStatus.BAD_REQUEST, f"Extension '{ext_id}' doesn't exist."
)
user_ext = await get_user_extension(user.id, ext_id)
user_ext = await get_user_extension(account_id.id, ext_id)
if not user_ext or not user_ext.active:
return SimpleStatus(
success=True, message=f"Extension '{ext_id}' already disabled."
@@ -376,7 +378,9 @@ async def get_pay_to_install_invoice(
@extension_router.put("/{ext_id}/invoice/enable")
async def get_pay_to_enable_invoice(
ext_id: str, data: PayToEnableInfo, user: User = Depends(check_user_exists)
ext_id: str,
data: PayToEnableInfo,
account_id: AccountId = Depends(check_account_id_exists),
):
if not data.amount or data.amount <= 0:
raise HTTPException(
@@ -422,9 +426,9 @@ async def get_pay_to_enable_invoice(
memo=f"Enable '{ext.name}' extension.",
)
user_ext = await get_user_extension(user.id, ext_id)
user_ext = await get_user_extension(account_id.id, ext_id)
if not user_ext:
user_ext = UserExtension(user=user.id, extension=ext_id, active=False)
user_ext = UserExtension(user=account_id.id, extension=ext_id, active=False)
await create_user_extension(user_ext)
user_ext_info = user_ext.extra if user_ext.extra else UserExtensionInfo()
user_ext_info.payment_hash_to_enable = payment.payment_hash
@@ -435,7 +439,7 @@ async def get_pay_to_enable_invoice(
@extension_router.get(
"/release/{org}/{repo}/{tag_name}",
dependencies=[Depends(check_user_exists)],
dependencies=[Depends(check_account_exists)],
)
async def get_extension_release(org: str, repo: str, tag_name: str):
try:
@@ -456,10 +460,12 @@ async def get_extension_release(org: str, repo: str, tag_name: str):
@extension_router.get("")
async def api_get_user_extensions(
user: User = Depends(check_user_exists),
account_id: AccountId = Depends(check_account_id_exists),
) -> list[Extension]:
user_extensions_ids = [ue.extension for ue in await get_user_extensions(user.id)]
user_extensions_ids = [
ue.extension for ue in await get_user_extensions(account_id.id)
]
return [
ext
for ext in await get_valid_extensions(False)
@@ -498,7 +504,7 @@ async def delete_extension_db(ext_id: str):
# TODO: create a response model for this
@extension_router.get("/all")
async def extensions(user: User = Depends(check_user_exists)):
async def extensions(account_id: AccountId = Depends(check_account_id_exists)):
installed_exts: list[InstallableExtension] = await get_installed_extensions()
installed_exts_ids = [e.id for e in installed_exts]
@@ -510,7 +516,7 @@ async def extensions(user: User = Depends(check_user_exists)):
installed_ext = next((ie for ie in installed_exts if e.id == ie.id), None)
if installed_ext and installed_ext.meta:
installed_release = installed_ext.meta.installed_release
if installed_ext.meta.pay_to_enable and not user.admin:
if installed_ext.meta.pay_to_enable and not account_id.is_admin_id:
# not a security leak, but better not to share the wallet id
installed_ext.meta.pay_to_enable.wallet = None
pay_to_enable = installed_ext.meta.pay_to_enable