use std::{fmt, rc::Rc, time::Duration};
use futures::future::LocalBoxFuture;
use lenso_kernel::RuntimeDriver;
#[derive(Clone)]
pub struct NativeHostClock {
now: Rc<dyn Fn() -> Duration>,
sleep_until: Rc<dyn Fn(Duration) -> LocalBoxFuture<'static, ()>>,
}
impl NativeHostClock {
pub fn from_driver<D: RuntimeDriver>(driver: D) -> Self {
let sleep_driver = driver.clone();
Self {
now: Rc::new(move || driver.now()),
sleep_until: Rc::new(move |deadline| sleep_driver.sleep_until(deadline)),
}
}
pub fn now(&self) -> Duration {
(self.now)()
}
pub fn sleep_until(&self, deadline: Duration) -> LocalBoxFuture<'static, ()> {
(self.sleep_until)(deadline)
}
}
impl fmt::Debug for NativeHostClock {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("NativeHostClock")
.finish_non_exhaustive()
}
}
#[cfg(test)]
mod tests {
use super::*;
use futures::{future::LocalBoxFuture, task::SpawnError};
use lenso_kernel::{DriverTask, LocalTask};
use std::cell::Cell;
#[derive(Clone)]
struct ControlledDriver(Rc<Cell<Duration>>);
impl RuntimeDriver for ControlledDriver {
fn now(&self) -> Duration {
self.0.get()
}
fn sleep_until(&self, _deadline: Duration) -> LocalBoxFuture<'static, ()> {
Box::pin(async {})
}
fn yield_now(&self) -> LocalBoxFuture<'static, ()> {
Box::pin(async {})
}
fn spawn_local(&self, _task: LocalTask) -> Result<DriverTask, SpawnError> {
Err(SpawnError::shutdown())
}
fn shutdown_requested(&self) -> bool {
false
}
}
#[test]
fn owner_clock_and_driver_share_one_controlled_domain() {
let driver = ControlledDriver(Rc::new(Cell::new(Duration::from_secs(7))));
let owner = NativeHostClock::from_driver(driver.clone());
let next_generation = owner.clone();
assert_eq!(owner.now(), driver.now());
driver.0.set(Duration::from_secs(23));
assert_eq!(owner.now(), Duration::from_secs(23));
assert_eq!(next_generation.now(), driver.now());
}
#[derive(Default)]
struct TimerState {
now: Cell<Duration>,
deadline: Cell<Option<Duration>>,
polls: Cell<usize>,
sleeps: Cell<usize>,
live: Cell<usize>,
spawns: Cell<usize>,
waker: std::cell::RefCell<Option<std::task::Waker>>,
}
#[derive(Clone, Default)]
struct TimerDriver(Rc<TimerState>);
impl TimerDriver {
fn advance_to(&self, now: Duration) {
self.0.now.set(now);
if self
.0
.deadline
.get()
.is_some_and(|deadline| now >= deadline)
{
let waker = self.0.waker.borrow_mut().take();
if let Some(waker) = waker {
waker.wake();
}
}
}
}
struct TimerWait(Rc<TimerState>);
impl Drop for TimerWait {
fn drop(&mut self) {
self.0.live.set(self.0.live.get() - 1);
self.0.deadline.set(None);
self.0.waker.borrow_mut().take();
}
}
impl RuntimeDriver for TimerDriver {
fn now(&self) -> Duration {
self.0.now.get()
}
fn sleep_until(&self, deadline: Duration) -> LocalBoxFuture<'static, ()> {
assert_eq!(self.0.live.get(), 0);
self.0.live.set(1);
self.0.sleeps.set(self.0.sleeps.get() + 1);
self.0.deadline.set(Some(deadline));
let wait = TimerWait(self.0.clone());
Box::pin(futures::future::poll_fn(move |cx| {
wait.0.polls.set(wait.0.polls.get() + 1);
if wait.0.now.get() >= deadline {
std::task::Poll::Ready(())
} else {
*wait.0.waker.borrow_mut() = Some(cx.waker().clone());
std::task::Poll::Pending
}
}))
}
fn yield_now(&self) -> LocalBoxFuture<'static, ()> {
Box::pin(async {})
}
fn spawn_local(&self, _task: LocalTask) -> Result<DriverTask, SpawnError> {
self.0.spawns.set(self.0.spawns.get() + 1);
Err(SpawnError::shutdown())
}
fn shutdown_requested(&self) -> bool {
false
}
}
#[derive(Default)]
struct WakeCount(std::sync::atomic::AtomicUsize);
impl std::task::Wake for WakeCount {
fn wake(self: std::sync::Arc<Self>) {
self.0.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
}
}
#[test]
fn owner_sleep_uses_selected_driver_and_is_woken_at_its_deadline() {
let driver = TimerDriver::default();
driver.advance_to(Duration::from_secs(7));
let clock = NativeHostClock::from_driver(driver.clone());
let wake = std::sync::Arc::new(WakeCount::default());
let waker = std::task::Waker::from(wake.clone());
let mut cx = std::task::Context::from_waker(&waker);
let deadline = Duration::from_secs(23);
let mut wait = clock.clone().sleep_until(deadline);
drop(clock);
assert_eq!(driver.0.deadline.get(), Some(deadline));
assert_eq!(driver.0.sleeps.get(), 1);
assert!(wait.as_mut().poll(&mut cx).is_pending());
driver.advance_to(Duration::from_secs(22));
assert_eq!(wake.0.load(std::sync::atomic::Ordering::SeqCst), 0);
assert!(wait.as_mut().poll(&mut cx).is_pending());
driver.advance_to(deadline);
assert_eq!(wake.0.load(std::sync::atomic::Ordering::SeqCst), 1);
assert!(wait.as_mut().poll(&mut cx).is_ready());
drop(wait);
assert_eq!(driver.0.live.get(), 0);
assert_eq!(driver.0.spawns.get(), 0);
}
#[test]
fn dropping_owner_sleep_drops_driver_wait_without_polling_or_spawning() {
let driver = TimerDriver::default();
let clock = NativeHostClock::from_driver(driver.clone());
let wake = std::sync::Arc::new(WakeCount::default());
let waker = std::task::Waker::from(wake.clone());
let mut cx = std::task::Context::from_waker(&waker);
let deadline = Duration::from_secs(23);
let mut wait = clock.sleep_until(deadline);
assert!(wait.as_mut().poll(&mut cx).is_pending());
let polls = driver.0.polls.get();
drop(wait);
assert_eq!(driver.0.live.get(), 0);
assert!(driver.0.waker.borrow().is_none());
assert_eq!(driver.0.deadline.get(), None);
driver.advance_to(deadline);
assert_eq!(wake.0.load(std::sync::atomic::Ordering::SeqCst), 0);
assert_eq!(driver.0.polls.get(), polls);
assert_eq!(driver.0.spawns.get(), 0);
}
}