Skip to main content

gwr_engine/
executor.rs

1// Copyright (c) 2023 Graphcore Ltd. All rights reserved.
2
3use std::cell::{Cell, RefCell};
4use std::future::Future;
5use std::mem;
6use std::pin::Pin;
7use std::rc::Rc;
8use std::task::{Context, Poll, RawWaker, RawWakerVTable, Waker};
9
10use gwr_track::entity::Entity;
11use rand::SeedableRng;
12use rand::rngs::StdRng;
13use rand::seq::SliceRandom;
14
15use crate::time::clock::Clock;
16use crate::time::simtime::SimTime;
17use crate::types::SimResult;
18
19fn no_op(_: *const ()) {}
20
21unsafe fn drop_task(data: *const ()) {
22    unsafe {
23        drop(Rc::from_raw(data as *const Task));
24    }
25}
26
27static VTABLE: RawWakerVTable = RawWakerVTable::new(clone_raw_waker, wake_task, no_op, drop_task);
28
29fn task_raw_waker(task: Rc<Task>) -> RawWaker {
30    let ptr = Rc::into_raw(task) as *const ();
31    RawWaker::new(ptr, &VTABLE)
32}
33
34fn waker_for_task(task: Rc<Task>) -> Waker {
35    unsafe { Waker::from_raw(task_raw_waker(task)) }
36}
37
38unsafe fn clone_raw_waker(data: *const ()) -> RawWaker {
39    unsafe {
40        // Tasks are always wrapped in a reference counter to allow them to be
41        // shared read-only. The input `data` pointer is borrowed — we must not
42        // decrement its refcount, so we mem::forget the reconstructed Rc rather
43        // than letting it drop.
44        let rc_task = Rc::from_raw(data as *const Task);
45        let clone = rc_task.clone();
46        mem::forget(rc_task);
47        let ptr = Rc::into_raw(clone) as *const ();
48        RawWaker::new(ptr, &VTABLE)
49    }
50}
51
52unsafe fn wake_task(data: *const ()) {
53    unsafe {
54        // Tasks are always wrapped in a reference counter to allow them to be
55        // shared read-only.
56        let rc_task = Rc::from_raw(data as *const Task);
57        let cloned = rc_task.clone();
58        rc_task.executor_state.new_tasks.borrow_mut().push(cloned);
59    }
60}
61
62struct Task {
63    future: RefCell<Option<Pin<Box<dyn Future<Output = SimResult>>>>>,
64    executor_state: Rc<ExecutorState>,
65}
66
67impl Task {
68    pub fn new(
69        future: impl Future<Output = SimResult> + 'static,
70        executor_state: Rc<ExecutorState>,
71    ) -> Task {
72        Task {
73            future: RefCell::new(Some(Box::pin(future))),
74            executor_state,
75        }
76    }
77
78    fn poll(&self, context: &mut Context) -> Poll<SimResult> {
79        let mut future_slot = self.future.borrow_mut();
80        let Some(future) = future_slot.as_mut() else {
81            return Poll::Ready(Ok(()));
82        };
83
84        let poll_result = future.as_mut().poll(context);
85        if poll_result.is_ready() {
86            future_slot.take();
87        }
88
89        poll_result
90    }
91}
92
93struct ExecutorState {
94    task_queue: RefCell<Vec<Rc<Task>>>,
95    new_tasks: RefCell<Vec<Rc<Task>>>,
96    time: RefCell<SimTime>,
97    randomize_task_order: Cell<bool>,
98    task_order_rng: RefCell<StdRng>,
99}
100
101impl ExecutorState {
102    pub fn new(top: &Rc<Entity>) -> Self {
103        Self {
104            task_queue: RefCell::new(Vec::new()),
105            new_tasks: RefCell::new(Vec::new()),
106            time: RefCell::new(SimTime::new(top)),
107            randomize_task_order: Cell::new(false),
108            task_order_rng: RefCell::new(StdRng::seed_from_u64(rand::random())),
109        }
110    }
111}
112
113/// Single-threaded executor
114///
115/// This is a thin-wrapper (using [`Rc`]) around the real executor, so that this
116/// struct can be cloned and passed around.
117///
118/// See the [module documentation] for more details.
119///
120/// [module documentation]: index.html
121#[derive(Clone)]
122pub struct Executor {
123    state: Rc<ExecutorState>,
124}
125
126impl Executor {
127    pub fn run(&self, finished: &Rc<RefCell<bool>>) -> SimResult {
128        loop {
129            self.step(finished)?;
130            if *finished.borrow() {
131                break;
132            }
133
134            if self.state.new_tasks.borrow().is_empty() {
135                if self.state.time.borrow().can_exit() {
136                    break;
137                }
138
139                if let Some(wakers) = self.state.time.borrow_mut().advance_time() {
140                    // No events left, advance time
141                    for task_waker in wakers.into_iter() {
142                        task_waker.waker.wake();
143                    }
144                } else {
145                    break;
146                }
147            }
148        }
149        Ok(())
150    }
151
152    pub fn step(&self, finished: &Rc<RefCell<bool>>) -> SimResult {
153        // Append new tasks created since the last step into the task queue
154        let mut task_queue = self.state.task_queue.borrow_mut();
155        task_queue.append(&mut self.state.new_tasks.borrow_mut());
156        if self.state.randomize_task_order.get() {
157            task_queue.shuffle(&mut *self.state.task_order_rng.borrow_mut());
158        }
159
160        // Loop over all tasks, polling them. If a task is not ready, add it to
161        // the pending tasks.
162        for task in task_queue.drain(..) {
163            if *finished.borrow() {
164                break;
165            }
166
167            // Dummy waker and context (not used as we poll all tasks)
168            let waker = waker_for_task(task.clone());
169            let mut context = Context::from_waker(&waker);
170
171            match task.poll(&mut context) {
172                Poll::Ready(Err(e)) => {
173                    // Error - return early
174                    return Err(e);
175                }
176                Poll::Ready(Ok(())) => {
177                    // Otherwise, drop task as it is complete
178                }
179                Poll::Pending => {
180                    // Task will have parked itself waiting somewhere
181                }
182            }
183        }
184        Ok(())
185    }
186
187    #[must_use]
188    pub fn get_clock(&self, freq_mhz: f64) -> Clock {
189        self.state.time.borrow_mut().get_clock(freq_mhz)
190    }
191
192    #[must_use]
193    pub fn time_now_ns(&self) -> f64 {
194        self.state.time.borrow().time_now_ns()
195    }
196
197    pub fn set_randomize_task_order(&self, randomize: bool) {
198        self.state.randomize_task_order.set(randomize);
199    }
200
201    pub fn set_task_order_seed(&self, seed: u64) {
202        *self.state.task_order_rng.borrow_mut() = StdRng::seed_from_u64(seed);
203    }
204}
205
206/// `Spawner` spawns new futures into the executor.
207#[derive(Clone)]
208pub struct Spawner {
209    state: Rc<ExecutorState>,
210}
211
212impl Spawner {
213    pub fn spawn(&self, future: impl Future<Output = SimResult> + 'static) {
214        self.state
215            .new_tasks
216            .borrow_mut()
217            .push(Rc::new(Task::new(future, self.state.clone())));
218    }
219}
220
221#[must_use]
222pub fn new_executor_and_spawner(top: &Rc<Entity>) -> (Executor, Spawner) {
223    let state = Rc::new(ExecutorState::new(top));
224    (
225        Executor {
226            state: state.clone(),
227        },
228        Spawner { state },
229    )
230}
231
232#[cfg(test)]
233mod tests {
234    use std::cell::RefCell;
235    use std::future::Future;
236    use std::pin::Pin;
237    use std::rc::Rc;
238    use std::task::{Context, Poll};
239
240    use futures::task::noop_waker;
241    use gwr_track::entity::toplevel;
242    use gwr_track::tracker::dev_null_tracker;
243
244    use super::*;
245    use crate::time::clock::TaskWaker;
246
247    struct PanicIfPolled;
248
249    impl Future for PanicIfPolled {
250        type Output = SimResult;
251
252        fn poll(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Self::Output> {
253            panic!("finished executor step polled a task");
254        }
255    }
256
257    #[test]
258    #[should_panic(expected = "finished executor step polled a task")]
259    fn panic_if_polled_panics_when_polled() {
260        let mut future = Box::pin(PanicIfPolled);
261        let waker = noop_waker();
262        let mut cx = Context::from_waker(&waker);
263
264        let _ = future.as_mut().poll(&mut cx);
265    }
266
267    #[test]
268    fn run_exits_when_time_cannot_advance() {
269        let tracker = dev_null_tracker();
270        let top = toplevel(&tracker, "top");
271        let (executor, _spawner) = new_executor_and_spawner(&top);
272        let clock = executor.get_clock(1000.0);
273
274        clock
275            .shared_state
276            .waiting
277            .borrow_mut()
278            .push(vec![TaskWaker {
279                id: 0,
280                waker: noop_waker(),
281                can_exit: false,
282            }]);
283
284        let finished = Rc::new(RefCell::new(false));
285
286        executor.run(&finished).unwrap();
287    }
288
289    #[test]
290    fn step_stops_polling_when_finished_is_set() {
291        let tracker = dev_null_tracker();
292        let top = toplevel(&tracker, "top");
293        let (executor, spawner) = new_executor_and_spawner(&top);
294
295        spawner.spawn(PanicIfPolled);
296
297        let finished = Rc::new(RefCell::new(true));
298
299        executor.step(&finished).unwrap();
300    }
301}