Files
chia-blockchain/chia/full_node/subscriptions.py
04eef678e9 Bump chia rs 0.28 (#19891)
* bump chia_rs to 0.29

* correct return value from get_spends_for_trusted_block()

* remove unnecessary casts

* adjust test for the new error messages in chia_rs

* Update chia/full_node/full_node_rpc_api.py

Co-authored-by: Kyle Altendorf <sda@fstab.net>

---------

Co-authored-by: Kyle Altendorf <sda@fstab.net>
2025-08-05 11:20:27 -07:00

246 lines
8.4 KiB
Python

from __future__ import annotations
import logging
from dataclasses import dataclass, field
from chia_rs import Coin, SpendBundleConditions
from chia_rs.sized_bytes import bytes32
from chia_rs.sized_ints import uint64
log = logging.getLogger(__name__)
@dataclass(frozen=True)
class SubscriptionSet:
_subscriptions_for_peer: dict[bytes32, set[bytes32]] = field(default_factory=dict, init=False)
_peers_for_subscription: dict[bytes32, set[bytes32]] = field(default_factory=dict, init=False)
def add_subscription(self, peer_id: bytes32, item: bytes32) -> bool:
peers = self._peers_for_subscription.setdefault(item, set())
if peer_id in peers:
return False
subscriptions = self._subscriptions_for_peer.setdefault(peer_id, set())
subscriptions.add(item)
peers.add(peer_id)
return True
def remove_subscription(self, peer_id: bytes32, item: bytes32) -> bool:
subscriptions = self._subscriptions_for_peer.get(peer_id)
if subscriptions is None or item not in subscriptions:
return False
peers = self._peers_for_subscription[item]
peers.remove(peer_id)
subscriptions.remove(item)
if len(subscriptions) == 0:
self._subscriptions_for_peer.pop(peer_id)
if len(peers) == 0:
self._peers_for_subscription.pop(item)
return True
def has_subscription(self, item: bytes32) -> bool:
return item in self._peers_for_subscription
def count_subscriptions(self, peer_id: bytes32) -> int:
return len(self._subscriptions_for_peer.get(peer_id, {}))
def remove_peer(self, peer_id: bytes32) -> None:
for item in self._subscriptions_for_peer.pop(peer_id, {}):
self._peers_for_subscription[item].remove(peer_id)
if len(self._peers_for_subscription[item]) == 0:
self._peers_for_subscription.pop(item)
def subscriptions(self, peer_id: bytes32) -> set[bytes32]:
return self._subscriptions_for_peer.get(peer_id, set())
def peers(self, item: bytes32) -> set[bytes32]:
return self._peers_for_subscription.get(item, set())
def total_count(self) -> int:
return len(self._peers_for_subscription)
@dataclass(frozen=True)
class PeerSubscriptions:
_puzzle_subscriptions: SubscriptionSet = field(default_factory=SubscriptionSet)
_coin_subscriptions: SubscriptionSet = field(default_factory=SubscriptionSet)
def has_puzzle_subscription(self, puzzle_hash: bytes32) -> bool:
return self._puzzle_subscriptions.has_subscription(puzzle_hash)
def has_coin_subscription(self, coin_id: bytes32) -> bool:
return self._coin_subscriptions.has_subscription(coin_id)
def peer_subscription_count(self, peer_id: bytes32) -> int:
puzzle_subscriptions = self._puzzle_subscriptions.count_subscriptions(peer_id)
coin_subscriptions = self._coin_subscriptions.count_subscriptions(peer_id)
return puzzle_subscriptions + coin_subscriptions
def add_puzzle_subscriptions(self, peer_id: bytes32, puzzle_hashes: list[bytes32], max_items: int) -> set[bytes32]:
"""
Adds subscriptions until max_items is reached. Filters out duplicates and returns all additions.
"""
subscription_count = self.peer_subscription_count(peer_id)
added: set[bytes32] = set()
def limit_reached() -> set[bytes32]:
log.info(
"Peer %s attempted to exceed the subscription limit while adding puzzle subscriptions.",
peer_id,
)
return added
# If the subscription limit is reached, bail.
if subscription_count >= max_items:
return limit_reached()
# Decrement this counter to know if we've hit the subscription limit.
subscriptions_left = max_items - subscription_count
for puzzle_hash in puzzle_hashes:
if not self._puzzle_subscriptions.add_subscription(peer_id, puzzle_hash):
continue
subscriptions_left -= 1
added.add(puzzle_hash)
if subscriptions_left == 0:
return limit_reached()
return added
def add_coin_subscriptions(self, peer_id: bytes32, coin_ids: list[bytes32], max_items: int) -> set[bytes32]:
"""
Adds subscriptions until max_items is reached. Filters out duplicates and returns all additions.
"""
subscription_count = self.peer_subscription_count(peer_id)
added: set[bytes32] = set()
def limit_reached() -> set[bytes32]:
log.info(
"Peer %s attempted to exceed the subscription limit while adding coin subscriptions.",
peer_id,
)
return added
# If the subscription limit is reached, bail.
if subscription_count >= max_items:
return limit_reached()
# Decrement this counter to know if we've hit the subscription limit.
subscriptions_left = max_items - subscription_count
for coin_id in coin_ids:
if not self._coin_subscriptions.add_subscription(peer_id, coin_id):
continue
subscriptions_left -= 1
added.add(coin_id)
if subscriptions_left == 0:
return limit_reached()
return added
def remove_puzzle_subscriptions(self, peer_id: bytes32, puzzle_hashes: list[bytes32]) -> set[bytes32]:
"""
Removes subscriptions. Filters out duplicates and returns all removals.
"""
removed: set[bytes32] = set()
for puzzle_hash in puzzle_hashes:
if not self._puzzle_subscriptions.remove_subscription(peer_id, puzzle_hash):
continue
removed.add(puzzle_hash)
return removed
def remove_coin_subscriptions(self, peer_id: bytes32, coin_ids: list[bytes32]) -> set[bytes32]:
"""
Removes subscriptions. Filters out duplicates and returns all removals.
"""
removed: set[bytes32] = set()
for coin_id in coin_ids:
if not self._coin_subscriptions.remove_subscription(peer_id, coin_id):
continue
removed.add(coin_id)
return removed
def clear_puzzle_subscriptions(self, peer_id: bytes32) -> None:
self._puzzle_subscriptions.remove_peer(peer_id)
def clear_coin_subscriptions(self, peer_id: bytes32) -> None:
self._coin_subscriptions.remove_peer(peer_id)
def remove_peer(self, peer_id: bytes32) -> None:
self._puzzle_subscriptions.remove_peer(peer_id)
self._coin_subscriptions.remove_peer(peer_id)
def coin_subscriptions(self, peer_id: bytes32) -> set[bytes32]:
return self._coin_subscriptions.subscriptions(peer_id)
def puzzle_subscriptions(self, peer_id: bytes32) -> set[bytes32]:
return self._puzzle_subscriptions.subscriptions(peer_id)
def peers_for_coin_id(self, coin_id: bytes32) -> set[bytes32]:
return self._coin_subscriptions.peers(coin_id)
def peers_for_puzzle_hash(self, puzzle_hash: bytes32) -> set[bytes32]:
return self._puzzle_subscriptions.peers(puzzle_hash)
def coin_subscription_count(self) -> int:
return self._coin_subscriptions.total_count()
def puzzle_subscription_count(self) -> int:
return self._puzzle_subscriptions.total_count()
def peers_for_spend_bundle(
peer_subscriptions: PeerSubscriptions, conds: SpendBundleConditions, hints_for_removals: set[bytes32]
) -> set[bytes32]:
"""
Returns a list of peer ids that are subscribed to any of the created or
spent coins, puzzle hashes, or hints in the spend bundle. To avoid repeated
lookups, `hints_for_removals` should be a set of all puzzle hashes that are being removed.
"""
coin_ids: set[bytes32] = set()
puzzle_hashes: set[bytes32] = hints_for_removals.copy()
for spend in conds.spends:
coin_ids.add(bytes32(spend.coin_id))
puzzle_hashes.add(spend.puzzle_hash)
for puzzle_hash, amount, memo in spend.create_coin:
coin_ids.add(Coin(spend.coin_id, puzzle_hash, uint64(amount)).name())
puzzle_hashes.add(puzzle_hash)
if memo is not None and len(memo) == 32:
puzzle_hashes.add(bytes32(memo))
peers: set[bytes32] = set()
for coin_id in coin_ids:
peers |= peer_subscriptions.peers_for_coin_id(coin_id)
for puzzle_hash in puzzle_hashes:
peers |= peer_subscriptions.peers_for_puzzle_hash(puzzle_hash)
return peers