use yo_common::Addr;
const MIN_SLOTS: usize = 64;
const LOAD_NUM: usize = 3;
const LOAD_DEN: usize = 4;
const MIX: u64 = 0x9e37_79b9_7f4a_7c15;
#[derive(Debug)]
pub(crate) struct Tagged {
slots: Vec<Addr>,
len: usize,
}
impl Tagged {
pub(crate) fn new() -> Tagged {
Tagged {
slots: Vec::new(),
len: 0,
}
}
#[inline]
pub(crate) const fn len(&self) -> usize {
self.len
}
#[inline]
pub(crate) fn memory_bytes(&self) -> usize {
self.slots.capacity() * size_of::<Addr>()
}
#[inline]
fn home(&self, addr: Addr) -> usize {
let mask = self.slots.len() - 1;
((addr.to_bits().wrapping_mul(MIX) >> 32) as usize) & mask
}
#[inline]
fn find(&self, addr: Addr) -> Option<usize> {
if self.len == 0 {
return None;
}
let mask = self.slots.len() - 1;
let mut at = self.home(addr);
loop {
let s = self.slots[at];
if s == addr {
return Some(at);
}
if s == Addr::NONE {
return None;
}
at = (at + 1) & mask;
}
}
#[inline]
pub(crate) fn contains(&self, addr: Addr) -> bool {
self.find(addr).is_some()
}
pub(crate) fn insert(&mut self, addr: Addr) -> bool {
debug_assert_ne!(addr, Addr::NONE, "the empty slot is not an address");
if (self.len + 1) * LOAD_DEN > self.slots.len() * LOAD_NUM {
self.resize(self.slots.len().max(MIN_SLOTS / 2) * 2);
}
let mask = self.slots.len() - 1;
let mut at = self.home(addr);
loop {
let s = self.slots[at];
if s == addr {
return false;
}
if s == Addr::NONE {
self.slots[at] = addr;
self.len += 1;
return true;
}
at = (at + 1) & mask;
}
}
pub(crate) fn remove(&mut self, addr: Addr) -> bool {
let Some(at) = self.find(addr) else {
return false;
};
self.len -= 1;
self.shift_back(at);
if self.slots.len() > MIN_SLOTS && self.len * LOAD_DEN < self.slots.len() {
self.resize(self.slots.len() / 2);
}
true
}
fn shift_back(&mut self, at: usize) {
let mask = self.slots.len() - 1;
let mut hole = at;
loop {
self.slots[hole] = Addr::NONE;
let mut j = hole;
loop {
j = (j + 1) & mask;
let s = self.slots[j];
if s == Addr::NONE {
return;
}
if !between(hole, self.home(s), j) {
break;
}
}
self.slots[hole] = self.slots[j];
hole = j;
}
}
fn resize(&mut self, slots: usize) {
debug_assert!(slots.is_power_of_two());
debug_assert!(self.len * LOAD_DEN <= slots * LOAD_NUM);
let fresh = yo_alloc::for_the_data(|| vec![Addr::NONE; slots]);
let old = core::mem::replace(&mut self.slots, fresh);
let mask = slots - 1;
for a in old {
if a == Addr::NONE {
continue;
}
let mut at = self.home(a);
while self.slots[at] != Addr::NONE {
at = (at + 1) & mask;
}
self.slots[at] = a;
}
}
pub(crate) fn sample(&self, r: u64, mut out: impl FnMut(Addr) -> bool) {
if self.len == 0 {
return;
}
let mask = self.slots.len() - 1;
let mut at = (r as usize) & mask;
for _ in 0..self.slots.len() {
let s = self.slots[at];
if s != Addr::NONE && !out(s) {
return;
}
at = (at + 1) & mask;
}
}
}
#[inline]
const fn between(lo: usize, x: usize, hi: usize) -> bool {
if lo <= hi {
lo < x && x <= hi
} else {
x > lo || x <= hi
}
}
#[cfg(test)]
mod tests {
use super::*;
use yo_common::Space;
fn a(n: u64) -> Addr {
Addr::new(Space::Arena, n * 64)
}
#[test]
fn an_empty_set_holds_nothing_and_costs_nothing() {
let t = Tagged::new();
assert_eq!(t.len(), 0);
assert_eq!(t.memory_bytes(), 0);
assert!(!t.contains(a(1)));
let mut seen = 0;
t.sample(0, |_| {
seen += 1;
true
});
assert_eq!(seen, 0);
}
#[test]
fn what_goes_in_comes_back_out() {
let mut t = Tagged::new();
for i in 1..1_000 {
assert!(t.insert(a(i)), "{i} was not there");
}
assert_eq!(t.len(), 999);
for i in 1..1_000 {
assert!(t.contains(a(i)), "{i} went missing");
}
assert!(!t.contains(a(1_000)));
}
#[test]
fn tagging_twice_is_tagging_once() {
let mut t = Tagged::new();
assert!(t.insert(a(7)));
assert!(!t.insert(a(7)));
assert_eq!(t.len(), 1);
assert!(t.remove(a(7)));
assert!(!t.remove(a(7)));
assert_eq!(t.len(), 0);
}
#[test]
fn removing_the_middle_leaves_the_rest_findable() {
let mut t = Tagged::new();
for i in 1..2_000 {
t.insert(a(i));
}
for i in (1..2_000).step_by(3) {
assert!(t.remove(a(i)), "{i} should have been there");
}
for i in 1..2_000 {
let want = i % 3 != 1;
assert_eq!(t.contains(a(i)), want, "{i}");
}
assert_eq!(t.len(), 1_999 - (1..2_000).step_by(3).count());
}
#[test]
fn churn_does_not_leave_anything_behind() {
let mut t = Tagged::new();
for round in 0..50u64 {
for i in 1..200 {
t.insert(a(round * 1_000 + i));
}
for i in 1..200 {
assert!(t.remove(a(round * 1_000 + i)));
}
assert_eq!(t.len(), 0, "round {round}");
}
assert!(t.memory_bytes() <= MIN_SLOTS * size_of::<Addr>());
}
#[test]
fn a_sample_offers_every_tagged_address_and_no_others() {
let mut t = Tagged::new();
for i in 1..500 {
t.insert(a(i));
}
let mut seen = std::collections::HashSet::new();
t.sample(12_345, |x| {
assert!(seen.insert(x), "{x:?} came round twice");
true
});
assert_eq!(seen.len(), 499);
for i in 1..500 {
assert!(seen.contains(&a(i)));
}
}
#[test]
fn a_sample_stops_when_it_is_told_to() {
let mut t = Tagged::new();
for i in 1..500 {
t.insert(a(i));
}
let mut seen = 0;
t.sample(99, |_| {
seen += 1;
seen < 20
});
assert_eq!(seen, 20);
}
#[test]
fn a_sample_starts_where_it_is_told_to() {
let mut t = Tagged::new();
for i in 1..500 {
t.insert(a(i));
}
let first = |r| {
let mut got = None;
t.sample(r, |x| {
got = Some(x);
false
});
got.expect("something is tagged")
};
let mut starts = std::collections::HashSet::new();
for r in 0..64u64 {
starts.insert(first(r * 7));
}
assert!(starts.len() > 8, "{} distinct starts", starts.len());
}
#[test]
fn the_stretch_test_wraps() {
assert!(between(2, 3, 5));
assert!(between(2, 5, 5));
assert!(!between(2, 2, 5));
assert!(!between(2, 6, 5));
assert!(between(6, 7, 2));
assert!(between(6, 0, 2));
assert!(between(6, 2, 2));
assert!(!between(6, 6, 2));
assert!(!between(6, 3, 2));
}
}