#![allow(unsafe_code)]
use std::collections::HashMap;
use std::net::{TcpStream, UdpSocket};
#[cfg(windows)]
pub(crate) type Fd = std::os::windows::io::RawSocket;
#[cfg(not(windows))]
pub(crate) type Fd = std::os::fd::RawFd;
pub(crate) const WAKE_ID: usize = 0;
#[cfg(windows)]
pub(crate) fn ensure_high_resolution_timer() {
use std::sync::Once;
static INIT: Once = Once::new();
INIT.call_once(|| {
unsafe {
timeBeginPeriod(1);
}
});
}
#[cfg(windows)]
#[link(name = "winmm")]
extern "system" {
fn timeBeginPeriod(period: u32) -> u32;
}
pub(crate) fn fd_of(socket: &TcpStream) -> Fd {
#[cfg(windows)]
{
use std::os::windows::io::AsRawSocket;
socket.as_raw_socket()
}
#[cfg(not(windows))]
{
use std::os::fd::AsRawFd;
socket.as_raw_fd()
}
}
pub(crate) fn udp_fd_of(socket: &UdpSocket) -> Fd {
#[cfg(windows)]
{
use std::os::windows::io::AsRawSocket;
socket.as_raw_socket()
}
#[cfg(not(windows))]
{
use std::os::fd::AsRawFd;
socket.as_raw_fd()
}
}
#[cfg(windows)]
mod ws {
use super::Fd;
use std::os::windows::io::RawSocket;
pub(super) const FD_SETSIZE: usize = 64;
#[repr(C)]
pub(super) struct FdSet {
fd_count: u32,
fd_array: [RawSocket; FD_SETSIZE],
}
impl FdSet {
pub(super) fn new() -> Self {
Self {
fd_count: 0,
fd_array: [0; FD_SETSIZE],
}
}
pub(super) fn insert(&mut self, fd: Fd) {
if (self.fd_count as usize) < FD_SETSIZE {
self.fd_array[self.fd_count as usize] = fd;
self.fd_count += 1;
}
}
pub(super) fn ready(&self) -> &[RawSocket] {
&self.fd_array[..self.fd_count as usize]
}
}
#[repr(C)]
pub(super) struct TimeVal {
pub(super) tv_sec: i32,
pub(super) tv_usec: i32,
}
#[link(name = "ws2_32")]
extern "system" {
pub(super) fn select(
nfds: i32,
readfds: *mut FdSet,
writefds: *mut FdSet,
exceptfds: *mut FdSet,
timeout: *const TimeVal,
) -> i32;
}
}
#[cfg(not(windows))]
mod posix {
use super::Fd;
pub(super) const POLLIN: i16 = 0x001;
pub(super) const POLLOUT: i16 = 0x004;
pub(super) const POLLERR: i16 = 0x008;
pub(super) const POLLHUP: i16 = 0x010;
pub(super) const POLLNVAL: i16 = 0x020;
#[repr(C)]
pub(super) struct PollFd {
pub(super) fd: Fd,
pub(super) events: i16,
pub(super) revents: i16,
}
#[cfg(any(target_os = "linux", target_os = "android"))]
pub(super) type NfdsT = u64;
#[cfg(not(any(target_os = "linux", target_os = "android")))]
pub(super) type NfdsT = u32;
extern "C" {
pub(super) fn poll(fds: *mut PollFd, nfds: NfdsT, timeout: i32) -> i32;
}
}
pub(crate) struct Poller {
fds: Vec<(usize, Fd, bool)>,
index: HashMap<usize, usize>,
}
impl Poller {
pub(crate) fn new() -> Self {
#[cfg(windows)]
ensure_high_resolution_timer();
Self {
fds: Vec::new(),
index: HashMap::new(),
}
}
pub(crate) fn is_empty(&self) -> bool {
self.fds.is_empty()
}
pub(crate) fn register(&mut self, id: usize, fd: Fd, want_write: bool) {
self.unregister(id);
self.fds.push((id, fd, want_write));
self.index.insert(id, self.fds.len() - 1);
}
pub(crate) fn unregister(&mut self, id: usize) {
if let Some(&idx) = self.index.get(&id) {
self.fds.swap_remove(idx);
if idx < self.fds.len() {
self.index.insert(self.fds[idx].0, idx);
}
self.index.remove(&id);
}
}
pub(crate) fn clear(&mut self) {
self.fds.clear();
self.index.clear();
}
pub(crate) fn wait(
&mut self,
timeout_ms: i32,
wake: Option<Fd>,
) -> std::io::Result<Vec<usize>> {
if self.fds.is_empty() && wake.is_none() {
return Ok(Vec::new());
}
#[cfg(windows)]
{
self.wait_select(timeout_ms, wake)
}
#[cfg(not(windows))]
{
self.wait_poll(timeout_ms, wake)
}
}
#[cfg(windows)]
fn wait_select(&self, timeout_ms: i32, wake: Option<Fd>) -> std::io::Result<Vec<usize>> {
use ws::*;
let full_tv = TimeVal {
tv_sec: timeout_ms / 1000,
tv_usec: (timeout_ms % 1000) * 1000,
};
let zero_tv = TimeVal {
tv_sec: 0,
tv_usec: 0,
};
let mut ready = Vec::new();
let batch_cap = if wake.is_some() {
FD_SETSIZE - 1
} else {
FD_SETSIZE
};
let batches = if self.fds.is_empty() {
1
} else {
self.fds.len().div_ceil(batch_cap).max(1)
};
for b in 0..batches {
let start = b * batch_cap;
let end = core::cmp::min(start + batch_cap, self.fds.len());
let mut readset = FdSet::new();
let mut writeset = FdSet::new();
if let Some(w) = wake {
readset.insert(w);
}
let mut wake_ready = false;
for &(_, fd, ww) in &self.fds[start..end] {
if ww {
writeset.insert(fd);
} else {
readset.insert(fd);
}
}
let tv = if b == 0 { &full_tv } else { &zero_tv };
let n = unsafe { select(0, &mut readset, &mut writeset, std::ptr::null_mut(), tv) };
if n < 0 {
return Err(std::io::Error::last_os_error());
}
if n > 0 {
for &fd in readset.ready() {
if Some(fd) == wake {
wake_ready = true;
continue;
}
if let Some((id, _, _)) = self.fds[start..end].iter().find(|(_, f, _)| *f == fd)
{
ready.push(*id);
}
}
for &fd in writeset.ready() {
if let Some((id, _, _)) = self.fds[start..end].iter().find(|(_, f, _)| *f == fd)
{
ready.push(*id);
}
}
}
if wake_ready {
ready.push(WAKE_ID);
}
}
ready.sort_unstable();
ready.dedup();
Ok(ready)
}
#[cfg(not(windows))]
fn wait_poll(&self, timeout_ms: i32, wake: Option<Fd>) -> std::io::Result<Vec<usize>> {
use posix::*;
let mut pfds: Vec<PollFd> = self
.fds
.iter()
.map(|&(_, fd, ww)| PollFd {
fd,
events: if ww { POLLOUT } else { POLLIN },
revents: 0,
})
.collect();
let wake_idx = match wake {
Some(w) => {
pfds.push(PollFd {
fd: w,
events: POLLIN,
revents: 0,
});
Some(pfds.len() - 1)
}
None => None,
};
let n = unsafe { poll(pfds.as_mut_ptr(), pfds.len() as NfdsT, timeout_ms) };
if n < 0 {
return Err(std::io::Error::last_os_error());
}
if n == 0 {
return Ok(Vec::new());
}
let mut ready = Vec::new();
for (idx, pfd) in pfds.iter().enumerate() {
if Some(idx) == wake_idx {
if pfd.revents & (POLLIN | POLLERR | POLLHUP | POLLNVAL) != 0 {
ready.push(WAKE_ID);
}
continue;
}
let (_, _, ww) = self.fds[idx];
let expected = if ww { POLLOUT } else { POLLIN };
if pfd.revents & (expected | POLLERR | POLLHUP | POLLNVAL) != 0 {
ready.push(self.fds[idx].0);
}
}
ready.sort_unstable();
ready.dedup();
Ok(ready)
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::{Read as _, Write as _};
use std::net::{TcpListener, TcpStream};
#[test]
fn poll_reports_readable_socket() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let mut client = TcpStream::connect(addr).unwrap();
let (server, _) = listener.accept().unwrap();
let mut p = Poller::new();
let fd = fd_of(&server);
p.register(7, fd, false);
let ready = p.wait(50, None).unwrap();
assert!(ready.is_empty(), "unexpected ready: {ready:?}");
client.write_all(b"hi").unwrap();
let ready = p.wait(2000, None).unwrap();
assert_eq!(ready, vec![7]);
let mut b = [0u8; 8];
let mut s = &server;
let n = s.read(&mut b).unwrap();
assert_eq!(n, 2);
}
#[test]
fn unregister_stops_reporting() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let mut client = TcpStream::connect(addr).unwrap();
let (server, _) = listener.accept().unwrap();
let mut p = Poller::new();
let fd = fd_of(&server);
p.register(7, fd, false);
p.unregister(7);
assert!(p.is_empty());
client.write_all(b"hi").unwrap();
let ready = p.wait(100, None).unwrap();
assert!(ready.is_empty(), "unregistered socket reported: {ready:?}");
}
#[test]
fn a_closed_descriptor_cannot_wedge_a_wait() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let client = TcpStream::connect(addr).unwrap();
let (server, _) = listener.accept().unwrap();
let mut p = Poller::new();
let fd = fd_of(&server);
p.register(7, fd, false);
drop(server);
let started = std::time::Instant::now();
let outcome = p.wait(200, None);
assert!(
started.elapsed() < std::time::Duration::from_secs(2),
"the wait did not return: a closed descriptor wedged it"
);
match outcome {
Err(_) => {}
Ok(ids) => {
for id in &ids {
assert_eq!(*id, 7, "an unregistered descriptor was reported: {ids:?}");
}
}
}
drop(client);
}
#[test]
fn wake_descriptor_interrupts_wait() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let writer = TcpStream::connect(addr).unwrap();
let (wake_reader, _) = listener.accept().unwrap();
wake_reader.set_nonblocking(true).unwrap();
writer.set_nonblocking(true).unwrap();
let mut p = Poller::new();
let wfd = fd_of(&wake_reader);
let mut w: &TcpStream = &writer;
std::io::Write::write_all(&mut w, b"\x01").unwrap();
let started = std::time::Instant::now();
let ready = p.wait(10_000, Some(wfd)).unwrap();
assert_eq!(ready, vec![WAKE_ID]);
assert!(
started.elapsed() < std::time::Duration::from_secs(1),
"wake did not interrupt the poll"
);
}
#[test]
fn wake_interrupts_the_poll_promptly() {
use crate::courierust_server::event::{drain_wake, wake_nudge, wakeup_pair};
let (reader, writer) = wakeup_pair().unwrap();
let mut p = Poller::new();
let wfd = fd_of(&reader);
let mut samples: Vec<std::time::Duration> = Vec::with_capacity(100);
for i in 0..100 {
wake_nudge(&writer);
let started = std::time::Instant::now();
let ready = p.wait(1000, Some(wfd)).unwrap();
let elapsed = started.elapsed();
samples.push(elapsed);
assert!(ready.contains(&WAKE_ID), "wake {i} lost: ready={ready:?}");
drain_wake(&reader);
}
samples.sort_unstable();
assert!(
samples[95] < std::time::Duration::from_millis(50),
"p95 wake latency too high: {:#?}",
samples[95]
);
assert!(
samples[99] < std::time::Duration::from_millis(250),
"p100 wake latency too high (wake likely lost): {:#?}",
samples[99]
);
}
#[test]
fn wake_interrupts_already_blocked_wait() {
use crate::courierust_server::event::{drain_wake, wake_nudge, wakeup_pair};
use std::sync::Arc;
use std::time::Instant;
let (reader, writer) = wakeup_pair().unwrap();
let writer = Arc::new(writer);
let mut p = Poller::new();
let wfd = fd_of(&reader);
let mut samples: Vec<std::time::Duration> = Vec::with_capacity(100);
for _ in 0..100 {
let writer = writer.clone();
let nudger = std::thread::spawn(move || {
std::thread::sleep(std::time::Duration::from_micros(200));
wake_nudge(&writer);
});
let started = Instant::now();
let ready = p.wait(1000, Some(wfd)).unwrap();
let elapsed = started.elapsed();
samples.push(elapsed);
assert!(
ready.contains(&WAKE_ID),
"wake lost while wait was blocked: {ready:?}"
);
nudger.join().unwrap();
drain_wake(&reader);
}
samples.sort_unstable();
assert!(
samples[49] < std::time::Duration::from_millis(10),
"p50 blocked-wait wake latency too high: {:#?}",
samples[49]
);
assert!(
samples[95] < std::time::Duration::from_millis(100),
"p95 blocked-wait wake latency too high: {:#?}",
samples[95]
);
assert!(
samples[99] < std::time::Duration::from_millis(900),
"p100 blocked-wait wake latency too high (wake likely lost): {:#?}",
samples[99]
);
}
#[test]
fn wake_fires_alongside_connection_ready() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let mut client = TcpStream::connect(addr).unwrap();
let (server, _) = listener.accept().unwrap();
let wl = TcpListener::bind("127.0.0.1:0").unwrap();
let wa = wl.local_addr().unwrap();
let mut writer = TcpStream::connect(wa).unwrap();
let (wake_reader, _) = wl.accept().unwrap();
wake_reader.set_nonblocking(true).unwrap();
writer.set_nonblocking(true).unwrap();
let mut p = Poller::new();
p.register(7, fd_of(&server), false);
std::io::Write::write_all(&mut writer, b"\x01").unwrap();
client.write_all(b"hi").unwrap();
let deadline = std::time::Instant::now() + std::time::Duration::from_secs(5);
let (mut saw_conn, mut saw_wake) = (false, false);
while !(saw_conn && saw_wake) {
assert!(
std::time::Instant::now() < deadline,
"timed out waiting for readiness: conn={saw_conn} wake={saw_wake}"
);
let ready = p.wait(100, Some(fd_of(&wake_reader))).unwrap();
saw_conn |= ready.contains(&7);
saw_wake |= ready.contains(&WAKE_ID);
}
}
}