From dcb740a48da8c55601c2cf007c8a6132dcd50617 Mon Sep 17 00:00:00 2001 From: Vlad Stan Date: Tue, 30 Jun 2026 11:13:36 +0300 Subject: [PATCH] fix: owner id --- lnbits/core/extensions/api.py | 36 ++++++++++-- lnbits/core/extensions/storage.py | 92 +++++++++++++++++++++++++++++-- lnbits/core/extensions/wasm.py | 2 + lnbits/core/tasks.py | 31 +++++++++++ 4 files changed, 152 insertions(+), 9 deletions(-) diff --git a/lnbits/core/extensions/api.py b/lnbits/core/extensions/api.py index bbef2ac8c..cd794cf5f 100644 --- a/lnbits/core/extensions/api.py +++ b/lnbits/core/extensions/api.py @@ -12,6 +12,8 @@ from typing import Any, TypeVar, cast, get_type_hints from pydantic import BaseModel +from lnbits.helpers import sha256s + from .models import ( CreateInvoicePublicRequest, CreateInvoiceRequest, @@ -36,6 +38,7 @@ from .models import ( from .storage import ( storage_delete_row, storage_get_paginated_rows, + storage_get_public_row, storage_get_row, storage_set_row, ) @@ -129,11 +132,13 @@ class ExtensionAPI: *, user_id: str | None = None, context: str = "user", + owner_id: str | None = None, ) -> None: self.extension_id = extension_id self.permissions, self.permission_policies = self._permission_data(permissions) self.user_id = user_id self.context = context + self.owner_id = sha256s(user_id) if user_id else owner_id self._uuid = secrets.token_urlsafe(12).replace("-", "_") def __repr__(self) -> str: @@ -154,6 +159,11 @@ class ExtensionAPI: def has_authenticated_context(self) -> bool: return bool(self.user_id) or self.context == "event" + def _require_owner_id(self) -> str: + if not self.owner_id: + raise PermissionError("Extension API method requires an owner context.") + return self.owner_id + @extension_api_method( method_id="storage.get", namespace="storage", @@ -165,7 +175,12 @@ class ExtensionAPI: require_auth=True, ) async def storage_get(self, request: StorageGetRequest) -> StorageGetResponse: - row = await storage_get_row(self.extension_id, request.table, request.id) + row = await storage_get_row( + self.extension_id, + request.table, + request.id, + self._require_owner_id(), + ) return StorageGetResponse(data_json=json.dumps(row) if row else None) @extension_api_method( @@ -182,7 +197,7 @@ class ExtensionAPI: self, request: StorageGetRequest ) -> StorageGetResponse: public_fields = self._public_storage_fields(request.table) - row = await storage_get_row(self.extension_id, request.table, request.id) + row = await storage_get_public_row(self.extension_id, request.table, request.id) if not row: return StorageGetResponse() public_row = { @@ -203,7 +218,12 @@ class ExtensionAPI: require_auth=True, ) async def storage_set(self, request: StorageSetRequest) -> StorageSetResponse: - await storage_set_row(self.extension_id, request.table, request.data) + await storage_set_row( + self.extension_id, + request.table, + request.data, + self._require_owner_id(), + ) return StorageSetResponse() @extension_api_method( @@ -223,6 +243,7 @@ class ExtensionAPI: self.extension_id, request.table, request.filters, + owner_id=self._require_owner_id(), search=request.search, search_fields=request.search_fields, sort_by=request.sort_by, @@ -248,7 +269,12 @@ class ExtensionAPI: async def storage_delete( self, request: StorageDeleteRequest ) -> StorageDeleteResponse: - await storage_delete_row(self.extension_id, request.table, request.id) + await storage_delete_row( + self.extension_id, + request.table, + request.id, + self._require_owner_id(), + ) return StorageDeleteResponse() @extension_api_method( @@ -312,7 +338,7 @@ class ExtensionAPI: from lnbits.core.services.payments import create_payment_request table, wallet_field = self._public_invoice_wallet_source() - row = await storage_get_row(self.extension_id, table, request.source_id) + row = await storage_get_public_row(self.extension_id, table, request.source_id) if not row: raise PermissionError("Public invoice source was not found.") diff --git a/lnbits/core/extensions/storage.py b/lnbits/core/extensions/storage.py index 5bca76f6c..a15d54259 100644 --- a/lnbits/core/extensions/storage.py +++ b/lnbits/core/extensions/storage.py @@ -17,12 +17,29 @@ 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_]*$") +OWNER_ID_FIELD = "__lnbits_owner_id__" async def storage_get_row( ext_id: str, table: str, row_id: str, + owner_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 AND {OWNER_ID_FIELD} = :owner_id + """ # noqa: S608 + async with Database(f"ext_{ext_id}").connect() as conn: + row = await conn.fetchone(query, {"id": row_id, "owner_id": owner_id}) + return _row_from_db(table_schema, row) if row else None + + +async def storage_get_public_row( + ext_id: str, + table: str, + row_id: str, ) -> dict[str, Any] | None: table_schema = _load_table_schema(ext_id, table) query = f""" @@ -34,17 +51,46 @@ async def storage_get_row( return _row_from_db(table_schema, row) if row else None +async def storage_get_row_owner_id( + ext_id: str, + table: str, + row_id: str, +) -> str | None: + _load_table_schema(ext_id, table) + query = f""" + SELECT {OWNER_ID_FIELD} 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}) + + owner_id = row[OWNER_ID_FIELD] if row else None + return owner_id if isinstance(owner_id, str) and owner_id else None + + async def storage_set_row( ext_id: str, table: str, data: dict[str, Any], + owner_id: str, ) -> None: table_schema = _load_table_schema(ext_id, table) clean_data = _data_to_db(table_schema, data, require_id=True) + clean_data[OWNER_ID_FIELD] = owner_id 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" + updates = [ + f"{column} = excluded.{column}" + for column in columns + if column not in ("id", OWNER_ID_FIELD) + ] + conflict_sql = ( + "DO UPDATE SET " + + ", ".join(updates) + + f" WHERE {OWNER_ID_FIELD} = :{OWNER_ID_FIELD}" + if updates + else "DO NOTHING" + ) query = f""" INSERT INTO {_table_ref_for_schema(ext_id, table)} ({", ".join(columns)}) @@ -62,6 +108,7 @@ async def storage_get_paginated_rows( table: str, filters: dict[str, Any], *, + owner_id: str, search: str | None, search_fields: list[str], sort_by: str | None, @@ -71,6 +118,8 @@ async def storage_get_paginated_rows( ) -> dict[str, Any]: table_schema = _load_table_schema(ext_id, table) where_sql, values = _where_sql(table_schema, filters, search, search_fields) + where_sql = _append_owner_where_sql(where_sql) + values[OWNER_ID_FIELD] = owner_id order_sql = _order_sql(table_schema, sort_by, descending) count_values = dict(values) values.update({"limit": min(limit, 1000), "offset": offset}) @@ -102,14 +151,15 @@ async def storage_delete_row( ext_id: str, table: str, row_id: str, + owner_id: str, ) -> None: _load_table_schema(ext_id, table) query = f""" DELETE FROM {_table_ref_for_schema(ext_id, table)} - WHERE id = :id + WHERE id = :id AND {OWNER_ID_FIELD} = :owner_id """ # noqa: S608 async with Database(f"ext_{ext_id}").connect() as conn: - await conn.execute(query, {"id": row_id}) + await conn.execute(query, {"id": row_id, "owner_id": owner_id}) async def migrate_wasm_extension_database( @@ -175,11 +225,16 @@ def _create_table_sql(db: Connection, operation: dict[str, Any]) -> str: fields = _require_fields(operation) if not any(field.get("name") == "id" for field in fields): raise ValueError(f"WASM storage table '{table}' must define an id field.") + if any(field.get("name") == OWNER_ID_FIELD for field in fields): + raise ValueError( + f"WASM storage table '{table}' defines reserved field '{OWNER_ID_FIELD}'." + ) columns = [ _column_sql(db, field, primary_key=field.get("name") == "id") for field in fields ] + columns.append(f"{OWNER_ID_FIELD} TEXT NOT NULL") return f""" CREATE TABLE IF NOT EXISTS {_table_ref(db, table)} ( {", ".join(columns)} @@ -190,6 +245,11 @@ def _create_table_sql(db: Connection, operation: dict[str, Any]) -> str: def _add_field_sql(db: Connection, operation: dict[str, Any]) -> str: table = _require_identifier(operation, "table") field = _field_from_add_field_operation(operation) + if field["name"] == OWNER_ID_FIELD: + raise ValueError( + f"WASM storage table '{table}' cannot add reserved field " + f"'{OWNER_ID_FIELD}'." + ) return f""" ALTER TABLE {_table_ref(db, table)} ADD COLUMN {_column_sql(db, field)}; @@ -200,6 +260,11 @@ def _create_index_sql(db: Connection, operation: dict[str, Any]) -> str: table = _require_identifier(operation, "table") name = _require_identifier(operation, "name") field = _require_identifier(operation, "field") + if field == OWNER_ID_FIELD: + raise ValueError( + f"WASM storage table '{table}' cannot index reserved field " + f"'{OWNER_ID_FIELD}'." + ) if db.type == SQLITE and db.schema: return f""" @@ -271,6 +336,11 @@ def _load_table_schema(ext_id: str, table: str) -> dict[str, Any]: if not isinstance(field, dict): raise ValueError(f"WASM storage table '{table}' has invalid field schema.") _require_identifier(field, "name") + if field["name"] == OWNER_ID_FIELD: + raise ValueError( + f"WASM storage table '{table}' defines reserved field " + f"'{OWNER_ID_FIELD}'." + ) return table_schema @@ -297,6 +367,7 @@ def _data_to_db( 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.") + _reject_reserved_owner_field(data, "row") fields = _fields_by_name(table_schema) unknown_fields = sorted(set(data) - set(fields)) @@ -317,6 +388,7 @@ def _filters_to_db( ) -> dict[str, Any]: if not isinstance(filters, dict): raise ValueError("WASM storage filters must be an object.") + _reject_reserved_owner_field(filters, "filters") fields = _fields_by_name(table_schema) unknown_fields = sorted(set(filters) - set(fields)) @@ -359,6 +431,13 @@ def _where_sql( return ("WHERE " + " AND ".join(clauses), values) if clauses else ("", values) +def _append_owner_where_sql(where_sql: str) -> str: + owner_clause = f"{OWNER_ID_FIELD} = :{OWNER_ID_FIELD}" + if where_sql: + return f"{where_sql} AND {owner_clause}" + return f"WHERE {owner_clause}" + + def _order_sql( table_schema: dict[str, Any], sort_by: str | None, @@ -392,6 +471,11 @@ def _fields_by_name(table_schema: dict[str, Any]) -> dict[str, dict[str, Any]]: return {field["name"]: field for field in fields} +def _reject_reserved_owner_field(data: dict[str, Any], value_name: str) -> None: + if OWNER_ID_FIELD in data: + raise ValueError(f"WASM storage {value_name} includes a reserved owner field.") + + def _value_to_db(field: dict[str, Any], value: Any) -> Any: # noqa: C901 if value is None: if field.get("nullable", False): diff --git a/lnbits/core/extensions/wasm.py b/lnbits/core/extensions/wasm.py index d9ba0555c..7717c17b3 100644 --- a/lnbits/core/extensions/wasm.py +++ b/lnbits/core/extensions/wasm.py @@ -24,6 +24,7 @@ async def invoke_wasm_extension_export( *, user: Any | None = None, context: str = "user", + owner_id: str | None = None, ) -> dict[str, Any]: extension = _get_registered_extension(app, ext_id) permissions = await _extension_permissions(extension) @@ -32,6 +33,7 @@ async def invoke_wasm_extension_export( permissions, user_id=_user_id(user), context=context, + owner_id=owner_id, ) print("### api", api) diff --git a/lnbits/core/tasks.py b/lnbits/core/tasks.py index 9c46ca3e8..392a7e151 100644 --- a/lnbits/core/tasks.py +++ b/lnbits/core/tasks.py @@ -139,6 +139,7 @@ async def dispatch_wasm_invoice_paid(app: FastAPI, payment: Any) -> None: export_name, _wasm_invoice_paid_payload(payment), context="event", + owner_id=await _wasm_invoice_paid_owner_id(extension, payment), ) except Exception as exc: logger.warning( @@ -156,6 +157,36 @@ def _payment_extension_id(payment: Any) -> str | None: return tag if isinstance(tag, str) and tag else None +async def _wasm_invoice_paid_owner_id(extension: Any, payment: Any) -> str | None: + source_id = _payment_source_id(payment) + source_table = _wasm_public_invoice_source_table(extension.config) + if not source_id or not source_table: + return None + + from lnbits.core.extensions.storage import storage_get_row_owner_id + + return await storage_get_row_owner_id(extension.id, source_table, source_id) + + +def _payment_source_id(payment: Any) -> str | None: + extra = payment.extra or {} + source_id = extra.get("source_id") + return source_id if isinstance(source_id, str) and source_id else None + + +def _wasm_public_invoice_source_table(config: dict[str, Any]) -> str | None: + permissions = config.get("permissions") or [] + for permission in permissions: + if not isinstance(permission, dict): + continue + if permission.get("id") != "wallet.create_invoice_public": + continue + policy = permission.get("policy") or {} + table = policy.get("table") + return table if isinstance(table, str) and table else None + return None + + def _wasm_invoice_paid_export(config: dict[str, Any]) -> str | None: events = config.get("events") or {} export_name = events.get("onInvoicePaid")