1use binius_frontend::{CircuitBuilder, Wire, WitnessFiller};
4use num_integer::Integer;
5
6use super::fixed_byte_vec::ByteVec;
7use crate::{
8 bignum::{BigUint, ModReduce, assert_eq, biguint_lt, optimal_mul, optimal_sqr},
9 bytes::swap_bytes,
10 sha256::sha256_varlen,
11};
12
13fn fixedbytevec_le_to_biguint(builder: &CircuitBuilder, byte_vec: &ByteVec) -> BigUint {
18 let limbs = byte_vec
21 .data
22 .iter()
23 .rev()
24 .map(|&packed_wire| swap_bytes(builder, packed_wire))
25 .collect();
26 BigUint { limbs }
27}
28
29pub struct Rs256Verify {
39 pub message: ByteVec,
41 pub signature: ByteVec,
43 pub modulus: ByteVec,
45 pub rsa_intermediates: RsaIntermediates,
47}
48
49impl Rs256Verify {
50 pub fn new(
74 builder: &mut CircuitBuilder,
75 message: ByteVec,
76 signature: ByteVec,
77 modulus: ByteVec,
78 ) -> Self {
79 assert!(
80 signature.data.len() >= 32,
81 "signature must have at least 256 bytes for 2048-bit RSA"
82 );
83 assert!(modulus.data.len() >= 32, "modulus must have at least 256 bytes for 2048-bit RSA");
84
85 let signature = if signature.data.len() > 32 {
91 signature.truncate(builder, 32)
92 } else {
93 signature
94 };
95
96 let signature_bignum = fixedbytevec_le_to_biguint(builder, &signature);
97 builder.assert_eq("signature_bytes_len", signature.len_bytes, builder.add_constant_64(256));
98
99 let modulus_bignum = fixedbytevec_le_to_biguint(builder, &modulus);
100 builder.assert_eq("modulus_bytes_len", modulus.len_bytes, builder.add_constant_64(256));
101
102 let signature_wide =
109 signature_bignum.pad_limbs_to(modulus_bignum.limbs.len(), builder.add_constant_64(0));
110 builder.assert_true(
111 "signature_below_modulus",
112 biguint_lt(builder, &signature_wide, &modulus_bignum),
113 );
114
115 let expected_hash_wires: [Wire; 4] = sha256_varlen(&builder.subcircuit("sha256"), &message);
116 let expected_hash = BigUint {
117 limbs: expected_hash_wires.to_vec(),
118 };
119
120 let rsa_intermediates = RsaIntermediates::new_witness(builder);
121
122 modexp_65537_verify(
123 builder,
124 &signature_bignum,
125 &modulus_bignum,
126 &rsa_intermediates.square_quotients,
127 &rsa_intermediates.square_remainders,
128 &rsa_intermediates.mul_quotient,
129 &rsa_intermediates.mul_remainder,
130 );
131
132 const EXPECTED_PREFIX_LIMBS: [u64; 28] = [
148 0x0304020105000420,
150 0x0d06096086480165,
151 0xffffffff00303130,
152 0xffffffffffffffff,
154 0xffffffffffffffff,
155 0xffffffffffffffff,
156 0xffffffffffffffff,
157 0xffffffffffffffff,
158 0xffffffffffffffff,
159 0xffffffffffffffff,
160 0xffffffffffffffff,
161 0xffffffffffffffff,
162 0xffffffffffffffff,
163 0xffffffffffffffff,
164 0xffffffffffffffff,
165 0xffffffffffffffff,
166 0xffffffffffffffff,
167 0xffffffffffffffff,
168 0xffffffffffffffff,
169 0xffffffffffffffff,
170 0xffffffffffffffff,
171 0xffffffffffffffff,
172 0xffffffffffffffff,
173 0xffffffffffffffff,
174 0xffffffffffffffff,
175 0xffffffffffffffff,
176 0xffffffffffffffff,
177 0x0001ffffffffffff,
179 ];
180
181 let prefix_wires = EXPECTED_PREFIX_LIMBS.map(|l| builder.add_constant_64(l));
183 let expected_em = BigUint {
184 limbs: expected_hash
185 .limbs
186 .iter()
187 .copied()
188 .rev()
189 .chain(prefix_wires)
190 .collect(),
191 };
192
193 assert_eq(
194 builder,
195 "mul_remainder_expected_em",
196 &rsa_intermediates.mul_remainder,
197 &expected_em,
198 );
199
200 Self {
201 message,
202 signature,
203 modulus,
204 rsa_intermediates,
205 }
206 }
207
208 pub fn populate_len_bytes(&self, w: &mut WitnessFiller<'_>, len_bytes: usize) {
210 self.message.populate_len_bytes(w, len_bytes);
211 }
212
213 pub fn populate_rsa(&self, w: &mut WitnessFiller<'_>, signature: &[u8], modulus: &[u8]) {
215 self.populate_signature(w, signature);
216 self.populate_modulus(w, modulus);
217 self.rsa_intermediates
218 .populate_witness(w, signature, modulus);
219 }
220
221 pub fn populate_intermediates(
222 &self,
223 w: &mut WitnessFiller<'_>,
224 signature: &[u8],
225 modulus: &[u8],
226 ) {
227 self.rsa_intermediates
228 .populate_witness(w, signature, modulus);
229 }
230
231 pub fn populate_message(&self, w: &mut WitnessFiller<'_>, message: &[u8]) {
236 self.message.populate_data(w, message);
237 }
238
239 pub fn populate_modulus(&self, w: &mut WitnessFiller<'_>, modulus_bytes: &[u8]) {
244 assert_eq!(modulus_bytes.len(), 256, "modulus must be exactly 256 bytes");
245 self.modulus.populate_bytes_le(w, modulus_bytes);
246 }
247
248 pub fn populate_signature(&self, w: &mut WitnessFiller<'_>, signature_bytes: &[u8]) {
257 assert_eq!(signature_bytes.len(), 256, "signature must be exactly 256 bytes");
258 self.signature.populate_bytes_le(w, signature_bytes);
259 }
260}
261
262fn modexp_65537_verify(
264 builder: &CircuitBuilder,
265 base: &BigUint,
266 modulus: &BigUint,
267 square_quotients: &[BigUint],
268 square_remainders: &[BigUint],
269 mul_quotient: &BigUint,
270 mul_remainder: &BigUint,
271) {
272 let mut result = base.clone();
273
274 for i in 0..16 {
275 let builder = builder.subcircuit(format!("square[{i}]"));
276 let squared = optimal_sqr(&builder, &result);
277 let circuit = ModReduce::new(
278 &builder,
279 squared,
280 modulus.clone(),
281 square_quotients[i].clone(),
282 square_remainders[i].clone(),
283 );
284 result = circuit.remainder;
285 }
286
287 let builder = builder.subcircuit("final_multiply");
288 let multiplied = optimal_mul(&builder, &result, base);
289 let _mod_reduce_multiplied = ModReduce::new(
290 &builder,
291 multiplied,
292 modulus.clone(),
293 mul_quotient.clone(),
294 mul_remainder.clone(),
295 );
296}
297
298pub struct RsaIntermediates {
300 square_quotients: Vec<BigUint>,
302 square_remainders: Vec<BigUint>,
304 mul_quotient: BigUint,
306 mul_remainder: BigUint,
308}
309
310impl RsaIntermediates {
311 fn new_witness(builder: &CircuitBuilder) -> Self {
312 let mut square_quotients = Vec::new();
313 let mut square_remainders = Vec::new();
314 for _ in 0..16 {
315 square_quotients.push(BigUint::new_witness(builder, 32));
316 square_remainders.push(BigUint::new_witness(builder, 32));
317 }
318 let mul_quotient = BigUint::new_witness(builder, 32);
319 let mul_remainder = BigUint::new_witness(builder, 32);
320
321 RsaIntermediates {
322 square_quotients,
323 square_remainders,
324 mul_quotient,
325 mul_remainder,
326 }
327 }
328
329 pub fn populate_witness(&self, w: &mut WitnessFiller<'_>, signature: &[u8], modulus: &[u8]) {
338 assert_eq!(signature.len(), 256, "signature must be exactly 256 bytes");
339 assert_eq!(modulus.len(), 256, "modulus must be exactly 256 bytes");
340
341 let signature_value = num_bigint::BigUint::from_bytes_be(signature);
342 let modulus_value = num_bigint::BigUint::from_bytes_be(modulus);
343
344 let mut square_quotients = Vec::new();
345 let mut square_remainders = Vec::new();
346
347 let mut result = signature_value.clone();
348 for _ in 0..16 {
349 let squared = &result * &result;
350 let (q, r) = squared.div_rem(&modulus_value);
351
352 let mut q_limbs = q.to_u64_digits();
353 q_limbs.resize(32, 0u64);
354 square_quotients.push(q_limbs);
355
356 let mut r_limbs = r.to_u64_digits();
357 r_limbs.resize(32, 0u64);
358 square_remainders.push(r_limbs);
359
360 result = r;
361 }
362
363 let multiplied = &result * &signature_value;
365 let (mul_q, mul_r) = multiplied.div_rem(&modulus_value);
366
367 let mut mul_quotient = mul_q.to_u64_digits();
368 mul_quotient.resize(32, 0u64);
369
370 let mut mul_remainder = mul_r.to_u64_digits();
371 mul_remainder.resize(32, 0u64);
372
373 self.populate_square_quotients(w, &square_quotients);
374 self.populate_square_remainders(w, &square_remainders);
375 self.populate_mul_quotient(w, &mul_quotient);
376 self.populate_mul_remainder(w, &mul_remainder);
377 }
378
379 fn populate_square_quotients(
384 &self,
385 w: &mut WitnessFiller<'_>,
386 square_quotient_limbs: &[Vec<u64>],
387 ) {
388 assert_eq!(square_quotient_limbs.len(), 16, "must provide 16 square quotients");
389 for (i, q_limbs) in square_quotient_limbs.iter().enumerate() {
390 assert_eq!(
391 q_limbs.len(),
392 self.square_quotients[i].limbs.len(),
393 "square_quotient[{i}] must have {} limbs",
394 self.square_quotients[i].limbs.len()
395 );
396 self.square_quotients[i].populate_limbs(w, q_limbs);
397 }
398 }
399
400 fn populate_square_remainders(
405 &self,
406 w: &mut WitnessFiller<'_>,
407 square_remainder_limbs: &[Vec<u64>],
408 ) {
409 assert_eq!(square_remainder_limbs.len(), 16, "must provide 16 square remainders");
410 for (i, r_limbs) in square_remainder_limbs.iter().enumerate() {
411 assert_eq!(r_limbs.len(), 32, "square_remainder[{i}] must have 32 limbs");
412 self.square_remainders[i].populate_limbs(w, r_limbs);
413 }
414 }
415
416 fn populate_mul_quotient(&self, w: &mut WitnessFiller<'_>, mul_quotient_limbs: &[u64]) {
421 assert_eq!(
422 mul_quotient_limbs.len(),
423 self.mul_quotient.limbs.len(),
424 "mul_quotient must have {} limbs",
425 self.mul_quotient.limbs.len()
426 );
427 self.mul_quotient.populate_limbs(w, mul_quotient_limbs);
428 }
429
430 fn populate_mul_remainder(&self, w: &mut WitnessFiller<'_>, mul_remainder_limbs: &[u64]) {
435 assert_eq!(mul_remainder_limbs.len(), 32, "mul_remainder must have 32 limbs");
436 self.mul_remainder.populate_limbs(w, mul_remainder_limbs);
437 }
438}
439
440#[cfg(test)]
441mod tests {
442 use hex_literal::hex;
443 use num_bigint::BigUint;
444 use rand::{TryRng, prelude::*};
445 use rsa::{
446 BigUint as RsaBigUint, RsaPrivateKey, RsaPublicKey,
447 sha2::{Digest, Sha256},
448 traits::{PrivateKeyParts, PublicKeyParts},
449 };
450
451 use super::*;
452
453 fn test_rsa_key() -> RsaPrivateKey {
456 let p = RsaBigUint::from_bytes_be(&hex!(
457 "c8b4e97508c3d0fad0062e8ee475909d5315bc9433e9b8a174a52b8f024e7d6b"
458 "ea80a56901555021b2d44f727aa287b84de8bac5ceef88d03b259f8ac91bda42"
459 "e653e27596d8090e08e9dac47dcd288e1c0e95ac74d7428cd0479c8514bc3538"
460 "7380a480873c7f519ece6f5ea4356c81bd7ec31c126c1f097b84bb33c8acd565"
461 ));
462 let q = RsaBigUint::from_bytes_be(&hex!(
463 "efffcc7f550f977db26971fb6a0f036d61cccde351c394fe177cd36a0a7dde60"
464 "8cd263d8ca382031fc0f16bef5ebb2125ab1b8e837c71c006a8639c090a7ebac"
465 "530de579bca2ea7ad175c8a31d45078130e0ad15cf23139d230f30c106259c7a"
466 "55024f4e51a97b1b38b7ed4dfe05a0706bf53a067e7f0ee18dc685b53300708b"
467 ));
468 let e = RsaBigUint::from(65537u32);
469 RsaPrivateKey::from_p_q(p, q, e).expect("valid key")
470 }
471
472 fn populate_circuit(
473 circuit: &Rs256Verify,
474 w: &mut WitnessFiller<'_>,
475 signature_bytes: &[u8],
476 message_bytes: &[u8],
477 modulus_bytes: &[u8],
478 ) {
479 circuit.populate_rsa(w, signature_bytes, modulus_bytes);
480 circuit.populate_len_bytes(w, message_bytes.len());
481 circuit.populate_message(w, message_bytes);
482 }
483
484 fn setup_circuit(builder: &mut CircuitBuilder, max_len: usize) -> Rs256Verify {
485 let signature_bytes = ByteVec::new_inout(builder, 32);
488 let modulus_bytes = ByteVec::new_inout(builder, 32);
489 let message = ByteVec::new_witness(builder, max_len);
490
491 Rs256Verify::new(builder, message, signature_bytes, modulus_bytes)
492 }
493
494 #[test]
495 fn test_real_rsa_signature_verification_with_message() {
496 let mut builder = CircuitBuilder::new();
497 let circuit = setup_circuit(&mut builder, 32);
498 let cs = builder.build();
499
500 let private_key = test_rsa_key();
501 let public_key = RsaPublicKey::from(&private_key);
502 let mut rng = StdRng::seed_from_u64(42);
503 let mut message_bytes = [0u8; 256];
504 rng.try_fill_bytes(&mut message_bytes).unwrap();
505
506 let digest = Sha256::digest(message_bytes);
508 let signature_bytes = private_key
509 .sign(rsa::Pkcs1v15Sign::new::<Sha256>(), &digest)
510 .expect("failed to sign");
511 let modulus_bytes = public_key.n().to_bytes_be();
512
513 let mut w = cs.new_witness_filler();
514 populate_circuit(&circuit, &mut w, &signature_bytes, &message_bytes, &modulus_bytes);
515
516 cs.populate_wire_witness(&mut w).unwrap();
517 cs.constraint_system().verify(&w.into_value_vec()).unwrap();
518 }
519
520 #[test]
521 fn test_real_rsa_signature_with_invalid_prefix() {
522 let mut builder = CircuitBuilder::new();
523 let max_message_len = 256;
524 let circuit = setup_circuit(&mut builder, max_message_len);
525 let cs = builder.build();
526
527 let private_key = test_rsa_key();
528 let public_key = RsaPublicKey::from(&private_key);
529
530 let message = b"Test message for RS256 verification with invalid prefix";
531
532 let corrupted_em = BigUint::ZERO;
535 let d_bytes = private_key.d().to_bytes_le();
536 let n_bytes = private_key.n().to_bytes_le();
537 let d = BigUint::from_bytes_le(&d_bytes);
538 let n = BigUint::from_bytes_le(&n_bytes);
539 let corrupted_signature = corrupted_em.modpow(&d, &n);
540
541 let mut signature_bytes = corrupted_signature.to_bytes_be();
542 signature_bytes.resize(256, 0u8);
543 let modulus_bytes = public_key.n().to_bytes_be();
544
545 let mut w = cs.new_witness_filler();
546 populate_circuit(&circuit, &mut w, &signature_bytes, message, &modulus_bytes);
547
548 let result = cs.populate_wire_witness(&mut w);
549 assert!(result.is_err(), "Circuit should fail when PKCS#1 v1.5 prefix is corrupted");
550 }
551
552 #[test]
553 fn test_real_rsa_signature_verification_with_wrong_message() {
554 let mut builder = CircuitBuilder::new();
555 let max_message_len = 256;
556 let circuit = setup_circuit(&mut builder, max_message_len);
557 let cs = builder.build();
558
559 let private_key = test_rsa_key();
560 let public_key = RsaPublicKey::from(&private_key);
561
562 let message = b"Test message for RS256 verification with wrong message";
563 let digest = Sha256::digest(message);
564 let signature_bytes = private_key
565 .sign(rsa::Pkcs1v15Sign::new::<Sha256>(), &digest)
566 .expect("failed to sign");
567
568 let signature_bytes = BigUint::from_bytes_be(&signature_bytes).to_bytes_be();
569 let modulus_bytes = public_key.n().to_bytes_be();
570
571 let wrong_message = b"This is a completely different message!";
573
574 let mut w = cs.new_witness_filler();
575 populate_circuit(&circuit, &mut w, &signature_bytes, wrong_message, &modulus_bytes);
576
577 let result = cs.populate_wire_witness(&mut w);
578 assert!(result.is_err(), "Circuit should fail when message doesn't match signature");
579 }
580
581 #[test]
584 fn test_modulus_wider_than_32_words_builds() {
585 let mut builder = CircuitBuilder::new();
586 let signature = ByteVec::new_inout(&builder, 32);
587 let modulus = ByteVec::new_inout(&builder, 33);
588 let message = ByteVec::new_witness(&builder, 8);
589 Rs256Verify::new(&mut builder, message, signature, modulus);
590 builder.build();
591 }
592
593 #[test]
598 fn test_real_rsa_signature_not_below_modulus_is_rejected() {
599 let private_key = test_rsa_key();
600 let public_key = RsaPublicKey::from(&private_key);
601 let n = BigUint::from_bytes_be(&public_key.n().to_bytes_be());
602 let bound = (BigUint::from(1u8) << 2048) - &n;
603
604 let mut rng = StdRng::seed_from_u64(7);
606 let (message_bytes, signature) = loop {
607 let mut message_bytes = [0u8; 64];
608 rng.try_fill_bytes(&mut message_bytes).unwrap();
609 let digest = Sha256::digest(message_bytes);
610 let signature_bytes = private_key
611 .sign(rsa::Pkcs1v15Sign::new::<Sha256>(), &digest)
612 .expect("failed to sign");
613 let signature = BigUint::from_bytes_be(&signature_bytes);
614 if signature < bound {
615 break (message_bytes, signature);
616 }
617 };
618
619 let mut lifted_bytes = (&signature + &n).to_bytes_be();
620 assert!(lifted_bytes.len() <= 256);
621 lifted_bytes.splice(0..0, std::iter::repeat_n(0u8, 256 - lifted_bytes.len()));
622
623 let digest = Sha256::digest(message_bytes);
625 assert!(
626 public_key
627 .verify(rsa::Pkcs1v15Sign::new::<Sha256>(), &digest, &lifted_bytes)
628 .is_err()
629 );
630
631 let mut builder = CircuitBuilder::new();
632 let circuit = setup_circuit(&mut builder, 8);
633 let cs = builder.build();
634
635 let modulus_bytes = public_key.n().to_bytes_be();
636 let mut w = cs.new_witness_filler();
637 populate_circuit(&circuit, &mut w, &lifted_bytes, &message_bytes, &modulus_bytes);
638
639 let accepted = cs.populate_wire_witness(&mut w).is_ok()
640 && cs.constraint_system().verify(&w.into_value_vec()).is_ok();
641 assert!(!accepted, "circuit accepted a signature representative not below the modulus");
642 }
643}