This commit is contained in:
Oleg Andreev 2019-05-21 13:31:50 -07:00
parent b52c2053c1
commit 42648aa460
5 changed files with 401 additions and 16 deletions

View file

@ -26,3 +26,6 @@ pub mod straus;
#[cfg(feature = "alloc")]
pub mod precomputed_straus;
#[cfg(feature = "alloc")]
pub mod pippenger;

View file

@ -0,0 +1,207 @@
// -*- mode: rust; -*-
//
// This file is part of curve25519-dalek.
// Copyright (c) 2016-2018 Isis Lovecruft, Henry de Valence
// See LICENSE for licensing information.
//
// Authors:
// - Isis Agora Lovecruft <isis@patternsinthevoid.net>
// - Henry de Valence <hdevalence@hdevalence.ca>
// - Oleg Andreev <oleganza@gmail.com>
//! Implementation of a variant of Pippenger's algorithm.
#![allow(non_snake_case)]
use core::borrow::Borrow;
use edwards::EdwardsPoint;
use scalar::Scalar;
use traits::VartimeMultiscalarMul;
#[allow(unused_imports)]
use prelude::*;
/// Implements a version of Pippenger's algorithm.
///
/// The algorithm works as follows:
///
/// Let `n` be a number of point-scalar pairs.
/// Let `w` be a window of bits (6..8, chosen based on `n`, see cost factor).
///
/// 1. Prepare `2^(w-1) - 1` buckets with indices `[1..2^(w-1))` initialized with identity points.
/// Bucket 0 is not needed as it would contain points multiplied by 0.
/// 2. Convert scalars to a radix-`2^w` representation with signed digits in `[-2^w/2, 2^w/2]`.
/// Note: only the last digit may equal `2^w/2`.
/// 3. Starting with the last window, for each point `i=[0..n)` add it to a a bucket indexed by
/// the point's scalar's value in the window.
/// 4. Once all points in a window are sorted into buckets, add buckets by multiplying each
/// by their index. Efficient way of doing it is to start with the last bucket and compute two sums:
/// intermediate sum from the last to the first, and the full sum made of all intermediate sums.
/// 5. Shift the resulting sum of buckets by `w` bits by using `w` doublings.
/// 6. Add to the return value.
/// 7. Repeat the loop.
///
/// Approximate cost w/o wNAF optimizations (A = addition, D = doubling):
///
/// ```ascii
/// cost = (n*A + 2*(2^w/2)*A + w*D + A)*256/w
/// | | | | |
/// | | | | looping over 256/w windows
/// | | | adding to the result
/// sorting points | shifting the sum by w bits (to the next window, starting from last window)
/// one by one |
/// into buckets adding/subtracting all buckets
/// multiplied by their indexes
/// using a sum of intermediate sums
/// ```
///
/// For large `n`, dominant factor is (n*256/w) additions.
/// However, if `w` is too big and `n` is not too big, then `(2^w/2)*A` could dominate.
/// Therefore, the optimal choice of `w` grows slowly as `n` grows.
///
pub struct Pippenger;
#[cfg(any(feature = "alloc", feature = "std"))]
impl VartimeMultiscalarMul for Pippenger {
type Point = EdwardsPoint;
fn optional_multiscalar_mul<I, J>(scalars: I, points: J) -> Option<EdwardsPoint>
where
I: IntoIterator,
I::Item: Borrow<Scalar>,
J: IntoIterator<Item = Option<EdwardsPoint>>,
{
use backend::serial::curve_models::{ProjectiveNielsPoint};
use traits::Identity;
let mut scalars = scalars.into_iter();
let size = scalars.by_ref().size_hint().0;
// Digit width in bits. As digit width grows,
// number of point additions goes down, but amount of
// buckets and bucket additions grows exponentially.
let w = if size < 500 {
6
} else if size < 800 {
7
} else {
8
};
let max_digit: usize = 1 << w;
let digits_count: usize = (256 + w - 1) / w; // == ceil(256/w)
let buckets_count: usize = max_digit / 2; // digits are signed+centered hence 2^w/2, excluding 0-th bucket
// Collect optimized scalars and points in buffers for repeated access
// (scanning the whole set per digit position).
let scalars = scalars.into_iter()
.map(|s| s.borrow().to_pippenger_radix(w).0 )
.collect::<Vec<_>>();
let points: Vec<ProjectiveNielsPoint> = match points
.into_iter()
.map(|p| p.map(|P| P.to_projective_niels()))
.collect::<Option<Vec<_>>>() {
Some(x) => x,
None => return None,
};
// Prepare 2^w/2 buckets.
// buckets[i] corresponds to a multiplication factor (i+1).
let mut buckets: Vec<_> = (0..buckets_count)
.map(|_| EdwardsPoint::identity())
.collect();
let columns: Vec<_> = (0..digits_count).map(|digit_index| {
// Clear the buckets when processing another digit.
for i in 0..buckets_count {
buckets[i] = EdwardsPoint::identity();
}
// Iterate over pairs of (point, scalar)
// and add/sub the point to the corresponding bucket.
// Note: if we add support for precomputed lookup tables,
// we'll be adding/subtractiong point premultiplied by `digits[i]` to buckets[0].
for (digits, pt) in scalars.iter().zip(points.iter()) {
let digit = digits[digit_index];
if digit > 0 {
let b = (digit - 1) as usize;
buckets[b] = (&buckets[b] + pt).to_extended();
} else if digit < 0 {
let b = (-digit - 1) as usize;
buckets[b] = (&buckets[b] - pt).to_extended();
}
}
// Add the buckets applying the multiplication factor to each bucket.
// The most efficient way to do that is to have a single sum with two running sums:
// an intermediate sum from last bucket to the first, and a sum of intermediate sums.
//
// For example, to add buckets 1*A, 2*B, 3*C we need to add these points:
// C
// C B
// C B A Sum = C + (C+B) + (C+B+A)
let mut buckets_intermediate_sum = buckets[buckets_count - 1];
let mut buckets_sum = buckets[buckets_count - 1];
for i in (0..(buckets_count - 1)).rev() {
buckets_intermediate_sum += buckets[i];
buckets_sum += buckets_intermediate_sum;
}
buckets_sum
})
.collect();
// ^ Note: we collect points because if we chain .rev().fold()
// then the .map() will run in reversed order, producing incorrect digit values
// (they can only be produced in lo->hi order).
// Add the intermediate per-digit results in hi->lo order
// so that we can minimize doublings.
Some(columns[0..(digits_count - 1)].iter().rev().fold(
columns[digits_count - 1],
|total, &p| total.mul_by_pow_2(w as u32) + p,
))
}
}
#[cfg(test)]
mod test {
use super::*;
use constants;
use scalar::Scalar;
#[test]
fn test_vartime_pippenger() {
// Reuse points across different tests
let mut n = 512;
let x = Scalar::from(2128506u64).invert();
let y = Scalar::from(4443282u64).invert();
let points: Vec<_> = (0..n)
.map(|i| {
constants::ED25519_BASEPOINT_POINT * Scalar::from(1 + i as u64)
})
.collect();
let scalars: Vec<_> = (0..n)
.map(|i| x + (Scalar::from(i as u64)*y)) // fast way to make ~random but deterministic scalars
.collect();
let premultiplied: Vec<EdwardsPoint> = scalars
.iter()
.zip(points.iter())
.map(|(sc, pt)| sc * pt)
.collect();
while n > 0 {
let scalars = &scalars[0..n].to_vec();
let points = &points[0..n].to_vec();
let control: EdwardsPoint = premultiplied[0..n].iter().sum();
let subject = Pippenger::vartime_multiscalar_mul(scalars.clone(), points.clone());
assert_eq!(subject.compress(), control.compress());
n = n / 2;
}
}
}

View file

@ -17,3 +17,6 @@ pub mod straus;
#[cfg(feature = "alloc")]
pub mod precomputed_straus;
#[cfg(feature = "alloc")]
pub mod pippenger;

View file

@ -0,0 +1,172 @@
// -*- mode: rust; -*-
//
// This file is part of curve25519-dalek.
// Copyright (c) 2016-2018 Isis Lovecruft, Henry de Valence
// See LICENSE for licensing information.
//
// Authors:
// - Isis Agora Lovecruft <isis@patternsinthevoid.net>
// - Henry de Valence <hdevalence@hdevalence.ca>
// - Oleg Andreev <oleganza@gmail.com>
#![allow(non_snake_case)]
use core::borrow::Borrow;
use clear_on_drop::ClearOnDrop;
use backend::vector::{CachedPoint, ExtendedPoint};
use edwards::EdwardsPoint;
use scalar::Scalar;
use window::{LookupTable, NafLookupTable5};
use traits::{Identity, MultiscalarMul, VartimeMultiscalarMul};
#[allow(unused_imports)]
use prelude::*;
/// Implements a version of Pippenger's algorithm.
///
/// See the documentation in the serial `scalar_mul::pippenger` module for details.
pub struct Pippenger;
#[cfg(any(feature = "alloc", feature = "std"))]
impl VartimeMultiscalarMul for Pippenger {
type Point = EdwardsPoint;
fn optional_multiscalar_mul<I, J>(scalars: I, points: J) -> Option<EdwardsPoint>
where
I: IntoIterator,
I::Item: Borrow<Scalar>,
J: IntoIterator<Item = Option<EdwardsPoint>>,
{
let mut scalars = scalars.into_iter();
let size = scalars.by_ref().size_hint().0;
let w = if size < 500 {
6
} else if size < 800 {
7
} else {
8
};
let max_digit: usize = 1 << w;
let digits_count: usize = (256 + w - 1) / w; // == ceil(256/w)
let buckets_count: usize = max_digit / 2; // digits are signed+centered hence 2^w/2, excluding 0-th bucket
// Collect optimized scalars and points in buffers for repeated access
// (scanning the whole set per digit position).
let scalars = scalars.into_iter()
.map(|s| s.borrow().to_pippenger_radix(w).0 )
.collect::<Vec<_>>();
let points: Vec<CachedPoint> = match points
.into_iter()
.map(|p| p.map(|P| CachedPoint::from(ExtendedPoint::from(P))))
.collect::<Option<Vec<_>>>() {
Some(x) => x,
None => return None,
};
// Prepare 2^w/2 buckets.
// buckets[i] corresponds to a multiplication factor (i+1).
let mut buckets: Vec<_> = (0..buckets_count)
.map(|_| ExtendedPoint::identity())
.collect();
let columns: Vec<ExtendedPoint> = (0..digits_count).map(|digit_index| {
// Clear the buckets when processing another digit.
for i in 0..buckets_count {
buckets[i] = ExtendedPoint::identity();
}
// Iterate over pairs of (point, scalar)
// and add/sub the point to the corresponding bucket.
// Note: if we add support for precomputed lookup tables,
// we'll be adding/subtractiong point premultiplied by `digits[i]` to buckets[0].
for (digits, pt) in scalars.iter().zip(points.iter()) {
let digit = digits[digit_index];
if digit > 0 {
let b = (digit - 1) as usize;
buckets[b] = &buckets[b] + pt;
} else if digit < 0 {
let b = (-digit - 1) as usize;
buckets[b] = &buckets[b] - pt;
}
}
// Add the buckets applying the multiplication factor to each bucket.
// The most efficient way to do that is to have a single sum with two running sums:
// an intermediate sum from last bucket to the first, and a sum of intermediate sums.
//
// For example, to add buckets 1*A, 2*B, 3*C we need to add these points:
// C
// C B
// C B A Sum = C + (C+B) + (C+B+A)
let mut buckets_intermediate_sum = buckets[buckets_count - 1];
let mut buckets_sum = buckets[buckets_count - 1];
for i in (0..(buckets_count - 1)).rev() {
buckets_intermediate_sum = &buckets_intermediate_sum + &buckets[i];
buckets_sum = &buckets_sum + &buckets_intermediate_sum;
}
buckets_sum
})
.collect();
// ^ Note: we collect points because if we chain .rev().fold()
// then the .map() will run in reversed order, producing incorrect digit values
// (they can only be produced in lo->hi order).
// Add the intermediate per-digit results in hi->lo order
// so that we can minimize doublings.
Some(
columns[0..(digits_count - 1)]
.iter()
.rev()
.fold(columns[digits_count - 1], |total, &p| {
&total.mul_by_pow_2(w as u32) + &p
})
.into(),
)
}
}
#[cfg(test)]
mod test {
use super::*;
use constants;
use scalar::Scalar;
#[test]
fn test_vartime_pippenger() {
// Reuse points across different tests
let mut n = 512;
let x = Scalar::from(2128506u64).invert();
let y = Scalar::from(4443282u64).invert();
let points: Vec<_> = (0..n)
.map(|i| {
constants::ED25519_BASEPOINT_POINT * Scalar::from(1 + i as u64)
})
.collect();
let scalars: Vec<_> = (0..n)
.map(|i| x + (Scalar::from(i as u64)*y)) // fast way to make ~random but deterministic scalars
.collect();
let premultiplied: Vec<EdwardsPoint> = scalars
.iter()
.zip(points.iter())
.map(|(sc, pt)| sc * pt)
.collect();
while n > 0 {
let scalars = &scalars[0..n].to_vec();
let points = &points[0..n].to_vec();
let control: EdwardsPoint = premultiplied[0..n].iter().sum();
let subject = Pippenger::vartime_multiscalar_mul(scalars.clone(), points.clone());
assert_eq!(subject.compress(), control.compress());
n = n / 2;
}
}
}

View file

@ -972,18 +972,18 @@ impl Scalar {
///
/// ## Scalar representation
///
/// Radix \\(2\^r\\), with \\(n = ceil(256/r)\\) coefficients in \\([-(2\^r)/2,(2\^r)/2)\\),
/// Radix \\(2\^w\\), with \\(n = ceil(256/w)\\) coefficients in \\([-(2\^w)/2,(2\^w)/2)\\),
/// i.e., scalar is represented using digits \\(a\_i\\) such that
/// $$
/// a = a\_0 + a\_1 2\^1r + \cdots + a_{n-1} 2\^{r*(n-1)},
/// a = a\_0 + a\_1 2\^1w + \cdots + a_{n-1} 2\^{w*(n-1)},
/// $$
/// with \\(-2\^r/2 \leq a_i < 2\^r/2\\) for \\(0 \leq i < (n-1)\\) and \\(-2\^r/2 \leq a_{n-1} \leq 2\^r/2\\).
/// with \\(-2\^w/2 \leq a_i < 2\^w/2\\) for \\(0 \leq i < (n-1)\\) and \\(-2\^w/2 \leq a_{n-1} \leq 2\^w/2\\).
///
pub(crate) fn to_pippenger_radix(&self, r: usize) -> ([i8; 43], usize) {
debug_assert!(r >= 6);
debug_assert!(r <= 8);
pub(crate) fn to_pippenger_radix(&self, w: usize) -> ([i8; 43], usize) {
debug_assert!(w >= 6);
debug_assert!(w <= 8);
let digits_count = (256 + r - 1)/r as usize;
let digits_count = (256 + w - 1)/w as usize;
debug_assert!(digits_count <= 43);
use byteorder::{ByteOrder, LittleEndian};
@ -992,20 +992,20 @@ impl Scalar {
let mut scalar64x4 = [0u64; 4];
LittleEndian::read_u64_into(&self.bytes, &mut scalar64x4[0..4]);
let radix: u64 = 1 << r;
let radix: u64 = 1 << w;
let window_mask: u64 = radix - 1;
let mut carry = 0u64;
let mut digits = [0i8; 43];
for i in 0..digits_count {
// Construct a buffer of bits of the scalar, starting at `bit_offset`.
let bit_offset = i*r;
let bit_offset = i*w;
let u64_idx = bit_offset / 64;
let bit_idx = bit_offset % 64;
// Read the bits from the scalar
let bit_buf: u64;
if bit_idx < 64 - r || u64_idx == 3 {
if bit_idx < 64 - w || u64_idx == 3 {
// This window's bits are contained in a single u64,
// or it's the last u64 anyway.
bit_buf = scalar64x4[u64_idx] >> bit_idx;
@ -1018,8 +1018,8 @@ impl Scalar {
let coef = carry + (bit_buf & window_mask); // coef = [0, 2^r)
// Recenter coefficients from [0,2^r) to [-2^r/2, 2^r/2)
carry = (coef + (radix/2) as u64) >> r;
digits[i] = ((coef as i64) - (carry << r) as i64) as i8;
carry = (coef + (radix/2) as u64) >> w;
digits[i] = ((coef as i64) - (carry << w) as i64) as i8;
}
// Apply the resulting carry to the last digit
@ -1030,7 +1030,7 @@ impl Scalar {
// we allow the last word to touch the value 2^r/2.
// XXX: make sure tests cover this case, so the carry is non-zero and this line matters.
// Maybe it never happens to be non-zero for r=6/7/8?...
digits[digits_count-1] += (carry << r) as i8;
digits[digits_count-1] += (carry << w) as i8;
(digits, digits_count)
}
@ -1517,11 +1517,11 @@ mod test {
use std::iter;
// For each valid radix it tests that 1000 random-ish scalars can be restored
// from the produced representation precisely.
for r in 6..9 {
for w in 6..9 {
for scalar in (2..100).map(|s| Scalar::from(s as u64).invert() ).chain(iter::once(-Scalar::one())) {
let (digits, digits_count) = scalar.to_pippenger_radix(r);
let (digits, digits_count) = scalar.to_pippenger_radix(w);
let radix = Scalar::from((1<<r) as u64);
let radix = Scalar::from((1<<w) as u64);
let mut term = Scalar::one();
let mut recovered_scalar = Scalar::zero();
for digit in &digits[0..digits_count] {