fix only have address updates sent if tracked

This commit is contained in:
dni
2026-07-13 09:09:42 +02:00
parent fe02f54021
commit 3f28d7a31a
+32 -10
View File
@@ -63,12 +63,15 @@ class Task:
class TaskManager: class TaskManager:
"""Singleton class to manage background tasks.""" """Singleton class to manage background tasks."""
ONCHAIN_ADDRESS_LISTENER_SUFFIX = "_onchain_address_listener"
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()
_address_tracker: "AddressTracker | None" = None _address_tracker: "AddressTracker | None" = None
_block_tracker: "BlockTracker | None" = None _block_tracker: "BlockTracker | None" = None
_tx_trackers: dict[str, "TransactionTracker"] = {} _tx_trackers: dict[str, "TransactionTracker"] = {}
_tracked_addresses_by_listener: dict[str, set[str]] = {}
def init(self) -> None: def init(self) -> None:
self.create_permanent_task( self.create_permanent_task(
@@ -177,14 +180,17 @@ class TaskManager:
name: str | None = None, name: str | None = None,
) -> Task: ) -> Task:
""" """
Register a callback for onchain address events dispatched by track_address. Register a callback for onchain address events. Only dispatches events
Will call the provided coroutine with an OnchainAddressEvent on each update. for addresses tracked under the same `name` via track_address, e.g. an
extension registering as "ext_satspay" only sees events for addresses
it tracked with that same name. Defaults to the shared "core" listener
if no name is given.
""" """
name = f"{name or uuid.uuid4()}_onchain_address_listener" name = name or "core"
queue: asyncio.Queue[OnchainAddressEvent] = asyncio.Queue() queue: asyncio.Queue[OnchainAddressEvent] = asyncio.Queue()
return self.create_permanent_task( return self.create_permanent_task(
self._onchain_address_listener_worker(func, queue), self._onchain_address_listener_worker(func, queue),
name=name, name=f"{name}{self.ONCHAIN_ADDRESS_LISTENER_SUFFIX}",
onchain_address_queue=queue, onchain_address_queue=queue,
) )
@@ -225,14 +231,24 @@ class TaskManager:
block_queue=queue, block_queue=queue,
) )
def track_address(self, address: str) -> None: def track_address(self, address: str, name: str) -> None:
"""Start tracking a Bitcoin address via Electrum (ref-counted).""" """Start tracking a Bitcoin address via Electrum (ref-counted).
self._get_address_tracker().add(address)
def untrack_address(self, address: str) -> None: `name` identifies the listener (see register_onchain_listener) that
should receive events for this address.
"""
self._get_address_tracker().add(address)
self._tracked_addresses_by_listener.setdefault(name, set()).add(address)
def untrack_address(self, address: str, name: str) -> None:
"""Decrement ref count; remove from shared tracker when last caller leaves.""" """Decrement ref count; remove from shared tracker when last caller leaves."""
if self._address_tracker: if self._address_tracker:
self._address_tracker.remove(address) self._address_tracker.remove(address)
tracked = self._tracked_addresses_by_listener.get(name)
if tracked:
tracked.discard(address)
if not tracked:
self._tracked_addresses_by_listener.pop(name, None)
def _get_address_tracker(self) -> "AddressTracker": def _get_address_tracker(self) -> "AddressTracker":
if self._address_tracker is None: if self._address_tracker is None:
@@ -434,12 +450,18 @@ 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 registered listeners. """Dispatches an onchain address event to listeners tracking that
address under their own name (see track_address).
Per-address WS queue fan-out is handled by AddressTracker itself. Per-address WS queue fan-out is handled by AddressTracker itself.
""" """
for task in self.tasks: for task in self.tasks:
if task.onchain_address_queue: if not task.onchain_address_queue:
continue
if not task.name.endswith(self.ONCHAIN_ADDRESS_LISTENER_SUFFIX):
continue
name = task.name[: -len(self.ONCHAIN_ADDRESS_LISTENER_SUFFIX)]
if event.address in self._tracked_addresses_by_listener.get(name, ()):
task.onchain_address_queue.put_nowait(event) task.onchain_address_queue.put_nowait(event)
async def _dispatch_onchain_tx_event(self, event: OnchainTxEvent) -> None: async def _dispatch_onchain_tx_event(self, event: OnchainTxEvent) -> None: