use std::{marker::PhantomData, mem::size_of_val, hash::Hash, sync::atomic::*, alloc::*, ptr::*};
pub trait Counter: Sized {
fn new() -> Self;
fn increment(&mut self);
fn decrement(&mut self) -> bool;
}
macro_rules! impl_ref_count_for_primitive {
($($t:ty),*) => {
$(
impl Counter for $t {
#[inline(always)] fn new() -> Self { 1 }
#[inline(always)] fn increment(&mut self) { *self += 1 }
#[inline(always)] fn decrement(&mut self) -> bool {
*self -= 1;
*self == 0
}
}
impl Counter for std::cell::Cell<$t> {
#[inline(always)] fn new() -> Self { std::cell::Cell::new(1) }
#[inline(always)] fn increment(&mut self) { self.set(self.get().checked_add(1).expect("RefCount overflow")); }
#[inline(always)] fn decrement(&mut self) -> bool {
let value = self.get().checked_sub(1).expect("RefCount underflow");
self.set(value);
value == 0
}
}
)*
};
}
macro_rules! impl_ref_count_for_atomic {
($($atomic:ty),*) => {
$(
impl Counter for $atomic {
#[inline(always)] fn new() -> Self { <$atomic>::new(1) }
#[inline(always)] fn increment(&mut self) { self.fetch_add(1, Ordering::Release); }
#[inline(always)] fn decrement(&mut self) -> bool {
if self.fetch_sub(1, Ordering::Release) == 1 {
fence(Ordering::Acquire); true
} else { false }
}
}
)*
};
}
impl_ref_count_for_primitive!(u8, u16, u32, u64, u128, usize);
impl_ref_count_for_atomic!(AtomicU8, AtomicU16, AtomicU32, AtomicU64, AtomicUsize);
#[derive(Debug)]
pub struct Rime<C: Counter, T: ?Sized> {
_marker: PhantomData<(C, T)>,
counter_ptr: *mut C,
inner_ptr: *const T,
}
impl<C: Counter, T: Sized> Rime<C, T> {
pub fn steal(value: T) -> Self {
unsafe {
let c_size = size_of::<C>();
let layout = Layout::from_size_align_unchecked(
size_of::<T>() + c_size,
align_of::<C>().max(align_of::<T>()));
let raw = alloc(layout);
if raw.is_null() {
dealloc(raw, layout);
handle_alloc_error(layout);
}
let counter_ptr = raw as *mut C;
write(counter_ptr, C::new());
let data_ptr = raw.add(c_size) as *mut T;
write(data_ptr, value);
Self::from_raw(counter_ptr, data_ptr as *const T)
}
}
}
impl<C: Counter, T: ?Sized> Rime<C, T> {
#[inline(always)]
pub fn from_raw(counter_ptr: *mut C, inner_ptr: *const T) -> Self {
Self { _marker: PhantomData, counter_ptr, inner_ptr }
}
#[inline(always)]
pub fn from_raw_parts(counter_ptr: *mut C, inner_ptr: *mut u8, metadata: <T as Pointee>::Metadata) -> Self {
Self::from_raw(counter_ptr, from_raw_parts::<T>(inner_ptr, metadata))
}
pub fn new(value: &T) -> Self {
unsafe {
let t_size = size_of_val(value);
let c_size = size_of::<C>();
let layout = Layout::from_size_align_unchecked(
c_size + t_size,
align_of::<C>().max(align_of_val(value))
);
let raw = alloc(layout);
if raw.is_null() {
dealloc(raw, layout);
handle_alloc_error(layout);
}
let counter_ptr = raw as *mut C;
write(counter_ptr, C::new());
let inner_ptr = raw.add(c_size);
copy_nonoverlapping(value as *const T as *const u8, inner_ptr, t_size);
Self::from_raw_parts(counter_ptr, inner_ptr, metadata(value))
}
}
#[inline(always)]
pub fn as_ptr(&self) -> *const T {
self.inner_ptr
}
#[inline(always)]
pub fn as_mut_ptr(&self) -> *mut T {
self.inner_ptr.cast_mut()
}
}
impl<C: Counter, T: ?Sized> Drop for Rime<C, T> {
#[inline(always)]
fn drop(&mut self) {
unsafe {
if (*self.counter_ptr).decrement() {
let inner = &*self.inner_ptr;
dealloc(
self.counter_ptr.cast(),
Layout::from_size_align_unchecked(
size_of::<C>() + size_of_val(inner),
align_of::<C>().max(align_of_val(inner))
)
);
}
}
}
}
impl<C: Counter, T: ?Sized> Clone for Rime<C, T> {
#[inline]
fn clone(&self) -> Self {
unsafe { (*self.counter_ptr).increment() }
Self {
inner_ptr: self.inner_ptr,
counter_ptr: self.counter_ptr,
_marker: PhantomData
}
}
}
impl<C: Counter, T: ?Sized> From<&Rime<C, T>> for Rime<C, T> {
#[inline(always)]
fn from(value: &Rime<C, T>) -> Self {
value.clone()
}
}
impl<C: Counter, T: ?Sized> AsRef<T> for Rime<C, T> {
#[inline]
fn as_ref(&self) -> &T {
unsafe { &*self.inner_ptr }
}
}
impl<C: Counter, T: ?Sized> std::ops::Deref for Rime<C, T> {
type Target = T;
#[inline]
fn deref(&self) -> &Self::Target {
unsafe { &*self.inner_ptr }
}
}
impl<C: Counter, T: ?Sized> Eq for Rime<C, T> { }
impl<C: Counter, T: ?Sized> PartialEq for Rime<C, T> {
#[inline(always)]
fn eq(&self, other: &Self) -> bool {
addr_eq(self.inner_ptr, other.inner_ptr)
}
}
impl<C: Counter, T: ?Sized + Ord> Ord for Rime<C, T> {
#[inline(always)]
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
unsafe { (&*self.inner_ptr).cmp(&*other) }
}
}
impl<C: Counter, T: ?Sized + PartialOrd> PartialOrd for Rime<C, T> {
#[inline(always)]
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
unsafe { (&*self.inner_ptr).partial_cmp(&*other) }
}
}
impl<C: Counter, T: ?Sized + Hash> Hash for Rime<C, T> {
#[inline(always)]
fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
unsafe { (&*self.inner_ptr).hash(state) }
}
}
unsafe impl<C: Counter + Send, T: ?Sized + Send> Send for Rime<C, T> {}
unsafe impl<C: Counter + Sync, T: ?Sized + Sync> Sync for Rime<C, T> {}
#[cfg(test)]
mod tests {
use std::sync::atomic::*;
use super::*;
#[test]
fn test_basic_clone_and_deref() {
let rime = Rime::<u8, str>::new("hello");
assert_eq!(&*rime, "hello");
let cloned = rime.clone();
assert_eq!(&*cloned, "hello");
assert_eq!(rime.as_ptr(), cloned.as_ptr()); }
#[test]
fn test_equality_and_ordering() {
let r1 = Rime::<u8, str>::new("abc");
let r2 = r1.clone();
let r3 = Rime::<u8, str>::new("abc");
assert_eq!(r1, r2);
assert_ne!(r1, r3); assert!(r1 <= r2);
assert!(r3 >= r1);
}
#[test]
fn test_drop_deallocates() {
use std::cell::RefCell;
use std::rc::Rc;
struct DropCounter(Rc<RefCell<u8>>);
impl Drop for DropCounter {
fn drop(&mut self) {
*self.0.borrow_mut() += 1;
}
}
let dropped = Rc::new(RefCell::new(0));
{
let counter = DropCounter(dropped.clone());
let r1 = Rime::<usize, _>::new(&counter);
let _r2 = r1.clone(); }
assert_eq!(*dropped.borrow(), 1); }
#[test]
fn test_atomic_clone_thread_safe() {
use std::thread;
let rime = Rime::<AtomicUsize, str>::new("multi");
let mut handles = vec![];
for _ in 0..10 {
let cloned = rime.clone();
handles.push(thread::spawn(move || {
assert_eq!(&*cloned, "multi");
}));
}
for handle in handles {
handle.join().unwrap();
}
assert_eq!(&*rime, "multi");
}
#[test]
fn test_as_ref_and_conversion() {
let rime = Rime::<u8, str>::new("as_ref test");
let rime2: Rime<u8, str> = (&rime).into();
assert_eq!(rime.as_ref(), rime2.as_ref());
}
}