#![deny(missing_docs)]
#![deny(rustdoc::broken_intra_doc_links)]
mod compute;
mod graphics;
mod ray_tracing;
mod shader;
pub use self::{
compute::HotComputePipeline,
graphics::HotGraphicsPipeline,
ray_tracing::HotRayTracingPipeline,
shader::{HotShader, HotShaderBuilder},
};
use {
log::{error, info},
notify::{Event, EventKind, RecommendedWatcher, recommended_watcher},
shader_prepper::{
BoxedIncludeProviderError, IncludeProvider, ResolvedInclude, ResolvedIncludePath,
process_file,
},
shaderc::{CompileOptions, Compiler, ShaderKind, SourceLanguage},
std::{
collections::HashSet,
fs::read_to_string,
io::{Error, ErrorKind},
path::{Path, PathBuf},
sync::{
Arc, OnceLock,
atomic::{AtomicBool, Ordering},
},
},
vk_graph::driver::{
DriverError,
shader::{Shader, ShaderBuilder},
},
};
struct CompiledShader {
files_included: HashSet<PathBuf>,
spirv_code: Vec<u8>,
}
fn compile_shader(
path: impl AsRef<Path>,
entry_name: &str,
shader_kind: Option<ShaderKind>,
additional_opts: Option<&CompileOptions<'_>>,
) -> anyhow::Result<CompiledShader> {
info!("Compiling: {}", path.as_ref().display());
let path = path.as_ref().to_path_buf();
let shader_kind = shader_kind.unwrap_or_else(|| guess_shader_kind(&path));
#[derive(Default)]
struct FileIncludeProvider(HashSet<PathBuf>);
impl IncludeProvider for FileIncludeProvider {
type IncludeContext = PathBuf;
fn get_include(
&mut self,
path: &ResolvedIncludePath,
) -> Result<String, BoxedIncludeProviderError> {
self.0.insert(PathBuf::from(&path.0));
Ok(read_to_string(&path.0)?)
}
fn resolve_path(
&self,
path: &str,
context: &Self::IncludeContext,
) -> Result<ResolvedInclude<Self::IncludeContext>, BoxedIncludeProviderError> {
let path = context.join(path);
Ok(ResolvedInclude {
resolved_path: ResolvedIncludePath(path.to_str().unwrap_or_default().to_string()),
context: path
.parent()
.map(|path| path.to_path_buf())
.unwrap_or_default(),
})
}
}
let mut file_include_provider = FileIncludeProvider::default();
let source_code = process_file(
path.to_string_lossy().as_ref(),
&mut file_include_provider,
PathBuf::new(),
)
.map_err(|err| {
error!("unable to process shader file: {err}");
Error::new(ErrorKind::InvalidData, err)
})?
.iter()
.map(|chunk| chunk.source.as_str())
.collect::<String>();
let files_included = file_include_provider.0;
static COMPILER: OnceLock<Compiler> = OnceLock::new();
let spirv_code = COMPILER
.get_or_init(|| Compiler::new().expect("invalid shaderc compiler"))
.compile_into_spirv(
&source_code,
shader_kind,
&path.to_string_lossy(),
entry_name,
additional_opts,
)
.inspect_err(|_| {
eprintln!("Shader: {}", path.display());
for (line_index, line) in source_code.split('\n').enumerate() {
let line_number = line_index + 1;
eprintln!("{line_number}: {line}");
}
})?
.as_binary_u8()
.to_vec();
Ok(CompiledShader {
files_included,
spirv_code,
})
}
fn compile_shader_and_watch(
shader: &HotShader,
watcher: &mut RecommendedWatcher,
) -> Result<ShaderBuilder, DriverError> {
let mut base_shader = Shader::new(shader.stage, shader.compile_and_watch(watcher)?.as_slice());
base_shader = base_shader.entry_name(shader.entry_name.clone());
if let Some(specialization) = &shader.specialization {
base_shader = base_shader.specialization(specialization.clone());
}
Ok(base_shader)
}
fn compile_shaders_and_watch(
shaders: &[HotShader],
watcher: &mut RecommendedWatcher,
) -> Result<Box<[ShaderBuilder]>, DriverError> {
shaders
.iter()
.map(|shader| compile_shader_and_watch(shader, watcher))
.collect()
}
fn create_watcher(has_changes: &Arc<AtomicBool>) -> RecommendedWatcher {
let has_changes = Arc::clone(has_changes);
recommended_watcher(move |event: notify::Result<Event>| {
let event = event.unwrap_or_else(|_| Event::new(EventKind::Any));
if matches!(
event.kind,
EventKind::Any | EventKind::Modify(_) | EventKind::Other
) {
has_changes.store(true, Ordering::Relaxed);
}
})
.expect("invalid shader watcher")
}
fn guess_shader_kind(path: impl AsRef<Path>) -> ShaderKind {
match path
.as_ref()
.extension()
.map(|ext| ext.to_string_lossy().to_string())
.unwrap_or_default()
.as_str()
{
"comp" => ShaderKind::Compute,
"task" => ShaderKind::Task,
"mesh" => ShaderKind::Mesh,
"vert" => ShaderKind::Vertex,
"geom" => ShaderKind::Geometry,
"tesc" => ShaderKind::TessControl,
"tese" => ShaderKind::TessEvaluation,
"frag" => ShaderKind::Fragment,
"rgen" => ShaderKind::RayGeneration,
"rahit" => ShaderKind::AnyHit,
"rchit" => ShaderKind::ClosestHit,
"rint" => ShaderKind::Intersection,
"rcall" => ShaderKind::Callable,
"rmiss" => ShaderKind::Miss,
_ => ShaderKind::InferFromSource,
}
}
fn guess_shader_source_language(path: impl AsRef<Path>) -> Option<SourceLanguage> {
match path
.as_ref()
.extension()
.map(|ext| ext.to_string_lossy().to_string())
.unwrap_or_default()
.as_str()
{
"comp" | "task" | "mesh" | "vert" | "geom" | "tesc" | "tese" | "frag" | "rgen"
| "rahit" | "rchit" | "rint" | "rcall" | "rmiss" | "glsl" => Some(SourceLanguage::GLSL),
"hlsl" => Some(SourceLanguage::HLSL),
_ => None,
}
}
#[cfg(test)]
mod test {
use super::{guess_shader_kind, guess_shader_source_language};
use shaderc::{ShaderKind, SourceLanguage};
#[test]
fn guess_shader_kind_from_known_extensions() {
assert_eq!(guess_shader_kind("shader.comp"), ShaderKind::Compute);
assert_eq!(guess_shader_kind("shader.vert"), ShaderKind::Vertex);
assert_eq!(guess_shader_kind("shader.frag"), ShaderKind::Fragment);
assert_eq!(guess_shader_kind("shader.rgen"), ShaderKind::RayGeneration);
}
#[test]
fn guess_shader_kind_defaults_to_infer_from_source() {
assert_eq!(
guess_shader_kind("shader.unknown"),
ShaderKind::InferFromSource
);
assert_eq!(guess_shader_kind("shader"), ShaderKind::InferFromSource);
}
#[test]
fn guess_shader_source_language_from_known_extensions() {
assert_eq!(
guess_shader_source_language("shader.comp"),
Some(SourceLanguage::GLSL)
);
assert_eq!(
guess_shader_source_language("shader.glsl"),
Some(SourceLanguage::GLSL)
);
assert_eq!(
guess_shader_source_language("shader.hlsl"),
Some(SourceLanguage::HLSL)
);
}
#[test]
fn guess_shader_source_language_returns_none_for_unknown_extensions() {
assert_eq!(guess_shader_source_language("shader.spv"), None);
assert_eq!(guess_shader_source_language("shader"), None);
}
}
macro_rules! pipeline {
($pipeline:ident) => {
paste::paste! {
impl [<Hot $pipeline>] {
fn cache(&self) -> ::std::sync::RwLockReadGuard<'_, HotPipeline<$pipeline>> {
self.cache.read().expect("poisoned hot pipeline lock")
}
fn cache_mut(&self) -> ::std::sync::RwLockWriteGuard<'_, HotPipeline<$pipeline>> {
self.cache.write().expect("poisoned hot pipeline lock")
}
}
impl [<Hot $pipeline>] {
pub fn device(&self) -> &Device {
&self.device
}
pub fn info(&self) -> [<$pipeline Info>] {
self.cache().pipeline.info()
}
pub fn set_debug_name(&self, name: impl AsRef<str>) {
self.cache().pipeline.set_debug_name(name);
}
pub fn with_debug_name(self, name: impl AsRef<str>) -> Self {
self.set_debug_name(name);
self
}
}
impl<'a> Pipeline<'a> for [<Hot $pipeline>] {
type Command = <$pipeline as Pipeline<'a>>::Command;
fn bind_cmd(self, cmd: Command<'a>) -> Self::Command {
self.compile_shader_and_bind_cmd(cmd)
}
}
impl<'a> Pipeline<'a> for &'a [<Hot $pipeline>] {
type Command = <$pipeline as Pipeline<'a>>::Command;
fn bind_cmd(self, cmd: Command<'a>) -> Self::Command {
self.compile_shader_and_bind_cmd(cmd)
}
}
}
};
}
use pipeline;
macro_rules! pipeline_handle {
($hot:ident) => {
impl $hot {
pub fn handle(&self) -> ::vk_graph::driver::ash::vk::Pipeline {
self.cache().pipeline.handle()
}
}
};
}
use pipeline_handle;
#[derive(Debug)]
struct HotPipeline<T> {
pipeline: T,
watcher: RecommendedWatcher,
}