refactor: move ws logic

This commit is contained in:
Vlad Stan 2026-07-17 10:56:31 +03:00
parent eac64dc6de
commit 486426e453
4 changed files with 45 additions and 50 deletions

View file

@ -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"])

View file

@ -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)

View file

@ -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:

View file

@ -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),
)