use std::future::Future;
use std::io;
use std::os::unix::io::RawFd;
use std::pin::Pin;
use std::task::{Context, Poll};
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
use crate::error::{Result, SaferRingError};
use crate::ownership::OwnedBuffer;
use crate::ring::Ring;
const ADAPTER_BUFFER_SIZE: usize = 8192;
type OperationFutureResult = Result<(usize, OwnedBuffer)>;
type OwnedOperationFuture<'a> = Pin<Box<dyn Future<Output = OperationFutureResult> + 'a>>;
pub struct AsyncReadAdapter<'ring> {
ring: &'ring Ring<'ring>,
fd: RawFd,
internal_buffer: Option<OwnedBuffer>,
cached_data: Vec<u8>,
cached_pos: usize,
read_future: Option<OwnedOperationFuture<'ring>>,
}
impl<'ring> AsyncReadAdapter<'ring> {
pub fn new(ring: &'ring Ring<'ring>, fd: RawFd) -> Self {
Self {
ring,
fd,
internal_buffer: Some(OwnedBuffer::new(ADAPTER_BUFFER_SIZE)),
cached_data: Vec::new(),
cached_pos: 0,
read_future: None,
}
}
pub fn ring(&self) -> &'ring Ring<'ring> {
self.ring
}
pub fn fd(&self) -> RawFd {
self.fd
}
}
impl<'ring> AsyncRead for AsyncReadAdapter<'ring> {
fn poll_read(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
if self.cached_pos < self.cached_data.len() {
let available = self.cached_data.len() - self.cached_pos;
let to_copy = available.min(buf.remaining());
if to_copy > 0 {
let data = &self.cached_data[self.cached_pos..self.cached_pos + to_copy];
buf.put_slice(data);
self.cached_pos += to_copy;
if self.cached_pos >= self.cached_data.len() {
self.cached_data.clear();
self.cached_pos = 0;
}
return Poll::Ready(Ok(()));
}
}
if let Some(mut future) = self.read_future.take() {
match future.as_mut().poll(cx) {
Poll::Ready(Ok((bytes_read, returned_buffer))) => {
self.internal_buffer = Some(returned_buffer);
if bytes_read > 0 {
if let Some(buffer_guard) =
self.internal_buffer.as_ref().and_then(|b| b.try_access())
{
let data = &buffer_guard[..bytes_read];
let to_copy = data.len().min(buf.remaining());
buf.put_slice(&data[..to_copy]);
if to_copy < data.len() {
self.cached_data.extend_from_slice(&data[to_copy..]);
self.cached_pos = 0;
}
}
}
Poll::Ready(Ok(()))
}
Poll::Ready(Err(e)) => {
Poll::Ready(Err(io::Error::other(e.to_string())))
}
Poll::Pending => {
self.read_future = Some(future);
Poll::Pending
}
}
} else if let Some(buffer) = self.internal_buffer.take() {
let future = self.ring.read_owned(self.fd, buffer);
self.read_future = Some(Box::pin(future));
self.poll_read(cx, buf)
} else {
Poll::Ready(Err(io::Error::other(
"AsyncReadAdapter internal buffer unavailable",
)))
}
}
}
pub struct AsyncWriteAdapter<'ring> {
ring: &'ring Ring<'ring>,
fd: RawFd,
#[allow(dead_code)]
write_buffer: Vec<u8>,
internal_buffer: Option<OwnedBuffer>,
write_future: Option<OwnedOperationFuture<'ring>>,
}
impl<'ring> AsyncWriteAdapter<'ring> {
pub fn new(ring: &'ring Ring<'ring>, fd: RawFd) -> Self {
Self {
ring,
fd,
write_buffer: Vec::new(),
internal_buffer: Some(OwnedBuffer::new(ADAPTER_BUFFER_SIZE)),
write_future: None,
}
}
pub fn ring(&self) -> &'ring Ring<'ring> {
self.ring
}
pub fn fd(&self) -> RawFd {
self.fd
}
}
impl<'ring> AsyncWrite for AsyncWriteAdapter<'ring> {
fn poll_write(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
if let Some(mut future) = self.write_future.take() {
match future.as_mut().poll(cx) {
Poll::Ready(Ok((_bytes_written, returned_buffer))) => {
self.internal_buffer = Some(returned_buffer);
}
Poll::Ready(Err(e)) => {
return Poll::Ready(Err(io::Error::other(e.to_string())));
}
Poll::Pending => {
self.write_future = Some(future);
return Poll::Pending;
}
}
}
if let Some(internal_buffer) = self.internal_buffer.take() {
if let Some(mut buffer_guard) = internal_buffer.try_access() {
let to_copy = buf.len().min(buffer_guard.len());
buffer_guard[..to_copy].copy_from_slice(&buf[..to_copy]);
drop(buffer_guard);
let future = self.ring.write_owned(self.fd, internal_buffer);
self.write_future = Some(Box::pin(future));
match self.poll_write(cx, &[]) {
Poll::Ready(Ok(_)) => Poll::Ready(Ok(to_copy)),
Poll::Ready(Err(e)) => Poll::Ready(Err(e)),
Poll::Pending => Poll::Ready(Ok(to_copy)), }
} else {
self.internal_buffer = Some(internal_buffer);
Poll::Pending
}
} else {
Poll::Ready(Err(io::Error::other(
"AsyncWriteAdapter internal buffer unavailable",
)))
}
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
if let Some(mut future) = self.write_future.take() {
match future.as_mut().poll(cx) {
Poll::Ready(Ok((_bytes_written, returned_buffer))) => {
self.internal_buffer = Some(returned_buffer);
Poll::Ready(Ok(()))
}
Poll::Ready(Err(e)) => {
Poll::Ready(Err(io::Error::other(e.to_string())))
}
Poll::Pending => {
self.write_future = Some(future);
Poll::Pending
}
}
} else {
Poll::Ready(Ok(()))
}
}
fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
match self.as_mut().poll_flush(cx) {
Poll::Ready(Ok(())) => {
Poll::Ready(Ok(()))
}
Poll::Ready(Err(e)) => Poll::Ready(Err(e)),
Poll::Pending => Poll::Pending,
}
}
}
pub struct File<'ring> {
ring: &'ring Ring<'ring>,
fd: RawFd,
read_adapter: Option<AsyncReadAdapter<'ring>>,
write_adapter: Option<AsyncWriteAdapter<'ring>>,
}
impl<'ring> File<'ring> {
pub fn new(ring: &'ring Ring<'ring>, fd: RawFd) -> Self {
Self {
ring,
fd,
read_adapter: Some(AsyncReadAdapter::new(ring, fd)),
write_adapter: Some(AsyncWriteAdapter::new(ring, fd)),
}
}
pub async fn create(ring: &'ring Ring<'ring>, path: &str) -> Result<Self> {
use std::ffi::CString;
let path_cstr = CString::new(path).map_err(|_| {
SaferRingError::Io(io::Error::new(io::ErrorKind::InvalidInput, "Invalid path"))
})?;
let fd = unsafe {
libc::open(
path_cstr.as_ptr(),
libc::O_CREAT | libc::O_WRONLY | libc::O_TRUNC,
0o644,
)
};
if fd == -1 {
return Err(SaferRingError::Io(io::Error::last_os_error()));
}
Ok(Self::new(ring, fd))
}
pub async fn open(ring: &'ring Ring<'ring>, path: &str) -> Result<Self> {
use std::ffi::CString;
let path_cstr = CString::new(path).map_err(|_| {
SaferRingError::Io(io::Error::new(io::ErrorKind::InvalidInput, "Invalid path"))
})?;
let fd = unsafe { libc::open(path_cstr.as_ptr(), libc::O_RDWR, 0) };
if fd == -1 {
return Err(SaferRingError::Io(io::Error::last_os_error()));
}
Ok(Self::new(ring, fd))
}
pub async fn sync_all(&self) -> Result<()> {
Ok(())
}
pub fn fd(&self) -> RawFd {
self.fd
}
pub fn ring(&self) -> &'ring Ring<'ring> {
self.ring
}
}
impl<'ring> Drop for File<'ring> {
fn drop(&mut self) {
unsafe {
libc::close(self.fd);
}
}
}
impl<'ring> AsyncRead for File<'ring> {
fn poll_read(
mut self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
if let Some(adapter) = &mut self.read_adapter {
Pin::new(adapter).poll_read(_cx, buf)
} else {
Poll::Ready(Err(io::Error::other("Read adapter not available")))
}
}
}
impl<'ring> AsyncWrite for File<'ring> {
fn poll_write(
mut self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
if let Some(adapter) = &mut self.write_adapter {
Pin::new(adapter).poll_write(_cx, buf)
} else {
Poll::Ready(Err(io::Error::other("Write adapter not available")))
}
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
if let Some(adapter) = &mut self.write_adapter {
Pin::new(adapter).poll_flush(cx)
} else {
Poll::Ready(Ok(()))
}
}
fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
if let Some(adapter) = &mut self.write_adapter {
Pin::new(adapter).poll_shutdown(cx)
} else {
Poll::Ready(Ok(()))
}
}
}
pub struct Socket<'ring> {
ring: &'ring Ring<'ring>,
fd: RawFd,
read_adapter: AsyncReadAdapter<'ring>,
write_adapter: AsyncWriteAdapter<'ring>,
}
impl<'ring> Socket<'ring> {
pub fn new(ring: &'ring Ring<'ring>, fd: RawFd) -> Self {
Self {
ring,
fd,
read_adapter: AsyncReadAdapter::new(ring, fd),
write_adapter: AsyncWriteAdapter::new(ring, fd),
}
}
pub fn fd(&self) -> RawFd {
self.fd
}
pub fn ring(&self) -> &'ring Ring<'ring> {
self.ring
}
}
impl<'ring> AsyncRead for Socket<'ring> {
fn poll_read(
mut self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
Pin::new(&mut self.read_adapter).poll_read(_cx, buf)
}
}
impl<'ring> AsyncWrite for Socket<'ring> {
fn poll_write(
mut self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
Pin::new(&mut self.write_adapter).poll_write(_cx, buf)
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Pin::new(&mut self.write_adapter).poll_flush(cx)
}
fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Pin::new(&mut self.write_adapter).poll_shutdown(cx)
}
}
pub trait AsyncCompat<'ring> {
fn async_read(&'ring self, fd: RawFd) -> AsyncReadAdapter<'ring>;
fn async_write(&'ring self, fd: RawFd) -> AsyncWriteAdapter<'ring>;
fn file(&'ring self, fd: RawFd) -> File<'ring>;
fn socket(&'ring self, fd: RawFd) -> Socket<'ring>;
}
impl<'ring> AsyncCompat<'ring> for Ring<'ring> {
fn async_read(&'ring self, fd: RawFd) -> AsyncReadAdapter<'ring> {
AsyncReadAdapter::new(self, fd)
}
fn async_write(&'ring self, fd: RawFd) -> AsyncWriteAdapter<'ring> {
AsyncWriteAdapter::new(self, fd)
}
fn file(&'ring self, fd: RawFd) -> File<'ring> {
File::new(self, fd)
}
fn socket(&'ring self, fd: RawFd) -> Socket<'ring> {
Socket::new(self, fd)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_async_read_adapter_creation() {
#[cfg(target_os = "linux")]
{
use crate::Ring;
let _can_create_adapter: for<'r> fn(&'r Ring<'r>) -> AsyncReadAdapter<'r> =
|ring| ring.async_read(0);
}
#[cfg(not(target_os = "linux"))]
{
println!("Skipping AsyncRead adapter test on non-Linux platform");
}
}
#[tokio::test]
async fn test_async_write_adapter_creation() {
#[cfg(target_os = "linux")]
{
use crate::Ring;
let _can_create_adapter: for<'r> fn(&'r Ring<'r>) -> AsyncWriteAdapter<'r> =
|ring| ring.async_write(1);
}
#[cfg(not(target_os = "linux"))]
{
println!("Skipping AsyncWrite adapter test on non-Linux platform");
}
}
#[tokio::test]
async fn test_file_wrapper_creation() {
#[cfg(target_os = "linux")]
{
use crate::Ring;
let _can_create_file: for<'r> fn(&'r Ring<'r>) -> File<'r> = |ring| ring.file(0);
println!("✓ File wrapper method accessible");
}
#[cfg(not(target_os = "linux"))]
{
println!("Skipping File wrapper test on non-Linux platform");
}
}
#[tokio::test]
async fn test_socket_wrapper_creation() {
#[cfg(target_os = "linux")]
{
use crate::Ring;
let _can_create_socket: for<'r> fn(&'r Ring<'r>) -> Socket<'r> = |ring| ring.socket(0);
println!("✓ Socket wrapper method accessible");
}
#[cfg(not(target_os = "linux"))]
{
println!("Skipping Socket wrapper test on non-Linux platform");
}
}
#[test]
fn test_adapter_buffer_size() {
const _: () = assert!(ADAPTER_BUFFER_SIZE > 0, "Buffer size must be positive");
const _: () = assert!(
ADAPTER_BUFFER_SIZE <= 65536,
"Buffer size should be reasonable"
); }
}