use super::{
AsyncInterval, AsyncTcpListener, AsyncTcpStream, AsyncUdpSocket, JoinHandle, RecvMeta, Runtime,
Transmit,
};
use std::collections::{HashMap, VecDeque};
use std::future::Future;
use std::io;
use std::io::IoSliceMut;
use std::net::SocketAddr;
use std::pin::Pin;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex};
use std::task::{Context, Poll, Waker};
use std::time::{Duration, Instant};
#[derive(Debug)]
pub struct VirtualClock {
state: Mutex<ClockState>,
base: Instant,
}
impl Default for VirtualClock {
fn default() -> Self {
Self {
state: ClockState::default().into(),
base: Instant::now(),
}
}
}
#[derive(Debug, Default)]
struct ClockState {
time_elapsed: Duration,
timers: HashMap<u64, Timer>,
next_id: u64,
}
#[derive(Debug)]
struct Timer {
deadline: Duration,
waker: Option<Waker>,
}
impl VirtualClock {
pub fn new() -> Self {
Self::default()
}
pub fn elapsed(&self) -> Duration {
self.state.lock().expect("clock poisoned").time_elapsed
}
pub fn now(&self) -> Instant {
self.base + self.elapsed()
}
pub fn advance(&self, delta: Duration) {
let wakers = {
let mut state = self.state.lock().expect("clock poisoned");
state.time_elapsed += delta;
let now = state.time_elapsed;
let due: Vec<u64> = state
.timers
.iter()
.filter(|(_, t)| t.deadline <= now)
.map(|(id, _)| *id)
.collect();
due.into_iter()
.filter_map(|id| state.timers.remove(&id).and_then(|t| t.waker))
.collect::<Vec<_>>()
};
for waker in wakers {
waker.wake();
}
}
pub fn pending_timers(&self) -> usize {
self.state.lock().expect("clock poisoned").timers.len()
}
fn register(&self, delay: Duration) -> Option<u64> {
let mut state = self.state.lock().expect("clock poisoned");
if delay.is_zero() {
return None;
}
let deadline = state.time_elapsed + delay;
let id = state.next_id;
state.next_id += 1;
state.timers.insert(
id,
Timer {
deadline,
waker: None,
},
);
Some(id)
}
fn poll_timer(&self, id: u64, cx: &mut Context<'_>) -> Poll<()> {
let mut state = self.state.lock().expect("clock poisoned");
match state.timers.get_mut(&id) {
None => Poll::Ready(()),
Some(timer) => {
timer.waker = Some(cx.waker().clone());
Poll::Pending
}
}
}
fn cancel(&self, id: u64) {
self.state
.lock()
.expect("clock poisoned")
.timers
.remove(&id);
}
}
struct Sleep {
clock: Arc<VirtualClock>,
id: Option<u64>,
}
impl Future for Sleep {
type Output = ();
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<()> {
match self.id {
None => Poll::Ready(()),
Some(id) => match self.clock.poll_timer(id, cx) {
Poll::Ready(()) => {
self.id = None;
Poll::Ready(())
}
Poll::Pending => Poll::Pending,
},
}
}
}
impl Drop for Sleep {
fn drop(&mut self) {
if let Some(id) = self.id {
self.clock.cancel(id);
}
}
}
struct MockInterval {
clock: Arc<VirtualClock>,
period: Duration,
first: bool,
}
impl AsyncInterval for MockInterval {
fn tick(&mut self) -> Pin<Box<dyn Future<Output = ()> + Send + '_>> {
if self.first {
self.first = false;
return Box::pin(std::future::ready(()));
}
let id = self.clock.register(self.period);
let clock = Arc::clone(&self.clock);
Box::pin(Sleep { clock, id })
}
}
struct MockJoinHandle {
abort: futures::future::AbortHandle,
finished: Arc<AtomicBool>,
}
impl super::JoinHandle for MockJoinHandle {
fn detach(&self) {
}
fn abort(&self) {
self.abort.abort();
}
fn is_finished(&self) -> bool {
self.finished.load(Ordering::SeqCst)
}
}
#[derive(Debug, Default)]
pub struct MockUDPNetwork {
inboxes: Mutex<HashMap<SocketAddr, Inbox>>,
}
#[derive(Debug, Default)]
struct Inbox {
packets: VecDeque<(SocketAddr, Vec<u8>)>,
waker: Option<Waker>,
}
impl MockUDPNetwork {
pub fn new() -> Self {
Self::default()
}
pub fn bind(self: &Arc<Self>, addr: SocketAddr) -> io::Result<Arc<MockUdpSocket>> {
let mut inboxes = self.inboxes.lock().expect("network poisoned");
if inboxes.contains_key(&addr) {
return Err(io::Error::new(
io::ErrorKind::AddrInUse,
format!("mock network: {addr} is already bound"),
));
}
inboxes.insert(addr, Inbox::default());
Ok(Arc::new(MockUdpSocket {
local_addr: addr,
udp_network: Arc::clone(self),
}))
}
fn deliver(&self, from: SocketAddr, to: SocketAddr, payload: &[u8]) {
let waker = {
let mut inboxes = self.inboxes.lock().expect("network poisoned");
let Some(inbox) = inboxes.get_mut(&to) else {
return; };
inbox.packets.push_back((from, payload.to_vec()));
inbox.waker.take()
};
if let Some(waker) = waker {
waker.wake();
}
}
}
#[derive(Debug)]
pub struct MockUdpSocket {
local_addr: SocketAddr,
udp_network: Arc<MockUDPNetwork>,
}
impl Drop for MockUdpSocket {
fn drop(&mut self) {
self.udp_network
.inboxes
.lock()
.expect("udp network poisoned")
.remove(&self.local_addr);
}
}
impl AsyncUdpSocket for MockUdpSocket {
fn local_addr(&self) -> io::Result<SocketAddr> {
Ok(self.local_addr)
}
fn poll_send(&self, _cx: &mut Context<'_>, transmit: &Transmit<'_>) -> Poll<io::Result<usize>> {
match transmit.segment_size {
Some(size) if size > 0 => {
for chunk in transmit.contents.chunks(size) {
self.udp_network
.deliver(self.local_addr, transmit.destination, chunk);
}
}
_ => self
.udp_network
.deliver(self.local_addr, transmit.destination, transmit.contents),
}
Poll::Ready(Ok(transmit.contents.len()))
}
fn poll_recv(
&self,
cx: &mut Context<'_>,
bufs: &mut [IoSliceMut<'_>],
meta: &mut [RecvMeta],
) -> Poll<io::Result<usize>> {
let capacity = bufs.len().min(meta.len());
if capacity == 0 {
return Poll::Ready(Ok(0));
}
let mut inboxes = self.udp_network.inboxes.lock().expect("network poisoned");
let inbox = inboxes.get_mut(&self.local_addr).ok_or_else(|| {
io::Error::new(io::ErrorKind::NotConnected, "mock socket is not bound")
})?;
if inbox.packets.is_empty() {
inbox.waker = Some(cx.waker().clone());
return Poll::Pending;
}
let mut n = 0;
while n < capacity {
let Some((from, payload)) = inbox.packets.pop_front() else {
break;
};
let len = payload.len().min(bufs[n].len());
bufs[n][..len].copy_from_slice(&payload[..len]);
let mut m = RecvMeta::default();
m.addr = from;
m.len = len;
m.stride = len.max(1);
m.dst_ip = Some(self.local_addr.ip());
meta[n] = m;
n += 1;
}
Poll::Ready(Ok(n))
}
}
#[derive(Debug, Default)]
pub struct MockRuntime {
clock: Arc<VirtualClock>,
udp_network: Arc<MockUDPNetwork>,
}
impl MockRuntime {
pub fn new() -> Self {
Self::default()
}
pub fn with_network(network: Arc<MockUDPNetwork>) -> Self {
Self {
clock: Arc::new(VirtualClock::new()),
udp_network: network,
}
}
pub fn clock(&self) -> Arc<VirtualClock> {
Arc::clone(&self.clock)
}
pub fn network(&self) -> Arc<MockUDPNetwork> {
Arc::clone(&self.udp_network)
}
}
impl Runtime for MockRuntime {
fn now(&self) -> Instant {
self.clock.now()
}
fn spawn(&self, future: Pin<Box<dyn Future<Output = ()> + Send>>) -> Box<dyn JoinHandle> {
let (abortable, abort) = futures::future::abortable(future);
let finished = Arc::new(AtomicBool::new(false));
let done = Arc::clone(&finished);
std::thread::spawn(move || {
let _ = futures::executor::block_on(abortable);
done.store(true, Ordering::SeqCst);
});
Box::new(MockJoinHandle { abort, finished })
}
fn wrap_udp_socket(&self, socket: std::net::UdpSocket) -> io::Result<Arc<dyn AsyncUdpSocket>> {
let addr = socket.local_addr()?;
drop(socket);
Ok(self.udp_network.bind(addr)? as Arc<dyn AsyncUdpSocket>)
}
fn wrap_tcp_listener(
&self,
_listener: std::net::TcpListener,
) -> io::Result<Arc<dyn AsyncTcpListener>> {
Err(unsupported("wrap_tcp_listener"))
}
fn connect_tcp<'a>(
&'a self,
_remote_addr: SocketAddr,
) -> Pin<Box<dyn Future<Output = io::Result<Arc<dyn AsyncTcpStream>>> + Send + 'a>> {
Box::pin(async move { Err(unsupported("connect_tcp")) })
}
fn resolve_host<'a>(
&'a self,
host: &'a str,
) -> Pin<Box<dyn Future<Output = io::Result<Vec<SocketAddr>>> + Send + 'a>> {
Box::pin(async move {
match host.parse::<SocketAddr>() {
Ok(addr) => Ok(vec![addr]),
Err(_) => Err(unsupported("resolve_host (non-literal address)")),
}
})
}
fn sleep(&self, duration: Duration) -> Pin<Box<dyn Future<Output = ()> + Send + 'static>> {
let id = self.clock.register(duration);
Box::pin(Sleep {
clock: Arc::clone(&self.clock),
id,
})
}
fn interval(&self, period: Duration) -> Box<dyn AsyncInterval> {
Box::new(MockInterval {
clock: Arc::clone(&self.clock),
period,
first: true,
})
}
fn block_on(&self, future: Pin<Box<dyn Future<Output = ()> + '_>>) {
futures::executor::block_on(future);
}
fn name(&self) -> &'static str {
"mock"
}
}
fn unsupported(what: &str) -> io::Error {
io::Error::new(
io::ErrorKind::Unsupported,
format!("MockRuntime does not implement {what}; it covers timers and tasks only"),
)
}
#[cfg(test)]
mod tests {
use super::*;
use futures::FutureExt;
#[test]
fn sleep_only_completes_when_time_advances() {
let rt = MockRuntime::new();
let clock = rt.clock();
let mut sleep = rt.sleep(Duration::from_secs(10));
assert!(sleep.as_mut().now_or_never().is_none(), "must not be ready");
clock.advance(Duration::from_secs(9));
assert!(
sleep.as_mut().now_or_never().is_none(),
"9s < 10s deadline, still pending"
);
clock.advance(Duration::from_secs(1));
assert!(sleep.as_mut().now_or_never().is_some(), "deadline reached");
}
#[test]
fn advance_is_instant_and_cumulative() {
let rt = MockRuntime::new();
let clock = rt.clock();
assert_eq!(clock.elapsed(), Duration::ZERO);
clock.advance(Duration::from_millis(250));
clock.advance(Duration::from_millis(750));
assert_eq!(clock.elapsed(), Duration::from_secs(1));
}
#[test]
fn zero_duration_sleep_is_immediately_ready() {
let rt = MockRuntime::new();
assert!(rt.sleep(Duration::ZERO).now_or_never().is_some());
}
#[test]
fn dropping_a_sleep_cancels_its_timer() {
let rt = MockRuntime::new();
let clock = rt.clock();
let sleep = rt.sleep(Duration::from_secs(5));
assert_eq!(clock.pending_timers(), 1);
drop(sleep);
assert_eq!(clock.pending_timers(), 0, "timer must not leak");
}
#[test]
fn interval_first_tick_is_immediate_then_gated_by_clock() {
let rt = MockRuntime::new();
let clock = rt.clock();
let mut interval = rt.interval(Duration::from_secs(1));
assert!(
interval.tick().now_or_never().is_some(),
"first tick fires immediately"
);
let mut second = interval.tick();
assert!(second.as_mut().now_or_never().is_none());
clock.advance(Duration::from_secs(1));
assert!(second.as_mut().now_or_never().is_some());
}
#[test]
fn independent_runtimes_have_independent_clocks() {
let a = MockRuntime::new();
let b = MockRuntime::new();
a.clock().advance(Duration::from_secs(5));
assert_eq!(a.clock().elapsed(), Duration::from_secs(5));
assert_eq!(b.clock().elapsed(), Duration::ZERO);
}
#[test]
fn timeout_helper_works_on_the_mock() {
let rt = MockRuntime::new();
let clock = rt.clock();
let mut fut = Box::pin(super::super::timeout(
&rt,
Duration::from_secs(3),
std::future::pending::<()>(),
));
assert!(fut.as_mut().now_or_never().is_none());
clock.advance(Duration::from_secs(3));
assert!(matches!(fut.as_mut().now_or_never(), Some(Err(_))));
}
#[test]
fn tcp_operations_report_unsupported() {
let rt = MockRuntime::new();
if let Ok(listener) = std::net::TcpListener::bind("127.0.0.1:0") {
let err = rt.wrap_tcp_listener(listener).unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::Unsupported);
}
let err = futures::executor::block_on(
rt.connect_tcp("127.0.0.1:1".parse().expect("literal addr")),
)
.unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::Unsupported);
}
#[test]
fn udp_datagrams_round_trip_in_memory() {
let network = Arc::new(MockUDPNetwork::new());
let a: SocketAddr = "127.0.0.1:4000".parse().expect("literal addr");
let b: SocketAddr = "127.0.0.1:4001".parse().expect("literal addr");
let sock_a = network.bind(a).expect("bind a");
let sock_b = network.bind(b).expect("bind b");
assert_eq!(
network.bind(a).unwrap_err().kind(),
io::ErrorKind::AddrInUse,
"a bound address cannot be bound twice"
);
let waker = futures::task::noop_waker();
let mut cx = Context::from_waker(&waker);
let mut buf = [0u8; 64];
let mut bufs = [IoSliceMut::new(&mut buf)];
let mut meta = [RecvMeta::default()];
assert!(sock_b.poll_recv(&mut cx, &mut bufs, &mut meta).is_pending());
let payload = b"hello";
let sent = sock_a.poll_send(
&mut cx,
&Transmit {
destination: b,
ecn: None,
contents: payload,
segment_size: None,
src_ip: None,
},
);
assert!(matches!(sent, Poll::Ready(Ok(n)) if n == payload.len()));
let got = sock_b.poll_recv(&mut cx, &mut bufs, &mut meta);
assert!(matches!(got, Poll::Ready(Ok(1))));
assert_eq!(meta[0].addr, a, "the source address is preserved");
assert_eq!(meta[0].len, payload.len());
assert!(meta[0].stride >= 1, "stride must never be zero");
assert_eq!(&buf[..payload.len()], payload);
}
#[test]
fn udp_send_to_unbound_address_is_dropped() {
let network = Arc::new(MockUDPNetwork::new());
let a: SocketAddr = "127.0.0.1:4002".parse().expect("literal addr");
let sock_a = network.bind(a).expect("bind a");
let waker = futures::task::noop_waker();
let mut cx = Context::from_waker(&waker);
let sent = sock_a.poll_send(
&mut cx,
&Transmit {
destination: "127.0.0.1:9999".parse().expect("literal addr"),
ecn: None,
contents: b"into the void",
segment_size: None,
src_ip: None,
},
);
assert!(matches!(sent, Poll::Ready(Ok(_))), "send still succeeds");
}
#[test]
fn resolve_host_accepts_literal_addresses() {
let rt = MockRuntime::new();
let addrs = futures::executor::block_on(rt.resolve_host("127.0.0.1:3478")).unwrap();
assert_eq!(addrs.len(), 1);
assert_eq!(addrs[0].port(), 3478);
}
}