use alloc::vec::Vec;
use super::{
algo::{Direction, PaddingMode, StreamingCipherAlgo},
mode::{aead, block, ctr},
padding::{pkcs7_pad, pkcs7_unpad},
};
use crate::error::{CryptoError, Result};
#[derive(Clone)]
pub struct StreamingCipherCtx {
pub(super) algo: StreamingCipherAlgo,
pub(super) key: Vec<u8>,
pub(super) iv: Vec<u8>,
pub(super) buffer: Vec<u8>,
pub(super) direction: Direction,
pub(super) padding: PaddingMode,
pub(super) stream_offset: usize,
pub(super) aad: Vec<u8>,
pub(super) ciphertext: Vec<u8>,
pub(super) plaintext: Vec<u8>,
pub(super) returned_len: usize,
pub(super) tag_len: usize,
}
impl StreamingCipherCtx {
pub fn new(
algo: StreamingCipherAlgo,
key: &[u8],
iv: &[u8],
direction: Direction,
padding: PaddingMode,
) -> Result<Self> {
let iv = if algo.is_ecb() {
Vec::new()
} else {
iv.to_vec()
};
Ok(Self {
algo,
key: key.to_vec(),
iv,
buffer: Vec::new(),
direction,
padding,
stream_offset: 0,
aad: Vec::new(),
ciphertext: Vec::new(),
plaintext: Vec::new(),
returned_len: 0,
tag_len: 16,
})
}
pub fn new_aead(
algo: StreamingCipherAlgo,
key: &[u8],
nonce: &[u8],
direction: Direction,
tag_len: usize,
) -> Result<Self> {
Ok(Self {
algo,
key: key.to_vec(),
iv: nonce.to_vec(),
buffer: Vec::new(),
direction,
padding: PaddingMode::None,
stream_offset: 0,
aad: Vec::new(),
ciphertext: Vec::new(),
plaintext: Vec::new(),
returned_len: 0,
tag_len,
})
}
pub fn update_aad(&mut self, aad: &[u8]) {
self.aad.extend_from_slice(aad);
}
pub fn payload_started(&self) -> bool {
!self.ciphertext.is_empty() || !self.plaintext.is_empty()
}
pub fn max_update_output_len(&self, input_len: usize) -> usize {
if self.algo.is_aead() {
if self.algo.is_ccm() {
return 0;
}
return input_len;
}
if self.algo.is_ctr() {
return input_len;
}
let bs = self.algo.block_size();
let buffered_len = self.buffer.len().saturating_add(input_len);
let keep = if self.direction.is_decrypting()
&& matches!(self.padding, PaddingMode::Pkcs7)
&& buffered_len >= bs
{
bs
} else {
0
};
buffered_len.saturating_sub(keep) / bs * bs
}
pub fn update(&mut self, data: &[u8]) -> Result<Vec<u8>> {
if self.algo.is_aead() {
return aead::update(self, data);
}
if self.algo.is_ctr() {
return ctr::process(self, data);
}
self.buffer.extend_from_slice(data);
let bs = self.algo.block_size();
let keep = if self.direction.is_decrypting()
&& matches!(self.padding, PaddingMode::Pkcs7)
&& self.buffer.len() >= bs
{
bs
} else {
0
};
let complete = self.buffer.len().saturating_sub(keep) / bs * bs;
if complete == 0 {
return Ok(Vec::new());
}
let chunk: Vec<u8> = self.buffer.drain(..complete).collect();
block::process(self, &chunk)
}
pub fn r#final(&mut self) -> Result<Vec<u8>> {
if self.algo.is_aead() {
return Ok(Vec::new());
}
if self.algo.is_ctr() {
let data = core::mem::take(&mut self.buffer);
return ctr::process(self, &data);
}
let bs = self.algo.block_size();
if self.direction.is_encrypting() {
let mut data = core::mem::take(&mut self.buffer);
if matches!(self.padding, PaddingMode::Pkcs7) {
pkcs7_pad(&mut data, bs)?;
}
block::process(self, &data)
} else {
let data = core::mem::take(&mut self.buffer);
let output = block::process(self, &data)?;
if matches!(self.padding, PaddingMode::Pkcs7) {
pkcs7_unpad(&output, bs)
} else {
Ok(output)
}
}
}
pub fn encrypt_final(mut self) -> Result<(Vec<u8>, Vec<u8>)> {
if self.algo.is_gcm() {
let tag = aead::compute_tag(&self)?;
let ct = core::mem::take(&mut self.ciphertext);
let pending = ct[self.returned_len..].to_vec();
Ok((pending, tag.to_vec()))
} else {
let tag_len = self.tag_len;
let ct_with_tag = aead::one_shot(self)?;
if ct_with_tag.len() < tag_len {
return Err(CryptoError::InvalidLength);
}
let split = ct_with_tag.len() - tag_len;
let tag = ct_with_tag[split..].to_vec();
let ct = ct_with_tag[..split].to_vec();
Ok((ct, tag))
}
}
pub fn encrypt_final_with_input(mut self, input: Option<&[u8]>) -> Result<(Vec<u8>, Vec<u8>)> {
let mut output = Vec::new();
if let Some(input) = input
&& !input.is_empty()
{
output.extend(self.update(input)?);
}
let (tail, tag) = self.encrypt_final()?;
output.extend(tail);
Ok((output, tag))
}
pub fn decrypt_final(mut self, tag: &[u8]) -> Result<Vec<u8>> {
if self.algo.is_gcm() {
let expected_tag = aead::compute_tag(&self)?;
use subtle::ConstantTimeEq;
if expected_tag.ct_eq(tag).into() {
let pt = core::mem::take(&mut self.plaintext);
Ok(pt[self.returned_len..].to_vec())
} else {
Err(CryptoError::VerificationFailed)
}
} else {
self.buffer.extend_from_slice(tag);
aead::one_shot(self)
}
}
pub fn decrypt_final_with_input(mut self, input: Option<&[u8]>, tag: &[u8]) -> Result<Vec<u8>> {
let mut output = Vec::new();
if let Some(input) = input
&& !input.is_empty()
{
output.extend(self.update(input)?);
}
output.extend(self.decrypt_final(tag)?);
Ok(output)
}
}