use std::{
collections::BTreeSet,
sync::atomic::{AtomicBool, AtomicU8, AtomicU64, Ordering},
};
use crate::{
backend::tokio::runtime::{Handle, RuntimeFlavor},
flash::ids::ThreadKey,
};
#[repr(u8)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(super) enum TaskState {
Parked = 0,
Runnable = 1,
Running = 2,
RunningNotified = 3,
Done = 4,
}
pub(super) struct AtomicTaskState(AtomicU8);
impl AtomicTaskState {
fn new(initial: TaskState) -> Self {
Self(AtomicU8::new(initial as u8))
}
pub(super) fn compare_exchange(&self, current: TaskState, new: TaskState) -> bool {
self.0
.compare_exchange(
current as u8,
new as u8,
Ordering::AcqRel,
Ordering::Acquire,
)
.is_ok()
}
pub(super) fn load(&self) -> TaskState {
Self::unpack(self.0.load(Ordering::Acquire))
}
pub(super) fn store(&self, new: TaskState) {
self.0.store(new as u8, Ordering::Release);
}
pub(super) fn swap(&self, new: TaskState) -> TaskState {
Self::unpack(self.0.swap(new as u8, Ordering::AcqRel))
}
fn unpack(v: u8) -> TaskState {
match v {
0 => TaskState::Parked,
1 => TaskState::Runnable,
2 => TaskState::Running,
3 => TaskState::RunningNotified,
4 => TaskState::Done,
_ => unreachable!("BUG: invalid TaskState discriminant {v}"),
}
}
}
pub(super) enum ParkOutcome {
Parked,
WokenMidPoll,
}
pub(super) enum WakeOutcome {
Resumed,
NotParked,
}
pub(super) struct TaskDiag {
pub(super) state: AtomicTaskState,
sole_poller: AtomicBool,
driver: AtomicU64,
polls: AtomicU64,
}
impl Default for TaskDiag {
fn default() -> Self {
Self {
state: AtomicTaskState::new(TaskState::Runnable),
polls: AtomicU64::new(0),
driver: AtomicU64::new(0),
sole_poller: AtomicBool::new(false),
}
}
}
impl TaskDiag {
pub(super) fn driver(&self) -> Option<ThreadKey> {
(self.polls() > 0).then(|| ThreadKey::from(self.driver.load(Ordering::Relaxed)))
}
pub(super) fn enter_poll(&self, driver: ThreadKey) {
self.driver.store(driver.raw(), Ordering::Relaxed);
self.sole_poller
.store(sole_poller_runtime(), Ordering::Relaxed);
self.polls.fetch_add(1, Ordering::Release);
}
pub(super) fn polls(&self) -> u64 {
self.polls.load(Ordering::Acquire)
}
pub(super) fn stranded_behind(&self, bridged: &BTreeSet<ThreadKey>) -> bool {
self.driver()
.is_some_and(|driver| bridged.contains(&driver))
&& self.sole_poller.load(Ordering::Relaxed)
&& self.state.load() == TaskState::Runnable
}
}
fn sole_poller_runtime() -> bool {
matches!(
Handle::try_current().map(|h| h.runtime_flavor()),
Ok(RuntimeFlavor::CurrentThread)
)
}