Skip to main content

voprf_vx/
voprf.rs

1// SPDX-License-Identifier: MIT OR Apache-2.0
2// Copyright (c) VexaHub and contributors.
3// Copyright (c) Meta Platforms, Inc. and affiliates.
4
5//! Contains the main VOPRF API
6
7#[cfg(feature = "alloc")]
8use alloc::vec::Vec;
9use core::iter::{self, Map, Repeat, Zip};
10
11use derive_where::derive_where;
12use digest::Output;
13use hybrid_array::Array;
14use rand_core::{TryCryptoRng, TryRng};
15
16use crate::common::{
17    BlindedElement, EvaluationElement, FinalizeAfterUnblindResult, Mode, PreparedEvaluationElement,
18    Proof, derive_keypair, deterministic_blind_unchecked, finalize_after_unblind, generate_proof,
19    hash_to_group, server_evaluate_hash_input, verify_proof,
20};
21#[cfg(feature = "serde")]
22use crate::serialization::serde::{Element, Scalar};
23use crate::{CipherSuite, Error, Group, Result};
24
25////////////////////////////
26// High-level API Structs //
27// ====================== //
28////////////////////////////
29
30/// A client which engages with a [VoprfServer] in verifiable mode, meaning
31/// that the OPRF outputs can be checked against a server public key.
32#[derive_where(Clone, ZeroizeOnDrop)]
33#[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; <CS::Group as Group>::Scalar, <CS::Group as Group>::Elem)]
34#[cfg_attr(
35    feature = "serde",
36    derive(serde::Deserialize, serde::Serialize),
37    serde(bound = "")
38)]
39pub struct VoprfClient<CS: CipherSuite> {
40    #[cfg_attr(feature = "serde", serde(with = "Scalar::<CS::Group>"))]
41    pub(crate) blind: <CS::Group as Group>::Scalar,
42    #[cfg_attr(feature = "serde", serde(with = "Element::<CS::Group>"))]
43    pub(crate) blinded_element: <CS::Group as Group>::Elem,
44}
45
46/// A server which engages with a [VoprfClient] in verifiable mode, meaning
47/// that the OPRF outputs can be checked against a server public key.
48#[derive_where(Clone, ZeroizeOnDrop)]
49#[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; <CS::Group as Group>::Scalar, <CS::Group as Group>::Elem)]
50#[cfg_attr(
51    feature = "serde",
52    derive(serde::Deserialize, serde::Serialize),
53    serde(bound = "")
54)]
55pub struct VoprfServer<CS: CipherSuite> {
56    #[cfg_attr(feature = "serde", serde(with = "Scalar::<CS::Group>"))]
57    pub(crate) sk: <CS::Group as Group>::Scalar,
58    #[cfg_attr(feature = "serde", serde(with = "Element::<CS::Group>"))]
59    pub(crate) pk: <CS::Group as Group>::Elem,
60}
61
62/////////////////////////
63// API Implementations //
64// =================== //
65/////////////////////////
66
67impl<CS: CipherSuite> VoprfClient<CS> {
68    /// Computes the first step for the multiplicative blinding version of
69    /// DH-OPRF.
70    ///
71    /// # Errors
72    /// [`Error::Input`] if the `input` is empty or longer then [`u16::MAX`].
73    pub fn blind<R: TryRng + TryCryptoRng>(
74        input: &[u8],
75        blinding_factor_rng: &mut R,
76    ) -> Result<VoprfClientBlindResult<CS>> {
77        let blind = CS::Group::random_scalar(blinding_factor_rng)?;
78        Self::deterministic_blind_unchecked_inner(input, blind)
79    }
80
81    /// Computes the first step for the multiplicative blinding version of
82    /// DH-OPRF, taking a blinding factor scalar as input instead of sampling
83    /// from an RNG.
84    ///
85    /// # Caution
86    ///
87    /// This should be used with caution, since it does not perform any checks
88    /// on the validity of the blinding factor!
89    ///
90    /// # Errors
91    /// [`Error::Input`] if the `input` is empty or longer then [`u16::MAX`].
92    #[cfg(any(feature = "danger", test))]
93    pub fn deterministic_blind_unchecked(
94        input: &[u8],
95        blind: <CS::Group as Group>::Scalar,
96    ) -> Result<VoprfClientBlindResult<CS>> {
97        Self::deterministic_blind_unchecked_inner(input, blind)
98    }
99
100    /// Can only fail with [`Error::Input`].
101    fn deterministic_blind_unchecked_inner(
102        input: &[u8],
103        blind: <CS::Group as Group>::Scalar,
104    ) -> Result<VoprfClientBlindResult<CS>> {
105        let blinded_element = deterministic_blind_unchecked::<CS>(input, &blind, Mode::Voprf)?;
106        Ok(VoprfClientBlindResult {
107            state: Self {
108                blind,
109                blinded_element,
110            },
111            message: BlindedElement(blinded_element),
112        })
113    }
114
115    /// Computes the third step for the multiplicative blinding version of
116    /// DH-OPRF, in which the client unblinds the server's message.
117    ///
118    /// # Errors
119    /// - [`Error::Input`] if the `input` is empty or longer then [`u16::MAX`].
120    /// - [`Error::ProofVerification`] if the `proof` failed to verify.
121    pub fn finalize(
122        &self,
123        input: &[u8],
124        evaluation_element: &EvaluationElement<CS>,
125        proof: &Proof<CS>,
126        pk: <CS::Group as Group>::Elem,
127    ) -> Result<Output<CS::Hash>> {
128        let inputs = core::array::from_ref(&input);
129        let clients = core::array::from_ref(self);
130        let messages = core::array::from_ref(evaluation_element);
131
132        let mut batch_result = Self::batch_finalize(inputs, clients, messages, proof, pk)?;
133        batch_result.next().unwrap()
134    }
135
136    /// Allows for batching of the finalization of multiple [VoprfClient]
137    /// and [EvaluationElement] pairs
138    ///
139    /// # Errors
140    /// - [`Error::Batch`] if the number of `clients` and `messages` don't match
141    ///   or is longer then [`u16::MAX`].
142    /// - [`Error::ProofVerification`] if the `proof` failed to verify.
143    ///
144    /// The resulting messages can each fail individually with [`Error::Input`]
145    /// if the `input` is empty or longer then [`u16::MAX`].
146    pub fn batch_finalize<'a, I, II, IC, IM>(
147        inputs: &'a II,
148        clients: &'a IC,
149        messages: &'a IM,
150        proof: &Proof<CS>,
151        pk: <CS::Group as Group>::Elem,
152    ) -> Result<VoprfClientBatchFinalizeResult<'a, CS, I, II, IC, IM>>
153    where
154        CS: 'a,
155        I: 'a + AsRef<[u8]>,
156        &'a II: 'a + IntoIterator<Item = I>,
157        <&'a II as IntoIterator>::IntoIter: ExactSizeIterator,
158        &'a IC: 'a + IntoIterator<Item = &'a VoprfClient<CS>>,
159        <&'a IC as IntoIterator>::IntoIter: ExactSizeIterator,
160        &'a IM: 'a + IntoIterator<Item = &'a EvaluationElement<CS>>,
161        <&'a IM as IntoIterator>::IntoIter: ExactSizeIterator,
162    {
163        let unblinded_elements = verifiable_unblind(clients, messages, pk, proof)?;
164        let inputs_and_unblinded_elements = inputs.into_iter().zip(unblinded_elements);
165        Ok(finalize_after_unblind::<CS, _, _>(
166            inputs_and_unblinded_elements,
167        ))
168    }
169
170    /// Only used for test functions
171    #[cfg(test)]
172    pub fn from_blind_and_element(
173        blind: <CS::Group as Group>::Scalar,
174        blinded_element: <CS::Group as Group>::Elem,
175    ) -> Self {
176        Self {
177            blind,
178            blinded_element,
179        }
180    }
181
182    /// Only used for test functions
183    #[cfg(test)]
184    pub fn get_blind(&self) -> <CS::Group as Group>::Scalar {
185        self.blind
186    }
187}
188
189impl<CS: CipherSuite> VoprfServer<CS> {
190    /// Produces a new instance of a [VoprfServer] using a supplied RNG
191    ///
192    /// # Errors
193    /// [`Error::Protocol`] if the protocol fails and can't be completed.
194    pub fn new<R: TryRng + TryCryptoRng>(rng: &mut R) -> Result<Self> {
195        let mut seed = Array::<_, <CS::Group as Group>::ScalarLen>::default();
196        rng.try_fill_bytes(&mut seed).map_err(|_| Error::Protocol)?;
197        // This can't fail as the hash output is type constrained.
198        Self::new_from_seed(&seed, &[])
199    }
200
201    /// Produces a new instance of a [VoprfServer] using a supplied set of
202    /// bytes to represent the server's private key
203    ///
204    /// # Errors
205    /// [`Error::Deserialization`] if the private key is not a valid point on
206    /// the group or zero.
207    pub fn new_with_key(key: &[u8]) -> Result<Self> {
208        let sk = CS::Group::deserialize_scalar(key)?;
209        let pk = CS::Group::base_elem() * &sk;
210        Ok(Self { sk, pk })
211    }
212
213    /// Produces a new instance of a [VoprfServer] using a supplied set of
214    /// bytes which are used as a seed to derive the server's private key.
215    ///
216    /// Corresponds to DeriveKeyPair() function from the VOPRF specification.
217    ///
218    /// # Errors
219    /// - [`Error::DeriveKeyPair`] if the `input` and `seed` together are longer
220    ///   then `u16::MAX - 3`.
221    /// - [`Error::Protocol`] if the protocol fails and can't be completed.
222    pub fn new_from_seed(seed: &[u8], info: &[u8]) -> Result<Self> {
223        let (sk, pk) = derive_keypair::<CS>(seed, info, Mode::Voprf)?;
224        Ok(Self { sk, pk })
225    }
226
227    /// Only used for tests
228    #[cfg(test)]
229    pub fn get_private_key(&self) -> <CS::Group as Group>::Scalar {
230        self.sk
231    }
232
233    /// Computes the second step for the multiplicative blinding version of
234    /// DH-OPRF. This message is sent from the server (who holds the OPRF key)
235    /// to the client.
236    pub fn blind_evaluate<R: TryRng + TryCryptoRng>(
237        &self,
238        rng: &mut R,
239        blinded_element: &BlindedElement<CS>,
240    ) -> VoprfServerEvaluateResult<CS> {
241        let mut prepared_evaluation_elements =
242            self.batch_blind_evaluate_prepare(iter::once(blinded_element));
243        let prepared_evaluation_element = [prepared_evaluation_elements.next().unwrap()];
244
245        // This can't fail because we know the size of the inputs.
246        let VoprfServerBatchEvaluateFinishResult {
247            mut messages,
248            proof,
249        } = self
250            .batch_blind_evaluate_finish(
251                rng,
252                iter::once(blinded_element),
253                &prepared_evaluation_element,
254            )
255            .unwrap();
256
257        let message = messages.next().unwrap();
258
259        VoprfServerEvaluateResult { message, proof }
260    }
261
262    /// Allows for batching of the evaluation of multiple [BlindedElement]
263    /// messages from a [VoprfClient]
264    ///
265    /// # Errors
266    /// [`Error::Batch`] if the number of `blinded_elements` and
267    /// `evaluation_elements` don't match or is longer then [`u16::MAX`]
268    #[cfg(feature = "alloc")]
269    pub fn batch_blind_evaluate<'a, R: TryRng + TryCryptoRng, I>(
270        &self,
271        rng: &mut R,
272        blinded_elements: &'a I,
273    ) -> Result<VoprfServerBatchEvaluateResult<CS>>
274    where
275        CS: 'a,
276        &'a I: IntoIterator<Item = &'a BlindedElement<CS>>,
277        <&'a I as IntoIterator>::IntoIter: ExactSizeIterator,
278    {
279        let prepared_evaluation_elements = self
280            .batch_blind_evaluate_prepare(blinded_elements.into_iter())
281            .collect();
282        let VoprfServerBatchEvaluateFinishResult { messages, proof } = self
283            .batch_blind_evaluate_finish::<_, _, Vec<_>>(
284                rng,
285                blinded_elements.into_iter(),
286                &prepared_evaluation_elements,
287            )?;
288        let messages = messages.collect();
289
290        Ok(VoprfServerBatchEvaluateResult { messages, proof })
291    }
292
293    /// Alternative version of `batch_blind_evaluate` without memory allocation.
294    /// Returned [`PreparedEvaluationElement`] have to be
295    /// [`collect`](Iterator::collect)ed and passed into
296    /// [`batch_blind_evaluate_finish`](Self::batch_blind_evaluate_finish).
297    pub fn batch_blind_evaluate_prepare<'a, I: Iterator<Item = &'a BlindedElement<CS>>>(
298        &self,
299        blinded_elements: I,
300    ) -> VoprfServerBatchEvaluatePreparedEvaluationElements<CS, I>
301    where
302        CS: 'a,
303    {
304        blinded_elements
305            .zip(iter::repeat(self.sk))
306            .map(|(blinded_element, sk)| {
307                PreparedEvaluationElement(EvaluationElement(blinded_element.0 * &sk))
308            })
309    }
310
311    /// See [`batch_blind_evaluate_prepare`](Self::batch_blind_evaluate_prepare)
312    /// for more details.
313    ///
314    /// # Errors
315    /// [`Error::Batch`] if the number of `blinded_elements` and
316    /// `evaluation_elements` don't match or is longer then [`u16::MAX`]
317    pub fn batch_blind_evaluate_finish<
318        'a,
319        'b,
320        R: TryRng + TryCryptoRng,
321        IB: Iterator<Item = &'a BlindedElement<CS>> + ExactSizeIterator,
322        IE,
323    >(
324        &self,
325        rng: &mut R,
326        blinded_elements: IB,
327        evaluation_elements: &'b IE,
328    ) -> Result<VoprfServerBatchEvaluateFinishResult<'b, CS, IE>>
329    where
330        CS: 'a + 'b,
331        &'b IE: IntoIterator<Item = &'b PreparedEvaluationElement<CS>>,
332        <&'b IE as IntoIterator>::IntoIter: ExactSizeIterator,
333    {
334        let g = CS::Group::base_elem();
335        let proof = generate_proof(
336            rng,
337            self.sk,
338            g,
339            self.pk,
340            blinded_elements.map(|element| element.0),
341            evaluation_elements.into_iter().map(|element| element.0.0),
342            Mode::Voprf,
343        )?;
344
345        let messages = evaluation_elements.into_iter().map(<fn(
346            &PreparedEvaluationElement<CS>,
347        ) -> EvaluationElement<CS>>::from(
348            |element| EvaluationElement(element.0.0),
349        ));
350
351        Ok(VoprfServerBatchEvaluateFinishResult { messages, proof })
352    }
353
354    /// Computes the output of the POPRF on the server side
355    ///
356    /// # Errors
357    /// [`Error::Input`]  if the `input` is longer then [`u16::MAX`].
358    pub fn evaluate(&self, input: &[u8]) -> Result<Output<<CS as CipherSuite>::Hash>> {
359        let input_element = hash_to_group::<CS>(input, Mode::Voprf)?;
360        if CS::Group::is_identity_elem(input_element).into() {
361            return Err(Error::Input);
362        };
363        let evaluated_element = input_element * &self.sk;
364
365        let issued_element = CS::Group::serialize_elem(evaluated_element);
366
367        server_evaluate_hash_input::<CS>(input, None, issued_element)
368    }
369
370    /// Retrieves the server's public key
371    pub fn get_public_key(&self) -> <CS::Group as Group>::Elem {
372        self.pk
373    }
374}
375
376/////////////////////////
377// Convenience Structs //
378//==================== //
379/////////////////////////
380
381/// Contains the fields that are returned by a verifiable client blind
382#[derive_where(Debug; <CS::Group as Group>::Scalar, <CS::Group as Group>::Elem)]
383pub struct VoprfClientBlindResult<CS: CipherSuite> {
384    /// The state to be persisted on the client
385    pub state: VoprfClient<CS>,
386    /// The message to send to the server
387    pub message: BlindedElement<CS>,
388}
389
390/// Concrete return type for [`VoprfClient::batch_finalize`].
391pub type VoprfClientBatchFinalizeResult<'a, C, I, II, IC, IM> = FinalizeAfterUnblindResult<
392    'a,
393    C,
394    I,
395    Zip<<&'a II as IntoIterator>::IntoIter, VoprfUnblindResult<'a, C, IC, IM>>,
396>;
397
398/// Contains the fields that are returned by a verifiable server evaluate
399#[derive_where(Debug; <CS::Group as Group>::Scalar, <CS::Group as Group>::Elem)]
400pub struct VoprfServerEvaluateResult<CS: CipherSuite> {
401    /// The message to send to the client
402    pub message: EvaluationElement<CS>,
403    /// The proof for the client to verify
404    pub proof: Proof<CS>,
405}
406
407/// Contains the fields that are returned by a verifiable server batch evaluate
408#[derive_where(Debug; <CS::Group as Group>::Scalar, <CS::Group as Group>::Elem)]
409#[cfg(feature = "alloc")]
410pub struct VoprfServerBatchEvaluateResult<CS: CipherSuite> {
411    /// The messages to send to the client
412    pub messages: Vec<EvaluationElement<CS>>,
413    /// The proof for the client to verify
414    pub proof: Proof<CS>,
415}
416
417/// Concrete type of [`EvaluationElement`]s returned by
418/// [`VoprfServer::batch_blind_evaluate_prepare`].
419pub type VoprfServerBatchEvaluatePreparedEvaluationElements<CS, I> = Map<
420    Zip<I, Repeat<<<CS as CipherSuite>::Group as Group>::Scalar>>,
421    fn(
422        (
423            &BlindedElement<CS>,
424            <<CS as CipherSuite>::Group as Group>::Scalar,
425        ),
426    ) -> PreparedEvaluationElement<CS>,
427>;
428
429/// Concrete type of [`EvaluationElement`]s in
430/// [`VoprfServerBatchEvaluateFinishResult`].
431pub type VoprfServerBatchEvaluateFinishedMessages<'a, CS, I> = Map<
432    <&'a I as IntoIterator>::IntoIter,
433    fn(&PreparedEvaluationElement<CS>) -> EvaluationElement<CS>,
434>;
435
436/// Contains the fields that are returned by a verifiable server batch evaluate
437/// finish.
438#[derive_where(Debug; <&'a I as IntoIterator>::IntoIter, <CS::Group as Group>::Scalar)]
439pub struct VoprfServerBatchEvaluateFinishResult<'a, CS: 'a + CipherSuite, I>
440where
441    &'a I: IntoIterator<Item = &'a PreparedEvaluationElement<CS>>,
442{
443    /// The [`EvaluationElement`]s to send to the client
444    pub messages: VoprfServerBatchEvaluateFinishedMessages<'a, CS, I>,
445    /// The proof for the client to verify
446    pub proof: Proof<CS>,
447}
448
449/////////////////////
450// Inner functions //
451// =============== //
452/////////////////////
453
454type VoprfUnblindResult<'a, CS, IC, IM> = Map<
455    Zip<
456        Map<
457            <&'a IC as IntoIterator>::IntoIter,
458            fn(&VoprfClient<CS>) -> <<CS as CipherSuite>::Group as Group>::Scalar,
459        >,
460        <&'a IM as IntoIterator>::IntoIter,
461    >,
462    fn(
463        (
464            <<CS as CipherSuite>::Group as Group>::Scalar,
465            &EvaluationElement<CS>,
466        ),
467    ) -> <<CS as CipherSuite>::Group as Group>::Elem,
468>;
469
470/// Can only fail with [`Error::Batch] or [`Error::ProofVerification`].
471fn verifiable_unblind<'a, CS: 'a + CipherSuite, IC, IM>(
472    clients: &'a IC,
473    messages: &'a IM,
474    pk: <CS::Group as Group>::Elem,
475    proof: &Proof<CS>,
476) -> Result<VoprfUnblindResult<'a, CS, IC, IM>>
477where
478    &'a IC: 'a + IntoIterator<Item = &'a VoprfClient<CS>>,
479    <&'a IC as IntoIterator>::IntoIter: ExactSizeIterator,
480    &'a IM: 'a + IntoIterator<Item = &'a EvaluationElement<CS>>,
481    <&'a IM as IntoIterator>::IntoIter: ExactSizeIterator,
482{
483    let g = CS::Group::base_elem();
484
485    let blinds = clients
486        .into_iter()
487        // Convert to `fn` pointer to make a return type possible.
488        .map(<fn(&VoprfClient<CS>) -> _>::from(|x| x.blind));
489    let evaluation_elements = messages.into_iter().map(|element| element.0);
490    let blinded_elements = clients.into_iter().map(|client| client.blinded_element);
491
492    verify_proof(
493        g,
494        pk,
495        blinded_elements,
496        evaluation_elements,
497        proof,
498        Mode::Voprf,
499    )?;
500
501    Ok(blinds
502        .zip(messages)
503        .map(|(blind, x)| x.0 * &CS::Group::invert_scalar(blind)))
504}
505
506///////////
507// Tests //
508// ===== //
509///////////
510
511#[cfg(test)]
512mod tests {
513    use core::ptr;
514
515    use ::alloc::vec;
516    use ::alloc::vec::Vec;
517    use rand::rngs::SysRng;
518
519    use super::*;
520    use crate::Group;
521    use crate::common::{Dst, STR_HASH_TO_GROUP};
522    use crate::tests::helpers::prf;
523
524    fn verifiable_retrieval<CS: CipherSuite>() {
525        let input = b"input";
526        let mut rng = SysRng;
527        let client_blind_result = VoprfClient::<CS>::blind(input, &mut rng).unwrap();
528        let server = VoprfServer::<CS>::new(&mut rng).unwrap();
529        let server_result = server.blind_evaluate(&mut rng, &client_blind_result.message);
530        let client_finalize_result = client_blind_result
531            .state
532            .finalize(
533                input,
534                &server_result.message,
535                &server_result.proof,
536                server.get_public_key(),
537            )
538            .unwrap();
539        let res2 = prf::<CS>(input, server.get_private_key(), Mode::Voprf);
540        assert_eq!(client_finalize_result, res2);
541    }
542
543    fn verifiable_batch_retrieval<CS: CipherSuite>() {
544        let mut rng = SysRng;
545        let mut inputs = vec![];
546        let mut client_states = vec![];
547        let mut client_messages = vec![];
548        let num_iterations = 10;
549        for _ in 0..num_iterations {
550            let mut input = [0u8; 32];
551            rng.try_fill_bytes(&mut input).unwrap();
552            let client_blind_result = VoprfClient::<CS>::blind(&input, &mut rng).unwrap();
553            inputs.push(input);
554            client_states.push(client_blind_result.state);
555            client_messages.push(client_blind_result.message);
556        }
557        let server = VoprfServer::<CS>::new(&mut rng).unwrap();
558        let prepared_evaluation_elements: Vec<_> = server
559            .batch_blind_evaluate_prepare(client_messages.iter())
560            .collect();
561        let VoprfServerBatchEvaluateFinishResult { messages, proof } = server
562            .batch_blind_evaluate_finish(
563                &mut rng,
564                client_messages.iter(),
565                &prepared_evaluation_elements,
566            )
567            .unwrap();
568        let messages: Vec<_> = messages.collect();
569        let client_finalize_result = VoprfClient::batch_finalize(
570            &inputs,
571            &client_states,
572            &messages,
573            &proof,
574            server.get_public_key(),
575        )
576        .unwrap()
577        .collect::<Result<Vec<_>>>()
578        .unwrap();
579        let mut res2 = vec![];
580        for input in inputs.iter().take(num_iterations) {
581            let output = prf::<CS>(input, server.get_private_key(), Mode::Voprf);
582            res2.push(output);
583        }
584        assert_eq!(client_finalize_result, res2);
585    }
586
587    fn verifiable_batch_bad_public_key<CS: CipherSuite>() {
588        let mut rng = SysRng;
589        let mut inputs = vec![];
590        let mut client_states = vec![];
591        let mut client_messages = vec![];
592        let num_iterations = 10;
593        for _ in 0..num_iterations {
594            let mut input = [0u8; 32];
595            rng.try_fill_bytes(&mut input).unwrap();
596            let client_blind_result = VoprfClient::<CS>::blind(&input, &mut rng).unwrap();
597            inputs.push(input);
598            client_states.push(client_blind_result.state);
599            client_messages.push(client_blind_result.message);
600        }
601        let server = VoprfServer::<CS>::new(&mut rng).unwrap();
602        let prepared_evaluation_elements: Vec<_> = server
603            .batch_blind_evaluate_prepare(client_messages.iter())
604            .collect();
605        let VoprfServerBatchEvaluateFinishResult { messages, proof } = server
606            .batch_blind_evaluate_finish(
607                &mut rng,
608                client_messages.iter(),
609                &prepared_evaluation_elements,
610            )
611            .unwrap();
612        let messages: Vec<_> = messages.collect();
613        let wrong_pk = {
614            let dst = Dst::new::<CS, _>(STR_HASH_TO_GROUP, Mode::Oprf);
615            // Choose a group element that is unlikely to be the right public key
616            CS::Group::hash_to_curve::<CS::Hash>(&[b"msg"], &dst.as_dst()).unwrap()
617        };
618        let client_finalize_result =
619            VoprfClient::batch_finalize(&inputs, &client_states, &messages, &proof, wrong_pk);
620        assert!(client_finalize_result.is_err());
621    }
622
623    fn verifiable_bad_public_key<CS: CipherSuite>() {
624        let input = b"input";
625        let mut rng = SysRng;
626        let client_blind_result = VoprfClient::<CS>::blind(input, &mut rng).unwrap();
627        let server = VoprfServer::<CS>::new(&mut rng).unwrap();
628        let server_result = server.blind_evaluate(&mut rng, &client_blind_result.message);
629        let wrong_pk = {
630            let dst = Dst::new::<CS, _>(STR_HASH_TO_GROUP, Mode::Oprf);
631            // Choose a group element that is unlikely to be the right public key
632            CS::Group::hash_to_curve::<CS::Hash>(&[b"msg"], &dst.as_dst()).unwrap()
633        };
634        let client_finalize_result = client_blind_result.state.finalize(
635            input,
636            &server_result.message,
637            &server_result.proof,
638            wrong_pk,
639        );
640        assert!(client_finalize_result.is_err());
641    }
642
643    fn verifiable_server_evaluate<CS: CipherSuite>() {
644        let input = b"input";
645        let mut rng = SysRng;
646        let client_blind_result = VoprfClient::<CS>::blind(input, &mut rng).unwrap();
647        let server = VoprfServer::<CS>::new(&mut rng).unwrap();
648        let server_result = server.blind_evaluate(&mut rng, &client_blind_result.message);
649
650        let client_finalize = client_blind_result
651            .state
652            .finalize(
653                input,
654                &server_result.message,
655                &server_result.proof,
656                server.get_public_key(),
657            )
658            .unwrap();
659
660        // We expect the outputs from client and server to be equal given an identical
661        // input
662        let server_evaluate = server.evaluate(input).unwrap();
663        assert_eq!(client_finalize, server_evaluate);
664
665        // We expect the outputs from client and server to be different given different
666        // inputs
667        let wrong_input = b"wrong input";
668        let server_evaluate = server.evaluate(wrong_input).unwrap();
669        assert_ne!(client_finalize, server_evaluate);
670    }
671
672    fn zeroize_voprf_client<CS: CipherSuite>() {
673        let input = b"input";
674        let mut rng = SysRng;
675        let client_blind_result = VoprfClient::<CS>::blind(input, &mut rng).unwrap();
676
677        let mut state = client_blind_result.state;
678        unsafe { ptr::drop_in_place(&mut state) };
679        assert!(state.serialize().iter().all(|&x| x == 0));
680
681        let mut message = client_blind_result.message;
682        unsafe { ptr::drop_in_place(&mut message) };
683        assert!(message.serialize().iter().all(|&x| x == 0));
684    }
685
686    fn zeroize_voprf_server<CS: CipherSuite>() {
687        let input = b"input";
688        let mut rng = SysRng;
689        let client_blind_result = VoprfClient::<CS>::blind(input, &mut rng).unwrap();
690        let server = VoprfServer::<CS>::new(&mut rng).unwrap();
691        let server_result = server.blind_evaluate(&mut rng, &client_blind_result.message);
692
693        let mut state = server;
694        unsafe { ptr::drop_in_place(&mut state) };
695        assert!(state.serialize().iter().all(|&x| x == 0));
696
697        let mut message = server_result.message;
698        unsafe { ptr::drop_in_place(&mut message) };
699        assert!(message.serialize().iter().all(|&x| x == 0));
700
701        let mut proof = server_result.proof;
702        unsafe { ptr::drop_in_place(&mut proof) };
703        assert!(proof.serialize().iter().all(|&x| x == 0));
704    }
705
706    crate::tests::test_all_curves!(
707        verifiable_retrieval,
708        verifiable_batch_retrieval,
709        verifiable_bad_public_key,
710        verifiable_batch_bad_public_key,
711        verifiable_server_evaluate,
712        zeroize_voprf_client,
713        zeroize_voprf_server,
714    );
715}