use super::*;
use crate::transports::AdmissionState;
use crate::transports::address::WorkerAddressBuilder;
use crate::transports::tcp::TcpFrameCodec;
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 make_handle(capacity: usize) -> (ConnectionHandle, flume::Receiver<SendTask>) {
let (tx, rx) = flume::bounded::<SendTask>(capacity);
let handle = ConnectionHandle {
gate: AdmissionGate::new(tx.clone(), tokio::runtime::Handle::current()),
tx,
};
(handle, rx)
}
fn insert_stale_handle(transport: &UdsTransport, instance_id: crate::InstanceId) {
let (handle, _rx) = make_handle(1);
transport.connections.insert(instance_id, handle);
}
fn task(on_error: Arc<dyn TransportErrorHandler>) -> SendTask {
SendTask {
msg_type: MessageType::Message,
header: Bytes::from_static(b"hdr"),
payload: Bytes::from_static(b"pay"),
on_error,
}
}
#[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 (handle, rx) = make_handle(8);
let tx = handle.tx.clone();
let connections: Arc<DashMap<crate::InstanceId, ConnectionHandle>> = Arc::new(DashMap::new());
connections.insert(iid, handle);
let conns = Arc::clone(&connections);
let cancel = CancellationToken::new();
let writer = tokio::spawn(connection_writer_task(
socket_path.clone(),
iid,
rx,
WriterTaskContext {
connections: conns,
cancel_token: cancel,
connect_timeout: Duration::from_secs(5),
reader_ctx: None,
metrics: 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());
assert!(
transport
.send_message(
iid,
Bytes::from_static(b"test-header"),
Bytes::from_static(b"test-payload"),
MessageType::Message,
error_handler.clone(),
)
.is_admitted(),
"a fresh connection's channel is empty, so the send admits immediately"
);
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_hook_does_not_flip_shared_state() {
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(),
"begin_drain hook must not flip the shared state"
);
_streams.shutdown_state.begin_drain();
assert!(
transport.shutdown_state.get().unwrap().is_draining(),
"transport observes the runtime's flip through the shared state"
);
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 (handle, rx) = make_handle(8);
let tx = handle.tx.clone();
let connections: Arc<DashMap<crate::InstanceId, ConnectionHandle>> = Arc::new(DashMap::new());
connections.insert(iid, handle);
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,
WriterTaskContext {
connections: conns,
cancel_token: cancel,
connect_timeout: Duration::from_secs(5),
reader_ctx: None,
metrics: 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();
}
#[tokio::test]
async fn stale_replacement_fails_the_old_epoch_and_admits_on_the_successor() {
let (transport, _socket_path) = make_transport();
let peer_dir = std::env::temp_dir().join(format!("uds-test-{}", crate::InstanceId::new_v4()));
std::fs::create_dir_all(&peer_dir).unwrap();
let peer_path = peer_dir.join("peer.sock");
let _peer_listener = tokio::net::UnixListener::bind(&peer_path).unwrap();
let peer = make_uds_peer(&peer_path);
let iid = peer.instance_id();
transport.register(peer).unwrap();
let (handle, rx) = make_handle(1);
transport.connections.insert(iid, handle.clone());
let errors = Arc::new(TrackingErrorHandler::new());
assert!(handle.gate.send(task(errors.clone())).is_admitted());
let queued = match handle.gate.send(task(errors.clone())) {
SendOutcome::Pending(admission) => admission,
SendOutcome::Admitted => panic!("a full channel must not admit"),
};
drop(rx);
assert!(handle.tx.is_disconnected());
let fresh = transport.get_or_create_connection(iid).unwrap();
assert_eq!(
queued.state(),
AdmissionState::Failed,
"the old epoch's queued frame must not survive the replacement"
);
assert!(!fresh.tx.is_disconnected(), "the successor should be live");
assert!(
fresh.gate.send(task(errors)).is_admitted(),
"the successor's gate is unaffected by the dead epoch"
);
}
#[tokio::test]
async fn max_message_size_is_exactly_what_the_codec_will_encode() {
let (transport, _socket_path) = make_transport();
let capacity = transport
.max_message_size(crate::InstanceId::new_v4())
.expect("UDS always knows its framed limit");
assert_eq!(capacity, 16 * 1024 * 1024);
let header_len = 1024u32;
let payload_len = capacity as u32 - header_len;
assert!(
TcpFrameCodec::build_preamble(MessageType::Message, header_len, payload_len).is_ok(),
"a frame of exactly the reported capacity must encode",
);
assert!(
TcpFrameCodec::build_preamble(MessageType::Message, header_len, payload_len + 1).is_err(),
"one byte past the reported capacity must not",
);
}