tasm_lib/verifier/fri/
derive_from_stark.rs1use triton_vm::prelude::*;
2use triton_vm::table::NUM_RANDOMIZED_QUOTIENT_SEGMENTS;
3
4use crate::arithmetic::bfe::primitive_root_of_unity::PrimitiveRootOfUnity;
5use crate::arithmetic::u32::next_power_of_two::NextPowerOfTwo;
6use crate::prelude::*;
7use crate::verifier::fri::verify::fri_verify_type;
8
9#[derive(Debug, Copy, Clone, Eq, PartialEq, Hash)]
12pub struct DeriveFriFromStark {
13 pub stark: Stark,
14}
15
16impl DeriveFriFromStark {
17 fn derive_fri_field_values(&self, library: &mut Library) -> Vec<LabelledInstruction> {
18 let next_power_of_two = library.import(Box::new(NextPowerOfTwo));
19 let domain_generator = library.import(Box::new(PrimitiveRootOfUnity));
20
21 let num_trace_randomizers = self.stark.num_trace_randomizers;
22 let fri_expansion_factor = self.stark.fri_expansion_factor;
23
24 let min_randomized_trace_len = (2 * num_trace_randomizers + 1)
28 .max((num_trace_randomizers + 1) * NUM_RANDOMIZED_QUOTIENT_SEGMENTS);
29 let interpolant_codeword_length_code = triton_asm!(
30 addi {num_trace_randomizers}
33 push {min_randomized_trace_len}
38 dup 1
39 lt
40 dup 0
43 push {min_randomized_trace_len}
44 mul
45 place 2
48 push -1
51 mul
52 addi 1
53 mul
56 add
57 call {next_power_of_two}
60 );
63 let fri_domain_length = triton_asm!(
64 push {fri_expansion_factor}
66 mul
67 );
70
71 let domain_offset = BFieldElement::generator();
72 let num_collinearity_checks = self.stark.num_collinearity_checks;
73 let expansion_factor = self.stark.fri_expansion_factor;
74 triton_asm!(
75 {&interpolant_codeword_length_code}
78 {&fri_domain_length}
79 push {num_collinearity_checks}
82 push {expansion_factor}
85 swap 2
88 push {domain_offset}
91 dup 1
94 split
95 call {domain_generator}
96 )
98 }
99}
100
101impl BasicSnippet for DeriveFriFromStark {
102 fn parameters(&self) -> Vec<(DataType, String)> {
103 vec![(DataType::U32, "padded_height".to_owned())]
104 }
105
106 fn return_values(&self) -> Vec<(DataType, String)> {
107 vec![(
108 DataType::StructRef(fri_verify_type()),
109 "*fri_verify".to_owned(),
110 )]
111 }
112
113 fn entrypoint(&self) -> String {
114 "tasmlib_verifier_fri_derive_from_stark".to_owned()
115 }
116
117 fn code(&self, library: &mut Library) -> Vec<LabelledInstruction> {
118 let entrypoint = self.entrypoint();
119 let derive_fri_field_values = self.derive_fri_field_values(library);
120 let dyn_malloc = library.import(Box::new(DynMalloc));
121
122 triton_asm!(
123 {entrypoint}:
124 {&derive_fri_field_values}
127 call {dyn_malloc}
130 write_mem 5
133 push -5
136 add
137 return
140 )
141 }
142}
143
144#[cfg(test)]
145mod tests {
146 use super::*;
147 use crate::U32_TO_USIZE_ERR;
148 use crate::rust_shadowing_helper_functions;
149 use crate::test_prelude::*;
150 use crate::verifier::fri::verify::FriVerify;
151
152 #[macro_rules_attr::apply(test)]
153 fn fri_param_derivation_default_stark_pbt() {
154 ShadowedFunction::new(DeriveFriFromStark {
155 stark: Stark::default(),
156 })
157 .test();
158 }
159
160 #[macro_rules_attr::apply(proptest(cases = 10))]
161 fn fri_param_derivation_pbt_pbt(#[strategy(arb())] stark: Stark) {
162 ShadowedFunction::new(DeriveFriFromStark { stark }).test();
163 }
164
165 impl Function for DeriveFriFromStark {
166 fn rust_shadow(
167 &self,
168 stack: &mut Vec<BFieldElement>,
169 memory: &mut HashMap<BFieldElement, BFieldElement>,
170 ) -> Result<(), RustShadowError> {
171 let padded_height: u32 = stack
172 .pop()
173 .ok_or(RustShadowError::StackUnderflow)?
174 .try_into()
175 .map_err(|_| RustShadowError::U64ToU32Error)?;
176 let fri_from_tvm = self
177 .stark
178 .fri(padded_height.try_into().expect(U32_TO_USIZE_ERR))
179 .map_err(|_| RustShadowError::Other)?;
180 let local_fri: FriVerify = fri_from_tvm.into();
181 let fri_pointer =
182 rust_shadowing_helper_functions::dyn_malloc::dynamic_allocator(memory);
183 encode_to_memory(memory, fri_pointer, &local_fri);
184 stack.push(fri_pointer);
185
186 Ok(())
187 }
188
189 fn pseudorandom_initial_state(
190 &self,
191 seed: [u8; 32],
192 bench_case: Option<BenchmarkCase>,
193 ) -> FunctionInitialState {
194 #[cfg(target_pointer_width = "32")]
201 const WORST_CASE_BENCH_SIZE: u32 = 21;
202 #[cfg(target_pointer_width = "64")]
203 const WORST_CASE_BENCH_SIZE: u32 = 23;
204
205 #[cfg(target_pointer_width = "32")]
206 const MAX_BENCH_SIZE: u32 = WORST_CASE_BENCH_SIZE;
207 #[cfg(target_pointer_width = "64")]
208 const MAX_BENCH_SIZE: u32 = 25;
209
210 let padded_height: u32 = match bench_case {
211 Some(BenchmarkCase::CommonCase) => 2u32.pow(21),
212 Some(BenchmarkCase::WorstCase) => 2u32.pow(WORST_CASE_BENCH_SIZE),
213 None => {
214 let mut rng = StdRng::from_seed(seed);
215 let mut padded_height = 2u32.pow(rng.random_range(8..=MAX_BENCH_SIZE));
216
217 while self.stark.fri(padded_height as usize * 2).is_err() {
223 padded_height /= 2;
224 }
225
226 assert!(padded_height >= 2u32.pow(8));
227
228 padded_height
229 }
230 };
231
232 FunctionInitialState {
233 stack: [
234 self.init_stack_for_isolated_run(),
235 vec![padded_height.into()],
236 ]
237 .concat(),
238 memory: HashMap::default(),
239 }
240 }
241 }
242}