use alloc::{collections::BTreeMap, sync::Arc, vec::Vec};
use super::{
DeferredError, DeferredStateWire, Digest, IntegrityError, Node, NodeType, PrecompileError,
PrecompileRegistry, TRUE_DIGEST, Tag,
};
#[derive(Debug, Clone)]
pub struct DeferredState {
registry: Arc<PrecompileRegistry>,
nodes: BTreeMap<Digest, Node>,
pub(super) root: Digest,
evals: BTreeMap<Digest, Digest>,
remaining_elements: usize,
}
impl Default for DeferredState {
fn default() -> Self {
Self::new(Arc::new(PrecompileRegistry::new()), usize::MAX)
.expect("empty registry initialization cannot fail")
}
}
impl DeferredState {
pub fn new(
registry: Arc<PrecompileRegistry>,
max_elements: usize,
) -> Result<Self, PrecompileError> {
let mut state = Self::empty(registry, max_elements);
state.initialize_precompile_nodes()?;
Ok(state)
}
fn empty(registry: Arc<PrecompileRegistry>, max_elements: usize) -> Self {
let mut nodes = BTreeMap::new();
nodes.insert(TRUE_DIGEST, Node::TRUE);
let mut evals = BTreeMap::new();
evals.insert(TRUE_DIGEST, TRUE_DIGEST);
Self {
registry,
nodes,
root: TRUE_DIGEST,
evals,
remaining_elements: max_elements,
}
}
fn initialize_precompile_nodes(&mut self) -> Result<(), PrecompileError> {
let init_nodes = self.registry.init_nodes();
let init_digests: Vec<Digest> = init_nodes.iter().map(Node::digest).collect();
for node in init_nodes {
self.registry.validate_node(&node)?;
self.insert_node(node)?;
}
for digest in init_digests {
self.evaluate_digest(digest)?;
}
Ok(())
}
pub fn extend_precompiles(
&mut self,
precompiles: PrecompileRegistry,
) -> Result<(), PrecompileError> {
let mut next = self.clone();
Arc::make_mut(&mut next.registry).merge(precompiles);
next.initialize_precompile_nodes()?;
*self = next;
Ok(())
}
pub fn registry(&self) -> &PrecompileRegistry {
&self.registry
}
pub fn root(&self) -> Digest {
self.root
}
pub fn get_node(&self, digest: &Digest) -> Option<&Node> {
self.nodes.get(digest)
}
pub fn get_canonical_digest(&self, digest: Digest) -> Option<Digest> {
let canonical_digest = self.evals.get(&digest).copied()?;
self.nodes.contains_key(&canonical_digest).then_some(canonical_digest)
}
pub fn get_canonical_node(&self, digest: Digest) -> Option<(Digest, &Node)> {
let canonical_digest = self.get_canonical_digest(digest)?;
self.nodes.get(&canonical_digest).map(|node| (canonical_digest, node))
}
pub fn require_canonical_node(
&self,
digest: Digest,
) -> Result<(Digest, &Node), PrecompileError> {
self.get_canonical_node(digest).ok_or(PrecompileError::MissingNode)
}
pub fn nodes(&self) -> &BTreeMap<Digest, Node> {
&self.nodes
}
pub fn remaining_elements(&self) -> usize {
self.remaining_elements
}
pub fn set_max_elements(&mut self, max_elements: usize) {
let used_elements = self
.nodes
.iter()
.filter_map(|(digest, node)| {
(*digest != TRUE_DIGEST).then_some(node.storage_felt_len())
})
.sum::<usize>();
self.remaining_elements = max_elements.saturating_sub(used_elements);
}
pub fn decode(&self, tag: Tag) -> Result<NodeType, PrecompileError> {
self.registry.decode_node_type(tag)
}
pub fn register(&mut self, node: Node) -> Result<Digest, PrecompileError> {
self.validate_node_for_insertion(&node)?;
let digest = self.insert_node(node)?;
self.evaluate_digest(digest)?;
Ok(digest)
}
pub fn log_statement(&mut self, statement_digest: Digest) -> Result<Digest, PrecompileError> {
let prev_root = self.root;
self.require_true_eval(prev_root)?;
self.require_true_eval(statement_digest)?;
let and_node = Node::and(prev_root, statement_digest);
let new_root = and_node.digest();
self.insert_node(and_node)?;
self.root = new_root;
self.record_eval(new_root, Node::TRUE)?;
Ok(new_root)
}
pub fn log_verified_statement(
&mut self,
statement_digest: Digest,
expected_new_root: Digest,
) -> Result<Digest, PrecompileError> {
let actual_new_root = Node::and(self.root, statement_digest).digest();
if actual_new_root != expected_new_root {
return Err(DeferredError::InvalidDeferredRootTransition {
expected: expected_new_root,
actual: actual_new_root,
}
.into());
}
self.log_statement(statement_digest)
}
pub fn evaluate_digest(&mut self, digest: Digest) -> Result<Digest, PrecompileError> {
let node = self.nodes.get(&digest).ok_or(PrecompileError::MissingNode)?.clone();
if let Some(canonical_digest) = self.evals.get(&digest) {
if self.nodes.contains_key(canonical_digest) {
return Ok(*canonical_digest);
}
return Err(PrecompileError::MissingNode);
}
self.validate_node_for_insertion(&node)?;
let canonical = if node.tag() == Tag::TRUE {
Node::TRUE
} else if node.tag() == Tag::AND {
let (lhs, rhs) = node.payload().as_join()?;
for child in [lhs, rhs] {
self.require_true_eval(child)?;
}
Node::TRUE
} else if node.tag() == Tag::CHUNKS {
node
} else {
let registry = Arc::clone(&self.registry);
let mut context = DeferredContext::new(self);
registry.evaluate(&node, &mut context)?
};
self.record_eval(digest, canonical)?;
self.evals.get(&digest).copied().ok_or(PrecompileError::MissingNode)
}
pub fn to_wire(&self) -> Result<DeferredStateWire, IntegrityError> {
DeferredStateWire::from_state(self)
}
pub fn from_wire(
registry: Arc<PrecompileRegistry>,
wire: &DeferredStateWire,
max_elements: usize,
) -> Result<Self, IntegrityError> {
wire.rehydrate(registry, max_elements)
}
fn validate_node_for_insertion(&self, node: &Node) -> Result<NodeType, PrecompileError> {
let node_type = self.registry.validate_node(node)?;
for child in node.children() {
if child != TRUE_DIGEST && !self.nodes.contains_key(&child) {
return Err(PrecompileError::MissingNode);
}
}
Ok(node_type)
}
fn insert_node(&mut self, node: Node) -> Result<Digest, PrecompileError> {
let digest = node.digest();
match self.nodes.get(&digest) {
Some(existing) if existing == &node => Ok(digest),
Some(_) => Err(DeferredError::ConflictingNode.into()),
None => {
let required = node.storage_felt_len();
self.remaining_elements = self.remaining_elements.checked_sub(required).ok_or(
DeferredError::DeferredStateTooLarge {
num_elements: required,
max: self.remaining_elements,
},
)?;
self.nodes.insert(digest, node);
Ok(digest)
},
}
}
fn record_eval(
&mut self,
input_digest: Digest,
canonical: Node,
) -> Result<(), PrecompileError> {
if !self.nodes.contains_key(&input_digest) {
return Err(PrecompileError::MissingNode);
}
self.validate_node_for_insertion(&canonical)?;
let canonical_digest = self.insert_node(canonical)?;
match self.evals.get(&input_digest) {
Some(existing) if *existing == canonical_digest => Ok(()),
Some(_) => Err(DeferredError::ConflictingNode.into()),
None => {
self.evals.insert(input_digest, canonical_digest);
Ok(())
},
}
}
fn require_true_eval(&mut self, digest: Digest) -> Result<(), PrecompileError> {
if self.evaluate_digest(digest)? != TRUE_DIGEST {
return Err(PrecompileError::AssertionFailed);
}
Ok(())
}
}
pub struct DeferredContext<'a> {
state: &'a mut DeferredState,
}
impl<'a> DeferredContext<'a> {
pub(crate) fn new(state: &'a mut DeferredState) -> Self {
Self { state }
}
pub fn get_node(&self, digest: &Digest) -> Option<&Node> {
self.state.get_node(digest)
}
pub fn evaluate_digest(&mut self, digest: Digest) -> Result<Digest, PrecompileError> {
self.state.evaluate_digest(digest)
}
pub fn evaluate_digest_pair(
&mut self,
lhs: Digest,
rhs: Digest,
) -> Result<(Digest, Digest), PrecompileError> {
Ok((self.evaluate_digest(lhs)?, self.evaluate_digest(rhs)?))
}
pub fn ensure_equal(&mut self, lhs: Digest, rhs: Digest) -> Result<(), PrecompileError> {
let (lhs, rhs) = self.evaluate_digest_pair(lhs, rhs)?;
if lhs != rhs {
return Err(PrecompileError::AssertionFailed);
}
Ok(())
}
pub fn register(&mut self, node: Node) -> Result<Digest, PrecompileError> {
self.state.register(node)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{
Felt, ZERO,
deferred::{Payload, Precompile, precompile_id},
};
#[derive(Debug, Clone, Copy)]
struct RejectingPrecompile;
impl Precompile for RejectingPrecompile {
fn name(&self) -> &'static str {
"rejecting-registration-fixture"
}
fn id(&self) -> Felt {
precompile_id(self.name())
}
fn decode(&self, args: [Felt; 3]) -> Option<NodeType> {
(args == [ZERO; 3]).then_some(NodeType::Data)
}
fn evaluate(
&self,
_args: [Felt; 3],
_payload: &Payload,
_context: &mut DeferredContext<'_>,
) -> Result<Node, PrecompileError> {
Err(PrecompileError::AssertionFailed)
}
}
#[test]
fn register_eagerly_propagates_precompile_evaluation_errors() {
let precompile = RejectingPrecompile;
let tag =
Tag::precompile(precompile.id(), [ZERO; 3]).expect("fixture id is precompile-owned");
let registry = Arc::new(PrecompileRegistry::new().with_precompile(precompile));
let mut state = DeferredState::new(registry, usize::MAX).unwrap();
let node = Node::value(tag, [ZERO; 8]).unwrap();
let digest = node.digest();
let error = state.register(node).unwrap_err();
assert!(matches!(error.root(), PrecompileError::AssertionFailed));
assert_eq!(state.get_canonical_digest(digest), None);
}
}