Skip to main content

base64_ng_tokio/
encoder_writer.rs

1use base64_ng::{Alphabet, Engine};
2use core::{
3    marker::PhantomData,
4    pin::Pin,
5    task::{Context, Poll, ready},
6};
7use tokio::io::{self, AsyncWrite};
8
9use crate::{encode_io_error, queue::OutputQueue, wipe_bytes};
10
11const ENCODE_INPUT_CAP: usize = 768;
12const ENCODE_OUTPUT_CAP: usize = 1024;
13/// Async writer that accepts raw bytes and writes Base64 to the wrapped writer.
14///
15/// `poll_write` may accept only part of the input, following normal
16/// [`AsyncWrite`] rules. Accepted bytes may remain buffered internally until
17/// a later write, [`AsyncWrite::poll_flush`], or [`AsyncWrite::poll_shutdown`].
18/// Shutdown is the finalization boundary: it encodes any trailing partial
19/// quantum, drains all buffered output, flushes, and then shuts down `inner`.
20///
21/// # Security
22///
23/// Internal cleanup is best-effort and limited to this adapter's fixed pending
24/// and output buffers. It cannot clear copies held by the wrapped writer, the
25/// caller's buffers, registers, caches, swap, or crash dumps.
26///
27/// I/O errors from the wrapped writer during drain do not set [`Self::is_failed`];
28/// only internal protocol or capacity violations latch a permanent failure.
29pub struct EncoderWriter<W, A, const PAD: bool>
30where
31    A: Alphabet,
32{
33    inner: Option<W>,
34    engine: Engine<A, PAD>,
35    pending: [u8; 2],
36    pending_len: usize,
37    output: OutputQueue<ENCODE_OUTPUT_CAP>,
38    finalized: bool,
39    failed: bool,
40    _alphabet: PhantomData<A>,
41}
42
43impl<W, A, const PAD: bool> EncoderWriter<W, A, PAD>
44where
45    A: Alphabet,
46{
47    /// Creates a new async Base64 encoder writer.
48    #[must_use]
49    pub fn new(inner: W, engine: Engine<A, PAD>) -> Self {
50        Self {
51            inner: Some(inner),
52            engine,
53            pending: [0; 2],
54            pending_len: 0,
55            output: OutputQueue::new(),
56            finalized: false,
57            failed: false,
58            _alphabet: PhantomData,
59        }
60    }
61
62    /// Returns a shared reference to the wrapped writer.
63    #[must_use]
64    pub fn get_ref(&self) -> &W {
65        self.inner_ref()
66    }
67
68    /// Returns a mutable reference to the wrapped writer.
69    pub fn get_mut(&mut self) -> &mut W {
70        self.inner_mut()
71    }
72
73    /// Consumes the adapter and returns the wrapped writer.
74    ///
75    /// This does not finalize pending input. Prefer
76    /// [`AsyncWriteExt::shutdown`](tokio::io::AsyncWriteExt::shutdown) before
77    /// calling this when the Base64 stream must be complete.
78    #[must_use]
79    pub fn into_inner(mut self) -> W {
80        self.take_inner()
81    }
82
83    /// Returns whether this adapter has encountered an unrecoverable error.
84    #[must_use]
85    pub const fn is_failed(&self) -> bool {
86        self.failed
87    }
88
89    /// Returns whether shutdown has finalized this adapter.
90    #[must_use]
91    pub const fn is_finalized(&self) -> bool {
92        self.finalized
93    }
94
95    /// Returns the number of raw bytes buffered until a full encode quantum is
96    /// available.
97    #[must_use]
98    pub const fn pending_len(&self) -> usize {
99        self.pending_len
100    }
101
102    /// Returns the number of encoded bytes currently buffered for `inner`.
103    #[must_use]
104    pub const fn buffered_output_len(&self) -> usize {
105        self.output.len()
106    }
107
108    fn clear_pending(&mut self) {
109        wipe_bytes(&mut self.pending);
110        self.pending_len = 0;
111    }
112
113    fn clear_output(&mut self) {
114        self.output.clear_all();
115    }
116
117    fn inner_ref(&self) -> &W {
118        match &self.inner {
119            Some(inner) => inner,
120            None => unreachable!("tokio encoder writer inner writer was already taken"),
121        }
122    }
123
124    fn inner_mut(&mut self) -> &mut W {
125        match &mut self.inner {
126            Some(inner) => inner,
127            None => unreachable!("tokio encoder writer inner writer was already taken"),
128        }
129    }
130
131    fn take_inner(&mut self) -> W {
132        match self.inner.take() {
133            Some(inner) => inner,
134            None => unreachable!("tokio encoder writer inner writer was already taken"),
135        }
136    }
137
138    fn queue_encoded_temp(&mut self, input: &[u8], encoded: &mut [u8]) -> io::Result<()> {
139        let written = match self.engine.encode_slice(input, encoded) {
140            Ok(written) => written,
141            Err(error) => {
142                wipe_bytes(encoded);
143                self.failed = true;
144                return Err(encode_io_error(error));
145            }
146        };
147
148        let result = self.output.push_slice(&encoded[..written]);
149        wipe_bytes(encoded);
150        if result.is_err() {
151            self.failed = true;
152        }
153        result
154    }
155
156    fn queue_pending_final(&mut self) -> io::Result<()> {
157        if self.pending_len == 0 {
158            return Ok(());
159        }
160
161        let mut pending = [0u8; 2];
162        pending[..self.pending_len].copy_from_slice(&self.pending[..self.pending_len]);
163        let pending_len = self.pending_len;
164        let mut encoded = [0u8; 4];
165        let result = self.queue_encoded_temp(&pending[..pending_len], &mut encoded);
166        wipe_bytes(&mut pending);
167        result?;
168        self.clear_pending();
169        Ok(())
170    }
171
172    fn process_input(&mut self, input: &[u8]) -> io::Result<usize> {
173        if input.is_empty() {
174            return Ok(0);
175        }
176
177        let mut consumed = 0;
178        if self.pending_len > 0 {
179            let needed = 3 - self.pending_len;
180            if input.len() < needed {
181                self.pending[self.pending_len..self.pending_len + input.len()]
182                    .copy_from_slice(input);
183                self.pending_len += input.len();
184                return Ok(input.len());
185            }
186
187            let mut quantum = [0u8; 3];
188            quantum[..self.pending_len].copy_from_slice(&self.pending[..self.pending_len]);
189            quantum[self.pending_len..].copy_from_slice(&input[..needed]);
190            let mut encoded = [0u8; 4];
191            let result = self.queue_encoded_temp(&quantum, &mut encoded);
192            wipe_bytes(&mut quantum);
193            result?;
194            self.clear_pending();
195            consumed += needed;
196        }
197
198        let remaining = &input[consumed..];
199        let full_len = remaining.len() / 3 * 3;
200        if full_len != 0 {
201            let max_by_queue = self.output.available_capacity() / 4 * 3;
202            let mut take = core::cmp::min(full_len, core::cmp::min(ENCODE_INPUT_CAP, max_by_queue));
203            take -= take % 3;
204
205            if take == 0 {
206                return Ok(consumed);
207            }
208
209            let mut encoded = [0u8; ENCODE_OUTPUT_CAP];
210            self.queue_encoded_temp(&remaining[..take], &mut encoded)?;
211            consumed += take;
212
213            if take < full_len {
214                return Ok(consumed);
215            }
216        }
217
218        let tail = &input[consumed..];
219        self.pending[..tail.len()].copy_from_slice(tail);
220        self.pending_len = tail.len();
221        consumed += tail.len();
222        Ok(consumed)
223    }
224}
225
226impl<W, A, const PAD: bool> Drop for EncoderWriter<W, A, PAD>
227where
228    A: Alphabet,
229{
230    fn drop(&mut self) {
231        self.clear_pending();
232        self.clear_output();
233    }
234}
235
236impl<W, A, const PAD: bool> EncoderWriter<W, A, PAD>
237where
238    W: AsyncWrite + Unpin,
239    A: Alphabet + Unpin,
240{
241    fn poll_drain_output(&mut self, context: &mut Context<'_>) -> Poll<io::Result<()>> {
242        let mut chunk = [0u8; ENCODE_OUTPUT_CAP];
243        while !self.output.is_empty() {
244            let pending = self.output.copy_front(&mut chunk);
245            let result = Pin::new(self.inner_mut()).poll_write(context, &chunk[..pending]);
246            wipe_bytes(&mut chunk[..pending]);
247            match result {
248                Poll::Pending => return Poll::Pending,
249                Poll::Ready(Ok(0)) => {
250                    return Poll::Ready(Err(io::Error::new(
251                        io::ErrorKind::WriteZero,
252                        "base64-ng-tokio encoder writer could not drain buffered output",
253                    )));
254                }
255                Poll::Ready(Ok(written)) => {
256                    if written > pending {
257                        self.failed = true;
258                        return Poll::Ready(Err(io::Error::new(
259                            io::ErrorKind::InvalidData,
260                            "wrapped async writer reported more bytes than provided",
261                        )));
262                    }
263                    self.output.discard_front(written);
264                }
265                Poll::Ready(Err(error)) => return Poll::Ready(Err(error)),
266            }
267        }
268
269        Poll::Ready(Ok(()))
270    }
271}
272
273impl<W, A, const PAD: bool> AsyncWrite for EncoderWriter<W, A, PAD>
274where
275    W: AsyncWrite + Unpin,
276    A: Alphabet + Unpin,
277{
278    fn poll_write(
279        mut self: Pin<&mut Self>,
280        context: &mut Context<'_>,
281        input: &[u8],
282    ) -> Poll<io::Result<usize>> {
283        if self.failed {
284            return Poll::Ready(Err(io::Error::other(
285                "base64-ng-tokio encoder writer is failed",
286            )));
287        }
288
289        ready!(self.poll_drain_output(context))?;
290        if self.finalized {
291            return Poll::Ready(Err(io::Error::new(
292                io::ErrorKind::InvalidInput,
293                "base64-ng-tokio encoder writer received input after shutdown",
294            )));
295        }
296
297        Poll::Ready(self.process_input(input))
298    }
299
300    fn poll_flush(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<io::Result<()>> {
301        if self.failed {
302            return Poll::Ready(Err(io::Error::other(
303                "base64-ng-tokio encoder writer is failed",
304            )));
305        }
306
307        ready!(self.poll_drain_output(context))?;
308        Pin::new(self.inner_mut()).poll_flush(context)
309    }
310
311    fn poll_shutdown(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<io::Result<()>> {
312        if self.failed {
313            return Poll::Ready(Err(io::Error::other(
314                "base64-ng-tokio encoder writer is failed",
315            )));
316        }
317
318        ready!(self.poll_drain_output(context))?;
319        if !self.finalized {
320            if let Err(error) = self.queue_pending_final() {
321                self.failed = true;
322                self.clear_output();
323                return Poll::Ready(Err(error));
324            }
325            self.finalized = true;
326        }
327        ready!(self.poll_drain_output(context))?;
328        ready!(Pin::new(self.inner_mut()).poll_flush(context))?;
329        Pin::new(self.inner_mut()).poll_shutdown(context)
330    }
331}