diff --git a/src/lib.rs b/src/lib.rs index 3230b25..82aa96d 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -331,13 +331,10 @@ //! this case. In the following example, we show how to use the batch API to //! produce a single proof for 10 parallel VOPRF evaluations. //! -//! This requires the crate feature `alloc`. -//! //! First, the client produces 10 blindings, storing their resulting states and //! messages: //! //! ``` -//! # #[cfg(feature = "alloc")] { //! # #[cfg(feature = "ristretto255")] //! # type Group = curve25519_dalek::ristretto::RistrettoPoint; //! # #[cfg(feature = "ristretto255")] @@ -358,16 +355,14 @@ //! client_states.push(client_blind_result.state); //! client_messages.push(client_blind_result.message); //! } -//! # } //! ``` //! -//! Next, the server calls the [VerifiableServer::batch_evaluate] function on a -//! set of client messages, to produce a corresponding set of messages to be -//! returned to the client (returned in the same order), along with a single -//! proof: +//! Next, the server calls the [VerifiableServer::batch_evaluate_prepare] and +//! [VerifiableServer::batch_evaluate_finish] function on a set of client +//! messages, to produce a corresponding set of messages to be returned to the +//! client (returned in the same order), along with a single proof: //! //! ``` -//! # #[cfg(feature = "alloc")] { //! # #[cfg(feature = "ristretto255")] //! # type Group = curve25519_dalek::ristretto::RistrettoPoint; //! # #[cfg(feature = "ristretto255")] @@ -376,7 +371,7 @@ //! # type Group = p256_::ProjectivePoint; //! # #[cfg(all(feature = "p256", not(feature = "ristretto255")))] //! # type Hash = sha2::Sha256; -//! # use voprf::VerifiableClient; +//! # use voprf::{VerifiableServerBatchEvaluatePrepareResult, VerifiableServerBatchEvaluateFinishResult, VerifiableClient}; //! # use rand::{rngs::OsRng, RngCore}; //! # //! # let mut client_rng = OsRng; @@ -394,7 +389,50 @@ //! let mut server_rng = OsRng; //! # let server = VerifiableServer::::new(&mut server_rng) //! # .expect("Unable to construct server"); -//! let server_batch_evaluate_result = server +//! let VerifiableServerBatchEvaluatePrepareResult { +//! prepared_evaluation_elements, +//! t, +//! } = server +//! .batch_evaluate_prepare(client_messages.iter(), None) +//! .expect("Unable to perform server batch evaluate"); +//! let prepared_elements: Vec<_> = prepared_evaluation_elements.collect(); +//! let VerifiableServerBatchEvaluateFinishResult { messages, proof } = VerifiableServer::batch_evaluate_finish(&mut server_rng, client_messages.iter(), &prepared_elements, &t) +//! .expect("Unable to perform server batch evaluate"); +//! let messages: Vec<_> = messages.collect(); +//! ``` +//! +//! If [`alloc`] is available, [VerifiableServer::batch_evaluate] can be called +//! to avoid having to collect output manually: +//! +//! ``` +//! # #[cfg(feature = "alloc")] { +//! # #[cfg(feature = "ristretto255")] +//! # type Group = curve25519_dalek::ristretto::RistrettoPoint; +//! # #[cfg(feature = "ristretto255")] +//! # type Hash = sha2::Sha512; +//! # #[cfg(all(feature = "p256", not(feature = "ristretto255")))] +//! # type Group = p256_::ProjectivePoint; +//! # #[cfg(all(feature = "p256", not(feature = "ristretto255")))] +//! # type Hash = sha2::Sha256; +//! # use voprf::{VerifiableServerBatchEvaluateResult, VerifiableClient}; +//! # use rand::{rngs::OsRng, RngCore}; +//! # +//! # let mut client_rng = OsRng; +//! # let mut client_states = vec![]; +//! # let mut client_messages = vec![]; +//! # for _ in 0..10 { +//! # let client_blind_result = VerifiableClient::::blind( +//! # b"input", +//! # &mut client_rng, +//! # ).expect("Unable to construct client"); +//! # client_states.push(client_blind_result.state); +//! # client_messages.push(client_blind_result.message); +//! # } +//! # use voprf::VerifiableServer; +//! let mut server_rng = OsRng; +//! # let server = VerifiableServer::::new(&mut server_rng) +//! # .expect("Unable to construct server"); +//! let VerifiableServerBatchEvaluateResult { messages, proof } = server //! .batch_evaluate(&mut server_rng, &client_messages, None) //! .expect("Unable to perform server batch evaluate"); //! # } @@ -415,7 +453,7 @@ //! # type Group = p256_::ProjectivePoint; //! # #[cfg(all(feature = "p256", not(feature = "ristretto255")))] //! # type Hash = sha2::Sha256; -//! # use voprf::VerifiableClient; +//! # use voprf::{VerifiableServerBatchEvaluateResult, VerifiableClient}; //! # use rand::{rngs::OsRng, RngCore}; //! # //! # let mut client_rng = OsRng; @@ -430,19 +468,17 @@ //! # client_messages.push(client_blind_result.message); //! # } //! # use voprf::VerifiableServer; -//! let mut server_rng = OsRng; +//! # 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, -//! # None, -//! # ).expect("Unable to perform server batch evaluate"); +//! # let VerifiableServerBatchEvaluateResult { messages, proof } = server +//! # .batch_evaluate(&mut server_rng, &client_messages, None) +//! # .expect("Unable to perform server batch evaluate"); //! let client_batch_finalize_result = VerifiableClient::batch_finalize( //! &[b"input"; 10], //! &client_states, -//! &server_batch_evaluate_result.messages, -//! &server_batch_evaluate_result.proof, +//! &messages, +//! &proof, //! server.get_public_key(), //! None, //! ) @@ -527,7 +563,8 @@ pub use crate::group::Group; pub use crate::voprf::VerifiableServerBatchEvaluateResult; pub use crate::voprf::{ BlindedElement, EvaluationElement, NonVerifiableClient, NonVerifiableClientBlindResult, - NonVerifiableServer, NonVerifiableServerEvaluateResult, Proof, VerifiableClient, - VerifiableClientBatchFinalizeResult, VerifiableClientBlindResult, VerifiableServer, - VerifiableServerEvaluateResult, + NonVerifiableServer, NonVerifiableServerEvaluateResult, PreparedEvaluationElement, + PreparedTscalar, Proof, VerifiableClient, VerifiableClientBatchFinalizeResult, + VerifiableClientBlindResult, VerifiableServer, VerifiableServerBatchEvaluateFinishResult, + VerifiableServerBatchEvaluatePrepareResult, VerifiableServerEvaluateResult, }; diff --git a/src/tests/mod.rs b/src/tests/mod.rs index 148aedc..7665c25 100644 --- a/src/tests/mod.rs +++ b/src/tests/mod.rs @@ -5,7 +5,6 @@ // License, Version 2.0 found in the LICENSE-APACHE file in the root directory // of this source tree. -#[cfg(feature = "alloc")] mod mock_rng; mod parser; mod voprf_test_vectors; diff --git a/src/tests/voprf_test_vectors.rs b/src/tests/voprf_test_vectors.rs index 3663f3a..dd949aa 100644 --- a/src/tests/voprf_test_vectors.rs +++ b/src/tests/voprf_test_vectors.rs @@ -8,18 +8,14 @@ use alloc::string::{String, ToString}; use alloc::vec; use alloc::vec::Vec; +use core::ops::Add; use digest::core_api::BlockSizeUser; use digest::{Digest, FixedOutputReset}; -use generic_array::GenericArray; +use generic_array::typenum::Sum; +use generic_array::{ArrayLength, GenericArray}; use json::JsonValue; -#[cfg(feature = "alloc")] -use ::{ - core::ops::Add, - generic_array::{typenum::Sum, ArrayLength}, -}; -#[cfg(feature = "alloc")] use crate::tests::mock_rng::CycleRng; use crate::tests::parser::*; use crate::{ @@ -38,7 +34,6 @@ struct VOPRFTestVectorParameters { blinded_element: Vec>, evaluation_element: Vec>, proof: Vec, - #[cfg(feature = "alloc")] proof_random_scalar: Vec, output: Vec>, } @@ -54,7 +49,6 @@ fn populate_test_vectors(values: &JsonValue) -> VOPRFTestVectorParameters { blinded_element: decode_vec(values, "BlindedElement"), evaluation_element: decode_vec(values, "EvaluationElement"), proof: decode(values, "Proof"), - #[cfg(feature = "alloc")] proof_random_scalar: decode(values, "ProofRandomScalar"), output: decode_vec(values, "Output"), } @@ -118,7 +112,6 @@ fn test_vectors() -> Result<()> { test_verifiable_seed_to_key::(&ristretto_verifiable_tvs)?; test_verifiable_blind::(&ristretto_verifiable_tvs)?; - #[cfg(feature = "alloc")] test_verifiable_evaluate::(&ristretto_verifiable_tvs)?; test_verifiable_finalize::(&ristretto_verifiable_tvs)?; } @@ -253,7 +246,6 @@ fn test_base_evaluate( Ok(()) } -#[cfg(feature = "alloc")] fn test_verifiable_evaluate( tvs: &[VOPRFTestVectorParameters], ) -> Result<()> @@ -261,6 +253,10 @@ where G::ScalarLen: Add, Sum: ArrayLength, { + use crate::{ + VerifiableServerBatchEvaluateFinishResult, VerifiableServerBatchEvaluatePrepareResult, + }; + for parameters in tvs { let mut rng = CycleRng::new(parameters.proof_random_scalar.clone()); let server = VerifiableServer::::new_with_key(¶meters.sksm)?; @@ -270,20 +266,25 @@ where blinded_elements.push(BlindedElement::deserialize(blinded_element_bytes)?); } - let batch_evaluate_result = - server.batch_evaluate(&mut rng, &blinded_elements, Some(¶meters.info))?; + let VerifiableServerBatchEvaluatePrepareResult { + prepared_evaluation_elements, + t, + } = server.batch_evaluate_prepare(blinded_elements.iter(), Some(¶meters.info))?; + let prepared_elements: Vec<_> = prepared_evaluation_elements.collect(); + let VerifiableServerBatchEvaluateFinishResult { messages, proof } = + VerifiableServer::batch_evaluate_finish( + &mut rng, + blinded_elements.iter(), + &prepared_elements, + &t, + )?; + let messages: Vec<_> = messages.collect(); - for i in 0..parameters.evaluation_element.len() { - assert_eq!( - ¶meters.evaluation_element[i], - &batch_evaluate_result.messages[i].serialize().as_slice(), - ); + for (parameter, message) in parameters.evaluation_element.iter().zip(messages) { + assert_eq!(¶meter, &message.serialize().as_slice(),); } - assert_eq!( - ¶meters.proof, - &batch_evaluate_result.proof.serialize().as_slice() - ); + assert_eq!(¶meters.proof, &proof.serialize().as_slice()); } Ok(()) } diff --git a/src/voprf.rs b/src/voprf.rs index 680b1b0..9215083 100644 --- a/src/voprf.rs +++ b/src/voprf.rs @@ -511,23 +511,27 @@ impl VerifiableServer, metadata: Option<&[u8]>, ) -> Result> { - let (mut evaluation_elements, t) = - self.batch_evaluate_1(Some(blinded_element.copy()).into_iter(), metadata)?; - - let evaluation_element = evaluation_elements.next().unwrap(); - - let proof = Self::batch_evaluate_2( - rng, - Some(blinded_element.copy()).into_iter(), - Some(evaluation_element.copy()).into_iter(), + let VerifiableServerBatchEvaluatePrepareResult { + prepared_evaluation_elements: mut evaluation_elements, t, + } = self.batch_evaluate_prepare(Some(blinded_element).into_iter(), metadata)?; + + let prepared_element = [evaluation_elements.next().unwrap()]; + + let VerifiableServerBatchEvaluateFinishResult { + mut messages, + proof, + } = Self::batch_evaluate_finish( + rng, + Some(blinded_element).into_iter(), + &prepared_element, + &t, )?; + let message = messages.next().unwrap(); + //let batch_result = self.batch_evaluate(rng, blinded_elements, metadata)?; - Ok(VerifiableServerEvaluateResult { - message: evaluation_element, - proof, - }) + Ok(VerifiableServerEvaluateResult { message, proof }) } /// Allows for batching of the evaluation of multiple [BlindedElement] @@ -545,37 +549,36 @@ impl VerifiableServer>, <&'a I as IntoIterator>::IntoIter: ExactSizeIterator, { - let (evaluation_elements, t) = self.batch_evaluate_1( - blinded_elements.into_iter().map(BlindedElement::copy), - metadata, - )?; - - let evaluation_elements: Vec<_> = evaluation_elements.collect(); - - let proof = Self::batch_evaluate_2( - rng, - blinded_elements.into_iter().map(BlindedElement::copy), - evaluation_elements.iter().map(EvaluationElement::copy), + let VerifiableServerBatchEvaluatePrepareResult { + prepared_evaluation_elements: evaluation_elements, t, - )?; + } = self.batch_evaluate_prepare(blinded_elements.into_iter(), metadata)?; + + let prepared_elements = evaluation_elements.collect(); + + let VerifiableServerBatchEvaluateFinishResult { messages, proof } = + Self::batch_evaluate_finish::<_, _, Vec<_>>( + rng, + blinded_elements.into_iter(), + &prepared_elements, + &t, + )?; Ok(VerifiableServerBatchEvaluateResult { - messages: evaluation_elements, + messages: messages.collect(), proof, }) } - fn batch_evaluate_1( + /// Alternative version of [`batch_evaluate`](Self::batch_evaluate) without + /// memory allocation. Returned [`PreparedEvaluationElement`] have to be + /// [`collect`](Iterator::collect)ed and passed into + /// [`batch_evaluate_finish`](Self::batch_evaluate_finish). + pub fn batch_evaluate_prepare<'a, I: Iterator>>( &self, blinded_elements: I, metadata: Option<&[u8]>, - ) -> Result<( - impl Iterator> + ExactSizeIterator, - G::Scalar, - )> - where - I: Iterator> + ExactSizeIterator, - { + ) -> Result> { chain!(context, STR_CONTEXT => |x| Some(x.as_ref()), get_context_string::(Mode::Verifiable)? => |x| Some(x.as_slice()), @@ -585,30 +588,62 @@ impl VerifiableServer(Mode::Verifiable)?); let m = G::hash_to_scalar::(context, dst)?; let t = self.sk + &m; - let evaluation_elements = blinded_elements.map(move |x| EvaluationElement { - value: x.value * &G::scalar_invert(&t), - hash: PhantomData, - }); + let evaluation_elements = blinded_elements + // To make a return type possible, we have to convert to a `fn` pointer, which isn't + // possible if we `move` from context. + .zip(iter::repeat(G::scalar_invert(&t))) + .map(, _)) -> _>::from(|(x, t)| { + PreparedEvaluationElement(EvaluationElement { + value: x.value * &t, + hash: PhantomData, + }) + })); - Ok((evaluation_elements, t)) + Ok(VerifiableServerBatchEvaluatePrepareResult { + prepared_evaluation_elements: evaluation_elements, + t: PreparedTscalar { + t, + hash: PhantomData, + }, + }) } - /// Allows for batching of the evaluation of multiple [BlindedElement] - /// messages from a [VerifiableClient] - fn batch_evaluate_2( + /// See [`batch_evaluate_prepare`](Self::batch_evaluate_prepare) for more + /// details. + pub fn batch_evaluate_finish<'a, 'b, R: RngCore + CryptoRng, IB, IE>( rng: &mut R, blinded_elements: IB, - evaluation_elements: IE, - t: G::Scalar, - ) -> Result> + evaluation_elements: &'b IE, + PreparedTscalar { t, .. }: &PreparedTscalar, + ) -> Result> where - IB: Iterator> + ExactSizeIterator, - IE: Iterator> + ExactSizeIterator, + G: 'a + 'b, + H: 'a + 'b, + IB: Iterator> + ExactSizeIterator, + &'b IE: IntoIterator>, + <&'b IE as IntoIterator>::IntoIter: ExactSizeIterator, { let g = G::base_point(); - let u = g * &t; + let u = g * t; - generate_proof(rng, t, g, u, evaluation_elements, blinded_elements) + let proof = generate_proof( + rng, + *t, + g, + u, + evaluation_elements + .into_iter() + .map(|element| element.0.copy()), + blinded_elements.map(BlindedElement::copy), + )?; + let messages = + evaluation_elements + .into_iter() + .map() -> _>::from( + |element| element.0.copy(), + )); + + Ok(VerifiableServerBatchEvaluateFinishResult { messages, proof }) } /// Retrieves the server's public key @@ -662,6 +697,59 @@ pub struct VerifiableServerEvaluateResult, } +/// Contains prepared [`EvaluationElement`]s by a verifiable server batch +/// evaluate preparation. +pub struct PreparedEvaluationElement( + EvaluationElement, +); + +/// Contains the prepared `t` by a verifiable server batch evaluate preparation. +#[derive(DeriveWhere)] +#[derive_where(Zeroize(drop))] +pub struct PreparedTscalar { + t: G::Scalar, + #[derive_where(skip)] + hash: PhantomData, +} + +/// Contains the fields that are returned by a verifiable server batch evaluate +/// preparation. +pub struct VerifiableServerBatchEvaluatePrepareResult< + 'a, + G: 'a + Group, + H: 'a + BlockSizeUser + Digest + FixedOutputReset, + I: Iterator>, +> { + /// Prepared [`EvaluationElement`]s that will become messages. + #[allow(clippy::type_complexity)] + pub prepared_evaluation_elements: Map< + Zip>, + fn((&BlindedElement, G::Scalar)) -> PreparedEvaluationElement, + >, + /// Prepared `t` needed to finish the verifiable server batch evaluation. + pub t: PreparedTscalar, +} + +/// Contains the fields that are returned by a verifiable server batch evaluate +/// finish. +pub struct VerifiableServerBatchEvaluateFinishResult< + 'a, + G: 'a + Group, + H: 'a + BlockSizeUser + Digest + FixedOutputReset, + I, +> where + &'a I: IntoIterator>, +{ + /// The messages to send to the client + #[allow(clippy::type_complexity)] + pub messages: Map< + <&'a I as IntoIterator>::IntoIter, + fn(&PreparedEvaluationElement) -> EvaluationElement, + >, + /// The proof for the client to verify + pub proof: Proof, +} + /// Contains the fields that are returned by a verifiable server batch evaluate #[cfg(feature = "alloc")] pub struct VerifiableServerBatchEvaluateResult< @@ -1008,12 +1096,12 @@ fn get_context_string(mode: Mode) -> Result> { mod tests { use core::ops::Add; + use ::alloc::vec; + use ::alloc::vec::Vec; use generic_array::typenum::Sum; use generic_array::{ArrayLength, GenericArray}; use rand::rngs::OsRng; use zeroize::Zeroize; - #[cfg(feature = "alloc")] - use ::{alloc::vec, alloc::vec::Vec}; use super::*; use crate::Group; @@ -1087,7 +1175,6 @@ mod tests { assert_eq!(client_finalize_result, res2); } - #[cfg(feature = "alloc")] fn verifiable_bad_public_key() { let input = b"input"; let info = b"info"; @@ -1111,7 +1198,6 @@ mod tests { assert!(client_finalize_result.is_err()); } - #[cfg(feature = "alloc")] fn verifiable_batch_retrieval() { let info = b"info"; let mut rng = OsRng; @@ -1128,14 +1214,27 @@ mod tests { client_messages.push(client_blind_result.message); } let server = VerifiableServer::::new(&mut rng).unwrap(); - let server_result = server - .batch_evaluate(&mut rng, &client_messages, Some(info)) + let VerifiableServerBatchEvaluatePrepareResult { + prepared_evaluation_elements, + t, + } = server + .batch_evaluate_prepare(client_messages.iter(), Some(info)) .unwrap(); + let prepared_elements: Vec<_> = prepared_evaluation_elements.collect(); + let VerifiableServerBatchEvaluateFinishResult { messages, proof } = + VerifiableServer::batch_evaluate_finish( + &mut rng, + client_messages.iter(), + &prepared_elements, + &t, + ) + .unwrap(); + let messages: Vec<_> = messages.collect(); let client_finalize_result = VerifiableClient::batch_finalize( &inputs, &client_states, - &server_result.messages, - &server_result.proof, + &messages, + &proof, server.get_public_key(), Some(info), ) @@ -1150,7 +1249,6 @@ mod tests { assert_eq!(client_finalize_result, res2); } - #[cfg(feature = "alloc")] fn verifiable_batch_bad_public_key() { let info = b"info"; let mut rng = OsRng; @@ -1167,9 +1265,22 @@ mod tests { client_messages.push(client_blind_result.message); } let server = VerifiableServer::::new(&mut rng).unwrap(); - let server_result = server - .batch_evaluate(&mut rng, &client_messages, Some(info)) + let VerifiableServerBatchEvaluatePrepareResult { + prepared_evaluation_elements, + t, + } = server + .batch_evaluate_prepare(client_messages.iter(), Some(info)) .unwrap(); + let prepared_elements: Vec<_> = prepared_evaluation_elements.collect(); + let VerifiableServerBatchEvaluateFinishResult { messages, proof } = + VerifiableServer::batch_evaluate_finish( + &mut rng, + client_messages.iter(), + &prepared_elements, + &t, + ) + .unwrap(); + let messages: Vec<_> = messages.collect(); let wrong_pk = { // Choose a group element that is unlikely to be the right public key G::hash_to_curve::(b"msg", (*b"dst").into()).unwrap() @@ -1177,8 +1288,8 @@ mod tests { let client_finalize_result = VerifiableClient::batch_finalize( &inputs, &client_states, - &server_result.messages, - &server_result.proof, + &messages, + &proof, wrong_pk, Some(info), ); @@ -1309,11 +1420,8 @@ mod tests { base_retrieval::(); base_inversion_unsalted::(); verifiable_retrieval::(); - #[cfg(feature = "alloc")] verifiable_batch_retrieval::(); - #[cfg(feature = "alloc")] verifiable_bad_public_key::(); - #[cfg(feature = "alloc")] verifiable_batch_bad_public_key::(); zeroize_base_client::();