1use std::borrow::Cow;
91use std::collections::{HashMap, VecDeque};
92use std::sync::{Arc, Condvar, Mutex, MutexGuard, PoisonError};
93
94use crate::shader::{ShaderId, ShaderLibrary};
95
96const PROGRAM_BITS: u32 = 16;
98const FORMAT_BITS: u32 = 12;
100const BLEND_BITS: u32 = 12;
102const DEPTH_BITS: u32 = 12;
104const SAMPLE_BITS: u32 = 4;
106
107const PROGRAM_SHIFT: u32 = 0;
108const FORMAT_SHIFT: u32 = PROGRAM_SHIFT + PROGRAM_BITS;
109const BLEND_SHIFT: u32 = FORMAT_SHIFT + FORMAT_BITS;
110const DEPTH_SHIFT: u32 = BLEND_SHIFT + BLEND_BITS;
111const SAMPLE_SHIFT: u32 = DEPTH_SHIFT + DEPTH_BITS;
112
113#[derive(Clone, Debug, Default, PartialEq, Eq, Hash)]
120pub struct VertexLayout {
121 pub array_stride: wgpu::BufferAddress,
123 pub step_mode: wgpu::VertexStepMode,
125 pub attributes: Vec<wgpu::VertexAttribute>,
127}
128
129impl VertexLayout {
130 #[must_use]
132 pub fn per_vertex(
133 array_stride: wgpu::BufferAddress,
134 attributes: Vec<wgpu::VertexAttribute>,
135 ) -> Self {
136 Self {
137 array_stride,
138 step_mode: wgpu::VertexStepMode::Vertex,
139 attributes,
140 }
141 }
142
143 #[must_use]
145 pub fn as_wgpu(&self) -> wgpu::VertexBufferLayout<'_> {
146 wgpu::VertexBufferLayout {
147 array_stride: self.array_stride,
148 step_mode: self.step_mode,
149 attributes: &self.attributes,
150 }
151 }
152}
153
154#[derive(Clone, Debug, PartialEq, Eq, Hash)]
165pub struct RenderPipelineDesc {
166 pub shader: ShaderId,
168 pub vs: Cow<'static, str>,
170 pub fs: Cow<'static, str>,
172 pub vertex_layouts: Vec<VertexLayout>,
175 pub blend: Option<wgpu::BlendState>,
177 pub format: wgpu::TextureFormat,
179 pub sample_count: u32,
181 pub depth: Option<wgpu::DepthStencilState>,
183 pub topology: wgpu::PrimitiveTopology,
186}
187
188impl RenderPipelineDesc {
189 #[must_use]
193 pub fn new(
194 shader: ShaderId,
195 vs: impl Into<Cow<'static, str>>,
196 fs: impl Into<Cow<'static, str>>,
197 format: wgpu::TextureFormat,
198 ) -> Self {
199 Self {
200 shader,
201 vs: vs.into(),
202 fs: fs.into(),
203 vertex_layouts: Vec::new(),
204 blend: None,
205 format,
206 sample_count: 1,
207 depth: None,
208 topology: wgpu::PrimitiveTopology::TriangleList,
209 }
210 }
211}
212
213#[derive(Clone, Debug, PartialEq, Eq, Hash)]
216struct ProgramKey {
217 shader: ShaderId,
218 vs: Cow<'static, str>,
219 fs: Cow<'static, str>,
220 vertex_layouts: Vec<VertexLayout>,
221 topology: wgpu::PrimitiveTopology,
222}
223
224impl ProgramKey {
225 fn of(desc: &RenderPipelineDesc) -> Self {
226 Self {
227 shader: desc.shader,
228 vs: desc.vs.clone(),
229 fs: desc.fs.clone(),
230 vertex_layouts: desc.vertex_layouts.clone(),
231 topology: desc.topology,
232 }
233 }
234
235 fn matches(&self, desc: &RenderPipelineDesc) -> bool {
238 self.shader == desc.shader
239 && self.vs == desc.vs
240 && self.fs == desc.fs
241 && self.topology == desc.topology
242 && self.vertex_layouts == desc.vertex_layouts
243 }
244}
245
246#[derive(Debug)]
251struct AxisTable<K> {
252 index: HashMap<K, u64>,
253 bits: u32,
254}
255
256impl<K: std::hash::Hash + Eq + Clone> AxisTable<K> {
257 fn new(bits: u32) -> Self {
258 Self {
259 index: HashMap::new(),
260 bits,
261 }
262 }
263
264 fn intern(&mut self, value: &K) -> Option<u64> {
265 if let Some(existing) = self.index.get(value) {
266 return Some(*existing);
267 }
268 let next = self.index.len() as u64;
269 if next >= 1 << self.bits {
270 return None;
271 }
272 self.index.insert(value.clone(), next);
273 Some(next)
274 }
275}
276
277#[derive(Debug, Default)]
280struct ProgramTable {
281 by_hash: HashMap<u64, Vec<u64>>,
285 entries: Vec<ProgramKey>,
286}
287
288impl ProgramTable {
289 fn intern(&mut self, desc: &RenderPipelineDesc) -> Option<u64> {
290 let hash = program_hash(desc);
291 let chain = self.by_hash.entry(hash).or_default();
292 for &candidate in chain.iter() {
293 if self.entries[candidate as usize].matches(desc) {
294 return Some(candidate);
295 }
296 }
297 let next = self.entries.len() as u64;
298 if next >= 1 << PROGRAM_BITS {
299 return None;
300 }
301 chain.push(next);
302 self.entries.push(ProgramKey::of(desc));
303 Some(next)
304 }
305}
306
307fn program_hash(desc: &RenderPipelineDesc) -> u64 {
309 use std::hash::{Hash, Hasher};
310 let mut hasher = std::collections::hash_map::DefaultHasher::new();
311 desc.shader.hash(&mut hasher);
312 desc.vs.hash(&mut hasher);
313 desc.fs.hash(&mut hasher);
314 desc.vertex_layouts.hash(&mut hasher);
315 desc.topology.hash(&mut hasher);
316 hasher.finish()
317}
318
319#[derive(Debug)]
321struct KeyPacker {
322 programs: ProgramTable,
323 formats: AxisTable<wgpu::TextureFormat>,
324 blends: AxisTable<Option<wgpu::BlendState>>,
325 depths: AxisTable<Option<wgpu::DepthStencilState>>,
326}
327
328impl Default for KeyPacker {
329 fn default() -> Self {
330 Self {
331 programs: ProgramTable::default(),
332 formats: AxisTable::new(FORMAT_BITS),
333 blends: AxisTable::new(BLEND_BITS),
334 depths: AxisTable::new(DEPTH_BITS),
335 }
336 }
337}
338
339impl KeyPacker {
340 fn key_for(&mut self, desc: &RenderPipelineDesc) -> Option<u64> {
343 let sample = sample_field(desc.sample_count)?;
346 let program = self.programs.intern(desc)?;
347 let format = self.formats.intern(&desc.format)?;
348 let blend = self.blends.intern(&desc.blend)?;
349 let depth = self.depths.intern(&desc.depth)?;
350 Some(
351 (program << PROGRAM_SHIFT)
352 | (format << FORMAT_SHIFT)
353 | (blend << BLEND_SHIFT)
354 | (depth << DEPTH_SHIFT)
355 | (sample << SAMPLE_SHIFT),
356 )
357 }
358}
359
360fn sample_field(sample_count: u32) -> Option<u64> {
362 if sample_count == 0 || !sample_count.is_power_of_two() {
363 return None;
364 }
365 let field = u64::from(sample_count.trailing_zeros());
366 (field < 1 << SAMPLE_BITS).then_some(field)
367}
368
369#[derive(Clone, Copy, Debug, PartialEq, Eq)]
371enum JobState {
372 Pending,
374 Building,
376 Done,
378 Failed,
382}
383
384#[derive(Debug)]
386struct Job {
387 desc: RenderPipelineDesc,
388 state: JobState,
389}
390
391#[derive(Debug)]
393struct Inner<P> {
394 built: HashMap<u64, Arc<P>>,
395 jobs: HashMap<u64, Job>,
396 queue: VecDeque<u64>,
399 compiles: u64,
401 worker_live: bool,
405 shutdown: bool,
408 worker_epoch: u64,
411}
412
413impl<P> Default for Inner<P> {
414 fn default() -> Self {
415 Self {
416 built: HashMap::new(),
417 jobs: HashMap::new(),
418 queue: VecDeque::new(),
419 compiles: 0,
420 worker_live: false,
421 shutdown: false,
422 worker_epoch: 0,
423 }
424 }
425}
426
427impl<P> Inner<P> {
428 fn claim_next_pending(&mut self) -> Option<(u64, RenderPipelineDesc)> {
435 while !self.shutdown
436 && let Some(key) = self.queue.pop_front()
437 {
438 if let Some(job) = self.jobs.get_mut(&key)
439 && job.state == JobState::Pending
440 {
441 job.state = JobState::Building;
442 return Some((key, job.desc.clone()));
443 }
444 }
445 self.worker_live = false;
446 None
447 }
448}
449
450#[derive(Debug)]
453struct Shared<P> {
454 inner: Mutex<Inner<P>>,
455 built: Condvar,
456}
457
458impl<P> Default for Shared<P> {
459 fn default() -> Self {
460 Self {
461 inner: Mutex::new(Inner::default()),
462 built: Condvar::new(),
463 }
464 }
465}
466
467fn lock<T>(mutex: &Mutex<T>) -> MutexGuard<'_, T> {
471 mutex.lock().unwrap_or_else(PoisonError::into_inner)
472}
473
474struct BuildGuard<'a, P> {
483 shared: &'a Shared<P>,
484 key: u64,
485 armed: bool,
486}
487
488impl<'a, P> BuildGuard<'a, P> {
489 fn new(shared: &'a Shared<P>, key: u64) -> Self {
490 Self {
491 shared,
492 key,
493 armed: true,
494 }
495 }
496
497 fn publish(mut self, pipeline: Arc<P>) {
499 self.armed = false;
500 let mut inner = lock(&self.shared.inner);
501 inner.built.insert(self.key, pipeline);
502 if let Some(job) = inner.jobs.get_mut(&self.key) {
503 job.state = JobState::Done;
504 }
505 inner.compiles += 1;
506 drop(inner);
507 self.shared.built.notify_all();
508 }
509}
510
511struct LiveWorkerGuard<'a, P> {
518 shared: &'a Shared<P>,
519 epoch: u64,
520}
521
522impl<P> Drop for LiveWorkerGuard<'_, P> {
523 fn drop(&mut self) {
524 let mut inner = lock(&self.shared.inner);
525 if inner.worker_epoch == self.epoch {
526 inner.worker_live = false;
527 }
528 }
529}
530
531impl<P> Drop for BuildGuard<'_, P> {
532 fn drop(&mut self) {
533 if !self.armed {
534 return;
535 }
536 let mut inner = lock(&self.shared.inner);
537 if let Some(job) = inner.jobs.get_mut(&self.key) {
538 job.state = JobState::Failed;
539 }
540 drop(inner);
541 self.shared.built.notify_all();
542 }
543}
544
545#[derive(Debug)]
552struct VariantCache<P> {
553 shared: Arc<Shared<P>>,
554 local: HashMap<u64, Arc<P>>,
557 overflow: Option<Arc<P>>,
560 keys: KeyPacker,
561 warned_key_space: bool,
562 worker: Option<std::thread::JoinHandle<()>>,
564}
565
566impl<P> Default for VariantCache<P> {
567 fn default() -> Self {
568 Self {
569 shared: Arc::new(Shared::default()),
570 local: HashMap::new(),
571 overflow: None,
572 keys: KeyPacker::default(),
573 warned_key_space: false,
574 worker: None,
575 }
576 }
577}
578
579impl<P> Drop for VariantCache<P> {
580 fn drop(&mut self) {
581 self.shutdown();
582 }
583}
584
585impl<P> VariantCache<P> {
586 fn shared(&self) -> Arc<Shared<P>> {
588 Arc::clone(&self.shared)
589 }
590
591 fn get_or_create<F>(&mut self, desc: &RenderPipelineDesc, compile: F) -> &P
593 where
594 F: Fn(&RenderPipelineDesc) -> P,
595 {
596 let Some(key) = self.keys.key_for(desc) else {
597 if !self.warned_key_space {
598 self.warned_key_space = true;
599 log::error!(
600 "frust-gpu: pipeline variant key space exhausted (or an unusable \
601 sample_count {}); this variant is compiled per request instead of cached",
602 desc.sample_count
603 );
604 }
605 self.overflow = Some(Arc::new(compile(desc)));
606 return self
607 .overflow
608 .as_deref()
609 .expect("the uncached variant was just stored");
610 };
611
612 if !self.local.contains_key(&key) {
613 let built = self.resolve(key, desc, &compile);
614 self.local.insert(key, built);
615 }
616 &self.local[&key]
617 }
618
619 fn resolve<F>(&self, key: u64, desc: &RenderPipelineDesc, compile: &F) -> Arc<P>
621 where
622 F: Fn(&RenderPipelineDesc) -> P,
623 {
624 let mut inner = lock(&self.shared.inner);
625 loop {
626 if let Some(built) = inner.built.get(&key) {
627 return Arc::clone(built);
628 }
629 match inner.jobs.get(&key).map(|job| job.state) {
630 Some(JobState::Building) => {
633 inner = self
634 .shared
635 .built
636 .wait(inner)
637 .unwrap_or_else(PoisonError::into_inner);
638 }
639 Some(JobState::Pending | JobState::Failed) => {
644 if let Some(job) = inner.jobs.get_mut(&key) {
645 job.state = JobState::Building;
646 }
647 drop(inner);
648 return self.build(key, desc, compile);
649 }
650 Some(JobState::Done) | None => {
653 inner.jobs.insert(
654 key,
655 Job {
656 desc: desc.clone(),
657 state: JobState::Building,
658 },
659 );
660 drop(inner);
661 return self.build(key, desc, compile);
662 }
663 }
664 }
665 }
666
667 fn build<F>(&self, key: u64, desc: &RenderPipelineDesc, compile: &F) -> Arc<P>
669 where
670 F: Fn(&RenderPipelineDesc) -> P,
671 {
672 let guard = BuildGuard::new(&self.shared, key);
673 let pipeline = Arc::new(compile(desc));
674 guard.publish(Arc::clone(&pipeline));
675 pipeline
676 }
677
678 fn enqueue(&mut self, descs: &[RenderPipelineDesc]) -> bool {
686 let mut inner = lock(&self.shared.inner);
687 for desc in descs {
688 let Some(key) = self.keys.key_for(desc) else {
689 log::error!(
690 "frust-gpu: warm-up variant has no representable key (sample_count {}); \
691 it will be compiled on first use instead",
692 desc.sample_count
693 );
694 continue;
695 };
696 if inner.built.contains_key(&key) || inner.jobs.contains_key(&key) {
697 continue;
698 }
699 inner.jobs.insert(
700 key,
701 Job {
702 desc: desc.clone(),
703 state: JobState::Pending,
704 },
705 );
706 inner.queue.push_back(key);
707 }
708 if inner.shutdown || inner.worker_live || inner.queue.is_empty() {
709 return false;
710 }
711 inner.worker_live = true;
712 inner.worker_epoch += 1;
713 true
714 }
715
716 fn drain_queue<F>(shared: &Arc<Shared<P>>, compile: F)
723 where
724 F: Fn(&RenderPipelineDesc) -> P,
725 {
726 let _live = LiveWorkerGuard {
727 shared,
728 epoch: lock(&shared.inner).worker_epoch,
729 };
730 loop {
731 let claimed = lock(&shared.inner).claim_next_pending();
732 let Some((key, desc)) = claimed else {
733 return;
734 };
735 let built = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
736 let guard = BuildGuard::new(shared, key);
737 guard.publish(Arc::new(compile(&desc)));
738 }));
739 if built.is_err() {
740 log::error!(
741 "frust-gpu: a warm-up compile panicked; that variant is left to be \
742 rebuilt on request and the rest of the queue continues"
743 );
744 }
745 }
746 }
747
748 fn warm_up<F>(&mut self, descs: &[RenderPipelineDesc], make_compile: impl Fn() -> F)
754 where
755 P: Send + Sync + 'static,
756 F: Fn(&RenderPipelineDesc) -> P + Send + 'static,
757 {
758 if !self.enqueue(descs) {
759 return;
760 }
761 self.join_worker();
765
766 let worker_shared = self.shared();
767 let worker_compile = make_compile();
768 let spawned = std::thread::Builder::new()
769 .name("frust-gpu pipeline warm-up".to_string())
770 .spawn(move || VariantCache::drain_queue(&worker_shared, worker_compile));
771 match spawned {
772 Ok(handle) => self.worker = Some(handle),
773 Err(err) => {
774 log::warn!(
775 "frust-gpu: could not spawn the pipeline warm-up thread ({err}); \
776 building the listed variants inline"
777 );
778 VariantCache::drain_queue(&self.shared(), make_compile());
779 }
780 }
781 }
782
783 fn shutdown(&mut self) {
790 lock(&self.shared.inner).shutdown = true;
791 self.join_worker();
792 }
793
794 fn join_worker(&mut self) {
796 if let Some(worker) = self.worker.take()
797 && worker.join().is_err()
798 {
799 log::error!("frust-gpu: the pipeline warm-up worker panicked");
800 }
801 }
802
803 fn compiled_variants(&self) -> u64 {
805 lock(&self.shared.inner).compiles
806 }
807
808 fn queued_variants(&self) -> usize {
810 let inner = lock(&self.shared.inner);
811 inner
812 .jobs
813 .values()
814 .filter(|job| job.state == JobState::Pending)
815 .count()
816 }
817}
818
819#[derive(Debug)]
828pub struct PipelineCache {
829 core: VariantCache<wgpu::RenderPipeline>,
830 shaders: Arc<ShaderLibrary>,
831 driver_cache: Option<wgpu::PipelineCache>,
832}
833
834impl PipelineCache {
835 #[must_use]
842 pub fn new(shaders: Arc<ShaderLibrary>, driver_cache: Option<wgpu::PipelineCache>) -> Self {
843 Self {
844 core: VariantCache::default(),
845 shaders,
846 driver_cache,
847 }
848 }
849
850 #[must_use]
852 pub fn shaders(&self) -> &ShaderLibrary {
853 &self.shaders
854 }
855
856 pub fn get_or_create(
868 &mut self,
869 device: &wgpu::Device,
870 desc: &RenderPipelineDesc,
871 ) -> &wgpu::RenderPipeline {
872 let shaders = &self.shaders;
873 let driver_cache = self.driver_cache.as_ref();
874 self.core.get_or_create(desc, |desc| {
875 build_render_pipeline(device, shaders, driver_cache, desc)
876 })
877 }
878
879 pub fn warm_up(&mut self, device: &wgpu::Device, descs: &[RenderPipelineDesc]) {
894 let shaders = Arc::clone(&self.shaders);
895 let driver_cache = self.driver_cache.clone();
896 self.core
897 .warm_up(descs, || compiler(device, &shaders, driver_cache.as_ref()));
898 }
899
900 pub fn shutdown(&mut self) {
908 self.core.shutdown();
909 }
910
911 #[must_use]
913 pub fn compiled_variants(&self) -> u64 {
914 self.core.compiled_variants()
915 }
916
917 #[must_use]
919 pub fn queued_variants(&self) -> usize {
920 self.core.queued_variants()
921 }
922}
923
924fn compiler(
928 device: &wgpu::Device,
929 shaders: &Arc<ShaderLibrary>,
930 driver_cache: Option<&wgpu::PipelineCache>,
931) -> impl Fn(&RenderPipelineDesc) -> wgpu::RenderPipeline + Send + 'static + use<> {
932 let device = device.clone();
933 let shaders = Arc::clone(shaders);
934 let driver_cache = driver_cache.cloned();
935 move |desc| build_render_pipeline(&device, &shaders, driver_cache.as_ref(), desc)
936}
937
938fn build_render_pipeline(
941 device: &wgpu::Device,
942 shaders: &ShaderLibrary,
943 driver_cache: Option<&wgpu::PipelineCache>,
944 desc: &RenderPipelineDesc,
945) -> wgpu::RenderPipeline {
946 let name = shaders.name_of(desc.shader).unwrap_or("<unknown shader>");
947 let module = shaders
948 .get(desc.shader)
949 .expect("a RenderPipelineDesc's ShaderId must come from this cache's ShaderLibrary");
950 let layouts: Vec<Option<wgpu::VertexBufferLayout<'_>>> = desc
955 .vertex_layouts
956 .iter()
957 .map(|layout| Some(VertexLayout::as_wgpu(layout)))
958 .collect();
959 let label = format!(
960 "frust-gpu pipeline: {name} [{:?}, msaa x{}]",
961 desc.format, desc.sample_count
962 );
963
964 device.create_render_pipeline(&wgpu::RenderPipelineDescriptor {
965 label: Some(&label),
966 layout: None,
968 vertex: wgpu::VertexState {
969 module,
970 entry_point: Some(desc.vs.as_ref()),
971 compilation_options: wgpu::PipelineCompilationOptions::default(),
972 buffers: &layouts,
973 },
974 primitive: wgpu::PrimitiveState {
975 topology: desc.topology,
976 ..Default::default()
977 },
978 depth_stencil: desc.depth.clone(),
979 multisample: wgpu::MultisampleState {
980 count: desc.sample_count,
981 mask: !0,
982 alpha_to_coverage_enabled: false,
983 },
984 fragment: Some(wgpu::FragmentState {
985 module,
986 entry_point: Some(desc.fs.as_ref()),
987 compilation_options: wgpu::PipelineCompilationOptions::default(),
988 targets: &[Some(wgpu::ColorTargetState {
989 format: desc.format,
990 blend: desc.blend,
991 write_mask: wgpu::ColorWrites::ALL,
992 })],
993 }),
994 multiview_mask: None,
995 cache: driver_cache,
996 })
997}
998
999#[cfg(test)]
1000mod tests {
1001 use super::*;
1002 use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
1003
1004 #[derive(Debug, PartialEq, Eq)]
1007 struct FakePipeline {
1008 sample_count: u32,
1009 serial: u64,
1010 }
1011
1012 #[derive(Default)]
1015 struct Counter(AtomicU64);
1016
1017 impl Counter {
1018 fn compile(&self, desc: &RenderPipelineDesc) -> FakePipeline {
1019 let serial = self.0.fetch_add(1, Ordering::SeqCst);
1020 FakePipeline {
1021 sample_count: desc.sample_count,
1022 serial,
1023 }
1024 }
1025
1026 fn count(&self) -> u64 {
1027 self.0.load(Ordering::SeqCst)
1028 }
1029 }
1030
1031 fn desc() -> RenderPipelineDesc {
1032 RenderPipelineDesc::new(
1033 ShaderId::from_raw(0),
1034 "vs_main",
1035 "fs_main",
1036 wgpu::TextureFormat::Rgba8Unorm,
1037 )
1038 }
1039
1040 #[test]
1041 fn a_repeated_request_compiles_once() {
1042 let mut cache = VariantCache::<FakePipeline>::default();
1043 let counter = Counter::default();
1044 let desc = desc();
1045
1046 let first = cache.get_or_create(&desc, |d| counter.compile(d)).serial;
1047 for _ in 0..8 {
1048 let again = cache.get_or_create(&desc, |d| counter.compile(d)).serial;
1049 assert_eq!(again, first, "every repeat must return the same pipeline");
1050 }
1051 assert_eq!(counter.count(), 1);
1052 assert_eq!(cache.compiled_variants(), 1);
1053 }
1054
1055 #[test]
1056 fn each_key_axis_is_its_own_variant() {
1057 let base = desc();
1058 let variants = [
1059 RenderPipelineDesc {
1060 blend: Some(wgpu::BlendState::ALPHA_BLENDING),
1061 ..base.clone()
1062 },
1063 RenderPipelineDesc {
1064 format: wgpu::TextureFormat::Bgra8Unorm,
1065 ..base.clone()
1066 },
1067 RenderPipelineDesc {
1068 sample_count: 4,
1069 ..base.clone()
1070 },
1071 RenderPipelineDesc {
1072 depth: Some(wgpu::DepthStencilState {
1073 format: wgpu::TextureFormat::Depth32Float,
1074 depth_write_enabled: Some(true),
1075 depth_compare: Some(wgpu::CompareFunction::Less),
1076 stencil: wgpu::StencilState::default(),
1077 bias: wgpu::DepthBiasState::default(),
1078 }),
1079 ..base.clone()
1080 },
1081 RenderPipelineDesc {
1082 shader: ShaderId::from_raw(1),
1083 ..base.clone()
1084 },
1085 RenderPipelineDesc {
1088 fs: "fs_other".into(),
1089 ..base.clone()
1090 },
1091 RenderPipelineDesc {
1093 vertex_layouts: vec![VertexLayout::per_vertex(
1094 16,
1095 vec![wgpu::VertexAttribute {
1096 format: wgpu::VertexFormat::Float32x4,
1097 offset: 0,
1098 shader_location: 0,
1099 }],
1100 )],
1101 ..base.clone()
1102 },
1103 RenderPipelineDesc {
1104 topology: wgpu::PrimitiveTopology::LineList,
1105 ..base.clone()
1106 },
1107 ];
1108
1109 let mut packer = KeyPacker::default();
1110 let base_key = packer.key_for(&base).expect("base key");
1111 let mut seen = vec![base_key];
1112 for variant in &variants {
1113 let key = packer.key_for(variant).expect("variant key");
1114 assert!(
1115 !seen.contains(&key),
1116 "{variant:?} must not reuse an existing key"
1117 );
1118 seen.push(key);
1119 }
1120
1121 let mut cache = VariantCache::<FakePipeline>::default();
1123 let counter = Counter::default();
1124 cache.get_or_create(&base, |d| counter.compile(d));
1125 for variant in &variants {
1126 cache.get_or_create(variant, |d| counter.compile(d));
1127 }
1128 assert_eq!(counter.count(), 1 + variants.len() as u64);
1129 }
1130
1131 #[test]
1132 fn a_key_is_stable_across_repeated_packing() {
1133 let mut packer = KeyPacker::default();
1134 let desc = desc();
1135 let first = packer.key_for(&desc).expect("key");
1136 for _ in 0..4 {
1137 assert_eq!(packer.key_for(&desc), Some(first));
1138 }
1139 }
1140
1141 #[test]
1142 fn an_unusable_sample_count_has_no_key_and_is_not_cached() {
1143 let mut packer = KeyPacker::default();
1144 assert_eq!(
1145 packer.key_for(&RenderPipelineDesc {
1146 sample_count: 3,
1147 ..desc()
1148 }),
1149 None
1150 );
1151 assert_eq!(
1152 packer.key_for(&RenderPipelineDesc {
1153 sample_count: 0,
1154 ..desc()
1155 }),
1156 None
1157 );
1158
1159 let mut cache = VariantCache::<FakePipeline>::default();
1162 let counter = Counter::default();
1163 let odd = RenderPipelineDesc {
1164 sample_count: 3,
1165 ..desc()
1166 };
1167 assert_eq!(
1168 cache
1169 .get_or_create(&odd, |d| counter.compile(d))
1170 .sample_count,
1171 3
1172 );
1173 cache.get_or_create(&odd, |d| counter.compile(d));
1174 assert_eq!(counter.count(), 2, "an unkeyable variant is never cached");
1175 }
1176
1177 #[test]
1178 fn draining_the_queue_builds_every_listed_variant_once() {
1179 let mut cache = VariantCache::<FakePipeline>::default();
1180 let counter = Counter::default();
1181 let listed: Vec<_> = [1u32, 2, 4]
1182 .into_iter()
1183 .map(|sample_count| RenderPipelineDesc {
1184 sample_count,
1185 ..desc()
1186 })
1187 .collect();
1188
1189 cache.enqueue(&listed);
1190 assert_eq!(cache.queued_variants(), 3);
1191 VariantCache::drain_queue(&cache.shared(), |d| counter.compile(d));
1192 assert_eq!(counter.count(), 3);
1193 assert_eq!(cache.queued_variants(), 0);
1194
1195 for desc in &listed {
1197 cache.get_or_create(desc, |d| counter.compile(d));
1198 }
1199 assert_eq!(counter.count(), 3, "a warmed variant must not recompile");
1200 }
1201
1202 #[test]
1203 fn enqueueing_the_same_variant_twice_queues_one_job() {
1204 let mut cache = VariantCache::<FakePipeline>::default();
1205 let listed = [desc(), desc()];
1206 cache.enqueue(&listed);
1207 cache.enqueue(&listed);
1208 assert_eq!(cache.queued_variants(), 1);
1209 }
1210
1211 #[test]
1212 fn a_render_thread_request_steals_a_queued_job_and_the_queue_skips_it() {
1213 let mut cache = VariantCache::<FakePipeline>::default();
1214 let counter = Counter::default();
1215 let listed: Vec<_> = [1u32, 2, 4]
1216 .into_iter()
1217 .map(|sample_count| RenderPipelineDesc {
1218 sample_count,
1219 ..desc()
1220 })
1221 .collect();
1222 cache.enqueue(&listed);
1223
1224 let stolen = cache
1227 .get_or_create(&listed[2], |d| counter.compile(d))
1228 .serial;
1229 assert_eq!(stolen, 0, "the steal is the first compile to run");
1230 assert_eq!(counter.count(), 1);
1231 assert_eq!(cache.queued_variants(), 2, "the stolen job is done");
1232
1233 VariantCache::drain_queue(&cache.shared(), |d| counter.compile(d));
1236 assert_eq!(
1237 counter.count(),
1238 3,
1239 "the stolen variant must not compile twice"
1240 );
1241 assert_eq!(
1242 cache
1243 .get_or_create(&listed[2], |d| counter.compile(d))
1244 .serial,
1245 stolen
1246 );
1247 assert_eq!(counter.count(), 3);
1248 }
1249
1250 #[test]
1251 fn a_worker_thread_drain_publishes_to_the_render_thread() {
1252 let mut cache = VariantCache::<FakePipeline>::default();
1253 let listed: Vec<_> = (0..6)
1254 .map(|i| RenderPipelineDesc {
1255 shader: ShaderId::from_raw(i),
1256 ..desc()
1257 })
1258 .collect();
1259 cache.enqueue(&listed);
1260
1261 let shared = cache.shared();
1262 let worker_counter = Arc::new(Counter::default());
1263 let thread_counter = Arc::clone(&worker_counter);
1264 let worker = std::thread::spawn(move || {
1265 VariantCache::drain_queue(&shared, move |d| thread_counter.compile(d));
1266 });
1267 worker.join().expect("the warm-up worker must not panic");
1268
1269 assert_eq!(worker_counter.count(), 6);
1270 assert_eq!(cache.compiled_variants(), 6);
1271
1272 let render_counter = Counter::default();
1275 for desc in &listed {
1276 cache.get_or_create(desc, |d| render_counter.compile(d));
1277 }
1278 assert_eq!(render_counter.count(), 0);
1279 }
1280
1281 #[test]
1282 fn an_abandoned_build_marks_the_job_failed_rather_than_pending() {
1283 let cache = VariantCache::<FakePipeline>::default();
1284 let shared = cache.shared();
1285 let guard = BuildGuard::new(&shared, 42);
1286 lock(&shared.inner).jobs.insert(
1287 42,
1288 Job {
1289 desc: desc(),
1290 state: JobState::Building,
1291 },
1292 );
1293 drop(guard);
1295 assert_eq!(
1296 lock(&shared.inner).jobs.get(&42).map(|job| job.state),
1297 Some(JobState::Failed),
1298 "an abandoned build must not leave a Pending job no worker can claim"
1299 );
1300 }
1301
1302 struct PanicOnce {
1305 counter: Arc<Counter>,
1306 armed: AtomicBool,
1307 }
1308
1309 impl PanicOnce {
1310 fn new(counter: &Arc<Counter>) -> Self {
1311 Self {
1312 counter: Arc::clone(counter),
1313 armed: AtomicBool::new(true),
1314 }
1315 }
1316
1317 fn compile(&self, desc: &RenderPipelineDesc) -> FakePipeline {
1318 assert!(
1319 !self.armed.swap(false, Ordering::SeqCst),
1320 "the fake compile fails its first call"
1321 );
1322 self.counter.compile(desc)
1323 }
1324 }
1325
1326 #[test]
1327 fn a_panicking_warm_up_compile_costs_one_variant_and_not_the_queue() {
1328 let mut cache = VariantCache::<FakePipeline>::default();
1329 let counter = Arc::new(Counter::default());
1330 let failing = PanicOnce::new(&counter);
1331 let listed: Vec<_> = [1u32, 2, 4]
1332 .into_iter()
1333 .map(|sample_count| RenderPipelineDesc {
1334 sample_count,
1335 ..desc()
1336 })
1337 .collect();
1338
1339 cache.enqueue(&listed);
1340 VariantCache::drain_queue(&cache.shared(), |d| failing.compile(d));
1341
1342 assert_eq!(counter.count(), 2);
1345 assert_eq!(
1346 cache.queued_variants(),
1347 0,
1348 "a failed job must not be counted as still queued forever"
1349 );
1350
1351 assert_eq!(
1353 cache
1354 .get_or_create(&listed[0], |d| counter.compile(d))
1355 .sample_count,
1356 1
1357 );
1358 assert_eq!(counter.count(), 3);
1359
1360 for desc in &listed {
1362 cache.get_or_create(desc, |d| counter.compile(d));
1363 }
1364 assert_eq!(counter.count(), 3);
1365 }
1366
1367 #[test]
1368 fn a_panicking_inline_build_surfaces_to_the_requester_and_leaves_it_rebuildable() {
1369 let mut cache = VariantCache::<FakePipeline>::default();
1370 let counter = Arc::new(Counter::default());
1371 let failing = PanicOnce::new(&counter);
1372 let desc = desc();
1373 cache.enqueue(std::slice::from_ref(&desc));
1374
1375 let stolen = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
1378 cache.get_or_create(&desc, |d| failing.compile(d));
1379 }));
1380 assert!(
1381 stolen.is_err(),
1382 "an inline compile must not swallow a panic"
1383 );
1384 assert_eq!(counter.count(), 0);
1385
1386 assert_eq!(cache.get_or_create(&desc, |d| counter.compile(d)).serial, 0);
1389 assert_eq!(counter.count(), 1);
1390 assert_eq!(cache.queued_variants(), 0);
1391 }
1392
1393 #[derive(Default)]
1397 struct Latch {
1398 open: Mutex<bool>,
1399 changed: Condvar,
1400 }
1401
1402 impl Latch {
1403 fn wait(&self) {
1404 let mut open = lock(&self.open);
1405 while !*open {
1406 open = self
1407 .changed
1408 .wait(open)
1409 .unwrap_or_else(PoisonError::into_inner);
1410 }
1411 }
1412
1413 fn open(&self) {
1414 *lock(&self.open) = true;
1415 self.changed.notify_all();
1416 }
1417 }
1418
1419 fn shader_variants(count: u32) -> Vec<RenderPipelineDesc> {
1420 (0..count)
1421 .map(|i| RenderPipelineDesc {
1422 shader: ShaderId::from_raw(i),
1423 ..desc()
1424 })
1425 .collect()
1426 }
1427
1428 #[test]
1429 fn repeated_warm_up_calls_keep_at_most_one_worker() {
1430 let mut cache = VariantCache::<FakePipeline>::default();
1431 let counter = Arc::new(Counter::default());
1432 let latch = Arc::new(Latch::default());
1433 let workers = AtomicU64::new(0);
1434 let listed = shader_variants(6);
1435
1436 let make_compile = || {
1439 workers.fetch_add(1, Ordering::SeqCst);
1440 let counter = Arc::clone(&counter);
1441 let latch = Arc::clone(&latch);
1442 let held = AtomicBool::new(false);
1443 move |d: &RenderPipelineDesc| {
1444 if !held.swap(true, Ordering::SeqCst) {
1445 latch.wait();
1446 }
1447 counter.compile(d)
1448 }
1449 };
1450
1451 cache.warm_up(&listed[..1], make_compile);
1454 for _ in 0..8 {
1455 cache.warm_up(&listed, make_compile);
1456 }
1457 assert_eq!(
1458 workers.load(Ordering::SeqCst),
1459 1,
1460 "a live worker must absorb further warm-up calls, not be joined by more"
1461 );
1462
1463 latch.open();
1464 cache.shutdown();
1465
1466 for desc in &listed {
1469 cache.get_or_create(desc, |d| counter.compile(d));
1470 }
1471 assert_eq!(counter.count(), listed.len() as u64);
1472 assert_eq!(cache.compiled_variants(), listed.len() as u64);
1473 }
1474
1475 #[test]
1476 fn dropping_the_cache_joins_its_worker() {
1477 let device = Arc::new(());
1481 let counter = Arc::new(Counter::default());
1482 let listed = shader_variants(4);
1483
1484 let mut cache = VariantCache::<FakePipeline>::default();
1485 cache.warm_up(&listed, || {
1486 let device = Arc::clone(&device);
1487 let counter = Arc::clone(&counter);
1488 move |d: &RenderPipelineDesc| {
1489 let _held = Arc::clone(&device);
1490 counter.compile(d)
1491 }
1492 });
1493 drop(cache);
1494
1495 assert_eq!(
1496 Arc::strong_count(&device),
1497 1,
1498 "no warm-up worker may outlive the cache that started it"
1499 );
1500 }
1501
1502 #[test]
1503 fn warming_up_after_shutdown_starts_no_worker_and_still_builds_on_request() {
1504 let mut cache = VariantCache::<FakePipeline>::default();
1505 let counter = Arc::new(Counter::default());
1506 let workers = AtomicU64::new(0);
1507 let listed = shader_variants(3);
1508
1509 cache.shutdown();
1510 cache.shutdown();
1512 cache.warm_up(&listed, || {
1513 workers.fetch_add(1, Ordering::SeqCst);
1514 let counter = Arc::clone(&counter);
1515 move |d: &RenderPipelineDesc| counter.compile(d)
1516 });
1517
1518 assert_eq!(workers.load(Ordering::SeqCst), 0);
1519 for desc in &listed {
1520 cache.get_or_create(desc, |d| counter.compile(d));
1521 }
1522 assert_eq!(counter.count(), listed.len() as u64);
1523 }
1524
1525 fn drain_error_scope(
1529 device: &wgpu::Device,
1530 scope: wgpu::ErrorScopeGuard,
1531 ) -> Option<wgpu::Error> {
1532 use std::task::{Context, Poll, Waker};
1533
1534 let waker = Waker::noop();
1535 let mut cx = Context::from_waker(waker);
1536 let mut future = std::pin::pin!(scope.pop());
1537 loop {
1538 match future.as_mut().poll(&mut cx) {
1539 Poll::Ready(error) => return error,
1540 Poll::Pending => {
1541 let _ = device.poll(wgpu::PollType::wait_indefinitely());
1542 }
1543 }
1544 }
1545 }
1546
1547 #[test]
1553 #[ignore = "requires a GPU (Metal/Vulkan); run locally with `cargo test -p frust-gpu -- --ignored`"]
1554 fn creates_one_real_pipeline() {
1555 const WGSL: &str = r#"
1556@vertex
1557fn vs_main(@builtin(vertex_index) i: u32) -> @builtin(position) vec4<f32> {
1558 let uv = vec2<f32>(f32((i << 1u) & 2u), f32(i & 2u));
1559 return vec4<f32>(uv * 2.0 - 1.0, 0.0, 1.0);
1560}
1561@fragment
1562fn fs_main() -> @location(0) vec4<f32> {
1563 return vec4<f32>(0.0, 1.0, 0.0, 1.0);
1564}
1565"#;
1566
1567 let (device, _queue) = pollster::block_on(async {
1568 let instance = wgpu::Instance::new(
1569 wgpu::InstanceDescriptor::new_without_display_handle_from_env(),
1570 );
1571 let adapter = wgpu::util::initialize_adapter_from_env_or_default(&instance, None)
1575 .await
1576 .expect("no compatible GPU adapter");
1577 println!("frust-gpu pipeline test adapter: {:?}", adapter.get_info());
1578 adapter
1579 .request_device(&wgpu::DeviceDescriptor {
1580 label: Some("frust-gpu pipeline test device"),
1581 required_features: wgpu::Features::empty(),
1582 required_limits: wgpu::Limits::default(),
1583 ..Default::default()
1584 })
1585 .await
1586 .expect("failed to create the device")
1587 });
1588
1589 let mut library = ShaderLibrary::new();
1590 let shader = library.insert_wgsl(&device, "fullscreen-green", WGSL);
1591 let mut cache = PipelineCache::new(Arc::new(library), None);
1592
1593 let desc = RenderPipelineDesc::new(
1594 shader,
1595 "vs_main",
1596 "fs_main",
1597 wgpu::TextureFormat::Rgba8Unorm,
1598 );
1599 let scope = device.push_error_scope(wgpu::ErrorFilter::Validation);
1600 for _ in 0..2 {
1601 let _pipeline = cache.get_or_create(&device, &desc);
1602 }
1603 let error = drain_error_scope(&device, scope);
1604
1605 assert!(error.is_none(), "pipeline creation raised {error:?}");
1606 assert_eq!(
1607 cache.compiled_variants(),
1608 1,
1609 "a repeat request must reuse the pipeline"
1610 );
1611 }
1612}