Skip to main content

tasm_lib/verifier/fri/
derive_from_stark.rs

1use 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/// Mimics Triton-VM's FRI parameter-derivation method, but doesn't allow for a FRI-domain length
10/// of 2^32 bc the domain length is stored in a single word/a `u32`.
11#[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        // padded-height independent lower bounds on the randomized trace
25        // length. From `Stark::randomized_trace_len`. The `+ 1` comes from
26        // (private) `NUM_OUT_OF_DOMAIN_QUOTIENTS`, defined upstream.
27        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            // _ padded_height
31
32            addi {num_trace_randomizers}
33            // _ (padded_height + num_trace_randomizers)
34
35            /* Clamp to at least `min_randomized_trace_len`:
36               a + (a < min)·min − (a < min)·a = max(a, min) */
37            push {min_randomized_trace_len}
38            dup 1
39            lt
40            // _ a (a < min)
41
42            dup 0
43            push {min_randomized_trace_len}
44            mul
45            // _ a (a < min) ((a < min)·min)
46
47            place 2
48            // _ ((a < min)·min) a (a < min)
49
50            push -1
51            mul
52            addi 1
53            // _ ((a < min)·min) a (1 - (a < min))
54
55            mul
56            add
57            // _ max(a, min_randomized_trace_len)
58
59            call {next_power_of_two}
60            // _ next_pow2(max(padded_height + num_trace_randomizers, min_randomized_trace_len))
61            // _ interpolant_codeword_length
62        );
63        let fri_domain_length = triton_asm!(
64            // _ interpolant_codeword_length
65            push {fri_expansion_factor}
66            mul
67            // _ (interpolant_codeword_length * fri_expansion_factor)
68            // _ fri_domain_length
69        );
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            // _ padded_height
76
77            {&interpolant_codeword_length_code}
78            {&fri_domain_length}
79            // _ fri_domain_length
80
81            push {num_collinearity_checks}
82            // _ fri_domain_length num_collinearity_checks
83
84            push {expansion_factor}
85            // _ fri_domain_length num_collinearity_checks expansion_factor
86
87            swap 2
88            // _ expansion_factor num_collinearity_checks fri_domain_length
89
90            push {domain_offset}
91            // _ expansion_factor num_collinearity_checks fri_domain_length domain_offset
92
93            dup 1
94            split
95            call {domain_generator}
96            // _ expansion_factor num_collinearity_checks fri_domain_length domain_offset domain_generator
97        )
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                // _ padded_height
125
126                {&derive_fri_field_values}
127                // _ fri_domain_length domain_offset domain_generator num_collinearity_checks expansion_factor
128
129                call {dyn_malloc}
130                // _ fri_domain_length domain_offset domain_generator num_collinearity_checks expansion_factor *fri_verify
131
132                write_mem 5
133                // _ (*fri_verify + 5)
134
135                push -5
136                add
137                // _ *fri_verify
138
139                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            // Due to an arithmetic-overflow bug in Triton VM v3.0.0, derivation
195            // of a FRI instance using values that are too close to `usize::MAX`
196            // (and what “too close” means depends on the FRI expansion factor)
197            // results in a `panic!`, not an `Err`. The workaround is not pretty
198            // but should be temporary.
199
200            #[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                    // Don't test parameters that result in too big FRI domains, i.e. larger
218                    // than 2^32. Note that this also excludes 2^32 as domain length because
219                    // the type used to hold this value is a `u32` in this repo. I think such a
220                    // large FRI domain is unfeasible anyway, so I'm reasonably comfortable
221                    // excluding that option.
222                    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}