1#[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#[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#[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
62impl<CS: CipherSuite> VoprfClient<CS> {
68 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 #[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 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 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 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 #[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 #[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 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 Self::new_from_seed(&seed, &[])
199 }
200
201 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 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 #[cfg(test)]
229 pub fn get_private_key(&self) -> <CS::Group as Group>::Scalar {
230 self.sk
231 }
232
233 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 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 #[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 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 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 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 pub fn get_public_key(&self) -> <CS::Group as Group>::Elem {
372 self.pk
373 }
374}
375
376#[derive_where(Debug; <CS::Group as Group>::Scalar, <CS::Group as Group>::Elem)]
383pub struct VoprfClientBlindResult<CS: CipherSuite> {
384 pub state: VoprfClient<CS>,
386 pub message: BlindedElement<CS>,
388}
389
390pub 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#[derive_where(Debug; <CS::Group as Group>::Scalar, <CS::Group as Group>::Elem)]
400pub struct VoprfServerEvaluateResult<CS: CipherSuite> {
401 pub message: EvaluationElement<CS>,
403 pub proof: Proof<CS>,
405}
406
407#[derive_where(Debug; <CS::Group as Group>::Scalar, <CS::Group as Group>::Elem)]
409#[cfg(feature = "alloc")]
410pub struct VoprfServerBatchEvaluateResult<CS: CipherSuite> {
411 pub messages: Vec<EvaluationElement<CS>>,
413 pub proof: Proof<CS>,
415}
416
417pub 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
429pub type VoprfServerBatchEvaluateFinishedMessages<'a, CS, I> = Map<
432 <&'a I as IntoIterator>::IntoIter,
433 fn(&PreparedEvaluationElement<CS>) -> EvaluationElement<CS>,
434>;
435
436#[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 pub messages: VoprfServerBatchEvaluateFinishedMessages<'a, CS, I>,
445 pub proof: Proof<CS>,
447}
448
449type 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
470fn 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 .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#[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 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 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 let server_evaluate = server.evaluate(input).unwrap();
663 assert_eq!(client_finalize, server_evaluate);
664
665 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}