use std::cell::UnsafeCell;
use std::ops::{Deref, DerefMut};
use std::pin::Pin;
use std::sync::atomic::{AtomicU32, AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use std::task::{Poll, Waker};
use std::time::Instant;
use polars_async::ASYNC;
use polars_utils::UnitVec;
use polars_utils::with_drop::WithDrop;
use crate::{SpillContextParam, Spillable, WeakSpillContext};
const SPILLED_BIT: u64 = 1; const DROPPED_BIT: u64 = 2; const LOCK_BIT: u64 = 4; const HAS_WAITERS_BIT: u64 = 8; const RO_PIN_COUNT_UNIT: u64 = 16; const RO_PIN_MASK: u64 = u64::MAX << 4;
enum ValueSlot<T> {
InMemory(T),
Spilled {
n_bytes: usize,
spill_ctx: WeakSpillContext,
reinsert_reg_id: u32,
spill_time_ns: u64,
spilled_start: Instant,
},
Dropped,
}
#[derive(Default)]
struct LockState {
waiters: UnitVec<Waker>,
cur_ctx: Option<(WeakSpillContext, SpillContextParam)>,
}
struct SpillTokenInner<T: Spillable> {
value_slot: UnsafeCell<ValueSlot<T>>,
spilled_value: UnsafeCell<Option<T::Spilled>>,
lock: Mutex<LockState>,
registration_id: AtomicU32,
state: AtomicU64,
}
unsafe impl<T: Spillable + Send> Send for SpillTokenInner<T> {}
unsafe impl<T: Spillable + Sync> Sync for SpillTokenInner<T> {}
impl<T: Spillable> SpillTokenInner<T> {
async fn wait(&self, mask: u64) -> u64 {
std::future::poll_fn(|ctx| {
let mut lock = self.lock.lock().unwrap();
let mut state = self.state.load(Ordering::Acquire);
if state & mask != 0 {
if state & HAS_WAITERS_BIT == 0 {
state = self.state.fetch_add(HAS_WAITERS_BIT, Ordering::AcqRel);
if state & mask == 0 {
self.state.fetch_sub(HAS_WAITERS_BIT, Ordering::Relaxed);
return Poll::Ready(state);
}
}
lock.waiters.push(ctx.waker().clone());
Poll::Pending
} else {
Poll::Ready(state)
}
})
.await
}
#[inline(always)]
fn wake_waiters(&self, state: u64) {
if state & HAS_WAITERS_BIT != 0 {
self.wake_waiters_slow();
}
}
#[inline(never)]
#[cold]
fn wake_waiters_slow(&self) {
let mut lock = self.lock.lock().unwrap();
let waiters = core::mem::take(&mut lock.waiters);
self.state.fetch_sub(HAS_WAITERS_BIT, Ordering::Relaxed);
drop(lock);
for w in waiters {
w.wake();
}
}
fn try_pin(&self) -> Option<PinnedRef<'_, T>> {
self.state
.try_update(Ordering::Acquire, Ordering::Relaxed, |state| {
if state & (SPILLED_BIT | LOCK_BIT | DROPPED_BIT) != 0 {
return None;
}
Some(state + RO_PIN_COUNT_UNIT)
})
.ok()
.map(|_| PinnedRef { inner: self })
}
async fn pin_or_lock(&self) -> Option<PinnedRef<'_, T>> {
let mut state = self.state.load(Ordering::Relaxed);
loop {
if state & (SPILLED_BIT | LOCK_BIT | DROPPED_BIT) == 0 {
match self.state.compare_exchange_weak(
state,
state + RO_PIN_COUNT_UNIT,
Ordering::Acquire,
Ordering::Relaxed,
) {
Ok(_) => return Some(PinnedRef { inner: self }),
Err(s) => state = s,
}
} else if state & (LOCK_BIT | RO_PIN_MASK | DROPPED_BIT) == 0 {
match self.state.compare_exchange_weak(
state,
state | LOCK_BIT,
Ordering::Acquire,
Ordering::Relaxed,
) {
Ok(_) => return None,
Err(s) => state = s,
}
} else {
assert!(state & DROPPED_BIT == 0);
state = self.wait(LOCK_BIT).await;
}
}
}
async fn lock(&self) {
let mut state = self.state.load(Ordering::Relaxed);
loop {
if state & (LOCK_BIT | RO_PIN_MASK | DROPPED_BIT) == 0 {
match self.state.compare_exchange_weak(
state,
state | LOCK_BIT,
Ordering::Acquire,
Ordering::Relaxed,
) {
Ok(_) => return,
Err(s) => state = s,
}
} else {
assert!(state & DROPPED_BIT == 0);
state = self.wait(LOCK_BIT | RO_PIN_MASK).await;
}
}
}
async fn pin(slf: &Arc<Self>) -> PinnedRef<'_, T> {
if let Some(r) = slf.try_pin() {
return r;
}
std::hint::cold_path();
if let Some(r) = slf.pin_or_lock().await {
return r;
}
unsafe {
debug_assert!(
slf.state.load(Ordering::Relaxed) & (SPILLED_BIT | LOCK_BIT | RO_PIN_MASK)
== (SPILLED_BIT | LOCK_BIT)
);
let lock_guard = WithDrop::new(slf, |slf| {
slf.wake_waiters(slf.state.fetch_and(!LOCK_BIT, Ordering::AcqRel));
});
let unspill_start = Instant::now();
let spilled = (*slf.spilled_value.get()).as_ref().unwrap();
let value = T::unspill(spilled).await;
let ValueSlot::Spilled {
n_bytes,
spill_ctx,
reinsert_reg_id,
spill_time_ns,
spilled_start,
} = slf.value_slot.get().replace(ValueSlot::InMemory(value))
else {
unreachable!()
};
if let Some(strong) = spill_ctx.upgrade() {
strong
.stats()
.add_unspill(n_bytes, spill_time_ns, spilled_start, unspill_start);
}
if reinsert_reg_id == slf.registration_id.load(Ordering::Relaxed) {
let dyn_slf: Arc<dyn DynSpillToken> = slf.clone();
spill_ctx.0.reinsert(&dyn_slf, reinsert_reg_id, spill_ctx.1);
}
WithDrop::dismiss(lock_guard);
slf.wake_waiters(
slf.state
.fetch_add(RO_PIN_COUNT_UNIT - LOCK_BIT - SPILLED_BIT, Ordering::AcqRel),
);
PinnedRef { inner: slf }
}
}
fn pin_blocking(slf: &Arc<Self>) -> PinnedRef<'_, T> {
if let Some(r) = slf.try_pin() {
return r;
}
std::hint::cold_path();
ASYNC.block_in_place_on(Self::pin(slf))
}
async fn pin_mut(slf: &Arc<Self>) -> PinnedMut<'_, T> {
unsafe {
slf.lock().await;
let lock_guard = WithDrop::new(slf, |slf| {
slf.wake_waiters(slf.state.fetch_and(!LOCK_BIT, Ordering::AcqRel));
});
let value_slot = &mut *slf.value_slot.get();
if let ValueSlot::Spilled {
n_bytes,
spill_ctx,
reinsert_reg_id,
spill_time_ns,
spilled_start,
} = value_slot
{
debug_assert!(slf.state.load(Ordering::Relaxed) & SPILLED_BIT == SPILLED_BIT);
let unspill_start = Instant::now();
let spilled = (*slf.spilled_value.get()).take().unwrap();
let value = T::unspill(&spilled).await;
if let Some(strong) = spill_ctx.upgrade() {
strong.stats().add_unspill(
*n_bytes,
*spill_time_ns,
*spilled_start,
unspill_start,
);
}
if *reinsert_reg_id == slf.registration_id.load(Ordering::Relaxed) {
let dyn_slf: Arc<dyn DynSpillToken> = slf.clone();
spill_ctx
.0
.reinsert(&dyn_slf, *reinsert_reg_id, spill_ctx.1);
}
*value_slot = ValueSlot::InMemory(value);
slf.state.fetch_sub(SPILLED_BIT, Ordering::Relaxed);
} else {
*slf.spilled_value.get() = None;
}
WithDrop::dismiss(lock_guard);
PinnedMut { inner: slf }
}
}
fn pin_mut_blocking(slf: &Arc<Self>) -> PinnedMut<'_, T> {
ASYNC.block_in_place_on(Self::pin_mut(slf))
}
unsafe fn unpin(&self) {
let old_s = self.state.fetch_sub(RO_PIN_COUNT_UNIT, Ordering::AcqRel);
if old_s & RO_PIN_MASK == RO_PIN_COUNT_UNIT {
self.wake_waiters(old_s);
}
}
unsafe fn unpin_mut(&self) {
self.wake_waiters(self.state.fetch_sub(LOCK_BIT, Ordering::AcqRel));
}
unsafe fn mark_as_dropped(&self) {
unsafe {
let old_state = self.state.fetch_or(DROPPED_BIT, Ordering::Acquire);
if old_state & (LOCK_BIT | RO_PIN_MASK) == 0 {
self.value_slot.get().replace(ValueSlot::Dropped);
self.spilled_value.get().replace(None);
}
let mut lock = self.lock.lock().unwrap();
self.registration_id.fetch_add(1, Ordering::Relaxed);
lock.cur_ctx = None;
}
}
}
impl<T, S> Clone for SpillTokenInner<T>
where
T: Clone + Spillable<Spilled = S>,
S: Clone,
{
fn clone(&self) -> Self {
if let Some(r) = ASYNC.block_in_place_on(self.pin_or_lock()) {
return SpillTokenInner {
value_slot: UnsafeCell::new(ValueSlot::InMemory(r.clone())),
spilled_value: UnsafeCell::new(None),
registration_id: AtomicU32::new(0),
state: AtomicU64::new(0),
lock: Mutex::default(),
};
}
unsafe {
let lock_guard = WithDrop::new(self, |slf| {
slf.wake_waiters(slf.state.fetch_and(!LOCK_BIT, Ordering::AcqRel));
});
let ValueSlot::Spilled {
n_bytes,
spill_ctx,
reinsert_reg_id: _,
spill_time_ns: _,
spilled_start: _,
} = &*lock_guard.value_slot.get()
else {
unreachable!()
};
let clone_spill_start = Instant::now();
let spilled_value = (&*lock_guard.spilled_value.get()).as_ref().unwrap().clone();
let (spill_time_ns, spilled_start) = if let Some(strong) = spill_ctx.upgrade() {
strong
.stats()
.add_successful_spill(*n_bytes, clone_spill_start)
} else {
(0, Instant::now())
};
SpillTokenInner {
value_slot: UnsafeCell::new(ValueSlot::Spilled {
n_bytes: *n_bytes,
spill_ctx: spill_ctx.clone(),
reinsert_reg_id: 0,
spill_time_ns,
spilled_start,
}),
spilled_value: UnsafeCell::new(Some(spilled_value)),
registration_id: AtomicU32::new(0),
state: AtomicU64::new(SPILLED_BIT),
lock: Mutex::default(),
}
}
}
}
pub enum TrySpillError {
AlreadySpilled,
Pinned,
}
pub(crate) trait DynSpillToken: Send + Sync + 'static {
fn register(&self, ctx: WeakSpillContext, param: SpillContextParam) -> u32;
fn unregister(&self) -> Option<(WeakSpillContext, SpillContextParam)>;
fn current_registration_id(&self) -> u32;
fn can_spill(&self) -> bool;
fn is_spilled_or_dropped(&self) -> bool;
fn estimate_byte_size(&self) -> Option<usize>;
fn try_spill(
&self,
context: WeakSpillContext,
registration_id: u32,
) -> Result<Pin<Box<dyn Future<Output = bool> + Send + '_>>, TrySpillError>;
}
impl<T: Spillable> DynSpillToken for SpillTokenInner<T> {
fn register(&self, ctx: WeakSpillContext, param: SpillContextParam) -> u32 {
let mut lock = self.lock.lock().unwrap();
lock.cur_ctx = Some((ctx, param));
self.registration_id.fetch_add(1, Ordering::Release) + 1
}
fn unregister(&self) -> Option<(WeakSpillContext, SpillContextParam)> {
let mut lock = self.lock.lock().unwrap();
self.registration_id.fetch_add(1, Ordering::Release);
lock.cur_ctx.take()
}
fn current_registration_id(&self) -> u32 {
self.registration_id.load(Ordering::Relaxed)
}
fn can_spill(&self) -> bool {
self.state.load(Ordering::Acquire) & (SPILLED_BIT | DROPPED_BIT | LOCK_BIT | RO_PIN_MASK)
== 0
}
fn is_spilled_or_dropped(&self) -> bool {
self.state.load(Ordering::Acquire) & (SPILLED_BIT | DROPPED_BIT) != 0
}
fn estimate_byte_size(&self) -> Option<usize> {
self.try_pin().map(|p| p.estimate_byte_size())
}
fn try_spill(
&self,
ctx: WeakSpillContext,
registration_id: u32,
) -> Result<Pin<Box<dyn Future<Output = bool> + Send + '_>>, TrySpillError> {
let pin_update = self
.state
.try_update(Ordering::Relaxed, Ordering::Acquire, |state| {
if state & (LOCK_BIT | DROPPED_BIT | RO_PIN_MASK) != 0 {
return None;
}
Some(state | LOCK_BIT)
});
let Ok(state) = pin_update else {
return Err(TrySpillError::Pinned);
};
if state & SPILLED_BIT != 0 {
let ValueSlot::Spilled {
spill_ctx: reinsert_ctx,
reinsert_reg_id: reinsert_id,
..
} = (unsafe { &mut *self.value_slot.get() })
else {
unreachable!()
};
*reinsert_ctx = ctx;
*reinsert_id = registration_id;
self.wake_waiters(self.state.fetch_sub(LOCK_BIT, Ordering::AcqRel));
return Err(TrySpillError::AlreadySpilled);
}
Ok(Box::pin(async move {
let spill_start = Instant::now();
let needs_spill = unsafe { (*self.spilled_value.get()).is_none() };
let is_exclusive = if needs_spill {
self.wake_waiters(
self.state
.fetch_add(RO_PIN_COUNT_UNIT - LOCK_BIT, Ordering::AcqRel),
);
let pin_guard = PinnedRef { inner: self };
let spilled = pin_guard.spill(&ctx.0.stats().name()).await;
core::mem::forget(pin_guard);
let old_state = self
.state
.fetch_add(LOCK_BIT.wrapping_sub(RO_PIN_COUNT_UNIT), Ordering::Acquire);
unsafe {
self.spilled_value.get().write(Some(spilled));
}
old_state & RO_PIN_MASK == RO_PIN_COUNT_UNIT
} else {
true
};
let state = if is_exclusive {
let n_bytes = match unsafe { &*self.value_slot.get() } {
ValueSlot::InMemory(val) => val.estimate_byte_size(),
_ => unreachable!(),
};
let (spill_time_ns, spilled_start) = if let Some(strong) = ctx.upgrade() {
strong.stats().add_successful_spill(n_bytes, spill_start)
} else {
(0, Instant::now())
};
unsafe {
self.value_slot.get().replace(ValueSlot::Spilled {
n_bytes,
spill_ctx: ctx,
reinsert_reg_id: registration_id,
spill_time_ns,
spilled_start,
})
};
self.state
.fetch_add(SPILLED_BIT.wrapping_sub(LOCK_BIT), Ordering::AcqRel)
} else {
if let Some(strong) = ctx.upgrade() {
strong.stats().add_failed_spill(spill_start);
}
self.state.fetch_sub(LOCK_BIT, Ordering::AcqRel)
};
self.wake_waiters(state);
is_exclusive
}))
}
}
pub struct SpillToken<T: Spillable> {
inner: Arc<SpillTokenInner<T>>,
}
impl<T: Spillable> SpillToken<T> {
pub fn new(value: T) -> Self {
let inner = Arc::new(SpillTokenInner {
value_slot: UnsafeCell::new(ValueSlot::InMemory(value)),
spilled_value: UnsafeCell::new(None),
registration_id: AtomicU32::new(0),
state: AtomicU64::new(0),
lock: Mutex::default(),
});
Self { inner }
}
pub(crate) fn upcast(&self) -> Arc<dyn DynSpillToken> {
let inner: Arc<SpillTokenInner<T>> = self.inner.clone();
inner
}
pub fn unregister(&mut self) -> Option<(WeakSpillContext, SpillContextParam)> {
self.inner.unregister()
}
pub fn try_get(&self) -> Option<PinnedRef<'_, T>> {
self.inner.try_pin()
}
pub async fn get(&self) -> PinnedRef<'_, T> {
SpillTokenInner::pin(&self.inner).await
}
pub fn get_blocking(&self) -> PinnedRef<'_, T> {
SpillTokenInner::pin_blocking(&self.inner)
}
pub async fn get_mut(&mut self) -> PinnedMut<'_, T> {
SpillTokenInner::pin_mut(&self.inner).await
}
pub fn get_mut_blocking(&mut self) -> PinnedMut<'_, T> {
SpillTokenInner::pin_mut_blocking(&self.inner)
}
pub async fn into_inner(mut self) -> T {
let pin = self.get_mut().await;
let slot = unsafe { pin.inner.value_slot.get().replace(ValueSlot::Dropped) };
let ValueSlot::InMemory(value) = slot else {
unreachable!()
};
value
}
pub fn into_inner_blocking(mut self) -> T {
let pin = self.get_mut_blocking();
let slot = unsafe { pin.inner.value_slot.get().replace(ValueSlot::Dropped) };
let ValueSlot::InMemory(value) = slot else {
unreachable!()
};
value
}
}
impl<T, S> Clone for SpillToken<T>
where
T: Clone + Spillable<Spilled = S>,
S: Clone,
{
fn clone(&self) -> Self {
Self {
inner: Arc::new((*self.inner).clone()),
}
}
}
impl<T: Spillable> Drop for SpillToken<T> {
fn drop(&mut self) {
unsafe { self.inner.mark_as_dropped() };
}
}
pub struct PinnedRef<'a, T: Spillable> {
inner: &'a SpillTokenInner<T>,
}
impl<'a, T: Spillable> Deref for PinnedRef<'a, T> {
type Target = T;
fn deref(&self) -> &Self::Target {
let slot = unsafe { &*self.inner.value_slot.get() };
let ValueSlot::InMemory(value) = slot else {
unreachable!()
};
value
}
}
impl<'a, T: Spillable> Drop for PinnedRef<'a, T> {
fn drop(&mut self) {
unsafe { self.inner.unpin() }
}
}
pub struct PinnedMut<'a, T: Spillable> {
inner: &'a SpillTokenInner<T>,
}
impl<'a, T: Spillable> Deref for PinnedMut<'a, T> {
type Target = T;
fn deref(&self) -> &Self::Target {
let slot = unsafe { &*self.inner.value_slot.get() };
let ValueSlot::InMemory(value) = slot else {
unreachable!()
};
value
}
}
impl<'a, T: Spillable> DerefMut for PinnedMut<'a, T> {
fn deref_mut(&mut self) -> &mut Self::Target {
let slot = unsafe { &mut *self.inner.value_slot.get() };
let ValueSlot::InMemory(value) = slot else {
unreachable!()
};
value
}
}
impl<'a, T: Spillable> Drop for PinnedMut<'a, T> {
fn drop(&mut self) {
unsafe { self.inner.unpin_mut() }
}
}