chore: clean-up
This commit is contained in:
@@ -9,7 +9,6 @@ from .api import (
|
||||
list_extension_api_methods,
|
||||
)
|
||||
from .loader import WasmExtension, load_wasm_extension, register_wasm_extension
|
||||
from .prototype import InMemoryExtensionAPI, InMemoryExtensionState
|
||||
from .runtime import ExtensionAPIHost
|
||||
from .wasm import invoke_wasm_extension_export
|
||||
|
||||
@@ -17,8 +16,6 @@ __all__ = [
|
||||
"ExtensionAPI",
|
||||
"ExtensionAPIHost",
|
||||
"ExtensionAPIMethod",
|
||||
"InMemoryExtensionAPI",
|
||||
"InMemoryExtensionState",
|
||||
"WasmExtension",
|
||||
"extension_api_contract",
|
||||
"extension_api_method",
|
||||
|
||||
@@ -1,13 +1,14 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
import json
|
||||
import logging
|
||||
import secrets
|
||||
import time
|
||||
from collections.abc import Awaitable, Callable, Iterable
|
||||
from dataclasses import dataclass
|
||||
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
|
||||
|
||||
@@ -29,8 +30,13 @@ from .models import (
|
||||
StoragePaginatedResponse,
|
||||
StorageSetRequest,
|
||||
StorageSetResponse,
|
||||
WatchPaymentRequest,
|
||||
WatchPaymentResponse,
|
||||
UserWalletSummary,
|
||||
)
|
||||
from .storage import (
|
||||
storage_delete_row,
|
||||
storage_get_paginated_rows,
|
||||
storage_get_row,
|
||||
storage_set_row,
|
||||
)
|
||||
|
||||
logger = logging.getLogger("lnbits.extensions")
|
||||
@@ -136,7 +142,8 @@ class ExtensionAPI:
|
||||
required_permission="ext.storage.read_write",
|
||||
)
|
||||
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(
|
||||
method_id="storage.set",
|
||||
@@ -148,7 +155,8 @@ class ExtensionAPI:
|
||||
required_permission="ext.storage.read_write",
|
||||
)
|
||||
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(
|
||||
method_id="storage.get_paginated",
|
||||
@@ -162,7 +170,21 @@ class ExtensionAPI:
|
||||
async def storage_get_paginated(
|
||||
self, request: StoragePaginatedRequest
|
||||
) -> 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(
|
||||
method_id="storage.delete",
|
||||
@@ -176,7 +198,8 @@ class ExtensionAPI:
|
||||
async def storage_delete(
|
||||
self, request: StorageDeleteRequest
|
||||
) -> StorageDeleteResponse:
|
||||
self._raise_unwired_runtime("storage_delete")
|
||||
await storage_delete_row(self.extension_id, request.table, request.id)
|
||||
return StorageDeleteResponse()
|
||||
|
||||
@extension_api_method(
|
||||
method_id="wallet.create_invoice",
|
||||
@@ -190,7 +213,36 @@ class ExtensionAPI:
|
||||
async def wallet_create_invoice(
|
||||
self, request: CreateInvoiceRequest
|
||||
) -> 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(
|
||||
method_id="wallet.list_user_wallets",
|
||||
@@ -204,21 +256,24 @@ class ExtensionAPI:
|
||||
async def wallet_list_user_wallets(
|
||||
self, request: EmptyRequest
|
||||
) -> 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(
|
||||
method_id="payments.watch",
|
||||
namespace="payments",
|
||||
name="Watch payment",
|
||||
host_name="watch_payment",
|
||||
sdk_name="watch",
|
||||
description="Subscribe the extension to a payment state callback.",
|
||||
required_permission="payments.watch",
|
||||
)
|
||||
async def payments_watch(
|
||||
self, request: WatchPaymentRequest
|
||||
) -> WatchPaymentResponse:
|
||||
self._raise_unwired_runtime("payments_watch")
|
||||
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="system.random_id",
|
||||
@@ -257,11 +312,6 @@ class ExtensionAPI:
|
||||
log("extension:%s %s", self.extension_id, request.message)
|
||||
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(
|
||||
api_cls: type[ExtensionAPI] = ExtensionAPI,
|
||||
|
||||
@@ -98,15 +98,6 @@ class ListUserWalletsResponse(BaseModel):
|
||||
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):
|
||||
prefix: str = Field(..., min_length=1, max_length=32)
|
||||
|
||||
|
||||
@@ -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()
|
||||
@@ -11,10 +11,8 @@ from fastapi import FastAPI
|
||||
|
||||
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 .models import UserWalletSummary
|
||||
from .prototype import InMemoryExtensionAPI, InMemoryExtensionState
|
||||
from .runtime import ExtensionAPIHost
|
||||
|
||||
|
||||
@@ -27,13 +25,10 @@ async def invoke_wasm_extension_export(
|
||||
user: Any | None = None,
|
||||
) -> dict[str, Any]:
|
||||
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)
|
||||
api = InMemoryExtensionAPI(
|
||||
api = ExtensionAPI(
|
||||
extension.id,
|
||||
permissions,
|
||||
state=state,
|
||||
user_id=_user_id(user),
|
||||
)
|
||||
|
||||
@@ -54,7 +49,7 @@ def _invoke_wasm_extension_export_sync(
|
||||
extension: WasmExtension,
|
||||
export_name: str,
|
||||
payload: Mapping[str, Any],
|
||||
api: InMemoryExtensionAPI,
|
||||
api: ExtensionAPI,
|
||||
) -> dict[str, Any]:
|
||||
try:
|
||||
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)
|
||||
|
||||
|
||||
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]:
|
||||
installed_extension = await get_installed_extension(extension.id)
|
||||
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
|
||||
|
||||
|
||||
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:
|
||||
return re.sub(r"([a-z0-9])([A-Z])", r"\1-\2", value).replace("_", "-").lower()
|
||||
|
||||
Reference in New Issue
Block a user