use std::collections::{BTreeMap, HashMap, VecDeque};
use std::sync::Arc;
use serde::{Deserialize, Serialize};
use super::bundle::{ReplayBundle, Transition};
use super::generator::Corpus;
use super::verifier::Bytes;
pub type StateHash = [u8; 32];
fn hash_of(bytes: &[u8]) -> StateHash {
*blake3::hash(bytes).as_bytes()
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
pub enum Stratum {
Earliest,
Recent,
Reservoir,
Largest,
Smallest,
}
impl Stratum {
pub const ALL: &'static [Stratum] = &[
Stratum::Earliest,
Stratum::Recent,
Stratum::Reservoir,
Stratum::Largest,
Stratum::Smallest,
];
pub fn as_str(self) -> &'static str {
match self {
Stratum::Earliest => "earliest",
Stratum::Recent => "recent",
Stratum::Reservoir => "reservoir",
Stratum::Largest => "largest",
Stratum::Smallest => "smallest",
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct SamplerConfig {
pub earliest: usize,
pub recent: usize,
pub reservoir: usize,
pub largest: usize,
pub smallest: usize,
pub transitions: usize,
pub max_bytes: usize,
pub max_state_bytes: usize,
pub seed: u64,
}
impl Default for SamplerConfig {
fn default() -> Self {
Self {
earliest: 4,
recent: 8,
reservoir: 8,
largest: 4,
smallest: 4,
transitions: 8,
max_bytes: 4 * 1024 * 1024,
max_state_bytes: 1024 * 1024,
seed: 0,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct TransitionRecord {
pub base: StateHash,
pub result: StateHash,
pub incoming_state: Option<StateHash>,
pub delta: Option<Vec<u8>>,
pub summary: Option<Vec<u8>>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Admission {
Stored,
Duplicate,
TooLarge,
NoBudget,
NotSelected,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct ContractSampler {
config: SamplerConfig,
blobs: BTreeMap<StateHash, Vec<u8>>,
earliest: Vec<StateHash>,
recent: Vec<StateHash>,
reservoir: Vec<StateHash>,
largest: Vec<StateHash>,
smallest: Vec<StateHash>,
transitions: VecDeque<TransitionRecord>,
distinct_seen: u64,
total_seen: u64,
}
impl ContractSampler {
pub fn new(config: SamplerConfig) -> Self {
Self {
config,
blobs: BTreeMap::new(),
earliest: Vec::new(),
recent: Vec::new(),
reservoir: Vec::new(),
largest: Vec::new(),
smallest: Vec::new(),
transitions: VecDeque::new(),
distinct_seen: 0,
total_seen: 0,
}
}
pub fn config(&self) -> &SamplerConfig {
&self.config
}
pub fn stored_bytes(&self) -> usize {
self.blobs.values().map(Vec::len).sum::<usize>() + self.transition_payload_bytes()
}
fn transition_payload_bytes(&self) -> usize {
self.transitions
.iter()
.map(|t| t.delta.as_ref().map_or(0, Vec::len) + t.summary.as_ref().map_or(0, Vec::len))
.sum()
}
pub fn distinct_states(&self) -> usize {
self.blobs.len()
}
pub fn distinct_seen(&self) -> u64 {
self.distinct_seen
}
pub fn total_seen(&self) -> u64 {
self.total_seen
}
pub fn members(&self, stratum: Stratum) -> &[StateHash] {
match stratum {
Stratum::Earliest => &self.earliest,
Stratum::Reservoir => &self.reservoir,
Stratum::Largest => &self.largest,
Stratum::Smallest => &self.smallest,
Stratum::Recent => &self.recent,
}
}
pub fn observe_state(&mut self, state: &[u8]) -> Admission {
self.total_seen += 1;
if state.len() > self.config.max_state_bytes {
return Admission::TooLarge;
}
let hash = hash_of(state);
if self.blobs.contains_key(&hash) {
self.touch_recent(hash);
return Admission::Duplicate;
}
self.distinct_seen += 1;
let wanted = self.strata_for(state.len(), hash);
if wanted.is_empty() {
return Admission::NotSelected;
}
if !self.make_room_for(state.len()) {
return Admission::NoBudget;
}
self.blobs.insert(hash, state.to_vec());
for stratum in wanted {
self.admit_to(stratum, hash, state.len());
}
self.collect_garbage();
Admission::Stored
}
pub fn observe_transition(
&mut self,
base: &[u8],
incoming_state: Option<&[u8]>,
delta: Option<&[u8]>,
summary: Option<&[u8]>,
result: &[u8],
) -> Admission {
let base_admission = self.observe_state(base);
let result_admission = self.observe_state(result);
let incoming_admission = incoming_state.map(|incoming| self.observe_state(incoming));
let kept = |a: Admission| matches!(a, Admission::Stored | Admission::Duplicate);
if !kept(base_admission)
|| !kept(result_admission)
|| incoming_admission.is_some_and(|a| !kept(a))
{
if [
Some(base_admission),
Some(result_admission),
incoming_admission,
]
.into_iter()
.flatten()
.any(|a| matches!(a, Admission::TooLarge))
{
return Admission::TooLarge;
}
return Admission::NotSelected;
}
let payload_bytes = delta.map_or(0, <[u8]>::len) + summary.map_or(0, <[u8]>::len);
if payload_bytes > self.config.max_state_bytes {
return Admission::TooLarge;
}
let record = TransitionRecord {
base: hash_of(base),
result: hash_of(result),
incoming_state: incoming_state.map(hash_of),
delta: delta.map(<[u8]>::to_vec),
summary: summary.map(<[u8]>::to_vec),
};
if self.transitions.contains(&record) {
return Admission::Duplicate;
}
self.transitions.push_back(record);
while self.transitions.len() > self.config.transitions {
self.transitions.pop_front();
}
while self.stored_bytes() > self.config.max_bytes {
if !self.drop_lowest_value() {
break;
}
}
self.collect_garbage();
Admission::Stored
}
pub fn corpus(&self) -> Corpus {
let mut ordered: Vec<StateHash> = Vec::new();
let mut round = 0usize;
loop {
let mut produced = false;
for stratum in Stratum::ALL {
let members = self.members(*stratum);
if let Some(hash) = members.get(round) {
produced = true;
if !ordered.contains(hash) {
ordered.push(*hash);
}
}
}
if !produced {
break;
}
round += 1;
}
let states: Vec<Bytes> = ordered
.iter()
.filter_map(|h| self.blobs.get(h))
.map(|b| Arc::from(b.as_slice()))
.collect();
let mut deltas: Vec<Bytes> = Vec::new();
let mut delta_bases: Vec<Option<Bytes>> = Vec::new();
for record in &self.transitions {
let Some(delta) = record.delta.as_ref() else {
continue;
};
deltas.push(Arc::from(delta.as_slice()));
delta_bases.push(
self.materialize(record)
.map(|m| Arc::from(m.base_state.as_slice())),
);
}
let summaries: Vec<Bytes> = self
.transitions
.iter()
.filter_map(|t| t.summary.as_ref())
.map(|s| Arc::from(s.as_slice()))
.collect();
let transitions: Vec<(Bytes, Bytes)> = self
.transitions
.iter()
.filter_map(|record| {
let materialized = self.materialize(record)?;
Some((
Arc::from(materialized.base_state.as_slice()),
Arc::from(materialized.result_state.as_slice()),
))
})
.collect();
Corpus {
states,
deltas,
delta_bases,
summaries,
transitions,
..Default::default()
}
.deduplicated()
}
pub fn to_bundle(
&self,
code: Option<Vec<u8>>,
code_hash: Option<[u8; 32]>,
parameters: Vec<u8>,
) -> ReplayBundle {
let code_hash = code
.as_ref()
.map(|c| *blake3::hash(c).as_bytes())
.or(code_hash);
let corpus = self.corpus();
ReplayBundle {
schema_version: super::bundle::BUNDLE_SCHEMA_VERSION,
code,
code_hash,
parameters,
instance: None,
states: corpus.states.iter().map(|s| s.to_vec()).collect(),
deltas: corpus
.deltas
.iter()
.enumerate()
.filter(|(i, _)| corpus.delta_base(*i).is_none())
.map(|(_, d)| d.to_vec())
.collect(),
summaries: corpus.summaries.iter().map(|s| s.to_vec()).collect(),
transitions: self
.transitions
.iter()
.filter_map(|t| self.materialize(t))
.collect(),
related: Vec::new(),
note: None,
}
}
fn materialize(&self, record: &TransitionRecord) -> Option<Transition> {
Some(Transition {
base_state: self.blobs.get(&record.base)?.clone(),
result_state: self.blobs.get(&record.result)?.clone(),
incoming_state: record
.incoming_state
.and_then(|h| self.blobs.get(&h).cloned()),
delta: record.delta.clone(),
summary: record.summary.clone(),
})
}
fn touch_recent(&mut self, hash: StateHash) {
if self.config.recent == 0 {
return;
}
if let Some(pos) = self.recent.iter().position(|h| *h == hash) {
self.recent.remove(pos);
}
self.recent.push(hash);
while self.recent.len() > self.config.recent {
self.recent.remove(0);
}
}
fn strata_for(&self, len: usize, hash: StateHash) -> Vec<Stratum> {
let mut wanted = Vec::new();
if self.earliest.len() < self.config.earliest {
wanted.push(Stratum::Earliest);
}
if self.config.recent > 0 {
wanted.push(Stratum::Recent);
}
if self.reservoir_admits(hash) {
wanted.push(Stratum::Reservoir);
}
if self.largest.len() < self.config.largest
|| self
.largest
.last()
.and_then(|h| self.blobs.get(h))
.is_some_and(|b| b.len() < len)
{
wanted.push(Stratum::Largest);
}
if self.smallest.len() < self.config.smallest
|| self
.smallest
.last()
.and_then(|h| self.blobs.get(h))
.is_some_and(|b| b.len() > len)
{
wanted.push(Stratum::Smallest);
}
wanted
}
fn reservoir_admits(&self, _hash: StateHash) -> bool {
if self.config.reservoir == 0 {
return false;
}
if self.reservoir.len() < self.config.reservoir {
return true;
}
let n = self.distinct_seen.max(1);
splitmix64(self.config.seed ^ n) % n < self.config.reservoir as u64
}
fn admit_to(&mut self, stratum: Stratum, hash: StateHash, len: usize) {
match stratum {
Stratum::Earliest => self.earliest.push(hash),
Stratum::Recent => {
self.recent.push(hash);
while self.recent.len() > self.config.recent {
self.recent.remove(0);
}
}
Stratum::Reservoir => {
if self.reservoir.len() < self.config.reservoir {
self.reservoir.push(hash);
} else {
let n = self.distinct_seen.max(1);
let victim = (splitmix64(self.config.seed ^ n.wrapping_mul(3)) as usize)
% self.reservoir.len();
self.reservoir[victim] = hash;
}
}
Stratum::Largest => {
self.largest.push(hash);
let blobs = &self.blobs;
self.largest
.sort_by_key(|h| std::cmp::Reverse(blobs.get(h).map_or(0, Vec::len)));
self.largest.truncate(self.config.largest);
}
Stratum::Smallest => {
self.smallest.push(hash);
let blobs = &self.blobs;
self.smallest
.sort_by_key(|h| blobs.get(h).map_or(usize::MAX, Vec::len));
self.smallest.truncate(self.config.smallest);
}
}
let _ = len;
}
fn make_room_for(&mut self, incoming: usize) -> bool {
if incoming > self.config.max_bytes {
return false;
}
while self.stored_bytes() + incoming > self.config.max_bytes {
let freed = self.drop_lowest_value();
if !freed {
return false;
}
}
true
}
fn drop_lowest_value(&mut self) -> bool {
if self.collect_garbage() {
return true;
}
if self.recent.len() > 1 {
self.recent.remove(0);
self.collect_garbage();
return true;
}
if self.transitions.pop_front().is_some() {
self.collect_garbage();
return true;
}
if self.reservoir.len() > 1 {
self.reservoir.pop();
self.collect_garbage();
return true;
}
if self.largest.len() > 1 {
self.largest.pop();
self.collect_garbage();
return true;
}
if self.smallest.len() > 1 {
self.smallest.pop();
self.collect_garbage();
return true;
}
if self.earliest.len() > 1 {
self.earliest.pop();
self.collect_garbage();
return true;
}
false
}
fn collect_garbage(&mut self) -> bool {
let mut referenced: HashMap<StateHash, ()> = HashMap::new();
for stratum in Stratum::ALL {
for hash in self.members(*stratum) {
referenced.insert(*hash, ());
}
}
for hash in &self.recent {
referenced.insert(*hash, ());
}
for transition in &self.transitions {
referenced.insert(transition.base, ());
referenced.insert(transition.result, ());
if let Some(incoming) = transition.incoming_state {
referenced.insert(incoming, ());
}
}
let before = self.blobs.len();
self.blobs.retain(|hash, _| referenced.contains_key(hash));
self.blobs.len() < before
}
}
fn splitmix64(mut x: u64) -> u64 {
x = x.wrapping_add(0x9E37_79B9_7F4A_7C15);
let mut z = x;
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
z ^ (z >> 31)
}