// SPDX-License-Identifier: MIT OR Apache-2.0 // Copyright (c) VexaHub and contributors. // Copyright (c) Meta Platforms, Inc. and affiliates. //! TripleDH-KEM is a variant of the OPAQUE Triple Diffie-Hellman handshake in //! which the client supplies a KEM public key in KE1 and the server performs a //! KEM encapsulation in KE2 instead of relying solely on the final Diffie- //! Hellman hop. The server bundles the KEM ciphertext alongside the classic //! `TripleDH` payload, both parties absorb the ciphertext into the transcript //! and mix the encapsulated shared secret with the three Diffie-Hellman //! products when deriving handshake keys, and the client decapsulates during //! KE3 to recover that shared secret before validating the server MAC. This //! file contains the data model and trait glue that layer //! the generic `ml-kem` abstractions into the existing OPAQUE key-exchange //! pipeline. use core::fmt::Debug; use core::marker::PhantomData; use core::ops::Add; use derive_where::derive_where; use digest::Output; use digest::block_api::{CoreProxy, SmallBlockSizeUser}; use generic_array::typenum::{Cmp, IsLess, Le, NonZero, Sum, U256}; use generic_array::{ArrayLength, GenericArray}; use hybrid_array::ArraySize; #[allow(deprecated)] use ml_kem::ExpandedKeyEncoding; use ml_kem::kem::{ Ciphertext as MlKemCiphertext, Decapsulate, Encapsulate, Kem as MlKemTrait, KeyExport, KeySizeUser, TryKeyInit, }; use rand::{CryptoRng, Rng}; use subtle::{ConstantTimeEq, CtOption}; use zeroize::{Zeroize, ZeroizeOnDrop}; use super::shared::{self, Ke1Message, Ke1State, NonceLen}; use super::{ Deserialize, GenerateKe1Result, GenerateKe2Result, GenerateKe3Result, KeyExchange, Serialize, SerializedContext, SerializedCredentialRequest, SerializedCredentialResponse, SerializedIdentifiers, }; use crate::ciphersuite::{CipherSuite, KeGroup}; use crate::errors::ProtocolError; use crate::hash::{Hash, OutputSize, ProxyHash}; use crate::key_exchange::group::Group; use crate::keypair::{PrivateKey, PublicKey}; use crate::opaque::Identifiers; use crate::serialization::{ConcatExt, SliceExt}; /// Adapter trait that augments the `ml-kem` core traits with the metadata /// required by OPAQUE (e.g. fixed lengths and serialization hooks). pub trait KemCoreWrapper { /// Public key type used for encapsulation operations. type EncapsulationKey: Clone; /// Secret key type used for decapsulation operations. type DecapsulationKey: Clone + ZeroizeOnDrop; /// Length (in bytes) of the serialized public key. type EncapsulationKeyLen: ArrayLength + ArraySize; /// Length (in bytes) of the serialized secret key. type DecapsulationKeyLen: ArrayLength + ArraySize; /// Length (in bytes) of the encapsulated ciphertext. type CiphertextLen: ArrayLength + ArraySize; /// Length (in bytes) of the shared secret output by the KEM. type SharedSecretLen: ArrayLength + ArraySize; /// Generates a fresh KEM key pair. fn generate( rng: &mut R, ) -> Result<(Self::DecapsulationKey, Self::EncapsulationKey), ProtocolError>; /// Serializes the public encapsulation key. fn serialize_encapsulation_key( key: &Self::EncapsulationKey, ) -> GenericArray; /// Deserializes the public encapsulation key, advancing the input slice. fn deserialize_encapsulation_key( input: &mut &[u8], ) -> Result; /// Serializes the secret decapsulation key. fn serialize_decapsulation_key( key: &Self::DecapsulationKey, ) -> GenericArray; /// Deserializes the secret decapsulation key, advancing the input slice. fn deserialize_decapsulation_key( input: &mut &[u8], ) -> Result; /// Encapsulates to the given public key, returning the ciphertext and /// shared secret. #[allow(clippy::type_complexity)] fn encapsulate( key: &Self::EncapsulationKey, rng: &mut R, ) -> Result< ( GenericArray, GenericArray, ), ProtocolError, >; /// Decapsulates the shared secret from the provided ciphertext. fn decapsulate( key: &Self::DecapsulationKey, encapsulated_key: &GenericArray, ) -> Result, ProtocolError>; } /// Adapter to bridge `rand 0.8` (`rand_core 0.6`) RNGs to `rand_core 0.10` /// which is required by `ml-kem 0.3.x`. struct RngCompat<'a, R>(&'a mut R); impl rand_core::TryRng for RngCompat<'_, R> { type Error = core::convert::Infallible; fn try_next_u32(&mut self) -> Result { Ok(self.0.next_u32()) } fn try_next_u64(&mut self) -> Result { Ok(self.0.next_u64()) } fn try_fill_bytes(&mut self, dst: &mut [u8]) -> Result<(), Self::Error> { self.0.fill_bytes(dst); Ok(()) } } impl rand_core::TryCryptoRng for RngCompat<'_, R> {} type RcEncapsulationKeyLen = <::EncapsulationKey as KeySizeUser>::KeySize; #[allow(deprecated)] type RcDecapsulationKeyLen = <::DecapsulationKey as ExpandedKeyEncoding>::EncodedSize; type RcCiphertextLen = ::CiphertextSize; type RcSharedSecretLen = ::SharedKeySize; #[allow(deprecated)] impl KemCoreWrapper for K where K: MlKemTrait, K::EncapsulationKey: Encapsulate + KeyExport + TryKeyInit + Clone, K::DecapsulationKey: Decapsulate + ExpandedKeyEncoding + Clone + ZeroizeOnDrop, RcEncapsulationKeyLen: ArrayLength + ArraySize, RcDecapsulationKeyLen: ArrayLength + ArraySize, RcCiphertextLen: ArrayLength + ArraySize, RcSharedSecretLen: ArrayLength + ArraySize, { type EncapsulationKey = K::EncapsulationKey; type DecapsulationKey = K::DecapsulationKey; type EncapsulationKeyLen = RcEncapsulationKeyLen; type DecapsulationKeyLen = RcDecapsulationKeyLen; type CiphertextLen = RcCiphertextLen; type SharedSecretLen = RcSharedSecretLen; fn generate( rng: &mut R, ) -> Result<(Self::DecapsulationKey, Self::EncapsulationKey), ProtocolError> { Ok(K::generate_keypair_from_rng(&mut RngCompat(rng))) } fn serialize_encapsulation_key( key: &Self::EncapsulationKey, ) -> GenericArray { GenericArray::from_slice(key.to_bytes().as_slice()).clone() } fn deserialize_encapsulation_key( input: &mut &[u8], ) -> Result { let bytes: GenericArray> = input.take_array("kem encapsulation key")?; let key = ml_kem::array::Array::try_from(bytes.as_slice()) .map_err(|_| ProtocolError::SerializationError)?; TryKeyInit::new(&key).map_err(|_| ProtocolError::SerializationError) } fn serialize_decapsulation_key( key: &Self::DecapsulationKey, ) -> GenericArray { GenericArray::from_slice(key.to_expanded_bytes().as_slice()).clone() } fn deserialize_decapsulation_key( input: &mut &[u8], ) -> Result { let bytes: GenericArray> = input.take_array("kem decapsulation key")?; let key = ml_kem::array::Array::try_from(bytes.as_slice()) .map_err(|_| ProtocolError::SerializationError)?; K::DecapsulationKey::from_expanded_bytes(&key) .map_err(|_| ProtocolError::SerializationError) } fn encapsulate( key: &Self::EncapsulationKey, rng: &mut R, ) -> Result< ( GenericArray, GenericArray, ), ProtocolError, > { let (ciphertext, shared) = key.encapsulate_with_rng(&mut RngCompat(rng)); Ok(( GenericArray::from_slice(ciphertext.as_slice()).clone(), GenericArray::from_slice(shared.as_slice()).clone(), )) } fn decapsulate( key: &Self::DecapsulationKey, encapsulated_key: &GenericArray, ) -> Result, ProtocolError> { let ciphertext = MlKemCiphertext::::try_from(encapsulated_key.as_slice()) .map_err(|_| ProtocolError::SerializationError)?; let shared = key.decapsulate(&ciphertext); Ok(GenericArray::from_slice(shared.as_slice()).clone()) } } /// Triple Diffie-Hellman-style key exchange that offloads the second hop to a /// generic KEM. #[derive(Clone, Debug)] pub struct TripleDhKem(PhantomData<(G, H, K)>); /// Client state combining the classic `TripleDH` state with a KEM secret key. #[cfg_attr( feature = "serde", derive(serde::Deserialize, serde::Serialize), serde(bound( deserialize = "Ke1State: serde::Deserialize<'de>, K::DecapsulationKey: \ serde::Deserialize<'de>", serialize = "Ke1State: serde::Serialize, K::DecapsulationKey: serde::Serialize", )) )] #[derive_where(Clone, ZeroizeOnDrop)] #[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; Ke1State, K::DecapsulationKey)] pub struct KemKe1State { dh_state: Ke1State, kem_decapsulation_key: K::DecapsulationKey, } /// Client message including the ephemeral Diffie-Hellman component alongside a /// serialized KEM public key. #[cfg_attr( feature = "serde", derive(serde::Deserialize, serde::Serialize), serde(bound( deserialize = "Ke1Message: serde::Deserialize<'de>", serialize = "Ke1Message: serde::Serialize", )) )] #[derive_where(Clone, ZeroizeOnDrop)] #[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; Ke1Message)] pub struct KemKe1Message { dh_message: Ke1Message, kem_encapsulation_key: GenericArray, } /// Server state mirrors the `TripleDH` state and carries the client’s KEM /// public key for later use. #[cfg_attr( feature = "serde", derive(serde::Deserialize, serde::Serialize), serde(bound = "") )] #[derive_where(Clone, ZeroizeOnDrop)] #[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd)] pub struct KemKe2State where H::Core: ProxyHash, <::Core as SmallBlockSizeUser>::_BlockSize: IsLess, Le<<::Core as SmallBlockSizeUser>::_BlockSize, U256>: NonZero, <::Core as SmallBlockSizeUser>::_BlockSize: Cmp, OutputSize: ArrayLength, { base_state: super::tripledh::Ke2State, kem_encapsulation_key: GenericArray, server_kem_ciphertext: GenericArray, } /// Server builder placeholder capturing the data needed to finish the KEM /// exchange. #[derive_where(Clone)] pub struct KemKe2Builder where H::Core: ProxyHash, <::Core as SmallBlockSizeUser>::_BlockSize: IsLess, Le<<::Core as SmallBlockSizeUser>::_BlockSize, U256>: NonZero, <::Core as SmallBlockSizeUser>::_BlockSize: Cmp, OutputSize: ArrayLength, { server_nonce: GenericArray, transcript_hasher: H, client_e_pk: PublicKey, server_e_pk: PublicKey, shared_secret_1: GenericArray, shared_secret_3: GenericArray, kem_encapsulation_key: GenericArray, kem_ciphertext: GenericArray, kem_shared_secret: GenericArray, } /// Server message bundles the `TripleDH` payload with the KEM encapsulation. #[cfg_attr( feature = "serde", derive(serde::Deserialize, serde::Serialize), serde(bound( deserialize = "super::tripledh::Ke2Message: serde::Deserialize<'de>", serialize = "super::tripledh::Ke2Message: serde::Serialize", )) )] #[derive_where(Clone, ZeroizeOnDrop)] #[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; super::tripledh::Ke2Message)] pub struct KemKe2Message where H::Core: ProxyHash, <::Core as SmallBlockSizeUser>::_BlockSize: IsLess, Le<<::Core as SmallBlockSizeUser>::_BlockSize, U256>: NonZero, <::Core as SmallBlockSizeUser>::_BlockSize: Cmp, OutputSize: ArrayLength, { dh_message: super::tripledh::Ke2Message, kem_ciphertext: GenericArray, } /// Third message remains the same as `TripleDH`. pub type KemKe3Message = super::tripledh::Ke3Message; impl Drop for KemKe2Builder where G: Group, H: Hash, H::Core: ProxyHash, <::Core as SmallBlockSizeUser>::_BlockSize: IsLess, Le<<::Core as SmallBlockSizeUser>::_BlockSize, U256>: NonZero, <::Core as SmallBlockSizeUser>::_BlockSize: Cmp, OutputSize: ArrayLength, K: KemCoreWrapper, { fn drop(&mut self) { self.server_nonce.zeroize(); digest::Digest::reset(&mut self.transcript_hasher); self.shared_secret_1.zeroize(); self.shared_secret_3.zeroize(); self.kem_shared_secret.zeroize(); self.kem_ciphertext.zeroize(); } } impl ZeroizeOnDrop for KemKe2Builder where G: Group, H: Hash, H::Core: ProxyHash, <::Core as SmallBlockSizeUser>::_BlockSize: IsLess, Le<<::Core as SmallBlockSizeUser>::_BlockSize, U256>: NonZero, <::Core as SmallBlockSizeUser>::_BlockSize: Cmp, OutputSize: ArrayLength, K: KemCoreWrapper, { } impl KeyExchange for TripleDhKem where G: Group + 'static, G::Sk: shared::DiffieHellman, H: Hash, H::Core: ProxyHash, <::Core as SmallBlockSizeUser>::_BlockSize: IsLess, Le<<::Core as SmallBlockSizeUser>::_BlockSize, U256>: NonZero, <::Core as SmallBlockSizeUser>::_BlockSize: Cmp, OutputSize: ArrayLength, K: KemCoreWrapper, NonceLen: Add, Sum: ArrayLength, { type Group = G; type Hash = H; type KE1State = KemKe1State; type KE2State = KemKe2State; type KE1Message = KemKe1Message; type KE2Builder<'a, CS: CipherSuite> = KemKe2Builder; type KE2BuilderData<'a, CS: 'static + CipherSuite> = ( &'a PublicKey, &'a GenericArray, ); type KE2BuilderInput = GenericArray; type KE2Message = KemKe2Message; type KE3Message = KemKe3Message; fn generate_ke1( rng: &mut R, ) -> Result, ProtocolError> { let base = super::tripledh::TripleDh::::generate_ke1(rng)?; let (kem_secret, kem_public) = K::generate(rng)?; let kem_encapsulation_key = K::serialize_encapsulation_key(&kem_public); Ok(GenerateKe1Result { state: KemKe1State { dh_state: base.state, kem_decapsulation_key: kem_secret, }, message: KemKe1Message { dh_message: base.message, kem_encapsulation_key, }, }) } fn ke2_builder<'a, CS: CipherSuite, R: Rng + CryptoRng>( rng: &mut R, credential_request: SerializedCredentialRequest, ke1_message: Self::KE1Message, credential_response: SerializedCredentialResponse, client_s_pk: PublicKey, identifiers: SerializedIdentifiers<'_, KeGroup>, context: SerializedContext<'a>, ) -> Result, ProtocolError> { let shared::Ke2BuilderCommon { server_nonce, transcript_hasher, client_e_pk, server_e_pk, shared_secret_1, shared_secret_3, } = shared::ke2_builder_common::( rng, credential_request, ke1_message.dh_message.clone(), credential_response, client_s_pk, identifiers, context, )?; let mut kem_bytes_slice: &[u8] = ke1_message.kem_encapsulation_key.as_slice(); let encapsulation_key = K::deserialize_encapsulation_key(&mut kem_bytes_slice)?; let (kem_ciphertext, kem_shared_secret) = K::encapsulate(&encapsulation_key, rng)?; let mut transcript_hasher = transcript_hasher; digest::Digest::update( &mut transcript_hasher, ke1_message.kem_encapsulation_key.as_slice(), ); digest::Digest::update(&mut transcript_hasher, kem_ciphertext.as_slice()); Ok(KemKe2Builder { server_nonce, transcript_hasher, client_e_pk, server_e_pk, shared_secret_1, shared_secret_3, kem_encapsulation_key: ke1_message.kem_encapsulation_key.clone(), kem_ciphertext, kem_shared_secret, }) } fn ke2_builder_data<'a, CS: 'static + CipherSuite>( builder: &'a Self::KE2Builder<'_, CS>, ) -> Self::KE2BuilderData<'a, CS> { (&builder.client_e_pk, &builder.kem_encapsulation_key) } fn generate_ke2_input, R: CryptoRng + Rng>( builder: &Self::KE2Builder<'_, CS>, _: &mut R, server_s_sk: &PrivateKey, ) -> Self::KE2BuilderInput { server_s_sk.ke_diffie_hellman(&builder.client_e_pk) } fn build_ke2>( mut builder: Self::KE2Builder<'_, CS>, shared_secret_2: Self::KE2BuilderInput, ) -> Result, ProtocolError> { let transcript_digest = builder.transcript_hasher.clone().finalize(); let derived_keys = shared::derive_keys::( [ builder.shared_secret_1.as_slice(), shared_secret_2.as_slice(), builder.shared_secret_3.as_slice(), builder.kem_shared_secret.as_slice(), ] .into_iter(), &transcript_digest, )?; let (mac, expected_mac) = shared::compute_ke2_macs( &mut builder.transcript_hasher, &derived_keys, &transcript_digest, )?; Ok(GenerateKe2Result { state: KemKe2State { base_state: super::tripledh::Ke2State { session_key: derived_keys.session_key.clone(), expected_mac, }, kem_encapsulation_key: builder.kem_encapsulation_key.clone(), server_kem_ciphertext: builder.kem_ciphertext.clone(), }, message: KemKe2Message { dh_message: super::tripledh::Ke2Message { server_nonce: builder.server_nonce, server_e_pk: builder.server_e_pk.clone(), mac, }, kem_ciphertext: builder.kem_ciphertext.clone(), }, #[cfg(test)] handshake_secret: derived_keys.handshake_secret, #[cfg(test)] km2: derived_keys.km2, }) } fn generate_ke3, R: CryptoRng + Rng>( _rng: &mut R, credential_request: SerializedCredentialRequest, ke1_message: Self::KE1Message, credential_response: SerializedCredentialResponse, ke1_state: &Self::KE1State, ke2_message: Self::KE2Message, server_s_pk: PublicKey, client_s_sk: PrivateKey, identifiers: SerializedIdentifiers<'_, KeGroup>, context: SerializedContext<'_>, ) -> Result, ProtocolError> { let mut transcript_hasher = shared::transcript( &context, &identifiers, &credential_request, &ke1_message.dh_message.to_iter(), &credential_response, ke2_message.dh_message.server_nonce, &ke2_message.dh_message.server_e_pk.serialize(), ); digest::Digest::update( &mut transcript_hasher, ke1_message.kem_encapsulation_key.as_slice(), ); digest::Digest::update( &mut transcript_hasher, ke2_message.kem_ciphertext.as_slice(), ); let shared_secret_1 = ke1_state .dh_state .client_e_sk .ke_diffie_hellman(&ke2_message.dh_message.server_e_pk); let shared_secret_2 = ke1_state .dh_state .client_e_sk .ke_diffie_hellman(&server_s_pk); let shared_secret_3 = client_s_sk.ke_diffie_hellman(&ke2_message.dh_message.server_e_pk); let kem_shared_secret = K::decapsulate( &ke1_state.kem_decapsulation_key, &ke2_message.kem_ciphertext, )?; let (derived_keys, client_mac) = shared::finalize_ke3_transcript( &mut transcript_hasher, [ shared_secret_1.as_slice(), shared_secret_2.as_slice(), shared_secret_3.as_slice(), kem_shared_secret.as_slice(), ] .into_iter(), &ke2_message.dh_message.mac, )?; Ok(GenerateKe3Result { session_key: derived_keys.session_key, message: super::tripledh::Ke3Message { mac: client_mac }, #[cfg(test)] handshake_secret: derived_keys.handshake_secret, #[cfg(test)] km3: derived_keys.km3, }) } fn finish_ke( ke2_state: &Self::KE2State, ke3_message: Self::KE3Message, _identifiers: Identifiers<'_>, _context: SerializedContext<'_>, ) -> Result, ProtocolError> { CtOption::new( ke2_state.base_state.session_key.clone(), ke2_state.base_state.expected_mac.ct_eq(&ke3_message.mac), ) .into_option() .ok_or(ProtocolError::InvalidLoginError) } } /// Serialization logic will be implemented once the concrete KEM wiring is in /// place. impl Deserialize for KemKe1State { fn deserialize_take(input: &mut &[u8]) -> Result { Ok(Self { dh_state: Ke1State::::deserialize_take(input)?, kem_decapsulation_key: K::deserialize_decapsulation_key(input)?, }) } } impl Serialize for KemKe1State where Ke1State: Serialize, as Serialize>::Len: Add, Sum< as Serialize>::Len, K::DecapsulationKeyLen>: ArrayLength, { type Len = Sum< as Serialize>::Len, K::DecapsulationKeyLen>; fn serialize(&self) -> GenericArray { self.dh_state .serialize() .cat(K::serialize_decapsulation_key(&self.kem_decapsulation_key)) } } impl Deserialize for KemKe1Message { fn deserialize_take(input: &mut &[u8]) -> Result { Ok(Self { dh_message: Ke1Message::::deserialize_take(input)?, kem_encapsulation_key: input.take_array("kem encapsulation key")?, }) } } impl Serialize for KemKe1Message where Ke1Message: Serialize, as Serialize>::Len: Add, Sum< as Serialize>::Len, K::EncapsulationKeyLen>: ArrayLength, { type Len = Sum< as Serialize>::Len, K::EncapsulationKeyLen>; fn serialize(&self) -> GenericArray { self.dh_message .serialize() .cat(self.kem_encapsulation_key.clone()) } } impl Deserialize for KemKe2State where H::Core: ProxyHash, <::Core as SmallBlockSizeUser>::_BlockSize: IsLess, Le<<::Core as SmallBlockSizeUser>::_BlockSize, U256>: NonZero, <::Core as SmallBlockSizeUser>::_BlockSize: Cmp, OutputSize: ArrayLength, { fn deserialize_take(input: &mut &[u8]) -> Result { Ok(Self { base_state: super::tripledh::Ke2State::::deserialize_take(input)?, kem_encapsulation_key: input.take_array("kem encapsulation key")?, server_kem_ciphertext: input.take_array("kem ciphertext")?, }) } } impl Serialize for KemKe2State where H::Core: ProxyHash, <::Core as SmallBlockSizeUser>::_BlockSize: IsLess, Le<<::Core as SmallBlockSizeUser>::_BlockSize, U256>: NonZero, <::Core as SmallBlockSizeUser>::_BlockSize: Cmp, OutputSize: ArrayLength, super::tripledh::Ke2State: Serialize, as Serialize>::Len: Add, Sum< as Serialize>::Len, K::EncapsulationKeyLen>: ArrayLength + Add, Sum< Sum< as Serialize>::Len, K::EncapsulationKeyLen>, K::CiphertextLen, >: ArrayLength, { type Len = Sum< Sum< as Serialize>::Len, K::EncapsulationKeyLen>, K::CiphertextLen, >; fn serialize(&self) -> GenericArray { self.base_state .serialize() .cat(self.kem_encapsulation_key.clone()) .cat(self.server_kem_ciphertext.clone()) } } impl Deserialize for KemKe2Message where H::Core: ProxyHash, <::Core as SmallBlockSizeUser>::_BlockSize: IsLess, Le<<::Core as SmallBlockSizeUser>::_BlockSize, U256>: NonZero, <::Core as SmallBlockSizeUser>::_BlockSize: Cmp, OutputSize: ArrayLength, { fn deserialize_take(input: &mut &[u8]) -> Result { Ok(Self { dh_message: super::tripledh::Ke2Message::::deserialize_take(input)?, kem_ciphertext: input.take_array("kem ciphertext")?, }) } } impl Serialize for KemKe2Message where H::Core: ProxyHash, <::Core as SmallBlockSizeUser>::_BlockSize: IsLess, Le<<::Core as SmallBlockSizeUser>::_BlockSize, U256>: NonZero, <::Core as SmallBlockSizeUser>::_BlockSize: Cmp, OutputSize: ArrayLength, NonceLen: Add, Sum: ArrayLength + Add>, Sum, OutputSize>: ArrayLength, super::tripledh::Ke2Message: Serialize, as Serialize>::Len: Add, < as Serialize>::Len as Add>::Output: ArrayLength, { type Len = Sum< as Serialize>::Len, K::CiphertextLen>; fn serialize(&self) -> GenericArray { self.dh_message.serialize().cat(self.kem_ciphertext.clone()) } }