Skip to content

Commit c87829b

Browse files
committed
fix(host): replace WSAPoll with select() on Windows
WSAPoll has known reliability issues on Windows — it fails to report POLLRDNORM on connected TCP sockets, causing outbound HTTP requests to time out waiting for response data. Replace both handle_net_poll and hl_sleep_poll_sockets with Winsock select(), which correctly detects read/write readiness on all socket types. Also ensures hl_sleep_poll_sockets (Unix and Windows) polls for both read and write so the Unikraft cooperative scheduler wakes the guest for send and receive readiness. Signed-off-by: danbugs <danilochiarlone@gmail.com>
1 parent 682ad20 commit c87829b

1 file changed

Lines changed: 113 additions & 65 deletions

File tree

host/src/lib.rs

Lines changed: 113 additions & 65 deletions
Original file line numberDiff line numberDiff line change
@@ -1706,47 +1706,14 @@ fn handle_net_poll(
17061706
use serde_json::json;
17071707
use std::os::windows::io::AsRawSocket;
17081708
use windows_sys::Win32::Networking::WinSock::{
1709-
WSAPoll, POLLERR as W_POLLERR, POLLHUP as W_POLLHUP, POLLNVAL as W_POLLNVAL, POLLRDNORM,
1710-
POLLWRNORM, WSAPOLLFD,
1709+
__WSAFDIsSet, select, FD_SET, FD_SETSIZE, SOCKET, SOCKET_ERROR, TIMEVAL,
17111710
};
17121711

17131712
const POSIX_POLLIN: i16 = 0x0001;
17141713
const POSIX_POLLOUT: i16 = 0x0004;
17151714
const POSIX_POLLERR: i16 = 0x0008;
1716-
const POSIX_POLLHUP: i16 = 0x0010;
17171715
const POSIX_POLLNVAL: i16 = 0x0020;
17181716

1719-
fn posix_to_win(posix: i16) -> i16 {
1720-
let mut win: i16 = 0;
1721-
if posix & POSIX_POLLIN != 0 {
1722-
win |= POLLRDNORM;
1723-
}
1724-
if posix & POSIX_POLLOUT != 0 {
1725-
win |= POLLWRNORM;
1726-
}
1727-
win
1728-
}
1729-
1730-
fn win_to_posix(win: i16) -> i16 {
1731-
let mut posix: i16 = 0;
1732-
if win & POLLRDNORM != 0 {
1733-
posix |= POSIX_POLLIN;
1734-
}
1735-
if win & POLLWRNORM != 0 {
1736-
posix |= POSIX_POLLOUT;
1737-
}
1738-
if win & W_POLLERR != 0 {
1739-
posix |= POSIX_POLLERR;
1740-
}
1741-
if win & W_POLLHUP != 0 {
1742-
posix |= POSIX_POLLHUP;
1743-
}
1744-
if win & W_POLLNVAL != 0 {
1745-
posix |= POSIX_POLLNVAL;
1746-
}
1747-
posix
1748-
}
1749-
17501717
let fds_val = args
17511718
.get("fds")
17521719
.and_then(|v| v.as_array())
@@ -1758,8 +1725,15 @@ fn handle_net_poll(
17581725
.clamp(0, i32::MAX as i64) as i32;
17591726

17601727
let tbl = table.lock().unwrap();
1761-
let mut pollfds: Vec<WSAPOLLFD> = Vec::new();
1762-
let mut guest_fds: Vec<u64> = Vec::new();
1728+
1729+
struct FdEntry {
1730+
raw: SOCKET,
1731+
guest_fd: u64,
1732+
want_read: bool,
1733+
want_write: bool,
1734+
}
1735+
1736+
let mut entries: Vec<FdEntry> = Vec::new();
17631737
let mut ready = Vec::new();
17641738

17651739
for entry in fds_val {
@@ -1771,36 +1745,90 @@ fn handle_net_poll(
17711745
if !(0..=i16::MAX as i64).contains(&raw_events) {
17721746
return Err(anyhow!("net_poll: events {raw_events} out of i16 range"));
17731747
}
1774-
let events = posix_to_win(raw_events as i16);
1748+
let events = raw_events as i16;
17751749
if let Ok(sock) = tbl.get_socket(fd) {
1776-
pollfds.push(WSAPOLLFD {
1777-
fd: sock.as_raw_socket() as usize,
1778-
events,
1779-
revents: 0,
1750+
entries.push(FdEntry {
1751+
raw: sock.as_raw_socket() as SOCKET,
1752+
guest_fd: fd,
1753+
want_read: events & POSIX_POLLIN != 0,
1754+
want_write: events & POSIX_POLLOUT != 0,
17801755
});
1781-
guest_fds.push(fd);
17821756
} else {
17831757
ready.push(json!({ "fd": fd, "revents": POSIX_POLLNVAL as i64 }));
17841758
}
17851759
}
17861760
drop(tbl);
17871761

1788-
if pollfds.is_empty() {
1762+
if entries.is_empty() {
17891763
return Ok(json!({"ready": ready}));
17901764
}
1765+
if entries.len() > FD_SETSIZE as usize {
1766+
return Err(anyhow!("net_poll: too many fds for select()"));
1767+
}
17911768

1792-
let ret = unsafe { WSAPoll(pollfds.as_mut_ptr(), pollfds.len() as u32, timeout_ms) };
1769+
let mut readfds: FD_SET = unsafe { std::mem::zeroed() };
1770+
let mut writefds: FD_SET = unsafe { std::mem::zeroed() };
1771+
let mut exceptfds: FD_SET = unsafe { std::mem::zeroed() };
17931772

1794-
if ret < 0 {
1773+
for e in &entries {
1774+
if e.want_read {
1775+
let c = readfds.fd_count as usize;
1776+
readfds.fd_array[c] = e.raw;
1777+
readfds.fd_count += 1;
1778+
}
1779+
if e.want_write {
1780+
let c = writefds.fd_count as usize;
1781+
writefds.fd_array[c] = e.raw;
1782+
writefds.fd_count += 1;
1783+
}
1784+
let c = exceptfds.fd_count as usize;
1785+
exceptfds.fd_array[c] = e.raw;
1786+
exceptfds.fd_count += 1;
1787+
}
1788+
1789+
let tv = TIMEVAL {
1790+
tv_sec: timeout_ms / 1000,
1791+
tv_usec: (timeout_ms % 1000) * 1000,
1792+
};
1793+
1794+
let ret = unsafe {
1795+
select(
1796+
0,
1797+
if readfds.fd_count > 0 {
1798+
&mut readfds
1799+
} else {
1800+
std::ptr::null_mut()
1801+
},
1802+
if writefds.fd_count > 0 {
1803+
&mut writefds
1804+
} else {
1805+
std::ptr::null_mut()
1806+
},
1807+
&mut exceptfds,
1808+
&tv,
1809+
)
1810+
};
1811+
1812+
if ret == SOCKET_ERROR {
17951813
let err = std::io::Error::last_os_error();
1796-
return Err(anyhow!("net_poll: WSAPoll() failed: {err}"));
1814+
return Err(anyhow!("net_poll: select() failed: {err}"));
17971815
}
17981816

1799-
for (i, pfd) in pollfds.iter().enumerate() {
1800-
if pfd.revents != 0 {
1817+
for e in &entries {
1818+
let mut revents: i16 = 0;
1819+
if unsafe { __WSAFDIsSet(e.raw, &mut readfds) } != 0 {
1820+
revents |= POSIX_POLLIN;
1821+
}
1822+
if unsafe { __WSAFDIsSet(e.raw, &mut writefds) } != 0 {
1823+
revents |= POSIX_POLLOUT;
1824+
}
1825+
if unsafe { __WSAFDIsSet(e.raw, &mut exceptfds) } != 0 {
1826+
revents |= POSIX_POLLERR;
1827+
}
1828+
if revents != 0 {
18011829
ready.push(json!({
1802-
"fd": guest_fds[i],
1803-
"revents": win_to_posix(pfd.revents) as i64,
1830+
"fd": e.guest_fd,
1831+
"revents": revents as i64,
18041832
}));
18051833
}
18061834
}
@@ -1824,7 +1852,7 @@ fn hl_sleep_poll_sockets(
18241852
.values()
18251853
.map(|hs| libc::pollfd {
18261854
fd: hs.socket.as_raw_fd(),
1827-
events: libc::POLLIN | libc::POLLOUT | libc::POLLERR,
1855+
events: libc::POLLIN | libc::POLLOUT,
18281856
revents: 0,
18291857
})
18301858
.collect();
@@ -1851,7 +1879,8 @@ fn hl_sleep_poll_sockets(
18511879
}
18521880
}
18531881

1854-
Ok(json!({"socket_ready": ret > 0}))
1882+
let ready = pollfds.iter().any(|p| p.revents != 0);
1883+
Ok(json!({"socket_ready": ready}))
18551884
}
18561885

18571886
#[cfg(windows)]
@@ -1862,35 +1891,54 @@ fn hl_sleep_poll_sockets(
18621891
) -> Result<serde_json::Value> {
18631892
use serde_json::json;
18641893
use std::os::windows::io::AsRawSocket;
1865-
use windows_sys::Win32::Networking::WinSock::{WSAPoll, POLLRDNORM, POLLWRNORM, WSAPOLLFD};
1894+
use windows_sys::Win32::Networking::WinSock::{
1895+
select, FD_SET, FD_SETSIZE, SOCKET, SOCKET_ERROR, TIMEVAL,
1896+
};
18661897

18671898
let tbl = table.lock().unwrap();
1868-
let mut pollfds: Vec<WSAPOLLFD> = tbl
1899+
let raw_sockets: Vec<SOCKET> = tbl
18691900
.sockets
18701901
.values()
1871-
.map(|hs| WSAPOLLFD {
1872-
fd: hs.socket.as_raw_socket() as usize,
1873-
// POLLERR is output-only on Windows; setting it in events causes WSAEINVAL
1874-
events: (POLLRDNORM | POLLWRNORM) as i16,
1875-
revents: 0,
1876-
})
1902+
.take(FD_SETSIZE as usize)
1903+
.map(|hs| hs.socket.as_raw_socket() as SOCKET)
18771904
.collect();
18781905
drop(tbl);
18791906

1880-
if pollfds.is_empty() {
1907+
if raw_sockets.is_empty() {
18811908
_sc.wait(Duration::from_nanos(ns));
18821909
return Ok(json!({}));
18831910
}
18841911

1912+
let mut readfds: FD_SET = unsafe { std::mem::zeroed() };
1913+
let mut writefds: FD_SET = unsafe { std::mem::zeroed() };
1914+
let mut exceptfds: FD_SET = unsafe { std::mem::zeroed() };
1915+
for &s in &raw_sockets {
1916+
let rc = readfds.fd_count as usize;
1917+
readfds.fd_array[rc] = s;
1918+
readfds.fd_count += 1;
1919+
let wc = writefds.fd_count as usize;
1920+
writefds.fd_array[wc] = s;
1921+
writefds.fd_count += 1;
1922+
let ec = exceptfds.fd_count as usize;
1923+
exceptfds.fd_array[ec] = s;
1924+
exceptfds.fd_count += 1;
1925+
}
1926+
18851927
let timeout_ms = ((ns / 1_000_000) as i32).clamp(1, 30_000);
1886-
let ret = unsafe { WSAPoll(pollfds.as_mut_ptr(), pollfds.len() as u32, timeout_ms) };
1928+
let tv = TIMEVAL {
1929+
tv_sec: timeout_ms / 1000,
1930+
tv_usec: (timeout_ms % 1000) * 1000,
1931+
};
18871932

1888-
if ret < 0 {
1933+
let ret = unsafe { select(0, &mut readfds, &mut writefds, &mut exceptfds, &tv) };
1934+
1935+
if ret == SOCKET_ERROR {
18891936
let err = std::io::Error::last_os_error();
1890-
return Err(anyhow!("hl_sleep WSAPoll failed: {err}"));
1937+
return Err(anyhow!("hl_sleep select() failed: {err}"));
18911938
}
18921939

1893-
Ok(json!({"socket_ready": ret > 0}))
1940+
let ready = ret > 0;
1941+
Ok(json!({"socket_ready": ready}))
18941942
}
18951943

18961944
// ---------------------------------------------------------------------------

0 commit comments

Comments
 (0)