#![allow(unsafe_code)]
use super::{
Base32,
alphabet::{INVALID, Tables},
encode::{GROUP, GROUP_SYMBOLS, TAIL_SYMBOLS},
};
use crate::Buffer;
use std::{fmt, io, iter::FusedIterator, mem::MaybeUninit, slice};
const BYTE_BITS: u32 = 8;
const SYMBOL_BITS: u32 = 5;
const GROUP_BITS: u32 = 40;
impl Base32 {
pub fn encoder<'x, B: Buffer>(&self, out: &'x mut B) -> Encoder<'x, B> {
Encoder {
out,
tables: self.tables(),
pad: self.pads(),
pending: 0,
count: 0,
}
}
pub fn decoder<'x>(&self, input: &'x (impl AsRef<[u8]> + ?Sized)) -> Decoder<'x> {
Decoder::new(*self, input.as_ref())
}
pub fn decoder_from_iter<'x>(&self, input: slice::Iter<'x, u8>) -> Decoder<'x> {
Decoder::new(*self, input.as_slice())
}
}
pub struct Encoder<'x, B: Buffer> {
out: &'x mut B,
tables: &'static Tables,
pad: bool,
pending: u64,
count: usize,
}
impl<B: Buffer> Encoder<'_, B> {
#[inline]
pub fn push(&mut self, input: &[u8]) {
if self.count + input.len() < GROUP {
for &byte in input {
self.pending = (self.pending << BYTE_BITS) | byte as u64;
}
self.count += input.len();
} else {
self.push_groups(input);
}
}
pub fn finish(mut self) {
self.flush_tail();
}
fn push_groups(&mut self, input: &[u8]) {
let (head, input) = input
.split_at_checked((GROUP - self.count) % GROUP)
.unwrap_or((input, &[]));
if self.count > 0 {
for &byte in head {
self.pending = (self.pending << BYTE_BITS) | byte as u64;
}
let block = self.tables.encode_block(self.pending);
unsafe {
self.out.append_ascii(GROUP_SYMBOLS, |dst| {
let Some(slots) = dst.first_chunk_mut::<GROUP_SYMBOLS>() else {
return 0;
};
*slots = block.map(MaybeUninit::new);
GROUP_SYMBOLS
})
};
}
let (groups, tail) = input.as_chunks::<GROUP>();
let groups = groups.as_flattened();
if !groups.is_empty() {
let len = groups.len() / GROUP * GROUP_SYMBOLS;
let tables = self.tables;
unsafe {
self.out
.append_ascii(len, |dst| tables.encode_groups(groups, dst).1)
};
}
self.pending = tail
.iter()
.fold(0, |pending, &byte| (pending << BYTE_BITS) | byte as u64);
self.count = tail.len();
}
fn flush_tail(&mut self) {
if self.count == 0 {
return;
}
let bytes = self.pending.to_be_bytes();
let tail = bytes.get(bytes.len() - self.count..).unwrap_or_default();
let (tables, pad) = (self.tables, self.pad);
let len = if pad {
GROUP_SYMBOLS
} else {
TAIL_SYMBOLS.get(self.count).copied().unwrap_or_default()
};
unsafe {
self.out
.append_ascii(len, |dst| tables.encode_tail(pad, tail, dst))
};
self.pending = 0;
self.count = 0;
}
}
impl<B: Buffer> io::Write for Encoder<'_, B> {
#[inline]
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
self.push(buf);
Ok(buf.len())
}
#[inline]
fn write_all(&mut self, buf: &[u8]) -> io::Result<()> {
self.push(buf);
Ok(())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
impl<B: Buffer> Drop for Encoder<'_, B> {
fn drop(&mut self) {
self.flush_tail();
}
}
impl<B: Buffer> fmt::Debug for Encoder<'_, B> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Encoder")
.field("pad", &self.pad)
.field("pending", &self.count)
.finish_non_exhaustive()
}
}
#[derive(Debug, Clone)]
pub struct Decoder<'x> {
engine: Base32,
input: &'x [u8],
position: usize,
bits: u64,
count: u32,
}
impl<'x> Decoder<'x> {
fn new(engine: Base32, input: &'x [u8]) -> Self {
Decoder {
engine,
input,
position: 0,
bits: 0,
count: 0,
}
}
pub fn remaining(&self) -> &'x [u8] {
self.input
.get(
self.position
.saturating_sub((self.count / SYMBOL_BITS) as usize)..,
)
.unwrap_or_default()
}
}
impl Iterator for Decoder<'_> {
type Item = u8;
#[inline]
fn next(&mut self) -> Option<u8> {
if self.count < BYTE_BITS {
let (bits, symbols) =
self.engine
.tables()
.read_symbols(self.input, self.position, self.bits, self.count);
self.bits = bits;
self.position += symbols;
self.count += SYMBOL_BITS * symbols as u32;
if self.count < BYTE_BITS {
return None;
}
}
self.count -= BYTE_BITS;
Some((self.bits >> self.count) as u8)
}
fn size_hint(&self) -> (usize, Option<usize>) {
let symbols = self.input.len().saturating_sub(self.position);
let bits = self.count as usize + symbols * SYMBOL_BITS as usize;
(0, Some(bits / BYTE_BITS as usize))
}
}
impl FusedIterator for Decoder<'_> {}
impl Tables {
#[inline(never)]
fn read_symbols(&self, input: &[u8], position: usize, bits: u64, count: u32) -> (u64, usize) {
let rest = input.get(position..).unwrap_or_default();
if let Some(block) = rest
.first_chunk::<GROUP_SYMBOLS>()
.and_then(|group| self.decode_block(group))
{
return ((bits << GROUP_BITS) | block, GROUP_SYMBOLS);
}
let mut bits = bits;
let mut count = count;
let mut symbols = 0;
for &byte in rest {
let value = self.decode[byte as usize];
if value == INVALID || count >= BYTE_BITS {
break;
}
bits = (bits << SYMBOL_BITS) | value as u64;
count += SYMBOL_BITS;
symbols += 1;
}
(bits, symbols)
}
}