diff --git a/lnbits/core/views/blockexplorer_api.py b/lnbits/core/views/blockexplorer_api.py index 9bb146fe4..9d033514f 100644 --- a/lnbits/core/views/blockexplorer_api.py +++ b/lnbits/core/views/blockexplorer_api.py @@ -3,12 +3,16 @@ from http import HTTPStatus from typing import Annotated from fastapi import APIRouter, Depends, HTTPException, Request, WebSocket -from loguru import logger from pydantic.types import UUID4 from lnbits.decorators import check_access_token, check_user_exists from lnbits.settings import settings -from lnbits.task_manager import OnchainAddressEvent, OnchainTxEvent, task_manager +from lnbits.task_manager import ( + OnchainAddressEvent, + OnchainTxEvent, + relay_ws_queue, + task_manager, +) from lnbits.utils.electrum import ( AddressResponse, Balance, @@ -154,37 +158,11 @@ async def ws_blocks(websocket: WebSocket) -> None: queue: asyncio.Queue[BlockInfo] = asyncio.Queue() task_manager.register_ws_block_queue(queue) try: - await _ws_blocks_loop(websocket, queue) + await relay_ws_queue(websocket, queue) finally: task_manager.unregister_ws_block_queue(queue) -async def _ws_blocks_loop( - websocket: WebSocket, - queue: asyncio.Queue[BlockInfo], -) -> None: - try: - while True: - recv_task = asyncio.create_task(websocket.receive()) - event_task = asyncio.create_task(queue.get()) - done, pending = await asyncio.wait( - [recv_task, event_task], return_when=asyncio.FIRST_COMPLETED - ) - for t in pending: - t.cancel() - if recv_task in done: - if recv_task.result().get("type") == "websocket.disconnect": - break - if event_task in done: - try: - await websocket.send_json(event_task.result().dict()) - except Exception as exc: - logger.debug(f"ws_blocks send error: {exc}") - break - except Exception as exc: - logger.debug(f"ws_blocks error: {exc}") - - def _address_event_to_response(event: OnchainAddressEvent) -> AddressResponse: return AddressResponse( balance=Balance(confirmed=event.confirmed, unconfirmed=event.unconfirmed), @@ -207,40 +185,11 @@ async def ws_address(websocket: WebSocket, address: str) -> None: await websocket.close(code=1008, reason=str(e)) return try: - await _ws_address_loop(websocket, address, queue) + await relay_ws_queue(websocket, queue, serialize=_address_event_to_response) finally: task_manager.unregister_ws_address_queue(address, queue) -async def _ws_address_loop( - websocket: WebSocket, - address: str, - queue: asyncio.Queue[OnchainAddressEvent], -) -> None: - try: - while True: - recv_task = asyncio.create_task(websocket.receive()) - event_task = asyncio.create_task(queue.get()) - done, pending = await asyncio.wait( - [recv_task, event_task], return_when=asyncio.FIRST_COMPLETED - ) - for t in pending: - t.cancel() - if recv_task in done: - if recv_task.result().get("type") == "websocket.disconnect": - break - if event_task in done: - try: - await websocket.send_json( - _address_event_to_response(event_task.result()).dict() - ) - except Exception as exc: - logger.debug(f"ws_address send error: {exc}") - break - except Exception as exc: - logger.debug(f"ws_address error: {exc}") - - @blockexplorer_router.websocket("/ws/tx/{txid}") async def ws_tx(websocket: WebSocket, txid: str) -> None: if not settings.lnbits_blockexplorer_enabled: @@ -251,35 +200,6 @@ async def ws_tx(websocket: WebSocket, txid: str) -> None: queue: asyncio.Queue[OnchainTxEvent] = asyncio.Queue() task_manager.register_ws_tx_queue(txid, queue) try: - await _ws_tx_loop(websocket, queue) + await relay_ws_queue(websocket, queue, stop_after=lambda e: e.confirmed) finally: task_manager.unregister_ws_tx_queue(txid, queue) - - -async def _ws_tx_loop( - websocket: WebSocket, - queue: asyncio.Queue[OnchainTxEvent], -) -> None: - try: - while True: - recv_task = asyncio.create_task(websocket.receive()) - event_task = asyncio.create_task(queue.get()) - done, pending = await asyncio.wait( - [recv_task, event_task], return_when=asyncio.FIRST_COMPLETED - ) - for t in pending: - t.cancel() - if recv_task in done: - if recv_task.result().get("type") == "websocket.disconnect": - break - if event_task in done: - event: OnchainTxEvent = event_task.result() - try: - await websocket.send_json(event.dict()) - except Exception as exc: - logger.debug(f"ws_tx send error: {exc}") - break - if event.confirmed: - break - except Exception as exc: - logger.debug(f"ws_tx error: {exc}") diff --git a/lnbits/task_manager.py b/lnbits/task_manager.py index 01ed0680a..43972aad5 100644 --- a/lnbits/task_manager.py +++ b/lnbits/task_manager.py @@ -3,7 +3,9 @@ import traceback import uuid from collections.abc import Callable, Coroutine from datetime import datetime, timezone +from typing import TypeVar +from fastapi import WebSocket from loguru import logger from pydantic import BaseModel @@ -403,4 +405,45 @@ class TaskManager: ) +T = TypeVar("T", bound=BaseModel) + + +async def relay_ws_queue( + websocket: WebSocket, + queue: "asyncio.Queue[T]", + serialize: Callable[[T], BaseModel] = lambda e: e, + stop_after: Callable[[T], bool] = lambda _: False, +) -> None: + """ + Pumps events from `queue` to `websocket` as JSON until the client + disconnects, sending fails, or `stop_after` returns True for an event. + Shared by the blockexplorer address/tx/block websocket endpoints. + """ + try: + while True: + recv_task = asyncio.create_task(websocket.receive()) + event_task = asyncio.create_task(queue.get()) + done, pending = await asyncio.wait( + [recv_task, event_task], return_when=asyncio.FIRST_COMPLETED + ) + for t in pending: + t.cancel() + disconnect = recv_task in done and ( + recv_task.result().get("type") == "websocket.disconnect" + ) + if disconnect: + break + if event_task in done: + event = event_task.result() + try: + await websocket.send_json(serialize(event).dict()) + except Exception as exc: + logger.debug(f"ws relay send error: {exc}") + break + if stop_after(event): + break + except Exception as exc: + logger.debug(f"ws relay error: {exc}") + + task_manager = TaskManager() diff --git a/lnbits/utils/electrum.py b/lnbits/utils/electrum.py index 0ae68070e..80092f3f5 100644 --- a/lnbits/utils/electrum.py +++ b/lnbits/utils/electrum.py @@ -827,9 +827,7 @@ class AddressTracker: return address = subscribed.get(params[0]) if address: - await self._fetch_and_dispatch( - client, address, params[0], callback - ) + await self._fetch_and_dispatch(client, address, params[0], callback) client.on("blockchain.scripthash.subscribe", on_status_change)