use std::ffi::OsString;
use std::io::Write as _;
use std::path::PathBuf;
use std::sync::{Arc, Mutex, OnceLock};
use std::time::Duration;
use tokio_util::sync::CancellationToken;
use windows_service::service::{ServiceControl, ServiceControlAccept, ServiceExitCode, ServiceState, ServiceStatus, ServiceType};
use windows_service::service_control_handler::{self, ServiceControlHandlerResult, ServiceStatusHandle};
use windows_service::service_dispatcher;
use zeroize::Zeroize;
use super::serve::{ServeArgs, host_options, run_listeners_with_shutdown};
const SERVICE_NAME: &str = "AuvDevice";
static CONFIG: OnceLock<(ServeArgs, PathBuf)> = OnceLock::new();
windows_service::define_windows_service!(service_entry, service_main);
pub fn run(args: ServeArgs, project_root: PathBuf) -> Result<i32, String> {
CONFIG.set((args, project_root)).map_err(|_| "AUV service was already dispatched".to_string())?;
service_dispatcher::start(SERVICE_NAME, service_entry)
.map_err(|error| format!("failed to join Windows Service Control Manager: {error}"))?;
Ok(0)
}
pub fn issue_bootstrap_token() -> Result<i32, String> {
require_local_system_session_zero()?;
let mut token = auv_daemon::issue_windows_bootstrap_token()?;
let output = (|| {
let mut stdout = std::io::stdout().lock();
stdout.write_all(token.as_bytes()).map_err(|error| format!("failed to write bootstrap token: {error}"))?;
stdout.write_all(b"\n").map_err(|error| format!("failed to terminate bootstrap token: {error}"))?;
stdout.flush().map_err(|error| format!("failed to flush bootstrap token: {error}"))
})();
token.zeroize();
output?;
Ok(0)
}
fn validate(args: &ServeArgs) -> Result<(), String> {
validate_with_root(args, &auv_daemon::windows_device_entry_store_root()?)
}
fn validate_with_root(args: &ServeArgs, expected_root: &std::path::Path) -> Result<(), String> {
let root = args.store_root.as_ref().ok_or("--windows-service requires --store-root")?;
let pairing = args.pairing_store.as_ref().ok_or("--windows-service requires --pairing-store")?;
if !root.is_absolute() || !pairing.is_absolute() {
return Err("Windows service store paths must be absolute".into());
}
if root != expected_root || pairing != &root.join("pairings.json") {
return Err("Windows service requires --store-root at ProgramData\\AUVDeviceEntry and --pairing-store at its pairings.json".into());
}
let [listener] = args.listeners.as_slice() else {
return Err("Windows service requires exactly one explicit http://LOOPBACK_IP:PORT --listen URI".into());
};
let address = listener
.strip_prefix("http://")
.ok_or("Windows service requires an http://LOOPBACK_IP:PORT --listen URI")?
.parse::<std::net::SocketAddr>()
.map_err(|error| format!("invalid Windows service --listen URI: {error}"))?;
if !address.ip().is_loopback() || address.port() == 0 {
return Err("Windows service --listen must use a loopback IP and nonzero port".into());
}
if args.discovery_file.is_some() || args.daemon_idle_timeout.is_some() || !args.runner_providers.is_empty() {
return Err("Windows service does not accept --discovery-file, --daemon-idle-timeout, or --runner-provider".into());
}
Ok(())
}
fn service_main(_scm_arguments: Vec<OsString>) {
if let Err(error) = serve_service() {
eprintln!("AUV Windows service failed: {error}");
}
}
fn serve_service() -> Result<(), String> {
let (args, project_root) = CONFIG.get().ok_or("AUV service configuration is missing")?.clone();
let shutdown = CancellationToken::new();
let handles = Arc::new(Mutex::new(None::<ServiceStatusHandle>));
let handler_shutdown = shutdown.clone();
let handler_handles = Arc::clone(&handles);
let status = service_control_handler::register(SERVICE_NAME, move |control| match control {
ServiceControl::Stop => {
if let Ok(guard) = handler_handles.lock()
&& let Some(handle) = *guard
{
let _ = report(&handle, ServiceState::StopPending, ServiceControlAccept::empty(), 0, 1, Duration::from_secs(30));
}
handler_shutdown.cancel();
ServiceControlHandlerResult::NoError
}
ServiceControl::Interrogate => ServiceControlHandlerResult::NoError,
_ => ServiceControlHandlerResult::NotImplemented,
})
.map_err(|error| format!("failed to register SCM control handler: {error}"))?;
*handles.lock().map_err(|_| "SCM status lock was poisoned")? = Some(status);
let result = (|| {
report(&status, ServiceState::StartPending, ServiceControlAccept::empty(), 0, 1, Duration::from_secs(30))?;
validate(&args)?;
require_local_system_session_zero()?;
let options = service_options(args)?;
let runtime = tokio::runtime::Builder::new_multi_thread()
.enable_all()
.build()
.map_err(|error| format!("failed to create Windows service runtime: {error}"))?;
let heartbeat_shutdown = shutdown.clone();
let heartbeat_status = status;
runtime.spawn(async move {
heartbeat_shutdown.cancelled().await;
let mut ticker = tokio::time::interval(Duration::from_secs(10));
ticker.tick().await;
for checkpoint in 2.. {
ticker.tick().await;
if report(&heartbeat_status, ServiceState::StopPending, ServiceControlAccept::empty(), 0, checkpoint, Duration::from_secs(30))
.is_err()
{
break;
}
}
});
runtime.block_on(run_listeners_with_shutdown(options, &project_root, shutdown.clone(), || {
if shutdown.is_cancelled() {
return Err("service stopped during startup".into());
}
report(&status, ServiceState::Running, ServiceControlAccept::STOP, 0, 0, Duration::ZERO)
}))?;
Ok::<(), String>(())
})();
let exit_code = if result.is_ok() { 0 } else { 1 };
let stopped = report(&status, ServiceState::Stopped, ServiceControlAccept::empty(), exit_code, 0, Duration::ZERO);
result.and(stopped)
}
fn service_options(args: ServeArgs) -> Result<super::serve::HostOptions, String> {
let mut options = host_options(args)?;
options.publish_discovery = false;
options.local_driver_runner = false;
options.emit_bound_endpoints = false;
options.enable_device_entry = true;
Ok(options)
}
fn report(
handle: &ServiceStatusHandle,
state: ServiceState,
controls: ServiceControlAccept,
exit_code: u32,
checkpoint: u32,
wait_hint: Duration,
) -> Result<(), String> {
handle
.set_service_status(ServiceStatus {
service_type: ServiceType::OWN_PROCESS,
current_state: state,
controls_accepted: controls,
exit_code: ServiceExitCode::Win32(exit_code),
checkpoint,
wait_hint,
process_id: None,
})
.map_err(|error| format!("failed to report SCM state {state:?}: {error}"))
}
fn require_local_system_session_zero() -> Result<(), String> {
use std::mem::{align_of, size_of};
use windows::Win32::Foundation::{CloseHandle, HANDLE};
use windows::Win32::Security::{GetTokenInformation, IsWellKnownSid, TOKEN_QUERY, TOKEN_USER, TokenUser, WinLocalSystemSid};
use windows::Win32::System::RemoteDesktop::ProcessIdToSessionId;
use windows::Win32::System::Threading::{GetCurrentProcess, GetCurrentProcessId, OpenProcessToken};
struct Token(HANDLE);
impl Drop for Token {
fn drop(&mut self) {
let _ = unsafe { CloseHandle(self.0) };
}
}
let mut session = u32::MAX;
unsafe { ProcessIdToSessionId(GetCurrentProcessId(), &mut session) }.map_err(|error| error.to_string())?;
if session != 0 {
return Err("Windows service must run in Session 0".into());
}
let mut raw = HANDLE::default();
unsafe { OpenProcessToken(GetCurrentProcess(), TOKEN_QUERY, &mut raw) }.map_err(|error| error.to_string())?;
let token = Token(raw);
let mut bytes = 0u32;
let _ = unsafe { GetTokenInformation(token.0, TokenUser, None, 0, &mut bytes) };
if bytes < size_of::<TOKEN_USER>() as u32 || align_of::<TOKEN_USER>() > align_of::<usize>() {
return Err("Windows service token user information is invalid".into());
}
let mut data = vec![0usize; (bytes as usize).div_ceil(size_of::<usize>())];
unsafe { GetTokenInformation(token.0, TokenUser, Some(data.as_mut_ptr().cast()), bytes, &mut bytes) }
.map_err(|error| error.to_string())?;
if (bytes as usize) < size_of::<TOKEN_USER>() {
return Err("Windows service token user information is truncated".into());
}
let user = unsafe { data.as_ptr().cast::<TOKEN_USER>().read() };
if user.User.Sid.0.is_null() || !unsafe { IsWellKnownSid(user.User.Sid, WinLocalSystemSid) }.as_bool() {
return Err("Windows service must run as LocalSystem".into());
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn service_requires_explicit_private_store_and_listener() {
let root = tempfile::tempdir().unwrap();
let root = root.path();
let mut args = ServeArgs {
id: None,
listeners: vec!["http://127.0.0.1:9847".into()],
pairing_store: Some(root.join("pairings.json")),
store_root: Some(root.to_path_buf()),
discovery_file: None,
no_discovery: false,
daemon_idle_timeout: None,
runner_providers: Vec::new(),
windows_service: true,
};
assert!(validate_with_root(&args, root).is_ok());
let service = service_options(args.clone()).unwrap();
assert!(!service.local_driver_runner);
assert!(!service.publish_discovery);
assert!(!service.emit_bound_endpoints);
assert!(service.enable_device_entry);
let absent_root = root.join("new-store");
args.store_root = Some(absent_root.clone());
args.pairing_store = Some(absent_root.join("pairings.json"));
assert!(validate_with_root(&args, &absent_root).is_ok());
args.store_root = Some(root.to_path_buf());
args.pairing_store = Some(root.join("pairings.json"));
args.listeners.clear();
assert!(validate_with_root(&args, root).unwrap_err().contains("--listen"));
args.listeners.push("http://127.0.0.1:9847".into());
args.listeners[0] = "http://0.0.0.0:9847".into();
assert!(validate_with_root(&args, root).unwrap_err().contains("loopback"));
args.listeners[0] = "http://127.0.0.1:9847".into();
args.pairing_store = Some(root.join("nested").join("pairings.json"));
assert!(validate_with_root(&args, root).unwrap_err().contains("pairings.json"));
args.pairing_store = Some(root.join("pairings.json"));
args.runner_providers.push(root.join("provider.json"));
assert!(validate_with_root(&args, root).unwrap_err().contains("--runner-provider"));
}
}