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_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)?;
Ok(Self::from_decryptor(inner, decryptor, header))
}
pub fn new_with_parameters(
mut inner: R,
key: &Key,
aad: &[u8],
parameters: Parameters,
) -> io::Result<Self> {
let header = read_header(&mut inner)?;
let decryptor = Decryptor::new_with_parameters(key, aad, parameters, &header)
.map_err(decryption_error)?;
Ok(Self::from_decryptor(inner, decryptor, header))
}
fn from_decryptor(inner: R, decryptor: Decryptor, header: Header) -> Self {
let parameters = decryptor.parameters();
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;
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"))
}
}
#[derive(Debug)]
struct FailingWriter {
accepted: Vec<u8>,
budget: usize,
write_calls: usize,
}
impl FailingWriter {
fn new(budget: usize) -> Self {
Self {
accepted: Vec::new(),
budget,
write_calls: 0,
}
}
}
impl Write for FailingWriter {
fn write(&mut self, input: &[u8]) -> io::Result<usize> {
self.write_calls += 1;
if self.accepted.len() + input.len() > self.budget {
return Err(io::Error::other("injected write failure"));
}
self.accepted.extend_from_slice(input);
Ok(input.len())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
#[derive(Debug)]
struct FailingReader {
data: Cursor<Vec<u8>>,
}
impl Read for FailingReader {
fn read(&mut self, output: &mut [u8]) -> io::Result<usize> {
let read = self.data.read(output)?;
if read == 0 {
return Err(io::Error::other("injected read failure"));
}
Ok(read)
}
}
#[derive(Debug)]
struct InterruptedReader {
data: Cursor<Vec<u8>>,
interrupt_next: bool,
}
impl Read for InterruptedReader {
fn read(&mut self, output: &mut [u8]) -> io::Result<usize> {
self.interrupt_next = !self.interrupt_next;
if self.interrupt_next {
return Err(io::Error::new(
io::ErrorKind::Interrupted,
"injected interruption",
));
}
self.data.read(output)
}
}
#[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::io_source(&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);
}
#[test]
fn writer_rejects_writes_after_finish() {
let mut writer =
EncryptWriter::new(Vec::new(), &test_key(), b"io", Parameters::SEGMENT_4_KIB).unwrap();
writer.write_all(b"message").unwrap();
writer.try_finish().unwrap();
assert!(writer.is_finished());
assert_eq!(
writer.write(b"more").unwrap_err().kind(),
io::ErrorKind::BrokenPipe
);
assert_eq!(writer.write(&[]).unwrap(), 0);
}
#[test]
fn writer_try_finish_is_idempotent_after_success() {
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();
writer.try_finish().unwrap();
let emitted = writer.get_ref().len();
writer.try_finish().unwrap();
assert_eq!(writer.get_ref().len(), emitted);
let ciphertext = writer.finish().unwrap();
assert_eq!(
decrypt(&test_key(), b"io", &ciphertext).unwrap(),
b"message"
);
}
#[test]
fn writer_empty_write_is_a_noop() {
let mut writer =
EncryptWriter::new(Vec::new(), &test_key(), b"io", Parameters::SEGMENT_4_KIB).unwrap();
assert_eq!(writer.write(&[]).unwrap(), 0);
writer.write_all(b"buffered").unwrap();
assert_eq!(writer.write(&[]).unwrap(), 0);
assert_eq!(writer.get_ref().len(), Header::LEN);
let ciphertext = writer.finish().unwrap();
assert_eq!(
decrypt(&test_key(), b"io", &ciphertext).unwrap(),
b"buffered"
);
}
#[test]
fn writer_is_poisoned_after_segment_write_failure() {
let parameters = Parameters::SEGMENT_4_KIB;
let capacity = parameters.plaintext_segment_length();
let mut writer = EncryptWriter::new(
FailingWriter::new(Header::LEN),
&test_key(),
b"io",
parameters,
)
.unwrap();
let error = writer.write(&vec![0x33; capacity]).unwrap_err();
assert!(error.to_string().contains("injected write failure"));
let calls = writer.get_ref().write_calls;
for error in [
writer.write(b"x").unwrap_err(),
writer.flush().unwrap_err(),
writer.try_finish().unwrap_err(),
] {
assert!(error.to_string().contains("poisoned"));
}
assert_eq!(writer.get_ref().write_calls, calls);
assert_eq!(writer.get_ref().accepted.len(), Header::LEN);
}
#[test]
fn writer_is_poisoned_after_final_segment_write_failure() {
let mut writer = EncryptWriter::new(
FailingWriter::new(Header::LEN),
&test_key(),
b"io",
Parameters::SEGMENT_4_KIB,
)
.unwrap();
writer.write_all(b"partial final").unwrap();
let error = writer.try_finish().unwrap_err();
assert!(error.to_string().contains("injected write failure"));
assert!(!writer.is_finished());
assert!(
writer
.try_finish()
.unwrap_err()
.to_string()
.contains("poisoned")
);
}
#[test]
fn finish_error_exposes_standard_error_interfaces() {
let mut writer = EncryptWriter::new(
FlushFails::default(),
&test_key(),
b"io",
Parameters::SEGMENT_4_KIB,
)
.unwrap();
writer.write_all(b"message").unwrap();
let failure = writer.finish().unwrap_err();
assert!(failure.to_string().contains("failed to finish FLOE stream"));
let debug = format!("{failure:?}");
assert!(debug.contains("FinishError"));
assert!(debug.contains("injected flush failure"));
let source = std::error::Error::source(&failure).unwrap();
assert!(source.to_string().contains("injected flush failure"));
let (error, sink) = failure.into_parts();
assert!(error.to_string().contains("injected flush failure"));
assert!(!sink.0.is_empty());
}
#[test]
fn encrypt_reader_empty_read_is_a_noop() {
let mut reader = EncryptReader::new(
Cursor::new(b"data".to_vec()),
&test_key(),
b"io",
Parameters::SEGMENT_4_KIB,
)
.unwrap();
assert_eq!(reader.read(&mut []).unwrap(), 0);
let mut ciphertext = Vec::new();
reader.read_to_end(&mut ciphertext).unwrap();
assert_eq!(decrypt(&test_key(), b"io", &ciphertext).unwrap(), b"data");
}
#[test]
fn encrypt_reader_poisoned_after_source_read_failure() {
let parameters = Parameters::SEGMENT_4_KIB;
for direct in [false, true] {
let source = FailingReader {
data: Cursor::new(vec![0x31; 10]),
};
let mut reader = EncryptReader::new(source, &test_key(), b"io", parameters).unwrap();
let mut header = [0u8; Header::LEN];
reader.read_exact(&mut header).unwrap();
let mut output = vec![
0u8;
if direct {
parameters.ciphertext_segment_length()
} else {
32
}
];
let error = reader.read(&mut output).unwrap_err();
assert!(
error.to_string().contains("injected read failure"),
"unexpected first error on direct={direct}: {error}"
);
let error = reader.read(&mut output).unwrap_err();
assert!(error.to_string().contains("poisoned"));
assert_eq!(reader.read(&mut []).unwrap(), 0);
}
}
#[test]
fn encrypt_reader_accessors_preserve_inner_reader() {
let parameters = Parameters::SEGMENT_4_KIB;
let mut reader = EncryptReader::new(
Cursor::new(b"accessor data".to_vec()),
&test_key(),
b"io",
parameters,
)
.unwrap();
assert_eq!(reader.parameters(), parameters);
assert!(!reader.is_finished());
let header = *reader.header();
let mut ciphertext = Vec::new();
reader.read_to_end(&mut ciphertext).unwrap();
assert!(reader.is_finished());
assert_eq!(&ciphertext[..Header::LEN], header.as_bytes());
assert_eq!(reader.get_ref().position(), 13);
reader.get_mut().set_position(0);
let mut more = [0u8; 8];
assert_eq!(reader.read(&mut more).unwrap(), 0);
let inner = reader.into_inner_unfinished();
assert_eq!(inner.position(), 0);
assert_eq!(
decrypt(&test_key(), b"io", &ciphertext).unwrap(),
b"accessor data"
);
}
#[test]
fn decrypt_reader_empty_read_is_a_noop() {
let parameters = Parameters::SEGMENT_4_KIB;
let ciphertext = encrypt(&test_key(), b"io", parameters, b"data").unwrap();
let mut reader = DecryptReader::new(Cursor::new(ciphertext), &test_key(), b"io").unwrap();
assert_eq!(reader.read(&mut []).unwrap(), 0);
let mut plaintext = Vec::new();
reader.read_to_end(&mut plaintext).unwrap();
assert_eq!(plaintext, b"data");
reader.try_finish().unwrap();
}
#[test]
fn decrypt_reader_poisoned_after_truncated_prefix() {
let parameters = Parameters::SEGMENT_4_KIB;
let capacity = parameters.plaintext_segment_length();
let plaintext = vec![0x37; capacity + 10];
let mut ciphertext = encrypt(&test_key(), b"io", parameters, &plaintext).unwrap();
ciphertext.truncate(Header::LEN + parameters.ciphertext_segment_length() + 2);
let mut reader = DecryptReader::new(Cursor::new(ciphertext), &test_key(), b"io").unwrap();
let mut recovered = Vec::new();
let mut chunk = [0u8; 1024];
let error = loop {
match reader.read(&mut chunk) {
Ok(read) => recovered.extend_from_slice(&chunk[..read]),
Err(error) => break error,
}
};
assert_eq!(recovered, plaintext[..capacity]);
assert_eq!(error.kind(), io::ErrorKind::InvalidData);
assert!(matches!(Error::io_source(&error), Some(Error::Truncated)));
let error = reader.read(&mut chunk).unwrap_err();
assert!(error.to_string().contains("poisoned"));
}
#[test]
fn decrypt_reader_poisoned_after_truncated_segment_body() {
let parameters = Parameters::SEGMENT_4_KIB;
let ciphertext = encrypt(&test_key(), b"io", parameters, b"message").unwrap();
let segment_length = ciphertext.len() - Header::LEN;
let mut truncated = ciphertext;
truncated.truncate(Header::LEN + segment_length - 3);
let mut reader = DecryptReader::new(Cursor::new(truncated), &test_key(), b"io").unwrap();
let mut chunk = [0u8; 4];
let error = reader.read(&mut chunk).unwrap_err();
assert_eq!(error.kind(), io::ErrorKind::InvalidData);
assert!(matches!(
Error::io_source(&error),
Some(Error::InvalidCiphertextLength {
actual,
required: LengthRequirement::Exactly(required),
}) if *actual == segment_length - 3 && *required == segment_length
));
assert!(
reader
.read(&mut chunk)
.unwrap_err()
.to_string()
.contains("poisoned")
);
}
#[test]
fn decrypt_reader_returns_no_plaintext_after_tag_corruption() {
let parameters = Parameters::SEGMENT_4_KIB;
let capacity = parameters.plaintext_segment_length();
let ciphertext = encrypt(&test_key(), b"io", parameters, b"secret payload").unwrap();
for direct in [false, true] {
let mut corrupted = ciphertext.clone();
*corrupted.last_mut().unwrap() ^= 0x01;
let mut reader =
DecryptReader::new(Cursor::new(corrupted), &test_key(), b"io").unwrap();
let mut output = vec![0u8; if direct { capacity } else { 4 }];
let error = reader.read(&mut output).unwrap_err();
assert_eq!(error.kind(), io::ErrorKind::InvalidData);
assert!(
matches!(Error::io_source(&error), Some(Error::AuthenticationFailed)),
"unexpected error on direct={direct}: {error}"
);
assert!(
reader
.read(&mut output)
.unwrap_err()
.to_string()
.contains("poisoned")
);
}
}
#[test]
fn decrypt_reader_finish_propagates_underlying_read_error() {
let parameters = Parameters::SEGMENT_4_KIB;
let ciphertext = encrypt(&test_key(), b"io", parameters, b"message").unwrap();
let source = FailingReader {
data: Cursor::new(ciphertext),
};
let mut reader = DecryptReader::new(source, &test_key(), b"io").unwrap();
let mut plaintext = Vec::new();
reader.read_to_end(&mut plaintext).unwrap();
assert_eq!(plaintext, b"message");
let error = reader.try_finish().unwrap_err();
assert!(error.to_string().contains("injected read failure"));
assert!(
reader
.try_finish()
.unwrap_err()
.to_string()
.contains("poisoned")
);
assert!(
reader
.read(&mut [0u8; 5])
.unwrap_err()
.to_string()
.contains("poisoned")
);
}
#[test]
fn decrypt_reader_try_finish_is_idempotent_after_success() {
let parameters = Parameters::SEGMENT_4_KIB;
let ciphertext = encrypt(&test_key(), b"io", parameters, b"message").unwrap();
let mut reader = DecryptReader::new(Cursor::new(ciphertext), &test_key(), b"io").unwrap();
reader.try_finish().unwrap();
reader.try_finish().unwrap();
assert!(reader.is_finished());
}
#[test]
fn decrypt_reader_accessors_and_unchecked_extraction_are_observable() {
let parameters = Parameters::SEGMENT_4_KIB;
let ciphertext = encrypt(&test_key(), b"io", parameters, b"peek").unwrap();
let mut reader =
DecryptReader::new(Cursor::new(ciphertext.clone()), &test_key(), b"io").unwrap();
assert_eq!(reader.header().as_bytes(), &ciphertext[..Header::LEN]);
assert_eq!(reader.parameters(), parameters);
assert!(!reader.is_finished());
assert_eq!(
reader.get_ref().position(),
u64::try_from(Header::LEN).unwrap()
);
reader.get_mut().set_position(0);
let inner = reader.into_inner_unchecked();
assert_eq!(inner.position(), 0);
}
#[test]
fn decrypt_reader_retries_interrupted_reads() {
let parameters = Parameters::SEGMENT_4_KIB;
let capacity = parameters.plaintext_segment_length();
let plaintext: Vec<u8> = (0..2 * capacity + 5)
.map(|index| u8::try_from(index % 251).unwrap())
.collect();
let ciphertext = encrypt(&test_key(), b"io", parameters, &plaintext).unwrap();
let source = InterruptedReader {
data: Cursor::new(ciphertext),
interrupt_next: false,
};
let mut reader = DecryptReader::new(source, &test_key(), b"io").unwrap();
let mut recovered = Vec::new();
reader.read_to_end(&mut recovered).unwrap();
reader.try_finish().unwrap();
assert_eq!(recovered, plaintext);
assert!(reader.is_finished());
}
}