use crate::pq::{self, MlKem768Keypair};
use sha3::{Digest, Sha3_256};
use x25519_dalek::{PublicKey, StaticSecret};
use rand::rngs::OsRng;
use zeroize::Zeroize;
const COMBINER_LABEL: &[u8] = b"ling-hybrid-x25519-mlkem768-v1";
const X25519_LEN: usize = 32;
pub const PUBLIC_KEY_LEN: usize = X25519_LEN + pq::ENCAPS_KEY_LEN; pub const CIPHERTEXT_LEN: usize = X25519_LEN + pq::CIPHERTEXT_LEN; pub const SHARED_SECRET_LEN: usize = 32;
fn combine(ss_mlkem: &[u8], ss_x25519: &[u8], eph_pk: &[u8], recipient_pk: &[u8]) -> [u8; 32] {
let mut h = Sha3_256::new();
h.update(COMBINER_LABEL);
h.update(ss_mlkem);
h.update(ss_x25519);
h.update(eph_pk);
h.update(recipient_pk);
h.finalize().into()
}
pub struct HybridKeypair {
x25519_secret: StaticSecret,
x25519_public: [u8; X25519_LEN],
mlkem: MlKem768Keypair,
}
impl HybridKeypair {
pub fn generate() -> Self {
let x25519_secret = StaticSecret::random_from_rng(OsRng);
let x25519_public = PublicKey::from(&x25519_secret).to_bytes();
Self { x25519_secret, x25519_public, mlkem: MlKem768Keypair::generate() }
}
pub fn public_key(&self) -> Vec<u8> {
let mut out = Vec::with_capacity(PUBLIC_KEY_LEN);
out.extend_from_slice(&self.x25519_public);
out.extend_from_slice(&self.mlkem.encapsulation_key());
out
}
pub fn decapsulate(&self, ciphertext: &[u8]) -> Result<[u8; SHARED_SECRET_LEN], &'static str> {
if ciphertext.len() != CIPHERTEXT_LEN {
return Err("hybrid ciphertext wrong length");
}
let (eph_pk_bytes, ct_pq) = ciphertext.split_at(X25519_LEN);
let mut eph_arr = [0u8; X25519_LEN];
eph_arr.copy_from_slice(eph_pk_bytes);
let eph_pk = PublicKey::from(eph_arr);
let mut ss_x = self.x25519_secret.diffie_hellman(&eph_pk).to_bytes();
let ss_pq = self.mlkem.decapsulate(ct_pq)?;
let out = combine(&ss_pq, &ss_x, eph_pk_bytes, &self.x25519_public);
ss_x.zeroize();
Ok(out)
}
}
pub fn encapsulate(hybrid_public_key: &[u8]) -> Result<(Vec<u8>, [u8; SHARED_SECRET_LEN]), &'static str> {
if hybrid_public_key.len() != PUBLIC_KEY_LEN {
return Err("hybrid public key wrong length");
}
let (x_pk_bytes, mlkem_ek) = hybrid_public_key.split_at(X25519_LEN);
let mut x_pk_arr = [0u8; X25519_LEN];
x_pk_arr.copy_from_slice(x_pk_bytes);
let eph_secret = StaticSecret::random_from_rng(OsRng);
let eph_public = PublicKey::from(&eph_secret).to_bytes();
let mut ss_x = eph_secret.diffie_hellman(&PublicKey::from(x_pk_arr)).to_bytes();
let (ct_pq, ss_pq) = pq::encapsulate(mlkem_ek)?;
let shared = combine(&ss_pq, &ss_x, &eph_public, x_pk_bytes);
ss_x.zeroize();
let mut ciphertext = Vec::with_capacity(CIPHERTEXT_LEN);
ciphertext.extend_from_slice(&eph_public);
ciphertext.extend_from_slice(&ct_pq);
Ok((ciphertext, shared))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn round_trip() {
let kp = HybridKeypair::generate();
let pk = kp.public_key();
assert_eq!(pk.len(), PUBLIC_KEY_LEN);
let (ct, ss_send) = encapsulate(&pk).expect("encapsulate");
assert_eq!(ct.len(), CIPHERTEXT_LEN);
let ss_recv = kp.decapsulate(&ct).expect("decapsulate");
assert_eq!(ss_send, ss_recv, "hybrid shared secrets agree");
}
#[test]
fn distinct_encapsulations_differ() {
let kp = HybridKeypair::generate();
let pk = kp.public_key();
let (_, a) = encapsulate(&pk).unwrap();
let (_, b) = encapsulate(&pk).unwrap();
assert_ne!(a, b, "fresh randomness yields distinct shared secrets");
}
#[test]
fn tampered_ciphertext_changes_secret() {
let kp = HybridKeypair::generate();
let pk = kp.public_key();
let (mut ct, ss_send) = encapsulate(&pk).unwrap();
let last = ct.len() - 1;
ct[last] ^= 0xFF;
let ss_recv = kp.decapsulate(&ct).unwrap();
assert_ne!(ss_send, ss_recv);
}
#[test]
fn wrong_lengths_rejected() {
let kp = HybridKeypair::generate();
assert!(encapsulate(&[0u8; 10]).is_err());
assert!(kp.decapsulate(&[0u8; 10]).is_err());
}
}