use core::fmt;
use core::sync::atomic::{AtomicU32, AtomicU64, Ordering};
use std::sync::{Mutex, MutexGuard, OnceLock, PoisonError};
use hashbrown::HashTable;
use hashbrown::hash_table::Entry;
use super::name::{InputPosition, SymbolName};
use crate::ids::FileId;
const CLAIM_SHARD_BITS: u32 = 8;
struct Claim<'a> {
key: SymbolName<'a>,
position: InputPosition,
owner: FileId,
round: u32,
}
#[inline]
fn lock<T>(mutex: &Mutex<T>) -> MutexGuard<'_, T> {
mutex.lock().unwrap_or_else(PoisonError::into_inner)
}
#[inline]
fn shard_of(hash: u64) -> usize {
(hash >> (64 - CLAIM_SHARD_BITS)) as usize
}
#[inline]
fn table_hash(hash: u64) -> u64 {
hash.rotate_left(CLAIM_SHARD_BITS)
}
pub struct GroupClaims<'a> {
shards: Box<[Mutex<HashTable<Claim<'a>>>]>,
round: u32,
}
impl Default for GroupClaims<'_> {
fn default() -> Self {
Self::new()
}
}
impl fmt::Debug for GroupClaims<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("GroupClaims")
.field("round", &self.round)
.field("len", &self.len())
.finish_non_exhaustive()
}
}
impl<'a> GroupClaims<'a> {
#[must_use]
pub fn new() -> Self {
Self {
shards: (0..1usize << CLAIM_SHARD_BITS)
.map(|_| Mutex::new(HashTable::new()))
.collect(),
round: 0,
}
}
pub fn begin_round(&mut self) -> ClaimRound<'_, 'a> {
self.round = self.round.saturating_add(1);
ClaimRound {
claims: self,
round: self.round,
}
}
#[must_use]
pub fn in_round(&self, round: u32) -> ClaimRound<'_, 'a> {
ClaimRound {
claims: self,
round,
}
}
#[must_use]
pub fn owner(&self, key: &SymbolName<'_>) -> Option<FileId> {
let shard = lock(&self.shards[shard_of(key.hash())]);
shard
.find(table_hash(key.hash()), |claim| claim.key == *key)
.map(|claim| claim.owner)
}
#[must_use]
pub fn len(&self) -> usize {
self.shards.iter().map(|shard| lock(shard).len()).sum()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.len() == 0
}
}
pub struct ClaimRound<'c, 'a> {
claims: &'c GroupClaims<'a>,
round: u32,
}
impl fmt::Debug for ClaimRound<'_, '_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ClaimRound")
.field("round", &self.round)
.finish_non_exhaustive()
}
}
impl<'a> ClaimRound<'_, 'a> {
pub fn offer(&self, key: SymbolName<'a>, position: InputPosition, file: FileId) -> bool {
let mut shard = lock(&self.claims.shards[shard_of(key.hash())]);
let entry = shard.entry(
table_hash(key.hash()),
|claim| claim.key == key,
|claim| table_hash(claim.key.hash()),
);
match entry {
Entry::Vacant(vacant) => {
vacant.insert(Claim {
key,
position,
owner: file,
round: self.round,
});
true
}
Entry::Occupied(mut occupied) => {
let claim = occupied.get_mut();
if claim.round == self.round && (position, file) < (claim.position, claim.owner) {
claim.position = position;
claim.owner = file;
}
claim.owner == file
}
}
}
#[must_use]
pub fn owner(&self, key: &SymbolName<'_>) -> Option<FileId> {
self.claims.owner(key)
}
#[must_use]
pub fn is_owner(&self, key: &SymbolName<'_>, file: FileId) -> bool {
self.owner(key) == Some(file)
}
}
const SLOT_BASE: usize = 1024;
const SLOT_SEGMENTS: usize = 23;
type KeyShard<'a> = Mutex<HashTable<(SymbolName<'a>, u32)>>;
pub struct GroupSlots<'a> {
keys: Box<[KeyShard<'a>]>,
next: AtomicU32,
claims: [OnceLock<Box<[AtomicU64]>>; SLOT_SEGMENTS],
}
impl Default for GroupSlots<'_> {
fn default() -> Self {
Self::new()
}
}
impl fmt::Debug for GroupSlots<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("GroupSlots")
.field("slots", &self.next.load(Ordering::Relaxed))
.finish_non_exhaustive()
}
}
impl<'a> GroupSlots<'a> {
#[must_use]
pub fn new() -> Self {
Self {
keys: (0..1usize << CLAIM_SHARD_BITS)
.map(|_| Mutex::new(HashTable::new()))
.collect(),
next: AtomicU32::new(0),
claims: [const { OnceLock::new() }; SLOT_SEGMENTS],
}
}
pub fn slot(&self, key: SymbolName<'a>) -> Option<u32> {
let mut shard = lock(&self.keys[shard_of(key.hash())]);
let entry = shard.entry(
table_hash(key.hash()),
|(claimed, _)| *claimed == key,
|(claimed, _)| table_hash(claimed.hash()),
);
match entry {
Entry::Occupied(occupied) => Some(occupied.get().1),
Entry::Vacant(vacant) => {
let slot = self
.next
.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |n| n.checked_add(1))
.ok()?;
vacant.insert((key, slot));
Some(slot)
}
}
}
fn cell(&self, slot: u32) -> &AtomicU64 {
debug_assert!(
slot < self.next.load(Ordering::Relaxed),
"slot {slot} never given out"
);
let scaled = slot as usize / SLOT_BASE + 1;
let segment = (usize::BITS - 1 - scaled.leading_zeros()) as usize;
let start = SLOT_BASE * ((1usize << segment) - 1);
let cells = self.claims[segment].get_or_init(|| {
(0..SLOT_BASE << segment)
.map(|_| AtomicU64::new(u64::MAX))
.collect()
});
&cells[slot as usize - start]
}
pub fn offer(&self, slot: u32, round: u32, rank: u32) {
self.cell(slot)
.fetch_min(u64::from(round) << 32 | u64::from(rank), Ordering::Relaxed);
}
#[must_use]
pub fn holds(&self, slot: u32, round: u32, rank: u32) -> bool {
self.cell(slot).load(Ordering::Relaxed) == u64::from(round) << 32 | u64::from(rank)
}
}
#[cfg(test)]
mod tests {
use super::*;
use rayon::prelude::*;
#[test]
fn slots_follow_the_same_rules_as_claims() {
let keys: Vec<String> = (0..1500).map(|k| format!("group{k}")).collect();
let mut claims = GroupClaims::new();
let slots = GroupSlots::new();
let offers: Vec<(u32, u32, usize)> = (0..3000u32)
.map(|i| (i % 3 + 1, (i * 7919) % 3000, (i as usize * 31) % 1500))
.collect();
for round in 1..=3 {
let this: Vec<&(u32, u32, usize)> =
offers.iter().filter(|(r, _, _)| *r == round).collect();
{
let claim = claims.begin_round();
this.par_iter().for_each(|&&(_, rank, k)| {
claim.offer(
SymbolName::new(keys[k].as_bytes()),
InputPosition::new(rank, 0),
FileId::new(rank as usize),
);
});
}
this.par_iter().for_each(|&&(_, rank, k)| {
let slot = slots.slot(SymbolName::new(keys[k].as_bytes())).unwrap();
slots.offer(slot, round, rank);
});
for &&(_, rank, k) in &this {
let key = SymbolName::new(keys[k].as_bytes());
let slot = slots.slot(key).unwrap();
assert_eq!(
slots.holds(slot, round, rank),
claims.owner(&key) == Some(FileId::new(rank as usize)),
"{} rank {rank} round {round}",
keys[k]
);
}
}
let names: Vec<String> = (0..5000).map(|k| format!("k{k}")).collect();
let many = GroupSlots::new();
for (i, name) in names.iter().enumerate() {
let slot = many.slot(SymbolName::new(name.as_bytes())).unwrap();
assert_eq!(slot as usize, i);
many.offer(slot, 1, 10);
many.offer(slot, 1, 5);
assert!(many.holds(slot, 1, 5) && !many.holds(slot, 1, 10));
}
}
fn key(name: &'static str) -> SymbolName<'static> {
SymbolName::new(name.as_bytes())
}
#[test]
fn lowest_offer_in_a_round_wins_regardless_of_order() {
let offers: Vec<(u32, usize)> = (0..200).map(|i| ((i * 37) % 200, i as usize)).collect();
for threads in [1, 4] {
let pool = rayon::ThreadPoolBuilder::new()
.num_threads(threads)
.build()
.unwrap();
let mut claims = GroupClaims::new();
pool.install(|| {
let round = claims.begin_round();
offers.par_iter().for_each(|&(position, file)| {
round.offer(
key(if file % 2 == 0 { "even" } else { "odd" }),
InputPosition::new(position + 1, 0),
FileId::new(file),
);
});
});
let lowest = |parity: usize| {
offers
.iter()
.filter(|(_, file)| file % 2 == parity)
.min()
.map(|&(_, file)| FileId::new(file))
};
assert_eq!(claims.owner(&key("even")), lowest(0));
assert_eq!(claims.owner(&key("odd")), lowest(1));
assert_eq!(claims.len(), 2);
}
}
#[test]
fn claims_from_earlier_rounds_are_final() {
let mut claims = GroupClaims::new();
assert!(claims.is_empty());
{
let round = claims.begin_round();
assert!(round.offer(key("g"), InputPosition::new(5, 0), FileId::new(5)));
assert!(round.offer(key("g"), InputPosition::new(3, 0), FileId::new(3)));
assert!(!round.offer(key("g"), InputPosition::new(4, 0), FileId::new(4)));
assert!(round.is_owner(&key("g"), FileId::new(3)));
}
{
let round = claims.begin_round();
assert!(!round.offer(key("g"), InputPosition::new(1, 0), FileId::new(1)));
assert!(round.offer(key("h"), InputPosition::new(9, 0), FileId::new(9)));
assert_eq!(round.owner(&key("g")), Some(FileId::new(3)));
}
assert_eq!(claims.owner(&key("h")), Some(FileId::new(9)));
assert_eq!(claims.owner(&key("missing")), None);
let round = claims.begin_round();
round.offer(key("tie"), InputPosition::new(2, 0), FileId::new(8));
round.offer(key("tie"), InputPosition::new(2, 0), FileId::new(7));
assert!(round.is_owner(&key("tie"), FileId::new(7)));
}
}