use ed25519_dalek::{Signature, Signer, SigningKey, Verifier, VerifyingKey};
use hkdf::Hkdf;
use ml_dsa::{
B32, EncodedSignature as MlDsaEncodedSignature,
EncodedVerifyingKey as MlDsaEncodedVerifyingKey, Generate as MlDsaGenerate, MlDsa65,
Signature as MlDsaSignatureT, SigningKey as MlDsaSigningKey, VerifyingKey as MlDsaVerifyingKey,
signature::{Keypair as MlDsaKeypair, Signer as MlDsaSigner, Verifier as MlDsaVerifier},
};
use ml_kem::{TryKeyInit, kem::Encapsulate, ml_kem_768::EncapsulationKey as MlKemEk};
use rand::rngs::OsRng;
use serde::{Deserialize, Serialize};
use serde_big_array::BigArray;
use sha2::Sha256;
const ML_KEM_768_PK_LEN: usize = 1184;
const ML_KEM_768_CT_LEN: usize = 1088;
const ML_DSA_65_PK_LEN: usize = 1952;
const ML_DSA_65_SIG_LEN: usize = 3309;
const HKDF_SALT: &[u8] = b"xenia-handshake-v1";
const HKDF_INFO: &[u8] = b"xenia-session-key";
const KEM_SUITE_LABEL: &str = "ml-kem-768-fips203";
const TRANSCRIPT_SIGNATURE_SUITE_LABEL: &str = "ed25519-rfc8032+ml-dsa-65-fips204";
const KDF_SUITE_LABEL: &str = "hkdf-sha256";
const HANDSHAKE_POLICY_PROFILE: &str = "hybrid-pq-transcript-v1";
const HANDSHAKE_TRANSCRIPT_SCHEMA: &str = "xenia-handshake-transcript-v1";
const HANDSHAKE_SIGNATURE_CONTEXT_V1: &str = "xenia-handshake-signature-v1";
const SESSION_KEY_SCHEDULE_SCHEMA: &str = "xenia-session-key-schedule-v1";
const SESSION_AEAD_KEY_LABEL: &[u8] = b"xenia/session/aead";
const SESSION_CONTROL_KEY_LABEL: &[u8] = b"xenia/session/control";
const SESSION_VIDEO_KEY_LABEL: &[u8] = b"xenia/session/video";
const SESSION_AUDIO_KEY_LABEL: &[u8] = b"xenia/session/audio";
const SESSION_TELEMETRY_KEY_LABEL: &[u8] = b"xenia/session/telemetry";
const SESSION_REKEY_KEY_LABEL: &[u8] = b"xenia/session/rekey";
const SESSION_CONTEXT_KEY_LABEL: &[u8] = b"xenia/session/context";
#[allow(clippy::large_enum_variant)]
#[derive(Debug, Clone, Serialize, Deserialize)]
enum HandshakeMessage {
HostHello {
ed25519_pk: [u8; 32],
#[serde(with = "BigArray")]
ml_dsa_pk: [u8; ML_DSA_65_PK_LEN],
#[serde(with = "BigArray")]
kem_pk: [u8; ML_KEM_768_PK_LEN],
nonce: [u8; 32],
negotiated_context_hash: Option<[u8; 32]>,
},
ViewerResponse {
ed25519_pk: [u8; 32],
#[serde(with = "BigArray")]
ml_dsa_pk: [u8; ML_DSA_65_PK_LEN],
#[serde(with = "BigArray")]
kem_ct: [u8; ML_KEM_768_CT_LEN],
nonce: [u8; 32],
#[serde(with = "BigArray")]
signature: [u8; 64],
#[serde(with = "BigArray")]
ml_dsa_signature: [u8; ML_DSA_65_SIG_LEN],
},
HostFinalize {
#[serde(with = "BigArray")]
signature: [u8; 64],
#[serde(with = "BigArray")]
ml_dsa_signature: [u8; ML_DSA_65_SIG_LEN],
},
}
#[derive(Debug, Clone, Serialize, Deserialize)]
struct HandshakeTranscriptV1 {
schema: String,
kem: String,
transcript_signature: String,
kdf: String,
negotiated_context_hash: Option<[u8; 32]>,
host_ed25519_pk: [u8; 32],
viewer_ed25519_pk: [u8; 32],
host_ml_dsa_pk: Vec<u8>,
viewer_ml_dsa_pk: Vec<u8>,
host_kem_pk: Vec<u8>,
kem_ciphertext: Vec<u8>,
host_nonce: [u8; 32],
viewer_nonce: [u8; 32],
viewer_signature: Vec<u8>,
host_signature: Vec<u8>,
viewer_ml_dsa_signature: Vec<u8>,
host_ml_dsa_signature: Vec<u8>,
}
fn append_len_prefixed(out: &mut Vec<u8>, bytes: &[u8]) {
let len = u32::try_from(bytes.len()).expect("handshake transcript component exceeds u32");
out.extend_from_slice(&len.to_be_bytes());
out.extend_from_slice(bytes);
}
fn signature_context_prefix() -> Vec<u8> {
let mut out = Vec::new();
append_len_prefixed(&mut out, HANDSHAKE_SIGNATURE_CONTEXT_V1.as_bytes());
append_len_prefixed(&mut out, HANDSHAKE_TRANSCRIPT_SCHEMA.as_bytes());
append_len_prefixed(&mut out, HANDSHAKE_POLICY_PROFILE.as_bytes());
append_len_prefixed(&mut out, KEM_SUITE_LABEL.as_bytes());
append_len_prefixed(&mut out, TRANSCRIPT_SIGNATURE_SUITE_LABEL.as_bytes());
append_len_prefixed(&mut out, KDF_SUITE_LABEL.as_bytes());
out
}
fn viewer_signature_transcript(
hello_bytes: &[u8],
viewer_ed25519_pk: &[u8; 32],
viewer_ml_dsa_pk: &[u8; ML_DSA_65_PK_LEN],
kem_ct: &[u8],
viewer_nonce: &[u8; 32],
) -> Vec<u8> {
let mut transcript = signature_context_prefix();
append_len_prefixed(&mut transcript, b"viewer-response");
append_len_prefixed(&mut transcript, hello_bytes);
append_len_prefixed(&mut transcript, viewer_ed25519_pk);
append_len_prefixed(&mut transcript, viewer_ml_dsa_pk);
append_len_prefixed(&mut transcript, kem_ct);
append_len_prefixed(&mut transcript, viewer_nonce);
transcript
}
fn host_signature_transcript(
hello_bytes: &[u8],
viewer_ed25519_pk: &[u8; 32],
viewer_ml_dsa_pk: &[u8; ML_DSA_65_PK_LEN],
kem_ct: &[u8],
viewer_nonce: &[u8; 32],
viewer_signature: &[u8; 64],
viewer_ml_dsa_signature: &[u8; ML_DSA_65_SIG_LEN],
) -> Vec<u8> {
let mut transcript = viewer_signature_transcript(
hello_bytes,
viewer_ed25519_pk,
viewer_ml_dsa_pk,
kem_ct,
viewer_nonce,
);
append_len_prefixed(&mut transcript, b"host-finalize");
append_len_prefixed(&mut transcript, viewer_signature);
append_len_prefixed(&mut transcript, viewer_ml_dsa_signature);
transcript
}
fn hkdf_derive(classical_nonce: &[u8], kem_shared_secret: &[u8]) -> [u8; 32] {
let mut ikm = Vec::with_capacity(classical_nonce.len() + kem_shared_secret.len());
ikm.extend_from_slice(classical_nonce);
ikm.extend_from_slice(kem_shared_secret);
let hk = Hkdf::<Sha256>::new(Some(HKDF_SALT), &ikm);
let mut okm = [0u8; 32];
hk.expand(HKDF_INFO, &mut okm)
.expect("HKDF-SHA256 32-byte expand cannot fail for 32-byte output");
okm
}
pub fn derive_labeled_session_key(
root_key: &[u8; 32],
transcript_hash: &[u8; 32],
label: &[u8],
) -> [u8; 32] {
let hk = Hkdf::<Sha256>::new(Some(SESSION_KEY_SCHEDULE_SCHEMA.as_bytes()), root_key);
let mut info = Vec::with_capacity(label.len() + 1 + transcript_hash.len());
info.extend_from_slice(label);
info.extend_from_slice(b":");
info.extend_from_slice(transcript_hash);
let mut okm = [0u8; 32];
hk.expand(&info, &mut okm)
.expect("HKDF-SHA256 32-byte expand cannot fail for labeled session key");
okm
}
#[derive(Debug, thiserror::Error)]
pub enum HandshakeError {
#[error("expected HostHello message")]
ExpectedHostHello,
#[error("expected HostFinalize message")]
ExpectedHostFinalize,
#[error("bincode decode/encode failed: {0}")]
Codec(#[from] bincode::Error),
#[error("invalid host Ed25519 verifying key")]
InvalidVerifyingKey,
#[error("invalid host ML-KEM-768 public key")]
InvalidKemPublicKey,
#[error("host signature verification failed")]
SignatureVerificationFailed,
#[error("invalid host ML-DSA-65 verifying key")]
InvalidMlDsaVerifyingKey,
#[error("invalid ML-DSA-65 signature encoding")]
InvalidMlDsaSignatureEncoding,
#[error("host ML-DSA-65 signature verification failed")]
MlDsaSignatureVerificationFailed,
#[error("called finish() before begin()")]
NotStarted,
#[error("identity seed must be exactly 32 bytes")]
InvalidSeedLength,
}
fn parse_peer_ml_dsa_public_key(
bytes: &[u8; ML_DSA_65_PK_LEN],
) -> Result<MlDsaVerifyingKey<MlDsa65>, HandshakeError> {
let encoded = MlDsaEncodedVerifyingKey::<MlDsa65>::try_from(bytes.as_slice())
.map_err(|_| HandshakeError::InvalidMlDsaVerifyingKey)?;
Ok(MlDsaVerifyingKey::<MlDsa65>::decode(&encoded))
}
fn verify_ml_dsa(
peer_pk: &[u8; ML_DSA_65_PK_LEN],
message: &[u8],
signature: &[u8; ML_DSA_65_SIG_LEN],
) -> Result<(), HandshakeError> {
let verifying_key = parse_peer_ml_dsa_public_key(peer_pk)?;
let encoded_sig = MlDsaEncodedSignature::<MlDsa65>::try_from(signature.as_slice())
.map_err(|_| HandshakeError::InvalidMlDsaSignatureEncoding)?;
let sig = MlDsaSignatureT::<MlDsa65>::decode(&encoded_sig)
.ok_or(HandshakeError::InvalidMlDsaSignatureEncoding)?;
verifying_key
.verify(message, &sig)
.map_err(|_| HandshakeError::MlDsaSignatureVerificationFailed)
}
struct PendingState {
hello_bytes: Vec<u8>,
host_ed25519_pk: [u8; 32],
host_verifying_key: VerifyingKey,
host_ml_dsa_pk: [u8; ML_DSA_65_PK_LEN],
host_kem_pk: [u8; ML_KEM_768_PK_LEN],
host_nonce: [u8; 32],
negotiated_context_hash: Option<[u8; 32]>,
viewer_nonce: [u8; 32],
kem_ct: [u8; ML_KEM_768_CT_LEN],
viewer_ed25519_pk: [u8; 32],
viewer_ml_dsa_pk: [u8; ML_DSA_65_PK_LEN],
viewer_signature: [u8; 64],
viewer_ml_dsa_signature: [u8; ML_DSA_65_SIG_LEN],
root_key: [u8; 32],
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct SessionKeySchedule {
pub aead: [u8; 32],
pub control: [u8; 32],
pub video: [u8; 32],
pub audio: [u8; 32],
pub telemetry: [u8; 32],
pub rekey: [u8; 32],
pub context: [u8; 32],
pub transcript_hash: [u8; 32],
pub host_identity_fingerprint: [u8; 32],
}
impl SessionKeySchedule {
fn derive(
root_key: &[u8; 32],
transcript_hash: [u8; 32],
host_identity_fingerprint: [u8; 32],
) -> Self {
Self {
host_identity_fingerprint,
aead: derive_labeled_session_key(root_key, &transcript_hash, SESSION_AEAD_KEY_LABEL),
control: derive_labeled_session_key(
root_key,
&transcript_hash,
SESSION_CONTROL_KEY_LABEL,
),
video: derive_labeled_session_key(root_key, &transcript_hash, SESSION_VIDEO_KEY_LABEL),
audio: derive_labeled_session_key(root_key, &transcript_hash, SESSION_AUDIO_KEY_LABEL),
telemetry: derive_labeled_session_key(
root_key,
&transcript_hash,
SESSION_TELEMETRY_KEY_LABEL,
),
rekey: derive_labeled_session_key(root_key, &transcript_hash, SESSION_REKEY_KEY_LABEL),
context: derive_labeled_session_key(
root_key,
&transcript_hash,
SESSION_CONTEXT_KEY_LABEL,
),
transcript_hash,
}
}
}
pub struct ViewerHandshake {
signing_key: SigningKey,
ml_dsa_signing_key: MlDsaSigningKey<MlDsa65>,
pending: Option<PendingState>,
}
impl Default for ViewerHandshake {
fn default() -> Self {
Self::new()
}
}
impl ViewerHandshake {
pub fn new() -> Self {
Self {
signing_key: SigningKey::generate(&mut OsRng),
ml_dsa_signing_key: MlDsaSigningKey::<MlDsa65>::generate(),
pending: None,
}
}
pub fn from_identity(
ed25519_secret: &[u8],
ml_dsa_seed: &[u8],
) -> Result<Self, HandshakeError> {
let ed: [u8; 32] = ed25519_secret
.try_into()
.map_err(|_| HandshakeError::InvalidSeedLength)?;
let ml: [u8; 32] = ml_dsa_seed
.try_into()
.map_err(|_| HandshakeError::InvalidSeedLength)?;
let seed: B32 = ml.into();
Ok(Self {
signing_key: SigningKey::from_bytes(&ed),
ml_dsa_signing_key: MlDsaSigningKey::<MlDsa65>::from_seed(&seed),
pending: None,
})
}
pub fn ed25519_public_key(&self) -> [u8; 32] {
self.signing_key.verifying_key().to_bytes()
}
pub fn begin(&mut self, hello_bytes: &[u8]) -> Result<Vec<u8>, HandshakeError> {
let hello: HandshakeMessage = bincode::deserialize(hello_bytes)?;
let HandshakeMessage::HostHello {
ed25519_pk,
ml_dsa_pk: host_ml_dsa_pk,
kem_pk,
nonce: host_nonce,
negotiated_context_hash,
} = hello
else {
return Err(HandshakeError::ExpectedHostHello);
};
let host_verifying_key = VerifyingKey::from_bytes(&ed25519_pk)
.map_err(|_| HandshakeError::InvalidVerifyingKey)?;
parse_peer_ml_dsa_public_key(&host_ml_dsa_pk)?;
let viewer_nonce = rand::random::<[u8; 32]>();
let ek: MlKemEk = <MlKemEk as TryKeyInit>::new_from_slice(&kem_pk)
.map_err(|_| HandshakeError::InvalidKemPublicKey)?;
let (kem_ct, shared) = <MlKemEk as Encapsulate>::encapsulate(&ek);
let kem_ct: [u8; ML_KEM_768_CT_LEN] = kem_ct
.as_slice()
.try_into()
.map_err(|_| HandshakeError::InvalidKemPublicKey)?;
let viewer_ed25519_pk = self.signing_key.verifying_key().to_bytes();
let viewer_ml_dsa_pk: [u8; ML_DSA_65_PK_LEN] = self
.ml_dsa_signing_key
.verifying_key()
.encode()
.as_slice()
.try_into()
.expect("ml-dsa-65 encoded verifying key is always ML_DSA_65_PK_LEN bytes");
let transcript = viewer_signature_transcript(
hello_bytes,
&viewer_ed25519_pk,
&viewer_ml_dsa_pk,
&kem_ct,
&viewer_nonce,
);
let viewer_signature = self.signing_key.sign(&transcript).to_bytes();
let viewer_ml_dsa_signature: [u8; ML_DSA_65_SIG_LEN] = {
let sig: MlDsaSignatureT<MlDsa65> = self.ml_dsa_signing_key.sign(&transcript);
sig.encode()
.as_slice()
.try_into()
.expect("ml-dsa-65 encoded signature is always ML_DSA_65_SIG_LEN bytes")
};
let mut combined_nonce = [0u8; 64];
combined_nonce[..32].copy_from_slice(&host_nonce);
combined_nonce[32..].copy_from_slice(&viewer_nonce);
let root_key = hkdf_derive(&combined_nonce, shared.as_slice());
self.pending = Some(PendingState {
hello_bytes: hello_bytes.to_vec(),
host_ed25519_pk: ed25519_pk,
host_verifying_key,
host_ml_dsa_pk,
host_kem_pk: kem_pk,
host_nonce,
negotiated_context_hash,
viewer_nonce,
kem_ct,
viewer_ed25519_pk,
viewer_ml_dsa_pk,
viewer_signature,
viewer_ml_dsa_signature,
root_key,
});
let response = HandshakeMessage::ViewerResponse {
ed25519_pk: viewer_ed25519_pk,
ml_dsa_pk: viewer_ml_dsa_pk,
kem_ct,
nonce: viewer_nonce,
signature: viewer_signature,
ml_dsa_signature: viewer_ml_dsa_signature,
};
Ok(bincode::serialize(&response)?)
}
pub fn finish(&mut self, finalize_bytes: &[u8]) -> Result<SessionKeySchedule, HandshakeError> {
let state = self.pending.take().ok_or(HandshakeError::NotStarted)?;
let finalize: HandshakeMessage = bincode::deserialize(finalize_bytes)?;
let HandshakeMessage::HostFinalize {
signature: host_sig_bytes,
ml_dsa_signature: host_ml_dsa_sig_bytes,
} = finalize
else {
return Err(HandshakeError::ExpectedHostFinalize);
};
let final_transcript = host_signature_transcript(
&state.hello_bytes,
&state.viewer_ed25519_pk,
&state.viewer_ml_dsa_pk,
&state.kem_ct,
&state.viewer_nonce,
&state.viewer_signature,
&state.viewer_ml_dsa_signature,
);
let host_sig = Signature::from_bytes(&host_sig_bytes);
state
.host_verifying_key
.verify(&final_transcript, &host_sig)
.map_err(|_| HandshakeError::SignatureVerificationFailed)?;
verify_ml_dsa(
&state.host_ml_dsa_pk,
&final_transcript,
&host_ml_dsa_sig_bytes,
)?;
let transcript = HandshakeTranscriptV1 {
schema: HANDSHAKE_TRANSCRIPT_SCHEMA.to_string(),
kem: KEM_SUITE_LABEL.to_string(),
transcript_signature: TRANSCRIPT_SIGNATURE_SUITE_LABEL.to_string(),
kdf: KDF_SUITE_LABEL.to_string(),
negotiated_context_hash: state.negotiated_context_hash,
host_ed25519_pk: state.host_ed25519_pk,
viewer_ed25519_pk: state.viewer_ed25519_pk,
host_ml_dsa_pk: state.host_ml_dsa_pk.to_vec(),
viewer_ml_dsa_pk: state.viewer_ml_dsa_pk.to_vec(),
host_kem_pk: state.host_kem_pk.to_vec(),
kem_ciphertext: state.kem_ct.to_vec(),
host_nonce: state.host_nonce,
viewer_nonce: state.viewer_nonce,
viewer_signature: state.viewer_signature.to_vec(),
host_signature: host_sig_bytes.to_vec(),
viewer_ml_dsa_signature: state.viewer_ml_dsa_signature.to_vec(),
host_ml_dsa_signature: host_ml_dsa_sig_bytes.to_vec(),
};
let transcript_bytes = bincode::serialize(&transcript)?;
let transcript_hash = *blake3::hash(&transcript_bytes).as_bytes();
let host_fingerprint =
host_identity_fingerprint(&state.host_ed25519_pk, &state.host_ml_dsa_pk);
Ok(SessionKeySchedule::derive(
&state.root_key,
transcript_hash,
host_fingerprint,
))
}
}
fn host_identity_fingerprint(ed25519_pk: &[u8; 32], ml_dsa_pk: &[u8]) -> [u8; 32] {
let mut hasher = blake3::Hasher::new();
hasher.update(b"xenia-host-identity-fingerprint-v1");
hasher.update(ed25519_pk);
hasher.update(ml_dsa_pk);
*hasher.finalize().as_bytes()
}