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::vec::Vec;
use zeroize::Zeroizing;
fn chunk_nonce(prefix: &[u8; STREAM_NONCE_PREFIX_LEN], index: u32, last: bool) -> [u8; NONCE_LEN] {
let mut nonce = [0u8; NONCE_LEN];
nonce[..STREAM_NONCE_PREFIX_LEN].copy_from_slice(prefix);
nonce[STREAM_NONCE_PREFIX_LEN..STREAM_NONCE_PREFIX_LEN + 4]
.copy_from_slice(&index.to_be_bytes());
nonce[NONCE_LEN - 1] = last as u8;
nonce
}
fn chunk_aad(header: &[u8], index: u32, last: bool) -> Vec<u8> {
let mut aad = Vec::with_capacity(header.len() + 5);
aad.extend_from_slice(header);
aad.extend_from_slice(&index.to_be_bytes());
aad.push(last as u8);
aad
}
pub struct StreamSealer {
cipher: Aes256Gcm,
nonce_prefix: [u8; STREAM_NONCE_PREFIX_LEN],
header: Vec<u8>,
index: u32,
finished: bool,
}
impl StreamSealer {
pub fn new(recipient: &PublicKeyBundle) -> Result<(Self, Vec<u8>)> {
let (kem_ct, ss) = hybrid_kem::encapsulate(recipient)?;
let mut nonce_prefix = [0u8; STREAM_NONCE_PREFIX_LEN];
getrandom::fill(&mut nonce_prefix).map_err(|_| Error::RandomnessUnavailable)?;
let mut header = Vec::with_capacity(STREAM_HEADER_LEN);
write_header(&mut header, MAGIC_STREAM);
header.extend_from_slice(&kem_ct.epk_x25519);
header.extend_from_slice(kem_ct.ct_mlkem.as_ref());
header.extend_from_slice(&nonce_prefix);
let cipher = Aes256Gcm::new((&*ss).into());
Ok((
Self {
cipher,
nonce_prefix,
header: header.clone(),
index: 0,
finished: false,
},
header,
))
}
pub fn seal_chunk(&mut self, plaintext: &[u8], last: bool) -> Result<Vec<u8>> {
if self.finished {
return Err(Error::StreamFinished);
}
if !last && self.index == u32::MAX {
self.finished = true;
return Err(Error::StreamFinished);
}
if plaintext.len() > (u32::MAX as usize - TAG_LEN) {
return Err(Error::MessageTooLarge {
len: plaintext.len(),
max: u32::MAX as usize - TAG_LEN,
});
}
let nonce = chunk_nonce(&self.nonce_prefix, self.index, last);
let aad = chunk_aad(&self.header, self.index, last);
let ct = self
.cipher
.encrypt(
(&nonce).into(),
Payload {
msg: plaintext,
aad: &aad,
},
)
.map_err(|_| Error::MessageTooLarge {
len: plaintext.len(),
max: u32::MAX as usize - TAG_LEN,
})?;
let mut frame = Vec::with_capacity(5 + ct.len());
frame.push(last as u8);
frame.extend_from_slice(&(ct.len() as u32).to_be_bytes());
frame.extend_from_slice(&ct);
if last {
self.finished = true;
} else {
self.index += 1;
}
Ok(frame)
}
}
pub struct StreamOpener {
cipher: Aes256Gcm,
nonce_prefix: [u8; STREAM_NONCE_PREFIX_LEN],
header: Vec<u8>,
index: u32,
finished: bool,
}
impl StreamOpener {
pub fn new(keypair: &KeyPair, header_bytes: &[u8]) -> Result<Self> {
let mut rest = read_header(header_bytes, MAGIC_STREAM, Error::InvalidEnvelope)?;
let epk_x25519: [u8; X25519_PK_LEN] = take(&mut rest, Error::InvalidEnvelope)?;
let ct_mlkem: [u8; MLKEM1024_CT_LEN] = take(&mut rest, Error::InvalidEnvelope)?;
let nonce_prefix: [u8; STREAM_NONCE_PREFIX_LEN] = take(&mut rest, Error::InvalidEnvelope)?;
if !rest.is_empty() {
return Err(Error::InvalidEnvelope);
}
let kem_ct = KemCiphertext {
epk_x25519,
ct_mlkem: alloc::boxed::Box::new(ct_mlkem),
};
let ss: Zeroizing<[u8; 32]> = hybrid_kem::decapsulate(keypair, &kem_ct);
let cipher = Aes256Gcm::new((&*ss).into());
Ok(Self {
cipher,
nonce_prefix,
header: header_bytes.to_vec(),
index: 0,
finished: false,
})
}
pub fn open_chunk(&mut self, frame: &[u8]) -> Result<(Vec<u8>, bool)> {
if self.finished {
return Err(Error::StreamFinished);
}
if frame.len() < 5 {
return Err(Error::DecryptionFailed);
}
let last = match frame[0] {
0 => false,
1 => true,
_ => return Err(Error::DecryptionFailed),
};
let ct_len = u32::from_be_bytes([frame[1], frame[2], frame[3], frame[4]]) as usize;
let ct = &frame[5..];
if ct.len() != ct_len || ct_len < TAG_LEN {
return Err(Error::DecryptionFailed);
}
let nonce = chunk_nonce(&self.nonce_prefix, self.index, last);
let aad = chunk_aad(&self.header, self.index, last);
let plaintext = self
.cipher
.decrypt((&nonce).into(), Payload { msg: ct, aad: &aad })
.map_err(|_| Error::DecryptionFailed)?;
if last {
self.finished = true;
} else {
self.index = self.index.checked_add(1).ok_or(Error::DecryptionFailed)?;
}
Ok((plaintext, last))
}
pub fn finish(self) -> Result<()> {
if self.finished {
Ok(())
} else {
Err(Error::StreamTruncated)
}
}
}