use alloc::vec::Vec;
use core::fmt;
use core::marker::PhantomData;
#[cfg(feature = "counters")]
use core::cell::Cell;
use crate::error::Error;
use crate::paged::Paged;
use crate::slot::Slot;
pub const PAGE_BYTES: usize = 4096;
const NONE: u32 = u32::MAX;
const FIB: u64 = 0x9E37_79B9_7F4A_7C15;
const IMAGE_HEADER: usize = 24;
const PAGE_IDX_BYTES: usize = core::mem::size_of::<u32>();
const COUNT_BYTES: usize = core::mem::size_of::<u16>();
const PAGE_META_BYTES: usize = PAGE_IDX_BYTES + COUNT_BYTES;
const COUNT: bool = cfg!(feature = "counters");
macro_rules! bump {
($self:ident, $field:ident, $n:expr) => {
#[cfg(feature = "counters")]
{
let mut c = $self.counters.get();
c.$field += $n as u64;
$self.counters.set(c);
}
};
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub enum ShardMode {
Uniform,
Ordered,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct ArenaCfg {
pub shards: usize,
pub max_bytes: usize,
pub mode: ShardMode,
}
impl ArenaCfg {
pub const fn new(shards: usize, mode: ShardMode) -> Self {
Self {
shards,
max_bytes: usize::MAX,
mode,
}
}
pub const fn with_max_bytes(mut self, max_bytes: usize) -> Self {
self.max_bytes = max_bytes;
self
}
}
impl Default for ArenaCfg {
fn default() -> Self {
Self::new(1024, ShardMode::Uniform)
}
}
#[cfg(feature = "counters")]
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct Counters {
pub cmp_ops: u64,
pub bytes_shifted: u64,
pub pages_allocated: u64,
pub chain_steps: u64,
pub splits: u64,
}
pub struct Arena<'a, T: Slot> {
pool: Paged<'a, PAGE_BYTES>,
heads: Vec<u32>,
next: Vec<u32>,
counts: Vec<u16>,
free_head: u32,
total: usize,
scratch: Vec<u8>,
split_buf: Vec<u8>,
cfg: ArenaCfg,
#[cfg(feature = "counters")]
counters: Cell<Counters>,
_marker: PhantomData<fn() -> T>,
}
struct Target {
prev: u32,
page: u32,
}
impl<'a, T: Slot> Arena<'a, T> {
pub fn new(cfg: ArenaCfg) -> Result<Self, Error> {
if T::SIZE == 0 || T::SIZE > PAGE_BYTES || T::KEY_LEN == 0 || T::KEY_LEN > T::SIZE {
return Err(Error::BadSlot {
size: T::SIZE,
key_len: T::KEY_LEN,
});
}
if cfg.shards == 0 || !cfg.shards.is_power_of_two() {
return Err(Error::BadShardCount { got: cfg.shards });
}
Ok(Self {
pool: Paged::owned_empty(),
heads: alloc::vec![NONE; cfg.shards],
next: Vec::new(),
counts: Vec::new(),
free_head: NONE,
total: 0,
scratch: Vec::new(),
split_buf: Vec::new(),
cfg,
#[cfg(feature = "counters")]
counters: Cell::new(Counters::default()),
_marker: PhantomData,
})
}
pub const fn slots_per_page() -> usize {
PAGE_BYTES / T::SIZE
}
pub fn len(&self) -> usize {
self.total
}
pub fn is_empty(&self) -> bool {
self.total == 0
}
pub fn pool_bytes(&self) -> usize {
self.pool.len()
}
pub fn cfg(&self) -> &ArenaCfg {
&self.cfg
}
pub fn insert(&mut self, value: &T) -> Result<bool, Error> {
let mut buf = core::mem::take(&mut self.scratch);
buf.resize(T::SIZE, 0);
value.write(&mut buf);
let result = self.insert_bytes(&buf);
self.scratch = buf;
result
}
fn insert_bytes(&mut self, slot: &[u8]) -> Result<bool, Error> {
let key = &slot[..T::KEY_LEN];
let shard = self.shard_of(key);
let target = match self.find_page(shard, key) {
Some(t) => t,
None => {
let page = self.alloc_page()?;
self.heads[shard] = page;
Target { prev: NONE, page }
}
};
let mut page = target.page;
let mut count = self.counts[page as usize] as usize;
let mut pos = {
let mut cmps = 0u64;
let found = self.search_in(page, count, key, &mut cmps);
bump!(self, cmp_ops, cmps);
match found {
Ok(_) => return Ok(false),
Err(pos) => pos,
}
};
if count == Self::slots_per_page() {
(page, pos, count) = self.split(page, pos)?;
}
let slot_start = pos * T::SIZE;
let used_end = count * T::SIZE;
let shifted = used_end - slot_start;
let bytes = self.pool.page_mut(page);
if shifted > 0 {
bytes.copy_within(slot_start..used_end, slot_start + T::SIZE);
}
bytes[slot_start..slot_start + T::SIZE].copy_from_slice(slot);
bump!(self, bytes_shifted, shifted);
self.counts[page as usize] += 1;
self.total += 1;
Ok(true)
}
fn split(&mut self, page: u32, pos: usize) -> Result<(u32, usize, usize), Error> {
let spp = Self::slots_per_page();
let fresh = self.alloc_page()?;
bump!(self, splits, 1);
self.next[fresh as usize] = self.next[page as usize];
self.next[page as usize] = fresh;
if spp == 1 {
return Ok(if pos == 0 {
self.move_slots(page, 0, fresh, T::SIZE);
bump!(self, bytes_shifted, T::SIZE);
self.counts[page as usize] = 0;
self.counts[fresh as usize] = 1;
(page, 0, 0)
} else {
(fresh, 0, 0)
});
}
let half = spp / 2;
let moved = spp - half;
self.move_slots(page, half * T::SIZE, fresh, moved * T::SIZE);
bump!(self, bytes_shifted, moved * T::SIZE);
self.counts[page as usize] = half as u16;
self.counts[fresh as usize] = moved as u16;
Ok(if pos <= half {
(page, pos, half)
} else {
(fresh, pos - half, moved)
})
}
fn move_slots(&mut self, src_page: u32, src_off: usize, dst_page: u32, len: usize) {
let mut buf = core::mem::take(&mut self.split_buf);
buf.clear();
buf.extend_from_slice(&self.pool.page(src_page)[src_off..src_off + len]);
self.pool.page_mut(dst_page)[..len].copy_from_slice(&buf);
self.split_buf = buf;
}
pub fn get(&self, key: &[u8]) -> Option<T> {
self.locate(key).map(|off| {
let (page, rel) = (off / PAGE_BYTES, off % PAGE_BYTES);
T::read(&self.pool.page(page as u32)[rel..rel + T::SIZE])
})
}
pub fn contains(&self, key: &[u8]) -> bool {
self.locate(key).is_some()
}
pub fn get_slot(&self, key: &[u8]) -> Option<&[u8]> {
self.locate(key).map(|off| {
let (page, rel) = (off / PAGE_BYTES, off % PAGE_BYTES);
&self.pool.page(page as u32)[rel..rel + T::SIZE]
})
}
pub fn payload_mut(&mut self, key: &[u8]) -> Option<&mut [u8]> {
let off = self.locate(key)?;
let (page, rel) = (off / PAGE_BYTES, off % PAGE_BYTES);
Some(&mut self.pool.page_mut(page as u32)[rel + T::KEY_LEN..rel + T::SIZE])
}
pub fn remove(&mut self, key: &[u8]) -> bool {
assert_eq!(key.len(), T::KEY_LEN, "key length must equal Slot::KEY_LEN");
let shard = self.shard_of(key);
let Some(Target { prev, page }) = self.find_page(shard, key) else {
return false;
};
let count = self.counts[page as usize] as usize;
let pos = {
let mut cmps = 0u64;
let found = self.search_in(page, count, key, &mut cmps);
bump!(self, cmp_ops, cmps);
match found {
Ok(pos) => pos,
Err(_) => return false,
}
};
let slot_start = pos * T::SIZE;
let used_end = count * T::SIZE;
let tail = used_end - (slot_start + T::SIZE);
if tail > 0 {
self.pool
.page_mut(page)
.copy_within(slot_start + T::SIZE..used_end, slot_start);
}
bump!(self, bytes_shifted, tail);
self.counts[page as usize] -= 1;
self.total -= 1;
if self.counts[page as usize] == 0 {
let successor = self.next[page as usize];
if prev == NONE {
self.heads[shard] = successor;
} else {
self.next[prev as usize] = successor;
}
self.next[page as usize] = self.free_head;
self.free_head = page;
}
true
}
pub fn iter(&self) -> Iter<'_, T> {
Iter {
arena: self,
shard: 0,
page: NONE,
idx: 0,
remaining: self.total,
}
}
pub fn range<'s>(&'s self, from: &[u8], to: &'s [u8]) -> Range<'s, T> {
assert_eq!(
self.cfg.mode,
ShardMode::Ordered,
"range scans require ShardMode::Ordered"
);
assert_eq!(
from.len(),
T::KEY_LEN,
"key length must equal Slot::KEY_LEN"
);
assert_eq!(to.len(), T::KEY_LEN, "key length must equal Slot::KEY_LEN");
let shard = self.shard_of(from);
let mut page = NONE;
let mut idx = 0usize;
if let Some(t) = self.find_page(shard, from) {
let count = self.counts[t.page as usize] as usize;
let mut cmps = 0u64;
let pos = match self.search_in(t.page, count, from, &mut cmps) {
Ok(p) | Err(p) => p,
};
bump!(self, cmp_ops, cmps);
if pos < count {
page = t.page;
idx = pos;
} else {
page = self.next[t.page as usize];
}
}
Range {
arena: self,
shard: shard + 1,
page,
idx,
to,
}
}
#[cfg(feature = "counters")]
pub fn counters(&self) -> Counters {
self.counters.get()
}
#[cfg(feature = "counters")]
pub fn reset_counters(&self) {
self.counters.set(Counters::default());
}
pub fn dump_meta(&self, out: &mut Vec<u8>) {
let pages = self.pool.len() / PAGE_BYTES;
debug_assert!(self.cfg.shards <= u32::MAX as usize && pages < u32::MAX as usize);
out.reserve(IMAGE_HEADER + self.heads.len() * PAGE_IDX_BYTES + pages * PAGE_META_BYTES);
out.extend_from_slice(&(self.cfg.shards as u32).to_le_bytes());
out.extend_from_slice(&(pages as u32).to_le_bytes());
out.extend_from_slice(&self.free_head.to_le_bytes());
out.extend_from_slice(&(self.total as u64).to_le_bytes());
out.push(match self.cfg.mode {
ShardMode::Uniform => 0,
ShardMode::Ordered => 1,
});
out.extend_from_slice(&[0u8; 3]);
for &head in &self.heads {
out.extend_from_slice(&head.to_le_bytes());
}
for &next in &self.next {
out.extend_from_slice(&next.to_le_bytes());
}
for &count in &self.counts {
out.extend_from_slice(&count.to_le_bytes());
}
}
pub fn dump_pool(&self, out: &mut Vec<u8>) {
out.reserve(self.pool.len());
for (page, &count) in self.counts.iter().enumerate() {
let used = count as usize * T::SIZE;
out.extend_from_slice(&self.pool.page(page as u32)[..used]);
out.resize(out.len() + (PAGE_BYTES - used), 0);
}
}
pub fn load(cfg: ArenaCfg, meta: &[u8], pool: &[u8]) -> Result<Self, Error> {
Self::load_impl(cfg, meta, pool, Paged::owned_from(pool.to_vec()))
}
pub fn load_borrowed(cfg: ArenaCfg, meta: &[u8], pool: &'a [u8]) -> Result<Self, Error> {
Self::load_impl(cfg, meta, pool, Paged::borrowed(pool))
}
pub fn load_overlay(cfg: ArenaCfg, meta: &[u8], pool: &'a [u8]) -> Result<Self, Error> {
Self::load_impl(cfg, meta, pool, Paged::borrowed(pool))
}
fn load_impl(
cfg: ArenaCfg,
meta: &[u8],
pool: &[u8],
backing: Paged<'a, PAGE_BYTES>,
) -> Result<Self, Error> {
let mut arena = Self::new(cfg)?;
if meta.len() < IMAGE_HEADER {
return Err(Error::Corrupt("arena meta shorter than its header"));
}
let shards = u32::from_le_bytes(meta[0..4].try_into().unwrap()) as usize;
let pages = u32::from_le_bytes(meta[4..8].try_into().unwrap()) as usize;
let free_head = u32::from_le_bytes(meta[8..12].try_into().unwrap());
let total = u64::from_le_bytes(meta[12..20].try_into().unwrap());
let mode = meta[20];
if meta[21..24] != [0u8; 3] {
return Err(Error::Corrupt("arena meta reserved bytes must be zero"));
}
if shards != cfg.shards {
return Err(Error::Corrupt("arena meta shard count disagrees with cfg"));
}
let want_mode = match cfg.mode {
ShardMode::Uniform => 0u8,
ShardMode::Ordered => 1,
};
if mode != want_mode {
return Err(Error::Corrupt("arena meta shard mode disagrees with cfg"));
}
if pages as u64 >= u64::from(NONE) {
return Err(Error::Corrupt("arena page count overflows the index space"));
}
let want_meta = IMAGE_HEADER as u64
+ shards as u64 * PAGE_IDX_BYTES as u64
+ pages as u64 * PAGE_META_BYTES as u64;
if meta.len() as u64 != want_meta {
return Err(Error::Corrupt("arena meta length mismatch"));
}
if pool.len() as u64 != pages as u64 * PAGE_BYTES as u64 {
return Err(Error::Corrupt("arena pool length mismatch"));
}
let mut heads = Vec::with_capacity(shards);
for i in 0..shards {
let at = IMAGE_HEADER + i * PAGE_IDX_BYTES;
heads.push(u32::from_le_bytes(
meta[at..at + PAGE_IDX_BYTES].try_into().unwrap(),
));
}
let next_base = IMAGE_HEADER + shards * PAGE_IDX_BYTES;
let mut next = Vec::with_capacity(pages);
for i in 0..pages {
let at = next_base + i * PAGE_IDX_BYTES;
next.push(u32::from_le_bytes(
meta[at..at + PAGE_IDX_BYTES].try_into().unwrap(),
));
}
let counts_base = next_base + pages * PAGE_IDX_BYTES;
let mut counts = Vec::with_capacity(pages);
for i in 0..pages {
let at = counts_base + i * COUNT_BYTES;
counts.push(u16::from_le_bytes(
meta[at..at + COUNT_BYTES].try_into().unwrap(),
));
}
let spp = Self::slots_per_page();
if counts.iter().any(|&c| c as usize > spp) {
return Err(Error::Corrupt("arena page count exceeds slots per page"));
}
let mut seen = alloc::vec![false; pages];
let mut live = 0u64;
for (shard, &head) in heads.iter().enumerate() {
let mut page = head;
let mut prev: Option<u32> = None;
while page != NONE {
let p = page as usize;
if p >= pages {
return Err(Error::Corrupt("arena chain page out of bounds"));
}
if core::mem::replace(&mut seen[p], true) {
return Err(Error::Corrupt("arena page linked more than once"));
}
let count = counts[p] as usize;
if count == 0 {
return Err(Error::Corrupt("arena chain contains an empty page"));
}
live += count as u64;
let first = &pool[p * PAGE_BYTES..p * PAGE_BYTES + T::KEY_LEN];
if arena.shard_of(first) != shard {
return Err(Error::Corrupt("arena page sits in the wrong shard"));
}
if let Some(pr) = prev {
let last_at =
pr as usize * PAGE_BYTES + (counts[pr as usize] as usize - 1) * T::SIZE;
if pool[last_at..last_at + T::KEY_LEN] >= *first {
return Err(Error::Corrupt("arena chain pages out of key order"));
}
}
prev = Some(page);
page = next[p];
}
}
let mut page = free_head;
while page != NONE {
let p = page as usize;
if p >= pages {
return Err(Error::Corrupt("arena free page out of bounds"));
}
if core::mem::replace(&mut seen[p], true) {
return Err(Error::Corrupt("arena page linked more than once"));
}
if counts[p] != 0 {
return Err(Error::Corrupt("arena free page has a nonzero count"));
}
page = next[p];
}
if seen.iter().any(|&s| !s) {
return Err(Error::Corrupt("arena has an orphan page"));
}
if live != total {
return Err(Error::Corrupt(
"arena record total disagrees with page counts",
));
}
arena.pool = backing;
arena.heads = heads;
arena.next = next;
arena.counts = counts;
arena.free_head = free_head;
arena.total = total as usize;
Ok(arena)
}
fn find_page(&self, shard: usize, key: &[u8]) -> Option<Target> {
let head = self.heads[shard];
if head == NONE {
return None;
}
let mut prev = NONE;
let mut page = head;
let mut steps = 0u64;
loop {
let nxt = self.next[page as usize];
if nxt == NONE || self.first_key(nxt) > key {
break;
}
steps += COUNT as u64;
prev = page;
page = nxt;
}
bump!(self, chain_steps, steps);
let _ = steps; Some(Target { prev, page })
}
fn first_key(&self, page: u32) -> &[u8] {
&self.pool.page(page)[..T::KEY_LEN]
}
fn search_in(
&self,
page: u32,
count: usize,
key: &[u8],
cmps: &mut u64,
) -> Result<usize, usize> {
search::<T>(self.pool.page(page), count, key, cmps)
}
fn locate(&self, key: &[u8]) -> Option<usize> {
assert_eq!(key.len(), T::KEY_LEN, "key length must equal Slot::KEY_LEN");
let shard = self.shard_of(key);
let Target { page, .. } = self.find_page(shard, key)?;
let count = self.counts[page as usize] as usize;
let mut cmps = 0u64;
let found = self.search_in(page, count, key, &mut cmps).ok();
bump!(self, cmp_ops, cmps);
found.map(|pos| page as usize * PAGE_BYTES + pos * T::SIZE)
}
fn shard_of(&self, key: &[u8]) -> usize {
let mut pad = [0u8; 8];
let n = key.len().min(8);
pad[..n].copy_from_slice(&key[..n]);
let v = u64::from_be_bytes(pad);
let bits = self.cfg.shards.trailing_zeros();
if bits == 0 {
return 0; }
let h = match self.cfg.mode {
ShardMode::Ordered => v,
ShardMode::Uniform => v.wrapping_mul(FIB),
};
(h >> (64 - bits)) as usize
}
fn alloc_page(&mut self) -> Result<u32, Error> {
if self.free_head != NONE {
let page = self.free_head;
self.free_head = self.next[page as usize];
self.next[page as usize] = NONE;
self.counts[page as usize] = 0;
bump!(self, pages_allocated, 1);
return Ok(page);
}
let old_len = self.pool.len();
let new_len = old_len + PAGE_BYTES;
if new_len > self.cfg.max_bytes {
return Err(Error::CapacityExceeded {
max_bytes: self.cfg.max_bytes,
});
}
let page = (old_len / PAGE_BYTES) as u32;
let tail = self.pool.grown_tail_mut();
let tail_len = tail.len() + PAGE_BYTES;
tail.reserve(PAGE_BYTES);
#[allow(clippy::uninit_vec)]
unsafe {
tail.set_len(tail_len);
}
self.next.push(NONE);
self.counts.push(0);
bump!(self, pages_allocated, 1);
Ok(page)
}
}
fn search<T: Slot>(page: &[u8], count: usize, key: &[u8], cmps: &mut u64) -> Result<usize, usize> {
let mut lo = 0usize;
let mut hi = count;
while lo < hi {
let mid = lo + (hi - lo) / 2;
let off = mid * T::SIZE;
*cmps += COUNT as u64;
match page[off..off + T::KEY_LEN].cmp(key) {
core::cmp::Ordering::Less => lo = mid + 1,
core::cmp::Ordering::Greater => hi = mid,
core::cmp::Ordering::Equal => return Ok(mid),
}
}
Err(lo)
}
impl<T: Slot> fmt::Debug for Arena<'_, T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Arena")
.field("len", &self.total)
.field("shards", &self.cfg.shards)
.field("mode", &self.cfg.mode)
.field("pages", &(self.pool.len() / PAGE_BYTES))
.field("slot_size", &T::SIZE)
.finish()
}
}
pub struct Iter<'a, T: Slot> {
arena: &'a Arena<'a, T>,
shard: usize,
page: u32,
idx: usize,
remaining: usize,
}
impl<T: Slot> Iterator for Iter<'_, T> {
type Item = T;
fn next(&mut self) -> Option<T> {
loop {
if self.page == NONE {
if self.shard >= self.arena.cfg.shards {
return None;
}
self.page = self.arena.heads[self.shard];
self.idx = 0;
self.shard += 1;
continue;
}
if self.idx < self.arena.counts[self.page as usize] as usize {
let rel = self.idx * T::SIZE;
self.idx += 1;
self.remaining -= 1;
return Some(T::read(
&self.arena.pool.page(self.page)[rel..rel + T::SIZE],
));
}
self.page = self.arena.next[self.page as usize];
self.idx = 0;
}
}
fn size_hint(&self) -> (usize, Option<usize>) {
(self.remaining, Some(self.remaining))
}
}
impl<T: Slot> ExactSizeIterator for Iter<'_, T> {}
impl<'a, T: Slot> IntoIterator for &'a Arena<'a, T> {
type Item = T;
type IntoIter = Iter<'a, T>;
fn into_iter(self) -> Self::IntoIter {
self.iter()
}
}
pub struct Range<'a, T: Slot> {
arena: &'a Arena<'a, T>,
shard: usize,
page: u32,
idx: usize,
to: &'a [u8],
}
impl<T: Slot> Iterator for Range<'_, T> {
type Item = T;
fn next(&mut self) -> Option<T> {
loop {
if self.page == NONE {
if self.shard >= self.arena.cfg.shards {
return None;
}
self.page = self.arena.heads[self.shard];
self.idx = 0;
self.shard += 1;
continue;
}
if self.idx < self.arena.counts[self.page as usize] as usize {
let rel = self.idx * T::SIZE;
let page = self.arena.pool.page(self.page);
if page[rel..rel + T::KEY_LEN] >= *self.to {
return None;
}
self.idx += 1;
return Some(T::read(&page[rel..rel + T::SIZE]));
}
self.page = self.arena.next[self.page as usize];
self.idx = 0;
}
}
}