diff --git a/lnbits/core/wasm_ext/__init__.py b/lnbits/core/wasm_ext/__init__.py index a2b554adf..3ac850ed1 100644 --- a/lnbits/core/wasm_ext/__init__.py +++ b/lnbits/core/wasm_ext/__init__.py @@ -1,8 +1,9 @@ -from .api.host import ( - ExtensionAPIMethod, - ExtensionHostAPI, +from .api.host import ExtensionHostAPI +from .api.models import ExtensionAPIMethod, ExtensionAPIMethodExport +from .api.registry import ( extension_api_contract, extension_api_method, + extension_api_permission_ids, get_extension_api_method, list_extension_api_methods, ) @@ -12,10 +13,12 @@ from .wasm.loader import WasmExtension __all__ = [ "ExtensionAPIHost", "ExtensionAPIMethod", + "ExtensionAPIMethodExport", "ExtensionHostAPI", "WasmExtension", "extension_api_contract", "extension_api_method", + "extension_api_permission_ids", "get_extension_api_method", "list_extension_api_methods", ] diff --git a/lnbits/core/wasm_ext/api/__init__.py b/lnbits/core/wasm_ext/api/__init__.py index 011dfe7aa..1ad668781 100644 --- a/lnbits/core/wasm_ext/api/__init__.py +++ b/lnbits/core/wasm_ext/api/__init__.py @@ -1,8 +1,9 @@ -from .host import ( - ExtensionAPIMethod, - ExtensionHostAPI, +from .host import ExtensionHostAPI +from .models import ExtensionAPIMethod, ExtensionAPIMethodExport +from .registry import ( extension_api_contract, extension_api_method, + extension_api_permission_ids, get_extension_api_method, list_extension_api_methods, ) @@ -11,9 +12,11 @@ from .runtime import ExtensionAPIHost __all__ = [ "ExtensionAPIHost", "ExtensionAPIMethod", + "ExtensionAPIMethodExport", "ExtensionHostAPI", "extension_api_contract", "extension_api_method", + "extension_api_permission_ids", "get_extension_api_method", "list_extension_api_methods", ] diff --git a/lnbits/core/wasm_ext/api/host.py b/lnbits/core/wasm_ext/api/host.py index bf39b642a..1e9dc8852 100644 --- a/lnbits/core/wasm_ext/api/host.py +++ b/lnbits/core/wasm_ext/api/host.py @@ -1,16 +1,11 @@ from __future__ import annotations -import inspect import json import logging import secrets import time -from collections.abc import Awaitable, Callable, Iterable, Mapping -from dataclasses import dataclass -from functools import wraps -from typing import Any, TypeVar, cast, get_type_hints - -from pydantic import BaseModel +from collections.abc import Iterable, Mapping +from typing import Any from lnbits.helpers import sha256s @@ -49,93 +44,10 @@ from .models import ( WalletBalanceRequest, WalletBalanceResponse, ) +from .registry import extension_api_method logger = logging.getLogger("lnbits.extensions") -_EXTENSION_API_METHOD_ATTR = "__lnbits_extension_api_method__" -_EXTENSION_RUNTIME_PERMISSION_IDS = {"ui.camera.scan_qr"} -_RequestModel = TypeVar("_RequestModel", bound=BaseModel) -_ResponseModel = TypeVar("_ResponseModel", bound=BaseModel) - - -@dataclass(frozen=True) -class ExtensionAPIMethodExport: - method_id: str - namespace: str - name: str - host_interface: str - host_name: str - sdk_name: str - description: str - required_permission: str | None = None - require_auth: bool = True - - -@dataclass(frozen=True) -class ExtensionAPIMethod: - method_id: str - namespace: str - name: str - python_name: str - host_interface: str - host_name: str - sdk_name: str - description: str - request_model: type[BaseModel] - response_model: type[BaseModel] - required_permission: str | None = None - require_auth: bool = True - - @property - def sdk_qualified_name(self) -> str: - return f"{self.namespace}.{self.sdk_name}" - - -def extension_api_method( - *, - method_id: str, - namespace: str, - name: str, - host_name: str, - sdk_name: str, - description: str, - host_interface: str = "host", - required_permission: str | None = None, - require_auth: bool = True, -) -> Callable[ - [Callable[[Any, _RequestModel], Awaitable[_ResponseModel]]], - Callable[[Any, _RequestModel], Awaitable[_ResponseModel]], -]: - export = ExtensionAPIMethodExport( - method_id=method_id, - namespace=namespace, - name=name, - host_interface=host_interface, - host_name=host_name, - sdk_name=sdk_name, - description=description, - required_permission=required_permission, - require_auth=require_auth, - ) - - def decorator( - function: Callable[[Any, _RequestModel], Awaitable[_ResponseModel]], - ) -> Callable[[Any, _RequestModel], Awaitable[_ResponseModel]]: - @wraps(function) - async def wrapper(self: Any, request: _RequestModel) -> _ResponseModel: - api = getattr(self, "api", self) - if require_auth and not api.has_authenticated_context(): - raise PermissionError( - f"Extension API method '{method_id}' requires authentication." - ) - api.require_permission(required_permission) - return await function(self, request) - - setattr(wrapper, _EXTENSION_API_METHOD_ATTR, export) - return wrapper - - return decorator - class ExtensionHostAPI: def __init__( @@ -154,34 +66,10 @@ class ExtensionHostAPI: self.access_token = access_token self.context = context self.owner_id = sha256s(user_id) if user_id else owner_id - self._uuid = secrets.token_urlsafe(12).replace("-", "_") from .utils import ExtensionAPIUtils self.utils = ExtensionAPIUtils(self) - def __repr__(self) -> str: - return ( - "ExtensionHostAPI(" - f"extension_id={self.extension_id!r}, " - f"context={self.context!r}, " - f"_uuid={self._uuid!r}" - ")" - ) - - def require_permission(self, permission: str | None) -> None: - if permission and permission not in self.permissions: - raise PermissionError( - f"Extension '{self.extension_id}' is missing permission '{permission}'." - ) - - def has_authenticated_context(self) -> bool: - return bool(self.user_id) or self.context == "event" - - def _require_owner_id(self) -> str: - if not self.owner_id: - raise PermissionError("Extension API method requires an owner context.") - return self.owner_id - @extension_api_method( method_id="storage.get", namespace="storage", @@ -223,6 +111,7 @@ class ExtensionHostAPI: for field_name, value in row.items() if field_name in public_fields } + # todo: check public fields filtering return StorageGetResponse(data_json=json.dumps(public_row)) @extension_api_method( @@ -647,131 +536,25 @@ class ExtensionHostAPI: ) return table, wallet_field - -def list_extension_api_methods( - api_cls: type[ExtensionHostAPI] = ExtensionHostAPI, -) -> list[ExtensionAPIMethod]: - methods: list[ExtensionAPIMethod] = [] - - for prefix, method_cls in _extension_api_method_sources(api_cls): - for python_name, function in inspect.getmembers(method_cls, inspect.isfunction): - export = getattr(function, _EXTENSION_API_METHOD_ATTR, None) - if not export: - continue - - request_model, response_model = _get_method_models(function) - methods.append( - ExtensionAPIMethod( - method_id=export.method_id, - namespace=export.namespace, - name=export.name, - python_name=f"{prefix}.{python_name}" if prefix else python_name, - host_interface=export.host_interface, - host_name=export.host_name, - sdk_name=export.sdk_name, - description=export.description, - request_model=request_model, - response_model=response_model, - required_permission=export.required_permission, - require_auth=export.require_auth, - ) + def require_permission(self, permission: str | None) -> None: + if permission and permission not in self.permissions: + raise PermissionError( + f"Extension '{self.extension_id}' is missing permission '{permission}'." ) - return sorted(methods, key=lambda method: method.method_id) + def has_authenticated_context(self) -> bool: + return bool(self.user_id) or self.context == "event" + def _require_owner_id(self) -> str: + if not self.owner_id: + raise PermissionError("Extension API method requires an owner context.") + return self.owner_id -def _extension_api_method_sources( - api_cls: type[ExtensionHostAPI], -) -> list[tuple[str, type[Any]]]: - sources: list[tuple[str, type[Any]]] = [("", api_cls)] - if issubclass(api_cls, ExtensionHostAPI): - from .utils import extension_api_utils_method_classes - - sources.extend(extension_api_utils_method_classes().items()) - return sources - - -def extension_api_permission_ids( - api_cls: type[ExtensionHostAPI] = ExtensionHostAPI, -) -> set[str]: - permissions = { - method.required_permission - for method in list_extension_api_methods(api_cls) - if method.required_permission - } - permissions.update(_EXTENSION_RUNTIME_PERMISSION_IDS) - return permissions - - -def get_extension_api_method( - method_id: str, - api_cls: type[ExtensionHostAPI] = ExtensionHostAPI, -) -> ExtensionAPIMethod: - for method in list_extension_api_methods(api_cls): - if method.method_id == method_id: - return method - raise KeyError(f"Unknown extension API method '{method_id}'.") - - -def extension_api_contract( - api_cls: type[ExtensionHostAPI] = ExtensionHostAPI, -) -> dict[str, object]: - return { - "version": 1, - "methods": [ - { - "id": method.method_id, - "namespace": method.namespace, - "name": method.name, - "python_name": method.python_name, - "host_interface": method.host_interface, - "host_name": method.host_name, - "sdk_name": method.sdk_name, - "sdk_qualified_name": method.sdk_qualified_name, - "description": method.description, - "required_permission": method.required_permission, - "require_auth": method.require_auth, - "request_schema": method.request_model.schema( - ref_template="#/definitions/{model}" - ), - "response_schema": method.response_model.schema( - ref_template="#/definitions/{model}" - ), - } - for method in list_extension_api_methods(api_cls) - ], - } - - -def _get_method_models( - function: Callable[..., object], -) -> tuple[type[BaseModel], type[BaseModel]]: - signature = inspect.signature(function) - request_parameters = [ - parameter - for parameter in signature.parameters.values() - if parameter.name != "self" - ] - if len(request_parameters) != 1: - raise TypeError( - f"Extension API method '{function.__name__}' must accept one request model." + def __repr__(self) -> str: + return ( + "ExtensionHostAPI(" + f"extension_id={self.extension_id!r}, " + f"context={self.context!r}, " + f"owner_id={self.owner_id!r}" + ")" ) - - hints = get_type_hints(function) - request_model = hints.get(request_parameters[0].name) - response_model = hints.get("return") - - if not _is_pydantic_model(request_model): - raise TypeError( - f"Extension API method '{function.__name__}' request must be a BaseModel." - ) - if not _is_pydantic_model(response_model): - raise TypeError( - f"Extension API method '{function.__name__}' response must be a BaseModel." - ) - - return cast(type[BaseModel], request_model), cast(type[BaseModel], response_model) - - -def _is_pydantic_model(value: object) -> bool: - return isinstance(value, type) and issubclass(value, BaseModel) diff --git a/lnbits/core/wasm_ext/api/permissions.py b/lnbits/core/wasm_ext/api/permissions.py index 1d3f345b5..151e608a7 100644 --- a/lnbits/core/wasm_ext/api/permissions.py +++ b/lnbits/core/wasm_ext/api/permissions.py @@ -2,7 +2,7 @@ from collections.abc import Iterable from typing import Any from lnbits.core.models.extensions import ExtensionPermission, InstallableExtension -from lnbits.core.wasm_ext.api.host import extension_api_permission_ids +from lnbits.core.wasm_ext.api.registry import extension_api_permission_ids def validate_extension_permissions( diff --git a/lnbits/core/wasm_ext/api/registry.py b/lnbits/core/wasm_ext/api/registry.py new file mode 100644 index 000000000..20415afa7 --- /dev/null +++ b/lnbits/core/wasm_ext/api/registry.py @@ -0,0 +1,199 @@ +from __future__ import annotations + +import inspect +from collections.abc import Awaitable, Callable +from functools import wraps +from typing import Any, TypeVar, cast, get_type_hints + +from pydantic import BaseModel + +from .models import ExtensionAPIMethod, ExtensionAPIMethodExport + +_EXTENSION_API_METHOD_ATTR = "__lnbits_extension_api_method__" +_EXTENSION_RUNTIME_PERMISSION_IDS = {"ui.camera.scan_qr"} +_RequestModel = TypeVar("_RequestModel", bound=BaseModel) +_ResponseModel = TypeVar("_ResponseModel", bound=BaseModel) + + +def extension_api_method( + *, + method_id: str, + namespace: str, + name: str, + host_name: str, + sdk_name: str, + description: str, + host_interface: str = "host", + required_permission: str | None = None, + require_auth: bool = True, +) -> Callable[ + [Callable[[Any, _RequestModel], Awaitable[_ResponseModel]]], + Callable[[Any, _RequestModel], Awaitable[_ResponseModel]], +]: + export = ExtensionAPIMethodExport( + method_id=method_id, + namespace=namespace, + name=name, + host_interface=host_interface, + host_name=host_name, + sdk_name=sdk_name, + description=description, + required_permission=required_permission, + require_auth=require_auth, + ) + + def decorator( + function: Callable[[Any, _RequestModel], Awaitable[_ResponseModel]], + ) -> Callable[[Any, _RequestModel], Awaitable[_ResponseModel]]: + @wraps(function) + async def wrapper(self: Any, request: _RequestModel) -> _ResponseModel: + api = getattr(self, "api", self) + if require_auth and not api.has_authenticated_context(): + raise PermissionError( + f"Extension API method '{method_id}' requires authentication." + ) + api.require_permission(required_permission) + return await function(self, request) + + setattr(wrapper, _EXTENSION_API_METHOD_ATTR, export) + return wrapper + + return decorator + + +def list_extension_api_methods( + api_cls: type[Any] | None = None, +) -> list[ExtensionAPIMethod]: + api_cls = _default_api_cls(api_cls) + methods: list[ExtensionAPIMethod] = [] + + for prefix, method_cls in _extension_api_method_sources(api_cls): + for python_name, function in inspect.getmembers(method_cls, inspect.isfunction): + export = getattr(function, _EXTENSION_API_METHOD_ATTR, None) + if not export: + continue + + request_model, response_model = _get_method_models(function) + methods.append( + ExtensionAPIMethod( + method_id=export.method_id, + namespace=export.namespace, + name=export.name, + python_name=f"{prefix}.{python_name}" if prefix else python_name, + host_interface=export.host_interface, + host_name=export.host_name, + sdk_name=export.sdk_name, + description=export.description, + request_model=request_model, + response_model=response_model, + required_permission=export.required_permission, + require_auth=export.require_auth, + ) + ) + + return sorted(methods, key=lambda method: method.method_id) + + +def extension_api_permission_ids(api_cls: type[Any] | None = None) -> set[str]: + permissions = { + method.required_permission + for method in list_extension_api_methods(api_cls) + if method.required_permission + } + permissions.update(_EXTENSION_RUNTIME_PERMISSION_IDS) + return permissions + + +def get_extension_api_method( + method_id: str, + api_cls: type[Any] | None = None, +) -> ExtensionAPIMethod: + for method in list_extension_api_methods(api_cls): + if method.method_id == method_id: + return method + raise KeyError(f"Unknown extension API method '{method_id}'.") + + +def extension_api_contract(api_cls: type[Any] | None = None) -> dict[str, object]: + return { + "version": 1, + "methods": [ + { + "id": method.method_id, + "namespace": method.namespace, + "name": method.name, + "python_name": method.python_name, + "host_interface": method.host_interface, + "host_name": method.host_name, + "sdk_name": method.sdk_name, + "sdk_qualified_name": method.sdk_qualified_name, + "description": method.description, + "required_permission": method.required_permission, + "require_auth": method.require_auth, + "request_schema": method.request_model.schema( + ref_template="#/definitions/{model}" + ), + "response_schema": method.response_model.schema( + ref_template="#/definitions/{model}" + ), + } + for method in list_extension_api_methods(api_cls) + ], + } + + +def _default_api_cls(api_cls: type[Any] | None) -> type[Any]: + if api_cls is not None: + return api_cls + + from .host import ExtensionHostAPI + + return ExtensionHostAPI + + +def _extension_api_method_sources( + api_cls: type[Any], +) -> list[tuple[str, type[Any]]]: + sources: list[tuple[str, type[Any]]] = [("", api_cls)] + + from .host import ExtensionHostAPI + + if issubclass(api_cls, ExtensionHostAPI): + from .utils import extension_api_utils_method_classes + + sources.extend(extension_api_utils_method_classes().items()) + return sources + + +def _get_method_models( + function: Callable[..., object], +) -> tuple[type[BaseModel], type[BaseModel]]: + signature = inspect.signature(function) + request_parameters = [ + parameter + for parameter in signature.parameters.values() + if parameter.name != "self" + ] + if len(request_parameters) != 1: + raise TypeError( + f"Extension API method '{function.__name__}' must accept one request model." + ) + + hints = get_type_hints(function) + request_model = hints.get(request_parameters[0].name) + response_model = hints.get("return") + + if not _is_pydantic_model(request_model): + raise TypeError( + f"Extension API method '{function.__name__}' request must be a BaseModel." + ) + if not _is_pydantic_model(response_model): + raise TypeError( + f"Extension API method '{function.__name__}' response must be a BaseModel." + ) + + return cast(type[BaseModel], request_model), cast(type[BaseModel], response_model) + + +def _is_pydantic_model(value: object) -> bool: + return isinstance(value, type) and issubclass(value, BaseModel) diff --git a/lnbits/core/wasm_ext/api/runtime.py b/lnbits/core/wasm_ext/api/runtime.py index b70892cfa..cc9edbb73 100644 --- a/lnbits/core/wasm_ext/api/runtime.py +++ b/lnbits/core/wasm_ext/api/runtime.py @@ -7,7 +7,9 @@ from typing import Any from pydantic import BaseModel -from .host import ExtensionAPIMethod, ExtensionHostAPI, list_extension_api_methods +from .host import ExtensionHostAPI +from .models import ExtensionAPIMethod +from .registry import list_extension_api_methods HostImport = Callable[..., Awaitable[dict[str, Any]]] diff --git a/lnbits/core/wasm_ext/api/utils.py b/lnbits/core/wasm_ext/api/utils.py index 245024452..f0c00f643 100644 --- a/lnbits/core/wasm_ext/api/utils.py +++ b/lnbits/core/wasm_ext/api/utils.py @@ -4,7 +4,6 @@ import time from datetime import datetime from typing import TYPE_CHECKING, Any -from .host import extension_api_method from .models import ( Bolt11Request, CurrencyConvertRequest, @@ -29,6 +28,7 @@ from .models import ( VerifyPreimageRequest, VerifyPreimageResponse, ) +from .registry import extension_api_method if TYPE_CHECKING: from .host import ExtensionHostAPI diff --git a/lnbits/core/wasm_ext/wasm/host.py b/lnbits/core/wasm_ext/wasm/host.py index 2b033a0da..fd9e3e5ba 100644 --- a/lnbits/core/wasm_ext/wasm/host.py +++ b/lnbits/core/wasm_ext/wasm/host.py @@ -5,8 +5,8 @@ import re from collections.abc import Mapping from typing import Any -from ..api.host import list_extension_api_methods from ..api.models import EmptyRequest +from ..api.registry import list_extension_api_methods from ..api.runtime import ExtensionAPIHost