1#![allow(unsafe_code)]
8
9use derive_where::derive_where;
10use digest::{Output, OutputSizeUser};
11use generic_array::{ArrayLength, GenericArray};
12use rand::{CryptoRng, Rng};
13
14use crate::ciphersuite::CipherSuite;
15use crate::errors::ProtocolError;
16use crate::key_exchange::group::Group;
17use crate::key_exchange::shared::DiffieHellman;
18use crate::key_exchange::sigma_i::{Message, MessageBuilder, SignatureProtocol};
19use crate::serialization::SliceExt;
20
21#[cfg_attr(
23 feature = "serde",
24 derive(serde::Deserialize, serde::Serialize),
25 serde(bound(
26 deserialize = "G::Pk: serde::Deserialize<'de>, SK: serde::Deserialize<'de>",
27 serialize = "G::Pk: serde::Serialize, SK: serde::Serialize"
28 ))
29)]
30#[derive_where(Clone)]
31#[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; G::Pk, SK)]
32pub struct KeyPair<G: Group, SK: Clone = PrivateKey<G>> {
33 pk: PublicKey<G>,
34 sk: SK,
35}
36
37impl<G: Group, SK: Clone> KeyPair<G, SK> {
38 pub fn new(sk: SK, pk: PublicKey<G>) -> Self {
40 Self { pk, sk }
41 }
42
43 pub fn public(&self) -> &PublicKey<G> {
45 &self.pk
46 }
47
48 pub fn private(&self) -> &SK {
50 &self.sk
51 }
52}
53
54impl<G: Group> KeyPair<G> {
55 pub(crate) fn random<R: Rng + CryptoRng>(rng: &mut R) -> Self {
56 let sk = G::random_sk(rng);
57 let pk = G::public_key(&sk);
58 Self {
59 pk: PublicKey(pk),
60 sk: PrivateKey(sk),
61 }
62 }
63
64 pub(crate) fn derive_random<R: Rng + CryptoRng>(rng: &mut R) -> Self {
66 let mut scalar_bytes = GenericArray::<_, <G as Group>::SkLen>::default();
67 rng.fill_bytes(&mut scalar_bytes);
68 let sk = G::derive_scalar(scalar_bytes).unwrap();
69 let pk = G::public_key(&sk);
70 Self {
71 pk: PublicKey(pk),
72 sk: PrivateKey(sk),
73 }
74 }
75}
76
77#[cfg_attr(
79 feature = "serde",
80 derive(serde::Deserialize, serde::Serialize),
81 serde(bound(
82 deserialize = "G::Sk: serde::Deserialize<'de>",
83 serialize = "G::Sk: serde::Serialize"
84 ))
85)]
86#[derive_where(Clone, ZeroizeOnDrop)]
87#[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; G::Sk)]
88pub struct PrivateKey<G: Group>(G::Sk);
89
90impl<G: Group> PrivateKey<G> {
91 pub(crate) fn new(key: G::Sk) -> Self {
92 Self(key)
93 }
94
95 pub fn public_key(&self) -> PublicKey<G> {
97 PublicKey(G::public_key(&self.0))
98 }
99
100 pub fn serialize(&self) -> GenericArray<u8, G::SkLen> {
102 G::serialize_sk(&self.0)
103 }
104
105 pub fn deserialize(mut input: &[u8]) -> Result<Self, ProtocolError> {
107 Self::deserialize_take(&mut input)
108 }
109
110 pub(crate) fn deserialize_take(input: &mut &[u8]) -> Result<Self, ProtocolError> {
111 G::deserialize_take_sk(input).map(Self)
112 }
113}
114
115impl<G: Group> PrivateKey<G>
116where
117 G::Sk: DiffieHellman<G>,
118{
119 pub(crate) fn ke_diffie_hellman(&self, pk: &PublicKey<G>) -> GenericArray<u8, G::PkLen> {
121 self.0.diffie_hellman(&pk.0)
122 }
123}
124
125impl<G: Group> PrivateKey<G> {
126 pub(crate) fn sign<
128 R: CryptoRng + Rng,
129 CS: CipherSuite,
130 SIG: SignatureProtocol<Group = G>,
131 KE: Group,
132 >(
133 &self,
134 rng: &mut R,
135 message: &Message<CS, KE>,
136 ) -> (SIG::Signature, SIG::VerifyState<CS, KE>) {
137 SIG::sign(&self.0, rng, message)
138 }
139}
140
141pub trait PrivateKeySerialization<G: Group>: Clone {
144 type Error;
146 type Len: ArrayLength;
148
149 fn serialize_key_pair(key_pair: &KeyPair<G, Self>) -> GenericArray<u8, Self::Len>;
151
152 fn deserialize_take_key_pair(
156 bytes: &mut &[u8],
157 ) -> Result<KeyPair<G, Self>, ProtocolError<Self::Error>>;
158}
159
160impl<G: Group> PrivateKeySerialization<G> for PrivateKey<G> {
161 type Error = core::convert::Infallible;
162 type Len = G::SkLen;
163
164 fn serialize_key_pair(key_pair: &KeyPair<G, Self>) -> GenericArray<u8, Self::Len> {
165 key_pair.private().serialize()
166 }
167
168 fn deserialize_take_key_pair(input: &mut &[u8]) -> Result<KeyPair<G, Self>, ProtocolError> {
169 let sk = PrivateKey::deserialize_take(input)?;
170 let pk = sk.public_key();
171
172 Ok(KeyPair::new(sk, pk))
173 }
174}
175
176#[cfg_attr(
178 feature = "serde",
179 derive(serde::Deserialize, serde::Serialize),
180 serde(bound(
181 deserialize = "G::Pk: serde::Deserialize<'de>",
182 serialize = "G::Pk: serde::Serialize"
183 ))
184)]
185#[derive_where(Clone)]
186#[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; G::Pk)]
187pub struct PublicKey<G: Group + ?Sized>(G::Pk);
188
189impl<G: Group> PublicKey<G> {
190 pub fn deserialize(mut key_bytes: &[u8]) -> Result<Self, ProtocolError> {
192 Self::deserialize_take(&mut key_bytes)
193 }
194
195 pub(crate) fn deserialize_take(key_bytes: &mut &[u8]) -> Result<Self, ProtocolError> {
196 G::deserialize_take_pk(key_bytes).map(Self)
197 }
198
199 pub fn serialize(&self) -> GenericArray<u8, G::PkLen> {
201 G::serialize_pk(&self.0)
202 }
203
204 pub fn to_group_type(&self) -> &G::Pk {
206 &self.0
207 }
208}
209
210impl<G: Group> PublicKey<G> {
211 pub(crate) fn verify<CS: CipherSuite, SIG: SignatureProtocol<Group = G>, KE: Group>(
213 &self,
214 message_builder: MessageBuilder<'_, CS>,
215 state: SIG::VerifyState<CS, KE>,
216 signature: &SIG::Signature,
217 ) -> Result<(), ProtocolError> {
218 SIG::verify(&self.0, message_builder, state, signature)
219 }
220}
221
222#[cfg_attr(
224 feature = "serde",
225 derive(serde::Deserialize, serde::Serialize),
226 serde(bound = "")
227)]
228#[derive_where(Clone, Debug, Eq, Hash, PartialEq, ZeroizeOnDrop)]
229pub struct OprfSeed<H: OutputSizeUser>(pub(crate) Output<H>);
230
231pub trait OprfSeedSerialization<H, E>: Sized {
236 type Len: ArrayLength;
238
239 fn serialize(&self) -> GenericArray<u8, Self::Len>;
241
242 fn deserialize_take(bytes: &mut &[u8]) -> Result<Self, ProtocolError<E>>;
246}
247
248impl<H: OutputSizeUser, E> OprfSeedSerialization<H, E> for OprfSeed<H>
249where
250 H::OutputSize: ArrayLength,
251{
252 type Len = H::OutputSize;
253
254 fn serialize(&self) -> GenericArray<u8, Self::Len> {
255 GenericArray::from_slice(self.0.as_slice()).clone()
256 }
257
258 fn deserialize_take(input: &mut &[u8]) -> Result<Self, ProtocolError<E>> {
259 Ok(Self(
260 input
261 .take_array("OPRF seed")
262 .map_err(ProtocolError::into_custom)?
263 .into_ha0_4(),
264 ))
265 }
266}
267
268#[cfg(test)]
274impl<G: Group> KeyPair<G>
275where
276 G::Pk: core::fmt::Debug,
277 G::Sk: core::fmt::Debug,
278{
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 super::*;
301 use crate::ciphersuite::{KeGroup, OprfHash};
302 use crate::{
303 CipherSuite, ClientLogin, ClientLoginFinishParameters, ClientLoginFinishResult,
304 ClientLoginStartResult, ClientRegistration, ClientRegistrationFinishParameters,
305 ClientRegistrationFinishResult, ClientRegistrationStartResult, ServerLogin,
306 ServerLoginParameters, ServerLoginStartResult, ServerRegistration,
307 ServerRegistrationStartResult, ServerSetup,
308 };
309 use hkdf::Hkdf;
310 use rand::rand_core::UnwrapErr;
311 use rand::rngs::SysRng;
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().serialize(), pk.serialize());
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 UnwrapErr(SysRng)));
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 UnwrapErr(SysRng), keypair);
390
391 let ClientRegistrationStartResult {
392 message,
393 state: client,
394 } = ClientRegistration::<Default>::start(&mut UnwrapErr(SysRng), PASSWORD.as_bytes())
395 .unwrap();
396 let ServerRegistrationStartResult { message, .. } =
397 ServerRegistration::start(&server_setup, message, &[]).unwrap();
398 let ClientRegistrationFinishResult { message, .. } = client
399 .finish(
400 &mut UnwrapErr(SysRng),
401 PASSWORD.as_bytes(),
402 message,
403 ClientRegistrationFinishParameters::default(),
404 )
405 .unwrap();
406 let file = ServerRegistration::finish(message);
407
408 let ClientLoginStartResult {
409 message,
410 state: client,
411 } = ClientLogin::<Default>::start(&mut UnwrapErr(SysRng), PASSWORD.as_bytes()).unwrap();
412 let builder = ServerLogin::builder(
413 &mut UnwrapErr(SysRng),
414 &server_setup,
415 Some(file),
416 message,
417 &[],
418 ServerLoginParameters::default(),
419 )
420 .unwrap();
421 let shared_secret = builder.private_key().0.ke_diffie_hellman(builder.data());
422 let ServerLoginStartResult {
423 message,
424 state: server,
425 ..
426 } = builder.build(shared_secret).unwrap();
427 let ClientLoginFinishResult { message, .. } = client
428 .finish(
429 &mut UnwrapErr(SysRng),
430 PASSWORD.as_bytes(),
431 message,
432 ClientLoginFinishParameters::default(),
433 )
434 .unwrap();
435 server
436 .finish(message, ServerLoginParameters::default())
437 .unwrap();
438 }
439
440 #[test]
441 fn remote_seed() {
442 let mut oprf_seed = RemoteSeed::<OprfHash<Default>>(GenericArray::default().into_ha0_4());
443 UnwrapErr(SysRng).fill_bytes(&mut oprf_seed.0);
444
445 let sk = PrivateKey(KeGroup::<Default>::random_sk(&mut UnwrapErr(SysRng)));
446 let pk = sk.public_key();
447 let sk = RemoteKey(sk);
448 let keypair = KeyPair::new(sk, pk);
449
450 let server_setup = ServerSetup::<Default, _, _>::new_with_key_pair_and_seed(
451 &mut UnwrapErr(SysRng),
452 keypair,
453 oprf_seed,
454 );
455
456 let ClientRegistrationStartResult {
457 message,
458 state: client,
459 } = ClientRegistration::<Default>::start(&mut UnwrapErr(SysRng), PASSWORD.as_bytes())
460 .unwrap();
461 let km = server_setup.key_material_info(&[]);
462 let mut ikm = GenericArray::default();
463 Hkdf::<OprfHash<Default>>::from_prk(&km.ikm.0)
464 .unwrap()
465 .expand_multi_info(&km.info, &mut ikm)
466 .unwrap();
467 let ServerRegistrationStartResult { message, .. } =
468 ServerRegistration::start_with_key_material(&server_setup, ikm, message).unwrap();
469 let ClientRegistrationFinishResult { message, .. } = client
470 .finish(
471 &mut UnwrapErr(SysRng),
472 PASSWORD.as_bytes(),
473 message,
474 ClientRegistrationFinishParameters::default(),
475 )
476 .unwrap();
477 let file = ServerRegistration::finish(message);
478
479 let ClientLoginStartResult {
480 message,
481 state: client,
482 } = ClientLogin::<Default>::start(&mut UnwrapErr(SysRng), PASSWORD.as_bytes()).unwrap();
483 let km = server_setup.key_material_info(&[]);
484 let mut ikm = GenericArray::default();
485 Hkdf::<OprfHash<Default>>::from_prk(&km.ikm.0)
486 .unwrap()
487 .expand_multi_info(&km.info, &mut ikm)
488 .unwrap();
489 let builder = ServerLogin::builder_with_key_material(
490 &mut UnwrapErr(SysRng),
491 &server_setup,
492 ikm,
493 Some(file),
494 message,
495 ServerLoginParameters::default(),
496 )
497 .unwrap();
498 let shared_secret = builder.private_key().0.ke_diffie_hellman(builder.data());
499 let ServerLoginStartResult {
500 message,
501 state: server,
502 ..
503 } = builder.build(shared_secret).unwrap();
504 let ClientLoginFinishResult { message, .. } = client
505 .finish(
506 &mut UnwrapErr(SysRng),
507 PASSWORD.as_bytes(),
508 message,
509 ClientLoginFinishParameters::default(),
510 )
511 .unwrap();
512 server
513 .finish(message, ServerLoginParameters::default())
514 .unwrap();
515 }
516}