openvm_deferral_circuit/
def_fn.rs

1use std::{array::from_fn, fmt::Debug};
2
3use openvm_circuit::arch::{
4    deferral::{DeferralResult, DeferralState, InputCommit, InputMapVal, OutputCommit, OutputRaw},
5    VmField,
6};
7use openvm_poseidon2_air::POSEIDON2_WIDTH;
8use openvm_stark_sdk::config::baby_bear_poseidon2::DIGEST_SIZE;
9
10use crate::{poseidon2::DeferralPoseidon2Chip, utils::f_commit_to_bytes};
11
12#[derive(Clone, Debug, derive_new::new)]
13pub struct RawDeferralResult {
14    pub input: InputCommit,
15    pub output_raw: OutputRaw,
16}
17
18#[allow(clippy::type_complexity)]
19pub struct DeferralFn {
20    f: Box<dyn Fn(&[u8]) -> OutputRaw + Send + Sync + 'static>,
21}
22
23impl Debug for DeferralFn {
24    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
25        f.debug_struct("DeferralFn").finish()
26    }
27}
28
29impl DeferralFn {
30    pub fn new<FN: Fn(&[u8]) -> OutputRaw + Send + Sync + 'static>(f: FN) -> Self {
31        Self { f: Box::new(f) }
32    }
33
34    pub fn execute<F: VmField>(
35        &self,
36        input_commit: &InputCommit,
37        state: &mut DeferralState,
38        deferral_idx: u32,
39        hasher: &DeferralPoseidon2Chip<F>,
40    ) -> (OutputCommit, u64) {
41        let value = state.get_input(input_commit);
42        match value {
43            InputMapVal::Raw(input_raw) => {
44                let output_raw = self.f.as_ref()(input_raw);
45                let output_commit = hash_output_raw(hasher, deferral_idx, &output_raw);
46                let output_len = output_raw.len();
47                state.store_output(input_commit, output_commit.clone(), output_raw);
48                (output_commit, output_len as u64)
49            }
50            InputMapVal::Output(output_commit) => {
51                let output_raw = state.get_output(output_commit);
52                (output_commit.clone(), output_raw.len() as u64)
53            }
54        }
55    }
56}
57
58pub fn generate_deferral_results<F: VmField>(
59    raw_results: Vec<RawDeferralResult>,
60    deferral_idx: u32,
61    hasher: &DeferralPoseidon2Chip<F>,
62) -> Vec<DeferralResult> {
63    raw_results
64        .into_iter()
65        .map(|r| {
66            let output_commit = hash_output_raw(hasher, deferral_idx, &r.output_raw);
67            DeferralResult {
68                input: r.input,
69                output_commit,
70                output_raw: r.output_raw,
71            }
72        })
73        .collect()
74}
75
76fn hash_output_raw<F: VmField>(
77    hasher: &DeferralPoseidon2Chip<F>,
78    deferral_idx: u32,
79    output_ref: &[u8],
80) -> OutputCommit {
81    assert!(output_ref.len().is_multiple_of(DIGEST_SIZE));
82
83    let mut state = [F::ZERO; POSEIDON2_WIDTH];
84    state[0] = F::from_u32(deferral_idx);
85    state[1] = F::from_usize(output_ref.len());
86
87    let (lhs, rhs) = state_to_chunks(&state);
88    if output_ref.is_empty() {
89        let res = hasher.perm(&lhs, &rhs, true);
90        return f_commit_to_bytes(&res).to_vec();
91    }
92
93    state[DIGEST_SIZE..].copy_from_slice(&hasher.perm(&lhs, &rhs, false));
94
95    let mut output_chunks = output_ref.chunks_exact(DIGEST_SIZE);
96    let last_chunk = output_chunks.next_back().unwrap();
97
98    for chunk in output_chunks {
99        let f_chunk = chunk.iter().map(|b| F::from_u8(*b)).collect::<Vec<_>>();
100        state[..DIGEST_SIZE].copy_from_slice(&f_chunk);
101        let (lhs, rhs) = state_to_chunks(&state);
102        let capacity = hasher.perm(&lhs, &rhs, false);
103        state[DIGEST_SIZE..].copy_from_slice(&capacity);
104    }
105
106    let (_, rhs) = state_to_chunks(&state);
107    let last_chunk_f = from_fn(|i| F::from_u8(last_chunk[i]));
108    let res = hasher.perm(&last_chunk_f, &rhs, true);
109    f_commit_to_bytes(&res).to_vec()
110}
111
112pub(crate) fn chunks_to_state<F: Copy>(
113    lhs: &[F; DIGEST_SIZE],
114    rhs: &[F; DIGEST_SIZE],
115) -> [F; POSEIDON2_WIDTH] {
116    from_fn(|i| {
117        if i < DIGEST_SIZE {
118            lhs[i]
119        } else {
120            rhs[i - DIGEST_SIZE]
121        }
122    })
123}
124
125pub(crate) fn state_to_chunks<F: Copy>(
126    state: &[F; POSEIDON2_WIDTH],
127) -> ([F; DIGEST_SIZE], [F; DIGEST_SIZE]) {
128    (from_fn(|i| state[i]), from_fn(|i| state[i + DIGEST_SIZE]))
129}