use anyhow::{Context, Result};
use bytes::Bytes;
use dashmap::DashMap;
use std::os::unix::fs::FileTypeExt;
use std::path::{Path, PathBuf};
use std::sync::{Arc, OnceLock};
use std::time::Duration;
use tokio::net::UnixStream;
use tokio_util::sync::CancellationToken;
use tracing::{debug, error, info, warn};
use crate::transports::transport::{
HealthCheckError, SendBackpressure, ShutdownState, TransportError, TransportErrorHandler,
try_send_or_backpressure,
};
use velo_ext::{MessageType, PeerInfo, Transport, TransportAdapter, TransportKey, WorkerAddress};
use super::listener::UdsListener;
use crate::transports::tcp::TcpFrameCodec;
pub struct UdsTransport {
key: TransportKey,
socket_path: PathBuf,
local_address: WorkerAddress,
peers: Arc<DashMap<crate::InstanceId, PathBuf>>,
connections: Arc<DashMap<crate::InstanceId, ConnectionHandle>>,
runtime: OnceLock<tokio::runtime::Handle>,
cancel_token: CancellationToken,
shutdown_state: OnceLock<ShutdownState>,
channel_capacity: usize,
connect_timeout: Duration,
metrics: OnceLock<std::sync::Arc<dyn velo_ext::TransportObservability>>,
}
#[derive(Clone)]
struct ConnectionHandle {
tx: flume::Sender<SendTask>,
}
struct SendTask {
msg_type: MessageType,
header: Bytes,
payload: Bytes,
on_error: Arc<dyn TransportErrorHandler>,
}
impl SendTask {
fn on_error(self, error: impl Into<String>) {
self.on_error
.on_error(self.header, self.payload, error.into());
}
}
impl UdsTransport {
pub fn new(
socket_path: PathBuf,
key: TransportKey,
local_address: WorkerAddress,
channel_capacity: usize,
connect_timeout: Duration,
) -> Self {
Self {
key,
socket_path,
local_address,
peers: Arc::new(DashMap::new()),
connections: Arc::new(DashMap::new()),
runtime: OnceLock::new(),
cancel_token: CancellationToken::new(),
shutdown_state: OnceLock::new(),
channel_capacity,
connect_timeout,
metrics: OnceLock::new(),
}
}
pub fn socket_path(&self) -> &Path {
&self.socket_path
}
pub fn ensure_connected(&self, instance_id: crate::InstanceId) -> Result<()> {
self.get_or_create_connection(instance_id)?;
Ok(())
}
fn get_or_create_connection(&self, instance_id: crate::InstanceId) -> Result<ConnectionHandle> {
if let Some(handle) = self.connections.get(&instance_id) {
if !handle.tx.is_disconnected() {
return Ok(handle.clone());
}
drop(handle);
self.connections
.remove_if(&instance_id, |_, h| h.tx.is_disconnected());
self.update_connection_gauge();
}
let rt = self.runtime.get().ok_or(TransportError::NotStarted)?;
let handle = match self.connections.entry(instance_id) {
dashmap::mapref::entry::Entry::Occupied(mut entry) => {
if !entry.get().tx.is_disconnected() {
entry.get().clone()
} else {
let handle = self.create_connection(instance_id, rt)?;
entry.insert(handle.clone());
self.update_connection_gauge();
handle
}
}
dashmap::mapref::entry::Entry::Vacant(entry) => {
let handle = self.create_connection(instance_id, rt)?;
entry.insert(handle.clone());
self.update_connection_gauge();
handle
}
};
Ok(handle)
}
fn create_connection(
&self,
instance_id: crate::InstanceId,
rt: &tokio::runtime::Handle,
) -> Result<ConnectionHandle> {
let path = self
.peers
.get(&instance_id)
.ok_or(TransportError::PeerNotRegistered(instance_id))?
.value()
.clone();
let (tx, rx) = flume::bounded(self.channel_capacity);
let handle = ConnectionHandle { tx };
let cancel = self.cancel_token.clone();
let conns = Arc::clone(&self.connections);
let connect_timeout = self.connect_timeout;
let metrics = self.metrics.get().cloned();
debug!("Created new UDS connection to {} ({:?})", instance_id, path);
rt.spawn(connection_writer_task(
path,
instance_id,
rx,
conns,
cancel,
connect_timeout,
metrics,
));
Ok(handle)
}
fn update_peer_gauge(&self) {
if let Some(metrics) = self.metrics.get() {
metrics.set_registered_peers(self.peers.len());
}
}
fn update_connection_gauge(&self) {
if let Some(metrics) = self.metrics.get() {
metrics.set_active_connections(self.connections.len());
}
}
fn slow_path_send(
&self,
instance_id: crate::InstanceId,
send_msg: SendTask,
) -> Result<(), SendBackpressure> {
if self.runtime.get().is_none() {
send_msg.on_error("Transport not started");
return Ok(());
}
let handle = match self.get_or_create_connection(instance_id) {
Ok(h) => h,
Err(e) => {
send_msg.on_error(format!("Failed to create connection: {}", e));
return Ok(());
}
};
let r = try_send_or_backpressure(
&handle.tx,
send_msg,
|msg| msg.on_error("Connection closed immediately"),
|msg| msg.on_error("Connection closed"),
);
if let Some(m) = self.metrics.get()
&& r.is_err()
{
m.record_send_backpressure();
}
r
}
}
impl Transport for UdsTransport {
fn key(&self) -> TransportKey {
self.key.clone()
}
fn address(&self) -> WorkerAddress {
self.local_address.clone()
}
fn register(&self, peer_info: PeerInfo) -> Result<(), TransportError> {
let endpoint = peer_info
.worker_address()
.get_entry(&self.key)
.map_err(|_| TransportError::NoEndpoint)?
.ok_or(TransportError::NoEndpoint)?;
let path = parse_uds_endpoint(&endpoint).map_err(|e| {
error!("Failed to parse UDS endpoint: {}", e);
TransportError::InvalidEndpoint
})?;
match std::fs::metadata(&path) {
Ok(m) if m.file_type().is_socket() => {}
Ok(_) => {
debug!(
"UDS path {:?} exists but is not a socket; rejecting UDS for peer {}",
path,
peer_info.instance_id()
);
return Err(TransportError::NoEndpoint);
}
Err(_) => {
debug!(
"UDS path {:?} not visible on this host; rejecting UDS for peer {}",
path,
peer_info.instance_id()
);
return Err(TransportError::NoEndpoint);
}
}
self.peers.insert(peer_info.instance_id(), path.clone());
self.update_peer_gauge();
debug!("Registered peer {} at {:?}", peer_info.instance_id(), path);
Ok(())
}
#[inline]
fn send_message(
&self,
instance_id: crate::InstanceId,
header: Bytes,
payload: Bytes,
message_type: MessageType,
on_error: Arc<dyn TransportErrorHandler>,
) -> Result<(), SendBackpressure> {
let send_msg = SendTask {
msg_type: message_type,
header,
payload,
on_error,
};
if let Some(handle) = self.connections.get(&instance_id) {
match handle.tx.try_send(send_msg) {
Ok(()) => return Ok(()),
Err(flume::TrySendError::Full(send_msg)) => {
if let Some(m) = self.metrics.get() {
m.record_send_backpressure();
}
let tx = handle.tx.clone();
return Err(SendBackpressure::new(Box::pin(async move {
if let Err(flume::SendError(m)) = tx.send_async(send_msg).await {
m.on_error("Connection closed");
}
})));
}
Err(flume::TrySendError::Disconnected(send_msg_out)) => {
drop(handle);
self.connections
.remove_if(&instance_id, |_, h| h.tx.is_disconnected());
self.update_connection_gauge();
return self.slow_path_send(instance_id, send_msg_out);
}
}
}
self.slow_path_send(instance_id, send_msg)
}
fn start(
&self,
_instance_id: crate::InstanceId,
channels: TransportAdapter,
rt: tokio::runtime::Handle,
) -> futures::future::BoxFuture<'_, anyhow::Result<()>> {
self.runtime.set(rt.clone()).ok();
self.shutdown_state
.set(channels.shutdown_state.clone())
.ok();
let socket_path = self.socket_path.clone();
let shutdown_state = channels.shutdown_state.clone();
Box::pin(async move {
struct DefaultErrorHandler;
impl TransportErrorHandler for DefaultErrorHandler {
fn on_error(&self, _header: Bytes, _payload: Bytes, error: String) {
warn!("UDS transport error: {}", error);
}
}
if socket_path.exists() {
let is_socket = std::fs::metadata(&socket_path)
.map(|m| m.file_type().is_socket())
.unwrap_or(false);
if !is_socket {
anyhow::bail!(
"path {:?} exists and is not a Unix domain socket",
socket_path
);
}
match tokio::time::timeout(
Duration::from_millis(100),
UnixStream::connect(&socket_path),
)
.await
{
Ok(Ok(_)) => {
anyhow::bail!(
"a live UDS listener is already running at {:?}",
socket_path
);
}
_ => {
std::fs::remove_file(&socket_path).ok();
}
}
}
let uds_listener = UdsListener::builder()
.socket_path(socket_path.clone())
.adapter(channels)
.error_handler(Arc::new(DefaultErrorHandler))
.shutdown_state(shutdown_state)
.transport_key(self.key.as_str())
.metrics(self.metrics.get().cloned())
.build()?;
let bound_listener = uds_listener.bind()?;
rt.spawn(async move {
if let Err(e) = bound_listener.serve().await {
error!("UDS listener error: {}", e);
}
});
info!("UDS transport started on {:?}", socket_path);
Ok(())
})
}
fn begin_drain(&self) {
if let Some(state) = self.shutdown_state.get() {
state.begin_drain();
}
}
fn shutdown(&self) {
info!("Shutting down UDS transport");
if let Some(state) = self.shutdown_state.get() {
state.teardown_token().cancel();
}
self.cancel_token.cancel();
self.connections.clear();
self.update_connection_gauge();
}
fn set_observability(
&self,
observability: std::sync::Arc<dyn velo_ext::TransportObservability>,
) {
let _ = self.metrics.set(observability);
self.update_peer_gauge();
self.update_connection_gauge();
}
fn check_health(
&self,
instance_id: crate::InstanceId,
timeout: Duration,
) -> std::pin::Pin<
Box<dyn std::future::Future<Output = Result<(), HealthCheckError>> + Send + '_>,
> {
Box::pin(async move {
let connection_exists = self.connections.contains_key(&instance_id);
if let Some(handle) = self.connections.get(&instance_id) {
if !handle.tx.is_disconnected() {
return Ok(());
}
drop(handle);
self.connections
.remove_if(&instance_id, |_, h| h.tx.is_disconnected());
}
let path = self
.peers
.get(&instance_id)
.ok_or(HealthCheckError::PeerNotRegistered)?
.value()
.clone();
match tokio::time::timeout(timeout, UnixStream::connect(&path)).await {
Ok(Ok(_stream)) => {
if connection_exists {
Ok(())
} else {
Err(HealthCheckError::NeverConnected)
}
}
Ok(Err(_)) => Err(HealthCheckError::ConnectionFailed),
Err(_) => Err(HealthCheckError::Timeout),
}
})
}
}
async fn connection_writer_task(
path: PathBuf,
instance_id: crate::InstanceId,
rx: flume::Receiver<SendTask>,
connections: Arc<DashMap<crate::InstanceId, ConnectionHandle>>,
cancel_token: CancellationToken,
connect_timeout: Duration,
metrics: Option<std::sync::Arc<dyn velo_ext::TransportObservability>>,
) -> Result<()> {
let result =
connection_writer_inner(&path, instance_id, &rx, &cancel_token, connect_timeout).await;
while let Ok(msg) = rx.try_recv() {
msg.on_error("Connection closed");
}
drop(rx);
connections.remove_if(&instance_id, |_, h| h.tx.is_disconnected());
if let Some(metrics) = metrics.as_ref() {
metrics.set_active_connections(connections.len());
}
debug!("UDS connection to {} ({:?}) closed", instance_id, path);
result
}
async fn connection_writer_inner(
path: &Path,
instance_id: crate::InstanceId,
rx: &flume::Receiver<SendTask>,
cancel_token: &CancellationToken,
connect_timeout: Duration,
) -> Result<()> {
debug!("Connecting to UDS {:?}", path);
let mut stream = tokio::select! {
_ = cancel_token.cancelled() => return Ok(()),
res = tokio::time::timeout(connect_timeout, UnixStream::connect(path)) => {
res.context("UDS connect timeout")?.context("UDS connect failed")?
},
};
let sock = socket2::SockRef::from(&stream);
if let Err(e) = sock.set_send_buffer_size(2_097_152) {
warn!("Failed to set UDS send buffer size: {}", e);
}
if let Err(e) = sock.set_recv_buffer_size(2_097_152) {
warn!("Failed to set UDS recv buffer size: {}", e);
}
debug!("Connected to UDS {:?}", path);
loop {
let msg = tokio::select! {
_ = cancel_token.cancelled() => break,
res = rx.recv_async() => match res {
Ok(msg) => msg,
Err(_) => break,
},
};
if let Err(e) =
TcpFrameCodec::encode_frame(&mut stream, msg.msg_type, &msg.header, &msg.payload).await
{
error!("Write error to {} ({:?}): {}", instance_id, path, e);
msg.on_error(format!("Failed to write to UDS stream: {}", e));
break;
}
}
Ok(())
}
fn parse_uds_endpoint(endpoint: &[u8]) -> Result<PathBuf> {
let endpoint_str = std::str::from_utf8(endpoint).context("endpoint is not valid UTF-8")?;
let path_str = endpoint_str.strip_prefix("uds://").unwrap_or(endpoint_str);
if path_str.is_empty() {
anyhow::bail!("empty UDS socket path");
}
Ok(PathBuf::from(path_str))
}
pub struct UdsTransportBuilder {
socket_path: Option<PathBuf>,
key: Option<TransportKey>,
channel_capacity: usize,
connect_timeout: Duration,
}
impl UdsTransportBuilder {
pub fn new() -> Self {
Self {
socket_path: None,
key: None,
channel_capacity: 256,
connect_timeout: Duration::from_secs(5),
}
}
pub fn socket_path(mut self, path: impl Into<PathBuf>) -> Self {
self.socket_path = Some(path.into());
self
}
pub fn key(mut self, key: TransportKey) -> Self {
self.key = Some(key);
self
}
pub fn channel_capacity(mut self, capacity: usize) -> Self {
self.channel_capacity = capacity;
self
}
pub fn connect_timeout(mut self, timeout: Duration) -> Self {
self.connect_timeout = timeout;
self
}
pub fn build(self) -> Result<UdsTransport> {
let socket_path = self
.socket_path
.ok_or_else(|| anyhow::anyhow!("socket_path is required"))?;
let key = self.key.unwrap_or_else(|| TransportKey::from("uds"));
let local_endpoint = format!("uds://{}", socket_path.display());
let mut addr_builder = crate::transports::address::WorkerAddressBuilder::new();
addr_builder.add_entry(key.clone(), local_endpoint.as_bytes().to_vec())?;
let local_address = addr_builder.build()?;
Ok(UdsTransport::new(
socket_path,
key,
local_address,
self.channel_capacity,
self.connect_timeout,
))
}
}
impl Default for UdsTransportBuilder {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::transports::address::WorkerAddressBuilder;
use std::sync::atomic::{AtomicUsize, Ordering};
use velo_ext::PeerInfo;
struct NullErrorHandler;
impl TransportErrorHandler for NullErrorHandler {
fn on_error(&self, _: Bytes, _: Bytes, _: String) {}
}
struct TrackingErrorHandler {
count: AtomicUsize,
}
impl TrackingErrorHandler {
fn new() -> Self {
Self {
count: AtomicUsize::new(0),
}
}
fn error_count(&self) -> usize {
self.count.load(Ordering::SeqCst)
}
}
impl TransportErrorHandler for TrackingErrorHandler {
fn on_error(&self, _: Bytes, _: Bytes, _: String) {
self.count.fetch_add(1, Ordering::SeqCst);
}
}
fn make_uds_peer(path: &Path) -> PeerInfo {
let instance_id = crate::InstanceId::new_v4();
let mut builder = WorkerAddressBuilder::new();
builder
.add_entry("uds", format!("uds://{}", path.display()).into_bytes())
.unwrap();
PeerInfo::new(instance_id, builder.build().unwrap())
}
fn make_transport() -> (UdsTransport, PathBuf) {
let dir = std::env::temp_dir().join(format!("uds-test-{}", crate::InstanceId::new_v4()));
std::fs::create_dir_all(&dir).unwrap();
let socket_path = dir.join("test.sock");
let transport = UdsTransportBuilder::new()
.socket_path(&socket_path)
.build()
.unwrap();
transport
.runtime
.set(tokio::runtime::Handle::current())
.ok();
(transport, socket_path)
}
fn insert_stale_handle(transport: &UdsTransport, instance_id: crate::InstanceId) {
let (tx, _rx) = flume::bounded::<SendTask>(1);
transport
.connections
.insert(instance_id, ConnectionHandle { tx });
}
#[test]
fn test_parse_uds_endpoint() {
let path = parse_uds_endpoint(b"uds:///tmp/test.sock").unwrap();
assert_eq!(path, PathBuf::from("/tmp/test.sock"));
let path = parse_uds_endpoint(b"/var/run/anvil.sock").unwrap();
assert_eq!(path, PathBuf::from("/var/run/anvil.sock"));
assert!(parse_uds_endpoint(b"").is_err());
}
#[test]
fn test_builder_requires_socket_path() {
let result = UdsTransportBuilder::new().build();
assert!(result.is_err());
}
#[test]
fn test_builder_with_socket_path() {
let result = UdsTransportBuilder::new()
.socket_path("/tmp/test.sock")
.build();
assert!(result.is_ok());
}
#[test]
fn test_builder_custom_key() {
let transport = UdsTransportBuilder::new()
.socket_path("/tmp/test.sock")
.key(TransportKey::from("custom-uds"))
.build()
.unwrap();
assert_eq!(transport.key(), TransportKey::from("custom-uds"));
}
#[test]
fn test_transport_socket_path() {
let transport = UdsTransportBuilder::new()
.socket_path("/tmp/test.sock")
.build()
.unwrap();
assert_eq!(transport.socket_path(), Path::new("/tmp/test.sock"));
}
#[tokio::test]
async fn test_get_or_create_connection_replaces_stale_handle() {
let (transport, _socket_path) = make_transport();
let dir = std::env::temp_dir().join(format!("uds-peer-{}", crate::InstanceId::new_v4()));
std::fs::create_dir_all(&dir).unwrap();
let peer_socket = dir.join("peer.sock");
let peer_listener = tokio::net::UnixListener::bind(&peer_socket).unwrap();
let peer = make_uds_peer(&peer_socket);
let iid = peer.instance_id();
transport.register(peer).unwrap();
insert_stale_handle(&transport, iid);
assert!(
transport
.connections
.get(&iid)
.unwrap()
.tx
.is_disconnected()
);
let handle = transport.get_or_create_connection(iid).unwrap();
assert!(!handle.tx.is_disconnected());
let entry = transport.connections.get(&iid).unwrap();
assert!(!entry.tx.is_disconnected());
drop(peer_listener);
std::fs::remove_file(&peer_socket).ok();
std::fs::remove_dir_all(&dir).ok();
}
#[tokio::test]
async fn test_check_health_removes_stale_entry() {
let (transport, _socket_path) = make_transport();
let dir = std::env::temp_dir().join(format!("uds-peer-{}", crate::InstanceId::new_v4()));
std::fs::create_dir_all(&dir).unwrap();
let peer_socket = dir.join("peer.sock");
let _peer_listener = tokio::net::UnixListener::bind(&peer_socket).unwrap();
let peer = make_uds_peer(&peer_socket);
let iid = peer.instance_id();
transport.register(peer).unwrap();
insert_stale_handle(&transport, iid);
assert!(transport.connections.contains_key(&iid));
let result = transport.check_health(iid, Duration::from_secs(2)).await;
assert!(!transport.connections.contains_key(&iid));
assert!(result.is_ok());
std::fs::remove_file(&peer_socket).ok();
std::fs::remove_dir_all(&dir).ok();
}
#[tokio::test]
async fn test_writer_task_cleans_up_on_write_error() {
let dir = std::env::temp_dir().join(format!("uds-test-{}", crate::InstanceId::new_v4()));
std::fs::create_dir_all(&dir).unwrap();
let socket_path = dir.join("writer-test.sock");
let listener = tokio::net::UnixListener::bind(&socket_path).unwrap();
let iid = crate::InstanceId::new_v4();
let (tx, rx) = flume::bounded::<SendTask>(8);
let connections: Arc<DashMap<crate::InstanceId, ConnectionHandle>> =
Arc::new(DashMap::new());
connections.insert(iid, ConnectionHandle { tx: tx.clone() });
let conns = Arc::clone(&connections);
let cancel = CancellationToken::new();
let writer = tokio::spawn(connection_writer_task(
socket_path.clone(),
iid,
rx,
conns,
cancel,
Duration::from_secs(5),
None,
));
let (stream, _) = listener.accept().await.unwrap();
drop(stream);
drop(listener);
tx.send(SendTask {
msg_type: MessageType::Message,
header: Bytes::from_static(b"hdr"),
payload: Bytes::from_static(b"pay"),
on_error: Arc::new(NullErrorHandler),
})
.unwrap();
let _ = writer.await;
assert!(
!connections.contains_key(&iid),
"writer task should clean up its DashMap entry on write error"
);
std::fs::remove_file(&socket_path).ok();
std::fs::remove_dir_all(&dir).ok();
}
#[tokio::test]
async fn test_send_message_does_not_fail_on_stale_handle() {
let (transport, _socket_path) = make_transport();
let dir = std::env::temp_dir().join(format!("uds-peer-{}", crate::InstanceId::new_v4()));
std::fs::create_dir_all(&dir).unwrap();
let peer_socket = dir.join("peer.sock");
let peer_listener = tokio::net::UnixListener::bind(&peer_socket).unwrap();
let peer = make_uds_peer(&peer_socket);
let iid = peer.instance_id();
transport.register(peer).unwrap();
insert_stale_handle(&transport, iid);
let error_handler = Arc::new(TrackingErrorHandler::new());
transport
.send_message(
iid,
Bytes::from_static(b"test-header"),
Bytes::from_static(b"test-payload"),
MessageType::Message,
error_handler.clone(),
)
.expect("slow-path send on fresh connection should enqueue synchronously");
let (mut stream, _) = peer_listener.accept().await.unwrap();
use tokio::io::AsyncReadExt;
let mut buf = [0u8; 256];
let n = tokio::time::timeout(Duration::from_secs(2), stream.read(&mut buf))
.await
.expect("timed out waiting for data")
.expect("read error");
assert!(n > 0, "expected data from the writer task");
assert_eq!(
error_handler.error_count(),
0,
"send_message should retry on stale handle, not fail"
);
let entry = transport.connections.get(&iid).unwrap();
assert!(
!entry.tx.is_disconnected(),
"stale handle should have been replaced with a live one"
);
std::fs::remove_file(&peer_socket).ok();
std::fs::remove_dir_all(&dir).ok();
}
#[tokio::test]
async fn test_double_bind_returns_err() {
use crate::transports::transport::make_channels;
let dir = std::env::temp_dir().join(format!("uds-test-{}", crate::InstanceId::new_v4()));
std::fs::create_dir_all(&dir).unwrap();
let socket_path = dir.join("double-bind.sock");
let transport1 = UdsTransportBuilder::new()
.socket_path(&socket_path)
.build()
.unwrap();
let instance_id = crate::InstanceId::new_v4();
let (adapter1, _streams1) = make_channels();
let rt = tokio::runtime::Handle::current();
transport1
.start(instance_id, adapter1, rt.clone())
.await
.unwrap();
let transport2 = UdsTransportBuilder::new()
.socket_path(&socket_path)
.build()
.unwrap();
let (adapter2, _streams2) = make_channels();
let result = transport2.start(instance_id, adapter2, rt).await;
assert!(
result.is_err(),
"start() should return Err when a live listener already owns the socket"
);
transport1.shutdown();
std::fs::remove_dir_all(&dir).ok();
}
#[tokio::test]
async fn test_begin_drain_activates_draining_flag() {
use crate::transports::transport::make_channels;
let dir = std::env::temp_dir().join(format!("uds-test-{}", crate::InstanceId::new_v4()));
std::fs::create_dir_all(&dir).unwrap();
let socket_path = dir.join("drain-test.sock");
let transport = UdsTransportBuilder::new()
.socket_path(&socket_path)
.build()
.unwrap();
let instance_id = crate::InstanceId::new_v4();
let (adapter, _streams) = make_channels();
let rt = tokio::runtime::Handle::current();
transport.start(instance_id, adapter, rt).await.unwrap();
assert!(
!transport.shutdown_state.get().unwrap().is_draining(),
"should not be draining before begin_drain()"
);
transport.begin_drain();
assert!(
transport.shutdown_state.get().unwrap().is_draining(),
"should be draining after begin_drain()"
);
transport.shutdown();
std::fs::remove_dir_all(&dir).ok();
}
#[tokio::test]
async fn test_writer_task_drains_on_connect_failure() {
let dir = std::env::temp_dir().join(format!("uds-test-{}", crate::InstanceId::new_v4()));
std::fs::create_dir_all(&dir).unwrap();
let dead_socket = dir.join("dead.sock");
let iid = crate::InstanceId::new_v4();
let (tx, rx) = flume::bounded::<SendTask>(8);
let connections: Arc<DashMap<crate::InstanceId, ConnectionHandle>> =
Arc::new(DashMap::new());
connections.insert(iid, ConnectionHandle { tx: tx.clone() });
let error_handler = Arc::new(TrackingErrorHandler::new());
tx.send(SendTask {
msg_type: MessageType::Message,
header: Bytes::from_static(b"hdr"),
payload: Bytes::from_static(b"pay"),
on_error: error_handler.clone(),
})
.unwrap();
let conns = Arc::clone(&connections);
let cancel = CancellationToken::new();
let writer = tokio::spawn(connection_writer_task(
dead_socket,
iid,
rx,
conns,
cancel,
Duration::from_secs(5),
None,
));
let _ = writer.await;
assert_eq!(
error_handler.error_count(),
1,
"queued message should have its on_error called when connect fails"
);
assert!(
!connections.contains_key(&iid),
"writer task should clean up its DashMap entry on connect failure"
);
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn test_register_rejects_missing_path() {
let dir = std::env::temp_dir().join(format!("uds-reject-{}", crate::InstanceId::new_v4()));
std::fs::create_dir_all(&dir).unwrap();
let socket_path = dir.join("self.sock");
let transport = UdsTransportBuilder::new()
.socket_path(&socket_path)
.build()
.unwrap();
let missing =
std::env::temp_dir().join(format!("uds-missing-{}.sock", crate::InstanceId::new_v4()));
assert!(!missing.exists());
let peer = make_uds_peer(&missing);
let peer_id = peer.instance_id();
let result = transport.register(peer);
assert!(matches!(result, Err(TransportError::NoEndpoint)));
assert!(!transport.peers.contains_key(&peer_id));
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn test_register_rejects_non_socket_file() {
let dir = std::env::temp_dir().join(format!("uds-nonsock-{}", crate::InstanceId::new_v4()));
std::fs::create_dir_all(&dir).unwrap();
let socket_path = dir.join("self.sock");
let transport = UdsTransportBuilder::new()
.socket_path(&socket_path)
.build()
.unwrap();
let regular_file = dir.join("not-a-socket");
std::fs::write(®ular_file, b"I am not a socket").unwrap();
let peer = make_uds_peer(®ular_file);
let peer_id = peer.instance_id();
let result = transport.register(peer);
assert!(matches!(result, Err(TransportError::NoEndpoint)));
assert!(!transport.peers.contains_key(&peer_id));
std::fs::remove_dir_all(&dir).ok();
}
#[tokio::test]
async fn test_register_accepts_bound_socket() {
let dir = std::env::temp_dir().join(format!("uds-accept-{}", crate::InstanceId::new_v4()));
std::fs::create_dir_all(&dir).unwrap();
let socket_path = dir.join("self.sock");
let transport = UdsTransportBuilder::new()
.socket_path(&socket_path)
.build()
.unwrap();
let peer_socket = dir.join("peer.sock");
let _peer_listener = tokio::net::UnixListener::bind(&peer_socket).unwrap();
let peer = make_uds_peer(&peer_socket);
let peer_id = peer.instance_id();
transport.register(peer).expect("register should succeed");
assert!(transport.peers.contains_key(&peer_id));
std::fs::remove_dir_all(&dir).ok();
}
}