// Copyright (c) Facebook, Inc. and its affiliates. // // This source code is licensed under both the MIT license found in the // LICENSE-MIT file in the root directory of this source tree and the Apache // License, Version 2.0 found in the LICENSE-APACHE file in the root directory // of this source tree. //! An implementation of the Triple Diffie-Hellman key exchange protocol use core::convert::TryFrom; use core::ops::Add; use derive_where::derive_where; use digest::core_api::BlockSizeUser; use digest::{Digest, Output}; use generic_array::sequence::Concat; use generic_array::typenum::{IsLess, Le, NonZero, Sum, Unsigned, U1, U2, U256, U32}; use generic_array::{ArrayLength, GenericArray}; use hkdf::{Hkdf, HkdfExtract}; use hmac::{Hmac, Mac}; use rand::{CryptoRng, RngCore}; use crate::errors::utils::{check_slice_size, check_slice_size_atleast}; use crate::errors::{InternalError, ProtocolError}; use crate::hash::{Hash, OutputSize, ProxyHash}; use crate::key_exchange::group::KeGroup; use crate::key_exchange::traits::{ Deserialize, GenerateKe2Result, GenerateKe3Result, KeyExchange, Serialize, }; use crate::keypair::{KeyPair, PrivateKey, PublicKey, SecretKey}; use crate::serialization::{Input, UpdateExt}; /////////////// // Constants // // ========= // /////////////// pub(crate) type NonceLen = U32; static STR_RFC: &[u8] = b"RFCXXXX"; static STR_CLIENT_MAC: &[u8] = b"ClientMAC"; static STR_HANDSHAKE_SECRET: &[u8] = b"HandshakeSecret"; static STR_SERVER_MAC: &[u8] = b"ServerMAC"; static STR_SESSION_KEY: &[u8] = b"SessionKey"; static STR_OPAQUE: &[u8] = b"OPAQUE-"; //////////////////////////// // High-level API Structs // // ====================== // //////////////////////////// /// The Triple Diffie-Hellman key exchange implementation pub struct TripleDh; /// The client state produced after the first key exchange message #[cfg_attr( feature = "serde", derive(serde::Deserialize, serde::Serialize), serde(bound = "", crate = "serde") )] #[derive_where(Clone, ZeroizeOnDrop)] #[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; KG::Sk)] pub struct Ke1State { client_e_sk: PrivateKey, client_nonce: GenericArray, } /// The first key exchange message #[cfg_attr( feature = "serde", derive(serde::Deserialize, serde::Serialize), serde(bound = "", crate = "serde") )] #[derive_where(Clone, ZeroizeOnDrop)] #[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; KG::Pk)] pub struct Ke1Message { pub(crate) client_nonce: GenericArray, pub(crate) client_e_pk: PublicKey, } /// The server state produced after the second key exchange message #[cfg_attr( feature = "serde", derive(serde::Deserialize, serde::Serialize), serde(bound = "", crate = "serde") )] #[derive_where(Clone, Debug, Eq, Hash, Ord, PartialEq, PartialOrd, ZeroizeOnDrop)] pub struct Ke2State where D::Core: ProxyHash, ::BlockSize: IsLess, Le<::BlockSize, U256>: NonZero, { km3: Output, hashed_transcript: Output, session_key: Output, } /// The second key exchange message #[cfg_attr( feature = "serde", derive(serde::Deserialize, serde::Serialize), serde(bound = "", crate = "serde") )] #[derive_where(Clone, ZeroizeOnDrop)] #[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; KG::Pk)] pub struct Ke2Message where D::Core: ProxyHash, ::BlockSize: IsLess, Le<::BlockSize, U256>: NonZero, { server_nonce: GenericArray, server_e_pk: PublicKey, mac: Output, } /// The third key exchange message #[cfg_attr( feature = "serde", derive(serde::Deserialize, serde::Serialize), serde(bound = "", crate = "serde") )] #[derive_where(Clone, Debug, Eq, Hash, Ord, PartialEq, PartialOrd, ZeroizeOnDrop)] pub struct Ke3Message where D::Core: ProxyHash, ::BlockSize: IsLess, Le<::BlockSize, U256>: NonZero, { mac: Output, } //////////////////////////////// // High-level Implementations // // ========================== // //////////////////////////////// impl KeyExchange for TripleDh where D::Core: ProxyHash, ::BlockSize: IsLess, Le<::BlockSize, U256>: NonZero, // Ke1State: KeSk + Nonce KG::SkLen: Add, Sum: ArrayLength, // Ke1Message: Nonce + KePk NonceLen: Add, Sum: ArrayLength, // Ke2State: (Hash + Hash) + Hash OutputSize: Add>, Sum, OutputSize>: ArrayLength + Add>, Sum, OutputSize>, OutputSize>: ArrayLength, // Ke2Message: (Nonce + KePk) + Hash NonceLen: Add, Sum: ArrayLength + Add>, Sum, OutputSize>: ArrayLength, { type KE1State = Ke1State; type KE2State = Ke2State; type KE1Message = Ke1Message; type KE2Message = Ke2Message; type KE3Message = Ke3Message; fn generate_ke1( rng: &mut R, ) -> Result<(Self::KE1State, Self::KE1Message), ProtocolError> { let client_e_kp = KeyPair::::generate_random(rng); let client_nonce = generate_nonce::(rng); let ke1_message = Ke1Message { client_nonce, client_e_pk: client_e_kp.public().clone(), }; Ok(( Ke1State { client_e_sk: client_e_kp.private().clone(), client_nonce, }, ke1_message, )) } #[allow(clippy::type_complexity)] fn generate_ke2<'a, 'b, 'c, 'd, R: RngCore + CryptoRng, S: SecretKey>( rng: &mut R, serialized_credential_request: impl Iterator, l2_bytes: impl Iterator, ke1_message: Self::KE1Message, client_s_pk: PublicKey, server_s_sk: S, id_u: impl Iterator, id_s: impl Iterator, context: &[u8], ) -> Result, ProtocolError> { let server_e_kp = KeyPair::::generate_random(rng); let server_nonce = generate_nonce::(rng); let mut transcript_hasher = D::new() .chain(STR_RFC) .chain_iter( Input::::from(context) .map_err(ProtocolError::into_custom)? .iter(), ) .chain_iter(id_u.into_iter()) .chain_iter(serialized_credential_request) .chain_iter(id_s.into_iter()) .chain_iter(l2_bytes) .chain(server_nonce) .chain(&server_e_kp.public().serialize()); let result = derive_3dh_keys::( TripleDhComponents { pk1: ke1_message.client_e_pk.clone(), sk1: server_e_kp.private().clone(), pk2: ke1_message.client_e_pk.clone(), sk2: server_s_sk, pk3: client_s_pk, sk3: server_e_kp.private().clone(), }, &transcript_hasher.clone().finalize(), )?; let mut mac_hasher = Hmac::::new_from_slice(&result.1).map_err(|_| InternalError::HmacError)?; mac_hasher.update(&transcript_hasher.clone().finalize()); let mac = mac_hasher.finalize().into_bytes(); Digest::update(&mut transcript_hasher, &mac); Ok(( Ke2State { km3: result.2, hashed_transcript: transcript_hasher.finalize(), session_key: result.0, }, Ke2Message { server_nonce, server_e_pk: server_e_kp.public().clone(), mac, }, #[cfg(test)] result.3, #[cfg(test)] result.1, )) } #[allow(clippy::type_complexity)] fn generate_ke3<'a, 'b, 'c, 'd>( l2_component: impl Iterator, ke2_message: Self::KE2Message, ke1_state: &Self::KE1State, serialized_credential_request: impl Iterator, server_s_pk: PublicKey, client_s_sk: PrivateKey, id_u: impl Iterator, id_s: impl Iterator, context: &[u8], ) -> Result, ProtocolError> { let mut transcript_hasher = D::new() .chain(STR_RFC) .chain_iter(Input::::from(context)?.iter()) .chain_iter(id_u) .chain_iter(serialized_credential_request) .chain_iter(id_s) .chain_iter(l2_component) .chain(ke2_message.to_bytes_without_mac()); let result = derive_3dh_keys::>( TripleDhComponents { pk1: ke2_message.server_e_pk.clone(), sk1: ke1_state.client_e_sk.clone(), pk2: server_s_pk, sk2: ke1_state.client_e_sk.clone(), pk3: ke2_message.server_e_pk.clone(), sk3: client_s_sk, }, &transcript_hasher.clone().finalize(), )?; let mut server_mac = Hmac::::new_from_slice(&result.1).map_err(|_| InternalError::HmacError)?; server_mac.update(&transcript_hasher.clone().finalize()); server_mac .verify(&ke2_message.mac) .map_err(|_| ProtocolError::InvalidLoginError)?; Digest::update(&mut transcript_hasher, &ke2_message.mac); let mut client_mac = Hmac::::new_from_slice(&result.2).map_err(|_| InternalError::HmacError)?; client_mac.update(&transcript_hasher.finalize()); Ok(( result.0, Ke3Message { mac: client_mac.finalize().into_bytes(), }, #[cfg(test)] result.3, #[cfg(test)] result.2, )) } fn finish_ke( ke3_message: Self::KE3Message, ke2_state: &Self::KE2State, ) -> Result, ProtocolError> { let mut client_mac = Hmac::::new_from_slice(&ke2_state.km3).map_err(|_| InternalError::HmacError)?; client_mac.update(&ke2_state.hashed_transcript); client_mac .verify(&ke3_message.mac) .map_err(|_| ProtocolError::InvalidLoginError)?; Ok(ke2_state.session_key.clone()) } } ///////////////////////// // Convenience Structs // //==================== // ///////////////////////// // The triple of public and private components used in the 3DH computation struct TripleDhComponents> { pk1: PublicKey, sk1: PrivateKey, pk2: PublicKey, sk2: S, pk3: PublicKey, sk3: PrivateKey, } // Consists of a session key, followed by two mac keys: (session_key, km2, km3) #[cfg(not(test))] type TripleDhDerivationResult = (Output, Output, Output); #[cfg(test)] type TripleDhDerivationResult = (Output, Output, Output, Output); //////////////////////////////////////////////// // Helper functions and Trait Implementations // // ========================================== // //////////////////////////////////////////////// // Helper functions // Internal function which takes the public and private components of the client // and server keypairs, along with some auxiliary metadata, to produce the // session key and two MAC keys fn derive_3dh_keys>( dh: TripleDhComponents, hashed_derivation_transcript: &[u8], ) -> Result, ProtocolError> where D::Core: ProxyHash, ::BlockSize: IsLess, Le<::BlockSize, U256>: NonZero, { let mut hkdf = HkdfExtract::::new(None); hkdf.input_ikm( &dh.sk1 .diffie_hellman(dh.pk1) .map_err(InternalError::into_custom)?, ); hkdf.input_ikm(&dh.sk2.diffie_hellman(dh.pk2)?); hkdf.input_ikm( &dh.sk3 .diffie_hellman(dh.pk3) .map_err(InternalError::into_custom)?, ); let (_, extracted_ikm) = hkdf.finalize(); let handshake_secret = derive_secrets::( &extracted_ikm, STR_HANDSHAKE_SECRET, hashed_derivation_transcript, ) .map_err(ProtocolError::into_custom)?; let session_key = derive_secrets::( &extracted_ikm, STR_SESSION_KEY, hashed_derivation_transcript, ) .map_err(ProtocolError::into_custom)?; let km2 = hkdf_expand_label::(&handshake_secret, STR_SERVER_MAC, b"") .map_err(ProtocolError::into_custom)?; let km3 = hkdf_expand_label::(&handshake_secret, STR_CLIENT_MAC, b"") .map_err(ProtocolError::into_custom)?; Ok(( GenericArray::clone_from_slice(&session_key), GenericArray::clone_from_slice(&km2), GenericArray::clone_from_slice(&km3), #[cfg(test)] handshake_secret, )) } fn hkdf_expand_label( secret: &[u8], label: &[u8], context: &[u8], ) -> Result, ProtocolError> where D::Core: ProxyHash, ::BlockSize: IsLess, Le<::BlockSize, U256>: NonZero, { let h = Hkdf::::from_prk(secret).map_err(|_| InternalError::HkdfError)?; hkdf_expand_label_extracted(&h, label, context) } fn hkdf_expand_label_extracted( hkdf: &Hkdf, label: &[u8], context: &[u8], ) -> Result, ProtocolError> where D::Core: ProxyHash, ::BlockSize: IsLess, Le<::BlockSize, U256>: NonZero, { let mut okm = GenericArray::default(); let length_u16: u16 = u16::try_from(OutputSize::::USIZE).map_err(|_| ProtocolError::SerializationError)?; let label = Input::::from_label(STR_OPAQUE, label)?; let label = label.to_array_3(); let context = Input::::from(context)?; let context = context.to_array_2(); let hkdf_label = [ &length_u16.to_be_bytes(), label[0], label[1], label[2], context[0], context[1], ]; hkdf.expand_multi_info(&hkdf_label, &mut okm) .map_err(|_| InternalError::HkdfError)?; Ok(okm) } fn derive_secrets( hkdf: &Hkdf, label: &[u8], hashed_derivation_transcript: &[u8], ) -> Result, ProtocolError> where D::Core: ProxyHash, ::BlockSize: IsLess, Le<::BlockSize, U256>: NonZero, { hkdf_expand_label_extracted::(hkdf, label, hashed_derivation_transcript) } // Generate a random nonce up to NonceLen::USIZE bytes. fn generate_nonce(rng: &mut R) -> GenericArray { let mut nonce_bytes = GenericArray::default(); rng.fill_bytes(&mut nonce_bytes); nonce_bytes } // Serialization and deserialization implementations impl Deserialize for Ke1State { fn deserialize(bytes: &[u8]) -> Result { let key_len = KG::SkLen::USIZE; let nonce_len = NonceLen::USIZE; let checked_bytes = check_slice_size_atleast(bytes, key_len + nonce_len, "ke1_state")?; Ok(Self { client_e_sk: PrivateKey::deserialize(&checked_bytes[..key_len])?, client_nonce: GenericArray::clone_from_slice( &checked_bytes[key_len..key_len + nonce_len], ), }) } } impl Serialize for Ke1State where // Ke1State: KeSk + Nonce KG::SkLen: Add, Sum: ArrayLength, { type Len = Sum; fn serialize(&self) -> GenericArray { self.client_e_sk.serialize().concat(self.client_nonce) } } impl Deserialize for Ke1Message { fn deserialize(ke1_message_bytes: &[u8]) -> Result { let nonce_len = NonceLen::USIZE; let checked_nonce = check_slice_size( ke1_message_bytes, nonce_len + ::PkLen::USIZE, "ke1_message nonce", )?; Ok(Self { client_nonce: GenericArray::clone_from_slice(&checked_nonce[..nonce_len]), client_e_pk: PublicKey::deserialize(&checked_nonce[nonce_len..])?, }) } } impl Serialize for Ke1Message where // Ke1Message: Nonce + KePk NonceLen: Add, Sum: ArrayLength, { type Len = Sum; fn serialize(&self) -> GenericArray { self.client_nonce.concat(self.client_e_pk.serialize()) } } impl Deserialize for Ke2State where D::Core: ProxyHash, ::BlockSize: IsLess, Le<::BlockSize, U256>: NonZero, { fn deserialize(input: &[u8]) -> Result { let hash_len = OutputSize::::USIZE; let checked_bytes = check_slice_size(input, 3 * hash_len, "ke2_state")?; Ok(Self { km3: GenericArray::clone_from_slice(&checked_bytes[..hash_len]), hashed_transcript: GenericArray::clone_from_slice( &checked_bytes[hash_len..2 * hash_len], ), session_key: GenericArray::clone_from_slice(&checked_bytes[2 * hash_len..3 * hash_len]), }) } } impl Serialize for Ke2State where D::Core: ProxyHash, ::BlockSize: IsLess, Le<::BlockSize, U256>: NonZero, // Ke2State: (Hash + Hash) + Hash OutputSize: Add>, Sum, OutputSize>: ArrayLength + Add>, Sum, OutputSize>, OutputSize>: ArrayLength, { type Len = Sum, OutputSize>, OutputSize>; fn serialize(&self) -> GenericArray { self.km3 .clone() .concat(self.hashed_transcript.clone()) .concat(self.session_key.clone()) } } impl Deserialize for Ke2Message where D::Core: ProxyHash, ::BlockSize: IsLess, Le<::BlockSize, U256>: NonZero, { fn deserialize(input: &[u8]) -> Result { let key_len = ::PkLen::USIZE; let nonce_len = NonceLen::USIZE; let checked_nonce = check_slice_size_atleast(input, nonce_len, "ke2_message nonce")?; let unchecked_server_e_pk = check_slice_size_atleast( &checked_nonce[nonce_len..], key_len, "ke2_message server_e_pk", )?; let checked_mac = check_slice_size( &unchecked_server_e_pk[key_len..], OutputSize::::USIZE, "ke1_message mac", )?; // Check the public key bytes let server_e_pk = PublicKey::deserialize(&unchecked_server_e_pk[..key_len])?; Ok(Self { server_nonce: GenericArray::clone_from_slice(&checked_nonce[..nonce_len]), server_e_pk, mac: GenericArray::clone_from_slice(checked_mac), }) } } impl Serialize for Ke2Message where D::Core: ProxyHash, ::BlockSize: IsLess, Le<::BlockSize, U256>: NonZero, // Ke2Message: (Nonce + KePk) + Hash NonceLen: Add, Sum: ArrayLength + Add>, Sum, OutputSize>: ArrayLength, { type Len = Sum, OutputSize>; fn serialize(&self) -> GenericArray { self.server_nonce .concat(self.server_e_pk.serialize()) .concat(self.mac.clone()) } } impl Ke2Message where D::Core: ProxyHash, ::BlockSize: IsLess, Le<::BlockSize, U256>: NonZero, NonceLen: Add, Sum: ArrayLength, { fn to_bytes_without_mac(&self) -> GenericArray> { self.server_nonce.concat(self.server_e_pk.serialize()) } } impl Deserialize for Ke3Message where D::Core: ProxyHash, ::BlockSize: IsLess, Le<::BlockSize, U256>: NonZero, { fn deserialize(bytes: &[u8]) -> Result { let checked_bytes = check_slice_size(bytes, OutputSize::::USIZE, "ke3_message")?; Ok(Self { mac: GenericArray::clone_from_slice(checked_bytes), }) } } impl Serialize for Ke3Message where D::Core: ProxyHash, ::BlockSize: IsLess, Le<::BlockSize, U256>: NonZero, { type Len = OutputSize; fn serialize(&self) -> GenericArray { self.mac.clone() } }