#![cfg(windows)]
use super::framed::{self, FramedIo};
use super::transport::TransportChannel;
use crate::CanonicalMessage;
use anyhow::{anyhow, Result};
use async_trait::async_trait;
use std::sync::{
atomic::{AtomicBool, Ordering},
Arc,
};
use tokio::net::windows::named_pipe::{
ClientOptions, NamedPipeClient, NamedPipeServer, ServerOptions,
};
use tokio::sync::Mutex;
use tracing::{debug, info, warn};
enum ServerState {
Idle(NamedPipeServer),
Connected(FramedIo<NamedPipeServer>),
Empty,
}
#[derive(Clone)]
pub struct WindowsIpcTransport {
inner: Arc<WindowsIpcTransportInner>,
}
struct WindowsIpcTransportInner {
pipe_name: String,
capacity: usize,
server: Mutex<ServerState>,
client: Mutex<Option<FramedIo<NamedPipeClient>>>,
closed: AtomicBool,
}
impl WindowsIpcTransport {
pub async fn new_server(pipe_name: impl AsRef<str>, capacity: usize) -> Result<Self> {
let pipe_name = pipe_name.as_ref();
let full_path = format!(r"\\.\pipe\{}", pipe_name);
let server = ServerOptions::new()
.first_pipe_instance(true)
.create(&full_path)?;
info!(pipe = %full_path, "Windows Named Pipe server created");
Ok(Self {
inner: Arc::new(WindowsIpcTransportInner {
pipe_name: full_path,
capacity,
server: Mutex::new(ServerState::Idle(server)),
client: Mutex::new(None),
closed: AtomicBool::new(false),
}),
})
}
pub async fn new_client(pipe_name: impl AsRef<str>, capacity: usize) -> Result<Self> {
let pipe_name = pipe_name.as_ref();
let full_path = format!(r"\\.\pipe\{}", pipe_name);
let client = ClientOptions::new().open(&full_path)?;
info!(pipe = %full_path, "Windows Named Pipe client connected");
Ok(Self {
inner: Arc::new(WindowsIpcTransportInner {
pipe_name: full_path,
capacity,
server: Mutex::new(ServerState::Empty),
client: Mutex::new(Some(framed::wrap(client))),
closed: AtomicBool::new(false),
}),
})
}
pub async fn wait_for_connection(&self) -> Result<()> {
let mut guard = self.inner.server.lock().await;
match std::mem::replace(&mut *guard, ServerState::Empty) {
ServerState::Connected(conn) => {
*guard = ServerState::Connected(conn);
Ok(())
}
ServerState::Idle(server) => {
match server.connect().await {
Ok(()) => {
debug!(pipe = %self.inner.pipe_name, "Client connected to Named Pipe");
*guard = ServerState::Connected(framed::wrap(server));
Ok(())
}
Err(e) => {
*guard = ServerState::Idle(server);
Err(e.into())
}
}
}
ServerState::Empty => Err(anyhow!(
"Windows Named Pipe transport not in server mode: {}",
self.inner.pipe_name
)),
}
}
fn is_disconnected(error: &anyhow::Error) -> bool {
error
.downcast_ref::<std::io::Error>()
.is_some_and(|io_error| {
matches!(
io_error.kind(),
std::io::ErrorKind::UnexpectedEof
| std::io::ErrorKind::BrokenPipe
| std::io::ErrorKind::ConnectionReset
| std::io::ErrorKind::ConnectionAborted
| std::io::ErrorKind::NotConnected
)
})
}
}
#[async_trait]
impl TransportChannel for WindowsIpcTransport {
async fn send_batch(&self, messages: Vec<CanonicalMessage>) -> Result<()> {
if self.inner.closed.load(Ordering::SeqCst) {
return Err(anyhow!("Windows Named Pipe transport is closed"));
}
let mut client_guard = self.inner.client.lock().await;
if let Some(conn) = client_guard.as_mut() {
let bytes = framed::send_batch(conn, &messages, &self.inner.pipe_name).await?;
debug!(
pipe = %self.inner.pipe_name,
count = messages.len(),
bytes,
"Sent batch via Named Pipe"
);
Ok(())
} else {
Err(anyhow!(
"Windows Named Pipe transport at '{}' is the consumer (server) side and cannot \
send; the pipe carries publisher -> consumer traffic only",
self.inner.pipe_name
))
}
}
async fn recv_batch(&self) -> Result<Vec<CanonicalMessage>> {
loop {
if self.inner.closed.load(Ordering::SeqCst) {
return Err(anyhow!("Windows Named Pipe transport is closed"));
}
self.wait_for_connection().await?;
let read_result = {
let mut guard = self.inner.server.lock().await;
let ServerState::Connected(conn) = &mut *guard else {
return Err(anyhow!(
"Windows Named Pipe transport not in server mode: {}",
self.inner.pipe_name
));
};
framed::recv_batch(conn).await
};
match read_result {
Ok(messages) => {
debug!(
pipe = %self.inner.pipe_name,
count = messages.len(),
"Received batch via Named Pipe"
);
return Ok(messages);
}
Err(error) if Self::is_disconnected(&error) => {
warn!(pipe = %self.inner.pipe_name, error = %error, "Named Pipe peer disconnected; waiting for a new connection");
let mut guard = self.inner.server.lock().await;
match ServerOptions::new().create(&self.inner.pipe_name) {
Ok(server) => *guard = ServerState::Idle(server),
Err(e) => {
*guard = ServerState::Empty;
return Err(anyhow!(
"Failed to recreate Named Pipe instance at '{}': {}",
self.inner.pipe_name,
e
));
}
}
}
Err(error) => return Err(error),
}
}
}
fn try_recv_batch(&self) -> Result<Option<Vec<CanonicalMessage>>> {
let Ok(mut guard) = self.inner.server.try_lock() else {
return Ok(None);
};
let ServerState::Connected(conn) = &mut *guard else {
return Ok(None);
};
framed::try_recv_batch(conn)
}
fn len(&self) -> usize {
match self.inner.server.try_lock() {
Ok(guard) => match &*guard {
ServerState::Connected(conn) => framed::buffered_frames(conn),
_ => 0,
},
Err(_) => 0,
}
}
fn capacity(&self) -> Option<usize> {
Some(self.inner.capacity)
}
fn is_closed(&self) -> bool {
self.inner.closed.load(Ordering::SeqCst)
}
fn close(&self) {
if !self.inner.closed.swap(true, Ordering::SeqCst) {
info!(pipe = %self.inner.pipe_name, "Closing Windows Named Pipe transport");
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_windows_pipe_roundtrip() {
let pipe_name = format!("mq-bridge-test-{}", fast_uuid_v7::gen_id_str());
let server = WindowsIpcTransport::new_server(&pipe_name, 10)
.await
.unwrap();
let pipe_name_clone = pipe_name.clone();
let client_task = tokio::spawn(async move {
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
WindowsIpcTransport::new_client(&pipe_name_clone, 10)
.await
.unwrap()
});
server.wait_for_connection().await.unwrap();
let client = client_task.await.unwrap();
let msg = CanonicalMessage::from_vec(b"test");
client.send_batch(vec![msg.clone()]).await.unwrap();
let received = server.recv_batch().await.unwrap();
assert_eq!(received.len(), 1);
}
#[tokio::test]
async fn test_windows_pipe_close() {
let pipe_name = format!("mq-bridge-test-{}", fast_uuid_v7::gen_id_str());
let server = WindowsIpcTransport::new_server(&pipe_name, 10)
.await
.unwrap();
assert!(!server.is_closed());
server.close();
assert!(server.is_closed());
}
#[tokio::test]
async fn test_windows_pipe_server_cannot_send() {
let pipe_name = format!("mq-bridge-test-{}", fast_uuid_v7::gen_id_str());
let server = WindowsIpcTransport::new_server(&pipe_name, 10)
.await
.unwrap();
let err = server
.send_batch(vec![CanonicalMessage::from_vec(b"nope")])
.await
.unwrap_err();
assert!(
err.to_string().contains("cannot send"),
"unexpected error: {err}"
);
}
#[tokio::test]
async fn test_windows_pipe_cancelled_receive_loses_nothing() {
let pipe_name = format!("mq-bridge-test-{}", fast_uuid_v7::gen_id_str());
let server = WindowsIpcTransport::new_server(&pipe_name, 10)
.await
.unwrap();
let client = WindowsIpcTransport::new_client(&pipe_name, 10)
.await
.unwrap();
server.wait_for_connection().await.unwrap();
tokio::select! {
_ = server.recv_batch() => panic!("nothing has been sent yet"),
_ = tokio::time::sleep(std::time::Duration::from_millis(50)) => {}
}
client
.send_batch(vec![CanonicalMessage::from_vec(b"after-cancel")])
.await
.unwrap();
let received = server.recv_batch().await.unwrap();
assert_eq!(received[0].payload.as_ref(), b"after-cancel");
}
}