use std::{
ptr::{null_mut, slice_from_raw_parts_mut},
sync::atomic::{AtomicPtr, AtomicU64, Ordering},
};
use crate::{Result, bucket::HashBucket, error::Error};
pub struct OverflowPool {
chunks: Box<[AtomicPtr<HashBucket>]>,
allocated: AtomicU64,
free_head: AtomicU64,
free_count: AtomicU64,
}
impl OverflowPool {
pub const CHUNK_BITS: usize = 10;
pub const CHUNK_SIZE: usize = 1 << Self::CHUNK_BITS;
pub const CHUNK_MASK: usize = Self::CHUNK_SIZE - 1;
pub const MAX_CHUNKS: usize = 4096;
pub const FREE_LIST_ID_MASK: u64 = 0xFFFF_FFFF;
pub fn new() -> Self {
let chunks = (0..Self::MAX_CHUNKS)
.map(|_| AtomicPtr::new(null_mut()))
.collect::<Box<[_]>>();
Self {
chunks,
allocated: AtomicU64::new(0),
free_head: AtomicU64::new(0),
free_count: AtomicU64::new(0),
}
}
pub fn allocate(&self) -> Result<u64> {
let mut curr = self.free_head.load(Ordering::Acquire);
while (curr & Self::FREE_LIST_ID_MASK) != 0 {
let head_id = curr & Self::FREE_LIST_ID_MASK;
let tag = (curr >> 32) as u32;
if let Some(bucket) = self.get(head_id) {
let next_id = bucket.entries[0].load(Ordering::Acquire) & Self::FREE_LIST_ID_MASK;
let next_val = (((tag.wrapping_add(1)) as u64) << 32) | next_id;
match self.free_head.compare_exchange_weak(
curr,
next_val,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => {
bucket.entries[0].store(0, Ordering::Release);
self.free_count.fetch_sub(1, Ordering::AcqRel);
return Ok(head_id);
}
Err(actual) => curr = actual,
}
} else {
break;
}
}
let prev = self.allocated.fetch_add(1, Ordering::AcqRel);
let id = prev + 1;
let zero_based = prev as usize;
let chunk_idx = zero_based >> Self::CHUNK_BITS;
if chunk_idx >= Self::MAX_CHUNKS {
self.allocated.fetch_sub(1, Ordering::AcqRel);
return Err(Error::OverflowPoolExhausted);
}
if self.chunks[chunk_idx].load(Ordering::Acquire).is_null()
&& let Err(e) = self.ensure_chunk(chunk_idx)
{
self.allocated.fetch_sub(1, Ordering::AcqRel);
return Err(e);
}
Ok(id)
}
pub fn free(&self, id: u64) {
if id == 0 || id > self.allocated.load(Ordering::Acquire) {
return;
}
if let Some(bucket) = self.get(id) {
for entry in &bucket.entries {
entry.store(0, Ordering::Relaxed);
}
let mut curr = self.free_head.load(Ordering::Acquire);
loop {
let tag = (curr >> 32) as u32;
let head_id = curr & Self::FREE_LIST_ID_MASK;
bucket.entries[0].store(head_id, Ordering::Release);
let next_val = (((tag.wrapping_add(1)) as u64) << 32) | (id & Self::FREE_LIST_ID_MASK);
match self.free_head.compare_exchange_weak(
curr,
next_val,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => {
self.free_count.fetch_add(1, Ordering::AcqRel);
break;
}
Err(actual) => curr = actual,
}
}
}
}
#[inline]
pub fn has_free(&self) -> bool {
self.free_count.load(Ordering::Acquire) != 0
}
#[inline]
pub fn free_count(&self) -> u64 {
self.free_count.load(Ordering::Acquire)
}
#[inline]
pub fn get(&self, id: u64) -> Option<&HashBucket> {
if id == 0 || id > self.allocated.load(Ordering::Acquire) {
return None;
}
let zero_based = (id - 1) as usize;
let chunk_idx = zero_based >> Self::CHUNK_BITS;
let slot_idx = zero_based & Self::CHUNK_MASK;
let ptr = self.chunks.get(chunk_idx)?.load(Ordering::Acquire);
if ptr.is_null() {
return None;
}
unsafe { Some(&*ptr.add(slot_idx)) }
}
#[inline]
pub unsafe fn get_unchecked(&self, id: u64) -> &HashBucket {
let zero_based = (id - 1) as usize;
let chunk = unsafe { self.chunks.get_unchecked(zero_based >> Self::CHUNK_BITS) };
debug_assert!(
!chunk.load(Ordering::Relaxed).is_null(),
"合法溢出索引对应的 chunk 必须已初始化"
);
unsafe {
&*chunk
.load(Ordering::Relaxed)
.add(zero_based & Self::CHUNK_MASK)
}
}
#[inline]
pub fn allocated_count(&self) -> u64 {
self.allocated.load(Ordering::Acquire)
}
fn ensure_chunk(&self, chunk_idx: usize) -> Result<()> {
if chunk_idx >= Self::MAX_CHUNKS {
return Err(Error::OverflowPoolExhausted);
}
if !self.chunks[chunk_idx].load(Ordering::Acquire).is_null() {
return Ok(());
}
let chunk = (0..Self::CHUNK_SIZE)
.map(|_| HashBucket::new())
.collect::<Box<[HashBucket]>>();
let raw = Box::into_raw(chunk) as *mut HashBucket;
match self.chunks[chunk_idx].compare_exchange(
null_mut(),
raw,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => Ok(()),
Err(_) => {
unsafe {
let slice_ptr = slice_from_raw_parts_mut(raw, Self::CHUNK_SIZE);
let _ = Box::from_raw(slice_ptr);
}
Ok(())
}
}
}
}
impl Default for OverflowPool {
fn default() -> Self {
Self::new()
}
}
impl Drop for OverflowPool {
fn drop(&mut self) {
for chunk in self.chunks.iter() {
let ptr = chunk.load(Ordering::Relaxed);
if !ptr.is_null() {
unsafe {
let slice_ptr = slice_from_raw_parts_mut(ptr, Self::CHUNK_SIZE);
let _ = Box::from_raw(slice_ptr);
}
}
}
}
}