use alloc::vec::Vec;
use core::fmt;
use xxhash_rust::xxh3::xxh3_64;
use crate::blob::{BlobHeap, BlobHeapCfg, BlobId};
use crate::error::Error;
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct TermId(pub u32);
const INITIAL_SLOTS: usize = 16;
const LOAD_NUM: usize = 7;
const LOAD_DEN: usize = 10;
const SLOT_BYTES: usize = core::mem::size_of::<u32>();
const TABLE_HEADER: usize = 2 * core::mem::size_of::<u32>();
#[derive(Clone)]
pub struct Interner<'a> {
heap: BlobHeap<'a>,
table: Vec<u32>,
len: u32,
#[cfg(feature = "counters")]
probes: u64,
}
impl<'a> Interner<'a> {
pub fn new(cfg: BlobHeapCfg) -> Self {
Self {
heap: BlobHeap::new(cfg),
table: alloc::vec![0; INITIAL_SLOTS],
len: 0,
#[cfg(feature = "counters")]
probes: 0,
}
}
pub fn intern(&mut self, s: &str) -> Result<TermId, Error> {
if (self.len as usize + 1) * LOAD_DEN > self.table.len() * LOAD_NUM {
self.rehash();
}
let bytes = s.as_bytes();
let mask = self.table.len() - 1;
let mut idx = xxh3_64(bytes) as usize & mask;
loop {
#[cfg(feature = "counters")]
{
self.probes += 1;
}
match self.table[idx] {
0 => {
let id = self.heap.push(bytes)?;
self.table[idx] = id.0 + 1;
self.len += 1;
return Ok(TermId(id.0));
}
entry if self.heap.get(BlobId(entry - 1)) == bytes => {
return Ok(TermId(entry - 1));
}
_ => idx = (idx + 1) & mask,
}
}
}
pub fn lookup(&self, s: &str) -> Option<TermId> {
let bytes = s.as_bytes();
let mask = self.table.len() - 1;
let mut idx = xxh3_64(bytes) as usize & mask;
loop {
match self.table[idx] {
0 => return None,
entry if self.heap.get(BlobId(entry - 1)) == bytes => {
return Some(TermId(entry - 1));
}
_ => idx = (idx + 1) & mask,
}
}
}
pub fn resolve(&self, id: TermId) -> &str {
core::str::from_utf8(self.heap.get(BlobId(id.0)))
.expect("interner heap holds only pushed &str bytes")
}
pub fn len(&self) -> usize {
self.len as usize
}
pub fn is_empty(&self) -> bool {
self.len == 0
}
pub fn pool_bytes(&self) -> usize {
self.heap.pool_bytes()
}
fn rehash(&mut self) {
let mask = self.table.len() * 2 - 1;
let mut table = alloc::vec![0u32; mask + 1];
for &entry in self.table.iter().filter(|&&e| e != 0) {
let mut idx = xxh3_64(self.heap.get(BlobId(entry - 1))) as usize & mask;
#[cfg(feature = "counters")]
{
self.probes += 1;
}
while table[idx] != 0 {
idx = (idx + 1) & mask;
#[cfg(feature = "counters")]
{
self.probes += 1;
}
}
table[idx] = entry;
}
self.table = table;
}
pub fn dump_index(&self, out: &mut Vec<u8>) {
self.heap.dump_index(out);
}
pub fn dump_pool(&self, out: &mut Vec<u8>) {
self.heap.dump_pool(out);
}
pub fn dump_table(&self, out: &mut Vec<u8>) {
out.reserve(TABLE_HEADER + self.table.len() * SLOT_BYTES);
out.extend_from_slice(&(self.table.len() as u32).to_le_bytes());
out.extend_from_slice(&self.len.to_le_bytes());
for &entry in &self.table {
out.extend_from_slice(&entry.to_le_bytes());
}
}
pub fn load(cfg: BlobHeapCfg, index: &[u8], pool: &[u8], table: &[u8]) -> Result<Self, Error> {
Self::from_heap_and_table(BlobHeap::load(cfg, index, pool)?, table)
}
pub fn load_borrowed(
cfg: BlobHeapCfg,
index: &[u8],
pool: &'a [u8],
table: &[u8],
) -> Result<Self, Error> {
Self::from_heap_and_table(BlobHeap::load_borrowed(cfg, index, pool)?, table)
}
fn from_heap_and_table(heap: BlobHeap<'a>, table: &[u8]) -> Result<Self, Error> {
for (_, blob) in heap.iter() {
if core::str::from_utf8(blob).is_err() {
return Err(Error::Corrupt("interned term is not valid UTF-8"));
}
}
if table.len() < TABLE_HEADER {
return Err(Error::Corrupt("interner table shorter than its header"));
}
let slots = u32::from_le_bytes(table[0..4].try_into().unwrap()) as usize;
let len = u32::from_le_bytes(table[4..TABLE_HEADER].try_into().unwrap());
if slots < INITIAL_SLOTS || !slots.is_power_of_two() {
return Err(Error::Corrupt("interner table size is not a power of two"));
}
if table.len() as u64 != TABLE_HEADER as u64 + slots as u64 * SLOT_BYTES as u64 {
return Err(Error::Corrupt("interner table length mismatch"));
}
if len as usize != heap.len() {
return Err(Error::Corrupt("interner length disagrees with its heap"));
}
if len as u64 * LOAD_DEN as u64 > slots as u64 * LOAD_NUM as u64 {
return Err(Error::Corrupt("interner table over the load factor"));
}
let mut entries = Vec::with_capacity(slots);
let mut seen = alloc::vec![false; heap.len()];
let mut filled = 0u64;
for i in 0..slots {
let at = TABLE_HEADER + i * SLOT_BYTES;
let entry = u32::from_le_bytes(table[at..at + SLOT_BYTES].try_into().unwrap());
if entry != 0 {
let id = (entry - 1) as usize;
if id >= heap.len() {
return Err(Error::Corrupt("interner table entry out of bounds"));
}
if core::mem::replace(&mut seen[id], true) {
return Err(Error::Corrupt("interner table stores an id twice"));
}
filled += 1;
}
entries.push(entry);
}
if filled != u64::from(len) {
return Err(Error::Corrupt(
"interner table entry count disagrees with len",
));
}
Ok(Self {
heap,
table: entries,
len,
#[cfg(feature = "counters")]
probes: 0,
})
}
#[cfg(feature = "counters")]
pub fn probes(&self) -> u64 {
self.probes
}
#[cfg(feature = "counters")]
pub fn reset_probes(&mut self) {
self.probes = 0;
}
}
impl fmt::Debug for Interner<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Interner")
.field("terms", &self.len)
.field("table_slots", &self.table.len())
.field("heap_bytes", &self.heap.pool_bytes())
.finish()
}
}