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