feat: store data
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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, {})
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user