mirror of
https://github.com/saymrwulf/curve25519-dalek-source.git
synced 2026-09-04 20:24:10 +00:00
Merge remote-tracking branch 'hdevalence/feature/constant-time-k-fold-scalar-mult' into develop
This commit is contained in:
commit
7202ab8e63
2 changed files with 148 additions and 16 deletions
142
src/curve.rs
142
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,75 @@ 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::multiscalar_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 multiscalar_mult<'a, 'b, I, J>(scalars: I, points: J) -> ExtendedPoint
|
||||
where I: IntoIterator<Item = &'a Scalar>,
|
||||
J: IntoIterator<Item = &'b ExtendedPoint>
|
||||
{
|
||||
//assert_eq!(scalars.len(), points.len());
|
||||
|
||||
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();
|
||||
|
||||
// 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();
|
||||
|
||||
// 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();
|
||||
// XXX this algorithm makes no effort to be cache-aware; maybe it could be improved?
|
||||
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)]
|
||||
|
|
@ -1195,7 +1282,7 @@ pub mod vartime {
|
|||
/// 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
|
||||
pub fn multiscalar_mult<'a, 'b, I, J>(scalars: I, points: J) -> ExtendedPoint
|
||||
where I: IntoIterator<Item = &'a Scalar>,
|
||||
J: IntoIterator<Item = &'b ExtendedPoint>
|
||||
{
|
||||
|
|
@ -1627,14 +1714,29 @@ mod test {
|
|||
}
|
||||
|
||||
#[test]
|
||||
fn k_fold_scalar_mult_vs_ed25519py() {
|
||||
fn multiscalar_mult_vs_ed25519py() {
|
||||
let A = A_TIMES_BASEPOINT.decompress().unwrap();
|
||||
let result = vartime::k_fold_scalar_mult(
|
||||
let result = vartime::multiscalar_mult(
|
||||
&[A_SCALAR, B_SCALAR],
|
||||
&[A, constants::ED25519_BASEPOINT]
|
||||
);
|
||||
assert_eq!(result.compress_edwards(), DOUBLE_SCALAR_MULT_RESULT);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn multiscalar_mult_vartime_vs_consttime() {
|
||||
let A = A_TIMES_BASEPOINT.decompress().unwrap();
|
||||
let result_vartime = vartime::multiscalar_mult(
|
||||
&[A_SCALAR, B_SCALAR],
|
||||
&[A, constants::ED25519_BASEPOINT]
|
||||
);
|
||||
let result_consttime = multiscalar_mult(
|
||||
&[A_SCALAR, B_SCALAR],
|
||||
&[A, constants::ED25519_BASEPOINT]
|
||||
);
|
||||
|
||||
assert_eq!(result_vartime.compress_edwards(), result_consttime.compress_edwards());
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "serde")]
|
||||
|
|
@ -1761,6 +1863,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(|| multiscalar_mult(&scalars, &points));
|
||||
}
|
||||
|
||||
mod vartime {
|
||||
use super::super::*;
|
||||
use super::super::test::{A_SCALAR, B_SCALAR, A_TIMES_BASEPOINT};
|
||||
|
|
@ -1787,7 +1901,7 @@ mod bench {
|
|||
//
|
||||
// Since this is a variable-time function, this means the
|
||||
// benchmark is only useful as a ballpark measurement.
|
||||
b.iter(|| vartime::k_fold_scalar_mult(&scalars, &points));
|
||||
b.iter(|| vartime::multiscalar_mult(&scalars, &points));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
22
src/decaf.rs
22
src/decaf.rs
|
|
@ -576,6 +576,24 @@ impl<'a, 'b> Mul<&'b DecafPoint> 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::multiscalar_mult` but is constant-time.
|
||||
///
|
||||
/// # Input
|
||||
///
|
||||
/// A vector of `Scalar`s and a vector of `DecafPoints`. It is an
|
||||
/// error to call this function with two vectors of different lengths.
|
||||
#[cfg(any(feature = "alloc", feature = "std"))]
|
||||
pub fn multiscalar_mult<'a, 'b, I, J>(scalars: I, points: J) -> DecafPoint
|
||||
where I: IntoIterator<Item = &'a Scalar>,
|
||||
J: IntoIterator<Item = &'b DecafPoint>,
|
||||
{
|
||||
let extended_points = points.into_iter().map(|P| &P.0);
|
||||
DecafPoint(curve::multiscalar_mult(scalars, extended_points))
|
||||
}
|
||||
|
||||
/// Precomputation
|
||||
#[derive(Clone)]
|
||||
|
|
@ -682,12 +700,12 @@ pub mod vartime {
|
|||
///
|
||||
/// A vector of `Scalar`s and a vector of `ExtendedPoints`. It is an
|
||||
/// error to call this function with two vectors of different lengths.
|
||||
pub fn k_fold_scalar_mult<'a, 'b, I, J>(scalars: I, points: J) -> DecafPoint
|
||||
pub fn multiscalar_mult<'a, 'b, I, J>(scalars: I, points: J) -> DecafPoint
|
||||
where I: IntoIterator<Item = &'a Scalar>,
|
||||
J: IntoIterator<Item = &'b DecafPoint>
|
||||
{
|
||||
let extended_points = points.into_iter().map(|P| &P.0);
|
||||
DecafPoint(curve::vartime::k_fold_scalar_mult(scalars, extended_points))
|
||||
DecafPoint(curve::vartime::multiscalar_mult(scalars, extended_points))
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
Loading…
Reference in a new issue