use crate::error::WslcError;
use core::ffi::c_void;
use std::ffi::CString;
use std::os::windows::raw::HANDLE;
use std::path::Path;
use std::sync::Arc;
use wslcsdk_sys::types::{
WslcCrashDumpSubscription, WslcSession, WslcSessionCrashDumpInfo, WslcSessionFeatureFlags,
WslcSessionSettings, WslcSessionTerminationReason, WslcVhdRequirements,
};
use wslcsdk_sys::*;
pub(crate) fn to_wide_null(s: &str) -> Result<Vec<u16>, WslcError> {
if s.contains('\0') {
return Err(WslcError::Utf8Error("字符串包含非法空字符".to_string()));
}
Ok(s.encode_utf16().chain(std::iter::once(0)).collect())
}
#[derive(Clone, Debug)]
pub struct VhdRequirementsData {
pub name: Option<String>,
pub size_bytes: u64,
pub vhd_type: wslcsdk_sys::types::WslcVhdType,
pub flags: wslcsdk_sys::types::WslcVhdRequirementsFlags,
pub uid: u32,
pub gid: u32,
}
#[derive(Clone, Debug)]
pub struct SessionBuilder {
name: String,
storage_path: String,
cpu_count: Option<u32>,
memory_mb: Option<u32>,
timeout_ms: Option<u32>,
vhd: Option<VhdRequirementsData>,
feature_flags: WslcSessionFeatureFlags,
}
impl SessionBuilder {
pub fn new_default(name: impl Into<String>) -> Result<Self, WslcError> {
let name_str = name.into();
let base_dir = std::env::var("LOCALAPPDATA")
.map(std::path::PathBuf::from)
.unwrap_or_else(|_| {
std::env::var("USERPROFILE")
.map(|p| std::path::PathBuf::from(p).join(".wslc"))
.unwrap_or_else(|_| std::path::PathBuf::from(r"C:\wslc"))
});
let storage_path = base_dir.join("wslc").join("sessions").join(&name_str);
if let Err(e) = std::fs::create_dir_all(&storage_path) {
return Err(WslcError::Io(format!("创建默认会话存储目录失败: {e}")));
}
Ok(Self::new(name_str, storage_path))
}
pub fn new(name: impl Into<String>, storage_path: impl AsRef<Path>) -> Self {
Self {
name: name.into(),
storage_path: storage_path.as_ref().to_string_lossy().to_string(),
cpu_count: None,
memory_mb: None,
timeout_ms: None,
vhd: None,
feature_flags: 0,
}
}
pub fn cpu_count(mut self, count: u32) -> Self {
self.cpu_count = Some(count);
self
}
pub fn memory_mb(mut self, mb: u32) -> Self {
self.memory_mb = Some(mb);
self
}
pub fn timeout_ms(mut self, ms: u32) -> Self {
self.timeout_ms = Some(ms);
self
}
pub fn feature_flags(mut self, flags: WslcSessionFeatureFlags) -> Self {
self.feature_flags = flags;
self
}
pub fn enable_gpu(mut self, enable: bool) -> Self {
if enable {
self.feature_flags |= wslcsdk_sys::types::WSLC_SESSION_FEATURE_FLAG_ENABLE_GPU;
} else {
self.feature_flags &= !wslcsdk_sys::types::WSLC_SESSION_FEATURE_FLAG_ENABLE_GPU;
}
self
}
pub fn vhd(mut self, vhd: VhdRequirementsData) -> Self {
self.vhd = Some(vhd);
self
}
fn apply_resource_limits(&self, settings: &mut WslcSessionSettings) -> Result<(), WslcError> {
if let Some(cpus) = self.cpu_count {
let hr = unsafe { WslcSetSessionSettingsCpuCount(settings, cpus) };
if hr < 0 {
return Err(WslcError::Win32(
hr as u32,
"设置会话 CPU 核心数失败".to_string(),
));
}
}
if let Some(mem) = self.memory_mb {
let hr = unsafe { WslcSetSessionSettingsMemory(settings, mem) };
if hr < 0 {
return Err(WslcError::Win32(
hr as u32,
"设置会话内存配额失败".to_string(),
));
}
}
if let Some(timeout) = self.timeout_ms {
let hr = unsafe { WslcSetSessionSettingsTimeout(settings, timeout) };
if hr < 0 {
return Err(WslcError::Win32(
hr as u32,
"设置会话超时时间失败".to_string(),
));
}
}
if self.feature_flags != 0 {
let hr = unsafe { WslcSetSessionSettingsFeatureFlags(settings, self.feature_flags) };
if hr < 0 {
return Err(WslcError::Win32(
hr as u32,
"设置会话特性标志位失败".to_string(),
));
}
}
Ok(())
}
fn apply_vhd_settings(
&self,
settings: &mut WslcSessionSettings,
) -> Result<Option<CString>, WslcError> {
let Some(ref v) = self.vhd else {
return Ok(None);
};
let c_str = match &v.name {
Some(n) => Some(
CString::new(n.as_str())
.map_err(|e| WslcError::Utf8Error(format!("VHD 名称非法: {e}")))?,
),
None => None,
};
let raw_req = WslcVhdRequirements {
name: c_str.as_ref().map_or(std::ptr::null(), |s| s.as_ptr()),
size_bytes: v.size_bytes,
vhd_type: v.vhd_type,
flags: v.flags,
uid: v.uid,
gid: v.gid,
};
let hr = unsafe { WslcSetSessionSettingsVhd(settings, &raw_req) };
if hr < 0 {
return Err(WslcError::Win32(
hr as u32,
"设置会话 VHD 存储规格失败".to_string(),
));
}
Ok(c_str)
}
pub fn build(self) -> Result<WslcSessionHandle, WslcError> {
let _com_guard = crate::com::try_initialize_mta()?;
let wide_name = to_wide_null(&self.name)?;
let wide_path = to_wide_null(&self.storage_path)?;
let mut settings = WslcSessionSettings::default();
let hr = unsafe {
WslcInitSessionSettings(wide_name.as_ptr(), wide_path.as_ptr(), &mut settings)
};
if hr < 0 {
return Err(WslcError::Win32(
hr as u32,
"初始化会话设置失败".to_string(),
));
}
self.apply_resource_limits(&mut settings)?;
let _retained_vhd_name = self.apply_vhd_settings(&mut settings)?;
let mut raw_session = WslcSession::NULL;
let mut err_msg: *mut u16 = std::ptr::null_mut();
let hr = unsafe { WslcCreateSession(&mut settings, &mut raw_session, &mut err_msg) };
unsafe {
WslcError::check(hr, err_msg)?;
}
if raw_session.is_null() {
return Err(WslcError::InvalidHandle);
}
Ok(WslcSessionHandle {
inner: Arc::new(SessionInner {
raw: raw_session,
name: self.name,
}),
})
}
}
struct SessionInner {
raw: WslcSession,
name: String,
}
impl Drop for SessionInner {
fn drop(&mut self) {
if !self.raw.is_null() {
unsafe {
let _ = WslcReleaseSession(self.raw);
}
self.raw = WslcSession::NULL;
}
}
}
unsafe impl Send for SessionInner {}
unsafe impl Sync for SessionInner {}
#[derive(Clone)]
pub struct WslcSessionHandle {
inner: Arc<SessionInner>,
}
impl WslcSessionHandle {
pub fn as_raw(&self) -> WslcSession {
self.inner.raw
}
pub fn name(&self) -> &str {
&self.inner.name
}
pub fn termination_event(&self) -> Result<HANDLE, WslcError> {
let _com_guard = crate::com::try_initialize_mta()?;
let mut event: HANDLE = std::ptr::null_mut();
let hr = unsafe { WslcGetSessionTerminationEvent(self.inner.raw, &mut event) };
if hr >= 0 {
Ok(event)
} else {
Err(WslcError::Win32(
hr as u32,
"获取会话终止事件句柄失败".to_string(),
))
}
}
pub fn termination_reason(&self) -> Result<WslcSessionTerminationReason, WslcError> {
let _com_guard = crate::com::try_initialize_mta()?;
let mut reason = WslcSessionTerminationReason::Unknown;
let hr = unsafe { WslcGetSessionTerminationReason(self.inner.raw, &mut reason) };
if hr >= 0 {
Ok(reason)
} else {
Err(WslcError::Win32(
hr as u32,
"获取会话终止原因失败".to_string(),
))
}
}
pub fn terminate(&self) -> Result<(), WslcError> {
let _com_guard = crate::com::try_initialize_mta()?;
let hr = unsafe { WslcTerminateSession(self.inner.raw) };
if hr >= 0 {
Ok(())
} else {
Err(WslcError::Win32(hr as u32, "终止会话失败".to_string()))
}
}
pub fn register_crash_dump_callback<F>(
&self,
callback: F,
) -> Result<CrashDumpSubscription, WslcError>
where
F: Fn(&WslcSessionCrashDumpInfo) + Send + 'static,
{
let _com_guard = crate::com::try_initialize_mta()?;
unsafe extern "system" fn crash_trampoline<F>(
info: *const WslcSessionCrashDumpInfo,
context: *mut c_void,
) where
F: Fn(&WslcSessionCrashDumpInfo),
{
if !context.is_null() && !info.is_null() {
let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
let cb = unsafe { &*(context as *const F) };
cb(unsafe { &*info });
}));
}
}
let boxed_cb = Box::new(callback);
let ctx = Box::into_raw(boxed_cb);
let mut sub = WslcCrashDumpSubscription::NULL;
let mut err_msg: *mut u16 = std::ptr::null_mut();
let hr = unsafe {
WslcRegisterSessionCrashDumpCallback(
self.inner.raw,
Some(crash_trampoline::<F>),
ctx as *mut c_void,
&mut sub,
&mut err_msg,
)
};
if let Err(e) = unsafe { WslcError::check(hr, err_msg) } {
let _ = unsafe { Box::from_raw(ctx) };
return Err(e);
}
unsafe fn drop_ctx<F>(context: *mut c_void) {
if !context.is_null() {
let _ = unsafe { Box::from_raw(context as *mut F) };
}
}
Ok(CrashDumpSubscription {
raw: sub,
ctx: ctx as *mut c_void,
drop_fn: drop_ctx::<F>,
})
}
}
pub struct CrashDumpSubscription {
raw: WslcCrashDumpSubscription,
ctx: *mut c_void,
drop_fn: unsafe fn(*mut c_void),
}
impl Drop for CrashDumpSubscription {
fn drop(&mut self) {
if !self.raw.is_null() {
unsafe {
let _ = WslcReleaseCrashDumpSubscription(self.raw);
}
self.raw = WslcCrashDumpSubscription::NULL;
}
if !self.ctx.is_null() {
unsafe {
(self.drop_fn)(self.ctx);
}
self.ctx = std::ptr::null_mut();
}
}
}
unsafe impl Send for CrashDumpSubscription {}
unsafe impl Sync for CrashDumpSubscription {}