diff --git a/.github/workflows/main.yml b/.github/workflows/main.yml index 15683e6..e3d989d 100644 --- a/.github/workflows/main.yml +++ b/.github/workflows/main.yml @@ -32,7 +32,6 @@ jobs: profile: minimal toolchain: ${{ matrix.toolchain }} override: true - components: rustfmt, clippy - name: Run cargo test uses: actions-rs/cargo@v1 @@ -83,7 +82,7 @@ jobs: profile: minimal toolchain: stable override: true - components: rustfmt, clippy + components: clippy - name: Run cargo clippy uses: actions-rs/cargo@v1 @@ -113,7 +112,7 @@ jobs: profile: minimal toolchain: stable override: true - components: rustfmt, clippy + components: rustfmt - name: Run cargo fmt uses: actions-rs/cargo@v1 diff --git a/src/lib.rs b/src/lib.rs index 077c375..f85d7f0 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -78,7 +78,7 @@ //! //! let mut client_rng = OsRng; //! let client_blind_result = NonVerifiableClient::::blind( -//! b"input", +//! b"input".to_vec(), //! &mut client_rng, //! ).expect("Unable to construct client"); //! ``` @@ -99,17 +99,16 @@ //! # //! # let mut client_rng = OsRng; //! # let client_blind_result = NonVerifiableClient::::blind( -//! # b"input", +//! # b"input".to_vec(), //! # &mut client_rng, //! # ).expect("Unable to construct client"); //! # use voprf::NonVerifiableServer; //! # let mut server_rng = OsRng; //! # let server = NonVerifiableServer::::new(&mut server_rng) //! # .expect("Unable to construct server"); -//! use voprf::Metadata; //! let server_evaluate_result = server.evaluate( //! client_blind_result.message, -//! &Metadata::none(), +//! None, //! ).expect("Unable to perform server evaluate"); //! ``` //! @@ -127,7 +126,7 @@ //! # //! # let mut client_rng = OsRng; //! # let client_blind_result = NonVerifiableClient::::blind( -//! # b"input", +//! # b"input".to_vec(), //! # &mut client_rng, //! # ).expect("Unable to construct client"); //! # use voprf::NonVerifiableServer; @@ -136,12 +135,11 @@ //! # .expect("Unable to construct server"); //! # let server_evaluate_result = server.evaluate( //! # client_blind_result.message, -//! # &Metadata::none(), +//! # None, //! # ).expect("Unable to perform server evaluate"); -//! use voprf::Metadata; //! let client_finalize_result = client_blind_result.state.finalize( //! server_evaluate_result.message, -//! &Metadata::none(), +//! None, //! ).expect("Unable to perform client finalization"); //! //! println!("VOPRF output: {:?}", client_finalize_result.to_vec()); @@ -200,7 +198,7 @@ //! //! let mut client_rng = OsRng; //! let client_blind_result = VerifiableClient::::blind( -//! b"input", +//! b"input".to_vec(), //! &mut client_rng, //! ).expect("Unable to construct client"); //! ``` @@ -221,18 +219,17 @@ //! # //! # let mut client_rng = OsRng; //! # let client_blind_result = VerifiableClient::::blind( -//! # b"input", +//! # b"input".to_vec(), //! # &mut client_rng, //! # ).expect("Unable to construct client"); //! # use voprf::VerifiableServer; //! # let mut server_rng = OsRng; //! # let server = VerifiableServer::::new(&mut server_rng) //! # .expect("Unable to construct server"); -//! use voprf::Metadata; //! let server_evaluate_result = server.evaluate( //! &mut server_rng, //! client_blind_result.message, -//! &Metadata::none(), +//! None, //! ).expect("Unable to perform server evaluate"); //! ``` //! @@ -251,7 +248,7 @@ //! # //! # let mut client_rng = OsRng; //! # let client_blind_result = VerifiableClient::::blind( -//! # b"input", +//! # b"input".to_vec(), //! # &mut client_rng, //! # ).expect("Unable to construct client"); //! # use voprf::VerifiableServer; @@ -261,14 +258,13 @@ //! # let server_evaluate_result = server.evaluate( //! # &mut server_rng, //! # client_blind_result.message, -//! # &Metadata::none(), +//! # None, //! # ).expect("Unable to perform server evaluate"); -//! use voprf::Metadata; //! let client_finalize_result = client_blind_result.state.finalize( //! server_evaluate_result.message, //! server_evaluate_result.proof, //! server.get_public_key(), -//! &Metadata::none(), +//! None, //! ).expect("Unable to perform client finalization"); //! //! println!("VOPRF output: {:?}", client_finalize_result.to_vec()); @@ -303,7 +299,7 @@ //! let mut client_messages = vec![]; //! for _ in 0..10 { //! let client_blind_result = VerifiableClient::::blind( -//! b"input", +//! b"input".to_vec(), //! &mut client_rng, //! ).expect("Unable to construct client"); //! client_states.push(client_blind_result.state); @@ -327,13 +323,12 @@ //! # let mut client_messages = vec![]; //! # for _ in 0..10 { //! # let client_blind_result = VerifiableClient::::blind( -//! # b"input", +//! # b"input".to_vec(), //! # &mut client_rng, //! # ).expect("Unable to construct client"); //! # client_states.push(client_blind_result.state); //! # client_messages.push(client_blind_result.message); //! # } -//! # use voprf::Metadata; //! # use voprf::VerifiableServer; //! let mut server_rng = OsRng; //! # let server = VerifiableServer::::new(&mut server_rng) @@ -341,15 +336,14 @@ //! let server_batch_evaluate_result = server.batch_evaluate( //! &mut server_rng, //! &client_messages, -//! &Metadata::none(), +//! None, //! ).expect("Unable to perform server batch evaluate"); //! ``` //! //! Then, the client calls [VerifiableClient::batch_finalize] on //! the client states saved from the first step, along with the messages -//! returned by the server (constructing a [BatchFinalizeInput]), along with the -//! server's proof, in order to produce a vector of outputs if the proof -//! verifies correctly. +//! returned by the server, along with the server's proof, in order to produce +//! a vector of outputs if the proof verifies correctly. //! //! ``` //! # type Group = curve25519_dalek::ristretto::RistrettoPoint; @@ -362,32 +356,27 @@ //! # let mut client_messages = vec![]; //! # for _ in 0..10 { //! # let client_blind_result = VerifiableClient::::blind( -//! # b"input", +//! # b"input".to_vec(), //! # &mut client_rng, //! # ).expect("Unable to construct client"); //! # client_states.push(client_blind_result.state); //! # client_messages.push(client_blind_result.message); //! # } -//! # use voprf::Metadata; //! # use voprf::VerifiableServer; -//! use voprf::BatchFinalizeInput; //! let mut server_rng = OsRng; //! # let server = VerifiableServer::::new(&mut server_rng) //! # .expect("Unable to construct server"); //! # let server_batch_evaluate_result = server.batch_evaluate( //! # &mut server_rng, //! # &client_messages, -//! # &Metadata::none(), +//! # None, //! # ).expect("Unable to perform server batch evaluate"); -//! let batch_finalize_input = BatchFinalizeInput::new( -//! client_states, -//! server_batch_evaluate_result.messages, -//! ); //! let client_batch_finalize_result = VerifiableClient::batch_finalize( -//! batch_finalize_input, +//! &client_states, +//! &server_batch_evaluate_result.messages, //! server_batch_evaluate_result.proof, //! server.get_public_key(), -//! &Metadata::none(), +//! None, //! ).expect("Unable to perform client batch finalization"); //! //! println!("VOPRF batch outputs: {:?}", client_batch_finalize_result); @@ -402,8 +391,7 @@ //! This metadata can be constructed with some type of higher-level domain separation //! to avoid cross-protocol attacks or related issues. //! -//! The default metadata simply consists of the empty vector of bytes, but a custom -//! metadata can be specified, for example, by: `Metadata(b"custom metadata")`. +//! A custom metadata can be specified, for example, by: `Some(b"custom metadata")`. //! //! # Features //! @@ -444,8 +432,7 @@ mod tests; pub use rand; pub use crate::voprf::{ - BatchFinalizeInput, BlindedElement, EvaluationElement, Metadata, NonVerifiableClient, - NonVerifiableClientBlindResult, NonVerifiableServer, NonVerifiableServerEvaluateResult, - VerifiableClient, VerifiableClientBlindResult, VerifiableServer, - VerifiableServerEvaluateResult, + BlindedElement, EvaluationElement, NonVerifiableClient, NonVerifiableClientBlindResult, + NonVerifiableServer, NonVerifiableServerEvaluateResult, VerifiableClient, + VerifiableClientBlindResult, VerifiableServer, VerifiableServerEvaluateResult, }; diff --git a/src/tests/voprf_test_vectors.rs b/src/tests/voprf_test_vectors.rs index ffbb8d1..3c7c51a 100644 --- a/src/tests/voprf_test_vectors.rs +++ b/src/tests/voprf_test_vectors.rs @@ -10,8 +10,8 @@ use crate::{ group::Group, tests::{mock_rng::CycleRng, parser::*}, voprf::{ - BatchFinalizeInput, BlindedElement, EvaluationElement, Metadata, NonVerifiableClient, - NonVerifiableServer, Proof, VerifiableClient, VerifiableServer, + BlindedElement, EvaluationElement, NonVerifiableClient, NonVerifiableServer, Proof, + VerifiableClient, VerifiableServer, }, }; use alloc::string::ToString; @@ -174,7 +174,8 @@ fn test_base_blind( for parameters in tvs { for i in 0..parameters.input.len() { let mut rng = CycleRng::new(parameters.blind[i].to_vec()); - let client_result = NonVerifiableClient::::blind(¶meters.input[i], &mut rng)?; + let client_result = + NonVerifiableClient::::blind(parameters.input[i].clone(), &mut rng)?; assert_eq!( ¶meters.blind[i], @@ -197,7 +198,7 @@ fn test_verifiable_blind( for i in 0..parameters.input.len() { let mut rng = CycleRng::new(parameters.blind[i].to_vec()); let client_blind_result = - VerifiableClient::::blind(¶meters.input[i], &mut rng)?; + VerifiableClient::::blind(parameters.input[i].clone(), &mut rng)?; assert_eq!( ¶meters.blind[i], @@ -221,7 +222,7 @@ fn test_base_evaluate( let server = NonVerifiableServer::::new_with_key(¶meters.sksm)?; let server_result = server.evaluate( BlindedElement::deserialize(¶meters.blinded_element[i])?, - &Metadata(parameters.info.clone()), + Some(¶meters.info), )?; assert_eq!( @@ -245,11 +246,8 @@ fn test_verifiable_evaluate( blinded_elements.push(BlindedElement::deserialize(blinded_element_bytes)?); } - let batch_evaluate_result = server.batch_evaluate( - &mut rng, - &blinded_elements, - &Metadata(parameters.info.clone()), - )?; + let batch_evaluate_result = + server.batch_evaluate(&mut rng, &blinded_elements, Some(¶meters.info))?; for i in 0..parameters.evaluation_element.len() { assert_eq!( @@ -278,7 +276,7 @@ fn test_base_finalize( let client_finalize_result = client.finalize( EvaluationElement::deserialize(¶meters.evaluation_element[i])?, - &Metadata(parameters.info.clone()), + Some(¶meters.info), )?; assert_eq!(¶meters.output[i], &client_finalize_result.to_vec()); @@ -305,20 +303,18 @@ fn test_verifiable_finalize( clients.push(client.clone()); } - let batch_finalize_input = BatchFinalizeInput::new( - clients, - parameters - .evaluation_element - .iter() - .map(|x| EvaluationElement::deserialize(x).unwrap()) - .collect(), - ); + let messages: Vec<_> = parameters + .evaluation_element + .iter() + .map(|x| EvaluationElement::deserialize(x).unwrap()) + .collect(); let batch_result = VerifiableClient::batch_finalize( - batch_finalize_input, + &clients, + &messages, Proof::deserialize(¶meters.proof)?, G::from_element_slice(GenericArray::from_slice(¶meters.pksm))?, - &Metadata(parameters.info.clone()), + Some(¶meters.info), )?; assert_eq!( diff --git a/src/voprf.rs b/src/voprf.rs index a53ca88..79e0423 100644 --- a/src/voprf.rs +++ b/src/voprf.rs @@ -12,12 +12,13 @@ use crate::{ group::Group, serialization::{i2osp, serialize}, }; -use alloc::vec; use alloc::vec::Vec; +use core::convert::TryInto; use core::marker::PhantomData; use digest::{BlockInput, Digest}; +use generic_array::sequence::Concat; use generic_array::{ - typenum::{Unsigned, U1, U2}, + typenum::{U1, U11, U2}, GenericArray, }; use rand::{CryptoRng, RngCore}; @@ -27,14 +28,14 @@ use rand::{CryptoRng, RngCore}; // ========= // /////////////// -static STR_HASH_TO_SCALAR: &[u8] = b"HashToScalar-"; -static STR_HASH_TO_GROUP: &[u8] = b"HashToGroup-"; -static STR_FINALIZE: &[u8] = b"Finalize-"; -static STR_SEED: &[u8] = b"Seed-"; +static STR_HASH_TO_SCALAR: &[u8; 13] = b"HashToScalar-"; +static STR_HASH_TO_GROUP: &[u8; 12] = b"HashToGroup-"; +static STR_FINALIZE: &[u8; 9] = b"Finalize-"; +static STR_SEED: &[u8; 5] = b"Seed-"; static STR_CONTEXT: &[u8] = b"Context-"; -static STR_COMPOSITE: &[u8] = b"Composite-"; -static STR_CHALLENGE: &[u8] = b"Challenge-"; -static STR_VOPRF: &[u8] = b"VOPRF07-"; +static STR_COMPOSITE: &[u8; 10] = b"Composite-"; +static STR_CHALLENGE: &[u8; 10] = b"Challenge-"; +static STR_VOPRF: &[u8; 8] = b"VOPRF07-"; /// Determines the mode of operation (either base mode or /// verifiable mode) @@ -71,7 +72,7 @@ impl_traits_for! { pub(crate) blind: ::Scalar, #[bind] pub(crate) blinded_element: G, - pub(crate) data: alloc::vec::Vec, + pub(crate) data: Vec, #[pd] pub(crate) hash: PhantomData, } @@ -146,13 +147,13 @@ impl_traits_for! { impl NonVerifiableClient { /// Computes the first step for the multiplicative blinding version of DH-OPRF. pub fn blind( - input: &[u8], + input: Vec, blinding_factor_rng: &mut R, ) -> Result, InternalError> { - let (blind, blinded_element) = blind::(input, blinding_factor_rng, Mode::Base)?; + let (blind, blinded_element) = blind::(&input, blinding_factor_rng, Mode::Base)?; Ok(NonVerifiableClientBlindResult { state: Self { - data: input.to_vec(), + data: input, blind, hash: PhantomData, }, @@ -168,13 +169,13 @@ impl NonVerifiableClient { pub fn finalize( &self, evaluation_element: EvaluationElement, - metadata: &Metadata, + metadata: Option<&[u8]>, ) -> Result::OutputSize>, InternalError> { let unblinded_element = evaluation_element.value * &::scalar_invert(&self.blind); - let outputs = finalize_after_unblind::( - &[(self.data.clone(), unblinded_element)], - &metadata.0, + let outputs = finalize_after_unblind::( + Some((self.data.as_slice(), unblinded_element)).into_iter(), + metadata.unwrap_or_default(), Mode::Base, )?; Ok(outputs[0].clone()) @@ -209,14 +210,14 @@ impl NonVerifiableClient { impl VerifiableClient { /// Computes the first step for the multiplicative blinding version of DH-OPRF. pub fn blind( - input: &[u8], + input: Vec, blinding_factor_rng: &mut R, ) -> Result, InternalError> { let (blind, blinded_element) = - blind::(input, blinding_factor_rng, Mode::Verifiable)?; + blind::(&input, blinding_factor_rng, Mode::Verifiable)?; Ok(VerifiableClientBlindResult { state: Self { - data: input.to_vec(), + data: input, blind, blinded_element, hash: PhantomData, @@ -235,50 +236,77 @@ impl VerifiableClient { evaluation_element: EvaluationElement, proof: Proof, pk: G, - metadata: &Metadata, + metadata: Option<&[u8]>, ) -> Result::OutputSize>, InternalError> { - let batch_finalize_input = - BatchFinalizeInput::new(vec![self.clone()], vec![evaluation_element]); - let batch_result = Self::batch_finalize(batch_finalize_input, proof, pk, metadata)?; + // circumvent `.clone()` + let clients: &[Self; 1] = core::slice::from_ref(self).try_into().unwrap(); + let batch_result = + Self::batch_finalize(clients, &[evaluation_element], proof, pk, metadata)?; Ok(batch_result[0].clone()) } /// Allows for batching of the finalization of multiple [VerifiableClient] and [EvaluationElement] pairs - #[allow(clippy::type_complexity)] - pub fn batch_finalize( - batch_finalize_input: BatchFinalizeInput, + pub fn batch_finalize<'a, IC, IM>( + clients: &'a IC, + messages: &'a IM, proof: Proof, pk: G, - metadata: &Metadata, - ) -> Result::OutputSize>>, InternalError> { - let batch_items: Vec> = batch_finalize_input - .clients - .iter() - .zip(batch_finalize_input.messages.iter()) - .map(|(client, evaluation_element)| BatchItems { - blind: client.blind, - evaluation_element: evaluation_element.clone(), - blinded_element: BlindedElement { - value: client.blinded_element, - hash: PhantomData, - }, - }) - .collect(); + metadata: Option<&[u8]>, + ) -> Result::OutputSize>>, InternalError> + where + G: 'a, + H: 'a, + &'a IC: 'a + IntoIterator>, + <&'a IC as IntoIterator>::IntoIter: ExactSizeIterator, + &'a IM: 'a + IntoIterator>, + <&'a IM as IntoIterator>::IntoIter: ExactSizeIterator, + { + struct Items { + clients: IC, + messages: IM, + } - let unblinded_elements = verifiable_unblind(&batch_items, pk, proof, &metadata.0)?; + impl<'a, G: 'a + Group, H: 'a + BlockInput + Digest, IC: Copy, IM: Copy> IntoIterator + for &Items + where + IC: IntoIterator>, + ::IntoIter: ExactSizeIterator, + IM: IntoIterator>, + ::IntoIter: ExactSizeIterator, + { + type Item = BatchItems; - let inputs_and_unblinded_elements: Vec<(Vec, G)> = batch_finalize_input - .clients - .iter() + #[allow(clippy::type_complexity)] + type IntoIter = core::iter::Map< + core::iter::Zip<::IntoIter, ::IntoIter>, + fn((&VerifiableClient, &EvaluationElement)) -> BatchItems, + >; + + fn into_iter(self) -> Self::IntoIter { + self.clients.into_iter().zip(self.messages.into_iter()).map( + |(client, evaluation_element)| BatchItems { + blind: client.blind, + evaluation_element: evaluation_element.copy(), + blinded_element: BlindedElement { + value: client.blinded_element, + hash: PhantomData, + }, + }, + ) + } + } + + let batch_items = Items { clients, messages }; + let metadata = metadata.unwrap_or_default(); + + let unblinded_elements = verifiable_unblind(&batch_items, pk, proof, metadata)?; + + let inputs_and_unblinded_elements = clients + .into_iter() .zip(unblinded_elements.iter()) - .map(|(client, &unblinded_element)| (client.data.clone(), unblinded_element)) - .collect(); + .map(|(client, &unblinded_element)| (client.data.as_slice(), unblinded_element)); - finalize_after_unblind::( - &inputs_and_unblinded_elements, - &metadata.0, - Mode::Verifiable, - ) + finalize_after_unblind::(inputs_and_unblinded_elements, metadata, Mode::Verifiable) } #[cfg(test)] @@ -316,7 +344,7 @@ impl VerifiableClient { impl NonVerifiableServer { /// Produces a new instance of a [NonVerifiableServer] using a supplied RNG pub fn new(rng: &mut R) -> Result { - let mut seed = vec![0u8; ::OutputSize::USIZE]; + let mut seed = GenericArray::<_, ::OutputSize>::default(); rng.fill_bytes(&mut seed); Self::new_from_seed(&seed) } @@ -324,7 +352,7 @@ impl NonVerifiableServer { /// Produces a new instance of a [NonVerifiableServer] using a supplied set of bytes to /// represent the server's private key pub fn new_with_key(private_key_bytes: &[u8]) -> Result { - let sk = G::from_scalar_slice(&GenericArray::clone_from_slice(private_key_bytes))?; + let sk = G::from_scalar_slice(private_key_bytes)?; Ok(Self { sk, hash: PhantomData, @@ -336,7 +364,8 @@ impl NonVerifiableServer { /// /// Corresponds to DeriveKeyPair() function from the VOPRF specification. pub fn new_from_seed(seed: &[u8]) -> Result { - let dst = [STR_HASH_TO_SCALAR, &get_context_string::(Mode::Base)?].concat(); + let dst = + GenericArray::from(*STR_HASH_TO_SCALAR).concat(get_context_string::(Mode::Base)?); let sk = G::hash_to_scalar::(seed, &dst)?; Ok(Self { sk, @@ -355,15 +384,16 @@ impl NonVerifiableServer { pub fn evaluate( &self, blinded_element: BlindedElement, - metadata: &Metadata, + metadata: Option<&[u8]>, ) -> Result, InternalError> { let context = [ STR_CONTEXT, &get_context_string::(Mode::Base)?, - &serialize::(&metadata.0)?, + &serialize::(metadata.unwrap_or_default())?, ] .concat(); - let dst = [STR_HASH_TO_SCALAR, &get_context_string::(Mode::Base)?].concat(); + let dst = + GenericArray::from(*STR_HASH_TO_SCALAR).concat(get_context_string::(Mode::Base)?); let m = G::hash_to_scalar::(&context, &dst)?; let t = self.sk + &m; let evaluation_element = blinded_element.value * &G::scalar_invert(&t); @@ -385,7 +415,7 @@ impl NonVerifiableServer { impl VerifiableServer { /// Produces a new instance of a [VerifiableServer] using a supplied RNG pub fn new(rng: &mut R) -> Result { - let mut seed = vec![0u8; ::OutputSize::USIZE]; + let mut seed = GenericArray::<_, ::OutputSize>::default(); rng.fill_bytes(&mut seed); Self::new_from_seed(&seed) } @@ -393,7 +423,7 @@ impl VerifiableServer { /// Produces a new instance of a [VerifiableServer] using a supplied set of bytes to /// represent the server's private key pub fn new_with_key(key: &[u8]) -> Result { - let sk = G::from_scalar_slice(&GenericArray::clone_from_slice(key))?; + let sk = G::from_scalar_slice(key)?; let pk = G::base_point() * &sk; Ok(Self { sk, @@ -407,11 +437,8 @@ impl VerifiableServer { /// /// Corresponds to DeriveKeyPair() function from the VOPRF specification. pub fn new_from_seed(seed: &[u8]) -> Result { - let dst = [ - STR_HASH_TO_SCALAR, - &get_context_string::(Mode::Verifiable)?, - ] - .concat(); + let dst = GenericArray::from(*STR_HASH_TO_SCALAR) + .concat(get_context_string::(Mode::Verifiable)?); let sk = G::hash_to_scalar::(seed, &dst)?; let pk = G::base_point() * &sk; Ok(Self { @@ -433,37 +460,40 @@ impl VerifiableServer { &self, rng: &mut R, blinded_element: BlindedElement, - metadata: &Metadata, + metadata: Option<&[u8]>, ) -> Result, InternalError> { let batch_result = self.batch_evaluate(rng, &[blinded_element], metadata)?; Ok(VerifiableServerEvaluateResult { - message: batch_result.messages[0].clone(), + message: batch_result.messages[0].copy(), proof: batch_result.proof, }) } /// Allows for batching of the evaluation of multiple [BlindedElement] messages from a [VerifiableClient] - pub fn batch_evaluate( + pub fn batch_evaluate<'a, R: RngCore + CryptoRng, I>( &self, rng: &mut R, - blinded_elements: &[BlindedElement], - metadata: &Metadata, - ) -> Result, InternalError> { + blinded_elements: &'a I, + metadata: Option<&[u8]>, + ) -> Result, InternalError> + where + G: 'a, + H: 'a, + &'a I: IntoIterator>, + <&'a I as IntoIterator>::IntoIter: ExactSizeIterator, + { let context = [ STR_CONTEXT, &get_context_string::(Mode::Verifiable)?, - &serialize::(&metadata.0)?, - ] - .concat(); - let dst = [ - STR_HASH_TO_SCALAR, - &get_context_string::(Mode::Verifiable)?, + &serialize::(metadata.unwrap_or_default())?, ] .concat(); + let dst = GenericArray::from(*STR_HASH_TO_SCALAR) + .concat(get_context_string::(Mode::Verifiable)?); let m = G::hash_to_scalar::(&context, &dst)?; let t = self.sk + &m; let evaluation_elements: Vec> = blinded_elements - .iter() + .into_iter() .map(|x| EvaluationElement { value: x.value * &G::scalar_invert(&t), hash: PhantomData, @@ -473,7 +503,14 @@ impl VerifiableServer { let g = G::base_point(); let u = g * &t; - let proof = generate_proof(rng, t, g, u, &evaluation_elements, blinded_elements)?; + let proof = generate_proof( + rng, + t, + g, + u, + evaluation_elements.iter().map(EvaluationElement::copy), + blinded_elements.into_iter().map(BlindedElement::copy), + )?; Ok(VerifiableServerBatchEvaluateResult { messages: evaluation_elements, @@ -496,23 +533,6 @@ impl VerifiableServer { } } -///////////////////////// -// Optional Parameters // -//==================== // -///////////////////////// - -/// Allows for implementations to specify an optional sequence of -/// public bytes that must be agreed-upon by the client and server -#[derive(Default)] -pub struct Metadata(pub Vec); - -impl Metadata { - /// Specifies no metadata (the default option) - pub fn none() -> Self { - Self::default() - } -} - ///////////////////////// // Convenience Structs // //==================== // @@ -556,23 +576,6 @@ pub struct VerifiableServerBatchEvaluateResult pub proof: Proof, } -/// An input to the verifiable client batch finalize function, constructed -/// by aggregating clients and server messages -pub struct BatchFinalizeInput { - clients: Vec>, - messages: Vec>, -} - -impl BatchFinalizeInput { - /// Create a new instance from a vector of clients and a vector of messages - pub fn new( - clients: Vec>, - messages: Vec>, - ) -> Self { - Self { clients, messages } - } -} - /////////////////////////////////////////////// // Inner functions and Trait Implementations // // ========================================= // @@ -585,9 +588,15 @@ struct BatchItems { blinded_element: BlindedElement, } -/// Convenience test functions for [BlindedElement], [EvaluationElement], and [Proof] - impl BlindedElement { + /// Only used to easier validate allocation + fn copy(&self) -> Self { + Self { + value: self.value, + hash: PhantomData, + } + } + #[cfg(test)] /// Only used for testing zeroize pub fn as_ptrs(&self) -> Vec> { @@ -596,6 +605,14 @@ impl BlindedElement { } impl EvaluationElement { + /// Only used to easier validate allocation + fn copy(&self) -> Self { + Self { + value: self.value, + hash: PhantomData, + } + } + #[cfg(test)] /// Only used for testing zeroize pub fn as_ptrs(&self) -> Vec> { @@ -622,18 +639,22 @@ fn blind( ) -> Result<(::Scalar, G), InternalError> { // Choose a random scalar that must be non-zero let blind = ::random_nonzero_scalar(blinding_factor_rng); - let dst = [STR_HASH_TO_GROUP, &get_context_string::(mode)?].concat(); + let dst = GenericArray::from(*STR_HASH_TO_GROUP).concat(get_context_string::(mode)?); let hashed_point = ::hash_to_curve::(input, &dst)?; let blinded_element = hashed_point * &blind; Ok((blind, blinded_element)) } -fn verifiable_unblind( - batch_items: &[BatchItems], +fn verifiable_unblind<'a, G: 'a + Group, H: 'a + BlockInput + Digest, I>( + batch_items: &'a I, pk: G, proof: Proof, info: &[u8], -) -> Result, InternalError> { +) -> Result, InternalError> +where + &'a I: IntoIterator>, + <&'a I as IntoIterator>::IntoIter: ExactSizeIterator, +{ let context = [ STR_CONTEXT, &get_context_string::(Mode::Verifiable)?, @@ -641,33 +662,23 @@ fn verifiable_unblind( ] .concat(); - let dst = [ - STR_HASH_TO_SCALAR, - &get_context_string::(Mode::Verifiable)?, - ] - .concat(); + let dst = + GenericArray::from(*STR_HASH_TO_SCALAR).concat(get_context_string::(Mode::Verifiable)?); let m = G::hash_to_scalar::(&context, &dst)?; let g = G::base_point(); let t = g * &m; let u = t + &pk; - let blinds: Vec<::Scalar> = batch_items.iter().map(|x| x.blind).collect(); - let evaluation_elements: Vec> = batch_items - .iter() - .map(|x| x.evaluation_element.clone()) - .collect(); - let blinded_elements: Vec> = batch_items - .iter() - .map(|x| x.blinded_element.clone()) - .collect(); + let blinds = batch_items.into_iter().map(|x| x.blind); + let evaluation_elements = batch_items.into_iter().map(|x| x.evaluation_element); + let blinded_elements = batch_items.into_iter().map(|x| x.blinded_element); - verify_proof(g, u, &evaluation_elements, &blinded_elements, proof)?; + verify_proof(g, u, evaluation_elements, blinded_elements, proof)?; let unblinded_elements = blinds - .iter() - .zip(evaluation_elements.iter()) - .map(|(&blind, x)| x.value * &G::scalar_invert(&blind)) + .zip(batch_items.into_iter().map(|x| x.evaluation_element)) + .map(|(blind, x)| x.value * &G::scalar_invert(&blind)) .collect(); Ok(unblinded_elements) } @@ -678,31 +689,29 @@ fn generate_proof( k: ::Scalar, a: G, b: G, - cs: &[EvaluationElement], - ds: &[BlindedElement], + cs: impl Iterator> + ExactSizeIterator, + ds: impl Iterator> + ExactSizeIterator, ) -> Result, InternalError> { - let (m, z) = compute_composites::(Some(k), b, cs, ds)?; + let (m, z) = compute_composites(Some(k), b, cs, ds)?; let r = G::random_nonzero_scalar(rng); let t2 = a * &r; let t3 = m * &r; - let challenge_dst = [STR_CHALLENGE, &get_context_string::(Mode::Verifiable)?].concat(); + let challenge_dst = + GenericArray::from(*STR_CHALLENGE).concat(get_context_string::(Mode::Verifiable)?); let h2_input = [ - serialize::(&b.to_arr().to_vec())?, - serialize::(&m.to_arr().to_vec())?, - serialize::(&z.to_arr().to_vec())?, - serialize::(&t2.to_arr().to_vec())?, - serialize::(&t3.to_arr().to_vec())?, + serialize::(&b.to_arr())?, + serialize::(&m.to_arr())?, + serialize::(&z.to_arr())?, + serialize::(&t2.to_arr())?, + serialize::(&t3.to_arr())?, serialize::(&challenge_dst)?, ] .concat(); - let hash_to_scalar_dst = [ - STR_HASH_TO_SCALAR, - &get_context_string::(Mode::Verifiable)?, - ] - .concat(); + let hash_to_scalar_dst = + GenericArray::from(*STR_HASH_TO_SCALAR).concat(get_context_string::(Mode::Verifiable)?); let c_scalar = G::hash_to_scalar::(&h2_input, &hash_to_scalar_dst)?; let s_scalar = r - &(c_scalar * &k); @@ -718,30 +727,28 @@ fn generate_proof( fn verify_proof( a: G, b: G, - cs: &[EvaluationElement], - ds: &[BlindedElement], + cs: impl Iterator> + ExactSizeIterator, + ds: impl Iterator> + ExactSizeIterator, proof: Proof, ) -> Result<(), InternalError> { - let (m, z) = compute_composites::(None, b, cs, ds)?; + let (m, z) = compute_composites(None, b, cs, ds)?; let t2 = (a * &proof.s_scalar) + &(b * &proof.c_scalar); let t3 = (m * &proof.s_scalar) + &(z * &proof.c_scalar); - let challenge_dst = [STR_CHALLENGE, &get_context_string::(Mode::Verifiable)?].concat(); + let challenge_dst = + GenericArray::from(*STR_CHALLENGE).concat(get_context_string::(Mode::Verifiable)?); let h2_input = [ - serialize::(&b.to_arr().to_vec())?, - serialize::(&m.to_arr().to_vec())?, - serialize::(&z.to_arr().to_vec())?, - serialize::(&t2.to_arr().to_vec())?, - serialize::(&t3.to_arr().to_vec())?, + serialize::(&b.to_arr())?, + serialize::(&m.to_arr())?, + serialize::(&z.to_arr())?, + serialize::(&t2.to_arr())?, + serialize::(&t3.to_arr())?, serialize::(&challenge_dst)?, ] .concat(); - let hash_to_scalar_dst = [ - STR_HASH_TO_SCALAR, - &get_context_string::(Mode::Verifiable)?, - ] - .concat(); + let hash_to_scalar_dst = + GenericArray::from(*STR_HASH_TO_SCALAR).concat(get_context_string::(Mode::Verifiable)?); let c = G::hash_to_scalar::(&h2_input, &hash_to_scalar_dst)?; match G::ct_equal_scalar(&c, &proof.c_scalar) { @@ -750,73 +757,69 @@ fn verify_proof( } } -#[allow(clippy::type_complexity)] -fn finalize_after_unblind( - inputs_and_unblinded_elements: &[(Vec, G)], +fn finalize_after_unblind< + 'a, + G: Group, + H: BlockInput + Digest, + I: Iterator, +>( + inputs_and_unblinded_elements: I, info: &[u8], mode: Mode, ) -> Result::OutputSize>>, InternalError> { - let finalize_dst = [STR_FINALIZE, &get_context_string::(mode)?].concat(); + let finalize_dst = GenericArray::from(*STR_FINALIZE).concat(get_context_string::(mode)?); - let mut outputs = vec![]; - - for (input, unblinded_element) in inputs_and_unblinded_elements { - outputs.push(::digest( - &[ - serialize::(input)?, - serialize::(info)?, - serialize::(&unblinded_element.to_arr().to_vec())?, - serialize::(&finalize_dst)?, - ] - .concat(), - )); - } - - Ok(outputs) + inputs_and_unblinded_elements + .map(|(input, unblinded_element)| { + Ok(::digest( + &[ + serialize::(input)?, + serialize::(info)?, + serialize::(&unblinded_element.to_arr())?, + serialize::(&finalize_dst)?, + ] + .concat(), + )) + }) + .collect() } fn compute_composites( k_option: Option<::Scalar>, b: G, - c_slice: &[EvaluationElement], - d_slice: &[BlindedElement], + c_slice: impl Iterator> + ExactSizeIterator, + d_slice: impl Iterator> + ExactSizeIterator, ) -> Result<(G, G), InternalError> { if c_slice.len() != d_slice.len() { return Err(InternalError::MismatchedLengthsForCompositeInputs); } - let seed_dst = [STR_SEED, &get_context_string::(Mode::Verifiable)?].concat(); - let composite_dst = [STR_COMPOSITE, &get_context_string::(Mode::Verifiable)?].concat(); + let seed_dst = GenericArray::from(*STR_SEED).concat(get_context_string::(Mode::Verifiable)?); + let composite_dst = + GenericArray::from(*STR_COMPOSITE).concat(get_context_string::(Mode::Verifiable)?); - let h1_input = [ - serialize::(&b.to_arr().to_vec())?, - serialize::(&seed_dst)?, - ] - .concat(); + let h1_input = [serialize::(&b.to_arr())?, serialize::(&seed_dst)?].concat(); let seed = ::digest(&h1_input); let mut m = G::identity(); let mut z = G::identity(); - for i in 0..c_slice.len() { + for (i, (c, d)) in c_slice.zip(d_slice).enumerate() { let h2_input = [ - serialize::(&seed)?, - i2osp::(i)?.to_vec(), - serialize::(&c_slice[i].value.to_arr().to_vec())?, - serialize::(&d_slice[i].value.to_arr().to_vec())?, - serialize::(&composite_dst)?, - ] - .concat(); - let dst = [ - STR_HASH_TO_SCALAR, - &get_context_string::(Mode::Verifiable)?, + serialize::(&seed)?.as_slice(), + &i2osp::(i)?, + &serialize::(&c.value.to_arr())?, + &serialize::(&d.value.to_arr())?, + &serialize::(&composite_dst)?, ] .concat(); + let dst = GenericArray::from(*STR_HASH_TO_SCALAR) + .concat(get_context_string::(Mode::Verifiable)?); let di = G::hash_to_scalar::(&h2_input, &dst)?; - m = c_slice[i].value * &di + &m; + m = c.value * &di + &m; z = match k_option { Some(_) => z, - None => d_slice[i].value * &di + &z, + None => d.value * &di + &z, }; } @@ -830,13 +833,10 @@ fn compute_composites( /// Generates the contextString parameter as defined in /// -fn get_context_string(mode: Mode) -> Result, InternalError> { - Ok([ - STR_VOPRF, - &i2osp::(mode as usize)?, - &i2osp::(G::SUITE_ID)?, - ] - .concat()) +fn get_context_string(mode: Mode) -> Result, InternalError> { + Ok(GenericArray::from(*STR_VOPRF) + .concat(i2osp::(mode as usize)?) + .concat(i2osp::(G::SUITE_ID)?)) } /////////// @@ -858,7 +858,8 @@ mod tests { info: &[u8], mode: Mode, ) -> GenericArray::OutputSize> { - let dst = [STR_HASH_TO_GROUP, &get_context_string::(mode).unwrap()].concat(); + let dst = + GenericArray::from(*STR_HASH_TO_GROUP).concat(get_context_string::(mode).unwrap()); let point = G::hash_to_curve::(input, &dst).unwrap(); let context = [ @@ -867,28 +868,31 @@ mod tests { &serialize::(info).unwrap(), ] .concat(); - let dst = [STR_HASH_TO_SCALAR, &get_context_string::(mode).unwrap()].concat(); + let dst = + GenericArray::from(*STR_HASH_TO_SCALAR).concat(get_context_string::(mode).unwrap()); let m = ::hash_to_scalar::(&context, &dst).unwrap(); let res = point * &::scalar_invert(&(key + &m)); - finalize_after_unblind::(&[(input.to_vec(), res)], info, mode).unwrap()[0].clone() + finalize_after_unblind::(Some((input, res)).into_iter(), info, mode).unwrap()[0] + .clone() } fn base_retrieval() { let input = b"input"; let info = b"info"; let mut rng = OsRng; - let client_blind_result = NonVerifiableClient::::blind(&input[..], &mut rng).unwrap(); + let client_blind_result = + NonVerifiableClient::::blind(input.to_vec(), &mut rng).unwrap(); let server = NonVerifiableServer::::new(&mut rng).unwrap(); let server_result = server - .evaluate(client_blind_result.message, &Metadata(info.to_vec())) + .evaluate(client_blind_result.message, Some(info)) .unwrap(); let client_finalize_result = client_blind_result .state - .finalize(server_result.message, &Metadata(info.to_vec())) + .finalize(server_result.message, Some(info)) .unwrap(); - let res2 = prf::(&input[..], server.get_private_key(), info, Mode::Base); + let res2 = prf::(input, server.get_private_key(), info, Mode::Base); assert_eq!(client_finalize_result, res2); } @@ -896,14 +900,11 @@ mod tests { let input = b"input"; let info = b"info"; let mut rng = OsRng; - let client_blind_result = VerifiableClient::::blind(&input[..], &mut rng).unwrap(); + let client_blind_result = + VerifiableClient::::blind(input.to_vec(), &mut rng).unwrap(); let server = VerifiableServer::::new(&mut rng).unwrap(); let server_result = server - .evaluate( - &mut rng, - client_blind_result.message, - &Metadata(info.to_vec()), - ) + .evaluate(&mut rng, client_blind_result.message, Some(info)) .unwrap(); let client_finalize_result = client_blind_result .state @@ -911,10 +912,10 @@ mod tests { server_result.message, server_result.proof, server.get_public_key(), - &Metadata(info.to_vec()), + Some(info), ) .unwrap(); - let res2 = prf::(&input[..], server.get_private_key(), info, Mode::Verifiable); + let res2 = prf::(input, server.get_private_key(), info, Mode::Verifiable); assert_eq!(client_finalize_result, res2); } @@ -922,14 +923,11 @@ mod tests { let input = b"input"; let info = b"info"; let mut rng = OsRng; - let client_blind_result = VerifiableClient::::blind(&input[..], &mut rng).unwrap(); + let client_blind_result = + VerifiableClient::::blind(input.to_vec(), &mut rng).unwrap(); let server = VerifiableServer::::new(&mut rng).unwrap(); let server_result = server - .evaluate( - &mut rng, - client_blind_result.message, - &Metadata(info.to_vec()), - ) + .evaluate(&mut rng, client_blind_result.message, Some(info)) .unwrap(); let wrong_pk = { // Choose a group element that is unlikely to be the right public key @@ -939,7 +937,7 @@ mod tests { server_result.message, server_result.proof, wrong_pk, - &Metadata(info.to_vec()), + Some(info), ); assert!(client_finalize_result.is_err()); } @@ -955,26 +953,26 @@ mod tests { let mut input = vec![0u8; 32]; rng.fill_bytes(&mut input); let client_blind_result = - VerifiableClient::::blind(&input[..], &mut rng).unwrap(); + VerifiableClient::::blind(input.clone(), &mut rng).unwrap(); inputs.push(input); client_states.push(client_blind_result.state); client_messages.push(client_blind_result.message); } let server = VerifiableServer::::new(&mut rng).unwrap(); let server_result = server - .batch_evaluate(&mut rng, &client_messages, &Metadata(info.to_vec())) + .batch_evaluate(&mut rng, &client_messages, Some(info)) .unwrap(); - let batch_finalize_input = BatchFinalizeInput::new(client_states, server_result.messages); let client_finalize_result = VerifiableClient::batch_finalize( - batch_finalize_input, + &client_states, + &server_result.messages, server_result.proof, server.get_public_key(), - &Metadata(info.to_vec()), + Some(info), ) .unwrap(); let mut res2 = vec![]; for input in inputs.iter().take(num_iterations) { - let output = prf::(&input[..], server.get_private_key(), info, Mode::Verifiable); + let output = prf::(input, server.get_private_key(), info, Mode::Verifiable); res2.push(output); } assert_eq!(client_finalize_result, res2); @@ -991,25 +989,25 @@ mod tests { let mut input = vec![0u8; 32]; rng.fill_bytes(&mut input); let client_blind_result = - VerifiableClient::::blind(&input[..], &mut rng).unwrap(); + VerifiableClient::::blind(input.clone(), &mut rng).unwrap(); inputs.push(input); client_states.push(client_blind_result.state); client_messages.push(client_blind_result.message); } let server = VerifiableServer::::new(&mut rng).unwrap(); let server_result = server - .batch_evaluate(&mut rng, &client_messages, &Metadata(info.to_vec())) + .batch_evaluate(&mut rng, &client_messages, Some(info)) .unwrap(); - let batch_finalize_input = BatchFinalizeInput::new(client_states, server_result.messages); let wrong_pk = { // Choose a group element that is unlikely to be the right public key G::hash_to_curve::(b"msg", b"dst").unwrap() }; let client_finalize_result = VerifiableClient::batch_finalize( - batch_finalize_input, + &client_states, + &server_result.messages, server_result.proof, wrong_pk, - &Metadata(info.to_vec()), + Some(info), ); assert!(client_finalize_result.is_err()); } @@ -1019,7 +1017,8 @@ mod tests { let mut input = alloc::vec![0u8; 64]; rng.fill_bytes(&mut input); let info = b"info"; - let client_blind_result = NonVerifiableClient::::blind(&input, &mut rng).unwrap(); + let client_blind_result = + NonVerifiableClient::::blind(input.clone(), &mut rng).unwrap(); let client_finalize_result = client_blind_result .state .finalize( @@ -1027,18 +1026,19 @@ mod tests { value: client_blind_result.message.value, hash: PhantomData, }, - &Metadata(info.to_vec()), + Some(info), ) .unwrap(); - let dst = [ - STR_HASH_TO_GROUP, - &get_context_string::(Mode::Base).unwrap(), - ] - .concat(); + let dst = GenericArray::from(*STR_HASH_TO_GROUP) + .concat(get_context_string::(Mode::Base).unwrap()); let point = G::hash_to_curve::(&input, &dst).unwrap(); - let res2 = finalize_after_unblind::(&[(input.to_vec(), point)], info, Mode::Base) - .unwrap()[0] + let res2 = finalize_after_unblind::( + Some((input.as_slice(), point)).into_iter(), + info, + Mode::Base, + ) + .unwrap()[0] .clone(); assert_eq!(client_finalize_result, res2); @@ -1047,7 +1047,8 @@ mod tests { fn zeroize_base_client() { let input = b"input"; let mut rng = OsRng; - let client_blind_result = NonVerifiableClient::::blind(&input[..], &mut rng).unwrap(); + let client_blind_result = + NonVerifiableClient::::blind(input.to_vec(), &mut rng).unwrap(); let mut state = client_blind_result.state; Zeroize::zeroize(&mut state); @@ -1065,7 +1066,8 @@ mod tests { fn zeroize_verifiable_client() { let input = b"input"; let mut rng = OsRng; - let client_blind_result = VerifiableClient::::blind(&input[..], &mut rng).unwrap(); + let client_blind_result = + VerifiableClient::::blind(input.to_vec(), &mut rng).unwrap(); let mut state = client_blind_result.state; Zeroize::zeroize(&mut state); @@ -1084,10 +1086,11 @@ mod tests { let input = b"input"; let info = b"info"; let mut rng = OsRng; - let client_blind_result = NonVerifiableClient::::blind(&input[..], &mut rng).unwrap(); + let client_blind_result = + NonVerifiableClient::::blind(input.to_vec(), &mut rng).unwrap(); let server = NonVerifiableServer::::new(&mut rng).unwrap(); let server_result = server - .evaluate(client_blind_result.message, &Metadata(info.to_vec())) + .evaluate(client_blind_result.message, Some(info)) .unwrap(); let mut state = server; @@ -1107,14 +1110,11 @@ mod tests { let input = b"input"; let info = b"info"; let mut rng = OsRng; - let client_blind_result = VerifiableClient::::blind(&input[..], &mut rng).unwrap(); + let client_blind_result = + VerifiableClient::::blind(input.to_vec(), &mut rng).unwrap(); let server = VerifiableServer::::new(&mut rng).unwrap(); let server_result = server - .evaluate( - &mut rng, - client_blind_result.message, - &Metadata(info.to_vec()), - ) + .evaluate(&mut rng, client_blind_result.message, Some(info)) .unwrap(); let mut state = server;