feat: store data

This commit is contained in:
Vlad Stan
2026-07-09 16:30:18 +03:00
parent da76d288ae
commit 3d4202b495
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 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:
+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):
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):
+37 -19
View File
@@ -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, {})
+253
View File
@@ -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