use anyhow::{Context, Result};
use std::ffi::{CStr, CString};
use std::os::raw::c_char;
use std::ptr;
use std::sync::{Arc, Mutex};
static SLANG_PROCESS_LOCK: Mutex<()> = Mutex::new(());
use super::ffi::*;
use super::loader::SlangLibrary;
use super::virtual_main::effective_slang_source_for_compile;
use crate::types::{OptimizationLevel, ResourceCategory};
use crate::{goldy_event, goldy_span};
pub fn layout_validation_enabled() -> bool {
crate::validation_env::layout_validation_enabled()
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub enum ResourceKind {
Buffer,
MutableBuffer,
Texture,
MutableTexture,
Sampler,
ConstantBuffer,
ParameterBlock,
Other,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct FieldLayout {
pub name: String,
pub offset: usize,
pub size: usize,
pub resource_kind: ResourceKind,
pub type_name: String,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct ParameterBlockLayout {
pub name: String,
pub binding_slot: u32,
pub binding_space: u32,
pub size: usize,
pub alignment: usize,
pub fields: Vec<FieldLayout>,
}
#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
pub struct ShaderReflection {
pub parameter_blocks: Vec<ParameterBlockLayout>,
pub push_constant_categories: Vec<Option<crate::types::ResourceCategory>>,
#[cfg(all(feature = "dx12", target_os = "windows"))]
#[serde(skip)]
pub(crate) push_constant_slot_kinds: Vec<Option<crate::types::BindlessSlotKind>>,
#[serde(default)]
pub binding_element_strides: Vec<Option<u32>>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct StructLayout {
pub name: String,
pub size: usize,
pub alignment: usize,
pub fields: Vec<StructFieldLayout>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct StructFieldLayout {
pub name: String,
pub offset: usize,
pub size: usize,
pub type_name: String,
}
impl StructLayout {
pub fn validate(&self, rust_size: usize, rust_fields: &[(&str, usize, usize)]) -> Result<()> {
let mut errors: Vec<String> = Vec::new();
let mut warnings: Vec<String> = Vec::new();
let slang_data_extent = self.fields.iter().map(|f| f.offset + f.size).max().unwrap_or(0);
if rust_size < slang_data_extent {
errors.push(format!(
"Rust struct ({rust_size} bytes) is smaller than the shader's data extent \
({slang_data_extent} bytes); all shader fields must fit inside the Rust struct"
));
}
for sf in &self.fields {
match rust_fields.iter().find(|&&(name, _, _)| name == sf.name) {
Some(&(_, rust_offset, rust_size_field)) => {
if sf.offset != rust_offset {
errors.push(format!(
"field `{}`: offset Slang {} vs Rust {}",
sf.name, sf.offset, rust_offset
));
}
if sf.size != rust_size_field {
errors.push(format!(
"field `{}`: size Slang {} vs Rust {}",
sf.name, sf.size, rust_size_field
));
}
}
None => {
errors.push(format!(
"field `{}` is declared in the shader but missing from the Rust struct",
sf.name
));
}
}
}
for &(name, _, _) in rust_fields {
if !self.fields.iter().any(|sf| sf.name == name) && !name.starts_with('_') {
warnings.push(format!(
"field `{name}` is in the Rust struct but not in the shader \
(prefix with `_` to suppress this warning)"
));
}
}
if !warnings.is_empty() {
tracing::warn!("Layout check for `{}`: {}", self.name, warnings.join("; "));
}
if errors.is_empty() {
Ok(())
} else {
anyhow::bail!("Struct layout mismatch for `{}`:\n{}", self.name, errors.join("\n"));
}
}
}
#[derive(Debug, Clone, Copy)]
pub struct LayoutCheck<'a> {
pub type_name: &'a str,
pub rust_size: usize,
pub rust_fields: &'a [(&'a str, usize, usize)],
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct OwnedLayoutCheck {
pub type_name: String,
pub rust_size: usize,
pub rust_fields: Vec<(String, usize, usize)>,
}
impl OwnedLayoutCheck {
pub fn from_layout_check(c: &LayoutCheck<'_>) -> Self {
Self {
type_name: c.type_name.to_string(),
rust_size: c.rust_size,
rust_fields: c
.rust_fields
.iter()
.map(|(n, o, s)| ((*n).to_string(), *o, *s))
.collect(),
}
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct CompiledShaderWithReflection {
pub shader: CompiledShader,
pub reflection: ShaderReflection,
}
#[repr(u8)]
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, serde::Serialize, serde::Deserialize)]
pub enum ShaderTarget {
Spirv,
Dxil,
Metal,
Wgsl,
Ptx,
}
impl ShaderTarget {
fn to_slang_target(self) -> SlangCompileTarget {
match self {
ShaderTarget::Spirv => SlangCompileTarget::Spirv,
ShaderTarget::Dxil => SlangCompileTarget::Dxil,
ShaderTarget::Metal => SlangCompileTarget::Metal,
ShaderTarget::Wgsl => SlangCompileTarget::Wgsl,
ShaderTarget::Ptx => SlangCompileTarget::Ptx,
}
}
pub fn is_binary(self) -> bool {
matches!(self, ShaderTarget::Spirv | ShaderTarget::Dxil)
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct CompiledShader {
pub data: Vec<u8>,
pub target: ShaderTarget,
}
impl CompiledShader {
pub fn as_str(&self) -> Option<&str> {
match self.target {
ShaderTarget::Metal | ShaderTarget::Wgsl | ShaderTarget::Ptx => std::str::from_utf8(&self.data).ok(),
ShaderTarget::Spirv | ShaderTarget::Dxil => None,
}
}
pub fn as_spirv(&self) -> Option<&[u32]> {
if self.target == ShaderTarget::Spirv && self.data.len().is_multiple_of(4) {
Some(bytemuck::cast_slice(&self.data))
} else {
None
}
}
pub fn as_dxil(&self) -> Option<&[u8]> {
if self.target == ShaderTarget::Dxil {
Some(&self.data)
} else {
None
}
}
}
pub fn builtin_type_stride(name: &str) -> Option<u32> {
match name {
"uint" | "int" | "float" | "bool" | "dword" => Some(4),
"half" | "float16_t" => Some(2),
"double" | "uint64_t" | "int64_t" => Some(8),
"uint2" | "int2" | "float2" => Some(8),
"half2" => Some(4),
"uint3" | "int3" | "float3" => Some(12),
"half3" => Some(6),
"uint4" | "int4" | "float4" => Some(16),
"half4" => Some(8),
"float2x2" => Some(16),
"float3x3" => Some(36),
"float4x4" => Some(64),
"DispatchShape" => Some(12),
_ => None,
}
}
pub struct SlangCompiler {
library: Arc<SlangLibrary>,
global_session: *mut IGlobalSession,
shader_disk_cache: std::sync::Mutex<crate::shader_cache::ShaderBytecodeDiskCache>,
}
unsafe impl Send for SlangCompiler {}
unsafe impl Sync for SlangCompiler {}
impl SlangCompiler {
pub fn new() -> Result<Self> {
let _span = goldy_span!("slang.compiler.init").entered();
let _guard = SLANG_PROCESS_LOCK.lock().unwrap();
let library = Arc::new(SlangLibrary::load()?);
let mut global_session: *mut IGlobalSession = ptr::null_mut();
let global_desc = SlangGlobalSessionDesc::default();
tracing::debug!(
"Creating global session with desc size: {}",
std::mem::size_of::<SlangGlobalSessionDesc>()
);
let result = unsafe { (library.create_global_session)(&global_desc, &mut global_session) };
if !slang_succeeded(result) || global_session.is_null() {
anyhow::bail!(
"Failed to create Slang global session (result={}, ptr={:?})",
result,
global_session
);
}
tracing::debug!("Global session created: {:?}", global_session);
goldy_event!("slang.session.create", success = true);
tracing::info!("Slang compiler initialized");
Ok(Self {
library,
global_session,
shader_disk_cache: std::sync::Mutex::new(crate::shader_cache::ShaderBytecodeDiskCache::new_load_or_empty()),
})
}
pub fn compile_bindless_with_reflection(
&self,
source: &str,
target: ShaderTarget,
entry_points: &[(&str, SlangStage)],
search_paths: &[&str],
) -> Result<CompiledShaderWithReflection> {
self.compile_bindless_with_reflection_and_defines(
source,
target,
entry_points,
search_paths,
&[],
&[],
OptimizationLevel::Default,
)
}
#[allow(clippy::too_many_arguments)] pub fn compile_bindless_with_reflection_and_defines(
&self,
source: &str,
target: ShaderTarget,
entry_points: &[(&str, SlangStage)],
search_paths: &[&str],
extra_defines: &[(&str, &str)],
layout_checks: &[OwnedLayoutCheck],
optimization_level: OptimizationLevel,
) -> Result<CompiledShaderWithReflection> {
let mut defines = Self::bindless_defines_for_target(target);
defines.extend_from_slice(extra_defines);
self.compile_with_reflection(
source,
target,
entry_points,
search_paths,
&defines,
layout_checks,
optimization_level,
)
}
fn bindless_defines_for_target(target: ShaderTarget) -> Vec<(&'static str, &'static str)> {
match target {
ShaderTarget::Spirv => vec![("__SPIRV__", "1")],
ShaderTarget::Dxil => vec![("__DX12__", "1")],
ShaderTarget::Metal => vec![("__METAL__", "1")],
ShaderTarget::Wgsl => vec![("__WGSL__", "1")],
ShaderTarget::Ptx => vec![("__CUDA__", "1")],
}
}
#[allow(clippy::too_many_arguments)] fn with_compiled_request<R>(
&self,
source: &str,
target: ShaderTarget,
entry_points: &[(&str, SlangStage)],
search_paths: &[&str],
defines: &[(&str, &str)],
optimization_level: OptimizationLevel,
f: impl FnOnce(&Self, *mut SlangCompileRequest, i32) -> Result<R>,
) -> Result<R> {
let _guard = SLANG_PROCESS_LOCK.lock().unwrap();
let define_names: Vec<CString> = defines.iter().map(|(k, _)| CString::new(*k).unwrap()).collect();
let define_values: Vec<CString> = defines.iter().map(|(_, v)| CString::new(*v).unwrap()).collect();
let macro_descs: Vec<PreprocessorMacroDesc> = define_names
.iter()
.zip(define_values.iter())
.map(|(name, value)| PreprocessorMacroDesc {
name: name.as_ptr(),
value: value.as_ptr(),
})
.collect();
let search_path_cstrings: Vec<CString> = search_paths.iter().map(|p| CString::new(*p).unwrap()).collect();
let search_path_ptrs: Vec<*const c_char> = search_path_cstrings.iter().map(|s| s.as_ptr()).collect();
let mut session_desc = SessionDesc::default();
if !search_path_ptrs.is_empty() {
session_desc.search_paths = search_path_ptrs.as_ptr();
session_desc.search_path_count = search_path_ptrs.len() as i64;
}
if !macro_descs.is_empty() {
session_desc.preprocessor_macros = macro_descs.as_ptr();
session_desc.preprocessor_macro_count = macro_descs.len() as i64;
}
tracing::debug!(
"Creating session with {} macros, SessionDesc size: {}",
macro_descs.len(),
std::mem::size_of::<SessionDesc>()
);
let mut session: *mut ISession = ptr::null_mut();
let result = unsafe { global_session_create_session(self.global_session, &session_desc, &mut session) };
if !slang_succeeded(result) || session.is_null() {
anyhow::bail!(
"Failed to create Slang session with preprocessor defines (result={}, ptr={:?})",
result,
session
);
}
tracing::debug!("Session with defines created: {:?}", session);
let _session_guard = scopeguard::guard(session, |s| {
unsafe { session_release(s) };
});
let mut request: *mut SlangCompileRequest = ptr::null_mut();
let result = unsafe { session_create_compile_request(session, &mut request) };
if !slang_succeeded(result) || request.is_null() {
anyhow::bail!(
"Failed to create Slang compile request (result={}, ptr={:?})",
result,
request
);
}
tracing::debug!("Compile request created: {:?}", request);
let library = self.library.clone();
let _guard = scopeguard::guard(request, |req| {
unsafe { (library.destroy_compile_request)(req) };
});
let target_index = unsafe { (self.library.add_code_gen_target)(request, target.to_slang_target() as i32) };
if target_index < 0 {
anyhow::bail!("Failed to add code generation target");
}
if target == ShaderTarget::Dxil {
let profile_name = CString::new("sm_6_6").unwrap();
let profile_id = unsafe { global_session_find_profile(self.global_session, profile_name.as_ptr()) };
if profile_id > 0 {
unsafe {
(self.library.set_target_profile)(request, target_index, profile_id);
}
tracing::debug!("Set DXIL target profile to sm_6_6 (id={})", profile_id);
} else {
tracing::warn!("Could not find sm_6_6 profile, using default");
}
unsafe {
(self.library.set_target_floating_point_mode)(request, target_index, SLANG_FLOATING_POINT_MODE_PRECISE);
}
}
let unit_name = CString::new("shader").unwrap();
let translation_unit = unsafe {
(self.library.add_translation_unit)(request, SlangSourceLanguage::Slang as i32, unit_name.as_ptr())
};
if translation_unit < 0 {
anyhow::bail!("Failed to add translation unit");
}
let source_path = CString::new("shader.slang").unwrap();
let source_cstr = CString::new(source).context("Source contains null bytes")?;
unsafe {
(self.library.add_translation_unit_source_string)(
request,
translation_unit,
source_path.as_ptr(),
source_cstr.as_ptr(),
);
}
for (name, stage) in entry_points {
let name_cstr = CString::new(*name).context("Entry point name contains null bytes")?;
let entry_index =
unsafe { (self.library.add_entry_point)(request, translation_unit, name_cstr.as_ptr(), *stage as i32) };
if entry_index < 0 {
anyhow::bail!("Failed to add entry point: {}", name);
}
}
if optimization_level != OptimizationLevel::Default {
let ffi_level = match optimization_level {
OptimizationLevel::None => SLANG_OPTIMIZATION_LEVEL_NONE,
OptimizationLevel::Default => unreachable!(),
OptimizationLevel::High => SLANG_OPTIMIZATION_LEVEL_HIGH,
OptimizationLevel::Maximal => SLANG_OPTIMIZATION_LEVEL_MAXIMAL,
};
unsafe { (self.library.set_optimization_level)(request, ffi_level) };
tracing::info!("Slang optimization level set to {:?}", optimization_level);
}
let result = unsafe { (self.library.compile)(request) };
if !slang_succeeded(result) {
let diag_ptr = unsafe { (self.library.get_diagnostic_output)(request) };
let diagnostic = if !diag_ptr.is_null() {
unsafe { CStr::from_ptr(diag_ptr) }.to_string_lossy().into_owned()
} else {
"Unknown compilation error".to_string()
};
anyhow::bail!("Slang compilation failed:\n{}", diagnostic);
}
f(self, request, target_index)
}
#[allow(clippy::too_many_arguments)] pub fn compile_with_reflection(
&self,
source: &str,
target: ShaderTarget,
entry_points: &[(&str, SlangStage)],
search_paths: &[&str],
defines: &[(&str, &str)],
layout_checks: &[OwnedLayoutCheck],
optimization_level: OptimizationLevel,
) -> Result<CompiledShaderWithReflection> {
let effective = effective_slang_source_for_compile(source);
let cache_key = crate::shader_cache::compile_cache_key(
effective.as_ref(),
target,
entry_points,
search_paths,
defines,
layout_checks,
optimization_level,
);
{
let mut disk = self.shader_disk_cache.lock().unwrap_or_else(|p| p.into_inner());
if let Some(hit) = disk.get(cache_key) {
return hit.with_context(|| "decode shader disk cache");
}
}
let binding_type_names = super::virtual_main::extract_binding_element_type_names(source);
let binding_categories = super::virtual_main::extract_push_constant_categories(source);
let out = self.with_compiled_request(
effective.as_ref(),
target,
entry_points,
search_paths,
defines,
optimization_level,
|slf, request, target_index| {
let mut blob: *mut ISlangBlob = ptr::null_mut();
let result = unsafe { (slf.library.get_entry_point_code_blob)(request, 0, target_index, &mut blob) };
if !slang_succeeded(result) || blob.is_null() {
anyhow::bail!("Failed to get compiled shader code");
}
let (data_ptr, data_size) = unsafe { blob_get_data(blob) };
let data = unsafe { std::slice::from_raw_parts(data_ptr, data_size) }.to_vec();
unsafe { blob_release(blob) };
let mut reflection = slf.extract_reflection(request)?;
if !layout_checks.is_empty() {
slf.validate_owned_layout_checks(request, layout_checks)?;
}
let strides: Vec<Option<u32>> = binding_type_names
.iter()
.enumerate()
.map(|(i, opt_name)| {
let cat = binding_categories.get(i).copied().unwrap_or(None);
opt_name.as_deref().and_then(|name| {
builtin_type_stride(name).or_else(|| slf.reflect_binding_element_stride(request, name, cat))
})
})
.collect();
reflection.binding_element_strides = strides;
Ok(CompiledShaderWithReflection {
shader: CompiledShader { data, target },
reflection,
})
},
)?;
{
let mut disk = self.shader_disk_cache.lock().unwrap_or_else(|p| p.into_inner());
if let Err(e) = disk.insert(cache_key, &out) {
tracing::warn!(?e, "failed to serialize shader disk cache entry");
}
}
Ok(out)
}
fn validate_owned_layout_checks(
&self,
request: *mut SlangCompileRequest,
checks: &[OwnedLayoutCheck],
) -> Result<()> {
for owned in checks {
let layout = self.reflect_named_struct_from_request(request, &owned.type_name)?;
let field_refs: Vec<(&str, usize, usize)> =
owned.rust_fields.iter().map(|(n, o, s)| (n.as_str(), *o, *s)).collect();
layout.validate(owned.rust_size, &field_refs)?;
}
Ok(())
}
pub fn reflect_struct_layout(
&self,
shader_source: &str,
target: ShaderTarget,
search_paths: &[&str],
type_name: &str,
) -> Result<StructLayout> {
let defines = Self::bindless_defines_for_target(target);
let entry_points = &[("vs_main", SlangStage::Vertex)];
let effective = effective_slang_source_for_compile(shader_source);
self.with_compiled_request(
effective.as_ref(),
target,
entry_points,
search_paths,
&defines,
OptimizationLevel::Default,
|slf, request, _target_index| slf.reflect_named_struct_from_request(request, type_name),
)
}
fn reflect_named_struct_from_request(
&self,
request: *mut SlangCompileRequest,
type_name: &str,
) -> Result<StructLayout> {
let reflection_ptr = unsafe { (self.library.get_reflection)(request) };
if reflection_ptr.is_null() {
anyhow::bail!("No Slang reflection available after compile");
}
let name_cstr = CString::new(type_name).context("type_name contains null bytes")?;
let ty = unsafe { (self.library.reflection_find_type_by_name)(reflection_ptr, name_cstr.as_ptr()) };
if ty.is_null() {
anyhow::bail!("Slang reflection: type `{type_name}` not found");
}
let layout_ptr =
unsafe { (self.library.reflection_get_type_layout)(reflection_ptr, ty, SlangLayoutRules::Default) };
if layout_ptr.is_null() {
anyhow::bail!("Slang reflection: failed to get layout for `{type_name}`");
}
self.extract_struct_layout_uniform(layout_ptr, type_name)
}
fn reflect_binding_element_stride(
&self,
request: *mut SlangCompileRequest,
type_name: &str,
category: Option<ResourceCategory>,
) -> Option<u32> {
if matches!(category, Some(ResourceCategory::Broadcast)) {
return self.reflect_struct_storage_stride(request, type_name, SlangParameterCategory::Uniform);
}
let layout_cat = match category {
Some(ResourceCategory::StorageImage) => SlangParameterCategory::UnorderedAccess,
Some(ResourceCategory::Scattered)
| Some(ResourceCategory::Texture)
| Some(ResourceCategory::Sampler)
| None => SlangParameterCategory::ShaderResource,
Some(ResourceCategory::Broadcast) => unreachable!("handled above"),
};
self.reflect_type_size_with_category(request, type_name, layout_cat)
.or_else(|| {
if matches!(
category,
Some(ResourceCategory::Scattered)
| Some(ResourceCategory::StorageImage)
| Some(ResourceCategory::Texture)
| None
) {
self.reflect_struct_storage_stride(request, type_name, layout_cat)
} else {
None
}
})
}
fn reflect_type_size_with_category(
&self,
request: *mut SlangCompileRequest,
type_name: &str,
layout_cat: SlangParameterCategory,
) -> Option<u32> {
let layout_ptr = self.reflect_type_layout_ptr(request, type_name)?;
let size = unsafe { (self.library.reflection_type_layout_get_size)(layout_ptr, layout_cat as i32) } as u32;
if size > 0 {
Some(size)
} else {
None
}
}
fn reflect_struct_storage_stride(
&self,
request: *mut SlangCompileRequest,
type_name: &str,
_layout_cat: SlangParameterCategory,
) -> Option<u32> {
let layout_ptr = self.reflect_type_layout_ptr(request, type_name)?;
let field_count = unsafe { (self.library.reflection_type_layout_get_field_count)(layout_ptr) };
if field_count == 0 {
return None;
}
let field_cat = SlangParameterCategory::Uniform as i32;
let mut extent = 0u32;
for i in 0..field_count {
let field_var = unsafe { (self.library.reflection_type_layout_get_field_by_index)(layout_ptr, i) };
if field_var.is_null() {
continue;
}
let field_type_layout = unsafe { (self.library.reflection_variable_layout_get_type_layout)(field_var) };
if field_type_layout.is_null() {
continue;
}
let offset = unsafe { (self.library.reflection_variable_layout_get_offset)(field_var, field_cat) } as u32;
let field_size =
unsafe { (self.library.reflection_type_layout_get_size)(field_type_layout, field_cat) } as u32;
let field_extent = offset.saturating_add(field_size.max(1));
extent = extent.max(field_extent);
}
if extent > 0 {
Some(extent)
} else {
None
}
}
fn reflect_type_layout_ptr(
&self,
request: *mut SlangCompileRequest,
type_name: &str,
) -> Option<*mut SlangReflectionTypeLayout> {
let reflection_ptr = unsafe { (self.library.get_reflection)(request) };
if reflection_ptr.is_null() {
return None;
}
let mut candidates = vec![type_name.to_string()];
if !type_name.contains('.') {
candidates.push(format!("shader.{type_name}"));
}
for candidate in &candidates {
let name_cstr = CString::new(candidate.as_str()).ok()?;
let ty = unsafe { (self.library.reflection_find_type_by_name)(reflection_ptr, name_cstr.as_ptr()) };
if ty.is_null() {
continue;
}
let layout_ptr =
unsafe { (self.library.reflection_get_type_layout)(reflection_ptr, ty, SlangLayoutRules::Default) };
if !layout_ptr.is_null() {
return Some(layout_ptr);
}
}
None
}
fn extract_struct_layout_uniform(
&self,
type_layout: *mut SlangReflectionTypeLayout,
struct_name: &str,
) -> Result<StructLayout> {
let cat = SlangParameterCategory::Uniform as i32;
let size = unsafe { (self.library.reflection_type_layout_get_size)(type_layout, cat) };
let alignment = unsafe { (self.library.reflection_type_layout_get_alignment)(type_layout, cat) };
let field_count = unsafe { (self.library.reflection_type_layout_get_field_count)(type_layout) };
let mut fields = Vec::new();
for i in 0..field_count {
let field_var = unsafe { (self.library.reflection_type_layout_get_field_by_index)(type_layout, i) };
if field_var.is_null() {
continue;
}
let variable = unsafe { (self.library.reflection_variable_layout_get_variable)(field_var) };
let name = if !variable.is_null() {
let name_ptr = unsafe { (self.library.reflection_variable_get_name)(variable) };
if !name_ptr.is_null() {
unsafe { CStr::from_ptr(name_ptr) }.to_string_lossy().into_owned()
} else {
format!("field_{i}")
}
} else {
format!("field_{i}")
};
let field_type_layout = unsafe { (self.library.reflection_variable_layout_get_type_layout)(field_var) };
if field_type_layout.is_null() {
continue;
}
let offset = unsafe { (self.library.reflection_variable_layout_get_offset)(field_var, cat) };
let fsize = unsafe { (self.library.reflection_type_layout_get_size)(field_type_layout, cat) };
let field_type = unsafe { (self.library.reflection_type_layout_get_type)(field_type_layout) };
let type_name = if !field_type.is_null() {
let type_name_ptr = unsafe { (self.library.reflection_type_get_name)(field_type) };
if !type_name_ptr.is_null() {
unsafe { CStr::from_ptr(type_name_ptr) }.to_string_lossy().into_owned()
} else {
String::new()
}
} else {
String::new()
};
fields.push(StructFieldLayout {
name,
offset,
size: fsize,
type_name,
});
}
Ok(StructLayout {
name: struct_name.to_string(),
size,
alignment,
fields,
})
}
fn extract_reflection(&self, request: *mut SlangCompileRequest) -> Result<ShaderReflection> {
let _span = goldy_span!("slang.reflection.extract").entered();
let reflection_ptr = unsafe { (self.library.get_reflection)(request) };
if reflection_ptr.is_null() {
return Ok(ShaderReflection::default());
}
let mut parameter_blocks = Vec::new();
let param_count = unsafe { (self.library.reflection_get_parameter_count)(reflection_ptr) };
for i in 0..param_count {
let param = unsafe { (self.library.reflection_get_parameter_by_index)(reflection_ptr, i) };
if param.is_null() {
continue;
}
let variable = unsafe { (self.library.reflection_variable_layout_get_variable)(param) };
let name = if !variable.is_null() {
let name_ptr = unsafe { (self.library.reflection_variable_get_name)(variable) };
if !name_ptr.is_null() {
unsafe { CStr::from_ptr(name_ptr) }.to_string_lossy().into_owned()
} else {
format!("param_{}", i)
}
} else {
format!("param_{}", i)
};
let type_layout = unsafe { (self.library.reflection_parameter_get_type_layout)(param) };
if type_layout.is_null() {
continue;
}
let type_ptr = unsafe { (self.library.reflection_type_layout_get_type)(type_layout) };
if type_ptr.is_null() {
continue;
}
let type_kind = unsafe { (self.library.reflection_type_get_kind)(type_ptr) };
if type_kind == SlangTypeKind::ParameterBlock as i32 {
let block_layout = self.extract_parameter_block_layout(param, type_layout, &name)?;
parameter_blocks.push(block_layout);
}
}
goldy_event!(
"slang.reflection.extract",
parameter_blocks = parameter_blocks.len(),
total_fields = parameter_blocks.iter().map(|pb| pb.fields.len()).sum::<usize>()
);
Ok(ShaderReflection {
parameter_blocks,
push_constant_categories: Vec::new(),
#[cfg(all(feature = "dx12", target_os = "windows"))]
push_constant_slot_kinds: Vec::new(),
binding_element_strides: Vec::new(),
})
}
fn extract_parameter_block_layout(
&self,
param: *mut SlangReflectionParameter,
type_layout: *mut SlangReflectionTypeLayout,
name: &str,
) -> Result<ParameterBlockLayout> {
let binding_slot = unsafe { (self.library.reflection_parameter_get_binding_index)(param) } as u32;
let binding_space = unsafe { (self.library.reflection_parameter_get_binding_space)(param) } as u32;
let element_type_layout = unsafe { (self.library.reflection_type_layout_get_element_type_layout)(type_layout) };
const SLOT_SIZE_BYTES: usize = 8;
let (mut size, alignment, fields) = if !element_type_layout.is_null() {
let size_slots = unsafe {
(self.library.reflection_type_layout_get_size)(
element_type_layout,
SlangParameterCategory::MetalArgumentBufferElement as i32,
)
};
let alignment = unsafe {
(self.library.reflection_type_layout_get_alignment)(
element_type_layout,
SlangParameterCategory::MetalArgumentBufferElement as i32,
)
};
let fields = self.extract_struct_fields(element_type_layout)?;
let size = size_slots * SLOT_SIZE_BYTES;
(size, alignment, fields)
} else {
let size = unsafe {
(self.library.reflection_type_layout_get_size)(type_layout, SlangParameterCategory::Uniform as i32)
};
let alignment = unsafe {
(self.library.reflection_type_layout_get_alignment)(type_layout, SlangParameterCategory::Uniform as i32)
};
(size, alignment, Vec::new())
};
if size == 0 && !fields.is_empty() {
size = fields.iter().map(|f| f.offset + f.size).max().unwrap_or(0);
}
let alignment_bytes = if alignment > 0 {
alignment * SLOT_SIZE_BYTES
} else {
SLOT_SIZE_BYTES };
Ok(ParameterBlockLayout {
name: name.to_string(),
binding_slot,
binding_space,
size,
alignment: alignment_bytes,
fields,
})
}
fn extract_struct_fields(&self, type_layout: *mut SlangReflectionTypeLayout) -> Result<Vec<FieldLayout>> {
let mut fields = Vec::new();
let field_count = unsafe { (self.library.reflection_type_layout_get_field_count)(type_layout) };
for i in 0..field_count {
let field_var = unsafe { (self.library.reflection_type_layout_get_field_by_index)(type_layout, i) };
if field_var.is_null() {
continue;
}
let variable = unsafe { (self.library.reflection_variable_layout_get_variable)(field_var) };
let name = if !variable.is_null() {
let name_ptr = unsafe { (self.library.reflection_variable_get_name)(variable) };
if !name_ptr.is_null() {
unsafe { CStr::from_ptr(name_ptr) }.to_string_lossy().into_owned()
} else {
format!("field_{}", i)
}
} else {
format!("field_{}", i)
};
let field_type_layout = unsafe { (self.library.reflection_variable_layout_get_type_layout)(field_var) };
if field_type_layout.is_null() {
continue;
}
let resource_kind = self.determine_resource_kind(field_type_layout);
let offset_slots = unsafe {
(self.library.reflection_variable_layout_get_offset)(
field_var,
SlangParameterCategory::MetalArgumentBufferElement as i32,
)
};
let size_slots = unsafe {
(self.library.reflection_type_layout_get_size)(
field_type_layout,
SlangParameterCategory::MetalArgumentBufferElement as i32,
)
};
const SLOT_SIZE_BYTES: usize = 8;
let offset = offset_slots * SLOT_SIZE_BYTES;
let size = if size_slots > 0 {
size_slots * SLOT_SIZE_BYTES
} else {
SLOT_SIZE_BYTES
};
tracing::trace!(
"Field {} (index {}): offset_slots={}, size_slots={} -> offset={}, size={}, resource_kind={:?}",
name,
i,
offset_slots,
size_slots,
offset,
size,
resource_kind
);
let field_type = unsafe { (self.library.reflection_type_layout_get_type)(field_type_layout) };
let type_name = if !field_type.is_null() {
let type_name_ptr = unsafe { (self.library.reflection_type_get_name)(field_type) };
if !type_name_ptr.is_null() {
unsafe { CStr::from_ptr(type_name_ptr) }.to_string_lossy().into_owned()
} else {
String::new()
}
} else {
String::new()
};
fields.push(FieldLayout {
name,
offset,
size,
resource_kind,
type_name,
});
}
Ok(fields)
}
fn determine_resource_kind(&self, type_layout: *mut SlangReflectionTypeLayout) -> ResourceKind {
let type_ptr = unsafe { (self.library.reflection_type_layout_get_type)(type_layout) };
if type_ptr.is_null() {
return ResourceKind::Other;
}
let type_kind = unsafe { (self.library.reflection_type_get_kind)(type_ptr) };
let binding_type = unsafe { (self.library.reflection_type_layout_get_binding_type)(type_layout) };
tracing::trace!(
"determine_resource_kind: type_kind={}, binding_type={}",
type_kind,
binding_type
);
match type_kind {
k if k == SlangTypeKind::SamplerState as i32 => ResourceKind::Sampler,
k if k == SlangTypeKind::ConstantBuffer as i32 => ResourceKind::ConstantBuffer,
k if k == SlangTypeKind::ParameterBlock as i32 => ResourceKind::ParameterBlock,
k if k == SlangTypeKind::Resource as i32 => {
match binding_type {
b if b == SlangBindingType::Texture as i32 => ResourceKind::Texture,
b if b == SlangBindingType::MutableTexture as i32 => ResourceKind::MutableTexture,
b if b == SlangBindingType::TypedBuffer as i32 => ResourceKind::Buffer,
b if b == SlangBindingType::MutableTypedBuffer as i32 => ResourceKind::MutableBuffer,
b if b == SlangBindingType::RawBuffer as i32 => ResourceKind::Buffer,
b if b == SlangBindingType::MutableRawBuffer as i32 => ResourceKind::MutableBuffer,
_ => ResourceKind::Other,
}
}
k if k == SlangTypeKind::ShaderStorageBuffer as i32 => ResourceKind::MutableBuffer,
_ => {
match binding_type {
b if b == SlangBindingType::TypedBuffer as i32 => ResourceKind::Buffer,
b if b == SlangBindingType::MutableTypedBuffer as i32 => ResourceKind::MutableBuffer,
b if b == SlangBindingType::RawBuffer as i32 => ResourceKind::Buffer,
b if b == SlangBindingType::MutableRawBuffer as i32 => ResourceKind::MutableBuffer,
b if b == SlangBindingType::Texture as i32 => ResourceKind::Texture,
b if b == SlangBindingType::MutableTexture as i32 => ResourceKind::MutableTexture,
b if b == SlangBindingType::Sampler as i32 => ResourceKind::Sampler,
b if b == SlangBindingType::ConstantBuffer as i32 => ResourceKind::ConstantBuffer,
_ => ResourceKind::Other,
}
}
}
}
}
#[cfg(test)]
mod builtin_stride_tests {
use super::builtin_type_stride;
#[test]
fn scalar_types() {
assert_eq!(builtin_type_stride("uint"), Some(4));
assert_eq!(builtin_type_stride("int"), Some(4));
assert_eq!(builtin_type_stride("float"), Some(4));
assert_eq!(builtin_type_stride("half"), Some(2));
assert_eq!(builtin_type_stride("double"), Some(8));
}
#[test]
fn vector_types() {
assert_eq!(builtin_type_stride("float2"), Some(8));
assert_eq!(builtin_type_stride("float3"), Some(12));
assert_eq!(builtin_type_stride("float4"), Some(16));
assert_eq!(builtin_type_stride("uint4"), Some(16));
}
#[test]
fn matrix_types() {
assert_eq!(builtin_type_stride("float4x4"), Some(64));
}
#[test]
fn user_struct_returns_none() {
assert_eq!(builtin_type_stride("MyStruct"), None);
assert_eq!(builtin_type_stride("Particle"), None);
}
#[test]
fn dispatch_shape_stride() {
assert_eq!(builtin_type_stride("DispatchShape"), Some(12));
}
}
#[cfg(test)]
mod struct_layout_validate_tests {
use super::{StructFieldLayout, StructLayout};
use crate as goldy;
fn two_float_layout() -> StructLayout {
StructLayout {
name: "S".into(),
size: 8,
alignment: 4,
fields: vec![
StructFieldLayout {
name: "a".into(),
offset: 0,
size: 4,
type_name: "float".into(),
},
StructFieldLayout {
name: "b".into(),
offset: 4,
size: 4,
type_name: "float".into(),
},
],
}
}
fn layout_time_only_cb_padded() -> StructLayout {
StructLayout {
name: "TimeUniforms".into(),
size: 16,
alignment: 16,
fields: vec![StructFieldLayout {
name: "time".into(),
offset: 0,
size: 4,
type_name: "float".into(),
}],
}
}
#[test]
fn validate_cb_padded_single_field_passes() {
let slang = layout_time_only_cb_padded();
let rust_fields = [("time", 0usize, 4usize)];
slang
.validate(4, &rust_fields)
.expect("Rust 4-byte struct should cover Slang data extent");
}
#[test]
fn validate_shader_field_missing_from_rust_errors() {
let slang = StructLayout {
name: "U".into(),
size: 8,
alignment: 4,
fields: vec![
StructFieldLayout {
name: "time".into(),
offset: 0,
size: 4,
type_name: "float".into(),
},
StructFieldLayout {
name: "brightness".into(),
offset: 4,
size: 4,
type_name: "float".into(),
},
],
};
let rust_fields = [("time", 0usize, 4usize)];
let err = slang.validate(4, &rust_fields).unwrap_err();
let s = err.to_string();
assert!(
s.contains("brightness") && s.contains("missing"),
"expected missing-field error, got: {s}"
);
}
#[test]
fn validate_rust_too_small_for_data_extent_errors() {
let slang = layout_time_only_cb_padded();
let rust_fields = [("time", 0usize, 4usize)];
let err = slang.validate(2, &rust_fields).unwrap_err();
assert!(
err.to_string().contains("smaller than the shader's data extent"),
"{err}"
);
}
#[test]
fn validate_extra_rust_field_without_underscore_passes() {
let slang = layout_time_only_cb_padded();
let rust_fields = [("time", 0usize, 4usize), ("brightness", 4usize, 4usize)];
slang
.validate(8, &rust_fields)
.expect("extra Rust field is not an error");
}
#[test]
fn validate_extra_rust_field_with_underscore_passes() {
let slang = layout_time_only_cb_padded();
let rust_fields = [("time", 0usize, 4usize), ("_pad0", 4usize, 4usize)];
slang
.validate(8, &rust_fields)
.expect("_prefixed extra field is silent");
}
#[test]
fn validate_ok_when_matching() {
two_float_layout().validate(8, &[("a", 0, 4), ("b", 4, 4)]).unwrap();
}
#[test]
fn validate_does_not_require_rust_struct_to_match_slang_cb_padding() {
let mut layout = two_float_layout();
layout.size = 16;
layout
.validate(8, &[("a", 0, 4), ("b", 4, 4)])
.expect("Slang padded size must not force Rust to pad");
}
#[test]
fn validate_err_on_field_count_mismatch() {
let err = two_float_layout().validate(8, &[("a", 0, 4)]).unwrap_err().to_string();
assert!(
err.contains("`b`") && err.contains("missing"),
"expected shader field b missing in Rust: {err}"
);
}
#[test]
fn validate_allows_extra_rust_fields_not_in_shader() {
two_float_layout()
.validate(12, &[("a", 0, 4), ("b", 4, 4), ("c", 8, 4)])
.expect("extra Rust-only field should not fail validation");
}
#[test]
fn validate_err_on_field_offset_mismatch() {
let err = two_float_layout()
.validate(8, &[("a", 0, 4), ("b", 0, 4)])
.unwrap_err()
.to_string();
assert!(err.contains("offset"), "expected offset mismatch: {err}");
assert!(err.contains("`b`"), "expected field name b: {err}");
}
#[test]
fn validate_err_on_field_size_mismatch() {
let err = two_float_layout()
.validate(8, &[("a", 0, 8), ("b", 4, 4)])
.unwrap_err()
.to_string();
assert!(
err.contains("size") && err.contains("`a`"),
"expected field size mismatch for a: {err}"
);
}
#[test]
fn validate_err_on_field_name_mismatch() {
let err = two_float_layout()
.validate(8, &[("x", 0, 4), ("b", 4, 4)])
.unwrap_err()
.to_string();
assert!(
err.contains("`a`") && err.contains("missing"),
"expected shader field `a` missing from Rust (got `x` instead): {err}"
);
}
#[test]
fn validate_reports_multiple_errors() {
let err = two_float_layout().validate(8, &[("a", 4, 4)]).unwrap_err().to_string();
assert!(
err.contains("offset") && err.contains("`a`"),
"expected offset error for a: {err}"
);
assert!(
err.contains("`b`") && err.contains("missing"),
"expected missing b: {err}"
);
}
#[test]
fn layout_checkable_derive_generates_correct_const() {
#[derive(goldy_derive::LayoutCheckable)]
#[repr(C)]
struct TestStruct {
pos: [f32; 2],
color: [f32; 4],
}
let check = TestStruct::LAYOUT_CHECK;
assert_eq!(check.type_name, "TestStruct");
assert_eq!(check.rust_size, std::mem::size_of::<TestStruct>());
assert_eq!(check.rust_fields.len(), 2);
let (name, offset, size) = check.rust_fields[0];
assert_eq!(name, "pos");
assert_eq!(offset, 0);
assert_eq!(size, std::mem::size_of::<[f32; 2]>());
let (name, offset, size) = check.rust_fields[1];
assert_eq!(name, "color");
assert_eq!(offset, std::mem::size_of::<[f32; 2]>());
assert_eq!(size, std::mem::size_of::<[f32; 4]>());
}
#[test]
fn layout_check_validates_against_matching_slang_layout() {
#[derive(goldy_derive::LayoutCheckable)]
#[repr(C)]
struct Uniforms {
x: f32,
y: f32,
}
let slang = StructLayout {
name: "Uniforms".into(),
size: 8,
alignment: 4,
fields: vec![
StructFieldLayout {
name: "x".into(),
offset: 0,
size: 4,
type_name: "float".into(),
},
StructFieldLayout {
name: "y".into(),
offset: 4,
size: 4,
type_name: "float".into(),
},
],
};
let check = Uniforms::LAYOUT_CHECK;
slang.validate(check.rust_size, check.rust_fields).unwrap();
}
#[test]
fn layout_check_detects_mismatch_against_slang_layout() {
#[derive(goldy_derive::LayoutCheckable)]
#[repr(C)]
struct Uniforms {
x: f32,
y: f32,
}
let slang_with_wrong_offset = StructLayout {
name: "Uniforms".into(),
size: 8,
alignment: 4,
fields: vec![
StructFieldLayout {
name: "x".into(),
offset: 0,
size: 4,
type_name: "float".into(),
},
StructFieldLayout {
name: "y".into(),
offset: 8,
size: 4,
type_name: "float".into(),
},
],
};
let check = Uniforms::LAYOUT_CHECK;
let err = slang_with_wrong_offset
.validate(check.rust_size, check.rust_fields)
.unwrap_err()
.to_string();
assert!(err.contains("offset"), "expected offset mismatch: {err}");
assert!(err.contains("`y`"), "expected field y: {err}");
}
#[test]
fn layout_validation_end_to_end_catches_mismatch() {
use super::{OwnedLayoutCheck, ShaderTarget, SlangCompiler, SlangStage};
use crate::types::OptimizationLevel;
let compiler = SlangCompiler::new().expect("Slang compiler unavailable; skipping");
let source = r#"
struct MyUniforms { float x; float y; };
[shader("compute")]
[numthreads(1, 1, 1)]
void cs_main() {}
"#;
let bad_check = OwnedLayoutCheck {
type_name: "MyUniforms".into(),
rust_size: 8,
rust_fields: vec![
("x".into(), 0, 4),
("y".into(), 8, 4), ],
};
let err = compiler
.compile_with_reflection(
source,
ShaderTarget::Spirv,
&[("cs_main", SlangStage::Compute)],
&[],
&[],
&[bad_check],
OptimizationLevel::None,
)
.unwrap_err()
.to_string();
assert!(
err.contains("MyUniforms"),
"error should name the mismatched struct: {err}"
);
assert!(
err.contains("offset"),
"error should describe the offset mismatch: {err}"
);
assert!(err.contains("`y`"), "error should name the offending field: {err}");
}
#[test]
fn stride_extraction_end_to_end() {
use super::{ShaderTarget, SlangCompiler, SlangStage};
use crate::types::OptimizationLevel;
let compiler = SlangCompiler::new().expect("Slang compiler unavailable; skipping");
let manifest_dir = std::path::Path::new(env!("CARGO_MANIFEST_DIR"));
let path = manifest_dir.join("shaders").to_string_lossy().into_owned();
let source = r#"
import goldy_exp;
struct Params { float x; float y; };
[goldy_compute]
[numthreads(64, 1, 1)]
void cs_main(Params cfg, Scattered<uint> data, ThreadId id) {
data[id.x] = uint(cfg.x);
}
"#;
let result = compiler
.compile_with_reflection(
source,
ShaderTarget::Spirv,
&[("cs_main", SlangStage::Compute)],
&[&path],
&[("__SPIRV__", "1")],
&[],
OptimizationLevel::None,
)
.expect("compilation failed");
let strides = &result.reflection.binding_element_strides;
assert_eq!(strides.len(), 2, "expected 2 binding slots: {strides:?}");
assert_eq!(
strides[0],
Some(8),
"Broadcast Params {{float x; float y}} natural stride should be 8 (not cbuffer 16): {strides:?}"
);
assert_eq!(
strides[1],
Some(4),
"Scattered<uint> element stride should be 4: {strides:?}"
);
}
#[test]
fn stride_extraction_structured_buffer_element_uses_storage_layout() {
use super::{ShaderTarget, SlangCompiler, SlangStage};
use crate::types::OptimizationLevel;
let compiler = SlangCompiler::new().expect("Slang compiler unavailable; skipping");
let manifest_dir = std::path::Path::new(env!("CARGO_MANIFEST_DIR"));
let path = manifest_dir.join("shaders").to_string_lossy().into_owned();
let source = r#"
import goldy_exp;
struct Pair { uint a; uint b; };
[goldy_compute]
[numthreads(64, 1, 1)]
void cs_main(BufRO<Pair> input, Scattered<Pair> output, ThreadId id) {
output[id.x] = input[id.x];
}
"#;
let result = compiler
.compile_with_reflection(
source,
ShaderTarget::Spirv,
&[("cs_main", SlangStage::Compute)],
&[&path],
&[("__SPIRV__", "1")],
&[],
OptimizationLevel::None,
)
.expect("compilation failed");
let strides = &result.reflection.binding_element_strides;
assert_eq!(strides.len(), 2, "expected 2 binding slots: {strides:?}");
assert_eq!(
strides[0],
Some(8),
"BufRO<Pair> element stride should be 8 (not uniform 16): {strides:?}"
);
assert_eq!(
strides[1],
Some(8),
"Scattered<Pair> element stride should be 8: {strides:?}"
);
}
#[test]
fn broadcast_param_stride_matches_natural_struct_size_not_cbuffer_padded() {
use super::{ShaderTarget, SlangCompiler, SlangStage};
use crate::types::OptimizationLevel;
let compiler = SlangCompiler::new().expect("Slang compiler unavailable; skipping");
let manifest_dir = std::path::Path::new(env!("CARGO_MANIFEST_DIR"));
let path = manifest_dir.join("shaders").to_string_lossy().into_owned();
let source = r#"
import goldy_exp;
struct SimParams { float deltaTime; };
[goldy_compute]
[numthreads(64, 1, 1)]
void cs_main(Scattered<uint> data, SimParams params, ThreadId id) {
data[id.x] = uint(params.deltaTime);
}
"#;
let result = compiler
.compile_with_reflection(
source,
ShaderTarget::Spirv,
&[("cs_main", SlangStage::Compute)],
&[&path],
&[("__SPIRV__", "1")],
&[],
OptimizationLevel::None,
)
.expect("compilation failed");
let strides = &result.reflection.binding_element_strides;
assert_eq!(strides.len(), 2, "expected 2 binding slots: {strides:?}");
assert_eq!(
strides[0],
Some(4),
"Scattered<uint> element stride should be 4: {strides:?}"
);
assert_eq!(
strides[1],
Some(4),
"Broadcast SimParams{{float deltaTime}} natural stride should be 4, not cbuffer 16: {strides:?}"
);
}
#[test]
fn broadcast_two_float_struct_stride_is_eight_not_sixteen() {
use super::{ShaderTarget, SlangCompiler, SlangStage};
use crate::types::OptimizationLevel;
let compiler = SlangCompiler::new().expect("Slang compiler unavailable; skipping");
let manifest_dir = std::path::Path::new(env!("CARGO_MANIFEST_DIR"));
let path = manifest_dir.join("shaders").to_string_lossy().into_owned();
let source = r#"
import goldy_exp;
struct Params { float x; float y; };
[goldy_compute]
[numthreads(64, 1, 1)]
void cs_main(Params cfg, Scattered<uint> data, ThreadId id) {
data[id.x] = uint(cfg.x + cfg.y);
}
"#;
let result = compiler
.compile_with_reflection(
source,
ShaderTarget::Spirv,
&[("cs_main", SlangStage::Compute)],
&[&path],
&[("__SPIRV__", "1")],
&[],
OptimizationLevel::None,
)
.expect("compilation failed");
let strides = &result.reflection.binding_element_strides;
assert_eq!(strides.len(), 2, "expected 2 binding slots: {strides:?}");
assert_eq!(
strides[0],
Some(8),
"Broadcast Params{{float x; float y}} natural stride = 8: {strides:?}"
);
assert_eq!(strides[1], Some(4), "Scattered<uint> = 4: {strides:?}");
}
#[test]
fn validate_binding_strides_passes_and_fails_correctly() {
use crate::backend::validate_binding_strides;
let actual = vec![Some(16u32), Some(4u32)];
let expected = vec![Some(16u32), Some(4u32)];
assert!(validate_binding_strides(&actual, &expected, "test").is_ok());
let actual_bad = vec![Some(16u32), Some(4u32)];
let expected_bad = vec![Some(16u32), Some(16u32)];
let err =
validate_binding_strides(&actual_bad, &expected_bad, "myshader").expect_err("should fail on mismatch");
let msg = err.to_string();
assert!(msg.contains("slot 1"), "error should name the slot: {msg}");
assert!(msg.contains("myshader"), "error should name the shader: {msg}");
}
}
impl Drop for SlangCompiler {
fn drop(&mut self) {
let _guard = SLANG_PROCESS_LOCK.lock().unwrap();
if !self.global_session.is_null() {
unsafe { global_session_release(self.global_session) };
self.global_session = std::ptr::null_mut();
}
}
}
#[cfg(test)]
mod uniform_entry_point_param_binding_tests {
use super::*;
const TEST_SHADER: &str = r#"
import goldy_exp;
[goldy_compute]
[numthreads(64, 1, 1)]
void cs_main(BufRO<uint> src, Scattered<uint> dst, uint base, ThreadId id) {
uint ix = id.x + base;
dst[ix] = src[ix];
}
"#;
fn shader_path() -> String {
let manifest_dir = std::path::Path::new(env!("CARGO_MANIFEST_DIR"));
manifest_dir.join("shaders").to_string_lossy().into_owned()
}
#[test]
fn uniform_params_compile_spirv() {
let compiler = SlangCompiler::new().expect("Slang unavailable");
let path = shader_path();
let output = compiler
.compile_bindless_with_reflection_and_defines(
TEST_SHADER,
ShaderTarget::Spirv,
&[("cs_main", SlangStage::Compute)],
&[&path],
&[],
&[],
OptimizationLevel::None,
)
.expect("SPIR-V compilation failed for uniform entry-point params");
assert!(!output.shader.data.is_empty(), "SPIR-V output is empty");
let words = output.shader.as_spirv().expect("should be valid SPIR-V");
assert_eq!(words[0], 0x07230203, "SPIR-V magic number mismatch");
assert!(
words.contains(&9),
"Expected PushConstant storage class (9) in SPIR-V for uniform params"
);
}
#[cfg(target_os = "windows")]
#[test]
fn uniform_params_compile_dxil() {
let compiler = SlangCompiler::new().expect("Slang unavailable");
let path = shader_path();
let output = compiler
.compile_bindless_with_reflection_and_defines(
TEST_SHADER,
ShaderTarget::Dxil,
&[("cs_main", SlangStage::Compute)],
&[&path],
&[],
&[],
OptimizationLevel::None,
)
.expect("DXIL compilation failed for uniform entry-point params");
assert!(!output.shader.data.is_empty(), "DXIL output is empty");
let magic = u32::from_le_bytes(output.shader.data[..4].try_into().unwrap());
assert_eq!(magic, 0x43425844, "DXIL magic 'DXBC' mismatch");
}
#[test]
fn uniform_params_compile_metal() {
let compiler = SlangCompiler::new().expect("Slang unavailable");
let path = shader_path();
let output = compiler
.compile_bindless_with_reflection_and_defines(
TEST_SHADER,
ShaderTarget::Metal,
&[("cs_main", SlangStage::Compute)],
&[&path],
&[],
&[],
OptimizationLevel::None,
)
.expect("Metal MSL compilation failed for uniform entry-point params");
let msl = String::from_utf8_lossy(&output.shader.data);
assert!(!msl.is_empty(), "Metal MSL output is empty");
assert!(
msl.contains("[[buffer(") || msl.contains("buffer("),
"Expected Metal buffer binding for uniform params in MSL:\n{msl}"
);
}
#[test]
fn no_ggoldydynamic_in_compiled_output() {
let compiler = SlangCompiler::new().expect("Slang unavailable");
let path = shader_path();
let output = compiler
.compile_bindless_with_reflection_and_defines(
TEST_SHADER,
ShaderTarget::Metal,
&[("cs_main", SlangStage::Compute)],
&[&path],
&[],
&[],
OptimizationLevel::None,
)
.expect("Metal MSL compilation failed");
let msl = String::from_utf8_lossy(&output.shader.data);
assert!(
!msl.contains("gGoldyDynamic"),
"gGoldyDynamic must not appear in output MSL"
);
assert!(
!msl.contains("GoldyDynamicSlots"),
"GoldyDynamicSlots must not appear in output MSL"
);
}
}