Compare commits

..
Author SHA1 Message Date
dni 39b1f540b5 make it configureable 2026-07-03 14:16:02 +02:00
dadofsambonzukianddni 34bf834dc6 fix: disable uvicorn ws_ping_interval/timeout to prevent WebSocket disconnects
Uvicorn's default ws_ping_interval=20 / ws_ping_timeout=20 can
non-deterministically close local WebSocket connections (e.g. NWC provider
↔ nostrclient relay) under event-loop load. The client side already sends
its own pings, so server-side pings are redundant here.

See discussion:
https://github.com/lnbits/nostrclient/pull/71#issuecomment-4803033307
2026-07-03 13:06:09 +02:00
blackcoffeexbtandGitHub 7db5c986b3 NWC funding source: Reduce polling (#4008) 2026-07-02 13:00:49 +01:00
42 changed files with 1127 additions and 5569 deletions
+2 -26
View File
@@ -24,12 +24,6 @@ from lnbits.core.crud import (
update_installed_extension_state,
)
from lnbits.core.crud.extensions import create_installed_extension
from lnbits.core.extensions.events import dispatch_wasm_invoice_paid
from lnbits.core.extensions.loader import (
is_wasm_extension_dir,
is_wasm_extension_id,
)
from lnbits.core.extensions.routes import register_wasm_extension
from lnbits.core.helpers import migrate_extension_database
from lnbits.core.models.notifications import NotificationType
from lnbits.core.services.extensions import deactivate_extension, get_valid_extensions
@@ -172,7 +166,6 @@ def create_app() -> FastAPI:
# Allow registering new extensions routes without direct access to the `app` object
core_app_extra.register_new_ext_routes = register_new_ext_routes(app)
core_app_extra.register_new_wasm_ext_routes = register_new_wasm_ext_routes(app)
core_app_extra.register_new_ratelimiter = register_new_ratelimiter(app)
# register static files
@@ -317,9 +310,8 @@ async def build_all_installed_extensions_list( # noqa: C901
installed_extensions.append(ext_info)
await create_installed_extension(ext_info)
if not is_wasm_extension_dir(ext_dir):
current_version = await get_db_version(ext_id)
await migrate_extension_database(ext_info, current_version)
current_version = await get_db_version(ext_id)
await migrate_extension_database(ext_info, current_version)
except Exception as e:
logger.warning(e)
@@ -420,13 +412,6 @@ def register_new_ext_routes(app: FastAPI) -> Callable:
return register_new_ext_routes_fn
def register_new_wasm_ext_routes(app: FastAPI) -> Callable:
def register_new_wasm_ext_routes_fn(ext_id: str):
register_wasm_extension(app, ext_id)
return register_new_wasm_ext_routes_fn
def register_new_ratelimiter(app: FastAPI) -> Callable:
def register_new_ratelimiter_fn():
limiter = Limiter(
@@ -480,14 +465,10 @@ async def check_and_register_extensions(app: FastAPI) -> None:
await check_installed_extensions(app)
for ext in await get_valid_extensions(False):
try:
if is_wasm_extension_id(ext.code):
register_wasm_extension(app, ext.code)
continue
register_ext_routes(app, ext)
register_ext_tasks(ext)
except Exception as exc:
logger.error(f"Could not load extension `{ext.code}`: {exc!s}")
await update_installed_extension_state(ext_id=ext.code, active=False)
def register_async_tasks() -> None:
@@ -508,11 +489,6 @@ def register_async_tasks() -> None:
# core invoice listener
invoice_queue: asyncio.Queue = asyncio.Queue()
register_invoice_listener(invoice_queue, "core")
async def dispatch_extension_invoice_paid(payment) -> None:
await dispatch_wasm_invoice_paid(payment)
core_app_extra.dispatch_extension_invoice_paid = dispatch_extension_invoice_paid
create_permanent_task(lambda: wait_for_paid_invoices(invoice_queue))
create_permanent_task(run_by_the_minute_tasks)
-28
View File
@@ -1,28 +0,0 @@
"""Extension runtime contracts."""
from .api import (
ExtensionAPI,
ExtensionAPIMethod,
extension_api_contract,
extension_api_method,
get_extension_api_method,
list_extension_api_methods,
)
from .loader import WasmExtension, load_wasm_extension
from .routes import register_wasm_extension
from .runtime import ExtensionAPIHost
from .wasm import invoke_wasm_extension_export
__all__ = [
"ExtensionAPI",
"ExtensionAPIHost",
"ExtensionAPIMethod",
"WasmExtension",
"extension_api_contract",
"extension_api_method",
"get_extension_api_method",
"invoke_wasm_extension_export",
"list_extension_api_methods",
"load_wasm_extension",
"register_wasm_extension",
]
-777
View File
@@ -1,777 +0,0 @@
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 lnbits.helpers import sha256s
from .models import (
CreateInvoicePublicRequest,
CreateInvoiceRequest,
CreateInvoiceResponse,
EmptyRequest,
ExtensionApiRequest,
HttpRequest,
HttpResponse,
ListUserWalletsResponse,
LogRequest,
LogResponse,
NowResponse,
PayInvoiceRequest,
PayInvoiceResponse,
RandomIdRequest,
RandomIdResponse,
StorageDeleteRequest,
StorageDeleteResponse,
StorageGetRequest,
StorageGetResponse,
StoragePaginatedRequest,
StoragePaginatedResponse,
StorageSetRequest,
StorageSetResponse,
UserWalletSummary,
WalletBalanceRequest,
WalletBalanceResponse,
)
from .storage import (
storage_delete_row,
storage_get_paginated_rows,
storage_get_public_row,
storage_get_row,
storage_set_row,
)
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 ExtensionAPI:
def __init__(
self,
extension_id: str,
permissions: Iterable[Any],
*,
user_id: str | None = None,
access_token: str | None = None,
context: str = "user",
owner_id: str | None = None,
) -> None:
self.extension_id = extension_id
self.permissions, self.permission_policies = self._permission_data(permissions)
self.user_id = user_id
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 .api_utils import ExtensionAPIUtils
self.utils = ExtensionAPIUtils(self)
def __repr__(self) -> str:
return (
"ExtensionAPI("
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",
name="Get storage row",
host_name="storage_get",
sdk_name="get",
description="Read one row from an extension storage table.",
required_permission="ext.storage.read",
require_auth=True,
)
async def storage_get(self, request: StorageGetRequest) -> StorageGetResponse:
row = await storage_get_row(
self.extension_id,
request.table,
request.id,
self._require_owner_id(),
)
return StorageGetResponse(data_json=json.dumps(row) if row else None)
@extension_api_method(
method_id="storage.get_public",
namespace="storage",
name="Get public storage row",
host_name="storage_get_public",
sdk_name="getPublic",
description="Read one public row from an extension storage table.",
required_permission="ext.storage.read_public",
require_auth=False,
)
async def storage_get_public(
self, request: StorageGetRequest
) -> StorageGetResponse:
public_fields = self._public_storage_fields(request.table)
row = await storage_get_public_row(self.extension_id, request.table, request.id)
if not row:
return StorageGetResponse()
public_row = {
field_name: value
for field_name, value in row.items()
if field_name in public_fields
}
return StorageGetResponse(data_json=json.dumps(public_row))
@extension_api_method(
method_id="storage.set",
namespace="storage",
name="Set storage row",
host_name="storage_set",
sdk_name="set",
description="Create or update one row in an extension storage table.",
required_permission="ext.storage.write",
require_auth=True,
)
async def storage_set(self, request: StorageSetRequest) -> StorageSetResponse:
await storage_set_row(
self.extension_id,
request.table,
request.data,
self._require_owner_id(),
)
return StorageSetResponse()
@extension_api_method(
method_id="storage.get_paginated",
namespace="storage",
name="Get paginated storage rows",
host_name="storage_get_paginated",
sdk_name="getPaginated",
description="Get filtered, searched, sorted, paginated storage rows.",
required_permission="ext.storage.read",
require_auth=True,
)
async def storage_get_paginated(
self, request: StoragePaginatedRequest
) -> StoragePaginatedResponse:
page = await storage_get_paginated_rows(
self.extension_id,
request.table,
request.filters,
owner_id=self._require_owner_id(),
search=request.search,
search_fields=request.search_fields,
sort_by=request.sort_by,
descending=request.descending,
limit=request.limit,
offset=request.offset,
)
return StoragePaginatedResponse(
rows_json=json.dumps(page["data"]),
total=page["total"],
)
@extension_api_method(
method_id="storage.delete",
namespace="storage",
name="Delete storage row",
host_name="storage_delete",
sdk_name="delete",
description="Delete one row from an extension storage table.",
required_permission="ext.storage.write",
require_auth=True,
)
async def storage_delete(
self, request: StorageDeleteRequest
) -> StorageDeleteResponse:
await storage_delete_row(
self.extension_id,
request.table,
request.id,
self._require_owner_id(),
)
return StorageDeleteResponse()
@extension_api_method(
method_id="wallet.create_invoice",
namespace="wallet",
name="Create invoice",
host_name="create_invoice",
sdk_name="createInvoice",
description="Create an incoming Lightning invoice for an allowed wallet.",
required_permission="wallet.create_invoice",
require_auth=True,
)
async def wallet_create_invoice(
self, request: CreateInvoiceRequest
) -> CreateInvoiceResponse:
from lnbits.core.crud.wallets import get_wallet
from lnbits.core.models.payments import CreateInvoice
from lnbits.core.services.payments import create_payment_request
if self.user_id:
wallet = await get_wallet(request.wallet_id)
if wallet is None or wallet.user != self.user_id:
raise PermissionError(
"Creating an invoice for this wallet requires an "
"authenticated user context."
)
else:
pass
# todo: security stuff here
payment = await create_payment_request(
request.wallet_id,
CreateInvoice(
amount=request.amount,
unit=request.currency,
memo=request.memo,
extra=request.extra,
extension=self.extension_id,
),
)
return CreateInvoiceResponse(
payment_hash=payment.payment_hash,
payment_request=payment.payment_request or payment.bolt11,
checking_id=payment.checking_id,
)
@extension_api_method(
method_id="wallet.create_invoice_public",
namespace="wallet",
name="Create public invoice",
host_name="create_invoice_public",
sdk_name="createInvoicePublic",
description="Create a public incoming Lightning invoice.",
required_permission="wallet.create_invoice_public",
require_auth=False,
)
async def wallet_create_invoice_public(
self, request: CreateInvoicePublicRequest
) -> CreateInvoiceResponse:
from lnbits.core.models.payments import CreateInvoice
from lnbits.core.services.payments import create_payment_request
table, wallet_field = self._public_invoice_wallet_source()
row = await storage_get_public_row(self.extension_id, table, request.source_id)
if not row:
raise PermissionError("Public invoice source was not found.")
wallet_id = row.get(wallet_field)
if not isinstance(wallet_id, str) or not wallet_id:
raise PermissionError("Public invoice source has no valid wallet.")
payment = await create_payment_request(
wallet_id,
CreateInvoice(
amount=request.amount,
unit=request.currency,
memo=request.memo,
extra={
"tag": self.extension_id,
"source_id": request.source_id,
f"extra_{self.extension_id}": request.extra,
},
extension=self.extension_id,
),
)
return CreateInvoiceResponse(
payment_hash=payment.payment_hash,
payment_request=payment.payment_request or payment.bolt11,
checking_id=payment.checking_id,
)
@extension_api_method(
method_id="wallet.list_user_wallets",
namespace="wallet",
name="List user wallets",
host_name="list_user_wallets",
sdk_name="listUserWallets",
description="List wallets available to the authenticated extension user.",
required_permission="wallet.list",
)
async def wallet_list_user_wallets(
self, request: EmptyRequest
) -> ListUserWalletsResponse:
if not self.user_id:
raise PermissionError(
"Listing user wallets requires an authenticated user context."
)
from lnbits.core.crud.wallets import get_wallets
user_wallets = await get_wallets(self.user_id)
if user_wallets is None:
raise PermissionError(
"Listing user wallets requires an authenticated user context."
)
return ListUserWalletsResponse(
wallets=[
UserWalletSummary(id=w.id, name=w.name, currency=w.currency)
for w in user_wallets
]
)
@extension_api_method(
method_id="wallet.balance",
namespace="wallet",
name="Read wallet balance",
host_name="wallet_balance",
sdk_name="balance",
description="Read the balance of a wallet available to the user.",
required_permission="wallet.balance.read",
)
async def wallet_balance(
self, request: WalletBalanceRequest
) -> WalletBalanceResponse:
from lnbits.core.crud.wallets import get_wallet
if not self.user_id:
raise PermissionError(
"Reading a wallet balance requires an authenticated user context."
)
wallet = await get_wallet(request.wallet_id)
if wallet is None or wallet.user != self.user_id:
raise PermissionError("Reading this wallet balance is not allowed.")
withdrawable_msat = max(wallet.withdrawable_balance, 0)
fee_reserve_msat = max(wallet.balance_msat - withdrawable_msat, 0)
return WalletBalanceResponse(
wallet_id=wallet.id,
name=wallet.name,
currency=wallet.currency,
balance_msat=wallet.balance_msat,
balance_sat=wallet.balance,
withdrawable_msat=withdrawable_msat,
withdrawable_sat=withdrawable_msat // 1000,
fee_reserve_msat=fee_reserve_msat,
fee_reserve_sat=fee_reserve_msat // 1000,
can_send_payments=wallet.can_send_payments,
)
@extension_api_method(
method_id="wallet.pay_invoice",
namespace="wallet",
name="Pay invoice",
host_name="pay_invoice",
sdk_name="payInvoice",
description="Pay a Lightning invoice from a wallet available to the user.",
required_permission="wallet.pay_invoice",
)
async def wallet_pay_invoice(
self, request: PayInvoiceRequest
) -> PayInvoiceResponse:
from lnbits.core.crud.wallets import get_wallet
from lnbits.core.services.payments import pay_invoice
from lnbits.exceptions import PaymentError
if not self.user_id:
raise PermissionError(
"Paying an invoice requires an authenticated user context."
)
wallet = await get_wallet(request.wallet_id)
if wallet is None or wallet.user != self.user_id:
raise PermissionError("Paying invoices from this wallet is not allowed.")
try:
payment = await pay_invoice(
wallet_id=request.wallet_id,
payment_request=request.payment_request,
max_sat=request.max_sat,
extra={"tag": self.extension_id, **request.extra},
description=request.description,
tag=self.extension_id,
)
except (PaymentError, ValueError) as exc:
return PayInvoiceResponse(ok=False, error=str(exc))
return PayInvoiceResponse(
ok=True,
checking_id=payment.checking_id,
payment_hash=payment.payment_hash,
status=payment.status,
amount_msat=abs(payment.amount),
fee_msat=abs(payment.fee),
pending=payment.pending,
success=payment.success,
)
@extension_api_method(
method_id="http.request",
namespace="http",
name="HTTP request",
host_name="http_request",
sdk_name="request",
description="Make an outbound HTTP request to an allowed host.",
required_permission="http.request",
require_auth=True,
)
async def http_request(self, request: HttpRequest) -> HttpResponse:
from .http_client import send_extension_http_request
policy = self.permission_policies.get("http.request") or {}
return await send_extension_http_request(self.extension_id, policy, request)
@extension_api_method(
method_id="extension.api.request",
namespace="extension",
name="Extension API request",
host_name="extension_api_request",
sdk_name="request",
description="Call an allowed installed extension API.",
required_permission="extension.api.request",
require_auth=True,
)
async def extension_api_request(self, request: ExtensionApiRequest) -> HttpResponse:
from .extension_client import send_extension_api_request
policy = self.permission_policies.get("extension.api.request") or {}
return await send_extension_api_request(
self.extension_id,
policy,
self.user_id,
self.access_token,
request,
)
@extension_api_method(
method_id="system.random_id",
namespace="system",
name="Random ID",
host_name="random_id",
sdk_name="id",
description="Create a random extension-local identifier.",
require_auth=False,
)
async def system_random_id(self, request: RandomIdRequest) -> RandomIdResponse:
return RandomIdResponse(
id=f"{request.prefix}_{secrets.token_urlsafe(12).replace('-', '_')}"
)
@extension_api_method(
method_id="system.now",
namespace="system",
name="Current timestamp",
host_name="now",
sdk_name="now",
description="Return the current Unix timestamp.",
require_auth=False,
)
async def system_now(self, request: EmptyRequest) -> NowResponse:
return NowResponse(timestamp=int(time.time()))
@extension_api_method(
method_id="system.log",
namespace="system",
name="Log message",
host_name="log",
sdk_name="log",
description="Write a bounded message to the extension log.",
require_auth=False,
)
async def system_log(self, request: LogRequest) -> LogResponse:
log = getattr(logger, request.level)
log("extension:%s %s", self.extension_id, request.message)
return LogResponse()
@staticmethod
def _permission_data(
permissions: Iterable[Any],
) -> tuple[set[str], dict[str, dict[str, Any]]]:
permission_ids: set[str] = set()
policies: dict[str, dict[str, Any]] = {}
for permission in permissions:
if isinstance(permission, str):
permission_ids.add(permission)
continue
permission_id: str | None = None
policy: Any = None
if isinstance(permission, Mapping):
permission_id = permission.get("id") # type: ignore[assignment]
policy = permission.get("policy")
else:
permission_id = getattr(permission, "id", None)
policy = getattr(permission, "policy", None)
if not permission_id:
continue
permission_ids.add(permission_id)
if isinstance(policy, dict):
policies[permission_id] = policy
return permission_ids, policies
def _public_storage_fields(self, table: str) -> set[str]:
policy = self.permission_policies.get("ext.storage.read_public") or {}
tables = policy.get("tables")
if not isinstance(tables, list):
raise PermissionError(
"Public storage reads require a tables policy for "
"'ext.storage.read_public'."
)
for table_policy in tables:
if not isinstance(table_policy, dict):
continue
if table_policy.get("table_name") != table:
continue
public_fields = table_policy.get("public_fields")
if not isinstance(public_fields, list) or not all(
isinstance(field, str) and field for field in public_fields
):
raise PermissionError(
f"Public storage table '{table}' has no valid public fields."
)
return set(public_fields)
raise PermissionError(f"Storage table '{table}' is not publicly readable.")
def _public_invoice_wallet_source(self) -> tuple[str, str]:
policy = self.permission_policies.get("wallet.create_invoice_public") or {}
table = policy.get("table")
wallet_field = policy.get("wallet_field")
if not isinstance(table, str) or not table:
raise PermissionError(
"Public invoice creation requires a storage table policy."
)
if not isinstance(wallet_field, str) or not wallet_field:
raise PermissionError(
"Public invoice creation requires a wallet field policy."
)
return table, wallet_field
def list_extension_api_methods(
api_cls: type[ExtensionAPI] = ExtensionAPI,
) -> 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,
)
)
return sorted(methods, key=lambda method: method.method_id)
def _extension_api_method_sources(
api_cls: type[ExtensionAPI],
) -> list[tuple[str, type[Any]]]:
sources: list[tuple[str, type[Any]]] = [("", api_cls)]
if issubclass(api_cls, ExtensionAPI):
from .api_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[ExtensionAPI] = ExtensionAPI,
) -> 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[ExtensionAPI] = ExtensionAPI,
) -> 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[ExtensionAPI] = ExtensionAPI,
) -> 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."
)
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)
-381
View File
@@ -1,381 +0,0 @@
from __future__ import annotations
import time
from datetime import datetime
from typing import TYPE_CHECKING, Any
from .api import extension_api_method
from .models import (
Bolt11Request,
CurrencyConvertRequest,
CurrencyConvertResponse,
CurrencyListResponse,
CurrencyRateRequest,
CurrencyRateResponse,
DecodeInvoiceResponse,
EmptyRequest,
FiatToSatsRequest,
FiatToSatsResponse,
InvoiceAmountMsatResponse,
InvoiceExpiryResponse,
InvoiceMemoResponse,
InvoicePaymentHashResponse,
RandomSecretAndHashRequest,
RandomSecretAndHashResponse,
SatsToFiatRequest,
SatsToFiatResponse,
ServerHealthResponse,
ValidateInvoiceResponse,
VerifyPreimageRequest,
VerifyPreimageResponse,
)
if TYPE_CHECKING:
from .api import ExtensionAPI
class ExtensionAPIUtils:
def __init__(self, api: ExtensionAPI) -> None:
self.api = api
self.currencies = ExtensionCurrencyUtils(api)
self.server = ExtensionServerUtils(api)
self.lightning = ExtensionLightningUtils(api)
class _ExtensionAPIUtilsGroup:
def __init__(self, api: ExtensionAPI) -> None:
self.api = api
class ExtensionCurrencyUtils(_ExtensionAPIUtilsGroup):
@extension_api_method(
method_id="utils.currencies.list",
namespace="utils.currencies",
name="List currencies",
host_interface="utils-currencies",
host_name="list_currencies",
sdk_name="list",
description="List currencies supported by LNbits exchange-rate conversion.",
required_permission="utils.basic",
require_auth=False,
)
async def list(self, request: EmptyRequest) -> CurrencyListResponse:
from lnbits.utils.exchange_rates import allowed_currencies
return CurrencyListResponse(currencies=allowed_currencies())
@extension_api_method(
method_id="utils.currencies.rate",
namespace="utils.currencies",
name="Get currency rate",
host_interface="utils-currencies",
host_name="rate",
sdk_name="rate",
description="Get sats-per-fiat and BTC price for a currency.",
required_permission="utils.basic",
require_auth=False,
)
async def rate(self, request: CurrencyRateRequest) -> CurrencyRateResponse:
from lnbits.utils.exchange_rates import get_fiat_rate_and_price_satoshis
rate, price = await get_fiat_rate_and_price_satoshis(request.currency)
return CurrencyRateResponse(rate=rate, price=price)
@extension_api_method(
method_id="utils.currencies.convert",
namespace="utils.currencies",
name="Convert currency amount",
host_interface="utils-currencies",
host_name="convert",
sdk_name="convert",
description="Convert between sats, BTC, and supported fiat currencies.",
required_permission="utils.basic",
require_auth=False,
)
async def convert(self, request: CurrencyConvertRequest) -> CurrencyConvertResponse:
from lnbits.utils.exchange_rates import (
fiat_amount_as_satoshis,
satoshis_amount_as_fiat,
)
from_currency = request.from_currency
if from_currency == "sats":
from_currency = "sat"
amounts: list[tuple[str, float]] = []
if from_currency == "sat":
sats = int(request.amount)
amounts.append(("BTC", sats / 100_000_000))
amounts.append(("sats", sats))
for currency in request.to.split(","):
currency = currency.strip()
if currency:
amounts.append(
(
currency.upper(),
await satoshis_amount_as_fiat(sats, currency),
)
)
else:
sats = await fiat_amount_as_satoshis(request.amount, from_currency)
amounts.append((from_currency.upper(), request.amount))
amounts.append(("sats", sats))
amounts.append(("BTC", sats / 100_000_000))
return CurrencyConvertResponse(amounts=amounts)
@extension_api_method(
method_id="utils.currencies.fiat_to_sats",
namespace="utils.currencies",
name="Convert fiat to sats",
host_interface="utils-currencies",
host_name="fiat_to_sats",
sdk_name="fiatToSats",
description="Convert a fiat amount to sats.",
required_permission="utils.basic",
require_auth=False,
)
async def fiat_to_sats(self, request: FiatToSatsRequest) -> FiatToSatsResponse:
from lnbits.utils.exchange_rates import fiat_amount_as_satoshis
return FiatToSatsResponse(
amount_sat=await fiat_amount_as_satoshis(
request.amount,
request.currency,
)
)
@extension_api_method(
method_id="utils.currencies.sats_to_fiat",
namespace="utils.currencies",
name="Convert sats to fiat",
host_interface="utils-currencies",
host_name="sats_to_fiat",
sdk_name="satsToFiat",
description="Convert a sats amount to fiat.",
required_permission="utils.basic",
require_auth=False,
)
async def sats_to_fiat(self, request: SatsToFiatRequest) -> SatsToFiatResponse:
from lnbits.utils.exchange_rates import satoshis_amount_as_fiat
return SatsToFiatResponse(
amount=await satoshis_amount_as_fiat(request.amount, request.currency)
)
class ExtensionServerUtils(_ExtensionAPIUtilsGroup):
@extension_api_method(
method_id="utils.server.health",
namespace="utils.server",
name="Server health",
host_interface="utils-server",
host_name="health",
sdk_name="health",
description="Return basic public LNbits server health data.",
required_permission="utils.basic",
require_auth=False,
)
async def health(self, request: EmptyRequest) -> ServerHealthResponse:
from lnbits.settings import settings
return ServerHealthResponse(
server_time=int(time.time()),
up_time=settings.lnbits_server_up_time,
)
class ExtensionLightningUtils(_ExtensionAPIUtilsGroup):
@extension_api_method(
method_id="utils.lightning.decode_invoice",
namespace="utils.lightning",
name="Decode Lightning invoice",
host_interface="utils-lightning",
host_name="decode_invoice",
sdk_name="decodeInvoice",
description="Decode a BOLT11 Lightning invoice.",
required_permission="utils.basic",
require_auth=False,
)
async def decode_invoice(self, request: Bolt11Request) -> DecodeInvoiceResponse:
invoice = _decode_bolt11(request.bolt11)
return _decoded_invoice_response(invoice)
@extension_api_method(
method_id="utils.lightning.validate_invoice",
namespace="utils.lightning",
name="Validate Lightning invoice",
host_interface="utils-lightning",
host_name="validate_invoice",
sdk_name="validateInvoice",
description="Validate whether a string is a BOLT11 Lightning invoice.",
required_permission="utils.basic",
require_auth=False,
)
async def validate_invoice(self, request: Bolt11Request) -> ValidateInvoiceResponse:
try:
_decode_bolt11(request.bolt11)
return ValidateInvoiceResponse(valid=True)
except Exception as exc:
return ValidateInvoiceResponse(valid=False, error=str(exc))
@extension_api_method(
method_id="utils.lightning.invoice_payment_hash",
namespace="utils.lightning",
name="Get Lightning invoice payment hash",
host_interface="utils-lightning",
host_name="invoice_payment_hash",
sdk_name="invoicePaymentHash",
description="Get the payment hash from a BOLT11 Lightning invoice.",
required_permission="utils.basic",
require_auth=False,
)
async def invoice_payment_hash(
self, request: Bolt11Request
) -> InvoicePaymentHashResponse:
return InvoicePaymentHashResponse(
payment_hash=str(_decode_bolt11(request.bolt11).payment_hash)
)
@extension_api_method(
method_id="utils.lightning.invoice_amount_msat",
namespace="utils.lightning",
name="Get Lightning invoice amount",
host_interface="utils-lightning",
host_name="invoice_amount_msat",
sdk_name="invoiceAmountMsat",
description="Get the amount in msat from a BOLT11 Lightning invoice.",
required_permission="utils.basic",
require_auth=False,
)
async def invoice_amount_msat(
self, request: Bolt11Request
) -> InvoiceAmountMsatResponse:
return InvoiceAmountMsatResponse(
amount_msat=_invoice_amount_msat(_decode_bolt11(request.bolt11))
)
@extension_api_method(
method_id="utils.lightning.invoice_expiry",
namespace="utils.lightning",
name="Get Lightning invoice expiry",
host_interface="utils-lightning",
host_name="invoice_expiry",
sdk_name="invoiceExpiry",
description="Get the expiry timestamp from a BOLT11 Lightning invoice.",
required_permission="utils.basic",
require_auth=False,
)
async def invoice_expiry(self, request: Bolt11Request) -> InvoiceExpiryResponse:
return InvoiceExpiryResponse(
expires_at=_invoice_expires_at(_decode_bolt11(request.bolt11))
)
@extension_api_method(
method_id="utils.lightning.invoice_memo",
namespace="utils.lightning",
name="Get Lightning invoice memo",
host_interface="utils-lightning",
host_name="invoice_memo",
sdk_name="invoiceMemo",
description="Get the memo from a BOLT11 Lightning invoice.",
required_permission="utils.basic",
require_auth=False,
)
async def invoice_memo(self, request: Bolt11Request) -> InvoiceMemoResponse:
return InvoiceMemoResponse(memo=_invoice_memo(_decode_bolt11(request.bolt11)))
@extension_api_method(
method_id="utils.lightning.verify_preimage",
namespace="utils.lightning",
name="Verify Lightning preimage",
host_interface="utils-lightning",
host_name="verify_preimage",
sdk_name="verifyPreimage",
description="Verify that a preimage matches a payment hash.",
required_permission="utils.basic",
require_auth=False,
)
async def verify_preimage(
self, request: VerifyPreimageRequest
) -> VerifyPreimageResponse:
from lnbits.utils.crypto import verify_preimage
return VerifyPreimageResponse(
valid=verify_preimage(request.preimage, request.payment_hash)
)
@extension_api_method(
method_id="utils.lightning.random_secret_and_hash",
namespace="utils.lightning",
name="Random Lightning secret and hash",
host_interface="utils-lightning",
host_name="random_secret_and_hash",
sdk_name="randomSecretAndHash",
description="Create a random secret and matching SHA256 hash.",
required_permission="utils.basic",
require_auth=False,
)
async def random_secret_and_hash(
self, request: RandomSecretAndHashRequest
) -> RandomSecretAndHashResponse:
from lnbits.utils.crypto import random_secret_and_hash
secret, payment_hash = random_secret_and_hash(request.length)
return RandomSecretAndHashResponse(secret=secret, hash=payment_hash)
def extension_api_utils_method_classes() -> dict[str, type[_ExtensionAPIUtilsGroup]]:
return {
"utils.currencies": ExtensionCurrencyUtils,
"utils.server": ExtensionServerUtils,
"utils.lightning": ExtensionLightningUtils,
}
def _decode_bolt11(payment_request: str) -> Any:
from lnbits import bolt11
return bolt11.decode(payment_request)
def _decoded_invoice_response(invoice: Any) -> DecodeInvoiceResponse:
return DecodeInvoiceResponse(
payment_hash=str(getattr(invoice, "payment_hash", "")) or None,
amount_msat=_invoice_amount_msat(invoice),
expiry=_invoice_expiry(invoice),
expires_at=_invoice_expires_at(invoice),
memo=_invoice_memo(invoice),
)
def _invoice_amount_msat(invoice: Any) -> int | None:
amount_msat = getattr(invoice, "amount_msat", None)
if amount_msat is None:
return None
return int(amount_msat)
def _invoice_expiry(invoice: Any) -> int | None:
expiry = getattr(invoice, "expiry", None)
if expiry is None:
return None
return int(expiry)
def _invoice_expires_at(invoice: Any) -> int | None:
expiry_date = getattr(invoice, "expiry_date", None)
if isinstance(expiry_date, datetime):
return int(expiry_date.timestamp())
date = getattr(invoice, "date", None)
expiry = getattr(invoice, "expiry", None)
if isinstance(date, datetime) and expiry is not None:
return int(date.timestamp() + int(expiry))
if isinstance(date, (int, float)) and expiry is not None:
return int(date + int(expiry))
return None
def _invoice_memo(invoice: Any) -> str | None:
memo = getattr(invoice, "description", None)
return str(memo) if memo is not None else None
-113
View File
@@ -1,113 +0,0 @@
import json
from typing import Any
from loguru import logger
from lnbits.core.db import core_app_extra
async def dispatch_wasm_invoice_paid(payment: Any) -> None:
extension_id = _payment_extension_id(payment)
if not extension_id:
return
extension = core_app_extra.wasm_extension_registry.get(extension_id)
if not extension:
return
export_name = _wasm_invoice_paid_export(extension.config)
if not export_name:
return
if not _is_wasm_event_export(extension, export_name):
logger.warning(
f"WASM extension '{extension.id}' declares invalid onInvoicePaid "
f"export '{export_name}'."
)
return
try:
from lnbits.core.extensions.wasm import invoke_wasm_extension_export
await invoke_wasm_extension_export(
extension.id,
export_name,
_wasm_invoice_paid_payload(payment),
context="event",
owner_id=await _wasm_invoice_paid_owner_id(extension, payment),
)
except Exception as exc:
logger.warning(
f"WASM extension '{extension.id}' failed to handle paid invoice "
f"'{payment.payment_hash}': {exc!s}"
)
def _payment_extension_id(payment: Any) -> str | None:
if isinstance(payment.extension, str) and payment.extension:
return payment.extension
extra = payment.extra or {}
tag = extra.get("tag") or payment.tag
return tag if isinstance(tag, str) and tag else None
async def _wasm_invoice_paid_owner_id(extension: Any, payment: Any) -> str | None:
source_id = _payment_source_id(payment)
source_table = _wasm_public_invoice_source_table(extension.config)
if not source_id or not source_table:
return None
from lnbits.core.extensions.storage import storage_get_row_owner_id
return await storage_get_row_owner_id(extension.id, source_table, source_id)
def _payment_source_id(payment: Any) -> str | None:
extra = payment.extra or {}
source_id = extra.get("source_id")
return source_id if isinstance(source_id, str) and source_id else None
def _wasm_public_invoice_source_table(config: dict[str, Any]) -> str | None:
permissions = config.get("permissions") or []
for permission in permissions:
if not isinstance(permission, dict):
continue
if permission.get("id") != "wallet.create_invoice_public":
continue
policy = permission.get("policy") or {}
table = policy.get("table")
return table if isinstance(table, str) and table else None
return None
def _wasm_invoice_paid_export(config: dict[str, Any]) -> str | None:
events = config.get("events") or {}
export_name = events.get("onInvoicePaid")
return export_name if isinstance(export_name, str) and export_name else None
def _is_wasm_event_export(extension: Any, export_name: str) -> bool:
for export in extension.exports:
if export.get("name") == export_name:
return export.get("visibility") == "event"
return False
def _wasm_invoice_paid_payload(payment: Any) -> dict[str, Any]:
return {
"checkingId": payment.checking_id,
"paymentHash": payment.payment_hash,
"walletId": payment.wallet_id,
"amount": payment.amount,
"fee": payment.fee,
"bolt11": payment.bolt11,
"memo": payment.memo,
"pending": payment.pending,
"status": payment.status,
"tag": payment.tag,
"extension": payment.extension,
"extra": payment.extra or {},
"payment": json.loads(payment.json()),
}
-199
View File
@@ -1,199 +0,0 @@
from __future__ import annotations
import posixpath
import re
from typing import Any
from urllib.parse import unquote, urlsplit, urlunsplit
import httpx
from lnbits.core.crud.extensions import (
get_installed_extension,
get_user_active_extensions_ids,
)
from lnbits.settings import settings
from .models import ExtensionApiRequest, HttpResponse
EXTENSION_API_TIMEOUT_SECONDS = 10.0
EXTENSION_API_MAX_RESPONSE_BYTES = 262_144
_READ_METHODS = {"GET", "HEAD"}
_WRITE_METHODS = {"DELETE", "PATCH", "POST", "PUT"}
_EXTENSION_ID_RE = re.compile(r"^[A-Za-z0-9_-]+$")
_FORBIDDEN_RESPONSE_HEADERS = {
"connection",
"content-length",
"set-cookie",
"transfer-encoding",
}
async def send_extension_api_request(
caller_extension_id: str,
policy: dict[str, Any],
user_id: str | None,
access_token: str | None,
request: ExtensionApiRequest,
) -> HttpResponse:
if not user_id:
raise PermissionError("Extension API requests require authentication.")
if not access_token:
raise PermissionError("Extension API requests require an account access token.")
target_extension_id = _target_extension_id(request.extension_id)
access = _target_extension_access(policy, target_extension_id)
_require_method_access(caller_extension_id, target_extension_id, access, request)
await _require_enabled_extension(target_extension_id, user_id)
path = _extension_api_path(request.path)
body = request.body.encode() if request.body is not None else b""
if len(body) > 65_536:
raise ValueError("Extension API request body is too large.")
url = f"http://{settings.host}:{settings.port}/{target_extension_id}{path}"
try:
async with httpx.AsyncClient(
follow_redirects=False,
timeout=EXTENSION_API_TIMEOUT_SECONDS,
trust_env=False,
) as client:
async with client.stream(
request.method,
url,
headers={"Authorization": f"Bearer {access_token}"},
content=body,
) as response:
response_body = await _read_limited_response(response)
return HttpResponse(
status_code=response.status_code,
headers=_response_headers(dict(response.headers)),
body=response_body.decode(response.encoding or "utf-8", "replace"),
)
except httpx.RequestError as exc:
raise ValueError("Extension API request failed.") from exc
def _target_extension_id(extension_id: str) -> str:
target = extension_id.strip()
if not target or not _EXTENSION_ID_RE.match(target):
raise PermissionError("Extension API request has an invalid target extension.")
return target
def _target_extension_access(
policy: dict[str, Any], target_extension_id: str
) -> set[str]:
extensions = policy.get("extensions")
if not isinstance(extensions, list) or not extensions:
raise PermissionError(
"Extension API requests require a non-empty extensions policy."
)
for extension in extensions:
if isinstance(extension, str):
extension_id = extension
access = ["read"]
elif isinstance(extension, dict):
raw_extension_id = extension.get("id")
raw_access = extension.get("access")
if not isinstance(raw_extension_id, str):
continue
if not isinstance(raw_access, list):
raise PermissionError(
f"Extension API target '{target_extension_id}' "
"has no access policy."
)
extension_id = raw_extension_id
access = raw_access
else:
continue
if extension_id != target_extension_id:
continue
clean_access = {
item
for item in access
if isinstance(item, str) and item in {"read", "write"}
}
if clean_access:
return clean_access
break
raise PermissionError(
f"Extension API target '{target_extension_id}' is not allowed."
)
def _require_method_access(
caller_extension_id: str,
target_extension_id: str,
access: set[str],
request: ExtensionApiRequest,
) -> None:
if request.method in _READ_METHODS:
required_access = "read"
elif request.method in _WRITE_METHODS:
required_access = "write"
else:
raise PermissionError("Extension API request method is not allowed.")
if required_access not in access:
raise PermissionError(
f"Extension '{caller_extension_id}' cannot {required_access} "
f"extension '{target_extension_id}'."
)
async def _require_enabled_extension(target_extension_id: str, user_id: str) -> None:
extension = await get_installed_extension(target_extension_id)
if not extension or not extension.active:
raise PermissionError(
f"Target extension '{target_extension_id}' is not installed or enabled."
)
active_extensions = await get_user_active_extensions_ids(user_id)
if target_extension_id not in active_extensions:
raise PermissionError(
f"Target extension '{target_extension_id}' is not active for this user."
)
def _extension_api_path(path: str) -> str:
parts = urlsplit(path)
if parts.scheme or parts.netloc:
raise PermissionError("Extension API request path must be relative.")
if parts.fragment:
raise PermissionError("Extension API request path cannot include a fragment.")
if not parts.path.startswith("/api/"):
raise PermissionError("Extension API request path must start with '/api/'.")
decoded_path = unquote(parts.path)
path_parts = decoded_path.split("/")
if any(part == ".." for part in path_parts):
raise PermissionError("Extension API request path cannot traverse directories.")
normalized = posixpath.normpath(decoded_path)
if normalized != decoded_path.rstrip("/") or not normalized.startswith("/api/"):
raise PermissionError("Extension API request path is invalid.")
return urlunsplit(("", "", parts.path, parts.query, ""))
async def _read_limited_response(response: httpx.Response) -> bytes:
chunks: list[bytes] = []
size = 0
async for chunk in response.aiter_bytes():
size += len(chunk)
if size > EXTENSION_API_MAX_RESPONSE_BYTES:
raise ValueError("Extension API response is too large.")
chunks.append(chunk)
return b"".join(chunks)
def _response_headers(headers: dict[str, str]) -> dict[str, str]:
return {
key: value
for key, value in headers.items()
if key.lower() not in _FORBIDDEN_RESPONSE_HEADERS
}
-183
View File
@@ -1,183 +0,0 @@
from __future__ import annotations
import ipaddress
import socket
from typing import Any
from urllib.parse import urlparse
import httpx
from .models import HttpRequest, HttpResponse
HTTP_REQUEST_TIMEOUT_SECONDS = 10.0
HTTP_MAX_RESPONSE_BYTES = 262_144
_FORBIDDEN_REQUEST_HEADERS = {
"connection",
"content-length",
"cookie",
"host",
"proxy-authorization",
"transfer-encoding",
}
_FORBIDDEN_RESPONSE_HEADERS = {
"connection",
"content-length",
"set-cookie",
"transfer-encoding",
}
async def send_extension_http_request(
extension_id: str,
policy: dict[str, Any],
request: HttpRequest,
) -> HttpResponse:
allowed_origins = _allowed_origins(policy)
origin = _request_origin(request.url)
if origin not in allowed_origins:
raise PermissionError(
f"Extension '{extension_id}' is not allowed to request '{origin}'."
)
await _reject_internal_host(request.url)
headers = _request_headers(request.headers)
body = request.body.encode() if request.body is not None else b""
if len(body) > 65_536:
raise ValueError("HTTP request body is too large.")
try:
async with httpx.AsyncClient(
follow_redirects=False,
timeout=HTTP_REQUEST_TIMEOUT_SECONDS,
trust_env=False,
) as client:
async with client.stream(
request.method,
request.url,
headers=headers,
content=body,
) as response:
response_body = await _read_limited_response(response)
return HttpResponse(
status_code=response.status_code,
headers=_response_headers(dict(response.headers)),
body=response_body.decode(response.encoding or "utf-8", "replace"),
)
except httpx.RequestError as exc:
raise ValueError("HTTP request failed.") from exc
def _allowed_origins(policy: dict[str, Any]) -> set[str]:
hosts = policy.get("hosts")
if not isinstance(hosts, list) or not hosts:
raise PermissionError("HTTP requests require a non-empty hosts policy.")
origins: set[str] = set()
for host in hosts:
if not isinstance(host, str) or not host:
continue
origins.add(_request_origin(host))
if not origins:
raise PermissionError("HTTP requests require at least one valid host.")
return origins
def _request_origin(url: str) -> str:
parsed = urlparse(url)
if parsed.scheme != "https":
raise PermissionError("HTTP requests require https URLs.")
if parsed.username or parsed.password:
raise PermissionError("HTTP requests cannot include credentials in URLs.")
if not parsed.hostname:
raise PermissionError("HTTP requests require a hostname.")
hostname = parsed.hostname.lower()
port = _url_port(parsed)
if port is None or port == 443:
return f"https://{hostname}"
return f"https://{hostname}:{port}"
def _url_port(parsed: Any) -> int | None:
try:
return parsed.port
except ValueError as exc:
raise PermissionError("HTTP request URL has an invalid port.") from exc
async def _reject_internal_host(url: str) -> None:
parsed = urlparse(url)
hostname = parsed.hostname
if not hostname:
raise PermissionError("HTTP requests require a hostname.")
if hostname == "localhost" or hostname.endswith(".localhost"):
raise PermissionError("HTTP requests cannot target localhost.")
try:
address = ipaddress.ip_address(hostname)
_reject_internal_address(address)
return
except ValueError:
pass
for address in await _resolve_host(hostname):
_reject_internal_address(address)
async def _resolve_host(
hostname: str,
) -> list[ipaddress.IPv4Address | ipaddress.IPv6Address]:
import asyncio
def resolve() -> list[ipaddress.IPv4Address | ipaddress.IPv6Address]:
try:
infos = socket.getaddrinfo(hostname, None, type=socket.SOCK_STREAM)
except socket.gaierror as exc:
raise PermissionError("HTTP request host could not be resolved.") from exc
addresses: list[ipaddress.IPv4Address | ipaddress.IPv6Address] = []
for info in infos:
sockaddr = info[4]
addresses.append(ipaddress.ip_address(sockaddr[0]))
return addresses
return await asyncio.to_thread(resolve)
def _reject_internal_address(
address: ipaddress.IPv4Address | ipaddress.IPv6Address,
) -> None:
if not address.is_global:
raise PermissionError("HTTP requests cannot target internal network addresses.")
def _request_headers(headers: dict[str, str]) -> dict[str, str]:
clean: dict[str, str] = {}
for key, value in headers.items():
header = key.strip()
if not header:
continue
if header.lower() in _FORBIDDEN_REQUEST_HEADERS:
continue
clean[header] = value
return clean
async def _read_limited_response(response: httpx.Response) -> bytes:
chunks: list[bytes] = []
size = 0
async for chunk in response.aiter_bytes():
size += len(chunk)
if size > HTTP_MAX_RESPONSE_BYTES:
raise ValueError("HTTP response is too large.")
chunks.append(chunk)
return b"".join(chunks)
def _response_headers(headers: dict[str, str]) -> dict[str, str]:
return {
key: value
for key, value in headers.items()
if key.lower() not in _FORBIDDEN_RESPONSE_HEADERS
}
-100
View File
@@ -1,100 +0,0 @@
from __future__ import annotations
import json
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from lnbits.settings import settings
@dataclass(frozen=True)
class WasmExtension:
id: str
name: str
version: str
root_path: Path
module_path: Path
wit_path: Path | None
world: str
host_api: str
exports: list[dict[str, Any]]
config: dict[str, Any]
def is_wasm_extension_id(ext_id: str) -> bool:
config = load_wasm_extension_config(ext_id)
return bool(config and config.get("extension_type") == "wasm")
def is_wasm_extension_dir(ext_dir: Path) -> bool:
config = _load_json(ext_dir / "config.json")
return bool(config and config.get("extension_type") == "wasm")
def load_wasm_extension_config(ext_id: str) -> dict[str, Any] | None:
ext_dir = Path(settings.lnbits_extensions_path, "extensions", ext_id)
return _load_json(ext_dir / "config.json")
def load_wasm_extension(ext_id: str) -> WasmExtension:
ext_dir = Path(settings.lnbits_extensions_path, "extensions", ext_id)
config = load_wasm_extension_config(ext_id)
if not config:
raise FileNotFoundError(f"Missing WASM extension config for '{ext_id}'.")
if config.get("extension_type") != "wasm":
raise ValueError(f"Extension '{ext_id}' is not a WASM extension.")
wasm_config = config.get("wasm") or {}
module_path = _extension_path(ext_dir, wasm_config.get("module"))
wit_path = _optional_extension_path(ext_dir, wasm_config.get("wit"))
_check_wasm_module(module_path)
if wit_path and not wit_path.is_file():
raise FileNotFoundError(f"WIT file not found: {wit_path}")
return WasmExtension(
id=config.get("id") or ext_id,
name=config.get("name") or ext_id,
version=config.get("version") or "0.0",
root_path=ext_dir,
module_path=module_path,
wit_path=wit_path,
world=wasm_config.get("world") or "",
host_api=wasm_config.get("host_api") or "lnbits.core.extensions.ExtensionAPI",
exports=wasm_config.get("exports") or [],
config=config,
)
def _load_json(path: Path) -> dict[str, Any] | None:
if not path.is_file():
return None
with path.open("r", encoding="utf-8") as config_file:
value = json.load(config_file)
if not isinstance(value, dict):
raise ValueError(f"Expected JSON object in '{path}'.")
return value
def _extension_path(ext_dir: Path, value: Any) -> Path:
if not isinstance(value, str) or not value:
raise ValueError(f"Missing relative path for extension '{ext_dir.name}'.")
path = (ext_dir / value).resolve()
if ext_dir.resolve() not in path.parents:
raise ValueError(f"Extension path escapes extension root: {value}")
return path
def _optional_extension_path(ext_dir: Path, value: Any) -> Path | None:
if value is None:
return None
return _extension_path(ext_dir, value)
def _check_wasm_module(path: Path) -> None:
if not path.is_file():
raise FileNotFoundError(f"WASM module not found: {path}")
with path.open("rb") as wasm_file:
magic = wasm_file.read(4)
if magic != b"\0asm":
raise ValueError(f"Invalid WASM module: {path}")
-321
View File
@@ -1,321 +0,0 @@
import json
from typing import Any, Literal
from pydantic import BaseModel, Field, root_validator
class EmptyRequest(BaseModel):
pass
class StorageGetRequest(BaseModel):
table: str = Field(..., min_length=1, max_length=128)
id: str = Field(..., min_length=1, max_length=512)
class StorageGetResponse(BaseModel):
data_json: str | None = None
class StorageSetRequest(BaseModel):
table: str = Field(..., min_length=1, max_length=128)
data: dict[str, Any] = Field(default_factory=dict)
@root_validator(pre=True)
def parse_data_json(cls, values: dict[str, Any]) -> dict[str, Any]:
data_json = values.get("data_json")
if data_json is not None and "data" not in values:
values["data"] = json.loads(data_json)
return values
class StorageSetResponse(BaseModel):
ok: bool = True
class StoragePaginatedRequest(BaseModel):
table: str = Field(..., min_length=1, max_length=128)
filters: dict[str, Any] = Field(default_factory=dict)
search: str | None = Field(None, max_length=256)
search_fields: list[str] = Field(default_factory=list)
sort_by: str | None = Field(None, min_length=1, max_length=128)
descending: bool = False
limit: int = Field(25, ge=1, le=1000)
offset: int = Field(0, ge=0)
@root_validator(pre=True)
def parse_json_fields(cls, values: dict[str, Any]) -> dict[str, Any]:
filters_json = values.get("filters_json")
if filters_json is not None and "filters" not in values:
values["filters"] = json.loads(filters_json)
search_fields_json = values.get("search_fields_json")
if search_fields_json is not None and "search_fields" not in values:
values["search_fields"] = json.loads(search_fields_json)
if values.get("sort_by") == "":
values["sort_by"] = None
return values
class StoragePaginatedResponse(BaseModel):
rows_json: str = "[]"
total: int = 0
class StorageDeleteRequest(BaseModel):
table: str = Field(..., min_length=1, max_length=128)
id: str = Field(..., min_length=1, max_length=512)
class StorageDeleteResponse(BaseModel):
ok: bool = True
class CreateInvoiceRequest(BaseModel):
wallet_id: str = Field(..., min_length=1, max_length=128)
amount: float = Field(..., gt=0)
currency: str = Field("sat", min_length=1, max_length=8)
memo: str = Field(..., max_length=512)
tag: str = Field(..., min_length=1, max_length=64)
extra: dict[str, str] = Field(default_factory=dict)
class CreateInvoicePublicRequest(BaseModel):
source_id: str = Field(..., min_length=1, max_length=512)
amount: float = Field(..., gt=0)
currency: str = Field(..., min_length=1, max_length=8)
memo: str = Field("", max_length=512)
extra: dict[str, Any] = Field(default_factory=dict)
@root_validator
def validate_extra_size(cls, values: dict[str, Any]) -> dict[str, Any]:
extra = values.get("extra") or {}
try:
encoded = json.dumps(extra, separators=(",", ":"))
except TypeError as exc:
raise ValueError("extra must be JSON serializable.") from exc
if len(encoded.encode()) > 4096:
raise ValueError("extra must not exceed 4096 bytes.")
values["extra"] = extra
return values
class CreateInvoiceResponse(BaseModel):
payment_hash: str
payment_request: str
checking_id: str
class UserWalletSummary(BaseModel):
id: str
name: str
currency: str | None = None
class ListUserWalletsResponse(BaseModel):
wallets: list[UserWalletSummary] = Field(default_factory=list)
class WalletBalanceRequest(BaseModel):
wallet_id: str = Field(..., min_length=1, max_length=128)
class WalletBalanceResponse(BaseModel):
wallet_id: str
name: str
currency: str | None = None
balance_msat: int
balance_sat: int
withdrawable_msat: int
withdrawable_sat: int
fee_reserve_msat: int
fee_reserve_sat: int
can_send_payments: bool
class PayInvoiceRequest(BaseModel):
wallet_id: str = Field(..., min_length=1, max_length=128)
payment_request: str = Field(..., min_length=1, max_length=8192)
max_sat: int | None = Field(None, gt=0)
description: str = Field("", max_length=512)
extra: dict[str, str] = Field(default_factory=dict)
class PayInvoiceResponse(BaseModel):
ok: bool = True
error: str | None = None
checking_id: str | None = None
payment_hash: str | None = None
status: str | None = None
amount_msat: int = 0
fee_msat: int = 0
pending: bool = False
success: bool = False
class HttpRequest(BaseModel):
method: Literal["DELETE", "GET", "HEAD", "PATCH", "POST", "PUT"] = "GET"
url: str = Field(..., min_length=1, max_length=2048)
headers: dict[str, str] = Field(default_factory=dict)
body: str | None = Field(None, max_length=65536)
@root_validator(pre=True)
def normalize_method(cls, values: dict[str, Any]) -> dict[str, Any]:
method = values.get("method")
if isinstance(method, str):
values["method"] = method.upper()
return values
@root_validator
def validate_headers_size(cls, values: dict[str, Any]) -> dict[str, Any]:
headers = values.get("headers") or {}
if len(headers) > 32:
raise ValueError("headers must not contain more than 32 entries.")
for key, value in headers.items():
if len(key) > 128 or len(value) > 4096:
raise ValueError("headers are too large.")
values["headers"] = headers
return values
class HttpResponse(BaseModel):
status_code: int
headers: dict[str, str] = Field(default_factory=dict)
body: str = ""
class ExtensionApiRequest(BaseModel):
extension_id: str = Field(..., min_length=1, max_length=128)
method: Literal["DELETE", "GET", "HEAD", "PATCH", "POST", "PUT"] = "GET"
path: str = Field(..., min_length=1, max_length=2048)
body: str | None = Field(None, max_length=65536)
@root_validator(pre=True)
def normalize_method(cls, values: dict[str, Any]) -> dict[str, Any]:
method = values.get("method")
if isinstance(method, str):
values["method"] = method.upper()
return values
class CurrencyListResponse(BaseModel):
currencies: list[str] = Field(default_factory=list)
class CurrencyRateRequest(BaseModel):
currency: str = Field(..., min_length=1, max_length=8)
class CurrencyRateResponse(BaseModel):
rate: float
price: float
class CurrencyConvertRequest(BaseModel):
amount: float = Field(..., gt=0)
from_currency: str = Field(..., alias="from", min_length=1, max_length=8)
to: str = Field(..., min_length=1, max_length=256)
class Config:
allow_population_by_field_name = True
class CurrencyConvertResponse(BaseModel):
amounts: list[tuple[str, float]] = Field(default_factory=list)
class FiatToSatsRequest(BaseModel):
amount: float = Field(..., gt=0)
currency: str = Field(..., min_length=1, max_length=8)
class FiatToSatsResponse(BaseModel):
amount_sat: int
class SatsToFiatRequest(BaseModel):
amount: float = Field(..., gt=0)
currency: str = Field(..., min_length=1, max_length=8)
class SatsToFiatResponse(BaseModel):
amount: float
class ServerHealthResponse(BaseModel):
server_time: int
up_time: str
class Bolt11Request(BaseModel):
bolt11: str = Field(..., min_length=1, max_length=8192)
class DecodeInvoiceResponse(BaseModel):
valid: bool = True
payment_hash: str | None = None
amount_msat: int | None = None
expiry: int | None = None
expires_at: int | None = None
memo: str | None = None
class ValidateInvoiceResponse(BaseModel):
valid: bool
error: str | None = None
class InvoicePaymentHashResponse(BaseModel):
payment_hash: str
class InvoiceAmountMsatResponse(BaseModel):
amount_msat: int | None = None
class InvoiceExpiryResponse(BaseModel):
expires_at: int | None = None
class InvoiceMemoResponse(BaseModel):
memo: str | None = None
class VerifyPreimageRequest(BaseModel):
preimage: str = Field(..., min_length=64, max_length=64)
payment_hash: str = Field(..., min_length=64, max_length=64)
class VerifyPreimageResponse(BaseModel):
valid: bool
class RandomSecretAndHashRequest(BaseModel):
length: int = Field(32, ge=16, le=64)
class RandomSecretAndHashResponse(BaseModel):
secret: str
hash: str
class RandomIdRequest(BaseModel):
prefix: str = Field(..., min_length=1, max_length=32)
class RandomIdResponse(BaseModel):
id: str
class NowResponse(BaseModel):
timestamp: int
class LogRequest(BaseModel):
level: Literal["debug", "info", "warning", "error"] = "info"
message: str = Field(..., min_length=1, max_length=2048)
class LogResponse(BaseModel):
ok: bool = True
-59
View File
@@ -1,59 +0,0 @@
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
-777
View File
@@ -1,777 +0,0 @@
from __future__ import annotations
import json
import os
import re
from pathlib import Path
from typing import Annotated, Any, NoReturn
from uuid import uuid4
from fastapi import Depends, FastAPI, HTTPException, Request
from fastapi.responses import FileResponse, Response
from fastapi.staticfiles import StaticFiles
from loguru import logger
from pydantic import UUID4
from starlette.staticfiles import PathLike as StaticFilesPathLike
from starlette.types import Scope
from lnbits.core.crud import get_installed_extension, get_user_from_account
from lnbits.core.db import core_app_extra
from lnbits.core.models import Account
from lnbits.decorators import (
check_access_token,
check_account_exists,
optional_user_id,
)
from lnbits.helpers import template_renderer
from lnbits.settings import settings
from lnbits.utils.cache import cache
from .loader import WasmExtension, load_wasm_extension
from .wasm import invoke_wasm_extension_export, warm_wasm_extension
WASM_FRAME_TOKEN_EXPIRY_SECONDS = 60
WASM_EXTENSION_CORE_ASSET_PREFIX = "_lnbits"
WASM_EXTENSION_CORE_STATIC_ASSETS = {
"bundle.min.css": ("static/bundle.min.css", "text/css; charset=utf-8"),
"material-icons-v50.woff2": (
"static/fonts/material-icons-v50.woff2",
"font/woff2",
),
"quasar.css": ("static/vendor/quasar.css", "text/css; charset=utf-8"),
"quasar.umd.prod.js": (
"static/vendor/quasar.umd.prod.js",
"text/javascript; charset=utf-8",
),
"qrcode.vue.browser.js": (
"static/vendor/qrcode.vue.browser.js",
"text/javascript; charset=utf-8",
),
"vue.global.prod.js": (
"static/vendor/vue.global.prod.js",
"text/javascript; charset=utf-8",
),
}
WASM_EXTENSION_GENERATED_CORE_ASSETS = {
"material-icons.css": (
"""
@font-face {
font-family: 'Material Icons';
font-style: normal;
font-weight: 400;
src: url('./material-icons-v50.woff2') format('woff2');
}
""",
"text/css; charset=utf-8",
)
}
WASM_EXTENSION_STATIC_MIME_TYPES = {
".css": "text/css; charset=utf-8",
".gif": "image/gif",
".ico": "image/x-icon",
".jpeg": "image/jpeg",
".jpg": "image/jpeg",
".js": "text/javascript; charset=utf-8",
".png": "image/png",
".webp": "image/webp",
".woff": "font/woff",
".woff2": "font/woff2",
}
WASM_EXTENSION_TEXT_STATIC_EXTENSIONS = {".css", ".js"}
WASM_EXTENSION_HTML_PREFIXES = (b"<!doctype", b"<html", b"<script")
class GuardedWasmExtensionStaticFiles(StaticFiles):
async def get_response(self, path: str, scope: Scope) -> Response:
if path.startswith(f"{WASM_EXTENSION_CORE_ASSET_PREFIX}/"):
return _wasm_extension_core_asset_response(path)
if Path(path).suffix.lower() not in WASM_EXTENSION_STATIC_MIME_TYPES:
raise HTTPException(status_code=404)
return await super().get_response(path, scope)
def file_response(
self,
full_path: StaticFilesPathLike,
stat_result: os.stat_result,
scope: Scope,
status_code: int = 200,
) -> Response:
suffix = Path(full_path).suffix.lower()
if suffix in WASM_EXTENSION_TEXT_STATIC_EXTENSIONS:
_reject_html_like_wasm_static_asset(Path(full_path))
response = super().file_response(full_path, stat_result, scope, status_code)
response.headers["Content-Type"] = WASM_EXTENSION_STATIC_MIME_TYPES[suffix]
response.headers["X-Content-Type-Options"] = "nosniff"
response.headers["Cache-Control"] = "no-store"
return response
def register_wasm_extension(app: FastAPI, ext_id: str) -> WasmExtension:
loaded = load_wasm_extension(ext_id)
warm_wasm_extension(loaded)
_mount_wasm_extension_static(app, loaded)
_register_wasm_extension_ui_routes(app, loaded)
_register_wasm_extension_api_routes(app, loaded)
core_app_extra.wasm_extension_registry.register(loaded)
settings.activate_extension_paths(ext_id, "", [])
logger.info(
f"Loaded WASM extension '{loaded.id}' "
f"({loaded.module_path.stat().st_size} bytes)."
)
return loaded
def _mount_wasm_extension_static(app: FastAPI, extension: WasmExtension) -> None:
static_path = extension.root_path / "static"
mount_path = f"/ext-assets/{extension.id}"
if any(getattr(route, "path", None) == mount_path for route in app.routes):
return
app.mount(
mount_path,
GuardedWasmExtensionStaticFiles(directory=static_path, check_dir=False),
name=f"{extension.id}-static",
)
def _register_wasm_extension_ui_routes(app: FastAPI, extension: WasmExtension) -> None:
_add_wasm_extension_frame_config_route(app, extension)
for route_index, route_config in enumerate(extension.config.get("ui_routes") or []):
route_path = _wasm_extension_ui_route_path(extension, route_config.get("path"))
entrypoint = _wasm_extension_entrypoint(
extension, route_config.get("entrypoint")
)
frame_path = f"/ext-frame/{extension.id}/{route_index}"
auth = _wasm_extension_route_auth(extension, route_config.get("auth"))
_add_wasm_extension_frame_route(app, extension, frame_path, entrypoint)
_add_wasm_extension_wrapper_route(
app,
extension,
route_path,
auth,
)
def _register_wasm_extension_api_routes(app: FastAPI, extension: WasmExtension) -> None:
for route_config in extension.config.get("api_routes") or []:
_add_wasm_extension_api_route(app, extension, route_config)
def _add_wasm_extension_api_route(
app: FastAPI,
extension: WasmExtension,
route_config: dict[str, Any],
) -> None:
method = _wasm_extension_api_method(extension, route_config.get("method"))
route_path = _wasm_extension_api_path(extension, route_config.get("path"))
export_name = _wasm_extension_api_export(extension, route_config.get("export"))
path_params = route_config.get("path_params") or {}
auth = _wasm_extension_route_auth(extension, route_config.get("auth"))
if _has_route(app, route_path, method):
return
async def invoke_wasm_api_request(
request: Request,
account: Account | None = None,
access_token: str | None = None,
) -> dict[str, Any]:
try:
payload = await _read_api_payload(request, path_params)
return await invoke_wasm_extension_export(
extension.id,
export_name,
payload,
user=account,
access_token=access_token,
)
except KeyError as exc:
raise HTTPException(status_code=404, detail=str(exc)) from exc
except PermissionError as exc:
raise HTTPException(status_code=403, detail=str(exc)) from exc
except (TypeError, ValueError) as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
async def invoke_private_wasm_extension_export(
request: Request,
access_token: Annotated[str | None, Depends(check_access_token)],
account: Account = Depends(check_account_exists),
) -> dict[str, Any]:
return await invoke_wasm_api_request(request, account, access_token)
async def invoke_public_wasm_extension_export(request: Request) -> dict[str, Any]:
return await invoke_wasm_api_request(request)
app.add_api_route(
route_path,
(
invoke_public_wasm_extension_export
if auth == "public"
else invoke_private_wasm_extension_export
),
methods=[method],
name=f"{extension.id}:{method}:{route_path}",
include_in_schema=False,
)
def _add_wasm_extension_frame_config_route(
app: FastAPI,
extension: WasmExtension,
) -> None:
route_path = _wasm_extension_frame_config_path(extension)
if _has_route(app, route_path, "POST"):
return
async def create_wasm_extension_frame_config(
request: Request,
access_token: Annotated[str | None, Depends(check_access_token)],
usr: UUID4 | None = None,
) -> dict[str, Any]:
try:
body = await _read_json_object(request)
except (TypeError, ValueError) as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
ui_route = _match_wasm_extension_ui_route(extension, body.get("path"))
auth = ui_route["auth"]
if auth == "user":
account = await check_account_exists(request, access_token, usr)
user_id: str | None = account.id
else:
user_id = await _optional_wasm_user_id(request, access_token, usr)
granted_permission_ids = await _wasm_extension_granted_permission_ids(extension)
return _wasm_extension_frame_config(
extension,
ui_route["frame_path"],
auth,
ui_route["path_params"],
ui_route["route_params"],
_read_wasm_extension_route_query(body.get("query")),
user_id,
granted_permission_ids,
)
app.add_api_route(
route_path,
create_wasm_extension_frame_config,
methods=["POST"],
name=f"{extension.id}:frame-config",
include_in_schema=False,
)
async def _read_api_payload(
request: Request,
path_params: dict[str, str],
) -> dict[str, Any]:
payload = _read_api_path_params(request, path_params)
payload.update(_read_api_query_params(request))
if request.method in {"POST", "PUT", "PATCH"}:
payload.update(await _read_json_object(request))
return payload
async def _read_json_object(request: Request) -> dict[str, Any]:
body = await request.body()
if not body:
return {}
value = json.loads(body)
if not isinstance(value, dict):
raise TypeError("WASM extension API payload must be a JSON object.")
return value
def _read_api_path_params(
request: Request,
path_params: dict[str, str],
) -> dict[str, Any]:
payload: dict[str, Any] = {}
for key, value in request.path_params.items():
target = path_params.get(key) or _snake_to_camel(key)
payload[target] = value
return payload
def _read_api_query_params(request: Request) -> dict[str, Any]:
return {_snake_to_camel(key): value for key, value in request.query_params.items()}
def _wasm_extension_api_export(extension: WasmExtension, export_name: Any) -> str:
if not isinstance(export_name, str) or not export_name:
raise ValueError(f"Invalid API export for WASM extension '{extension.id}'.")
for export in extension.exports:
if export.get("name") != export_name:
continue
if export.get("visibility") in {"public", "authenticated"}:
return export_name
raise PermissionError(f"WASM export '{export_name}' is not callable over HTTP.")
raise KeyError(f"WASM extension '{extension.id}' has no export '{export_name}'.")
def _wasm_extension_api_method(extension: WasmExtension, method: Any) -> str:
if not isinstance(method, str):
raise ValueError(f"Invalid API method for WASM extension '{extension.id}'.")
method = method.upper()
if method not in {"GET", "POST", "PUT", "PATCH", "DELETE"}:
raise ValueError(f"Unsupported API method for WASM extension '{extension.id}'.")
return method
def _wasm_extension_api_path(extension: WasmExtension, path: Any) -> str:
if not isinstance(path, str) or not path.startswith("/"):
raise ValueError(f"Invalid API path for WASM extension '{extension.id}'.")
if path == "/":
return f"/api/v1/ext/{extension.id}"
return f"/api/v1/ext/{extension.id}{path}"
def _has_route(app: FastAPI, route_path: str, method: str) -> bool:
for route in app.routes:
if getattr(route, "path", None) != route_path:
continue
methods = getattr(route, "methods", set()) or set()
if method in methods:
return True
return False
def _snake_to_camel(value: str) -> str:
head, *tail = value.split("_")
return head + "".join(part.capitalize() for part in tail)
def _add_wasm_extension_wrapper_route(
app: FastAPI,
extension: WasmExtension,
route_path: str,
auth: str,
) -> None:
if _has_route(app, route_path, "GET"):
return
async def serve_private_wasm_extension_page(
request: Request,
account: Account = Depends(check_account_exists),
) -> Any:
user = await get_user_from_account(account)
return _wasm_extension_wrapper_response(
request,
extension,
auth,
user.json() if user else None,
)
async def serve_public_wasm_extension_page(request: Request) -> Any:
return _wasm_extension_wrapper_response(
request,
extension,
auth,
None,
)
app.add_api_route(
route_path,
(
serve_public_wasm_extension_page
if auth == "public"
else serve_private_wasm_extension_page
),
methods=["GET"],
name=f"{extension.id}:{route_path}",
include_in_schema=False,
)
def _add_wasm_extension_frame_route(
app: FastAPI,
extension: WasmExtension,
frame_path: str,
entrypoint: Path,
) -> None:
if _has_route(app, frame_path, "GET"):
return
async def serve_wasm_extension_frame(
request: Request,
user_id: str | None = Depends(_optional_wasm_user_id),
) -> FileResponse:
_consume_wasm_extension_frame_token(request, extension, frame_path, user_id)
response = FileResponse(entrypoint)
response.headers["Content-Security-Policy"] = _wasm_extension_frame_csp(
request, extension
)
response.headers["Cache-Control"] = "no-store"
response.headers["Cross-Origin-Opener-Policy"] = "same-origin"
response.headers["Cross-Origin-Resource-Policy"] = "same-origin"
# Extension access goes through the parent bridge.
response.headers["Permissions-Policy"] = (
"camera=(), microphone=(), geolocation=(), payment=(), "
"clipboard-read=(), usb=()"
)
response.headers["Referrer-Policy"] = "no-referrer"
response.headers["X-Content-Type-Options"] = "nosniff"
return response
app.add_api_route(
frame_path,
serve_wasm_extension_frame,
methods=["GET"],
name=f"{extension.id}:frame:{frame_path}",
include_in_schema=False,
)
def _wasm_extension_wrapper_response(
request: Request,
extension: WasmExtension,
auth: str,
user_json: str | None,
) -> Any:
public = auth == "public"
response = template_renderer().TemplateResponse(
request,
"wasm_extension.html",
{
"extension": extension,
"public": public,
"user": user_json,
},
)
response.headers["Content-Security-Policy"] = "frame-ancestors 'self'"
response.headers["X-Frame-Options"] = "SAMEORIGIN"
return response
def _wasm_extension_frame_csp(request: Request, extension: WasmExtension) -> str:
origin = str(request.base_url).rstrip("/")
extension_assets = f"{origin}/ext-assets/{extension.id}/"
return (
"sandbox allow-scripts; "
"default-src 'none'; "
f"script-src {extension_assets}; "
"script-src-attr 'none'; "
f"style-src {extension_assets}; "
"style-src-attr 'none'; "
f"img-src {extension_assets} data:; "
f"font-src {extension_assets}; "
"connect-src 'none'; "
"form-action 'none'; "
"object-src 'none'; "
"base-uri 'none'; "
"frame-src 'none'; "
"worker-src 'none'; "
"media-src 'none'; "
"manifest-src 'none'; "
"frame-ancestors 'self'"
)
def _wasm_extension_frame_url(
extension: WasmExtension, frame_path: str, user_id: str | None
) -> str:
token = _create_wasm_extension_frame_token(extension, frame_path, user_id)
return f"{frame_path}?frame_token={token}"
def _create_wasm_extension_frame_token(
extension: WasmExtension,
frame_path: str,
user_id: str | None,
) -> str:
token = uuid4().hex
cache.set(
_wasm_extension_frame_token_cache_key(token),
{
"extension_id": extension.id,
"frame_path": frame_path,
"user_id": user_id,
},
expiry=WASM_FRAME_TOKEN_EXPIRY_SECONDS,
)
return token
def _consume_wasm_extension_frame_token(
request: Request,
extension: WasmExtension,
frame_path: str,
user_id: str | None,
) -> None:
token = request.query_params.get("frame_token")
if not token:
_raise_wasm_extension_frame_not_found(extension, frame_path, "missing")
cache_key = _wasm_extension_frame_token_cache_key(token)
token_data = cache.get(cache_key)
if (
not isinstance(token_data, dict)
or token_data.get("extension_id") != extension.id
or token_data.get("frame_path") != frame_path
):
_raise_wasm_extension_frame_not_found(
extension, frame_path, "unknown or expired"
)
token_user_id = token_data.get("user_id")
if token_user_id and token_user_id != user_id:
_raise_wasm_extension_frame_not_found(extension, frame_path, "wrong user")
cache.pop(cache_key)
def _wasm_extension_frame_token_cache_key(token: str) -> str:
return f"wasm-frame-token:{token}"
def _raise_wasm_extension_frame_not_found(
extension: WasmExtension,
frame_path: str,
reason: str,
) -> NoReturn:
logger.warning(
f"WASM frame token {reason} for extension '{extension.id}' at '{frame_path}'."
)
raise HTTPException(status_code=404, detail="Not found")
def _wasm_extension_bridge_api_routes(
extension: WasmExtension,
public: bool,
) -> list[dict[str, str]]:
routes: list[dict[str, str]] = []
for route_config in extension.config.get("api_routes") or []:
auth = _wasm_extension_route_auth(extension, route_config.get("auth"))
if public and auth != "public":
continue
method = _wasm_extension_api_method(extension, route_config.get("method"))
path = _wasm_extension_api_path(extension, route_config.get("path"))
_wasm_extension_api_export(extension, route_config.get("export"))
routes.append(
{
"method": method,
"path": path,
"pattern": _path_template_pattern(path),
}
)
return routes
def _wasm_extension_frame_config_path(extension: WasmExtension) -> str:
return f"/api/v1/ext/{extension.id}/_ui/frame"
def _match_wasm_extension_ui_route(
extension: WasmExtension,
path: Any,
) -> dict[str, Any]:
if not isinstance(path, str) or not path.startswith("/"):
raise HTTPException(status_code=404, detail="Not found")
for route_index, route_config in enumerate(extension.config.get("ui_routes") or []):
route_path = _wasm_extension_ui_route_path(extension, route_config.get("path"))
route_params = _path_template_params(route_path, path)
if route_params is None:
continue
return {
"frame_path": f"/ext-frame/{extension.id}/{route_index}",
"auth": _wasm_extension_route_auth(extension, route_config.get("auth")),
"path_params": route_config.get("path_params") or {},
"route_params": route_params,
}
raise HTTPException(status_code=404, detail="Not found")
def _path_template_params(template: str, path: str) -> dict[str, str] | None:
template_parts = _path_parts(template)
path_parts = _path_parts(path)
if len(template_parts) != len(path_parts):
return None
params: dict[str, str] = {}
for template_part, path_part in zip(template_parts, path_parts, strict=False):
if template_part.startswith("{") and template_part.endswith("}"):
param_name = template_part[1:-1]
if not param_name:
return None
params[param_name] = path_part
continue
if template_part != path_part:
return None
return params
def _path_parts(path: str) -> list[str]:
return [part for part in path.strip("/").split("/") if part]
def _wasm_extension_frame_config(
extension: WasmExtension,
frame_path: str,
auth: str,
path_params: dict[str, str],
route_params: dict[str, str],
query: dict[str, Any],
user_id: str | None,
permissions: set[str],
) -> dict[str, Any]:
public = auth == "public"
return {
"extension": {
"id": extension.id,
"name": extension.name,
},
"frameUrl": _wasm_extension_frame_url(extension, frame_path, user_id),
"bridge": {
"extensionId": extension.id,
"public": public,
"routeParams": _map_wasm_extension_route_params(route_params, path_params),
"query": query,
"permissions": sorted(permissions),
"apiRoutes": _wasm_extension_bridge_api_routes(extension, public),
},
}
async def _wasm_extension_granted_permission_ids(
extension: WasmExtension,
) -> set[str]:
installed_extension = await get_installed_extension(extension.id)
if not installed_extension:
return set()
return {permission.id for permission in installed_extension.permissions}
def _map_wasm_extension_route_params(
route_params: dict[str, str],
path_params: dict[str, str],
) -> dict[str, str]:
payload: dict[str, str] = {}
for key, value in route_params.items():
target = path_params.get(key) or _snake_to_camel(key)
payload[target] = value
return payload
def _read_wasm_extension_route_query(query: Any) -> dict[str, Any]:
if not isinstance(query, dict):
return {}
payload: dict[str, Any] = {}
for key, value in query.items():
if value is None:
continue
payload[_snake_to_camel(str(key))] = value
return payload
def _path_template_pattern(path: str) -> str:
pattern = re.sub(r"\\{[^/{}]+\\}", r"[^/]+", re.escape(path))
return f"^{pattern}$"
async def _optional_wasm_user_id(
request: Request,
access_token: Annotated[str | None, Depends(check_access_token)],
usr: UUID4 | None = None,
) -> str | None:
try:
return await optional_user_id(request, access_token, usr)
except HTTPException:
return None
def _wasm_extension_route_auth(extension: WasmExtension, auth: Any) -> str:
if auth in {"public", "user"}:
return auth
raise ValueError(f"Invalid route auth for WASM extension '{extension.id}'.")
def _wasm_extension_ui_route_path(extension: WasmExtension, path: Any) -> str:
if not isinstance(path, str) or not path.startswith("/"):
raise ValueError(f"Invalid route path for WASM extension '{extension.id}'.")
if path == "/":
return "/ext"
return f"/ext{path}"
def _wasm_extension_entrypoint(extension: WasmExtension, entrypoint: Any) -> Path:
if not isinstance(entrypoint, str) or not entrypoint:
raise ValueError(
f"Invalid route entrypoint for WASM extension '{extension.id}'."
)
if entrypoint.startswith("/"):
raise ValueError(
f"Route entrypoint for WASM extension '{extension.id}' must be a "
"relative extension path."
)
path = (extension.root_path / entrypoint).resolve()
root_path = extension.root_path.resolve()
if path != root_path and root_path not in path.parents:
raise ValueError(f"Route entrypoint escapes extension root: {entrypoint}")
static_path = (extension.root_path / "static").resolve()
if path == static_path or static_path in path.parents:
raise ValueError(
f"Route entrypoint for WASM extension '{extension.id}' must not be "
"inside the static asset directory."
)
if path.suffix.lower() != ".html":
raise ValueError(
f"Route entrypoint for WASM extension '{extension.id}' must be "
"an HTML file."
)
if not path.is_file():
raise FileNotFoundError(f"Route entrypoint not found: {path}")
return path
def _reject_html_like_wasm_static_asset(path: Path) -> None:
with path.open("rb") as asset_file:
prefix = asset_file.read(512).lstrip().lower()
if prefix.startswith(WASM_EXTENSION_HTML_PREFIXES):
raise HTTPException(status_code=404)
def _wasm_extension_core_asset_response(path: str) -> Response:
asset_name = path.removeprefix(f"{WASM_EXTENSION_CORE_ASSET_PREFIX}/")
if not asset_name or "/" in asset_name or "\\" in asset_name:
raise HTTPException(status_code=404)
generated_asset = WASM_EXTENSION_GENERATED_CORE_ASSETS.get(asset_name)
if generated_asset:
content, content_type = generated_asset
response = Response(content=content, media_type=content_type)
response.headers["X-Content-Type-Options"] = "nosniff"
response.headers["Cache-Control"] = "no-store"
return response
asset_config = WASM_EXTENSION_CORE_STATIC_ASSETS.get(asset_name)
if not asset_config:
raise HTTPException(status_code=404)
relative_path, content_type = asset_config
asset_path = Path(settings.lnbits_path, relative_path)
if not asset_path.is_file():
raise HTTPException(status_code=404)
response = FileResponse(asset_path)
response.headers["Content-Type"] = content_type
response.headers["X-Content-Type-Options"] = "nosniff"
response.headers["Cache-Control"] = "no-store"
return response
-136
View File
@@ -1,136 +0,0 @@
from __future__ import annotations
import inspect
import re
from collections.abc import Awaitable, Callable, Mapping
from typing import Any
from pydantic import BaseModel
from .api import ExtensionAPI, ExtensionAPIMethod, list_extension_api_methods
HostImport = Callable[..., Awaitable[dict[str, Any]]]
class ExtensionAPIHost:
def __init__(
self,
api: ExtensionAPI,
*,
api_cls: type[ExtensionAPI] = ExtensionAPI,
) -> None:
self.api = api
self.methods = list_extension_api_methods(api_cls)
self._methods_by_host_name = self._index_methods(self.methods)
async def invoke(
self,
host_name: str,
payload: Mapping[str, Any] | BaseModel | None = None,
) -> dict[str, Any]:
method = self._require_method(host_name)
request = self._request_model(method, payload)
handler = _resolve_attr_path(self.api, method.python_name)
response = handler(request)
if inspect.isawaitable(response):
response = await response
return self._response_payload(method, response)
def imports(self) -> dict[str, HostImport]:
return self.imports_for_interface("host")
def import_object(self) -> dict[str, dict[str, HostImport]]:
interfaces = sorted({method.host_interface for method in self.methods})
return {
f"lnbits:extension/{interface}": self.imports_for_interface(interface)
for interface in interfaces
}
def imports_for_interface(self, host_interface: str) -> dict[str, HostImport]:
return {
_snake_to_camel(method.host_name): self._make_import(method)
for method in self.methods
if method.host_interface == host_interface
}
def _make_import(self, method: ExtensionAPIMethod) -> HostImport:
async def host_import(
payload: Mapping[str, Any] | BaseModel | None = None,
) -> dict[str, Any]:
return await self.invoke(method.method_id, payload)
return host_import
def _require_method(self, host_name: str) -> ExtensionAPIMethod:
method = self._methods_by_host_name.get(host_name)
if not method:
raise KeyError(f"Unknown extension host function '{host_name}'.")
return method
@staticmethod
def _index_methods(
methods: list[ExtensionAPIMethod],
) -> dict[str, ExtensionAPIMethod]:
index: dict[str, ExtensionAPIMethod] = {}
for method in methods:
for host_name in {
method.method_id,
f"{method.host_interface}:{method.host_name}",
method.host_name,
_snake_to_camel(method.host_name),
method.host_name.replace("_", "-"),
}:
index[host_name] = method
return index
@staticmethod
def _request_model(
method: ExtensionAPIMethod,
payload: Mapping[str, Any] | BaseModel | None,
) -> BaseModel:
if isinstance(payload, method.request_model):
return payload
if isinstance(payload, BaseModel):
payload = payload.dict()
if payload is None:
payload = {}
if not isinstance(payload, Mapping):
raise TypeError(
f"Host function '{method.host_name}' expects an object payload."
)
data = {_to_snake(key): value for key, value in payload.items()}
if isinstance(data.get("extra"), list):
data["extra"] = dict(data["extra"])
if isinstance(data.get("headers"), list):
data["headers"] = dict(data["headers"])
return method.request_model.parse_obj(data)
@staticmethod
def _response_payload(
method: ExtensionAPIMethod,
response: Any,
) -> dict[str, Any]:
if not isinstance(response, method.response_model):
response = method.response_model.parse_obj(response)
payload = response.dict()
if method.method_id in {"http.request", "extension.api.request"} and isinstance(
payload.get("headers"), Mapping
):
payload["headers"] = list(payload["headers"].items())
return {_snake_to_camel(key): value for key, value in payload.items()}
def _snake_to_camel(value: str) -> str:
head, *tail = value.split("_")
return head + "".join(part.capitalize() for part in tail)
def _to_snake(value: str) -> str:
value = value.replace("-", "_")
return re.sub(r"([a-z0-9])([A-Z])", r"\1_\2", value).lower()
def _resolve_attr_path(value: Any, path: str) -> Any:
for part in path.split("."):
value = getattr(value, part)
return value
-627
View File
@@ -1,627 +0,0 @@
from __future__ import annotations
import json
import re
from datetime import datetime, timezone
from pathlib import Path
from typing import Any
from loguru import logger
from lnbits.core.crud import update_migration_version
from lnbits.core.db import db as core_db
from lnbits.core.models import DbVersion
from lnbits.core.models.extensions import InstallableExtension
from lnbits.db import POSTGRES, SQLITE, Connection, Database
from lnbits.settings import settings
_MIGRATION_FILE_RE = re.compile(r"^(\d+)_.*\.json$")
_SQL_IDENTIFIER_RE = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$")
OWNER_ID_FIELD = "__lnbits_owner_id__"
async def storage_get_row(
ext_id: str,
table: str,
row_id: str,
owner_id: str,
) -> dict[str, Any] | None:
table_schema = _load_table_schema(ext_id, table)
query = f"""
SELECT * FROM {_table_ref_for_schema(ext_id, table)}
WHERE id = :id AND {OWNER_ID_FIELD} = :owner_id
""" # noqa: S608
async with Database(f"ext_{ext_id}").connect() as conn:
row = await conn.fetchone(query, {"id": row_id, "owner_id": owner_id})
return _row_from_db(table_schema, row) if row else None
async def storage_get_public_row(
ext_id: str,
table: str,
row_id: str,
) -> dict[str, Any] | None:
table_schema = _load_table_schema(ext_id, table)
query = f"""
SELECT * FROM {_table_ref_for_schema(ext_id, table)}
WHERE id = :id
""" # noqa: S608
async with Database(f"ext_{ext_id}").connect() as conn:
row = await conn.fetchone(query, {"id": row_id})
return _row_from_db(table_schema, row) if row else None
async def storage_get_row_owner_id(
ext_id: str,
table: str,
row_id: str,
) -> str | None:
_load_table_schema(ext_id, table)
query = f"""
SELECT {OWNER_ID_FIELD} FROM {_table_ref_for_schema(ext_id, table)}
WHERE id = :id
""" # noqa: S608
async with Database(f"ext_{ext_id}").connect() as conn:
row = await conn.fetchone(query, {"id": row_id})
owner_id = row[OWNER_ID_FIELD] if row else None
return owner_id if isinstance(owner_id, str) and owner_id else None
async def storage_set_row(
ext_id: str,
table: str,
data: dict[str, Any],
owner_id: str,
) -> None:
table_schema = _load_table_schema(ext_id, table)
clean_data = _data_to_db(table_schema, data, require_id=True)
clean_data[OWNER_ID_FIELD] = owner_id
columns = list(clean_data.keys())
placeholders = [f":{column}" for column in columns]
updates = [
f"{column} = excluded.{column}"
for column in columns
if column not in ("id", OWNER_ID_FIELD)
]
conflict_sql = (
"DO UPDATE SET "
+ ", ".join(updates)
+ f" WHERE {OWNER_ID_FIELD} = :{OWNER_ID_FIELD}"
if updates
else "DO NOTHING"
)
query = f"""
INSERT INTO {_table_ref_for_schema(ext_id, table)}
({", ".join(columns)})
VALUES
({", ".join(placeholders)})
ON CONFLICT (id) {conflict_sql}
""" # noqa: S608
async with Database(f"ext_{ext_id}").connect() as conn:
await conn.execute(query, clean_data)
async def storage_get_paginated_rows(
ext_id: str,
table: str,
filters: dict[str, Any],
*,
owner_id: str,
search: str | None,
search_fields: list[str],
sort_by: str | None,
descending: bool,
limit: int,
offset: int,
) -> dict[str, Any]:
table_schema = _load_table_schema(ext_id, table)
where_sql, values = _where_sql(table_schema, filters, search, search_fields)
where_sql = _append_owner_where_sql(where_sql)
values[OWNER_ID_FIELD] = owner_id
order_sql = _order_sql(table_schema, sort_by, descending)
count_values = dict(values)
values.update({"limit": min(limit, 1000), "offset": offset})
table_ref = _table_ref_for_schema(ext_id, table)
rows_query = f"""
SELECT * FROM {table_ref}
{where_sql}
{order_sql}
LIMIT :limit
OFFSET :offset
""" # noqa: S608
count_query = f"""
SELECT COUNT(*) AS count FROM {table_ref}
{where_sql}
""" # noqa: S608
async with Database(f"ext_{ext_id}").connect() as conn:
rows = await conn.fetchall(rows_query, values)
count_row = await conn.fetchone(count_query, count_values)
return {
"data": [_row_from_db(table_schema, row) for row in rows],
"total": int(count_row["count"]) if count_row else 0,
}
async def storage_delete_row(
ext_id: str,
table: str,
row_id: str,
owner_id: str,
) -> None:
_load_table_schema(ext_id, table)
query = f"""
DELETE FROM {_table_ref_for_schema(ext_id, table)}
WHERE id = :id AND {OWNER_ID_FIELD} = :owner_id
""" # noqa: S608
async with Database(f"ext_{ext_id}").connect() as conn:
await conn.execute(query, {"id": row_id, "owner_id": owner_id})
async def migrate_wasm_extension_database(
ext: InstallableExtension,
current_version: DbVersion | None = None,
) -> None:
migrations_dir = ext.ext_dir / "storage" / "migrations"
migration_files = _migration_files(migrations_dir)
if not migration_files:
logger.debug(f"No storage migrations for WASM extension '{ext.id}'.")
return
ext_db = Database(f"ext_{ext.id}")
async with ext_db.connect() as conn:
for version, path in migration_files:
if current_version and version <= current_version.version:
continue
logger.debug(f"running WASM storage migration {ext.id}.{version}")
print(f"running migration {ext.id}.{version}")
await _run_storage_migration(conn, path)
await _update_wasm_migration_version(conn, ext.id, version)
def _migration_files(migrations_dir: Path) -> list[tuple[int, Path]]:
if not migrations_dir.is_dir():
return []
files: list[tuple[int, Path]] = []
for path in migrations_dir.glob("*.json"):
match = _MIGRATION_FILE_RE.match(path.name)
if not match:
raise ValueError(f"Invalid WASM storage migration filename: {path.name}")
files.append((int(match.group(1)), path))
return sorted(files)
async def _run_storage_migration(db: Connection, path: Path) -> None:
migration = _load_json(path)
operations = migration.get("operations")
if not isinstance(operations, list):
raise ValueError(f"WASM storage migration '{path}' has no operations list.")
for operation in operations:
if not isinstance(operation, dict):
raise ValueError(f"WASM storage migration '{path}' has invalid operation.")
sql = _operation_sql(db, operation)
await db.execute(sql)
def _operation_sql(db: Connection, operation: dict[str, Any]) -> str:
op = operation.get("op")
if op == "create_table":
return _create_table_sql(db, operation)
if op == "add_field":
return _add_field_sql(db, operation)
if op == "create_index":
return _create_index_sql(db, operation)
raise ValueError(f"Unsupported WASM storage migration operation: {op}")
def _create_table_sql(db: Connection, operation: dict[str, Any]) -> str:
table = _require_identifier(operation, "table")
fields = _require_fields(operation)
if not any(field.get("name") == "id" for field in fields):
raise ValueError(f"WASM storage table '{table}' must define an id field.")
if any(field.get("name") == OWNER_ID_FIELD for field in fields):
raise ValueError(
f"WASM storage table '{table}' defines reserved field '{OWNER_ID_FIELD}'."
)
columns = [
_column_sql(db, field, primary_key=field.get("name") == "id")
for field in fields
]
columns.append(f"{OWNER_ID_FIELD} TEXT NOT NULL")
return f"""
CREATE TABLE IF NOT EXISTS {_table_ref(db, table)} (
{", ".join(columns)}
);
"""
def _add_field_sql(db: Connection, operation: dict[str, Any]) -> str:
table = _require_identifier(operation, "table")
field = _field_from_add_field_operation(operation)
if field["name"] == OWNER_ID_FIELD:
raise ValueError(
f"WASM storage table '{table}' cannot add reserved field "
f"'{OWNER_ID_FIELD}'."
)
return f"""
ALTER TABLE {_table_ref(db, table)}
ADD COLUMN {_column_sql(db, field)};
"""
def _create_index_sql(db: Connection, operation: dict[str, Any]) -> str:
table = _require_identifier(operation, "table")
name = _require_identifier(operation, "name")
field = _require_identifier(operation, "field")
if field == OWNER_ID_FIELD:
raise ValueError(
f"WASM storage table '{table}' cannot index reserved field "
f"'{OWNER_ID_FIELD}'."
)
if db.type == SQLITE and db.schema:
return f"""
CREATE INDEX IF NOT EXISTS {_schema_ref(db, name)}
ON {table} ({field});
"""
return f"""
CREATE INDEX IF NOT EXISTS {name}
ON {_table_ref(db, table)} ({field});
"""
def _column_sql(
db: Connection,
field: dict[str, Any],
*,
primary_key: bool = False,
) -> str:
name = _require_identifier(field, "name")
column_type = _field_type_sql(db, field)
parts = [name, column_type]
if primary_key:
parts.append("PRIMARY KEY")
elif not field.get("nullable", False):
parts.append("NOT NULL")
if "default" in field:
parts.append(f"DEFAULT {_default_sql(field['default'])}")
return " ".join(parts)
def _field_type_sql(db: Connection, field: dict[str, Any]) -> str:
if field.get("list") is True:
return "TEXT"
field_type = field.get("type")
if field_type == "string":
return "TEXT"
if field_type == "integer":
return db.big_int
if field_type == "number":
return "DOUBLE PRECISION" if db.type == POSTGRES else "REAL"
if field_type == "boolean":
return "BOOLEAN"
if field_type == "datetime":
return "TIMESTAMP"
raise ValueError(f"Unsupported WASM storage field type: {field_type}")
def _load_table_schema(ext_id: str, table: str) -> dict[str, Any]:
schema = _load_storage_schema(ext_id)
tables = schema.get("tables")
if not isinstance(tables, dict):
raise ValueError(f"WASM extension '{ext_id}' has no storage tables schema.")
_require_identifier({"table": table}, "table")
table_schema = tables.get(table)
if not isinstance(table_schema, dict):
raise ValueError(f"WASM extension '{ext_id}' has no storage table '{table}'.")
fields = table_schema.get("fields")
if not isinstance(fields, list) or not fields:
raise ValueError(f"WASM storage table '{table}' has no fields schema.")
for field in fields:
if not isinstance(field, dict):
raise ValueError(f"WASM storage table '{table}' has invalid field schema.")
_require_identifier(field, "name")
if field["name"] == OWNER_ID_FIELD:
raise ValueError(
f"WASM storage table '{table}' defines reserved field "
f"'{OWNER_ID_FIELD}'."
)
return table_schema
def _load_storage_schema(ext_id: str) -> dict[str, Any]:
schema_path = (
Path(settings.lnbits_extensions_path)
/ "extensions"
/ ext_id
/ "storage"
/ "schema.json"
)
if not schema_path.is_file():
raise ValueError(f"WASM extension '{ext_id}' has no storage schema.")
return _load_json(schema_path)
def _data_to_db(
table_schema: dict[str, Any],
data: dict[str, Any],
*,
require_id: bool,
) -> dict[str, Any]:
if not isinstance(data, dict):
raise ValueError("WASM storage row data must be an object.")
if require_id and not data.get("id"):
raise ValueError("WASM storage row data must include an id.")
_reject_reserved_owner_field(data, "row")
fields = _fields_by_name(table_schema)
unknown_fields = sorted(set(data) - set(fields))
if unknown_fields:
raise ValueError(
"WASM storage row has unknown fields: " + ", ".join(unknown_fields)
)
return {
field_name: _value_to_db(fields[field_name], value)
for field_name, value in data.items()
}
def _filters_to_db(
table_schema: dict[str, Any],
filters: dict[str, Any],
) -> dict[str, Any]:
if not isinstance(filters, dict):
raise ValueError("WASM storage filters must be an object.")
_reject_reserved_owner_field(filters, "filters")
fields = _fields_by_name(table_schema)
unknown_fields = sorted(set(filters) - set(fields))
if unknown_fields:
raise ValueError(
"WASM storage filters have unknown fields: " + ", ".join(unknown_fields)
)
return {
field_name: _value_to_db(fields[field_name], value)
for field_name, value in filters.items()
}
def _where_sql(
table_schema: dict[str, Any],
filters: dict[str, Any],
search: str | None,
search_fields: list[str],
) -> tuple[str, dict[str, Any]]:
clean_filters = _filters_to_db(table_schema, filters)
clauses = [f"{field} = :filter_{field}" for field in clean_filters]
values = {f"filter_{field}": value for field, value in clean_filters.items()}
clean_search = search.strip().lower() if search else ""
if clean_search:
fields = _fields_by_name(table_schema)
invalid_fields = sorted(set(search_fields) - set(fields))
if invalid_fields:
raise ValueError(
"WASM storage search has unknown fields: " + ", ".join(invalid_fields)
)
if search_fields:
search_clause = " OR ".join(
f"LOWER(CAST({field} AS TEXT)) LIKE :search" for field in search_fields
)
clauses.append(f"({search_clause})")
values["search"] = f"%{clean_search}%"
return ("WHERE " + " AND ".join(clauses), values) if clauses else ("", values)
def _append_owner_where_sql(where_sql: str) -> str:
owner_clause = f"{OWNER_ID_FIELD} = :{OWNER_ID_FIELD}"
if where_sql:
return f"{where_sql} AND {owner_clause}"
return f"WHERE {owner_clause}"
def _order_sql(
table_schema: dict[str, Any],
sort_by: str | None,
descending: bool,
) -> str:
if not sort_by:
return ""
fields = _fields_by_name(table_schema)
if sort_by not in fields:
raise ValueError(f"WASM storage sort field is unknown: {sort_by}")
direction = "DESC" if descending else "ASC"
return f"ORDER BY {sort_by} {direction}"
def _row_from_db(
table_schema: dict[str, Any],
row: dict[str, Any],
) -> dict[str, Any]:
fields = _fields_by_name(table_schema)
return {
field_name: _value_from_db(fields[field_name], value)
for field_name, value in dict(row).items()
if field_name in fields
}
def _fields_by_name(table_schema: dict[str, Any]) -> dict[str, dict[str, Any]]:
fields = table_schema.get("fields")
if not isinstance(fields, list):
raise ValueError("WASM storage table schema fields must be a list.")
return {field["name"]: field for field in fields}
def _reject_reserved_owner_field(data: dict[str, Any], value_name: str) -> None:
if OWNER_ID_FIELD in data:
raise ValueError(f"WASM storage {value_name} includes a reserved owner field.")
def _value_to_db(field: dict[str, Any], value: Any) -> Any: # noqa: C901
if value is None:
if field.get("nullable", False):
return None
raise ValueError(f"WASM storage field '{field['name']}' cannot be null.")
if field.get("list") is True:
if not isinstance(value, list):
raise ValueError(f"WASM storage field '{field['name']}' must be a list.")
return json.dumps(value)
field_type = field.get("type")
if field_type == "string":
if not isinstance(value, str):
raise ValueError(f"WASM storage field '{field['name']}' must be a string.")
return value
if field_type == "integer":
if isinstance(value, bool) or not isinstance(value, int):
raise ValueError(
f"WASM storage field '{field['name']}' must be an integer."
)
return value
if field_type == "number":
if isinstance(value, bool) or not isinstance(value, int | float):
raise ValueError(f"WASM storage field '{field['name']}' must be a number.")
return value
if field_type == "boolean":
if not isinstance(value, bool):
raise ValueError(f"WASM storage field '{field['name']}' must be a boolean.")
return value
if field_type == "datetime":
if isinstance(value, int | float):
return datetime.fromtimestamp(value, tz=timezone.utc)
if isinstance(value, datetime):
return value
raise ValueError(
f"WASM storage field '{field['name']}' must be a Unix timestamp."
)
raise ValueError(f"Unsupported WASM storage field type: {field_type}")
def _value_from_db(field: dict[str, Any], value: Any) -> Any:
if value is None:
return None
if field.get("list") is True:
if isinstance(value, str):
return json.loads(value)
return value
field_type = field.get("type")
if field_type == "boolean":
return bool(value)
if field_type == "datetime":
if isinstance(value, datetime):
return int(value.replace(tzinfo=timezone.utc).timestamp())
if isinstance(value, int | float):
return int(value)
if isinstance(value, str):
try:
return int(datetime.fromisoformat(value).timestamp())
except ValueError:
return value
return value
def _default_sql(value: Any) -> str:
if value is None:
return "NULL"
if isinstance(value, bool):
return "true" if value else "false"
if isinstance(value, int | float):
return str(value)
if isinstance(value, str):
return _quote_sql_string(value)
if isinstance(value, list | dict):
return _quote_sql_string(json.dumps(value))
raise ValueError(f"Unsupported WASM storage default value: {value}")
def _field_from_add_field_operation(operation: dict[str, Any]) -> dict[str, Any]:
field = {
"name": operation.get("field"),
"type": operation.get("type"),
}
for key in ("default", "list", "nullable"):
if key in operation:
field[key] = operation[key]
return field
def _require_fields(operation: dict[str, Any]) -> list[dict[str, Any]]:
fields = operation.get("fields")
if not isinstance(fields, list) or not fields:
raise ValueError("WASM storage create_table operation requires fields.")
if not all(isinstance(field, dict) for field in fields):
raise ValueError("WASM storage fields must be objects.")
return fields
def _require_identifier(data: dict[str, Any], key: str) -> str:
value = data.get(key)
if not isinstance(value, str) or not _SQL_IDENTIFIER_RE.match(value):
raise ValueError(f"Invalid WASM storage SQL identifier for '{key}': {value}")
return value
def _table_ref(db: Connection, table: str) -> str:
if db.schema:
return f"{_schema_ref(db, table)}"
return table
def _table_ref_for_schema(ext_id: str, table: str) -> str:
_require_identifier({"schema": ext_id}, "schema")
_require_identifier({"table": table}, "table")
return f"{ext_id}.{table}"
def _schema_ref(db: Connection, name: str) -> str:
if not db.schema:
return name
if not _SQL_IDENTIFIER_RE.match(db.schema):
raise ValueError(f"Invalid WASM extension storage schema: {db.schema}")
return f"{db.schema}.{name}"
def _quote_sql_string(value: str) -> str:
return "'" + value.replace("'", "''") + "'"
def _load_json(path: Path) -> dict[str, Any]:
with open(path, encoding="utf-8") as json_file:
data = json.load(json_file)
if not isinstance(data, dict):
raise ValueError(f"WASM storage migration '{path}' must be a JSON object.")
return data
async def _update_wasm_migration_version(
db: Connection,
ext_id: str,
version: int,
) -> None:
if db.schema is None:
await update_migration_version(db, ext_id, version)
else:
async with core_db.connect() as conn:
await update_migration_version(conn, ext_id, version)
-230
View File
@@ -1,230 +0,0 @@
from __future__ import annotations
import asyncio
import json
import re
from collections.abc import Mapping
from functools import lru_cache
from typing import Any
from lnbits.core.crud.extensions import get_installed_extension
from lnbits.core.db import core_app_extra
from .api import ExtensionAPI, list_extension_api_methods
from .loader import WasmExtension
from .runtime import ExtensionAPIHost
async def invoke_wasm_extension_export(
ext_id: str,
export_name: str,
payload: Mapping[str, Any] | None = None,
*,
user: Any | None = None,
access_token: str | None = None,
context: str = "user",
owner_id: str | None = None,
) -> dict[str, Any]:
extension = _get_registered_extension(ext_id)
permissions = await _extension_permissions(extension)
api = ExtensionAPI(
extension.id,
permissions,
user_id=_user_id(user),
access_token=access_token,
context=context,
owner_id=owner_id,
)
event_loop = asyncio.get_running_loop()
return await asyncio.to_thread(
_invoke_wasm_extension_export_sync,
extension,
export_name,
payload or {},
api,
event_loop,
)
def warm_wasm_extension(extension: WasmExtension) -> None:
_wasm_component(extension)
def _invoke_wasm_extension_export_sync(
extension: WasmExtension,
export_name: str,
payload: Mapping[str, Any],
api: ExtensionAPI,
event_loop: asyncio.AbstractEventLoop,
) -> dict[str, Any]:
try:
from wasmtime import Store, WasiConfig, component
except ImportError as exc:
raise RuntimeError(
"WASM extension runtime is not installed. Install the 'wasmtime' "
"Python package to run WASM extensions."
) from exc
engine = _wasm_engine()
store = Store(engine)
store.set_wasi(WasiConfig())
linker = component.Linker(engine)
linker.add_wasip2()
_add_extension_host_imports(linker, ExtensionAPIHost(api), event_loop)
wasm_component = _wasm_component(extension)
instance = linker.instantiate(store, wasm_component)
function = instance.get_func(store, export_name)
if not function:
raise KeyError(
f"WASM extension '{extension.id}' has no export '{export_name}'."
)
result = function(store, json.dumps(payload))
function.post_return(store)
return _parse_wasm_export_result(extension, result)
@lru_cache(maxsize=1)
def _wasm_engine() -> Any:
try:
from wasmtime import Config, Engine
except ImportError as exc:
raise RuntimeError(
"WASM extension runtime is not installed. Install the 'wasmtime' "
"Python package to run WASM extensions."
) from exc
config = Config()
config.wasm_component_model = True
return Engine(config)
def _wasm_component(extension: WasmExtension) -> Any:
stat = extension.module_path.stat()
return _cached_wasm_component(
str(extension.module_path),
stat.st_mtime_ns,
stat.st_size,
)
@lru_cache(maxsize=32)
def _cached_wasm_component(
module_path: str,
mtime_ns: int,
size: int,
) -> Any:
from wasmtime import component
return component.Component.from_file(_wasm_engine(), module_path)
def _add_extension_host_imports(
linker: Any,
api_host: ExtensionAPIHost,
event_loop: asyncio.AbstractEventLoop,
) -> None:
with linker.root() as root:
methods_by_interface: dict[str, list[Any]] = {}
for method in list_extension_api_methods():
methods_by_interface.setdefault(method.host_interface, []).append(method)
for host_interface, methods in methods_by_interface.items():
with root.add_instance(f"lnbits:extension/{host_interface}") as host:
for method in methods:
host.add_func(
method.host_name.replace("_", "-"),
_make_host_import(api_host, method.method_id, event_loop),
)
def _make_host_import(
api_host: ExtensionAPIHost,
host_name: str,
event_loop: asyncio.AbstractEventLoop,
) -> Any:
def host_import(_store: Any, request: Any = None) -> Any:
payload = _component_payload_to_dict(request)
future = asyncio.run_coroutine_threadsafe(
api_host.invoke(host_name, payload), event_loop
)
response = future.result()
return _dict_to_component_record(response)
return host_import
def _component_payload_to_dict(value: Any) -> dict[str, Any]:
if value is None:
return {}
if hasattr(value, "__dict__"):
return dict(value.__dict__)
if isinstance(value, Mapping):
return dict(value)
raise TypeError("WASM host function payload must be a record.")
def _dict_to_component_record(value: Mapping[str, Any]) -> Any:
from wasmtime import component
record = component.Record()
for key, item in value.items():
setattr(record, _camel_to_kebab(key), _to_component_value(item))
return record
def _to_component_value(value: Any) -> Any:
if isinstance(value, Mapping):
return _dict_to_component_record(value)
if isinstance(value, list):
return [_to_component_value(item) for item in value]
return value
def _parse_wasm_export_result(extension: WasmExtension, value: Any) -> dict[str, Any]:
if isinstance(value, bytes):
value = value.decode()
if not isinstance(value, str):
return {"ok": True, "data": value}
max_response_bytes = (
(extension.config.get("wasm") or {})
.get("resource_limits", {})
.get("max_response_bytes")
)
if isinstance(max_response_bytes, int):
response_size = len(value.encode())
if response_size > max_response_bytes:
raise ValueError(
f"WASM extension response is too large: {response_size} bytes."
)
parsed = json.loads(value)
if isinstance(parsed, dict):
return parsed
return {"ok": True, "data": parsed}
def _get_registered_extension(ext_id: str) -> WasmExtension:
extension = core_app_extra.wasm_extension_registry.get(ext_id)
if extension:
return extension
raise RuntimeError(f"WASM extension '{ext_id}' is not registered.")
async def _extension_permissions(extension: WasmExtension) -> list[Any]:
installed_extension = await get_installed_extension(extension.id)
if not installed_extension:
return []
return installed_extension.permissions
def _user_id(user: Any | None) -> str | None:
return getattr(user, "id", None) if user else None
def _camel_to_kebab(value: str) -> str:
return re.sub(r"([a-z0-9])([A-Z])", r"\1-\2", value).replace("_", "-").lower()
-5
View File
@@ -13,8 +13,6 @@ from lnbits.core.crud import (
update_migration_version,
)
from lnbits.core.db import db as core_db
from lnbits.core.extensions.loader import is_wasm_extension_id
from lnbits.core.extensions.storage import migrate_wasm_extension_database
from lnbits.core.models import DbVersion
from lnbits.core.models.extensions import InstallableExtension
from lnbits.db import COCKROACH, POSTGRES, SQLITE, Connection
@@ -24,9 +22,6 @@ from lnbits.settings import settings
async def migrate_extension_database(
ext: InstallableExtension, current_version: DbVersion | None = None
):
if is_wasm_extension_id(ext.id):
await migrate_wasm_extension_database(ext, current_version)
return
try:
ext_migrations = importlib.import_module(f"{ext.module_name}.migrations")
-9
View File
@@ -815,12 +815,3 @@ async def m045_add_external_id_to_payments(db: Connection):
CREATE INDEX IF NOT EXISTS idx_payments_external_id
ON apipayments (external_id);
""")
async def m046_add_permissions_to_installed_extensions(db: Connection):
"""
Adds granted permissions to installed extensions.
"""
await db.execute(
"ALTER TABLE installed_extensions ADD COLUMN permissions TEXT DEFAULT '[]'"
)
+2 -69
View File
@@ -7,8 +7,7 @@ import os
import shutil
import zipfile
from asyncio.tasks import create_task
from collections.abc import Mapping
from pathlib import Path, PurePosixPath
from pathlib import Path
from typing import Any
import httpx
@@ -78,21 +77,6 @@ class GitHubRepo(BaseModel):
default_branch: str
class ExtensionPermission(BaseModel):
id: str
label: str | None = None
description: str | None = None
policy: dict[str, Any] | None = None
@staticmethod
def list_from_config(config_json: Mapping[str, Any]) -> list[ExtensionPermission]:
return [
ExtensionPermission.parse_obj(permission)
for permission in config_json.get("permissions") or []
if isinstance(permission, dict) and permission.get("id")
]
class ExtensionConfig(BaseModel):
name: str
short_description: str
@@ -100,8 +84,6 @@ class ExtensionConfig(BaseModel):
warning: str | None = ""
min_lnbits_version: str | None
max_lnbits_version: str | None
extension_type: str | None = None
permissions: list[ExtensionPermission] = []
def is_version_compatible(self) -> bool:
return is_lnbits_version_ok(self.min_lnbits_version, self.max_lnbits_version)
@@ -162,7 +144,6 @@ class UserExtension(BaseModel):
class Extension(BaseModel):
code: str
is_valid: bool
is_wasm: bool = False
name: str | None = None
short_description: str | None = None
tile: str | None = None
@@ -186,22 +167,13 @@ class Extension(BaseModel):
return Extension(
code=ext_info.id,
is_valid=True,
is_wasm=ext_info.is_wasm,
name=ext_info.name,
short_description=ext_info.short_description,
tile=(
wasm_extension_icon_url(ext_info.id)
if ext_info.is_wasm
else ext_info.icon
),
tile=ext_info.icon,
upgrade_hash=ext_info.hash if ext_info.ext_upgrade_dir.is_dir() else "",
)
def wasm_extension_icon_url(ext_id: str) -> str:
return f"/ext-assets/{ext_id}/assets/icon.png"
class ExtensionRelease(BaseModel):
name: str
version: str
@@ -376,7 +348,6 @@ class InstallableExtension(BaseModel):
icon: str | None = None
stars: int = 0
meta: ExtensionMeta | None = None
permissions: list[ExtensionPermission] = []
@property
def hash(self) -> str:
@@ -429,18 +400,6 @@ class InstallableExtension(BaseModel):
return False
return self.meta.pay_to_enable.required is True
@property
def is_wasm(self) -> bool:
config_path = Path(self.ext_dir, "config.json")
if not config_path.is_file():
return False
try:
with open(config_path, encoding="utf-8") as json_file:
config_json = json.load(json_file)
except Exception:
return False
return config_json.get("extension_type") == "wasm"
async def download_archive(self):
logger.info(f"Downloading extension {self.name} ({self.installed_version}).")
ext_zip_file = self.zip_path
@@ -473,30 +432,6 @@ class InstallableExtension(BaseModel):
os.remove(ext_zip_file)
raise AssertionError("File hash missmatch. Will not install.")
def load_archive_config(self) -> dict[str, Any]:
if not self.zip_path.is_file():
return {}
try:
with zipfile.ZipFile(self.zip_path, "r") as archive:
config_name = self._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 '{self.id}'.") from exc
return config if isinstance(config, dict) else {}
@staticmethod
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
def extract_archive(self):
logger.info(f"Extracting extension {self.name} ({self.installed_version}).")
Path(settings.lnbits_extensions_upgrade_path).mkdir(parents=True, exist_ok=True)
@@ -675,7 +610,6 @@ class InstallableExtension(BaseModel):
version=version,
short_description=config_json.get("short_description"),
icon=config_json.get("tile"),
permissions=ExtensionPermission.list_from_config(config_json),
meta=ExtensionMeta(
installed_release=ExtensionRelease(
name=ext_id,
@@ -866,7 +800,6 @@ class CreateExtension(BaseModel):
version: str
cost_sats: int | None = 0
payment_hash: str | None = None
permissions: list[ExtensionPermission] = []
class ExtensionDetailsRequest(BaseModel):
+1 -27
View File
@@ -1,7 +1,6 @@
from __future__ import annotations
from collections.abc import Awaitable, Callable
from typing import Any
from collections.abc import Callable
from pydantic import BaseModel
@@ -10,34 +9,9 @@ def _do_nothing(*_):
pass
async def _do_nothing_async(_: Any) -> None:
pass
class CoreAppExtra:
register_new_ext_routes: Callable = _do_nothing
register_new_wasm_ext_routes: Callable = _do_nothing
register_new_ratelimiter: Callable
dispatch_extension_invoice_paid: Callable[[Any], Awaitable[None]] = (
_do_nothing_async
)
def __init__(self) -> None:
self.wasm_extension_registry = WasmExtensionRegistry()
class WasmExtensionRegistry:
def __init__(self) -> None:
self._extensions: dict[str, Any] = {}
def register(self, extension: Any) -> None:
self._extensions[extension.id] = extension
def get(self, ext_id: str) -> Any | None:
return self._extensions.get(ext_id)
def list(self) -> list[Any]:
return list(self._extensions.values())
class ConversionData(BaseModel):
+4 -23
View File
@@ -16,23 +16,15 @@ from lnbits.core.crud.extensions import (
get_installed_extensions,
update_installed_extension,
)
from lnbits.core.extensions.permissions import validate_wasm_extension_permissions
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,
)
from ..models.extensions import Extension, ExtensionMeta, InstallableExtension
async def install_extension(
ext_info: InstallableExtension,
skip_download: bool | None = False,
granted_permissions: list[ExtensionPermission] | None = None,
ext_info: InstallableExtension, skip_download: bool | None = False
) -> Extension:
ext_info.meta = ext_info.meta or ExtensionMeta()
@@ -52,11 +44,6 @@ async def install_extension(
if not skip_download:
await ext_info.download_archive()
extension_config = ext_info.load_archive_config()
ext_info.permissions = validate_wasm_extension_permissions(
ext_info, granted_permissions, extension_config
)
ext_info.extract_archive()
db_version = await get_db_version(ext_info.id)
@@ -70,12 +57,11 @@ async def install_extension(
await update_installed_extension(ext_info)
extension = Extension.from_installable_ext(ext_info)
if extension.is_upgrade_extension and not extension.is_wasm:
if extension.is_upgrade_extension:
# call stop while the old routes are still active
await stop_extension_background_work(ext_info.id)
if not extension.is_wasm:
await start_extension_background_work(ext_info.id)
await start_extension_background_work(ext_info.id)
return extension
@@ -101,11 +87,6 @@ async def uninstall_extension(ext_id: str):
async def activate_extension(ext: Extension):
if ext.is_wasm:
core_app_extra.register_new_wasm_ext_routes(ext.code)
await update_installed_extension_state(ext_id=ext.code, active=True)
return
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)
-2
View File
@@ -10,7 +10,6 @@ from lnbits.core.crud.audit import delete_expired_audit_entries
from lnbits.core.crud.payments import get_payments_status_count
from lnbits.core.crud.users import get_accounts
from lnbits.core.crud.wallets import get_wallets_count
from lnbits.core.db import core_app_extra
from lnbits.core.models.audit import AuditEntry
from lnbits.core.models.extensions import InstallableExtension
from lnbits.core.models.notifications import NotificationType
@@ -101,7 +100,6 @@ async def wait_for_paid_invoices(invoice_paid_queue: asyncio.Queue) -> None:
wallet = await get_wallet(payment.wallet_id)
if wallet:
await send_payment_notification(wallet, payment)
await core_app_extra.dispatch_extension_invoice_paid(payment)
async def wait_for_audit_data() -> None:
+43 -69
View File
@@ -11,7 +11,6 @@ from loguru import logger
from lnbits.core.crud.extensions import get_user_extensions
from lnbits.core.crud.wallets import get_wallets_ids
from lnbits.core.db import db
from lnbits.core.extensions.permissions import validate_extension_permissions
from lnbits.core.models import (
SimpleStatus,
)
@@ -30,7 +29,6 @@ from lnbits.core.models.extensions import (
ReleasePaymentInfo,
UserExtension,
UserExtensionInfo,
wasm_extension_icon_url,
)
from lnbits.core.models.users import Account, AccountId
from lnbits.core.services import check_transaction_status, create_invoice
@@ -91,9 +89,7 @@ async def api_install_extension(data: CreateExtension):
)
try:
extension = await install_extension(
ext_info, granted_permissions=data.permissions
)
extension = await install_extension(ext_info)
except Exception as exc:
logger.warning(exc)
@@ -463,19 +459,11 @@ async def get_extension_release(org: str, repo: str, tag_name: str):
if not config:
return {}
permissions = validate_extension_permissions(config.name, config.permissions)
return {
"min_lnbits_version": config.min_lnbits_version,
"is_version_compatible": config.is_version_compatible(),
"warning": config.warning,
"extension_type": config.extension_type,
"permissions": [dict(permission) for permission in permissions],
}
except ValueError as exc:
raise HTTPException(
status_code=HTTPStatus.BAD_REQUEST, detail=str(exc)
) from exc
except Exception as exc:
raise HTTPException(
status_code=HTTPStatus.INTERNAL_SERVER_ERROR, detail=str(exc)
@@ -547,10 +535,9 @@ async def extensions(account_id: AccountId = Depends(check_account_id_exists)):
)
installable_exts_ids = [e.id for e in installable_exts]
installable_exts += [e for e in installed_exts if e.id not in installable_exts_ids]
installed_exts_by_id = {e.id: e for e in installed_exts}
for e in installable_exts:
installed_ext = installed_exts_by_id.get(e.id)
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 account_id.is_admin_id:
@@ -571,60 +558,47 @@ async def extensions(account_id: AccountId = Depends(check_account_id_exists)):
e.short_description = installed_ext.short_description
e.icon = installed_ext.icon
extension_data = []
for ext in installable_exts:
installed_ext = installed_exts_by_id.get(ext.id)
is_wasm = installed_ext.is_wasm if installed_ext else ext.is_wasm
icon = wasm_extension_icon_url(ext.id) if is_wasm else ext.icon
permissions = (
validate_extension_permissions(
installed_ext.id, installed_ext.permissions, strict=False
)
if installed_ext
else []
)
extension_data.append(
{
"id": ext.id,
"name": ext.name,
"icon": icon,
"shortDescription": ext.short_description,
"stars": ext.stars,
"isFeatured": ext.meta.featured if ext.meta else False,
"categories": ext.meta.categories if ext.meta else [],
"dependencies": ext.meta.dependencies if ext.meta else "",
"isInstalled": ext.id in installed_exts_ids,
"hasDatabaseTables": next(
(True for version in db_versions if version.db == ext.id), False
),
"isAvailable": ext.id in all_ext_ids,
"isAdminOnly": ext.id in settings.lnbits_admin_extensions,
"isActive": ext.id not in inactive_extensions,
"latestRelease": (
dict(ext.meta.latest_release)
if ext.meta and ext.meta.latest_release
else None
),
"hasPaidRelease": ext.meta.has_paid_release if ext.meta else False,
"hasFreeRelease": ext.meta.has_free_release if ext.meta else False,
"paidFeatures": ext.meta.paid_features if ext.meta else False,
"installedRelease": (
dict(ext.meta.installed_release)
if ext.meta and ext.meta.installed_release
else None
),
"payToEnable": (
dict(ext.meta.pay_to_enable)
if ext.meta and ext.meta.pay_to_enable
else {}
),
"isPaymentRequired": ext.requires_payment,
"isWasm": is_wasm,
"permissions": [dict(permission) for permission in permissions],
"inProgress": False,
"selectedForUpdate": False,
}
)
extension_data = [
{
"id": ext.id,
"name": ext.name,
"icon": ext.icon,
"shortDescription": ext.short_description,
"stars": ext.stars,
"isFeatured": ext.meta.featured if ext.meta else False,
"categories": ext.meta.categories if ext.meta else [],
"dependencies": ext.meta.dependencies if ext.meta else "",
"isInstalled": ext.id in installed_exts_ids,
"hasDatabaseTables": next(
(True for version in db_versions if version.db == ext.id), False
),
"isAvailable": ext.id in all_ext_ids,
"isAdminOnly": ext.id in settings.lnbits_admin_extensions,
"isActive": ext.id not in inactive_extensions,
"latestRelease": (
dict(ext.meta.latest_release)
if ext.meta and ext.meta.latest_release
else None
),
"hasPaidRelease": ext.meta.has_paid_release if ext.meta else False,
"hasFreeRelease": ext.meta.has_free_release if ext.meta else False,
"paidFeatures": ext.meta.paid_features if ext.meta else False,
"installedRelease": (
dict(ext.meta.installed_release)
if ext.meta and ext.meta.installed_release
else None
),
"payToEnable": (
dict(ext.meta.pay_to_enable)
if ext.meta and ext.meta.pay_to_enable
else {}
),
"isPaymentRequired": ext.requires_payment,
"inProgress": False,
"selectedForUpdate": False,
}
for ext in installable_exts
]
return extension_data
+1 -10
View File
@@ -448,7 +448,7 @@ async def _check_user_access(r: Request, user_id: str, conn: Connection | None =
async def _check_user_extension_access(
user_id: str, path: str, conn: Connection | None = None
):
ext_id = _extension_id_from_request_path(path)
ext_id = path_segments(path)[0]
status = await check_user_extension_access(user_id, ext_id, conn=conn)
if not status.success:
raise HTTPException(
@@ -457,15 +457,6 @@ async def _check_user_extension_access(
)
def _extension_id_from_request_path(path: str) -> str:
segments = path_segments(path)
if len(segments) >= 2 and segments[0] == "ext":
return segments[1]
if len(segments) >= 4 and segments[:3] == ["api", "v1", "ext"]:
return segments[3]
return segments[0]
async def _get_account_from_token(
access_token: str, path: str, method: str, conn: Connection | None = None
) -> Account | None:
+6
View File
@@ -27,6 +27,8 @@ from lnbits.settings import set_cli_settings, settings
@click.option(
"--reload", is_flag=True, default=False, help="Enable auto-reload for development"
)
@click.option("--ws-max-queue", default=128, help="Websocket max queue size")
@click.option("--ws-ping-timeout", default=900.0, help="Websocket ping timeout")
def main(
port: int,
host: str,
@@ -34,6 +36,8 @@ def main(
ssl_keyfile: str,
ssl_certfile: str,
reload: bool,
ws_max_queue: int,
ws_ping_timeout: float
):
"""Launched with `uv run lnbits` at root level"""
@@ -58,6 +62,8 @@ def main(
ssl_keyfile=ssl_keyfile,
ssl_certfile=ssl_certfile,
reload=reload or False,
ws_ping_timeout=ws_ping_timeout,
ws_max_queue=ws_max_queue,
)
server = uvicorn.Server(config=config)
File diff suppressed because one or more lines are too long
+1 -1
View File
File diff suppressed because one or more lines are too long
-33
View File
@@ -519,39 +519,6 @@ window.localisation.en = {
admin_settings: 'Admin Settings',
extension_cost: 'This release requires a payment of minimum {cost} sats.',
extension_paid_sats: 'You have already paid {paid_sats} sats.',
extension_permissions_title: 'Grant extension permissions',
extension_permissions_request: 'This extension requests these permissions:',
extension_permissions_grant_install: 'Grant and install',
extension_permission_ext_storage_read: 'Read extension storage',
extension_permission_ext_storage_read_public: 'Read public extension storage',
extension_permission_ext_storage_write: 'Write extension storage',
extension_permission_extension_api_request: 'Use other extensions',
extension_permission_extension_api_request_desc:
'Call approved installed extensions using your account permissions.',
extension_permission_extension_api_request_extensions: 'Allowed extensions',
extension_permission_http_request: 'Connect to external websites',
extension_permission_http_request_desc:
'Make HTTP requests to approved external hosts.',
extension_permission_http_request_hosts: 'Allowed hosts',
extension_permission_utils_basic: 'Use basic LNbits utilities',
extension_permission_utils_basic_desc:
'Use public currency conversion, server health, and Lightning invoice helper functions.',
extension_permission_ui_camera_scan_qr: 'Scan QR codes',
extension_permission_ui_camera_scan_qr_desc:
'Use the LNbits scanner to read QR codes when you choose to scan.',
extension_permission_payments_watch: 'Watch payments',
extension_permission_wallet_create_invoice: 'Create invoices',
extension_permission_wallet_create_invoice_public:
'Create Lightning invoices from public pages',
extension_permission_wallet_create_invoice_public_desc:
'Create incoming Lightning invoices from public pages.',
extension_permission_wallet_balance_read: 'View wallet balances',
extension_permission_wallet_balance_read_desc:
'Read balances of wallets available to your account.',
extension_permission_wallet_list: 'List wallets',
extension_permission_wallet_pay_invoice: 'Pay invoices',
extension_permission_wallet_pay_invoice_desc:
'Send Lightning payments from wallets available to your account.',
create_extension: 'Create Extension',
release_details_error: 'Cannot get the release details.',
pay_from_wallet: 'Pay from Wallet',
@@ -35,18 +35,6 @@ window.app.component('lnbits-manage-extension-list', {
.toLocaleLowerCase()
.includes(this.searchTerm.toLocaleLowerCase())
})
},
extensionUrl(extension) {
if (extension.is_wasm) {
return `/ext/${extension.code}`
}
return `/${extension.code}/`
},
extensionActive(extension) {
const extensionPath = extension.is_wasm
? `/ext/${extension.code}`
: `/${extension.code}`
return this.$route.path.startsWith(extensionPath)
}
},
async created() {
-10
View File
@@ -139,16 +139,6 @@ const routes = [
name: 'PageError',
component: PageError
},
{
path: '/ext/:extId',
name: 'WasmExtensionRoot',
component: window.WasmExtensionComponent
},
{
path: '/ext/:extId/:pathMatch(.*)*',
name: 'WasmExtension',
component: window.WasmExtensionComponent
},
{
path: '/:pathMatch(.*)*',
name: 'DynamicComponent',
+1 -133
View File
@@ -24,11 +24,6 @@ window.PageExtensions = {
selectedExtensionDetails: null,
selectedExtensionRepos: null,
selectedRelease: null,
permissionGrant: {
show: false,
permissions: [],
resolve: null
},
uninstallAndDropDb: false,
maxStars: 5,
paylinkWebsocket: null,
@@ -137,12 +132,6 @@ window.PageExtensions = {
// the install logic has been triggered one way or another
this.unsubscribeFromPaylinkWs()
const grantedPermissions =
await this.resolveExtensionPermissionGrant(release)
if (grantedPermissions === null) {
return
}
this.selectedExtension.inProgress = true
this.showManageExtensionDialog = false
release.payment_hash =
@@ -154,8 +143,7 @@ window.PageExtensions = {
archive: release.archive,
source_repo: release.source_repo,
payment_hash: release.payment_hash,
version: release.version,
permissions: grantedPermissions
version: release.version
})
.then(response => {
this.selectedExtension.inProgress = false
@@ -164,7 +152,6 @@ window.PageExtensions = {
)
extension.isAvailable = true
extension.isInstalled = true
extension.icon = response.data.icon || extension.icon
extension.installedRelease = release
this.toggleExtension(extension)
extension.inProgress = false
@@ -419,10 +406,6 @@ window.PageExtensions = {
},
async payAndInstall(release) {
try {
if ((await this.resolveExtensionPermissionGrant(release)) === null) {
return
}
this.selectedExtension.inProgress = true
this.showManageExtensionDialog = false
const paymentInfo = await this.requestPaymentForInstall(
@@ -468,10 +451,6 @@ window.PageExtensions = {
}
},
async showInstallQRCode(release) {
if ((await this.resolveExtensionPermissionGrant(release)) === null) {
return
}
this.selectedRelease = release
try {
@@ -634,9 +613,6 @@ window.PageExtensions = {
return ''
},
extensionOpenUrl(extension) {
return extension.isWasm ? `/ext/${extension.id}` : `/${extension.id}`
},
async getGitHubReleaseDetails(release) {
if (!release.is_github_release || release.loaded) {
return
@@ -652,8 +628,6 @@ window.PageExtensions = {
release.is_version_compatible = data.is_version_compatible
release.min_lnbits_version = data.min_lnbits_version
release.warning = data.warning
release.extension_type = data.extension_type
release.permissions = data.permissions || []
} catch (error) {
console.warn(error)
release.error = error
@@ -662,105 +636,6 @@ window.PageExtensions = {
release.inProgress = false
}
},
async resolveExtensionPermissionGrant(release) {
const permissions = this.extensionPermissionsForRelease(release)
if (
!this.releaseRequiresPermissionGrant(release) ||
!permissions.length
) {
return []
}
if (release.grantedPermissions) {
return release.grantedPermissions
}
const grantedPermissions =
await this.confirmExtensionPermissions(permissions)
if (!grantedPermissions) {
return null
}
release.grantedPermissions = grantedPermissions
return grantedPermissions
},
extensionPermissionsForRelease(release) {
return release.permissions || this.selectedExtension?.permissions || []
},
releaseRequiresPermissionGrant(release) {
return (
release.extension_type === 'wasm' ||
this.selectedExtension?.isWasm === true
)
},
confirmExtensionPermissions(permissions) {
return new Promise(resolve => {
this.selectedRelease = null
this.permissionGrant = {
show: true,
permissions,
resolve
}
this.showManageExtensionDialog = true
})
},
grantExtensionPermissions() {
this.resolveExtensionPermissionDialog(this.permissionGrant.permissions)
},
cancelExtensionPermissions() {
this.resolveExtensionPermissionDialog(null)
},
onManageExtensionDialogHide() {
if (this.permissionGrant.show) {
this.resolveExtensionPermissionDialog(null)
}
},
resolveExtensionPermissionDialog(grantedPermissions) {
const resolve = this.permissionGrant.resolve
this.permissionGrant = {
show: false,
permissions: [],
resolve: null
}
this.showManageExtensionDialog = false
if (resolve) {
resolve(grantedPermissions)
}
},
permissionI18nKey(permission) {
return `extension_permission_${permission.id.replace(/[^A-Za-z0-9]/g, '_')}`
},
permissionLabel(permission) {
const key = this.permissionI18nKey(permission)
const label = this.$t(key)
return label === key ? permission.id : label
},
permissionDescription(permission) {
const key = `${this.permissionI18nKey(permission)}_desc`
const description = this.$t(key)
return description === key ? permission.description : description
},
permissionPolicyDetails(permission) {
if (permission.id === 'http.request') {
const hosts = permission.policy?.hosts
if (!Array.isArray(hosts) || !hosts.length) return ''
return `${this.$t('extension_permission_http_request_hosts')}: ${hosts.join(', ')}`
}
if (permission.id === 'extension.api.request') {
const extensions = permission.policy?.extensions
if (!Array.isArray(extensions) || !extensions.length) return ''
const targets = extensions
.map(extension => {
if (typeof extension === 'string') return `${extension} (read)`
if (!extension?.id) return null
const access = Array.isArray(extension.access)
? extension.access.join(', ')
: 'read'
return `${extension.id} (${access})`
})
.filter(Boolean)
if (!targets.length) return ''
return `${this.$t('extension_permission_extension_api_request_extensions')}: ${targets.join(', ')}`
}
return ''
},
async selectAllUpdatableExtensionss() {
this.updatableExtensions.forEach(e => (e.selectedForUpdate = true))
},
@@ -771,13 +646,6 @@ window.PageExtensions = {
if (!ext.selectedForUpdate) {
continue
}
if (ext.isWasm) {
Quasar.Notify.create({
type: 'warning',
message: `Skipping ${ext.id}; this extension update requires permission approval.`
})
continue
}
ext.inProgress = true
await LNbits.api.request('POST', `/api/v1/extension`, null, {
ext_id: ext.id,
@@ -1,566 +0,0 @@
window.WasmExtensionComponent = {
template: `
<div class="wasm-extension-page relative-position">
<q-inner-loading :showing="loading && !frameUrl">
<q-spinner-dots size="40px"></q-spinner-dots>
</q-inner-loading>
<q-banner v-if="error" class="q-ma-md bg-negative text-white">
{{ error }}
</q-banner>
<iframe
v-else-if="frameUrl"
ref="frame"
:key="frameUrl"
class="wasm-extension-frame"
:src="frameUrl"
:title="extensionName || 'Extension'"
sandbox="allow-scripts"
allow="clipboard-write"
referrerpolicy="no-referrer"
></iframe>
<q-dialog v-model="cameraPrompt.show" persistent>
<q-card style="width: min(520px, calc(100vw - 32px)); max-width: 520px">
<q-card-section>
<div class="text-h6">Camera access</div>
</q-card-section>
<q-card-section class="q-pt-none">
{{ cameraPrompt.extensionName }} wants to access the camera to scan a QR code.
</q-card-section>
<q-card-actions align="right">
<q-btn
flat
color="negative"
label="Deny"
@click="resolveCameraPrompt('deny')"
></q-btn>
<q-btn
flat
color="primary"
label="Allow"
@click="resolveCameraPrompt('allow')"
></q-btn>
<q-btn
unelevated
color="primary"
label="Allow and Remember"
@click="resolveCameraPrompt('allow_remember')"
></q-btn>
</q-card-actions>
</q-card>
</q-dialog>
</div>
`,
data() {
return {
allowedPaymentHashes: new Set(),
bridge: {
apiRoutes: [],
extensionId: '',
permissions: [],
public: false,
query: {},
routeParams: {}
},
bridgePort: null,
cameraPrompt: {
extensionName: '',
reject: null,
resolve: null,
show: false
},
error: '',
extensionName: '',
frameUrl: '',
handleWindowMessage: null,
loading: false,
loadId: 0,
paymentSubscriptions: new Map()
}
},
created() {
this.handleWindowMessage = event => this.onWindowMessage(event)
window.addEventListener('message', this.handleWindowMessage)
},
unmounted() {
window.removeEventListener('message', this.handleWindowMessage)
this.rejectCameraPrompt('Camera scan cancelled.')
this.closeBridgePort()
},
watch: {
'$route.fullPath': {
immediate: true,
handler() {
this.loadFrameConfig()
}
}
},
methods: {
emptyBridge() {
return {
apiRoutes: [],
extensionId: '',
permissions: [],
public: false,
query: {},
routeParams: {}
}
},
plainBridgeContext() {
return {
extensionId: String(this.bridge.extensionId || ''),
public: Boolean(this.bridge.public),
routeParams: this.plainValue(this.bridge.routeParams || {}),
query: this.plainValue(this.bridge.query || {})
}
},
hasBridgePermission(permission) {
return (this.bridge.permissions || []).includes(permission)
},
cameraPromptStorageKey() {
return `lnbits.ext.permissions.${this.bridge.extensionId}.ui.camera.scan_qr`
},
emptyCameraPrompt() {
return {
extensionName: '',
reject: null,
resolve: null,
show: false
}
},
plainValue(value) {
try {
return JSON.parse(JSON.stringify(value))
} catch (_error) {
return {}
}
},
async loadFrameConfig() {
const extId = String(this.$route.params.extId || '')
const loadId = ++this.loadId
this.loading = true
this.error = ''
this.frameUrl = ''
this.bridge = this.emptyBridge()
this.allowedPaymentHashes.clear()
this.rejectCameraPrompt('Camera scan cancelled.')
this.closeBridgePort()
try {
const response = await fetch(
`/api/v1/ext/${encodeURIComponent(extId)}/_ui/frame`,
{
method: 'POST',
headers: {'content-type': 'application/json'},
credentials: 'same-origin',
body: JSON.stringify({
path: this.$route.path,
query: this.$route.query || {}
})
}
)
const text = await response.text()
let data = {}
if (text) {
try {
data = JSON.parse(text)
} catch (_error) {
data = {detail: text}
}
}
if (!response.ok) {
throw new Error(data?.detail || 'Failed to load extension page.')
}
if (loadId !== this.loadId) return
this.bridge = data.bridge || this.emptyBridge()
this.extensionName = data.extension?.name || extId
this.frameUrl = data.frameUrl
} catch (error) {
if (loadId !== this.loadId) return
console.error('[lnbits wasm extension] Failed to load frame.', error)
this.error = error instanceof Error ? error.message : String(error)
} finally {
if (loadId === this.loadId) {
this.loading = false
}
}
},
extensionFrameWindow() {
return this.$refs.frame?.contentWindow
},
sendResponse(reply, id, payload) {
reply({
type: 'lnbits-extension:response',
id,
...payload
})
},
allowedApiRoute(method, path) {
let url
try {
url = new URL(path, window.location.origin)
} catch (_error) {
return false
}
if (url.origin !== window.location.origin) return false
method = String(method || 'GET').toUpperCase()
return (this.bridge.apiRoutes || []).some(route => {
return (
route.method === method &&
new RegExp(route.pattern).test(url.pathname)
)
})
},
async callApi(message) {
const method = String(message.method || 'GET').toUpperCase()
const path = String(message.path || '')
if (!this.allowedApiRoute(method, path)) {
throw new Error('Extension API route is not allowed.')
}
const options = {
method,
headers: {},
credentials: 'same-origin'
}
if (message.body !== undefined && message.body !== null) {
options.headers['content-type'] = 'application/json'
options.body = JSON.stringify(message.body)
}
const response = await fetch(path, options)
const text = await response.text()
let data = text
if (text) {
try {
data = JSON.parse(text)
} catch (_error) {
data = text
}
}
if (!response.ok) {
throw new Error(
typeof data === 'object' && data.detail ? data.detail : text
)
}
this.rememberPaymentHashes(data)
return data
},
notify(message) {
const level = ['positive', 'negative', 'warning', 'info'].includes(
message.level
)
? message.level
: 'info'
if (window.Quasar?.Notify) {
window.Quasar.Notify.create({
color: level,
message: String(message.message || '')
})
}
},
async scanQrCode() {
if (!this.hasBridgePermission('ui.camera.scan_qr')) {
throw new Error('Extension is missing scanner permission.')
}
if (!this.g) {
throw new Error('LNbits scanner is not available.')
}
if (this.g.scanner) {
throw new Error('A scanner is already active.')
}
await this.requireCameraScanApproval()
if (this.g.scanner) {
throw new Error('A scanner is already active.')
}
return new Promise((resolve, reject) => {
let completed = false
const cleanup = () => {
window.clearTimeout(timeout)
window.clearInterval(cancelPoll)
if (this.g.scanner === onScan) {
this.g.scanner = null
}
}
const complete = callback => value => {
if (completed) return
completed = true
cleanup()
callback(value)
}
const onScan = value => {
complete(resolve)({value: String(value || '')})
}
const timeout = window.setTimeout(() => {
complete(reject)(new Error('QR scan timed out.'))
}, 120000)
const cancelPoll = window.setInterval(() => {
if (!completed && this.g.scanner !== onScan) {
complete(reject)(new Error('QR scan cancelled.'))
}
}, 250)
this.g.scanner = onScan
})
},
requireCameraScanApproval() {
if (this.isCameraScanRemembered()) return Promise.resolve()
if (this.cameraPrompt.show) {
return Promise.reject(
new Error('Camera access prompt is already open.')
)
}
return new Promise((resolve, reject) => {
this.cameraPrompt = {
extensionName:
this.extensionName || this.bridge.extensionId || 'This extension',
reject,
resolve,
show: true
}
})
},
isCameraScanRemembered() {
try {
return (
this.$q.localStorage.getItem(this.cameraPromptStorageKey()) ===
'allow'
)
} catch (_error) {
return false
}
},
rememberCameraScanApproval() {
try {
this.$q.localStorage.set(this.cameraPromptStorageKey(), 'allow')
} catch (_error) {}
},
resolveCameraPrompt(decision) {
const resolve = this.cameraPrompt.resolve
const reject = this.cameraPrompt.reject
this.cameraPrompt = this.emptyCameraPrompt()
if (decision === 'allow_remember') {
this.rememberCameraScanApproval()
resolve?.()
return
}
if (decision === 'allow') {
resolve?.()
return
}
reject?.(new Error('Camera scan denied by user.'))
},
rejectCameraPrompt(message) {
const reject = this.cameraPrompt.reject
this.cameraPrompt = this.emptyCameraPrompt()
reject?.(new Error(message))
},
rememberPaymentHashes(value) {
if (!value || typeof value !== 'object') return
if (Array.isArray(value)) {
value.forEach(item => this.rememberPaymentHashes(item))
return
}
for (const [key, item] of Object.entries(value)) {
if (
['paymentHash', 'payment_hash'].includes(key) &&
this.isPaymentHash(item)
) {
this.allowedPaymentHashes.add(item)
}
this.rememberPaymentHashes(item)
}
},
isPaymentHash(value) {
return typeof value === 'string' && /^[a-f0-9]{64}$/i.test(value)
},
websocketUrl(path) {
const url = new URL(window.location.href)
url.protocol = url.protocol === 'https:' ? 'wss:' : 'ws:'
url.pathname = path
url.search = ''
url.hash = ''
return url.toString()
},
sendBridgeEvent(message) {
if (!this.bridgePort) return
this.bridgePort.postMessage({
type: 'lnbits-extension:event',
...message
})
},
closePaymentSubscription(subscriptionId) {
const subscription = this.paymentSubscriptions.get(subscriptionId)
if (!subscription) return
this.paymentSubscriptions.delete(subscriptionId)
try {
subscription.socket.close()
} catch (_error) {}
},
closePaymentSubscriptions() {
for (const subscriptionId of Array.from(
this.paymentSubscriptions.keys()
)) {
this.closePaymentSubscription(subscriptionId)
}
},
closeBridgePort() {
this.closePaymentSubscriptions()
this.bridgePort?.close()
this.bridgePort = null
},
subscribePayment(message) {
const subscriptionId = String(message.subscriptionId || '')
const paymentHash = String(message.paymentHash || '')
if (!subscriptionId || !this.isPaymentHash(paymentHash)) {
throw new Error('Invalid payment subscription.')
}
if (!this.allowedPaymentHashes.has(paymentHash)) {
throw new Error('Payment subscription is not allowed.')
}
this.closePaymentSubscription(subscriptionId)
const socket = new WebSocket(
this.websocketUrl(`/api/v1/ws/${encodeURIComponent(paymentHash)}`)
)
this.paymentSubscriptions.set(subscriptionId, {paymentHash, socket})
socket.addEventListener('message', event => {
let data = event.data
try {
data = JSON.parse(event.data)
} catch (_error) {}
this.sendBridgeEvent({
event: 'payment.update',
subscriptionId,
paymentHash,
data
})
if (
data &&
typeof data === 'object' &&
(data.pending === false ||
['success', 'settled', 'paid'].includes(String(data.status || '')))
) {
this.sendBridgeEvent({
event: 'payment.settled',
subscriptionId,
paymentHash,
data
})
this.closePaymentSubscription(subscriptionId)
}
})
socket.addEventListener('error', () => {
this.sendBridgeEvent({
event: 'payment.error',
subscriptionId,
paymentHash
})
this.closePaymentSubscription(subscriptionId)
})
socket.addEventListener('close', () => {
this.paymentSubscriptions.delete(subscriptionId)
})
},
async handleBridgeRequest(message, reply) {
if (!message || message.type !== 'lnbits-extension:request') return
try {
if (message.action === 'context') {
this.sendResponse(reply, message.id, {
ok: true,
data: this.plainBridgeContext()
})
return
}
if (message.action === 'api') {
this.sendResponse(reply, message.id, {
ok: true,
data: await this.callApi(message)
})
return
}
if (message.action === 'ui.notify') {
this.notify(message)
this.sendResponse(reply, message.id, {
ok: true,
data: {ok: true}
})
return
}
if (message.action === 'ui.scan_qr') {
this.sendResponse(reply, message.id, {
ok: true,
data: await this.scanQrCode()
})
return
}
if (message.action === 'payment.subscribe') {
this.subscribePayment(message)
this.sendResponse(reply, message.id, {
ok: true,
data: {ok: true}
})
return
}
if (message.action === 'payment.unsubscribe') {
this.closePaymentSubscription(String(message.subscriptionId || ''))
this.sendResponse(reply, message.id, {
ok: true,
data: {ok: true}
})
return
}
throw new Error('Unknown extension bridge action.')
} catch (error) {
this.sendResponse(reply, message.id, {
ok: false,
error: error instanceof Error ? error.message : String(error)
})
}
},
onWindowMessage(event) {
if (event.source !== this.extensionFrameWindow()) return
const message = event.data
if (!message || message.type !== 'lnbits-extension:connect') return
const port = event.ports?.[0]
if (!port) return
this.closeBridgePort()
this.bridgePort = port
this.bridgePort.addEventListener('message', portEvent => {
this.handleBridgeRequest(portEvent.data, response => {
port.postMessage(response)
})
})
this.bridgePort.start()
this.bridgePort.postMessage({
type: 'lnbits-extension:connected',
id: message.id
})
}
}
}
-1
View File
@@ -96,7 +96,6 @@
"js/components/extension-settings.js",
"js/components/data-fields.js",
"js/components.js",
"js/wasm-extension-component.js",
"js/init-app.js"
],
"css": ["vendor/quasar.css", "css/base.css"]
+3 -17
View File
@@ -16,18 +16,6 @@
src: url("{{ static_url_for('static', 'fonts/material-icons-v50.woff2') }}")
format('woff2');
}
.wasm-extension-page,
.wasm-extension-frame {
height: calc(100vh - 56px);
min-height: calc(100vh - 56px);
width: 100%;
}
.wasm-extension-frame {
border: 0;
display: block;
}
</style>
<title>{% block title %}{{ SITE_TITLE }}{% endblock %}</title>
<meta charset="utf-8" />
@@ -56,14 +44,12 @@
<lnbits-drawer v-if="g.user && !g.isPublicPage"></lnbits-drawer>
{% block page_container %}
<q-page-container>
<q-page
:class="$route.path.startsWith('/ext/') ? 'q-pa-none' : ['q-px-md', 'q-py-lg', {'q-px-lg': $q.screen.gt.xs}]"
>
<q-page class="q-px-md q-py-lg" :class="{'q-px-lg': $q.screen.gt.xs}">
<lnbits-wallet-new
v-if="g.user && !g.isPublicPage && !$route.path.startsWith('/ext/')"
v-if="g.user && !g.isPublicPage"
></lnbits-wallet-new>
<lnbits-header-wallets
v-if="g.user && !g.isPublicPage && !$route.path.startsWith('/ext/')"
v-if="g.user && !g.isPublicPage"
></lnbits-header-wallets>
<!-- block page content from static extensions -->
<div
@@ -18,18 +18,13 @@
<q-item
v-for="extension in userExtensions"
clickable
:active="extensionActive(extension)"
:active="$route.path.startsWith('/' + extension.code)"
tag="a"
:to="extensionUrl(extension)"
:to="'/' + extension.code + '/'"
>
<q-item-section side>
<q-avatar size="md">
<img
v-if="extension.tile"
:src="extension.tile"
style="width: 20px; height: 20px; object-fit: contain"
/>
<q-icon v-else name="extension" size="20px"></q-icon>
<q-img :src="extension.tile" style="max-width: 20px"></q-img>
</q-avatar>
</q-item-section>
<q-item-section>
@@ -37,7 +32,10 @@
><span v-text="extension.name"></span>
</q-item-label>
</q-item-section>
<q-item-section side v-show="extensionActive(extension)">
<q-item-section
side
v-show="$route.path.startsWith('/' + extension.code)"
>
<q-icon name="chevron_right" color="grey-5" size="md"></q-icon>
</q-item-section>
</q-item>
+3 -55
View File
@@ -314,7 +314,7 @@
flat
color="primary"
type="a"
:href="extensionOpenUrl(extension)"
:href="extension.id + '/'"
:label="$t('open')"
></q-btn>
<q-btn
@@ -455,60 +455,8 @@
</q-card>
</q-dialog>
<q-dialog
v-model="showManageExtensionDialog"
position="top"
@hide="onManageExtensionDialogHide"
>
<q-card v-if="permissionGrant.show" class="q-pa-lg lnbits__dialog-card">
<q-card-section>
<div class="text-h6" v-text="$t('extension_permissions_title')"></div>
<div
class="text-body2 q-mt-sm"
v-text="$t('extension_permissions_request')"
></div>
</q-card-section>
<q-list bordered separator class="q-mt-md">
<q-item
v-for="permission of permissionGrant.permissions"
:key="permission.id"
>
<q-item-section>
<q-item-label>
<li><strong v-text="permissionLabel(permission)"></strong></li>
</q-item-label>
<q-item-label
v-if="permissionDescription(permission)"
caption
v-text="permissionDescription(permission)"
></q-item-label>
<q-item-label
v-if="permissionPolicyDetails(permission)"
caption
v-text="permissionPolicyDetails(permission)"
></q-item-label>
</q-item-section>
</q-item>
</q-list>
<div class="row q-mt-lg">
<q-btn
flat
color="grey"
v-text="$t('cancel')"
@click="cancelExtensionPermissions"
></q-btn>
<q-btn
flat
color="primary"
class="q-ml-auto"
v-text="$t('extension_permissions_grant_install')"
@click="grantExtensionPermissions"
></q-btn>
</div>
</q-card>
<q-card v-else-if="selectedRelease" class="q-pa-lg lnbits__dialog-card">
<q-dialog v-model="showManageExtensionDialog" position="top">
<q-card v-if="selectedRelease" class="q-pa-lg lnbits__dialog-card">
<q-card-section>
<div v-if="selectedRelease.paymentRequest">
<lnbits-qrcode
-1
View File
@@ -1 +0,0 @@
{% extends "base.html" %}
+921 -138
View File
File diff suppressed because it is too large Load Diff
-1
View File
@@ -149,7 +149,6 @@
"js/components/extension-settings.js",
"js/components/data-fields.js",
"js/components.js",
"js/wasm-extension-component.js",
"js/init-app.js"
],
"css": [
Generated
+4 -29
View File
@@ -1,4 +1,4 @@
# This file is automatically @generated by Poetry 2.4.1 and should not be changed by hand.
# This file is automatically @generated by Poetry 2.2.1 and should not be changed by hand.
[[package]]
name = "aiohappyeyeballs"
@@ -2147,7 +2147,7 @@ files = [
[package.dependencies]
attrs = ">=22.2.0"
jsonschema-specifications = ">=2023.3.6"
jsonschema-specifications = ">=2023.03.6"
referencing = ">=0.28.4"
rpds-py = ">=0.25.0"
@@ -2308,7 +2308,7 @@ colorama = {version = ">=0.3.4", markers = "sys_platform == \"win32\""}
win32-setctime = {version = ">=1.0.0", markers = "sys_platform == \"win32\""}
[package.extras]
dev = ["Sphinx (==8.1.3) ; python_version >= \"3.11\"", "build (==1.2.2) ; python_version >= \"3.11\"", "colorama (==0.4.5) ; python_version < \"3.8\"", "colorama (==0.4.6) ; python_version >= \"3.8\"", "exceptiongroup (==1.1.3) ; python_version >= \"3.7\" and python_version < \"3.11\"", "freezegun (==1.1.0) ; python_version < \"3.8\"", "freezegun (==1.5.0) ; python_version >= \"3.8\"", "mypy (==0.910) ; python_version < \"3.6\"", "mypy (==0.971) ; python_version == \"3.6\"", "mypy (==1.13.0) ; python_version >= \"3.8\"", "mypy (==1.4.1) ; python_version == \"3.7\"", "myst-parser (==4.0.0) ; python_version >= \"3.11\"", "pre-commit (==4.0.1) ; python_version >= \"3.9\"", "pytest (==6.1.2) ; python_version < \"3.8\"", "pytest (==8.3.2) ; python_version >= \"3.8\"", "pytest-cov (==2.12.1) ; python_version < \"3.8\"", "pytest-cov (==5.0.0) ; python_version == \"3.8\"", "pytest-cov (==6.0.0) ; python_version >= \"3.9\"", "pytest-mypy-plugins (==1.9.3) ; python_version >= \"3.6\" and python_version < \"3.8\"", "pytest-mypy-plugins (==3.1.0) ; python_version >= \"3.8\"", "sphinx-rtd-theme (==3.0.2) ; python_version >= \"3.11\"", "tox (==3.27.1) ; python_version < \"3.8\"", "tox (==4.23.2) ; python_version >= \"3.8\"", "twine (==6.0.1) ; python_version >= \"3.11\""]
dev = ["Sphinx (==8.1.3) ; python_version >= \"3.11\"", "build (==1.2.2) ; python_version >= \"3.11\"", "colorama (==0.4.5) ; python_version < \"3.8\"", "colorama (==0.4.6) ; python_version >= \"3.8\"", "exceptiongroup (==1.1.3) ; python_version >= \"3.7\" and python_version < \"3.11\"", "freezegun (==1.1.0) ; python_version < \"3.8\"", "freezegun (==1.5.0) ; python_version >= \"3.8\"", "mypy (==v0.910) ; python_version < \"3.6\"", "mypy (==v0.971) ; python_version == \"3.6\"", "mypy (==v1.13.0) ; python_version >= \"3.8\"", "mypy (==v1.4.1) ; python_version == \"3.7\"", "myst-parser (==4.0.0) ; python_version >= \"3.11\"", "pre-commit (==4.0.1) ; python_version >= \"3.9\"", "pytest (==6.1.2) ; python_version < \"3.8\"", "pytest (==8.3.2) ; python_version >= \"3.8\"", "pytest-cov (==2.12.1) ; python_version < \"3.8\"", "pytest-cov (==5.0.0) ; python_version == \"3.8\"", "pytest-cov (==6.0.0) ; python_version >= \"3.9\"", "pytest-mypy-plugins (==1.9.3) ; python_version >= \"3.6\" and python_version < \"3.8\"", "pytest-mypy-plugins (==3.1.0) ; python_version >= \"3.8\"", "sphinx-rtd-theme (==3.0.2) ; python_version >= \"3.11\"", "tox (==3.27.1) ; python_version < \"3.8\"", "tox (==4.23.2) ; python_version >= \"3.8\"", "twine (==6.0.1) ; python_version >= \"3.11\""]
[[package]]
name = "markdown-it-py"
@@ -4739,31 +4739,6 @@ files = [
{file = "wallycore-1.5.2.tar.gz", hash = "sha256:352b7b215046d4fd0333a82dc4b3936b749e5b9ae22a078b991a6b10712b3748"},
]
[[package]]
name = "wasmtime"
version = "46.0.1"
description = "A WebAssembly runtime powered by Wasmtime"
optional = false
python-versions = ">=3.9"
groups = ["main"]
files = [
{file = "wasmtime-46.0.1-py3-none-android_26_arm64_v8a.whl", hash = "sha256:cc8d52f9ad3bedc1e4de5002f7b22d7cac400be046711d177f6ce20a11eacb31"},
{file = "wasmtime-46.0.1-py3-none-android_26_x86_64.whl", hash = "sha256:f8e4b0ec402b84b856d3c6c2f96c6e6889c56e685b5cfed0a4040775c751ca38"},
{file = "wasmtime-46.0.1-py3-none-any.whl", hash = "sha256:85a092a63c20ccecb965b9aa12a19368d2e06203436d19701068efc390efa678"},
{file = "wasmtime-46.0.1-py3-none-macosx_10_13_x86_64.whl", hash = "sha256:9b46c546bf73ece2600403db7dc604c3ef12046ccf2fabe07d7bfaa00453ce8b"},
{file = "wasmtime-46.0.1-py3-none-macosx_11_0_arm64.whl", hash = "sha256:de1a69573a173b5171f9413bcf0b88f4fed2721ed02c842fc25de5358730ccdf"},
{file = "wasmtime-46.0.1-py3-none-manylinux1_x86_64.whl", hash = "sha256:e53c65abe31aeeb19a3f794b6e53140d401c4c79ad91c89caefdc502ee2b10c1"},
{file = "wasmtime-46.0.1-py3-none-manylinux2014_aarch64.whl", hash = "sha256:841b53fc17eedabaa6deb1e062a04a0a8953908d540fadb4149bc55c3f6d3e50"},
{file = "wasmtime-46.0.1-py3-none-musllinux_1_2_aarch64.whl", hash = "sha256:f295d10012b6ca6ecffa7757eed70b84ebaa2e33dc39275e7b6bed5eed5130b9"},
{file = "wasmtime-46.0.1-py3-none-musllinux_1_2_x86_64.whl", hash = "sha256:05fc65164b4825bf1c3c8f007df070b25accc974cb81d0bfc1c5d76fdf6f70e0"},
{file = "wasmtime-46.0.1-py3-none-win_amd64.whl", hash = "sha256:559b0753e3ea311fd16000fe51c08592a625e61ebb8640601ae7173fc516e430"},
{file = "wasmtime-46.0.1-py3-none-win_arm64.whl", hash = "sha256:967625406fde8fc3c9d795ffbb7bdde77d0a38de776e8eb7a0416ca6d811a27a"},
{file = "wasmtime-46.0.1.tar.gz", hash = "sha256:0da0388c21bc0f0e633c7a30f2b7939a657f5019258c9a3a63fd37298a0dbb8b"},
]
[package.extras]
testing = ["coverage", "pycparser", "pytest", "pytest-mypy"]
[[package]]
name = "websocket-client"
version = "1.9.0"
@@ -5148,4 +5123,4 @@ migration = ["psycopg2-binary"]
[metadata]
lock-version = "2.1"
python-versions = ">=3.10,<3.13"
content-hash = "3097673d0cd279b0bf2b8fa59c1e523273f63d430f2c68ccc268a3cf232068af"
content-hash = "7c70bdad0089089d383cf8f54c37c4473180cea175689b91322beaed1ab42e94"
-1
View File
@@ -53,7 +53,6 @@ dependencies = [
"greenlet~=3.3.0",
"urllib3>=2.7.0",
"pyinstrument>=5.1.2",
"wasmtime>=45.0.0",
]
[project.scripts]
+127 -1
View File
@@ -1,3 +1,4 @@
import asyncio
import base64
import hashlib
import json
@@ -12,7 +13,7 @@ from Cryptodome.Util.Padding import pad, unpad
from websockets import ServerConnection
from websockets import serve as ws_serve
from lnbits.wallets.nwc import NWCWallet
from lnbits.wallets.nwc import NWCConnection, NWCWallet
from tests.wallets.helpers import (
WalletTest,
build_test_id,
@@ -99,6 +100,8 @@ async def handle( # noqa: C901
event,
)
await websocket.send(json.dumps(["EVENT", sub_id, event]))
elif 23195 in kinds:
assert sub_filter["authors"] == [mock_settings["service_public_key"]]
elif msg[0] == "EVENT":
event = msg[1]
decrypted_content = decrypt_content(
@@ -177,6 +180,129 @@ async def run(data: WalletTest):
await nwcwallet.cleanup()
@pytest.mark.anyio
async def test_nwc_rejects_event_from_unexpected_pubkey(mocker):
async def _noop(*args, **kwargs):
return None
mocker.patch("lnbits.wallets.nwc.NWCConnection._connect_to_relay", new=_noop)
mocker.patch("lnbits.wallets.nwc.NWCConnection._handle_timeouts", new=_noop)
service_private_key = PrivateKey()
service_public_key = service_private_key.public_key.format().hex()[2:]
attacker_private_key = PrivateKey()
attacker_public_key = attacker_private_key.public_key.format().hex()[2:]
account_private_key = PrivateKey()
conn = NWCConnection(
service_public_key,
account_private_key.secret.hex(),
"ws://127.0.0.1:8555",
)
try:
event = {
"kind": 23195,
"content": "{}",
"created_at": int(time.time()),
"tags": [["e", "request-event-id"]],
}
sign_event(attacker_public_key, attacker_private_key.secret.hex(), event)
with pytest.raises(Exception, match="Invalid event signature"):
await conn._on_event_message(["EVENT", "subid", event])
finally:
await conn.close()
@pytest.mark.anyio
async def test_nwc_marks_pending_invoice_settled_only_once():
wallet = NWCWallet.__new__(NWCWallet)
wallet.pending_invoice_details = {"checking-id": {"checking_id": "checking-id"}}
wallet.pending_invoices = ["checking-id"]
wallet.paid_invoices_queue = asyncio.Queue(0)
wallet._mark_invoice_settled("checking-id", source="notification")
wallet._mark_invoice_settled("checking-id", source="notification")
assert wallet.paid_invoices_queue.qsize() == 1
assert await wallet.paid_invoices_queue.get() == "checking-id"
@pytest.mark.anyio
async def test_nwc_registers_notification_subscriptions(mocker):
async def _noop(*args, **kwargs):
return None
mocker.patch("lnbits.wallets.nwc.NWCConnection._connect_to_relay", new=_noop)
mocker.patch("lnbits.wallets.nwc.NWCConnection._handle_timeouts", new=_noop)
service_private_key = PrivateKey()
service_public_key = service_private_key.public_key.format().hex()[2:]
account_private_key = PrivateKey()
conn = NWCConnection(
service_public_key,
account_private_key.secret.hex(),
"ws://127.0.0.1:8555",
)
send_mock = mocker.patch.object(conn, "_send", mocker.AsyncMock())
try:
await conn._subscribe_to_notifications()
assert len(conn.notification_subscription_ids) == 2
assert len(conn.subscriptions) == 2
assert set(conn.subscriptions.keys()) == conn.notification_subscription_ids
assert all(
subscription["method"] == "notification_sub"
and subscription["event_id"] == subscription["sub_id"]
for subscription in conn.subscriptions.values()
)
assert send_mock.await_count == 2
finally:
await conn.close()
@pytest.mark.anyio
async def test_nwc_spreads_fallback_lookups_with_cooldown(mocker):
def _schedule_next_lookup(
invoice: dict[str, object], now: float | None = None
) -> None:
assert now is not None
invoice["next_lookup_at"] = now + 1
wallet = NWCWallet.__new__(NWCWallet)
wallet.shutdown = False
wallet.pending_invoices = ["checking-1", "checking-2"]
wallet.pending_invoice_details = {
"checking-1": {
"checking_id": "checking-1",
"next_lookup_at": 0.0,
"lookup_attempts": 0,
},
"checking-2": {
"checking_id": "checking-2",
"next_lookup_at": 0.0,
"lookup_attempts": 0,
},
}
wallet.pending_invoices_lookup_cooldown = 1.0
wallet._is_shutting_down = lambda: False
wallet._payment_data_is_settled = lambda payment_data: False
wallet._cache_payment_data = lambda *args, **kwargs: None
wallet._schedule_next_lookup = _schedule_next_lookup
wallet.conn = mocker.Mock()
wallet.conn.get_info = mocker.AsyncMock()
wallet.conn.supports_method = mocker.Mock(return_value=True)
wallet.conn.call = mocker.AsyncMock(return_value={"settled_at": None})
sleep_mock = mocker.patch("lnbits.wallets.nwc.asyncio.sleep", mocker.AsyncMock())
await wallet._run_fallback_lookups(100.0)
assert wallet.conn.call.await_count == 2
sleep_mock.assert_awaited_once_with(1.0)
@pytest.mark.anyio
@pytest.mark.parametrize(
"test_data",
-368
View File
@@ -1,368 +0,0 @@
from __future__ import annotations
import argparse
import re
from collections import defaultdict
from collections.abc import Sequence
from pathlib import Path
from types import UnionType
from typing import Any, Literal, Union, get_args, get_origin
from pydantic import BaseModel
from lnbits.core.extensions import (
ExtensionAPI,
ExtensionAPIMethod,
get_extension_api_method,
list_extension_api_methods,
)
def generate_typescript_sdk(
api_cls: type[ExtensionAPI] | None = None,
method_ids: Sequence[str] | None = None,
) -> str:
api_cls = api_cls or ExtensionAPI
methods = _select_methods(api_cls, method_ids)
models = _collect_models(methods)
lines = [
"/* Generated by LNbits ExtensionAPI codegen. */",
"/* Do not edit by hand. */",
"",
"export type MaybePromise<T> = T | Promise<T>",
"",
]
for model in models:
lines.extend(_render_model_type(model))
lines.append("")
lines.extend(_render_method_metadata(methods))
lines.append("")
lines.extend(_render_host_type(methods))
lines.append("")
lines.extend(_render_sdk_type(methods))
lines.append("")
lines.extend(_render_create_sdk(methods))
return "\n".join(lines).rstrip() + "\n"
def write_typescript_sdk(
path: str | Path,
api_cls: type[ExtensionAPI] | None = None,
method_ids: Sequence[str] | None = None,
) -> None:
Path(path).write_text(
generate_typescript_sdk(api_cls, method_ids), encoding="utf-8"
)
def _select_methods(
api_cls: type[ExtensionAPI], method_ids: Sequence[str] | None
) -> list[ExtensionAPIMethod]:
if not method_ids:
return list_extension_api_methods(api_cls)
return [get_extension_api_method(method_id, api_cls) for method_id in method_ids]
def _collect_models(methods: Sequence[ExtensionAPIMethod]) -> list[type[BaseModel]]:
models: dict[str, type[BaseModel]] = {}
pending = [
model
for method in methods
for model in (method.request_model, method.response_model)
]
while pending:
model = pending.pop()
if model.__name__ in models:
continue
models[model.__name__] = model
for field in model.__fields__.values():
pending.extend(_nested_model_types(field.outer_type_))
return [models[name] for name in sorted(models)]
def _nested_model_types(type_: Any) -> list[type[BaseModel]]:
models: list[type[BaseModel]] = []
if _is_model_type(type_):
models.append(type_)
for arg in get_args(type_):
models.extend(_nested_model_types(arg))
return models
def _render_model_type(model: type[BaseModel]) -> list[str]:
name = _model_name(model)
fields = model.__fields__
if not fields:
return [f"export type {name} = Record<string, never>"]
lines = [f"export type {name} = {{"]
for field_name, field in fields.items():
optional = "?" if not field.required else ""
ts_type = _python_type_to_ts(field.outer_type_, field.allow_none)
lines.append(f" {_camel(field_name)}{optional}: {ts_type}")
lines.append("}")
return lines
def _python_type_to_ts(type_: Any, allow_none: bool = False) -> str:
origin = get_origin(type_)
args = get_args(type_)
if origin in (UnionType, Union):
ts = " | ".join(
_python_type_to_ts(arg) for arg in args if arg is not type(None)
)
if type(None) in args:
ts = f"{ts} | null"
return ts
if origin is Literal:
return " | ".join(_literal_to_ts(arg) for arg in args)
if _is_model_type(type_):
ts = _model_name(type_)
elif origin in (list, Sequence):
item_type = _python_type_to_ts(args[0]) if args else "unknown"
ts = f"{item_type}[]"
elif origin is dict:
key_type = _python_type_to_ts(args[0]) if args else "string"
value_type = _python_type_to_ts(args[1]) if len(args) > 1 else "unknown"
ts = (
f"Record<{key_type}, {value_type}>"
if key_type == "string"
else f"{{ [key: string]: {value_type} }}"
)
elif _is_subclass(type_, str):
ts = "string"
elif _is_subclass(type_, bool):
ts = "boolean"
elif _is_subclass(type_, int) or _is_subclass(type_, float):
ts = "number"
else:
ts = "unknown"
return f"{ts} | null" if allow_none else ts
def _literal_to_ts(value: Any) -> str:
if isinstance(value, str):
return f'"{value}"'
if isinstance(value, bool):
return "true" if value else "false"
if value is None:
return "null"
return str(value)
def _is_subclass(type_: Any, class_: type) -> bool:
try:
return isinstance(type_, type) and issubclass(type_, class_)
except TypeError:
return False
def _is_model_type(type_: Any) -> bool:
return _is_subclass(type_, BaseModel)
def _render_method_metadata(
methods: Sequence[ExtensionAPIMethod],
) -> list[str]:
lines = ["export const extensionApiMethods = ["]
for method in methods:
permission = (
f'"{method.required_permission}"' if method.required_permission else "null"
)
lines.extend(
[
" {",
f' id: "{method.method_id}",',
f' namespace: "{method.namespace}",',
f' sdkName: "{method.sdk_name}",',
f' pythonName: "{method.python_name}",',
f' hostInterface: "{method.host_interface}",',
f' hostName: "{method.host_name}",',
f' hostJsName: "{_camel(method.host_name)}",',
f" requiredPermission: {permission},",
" },",
]
)
lines.append("] as const")
return lines
def _render_host_type(methods: Sequence[ExtensionAPIMethod]) -> list[str]:
lines = ["export type ExtensionHost = {"]
for host_interface, interface_methods in _methods_by_host_interface(
methods
).items():
lines.append(f" {_ts_property(host_interface)}: {{")
for method in sorted(interface_methods, key=lambda item: item.host_name):
request = _model_name(method.request_model)
response = _model_name(method.response_model)
if _is_empty_model(method.request_model):
lines.append(
f" {_camel(method.host_name)}(): MaybePromise<{response}>"
)
else:
lines.append(
f" {_camel(method.host_name)}"
f"(input: {request}): MaybePromise<{response}>"
)
lines.append(" }")
lines.append("}")
return lines
def _render_sdk_type(methods: Sequence[ExtensionAPIMethod]) -> list[str]:
lines = ["export type ExtensionSdk = {"]
_render_sdk_type_node(lines, _namespace_tree(methods), 1)
lines.append("}")
return lines
def _render_create_sdk(methods: Sequence[ExtensionAPIMethod]) -> list[str]:
lines = [
"export function createExtensionSdk(",
" host: ExtensionHost",
"): ExtensionSdk {",
" return {",
]
_render_create_sdk_node(lines, _namespace_tree(methods), 2)
lines.extend([" }", "}"])
return lines
def _render_sdk_type_node(lines: list[str], node: dict[str, Any], level: int) -> None:
indent = " " * level
for namespace, child in _iter_child_namespaces(node):
lines.append(f"{indent}{namespace}: {{")
_render_sdk_type_node(lines, child, level + 1)
lines.append(f"{indent}}}")
for method in node.get("__methods__", []):
request = _model_name(method.request_model)
response = _model_name(method.response_model)
if _is_empty_model(method.request_model):
lines.append(f"{indent}{method.sdk_name}(): Promise<{response}>")
else:
lines.append(
f"{indent}{method.sdk_name}(input: {request}): Promise<{response}>"
)
def _render_create_sdk_node(lines: list[str], node: dict[str, Any], level: int) -> None:
indent = " " * level
for namespace, child in _iter_child_namespaces(node):
lines.append(f"{indent}{namespace}: {{")
_render_create_sdk_node(lines, child, level + 1)
lines.append(f"{indent}}},")
for method in node.get("__methods__", []):
host_call_target = (
f"host{_ts_access(method.host_interface)}"
f"{_ts_access(_camel(method.host_name))}"
)
if _is_empty_model(method.request_model):
signature = f"{method.sdk_name}()"
host_call = f"{host_call_target}()"
else:
signature = f"{method.sdk_name}(input)"
host_call = f"{host_call_target}(input)"
lines.extend(
[
f"{indent}async {signature} {{",
f"{indent} return {host_call}",
f"{indent}}},",
]
)
def _methods_by_host_interface(
methods: Sequence[ExtensionAPIMethod],
) -> dict[str, list[ExtensionAPIMethod]]:
interfaces: dict[str, list[ExtensionAPIMethod]] = defaultdict(list)
for method in sorted(
methods, key=lambda item: (item.host_interface, item.host_name)
):
interfaces[method.host_interface].append(method)
return dict(sorted(interfaces.items()))
def _namespace_tree(methods: Sequence[ExtensionAPIMethod]) -> dict[str, Any]:
tree: dict[str, Any] = {}
for method in sorted(methods, key=lambda item: (item.namespace, item.sdk_name)):
node = tree
for part in method.namespace.split("."):
node = node.setdefault(part, {})
node.setdefault("__methods__", []).append(method)
return tree
def _iter_child_namespaces(
node: dict[str, Any],
) -> list[tuple[str, dict[str, Any]]]:
return [
(key, value)
for key, value in sorted(node.items())
if key != "__methods__" and isinstance(value, dict)
]
def _model_name(model: type[BaseModel]) -> str:
return model.__name__
def _camel(value: str) -> str:
head, *tail = value.split("_")
return head + "".join(part.capitalize() for part in tail)
def _ts_property(value: str) -> str:
if re.match(r"^[A-Za-z_$][A-Za-z0-9_$]*$", value):
return value
return f'"{value}"'
def _ts_access(value: str) -> str:
if re.match(r"^[A-Za-z_$][A-Za-z0-9_$]*$", value):
return f".{value}"
return f'["{value}"]'
def _is_empty_model(model: type[BaseModel]) -> bool:
return not model.__fields__
def main(argv: Sequence[str] | None = None) -> int:
parser = argparse.ArgumentParser(
description="Generate a TypeScript SDK from the LNbits ExtensionAPI."
)
parser.add_argument(
"--method",
action="append",
dest="method_ids",
help="ExtensionAPI method id to include. Can be passed multiple times.",
)
parser.add_argument(
"--out",
help="Output file. If omitted, the generated SDK is printed to stdout.",
)
args = parser.parse_args(argv)
sdk = generate_typescript_sdk(method_ids=args.method_ids)
if args.out:
Path(args.out).write_text(sdk, encoding="utf-8")
else:
print(sdk, end="")
return 0
if __name__ == "__main__":
raise SystemExit(main())
Generated
-21
View File
@@ -1333,7 +1333,6 @@ dependencies = [
{ name = "urllib3" },
{ name = "uvicorn" },
{ name = "uvloop" },
{ name = "wasmtime" },
{ name = "websocket-client" },
{ name = "websockets" },
]
@@ -1423,7 +1422,6 @@ requires-dist = [
{ name = "uvicorn", specifier = "~=0.40.0" },
{ name = "uvloop", specifier = "~=0.22.1" },
{ name = "wallycore", marker = "extra == 'liquid'", specifier = "~=1.5.1" },
{ name = "wasmtime", specifier = ">=45.0.0" },
{ name = "websocket-client", specifier = "~=1.9.0" },
{ name = "websockets", specifier = "~=15.0.1" },
]
@@ -2860,25 +2858,6 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/41/e5/1c1d2979d074cba4f0a1516f2bf4c3ee66067b6ba3b96e5cb4460555d474/wallycore-1.5.3-cp312-cp312-win_amd64.whl", hash = "sha256:c3343a1df11a7ef4572521e002f4550a0157aea2c749dcfd9be62fb5babe0d03", size = 1728038, upload-time = "2026-04-15T22:04:36.434Z" },
]
[[package]]
name = "wasmtime"
version = "45.0.0"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/b1/ff/db9cfc61d988bc15303134bb174176a29839976876dfd18c3a12548ad291/wasmtime-45.0.0.tar.gz", hash = "sha256:2ad4bf7ca286ceea35c1e420d10b368d7f83faf9a5ffde87b4ee334a9b7f55f3", size = 128297, upload-time = "2026-05-26T17:57:39.131Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/75/56/7d941adba273210dcf4198266a47f472a5eeca20172005b443af71a9a3e7/wasmtime-45.0.0-py3-none-android_26_arm64_v8a.whl", hash = "sha256:4e843795b53e66c71313f2254731467372e5e1549227cf14accb9e2d57701c10", size = 8659052, upload-time = "2026-05-26T17:57:12.338Z" },
{ url = "https://files.pythonhosted.org/packages/3f/81/c4d81ebf3db8aa28f789a9569640f30790d5234c509a0234cd502aa2638b/wasmtime-45.0.0-py3-none-android_26_x86_64.whl", hash = "sha256:35e713f907264e470f3bc9b592b81b8ed0f8f5651725d9f07a5d52beb0642e38", size = 9619373, upload-time = "2026-05-26T17:57:14.979Z" },
{ url = "https://files.pythonhosted.org/packages/e1/c7/7594da7fa8a3bc5e765733ad57aac9b7b27262c4afa47521bd500e4a4574/wasmtime-45.0.0-py3-none-any.whl", hash = "sha256:6251ee5074a8b8bfaa98e6e99cb5d49d6d0f2320b3265d5aa6c2ee5df5fb4519", size = 8019034, upload-time = "2026-05-26T17:57:20.138Z" },
{ url = "https://files.pythonhosted.org/packages/75/76/7d0e440ca03a717a97889dbb7b68f952c20ed4ffd3f59addf9553579e1d5/wasmtime-45.0.0-py3-none-macosx_10_13_x86_64.whl", hash = "sha256:3579b0ec6d001750d66ec7089aaeee2c048f88328c82743e15f099af01b0cf84", size = 9401625, upload-time = "2026-05-26T17:57:22.149Z" },
{ url = "https://files.pythonhosted.org/packages/5b/0b/a81b5daf5adea482ecb68d9615f6a348486ab4d8e980a915d4420e57ee4d/wasmtime-45.0.0-py3-none-macosx_11_0_arm64.whl", hash = "sha256:31d10f25c330cebcfb364e9a357123deeec96c41725ff2bba91b705587f38a93", size = 8255954, upload-time = "2026-05-26T17:57:24.769Z" },
{ url = "https://files.pythonhosted.org/packages/d7/8c/e9019a28e908214031310aefd78e4755221d02303190b54b2c85cb69573e/wasmtime-45.0.0-py3-none-manylinux1_x86_64.whl", hash = "sha256:5d1416ec6da8cd87c29e2e9eb074358c91839c2fff971fe428c8921eaae68e73", size = 9681185, upload-time = "2026-05-26T17:57:26.641Z" },
{ url = "https://files.pythonhosted.org/packages/42/56/ed5f492bd553a31c8e28d621f8256f2c7b1a133b28f73525d96ca355891a/wasmtime-45.0.0-py3-none-manylinux2014_aarch64.whl", hash = "sha256:a499f6ab0eebb70dca83d6a4904b743cd122f322af3abe86af08ad753533d946", size = 8582001, upload-time = "2026-05-26T17:57:28.883Z" },
{ url = "https://files.pythonhosted.org/packages/62/12/9b41740da83f51014b88181c9086de0ed75d736a5329baff7323c4fb6eff/wasmtime-45.0.0-py3-none-musllinux_1_2_aarch64.whl", hash = "sha256:bef65282b7de744106a91da43e4d06ba19d2d587bc54abb83b3e757f0c4fc030", size = 8633462, upload-time = "2026-05-26T17:57:31.423Z" },
{ url = "https://files.pythonhosted.org/packages/ea/63/49d8317706a108d9ed1d4166d0fc710796da1b20e591a98a96575dec367a/wasmtime-45.0.0-py3-none-musllinux_1_2_x86_64.whl", hash = "sha256:a0b6ca14b4628a5d1ffa91ccf2c0f2c58fa171f126ec085d564b09d5795395dd", size = 9712524, upload-time = "2026-05-26T17:57:33.839Z" },
{ url = "https://files.pythonhosted.org/packages/20/71/8e31ea472ceb934e7261ac59a786e82cd82b4d4dcb7c870d498aa9c3c21e/wasmtime-45.0.0-py3-none-win_amd64.whl", hash = "sha256:1736a70a48f713aaf1a878514d29cc6f554213b5431e04447813a3b9b4320381", size = 8019039, upload-time = "2026-05-26T17:57:36.04Z" },
{ url = "https://files.pythonhosted.org/packages/9c/d1/ac536e92ac95a02e137be5b6829f15b87d5eef93ace32e5ee8035155b839/wasmtime-45.0.0-py3-none-win_arm64.whl", hash = "sha256:ae9726590e6d90c6305b8b507c93468b145204d4390aa9a2e29e26babcae110e", size = 6845659, upload-time = "2026-05-26T17:57:37.696Z" },
]
[[package]]
name = "websocket-client"
version = "1.9.0"