use std::cell::RefCell;
use std::collections::BTreeMap;
use std::io;
use std::os::unix::io::RawFd;
use std::rc::Rc;
use io_uring::{opcode, types, IoUring};
use crate::scheduler::{Poll, Task};
enum IoOutcome {
Pending,
Done { res: i32 },
}
pub struct Reactor {
ring: IoUring,
next_op_id: u64,
pending: BTreeMap<u64, Rc<RefCell<IoOutcome>>>,
}
impl Reactor {
pub fn new(entries: u32) -> io::Result<Self> {
Ok(Self {
ring: IoUring::new(entries)?,
next_op_id: 0,
pending: BTreeMap::new(),
})
}
fn next_id(&mut self) -> u64 {
let id = self.next_op_id;
self.next_op_id += 1;
id
}
pub unsafe fn submit_read(&mut self, fd: RawFd, buf: &mut [u8], offset: u64) -> u64 {
let op_id = self.next_id();
let entry = opcode::Read::new(types::Fd(fd), buf.as_mut_ptr(), buf.len() as u32)
.offset(offset)
.build()
.user_data(op_id);
self.push(entry);
self.pending
.insert(op_id, Rc::new(RefCell::new(IoOutcome::Pending)));
op_id
}
pub unsafe fn submit_write(&mut self, fd: RawFd, buf: &[u8], offset: u64) -> u64 {
let op_id = self.next_id();
let entry = opcode::Write::new(types::Fd(fd), buf.as_ptr(), buf.len() as u32)
.offset(offset)
.build()
.user_data(op_id);
self.push(entry);
self.pending
.insert(op_id, Rc::new(RefCell::new(IoOutcome::Pending)));
op_id
}
fn push(&mut self, entry: io_uring::squeue::Entry) {
while unsafe { self.ring.submission().push(&entry) }.is_err() {
let _ = self.ring.submit();
}
}
pub fn poll_completions(&mut self) -> io::Result<usize> {
self.ring.submit()?;
let mut completed = 0;
let cqes: std::vec::Vec<(u64, i32)> = self
.ring
.completion()
.map(|cqe| (cqe.user_data(), cqe.result()))
.collect();
for (op_id, res) in cqes {
if let Some(outcome) = self.pending.remove(&op_id) {
*outcome.borrow_mut() = IoOutcome::Done { res };
}
completed += 1;
}
Ok(completed)
}
fn take_result(&mut self, op_id: u64) -> Option<i32> {
match self.pending.get(&op_id) {
Some(outcome) => match &*outcome.borrow() {
IoOutcome::Done { res } => Some(*res),
IoOutcome::Pending => None,
},
None => None,
}
}
}
enum Phase {
NotSubmitted,
Submitted(u64),
Done(i32),
}
pub struct IoReadTask {
reactor: Rc<RefCell<Reactor>>,
fd: RawFd,
offset: u64,
buf: std::vec::Vec<u8>,
phase: Phase,
}
impl IoReadTask {
pub fn new(reactor: Rc<RefCell<Reactor>>, fd: RawFd, offset: u64, len: usize) -> Self {
Self {
reactor,
fd,
offset,
buf: alloc::vec![0u8; len],
phase: Phase::NotSubmitted,
}
}
pub fn result(&self) -> i32 {
match self.phase {
Phase::Done(res) => res,
_ => panic!("IoReadTask::result called before completion"),
}
}
pub fn buffer(&self) -> &[u8] {
&self.buf
}
}
impl Task for IoReadTask {
fn poll(&mut self) -> Poll {
match self.phase {
Phase::NotSubmitted => {
let op_id = unsafe {
self.reactor
.borrow_mut()
.submit_read(self.fd, &mut self.buf, self.offset)
};
self.phase = Phase::Submitted(op_id);
let _ = self.reactor.borrow_mut().poll_completions();
Poll::Pending
}
Phase::Submitted(op_id) => {
let _ = self.reactor.borrow_mut().poll_completions();
match self.reactor.borrow_mut().take_result(op_id) {
Some(res) => {
self.phase = Phase::Done(res);
Poll::Ready
}
None => Poll::Pending,
}
}
Phase::Done(_) => Poll::Ready,
}
}
}
pub struct IoWriteTask {
reactor: Rc<RefCell<Reactor>>,
fd: RawFd,
offset: u64,
buf: std::vec::Vec<u8>,
phase: Phase,
}
impl IoWriteTask {
pub fn new(
reactor: Rc<RefCell<Reactor>>,
fd: RawFd,
offset: u64,
data: std::vec::Vec<u8>,
) -> Self {
Self {
reactor,
fd,
offset,
buf: data,
phase: Phase::NotSubmitted,
}
}
pub fn result(&self) -> i32 {
match self.phase {
Phase::Done(res) => res,
_ => panic!("IoWriteTask::result called before completion"),
}
}
}
impl Task for IoWriteTask {
fn poll(&mut self) -> Poll {
match self.phase {
Phase::NotSubmitted => {
let op_id = unsafe {
self.reactor
.borrow_mut()
.submit_write(self.fd, &self.buf, self.offset)
};
self.phase = Phase::Submitted(op_id);
let _ = self.reactor.borrow_mut().poll_completions();
Poll::Pending
}
Phase::Submitted(op_id) => {
let _ = self.reactor.borrow_mut().poll_completions();
match self.reactor.borrow_mut().take_result(op_id) {
Some(res) => {
self.phase = Phase::Done(res);
Poll::Ready
}
None => Poll::Pending,
}
}
Phase::Done(_) => Poll::Ready,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::scheduler::Scheduler;
use std::io::{Seek, SeekFrom, Write};
use std::os::unix::io::AsRawFd;
fn temp_file_with(contents: &[u8]) -> std::fs::File {
let path = std::env::temp_dir().join(format!(
"tpt-archon-kernel-io-uring-test-{}-{:?}",
std::process::id(),
std::thread::current().id()
));
let mut f = std::fs::OpenOptions::new()
.read(true)
.write(true)
.create(true)
.truncate(true)
.open(&path)
.unwrap();
std::fs::remove_file(&path).unwrap();
f.write_all(contents).unwrap();
f.seek(SeekFrom::Start(0)).unwrap();
f
}
#[test]
fn read_task_yields_correct_bytes() {
let f = temp_file_with(b"hello io_uring");
let reactor = Rc::new(RefCell::new(
Reactor::new(8).expect("io_uring not available"),
));
let mut task = IoReadTask::new(reactor, f.as_raw_fd(), 0, b"hello io_uring".len());
let mut ticks = 0;
while task.poll() == Poll::Pending {
ticks += 1;
assert!(ticks < 100_000, "read never completed");
}
assert_eq!(task.result(), b"hello io_uring".len() as i32);
assert_eq!(task.buffer(), b"hello io_uring");
drop(f);
}
#[test]
fn read_task_completes_through_scheduler() {
let f = temp_file_with(b"round trip through the scheduler");
let reactor = Rc::new(RefCell::new(
Reactor::new(8).expect("io_uring not available"),
));
let mut scheduler = Scheduler::new();
let fd = f.as_raw_fd();
let task = IoReadTask::new(reactor, fd, 0, b"round trip through the scheduler".len());
let id = scheduler.spawn(alloc::boxed::Box::new(task));
let mut ticks = 0;
loop {
match scheduler.tick() {
Some((got_id, Poll::Ready)) => {
assert_eq!(got_id, id);
break;
}
Some((_, Poll::Pending)) => {
ticks += 1;
assert!(ticks < 100_000, "read never completed");
}
None => panic!("task disappeared before completing"),
}
}
drop(f);
}
#[test]
fn write_task_persists_data() {
let f = temp_file_with(b"");
let reactor = Rc::new(RefCell::new(
Reactor::new(8).expect("io_uring not available"),
));
let mut write = IoWriteTask::new(
reactor.clone(),
f.as_raw_fd(),
0,
b"written via io_uring".to_vec(),
);
let mut ticks = 0;
while write.poll() == Poll::Pending {
ticks += 1;
assert!(ticks < 100_000, "write never completed");
}
assert_eq!(write.result(), b"written via io_uring".len() as i32);
let mut read = IoReadTask::new(reactor, f.as_raw_fd(), 0, b"written via io_uring".len());
ticks = 0;
while read.poll() == Poll::Pending {
ticks += 1;
assert!(ticks < 100_000, "read-back never completed");
}
assert_eq!(read.buffer(), b"written via io_uring");
drop(f);
}
}