use std::io::{self, Read, Write};
use crate::online::{Decryptor, Encryptor};
use crate::wire::{decryption_error, encryption_error, read_fully, read_header};
use crate::{Error, Header, Key, LengthRequirement, Parameters, SegmentBuffer, SegmentFraming};
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum StreamStatus {
Open,
Finished,
Failed,
}
pub struct FinishError<W> {
inner: W,
error: io::Error,
}
impl<W> FinishError<W> {
#[must_use]
pub const fn error(&self) -> &io::Error {
&self.error
}
#[must_use]
pub fn into_inner(self) -> W {
self.inner
}
#[must_use]
pub fn into_parts(self) -> (io::Error, W) {
(self.error, self.inner)
}
fn wrap(result: io::Result<()>, inner: W) -> Result<W, FinishError<W>> {
match result {
Ok(()) => Ok(inner),
Err(error) => Err(Self { inner, error }),
}
}
}
impl<W> core::fmt::Display for FinishError<W> {
fn fmt(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
write!(formatter, "failed to finish FLOE stream: {}", self.error)
}
}
impl<W> core::fmt::Debug for FinishError<W> {
fn fmt(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
formatter
.debug_struct("FinishError")
.field("error", &self.error)
.finish_non_exhaustive()
}
}
impl<W> std::error::Error for FinishError<W> {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
Some(&self.error)
}
}
#[derive(Debug)]
pub struct EncryptWriter<W> {
inner: W,
encryptor: Option<Encryptor>,
header: Header,
buffer: SegmentBuffer,
status: StreamStatus,
}
impl<W: Write> EncryptWriter<W> {
pub fn new(mut inner: W, key: &Key, aad: &[u8], parameters: Parameters) -> io::Result<Self> {
let encryptor = Encryptor::new(key, aad, parameters).map_err(encryption_error)?;
let header = *encryptor.header();
inner.write_all(header.as_ref())?;
Ok(Self {
inner,
encryptor: Some(encryptor),
header,
buffer: SegmentBuffer::new(parameters),
status: StreamStatus::Open,
})
}
#[must_use]
pub const fn header(&self) -> &Header {
&self.header
}
#[must_use]
pub const fn parameters(&self) -> Parameters {
self.buffer.parameters()
}
#[must_use]
pub const fn get_ref(&self) -> &W {
&self.inner
}
pub const fn get_mut(&mut self) -> &mut W {
&mut self.inner
}
#[must_use]
pub const fn is_finished(&self) -> bool {
matches!(self.status, StreamStatus::Finished)
}
pub fn try_finish(&mut self) -> io::Result<()> {
match self.status {
StreamStatus::Finished => return Ok(()),
StreamStatus::Failed => return Err(poisoned_error()),
StreamStatus::Open => {}
}
self.poison_on_error(Self::finish_message)
}
pub fn finish(mut self) -> Result<W, FinishError<W>> {
let result = self.try_finish();
FinishError::wrap(result, self.inner)
}
#[must_use]
pub fn into_inner_unfinished(self) -> W {
self.inner
}
fn poison_on_error(&mut self, op: impl FnOnce(&mut Self) -> io::Result<()>) -> io::Result<()> {
let result = op(self);
if result.is_err() {
self.status = StreamStatus::Failed;
}
result
}
fn finish_message(&mut self) -> io::Result<()> {
let encryptor = self.encryptor.take().ok_or_else(poisoned_error)?;
if self.buffer.plaintext_length().is_err() {
self.buffer.prepare_plaintext(0).map_err(encryption_error)?;
}
encryptor
.encrypt_final_segment_in_place(&mut self.buffer)
.map_err(|error| encryption_error(error.into_error()))?;
self.inner
.write_all(self.buffer.ciphertext().map_err(encryption_error)?)?;
self.inner.flush()?;
self.status = StreamStatus::Finished;
Ok(())
}
fn emit_non_final(&mut self) -> io::Result<()> {
self.poison_on_error(|writer| {
writer
.encryptor
.as_mut()
.ok_or_else(poisoned_error)?
.encrypt_non_final_segment_in_place(&mut writer.buffer)
.map_err(encryption_error)?;
writer
.inner
.write_all(writer.buffer.ciphertext().map_err(encryption_error)?)?;
Ok(())
})
}
fn write_segment_direct(&mut self, chunk: &[u8]) -> io::Result<usize> {
self.poison_on_error(|writer| {
let written = writer
.encryptor
.as_mut()
.ok_or_else(poisoned_error)?
.encrypt_non_final_segment_into(chunk, writer.buffer.raw_mut())
.map_err(encryption_error)?;
writer.buffer.mark_ciphertext(written);
writer
.inner
.write_all(writer.buffer.ciphertext().map_err(encryption_error)?)?;
Ok(())
})?;
Ok(chunk.len())
}
}
impl<W: Write> Write for EncryptWriter<W> {
fn write(&mut self, input: &[u8]) -> io::Result<usize> {
if input.is_empty() {
return Ok(0);
}
match self.status {
StreamStatus::Open => {}
StreamStatus::Finished => {
return Err(io::Error::new(
io::ErrorKind::BrokenPipe,
"FLOE encrypting writer is already finished",
));
}
StreamStatus::Failed => return Err(poisoned_error()),
}
let capacity = self.parameters().plaintext_segment_length();
if input.len() >= capacity && self.buffer.plaintext_length().is_err() {
return self.write_segment_direct(&input[..capacity]);
}
let copied = self.buffer.extend_plaintext(input);
if copied > 0 {
return Ok(copied);
}
self.emit_non_final()?;
if input.len() >= capacity {
return self.write_segment_direct(&input[..capacity]);
}
Ok(self.buffer.extend_plaintext(input))
}
fn flush(&mut self) -> io::Result<()> {
if self.status == StreamStatus::Failed {
return Err(poisoned_error());
}
self.inner.flush()
}
}
#[derive(Debug)]
pub struct EncryptReader<R> {
inner: R,
encryptor: Option<Encryptor>,
header: Header,
header_position: usize,
buffer: SegmentBuffer,
ciphertext_position: usize,
ciphertext_length: usize,
lookahead: Option<u8>,
status: StreamStatus,
}
impl<R: Read> EncryptReader<R> {
pub fn new(inner: R, key: &Key, aad: &[u8], parameters: Parameters) -> io::Result<Self> {
let encryptor = Encryptor::new(key, aad, parameters).map_err(encryption_error)?;
let header = *encryptor.header();
Ok(Self {
inner,
encryptor: Some(encryptor),
header,
header_position: 0,
buffer: SegmentBuffer::new(parameters),
ciphertext_position: 0,
ciphertext_length: 0,
lookahead: None,
status: StreamStatus::Open,
})
}
#[must_use]
pub const fn header(&self) -> &Header {
&self.header
}
#[must_use]
pub const fn parameters(&self) -> Parameters {
self.buffer.parameters()
}
#[must_use]
pub const fn get_ref(&self) -> &R {
&self.inner
}
pub const fn get_mut(&mut self) -> &mut R {
&mut self.inner
}
#[must_use]
pub const fn is_finished(&self) -> bool {
matches!(self.status, StreamStatus::Finished)
&& self.ciphertext_position == self.ciphertext_length
}
#[must_use]
pub fn into_inner_unfinished(self) -> R {
self.inner
}
fn load_segment(&mut self) -> io::Result<()> {
let result = self.try_load_segment();
if result.is_err() {
self.status = StreamStatus::Failed;
}
result
}
fn try_load_segment(&mut self) -> io::Result<()> {
let is_final = self.fill_plaintext()?;
self.ciphertext_position = 0;
self.ciphertext_length = if is_final {
let encryptor = self.encryptor.take().ok_or_else(poisoned_error)?;
match encryptor.encrypt_final_segment_in_place(&mut self.buffer) {
Ok(encrypted) => {
self.status = StreamStatus::Finished;
encrypted.len()
}
Err(error) => return Err(encryption_error(error.into_error())),
}
} else {
self.encryptor
.as_mut()
.ok_or_else(poisoned_error)?
.encrypt_non_final_segment_in_place(&mut self.buffer)
.map_err(encryption_error)?
.len()
};
Ok(())
}
fn fill_plaintext(&mut self) -> io::Result<bool> {
let capacity = self.parameters().plaintext_segment_length();
let plaintext = self
.buffer
.prepare_plaintext(capacity)
.map_err(encryption_error)?;
let mut plaintext_length = 0;
if let Some(byte) = self.lookahead.take() {
plaintext[0] = byte;
plaintext_length = 1;
}
plaintext_length +=
read_fully(&mut self.inner, &mut plaintext[plaintext_length..capacity])?;
let is_final = if plaintext_length < capacity {
true
} else {
let mut trailing = [0; 1];
if read_fully(&mut self.inner, &mut trailing)? == 0 {
true
} else {
self.lookahead = Some(trailing[0]);
false
}
};
self.buffer
.truncate_plaintext(plaintext_length)
.map_err(encryption_error)?;
Ok(is_final)
}
fn emit_segment_direct(&mut self, output: &mut [u8]) -> io::Result<usize> {
let result = self.try_emit_segment_direct(output);
if result.is_err() {
self.status = StreamStatus::Failed;
}
result
}
fn try_emit_segment_direct(&mut self, output: &mut [u8]) -> io::Result<usize> {
let is_final = self.fill_plaintext()?;
if is_final {
let encryptor = self.encryptor.take().ok_or_else(poisoned_error)?;
let written = encryptor
.encrypt_final_segment_into(
self.buffer.plaintext().map_err(encryption_error)?,
output,
)
.map_err(|error| encryption_error(error.into_error()))?;
self.status = StreamStatus::Finished;
Ok(written)
} else {
self.encryptor
.as_mut()
.ok_or_else(poisoned_error)?
.encrypt_non_final_segment_into(
self.buffer.plaintext().map_err(encryption_error)?,
output,
)
.map_err(encryption_error)
}
}
}
impl<R: Read> Read for EncryptReader<R> {
fn read(&mut self, output: &mut [u8]) -> io::Result<usize> {
if output.is_empty() {
return Ok(0);
}
if self.status == StreamStatus::Failed {
return Err(poisoned_error());
}
if self.header_position < Header::LEN {
let read = output.len().min(Header::LEN - self.header_position);
let end = self.header_position + read;
output[..read].copy_from_slice(&self.header.as_bytes()[self.header_position..end]);
self.header_position = end;
return Ok(read);
}
if self.ciphertext_position == self.ciphertext_length {
if self.status == StreamStatus::Finished {
return Ok(0);
}
if output.len() >= self.parameters().ciphertext_segment_length() {
return self.emit_segment_direct(output);
}
self.load_segment()?;
}
let ciphertext = self.buffer.ciphertext().map_err(encryption_error)?;
let read = output
.len()
.min(self.ciphertext_length - self.ciphertext_position);
let end = self.ciphertext_position + read;
output[..read].copy_from_slice(&ciphertext[self.ciphertext_position..end]);
self.ciphertext_position = end;
Ok(read)
}
}
#[derive(Debug)]
pub struct DecryptReader<R> {
inner: R,
decryptor: Option<Decryptor>,
header: Header,
buffer: SegmentBuffer,
plaintext_consumed: usize,
plaintext_available: usize,
status: StreamStatus,
}
impl<R: Read> DecryptReader<R> {
pub fn new(mut inner: R, key: &Key, aad: &[u8]) -> io::Result<Self> {
let header = read_header(&mut inner)?;
let decryptor = Decryptor::new(key, aad, &header).map_err(decryption_error)?;
let parameters = decryptor.parameters();
Ok(Self {
inner,
decryptor: Some(decryptor),
header,
buffer: SegmentBuffer::new(parameters),
plaintext_consumed: 0,
plaintext_available: 0,
status: StreamStatus::Open,
})
}
#[must_use]
pub const fn header(&self) -> &Header {
&self.header
}
#[must_use]
pub const fn parameters(&self) -> Parameters {
self.buffer.parameters()
}
#[must_use]
pub const fn get_ref(&self) -> &R {
&self.inner
}
pub const fn get_mut(&mut self) -> &mut R {
&mut self.inner
}
#[must_use]
pub const fn is_finished(&self) -> bool {
matches!(self.status, StreamStatus::Finished)
&& self.plaintext_consumed == self.plaintext_available
}
pub fn try_finish(&mut self) -> io::Result<()> {
self.drain_message()?;
let mut trailing = [0; 1];
match read_fully(&mut self.inner, &mut trailing) {
Ok(0) => Ok(()),
Ok(_) => {
self.status = StreamStatus::Failed;
Err(io::Error::new(
io::ErrorKind::InvalidData,
"FLOE ciphertext has data after its final segment",
))
}
Err(error) => {
self.status = StreamStatus::Failed;
Err(error)
}
}
}
pub fn try_finish_frame(&mut self) -> io::Result<()> {
self.drain_message()
}
pub fn finish(mut self) -> Result<R, FinishError<R>> {
let result = self.try_finish();
FinishError::wrap(result, self.inner)
}
pub fn finish_frame(mut self) -> Result<R, FinishError<R>> {
let result = self.try_finish_frame();
FinishError::wrap(result, self.inner)
}
#[must_use]
pub fn into_inner_unchecked(self) -> R {
self.inner
}
fn drain_message(&mut self) -> io::Result<()> {
self.plaintext_consumed = self.plaintext_available;
while self.status == StreamStatus::Open {
self.load_segment()?;
self.plaintext_consumed = self.plaintext_available;
}
if self.status == StreamStatus::Failed {
Err(poisoned_error())
} else {
Ok(())
}
}
fn load_segment(&mut self) -> io::Result<()> {
let result = self.try_load_segment();
if result.is_err() {
self.status = StreamStatus::Failed;
}
result
}
fn try_load_segment(&mut self) -> io::Result<()> {
let framing = self.fetch_segment()?;
let decryptor = self.decryptor.as_mut().ok_or_else(poisoned_error)?;
let plaintext_length = decryptor
.decrypt_segment_in_place(&mut self.buffer)
.map_err(decryption_error)?
.len();
self.finish_if_final(framing)?;
self.plaintext_consumed = 0;
self.plaintext_available = plaintext_length;
Ok(())
}
#[inline]
fn fetch_segment(&mut self) -> io::Result<SegmentFraming> {
let mut prefix = [0; crate::SEGMENT_PREFIX_LENGTH];
let actual = read_fully(&mut self.inner, &mut prefix)?;
if actual != prefix.len() {
return Err(decryption_error(Error::Truncated));
}
let framing =
SegmentFraming::decode(self.parameters(), prefix).map_err(decryption_error)?;
let ciphertext_length = framing.ciphertext_length();
let ciphertext = self
.buffer
.prepare_ciphertext(ciphertext_length)
.map_err(decryption_error)?;
ciphertext[..prefix.len()].copy_from_slice(&prefix);
let actual = read_fully(&mut self.inner, &mut ciphertext[prefix.len()..])?;
let expected_remainder = ciphertext_length - prefix.len();
if actual != expected_remainder {
return Err(decryption_error(Error::InvalidCiphertextLength {
actual: prefix.len() + actual,
required: LengthRequirement::Exactly(ciphertext_length),
}));
}
Ok(framing)
}
#[inline]
fn finish_if_final(&mut self, framing: SegmentFraming) -> io::Result<()> {
if framing.is_final() {
self.decryptor
.take()
.ok_or_else(poisoned_error)?
.finish()
.map_err(decryption_error)?;
self.status = StreamStatus::Finished;
}
Ok(())
}
fn read_segment_direct(&mut self, output: &mut [u8]) -> io::Result<usize> {
let result = self.try_read_segment_direct(output);
if result.is_err() {
self.status = StreamStatus::Failed;
}
result
}
fn try_read_segment_direct(&mut self, output: &mut [u8]) -> io::Result<usize> {
let framing = self.fetch_segment()?;
let written = self
.decryptor
.as_mut()
.ok_or_else(poisoned_error)?
.decrypt_segment_into_framed(
self.buffer.ciphertext().map_err(decryption_error)?,
framing,
output,
)
.map_err(decryption_error)?;
self.finish_if_final(framing)?;
self.plaintext_consumed = 0;
self.plaintext_available = 0;
Ok(written)
}
}
impl<R: Read> Read for DecryptReader<R> {
fn read(&mut self, output: &mut [u8]) -> io::Result<usize> {
if output.is_empty() {
return Ok(0);
}
if self.status == StreamStatus::Failed {
return Err(poisoned_error());
}
if self.plaintext_consumed == self.plaintext_available {
if self.status == StreamStatus::Finished {
return Ok(0);
}
if output.len() >= self.parameters().plaintext_segment_length() {
return self.read_segment_direct(output);
}
self.load_segment()?;
if self.plaintext_consumed == self.plaintext_available {
debug_assert_eq!(self.status, StreamStatus::Finished);
return Ok(0);
}
}
let plaintext = self.buffer.plaintext().map_err(decryption_error)?;
let read = output
.len()
.min(self.plaintext_available - self.plaintext_consumed);
let end = self.plaintext_consumed + read;
output[..read].copy_from_slice(&plaintext[self.plaintext_consumed..end]);
self.plaintext_consumed = end;
Ok(read)
}
}
fn poisoned_error() -> io::Error {
io::Error::other("FLOE stream is poisoned after an earlier error")
}
#[cfg(test)]
mod tests {
use std::io::{Cursor, Read as _, Write as _};
use super::*;
use crate::key::test_key;
use crate::{decrypt, decrypt_with_parameters, encrypt, random_access};
#[derive(Debug, Default)]
struct FlushFails(Vec<u8>);
impl Write for FlushFails {
fn write(&mut self, input: &[u8]) -> io::Result<usize> {
self.0.extend_from_slice(input);
Ok(input.len())
}
fn flush(&mut self) -> io::Result<()> {
Err(io::Error::other("injected flush failure"))
}
}
#[test]
fn encrypt_reader_round_trips_plaintext_boundaries() {
let parameters = Parameters::SEGMENT_4_KIB;
let capacity = parameters.plaintext_segment_length();
for length in [0, 1, capacity - 1, capacity, capacity + 1, 2 * capacity] {
let plaintext = vec![0x31; length];
let mut reader =
EncryptReader::new(Cursor::new(&plaintext), &test_key(), b"io", parameters)
.unwrap();
let mut ciphertext = Vec::new();
reader.read_to_end(&mut ciphertext).unwrap();
assert!(
reader.is_finished(),
"reader did not finish at length {length}"
);
assert_eq!(
decrypt(&test_key(), b"io", &ciphertext).unwrap(),
plaintext,
"reader round trip failed at length {length}"
);
}
}
#[test]
fn encrypt_reader_serves_whole_segments_into_large_buffers() {
let parameters = Parameters::SEGMENT_4_KIB;
let capacity = parameters.plaintext_segment_length();
for length in [0, 1, capacity - 1, capacity, capacity + 1, 3 * capacity + 7] {
let plaintext: Vec<u8> = (0..length)
.map(|index| u8::try_from(index % 251).unwrap())
.collect();
let mut reader =
EncryptReader::new(Cursor::new(&plaintext), &test_key(), b"io", parameters)
.unwrap();
let mut ciphertext = Vec::new();
let mut chunk = vec![0u8; parameters.ciphertext_segment_length()];
loop {
let read = reader.read(&mut chunk).unwrap();
if read == 0 {
break;
}
ciphertext.extend_from_slice(&chunk[..read]);
}
assert!(reader.is_finished(), "not finished at length {length}");
assert_eq!(
decrypt(&test_key(), b"io", &ciphertext).unwrap(),
plaintext,
"direct-read round trip failed at length {length}"
);
}
}
#[test]
fn encrypt_reader_alternates_small_and_segment_sized_reads() {
let parameters = Parameters::SEGMENT_4_KIB;
let capacity = parameters.plaintext_segment_length();
let plaintext: Vec<u8> = (0..4 * capacity + 13)
.map(|index| u8::try_from(index % 249).unwrap())
.collect();
let mut reader =
EncryptReader::new(Cursor::new(&plaintext), &test_key(), b"io", parameters).unwrap();
let mut ciphertext = Vec::new();
let mut small = [0u8; 7];
let mut large = vec![0u8; parameters.ciphertext_segment_length()];
let mut use_small = true;
loop {
let read = if use_small {
let read = reader.read(&mut small).unwrap();
ciphertext.extend_from_slice(&small[..read]);
read
} else {
let read = reader.read(&mut large).unwrap();
ciphertext.extend_from_slice(&large[..read]);
read
};
if read == 0 {
break;
}
use_small = !use_small;
}
assert!(reader.is_finished());
assert_eq!(decrypt(&test_key(), b"io", &ciphertext).unwrap(), plaintext);
}
#[test]
fn adapters_round_trip_across_parameter_sets() {
for parameters in [Parameters::SEGMENT_4_KIB, Parameters::SEGMENT_1_MIB] {
let plaintext = vec![0x31; parameters.plaintext_segment_length() + 7];
let mut writer =
EncryptWriter::new(Vec::new(), &test_key(), b"io", parameters).unwrap();
writer.write_all(&plaintext).unwrap();
let ciphertext = writer.finish().unwrap();
let mut reader =
DecryptReader::new(Cursor::new(ciphertext), &test_key(), b"io").unwrap();
let mut recovered = Vec::new();
reader.read_to_end(&mut recovered).unwrap();
assert_eq!(recovered, plaintext);
assert!(reader.is_finished());
}
}
#[test]
fn decrypt_reader_serves_whole_segments_into_large_buffers() {
let parameters = Parameters::SEGMENT_4_KIB;
let capacity = parameters.plaintext_segment_length();
for length in [0, 1, capacity - 1, capacity, capacity + 1, 3 * capacity + 7] {
let plaintext: Vec<u8> = (0..length)
.map(|index| u8::try_from(index % 251).unwrap())
.collect();
let ciphertext = encrypt(&test_key(), b"io", parameters, &plaintext).unwrap();
let mut reader =
DecryptReader::new(Cursor::new(&ciphertext), &test_key(), b"io").unwrap();
let mut recovered = Vec::new();
let mut chunk = vec![0u8; capacity];
loop {
let read = reader.read(&mut chunk).unwrap();
if read == 0 {
break;
}
recovered.extend_from_slice(&chunk[..read]);
}
assert_eq!(recovered, plaintext, "direct read failed at {length}");
assert!(reader.is_finished(), "not finished at length {length}");
reader.try_finish().unwrap();
}
}
#[test]
fn decrypt_reader_alternates_small_and_segment_sized_reads() {
let parameters = Parameters::SEGMENT_4_KIB;
let capacity = parameters.plaintext_segment_length();
let plaintext: Vec<u8> = (0..4 * capacity + 13)
.map(|index| u8::try_from(index % 249).unwrap())
.collect();
let ciphertext = encrypt(&test_key(), b"io", parameters, &plaintext).unwrap();
let mut reader = DecryptReader::new(Cursor::new(&ciphertext), &test_key(), b"io").unwrap();
let mut recovered = Vec::new();
let mut small = [0u8; 7];
let mut large = vec![0u8; capacity];
let mut use_small = true;
loop {
let read = if use_small {
let read = reader.read(&mut small).unwrap();
recovered.extend_from_slice(&small[..read]);
read
} else {
let read = reader.read(&mut large).unwrap();
recovered.extend_from_slice(&large[..read]);
read
};
if read == 0 {
break;
}
use_small = !use_small;
}
assert_eq!(recovered, plaintext);
assert!(reader.is_finished());
reader.try_finish().unwrap();
}
#[test]
fn direct_segment_reads_preserve_a_following_frame() {
let parameters = Parameters::SEGMENT_4_KIB;
let capacity = parameters.plaintext_segment_length();
let plaintext = vec![0x42; capacity + 9];
let mut framed = encrypt(&test_key(), b"io", parameters, &plaintext).unwrap();
framed.extend_from_slice(b"next frame");
let mut reader = DecryptReader::new(Cursor::new(framed), &test_key(), b"io").unwrap();
let mut recovered = Vec::new();
let mut chunk = vec![0u8; capacity];
loop {
let read = reader.read(&mut chunk).unwrap();
if read == 0 {
break;
}
recovered.extend_from_slice(&chunk[..read]);
}
assert_eq!(recovered, plaintext);
let inner = reader.finish_frame().unwrap();
assert_eq!(
&inner.get_ref()[usize::try_from(inner.position()).unwrap()..],
b"next frame"
);
}
#[test]
fn try_finish_verifies_complete_message_without_reading() {
let parameters = Parameters::SEGMENT_4_KIB;
let plaintext = vec![0x31; parameters.plaintext_segment_length() + 7];
let mut reader = DecryptReader::new(
Cursor::new(encrypt(&test_key(), b"io", parameters, &plaintext).unwrap()),
&test_key(),
b"io",
)
.unwrap();
reader.try_finish().unwrap();
assert!(reader.is_finished());
}
#[test]
fn writer_boundary_lengths_round_trip() {
let parameters = Parameters::SEGMENT_4_KIB;
let capacity = parameters.plaintext_segment_length();
for length in [0, 1, capacity - 1, capacity, capacity + 1, 2 * capacity] {
let plaintext = vec![0x32; length];
let mut writer =
EncryptWriter::new(Vec::new(), &test_key(), b"io", parameters).unwrap();
writer.write_all(&plaintext).unwrap();
let ciphertext = writer.finish().unwrap();
assert_eq!(
decrypt(&test_key(), b"io", &ciphertext).unwrap(),
plaintext,
"writer round trip failed at length {length}"
);
}
}
#[test]
fn writer_mixes_buffered_and_full_segment_writes() {
let parameters = Parameters::SEGMENT_4_KIB;
let capacity = parameters.plaintext_segment_length();
let plaintext: Vec<u8> = (0..3 * capacity + 41)
.map(|index| u8::try_from(index % 251).unwrap())
.collect();
let mut writer = EncryptWriter::new(Vec::new(), &test_key(), b"io", parameters).unwrap();
writer.write_all(&plaintext[..17]).unwrap();
writer.write_all(&plaintext[17..3 * capacity]).unwrap();
writer.write_all(&plaintext[3 * capacity..]).unwrap();
let ciphertext = writer.finish().unwrap();
assert_eq!(decrypt(&test_key(), b"io", &ciphertext).unwrap(), plaintext);
}
#[test]
fn writer_consumes_at_most_one_segment_per_write_call() {
let parameters = Parameters::SEGMENT_4_KIB;
let capacity = parameters.plaintext_segment_length();
let input = vec![0x31; capacity + 100];
let mut writer = EncryptWriter::new(Vec::new(), &test_key(), b"io", parameters).unwrap();
assert_eq!(writer.write(&input).unwrap(), capacity);
assert_eq!(
writer.get_ref().len(),
Header::LEN + parameters.ciphertext_segment_length()
);
assert_eq!(writer.write(&input[capacity..]).unwrap(), 100);
let ciphertext = writer.finish().unwrap();
assert_eq!(
decrypt(&test_key(), b"io", &ciphertext).unwrap(),
input[..capacity + 100]
);
}
#[test]
fn writer_flush_does_not_emit_partial_segment() {
let mut writer =
EncryptWriter::new(Vec::new(), &test_key(), b"io", Parameters::SEGMENT_4_KIB).unwrap();
writer.write_all(b"buffered plaintext").unwrap();
writer.flush().unwrap();
assert_eq!(writer.get_ref().len(), Header::LEN);
let ciphertext = writer.finish().unwrap();
assert_eq!(
decrypt(&test_key(), b"io", &ciphertext).unwrap(),
b"buffered plaintext"
);
}
#[test]
fn frame_finish_preserves_trailing_data() {
let parameters = Parameters::SEGMENT_4_KIB;
let mut writer = EncryptWriter::new(Vec::new(), &test_key(), b"io", parameters).unwrap();
writer.write_all(b"message").unwrap();
let mut framed = writer.finish().unwrap();
framed.extend_from_slice(b"next frame");
let mut reader = DecryptReader::new(Cursor::new(framed), &test_key(), b"io").unwrap();
let mut plaintext = Vec::new();
reader.read_to_end(&mut plaintext).unwrap();
assert_eq!(plaintext, b"message");
let inner = reader.finish_frame().unwrap();
assert_eq!(
&inner.get_ref()[usize::try_from(inner.position()).unwrap()..],
b"next frame"
);
}
#[test]
fn default_finish_rejects_trailing_data() {
let parameters = Parameters::SEGMENT_4_KIB;
let mut writer = EncryptWriter::new(Vec::new(), &test_key(), b"io", parameters).unwrap();
writer.write_all(b"message").unwrap();
let mut bytes = writer.finish().unwrap();
bytes.push(0);
let mut reader = DecryptReader::new(Cursor::new(bytes), &test_key(), b"io").unwrap();
let mut plaintext = Vec::new();
reader.read_to_end(&mut plaintext).unwrap();
assert_eq!(
reader.finish().unwrap_err().error().kind(),
io::ErrorKind::InvalidData
);
}
#[test]
fn consuming_reader_finish_error_preserves_wrapped_reader() {
let mut ciphertext = {
let mut writer =
EncryptWriter::new(Vec::new(), &test_key(), b"io", Parameters::SEGMENT_4_KIB)
.unwrap();
writer.write_all(b"message").unwrap();
writer.finish().unwrap()
};
ciphertext.pop();
let reader = DecryptReader::new(Cursor::new(ciphertext), &test_key(), b"io").unwrap();
let failure = reader.finish().unwrap_err();
assert_eq!(failure.error().kind(), io::ErrorKind::InvalidData);
let inner = failure.into_inner();
assert!(inner.position() >= u64::try_from(Header::LEN).unwrap());
}
#[test]
fn reader_constructors_classify_short_headers_consistently() {
for length in 0..Header::LEN {
let input = vec![0u8; length];
let stream_error =
DecryptReader::new(Cursor::new(&input), &test_key(), b"io").unwrap_err();
let random_error =
random_access::Reader::new(Cursor::new(&input), &test_key(), b"io").unwrap_err();
for error in [stream_error, random_error] {
assert_eq!(error.kind(), io::ErrorKind::InvalidData);
assert!(matches!(
error
.get_ref()
.and_then(|source| source.downcast_ref::<Error>()),
Some(Error::InvalidHeaderLength { actual }) if *actual == length
));
}
}
}
#[test]
fn unfinished_extraction_is_explicit_and_truncated() {
let writer =
EncryptWriter::new(Vec::new(), &test_key(), b"io", Parameters::SEGMENT_4_KIB).unwrap();
let ciphertext = writer.into_inner_unfinished();
assert_eq!(
decrypt(&test_key(), b"io", &ciphertext),
Err(Error::Truncated)
);
}
#[test]
fn consuming_finish_error_preserves_wrapped_writer() {
let mut writer = EncryptWriter::new(
FlushFails::default(),
&test_key(),
b"io",
Parameters::SEGMENT_4_KIB,
)
.unwrap();
writer.write_all(b"message").unwrap();
let error = writer.finish().unwrap_err();
assert_eq!(error.error().kind(), io::ErrorKind::Other);
assert!(!error.into_inner().0.is_empty());
}
#[test]
fn finish_error_does_not_require_inner_type_to_implement_debug() {
struct NoDebug;
fn assert_error<T: std::error::Error>() {}
assert_error::<FinishError<NoDebug>>();
}
#[test]
fn parameter_policy_checkable_after_construction() {
let parameters = Parameters::SEGMENT_1_MIB;
let mut writer = EncryptWriter::new(Vec::new(), &test_key(), b"io", parameters).unwrap();
writer.write_all(b"message").unwrap();
let ciphertext = writer.finish().unwrap();
assert_eq!(
decrypt_with_parameters(&test_key(), b"io", parameters, &ciphertext).unwrap(),
b"message"
);
let reader = DecryptReader::new(Cursor::new(ciphertext), &test_key(), b"io").unwrap();
assert_eq!(reader.parameters(), parameters);
}
}