use super::{Future, Pin};
use std::cell::RefCell;
use std::os::fd::{AsRawFd, FromRawFd, RawFd};
use std::sync::OnceLock;
use std::sync::atomic::{
AtomicBool, AtomicI32, AtomicPtr, AtomicU32, AtomicUsize, Ordering, fence,
};
use std::task::{Context, Poll, RawWaker, RawWakerVTable, Waker};
use crate::io::trace::io_trace;
#[allow(unused_imports)]
use crate::io::trace::trace_now_us;
static WORKER_ROUND_ROBIN: AtomicUsize = AtomicUsize::new(0);
use crate::lockfree::{BufferPool, SpscQueue, TreiberStack};
#[repr(align(64))]
struct LocalAllocator {
thread_idx: usize,
local_chunks: Vec<u32>,
}
thread_local! {
static LOCAL_ALLOCATOR: RefCell<Option<LocalAllocator>> = const { RefCell::new(None) };
static THREAD_ID: usize = {
static NEXT_ID: AtomicUsize = AtomicUsize::new(0);
NEXT_ID.fetch_add(1, Ordering::Relaxed)
};
}
#[doc(hidden)]
#[must_use]
pub fn get_local_thread_id() -> usize {
THREAD_ID.with(|id| *id)
}
static THREAD_RETURNED_STACKS: OnceLock<Box<[TreiberStack]>> = OnceLock::new();
static GLOBAL_BUFFER_POOL: OnceLock<BufferPool> = OnceLock::new();
static CHUNK_OWNERS: OnceLock<Box<[AtomicU32]>> = OnceLock::new();
fn get_or_init_local_allocator() -> Option<usize> {
LOCAL_ALLOCATOR.with(|cell| {
let mut borrow = cell.borrow_mut();
if borrow.is_none() {
let idx = get_local_thread_id();
if idx < 512 {
*borrow = Some(LocalAllocator {
thread_idx: idx,
local_chunks: Vec::new(),
});
}
}
borrow.as_ref().map(|alloc| alloc.thread_idx)
})
}
#[doc(hidden)]
pub fn allocate_buffer() -> Option<u32> {
let t_idx_opt = get_or_init_local_allocator();
if let Some(t_idx) = t_idx_opt {
LOCAL_ALLOCATOR.with(|cell| {
let mut borrow = cell.borrow_mut();
let alloc = borrow.as_mut().unwrap();
if let Some(idx) = alloc.local_chunks.pop() {
return Some(idx);
}
if let Some(stacks) = THREAD_RETURNED_STACKS.get()
&& let Some(stack) = stacks.get(t_idx)
{
while let Some(idx) = stack.pop() {
alloc.local_chunks.push(idx);
}
if let Some(idx) = alloc.local_chunks.pop() {
return Some(idx);
}
}
if let Some(pool) = GLOBAL_BUFFER_POOL.get()
&& let Some(idx) = pool.acquire()
{
if let Some(owners) = CHUNK_OWNERS.get() {
owners[idx as usize].store(t_idx as u32, Ordering::Release);
}
return Some(idx);
}
None
})
} else if let Some(pool) = GLOBAL_BUFFER_POOL.get() {
if let Some(idx) = pool.acquire() {
if let Some(owners) = CHUNK_OWNERS.get() {
owners[idx as usize].store(u32::MAX, Ordering::Release);
}
return Some(idx);
}
None
} else {
None
}
}
#[doc(hidden)]
pub fn free_buffer(idx: u32) {
if let Some(pool) = GLOBAL_BUFFER_POOL.get() {
let owner = CHUNK_OWNERS.get().map_or(u32::MAX, |owners| {
owners[idx as usize].load(Ordering::Acquire)
});
if owner == u32::MAX {
pool.release(idx);
return;
}
let current_thread_idx = get_or_init_local_allocator();
if Some(owner as usize) == current_thread_idx {
LOCAL_ALLOCATOR.with(|cell| {
if let Some(alloc) = cell.borrow_mut().as_mut() {
alloc.local_chunks.push(idx);
}
});
} else if let Some(stacks) = THREAD_RETURNED_STACKS.get() {
if let Some(stack) = stacks.get(owner as usize) {
stack.push(idx);
} else {
pool.release(idx);
}
} else {
pool.release(idx);
}
}
}
#[doc(hidden)]
#[repr(align(64))]
pub struct BufferSlice {
pub buf_idx: u32,
pub read_pos: usize,
pub write_pos: usize,
}
impl BufferSlice {
#[must_use]
pub const fn new(buf_idx: u32, len: usize) -> Self {
Self {
buf_idx,
read_pos: 0,
write_pos: len,
}
}
#[inline]
pub fn data(&self) -> *mut u8 {
GLOBAL_BUFFER_POOL.get().unwrap().get_ptr(self.buf_idx)
}
#[inline]
#[must_use]
pub const fn remaining(&self) -> usize {
self.write_pos.saturating_sub(self.read_pos)
}
}
impl Drop for BufferSlice {
fn drop(&mut self) {
free_buffer(self.buf_idx);
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum OpCode {
Read,
Write,
Accept,
Connect,
SendTo,
RecvFrom,
}
#[derive(Clone, Copy)]
pub enum IoRequest {
Read {
fd: u32,
direct_fd_idx: u32,
buf_ptr: *mut u8,
len: usize,
offset: i64,
slot_idx: usize,
},
Write {
fd: u32,
direct_fd_idx: u32,
buf_ptr: *const u8,
len: usize,
offset: i64,
slot_idx: usize,
},
Accept {
fd: u32,
direct_fd_idx: u32,
slot_idx: usize,
},
Connect {
fd: u32,
direct_fd_idx: u32,
addr: libc::sockaddr_storage,
addr_len: libc::socklen_t,
slot_idx: usize,
},
SendTo {
fd: u32,
direct_fd_idx: u32,
msg_ptr: *mut libc::msghdr,
slot_idx: usize,
},
RecvFrom {
fd: u32,
direct_fd_idx: u32,
msg_ptr: *mut libc::msghdr,
slot_idx: usize,
},
RegisterFile {
fd: RawFd,
slot_idx: usize,
},
UnregisterFile {
direct_fd_idx: u32,
slot_idx: usize,
},
}
#[repr(align(64))]
struct WakerSlot {
waker_data: AtomicPtr<()>,
waker_vtable: AtomicPtr<RawWakerVTable>,
waker_lock: AtomicBool,
result: AtomicI32,
completed: AtomicBool,
dropped: AtomicBool,
origin_fd: AtomicU32,
}
#[repr(align(64))]
struct WaitSlot {
waker_data: AtomicPtr<()>,
waker_vtable: AtomicPtr<RawWakerVTable>,
}
impl WakerSlot {
#[inline(always)]
fn lock_waker(&self) {
while self
.waker_lock
.compare_exchange_weak(false, true, Ordering::Acquire, Ordering::Relaxed)
.is_err()
{
core::hint::spin_loop();
}
}
#[inline(always)]
fn unlock_waker(&self) {
self.waker_lock.store(false, Ordering::Release);
}
}
#[inline(always)]
fn wake_next_waiting_fiber(state: &WorkerState) {
if let Some(wait_idx) = state.waiting_queue.pop() {
let wait_slot = &state.wait_slots[wait_idx as usize];
let data = wait_slot
.waker_data
.swap(std::ptr::null_mut(), Ordering::Relaxed);
let vtable = wait_slot
.waker_vtable
.swap(std::ptr::null_mut(), Ordering::Relaxed);
state.free_wait_slots.push(wait_idx);
if !data.is_null() && !vtable.is_null() {
let raw = RawWaker::new(data.cast_const(), unsafe { &*vtable });
let w = unsafe { Waker::from_raw(raw) };
w.wake();
}
}
}
pub struct WorkerState {
#[cfg(target_os = "linux")]
ring: std::cell::UnsafeCell<io_uring::IoUring>,
#[cfg(not(target_os = "linux"))]
poll: std::cell::UnsafeCell<mio::Poll>,
queues: Box<[SpscQueue<IoRequest>]>,
slots: Box<[WakerSlot]>,
free_slots: TreiberStack,
wait_slots: Box<[WaitSlot]>,
free_wait_slots: TreiberStack,
waiting_queue: TreiberStack,
is_sleeping: AtomicBool,
cancel_queue: TreiberStack,
#[cfg(target_os = "linux")]
wake_eventfd: RawFd,
#[cfg(target_os = "linux")]
sqpoll_enabled: bool,
#[cfg(not(target_os = "linux"))]
waker: std::sync::Arc<mio::Waker>,
direct_fd_free: TreiberStack,
}
#[allow(clippy::non_send_fields_in_send_ty)]
unsafe impl Send for WorkerState {}
unsafe impl Sync for WorkerState {}
struct GlobalConfig {
workers: usize,
pin_cpus: Vec<usize>,
}
static GLOBAL_CONFIG: OnceLock<GlobalConfig> = OnceLock::new();
static WORKERS: OnceLock<Box<[WorkerState]>> = OnceLock::new();
static SHUTDOWN: AtomicBool = AtomicBool::new(false);
#[cfg(target_os = "linux")]
fn pin_thread_to_cpu(cpu_id: usize) -> Result<(), &'static str> {
unsafe {
let mut cpuset: libc::cpu_set_t = std::mem::zeroed();
libc::CPU_SET(cpu_id, &mut cpuset);
let thread = libc::pthread_self();
let res = libc::pthread_setaffinity_np(
thread,
std::mem::size_of::<libc::cpu_set_t>(),
&raw const cpuset,
);
if res == 0 {
Ok(())
} else {
Err("pthread_setaffinity_np failed")
}
}
}
#[cfg(not(target_os = "linux"))]
fn pin_thread_to_cpu(_cpu_id: usize) -> Result<(), &'static str> {
Ok(())
}
#[cfg(target_os = "linux")]
fn pick_direct_fd_count(desired_max: usize) -> usize {
unsafe {
let mut lim: libc::rlimit = std::mem::zeroed();
if libc::getrlimit(libc::RLIMIT_NOFILE, &raw mut lim) == 0 {
if lim.rlim_cur < lim.rlim_max {
let raised = libc::rlimit {
rlim_cur: lim.rlim_max,
rlim_max: lim.rlim_max,
};
let _ = libc::setrlimit(libc::RLIMIT_NOFILE, &raw const raised);
let _ = libc::getrlimit(libc::RLIMIT_NOFILE, &raw mut lim);
}
let headroom = 256usize;
let available = (lim.rlim_cur as usize).saturating_sub(headroom);
return available.clamp(64, desired_max);
}
}
64
}
#[allow(clippy::too_many_lines)]
pub fn init_runtime(
workers: usize,
ring_depth: u32,
buffer_pool_size: usize,
chunk_size: usize,
pin_cpus: &[usize],
) {
let config = GlobalConfig {
workers,
pin_cpus: pin_cpus.to_vec(),
};
if GLOBAL_CONFIG.set(config).is_err() {
return;
}
let pool = BufferPool::new(buffer_pool_size, chunk_size);
let _ = GLOBAL_BUFFER_POOL.set(pool);
let owners: Vec<AtomicU32> = (0..buffer_pool_size)
.map(|_| AtomicU32::new(u32::MAX))
.collect();
let _ = CHUNK_OWNERS.set(owners.into_boxed_slice());
let mut returned_stacks = Vec::with_capacity(512);
for _ in 0..512 {
returned_stacks.push(TreiberStack::new(0));
}
let _ = THREAD_RETURNED_STACKS.set(returned_stacks.into_boxed_slice());
let mut worker_states = Vec::with_capacity(workers);
for _worker_idx in 0..workers {
let mut queues = Vec::with_capacity(512);
for _ in 0..512 {
queues.push(SpscQueue::new(256));
}
let queues = queues.into_boxed_slice();
let mut slots = Vec::with_capacity(ring_depth as usize);
for _ in 0..ring_depth {
slots.push(WakerSlot {
waker_data: AtomicPtr::new(std::ptr::null_mut()),
waker_vtable: AtomicPtr::new(std::ptr::null_mut()),
waker_lock: AtomicBool::new(false),
result: AtomicI32::new(0),
completed: AtomicBool::new(false),
dropped: AtomicBool::new(false),
origin_fd: AtomicU32::new(u32::MAX),
});
}
let slots = slots.into_boxed_slice();
let free_slots = TreiberStack::new(ring_depth as usize);
for i in 0..ring_depth {
free_slots.push(i);
}
let cancel_queue = TreiberStack::new(ring_depth as usize);
let wait_slots_depth = 65536;
let mut wait_slots = Vec::with_capacity(wait_slots_depth);
for _ in 0..wait_slots_depth {
wait_slots.push(WaitSlot {
waker_data: AtomicPtr::new(std::ptr::null_mut()),
waker_vtable: AtomicPtr::new(std::ptr::null_mut()),
});
}
let wait_slots = wait_slots.into_boxed_slice();
let free_wait_slots = TreiberStack::new(wait_slots_depth);
for i in 0..wait_slots_depth {
free_wait_slots.push(i as u32);
}
let waiting_queue = TreiberStack::new(wait_slots_depth);
let is_sleeping = AtomicBool::new(false);
#[cfg(target_os = "linux")]
let direct_fd_count = pick_direct_fd_count(4096);
#[cfg(not(target_os = "linux"))]
let direct_fd_count = 4096usize;
let direct_fd_free = TreiberStack::new(direct_fd_count);
for i in 0..direct_fd_count as u32 {
direct_fd_free.push(i);
}
#[cfg(target_os = "linux")]
{
let wake_eventfd = unsafe { libc::eventfd(0, libc::EFD_CLOEXEC | libc::EFD_NONBLOCK) };
assert!(wake_eventfd >= 0, "Failed to create eventfd");
let (ring, sqpoll_enabled) = io_uring::IoUring::builder()
.setup_sqpoll(2000)
.build(ring_depth)
.map_or_else(
|_| {
(
io_uring::IoUring::new(ring_depth)
.expect("Failed to initialize io_uring fallback"),
false,
)
},
|r| (r, true),
);
let initial_fds = vec![-1; direct_fd_count];
ring.submitter()
.register_files(&initial_fds)
.expect("Failed to register direct FDs");
worker_states.push(WorkerState {
ring: std::cell::UnsafeCell::new(ring),
queues,
slots,
free_slots,
wait_slots,
free_wait_slots,
waiting_queue,
is_sleeping,
cancel_queue,
wake_eventfd,
sqpoll_enabled,
direct_fd_free,
});
}
#[cfg(not(target_os = "linux"))]
{
let poll = mio::Poll::new().expect("Failed to initialize mio Poll");
let waker = std::sync::Arc::new(
mio::Waker::new(poll.registry(), mio::Token(0))
.expect("Failed to create mio waker"),
);
worker_states.push(WorkerState {
poll: std::cell::UnsafeCell::new(poll),
queues,
slots,
free_slots,
wait_slots,
free_wait_slots,
waiting_queue,
is_sleeping,
cancel_queue,
waker,
direct_fd_free,
});
}
}
let worker_states = worker_states.into_boxed_slice();
let _ = WORKERS.set(worker_states);
for worker_idx in 0..workers {
std::thread::Builder::new()
.name(format!("dtact-io-worker-{worker_idx}"))
.spawn(move || {
LOCAL_ALLOCATOR.with(|cell| {
*cell.borrow_mut() = Some(LocalAllocator {
thread_idx: worker_idx,
local_chunks: Vec::new(),
});
});
let state = &WORKERS.get().unwrap()[worker_idx];
#[cfg(target_os = "linux")]
run_linux_worker_loop(worker_idx, state);
#[cfg(not(target_os = "linux"))]
run_mio_worker_loop(worker_idx, state);
})
.expect("Failed to spawn dtact-io worker thread");
}
}
pub fn init(workers: usize) {
init_runtime(workers, 1024, 65536, 4096, &[]);
}
pub fn shutdown_runtime() {
SHUTDOWN.store(true, Ordering::Release);
if let Some(workers) = WORKERS.get() {
for state in workers {
#[cfg(target_os = "linux")]
let _ = unsafe {
libc::write(
state.wake_eventfd,
std::ptr::from_ref::<u64>(&1u64).cast::<libc::c_void>(),
8,
)
};
#[cfg(not(target_os = "linux"))]
state.waker.wake();
}
}
}
#[allow(clippy::too_many_lines)]
#[cfg(target_os = "linux")]
fn run_linux_worker_loop(worker_idx: usize, state: &WorkerState) {
if let Some(config) = GLOBAL_CONFIG.get()
&& let Some(&cpu_id) = config.pin_cpus.get(worker_idx)
{
let _ = pin_thread_to_cpu(cpu_id);
}
let ring = unsafe { &mut *state.ring.get() };
let mut eventfd_buf = 0u64;
let mut eventfd_submitted = false;
#[cfg(feature = "spin")]
let mut idle_streak: u32 = 0;
loop {
if SHUTDOWN.load(Ordering::Relaxed) {
break;
}
if !eventfd_submitted {
let sqe = io_uring::opcode::Read::new(
io_uring::types::Fd(state.wake_eventfd),
(&raw mut eventfd_buf).cast::<u8>(),
8,
)
.build()
.user_data(u64::MAX);
unsafe {
if ring.submission().push(&sqe).is_ok() {
eventfd_submitted = true;
}
}
}
let mut pushed_sqe = false;
for q in &state.queues {
while let Some(req) = q.pop() {
pushed_sqe = true;
let _ = submit_linux_request(state, &req);
}
}
while let Some(slot_idx) = state.cancel_queue.pop() {
pushed_sqe = true;
let sqe = io_uring::opcode::AsyncCancel::new(u64::from(slot_idx))
.build()
.user_data(u64::MAX - 1);
unsafe {
let _ = push_sqe(ring, &sqe);
}
}
let any_pending = state.queues.iter().any(|q| !q.is_empty());
let mut folded_wait = false;
if pushed_sqe || eventfd_submitted {
ring.submission().sync();
if pushed_sqe && !any_pending {
state.is_sleeping.store(true, Ordering::SeqCst);
fence(Ordering::SeqCst);
let missed = state.queues.iter().any(|q| !q.is_empty());
if missed {
state.is_sleeping.store(false, Ordering::SeqCst);
let sr = ring.submit();
io_trace!(
"[dtact-io] t={} loop submit(folded-missed) result={:?}",
trace_now_us(),
sr
);
} else {
let sr = ring.submit_and_wait(1);
state.is_sleeping.store(false, Ordering::Release);
io_trace!(
"[dtact-io] t={} loop submit_and_wait(folded) result={:?}",
trace_now_us(),
sr
);
}
folded_wait = true;
} else {
let should_enter = if state.sqpoll_enabled {
ring.submission().need_wakeup()
} else {
true
};
if should_enter {
let sr = ring.submit();
io_trace!(
"[dtact-io] t={} loop submit() result={:?}",
trace_now_us(),
sr
);
}
}
}
let mut has_completions = false;
let mut cq = ring.completion();
cq.sync();
let cq_len = cq.len();
if cq_len > 0 {
io_trace!("[dtact-io] t={} loop cq_len={}", trace_now_us(), cq_len);
}
for cqe in cq {
has_completions = true;
let user_data = cqe.user_data();
let res = cqe.result();
if user_data == u64::MAX {
eventfd_submitted = false;
} else if user_data == u64::MAX - 1 {
} else {
process_linux_completion(state, user_data as usize, res);
}
}
#[cfg(feature = "spin")]
if folded_wait || pushed_sqe || has_completions {
idle_streak = 0;
}
if !folded_wait && !pushed_sqe && !has_completions {
#[cfg(feature = "spin")]
{
if adaptive_idle_spin(state, idle_streak) {
continue;
}
}
#[cfg(feature = "spin")]
{
idle_streak = idle_streak.saturating_add(1);
}
state.is_sleeping.store(true, Ordering::SeqCst);
fence(Ordering::SeqCst);
let missed = state.queues.iter().any(|q| !q.is_empty());
if !any_pending && !missed {
let sr = ring.submit_and_wait(1);
io_trace!(
"[dtact-io] t={} loop submit_and_wait(idle) result={:?}",
trace_now_us(),
sr
);
} else if missed {
let sr = ring.submit();
io_trace!(
"[dtact-io] t={} loop submit(idle-missed) result={:?}",
trace_now_us(),
sr
);
}
state.is_sleeping.store(false, Ordering::Release);
}
}
}
#[cfg(all(target_os = "linux", feature = "spin"))]
fn adaptive_idle_spin(state: &WorkerState, idle_streak: u32) -> bool {
const ARM_AFTER: u32 = 2;
if idle_streak < ARM_AFTER {
return false;
}
let budget = 256u32.saturating_mul(idle_streak.saturating_sub(ARM_AFTER) + 1);
let budget = budget.min(4096);
for _ in 0..budget {
if state.queues.iter().any(|q| !q.is_empty()) {
return true;
}
core::hint::spin_loop();
}
false
}
#[cfg(target_os = "linux")]
unsafe fn push_sqe(
ring: &mut io_uring::IoUring,
sqe: &io_uring::squeue::Entry,
) -> Result<(), &'static str> {
loop {
let res = unsafe { ring.submission().push(sqe) };
if res == Ok(()) {
return Ok(());
}
let _ = ring.submit();
core::hint::spin_loop();
}
}
#[allow(clippy::too_many_lines)]
#[cfg(target_os = "linux")]
fn submit_linux_request(state: &WorkerState, req: &IoRequest) -> Result<(), &'static str> {
let ring = unsafe { &mut *state.ring.get() };
let sqe = match *req {
IoRequest::Read {
fd,
direct_fd_idx,
buf_ptr,
len,
offset,
slot_idx,
} => {
let use_fixed = direct_fd_idx != u32::MAX;
let target_fd = if use_fixed {
direct_fd_idx as i32
} else {
fd as i32
};
let mut s =
io_uring::opcode::Read::new(io_uring::types::Fd(target_fd), buf_ptr, len as u32)
.offset(offset as u64)
.build()
.user_data(slot_idx as u64);
if use_fixed {
s = s.flags(io_uring::squeue::Flags::FIXED_FILE);
}
s
}
IoRequest::Write {
fd,
direct_fd_idx,
buf_ptr,
len,
offset,
slot_idx,
} => {
let use_fixed = direct_fd_idx != u32::MAX;
let target_fd = if use_fixed {
direct_fd_idx as i32
} else {
fd as i32
};
let mut s =
io_uring::opcode::Write::new(io_uring::types::Fd(target_fd), buf_ptr, len as u32)
.offset(offset as u64)
.build()
.user_data(slot_idx as u64);
if use_fixed {
s = s.flags(io_uring::squeue::Flags::FIXED_FILE);
}
s
}
IoRequest::Accept {
fd,
direct_fd_idx,
slot_idx,
} => {
let use_fixed = direct_fd_idx != u32::MAX;
let target_fd = if use_fixed {
direct_fd_idx as i32
} else {
fd as i32
};
let mut s = io_uring::opcode::Accept::new(
io_uring::types::Fd(target_fd),
std::ptr::null_mut(),
std::ptr::null_mut(),
)
.build()
.user_data(slot_idx as u64);
if use_fixed {
s = s.flags(io_uring::squeue::Flags::FIXED_FILE);
}
s
}
IoRequest::Connect {
fd,
direct_fd_idx,
addr,
addr_len,
slot_idx,
} => {
let addr_ptr = (&raw const addr).cast::<libc::sockaddr>();
let use_fixed = direct_fd_idx != u32::MAX;
let target_fd = if use_fixed {
direct_fd_idx as i32
} else {
fd as i32
};
let mut s =
io_uring::opcode::Connect::new(io_uring::types::Fd(target_fd), addr_ptr, addr_len)
.build()
.user_data(slot_idx as u64);
if use_fixed {
s = s.flags(io_uring::squeue::Flags::FIXED_FILE);
}
s
}
IoRequest::SendTo {
fd,
direct_fd_idx,
msg_ptr,
slot_idx,
} => {
let use_fixed = direct_fd_idx != u32::MAX;
let target_fd = if use_fixed {
direct_fd_idx as i32
} else {
fd as i32
};
let mut s = io_uring::opcode::SendMsg::new(
io_uring::types::Fd(target_fd),
msg_ptr.cast_const(),
)
.build()
.user_data(slot_idx as u64);
if use_fixed {
s = s.flags(io_uring::squeue::Flags::FIXED_FILE);
}
s
}
IoRequest::RecvFrom {
fd,
direct_fd_idx,
msg_ptr,
slot_idx,
} => {
let use_fixed = direct_fd_idx != u32::MAX;
let target_fd = if use_fixed {
direct_fd_idx as i32
} else {
fd as i32
};
let mut s = io_uring::opcode::RecvMsg::new(io_uring::types::Fd(target_fd), msg_ptr)
.build()
.user_data(slot_idx as u64);
if use_fixed {
s = s.flags(io_uring::squeue::Flags::FIXED_FILE);
}
s
}
IoRequest::RegisterFile { fd, slot_idx } => {
if let Some(direct_idx) = state.direct_fd_free.pop() {
let fds = [fd];
let res = ring.submitter().register_files_update(direct_idx, &fds);
let out_res = match res {
Ok(_) => direct_idx as i32,
Err(e) => -(e.raw_os_error().unwrap_or(libc::EINVAL)),
};
process_linux_completion(state, slot_idx, out_res);
} else {
process_linux_completion(state, slot_idx, -libc::ENFILE);
}
return Ok(());
}
IoRequest::UnregisterFile {
direct_fd_idx,
slot_idx,
} => {
let fds = [-1];
let res = ring.submitter().register_files_update(direct_fd_idx, &fds);
state.direct_fd_free.push(direct_fd_idx);
let out_res = match res {
Ok(_) => 0,
Err(e) => -(e.raw_os_error().unwrap_or(libc::EINVAL)),
};
process_linux_completion(state, slot_idx, out_res);
return Ok(());
}
};
let user_data = sqe.get_user_data();
let r = unsafe { push_sqe(ring, &sqe) };
io_trace!(
"[dtact-io] t={} slot={} submit_linux_request pushed_local ok={}",
trace_now_us(),
user_data,
r.is_ok()
);
r
}
#[cfg(target_os = "linux")]
fn process_linux_completion(state: &WorkerState, slot_idx: usize, res: i32) {
let slot = &state.slots[slot_idx];
io_trace!(
"[dtact-io] t={} slot={} res={} B_kernel_complete",
trace_now_us(),
slot_idx,
res
);
slot.result.store(res, Ordering::Release);
slot.lock_waker();
let data = slot
.waker_data
.swap(std::ptr::null_mut(), Ordering::Relaxed);
let vtable = slot
.waker_vtable
.swap(std::ptr::null_mut(), Ordering::Relaxed);
slot.unlock_waker();
slot.completed.store(true, Ordering::Release);
if slot.dropped.load(Ordering::Acquire) {
state.free_slots.push(slot_idx as u32);
wake_next_waiting_fiber(state);
} else if !data.is_null() && !vtable.is_null() {
let raw = RawWaker::new(data.cast_const(), unsafe { &*vtable });
let w = unsafe { Waker::from_raw(raw) };
w.wake();
}
}
#[cfg(not(target_os = "linux"))]
struct FdState {
reader_waker: Option<Waker>,
writer_waker: Option<Waker>,
reader_slot: Option<usize>,
writer_slot: Option<usize>,
registered_interest: Option<mio::Interest>,
}
#[cfg(not(target_os = "linux"))]
impl FdState {
const fn new() -> Self {
Self {
reader_waker: None,
writer_waker: None,
reader_slot: None,
writer_slot: None,
registered_interest: None,
}
}
}
#[cfg(not(target_os = "linux"))]
fn ensure_fd_state(fd_states: &mut Vec<FdState>, fd: usize) {
if fd_states.len() <= fd {
fd_states.resize_with(fd + 1, FdState::new);
}
}
#[cfg(not(target_os = "linux"))]
fn install_interest(
state: &WorkerState,
fd_state: &mut FdState,
fd: u32,
slot_idx: usize,
is_reader: bool,
) -> bool {
let slot = &state.slots[slot_idx];
slot.lock_waker();
let data = slot
.waker_data
.swap(std::ptr::null_mut(), Ordering::Relaxed);
let vtable = slot
.waker_vtable
.swap(std::ptr::null_mut(), Ordering::Relaxed);
slot.unlock_waker();
let waker = if !data.is_null() && !vtable.is_null() {
let raw = RawWaker::new(data as *const (), unsafe { &*vtable });
Some(unsafe { Waker::from_raw(raw) })
} else {
None
};
if is_reader {
fd_state.reader_waker = waker;
fd_state.reader_slot = if fd_state.reader_waker.is_some() {
Some(slot_idx)
} else {
None
};
} else {
fd_state.writer_waker = waker;
fd_state.writer_slot = if fd_state.writer_waker.is_some() {
Some(slot_idx)
} else {
None
};
}
let interest = get_mio_interest(fd_state);
if fd_state.registered_interest == Some(interest) {
return true;
}
let res = unsafe {
let poll = &mut *state.poll.get();
poll.registry().reregister(
&mut mio::unix::SourceFd(&(fd as i32)),
mio::Token(fd as usize),
interest,
)
};
match res {
Ok(()) => {
fd_state.registered_interest = Some(interest);
true
}
Err(e) => {
io_trace!(
"[dtact-io] t={} fd={} reregister failed: {e}",
trace_now_us(),
fd
);
let woken = if is_reader {
fd_state.reader_waker.take()
} else {
fd_state.writer_waker.take()
};
if is_reader {
fd_state.reader_slot = None;
} else {
fd_state.writer_slot = None;
}
if let Some(w) = woken {
w.wake();
}
false
}
}
}
#[cfg(not(target_os = "linux"))]
fn get_mio_interest(fd_state: &FdState) -> mio::Interest {
let r = fd_state.reader_waker.is_some();
let w = fd_state.writer_waker.is_some();
if r && w {
mio::Interest::READABLE | mio::Interest::WRITABLE
} else if r {
mio::Interest::READABLE
} else if w {
mio::Interest::WRITABLE
} else {
mio::Interest::READABLE
}
}
#[cfg(not(target_os = "linux"))]
fn run_mio_worker_loop(worker_idx: usize, state: &WorkerState) {
if let Some(config) = GLOBAL_CONFIG.get() {
if let Some(&cpu_id) = config.pin_cpus.get(worker_idx) {
let _ = pin_thread_to_cpu(cpu_id);
}
}
let poll = unsafe { &mut *state.poll.get() };
let mut events = mio::Events::with_capacity(256);
let mut fd_states: Vec<FdState> = Vec::with_capacity(256);
loop {
if SHUTDOWN.load(Ordering::Relaxed) {
break;
}
let mut processed_any = false;
for q in state.queues.iter() {
while let Some(req) = q.pop() {
processed_any = true;
process_mio_request(state, &mut fd_states, req);
}
}
while let Some(slot_idx) = state.cancel_queue.pop() {
processed_any = true;
cancel_mio_slot(state, &mut fd_states, slot_idx as usize);
}
state.is_sleeping.store(true, Ordering::SeqCst);
fence(Ordering::SeqCst);
let mut any_pending = false;
for q in state.queues.iter() {
if !q.is_empty() {
any_pending = true;
break;
}
}
let poll_res = if !any_pending {
poll.poll(&mut events, Some(std::time::Duration::from_millis(10)))
} else {
poll.poll(&mut events, Some(std::time::Duration::from_millis(0)))
};
state.is_sleeping.store(false, Ordering::Release);
if poll_res.is_err() {
continue;
}
for event in events.iter() {
let token = event.token();
if token == mio::Token(0) {
continue;
}
let fd = token.0;
process_mio_event(
state,
&mut fd_states,
fd,
event.is_readable(),
event.is_writable(),
);
}
}
}
#[cfg(not(target_os = "linux"))]
fn process_mio_request(state: &WorkerState, fd_states: &mut Vec<FdState>, req: IoRequest) {
match req {
IoRequest::Read { fd, slot_idx, .. }
| IoRequest::Accept { fd, slot_idx, .. }
| IoRequest::RecvFrom { fd, slot_idx, .. } => {
ensure_fd_state(fd_states, fd as usize);
install_interest(state, &mut fd_states[fd as usize], fd, slot_idx, true);
}
IoRequest::Write { fd, slot_idx, .. }
| IoRequest::Connect { fd, slot_idx, .. }
| IoRequest::SendTo { fd, slot_idx, .. } => {
ensure_fd_state(fd_states, fd as usize);
install_interest(state, &mut fd_states[fd as usize], fd, slot_idx, false);
}
IoRequest::RegisterFile { fd, slot_idx } => {
let res = unsafe {
let poll = &mut *state.poll.get();
poll.registry().register(
&mut mio::unix::SourceFd(&fd),
mio::Token(fd as usize),
mio::Interest::READABLE | mio::Interest::WRITABLE,
)
};
match res {
Ok(()) => complete_mio_slot(state, slot_idx, fd),
Err(e) => {
let os_err = e.raw_os_error().unwrap_or(libc::EINVAL);
complete_mio_slot(state, slot_idx, -os_err);
}
}
}
IoRequest::UnregisterFile {
direct_fd_idx,
slot_idx,
} => {
let _ = unsafe {
let poll = &mut *state.poll.get();
poll.registry()
.deregister(&mut mio::unix::SourceFd(&(direct_fd_idx as i32)))
};
if let Some(fd_state) = fd_states.get_mut(direct_fd_idx as usize) {
fd_state.reader_waker = None;
fd_state.writer_waker = None;
fd_state.reader_slot = None;
fd_state.writer_slot = None;
fd_state.registered_interest = None;
}
complete_mio_slot(state, slot_idx, 0);
}
}
}
#[cfg(not(target_os = "linux"))]
fn cancel_mio_slot(state: &WorkerState, fd_states: &mut Vec<FdState>, slot_idx: usize) {
let slot = &state.slots[slot_idx];
let fd = slot.origin_fd.load(Ordering::Relaxed);
if fd != u32::MAX {
ensure_fd_state(fd_states, fd as usize);
let fd_state = &mut fd_states[fd as usize];
let mut touched = false;
if fd_state.reader_slot == Some(slot_idx) {
fd_state.reader_waker = None;
fd_state.reader_slot = None;
touched = true;
}
if fd_state.writer_slot == Some(slot_idx) {
fd_state.writer_waker = None;
fd_state.writer_slot = None;
touched = true;
}
if touched {
let interest = get_mio_interest(fd_state);
if fd_state.registered_interest != Some(interest) {
let res = unsafe {
let poll = &mut *state.poll.get();
poll.registry().reregister(
&mut mio::unix::SourceFd(&(fd as i32)),
mio::Token(fd as usize),
interest,
)
};
if res.is_ok() {
fd_state.registered_interest = Some(interest);
}
}
}
}
state.free_slots.push(slot_idx as u32);
wake_next_waiting_fiber(state);
}
#[cfg(not(target_os = "linux"))]
fn process_mio_event(
_state: &WorkerState,
fd_states: &mut Vec<FdState>,
fd: usize,
readable: bool,
writable: bool,
) {
ensure_fd_state(fd_states, fd);
let fd_state = &mut fd_states[fd];
if readable {
if let Some(w) = fd_state.reader_waker.take() {
w.wake();
}
fd_state.reader_slot = None;
}
if writable {
if let Some(w) = fd_state.writer_waker.take() {
w.wake();
}
fd_state.writer_slot = None;
}
let interest = get_mio_interest(fd_state);
if fd_state.registered_interest != Some(interest) {
let res = unsafe {
let poll = &mut *_state.poll.get();
poll.registry().reregister(
&mut mio::unix::SourceFd(&(fd as i32)),
mio::Token(fd),
interest,
)
};
if res.is_ok() {
fd_state.registered_interest = Some(interest);
}
}
}
#[cfg(not(target_os = "linux"))]
fn complete_mio_slot(state: &WorkerState, slot_idx: usize, res: i32) {
let slot = &state.slots[slot_idx];
slot.result.store(res, Ordering::Release);
slot.lock_waker();
let data = slot
.waker_data
.swap(std::ptr::null_mut(), Ordering::Relaxed);
let vtable = slot
.waker_vtable
.swap(std::ptr::null_mut(), Ordering::Relaxed);
slot.unlock_waker();
slot.completed.store(true, Ordering::Release);
if !data.is_null() && !vtable.is_null() {
let raw = RawWaker::new(data as *const (), unsafe { &*vtable });
let w = unsafe { Waker::from_raw(raw) };
w.wake();
}
}
pub struct DtactIoFuture {
pub worker_idx: usize,
pub fd: u32,
pub direct_fd_idx: u32,
pub op: OpCode,
pub buf_ptr: *mut u8,
pub len: usize,
pub offset: i64,
pub addr: Option<libc::sockaddr_storage>,
pub addr_len: libc::socklen_t,
pub slot_idx: Option<usize>,
pub msg_ptr: *mut libc::msghdr,
}
unsafe impl Send for DtactIoFuture {}
unsafe impl Sync for DtactIoFuture {}
impl DtactIoFuture {
#[allow(clippy::too_many_arguments)]
pub const fn new(
worker_idx: usize,
fd: u32,
direct_fd_idx: u32,
op: OpCode,
buf_ptr: *mut u8,
len: usize,
offset: i64,
addr: Option<libc::sockaddr_storage>,
addr_len: libc::socklen_t,
slot_idx: Option<usize>,
) -> Self {
Self {
worker_idx,
fd,
direct_fd_idx,
op,
buf_ptr,
len,
offset,
addr,
addr_len,
slot_idx,
msg_ptr: std::ptr::null_mut(),
}
}
const fn create_io_request(&self, slot_idx: usize) -> IoRequest {
match self.op {
OpCode::SendTo => IoRequest::SendTo {
fd: self.fd,
direct_fd_idx: self.direct_fd_idx,
msg_ptr: self.msg_ptr,
slot_idx,
},
OpCode::RecvFrom => IoRequest::RecvFrom {
fd: self.fd,
direct_fd_idx: self.direct_fd_idx,
msg_ptr: self.msg_ptr,
slot_idx,
},
OpCode::Read => IoRequest::Read {
fd: self.fd,
direct_fd_idx: self.direct_fd_idx,
buf_ptr: self.buf_ptr,
len: self.len,
offset: self.offset,
slot_idx,
},
OpCode::Write => IoRequest::Write {
fd: self.fd,
direct_fd_idx: self.direct_fd_idx,
buf_ptr: self.buf_ptr,
len: self.len,
offset: self.offset,
slot_idx,
},
OpCode::Accept => IoRequest::Accept {
fd: self.fd,
direct_fd_idx: self.direct_fd_idx,
slot_idx,
},
OpCode::Connect => IoRequest::Connect {
fd: self.fd,
direct_fd_idx: self.direct_fd_idx,
addr: self.addr.unwrap(),
addr_len: self.addr_len,
slot_idx,
},
}
}
#[cfg(not(target_os = "linux"))]
fn execute_syscall(&self) -> std::io::Result<usize> {
let res = match self.op {
OpCode::Read => {
let buf_ptr = self.buf_ptr;
let len = self.len;
unsafe { libc::read(self.fd as i32, buf_ptr as *mut libc::c_void, len) }
}
OpCode::Write => {
let buf_ptr = self.buf_ptr;
let len = self.len;
unsafe { libc::write(self.fd as i32, buf_ptr as *const libc::c_void, len) }
}
OpCode::Accept => unsafe {
libc::accept(self.fd as i32, std::ptr::null_mut(), std::ptr::null_mut()) as isize
},
OpCode::SendTo => unsafe { libc::sendmsg(self.fd as i32, self.msg_ptr, 0) },
OpCode::RecvFrom => unsafe { libc::recvmsg(self.fd as i32, self.msg_ptr, 0) },
OpCode::Connect => {
let addr_ptr =
&self.addr.unwrap() as *const libc::sockaddr_storage as *const libc::sockaddr;
let res = unsafe { libc::connect(self.fd as i32, addr_ptr, self.addr_len) };
if res < 0 {
let err = std::io::Error::last_os_error();
if err.raw_os_error() == Some(libc::EISCONN) {
return Ok(0);
}
return Err(err);
}
res as isize
}
};
if res < 0 {
Err(std::io::Error::last_os_error())
} else {
Ok(res as usize)
}
}
}
impl Future for DtactIoFuture {
type Output = std::io::Result<usize>;
#[allow(clippy::too_many_lines)]
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
#[cfg(target_os = "linux")]
{
let slot_idx = if let Some(idx) = self.slot_idx {
idx
} else {
let state = &WORKERS.get().unwrap()[self.worker_idx];
let idx = match state.free_slots.pop() {
Some(i) => i as usize,
None => {
if let Some(wait_idx) = state.free_wait_slots.pop() {
let wait_slot = &state.wait_slots[wait_idx as usize];
wait_slot
.waker_data
.store(cx.waker().data().cast_mut(), Ordering::Relaxed);
wait_slot.waker_vtable.store(
std::ptr::from_ref::<RawWakerVTable>(cx.waker().vtable())
.cast_mut(),
Ordering::Relaxed,
);
state.waiting_queue.push(wait_idx);
if let Some(i) = state.free_slots.pop() {
wait_slot
.waker_data
.store(std::ptr::null_mut(), Ordering::Relaxed);
wait_slot
.waker_vtable
.store(std::ptr::null_mut(), Ordering::Relaxed);
i as usize
} else {
return Poll::Pending;
}
} else {
cx.waker().wake_by_ref();
return Poll::Pending;
}
}
};
let slot = &state.slots[idx];
slot.completed.store(false, Ordering::Relaxed);
slot.dropped.store(false, Ordering::Relaxed);
slot.origin_fd.store(self.fd, Ordering::Relaxed);
slot.lock_waker();
slot.waker_data
.store(cx.waker().data().cast_mut(), Ordering::Relaxed);
slot.waker_vtable.store(
std::ptr::from_ref::<RawWakerVTable>(cx.waker().vtable()).cast_mut(),
Ordering::Relaxed,
);
slot.unlock_waker();
let req = self.create_io_request(idx);
let q_idx = get_or_init_local_allocator().unwrap_or(0);
let queue = &state.queues[q_idx];
if queue.push(req).is_err() {
slot.lock_waker();
slot.waker_data
.store(std::ptr::null_mut(), Ordering::Relaxed);
slot.waker_vtable
.store(std::ptr::null_mut(), Ordering::Relaxed);
slot.unlock_waker();
state.free_slots.push(idx as u32);
wake_next_waiting_fiber(state);
cx.waker().wake_by_ref();
return Poll::Pending;
}
fence(Ordering::SeqCst);
if state.is_sleeping.load(Ordering::SeqCst) {
unsafe {
let _ = libc::write(
state.wake_eventfd,
std::ptr::from_ref::<u64>(&1u64).cast::<libc::c_void>(),
8,
);
}
}
io_trace!(
"[dtact-io] t={} slot={} fd={} op={:?} A_submit",
trace_now_us(),
idx,
self.fd,
self.op
);
self.slot_idx = Some(idx);
idx
};
let state = &WORKERS.get().unwrap()[self.worker_idx];
let slot = &state.slots[slot_idx];
if slot.completed.load(Ordering::Acquire) {
let res = slot.result.load(Ordering::Acquire);
io_trace!(
"[dtact-io] t={} slot={} res={} C_fiber_resumed",
trace_now_us(),
slot_idx,
res
);
slot.lock_waker();
slot.waker_data
.store(std::ptr::null_mut(), Ordering::Relaxed);
slot.waker_vtable
.store(std::ptr::null_mut(), Ordering::Relaxed);
slot.unlock_waker();
state.free_slots.push(slot_idx as u32);
self.slot_idx = None;
wake_next_waiting_fiber(state);
if res < 0 {
Poll::Ready(Err(std::io::Error::from_raw_os_error(-res)))
} else {
Poll::Ready(Ok(res as usize))
}
} else {
let new_data = cx.waker().data().cast_mut();
let new_vtable =
std::ptr::from_ref::<RawWakerVTable>(cx.waker().vtable()).cast_mut();
slot.lock_waker();
let old_data = slot.waker_data.load(Ordering::Relaxed);
let old_vtable = slot.waker_vtable.load(Ordering::Relaxed);
if old_data != new_data || old_vtable != new_vtable {
slot.waker_data.store(new_data, Ordering::Relaxed);
slot.waker_vtable.store(new_vtable, Ordering::Relaxed);
}
slot.unlock_waker();
Poll::Pending
}
}
#[cfg(not(target_os = "linux"))]
{
let res = self.execute_syscall();
if self.slot_idx.is_some()
&& !matches!(res, Err(ref e) if e.kind() == std::io::ErrorKind::WouldBlock)
{
let state = &WORKERS.get().unwrap()[self.worker_idx];
state.free_slots.push(self.slot_idx.unwrap() as u32);
self.slot_idx = None;
wake_next_waiting_fiber(state);
}
match res {
Ok(n) => Poll::Ready(Ok(n)),
Err(ref e) if e.kind() == std::io::ErrorKind::WouldBlock => {
let slot_idx = match self.slot_idx {
Some(idx) => idx,
None => {
let state = &WORKERS.get().unwrap()[self.worker_idx];
let idx = match state.free_slots.pop() {
Some(i) => i as usize,
None => {
if let Some(wait_idx) = state.free_wait_slots.pop() {
let wait_slot = &state.wait_slots[wait_idx as usize];
wait_slot
.waker_data
.store(cx.waker().data() as *mut (), Ordering::Relaxed);
wait_slot.waker_vtable.store(
cx.waker().vtable() as *const RawWakerVTable as *mut _,
Ordering::Relaxed,
);
state.waiting_queue.push(wait_idx);
if let Some(i) = state.free_slots.pop() {
wait_slot
.waker_data
.store(std::ptr::null_mut(), Ordering::Relaxed);
wait_slot
.waker_vtable
.store(std::ptr::null_mut(), Ordering::Relaxed);
i as usize
} else {
return Poll::Pending;
}
} else {
cx.waker().wake_by_ref();
return Poll::Pending;
}
}
};
let slot = &state.slots[idx];
slot.completed.store(false, Ordering::Relaxed);
slot.dropped.store(false, Ordering::Relaxed);
slot.origin_fd.store(self.fd, Ordering::Relaxed);
let raw = cx.waker().as_raw();
slot.lock_waker();
slot.waker_data
.store(raw.data() as *mut (), Ordering::Relaxed);
slot.waker_vtable.store(
raw.vtable() as *const RawWakerVTable as *mut _,
Ordering::Relaxed,
);
slot.unlock_waker();
let req = self.create_io_request(idx);
let q_idx = get_or_init_local_allocator().unwrap_or(0);
let queue = &state.queues[q_idx];
if queue.push(req).is_err() {
slot.lock_waker();
slot.waker_data
.store(std::ptr::null_mut(), Ordering::Relaxed);
slot.waker_vtable
.store(std::ptr::null_mut(), Ordering::Relaxed);
slot.unlock_waker();
state.free_slots.push(idx as u32);
wake_next_waiting_fiber(state);
cx.waker().wake_by_ref();
return Poll::Pending;
}
fence(Ordering::SeqCst);
if state.is_sleeping.load(Ordering::SeqCst) {
state.waker.wake();
}
self.slot_idx = Some(idx);
idx
}
};
let state = &WORKERS.get().unwrap()[self.worker_idx];
let slot = &state.slots[slot_idx];
let raw = cx.waker().as_raw();
let new_data = raw.data() as *mut ();
let new_vtable = raw.vtable() as *const RawWakerVTable as *mut _;
slot.lock_waker();
let old_data = slot.waker_data.load(Ordering::Relaxed);
let old_vtable = slot.waker_vtable.load(Ordering::Relaxed);
let mut changed = false;
if old_data != new_data || old_vtable != new_vtable {
slot.waker_data.store(new_data, Ordering::Relaxed);
slot.waker_vtable.store(new_vtable, Ordering::Relaxed);
changed = true;
}
slot.unlock_waker();
if changed {
let req = self.create_io_request(slot_idx);
let q_idx = get_or_init_local_allocator().unwrap_or(0);
let _ = state.queues[q_idx].push(req);
fence(Ordering::SeqCst);
if state.is_sleeping.load(Ordering::SeqCst) {
state.waker.wake();
}
}
Poll::Pending
}
Err(e) => Poll::Ready(Err(e)),
}
}
}
}
impl Drop for DtactIoFuture {
fn drop(&mut self) {
let Some(idx) = self.slot_idx else { return };
let Some(state) = WORKERS.get().and_then(|w| w.get(self.worker_idx)) else {
return;
};
let slot = &state.slots[idx];
slot.lock_waker();
slot.waker_data
.store(std::ptr::null_mut(), Ordering::Relaxed);
slot.waker_vtable
.store(std::ptr::null_mut(), Ordering::Relaxed);
slot.unlock_waker();
if slot.completed.load(Ordering::Acquire) {
state.free_slots.push(idx as u32);
wake_next_waiting_fiber(state);
return;
}
slot.dropped.store(true, Ordering::Release);
state.cancel_queue.push(idx as u32);
fence(Ordering::SeqCst);
#[cfg(target_os = "linux")]
{
if state.is_sleeping.load(Ordering::SeqCst) {
unsafe {
let _ = libc::write(
state.wake_eventfd,
std::ptr::from_ref::<u64>(&1u64).cast::<libc::c_void>(),
8,
);
}
}
}
#[cfg(not(target_os = "linux"))]
{
if state.is_sleeping.load(Ordering::SeqCst) {
state.waker.wake();
}
}
}
}
pub struct DtactTcpStream {
inner: std::net::TcpStream,
direct_fd_idx: u32,
worker_idx: usize,
}
impl DtactTcpStream {
pub fn from_std(stream: std::net::TcpStream) -> std::io::Result<Self> {
let fd = stream.as_raw_fd();
stream.set_nonblocking(true)?;
stream.set_nodelay(true)?;
let num_workers = GLOBAL_CONFIG.get().map_or(1, |c| c.workers);
let worker_idx = fd as usize % num_workers;
let state = &WORKERS.get().unwrap()[worker_idx];
let direct_fd_idx = register_fd_sync(state, fd);
Ok(Self {
inner: stream,
direct_fd_idx,
worker_idx,
})
}
pub async fn read(&self, buf: &mut [u8]) -> std::io::Result<usize> {
if buf.is_empty() {
return Ok(0);
}
let res = unsafe {
let r = libc::read(
self.inner.as_raw_fd(),
buf.as_mut_ptr().cast::<libc::c_void>(),
buf.len(),
);
match r.cmp(&0) {
std::cmp::Ordering::Greater => Ok(r as usize),
std::cmp::Ordering::Equal => Ok(0), std::cmp::Ordering::Less => Err(std::io::Error::last_os_error()),
}
};
match res {
Ok(n) => return Ok(n),
Err(e) => {
if e.kind() != std::io::ErrorKind::WouldBlock {
return Err(e);
}
}
}
DtactIoFuture {
worker_idx: self.worker_idx,
fd: self.inner.as_raw_fd() as u32,
direct_fd_idx: self.direct_fd_idx,
op: OpCode::Read,
buf_ptr: buf.as_mut_ptr(),
len: buf.len(),
offset: 0,
addr: None,
addr_len: 0,
slot_idx: None,
msg_ptr: std::ptr::null_mut(),
}
.await
.map(|n| n.min(buf.len()))
}
pub async fn write(&self, buf: &[u8]) -> std::io::Result<usize> {
if buf.is_empty() {
return Ok(0);
}
let res = unsafe {
let r = libc::write(
self.inner.as_raw_fd(),
buf.as_ptr().cast::<libc::c_void>(),
buf.len(),
);
if r >= 0 {
Ok(r as usize)
} else {
Err(std::io::Error::last_os_error())
}
};
match res {
Ok(n) => return Ok(n),
Err(e) => {
if e.kind() != std::io::ErrorKind::WouldBlock {
return Err(e);
}
}
}
DtactIoFuture {
worker_idx: self.worker_idx,
fd: self.inner.as_raw_fd() as u32,
direct_fd_idx: self.direct_fd_idx,
op: OpCode::Write,
buf_ptr: buf.as_ptr().cast_mut(),
len: buf.len(),
offset: 0,
addr: None,
addr_len: 0,
slot_idx: None,
msg_ptr: std::ptr::null_mut(),
}
.await
}
#[allow(clippy::too_many_lines)]
pub async fn connect(addr: std::net::SocketAddr) -> std::io::Result<Self> {
let domain = match addr {
std::net::SocketAddr::V4(_) => libc::AF_INET,
std::net::SocketAddr::V6(_) => libc::AF_INET6,
};
let fd = unsafe {
libc::socket(
domain,
libc::SOCK_STREAM | libc::SOCK_CLOEXEC | libc::SOCK_NONBLOCK,
0,
)
};
if fd < 0 {
return Err(std::io::Error::last_os_error());
}
let stream = unsafe { std::net::TcpStream::from_raw_fd(fd) };
stream.set_nodelay(true)?;
let num_workers = GLOBAL_CONFIG.get().map_or(1, |c| c.workers);
let worker_idx = fd as usize % num_workers;
let state = &WORKERS.get().unwrap()[worker_idx];
let direct_fd_idx = register_fd_sync(state, fd);
let (libc_addr, addr_len) = socket_addr_to_libc(addr);
let connect_res = unsafe {
libc::connect(
fd,
(&raw const libc_addr).cast::<libc::sockaddr>(),
addr_len,
)
};
if connect_res == 0 {
return Ok(Self {
inner: stream,
direct_fd_idx,
worker_idx,
});
}
let err = std::io::Error::last_os_error();
#[cfg(target_os = "windows")]
let is_in_progress = err.raw_os_error()
== Some(windows_sys::Win32::Networking::WinSock::WSAEWOULDBLOCK as i32);
#[cfg(not(target_os = "windows"))]
let is_in_progress = err.raw_os_error() == Some(libc::EINPROGRESS);
if !is_in_progress {
return Err(err);
}
let mut pollfd = libc::pollfd {
fd,
events: libc::POLLOUT,
revents: 0,
};
let poll_res = unsafe { libc::poll(&raw mut pollfd, 1, 0) };
if poll_res > 0 {
if (pollfd.revents & libc::POLLOUT) != 0 {
let mut err_code: libc::c_int = 0;
let mut err_len = std::mem::size_of::<libc::c_int>() as libc::socklen_t;
let sockopt_res = unsafe {
libc::getsockopt(
fd,
libc::SOL_SOCKET,
libc::SO_ERROR,
(&raw mut err_code).cast::<libc::c_void>(),
&raw mut err_len,
)
};
if sockopt_res == 0 && err_code == 0 {
return Ok(Self {
inner: stream,
direct_fd_idx,
worker_idx,
});
}
let os_err = if err_code != 0 {
err_code
} else {
libc::ECONNREFUSED
};
return Err(std::io::Error::from_raw_os_error(os_err));
} else if (pollfd.revents & (libc::POLLERR | libc::POLLHUP)) != 0 {
return Err(std::io::Error::new(
std::io::ErrorKind::ConnectionRefused,
"connect failed",
));
}
}
let connect_res = DtactIoFuture {
worker_idx,
fd: fd as u32,
direct_fd_idx,
op: OpCode::Connect,
buf_ptr: std::ptr::null_mut(),
len: 0,
offset: 0,
addr: Some(libc_addr),
addr_len,
slot_idx: None,
msg_ptr: std::ptr::null_mut(),
}
.await;
match connect_res {
Ok(_) => Ok(Self {
inner: stream,
direct_fd_idx,
worker_idx,
}),
Err(e) => Err(e),
}
}
}
impl Drop for DtactTcpStream {
fn drop(&mut self) {
if let Some(workers) = WORKERS.get()
&& let Some(state) = workers.get(self.worker_idx)
{
unregister_fd_sync(state, self.direct_fd_idx);
}
}
}
impl crate::io::AsyncRead for DtactTcpStream {
async fn read(&self, buf: &mut [u8]) -> std::io::Result<usize> {
self.read(buf).await
}
}
impl crate::io::AsyncWrite for DtactTcpStream {
async fn write(&self, buf: &[u8]) -> std::io::Result<usize> {
self.write(buf).await
}
}
pub struct DtactTcpListener {
inner: std::net::TcpListener,
direct_fd_idx: u32,
worker_idx: usize,
}
impl DtactTcpListener {
pub fn from_std(listener: std::net::TcpListener) -> std::io::Result<Self> {
let fd = listener.as_raw_fd();
listener.set_nonblocking(true)?;
let num_workers = GLOBAL_CONFIG.get().map_or(1, |c| c.workers);
let worker_idx = fd as usize % num_workers;
let state = &WORKERS.get().unwrap()[worker_idx];
let direct_fd_idx = register_fd_sync(state, fd);
Ok(Self {
inner: listener,
direct_fd_idx,
worker_idx,
})
}
pub async fn accept(&self) -> std::io::Result<(DtactTcpStream, std::net::SocketAddr)> {
let res = unsafe {
let mut addr: libc::sockaddr_storage = std::mem::zeroed();
let mut len = std::mem::size_of::<libc::sockaddr_storage>() as libc::socklen_t;
let r = libc::accept(
self.inner.as_raw_fd(),
(&raw mut addr).cast::<libc::sockaddr>(),
&raw mut len,
);
if r >= 0 {
Ok((r, addr, len))
} else {
Err(std::io::Error::last_os_error())
}
};
match res {
Ok((client_fd, addr, len)) => {
let peer_addr = sockaddr_storage_to_socketaddr(&addr, len);
unsafe { libc::fcntl(client_fd, libc::F_SETFL, libc::O_NONBLOCK) };
let stream = unsafe { std::net::TcpStream::from_raw_fd(client_fd) };
let client_stream = DtactTcpStream::from_std(stream)?;
return Ok((client_stream, peer_addr));
}
Err(e) => {
if e.kind() != std::io::ErrorKind::WouldBlock {
return Err(e);
}
}
}
let res = DtactIoFuture {
worker_idx: self.worker_idx,
fd: self.inner.as_raw_fd() as u32,
direct_fd_idx: self.direct_fd_idx,
op: OpCode::Accept,
buf_ptr: std::ptr::null_mut(),
len: 0,
offset: 0,
addr: None,
addr_len: 0,
slot_idx: None,
msg_ptr: std::ptr::null_mut(),
}
.await?;
let client_fd = res as RawFd;
unsafe { libc::fcntl(client_fd, libc::F_SETFL, libc::O_NONBLOCK) };
let stream = unsafe { std::net::TcpStream::from_raw_fd(client_fd) };
let peer_addr = stream.peer_addr()?;
let client_stream = DtactTcpStream::from_std(stream)?;
Ok((client_stream, peer_addr))
}
}
impl Drop for DtactTcpListener {
fn drop(&mut self) {
if let Some(workers) = WORKERS.get()
&& let Some(state) = workers.get(self.worker_idx)
{
unregister_fd_sync(state, self.direct_fd_idx);
}
}
}
pub struct DtactUdpSocket {
inner: std::net::UdpSocket,
direct_fd_idx: u32,
worker_idx: usize,
}
impl DtactUdpSocket {
pub async fn bind(addr: std::net::SocketAddr) -> std::io::Result<Self> {
let sock = std::net::UdpSocket::bind(addr)?;
Self::from_std(sock)
}
pub fn from_std(socket: std::net::UdpSocket) -> std::io::Result<Self> {
let fd = socket.as_raw_fd();
socket.set_nonblocking(true)?;
let num_workers = GLOBAL_CONFIG.get().map_or(1, |c| c.workers);
let worker_idx = fd as usize % num_workers;
let state = &WORKERS.get().unwrap()[worker_idx];
let direct_fd_idx = register_fd_sync(state, fd);
Ok(Self {
inner: socket,
direct_fd_idx,
worker_idx,
})
}
pub fn local_addr(&self) -> std::io::Result<std::net::SocketAddr> {
self.inner.local_addr()
}
pub async fn send_to(
&self,
buf: &[u8],
target: std::net::SocketAddr,
) -> std::io::Result<usize> {
struct SendToState {
storage: libc::sockaddr_storage,
iov: libc::iovec,
msg: libc::msghdr,
}
unsafe impl Send for SendToState {}
let (storage, addr_len) = socket_addr_to_libc(target);
let mut state = SendToState {
storage,
iov: libc::iovec {
iov_base: buf.as_ptr().cast_mut().cast::<libc::c_void>(),
iov_len: buf.len(),
},
msg: unsafe { std::mem::zeroed() },
};
state.msg.msg_name = std::ptr::addr_of_mut!(state.storage).cast::<libc::c_void>();
state.msg.msg_namelen = addr_len;
state.msg.msg_iov = &raw mut state.iov;
state.msg.msg_iovlen = 1;
let r = unsafe { libc::sendmsg(self.inner.as_raw_fd(), &raw const state.msg, 0) };
if r >= 0 {
return Ok(r as usize);
}
let e = std::io::Error::last_os_error();
if e.kind() != std::io::ErrorKind::WouldBlock {
return Err(e);
}
let mut fut = DtactIoFuture::new(
self.worker_idx,
self.inner.as_raw_fd() as u32,
self.direct_fd_idx,
OpCode::SendTo,
std::ptr::null_mut(),
0,
0,
None,
0,
None,
);
fut.msg_ptr = &raw mut state.msg;
fut.await
}
pub async fn recv_from(
&self,
buf: &mut [u8],
) -> std::io::Result<(usize, std::net::SocketAddr)> {
struct RecvFromState {
storage: libc::sockaddr_storage,
iov: libc::iovec,
msg: libc::msghdr,
}
unsafe impl Send for RecvFromState {}
let mut state = RecvFromState {
storage: unsafe { std::mem::zeroed() },
iov: libc::iovec {
iov_base: buf.as_mut_ptr().cast::<libc::c_void>(),
iov_len: buf.len(),
},
msg: unsafe { std::mem::zeroed() },
};
state.msg.msg_name = std::ptr::addr_of_mut!(state.storage).cast::<libc::c_void>();
state.msg.msg_namelen = std::mem::size_of::<libc::sockaddr_storage>() as libc::socklen_t;
state.msg.msg_iov = &raw mut state.iov;
state.msg.msg_iovlen = 1;
let r = unsafe { libc::recvmsg(self.inner.as_raw_fd(), &raw mut state.msg, 0) };
if r >= 0 {
let from = sockaddr_storage_to_socketaddr(&state.storage, state.msg.msg_namelen);
return Ok((r as usize, from));
}
let e = std::io::Error::last_os_error();
if e.kind() != std::io::ErrorKind::WouldBlock {
return Err(e);
}
let mut fut = DtactIoFuture::new(
self.worker_idx,
self.inner.as_raw_fd() as u32,
self.direct_fd_idx,
OpCode::RecvFrom,
std::ptr::null_mut(),
0,
0,
None,
0,
None,
);
fut.msg_ptr = &raw mut state.msg;
let n = fut.await?;
let from = sockaddr_storage_to_socketaddr(&state.storage, state.msg.msg_namelen);
Ok((n, from))
}
pub async fn connect(&self, addr: std::net::SocketAddr) -> std::io::Result<()> {
self.inner.connect(addr)
}
pub async fn send(&self, buf: &[u8]) -> std::io::Result<usize> {
if buf.is_empty() {
return Ok(0);
}
let r = unsafe {
libc::send(
self.inner.as_raw_fd(),
buf.as_ptr().cast::<libc::c_void>(),
buf.len(),
0,
)
};
if r >= 0 {
return Ok(r as usize);
}
let e = std::io::Error::last_os_error();
if e.kind() != std::io::ErrorKind::WouldBlock {
return Err(e);
}
DtactIoFuture {
worker_idx: self.worker_idx,
fd: self.inner.as_raw_fd() as u32,
direct_fd_idx: self.direct_fd_idx,
op: OpCode::Write,
buf_ptr: buf.as_ptr().cast_mut(),
len: buf.len(),
offset: 0,
addr: None,
addr_len: 0,
slot_idx: None,
msg_ptr: std::ptr::null_mut(),
}
.await
}
pub async fn recv(&self, buf: &mut [u8]) -> std::io::Result<usize> {
if buf.is_empty() {
return Ok(0);
}
let r = unsafe {
libc::recv(
self.inner.as_raw_fd(),
buf.as_mut_ptr().cast::<libc::c_void>(),
buf.len(),
0,
)
};
if r >= 0 {
return Ok(r as usize);
}
let e = std::io::Error::last_os_error();
if e.kind() != std::io::ErrorKind::WouldBlock {
return Err(e);
}
DtactIoFuture {
worker_idx: self.worker_idx,
fd: self.inner.as_raw_fd() as u32,
direct_fd_idx: self.direct_fd_idx,
op: OpCode::Read,
buf_ptr: buf.as_mut_ptr(),
len: buf.len(),
offset: 0,
addr: None,
addr_len: 0,
slot_idx: None,
msg_ptr: std::ptr::null_mut(),
}
.await
}
}
impl Drop for DtactUdpSocket {
fn drop(&mut self) {
if let Some(workers) = WORKERS.get()
&& let Some(state) = workers.get(self.worker_idx)
{
unregister_fd_sync(state, self.direct_fd_idx);
}
}
}
pub struct DtactUnixStream {
inner: std::os::unix::net::UnixStream,
direct_fd_idx: u32,
worker_idx: usize,
read_backpressured: std::sync::atomic::AtomicBool,
write_backpressured: std::sync::atomic::AtomicBool,
}
impl DtactUnixStream {
pub fn from_std(stream: std::os::unix::net::UnixStream) -> std::io::Result<Self> {
let fd = stream.as_raw_fd();
stream.set_nonblocking(true)?;
let num_workers = GLOBAL_CONFIG.get().map_or(1, |c| c.workers);
let worker_idx = WORKER_ROUND_ROBIN.fetch_add(1, Ordering::Relaxed) % num_workers;
let state = &WORKERS.get().unwrap()[worker_idx];
let direct_fd_idx = register_fd_sync(state, fd);
Ok(Self {
inner: stream,
direct_fd_idx,
worker_idx,
read_backpressured: std::sync::atomic::AtomicBool::new(false),
write_backpressured: std::sync::atomic::AtomicBool::new(false),
})
}
pub async fn read(&self, buf: &mut [u8]) -> std::io::Result<usize> {
if buf.is_empty() {
return Ok(0);
}
if !self.read_backpressured.load(Ordering::Relaxed) {
let res = unsafe {
let r = libc::read(
self.inner.as_raw_fd(),
buf.as_mut_ptr().cast::<libc::c_void>(),
buf.len(),
);
match r.cmp(&0) {
std::cmp::Ordering::Greater => Ok(r as usize),
std::cmp::Ordering::Equal => Ok(0), std::cmp::Ordering::Less => Err(std::io::Error::last_os_error()),
}
};
match res {
Ok(n) => return Ok(n),
Err(e) => {
if e.kind() != std::io::ErrorKind::WouldBlock {
return Err(e);
}
}
}
self.read_backpressured.store(true, Ordering::Relaxed);
}
let future = DtactIoFuture {
worker_idx: self.worker_idx,
fd: self.inner.as_raw_fd() as u32,
direct_fd_idx: self.direct_fd_idx,
op: OpCode::Read,
buf_ptr: buf.as_mut_ptr(),
len: buf.len(),
offset: 0,
addr: None,
addr_len: 0,
slot_idx: None,
msg_ptr: std::ptr::null_mut(),
}
.await;
self.read_backpressured.store(false, Ordering::Relaxed);
future.map(|n| n.min(buf.len()))
}
pub async fn write(&self, buf: &[u8]) -> std::io::Result<usize> {
if buf.is_empty() {
return Ok(0);
}
if !self.write_backpressured.load(Ordering::Relaxed) {
let res = unsafe {
let r = libc::write(
self.inner.as_raw_fd(),
buf.as_ptr().cast::<libc::c_void>(),
buf.len(),
);
if r >= 0 {
Ok(r as usize)
} else {
Err(std::io::Error::last_os_error())
}
};
match res {
Ok(n) => return Ok(n),
Err(e) => {
if e.kind() != std::io::ErrorKind::WouldBlock {
return Err(e);
}
}
}
self.write_backpressured.store(true, Ordering::Relaxed);
}
let future = DtactIoFuture {
worker_idx: self.worker_idx,
fd: self.inner.as_raw_fd() as u32,
direct_fd_idx: self.direct_fd_idx,
op: OpCode::Write,
buf_ptr: buf.as_ptr().cast_mut(),
len: buf.len(),
offset: 0,
addr: None,
addr_len: 0,
slot_idx: None,
msg_ptr: std::ptr::null_mut(),
};
self.write_backpressured.store(false, Ordering::Relaxed);
future.await
}
#[allow(clippy::too_many_lines)]
pub async fn connect(path: impl AsRef<std::path::Path>) -> std::io::Result<Self> {
let (libc_addr, addr_len) = unix_path_to_libc(path.as_ref())?;
let fd = unsafe {
libc::socket(
libc::AF_UNIX,
libc::SOCK_STREAM | libc::SOCK_CLOEXEC | libc::SOCK_NONBLOCK,
0,
)
};
if fd < 0 {
return Err(std::io::Error::last_os_error());
}
let stream = unsafe { std::os::unix::net::UnixStream::from_raw_fd(fd) };
let num_workers = GLOBAL_CONFIG.get().map_or(1, |c| c.workers);
let worker_idx = fd as usize % num_workers;
let state = &WORKERS.get().unwrap()[worker_idx];
let direct_fd_idx = register_fd_sync(state, fd);
let connect_res = unsafe {
libc::connect(
fd,
(&raw const libc_addr).cast::<libc::sockaddr>(),
addr_len,
)
};
if connect_res == 0 {
return Ok(Self {
inner: stream,
direct_fd_idx,
worker_idx,
read_backpressured: std::sync::atomic::AtomicBool::new(false),
write_backpressured: std::sync::atomic::AtomicBool::new(false),
});
}
let err = std::io::Error::last_os_error();
if err.raw_os_error() != Some(libc::EINPROGRESS) {
return Err(err);
}
let mut pollfd = libc::pollfd {
fd,
events: libc::POLLOUT,
revents: 0,
};
let poll_res = unsafe { libc::poll(&raw mut pollfd, 1, 0) };
if poll_res > 0 {
if (pollfd.revents & libc::POLLOUT) != 0 {
let mut err_code: libc::c_int = 0;
let mut err_len = std::mem::size_of::<libc::c_int>() as libc::socklen_t;
let sockopt_res = unsafe {
libc::getsockopt(
fd,
libc::SOL_SOCKET,
libc::SO_ERROR,
(&raw mut err_code).cast::<libc::c_void>(),
&raw mut err_len,
)
};
if sockopt_res == 0 && err_code == 0 {
return Ok(Self {
inner: stream,
direct_fd_idx,
worker_idx,
read_backpressured: std::sync::atomic::AtomicBool::new(false),
write_backpressured: std::sync::atomic::AtomicBool::new(false),
});
}
let os_err = if err_code != 0 {
err_code
} else {
libc::ECONNREFUSED
};
return Err(std::io::Error::from_raw_os_error(os_err));
} else if (pollfd.revents & (libc::POLLERR | libc::POLLHUP)) != 0 {
return Err(std::io::Error::new(
std::io::ErrorKind::ConnectionRefused,
"connect failed",
));
}
}
let connect_res = DtactIoFuture {
worker_idx,
fd: fd as u32,
direct_fd_idx,
op: OpCode::Connect,
buf_ptr: std::ptr::null_mut(),
len: 0,
offset: 0,
addr: Some(libc_addr),
addr_len,
slot_idx: None,
msg_ptr: std::ptr::null_mut(),
}
.await;
match connect_res {
Ok(_) => Ok(Self {
inner: stream,
direct_fd_idx,
worker_idx,
read_backpressured: std::sync::atomic::AtomicBool::new(false),
write_backpressured: std::sync::atomic::AtomicBool::new(false),
}),
Err(e) => Err(e),
}
}
pub fn peer_cred(&self) -> std::io::Result<DtactUCred> {
peer_cred_impl(self.inner.as_raw_fd())
}
}
impl Drop for DtactUnixStream {
fn drop(&mut self) {
if let Some(workers) = WORKERS.get()
&& let Some(state) = workers.get(self.worker_idx)
{
unregister_fd_sync(state, self.direct_fd_idx);
}
}
}
impl crate::io::AsyncRead for DtactUnixStream {
async fn read(&self, buf: &mut [u8]) -> std::io::Result<usize> {
self.read(buf).await
}
}
impl crate::io::AsyncWrite for DtactUnixStream {
async fn write(&self, buf: &[u8]) -> std::io::Result<usize> {
self.write(buf).await
}
}
pub struct DtactUnixListener {
inner: std::os::unix::net::UnixListener,
direct_fd_idx: u32,
worker_idx: usize,
}
impl DtactUnixListener {
pub fn bind(path: impl AsRef<std::path::Path>) -> std::io::Result<Self> {
let listener = std::os::unix::net::UnixListener::bind(path)?;
Self::from_std(listener)
}
pub fn from_std(listener: std::os::unix::net::UnixListener) -> std::io::Result<Self> {
let fd = listener.as_raw_fd();
listener.set_nonblocking(true)?;
let num_workers = GLOBAL_CONFIG.get().map_or(1, |c| c.workers);
let worker_idx = WORKER_ROUND_ROBIN.fetch_add(1, Ordering::Relaxed) % num_workers;
let state = &WORKERS.get().unwrap()[worker_idx];
let direct_fd_idx = register_fd_sync(state, fd);
Ok(Self {
inner: listener,
direct_fd_idx,
worker_idx,
})
}
pub async fn accept(
&self,
) -> std::io::Result<(DtactUnixStream, std::os::unix::net::SocketAddr)> {
let res = unsafe {
libc::accept4(
self.inner.as_raw_fd(),
std::ptr::null_mut(),
std::ptr::null_mut(),
libc::SOCK_NONBLOCK | libc::SOCK_CLOEXEC,
)
};
if res >= 0 {
let stream = unsafe { std::os::unix::net::UnixStream::from_raw_fd(res) };
let peer_addr = stream.peer_addr()?;
let client_stream = DtactUnixStream::from_std(stream)?;
return Ok((client_stream, peer_addr));
}
let err = std::io::Error::last_os_error();
if err.kind() != std::io::ErrorKind::WouldBlock {
return Err(err);
}
let res = DtactIoFuture {
worker_idx: self.worker_idx,
fd: self.inner.as_raw_fd() as u32,
direct_fd_idx: self.direct_fd_idx,
op: OpCode::Accept,
buf_ptr: std::ptr::null_mut(),
len: 0,
offset: 0,
addr: None,
addr_len: 0,
slot_idx: None,
msg_ptr: std::ptr::null_mut(),
}
.await?;
let client_fd = res as RawFd;
let stream = unsafe { std::os::unix::net::UnixStream::from_raw_fd(client_fd) };
let peer_addr = stream.peer_addr()?;
let client_stream = DtactUnixStream::from_std(stream)?;
Ok((client_stream, peer_addr))
}
}
impl Drop for DtactUnixListener {
fn drop(&mut self) {
if let Some(workers) = WORKERS.get()
&& let Some(state) = workers.get(self.worker_idx)
{
unregister_fd_sync(state, self.direct_fd_idx);
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct DtactUCred {
uid: u32,
gid: u32,
pid: Option<i32>,
}
impl DtactUCred {
#[must_use]
pub const fn uid(&self) -> u32 {
self.uid
}
#[must_use]
pub const fn gid(&self) -> u32 {
self.gid
}
#[must_use]
pub const fn pid(&self) -> Option<i32> {
self.pid
}
}
#[cfg(target_os = "linux")]
fn peer_cred_impl(fd: RawFd) -> std::io::Result<DtactUCred> {
let mut cred: libc::ucred = unsafe { std::mem::zeroed() };
let mut len = std::mem::size_of::<libc::ucred>() as libc::socklen_t;
let r = unsafe {
libc::getsockopt(
fd,
libc::SOL_SOCKET,
libc::SO_PEERCRED,
(&raw mut cred).cast::<libc::c_void>(),
&raw mut len,
)
};
if r != 0 {
return Err(std::io::Error::last_os_error());
}
Ok(DtactUCred {
uid: cred.uid,
gid: cred.gid,
pid: Some(cred.pid),
})
}
#[cfg(not(target_os = "linux"))]
fn peer_cred_impl(fd: RawFd) -> std::io::Result<DtactUCred> {
let mut uid: libc::uid_t = 0;
let mut gid: libc::gid_t = 0;
let r = unsafe { libc::getpeereid(fd, &raw mut uid, &raw mut gid) };
if r != 0 {
return Err(std::io::Error::last_os_error());
}
Ok(DtactUCred {
uid,
gid,
pid: None,
})
}
fn unix_path_to_libc(
path: &std::path::Path,
) -> std::io::Result<(libc::sockaddr_storage, libc::socklen_t)> {
use std::os::unix::ffi::OsStrExt;
let bytes = path.as_os_str().as_bytes();
let mut storage: libc::sockaddr_storage = unsafe { std::mem::zeroed() };
let sun_ptr = (&raw mut storage).cast::<libc::sockaddr_un>();
let sun_path_cap = unsafe { (*sun_ptr).sun_path.len() };
if bytes.len() >= sun_path_cap {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"unix socket path too long for sockaddr_un::sun_path",
));
}
unsafe {
(*sun_ptr).sun_family = libc::AF_UNIX as libc::sa_family_t;
std::ptr::copy_nonoverlapping(
bytes.as_ptr(),
(*sun_ptr).sun_path.as_mut_ptr().cast::<u8>(),
bytes.len(),
);
}
let len = (std::mem::size_of::<libc::sa_family_t>() + bytes.len() + 1) as libc::socklen_t;
Ok((storage, len))
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct DtactUnixSocketAddr(Option<std::path::PathBuf>);
impl DtactUnixSocketAddr {
#[must_use]
pub fn as_pathname(&self) -> Option<&std::path::Path> {
self.0.as_deref()
}
#[must_use]
pub const fn is_unnamed(&self) -> bool {
self.0.is_none()
}
}
fn sockaddr_un_to_addr(
storage: &libc::sockaddr_storage,
len: libc::socklen_t,
) -> DtactUnixSocketAddr {
let family_len = std::mem::size_of::<libc::sa_family_t>();
if (len as usize) <= family_len {
return DtactUnixSocketAddr(None); }
let sun = unsafe { &*std::ptr::from_ref(storage).cast::<libc::sockaddr_un>() };
let path_len = (len as usize) - family_len;
let path_len = path_len.min(sun.sun_path.len());
let bytes = unsafe { std::slice::from_raw_parts(sun.sun_path.as_ptr().cast::<u8>(), path_len) };
let bytes = bytes.split(|&b| b == 0).next().unwrap_or(&[]);
if bytes.is_empty() {
DtactUnixSocketAddr(None)
} else {
use std::os::unix::ffi::OsStrExt;
DtactUnixSocketAddr(Some(std::path::PathBuf::from(std::ffi::OsStr::from_bytes(
bytes,
))))
}
}
pub struct DtactUnixDatagram {
inner: std::os::unix::net::UnixDatagram,
direct_fd_idx: u32,
worker_idx: usize,
read_backpressured: std::sync::atomic::AtomicBool,
write_backpressured: std::sync::atomic::AtomicBool,
}
impl DtactUnixDatagram {
pub fn bind(path: impl AsRef<std::path::Path>) -> std::io::Result<Self> {
let sock = std::os::unix::net::UnixDatagram::bind(path)?;
Self::from_std(sock)
}
pub fn unbound() -> std::io::Result<Self> {
let sock = std::os::unix::net::UnixDatagram::unbound()?;
Self::from_std(sock)
}
pub fn from_std(socket: std::os::unix::net::UnixDatagram) -> std::io::Result<Self> {
let fd = socket.as_raw_fd();
socket.set_nonblocking(true)?;
let num_workers = GLOBAL_CONFIG.get().map_or(1, |c| c.workers);
let worker_idx = fd as usize % num_workers;
let state = &WORKERS.get().unwrap()[worker_idx];
let direct_fd_idx = register_fd_sync(state, fd);
Ok(Self {
inner: socket,
direct_fd_idx,
worker_idx,
read_backpressured: std::sync::atomic::AtomicBool::new(false),
write_backpressured: std::sync::atomic::AtomicBool::new(false),
})
}
pub async fn send_to(
&self,
buf: &[u8],
target: impl AsRef<std::path::Path>,
) -> std::io::Result<usize> {
struct SendToState {
storage: libc::sockaddr_storage,
iov: libc::iovec,
msg: libc::msghdr,
}
unsafe impl Send for SendToState {}
if !self.write_backpressured.load(Ordering::Relaxed) {
let (storage, addr_len) = unix_path_to_libc(target.as_ref())?;
let mut state = SendToState {
storage,
iov: libc::iovec {
iov_base: buf.as_ptr().cast_mut().cast::<libc::c_void>(),
iov_len: buf.len(),
},
msg: unsafe { std::mem::zeroed() },
};
state.msg.msg_name = std::ptr::addr_of_mut!(state.storage).cast::<libc::c_void>();
state.msg.msg_namelen = addr_len;
state.msg.msg_iov = &raw mut state.iov;
state.msg.msg_iovlen = 1;
let r = unsafe { libc::sendmsg(self.inner.as_raw_fd(), &raw const state.msg, 0) };
if r >= 0 {
return Ok(r as usize);
}
let e = std::io::Error::last_os_error();
if e.kind() != std::io::ErrorKind::WouldBlock {
return Err(e);
}
self.write_backpressured.store(true, Ordering::Relaxed);
}
let (storage, addr_len) = unix_path_to_libc(target.as_ref())?;
let mut state = SendToState {
storage,
iov: libc::iovec {
iov_base: buf.as_ptr().cast_mut().cast::<libc::c_void>(),
iov_len: buf.len(),
},
msg: unsafe { std::mem::zeroed() },
};
state.msg.msg_name = std::ptr::addr_of_mut!(state.storage).cast::<libc::c_void>();
state.msg.msg_namelen = addr_len;
state.msg.msg_iov = &raw mut state.iov;
state.msg.msg_iovlen = 1;
let mut fut = DtactIoFuture::new(
self.worker_idx,
self.inner.as_raw_fd() as u32,
self.direct_fd_idx,
OpCode::SendTo,
std::ptr::null_mut(),
0,
0,
None,
0,
None,
);
fut.msg_ptr = &raw mut state.msg;
let res = fut.await;
self.write_backpressured.store(false, Ordering::Relaxed);
res
}
pub async fn recv_from(&self, buf: &mut [u8]) -> std::io::Result<(usize, DtactUnixSocketAddr)> {
struct RecvFromState {
storage: libc::sockaddr_storage,
iov: libc::iovec,
msg: libc::msghdr,
}
unsafe impl Send for RecvFromState {}
if !self.read_backpressured.load(Ordering::Relaxed) {
let mut state = RecvFromState {
storage: unsafe { std::mem::zeroed() },
iov: libc::iovec {
iov_base: buf.as_mut_ptr().cast::<libc::c_void>(),
iov_len: buf.len(),
},
msg: unsafe { std::mem::zeroed() },
};
state.msg.msg_name = std::ptr::addr_of_mut!(state.storage).cast::<libc::c_void>();
state.msg.msg_namelen =
std::mem::size_of::<libc::sockaddr_storage>() as libc::socklen_t;
state.msg.msg_iov = &raw mut state.iov;
state.msg.msg_iovlen = 1;
let r = unsafe { libc::recvmsg(self.inner.as_raw_fd(), &raw mut state.msg, 0) };
if r >= 0 {
let from = sockaddr_un_to_addr(&state.storage, state.msg.msg_namelen);
return Ok((r as usize, from));
}
let e = std::io::Error::last_os_error();
if e.kind() != std::io::ErrorKind::WouldBlock {
return Err(e);
}
self.read_backpressured.store(true, Ordering::Relaxed);
}
let mut state = RecvFromState {
storage: unsafe { std::mem::zeroed() },
iov: libc::iovec {
iov_base: buf.as_mut_ptr().cast::<libc::c_void>(),
iov_len: buf.len(),
},
msg: unsafe { std::mem::zeroed() },
};
state.msg.msg_name = std::ptr::addr_of_mut!(state.storage).cast::<libc::c_void>();
state.msg.msg_namelen = std::mem::size_of::<libc::sockaddr_storage>() as libc::socklen_t;
state.msg.msg_iov = &raw mut state.iov;
state.msg.msg_iovlen = 1;
let mut fut = DtactIoFuture::new(
self.worker_idx,
self.inner.as_raw_fd() as u32,
self.direct_fd_idx,
OpCode::RecvFrom,
std::ptr::null_mut(),
0,
0,
None,
0,
None,
);
fut.msg_ptr = &raw mut state.msg;
let n = fut.await?;
self.read_backpressured.store(false, Ordering::Relaxed);
let from = sockaddr_un_to_addr(&state.storage, state.msg.msg_namelen);
Ok((n, from))
}
pub async fn connect(&self, target: impl AsRef<std::path::Path>) -> std::io::Result<()> {
self.inner.connect(target)
}
pub async fn send(&self, buf: &[u8]) -> std::io::Result<usize> {
if buf.is_empty() {
return Ok(0);
}
let r = unsafe {
libc::send(
self.inner.as_raw_fd(),
buf.as_ptr().cast::<libc::c_void>(),
buf.len(),
0,
)
};
if r >= 0 {
return Ok(r as usize);
}
let e = std::io::Error::last_os_error();
if e.kind() != std::io::ErrorKind::WouldBlock {
return Err(e);
}
DtactIoFuture {
worker_idx: self.worker_idx,
fd: self.inner.as_raw_fd() as u32,
direct_fd_idx: self.direct_fd_idx,
op: OpCode::Write,
buf_ptr: buf.as_ptr().cast_mut(),
len: buf.len(),
offset: 0,
addr: None,
addr_len: 0,
slot_idx: None,
msg_ptr: std::ptr::null_mut(),
}
.await
}
pub async fn recv(&self, buf: &mut [u8]) -> std::io::Result<usize> {
if buf.is_empty() {
return Ok(0);
}
let r = unsafe {
libc::recv(
self.inner.as_raw_fd(),
buf.as_mut_ptr().cast::<libc::c_void>(),
buf.len(),
0,
)
};
if r >= 0 {
return Ok(r as usize);
}
let e = std::io::Error::last_os_error();
if e.kind() != std::io::ErrorKind::WouldBlock {
return Err(e);
}
DtactIoFuture {
worker_idx: self.worker_idx,
fd: self.inner.as_raw_fd() as u32,
direct_fd_idx: self.direct_fd_idx,
op: OpCode::Read,
buf_ptr: buf.as_mut_ptr(),
len: buf.len(),
offset: 0,
addr: None,
addr_len: 0,
slot_idx: None,
msg_ptr: std::ptr::null_mut(),
}
.await
}
}
impl Drop for DtactUnixDatagram {
fn drop(&mut self) {
if let Some(workers) = WORKERS.get()
&& let Some(state) = workers.get(self.worker_idx)
{
unregister_fd_sync(state, self.direct_fd_idx);
}
}
}
pub struct DtactFifoReader {
inner: std::fs::File,
direct_fd_idx: u32,
worker_idx: usize,
backpressured: std::sync::atomic::AtomicBool,
}
unsafe impl Send for DtactFifoReader {}
unsafe impl Sync for DtactFifoReader {}
impl DtactFifoReader {
pub async fn read(&self, buf: &mut [u8]) -> std::io::Result<usize> {
if buf.is_empty() {
return Ok(0);
}
if !self.backpressured.load(Ordering::Relaxed) {
let res = unsafe {
let r = libc::read(
self.inner.as_raw_fd(),
buf.as_mut_ptr().cast::<libc::c_void>(),
buf.len(),
);
match r.cmp(&0) {
std::cmp::Ordering::Greater => Ok(r as usize),
std::cmp::Ordering::Equal => Ok(0),
std::cmp::Ordering::Less => Err(std::io::Error::last_os_error()),
}
};
match res {
Ok(n) => return Ok(n),
Err(e) => {
if e.kind() != std::io::ErrorKind::WouldBlock {
return Err(e);
}
}
}
self.backpressured.store(true, Ordering::Relaxed);
}
let future = DtactIoFuture {
worker_idx: self.worker_idx,
fd: self.inner.as_raw_fd() as u32,
direct_fd_idx: self.direct_fd_idx,
op: OpCode::Read,
buf_ptr: buf.as_mut_ptr(),
len: buf.len(),
offset: 0,
addr: None,
addr_len: 0,
slot_idx: None,
msg_ptr: std::ptr::null_mut(),
}
.await;
self.backpressured.store(false, Ordering::Relaxed);
future.map(|n| n.min(buf.len()))
}
}
impl Drop for DtactFifoReader {
fn drop(&mut self) {
if let Some(workers) = WORKERS.get()
&& let Some(state) = workers.get(self.worker_idx)
{
unregister_fd_sync(state, self.direct_fd_idx);
}
}
}
pub struct DtactFifoWriter {
inner: std::fs::File,
direct_fd_idx: u32,
worker_idx: usize,
backpressured: std::sync::atomic::AtomicBool,
}
unsafe impl Send for DtactFifoWriter {}
unsafe impl Sync for DtactFifoWriter {}
impl DtactFifoWriter {
pub async fn write(&self, buf: &[u8]) -> std::io::Result<usize> {
if buf.is_empty() {
return Ok(0);
}
if !self.backpressured.load(Ordering::Relaxed) {
let res = unsafe {
let r = libc::write(
self.inner.as_raw_fd(),
buf.as_ptr().cast::<libc::c_void>(),
buf.len(),
);
if r >= 0 {
Ok(r as usize)
} else {
Err(std::io::Error::last_os_error())
}
};
match res {
Ok(n) => return Ok(n),
Err(e) => {
if e.kind() != std::io::ErrorKind::WouldBlock {
return Err(e);
}
}
}
self.backpressured.store(true, Ordering::Relaxed);
}
let future = DtactIoFuture {
worker_idx: self.worker_idx,
fd: self.inner.as_raw_fd() as u32,
direct_fd_idx: self.direct_fd_idx,
op: OpCode::Write,
buf_ptr: buf.as_ptr().cast_mut(),
len: buf.len(),
offset: 0,
addr: None,
addr_len: 0,
slot_idx: None,
msg_ptr: std::ptr::null_mut(),
}
.await;
self.backpressured.store(false, Ordering::Relaxed);
future.map(|n| n.min(buf.len()))
}
}
impl Drop for DtactFifoWriter {
fn drop(&mut self) {
if let Some(workers) = WORKERS.get()
&& let Some(state) = workers.get(self.worker_idx)
{
unregister_fd_sync(state, self.direct_fd_idx);
}
}
}
fn register_fifo_fd(fd: RawFd) -> (u32, usize) {
let num_workers = GLOBAL_CONFIG.get().map_or(1, |c| c.workers);
let worker_idx = fd as usize % num_workers;
let state = &WORKERS.get().unwrap()[worker_idx];
(register_fd_sync(state, fd), worker_idx)
}
pub async fn open_fifo_read(
path: impl Into<std::path::PathBuf>,
) -> std::io::Result<DtactFifoReader> {
use std::os::unix::fs::OpenOptionsExt;
let path = path.into();
let file = std::fs::OpenOptions::new()
.read(true)
.custom_flags(libc::O_NONBLOCK)
.open(&path)?;
let (direct_fd_idx, worker_idx) = register_fifo_fd(file.as_raw_fd());
Ok(DtactFifoReader {
inner: file,
direct_fd_idx,
worker_idx,
backpressured: AtomicBool::new(false),
})
}
pub async fn open_fifo_write(
path: impl Into<std::path::PathBuf>,
) -> std::io::Result<DtactFifoWriter> {
use std::os::unix::fs::OpenOptionsExt;
let path = path.into();
let file = std::fs::OpenOptions::new()
.write(true)
.custom_flags(libc::O_NONBLOCK)
.open(&path)?;
let (direct_fd_idx, worker_idx) = register_fifo_fd(file.as_raw_fd());
Ok(DtactFifoWriter {
inner: file,
direct_fd_idx,
worker_idx,
backpressured: AtomicBool::new(false),
})
}
const fn register_fd_sync(_state: &WorkerState, _fd: RawFd) -> u32 {
u32::MAX
}
const fn unregister_fd_sync(_state: &WorkerState, _direct_fd_idx: u32) {}
const fn socket_addr_to_libc(
addr: std::net::SocketAddr,
) -> (libc::sockaddr_storage, libc::socklen_t) {
let mut storage: libc::sockaddr_storage = unsafe { std::mem::zeroed() };
let len = match addr {
std::net::SocketAddr::V4(a) => {
let sin = libc::sockaddr_in {
sin_family: libc::AF_INET as libc::sa_family_t,
sin_port: a.port().to_be(),
sin_addr: libc::in_addr {
s_addr: u32::from_ne_bytes(a.ip().octets()),
},
sin_zero: [0; 8],
};
unsafe {
std::ptr::copy_nonoverlapping(
(&raw const sin).cast::<u8>(),
(&raw mut storage).cast::<u8>(),
std::mem::size_of::<libc::sockaddr_in>(),
);
}
std::mem::size_of::<libc::sockaddr_in>() as libc::socklen_t
}
std::net::SocketAddr::V6(a) => {
let sin6 = libc::sockaddr_in6 {
sin6_family: libc::AF_INET6 as libc::sa_family_t,
sin6_port: a.port().to_be(),
sin6_flowinfo: a.flowinfo(),
sin6_addr: libc::in6_addr {
s6_addr: a.ip().octets(),
},
sin6_scope_id: a.scope_id(),
};
unsafe {
std::ptr::copy_nonoverlapping(
(&raw const sin6).cast::<u8>(),
(&raw mut storage).cast::<u8>(),
std::mem::size_of::<libc::sockaddr_in6>(),
);
}
std::mem::size_of::<libc::sockaddr_in6>() as libc::socklen_t
}
};
(storage, len)
}
fn sockaddr_storage_to_socketaddr(
storage: &libc::sockaddr_storage,
_len: libc::socklen_t,
) -> std::net::SocketAddr {
match libc::c_int::from(storage.ss_family) {
libc::AF_INET => {
let sin = unsafe { &*std::ptr::from_ref(storage).cast::<libc::sockaddr_in>() };
let ip = std::net::Ipv4Addr::from(u32::from_be(sin.sin_addr.s_addr));
let port = u16::from_be(sin.sin_port);
std::net::SocketAddr::V4(std::net::SocketAddrV4::new(ip, port))
}
libc::AF_INET6 => {
let sin6 = unsafe { &*std::ptr::from_ref(storage).cast::<libc::sockaddr_in6>() };
let ip = std::net::Ipv6Addr::from(sin6.sin6_addr.s6_addr);
let port = u16::from_be(sin6.sin6_port);
std::net::SocketAddr::V6(std::net::SocketAddrV6::new(
ip,
port,
sin6.sin6_flowinfo,
sin6.sin6_scope_id,
))
}
_ => {
panic!("Unsupported address family: {}", storage.ss_family);
}
}
}