use crate::iou::sqe::SockAddrStorage;
use ahash::AHashMap;
use log::error;
use nix::sys::socket::{MsgFlags, SockAddr};
use std::cell::{Cell, RefCell};
use std::collections::{BTreeMap, VecDeque};
use std::ffi::CString;
use std::fmt;
use std::io;
use std::mem;
use std::os::unix::ffi::OsStrExt;
use std::os::unix::io::RawFd;
use std::panic::{self, RefUnwindSafe, UnwindSafe};
use std::path::Path;
use std::rc::Rc;
use std::sync::atomic::{AtomicU32, Ordering};
use std::sync::Arc;
use std::task::{Poll, Waker};
use std::time::{Duration, Instant};
use futures_lite::*;
use crate::sys;
use crate::sys::{
DirectIO, DmaBuffer, IOBuffer, PollableStatus, SleepNotifier, Source, SourceType,
};
use crate::Local;
use crate::{IoRequirements, Latency};
pub(crate) struct Parker {
inner: Rc<Inner>,
}
impl UnwindSafe for Parker {}
impl RefUnwindSafe for Parker {}
impl Parker {
pub(crate) fn new() -> Parker {
Parker {
inner: Rc::new(Inner {}),
}
}
pub(crate) fn park(&self) {
self.inner.park(None);
}
pub(crate) fn poll_io(&self, timeout: Duration) {
self.inner.park(Some(timeout));
}
}
impl Drop for Parker {
fn drop(&mut self) {}
}
impl Default for Parker {
fn default() -> Parker {
Parker::new()
}
}
impl fmt::Debug for Parker {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.pad("Parker { .. }")
}
}
impl fmt::Debug for Reactor {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.pad("Reactor { .. }")
}
}
struct Inner {}
impl Inner {
fn park(&self, timeout: Option<Duration>) -> bool {
let _ = Local::get_reactor().react(timeout);
false
}
}
struct Timers {
timer_id: u64,
timers_by_id: AHashMap<u64, Instant>,
timers: BTreeMap<(Instant, u64), Waker>,
}
impl Timers {
fn new() -> Timers {
Timers {
timer_id: 0,
timers_by_id: AHashMap::new(),
timers: BTreeMap::new(),
}
}
fn new_id(&mut self) -> u64 {
self.timer_id += 1;
self.timer_id
}
fn remove(&mut self, id: u64) -> Option<Waker> {
if let Some(when) = self.timers_by_id.remove(&id) {
return self.timers.remove(&(when, id));
}
None
}
fn insert(&mut self, id: u64, when: Instant, waker: Waker) {
if let Some(when) = self.timers_by_id.get_mut(&id) {
self.timers.remove(&(*when, id));
}
self.timers_by_id.insert(id, when);
self.timers.insert((when, id), waker);
}
fn process_timers(&mut self, wakers: &mut Vec<Waker>) -> Option<Duration> {
let now = Instant::now();
let pending = self.timers.split_off(&(now, 0));
let ready = mem::replace(&mut self.timers, pending);
let dur = if ready.is_empty() {
self.timers
.keys()
.next()
.map(|(when, _)| when.saturating_duration_since(now))
} else {
Some(Duration::from_secs(0))
};
for (_, waker) in ready {
wakers.push(waker);
}
dur
}
}
struct SharedChannels {
id: u64,
check_map: BTreeMap<u64, Box<dyn Fn() -> usize>>,
wakers_map: BTreeMap<u64, VecDeque<Waker>>,
connection_wakers: Vec<Waker>,
}
impl SharedChannels {
fn new() -> SharedChannels {
SharedChannels {
id: 0,
connection_wakers: Vec::new(),
check_map: BTreeMap::new(),
wakers_map: BTreeMap::new(),
}
}
fn process_shared_channels(&mut self, wakers: &mut Vec<Waker>) -> usize {
let mut added = self.connection_wakers.len();
wakers.append(&mut self.connection_wakers);
let current_wakers = mem::take(&mut self.wakers_map);
for (id, mut pending) in current_wakers.into_iter() {
let room = self.check_map.get(&id).unwrap()();
let room = std::cmp::min(room, pending.len());
for w in pending.drain(0..room) {
added += 1;
wakers.push(w);
}
if !pending.is_empty() {
self.wakers_map.insert(id, pending);
}
}
added
}
}
pub(crate) struct Reactor {
pub(crate) sys: sys::Reactor,
timers: RefCell<Timers>,
shared_channels: RefCell<SharedChannels>,
current_io_requirements: Cell<IoRequirements>,
wakers: RefCell<Vec<Waker>>,
preempt_ptr_head: *const u32,
preempt_ptr_tail: *const AtomicU32,
}
impl Reactor {
pub(crate) fn new(notifier: Arc<SleepNotifier>, io_memory: usize) -> Reactor {
let sys = sys::Reactor::new(notifier, io_memory)
.expect("cannot initialize I/O event notification");
let (preempt_ptr_head, preempt_ptr_tail) = sys.preempt_pointers();
Reactor {
sys,
timers: RefCell::new(Timers::new()),
shared_channels: RefCell::new(SharedChannels::new()),
current_io_requirements: Cell::new(IoRequirements::default()),
wakers: RefCell::new(Vec::with_capacity(256)),
preempt_ptr_head,
preempt_ptr_tail: preempt_ptr_tail as _,
}
}
#[inline(always)]
pub(crate) fn need_preempt(&self) -> bool {
unsafe { *self.preempt_ptr_head != (*self.preempt_ptr_tail).load(Ordering::Acquire) }
}
pub(crate) fn id(&self) -> usize {
self.sys.id()
}
pub(crate) fn notify(&self, remote: RawFd) {
sys::write_eventfd(remote);
}
fn new_source(&self, raw: RawFd, stype: SourceType) -> Source {
let ioreq = self.current_io_requirements.get();
sys::Source::new(ioreq, raw, stype)
}
pub(crate) fn inform_io_requirements(&self, req: IoRequirements) {
self.current_io_requirements.set(req);
}
pub(crate) fn register_shared_channel<F>(&self, test_function: Box<F>) -> u64
where
F: Fn() -> usize + 'static,
{
let mut channels = self.shared_channels.borrow_mut();
let id = channels.id;
channels.id += 1;
let ret = channels.check_map.insert(id, test_function);
assert_eq!(ret.is_none(), true);
id
}
pub(crate) fn unregister_shared_channel(&self, id: u64) {
let mut channels = self.shared_channels.borrow_mut();
channels.wakers_map.remove(&id);
channels.check_map.remove(&id);
}
pub(crate) fn add_shared_channel_connection_waker(&self, waker: Waker) {
let mut channels = self.shared_channels.borrow_mut();
channels.connection_wakers.push(waker);
}
pub(crate) fn add_shared_channel_waker(&self, id: u64, waker: Waker) {
let mut channels = self.shared_channels.borrow_mut();
let map = channels.wakers_map.entry(id).or_insert_with(VecDeque::new);
map.push_back(waker);
}
pub(crate) fn alloc_dma_buffer(&self, size: usize) -> DmaBuffer {
self.sys.alloc_dma_buffer(size)
}
pub(crate) fn write_dma(
&self,
raw: RawFd,
buf: DmaBuffer,
pos: u64,
pollable: PollableStatus,
) -> Source {
let source = self.new_source(raw, SourceType::Write(pollable, IOBuffer::Dma(buf)));
self.sys.write_dma(&source, pos);
source
}
pub(crate) fn write_buffered(&self, raw: RawFd, buf: Vec<u8>, pos: u64) -> Source {
let source = self.new_source(
raw,
SourceType::Write(
PollableStatus::NonPollable(DirectIO::Disabled),
IOBuffer::Buffered(buf),
),
);
self.sys.write_buffered(&source, pos);
source
}
pub(crate) fn connect(&self, raw: RawFd, addr: SockAddr) -> Source {
let source = self.new_source(raw, SourceType::Connect(addr));
self.sys.connect(&source);
source
}
pub(crate) fn accept(&self, raw: RawFd) -> Source {
let addr = SockAddrStorage::uninit();
let source = self.new_source(raw, SourceType::Accept(addr));
self.sys.accept(&source);
source
}
pub(crate) fn rushed_send(&self, fd: RawFd, buf: DmaBuffer) -> io::Result<Source> {
let source = self.new_source(fd, SourceType::SockSend(buf));
self.sys.send(&source, MsgFlags::empty());
self.rush_dispatch(&source)?;
Ok(source)
}
pub(crate) fn rushed_sendmsg(
&self,
fd: RawFd,
buf: DmaBuffer,
addr: nix::sys::socket::SockAddr,
) -> io::Result<Source> {
let iov = libc::iovec {
iov_base: buf.as_ptr() as *mut libc::c_void,
iov_len: 1,
};
let hdr = libc::msghdr {
msg_name: std::ptr::null_mut(),
msg_namelen: 0,
msg_iov: std::ptr::null_mut(),
msg_iovlen: 0,
msg_control: std::ptr::null_mut(),
msg_controllen: 0,
msg_flags: 0,
};
let source = self.new_source(fd, SourceType::SockSendMsg(buf, iov, hdr, addr));
self.sys.sendmsg(&source, MsgFlags::empty());
self.rush_dispatch(&source)?;
Ok(source)
}
pub(crate) fn rushed_recvmsg(
&self,
fd: RawFd,
size: usize,
flags: MsgFlags,
) -> io::Result<Source> {
let hdr = libc::msghdr {
msg_name: std::ptr::null_mut(),
msg_namelen: 0,
msg_iov: std::ptr::null_mut(),
msg_iovlen: 0,
msg_control: std::ptr::null_mut(),
msg_controllen: 0,
msg_flags: 0,
};
let iov = libc::iovec {
iov_base: std::ptr::null_mut(),
iov_len: 0,
};
let source = self.new_source(
fd,
SourceType::SockRecvMsg(
None,
iov,
hdr,
std::mem::MaybeUninit::<nix::sys::socket::sockaddr_storage>::uninit(),
),
);
self.sys.recvmsg(&source, size, flags);
self.rush_dispatch(&source)?;
Ok(source)
}
pub(crate) fn rushed_recv(&self, fd: RawFd, size: usize) -> io::Result<Source> {
let source = self.new_source(fd, SourceType::SockRecv(None));
self.sys.recv(&source, size, MsgFlags::empty());
self.rush_dispatch(&source)?;
Ok(source)
}
pub(crate) fn recv(&self, fd: RawFd, size: usize, flags: MsgFlags) -> Source {
let source = self.new_source(fd, SourceType::SockRecv(None));
self.sys.recv(&source, size, flags);
source
}
pub(crate) fn read_dma(
&self,
raw: RawFd,
pos: u64,
size: usize,
pollable: PollableStatus,
) -> Source {
let source = self.new_source(raw, SourceType::Read(pollable, None));
self.sys.read_dma(&source, pos, size);
source
}
pub(crate) fn read_buffered(&self, raw: RawFd, pos: u64, size: usize) -> Source {
let source = self.new_source(
raw,
SourceType::Read(PollableStatus::NonPollable(DirectIO::Disabled), None),
);
self.sys.read_buffered(&source, pos, size);
source
}
pub(crate) fn fdatasync(&self, raw: RawFd) -> Source {
let source = self.new_source(raw, SourceType::FdataSync);
self.sys.fdatasync(&source);
source
}
pub(crate) fn fallocate(
&self,
raw: RawFd,
position: u64,
size: u64,
flags: libc::c_int,
) -> Source {
let source = self.new_source(raw, SourceType::Fallocate);
self.sys.fallocate(&source, position, size, flags);
source
}
pub(crate) fn close(&self, raw: RawFd) -> Source {
let source = self.new_source(raw, SourceType::Close);
self.sys.close(&source);
source
}
pub(crate) fn statx(&self, raw: RawFd, path: &Path) -> Source {
let path = CString::new(path.as_os_str().as_bytes()).expect("path contained null!");
let statx_buf = unsafe {
let statx_buf = mem::MaybeUninit::<libc::statx>::zeroed();
statx_buf.assume_init()
};
let source = self.new_source(
raw,
SourceType::Statx(path, Box::new(RefCell::new(statx_buf))),
);
self.sys.statx(&source);
source
}
pub(crate) fn open_at(
&self,
dir: RawFd,
path: &Path,
flags: libc::c_int,
mode: libc::mode_t,
) -> Source {
let path = CString::new(path.as_os_str().as_bytes()).expect("path contained null!");
let source = self.new_source(dir, SourceType::Open(path));
self.sys.open_at(&source, flags, mode);
source
}
#[cfg(feature = "bench")]
pub(crate) fn nop(&self) -> Source {
let source = self.new_source(-1, SourceType::Noop);
self.sys.nop(&source);
source
}
pub(crate) fn register_timer(&self) -> u64 {
let mut timers = self.timers.borrow_mut();
timers.new_id()
}
pub(crate) fn insert_timer(&self, id: u64, when: Instant, waker: Waker) {
let mut timers = self.timers.borrow_mut();
timers.insert(id, when, waker);
}
pub(crate) fn remove_timer(&self, id: u64) -> Option<Waker> {
let mut timers = self.timers.borrow_mut();
timers.remove(id)
}
fn process_timers(&self, wakers: &mut Vec<Waker>) -> Option<Duration> {
let mut timers = self.timers.borrow_mut();
timers.process_timers(wakers)
}
fn process_shared_channels(&self, wakers: &mut Vec<Waker>) -> usize {
let mut channels = self.shared_channels.borrow_mut();
channels.process_shared_channels(wakers)
}
fn rush_dispatch(&self, source: &Source) -> io::Result<()> {
let mut wakers = self.wakers.borrow_mut();
self.sys
.rush_dispatch(Some(source.latency_req()), &mut wakers)?;
for waker in wakers.drain(..) {
let _ = panic::catch_unwind(|| waker.wake());
}
Ok(())
}
pub(crate) fn spin_poll_io(&self) -> io::Result<bool> {
let mut wakers = self.wakers.borrow_mut();
self.sys
.rush_dispatch(Some(Latency::Matters(Duration::from_secs(1))), &mut wakers)?;
self.sys
.rush_dispatch(Some(Latency::NotImportant), &mut wakers)?;
self.sys.rush_dispatch(None, &mut wakers)?;
self.process_timers(&mut wakers);
self.process_shared_channels(&mut wakers);
let woke = wakers.len();
for waker in wakers.drain(..) {
let _ = panic::catch_unwind(|| waker.wake());
}
Ok(woke > 0)
}
fn react(&self, timeout: Option<Duration>) -> io::Result<()> {
let mut wakers = self.wakers.borrow_mut();
let next_timer = self.process_timers(&mut wakers);
self.process_shared_channels(&mut wakers);
let res = match self.sys.wait(&mut wakers, timeout, next_timer, |wakers| {
self.process_shared_channels(wakers)
}) {
Ok(true) => {
self.process_timers(&mut wakers);
Ok(())
}
Ok(false) => Ok(()),
Err(err) if err.kind() == io::ErrorKind::Interrupted => Ok(()),
Err(err) => Err(err),
};
for waker in wakers.drain(..) {
if let Err(x) = panic::catch_unwind(|| waker.wake()) {
error!("Panic while calling waker! {:?}", x);
}
}
res
}
}
impl Source {
pub(crate) async fn collect_rw(&self) -> io::Result<usize> {
future::poll_fn(|cx| {
if let Some(result) = self.take_result() {
return Poll::Ready(result);
}
self.add_waiter(cx.waker().clone());
Poll::Pending
})
.await
}
}