use std::mem::MaybeUninit;
use crate::core::dictionary::{CompactDictionary, CompactDictionaryView, Dictionary};
use crate::core::offset::Offset;
use crate::core::types::Token;
use crate::core::validate::{InvalidColumn, panic_malformed};
use crate::decoding;
use crate::search::{ContainsTable, PrefixQuery, contains, equals, starts_with, tokenize};
#[derive(Debug, Clone)]
pub struct Column<O: Offset> {
pub dict: CompactDictionary,
pub codes: Vec<Token>,
pub row_offsets: Vec<O>,
}
impl<O: Offset> Column<O> {
pub fn compress(bytes: &[u8], offsets: &[O], cfg: crate::Config) -> Result<Self, crate::Error> {
crate::compress(bytes, offsets, cfg)
}
#[inline]
pub fn view(&self) -> ColumnView<'_, O> {
ColumnView {
dict: self.dict.as_view(),
codes: &self.codes,
row_offsets: &self.row_offsets,
}
}
#[inline]
pub fn into_raw(self) -> (CompactDictionary, Vec<Token>, Vec<O>) {
(self.dict, self.codes, self.row_offsets)
}
}
#[derive(Copy, Clone, Debug)]
pub struct ColumnView<'a, O: Offset> {
pub dict: CompactDictionaryView<'a>,
pub codes: &'a [Token],
pub row_offsets: &'a [O],
}
impl<'a, O: Offset> ColumnView<'a, O> {
#[inline]
pub fn num_rows(&self) -> usize {
self.row_offsets.len().saturating_sub(1)
}
#[inline]
pub fn row_codes(&self, k: usize) -> &'a [Token] {
let a = self.row_offsets[k].to_usize();
let b = self.row_offsets[k + 1].to_usize();
if b < a || b > self.codes.len() {
panic_malformed(InvalidColumn::BadRowOffsets);
}
unsafe { self.codes.get_unchecked(a..b) }
}
#[inline]
pub fn decoded_len(&self) -> usize {
decoding::decoded_len(self.codes, self.dict)
}
#[inline]
pub fn row_decoded_len(&self, k: usize) -> usize {
decoding::decoded_len(self.row_codes(k), self.dict)
}
#[inline]
pub unsafe fn decompress_into(&self, out: &mut [MaybeUninit<u8>]) -> usize {
let wide = self.dict.to_wide();
unsafe { decoding::decode_into(self.codes, wide.as_view(), out) }
}
#[inline]
pub unsafe fn decompress_row_into(&self, k: usize, out: &mut [MaybeUninit<u8>]) -> usize {
unsafe { decoding::decode_into(self.row_codes(k), self.dict, out) }
}
pub fn rows_equal_to(&self, needle: &[u8]) -> Vec<usize> {
let query = tokenize(needle, self.dict);
self.select(|codes| equals(codes, &query))
}
pub fn rows_starting_with(&self, prefix: &[u8]) -> Vec<usize> {
let query = PrefixQuery::new(prefix, self.dict);
self.select(|codes| starts_with(codes, &query))
}
pub fn rows_containing(&self, pattern: &[u8]) -> Vec<usize> {
let table = ContainsTable::new(pattern, self.dict);
self.select(|codes| contains(codes, &table))
}
fn select(&self, pred: impl Fn(&[Token]) -> bool) -> Vec<usize> {
(0..self.num_rows())
.filter(|&k| pred(self.row_codes(k)))
.collect()
}
}
#[cfg(test)]
mod tests {
use crate::{ColumnView, Config, DECODE_PADDING, DEFAULT_CONFIG, MaxDictBits, compress};
fn pack(rows: &[&[u8]]) -> (Vec<u8>, Vec<u32>) {
let mut bytes = Vec::new();
let mut offsets = vec![0u32];
for r in rows {
bytes.extend_from_slice(r);
offsets.push(bytes.len() as u32);
}
(bytes, offsets)
}
fn decode_all(view: ColumnView<'_, u32>) -> Vec<u8> {
let mut buf = vec![std::mem::MaybeUninit::uninit(); view.decoded_len() + DECODE_PADDING];
let w = unsafe { view.decompress_into(&mut buf) };
let got = unsafe { std::slice::from_raw_parts(buf.as_ptr().cast::<u8>(), w) };
got.to_vec()
}
fn decode_row(view: ColumnView<'_, u32>, k: usize) -> Vec<u8> {
let mut buf =
vec![std::mem::MaybeUninit::uninit(); view.row_decoded_len(k) + DECODE_PADDING];
let w = unsafe { view.decompress_row_into(k, &mut buf) };
let got = unsafe { std::slice::from_raw_parts(buf.as_ptr().cast::<u8>(), w) };
got.to_vec()
}
#[test]
fn roundtrip_bulk_and_per_row() {
let rows: &[&[u8]] = &[b"alpha", b"", b"beta beta", b"gamma"];
let (bytes, offsets) = pack(rows);
let col = compress(&bytes, &offsets, DEFAULT_CONFIG).unwrap();
let view = col.view();
assert_eq!(view.decoded_len(), bytes.len());
assert_eq!(decode_all(view), bytes);
assert_eq!(view.num_rows(), rows.len());
for (k, row) in rows.iter().enumerate() {
assert_eq!(decode_row(view, k), *row, "row {k}");
}
}
#[test]
fn compact_and_wide_decode_agree() {
use crate::decode_into;
let rows: &[&[u8]] = &[b"alpha", b"", b"beta beta", b"gamma", b"alpha"];
let (bytes, offsets) = pack(rows);
let col = compress(&bytes, &offsets, DEFAULT_CONFIG).unwrap();
let view = col.view();
assert_eq!(decode_all(view), bytes);
let mut buf = vec![std::mem::MaybeUninit::uninit(); view.decoded_len() + DECODE_PADDING];
let w = unsafe { decode_into(view.codes, view.dict, &mut buf) };
let got = unsafe { std::slice::from_raw_parts(buf.as_ptr().cast::<u8>(), w) };
assert_eq!(got, bytes.as_slice());
}
#[test]
fn roundtrip_across_bit_widths() {
let mut bytes = Vec::new();
let mut offsets = vec![0u32];
for i in 0..5000u32 {
let row = format!("row-{i:04}-https://example.com/path/{}", i % 37);
bytes.extend_from_slice(row.as_bytes());
offsets.push(bytes.len() as u32);
}
for bits in 9..=16u8 {
let cfg = Config {
max_dict_bits: MaxDictBits::new(bits).unwrap(),
..DEFAULT_CONFIG
};
let col = compress(&bytes, &offsets, cfg).unwrap();
assert_eq!(decode_all(col.view()), bytes, "bits={bits}");
}
}
#[test]
fn code_bits_is_within_capacity() {
let (bytes, offsets) = pack(&[b"hello world", b"hello there", b"world peace"]);
let cfg = Config {
max_dict_bits: MaxDictBits::new(12).unwrap(),
..DEFAULT_CONFIG
};
let col = compress(&bytes, &offsets, cfg).unwrap();
assert!(col.dict.code_bits() <= 12);
}
#[test]
#[should_panic(expected = "row offsets must be non-decreasing")]
fn row_codes_panics_typed_on_bad_row_offsets() {
let (bytes, offsets) = pack(&[b"alpha", b"beta"]);
let col = compress(&bytes, &offsets, DEFAULT_CONFIG).unwrap();
let view = col.view();
let ro = vec![0u32, (view.codes.len() + 1) as u32];
let bad = ColumnView {
dict: view.dict,
codes: view.codes,
row_offsets: &ro,
};
let _ = bad.row_codes(0);
}
#[test]
fn search_selects_matching_rows() {
let rows: &[&[u8]] = &[b"apple", b"banana", b"apricot", b"cherry", b"apple"];
let (bytes, offsets) = pack(rows);
let col = compress(&bytes, &offsets, DEFAULT_CONFIG).unwrap();
let view = col.view();
assert_eq!(view.rows_equal_to(b"apple"), vec![0, 4]);
assert_eq!(view.rows_starting_with(b"ap"), vec![0, 2, 4]);
assert_eq!(view.rows_containing(b"an"), vec![1]);
assert_eq!(view.rows_equal_to(b"grape"), Vec::<usize>::new());
}
#[test]
fn search_empty_needle_semantics() {
let rows: &[&[u8]] = &[b"a", b"", b"abc", b""];
let (bytes, offsets) = pack(rows);
let col = compress(&bytes, &offsets, DEFAULT_CONFIG).unwrap();
let view = col.view();
assert_eq!(view.rows_equal_to(b""), vec![1, 3]);
assert_eq!(view.rows_starting_with(b""), vec![0, 1, 2, 3]);
assert_eq!(view.rows_containing(b""), vec![0, 1, 2, 3]);
}
#[test]
fn search_agrees_with_decode_oracle() {
use crate::test_corpus::user_strings;
let corpus: Vec<Vec<u8>> = user_strings(60)
.into_iter()
.map(String::into_bytes)
.collect();
let rows: Vec<&[u8]> = corpus.iter().map(Vec::as_slice).collect();
let (bytes, offsets) = pack(&rows);
let col = compress(&bytes, &offsets, DEFAULT_CONFIG).unwrap();
let view = col.view();
let needles: &[&[u8]] = &[
b"",
b"h",
b"https",
b"https://www.example.com/",
b"example",
b".com",
b"://",
b"zzz",
];
for &needle in needles {
let eq: Vec<usize> = (0..view.num_rows())
.filter(|&k| decode_row(view, k).as_slice() == needle)
.collect();
assert_eq!(view.rows_equal_to(needle), eq, "equals {needle:?}");
let pre: Vec<usize> = (0..view.num_rows())
.filter(|&k| decode_row(view, k).starts_with(needle))
.collect();
assert_eq!(
view.rows_starting_with(needle),
pre,
"starts_with {needle:?}"
);
let con: Vec<usize> = (0..view.num_rows())
.filter(|&k| {
let r = decode_row(view, k);
needle.is_empty() || r.windows(needle.len()).any(|w| w == needle)
})
.collect();
assert_eq!(view.rows_containing(needle), con, "contains {needle:?}");
}
}
}