use super::{
DecoderDriver, OutputQueue, redacted_inner_state, stream_decoder_failed_error,
trailing_input_after_padding_error,
};
use crate::{Alphabet, Engine};
use std::io::{self, Write};
pub struct Decoder<W, A, const PAD: bool>
where
A: Alphabet,
{
inner: Option<W>,
engine: Engine<A, PAD>,
driver: DecoderDriver,
output: OutputQueue<1024>,
finished: bool,
failed: bool,
finalized: bool,
}
impl<W, A, const PAD: bool> Decoder<W, A, PAD>
where
A: Alphabet,
{
#[must_use]
pub const fn new(inner: W, engine: Engine<A, PAD>) -> Self {
Self {
inner: Some(inner),
engine,
driver: DecoderDriver::new::<A, PAD>(),
output: OutputQueue::new(),
finished: false,
finalized: false,
failed: false,
}
}
#[must_use]
pub fn get_ref(&self) -> &W {
self.inner_ref()
}
pub fn get_mut(&mut self) -> &mut W {
self.inner_mut()
}
#[must_use]
pub const fn engine(&self) -> Engine<A, PAD> {
self.engine
}
#[must_use]
pub const fn is_padded(&self) -> bool {
PAD
}
#[must_use]
pub const fn pending_len(&self) -> usize {
self.driver.pending_input_len()
}
#[must_use]
pub const fn has_pending_input(&self) -> bool {
self.pending_len() != 0
}
#[must_use]
pub const fn pending_input_needed_len(&self) -> usize {
if self.has_pending_input() {
4 - self.pending_len()
} else {
0
}
}
#[must_use]
pub const fn buffered_output_len(&self) -> usize {
self.output.len()
}
#[must_use]
pub const fn buffered_output_capacity(&self) -> usize {
self.output.capacity()
}
#[must_use]
pub const fn buffered_output_remaining_capacity(&self) -> usize {
self.output.available_capacity()
}
#[must_use]
pub const fn has_buffered_output(&self) -> bool {
!self.output.is_empty()
}
#[must_use]
pub const fn has_terminal_padding(&self) -> bool {
self.finished
}
#[must_use]
pub const fn is_finalized(&self) -> bool {
self.finalized
}
#[must_use]
pub const fn is_failed(&self) -> bool {
self.failed
}
#[must_use]
pub const fn can_into_inner(&self) -> bool {
!self.is_failed() && !self.has_pending_input() && !self.has_buffered_output()
}
#[must_use]
pub fn into_inner(mut self) -> W {
self.take_inner()
}
#[allow(clippy::result_large_err)]
pub fn try_into_inner(mut self) -> Result<W, Self> {
if !self.can_into_inner() {
return Err(self);
}
Ok(self.take_inner())
}
fn inner_ref(&self) -> &W {
match &self.inner {
Some(inner) => inner,
None => unreachable!("stream decoder inner writer was already taken"),
}
}
fn inner_mut(&mut self) -> &mut W {
match &mut self.inner {
Some(inner) => inner,
None => unreachable!("stream decoder inner writer was already taken"),
}
}
fn take_inner(&mut self) -> W {
match self.inner.take() {
Some(inner) => inner,
None => unreachable!("stream decoder inner writer was already taken"),
}
}
fn clear_pending(&mut self) {
self.driver.wipe();
}
fn clear_output(&mut self) {
self.output.clear_all();
}
}
impl<W, A, const PAD: bool> Drop for Decoder<W, A, PAD>
where
A: Alphabet,
{
fn drop(&mut self) {
self.clear_pending();
self.clear_output();
}
}
impl<W, A, const PAD: bool> core::fmt::Debug for Decoder<W, A, PAD>
where
A: Alphabet,
{
fn fmt(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
formatter
.debug_struct("Decoder")
.field("inner", &redacted_inner_state(self.inner.is_some()))
.field("engine", &self.engine)
.field("driver", &"<redacted>")
.field("pending", &"<redacted>")
.field("pending_len", &self.pending_len())
.field("pending_input_needed_len", &self.pending_input_needed_len())
.field("buffered_output_len", &self.output.len())
.field("buffered_output_capacity", &self.output.capacity())
.field(
"buffered_output_remaining_capacity",
&self.output.available_capacity(),
)
.field("can_into_inner", &self.can_into_inner())
.field("terminal_padding", &self.finished)
.field("finalized", &self.finalized)
.field("failed", &self.failed)
.finish()
}
}
impl<W, A, const PAD: bool> Decoder<W, A, PAD>
where
W: Write,
A: Alphabet,
{
pub fn try_finish(&mut self) -> io::Result<()> {
if self.failed {
return Err(stream_decoder_failed_error());
}
if !self.finalized {
self.queue_pending_final()?;
self.finalized = true;
}
self.flush()
}
pub fn finish(mut self) -> io::Result<W> {
self.try_finish()?;
Ok(self.take_inner())
}
fn queue_pending_final(&mut self) -> io::Result<()> {
let mut decoded = [0u8; 3];
let step = match self.driver.finish(&mut decoded) {
Ok(step) => step,
Err(err) => {
crate::wipe_bytes(&mut decoded);
self.failed = true;
self.clear_pending();
return Err(err);
}
};
let produced = step.progress().output_produced();
let result = self.output.push_slice(&decoded[..produced]);
crate::wipe_bytes(&mut decoded);
if result.is_err() {
self.failed = true;
}
result
}
fn queue_update(&mut self, input: &[u8]) -> io::Result<usize> {
let mut decoded = [0u8; 3];
let step = match self.driver.update(input, &mut decoded) {
Ok(step) => step,
Err(err) => {
crate::wipe_bytes(&mut decoded);
self.failed = true;
self.clear_pending();
return Err(err);
}
};
let progress = step.progress();
let result = self
.output
.push_slice(&decoded[..progress.output_produced()]);
crate::wipe_bytes(&mut decoded);
if result.is_err() {
self.failed = true;
}
result?;
self.finished = self.driver.has_terminal_padding();
Ok(progress.input_consumed())
}
fn drain_output(&mut self) -> io::Result<()> {
let mut chunk = [0u8; 1024];
while !self.output.is_empty() {
let pending = self.output.copy_front(&mut chunk);
let result = self.inner_mut().write(&chunk[..pending]);
crate::wipe_bytes(&mut chunk[..pending]);
match result {
Ok(0) => {
return Err(io::Error::new(
io::ErrorKind::WriteZero,
"base64 stream decoder could not drain buffered output",
));
}
Ok(written) => {
if written > pending {
self.failed = true;
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"wrapped writer reported more bytes than provided",
));
}
self.output.discard_front(written);
}
Err(err) => return Err(err),
}
}
Ok(())
}
}
impl<W, A, const PAD: bool> Write for Decoder<W, A, PAD>
where
W: Write,
A: Alphabet,
{
fn write(&mut self, input: &[u8]) -> io::Result<usize> {
if self.failed {
return Err(stream_decoder_failed_error());
}
if input.is_empty() {
self.drain_output()?;
return Ok(0);
}
self.drain_output()?;
if self.finalized {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
"base64 stream decoder received input after finalization",
));
}
if self.finished {
self.failed = true;
return Err(trailing_input_after_padding_error());
}
let mut consumed = 0;
while consumed < input.len() {
let pending = self.pending_len();
let take = (4 - pending).min(input.len() - consumed);
if pending + take == 4 && self.output.available_capacity() < 3 {
break;
}
match self.queue_update(&input[consumed..consumed + take]) {
Ok(accepted) => consumed += accepted,
Err(_) if consumed != 0 => return Ok(consumed),
Err(err) => return Err(err),
}
if self.finished || take < 4 - pending {
break;
}
}
Ok(consumed)
}
fn flush(&mut self) -> io::Result<()> {
if self.failed {
return Err(stream_decoder_failed_error());
}
self.drain_output()?;
self.inner_mut().flush()
}
}