use std::ffi::c_void;
use windows::{
Win32::Graphics::{
Direct3D::{Fxc::*, ID3DBlob, ID3DInclude},
Direct3D11::{
D3D11_CREATE_DEVICE_SINGLETHREADED, ID3D11Device, ID3D11DeviceContext,
ID3D11Multithread,
},
},
core::Interface,
};
use crate::error::D3d11SharedDeviceError;
pub(crate) fn protect_shared_device(
device: &ID3D11Device,
) -> Result<ID3D11DeviceContext, D3d11SharedDeviceError> {
let flags = unsafe { device.GetCreationFlags() };
if flags & D3D11_CREATE_DEVICE_SINGLETHREADED.0 != 0 {
return Err(D3d11SharedDeviceError::SingleThreaded);
}
let context = unsafe { device.GetImmediateContext()? };
let multithread: ID3D11Multithread = context.cast()?;
unsafe {
let _ = multithread.SetMultithreadProtected(true);
if !multithread.GetMultithreadProtected().as_bool() {
return Err(D3d11SharedDeviceError::ProtectionRefused);
}
}
Ok(context)
}
pub(crate) unsafe fn compile_shader(
source: &[u8],
file_name: windows::core::PCSTR,
entry: windows::core::PCSTR,
target: windows::core::PCSTR,
) -> windows::core::Result<ID3DBlob> {
let mut shader = None;
let mut errors = None;
let flags = if cfg!(debug_assertions) {
D3DCOMPILE_DEBUG | D3DCOMPILE_SKIP_OPTIMIZATION
} else {
D3DCOMPILE_OPTIMIZATION_LEVEL3
};
let result = unsafe {
D3DCompile(
source.as_ptr().cast::<c_void>(),
source.len(),
file_name,
None,
None::<&ID3DInclude>,
entry,
target,
flags,
0,
&mut shader,
Some(&mut errors),
)
};
if let Err(error) = result {
let message = errors
.map(|blob| unsafe {
let bytes = std::slice::from_raw_parts(
blob.GetBufferPointer().cast::<u8>(),
blob.GetBufferSize(),
);
String::from_utf8_lossy(bytes).into_owned()
})
.unwrap_or_else(|| error.message());
return Err(windows::core::Error::new(
windows::Win32::Foundation::E_FAIL,
message,
));
}
Ok(shader.expect("D3DCompile succeeded without producing a blob"))
}
#[cfg(test)]
mod tests {
use windows::Win32::Graphics::Direct3D11::ID3D11Multithread;
use super::*;
use crate::test_support::{try_d3d11_device, try_single_threaded_d3d11_device};
#[test]
fn enables_runtime_protection_on_a_multithread_capable_device() {
let Some((device, _context)) = try_d3d11_device() else {
return;
};
let context = protect_shared_device(&device).expect("protect the shared device");
let multithread: ID3D11Multithread = context.cast().expect("multithread interface");
assert!(unsafe { multithread.GetMultithreadProtected() }.as_bool());
protect_shared_device(&device).expect("protect an already protected device");
}
#[test]
fn rejects_a_device_that_promised_single_threaded_use() {
let Some(device) = try_single_threaded_d3d11_device() else {
return;
};
assert!(matches!(
protect_shared_device(&device),
Err(D3d11SharedDeviceError::SingleThreaded)
));
}
}