use std::io::{self, BufWriter};
use std::os::fd::OwnedFd;
use std::os::unix::fs::{FileTypeExt, MetadataExt, PermissionsExt as _};
use std::os::unix::net::{UnixListener, UnixStream};
use std::path::{Path, PathBuf};
use std::sync::mpsc::RecvTimeoutError;
#[cfg(test)]
use std::sync::mpsc::SyncSender;
use std::time::Duration;
use std::{fmt, fs, net as path_std_net};
use tau_proto::{
DecodeError, HarnessInputMessage, HarnessInputReader, HarnessOutputMessage,
HarnessOutputWriter, PeerOutputWriter,
};
use self::reader_worker::ReaderWorker;
mod reader_worker;
#[derive(Debug)]
pub enum SocketTransportError {
CreateParentDirectory {
path: PathBuf,
source: io::Error,
},
RefuseNonSocketPath {
path: PathBuf,
},
ActiveSocketExists {
path: PathBuf,
},
ProbeExistingSocket {
path: PathBuf,
source: io::Error,
},
RemoveStaleSocket {
path: PathBuf,
source: io::Error,
},
Bind {
path: PathBuf,
source: io::Error,
},
BoundSocketMetadata {
path: PathBuf,
source: io::Error,
},
Accept {
source: io::Error,
},
Connect {
path: PathBuf,
source: io::Error,
},
Clone {
source: io::Error,
},
SpawnReader {
source: io::Error,
},
Encode {
source: tau_proto::EncodeError,
},
Flush {
source: io::Error,
},
Decode {
source: DecodeError,
},
}
impl fmt::Display for SocketTransportError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::CreateParentDirectory { path, source } => write!(
f,
"failed to create socket parent directory {}: {source}",
path.display()
),
Self::RefuseNonSocketPath { path } => {
write!(f, "refusing to replace non-socket path {}", path.display())
}
Self::ActiveSocketExists { path } => {
write!(
f,
"refusing to replace active Unix socket {}",
path.display()
)
}
Self::ProbeExistingSocket { path, source } => write!(
f,
"refusing to replace Unix socket {} after liveness probe failed: {source}",
path.display()
),
Self::RemoveStaleSocket { path, source } => write!(
f,
"failed to remove stale socket {}: {source}",
path.display()
),
Self::Bind { path, source } => {
write!(f, "failed to bind Unix socket {}: {source}", path.display())
}
Self::BoundSocketMetadata { path, source } => write!(
f,
"failed to inspect bound Unix socket {}: {source}",
path.display()
),
Self::Accept { source } => write!(f, "failed to accept Unix socket client: {source}"),
Self::Connect { path, source } => {
write!(
f,
"failed to connect to Unix socket {}: {source}",
path.display()
)
}
Self::Clone { source } => write!(f, "failed to clone Unix socket stream: {source}"),
Self::SpawnReader { source } => {
write!(f, "failed to spawn Unix socket reader: {source}")
}
Self::Encode { source } => write!(f, "failed to encode socket event: {source}"),
Self::Flush { source } => write!(f, "failed to flush socket stream: {source}"),
Self::Decode { source } => write!(f, "failed to decode socket event: {source}"),
}
}
}
impl std::error::Error for SocketTransportError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::CreateParentDirectory { source, .. } => Some(source),
Self::RefuseNonSocketPath { .. } | Self::ActiveSocketExists { .. } => None,
Self::ProbeExistingSocket { source, .. } => Some(source),
Self::RemoveStaleSocket { source, .. } => Some(source),
Self::Bind { source, .. } => Some(source),
Self::BoundSocketMetadata { source, .. } => Some(source),
Self::Accept { source } => Some(source),
Self::Connect { source, .. } => Some(source),
Self::Clone { source } => Some(source),
Self::SpawnReader { source } => Some(source),
Self::Encode { source } => Some(source),
Self::Flush { source } => Some(source),
Self::Decode { source } => Some(source),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct SocketIdentity {
dev: u64,
ino: u64,
}
impl SocketIdentity {
fn from_metadata(metadata: &fs::Metadata) -> Self {
Self {
dev: metadata.dev(),
ino: metadata.ino(),
}
}
}
pub struct SocketListener {
path: PathBuf,
listener: UnixListener,
socket_identity: SocketIdentity,
}
impl SocketListener {
pub fn bind_fresh(path: impl Into<PathBuf>) -> Result<Self, SocketTransportError> {
let path = path.into();
let listener = UnixListener::bind(&path).map_err(|source| SocketTransportError::Bind {
path: path.clone(),
source,
})?;
let metadata = fs::symlink_metadata(&path).map_err(|source| {
SocketTransportError::BoundSocketMetadata {
path: path.clone(),
source,
}
})?;
fs::set_permissions(&path, fs::Permissions::from_mode(0o600)).map_err(|source| {
SocketTransportError::BoundSocketMetadata {
path: path.clone(),
source,
}
})?;
Ok(Self {
path,
listener,
socket_identity: SocketIdentity::from_metadata(&metadata),
})
}
pub fn bind(path: impl Into<PathBuf>) -> Result<Self, SocketTransportError> {
let path = path.into();
create_socket_parent_if_needed(&path)?;
remove_inactive_stale_socket(&path)?;
Self::bind_fresh(path)
}
#[must_use]
pub fn path(&self) -> &Path {
&self.path
}
pub fn try_clone_raw_listener(&self) -> Result<UnixListener, SocketTransportError> {
self.listener
.try_clone()
.map_err(|source| SocketTransportError::Clone { source })
}
pub fn accept(&self) -> Result<SocketAcceptedClient, SocketTransportError> {
let (stream, _) = self
.listener
.accept()
.map_err(|source| SocketTransportError::Accept { source })?;
SocketAcceptedClient::new(stream)
}
}
impl Drop for SocketListener {
fn drop(&mut self) {
let Ok(metadata) = fs::symlink_metadata(&self.path) else {
return;
};
if !metadata.file_type().is_socket() {
return;
}
if SocketIdentity::from_metadata(&metadata) == self.socket_identity {
let _ = fs::remove_file(&self.path);
}
}
}
pub struct SocketAcceptedClient {
reader: HarnessInputReader<UnixStream>,
writer: HarnessOutputWriter<BufWriter<UnixStream>>,
}
impl SocketAcceptedClient {
fn new(stream: UnixStream) -> Result<Self, SocketTransportError> {
let writer_stream = stream
.try_clone()
.map_err(|source| SocketTransportError::Clone { source })?;
Ok(Self {
reader: HarnessInputReader::new(stream),
writer: HarnessOutputWriter::new(BufWriter::new(writer_stream)),
})
}
pub fn recv(&mut self) -> Result<Option<HarnessInputMessage>, SocketTransportError> {
self.reader
.read_message()
.map_err(|source| SocketTransportError::Decode { source })
}
pub fn send(&mut self, message: &HarnessOutputMessage) -> Result<(), SocketTransportError> {
self.writer
.write_message(message)
.map_err(|source| SocketTransportError::Encode { source })?;
self.writer
.flush()
.map_err(|source| SocketTransportError::Flush { source })
}
}
#[derive(Debug, PartialEq)]
pub enum SocketReceive {
Message {
message: HarnessOutputMessage,
},
Timeout,
Closed,
}
pub struct SocketPeer {
writer: PeerOutputWriter<BufWriter<UnixStream>>,
reader_worker: Option<ReaderWorker>,
shutdown_stream: UnixStream,
}
impl SocketPeer {
pub fn connect(path: impl Into<PathBuf>) -> Result<Self, SocketTransportError> {
let path = path.into();
let stream =
UnixStream::connect(&path).map_err(|source| SocketTransportError::Connect {
path: path.clone(),
source,
})?;
Self::new(stream)
}
pub fn connect_with_io_timeout(
path: impl Into<PathBuf>,
timeout: Duration,
) -> Result<Self, SocketTransportError> {
Self::connect_with_timeouts(path, timeout, timeout)
}
pub fn connect_with_timeouts(
path: impl Into<PathBuf>,
connect_timeout: Duration,
io_timeout: Duration,
) -> Result<Self, SocketTransportError> {
let path = path.into();
let stream = connect_unix_with_timeout(&path, connect_timeout).map_err(|source| {
SocketTransportError::Connect {
path: path.clone(),
source,
}
})?;
stream
.set_read_timeout(Some(io_timeout))
.map_err(|source| SocketTransportError::Connect {
path: path.clone(),
source,
})?;
stream
.set_write_timeout(Some(io_timeout))
.map_err(|source| SocketTransportError::Connect { path, source })?;
Self::new(stream)
}
fn new(stream: UnixStream) -> Result<Self, SocketTransportError> {
let writer_stream = stream
.try_clone()
.map_err(|source| SocketTransportError::Clone { source })?;
let shutdown_stream = stream
.try_clone()
.map_err(|source| SocketTransportError::Clone { source })?;
let reader_worker = ReaderWorker::spawn(stream)?;
Ok(Self {
writer: PeerOutputWriter::new(BufWriter::new(writer_stream)),
reader_worker: Some(reader_worker),
shutdown_stream,
})
}
#[cfg(test)]
fn new_with_blocked_enqueue_hook(
stream: UnixStream,
blocked_enqueue: SyncSender<()>,
) -> Result<Self, SocketTransportError> {
let writer_stream = stream
.try_clone()
.map_err(|source| SocketTransportError::Clone { source })?;
let shutdown_stream = stream
.try_clone()
.map_err(|source| SocketTransportError::Clone { source })?;
let reader_worker = ReaderWorker::spawn_with_blocked_enqueue_hook(stream, blocked_enqueue)?;
Ok(Self {
writer: PeerOutputWriter::new(BufWriter::new(writer_stream)),
reader_worker: Some(reader_worker),
shutdown_stream,
})
}
pub fn send(&mut self, message: &HarnessInputMessage) -> Result<(), SocketTransportError> {
self.writer
.write_message(message)
.map_err(|source| SocketTransportError::Encode { source })?;
self.writer
.flush()
.map_err(|source| SocketTransportError::Flush { source })
}
pub fn set_write_timeout(&self, timeout: Duration) -> Result<(), SocketTransportError> {
self.writer
.get_ref()
.get_ref()
.set_write_timeout(Some(timeout))
.map_err(|source| SocketTransportError::Flush { source })
}
pub fn recv_timeout(
&mut self,
timeout: Duration,
) -> Result<SocketReceive, SocketTransportError> {
let reader_worker = self
.reader_worker
.as_ref()
.expect("socket peer reader missing before drop");
match reader_worker.frames.recv_timeout(timeout) {
Ok(Ok(frame)) => Ok(SocketReceive::Message { message: frame }),
Ok(Err(error)) => Err(SocketTransportError::Decode { source: error }),
Err(RecvTimeoutError::Timeout) => Ok(SocketReceive::Timeout),
Err(RecvTimeoutError::Disconnected) => Ok(SocketReceive::Closed),
}
}
}
fn connect_unix_with_timeout(path: &Path, timeout: Duration) -> io::Result<UnixStream> {
let socket = socket2::Socket::new(socket2::Domain::UNIX, socket2::Type::STREAM, None)?;
socket.connect_timeout(&socket2::SockAddr::unix(path)?, timeout)?;
let fd: OwnedFd = socket.into();
Ok(fd.into())
}
impl Drop for SocketPeer {
fn drop(&mut self) {
let ReaderWorker { frames, thread } = self
.reader_worker
.take()
.expect("socket peer reader missing during drop");
drop(frames);
let _ = self.shutdown_stream.shutdown(path_std_net::Shutdown::Both);
let _ = thread.join();
}
}
fn create_socket_parent_if_needed(path: &Path) -> Result<(), SocketTransportError> {
let Some(parent) = path
.parent()
.filter(|parent| !parent.as_os_str().is_empty())
else {
return Ok(());
};
fs::create_dir_all(parent).map_err(|source| SocketTransportError::CreateParentDirectory {
path: parent.to_path_buf(),
source,
})
}
fn remove_inactive_stale_socket(path: &Path) -> Result<(), SocketTransportError> {
let Ok(metadata) = fs::symlink_metadata(path) else {
return Ok(());
};
if !metadata.file_type().is_socket() {
return Err(SocketTransportError::RefuseNonSocketPath {
path: path.to_path_buf(),
});
}
match UnixStream::connect(path) {
Ok(_) => {
return Err(SocketTransportError::ActiveSocketExists {
path: path.to_path_buf(),
});
}
Err(error) if error.kind() == io::ErrorKind::ConnectionRefused => {}
Err(source) => {
return Err(SocketTransportError::ProbeExistingSocket {
path: path.to_path_buf(),
source,
});
}
}
fs::remove_file(path).map_err(|source| SocketTransportError::RemoveStaleSocket {
path: path.to_path_buf(),
source,
})
}
#[cfg(test)]
mod tests;