lenso_native_adapter/
host_clock.rs1use std::{fmt, rc::Rc, time::Duration};
3
4use futures::future::LocalBoxFuture;
5use lenso_kernel::RuntimeDriver;
6
7#[derive(Clone)]
9pub struct NativeHostClock {
10 now: Rc<dyn Fn() -> Duration>,
11 sleep_until: Rc<dyn Fn(Duration) -> LocalBoxFuture<'static, ()>>,
12}
13
14impl NativeHostClock {
15 pub fn from_driver<D: RuntimeDriver>(driver: D) -> Self {
17 let sleep_driver = driver.clone();
18 Self {
19 now: Rc::new(move || driver.now()),
20 sleep_until: Rc::new(move |deadline| sleep_driver.sleep_until(deadline)),
21 }
22 }
23
24 pub fn now(&self) -> Duration {
26 (self.now)()
27 }
28
29 pub fn sleep_until(&self, deadline: Duration) -> LocalBoxFuture<'static, ()> {
35 (self.sleep_until)(deadline)
36 }
37}
38
39impl fmt::Debug for NativeHostClock {
40 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
41 formatter
42 .debug_struct("NativeHostClock")
43 .finish_non_exhaustive()
44 }
45}
46
47#[cfg(test)]
48mod tests {
49 use super::*;
50 use futures::{future::LocalBoxFuture, task::SpawnError};
51 use lenso_kernel::{DriverTask, LocalTask};
52 use std::cell::Cell;
53
54 #[derive(Clone)]
55 struct ControlledDriver(Rc<Cell<Duration>>);
56
57 impl RuntimeDriver for ControlledDriver {
58 fn now(&self) -> Duration {
59 self.0.get()
60 }
61 fn sleep_until(&self, _deadline: Duration) -> LocalBoxFuture<'static, ()> {
62 Box::pin(async {})
63 }
64 fn yield_now(&self) -> LocalBoxFuture<'static, ()> {
65 Box::pin(async {})
66 }
67 fn spawn_local(&self, _task: LocalTask) -> Result<DriverTask, SpawnError> {
68 Err(SpawnError::shutdown())
69 }
70 fn shutdown_requested(&self) -> bool {
71 false
72 }
73 }
74
75 #[test]
76 fn owner_clock_and_driver_share_one_controlled_domain() {
77 let driver = ControlledDriver(Rc::new(Cell::new(Duration::from_secs(7))));
78 let owner = NativeHostClock::from_driver(driver.clone());
79 let next_generation = owner.clone();
80 assert_eq!(owner.now(), driver.now());
81 driver.0.set(Duration::from_secs(23));
82 assert_eq!(owner.now(), Duration::from_secs(23));
83 assert_eq!(next_generation.now(), driver.now());
84 }
85 #[derive(Default)]
86 struct TimerState {
87 now: Cell<Duration>,
88 deadline: Cell<Option<Duration>>,
89 polls: Cell<usize>,
90 sleeps: Cell<usize>,
91 live: Cell<usize>,
92 spawns: Cell<usize>,
93 waker: std::cell::RefCell<Option<std::task::Waker>>,
94 }
95
96 #[derive(Clone, Default)]
97 struct TimerDriver(Rc<TimerState>);
98
99 impl TimerDriver {
100 fn advance_to(&self, now: Duration) {
101 self.0.now.set(now);
102 if self
103 .0
104 .deadline
105 .get()
106 .is_some_and(|deadline| now >= deadline)
107 {
108 let waker = self.0.waker.borrow_mut().take();
109 if let Some(waker) = waker {
110 waker.wake();
111 }
112 }
113 }
114 }
115
116 struct TimerWait(Rc<TimerState>);
117
118 impl Drop for TimerWait {
119 fn drop(&mut self) {
120 self.0.live.set(self.0.live.get() - 1);
121 self.0.deadline.set(None);
122 self.0.waker.borrow_mut().take();
123 }
124 }
125
126 impl RuntimeDriver for TimerDriver {
127 fn now(&self) -> Duration {
128 self.0.now.get()
129 }
130 fn sleep_until(&self, deadline: Duration) -> LocalBoxFuture<'static, ()> {
131 assert_eq!(self.0.live.get(), 0);
132 self.0.live.set(1);
133 self.0.sleeps.set(self.0.sleeps.get() + 1);
134 self.0.deadline.set(Some(deadline));
135 let wait = TimerWait(self.0.clone());
136 Box::pin(futures::future::poll_fn(move |cx| {
137 wait.0.polls.set(wait.0.polls.get() + 1);
138 if wait.0.now.get() >= deadline {
139 std::task::Poll::Ready(())
140 } else {
141 *wait.0.waker.borrow_mut() = Some(cx.waker().clone());
142 std::task::Poll::Pending
143 }
144 }))
145 }
146 fn yield_now(&self) -> LocalBoxFuture<'static, ()> {
147 Box::pin(async {})
148 }
149 fn spawn_local(&self, _task: LocalTask) -> Result<DriverTask, SpawnError> {
150 self.0.spawns.set(self.0.spawns.get() + 1);
151 Err(SpawnError::shutdown())
152 }
153 fn shutdown_requested(&self) -> bool {
154 false
155 }
156 }
157
158 #[derive(Default)]
159 struct WakeCount(std::sync::atomic::AtomicUsize);
160
161 impl std::task::Wake for WakeCount {
162 fn wake(self: std::sync::Arc<Self>) {
163 self.0.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
164 }
165 }
166
167 #[test]
168 fn owner_sleep_uses_selected_driver_and_is_woken_at_its_deadline() {
169 let driver = TimerDriver::default();
170 driver.advance_to(Duration::from_secs(7));
171 let clock = NativeHostClock::from_driver(driver.clone());
172 let wake = std::sync::Arc::new(WakeCount::default());
173 let waker = std::task::Waker::from(wake.clone());
174 let mut cx = std::task::Context::from_waker(&waker);
175 let deadline = Duration::from_secs(23);
176 let mut wait = clock.clone().sleep_until(deadline);
177 drop(clock);
178 assert_eq!(driver.0.deadline.get(), Some(deadline));
179 assert_eq!(driver.0.sleeps.get(), 1);
180 assert!(wait.as_mut().poll(&mut cx).is_pending());
181 driver.advance_to(Duration::from_secs(22));
182 assert_eq!(wake.0.load(std::sync::atomic::Ordering::SeqCst), 0);
183 assert!(wait.as_mut().poll(&mut cx).is_pending());
184 driver.advance_to(deadline);
185 assert_eq!(wake.0.load(std::sync::atomic::Ordering::SeqCst), 1);
186 assert!(wait.as_mut().poll(&mut cx).is_ready());
187 drop(wait);
188 assert_eq!(driver.0.live.get(), 0);
189 assert_eq!(driver.0.spawns.get(), 0);
190 }
191
192 #[test]
193 fn dropping_owner_sleep_drops_driver_wait_without_polling_or_spawning() {
194 let driver = TimerDriver::default();
195 let clock = NativeHostClock::from_driver(driver.clone());
196 let wake = std::sync::Arc::new(WakeCount::default());
197 let waker = std::task::Waker::from(wake.clone());
198 let mut cx = std::task::Context::from_waker(&waker);
199 let deadline = Duration::from_secs(23);
200 let mut wait = clock.sleep_until(deadline);
201 assert!(wait.as_mut().poll(&mut cx).is_pending());
202 let polls = driver.0.polls.get();
203 drop(wait);
204 assert_eq!(driver.0.live.get(), 0);
205 assert!(driver.0.waker.borrow().is_none());
206 assert_eq!(driver.0.deadline.get(), None);
207 driver.advance_to(deadline);
208 assert_eq!(wake.0.load(std::sync::atomic::Ordering::SeqCst), 0);
209 assert_eq!(driver.0.polls.get(), polls);
210 assert_eq!(driver.0.spawns.get(), 0);
211 }
212}