use super::super::{DecodeError, DecodeOperation, DecoderSession, DecoderStatus};
use crate::io::FinishError;
use std::io::{self, ErrorKind, Write};
#[derive(Debug, Clone, Copy, Eq, PartialEq)]
enum Lifecycle {
Accepting,
Finishing,
Finished,
Failed,
}
#[derive(Debug)]
pub struct DecoderWriter<'d, 'dict, W> {
session: DecoderSession<'d, 'dict>,
writer: W,
buffer: Vec<u8>,
cursor: u16,
filled: u16,
lifecycle: Lifecycle,
pending_error: Option<io::Error>,
}
impl<'d, 'dict, W: Write> DecoderWriter<'d, 'dict, W> {
pub(super) fn new(session: DecoderSession<'d, 'dict>, writer: W) -> Result<Self, DecodeError> {
Ok(Self {
session,
writer,
buffer: super::buffer()?,
cursor: 0,
filled: 0,
lifecycle: Lifecycle::Accepting,
pending_error: None,
})
}
pub const fn get_ref(&self) -> &W {
&self.writer
}
pub const fn get_mut(&mut self) -> &mut W {
&mut self.writer
}
pub const fn is_finished(&self) -> bool {
matches!(self.lifecycle, Lifecycle::Finished)
}
fn check_error(&mut self) -> io::Result<()> {
if let Some(error) = self.pending_error.take() {
return Err(error);
}
self.drain()?;
if self.lifecycle == Lifecycle::Failed {
return Err(DecodeError::InvalidState.into());
}
Ok(())
}
fn drain(&mut self) -> io::Result<()> {
while self.cursor < self.filled {
match self
.writer
.write(&self.buffer[usize::from(self.cursor)..usize::from(self.filled)])
{
Ok(0) => return Err(io::Error::from(ErrorKind::WriteZero)),
Ok(written) => self.cursor += written as u16,
Err(error) if error.kind() == ErrorKind::Interrupted => continue,
Err(error) => return Err(error),
}
}
self.cursor = 0;
self.filled = 0;
Ok(())
}
fn pump(
&mut self,
input: &[u8],
operation: DecodeOperation,
) -> io::Result<(usize, DecoderStatus)> {
match self.session.process(input, &mut self.buffer, operation) {
Ok(progress) => {
self.filled = progress.produced as u16;
match self.drain() {
Err(error) if progress.consumed != 0 => {
self.pending_error = Some(error);
Ok((progress.consumed, progress.status))
}
Err(error) => Err(error),
Ok(()) => Ok((progress.consumed, progress.status)),
}
}
Err(failure) => {
self.filled = failure.produced as u16;
self.lifecycle = Lifecycle::Failed;
let error = failure.error.into();
if failure.consumed != 0 {
self.pending_error = Some(error);
Ok((failure.consumed, DecoderStatus::NeedsInput))
} else {
Err(error)
}
}
}
}
pub fn try_finish(&mut self) -> io::Result<()> {
if self.lifecycle == Lifecycle::Finished {
return Ok(());
}
if self.lifecycle == Lifecycle::Accepting {
self.lifecycle = Lifecycle::Finishing;
}
self.check_error()?;
self.drain()?;
while !self.session.is_finished() {
self.pump(&[], DecodeOperation::Finish)?;
}
self.writer.flush()?;
self.lifecycle = Lifecycle::Finished;
Ok(())
}
pub fn finish(mut self) -> Result<W, FinishError<Self>> {
match self.try_finish() {
Ok(()) => Ok(self.writer),
Err(error) => Err(FinishError::from_parts(error, self)),
}
}
}
impl<W: Write> Write for DecoderWriter<'_, '_, W> {
fn write(&mut self, input: &[u8]) -> io::Result<usize> {
if input.is_empty() {
return Ok(0);
}
self.check_error()?;
if self.lifecycle != Lifecycle::Accepting {
return Err(DecodeError::InvalidState.into());
}
if self.session.is_finished() {
return Err(DecodeError::TrailingData {
offset: self.session.total_in(),
}
.into());
}
self.drain()?;
loop {
let (consumed, status) = self.pump(input, DecodeOperation::Process)?;
if consumed != 0 {
return Ok(consumed);
}
if status == DecoderStatus::Finished {
return Err(DecodeError::TrailingData {
offset: self.session.total_in(),
}
.into());
}
}
}
fn flush(&mut self) -> io::Result<()> {
self.check_error()?;
self.drain()?;
if self.lifecycle == Lifecycle::Finishing {
return self.try_finish();
}
while !self.session.is_finished() {
let (_, status) = self.pump(&[], DecodeOperation::Process)?;
if status == DecoderStatus::NeedsInput {
break;
}
}
self.writer.flush()
}
}