use super::common::{Interned, RawInterned, UnsafeLock};
use bumpalo::Bump;
use core::str;
use hashbrown::{hash_table::Entry, HashTable};
use std::{
alloc::Layout,
fmt::{self, Debug, Display, Write},
hash::{BuildHasher, Hash, Hasher},
mem::MaybeUninit,
ptr::{self, NonNull},
slice,
};
pub trait Dropless {
fn as_byte_ptr(&self) -> NonNull<u8> {
unsafe { NonNull::new_unchecked(self as *const Self as *mut Self).cast::<u8>() }
}
fn layout(&self) -> Layout {
Layout::for_value(self)
}
unsafe fn from_bytes(bytes: &[u8]) -> &Self;
unsafe fn hash<S: BuildHasher>(hash_builder: &S, bytes: &[u8]) -> u64;
unsafe fn eq(a: &[u8], b: &[u8]) -> bool;
}
macro_rules! simple_impl_dropless_fn {
(from_bytes) => {
unsafe fn from_bytes(bytes: &[u8]) -> &Self {
let ptr = bytes.as_ptr().cast::<Self>();
unsafe { ptr.as_ref().unwrap_unchecked() }
}
};
(hash) => {
unsafe fn hash<S: BuildHasher>(hash_builder: &S, bytes: &[u8]) -> u64 {
let this = unsafe { Self::from_bytes(bytes) };
let mut hasher = hash_builder.build_hasher();
this.hash(&mut hasher);
hasher.finish()
}
};
(eq) => {
unsafe fn eq(a: &[u8], b: &[u8]) -> bool {
unsafe { Self::from_bytes(a) == Self::from_bytes(b) }
}
};
}
macro_rules! impl_dropless {
($($ty:ty)*) => {
$(
impl Dropless for $ty {
simple_impl_dropless_fn!(from_bytes);
simple_impl_dropless_fn!(hash);
simple_impl_dropless_fn!(eq);
}
)*
};
}
impl_dropless!(i8 i16 i32 i64 i128 isize u8 u16 u32 u64 u128 usize bool char);
impl<T: Dropless + Hash + Eq> Dropless for [T] {
unsafe fn from_bytes(bytes: &[u8]) -> &Self {
let ptr = bytes.as_ptr().cast::<T>();
let stride = Layout::new::<T>().pad_to_align().size();
let num_elems = bytes.len() / stride;
unsafe { slice::from_raw_parts(ptr, num_elems) }
}
simple_impl_dropless_fn!(hash);
simple_impl_dropless_fn!(eq);
}
impl<T: Dropless + Hash + Eq, const N: usize> Dropless for [T; N] {
simple_impl_dropless_fn!(from_bytes);
simple_impl_dropless_fn!(hash);
simple_impl_dropless_fn!(eq);
}
impl Dropless for str {
unsafe fn from_bytes(bytes: &[u8]) -> &Self {
unsafe { std::str::from_utf8_unchecked(bytes) }
}
simple_impl_dropless_fn!(hash);
simple_impl_dropless_fn!(eq);
}
impl<T: Dropless + Hash + Eq> Dropless for Option<T> {
simple_impl_dropless_fn!(from_bytes);
simple_impl_dropless_fn!(hash);
simple_impl_dropless_fn!(eq);
}
impl<T: Dropless + Hash + Eq, E: Dropless + Hash + Eq> Dropless for Result<T, E> {
simple_impl_dropless_fn!(from_bytes);
simple_impl_dropless_fn!(hash);
simple_impl_dropless_fn!(eq);
}
#[derive(Debug)]
pub struct DroplessInterner<S = fxhash::FxBuildHasher> {
inner: UnsafeLock<DroplessInternSet<S>>,
}
impl DroplessInterner {
pub fn new() -> Self {
Self::default()
}
}
impl<S: BuildHasher> DroplessInterner<S> {
pub fn with_hasher(hash_builder: S) -> Self {
let inner = unsafe { UnsafeLock::new(DroplessInternSet::with_hasher(hash_builder)) };
Self { inner }
}
pub fn len(&self) -> usize {
self.with_inner(|set| set.len())
}
pub fn is_empty(&self) -> bool {
self.with_inner(|set| set.is_empty())
}
pub fn intern<K: Dropless + ?Sized>(&self, value: &K) -> Interned<'_, K> {
self.with_inner(|set| set.intern(value))
}
pub fn intern_formatted_str<K: Display + ?Sized>(
&self,
value: &K,
upper_size: usize,
) -> Result<Interned<'_, str>, fmt::Error> {
self.with_inner(|set| set.intern_formatted_str(value, upper_size))
}
pub fn get<K: Dropless + ?Sized>(&self, value: &K) -> Option<Interned<'_, K>> {
self.with_inner(|set| set.get(value))
}
pub fn clear(&mut self) {
self.with_inner(|set| set.clear())
}
fn with_inner<'this, F, R>(&'this self, f: F) -> R
where
F: FnOnce(&'this mut DroplessInternSet<S>) -> R,
R: 'this,
{
unsafe {
let set = self.inner.lock().as_mut();
let ret = f(set);
self.inner.unlock();
ret
}
}
}
impl<S: Default> Default for DroplessInterner<S> {
fn default() -> Self {
let inner = unsafe { UnsafeLock::new(DroplessInternSet::default()) };
Self { inner }
}
}
#[derive(Debug, Default)]
pub struct DroplessInternSet<S = fxhash::FxBuildHasher> {
bump: Bump,
str_buf: StringBuffer,
set: HashTable<DynInternEntry<S>>,
hash_builder: S,
}
impl DroplessInternSet {
pub fn new() -> Self {
Self::default()
}
}
impl<S: BuildHasher> DroplessInternSet<S> {
pub fn with_hasher(hash_builder: S) -> Self {
Self {
bump: Bump::new(),
str_buf: StringBuffer::default(),
set: HashTable::new(),
hash_builder,
}
}
pub fn intern<K: Dropless + ?Sized>(&mut self, value: &K) -> Interned<'_, K> {
let src = value.as_byte_ptr();
let layout = value.layout();
unsafe {
if layout.size() == 0 {
let mut ptr = value as *const K;
let addr_part = &mut ptr as *mut *const K as *mut *const () as *mut usize;
*addr_part = layout.align(); return Interned::unique(&*ptr);
}
let bytes = slice::from_raw_parts(src.as_ptr(), layout.size());
let hash = <K as Dropless>::hash(&self.hash_builder, bytes);
let eq = Self::table_eq::<K>(bytes, layout);
let hasher = Self::table_hasher(&self.hash_builder);
let ref_ = match self.set.entry(hash, eq, hasher) {
Entry::Occupied(entry) => Self::ref_from_entry(entry.get()),
Entry::Vacant(entry) => {
let dst = self.bump.alloc_layout(layout);
ptr::copy_nonoverlapping(src.as_ptr(), dst.as_ptr(), layout.size());
let occupied = entry.insert(DynInternEntry {
data: RawInterned(dst).cast(),
layout,
hash: <K as Dropless>::hash,
});
Self::ref_from_entry(occupied.get())
}
};
Interned::unique(ref_)
}
}
pub fn intern_formatted_str<K: Display + ?Sized>(
&mut self,
value: &K,
upper_size: usize,
) -> Result<Interned<'_, str>, fmt::Error> {
let mut write_buf = self.str_buf.speculative_alloc(upper_size);
write!(write_buf, "{value}")?;
let bytes = write_buf.as_bytes();
unsafe {
let layout = Layout::from_size_align_unchecked(bytes.len(), 1);
if bytes.is_empty() {
let dangling_ptr = NonNull::<u8>::dangling().as_ptr().cast_const();
let empty_slice = slice::from_raw_parts(dangling_ptr, 0);
let empty_str = std::str::from_utf8_unchecked(empty_slice);
return Ok(Interned::unique(empty_str));
}
let hash = <str as Dropless>::hash(&self.hash_builder, bytes);
let eq = Self::table_eq::<str>(bytes, layout);
let hasher = Self::table_hasher(&self.hash_builder);
let ref_ = match self.set.entry(hash, eq, hasher) {
Entry::Occupied(entry) => {
Self::ref_from_entry(entry.get())
}
Entry::Vacant(entry) => {
let bytes = write_buf.commit();
let ptr = NonNull::new_unchecked(bytes.as_ptr().cast_mut());
let occupied = entry.insert(DynInternEntry {
data: RawInterned(ptr).cast(),
layout,
hash: <str as Dropless>::hash,
});
Self::ref_from_entry(occupied.get())
}
};
Ok(Interned::unique(ref_))
}
}
pub fn get<K: Dropless + ?Sized>(&self, value: &K) -> Option<Interned<'_, K>> {
let ptr = value.as_byte_ptr();
let layout = value.layout();
unsafe {
let bytes = slice::from_raw_parts(ptr.as_ptr(), layout.size());
let hash = <K as Dropless>::hash(&self.hash_builder, bytes);
let eq = Self::table_eq::<K>(bytes, layout);
let entry = self.set.find(hash, eq)?;
let ref_ = Self::ref_from_entry(entry);
Some(Interned::unique(ref_))
}
}
pub fn len(&self) -> usize {
self.set.len()
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn clear(&mut self) {
self.bump.reset();
self.set.clear();
}
#[inline]
fn table_eq<'a, K: Dropless + ?Sized>(
key: &'a [u8],
layout: Layout,
) -> impl FnMut(&DynInternEntry<S>) -> bool + 'a {
move |entry: &DynInternEntry<S>| unsafe {
if layout == entry.layout {
let entry_bytes =
slice::from_raw_parts(entry.data.cast::<u8>().as_ptr(), entry.layout.size());
<K as Dropless>::eq(key, entry_bytes)
} else {
false
}
}
}
#[inline]
fn table_hasher<'a>(hash_builder: &'a S) -> impl Fn(&DynInternEntry<S>) -> u64 + 'a {
|entry: &DynInternEntry<S>| unsafe {
let entry_bytes =
slice::from_raw_parts(entry.data.cast::<u8>().as_ptr(), entry.layout.size());
(entry.hash)(hash_builder, entry_bytes)
}
}
#[inline]
unsafe fn ref_from_entry<'a, K: Dropless + ?Sized>(entry: &DynInternEntry<S>) -> &'a K {
unsafe {
let bytes =
slice::from_raw_parts(entry.data.cast::<u8>().as_ptr(), entry.layout.size());
K::from_bytes(bytes)
}
}
}
#[derive(Debug, Default)]
struct StringBuffer {
chunks: Vec<Vec<MaybeUninit<u8>>>,
last_chunk_start: usize,
}
const INIT_CHUNK_SIZE: usize = 1 << 5;
const GROW_MAX_CHUNK_SIZE: usize = 1 << 12;
impl StringBuffer {
fn speculative_alloc(&mut self, upper_size: usize) -> StringWriteBuffer<'_> {
match self.chunks.last() {
None => {
let chunk_size = INIT_CHUNK_SIZE.max(upper_size.next_power_of_two());
self.append_new_chunk(chunk_size)
}
Some(last_chunk) if last_chunk.len() - self.last_chunk_start < upper_size => {
let chunk_size = (last_chunk.len() * 2)
.min(GROW_MAX_CHUNK_SIZE)
.max(upper_size.next_power_of_two());
self.append_new_chunk(chunk_size);
}
_ => {}
}
let buf = unsafe {
let last_chunk = self.chunks.last_mut().unwrap_unchecked();
last_chunk.as_mut_ptr().add(self.last_chunk_start)
};
StringWriteBuffer {
buf,
buf_cap: upper_size,
last_chuck_start: &mut self.last_chunk_start,
written: 0,
}
}
#[inline]
fn append_new_chunk(&mut self, chunk_size: usize) {
let mut chunk: Vec<MaybeUninit<u8>> = Vec::with_capacity(chunk_size);
unsafe { chunk.set_len(chunk_size) };
self.chunks.push(chunk);
self.last_chunk_start = 0;
}
}
struct StringWriteBuffer<'a> {
buf: *mut MaybeUninit<u8>,
buf_cap: usize,
last_chuck_start: &'a mut usize,
written: usize,
}
impl<'a> StringWriteBuffer<'a> {
#[inline]
fn as_bytes(&self) -> &[u8] {
unsafe { slice::from_raw_parts(self.buf.cast::<u8>(), self.written) }
}
#[inline]
fn commit(self) -> &'a [u8] {
*self.last_chuck_start += self.written;
unsafe { slice::from_raw_parts(self.buf.cast::<u8>(), self.written) }
}
}
impl Write for StringWriteBuffer<'_> {
fn write_str(&mut self, s: &str) -> std::fmt::Result {
let size = s.len();
if self.buf_cap - self.written >= size {
let src = s.as_ptr();
let dst = unsafe { self.buf.add(self.written).cast::<u8>() };
unsafe { ptr::copy_nonoverlapping(src, dst, size) };
self.written += size;
Ok(())
} else {
Err(std::fmt::Error)
}
}
}
struct DynInternEntry<S> {
data: RawInterned,
layout: Layout,
hash: unsafe fn(&S, &[u8]) -> u64,
}
impl<S> Debug for DynInternEntry<S> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("DynInternEntry")
.field("data", &self.data)
.field("layout", &self.layout)
.finish_non_exhaustive()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::common;
use std::num::NonZeroU32;
#[test]
fn test_dropless_interner() {
test_dropless_interner_int();
test_dropless_interner_str();
test_dropless_interner_bytes();
test_dropless_interner_mixed();
test_dropless_interner_many();
test_dropless_interner_alignment_handling();
test_dropless_interner_complex_display_type();
test_dropless_interner_consecutive_formatted_strs();
}
fn test_dropless_interner_int() {
let interner = DroplessInterner::new();
let a = interner.intern(&0_u32).erased_raw();
let b = interner.intern(&0_u32).erased_raw();
let c = interner.intern(&1_u32).erased_raw();
let d = interner.intern(&1_i32).erased_raw();
let groups: [&[RawInterned]; _] = [&[a, b], &[c, d]];
common::assert_group_addr_eq(&groups);
}
fn test_dropless_interner_str() {
let interner = DroplessInterner::new();
let a = interner.intern("apple").erased_raw();
let b = interner.intern(*Box::new("apple")).erased_raw();
let c = interner.intern(&*String::from("apple")).erased_raw();
let d = interner.intern("banana").erased_raw();
let e = interner.intern("").erased_raw();
let f = interner.intern(&*String::from("")).erased_raw();
let g = interner.intern("42").erased_raw();
let h = interner.intern_formatted_str(&42, 2).unwrap().erased_raw();
let i = interner
.intern_formatted_str(&NonZeroU32::new(43).unwrap(), 2)
.unwrap()
.erased_raw();
let j = interner.intern(&*43.to_string()).erased_raw();
let groups: [&[RawInterned]; _] = [&[a, b, c], &[d], &[e, f], &[g, h], &[i, j]];
common::assert_group_addr_eq(&groups);
}
fn test_dropless_interner_bytes() {
let interner = DroplessInterner::new();
let a = interner.intern(&[0, 1]).erased_raw();
let boxed: Box<[i32]> = Box::new([0, 1]);
let b = interner.intern(&*boxed).erased_raw();
let c = interner.intern(&*vec![0, 1]).erased_raw();
let d = interner.intern(&[2, 3]).erased_raw();
let groups: [&[RawInterned]; _] = [&[a, b, c], &[d]];
common::assert_group_addr_eq(&groups);
}
fn test_dropless_interner_mixed() {
let interner = DroplessInterner::new();
interner.intern("apple");
interner.intern(&1_u32);
interner.intern(&[2, 3]);
interner.intern_formatted_str(&42, 10).unwrap();
assert_eq!(interner.get("apple").as_deref(), Some("apple"));
assert!(interner.get("banana").is_none());
assert_eq!(interner.get(&1_u32).as_deref(), Some(&1_u32));
assert!(interner.get(&2_u32).is_none());
assert_eq!(interner.get(&[2, 3]).as_deref(), Some(&[2, 3]));
assert!(interner.get(&[2]).is_none());
assert_eq!(interner.get("42").as_deref(), Some("42"));
}
fn test_dropless_interner_many() {
let interner = DroplessInterner::new();
let mut interned_int = Vec::new();
let mut interned_str = Vec::new();
let mut interned_bytes = Vec::new();
#[cfg(not(miri))]
const N: usize = 1000;
#[cfg(miri)]
const N: usize = 50;
let strs = (0..N).map(|i| (i * 10_000).to_string()).collect::<Vec<_>>();
for (i, str_) in strs.iter().enumerate().take(N) {
let int = &(i as u32); let str_ = str_.as_str(); let bytes = &[i as u16]; interned_int.push(interner.intern(int).erased_raw());
interned_str.push(interner.intern(str_).erased_raw());
interned_bytes.push(interner.intern(bytes).erased_raw());
}
interned_int.sort_unstable();
interned_int.dedup();
interned_str.sort_unstable();
interned_str.dedup();
interned_bytes.sort_unstable();
interned_bytes.dedup();
assert_eq!(interned_int.len(), N);
assert_eq!(interned_str.len(), N);
assert_eq!(interned_bytes.len(), N);
let whole = interned_int
.iter()
.chain(&interned_str)
.chain(&interned_bytes)
.cloned()
.collect::<fxhash::FxHashSet<_>>();
assert_eq!(whole.len(), N * 3);
for (i, str_) in strs.iter().enumerate().take(N) {
let int = &(i as u32); let str_ = str_.as_str(); let bytes = &[i as u16]; assert_eq!(interner.get(int).as_deref(), Some(int));
assert_eq!(interner.get(str_).as_deref(), Some(str_));
assert_eq!(interner.get(bytes).as_deref(), Some(bytes));
}
}
#[rustfmt::skip]
fn test_dropless_interner_alignment_handling() {
#[derive(PartialEq, Eq, Hash, Debug)] #[repr(C, align(4))] struct T4 (u16, [u8; 2]);
#[derive(PartialEq, Eq, Hash, Debug)] #[repr(C, align(8))] struct T8 (u16, [u8; 6]);
#[derive(PartialEq, Eq, Hash, Debug)] #[repr(C, align(16))] struct T16 (u16, [u8; 14]);
macro_rules! impl_for_first_2bytes {
() => {
unsafe fn hash<S: BuildHasher>(hash_builder: &S, bytes: &[u8]) -> u64 {
hash_builder.hash_one(&bytes[0..2])
}
unsafe fn eq(a: &[u8], b: &[u8]) -> bool {
a[0..2] == b[0..2]
}
};
}
impl Dropless for T4 { simple_impl_dropless_fn!(from_bytes); impl_for_first_2bytes!(); }
impl Dropless for T8 { simple_impl_dropless_fn!(from_bytes); impl_for_first_2bytes!(); }
impl Dropless for T16 { simple_impl_dropless_fn!(from_bytes); impl_for_first_2bytes!(); }
let interner = DroplessInterner::new();
let mut interned_4 = Vec::new();
let mut interned_8 = Vec::new();
let mut interned_16 = Vec::new();
#[cfg(not(miri))]
const N: usize = 1000;
#[cfg(miri)]
const N: usize = 50;
for i in 0..N {
let t4 = T4(i as u16, [4; _]);
let t8 = T8(i as u16, [8; _]);
let t16 = T16(i as u16, [16; _]);
interned_4.push(interner.intern(&t4).erased_raw());
interned_8.push(interner.intern(&t8).erased_raw());
interned_16.push(interner.intern(&t16).erased_raw());
}
for i in 0..N {
let t4 = T4(i as u16, [4; _]);
let t8 = T8(i as u16, [8; _]);
let t16 = T16(i as u16, [16; _]);
unsafe {
assert_eq!(*<Interned<'_, T4>>::from_erased_raw(interned_4[i]), t4);
assert_eq!(*<Interned<'_, T8>>::from_erased_raw(interned_8[i]), t8);
assert_eq!(*<Interned<'_, T16>>::from_erased_raw(interned_16[i]), t16);
}
}
}
fn test_dropless_interner_complex_display_type() {
let interner = DroplessInterner::new();
#[allow(unused)]
#[derive(Debug)]
struct A<'a> {
int: i32,
float: f32,
text: &'a str,
bytes: &'a [u8],
}
impl Display for A<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt::Debug::fmt(self, f)
}
}
let a = A {
int: 123,
float: 456.789,
text: "this is a text",
bytes: &[1, 2, 3, 4, 5, 6, 7, 8, 9, 0],
};
let interned = interner.intern_formatted_str(&a, 1_000).unwrap();
let mut s = String::new();
write!(&mut s, "{a}").unwrap();
assert_eq!(&*interned, s.as_str());
}
fn test_dropless_interner_consecutive_formatted_strs() {
let dropless = DroplessInterner::new();
let value = 0;
let i1 = dropless.intern_formatted_str(&value, 1).unwrap();
let value = 1;
let i2 = dropless.intern_formatted_str(&value, 1).unwrap();
let _c = format!("{i1}{i2}");
}
#[test]
fn test_string_buffer() {
test_string_buffer_chunk();
test_string_buffer_discard();
test_string_buffer_insufficient_upper_size();
test_string_buffer_long_string();
}
fn test_string_buffer_chunk() {
let mut buf = StringBuffer::default();
assert_eq!(buf.chunks.len(), 0);
let s = "a".repeat(INIT_CHUNK_SIZE);
let mut write_buf = buf.speculative_alloc(s.len());
write_buf.write_str(&s).unwrap();
assert_eq!(write_buf.as_bytes(), s.as_bytes());
write_buf.commit();
assert_eq!(buf.chunks.len(), 1);
assert_eq!(buf.last_chunk_start, s.len());
let mut write_buf = buf.speculative_alloc(1);
write_buf.write_str("a").unwrap();
assert_eq!(write_buf.as_bytes(), b"a");
write_buf.commit();
assert_eq!(buf.chunks.len(), 2);
assert_eq!(buf.last_chunk_start, 1);
let mut write_buf = buf.speculative_alloc(GROW_MAX_CHUNK_SIZE);
write_buf.write_str("aa").unwrap();
assert_eq!(write_buf.as_bytes(), b"aa");
write_buf.commit();
assert_eq!(buf.chunks.len(), 3);
assert_eq!(buf.last_chunk_start, 2);
}
fn test_string_buffer_discard() {
let mut buf = StringBuffer::default();
for _ in 0..10 {
let s = "a".repeat(INIT_CHUNK_SIZE);
let mut write_buf = buf.speculative_alloc(s.len());
write_buf.write_str(&s).unwrap();
assert_eq!(write_buf.as_bytes(), s.as_bytes());
assert_eq!(buf.chunks.len(), 1);
assert_eq!(buf.last_chunk_start, 0);
}
}
fn test_string_buffer_insufficient_upper_size() {
let mut buf = StringBuffer::default();
for _ in 0..10 {
let mut write_buf = buf.speculative_alloc(5);
let res = write_buf.write_str("this is longer than 5");
assert!(res.is_err());
write_buf.commit();
assert_eq!(buf.chunks.len(), 1);
assert_eq!(buf.last_chunk_start, 0);
}
}
fn test_string_buffer_long_string() {
let mut buf = StringBuffer::default();
let s = "a".repeat(GROW_MAX_CHUNK_SIZE * 10);
let mut write_buf = buf.speculative_alloc(s.len());
write_buf.write_str(&s).unwrap();
assert_eq!(write_buf.as_bytes(), s.as_bytes());
write_buf.commit();
assert_eq!(buf.chunks.len(), 1);
assert_eq!(buf.last_chunk_start, s.len());
}
}