extern crate aio_bindings;
extern crate futures;
extern crate libc;
extern crate mio;
extern crate rand;
extern crate tokio;
extern crate futures_cpupool;
extern crate memmap;
use std::collections;
use std::io;
use std::mem;
use std::ops;
use std::ptr;
use std::sync;
use std::os::unix::io::RawFd;
use libc::{c_long, c_void, mlock};
use futures::Future;
use ops::Deref;
use aio_bindings::{aio_context_t, io_event, iocb, syscall, timespec, __NR_io_destroy,
__NR_io_getevents, __NR_io_setup, __NR_io_submit, IOCB_CMD_PREAD,
IOCB_CMD_PWRITE, IOCB_FLAG_RESFD};
mod eventfd;
#[inline(always)]
unsafe fn io_setup(nr: c_long, ctxp: *mut aio_context_t) -> c_long {
syscall(__NR_io_setup as c_long, nr, ctxp)
}
#[inline(always)]
unsafe fn io_destroy(ctx: aio_context_t) -> c_long {
syscall(__NR_io_destroy as c_long, ctx)
}
#[inline(always)]
unsafe fn io_submit(ctx: aio_context_t, nr: c_long, iocbpp: *mut *mut iocb) -> c_long {
syscall(__NR_io_submit as c_long, ctx, nr, iocbpp)
}
#[inline(always)]
unsafe fn io_getevents(
ctx: aio_context_t,
min_nr: c_long,
max_nr: c_long,
events: *mut io_event,
timeout: *mut timespec,
) -> c_long {
syscall(
__NR_io_getevents as c_long,
ctx,
min_nr,
max_nr,
events,
timeout,
)
}
struct SemaphoreInner {
capacity: usize,
waiters: collections::VecDeque<futures::sync::oneshot::Sender<()>>,
}
struct Semaphore {
inner: sync::RwLock<SemaphoreInner>,
}
impl Semaphore {
fn new(initial: usize) -> Semaphore {
Semaphore {
inner: sync::RwLock::new(SemaphoreInner {
capacity: initial,
waiters: collections::VecDeque::new(),
}),
}
}
fn acquire(&self) -> SemaphoreHandle {
let mut lock_result = self.inner.write();
match lock_result {
Ok(ref mut guard) => {
if guard.capacity > 0 {
guard.capacity -= 1;
SemaphoreHandle::Completed(futures::future::result(Ok(())))
} else {
let (sender, receiver) = futures::sync::oneshot::channel();
guard.waiters.push_back(sender);
SemaphoreHandle::Waiting(receiver)
}
}
Err(err) => panic!("Lock failure {:?}", err),
}
}
fn release(&self) {
let mut lock_result = self.inner.write();
match lock_result {
Ok(ref mut guard) => {
if !guard.waiters.is_empty() {
guard.waiters.pop_front().unwrap().send(()).unwrap();
} else {
guard.capacity += 1;
}
}
Err(err) => panic!("Lock failure {:?}", err),
}
}
fn current_capacity(&self) -> usize {
let mut lock_result = self.inner.read();
match lock_result {
Ok(ref mut guard) => guard.capacity,
Err(_) => panic!("Lock failure"),
}
}
}
enum SemaphoreHandle {
Waiting(futures::sync::oneshot::Receiver<()>),
Completed(futures::future::FutureResult<(), io::Error>),
}
impl futures::Future for SemaphoreHandle {
type Item = ();
type Error = io::Error;
fn poll(&mut self) -> Result<futures::Async<()>, io::Error> {
match self {
&mut SemaphoreHandle::Completed(_) => Ok(futures::Async::Ready(())),
&mut SemaphoreHandle::Waiting(ref mut receiver) => receiver
.poll()
.map_err(|err| io::Error::new(io::ErrorKind::Other, err)),
}
}
}
struct AioBaseFuture {
context: sync::Arc<AioContextInner>,
opcode: u32,
fd: RawFd,
offset: u64,
buf: u64,
len: u64,
submitted: sync::atomic::AtomicBool,
state: Option<Box<RequestState>>,
acquire_state: Option<SemaphoreHandle>,
}
impl AioBaseFuture {
fn poll(&mut self) -> Result<futures::Async<()>, io::Error> {
if !self.submitted.load(sync::atomic::Ordering::Acquire) {
assert!(self.state.is_none());
if self.acquire_state.is_none() {
self.acquire_state = Some(self.context.have_capacity.acquire());
}
match self.acquire_state.as_mut().unwrap().poll() {
Err(err) => return Err(err),
Ok(futures::Async::NotReady) => return Ok(futures::Async::NotReady),
Ok(futures::Async::Ready(_)) => {
let mut guard = self.context.capacity.write();
match guard {
Ok(ref mut guard) => {
self.state = guard.state.pop();
}
Err(_) => panic!("TODO: Figure out how to handle this kind of error"),
}
}
}
{
assert!(self.state.is_some());
let state = self.state.as_mut().unwrap();
let state_addr = state.deref().deref() as *const RequestState;
state.request.aio_data = unsafe { mem::transmute(state_addr) };
state.request.aio_resfd = self.context.completed_fd as u32;
state.request.aio_flags = IOCB_FLAG_RESFD;
state.request.aio_fildes = self.fd as u32;
state.request.aio_offset = self.offset as i64;
state.request.aio_buf = self.buf;
state.request.aio_nbytes = self.len;
state.request.aio_lio_opcode = self.opcode as u16;
let (sender, receiver) = futures::sync::oneshot::channel();
state.completed_receiver = receiver;
state.completed_sender = Some(sender);
}
let mut request_ptr_array: [*mut iocb; 1] =
[&mut self.state.as_mut().unwrap().request as *mut iocb; 1];
let result = unsafe {
io_submit(
self.context.context,
1,
&mut request_ptr_array[0] as *mut *mut iocb,
)
};
self.submitted.store(true, sync::atomic::Ordering::Release);
if result != 1 {
return Err(io::Error::last_os_error());
}
}
let result_code = match self.state.as_mut().unwrap().completed_receiver.poll() {
Err(err) => return Err(io::Error::new(io::ErrorKind::Other, err)),
Ok(futures::Async::NotReady) => return Ok(futures::Async::NotReady),
Ok(futures::Async::Ready(n)) => n,
};
match self.context.capacity.write() {
Ok(ref mut guard) => {
guard.state.push(self.state.take().unwrap());
}
Err(_) => panic!("TODO: Figure out how to handle this kind of error"),
}
self.context.have_capacity.release();
if result_code < 0 {
Err(io::Error::from_raw_os_error(result_code as i32))
} else {
Ok(futures::Async::Ready(()))
}
}
}
pub struct AioReadResultFuture<ReadWriteHandle>
where
ReadWriteHandle: ops::DerefMut<Target = [u8]>,
{
base: AioBaseFuture,
buffer: ReadWriteHandle,
}
impl<ReadWriteHandle> futures::Future for AioReadResultFuture<ReadWriteHandle>
where
ReadWriteHandle: ops::DerefMut<Target = [u8]>,
{
type Item = ();
type Error = io::Error;
fn poll(&mut self) -> Result<futures::Async<Self::Item>, Self::Error> {
self.base.poll()
}
}
pub struct AioWriteResultFuture<ReadOnlyHandle>
where
ReadOnlyHandle: ops::Deref<Target = [u8]>,
{
base: AioBaseFuture,
buffer: ReadOnlyHandle,
}
impl<ReadOnlyHandle> futures::Future for AioWriteResultFuture<ReadOnlyHandle>
where
ReadOnlyHandle: ops::Deref<Target = [u8]>,
{
type Item = ();
type Error = io::Error;
fn poll(&mut self) -> Result<futures::Async<Self::Item>, Self::Error> {
self.base.poll()
}
}
struct RequestState {
request: iocb,
completed_receiver: futures::sync::oneshot::Receiver<c_long>,
completed_sender: Option<futures::sync::oneshot::Sender<c_long>>,
}
struct Capacity {
state: Vec<Box<RequestState>>,
}
impl Capacity {
fn new(nr: usize) -> Result<Capacity, io::Error> {
let mut state = Vec::with_capacity(nr);
for _ in 0..nr {
let (sender, receiver) = futures::sync::oneshot::channel();
state.push(Box::new(RequestState {
request: unsafe { mem::zeroed() },
completed_receiver: receiver,
completed_sender: None,
}));
}
Ok(Capacity { state })
}
}
struct AioContextInner {
context: aio_context_t,
completed_fd: RawFd,
have_capacity: Semaphore,
capacity: sync::RwLock<Capacity>,
}
impl AioContextInner {
fn new(fd: RawFd, nr: usize) -> Result<AioContextInner, io::Error> {
let mut context: aio_context_t = 0;
unsafe {
if io_setup(nr as c_long, &mut context) != 0 {
return Err(io::Error::last_os_error());
}
};
Ok(AioContextInner {
context,
capacity: sync::RwLock::new(Capacity::new(nr)?),
have_capacity: Semaphore::new(nr),
completed_fd: fd,
})
}
}
impl Drop for AioContextInner {
fn drop(&mut self) {
let result = unsafe { io_destroy(self.context) };
assert!(result == 0);
}
}
pub struct AioContext {
inner: sync::Arc<AioContextInner>,
poll_task_handle: futures::sync::oneshot::SpawnHandle<(), io::Error>,
}
impl AioContext {
pub fn new<E>(executor: &E, nr: usize) -> Result<AioContext, io::Error>
where
E: futures::future::Executor<futures::sync::oneshot::Execute<AioPollFuture>>,
{
let eventfd = eventfd::EventFd::create(0, false)?;
let fd = eventfd.evented.get_ref().fd;
let inner = AioContextInner::new(fd, nr)?;
let context = inner.context;
let poll_future = AioPollFuture {
context,
eventfd,
events: Vec::with_capacity(nr),
};
Ok(AioContext {
inner: sync::Arc::new(inner),
poll_task_handle: futures::sync::oneshot::spawn(poll_future, executor),
})
}
pub fn read<ReadWriteHandle>(
&self,
fd: RawFd,
offset: u64,
buffer: ReadWriteHandle,
) -> AioReadResultFuture<ReadWriteHandle>
where
ReadWriteHandle: ops::DerefMut<Target = [u8]>,
{
let len = buffer.len() as u64;
AioReadResultFuture {
base: AioBaseFuture {
context: self.inner.clone(),
opcode: IOCB_CMD_PREAD,
fd,
offset,
len,
buf: unsafe { mem::transmute(buffer.as_ptr()) },
submitted: sync::atomic::AtomicBool::new(false),
state: None,
acquire_state: None,
},
buffer,
}
}
pub fn write<ReadOnlyHandle>(
&self,
fd: RawFd,
offset: u64,
buffer: ReadOnlyHandle,
) -> AioWriteResultFuture<ReadOnlyHandle>
where
ReadOnlyHandle: ops::Deref<Target = [u8]>,
{
let len = buffer.len() as u64;
AioWriteResultFuture {
base: AioBaseFuture {
context: self.inner.clone(),
opcode: IOCB_CMD_PWRITE,
fd,
offset,
len,
buf: unsafe { mem::transmute(buffer.as_ptr()) },
submitted: sync::atomic::AtomicBool::new(false),
state: None,
acquire_state: None,
},
buffer,
}
}
}
pub struct AioPollFuture {
context: aio_context_t,
eventfd: eventfd::EventFd,
events: Vec<io_event>,
}
impl futures::Future for AioPollFuture {
type Item = ();
type Error = io::Error;
fn poll(&mut self) -> Result<futures::Async<Self::Item>, Self::Error> {
loop {
let available = match self.eventfd.read() {
Err(err) => return Err(err),
Ok(futures::Async::NotReady) => return Ok(futures::Async::NotReady),
Ok(futures::Async::Ready(value)) => value as usize,
};
assert!(available > 0);
self.events.clear();
unsafe {
let result = io_getevents(
self.context,
available as c_long,
available as c_long,
self.events.as_mut_ptr(),
ptr::null_mut::<timespec>(),
);
if result < 0 {
return Err(io::Error::last_os_error());
}
assert!(result as usize == available);
self.events.set_len(available);
};
for ref event in &self.events {
let request_state: &mut RequestState = unsafe { mem::transmute(event.data) };
request_state
.completed_sender
.take()
.unwrap()
.send(event.res)
.unwrap();
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::env;
use std::fs;
use std::io::Write;
use std::os::unix::ffi::OsStrExt;
use std::path;
use std::sync;
use rand::Rng;
use tokio::executor::current_thread;
use memmap;
use futures_cpupool;
use libc::{close, open, O_DIRECT, O_RDWR};
const FILE_SIZE: u64 = 1024 * 512;
fn temp_file_name() -> path::PathBuf {
let mut rng = rand::thread_rng();
let mut result = env::temp_dir();
let filename = format!("test-aio-{}.dat", rng.gen::<u64>());
result.push(filename);
result
}
fn create_temp_file(path: &path::Path) {
let mut file = fs::File::create(path).unwrap();
let mut data: [u8; FILE_SIZE as usize] = [0; FILE_SIZE as usize];
for index in 0..data.len() {
data[index] = index as u8;
}
let result = file.write(&data).and_then(|_| file.sync_all());
assert!(result.is_ok());
}
fn remove_file(path: &path::Path) {
let _ = fs::remove_file(path);
}
#[test]
fn create_and_drop() {
let pool = futures_cpupool::CpuPool::new(3);
let _context = AioContext::new(&pool, 10).unwrap();
}
struct MemoryBlock {
bytes: sync::RwLock<memmap::MmapMut>,
}
impl MemoryBlock {
fn new() -> MemoryBlock {
let map = memmap::MmapMut::map_anon(8192).unwrap();
unsafe { mlock(map.as_ref().as_ptr() as *const c_void, map.len()) };
MemoryBlock {
bytes: sync::RwLock::new(map),
}
}
}
struct MemoryHandle {
block: sync::Arc<MemoryBlock>,
}
impl MemoryHandle {
fn new() -> MemoryHandle {
MemoryHandle {
block: sync::Arc::new(MemoryBlock::new()),
}
}
}
impl Clone for MemoryHandle {
fn clone(&self) -> MemoryHandle {
MemoryHandle {
block: self.block.clone(),
}
}
}
impl ops::Deref for MemoryHandle {
type Target = [u8];
fn deref(&self) -> &Self::Target {
unsafe { mem::transmute(&(*self.block.bytes.read().unwrap())[..]) }
}
}
impl ops::DerefMut for MemoryHandle {
fn deref_mut(&mut self) -> &mut Self::Target {
unsafe { mem::transmute(&mut (*self.block.bytes.write().unwrap())[..]) }
}
}
#[test]
fn read_block_mt() {
let file_name = temp_file_name();
create_temp_file(&file_name);
{
let owned_fd = OwnedFd::new_from_raw_fd(unsafe {
open(
mem::transmute(file_name.as_os_str().as_bytes().as_ptr()),
O_DIRECT | O_RDWR,
)
});
let fd = owned_fd.fd;
let pool = futures_cpupool::CpuPool::new(5);
let buffer = MemoryHandle::new();
let result_buffer = buffer.clone();
{
let context = AioContext::new(&pool, 10).unwrap();
let read_future = context
.read(fd, 0, buffer)
.map(move |_| {
assert!(validate_block(&result_buffer));
})
.map_err(|err| {
panic!("{:?}", err);
});
let cpu_future = pool.spawn(read_future);
let result = cpu_future.wait();
assert!(result.is_ok());
}
}
remove_file(&file_name);
}
#[test]
fn read_invalid_fd() {
let fd = 2431;
let pool = futures_cpupool::CpuPool::new(5);
let buffer = MemoryHandle::new();
let result_buffer = buffer.clone();
{
let context = AioContext::new(&pool, 10).unwrap();
let read_future = context
.read(fd, 0, buffer)
.map(move |_| {
assert!(false);
})
.map_err(|err| {
assert!(err.kind() == io::ErrorKind::Other);
err
});
let cpu_future = pool.spawn(read_future);
let result = cpu_future.wait();
assert!(result.is_err());
}
}
#[test]
fn read_many_blocks_mt() {
let file_name = temp_file_name();
create_temp_file(&file_name);
{
let owned_fd = OwnedFd::new_from_raw_fd(unsafe {
open(
mem::transmute(file_name.as_os_str().as_bytes().as_ptr()),
O_DIRECT | O_RDWR,
)
});
let fd = owned_fd.fd;
let pool = futures_cpupool::CpuPool::new(5);
{
let num_slots = 7;
let context = AioContext::new(&pool, num_slots).unwrap();
for _wave in 0..50 {
let mut futures = Vec::new();
for index in 0..100 {
let buffer = MemoryHandle::new();
let result_buffer = buffer.clone();
let read_future = context
.read(fd, (index * 8192) % FILE_SIZE, buffer)
.map(move |_| {
assert!(validate_block(&result_buffer));
})
.map_err(|err| {
panic!("{:?}", err);
});
futures.push(pool.spawn(read_future));
}
let result = futures::future::join_all(futures).wait();
assert!(result.is_ok());
assert!(context.inner.have_capacity.current_capacity() == num_slots);
}
}
}
remove_file(&file_name);
}
fn validate_block(data: &[u8]) -> bool {
for index in 0..data.len() {
if data[index] != index as u8 {
return false;
}
}
true
}
struct OwnedFd {
fd: RawFd,
}
impl OwnedFd {
fn new_from_raw_fd(fd: RawFd) -> OwnedFd {
OwnedFd { fd }
}
}
impl Drop for OwnedFd {
fn drop(&mut self) {
let result = unsafe { close(self.fd) };
assert!(result == 0);
}
}
}