use crate::constants::*;
use crate::error::{Error, Result};
use crate::hybrid_kem::{self, KemCiphertext};
use crate::keys::{KeyPair, PublicKeyBundle};
use crate::wire::{read_header, take, write_header};
use aes_gcm::aead::{Aead, Payload};
use aes_gcm::{Aes256Gcm, KeyInit};
use alloc::boxed::Box;
use alloc::vec::Vec;
use sha3::{Digest, Sha3_256};
use subtle::ConstantTimeEq;
use zeroize::Zeroizing;
fn cek_commitment(cek: &[u8; CEK_LEN]) -> [u8; CEK_COMMIT_LEN] {
let mut hasher = Sha3_256::new();
hasher.update(MULTI_CEK_COMMIT_LABEL);
hasher.update(cek);
hasher.finalize().into()
}
#[derive(Clone, Debug, PartialEq, Eq)]
struct Wrap {
epk_x25519: [u8; X25519_PK_LEN],
ct_mlkem: Box<[u8; MLKEM1024_CT_LEN]>,
wrap_nonce: [u8; NONCE_LEN],
wrapped_cek: [u8; CEK_LEN + TAG_LEN],
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct MultiRecipientEnvelope {
cek_commitment: [u8; CEK_COMMIT_LEN],
wraps: Vec<Wrap>,
payload_nonce: [u8; NONCE_LEN],
payload_ct: Vec<u8>,
}
impl MultiRecipientEnvelope {
pub fn recipient_count(&self) -> usize {
self.wraps.len()
}
fn write_prefix(&self, out: &mut Vec<u8>) {
debug_assert!(self.wraps.len() <= MAX_RECIPIENTS);
write_header(out, MAGIC_MULTI);
out.extend_from_slice(&(self.wraps.len() as u16).to_be_bytes());
out.extend_from_slice(&self.cek_commitment);
for wrap in &self.wraps {
out.extend_from_slice(&wrap.epk_x25519);
out.extend_from_slice(wrap.ct_mlkem.as_ref());
out.extend_from_slice(&wrap.wrap_nonce);
out.extend_from_slice(&wrap.wrapped_cek);
}
out.extend_from_slice(&self.payload_nonce);
}
pub fn to_bytes(&self) -> Vec<u8> {
let mut out = Vec::with_capacity(
HEADER_LEN
+ 2
+ CEK_COMMIT_LEN
+ self.wraps.len() * WRAP_LEN
+ NONCE_LEN
+ self.payload_ct.len(),
);
self.write_prefix(&mut out);
out.extend_from_slice(&self.payload_ct);
out
}
pub fn from_bytes(bytes: &[u8]) -> Result<Self> {
let mut rest = read_header(bytes, MAGIC_MULTI, Error::InvalidEnvelope)?;
let count_bytes: [u8; 2] = take(&mut rest, Error::InvalidEnvelope)?;
let count = u16::from_be_bytes(count_bytes) as usize;
if count == 0 {
return Err(Error::NoRecipients);
}
if count > MAX_RECIPIENTS {
return Err(Error::TooManyRecipients {
count,
max: MAX_RECIPIENTS,
});
}
let cek_commitment = take(&mut rest, Error::InvalidEnvelope)?;
let mut wraps = Vec::with_capacity(count);
for _ in 0..count {
let epk_x25519 = take(&mut rest, Error::InvalidEnvelope)?;
let ct_mlkem: [u8; MLKEM1024_CT_LEN] = take(&mut rest, Error::InvalidEnvelope)?;
let wrap_nonce = take(&mut rest, Error::InvalidEnvelope)?;
let wrapped_cek = take(&mut rest, Error::InvalidEnvelope)?;
wraps.push(Wrap {
epk_x25519,
ct_mlkem: Box::new(ct_mlkem),
wrap_nonce,
wrapped_cek,
});
}
let payload_nonce = take(&mut rest, Error::InvalidEnvelope)?;
if rest.len() < TAG_LEN {
return Err(Error::InvalidEnvelope);
}
Ok(Self {
cek_commitment,
wraps,
payload_nonce,
payload_ct: rest.to_vec(),
})
}
}
fn wrap_aad(count: usize) -> Vec<u8> {
let mut aad = Vec::with_capacity(HEADER_LEN + 2);
write_header(&mut aad, MAGIC_MULTI);
aad.extend_from_slice(&(count as u16).to_be_bytes());
aad
}
pub fn seal_multi(
plaintext: &[u8],
recipients: &[&PublicKeyBundle],
) -> Result<MultiRecipientEnvelope> {
if recipients.is_empty() {
return Err(Error::NoRecipients);
}
if recipients.len() > MAX_RECIPIENTS {
return Err(Error::TooManyRecipients {
count: recipients.len(),
max: MAX_RECIPIENTS,
});
}
if plaintext.len() > MAX_PLAINTEXT_LEN {
return Err(Error::MessageTooLarge {
len: plaintext.len(),
max: MAX_PLAINTEXT_LEN,
});
}
let mut cek = Zeroizing::new([0u8; CEK_LEN]);
getrandom::fill(cek.as_mut()).map_err(|_| Error::RandomnessUnavailable)?;
let aad = wrap_aad(recipients.len());
let mut wraps = Vec::with_capacity(recipients.len());
for recipient in recipients {
let (kem_ct, ss) = hybrid_kem::encapsulate(recipient)?;
let cipher = Aes256Gcm::new((&*ss).into());
let mut wrap_nonce = [0u8; NONCE_LEN];
getrandom::fill(&mut wrap_nonce).map_err(|_| Error::RandomnessUnavailable)?;
let wrapped = cipher
.encrypt(
(&wrap_nonce).into(),
Payload {
msg: &*cek,
aad: &aad,
},
)
.expect("AES-GCM wrap of a 32-byte CEK is infallible");
let wrapped_cek: [u8; CEK_LEN + TAG_LEN] = wrapped
.try_into()
.expect("AES-256-GCM output is plaintext length + 16-byte tag");
wraps.push(Wrap {
epk_x25519: kem_ct.epk_x25519,
ct_mlkem: kem_ct.ct_mlkem,
wrap_nonce,
wrapped_cek,
});
}
let mut payload_nonce = [0u8; NONCE_LEN];
getrandom::fill(&mut payload_nonce).map_err(|_| Error::RandomnessUnavailable)?;
let mut envelope = MultiRecipientEnvelope {
cek_commitment: cek_commitment(&cek),
wraps,
payload_nonce,
payload_ct: Vec::new(),
};
let mut payload_aad = Vec::new();
envelope.write_prefix(&mut payload_aad);
let cipher = Aes256Gcm::new((&*cek).into());
envelope.payload_ct = cipher
.encrypt(
(&payload_nonce).into(),
Payload {
msg: plaintext,
aad: &payload_aad,
},
)
.map_err(|_| Error::MessageTooLarge {
len: plaintext.len(),
max: MAX_PLAINTEXT_LEN,
})?;
Ok(envelope)
}
pub fn open_multi(keypair: &KeyPair, envelope: &MultiRecipientEnvelope) -> Result<Vec<u8>> {
let aad = wrap_aad(envelope.wraps.len());
let mut payload_aad = Vec::new();
envelope.write_prefix(&mut payload_aad);
for wrap in &envelope.wraps {
let kem_ct = KemCiphertext {
epk_x25519: wrap.epk_x25519,
ct_mlkem: wrap.ct_mlkem.clone(),
};
let ss = hybrid_kem::decapsulate(keypair, &kem_ct);
let cipher = Aes256Gcm::new((&*ss).into());
let Ok(cek_vec) = cipher.decrypt(
(&wrap.wrap_nonce).into(),
Payload {
msg: &wrap.wrapped_cek,
aad: &aad,
},
) else {
continue;
};
let cek_vec = Zeroizing::new(cek_vec);
let cek: Zeroizing<[u8; CEK_LEN]> = Zeroizing::new(
cek_vec
.as_slice()
.try_into()
.map_err(|_| Error::DecryptionFailed)?,
);
let commit_ok: bool = cek_commitment(&cek).ct_eq(&envelope.cek_commitment).into();
if !commit_ok {
return Err(Error::DecryptionFailed);
}
let payload_cipher = Aes256Gcm::new((&*cek).into());
return payload_cipher
.decrypt(
(&envelope.payload_nonce).into(),
Payload {
msg: &envelope.payload_ct,
aad: &payload_aad,
},
)
.map_err(|_| Error::DecryptionFailed);
}
Err(Error::DecryptionFailed)
}