use alloc::vec::Vec;
use miden_core::{
Word, ZERO,
deferred::{DataChunk, DeferredError, Digest, Node, NodeType, PrecompileError, Tag},
};
use super::SystemEventError;
use crate::{AdviceProvider, MemoryError, fast::FastProcessor};
const DEFERRED_PAYLOAD_LO_OFFSET: usize = 1;
const DEFERRED_PAYLOAD_HI_OFFSET: usize = 5;
const DEFERRED_TAG_OFFSET: usize = 9;
const DEFERRED_NODE_DIGEST_OFFSET: usize = 1;
const DATA_TAG_OFFSET: usize = 1;
const DATA_PTR_OFFSET: usize = 5;
const DATA_N_CHUNKS_OFFSET: usize = 6;
const TAG_NUM_ELEMENTS: usize = 4;
const PAYLOAD_BLOCK_NUM_ELEMENTS: usize = 8;
fn payload_node_num_elements(n_blocks: u32) -> usize {
(n_blocks as usize)
.checked_mul(PAYLOAD_BLOCK_NUM_ELEMENTS)
.and_then(|payload_elements| payload_elements.checked_add(TAG_NUM_ELEMENTS))
.unwrap_or(usize::MAX)
}
pub(super) fn handle_deferred_register(
processor: &mut FastProcessor,
) -> Result<(), SystemEventError> {
let lo = processor.stack_get_word(DEFERRED_PAYLOAD_LO_OFFSET);
let hi = processor.stack_get_word(DEFERRED_PAYLOAD_HI_OFFSET);
let tag = Tag::from_word(processor.stack_get_word(DEFERRED_TAG_OFFSET).into());
let block: DataChunk = [lo[0], lo[1], lo[2], lo[3], hi[0], hi[1], hi[2], hi[3]];
let node = match processor.deferred_state().decode(tag)? {
NodeType::Data if tag == Tag::CHUNKS => {
Node::chunks(Vec::from([block])).map_err(PrecompileError::from)?
},
NodeType::Data => Node::value(tag, block).map_err(PrecompileError::from)?,
NodeType::Join => {
let lhs = Digest::new([block[0], block[1], block[2], block[3]]);
let rhs = Digest::new([block[4], block[5], block[6], block[7]]);
Node::join(tag, lhs, rhs).map_err(PrecompileError::from)?
},
NodeType::PairList => {
let lhs = Digest::new([block[0], block[1], block[2], block[3]]);
let rhs = Digest::new([block[4], block[5], block[6], block[7]]);
Node::try_pair_list(tag, vec![(lhs, rhs)]).map_err(PrecompileError::from)?
},
NodeType::True => return Err(PrecompileError::InvalidNode.into()),
};
processor.deferred_state_mut().register(node)?;
Ok(())
}
pub(super) fn handle_deferred_evaluate(
processor: &mut FastProcessor,
) -> Result<(), SystemEventError> {
let canonical_node = evaluate_canonical_node(processor)?;
push_evaluated_payload(&mut processor.advice, &canonical_node)?;
push_evaluated_tag(&mut processor.advice, &canonical_node)?;
Ok(())
}
pub(super) fn handle_deferred_evaluate_tag(
processor: &mut FastProcessor,
) -> Result<(), SystemEventError> {
let canonical_node = evaluate_canonical_node(processor)?;
push_evaluated_tag(&mut processor.advice, &canonical_node)?;
Ok(())
}
pub(super) fn handle_deferred_evaluate_payload(
processor: &mut FastProcessor,
) -> Result<(), SystemEventError> {
let canonical_node = evaluate_canonical_node(processor)?;
push_evaluated_payload(&mut processor.advice, &canonical_node)?;
Ok(())
}
fn evaluate_canonical_node(processor: &mut FastProcessor) -> Result<Node, SystemEventError> {
let digest: Digest = processor.stack_get_word(DEFERRED_NODE_DIGEST_OFFSET);
let canonical_digest = processor.deferred_state_mut().evaluate_digest(digest)?;
processor
.deferred_state()
.get_node(&canonical_digest)
.cloned()
.ok_or(PrecompileError::MissingNode.into())
}
fn push_evaluated_tag(advice: &mut AdviceProvider, node: &Node) -> Result<(), SystemEventError> {
let tag = Word::from(node.tag().as_word());
advice.push_stack_word(&tag)?;
Ok(())
}
fn push_evaluated_payload(
advice: &mut AdviceProvider,
node: &Node,
) -> Result<(), SystemEventError> {
for chunk in node.payload().as_chunks().iter().rev() {
let [lo0, lo1, lo2, lo3, hi0, hi1, hi2, hi3] = *chunk;
advice.push_stack_word(&Word::new([lo0, lo1, lo2, lo3]))?;
advice.push_stack_word(&Word::new([hi0, hi1, hi2, hi3]))?;
}
Ok(())
}
pub(super) fn handle_deferred_register_data(
processor: &mut FastProcessor,
) -> Result<(), SystemEventError> {
let tag = Tag::from_word(processor.stack_get_word(DATA_TAG_OFFSET).into());
let ptr = processor.stack_get(DATA_PTR_OFFSET).as_canonical_u64();
let n_chunks_felt = processor.stack_get(DATA_N_CHUNKS_OFFSET).as_canonical_u64();
let n = u32::try_from(n_chunks_felt).map_err(|_| PrecompileError::InvalidNode)?;
if n == 0 {
return Err(PrecompileError::InvalidNode.into());
}
let node_type = processor.deferred_state().decode(tag)?;
match node_type {
NodeType::Data | NodeType::PairList => {},
NodeType::Join if n == 1 => {},
NodeType::Join | NodeType::True => {
return Err(PrecompileError::InvalidNode.into());
},
}
let num_elements = payload_node_num_elements(n);
let max_deferred_elements = processor.options.max_deferred_elements();
if num_elements > max_deferred_elements {
return Err(PrecompileError::from(DeferredError::DeferredStateTooLarge {
num_elements,
max: max_deferred_elements,
})
.into());
}
if ptr > u32::MAX as u64 {
return Err(MemoryError::AddressOutOfBounds { addr: ptr }.into());
}
if !ptr.is_multiple_of(4) {
return Err(
MemoryError::UnalignedWordAccess { addr: ptr as u32, ctx: processor.ctx }.into()
);
}
let total = 8u64 * n as u64;
let end = ptr
.checked_add(total)
.ok_or(MemoryError::AddressOutOfBounds { addr: u64::MAX })?;
if end > u32::MAX as u64 {
return Err(MemoryError::AddressOutOfBounds { addr: end }.into());
}
let ctx = processor.ctx;
let mut chunks: Vec<DataChunk> = Vec::with_capacity(n as usize);
for k in 0..n {
let base = ptr as u32 + k * 8;
let mut chunk = [ZERO; 8];
for (i, felt) in chunk.iter_mut().enumerate() {
*felt = processor.memory().read_element_impl(ctx, base + i as u32).unwrap_or(ZERO);
}
chunks.push(chunk);
}
let node = match node_type {
NodeType::Data if tag == Tag::CHUNKS => {
Node::chunks(chunks).map_err(PrecompileError::from)?
},
NodeType::Data => Node::try_data(tag, chunks).map_err(PrecompileError::from)?,
NodeType::Join => {
let block = chunks.into_iter().next().ok_or(PrecompileError::InvalidNode)?;
let lhs = Digest::new([block[0], block[1], block[2], block[3]]);
let rhs = Digest::new([block[4], block[5], block[6], block[7]]);
Node::join(tag, lhs, rhs).map_err(PrecompileError::from)?
},
NodeType::PairList => {
Node::try_pair_list_chunks(tag, chunks).map_err(PrecompileError::from)?
},
NodeType::True => unreachable!("TRUE was rejected before memory reads"),
};
processor.deferred_state_mut().register(node)?;
Ok(())
}