298 lines
9.5 KiB
Python
298 lines
9.5 KiB
Python
import asyncio
|
|
import importlib
|
|
import json
|
|
import zipfile
|
|
from collections.abc import Iterable
|
|
from pathlib import PurePosixPath
|
|
from typing import Any
|
|
|
|
from loguru import logger
|
|
|
|
from lnbits.core import core_app_extra
|
|
from lnbits.core.crud import (
|
|
create_installed_extension,
|
|
delete_installed_extension,
|
|
get_db_version,
|
|
get_installed_extension,
|
|
get_installed_extensions_count,
|
|
update_installed_extension_state,
|
|
)
|
|
from lnbits.core.crud.extensions import (
|
|
get_installed_extensions,
|
|
update_installed_extension,
|
|
)
|
|
from lnbits.core.extensions.api import extension_api_permission_ids
|
|
from lnbits.core.helpers import migrate_extension_database
|
|
from lnbits.db import Connection
|
|
from lnbits.settings import settings
|
|
|
|
from ..models.extensions import (
|
|
Extension,
|
|
ExtensionMeta,
|
|
ExtensionPermission,
|
|
InstallableExtension,
|
|
)
|
|
|
|
|
|
async def install_extension(
|
|
ext_info: InstallableExtension,
|
|
skip_download: bool | None = False,
|
|
granted_permissions: list[ExtensionPermission] | None = None,
|
|
) -> Extension:
|
|
|
|
ext_info.meta = ext_info.meta or ExtensionMeta()
|
|
|
|
if (
|
|
ext_info.meta.installed_release
|
|
and not ext_info.meta.installed_release.is_version_compatible
|
|
):
|
|
raise ValueError("Incompatible extension version")
|
|
|
|
installed_ext = await get_installed_extension(ext_info.id)
|
|
if installed_ext and installed_ext.meta:
|
|
ext_info.meta.payments = installed_ext.meta.payments
|
|
|
|
await check_extensions_limit(installed_ext)
|
|
|
|
if not skip_download:
|
|
await ext_info.download_archive()
|
|
|
|
extension_config = _load_extension_archive_config(ext_info)
|
|
ext_info.permissions = _validate_extension_permissions(
|
|
ext_info, granted_permissions, extension_config
|
|
)
|
|
|
|
ext_info.extract_archive()
|
|
|
|
db_version = await get_db_version(ext_info.id)
|
|
await migrate_extension_database(ext_info, db_version)
|
|
|
|
# if the extensions does not exist in the installed extensions table, create it
|
|
# if it does exist, it will be activated later in the code
|
|
if not installed_ext:
|
|
await create_installed_extension(ext_info)
|
|
else:
|
|
await update_installed_extension(ext_info)
|
|
|
|
extension = Extension.from_installable_ext(ext_info)
|
|
if extension.is_upgrade_extension:
|
|
# call stop while the old routes are still active
|
|
await stop_extension_background_work(ext_info.id)
|
|
|
|
await start_extension_background_work(ext_info.id)
|
|
|
|
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.parse_obj(permission)
|
|
for permission in extension_config.get("permissions") or []
|
|
if isinstance(permission, dict) and permission.get("id")
|
|
],
|
|
)
|
|
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]:
|
|
if not ext_info.zip_path.is_file():
|
|
return {}
|
|
|
|
try:
|
|
with zipfile.ZipFile(ext_info.zip_path, "r") as archive:
|
|
config_name = _archive_config_name(archive.namelist())
|
|
if not config_name:
|
|
return {}
|
|
with archive.open(config_name) as config_file:
|
|
config = json.load(config_file)
|
|
except Exception as exc:
|
|
raise ValueError(f"Cannot read extension config for '{ext_info.id}'.") from exc
|
|
|
|
return config if isinstance(config, dict) else {}
|
|
|
|
|
|
def _archive_config_name(names: list[str]) -> str | None:
|
|
for name in names:
|
|
path = PurePosixPath(name)
|
|
if len(path.parts) == 2 and path.name == "config.json":
|
|
return name
|
|
return None
|
|
|
|
|
|
async def check_extensions_limit(installed_ext: InstallableExtension | None = None):
|
|
if settings.lnbits_max_extensions == 0 or installed_ext:
|
|
return
|
|
|
|
extensions_count = await get_installed_extensions_count()
|
|
if extensions_count >= settings.lnbits_max_extensions:
|
|
raise ValueError("Max amount of extensions have been installed")
|
|
|
|
|
|
async def uninstall_extension(ext_id: str):
|
|
await stop_extension_background_work(ext_id)
|
|
|
|
settings.deactivate_extension_paths(ext_id)
|
|
|
|
extension = await get_installed_extension(ext_id)
|
|
if extension:
|
|
extension.clean_extension_files()
|
|
await delete_installed_extension(ext_id=ext_id)
|
|
|
|
|
|
async def activate_extension(ext: Extension):
|
|
core_app_extra.register_new_ext_routes(ext)
|
|
await update_installed_extension_state(ext_id=ext.code, active=True)
|
|
await start_extension_background_work(ext.code)
|
|
|
|
|
|
async def deactivate_extension(ext_id: str):
|
|
settings.deactivate_extension_paths(ext_id)
|
|
await update_installed_extension_state(ext_id=ext_id, active=False)
|
|
await stop_extension_background_work(ext_id)
|
|
|
|
|
|
async def stop_extension_background_work(ext_id: str) -> bool:
|
|
"""
|
|
Stop background work for extension (like asyncio.Tasks, WebSockets, etc).
|
|
Extension must expose a `myextension_stop()` function if it is starting tasks.
|
|
"""
|
|
upgrade_hash = settings.extension_upgrade_hash(ext_id)
|
|
ext = Extension(code=ext_id, is_valid=True, upgrade_hash=upgrade_hash)
|
|
|
|
try:
|
|
logger.info(f"Stopping background work for extension '{ext.module_name}'.")
|
|
old_module = importlib.import_module(ext.module_name)
|
|
|
|
stop_fn_name = f"{ext_id}_stop"
|
|
if not hasattr(old_module, stop_fn_name):
|
|
raise ValueError(f"No stop function found for '{ext.module_name}'.")
|
|
|
|
stop_fn = getattr(old_module, stop_fn_name)
|
|
if stop_fn:
|
|
if asyncio.iscoroutinefunction(stop_fn):
|
|
await stop_fn()
|
|
else:
|
|
stop_fn()
|
|
logger.info(f"Stopped background work for extension '{ext.module_name}'.")
|
|
except Exception as ex:
|
|
logger.warning(f"Failed to stop background work for '{ext.module_name}'.")
|
|
logger.warning(ex)
|
|
return False
|
|
|
|
return True
|
|
|
|
|
|
async def start_extension_background_work(ext_id: str) -> bool:
|
|
"""
|
|
Start background work for extension (like asyncio.Tasks, WebSockets, etc).
|
|
Extension CAN expose a `myextension_start()` function if it is starting tasks.
|
|
Extension MUST expose a `myextension_stop()` in that case.
|
|
"""
|
|
upgrade_hash = settings.extension_upgrade_hash(ext_id)
|
|
ext = Extension(code=ext_id, is_valid=True, upgrade_hash=upgrade_hash)
|
|
|
|
try:
|
|
logger.info(f"Starting background work for extension '{ext.module_name}'.")
|
|
new_module = importlib.import_module(ext.module_name)
|
|
start_fn_name = f"{ext_id}_start"
|
|
|
|
# start function is optional, return False if not found
|
|
if not hasattr(new_module, start_fn_name):
|
|
return False
|
|
|
|
start_fn = getattr(new_module, start_fn_name)
|
|
if start_fn:
|
|
if asyncio.iscoroutinefunction(start_fn):
|
|
await start_fn()
|
|
else:
|
|
start_fn()
|
|
logger.info(f"Started background work for extension '{ext.module_name}'.")
|
|
return True
|
|
except Exception as ex:
|
|
logger.warning(f"Failed to start background work for '{ext.module_name}'.")
|
|
logger.warning(ex)
|
|
return False
|
|
|
|
|
|
async def get_valid_extensions(
|
|
include_deactivated: bool | None = True, conn: Connection | None = None
|
|
) -> list[Extension]:
|
|
installed_extensions = await get_installed_extensions(conn=conn)
|
|
valid_extensions = [Extension.from_installable_ext(e) for e in installed_extensions]
|
|
|
|
if include_deactivated:
|
|
return valid_extensions
|
|
|
|
if settings.lnbits_extensions_deactivate_all:
|
|
return []
|
|
|
|
return [
|
|
e
|
|
for e in valid_extensions
|
|
if e.code not in settings.lnbits_deactivated_extensions
|
|
]
|
|
|
|
|
|
async def get_valid_extension(
|
|
ext_id: str, include_deactivated: bool | None = True
|
|
) -> Extension | None:
|
|
ext = await get_installed_extension(ext_id)
|
|
if not ext:
|
|
return None
|
|
|
|
if include_deactivated:
|
|
return Extension.from_installable_ext(ext)
|
|
|
|
if settings.lnbits_extensions_deactivate_all:
|
|
return None
|
|
|
|
return Extension.from_installable_ext(ext)
|