extern crate alloc;
use alloc::sync::{Arc, Weak};
use core::{fmt, ops::Range, sync::atomic::Ordering};
use portable_atomic::AtomicU128;
#[derive(Debug, Clone)]
pub struct Task(Arc<TaskInner>);
#[derive(Debug)]
struct TaskInner {
state: Arc<AtomicU128>,
}
#[derive(Debug, Clone)]
pub struct WeakTask(Weak<TaskInner>);
impl WeakTask {
#[must_use]
pub fn upgrade(&self) -> Option<Task> {
self.0.upgrade().map(Task)
}
#[must_use]
pub fn strong_count(&self) -> usize {
self.0.strong_count()
}
#[must_use]
pub fn weak_count(&self) -> usize {
self.0.weak_count()
}
#[must_use]
pub fn is_alive(&self) -> bool {
self.0.strong_count() > 0
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct RangeError;
impl fmt::Display for RangeError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "Range invariant violated: start > end")
}
}
impl core::error::Error for RangeError {}
impl Task {
#[allow(clippy::inline_always)]
#[inline(always)]
const fn pack(range: Range<u64>) -> u128 {
((range.start as u128) << 64) | (range.end as u128)
}
#[allow(clippy::inline_always)]
#[inline(always)]
const fn unpack(state: u128) -> Range<u64> {
#[allow(clippy::cast_possible_truncation)]
let end = state as u64;
(state >> 64) as u64..end
}
#[must_use]
pub fn new(range: Range<u64>) -> Self {
assert!(range.start <= range.end);
Self(Arc::new(TaskInner {
state: Arc::new(AtomicU128::new(Self::pack(range))),
}))
}
#[must_use]
pub fn get(&self) -> Range<u64> {
let state = self.0.state.load(Ordering::Acquire);
Self::unpack(state)
}
#[must_use]
pub fn start(&self) -> u64 {
(self.0.state.load(Ordering::Acquire) >> 64) as u64
}
pub fn safe_add_start(&self, start: u64, bias: u64) -> Result<Range<u64>, RangeError> {
let new_start = start.saturating_add(bias);
let mut old_state = self.0.state.load(Ordering::Acquire);
loop {
let mut range = Self::unpack(old_state);
if start > range.start {
break Err(RangeError);
}
let new_start = new_start.min(range.end);
if new_start <= range.start {
break Err(RangeError);
}
let span = range.start..new_start;
range.start = new_start;
let new_state = Self::pack(range);
match self.0.state.compare_exchange_weak(
old_state,
new_state,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => break Ok(span),
Err(x) => old_state = x,
}
}
}
#[must_use]
pub fn end(&self) -> u64 {
let state = self.0.state.load(Ordering::Acquire);
#[allow(clippy::cast_possible_truncation)]
let end = state as u64;
end
}
#[must_use]
pub fn remain(&self) -> u64 {
let range = self.get();
range.end.saturating_sub(range.start)
}
pub fn split_two(&self, min_chunk_size: u64) -> Result<Option<Range<u64>>, RangeError> {
let mut old_state = self.0.state.load(Ordering::Acquire);
loop {
let range = Self::unpack(old_state);
if range.start > range.end {
return Err(RangeError);
}
if range.end - range.start < min_chunk_size.saturating_mul(2) {
return Ok(None);
}
let mid = range.start.midpoint(range.end);
let new_state = Self::pack(range.start..mid);
match self.0.state.compare_exchange_weak(
old_state,
new_state,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => return Ok(Some(mid..range.end)),
Err(x) => old_state = x,
}
}
}
pub fn take(&self) -> Result<Option<Range<u64>>, RangeError> {
let mut old_state = self.0.state.load(Ordering::Acquire);
loop {
let range = Self::unpack(old_state);
if range.start > range.end {
return Err(RangeError);
}
if range.start == range.end {
return Ok(None);
}
let new_state = Self::pack(range.start..range.start);
match self.0.state.compare_exchange_weak(
old_state,
new_state,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => return Ok(Some(range)),
Err(x) => old_state = x,
}
}
}
#[must_use]
pub fn downgrade(&self) -> WeakTask {
WeakTask(Arc::downgrade(&self.0))
}
#[must_use]
pub fn strong_count(&self) -> usize {
Arc::strong_count(&self.0)
}
#[must_use]
pub fn weak_count(&self) -> usize {
Arc::weak_count(&self.0)
}
#[must_use]
pub(crate) fn sharer_count(&self) -> usize {
Arc::strong_count(&self.0.state)
}
pub(crate) fn share_state(&mut self, other: &Self) {
*self = Self(Arc::new(TaskInner {
state: other.0.state.clone(),
}));
}
#[cfg(test)]
#[must_use]
pub(crate) fn from_raw_state(state: Arc<AtomicU128>) -> Self {
Self(Arc::new(TaskInner { state }))
}
}
impl From<Range<u64>> for Task {
fn from(value: Range<u64>) -> Self {
Self::new(value)
}
}
impl PartialEq for Task {
fn eq(&self, other: &Self) -> bool {
Arc::ptr_eq(&self.0.state, &other.0.state)
}
}
impl Eq for Task {}
#[cfg(test)]
mod tests {
#![allow(clippy::unwrap_used)]
extern crate std;
use super::*;
use std::sync::Arc;
use std::thread;
use std::vec::Vec;
#[test]
fn test_new_task() {
let task = Task::new(10..20);
assert_eq!(task.start(), 10);
assert_eq!(task.end(), 20);
assert_eq!(task.remain(), 10);
}
#[test]
fn test_remain() {
let task = Task::new(10..25);
assert_eq!(task.remain(), 15);
}
#[test]
fn test_split_two() {
let task = Task::new(1..6); let range = task.split_two(1).unwrap().unwrap();
assert_eq!(task.start(), 1);
assert_eq!(task.end(), 3);
assert_eq!(range.start, 3);
assert_eq!(range.end, 6);
}
#[test]
fn test_split_empty() {
let task = Task::new(1..1);
let range = task.split_two(1).unwrap();
assert_eq!(task.start(), 1);
assert_eq!(task.end(), 1);
assert_eq!(range, None);
}
#[test]
fn test_split_one() {
let task = Task::new(1..2);
let range = task.split_two(1).unwrap();
assert_eq!(task.start(), 1);
assert_eq!(task.end(), 2);
assert_eq!(range, None);
}
#[test]
fn test_safe_add_start_no_progress() {
let task = Task::new(10..20);
assert_eq!(task.safe_add_start(10, 0), Err(RangeError));
assert_eq!(task.safe_add_start(8, 2), Err(RangeError));
}
#[test]
fn test_safe_add_start_claims_span() {
let task = Task::new(10..20);
let span = task.safe_add_start(10, 5).unwrap();
assert_eq!(span, 10..15);
assert_eq!(task.start(), 15);
assert_eq!(task.remain(), 5);
}
#[test]
fn test_safe_add_start_capped_at_end() {
let task = Task::new(10..12);
let span = task.safe_add_start(10, 100).unwrap();
assert_eq!(span, 10..12);
assert_eq!(task.remain(), 0);
}
#[test]
fn test_take_empties() {
let task = Task::new(5..9);
assert_eq!(task.take(), Ok(Some(5..9)));
assert_eq!(task.take(), Ok(None));
assert_eq!(task.remain(), 0);
}
#[test]
fn test_downgrade_upgrade() {
let task = Task::new(1..10);
let weak = task.downgrade();
assert_eq!(weak.strong_count(), 1);
assert_eq!(weak.upgrade().unwrap().get(), 1..10);
drop(task);
assert_eq!(weak.upgrade(), None);
}
#[test]
fn test_partial_eq_by_ptr() {
let a = Task::new(1..10);
let b = a.clone();
assert_eq!(a, b);
let c = Task::new(1..10);
assert_ne!(a, c);
}
#[test]
fn test_split_two_halves() {
let task = Task::new(0..100);
let range = task.split_two(1).unwrap().unwrap();
assert_eq!(range, 50..100);
assert_eq!(task.get(), 0..50);
}
#[test]
fn split_two_respects_min_chunk_size() {
let task = Task::new(0..(2 * 8 - 1)); assert_eq!(task.split_two(8), Ok(None));
assert_eq!(task.get(), 0..15);
let task = Task::new(0..16);
let range = task.split_two(8).unwrap().unwrap();
assert_eq!(range, 8..16);
assert_eq!(task.get(), 0..8);
let task = Task::new(0..10);
assert_eq!(task.split_two(8), Ok(None));
}
#[test]
fn weak_task_reports_strong_and_weak_counts() {
let task = Task::new(1..10);
let weak = task.downgrade();
assert_eq!(weak.strong_count(), 1);
assert_eq!(weak.weak_count(), 1);
let weak2 = task.downgrade();
assert_eq!(weak2.weak_count(), 2);
drop(weak2);
assert_eq!(weak.weak_count(), 1);
}
#[test]
fn task_weak_count_reflects_weak_refs() {
let task = Task::new(1..10);
assert_eq!(task.weak_count(), 0);
let _w1 = task.downgrade();
assert_eq!(task.weak_count(), 1);
let _w2 = task.downgrade();
assert_eq!(task.weak_count(), 2);
}
#[test]
fn safe_add_start_survives_contention() {
let task = Arc::new(Task::new(0..2_000));
let mut handles = Vec::new();
for _ in 0..4 {
let t = task.clone();
handles.push(thread::spawn(move || {
loop {
let s = t.start();
if s >= 2_000 {
break;
}
if t.safe_add_start(s, 1).is_err() {
}
}
}));
}
for h in handles {
h.join().unwrap();
}
assert_eq!(task.get(), 2_000..2_000);
assert_eq!(task.remain(), 0);
}
#[test]
fn split_two_survives_contention() {
let task = Arc::new(Task::new(0..2_000));
let mut handles = Vec::new();
for _ in 0..4 {
let t = task.clone();
handles.push(thread::spawn(
move || {
while t.split_two(1).unwrap().is_some() {}
},
));
}
for h in handles {
h.join().unwrap();
}
assert_eq!(task.remain(), 1);
}
#[test]
fn take_survives_contention() {
let task = Arc::new(Task::new(0..2_000));
let mut handles = Vec::new();
for _ in 0..4 {
let t = task.clone();
handles.push(thread::spawn(
move || {
while matches!(t.take(), Ok(Some(_))) {}
},
));
}
for h in handles {
h.join().unwrap();
}
assert_eq!(task.remain(), 0);
}
#[test]
fn split_two_reports_invariant_violation_when_start_gt_end() {
let bad = Task::from_raw_state(std::sync::Arc::new(portable_atomic::AtomicU128::new(
(20u128 << 64) | 0xA,
)));
assert_eq!(bad.split_two(1), Err(RangeError));
}
#[test]
fn range_error_display_and_error_impl() {
use std::{error::Error, format, string::ToString};
let e = RangeError;
assert_eq!(e.to_string(), "Range invariant violated: start > end");
assert_eq!(format!("{e}"), "Range invariant violated: start > end");
let as_dyn: &dyn Error = &e;
assert!(as_dyn.source().is_none());
}
#[test]
#[should_panic(expected = "assertion failed")]
fn new_panics_when_start_gt_end() {
let _ = Task::new(core::ops::Range { start: 10, end: 5 });
}
#[test]
#[should_panic(expected = "assertion failed")]
fn from_range_panics_when_start_gt_end() {
let _ = Task::from(core::ops::Range { start: 10, end: 5 });
}
#[test]
fn from_range_builds_equivalent_task() {
let t = Task::from(3..9);
assert_eq!(t.get(), 3..9);
let t2: Task = (0..0).into();
assert_eq!(t2.get(), 0..0);
assert_eq!(t2.remain(), 0);
}
#[test]
fn pack_unpack_round_trips_at_u64_bounds() {
for range in [
0..0,
0..u64::MAX,
u64::MAX..u64::MAX,
(u64::MAX - 1)..u64::MAX,
] {
let t = Task::new(range.clone());
assert_eq!(t.get(), range, "round-trip lost bits");
assert_eq!(t.start(), range.start);
assert_eq!(t.end(), range.end);
}
assert_eq!(Task::new(0..u64::MAX).remain(), u64::MAX);
}
#[test]
fn safe_add_start_saturates_on_u64_overflow() {
let task = Task::new(0..u64::MAX);
let span = task.safe_add_start(0, u64::MAX).unwrap();
assert_eq!(span, 0..u64::MAX);
assert_eq!(task.get(), u64::MAX..u64::MAX);
assert_eq!(task.remain(), 0);
let task = Task::new((u64::MAX - 5)..u64::MAX);
let span = task.safe_add_start(u64::MAX - 5, u64::MAX).unwrap();
assert_eq!(span, (u64::MAX - 5)..u64::MAX);
assert_eq!(task.remain(), 0);
}
#[test]
fn safe_add_start_rejects_caller_start_ahead_of_cursor() {
let task = Task::new(0..100);
assert_eq!(task.safe_add_start(50, 1), Err(RangeError));
assert_eq!(task.get(), 0..100, "a rejected call must not mutate state");
assert_eq!(task.safe_add_start(500, 1), Err(RangeError));
assert_eq!(task.get(), 0..100);
task.safe_add_start(0, 10).unwrap(); assert_eq!(task.safe_add_start(11, 1), Err(RangeError));
assert_eq!(task.get(), 10..100);
assert_eq!(task.safe_add_start(10, 1), Ok(10..11));
}
#[test]
fn safe_add_start_span_shrinks_when_caller_start_is_stale() {
let task = Task::new(0..100);
task.safe_add_start(0, 5).unwrap(); let span = task.safe_add_start(3, 5).unwrap(); assert_eq!(
span,
5..8,
"span must start at the real cursor, not the stale one"
);
assert_eq!(task.start(), 8);
}
#[test]
fn take_reports_invariant_violation_on_corrupted_state() {
let bad = Task::from_raw_state(Arc::new(portable_atomic::AtomicU128::new(
(20u128 << 64) | 0xA,
)));
let inverted = core::ops::Range {
start: 20u64,
end: 10u64,
};
assert_eq!(bad.get(), inverted);
assert_eq!(bad.take(), Err(RangeError));
assert_eq!(bad.get(), inverted);
assert_eq!(bad.split_two(1), Err(RangeError));
}
#[test]
fn remain_saturates_to_zero_on_corrupted_state() {
let bad = Task::from_raw_state(Arc::new(portable_atomic::AtomicU128::new(
(20u128 << 64) | 0xA,
)));
assert_eq!(bad.remain(), 0);
}
#[test]
fn weak_task_weak_count_collapses_to_zero_without_strong_refs() {
let task = Task::new(0..10);
let w1 = task.downgrade();
let _w2 = task.downgrade();
assert_eq!(w1.weak_count(), 2);
assert_eq!(w1.strong_count(), 1);
drop(task);
assert_eq!(w1.strong_count(), 0);
assert_eq!(
w1.weak_count(),
0,
"weak_count collapses once the strong count hits 0"
);
assert!(w1.upgrade().is_none());
}
#[test]
fn task_sharer_count_tracks_cursor_sharers() {
let task = Task::new(0..10);
assert_eq!(task.sharer_count(), 1);
let mut twin = Task::new(0..0);
twin.share_state(&task);
assert_eq!(
task.sharer_count(),
2,
"the twin holds a strong ref to the cursor"
);
assert_eq!(twin, task, "twins compare equal via the shared cursor");
drop(twin);
assert_eq!(
task.sharer_count(),
1,
"dropping the twin releases its cursor ref"
);
let _alias = task.clone();
assert_eq!(
task.sharer_count(),
1,
"clone shares identity, not a cursor ref"
);
let w = task.downgrade();
let _up = w.upgrade().unwrap();
assert_eq!(
task.sharer_count(),
1,
"upgrade keeps the cursor count unchanged"
);
}
#[test]
fn task_strong_count_counts_identity_not_cursor() {
let task = Task::new(0..10);
let weak = task.downgrade();
assert_eq!(task.strong_count(), 1);
assert_eq!(weak.strong_count(), 1);
let _alias = task.clone();
assert_eq!(task.strong_count(), 2);
assert_eq!(weak.strong_count(), 2);
assert_eq!(task.sharer_count(), 1, "clone does not alias the cursor");
let mut twin = Task::new(0..0);
twin.share_state(&task);
assert_eq!(task.sharer_count(), 2, "share_state aliases the cursor");
assert_eq!(task.strong_count(), 2, "share_state keeps its own identity");
}
#[test]
fn weak_task_is_alive_tracks_identity_not_cursor() {
let task = Task::new(0..10);
let weak = task.downgrade();
assert!(weak.is_alive());
drop(task);
assert!(!weak.is_alive(), "identity dropped -> not alive");
let victim = Task::new(0..10);
let mut twin = Task::new(0..0);
twin.share_state(&victim);
let twin_weak = twin.downgrade();
assert!(twin_weak.is_alive());
drop(victim);
assert!(
twin_weak.is_alive(),
"twin keeps its own identity alive after the victim dies"
);
}
#[test]
fn split_two_handles_u64_extremes_without_overflow() {
let task = Task::new(0..u64::MAX);
let hi = task.split_two(1).unwrap().unwrap();
assert_eq!(hi, (u64::MAX / 2)..u64::MAX);
assert_eq!(task.get(), 0..(u64::MAX / 2));
let task = Task::new((u64::MAX - 3)..u64::MAX);
let hi = task.split_two(1).unwrap().unwrap();
assert_eq!(hi, (u64::MAX - 2)..u64::MAX);
assert_eq!(task.get(), (u64::MAX - 3)..(u64::MAX - 2));
}
}