use std::hash::Hasher;
use rustc_hash::{FxHashMap, FxHasher};
#[derive(Clone, Copy)]
struct Entry {
hash: u64,
offset: u32,
len: u32,
id: u32,
}
const VACANT: Entry = Entry {
hash: 0,
offset: 0,
len: u32::MAX,
id: 0,
};
#[derive(Clone)]
pub struct Encoder {
arena: Vec<u8>,
slots: Box<[Entry]>,
len: usize,
hashes_unique: bool,
}
impl Default for Encoder {
fn default() -> Self {
Self::with_capacity(0)
}
}
const MIN_SLOTS: usize = 16;
impl Encoder {
pub fn with_capacity(tokens: usize) -> Self {
let slots = (tokens * 2).next_power_of_two().max(MIN_SLOTS);
Self {
arena: Vec::new(),
slots: vec![VACANT; slots].into_boxed_slice(),
len: 0,
hashes_unique: true,
}
}
pub fn with_arena(arena: Vec<u8>, tokens: usize) -> Self {
let slots = (tokens * 2).next_power_of_two().max(MIN_SLOTS);
Self {
arena,
slots: vec![VACANT; slots].into_boxed_slice(),
len: 0,
hashes_unique: true,
}
}
pub fn insert_span(&mut self, offset: u32, len: u32, id: u32) {
assert!(
offset as usize + len as usize <= self.arena.len(),
"span runs past the arena"
);
if self.len * 2 >= self.slots.len() {
self.resize(self.slots.len() * 2);
}
let mask = self.mask();
let key_start = offset as usize;
let key_end = key_start + len as usize;
let hash = Self::hash(&self.arena[key_start..key_end]);
let mut slot = Self::slot_of(hash, mask);
loop {
let entry = self.slots[slot];
if entry.len == u32::MAX {
break;
}
if entry.len == len
&& Self::key_at(&self.arena, &entry) == &self.arena[key_start..key_end]
{
return;
}
self.hashes_unique &= entry.hash != hash;
slot = (slot + 1) & mask;
}
self.slots[slot] = Entry {
hash,
offset,
len,
id,
};
self.len += 1;
}
pub fn reserve(&mut self, additional: usize) {
let wanted = (self.len + additional) * 2;
if wanted > self.slots.len() {
self.resize(wanted.next_power_of_two().max(MIN_SLOTS));
}
self.arena.reserve(additional * 8);
}
#[inline]
pub fn hash_of(key: &[u8]) -> u64 {
Self::hash(key)
}
#[inline]
fn hash(key: &[u8]) -> u64 {
let mut hasher = FxHasher::default();
hasher.write(key);
hasher.finish()
}
#[inline]
fn mask(&self) -> usize {
self.slots.len() - 1
}
#[inline]
fn slot_of(hash: u64, mask: usize) -> usize {
(hash.wrapping_mul(0x9E37_79B9_7F4A_7C15) >> 32) as usize & mask
}
#[inline]
fn same_bytes(a: &[u8], b: &[u8]) -> bool {
match a.len() {
0 => true,
1 => a[..1] == b[..1],
2 => a[..2] == b[..2],
3 => a[..3] == b[..3],
4 => a[..4] == b[..4],
5 => a[..5] == b[..5],
6 => a[..6] == b[..6],
7 => a[..7] == b[..7],
8 => a[..8] == b[..8],
9 => a[..9] == b[..9],
10 => a[..10] == b[..10],
11 => a[..11] == b[..11],
12 => a[..12] == b[..12],
13 => a[..13] == b[..13],
14 => a[..14] == b[..14],
15 => a[..15] == b[..15],
16 => a[..16] == b[..16],
_ => a == b,
}
}
#[inline]
fn key_at<'a>(arena: &'a [u8], entry: &Entry) -> &'a [u8] {
let start = entry.offset as usize;
&arena[start..start + entry.len as usize]
}
#[inline]
pub fn get(&self, key: &[u8]) -> Option<u32> {
self.get_with_hash(key, Self::hash_of(key))
}
#[inline]
pub fn get_with_hash(&self, key: &[u8], hash: u64) -> Option<u32> {
let mask = self.mask();
let mut slot = Self::slot_of(hash, mask);
if self.hashes_unique {
loop {
let entry = self.slots[slot];
if entry.len == u32::MAX {
return None;
}
if entry.hash == hash {
return Some(entry.id);
}
slot = (slot + 1) & mask;
}
}
loop {
let entry = self.slots[slot];
if entry.len == u32::MAX {
return None;
}
if entry.len as usize == key.len()
&& Self::same_bytes(Self::key_at(&self.arena, &entry), key)
{
return Some(entry.id);
}
slot = (slot + 1) & mask;
}
}
#[inline]
pub fn contains_key(&self, key: &[u8]) -> bool {
self.get(key).is_some()
}
pub fn insert(&mut self, key: &[u8], id: u32) -> Option<u32> {
if self.len * 2 >= self.slots.len() {
self.resize(self.slots.len() * 2);
}
let mask = self.mask();
let hash = Self::hash(key);
let mut slot = Self::slot_of(hash, mask);
loop {
let entry = self.slots[slot];
if entry.len == u32::MAX {
break;
}
if entry.len as usize == key.len() && Self::key_at(&self.arena, &entry) == key {
let previous = entry.id;
self.slots[slot].id = id;
return Some(previous);
}
slot = (slot + 1) & mask;
}
let offset = self.arena.len() as u32;
self.arena.extend_from_slice(key);
self.slots[slot] = Entry {
hash,
offset,
len: key.len() as u32,
id,
};
self.len += 1;
None
}
pub fn insert_if_absent(&mut self, key: &[u8], id: u32) -> u32 {
match self.get(key) {
Some(existing) => existing,
None => {
self.insert(key, id);
id
}
}
}
fn resize(&mut self, slots: usize) {
let mask = slots - 1;
let mut fresh = vec![VACANT; slots];
for entry in self.slots.iter().filter(|e| e.len != u32::MAX) {
let mut slot = Self::slot_of(entry.hash, mask);
while fresh[slot].len != u32::MAX {
slot = (slot + 1) & mask;
}
fresh[slot] = *entry;
}
self.slots = fresh.into_boxed_slice();
}
pub fn len(&self) -> usize {
self.len
}
pub fn is_empty(&self) -> bool {
self.len == 0
}
pub fn iter(&self) -> impl Iterator<Item = (&[u8], u32)> + '_ {
self.slots
.iter()
.filter(|entry| entry.len != u32::MAX)
.map(move |entry| (Self::key_at(&self.arena, entry), entry.id))
}
pub fn keys(&self) -> impl Iterator<Item = &[u8]> + '_ {
self.iter().map(|(key, _)| key)
}
pub fn values(&self) -> impl Iterator<Item = u32> + '_ {
self.iter().map(|(_, id)| id)
}
}
impl<'a> IntoIterator for &'a Encoder {
type Item = (&'a [u8], u32);
type IntoIter = Box<dyn Iterator<Item = (&'a [u8], u32)> + 'a>;
fn into_iter(self) -> Self::IntoIter {
Box::new(self.iter())
}
}
impl FromIterator<(Vec<u8>, u32)> for Encoder {
fn from_iter<I: IntoIterator<Item = (Vec<u8>, u32)>>(entries: I) -> Self {
let entries = entries.into_iter();
let mut encoder = Self::with_capacity(entries.size_hint().0);
for (key, id) in entries {
encoder.insert(&key, id);
}
encoder
}
}
impl<'a> FromIterator<(&'a [u8], u32)> for Encoder {
fn from_iter<I: IntoIterator<Item = (&'a [u8], u32)>>(entries: I) -> Self {
let entries = entries.into_iter();
let mut encoder = Self::with_capacity(entries.size_hint().0);
for (key, id) in entries {
encoder.insert(key, id);
}
encoder
}
}
impl std::fmt::Debug for Encoder {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Encoder")
.field("len", &self.len)
.field("bytes", &self.arena.len())
.finish()
}
}
pub fn encoder_from_owned(map: FxHashMap<Vec<u8>, u32>) -> Encoder {
let mut encoder = Encoder::with_capacity(map.len());
for (key, id) in map {
encoder.insert(&key, id);
}
encoder
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn resolves_what_it_was_given() {
let mut encoder = Encoder::default();
encoder.insert(b"hello", 1);
encoder.insert(b"world", 2);
assert_eq!(encoder.get(b"hello"), Some(1));
assert_eq!(encoder.get(b"world"), Some(2));
assert_eq!(encoder.get(b"absent"), None);
assert_eq!(encoder.len(), 2);
}
#[test]
fn the_empty_token_is_a_usable_key() {
let mut encoder = Encoder::default();
encoder.insert(b"", 50256);
assert_eq!(encoder.get(b""), Some(50256));
}
#[test]
fn the_first_key_inserted_is_found() {
let mut encoder = Encoder::default();
encoder.insert(b"first", 7);
for i in 0..500u32 {
encoder.insert(format!("filler{i}").as_bytes(), i + 100);
}
assert_eq!(encoder.get(b"first"), Some(7));
}
#[test]
fn growing_preserves_every_entry() {
let mut encoder = Encoder::with_capacity(4);
for i in 0..2000u32 {
encoder.insert(format!("token{i}").as_bytes(), i);
}
assert_eq!(encoder.len(), 2000);
for i in 0..2000u32 {
assert_eq!(encoder.get(format!("token{i}").as_bytes()), Some(i));
}
}
#[test]
fn reinserting_replaces_the_id_without_growing() {
let mut encoder = Encoder::default();
encoder.insert(b"key", 1);
let arena = encoder.arena.len();
assert_eq!(encoder.insert(b"key", 2), Some(1));
assert_eq!(encoder.get(b"key"), Some(2));
assert_eq!(encoder.len(), 1);
assert_eq!(encoder.arena.len(), arena, "arena grew on replacement");
}
#[test]
fn insert_if_absent_keeps_the_first_id() {
let mut encoder = Encoder::default();
assert_eq!(encoder.insert_if_absent(b"key", 1), 1);
assert_eq!(encoder.insert_if_absent(b"key", 9), 1);
assert_eq!(encoder.get(b"key"), Some(1));
}
#[test]
fn iteration_yields_every_entry_once() {
let mut encoder = Encoder::default();
for i in 0..100u32 {
encoder.insert(format!("t{i}").as_bytes(), i);
}
let mut seen: Vec<(Vec<u8>, u32)> = encoder.iter().map(|(k, v)| (k.to_vec(), v)).collect();
seen.sort();
assert_eq!(seen.len(), 100);
assert_eq!(encoder.keys().count(), 100);
assert_eq!(encoder.values().sum::<u32>(), (0..100).sum::<u32>());
}
#[test]
fn distinct_keys_of_equal_length_do_not_alias() {
let mut encoder = Encoder::default();
for i in 0..1000u32 {
encoder.insert(&i.to_le_bytes(), i);
}
for i in 0..1000u32 {
assert_eq!(encoder.get(&i.to_le_bytes()), Some(i));
}
}
}