feat: task manager onchain listeners
This commit is contained in:
@@ -1092,6 +1092,9 @@ class EnvSettings(LNbitsSettings):
|
|||||||
lnbits_max_extensions: int = Field(default=0, ge=0)
|
lnbits_max_extensions: int = Field(default=0, ge=0)
|
||||||
task_heart_beat_verbose: bool = Field(default=False)
|
task_heart_beat_verbose: bool = Field(default=False)
|
||||||
task_heart_beat_interval: int = Field(default=30)
|
task_heart_beat_interval: int = Field(default=30)
|
||||||
|
lnbits_blockexplorer_electrum_url: str = Field(
|
||||||
|
default="ssl://electrum.blockstream.info:50002"
|
||||||
|
)
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def has_default_extension_path(self) -> bool:
|
def has_default_extension_path(self) -> bool:
|
||||||
|
|||||||
+169
-5
@@ -3,6 +3,7 @@ import traceback
|
|||||||
import uuid
|
import uuid
|
||||||
from collections.abc import Callable, Coroutine
|
from collections.abc import Callable, Coroutine
|
||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
@@ -18,6 +19,13 @@ class PublicTask(BaseModel):
|
|||||||
created_at: datetime
|
created_at: datetime
|
||||||
|
|
||||||
|
|
||||||
|
class OnchainAddressEvent(BaseModel):
|
||||||
|
address: str
|
||||||
|
confirmed: int # satoshis
|
||||||
|
unconfirmed: int # satoshis
|
||||||
|
txids: list[str]
|
||||||
|
|
||||||
|
|
||||||
class Task:
|
class Task:
|
||||||
"""Model used on the backend to keep track of background tasks."""
|
"""Model used on the backend to keep track of background tasks."""
|
||||||
|
|
||||||
@@ -26,18 +34,21 @@ class Task:
|
|||||||
created_at: datetime
|
created_at: datetime
|
||||||
task: asyncio.Task
|
task: asyncio.Task
|
||||||
invoice_queue: asyncio.Queue[Payment] | None = None
|
invoice_queue: asyncio.Queue[Payment] | None = None
|
||||||
|
onchain_queue: asyncio.Queue[OnchainAddressEvent] | None = None
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
coro: Coroutine,
|
coro: Coroutine,
|
||||||
name: str | None = None,
|
name: str | None = None,
|
||||||
invoice_queue: asyncio.Queue | None = None,
|
invoice_queue: asyncio.Queue | None = None,
|
||||||
|
onchain_queue: asyncio.Queue | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
self.coro = coro
|
self.coro = coro
|
||||||
self.name = name or f"task_{uuid.uuid4()}"
|
self.name = name or f"task_{uuid.uuid4()}"
|
||||||
self.created_at = datetime.now(timezone.utc)
|
self.created_at = datetime.now(timezone.utc)
|
||||||
self.task = asyncio.create_task(self.coro, name=self.name)
|
self.task = asyncio.create_task(self.coro, name=self.name)
|
||||||
self.invoice_queue = invoice_queue
|
self.invoice_queue = invoice_queue
|
||||||
|
self.onchain_queue = onchain_queue
|
||||||
|
|
||||||
|
|
||||||
class TaskManager:
|
class TaskManager:
|
||||||
@@ -46,6 +57,7 @@ class TaskManager:
|
|||||||
tasks: list[Task] = []
|
tasks: list[Task] = []
|
||||||
invoice_queue: asyncio.Queue[Payment] = asyncio.Queue()
|
invoice_queue: asyncio.Queue[Payment] = asyncio.Queue()
|
||||||
internal_invoice_queue: asyncio.Queue[Payment] = asyncio.Queue()
|
internal_invoice_queue: asyncio.Queue[Payment] = asyncio.Queue()
|
||||||
|
_tracked_addresses: dict[str, str] = {} # address -> task_name
|
||||||
|
|
||||||
def init(self) -> None:
|
def init(self) -> None:
|
||||||
self.create_permanent_task(
|
self.create_permanent_task(
|
||||||
@@ -84,13 +96,14 @@ class TaskManager:
|
|||||||
coro: Coroutine,
|
coro: Coroutine,
|
||||||
name: str | None = None,
|
name: str | None = None,
|
||||||
invoice_queue: asyncio.Queue | None = None,
|
invoice_queue: asyncio.Queue | None = None,
|
||||||
|
onchain_queue: asyncio.Queue | None = None,
|
||||||
) -> Task:
|
) -> Task:
|
||||||
"""Create a task. If a task with the same name exists, it will be cancelled."""
|
"""Create a task. If a task with the same name exists, it will be cancelled."""
|
||||||
if name:
|
if name:
|
||||||
task = self.get_task(name)
|
task = self.get_task(name)
|
||||||
if task:
|
if task:
|
||||||
self.cancel_task(task)
|
self.cancel_task(task)
|
||||||
task = Task(coro=coro, name=name, invoice_queue=invoice_queue)
|
task = Task(coro=coro, name=name, invoice_queue=invoice_queue, onchain_queue=onchain_queue)
|
||||||
self.tasks.append(task)
|
self.tasks.append(task)
|
||||||
return task
|
return task
|
||||||
|
|
||||||
@@ -98,6 +111,7 @@ class TaskManager:
|
|||||||
self,
|
self,
|
||||||
func: Callable[[], Coroutine],
|
func: Callable[[], Coroutine],
|
||||||
invoice_queue: asyncio.Queue | None = None,
|
invoice_queue: asyncio.Queue | None = None,
|
||||||
|
onchain_queue: asyncio.Queue | None = None,
|
||||||
name: str | None = None,
|
name: str | None = None,
|
||||||
interval: int = 0,
|
interval: int = 0,
|
||||||
) -> Task:
|
) -> Task:
|
||||||
@@ -110,7 +124,10 @@ class TaskManager:
|
|||||||
await asyncio.sleep(interval)
|
await asyncio.sleep(interval)
|
||||||
|
|
||||||
return self.create_task(
|
return self.create_task(
|
||||||
coro=wrapper(), name=name or func.__name__, invoice_queue=invoice_queue
|
coro=wrapper(),
|
||||||
|
name=name or func.__name__,
|
||||||
|
invoice_queue=invoice_queue,
|
||||||
|
onchain_queue=onchain_queue,
|
||||||
)
|
)
|
||||||
|
|
||||||
def register_invoice_listener(
|
def register_invoice_listener(
|
||||||
@@ -130,6 +147,51 @@ class TaskManager:
|
|||||||
invoice_queue=queue,
|
invoice_queue=queue,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def register_onchain_listener(
|
||||||
|
self,
|
||||||
|
func: Callable[[OnchainAddressEvent], Coroutine],
|
||||||
|
name: str | None = None,
|
||||||
|
) -> Task:
|
||||||
|
"""
|
||||||
|
Register a callback for onchain address events dispatched by track_address.
|
||||||
|
Will call the provided coroutine with an OnchainAddressEvent on each update.
|
||||||
|
"""
|
||||||
|
name = f"{name or uuid.uuid4()}_onchain_listener"
|
||||||
|
queue: asyncio.Queue[OnchainAddressEvent] = asyncio.Queue()
|
||||||
|
return self.create_permanent_task(
|
||||||
|
self._onchain_listener_worker(func, queue),
|
||||||
|
name=name,
|
||||||
|
onchain_queue=queue,
|
||||||
|
)
|
||||||
|
|
||||||
|
def track_address(self, address: str) -> None:
|
||||||
|
"""Start tracking a Bitcoin address via Electrum. Dispatches OnchainAddressEvents."""
|
||||||
|
if address in self._tracked_addresses:
|
||||||
|
return
|
||||||
|
task_name = f"onchain_address_{address}"
|
||||||
|
self._tracked_addresses[address] = task_name
|
||||||
|
self.create_task(self._address_tracker(address), name=task_name)
|
||||||
|
|
||||||
|
def untrack_address(self, address: str) -> None:
|
||||||
|
"""Stop tracking a Bitcoin address."""
|
||||||
|
task_name = self._tracked_addresses.pop(address, None)
|
||||||
|
if task_name:
|
||||||
|
task = self.get_task(task_name)
|
||||||
|
if task:
|
||||||
|
self.cancel_task(task)
|
||||||
|
|
||||||
|
def track_transaction(
|
||||||
|
self,
|
||||||
|
txid: str,
|
||||||
|
callback: Callable[[str, int], Coroutine],
|
||||||
|
) -> Task:
|
||||||
|
"""
|
||||||
|
Poll until a transaction is confirmed, then call callback(txid, block_height).
|
||||||
|
The task cancels itself after the callback fires.
|
||||||
|
"""
|
||||||
|
task_name = f"onchain_tx_{txid}"
|
||||||
|
return self.create_task(self._transaction_tracker(txid, callback), name=task_name)
|
||||||
|
|
||||||
async def _heart_beat(self) -> None:
|
async def _heart_beat(self) -> None:
|
||||||
"""A heartbeat that removes done tasks logs the number of tasks."""
|
"""A heartbeat that removes done tasks logs the number of tasks."""
|
||||||
for task in self.tasks:
|
for task in self.tasks:
|
||||||
@@ -142,10 +204,13 @@ class TaskManager:
|
|||||||
if task.task and task.task.done():
|
if task.task and task.task.done():
|
||||||
logger.debug(f"Task Manager: task `{task.name}` is done.")
|
logger.debug(f"Task Manager: task `{task.name}` is done.")
|
||||||
self.cancel_task(task)
|
self.cancel_task(task)
|
||||||
listeners_count = sum(1 for task in self.tasks if task.invoice_queue)
|
invoice_listeners = sum(1 for task in self.tasks if task.invoice_queue)
|
||||||
|
onchain_listeners = sum(1 for task in self.tasks if task.onchain_queue)
|
||||||
|
other_tasks = len(self.tasks) - invoice_listeners - onchain_listeners
|
||||||
logger.debug(
|
logger.debug(
|
||||||
f"Task Manager: {len(self.tasks) - listeners_count} tasks "
|
f"Task Manager: {other_tasks} tasks, "
|
||||||
f"and {listeners_count} invoice listeners."
|
f"{invoice_listeners} invoice listeners, "
|
||||||
|
f"{onchain_listeners} onchain listeners."
|
||||||
)
|
)
|
||||||
|
|
||||||
async def _catch_everything_and_restart(
|
async def _catch_everything_and_restart(
|
||||||
@@ -178,6 +243,17 @@ class TaskManager:
|
|||||||
|
|
||||||
return wrapper
|
return wrapper
|
||||||
|
|
||||||
|
def _onchain_listener_worker(
|
||||||
|
self,
|
||||||
|
func: Callable[[OnchainAddressEvent], Coroutine],
|
||||||
|
queue: asyncio.Queue[OnchainAddressEvent],
|
||||||
|
) -> Callable:
|
||||||
|
async def wrapper() -> None:
|
||||||
|
event: OnchainAddressEvent = await queue.get()
|
||||||
|
await func(event)
|
||||||
|
|
||||||
|
return wrapper
|
||||||
|
|
||||||
def _invoice_dispatcher(self, payment: Payment) -> None:
|
def _invoice_dispatcher(self, payment: Payment) -> None:
|
||||||
"""Dispatches a payment to all registered invoice listeners."""
|
"""Dispatches a payment to all registered invoice listeners."""
|
||||||
for task in self.tasks:
|
for task in self.tasks:
|
||||||
@@ -186,6 +262,13 @@ class TaskManager:
|
|||||||
logger.debug(f"Enqueing payment to task {task.name}")
|
logger.debug(f"Enqueing payment to task {task.name}")
|
||||||
task.invoice_queue.put_nowait(payment)
|
task.invoice_queue.put_nowait(payment)
|
||||||
|
|
||||||
|
def _dispatch_onchain_event(self, event: OnchainAddressEvent) -> None:
|
||||||
|
"""Dispatches an onchain address event to all registered onchain listeners."""
|
||||||
|
for task in self.tasks:
|
||||||
|
if not task.onchain_queue:
|
||||||
|
continue
|
||||||
|
task.onchain_queue.put_nowait(event)
|
||||||
|
|
||||||
async def _invoice_listener_consumer(self) -> None:
|
async def _invoice_listener_consumer(self) -> None:
|
||||||
payment = await self.invoice_queue.get()
|
payment = await self.invoice_queue.get()
|
||||||
logger.info(f"got a payment notification {payment.checking_id}")
|
logger.info(f"got a payment notification {payment.checking_id}")
|
||||||
@@ -196,5 +279,86 @@ class TaskManager:
|
|||||||
logger.info(f"got an internal payment notification {payment.checking_id}")
|
logger.info(f"got an internal payment notification {payment.checking_id}")
|
||||||
self._invoice_dispatcher(payment)
|
self._invoice_dispatcher(payment)
|
||||||
|
|
||||||
|
async def _address_tracker(self, address: str) -> None:
|
||||||
|
"""Track an address via Electrum subscription; dispatches OnchainAddressEvents."""
|
||||||
|
from lnbits.utils.electrum import ElectrumClient, ElectrumError, scripthash_from_address
|
||||||
|
|
||||||
|
electrum_url = settings.lnbits_blockexplorer_electrum_url
|
||||||
|
scripthash = scripthash_from_address(address)
|
||||||
|
|
||||||
|
while address in self._tracked_addresses and settings.lnbits_running:
|
||||||
|
try:
|
||||||
|
async with ElectrumClient(electrum_url) as client:
|
||||||
|
await self._fetch_and_dispatch_address(client, address, scripthash)
|
||||||
|
|
||||||
|
async def on_status_change(params: list[Any]) -> None:
|
||||||
|
if params and params[0] == scripthash:
|
||||||
|
await self._fetch_and_dispatch_address(client, address, scripthash)
|
||||||
|
|
||||||
|
await client.subscribe_scripthash(scripthash, on_status_change)
|
||||||
|
|
||||||
|
while address in self._tracked_addresses and settings.lnbits_running:
|
||||||
|
await asyncio.sleep(30)
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
raise
|
||||||
|
except Exception as exc:
|
||||||
|
if not settings.lnbits_running:
|
||||||
|
return
|
||||||
|
logger.warning(f"Address tracker {address}: {exc!s}, retrying in 5s")
|
||||||
|
await asyncio.sleep(5)
|
||||||
|
|
||||||
|
async def _fetch_and_dispatch_address(
|
||||||
|
self, client: Any, address: str, scripthash: str
|
||||||
|
) -> None:
|
||||||
|
"""Fetch balance + history for a scripthash and dispatch an OnchainAddressEvent."""
|
||||||
|
from lnbits.utils.electrum import ElectrumError
|
||||||
|
|
||||||
|
balance = await client.get_balance(scripthash)
|
||||||
|
txids: list[str] = []
|
||||||
|
try:
|
||||||
|
history = await client.get_history(scripthash)
|
||||||
|
txids = [e.tx_hash for e in history]
|
||||||
|
except ElectrumError:
|
||||||
|
pass
|
||||||
|
try:
|
||||||
|
mempool = await client.get_mempool(scripthash)
|
||||||
|
for e in mempool:
|
||||||
|
if e.tx_hash not in txids:
|
||||||
|
txids.append(e.tx_hash)
|
||||||
|
except ElectrumError:
|
||||||
|
pass
|
||||||
|
self._dispatch_onchain_event(
|
||||||
|
OnchainAddressEvent(
|
||||||
|
address=address,
|
||||||
|
confirmed=balance.confirmed,
|
||||||
|
unconfirmed=balance.unconfirmed,
|
||||||
|
txids=txids,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _transaction_tracker(
|
||||||
|
self, txid: str, callback: Callable[[str, int], Coroutine]
|
||||||
|
) -> None:
|
||||||
|
"""Poll Electrum until txid is confirmed, then fire callback(txid, height)."""
|
||||||
|
from lnbits.utils.electrum import ElectrumClient
|
||||||
|
|
||||||
|
electrum_url = settings.lnbits_blockexplorer_electrum_url
|
||||||
|
while settings.lnbits_running:
|
||||||
|
try:
|
||||||
|
async with ElectrumClient(electrum_url) as client:
|
||||||
|
tx = await client.get_transaction(txid, verbose=True)
|
||||||
|
if isinstance(tx, dict):
|
||||||
|
height = tx.get("blockheight") or tx.get("block_height", 0)
|
||||||
|
if height and height > 0:
|
||||||
|
await callback(txid, height)
|
||||||
|
return
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
raise
|
||||||
|
except Exception as exc:
|
||||||
|
if not settings.lnbits_running:
|
||||||
|
return
|
||||||
|
logger.warning(f"Tx tracker {txid[:8]}: {exc!s}")
|
||||||
|
await asyncio.sleep(60)
|
||||||
|
|
||||||
|
|
||||||
task_manager = TaskManager()
|
task_manager = TaskManager()
|
||||||
|
|||||||
Reference in New Issue
Block a user