use std::ffi::OsStr;
use std::io;
use std::os::windows::ffi::OsStrExt;
use std::time::{Duration, Instant};
use windows::Win32::Foundation::{CloseHandle, HANDLE, WAIT_OBJECT_0, WAIT_TIMEOUT};
use windows::Win32::Storage::FileSystem::{
CreateFileW, FILE_ATTRIBUTE_NORMAL, FILE_FLAG_FIRST_PIPE_INSTANCE, FILE_FLAG_OVERLAPPED,
FILE_GENERIC_READ, FILE_GENERIC_WRITE, FILE_SHARE_READ, FILE_SHARE_WRITE, OPEN_EXISTING,
PIPE_ACCESS_DUPLEX, ReadFile, WriteFile,
};
use windows::Win32::System::IO::{CancelIo, GetOverlappedResult, OVERLAPPED};
use windows::Win32::System::Pipes::{
ConnectNamedPipe, CreateNamedPipeW, DisconnectNamedPipe, PIPE_READMODE_BYTE, PIPE_TYPE_BYTE,
PIPE_UNLIMITED_INSTANCES, PIPE_WAIT,
};
use windows::Win32::System::Threading::{CreateEventW, ResetEvent, SetEvent, WaitForSingleObject};
use windows::core::HRESULT;
use super::message::{self, SocketMessage, SocketResponse};
const BUF_SIZE: u32 = 8192;
const IPC_TIMEOUT: Duration = Duration::from_secs(30);
const ERROR_IO_PENDING_HRESULT: HRESULT = HRESULT(0x8007_03E5_u32 as i32);
#[derive(Debug)]
struct PipeHandle(HANDLE);
impl PipeHandle {
fn raw(&self) -> HANDLE {
self.0
}
}
impl Drop for PipeHandle {
fn drop(&mut self) {
unsafe {
let _ = CloseHandle(self.0);
}
}
}
unsafe impl Send for PipeHandle {}
#[derive(Debug)]
struct EventHandle(HANDLE);
impl EventHandle {
fn new() -> io::Result<Self> {
let handle = unsafe { CreateEventW(None, true, false, windows::core::PCWSTR::null()) }
.map_err(|e| io::Error::other(format!("CreateEventW failed: {e}")))?;
if handle.is_invalid() {
return Err(io::Error::other("CreateEventW returned invalid handle"));
}
Ok(Self(handle))
}
fn raw(&self) -> HANDLE {
self.0
}
}
impl Drop for EventHandle {
fn drop(&mut self) {
unsafe {
let _ = CloseHandle(self.0);
}
}
}
unsafe impl Send for EventHandle {}
pub struct PipeServer {
handle: PipeHandle,
connected_event: EventHandle,
accept_thread: Option<std::thread::JoinHandle<()>>,
}
impl PipeServer {
pub fn create() -> io::Result<Self> {
let name = wide(&message::pipe_name());
let handle = unsafe {
CreateNamedPipeW(
windows::core::PCWSTR(name.as_ptr()),
PIPE_ACCESS_DUPLEX | FILE_FLAG_FIRST_PIPE_INSTANCE,
PIPE_TYPE_BYTE | PIPE_READMODE_BYTE | PIPE_WAIT,
PIPE_UNLIMITED_INSTANCES,
BUF_SIZE,
BUF_SIZE,
0,
None,
)
};
if handle.is_invalid() {
return Err(io::Error::new(
io::ErrorKind::AddrInUse,
"pipe already in use (is another daemon running?)",
));
}
let connected_event = EventHandle::new()?;
Ok(Self {
handle: PipeHandle(handle),
connected_event,
accept_thread: None,
})
}
pub fn start_accept(&mut self) {
if let Some(handle) = self.accept_thread.take()
&& let Err(e) = handle.join()
{
log::warn!("PipeServer: previous accept thread panicked: {e:?}");
}
let _ = unsafe { ResetEvent(self.connected_event.raw()) };
let handle = self.handle.raw().0 as isize;
let event = self.connected_event.raw().0 as isize;
self.accept_thread = Some(std::thread::spawn(move || {
let pipe = HANDLE(handle as *mut core::ffi::c_void);
let event = HANDLE(event as *mut core::ffi::c_void);
const ERROR_PIPE_CONNECTED_HRESULT: windows::core::HRESULT =
windows::core::HRESULT(0x8007_0217_u32 as i32);
loop {
let result = unsafe { ConnectNamedPipe(pipe, None) };
match result {
Ok(()) => break,
Err(e) if e.code() == ERROR_PIPE_CONNECTED_HRESULT => break,
Err(e) => {
log::warn!(
"PipeServer accept thread: ConnectNamedPipe failed ({e:#?}); \
resetting pipe instance and retrying"
);
let _ = unsafe { DisconnectNamedPipe(pipe) };
std::thread::sleep(std::time::Duration::from_millis(20));
}
}
}
let _ = unsafe { SetEvent(event) };
}));
}
pub fn connected_event_handle(&self) -> HANDLE {
self.connected_event.raw()
}
pub fn read_message(&self) -> io::Result<SocketMessage> {
let line = read_line(self.handle.raw())?;
message::decode_message(&line).ok_or_else(|| {
io::Error::new(
io::ErrorKind::InvalidData,
format!("failed to parse message: {line:?}"),
)
})
}
pub fn write_response(&self, response: &SocketResponse) -> io::Result<()> {
let wire = message::encode_message(response)?;
write_all(self.handle.raw(), wire.as_bytes())
}
pub fn disconnect(&self) -> io::Result<()> {
unsafe { DisconnectNamedPipe(self.handle.raw()) }
.map_err(|e| io::Error::other(format!("DisconnectNamedPipe: {e}")))
}
}
fn read_line(handle: HANDLE) -> io::Result<String> {
let mut buf = vec![0u8; BUF_SIZE as usize];
let mut raw: Vec<u8> = Vec::new();
loop {
let mut bytes_read = 0u32;
unsafe {
ReadFile(
handle,
Some(buf.as_mut_slice()),
Some(&mut bytes_read),
None,
)
}
.map_err(|e| io::Error::new(io::ErrorKind::BrokenPipe, format!("ReadFile: {e}")))?;
if bytes_read == 0 {
return Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
"peer disconnected",
));
}
raw.extend_from_slice(&buf[..bytes_read as usize]);
if raw.contains(&b'\n') {
break;
}
}
String::from_utf8(raw).map_err(|e| {
io::Error::new(
io::ErrorKind::InvalidData,
format!("non-UTF-8 data from pipe: {e}"),
)
})
}
fn write_all(handle: HANDLE, data: &[u8]) -> io::Result<()> {
let mut total_written = 0;
while total_written < data.len() {
let mut bytes_written = 0u32;
unsafe {
WriteFile(
handle,
Some(&data[total_written..]),
Some(&mut bytes_written),
None,
)
}
.map_err(|e| io::Error::new(io::ErrorKind::BrokenPipe, format!("WriteFile: {e}")))?;
if bytes_written == 0 {
return Err(io::Error::new(
io::ErrorKind::BrokenPipe,
"WriteFile wrote zero bytes",
));
}
total_written += bytes_written as usize;
}
Ok(())
}
fn ms_until(deadline: Instant) -> Option<u32> {
let dur = deadline.checked_duration_since(Instant::now())?;
let ms = dur.as_millis().min(u128::from(u32::MAX)) as u32;
Some(ms.max(1)) }
unsafe fn await_overlapped(
handle: HANDLE,
event: HANDLE,
overlapped: &OVERLAPPED,
deadline: Instant,
) -> io::Result<u32> {
let ms = ms_until(deadline).ok_or_else(|| {
io::Error::new(
io::ErrorKind::TimedOut,
"IPC operation timed out (deadline passed)",
)
})?;
match unsafe { WaitForSingleObject(event, ms) } {
WAIT_OBJECT_0 => {
let mut transferred = 0u32;
unsafe { GetOverlappedResult(handle, overlapped, &mut transferred, false) }.map_err(
|e| {
io::Error::new(
io::ErrorKind::BrokenPipe,
format!("GetOverlappedResult: {e}"),
)
},
)?;
Ok(transferred)
}
WAIT_TIMEOUT => {
let _ = unsafe { CancelIo(handle) };
Err(io::Error::new(
io::ErrorKind::TimedOut,
"IPC operation timed out waiting for the daemon",
))
}
other => {
let _ = unsafe { CancelIo(handle) };
Err(io::Error::new(
io::ErrorKind::BrokenPipe,
format!("WaitForSingleObject returned {other:?}"),
))
}
}
}
fn write_all_overlapped(handle: HANDLE, data: &[u8], deadline: Instant) -> io::Result<()> {
let event = EventHandle::new()?;
let mut total = 0usize;
while total < data.len() {
let _ = unsafe { ResetEvent(event.raw()) };
let mut overlapped = OVERLAPPED {
hEvent: event.raw(),
..Default::default()
};
let mut written = 0u32;
let result = unsafe {
WriteFile(
handle,
Some(&data[total..]),
Some(&mut written),
Some(&mut overlapped),
)
};
match result {
Ok(()) => {}
Err(e) if e.code() == ERROR_IO_PENDING_HRESULT => {
written = unsafe { await_overlapped(handle, event.raw(), &overlapped, deadline)? };
}
Err(e) => {
return Err(io::Error::new(
io::ErrorKind::BrokenPipe,
format!("WriteFile: {e}"),
));
}
}
if written == 0 {
return Err(io::Error::new(
io::ErrorKind::BrokenPipe,
"WriteFile wrote zero bytes",
));
}
total += written as usize;
}
Ok(())
}
fn read_line_overlapped(handle: HANDLE, deadline: Instant) -> io::Result<String> {
let event = EventHandle::new()?;
let mut buf = vec![0u8; BUF_SIZE as usize];
let mut raw: Vec<u8> = Vec::new();
loop {
let _ = unsafe { ResetEvent(event.raw()) };
let mut overlapped = OVERLAPPED {
hEvent: event.raw(),
..Default::default()
};
let mut bytes_read = 0u32;
let result = unsafe {
ReadFile(
handle,
Some(buf.as_mut_slice()),
Some(&mut bytes_read),
Some(&mut overlapped),
)
};
match result {
Ok(()) => {}
Err(e) if e.code() == ERROR_IO_PENDING_HRESULT => {
bytes_read =
unsafe { await_overlapped(handle, event.raw(), &overlapped, deadline)? };
}
Err(e) => {
return Err(io::Error::new(
io::ErrorKind::BrokenPipe,
format!("ReadFile: {e}"),
));
}
}
if bytes_read == 0 {
return Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
"peer disconnected",
));
}
raw.extend_from_slice(&buf[..bytes_read as usize]);
if raw.contains(&b'\n') {
break;
}
}
String::from_utf8(raw).map_err(|e| {
io::Error::new(
io::ErrorKind::InvalidData,
format!("non-UTF-8 data from pipe: {e}"),
)
})
}
pub fn send_message(msg: &SocketMessage) -> io::Result<SocketResponse> {
send_message_to(&message::pipe_name(), msg)
}
pub fn send_message_to(pipe_name: &str, msg: &SocketMessage) -> io::Result<SocketResponse> {
let handle = connect_to_named_pipe(pipe_name)?;
let deadline = Instant::now() + IPC_TIMEOUT;
let wire = message::encode_message(msg)?;
write_all_overlapped(handle.raw(), wire.as_bytes(), deadline)?;
let line = read_line_overlapped(handle.raw(), deadline)?;
message::decode_message(&line).ok_or_else(|| {
io::Error::new(
io::ErrorKind::InvalidData,
format!("failed to parse response: {line:?}"),
)
})
}
#[must_use]
pub fn is_daemon_running() -> bool {
for _ in 0..3 {
if connect_to_pipe().is_ok() {
return true;
}
std::thread::sleep(std::time::Duration::from_millis(50));
}
false
}
fn connect_to_pipe() -> io::Result<PipeHandle> {
connect_to_named_pipe(&message::pipe_name())
}
fn connect_to_named_pipe(pipe_name: &str) -> io::Result<PipeHandle> {
let name = wide(pipe_name);
let handle = unsafe {
CreateFileW(
windows::core::PCWSTR(name.as_ptr()),
FILE_GENERIC_READ.0 | FILE_GENERIC_WRITE.0,
FILE_SHARE_READ | FILE_SHARE_WRITE,
None,
OPEN_EXISTING,
FILE_ATTRIBUTE_NORMAL | FILE_FLAG_OVERLAPPED,
None,
)
}
.map_err(|_| io::Error::new(io::ErrorKind::ConnectionRefused, "daemon not running"))?;
Ok(PipeHandle(handle))
}
fn wide(s: &str) -> Vec<u16> {
OsStr::new(s)
.encode_wide()
.chain(std::iter::once(0))
.collect()
}
#[cfg(test)]
impl PipeServer {
pub(crate) fn test_dummy() -> Self {
let dummy = unsafe {
CreateEventW(None, true, false, windows::core::PCWSTR::null())
.expect("CreateEventW should not fail in test_dummy")
};
Self {
handle: PipeHandle(dummy),
connected_event: EventHandle::new().expect("EventHandle::new in test_dummy"),
accept_thread: None,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::{Duration, Instant};
#[test]
fn ms_until_future_deadline_returns_positive_milliseconds() {
let deadline = Instant::now() + Duration::from_millis(500);
let result = ms_until(deadline);
assert!(result.is_some());
let ms = result.unwrap();
assert!((499..=501).contains(&ms), "expected ~500 ms, got {ms}");
}
#[test]
fn ms_until_past_deadline_returns_none() {
let deadline = Instant::now() - Duration::from_millis(100);
assert_eq!(ms_until(deadline), None);
}
#[test]
fn ms_until_nearly_expired_floor_at_one() {
let deadline = Instant::now() + Duration::from_micros(500);
let result = ms_until(deadline);
assert_eq!(result, Some(1));
}
#[test]
fn ms_until_large_deadline_capped_at_u32_max() {
let deadline = Instant::now() + Duration::from_secs(60 * 24 * 3600);
let result = ms_until(deadline);
assert_eq!(result, Some(u32::MAX));
}
#[test]
fn ms_until_zero_remaining_floors_to_one() {
let deadline = Instant::now();
let result = ms_until(deadline);
assert!(
result.is_some(),
"deadline at now should yield Some(1), got {result:?}"
);
assert_eq!(result.unwrap(), 1);
}
#[test]
fn wide_empty_string_produces_single_null() {
let result = wide("");
assert_eq!(result, vec![0u16]);
}
#[test]
fn wide_ascii_string_includes_null_terminator() {
let result = wide("abc");
assert_eq!(result, vec![b'a' as u16, b'b' as u16, b'c' as u16, 0]);
}
#[test]
fn wide_unicode_produces_correct_code_units() {
let result = wide("é");
assert_eq!(result, vec![0x00E9, 0]);
let result_surrogate = wide("𐍈");
assert_eq!(result_surrogate, vec![0xD800, 0xDF48, 0]);
}
#[test]
fn wide_never_produces_empty_vec() {
let result = wide("");
assert!(!result.is_empty());
assert_eq!(*result.last().unwrap(), 0);
}
}