use rudb_common::{Error, Result};
use crate::fsst::{ESCAPE, SymbolTable};
const UNKNOWN: u8 = u8::MAX;
#[derive(Debug, Clone)]
pub struct Sequence {
next: Vec<u8>,
done: u8,
needs: u64,
}
#[inline]
fn gram(a: u8, b: u8, c: u8) -> u64 {
let run = u32::from(a) << 16 | u32::from(b) << 8 | u32::from(c);
1 << (run.wrapping_mul(0x9E37_79B1) >> 26)
}
#[must_use]
pub fn grams(text: &[u8]) -> u64 {
text.windows(3).fold(0, |bits, run| bits | gram(run[0], run[1], run[2]))
}
impl Sequence {
#[must_use]
pub fn new(pieces: &[&[u8]]) -> Option<Self> {
let pieces: Vec<&[u8]> = pieces.iter().copied().filter(|piece| !piece.is_empty()).collect();
let total: usize = pieces.iter().map(|piece| piece.len()).sum();
if pieces.is_empty() || total >= usize::from(UNKNOWN) {
return None;
}
let states = total + 1;
let mut next = vec![0_u8; states * 256];
let mut base = 0;
for piece in &pieces {
let first = usize::from(piece[0]);
for byte in 0..256 {
next[base * 256 + byte] = base as u8;
}
next[base * 256 + first] = (base + 1) as u8;
let mut restart = base;
for (offset, &byte) in piece.iter().enumerate().skip(1) {
let state = base + offset;
let (before, row) = next.split_at_mut(state * 256);
row[..256].copy_from_slice(&before[restart * 256..restart * 256 + 256]);
row[usize::from(byte)] = (state + 1) as u8;
restart = usize::from(next[restart * 256 + usize::from(byte)]);
}
base += piece.len();
}
for byte in 0..256 {
next[total * 256 + byte] = total as u8;
}
let needs = pieces.iter().fold(0, |bits, piece| bits | grams(piece));
Some(Self { next, done: total as u8, needs })
}
#[must_use]
pub fn needs(&self) -> u64 {
self.needs
}
fn states(&self) -> usize {
usize::from(self.done) + 1
}
#[must_use]
pub fn holds(&self, text: &[u8]) -> bool {
let mut state = 0_u8;
for &byte in text {
state = self.next[usize::from(state) * 256 + usize::from(byte)];
if state == self.done {
return true;
}
}
false
}
#[must_use]
pub fn over<'a>(&'a self, table: &'a SymbolTable) -> Coded<'a> {
Coded { sequence: self, table, steps: vec![UNKNOWN; self.states() * 256] }
}
}
#[derive(Debug)]
pub struct Coded<'a> {
sequence: &'a Sequence,
table: &'a SymbolTable,
steps: Vec<u8>,
}
impl Coded<'_> {
pub fn holds(&mut self, codes: &[u8]) -> Result<bool> {
let done = self.sequence.done;
let mut state = 0_u8;
let mut at = 0;
while at < codes.len() {
let code = codes[at];
at += 1;
let slot = usize::from(state) * 256 + usize::from(code);
state = if code == ESCAPE {
let Some(&byte) = codes.get(at) else {
return Err(Error::internal("a compressed string ends in an escape"));
};
at += 1;
self.sequence.next[usize::from(state) * 256 + usize::from(byte)]
} else if self.steps[slot] != UNKNOWN {
self.steps[slot]
} else {
self.learn(state, code)?
};
if state == done {
return Ok(true);
}
}
Ok(false)
}
#[cold]
fn learn(&mut self, from: u8, code: u8) -> Result<u8> {
let Some((bytes, len)) = self.table.symbol(code) else {
return Err(Error::internal(format!("code {code} is not in the table")));
};
let mut state = from;
for &byte in &bytes[..len] {
state = self.sequence.next[usize::from(state) * 256 + usize::from(byte)];
}
self.steps[usize::from(from) * 256 + usize::from(code)] = state;
Ok(state)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn searched(text: &[u8], pieces: &[&[u8]]) -> bool {
let mut rest = text;
for piece in pieces {
if piece.is_empty() {
continue;
}
match rest.windows(piece.len()).position(|window| window == *piece) {
Some(at) => rest = &rest[at + piece.len()..],
None => return false,
}
}
true
}
#[test]
fn a_string_that_holds_the_pieces_has_every_bit_they_need() {
let cases: [&[&[u8]]; 4] =
[&[b"special", b"requests"], &[b"furiously"], &[b"ab"], &[b"aab", b"sts", b"\xc3\xa9"]];
for pieces in cases {
let sequence = Sequence::new(pieces).expect("an automaton");
let mut kept = 0;
for text in texts() {
let has = grams(&text) & sequence.needs() == sequence.needs();
if searched(&text, pieces) {
assert!(has, "{:?} holds {pieces:?} and its sketch says not", text);
kept += 1;
}
}
assert!(kept > 0, "{pieces:?} is held somewhere, so the test tests something");
}
assert_eq!(Sequence::new(&[b"ab"]).expect("an automaton").needs(), 0);
assert_eq!(grams(b"ab"), 0, "no run of three");
assert_ne!(grams(b"special") & grams(b"requests"), grams(b"special"));
}
fn texts() -> Vec<Vec<u8>> {
let words = [
"special",
"requests",
"spec",
"specia",
"ial",
"requ",
"sts",
"the",
"furiously",
"aaa",
"aab",
"ab",
"é",
"ü",
"",
"s",
];
let mut texts = Vec::new();
let mut seed = 7_u64;
for _ in 0..3000 {
let mut text = Vec::new();
seed = seed.wrapping_mul(6_364_136_223_846_793_005).wrapping_add(1);
for step in 0..(seed >> 60) {
let pick = (seed >> (step * 4 % 56)) as usize % words.len();
text.extend_from_slice(words[pick].as_bytes());
if (seed >> (step % 60)) & 1 == 1 {
text.push(b' ');
}
}
texts.push(text);
}
texts
}
#[test]
fn the_walk_over_bytes_and_the_walk_over_codes_agree_with_a_search() {
let texts = texts();
let samples: Vec<&[u8]> = texts.iter().map(Vec::as_slice).collect();
let table = SymbolTable::train(&samples);
let patterns: [&[&[u8]]; 7] = [
&[b"special", b"requests"],
&[b"aab"],
&[b"ab", b"ab", b"ab"],
&[b"s", b"s"],
&["é".as_bytes(), b"ial"],
&[b"", b"spec", b""],
&[b"furiously the", b"sts"],
];
let mut found = 0;
for pieces in patterns {
let sequence = Sequence::new(pieces).expect("something to find");
let mut coded = sequence.over(&table);
for text in &texts {
let wanted = searched(text, pieces);
found += usize::from(wanted);
assert_eq!(sequence.holds(text), wanted, "{pieces:?} in {text:?}");
let mut codes = Vec::new();
table.compress(text, &mut codes);
assert_eq!(coded.holds(&codes).expect("codes"), wanted, "{pieces:?} in {text:?}");
}
}
assert!(found > 1000, "only {found} matches");
}
#[test]
fn nothing_to_find_or_too_much_is_no_automaton() {
assert!(Sequence::new(&[]).is_none());
assert!(Sequence::new(&[b"", b""]).is_none());
assert!(Sequence::new(&[&[b'x'; 255]]).is_none());
assert!(Sequence::new(&[&[b'x'; 254]]).is_some());
}
}