Skip to main content

vhe/
utils.rs

1//! Utility functions for cryptographic operations
2
3use 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
9/// Modular exponentiation: base^exp mod modulus
10pub fn mod_exp(base: &BigUint, exp: &BigUint, modulus: &BigUint) -> BigUint {
11    base.modpow(exp, modulus)
12}
13
14/// Compute modular inverse using extended Euclidean algorithm
15pub 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    // Convert back to BigUint, handling negative values
23    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
34/// Extended Euclidean algorithm (using BigInt to handle negative intermediate values)
35fn 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
47/// Generate a safe prime (p = 2q + 1 where q is also prime)
48pub 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    // For 512-bit keys, we need more iterations and flexibility
56    let max_iterations = if bit_size <= 512 {
57        500000 // Much higher for small primes
58    } else if bit_size <= 1024 {
59        200000
60    } else {
61        100000
62    };
63
64    let mut iterations = 0;
65
66    // Allow a small range of bit sizes for flexibility
67    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        // Generate a random odd number of approximately the right size
79        // For safe primes, q should be about (bit_size - 1) bits
80        let mut q = rng.gen_biguint(bit_size - 1);
81
82        // Ensure q is odd
83        q |= BigUint::one();
84
85        // Set high bit to ensure minimum size
86        if bit_size > 2 {
87            q |= BigUint::one() << (bit_size - 2);
88        }
89
90        // Quick pre-check: if q is even, skip
91        if q.is_even() {
92            continue;
93        }
94
95        // First check if q is prime (cheaper check)
96        if !is_probable_prime(&q, 20) {
97            continue;
98        }
99
100        // Calculate p = 2q + 1
101        let p = &q * 2u32 + 1u32;
102
103        // Check that p has approximately the right bit size (allow some flexibility)
104        let p_bits = p.bits();
105        if p_bits < min_bits || p_bits > max_bits {
106            continue;
107        }
108
109        // Check if p is also prime
110        if is_probable_prime(&p, 20) {
111            return Ok((p, q));
112        }
113    }
114}
115
116/// Generate a safe prime with more lenient bit size requirements (for easier generation)
117pub 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; // Very high limit for lenient generation
124    let mut iterations = 0;
125
126    // Allow wider range for lenient generation
127    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        // Try different bit sizes near the target
140        let bit_variation = iterations % 17; // Vary the size slightly
141        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(); // Make odd
150
151        if q_bits > 1 {
152            q |= BigUint::one() << (q_bits - 1); // Set high bit
153        }
154
155        if is_probable_prime(&q, 15) {
156            // Slightly fewer rounds for speed
157            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
167/// Miller-Rabin primality test
168pub 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        // For small n, we need to be careful with the range
195        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
225/// Factor out powers of 2 from n
226pub 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
238/// Find a generator for the multiplicative group modulo p
239pub 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
255/// Generate a random element in the range [1, n)
256pub 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        // Known small primes
277        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        // Known composites
285        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        // Test that lenient generation works for 512-bit primes
297        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            // Check that p = 2q + 1
305            assert_eq!(p, &q * 2u32 + 1u32);
306
307            // Check that both are prime
308            assert!(is_probable_prime(&p, 20));
309            assert!(is_probable_prime(&q, 20));
310
311            // Check that bit size is reasonable (within ±8 bits)
312            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}