diff --git a/lnbits/core/extensions/api.py b/lnbits/core/extensions/api.py index e07674947..b6f708528 100644 --- a/lnbits/core/extensions/api.py +++ b/lnbits/core/extensions/api.py @@ -7,7 +7,7 @@ import time from collections.abc import Awaitable, Callable, Iterable from dataclasses import dataclass from functools import wraps -from typing import TypeVar, get_type_hints +from typing import NoReturn, TypeVar, cast, get_type_hints from pydantic import BaseModel @@ -15,18 +15,20 @@ from .models import ( CreateInvoiceRequest, CreateInvoiceResponse, EmptyRequest, - KvGetRequest, - KvGetResponse, - KvListRequest, - KvListResponse, - KvSetRequest, - KvSetResponse, ListUserWalletsResponse, LogRequest, LogResponse, NowResponse, RandomIdRequest, RandomIdResponse, + StorageDeleteRequest, + StorageDeleteResponse, + StorageGetRequest, + StorageGetResponse, + StorageListRequest, + StorageListResponse, + StorageSetRequest, + StorageSetResponse, WatchPaymentRequest, WatchPaymentResponse, ) @@ -127,39 +129,53 @@ class ExtensionAPI: @extension_api_method( method_id="storage.get", namespace="storage", - name="Get storage value", - host_name="kv_get", + name="Get storage row", + host_name="storage_get", sdk_name="get", - description="Read one value from the extension storage namespace.", + description="Read one row from an extension storage table.", required_permission="ext.storage.read_write", ) - async def storage_get(self, request: KvGetRequest) -> KvGetResponse: + async def storage_get(self, request: StorageGetRequest) -> StorageGetResponse: self._raise_unwired_runtime("storage_get") @extension_api_method( method_id="storage.set", namespace="storage", - name="Set storage value", - host_name="kv_set", + name="Set storage row", + host_name="storage_set", sdk_name="set", - description="Write one value to the extension storage namespace.", + description="Create or update one row in an extension storage table.", required_permission="ext.storage.read_write", ) - async def storage_set(self, request: KvSetRequest) -> KvSetResponse: + async def storage_set(self, request: StorageSetRequest) -> StorageSetResponse: self._raise_unwired_runtime("storage_set") @extension_api_method( method_id="storage.list", namespace="storage", - name="List storage keys", - host_name="kv_list", + name="List storage rows", + host_name="storage_list", sdk_name="list", - description="List keys under a prefix in the extension storage namespace.", + description="List rows from an extension storage table.", required_permission="ext.storage.read_write", ) - async def storage_list(self, request: KvListRequest) -> KvListResponse: + async def storage_list(self, request: StorageListRequest) -> StorageListResponse: self._raise_unwired_runtime("storage_list") + @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.read_write", + ) + async def storage_delete( + self, request: StorageDeleteRequest + ) -> StorageDeleteResponse: + self._raise_unwired_runtime("storage_delete") + @extension_api_method( method_id="wallet.create_invoice", namespace="wallet", @@ -239,7 +255,7 @@ class ExtensionAPI: log("extension:%s %s", self.extension_id, request.message) return LogResponse() - def _raise_unwired_runtime(self, method_name: str) -> None: + def _raise_unwired_runtime(self, method_name: str) -> NoReturn: raise NotImplementedError( f"ExtensionAPI.{method_name} must be wired to LNbits services before use." ) @@ -339,7 +355,7 @@ def _get_method_models( f"Extension API method '{function.__name__}' response must be a BaseModel." ) - return request_model, response_model + return cast(type[BaseModel], request_model), cast(type[BaseModel], response_model) def _is_pydantic_model(value: object) -> bool: diff --git a/lnbits/core/extensions/models.py b/lnbits/core/extensions/models.py index 3ba9f28cd..d6cd24ad9 100644 --- a/lnbits/core/extensions/models.py +++ b/lnbits/core/extensions/models.py @@ -1,35 +1,63 @@ -from typing import Literal +import json +from typing import Any, Literal -from pydantic import BaseModel, Field +from pydantic import BaseModel, Field, root_validator class EmptyRequest(BaseModel): pass -class KvGetRequest(BaseModel): - key: str = Field(..., min_length=1, max_length=512) +class StorageGetRequest(BaseModel): + table: str = Field(..., min_length=1, max_length=128) + id: str = Field(..., min_length=1, max_length=512) -class KvGetResponse(BaseModel): - value: str | None = None +class StorageGetResponse(BaseModel): + data_json: str | None = None -class KvSetRequest(BaseModel): - key: str = Field(..., min_length=1, max_length=512) - value: str = Field(..., max_length=65536) +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 KvSetResponse(BaseModel): +class StorageSetResponse(BaseModel): ok: bool = True -class KvListRequest(BaseModel): - prefix: str = Field(..., min_length=1, max_length=512) +class StorageListRequest(BaseModel): + table: str = Field(..., min_length=1, max_length=128) + filters: dict[str, Any] = Field(default_factory=dict) + limit: int = Field(100, ge=1, le=1000) + offset: int = Field(0, ge=0) + + @root_validator(pre=True) + def parse_filters_json(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) + return values -class KvListResponse(BaseModel): - keys: list[str] = Field(default_factory=list) +class StorageListResponse(BaseModel): + rows_json: str = "[]" + + +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): diff --git a/lnbits/core/extensions/prototype.py b/lnbits/core/extensions/prototype.py index db1543e14..a970452ea 100644 --- a/lnbits/core/extensions/prototype.py +++ b/lnbits/core/extensions/prototype.py @@ -1,5 +1,6 @@ from __future__ import annotations +import json from dataclasses import dataclass, field from .api import ExtensionAPI @@ -7,22 +8,29 @@ from .models import ( CreateInvoiceRequest, CreateInvoiceResponse, EmptyRequest, - KvGetRequest, - KvGetResponse, - KvListRequest, - KvListResponse, - KvSetRequest, - KvSetResponse, ListUserWalletsResponse, + StorageDeleteRequest, + StorageDeleteResponse, + StorageGetRequest, + StorageGetResponse, + StorageListRequest, + StorageListResponse, + StorageSetRequest, + StorageSetResponse, UserWalletSummary, WatchPaymentRequest, WatchPaymentResponse, ) +from .storage import ( + storage_delete_row, + storage_get_row, + storage_list_rows, + storage_set_row, +) @dataclass class InMemoryExtensionState: - storage: dict[str, dict[str, str]] = field(default_factory=dict) payment_watchers: dict[str, dict[str, str]] = field(default_factory=dict) user_wallets: dict[str, list[UserWalletSummary] | None] = field( default_factory=dict @@ -41,19 +49,33 @@ class InMemoryExtensionAPI(ExtensionAPI): super().__init__(extension_id, permissions, user_id=user_id) self.state = state or InMemoryExtensionState() - async def storage_get(self, request: KvGetRequest) -> KvGetResponse: + async def storage_get(self, request: StorageGetRequest) -> StorageGetResponse: self.require_permission("ext.storage.read_write") - return KvGetResponse(value=self._storage.get(request.key)) + 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: KvSetRequest) -> KvSetResponse: + async def storage_set(self, request: StorageSetRequest) -> StorageSetResponse: self.require_permission("ext.storage.read_write") - self._storage[request.key] = request.value - return KvSetResponse() + await storage_set_row(self.extension_id, request.table, request.data) + return StorageSetResponse() - async def storage_list(self, request: KvListRequest) -> KvListResponse: + async def storage_list(self, request: StorageListRequest) -> StorageListResponse: self.require_permission("ext.storage.read_write") - keys = sorted(key for key in self._storage if key.startswith(request.prefix)) - return KvListResponse(keys=keys) + rows = await storage_list_rows( + self.extension_id, + request.table, + request.filters, + limit=request.limit, + offset=request.offset, + ) + return StorageListResponse(rows_json=json.dumps(rows)) + + 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 @@ -121,7 +143,3 @@ class InMemoryExtensionAPI(ExtensionAPI): request.payment_hash ] = request.callback_export return WatchPaymentResponse() - - @property - def _storage(self) -> dict[str, str]: - return self.state.storage.setdefault(self.extension_id, {}) diff --git a/lnbits/core/extensions/storage.py b/lnbits/core/extensions/storage.py index 833e394a6..5b17d9327 100644 --- a/lnbits/core/extensions/storage.py +++ b/lnbits/core/extensions/storage.py @@ -2,6 +2,7 @@ from __future__ import annotations import json import re +from datetime import datetime, timezone from pathlib import Path from typing import Any @@ -12,11 +13,93 @@ 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_]*$") +async def storage_get_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_set_row( + ext_id: str, + table: str, + data: dict[str, Any], +) -> None: + table_schema = _load_table_schema(ext_id, table) + clean_data = _data_to_db(table_schema, data, require_id=True) + columns = list(clean_data.keys()) + placeholders = [f":{column}" for column in columns] + updates = [f"{column} = excluded.{column}" for column in columns if column != "id"] + conflict_sql = f"DO UPDATE SET {', '.join(updates)}" 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_list_rows( + ext_id: str, + table: str, + filters: dict[str, Any], + *, + limit: int, + offset: int, +) -> list[dict[str, Any]]: + table_schema = _load_table_schema(ext_id, table) + clean_filters = _filters_to_db(table_schema, filters) + where_sql = "" + if clean_filters: + clauses = [f"{field} = :filter_{field}" for field in clean_filters] + where_sql = "WHERE " + " AND ".join(clauses) + + values = {f"filter_{field}": value for field, value in clean_filters.items()} + values.update({"limit": min(limit, 1000), "offset": offset}) + query = f""" + SELECT * FROM {_table_ref_for_schema(ext_id, table)} + {where_sql} + LIMIT :limit + OFFSET :offset + """ # noqa: S608 + + async with Database(f"ext_{ext_id}").connect() as conn: + rows = await conn.fetchall(query, values) + return [_row_from_db(table_schema, row) for row in rows] + + +async def storage_delete_row( + ext_id: str, + table: str, + row_id: str, +) -> None: + _load_table_schema(ext_id, table) + query = f""" + DELETE FROM {_table_ref_for_schema(ext_id, table)} + WHERE id = :id + """ # noqa: S608 + async with Database(f"ext_{ext_id}").connect() as conn: + await conn.execute(query, {"id": row_id}) + + async def migrate_wasm_extension_database( ext: InstallableExtension, current_version: DbVersion | None = None, @@ -157,6 +240,170 @@ def _field_type_sql(db: Connection, field: dict[str, Any]) -> str: 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") + 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.") + + 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.") + + 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 _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 _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" @@ -204,6 +451,12 @@ def _table_ref(db: Connection, table: str) -> str: 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