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 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 ModularAddSubRv32_32(ModularExecutor<MODULAR_BLOCKS_32, DEFAULT_BLOCK_SIZE>), ModularMulDivRv32_32(ModularExecutor<MODULAR_BLOCKS_32, DEFAULT_BLOCK_SIZE>), ModularIsEqualRv32_32(
80 VmModularIsEqualExecutor<MODULAR_BLOCKS_32, DEFAULT_BLOCK_SIZE, NUM_LIMBS_32>,
81 ), ModularAddSubRv32_48(ModularExecutor<MODULAR_BLOCKS_48, DEFAULT_BLOCK_SIZE>), ModularMulDivRv32_48(ModularExecutor<MODULAR_BLOCKS_48, DEFAULT_BLOCK_SIZE>), ModularIsEqualRv32_48(
86 VmModularIsEqualExecutor<MODULAR_BLOCKS_48, DEFAULT_BLOCK_SIZE, NUM_LIMBS_48>,
87 ), }
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 let dummy_range_checker_bus = VariableRangeCheckerBus::new(u16::MAX, 16);
100 for (i, modulus) in self.supported_moduli.iter().enumerate() {
101 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 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 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
354impl<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 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 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 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 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
671pub fn mod_sqrt(x: &BigUint, modulus: &BigUint, non_qr: &BigUint) -> Option<BigUint> {
674 if modulus % 4u32 == BigUint::from_u8(3).unwrap() {
675 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 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 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
726pub fn find_non_qr(modulus: &BigUint, rng: &mut impl RngCore) -> BigUint {
728 if modulus % 4u32 == BigUint::from(3u8) {
729 modulus - BigUint::one()
731 } else if modulus % 8u32 == BigUint::from(5u8) {
732 BigUint::from_u8(2u8).unwrap()
735 } else {
736 let range = modulus - 3u32; let mut buf = vec![0u8; modulus.to_bytes_be().len()];
739 let exponent = (modulus - BigUint::one()) >> 1;
740 loop {
741 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 if non_qr.modpow(&exponent, modulus) == modulus - BigUint::one() {
751 return non_qr;
752 }
753 }
754 }
755}