1use crate::error::{ElGamalError, Result};
4use num_bigint::{BigInt, BigUint, RandBigInt, ToBigInt, ToBigUint};
5use num_integer::Integer;
6use num_traits::{One, Zero};
7use rand::thread_rng;
8
9pub fn mod_exp(base: &BigUint, exp: &BigUint, modulus: &BigUint) -> BigUint {
11 base.modpow(exp, modulus)
12}
13
14pub fn mod_inverse(a: &BigUint, m: &BigUint) -> Option<BigUint> {
16 let (gcd, x, _) = extended_gcd(&a.to_bigint().unwrap(), &m.to_bigint().unwrap());
17
18 if gcd != BigInt::one() {
19 return None;
20 }
21
22 let result = if x < BigInt::zero() {
24 let m_bigint = m.to_bigint().unwrap();
25 let positive_x = ((x % &m_bigint) + &m_bigint) % &m_bigint;
26 positive_x.to_biguint().unwrap()
27 } else {
28 (x % m.to_bigint().unwrap()).to_biguint().unwrap()
29 };
30
31 Some(result)
32}
33
34fn extended_gcd(a: &BigInt, b: &BigInt) -> (BigInt, BigInt, BigInt) {
36 if a == &BigInt::zero() {
37 return (b.clone(), BigInt::zero(), BigInt::one());
38 }
39
40 let (gcd, x1, y1) = extended_gcd(&(b % a), a);
41 let x = y1 - (b / a) * &x1;
42 let y = x1;
43
44 (gcd, x, y)
45}
46
47pub fn generate_safe_prime(bit_size: u64) -> Result<(BigUint, BigUint)> {
49 if bit_size < 512 {
50 return Err(ElGamalError::InvalidKeySize(bit_size));
51 }
52
53 let mut rng = thread_rng();
54
55 let max_iterations = if bit_size <= 512 {
57 500000 } else if bit_size <= 1024 {
59 200000
60 } else {
61 100000
62 };
63
64 let mut iterations = 0;
65
66 let min_bits = bit_size.saturating_sub(1);
68 let max_bits = bit_size + 1;
69
70 loop {
71 iterations += 1;
72 if iterations > max_iterations {
73 return Err(ElGamalError::CryptoError(
74 format!("Failed to generate {}-bit safe prime after {} iterations. Consider using generate_for_testing() for tests or a larger key size for production.", bit_size, max_iterations)
75 ));
76 }
77
78 let mut q = rng.gen_biguint(bit_size - 1);
81
82 q |= BigUint::one();
84
85 if bit_size > 2 {
87 q |= BigUint::one() << (bit_size - 2);
88 }
89
90 if q.is_even() {
92 continue;
93 }
94
95 if !is_probable_prime(&q, 20) {
97 continue;
98 }
99
100 let p = &q * 2u32 + 1u32;
102
103 let p_bits = p.bits();
105 if p_bits < min_bits || p_bits > max_bits {
106 continue;
107 }
108
109 if is_probable_prime(&p, 20) {
111 return Ok((p, q));
112 }
113 }
114}
115
116pub fn generate_safe_prime_lenient(target_bit_size: u64) -> Result<(BigUint, BigUint)> {
118 if target_bit_size < 512 {
119 return Err(ElGamalError::InvalidKeySize(target_bit_size));
120 }
121
122 let mut rng = thread_rng();
123 let max_iterations = 1000000; let mut iterations = 0;
125
126 let min_bits = target_bit_size.saturating_sub(8);
128 let max_bits = target_bit_size + 8;
129
130 loop {
131 iterations += 1;
132 if iterations > max_iterations {
133 return Err(ElGamalError::CryptoError(format!(
134 "Failed to generate safe prime near {} bits after {} iterations",
135 target_bit_size, max_iterations
136 )));
137 }
138
139 let bit_variation = iterations % 17; let attempt_bits = if bit_variation < 8 {
142 target_bit_size.saturating_sub(bit_variation / 2)
143 } else {
144 target_bit_size + (bit_variation - 8) / 2
145 };
146
147 let q_bits = attempt_bits.saturating_sub(1);
148 let mut q = rng.gen_biguint(q_bits);
149 q |= BigUint::one(); if q_bits > 1 {
152 q |= BigUint::one() << (q_bits - 1); }
154
155 if is_probable_prime(&q, 15) {
156 let p = &q * 2u32 + 1u32;
158 let p_bits = p.bits();
159
160 if p_bits >= min_bits && p_bits <= max_bits && is_probable_prime(&p, 15) {
161 return Ok((p, q));
162 }
163 }
164 }
165}
166
167pub fn is_probable_prime(n: &BigUint, k: usize) -> bool {
169 if n <= &BigUint::one() {
170 return false;
171 }
172
173 let two = 2u32.to_biguint().unwrap();
174 let three = 3u32.to_biguint().unwrap();
175
176 if n == &two {
177 return true;
178 }
179 if n == &three {
180 return true;
181 }
182 if n.is_even() {
183 return false;
184 }
185 if n < &two {
186 return false;
187 }
188
189 let mut rng = thread_rng();
190 let n_minus_1 = n - BigUint::one();
191 let (s, d) = factor_powers_of_two(&n_minus_1);
192
193 'witness: for _ in 0..k {
194 let a = if n == &three {
196 two.clone()
197 } else {
198 let upper = n_minus_1.clone();
199 if upper <= two {
200 two.clone()
201 } else {
202 rng.gen_biguint_range(&two, &upper)
203 }
204 };
205
206 let mut x = mod_exp(&a, &d, n);
207
208 if x == BigUint::one() || x == n_minus_1 {
209 continue;
210 }
211
212 for _ in 0..s - 1 {
213 x = mod_exp(&x, &two, n);
214 if x == n_minus_1 {
215 continue 'witness;
216 }
217 }
218
219 return false;
220 }
221
222 true
223}
224
225pub fn factor_powers_of_two(n: &BigUint) -> (u64, BigUint) {
227 let mut s = 0;
228 let mut d = n.clone();
229
230 while d.is_even() {
231 d >>= 1;
232 s += 1;
233 }
234
235 (s, d)
236}
237
238pub fn find_generator(p: &BigUint, q: &BigUint) -> BigUint {
240 let mut rng = thread_rng();
241 let p_minus_1 = p - BigUint::one();
242
243 loop {
244 let g = rng.gen_biguint_range(&2u32.to_biguint().unwrap(), &p_minus_1);
245
246 let g_squared = mod_exp(&g, &2u32.to_biguint().unwrap(), p);
247 let g_to_q = mod_exp(&g, q, p);
248
249 if g_squared != BigUint::one() && g_to_q != BigUint::one() {
250 return g;
251 }
252 }
253}
254
255pub fn random_in_range(n: &BigUint) -> BigUint {
257 let mut rng = thread_rng();
258 rng.gen_biguint_range(&BigUint::one(), n)
259}
260
261#[cfg(test)]
262mod tests {
263 use super::*;
264
265 #[test]
266 fn test_mod_inverse() {
267 let a = 3u32.to_biguint().unwrap();
268 let m = 11u32.to_biguint().unwrap();
269 let inv = mod_inverse(&a, &m).unwrap();
270
271 assert_eq!((a * inv) % m, BigUint::one());
272 }
273
274 #[test]
275 fn test_is_probable_prime() {
276 assert!(is_probable_prime(&2u32.to_biguint().unwrap(), 20));
278 assert!(is_probable_prime(&3u32.to_biguint().unwrap(), 20));
279 assert!(is_probable_prime(&5u32.to_biguint().unwrap(), 20));
280 assert!(is_probable_prime(&7u32.to_biguint().unwrap(), 20));
281 assert!(is_probable_prime(&11u32.to_biguint().unwrap(), 20));
282 assert!(is_probable_prime(&13u32.to_biguint().unwrap(), 20));
283
284 assert!(!is_probable_prime(&4u32.to_biguint().unwrap(), 20));
286 assert!(!is_probable_prime(&6u32.to_biguint().unwrap(), 20));
287 assert!(!is_probable_prime(&8u32.to_biguint().unwrap(), 20));
288 assert!(!is_probable_prime(&9u32.to_biguint().unwrap(), 20));
289 assert!(!is_probable_prime(&10u32.to_biguint().unwrap(), 20));
290 assert!(!is_probable_prime(&12u32.to_biguint().unwrap(), 20));
291 assert!(!is_probable_prime(&15u32.to_biguint().unwrap(), 20));
292 }
293
294 #[test]
295 fn test_safe_prime_generation_lenient() {
296 let result = generate_safe_prime_lenient(512);
298 assert!(
299 result.is_ok(),
300 "Lenient safe prime generation should succeed"
301 );
302
303 if let Ok((p, q)) = result {
304 assert_eq!(p, &q * 2u32 + 1u32);
306
307 assert!(is_probable_prime(&p, 20));
309 assert!(is_probable_prime(&q, 20));
310
311 let p_bits = p.bits();
313 assert!(
314 p_bits >= 504 && p_bits <= 520,
315 "Prime should be close to 512 bits, got {}",
316 p_bits
317 );
318 }
319 }
320}