use crate::CryptoError;
use crate::agreement::{KeyAgreement, P256Agreement};
use ml_kem::array::{Array, ArrayN};
use ml_kem::{Decapsulate, DecapsulationKey768, EncapsulationKey768, FromSeed, KeyExport};
use ml_kem::{MlKem768, TryKeyInit};
use sha3::digest::{ExtendableOutput, Update, XofReader};
use sha3::{Digest, Sha3_256, Shake256};
use zeroize::Zeroizing;
pub const PUBLIC_KEY_LEN: usize = 1249;
pub const CIPHERTEXT_LEN: usize = 1153;
pub const SHARED_LEN: usize = 32;
pub const ML_KEM_SEED_LEN: usize = 64;
pub const ENCAPS_SEED_LEN: usize = 160;
pub const KEYPAIR_SEED_LEN: usize = 32;
const ML_KEM_PUBLIC_LEN: usize = 1184;
const ML_KEM_CIPHERTEXT_LEN: usize = 1088;
const P256_POINT_LEN: usize = 65;
const SCALAR_LEN: usize = 32;
const SCALAR_TRIES_LEN: usize = 128;
const EXPANDED_LEN: usize = ML_KEM_SEED_LEN + SCALAR_TRIES_LEN;
const LABEL: &[u8] = b"MLKEM768-P256";
fn combine(
ss_pq: &[u8],
ss_t: &[u8],
ct_t: &[u8],
pk_t: &[u8],
) -> Zeroizing<[u8; SHARED_LEN]> {
let mut h = Sha3_256::new();
Digest::update(&mut h, ss_pq);
Digest::update(&mut h, ss_t);
Digest::update(&mut h, ct_t);
Digest::update(&mut h, pk_t);
Digest::update(&mut h, LABEL);
Zeroizing::new(h.finalize().into())
}
fn scalar_from(bytes: &[u8]) -> Result<P256Agreement, CryptoError> {
for part in bytes.chunks_exact(SCALAR_LEN) {
let candidate: [u8; SCALAR_LEN] =
part.try_into().map_err(|_| CryptoError::BadLength)?;
if let Ok(pair) = P256Agreement::from_be_bytes(&candidate) {
return Ok(pair);
}
}
Err(CryptoError::BadKey)
}
pub struct Keypair {
pub ml_kem_seed: Zeroizing<[u8; ML_KEM_SEED_LEN]>,
pub classical: P256Agreement,
pub public_key: [u8; PUBLIC_KEY_LEN],
}
impl core::fmt::Debug for Keypair {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("Keypair")
.field("ml_kem_seed", &"<redacted>")
.field("classical", &"<redacted>")
.field("public_key", &self.public_key)
.finish()
}
}
pub fn keypair_from_seed(seed: &[u8; KEYPAIR_SEED_LEN]) -> Result<Keypair, CryptoError> {
let mut xof = Shake256::default();
Update::update(&mut xof, seed);
let mut expanded = Zeroizing::new([0u8; EXPANDED_LEN]);
xof.finalize_xof().read(expanded.as_mut_slice());
let pq_seed: [u8; ML_KEM_SEED_LEN] = expanded
.get(..ML_KEM_SEED_LEN)
.ok_or(CryptoError::BadLength)?
.try_into()
.map_err(|_| CryptoError::BadLength)?;
let pq_seed = Zeroizing::new(pq_seed);
let classical =
scalar_from(expanded.get(ML_KEM_SEED_LEN..EXPANDED_LEN).ok_or(CryptoError::BadLength)?)?;
let public_key = public_key_from_parts(&pq_seed, &classical.public_key())?;
Ok(Keypair { ml_kem_seed: pq_seed, classical, public_key })
}
#[must_use]
pub fn ml_kem_seed_from_stored(stored: &[u8; 32]) -> Zeroizing<[u8; ML_KEM_SEED_LEN]> {
let mut xof = Shake256::default();
Update::update(&mut xof, stored);
let mut out = Zeroizing::new([0u8; ML_KEM_SEED_LEN]);
xof.finalize_xof().read(out.as_mut_slice());
out
}
pub fn public_key_from_parts(
ml_kem_seed: &[u8; ML_KEM_SEED_LEN],
p256_public: &[u8],
) -> Result<[u8; PUBLIC_KEY_LEN], CryptoError> {
if p256_public.len() != P256_POINT_LEN {
return Err(CryptoError::BadLength);
}
let seed = ArrayN::<u8, ML_KEM_SEED_LEN>::try_from(ml_kem_seed.as_slice())
.map_err(|_| CryptoError::BadLength)?;
let (_, ek) = MlKem768::from_seed(&seed);
let mut out = [0u8; PUBLIC_KEY_LEN];
let (head, tail) = out.split_at_mut(ML_KEM_PUBLIC_LEN);
head.copy_from_slice(ek.to_bytes().as_slice());
tail.copy_from_slice(p256_public);
Ok(out)
}
pub fn encapsulate_derand(
public_key: &[u8],
seed: &[u8; ENCAPS_SEED_LEN],
) -> Result<(Zeroizing<[u8; SHARED_LEN]>, [u8; CIPHERTEXT_LEN]), CryptoError> {
if public_key.len() != PUBLIC_KEY_LEN {
return Err(CryptoError::BadLength);
}
let pk_pq = public_key.get(..ML_KEM_PUBLIC_LEN).ok_or(CryptoError::BadLength)?;
let pk_t = public_key.get(ML_KEM_PUBLIC_LEN..PUBLIC_KEY_LEN).ok_or(CryptoError::BadLength)?;
let ek = EncapsulationKey768::new_from_slice(pk_pq).map_err(|_| CryptoError::BadKey)?;
let m_bytes = seed.get(..SCALAR_LEN).ok_or(CryptoError::BadLength)?;
let m = ArrayN::<u8, SCALAR_LEN>::try_from(m_bytes).map_err(|_| CryptoError::BadLength)?;
let (ct_pq, ss_pq) = ek.encapsulate_deterministic(&m);
let ephemeral =
scalar_from(seed.get(SCALAR_LEN..ENCAPS_SEED_LEN).ok_or(CryptoError::BadLength)?)?;
let ct_t = ephemeral.public_key();
let ss_t = ephemeral.agree(pk_t)?;
let shared = combine(&ss_pq, ss_t.expose(), &ct_t, pk_t);
let mut ciphertext = [0u8; CIPHERTEXT_LEN];
let (head, tail) = ciphertext.split_at_mut(ML_KEM_CIPHERTEXT_LEN);
head.copy_from_slice(&ct_pq);
tail.copy_from_slice(&ct_t);
Ok((shared, ciphertext))
}
pub fn encapsulate<R: rand_core::CryptoRng + ?Sized>(
public_key: &[u8],
rng: &mut R,
) -> Result<(Zeroizing<[u8; SHARED_LEN]>, [u8; CIPHERTEXT_LEN]), CryptoError> {
let mut seed = Zeroizing::new([0u8; ENCAPS_SEED_LEN]);
rng.fill_bytes(seed.as_mut_slice());
encapsulate_derand(public_key, &seed)
}
pub fn decapsulate_with(
ml_kem_seed: &[u8; ML_KEM_SEED_LEN],
classical: &dyn KeyAgreement,
ciphertext: &[u8],
) -> Result<Zeroizing<[u8; SHARED_LEN]>, CryptoError> {
if ciphertext.len() != CIPHERTEXT_LEN {
return Err(CryptoError::BadLength);
}
let seed = ArrayN::<u8, ML_KEM_SEED_LEN>::try_from(ml_kem_seed.as_slice())
.map_err(|_| CryptoError::BadLength)?;
let (dk, _) = MlKem768::from_seed(&seed);
let ct_pq_bytes = ciphertext.get(..ML_KEM_CIPHERTEXT_LEN).ok_or(CryptoError::BadLength)?;
let ct_pq = Array::try_from(ct_pq_bytes).map_err(|_| CryptoError::BadLength)?;
let ct_t = ciphertext.get(ML_KEM_CIPHERTEXT_LEN..CIPHERTEXT_LEN).ok_or(CryptoError::BadLength)?;
let ss_pq = DecapsulationKey768::decapsulate(&dk, &ct_pq);
let ss_t = classical.agree(ct_t)?;
let pk_t = classical.public_key();
Ok(combine(&ss_pq, ss_t.expose(), ct_t, &pk_t))
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
mod tests {
use super::*;
#[test]
fn what_is_encapsulated_comes_back_out_of_decapsulation() {
let pair = keypair_from_seed(&[0x31; KEYPAIR_SEED_LEN]).unwrap();
let (pq_seed, classical, public) = (pair.ml_kem_seed, pair.classical, pair.public_key);
let (sent, ciphertext) =
encapsulate_derand(&public, &[0x5c; ENCAPS_SEED_LEN]).unwrap();
let received = decapsulate_with(&pq_seed, &classical, &ciphertext).unwrap();
assert_eq!(sent.as_slice(), received.as_slice(), "секрет не сошёлся с обеих сторон");
}
#[test]
fn substituting_either_half_yields_a_different_secret() {
let pair = keypair_from_seed(&[0x31; KEYPAIR_SEED_LEN]).unwrap();
let (pq_seed, classical, public) = (pair.ml_kem_seed, pair.classical, pair.public_key);
let (sent, ciphertext) =
encapsulate_derand(&public, &[0x5c; ENCAPS_SEED_LEN]).unwrap();
let mut broken = ciphertext;
if let Some(byte) = broken.get_mut(0) {
*byte ^= 1;
}
let got = decapsulate_with(&pq_seed, &classical, &broken).unwrap();
assert_ne!(got.as_slice(), sent.as_slice(), "порча ML-KEM не повлияла на секрет");
let other = P256Agreement::from_be_bytes(&[0x42; SCALAR_LEN]).unwrap();
let mut swapped = ciphertext;
if let Some(tail) = swapped.get_mut(ML_KEM_CIPHERTEXT_LEN..) {
tail.copy_from_slice(&other.public_key());
}
let got = decapsulate_with(&pq_seed, &classical, &swapped).unwrap();
assert_ne!(got.as_slice(), sent.as_slice(), "подмена точки не повлияла на секрет");
}
#[test]
fn an_off_curve_point_and_a_foreign_length_are_refused() {
let pair = keypair_from_seed(&[0x31; KEYPAIR_SEED_LEN]).unwrap();
let (pq_seed, classical, public) = (pair.ml_kem_seed, pair.classical, pair.public_key);
let (_, ciphertext) = encapsulate_derand(&public, &[0x5c; ENCAPS_SEED_LEN]).unwrap();
let mut zeroed = ciphertext;
if let Some(tail) = zeroed.get_mut(ML_KEM_CIPHERTEXT_LEN..) {
tail.fill(0);
}
assert!(decapsulate_with(&pq_seed, &classical, &zeroed).is_err(), "нулевая точка принята");
assert!(
decapsulate_with(&pq_seed, &classical, ciphertext.get(..1152).unwrap()).is_err(),
"короткий шифротекст принят"
);
}
#[test]
fn the_wire_lengths_are_the_ones_the_format_promises() {
let public = keypair_from_seed(&[0x01; KEYPAIR_SEED_LEN]).unwrap().public_key;
assert_eq!(public.len(), 1249);
let (_, ciphertext) = encapsulate_derand(&public, &[0x02; ENCAPS_SEED_LEN]).unwrap();
assert_eq!(ciphertext.len(), 1153);
}
#[test]
fn the_label_terminates_the_combiner_input() {
let mut expected = Sha3_256::new();
Digest::update(&mut expected, [1_u8; 32]);
Digest::update(&mut expected, [2_u8; 32]);
Digest::update(&mut expected, [3_u8; 65]);
Digest::update(&mut expected, [4_u8; 65]);
Digest::update(&mut expected, b"MLKEM768-P256");
let expected: [u8; 32] = expected.finalize().into();
assert_eq!(
combine(&[1_u8; 32], &[2_u8; 32], &[3_u8; 65], &[4_u8; 65]).as_slice(),
expected
);
}
#[test]
fn parts_assembled_separately_give_the_same_public_half() {
let pair = keypair_from_seed(&[0x77; KEYPAIR_SEED_LEN]).unwrap();
let from_parts =
public_key_from_parts(&pair.ml_kem_seed, &pair.classical.public_key()).unwrap();
assert_eq!(pair.public_key, from_parts);
}
}