use zeroize::Zeroize;
use crate::{
Error, LengthRequirement, Parameters, Result, SEGMENT_OVERHEAD, SEGMENT_PAYLOAD_OFFSET,
};
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum BufferState {
Empty,
Plaintext { offset: usize, length: usize },
Ciphertext { length: usize },
}
#[derive(Debug)]
pub struct SegmentBuffer {
parameters: Parameters,
bytes: Vec<u8>,
state: BufferState,
}
impl SegmentBuffer {
#[must_use]
pub fn new(parameters: Parameters) -> Self {
Self {
parameters,
bytes: vec![0u8; parameters.ciphertext_segment_length()],
state: BufferState::Empty,
}
}
#[must_use]
pub const fn parameters(&self) -> Parameters {
self.parameters
}
#[must_use]
pub const fn capacity(&self) -> usize {
self.bytes.len()
}
pub fn prepare_plaintext(&mut self, length: usize) -> Result<&mut [u8]> {
let maximum = self.parameters.plaintext_segment_length();
if length > maximum {
return Err(Error::InvalidPlaintextLength {
actual: length,
required: LengthRequirement::AtMost(maximum),
});
}
self.state = BufferState::Plaintext {
offset: SEGMENT_PAYLOAD_OFFSET,
length,
};
let end = SEGMENT_PAYLOAD_OFFSET + length;
Ok(&mut self.bytes[SEGMENT_PAYLOAD_OFFSET..end])
}
pub fn prepare_ciphertext(&mut self, length: usize) -> Result<&mut [u8]> {
let maximum = self.parameters.ciphertext_segment_length();
if !(SEGMENT_OVERHEAD..=maximum).contains(&length) {
return Err(Error::InvalidCiphertextLength {
actual: length,
required: LengthRequirement::Between {
minimum: SEGMENT_OVERHEAD,
maximum,
},
});
}
self.state = BufferState::Ciphertext { length };
Ok(&mut self.bytes[..length])
}
pub fn plaintext(&self) -> Result<&[u8]> {
match self.state {
BufferState::Plaintext { offset, length } => Ok(&self.bytes[offset..offset + length]),
BufferState::Empty | BufferState::Ciphertext { .. } => Err(Error::InvalidBufferState),
}
}
pub fn ciphertext(&self) -> Result<&[u8]> {
match self.state {
BufferState::Ciphertext { length } => Ok(&self.bytes[..length]),
BufferState::Empty | BufferState::Plaintext { .. } => Err(Error::InvalidBufferState),
}
}
pub fn clear(&mut self) {
self.bytes.as_mut_slice().zeroize();
self.state = BufferState::Empty;
}
pub(crate) const fn plaintext_length(&self) -> Result<usize> {
match self.state {
BufferState::Plaintext { length, .. } => Ok(length),
BufferState::Empty | BufferState::Ciphertext { .. } => Err(Error::InvalidBufferState),
}
}
pub(crate) const fn ciphertext_length(&self) -> Result<usize> {
match self.state {
BufferState::Ciphertext { length } => Ok(length),
BufferState::Empty | BufferState::Plaintext { .. } => Err(Error::InvalidBufferState),
}
}
pub(crate) fn raw_mut(&mut self) -> &mut [u8] {
&mut self.bytes
}
pub(crate) fn matches(&self, parameters: Parameters) -> bool {
self.parameters == parameters
}
pub(crate) const fn mark_ciphertext(&mut self, length: usize) {
self.state = BufferState::Ciphertext { length };
}
pub(crate) const fn mark_plaintext(&mut self, offset: usize, length: usize) {
self.state = BufferState::Plaintext { offset, length };
}
pub(crate) const fn mark_empty(&mut self) {
self.state = BufferState::Empty;
}
}
impl Drop for SegmentBuffer {
fn drop(&mut self) {
self.bytes.zeroize();
}
}