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.core.db import core_app_extra
|
||||||
from lnbits.decorators import (
|
from lnbits.decorators import (
|
||||||
check_access_token,
|
check_access_token,
|
||||||
check_user_exists,
|
check_account_exists,
|
||||||
check_user_extension_access,
|
|
||||||
optional_user_id,
|
optional_user_id,
|
||||||
)
|
)
|
||||||
from lnbits.helpers import template_renderer
|
from lnbits.helpers import template_renderer
|
||||||
@@ -234,8 +233,6 @@ def _add_wasm_extension_api_route(
|
|||||||
if _has_route(app, route_path, method):
|
if _has_route(app, route_path, method):
|
||||||
return
|
return
|
||||||
|
|
||||||
require_user = _require_wasm_user_extension(extension.id)
|
|
||||||
|
|
||||||
async def invoke_wasm_api_request(
|
async def invoke_wasm_api_request(
|
||||||
request: Request, user: Any | None = None
|
request: Request, user: Any | None = None
|
||||||
) -> dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
@@ -258,9 +255,9 @@ def _add_wasm_extension_api_route(
|
|||||||
|
|
||||||
async def invoke_private_wasm_extension_export(
|
async def invoke_private_wasm_extension_export(
|
||||||
request: Request,
|
request: Request,
|
||||||
user: Any = Depends(require_user),
|
account: Any = Depends(check_account_exists),
|
||||||
) -> dict[str, Any]:
|
) -> 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]:
|
async def invoke_public_wasm_extension_export(request: Request) -> dict[str, Any]:
|
||||||
return await invoke_wasm_api_request(request)
|
return await invoke_wasm_api_request(request)
|
||||||
@@ -370,11 +367,9 @@ def _add_wasm_extension_wrapper_route(
|
|||||||
if _has_route(app, route_path, "GET"):
|
if _has_route(app, route_path, "GET"):
|
||||||
return
|
return
|
||||||
|
|
||||||
require_user = _require_wasm_user_extension(extension.id)
|
|
||||||
|
|
||||||
async def serve_private_wasm_extension_page(
|
async def serve_private_wasm_extension_page(
|
||||||
request: Request,
|
request: Request,
|
||||||
user: Any = Depends(require_user),
|
account: Any = Depends(check_account_exists),
|
||||||
) -> Any:
|
) -> Any:
|
||||||
return _wasm_extension_wrapper_response(
|
return _wasm_extension_wrapper_response(
|
||||||
request,
|
request,
|
||||||
@@ -382,8 +377,8 @@ def _add_wasm_extension_wrapper_route(
|
|||||||
frame_path,
|
frame_path,
|
||||||
auth,
|
auth,
|
||||||
path_params,
|
path_params,
|
||||||
user.json(),
|
None,
|
||||||
user.id,
|
account.id,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def serve_public_wasm_extension_page(
|
async def serve_public_wasm_extension_page(
|
||||||
@@ -603,18 +598,6 @@ def _path_template_pattern(path: str) -> str:
|
|||||||
return f"^{pattern}$"
|
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(
|
async def _optional_wasm_user_id(
|
||||||
request: Request,
|
request: Request,
|
||||||
access_token: Annotated[str | None, Depends(check_access_token)],
|
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(
|
async def _check_user_extension_access(
|
||||||
user_id: str, path: str, conn: Connection | None = None
|
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)
|
status = await check_user_extension_access(user_id, ext_id, conn=conn)
|
||||||
if not status.success:
|
if not status.success:
|
||||||
raise HTTPException(
|
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(
|
async def _get_account_from_token(
|
||||||
access_token: str, path: str, method: str, conn: Connection | None = None
|
access_token: str, path: str, method: str, conn: Connection | None = None
|
||||||
) -> Account | None:
|
) -> Account | None:
|
||||||
|
|||||||
Reference in New Issue
Block a user