1use crate::config::{
2    AntiDitherShader, FrameRotation, PreprocessShader, PrescaleMode, RendererConfig,
3};
4use crate::renderer::{PipelineShader, REQUIRED_TEXTURE_USAGES, Shaders};
5use jgenesis_common::frontend::{ColorCorrection, DisplayArea, FiniteF64, FrameSize};
6use std::sync::Arc;
7use thiserror::Error;
8use wgpu::util::DeviceExt;
9
10const IDENTITY_VERTICES: u32 = 4;
11
12const SRGB_TEX_VIEW_DESCRIPTOR: wgpu::TextureViewDescriptor<'static> =
13    wgpu::TextureViewDescriptor {
14        label: None,
15        format: Some(wgpu::TextureFormat::Rgba8UnormSrgb),
16        dimension: None,
17        usage: Some(wgpu::TextureUsages::TEXTURE_BINDING),
18        aspect: wgpu::TextureAspect::All,
19        base_mip_level: 0,
20        mip_level_count: None,
21        base_array_layer: 0,
22        array_layer_count: None,
23    };
24
25fn basic_render_pass<'encoder, 'label>(
26    encoder: &'encoder mut wgpu::CommandEncoder,
27    output: &wgpu::Texture,
28    output_format: wgpu::TextureFormat,
29    label: impl Into<wgpu::Label<'label>>,
30) -> wgpu::RenderPass<'encoder> {
31    let output_view = output.create_view(&wgpu::TextureViewDescriptor {
32        format: Some(output_format),
33        usage: Some(wgpu::TextureUsages::RENDER_ATTACHMENT),
34        ..wgpu::TextureViewDescriptor::default()
35    });
36
37    encoder.begin_render_pass(&wgpu::RenderPassDescriptor {
38        label: label.into(),
39        color_attachments: &[Some(wgpu::RenderPassColorAttachment {
40            view: &output_view,
41            depth_slice: None,
42            resolve_target: None,
43            ops: wgpu::Operations {
44                load: wgpu::LoadOp::Clear(wgpu::Color::BLACK),
45                store: wgpu::StoreOp::Store,
46            },
47        })],
48        ..wgpu::RenderPassDescriptor::default()
49    })
50}
51
52pub struct ColorCorrectionShader {
53    output: Arc<wgpu::Texture>,
54    bind_group: wgpu::BindGroup,
55    pipeline: wgpu::RenderPipeline,
56}
57
58impl ColorCorrectionShader {
59    pub fn create(
60        correction: ColorCorrection,
61        input: &wgpu::Texture,
62        device: &wgpu::Device,
63        shaders: &Shaders,
64    ) -> Option<Self> {
65        let (fs_main, screen_gamma) = match correction {
66            ColorCorrection::GbcLcd { screen_gamma } => ("gbc_color_correction", screen_gamma),
67            ColorCorrection::GbaLcd { screen_gamma } => ("gba_color_correction", screen_gamma),
68            ColorCorrection::None => return None,
69        };
70
71        let output = device.create_texture(&wgpu::TextureDescriptor {
72            label: "color_correction_texture".into(),
73            size: input.size(),
74            mip_level_count: 1,
75            sample_count: 1,
76            dimension: wgpu::TextureDimension::D2,
77            format: wgpu::TextureFormat::Rgba8Unorm,
78            usage: *REQUIRED_TEXTURE_USAGES | wgpu::TextureUsages::RENDER_ATTACHMENT,
79            view_formats: &[wgpu::TextureFormat::Rgba8UnormSrgb],
80        });
81
82        let pipeline = device.create_render_pipeline(&wgpu::RenderPipelineDescriptor {
83            label: "color_correction_pipeline".into(),
84            layout: None,
85            vertex: wgpu::VertexState {
86                module: &shaders.identity,
87                entry_point: None,
88                compilation_options: wgpu::PipelineCompilationOptions::default(),
89                buffers: &[],
90            },
91            primitive: wgpu::PrimitiveState {
92                topology: wgpu::PrimitiveTopology::TriangleStrip,
93                strip_index_format: None,
94                front_face: wgpu::FrontFace::Ccw,
95                cull_mode: None,
96                unclipped_depth: false,
97                polygon_mode: wgpu::PolygonMode::Fill,
98                conservative: false,
99            },
100            depth_stencil: None,
101            multisample: wgpu::MultisampleState::default(),
102            fragment: Some(wgpu::FragmentState {
103                module: &shaders.gb_color,
104                entry_point: Some(fs_main),
105                compilation_options: wgpu::PipelineCompilationOptions::default(),
106                targets: &[Some(wgpu::ColorTargetState {
107                    format: wgpu::TextureFormat::Rgba8UnormSrgb,
108                    blend: Some(wgpu::BlendState::REPLACE),
109                    write_mask: wgpu::ColorWrites::ALL,
110                })],
111            }),
112            multiview_mask: None,
113            cache: None,
114        });
115
116        let input_view = input.create_view(&SRGB_TEX_VIEW_DESCRIPTOR);
117        let gamma_buffer = device.create_buffer_init(&wgpu::util::BufferInitDescriptor {
118            label: "color_correction_gamma_buffer".into(),
119            contents: bytemuck::cast_slice(&[f32::from(screen_gamma)]),
120            usage: wgpu::BufferUsages::UNIFORM,
121        });
122
123        let bind_group = device.create_bind_group(&wgpu::BindGroupDescriptor {
124            label: "color_correction_bind_group".into(),
125            layout: &pipeline.get_bind_group_layout(0),
126            entries: &[
127                wgpu::BindGroupEntry {
128                    binding: 0,
129                    resource: wgpu::BindingResource::TextureView(&input_view),
130                },
131                wgpu::BindGroupEntry {
132                    binding: 1,
133                    resource: wgpu::BindingResource::Buffer(
134                        gamma_buffer.as_entire_buffer_binding(),
135                    ),
136                },
137            ],
138        });
139
140        Some(Self { output: Arc::new(output), bind_group, pipeline })
141    }
142}
143
144impl PipelineShader for ColorCorrectionShader {
145    fn draw(&mut self, encoder: &mut wgpu::CommandEncoder) {
146        let mut render_pass = basic_render_pass(
147            encoder,
148            &self.output,
149            wgpu::TextureFormat::Rgba8UnormSrgb,
150            "color_correction_render_pass",
151        );
152
153        render_pass.set_bind_group(0, &self.bind_group, &[]);
154        render_pass.set_pipeline(&self.pipeline);
155
156        render_pass.draw(0..IDENTITY_VERTICES, 0..1);
157    }
158
159    fn output_texture(&self) -> &Arc<wgpu::Texture> {
160        &self.output
161    }
162}
163
164pub struct FrameBlendShader {
165    previous_frame: Arc<wgpu::Texture>,
166    input: Arc<wgpu::Texture>,
167    output: Arc<wgpu::Texture>,
168    bind_group: wgpu::BindGroup,
169    pipeline: wgpu::RenderPipeline,
170    skip_next_frame: bool,
171}
172
173impl FrameBlendShader {
174    pub fn create(input: Arc<wgpu::Texture>, device: &wgpu::Device, shaders: &Shaders) -> Self {
175        let previous_frame_texture = device.create_texture(&wgpu::TextureDescriptor {
176            label: "blend_previous_frame_texture".into(),
177            size: input.size(),
178            mip_level_count: 1,
179            sample_count: 1,
180            dimension: wgpu::TextureDimension::D2,
181            format: wgpu::TextureFormat::Rgba8Unorm,
182            usage: wgpu::TextureUsages::TEXTURE_BINDING | wgpu::TextureUsages::COPY_DST,
183            view_formats: &[wgpu::TextureFormat::Rgba8UnormSrgb],
184        });
185
186        let output_texture = device.create_texture(&wgpu::TextureDescriptor {
187            label: "blend_output_texture".into(),
188            size: input.size(),
189            mip_level_count: 1,
190            sample_count: 1,
191            dimension: wgpu::TextureDimension::D2,
192            format: wgpu::TextureFormat::Rgba8Unorm,
193            usage: *REQUIRED_TEXTURE_USAGES
194                | wgpu::TextureUsages::RENDER_ATTACHMENT
195                | wgpu::TextureUsages::COPY_DST,
196            view_formats: &[wgpu::TextureFormat::Rgba8UnormSrgb],
197        });
198
199        let pipeline = device.create_render_pipeline(&wgpu::RenderPipelineDescriptor {
200            label: "blend_pipeline".into(),
201            layout: None,
202            vertex: wgpu::VertexState {
203                module: &shaders.identity,
204                entry_point: None,
205                compilation_options: wgpu::PipelineCompilationOptions::default(),
206                buffers: &[],
207            },
208            primitive: wgpu::PrimitiveState {
209                topology: wgpu::PrimitiveTopology::TriangleStrip,
210                strip_index_format: None,
211                front_face: wgpu::FrontFace::Ccw,
212                cull_mode: None,
213                unclipped_depth: false,
214                polygon_mode: wgpu::PolygonMode::Fill,
215                conservative: false,
216            },
217            depth_stencil: None,
218            multisample: wgpu::MultisampleState::default(),
219            fragment: Some(wgpu::FragmentState {
220                module: &shaders.frame_blend,
221                entry_point: None,
222                compilation_options: wgpu::PipelineCompilationOptions::default(),
223                targets: &[Some(wgpu::ColorTargetState {
224                    format: wgpu::TextureFormat::Rgba8UnormSrgb,
225                    blend: Some(wgpu::BlendState::REPLACE),
226                    write_mask: wgpu::ColorWrites::ALL,
227                })],
228            }),
229            multiview_mask: None,
230            cache: None,
231        });
232
233        let bind_group = device.create_bind_group(&wgpu::BindGroupDescriptor {
234            label: "blend_bind_group".into(),
235            layout: &pipeline.get_bind_group_layout(0),
236            entries: &[
237                wgpu::BindGroupEntry {
238                    binding: 0,
239                    resource: wgpu::BindingResource::TextureView(
240                        &input.create_view(&SRGB_TEX_VIEW_DESCRIPTOR),
241                    ),
242                },
243                wgpu::BindGroupEntry {
244                    binding: 1,
245                    resource: wgpu::BindingResource::TextureView(
246                        &previous_frame_texture.create_view(&SRGB_TEX_VIEW_DESCRIPTOR),
247                    ),
248                },
249            ],
250        });
251
252        Self {
253            previous_frame: Arc::new(previous_frame_texture),
254            input,
255            output: Arc::new(output_texture),
256            bind_group,
257            pipeline,
258            skip_next_frame: true,
259        }
260    }
261}
262
263impl PipelineShader for FrameBlendShader {
264    fn draw(&mut self, encoder: &mut wgpu::CommandEncoder) {
265        if !self.skip_next_frame {
266            let mut render_pass = basic_render_pass(
267                encoder,
268                &self.output,
269                wgpu::TextureFormat::Rgba8UnormSrgb,
270                "blend_render_pass",
271            );
272
273            render_pass.set_bind_group(0, &self.bind_group, &[]);
274            render_pass.set_pipeline(&self.pipeline);
275
276            render_pass.draw(0..IDENTITY_VERTICES, 0..1);
277        } else {
278            encoder.copy_texture_to_texture(
279                self.input.as_image_copy(),
280                self.output.as_image_copy(),
281                self.input.size(),
282            );
283        }
284        self.skip_next_frame = false;
285
286        encoder.copy_texture_to_texture(
287            self.input.as_image_copy(),
288            self.previous_frame.as_image_copy(),
289            self.input.size(),
290        );
291    }
292
293    fn output_texture(&self) -> &Arc<wgpu::Texture> {
294        &self.output
295    }
296
297    fn reset_interframe_state(&mut self) {
298        self.skip_next_frame = true;
299    }
300}
301
302pub struct BlurShader {
303    output: Arc<wgpu::Texture>,
304    bind_groups: Vec<wgpu::BindGroup>,
305    pipeline: wgpu::RenderPipeline,
306}
307
308impl BlurShader {
309    pub fn create_horizontal_blur(
310        preprocess_shader: PreprocessShader,
311        device: &wgpu::Device,
312        input_texture: &wgpu::Texture,
313        shaders: &Shaders,
314    ) -> Option<Self> {
315        let fs_main = match preprocess_shader {
316            PreprocessShader::HorizontalBlurTwoPixels => "hblur_2px",
317            PreprocessShader::HorizontalBlurThreePixels => "hblur_3px",
318            PreprocessShader::HorizontalBlurSnesAdaptive => "hblur_snes",
319            _ => return None,
320        };
321
322        let width_scale_factor = match preprocess_shader {
323            PreprocessShader::HorizontalBlurSnesAdaptive if input_texture.width() >= 512 => 1,
324            PreprocessShader::HorizontalBlurSnesAdaptive => 2,
325            _ => 1,
326        };
327
328        Some(Self::create(device, input_texture, shaders, fs_main, width_scale_factor))
329    }
330
331    pub fn create_anti_dither(
332        anti_dither_shader: AntiDitherShader,
333        device: &wgpu::Device,
334        input_texture: &wgpu::Texture,
335        shaders: &Shaders,
336    ) -> Option<Self> {
337        let fs_main = match anti_dither_shader {
338            AntiDitherShader::Weak => "anti_dither_weak",
339            AntiDitherShader::Strong => "anti_dither_strong",
340            AntiDitherShader::None => return None,
341        };
342
343        Some(Self::create(device, input_texture, shaders, fs_main, 1))
344    }
345
346    pub fn create(
347        device: &wgpu::Device,
348        input_texture: &wgpu::Texture,
349        shaders: &Shaders,
350        fragment_entry_point: &str,
351        width_scale_factor: u32,
352    ) -> Self {
353        let input_texture_view = input_texture.create_view(&wgpu::TextureViewDescriptor {
354            format: Some(wgpu::TextureFormat::Rgba8Unorm),
355            ..wgpu::TextureViewDescriptor::default()
356        });
357
358        let output_texture = device.create_texture(&wgpu::TextureDescriptor {
359            label: "preprocess_output_texture".into(),
360            size: wgpu::Extent3d {
361                width: input_texture.width() * width_scale_factor,
362                height: input_texture.height(),
363                depth_or_array_layers: 1,
364            },
365            mip_level_count: 1,
366            sample_count: 1,
367            dimension: wgpu::TextureDimension::D2,
368            format: wgpu::TextureFormat::Rgba8Unorm,
369            usage: *REQUIRED_TEXTURE_USAGES | wgpu::TextureUsages::RENDER_ATTACHMENT,
370            view_formats: &[wgpu::TextureFormat::Rgba8UnormSrgb],
371        });
372
373        let pipeline = device.create_render_pipeline(&wgpu::RenderPipelineDescriptor {
374            label: "hblur_pipeline".into(),
375            layout: None,
376            vertex: wgpu::VertexState {
377                module: &shaders.identity,
378                entry_point: None,
379                compilation_options: wgpu::PipelineCompilationOptions::default(),
380                buffers: &[],
381            },
382            primitive: wgpu::PrimitiveState {
383                topology: wgpu::PrimitiveTopology::TriangleStrip,
384                strip_index_format: None,
385                front_face: wgpu::FrontFace::Ccw,
386                cull_mode: None,
387                unclipped_depth: false,
388                polygon_mode: wgpu::PolygonMode::Fill,
389                conservative: false,
390            },
391            depth_stencil: None,
392            multisample: wgpu::MultisampleState {
393                count: 1,
394                mask: !0,
395                alpha_to_coverage_enabled: false,
396            },
397            fragment: Some(wgpu::FragmentState {
398                module: &shaders.hblur,
399                entry_point: Some(fragment_entry_point),
400                compilation_options: wgpu::PipelineCompilationOptions::default(),
401                targets: &[Some(wgpu::ColorTargetState {
402                    format: wgpu::TextureFormat::Rgba8Unorm,
403                    blend: Some(wgpu::BlendState::REPLACE),
404                    write_mask: wgpu::ColorWrites::ALL,
405                })],
406            }),
407            multiview_mask: None,
408            cache: None,
409        });
410
411        let texture_width_buffer = device.create_buffer_init(&wgpu::util::BufferInitDescriptor {
412            label: "hblur_texture_width_buffer".into(),
413            contents: bytemuck::cast_slice(&[input_texture.size().width]),
414            usage: wgpu::BufferUsages::UNIFORM,
415        });
416        let bind_group = device.create_bind_group(&wgpu::BindGroupDescriptor {
417            label: "hblur_bind_group".into(),
418            layout: &pipeline.get_bind_group_layout(0),
419            entries: &[
420                wgpu::BindGroupEntry {
421                    binding: 0,
422                    resource: wgpu::BindingResource::TextureView(&input_texture_view),
423                },
424                wgpu::BindGroupEntry {
425                    binding: 1,
426                    resource: wgpu::BindingResource::Buffer(
427                        texture_width_buffer.as_entire_buffer_binding(),
428                    ),
429                },
430            ],
431        });
432
433        Self { output: Arc::new(output_texture), bind_groups: vec![bind_group], pipeline }
434    }
435}
436
437impl PipelineShader for BlurShader {
438    fn draw(&mut self, encoder: &mut wgpu::CommandEncoder) {
439        let mut render_pass = basic_render_pass(
440            encoder,
441            &self.output,
442            wgpu::TextureFormat::Rgba8Unorm,
443            "preprocess_render_pass",
444        );
445
446        for (i, bind_group) in self.bind_groups.iter().enumerate() {
447            render_pass.set_bind_group(i as u32, bind_group, &[]);
448        }
449        render_pass.set_pipeline(&self.pipeline);
450
451        render_pass.draw(0..IDENTITY_VERTICES, 0..1);
452    }
453
454    fn output_texture(&self) -> &Arc<wgpu::Texture> {
455        &self.output
456    }
457}
458
459pub struct PrescaleShader {
460    bind_group: wgpu::BindGroup,
461    pipeline: wgpu::RenderPipeline,
462    output: Arc<wgpu::Texture>,
463}
464
465impl PrescaleShader {
466    #[allow(clippy::too_many_arguments)]
467    pub fn create(
468        renderer_config: RendererConfig,
469        frame_size: FrameSize,
470        display_area: DisplayArea,
471        pixel_aspect_ratio: Option<FiniteF64>,
472        input: &wgpu::Texture,
473        device: &wgpu::Device,
474        limits: &wgpu::Limits,
475        shaders: &Shaders,
476    ) -> Option<Self> {
477        let (prescale_width, prescale_height) = determine_prescale_factors(
478            renderer_config.prescale_mode,
479            frame_size,
480            pixel_aspect_ratio,
481            display_area,
482            renderer_config.frame_rotation,
483            input.size(),
484            limits,
485        );
486
487        if prescale_width <= 1 && prescale_height <= 1 && !renderer_config.scanlines_enabled {
488            return None;
489        }
490
491        log::debug!(
492            "Creating prescale shader with width factor {prescale_width}x and height factor {prescale_height}x",
493        );
494
495        let scaled_texture = device.create_texture(&wgpu::TextureDescriptor {
496            label: "scaled_texture".into(),
497            size: wgpu::Extent3d {
498                width: prescale_width * input.width(),
499                height: prescale_height * input.height(),
500                depth_or_array_layers: 1,
501            },
502            mip_level_count: 1,
503            sample_count: 1,
504            dimension: wgpu::TextureDimension::D2,
505            format: wgpu::TextureFormat::Rgba8Unorm,
506            usage: *REQUIRED_TEXTURE_USAGES | wgpu::TextureUsages::RENDER_ATTACHMENT,
507            view_formats: &[wgpu::TextureFormat::Rgba8UnormSrgb],
508        });
509
510        let fs_main =
511            if renderer_config.scanlines_enabled { "scanlines" } else { "basic_prescale" };
512        let scanline_multiplier = renderer_config.scanlines_brightness.clamp(0.0, 1.0);
513
514        let pipeline = device.create_render_pipeline(&wgpu::RenderPipelineDescriptor {
515            label: "prescale_pipeline".into(),
516            layout: None,
517            vertex: wgpu::VertexState {
518                module: &shaders.identity,
519                entry_point: None,
520                compilation_options: wgpu::PipelineCompilationOptions::default(),
521                buffers: &[],
522            },
523            primitive: wgpu::PrimitiveState {
524                topology: wgpu::PrimitiveTopology::TriangleStrip,
525                strip_index_format: None,
526                front_face: wgpu::FrontFace::Ccw,
527                cull_mode: None,
528                unclipped_depth: false,
529                polygon_mode: wgpu::PolygonMode::Fill,
530                conservative: false,
531            },
532            depth_stencil: None,
533            multisample: wgpu::MultisampleState {
534                count: 1,
535                mask: !0,
536                alpha_to_coverage_enabled: false,
537            },
538            fragment: Some(wgpu::FragmentState {
539                module: &shaders.prescale,
540                entry_point: Some(fs_main),
541                compilation_options: wgpu::PipelineCompilationOptions::default(),
542                targets: &[Some(wgpu::ColorTargetState {
543                    format: wgpu::TextureFormat::Rgba8UnormSrgb,
544                    blend: Some(wgpu::BlendState::REPLACE),
545                    write_mask: wgpu::ColorWrites::ALL,
546                })],
547            }),
548            multiview_mask: None,
549            cache: None,
550        });
551
552        let prescale_params = [
553            prescale_width,
554            prescale_height,
555            frame_size.height,
556            scaled_texture.height(),
557            (scanline_multiplier as f32).to_bits(),
558        ];
559
560        let prescale_factor_buffer = device.create_buffer_init(&wgpu::util::BufferInitDescriptor {
561            label: "prescale_factor_buffer".into(),
562            contents: bytemuck::cast_slice(&prescale_params),
563            usage: wgpu::BufferUsages::COPY_DST | wgpu::BufferUsages::UNIFORM,
564        });
565
566        let input_view = input.create_view(&SRGB_TEX_VIEW_DESCRIPTOR);
567        let bind_group = device.create_bind_group(&wgpu::BindGroupDescriptor {
568            label: "prescale_bind_group".into(),
569            layout: &pipeline.get_bind_group_layout(0),
570            entries: &[
571                wgpu::BindGroupEntry {
572                    binding: 0,
573                    resource: wgpu::BindingResource::TextureView(&input_view),
574                },
575                wgpu::BindGroupEntry {
576                    binding: 1,
577                    resource: wgpu::BindingResource::Buffer(
578                        prescale_factor_buffer.as_entire_buffer_binding(),
579                    ),
580                },
581            ],
582        });
583
584        Some(Self { bind_group, pipeline, output: Arc::new(scaled_texture) })
585    }
586}
587
588fn determine_prescale_factors(
589    mode: PrescaleMode,
590    frame_size: FrameSize,
591    pixel_aspect_ratio: Option<FiniteF64>,
592    display_area: DisplayArea,
593    rotation: FrameRotation,
594    input_size: wgpu::Extent3d,
595    limits: &wgpu::Limits,
596) -> (u32, u32) {
597    let (target_width, target_height) = match mode {
598        PrescaleMode::Auto => {
599            // For 90/270 degree rotations, display area is based on rotated frame size and aspect
600            // ratio, so rotate the display size back when computing auto-prescale factors
601            let (display_width, display_height) = rotation.rotate_display_area_size(display_area);
602
603            let width = match pixel_aspect_ratio {
604                Some(par) => {
605                    let frame_aspect_ratio =
606                        f64::from(frame_size.width) / f64::from(frame_size.height);
607                    let screen_aspect_ratio = f64::from(par) * frame_aspect_ratio;
608                    f64::from(display_height) * screen_aspect_ratio
609                }
610                None => f64::from(display_width),
611            };
612            let height = f64::from(display_height);
613            (width, height)
614        }
615        PrescaleMode::Manual { width, height } => {
616            let width = f64::from(width.get() * frame_size.width);
617            let height = f64::from(height.get() * frame_size.height);
618            (width, height)
619        }
620    };
621
622    let width_ratio = (target_width / f64::from(input_size.width)) as u32;
623    let height_ratio = (target_height / f64::from(input_size.height)) as u32;
624    let prescale_width = clamp_prescale_factor(width_ratio, input_size.width, limits);
625    let prescale_height = clamp_prescale_factor(height_ratio, input_size.height, limits);
626
627    (prescale_width, prescale_height)
628}
629
630fn clamp_prescale_factor(prescale_factor: u32, input_dimension: u32, limits: &wgpu::Limits) -> u32 {
631    let max_dimension = limits.max_texture_dimension_2d;
632    let max_prescale_factor = max_dimension / input_dimension;
633
634    if max_prescale_factor < prescale_factor {
635        log::warn!(
636            "Prescale factor {prescale_factor} is too high for frame dimension {input_dimension}; reducing to {max_prescale_factor}",
637        );
638    }
639
640    prescale_factor.clamp(1, max_prescale_factor)
641}
642
643impl PipelineShader for PrescaleShader {
644    fn draw(&mut self, encoder: &mut wgpu::CommandEncoder) {
645        let mut prescale_pass = basic_render_pass(
646            encoder,
647            &self.output,
648            wgpu::TextureFormat::Rgba8UnormSrgb,
649            "prescale_render_pass",
650        );
651
652        prescale_pass.set_bind_group(0, &self.bind_group, &[]);
653        prescale_pass.set_pipeline(&self.pipeline);
654
655        prescale_pass.draw(0..IDENTITY_VERTICES, 0..1);
656    }
657
658    fn output_texture(&self) -> &Arc<wgpu::Texture> {
659        &self.output
660    }
661}
662
663#[derive(Debug, Error)]
664pub enum UpscaleShaderError {
665    #[error(
666        "Scaled texture size of {scaled_width}x{scaled_height} exceeds GPU device's maximum of {max_dimension}x{max_dimension}"
667    )]
668    TextureTooLarge { scaled_width: u32, scaled_height: u32, max_dimension: u32 },
669}
670
671pub struct UpscaleShader {
672    output: Arc<wgpu::Texture>,
673    bind_group: wgpu::BindGroup,
674    pipeline: wgpu::ComputePipeline,
675    x_workgroups: u32,
676    y_workgroups: u32,
677}
678
679impl UpscaleShader {
680    pub fn create_xbrz(
681        device: &wgpu::Device,
682        shaders: &Shaders,
683        input: &wgpu::Texture,
684        scale_factor: u32,
685    ) -> Option<Self> {
686        let shader = (&shaders.xbrz, None);
687        let shader_constants = [("scale_factor", scale_factor.into())];
688
689        match Self::create(device, shader, &shader_constants, input, scale_factor) {
690            Ok(shader) => Some(shader),
691            Err(err) => {
692                log::error!("Error creating xBRZ {scale_factor}x shader: {err}");
693                if scale_factor > 2 {
694                    log::info!("Attempting to create an xBRZ {}x shader instead", scale_factor - 1);
695                    Self::create_xbrz(device, shaders, input, scale_factor - 1)
696                } else {
697                    None
698                }
699            }
700        }
701    }
702
703    pub fn create_mmpx(
704        device: &wgpu::Device,
705        shaders: &Shaders,
706        input: &wgpu::Texture,
707    ) -> Option<Self> {
708        match Self::create(device, (&shaders.mmpx, None), &[], input, 2) {
709            Ok(shader) => Some(shader),
710            Err(err) => {
711                log::error!("Error creating MMPX shader: {err}");
712                None
713            }
714        }
715    }
716
717    pub fn create_mmpx_enhanced(
718        device: &wgpu::Device,
719        shaders: &Shaders,
720        input: &wgpu::Texture,
721    ) -> Option<Self> {
722        match Self::create(device, (&shaders.mmpx_enhanced, None), &[], input, 2) {
723            Ok(shader) => Some(shader),
724            Err(err) => {
725                log::error!("Error creating MMPX Enhanced shader: {err}");
726                None
727            }
728        }
729    }
730
731    pub fn create(
732        device: &wgpu::Device,
733        (shader_module, shader_entry_point): (&wgpu::ShaderModule, Option<&str>),
734        shader_constants: &[(&str, f64)],
735        input: &wgpu::Texture,
736        scale_factor: u32,
737    ) -> Result<Self, UpscaleShaderError> {
738        let scaled_width = scale_factor * input.width();
739        let scaled_height = scale_factor * input.height();
740        let max_dimension = device.limits().max_texture_dimension_2d;
741
742        if scaled_width > max_dimension || scaled_height > max_dimension {
743            return Err(UpscaleShaderError::TextureTooLarge {
744                scaled_width,
745                scaled_height,
746                max_dimension,
747            });
748        }
749
750        let output = device.create_texture(&wgpu::TextureDescriptor {
751            label: "xbrz_texture".into(),
752            size: wgpu::Extent3d {
753                width: scaled_width,
754                height: scaled_height,
755                depth_or_array_layers: 1,
756            },
757            mip_level_count: 1,
758            sample_count: 1,
759            dimension: wgpu::TextureDimension::D2,
760            format: wgpu::TextureFormat::Rgba8Unorm,
761            usage: *REQUIRED_TEXTURE_USAGES,
762            view_formats: &[wgpu::TextureFormat::Rgba8UnormSrgb],
763        });
764
765        let pipeline = device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
766            label: "xbrz_pipeline".into(),
767            layout: None,
768            module: shader_module,
769            entry_point: shader_entry_point,
770            compilation_options: wgpu::PipelineCompilationOptions {
771                constants: shader_constants,
772                ..wgpu::PipelineCompilationOptions::default()
773            },
774            cache: None,
775        });
776
777        let input_view = input.create_view(&wgpu::TextureViewDescriptor::default());
778        let output_view = output.create_view(&wgpu::TextureViewDescriptor::default());
779
780        let bind_group = device.create_bind_group(&wgpu::BindGroupDescriptor {
781            label: "xbrz_bind_group".into(),
782            layout: &pipeline.get_bind_group_layout(0),
783            entries: &[
784                wgpu::BindGroupEntry {
785                    binding: 0,
786                    resource: wgpu::BindingResource::TextureView(&input_view),
787                },
788                wgpu::BindGroupEntry {
789                    binding: 1,
790                    resource: wgpu::BindingResource::TextureView(&output_view),
791                },
792            ],
793        });
794
795        let x_workgroups = input.width().div_ceil(16);
796        let y_workgroups = input.height().div_ceil(16);
797
798        Ok(Self { output: Arc::new(output), bind_group, pipeline, x_workgroups, y_workgroups })
799    }
800}
801
802impl PipelineShader for UpscaleShader {
803    fn draw(&mut self, encoder: &mut wgpu::CommandEncoder) {
804        let mut compute_pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
805
806        compute_pass.set_bind_group(0, &self.bind_group, &[]);
807        compute_pass.set_pipeline(&self.pipeline);
808
809        compute_pass.dispatch_workgroups(self.x_workgroups, self.y_workgroups, 1);
810    }
811
812    fn output_texture(&self) -> &Arc<wgpu::Texture> {
813        &self.output
814    }
815}
816
817#[cfg(test)]
818mod tests {
819    use super::*;
820    use crate::config::PrescaleFactor;
821
822    fn display_area(width: u32, height: u32) -> DisplayArea {
823        DisplayArea { width, height, x: 0, y: 0, pixel_density: 1.0 }
824    }
825
826    fn basic_auto_prescale_test(
827        width: u32,
828        height: u32,
829        width_scale: u32,
830        height_scale: u32,
831    ) -> (u32, u32) {
832        determine_prescale_factors(
833            PrescaleMode::Auto,
834            FrameSize { width, height },
835            None,
836            display_area(width * width_scale, height * height_scale),
837            FrameRotation::None,
838            wgpu::Extent3d { width, height, depth_or_array_layers: 1 },
839            &wgpu::Limits::default(),
840        )
841    }
842
843    #[test]
844    fn auto_prescale_square() {
845        let (width, height) = basic_auto_prescale_test(320, 240, 4, 4);
846
847        assert_eq!(width, 4);
848        assert_eq!(height, 4);
849    }
850
851    #[test]
852    fn auto_prescale_horizontal_rect() {
853        let (width, height) = basic_auto_prescale_test(320, 240, 4, 2);
854
855        assert_eq!(width, 4);
856        assert_eq!(height, 2);
857    }
858
859    #[test]
860    fn auto_prescale_vertical_rect() {
861        let (width, height) = basic_auto_prescale_test(320, 240, 2, 4);
862
863        assert_eq!(width, 2);
864        assert_eq!(height, 4);
865    }
866
867    #[test]
868    fn auto_prescale_squish_vertical() {
869        let (width, height) = determine_prescale_factors(
870            PrescaleMode::Auto,
871            FrameSize { width: 320, height: 480 },
872            Some(FiniteF64::try_from(2.0).unwrap()),
873            display_area(320 * 4, 240 * 4),
874            FrameRotation::None,
875            wgpu::Extent3d { width: 320, height: 480, depth_or_array_layers: 1 },
876            &wgpu::Limits::default(),
877        );
878
879        assert_eq!(width, 4);
880        assert_eq!(height, 2);
881    }
882
883    #[test]
884    fn auto_prescale_squish_horizontal() {
885        let (width, height) = determine_prescale_factors(
886            PrescaleMode::Auto,
887            FrameSize { width: 512, height: 240 },
888            Some(FiniteF64::try_from(0.5).unwrap()),
889            display_area(256 * 4, 240 * 4),
890            FrameRotation::None,
891            wgpu::Extent3d { width: 512, height: 240, depth_or_array_layers: 1 },
892            &wgpu::Limits::default(),
893        );
894
895        assert_eq!(width, 2);
896        assert_eq!(height, 4);
897    }
898
899    #[test]
900    fn auto_prescale_scaled_input() {
901        let (width, height) = determine_prescale_factors(
902            PrescaleMode::Auto,
903            FrameSize { width: 320, height: 240 },
904            None,
905            display_area(320 * 4, 240 * 4),
906            FrameRotation::None,
907            wgpu::Extent3d { width: 320 * 2, height: 240, depth_or_array_layers: 1 },
908            &wgpu::Limits::default(),
909        );
910
911        assert_eq!(width, 2);
912        assert_eq!(height, 4);
913    }
914
915    #[test]
916    fn auto_prescale_round_down() {
917        let (width, height) = determine_prescale_factors(
918            PrescaleMode::Auto,
919            FrameSize { width: 320, height: 240 },
920            None,
921            display_area(320 * 11 / 4, 240 * 7 / 4),
922            FrameRotation::None,
923            wgpu::Extent3d { width: 320, height: 240, depth_or_array_layers: 1 },
924            &wgpu::Limits::default(),
925        );
926
927        assert_eq!(width, 2);
928        assert_eq!(height, 1);
929    }
930
931    #[test]
932    fn auto_prescale_pixel_aspect_ratio() {
933        let (width, height) = determine_prescale_factors(
934            PrescaleMode::Auto,
935            FrameSize { width: 320, height: 240 },
936            Some(FiniteF64::try_from(0.9).unwrap()),
937            display_area(320 * 2 * 9 / 10, 240 * 2),
938            FrameRotation::None,
939            wgpu::Extent3d { width: 320, height: 240, depth_or_array_layers: 1 },
940            &wgpu::Limits::default(),
941        );
942
943        // Sub-1 pixel aspect ratio should drop prescale factor
944        assert_eq!(width, 1);
945        assert_eq!(height, 2);
946    }
947
948    #[test]
949    fn manual_prescale_basic() {
950        let factor = PrescaleFactor::try_from(5).unwrap();
951        let (width, height) = determine_prescale_factors(
952            PrescaleMode::Manual { width: factor, height: factor },
953            FrameSize { width: 320, height: 240 },
954            None,
955            display_area(320 * 5, 240 * 5),
956            FrameRotation::None,
957            wgpu::Extent3d { width: 320, height: 240, depth_or_array_layers: 1 },
958            &wgpu::Limits::default(),
959        );
960
961        assert_eq!(width, 5);
962        assert_eq!(height, 5);
963    }
964
965    #[test]
966    fn manual_prescale_scaled_input() {
967        let factor = PrescaleFactor::try_from(5).unwrap();
968        let (width, height) = determine_prescale_factors(
969            PrescaleMode::Manual { width: factor, height: factor },
970            FrameSize { width: 320, height: 240 },
971            None,
972            display_area(320 * 5, 240 * 5),
973            FrameRotation::None,
974            wgpu::Extent3d { width: 320 * 2, height: 240, depth_or_array_layers: 1 },
975            &wgpu::Limits::default(),
976        );
977
978        assert_eq!(width, 2);
979        assert_eq!(height, 5);
980    }
981
982    #[test]
983    fn auto_prescale_rotated() {
984        let (width, height) = determine_prescale_factors(
985            PrescaleMode::Auto,
986            FrameSize { width: 200, height: 400 },
987            None,
988            display_area(1000, 600),
989            FrameRotation::Clockwise,
990            wgpu::Extent3d { width: 200, height: 400, depth_or_array_layers: 1 },
991            &wgpu::Limits::default(),
992        );
993
994        assert_eq!(width, 3);
995        assert_eq!(height, 2);
996    }
997}