openvm_algebra_circuit/extension/
fp2.rs

1use std::sync::Arc;
2
3use num_bigint::BigUint;
4use openvm_algebra_transpiler::Fp2Opcode;
5use openvm_circuit::{
6    arch::{
7        AirInventory, AirInventoryError, ChipInventory, ChipInventoryError, ExecutionBridge,
8        ExecutorInventoryBuilder, ExecutorInventoryError, RowMajorMatrixArena, VmCircuitExtension,
9        VmExecutionExtension, VmProverExtension, DEFAULT_BLOCK_SIZE,
10    },
11    system::{memory::SharedMemoryHelper, SystemPort},
12};
13use openvm_circuit_derive::{AnyEnum, Executor, MeteredExecutor, PreflightExecutor};
14use openvm_circuit_primitives::{
15    bitwise_op_lookup::{
16        BitwiseOperationLookupAir, BitwiseOperationLookupBus, BitwiseOperationLookupChip,
17        SharedBitwiseOperationLookupChip,
18    },
19    var_range::VariableRangeCheckerBus,
20};
21use openvm_cpu_backend::{CpuBackend, CpuDevice};
22use openvm_instructions::{LocalOpcode, VmOpcode};
23use openvm_mod_circuit_builder::ExprBuilderConfig;
24use openvm_stark_backend::{p3_field::PrimeField32, StarkEngine, StarkProtocolConfig, Val};
25use serde::{Deserialize, Serialize};
26use serde_with::{serde_as, DisplayFromStr};
27use strum::EnumCount;
28
29use crate::{
30    fp2_chip::{
31        get_fp2_addsub_air, get_fp2_addsub_chip, get_fp2_addsub_executor, get_fp2_muldiv_air,
32        get_fp2_muldiv_chip, get_fp2_muldiv_executor, Fp2Air, Fp2Executor,
33    },
34    AlgebraCpuProverExt, ModularExtension, FP2_BLOCKS_32, FP2_BLOCKS_48, NUM_LIMBS_32,
35    NUM_LIMBS_48,
36};
37
38#[serde_as]
39#[derive(Clone, Debug, derive_new::new, Serialize, Deserialize)]
40pub struct Fp2Extension {
41    // (name, modulus)
42    // name must match the struct name defined by complex_declare
43    #[serde_as(as = "Vec<(_, DisplayFromStr)>")]
44    pub supported_moduli: Vec<(String, BigUint)>,
45}
46
47impl Fp2Extension {
48    pub fn generate_complex_init(&self, modular_config: &ModularExtension) -> String {
49        fn get_index_of_modulus(modulus: &BigUint, modular_config: &ModularExtension) -> usize {
50            modular_config
51                .supported_moduli
52                .iter()
53                .position(|m| m == modulus)
54                .expect("Modulus used in Fp2Extension not found in ModularExtension")
55        }
56
57        let supported_moduli = self
58            .supported_moduli
59            .iter()
60            .map(|(name, modulus)| {
61                format!(
62                    "\"{}\" {{ mod_idx = {} }}",
63                    name,
64                    get_index_of_modulus(modulus, modular_config)
65                )
66            })
67            .collect::<Vec<String>>()
68            .join(", ");
69
70        format!("openvm_algebra_guest::complex_macros::complex_init! {{ {supported_moduli} }}")
71    }
72}
73
74#[derive(Clone, AnyEnum, Executor, MeteredExecutor, PreflightExecutor)]
75#[cfg_attr(
76    feature = "aot",
77    derive(
78        openvm_circuit_derive::AotExecutor,
79        openvm_circuit_derive::AotMeteredExecutor
80    )
81)]
82pub enum Fp2ExtensionExecutor {
83    // 32 limbs prime
84    Fp2AddSubRv32_32(Fp2Executor<FP2_BLOCKS_32, DEFAULT_BLOCK_SIZE>), // Fp2AddSub
85    Fp2MulDivRv32_32(Fp2Executor<FP2_BLOCKS_32, DEFAULT_BLOCK_SIZE>), // Fp2MulDiv
86    // 48 limbs prime
87    Fp2AddSubRv32_48(Fp2Executor<FP2_BLOCKS_48, DEFAULT_BLOCK_SIZE>), // Fp2AddSub
88    Fp2MulDivRv32_48(Fp2Executor<FP2_BLOCKS_48, DEFAULT_BLOCK_SIZE>), // Fp2MulDiv
89}
90
91impl<F: PrimeField32> VmExecutionExtension<F> for Fp2Extension {
92    type Executor = Fp2ExtensionExecutor;
93
94    fn extend_execution(
95        &self,
96        inventory: &mut ExecutorInventoryBuilder<F, Fp2ExtensionExecutor>,
97    ) -> Result<(), ExecutorInventoryError> {
98        let pointer_max_bits = inventory.pointer_max_bits();
99        // TODO: somehow get the range checker bus from `ExecutorInventory`
100        let dummy_range_checker_bus = VariableRangeCheckerBus::new(u16::MAX, 16);
101        for (i, (_, modulus)) in self.supported_moduli.iter().enumerate() {
102            // determine the number of bytes needed to represent a prime field element
103            let bytes = modulus.bits().div_ceil(8) as usize;
104            let start_offset = Fp2Opcode::CLASS_OFFSET + i * Fp2Opcode::COUNT;
105
106            if bytes <= NUM_LIMBS_32 {
107                let config = ExprBuilderConfig {
108                    modulus: modulus.clone(),
109                    num_limbs: NUM_LIMBS_32,
110                    limb_bits: 8,
111                };
112                let addsub = get_fp2_addsub_executor(
113                    config.clone(),
114                    dummy_range_checker_bus,
115                    pointer_max_bits,
116                    start_offset,
117                );
118
119                inventory.add_executor(
120                    Fp2ExtensionExecutor::Fp2AddSubRv32_32(addsub),
121                    ((Fp2Opcode::ADD as usize)..=(Fp2Opcode::SETUP_ADDSUB as usize))
122                        .map(|x| VmOpcode::from_usize(x + start_offset)),
123                )?;
124
125                let muldiv = get_fp2_muldiv_executor(
126                    config,
127                    dummy_range_checker_bus,
128                    pointer_max_bits,
129                    start_offset,
130                );
131
132                inventory.add_executor(
133                    Fp2ExtensionExecutor::Fp2MulDivRv32_32(muldiv),
134                    ((Fp2Opcode::MUL as usize)..=(Fp2Opcode::SETUP_MULDIV as usize))
135                        .map(|x| VmOpcode::from_usize(x + start_offset)),
136                )?;
137            } else if bytes <= NUM_LIMBS_48 {
138                let config = ExprBuilderConfig {
139                    modulus: modulus.clone(),
140                    num_limbs: NUM_LIMBS_48,
141                    limb_bits: 8,
142                };
143                let addsub = get_fp2_addsub_executor(
144                    config.clone(),
145                    dummy_range_checker_bus,
146                    pointer_max_bits,
147                    start_offset,
148                );
149
150                inventory.add_executor(
151                    Fp2ExtensionExecutor::Fp2AddSubRv32_48(addsub),
152                    ((Fp2Opcode::ADD as usize)..=(Fp2Opcode::SETUP_ADDSUB as usize))
153                        .map(|x| VmOpcode::from_usize(x + start_offset)),
154                )?;
155
156                let muldiv = get_fp2_muldiv_executor(
157                    config,
158                    dummy_range_checker_bus,
159                    pointer_max_bits,
160                    start_offset,
161                );
162
163                inventory.add_executor(
164                    Fp2ExtensionExecutor::Fp2MulDivRv32_48(muldiv),
165                    ((Fp2Opcode::MUL as usize)..=(Fp2Opcode::SETUP_MULDIV as usize))
166                        .map(|x| VmOpcode::from_usize(x + start_offset)),
167                )?;
168            } else {
169                panic!("Modulus too large");
170            }
171        }
172        Ok(())
173    }
174}
175
176impl<SC: StarkProtocolConfig> VmCircuitExtension<SC> for Fp2Extension {
177    fn extend_circuit(&self, inventory: &mut AirInventory<SC>) -> Result<(), AirInventoryError> {
178        let SystemPort {
179            execution_bus,
180            program_bus,
181            memory_bridge,
182        } = inventory.system().port();
183
184        let exec_bridge = ExecutionBridge::new(execution_bus, program_bus);
185        let range_checker_bus = inventory.range_checker().bus;
186        let pointer_max_bits = inventory.pointer_max_bits();
187
188        let bitwise_lu = {
189            // A trick to get around Rust's borrow rules
190            let existing_air = inventory.find_air::<BitwiseOperationLookupAir<8>>().next();
191            if let Some(air) = existing_air {
192                air.bus
193            } else {
194                let bus = BitwiseOperationLookupBus::new(inventory.new_bus_idx());
195                let air = BitwiseOperationLookupAir::<8>::new(bus);
196                inventory.add_air(air);
197                air.bus
198            }
199        };
200        for (i, (_, modulus)) in self.supported_moduli.iter().enumerate() {
201            // determine the number of bytes needed to represent a prime field element
202            let bytes = modulus.bits().div_ceil(8) as usize;
203            let start_offset = Fp2Opcode::CLASS_OFFSET + i * Fp2Opcode::COUNT;
204
205            if bytes <= NUM_LIMBS_32 {
206                let config = ExprBuilderConfig {
207                    modulus: modulus.clone(),
208                    num_limbs: NUM_LIMBS_32,
209                    limb_bits: 8,
210                };
211
212                let addsub = get_fp2_addsub_air::<FP2_BLOCKS_32, DEFAULT_BLOCK_SIZE>(
213                    exec_bridge,
214                    memory_bridge,
215                    config.clone(),
216                    range_checker_bus,
217                    bitwise_lu,
218                    pointer_max_bits,
219                    start_offset,
220                );
221                inventory.add_air(addsub);
222
223                let muldiv = get_fp2_muldiv_air::<FP2_BLOCKS_32, DEFAULT_BLOCK_SIZE>(
224                    exec_bridge,
225                    memory_bridge,
226                    config,
227                    range_checker_bus,
228                    bitwise_lu,
229                    pointer_max_bits,
230                    start_offset,
231                );
232                inventory.add_air(muldiv);
233            } else if bytes <= NUM_LIMBS_48 {
234                let config = ExprBuilderConfig {
235                    modulus: modulus.clone(),
236                    num_limbs: NUM_LIMBS_48,
237                    limb_bits: 8,
238                };
239
240                let addsub = get_fp2_addsub_air::<FP2_BLOCKS_48, DEFAULT_BLOCK_SIZE>(
241                    exec_bridge,
242                    memory_bridge,
243                    config.clone(),
244                    range_checker_bus,
245                    bitwise_lu,
246                    pointer_max_bits,
247                    start_offset,
248                );
249                inventory.add_air(addsub);
250
251                let muldiv = get_fp2_muldiv_air::<FP2_BLOCKS_48, DEFAULT_BLOCK_SIZE>(
252                    exec_bridge,
253                    memory_bridge,
254                    config,
255                    range_checker_bus,
256                    bitwise_lu,
257                    pointer_max_bits,
258                    start_offset,
259                );
260                inventory.add_air(muldiv);
261            } else {
262                panic!("Modulus too large");
263            }
264        }
265
266        Ok(())
267    }
268}
269
270// This implementation is specific to CpuBackend because the lookup chips (VariableRangeChecker,
271// BitwiseOperationLookupChip) are specific to CpuBackend.
272impl<SC, E, RA> VmProverExtension<E, RA, Fp2Extension> for AlgebraCpuProverExt
273where
274    SC: StarkProtocolConfig,
275    E: StarkEngine<SC = SC, PB = CpuBackend<SC>, PD = CpuDevice<SC>>,
276    RA: RowMajorMatrixArena<Val<SC>>,
277    Val<SC>: PrimeField32,
278    SC::EF: Ord,
279{
280    fn extend_prover(
281        &self,
282        extension: &Fp2Extension,
283        inventory: &mut ChipInventory<SC, RA, CpuBackend<SC>>,
284    ) -> Result<(), ChipInventoryError> {
285        let range_checker = inventory.range_checker()?.clone();
286        let timestamp_max_bits = inventory.timestamp_max_bits();
287        let pointer_max_bits = inventory.airs().pointer_max_bits();
288        let mem_helper = SharedMemoryHelper::new(range_checker.clone(), timestamp_max_bits);
289        let bitwise_lu = {
290            let existing_chip = inventory
291                .find_chip::<SharedBitwiseOperationLookupChip<8>>()
292                .next();
293            if let Some(chip) = existing_chip {
294                chip.clone()
295            } else {
296                let air: &BitwiseOperationLookupAir<8> = inventory.next_air()?;
297                let chip = Arc::new(BitwiseOperationLookupChip::new(air.bus));
298                inventory.add_periphery_chip(chip.clone());
299                chip
300            }
301        };
302        for (_, modulus) in extension.supported_moduli.iter() {
303            // determine the number of bytes needed to represent a prime field element
304            let bytes = modulus.bits().div_ceil(8) as usize;
305
306            if bytes <= NUM_LIMBS_32 {
307                let config = ExprBuilderConfig {
308                    modulus: modulus.clone(),
309                    num_limbs: NUM_LIMBS_32,
310                    limb_bits: 8,
311                };
312
313                inventory.next_air::<Fp2Air<FP2_BLOCKS_32, DEFAULT_BLOCK_SIZE>>()?;
314                let addsub = get_fp2_addsub_chip::<Val<SC>, FP2_BLOCKS_32, DEFAULT_BLOCK_SIZE>(
315                    config.clone(),
316                    mem_helper.clone(),
317                    range_checker.clone(),
318                    bitwise_lu.clone(),
319                    pointer_max_bits,
320                );
321                inventory.add_executor_chip(addsub);
322
323                inventory.next_air::<Fp2Air<FP2_BLOCKS_32, DEFAULT_BLOCK_SIZE>>()?;
324                let muldiv = get_fp2_muldiv_chip::<Val<SC>, FP2_BLOCKS_32, DEFAULT_BLOCK_SIZE>(
325                    config,
326                    mem_helper.clone(),
327                    range_checker.clone(),
328                    bitwise_lu.clone(),
329                    pointer_max_bits,
330                );
331                inventory.add_executor_chip(muldiv);
332            } else if bytes <= NUM_LIMBS_48 {
333                let config = ExprBuilderConfig {
334                    modulus: modulus.clone(),
335                    num_limbs: NUM_LIMBS_48,
336                    limb_bits: 8,
337                };
338
339                inventory.next_air::<Fp2Air<FP2_BLOCKS_48, DEFAULT_BLOCK_SIZE>>()?;
340                let addsub = get_fp2_addsub_chip::<Val<SC>, FP2_BLOCKS_48, DEFAULT_BLOCK_SIZE>(
341                    config.clone(),
342                    mem_helper.clone(),
343                    range_checker.clone(),
344                    bitwise_lu.clone(),
345                    pointer_max_bits,
346                );
347                inventory.add_executor_chip(addsub);
348
349                inventory.next_air::<Fp2Air<FP2_BLOCKS_48, DEFAULT_BLOCK_SIZE>>()?;
350                let muldiv = get_fp2_muldiv_chip::<Val<SC>, FP2_BLOCKS_48, DEFAULT_BLOCK_SIZE>(
351                    config,
352                    mem_helper.clone(),
353                    range_checker.clone(),
354                    bitwise_lu.clone(),
355                    pointer_max_bits,
356                );
357                inventory.add_executor_chip(muldiv);
358            } else {
359                panic!("Modulus too large");
360            }
361        }
362
363        Ok(())
364    }
365}