refactor: extract permissions logic

This commit is contained in:
Vlad Stan
2026-07-09 16:35:53 +03:00
parent b5506e7a0b
commit e45f0cf73d
3 changed files with 62 additions and 58 deletions
+59
View File
@@ -0,0 +1,59 @@
from collections.abc import Iterable
from typing import Any
from lnbits.core.extensions.api import extension_api_permission_ids
from lnbits.core.models.extensions import ExtensionPermission, InstallableExtension
def validate_extension_permissions(
ext_id: str,
permissions: Iterable[ExtensionPermission],
*,
strict: bool = True,
) -> list[ExtensionPermission]:
known_permission_ids = extension_api_permission_ids()
normalized_permissions: list[ExtensionPermission] = []
unknown_ids: list[str] = []
for permission in permissions:
if permission.id not in known_permission_ids:
unknown_ids.append(permission.id)
if strict:
continue
normalized_permissions.append(permission.copy(update={"label": None}))
if unknown_ids and strict:
raise ValueError(
f"Extension '{ext_id}' requests unknown permissions: "
+ ", ".join(sorted(set(unknown_ids)))
)
return normalized_permissions
def validate_wasm_extension_permissions(
ext_info: InstallableExtension,
granted_permissions: list[ExtensionPermission] | None,
extension_config: dict[str, Any],
) -> list[ExtensionPermission]:
if extension_config.get("extension_type") != "wasm":
return []
requested_permissions = validate_extension_permissions(
ext_info.id,
ExtensionPermission.list_from_config(extension_config),
)
if not requested_permissions:
return []
if granted_permissions is None:
raise ValueError(f"Extension '{ext_info.id}' requires permission approval.")
requested_ids = {permission.id for permission in requested_permissions}
granted_ids = {permission.id for permission in granted_permissions}
if requested_ids != granted_ids:
raise ValueError(
f"Extension '{ext_info.id}' was not granted all requested permissions."
)
return requested_permissions
+2 -57
View File
@@ -2,7 +2,6 @@ import asyncio
import importlib import importlib
import json import json
import zipfile import zipfile
from collections.abc import Iterable
from pathlib import PurePosixPath from pathlib import PurePosixPath
from typing import Any from typing import Any
@@ -21,7 +20,7 @@ from lnbits.core.crud.extensions import (
get_installed_extensions, get_installed_extensions,
update_installed_extension, update_installed_extension,
) )
from lnbits.core.extensions.api import extension_api_permission_ids from lnbits.core.extensions.permissions import validate_wasm_extension_permissions
from lnbits.core.helpers import migrate_extension_database from lnbits.core.helpers import migrate_extension_database
from lnbits.db import Connection from lnbits.db import Connection
from lnbits.settings import settings from lnbits.settings import settings
@@ -58,7 +57,7 @@ async def install_extension(
await ext_info.download_archive() await ext_info.download_archive()
extension_config = _load_extension_archive_config(ext_info) extension_config = _load_extension_archive_config(ext_info)
ext_info.permissions = _validate_extension_permissions( ext_info.permissions = validate_wasm_extension_permissions(
ext_info, granted_permissions, extension_config ext_info, granted_permissions, extension_config
) )
@@ -85,60 +84,6 @@ async def install_extension(
return extension return extension
def validate_extension_permissions(
ext_id: str,
permissions: Iterable[ExtensionPermission],
*,
strict: bool = True,
) -> list[ExtensionPermission]:
known_permission_ids = extension_api_permission_ids()
normalized_permissions: list[ExtensionPermission] = []
unknown_ids: list[str] = []
for permission in permissions:
if permission.id not in known_permission_ids:
unknown_ids.append(permission.id)
if strict:
continue
normalized_permissions.append(permission.copy(update={"label": None}))
if unknown_ids and strict:
raise ValueError(
f"Extension '{ext_id}' requests unknown permissions: "
+ ", ".join(sorted(set(unknown_ids)))
)
return normalized_permissions
def _validate_extension_permissions(
ext_info: InstallableExtension,
granted_permissions: list[ExtensionPermission] | None,
extension_config: dict[str, Any],
) -> list[ExtensionPermission]:
if extension_config.get("extension_type") != "wasm":
return []
requested_permissions = validate_extension_permissions(
ext_info.id,
ExtensionPermission.list_from_config(extension_config),
)
if not requested_permissions:
return []
if granted_permissions is None:
raise ValueError(f"Extension '{ext_info.id}' requires permission approval.")
requested_ids = {permission.id for permission in requested_permissions}
granted_ids = {permission.id for permission in granted_permissions}
if requested_ids != granted_ids:
raise ValueError(
f"Extension '{ext_info.id}' was not granted all requested permissions."
)
return requested_permissions
def _load_extension_archive_config(ext_info: InstallableExtension) -> dict[str, Any]: def _load_extension_archive_config(ext_info: InstallableExtension) -> dict[str, Any]:
if not ext_info.zip_path.is_file(): if not ext_info.zip_path.is_file():
return {} return {}
+1 -1
View File
@@ -11,6 +11,7 @@ from loguru import logger
from lnbits.core.crud.extensions import get_user_extensions from lnbits.core.crud.extensions import get_user_extensions
from lnbits.core.crud.wallets import get_wallets_ids from lnbits.core.crud.wallets import get_wallets_ids
from lnbits.core.db import db from lnbits.core.db import db
from lnbits.core.extensions.permissions import validate_extension_permissions
from lnbits.core.models import ( from lnbits.core.models import (
SimpleStatus, SimpleStatus,
) )
@@ -40,7 +41,6 @@ from lnbits.core.services.extensions import (
get_valid_extensions, get_valid_extensions,
install_extension, install_extension,
uninstall_extension, uninstall_extension,
validate_extension_permissions,
) )
from lnbits.db import Page from lnbits.db import Page
from lnbits.decorators import ( from lnbits.decorators import (