use std::cell::Cell;
use std::num::NonZeroUsize;
use std::sync::atomic::{AtomicUsize, Ordering};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct LocalSharding {
shards: NonZeroUsize,
mask: usize,
}
impl LocalSharding {
pub const SINGLE: Self = Self::new(NonZeroUsize::MIN);
const NOT_A_MASK: usize = usize::MAX;
#[must_use]
pub const fn new(shards: NonZeroUsize) -> Self {
Self {
shards,
mask: if shards.get().is_power_of_two() {
shards.get() - 1
} else {
Self::NOT_A_MASK
},
}
}
#[must_use]
pub fn available_parallelism() -> Self {
Self::new(std::thread::available_parallelism().unwrap_or(NonZeroUsize::MIN))
}
#[must_use]
pub const fn get(self) -> usize {
self.shards.get()
}
#[must_use]
pub fn occupancy(self) -> ShardOccupancy {
ShardOccupancy {
shards: self.get(),
affinities_assigned: Locality::assigned(),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ShardOccupancy {
pub shards: usize,
pub affinities_assigned: usize,
}
impl ShardOccupancy {
#[must_use]
pub fn crowded_shards(self) -> usize {
self.shards
.min(self.affinities_assigned.saturating_sub(self.shards))
}
#[must_use]
pub fn is_crowded(self) -> bool {
self.affinities_assigned > self.shards
}
}
impl Default for LocalSharding {
fn default() -> Self {
Self::SINGLE
}
}
static NEXT_LOCALITY: AtomicUsize = AtomicUsize::new(0);
thread_local! {
static LOCALITY: Cell<usize> = Cell::new(NEXT_LOCALITY.fetch_add(1, Ordering::Relaxed));
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Locality(usize);
impl Locality {
pub const OBSERVER: Self = Self(0);
#[inline]
#[must_use]
pub fn current() -> Self {
LOCALITY.with(|locality| Self(locality.get()))
}
#[must_use]
pub fn assigned() -> usize {
NEXT_LOCALITY.load(Ordering::Relaxed)
}
#[inline]
#[must_use]
pub fn index(self, sharding: LocalSharding) -> usize {
if sharding.mask == LocalSharding::NOT_A_MASK {
self.0 % sharding.shards
} else {
self.0 & sharding.mask
}
}
#[cfg(test)]
pub(crate) const fn for_test(value: usize) -> Self {
Self(value)
}
}
#[cfg(test)]
#[path = "../tests/support/isolated.rs"]
mod isolated;
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn single_sharding_always_selects_the_only_shard() {
assert_eq!(Locality::current().index(LocalSharding::SINGLE), 0);
}
#[test]
fn one_thread_keeps_one_affinity() {
assert_eq!(Locality::current(), Locality::current());
}
#[test]
fn host_parallelism_helper_preserves_the_advertised_count() {
let expected = std::thread::available_parallelism().unwrap_or(NonZeroUsize::MIN);
assert_eq!(LocalSharding::available_parallelism().get(), expected.get());
}
#[test]
fn power_of_two_and_arbitrary_counts_select_the_same_modulo_index() {
let locality = Locality(13);
assert_eq!(
locality.index(LocalSharding::new(NonZeroUsize::new(8).unwrap())),
5
);
assert_eq!(
locality.index(LocalSharding::new(NonZeroUsize::new(5).unwrap())),
3
);
}
#[test]
fn a_masked_reduction_agrees_with_the_modulo_it_replaces() {
for shards in 1..=64usize {
let sharding = LocalSharding::new(NonZeroUsize::new(shards).unwrap());
assert_eq!(sharding.get(), shards, "the count itself must not move");
for value in [0usize, 1, 7, 13, 64, 255, 4_096, usize::MAX - 1, usize::MAX] {
assert_eq!(
Locality(value).index(sharding),
value % shards,
"{value} on {shards} shards"
);
assert!(Locality(value).index(sharding) < shards);
}
}
}
#[test]
fn a_power_of_two_count_masks_and_any_other_divides() {
for shards in [1usize, 2, 4, 8, 16, 1_024] {
let sharding = LocalSharding::new(NonZeroUsize::new(shards).unwrap());
assert_eq!(
sharding.mask,
shards - 1,
"{shards} is a power of two and must reduce by mask"
);
assert_ne!(sharding.mask, LocalSharding::NOT_A_MASK);
}
for shards in [3usize, 5, 6, 7, 10, 100] {
let sharding = LocalSharding::new(NonZeroUsize::new(shards).unwrap());
assert_eq!(
sharding.mask,
LocalSharding::NOT_A_MASK,
"{shards} is not a power of two and must reduce by modulo"
);
}
}
#[test]
fn the_modulo_sentinel_is_not_a_reachable_mask() {
let largest = NonZeroUsize::new(1usize << (usize::BITS - 1)).unwrap();
let sharding = LocalSharding::new(largest);
assert_ne!(sharding.mask, LocalSharding::NOT_A_MASK);
assert_eq!(sharding.mask, largest.get() - 1);
assert_eq!(LocalSharding::SINGLE.mask, 0);
assert_eq!(LocalSharding::SINGLE.get(), 1);
}
#[test]
fn shardings_compare_by_the_count_they_were_built_from() {
let four = LocalSharding::new(NonZeroUsize::new(4).unwrap());
assert_eq!(four, LocalSharding::new(NonZeroUsize::new(4).unwrap()));
assert_ne!(four, LocalSharding::new(NonZeroUsize::new(5).unwrap()));
assert_eq!(LocalSharding::SINGLE, LocalSharding::default());
}
#[test]
fn crowded_shards_counts_the_residues_the_counter_hands_out() {
for shards in 1..=16usize {
let sharding = LocalSharding::new(NonZeroUsize::new(shards).unwrap());
for assigned in 0..=(3 * shards) {
let occupancy = ShardOccupancy {
shards,
affinities_assigned: assigned,
};
let mut carried = vec![0usize; shards];
for affinity in 0..assigned {
carried[Locality(affinity).index(sharding)] += 1;
}
let expected = carried.iter().filter(|held| **held > 1).count();
assert_eq!(
occupancy.crowded_shards(),
expected,
"{assigned} affinities on {shards} shard(s) crowd {expected} of them"
);
assert_eq!(
occupancy.is_crowded(),
expected > 0,
"{assigned} on {shards}: crowding must agree with the count"
);
assert!(occupancy.crowded_shards() <= shards);
}
}
}
#[test]
fn a_layout_with_room_reports_every_affinity_distinct() {
for shards in 1..=16usize {
let sharding = LocalSharding::new(NonZeroUsize::new(shards).unwrap());
for assigned in 0..=shards {
let occupancy = ShardOccupancy {
shards,
affinities_assigned: assigned,
};
assert!(
!occupancy.is_crowded(),
"{assigned} affinities fit {shards} shard(s)"
);
assert_eq!(occupancy.crowded_shards(), 0);
}
let over = ShardOccupancy {
shards,
affinities_assigned: shards + 1,
};
assert!(over.is_crowded());
assert_eq!(over.crowded_shards(), 1);
assert_eq!(sharding.get(), shards);
}
}
#[test]
fn occupancy_reports_the_counter_without_consuming_from_it() {
if isolated::rerun_in_child() {
return;
}
let sharding = LocalSharding::new(NonZeroUsize::new(4).unwrap());
assert_eq!(Locality::assigned(), 0);
for expected in 0..=5 {
for _ in 0..2 {
let now = sharding.occupancy();
assert_eq!(now.shards, 4);
assert_eq!(now.affinities_assigned, expected, "the counter is live");
assert_eq!(Locality::assigned(), expected, "reporting spends nothing");
}
std::thread::spawn(Locality::current).join().unwrap();
}
}
#[test]
fn the_observer_affinity_costs_nothing_and_never_moves() {
if isolated::rerun_in_child() {
return;
}
let before = Locality::assigned();
assert_eq!(before, 0, "the thread has never claimed an affinity");
for shards in 1..=16usize {
let sharding = LocalSharding::new(NonZeroUsize::new(shards).unwrap());
for _ in 0..64 {
assert_eq!(Locality::OBSERVER.index(sharding), 0, "{shards} shards");
}
}
assert_eq!(Locality::assigned(), before, "observing spends nothing");
}
}