#![allow(clippy::cast_possible_truncation, clippy::cast_sign_loss)]
#![allow(unsafe_code)]
mod simd;
use std::fmt;
pub const SYMBOL_COUNT: u8 = 94;
const ESCAPE_RADIX: u8 = 93;
pub const DECIMAL_MAX_LEN: usize = 10;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct DecodeError;
impl fmt::Display for DecodeError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("invalid base94 data")
}
}
impl std::error::Error for DecodeError {}
pub fn encode_into(out: &mut Vec<u8>, src: &[u8], kf: u32) {
const UNROLL: usize = 8;
let kf8 = kf as u8;
let total = src.len() + simd::count_sub_ge(src, kf8, ESCAPE_RADIX);
let start = out.len();
out.reserve(total + 16);
let dst = out.as_mut_ptr();
let mut p = start;
let simd_consumed = simd::encode_simd(dst, &mut p, src, kf8);
let mut idx = simd_consumed;
while idx + UNROLL <= src.len() {
let chunk: [u8; UNROLL] = src[idx..idx + UNROLL].try_into().expect("fixed size");
let mut offs = [0usize; UNROLL];
let mut cursor = 0usize;
for (k, &b) in chunk.iter().enumerate() {
offs[k] = cursor;
cursor += 1 + usize::from(b.wrapping_sub(kf8) >= ESCAPE_RADIX);
}
for (k, &b) in chunk.iter().enumerate() {
let v = b.wrapping_sub(kf8);
let esc = u8::from(v >= ESCAPE_RADIX);
let q2 = u8::from(v >= 2 * ESCAPE_RADIX);
let c1 = if esc != 0 {
0x7d + q2
} else {
0x20 + v
};
let c2 = 0x20u8.wrapping_add(
v.wrapping_sub(ESCAPE_RADIX)
.wrapping_sub(ESCAPE_RADIX.wrapping_mul(q2)),
);
unsafe {
let q = dst.add(p + offs[k]);
*q = c1;
*q.add(1) = c2;
}
}
p += cursor;
idx += UNROLL;
}
for &b in &src[idx..] {
let v = b.wrapping_sub(kf8);
let esc = u8::from(v >= ESCAPE_RADIX);
let q2 = u8::from(v >= 2 * ESCAPE_RADIX);
let c1 = if esc != 0 {
0x7d + q2
} else {
0x20 + v
};
let c2 = 0x20u8.wrapping_add(
v.wrapping_sub(ESCAPE_RADIX)
.wrapping_sub(ESCAPE_RADIX.wrapping_mul(q2)),
);
unsafe {
*dst.add(p) = c1;
*dst.add(p + 1) = c2;
}
p += 1 + usize::from(esc);
}
unsafe { out.set_len(start + total) };
}
#[must_use]
pub fn encoded_len(src: &[u8], kf: u32) -> usize {
src.len() + simd::count_sub_ge(src, kf as u8, ESCAPE_RADIX)
}
pub fn decode_into(out: &mut Vec<u8>, src: &[u8], kf: u32) -> Result<(), DecodeError> {
if !simd::all_ge(src, 0x20) {
return Err(DecodeError);
}
let kf8 = kf as u8;
let kf16 = u16::from(kf8);
let start = out.len();
out.reserve(src.len() + 16);
let dst = out.as_mut_ptr();
let mut i;
let mut p = start;
let n = src.len();
match simd::decode_simd(dst, &mut p, src, kf8) {
Ok((consumed, _)) | Err((consumed, _)) => i = consumed,
}
while i + 1 < n {
let b = u16::from(src[i]) - 0x20;
let b2 = u16::from(src[i + 1]) - 0x20;
let esc = u16::from(b >= u16::from(ESCAPE_RADIX));
let v_esc = (b.wrapping_sub(92))
.wrapping_mul(u16::from(ESCAPE_RADIX))
.wrapping_add(b2);
let bad = esc
& (u16::from(b > 94)
| u16::from(b2 > u16::from(ESCAPE_RADIX))
| u16::from(v_esc > 0xff));
if bad != 0 {
return Err(DecodeError);
}
let val = if esc != 0 {
v_esc
} else {
b
};
unsafe { *dst.add(p) = val.wrapping_add(kf16) as u8 };
p += 1;
i += 1 + esc as usize;
}
if i < n {
let b = src[i] - 0x20;
if b >= ESCAPE_RADIX {
return Err(DecodeError);
}
unsafe { *dst.add(p) = b.wrapping_add(kf8) };
p += 1;
}
unsafe { out.set_len(p) };
Ok(())
}
#[must_use]
pub fn decimal_encode(v: u64, out: &mut [u8; DECIMAL_MAX_LEN]) -> usize {
let mut n = v;
let mut len = 0;
loop {
out[len] = (n % u64::from(SYMBOL_COUNT)) as u8 + 0x20;
len += 1;
n /= u64::from(SYMBOL_COUNT);
if n == 0 {
break;
}
}
out[..len].reverse();
len
}
pub fn decimal_decode(s: &[u8]) -> Result<u64, DecodeError> {
if s.is_empty() {
return Err(DecodeError);
}
let mut n: u64 = 0;
for &c in s {
if c < 0x20 {
return Err(DecodeError);
}
let d = c - 0x20;
if d >= SYMBOL_COUNT {
return Err(DecodeError);
}
n = n
.checked_mul(u64::from(SYMBOL_COUNT))
.and_then(|n| n.checked_add(u64::from(d)))
.ok_or(DecodeError)?;
}
Ok(n)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn roundtrip_and_printable() {
for kf in [0u32, 1, 93, 94, 0xff, 0xdead_beef] {
let original: Vec<u8> = (0..=u8::MAX).collect();
let mut encoded = Vec::new();
encode_into(&mut encoded, &original, kf);
assert!(encoded.iter().all(|&c| (0x20..=0x7e).contains(&c)));
assert_eq!(encoded.len(), encoded_len(&original, kf));
let mut decoded = Vec::new();
decode_into(&mut decoded, &encoded, kf).unwrap();
assert_eq!(decoded, original);
}
}
#[test]
fn roundtrip_large_and_edge_shapes() {
let mut s = 0x0bad_c0deu64;
let mut step = || {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
(s >> 56) as u8
};
let all_escape: Vec<u8> = (0..2000)
.map(|i| 93u8.wrapping_add(i as u8 % 163))
.collect();
let no_escape: Vec<u8> = vec![7u8; 2001];
for dataset in [all_escape, no_escape] {
for kf in [0u32, 0x5a5a_5a5a] {
let mut encoded = Vec::new();
encode_into(&mut encoded, &dataset, kf);
assert_eq!(encoded.len(), encoded_len(&dataset, kf));
let mut decoded = Vec::new();
decode_into(&mut decoded, &encoded, kf).unwrap();
assert_eq!(decoded, dataset);
}
}
let mut random: Vec<u8> = (0..65_537).map(|_| step()).collect();
random[0] = 0x93; let mut encoded = Vec::new();
encode_into(&mut encoded, &random, 154_543_927);
let mut decoded = Vec::new();
decode_into(&mut decoded, &encoded, 154_543_927).unwrap();
assert_eq!(decoded, random);
}
#[test]
fn decode_uses_existing_out_prefix() {
let original = vec![0xde, 0xad, 0xbe, 0xef];
let mut encoded = Vec::new();
encode_into(&mut encoded, &original, 123);
let mut out = b"prefix".to_vec();
decode_into(&mut out, &encoded, 123).unwrap();
assert_eq!(out, [b"prefix".as_slice(), original.as_slice()].concat());
}
#[test]
fn escape_boundaries() {
let mut encoded = Vec::new();
encode_into(&mut encoded, &[0, 92, 93, 94, 255], 0);
assert_eq!(encoded, [0x20, 0x7c, 0x7d, 0x20, 0x7d, 0x21, 0x7e, 0x65]);
}
#[test]
fn decode_rejects_garbage() {
let mut out = Vec::new();
assert!(decode_into(&mut out, &[0x1f], 0).is_err());
assert!(decode_into(&mut out, &[0x7f], 0).is_err());
assert!(decode_into(&mut out, &[0x7d], 0).is_err()); assert!(decode_into(&mut out, &[0x7d, 0x7e], 0).is_err()); assert!(decode_into(&mut out, &[0x7e, 0x7e], 0).is_err()); assert!(decode_into(&mut out, &[0x7f, 0x20], 0).is_err());
assert!(out.is_empty(), "no partial output on failure");
}
fn decode_reference(out: &mut Vec<u8>, src: &[u8], kf: u32) -> Result<(), DecodeError> {
let kf8 = kf as u8;
let start = out.len();
let mut i = 0;
while i < src.len() {
let c = src[i];
if c < 0x20 {
out.truncate(start);
return Err(DecodeError);
}
let b = c - 0x20;
if b < ESCAPE_RADIX {
out.push(b.wrapping_add(kf8));
i += 1;
continue;
}
if b > 94 {
out.truncate(start);
return Err(DecodeError);
}
let Some(&c2) = src.get(i + 1) else {
out.truncate(start);
return Err(DecodeError);
};
if c2 < 0x20 {
out.truncate(start);
return Err(DecodeError);
}
let b2 = c2 - 0x20;
if b2 > ESCAPE_RADIX {
out.truncate(start);
return Err(DecodeError);
}
if b == 94 && b2 > 0xff - 2 * ESCAPE_RADIX {
out.truncate(start);
return Err(DecodeError);
}
let v = u32::from(b - ESCAPE_RADIX + 1) * u32::from(ESCAPE_RADIX) + u32::from(b2);
out.push((v as u8).wrapping_add(kf8));
i += 2;
}
Ok(())
}
#[test]
fn decode_fuzz_matches_reference() {
let mut s = 0x00c0_ffee_d00d_feedu64;
let mut next = || {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
s
};
for len in [
0usize, 1, 15, 16, 17, 31, 33, 64, 100, 199, 200, 201, 255, 256, 257, 513,
] {
for _ in 0..40 {
let input: Vec<u8> = (0..len)
.map(|_| {
let r = (next() >> 32) as u8;
match r % 8 {
0..=4 => 0x20 + (r % 93),
5 => 0x7b,
6 => 0x7d + (r % 2),
_ => r, }
})
.collect();
for kf in [0u32, 77, 0x5a5a_5a5a] {
let mut got = Vec::new();
let mut want = Vec::new();
let a = decode_into(&mut got, &input, kf);
let b = decode_reference(&mut want, &input, kf);
assert_eq!(a.is_ok(), b.is_ok(), "len={len} kf={kf} ok-ness");
if let (Ok(()), Ok(())) = (a, b) {
assert_eq!(got, want, "len={len} kf={kf}");
}
}
}
}
}
#[test]
fn decode_rejects_7e_follower_in_simd_block() {
let input: Vec<u8> = b"ppppppppp#i}~C}(".to_vec();
assert_eq!(input.len(), 16);
let mut out = Vec::new();
assert!(decode_into(&mut out, &input, 125).is_err());
assert_eq!(out, [] as [u8; 0]);
for shift in 0..15usize {
let mut v = vec![0x41u8; 16];
v[shift] = 0x7d;
v[shift + 1] = 0x7e;
let mut out = Vec::new();
assert!(decode_into(&mut out, &v, 0).is_err(), "shift {shift}");
}
let mut v = vec![0x41u8; 16];
v[0] = 0x7e;
v[1] = 0x7d;
let mut out = Vec::new();
assert!(decode_into(&mut out, &v, 0).is_err());
assert_eq!(out, [] as [u8; 0]);
}
#[test]
fn decode_legal_adjacent_escape_pair() {
let mut big = vec![0x7du8; 64];
big.extend_from_slice(&[0x41; 8]);
for kf in [0u32, 200] {
let mut got = Vec::new();
decode_into(&mut got, &big, kf).unwrap();
let mut want = Vec::new();
decode_reference(&mut want, &big, kf).unwrap();
assert_eq!(got, want);
assert_eq!(got.len(), 32 + 8);
}
}
#[test]
fn decimal_roundtrip() {
let mut buf = [0u8; DECIMAL_MAX_LEN];
for v in [0u64, 1, 93, 94, 830_583, u64::from(u32::MAX), u64::MAX] {
let len = decimal_encode(v, &mut buf);
let digits = &buf[..len];
assert!(digits.iter().all(|&c| c >= 0x20));
if v > 0 {
assert_ne!(digits[0], 0x20, "no leading zero");
}
assert_eq!(decimal_decode(digits).unwrap(), v);
if len <= 3 {
let mut padded = [0x20u8; 3];
padded[3 - len..].copy_from_slice(digits);
assert_eq!(decimal_decode(&padded).unwrap(), v);
}
}
}
#[test]
fn decimal_known_value() {
let mut buf = [0u8; DECIMAL_MAX_LEN];
let len = decimal_encode(8836, &mut buf);
assert_eq!(len, 3);
assert_eq!(&buf[..3], &[0x21, 0x20, 0x20]);
}
}