use ahash::RandomState;
use hashbrown::HashTable;
use std::str;
use std::sync::Mutex;
use std::sync::atomic::{AtomicU32, Ordering};
const SHARD_COUNT: usize = 64;
#[derive(Default)]
struct Shard {
table: HashTable<u32>,
bytes: Vec<u8>,
spans: Vec<(usize, u32)>,
globals: Vec<u32>,
}
pub(crate) struct Interner {
shards: Box<[Mutex<Shard>]>,
next: AtomicU32,
hasher: RandomState,
}
impl Interner {
pub fn new() -> Self {
let shards = (0..SHARD_COUNT).map(|_| Mutex::new(Shard::default())).collect();
Self { shards, next: AtomicU32::new(0), hasher: RandomState::new() }
}
pub fn get_or_intern(&self, s: &str) -> u32 {
let hash = self.hasher.hash_one(s.as_bytes());
let shard_idx = (hash as usize) & (SHARD_COUNT - 1);
let mut guard = self.shards[shard_idx].lock().expect("interner shard poisoned");
let Shard { table, bytes, spans, globals } = &mut *guard;
if let Some(&local) = table.find(hash, |&l| {
let (off, len) = spans[l as usize];
&bytes[off..off + len as usize] == s.as_bytes()
}) {
return globals[local as usize];
}
let off = bytes.len();
bytes.extend_from_slice(s.as_bytes());
let local = u32::try_from(spans.len()).expect("more than u32::MAX terms in one shard");
spans.push((off, s.len() as u32));
let global = self.next.fetch_add(1, Ordering::Relaxed);
globals.push(global);
let hasher = &self.hasher;
table.insert_unique(hash, local, |&l| {
let (o, len) = spans[l as usize];
hasher.hash_one(&bytes[o..o + len as usize])
});
global
}
pub fn len(&self) -> usize {
self.next.load(Ordering::Relaxed) as usize
}
pub fn into_terms(self) -> Terms {
let n = self.len();
let shards: Vec<Shard> =
self.shards.into_vec().into_iter().map(|m| m.into_inner().expect("interner shard poisoned")).collect();
let mut locations = vec![(0u32, 0usize, 0u32); n];
for (sid, shard) in shards.iter().enumerate() {
for (local, &global) in shard.globals.iter().enumerate() {
let (off, len) = shard.spans[local];
locations[global as usize] = (sid as u32, off, len);
}
}
let bytes: Vec<Vec<u8>> = shards.into_iter().map(|s| s.bytes).collect();
Terms { bytes, locations }
}
}
impl std::fmt::Debug for Interner {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Interner")
.field("len", &self.len())
.field("shards", &self.shards.len())
.finish_non_exhaustive()
}
}
pub(crate) struct Terms {
bytes: Vec<Vec<u8>>,
locations: Vec<(u32, usize, u32)>,
}
impl Terms {
pub const fn len(&self) -> usize {
self.locations.len()
}
fn raw(&self, idx: u32) -> &[u8] {
let (sid, off, len) = self.locations[idx as usize];
&self.bytes[sid as usize][off..off + len as usize]
}
pub fn cmp(&self, a: u32, b: u32) -> std::cmp::Ordering {
self.raw(a).cmp(self.raw(b))
}
pub fn get(&self, idx: u32) -> &str {
str::from_utf8(self.raw(idx)).expect("interned bytes are valid UTF-8")
}
}
#[cfg(test)]
mod tests {
use super::Interner;
use rayon::prelude::*;
#[test]
fn round_trip_indices() {
let interner = Interner::new();
let a = interner.get_or_intern("alpha");
let b = interner.get_or_intern("beta");
let a2 = interner.get_or_intern("alpha");
assert_eq!(a, a2);
assert_ne!(a, b);
assert_eq!(interner.len(), 2);
let terms = interner.into_terms();
assert_eq!(terms.get(a), "alpha");
assert_eq!(terms.get(b), "beta");
}
#[test]
fn parallel_inserts_are_consistent() {
let interner = Interner::new();
let inputs: Vec<String> = (0..10_000).map(|i| format!("term-{}", i % 1000)).collect();
let indices: Vec<u32> = inputs.par_iter().map(|s| interner.get_or_intern(s)).collect();
assert_eq!(interner.len(), 1000);
let terms = interner.into_terms();
for (input, idx) in inputs.iter().zip(indices.iter()) {
assert_eq!(terms.get(*idx), input.as_str());
}
}
#[test]
fn empty_interner() {
let interner = Interner::new();
assert_eq!(interner.len(), 0);
assert_eq!(interner.into_terms().len(), 0);
}
#[test]
fn concurrent_growth_bounded() {
use std::sync::{Arc, Barrier};
use std::thread;
const THREADS: usize = 64;
const STRINGS_PER_THREAD: usize = 200;
const SHARED_KEYS: usize = 50;
let interner = Arc::new(Interner::new());
let barrier = Arc::new(Barrier::new(THREADS));
let handles: Vec<_> = (0..THREADS)
.map(|tid| {
let interner = Arc::clone(&interner);
let barrier = Arc::clone(&barrier);
thread::spawn(move || {
barrier.wait();
let mut local = Vec::with_capacity(STRINGS_PER_THREAD);
for i in 0..STRINGS_PER_THREAD {
let s = if i < SHARED_KEYS { format!("shared-{i}") } else { format!("thread-{tid}-{i}") };
let idx = interner.get_or_intern(&s);
local.push((s, idx));
}
local
})
})
.collect();
let mut results: Vec<(String, u32)> = Vec::new();
for h in handles {
results.extend(h.join().unwrap());
}
let interner = Arc::try_unwrap(interner).unwrap();
let expected_unique = SHARED_KEYS + THREADS * (STRINGS_PER_THREAD - SHARED_KEYS);
assert_eq!(interner.len(), expected_unique, "racing inserts produced duplicates in storage");
let terms = interner.into_terms();
assert_eq!(terms.len(), expected_unique);
for (s, idx) in &results {
assert_eq!(terms.get(*idx), s.as_str(), "round-trip failed for {s}");
}
}
}