use safer_ring::Ring;
#[cfg(target_os = "linux")]
use std::io::{self, Read, Write};
#[cfg(target_os = "linux")]
use std::net::{TcpListener, TcpStream};
#[cfg(target_os = "linux")]
use std::os::unix::io::{AsRawFd, RawFd};
#[cfg(target_os = "linux")]
use std::time::Duration;
#[cfg(target_os = "linux")]
use tokio::time::timeout;
#[cfg(target_os = "linux")]
use safer_ring::PinnedBuffer;
#[cfg(target_os = "linux")]
fn create_test_listener() -> io::Result<(TcpListener, u16)> {
let listener = TcpListener::bind("127.0.0.1:0")?;
let port = listener.local_addr()?.port();
listener.set_nonblocking(true)?;
Ok((listener, port))
}
#[cfg(target_os = "linux")]
fn create_connected_pair() -> io::Result<(TcpStream, TcpStream)> {
let listener = TcpListener::bind("127.0.0.1:0")?;
let addr = listener.local_addr()?;
let client = TcpStream::connect(addr)?;
let (server, _) = listener.accept()?;
server.set_nonblocking(true)?;
client.set_nonblocking(true)?;
Ok((server, client))
}
#[cfg(target_os = "linux")]
#[tokio::test]
async fn test_accept_operation() -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
let ring = Ring::new(32)?;
let (listener, port) = create_test_listener()?;
let listener_fd = listener.as_raw_fd();
let connect_handle =
tokio::spawn(async move { TcpStream::connect(format!("127.0.0.1:{port}")) });
let client_fd = timeout(Duration::from_secs(5), ring.accept_safe(listener_fd)).await??;
assert!(client_fd >= 0, "Accepted file descriptor should be valid");
connect_handle.await??;
unsafe { libc::close(client_fd) };
Ok(())
}
#[cfg(target_os = "linux")]
#[tokio::test]
async fn test_send_operation() -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
let (mut server, client) = create_connected_pair()?;
let client_fd = client.as_raw_fd();
let test_data = b"Hello, safer-ring!";
let ring = Ring::new(32)?;
let ring: &'static mut Ring = Box::leak(Box::new(ring));
let send_buffer = PinnedBuffer::from_slice(test_data);
let send_buffer: &'static mut PinnedBuffer<_> = Box::leak(Box::new(send_buffer));
let (bytes_sent, _) = ring.send(client_fd, send_buffer.as_mut_slice())?.await?;
assert_eq!(bytes_sent, test_data.len());
let mut received_data = vec![0u8; test_data.len()];
server.read_exact(&mut received_data)?;
assert_eq!(received_data.as_slice(), test_data);
Ok(())
}
#[cfg(target_os = "linux")]
#[tokio::test]
async fn test_recv_operation() -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
let (mut server, client) = create_connected_pair()?;
let client_fd = client.as_raw_fd();
let test_data = b"Hello, safer-ring!";
server.write_all(test_data)?;
server.flush()?;
let ring = Ring::new(32)?;
let ring: &'static mut Ring = Box::leak(Box::new(ring));
let recv_buffer = PinnedBuffer::with_capacity(1024);
let recv_buffer: &'static mut PinnedBuffer<_> = Box::leak(Box::new(recv_buffer));
let (bytes_received, received_slice) =
ring.recv(client_fd, recv_buffer.as_mut_slice())?.await?;
assert_eq!(bytes_received, test_data.len());
assert_eq!(&received_slice[..bytes_received], test_data);
Ok(())
}
#[cfg(target_os = "linux")]
#[tokio::test(flavor = "multi_thread")]
async fn test_network_echo_server() -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
let (listener, port) = create_test_listener()?;
let listener_fd = listener.as_raw_fd();
let client_handle = tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(50)).await;
let mut client = TcpStream::connect(format!("127.0.0.1:{port}"))?;
let test_data = b"Echo test message";
client.write_all(test_data)?;
client.flush()?;
let mut received = vec![0u8; test_data.len()];
client.read_exact(&mut received)?;
Ok::<_, Box<dyn std::error::Error + Send + Sync>>(received)
});
let ring1 = Ring::new(32)?;
let ring1: &'static Ring = Box::leak(Box::new(ring1));
let client_fd = timeout(Duration::from_secs(5), ring1.accept_safe(listener_fd)).await??;
let ring2 = Ring::new(32)?;
let ring2: &'static mut Ring = Box::leak(Box::new(ring2));
let recv_buffer = PinnedBuffer::with_capacity(1024);
let recv_buffer: &'static mut PinnedBuffer<_> = Box::leak(Box::new(recv_buffer));
let (bytes_received, received_slice) = timeout(
Duration::from_secs(5),
ring2.recv(client_fd, recv_buffer.as_mut_slice())?,
)
.await??;
let received_data = received_slice[..bytes_received].to_vec();
let ring3 = Ring::new(32)?;
let ring3: &'static mut Ring = Box::leak(Box::new(ring3));
let send_buffer = PinnedBuffer::from_slice(&received_data);
let send_buffer: &'static mut PinnedBuffer<_> = Box::leak(Box::new(send_buffer));
let (bytes_sent, _) = timeout(
Duration::from_secs(5),
ring3.send(client_fd, send_buffer.as_mut_slice())?,
)
.await??;
unsafe { libc::close(client_fd) };
let client_received = client_handle.await??;
assert_eq!(bytes_received, bytes_sent);
assert_eq!(client_received.as_slice(), b"Echo test message");
Ok(())
}
#[cfg(target_os = "linux")]
#[tokio::test(flavor = "multi_thread")]
async fn test_multiple_concurrent_connections(
) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
let (listener, port) = create_test_listener()?;
let listener_fd = listener.as_raw_fd();
const NUM_CONNECTIONS: usize = 3;
let mut client_tasks = Vec::new();
for i in 0..NUM_CONNECTIONS {
client_tasks.push(tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(i as u64 * 30)).await;
let mut client = TcpStream::connect(format!("127.0.0.1:{port}"))?;
let test_data = format!("Message from client {i}");
client.write_all(test_data.as_bytes())?;
client.flush()?;
let mut received = vec![0; test_data.len()];
client.read_exact(&mut received)?;
Ok::<_, Box<dyn std::error::Error + Send + Sync>>(String::from_utf8(received)?)
}));
}
let accept_ring = Ring::new(32)?;
let accept_ring: &'static Ring = Box::leak(Box::new(accept_ring));
let mut server_results = Vec::new();
for _ in 0..NUM_CONNECTIONS {
let client_fd =
timeout(Duration::from_secs(5), accept_ring.accept_safe(listener_fd)).await??;
let conn_ring_recv = Ring::new(32)?;
let conn_ring_recv: &'static mut Ring = Box::leak(Box::new(conn_ring_recv));
let recv_buffer = PinnedBuffer::with_capacity(1024);
let recv_buffer: &'static mut PinnedBuffer<_> = Box::leak(Box::new(recv_buffer));
let (bytes_received, received_slice) = timeout(
Duration::from_secs(5),
conn_ring_recv.recv(client_fd, recv_buffer.as_mut_slice())?,
)
.await??;
let received_data = received_slice[..bytes_received].to_vec();
let conn_ring_send = Ring::new(32)?;
let conn_ring_send: &'static mut Ring = Box::leak(Box::new(conn_ring_send));
let send_buffer = PinnedBuffer::from_slice(&received_data);
let send_buffer: &'static mut PinnedBuffer<_> = Box::leak(Box::new(send_buffer));
let (bytes_sent, _) = timeout(
Duration::from_secs(5),
conn_ring_send.send(client_fd, send_buffer.as_mut_slice())?,
)
.await??;
server_results.push((bytes_received, bytes_sent));
unsafe { libc::close(client_fd) };
}
let mut client_results = Vec::new();
for task in client_tasks {
client_results.push(task.await??);
}
assert_eq!(server_results.len(), NUM_CONNECTIONS);
assert_eq!(client_results.len(), NUM_CONNECTIONS);
client_results.sort();
let mut server_results_sorted = server_results;
server_results_sorted.sort_by_key(|(r, _)| *r);
for i in 0..NUM_CONNECTIONS {
let expected_message = format!("Message from client {i}");
assert_eq!(client_results[i], expected_message);
assert_eq!(server_results_sorted[i].0, expected_message.len());
assert_eq!(server_results_sorted[i].1, expected_message.len());
}
Ok(())
}
#[cfg(target_os = "linux")]
#[tokio::test]
async fn test_network_send_error_handling() -> Result<(), Box<dyn std::error::Error + Send + Sync>>
{
let invalid_fd: RawFd = -1;
let ring = Ring::new(32)?;
let ring: &'static mut Ring = Box::leak(Box::new(ring));
let send_buffer = PinnedBuffer::with_capacity(1024);
let send_buffer: &'static mut PinnedBuffer<_> = Box::leak(Box::new(send_buffer));
let send_result = ring.send(invalid_fd, send_buffer.as_mut_slice());
assert!(
send_result.is_err(),
"send with invalid fd should fail immediately"
);
Ok(())
}
#[cfg(target_os = "linux")]
#[tokio::test]
async fn test_network_recv_error_handling() -> Result<(), Box<dyn std::error::Error + Send + Sync>>
{
let invalid_fd: RawFd = -1;
let ring = Ring::new(32)?;
let ring: &'static mut Ring = Box::leak(Box::new(ring));
let recv_buffer = PinnedBuffer::with_capacity(1024);
let recv_buffer: &'static mut PinnedBuffer<_> = Box::leak(Box::new(recv_buffer));
let recv_result = ring.recv(invalid_fd, recv_buffer.as_mut_slice());
assert!(
recv_result.is_err(),
"recv with invalid fd should fail immediately"
);
Ok(())
}
#[cfg(target_os = "linux")]
#[tokio::test]
async fn test_network_accept_error_handling() -> Result<(), Box<dyn std::error::Error + Send + Sync>>
{
let invalid_fd: RawFd = -1;
let ring = Ring::new(32)?;
let ring: &'static Ring = Box::leak(Box::new(ring));
let accept_result = ring.accept_safe(invalid_fd).await;
assert!(accept_result.is_err(), "accept with invalid fd should fail");
Ok(())
}
#[cfg(not(target_os = "linux"))]
#[tokio::test]
async fn test_network_operations_unsupported_platform() {
let ring_result = Ring::new(32);
assert!(
ring_result.is_err(),
"Ring creation should fail on non-Linux platforms."
);
}