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 #[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 Fp2AddSubRv32_32(Fp2Executor<FP2_BLOCKS_32, DEFAULT_BLOCK_SIZE>), Fp2MulDivRv32_32(Fp2Executor<FP2_BLOCKS_32, DEFAULT_BLOCK_SIZE>), Fp2AddSubRv32_48(Fp2Executor<FP2_BLOCKS_48, DEFAULT_BLOCK_SIZE>), Fp2MulDivRv32_48(Fp2Executor<FP2_BLOCKS_48, DEFAULT_BLOCK_SIZE>), }
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 let dummy_range_checker_bus = VariableRangeCheckerBus::new(u16::MAX, 16);
101 for (i, (_, modulus)) in self.supported_moduli.iter().enumerate() {
102 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 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 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
270impl<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 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}