use super::lookup::prefix_range;
use super::tokenize::tokenize;
use crate::core::dictionary::DictionaryView;
use crate::core::types::{Token, TokenRange};
#[derive(Debug, Clone)]
pub struct PrefixQuery {
tokens: Vec<Token>,
intervals: Vec<TokenRange>,
}
impl PrefixQuery {
pub fn new<V: DictionaryView>(prefix: &[u8], dict: V) -> Self {
let tokens = tokenize(prefix, dict);
let mut intervals = Vec::with_capacity(tokens.len());
let mut byte_off = 0usize;
for &t in &tokens {
intervals.push(prefix_range(dict, &prefix[byte_off..]));
byte_off += dict.token_len(t);
}
Self { tokens, intervals }
}
}
pub fn starts_with(codes: &[Token], query: &PrefixQuery) -> bool {
let needle = &query.tokens;
let min_len = codes.len().min(needle.len());
let diff = std::iter::zip(&codes[..min_len], &needle[..min_len])
.position(|(c, t)| c != t)
.unwrap_or(min_len);
if diff == needle.len() {
return true; }
if diff < codes.len() {
return query.intervals[diff].contains(codes[diff]);
}
false }
#[cfg(test)]
mod tests {
use super::*;
use crate::core::dictionary::Dictionary;
use crate::{Column, DEFAULT_CONFIG, compress};
fn compress_rows(rows: &[&[u8]]) -> Column<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);
}
compress(&bytes, &offsets, DEFAULT_CONFIG).unwrap()
}
fn decode_row(view: crate::ColumnView<'_, u32>, k: usize) -> Vec<u8> {
let mut buf =
vec![std::mem::MaybeUninit::uninit(); view.row_decoded_len(k) + crate::DECODE_PADDING];
let w = unsafe { view.decompress_row_into(k, &mut buf) };
unsafe { std::slice::from_raw_parts(buf.as_ptr().cast::<u8>(), w) }.to_vec()
}
fn check(rows: &[&[u8]], prefixes: &[&[u8]]) {
let col = compress_rows(rows);
let view = col.view();
let wide = view.wide_dict();
for &prefix in prefixes {
let want: Vec<usize> = (0..view.num_rows())
.filter(|&k| decode_row(view, k).starts_with(prefix))
.collect();
for query in [
PrefixQuery::new(prefix, view.dict),
PrefixQuery::new(prefix, wide.as_view()),
] {
let got: Vec<usize> = (0..view.num_rows())
.filter(|&k| starts_with(view.row_codes(k), &query))
.collect();
assert_eq!(got, want, "prefix {prefix:?}");
}
}
}
#[test]
fn empty_prefix_matches_all_rows() {
let rows: &[&[u8]] = &[b"a", b"", b"abc"];
check(rows, &[b""]);
}
#[test]
fn whole_string_and_exact_prefixes() {
let rows: &[&[u8]] = &[b"alpha", b"alpine", b"beta", b"al"];
check(
rows,
&[b"al", b"alp", b"alpha", b"beta", b"b", b"alphas", b"z"],
);
}
#[test]
fn prefix_ending_inside_a_token() {
let rows: &[&[u8]] = &[b"abcdef", b"abcxyz", b"abxxxx", b"abc", b"ab"];
check(rows, &[b"a", b"ab", b"abc", b"abcd", b"abcx", b"abz"]);
}
#[test]
fn matches_brute_force_on_repetitive_corpus() {
use crate::test_corpus::user_strings;
let corpus: Vec<Vec<u8>> = user_strings(50)
.into_iter()
.map(String::into_bytes)
.collect();
let rows: Vec<&[u8]> = corpus.iter().map(Vec::as_slice).collect();
let prefixes: &[&[u8]] = &[
b"",
b"h",
b"https",
b"https://www.",
b"https://www.example.com/",
b"ftp://",
b"zzz",
];
check(&rows, prefixes);
}
#[test]
fn matches_brute_force_on_binary_corpus() {
use crate::test_corpus::binary_strings;
let corpus = binary_strings(40, 24, 19);
let rows: Vec<&[u8]> = corpus.iter().map(Vec::as_slice).collect();
let mut prefixes: Vec<&[u8]> = rows.iter().map(|r| &r[..r.len().min(3)]).collect();
prefixes.push(b"");
prefixes.push(b"\x00\x00");
check(&rows, &prefixes);
}
}