// 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::{ ciphersuite::CipherSuite, errors::{ utils::{check_slice_size, check_slice_size_atleast}, InternalPakeError, PakeError, ProtocolError, }, group::Group, hash::Hash, key_exchange::traits::{FromBytes, KeyExchange, ToBytes, ToBytesWithPointers}, keypair::{KeyPair, PrivateKey, PublicKey, SizedBytesExt}, serialization::serialize, }; use digest::{Digest, FixedOutput}; use generic_array::{ typenum::{Unsigned, U32}, ArrayLength, GenericArray, }; use generic_bytes::SizedBytes; use hkdf::Hkdf; use hmac::{Hmac, Mac, NewMac}; use rand::{CryptoRng, RngCore}; use std::convert::TryFrom; use zeroize::Zeroize; const KEY_LEN: usize = 32; 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-"; #[allow(clippy::upper_case_acronyms)] /// The Triple Diffie-Hellman key exchange implementation pub struct TripleDH; impl KeyExchange for TripleDH { type KE1State = Ke1State; type KE2State = Ke2State<::OutputSize>; type KE1Message = Ke1Message; type KE2Message = Ke2Message::OutputSize>; type KE3Message = Ke3Message<::OutputSize>; 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( rng: &mut R, serialized_credential_request: Vec, l2_bytes: Vec, ke1_message: Self::KE1Message, client_s_pk: PublicKey, server_s_sk: PrivateKey, id_u: Vec, id_s: Vec, context: Vec, ) -> Result<(Self::KE2State, Self::KE2Message), 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(&serialize(&context, 2)?) .chain(&id_u) .chain(&serialized_credential_request[..]) .chain(&id_s) .chain(&l2_bytes[..]) .chain(&server_nonce[..]) .chain(&server_e_kp.public().to_arr()); let (session_key, 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, pk3: client_s_pk, sk3: server_e_kp.private().clone(), }, &transcript_hasher.clone().finalize(), )?; let mut mac_hasher = Hmac::::new_from_slice(&km2).map_err(|_| InternalPakeError::HmacError)?; mac_hasher.update(&transcript_hasher.clone().finalize()); let mac = mac_hasher.finalize().into_bytes(); transcript_hasher.update(&mac); Ok(( Ke2State { km3, hashed_transcript: transcript_hasher.finalize(), session_key, }, Ke2Message { server_nonce, server_e_pk: server_e_kp.public().clone(), mac, }, )) } #[allow(clippy::type_complexity)] fn generate_ke3( l2_component: Vec, ke2_message: Self::KE2Message, ke1_state: &Self::KE1State, serialized_credential_request: &[u8], server_s_pk: PublicKey, client_s_sk: PrivateKey, id_u: Vec, id_s: Vec, context: Vec, ) -> Result<(Vec, Self::KE3Message), ProtocolError> { let mut transcript_hasher = D::new() .chain(STR_RFC) .chain(&serialize(&context, 2)?) .chain(&id_u) .chain(&serialized_credential_request) .chain(&id_s) .chain(&l2_component[..]) .chain(&ke2_message.to_bytes_without_info_or_mac()); let (session_key, km2, km3) = 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(&km2).map_err(|_| InternalPakeError::HmacError)?; server_mac.update(&transcript_hasher.clone().finalize()); if server_mac.verify(&ke2_message.mac).is_err() { return Err(ProtocolError::VerificationError( PakeError::KeyExchangeMacValidationError, )); } transcript_hasher.update(ke2_message.mac.to_vec()); let mut client_mac = Hmac::::new_from_slice(&km3).map_err(|_| InternalPakeError::HmacError)?; client_mac.update(&transcript_hasher.finalize()); Ok(( session_key.to_vec(), Ke3Message { mac: client_mac.finalize().into_bytes(), }, )) } #[allow(clippy::type_complexity)] 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(|_| InternalPakeError::HmacError)?; client_mac.update(&ke2_state.hashed_transcript); if client_mac.verify(&ke3_message.mac).is_err() { return Err(ProtocolError::VerificationError( PakeError::KeyExchangeMacValidationError, )); } Ok(ke2_state.session_key.to_vec()) } fn ke2_message_size() -> usize { NonceLen::to_usize() + KEY_LEN + <::OutputSize as Unsigned>::to_usize() } } /// The client state produced after the first key exchange message #[cfg_attr(feature = "serialize", derive(serde::Deserialize, serde::Serialize))] pub struct Ke1State { client_e_sk: PrivateKey, client_nonce: GenericArray, } impl_clone_for!( struct Ke1State, [client_e_sk, client_nonce], ); impl_debug_eq_hash_for!( struct Ke1State, [client_e_sk, client_nonce], ); // This can't be derived because of the use of a generic parameter impl Zeroize for Ke1State { fn zeroize(&mut self) { self.client_e_sk.zeroize(); self.client_nonce.zeroize(); } } impl Drop for Ke1State { fn drop(&mut self) { self.zeroize(); } } /// The first key exchange message #[derive(PartialEq, Eq, Debug, Hash, Clone)] #[cfg_attr(feature = "serialize", derive(serde::Deserialize, serde::Serialize))] pub struct Ke1Message { pub(crate) client_nonce: GenericArray, pub(crate) client_e_pk: PublicKey, } impl FromBytes for Ke1State { fn from_bytes(bytes: &[u8]) -> Result { let nonce_len = NonceLen::to_usize(); let checked_bytes = check_slice_size_atleast(bytes, KEY_LEN + nonce_len, "ke1_state")?; Ok(Self { client_e_sk: PrivateKey::from_bytes(&checked_bytes[..KEY_LEN])?, client_nonce: GenericArray::clone_from_slice( &checked_bytes[KEY_LEN..KEY_LEN + nonce_len], ), }) } } impl ToBytesWithPointers for Ke1State { fn to_bytes(&self) -> Vec { let output: Vec = [&self.client_e_sk.to_arr(), &self.client_nonce[..]].concat(); output } #[cfg(test)] fn as_byte_ptrs(&self) -> Vec<(*const u8, usize)> { vec![ ( self.client_e_sk.as_ptr(), as SizedBytes>::Len::to_usize(), ), (self.client_nonce.as_ptr(), NonceLen::to_usize()), ] } } impl ToBytes for Ke1Message { fn to_bytes(&self) -> Vec { [&self.client_nonce[..], &self.client_e_pk.to_arr()].concat() } } impl FromBytes for Ke1Message { fn from_bytes(ke1_message_bytes: &[u8]) -> Result { let nonce_len = NonceLen::to_usize(); let checked_nonce = check_slice_size(ke1_message_bytes, nonce_len + KEY_LEN, "ke1_message nonce")?; Ok(Self { client_nonce: GenericArray::clone_from_slice(&checked_nonce[..nonce_len]), client_e_pk: PublicKey::from_bytes(&checked_nonce[nonce_len..])?, }) } } /// The server state produced after the second key exchange message #[derive(Clone, Debug, Eq, Hash, PartialEq)] #[cfg_attr(feature = "serialize", derive(serde::Deserialize, serde::Serialize))] #[cfg_attr(feature = "serialize", serde(bound = ""))] pub struct Ke2State> { km3: GenericArray, hashed_transcript: GenericArray, session_key: GenericArray, } // This can't be derived because of the use of a phantom parameter impl> Zeroize for Ke2State { fn zeroize(&mut self) { self.km3.zeroize(); self.hashed_transcript.zeroize(); self.session_key.zeroize(); } } impl> Drop for Ke2State { fn drop(&mut self) { self.zeroize(); } } impl> ToBytesWithPointers for Ke2State { fn to_bytes(&self) -> Vec { [ &self.km3[..], &self.hashed_transcript[..], &self.session_key[..], ] .concat() } #[cfg(test)] fn as_byte_ptrs(&self) -> Vec<(*const u8, usize)> { vec![ (self.km3.as_ptr(), HashLen::to_usize()), (self.hashed_transcript.as_ptr(), HashLen::to_usize()), (self.session_key.as_ptr(), HashLen::to_usize()), ] } } /// The second key exchange message #[derive(Clone, Debug, Eq, Hash, PartialEq)] #[cfg_attr(feature = "serialize", derive(serde::Deserialize, serde::Serialize))] #[cfg_attr(feature = "serialize", serde(bound = ""))] pub struct Ke2Message> { server_nonce: GenericArray, server_e_pk: PublicKey, mac: GenericArray, } impl> FromBytes for Ke2State { fn from_bytes(input: &[u8]) -> Result { let hash_len = HashLen::to_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> ToBytes for Ke2Message { fn to_bytes(&self) -> Vec { [&self.to_bytes_without_info_or_mac(), &self.mac[..]].concat() } } impl> Ke2Message { fn to_bytes_without_info_or_mac(&self) -> Vec { [&self.server_nonce[..], &self.server_e_pk.to_arr()].concat() } } impl> FromBytes for Ke2Message { fn from_bytes(input: &[u8]) -> Result { let nonce_len = NonceLen::to_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..], HashLen::to_usize(), "ke1_message mac", )?; // Check the public key bytes let server_e_pk = KeyPair::::check_public_key(PublicKey::from_bytes( &unchecked_server_e_pk[..KEY_LEN], )?)?; Ok(Self { server_nonce: GenericArray::clone_from_slice(&checked_nonce[..nonce_len]), server_e_pk: PublicKey::from_bytes(&server_e_pk)?, mac: GenericArray::clone_from_slice(checked_mac), }) } } #[allow(clippy::upper_case_acronyms)] // The triple of public and private components used in the 3DH computation struct TripleDHComponents { pk1: PublicKey, sk1: PrivateKey, pk2: PublicKey, sk2: PrivateKey, pk3: PublicKey, sk3: PrivateKey, } #[allow(clippy::upper_case_acronyms)] // Consists of a session key, followed by two mac keys: (session_key, km2, km3) type TripleDHDerivationResult = ( GenericArray::OutputSize>, GenericArray::OutputSize>, GenericArray::OutputSize>, ); /// The third key exchange message #[derive(Clone, Debug, Eq, Hash, PartialEq)] #[cfg_attr(feature = "serialize", derive(serde::Deserialize, serde::Serialize))] #[cfg_attr(feature = "serialize", serde(bound = ""))] pub struct Ke3Message> { mac: GenericArray, } impl> ToBytes for Ke3Message { fn to_bytes(&self) -> Vec { self.mac.to_vec() } } impl> FromBytes for Ke3Message { fn from_bytes(bytes: &[u8]) -> Result { let checked_bytes = check_slice_size(bytes, HashLen::to_usize(), "ke3_message")?; Ok(Self { mac: GenericArray::clone_from_slice(checked_bytes), }) } } // 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> { let ikm: Vec = [ &KeyPair::::diffie_hellman(dh.pk1, dh.sk1)?[..], &KeyPair::::diffie_hellman(dh.pk2, dh.sk2)?[..], &KeyPair::::diffie_hellman(dh.pk3, dh.sk3)?[..], ] .concat(); let extracted_ikm = Hkdf::::new(None, &ikm); let handshake_secret = derive_secrets::( &extracted_ikm, STR_HANDSHAKE_SECRET, hashed_derivation_transcript, )?; let session_key = derive_secrets::( &extracted_ikm, STR_SESSION_KEY, hashed_derivation_transcript, )?; let km2 = hkdf_expand_label::( &handshake_secret, STR_SERVER_MAC, b"", ::OutputSize::to_usize(), )?; let km3 = hkdf_expand_label::( &handshake_secret, STR_CLIENT_MAC, b"", ::OutputSize::to_usize(), )?; Ok(( GenericArray::clone_from_slice(&session_key), GenericArray::clone_from_slice(&km2), GenericArray::clone_from_slice(&km3), )) } fn hkdf_expand_label( secret: &[u8], label: &[u8], context: &[u8], length: usize, ) -> Result, ProtocolError> { let h = Hkdf::::from_prk(secret).map_err(|_| InternalPakeError::HkdfError)?; hkdf_expand_label_extracted(&h, label, context, length) } fn hkdf_expand_label_extracted( hkdf: &Hkdf, label: &[u8], context: &[u8], length: usize, ) -> Result, ProtocolError> { let mut okm = vec![0u8; length]; let mut hkdf_label: Vec = Vec::new(); let length_u16: u16 = u16::try_from(length).map_err(|_| PakeError::SerializationError)?; hkdf_label.extend_from_slice(&length_u16.to_be_bytes()); let mut opaque_label: Vec = Vec::new(); opaque_label.extend_from_slice(STR_OPAQUE); opaque_label.extend_from_slice(label); hkdf_label.extend_from_slice(&serialize(&opaque_label, 1)?); hkdf_label.extend_from_slice(&serialize(context, 1)?); hkdf.expand(&hkdf_label, &mut okm) .map_err(|_| InternalPakeError::HkdfError)?; Ok(okm) } fn derive_secrets( hkdf: &Hkdf, label: &[u8], hashed_derivation_transcript: &[u8], ) -> Result, ProtocolError> { hkdf_expand_label_extracted::( hkdf, label, hashed_derivation_transcript, ::OutputSize::to_usize(), ) } // Generate a random nonce up to NonceLen::to_usize() bytes. fn generate_nonce(rng: &mut R) -> GenericArray { let mut nonce_bytes = vec![0u8; NonceLen::to_usize()]; rng.fill_bytes(&mut nonce_bytes); GenericArray::clone_from_slice(&nonce_bytes) }