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,
}
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;
}
Some(Self { next, done: total as u8 })
}
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
}
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());
}
}