use super::{Ipc, IpcListener, Liveness, PeerReject};
use std::ffi::{OsStr, OsString};
use std::io;
use std::os::windows::ffi::{OsStrExt, OsStringExt};
use std::path::PathBuf;
use std::ptr;
use tokio::net::windows::named_pipe::{ClientOptions, NamedPipeServer, ServerOptions};
use windows_sys::Win32::Foundation::{ERROR_PIPE_BUSY, HANDLE};
use windows_sys::Win32::Security::{
EqualSid, GetTokenInformation, RevertToSelf, TokenUser, PSID, TOKEN_QUERY, TOKEN_USER,
};
use windows_sys::Win32::Storage::FileSystem::{
FindClose, FindFirstFileW, FindNextFileW, WIN32_FIND_DATAW,
};
use windows_sys::Win32::System::Pipes::ImpersonateNamedPipeClient;
use windows_sys::Win32::System::Threading::{
GetCurrentProcess, GetCurrentThread, OpenProcessToken, OpenThreadToken,
};
pub type WindowsIpcStream = NamedPipeServerOrClient;
#[derive(Debug, Clone, Copy, Default)]
pub struct WindowsIpc;
impl WindowsIpc {
pub const fn new() -> Self {
Self
}
}
impl crate::sealed::Sealed for WindowsIpc {}
fn pipe_name(id: &str) -> String {
format!(r"\\.\pipe\hotl-{id}")
}
const PREFIX: &str = "hotl-";
pub enum NamedPipeServerOrClient {
Server(NamedPipeServer),
Client(tokio::net::windows::named_pipe::NamedPipeClient),
}
macro_rules! delegate {
($self:ident, $inner:ident, $body:expr) => {
match $self.get_mut() {
NamedPipeServerOrClient::Server($inner) => $body,
NamedPipeServerOrClient::Client($inner) => $body,
}
};
}
impl tokio::io::AsyncRead for NamedPipeServerOrClient {
fn poll_read(
self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
buf: &mut tokio::io::ReadBuf<'_>,
) -> std::task::Poll<io::Result<()>> {
delegate!(self, s, std::pin::Pin::new(s).poll_read(cx, buf))
}
}
impl tokio::io::AsyncWrite for NamedPipeServerOrClient {
fn poll_write(
self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
buf: &[u8],
) -> std::task::Poll<io::Result<usize>> {
delegate!(self, s, std::pin::Pin::new(s).poll_write(cx, buf))
}
fn poll_flush(
self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<io::Result<()>> {
delegate!(self, s, std::pin::Pin::new(s).poll_flush(cx))
}
fn poll_shutdown(
self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<io::Result<()>> {
delegate!(self, s, std::pin::Pin::new(s).poll_shutdown(cx))
}
}
pub struct WindowsIpcListener {
name: String,
next: Option<NamedPipeServer>,
}
impl IpcListener for WindowsIpcListener {
type Stream = WindowsIpcStream;
async fn accept(&mut self) -> io::Result<Self::Stream> {
let server = match self.next.take() {
Some(s) => s,
None => new_server_instance(&self.name, false)?,
};
server.connect().await?;
self.next = Some(new_server_instance(&self.name, false)?);
Ok(NamedPipeServerOrClient::Server(server))
}
}
fn new_server_instance(name: &str, first: bool) -> io::Result<NamedPipeServer> {
let mut opts = ServerOptions::new();
opts.first_pipe_instance(first)
.reject_remote_clients(true);
let mut attrs = crate::privatefs::owner_only_attributes()?;
unsafe { opts.create_with_security_attributes_raw(name, attrs.as_ptr()) }
}
impl Ipc for WindowsIpc {
type Listener = WindowsIpcListener;
type Stream = WindowsIpcStream;
const LEAVES_STALE_ARTIFACT: bool = false;
fn bind_private(&self, id: &str) -> io::Result<Self::Listener> {
let name = pipe_name(id);
let first = new_server_instance(&name, true)?;
Ok(WindowsIpcListener {
name,
next: Some(first),
})
}
async fn connect(&self, id: &str) -> io::Result<Self::Stream> {
const SECURITY_IDENTIFICATION: u32 = 0x0001_0000;
let name = pipe_name(id);
let client = tokio::time::timeout(std::time::Duration::from_secs(2), async {
loop {
match ClientOptions::new()
.security_qos_flags(SECURITY_IDENTIFICATION)
.open(&name)
{
Ok(c) => return Ok::<_, io::Error>(c),
Err(e) if e.raw_os_error() == Some(ERROR_PIPE_BUSY as i32) => {
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
}
Err(e) => return Err(e),
}
}
})
.await
.map_err(|_| io::Error::new(io::ErrorKind::TimedOut, "the named pipe stayed busy"))??;
Ok(NamedPipeServerOrClient::Client(client))
}
fn authenticate_peer(&self, stream: &Self::Stream) -> Result<(), PeerReject> {
use std::os::windows::io::AsRawHandle;
let NamedPipeServerOrClient::Server(server) = stream else {
return Ok(());
};
let handle = server.as_raw_handle() as HANDLE;
if unsafe { ImpersonateNamedPipeClient(handle) } == 0 {
return Err(PeerReject(format!(
"could not impersonate the peer: {}",
io::Error::last_os_error()
)));
}
let peer = token_user(true);
unsafe { RevertToSelf() };
let peer = peer.map_err(|e| PeerReject(format!("the peer's token is unreadable: {e}")))?;
let me = token_user(false)
.map_err(|e| PeerReject(format!("our own token is unreadable: {e}")))?;
let same = unsafe { EqualSid(sid_of(&peer), sid_of(&me)) } != 0;
if !same {
return Err(PeerReject("the peer runs as a different user".to_string()));
}
Ok(())
}
fn liveness(&self, id: &str) -> Liveness {
match ClientOptions::new().open(pipe_name(id)) {
Ok(_) => Liveness::Live,
Err(e) if e.raw_os_error() == Some(ERROR_PIPE_BUSY as i32) => {
Liveness::Live
}
Err(_) => Liveness::Dead,
}
}
fn list_live(&self) -> Vec<String> {
let pattern: Vec<u16> = OsStr::new(r"\\.\pipe\*")
.encode_wide()
.chain(Some(0))
.collect();
let mut data: WIN32_FIND_DATAW = unsafe { std::mem::zeroed() };
let find = unsafe { FindFirstFileW(pattern.as_ptr(), &mut data) };
if find.is_null() || find as isize == -1 {
return Vec::new();
}
let mut out = Vec::new();
loop {
let len = data
.cFileName
.iter()
.position(|&c| c == 0)
.unwrap_or(data.cFileName.len());
let name = OsString::from_wide(&data.cFileName[..len]);
if let Some(id) = name.to_string_lossy().strip_prefix(PREFIX) {
out.push(id.to_string());
}
if unsafe { FindNextFileW(find, &mut data) } == 0 {
break;
}
}
unsafe { FindClose(find) };
out
}
fn artifact_path(&self, _id: &str) -> Option<PathBuf> {
None
}
}
fn token_user(thread: bool) -> io::Result<Vec<u8>> {
let mut token: HANDLE = ptr::null_mut();
let opened = if thread {
unsafe { OpenThreadToken(GetCurrentThread(), TOKEN_QUERY, 1, &mut token) }
} else {
unsafe { OpenProcessToken(GetCurrentProcess(), TOKEN_QUERY, &mut token) }
};
if opened == 0 {
return Err(io::Error::last_os_error());
}
let mut needed = 0u32;
unsafe { GetTokenInformation(token, TokenUser, ptr::null_mut(), 0, &mut needed) };
let mut buf = vec![0u8; needed as usize];
let ok = unsafe {
GetTokenInformation(
token,
TokenUser,
buf.as_mut_ptr().cast(),
needed,
&mut needed,
)
};
unsafe { windows_sys::Win32::Foundation::CloseHandle(token) };
if ok == 0 {
return Err(io::Error::last_os_error());
}
Ok(buf)
}
fn sid_of(buf: &[u8]) -> PSID {
unsafe { (*(buf.as_ptr() as *const TOKEN_USER)).User.Sid }
}