fix: owner id

This commit is contained in:
Vlad Stan
2026-07-09 16:33:39 +03:00
parent a2a12a25a5
commit 0fce2481cb
3 changed files with 121 additions and 9 deletions
+31 -5
View File
@@ -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.")
+88 -4
View File
@@ -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):
+2
View File
@@ -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)