use std::{
io,
net::{Shutdown, SocketAddr},
ptr::NonNull,
task::Waker,
time::Duration,
};
use bitmask_enum::bitmask;
use crate::{Description, Handle, Interest};
#[bitmask]
pub enum FileMode {
Read,
Write,
Create,
Truncate,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
pub enum OpenFlags<'a> {
None,
OpenFile(&'a str, FileMode),
Bind(&'a [SocketAddr]),
Connect(&'a [SocketAddr]),
Duration(Duration),
UserDefined(&'a [u8]),
}
impl<'a> OpenFlags<'a> {
pub fn try_into_open_file(self) -> io::Result<(&'a str, FileMode)> {
match self {
Self::OpenFile(path, mode) => Ok((path, mode)),
_ => Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!("Expect OpenFile, but got {:?}", self),
)),
}
}
pub fn try_into_bind(self) -> io::Result<&'a [SocketAddr]> {
match self {
Self::Bind(laddrs) => Ok(laddrs),
_ => Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!("Expect Bind, but got {:?}", self),
)),
}
}
pub fn try_into_connect(self) -> io::Result<&'a [SocketAddr]> {
match self {
Self::Connect(raddrs) => Ok(raddrs),
_ => Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!("Expect Bind, but got {:?}", self),
)),
}
}
pub fn try_into_duration(self) -> io::Result<Duration> {
match self {
Self::Duration(duration) => Ok(duration),
_ => Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!("Expect Bind, but got {:?}", self),
)),
}
}
pub fn try_into_user_defined(self) -> io::Result<&'a [u8]> {
match self {
Self::UserDefined(buf) => Ok(buf),
_ => Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!("Expect UserDefined, but got {:?}", self),
)),
}
}
}
#[derive(Debug)]
pub enum Cmd<'a> {
Read {
waker: Waker,
buf: &'a mut [u8],
},
Write {
waker: Waker,
buf: &'a [u8],
},
SendTo {
waker: Waker,
buf: &'a [u8],
raddr: SocketAddr,
},
RecvFrom {
waker: Waker,
buf: &'a mut [u8],
},
Register {
source: Handle,
interests: Interest,
},
ReRegister {
source: Handle,
interests: Interest,
},
Accept(Waker),
PollOnce(Option<Duration>),
TryClone,
Timeout(Waker),
LocalAddr,
RemoteAddr,
Shutdown(Shutdown),
}
#[derive(Debug, Clone)]
pub enum CmdResp {
None,
RecvFrom(usize, SocketAddr),
Incoming(Handle, SocketAddr),
DataLen(usize),
Timeout(bool),
Cloned(Handle),
SockAddr(SocketAddr),
}
impl CmdResp {
pub fn try_into_incoming(self) -> io::Result<(Handle, SocketAddr)> {
match self {
Self::Incoming(handle, raddr) => Ok((handle, raddr)),
_ => Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!("Expect Incoming, but got {:?}", self),
)),
}
}
pub fn try_into_recv_from(self) -> io::Result<(usize, SocketAddr)> {
match self {
Self::RecvFrom(len, raddr) => Ok((len, raddr)),
_ => Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!("Expect Incoming, but got {:?}", self),
)),
}
}
pub fn try_into_sockaddr(self) -> io::Result<SocketAddr> {
match self {
Self::SockAddr(addr) => Ok(addr),
_ => Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!("Expect SockAddr, but got {:?}", self),
)),
}
}
pub fn try_into_datalen(self) -> io::Result<usize> {
match self {
Self::DataLen(len) => Ok(len),
_ => Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!("Expect SockAddr, but got {:?}", self),
)),
}
}
pub fn try_into_timeout(self) -> io::Result<bool> {
match self {
Self::Timeout(status) => Ok(status),
_ => Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!("Expect Timeout, but got {:?}", self),
)),
}
}
}
pub trait RawDriver {
fn fd_open(&self, desc: Description, open_flags: OpenFlags) -> io::Result<Handle>;
fn fd_cntl(&self, handle: Handle, cmd: Cmd) -> io::Result<CmdResp>;
fn fd_close(&self, handle: Handle) -> io::Result<()>;
}
#[repr(C)]
#[derive(Clone)]
struct DriverVTable {
fd_open: unsafe fn(NonNull<DriverVTable>, Description, OpenFlags) -> io::Result<Handle>,
fd_cntl: unsafe fn(NonNull<DriverVTable>, Handle, Cmd) -> io::Result<CmdResp>,
fd_close: unsafe fn(NonNull<DriverVTable>, Handle) -> io::Result<()>,
clone: unsafe fn(NonNull<DriverVTable>) -> Driver,
drop: unsafe fn(NonNull<DriverVTable>),
}
impl DriverVTable {
fn new<R: RawDriver + Clone>() -> Self {
fn fd_open<R: RawDriver + Clone>(
ptr: NonNull<DriverVTable>,
desc: Description,
flags: OpenFlags,
) -> io::Result<Handle> {
let header = ptr.cast::<DriverHeader<R>>();
unsafe { header.as_ref().data.fd_open(desc, flags) }
}
fn fd_cntl<R: RawDriver + Clone>(
ptr: NonNull<DriverVTable>,
handle: Handle,
cmd: Cmd,
) -> io::Result<CmdResp> {
let header = ptr.cast::<DriverHeader<R>>();
let result = unsafe { header.as_ref().data.fd_cntl(handle, cmd) };
result
}
fn fd_close<R: RawDriver + Clone>(
ptr: NonNull<DriverVTable>,
handle: Handle,
) -> io::Result<()> {
let header = ptr.cast::<DriverHeader<R>>();
unsafe { header.as_ref().data.fd_close(handle) }
}
fn clone<R: RawDriver + Clone>(ptr: NonNull<DriverVTable>) -> Driver {
let driver = unsafe { ptr.cast::<DriverHeader<R>>().as_ref().clone() };
let ptr = unsafe {
NonNull::new_unchecked(Box::into_raw(Box::new(driver)) as *mut DriverVTable)
};
Driver { ptr }
}
fn drop<R: RawDriver + Clone>(ptr: NonNull<DriverVTable>) {
_ = unsafe { Box::from_raw(ptr.cast::<DriverHeader<R>>().as_ptr()) };
}
Self {
fd_open: fd_open::<R>,
fd_cntl: fd_cntl::<R>,
fd_close: fd_close::<R>,
clone: clone::<R>,
drop: drop::<R>,
}
}
}
pub trait IntoRawDriver {
type Driver: RawDriver + Clone;
fn into_raw_driver(self) -> Self::Driver;
}
#[repr(C)]
#[derive(Clone)]
struct DriverHeader<R: RawDriver + Clone> {
vtable: DriverVTable,
data: R,
}
#[derive(Debug)]
pub struct Driver {
ptr: NonNull<DriverVTable>,
}
unsafe impl Send for Driver {}
unsafe impl Sync for Driver {}
impl<R: RawDriver + Clone> From<R> for Driver {
fn from(value: R) -> Self {
Self::new(value)
}
}
impl Driver {
pub fn new<R: RawDriver + Clone>(raw: R) -> Self {
let boxed = Box::new(DriverHeader::<R> {
data: raw,
vtable: DriverVTable::new::<R>(),
});
let ptr = unsafe { NonNull::new_unchecked(Box::into_raw(boxed) as *mut DriverVTable) };
Self { ptr }
}
pub fn fd_open(&self, desc: Description, open_flags: OpenFlags) -> io::Result<Handle> {
unsafe { (self.ptr.as_ref().fd_open)(self.ptr, desc, open_flags) }
}
pub fn fd_cntl(&self, handle: Handle, cmd: Cmd) -> io::Result<CmdResp> {
unsafe { (self.ptr.as_ref().fd_cntl)(self.ptr, handle, cmd) }
}
pub fn fd_close(&self, handle: Handle) -> io::Result<()> {
unsafe { (self.ptr.as_ref().fd_close)(self.ptr, handle) }
}
}
impl Clone for Driver {
fn clone(&self) -> Self {
unsafe { (self.ptr.as_ref().clone)(self.ptr) }
}
}
impl Drop for Driver {
fn drop(&mut self) {
unsafe { (self.ptr.as_ref().drop)(self.ptr) }
}
}
#[cfg(test)]
mod tests {
use super::*;
#[derive(Clone)]
struct MockDriver {}
struct MockFile {}
impl RawDriver for MockDriver {
fn fd_open(
&self,
desc: crate::Description,
_open_flags: crate::OpenFlags,
) -> std::io::Result<crate::Handle> {
Ok(Handle::from((desc, MockFile {})))
}
fn fd_cntl(
&self,
_handle: crate::Handle,
_cmd: crate::Cmd,
) -> std::io::Result<crate::CmdResp> {
Ok(CmdResp::None)
}
fn fd_close(&self, handle: crate::Handle) -> std::io::Result<()> {
handle.drop_as::<MockFile>();
Ok(())
}
}
#[test]
fn test_driver_vtable() {
let driver = Driver::new(MockDriver {});
let driver = driver.clone();
let handle = driver.fd_open(Description::File, OpenFlags::None).unwrap();
driver.fd_cntl(handle, Cmd::PollOnce(None)).unwrap();
driver.fd_close(handle).unwrap();
}
}