addresstracker
This commit is contained in:
+11
-26
@@ -60,12 +60,9 @@ 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, int] = {} # address -> ref count
|
|
||||||
_address_tracker: "AddressTracker | None" = None
|
_address_tracker: "AddressTracker | None" = None
|
||||||
_block_tracker: "BlockTracker | None" = None
|
_block_tracker: "BlockTracker | None" = None
|
||||||
_ws_address_queues: dict[str, list[asyncio.Queue[OnchainAddressEvent]]] = {}
|
_tx_trackers: dict[str, "TransactionTracker"] = {}
|
||||||
_ws_tx_queues: dict[str, list[asyncio.Queue[OnchainTxEvent]]] = {}
|
|
||||||
_ws_block_queues: list[asyncio.Queue[BlockInfo]] = []
|
|
||||||
|
|
||||||
def init(self) -> None:
|
def init(self) -> None:
|
||||||
self.create_permanent_task(
|
self.create_permanent_task(
|
||||||
@@ -179,20 +176,12 @@ class TaskManager:
|
|||||||
|
|
||||||
def track_address(self, address: str) -> None:
|
def track_address(self, address: str) -> None:
|
||||||
"""Start tracking a Bitcoin address via Electrum (ref-counted)."""
|
"""Start tracking a Bitcoin address via Electrum (ref-counted)."""
|
||||||
count = self._tracked_addresses.get(address, 0)
|
self._get_address_tracker().add(address)
|
||||||
self._tracked_addresses[address] = count + 1
|
|
||||||
if count == 0:
|
|
||||||
self._get_address_tracker().add(address)
|
|
||||||
|
|
||||||
def untrack_address(self, address: str) -> None:
|
def untrack_address(self, address: str) -> None:
|
||||||
"""Decrement ref count; remove from shared tracker when last caller leaves."""
|
"""Decrement ref count; remove from shared tracker when last caller leaves."""
|
||||||
count = self._tracked_addresses.get(address, 0)
|
if self._address_tracker:
|
||||||
if count <= 1:
|
self._address_tracker.remove(address)
|
||||||
self._tracked_addresses.pop(address, None)
|
|
||||||
if self._address_tracker:
|
|
||||||
self._address_tracker.remove(address)
|
|
||||||
else:
|
|
||||||
self._tracked_addresses[address] = count - 1
|
|
||||||
|
|
||||||
def _get_address_tracker(self) -> "AddressTracker":
|
def _get_address_tracker(self) -> "AddressTracker":
|
||||||
if self._address_tracker is None:
|
if self._address_tracker is None:
|
||||||
@@ -217,19 +206,14 @@ class TaskManager:
|
|||||||
Raises ValueError if the address is invalid.
|
Raises ValueError if the address is invalid.
|
||||||
"""
|
"""
|
||||||
scripthash_from_address(address)
|
scripthash_from_address(address)
|
||||||
self._ws_address_queues.setdefault(address, []).append(queue)
|
self._get_address_tracker().register_queue(address, queue)
|
||||||
self.track_address(address)
|
|
||||||
|
|
||||||
def unregister_ws_address_queue(
|
def unregister_ws_address_queue(
|
||||||
self, address: str, queue: asyncio.Queue[OnchainAddressEvent]
|
self, address: str, queue: asyncio.Queue[OnchainAddressEvent]
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Deregister a per-connection queue and decrement the address ref count."""
|
"""Deregister a per-connection queue and decrement the address ref count."""
|
||||||
queues = self._ws_address_queues.get(address, [])
|
if self._address_tracker:
|
||||||
if queue in queues:
|
self._address_tracker.unregister_queue(address, queue)
|
||||||
queues.remove(queue)
|
|
||||||
if not queues:
|
|
||||||
self._ws_address_queues.pop(address, None)
|
|
||||||
self.untrack_address(address)
|
|
||||||
|
|
||||||
def register_ws_tx_queue(
|
def register_ws_tx_queue(
|
||||||
self, txid: str, queue: asyncio.Queue[OnchainTxEvent]
|
self, txid: str, queue: asyncio.Queue[OnchainTxEvent]
|
||||||
@@ -361,12 +345,13 @@ class TaskManager:
|
|||||||
task.invoice_queue.put_nowait(payment)
|
task.invoice_queue.put_nowait(payment)
|
||||||
|
|
||||||
async def _dispatch_onchain_event(self, event: OnchainAddressEvent) -> None:
|
async def _dispatch_onchain_event(self, event: OnchainAddressEvent) -> None:
|
||||||
"""Dispatches an onchain address event to listeners and WS queues."""
|
"""Dispatches an onchain address event to registered listeners.
|
||||||
|
|
||||||
|
Per-address WS queue fan-out is handled by AddressTracker itself.
|
||||||
|
"""
|
||||||
for task in self.tasks:
|
for task in self.tasks:
|
||||||
if task.onchain_queue:
|
if task.onchain_queue:
|
||||||
task.onchain_queue.put_nowait(event)
|
task.onchain_queue.put_nowait(event)
|
||||||
for q in list(self._ws_address_queues.get(event.address, [])):
|
|
||||||
q.put_nowait(event)
|
|
||||||
|
|
||||||
async def _dispatch_onchain_tx_event(self, event: OnchainTxEvent) -> None:
|
async def _dispatch_onchain_tx_event(self, event: OnchainTxEvent) -> None:
|
||||||
"""Dispatches a tx event to all WS queues watching that txid."""
|
"""Dispatches a tx event to all WS queues watching that txid."""
|
||||||
|
|||||||
+95
-35
@@ -774,7 +774,9 @@ class AddressTracker:
|
|||||||
"""
|
"""
|
||||||
Subscribes to a set of Bitcoin addresses over a single shared Electrum
|
Subscribes to a set of Bitcoin addresses over a single shared Electrum
|
||||||
connection and calls a callback on every balance/history change.
|
connection and calls a callback on every balance/history change.
|
||||||
Addresses can be added/removed at runtime via :meth:`add`/:meth:`remove`.
|
Addresses can be added/removed at runtime via :meth:`add`/:meth:`remove`,
|
||||||
|
and per-connection queues can be attached via :meth:`register_queue` for
|
||||||
|
consumers (e.g. websockets) that want events for one specific address.
|
||||||
Reconnects automatically on failure.
|
Reconnects automatically on failure.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -783,20 +785,43 @@ class AddressTracker:
|
|||||||
|
|
||||||
def __init__(self, url: str) -> None:
|
def __init__(self, url: str) -> None:
|
||||||
self.url = url
|
self.url = url
|
||||||
self._addresses: set[str] = set()
|
self._ref_counts: dict[str, int] = {}
|
||||||
|
self._queues: dict[str, list[asyncio.Queue[OnchainAddressEvent]]] = {}
|
||||||
self._updated = asyncio.Event()
|
self._updated = asyncio.Event()
|
||||||
|
|
||||||
def add(self, address: str) -> None:
|
def add(self, address: str) -> None:
|
||||||
"""Start tracking an address on the shared connection."""
|
"""Start tracking an address on the shared connection (ref-counted)."""
|
||||||
if address not in self._addresses:
|
count = self._ref_counts.get(address, 0)
|
||||||
self._addresses.add(address)
|
self._ref_counts[address] = count + 1
|
||||||
|
if count == 0:
|
||||||
self._updated.set()
|
self._updated.set()
|
||||||
|
|
||||||
def remove(self, address: str) -> None:
|
def remove(self, address: str) -> None:
|
||||||
"""Stop tracking an address on the shared connection."""
|
"""Decrement ref count; drop the subscription once the last caller leaves."""
|
||||||
if address in self._addresses:
|
count = self._ref_counts.get(address, 0)
|
||||||
self._addresses.discard(address)
|
if count <= 1:
|
||||||
|
self._ref_counts.pop(address, None)
|
||||||
self._updated.set()
|
self._updated.set()
|
||||||
|
else:
|
||||||
|
self._ref_counts[address] = count - 1
|
||||||
|
|
||||||
|
def register_queue(
|
||||||
|
self, address: str, queue: "asyncio.Queue[OnchainAddressEvent]"
|
||||||
|
) -> None:
|
||||||
|
"""Register a per-connection queue to receive events for `address`."""
|
||||||
|
self._queues.setdefault(address, []).append(queue)
|
||||||
|
self.add(address)
|
||||||
|
|
||||||
|
def unregister_queue(
|
||||||
|
self, address: str, queue: "asyncio.Queue[OnchainAddressEvent]"
|
||||||
|
) -> None:
|
||||||
|
"""Deregister a per-connection queue for `address`."""
|
||||||
|
queues = self._queues.get(address, [])
|
||||||
|
if queue in queues:
|
||||||
|
queues.remove(queue)
|
||||||
|
if not queues:
|
||||||
|
self._queues.pop(address, None)
|
||||||
|
self.remove(address)
|
||||||
|
|
||||||
async def run(
|
async def run(
|
||||||
self,
|
self,
|
||||||
@@ -843,7 +868,7 @@ class AddressTracker:
|
|||||||
subscribed: dict[str, str],
|
subscribed: dict[str, str],
|
||||||
callback: Callable[[OnchainAddressEvent], Coroutine[Any, Any, None]],
|
callback: Callable[[OnchainAddressEvent], Coroutine[Any, Any, None]],
|
||||||
) -> None:
|
) -> None:
|
||||||
wanted = {a: scripthash_from_address(a) for a in self._addresses}
|
wanted = {a: scripthash_from_address(a) for a in self._ref_counts}
|
||||||
for address, scripthash in wanted.items():
|
for address, scripthash in wanted.items():
|
||||||
if scripthash not in subscribed:
|
if scripthash not in subscribed:
|
||||||
subscribed[scripthash] = address
|
subscribed[scripthash] = address
|
||||||
@@ -872,8 +897,8 @@ class AddressTracker:
|
|||||||
t.cancel()
|
t.cancel()
|
||||||
return closed_task in done
|
return closed_task in done
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
async def _fetch_and_dispatch(
|
async def _fetch_and_dispatch(
|
||||||
|
self,
|
||||||
client: ElectrumClient,
|
client: ElectrumClient,
|
||||||
address: str,
|
address: str,
|
||||||
scripthash: str,
|
scripthash: str,
|
||||||
@@ -898,15 +923,16 @@ class AddressTracker:
|
|||||||
for m in mempool_r:
|
for m in mempool_r:
|
||||||
if m.tx_hash not in seen:
|
if m.tx_hash not in seen:
|
||||||
history.append(HistoryEntry(tx_hash=m.tx_hash, height=0, fee=m.fee))
|
history.append(HistoryEntry(tx_hash=m.tx_hash, height=0, fee=m.fee))
|
||||||
await callback(
|
event = OnchainAddressEvent(
|
||||||
OnchainAddressEvent(
|
address=address,
|
||||||
address=address,
|
confirmed=balance_r.confirmed,
|
||||||
confirmed=balance_r.confirmed,
|
unconfirmed=balance_r.unconfirmed,
|
||||||
unconfirmed=balance_r.unconfirmed,
|
history=history,
|
||||||
history=history,
|
history_error=history_error,
|
||||||
history_error=history_error,
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
|
for q in list(self._queues.get(address, [])):
|
||||||
|
q.put_nowait(event)
|
||||||
|
await callback(event)
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
@@ -935,6 +961,8 @@ class TransactionTracker:
|
|||||||
Subscribes to a Bitcoin transaction via Electrum and calls a callback on
|
Subscribes to a Bitcoin transaction via Electrum and calls a callback on
|
||||||
each status change (unconfirmed → confirmed). Stops automatically once
|
each status change (unconfirmed → confirmed). Stops automatically once
|
||||||
the transaction is confirmed or ``is_active()`` returns ``False``.
|
the transaction is confirmed or ``is_active()`` returns ``False``.
|
||||||
|
Per-connection queues can be attached via :meth:`register_queue` for
|
||||||
|
consumers (e.g. websockets) that want events for this transaction.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
url: Electrum server URL (e.g. ``ssl://electrum.blockstream.info:50002``).
|
url: Electrum server URL (e.g. ``ssl://electrum.blockstream.info:50002``).
|
||||||
@@ -942,6 +970,19 @@ class TransactionTracker:
|
|||||||
|
|
||||||
def __init__(self, url: str) -> None:
|
def __init__(self, url: str) -> None:
|
||||||
self.url = url
|
self.url = url
|
||||||
|
self._queues: list[asyncio.Queue[OnchainTxEvent]] = []
|
||||||
|
|
||||||
|
def register_queue(self, queue: "asyncio.Queue[OnchainTxEvent]") -> None:
|
||||||
|
"""Register a per-connection queue to receive events for this tx."""
|
||||||
|
self._queues.append(queue)
|
||||||
|
|
||||||
|
def unregister_queue(self, queue: "asyncio.Queue[OnchainTxEvent]") -> None:
|
||||||
|
"""Deregister a per-connection queue."""
|
||||||
|
if queue in self._queues:
|
||||||
|
self._queues.remove(queue)
|
||||||
|
|
||||||
|
def has_queues(self) -> bool:
|
||||||
|
return bool(self._queues)
|
||||||
|
|
||||||
async def track(
|
async def track(
|
||||||
self,
|
self,
|
||||||
@@ -989,7 +1030,7 @@ class TransactionTracker:
|
|||||||
) -> None:
|
) -> None:
|
||||||
if params and params[0] == _sh:
|
if params and params[0] == _sh:
|
||||||
ev = await self._fetch_status(client, txid, _sh)
|
ev = await self._fetch_status(client, txid, _sh)
|
||||||
await callback(ev)
|
await self._dispatch(ev, callback)
|
||||||
if ev.confirmed:
|
if ev.confirmed:
|
||||||
_done.set()
|
_done.set()
|
||||||
|
|
||||||
@@ -997,7 +1038,7 @@ class TransactionTracker:
|
|||||||
await client.subscribe_scripthash(scripthash, on_change)
|
await client.subscribe_scripthash(scripthash, on_change)
|
||||||
|
|
||||||
event = await self._fetch_status(client, txid, scripthash)
|
event = await self._fetch_status(client, txid, scripthash)
|
||||||
await callback(event)
|
await self._dispatch(event, callback)
|
||||||
if event.confirmed:
|
if event.confirmed:
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@@ -1009,6 +1050,15 @@ class TransactionTracker:
|
|||||||
pass
|
pass
|
||||||
return confirmed_event.is_set()
|
return confirmed_event.is_set()
|
||||||
|
|
||||||
|
async def _dispatch(
|
||||||
|
self,
|
||||||
|
event: OnchainTxEvent,
|
||||||
|
callback: Callable[[OnchainTxEvent], Coroutine[Any, Any, None]],
|
||||||
|
) -> None:
|
||||||
|
for q in list(self._queues):
|
||||||
|
q.put_nowait(event)
|
||||||
|
await callback(event)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
async def _fetch_status(
|
async def _fetch_status(
|
||||||
client: ElectrumClient, txid: str, scripthash: str | None
|
client: ElectrumClient, txid: str, scripthash: str | None
|
||||||
@@ -1041,8 +1091,9 @@ class TransactionTracker:
|
|||||||
|
|
||||||
class BlockTracker:
|
class BlockTracker:
|
||||||
"""
|
"""
|
||||||
Subscribes to new block headers via Electrum and calls a callback on
|
Subscribes to new block headers via Electrum and dispatches them to
|
||||||
every new block. Reconnects automatically on failure.
|
registered queues. Per-connection queues can be attached via
|
||||||
|
:meth:`register_queue`. Reconnects automatically on failure.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
url: Electrum server URL (e.g. ``ssl://electrum.blockstream.info:50002``).
|
url: Electrum server URL (e.g. ``ssl://electrum.blockstream.info:50002``).
|
||||||
@@ -1050,15 +1101,24 @@ class BlockTracker:
|
|||||||
|
|
||||||
def __init__(self, url: str) -> None:
|
def __init__(self, url: str) -> None:
|
||||||
self.url = url
|
self.url = url
|
||||||
|
self._queues: list[asyncio.Queue[BlockInfo]] = []
|
||||||
|
|
||||||
async def run(
|
def register_queue(self, queue: "asyncio.Queue[BlockInfo]") -> None:
|
||||||
self,
|
"""Register a per-connection queue to receive new block events."""
|
||||||
callback: Callable[[BlockInfo], Coroutine[Any, Any, None]],
|
self._queues.append(queue)
|
||||||
is_active: Callable[[], bool],
|
|
||||||
) -> None:
|
def unregister_queue(self, queue: "asyncio.Queue[BlockInfo]") -> None:
|
||||||
|
"""Deregister a per-connection queue."""
|
||||||
|
if queue in self._queues:
|
||||||
|
self._queues.remove(queue)
|
||||||
|
|
||||||
|
def has_queues(self) -> bool:
|
||||||
|
return bool(self._queues)
|
||||||
|
|
||||||
|
async def run(self, is_active: Callable[[], bool]) -> None:
|
||||||
while is_active():
|
while is_active():
|
||||||
try:
|
try:
|
||||||
await self._run_once(callback, is_active)
|
await self._run_once(is_active)
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
raise
|
raise
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
@@ -1067,19 +1127,15 @@ class BlockTracker:
|
|||||||
logger.warning(f"BlockTracker: {exc!s}, retrying in 5s")
|
logger.warning(f"BlockTracker: {exc!s}, retrying in 5s")
|
||||||
await asyncio.sleep(5)
|
await asyncio.sleep(5)
|
||||||
|
|
||||||
async def _run_once(
|
async def _run_once(self, is_active: Callable[[], bool]) -> None:
|
||||||
self,
|
|
||||||
callback: Callable[[BlockInfo], Coroutine[Any, Any, None]],
|
|
||||||
is_active: Callable[[], bool],
|
|
||||||
) -> None:
|
|
||||||
async with ElectrumClient(self.url) as client:
|
async with ElectrumClient(self.url) as client:
|
||||||
|
|
||||||
async def on_header(params: list[Any]) -> None:
|
async def on_header(params: list[Any]) -> None:
|
||||||
h = params[0]
|
h = params[0]
|
||||||
await callback(parse_block_header(h["hex"], h["height"]))
|
self._dispatch(parse_block_header(h["hex"], h["height"]))
|
||||||
|
|
||||||
tip = await client.subscribe_headers(on_header)
|
tip = await client.subscribe_headers(on_header)
|
||||||
await callback(parse_block_header(tip.hex, tip.height))
|
self._dispatch(parse_block_header(tip.hex, tip.height))
|
||||||
|
|
||||||
while is_active():
|
while is_active():
|
||||||
try:
|
try:
|
||||||
@@ -1087,3 +1143,7 @@ class BlockTracker:
|
|||||||
break # connection closed; reconnect
|
break # connection closed; reconnect
|
||||||
except asyncio.TimeoutError:
|
except asyncio.TimeoutError:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
def _dispatch(self, event: BlockInfo) -> None:
|
||||||
|
for q in list(self._queues):
|
||||||
|
q.put_nowait(event)
|
||||||
|
|||||||
Reference in New Issue
Block a user