use alloc::boxed::Box;
use alloc::vec::Vec;
use core::fmt::Debug;
use core::ops::Deref;
use pki_types::FipsStatus;
use zeroize::Zeroize;
use crate::enums::ProtocolVersion;
use crate::error::{Error, PeerMisbehaved};
pub mod ffdhe;
use ffdhe::FfdheGroup;
#[expect(clippy::exhaustive_structs)]
#[derive(Debug)]
pub struct Hybrid {
pub classical: &'static dyn SupportedKxGroup,
pub post_quantum: &'static dyn SupportedKxGroup,
pub name: NamedGroup,
pub layout: HybridLayout,
}
impl SupportedKxGroup for Hybrid {
fn start(&self) -> Result<StartedKeyExchange, Error> {
let classical = self.classical.start()?.into_single();
let post_quantum = self.post_quantum.start()?.into_single();
let combined_pub_key = self
.layout
.concat(post_quantum.pub_key(), classical.pub_key());
Ok(StartedKeyExchange::Hybrid(Box::new(ActiveHybrid {
classical,
post_quantum,
name: self.name,
layout: self.layout,
combined_pub_key,
})))
}
fn start_and_complete(&self, client_share: &[u8]) -> Result<CompletedKeyExchange, Error> {
let (post_quantum_share, classical_share) = self
.layout
.split_received_client_share(client_share)
.ok_or(PeerMisbehaved::InvalidKeyShare)?;
let cl = self
.classical
.start_and_complete(classical_share)?;
let pq = self
.post_quantum
.start_and_complete(post_quantum_share)?;
let combined_pub_key = self
.layout
.concat(&pq.pub_key, &cl.pub_key);
let secret = self
.layout
.concat(pq.secret.secret_bytes(), cl.secret.secret_bytes());
Ok(CompletedKeyExchange {
group: self.name,
pub_key: combined_pub_key,
secret: SharedSecret::from(secret),
})
}
fn name(&self) -> NamedGroup {
self.name
}
fn fips(&self) -> FipsStatus {
match self.layout.post_quantum_first {
true => self.post_quantum.fips(),
false => self.classical.fips(),
}
}
}
struct ActiveHybrid {
classical: Box<dyn ActiveKeyExchange>,
post_quantum: Box<dyn ActiveKeyExchange>,
name: NamedGroup,
layout: HybridLayout,
combined_pub_key: Vec<u8>,
}
impl ActiveKeyExchange for ActiveHybrid {
fn complete(self: Box<Self>, peer_pub_key: &[u8]) -> Result<SharedSecret, Error> {
let (post_quantum_share, classical_share) = self
.layout
.split_received_server_share(peer_pub_key)
.ok_or(PeerMisbehaved::InvalidKeyShare)?;
let cl = self
.classical
.complete(classical_share)?;
let pq = self
.post_quantum
.complete(post_quantum_share)?;
let secret = self
.layout
.concat(pq.secret_bytes(), cl.secret_bytes());
Ok(SharedSecret::from(secret))
}
fn pub_key(&self) -> &[u8] {
&self.combined_pub_key
}
fn group(&self) -> NamedGroup {
self.name
}
}
impl HybridKeyExchange for ActiveHybrid {
fn component(&self) -> (NamedGroup, &[u8]) {
(self.classical.group(), self.classical.pub_key())
}
fn complete_component(self: Box<Self>, peer_pub_key: &[u8]) -> Result<SharedSecret, Error> {
self.classical.complete(peer_pub_key)
}
fn into_key_exchange(self: Box<Self>) -> Box<dyn ActiveKeyExchange> {
self
}
fn as_key_exchange(&self) -> &(dyn ActiveKeyExchange + 'static) {
self
}
}
#[expect(clippy::exhaustive_structs)]
#[derive(Clone, Copy, Debug)]
pub struct HybridLayout {
pub classical_share_len: usize,
pub post_quantum_client_share_len: usize,
pub post_quantum_server_share_len: usize,
pub post_quantum_first: bool,
}
impl HybridLayout {
fn split_received_client_share<'a>(&self, share: &'a [u8]) -> Option<(&'a [u8], &'a [u8])> {
self.split(share, self.post_quantum_client_share_len)
}
fn split_received_server_share<'a>(&self, share: &'a [u8]) -> Option<(&'a [u8], &'a [u8])> {
self.split(share, self.post_quantum_server_share_len)
}
fn split<'a>(
&self,
share: &'a [u8],
post_quantum_share_len: usize,
) -> Option<(&'a [u8], &'a [u8])> {
if share.len() != self.classical_share_len + post_quantum_share_len {
return None;
}
Some(match self.post_quantum_first {
true => {
let (first_share, second_share) = share.split_at(post_quantum_share_len);
(first_share, second_share)
}
false => {
let (first_share, second_share) = share.split_at(self.classical_share_len);
(second_share, first_share)
}
})
}
fn concat(&self, post_quantum: &[u8], classical: &[u8]) -> Vec<u8> {
match self.post_quantum_first {
true => [post_quantum, classical].concat(),
false => [classical, post_quantum].concat(),
}
}
}
pub trait SupportedKxGroup: Send + Sync + Debug {
fn start(&self) -> Result<StartedKeyExchange, Error>;
fn start_and_complete(&self, peer_pub_key: &[u8]) -> Result<CompletedKeyExchange, Error> {
let kx = self.start()?.into_single();
Ok(CompletedKeyExchange {
group: kx.group(),
pub_key: kx.pub_key().to_vec(),
secret: kx.complete(peer_pub_key)?,
})
}
fn ffdhe_group(&self) -> Option<FfdheGroup<'static>> {
None
}
fn name(&self) -> NamedGroup;
fn fips(&self) -> FipsStatus {
FipsStatus::Unvalidated
}
}
#[non_exhaustive]
pub enum StartedKeyExchange {
Single(Box<dyn ActiveKeyExchange>),
Hybrid(Box<dyn HybridKeyExchange>),
}
impl StartedKeyExchange {
pub fn into_single(self) -> Box<dyn ActiveKeyExchange> {
match self {
Self::Single(s) => s,
Self::Hybrid(h) => h.into_key_exchange(),
}
}
pub(crate) fn as_hybrid_checked(
&self,
supported: &[&'static dyn SupportedKxGroup],
version: ProtocolVersion,
) -> Option<(&dyn HybridKeyExchange, &'static dyn SupportedKxGroup)> {
let Self::Hybrid(hybrid) = self else {
return None;
};
let component_group = hybrid.component().0;
if !component_group.usable_for_version(version) {
return None;
}
supported
.iter()
.find(|g| g.name() == component_group)
.copied()
.map(|g| (hybrid.as_ref(), g))
}
}
impl Deref for StartedKeyExchange {
type Target = dyn ActiveKeyExchange;
fn deref(&self) -> &Self::Target {
match self {
Self::Single(s) => s.as_ref(),
Self::Hybrid(h) => h.as_key_exchange(),
}
}
}
pub trait ActiveKeyExchange: Send + Sync {
fn complete(self: Box<Self>, peer_pub_key: &[u8]) -> Result<SharedSecret, Error>;
fn complete_for_tls_version(
self: Box<Self>,
peer_pub_key: &[u8],
tls_version: ProtocolVersion,
) -> Result<SharedSecret, Error> {
if tls_version == ProtocolVersion::TLSv1_3 {
return self.complete(peer_pub_key);
}
let group = self.group();
let mut complete_res = self.complete(peer_pub_key)?;
if group.key_exchange_algorithm() == KeyExchangeAlgorithm::DHE {
complete_res.strip_leading_zeros();
}
Ok(complete_res)
}
fn pub_key(&self) -> &[u8];
fn ffdhe_group(&self) -> Option<FfdheGroup<'static>> {
None
}
fn group(&self) -> NamedGroup;
}
pub trait HybridKeyExchange: ActiveKeyExchange {
fn component(&self) -> (NamedGroup, &[u8]);
fn complete_component(self: Box<Self>, peer_pub_key: &[u8]) -> Result<SharedSecret, Error>;
fn as_key_exchange(&self) -> &(dyn ActiveKeyExchange + 'static);
fn into_key_exchange(self: Box<Self>) -> Box<dyn ActiveKeyExchange>;
}
#[expect(clippy::exhaustive_structs)]
pub struct CompletedKeyExchange {
pub group: NamedGroup,
pub pub_key: Vec<u8>,
pub secret: SharedSecret,
}
enum_builder! {
pub struct NamedGroup(pub u16);
enum NamedGroupName {
secp256r1 => 0x0017,
secp384r1 => 0x0018,
secp521r1 => 0x0019,
X25519 => 0x001d,
X448 => 0x001e,
brainpoolP256r1tls13 => 0x001f,
brainpoolP384r1tls13 => 0x0020,
brainpoolP512r1tls13 => 0x0021,
curveSM2 => 0x0029,
FFDHE2048 => 0x0100,
FFDHE3072 => 0x0101,
FFDHE4096 => 0x0102,
FFDHE6144 => 0x0103,
FFDHE8192 => 0x0104,
MLKEM512 => 0x0200,
MLKEM768 => 0x0201,
MLKEM1024 => 0x0202,
secp256r1MLKEM768 => 0x11eb,
X25519MLKEM768 => 0x11ec,
secp384r1MLKEM1024 => 0x11ed,
}
}
impl NamedGroup {
pub fn key_exchange_algorithm(self) -> KeyExchangeAlgorithm {
match u16::from(self) {
x if (0x100..0x200).contains(&x) => KeyExchangeAlgorithm::DHE,
_ => KeyExchangeAlgorithm::ECDHE,
}
}
pub fn usable_for_version(&self, version: ProtocolVersion) -> bool {
match version {
ProtocolVersion::TLSv1_3 => true,
_ => !matches!(
*self,
Self::MLKEM512
| Self::MLKEM768
| Self::MLKEM1024
| Self::X25519MLKEM768
| Self::secp256r1MLKEM768
| Self::secp384r1MLKEM1024
| Self::brainpoolP256r1tls13
| Self::brainpoolP384r1tls13
| Self::brainpoolP512r1tls13
| Self::curveSM2
),
}
}
}
pub struct SharedSecret {
buf: Vec<u8>,
offset: usize,
}
impl SharedSecret {
pub fn secret_bytes(&self) -> &[u8] {
&self.buf[self.offset..]
}
fn strip_leading_zeros(&mut self) {
let start = self
.secret_bytes()
.iter()
.enumerate()
.find(|(_i, x)| **x != 0)
.map(|(i, _x)| i)
.unwrap_or_else(|| self.secret_bytes().len());
self.offset += start;
}
}
impl Drop for SharedSecret {
#[inline(never)]
fn drop(&mut self) {
self.buf.zeroize();
}
}
impl From<&[u8]> for SharedSecret {
fn from(source: &[u8]) -> Self {
Self {
buf: source.to_vec(),
offset: 0,
}
}
}
impl From<Vec<u8>> for SharedSecret {
fn from(buf: Vec<u8>) -> Self {
Self { buf, offset: 0 }
}
}
#[derive(Clone, Copy, Debug, PartialEq)]
#[non_exhaustive]
pub enum KeyExchangeAlgorithm {
DHE,
ECDHE,
}
#[cfg(test)]
mod tests {
use std::vec;
use super::{NamedGroup, SharedSecret};
use crate::msgs::test_enum16;
#[test]
fn test_shared_secret_strip_leading_zeros() {
let test_cases = [
(vec![0, 1], vec![1]),
(vec![1], vec![1]),
(vec![1, 0, 2], vec![1, 0, 2]),
(vec![0, 0, 1, 2], vec![1, 2]),
(vec![0, 0, 0], vec![]),
(vec![], vec![]),
];
for (buf, expected) in test_cases {
let mut secret = SharedSecret::from(&buf[..]);
assert_eq!(secret.secret_bytes(), buf);
secret.strip_leading_zeros();
assert_eq!(secret.secret_bytes(), expected);
}
}
#[test]
fn test_enums() {
test_enum16::<NamedGroup>(NamedGroup::secp256r1, NamedGroup::FFDHE8192);
}
}