openvm_deferral_circuit/
def_fn.rs1use 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}