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;
13pub 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 #[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 #[must_use]
64 pub fn get_ref(&self) -> &W {
65 self.inner_ref()
66 }
67
68 pub fn get_mut(&mut self) -> &mut W {
70 self.inner_mut()
71 }
72
73 #[must_use]
79 pub fn into_inner(mut self) -> W {
80 self.take_inner()
81 }
82
83 #[must_use]
85 pub const fn is_failed(&self) -> bool {
86 self.failed
87 }
88
89 #[must_use]
91 pub const fn is_finalized(&self) -> bool {
92 self.finalized
93 }
94
95 #[must_use]
98 pub const fn pending_len(&self) -> usize {
99 self.pending_len
100 }
101
102 #[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}