use crate::EchoError;
use crate::stream::StreamConfig;
use crate::stream::StreamProtocol;
use async_trait::async_trait;
use std::net::SocketAddr;
use std::path::PathBuf;
use std::pin::Pin;
use std::task::{Context, Poll};
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, ReadBuf};
use tokio::net::{UnixListener, UnixStream};
pub struct Protocol;
pub struct ManagedUnixStream {
inner: UnixStream,
socket_path: Option<PathBuf>, }
impl ManagedUnixStream {
fn new(stream: UnixStream) -> Self {
Self {
inner: stream,
socket_path: None,
}
}
fn with_path(stream: UnixStream, path: PathBuf) -> Self {
Self {
inner: stream,
socket_path: Some(path),
}
}
pub fn inner(&self) -> &UnixStream {
&self.inner
}
pub fn inner_mut(&mut self) -> &mut UnixStream {
&mut self.inner
}
}
impl AsyncRead for ManagedUnixStream {
fn poll_read(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<std::io::Result<()>> {
Pin::new(&mut self.inner).poll_read(cx, buf)
}
}
impl AsyncWrite for ManagedUnixStream {
fn poll_write(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<Result<usize, std::io::Error>> {
Pin::new(&mut self.inner).poll_write(cx, buf)
}
fn poll_flush(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Result<(), std::io::Error>> {
Pin::new(&mut self.inner).poll_flush(cx)
}
fn poll_shutdown(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Result<(), std::io::Error>> {
Pin::new(&mut self.inner).poll_shutdown(cx)
}
}
impl Drop for ManagedUnixStream {
fn drop(&mut self) {
if let Some(path) = &self.socket_path {
let _ = std::fs::remove_file(path);
}
}
}
pub struct ManagedUnixListener {
inner: UnixListener,
socket_path: PathBuf,
}
impl ManagedUnixListener {
fn new(listener: UnixListener, path: PathBuf) -> Self {
Self {
inner: listener,
socket_path: path,
}
}
pub async fn accept(&mut self) -> Result<(ManagedUnixStream, SocketAddr), std::io::Error> {
let (stream, _) = self.inner.accept().await?;
let dummy_addr = SocketAddr::new(std::net::IpAddr::V4(std::net::Ipv4Addr::UNSPECIFIED), 0);
Ok((ManagedUnixStream::new(stream), dummy_addr))
}
}
impl Drop for ManagedUnixListener {
fn drop(&mut self) {
let _ = std::fs::remove_file(&self.socket_path);
}
}
#[async_trait]
impl StreamProtocol for Protocol {
type Error = EchoError;
type Listener = ManagedUnixListener;
type Stream = ManagedUnixStream;
async fn bind(_config: &StreamConfig) -> std::result::Result<Self::Listener, EchoError> {
let socket_path = PathBuf::from("/tmp/echosrv_stream.sock");
let _ = std::fs::remove_file(&socket_path);
if let Some(parent) = socket_path.parent() {
if !parent.exists() {
std::fs::create_dir_all(parent).map_err(|e| {
EchoError::Unix(std::io::Error::other(
format!("Failed to create directory {}: {}", parent.display(), e),
))
})?;
}
}
let listener = UnixListener::bind(&socket_path).map_err(EchoError::Unix)?;
Ok(ManagedUnixListener::new(listener, socket_path))
}
async fn accept(
listener: &mut Self::Listener,
) -> std::result::Result<(Self::Stream, SocketAddr), EchoError> {
listener.accept().await.map_err(EchoError::Unix)
}
async fn connect(_addr: SocketAddr) -> std::result::Result<Self::Stream, EchoError> {
Err(EchoError::Unsupported(
"Use connect_unix method for Unix domain socket connections".to_string(),
))
}
async fn read(
stream: &mut Self::Stream,
buffer: &mut [u8],
) -> std::result::Result<usize, EchoError> {
stream.inner.read(buffer).await.map_err(EchoError::Unix)
}
async fn write(stream: &mut Self::Stream, data: &[u8]) -> std::result::Result<(), EchoError> {
stream.inner.write_all(data).await.map_err(EchoError::Unix)
}
async fn flush(stream: &mut Self::Stream) -> std::result::Result<(), EchoError> {
stream.inner.flush().await.map_err(EchoError::Unix)
}
fn map_io_error(err: std::io::Error) -> EchoError {
EchoError::Unix(err)
}
}
#[async_trait]
pub trait StreamExt {
async fn bind_unix(
socket_path: &PathBuf,
) -> std::result::Result<ManagedUnixListener, EchoError>;
async fn connect_unix(
socket_path: &PathBuf,
) -> std::result::Result<ManagedUnixStream, EchoError>;
async fn connect_anonymous() -> std::result::Result<ManagedUnixStream, EchoError>;
}
#[async_trait]
impl StreamExt for Protocol {
async fn bind_unix(
socket_path: &PathBuf,
) -> std::result::Result<ManagedUnixListener, EchoError> {
let _ = std::fs::remove_file(socket_path);
if let Some(parent) = socket_path.parent() {
if !parent.exists() {
std::fs::create_dir_all(parent).map_err(|e| {
EchoError::Unix(std::io::Error::other(
format!("Failed to create directory {}: {}", parent.display(), e),
))
})?;
}
}
let listener = UnixListener::bind(socket_path).map_err(EchoError::Unix)?;
Ok(ManagedUnixListener::new(listener, socket_path.clone()))
}
async fn connect_unix(
socket_path: &PathBuf,
) -> std::result::Result<ManagedUnixStream, EchoError> {
let stream = UnixStream::connect(socket_path)
.await
.map_err(EchoError::Unix)?;
Ok(ManagedUnixStream::new(stream))
}
async fn connect_anonymous() -> std::result::Result<ManagedUnixStream, EchoError> {
let temp_dir = std::env::temp_dir();
let client_socket_path = temp_dir.join(format!(
"client_{}_{}.sock",
std::process::id(),
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
));
let _ = std::fs::remove_file(&client_socket_path);
let stream = UnixStream::connect(&client_socket_path)
.await
.map_err(EchoError::Unix)?;
Ok(ManagedUnixStream::with_path(stream, client_socket_path))
}
}
#[cfg(test)]
mod tests {
use super::*;
use tempfile::tempdir;
#[tokio::test]
async fn test_unix_stream_bind_and_cleanup() {
let temp_dir = tempdir().unwrap();
let socket_path = temp_dir.path().join("test.sock");
assert!(!socket_path.exists());
let _socket_path_clone = socket_path.clone();
}
#[tokio::test]
async fn test_unix_stream_connect() {
let temp_dir = tempdir().unwrap();
let socket_path = temp_dir.path().join("test_connect.sock");
let connect_result = Protocol::connect_unix(&socket_path).await;
assert!(connect_result.is_err());
}
}