use crate::pod::Pod;
use core::cell::UnsafeCell;
use core::mem::MaybeUninit;
use core::sync::atomic::{fence, AtomicU64, Ordering};
#[cfg(feature = "atomic-slots")]
use core::sync::atomic::{AtomicU16, AtomicU32, AtomicU8};
#[cfg(feature = "atomic-slots")]
#[inline(always)]
unsafe fn store_payload(dst: *mut u8, src: *const u8, size: usize) {
let mut off = 0usize;
while size - off >= 8 {
let v = (src.add(off) as *const u64).read_unaligned();
AtomicU64::from_ptr(dst.add(off) as *mut u64).store(v, Ordering::Relaxed);
off += 8;
}
if size - off >= 4 {
let v = (src.add(off) as *const u32).read_unaligned();
AtomicU32::from_ptr(dst.add(off) as *mut u32).store(v, Ordering::Relaxed);
off += 4;
}
if size - off >= 2 {
let v = (src.add(off) as *const u16).read_unaligned();
AtomicU16::from_ptr(dst.add(off) as *mut u16).store(v, Ordering::Relaxed);
off += 2;
}
if size - off >= 1 {
AtomicU8::from_ptr(dst.add(off)).store(*src.add(off), Ordering::Relaxed);
}
}
#[cfg(feature = "atomic-slots")]
#[inline(always)]
unsafe fn load_payload(dst: *mut u8, src: *mut u8, size: usize) {
let mut off = 0usize;
while size - off >= 8 {
let v = AtomicU64::from_ptr(src.add(off) as *mut u64).load(Ordering::Relaxed);
(dst.add(off) as *mut u64).write_unaligned(v);
off += 8;
}
if size - off >= 4 {
let v = AtomicU32::from_ptr(src.add(off) as *mut u32).load(Ordering::Relaxed);
(dst.add(off) as *mut u32).write_unaligned(v);
off += 4;
}
if size - off >= 2 {
let v = AtomicU16::from_ptr(src.add(off) as *mut u16).load(Ordering::Relaxed);
(dst.add(off) as *mut u16).write_unaligned(v);
off += 2;
}
if size - off >= 1 {
*dst.add(off) = AtomicU8::from_ptr(src.add(off)).load(Ordering::Relaxed);
}
}
#[repr(C, align(64))]
pub(crate) struct Slot<T> {
stamp: AtomicU64,
value: UnsafeCell<MaybeUninit<T>>,
}
unsafe impl<T: Send> Sync for Slot<T> {}
unsafe impl<T: Send> Send for Slot<T> {}
const _: () = assert!(core::mem::align_of::<Slot<u64>>() == 64);
impl<T> Slot<T> {
pub(crate) fn new() -> Self {
Slot {
stamp: AtomicU64::new(0),
value: UnsafeCell::new(MaybeUninit::uninit()),
}
}
#[inline]
pub(crate) fn stamp_load(&self) -> u64 {
self.stamp.load(Ordering::Acquire)
}
}
#[cfg(not(feature = "atomic-slots"))]
impl<T: Pod> Slot<T> {
#[inline]
pub(crate) fn write(&self, seq: u64, value: T) {
let writing = seq * 2 + 1;
let done = seq * 2 + 2;
self.stamp.store(writing, Ordering::Relaxed);
fence(Ordering::Release);
unsafe { core::ptr::write_volatile(self.value.get() as *mut T, value) };
self.stamp.store(done, Ordering::Release);
}
#[inline]
pub(crate) fn write_with(&self, seq: u64, f: impl FnOnce() -> T) {
let tmp = MaybeUninit::new(f());
let writing = seq * 2 + 1;
let done = seq * 2 + 2;
self.stamp.store(writing, Ordering::Relaxed);
fence(Ordering::Release);
unsafe { core::ptr::write_volatile(self.value.get() as *mut T, tmp.assume_init()) };
self.stamp.store(done, Ordering::Release);
}
#[inline]
pub(crate) fn try_read(&self, seq: u64) -> Result<Option<T>, u64> {
let expected = seq * 2 + 2;
let s1 = self.stamp.load(Ordering::Acquire);
if s1 == expected {
let value = unsafe { core::ptr::read_volatile((*self.value.get()).as_ptr()) };
fence(Ordering::Acquire);
let s2 = self.stamp.load(Ordering::Relaxed);
if s1 == s2 {
return Ok(Some(value));
}
return Ok(None); }
if s1 & 1 != 0 {
return Ok(None);
}
Err(s1)
}
}
#[cfg(feature = "atomic-slots")]
impl<T: Pod> Slot<T> {
#[inline]
pub(crate) fn write(&self, seq: u64, value: T) {
let writing = seq * 2 + 1;
let done = seq * 2 + 2;
self.stamp.store(writing, Ordering::Relaxed);
fence(Ordering::Release);
unsafe {
store_payload(
self.value.get() as *mut u8,
&value as *const T as *const u8,
core::mem::size_of::<T>(),
)
};
self.stamp.store(done, Ordering::Release);
}
#[inline]
pub(crate) fn write_with(&self, seq: u64, f: impl FnOnce() -> T) {
self.write(seq, f());
}
#[inline]
pub(crate) fn try_read(&self, seq: u64) -> Result<Option<T>, u64> {
let expected = seq * 2 + 2;
let s1 = self.stamp.load(Ordering::Acquire);
if s1 == expected {
let mut buf = MaybeUninit::<T>::uninit();
unsafe {
load_payload(
buf.as_mut_ptr() as *mut u8,
self.value.get() as *mut u8,
core::mem::size_of::<T>(),
)
};
fence(Ordering::Acquire);
let s2 = self.stamp.load(Ordering::Relaxed);
if s1 == s2 {
return Ok(Some(unsafe { buf.assume_init() }));
}
return Ok(None); }
if s1 & 1 != 0 {
return Ok(None);
}
Err(s1)
}
}