Skip to main content

tasm_lib/verifier/
out_of_domain_points.rs

1use triton_vm::prelude::*;
2use triton_vm::table::NUM_QUOTIENT_SEGMENTS;
3use twenty_first::math::x_field_element::EXTENSION_DEGREE;
4
5use crate::data_type::ArrayType;
6use crate::prelude::*;
7
8/// Calculate the four needed values related to out-of-domain points and store them in a statically
9/// allocated array. Return the pointer to this array.
10#[derive(Debug, Clone, Copy)]
11pub struct OutOfDomainPoints;
12
13pub const NUM_OF_OUT_OF_DOMAIN_POINTS: usize = 4;
14
15#[derive(Debug, Clone, Copy)]
16pub enum OodPoint {
17    CurrentRow,
18    NextRow,
19    CurrentRowPowNumSegments,
20    CurrentRowTimesZetaPowNumSegments,
21}
22
23impl OutOfDomainPoints {
24    /// Push the requested OOD point to the stack, pop the pointer.
25    pub fn read_ood_point(ood_point_type: OodPoint) -> Vec<LabelledInstruction> {
26        let address_offset = (ood_point_type as usize) * EXTENSION_DEGREE + (EXTENSION_DEGREE - 1);
27        triton_asm!(
28            // _ *ood_points // of type same as the output value of this snippet
29
30            push {address_offset}
31            add
32            // _ (*ood_points[n] + 2)
33
34            read_mem {EXTENSION_DEGREE}
35            // _ [ood_point] (*ood_points[n] - 1)
36
37            pop 1
38        )
39    }
40}
41
42impl BasicSnippet for OutOfDomainPoints {
43    fn parameters(&self) -> Vec<(DataType, String)> {
44        vec![
45            (DataType::Bfe, "trace_domain_generator".to_owned()),
46            (DataType::Xfe, "out_of_domain_curr_row".to_owned()),
47        ]
48    }
49
50    fn return_values(&self) -> Vec<(DataType, String)> {
51        vec![(
52            DataType::Array(Box::new(ArrayType {
53                element_type: DataType::Xfe,
54                length: NUM_OF_OUT_OF_DOMAIN_POINTS,
55            })),
56            "out_of_domain_points".to_owned(),
57        )]
58    }
59
60    fn entrypoint(&self) -> String {
61        "tasmlib_verifier_out_of_domain_points".to_owned()
62    }
63
64    fn code(&self, library: &mut Library) -> Vec<LabelledInstruction> {
65        let entrypoint = self.entrypoint();
66
67        // Snippet for sampling *one* scalar, and holding the values:
68        // - `out_of_domain_point_curr_row`
69        // - `out_of_domain_point_next_row`
70        // - `out_of_domain_point_curr_row_pow_num_segments`
71        // - `out_of_domain_point_curr_row_times_zeta_pow_num_segments`
72        let num_words_for_out_of_domain_points = (NUM_OF_OUT_OF_DOMAIN_POINTS * EXTENSION_DEGREE)
73            .try_into()
74            .unwrap();
75        let ood_points_alloc = library.kmalloc(num_words_for_out_of_domain_points);
76
77        triton_asm!(
78            {entrypoint}:
79                // _ trace_domain_generator [ood_curr_row]
80
81                dup 2
82                dup 2
83                dup 2
84                dup 2
85                dup 2
86                dup 2
87                push {ood_points_alloc.write_address()}
88                write_mem {EXTENSION_DEGREE}
89                // _ trace_domain_generator [ood_curr_row] [ood_curr_row] *ood_points[1]
90
91                swap 7
92                // _ *ood_points[1] [ood_curr_row] [ood_curr_row] trace_domain_generator
93
94                xb_mul
95                // _ *ood_points[1] [ood_curr_row] [ood_next_row]
96
97                dup 6
98                write_mem {EXTENSION_DEGREE}
99                // _ *ood_points[1] [ood_curr_row] *ood_points[2]
100
101                swap 4
102                pop 1
103                // _ *ood_points[2] [ood_curr_row]
104
105                dup 2 dup 2 dup 2
106                xx_mul
107                dup 2 dup 2 dup 2
108                xx_mul
109                // _ *ood_points[2] [ood_curr_row**4]
110
111                dup 2 dup 2 dup 2
112                // _ *ood_points[2] [ood_curr_row**4] [ood_curr_row**4]
113
114                pick 6
115                write_mem {EXTENSION_DEGREE}
116                // _ [ood_curr_row**4] *ood_points[3]
117
118                place 3
119                // _ *ood_points[3] [ood_curr_row**4]
120
121                push {Stark::ZETA.mod_pow(NUM_QUOTIENT_SEGMENTS as u64)}
122                xb_mul
123                // _ *ood_points[3] [(ood_curr_row·ζ)**4]
124
125                pick 3
126                write_mem {EXTENSION_DEGREE}
127                // _ *ood_points[4]
128
129                addi {-((4 * EXTENSION_DEGREE) as i32)}
130                // _ *ood_points
131
132                return
133        )
134    }
135}
136
137#[cfg(test)]
138mod tests {
139    use twenty_first::math::traits::ModPowU32;
140    use twenty_first::math::traits::PrimitiveRootOfUnity;
141
142    use super::*;
143    use crate::rust_shadowing_helper_functions::array::insert_as_array;
144    use crate::test_prelude::*;
145
146    #[macro_rules_attr::apply(test)]
147    fn ood_points_pbt() {
148        ShadowedFunction::new(OutOfDomainPoints).test();
149    }
150
151    impl Function for OutOfDomainPoints {
152        fn rust_shadow(
153            &self,
154            stack: &mut Vec<BFieldElement>,
155            memory: &mut HashMap<BFieldElement, BFieldElement>,
156        ) -> Result<(), RustShadowError> {
157            let ood_curr_row = XFieldElement::new([
158                stack.pop().ok_or(RustShadowError::StackUnderflow)?,
159                stack.pop().ok_or(RustShadowError::StackUnderflow)?,
160                stack.pop().ok_or(RustShadowError::StackUnderflow)?,
161            ]);
162            let domain_generator = stack.pop().ok_or(RustShadowError::StackUnderflow)?;
163            let ood_next_row = ood_curr_row * domain_generator;
164            let num_quotient_segments: u32 = NUM_QUOTIENT_SEGMENTS
165                .try_into()
166                .map_err(|_| RustShadowError::UsizeToU32Error)?;
167            let ood_curr_row_pow_num_segments = ood_curr_row.mod_pow_u32(num_quotient_segments);
168            let ood_curr_row_times_zeta_pow_num_segments =
169                (ood_curr_row * Stark::ZETA).mod_pow_u32(num_quotient_segments);
170            let static_malloc_size: i32 = (EXTENSION_DEGREE * NUM_OF_OUT_OF_DOMAIN_POINTS)
171                .try_into()
172                .map_err(|_| RustShadowError::Other)?;
173            let ood_points_pointer = bfe!(-static_malloc_size - 1);
174            insert_as_array(
175                ood_points_pointer,
176                memory,
177                vec![
178                    ood_curr_row,
179                    ood_next_row,
180                    ood_curr_row_pow_num_segments,
181                    ood_curr_row_times_zeta_pow_num_segments,
182                ],
183            );
184
185            stack.push(ood_points_pointer);
186
187            Ok(())
188        }
189
190        fn pseudorandom_initial_state(
191            &self,
192            seed: [u8; 32],
193            bench_case: Option<BenchmarkCase>,
194        ) -> FunctionInitialState {
195            let domain_length = match bench_case {
196                Some(BenchmarkCase::CommonCase) => 1u64 << 20,
197                Some(BenchmarkCase::WorstCase) => 1u64 << 24,
198                None => {
199                    let mut rng = StdRng::from_seed(seed);
200                    1u64 << rng.random_range(8..=32)
201                }
202            };
203            println!("domain_length: {domain_length}");
204
205            let domain_generator = BFieldElement::primitive_root_of_unity(domain_length).unwrap();
206            let ood_curr_row: XFieldElement = rand::random();
207
208            FunctionInitialState {
209                stack: [
210                    self.init_stack_for_isolated_run(),
211                    vec![
212                        domain_generator,
213                        ood_curr_row.coefficients[2],
214                        ood_curr_row.coefficients[1],
215                        ood_curr_row.coefficients[0],
216                    ],
217                ]
218                .concat(),
219                memory: HashMap::default(),
220            }
221        }
222    }
223}
224
225#[cfg(test)]
226mod benches {
227    use super::*;
228    use crate::test_prelude::*;
229
230    #[macro_rules_attr::apply(test)]
231    fn benchmark() {
232        ShadowedFunction::new(OutOfDomainPoints).bench();
233    }
234}