From 9c4046c3d9a95a4ef3a8152a90411a3c9b19c59a Mon Sep 17 00:00:00 2001 From: Henry de Valence Date: Sun, 30 Jul 2017 21:13:56 -0700 Subject: [PATCH] Add constant-time k-fold scalar multiplication --- src/curve.rs | 133 +++++++++++++++++++++++++++++++++++++++++++++++---- 1 file changed, 123 insertions(+), 10 deletions(-) diff --git a/src/curve.rs b/src/curve.rs index deac819..1cf1739 100644 --- a/src/curve.rs +++ b/src/curve.rs @@ -884,20 +884,38 @@ impl<'a, 'b> Mul<&'b Scalar> for &'a ExtendedPoint { /// Uses a window of size 4. Note: for scalar multiplication of /// the basepoint, `basepoint_mult` is approximately 4x faster. fn mul(self, scalar: &'b Scalar) -> ExtendedPoint { - let A = self.to_projective_niels(); - let mut As: [ProjectiveNielsPoint; 8] = [A; 8]; + // Construct a lookup table of [P,2P,3P,4P,5P,6P,7P,8P] + let P = self.to_projective_niels(); + let mut lookup_table: [ProjectiveNielsPoint; 8] = [P; 8]; for i in 0..7 { - As[i+1] = (self + &As[i]).to_extended().to_projective_niels(); + lookup_table[i+1] = (self + &lookup_table[i]) + .to_extended().to_projective_niels(); } - let e = scalar.to_radix_16(); - let mut h = ExtendedPoint::identity(); - let mut t: CompletedPoint; + + // Setting s = scalar, compute + // + // s = s_0 + s_1*16^1 + ... + s_63*16^63, + // + // with `-8 ≤ s_i < 8` for `0 ≤ i < 63` and `-8 ≤ s_63 ≤ 8`. + let scalar_digits = scalar.to_radix_16(); + + // Compute s*P as + // + // s*P = P*(s_0 + s_1*16^1 + s_2*16^2 + ... + s_63*16^63) + // s*P = P*s_0 + P*s_1*16^1 + P*s_2*16^2 + ... + P*s_63*16^63 + // s*P = P*s_0 + 16*(P*s_1 + 16*(P*s_2 + 16*( ... + P*s_63)...)) + // + // We sum right-to-left. + let mut Q = ExtendedPoint::identity(); for i in (0..64).rev() { - h = h.mult_by_pow_2(4); - t = &h + &select_precomputed_point(e[i], &As); - h = t.to_extended(); + // Q = 16*Q + Q = Q.mult_by_pow_2(4); + // R = s_i * Q + let R = select_precomputed_point(scalar_digits[i], &lookup_table); + // Q = Q + R + Q = (&Q + &R).to_extended(); } - h + Q } } @@ -913,6 +931,74 @@ impl<'a, 'b> Mul<&'b ExtendedPoint> for &'a Scalar { } } +/// Given a vector of (possibly secret) scalars and a vector of +/// (possibly secret) points, compute `c_1 P_1 + ... + c_n P_n`. +/// +/// This function has the same behaviour as +/// `vartime::k_fold_scalar_mult` but is constant-time. +/// +/// # Input +/// +/// A vector of `Scalar`s and a vector of `ExtendedPoints`. It is an +/// error to call this function with two vectors of different lengths. +#[cfg(any(feature = "alloc", feature = "std"))] +pub fn k_fold_scalar_mult<'a, 'b, I, J>(scalars: I, points: J) -> ExtendedPoint + where I: IntoIterator, + J: IntoIterator +{ + //assert_eq!(scalars.len(), points.len()); + + // Setting s_i = i-th scalar, compute + // + // s_i = s_{i,0} + s_{i,1}*16^1 + ... + s_{i,63}*16^63, + // + // with `-8 ≤ s_{i,j} < 8` for `0 ≤ j < 63` and `-8 ≤ s_{i,63} ≤ 8`. + let scalar_digits_list: Vec<_> = scalars.into_iter() + .map(|c| c.to_radix_16()).collect(); + + let lookup_tables: Vec<_> = points.into_iter() + .map(|P_i| { + // Construct a lookup table of [P_i,2*P_i,3*P_i,4*P_i,5*P_i,6*P_i,7*P_i] + let mut lookup_table = [P_i.to_projective_niels(); 8]; + for j in 0..7 { + lookup_table[j+1] = (P_i + &lookup_table[j]) + .to_extended().to_projective_niels(); + } + lookup_table + }).collect(); + + // Compute s_1*P_1 + ... + s_n*P_n: since + // + // s_i*P_i = P_i*(s_{i,0} + s_{i,1}*16^1 + ... + s_{i,63}*16^63) + // s_i*P_i = P_i*s_{i,0} + P_i*s_{i,1}*16^1 + ... + P_i*s_{i,63}*16^63 + // s_i*P_i = P_i*s_{i,0} + 16*(P_i*s_{i,1} + 16*( ... + 16*P_i*s_{i,63})...) + // + // we have the two-dimensional sum + // + // s_1*P_1 = P_1*s_{1,0} + 16*(P_1*s_{1,1} + 16*( ... + 16*P_1*s_{1,63})...) + // + s_2*P_2 = + P_2*s_{2,0} + 16*(P_2*s_{2,1} + 16*( ... + 16*P_2*s_{2,63})...) + // ... + // + s_n*P_n = + P_n*s_{n,0} + 16*(P_n*s_{n,1} + 16*( ... + 16*P_n*s_{n,63})...) + // + // We sum column-wise top-to-bottom, then right-to-left, + // multiplying by 16 only once per column. + // + // This provides the speedup over doing n independent scalar + // mults: we perform 63 multiplications by 16 instead of 63*n + // multiplications, saving 252*(n-1) doublings. + let mut Q = ExtendedPoint::identity(); + for j in (0..64).rev() { + Q = Q.mult_by_pow_2(4); + let it = scalar_digits_list.iter().zip(lookup_tables.iter()); + for (s_i, lookup_table_i) in it { + // R_i = s_{i,j} * P_i + let R_i = select_precomputed_point(s_i[j], lookup_table_i); + // Q = Q + R_i + Q = (&Q + &R_i).to_extended(); + } + } + Q +} /// Precomputation #[derive(Clone)] @@ -1635,6 +1721,21 @@ mod test { ); assert_eq!(result.compress_edwards(), DOUBLE_SCALAR_MULT_RESULT); } + + #[test] + fn k_fold_scalar_mult_vartime_vs_consttime() { + let A = A_TIMES_BASEPOINT.decompress().unwrap(); + let result_vartime = vartime::k_fold_scalar_mult( + &[A_SCALAR, B_SCALAR], + &[A, constants::ED25519_BASEPOINT] + ); + let result_consttime = k_fold_scalar_mult( + &[A_SCALAR, B_SCALAR], + &[A, constants::ED25519_BASEPOINT] + ); + + assert_eq!(result_vartime.compress_edwards(), result_consttime.compress_edwards()); + } } #[cfg(feature = "serde")] @@ -1761,6 +1862,18 @@ mod bench { b.iter(|| EdwardsBasepointTable::create(&aB)); } + #[bench] + fn ten_fold_scalar_mult(b: &mut Bencher) { + let mut csprng: OsRng = OsRng::new().unwrap(); + // Create 10 random scalars + let scalars: Vec<_> = (0..10).map(|_| Scalar::random(&mut csprng)).collect(); + // Create 10 points (by doing scalar mults) + let B = &constants::ED25519_BASEPOINT_TABLE; + let points: Vec<_> = scalars.iter().map(|s| B * &s).collect(); + + b.iter(|| k_fold_scalar_mult(&scalars, &points)); + } + mod vartime { use super::super::*; use super::super::test::{A_SCALAR, B_SCALAR, A_TIMES_BASEPOINT};