use zeroize::Zeroize;
use crate::{Error, LengthRequirement, Parameters, Result, SEGMENT_PAYLOAD_OFFSET};
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum BufferState {
Empty,
Plaintext { 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 { 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]> {
self.parameters.validate_ciphertext_segment_length(length)?;
self.state = BufferState::Ciphertext { length };
Ok(&mut self.bytes[..length])
}
pub fn plaintext(&self) -> Result<&[u8]> {
match self.state {
BufferState::Plaintext { length } => {
Ok(&self.bytes[SEGMENT_PAYLOAD_OFFSET..SEGMENT_PAYLOAD_OFFSET + length])
}
BufferState::Empty | BufferState::Ciphertext { .. } => Err(Error::InvalidBufferState),
}
}
pub(crate) fn plaintext_mut(&mut self) -> Result<&mut [u8]> {
match self.state {
BufferState::Plaintext { length } => {
Ok(&mut self.bytes[SEGMENT_PAYLOAD_OFFSET..SEGMENT_PAYLOAD_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 extend_plaintext(&mut self, input: &[u8]) -> usize {
let length = match self.state {
BufferState::Plaintext { length } => length,
BufferState::Empty | BufferState::Ciphertext { .. } => 0,
};
let capacity = self.parameters.plaintext_segment_length();
let copied = input.len().min(capacity - length);
let start = SEGMENT_PAYLOAD_OFFSET + length;
self.bytes[start..start + copied].copy_from_slice(&input[..copied]);
self.state = BufferState::Plaintext {
length: length + copied,
};
copied
}
pub(crate) fn truncate_plaintext(&mut self, length: usize) -> Result<()> {
let current = self.plaintext_length()?;
if length > current {
return Err(Error::InvalidPlaintextLength {
actual: length,
required: LengthRequirement::AtMost(current),
});
}
self.state = BufferState::Plaintext { length };
Ok(())
}
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, length: usize) {
self.state = BufferState::Plaintext { length };
}
}
impl Drop for SegmentBuffer {
fn drop(&mut self) {
self.bytes.zeroize();
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::SEGMENT_OVERHEAD;
#[test]
fn extend_plaintext_appends_and_saturates_at_capacity() {
let parameters = Parameters::SEGMENT_4_KIB;
let capacity = parameters.plaintext_segment_length();
let mut buffer = SegmentBuffer::new(parameters);
assert_eq!(buffer.extend_plaintext(b"abc"), 3);
assert_eq!(buffer.extend_plaintext(b"de"), 2);
assert_eq!(buffer.plaintext().unwrap(), b"abcde");
assert_eq!(buffer.plaintext_length().unwrap(), 5);
let oversized = vec![0x5a; capacity];
assert_eq!(buffer.extend_plaintext(&oversized), capacity - 5);
assert_eq!(buffer.plaintext_length().unwrap(), capacity);
assert_eq!(buffer.extend_plaintext(b"x"), 0);
assert_eq!(buffer.plaintext_length().unwrap(), capacity);
}
#[test]
fn extend_plaintext_starts_fresh_after_ciphertext_or_clear() {
let parameters = Parameters::SEGMENT_4_KIB;
let mut buffer = SegmentBuffer::new(parameters);
buffer.prepare_ciphertext(SEGMENT_OVERHEAD).unwrap();
assert_eq!(buffer.extend_plaintext(b"fresh"), 5);
assert_eq!(buffer.plaintext().unwrap(), b"fresh");
buffer.clear();
assert_eq!(buffer.extend_plaintext(b"again"), 5);
assert_eq!(buffer.plaintext().unwrap(), b"again");
}
#[test]
fn truncate_plaintext_shrinks_but_never_grows() {
let parameters = Parameters::SEGMENT_4_KIB;
let mut buffer = SegmentBuffer::new(parameters);
buffer.extend_plaintext(b"abcde");
buffer.truncate_plaintext(3).unwrap();
assert_eq!(buffer.plaintext().unwrap(), b"abc");
assert!(matches!(
buffer.truncate_plaintext(4),
Err(Error::InvalidPlaintextLength { .. })
));
buffer.clear();
assert_eq!(buffer.truncate_plaintext(0), Err(Error::InvalidBufferState));
}
#[test]
fn prepare_plaintext_rejects_length_above_capacity() {
let parameters = Parameters::SEGMENT_4_KIB;
let capacity = parameters.plaintext_segment_length();
let mut buffer = SegmentBuffer::new(parameters);
buffer.prepare_plaintext(3).unwrap().copy_from_slice(b"abc");
assert!(matches!(
buffer.prepare_plaintext(capacity + 1),
Err(Error::InvalidPlaintextLength {
actual,
required: LengthRequirement::AtMost(maximum),
}) if actual == capacity + 1 && maximum == capacity
));
assert_eq!(buffer.plaintext().unwrap(), b"abc");
}
#[test]
fn prepare_ciphertext_rejects_lengths_outside_segment_bounds() {
let parameters = Parameters::SEGMENT_4_KIB;
let maximum = parameters.ciphertext_segment_length();
let mut buffer = SegmentBuffer::new(parameters);
assert_eq!(
buffer.prepare_ciphertext(SEGMENT_OVERHEAD).unwrap().len(),
SEGMENT_OVERHEAD
);
for length in [SEGMENT_OVERHEAD - 1, maximum + 1] {
assert!(matches!(
buffer.prepare_ciphertext(length),
Err(Error::InvalidCiphertextLength {
actual,
required: LengthRequirement::Between {
minimum: SEGMENT_OVERHEAD,
maximum: reported,
},
}) if actual == length && reported == maximum
));
assert_eq!(buffer.ciphertext().unwrap().len(), SEGMENT_OVERHEAD);
}
assert_eq!(buffer.prepare_ciphertext(maximum).unwrap().len(), maximum);
}
#[test]
fn plaintext_rejects_empty_and_ciphertext_states() {
let mut buffer = SegmentBuffer::new(Parameters::SEGMENT_4_KIB);
assert_eq!(buffer.plaintext(), Err(Error::InvalidBufferState));
buffer.prepare_ciphertext(SEGMENT_OVERHEAD + 3).unwrap();
assert_eq!(buffer.plaintext(), Err(Error::InvalidBufferState));
}
#[test]
fn ciphertext_rejects_empty_and_plaintext_states() {
let mut buffer = SegmentBuffer::new(Parameters::SEGMENT_4_KIB);
assert_eq!(buffer.ciphertext(), Err(Error::InvalidBufferState));
buffer.prepare_plaintext(3).unwrap().copy_from_slice(b"abc");
assert_eq!(buffer.ciphertext(), Err(Error::InvalidBufferState));
}
#[test]
fn capacity_matches_parameters() {
for parameters in [
Parameters::SEGMENT_64_B,
Parameters::SEGMENT_4_KIB,
Parameters::SEGMENT_1_MIB,
] {
let buffer = SegmentBuffer::new(parameters);
assert_eq!(buffer.capacity(), parameters.ciphertext_segment_length());
assert_eq!(buffer.parameters(), parameters);
}
}
}