use crate::com::HandleMtaLease;
use crate::error::WslcError;
use crate::session::handle::WslcSessionHandle;
use std::ffi::CString;
use std::os::windows::ffi::OsStrExt;
use std::path::{Path, PathBuf};
use wslcsdk_sys::types::{WslcSessionFeatureFlags, WslcSessionSettings, WslcVhdRequirements};
use wslcsdk_sys::*;
fn to_wide_null(s: &str) -> Result<Vec<u16>, WslcError> {
if s.contains('\0') {
return Err(WslcError::NulError("字符串包含非法空字符".to_string()));
}
Ok(s.encode_utf16().chain(std::iter::once(0)).collect())
}
pub(crate) fn path_to_wide_null(path: impl AsRef<Path>) -> Result<Vec<u16>, WslcError> {
let mut wide: Vec<u16> = path.as_ref().as_os_str().encode_wide().collect();
if wide.contains(&0) {
return Err(WslcError::NulError(format!(
"路径中包含非法空字符: {}",
path.as_ref().display()
)));
}
wide.push(0);
Ok(wide)
}
#[derive(Clone, Debug)]
pub struct VhdRequirementsData {
pub name: Option<String>,
pub size_bytes: u64,
pub vhd_type: WslcVhdType,
pub flags: WslcVhdRequirementsFlags,
pub uid: u32,
pub gid: u32,
}
#[derive(Clone, Debug)]
pub struct SessionBuilder {
name: String,
storage_path: PathBuf,
cpu_count: Option<u32>,
memory_mb: Option<u32>,
timeout_ms: Option<u32>,
vhd: Option<VhdRequirementsData>,
feature_flags: WslcSessionFeatureFlags,
}
pub fn default_session_root() -> PathBuf {
let base_dir = std::env::var("LOCALAPPDATA")
.map(|p| PathBuf::from(p).join("wslc"))
.unwrap_or_else(|_| {
std::env::var("USERPROFILE")
.map(|p| PathBuf::from(p).join(".wslc"))
.unwrap_or_else(|_| PathBuf::from(r"C:\wslc"))
});
base_dir.join("sessions")
}
pub(crate) fn default_storage_path(name: &str) -> PathBuf {
default_session_root().join(name)
}
impl SessionBuilder {
pub fn new_default(name: impl Into<String>) -> Result<Self, WslcError> {
let name_str = name.into();
let storage_path = default_storage_path(&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_path_buf(),
cpu_count: None,
memory_mb: None,
timeout_ms: None,
vhd: None,
feature_flags: 0,
}
}
pub fn name(&self) -> &str {
&self.name
}
#[must_use]
pub fn storage_path(&self) -> &Path {
&self.storage_path
}
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 |= WSLC_SESSION_FEATURE_FLAG_ENABLE_GPU;
} else {
self.feature_flags &= !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) };
WslcError::check_hr(hr, "设置会话 CPU 核心数失败")?;
}
if let Some(mem) = self.memory_mb {
let hr = unsafe { WslcSetSessionSettingsMemory(settings, mem) };
WslcError::check_hr(hr, "设置会话内存配额失败")?;
}
if let Some(timeout) = self.timeout_ms {
let hr = unsafe { WslcSetSessionSettingsTimeout(settings, timeout) };
WslcError::check_hr(hr, "设置会话超时时间失败")?;
}
if self.feature_flags != 0 {
let hr = unsafe { WslcSetSessionSettingsFeatureFlags(settings, self.feature_flags) };
WslcError::check_hr(hr, "设置会话特性标志位失败")?;
}
Ok(())
}
fn apply_vhd_settings(
&self,
settings: &mut WslcSessionSettings,
) -> Result<Option<CString>, WslcError> {
let Some(ref v) = self.vhd else {
return Ok(None);
};
if v.flags != WSLC_VHD_REQ_FLAG_NONE {
return Err(WslcError::InvalidConfiguration(
"会话根 VHD 规格不支持指定所有者等标志位,flags 必须为 NONE".to_string(),
));
}
let c_str = match &v.name {
Some(n) => Some(
CString::new(n.as_str())
.map_err(|e| WslcError::NulError(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: WSLC_VHD_REQ_FLAG_NONE,
uid: 0,
gid: 0,
};
let hr = unsafe { WslcSetSessionSettingsVhd(settings, &raw_req) };
WslcError::check_hr(hr, "设置会话 VHD 存储规格失败")?;
Ok(c_str)
}
fn create_raw_session(
&self,
settings: &mut WslcSessionSettings,
) -> Result<WslcSession, WslcError> {
let mut raw_session = WslcSession::NULL;
let mut err_msg: *mut u16 = std::ptr::null_mut();
let hr = unsafe { WslcCreateSession(settings, &mut raw_session, &mut err_msg) };
unsafe {
WslcError::check(hr, err_msg, format!("创建会话失败,名称: '{}'", self.name))?;
}
if raw_session.is_null() {
Err(WslcError::InvalidHandle)
} else {
Ok(raw_session)
}
}
pub fn build(self) -> Result<WslcSessionHandle, WslcError> {
crate::system::WslcSystem::ensure_sdk_available()?;
let wide_name = to_wide_null(&self.name)?;
let wide_path = path_to_wide_null(&self.storage_path)?;
let _com_guard = crate::com::try_initialize_mta()?;
let mut settings = WslcSessionSettings::default();
let hr = unsafe {
WslcInitSessionSettings(wide_name.as_ptr(), wide_path.as_ptr(), &mut settings)
};
WslcError::check_hr(hr, "初始化会话设置失败")?;
self.apply_resource_limits(&mut settings)?;
let _retained_vhd_name = self.apply_vhd_settings(&mut settings)?;
let raw_session = self.create_raw_session(&mut settings)?;
log::info!("WSLC 会话创建成功,会话名称: '{}'", self.name);
Self::wrap_created_session(raw_session, self.name)
}
fn wrap_created_session(
raw: WslcSession,
name: String,
) -> Result<WslcSessionHandle, WslcError> {
unsafe {
HandleMtaLease::wrap_raw(
raw,
|h| {
let _ = WslcReleaseSession(h);
},
|raw, mta| WslcSessionHandle::from_acquired_lease(raw, name, mta),
)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_to_wide_null_validation() {
let valid = to_wide_null("valid_path");
assert!(valid.is_ok());
let invalid = to_wide_null("invalid\0path");
assert!(invalid.is_err());
let valid_path = path_to_wide_null(r"C:\valid\path");
assert!(valid_path.is_ok());
let invalid_path = path_to_wide_null("C:\\invalid\0\\path");
assert!(invalid_path.is_err());
}
#[test]
fn test_path_to_wide_null_is_nul_terminated() {
let wide = path_to_wide_null(r"C:\temp").expect("路径转换失败");
assert_eq!(*wide.last().expect("宽字符向量不应为空"), 0);
assert_eq!(wide.len(), r"C:\temp".encode_utf16().count() + 1);
}
#[test]
fn test_storage_path_preserves_non_utf8_path() {
use std::ffi::OsString;
use std::os::windows::ffi::OsStringExt;
let raw_wide: Vec<u16> = vec![0x0043, 0x003A, 0x005C, 0xD800, 0x005C, 0x0074];
let path = PathBuf::from(OsString::from_wide(&raw_wide));
let builder = SessionBuilder::new("test-session", &path);
assert_eq!(builder.storage_path().as_os_str(), path.as_os_str());
let wide = path_to_wide_null(builder.storage_path()).expect("路径转换失败");
assert_eq!(&wide[..raw_wide.len()], &raw_wide[..]);
assert_eq!(wide[raw_wide.len()], 0);
}
#[test]
fn test_default_storage_path_appends_session_name() {
let path = default_storage_path("unit-test-session");
assert!(
path.ends_with("sessions/unit-test-session"),
"默认路径应以 sessions 目录加会话名结尾,实际为: {}",
path.display()
);
}
#[test]
fn test_new_default_reports_unwritable_location_instead_of_panicking() {
let builder = SessionBuilder::new("dup", r"NUL\invalid");
assert_eq!(builder.name(), "dup");
}
#[test]
fn test_session_builder_rejects_non_none_vhd_flags() {
let builder = SessionBuilder::new("test-session", r"C:\temp").vhd(VhdRequirementsData {
name: None,
size_bytes: 1024 * 1024,
vhd_type: WslcVhdType::Dynamic,
flags: WSLC_VHD_REQ_FLAG_OWNER,
uid: 1000,
gid: 1000,
});
match builder.build() {
Err(WslcError::InvalidConfiguration(_)) => {}
other => panic!("预期返回 InvalidConfiguration 错误,实际为: {other:?}"),
}
}
}