#![cfg(windows)]
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::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
use tokio::net::windows::named_pipe::{
ClientOptions, NamedPipeClient, NamedPipeServer, ServerOptions,
};
use tokio::sync::Mutex;
use tracing::{debug, info};
#[derive(Clone)]
pub struct WindowsIpcTransport {
inner: Arc<WindowsIpcTransportInner>,
}
struct WindowsIpcTransportInner {
pipe_name: String,
capacity: usize,
server: Mutex<Option<NamedPipeServer>>,
client: Mutex<Option<NamedPipeClient>>,
connected: Mutex<bool>,
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(Some(server)),
client: Mutex::new(None),
connected: Mutex::new(false),
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(None),
client: Mutex::new(Some(client)),
connected: Mutex::new(true),
closed: AtomicBool::new(false),
}),
})
}
async fn send_frame<W>(pipe: &mut W, data: &[u8]) -> Result<()>
where
W: AsyncWrite + Unpin + Send,
{
let len = data.len() as u32;
pipe.write_all(&len.to_be_bytes()).await?;
pipe.write_all(data).await?;
pipe.flush().await?;
Ok(())
}
async fn recv_frame<R>(pipe: &mut R) -> Result<Vec<u8>>
where
R: AsyncRead + Unpin + Send,
{
let mut len_bytes = [0u8; 4];
pipe.read_exact(&mut len_bytes).await?;
let len = u32::from_be_bytes(len_bytes) as usize;
if len > 100 * 1024 * 1024 {
return Err(anyhow!("Frame too large: {} bytes", len));
}
let mut data = vec![0u8; len];
pipe.read_exact(&mut data).await?;
Ok(data)
}
pub async fn wait_for_connection(&self) -> Result<()> {
if *self.inner.connected.lock().await {
return Ok(());
}
let mut server_guard = self.inner.server.lock().await;
if let Some(server) = server_guard.as_mut() {
server.connect().await?;
*self.inner.connected.lock().await = true;
debug!(pipe = %self.inner.pipe_name, "Client connected to Named Pipe");
Ok(())
} else {
Err(anyhow!("Windows Named Pipe transport not in server mode"))
}
}
}
#[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 data = rmp_serde::to_vec(&messages)?;
let mut client_guard = self.inner.client.lock().await;
if let Some(client) = client_guard.as_mut() {
Self::send_frame(client, &data).await?;
debug!(
pipe = %self.inner.pipe_name,
count = messages.len(),
bytes = data.len(),
"Sent batch via Named Pipe"
);
Ok(())
} else {
Err(anyhow!("Windows Named Pipe transport not in client mode"))
}
}
async fn recv_batch(&self) -> Result<Vec<CanonicalMessage>> {
if self.inner.closed.load(Ordering::SeqCst) {
return Err(anyhow!("Windows Named Pipe transport is closed"));
}
self.wait_for_connection().await?;
let mut server_guard = self.inner.server.lock().await;
if let Some(server) = server_guard.as_mut() {
let data = Self::recv_frame(server).await?;
let messages: Vec<CanonicalMessage> = rmp_serde::from_slice(&data)?;
debug!(
pipe = %self.inner.pipe_name,
count = messages.len(),
bytes = data.len(),
"Received batch via Named Pipe"
);
Ok(messages)
} else {
Err(anyhow!("Windows Named Pipe transport not in server mode"))
}
}
fn try_recv_batch(&self) -> Result<Option<Vec<CanonicalMessage>>> {
Ok(None)
}
fn len(&self) -> usize {
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-{}", uuid::Uuid::new_v4());
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-{}", uuid::Uuid::new_v4());
let server = WindowsIpcTransport::new_server(&pipe_name, 10)
.await
.unwrap();
assert!(!server.is_closed());
server.close();
assert!(server.is_closed());
}
}