use std::io;
use std::io::Error;
use std::ops::Index;
use std::sync::OnceLock;
use crate::signal::RxFuture;
use crate::sync::watch;
use windows_sys::core::BOOL;
use windows_sys::Win32::System::Console as console;
type EventInfo = watch::Sender<()>;
#[derive(Clone, Copy)]
#[repr(u32)]
enum SignalKind {
CtrlC = console::CTRL_C_EVENT,
CtrlBreak = console::CTRL_BREAK_EVENT,
CtrlClose = console::CTRL_CLOSE_EVENT,
CtrlLogoff = console::CTRL_LOGOFF_EVENT,
CtrlShutdown = console::CTRL_SHUTDOWN_EVENT,
}
impl SignalKind {
const fn terminates(&self) -> bool {
matches!(
self,
Self::CtrlClose | Self::CtrlLogoff | Self::CtrlShutdown
)
}
}
pub(super) fn ctrl_break() -> io::Result<RxFuture> {
new(SignalKind::CtrlBreak)
}
pub(super) fn ctrl_close() -> io::Result<RxFuture> {
new(SignalKind::CtrlClose)
}
pub(super) fn ctrl_c() -> io::Result<RxFuture> {
new(SignalKind::CtrlC)
}
pub(super) fn ctrl_logoff() -> io::Result<RxFuture> {
new(SignalKind::CtrlLogoff)
}
pub(super) fn ctrl_shutdown() -> io::Result<RxFuture> {
new(SignalKind::CtrlShutdown)
}
fn new(signal: SignalKind) -> io::Result<RxFuture> {
let registry = REGISTRY
.get_or_init(
|| match unsafe { console::SetConsoleCtrlHandler(Some(handler), 1) } {
0 => Err(Error::last_os_error().raw_os_error().expect("unreachable")),
_ => Ok(Registry::default()),
},
)
.as_ref()
.map_err(|&code| Error::from_raw_os_error(code))?;
let rx = registry[signal].subscribe();
Ok(RxFuture::new(rx))
}
#[derive(Debug, Default)]
struct Registry {
ctrl_break: EventInfo,
ctrl_close: EventInfo,
ctrl_c: EventInfo,
ctrl_logoff: EventInfo,
ctrl_shutdown: EventInfo,
}
impl Index<SignalKind> for Registry {
type Output = EventInfo;
fn index(&self, signal: SignalKind) -> &Self::Output {
match signal {
SignalKind::CtrlC => &self.ctrl_c,
SignalKind::CtrlBreak => &self.ctrl_break,
SignalKind::CtrlClose => &self.ctrl_close,
SignalKind::CtrlLogoff => &self.ctrl_logoff,
SignalKind::CtrlShutdown => &self.ctrl_shutdown,
}
}
}
static REGISTRY: OnceLock<Result<Registry, i32>> = OnceLock::new();
unsafe extern "system" fn handler(ty: u32) -> BOOL {
let signal = match ty {
console::CTRL_C_EVENT => SignalKind::CtrlC,
console::CTRL_BREAK_EVENT => SignalKind::CtrlBreak,
console::CTRL_CLOSE_EVENT => SignalKind::CtrlClose,
console::CTRL_LOGOFF_EVENT => SignalKind::CtrlLogoff,
console::CTRL_SHUTDOWN_EVENT => SignalKind::CtrlShutdown,
_ => return 0,
};
let Ok(registry) = REGISTRY.wait().as_ref() else {
return 0;
};
match registry[signal].send(()) {
Ok(_) if signal.terminates() => loop {
std::thread::park();
},
Ok(_) => 1,
Err(_) => 0,
}
}
#[cfg(all(test, not(loom)))]
mod tests {
use super::*;
use crate::runtime::Runtime;
use tokio_test::{assert_ok, assert_pending, assert_ready_ok, task};
unsafe fn raise_event(signal: SignalKind) {
if signal.terminates() {
std::thread::spawn(move || unsafe { super::handler(signal as u32) });
} else {
unsafe { super::handler(signal as u32) };
}
}
#[test]
fn ctrl_c() {
let rt = rt();
let _enter = rt.enter();
let mut ctrl_c = task::spawn(crate::signal::ctrl_c());
assert_pending!(ctrl_c.poll());
unsafe {
raise_event(SignalKind::CtrlC);
}
assert_ready_ok!(ctrl_c.poll());
}
#[test]
fn ctrl_break() {
let rt = rt();
rt.block_on(async {
let mut ctrl_break = assert_ok!(crate::signal::windows::ctrl_break());
unsafe {
raise_event(SignalKind::CtrlBreak);
}
ctrl_break.recv().await.unwrap();
});
}
#[test]
fn ctrl_close() {
let rt = rt();
rt.block_on(async {
let mut ctrl_close = assert_ok!(crate::signal::windows::ctrl_close());
unsafe {
raise_event(SignalKind::CtrlClose);
}
ctrl_close.recv().await.unwrap();
});
}
#[test]
fn ctrl_shutdown() {
let rt = rt();
rt.block_on(async {
let mut ctrl_shutdown = assert_ok!(crate::signal::windows::ctrl_shutdown());
unsafe {
raise_event(SignalKind::CtrlShutdown);
}
ctrl_shutdown.recv().await.unwrap();
});
}
#[test]
fn ctrl_logoff() {
let rt = rt();
rt.block_on(async {
let mut ctrl_logoff = assert_ok!(crate::signal::windows::ctrl_logoff());
unsafe {
raise_event(SignalKind::CtrlLogoff);
}
ctrl_logoff.recv().await.unwrap();
});
}
fn rt() -> Runtime {
crate::runtime::Builder::new_current_thread()
.build()
.unwrap()
}
}