Skip to main content

opaque_ke/
keypair.rs

1// Copyright (c) Meta Platforms, Inc. and affiliates.
2//
3// This source code is dual-licensed under either the MIT license found in the
4// LICENSE-MIT file in the root directory of this source tree or the Apache
5// License, Version 2.0 found in the LICENSE-APACHE file in the root directory
6// of this source tree. You may select, at your option, one of the above-listed
7// licenses.
8
9//! Contains the keypair types that must be supplied for the OPAQUE API
10
11#![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/// A Keypair trait with public-private verification
26#[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// `NonZeroScalar` doesn't implement `Debug`.
37// TODO: remove after `elliptic-curve` bump to v0.14.
38#[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    /// Creates a new [`KeyPair`] from the given keys.
47    pub fn new(sk: SK, pk: PublicKey<G>) -> Self {
48        Self { pk, sk }
49    }
50
51    /// The public key component
52    pub fn public(&self) -> &PublicKey<G> {
53        &self.pk
54    }
55
56    /// The private key component
57    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    /// Generating a random key pair given a cryptographic rng
73    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/// Wrapper around a Key to enforce that it's a private one.
86#[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    /// Returns public key from private key
104    pub fn public_key(&self) -> PublicKey<G> {
105        PublicKey(G::public_key(&self.0))
106    }
107
108    /// Serializes this private key to a fixed-length byte array.
109    pub fn serialize(&self) -> GenericArray<u8, G::SkLen> {
110        G::serialize_sk(&self.0)
111    }
112
113    /// Creates a [`PrivateKey`] from the given bytes.
114    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    /// Diffie-Hellman key exchange implementation
128    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    /// Private-key signing implementation
135    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
149/// A trait to facilitate
150/// [`ServerSetup::de/serialize`](crate::ServerSetup::serialize).
151pub trait PrivateKeySerialization<G: Group>: Clone {
152    /// Custom error type that can be passed down to `ProtocolError::Custom`
153    type Error;
154    /// Serialization size in bytes.
155    type Len: ArrayLength<u8>;
156
157    /// Serialization into bytes
158    fn serialize_key_pair(key_pair: &KeyPair<G, Self>) -> GenericArray<u8, Self::Len>;
159
160    /// Deserialization from bytes
161    ///
162    /// The deserialized bytes must be taken from `bytes`.
163    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/// Wrapper around a Key to enforce that it's a public one.
185#[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    /// Convert from bytes
199    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    /// Convert to bytes
208    pub fn serialize(&self) -> GenericArray<u8, G::PkLen> {
209        G::serialize_pk(&self.0)
210    }
211
212    /// Returns the inner [`Group::Pk`].
213    pub fn to_group_type(&self) -> &G::Pk {
214        &self.0
215    }
216}
217
218impl<G: Group> PublicKey<G> {
219    /// Public-key verifying implementation
220    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/// Default OPRF seed container.
231#[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
239/// A trait to facilitate
240/// [`ServerSetup::de/serialize`](crate::ServerSetup::serialize).
241///
242/// Will be called with `E` being [`PrivateKeySerialization::Error`].
243pub trait OprfSeedSerialization<H, E>: Sized {
244    /// Serialization size in bytes.
245    type Len: ArrayLength<u8>;
246
247    /// Serialization into bytes
248    fn serialize(&self) -> GenericArray<u8, Self::Len>;
249
250    /// Deserialization from bytes
251    ///
252    /// The deserialized bytes must be taken from `bytes`.
253    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//////////////////////////
273// Test Implementations //
274//===================== //
275//////////////////////////
276
277#[cfg(test)]
278impl<G: Group> KeyPair<G> {
279    /// Test-only strategy returning a proptest Strategy based on
280    /// [`Self::derive_random`]
281    fn uniform_keypair_strategy() -> proptest::prelude::BoxedStrategy<Self> {
282        use proptest::prelude::*;
283        use rand::SeedableRng;
284        use rand::rngs::StdRng;
285
286        // The no_shrink is because keypairs should be fixed -- shrinking would cause a
287        // different keypair to be generated, which appears to not be very useful.
288        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}