From b20ccbd6854e5c487e7f06d4f3671b27ba6821b8 Mon Sep 17 00:00:00 2001 From: Henry de Valence Date: Tue, 16 Jan 2018 15:41:05 -0800 Subject: [PATCH] Implement batched encoding for RistrettoPoints --- src/ristretto.rs | 118 +++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 118 insertions(+) diff --git a/src/ristretto.rs b/src/ristretto.rs index 61b451c..50f9b96 100644 --- a/src/ristretto.rs +++ b/src/ristretto.rs @@ -666,6 +666,86 @@ impl RistrettoPoint { CompressedRistretto(s.to_bytes()) } + /// Double-and-compress a batch of points. + pub fn double_and_compress_batch<'a, I>(points: I) -> Vec + where I: IntoIterator + { + #[derive(Copy, Clone, Debug)] + struct BatchCompressState { + e: FieldElement, + f: FieldElement, + g: FieldElement, + h: FieldElement, + eg: FieldElement, + fh: FieldElement, + } + + impl BatchCompressState { + fn efgh(&self) -> FieldElement { + &self.eg * &self.fh + } + } + + impl<'a> From<&'a RistrettoPoint> for BatchCompressState { + fn from(P: &'a RistrettoPoint) -> BatchCompressState { + let XX = P.0.X.square(); + let YY = P.0.Y.square(); + let ZZ = P.0.Z.square(); + let dTT = &P.0.T.square() * &constants::EDWARDS_D; + + let e = &P.0.X * &(&P.0.Y + &P.0.Y); // = 2*X*Y + let f = &ZZ + &dTT; // = Z^2 + d*T^2 + let g = &YY + &XX; // = Y^2 - a*X^2 + let h = &ZZ - &dTT; // = Z^2 - d*T^2 + + let eg = &e * &g; + let fh = &f * &h; + + BatchCompressState{ e: e, f: f, g: g, h: h, eg: eg, fh: fh } + } + } + + let states: Vec = points.into_iter().map(|P| BatchCompressState::from(P)).collect(); + + let mut invs: Vec = states.iter().map(|state| state.efgh()).collect(); + + FieldElement::batch_invert(&mut invs[..]); + + states.iter().zip(invs.iter()).map(|(state, inv): (&BatchCompressState, &FieldElement)| { + let Zinv = &state.eg * &inv; + let Tinv = &state.fh * &inv; + + let mut magic = constants::INVSQRT_A_MINUS_D; + + let negcheck1 = (&state.eg * &Zinv).is_negative(); + + let mut e = state.e; + let mut g = state.g; + let mut h = state.h; + + let minus_e = -&e; + let f_times_sqrta = &state.f * &constants::SQRT_M1; + + e.conditional_assign(&state.g, negcheck1); + g.conditional_assign(&minus_e, negcheck1); + h.conditional_assign(&f_times_sqrta, negcheck1); + + magic.conditional_assign(&constants::SQRT_M1, negcheck1); + + let negcheck2 = (&(&h * &e) * &Zinv).is_negative(); + + g.conditional_negate(negcheck2); + + let mut s = &(&h - &g) * &(&magic * &(&g * &Tinv)); + + let s_is_negative = s.is_negative(); + s.conditional_negate(s_is_negative); + + CompressedRistretto(s.to_bytes()) + }).collect() + } + + /// Return the coset self + E[4], for debugging. fn coset4(&self) -> [ExtendedPoint; 4] { [ self.0 @@ -1216,6 +1296,20 @@ mod test { } } + #[test] + fn double_and_compress_1024_random_points() { + let mut rng = OsRng::new().unwrap(); + + let points: Vec = + (0..1024).map(|_| RistrettoPoint::random(&mut rng)).collect(); + + let compressed = RistrettoPoint::double_and_compress_batch(&points); + + for (P, P2_compressed) in points.iter().zip(compressed.iter()) { + assert_eq!(*P2_compressed, (P + P).compress()); + } + } + #[test] fn random_is_valid() { let mut rng = OsRng::new().unwrap(); @@ -1254,4 +1348,28 @@ mod bench { let P = B * &Scalar::random(&mut rng); b.iter(|| P.compress()); } + + fn double_and_compress_n_random_points(n: usize, b: &mut Bencher) { + let mut rng = OsRng::new().unwrap(); + + let points: Vec = + (0..n).map(|_| RistrettoPoint::random(&mut rng)).collect(); + + b.iter(|| RistrettoPoint::double_and_compress_batch(&points) ); + } + + #[bench] + fn double_and_compress_16_random_points(b: &mut Bencher) { + double_and_compress_n_random_points(16, b); + } + + #[bench] + fn double_and_compress_128_random_points(b: &mut Bencher) { + double_and_compress_n_random_points(128, b); + } + + #[bench] + fn double_and_compress_1024_random_points(b: &mut Bencher) { + double_and_compress_n_random_points(1024, b); + } }