use crate::error::WslcError;
use std::cell::{Cell, RefCell};
use std::rc::{Rc, Weak};
use windows_sys::Win32::System::Com::{
CO_MTA_USAGE_COOKIE, COINIT_MULTITHREADED, CoDecrementMTAUsage, CoIncrementMTAUsage,
CoInitializeEx, CoUninitialize,
};
use windows_sys::core::HRESULT;
const RPC_E_CHANGED_MODE: HRESULT = 0x8001_0106_u32 as HRESULT;
thread_local! {
static THREAD_MTA: RefCell<Weak<ComInitialization>> = const { RefCell::new(Weak::new()) };
static THREAD_STA_DEGRADED_WARNED: Cell<bool> = const { Cell::new(false) };
}
#[derive(Debug)]
struct ComInitialization;
impl Drop for ComInitialization {
fn drop(&mut self) {
unsafe {
CoUninitialize();
}
}
}
#[derive(Debug)]
pub struct ProcessMtaGuard {
cookie: CO_MTA_USAGE_COOKIE,
}
impl Drop for ProcessMtaGuard {
fn drop(&mut self) {
if !self.cookie.is_null() {
unsafe {
CoDecrementMTAUsage(self.cookie);
}
self.cookie = std::ptr::null_mut();
}
}
}
unsafe impl Send for ProcessMtaGuard {}
unsafe impl Sync for ProcessMtaGuard {}
pub fn init_process_mta() -> Result<ProcessMtaGuard, WslcError> {
let mut cookie: CO_MTA_USAGE_COOKIE = std::ptr::null_mut();
let hr = unsafe { CoIncrementMTAUsage(&mut cookie) };
if hr < 0 {
Err(WslcError::from_hresult(
hr,
"进程级 CoIncrementMTAUsage 初始化失败",
))
} else if cookie.is_null() {
Err(WslcError::missing_output("CoIncrementMTAUsage", hr))
} else {
static FIRST_LOGGED: std::sync::atomic::AtomicBool =
std::sync::atomic::AtomicBool::new(false);
if !FIRST_LOGGED.swap(true, std::sync::atomic::Ordering::Relaxed) {
log::info!("进程级 COM 多线程套间 (MTA) 首次初始化成功");
} else {
log::debug!("增加进程级 COM 多线程套间 (MTA) 引用计数");
}
Ok(ProcessMtaGuard { cookie })
}
}
#[derive(Debug)]
pub(crate) struct HandleMtaLease {
_mta: ProcessMtaGuard,
}
impl HandleMtaLease {
pub(crate) fn acquire() -> Result<Self, WslcError> {
Ok(Self {
_mta: init_process_mta()?,
})
}
pub(crate) unsafe fn wrap_raw<T, H>(
raw: H,
release: unsafe fn(H),
wrap: impl FnOnce(H, Self) -> T,
) -> Result<T, WslcError>
where
H: Copy,
{
match Self::acquire() {
Ok(lease) => Ok(wrap(raw, lease)),
Err(e) => {
unsafe { release(raw) };
Err(e)
}
}
}
pub(crate) fn degraded() -> Self {
Self {
_mta: ProcessMtaGuard {
cookie: std::ptr::null_mut(),
},
}
}
}
#[derive(Debug)]
pub struct ComGuard {
_init: Rc<ComInitialization>,
}
pub fn try_initialize_mta() -> Result<Option<ComGuard>, WslcError> {
if let Some(existing) = THREAD_MTA.with(|cell| cell.borrow().upgrade()) {
return Ok(Some(ComGuard { _init: existing }));
}
let hr = unsafe { CoInitializeEx(std::ptr::null_mut(), COINIT_MULTITHREADED as u32) };
match hr {
0 | 1 => {
let init = Rc::new(ComInitialization);
THREAD_MTA.with(|cell| *cell.borrow_mut() = Rc::downgrade(&init));
Ok(Some(ComGuard { _init: init }))
}
RPC_E_CHANGED_MODE => {
THREAD_STA_DEGRADED_WARNED.with(|warned| {
if !warned.get() {
warned.set(true);
log::warn!(
"当前线程已被外部宿主初始化为 STA 单线程套间,无法切换为 WSLC 所需的 MTA,\
将复用既有 COM 环境继续执行;若后续调用失败,请将相关操作移至独立线程执行"
);
}
});
Ok(None)
}
code => Err(WslcError::from_hresult(code, "CoInitializeEx 初始化失败")),
}
}
pub fn with_mta<F, R>(f: F) -> Result<R, WslcError>
where
F: FnOnce() -> Result<R, WslcError>,
{
let _guard = try_initialize_mta()?;
f()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_mta_guard_nesting_and_out_of_order_drop() {
let outer = try_initialize_mta()
.expect("COM MTA 初始化失败")
.expect("预期返回守卫");
let inner = try_initialize_mta()
.expect("嵌套 COM MTA 初始化失败")
.expect("预期返回守卫");
drop(outer);
drop(inner);
let reinit = try_initialize_mta()
.expect("重新初始化失败")
.expect("预期返回守卫");
drop(reinit);
}
#[test]
fn test_with_mta_returns_closure_result() {
assert_eq!(with_mta(|| Ok(42)).expect("闭包执行失败"), 42);
let err = with_mta(|| Err::<(), _>(WslcError::InvalidHandle));
assert_eq!(err.unwrap_err(), WslcError::InvalidHandle);
}
#[test]
fn test_process_mta_guard_is_reentrant() {
let first = init_process_mta();
assert!(first.is_ok());
let second = init_process_mta();
assert!(second.is_ok());
drop(second);
drop(first);
}
#[test]
fn test_handle_lease_keeps_mta_alive_independently_of_thread_guard() {
{
let thread_guard = try_initialize_mta()
.expect("线程级 MTA 初始化失败")
.expect("预期返回守卫");
let _ = &thread_guard;
}
let lease = HandleMtaLease::acquire().expect("句柄租约获取失败");
assert!(!lease._mta.cookie.is_null(), "租约应持有有效 cookie");
drop(lease);
let again = HandleMtaLease::acquire().expect("重新获取租约失败");
drop(again);
}
trait CloneDetector {
fn clone_detected(&self) -> bool {
true
}
}
impl<T: Clone> CloneDetector for T {}
trait NotCloneFallback {
fn clone_detected(&self) -> bool {
false
}
}
impl<T> NotCloneFallback for T {}
#[test]
fn test_clone_detector_is_live_and_correct() {
let sample = "probe".to_string();
assert!(
CloneDetector::clone_detected(&sample),
"`String` 实现了 Clone,探测应返回 true;否则探测器已失效,\
`assert_not_clone!` 的判定不可信"
);
assert!(
!NotCloneFallback::clone_detected(&sample),
"兜底实现应对任意类型返回 false"
);
}
macro_rules! assert_not_clone {
($($t:ty => $ctor:expr),+ $(,)?) => {$(
{
fn probe(value: &$t) -> bool {
value.clone_detected()
}
assert!(
!probe(&$ctor),
concat!(stringify!($t), " 绝不可实现 Clone:副本析构会多减一次 MTA 计数,\
提前拆掉套间并使存活句柄失效")
);
}
)+};
}
#[test]
fn test_degraded_lease_is_neutral() {
let degraded = HandleMtaLease::degraded();
assert!(degraded._mta.cookie.is_null());
drop(degraded);
let lease = HandleMtaLease::acquire().expect("获取租约失败");
drop(lease);
}
#[test]
fn test_handles_are_not_cloneable() {
assert_not_clone!(
HandleMtaLease => HandleMtaLease::degraded(),
ProcessMtaGuard => ProcessMtaGuard { cookie: std::ptr::null_mut() },
);
}
#[test]
fn test_lease_acquires_independent_cookie_each_time() {
let first = HandleMtaLease::acquire().expect("获取租约失败");
let second = HandleMtaLease::acquire().expect("再次获取租约失败");
assert_ne!(
first._mta.cookie, second._mta.cookie,
"每次获取租约都应签发独立 cookie,否则计数将被重复持有"
);
drop(first);
drop(second);
}
}