diff --git a/lnbits/core/__init__.py b/lnbits/core/__init__.py index 5b70a1078..ab82314e2 100644 --- a/lnbits/core/__init__.py +++ b/lnbits/core/__init__.py @@ -8,7 +8,6 @@ from .views.audit_api import audit_router from .views.auth_api import auth_router from .views.callback_api import callback_router from .views.extension_api import extension_router -from .views.extension_websocket_api import extension_websocket_router from .views.extensions_builder_api import extension_builder_router from .views.fiat_api import fiat_router @@ -21,7 +20,7 @@ from .views.tinyurl_api import tinyurl_router from .views.user_api import users_router from .views.wallet_api import wallet_router from .views.webpush_api import webpush_router -from .views.websocket_api import websocket_router +from .views.websocket_api import extension_websocket_router, websocket_router # backwards compatibility for extensions core_app = APIRouter(tags=["Core"]) diff --git a/lnbits/core/views/extension_websocket_api.py b/lnbits/core/views/extension_websocket_api.py deleted file mode 100644 index afa6f01e1..000000000 --- a/lnbits/core/views/extension_websocket_api.py +++ /dev/null @@ -1,38 +0,0 @@ -from fastapi import APIRouter, WebSocket, status - -from lnbits.core.crud import get_installed_extension -from lnbits.core.wasm_ext.api.websockets import wasm_extension_websocket_hub - -extension_websocket_router = APIRouter( - prefix="/api/v1/ext/ws", - tags=["Extension Websocket"], -) - - -@extension_websocket_router.websocket("/{ext_id}/{item_id}") -async def extension_websocket_connect( - websocket: WebSocket, - ext_id: str, - item_id: str, -) -> None: - installed_ext = await get_installed_extension(ext_id) - installed_permission_ids = ( - {permission.id for permission in installed_ext.permissions or []} - if installed_ext - else set() - ) - if ( - not installed_ext - or not installed_ext.active - or not installed_ext.is_wasm - or "websocket.subscribe" not in installed_permission_ids - ): - await websocket.close(code=status.WS_1008_POLICY_VIOLATION) - return - - try: - conn = await wasm_extension_websocket_hub.connect(ext_id, item_id, websocket) - except ValueError: - await websocket.close(code=status.WS_1008_POLICY_VIOLATION) - return - await wasm_extension_websocket_hub.listen(conn) diff --git a/lnbits/core/views/websocket_api.py b/lnbits/core/views/websocket_api.py index 956e2cbda..ee817641f 100644 --- a/lnbits/core/views/websocket_api.py +++ b/lnbits/core/views/websocket_api.py @@ -1,8 +1,15 @@ -from fastapi import APIRouter, WebSocket +from fastapi import APIRouter, WebSocket, status + +from lnbits.core.crud import get_installed_extension +from lnbits.core.wasm_ext.api.websockets import wasm_extension_websocket_hub from ..services import websocket_manager websocket_router = APIRouter(prefix="/api/v1/ws", tags=["Websocket"]) +extension_websocket_router = APIRouter( + prefix="/api/v1/ext/ws", + tags=["Extension Websocket"], +) @websocket_router.websocket("/{item_id}") @@ -11,6 +18,35 @@ async def websocket_connect(websocket: WebSocket, item_id: str) -> None: await websocket_manager.listen(conn) +@extension_websocket_router.websocket("/{ext_id}/{item_id}") +async def extension_websocket_connect( + websocket: WebSocket, + ext_id: str, + item_id: str, +) -> None: + installed_ext = await get_installed_extension(ext_id) + installed_permission_ids = ( + {permission.id for permission in installed_ext.permissions or []} + if installed_ext + else set() + ) + if ( + not installed_ext + or not installed_ext.active + or not installed_ext.is_wasm + or "websocket.subscribe" not in installed_permission_ids + ): + await websocket.close(code=status.WS_1008_POLICY_VIOLATION) + return + + try: + conn = await wasm_extension_websocket_hub.connect(ext_id, item_id, websocket) + except ValueError: + await websocket.close(code=status.WS_1008_POLICY_VIOLATION) + return + await wasm_extension_websocket_hub.listen(conn) + + @websocket_router.post("/{item_id}") async def websocket_update_post(item_id: str, data: str): try: diff --git a/tests/unit/test_wasm_extension_websocket_api.py b/tests/unit/test_wasm_extension_websocket_api.py index e725245d0..781762265 100644 --- a/tests/unit/test_wasm_extension_websocket_api.py +++ b/tests/unit/test_wasm_extension_websocket_api.py @@ -4,11 +4,11 @@ from unittest.mock import AsyncMock import pytest from lnbits.core.models.extensions import ExtensionPermission -from lnbits.core.views.extension_websocket_api import extension_websocket_connect +from lnbits.core.views.websocket_api import extension_websocket_connect @pytest.mark.anyio -async def test_extension_websocket_api_delegates_installed_wasm_subscription(mocker): +async def test_wasm_extension_websocket_delegates_installed_wasm_subscription(mocker): websocket = AsyncMock() conn = SimpleNamespace() installed_ext = SimpleNamespace( @@ -17,17 +17,15 @@ async def test_extension_websocket_api_delegates_installed_wasm_subscription(moc permissions=[ExtensionPermission(id="websocket.subscribe")], ) mocker.patch( - "lnbits.core.views.extension_websocket_api.get_installed_extension", + "lnbits.core.views.websocket_api.get_installed_extension", AsyncMock(return_value=installed_ext), ) connect = mocker.patch( - "lnbits.core.views.extension_websocket_api." - "wasm_extension_websocket_hub.connect", + "lnbits.core.views.websocket_api.wasm_extension_websocket_hub.connect", AsyncMock(return_value=conn), ) listen = mocker.patch( - "lnbits.core.views.extension_websocket_api." - "wasm_extension_websocket_hub.listen", + "lnbits.core.views.websocket_api.wasm_extension_websocket_hub.listen", AsyncMock(), ) @@ -39,11 +37,11 @@ async def test_extension_websocket_api_delegates_installed_wasm_subscription(moc @pytest.mark.anyio -async def test_extension_websocket_api_rejects_missing_subscribe_permission(mocker): +async def test_wasm_extension_websocket_rejects_missing_subscribe_permission(mocker): websocket = AsyncMock() installed_ext = SimpleNamespace(active=True, is_wasm=True, permissions=[]) mocker.patch( - "lnbits.core.views.extension_websocket_api.get_installed_extension", + "lnbits.core.views.websocket_api.get_installed_extension", AsyncMock(return_value=installed_ext), )