chore: clean-up

This commit is contained in:
Vlad Stan
2026-07-08 11:54:57 +03:00
parent b74b75d303
commit 2b791d1c23
5 changed files with 80 additions and 232 deletions
-3
View File
@@ -9,7 +9,6 @@ from .api import (
list_extension_api_methods, list_extension_api_methods,
) )
from .loader import WasmExtension, load_wasm_extension, register_wasm_extension from .loader import WasmExtension, load_wasm_extension, register_wasm_extension
from .prototype import InMemoryExtensionAPI, InMemoryExtensionState
from .runtime import ExtensionAPIHost from .runtime import ExtensionAPIHost
from .wasm import invoke_wasm_extension_export from .wasm import invoke_wasm_extension_export
@@ -17,8 +16,6 @@ __all__ = [
"ExtensionAPI", "ExtensionAPI",
"ExtensionAPIHost", "ExtensionAPIHost",
"ExtensionAPIMethod", "ExtensionAPIMethod",
"InMemoryExtensionAPI",
"InMemoryExtensionState",
"WasmExtension", "WasmExtension",
"extension_api_contract", "extension_api_contract",
"extension_api_method", "extension_api_method",
+77 -27
View File
@@ -1,13 +1,14 @@
from __future__ import annotations from __future__ import annotations
import inspect import inspect
import json
import logging import logging
import secrets import secrets
import time import time
from collections.abc import Awaitable, Callable, Iterable from collections.abc import Awaitable, Callable, Iterable
from dataclasses import dataclass from dataclasses import dataclass
from functools import wraps from functools import wraps
from typing import NoReturn, TypeVar, cast, get_type_hints from typing import TypeVar, cast, get_type_hints
from pydantic import BaseModel from pydantic import BaseModel
@@ -29,8 +30,13 @@ from .models import (
StoragePaginatedResponse, StoragePaginatedResponse,
StorageSetRequest, StorageSetRequest,
StorageSetResponse, StorageSetResponse,
WatchPaymentRequest, UserWalletSummary,
WatchPaymentResponse, )
from .storage import (
storage_delete_row,
storage_get_paginated_rows,
storage_get_row,
storage_set_row,
) )
logger = logging.getLogger("lnbits.extensions") logger = logging.getLogger("lnbits.extensions")
@@ -136,7 +142,8 @@ class ExtensionAPI:
required_permission="ext.storage.read_write", required_permission="ext.storage.read_write",
) )
async def storage_get(self, request: StorageGetRequest) -> StorageGetResponse: async def storage_get(self, request: StorageGetRequest) -> StorageGetResponse:
self._raise_unwired_runtime("storage_get") row = await storage_get_row(self.extension_id, request.table, request.id)
return StorageGetResponse(data_json=json.dumps(row) if row else None)
@extension_api_method( @extension_api_method(
method_id="storage.set", method_id="storage.set",
@@ -148,7 +155,8 @@ class ExtensionAPI:
required_permission="ext.storage.read_write", required_permission="ext.storage.read_write",
) )
async def storage_set(self, request: StorageSetRequest) -> StorageSetResponse: async def storage_set(self, request: StorageSetRequest) -> StorageSetResponse:
self._raise_unwired_runtime("storage_set") await storage_set_row(self.extension_id, request.table, request.data)
return StorageSetResponse()
@extension_api_method( @extension_api_method(
method_id="storage.get_paginated", method_id="storage.get_paginated",
@@ -162,7 +170,21 @@ class ExtensionAPI:
async def storage_get_paginated( async def storage_get_paginated(
self, request: StoragePaginatedRequest self, request: StoragePaginatedRequest
) -> StoragePaginatedResponse: ) -> StoragePaginatedResponse:
self._raise_unwired_runtime("storage_get_paginated") page = await storage_get_paginated_rows(
self.extension_id,
request.table,
request.filters,
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( @extension_api_method(
method_id="storage.delete", method_id="storage.delete",
@@ -176,7 +198,8 @@ class ExtensionAPI:
async def storage_delete( async def storage_delete(
self, request: StorageDeleteRequest self, request: StorageDeleteRequest
) -> StorageDeleteResponse: ) -> StorageDeleteResponse:
self._raise_unwired_runtime("storage_delete") await storage_delete_row(self.extension_id, request.table, request.id)
return StorageDeleteResponse()
@extension_api_method( @extension_api_method(
method_id="wallet.create_invoice", method_id="wallet.create_invoice",
@@ -190,7 +213,36 @@ class ExtensionAPI:
async def wallet_create_invoice( async def wallet_create_invoice(
self, request: CreateInvoiceRequest self, request: CreateInvoiceRequest
) -> CreateInvoiceResponse: ) -> CreateInvoiceResponse:
self._raise_unwired_runtime("wallet_create_invoice") 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_sat,
unit=request.currency or "sat",
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( @extension_api_method(
method_id="wallet.list_user_wallets", method_id="wallet.list_user_wallets",
@@ -204,21 +256,24 @@ class ExtensionAPI:
async def wallet_list_user_wallets( async def wallet_list_user_wallets(
self, request: EmptyRequest self, request: EmptyRequest
) -> ListUserWalletsResponse: ) -> ListUserWalletsResponse:
self._raise_unwired_runtime("wallet_list_user_wallets") if not self.user_id:
raise PermissionError(
"Listing user wallets requires an authenticated user context."
)
@extension_api_method( from lnbits.core.crud.wallets import get_wallets
method_id="payments.watch",
namespace="payments", user_wallets = await get_wallets(self.user_id)
name="Watch payment", if user_wallets is None:
host_name="watch_payment", raise PermissionError(
sdk_name="watch", "Listing user wallets requires an authenticated user context."
description="Subscribe the extension to a payment state callback.", )
required_permission="payments.watch", return ListUserWalletsResponse(
) wallets=[
async def payments_watch( UserWalletSummary(id=w.id, name=w.name, currency=w.currency)
self, request: WatchPaymentRequest for w in user_wallets
) -> WatchPaymentResponse: ]
self._raise_unwired_runtime("payments_watch") )
@extension_api_method( @extension_api_method(
method_id="system.random_id", method_id="system.random_id",
@@ -257,11 +312,6 @@ class ExtensionAPI:
log("extension:%s %s", self.extension_id, request.message) log("extension:%s %s", self.extension_id, request.message)
return LogResponse() return LogResponse()
def _raise_unwired_runtime(self, method_name: str) -> NoReturn:
raise NotImplementedError(
f"ExtensionAPI.{method_name} must be wired to LNbits services before use."
)
def list_extension_api_methods( def list_extension_api_methods(
api_cls: type[ExtensionAPI] = ExtensionAPI, api_cls: type[ExtensionAPI] = ExtensionAPI,
-9
View File
@@ -98,15 +98,6 @@ class ListUserWalletsResponse(BaseModel):
wallets: list[UserWalletSummary] = Field(default_factory=list) wallets: list[UserWalletSummary] = Field(default_factory=list)
class WatchPaymentRequest(BaseModel):
payment_hash: str = Field(..., min_length=1, max_length=128)
callback_export: str = Field(..., min_length=1, max_length=128)
class WatchPaymentResponse(BaseModel):
ok: bool = True
class RandomIdRequest(BaseModel): class RandomIdRequest(BaseModel):
prefix: str = Field(..., min_length=1, max_length=32) prefix: str = Field(..., min_length=1, max_length=32)
-155
View File
@@ -1,155 +0,0 @@
from __future__ import annotations
import json
from dataclasses import dataclass, field
from .api import ExtensionAPI
from .models import (
CreateInvoiceRequest,
CreateInvoiceResponse,
EmptyRequest,
ListUserWalletsResponse,
StorageDeleteRequest,
StorageDeleteResponse,
StorageGetRequest,
StorageGetResponse,
StoragePaginatedRequest,
StoragePaginatedResponse,
StorageSetRequest,
StorageSetResponse,
UserWalletSummary,
WatchPaymentRequest,
WatchPaymentResponse,
)
from .storage import (
storage_delete_row,
storage_get_paginated_rows,
storage_get_row,
storage_set_row,
)
@dataclass
class InMemoryExtensionState:
payment_watchers: dict[str, dict[str, str]] = field(default_factory=dict)
user_wallets: dict[str, list[UserWalletSummary] | None] = field(
default_factory=dict
)
class InMemoryExtensionAPI(ExtensionAPI):
def __init__(
self,
extension_id: str,
permissions: set[str],
*,
state: InMemoryExtensionState | None = None,
user_id: str | None = None,
) -> None:
super().__init__(extension_id, permissions, user_id=user_id)
self.state = state or InMemoryExtensionState()
async def storage_get(self, request: StorageGetRequest) -> StorageGetResponse:
self.require_permission("ext.storage.read_write")
row = await storage_get_row(self.extension_id, request.table, request.id)
return StorageGetResponse(data_json=json.dumps(row) if row else None)
async def storage_set(self, request: StorageSetRequest) -> StorageSetResponse:
self.require_permission("ext.storage.read_write")
await storage_set_row(self.extension_id, request.table, request.data)
return StorageSetResponse()
async def storage_get_paginated(
self, request: StoragePaginatedRequest
) -> StoragePaginatedResponse:
self.require_permission("ext.storage.read_write")
page = await storage_get_paginated_rows(
self.extension_id,
request.table,
request.filters,
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"],
)
async def storage_delete(
self, request: StorageDeleteRequest
) -> StorageDeleteResponse:
self.require_permission("ext.storage.read_write")
await storage_delete_row(self.extension_id, request.table, request.id)
return StorageDeleteResponse()
async def wallet_create_invoice(
self, request: CreateInvoiceRequest
) -> CreateInvoiceResponse:
self.require_permission("wallet.create_invoice")
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_sat,
unit=request.currency or "sat",
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,
)
async def wallet_list_user_wallets(
self, _request: EmptyRequest
) -> ListUserWalletsResponse:
self.require_permission("wallet.list")
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
]
)
async def payments_watch(
self, request: WatchPaymentRequest
) -> WatchPaymentResponse:
self.require_permission("payments.watch")
self.state.payment_watchers.setdefault(self.extension_id, {})[
request.payment_hash
] = request.callback_export
return WatchPaymentResponse()
+3 -38
View File
@@ -11,10 +11,8 @@ from fastapi import FastAPI
from lnbits.core.crud.extensions import get_installed_extension from lnbits.core.crud.extensions import get_installed_extension
from .api import list_extension_api_methods from .api import ExtensionAPI, list_extension_api_methods
from .loader import WasmExtension, register_wasm_extension from .loader import WasmExtension, register_wasm_extension
from .models import UserWalletSummary
from .prototype import InMemoryExtensionAPI, InMemoryExtensionState
from .runtime import ExtensionAPIHost from .runtime import ExtensionAPIHost
@@ -27,13 +25,10 @@ async def invoke_wasm_extension_export(
user: Any | None = None, user: Any | None = None,
) -> dict[str, Any]: ) -> dict[str, Any]:
extension = _get_registered_extension(app, ext_id) extension = _get_registered_extension(app, ext_id)
state = _get_extension_state(app)
state.user_wallets[extension.id] = _user_wallet_summaries(user)
permissions = await _extension_permissions(extension) permissions = await _extension_permissions(extension)
api = InMemoryExtensionAPI( api = ExtensionAPI(
extension.id, extension.id,
permissions, permissions,
state=state,
user_id=_user_id(user), user_id=_user_id(user),
) )
@@ -54,7 +49,7 @@ def _invoke_wasm_extension_export_sync(
extension: WasmExtension, extension: WasmExtension,
export_name: str, export_name: str,
payload: Mapping[str, Any], payload: Mapping[str, Any],
api: InMemoryExtensionAPI, api: ExtensionAPI,
) -> dict[str, Any]: ) -> dict[str, Any]:
try: try:
from wasmtime import Store, WasiConfig, component from wasmtime import Store, WasiConfig, component
@@ -198,14 +193,6 @@ def _get_registered_extension(app: FastAPI, ext_id: str) -> WasmExtension:
return register_wasm_extension(app, ext_id) return register_wasm_extension(app, ext_id)
def _get_extension_state(app: FastAPI) -> InMemoryExtensionState:
state = getattr(app.state, "lnbits_extension_state", None)
if not state:
state = InMemoryExtensionState()
app.state.lnbits_extension_state = state
return state
async def _extension_permissions(extension: WasmExtension) -> set[str]: async def _extension_permissions(extension: WasmExtension) -> set[str]:
installed_extension = await get_installed_extension(extension.id) installed_extension = await get_installed_extension(extension.id)
if not installed_extension: if not installed_extension:
@@ -217,27 +204,5 @@ def _user_id(user: Any | None) -> str | None:
return getattr(user, "id", None) if user else None return getattr(user, "id", None) if user else None
def _user_wallet_summaries(user: Any | None) -> list[UserWalletSummary] | None:
if user is None:
return None
summaries: list[UserWalletSummary] = []
for wallet in getattr(user, "wallets", []) or []:
if not getattr(wallet, "can_receive_payments", False):
continue
wallet_id = getattr(wallet, "id", None)
wallet_name = getattr(wallet, "name", None)
if not wallet_id or not wallet_name:
continue
summaries.append(
UserWalletSummary(
id=wallet_id,
name=wallet_name,
currency=getattr(wallet, "currency", None),
)
)
return summaries
def _camel_to_kebab(value: str) -> str: def _camel_to_kebab(value: str) -> str:
return re.sub(r"([a-z0-9])([A-Z])", r"\1-\2", value).replace("_", "-").lower() return re.sub(r"([a-z0-9])([A-Z])", r"\1-\2", value).replace("_", "-").lower()