use std::sync::Arc;
use frust_gpu::lint::PipelineLayoutDesc;
use frust_gpu::{PipelineCache, RenderPipelineDesc, ShaderId, ShaderLibrary};
use super::pipelines::{FS_MAIN, VS_MAIN};
use crate::OutputAlpha;
pub const UNPREMULTIPLY: &str = include_str!("../../shaders/unpremultiply.wgsl");
pub const UNPREMULTIPLY_NAME: &str = "frust-engine unpremultiply";
pub const ALPHA_FLOOR: f32 = 1e-4;
const FULLSCREEN_VERTICES: u32 = 3;
const UNPREMULTIPLY_BIND_GROUPS: usize = 1;
#[derive(Debug)]
pub struct UnpremultiplyPass {
pipeline: wgpu::RenderPipeline,
bind_group_layout: wgpu::BindGroupLayout,
format: wgpu::TextureFormat,
}
impl UnpremultiplyPass {
#[must_use]
pub fn selected_by(output: OutputAlpha) -> bool {
matches!(output, OutputAlpha::Straight)
}
#[must_use]
pub fn new(
device: &wgpu::Device,
format: wgpu::TextureFormat,
driver_cache: Option<&wgpu::PipelineCache>,
) -> Self {
let mut library = ShaderLibrary::new();
let shader = library.insert_wgsl(device, UNPREMULTIPLY_NAME, UNPREMULTIPLY);
let mut pipelines = PipelineCache::new(Arc::new(library), driver_cache.cloned());
let pipeline = pipelines
.get_or_create(device, &Self::desc(shader, format))
.clone();
let bind_group_layout = pipeline.get_bind_group_layout(0);
Self {
pipeline,
bind_group_layout,
format,
}
}
#[must_use]
pub fn format(&self) -> wgpu::TextureFormat {
self.format
}
pub fn record(
&self,
device: &wgpu::Device,
encoder: &mut wgpu::CommandEncoder,
source_view: &wgpu::TextureView,
target_view: &wgpu::TextureView,
timestamp_writes: Option<wgpu::RenderPassTimestampWrites<'_>>,
) {
let bind_group = device.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("frust-engine unpremultiply bind group"),
layout: &self.bind_group_layout,
entries: &[wgpu::BindGroupEntry {
binding: 0,
resource: wgpu::BindingResource::TextureView(source_view),
}],
});
let mut pass = encoder.begin_render_pass(&wgpu::RenderPassDescriptor {
label: Some("frust-engine unpremultiply pass"),
color_attachments: &[Some(wgpu::RenderPassColorAttachment {
view: target_view,
depth_slice: None,
resolve_target: None,
ops: wgpu::Operations {
load: wgpu::LoadOp::Clear(wgpu::Color::TRANSPARENT),
store: wgpu::StoreOp::Store,
},
})],
depth_stencil_attachment: None,
timestamp_writes,
occlusion_query_set: None,
multiview_mask: None,
});
pass.set_pipeline(&self.pipeline);
pass.set_bind_group(0, &bind_group, &[]);
pass.draw(0..FULLSCREEN_VERTICES, 0..1);
}
#[must_use]
pub fn layout_desc() -> PipelineLayoutDesc {
PipelineLayoutDesc {
bind_group_count: UNPREMULTIPLY_BIND_GROUPS,
max_vertex_buffers: 0,
total_vertex_attributes: 0,
max_vertex_buffer_stride: 0,
sample_count: 1,
uniform_buffer_sizes: Vec::new(),
}
}
fn desc(shader: ShaderId, format: wgpu::TextureFormat) -> RenderPipelineDesc {
RenderPipelineDesc::new(shader, VS_MAIN, FS_MAIN, format)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn only_a_straight_alpha_target_selects_the_pass() {
assert!(UnpremultiplyPass::selected_by(OutputAlpha::Straight));
assert!(!UnpremultiplyPass::selected_by(OutputAlpha::Premultiplied));
}
fn shader_alpha_floor() -> f32 {
let declaration = UNPREMULTIPLY
.lines()
.find(|line| line.trim_start().starts_with("const ALPHA_FLOOR"))
.expect("the shader declares an ALPHA_FLOOR constant");
declaration
.split('=')
.nth(1)
.expect("the declaration assigns a value")
.trim()
.trim_end_matches(';')
.parse()
.expect("the shader's alpha floor is a float literal")
}
#[test]
fn the_shader_and_rust_alpha_floors_agree() {
let floor = shader_alpha_floor();
assert_eq!(
floor, ALPHA_FLOOR,
"the shader's floor must stay in step with `ALPHA_FLOOR`"
);
assert!(
floor > 0.0,
"a floor of zero would leave `0 / 0` NaNs in the swapchain"
);
assert!(
floor < 1.0 / 255.0,
"the floor must sit below the smallest representable 8-bit alpha, so it never clamps \
a pixel that carries colour"
);
}
#[test]
fn the_program_is_loaded_from_the_shader_directory() {
assert!(UNPREMULTIPLY.contains("fn vs_main"));
assert!(UNPREMULTIPLY.contains("fn fs_main"));
assert!(
!UNPREMULTIPLY.contains("@compute"),
"the conversion is a render pass, never a compute one (E1)"
);
assert!(
!UNPREMULTIPLY.contains("texture_storage_"),
"a swapchain on this tier is RENDER_ATTACHMENT-only (E2)"
);
}
#[test]
fn the_pass_layout_passes_the_downlevel_lint() {
let violations = frust_gpu::lint_pipeline_layout(&UnpremultiplyPass::layout_desc());
assert!(
violations.is_empty(),
"the present pass violates a downlevel design rule: {violations:?}"
);
assert_eq!(
UnpremultiplyPass::layout_desc().bind_group_count,
UNPREMULTIPLY_BIND_GROUPS,
"one bind group: the source texture"
);
}
}