use std::marker::PhantomData;
use std::sync::atomic::AtomicU8;
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering;
pub trait AtomicU8Like {
fn new(value: u8) -> Self;
fn load(&self) -> u8;
fn compare_exchange(&self, current: u8, new: u8) -> Result<u8, u8>;
fn store(&self, value: u8);
}
pub trait AtomicUsizeLike {
fn new(value: usize) -> Self;
fn load(&self) -> usize;
fn fetch_add(&self, value: usize);
fn fetch_sub(&self, value: usize);
}
pub trait YieldLike {
fn spin_loop();
fn yield_now();
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum GateLifecycle {
Open = 0,
Finishing = 1,
Closed = 2,
}
pub struct OperationGate<A8, ASize, Scheduler> {
lifecycle: A8,
active_updates: ASize,
scheduler: PhantomData<Scheduler>,
}
impl<A8, ASize, Scheduler> OperationGate<A8, ASize, Scheduler>
where
A8: AtomicU8Like,
ASize: AtomicUsizeLike,
Scheduler: YieldLike,
{
#[inline]
#[must_use]
pub fn new() -> Self {
Self {
lifecycle: A8::new(GateLifecycle::Open as u8),
active_updates: ASize::new(0),
scheduler: PhantomData,
}
}
#[inline]
#[must_use]
pub fn lifecycle(&self) -> GateLifecycle {
match self.lifecycle.load() {
value if value == GateLifecycle::Open as u8 => GateLifecycle::Open,
value if value == GateLifecycle::Finishing as u8 => {
GateLifecycle::Finishing
}
value if value == GateLifecycle::Closed as u8 => {
GateLifecycle::Closed
}
_ => unreachable!("operation lifecycle must contain a known value"),
}
}
#[inline]
pub fn enter_update(&self) -> Result<(), GateLifecycle> {
let state = self.lifecycle();
if state != GateLifecycle::Open {
return Err(state);
}
self.active_updates.fetch_add(1);
let state = self.lifecycle();
if state == GateLifecycle::Open {
Ok(())
} else {
self.active_updates.fetch_sub(1);
Err(state)
}
}
#[inline]
pub fn leave_update(&self) {
self.active_updates.fetch_sub(1);
}
#[inline]
#[must_use]
pub fn active_updates(&self) -> usize {
self.active_updates.load()
}
#[must_use]
pub fn try_begin_finish(&self) -> bool {
if self
.lifecycle
.compare_exchange(
GateLifecycle::Open as u8,
GateLifecycle::Finishing as u8,
)
.is_err()
{
return false;
}
let mut attempts = 0u32;
while self.active_updates() != 0 {
if attempts > 0 && attempts.is_multiple_of(16) {
Scheduler::yield_now();
} else {
Scheduler::spin_loop();
}
attempts += 1;
}
true
}
#[inline]
pub fn reopen(&self) {
self.lifecycle.store(GateLifecycle::Open as u8);
}
#[inline]
pub fn close(&self) {
self.lifecycle.store(GateLifecycle::Closed as u8);
}
}
impl<A8, ASize, Scheduler> Default for OperationGate<A8, ASize, Scheduler>
where
A8: AtomicU8Like,
ASize: AtomicUsizeLike,
Scheduler: YieldLike,
{
fn default() -> Self {
Self::new()
}
}
#[allow(dead_code)]
pub struct StdScheduler;
impl YieldLike for StdScheduler {
fn spin_loop() {
std::hint::spin_loop();
}
fn yield_now() {
std::thread::yield_now();
}
}
impl AtomicU8Like for AtomicU8 {
#[inline]
fn new(value: u8) -> Self {
Self::new(value)
}
#[inline]
fn load(&self) -> u8 {
self.load(Ordering::Acquire)
}
#[inline]
fn compare_exchange(&self, current: u8, new: u8) -> Result<u8, u8> {
self.compare_exchange(current, new, Ordering::AcqRel, Ordering::Acquire)
}
#[inline]
fn store(&self, value: u8) {
self.store(value, Ordering::Release);
}
}
impl AtomicUsizeLike for AtomicUsize {
#[inline]
fn new(value: usize) -> Self {
Self::new(value)
}
#[inline]
fn load(&self) -> usize {
self.load(Ordering::Acquire)
}
#[inline]
fn fetch_add(&self, value: usize) {
self.fetch_add(value, Ordering::AcqRel);
}
#[inline]
fn fetch_sub(&self, value: usize) {
self.fetch_sub(value, Ordering::Release);
}
}
#[cfg(coverage)]
#[doc(hidden)]
pub fn __coverage_operation_gate() {
use std::sync::Arc;
type StandardGate = OperationGate<AtomicU8, AtomicUsize, StdScheduler>;
let gate: StandardGate = Default::default();
let default_constructor: fn() -> StandardGate = Default::default;
let _ = default_constructor();
let spin: fn() = <StdScheduler as YieldLike>::spin_loop;
let yield_now: fn() = <StdScheduler as YieldLike>::yield_now;
spin();
yield_now();
let gate = Arc::new(gate);
assert_eq!(gate.enter_update(), Ok(()));
let finisher_gate = Arc::clone(&gate);
let finisher = std::thread::spawn(move || finisher_gate.try_begin_finish());
while gate.lifecycle() == GateLifecycle::Open {
std::thread::yield_now();
}
assert_eq!(gate.lifecycle(), GateLifecycle::Finishing);
assert_eq!(gate.enter_update(), Err(GateLifecycle::Finishing));
gate.leave_update();
assert!(finisher.join().expect("coverage finisher must join"));
gate.close();
assert_eq!(gate.enter_update(), Err(GateLifecycle::Closed));
assert!(!gate.try_begin_finish());
}