use std::collections::{HashMap, HashSet};
use chia_protocol::{Bytes32, CoinState};
const MAX_COINS_PER_PUZZLE_HASH: usize = 10_000;
#[derive(Debug, Default)]
pub struct CoinStateCache {
coins: HashMap<Bytes32, CoinState>,
subscribed_coins: HashSet<Bytes32>,
subscribed_puzzle_hashes: HashSet<Bytes32>,
peak: Option<(u32, Bytes32)>,
}
impl CoinStateCache {
pub fn new() -> Self {
Self::default()
}
pub fn peak(&self) -> Option<(u32, Bytes32)> {
self.peak
}
pub fn set_peak(&mut self, height: u32, header_hash: Bytes32) {
if self.peak.is_none_or(|(h, _)| height >= h) {
self.peak = Some((height, header_hash));
}
}
pub fn get(&self, coin_id: Bytes32) -> Option<CoinState> {
self.coins.get(&coin_id).cloned()
}
pub fn seed(&mut self, states: impl IntoIterator<Item = CoinState>) {
for state in states {
self.coins.insert(state.coin.coin_id(), state);
}
}
pub fn track_coins(&mut self, coin_ids: impl IntoIterator<Item = Bytes32>) {
self.subscribed_coins.extend(coin_ids);
}
pub fn track_puzzle_hashes(&mut self, puzzle_hashes: impl IntoIterator<Item = Bytes32>) {
self.subscribed_puzzle_hashes.extend(puzzle_hashes);
}
pub fn untrack_coins(&mut self, coin_ids: &[Bytes32]) {
for id in coin_ids {
self.subscribed_coins.remove(id);
}
}
pub fn is_subscribed_coin(&self, coin_id: Bytes32) -> bool {
self.subscribed_coins.contains(&coin_id)
}
pub fn subscribed_coins(&self) -> Vec<Bytes32> {
self.subscribed_coins.iter().copied().collect()
}
pub fn subscribed_puzzle_hashes(&self) -> Vec<Bytes32> {
self.subscribed_puzzle_hashes.iter().copied().collect()
}
pub fn apply_update(
&mut self,
items: &[CoinState],
height: u32,
fork_height: u32,
peak_hash: Bytes32,
) {
let reasserted: HashSet<Bytes32> = items
.iter()
.filter(|s| self.is_subscribed(s))
.map(|s| s.coin.coin_id())
.collect();
self.coins.retain(|id, state| {
if reasserted.contains(id) {
return true; }
if state.created_height.is_some_and(|h| h > fork_height) {
return false; }
if state.spent_height.is_some_and(|h| h > fork_height) {
state.spent_height = None; }
true
});
let max_coins = self.max_cached_coins();
for state in items {
if !self.is_subscribed(state) {
continue; }
let id = state.coin.coin_id();
if !self.coins.contains_key(&id) && self.coins.len() >= max_coins {
log::warn!("chia-peer cache at cap ({max_coins}); dropping overflow coin state");
continue;
}
self.coins.insert(id, *state);
}
self.set_peak(height, peak_hash);
}
fn is_subscribed(&self, state: &CoinState) -> bool {
self.subscribed_coins.contains(&state.coin.coin_id())
|| self
.subscribed_puzzle_hashes
.contains(&state.coin.puzzle_hash)
}
fn max_cached_coins(&self) -> usize {
self.subscribed_coins.len().saturating_add(
self.subscribed_puzzle_hashes
.len()
.saturating_mul(MAX_COINS_PER_PUZZLE_HASH),
)
}
}
#[cfg(test)]
mod tests {
use super::*;
use chia_protocol::Coin;
fn coin(seed: u8, amount: u64) -> Coin {
Coin::new(
Bytes32::new([seed; 32]),
Bytes32::new([seed ^ 0xff; 32]),
amount,
)
}
fn state(seed: u8, created: Option<u32>, spent: Option<u32>) -> CoinState {
CoinState {
coin: coin(seed, 1),
created_height: created,
spent_height: spent,
}
}
#[test]
fn seed_then_get_returns_state() {
let mut cache = CoinStateCache::new();
let s = state(1, Some(100), None);
let id = s.coin.coin_id();
cache.seed([s]);
assert_eq!(cache.get(id), Some(s));
}
#[test]
fn peak_only_advances() {
let mut cache = CoinStateCache::new();
cache.set_peak(100, Bytes32::new([1; 32]));
cache.set_peak(90, Bytes32::new([2; 32])); assert_eq!(cache.peak().map(|(h, _)| h), Some(100));
}
#[test]
fn subscription_tracking_roundtrips() {
let mut cache = CoinStateCache::new();
let id = coin(5, 1).coin_id();
cache.track_coins([id]);
assert!(cache.is_subscribed_coin(id));
assert_eq!(cache.subscribed_coins(), vec![id]);
cache.untrack_coins(&[id]);
assert!(!cache.is_subscribed_coin(id));
cache.track_puzzle_hashes([Bytes32::new([7; 32])]);
assert_eq!(cache.subscribed_puzzle_hashes().len(), 1);
}
#[test]
fn reorg_update_rolls_cache_back_across_fork() {
let mut cache = CoinStateCache::new();
let x_pre = state(1, Some(80), Some(92));
let x_id = x_pre.coin.coin_id();
let y_pre = state(2, Some(95), None);
let y_id = y_pre.coin.coin_id();
let z_pre = state(3, Some(50), Some(93));
let z_id = z_pre.coin.coin_id();
cache.track_coins([x_id]); cache.seed([x_pre, y_pre, z_pre]);
cache.set_peak(100, Bytes32::new([0xaa; 32]));
let x_post = state(1, Some(80), None);
cache.apply_update(&[x_post], 90, 89, Bytes32::new([0xbb; 32]));
assert_eq!(cache.get(x_id), Some(x_post));
assert_eq!(cache.get(y_id), None);
assert_eq!(cache.get(z_id).and_then(|s| s.spent_height), None);
assert_eq!(cache.peak(), Some((100, Bytes32::new([0xaa; 32]))));
}
#[test]
fn normal_update_inserts_without_dropping_lower_coins() {
let mut cache = CoinStateCache::new();
let old = state(1, Some(10), None);
let old_id = old.coin.coin_id();
cache.seed([old]);
let fresh = state(2, Some(200), None);
let fresh_id = fresh.coin.coin_id();
cache.track_coins([old_id, fresh_id]); cache.apply_update(&[fresh], 200, 199, Bytes32::new([0xcc; 32]));
assert!(
cache.get(old_id).is_some(),
"coin below the fork is retained"
);
assert!(cache.get(fresh_id).is_some(), "new coin is inserted");
}
#[test]
fn unsolicited_update_item_is_dropped_not_cached() {
let mut cache = CoinStateCache::new();
let injected = state(9, Some(10), None);
let injected_id = injected.coin.coin_id();
cache.apply_update(&[injected], 10, 9, Bytes32::new([1; 32]));
assert_eq!(
cache.get(injected_id),
None,
"an unsubscribed coin must never be cached (nor served on a read)"
);
}
#[test]
fn update_item_matching_a_subscribed_puzzle_hash_is_accepted() {
let mut cache = CoinStateCache::new();
let watched = state(4, Some(20), None);
let ph = watched.coin.puzzle_hash;
cache.track_puzzle_hashes([ph]);
cache.apply_update(&[watched], 20, 19, Bytes32::new([2; 32]));
assert_eq!(
cache.get(watched.coin.coin_id()),
Some(watched),
"a coin discovered under a subscribed puzzle hash is accepted"
);
}
}