use std::borrow::Cow;
use std::collections::{HashMap, VecDeque};
use std::sync::{Arc, Condvar, Mutex, MutexGuard, PoisonError};
use crate::shader::{ShaderId, ShaderLibrary};
const PROGRAM_BITS: u32 = 16;
const FORMAT_BITS: u32 = 12;
const BLEND_BITS: u32 = 12;
const DEPTH_BITS: u32 = 12;
const SAMPLE_BITS: u32 = 4;
const PROGRAM_SHIFT: u32 = 0;
const FORMAT_SHIFT: u32 = PROGRAM_SHIFT + PROGRAM_BITS;
const BLEND_SHIFT: u32 = FORMAT_SHIFT + FORMAT_BITS;
const DEPTH_SHIFT: u32 = BLEND_SHIFT + BLEND_BITS;
const SAMPLE_SHIFT: u32 = DEPTH_SHIFT + DEPTH_BITS;
#[derive(Clone, Debug, Default, PartialEq, Eq, Hash)]
pub struct VertexLayout {
pub array_stride: wgpu::BufferAddress,
pub step_mode: wgpu::VertexStepMode,
pub attributes: Vec<wgpu::VertexAttribute>,
}
impl VertexLayout {
#[must_use]
pub fn per_vertex(
array_stride: wgpu::BufferAddress,
attributes: Vec<wgpu::VertexAttribute>,
) -> Self {
Self {
array_stride,
step_mode: wgpu::VertexStepMode::Vertex,
attributes,
}
}
#[must_use]
pub fn as_wgpu(&self) -> wgpu::VertexBufferLayout<'_> {
wgpu::VertexBufferLayout {
array_stride: self.array_stride,
step_mode: self.step_mode,
attributes: &self.attributes,
}
}
}
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
pub struct RenderPipelineDesc {
pub shader: ShaderId,
pub vs: Cow<'static, str>,
pub fs: Cow<'static, str>,
pub vertex_layouts: Vec<VertexLayout>,
pub blend: Option<wgpu::BlendState>,
pub format: wgpu::TextureFormat,
pub sample_count: u32,
pub depth: Option<wgpu::DepthStencilState>,
pub topology: wgpu::PrimitiveTopology,
}
impl RenderPipelineDesc {
#[must_use]
pub fn new(
shader: ShaderId,
vs: impl Into<Cow<'static, str>>,
fs: impl Into<Cow<'static, str>>,
format: wgpu::TextureFormat,
) -> Self {
Self {
shader,
vs: vs.into(),
fs: fs.into(),
vertex_layouts: Vec::new(),
blend: None,
format,
sample_count: 1,
depth: None,
topology: wgpu::PrimitiveTopology::TriangleList,
}
}
}
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
struct ProgramKey {
shader: ShaderId,
vs: Cow<'static, str>,
fs: Cow<'static, str>,
vertex_layouts: Vec<VertexLayout>,
topology: wgpu::PrimitiveTopology,
}
impl ProgramKey {
fn of(desc: &RenderPipelineDesc) -> Self {
Self {
shader: desc.shader,
vs: desc.vs.clone(),
fs: desc.fs.clone(),
vertex_layouts: desc.vertex_layouts.clone(),
topology: desc.topology,
}
}
fn matches(&self, desc: &RenderPipelineDesc) -> bool {
self.shader == desc.shader
&& self.vs == desc.vs
&& self.fs == desc.fs
&& self.topology == desc.topology
&& self.vertex_layouts == desc.vertex_layouts
}
}
#[derive(Debug)]
struct AxisTable<K> {
index: HashMap<K, u64>,
bits: u32,
}
impl<K: std::hash::Hash + Eq + Clone> AxisTable<K> {
fn new(bits: u32) -> Self {
Self {
index: HashMap::new(),
bits,
}
}
fn intern(&mut self, value: &K) -> Option<u64> {
if let Some(existing) = self.index.get(value) {
return Some(*existing);
}
let next = self.index.len() as u64;
if next >= 1 << self.bits {
return None;
}
self.index.insert(value.clone(), next);
Some(next)
}
}
#[derive(Debug, Default)]
struct ProgramTable {
by_hash: HashMap<u64, Vec<u64>>,
entries: Vec<ProgramKey>,
}
impl ProgramTable {
fn intern(&mut self, desc: &RenderPipelineDesc) -> Option<u64> {
let hash = program_hash(desc);
let chain = self.by_hash.entry(hash).or_default();
for &candidate in chain.iter() {
if self.entries[candidate as usize].matches(desc) {
return Some(candidate);
}
}
let next = self.entries.len() as u64;
if next >= 1 << PROGRAM_BITS {
return None;
}
chain.push(next);
self.entries.push(ProgramKey::of(desc));
Some(next)
}
}
fn program_hash(desc: &RenderPipelineDesc) -> u64 {
use std::hash::{Hash, Hasher};
let mut hasher = std::collections::hash_map::DefaultHasher::new();
desc.shader.hash(&mut hasher);
desc.vs.hash(&mut hasher);
desc.fs.hash(&mut hasher);
desc.vertex_layouts.hash(&mut hasher);
desc.topology.hash(&mut hasher);
hasher.finish()
}
#[derive(Debug)]
struct KeyPacker {
programs: ProgramTable,
formats: AxisTable<wgpu::TextureFormat>,
blends: AxisTable<Option<wgpu::BlendState>>,
depths: AxisTable<Option<wgpu::DepthStencilState>>,
}
impl Default for KeyPacker {
fn default() -> Self {
Self {
programs: ProgramTable::default(),
formats: AxisTable::new(FORMAT_BITS),
blends: AxisTable::new(BLEND_BITS),
depths: AxisTable::new(DEPTH_BITS),
}
}
}
impl KeyPacker {
fn key_for(&mut self, desc: &RenderPipelineDesc) -> Option<u64> {
let sample = sample_field(desc.sample_count)?;
let program = self.programs.intern(desc)?;
let format = self.formats.intern(&desc.format)?;
let blend = self.blends.intern(&desc.blend)?;
let depth = self.depths.intern(&desc.depth)?;
Some(
(program << PROGRAM_SHIFT)
| (format << FORMAT_SHIFT)
| (blend << BLEND_SHIFT)
| (depth << DEPTH_SHIFT)
| (sample << SAMPLE_SHIFT),
)
}
}
fn sample_field(sample_count: u32) -> Option<u64> {
if sample_count == 0 || !sample_count.is_power_of_two() {
return None;
}
let field = u64::from(sample_count.trailing_zeros());
(field < 1 << SAMPLE_BITS).then_some(field)
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum JobState {
Pending,
Building,
Done,
Failed,
}
#[derive(Debug)]
struct Job {
desc: RenderPipelineDesc,
state: JobState,
}
#[derive(Debug)]
struct Inner<P> {
built: HashMap<u64, Arc<P>>,
jobs: HashMap<u64, Job>,
queue: VecDeque<u64>,
compiles: u64,
worker_live: bool,
shutdown: bool,
worker_epoch: u64,
}
impl<P> Default for Inner<P> {
fn default() -> Self {
Self {
built: HashMap::new(),
jobs: HashMap::new(),
queue: VecDeque::new(),
compiles: 0,
worker_live: false,
shutdown: false,
worker_epoch: 0,
}
}
}
impl<P> Inner<P> {
fn claim_next_pending(&mut self) -> Option<(u64, RenderPipelineDesc)> {
while !self.shutdown
&& let Some(key) = self.queue.pop_front()
{
if let Some(job) = self.jobs.get_mut(&key)
&& job.state == JobState::Pending
{
job.state = JobState::Building;
return Some((key, job.desc.clone()));
}
}
self.worker_live = false;
None
}
}
#[derive(Debug)]
struct Shared<P> {
inner: Mutex<Inner<P>>,
built: Condvar,
}
impl<P> Default for Shared<P> {
fn default() -> Self {
Self {
inner: Mutex::new(Inner::default()),
built: Condvar::new(),
}
}
}
fn lock<T>(mutex: &Mutex<T>) -> MutexGuard<'_, T> {
mutex.lock().unwrap_or_else(PoisonError::into_inner)
}
struct BuildGuard<'a, P> {
shared: &'a Shared<P>,
key: u64,
armed: bool,
}
impl<'a, P> BuildGuard<'a, P> {
fn new(shared: &'a Shared<P>, key: u64) -> Self {
Self {
shared,
key,
armed: true,
}
}
fn publish(mut self, pipeline: Arc<P>) {
self.armed = false;
let mut inner = lock(&self.shared.inner);
inner.built.insert(self.key, pipeline);
if let Some(job) = inner.jobs.get_mut(&self.key) {
job.state = JobState::Done;
}
inner.compiles += 1;
drop(inner);
self.shared.built.notify_all();
}
}
struct LiveWorkerGuard<'a, P> {
shared: &'a Shared<P>,
epoch: u64,
}
impl<P> Drop for LiveWorkerGuard<'_, P> {
fn drop(&mut self) {
let mut inner = lock(&self.shared.inner);
if inner.worker_epoch == self.epoch {
inner.worker_live = false;
}
}
}
impl<P> Drop for BuildGuard<'_, P> {
fn drop(&mut self) {
if !self.armed {
return;
}
let mut inner = lock(&self.shared.inner);
if let Some(job) = inner.jobs.get_mut(&self.key) {
job.state = JobState::Failed;
}
drop(inner);
self.shared.built.notify_all();
}
}
#[derive(Debug)]
struct VariantCache<P> {
shared: Arc<Shared<P>>,
local: HashMap<u64, Arc<P>>,
overflow: Option<Arc<P>>,
keys: KeyPacker,
warned_key_space: bool,
worker: Option<std::thread::JoinHandle<()>>,
}
impl<P> Default for VariantCache<P> {
fn default() -> Self {
Self {
shared: Arc::new(Shared::default()),
local: HashMap::new(),
overflow: None,
keys: KeyPacker::default(),
warned_key_space: false,
worker: None,
}
}
}
impl<P> Drop for VariantCache<P> {
fn drop(&mut self) {
self.shutdown();
}
}
impl<P> VariantCache<P> {
fn shared(&self) -> Arc<Shared<P>> {
Arc::clone(&self.shared)
}
fn get_or_create<F>(&mut self, desc: &RenderPipelineDesc, compile: F) -> &P
where
F: Fn(&RenderPipelineDesc) -> P,
{
let Some(key) = self.keys.key_for(desc) else {
if !self.warned_key_space {
self.warned_key_space = true;
log::error!(
"frust-gpu: pipeline variant key space exhausted (or an unusable \
sample_count {}); this variant is compiled per request instead of cached",
desc.sample_count
);
}
self.overflow = Some(Arc::new(compile(desc)));
return self
.overflow
.as_deref()
.expect("the uncached variant was just stored");
};
if !self.local.contains_key(&key) {
let built = self.resolve(key, desc, &compile);
self.local.insert(key, built);
}
&self.local[&key]
}
fn resolve<F>(&self, key: u64, desc: &RenderPipelineDesc, compile: &F) -> Arc<P>
where
F: Fn(&RenderPipelineDesc) -> P,
{
let mut inner = lock(&self.shared.inner);
loop {
if let Some(built) = inner.built.get(&key) {
return Arc::clone(built);
}
match inner.jobs.get(&key).map(|job| job.state) {
Some(JobState::Building) => {
inner = self
.shared
.built
.wait(inner)
.unwrap_or_else(PoisonError::into_inner);
}
Some(JobState::Pending | JobState::Failed) => {
if let Some(job) = inner.jobs.get_mut(&key) {
job.state = JobState::Building;
}
drop(inner);
return self.build(key, desc, compile);
}
Some(JobState::Done) | None => {
inner.jobs.insert(
key,
Job {
desc: desc.clone(),
state: JobState::Building,
},
);
drop(inner);
return self.build(key, desc, compile);
}
}
}
}
fn build<F>(&self, key: u64, desc: &RenderPipelineDesc, compile: &F) -> Arc<P>
where
F: Fn(&RenderPipelineDesc) -> P,
{
let guard = BuildGuard::new(&self.shared, key);
let pipeline = Arc::new(compile(desc));
guard.publish(Arc::clone(&pipeline));
pipeline
}
fn enqueue(&mut self, descs: &[RenderPipelineDesc]) -> bool {
let mut inner = lock(&self.shared.inner);
for desc in descs {
let Some(key) = self.keys.key_for(desc) else {
log::error!(
"frust-gpu: warm-up variant has no representable key (sample_count {}); \
it will be compiled on first use instead",
desc.sample_count
);
continue;
};
if inner.built.contains_key(&key) || inner.jobs.contains_key(&key) {
continue;
}
inner.jobs.insert(
key,
Job {
desc: desc.clone(),
state: JobState::Pending,
},
);
inner.queue.push_back(key);
}
if inner.shutdown || inner.worker_live || inner.queue.is_empty() {
return false;
}
inner.worker_live = true;
inner.worker_epoch += 1;
true
}
fn drain_queue<F>(shared: &Arc<Shared<P>>, compile: F)
where
F: Fn(&RenderPipelineDesc) -> P,
{
let _live = LiveWorkerGuard {
shared,
epoch: lock(&shared.inner).worker_epoch,
};
loop {
let claimed = lock(&shared.inner).claim_next_pending();
let Some((key, desc)) = claimed else {
return;
};
let built = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
let guard = BuildGuard::new(shared, key);
guard.publish(Arc::new(compile(&desc)));
}));
if built.is_err() {
log::error!(
"frust-gpu: a warm-up compile panicked; that variant is left to be \
rebuilt on request and the rest of the queue continues"
);
}
}
}
fn warm_up<F>(&mut self, descs: &[RenderPipelineDesc], make_compile: impl Fn() -> F)
where
P: Send + Sync + 'static,
F: Fn(&RenderPipelineDesc) -> P + Send + 'static,
{
if !self.enqueue(descs) {
return;
}
self.join_worker();
let worker_shared = self.shared();
let worker_compile = make_compile();
let spawned = std::thread::Builder::new()
.name("frust-gpu pipeline warm-up".to_string())
.spawn(move || VariantCache::drain_queue(&worker_shared, worker_compile));
match spawned {
Ok(handle) => self.worker = Some(handle),
Err(err) => {
log::warn!(
"frust-gpu: could not spawn the pipeline warm-up thread ({err}); \
building the listed variants inline"
);
VariantCache::drain_queue(&self.shared(), make_compile());
}
}
}
fn shutdown(&mut self) {
lock(&self.shared.inner).shutdown = true;
self.join_worker();
}
fn join_worker(&mut self) {
if let Some(worker) = self.worker.take()
&& worker.join().is_err()
{
log::error!("frust-gpu: the pipeline warm-up worker panicked");
}
}
fn compiled_variants(&self) -> u64 {
lock(&self.shared.inner).compiles
}
fn queued_variants(&self) -> usize {
let inner = lock(&self.shared.inner);
inner
.jobs
.values()
.filter(|job| job.state == JobState::Pending)
.count()
}
}
#[derive(Debug)]
pub struct PipelineCache {
core: VariantCache<wgpu::RenderPipeline>,
shaders: Arc<ShaderLibrary>,
driver_cache: Option<wgpu::PipelineCache>,
}
impl PipelineCache {
#[must_use]
pub fn new(shaders: Arc<ShaderLibrary>, driver_cache: Option<wgpu::PipelineCache>) -> Self {
Self {
core: VariantCache::default(),
shaders,
driver_cache,
}
}
#[must_use]
pub fn shaders(&self) -> &ShaderLibrary {
&self.shaders
}
pub fn get_or_create(
&mut self,
device: &wgpu::Device,
desc: &RenderPipelineDesc,
) -> &wgpu::RenderPipeline {
let shaders = &self.shaders;
let driver_cache = self.driver_cache.as_ref();
self.core.get_or_create(desc, |desc| {
build_render_pipeline(device, shaders, driver_cache, desc)
})
}
pub fn warm_up(&mut self, device: &wgpu::Device, descs: &[RenderPipelineDesc]) {
let shaders = Arc::clone(&self.shaders);
let driver_cache = self.driver_cache.clone();
self.core
.warm_up(descs, || compiler(device, &shaders, driver_cache.as_ref()));
}
pub fn shutdown(&mut self) {
self.core.shutdown();
}
#[must_use]
pub fn compiled_variants(&self) -> u64 {
self.core.compiled_variants()
}
#[must_use]
pub fn queued_variants(&self) -> usize {
self.core.queued_variants()
}
}
fn compiler(
device: &wgpu::Device,
shaders: &Arc<ShaderLibrary>,
driver_cache: Option<&wgpu::PipelineCache>,
) -> impl Fn(&RenderPipelineDesc) -> wgpu::RenderPipeline + Send + 'static + use<> {
let device = device.clone();
let shaders = Arc::clone(shaders);
let driver_cache = driver_cache.cloned();
move |desc| build_render_pipeline(&device, &shaders, driver_cache.as_ref(), desc)
}
fn build_render_pipeline(
device: &wgpu::Device,
shaders: &ShaderLibrary,
driver_cache: Option<&wgpu::PipelineCache>,
desc: &RenderPipelineDesc,
) -> wgpu::RenderPipeline {
let name = shaders.name_of(desc.shader).unwrap_or("<unknown shader>");
let module = shaders
.get(desc.shader)
.expect("a RenderPipelineDesc's ShaderId must come from this cache's ShaderLibrary");
let layouts: Vec<Option<wgpu::VertexBufferLayout<'_>>> = desc
.vertex_layouts
.iter()
.map(|layout| Some(VertexLayout::as_wgpu(layout)))
.collect();
let label = format!(
"frust-gpu pipeline: {name} [{:?}, msaa x{}]",
desc.format, desc.sample_count
);
device.create_render_pipeline(&wgpu::RenderPipelineDescriptor {
label: Some(&label),
layout: None,
vertex: wgpu::VertexState {
module,
entry_point: Some(desc.vs.as_ref()),
compilation_options: wgpu::PipelineCompilationOptions::default(),
buffers: &layouts,
},
primitive: wgpu::PrimitiveState {
topology: desc.topology,
..Default::default()
},
depth_stencil: desc.depth.clone(),
multisample: wgpu::MultisampleState {
count: desc.sample_count,
mask: !0,
alpha_to_coverage_enabled: false,
},
fragment: Some(wgpu::FragmentState {
module,
entry_point: Some(desc.fs.as_ref()),
compilation_options: wgpu::PipelineCompilationOptions::default(),
targets: &[Some(wgpu::ColorTargetState {
format: desc.format,
blend: desc.blend,
write_mask: wgpu::ColorWrites::ALL,
})],
}),
multiview_mask: None,
cache: driver_cache,
})
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
#[derive(Debug, PartialEq, Eq)]
struct FakePipeline {
sample_count: u32,
serial: u64,
}
#[derive(Default)]
struct Counter(AtomicU64);
impl Counter {
fn compile(&self, desc: &RenderPipelineDesc) -> FakePipeline {
let serial = self.0.fetch_add(1, Ordering::SeqCst);
FakePipeline {
sample_count: desc.sample_count,
serial,
}
}
fn count(&self) -> u64 {
self.0.load(Ordering::SeqCst)
}
}
fn desc() -> RenderPipelineDesc {
RenderPipelineDesc::new(
ShaderId::from_raw(0),
"vs_main",
"fs_main",
wgpu::TextureFormat::Rgba8Unorm,
)
}
#[test]
fn a_repeated_request_compiles_once() {
let mut cache = VariantCache::<FakePipeline>::default();
let counter = Counter::default();
let desc = desc();
let first = cache.get_or_create(&desc, |d| counter.compile(d)).serial;
for _ in 0..8 {
let again = cache.get_or_create(&desc, |d| counter.compile(d)).serial;
assert_eq!(again, first, "every repeat must return the same pipeline");
}
assert_eq!(counter.count(), 1);
assert_eq!(cache.compiled_variants(), 1);
}
#[test]
fn each_key_axis_is_its_own_variant() {
let base = desc();
let variants = [
RenderPipelineDesc {
blend: Some(wgpu::BlendState::ALPHA_BLENDING),
..base.clone()
},
RenderPipelineDesc {
format: wgpu::TextureFormat::Bgra8Unorm,
..base.clone()
},
RenderPipelineDesc {
sample_count: 4,
..base.clone()
},
RenderPipelineDesc {
depth: Some(wgpu::DepthStencilState {
format: wgpu::TextureFormat::Depth32Float,
depth_write_enabled: Some(true),
depth_compare: Some(wgpu::CompareFunction::Less),
stencil: wgpu::StencilState::default(),
bias: wgpu::DepthBiasState::default(),
}),
..base.clone()
},
RenderPipelineDesc {
shader: ShaderId::from_raw(1),
..base.clone()
},
RenderPipelineDesc {
fs: "fs_other".into(),
..base.clone()
},
RenderPipelineDesc {
vertex_layouts: vec![VertexLayout::per_vertex(
16,
vec![wgpu::VertexAttribute {
format: wgpu::VertexFormat::Float32x4,
offset: 0,
shader_location: 0,
}],
)],
..base.clone()
},
RenderPipelineDesc {
topology: wgpu::PrimitiveTopology::LineList,
..base.clone()
},
];
let mut packer = KeyPacker::default();
let base_key = packer.key_for(&base).expect("base key");
let mut seen = vec![base_key];
for variant in &variants {
let key = packer.key_for(variant).expect("variant key");
assert!(
!seen.contains(&key),
"{variant:?} must not reuse an existing key"
);
seen.push(key);
}
let mut cache = VariantCache::<FakePipeline>::default();
let counter = Counter::default();
cache.get_or_create(&base, |d| counter.compile(d));
for variant in &variants {
cache.get_or_create(variant, |d| counter.compile(d));
}
assert_eq!(counter.count(), 1 + variants.len() as u64);
}
#[test]
fn a_key_is_stable_across_repeated_packing() {
let mut packer = KeyPacker::default();
let desc = desc();
let first = packer.key_for(&desc).expect("key");
for _ in 0..4 {
assert_eq!(packer.key_for(&desc), Some(first));
}
}
#[test]
fn an_unusable_sample_count_has_no_key_and_is_not_cached() {
let mut packer = KeyPacker::default();
assert_eq!(
packer.key_for(&RenderPipelineDesc {
sample_count: 3,
..desc()
}),
None
);
assert_eq!(
packer.key_for(&RenderPipelineDesc {
sample_count: 0,
..desc()
}),
None
);
let mut cache = VariantCache::<FakePipeline>::default();
let counter = Counter::default();
let odd = RenderPipelineDesc {
sample_count: 3,
..desc()
};
assert_eq!(
cache
.get_or_create(&odd, |d| counter.compile(d))
.sample_count,
3
);
cache.get_or_create(&odd, |d| counter.compile(d));
assert_eq!(counter.count(), 2, "an unkeyable variant is never cached");
}
#[test]
fn draining_the_queue_builds_every_listed_variant_once() {
let mut cache = VariantCache::<FakePipeline>::default();
let counter = Counter::default();
let listed: Vec<_> = [1u32, 2, 4]
.into_iter()
.map(|sample_count| RenderPipelineDesc {
sample_count,
..desc()
})
.collect();
cache.enqueue(&listed);
assert_eq!(cache.queued_variants(), 3);
VariantCache::drain_queue(&cache.shared(), |d| counter.compile(d));
assert_eq!(counter.count(), 3);
assert_eq!(cache.queued_variants(), 0);
for desc in &listed {
cache.get_or_create(desc, |d| counter.compile(d));
}
assert_eq!(counter.count(), 3, "a warmed variant must not recompile");
}
#[test]
fn enqueueing_the_same_variant_twice_queues_one_job() {
let mut cache = VariantCache::<FakePipeline>::default();
let listed = [desc(), desc()];
cache.enqueue(&listed);
cache.enqueue(&listed);
assert_eq!(cache.queued_variants(), 1);
}
#[test]
fn a_render_thread_request_steals_a_queued_job_and_the_queue_skips_it() {
let mut cache = VariantCache::<FakePipeline>::default();
let counter = Counter::default();
let listed: Vec<_> = [1u32, 2, 4]
.into_iter()
.map(|sample_count| RenderPipelineDesc {
sample_count,
..desc()
})
.collect();
cache.enqueue(&listed);
let stolen = cache
.get_or_create(&listed[2], |d| counter.compile(d))
.serial;
assert_eq!(stolen, 0, "the steal is the first compile to run");
assert_eq!(counter.count(), 1);
assert_eq!(cache.queued_variants(), 2, "the stolen job is done");
VariantCache::drain_queue(&cache.shared(), |d| counter.compile(d));
assert_eq!(
counter.count(),
3,
"the stolen variant must not compile twice"
);
assert_eq!(
cache
.get_or_create(&listed[2], |d| counter.compile(d))
.serial,
stolen
);
assert_eq!(counter.count(), 3);
}
#[test]
fn a_worker_thread_drain_publishes_to_the_render_thread() {
let mut cache = VariantCache::<FakePipeline>::default();
let listed: Vec<_> = (0..6)
.map(|i| RenderPipelineDesc {
shader: ShaderId::from_raw(i),
..desc()
})
.collect();
cache.enqueue(&listed);
let shared = cache.shared();
let worker_counter = Arc::new(Counter::default());
let thread_counter = Arc::clone(&worker_counter);
let worker = std::thread::spawn(move || {
VariantCache::drain_queue(&shared, move |d| thread_counter.compile(d));
});
worker.join().expect("the warm-up worker must not panic");
assert_eq!(worker_counter.count(), 6);
assert_eq!(cache.compiled_variants(), 6);
let render_counter = Counter::default();
for desc in &listed {
cache.get_or_create(desc, |d| render_counter.compile(d));
}
assert_eq!(render_counter.count(), 0);
}
#[test]
fn an_abandoned_build_marks_the_job_failed_rather_than_pending() {
let cache = VariantCache::<FakePipeline>::default();
let shared = cache.shared();
let guard = BuildGuard::new(&shared, 42);
lock(&shared.inner).jobs.insert(
42,
Job {
desc: desc(),
state: JobState::Building,
},
);
drop(guard);
assert_eq!(
lock(&shared.inner).jobs.get(&42).map(|job| job.state),
Some(JobState::Failed),
"an abandoned build must not leave a Pending job no worker can claim"
);
}
struct PanicOnce {
counter: Arc<Counter>,
armed: AtomicBool,
}
impl PanicOnce {
fn new(counter: &Arc<Counter>) -> Self {
Self {
counter: Arc::clone(counter),
armed: AtomicBool::new(true),
}
}
fn compile(&self, desc: &RenderPipelineDesc) -> FakePipeline {
assert!(
!self.armed.swap(false, Ordering::SeqCst),
"the fake compile fails its first call"
);
self.counter.compile(desc)
}
}
#[test]
fn a_panicking_warm_up_compile_costs_one_variant_and_not_the_queue() {
let mut cache = VariantCache::<FakePipeline>::default();
let counter = Arc::new(Counter::default());
let failing = PanicOnce::new(&counter);
let listed: Vec<_> = [1u32, 2, 4]
.into_iter()
.map(|sample_count| RenderPipelineDesc {
sample_count,
..desc()
})
.collect();
cache.enqueue(&listed);
VariantCache::drain_queue(&cache.shared(), |d| failing.compile(d));
assert_eq!(counter.count(), 2);
assert_eq!(
cache.queued_variants(),
0,
"a failed job must not be counted as still queued forever"
);
assert_eq!(
cache
.get_or_create(&listed[0], |d| counter.compile(d))
.sample_count,
1
);
assert_eq!(counter.count(), 3);
for desc in &listed {
cache.get_or_create(desc, |d| counter.compile(d));
}
assert_eq!(counter.count(), 3);
}
#[test]
fn a_panicking_inline_build_surfaces_to_the_requester_and_leaves_it_rebuildable() {
let mut cache = VariantCache::<FakePipeline>::default();
let counter = Arc::new(Counter::default());
let failing = PanicOnce::new(&counter);
let desc = desc();
cache.enqueue(std::slice::from_ref(&desc));
let stolen = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
cache.get_or_create(&desc, |d| failing.compile(d));
}));
assert!(
stolen.is_err(),
"an inline compile must not swallow a panic"
);
assert_eq!(counter.count(), 0);
assert_eq!(cache.get_or_create(&desc, |d| counter.compile(d)).serial, 0);
assert_eq!(counter.count(), 1);
assert_eq!(cache.queued_variants(), 0);
}
#[derive(Default)]
struct Latch {
open: Mutex<bool>,
changed: Condvar,
}
impl Latch {
fn wait(&self) {
let mut open = lock(&self.open);
while !*open {
open = self
.changed
.wait(open)
.unwrap_or_else(PoisonError::into_inner);
}
}
fn open(&self) {
*lock(&self.open) = true;
self.changed.notify_all();
}
}
fn shader_variants(count: u32) -> Vec<RenderPipelineDesc> {
(0..count)
.map(|i| RenderPipelineDesc {
shader: ShaderId::from_raw(i),
..desc()
})
.collect()
}
#[test]
fn repeated_warm_up_calls_keep_at_most_one_worker() {
let mut cache = VariantCache::<FakePipeline>::default();
let counter = Arc::new(Counter::default());
let latch = Arc::new(Latch::default());
let workers = AtomicU64::new(0);
let listed = shader_variants(6);
let make_compile = || {
workers.fetch_add(1, Ordering::SeqCst);
let counter = Arc::clone(&counter);
let latch = Arc::clone(&latch);
let held = AtomicBool::new(false);
move |d: &RenderPipelineDesc| {
if !held.swap(true, Ordering::SeqCst) {
latch.wait();
}
counter.compile(d)
}
};
cache.warm_up(&listed[..1], make_compile);
for _ in 0..8 {
cache.warm_up(&listed, make_compile);
}
assert_eq!(
workers.load(Ordering::SeqCst),
1,
"a live worker must absorb further warm-up calls, not be joined by more"
);
latch.open();
cache.shutdown();
for desc in &listed {
cache.get_or_create(desc, |d| counter.compile(d));
}
assert_eq!(counter.count(), listed.len() as u64);
assert_eq!(cache.compiled_variants(), listed.len() as u64);
}
#[test]
fn dropping_the_cache_joins_its_worker() {
let device = Arc::new(());
let counter = Arc::new(Counter::default());
let listed = shader_variants(4);
let mut cache = VariantCache::<FakePipeline>::default();
cache.warm_up(&listed, || {
let device = Arc::clone(&device);
let counter = Arc::clone(&counter);
move |d: &RenderPipelineDesc| {
let _held = Arc::clone(&device);
counter.compile(d)
}
});
drop(cache);
assert_eq!(
Arc::strong_count(&device),
1,
"no warm-up worker may outlive the cache that started it"
);
}
#[test]
fn warming_up_after_shutdown_starts_no_worker_and_still_builds_on_request() {
let mut cache = VariantCache::<FakePipeline>::default();
let counter = Arc::new(Counter::default());
let workers = AtomicU64::new(0);
let listed = shader_variants(3);
cache.shutdown();
cache.shutdown();
cache.warm_up(&listed, || {
workers.fetch_add(1, Ordering::SeqCst);
let counter = Arc::clone(&counter);
move |d: &RenderPipelineDesc| counter.compile(d)
});
assert_eq!(workers.load(Ordering::SeqCst), 0);
for desc in &listed {
cache.get_or_create(desc, |d| counter.compile(d));
}
assert_eq!(counter.count(), listed.len() as u64);
}
fn drain_error_scope(
device: &wgpu::Device,
scope: wgpu::ErrorScopeGuard,
) -> Option<wgpu::Error> {
use std::task::{Context, Poll, Waker};
let waker = Waker::noop();
let mut cx = Context::from_waker(waker);
let mut future = std::pin::pin!(scope.pop());
loop {
match future.as_mut().poll(&mut cx) {
Poll::Ready(error) => return error,
Poll::Pending => {
let _ = device.poll(wgpu::PollType::wait_indefinitely());
}
}
}
}
#[test]
#[ignore = "requires a GPU (Metal/Vulkan); run locally with `cargo test -p frust-gpu -- --ignored`"]
fn creates_one_real_pipeline() {
const WGSL: &str = r#"
@vertex
fn vs_main(@builtin(vertex_index) i: u32) -> @builtin(position) vec4<f32> {
let uv = vec2<f32>(f32((i << 1u) & 2u), f32(i & 2u));
return vec4<f32>(uv * 2.0 - 1.0, 0.0, 1.0);
}
@fragment
fn fs_main() -> @location(0) vec4<f32> {
return vec4<f32>(0.0, 1.0, 0.0, 1.0);
}
"#;
let (device, _queue) = pollster::block_on(async {
let instance = wgpu::Instance::new(
wgpu::InstanceDescriptor::new_without_display_handle_from_env(),
);
let adapter = wgpu::util::initialize_adapter_from_env_or_default(&instance, None)
.await
.expect("no compatible GPU adapter");
println!("frust-gpu pipeline test adapter: {:?}", adapter.get_info());
adapter
.request_device(&wgpu::DeviceDescriptor {
label: Some("frust-gpu pipeline test device"),
required_features: wgpu::Features::empty(),
required_limits: wgpu::Limits::default(),
..Default::default()
})
.await
.expect("failed to create the device")
});
let mut library = ShaderLibrary::new();
let shader = library.insert_wgsl(&device, "fullscreen-green", WGSL);
let mut cache = PipelineCache::new(Arc::new(library), None);
let desc = RenderPipelineDesc::new(
shader,
"vs_main",
"fs_main",
wgpu::TextureFormat::Rgba8Unorm,
);
let scope = device.push_error_scope(wgpu::ErrorFilter::Validation);
for _ in 0..2 {
let _pipeline = cache.get_or_create(&device, &desc);
}
let error = drain_error_scope(&device, scope);
assert!(error.is_none(), "pipeline creation raised {error:?}");
assert_eq!(
cache.compiled_variants(),
1,
"a repeat request must reuse the pipeline"
);
}
}