Skip to main content

opaque_vx/
keypair.rs

1// SPDX-License-Identifier: MIT OR Apache-2.0
2// Copyright (c) VexaHub and contributors.
3// Copyright (c) Meta Platforms, Inc. and affiliates.
4
5//! Contains the keypair types that must be supplied for the OPAQUE API
6
7#![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/// A Keypair trait with public-private verification
22#[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    /// Creates a new [`KeyPair`] from the given keys.
39    pub fn new(sk: SK, pk: PublicKey<G>) -> Self {
40        Self { pk, sk }
41    }
42
43    /// The public key component
44    pub fn public(&self) -> &PublicKey<G> {
45        &self.pk
46    }
47
48    /// The private key component
49    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    /// Generating a random key pair given a cryptographic rng
65    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/// Wrapper around a Key to enforce that it's a private one.
78#[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    /// Returns public key from private key
96    pub fn public_key(&self) -> PublicKey<G> {
97        PublicKey(G::public_key(&self.0))
98    }
99
100    /// Serializes this private key to a fixed-length byte array.
101    pub fn serialize(&self) -> GenericArray<u8, G::SkLen> {
102        G::serialize_sk(&self.0)
103    }
104
105    /// Creates a [`PrivateKey`] from the given bytes.
106    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    /// Diffie-Hellman key exchange implementation
120    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    /// Private-key signing implementation
127    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
141/// A trait to facilitate
142/// [`ServerSetup::de/serialize`](crate::ServerSetup::serialize).
143pub trait PrivateKeySerialization<G: Group>: Clone {
144    /// Custom error type that can be passed down to `ProtocolError::Custom`
145    type Error;
146    /// Serialization size in bytes.
147    type Len: ArrayLength;
148
149    /// Serialization into bytes
150    fn serialize_key_pair(key_pair: &KeyPair<G, Self>) -> GenericArray<u8, Self::Len>;
151
152    /// Deserialization from bytes
153    ///
154    /// The deserialized bytes must be taken from `bytes`.
155    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/// Wrapper around a Key to enforce that it's a public one.
177#[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    /// Convert from bytes
191    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    /// Convert to bytes
200    pub fn serialize(&self) -> GenericArray<u8, G::PkLen> {
201        G::serialize_pk(&self.0)
202    }
203
204    /// Returns the inner [`Group::Pk`].
205    pub fn to_group_type(&self) -> &G::Pk {
206        &self.0
207    }
208}
209
210impl<G: Group> PublicKey<G> {
211    /// Public-key verifying implementation
212    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/// Default OPRF seed container.
223#[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
231/// A trait to facilitate
232/// [`ServerSetup::de/serialize`](crate::ServerSetup::serialize).
233///
234/// Will be called with `E` being [`PrivateKeySerialization::Error`].
235pub trait OprfSeedSerialization<H, E>: Sized {
236    /// Serialization size in bytes.
237    type Len: ArrayLength;
238
239    /// Serialization into bytes
240    fn serialize(&self) -> GenericArray<u8, Self::Len>;
241
242    /// Deserialization from bytes
243    ///
244    /// The deserialized bytes must be taken from `bytes`.
245    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//////////////////////////
269// Test Implementations //
270//===================== //
271//////////////////////////
272
273#[cfg(test)]
274impl<G: Group> KeyPair<G>
275where
276    G::Pk: core::fmt::Debug,
277    G::Sk: core::fmt::Debug,
278{
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 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}