use std::num::NonZeroUsize;
use std::sync::Arc;
use simplicity::node::CoreConstructible;
use super::ProgNode;
use crate::named::CoreExt;
pub fn array_fold<'brand>(
size: NonZeroUsize,
f: &ProgNode<'brand>,
) -> Result<ProgNode<'brand>, simplicity::types::Error> {
fn tree_fold<'brand>(
n: usize,
f_powers_of_two: &[ProgNode<'brand>],
) -> Result<ProgNode<'brand>, simplicity::types::Error> {
let max_pow2 = n.ilog2() as usize;
debug_assert!(max_pow2 < f_powers_of_two.len());
let f_right = &f_powers_of_two[max_pow2];
let size_right = 1 << max_pow2;
if n == size_right {
return Ok(Arc::clone(f_right));
}
debug_assert!(size_right < n);
let f_left = tree_fold(n - size_right, f_powers_of_two)?;
f_array_fold(&f_left, f_right)
}
fn f_array_fold<'brand>(
f_left: &ProgNode<'brand>,
f_right: &ProgNode<'brand>,
) -> Result<ProgNode<'brand>, simplicity::types::Error> {
let ctx = f_left.inference_context();
let left_arr = ProgNode::o().o().h(ctx);
let right_arr = ProgNode::o().i().h(ctx);
let acc = ProgNode::i().h(ctx);
let left_res = left_arr.pair(acc).comp(f_left)?;
let right_res = right_arr.pair(left_res).comp(f_right)?;
Ok(right_res.build())
}
let n = size.get();
let mut f_powers_of_two: Vec<ProgNode> = Vec::with_capacity(1 + n.ilog2() as usize);
let mut f_prev = f.clone();
f_powers_of_two.push(f_prev.clone());
let mut i = 1;
while i < n {
f_prev = f_array_fold(&f_prev, &f_prev)?;
f_powers_of_two.push(Arc::clone(&f_prev));
i *= 2;
}
tree_fold(n, &f_powers_of_two)
}
#[cfg(test)]
mod tests {
use crate::{tests::TestCase, WitnessValues};
#[test]
fn array_fold() {
TestCase::program_file("./examples/array_fold.simf")
.with_witness_values(WitnessValues::default())
.assert_run_success();
}
}