From c18627f7c2d103a5dd232cf79b8202ade5361848 Mon Sep 17 00:00:00 2001 From: Henry de Valence Date: Thu, 4 May 2017 00:02:29 -0700 Subject: [PATCH] Generalize k_fold_scalar_mult --- src/curve.rs | 20 ++++++++++++-------- src/decaf.rs | 9 +++++---- 2 files changed, 17 insertions(+), 12 deletions(-) diff --git a/src/curve.rs b/src/curve.rs index ca3e4a1..78ff254 100644 --- a/src/curve.rs +++ b/src/curve.rs @@ -1066,12 +1066,15 @@ 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(scalars: &Vec, - points: &Vec) -> ExtendedPoint { - assert_eq!(scalars.len(), points.len()); + 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()); - let nafs: Vec<_> = scalars.iter().map(|c| c.non_adjacent_form()).collect(); - let odd_multiples: Vec<_> = points.iter().map(|P| OddMultiples::create(&P)).collect(); + let nafs: Vec<_> = scalars.into_iter() + .map(|c| c.non_adjacent_form()).collect(); + let odd_multiples: Vec<_> = points.into_iter() + .map(|P| OddMultiples::create(P)).collect(); let mut r = ProjectivePoint::identity(); @@ -1471,9 +1474,10 @@ mod test { #[test] fn k_fold_scalar_mult_vs_ed25519py() { let A = A_TIMES_BASEPOINT.decompress().unwrap(); - let points = vec![A,constants::ED25519_BASEPOINT]; - let scalars = vec![A_SCALAR, B_SCALAR]; - let result = vartime::k_fold_scalar_mult(&scalars, &points); + let result = vartime::k_fold_scalar_mult( + &[A_SCALAR, B_SCALAR], + &[A, constants::ED25519_BASEPOINT] + ); assert_eq!(result.compress_edwards(), DOUBLE_SCALAR_MULT_RESULT); } } diff --git a/src/decaf.rs b/src/decaf.rs index 9d1a830..4fba23f 100644 --- a/src/decaf.rs +++ b/src/decaf.rs @@ -318,10 +318,11 @@ 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(scalars: &Vec, - points: &Vec) -> DecafPoint { - let extended_points: Vec = points.iter().map(|P| P.0).collect(); - DecafPoint(curve::vartime::k_fold_scalar_mult(scalars, &extended_points)) + pub fn k_fold_scalar_mult<'a,'b,I,J>(scalars: I, points: J) -> DecafPoint + where I: IntoIterator, J: IntoIterator + { + let extended_points = points.into_iter().map(|P| &P.0); + DecafPoint(curve::vartime::k_fold_scalar_mult(scalars, extended_points)) } }