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};
#[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)
}
}
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,
parameters: Parameters,
buffer: SegmentBuffer,
plaintext_length: usize,
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,
parameters,
buffer: SegmentBuffer::new(parameters),
plaintext_length: 0,
status: StreamStatus::Open,
})
}
#[must_use]
pub const fn header(&self) -> &Header {
&self.header
}
#[must_use]
pub const fn parameters(&self) -> Parameters {
self.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 => {}
}
let encryptor = self.encryptor.take().ok_or_else(poisoned_error)?;
self.buffer
.prepare_plaintext(self.plaintext_length)
.map_err(encryption_error)?;
if let Err(error) = encryptor.encrypt_final_segment_in_place(&mut self.buffer) {
self.status = StreamStatus::Failed;
return Err(encryption_error(error.into_error()));
}
let ciphertext = self.buffer.ciphertext().map_err(encryption_error)?;
if let Err(error) = self.inner.write_all(ciphertext) {
self.status = StreamStatus::Failed;
return Err(error);
}
self.plaintext_length = 0;
if let Err(error) = self.inner.flush() {
self.status = StreamStatus::Failed;
return Err(error);
}
self.status = StreamStatus::Finished;
Ok(())
}
pub fn finish(mut self) -> Result<W, FinishError<W>> {
match self.try_finish() {
Ok(()) => Ok(self.inner),
Err(error) => Err(FinishError {
inner: self.inner,
error,
}),
}
}
#[must_use]
pub fn into_inner_unfinished(self) -> W {
self.inner
}
fn emit_non_final(&mut self) -> io::Result<()> {
self.buffer
.prepare_plaintext(self.plaintext_length)
.map_err(encryption_error)?;
if let Err(error) = self
.encryptor
.as_mut()
.ok_or_else(poisoned_error)?
.encrypt_non_final_segment_in_place(&mut self.buffer)
{
self.status = StreamStatus::Failed;
return Err(encryption_error(error));
}
let ciphertext = self.buffer.ciphertext().map_err(encryption_error)?;
if let Err(error) = self.inner.write_all(ciphertext) {
self.status = StreamStatus::Failed;
return Err(error);
}
self.plaintext_length = 0;
Ok(())
}
}
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 self.plaintext_length == capacity {
self.emit_non_final()?;
}
let written = input.len().min(capacity - self.plaintext_length);
let start = self.plaintext_length;
let plaintext = self
.buffer
.prepare_plaintext(capacity)
.map_err(encryption_error)?;
plaintext[start..start + written].copy_from_slice(&input[..written]);
self.plaintext_length += written;
Ok(written)
}
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,
parameters: Parameters,
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,
parameters,
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.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 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
.prepare_plaintext(plaintext_length)
.map_err(encryption_error)?;
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(())
}
}
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);
}
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,
parameters: Parameters,
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,
parameters,
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.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>> {
match self.try_finish() {
Ok(()) => Ok(self.inner),
Err(error) => Err(FinishError {
inner: self.inner,
error,
}),
}
}
pub fn finish_frame(mut self) -> Result<R, FinishError<R>> {
match self.try_finish_frame() {
Ok(()) => Ok(self.inner),
Err(error) => Err(FinishError {
inner: self.inner,
error,
}),
}
}
#[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 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 =
crate::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),
}));
}
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();
if framing.is_final() {
self.decryptor
.take()
.ok_or_else(poisoned_error)?
.finish()
.map_err(decryption_error)?;
self.status = StreamStatus::Finished;
}
self.plaintext_consumed = 0;
self.plaintext_available = plaintext_length;
Ok(())
}
}
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);
}
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::{decrypt, decrypt_with_parameters, encrypt, random_access};
fn key() -> Key {
Key::from_bytes_with_provider([0x5a; Key::LEN], crate::Provider::COMPILED[0])
}
#[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), &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(&key(), b"io", &ciphertext).unwrap(),
plaintext,
"reader round trip failed at length {length}"
);
}
}
#[test]
fn adapters_round_trip_and_consuming_finish_returns_inner() {
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(), &key(), b"io", parameters).unwrap();
writer.write_all(&plaintext).unwrap();
let ciphertext = writer.finish().unwrap();
let mut reader = DecryptReader::new(Cursor::new(ciphertext), &key(), b"io").unwrap();
let mut recovered = Vec::new();
reader.read_to_end(&mut recovered).unwrap();
assert_eq!(recovered, plaintext);
assert!(reader.is_finished());
let mut reader = DecryptReader::new(
Cursor::new(encrypt(&key(), b"io", parameters, &plaintext).unwrap()),
&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(), &key(), b"io", parameters).unwrap();
writer.write_all(&plaintext).unwrap();
let ciphertext = writer.finish().unwrap();
assert_eq!(
decrypt(&key(), b"io", &ciphertext).unwrap(),
plaintext,
"writer round trip failed at length {length}"
);
}
}
#[test]
fn writer_flush_does_not_emit_partial_segment() {
let mut writer =
EncryptWriter::new(Vec::new(), &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(&key(), b"io", &ciphertext).unwrap(),
b"buffered plaintext"
);
}
#[test]
fn frame_finish_preserves_trailing_data_and_default_finish_rejects_it() {
let parameters = Parameters::SEGMENT_4_KIB;
let ciphertext = {
let mut writer = EncryptWriter::new(Vec::new(), &key(), b"io", parameters).unwrap();
writer.write_all(b"message").unwrap();
writer.finish().unwrap()
};
let mut framed = ciphertext.clone();
framed.extend_from_slice(b"next frame");
let mut reader = DecryptReader::new(Cursor::new(framed), &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"
);
let mut reader = DecryptReader::new(
Cursor::new({
let mut bytes = ciphertext;
bytes.push(0);
bytes
}),
&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(), &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), &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), &key(), b"io").unwrap_err();
let random_error =
random_access::Reader::new(Cursor::new(&input), &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(), &key(), b"io", Parameters::SEGMENT_4_KIB).unwrap();
let ciphertext = writer.into_inner_unfinished();
assert_eq!(decrypt(&key(), b"io", &ciphertext), Err(Error::Truncated));
}
#[test]
fn consuming_finish_error_preserves_wrapped_writer() {
let mut writer = EncryptWriter::new(
FlushFails::default(),
&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(), &key(), b"io", parameters).unwrap();
writer.write_all(b"message").unwrap();
let ciphertext = writer.finish().unwrap();
assert_eq!(
decrypt_with_parameters(&key(), b"io", parameters, &ciphertext).unwrap(),
b"message"
);
let reader = DecryptReader::new(Cursor::new(ciphertext), &key(), b"io").unwrap();
assert_eq!(reader.parameters(), parameters);
}
}