1#![allow(unsafe_code)]
12
13use derive_where::derive_where;
14use digest::{Output, OutputSizeUser};
15use generic_array::{ArrayLength, GenericArray};
16use rand::{CryptoRng, RngCore};
17
18use crate::ciphersuite::CipherSuite;
19use crate::errors::ProtocolError;
20use crate::key_exchange::group::Group;
21use crate::key_exchange::shared::DiffieHellman;
22use crate::key_exchange::sigma_i::{Message, MessageBuilder, SignatureProtocol};
23use crate::serialization::SliceExt;
24
25#[cfg_attr(
27 feature = "serde",
28 derive(serde::Deserialize, serde::Serialize),
29 serde(bound(
30 deserialize = "G::Pk: serde::Deserialize<'de>, SK: serde::Deserialize<'de>",
31 serialize = "G::Pk: serde::Serialize, SK: serde::Serialize"
32 ))
33)]
34#[derive_where(Clone)]
35#[derive_where(Eq, Hash, Ord, PartialEq, PartialOrd; G::Pk, SK)]
36#[cfg_attr(not(test), derive_where(Debug; G::Pk, SK))]
39#[cfg_attr(test, derive_where(Debug), derive_where(skip_inner(Debug)))]
40pub struct KeyPair<G: Group, SK: Clone = PrivateKey<G>> {
41 pk: PublicKey<G>,
42 sk: SK,
43}
44
45impl<G: Group, SK: Clone> KeyPair<G, SK> {
46 pub fn new(sk: SK, pk: PublicKey<G>) -> Self {
48 Self { pk, sk }
49 }
50
51 pub fn public(&self) -> &PublicKey<G> {
53 &self.pk
54 }
55
56 pub fn private(&self) -> &SK {
58 &self.sk
59 }
60}
61
62impl<G: Group> KeyPair<G> {
63 pub(crate) fn random<R: RngCore + CryptoRng>(rng: &mut R) -> Self {
64 let sk = G::random_sk(rng);
65 let pk = G::public_key(&sk);
66 Self {
67 pk: PublicKey(pk),
68 sk: PrivateKey(sk),
69 }
70 }
71
72 pub(crate) fn derive_random<R: RngCore + CryptoRng>(rng: &mut R) -> Self {
74 let mut scalar_bytes = GenericArray::<_, <G as Group>::SkLen>::default();
75 rng.fill_bytes(&mut scalar_bytes);
76 let sk = G::derive_scalar(scalar_bytes).unwrap();
77 let pk = G::public_key(&sk);
78 Self {
79 pk: PublicKey(pk),
80 sk: PrivateKey(sk),
81 }
82 }
83}
84
85#[cfg_attr(
87 feature = "serde",
88 derive(serde::Deserialize, serde::Serialize),
89 serde(bound(
90 deserialize = "G::Sk: serde::Deserialize<'de>",
91 serialize = "G::Sk: serde::Serialize"
92 ))
93)]
94#[derive_where(Clone, ZeroizeOnDrop)]
95#[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; G::Sk)]
96pub struct PrivateKey<G: Group>(G::Sk);
97
98impl<G: Group> PrivateKey<G> {
99 pub(crate) fn new(key: G::Sk) -> Self {
100 Self(key)
101 }
102
103 pub fn public_key(&self) -> PublicKey<G> {
105 PublicKey(G::public_key(&self.0))
106 }
107
108 pub fn serialize(&self) -> GenericArray<u8, G::SkLen> {
110 G::serialize_sk(&self.0)
111 }
112
113 pub fn deserialize(mut input: &[u8]) -> Result<Self, ProtocolError> {
115 Self::deserialize_take(&mut input)
116 }
117
118 pub(crate) fn deserialize_take(input: &mut &[u8]) -> Result<Self, ProtocolError> {
119 G::deserialize_take_sk(input).map(Self)
120 }
121}
122
123impl<G: Group> PrivateKey<G>
124where
125 G::Sk: DiffieHellman<G>,
126{
127 pub(crate) fn ke_diffie_hellman(&self, pk: &PublicKey<G>) -> GenericArray<u8, G::PkLen> {
129 self.0.diffie_hellman(&pk.0)
130 }
131}
132
133impl<G: Group> PrivateKey<G> {
134 pub(crate) fn sign<
136 R: CryptoRng + RngCore,
137 CS: CipherSuite,
138 SIG: SignatureProtocol<Group = G>,
139 KE: Group,
140 >(
141 &self,
142 rng: &mut R,
143 message: &Message<CS, KE>,
144 ) -> (SIG::Signature, SIG::VerifyState<CS, KE>) {
145 SIG::sign(&self.0, rng, message)
146 }
147}
148
149pub trait PrivateKeySerialization<G: Group>: Clone {
152 type Error;
154 type Len: ArrayLength<u8>;
156
157 fn serialize_key_pair(key_pair: &KeyPair<G, Self>) -> GenericArray<u8, Self::Len>;
159
160 fn deserialize_take_key_pair(
164 bytes: &mut &[u8],
165 ) -> Result<KeyPair<G, Self>, ProtocolError<Self::Error>>;
166}
167
168impl<G: Group> PrivateKeySerialization<G> for PrivateKey<G> {
169 type Error = core::convert::Infallible;
170 type Len = G::SkLen;
171
172 fn serialize_key_pair(key_pair: &KeyPair<G, Self>) -> GenericArray<u8, Self::Len> {
173 key_pair.private().serialize()
174 }
175
176 fn deserialize_take_key_pair(input: &mut &[u8]) -> Result<KeyPair<G, Self>, ProtocolError> {
177 let sk = PrivateKey::deserialize_take(input)?;
178 let pk = sk.public_key();
179
180 Ok(KeyPair::new(sk, pk))
181 }
182}
183
184#[cfg_attr(
186 feature = "serde",
187 derive(serde::Deserialize, serde::Serialize),
188 serde(bound(
189 deserialize = "G::Pk: serde::Deserialize<'de>",
190 serialize = "G::Pk: serde::Serialize"
191 ))
192)]
193#[derive_where(Clone)]
194#[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; G::Pk)]
195pub struct PublicKey<G: Group + ?Sized>(G::Pk);
196
197impl<G: Group> PublicKey<G> {
198 pub fn deserialize(mut key_bytes: &[u8]) -> Result<Self, ProtocolError> {
200 Self::deserialize_take(&mut key_bytes)
201 }
202
203 pub(crate) fn deserialize_take(key_bytes: &mut &[u8]) -> Result<Self, ProtocolError> {
204 G::deserialize_take_pk(key_bytes).map(Self)
205 }
206
207 pub fn serialize(&self) -> GenericArray<u8, G::PkLen> {
209 G::serialize_pk(&self.0)
210 }
211
212 pub fn to_group_type(&self) -> &G::Pk {
214 &self.0
215 }
216}
217
218impl<G: Group> PublicKey<G> {
219 pub(crate) fn verify<CS: CipherSuite, SIG: SignatureProtocol<Group = G>, KE: Group>(
221 &self,
222 message_builder: MessageBuilder<'_, CS>,
223 state: SIG::VerifyState<CS, KE>,
224 signature: &SIG::Signature,
225 ) -> Result<(), ProtocolError> {
226 SIG::verify(&self.0, message_builder, state, signature)
227 }
228}
229
230#[cfg_attr(
232 feature = "serde",
233 derive(serde::Deserialize, serde::Serialize),
234 serde(bound = "")
235)]
236#[derive_where(Clone, Debug, Eq, Hash, PartialEq, ZeroizeOnDrop)]
237pub struct OprfSeed<H: OutputSizeUser>(pub(crate) Output<H>);
238
239pub trait OprfSeedSerialization<H, E>: Sized {
244 type Len: ArrayLength<u8>;
246
247 fn serialize(&self) -> GenericArray<u8, Self::Len>;
249
250 fn deserialize_take(bytes: &mut &[u8]) -> Result<Self, ProtocolError<E>>;
254}
255
256impl<H: OutputSizeUser, E> OprfSeedSerialization<H, E> for OprfSeed<H> {
257 type Len = H::OutputSize;
258
259 fn serialize(&self) -> GenericArray<u8, Self::Len> {
260 self.0.clone()
261 }
262
263 fn deserialize_take(input: &mut &[u8]) -> Result<Self, ProtocolError<E>> {
264 Ok(Self(
265 input
266 .take_array("OPRF seed")
267 .map_err(ProtocolError::into_custom)?,
268 ))
269 }
270}
271
272#[cfg(test)]
278impl<G: Group> KeyPair<G> {
279 fn uniform_keypair_strategy() -> proptest::prelude::BoxedStrategy<Self> {
282 use proptest::prelude::*;
283 use rand::SeedableRng;
284 use rand::rngs::StdRng;
285
286 any::<[u8; 32]>()
289 .prop_filter_map("valid random keypair", |seed| {
290 let mut rng = StdRng::from_seed(seed);
291 Some(Self::derive_random(&mut rng))
292 })
293 .no_shrink()
294 .boxed()
295 }
296}
297
298#[cfg(test)]
299mod tests {
300 use hkdf::Hkdf;
301 use rand::rngs::OsRng;
302
303 use super::*;
304 use crate::ciphersuite::{KeGroup, OprfHash};
305 use crate::{
306 CipherSuite, ClientLogin, ClientLoginFinishParameters, ClientLoginFinishResult,
307 ClientLoginStartResult, ClientRegistration, ClientRegistrationFinishParameters,
308 ClientRegistrationFinishResult, ClientRegistrationStartResult, ServerLogin,
309 ServerLoginParameters, ServerLoginStartResult, ServerRegistration,
310 ServerRegistrationStartResult, ServerSetup,
311 };
312
313 macro_rules! test {
314 ($mod:ident, $point:ty) => {
315 mod $mod {
316
317 use proptest::prelude::*;
318
319 use super::*;
320
321 proptest! {
322 #[test]
323 fn pub_from_priv(kp in KeyPair::<$point>::uniform_keypair_strategy()) {
324 let pk = kp.public();
325 let sk = kp.private();
326 prop_assert_eq!(&sk.public_key(), pk);
327 }
328
329 #[test]
330 fn dh(kp1 in KeyPair::<$point>::uniform_keypair_strategy(),
331 kp2 in KeyPair::<$point>::uniform_keypair_strategy()) {
332
333 let dh1 = kp2.private().ke_diffie_hellman(&kp1.public());
334 let dh2 = kp1.private().ke_diffie_hellman(kp2.public());
335
336 prop_assert_eq!(dh1, dh2);
337 }
338
339 #[test]
340 fn private_key_slice(kp in KeyPair::<$point>::uniform_keypair_strategy()) {
341 let sk_bytes = kp.private().serialize().to_vec();
342
343 let kp2 = PrivateKey::<$point>::deserialize_take_key_pair(&mut (sk_bytes.as_slice()))?;
344 let kp2_private_bytes = kp2.private().serialize().to_vec();
345
346 prop_assert_eq!(sk_bytes, kp2_private_bytes);
347 }
348 }
349 }
350 };
351 }
352
353 #[cfg(feature = "ristretto255")]
354 test!(ristretto, crate::Ristretto255);
355 test!(p256, ::p256::NistP256);
356 test!(p384, ::p384::NistP384);
357 test!(p521, ::p521::NistP521);
358
359 struct Default;
360
361 impl CipherSuite for Default {
362 #[cfg(feature = "ristretto255")]
363 type OprfCs = crate::Ristretto255;
364 #[cfg(not(feature = "ristretto255"))]
365 type OprfCs = ::p256::NistP256;
366 #[cfg(feature = "ristretto255")]
367 type KeyExchange = crate::TripleDh<crate::Ristretto255, sha2::Sha512>;
368 #[cfg(not(feature = "ristretto255"))]
369 type KeyExchange = crate::TripleDh<::p256::NistP256, sha2::Sha256>;
370 type Ksf = crate::ksf::Identity;
371 }
372
373 #[derive(Clone)]
374 struct RemoteSeed<H: OutputSizeUser>(Output<H>);
375
376 #[derive(Clone)]
377 struct RemoteKey(PrivateKey<KeGroup<Default>>);
378
379 const PASSWORD: &str = "password";
380
381 #[test]
382 fn remote_key() {
383 let sk = PrivateKey(KeGroup::<Default>::random_sk(&mut OsRng));
384 let pk = sk.public_key();
385 let sk = RemoteKey(sk);
386 let keypair = KeyPair::new(sk, pk);
387
388 let server_setup =
389 ServerSetup::<Default, RemoteKey>::new_with_key_pair(&mut OsRng, keypair);
390
391 let ClientRegistrationStartResult {
392 message,
393 state: client,
394 } = ClientRegistration::<Default>::start(&mut OsRng, PASSWORD.as_bytes()).unwrap();
395 let ServerRegistrationStartResult { message, .. } =
396 ServerRegistration::start(&server_setup, message, &[]).unwrap();
397 let ClientRegistrationFinishResult { message, .. } = client
398 .finish(
399 &mut OsRng,
400 PASSWORD.as_bytes(),
401 message,
402 ClientRegistrationFinishParameters::default(),
403 )
404 .unwrap();
405 let file = ServerRegistration::finish(message);
406
407 let ClientLoginStartResult {
408 message,
409 state: client,
410 } = ClientLogin::<Default>::start(&mut OsRng, PASSWORD.as_bytes()).unwrap();
411 let builder = ServerLogin::builder(
412 &mut OsRng,
413 &server_setup,
414 Some(file),
415 message,
416 &[],
417 ServerLoginParameters::default(),
418 )
419 .unwrap();
420 let shared_secret = builder.private_key().0.ke_diffie_hellman(builder.data());
421 let ServerLoginStartResult {
422 message,
423 state: server,
424 ..
425 } = builder.build(shared_secret).unwrap();
426 let ClientLoginFinishResult { message, .. } = client
427 .finish(
428 &mut OsRng,
429 PASSWORD.as_bytes(),
430 message,
431 ClientLoginFinishParameters::default(),
432 )
433 .unwrap();
434 server
435 .finish(message, ServerLoginParameters::default())
436 .unwrap();
437 }
438
439 #[test]
440 fn remote_seed() {
441 let mut oprf_seed = RemoteSeed::<OprfHash<Default>>(GenericArray::default());
442 OsRng.fill_bytes(&mut oprf_seed.0);
443
444 let sk = PrivateKey(KeGroup::<Default>::random_sk(&mut OsRng));
445 let pk = sk.public_key();
446 let sk = RemoteKey(sk);
447 let keypair = KeyPair::new(sk, pk);
448
449 let server_setup = ServerSetup::<Default, _, _>::new_with_key_pair_and_seed(
450 &mut OsRng, keypair, oprf_seed,
451 );
452
453 let ClientRegistrationStartResult {
454 message,
455 state: client,
456 } = ClientRegistration::<Default>::start(&mut OsRng, PASSWORD.as_bytes()).unwrap();
457 let km = server_setup.key_material_info(&[]);
458 let mut ikm = GenericArray::default();
459 Hkdf::<OprfHash<Default>>::from_prk(&km.ikm.0)
460 .unwrap()
461 .expand_multi_info(&km.info, &mut ikm)
462 .unwrap();
463 let ServerRegistrationStartResult { message, .. } =
464 ServerRegistration::start_with_key_material(&server_setup, ikm, message).unwrap();
465 let ClientRegistrationFinishResult { message, .. } = client
466 .finish(
467 &mut OsRng,
468 PASSWORD.as_bytes(),
469 message,
470 ClientRegistrationFinishParameters::default(),
471 )
472 .unwrap();
473 let file = ServerRegistration::finish(message);
474
475 let ClientLoginStartResult {
476 message,
477 state: client,
478 } = ClientLogin::<Default>::start(&mut OsRng, PASSWORD.as_bytes()).unwrap();
479 let km = server_setup.key_material_info(&[]);
480 let mut ikm = GenericArray::default();
481 Hkdf::<OprfHash<Default>>::from_prk(&km.ikm.0)
482 .unwrap()
483 .expand_multi_info(&km.info, &mut ikm)
484 .unwrap();
485 let builder = ServerLogin::builder_with_key_material(
486 &mut OsRng,
487 &server_setup,
488 ikm,
489 Some(file),
490 message,
491 ServerLoginParameters::default(),
492 )
493 .unwrap();
494 let shared_secret = builder.private_key().0.ke_diffie_hellman(builder.data());
495 let ServerLoginStartResult {
496 message,
497 state: server,
498 ..
499 } = builder.build(shared_secret).unwrap();
500 let ClientLoginFinishResult { message, .. } = client
501 .finish(
502 &mut OsRng,
503 PASSWORD.as_bytes(),
504 message,
505 ClientLoginFinishParameters::default(),
506 )
507 .unwrap();
508 server
509 .finish(message, ServerLoginParameters::default())
510 .unwrap();
511 }
512}