use crate::SimCxl;
use crate::event::Event;
use crate::id::Id;
use crate::node_id::NodeId;
use crate::simulator::for_all_simulators;
use crate::{SimCx, time::TimeScheduler};
use cooked_waker::{IntoWaker, WakeRef};
use futures::channel::oneshot;
use std::collections::HashSet;
use std::error::Error;
use std::fmt::{Debug, Display};
use std::io;
use std::marker::PhantomData;
use std::sync::atomic::Ordering::Relaxed;
use std::task::Poll;
use std::{
cell::Cell,
collections::{HashMap, VecDeque},
pin::Pin,
sync::{Arc, atomic::AtomicUsize},
task::Context,
};
pub(crate) struct ExecutorQueue {
ready_queue: VecDeque<Id>,
}
impl ExecutorQueue {
pub fn none_ready(&self) -> bool {
self.ready_queue.is_empty()
}
pub fn new() -> Self {
ExecutorQueue {
ready_queue: VecDeque::new(),
}
}
pub(crate) fn executor(&self) -> Executor {
Executor {
final_stopped: false,
tasks: HashMap::new(),
nodes: vec![NodeData {
run_level: NodeRunLevel::Running,
tasks: HashSet::new(),
}],
time_scheduler: TimeScheduler::new(),
}
}
}
pub(crate) struct Executor {
#[allow(clippy::type_complexity)]
tasks: HashMap<Id, TaskEntry>,
final_stopped: bool,
nodes: Vec<NodeData>,
pub(crate) time_scheduler: TimeScheduler,
}
enum NodeRunLevel {
Running,
Stopped,
FinalStopped,
}
struct NodeData {
run_level: NodeRunLevel,
tasks: HashSet<Id>,
}
struct TaskEntry {
shared: Arc<TaskShared>,
task: Cell<Option<Pin<Box<dyn TaskDyn>>>>,
}
impl TaskEntry {
#[cfg_attr(not(feature = "emit-tracing"), allow(dead_code))]
fn is_task_none(&self) -> bool {
let task = self.task.take();
let is_none = task.is_none();
self.task.set(task);
is_none
}
}
pin_project_lite::pin_project! {
struct Task<F:Future>{
shared:Arc<TaskShared>,
snd:Option<oneshot::Sender<F::Output>>,
#[pin]
fut: F,
}
}
trait TaskDyn {
fn run(self: Pin<&mut Self>) -> Poll<()>;
fn as_base(&self) -> &Arc<TaskShared>;
}
impl<F: Future> TaskDyn for Task<F> {
fn run(self: Pin<&mut Self>) -> Poll<()> {
let this = self.project();
let waker = this.shared.clone().into_waker();
let poll_result = {
#[cfg(feature = "emit-tracing")]
let _guard = this.shared.tracing_span.enter();
this.fut.poll(&mut Context::from_waker(&waker))
};
match poll_result {
Poll::Ready(x) => {
this.snd.take().unwrap().send(x).ok();
Poll::Ready(())
}
Poll::Pending => Poll::Pending,
}
}
fn as_base(&self) -> &Arc<TaskShared> {
&self.shared
}
}
pin_project_lite::pin_project! {
#[must_use]
pub struct TaskHandle<T> {
#[pin]
result: oneshot::Receiver<T>,
abort: AbortHandle,
}
}
#[derive(Debug)]
pub struct TaskAborted(());
impl From<TaskAborted> for io::Error {
fn from(_value: TaskAborted) -> Self {
io::Error::other("aborted")
}
}
impl Display for TaskAborted {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
Debug::fmt(self, f)
}
}
impl Error for TaskAborted {}
impl<T> TaskHandle<T> {
pub fn abort(self) {
self.abort.abort_inner();
}
pub fn detach(self) {}
pub fn abort_handle(&self) -> AbortHandle {
self.abort.clone()
}
}
#[derive(Clone)]
pub struct AbortHandle {
shared: Arc<TaskShared>,
_unsend: PhantomData<*const u8>,
}
#[derive(Clone)]
pub struct AbortGuard(AbortHandle);
impl Drop for AbortGuard {
fn drop(&mut self) {
self.0.abort_inner();
}
}
const TASK_COMPLETE_READY: usize = 0;
const TASK_COMPLETE: usize = 1;
const TASK_ABORTED_READY: usize = 2;
const TASK_ABORTED: usize = 3;
const TASK_READY: usize = 4;
const TASK_WAITING: usize = 5;
const TASK_END: usize = 6;
struct TaskShared {
state: AtomicUsize,
id: Id,
#[cfg(feature = "emit-tracing")]
tracing_span: tracing::Span,
node: crate::node_id::NodeId,
}
impl WakeRef for TaskShared {
fn wake_by_ref(&self) {
SimCx::with(|cx| {
let mut ex = cx.queue.borrow_mut();
let ex = ex.as_mut().unwrap();
match self.state.load(Relaxed) {
TASK_COMPLETE | TASK_COMPLETE_READY | TASK_ABORTED | TASK_ABORTED_READY
| TASK_READY => (),
TASK_WAITING => {
#[cfg(feature = "emit-tracing")]
tracing::debug!(task = self.id.tv(), "wake");
self.state.store(TASK_READY, Relaxed);
ex.ready_queue.push_back(self.id);
}
TASK_END.. => unreachable!(),
};
});
}
}
pub fn spawn_task_on_node<F: Future + 'static>(
node: crate::node_id::NodeId,
future: F,
) -> TaskHandle<F::Output> {
SimCxl::with(|cx| {
cx.event_handler.handle_event(Event::TaskSpawned);
cx.executor.spawn(node, future)
})
}
impl Executor {
pub fn node_count(&self) -> usize {
self.nodes.len()
}
pub fn is_final_stopping(&self) -> bool {
self.final_stopped
}
pub fn push_new_node(&mut self) -> crate::node_id::NodeId {
let id = NodeId::from_index(self.node_count());
assert!(!self.final_stopped, "node spawn during simulation shutdown");
self.nodes.push(NodeData {
run_level: NodeRunLevel::Running,
tasks: HashSet::new(),
});
id
}
fn spawn<F: Future + 'static>(
&mut self,
node: crate::node_id::NodeId,
future: F,
) -> TaskHandle<F::Output> {
match self.nodes[node.0.get() - 1].run_level {
NodeRunLevel::Running => (),
NodeRunLevel::Stopped | NodeRunLevel::FinalStopped => {
panic!("node {node:?} is stopped")
}
}
let (snd, rcv) = oneshot::channel();
let task_id = Id::new();
#[cfg(feature = "emit-tracing")]
tracing::info!(task = task_id.tv(), node = node.tv(), "spawn");
let task = Box::pin(Task {
snd: Some(snd),
fut: future,
shared: Arc::new(TaskShared {
state: AtomicUsize::new(TASK_WAITING),
id: task_id,
node,
#[cfg(feature = "emit-tracing")]
tracing_span: tracing::info_span!(
parent:None,
"task",
task = task_id.tv()
),
}),
});
let task_entry = TaskEntry {
shared: task.shared.clone(),
task: Cell::new(Some(task)),
};
let task_handle = TaskHandle {
result: rcv,
abort: AbortHandle {
shared: task_entry.shared.clone(),
_unsend: PhantomData,
},
};
self.tasks.insert(task_id, task_entry);
self.nodes[node.0.get() - 1].tasks.insert(task_id);
<TaskShared as WakeRef>::wake_by_ref(&task_handle.abort.shared);
task_handle
}
fn remove_task_entry(&mut self, task_id: Id) {
let removed = self.tasks.remove(&task_id);
let task_entry = removed.unwrap();
debug_assert_eq!(task_entry.shared.id, task_id);
assert!(task_entry.task.into_inner().is_none());
let node = &mut self.nodes[task_entry.shared.node.0.get() - 1];
let removed = node.tasks.remove(&task_id);
debug_assert!(removed);
}
pub(crate) fn run_current_context(cx: &SimCx) {
loop {
let Some(mut task) = cx.with_cx(|cxl| {
if cxl.executor.final_stopped {
return None;
}
loop {
let task_id = cx
.queue
.borrow_mut()
.as_mut()
.unwrap()
.ready_queue
.pop_front()?;
#[cfg(feature = "emit-tracing")]
tracing::trace!(task = task_id.tv(), "popped");
let task_entry = cxl.executor.tasks.get_mut(&task_id).unwrap();
match task_entry.shared.state.load(Relaxed) {
TASK_READY => {
let task = task_entry.task.take().unwrap();
cxl.event_handler
.handle_event(Event::TaskRun(task_entry.shared.id));
task_entry.shared.state.store(TASK_WAITING, Relaxed);
return Some(task);
}
state @ (TASK_COMPLETE_READY | TASK_ABORTED_READY) => {
let id = task_entry.shared.id;
task_entry.shared.state.store(
match state {
TASK_COMPLETE_READY => TASK_COMPLETE,
TASK_ABORTED_READY => TASK_ABORTED,
_ => unreachable!(),
},
Relaxed,
);
cxl.executor.remove_task_entry(id);
continue;
}
TASK_COMPLETE | TASK_ABORTED | TASK_WAITING | TASK_END.. => unreachable!(),
}
}
}) else {
if cx.with_cx(|cxl| {
cxl.executor
.time_scheduler
.wait_until_next_future_ready(&cx.cxu().time, &mut *cxl.event_handler)
}) {
continue;
} else {
break;
}
};
let &TaskShared { node, id, .. } = &**task.as_base();
cx.node_scope(node, move || {
#[cfg(feature = "emit-tracing")]
tracing::debug!(task = id.tv(), "poll");
let poll_result = task.as_mut().run();
cx.with_cx(|cxl| {
match poll_result {
Poll::Pending => {
match task.as_base().state.load(Relaxed) {
TASK_ABORTED_READY => {
drop(task);
}
TASK_WAITING | TASK_READY => {
let task_entry = cxl.executor.tasks.get_mut(&id).unwrap();
assert!(task_entry.task.replace(Some(task)).is_none());
}
TASK_ABORTED | TASK_COMPLETE | TASK_COMPLETE_READY | TASK_END.. => {
unreachable!()
}
}
}
Poll::Ready(()) => {
let shared = task.as_base();
#[cfg(feature = "emit-tracing")]
tracing::info!(task = shared.id.tv(), "complete");
match shared.state.load(Relaxed) {
TASK_ABORTED | TASK_COMPLETE | TASK_COMPLETE_READY | TASK_END.. => {
unreachable!()
}
TASK_WAITING => {
shared.state.store(TASK_COMPLETE, Relaxed);
cxl.executor.remove_task_entry(id);
}
TASK_READY | TASK_ABORTED_READY => {
shared.state.store(TASK_COMPLETE_READY, Relaxed);
}
};
drop(task);
}
}
})
});
}
}
}
fn abort_local(
task_id: Id,
queue: &mut ExecutorQueue,
tasks: &mut HashMap<Id, TaskEntry>,
) -> Option<Pin<Box<dyn TaskDyn>>> {
let task_entry = tasks.get_mut(&task_id)?;
let state = task_entry.shared.state.load(Relaxed);
match state {
TASK_ABORTED | TASK_ABORTED_READY | TASK_COMPLETE | TASK_COMPLETE_READY => None,
TASK_READY => {
#[cfg(feature = "emit-tracing")]
tracing::info!(
task = task_id.tv(),
is_ready = true,
is_running = task_entry.is_task_none(),
"abort"
);
task_entry.shared.state.store(TASK_ABORTED_READY, Relaxed);
task_entry.task.take()
}
TASK_WAITING => {
#[cfg(feature = "emit-tracing")]
tracing::info!(
task = task_id.tv(),
is_ready = false,
is_running = task_entry.is_task_none(),
"abort"
);
task_entry.shared.state.store(TASK_ABORTED_READY, Relaxed);
queue.ready_queue.push_back(task_id);
task_entry.task.take()
}
TASK_END.. => unreachable!(),
}
}
impl AbortHandle {
pub fn abort(self) {
self.abort_inner();
}
pub fn abort_on_drop(self) -> AbortGuard {
AbortGuard(self)
}
fn abort_inner(&self) {
let this = &*self.shared;
(SimCx::with_in_node(this.node, |cx| {
let task = abort_local(
this.id,
cx.queue.borrow_mut().as_mut().unwrap(),
&mut cx.context.borrow_mut().as_mut().unwrap().executor.tasks,
);
drop(task);
}));
}
}
impl<T> Future for TaskHandle<T> {
type Output = Result<T, TaskAborted>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let this = self.project();
match this.result.poll(cx) {
Poll::Ready(x) => Poll::Ready(x.map_err(|_| TaskAborted(()))),
Poll::Pending => Poll::Pending,
}
}
}
pub fn spawn<F: Future + 'static>(future: F) -> TaskHandle<F::Output> {
spawn_task_on_node(NodeId::current(), future)
}
pub(crate) fn stop_node(node: NodeId, is_final: bool) {
SimCx::with_in_node(node, |cx| {
let was_running = cx.with_cx(|cxl| {
let node = &mut cxl.executor.nodes[node.to_index()];
match node.run_level {
NodeRunLevel::Running => {
node.run_level = if is_final {
NodeRunLevel::FinalStopped
} else {
NodeRunLevel::Stopped
};
true
}
NodeRunLevel::FinalStopped | NodeRunLevel::Stopped => {
if is_final {
node.run_level = NodeRunLevel::FinalStopped;
}
false
}
}
});
if !was_running {
return;
}
#[cfg(feature = "emit-tracing")]
let span = tracing::info_span!("stop-node", node = node.tv());
#[cfg(feature = "emit-tracing")]
let _guard = span.enter();
let task_ids = {
cx.with_cx(|context| {
let node = &mut context.executor.nodes[node.to_index()];
node.tasks.iter().copied().collect::<Vec<Id>>()
})
};
for &task in &task_ids {
drop(cx.with_cx(|cxl| {
abort_local(
task,
cx.queue.borrow_mut().as_mut().unwrap(),
&mut cxl.executor.tasks,
)
}))
}
for_all_simulators(cx, false, |x| x.stop_node());
});
}
pub(crate) fn start_node(node: NodeId) {
SimCx::with_in_node(node, |cx| {
let was_running = cx.with_cx(|cxl| {
let node = &mut cxl.executor.nodes[node.to_index()];
match node.run_level {
NodeRunLevel::Running => true,
NodeRunLevel::FinalStopped => {
panic!("cannot start node because simulation is shutting down.")
}
NodeRunLevel::Stopped => {
node.run_level = NodeRunLevel::Running;
false
}
}
});
if was_running {
return;
}
#[cfg(feature = "emit-tracing")]
let span = tracing::info_span!("start-node", node = node.tv());
#[cfg(feature = "emit-tracing")]
let _guard = span.enter();
for_all_simulators(cx, false, |x| x.start_node());
});
}
pub fn stop_simulation() {
SimCxl::with(|cx| cx.executor.final_stopped = true);
for node in NodeId::all() {
stop_node(node, true);
}
}