// Copyright (c) Facebook, Inc. and its affiliates. // // This source code is licensed under the MIT license found in the // LICENSE file in the root directory of this source tree. //! An implementation of the Triple Diffie-Hellman key exchange protocol use crate::{ errors::{utils::check_slice_size, InternalPakeError, PakeError, ProtocolError}, hash::Hash, key_exchange::traits::{KeyExchange, ToBytes}, keypair::{KeyPair, SizedBytes}, }; use digest::{Digest, FixedOutput}; use generic_array::{ typenum::{Unsigned, U32}, ArrayLength, GenericArray, }; use hkdf::Hkdf; use hmac::{Hmac, Mac, NewMac}; use rand_core::{CryptoRng, RngCore}; use std::convert::TryFrom; const KEY_LEN: usize = 32; pub(crate) const NONCE_LEN: usize = 32; pub(crate) type NonceLen = U32; const KE1_STATE_LEN: usize = KEY_LEN + KEY_LEN + NONCE_LEN; static STR_3DH: &[u8] = b"3DH keys"; /// The Triple Diffie-Hellman key exchange implementation pub struct TripleDH; impl KeyExchange for TripleDH { type KE1State = KE1State<::OutputSize, KeyFormat>; type KE2State = KE2State<::OutputSize>; type KE1Message = KE1Message; type KE2Message = KE2Message<::OutputSize, KeyFormat>; type KE3Message = KE3Message<::OutputSize>; fn generate_ke1( l1_component: Vec, rng: &mut R, ) -> Result<(Self::KE1State, Self::KE1Message), ProtocolError> { let client_e_kp = KeyFormat::generate_random(rng)?; let client_nonce: GenericArray = { let mut client_nonce_bytes = [0u8; NONCE_LEN]; rng.fill_bytes(&mut client_nonce_bytes); client_nonce_bytes.into() }; let ke1_message = KE1Message { client_nonce, client_e_pk: client_e_kp.public().clone(), }; let l1_data: Vec = [&l1_component[..], &ke1_message.to_bytes()].concat(); let mut hasher = D::new(); hasher.update(&l1_data); let hashed_l1 = hasher.finalize(); Ok(( KE1State { client_e_sk: client_e_kp.private().clone(), client_nonce, hashed_l1, }, ke1_message, )) } fn generate_ke2( rng: &mut R, l1_bytes: Vec, l2_bytes: Vec, ke1_message: Self::KE1Message, client_s_pk: KeyFormat::Repr, server_s_sk: KeyFormat::Repr, ) -> Result<(Self::KE2State, Self::KE2Message), ProtocolError> { let server_e_kp = KeyFormat::generate_random(rng)?; let server_nonce: GenericArray = { let mut server_nonce_bytes = [0u8; NONCE_LEN]; rng.fill_bytes(&mut server_nonce_bytes); server_nonce_bytes.into() }; let (shared_secret, km2, km3) = derive_3dh_keys::( TripleDHComponents { pk1: ke1_message.client_e_pk.clone(), sk1: server_e_kp.private().clone(), pk2: ke1_message.client_e_pk, sk2: server_s_sk.clone(), pk3: client_s_pk.clone(), sk3: server_e_kp.private().clone(), }, &ke1_message.client_nonce, &server_nonce, client_s_pk, KeyFormat::public_from_private(&server_s_sk), )?; let mut hasher = D::new(); hasher.update(&l1_bytes); let hashed_l1 = hasher.finalize(); let transcript2: Vec = [ &hashed_l1[..], &l2_bytes[..], &server_nonce[..], &server_e_kp.public().to_arr(), ] .concat(); let mut hasher2 = D::new(); hasher2.update(&transcript2); let hashed_transcript = hasher2.finalize(); let mut mac = Hmac::::new_varkey(&km2).map_err(|_| InternalPakeError::HmacError)?; mac.update(&hashed_transcript); Ok(( KE2State { km3, hashed_transcript, shared_secret, }, KE2Message { server_nonce, server_e_pk: server_e_kp.public().clone(), mac: mac.finalize().into_bytes(), }, )) } fn generate_ke3( l2_component: Vec, ke2_message: Self::KE2Message, ke1_state: &Self::KE1State, server_s_pk: KeyFormat::Repr, client_s_sk: KeyFormat::Repr, ) -> Result<(Vec, Self::KE3Message), ProtocolError> { let (shared_secret, km2, km3) = derive_3dh_keys::( TripleDHComponents { pk1: ke2_message.server_e_pk.clone(), sk1: ke1_state.client_e_sk.clone(), pk2: server_s_pk.clone(), sk2: ke1_state.client_e_sk.clone(), pk3: ke2_message.server_e_pk.clone(), sk3: client_s_sk.clone(), }, &ke1_state.client_nonce, &ke2_message.server_nonce, KeyFormat::public_from_private(&client_s_sk), server_s_pk, )?; let transcript: Vec = [ &ke1_state.hashed_l1[..], &l2_component[..], &ke2_message.server_nonce[..], &ke2_message.server_e_pk.to_arr(), ] .concat(); let mut hasher = D::new(); hasher.update(&transcript); let hashed_transcript = hasher.finalize(); let mut server_mac = Hmac::::new_varkey(&km2).map_err(|_| InternalPakeError::HmacError)?; server_mac.update(&hashed_transcript); if ke2_message.mac != server_mac.finalize().into_bytes() { return Err(ProtocolError::VerificationError( PakeError::KeyExchangeMacValidationError, )); } let mut client_mac = Hmac::::new_varkey(&km3).map_err(|_| InternalPakeError::HmacError)?; client_mac.update(&hashed_transcript); Ok(( shared_secret.to_vec(), KE3Message { mac: client_mac.finalize().into_bytes(), }, )) } fn finish_ke( ke3_message: Self::KE3Message, ke2_state: &Self::KE2State, ) -> Result, ProtocolError> { let mut client_mac = Hmac::::new_varkey(&ke2_state.km3).map_err(|_| InternalPakeError::HmacError)?; client_mac.update(&ke2_state.hashed_transcript); if ke3_message.mac != client_mac.finalize().into_bytes() { return Err(ProtocolError::VerificationError( PakeError::KeyExchangeMacValidationError, )); } Ok(ke2_state.shared_secret.to_vec()) } fn ke1_state_size() -> usize { KE1_STATE_LEN } fn ke2_message_size() -> usize { NONCE_LEN + KEY_LEN + <::OutputSize as Unsigned>::to_usize() } } /// The client state produced after the first key exchange message #[derive(PartialEq, Eq)] pub struct KE1State, KeyFormat: KeyPair> { client_e_sk: KeyFormat::Repr, client_nonce: GenericArray, hashed_l1: GenericArray, } /// The first key exchange message #[derive(PartialEq, Eq)] pub struct KE1Message { pub(crate) client_nonce: GenericArray, pub(crate) client_e_pk: KeyFormat::Repr, } impl, KeyFormat: KeyPair> TryFrom<&[u8]> for KE1State { type Error = InternalPakeError; fn try_from(bytes: &[u8]) -> Result { let checked_bytes = check_slice_size( bytes, KEY_LEN + NONCE_LEN + HashLen::to_usize(), "ke1_state", )?; Ok(Self { client_e_sk: KeyFormat::Repr::from_bytes(&checked_bytes[..KEY_LEN])?, client_nonce: GenericArray::clone_from_slice( &checked_bytes[KEY_LEN..KEY_LEN + NONCE_LEN], ), hashed_l1: GenericArray::clone_from_slice(&checked_bytes[KEY_LEN + NONCE_LEN..]), }) } } impl, KeyFormat: KeyPair> ToBytes for KE1State { fn to_bytes(&self) -> Vec { let output: Vec = [ &self.client_e_sk.to_arr(), &self.client_nonce[..], &self.hashed_l1[..], ] .concat(); output } } impl ToBytes for KE1Message { fn to_bytes(&self) -> Vec { [&self.client_nonce[..], &self.client_e_pk.to_arr()].concat() } } impl TryFrom<&[u8]> for KE1Message { type Error = InternalPakeError; fn try_from(ke1_message_bytes: &[u8]) -> Result { let checked_bytes = check_slice_size(ke1_message_bytes, NONCE_LEN + KEY_LEN, "ke1_message")?; Ok(Self { client_nonce: GenericArray::clone_from_slice(&checked_bytes[..NONCE_LEN]), client_e_pk: KeyFormat::Repr::from_bytes(&checked_bytes[NONCE_LEN..])?, }) } } /// The server state produced after the second key exchange message pub struct KE2State> { km3: GenericArray, hashed_transcript: GenericArray, shared_secret: GenericArray, } /// The second key exchange message pub struct KE2Message, KeyFormat: KeyPair> { server_nonce: GenericArray, server_e_pk: KeyFormat::Repr, mac: GenericArray, } impl> ToBytes for KE2State { fn to_bytes(&self) -> Vec { let output: Vec = [ &self.km3[..], &self.hashed_transcript[..], &self.shared_secret[..], ] .concat(); output } } impl> TryFrom<&[u8]> for KE2State { type Error = InternalPakeError; fn try_from(ke1_message_bytes: &[u8]) -> Result { let checked_bytes = check_slice_size(ke1_message_bytes, 3 * KEY_LEN, "ke2_state")?; Ok(Self { km3: GenericArray::clone_from_slice(&checked_bytes[..KEY_LEN]), hashed_transcript: GenericArray::clone_from_slice(&checked_bytes[KEY_LEN..2 * KEY_LEN]), shared_secret: GenericArray::clone_from_slice(&checked_bytes[2 * KEY_LEN..]), }) } } impl, KeyFormat: KeyPair> ToBytes for KE2Message { fn to_bytes(&self) -> Vec { let output: Vec = [ &self.server_nonce[..], &self.server_e_pk.to_arr(), &self.mac[..], ] .concat(); output } } impl, KeyFormat: KeyPair> TryFrom<&[u8]> for KE2Message { type Error = InternalPakeError; fn try_from(ke2_message_bytes: &[u8]) -> Result { let ke2_message_len = NONCE_LEN + KEY_LEN + HashLen::to_usize(); let checked_bytes = check_slice_size(ke2_message_bytes, ke2_message_len, "ke2_message")?; Ok(Self { server_nonce: GenericArray::clone_from_slice(&checked_bytes[..NONCE_LEN]), server_e_pk: KeyFormat::Repr::from_bytes( &checked_bytes[NONCE_LEN..NONCE_LEN + KEY_LEN], )?, mac: GenericArray::clone_from_slice(&checked_bytes[NONCE_LEN + KEY_LEN..]), }) } } // The triple of public and private components used in the 3DH computation struct TripleDHComponents { pk1: KeyFormat::Repr, sk1: KeyFormat::Repr, pk2: KeyFormat::Repr, sk2: KeyFormat::Repr, pk3: KeyFormat::Repr, sk3: KeyFormat::Repr, } // Consists of a shared secret, followed by two mac keys type TripleDHDerivationResult = ( GenericArray::OutputSize>, GenericArray::OutputSize>, GenericArray::OutputSize>, ); // Internal function which takes the public and private components of the client and server keypairs, along // with some auxiliary metadata, to produce the shared secret and two MAC keys fn derive_3dh_keys( dh: TripleDHComponents, client_nonce: &GenericArray, server_nonce: &GenericArray, client_s_pk: KeyFormat::Repr, server_s_pk: KeyFormat::Repr, ) -> Result, ProtocolError> { let ikm: Vec = [ &KeyFormat::diffie_hellman(dh.pk1, dh.sk1)[..], &KeyFormat::diffie_hellman(dh.pk2, dh.sk2)[..], &KeyFormat::diffie_hellman(dh.pk3, dh.sk3)[..], ] .concat(); let info: Vec = [ STR_3DH, &client_nonce, &server_nonce, &client_s_pk.to_arr(), &server_s_pk.to_arr(), ] .concat(); const OUTPUT_SIZE: usize = 32; let mut okm = [0u8; 3 * OUTPUT_SIZE]; let h = Hkdf::::new(None, &ikm); h.expand(&info, &mut okm) .map_err(|_| InternalPakeError::HkdfError)?; Ok(( GenericArray::clone_from_slice(&okm[..OUTPUT_SIZE]), GenericArray::clone_from_slice(&okm[OUTPUT_SIZE..2 * OUTPUT_SIZE]), GenericArray::clone_from_slice(&okm[2 * OUTPUT_SIZE..]), )) } /// The third key exchange message pub struct KE3Message> { mac: GenericArray, } impl> ToBytes for KE3Message { fn to_bytes(&self) -> Vec { self.mac.to_vec() } } impl> TryFrom<&[u8]> for KE3Message { type Error = InternalPakeError; fn try_from(bytes: &[u8]) -> Result { let checked_bytes = check_slice_size(bytes, KEY_LEN, "ke3_message")?; Ok(Self { mac: GenericArray::clone_from_slice(&checked_bytes), }) } }