use std::mem::MaybeUninit;
use crate::core::dictionary::DictionaryView;
use crate::core::types::{MAX_TOKEN_SIZE, Token};
use crate::core::validate::{InvalidColumn, panic_malformed};
mod copy;
pub const DECODE_PADDING: usize = MAX_TOKEN_SIZE;
#[inline]
pub fn decoded_len<V: DictionaryView>(codes: &[Token], dict: V) -> usize {
let n = dict.num_tokens();
let mut sum = 0usize;
for &c in codes {
if (c as usize) >= n {
panic_malformed(InvalidColumn::CodeOutOfRange);
}
let len = unsafe { dict.token_len_unchecked(c) };
sum = sum
.checked_add(len)
.unwrap_or_else(|| panic_malformed(InvalidColumn::DecodedLenOverflow));
}
sum
}
pub unsafe fn decode_into<V: DictionaryView>(
codes: &[Token],
dict: V,
out: &mut [MaybeUninit<u8>],
) -> usize {
let ntok = dict.num_tokens();
let dst = out.as_mut_ptr().cast::<u8>();
let mut w = 0usize;
for &code in codes {
if code as usize >= ntok {
panic_malformed(InvalidColumn::CodeOutOfRange);
}
unsafe {
let src = dict.token_ptr(code);
let len = dict.token_len_unchecked(code);
copy::copy16(src, dst.add(w));
w += len;
}
}
w
}
#[derive(Debug, Copy, Clone, Eq, PartialEq)]
pub struct OutputTooSmall;
impl std::fmt::Display for OutputTooSmall {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("output buffer too small for decoded bytes")
}
}
impl std::error::Error for OutputTooSmall {}
pub fn try_decode_into<V: DictionaryView>(
codes: &[Token],
dict: V,
out: &mut [MaybeUninit<u8>],
) -> Result<usize, OutputTooSmall> {
let ntok = dict.num_tokens();
let cap = out.len();
let dst = out.as_mut_ptr().cast::<u8>();
let mut w = 0usize;
let mut i = 0usize;
while i < codes.len() {
let batch = (cap - w) / MAX_TOKEN_SIZE;
if batch == 0 {
break;
}
let end = (i + batch).min(codes.len());
for &code in &codes[i..end] {
if code as usize >= ntok {
panic_malformed(InvalidColumn::CodeOutOfRange);
}
unsafe {
let len = dict.token_len_unchecked(code);
copy::copy16(dict.token_ptr(code), dst.add(w));
w += len;
}
}
i = end;
}
for &code in &codes[i..] {
if code as usize >= ntok {
panic_malformed(InvalidColumn::CodeOutOfRange);
}
let len = unsafe { dict.token_len_unchecked(code) };
if w + len > cap {
return Err(OutputTooSmall);
}
unsafe {
copy::copy_token_bytes(dict.token_ptr(code), dst.add(w), len);
}
w += len;
}
Ok(w)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::core::dictionary::{CompactDictionaryView, Dictionary};
fn padded(tokens: &[&[u8]]) -> (Vec<u8>, Vec<u32>) {
let mut bytes = Vec::new();
let mut offsets = vec![0u32];
for t in tokens {
bytes.extend_from_slice(t);
offsets.push(bytes.len() as u32);
}
bytes.resize(bytes.len() + MAX_TOKEN_SIZE, 0); (bytes, offsets)
}
fn expected(tokens: &[&[u8]], codes: &[Token]) -> Vec<u8> {
codes
.iter()
.flat_map(|&c| tokens[c as usize].iter().copied())
.collect()
}
fn vec_decode<V: DictionaryView>(codes: &[Token], dict: V) -> Vec<u8> {
let n = decoded_len(codes, dict);
let mut out = Vec::with_capacity(n + DECODE_PADDING);
let w = unsafe { decode_into(codes, dict, out.spare_capacity_mut()) };
unsafe { out.set_len(w) };
out
}
fn check(tokens: &[&[u8]], codes: &[Token]) {
let (bytes, offsets) = padded(tokens);
let view = CompactDictionaryView::from_raw(&bytes, &offsets);
let want = expected(tokens, codes);
assert_eq!(decoded_len(codes, view), want.len());
assert_eq!(vec_decode(codes, view), want, "compact");
let wide = view.to_wide();
assert_eq!(vec_decode(codes, wide.as_view()), want, "wide");
}
#[test]
fn decodes_mixed_length_tokens() {
let tokens: &[&[u8]] = &[b"a", b"bc", b"def", b"ghij"];
let codes: Vec<Token> = (0..40).map(|i| (i % 4) as Token).collect();
check(tokens, &codes);
}
#[test]
fn decodes_full_width_last_token() {
let full = vec![b'z'; MAX_TOKEN_SIZE];
let tokens: &[&[u8]] = &[b"x", &full];
let codes: Vec<Token> = (0..40).map(|i| (i % 2) as Token).collect();
check(tokens, &codes);
}
#[test]
fn decodes_short_final_token() {
let tokens: &[&[u8]] = &[b"a", b"bcde"];
let mut codes: Vec<Token> = (0..40).map(|i| (i % 2) as Token).collect();
*codes.last_mut().unwrap() = 0;
check(tokens, &codes);
}
#[test]
fn decodes_all_tail_length_buckets() {
let t = [
vec![b'a'; 1],
vec![b'b'; 3],
vec![b'c'; 5],
vec![b'd'; 11],
vec![b'e'; 15],
vec![b'f'; 16],
];
let tokens: Vec<&[u8]> = t.iter().map(Vec::as_slice).collect();
let codes: Vec<Token> = (0..40).map(|i| (i % t.len()) as Token).collect();
check(&tokens, &codes);
}
#[test]
fn decodes_empty_code_stream() {
check(&[b"a", b"b"], &[]);
}
#[test]
fn single_token_decode() {
let tokens: &[&[u8]] = &[b"hello", b"world"];
check(tokens, &[0]);
check(tokens, &[1]);
}
#[test]
#[should_panic(expected = "code index out of range")]
fn decode_into_panics_on_out_of_range_code() {
let (bytes, offsets) = padded(&[b"a", b"b"]);
let view = CompactDictionaryView::from_raw(&bytes, &offsets);
let mut out = vec![MaybeUninit::uninit(); 64];
unsafe { decode_into(&[0, 5], view, &mut out) };
}
#[test]
#[should_panic(expected = "code index out of range")]
fn decoded_len_panics_on_out_of_range_code() {
let (bytes, offsets) = padded(&[b"a", b"b"]);
let view = CompactDictionaryView::from_raw(&bytes, &offsets);
let _ = decoded_len(&[0, 5], view);
}
#[test]
fn try_decode_matches_decode_into() {
let tokens: &[&[u8]] = &[b"a", b"bc", b"def", b"ghij"];
let codes: Vec<Token> = (0..40).map(|i| (i % 4) as Token).collect();
let (bytes, offsets) = padded(tokens);
let view = CompactDictionaryView::from_raw(&bytes, &offsets);
let want = expected(tokens, &codes);
for cap in [want.len(), want.len() + DECODE_PADDING] {
let mut out = vec![MaybeUninit::uninit(); cap];
let w = try_decode_into(&codes, view, &mut out).unwrap();
assert_eq!(w, want.len());
let got = unsafe { std::slice::from_raw_parts(out.as_ptr().cast::<u8>(), w) };
assert_eq!(got, want.as_slice(), "cap {cap}");
}
}
#[test]
fn try_decode_fits_exact_and_rejects_one_short() {
let tokens: &[&[u8]] = &[b"a", b"bcde"];
let codes: Vec<Token> = vec![1, 1, 1]; let (bytes, offsets) = padded(tokens);
let view = CompactDictionaryView::from_raw(&bytes, &offsets);
assert_eq!(
try_decode_into(&codes, view, &mut [MaybeUninit::uninit(); 12]),
Ok(12)
);
assert_eq!(
try_decode_into(&codes, view, &mut [MaybeUninit::uninit(); 11]),
Err(OutputTooSmall)
);
}
#[test]
#[should_panic(expected = "code index out of range")]
fn try_decode_panics_on_out_of_range_code() {
let (bytes, offsets) = padded(&[b"a", b"b"]);
let view = CompactDictionaryView::from_raw(&bytes, &offsets);
let mut out = vec![MaybeUninit::uninit(); 64];
let _ = try_decode_into(&[0, 5], view, &mut out);
}
#[test]
fn try_decode_empty_codes_into_empty_buffer() {
let (bytes, offsets) = padded(&[b"a", b"b"]);
let view = CompactDictionaryView::from_raw(&bytes, &offsets);
assert_eq!(try_decode_into(&[], view, &mut []), Ok(0));
}
#[test]
fn try_decode_tail_reproduces_every_length_bucket() {
for len in [1usize, 2, 3, 4, 5, 7, 8, 11, 15] {
let tok = vec![b'a' + (len as u8); len];
let (bytes, offsets) = padded(&[tok.as_slice()]);
let view = CompactDictionaryView::from_raw(&bytes, &offsets);
let mut out = vec![MaybeUninit::uninit(); len]; let w = try_decode_into(&[0], view, &mut out).unwrap();
assert_eq!(w, len, "len {len}");
let got = unsafe { std::slice::from_raw_parts(out.as_ptr().cast::<u8>(), w) };
assert_eq!(got, tok.as_slice(), "len {len}");
}
}
}