fix only have address updates sent if tracked
This commit is contained in:
+32
-10
@@ -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:
|
||||||
|
|||||||
Reference in New Issue
Block a user