feat: store data

This commit is contained in:
Vlad Stan
2026-07-08 11:54:57 +03:00
parent 5fe4752ebf
commit c4ee44ed4b
4 changed files with 369 additions and 54 deletions
+37 -21
View File
@@ -7,7 +7,7 @@ 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 TypeVar, get_type_hints from typing import NoReturn, TypeVar, cast, get_type_hints
from pydantic import BaseModel from pydantic import BaseModel
@@ -15,18 +15,20 @@ from .models import (
CreateInvoiceRequest, CreateInvoiceRequest,
CreateInvoiceResponse, CreateInvoiceResponse,
EmptyRequest, EmptyRequest,
KvGetRequest,
KvGetResponse,
KvListRequest,
KvListResponse,
KvSetRequest,
KvSetResponse,
ListUserWalletsResponse, ListUserWalletsResponse,
LogRequest, LogRequest,
LogResponse, LogResponse,
NowResponse, NowResponse,
RandomIdRequest, RandomIdRequest,
RandomIdResponse, RandomIdResponse,
StorageDeleteRequest,
StorageDeleteResponse,
StorageGetRequest,
StorageGetResponse,
StorageListRequest,
StorageListResponse,
StorageSetRequest,
StorageSetResponse,
WatchPaymentRequest, WatchPaymentRequest,
WatchPaymentResponse, WatchPaymentResponse,
) )
@@ -127,39 +129,53 @@ class ExtensionAPI:
@extension_api_method( @extension_api_method(
method_id="storage.get", method_id="storage.get",
namespace="storage", namespace="storage",
name="Get storage value", name="Get storage row",
host_name="kv_get", host_name="storage_get",
sdk_name="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", 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") self._raise_unwired_runtime("storage_get")
@extension_api_method( @extension_api_method(
method_id="storage.set", method_id="storage.set",
namespace="storage", namespace="storage",
name="Set storage value", name="Set storage row",
host_name="kv_set", host_name="storage_set",
sdk_name="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", 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") self._raise_unwired_runtime("storage_set")
@extension_api_method( @extension_api_method(
method_id="storage.list", method_id="storage.list",
namespace="storage", namespace="storage",
name="List storage keys", name="List storage rows",
host_name="kv_list", host_name="storage_list",
sdk_name="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", 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") 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( @extension_api_method(
method_id="wallet.create_invoice", method_id="wallet.create_invoice",
namespace="wallet", namespace="wallet",
@@ -239,7 +255,7 @@ 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) -> None: def _raise_unwired_runtime(self, method_name: str) -> NoReturn:
raise NotImplementedError( raise NotImplementedError(
f"ExtensionAPI.{method_name} must be wired to LNbits services before use." 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." 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: def _is_pydantic_model(value: object) -> bool:
+42 -14
View File
@@ -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): class EmptyRequest(BaseModel):
pass pass
class KvGetRequest(BaseModel): class StorageGetRequest(BaseModel):
key: str = Field(..., min_length=1, max_length=512) table: str = Field(..., min_length=1, max_length=128)
id: str = Field(..., min_length=1, max_length=512)
class KvGetResponse(BaseModel): class StorageGetResponse(BaseModel):
value: str | None = None data_json: str | None = None
class KvSetRequest(BaseModel): class StorageSetRequest(BaseModel):
key: str = Field(..., min_length=1, max_length=512) table: str = Field(..., min_length=1, max_length=128)
value: str = Field(..., max_length=65536) 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 ok: bool = True
class KvListRequest(BaseModel): class StorageListRequest(BaseModel):
prefix: str = Field(..., min_length=1, max_length=512) 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): class StorageListResponse(BaseModel):
keys: list[str] = Field(default_factory=list) 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): class CreateInvoiceRequest(BaseModel):
+37 -19
View File
@@ -1,5 +1,6 @@
from __future__ import annotations from __future__ import annotations
import json
from dataclasses import dataclass, field from dataclasses import dataclass, field
from .api import ExtensionAPI from .api import ExtensionAPI
@@ -7,22 +8,29 @@ from .models import (
CreateInvoiceRequest, CreateInvoiceRequest,
CreateInvoiceResponse, CreateInvoiceResponse,
EmptyRequest, EmptyRequest,
KvGetRequest,
KvGetResponse,
KvListRequest,
KvListResponse,
KvSetRequest,
KvSetResponse,
ListUserWalletsResponse, ListUserWalletsResponse,
StorageDeleteRequest,
StorageDeleteResponse,
StorageGetRequest,
StorageGetResponse,
StorageListRequest,
StorageListResponse,
StorageSetRequest,
StorageSetResponse,
UserWalletSummary, UserWalletSummary,
WatchPaymentRequest, WatchPaymentRequest,
WatchPaymentResponse, WatchPaymentResponse,
) )
from .storage import (
storage_delete_row,
storage_get_row,
storage_list_rows,
storage_set_row,
)
@dataclass @dataclass
class InMemoryExtensionState: class InMemoryExtensionState:
storage: dict[str, dict[str, str]] = field(default_factory=dict)
payment_watchers: 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( user_wallets: dict[str, list[UserWalletSummary] | None] = field(
default_factory=dict default_factory=dict
@@ -41,19 +49,33 @@ class InMemoryExtensionAPI(ExtensionAPI):
super().__init__(extension_id, permissions, user_id=user_id) super().__init__(extension_id, permissions, user_id=user_id)
self.state = state or InMemoryExtensionState() 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") 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.require_permission("ext.storage.read_write")
self._storage[request.key] = request.value await storage_set_row(self.extension_id, request.table, request.data)
return KvSetResponse() 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") self.require_permission("ext.storage.read_write")
keys = sorted(key for key in self._storage if key.startswith(request.prefix)) rows = await storage_list_rows(
return KvListResponse(keys=keys) 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( async def wallet_create_invoice(
self, request: CreateInvoiceRequest self, request: CreateInvoiceRequest
@@ -121,7 +143,3 @@ class InMemoryExtensionAPI(ExtensionAPI):
request.payment_hash request.payment_hash
] = request.callback_export ] = request.callback_export
return WatchPaymentResponse() return WatchPaymentResponse()
@property
def _storage(self) -> dict[str, str]:
return self.state.storage.setdefault(self.extension_id, {})
+253
View File
@@ -2,6 +2,7 @@ from __future__ import annotations
import json import json
import re import re
from datetime import datetime, timezone
from pathlib import Path from pathlib import Path
from typing import Any 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 import DbVersion
from lnbits.core.models.extensions import InstallableExtension from lnbits.core.models.extensions import InstallableExtension
from lnbits.db import POSTGRES, SQLITE, Connection, Database from lnbits.db import POSTGRES, SQLITE, Connection, Database
from lnbits.settings import settings
_MIGRATION_FILE_RE = re.compile(r"^(\d+)_.*\.json$") _MIGRATION_FILE_RE = re.compile(r"^(\d+)_.*\.json$")
_SQL_IDENTIFIER_RE = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") _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( async def migrate_wasm_extension_database(
ext: InstallableExtension, ext: InstallableExtension,
current_version: DbVersion | None = None, 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}") 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: def _default_sql(value: Any) -> str:
if value is None: if value is None:
return "NULL" return "NULL"
@@ -204,6 +451,12 @@ def _table_ref(db: Connection, table: str) -> str:
return 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: def _schema_ref(db: Connection, name: str) -> str:
if not db.schema: if not db.schema:
return name return name