use crate::{QsshError, Result};
use crate::crypto::mlkem::{MlKem768KeyPair, mlkem768_encapsulate, mlkem768};
use x25519_dalek::{PublicKey, StaticSecret};
use rand::RngCore;
use sha3::{Sha3_256, Digest};
use zeroize::{Zeroize, ZeroizeOnDrop};
pub const X25519_PUBLIC_KEY_SIZE: usize = 32;
pub const X25519_SECRET_KEY_SIZE: usize = 32;
#[derive(Clone, Zeroize, ZeroizeOnDrop)]
pub struct HybridKeyPair {
x25519_secret: [u8; X25519_SECRET_KEY_SIZE],
#[zeroize(skip)]
x25519_public: [u8; X25519_PUBLIC_KEY_SIZE],
mlkem: MlKem768KeyPair,
}
impl HybridKeyPair {
pub fn generate() -> Result<Self> {
let mut rng = rand::thread_rng();
let mut x25519_secret = [0u8; X25519_SECRET_KEY_SIZE];
rng.fill_bytes(&mut x25519_secret);
let x25519_static = StaticSecret::from(x25519_secret);
let x25519_public = *PublicKey::from(&x25519_static).as_bytes();
let mlkem = MlKem768KeyPair::generate()?;
Ok(Self {
x25519_secret,
x25519_public,
mlkem,
})
}
pub fn x25519_public_key(&self) -> &[u8; X25519_PUBLIC_KEY_SIZE] {
&self.x25519_public
}
pub fn mlkem_encapsulation_key(&self) -> &[u8] {
self.mlkem.encapsulation_key()
}
pub fn process_response(
&self,
client_x25519_public: &[u8],
mlkem_ciphertext: &[u8],
) -> Result<Vec<u8>> {
if client_x25519_public.len() != X25519_PUBLIC_KEY_SIZE {
return Err(QsshError::Crypto(format!(
"Invalid X25519 public key size: expected {}, got {}",
X25519_PUBLIC_KEY_SIZE,
client_x25519_public.len()
)));
}
let x25519_static = StaticSecret::from(self.x25519_secret);
let client_public: [u8; 32] = client_x25519_public.try_into()
.map_err(|_| QsshError::Crypto("X25519 public key conversion failed".into()))?;
let client_public = PublicKey::from(client_public);
let x25519_shared = x25519_static.diffie_hellman(&client_public);
let mlkem_shared = self.mlkem.decapsulate(mlkem_ciphertext)?;
Ok(combine_secrets(
x25519_shared.as_bytes(),
&mlkem_shared,
))
}
}
pub struct HybridClientExchange {
x25519_secret: [u8; X25519_SECRET_KEY_SIZE],
x25519_public: [u8; X25519_PUBLIC_KEY_SIZE],
}
impl HybridClientExchange {
pub fn new() -> Self {
let mut rng = rand::thread_rng();
let mut x25519_secret = [0u8; X25519_SECRET_KEY_SIZE];
rng.fill_bytes(&mut x25519_secret);
let x25519_static = StaticSecret::from(x25519_secret);
let x25519_public = *PublicKey::from(&x25519_static).as_bytes();
Self {
x25519_secret,
x25519_public,
}
}
pub fn x25519_public_key(&self) -> &[u8; X25519_PUBLIC_KEY_SIZE] {
&self.x25519_public
}
pub fn complete(
&self,
server_x25519_public: &[u8],
server_mlkem_ek: &[u8],
) -> Result<(Vec<u8>, Vec<u8>)> {
if server_x25519_public.len() != X25519_PUBLIC_KEY_SIZE {
return Err(QsshError::Crypto(format!(
"Invalid server X25519 public key size: expected {}, got {}",
X25519_PUBLIC_KEY_SIZE,
server_x25519_public.len()
)));
}
let x25519_static = StaticSecret::from(self.x25519_secret);
let server_public: [u8; 32] = server_x25519_public.try_into()
.map_err(|_| QsshError::Crypto("X25519 public key conversion failed".into()))?;
let server_public = PublicKey::from(server_public);
let x25519_shared = x25519_static.diffie_hellman(&server_public);
let (mlkem_shared, mlkem_ciphertext) = mlkem768_encapsulate(server_mlkem_ek)?;
let combined = combine_secrets(x25519_shared.as_bytes(), &mlkem_shared);
Ok((combined, mlkem_ciphertext))
}
}
impl Drop for HybridClientExchange {
fn drop(&mut self) {
self.x25519_secret.zeroize();
}
}
impl Default for HybridClientExchange {
fn default() -> Self {
Self::new()
}
}
pub fn combine_secrets(x25519_shared: &[u8], mlkem_shared: &[u8]) -> Vec<u8> {
let mut hasher = Sha3_256::new();
hasher.update(b"QSSH-HYBRID-X25519-MLKEM768-v1");
hasher.update(x25519_shared);
hasher.update(mlkem_shared);
hasher.finalize().to_vec()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_hybrid_key_exchange() {
let server = HybridKeyPair::generate().unwrap();
let client = HybridClientExchange::new();
let (client_shared, mlkem_ct) = client.complete(
server.x25519_public_key(),
server.mlkem_encapsulation_key(),
).unwrap();
let server_shared = server.process_response(
client.x25519_public_key(),
&mlkem_ct,
).unwrap();
assert_eq!(client_shared, server_shared);
assert_eq!(client_shared.len(), 32);
}
#[test]
fn test_hybrid_key_sizes() {
let keypair = HybridKeyPair::generate().unwrap();
assert_eq!(keypair.x25519_public_key().len(), X25519_PUBLIC_KEY_SIZE);
assert_eq!(keypair.mlkem_encapsulation_key().len(), mlkem768::EK_SIZE);
}
#[test]
fn test_combine_secrets_deterministic() {
let x25519 = [0x11u8; 32];
let mlkem = [0x22u8; 32];
let combined1 = combine_secrets(&x25519, &mlkem);
let combined2 = combine_secrets(&x25519, &mlkem);
assert_eq!(combined1, combined2);
}
#[test]
fn test_combine_secrets_different_inputs() {
let x25519 = [0x11u8; 32];
let mlkem1 = [0x22u8; 32];
let mlkem2 = [0x33u8; 32];
let combined1 = combine_secrets(&x25519, &mlkem1);
let combined2 = combine_secrets(&x25519, &mlkem2);
assert_ne!(combined1, combined2);
}
#[test]
fn test_invalid_x25519_public_key_size() {
let server = HybridKeyPair::generate().unwrap();
let client = HybridClientExchange::new();
let bad_pk = vec![0u8; 16];
let result = client.complete(&bad_pk, server.mlkem_encapsulation_key());
assert!(result.is_err());
}
}