use std::collections::HashMap;
use std::io;
use std::net::SocketAddr;
use std::os::fd::{IntoRawFd, RawFd};
use std::path::Path;
use std::time::Duration;
use io_uring::types::{Fd, SubmitArgs, Timespec};
use super::{Cqe, Engine, Poll};
use crate::buffer::BufferPool;
use crate::splice;
use crate::token::Token;
struct Listener {
fd: RawFd,
#[allow(dead_code)] token: Token,
sa: Box<[u8; 128]>,
sa_len: Box<libc::socklen_t>,
}
pub struct UringEngine {
ring: io_uring::IoUring,
slot_bases: Vec<*mut u8>,
buf_size: usize,
listeners: HashMap<u64, Listener>,
splice_dirs: HashMap<u64, (RawFd, RawFd)>,
accepted: HashMap<RawFd, SocketAddr>,
connect_addrs: HashMap<u64, Box<ConnectAddr>>,
}
struct ConnectAddr {
storage: libc::sockaddr_storage,
len: libc::socklen_t,
}
unsafe fn connect_addr_boxed(
addr: SocketAddr,
) -> (Box<ConnectAddr>, *const libc::sockaddr, libc::socklen_t) {
let mut storage: libc::sockaddr_storage = unsafe { std::mem::zeroed() };
let len = match addr {
SocketAddr::V4(v4) => {
let sa: &mut libc::sockaddr_in =
unsafe { &mut *std::ptr::addr_of_mut!(storage).cast::<libc::sockaddr_in>() };
sa.sin_family = libc::AF_INET as _;
sa.sin_port = v4.port().to_be();
sa.sin_addr.s_addr = u32::from_ne_bytes(v4.ip().octets());
std::mem::size_of::<libc::sockaddr_in>() as libc::socklen_t
}
SocketAddr::V6(v6) => {
let sa: &mut libc::sockaddr_in6 =
unsafe { &mut *std::ptr::addr_of_mut!(storage).cast::<libc::sockaddr_in6>() };
sa.sin6_family = libc::AF_INET6 as _;
sa.sin6_port = v6.port().to_be();
sa.sin6_addr.s6_addr = v6.ip().octets();
std::mem::size_of::<libc::sockaddr_in6>() as libc::socklen_t
}
};
let boxed = Box::new(ConnectAddr { storage, len });
let ptr = std::ptr::addr_of!(boxed.storage).cast::<libc::sockaddr>();
let out_len = boxed.len;
(boxed, ptr, out_len)
}
unsafe impl Send for UringEngine {}
impl UringEngine {
pub fn new(entries: u32, pool: Option<&BufferPool>, sqpoll: bool) -> io::Result<Self> {
let ring = if sqpoll {
io_uring::IoUring::builder()
.setup_sqpoll(2_000)
.build(entries)?
} else {
io_uring::IoUring::new(entries)?
};
let (slot_bases, buf_size) = match pool {
Some(pool) => {
let bases = (0..pool.capacity() as u32)
.map(|i| pool.slot(i).as_ptr() as *mut u8)
.collect::<Vec<_>>();
let iovecs: Vec<libc::iovec> = bases
.iter()
.map(|p| libc::iovec {
iov_base: (*p).cast(),
iov_len: pool.buf_size(),
})
.collect();
unsafe {
ring.submitter().register_buffers(&iovecs)?;
}
(bases, pool.buf_size())
}
None => (Vec::new(), 0),
};
Ok(Self {
ring,
slot_bases,
buf_size,
listeners: HashMap::new(),
splice_dirs: HashMap::new(),
accepted: HashMap::new(),
connect_addrs: HashMap::new(),
})
}
fn push(&mut self, entry: io_uring::squeue::Entry, token: Token) {
let entry = entry.user_data(token.bits());
loop {
unsafe {
if self.ring.submission().push(&entry).is_ok() {
return;
}
}
let _ = self.ring.submit();
std::hint::spin_loop();
}
}
fn slot_ptr(&self, slot: u32) -> *mut u8 {
self.slot_bases[slot as usize]
}
fn arm_accept(&mut self, bits: u64) {
let Some(l) = self.listeners.get_mut(&bits) else {
return;
};
let fd = l.fd;
let sa_ptr = l.sa.as_mut_ptr();
let len_ptr: *mut libc::socklen_t = &mut *l.sa_len;
let entry = io_uring::opcode::Accept::new(Fd(fd), sa_ptr.cast(), len_ptr)
.flags(libc::SOCK_NONBLOCK | libc::SOCK_CLOEXEC)
.build();
self.push(entry, Token::from_bits(bits));
}
fn arm_readiness(&mut self, fd: RawFd, token: Token) {
let entry = io_uring::opcode::PollAdd::new(Fd(fd), libc::POLLIN as u32)
.multi(true)
.build();
self.push(entry, token);
}
}
impl Engine for UringEngine {
fn kind(&self) -> &'static str {
"io_uring"
}
fn add_listener(&mut self, fd: RawFd, token: Token) -> io::Result<()> {
let bits = token.bits();
self.listeners.insert(
bits,
Listener {
fd,
token,
sa: Box::new([0u8; 128]),
sa_len: Box::new(128),
},
);
self.arm_accept(bits);
Ok(())
}
fn add_stream(&mut self, _fd: RawFd, _token: Token) -> io::Result<()> {
Ok(())
}
fn read(&mut self, token: Token, fd: RawFd, slot: u32) -> io::Result<Poll> {
let ptr = self.slot_ptr(slot);
let entry =
io_uring::opcode::ReadFixed::new(Fd(fd), ptr, self.buf_size as u32, slot as u16)
.offset(0)
.build();
self.push(entry, token);
Ok(Poll::Pending)
}
fn write(
&mut self,
token: Token,
fd: RawFd,
slot: u32,
len: usize,
offset: usize,
) -> io::Result<Poll> {
let ptr = self.slot_ptr(slot);
let entry = unsafe {
io_uring::opcode::WriteFixed::new(
Fd(fd),
ptr.add(offset),
(len - offset) as u32,
slot as u16,
)
.offset(0)
.build()
};
self.push(entry, token);
Ok(Poll::Pending)
}
fn connect(&mut self, token: Token, addr: SocketAddr) -> io::Result<(RawFd, Poll)> {
let domain = if addr.is_ipv4() {
socket2::Domain::IPV4
} else {
socket2::Domain::IPV6
};
let sock =
socket2::Socket::new(domain, socket2::Type::STREAM, Some(socket2::Protocol::TCP))?;
sock.set_nonblocking(true)?;
sock.set_tcp_nodelay(true)?;
let fd = sock.into_raw_fd();
let (owned, ptr, len) = unsafe { connect_addr_boxed(addr) };
let entry = io_uring::opcode::Connect::new(Fd(fd), ptr, len).build();
self.connect_addrs.insert(token.bits(), owned);
self.push(entry, token);
Ok((fd, Poll::Pending))
}
fn connect_unix(&mut self, token: Token, path: &Path) -> io::Result<(RawFd, Poll)> {
let sock = socket2::Socket::new(socket2::Domain::UNIX, socket2::Type::STREAM, None)?;
sock.set_nonblocking(true)?;
let fd = sock.into_raw_fd();
let sa = socket2::SockAddr::unix(path)?;
let mut storage: libc::sockaddr_storage = unsafe { std::mem::zeroed() };
let bytes = sa.as_ptr().cast::<u8>();
let copy_len = (sa.len() as usize).min(std::mem::size_of::<libc::sockaddr_storage>());
unsafe {
std::ptr::copy_nonoverlapping(bytes, std::ptr::addr_of_mut!(storage).cast(), copy_len)
};
let len = sa.len();
let owned = Box::new(ConnectAddr { storage, len });
let ptr = std::ptr::addr_of!(owned.storage).cast::<libc::sockaddr>();
let entry = io_uring::opcode::Connect::new(Fd(fd), ptr, len).build();
self.connect_addrs.insert(token.bits(), owned);
self.push(entry, token);
Ok((fd, Poll::Pending))
}
fn accept(&mut self, _lfd: RawFd, ltoken: Token) -> io::Result<Option<(RawFd, SocketAddr)>> {
let keys: Vec<RawFd> = self.accepted.keys().copied().collect();
if let Some(fd) = keys.into_iter().next() {
let addr = self.accepted.remove(&fd).expect("just listed");
return Ok(Some((fd, addr)));
}
let _ = ltoken;
Ok(None)
}
fn splice_pump(&mut self, a: Token, afd: i32, b: Token, bfd: i32) -> io::Result<()> {
debug_assert_ne!(a.bits(), b.bits(), "splice directions need distinct tokens");
self.splice_dirs.insert(a.bits(), (afd, bfd));
self.splice_dirs.insert(b.bits(), (bfd, afd));
self.arm_readiness(afd, a);
self.arm_readiness(bfd, b);
Ok(())
}
fn remove(&mut self, fd: RawFd) {
self.accepted.remove(&fd);
}
fn poll(&mut self, timeout: Option<Duration>, out: &mut Vec<Cqe>) -> io::Result<()> {
self.ring.submit()?;
match timeout {
None => match self.ring.submit_and_wait(1) {
Ok(_) => {}
Err(e) if e.kind() == io::ErrorKind::Interrupted => {}
Err(e) => return Err(e),
},
Some(d) => {
let ms = d.as_millis().min(60_000) as u64;
let ts = Timespec::new()
.sec(ms / 1_000)
.nsec((ms % 1_000) as u32 * 1_000_000);
let args = SubmitArgs::new().timespec(&ts);
match self.ring.submitter().submit_with_args(1, &args) {
Ok(_) => {}
Err(e) if e.raw_os_error() == Some(libc::ETIME) => return Ok(()),
Err(e) if e.kind() == io::ErrorKind::Interrupted => {}
Err(_) => {
self.ring.submit_and_wait(0)?;
}
}
}
}
let mut accept_hits: Vec<(Token, RawFd)> = Vec::new();
for cqe in self.ring.completion() {
let token = Token::from_bits(cqe.user_data());
let raw = cqe.result();
let result = if raw >= 0 {
Ok(raw as u32)
} else {
Err(io::Error::from_raw_os_error(-raw))
};
if token.op() == crate::token::Op::Connect {
self.connect_addrs.remove(&token.bits());
}
match token.op() {
crate::token::Op::Accept => {
if raw >= 0 {
accept_hits.push((token, raw as RawFd));
} else if -raw != libc::ECANCELED {
out.push(Cqe { token, result });
}
}
crate::token::Op::Splice => {
if raw < 0 {
if -raw != libc::ECANCELED {
out.push(Cqe { token, result });
}
} else if let Some(&(from, to)) = self.splice_dirs.get(&token.bits()) {
match splice::pump(from, to, 1 << 20) {
splice::PumpResult::Moved(n) => {
out.push(Cqe {
token,
result: Ok(n as u32),
});
}
splice::PumpResult::Eof => {
out.push(Cqe {
token,
result: Ok(0),
});
}
splice::PumpResult::WouldBlock => {}
splice::PumpResult::Err(code) => out.push(Cqe {
token,
result: Err(io::Error::from_raw_os_error(code)),
}),
}
}
}
_ => out.push(Cqe { token, result }),
}
}
for (token, fd) in accept_hits {
let bits = token.bits();
let addr = self.listeners.get(&bits).map_or_else(
|| SocketAddr::from(([0, 0, 0, 0], 0)),
|l| parse_sockaddr(&l.sa),
);
self.accepted.insert(fd, addr);
self.arm_accept(bits);
}
Ok(())
}
fn take_accepted(&mut self, fd: RawFd) -> Option<SocketAddr> {
self.accepted.remove(&fd)
}
}
fn parse_sockaddr(buf: &[u8; 128]) -> SocketAddr {
let sa: &libc::sockaddr_storage = unsafe { &*buf.as_ptr().cast() };
match sa.ss_family as i32 {
libc::AF_INET => {
let a: &libc::sockaddr_in =
unsafe { &*(sa as *const libc::sockaddr_storage).cast::<libc::sockaddr_in>() };
SocketAddr::from((
std::net::Ipv4Addr::from(u32::from_be(a.sin_addr.s_addr)),
u16::from_be(a.sin_port),
))
}
_ => {
let a: &libc::sockaddr_in6 =
unsafe { &*(sa as *const libc::sockaddr_storage).cast::<libc::sockaddr_in6>() };
SocketAddr::from((
std::net::Ipv6Addr::from(a.sin6_addr.s6_addr),
u16::from_be(a.sin6_port),
))
}
}
}
#[cfg(test)]
mod fault_tests {
use super::*;
use crate::buffer::DEFAULT_BUF_SIZE;
use crate::token::Op;
fn test_engine() -> (BufferPool, UringEngine) {
let pool = BufferPool::new(8, DEFAULT_BUF_SIZE).expect("pool");
let engine = UringEngine::new(64, Some(&pool), false).expect("uring available");
(pool, engine)
}
fn tok(op: Op) -> Token {
Token::new(op, 0, 0, 0)
}
fn sockpair() -> (RawFd, RawFd) {
let mut fds = [0 as RawFd; 2];
let rc = unsafe { libc::socketpair(libc::AF_UNIX, libc::SOCK_STREAM, 0, fds.as_mut_ptr()) };
assert_eq!(rc, 0);
for fd in fds {
let flags = unsafe { libc::fcntl(fd, libc::F_GETFL) };
unsafe { libc::fcntl(fd, libc::F_SETFL, flags | libc::O_NONBLOCK) };
}
(fds[0], fds[1])
}
fn close(fd: RawFd) {
unsafe { libc::close(fd) };
}
fn drain(engine: &mut UringEngine, secs: u64) -> Vec<Cqe> {
let mut out = Vec::new();
engine
.poll(Some(std::time::Duration::from_secs(secs)), &mut out)
.expect("poll");
out
}
fn drain_until(engine: &mut UringEngine, mut pred: impl FnMut(&Cqe) -> bool) -> Vec<Cqe> {
let mut all = Vec::new();
for _ in 0..60 {
let mut out = Vec::new();
engine
.poll(Some(std::time::Duration::from_millis(100)), &mut out)
.expect("poll");
if out.iter().any(&mut pred) {
all.extend(out);
return all;
}
all.extend(out);
}
all
}
#[test]
fn write_then_read_roundtrip() {
let (_pool, mut engine) = test_engine();
let (a, b) = sockpair();
let wt = tok(Op::DownstreamWrite);
assert!(matches!(
engine.write(wt, a, 0, 32, 0).expect("write"),
Poll::Pending
));
let cqes = drain(&mut engine, 2);
assert!(
cqes.iter().any(|c| c.token == wt && c.result.is_ok()),
"write CQE missing: {cqes:?}"
);
let rt = tok(Op::DownstreamRead);
assert!(matches!(
engine.read(rt, b, 1).expect("read"),
Poll::Pending
));
let cqes = drain(&mut engine, 2);
let got = cqes.iter().find(|c| c.token == rt).expect("read CQE");
assert!(matches!(got.result, Ok(32)));
close(a);
close(b);
}
#[test]
fn connect_refused_completes_with_error() {
let (_pool, mut engine) = test_engine();
let addr: SocketAddr = "127.0.0.1:1".parse().expect("addr");
let t = tok(Op::Connect);
let (fd, poll) = engine.connect(t, addr).expect("connect issued");
match poll {
Poll::Done(_) => {}
Poll::Pending => {
let cqes = drain_until(&mut engine, |c| c.token == t);
assert!(
cqes.iter().any(|c| c.token == t && c.result.is_err()),
"refused connect must error: {cqes:?}"
);
}
}
close(fd);
}
#[test]
fn connect_unix_missing_path_errors() {
let (_pool, mut engine) = test_engine();
let t = tok(Op::Connect);
let dir = tempfile::tempdir().expect("dir");
let missing = dir.path().join("no.sock");
let res = engine.connect_unix(t, &missing);
assert!(res.is_err() || matches!(res, Ok((_, Poll::Pending))));
}
#[test]
fn accept_flow_materializes_connection() {
let (_pool, mut engine) = test_engine();
let listener =
crate::tcp_listener("127.0.0.1:0".parse().expect("addr"), true, 64).expect("bind");
let lfd = std::os::fd::AsRawFd::as_raw_fd(&listener);
engine.add_listener(lfd, Token::accept(0)).expect("add");
assert!(
engine
.accept(lfd, Token::accept(0))
.expect("accept")
.is_none()
);
let addr = listener.local_addr().expect("addr");
let _client = std::net::TcpStream::connect(addr).expect("connect");
let mut got = None;
for _ in 0..60 {
let mut out = Vec::new();
engine
.poll(Some(std::time::Duration::from_millis(100)), &mut out)
.expect("poll");
got = engine.accept(lfd, Token::accept(0)).expect("accept2");
if got.is_some() {
break;
}
}
assert!(got.is_some(), "materialized connection expected");
let (fd, peer) = got.expect("some");
assert!(peer.port() != 0);
close(fd);
}
#[test]
fn splice_moved_and_eof_paths() {
let (_pool, mut engine) = test_engine();
let mut fds = [0 as RawFd; 2];
assert_eq!(
unsafe { libc::pipe2(fds.as_mut_ptr(), libc::O_NONBLOCK | libc::O_CLOEXEC) },
0
);
let (pr, pw) = (fds[0], fds[1]);
let payload = b"uring-splice";
unsafe { libc::write(pw, payload.as_ptr().cast(), payload.len()) };
let (sa, sb) = sockpair();
let t = tok(Op::Splice);
let t_rev = Token::new(Op::Splice, 0, 0, 1);
engine.splice_pump(t, pr, t_rev, sb).expect("splice_pump");
let cqes = drain_until(&mut engine, |c| c.token == t);
assert!(
cqes.iter()
.any(|c| c.token == t && matches!(c.result, Ok(n) if n as usize == payload.len())),
"splice moved CQE: {cqes:?}"
);
let mut buf = [0u8; 64];
let n = unsafe { libc::read(sa, buf.as_mut_ptr().cast(), 64) };
assert_eq!(&buf[..n as usize], payload);
close(pw); let t2 = Token::new(Op::Splice, 1, 0, 0);
let t2_rev = Token::new(Op::Splice, 1, 0, 1);
engine
.splice_pump(t2, pr, t2_rev, sb)
.expect("splice_pump2");
let cqes = drain_until(&mut engine, |c| c.token == t2);
assert!(
cqes.iter()
.any(|c| c.token == t2 && matches!(c.result, Ok(0))),
"splice EOF CQE: {cqes:?}"
);
close(pr);
close(sa);
close(sb);
}
#[test]
fn remove_clears_tracked_fd() {
let (_pool, mut engine) = test_engine();
let (a, b) = sockpair();
engine
.add_stream(a, tok(Op::DownstreamRead))
.expect("add_stream");
engine.remove(a);
close(a);
close(b);
}
}