use alloc::{
collections::{BTreeMap, BTreeSet},
string::ToString,
sync::Arc,
vec::Vec,
};
use miden_utils_indexing::newtype_id;
use crate::{
Word,
advice::AdviceMap,
mast::{ExecutableMastForest, MastForest, MastNode, MastNodeExt, MastNodeId},
serde::DeserializationError,
utils::Idx,
};
newtype_id!(MastForestId);
#[derive(Debug)]
pub struct SparseMastForest {
nodes: BTreeMap<MastNodeId, MastNode>,
digests: BTreeMap<MastNodeId, Word>,
roots: Vec<MastNodeId>,
advice_map: AdviceMap,
}
impl SparseMastForest {
pub fn nodes(&self) -> &BTreeMap<MastNodeId, MastNode> {
&self.nodes
}
pub fn num_nodes(&self) -> usize {
self.nodes
.keys()
.chain(self.digests.keys())
.chain(self.roots.iter())
.map(|id| id.to_usize() + 1)
.max()
.unwrap_or(0)
}
pub fn procedure_roots(&self) -> &[MastNodeId] {
&self.roots
}
pub fn advice_map(&self) -> &AdviceMap {
&self.advice_map
}
pub(in crate::mast) fn digest_entries(&self) -> &BTreeMap<MastNodeId, Word> {
&self.digests
}
pub(in crate::mast) fn from_serialized_parts(
nodes: Vec<(MastNodeId, MastNode)>,
digests: Vec<(MastNodeId, Word)>,
roots: Vec<MastNodeId>,
advice_map: AdviceMap,
) -> Result<Self, DeserializationError> {
if !advice_map.is_empty() {
return Err(DeserializationError::InvalidValue(
"sparse MAST replay payload must not carry advice map entries".to_string(),
));
}
let nodes = collect_unique_nodes(nodes)?;
let digests = collect_unique_digests(digests)?;
for &root in &roots {
validate_sparse_id(root, "procedure root")?;
}
for node_id in nodes.keys() {
if digests.contains_key(node_id) {
return Err(DeserializationError::InvalidValue(format!(
"sparse full-node id {} overlaps a digest-only entry",
node_id.0
)));
}
}
validate_full_node_child_digests(&nodes, &digests)?;
Ok(Self {
nodes,
digests,
roots,
advice_map: AdviceMap::default(),
})
}
}
fn validate_sparse_id(id: MastNodeId, label: &str) -> Result<(), DeserializationError> {
if id.to_usize() >= MastForest::MAX_NODES {
return Err(DeserializationError::InvalidValue(format!(
"{label} id {} exceeds maximum sparse MAST node id {}",
id.0,
MastForest::MAX_NODES - 1
)));
}
Ok(())
}
fn collect_unique_nodes(
nodes: Vec<(MastNodeId, MastNode)>,
) -> Result<BTreeMap<MastNodeId, MastNode>, DeserializationError> {
let mut result = BTreeMap::new();
for (id, node) in nodes {
validate_sparse_id(id, "full node")?;
if result.insert(id, node).is_some() {
return Err(DeserializationError::InvalidValue(format!(
"duplicate sparse full-node id {}",
id.0
)));
}
}
Ok(result)
}
fn collect_unique_digests(
digests: Vec<(MastNodeId, Word)>,
) -> Result<BTreeMap<MastNodeId, Word>, DeserializationError> {
let mut result = BTreeMap::new();
for (id, digest) in digests {
validate_sparse_id(id, "digest-only node")?;
if result.insert(id, digest).is_some() {
return Err(DeserializationError::InvalidValue(format!(
"duplicate sparse digest-only id {}",
id.0
)));
}
}
Ok(result)
}
fn validate_full_node_child_digests(
nodes: &BTreeMap<MastNodeId, MastNode>,
digests: &BTreeMap<MastNodeId, Word>,
) -> Result<(), DeserializationError> {
for (&node_id, node) in nodes {
validate_sparse_id(node_id, "full node")?;
match node {
MastNode::Block(block) => {
block.validate_batch_invariants().map_err(|error_msg| {
DeserializationError::InvalidValue(format!(
"invalid sparse basic block {}: {error_msg}",
node_id.0
))
})?;
},
MastNode::External(_) | MastNode::Dyn(_) => {},
MastNode::Join(join) => {
require_child_digest(node_id, join.first(), nodes, digests)?;
require_child_digest(node_id, join.second(), nodes, digests)?;
},
MastNode::Split(split) => {
require_child_digest(node_id, split.on_true(), nodes, digests)?;
require_child_digest(node_id, split.on_false(), nodes, digests)?;
},
MastNode::Loop(loop_node) => {
require_child_digest(node_id, loop_node.body(), nodes, digests)?;
},
MastNode::Call(call) => {
require_child_digest(node_id, call.callee(), nodes, digests)?;
},
}
}
Ok(())
}
fn require_child_digest(
parent_id: MastNodeId,
child_id: MastNodeId,
nodes: &BTreeMap<MastNodeId, MastNode>,
digests: &BTreeMap<MastNodeId, Word>,
) -> Result<(), DeserializationError> {
validate_sparse_id(child_id, "child")?;
if !nodes.contains_key(&child_id) && !digests.contains_key(&child_id) {
return Err(DeserializationError::InvalidValue(format!(
"sparse full node {} references child {} without a full node or digest-only entry",
parent_id.0, child_id.0
)));
}
Ok(())
}
impl ExecutableMastForest for SparseMastForest {
#[inline(always)]
fn get_node_by_id(&self, node_id: MastNodeId) -> Option<&MastNode> {
self.nodes.get(&node_id)
}
#[inline(always)]
fn get_digest_by_id(&self, node_id: MastNodeId) -> Option<Word> {
if let Some(node) = self.nodes.get(&node_id) {
return Some(node.digest());
}
self.digests.get(&node_id).copied()
}
#[inline(always)]
fn find_procedure_root(&self, digest: Word) -> Option<MastNodeId> {
self.roots.iter().find_map(|&root_id| {
let node = self.nodes.get(&root_id)?;
(node.digest() == digest).then_some(root_id)
})
}
#[inline(always)]
fn advice_map(&self) -> &AdviceMap {
&self.advice_map
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum VisitKind {
FullVisit,
DigestOnly,
}
#[derive(Debug)]
pub struct SparseMastForestBuilder {
source: Arc<MastForest>,
full_visits: BTreeSet<MastNodeId>,
digest_only_visits: BTreeSet<MastNodeId>,
}
impl SparseMastForestBuilder {
pub fn new(source: Arc<MastForest>) -> Self {
Self {
source,
full_visits: BTreeSet::new(),
digest_only_visits: BTreeSet::new(),
}
}
pub fn record_visit(&mut self, node_id: MastNodeId, kind: VisitKind) {
match kind {
VisitKind::FullVisit => {
self.full_visits.insert(node_id);
},
VisitKind::DigestOnly => {
self.digest_only_visits.insert(node_id);
},
}
}
pub fn source(&self) -> &Arc<MastForest> {
&self.source
}
pub fn finalize(self) -> SparseMastForest {
let SparseMastForestBuilder { source, full_visits, digest_only_visits } = self;
let mut nodes = BTreeMap::new();
for node_id in &full_visits {
let node = source
.get_node_by_id(*node_id)
.expect("recorded full-visit id must exist in source forest");
nodes.insert(*node_id, node.clone());
}
let mut digests = BTreeMap::new();
for node_id in digest_only_visits {
if full_visits.contains(&node_id) {
continue;
}
let node = source
.get_node_by_id(node_id)
.expect("recorded digest-only id must exist in source forest");
digests.insert(node_id, node.digest());
}
SparseMastForest {
nodes,
digests,
roots: source.procedure_roots().to_vec(),
advice_map: AdviceMap::default(),
}
}
}