diff --git a/lnbits/core/views/extension_api.py b/lnbits/core/views/extension_api.py index 8b68fed71..58c950311 100644 --- a/lnbits/core/views/extension_api.py +++ b/lnbits/core/views/extension_api.py @@ -705,8 +705,8 @@ async def extensions(account_id: AccountId = Depends(check_account_id_exists)): ), "kvSchema": _load_kv_schema_from_config(ext.id), "grantedPermissions": ( - user_exts_map.get(ext.id).extra.granted_permissions - if user_exts_map.get(ext.id) and user_exts_map.get(ext.id).extra + user_ext.extra.granted_permissions + if (user_ext := user_exts_map.get(ext.id)) and user_ext.extra else [] ), "latestRelease": ( diff --git a/lnbits/core/wasm/extension_host.py b/lnbits/core/wasm/extension_host.py index c7ce0dbea..d123a5518 100644 --- a/lnbits/core/wasm/extension_host.py +++ b/lnbits/core/wasm/extension_host.py @@ -2,8 +2,10 @@ from __future__ import annotations import asyncio import json +import re import time from pathlib import Path +from typing import Any import httpx from fastapi import APIRouter, Depends, HTTPException, Request @@ -37,7 +39,7 @@ def _ext_static_dir(ext_id: str, upgrade_hash: str | None = None) -> Path: def _ensure_kv_table(db: Database, ext_id: str) -> str: - table = f"{ext_id}.kv" if db.schema else "kv" + table = _kv_table_name(db, ext_id) query = f""" CREATE TABLE IF NOT EXISTS {table} ( key TEXT PRIMARY KEY, @@ -50,6 +52,16 @@ def _ensure_kv_table(db: Database, ext_id: str) -> str: _kv_schema_cache: dict[str, dict] = {} +def _kv_table_name(db: Database, ext_id: str) -> str: + table = f"{ext_id}.kv" if db.schema else "kv" + if ( + re.fullmatch(r"[a-zA-Z_][a-zA-Z0-9_]*(\.[a-zA-Z_][a-zA-Z0-9_]*)?", table) + is None + ): + raise ValueError("Invalid KV table name") + return table + + def _load_kv_schema(ext_id: str) -> dict: if ext_id in _kv_schema_cache: return _kv_schema_cache[ext_id] @@ -102,9 +114,9 @@ def _coerce_schema_value(schema_entry: dict, value): async def _kv_get(db: Database, ext_id: str, key: str) -> str | None: await db.execute(_ensure_kv_table(db, ext_id)) - table = f"{ext_id}.kv" if db.schema else "kv" - row = await db.fetchone( - f"SELECT value FROM {table} WHERE key = :key", + table = _kv_table_name(db, ext_id) + row: dict[str, Any] | None = await db.fetchone( + f"SELECT value FROM {table} WHERE key = :key", # noqa: S608 {"key": key}, ) if not row: @@ -125,19 +137,19 @@ async def _kv_get(db: Database, ext_id: str, key: str) -> str | None: async def _kv_set(db: Database, ext_id: str, key: str, value: str) -> None: await db.execute(_ensure_kv_table(db, ext_id)) - table = f"{ext_id}.kv" if db.schema else "kv" - existing = await db.fetchone( - f"SELECT key FROM {table} WHERE key = :key", + table = _kv_table_name(db, ext_id) + existing: dict[str, Any] | None = await db.fetchone( + f"SELECT key FROM {table} WHERE key = :key", # noqa: S608 {"key": key}, ) if existing: await db.execute( - f"UPDATE {table} SET value = :value WHERE key = :key", + f"UPDATE {table} SET value = :value WHERE key = :key", # noqa: S608 {"key": key, "value": value}, ) else: await db.execute( - f"INSERT INTO {table} (key, value) VALUES (:key, :value)", + f"INSERT INTO {table} (key, value) VALUES (:key, :value)", # noqa: S608 {"key": key, "value": value}, ) @@ -174,6 +186,12 @@ def _register_pages_routes(router: APIRouter, ext_id: str) -> None: def _register_kv_routes(router: APIRouter, ext_id: str, db: Database, ext) -> None: + _register_kv_read_routes(router, ext_id, db, ext) + _register_kv_write_routes(router, ext_id, db) + _register_kv_increment_route(router, ext_id, db, ext) + + +def _register_kv_read_routes(router: APIRouter, ext_id: str, db: Database, ext) -> None: @router.get("/api/v1/kv/{key}") async def api_kv_get(key: str, user: User = Depends(check_user_exists)): await _require_permission(user.id, ext_id, "ext.db.read_write") @@ -190,6 +208,8 @@ def _register_kv_routes(router: APIRouter, ext_id: str, db: Database, ext) -> No value = await _kv_get(db, ext_id, key) return {"key": key, "value": value} + +def _register_kv_write_routes(router: APIRouter, ext_id: str, db: Database) -> None: @router.post("/api/v1/kv/{key}") async def api_kv_set( key: str, payload: dict, user: User = Depends(check_user_exists) @@ -213,6 +233,10 @@ def _register_kv_routes(router: APIRouter, ext_id: str, db: Database, ext) -> No await websocket_updater(f"{ext_id}:{key}", str(value)) return {"key": key, "value": value} + +def _register_kv_increment_route( + router: APIRouter, ext_id: str, db: Database, ext +) -> None: @router.post("/api/v1/kv/{key}/increment") async def api_kv_increment(key: str, user: User = Depends(check_user_exists)): await _require_permission(user.id, ext_id, "ext.db.read_write") diff --git a/lnbits/core/wasm/runner.py b/lnbits/core/wasm/runner.py index 0c753fac2..4460454ee 100644 --- a/lnbits/core/wasm/runner.py +++ b/lnbits/core/wasm/runner.py @@ -1,7 +1,9 @@ import asyncio import json +import re import sys from pathlib import Path +from typing import cast from wasmtime import ( Caller, @@ -72,7 +74,7 @@ def _run(coro): def _ensure_kv_table(db: Database, ext_id: str) -> str: - table = f"{ext_id}.kv" if db.schema else "kv" + table = _kv_table_name(db, ext_id) return f""" CREATE TABLE IF NOT EXISTS {table} ( key TEXT PRIMARY KEY, @@ -81,8 +83,18 @@ def _ensure_kv_table(db: Database, ext_id: str) -> str: """ +def _kv_table_name(db: Database, ext_id: str) -> str: + table = f"{ext_id}.kv" if db.schema else "kv" + if ( + re.fullmatch(r"[a-zA-Z_][a-zA-Z0-9_]*(\.[a-zA-Z_][a-zA-Z0-9_]*)?", table) + is None + ): + raise RuntimeError("Invalid KV table name") + return table + + def _get_memory(caller: Caller): - memory = caller.get_export("memory") + memory = caller.get_export("memory") # type: ignore[attr-defined] if memory is None: raise RuntimeError("WASM module does not export memory") return memory @@ -109,10 +121,11 @@ def _db_get( ) -> int: key = _read_bytes(caller, key_ptr, key_len).decode(errors="ignore") _run(db.execute(_ensure_kv_table(db, ext_id))) - table = f"{ext_id}.kv" if db.schema else "kv" + table = _kv_table_name(db, ext_id) row = _run( db.fetchone( - f"SELECT value FROM {table} WHERE key = :key", {"key": key} + f"SELECT value FROM {table} WHERE key = :key", # noqa: S608 + {"key": key}, ) ) if not row: @@ -145,23 +158,24 @@ def _db_set( raise RuntimeError("Invalid value for schema") from exc value = json.dumps(coerced) if entry.get("type") == "json" else str(coerced) _run(db.execute(_ensure_kv_table(db, ext_id))) - table = f"{ext_id}.kv" if db.schema else "kv" + table = _kv_table_name(db, ext_id) row = _run( db.fetchone( - f"SELECT key FROM {table} WHERE key = :key", {"key": key} + f"SELECT key FROM {table} WHERE key = :key", # noqa: S608 + {"key": key}, ) ) if row: _run( db.execute( - f"UPDATE {table} SET value = :value WHERE key = :key", + f"UPDATE {table} SET value = :value WHERE key = :key", # noqa: S608 {"key": key, "value": value}, ) ) else: _run( db.execute( - f"INSERT INTO {table} (key, value) VALUES (:key, :value)", + f"INSERT INTO {table} (key, value) VALUES (:key, :value)", # noqa: S608 {"key": key, "value": value}, ) ) @@ -174,7 +188,7 @@ def _load_module(module_path: Path, ext_id: str): engine = Engine(config) store = Store(engine) if hasattr(store, "add_fuel"): - store.add_fuel(settings.lnbits_wasm_fuel) + store.add_fuel(settings.lnbits_wasm_fuel) # type: ignore[attr-defined] module = Module.from_file(engine, str(module_path)) db = Database(f"ext_{ext_id}") @@ -188,8 +202,15 @@ def _load_module(module_path: Path, ext_id: str): ) -> int: return _db_set(db, ext_id, caller, key_ptr, key_len, val_ptr, val_len) + def linker_define(linker: Linker, module: str, name: str, func: Func) -> None: + try: + linker.define(store, module, name, func) + except TypeError: + linker.define(module, name, func) # type: ignore[call-arg, arg-type] + linker = Linker(engine) - linker.define( + linker_define( + linker, "host", "db_get", Func( @@ -201,7 +222,8 @@ def _load_module(module_path: Path, ext_id: str): db_get, ), ) - linker.define( + linker_define( + linker, "host", "db_set", Func( @@ -236,9 +258,15 @@ def main() -> int: try: store, instance = _load_module(module_path, ext_id) - func = instance.exports(store)[function_name] + export = instance.exports(store)[function_name] + if not isinstance(export, Func): + raise RuntimeError(f"Export '{function_name}' is not callable") + func = cast(Func, export) int_args = [int(a) for a in args] - result = func(store, *int_args) + if hasattr(func, "call"): + result = func.call(store, *int_args) # type: ignore[attr-defined] + else: + result = func(store, *int_args) # type: ignore[operator] except Exception as exc: payload = {"ok": False, "error": str(exc)} sys.stdout.write(json.dumps(payload))