use std::collections::VecDeque;
pub(crate) const NO_TOKEN: u32 = u32::MAX;
pub(crate) const ROOT: u32 = 0;
const FREE: u32 = u32::MAX;
const ROOT_OWNED: u32 = u32::MAX - 1;
pub(crate) struct ByteTrie {
base: Vec<u32>,
check: Vec<u32>,
values: Vec<u32>,
}
impl ByteTrie {
pub(crate) fn build<'a>(entries: impl Iterator<Item = (&'a str, u32)>) -> Self {
let mut children: Vec<Vec<(u8, u32)>> = vec![Vec::new()];
let mut node_value = vec![NO_TOKEN];
for (surface, id) in entries {
if surface.is_empty() {
continue;
}
let mut node = 0usize;
for &byte in surface.as_bytes() {
node = match children[node].iter().find(|(label, _)| *label == byte) {
Some(&(_, next)) => next as usize,
None => {
children.push(Vec::new());
node_value.push(NO_TOKEN);
let next = children.len() - 1;
children[node].push((byte, next as u32));
next
}
};
}
node_value[node] = id;
}
let mut trie = Self {
base: vec![0; 256],
check: vec![FREE; 256],
values: vec![NO_TOKEN; 256],
};
trie.check[ROOT as usize] = ROOT_OWNED;
trie.values[ROOT as usize] = node_value[0];
let mut slot_of = vec![0u32; children.len()];
let mut queue = VecDeque::from([0usize]);
let mut cursor = 1usize;
while let Some(node) = queue.pop_front() {
let edges = std::mem::take(&mut children[node]);
if edges.is_empty() {
continue;
}
let base = trie.find_base(&edges, &mut cursor);
let state = slot_of[node];
trie.base[state as usize] = base;
for &(label, child) in &edges {
let slot = (base ^ label as u32) as usize;
trie.reserve(slot + 1);
trie.check[slot] = state;
trie.values[slot] = node_value[child as usize];
slot_of[child as usize] = slot as u32;
queue.push_back(child as usize);
}
}
trie
}
fn reserve(&mut self, len: usize) {
if len <= self.check.len() {
return;
}
let len = len.next_power_of_two();
self.base.resize(len, 0);
self.check.resize(len, FREE);
self.values.resize(len, NO_TOKEN);
}
#[inline]
fn is_free(&self, slot: usize) -> bool {
slot != ROOT as usize && self.check.get(slot).is_none_or(|&owner| owner == FREE)
}
fn find_base(&self, edges: &[(u8, u32)], cursor: &mut usize) -> u32 {
while !self.is_free(*cursor) {
*cursor += 1;
}
let first = edges[0].0 as u32;
let mut candidate = *cursor;
loop {
let base = candidate as u32 ^ first;
if edges
.iter()
.all(|&(label, _)| self.is_free((base ^ label as u32) as usize))
{
return base;
}
candidate += 1;
while !self.is_free(candidate) {
candidate += 1;
}
}
}
#[inline]
pub(crate) fn step(&self, state: u32, byte: u8) -> Option<u32> {
let slot = (self.base[state as usize] ^ byte as u32) as usize;
(self.check.get(slot) == Some(&state)).then_some(slot as u32)
}
#[inline]
pub(crate) fn value(&self, state: u32) -> u32 {
self.values[state as usize]
}
#[inline]
pub(crate) fn longest_prefix(&self, bytes: &[u8]) -> Option<(usize, u32)> {
let mut state = ROOT;
let mut best = None;
for (i, &byte) in bytes.iter().enumerate() {
let Some(next) = self.step(state, byte) else {
break;
};
state = next;
let value = self.value(state);
if value != NO_TOKEN {
best = Some((i + 1, value));
}
}
best
}
}
#[cfg(test)]
mod tests {
use super::*;
fn trie(entries: &[(&str, u32)]) -> ByteTrie {
ByteTrie::build(entries.iter().copied())
}
#[test]
fn longest_prefix_finds_the_deepest_token_it_passes() {
let t = trie(&[("a", 1), ("ab", 2), ("abcd", 3)]);
assert_eq!(t.longest_prefix(b"a"), Some((1, 1)));
assert_eq!(t.longest_prefix(b"abc"), Some((2, 2)));
assert_eq!(t.longest_prefix(b"abcde"), Some((4, 3)));
assert_eq!(t.longest_prefix(b"z"), None);
assert_eq!(t.longest_prefix(b""), None);
}
#[test]
fn a_missing_edge_is_not_answered_by_another_states_slot() {
let leads: Vec<u8> = (b'a'..=b'z')
.chain(b'A'..=b'Z')
.chain(b'0'..=b'9')
.collect();
let entries: Vec<(String, u32)> = leads
.iter()
.map(|&b| (format!("{}x", b as char), b as u32))
.collect();
let t = ByteTrie::build(entries.iter().map(|(s, id)| (s.as_str(), *id)));
for &b in &leads {
assert_eq!(t.longest_prefix(&[b, b'x']), Some((2, b as u32)), "{b}");
assert_eq!(t.longest_prefix(&[b, b'y']), None, "{b} followed by y");
}
}
#[test]
fn stepping_reports_every_token_prefix_in_order() {
let t = trie(&[("ab", 7), ("abc", 8), ("abcd", 9)]);
let mut state = ROOT;
let mut seen = Vec::new();
for (i, &byte) in b"abcd".iter().enumerate() {
state = t.step(state, byte).expect("edge exists");
if t.value(state) != NO_TOKEN {
seen.push((i + 1, t.value(state)));
}
}
assert_eq!(seen, vec![(2, 7), (3, 8), (4, 9)]);
}
#[test]
fn a_repeated_surface_keeps_the_later_id() {
let t = trie(&[("a", 1), ("a", 5)]);
assert_eq!(t.longest_prefix(b"a"), Some((1, 5)));
}
}