fix: simplify user check
This commit is contained in:
@@ -19,8 +19,7 @@ from starlette.types import Scope
|
||||
from lnbits.core.db import core_app_extra
|
||||
from lnbits.decorators import (
|
||||
check_access_token,
|
||||
check_user_exists,
|
||||
check_user_extension_access,
|
||||
check_account_exists,
|
||||
optional_user_id,
|
||||
)
|
||||
from lnbits.helpers import template_renderer
|
||||
@@ -234,8 +233,6 @@ def _add_wasm_extension_api_route(
|
||||
if _has_route(app, route_path, method):
|
||||
return
|
||||
|
||||
require_user = _require_wasm_user_extension(extension.id)
|
||||
|
||||
async def invoke_wasm_api_request(
|
||||
request: Request, user: Any | None = None
|
||||
) -> dict[str, Any]:
|
||||
@@ -258,9 +255,9 @@ def _add_wasm_extension_api_route(
|
||||
|
||||
async def invoke_private_wasm_extension_export(
|
||||
request: Request,
|
||||
user: Any = Depends(require_user),
|
||||
account: Any = Depends(check_account_exists),
|
||||
) -> dict[str, Any]:
|
||||
return await invoke_wasm_api_request(request, user)
|
||||
return await invoke_wasm_api_request(request, account)
|
||||
|
||||
async def invoke_public_wasm_extension_export(request: Request) -> dict[str, Any]:
|
||||
return await invoke_wasm_api_request(request)
|
||||
@@ -370,11 +367,9 @@ def _add_wasm_extension_wrapper_route(
|
||||
if _has_route(app, route_path, "GET"):
|
||||
return
|
||||
|
||||
require_user = _require_wasm_user_extension(extension.id)
|
||||
|
||||
async def serve_private_wasm_extension_page(
|
||||
request: Request,
|
||||
user: Any = Depends(require_user),
|
||||
account: Any = Depends(check_account_exists),
|
||||
) -> Any:
|
||||
return _wasm_extension_wrapper_response(
|
||||
request,
|
||||
@@ -382,8 +377,8 @@ def _add_wasm_extension_wrapper_route(
|
||||
frame_path,
|
||||
auth,
|
||||
path_params,
|
||||
user.json(),
|
||||
user.id,
|
||||
None,
|
||||
account.id,
|
||||
)
|
||||
|
||||
async def serve_public_wasm_extension_page(
|
||||
@@ -603,18 +598,6 @@ def _path_template_pattern(path: str) -> str:
|
||||
return f"^{pattern}$"
|
||||
|
||||
|
||||
def _require_wasm_user_extension(ext_id: str) -> Any:
|
||||
async def require_wasm_user_extension(
|
||||
user: Any = Depends(check_user_exists),
|
||||
) -> Any:
|
||||
status = await check_user_extension_access(user.id, ext_id)
|
||||
if not status.success:
|
||||
raise HTTPException(status_code=403, detail=status.message)
|
||||
return user
|
||||
|
||||
return require_wasm_user_extension
|
||||
|
||||
|
||||
async def _optional_wasm_user_id(
|
||||
request: Request,
|
||||
access_token: Annotated[str | None, Depends(check_access_token)],
|
||||
|
||||
+10
-1
@@ -448,7 +448,7 @@ async def _check_user_access(r: Request, user_id: str, conn: Connection | None =
|
||||
async def _check_user_extension_access(
|
||||
user_id: str, path: str, conn: Connection | None = None
|
||||
):
|
||||
ext_id = path_segments(path)[0]
|
||||
ext_id = _extension_id_from_request_path(path)
|
||||
status = await check_user_extension_access(user_id, ext_id, conn=conn)
|
||||
if not status.success:
|
||||
raise HTTPException(
|
||||
@@ -457,6 +457,15 @@ async def _check_user_extension_access(
|
||||
)
|
||||
|
||||
|
||||
def _extension_id_from_request_path(path: str) -> str:
|
||||
segments = path_segments(path)
|
||||
if len(segments) >= 2 and segments[0] == "ext":
|
||||
return segments[1]
|
||||
if len(segments) >= 4 and segments[:3] == ["api", "v1", "ext"]:
|
||||
return segments[3]
|
||||
return segments[0]
|
||||
|
||||
|
||||
async def _get_account_from_token(
|
||||
access_token: str, path: str, method: str, conn: Connection | None = None
|
||||
) -> Account | None:
|
||||
|
||||
Reference in New Issue
Block a user