use crate::atomic::{Ordering, PyAtomic, Radium};
const FLAG_BITS: u32 = 4;
const DESTRUCTED: usize = 1 << (usize::BITS - 1);
const PUBLISHED: usize = 1 << (usize::BITS - 2);
const LEAKED: usize = 1 << (usize::BITS - 3);
const IMMORTAL: usize = 1 << (usize::BITS - 4);
const STRONG_WIDTH: u32 = usize::BITS - FLAG_BITS;
const STRONG: usize = (1 << STRONG_WIDTH) - 1;
const COUNT: usize = 1;
const IMMORTAL_COUNT: usize = if (u32::MAX as u64) < STRONG as u64 {
u32::MAX as usize
} else {
STRONG / 2
};
#[inline(never)]
#[cold]
#[allow(
clippy::disallowed_methods,
reason = "refcount overflow must preserve upstream abort semantics"
)]
fn refcount_overflow() -> ! {
cfg_select! {
feature = "std" => std::process::abort(),
_ => core::panic!("refcount overflow"),
}
}
#[derive(Clone, Copy)]
struct State {
inner: usize,
}
impl State {
#[inline]
fn from_raw(inner: usize) -> Self {
Self { inner }
}
#[inline]
fn as_raw(self) -> usize {
self.inner
}
#[inline]
fn strong(self) -> usize {
(self.inner & STRONG) / COUNT
}
#[inline]
fn destructed(self) -> bool {
(self.inner & DESTRUCTED) != 0
}
#[inline]
fn leaked(self) -> bool {
(self.inner & LEAKED) != 0
}
#[inline]
const fn immortal(self) -> bool {
(self.inner & IMMORTAL) != 0
}
#[inline]
fn add_strong(self, val: u32) -> Self {
Self::from_raw(self.inner + (val as usize) * COUNT)
}
#[inline]
fn with_leaked(self, leaked: bool) -> Self {
Self::from_raw((self.inner & !LEAKED) | if leaked { LEAKED } else { 0 })
}
#[inline]
fn immortalized(self) -> Self {
Self::from_raw((self.inner & !STRONG) | IMMORTAL | (IMMORTAL_COUNT * COUNT))
}
}
pub struct RefCount {
state: PyAtomic<usize>,
}
impl Default for RefCount {
fn default() -> Self {
Self::new()
}
}
impl RefCount {
#[must_use]
pub fn new() -> Self {
Self {
state: Radium::new(COUNT),
}
}
#[inline]
pub fn get(&self) -> usize {
State::from_raw(self.state.load(Ordering::Relaxed)).strong()
}
#[inline(always)]
#[must_use]
pub fn is_immortal(&self) -> bool {
State::from_raw(self.state.load(Ordering::Relaxed)).immortal()
}
#[inline(always)]
pub fn inc(&self) {
if self.is_immortal() {
return;
}
let val = State::from_raw(self.state.fetch_add(COUNT, Ordering::Relaxed));
if (val.as_raw() & (DESTRUCTED | STRONG)).wrapping_sub(COUNT) >= STRONG - COUNT {
self.inc_uncommon(val);
}
}
#[cold]
#[inline(never)]
fn inc_uncommon(&self, val: State) {
if val.destructed() || val.strong() > STRONG - 1 {
refcount_overflow();
}
self.state.fetch_add(COUNT, Ordering::Relaxed);
}
#[inline(always)]
pub fn inc_by(&self, n: usize) {
debug_assert!(n <= STRONG);
if self.is_immortal() {
return;
}
let val = State::from_raw(self.state.fetch_add(n * COUNT, Ordering::Relaxed));
if val.destructed() || val.strong() > STRONG - n {
refcount_overflow();
}
}
#[inline]
#[must_use]
pub fn safe_inc(&self) -> bool {
let mut old = State::from_raw(self.state.load(Ordering::Relaxed));
loop {
if old.immortal() {
return true;
}
if old.destructed() || old.strong() == 0 {
return false;
}
if old.strong() >= STRONG {
refcount_overflow();
}
let new_state = old.add_strong(1);
match self.state.compare_exchange_weak(
old.as_raw(),
new_state.as_raw(),
Ordering::Relaxed,
Ordering::Relaxed,
) {
Ok(_) => return true,
Err(curr) => old = State::from_raw(curr),
}
}
}
#[inline(always)]
#[must_use]
pub fn dec(&self) -> bool {
if self.is_immortal() {
return false;
}
let old = State::from_raw(self.state.fetch_sub(COUNT, Ordering::Release));
debug_assert!(!old.leaked(), "a leaked object must also be immortal");
if old.strong() == 1 {
core::sync::atomic::fence(Ordering::Acquire);
return true;
}
false
}
pub fn leak(&self) {
debug_assert!(!self.is_leaked());
let mut old = State::from_raw(self.state.load(Ordering::Relaxed));
loop {
let new_state = old.with_leaked(true).immortalized();
match self.state.compare_exchange_weak(
old.as_raw(),
new_state.as_raw(),
Ordering::AcqRel,
Ordering::Relaxed,
) {
Ok(_) => return,
Err(curr) => old = State::from_raw(curr),
}
}
}
pub fn make_immortal(&self) {
let mut old = State::from_raw(self.state.load(Ordering::Relaxed));
loop {
if old.immortal() {
return;
}
debug_assert!(!old.destructed() && old.strong() > 0);
match self.state.compare_exchange_weak(
old.as_raw(),
old.immortalized().as_raw(),
Ordering::AcqRel,
Ordering::Relaxed,
) {
Ok(_) => return,
Err(curr) => old = State::from_raw(curr),
}
}
}
pub fn is_leaked(&self) -> bool {
State::from_raw(self.state.load(Ordering::Acquire)).leaked()
}
#[inline]
pub fn mark_published(&self) {
self.state.fetch_or(PUBLISHED, Ordering::Release);
}
#[inline]
pub fn is_published(&self) -> bool {
(self.state.load(Ordering::Acquire) & PUBLISHED) != 0
}
}
#[cfg(feature = "std")]
use core::cell::{Cell, RefCell};
#[cfg(feature = "std")]
thread_local! {
static IN_DEFERRED_CONTEXT: Cell<bool> = const { Cell::new(false) };
static DEFERRED_QUEUE: RefCell<Vec<Box<dyn FnOnce()>>> = const { RefCell::new(Vec::new()) };
}
#[cfg(feature = "std")]
struct DeferredDropGuard {
was_in_context: bool,
}
#[cfg(feature = "std")]
impl Drop for DeferredDropGuard {
fn drop(&mut self) {
IN_DEFERRED_CONTEXT.with(|in_ctx| {
in_ctx.set(self.was_in_context);
});
if !self.was_in_context && !std::thread::panicking() {
flush_deferred_drops();
}
}
}
#[cfg(feature = "std")]
#[inline]
pub fn with_deferred_drops<F, R>(f: F) -> R
where
F: FnOnce() -> R,
{
let _guard = IN_DEFERRED_CONTEXT.with(|in_ctx| {
let was_in_context = in_ctx.get();
in_ctx.set(true);
DeferredDropGuard { was_in_context }
});
f()
}
#[cfg(feature = "std")]
#[inline]
pub fn try_defer_drop<F>(f: F)
where
F: FnOnce() + 'static,
{
let should_defer = IN_DEFERRED_CONTEXT.with(|in_ctx| in_ctx.get());
if should_defer {
DEFERRED_QUEUE.with(|q| {
q.borrow_mut().push(Box::new(f));
});
} else {
f();
}
}
#[cfg(feature = "std")]
#[inline]
pub fn flush_deferred_drops() {
DEFERRED_QUEUE.with(|q| {
let ops: Vec<_> = q.borrow_mut().drain(..).collect();
for op in ops {
op();
}
});
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn strong_count_reaches_past_a_16_bit_ceiling() {
const REFERENCES: usize = 1 << 20;
let rc = RefCount::new();
rc.inc_by(REFERENCES);
assert_eq!(rc.get(), REFERENCES + 1);
}
#[test]
fn inc_and_dec_reach_past_a_16_bit_ceiling() {
const REFERENCES: usize = 1 << 20;
let rc = RefCount::new(); for _ in 1..REFERENCES {
rc.inc();
}
assert_eq!(rc.get(), REFERENCES);
for _ in 1..REFERENCES {
assert!(!rc.dec());
}
assert_eq!(rc.get(), 1);
assert!(rc.dec());
}
#[test]
fn a_new_refcount_holds_one_strong_reference_and_nothing_else() {
let rc = RefCount::new();
assert_eq!(rc.get(), 1);
assert_eq!(rc.state.load(Ordering::Relaxed), COUNT);
}
#[test]
fn an_immortal_count_never_moves() {
let rc = RefCount::new(); assert!(!rc.is_immortal());
rc.make_immortal();
assert!(rc.is_immortal());
assert_eq!(rc.get(), IMMORTAL_COUNT);
rc.inc();
rc.inc_by(1000);
assert!(rc.safe_inc());
assert_eq!(rc.get(), IMMORTAL_COUNT);
for _ in 0..1000 {
assert!(!rc.dec());
}
assert_eq!(rc.get(), IMMORTAL_COUNT);
rc.make_immortal();
assert_eq!(rc.get(), IMMORTAL_COUNT);
}
#[test]
fn interning_implies_immortality_but_not_the_other_way() {
let immortal = RefCount::new();
immortal.make_immortal();
assert!(immortal.is_immortal());
assert!(!immortal.is_leaked());
let interned = RefCount::new();
interned.leak();
assert!(interned.is_leaked());
assert!(interned.is_immortal());
assert_eq!(interned.get(), IMMORTAL_COUNT);
interned.inc();
assert!(!interned.dec());
assert!(!interned.dec());
assert_eq!(interned.get(), IMMORTAL_COUNT);
}
#[test]
fn the_parked_count_outruns_any_real_reference_total() {
const {
assert!(IMMORTAL_COUNT > 1);
assert!(IMMORTAL_COUNT <= STRONG);
}
if usize::BITS >= 64 {
assert_eq!(IMMORTAL_COUNT, u32::MAX as usize);
}
}
#[test]
fn published_bit_survives_refcount_traffic() {
let rc = RefCount::new(); assert!(!rc.is_published());
rc.mark_published();
assert!(rc.is_published());
rc.inc(); assert!(rc.is_published());
assert!(!rc.dec()); assert!(rc.is_published());
assert!(rc.safe_inc()); assert!(!rc.dec()); assert!(rc.dec()); assert!(rc.is_published());
}
}