openvm_algebra_circuit/extension/
modular.rs

1use std::{array, sync::Arc};
2
3use num_bigint::BigUint;
4use num_traits::{FromPrimitive, One};
5use openvm_algebra_transpiler::{ModularPhantom, Rv32ModularArithmeticOpcode};
6use openvm_circuit::{
7    self,
8    arch::{
9        AirInventory, AirInventoryError, ChipInventory, ChipInventoryError, ExecutionBridge,
10        ExecutorInventoryBuilder, ExecutorInventoryError, RowMajorMatrixArena, VmCircuitExtension,
11        VmExecutionExtension, VmProverExtension, DEFAULT_BLOCK_SIZE,
12    },
13    system::{memory::SharedMemoryHelper, SystemPort},
14};
15use openvm_circuit_derive::{AnyEnum, Executor, MeteredExecutor, PreflightExecutor};
16use openvm_circuit_primitives::{
17    bigint::utils::big_uint_to_limbs,
18    bitwise_op_lookup::{
19        BitwiseOperationLookupAir, BitwiseOperationLookupBus, BitwiseOperationLookupChip,
20        SharedBitwiseOperationLookupChip,
21    },
22    var_range::VariableRangeCheckerBus,
23};
24use openvm_cpu_backend::{CpuBackend, CpuDevice};
25use openvm_instructions::{LocalOpcode, PhantomDiscriminant, VmOpcode};
26use openvm_mod_circuit_builder::ExprBuilderConfig;
27use openvm_rv32_adapters::{
28    Rv32IsEqualModAdapterAir, Rv32IsEqualModAdapterExecutor, Rv32IsEqualModAdapterFiller,
29};
30use openvm_stark_backend::{p3_field::PrimeField32, StarkEngine, StarkProtocolConfig, Val};
31use rand::RngCore;
32use serde::{Deserialize, Serialize};
33use serde_with::{serde_as, DisplayFromStr};
34use strum::EnumCount;
35
36use crate::{
37    modular_chip::{
38        get_modular_addsub_air, get_modular_addsub_chip, get_modular_addsub_executor,
39        get_modular_muldiv_air, get_modular_muldiv_chip, get_modular_muldiv_executor, ModularAir,
40        ModularExecutor, ModularIsEqualAir, ModularIsEqualChip, ModularIsEqualCoreAir,
41        ModularIsEqualFiller, VmModularIsEqualExecutor,
42    },
43    AlgebraCpuProverExt, MODULAR_BLOCKS_32, MODULAR_BLOCKS_48, NUM_LIMBS_32, NUM_LIMBS_48,
44};
45
46#[serde_as]
47#[derive(Clone, Debug, derive_new::new, Serialize, Deserialize)]
48pub struct ModularExtension {
49    #[serde_as(as = "Vec<DisplayFromStr>")]
50    pub supported_moduli: Vec<BigUint>,
51}
52
53impl ModularExtension {
54    // Generates a call to the moduli_init! macro with moduli in the correct order
55    pub fn generate_moduli_init(&self) -> String {
56        let supported_moduli = self
57            .supported_moduli
58            .iter()
59            .map(|modulus| format!("\"{modulus}\""))
60            .collect::<Vec<String>>()
61            .join(", ");
62
63        format!("openvm_algebra_guest::moduli_macros::moduli_init! {{ {supported_moduli} }}",)
64    }
65}
66
67#[derive(Clone, AnyEnum, Executor, MeteredExecutor, PreflightExecutor)]
68#[cfg_attr(
69    feature = "aot",
70    derive(
71        openvm_circuit_derive::AotExecutor,
72        openvm_circuit_derive::AotMeteredExecutor
73    )
74)]
75pub enum ModularExtensionExecutor {
76    // 32 limbs prime
77    ModularAddSubRv32_32(ModularExecutor<MODULAR_BLOCKS_32, DEFAULT_BLOCK_SIZE>), // ModularAddSub
78    ModularMulDivRv32_32(ModularExecutor<MODULAR_BLOCKS_32, DEFAULT_BLOCK_SIZE>), // ModularMulDiv
79    ModularIsEqualRv32_32(
80        VmModularIsEqualExecutor<MODULAR_BLOCKS_32, DEFAULT_BLOCK_SIZE, NUM_LIMBS_32>,
81    ), /* ModularIsEqual */
82    // 48 limbs prime
83    ModularAddSubRv32_48(ModularExecutor<MODULAR_BLOCKS_48, DEFAULT_BLOCK_SIZE>), // ModularAddSub
84    ModularMulDivRv32_48(ModularExecutor<MODULAR_BLOCKS_48, DEFAULT_BLOCK_SIZE>), // ModularMulDiv
85    ModularIsEqualRv32_48(
86        VmModularIsEqualExecutor<MODULAR_BLOCKS_48, DEFAULT_BLOCK_SIZE, NUM_LIMBS_48>,
87    ), /* ModularIsEqual */
88}
89
90impl<F: PrimeField32> VmExecutionExtension<F> for ModularExtension {
91    type Executor = ModularExtensionExecutor;
92
93    fn extend_execution(
94        &self,
95        inventory: &mut ExecutorInventoryBuilder<F, ModularExtensionExecutor>,
96    ) -> Result<(), ExecutorInventoryError> {
97        let pointer_max_bits = inventory.pointer_max_bits();
98        // TODO: somehow get the range checker bus from `ExecutorInventory`
99        let dummy_range_checker_bus = VariableRangeCheckerBus::new(u16::MAX, 16);
100        for (i, modulus) in self.supported_moduli.iter().enumerate() {
101            // determine the number of bytes needed to represent a prime field element
102            let bytes = modulus.bits().div_ceil(8) as usize;
103            let start_offset =
104                Rv32ModularArithmeticOpcode::CLASS_OFFSET + i * Rv32ModularArithmeticOpcode::COUNT;
105            let modulus_limbs = big_uint_to_limbs(modulus, 8);
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_modular_addsub_executor::<MODULAR_BLOCKS_32, DEFAULT_BLOCK_SIZE>(
113                    config.clone(),
114                    dummy_range_checker_bus,
115                    pointer_max_bits,
116                    start_offset,
117                );
118
119                inventory.add_executor(
120                    ModularExtensionExecutor::ModularAddSubRv32_32(addsub),
121                    ((Rv32ModularArithmeticOpcode::ADD as usize)
122                        ..=(Rv32ModularArithmeticOpcode::SETUP_ADDSUB as usize))
123                        .map(|x| VmOpcode::from_usize(x + start_offset)),
124                )?;
125
126                let muldiv = get_modular_muldiv_executor::<MODULAR_BLOCKS_32, DEFAULT_BLOCK_SIZE>(
127                    config,
128                    dummy_range_checker_bus,
129                    pointer_max_bits,
130                    start_offset,
131                );
132
133                inventory.add_executor(
134                    ModularExtensionExecutor::ModularMulDivRv32_32(muldiv),
135                    ((Rv32ModularArithmeticOpcode::MUL as usize)
136                        ..=(Rv32ModularArithmeticOpcode::SETUP_MULDIV as usize))
137                        .map(|x| VmOpcode::from_usize(x + start_offset)),
138                )?;
139
140                let modulus_limbs = array::from_fn(|i| {
141                    if i < modulus_limbs.len() {
142                        modulus_limbs[i] as u8
143                    } else {
144                        0
145                    }
146                });
147
148                let is_eq = VmModularIsEqualExecutor::new(
149                    Rv32IsEqualModAdapterExecutor::new(pointer_max_bits),
150                    start_offset,
151                    modulus_limbs,
152                );
153
154                inventory.add_executor(
155                    ModularExtensionExecutor::ModularIsEqualRv32_32(is_eq),
156                    ((Rv32ModularArithmeticOpcode::IS_EQ as usize)
157                        ..=(Rv32ModularArithmeticOpcode::SETUP_ISEQ as usize))
158                        .map(|x| VmOpcode::from_usize(x + start_offset)),
159                )?;
160            } else if bytes <= NUM_LIMBS_48 {
161                let config = ExprBuilderConfig {
162                    modulus: modulus.clone(),
163                    num_limbs: NUM_LIMBS_48,
164                    limb_bits: 8,
165                };
166                let addsub = get_modular_addsub_executor::<MODULAR_BLOCKS_48, DEFAULT_BLOCK_SIZE>(
167                    config.clone(),
168                    dummy_range_checker_bus,
169                    pointer_max_bits,
170                    start_offset,
171                );
172
173                inventory.add_executor(
174                    ModularExtensionExecutor::ModularAddSubRv32_48(addsub),
175                    ((Rv32ModularArithmeticOpcode::ADD as usize)
176                        ..=(Rv32ModularArithmeticOpcode::SETUP_ADDSUB as usize))
177                        .map(|x| VmOpcode::from_usize(x + start_offset)),
178                )?;
179
180                let muldiv = get_modular_muldiv_executor::<MODULAR_BLOCKS_48, DEFAULT_BLOCK_SIZE>(
181                    config,
182                    dummy_range_checker_bus,
183                    pointer_max_bits,
184                    start_offset,
185                );
186
187                inventory.add_executor(
188                    ModularExtensionExecutor::ModularMulDivRv32_48(muldiv),
189                    ((Rv32ModularArithmeticOpcode::MUL as usize)
190                        ..=(Rv32ModularArithmeticOpcode::SETUP_MULDIV as usize))
191                        .map(|x| VmOpcode::from_usize(x + start_offset)),
192                )?;
193
194                let modulus_limbs = array::from_fn(|i| {
195                    if i < modulus_limbs.len() {
196                        modulus_limbs[i] as u8
197                    } else {
198                        0
199                    }
200                });
201
202                let is_eq = VmModularIsEqualExecutor::new(
203                    Rv32IsEqualModAdapterExecutor::new(pointer_max_bits),
204                    start_offset,
205                    modulus_limbs,
206                );
207
208                inventory.add_executor(
209                    ModularExtensionExecutor::ModularIsEqualRv32_48(is_eq),
210                    ((Rv32ModularArithmeticOpcode::IS_EQ as usize)
211                        ..=(Rv32ModularArithmeticOpcode::SETUP_ISEQ as usize))
212                        .map(|x| VmOpcode::from_usize(x + start_offset)),
213                )?;
214            } else {
215                panic!("Modulus too large");
216            }
217        }
218
219        let non_qr_hint_sub_ex = phantom::NonQrHintSubEx::new(self.supported_moduli.clone());
220        inventory.add_phantom_sub_executor(
221            non_qr_hint_sub_ex.clone(),
222            PhantomDiscriminant(ModularPhantom::HintNonQr as u16),
223        )?;
224
225        let sqrt_hint_sub_ex = phantom::SqrtHintSubEx::new(non_qr_hint_sub_ex);
226        inventory.add_phantom_sub_executor(
227            sqrt_hint_sub_ex,
228            PhantomDiscriminant(ModularPhantom::HintSqrt as u16),
229        )?;
230
231        Ok(())
232    }
233}
234
235impl<SC: StarkProtocolConfig> VmCircuitExtension<SC> for ModularExtension {
236    fn extend_circuit(&self, inventory: &mut AirInventory<SC>) -> Result<(), AirInventoryError> {
237        let SystemPort {
238            execution_bus,
239            program_bus,
240            memory_bridge,
241        } = inventory.system().port();
242
243        let exec_bridge = ExecutionBridge::new(execution_bus, program_bus);
244        let range_checker_bus = inventory.range_checker().bus;
245        let pointer_max_bits = inventory.pointer_max_bits();
246
247        let bitwise_lu = {
248            // A trick to get around Rust's borrow rules
249            let existing_air = inventory.find_air::<BitwiseOperationLookupAir<8>>().next();
250            if let Some(air) = existing_air {
251                air.bus
252            } else {
253                let bus = BitwiseOperationLookupBus::new(inventory.new_bus_idx());
254                let air = BitwiseOperationLookupAir::<8>::new(bus);
255                inventory.add_air(air);
256                air.bus
257            }
258        };
259        for (i, modulus) in self.supported_moduli.iter().enumerate() {
260            // determine the number of bytes needed to represent a prime field element
261            let bytes = modulus.bits().div_ceil(8) as usize;
262            let start_offset =
263                Rv32ModularArithmeticOpcode::CLASS_OFFSET + i * Rv32ModularArithmeticOpcode::COUNT;
264
265            if bytes <= NUM_LIMBS_32 {
266                let config = ExprBuilderConfig {
267                    modulus: modulus.clone(),
268                    num_limbs: NUM_LIMBS_32,
269                    limb_bits: 8,
270                };
271
272                let addsub = get_modular_addsub_air::<MODULAR_BLOCKS_32, DEFAULT_BLOCK_SIZE>(
273                    exec_bridge,
274                    memory_bridge,
275                    config.clone(),
276                    range_checker_bus,
277                    bitwise_lu,
278                    pointer_max_bits,
279                    start_offset,
280                );
281                inventory.add_air(addsub);
282
283                let muldiv = get_modular_muldiv_air::<MODULAR_BLOCKS_32, DEFAULT_BLOCK_SIZE>(
284                    exec_bridge,
285                    memory_bridge,
286                    config,
287                    range_checker_bus,
288                    bitwise_lu,
289                    pointer_max_bits,
290                    start_offset,
291                );
292                inventory.add_air(muldiv);
293
294                let is_eq =
295                    ModularIsEqualAir::<MODULAR_BLOCKS_32, DEFAULT_BLOCK_SIZE, NUM_LIMBS_32>::new(
296                        Rv32IsEqualModAdapterAir::new(
297                            exec_bridge,
298                            memory_bridge,
299                            bitwise_lu,
300                            pointer_max_bits,
301                        ),
302                        ModularIsEqualCoreAir::new(modulus.clone(), bitwise_lu, start_offset),
303                    );
304                inventory.add_air(is_eq);
305            } else if bytes <= NUM_LIMBS_48 {
306                let config = ExprBuilderConfig {
307                    modulus: modulus.clone(),
308                    num_limbs: NUM_LIMBS_48,
309                    limb_bits: 8,
310                };
311
312                let addsub = get_modular_addsub_air::<MODULAR_BLOCKS_48, DEFAULT_BLOCK_SIZE>(
313                    exec_bridge,
314                    memory_bridge,
315                    config.clone(),
316                    range_checker_bus,
317                    bitwise_lu,
318                    pointer_max_bits,
319                    start_offset,
320                );
321                inventory.add_air(addsub);
322
323                let muldiv = get_modular_muldiv_air::<MODULAR_BLOCKS_48, DEFAULT_BLOCK_SIZE>(
324                    exec_bridge,
325                    memory_bridge,
326                    config,
327                    range_checker_bus,
328                    bitwise_lu,
329                    pointer_max_bits,
330                    start_offset,
331                );
332                inventory.add_air(muldiv);
333
334                let is_eq =
335                    ModularIsEqualAir::<MODULAR_BLOCKS_48, DEFAULT_BLOCK_SIZE, NUM_LIMBS_48>::new(
336                        Rv32IsEqualModAdapterAir::new(
337                            exec_bridge,
338                            memory_bridge,
339                            bitwise_lu,
340                            pointer_max_bits,
341                        ),
342                        ModularIsEqualCoreAir::new(modulus.clone(), bitwise_lu, start_offset),
343                    );
344                inventory.add_air(is_eq);
345            } else {
346                panic!("Modulus too large");
347            }
348        }
349
350        Ok(())
351    }
352}
353
354// This implementation is specific to CpuBackend because the lookup chips (VariableRangeChecker,
355// BitwiseOperationLookupChip) are specific to CpuBackend.
356impl<SC, E, RA> VmProverExtension<E, RA, ModularExtension> for AlgebraCpuProverExt
357where
358    SC: StarkProtocolConfig,
359    E: StarkEngine<SC = SC, PB = CpuBackend<SC>, PD = CpuDevice<SC>>,
360    RA: RowMajorMatrixArena<Val<SC>>,
361    Val<SC>: PrimeField32,
362    SC::EF: Ord,
363{
364    fn extend_prover(
365        &self,
366        extension: &ModularExtension,
367        inventory: &mut ChipInventory<SC, RA, CpuBackend<SC>>,
368    ) -> Result<(), ChipInventoryError> {
369        let range_checker = inventory.range_checker()?.clone();
370        let timestamp_max_bits = inventory.timestamp_max_bits();
371        let pointer_max_bits = inventory.airs().pointer_max_bits();
372        let mem_helper = SharedMemoryHelper::new(range_checker.clone(), timestamp_max_bits);
373        let bitwise_lu = {
374            let existing_chip = inventory
375                .find_chip::<SharedBitwiseOperationLookupChip<8>>()
376                .next();
377            if let Some(chip) = existing_chip {
378                chip.clone()
379            } else {
380                let air: &BitwiseOperationLookupAir<8> = inventory.next_air()?;
381                let chip = Arc::new(BitwiseOperationLookupChip::new(air.bus));
382                inventory.add_periphery_chip(chip.clone());
383                chip
384            }
385        };
386        for (i, modulus) in extension.supported_moduli.iter().enumerate() {
387            // determine the number of bytes needed to represent a prime field element
388            let bytes = modulus.bits().div_ceil(8) as usize;
389            let start_offset =
390                Rv32ModularArithmeticOpcode::CLASS_OFFSET + i * Rv32ModularArithmeticOpcode::COUNT;
391
392            let modulus_limbs = big_uint_to_limbs(modulus, 8);
393
394            if bytes <= NUM_LIMBS_32 {
395                let config = ExprBuilderConfig {
396                    modulus: modulus.clone(),
397                    num_limbs: NUM_LIMBS_32,
398                    limb_bits: 8,
399                };
400
401                inventory.next_air::<ModularAir<MODULAR_BLOCKS_32, DEFAULT_BLOCK_SIZE>>()?;
402                let addsub =
403                    get_modular_addsub_chip::<Val<SC>, MODULAR_BLOCKS_32, DEFAULT_BLOCK_SIZE>(
404                        config.clone(),
405                        mem_helper.clone(),
406                        range_checker.clone(),
407                        bitwise_lu.clone(),
408                        pointer_max_bits,
409                    );
410                inventory.add_executor_chip(addsub);
411
412                inventory.next_air::<ModularAir<MODULAR_BLOCKS_32, DEFAULT_BLOCK_SIZE>>()?;
413                let muldiv =
414                    get_modular_muldiv_chip::<Val<SC>, MODULAR_BLOCKS_32, DEFAULT_BLOCK_SIZE>(
415                        config,
416                        mem_helper.clone(),
417                        range_checker.clone(),
418                        bitwise_lu.clone(),
419                        pointer_max_bits,
420                    );
421                inventory.add_executor_chip(muldiv);
422
423                let modulus_limbs = array::from_fn(|i| {
424                    if i < modulus_limbs.len() {
425                        modulus_limbs[i] as u8
426                    } else {
427                        0
428                    }
429                });
430                inventory.next_air::<ModularIsEqualAir<MODULAR_BLOCKS_32, DEFAULT_BLOCK_SIZE, NUM_LIMBS_32>>()?;
431                let is_eq = ModularIsEqualChip::<
432                    Val<SC>,
433                    MODULAR_BLOCKS_32,
434                    DEFAULT_BLOCK_SIZE,
435                    NUM_LIMBS_32,
436                >::new(
437                    ModularIsEqualFiller::new(
438                        Rv32IsEqualModAdapterFiller::new(pointer_max_bits, bitwise_lu.clone()),
439                        start_offset,
440                        modulus_limbs,
441                        bitwise_lu.clone(),
442                    ),
443                    mem_helper.clone(),
444                );
445                inventory.add_executor_chip(is_eq);
446            } else if bytes <= NUM_LIMBS_48 {
447                let config = ExprBuilderConfig {
448                    modulus: modulus.clone(),
449                    num_limbs: NUM_LIMBS_48,
450                    limb_bits: 8,
451                };
452
453                inventory.next_air::<ModularAir<MODULAR_BLOCKS_48, DEFAULT_BLOCK_SIZE>>()?;
454                let addsub =
455                    get_modular_addsub_chip::<Val<SC>, MODULAR_BLOCKS_48, DEFAULT_BLOCK_SIZE>(
456                        config.clone(),
457                        mem_helper.clone(),
458                        range_checker.clone(),
459                        bitwise_lu.clone(),
460                        pointer_max_bits,
461                    );
462                inventory.add_executor_chip(addsub);
463
464                inventory.next_air::<ModularAir<MODULAR_BLOCKS_48, DEFAULT_BLOCK_SIZE>>()?;
465                let muldiv =
466                    get_modular_muldiv_chip::<Val<SC>, MODULAR_BLOCKS_48, DEFAULT_BLOCK_SIZE>(
467                        config,
468                        mem_helper.clone(),
469                        range_checker.clone(),
470                        bitwise_lu.clone(),
471                        pointer_max_bits,
472                    );
473                inventory.add_executor_chip(muldiv);
474
475                let modulus_limbs = array::from_fn(|i| {
476                    if i < modulus_limbs.len() {
477                        modulus_limbs[i] as u8
478                    } else {
479                        0
480                    }
481                });
482                inventory.next_air::<ModularIsEqualAir<MODULAR_BLOCKS_48, DEFAULT_BLOCK_SIZE, NUM_LIMBS_48>>()?;
483                let is_eq = ModularIsEqualChip::<
484                    Val<SC>,
485                    MODULAR_BLOCKS_48,
486                    DEFAULT_BLOCK_SIZE,
487                    NUM_LIMBS_48,
488                >::new(
489                    ModularIsEqualFiller::new(
490                        Rv32IsEqualModAdapterFiller::new(pointer_max_bits, bitwise_lu.clone()),
491                        start_offset,
492                        modulus_limbs,
493                        bitwise_lu.clone(),
494                    ),
495                    mem_helper.clone(),
496                );
497                inventory.add_executor_chip(is_eq);
498            } else {
499                panic!("Modulus too large");
500            }
501        }
502
503        Ok(())
504    }
505}
506
507pub(crate) mod phantom {
508    use std::{
509        iter::{once, repeat},
510        ops::Deref,
511    };
512
513    use eyre::bail;
514    use num_bigint::BigUint;
515    use openvm_circuit::{
516        arch::{PhantomSubExecutor, Streams},
517        system::memory::online::GuestMemory,
518    };
519    use openvm_instructions::{riscv::RV32_MEMORY_AS, PhantomDiscriminant};
520    use openvm_rv32im_circuit::adapters::read_rv32_register;
521    use openvm_stark_backend::p3_field::PrimeField32;
522    use rand::{rngs::StdRng, SeedableRng};
523
524    use super::{find_non_qr, mod_sqrt};
525    use crate::{NUM_LIMBS_32, NUM_LIMBS_48};
526
527    #[derive(derive_new::new)]
528    pub struct SqrtHintSubEx(NonQrHintSubEx);
529
530    impl Deref for SqrtHintSubEx {
531        type Target = NonQrHintSubEx;
532
533        fn deref(&self) -> &NonQrHintSubEx {
534            &self.0
535        }
536    }
537
538    // Given x returns either a sqrt of x or a sqrt of x * non_qr, whichever exists.
539    // Note that non_qr is fixed for each modulus.
540    impl<F: PrimeField32> PhantomSubExecutor<F> for SqrtHintSubEx {
541        fn phantom_execute(
542            &self,
543            memory: &GuestMemory,
544            streams: &mut Streams<F>,
545            _: &mut StdRng,
546            _: PhantomDiscriminant,
547            a: u32,
548            _: u32,
549            c_upper: u16,
550        ) -> eyre::Result<()> {
551            let mod_idx = c_upper as usize;
552            if mod_idx >= self.supported_moduli.len() {
553                bail!(
554                    "Modulus index {mod_idx} out of range: {} supported moduli",
555                    self.supported_moduli.len()
556                );
557            }
558            let modulus = &self.supported_moduli[mod_idx];
559            let bytes = modulus.bits().div_ceil(8) as usize;
560            let num_limbs: usize = if bytes <= NUM_LIMBS_32 {
561                NUM_LIMBS_32
562            } else if bytes <= NUM_LIMBS_48 {
563                NUM_LIMBS_48
564            } else {
565                bail!("Modulus too large")
566            };
567
568            let rs1 = read_rv32_register(memory, a);
569            // SAFETY:
570            // - MEMORY_AS consists of `u8`s
571            // - MEMORY_AS is in bounds
572            let x_limbs: Vec<u8> =
573                unsafe { memory.memory.get_slice((RV32_MEMORY_AS, rs1), num_limbs) }.to_vec();
574            let x = BigUint::from_bytes_le(&x_limbs);
575
576            let (success, sqrt) = match mod_sqrt(&x, modulus, &self.non_qrs[mod_idx]) {
577                Some(sqrt) => (true, sqrt),
578                None => {
579                    let sqrt = mod_sqrt(
580                        &(&x * &self.non_qrs[mod_idx]),
581                        modulus,
582                        &self.non_qrs[mod_idx],
583                    )
584                    .expect("Either x or x * non_qr should be a square");
585                    (false, sqrt)
586                }
587            };
588
589            let hint_bytes = once(F::from_bool(success))
590                .chain(repeat(F::ZERO))
591                .take(4)
592                .chain(
593                    sqrt.to_bytes_le()
594                        .into_iter()
595                        .map(F::from_u8)
596                        .chain(repeat(F::ZERO))
597                        .take(num_limbs),
598                )
599                .collect();
600            streams.hint_stream = hint_bytes;
601            Ok(())
602        }
603    }
604
605    #[derive(Clone)]
606    pub struct NonQrHintSubEx {
607        pub supported_moduli: Vec<BigUint>,
608        pub non_qrs: Vec<BigUint>,
609    }
610
611    impl NonQrHintSubEx {
612        pub fn new(supported_moduli: Vec<BigUint>) -> Self {
613            // Use deterministic seed so that the non-QR are deterministic between different
614            // instances of the VM. The seed determines the runtime of Tonelli-Shanks, if the
615            // algorithm is necessary, which affects the time it takes to construct and initialize
616            // the VM but does not affect the runtime.
617            let mut rng = StdRng::from_seed([0u8; 32]);
618            let non_qrs = supported_moduli
619                .iter()
620                .map(|modulus| find_non_qr(modulus, &mut rng))
621                .collect();
622            Self {
623                supported_moduli,
624                non_qrs,
625            }
626        }
627    }
628
629    impl<F: PrimeField32> PhantomSubExecutor<F> for NonQrHintSubEx {
630        fn phantom_execute(
631            &self,
632            _: &GuestMemory,
633            streams: &mut Streams<F>,
634            _: &mut StdRng,
635            _: PhantomDiscriminant,
636            _: u32,
637            _: u32,
638            c_upper: u16,
639        ) -> eyre::Result<()> {
640            let mod_idx = c_upper as usize;
641            if mod_idx >= self.supported_moduli.len() {
642                bail!(
643                    "Modulus index {mod_idx} out of range: {} supported moduli",
644                    self.supported_moduli.len()
645                );
646            }
647            let modulus = &self.supported_moduli[mod_idx];
648
649            let bytes = modulus.bits().div_ceil(8) as usize;
650            let num_limbs: usize = if bytes <= NUM_LIMBS_32 {
651                NUM_LIMBS_32
652            } else if bytes <= NUM_LIMBS_48 {
653                NUM_LIMBS_48
654            } else {
655                bail!("Modulus too large")
656            };
657
658            let hint_bytes = self.non_qrs[mod_idx]
659                .to_bytes_le()
660                .into_iter()
661                .map(F::from_u8)
662                .chain(repeat(F::ZERO))
663                .take(num_limbs)
664                .collect();
665            streams.hint_stream = hint_bytes;
666            Ok(())
667        }
668    }
669}
670
671/// Find the square root of `x` modulo `modulus` with `non_qr` a
672/// quadratic nonresidue of the field.
673pub fn mod_sqrt(x: &BigUint, modulus: &BigUint, non_qr: &BigUint) -> Option<BigUint> {
674    if modulus % 4u32 == BigUint::from_u8(3).unwrap() {
675        // x^(1/2) = x^((p+1)/4) when p = 3 mod 4
676        let exponent = (modulus + BigUint::one()) >> 2;
677        let ret = x.modpow(&exponent, modulus);
678        if &ret * &ret % modulus == x % modulus {
679            Some(ret)
680        } else {
681            None
682        }
683    } else {
684        // Tonelli-Shanks algorithm
685        // https://en.wikipedia.org/wiki/Tonelli%E2%80%93Shanks_algorithm#The_algorithm
686        let mut q = modulus - BigUint::one();
687        let mut s = 0;
688        while &q % 2u32 == BigUint::ZERO {
689            s += 1;
690            q /= 2u32;
691        }
692        let z = non_qr;
693        let mut m = s;
694        let mut c = z.modpow(&q, modulus);
695        let mut t = x.modpow(&q, modulus);
696        let mut r = x.modpow(&((q + BigUint::one()) >> 1), modulus);
697        loop {
698            if t == BigUint::ZERO {
699                return Some(BigUint::ZERO);
700            }
701            if t == BigUint::one() {
702                return Some(r);
703            }
704            let mut i = 0;
705            let mut tmp = t.clone();
706            while tmp != BigUint::one() && i < m {
707                tmp = &tmp * &tmp % modulus;
708                i += 1;
709            }
710            if i == m {
711                // self is not a quadratic residue
712                return None;
713            }
714            for _ in 0..m - i - 1 {
715                c = &c * &c % modulus;
716            }
717            let b = c;
718            m = i;
719            c = &b * &b % modulus;
720            t = ((t * &b % modulus) * &b) % modulus;
721            r = (r * b) % modulus;
722        }
723    }
724}
725
726// Returns a non-quadratic residue in the field
727pub fn find_non_qr(modulus: &BigUint, rng: &mut impl RngCore) -> BigUint {
728    if modulus % 4u32 == BigUint::from(3u8) {
729        // p = 3 mod 4 then -1 is a quadratic residue
730        modulus - BigUint::one()
731    } else if modulus % 8u32 == BigUint::from(5u8) {
732        // p = 5 mod 8 then 2 is a non-quadratic residue
733        // since 2^((p-1)/2) = (-1)^((p^2-1)/8)
734        BigUint::from_u8(2u8).unwrap()
735    } else {
736        // Sample uniformly from [2, modulus - 1) using rejection sampling
737        let range = modulus - 3u32; // number of values in [2, modulus-1)
738        let mut buf = vec![0u8; modulus.to_bytes_be().len()];
739        let exponent = (modulus - BigUint::one()) >> 1;
740        loop {
741            // Rejection sample for uniform distribution
742            rng.fill_bytes(&mut buf);
743            let val = BigUint::from_bytes_be(&buf);
744            if val >= range {
745                continue;
746            }
747            let non_qr = val + 2u32;
748            // To check if non_qr is a quadratic nonresidue, we compute non_qr^((p-1)/2)
749            // If the result is p-1, then non_qr is a quadratic nonresidue
750            if non_qr.modpow(&exponent, modulus) == modulus - BigUint::one() {
751                return non_qr;
752            }
753        }
754    }
755}