use alloc::vec::Vec;
use purecrypto::ec::x25519::X25519PrivateKey;
use purecrypto::hash::{Digest, Sha256};
use purecrypto::mlkem::{
MlKem768Ciphertext, MlKem768DecapsKey, MlKem768EncapsKey, SHARED_SECRET_BYTES,
};
use purecrypto::rng::{CryptoRng, RngCore};
use zeroize::Zeroizing;
use super::Kex;
use super::common::{
KexContext, KexInitOut, KexOutput, SSH_MSG_KEX_ECDH_INIT, SSH_MSG_KEX_ECDH_REPLY,
};
use super::hash::ExchangeHash;
use crate::error::{Error, Result};
use crate::format::Reader;
use crate::hostkey::HostKeyVerify;
pub struct MlKem768X25519Sha256;
impl Kex for MlKem768X25519Sha256 {
const NAME: &'static str = "mlkem768x25519-sha256";
const HASH_LEN: usize = 32;
}
const EK_PQ_LEN: usize = MlKem768DecapsKey::ENCAPS_KEY_BYTES; const CT_PQ_LEN: usize = MlKem768DecapsKey::CIPHERTEXT_BYTES; const X25519_PUB_LEN: usize = 32;
const C_INIT_LEN: usize = EK_PQ_LEN + X25519_PUB_LEN;
const S_REPLY_LEN: usize = CT_PQ_LEN + X25519_PUB_LEN;
impl MlKem768X25519Sha256 {
pub const NAME: &'static str = <Self as Kex>::NAME;
pub const HASH_LEN: usize = <Self as Kex>::HASH_LEN;
pub const C_INIT_LEN: usize = C_INIT_LEN;
pub const S_REPLY_LEN: usize = S_REPLY_LEN;
}
pub struct ClientState {
x25519_secret: X25519PrivateKey,
pq_secret: MlKem768DecapsKey,
c_init: Vec<u8>,
}
pub struct ServerReplyOut {
pub payload: Vec<u8>,
pub kex: KexOutput,
}
fn combine_secrets(
k_pq: &Zeroizing<[u8; SHARED_SECRET_BYTES]>,
k_ecdh: &Zeroizing<[u8; 32]>,
) -> Zeroizing<[u8; 32]> {
let mut h = Sha256::new();
h.update(&**k_pq);
h.update(&**k_ecdh);
let digest = h.finalize();
let mut out = Zeroizing::new([0u8; 32]);
out.copy_from_slice(digest.as_ref());
out
}
fn k_string_bytes(k: &[u8; 32]) -> Vec<u8> {
let mut out = Vec::with_capacity(4 + 32);
out.extend_from_slice(&(k.len() as u32).to_be_bytes());
out.extend_from_slice(k);
out
}
impl MlKem768X25519Sha256 {
pub fn client_init<R: RngCore + CryptoRng>(rng: &mut R) -> (ClientState, KexInitOut) {
let x25519_secret = X25519PrivateKey::generate(rng);
let q_c = x25519_secret.public_key();
let (pq_secret, pq_public) = MlKem768DecapsKey::generate(rng);
let ek_bytes = pq_public.to_bytes();
let mut c_init = Vec::with_capacity(C_INIT_LEN);
c_init.extend_from_slice(&ek_bytes);
c_init.extend_from_slice(&q_c);
debug_assert_eq!(c_init.len(), C_INIT_LEN);
let mut payload = Vec::with_capacity(1 + 4 + C_INIT_LEN);
payload.push(SSH_MSG_KEX_ECDH_INIT);
payload.extend_from_slice(&(C_INIT_LEN as u32).to_be_bytes());
payload.extend_from_slice(&c_init);
(
ClientState {
x25519_secret,
pq_secret,
c_init,
},
KexInitOut { payload },
)
}
pub fn server_reply<R, S>(
rng: &mut R,
init_payload: &[u8],
host_key: &S,
ctx: &KexContext<'_>,
) -> Result<ServerReplyOut>
where
R: RngCore + CryptoRng,
S: crate::hostkey::HostKey + ?Sized,
{
let mut r = Reader::new(init_payload);
let msg = r.read_u8()?;
if msg != SSH_MSG_KEX_ECDH_INIT {
return Err(Error::Protocol("expected SSH_MSG_KEX_HYBRID_INIT"));
}
let c_init = r.read_string()?;
if c_init.len() != C_INIT_LEN {
return Err(Error::Format("hybrid C_INIT wrong length"));
}
let (ek_bytes, q_c_bytes) = c_init.split_at(EK_PQ_LEN);
let mut ek_arr = [0u8; EK_PQ_LEN];
ek_arr.copy_from_slice(ek_bytes);
let mut q_c = [0u8; X25519_PUB_LEN];
q_c.copy_from_slice(q_c_bytes);
let ek = MlKem768EncapsKey::from_bytes_validated(ek_arr)
.map_err(|_| Error::Crypto("hybrid KEX agreement failed"))?;
let secret = X25519PrivateKey::generate(rng);
let q_s = secret.public_key();
let mut k_pq_raw = Zeroizing::new([0u8; SHARED_SECRET_BYTES]);
let (ct, k_pq_bytes) = ek.encapsulate(rng);
k_pq_raw.copy_from_slice(&k_pq_bytes);
let k_ecdh_raw: Zeroizing<[u8; 32]> = Zeroizing::new(
secret
.diffie_hellman(&q_c)
.map_err(|_| Error::Crypto("hybrid KEX agreement failed"))?,
);
let k_combined = combine_secrets(&k_pq_raw, &k_ecdh_raw);
let ct_bytes = ct.to_bytes();
let mut s_reply = Vec::with_capacity(S_REPLY_LEN);
s_reply.extend_from_slice(&ct_bytes);
s_reply.extend_from_slice(&q_s);
debug_assert_eq!(s_reply.len(), S_REPLY_LEN);
let k_s = host_key.public_blob();
let mut eh = ExchangeHash::<Sha256>::new();
eh.write_string(ctx.v_c);
eh.write_string(ctx.v_s);
eh.write_string(ctx.i_c);
eh.write_string(ctx.i_s);
eh.write_string(&k_s);
eh.write_string(c_init);
eh.write_string(&s_reply);
eh.write_string(&*k_combined);
let h = eh.finalize();
let sig = host_key.sign(&h)?;
let mut payload = Vec::with_capacity(1 + 4 + k_s.len() + 4 + s_reply.len() + 4 + sig.len());
payload.push(SSH_MSG_KEX_ECDH_REPLY);
payload.extend_from_slice(&(k_s.len() as u32).to_be_bytes());
payload.extend_from_slice(&k_s);
payload.extend_from_slice(&(s_reply.len() as u32).to_be_bytes());
payload.extend_from_slice(&s_reply);
payload.extend_from_slice(&(sig.len() as u32).to_be_bytes());
payload.extend_from_slice(&sig);
let k = k_string_bytes(&k_combined);
Ok(ServerReplyOut {
payload,
kex: KexOutput { k, h },
})
}
pub fn client_finish(
state: ClientState,
reply_payload: &[u8],
verifier: &dyn HostKeyVerify,
ctx: &KexContext<'_>,
) -> Result<KexOutput> {
let mut r = Reader::new(reply_payload);
let msg = r.read_u8()?;
if msg != SSH_MSG_KEX_ECDH_REPLY {
return Err(Error::Protocol("expected SSH_MSG_KEX_HYBRID_REPLY"));
}
let k_s = r.read_string()?;
let s_reply = r.read_string()?;
if s_reply.len() != S_REPLY_LEN {
return Err(Error::Format("hybrid S_REPLY wrong length"));
}
let sig = r.read_string()?;
let (ct_bytes, q_s_bytes) = s_reply.split_at(CT_PQ_LEN);
let mut ct_arr = [0u8; CT_PQ_LEN];
ct_arr.copy_from_slice(ct_bytes);
let mut q_s = [0u8; X25519_PUB_LEN];
q_s.copy_from_slice(q_s_bytes);
let ct = MlKem768Ciphertext::from_bytes(ct_arr);
let mut k_pq_raw = Zeroizing::new([0u8; SHARED_SECRET_BYTES]);
k_pq_raw.copy_from_slice(&state.pq_secret.decapsulate(&ct));
let k_ecdh_raw: Zeroizing<[u8; 32]> = Zeroizing::new(
state
.x25519_secret
.diffie_hellman(&q_s)
.map_err(|_| Error::Crypto("hybrid KEX agreement failed"))?,
);
let k_combined = combine_secrets(&k_pq_raw, &k_ecdh_raw);
let mut eh = ExchangeHash::<Sha256>::new();
eh.write_string(ctx.v_c);
eh.write_string(ctx.v_s);
eh.write_string(ctx.i_c);
eh.write_string(ctx.i_s);
eh.write_string(k_s);
eh.write_string(&state.c_init);
eh.write_string(s_reply);
eh.write_string(&*k_combined);
let h = eh.finalize();
verifier.verify(&h, sig)?;
let k = k_string_bytes(&k_combined);
Ok(KexOutput { k, h })
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::hostkey::{Ed25519HostKey, HostKey};
use crate::transport::version::LOCAL_VERSION;
use purecrypto::rng::HmacDrbg;
fn ctx() -> ([u8; 8], [u8; 8], [u8; 4], [u8; 4]) {
let mut v_c = [0u8; 8];
v_c.copy_from_slice(b"SSH-2.0_");
let mut v_s = [0u8; 8];
v_s.copy_from_slice(b"SSH-2.0&");
(v_c, v_s, [0x11, 0x22, 0x33, 0x44], [0x55, 0x66, 0x77, 0x88])
}
#[test]
fn algorithm_constants() {
assert_eq!(MlKem768X25519Sha256::NAME, "mlkem768x25519-sha256");
assert_eq!(MlKem768X25519Sha256::HASH_LEN, 32);
assert_eq!(MlKem768X25519Sha256::C_INIT_LEN, 1216);
assert_eq!(MlKem768X25519Sha256::S_REPLY_LEN, 1120);
}
#[test]
fn init_payload_layout() {
let mut rng = HmacDrbg::<Sha256>::new(b"hybrid-init", b"nonce", &[]);
let (_state, init) = MlKem768X25519Sha256::client_init(&mut rng);
assert_eq!(init.payload.len(), 1 + 4 + C_INIT_LEN);
assert_eq!(init.payload[0], SSH_MSG_KEX_ECDH_INIT);
assert_eq!(&init.payload[1..5], &(C_INIT_LEN as u32).to_be_bytes());
}
#[test]
fn round_trip_shared_secret_matches() {
let mut rng = HmacDrbg::<Sha256>::new(b"hybrid-roundtrip", b"nonce", &[]);
let mut seed = [0u8; 32];
rng.fill_bytes(&mut seed);
let server_hk = Ed25519HostKey::from_seed(seed);
let public = server_hk.public_bytes();
let client_verifier = Ed25519HostKey::from_public(public);
let (v_c, v_s, i_c, i_s) = ctx();
let kex_ctx = KexContext {
v_c: &v_c,
v_s: &v_s,
i_c: &i_c,
i_s: &i_s,
};
let (state, init) = MlKem768X25519Sha256::client_init(&mut rng);
let reply =
MlKem768X25519Sha256::server_reply(&mut rng, &init.payload, &server_hk, &kex_ctx)
.expect("server_reply");
let client_out =
MlKem768X25519Sha256::client_finish(state, &reply.payload, &client_verifier, &kex_ctx)
.expect("client_finish");
assert_eq!(client_out.k, reply.kex.k);
assert_eq!(client_out.h, reply.kex.h);
assert_eq!(client_out.h.len(), 32);
assert_eq!(client_out.k.len(), 4 + 32);
assert_eq!(&client_out.k[..4], &32u32.to_be_bytes());
let mut expected = Vec::new();
expected.extend_from_slice(&32u32.to_be_bytes());
expected.extend_from_slice(&client_out.k[4..]);
assert_eq!(client_out.k, expected);
assert!(!LOCAL_VERSION.is_empty());
}
#[test]
fn k_string_framing_keeps_high_bit_byte() {
let mut k = [0u8; 32];
k[0] = 0x80;
k[31] = 0x01;
let framed = k_string_bytes(&k);
assert_eq!(framed.len(), 4 + 32);
assert_eq!(&framed[..4], &32u32.to_be_bytes());
assert_eq!(framed[4], 0x80);
assert_eq!(&framed[4..], &k[..]);
let as_mpint = super::super::hash::mpint_bytes(&k);
assert_ne!(framed, as_mpint);
assert_eq!(&as_mpint[..4], &33u32.to_be_bytes());
}
#[test]
fn k_string_framing_keeps_leading_zero_byte() {
let mut k = [0u8; 32];
k[1] = 0x7F;
k[31] = 0xCD;
let framed = k_string_bytes(&k);
assert_eq!(framed.len(), 4 + 32);
assert_eq!(&framed[..4], &32u32.to_be_bytes());
assert_eq!(framed[4], 0x00);
assert_eq!(&framed[4..], &k[..]);
let as_mpint = super::super::hash::mpint_bytes(&k);
assert_ne!(framed, as_mpint);
assert_eq!(&as_mpint[..4], &31u32.to_be_bytes());
}
#[test]
fn c_init_wrong_length_rejected() {
let mut rng = HmacDrbg::<Sha256>::new(b"hybrid-bad-c", b"nonce", &[]);
let mut seed = [0u8; 32];
rng.fill_bytes(&mut seed);
let server_hk = Ed25519HostKey::from_seed(seed);
let (v_c, v_s, i_c, i_s) = ctx();
let kex_ctx = KexContext {
v_c: &v_c,
v_s: &v_s,
i_c: &i_c,
i_s: &i_s,
};
let mut bad = Vec::with_capacity(1 + 4 + 100);
bad.push(SSH_MSG_KEX_ECDH_INIT);
bad.extend_from_slice(&100u32.to_be_bytes());
bad.extend(core::iter::repeat_n(0u8, 100));
let result = MlKem768X25519Sha256::server_reply(&mut rng, &bad, &server_hk, &kex_ctx);
assert!(matches!(
result.map(|_| ()),
Err(Error::Format("hybrid C_INIT wrong length"))
));
}
#[test]
fn s_reply_wrong_length_rejected() {
let mut rng = HmacDrbg::<Sha256>::new(b"hybrid-bad-s", b"nonce", &[]);
let mut seed = [0u8; 32];
rng.fill_bytes(&mut seed);
let server_hk = Ed25519HostKey::from_seed(seed);
let public = server_hk.public_bytes();
let client_verifier = Ed25519HostKey::from_public(public);
let (v_c, v_s, i_c, i_s) = ctx();
let kex_ctx = KexContext {
v_c: &v_c,
v_s: &v_s,
i_c: &i_c,
i_s: &i_s,
};
let (state, _init) = MlKem768X25519Sha256::client_init(&mut rng);
let k_s = server_hk.public_blob();
let s_reply_bad = [0u8; 200];
let sig = [0u8; 64];
let mut reply = Vec::new();
reply.push(SSH_MSG_KEX_ECDH_REPLY);
reply.extend_from_slice(&(k_s.len() as u32).to_be_bytes());
reply.extend_from_slice(&k_s);
reply.extend_from_slice(&(s_reply_bad.len() as u32).to_be_bytes());
reply.extend_from_slice(&s_reply_bad);
reply.extend_from_slice(&(sig.len() as u32).to_be_bytes());
reply.extend_from_slice(&sig);
let result = MlKem768X25519Sha256::client_finish(state, &reply, &client_verifier, &kex_ctx);
assert!(matches!(
result.map(|_| ()),
Err(Error::Format("hybrid S_REPLY wrong length"))
));
}
#[test]
fn exchange_hash_uses_sha256() {
let mut rng = HmacDrbg::<Sha256>::new(b"hybrid-hash", b"nonce", &[]);
let mut seed = [0u8; 32];
rng.fill_bytes(&mut seed);
let server_hk = Ed25519HostKey::from_seed(seed);
let (v_c, v_s, i_c, i_s) = ctx();
let kex_ctx = KexContext {
v_c: &v_c,
v_s: &v_s,
i_c: &i_c,
i_s: &i_s,
};
let (_state, init) = MlKem768X25519Sha256::client_init(&mut rng);
let reply =
MlKem768X25519Sha256::server_reply(&mut rng, &init.payload, &server_hk, &kex_ctx)
.expect("server_reply");
assert_eq!(reply.kex.h.len(), 32);
}
}