use crate::MAX_SPARSE_GRAM_SIZE;
use crate::ngram::NGram;
use crate::table::{bigram_h, bigram_priority_rolling};
#[inline]
pub const fn max_sparse_grams(content_len: usize) -> usize {
if content_len < 2 {
0
} else {
(content_len - 1) * 3
}
}
#[inline]
fn window_to_gram(window: u64, len: usize) -> NGram {
debug_assert!(len <= MAX_SPARSE_GRAM_SIZE);
NGram::from_window(window << ((MAX_SPARSE_GRAM_SIZE - len) * 8), len)
}
pub fn collect_sparse_grams(content: &[u8]) -> Vec<NGram> {
let mut out = Vec::with_capacity(max_sparse_grams(content.len()));
let spare = out.spare_capacity_mut();
let mut w = 0;
collect_sparse_grams_deque(content, |gram, _idx| {
spare[w].write(gram);
w += 1;
});
unsafe { out.set_len(w) };
out
}
pub fn collect_sparse_grams_deque(content: &[u8], mut emit: impl FnMut(NGram, u32)) {
let n = content.len();
if n < 2 {
return;
}
const MASK: usize = MAX_SPARSE_GRAM_SIZE - 1;
const EMPTY: u32 = 0u32.wrapping_sub(MAX_SPARSE_GRAM_SIZE as u32);
let mut idx_buf = [EMPTY; MAX_SPARSE_GRAM_SIZE];
let mut val_buf = [0u32; MAX_SPARSE_GRAM_SIZE];
let mut tail = 0usize;
let mut window = content[0] as u64;
let mut h = bigram_h(content[0]);
for idx in 1..n as u32 {
window = (window << 8) | content[idx as usize] as u64;
let (value, h_b) = bigram_priority_rolling((window >> 8) as u8, window as u8, h);
h = h_b;
emit(window_to_gram(window, 2), idx + 1);
let mut t = tail;
loop {
let slot = t.wrapping_sub(1) & MASK;
let begin = idx_buf[slot];
if idx.wrapping_sub(begin) + 1 >= MAX_SPARSE_GRAM_SIZE as u32 {
break;
}
emit(
window_to_gram(window, (idx.wrapping_sub(begin) + 2) as usize),
idx + 1,
);
let bval = val_buf[slot];
if bval < value {
break;
}
t -= 1;
if bval == value {
break;
}
}
idx_buf[t & MASK] = idx;
val_buf[t & MASK] = value;
tail = t + 1;
}
}
pub fn collect_sparse_grams_scan(content: &[u8], mut emit: impl FnMut(NGram, u32)) {
let n = content.len();
if n < 2 {
return;
}
const MASK: usize = MAX_SPARSE_GRAM_SIZE - 1;
let mut window = content[0] as u64;
let mut h = bigram_h(content[0]);
let mut priorities = [0u32; MAX_SPARSE_GRAM_SIZE];
for idx in 1..n as u32 {
window = (window << 8) | content[idx as usize] as u64;
let (v1, h_b) = bigram_priority_rolling((window >> 8) as u8, window as u8, h);
h = h_b;
priorities[idx as usize & MASK] = v1;
emit(window_to_gram(window, 2), idx + 1);
let mut running_min = u32::MAX;
for d in 1..=(MAX_SPARSE_GRAM_SIZE as u32 - 2) {
if d >= idx {
break;
}
let v_p = priorities[(idx - d) as usize & MASK];
if v_p < running_min {
if running_min <= v1 {
break;
}
let len = d as usize + 2;
emit(window_to_gram(window, len), idx + 1);
running_min = v_p;
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::table::bigram_priority;
use std::collections::HashSet;
fn collect_to_vec(run: impl FnOnce(&mut dyn FnMut(NGram, u32))) -> Vec<NGram> {
let mut out = Vec::new();
run(&mut |gram, _idx| out.push(gram));
out
}
fn brute_force_sparse_grams(content: &[u8]) -> HashSet<NGram> {
let n = content.len();
let mut result = HashSet::new();
if n < 2 {
return result;
}
for i in 0..n - 1 {
result.insert(NGram::from_bytes(&content[i..i + 2]));
}
for len in 3..=MAX_SPARSE_GRAM_SIZE {
for start in 0..=n.saturating_sub(len) {
if start + len > n {
break;
}
let left = bigram_priority(content[start], content[start + 1]);
let right = bigram_priority(content[start + len - 2], content[start + len - 1]);
let mut min_interior = u32::MAX;
for k in 1..len - 2 {
let p = bigram_priority(content[start + k], content[start + k + 1]);
min_interior = min_interior.min(p);
}
if left < min_interior && right < min_interior {
result.insert(NGram::from_bytes(&content[start..start + len]));
}
}
}
result
}
#[test]
fn test_empty_input() {
assert!(collect_sparse_grams(b"").is_empty());
}
#[test]
fn test_single_byte() {
assert!(collect_sparse_grams(b"a").is_empty());
}
#[test]
fn test_two_bytes() {
let grams = collect_sparse_grams(b"ab");
assert_eq!(grams.len(), 1);
assert_eq!(grams[0], NGram::from_bytes(b"ab"));
}
#[test]
fn test_three_bytes() {
let grams = collect_sparse_grams(b"abc");
assert!(grams.len() >= 2);
assert_eq!(grams[0], NGram::from_bytes(b"ab"));
assert_eq!(grams[1], NGram::from_bytes(b"bc"));
}
#[test]
fn test_gram_lengths_bounded() {
let input = b"self.reset_states(the_quick_brown_fox_jumps";
let grams = collect_sparse_grams(input);
for gram in &grams {
assert!(gram.len() >= 2, "gram too short: {gram:?}");
assert!(
gram.len() <= MAX_SPARSE_GRAM_SIZE,
"gram too long: {gram:?}"
);
}
}
#[test]
fn test_produces_longer_grams() {
let grams = collect_sparse_grams(b"self.reset_states(");
assert!(grams.iter().any(|g| g.len() > 2));
}
#[test]
fn test_max_gram_size_boundary() {
let grams = collect_sparse_grams(b"abcdefgh");
for gram in &grams {
assert!(gram.len() <= MAX_SPARSE_GRAM_SIZE);
}
}
#[test]
fn test_repeated_bytes() {
let grams = collect_sparse_grams(b"aaaaaaaaaa");
assert!(grams.iter().filter(|g| g.len() == 2).count() >= 9);
}
#[test]
fn test_gram_count_scales_linearly() {
let input: Vec<u8> = (0..1000).map(|i| (i % 256) as u8).collect();
let grams = collect_sparse_grams(&input);
assert!(grams.len() >= input.len() - 1);
assert!(grams.len() <= input.len() * 3);
}
#[test]
fn test_scan_equivalence_small() {
for input in [b"" as &[u8], b"x", b"ab", b"abc", b"abcdefgh", b"abcdefghi"] {
assert_eq!(
collect_to_vec(|emit| collect_sparse_grams_deque(input, emit)),
collect_to_vec(|emit| collect_sparse_grams_scan(input, emit)),
"mismatch on {:?}",
std::str::from_utf8(input).unwrap_or("?")
);
}
}
#[test]
fn test_scan_equivalence_hello_world() {
let input = b"hello world";
assert_eq!(
collect_to_vec(|emit| collect_sparse_grams_deque(input, emit)),
collect_to_vec(|emit| collect_sparse_grams_scan(input, emit)),
);
}
#[test]
fn test_scan_equivalence_large() {
let input: Vec<u8> = (0..1000).map(|i| (i % 256) as u8).collect();
assert_eq!(
collect_to_vec(|emit| collect_sparse_grams_deque(&input, emit)),
collect_to_vec(|emit| collect_sparse_grams_scan(&input, emit)),
);
}
#[test]
fn test_scan_equivalence_source_code() {
let input = include_bytes!("extract.rs");
assert_eq!(
collect_to_vec(|emit| collect_sparse_grams_deque(input, emit)),
collect_to_vec(|emit| collect_sparse_grams_scan(input, emit)),
);
}
fn assert_matches_brute_force(input: &[u8]) {
let grams = collect_sparse_grams(input);
let actual: HashSet<NGram> = grams.into_iter().collect();
let expected = brute_force_sparse_grams(input);
let only_actual: Vec<_> = actual.difference(&expected).collect();
let only_expected: Vec<_> = expected.difference(&actual).collect();
if !only_actual.is_empty() || !only_expected.is_empty() {
panic!(
"mismatch on input len={}\n only in algorithm: {:?}\n only in brute force: {:?}",
input.len(),
only_actual,
only_expected
);
}
}
#[test]
fn test_brute_force_small() {
for input in [
b"" as &[u8],
b"x",
b"ab",
b"abc",
b"abcd",
b"abcdefgh",
b"abcdefghi",
] {
assert_matches_brute_force(input);
}
}
#[test]
fn test_brute_force_hello_world() {
assert_matches_brute_force(b"hello world");
}
#[test]
fn test_brute_force_repeated() {
assert_matches_brute_force(b"aaaaaaaaaa");
}
#[test]
fn test_brute_force_code_snippet() {
assert_matches_brute_force(b"self.reset_states(the_quick_brown_fox_jumps");
}
#[test]
fn test_brute_force_tie_break() {
assert_matches_brute_force(b"ababababab");
assert_matches_brute_force(b"the the the the");
assert_matches_brute_force(b"a.b.a.b.a.b.");
}
#[test]
fn test_brute_force_diverse() {
let input: Vec<u8> = (0..200).map(|i| (i % 256) as u8).collect();
assert_matches_brute_force(&input);
}
#[test]
fn test_brute_force_long_ascending() {
let input: Vec<u8> = (0..3000u32).map(|i| 33 + (i % 90) as u8).collect();
assert_matches_brute_force(&input);
assert_eq!(
collect_to_vec(|emit| collect_sparse_grams_deque(&input, emit)),
collect_to_vec(|emit| collect_sparse_grams_scan(&input, emit)),
);
}
#[test]
fn test_brute_force_source_code() {
let input = include_bytes!("extract.rs");
assert_matches_brute_force(input);
}
}