diff --git a/lnbits/core/wasm_ext/api/host.py b/lnbits/core/wasm_ext/api/host.py index cae8e55da..30b7a934b 100644 --- a/lnbits/core/wasm_ext/api/host.py +++ b/lnbits/core/wasm_ext/api/host.py @@ -52,6 +52,7 @@ from .models import ( StorageGetResponse, StoragePaginatedRequest, StoragePaginatedResponse, + StoragePublicPaginatedRequest, StorageSetRequest, StorageSetResponse, UserWalletSummary, @@ -128,7 +129,7 @@ class ExtensionHostAPI: async def storage_get_public( self, request: StorageGetRequest ) -> StorageGetResponse: - public_fields = self._public_storage_fields(request.table) + public_fields = self._public_storage_policy(request.table)["public_fields"] row = await storage_get_public_row(self.extension_id, request.table, request.id) if not row: return StorageGetResponse() @@ -239,14 +240,23 @@ class ExtensionHostAPI: require_auth=False, ) async def storage_get_public_paginated( - self, request: StoragePaginatedRequest + self, request: StoragePublicPaginatedRequest ) -> StoragePaginatedResponse: - public_fields = self._public_storage_fields(request.table) - self._validate_public_storage_query_fields(request, public_fields) + policy = self._public_storage_policy(request.table) + public_fields = policy["public_fields"] + source_id_field = policy["source_id_field"] + if not isinstance(source_id_field, str) or not source_id_field: + raise PermissionError( + "Public paginated storage reads require a source ID field policy." + ) + + filters = self._public_storage_paginated_filters( + request, public_fields, source_id_field + ) page = await storage_get_public_paginated_rows( self.extension_id, request.table, - request.filters, + filters, search=request.search, search_fields=request.search_fields, sort_by=request.sort_by, @@ -737,7 +747,7 @@ class ExtensionHostAPI: return permission_ids, policies - def _public_storage_fields(self, table: str) -> set[str]: + def _public_storage_policy(self, table: str) -> dict[str, Any]: tables = self.permission_policies.get("ext.storage.read_public") if not isinstance(tables, list) or not tables: raise PermissionError( @@ -757,24 +767,56 @@ class ExtensionHostAPI: raise PermissionError( f"Public storage table '{table}' has no valid public fields." ) - return set(public_fields) + source_id_field = table_policy.get("source_id_field") + if source_id_field is not None and ( + not isinstance(source_id_field, str) or not source_id_field + ): + raise PermissionError( + f"Public storage table '{table}' has no valid source ID field." + ) + return { + "public_fields": set(public_fields), + "source_id_field": source_id_field, + } raise PermissionError(f"Storage table '{table}' is not publicly readable.") def _validate_public_storage_query_fields( - self, request: StoragePaginatedRequest, public_fields: set[str] + self, + request: StoragePaginatedRequest, + public_fields: set[str], + allowed_private_fields: set[str] | None = None, ) -> None: + allowed_private_fields = allowed_private_fields or set() query_fields = set(request.filters) query_fields.update(request.search_fields) if request.sort_by: query_fields.add(request.sort_by) - private_fields = sorted(query_fields - public_fields) + private_fields = sorted(query_fields - public_fields - allowed_private_fields) if private_fields: raise PermissionError( "Public storage query uses non-public fields: " + ", ".join(private_fields) ) + def _public_storage_paginated_filters( + self, + request: StoragePublicPaginatedRequest, + public_fields: set[str], + source_id_field: str, + ) -> dict[str, Any]: + self._validate_public_storage_query_fields( + request, public_fields, {source_id_field} + ) + filters = dict(request.filters) + requested_source_id = filters.get(source_id_field) + if requested_source_id is not None and requested_source_id != request.source_id: + raise PermissionError( + "Public storage source filter does not match source_id." + ) + filters[source_id_field] = request.source_id + return filters + async def _public_storage_append_policy( self, table: str, source_id: str ) -> tuple[dict[str, Any], str]: diff --git a/lnbits/core/wasm_ext/api/models.py b/lnbits/core/wasm_ext/api/models.py index b44b7b54c..de518af11 100644 --- a/lnbits/core/wasm_ext/api/models.py +++ b/lnbits/core/wasm_ext/api/models.py @@ -109,6 +109,10 @@ class StoragePaginatedRequest(BaseModel): return values +class StoragePublicPaginatedRequest(StoragePaginatedRequest): + source_id: str = Field(..., min_length=1, max_length=512) + + class StoragePaginatedResponse(BaseModel): rows_json: str = "[]" total: int = 0 diff --git a/lnbits/core/wasm_ext/api/permissions.py b/lnbits/core/wasm_ext/api/permissions.py index a192e34b5..e8a30d955 100644 --- a/lnbits/core/wasm_ext/api/permissions.py +++ b/lnbits/core/wasm_ext/api/permissions.py @@ -210,29 +210,47 @@ def _public_storage_grant_is_subset( ) -> bool: requested_tables = _public_storage_tables(requested_policies) granted_tables = _public_storage_tables(granted_policies) - for table_name, granted_fields in granted_tables.items(): - requested_fields = requested_tables.get(table_name) - if requested_fields is None or not granted_fields.issubset(requested_fields): + for table_name, granted_policy in granted_tables.items(): + requested_policy = requested_tables.get(table_name) + if requested_policy is None: + return False + if not granted_policy["public_fields"].issubset( + requested_policy["public_fields"] + ): + return False + requested_source_id_field = requested_policy["source_id_field"] + if ( + requested_source_id_field + and granted_policy["source_id_field"] != requested_source_id_field + ): return False return True -def _public_storage_tables(policies: list[Any] | None) -> dict[str, set[str]]: - tables: dict[str, set[str]] = {} +def _public_storage_tables(policies: list[Any] | None) -> dict[str, dict[str, Any]]: + tables: dict[str, dict[str, Any]] = {} for policy in _policy_list(policies): if not isinstance(policy, dict): continue table_name = policy.get("table_name") public_fields = policy.get("public_fields") + source_id_field = policy.get("source_id_field") if ( not isinstance(table_name, str) or table_name in tables or not isinstance(public_fields, list) + or ( + source_id_field is not None + and (not isinstance(source_id_field, str) or not source_id_field) + ) ): continue fields = {field for field in public_fields if isinstance(field, str) and field} if fields: - tables[table_name] = fields + tables[table_name] = { + "public_fields": fields, + "source_id_field": source_id_field, + } return tables diff --git a/lnbits/static/i18n/en.js b/lnbits/static/i18n/en.js index 74cf7bab6..b43a1442b 100644 --- a/lnbits/static/i18n/en.js +++ b/lnbits/static/i18n/en.js @@ -549,6 +549,8 @@ window.localisation.en = { extension_permission_ext_storage_append_public_sources: 'Allowed append targets', extension_permission_ext_storage_read_public: 'Read public extension storage', + extension_permission_ext_storage_read_public_source_required: + 'required to read', extension_permission_ext_storage_write: 'Write extension storage', extension_permission_ext_storage_read_write: 'Read & Write extension storage', extension_permission_extension_api_request: 'Use other extensions', diff --git a/lnbits/static/js/components/lnbits-extension-permissions.js b/lnbits/static/js/components/lnbits-extension-permissions.js index a0747fe8f..fbfb8400e 100644 --- a/lnbits/static/js/components/lnbits-extension-permissions.js +++ b/lnbits/static/js/components/lnbits-extension-permissions.js @@ -174,7 +174,12 @@ : table.public_fields.filter( field => typeof field === 'string' && field ) - return tableName ? {table: tableName, fields} : null + const sourceIdField = + typeof table === 'string' || + typeof table?.source_id_field !== 'string' + ? '' + : table.source_id_field + return tableName ? {table: tableName, fields, sourceIdField} : null }) .filter(Boolean) } diff --git a/lnbits/templates/components/lnbits-extension-permissions.vue b/lnbits/templates/components/lnbits-extension-permissions.vue index 2515a5509..3b79e4468 100644 --- a/lnbits/templates/components/lnbits-extension-permissions.vue +++ b/lnbits/templates/components/lnbits-extension-permissions.vue @@ -9,28 +9,31 @@ > @@ -57,7 +60,24 @@ >