1use core::iter::{self};
8
9use derive_where::derive_where;
10use digest::Output;
11use hybrid_array::Array;
12use rand_core::{TryCryptoRng, TryRng};
13
14use crate::common::{
15 BlindedElement, EvaluationElement, Mode, derive_key_internal, deterministic_blind_unchecked,
16 finalize_after_unblind, hash_to_group, server_evaluate_hash_input,
17};
18#[cfg(feature = "serde")]
19use crate::serialization::serde::Scalar;
20use crate::{CipherSuite, Error, Group, Result};
21
22#[derive_where(Clone, ZeroizeOnDrop)]
35#[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; <CS::Group as Group>::Scalar)]
36#[cfg_attr(
37 feature = "serde",
38 derive(serde::Deserialize, serde::Serialize),
39 serde(bound = "")
40)]
41pub struct OprfClient<CS: CipherSuite> {
42 #[cfg_attr(feature = "serde", serde(with = "Scalar::<CS::Group>"))]
43 pub(crate) blind: <CS::Group as Group>::Scalar,
44}
45
46#[derive_where(Clone, ZeroizeOnDrop)]
49#[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; <CS::Group as Group>::Scalar)]
50#[cfg_attr(
51 feature = "serde",
52 derive(serde::Deserialize, serde::Serialize),
53 serde(bound = "")
54)]
55pub struct OprfServer<CS: CipherSuite> {
56 #[cfg_attr(feature = "serde", serde(with = "Scalar::<CS::Group>"))]
57 pub(crate) sk: <CS::Group as Group>::Scalar,
58}
59
60impl<CS: CipherSuite> OprfClient<CS> {
66 pub fn blind<R: TryRng + TryCryptoRng>(
72 input: &[u8],
73 blinding_factor_rng: &mut R,
74 ) -> Result<OprfClientBlindResult<CS>> {
75 let blind = CS::Group::random_scalar(blinding_factor_rng)?;
76 Self::deterministic_blind_unchecked_inner(input, blind)
77 }
78
79 #[cfg(any(feature = "danger", test))]
91 pub fn deterministic_blind_unchecked(
92 input: &[u8],
93 blind: <CS::Group as Group>::Scalar,
94 ) -> Result<OprfClientBlindResult<CS>> {
95 Self::deterministic_blind_unchecked_inner(input, blind)
96 }
97
98 fn deterministic_blind_unchecked_inner(
100 input: &[u8],
101 blind: <CS::Group as Group>::Scalar,
102 ) -> Result<OprfClientBlindResult<CS>> {
103 let blinded_element = deterministic_blind_unchecked::<CS>(input, &blind, Mode::Oprf)?;
104 Ok(OprfClientBlindResult {
105 state: Self { blind },
106 message: BlindedElement(blinded_element),
107 })
108 }
109
110 pub fn finalize(
116 &self,
117 input: &[u8],
118 evaluation_element: &EvaluationElement<CS>,
119 ) -> Result<Output<CS::Hash>> {
120 let unblinded_element = evaluation_element.0 * &CS::Group::invert_scalar(self.blind);
121 let mut outputs =
122 finalize_after_unblind::<CS, _, _>(iter::once((input, unblinded_element)));
123 outputs.next().unwrap()
124 }
125
126 #[cfg(test)]
128 pub fn from_blind(blind: <CS::Group as Group>::Scalar) -> Self {
129 Self { blind }
130 }
131
132 #[cfg(feature = "danger")]
134 pub fn get_blind(&self) -> <CS::Group as Group>::Scalar {
135 self.blind
136 }
137}
138
139impl<CS: CipherSuite> OprfServer<CS> {
140 pub fn new<R: TryRng + TryCryptoRng>(rng: &mut R) -> Result<Self> {
145 let mut seed = Array::<_, <CS::Group as Group>::ScalarLen>::default();
146 rng.try_fill_bytes(&mut seed).map_err(|_| Error::Protocol)?;
147 Self::new_from_seed(&seed, &[])
148 }
149
150 pub fn new_with_key(private_key_bytes: &[u8]) -> Result<Self> {
157 let sk = CS::Group::deserialize_scalar(private_key_bytes)?;
158 Ok(Self { sk })
159 }
160
161 pub fn new_from_seed(seed: &[u8], info: &[u8]) -> Result<Self> {
171 let sk = derive_key_internal::<CS>(seed, info, Mode::Oprf)?;
172 Ok(Self { sk })
173 }
174
175 #[cfg(test)]
177 pub fn get_private_key(&self) -> <CS::Group as Group>::Scalar {
178 self.sk
179 }
180
181 pub fn blind_evaluate(&self, blinded_element: &BlindedElement<CS>) -> EvaluationElement<CS> {
185 EvaluationElement(blinded_element.0 * &self.sk)
186 }
187
188 pub fn evaluate(&self, input: &[u8]) -> Result<Output<<CS as CipherSuite>::Hash>> {
193 let input_element = hash_to_group::<CS>(input, Mode::Oprf)?;
194 if CS::Group::is_identity_elem(input_element).into() {
195 return Err(Error::Input);
196 };
197 let evaluated_element = input_element * &self.sk;
198
199 let issued_element = CS::Group::serialize_elem(evaluated_element);
200
201 server_evaluate_hash_input::<CS>(input, None, issued_element)
202 }
203}
204
205#[derive_where(Debug; <CS::Group as Group>::Scalar, <CS::Group as Group>::Elem)]
212pub struct OprfClientBlindResult<CS: CipherSuite> {
213 pub state: OprfClient<CS>,
215 pub message: BlindedElement<CS>,
217}
218
219#[cfg(test)]
225mod tests {
226 use core::ptr;
227
228 use rand::TryRng;
229 use rand::rngs::SysRng;
230
231 use super::*;
232 use crate::Group;
233 use crate::common::{Dst, STR_HASH_TO_GROUP};
234 use crate::tests::helpers::prf;
235
236 fn base_retrieval<CS: CipherSuite>() {
237 let input = b"input";
238 let mut rng = SysRng;
239 let client_blind_result = OprfClient::<CS>::blind(input, &mut rng).unwrap();
240 let server = OprfServer::<CS>::new(&mut rng).unwrap();
241 let message = server.blind_evaluate(&client_blind_result.message);
242 let client_finalize_result = client_blind_result.state.finalize(input, &message).unwrap();
243 let res2 = prf::<CS>(input, server.get_private_key(), Mode::Oprf);
244 assert_eq!(client_finalize_result, res2);
245 }
246
247 fn base_inversion_unsalted<CS: CipherSuite>() {
248 let mut rng = SysRng;
249 let mut input = [0u8; 64];
250 rng.try_fill_bytes(&mut input).unwrap();
251 let client_blind_result = OprfClient::<CS>::blind(&input, &mut rng).unwrap();
252 let client_finalize_result = client_blind_result
253 .state
254 .finalize(&input, &EvaluationElement(client_blind_result.message.0))
255 .unwrap();
256
257 let dst = Dst::new::<CS, _>(STR_HASH_TO_GROUP, Mode::Oprf);
258 let point = CS::Group::hash_to_curve::<CS::Hash>(&[&input], &dst.as_dst()).unwrap();
259 let res2 = finalize_after_unblind::<CS, _, _>(iter::once((input.as_ref(), point)))
260 .next()
261 .unwrap()
262 .unwrap();
263
264 assert_eq!(client_finalize_result, res2);
265 }
266
267 fn server_evaluate<CS: CipherSuite>() {
268 let input = b"input";
269 let mut rng = SysRng;
270 let client_blind_result = OprfClient::<CS>::blind(input, &mut rng).unwrap();
271 let server = OprfServer::<CS>::new(&mut rng).unwrap();
272 let server_result = server.blind_evaluate(&client_blind_result.message);
273
274 let client_finalize = client_blind_result
275 .state
276 .finalize(input, &server_result)
277 .unwrap();
278
279 let server_evaluate = server.evaluate(input).unwrap();
282 assert_eq!(client_finalize, server_evaluate);
283
284 let wrong_input = b"wrong input";
287 let server_evaluate = server.evaluate(wrong_input).unwrap();
288 assert!(client_finalize != server_evaluate);
289 }
290
291 fn zeroize_oprf_client<CS: CipherSuite>() {
292 let input = b"input";
293 let mut rng = SysRng;
294 let client_blind_result = OprfClient::<CS>::blind(input, &mut rng).unwrap();
295
296 let mut state = client_blind_result.state;
297 unsafe { ptr::drop_in_place(&mut state) };
298 assert!(state.serialize().iter().all(|&x| x == 0));
299
300 let mut message = client_blind_result.message;
301 unsafe { ptr::drop_in_place(&mut message) };
302 assert!(message.serialize().iter().all(|&x| x == 0));
303 }
304
305 fn zeroize_oprf_server<CS: CipherSuite>() {
306 let input = b"input";
307 let mut rng = SysRng;
308 let client_blind_result = OprfClient::<CS>::blind(input, &mut rng).unwrap();
309 let server = OprfServer::<CS>::new(&mut rng).unwrap();
310 let mut message = server.blind_evaluate(&client_blind_result.message);
311
312 let mut state = server;
313 unsafe { ptr::drop_in_place(&mut state) };
314 assert!(state.serialize().iter().all(|&x| x == 0));
315
316 unsafe { ptr::drop_in_place(&mut message) };
317 assert!(message.serialize().iter().all(|&x| x == 0));
318 }
319
320 crate::tests::test_all_curves!(
321 base_retrieval,
322 base_inversion_unsalted,
323 server_evaluate,
324 zeroize_oprf_client,
325 zeroize_oprf_server,
326 );
327}