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) {
self.update_peak(height, header_hash, false);
}
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.cache_coin(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();
let mut rolled_back = false;
self.coins.retain(|id, state| {
if reasserted.contains(id) {
return true; }
if state.created_height.is_some_and(|h| h > fork_height) {
rolled_back = true;
return false; }
if state.spent_height.is_some_and(|h| h > fork_height) {
state.spent_height = None; rolled_back = true;
}
true
});
let is_genuine_reorg = rolled_back
&& height >= fork_height
&& self.peak.is_some_and(|(current, _)| fork_height < current);
self.update_peak(height, peak_hash, is_genuine_reorg);
for state in items {
if !self.is_subscribed(state) {
continue; }
self.cache_coin(*state);
}
}
fn cache_coin(&mut self, state: CoinState) {
if let Some(created) = state.created_height {
if self
.peak
.is_some_and(|(peak_height, _)| created > peak_height)
{
return; }
}
let id = state.coin.coin_id();
if !self.coins.contains_key(&id) && self.coins.len() >= self.max_cached_coins() {
log::warn!("chia-peer cache at cap; dropping overflow coin state");
return;
}
self.coins.insert(id, state);
}
fn update_peak(&mut self, height: u32, header_hash: Bytes32, allow_lower: bool) {
let changed = match self.peak {
None => true,
Some((current, _)) => height >= current || allow_lower,
};
if !changed {
return;
}
self.peak = Some((height, header_hash));
self.coins
.retain(|_, state| state.created_height.is_none_or(|h| h <= height));
}
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.track_coins([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, y_id, z_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((90, Bytes32::new([0xbb; 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();
let fresh = state(2, Some(200), None);
let fresh_id = fresh.coin.coin_id();
cache.track_coins([old_id, fresh_id]); cache.seed([old]);
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"
);
}
#[test]
fn genuine_reorg_rollback_lowers_peak() {
let mut cache = CoinStateCache::new();
let orphaned = state(1, Some(95), None); cache.track_coins([orphaned.coin.coin_id()]);
cache.seed([orphaned]);
cache.set_peak(100, Bytes32::new([0xaa; 32]));
cache.apply_update(&[], 92, 90, Bytes32::new([0xbb; 32]));
assert_eq!(cache.peak(), Some((92, Bytes32::new([0xbb; 32]))));
}
#[test]
fn authoritative_reorg_lowers_peak() {
let mut cache = CoinStateCache::new();
let orphaned = state(2, Some(95), None);
cache.track_coins([orphaned.coin.coin_id()]);
cache.seed([orphaned]);
cache.set_peak(100, Bytes32::new([0xaa; 32]));
cache.apply_update(&[], 90, 89, Bytes32::new([0xbb; 32]));
assert_eq!(cache.peak(), Some((90, Bytes32::new([0xbb; 32]))));
}
#[test]
fn empty_update_with_low_fork_does_not_lower_peak() {
let mut cache = CoinStateCache::new();
cache.set_peak(1_000_000, Bytes32::new([0xaa; 32]));
cache.apply_update(&[], 7, 5, Bytes32::new([0xbb; 32]));
assert_eq!(cache.peak(), Some((1_000_000, Bytes32::new([0xaa; 32]))));
}
#[test]
fn peak_height_below_fork_height_is_rejected() {
let mut cache = CoinStateCache::new();
let orphaned = state(3, Some(60), None); cache.track_coins([orphaned.coin.coin_id()]);
cache.seed([orphaned]);
cache.set_peak(100, Bytes32::new([0xaa; 32]));
cache.apply_update(&[], 40, 50, Bytes32::new([0xbb; 32]));
assert_eq!(cache.peak(), Some((100, Bytes32::new([0xaa; 32]))));
}
#[test]
fn bare_new_peak_lower_does_not_lower_peak() {
let mut cache = CoinStateCache::new();
cache.set_peak(100, Bytes32::new([0xaa; 32]));
cache.set_peak(90, Bytes32::new([0xbb; 32])); assert_eq!(cache.peak(), Some((100, Bytes32::new([0xaa; 32]))));
}
#[test]
fn normal_forward_update_advances_peak() {
let mut cache = CoinStateCache::new();
cache.set_peak(100, Bytes32::new([0xaa; 32]));
cache.apply_update(&[], 101, 100, Bytes32::new([0xcc; 32]));
assert_eq!(cache.peak(), Some((101, Bytes32::new([0xcc; 32]))));
}
#[test]
fn forward_update_with_item_created_above_tip_is_refused() {
let mut cache = CoinStateCache::new();
cache.set_peak(1_000_000, Bytes32::new([0xaa; 32]));
let liar = state(1, Some(5_000_000), None);
let liar_id = liar.coin.coin_id();
cache.track_coins([liar_id]);
cache.apply_update(&[liar], 1_000_001, 1_000_000, Bytes32::new([0xbb; 32]));
assert_eq!(
cache.get(liar_id),
None,
"above-tip coin refused on forward path"
);
}
#[test]
fn invariant_no_cached_coin_above_peak_on_forward_path() {
let mut cache = CoinStateCache::new();
cache.set_peak(1_000_000, Bytes32::new([0xaa; 32]));
let honest = state(1, Some(999_999), None); let liar = state(2, Some(5_000_000), None); cache.track_coins([honest.coin.coin_id(), liar.coin.coin_id()]);
cache.apply_update(
&[honest, liar],
1_000_001,
1_000_000,
Bytes32::new([0xbb; 32]),
);
let (peak_height, _) = cache.peak().expect("peak set");
for id in [honest.coin.coin_id(), liar.coin.coin_id()] {
if let Some(cs) = cache.get(id) {
assert!(
cs.created_height.is_none_or(|h| h <= peak_height),
"cached coin {id:?} created above peak {peak_height}"
);
}
}
}
#[test]
fn invariant_no_cached_coin_created_above_peak_after_reorg() {
let mut cache = CoinStateCache::new();
let below = state(1, Some(50), None); let orphaned = state(2, Some(95), None); cache.track_coins([below.coin.coin_id(), orphaned.coin.coin_id()]);
cache.seed([below, orphaned]);
cache.set_peak(100, Bytes32::new([0xaa; 32]));
let above_tip = state(3, Some(99), None);
let above_tip_id = above_tip.coin.coin_id();
cache.track_coins([above_tip_id]);
cache.apply_update(&[above_tip], 92, 90, Bytes32::new([0xbb; 32]));
let (peak_height, _) = cache.peak().expect("peak set");
assert_eq!(peak_height, 92);
assert_eq!(cache.get(above_tip_id), None, "above-tip coin is refused");
for id in [below.coin.coin_id(), orphaned.coin.coin_id(), above_tip_id] {
if let Some(cs) = cache.get(id) {
assert!(
cs.created_height.is_none_or(|h| h <= peak_height),
"cached coin {id:?} created above peak {peak_height}"
);
}
}
}
#[test]
fn reassert_above_new_tip_during_peak_down_is_swept() {
let mut cache = CoinStateCache::new();
let c = state(1, Some(100), None); let d = state(2, Some(150), None); cache.track_coins([c.coin.coin_id(), d.coin.coin_id()]);
cache.seed([c, d]);
cache.set_peak(200, Bytes32::new([0xaa; 32]));
cache.apply_update(&[c], 50, 40, Bytes32::new([0xbb; 32]));
assert_eq!(cache.peak(), Some((50, Bytes32::new([0xbb; 32]))));
assert_eq!(
cache.get(c.coin.coin_id()),
None,
"re-asserted coin above the lowered peak must be swept"
);
}
#[test]
fn seed_above_peak_is_refused_or_swept() {
let mut cache = CoinStateCache::new();
cache.set_peak(100, Bytes32::new([0xaa; 32]));
let high = state(1, Some(200), None);
cache.track_coins([high.coin.coin_id()]);
cache.seed([high]);
assert_eq!(
cache.get(high.coin.coin_id()),
None,
"refused at add boundary"
);
let mut cache = CoinStateCache::new();
let early = state(2, Some(200), None);
cache.track_coins([early.coin.coin_id()]);
cache.seed([early]); assert!(cache.get(early.coin.coin_id()).is_some());
cache.set_peak(100, Bytes32::new([0xaa; 32])); assert_eq!(
cache.get(early.coin.coin_id()),
None,
"seeded coin above the first peak must be swept"
);
}
#[test]
fn property_invariant_holds_across_random_op_sequences() {
use rand::{rngs::StdRng, Rng, SeedableRng};
fn random_state(rng: &mut StdRng) -> CoinState {
let seed: u8 = rng.gen_range(0..8); let created = rng.gen_bool(0.85).then(|| rng.gen_range(0..1_000u32));
let spent = rng.gen_bool(0.3).then(|| rng.gen_range(0..1_000u32));
CoinState {
coin: Coin::new(Bytes32::new([seed; 32]), Bytes32::new([seed ^ 0xff; 32]), 1),
created_height: created,
spent_height: spent,
}
}
const SEQUENCES: usize = 3_000;
const OPS_PER_SEQUENCE: usize = 8;
let mut rng = StdRng::seed_from_u64(1311);
for _ in 0..SEQUENCES {
let mut cache = CoinStateCache::new();
for _ in 0..OPS_PER_SEQUENCE {
match rng.gen_range(0..3) {
0 => {
let s = random_state(&mut rng);
cache.track_coins([s.coin.coin_id()]);
cache.seed([s]);
}
1 => {
let items: Vec<CoinState> = (0..rng.gen_range(0..4))
.map(|_| random_state(&mut rng))
.collect();
for it in &items {
cache.track_coins([it.coin.coin_id()]);
}
let height = rng.gen_range(0..1_000u32);
let fork = rng.gen_range(0..1_000u32);
cache.apply_update(&items, height, fork, Bytes32::new([rng.gen(); 32]));
}
_ => cache.set_peak(rng.gen_range(0..1_000u32), Bytes32::new([rng.gen(); 32])),
}
if let Some((peak_height, _)) = cache.peak() {
for state in cache.coins.values() {
assert!(
state.created_height.is_none_or(|h| h <= peak_height),
"invariant violated: coin created {:?} > peak {peak_height}",
state.created_height
);
}
}
}
}
}
}