use crate::{bits::ToBytesGadget, integers::uint::UInt8, traits::algorithms::MaskedCRHGadget};
use snarkvm_algorithms::traits::CRH;
use snarkvm_fields::PrimeField;
use snarkvm_r1cs::{errors::SynthesisError, ConstraintSystem};
pub fn compute_masked_root<
H: CRH,
HG: MaskedCRHGadget<H, F>,
F: PrimeField,
TB: ToBytesGadget<F>,
CS: ConstraintSystem<F>,
>(
mut cs: CS,
parameters: &HG,
mask_parameters: &HG::MaskParametersGadget,
mask: &TB,
leaves: &[HG::OutputGadget],
) -> Result<HG::OutputGadget, SynthesisError> {
let mask_bytes = mask.to_bytes(cs.ns(|| "mask to bytes"))?;
let mut current_leaves = leaves.to_vec();
let mut level = 0;
while current_leaves.len() != 1 {
current_leaves = current_leaves
.chunks(2)
.enumerate()
.map(|(i, left_right)| {
let inner_hash = hash_inner_node_gadget::<H, HG, F, _, _>(
cs.ns(|| format!("hash left right {} on level {}", i, level)),
parameters,
&left_right[0],
&left_right[1],
mask_parameters,
mask_bytes.clone(),
);
inner_hash
})
.collect::<Result<Vec<_>, _>>()?;
level += 1;
}
let computed_root = current_leaves[0].clone();
Ok(computed_root)
}
pub(crate) fn hash_inner_node_gadget<H, HG, F, TB, CS>(
mut cs: CS,
parameters: &HG,
left_child: &TB,
right_child: &TB,
mask_parameters: &HG::MaskParametersGadget,
mask: Vec<UInt8>,
) -> Result<HG::OutputGadget, SynthesisError>
where
F: PrimeField,
CS: ConstraintSystem<F>,
H: CRH,
HG: MaskedCRHGadget<H, F>,
TB: ToBytesGadget<F>,
{
let left_bytes = left_child.to_bytes(&mut cs.ns(|| "left_to_bytes"))?;
let right_bytes = right_child.to_bytes(&mut cs.ns(|| "right_to_bytes"))?;
let bytes = [left_bytes, right_bytes].concat();
parameters.check_evaluation_gadget_masked(cs, bytes, mask_parameters, mask)
}