use super::{Dictionary, DictionaryView, WideDictionary};
use crate::core::types::{MAX_TOKEN_SIZE, Token};
use crate::core::validate::InvalidColumn;
pub(crate) fn pad_raw(bytes: &mut Vec<u8>, offsets: &[u32]) {
if offsets.len() < 2 {
return;
}
let last_token_start = offsets[offsets.len() - 2] as usize;
let required = last_token_start + MAX_TOKEN_SIZE;
if bytes.len() < required {
bytes.resize(required, 0);
}
}
fn validate_compact(bytes: &[u8], offsets: &[u32]) -> Result<(), InvalidColumn> {
let mut last_offset = 0usize;
for w in offsets.windows(2) {
let (s, e) = (w[0], w[1]);
if e < s {
return Err(InvalidColumn::NonDecreasingOffsets);
}
if e == s {
return Err(InvalidColumn::EmptyToken);
}
if (e - s) as usize > MAX_TOKEN_SIZE {
return Err(InvalidColumn::TokenTooLarge);
}
last_offset = s as usize;
}
if offsets.len() >= 2 && last_offset + MAX_TOKEN_SIZE > bytes.len() {
return Err(InvalidColumn::MissingPadding);
}
let mut seen = [false; 256];
let mut prev: &[u8] = &[];
for w in offsets.windows(2) {
let token = &bytes[w[0] as usize..w[1] as usize];
if prev >= token {
return Err(InvalidColumn::UnsortedTokens);
}
if token.len() == 1 {
seen[token[0] as usize] = true;
}
prev = token;
}
if seen.iter().any(|&present| !present) {
return Err(InvalidColumn::IncompleteAlphabet);
}
Ok(())
}
pub fn code_bits_for_num_tokens(num_tokens: usize) -> u8 {
debug_assert!(
num_tokens >= 1,
"log2(0) is undefined; num_tokens must be >= 1"
);
if num_tokens <= 1 {
1
} else {
((num_tokens as u32 - 1).ilog2() + 1) as u8
}
}
#[derive(Default, Debug, Clone)]
pub struct CompactDictionary {
bytes: Vec<u8>,
offsets: Vec<u32>,
}
impl CompactDictionary {
#[inline]
pub fn num_tokens(&self) -> usize {
self.offsets.len().saturating_sub(1)
}
#[inline]
pub fn bytes(&self) -> &[u8] {
&self.bytes
}
#[inline]
pub fn offsets(&self) -> &[u32] {
&self.offsets
}
#[inline]
pub fn into_raw(self) -> (Vec<u8>, Vec<u32>) {
(self.bytes, self.offsets)
}
#[inline]
pub(crate) fn from_raw(bytes: Vec<u8>, offsets: Vec<u32>) -> Self {
Self { bytes, offsets }
}
pub fn validate(bytes: Vec<u8>, offsets: Vec<u32>) -> Result<Self, InvalidColumn> {
validate_compact(&bytes, &offsets)?;
Ok(Self::from_raw(bytes, offsets))
}
pub unsafe fn new_unchecked(bytes: Vec<u8>, offsets: Vec<u32>) -> Self {
Self::from_raw(bytes, offsets)
}
#[inline]
pub fn logical_len(&self) -> usize {
self.offsets.last().copied().unwrap_or(0) as usize
}
#[inline]
pub fn code_bits(&self) -> u8 {
code_bits_for_num_tokens(self.num_tokens())
}
#[inline]
pub fn to_wide(&self) -> WideDictionary {
self.as_view().to_wide()
}
}
impl Dictionary for CompactDictionary {
type View<'a> = CompactDictionaryView<'a>;
#[inline]
fn as_view(&self) -> CompactDictionaryView<'_> {
CompactDictionaryView {
bytes: &self.bytes,
offsets: &self.offsets,
}
}
}
#[derive(Copy, Clone, Debug)]
pub struct CompactDictionaryView<'a> {
bytes: &'a [u8],
offsets: &'a [u32],
}
impl<'a> CompactDictionaryView<'a> {
#[inline]
pub(crate) fn from_raw(bytes: &'a [u8], offsets: &'a [u32]) -> Self {
Self { bytes, offsets }
}
pub fn validate(bytes: &'a [u8], offsets: &'a [u32]) -> Result<Self, InvalidColumn> {
validate_compact(bytes, offsets)?;
Ok(Self::from_raw(bytes, offsets))
}
pub unsafe fn new_unchecked(bytes: &'a [u8], offsets: &'a [u32]) -> Self {
Self::from_raw(bytes, offsets)
}
#[inline]
pub fn code_bits(&self) -> u8 {
code_bits_for_num_tokens(self.num_tokens())
}
pub fn to_wide(&self) -> WideDictionary {
let n = self.num_tokens();
let mut data = vec![0u8; n * MAX_TOKEN_SIZE];
let mut lens = vec![0u8; n];
let src = self.bytes.as_ptr();
let dst = data.as_mut_ptr();
for id in 0..n {
let (off, end) = unsafe {
(
*self.offsets.get_unchecked(id) as usize,
*self.offsets.get_unchecked(id + 1) as usize,
)
};
unsafe { *lens.get_unchecked_mut(id) = (end - off) as u8 };
unsafe {
std::ptr::copy_nonoverlapping(
src.add(off),
dst.add(id * MAX_TOKEN_SIZE),
MAX_TOKEN_SIZE,
);
}
}
WideDictionary::from_raw(data, lens)
}
}
impl DictionaryView for CompactDictionaryView<'_> {
#[inline]
fn num_tokens(&self) -> usize {
self.offsets.len().saturating_sub(1)
}
#[inline]
fn token(&self, id: Token) -> &[u8] {
let begin = self.offsets[id as usize] as usize;
let end = self.offsets[id as usize + 1] as usize;
&self.bytes[begin..end]
}
#[inline]
fn token_len(&self, id: Token) -> usize {
(self.offsets[id as usize + 1] - self.offsets[id as usize]) as usize
}
#[inline]
unsafe fn token_ptr(&self, id: Token) -> *const u8 {
unsafe {
self.bytes
.as_ptr()
.add(*self.offsets.get_unchecked(id as usize) as usize)
}
}
#[inline]
unsafe fn token_len_unchecked(&self, id: Token) -> usize {
unsafe {
(*self.offsets.get_unchecked(id as usize + 1)
- *self.offsets.get_unchecked(id as usize)) as usize
}
}
}
impl<'a> From<&'a CompactDictionary> for CompactDictionaryView<'a> {
#[inline]
fn from(d: &'a CompactDictionary) -> Self {
d.as_view()
}
}
#[cfg(test)]
mod tests {
use super::*;
fn dict(offsets: Vec<u32>, bytes: &[u8]) -> CompactDictionary {
CompactDictionary::from_raw(bytes.to_vec(), offsets)
}
#[test]
fn num_tokens_zero_when_offsets_empty() {
assert_eq!(CompactDictionary::default().num_tokens(), 0);
}
#[test]
fn num_tokens_is_offsets_len_minus_one() {
assert_eq!(dict(vec![0, 3, 5, 8], b"").num_tokens(), 3);
}
#[test]
fn token_returns_correct_slice() {
let d = dict(vec![0, 1, 3, 6], b"abcdef");
let v = d.as_view();
assert_eq!(v.token(0), b"a");
assert_eq!(v.token(1), b"bc");
assert_eq!(v.token(2), b"def");
assert_eq!(v.token_len(2), 3);
}
#[test]
fn code_bits_is_ceil_log2_num_tokens() {
assert_eq!(dict(vec![0; 257], b"").code_bits(), 8); assert_eq!(dict(vec![0; 258], b"").code_bits(), 9); assert_eq!(dict(vec![0; 513], b"").code_bits(), 9); assert_eq!(dict(vec![0; 514], b"").code_bits(), 10); assert_eq!(dict(vec![0; 65_537], b"").code_bits(), 16); }
#[test]
fn pad_raw_extends_to_max_token_read() {
let mut bytes = b"abc".to_vec();
pad_raw(&mut bytes, &[0, 1, 3]);
assert_eq!(bytes.len(), 1 + MAX_TOKEN_SIZE); }
#[test]
fn pad_raw_is_idempotent() {
let mut bytes = b"abc".to_vec();
let offsets = [0u32, 1, 3];
pad_raw(&mut bytes, &offsets);
let len = bytes.len();
pad_raw(&mut bytes, &offsets);
assert_eq!(bytes.len(), len);
}
#[test]
fn pad_raw_tops_up_insufficient_trailing_bytes() {
let mut bytes = vec![b'a', b'b', b'c', 0];
pad_raw(&mut bytes, &[0, 1, 3]);
assert_eq!(bytes.len(), 1 + MAX_TOKEN_SIZE);
}
#[test]
fn pad_raw_noop_for_full_width_last_token() {
let mut bytes = vec![b'z'; MAX_TOKEN_SIZE];
pad_raw(&mut bytes, &[0, MAX_TOKEN_SIZE as u32]);
assert_eq!(bytes.len(), MAX_TOKEN_SIZE);
}
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 conformant(extra: &[&[u8]]) -> (Vec<u8>, Vec<u32>) {
let mut toks: Vec<Vec<u8>> = (0u16..256).map(|b| vec![b as u8]).collect();
for &t in extra {
toks.push(t.to_vec());
}
toks.sort();
toks.dedup();
let refs: Vec<&[u8]> = toks.iter().map(Vec::as_slice).collect();
padded(&refs)
}
fn check(bytes: Vec<u8>, offsets: Vec<u32>) -> Result<(), InvalidColumn> {
CompactDictionary::validate(bytes, offsets).map(|_| ())
}
#[test]
fn validate_accepts_conformant() {
let (bytes, offsets) = conformant(&[b"bc", b"def"]);
assert_eq!(check(bytes, offsets), Ok(()));
}
#[test]
fn validate_classifies_safety_corruption() {
let mut bytes = b"ab".to_vec();
bytes.resize(2 + MAX_TOKEN_SIZE, 0);
assert_eq!(
check(bytes, vec![0, 2, 1]),
Err(InvalidColumn::NonDecreasingOffsets)
);
assert_eq!(
check(vec![0u8; MAX_TOKEN_SIZE], vec![0, 0]),
Err(InvalidColumn::EmptyToken)
);
assert_eq!(
check(vec![b'x'; 20 + MAX_TOKEN_SIZE], vec![0, 20]),
Err(InvalidColumn::TokenTooLarge)
);
assert_eq!(
check(b"abc".to_vec(), vec![0, 1, 3]),
Err(InvalidColumn::MissingPadding)
);
}
#[test]
fn validate_classifies_conformance_corruption() {
let (bytes, offsets) = padded(&[&[1u8], &[0u8]]);
assert_eq!(check(bytes, offsets), Err(InvalidColumn::UnsortedTokens));
let (bytes, offsets) = padded(&[&[0u8], &[1u8], &[2u8]]);
assert_eq!(
check(bytes, offsets),
Err(InvalidColumn::IncompleteAlphabet)
);
}
#[test]
fn new_unchecked_matches_validate() {
let (bytes, offsets) = conformant(&[b"bc"]);
let checked = CompactDictionary::validate(bytes.clone(), offsets.clone()).unwrap();
let trusted = unsafe { CompactDictionary::new_unchecked(bytes, offsets) };
assert_eq!(checked.bytes(), trusted.bytes());
assert_eq!(checked.offsets(), trusted.offsets());
}
#[test]
fn into_raw_returns_buffers_and_round_trips() {
let (bytes, offsets) = conformant(&[b"bc", b"def"]);
let num_tokens = offsets.len() - 1;
let dict = CompactDictionary::validate(bytes.clone(), offsets.clone()).unwrap();
let (raw_bytes, raw_offsets) = dict.into_raw();
assert_eq!(raw_bytes, bytes);
assert_eq!(raw_offsets, offsets);
let rebuilt = CompactDictionary::validate(raw_bytes, raw_offsets).unwrap();
assert_eq!(rebuilt.num_tokens(), num_tokens);
}
#[test]
fn view_validate_yields_usable_view() {
let (bytes, offsets) = conformant(&[b"bc"]);
let view = CompactDictionaryView::validate(&bytes, &offsets).unwrap();
assert_eq!(view.num_tokens(), 257); assert_eq!(view.token(0), &[0u8]);
let raw: &[u8] = b"abc";
assert!(CompactDictionaryView::validate(raw, &[0, 1, 3]).is_err());
}
}