use std::io;
use std::path::Path;
#[cfg(windows)]
const PIPE_BUSY_RETRY: std::time::Duration = std::time::Duration::from_millis(20);
#[cfg(windows)]
const ERROR_PIPE_BUSY: i32 = 231;
#[cfg(unix)]
pub type ClientStream = tokio::net::UnixStream;
#[cfg(windows)]
pub type ClientStream = tokio::net::windows::named_pipe::NamedPipeClient;
#[cfg(unix)]
pub type ServerStream = tokio::net::UnixStream;
#[cfg(windows)]
pub type ServerStream = tokio::net::windows::named_pipe::NamedPipeServer;
pub type ServerReadHalf = tokio::io::ReadHalf<ServerStream>;
pub type ServerWriteHalf = tokio::io::WriteHalf<ServerStream>;
#[must_use]
pub fn split(stream: ServerStream) -> (ServerReadHalf, ServerWriteHalf) {
tokio::io::split(stream)
}
pub async fn connected_pair() -> io::Result<(ServerStream, ClientStream)> {
#[cfg(unix)]
{
tokio::net::UnixStream::pair()
}
#[cfg(windows)]
{
use core::sync::atomic::{AtomicU64, Ordering};
static NEXT: AtomicU64 = AtomicU64::new(0);
let name = std::path::PathBuf::from(format!(
r"\\.\pipe\shep-pair-{}-{}",
std::process::id(),
NEXT.fetch_add(1, Ordering::Relaxed)
));
let mut listener = Listener::bind(&name)?;
tokio::try_join!(listener.accept(), connect(&name))
}
}
pub async fn connect(addr: &Path) -> io::Result<ClientStream> {
#[cfg(unix)]
{
tokio::net::UnixStream::connect(addr).await
}
#[cfg(windows)]
{
use tokio::net::windows::named_pipe::ClientOptions;
loop {
match ClientOptions::new().open(addr) {
Ok(client) => return Ok(client),
Err(err) if err.raw_os_error() == Some(ERROR_PIPE_BUSY) => {
tokio::time::sleep(PIPE_BUSY_RETRY).await;
}
Err(err) => return Err(err),
}
}
}
}
#[derive(Debug)]
pub struct Listener {
#[cfg(unix)]
listener: tokio::net::UnixListener,
#[cfg(windows)]
server: Option<tokio::net::windows::named_pipe::NamedPipeServer>,
#[cfg(windows)]
addr: std::ffi::OsString,
}
impl Listener {
pub fn bind(addr: &Path) -> io::Result<Self> {
#[cfg(unix)]
{
Ok(Self {
listener: tokio::net::UnixListener::bind(addr)?,
})
}
#[cfg(windows)]
{
use tokio::net::windows::named_pipe::ServerOptions;
let server = ServerOptions::new()
.first_pipe_instance(true)
.reject_remote_clients(true)
.create(addr)?;
Ok(Self {
server: Some(server),
addr: addr.as_os_str().to_os_string(),
})
}
}
pub async fn accept(&mut self) -> io::Result<ServerStream> {
#[cfg(unix)]
{
let (stream, _addr) = self.listener.accept().await?;
Ok(stream)
}
#[cfg(windows)]
{
use tokio::net::windows::named_pipe::ServerOptions;
if self.server.is_none() {
self.server = Some(
ServerOptions::new()
.reject_remote_clients(true)
.create(&self.addr)?,
);
}
let Some(server) = self.server.as_ref() else {
unreachable!("the slot was just filled")
};
server.connect().await?;
let connected = self
.server
.take()
.unwrap_or_else(|| unreachable!("the slot was just filled"));
self.server = ServerOptions::new()
.reject_remote_clients(true)
.create(&self.addr)
.ok();
Ok(connected)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
fn address(dir: &Path, tag: &str) -> std::path::PathBuf {
#[cfg(unix)]
{
let _ = tag;
dir.join("shep.sock")
}
#[cfg(windows)]
{
let _ = dir;
std::path::PathBuf::from(format!(
r"\\.\pipe\shep-transport-test-{tag}-{}",
std::process::id()
))
}
}
#[tokio::test]
async fn a_client_and_the_daemon_exchange_bytes_over_the_platform_transport() {
let dir = tempfile::tempdir().unwrap();
let addr = address(dir.path(), "roundtrip");
let mut listener = Listener::bind(&addr).unwrap();
let server = tokio::spawn(async move {
let mut stream = listener.accept().await.unwrap();
let mut buf = [0u8; 5];
stream.read_exact(&mut buf).await.unwrap();
stream.write_all(b"world").await.unwrap();
stream.flush().await.unwrap();
buf
});
let mut client = connect(&addr).await.unwrap();
client.write_all(b"hello").await.unwrap();
client.flush().await.unwrap();
let mut reply = [0u8; 5];
client.read_exact(&mut reply).await.unwrap();
assert_eq!(&server.await.unwrap(), b"hello");
assert_eq!(&reply, b"world");
}
#[tokio::test]
async fn the_listener_serves_more_than_one_connection() {
let dir = tempfile::tempdir().unwrap();
let addr = address(dir.path(), "sequential");
let mut listener = Listener::bind(&addr).unwrap();
let server = tokio::spawn(async move {
let mut seen = Vec::new();
for _ in 0..3 {
let mut stream = listener.accept().await.unwrap();
let mut byte = [0u8; 1];
stream.read_exact(&mut byte).await.unwrap();
seen.push(byte[0]);
}
seen
});
for tag in [1u8, 2, 3] {
let mut client = connect(&addr).await.unwrap();
client.write_all(&[tag]).await.unwrap();
client.flush().await.unwrap();
}
assert_eq!(server.await.unwrap(), vec![1, 2, 3]);
}
#[tokio::test]
async fn dialing_an_address_with_no_listener_fails_rather_than_hanging() {
let dir = tempfile::tempdir().unwrap();
let addr = address(dir.path(), "absent");
let result = tokio::time::timeout(std::time::Duration::from_secs(5), connect(&addr))
.await
.expect("connect must fail fast, not hang, when nothing is listening");
assert!(result.is_err(), "no listener must not read as a connection");
}
#[tokio::test]
async fn a_second_bind_on_the_same_address_is_refused() {
let dir = tempfile::tempdir().unwrap();
let addr = address(dir.path(), "exclusive");
let _first = Listener::bind(&addr).unwrap();
assert!(
Listener::bind(&addr).is_err(),
"a second daemon must not be able to bind the same control address"
);
}
}