use crate::error::{WslcDomainError, WslcError};
use core::ffi::c_void;
use wslcsdk_sys::types::{WslcComponentFlags, WslcInstallOptions, WslcVersion};
use wslcsdk_sys::*;
#[derive(Debug, Clone, Copy, Default)]
pub struct WslcSystem;
impl WslcSystem {
pub fn is_sdk_available() -> bool {
let dll_name: Vec<u16> = "wslcsdk.dll"
.encode_utf16()
.chain(std::iter::once(0))
.collect();
let handle = unsafe {
windows_sys::Win32::System::LibraryLoader::LoadLibraryExW(
dll_name.as_ptr(),
std::ptr::null_mut(),
windows_sys::Win32::System::LibraryLoader::LOAD_LIBRARY_SEARCH_DEFAULT_DIRS,
)
};
if !handle.is_null() {
unsafe {
windows_sys::Win32::Foundation::FreeLibrary(handle);
}
true
} else {
false
}
}
pub fn ensure_sdk_available() -> Result<(), WslcError> {
if Self::is_sdk_available() {
Ok(())
} else {
Err(WslcDomainError::SdkUpdateNeeded(
"宿主机系统中未检测到可用的 wslcsdk.dll 运行时,请先安装 WSL Containers 支持组件"
.to_string(),
)
.into())
}
}
pub fn get_version() -> Result<WslcVersion, WslcError> {
Self::ensure_sdk_available()?;
let _com_guard = crate::com::try_initialize_mta()?;
let mut ver = WslcVersion::default();
let hr = unsafe { WslcGetVersion(&mut ver) };
WslcError::check_hr(hr, "获取 WSLC 版本失败")?;
Ok(ver)
}
pub fn get_missing_components() -> Result<WslcComponentFlags, WslcError> {
Self::ensure_sdk_available()?;
let _com_guard = crate::com::try_initialize_mta()?;
let mut missing = 0u32;
let hr = unsafe { WslcGetMissingComponents(&mut missing) };
WslcError::check_hr(hr, "检测缺失组件失败")?;
Ok(missing)
}
pub fn install_with_dependencies<F>(
components: WslcComponentFlags,
options: WslcInstallOptions,
on_progress: Option<F>,
) -> Result<(), WslcError>
where
F: FnMut(WslcComponentFlags, u32, u32) + Send,
{
Self::ensure_sdk_available()?;
let _com_guard = crate::com::try_initialize_mta()?;
unsafe extern "system" fn progress_trampoline<F>(
component: WslcComponentFlags,
progress_steps: u32,
total_steps: u32,
context: *mut c_void,
) where
F: FnMut(WslcComponentFlags, u32, u32),
{
if !context.is_null() {
let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| unsafe {
let mutex = &*(context as *const std::sync::Mutex<F>);
if let Ok(mut callback) = mutex.lock() {
callback(component, progress_steps, total_steps);
}
}));
}
}
let (cb, ctx) = match on_progress {
Some(f) => {
let boxed = Box::into_raw(Box::new(std::sync::Mutex::new(f)));
(Some(progress_trampoline::<F> as _), boxed as *mut c_void)
}
None => (None, std::ptr::null_mut()),
};
let hr = unsafe { WslcInstallWithDependencies(components, options, cb, ctx) };
if !ctx.is_null() {
let _ = unsafe { Box::from_raw(ctx as *mut std::sync::Mutex<F>) };
}
WslcError::check_hr(hr, "安装依赖组件失败")
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_sdk_probe_never_panics() {
match WslcSystem::get_version() {
Ok(version) => {
assert!(version.major > 0 || version.minor > 0 || version.revision > 0);
}
Err(WslcError::Domain(WslcDomainError::SdkUpdateNeeded(_))) => {}
Err(WslcError::Hresult(code, _)) => assert_ne!(code, 0),
Err(other) => panic!("预期返回 SdkUpdateNeeded 或 Hresult 错误,实际为: {other:?}"),
}
}
#[test]
fn test_missing_components_probe_never_panics() {
match WslcSystem::get_missing_components() {
Ok(flags) => {
let known = WSLC_COMPONENT_FLAG_VIRTUAL_MACHINE_PLATFORM
| WSLC_COMPONENT_FLAG_WSL_PACKAGE
| WSLC_COMPONENT_FLAG_SDK_NEEDS_UPDATE;
assert_eq!(flags & !known, 0, "返回了未定义的组件标志位: {flags}");
}
Err(WslcError::Domain(WslcDomainError::SdkUpdateNeeded(_))) => {}
Err(WslcError::Hresult(code, _)) => assert_ne!(code, 0),
Err(other) => panic!("预期返回 SdkUpdateNeeded 或 Hresult 错误,实际为: {other:?}"),
}
}
#[test]
fn test_ensure_sdk_available_agrees_with_probe() {
let available = WslcSystem::is_sdk_available();
let ensured = WslcSystem::ensure_sdk_available();
assert_eq!(
available,
ensured.is_ok(),
"is_sdk_available 与 ensure_sdk_available 的判定结果不一致"
);
}
}