1use std::cell::{Cell, RefCell};
9use std::collections::{HashMap, VecDeque};
10use std::fmt;
11use std::future::Future;
12use std::panic::{catch_unwind, AssertUnwindSafe};
13use std::pin::Pin;
14use std::rc::Rc;
15use std::sync::{Arc, Mutex};
16use std::task::{Context, Poll, Wake, Waker};
17
18pub type TaskId = u64;
19
20#[derive(Copy, Clone, Debug, PartialEq, Eq)]
22pub enum TaskState {
23 Unstarted,
24 Scheduled,
25 Running,
26 Pending,
27 Finished,
28 Cancelled,
29 Failed,
30}
31
32#[derive(Debug, Clone)]
33pub enum TaskError {
34 Cancelled,
35 Panicked(String),
36 InvalidState,
38}
39
40impl fmt::Display for TaskError {
41 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
42 match self {
43 TaskError::Cancelled => write!(f, "task was cancelled"),
44 TaskError::Panicked(m) => write!(f, "task panicked: {m}"),
45 TaskError::InvalidState => write!(f, "task result not available"),
46 }
47 }
48}
49impl std::error::Error for TaskError {}
50
51struct WokenQueue(Mutex<VecDeque<TaskId>>);
58
59struct TaskWaker {
60 id: TaskId,
61 woken: Arc<WokenQueue>,
62}
63
64impl Wake for TaskWaker {
65 fn wake(self: Arc<Self>) {
66 self.woken.0.lock().unwrap().push_back(self.id);
67 }
68}
69
70struct TaskShared<T> {
76 state: Cell<TaskState>,
77 result: RefCell<Option<Result<T, TaskError>>>,
78 joiners: RefCell<Vec<Waker>>,
79}
80
81impl<T> TaskShared<T> {
82 fn complete(&self, r: Result<T, TaskError>) {
83 *self.result.borrow_mut() = Some(r);
84 for w in self.joiners.borrow_mut().drain(..) {
85 w.wake();
86 }
87 }
88}
89
90struct TaskEntry {
91 fut: Pin<Box<dyn Future<Output = ()>>>,
92 name: String,
93 state: TaskState, on_abort: Rc<dyn Fn(TaskError)>,
96 state_cell: Rc<dyn Fn(TaskState)>,
98}
99
100struct ExecInner {
101 tasks: RefCell<HashMap<TaskId, TaskEntry>>,
102 next_id: Cell<TaskId>,
103 run_queue: RefCell<VecDeque<TaskId>>,
104 woken: Arc<WokenQueue>,
105 running: Cell<bool>,
106 currently_polling: Cell<Option<TaskId>>,
107 cancel_pending: RefCell<Vec<TaskId>>,
109 failure_sink: RefCell<Option<Box<dyn Fn(&str)>>>,
112}
113
114#[derive(Clone)]
116pub struct Executor {
117 inner: Rc<ExecInner>,
118}
119
120thread_local! {
121 static CURRENT: RefCell<Option<Executor>> = const { RefCell::new(None) };
122}
123
124pub fn init() -> Executor {
126 let ex = Executor {
127 inner: Rc::new(ExecInner {
128 tasks: RefCell::new(HashMap::new()),
129 next_id: Cell::new(1),
130 run_queue: RefCell::new(VecDeque::new()),
131 woken: Arc::new(WokenQueue(Mutex::new(VecDeque::new()))),
132 running: Cell::new(false),
133 currently_polling: Cell::new(None),
134 cancel_pending: RefCell::new(Vec::new()),
135 failure_sink: RefCell::new(None),
136 }),
137 };
138 CURRENT.with(|c| *c.borrow_mut() = Some(ex.clone()));
139 ex
140}
141
142pub fn current() -> Executor {
144 CURRENT.with(|c| c.borrow().clone().expect("rustdv executor not initialized"))
145}
146
147pub fn spawn<F>(fut: F) -> TaskHandle<F::Output>
149where
150 F: Future + 'static,
151{
152 current().spawn_named(fut, None)
153}
154
155pub fn spawn_named<F>(fut: F, name: &str) -> TaskHandle<F::Output>
156where
157 F: Future + 'static,
158{
159 current().spawn_named(fut, Some(name))
160}
161
162impl Executor {
163 pub fn spawn_named<F>(&self, fut: F, name: Option<&str>) -> TaskHandle<F::Output>
164 where
165 F: Future + 'static,
166 {
167 let id = self.inner.next_id.get();
168 self.inner.next_id.set(id + 1);
169 let name = name.map(|s| s.to_string()).unwrap_or_else(|| format!("task_{id}"));
170
171 let shared = Rc::new(TaskShared::<F::Output> {
172 state: Cell::new(TaskState::Unstarted),
173 result: RefCell::new(None),
174 joiners: RefCell::new(Vec::new()),
175 });
176
177 let sh = shared.clone();
179 let wrapped = async move {
180 let out = fut.await;
181 sh.state.set(TaskState::Finished);
182 sh.complete(Ok(out));
183 };
184
185 let sh_abort = shared.clone();
186 let on_abort: Rc<dyn Fn(TaskError)> = Rc::new(move |e: TaskError| {
187 sh_abort.state.set(match e {
188 TaskError::Cancelled => TaskState::Cancelled,
189 _ => TaskState::Failed,
190 });
191 sh_abort.complete(Err(e));
192 });
193 let sh_state = shared.clone();
194 let state_cell: Rc<dyn Fn(TaskState)> = Rc::new(move |s| sh_state.state.set(s));
195
196 shared.state.set(TaskState::Scheduled);
197 self.inner.tasks.borrow_mut().insert(
198 id,
199 TaskEntry {
200 fut: Box::pin(wrapped),
201 name,
202 state: TaskState::Scheduled,
203 on_abort,
204 state_cell,
205 },
206 );
207 self.inner.run_queue.borrow_mut().push_back(id);
208
209 TaskHandle { id, shared, exec: self.clone() }
210 }
211
212 pub fn run_until_idle(&self) {
215 if self.inner.running.get() {
216 return; }
218 self.inner.running.set(true);
219 loop {
220 self.drain_woken();
221 let next = self.inner.run_queue.borrow_mut().pop_front();
222 let Some(id) = next else { break };
223 self.poll_task(id);
224 }
225 self.inner.running.set(false);
226 }
227
228 fn drain_woken(&self) {
229 let ids: Vec<TaskId> = self.inner.woken.0.lock().unwrap().drain(..).collect();
230 for id in ids {
231 let mut tasks = self.inner.tasks.borrow_mut();
232 if let Some(t) = tasks.get_mut(&id) {
233 if t.state == TaskState::Pending {
234 t.state = TaskState::Scheduled;
235 (t.state_cell)(TaskState::Scheduled);
236 self.inner.run_queue.borrow_mut().push_back(id);
237 }
238 }
239 }
240 }
241
242 fn poll_task(&self, id: TaskId) {
243 let (mut fut, on_abort, state_cell) = {
247 let mut tasks = self.inner.tasks.borrow_mut();
248 let Some(t) = tasks.get_mut(&id) else { return };
249 if t.state != TaskState::Scheduled {
250 return; }
252 t.state = TaskState::Running;
253 (t.state_cell)(TaskState::Running);
254 let fut = std::mem::replace(&mut t.fut, Box::pin(async {}));
256 (fut, t.on_abort.clone(), t.state_cell.clone())
257 };
258
259 let waker = Waker::from(Arc::new(TaskWaker { id, woken: self.inner.woken.clone() }));
260 let mut cx = Context::from_waker(&waker);
261
262 self.inner.currently_polling.set(Some(id));
263 let polled = catch_unwind(AssertUnwindSafe(|| fut.as_mut().poll(&mut cx)));
264 self.inner.currently_polling.set(None);
265
266 match polled {
267 Ok(Poll::Ready(())) => {
268 self.inner.tasks.borrow_mut().remove(&id);
270 }
271 Ok(Poll::Pending) => {
272 let mut tasks = self.inner.tasks.borrow_mut();
273 if let Some(t) = tasks.get_mut(&id) {
274 t.fut = fut; t.state = TaskState::Pending;
276 (t.state_cell)(TaskState::Pending);
277 }
278 drop(tasks);
279 let pending: Vec<TaskId> = self.inner.cancel_pending.borrow_mut().drain(..).collect();
281 for cid in pending {
282 self.cancel(cid);
283 }
284 }
285 Err(p) => {
286 let msg = panic_message(p);
287 let name = self
288 .inner
289 .tasks
290 .borrow()
291 .get(&id)
292 .map(|t| t.name.clone())
293 .unwrap_or_default();
294 self.inner.tasks.borrow_mut().remove(&id);
295 drop(fut); state_cell(TaskState::Failed);
297 on_abort(TaskError::Panicked(msg.clone()));
298 if let Some(sink) = self.inner.failure_sink.borrow().as_ref() {
299 sink(&format!("task '{name}' panicked: {msg}"));
300 } else {
301 eprintln!("rustdv: unhandled task panic in '{name}': {msg}");
302 }
303 }
304 }
305 }
306
307 pub fn cancel(&self, id: TaskId) {
310 if self.inner.currently_polling.get() == Some(id) {
311 self.inner.cancel_pending.borrow_mut().push(id);
312 return;
313 }
314 let entry = self.inner.tasks.borrow_mut().remove(&id);
315 if let Some(t) = entry {
316 (t.on_abort)(TaskError::Cancelled);
317 drop(t.fut);
318 }
319 }
320
321 pub fn cancel_after(&self, watermark: TaskId) {
325 let ids: Vec<TaskId> = self
326 .inner
327 .tasks
328 .borrow()
329 .keys()
330 .copied()
331 .filter(|&id| id >= watermark)
332 .collect();
333 for id in ids {
334 self.cancel(id);
335 }
336 }
337
338 pub fn watermark(&self) -> TaskId {
340 self.inner.next_id.get()
341 }
342
343 pub fn set_failure_sink(&self, f: Box<dyn Fn(&str)>) {
344 *self.inner.failure_sink.borrow_mut() = Some(f);
345 }
346
347 pub fn live_tasks(&self) -> usize {
348 self.inner.tasks.borrow().len()
349 }
350}
351
352fn panic_message(p: Box<dyn std::any::Any + Send>) -> String {
353 if let Some(s) = p.downcast_ref::<&str>() {
354 s.to_string()
355 } else if let Some(s) = p.downcast_ref::<String>() {
356 s.clone()
357 } else {
358 "panic (non-string payload)".to_string()
359 }
360}
361
362pub struct TaskHandle<T> {
369 id: TaskId,
370 shared: Rc<TaskShared<T>>,
371 exec: Executor,
372}
373
374impl<T> Clone for TaskHandle<T> {
375 fn clone(&self) -> Self {
376 TaskHandle { id: self.id, shared: self.shared.clone(), exec: self.exec.clone() }
377 }
378}
379
380impl<T> TaskHandle<T> {
381 pub fn id(&self) -> TaskId {
382 self.id
383 }
384
385 pub fn state(&self) -> TaskState {
386 self.shared.state.get()
387 }
388
389 pub fn done(&self) -> bool {
390 matches!(
391 self.state(),
392 TaskState::Finished | TaskState::Cancelled | TaskState::Failed
393 )
394 }
395
396 pub fn cancel(&self) {
398 self.exec.cancel(self.id);
399 }
400
401 pub fn result(&self) -> Result<T, TaskError> {
403 match self.shared.result.borrow_mut().take() {
404 Some(r) => r,
405 None => Err(TaskError::InvalidState),
406 }
407 }
408}
409
410impl<T> Future for TaskHandle<T> {
411 type Output = Result<T, TaskError>;
412
413 fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
414 if self.done() {
415 return Poll::Ready(self.result());
416 }
417 self.shared.joiners.borrow_mut().push(cx.waker().clone());
418 Poll::Pending
419 }
420}