diff --git a/src/scalar.rs b/src/scalar.rs index d41d22b..eb14546 100644 --- a/src/scalar.rs +++ b/src/scalar.rs @@ -19,6 +19,8 @@ use core::ops::{Sub, SubAssign}; use core::ops::{Mul, MulAssign}; use core::ops::{Index}; use core::cmp::{Eq, PartialEq}; +use core::iter::Product; +use core::borrow::Borrow; #[cfg(feature = "std")] use rand::Rng; @@ -293,6 +295,18 @@ impl<'de> Deserialize<'de> for Scalar { } } +impl Product for Scalar +where + T: Borrow +{ + fn product(iter: I) -> Self + where + I: Iterator + { + iter.fold(Scalar::one(), |acc, item| acc * item.borrow()) + } +} + impl Scalar { /// Return a `Scalar` chosen uniformly at random using a user-provided RNG. /// @@ -912,6 +926,37 @@ mod test { assert_eq!(should_be_X_times_Y, X_TIMES_Y); } + #[allow(non_snake_case)] + #[test] + fn impl_product() { + // Test that product works for non-empty iterators + let X_Y_vector = vec![X, Y]; + let should_be_X_times_Y: Scalar = X_Y_vector.iter().product(); + assert_eq!(should_be_X_times_Y, X_TIMES_Y); + + // Test that product works for the empty iterator + let one = Scalar::one(); + let empty_vector = vec![]; + let should_be_one: Scalar = empty_vector.iter().product(); + assert_eq!(should_be_one, one); + + // Test that product works for iterators where Item = Scalar + let xs = [Scalar::from_u64(2); 10]; + let ys = [Scalar::from_u64(3); 10]; + // now zs is an iterator with Item = Scalar + let zs = xs.iter().zip(ys.iter()).map(|(x,y)| x * y); + + let x_prod: Scalar = xs.iter().product(); + let y_prod: Scalar = ys.iter().product(); + let z_prod: Scalar = zs.product(); + + assert_eq!(x_prod, Scalar::from_u64(1024)); + assert_eq!(y_prod, Scalar::from_u64(59049)); + assert_eq!(z_prod, Scalar::from_u64(60466176)); + assert_eq!(x_prod * y_prod, z_prod); + + } + #[test] fn square() { let expected = &X * &X;