Skip to main content

crush_gpu/backend/
wgpu_backend.rs

1//! wgpu compute shader backend (Vulkan/Metal/DX12)
2//!
3//! Provides GPU decompression via wgpu's cross-platform compute shader API.
4//! Requires Vulkan 1.2 / Metal 2 / DX12 + 2 GB VRAM minimum.
5
6use std::sync::atomic::{AtomicBool, Ordering};
7use std::time::Duration;
8
9use crush_core::error::{CrushError, PluginError, Result};
10
11use super::{CompressedTile, ComputeBackend, GpuInfo, GpuVendor, MIN_VRAM_BYTES};
12
13/// WGSL compute shader source — LZ77 (v1 compat, embedded at compile time).
14const DECOMPRESS_SHADER: &str = include_str!("../shader/decompress.wgsl");
15
16/// WGSL compute shader source — `GDeflate` (v2, embedded at compile time).
17const GDEFLATE_SHADER: &str = include_str!("../shader/gdeflate_decompress.wgsl");
18
19/// Maximum time to wait for GPU work to complete before treating it as a hang.
20/// 5 seconds is generous for a single tile (64KB) decompression dispatch.
21/// On Windows, TDR typically resets the GPU after ~2s, so this catches hangs
22/// that survive TDR as well.
23const GPU_POLL_TIMEOUT: Duration = Duration::from_secs(5);
24
25/// wgpu-backed GPU compute backend.
26///
27/// Holds two compute pipelines: one for LZ77 (v1) and one for `GDeflate` (v2).
28pub struct WgpuBackend {
29    info: GpuInfo,
30    device: wgpu::Device,
31    queue: wgpu::Queue,
32    // LZ77 pipeline (v1)
33    pipeline: wgpu::ComputePipeline,
34    bind_group_layout: wgpu::BindGroupLayout,
35    // GDeflate pipeline (v2)
36    gdeflate_pipeline: wgpu::ComputePipeline,
37    gdeflate_bgl: wgpu::BindGroupLayout,
38}
39
40/// Uniform struct matching the `TileMeta` layout in the LZ77 WGSL shader.
41/// 8 × u32 = 32 bytes.
42#[repr(C)]
43#[derive(Copy, Clone, bytemuck::Pod, bytemuck::Zeroable)]
44struct TileMeta {
45    compressed_offset: u32,
46    compressed_size: u32,
47    uncompressed_size: u32,
48    sub_stream_count: u32,
49    output_offset: u32,
50    tile_index: u32,
51    _pad0: u32,
52    _pad1: u32,
53}
54
55/// Metadata struct matching `GDeflateMeta` in the `GDeflate` WGSL shader.
56/// 4 × u32 = 16 bytes.
57#[repr(C)]
58#[derive(Copy, Clone, bytemuck::Pod, bytemuck::Zeroable)]
59struct GDeflateMeta {
60    payload_size: u32,
61    uncompressed_size: u32,
62    _pad0: u32,
63    _pad1: u32,
64}
65
66/// GPU buffers needed for a single tile dispatch.
67struct TileBuffers {
68    meta: wgpu::Buffer,
69    compressed: wgpu::Buffer,
70    output: wgpu::Buffer,
71    lengths: wgpu::Buffer,
72    out_staging: wgpu::Buffer,
73    len_staging: wgpu::Buffer,
74}
75
76/// Allocate all GPU buffers for a single tile dispatch.
77///
78/// Uses `catch_unwind` to prevent wgpu internal panics (e.g. on OOM)
79/// from crashing the entire process.
80fn create_tile_buffers(
81    device: &wgpu::Device,
82    comp_data: &[u8],
83    out_aligned: u64,
84    len_size: u64,
85) -> Result<TileBuffers> {
86    std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
87        let buf = |label, size, usage| {
88            device.create_buffer(&wgpu::BufferDescriptor {
89                label: Some(label),
90                size,
91                usage,
92                mapped_at_creation: false,
93            })
94        };
95        let storage_dst = wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_DST;
96        let storage_src = wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_SRC;
97        let map_read = wgpu::BufferUsages::COPY_DST | wgpu::BufferUsages::MAP_READ;
98
99        TileBuffers {
100            meta: buf(
101                "tile_meta",
102                std::mem::size_of::<TileMeta>() as u64,
103                storage_dst,
104            ),
105            compressed: buf("compressed_data", comp_data.len() as u64, storage_dst),
106            output: buf("decompressed_data", out_aligned, storage_src),
107            lengths: buf("sub_stream_lengths", len_size, storage_src),
108            out_staging: buf("out_staging", out_aligned, map_read),
109            len_staging: buf("len_staging", len_size, map_read),
110        }
111    }))
112    .map_err(|e| {
113        let msg = e
114            .downcast_ref::<String>()
115            .map(String::as_str)
116            .or_else(|| e.downcast_ref::<&str>().copied())
117            .unwrap_or("unknown GPU buffer allocation panic");
118        CrushError::from(PluginError::OperationFailed(format!(
119            "GPU buffer allocation failed: {msg}"
120        )))
121    })
122}
123
124/// Parse sub-stream length u32 values from raw bytes.
125fn parse_ss_lengths(len_bytes: &[u8], n: u32) -> Vec<u32> {
126    let mut ss_lengths = Vec::with_capacity(n as usize);
127    for i in 0..n as usize {
128        let off = i * 4;
129        if off + 4 <= len_bytes.len() {
130            ss_lengths.push(u32::from_le_bytes([
131                len_bytes[off],
132                len_bytes[off + 1],
133                len_bytes[off + 2],
134                len_bytes[off + 3],
135            ]));
136        } else {
137            ss_lengths.push(0);
138        }
139    }
140    ss_lengths
141}
142
143/// Create the bind group layout with 4 storage buffer bindings for the LZ77 shader.
144fn create_bind_group_layout(device: &wgpu::Device) -> wgpu::BindGroupLayout {
145    let storage_entry = |binding: u32, read_only: bool| wgpu::BindGroupLayoutEntry {
146        binding,
147        visibility: wgpu::ShaderStages::COMPUTE,
148        ty: wgpu::BindingType::Buffer {
149            ty: wgpu::BufferBindingType::Storage { read_only },
150            has_dynamic_offset: false,
151            min_binding_size: None,
152        },
153        count: None,
154    };
155
156    device.create_bind_group_layout(&wgpu::BindGroupLayoutDescriptor {
157        label: Some("decompress_bgl"),
158        entries: &[
159            storage_entry(0, true),  // tile_meta
160            storage_entry(1, true),  // compressed_data
161            storage_entry(2, false), // decompressed_data
162            storage_entry(3, false), // sub_stream_lengths
163        ],
164    })
165}
166
167/// Create the bind group layout with 3 storage buffer bindings for the `GDeflate` shader.
168fn create_gdeflate_bgl(device: &wgpu::Device) -> wgpu::BindGroupLayout {
169    let storage_entry = |binding: u32, read_only: bool| wgpu::BindGroupLayoutEntry {
170        binding,
171        visibility: wgpu::ShaderStages::COMPUTE,
172        ty: wgpu::BindingType::Buffer {
173            ty: wgpu::BufferBindingType::Storage { read_only },
174            has_dynamic_offset: false,
175            min_binding_size: None,
176        },
177        count: None,
178    };
179
180    device.create_bind_group_layout(&wgpu::BindGroupLayoutDescriptor {
181        label: Some("gdeflate_bgl"),
182        entries: &[
183            storage_entry(0, true),  // meta (GDeflateMeta)
184            storage_entry(1, true),  // compressed (GDeflate payload)
185            storage_entry(2, false), // output (decompressed bytes)
186        ],
187    })
188}
189
190impl WgpuBackend {
191    /// Attempt to create a wgpu backend by discovering a suitable GPU adapter.
192    ///
193    /// Returns `None` if no compatible GPU is found.
194    ///
195    /// # Errors
196    ///
197    /// Returns an error if wgpu initialisation fails unexpectedly.
198    pub fn try_new() -> Result<Option<Self>> {
199        // wgpu 30 takes the descriptor by value, and `InstanceDescriptor` no
200        // longer implements `Default` (it gained a non-`Default` boxed display
201        // handle). We never present to a surface, so the handle-less
202        // constructor is the right base.
203        let instance = wgpu::Instance::new(wgpu::InstanceDescriptor {
204            backends: wgpu::Backends::VULKAN | wgpu::Backends::METAL | wgpu::Backends::DX12,
205            ..wgpu::InstanceDescriptor::new_without_display_handle()
206        });
207
208        let adapter = pollster::block_on(instance.request_adapter(&wgpu::RequestAdapterOptions {
209            power_preference: wgpu::PowerPreference::HighPerformance,
210            compatible_surface: None,
211            force_fallback_adapter: false,
212            // Limit bucketing rounds reported limits down into coarse buckets
213            // to reduce fingerprinting when exposing wgpu to untrusted content.
214            // crush is trusted local code and uses `max_buffer_size` as its
215            // VRAM proxy below, so keep the true device limits.
216            apply_limit_buckets: false,
217        }));
218
219        let Ok(adapter) = adapter else {
220            return Ok(None);
221        };
222
223        let adapter_info = adapter.get_info();
224
225        // Reject software/CPU adapters — they won't provide GPU acceleration.
226        if adapter_info.device_type == wgpu::DeviceType::Cpu {
227            return Ok(None);
228        }
229
230        // Use max_buffer_size as a rough VRAM proxy. Note: this is the
231        // driver-reported maximum single-buffer size, not total VRAM.
232        // On discrete GPUs it's typically 2+ GB. We use a conservative
233        // check and rely on catch_unwind + CPU fallback for OOM safety.
234        let limits = adapter.limits();
235        let estimated_vram = limits.max_buffer_size;
236        if estimated_vram < MIN_VRAM_BYTES {
237            return Ok(None);
238        }
239
240        let vendor = match adapter_info.vendor {
241            0x10DE => GpuVendor::Nvidia,
242            0x1002 => GpuVendor::Amd,
243            0x8086 => GpuVendor::Intel,
244            _ if adapter_info.driver.to_lowercase().contains("apple")
245                || adapter_info.name.to_lowercase().contains("apple") =>
246            {
247                GpuVendor::Apple
248            }
249            _ => GpuVendor::Other,
250        };
251
252        let (device, queue) = pollster::block_on(adapter.request_device(&wgpu::DeviceDescriptor {
253            label: Some("crush-gpu"),
254            required_features: wgpu::Features::empty(),
255            required_limits: wgpu::Limits::default(),
256            memory_hints: wgpu::MemoryHints::Performance,
257            trace: wgpu::Trace::Off,
258            experimental_features: wgpu::ExperimentalFeatures::default(),
259        }))
260        .map_err(|e| PluginError::OperationFailed(format!("wgpu device request failed: {e}")))?;
261
262        // --- LZ77 pipeline (v1) ---
263        let shader_module = device.create_shader_module(wgpu::ShaderModuleDescriptor {
264            label: Some("decompress.wgsl"),
265            source: wgpu::ShaderSource::Wgsl(DECOMPRESS_SHADER.into()),
266        });
267
268        let bind_group_layout = create_bind_group_layout(&device);
269
270        let pipeline_layout = device.create_pipeline_layout(&wgpu::PipelineLayoutDescriptor {
271            label: Some("decompress_pipeline_layout"),
272            bind_group_layouts: &[Some(&bind_group_layout)],
273            immediate_size: 0,
274        });
275
276        let pipeline = device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
277            label: Some("decompress_pipeline"),
278            layout: Some(&pipeline_layout),
279            module: &shader_module,
280            entry_point: Some("main"),
281            compilation_options: wgpu::PipelineCompilationOptions::default(),
282            cache: None,
283        });
284
285        // --- GDeflate pipeline (v2) ---
286        let gdeflate_module = device.create_shader_module(wgpu::ShaderModuleDescriptor {
287            label: Some("gdeflate_decompress.wgsl"),
288            source: wgpu::ShaderSource::Wgsl(GDEFLATE_SHADER.into()),
289        });
290
291        let gdeflate_bgl = create_gdeflate_bgl(&device);
292
293        let gdeflate_pl = device.create_pipeline_layout(&wgpu::PipelineLayoutDescriptor {
294            label: Some("gdeflate_pipeline_layout"),
295            bind_group_layouts: &[Some(&gdeflate_bgl)],
296            immediate_size: 0,
297        });
298
299        let gdeflate_pipeline = device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
300            label: Some("gdeflate_pipeline"),
301            layout: Some(&gdeflate_pl),
302            module: &gdeflate_module,
303            entry_point: Some("main"),
304            compilation_options: wgpu::PipelineCompilationOptions::default(),
305            cache: None,
306        });
307
308        let info = GpuInfo {
309            name: adapter_info.name.clone(),
310            vendor,
311            vram_bytes: estimated_vram,
312            api_backend: format!("{:?}", adapter_info.backend),
313        };
314
315        Ok(Some(Self {
316            info,
317            device,
318            queue,
319            pipeline,
320            bind_group_layout,
321            gdeflate_pipeline,
322            gdeflate_bgl,
323        }))
324    }
325
326    /// Map two staging buffers, poll the device, and return their contents.
327    fn readback_buffers(
328        &self,
329        out_staging: &wgpu::Buffer,
330        lengths_staging: &wgpu::Buffer,
331    ) -> Result<(Vec<u8>, Vec<u8>)> {
332        let out_slice = out_staging.slice(..);
333        let lengths_slice = lengths_staging.slice(..);
334
335        let (out_tx, out_rx) = std::sync::mpsc::channel();
336        out_slice.map_async(wgpu::MapMode::Read, move |result| {
337            let _ = out_tx.send(result);
338        });
339        let (len_tx, len_rx) = std::sync::mpsc::channel();
340        lengths_slice.map_async(wgpu::MapMode::Read, move |result| {
341            let _ = len_tx.send(result);
342        });
343
344        self.device
345            .poll(wgpu::PollType::Wait {
346                submission_index: None,
347                timeout: Some(GPU_POLL_TIMEOUT),
348            })
349            .map_err(|e| {
350                PluginError::OperationFailed(format!(
351                    "GPU poll failed (timeout or device lost): {e}"
352                ))
353            })?;
354
355        out_rx
356            .recv()
357            .map_err(|e| PluginError::OperationFailed(format!("GPU readback channel error: {e}")))?
358            .map_err(|e| PluginError::OperationFailed(format!("GPU output map failed: {e}")))?;
359        len_rx
360            .recv()
361            .map_err(|e| PluginError::OperationFailed(format!("GPU readback channel error: {e}")))?
362            .map_err(|e| PluginError::OperationFailed(format!("GPU lengths map failed: {e}")))?;
363
364        // wgpu 30 returns `Result` here instead of panicking on a bad range.
365        let out_bytes = out_slice
366            .get_mapped_range()
367            .map_err(|e| PluginError::OperationFailed(format!("GPU output range map failed: {e}")))?
368            .to_vec();
369        let len_bytes = lengths_slice
370            .get_mapped_range()
371            .map_err(|e| {
372                PluginError::OperationFailed(format!("GPU lengths range map failed: {e}"))
373            })?
374            .to_vec();
375        Ok((out_bytes, len_bytes))
376    }
377
378    /// Decompress a single tile on the GPU and return the raw sub-stream outputs.
379    fn dispatch_tile(&self, tile: &CompressedTile, tile_index: u32) -> Result<(Vec<u8>, Vec<u32>)> {
380        let n = u32::from(tile.sub_stream_count);
381        if n == 0 {
382            return Err(CrushError::InvalidFormat(
383                "tile has zero sub-stream count".to_owned(),
384            ));
385        }
386
387        // Guard against crafted archives with absurd uncompressed_size that
388        // would cause u32 overflow in `n * max_per_ss`.
389        let max_tile_size: u32 = crate::format::DEFAULT_TILE_SIZE;
390        if tile.uncompressed_size > max_tile_size.saturating_mul(2) {
391            return Err(CrushError::InvalidFormat(format!(
392                "tile uncompressed_size {} exceeds maximum {}",
393                tile.uncompressed_size,
394                max_tile_size * 2,
395            )));
396        }
397
398        let max_per_ss = tile.uncompressed_size.div_ceil(n);
399        let output_buf_size = n.checked_mul(max_per_ss).ok_or_else(|| {
400            CrushError::InvalidFormat(format!("output buffer size overflow: {n} * {max_per_ss}"))
401        })?;
402
403        let mut comp_data = tile.data.clone();
404        while !comp_data.len().is_multiple_of(4) {
405            comp_data.push(0);
406        }
407
408        let meta = TileMeta {
409            compressed_offset: 0,
410            compressed_size: u32::try_from(tile.data.len())
411                .map_err(|e| PluginError::OperationFailed(e.to_string()))?,
412            uncompressed_size: tile.uncompressed_size,
413            sub_stream_count: n,
414            output_offset: 0,
415            tile_index,
416            _pad0: 0,
417            _pad1: 0,
418        };
419
420        let out_aligned = u64::from(output_buf_size.div_ceil(4) * 4).max(4);
421        let len_size = (u64::from(n) * 4).max(4);
422        let bufs = create_tile_buffers(&self.device, &comp_data, out_aligned, len_size)?;
423
424        self.queue
425            .write_buffer(&bufs.meta, 0, bytemuck::bytes_of(&meta));
426        self.queue.write_buffer(&bufs.compressed, 0, &comp_data);
427
428        let bind_group = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
429            label: Some("decompress_bg"),
430            layout: &self.bind_group_layout,
431            entries: &[
432                wgpu::BindGroupEntry {
433                    binding: 0,
434                    resource: bufs.meta.as_entire_binding(),
435                },
436                wgpu::BindGroupEntry {
437                    binding: 1,
438                    resource: bufs.compressed.as_entire_binding(),
439                },
440                wgpu::BindGroupEntry {
441                    binding: 2,
442                    resource: bufs.output.as_entire_binding(),
443                },
444                wgpu::BindGroupEntry {
445                    binding: 3,
446                    resource: bufs.lengths.as_entire_binding(),
447                },
448            ],
449        });
450
451        let mut encoder = self
452            .device
453            .create_command_encoder(&wgpu::CommandEncoderDescriptor {
454                label: Some("decompress_encoder"),
455            });
456        {
457            let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
458                label: Some("decompress_pass"),
459                timestamp_writes: None,
460            });
461            pass.set_pipeline(&self.pipeline);
462            pass.set_bind_group(0, &bind_group, &[]);
463            pass.dispatch_workgroups(1, 1, 1);
464        }
465        encoder.copy_buffer_to_buffer(&bufs.output, 0, &bufs.out_staging, 0, out_aligned);
466        encoder.copy_buffer_to_buffer(&bufs.lengths, 0, &bufs.len_staging, 0, len_size);
467        self.queue.submit(std::iter::once(encoder.finish()));
468
469        let (out_bytes, len_bytes) = self.readback_buffers(&bufs.out_staging, &bufs.len_staging)?;
470
471        Ok((out_bytes, parse_ss_lengths(&len_bytes, n)))
472    }
473}
474
475/// GPU buffers for a single `GDeflate` tile dispatch.
476struct GDeflateBuffers {
477    meta: wgpu::Buffer,
478    compressed: wgpu::Buffer,
479    output: wgpu::Buffer,
480    out_staging: wgpu::Buffer,
481}
482
483/// Allocate GPU buffers for a `GDeflate` tile dispatch.
484fn create_gdeflate_buffers(
485    device: &wgpu::Device,
486    comp_data: &[u8],
487    out_aligned: u64,
488) -> Result<GDeflateBuffers> {
489    std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
490        let buf = |label, size, usage| {
491            device.create_buffer(&wgpu::BufferDescriptor {
492                label: Some(label),
493                size,
494                usage,
495                mapped_at_creation: false,
496            })
497        };
498        let storage_dst = wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_DST;
499        let storage_src = wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_SRC;
500        let map_read = wgpu::BufferUsages::COPY_DST | wgpu::BufferUsages::MAP_READ;
501
502        GDeflateBuffers {
503            meta: buf(
504                "gdeflate_meta",
505                std::mem::size_of::<GDeflateMeta>() as u64,
506                storage_dst,
507            ),
508            compressed: buf("gdeflate_compressed", comp_data.len() as u64, storage_dst),
509            output: buf("gdeflate_output", out_aligned, storage_src),
510            out_staging: buf("gdeflate_staging", out_aligned, map_read),
511        }
512    }))
513    .map_err(|e| {
514        let msg = e
515            .downcast_ref::<String>()
516            .map(String::as_str)
517            .or_else(|| e.downcast_ref::<&str>().copied())
518            .unwrap_or("unknown GPU buffer allocation panic");
519        CrushError::from(PluginError::OperationFailed(format!(
520            "GDeflate GPU buffer allocation failed: {msg}"
521        )))
522    })
523}
524
525impl ComputeBackend for WgpuBackend {
526    #[allow(clippy::unnecessary_literal_bound)]
527    fn name(&self) -> &str {
528        "wgpu"
529    }
530
531    fn gpu_info(&self) -> &GpuInfo {
532        &self.info
533    }
534
535    fn decompress_tiles(
536        &self,
537        tiles: &[CompressedTile],
538        cancel: &AtomicBool,
539    ) -> Result<Vec<Vec<u8>>> {
540        // Wrap the entire GPU dispatch loop in catch_unwind so that panics
541        // inside wgpu (e.g. "device is lost", driver crashes, TDR) are
542        // converted to errors instead of crashing the process.
543        std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
544            self.decompress_tiles_inner(tiles, cancel)
545        }))
546        .unwrap_or_else(|e| {
547            let msg = e
548                .downcast_ref::<String>()
549                .map(String::as_str)
550                .or_else(|| e.downcast_ref::<&str>().copied())
551                .unwrap_or("unknown GPU panic");
552            Err(CrushError::from(PluginError::OperationFailed(format!(
553                "GPU panic caught (falling back to CPU): {msg}"
554            ))))
555        })
556    }
557
558    fn decompress_tiles_gdeflate(
559        &self,
560        tiles: &[CompressedTile],
561        cancel: &AtomicBool,
562    ) -> Result<Vec<Vec<u8>>> {
563        std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
564            self.decompress_tiles_gdeflate_inner(tiles, cancel)
565        }))
566        .unwrap_or_else(|e| {
567            let msg = e
568                .downcast_ref::<String>()
569                .map(String::as_str)
570                .or_else(|| e.downcast_ref::<&str>().copied())
571                .unwrap_or("unknown GPU panic");
572            Err(CrushError::from(PluginError::OperationFailed(format!(
573                "GDeflate GPU panic caught (falling back to CPU): {msg}"
574            ))))
575        })
576    }
577
578    fn release(&self) {
579        // wgpu resources are dropped automatically via RAII.
580    }
581}
582
583/// Validated and padded tile data ready for GPU upload.
584struct PreparedTile {
585    padded_data: Vec<u8>,
586    meta: GDeflateMeta,
587    out_aligned: u64,
588}
589
590/// Validate a tile and prepare its padded data + metadata for GPU dispatch.
591fn prepare_gdeflate_tile(tile: &CompressedTile) -> Result<PreparedTile> {
592    let max_tile_size: u32 = crate::format::DEFAULT_TILE_SIZE;
593    if tile.uncompressed_size > max_tile_size.saturating_mul(2) {
594        return Err(CrushError::InvalidFormat(format!(
595            "tile uncompressed_size {} exceeds maximum {}",
596            tile.uncompressed_size,
597            max_tile_size * 2,
598        )));
599    }
600
601    let mut padded_data = tile.data.clone();
602    while !padded_data.len().is_multiple_of(4) {
603        padded_data.push(0);
604    }
605
606    let meta = GDeflateMeta {
607        payload_size: u32::try_from(padded_data.len())
608            .map_err(|e| PluginError::OperationFailed(e.to_string()))?,
609        uncompressed_size: tile.uncompressed_size,
610        _pad0: 0,
611        _pad1: 0,
612    };
613
614    let out_aligned = u64::from(tile.uncompressed_size.div_ceil(4) * 4).max(4);
615
616    Ok(PreparedTile {
617        padded_data,
618        meta,
619        out_aligned,
620    })
621}
622
623/// Create a `GDeflate` bind group for a single tile's buffers.
624fn create_gdeflate_bind_group(
625    device: &wgpu::Device,
626    layout: &wgpu::BindGroupLayout,
627    bufs: &GDeflateBuffers,
628) -> wgpu::BindGroup {
629    device.create_bind_group(&wgpu::BindGroupDescriptor {
630        label: Some("gdeflate_bg"),
631        layout,
632        entries: &[
633            wgpu::BindGroupEntry {
634                binding: 0,
635                resource: bufs.meta.as_entire_binding(),
636            },
637            wgpu::BindGroupEntry {
638                binding: 1,
639                resource: bufs.compressed.as_entire_binding(),
640            },
641            wgpu::BindGroupEntry {
642                binding: 2,
643                resource: bufs.output.as_entire_binding(),
644            },
645        ],
646    })
647}
648
649impl WgpuBackend {
650    /// Dispatch a batch of `GDeflate` tiles in a single GPU submission.
651    ///
652    /// All tiles share one `CommandEncoder`, one `ComputePass`, one `queue.submit()`,
653    /// and one `device.poll()`. Each tile gets its own buffers and bind group since
654    /// buffer sizes vary per tile. This eliminates per-tile host-GPU synchronization
655    /// overhead.
656    fn dispatch_batch_gdeflate(&self, tiles: &[CompressedTile]) -> Result<Vec<Vec<u8>>> {
657        // Prepare all tiles (validate, pad, build metadata).
658        let prepared: Vec<PreparedTile> = tiles
659            .iter()
660            .map(prepare_gdeflate_tile)
661            .collect::<Result<Vec<_>>>()?;
662
663        // Allocate all GPU buffers upfront.
664        let all_bufs: Vec<GDeflateBuffers> = prepared
665            .iter()
666            .map(|p| create_gdeflate_buffers(&self.device, &p.padded_data, p.out_aligned))
667            .collect::<Result<Vec<_>>>()?;
668
669        // Upload all metadata and compressed data.
670        for (p, bufs) in prepared.iter().zip(all_bufs.iter()) {
671            self.queue
672                .write_buffer(&bufs.meta, 0, bytemuck::bytes_of(&p.meta));
673            self.queue.write_buffer(&bufs.compressed, 0, &p.padded_data);
674        }
675
676        // Create all bind groups.
677        let bind_groups: Vec<wgpu::BindGroup> = all_bufs
678            .iter()
679            .map(|bufs| create_gdeflate_bind_group(&self.device, &self.gdeflate_bgl, bufs))
680            .collect();
681
682        // One encoder, one compute pass, multiple dispatches.
683        let mut encoder = self
684            .device
685            .create_command_encoder(&wgpu::CommandEncoderDescriptor {
686                label: Some("gdeflate_batch_encoder"),
687            });
688        {
689            let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
690                label: Some("gdeflate_batch_pass"),
691                timestamp_writes: None,
692            });
693            pass.set_pipeline(&self.gdeflate_pipeline);
694            for bg in &bind_groups {
695                pass.set_bind_group(0, bg, &[]);
696                pass.dispatch_workgroups(1, 1, 1);
697            }
698        }
699
700        // Copy all output buffers to staging buffers.
701        for (p, bufs) in prepared.iter().zip(all_bufs.iter()) {
702            encoder.copy_buffer_to_buffer(&bufs.output, 0, &bufs.out_staging, 0, p.out_aligned);
703        }
704
705        // Single submit for all tiles.
706        self.queue.submit(std::iter::once(encoder.finish()));
707
708        // Map all staging buffers, single poll, collect results.
709        self.readback_batch(&all_bufs, tiles)
710    }
711
712    /// Map all staging buffers, poll once, and collect decompressed results.
713    fn readback_batch(
714        &self,
715        all_bufs: &[GDeflateBuffers],
716        tiles: &[CompressedTile],
717    ) -> Result<Vec<Vec<u8>>> {
718        let receivers: Vec<_> = all_bufs
719            .iter()
720            .map(|bufs| {
721                let slice = bufs.out_staging.slice(..);
722                let (tx, rx) = std::sync::mpsc::channel();
723                slice.map_async(wgpu::MapMode::Read, move |result| {
724                    let _ = tx.send(result);
725                });
726                rx
727            })
728            .collect();
729
730        self.device
731            .poll(wgpu::PollType::Wait {
732                submission_index: None,
733                timeout: Some(GPU_POLL_TIMEOUT),
734            })
735            .map_err(|e| {
736                PluginError::OperationFailed(format!(
737                    "GDeflate GPU poll failed (timeout or device lost): {e}"
738                ))
739            })?;
740
741        let mut results = Vec::with_capacity(tiles.len());
742        for (i, rx) in receivers.into_iter().enumerate() {
743            rx.recv()
744                .map_err(|e| {
745                    PluginError::OperationFailed(format!(
746                        "GDeflate GPU readback channel error: {e}"
747                    ))
748                })?
749                .map_err(|e| {
750                    PluginError::OperationFailed(format!("GDeflate GPU output map failed: {e}"))
751                })?;
752
753            let slice = all_bufs[i].out_staging.slice(..);
754            let out_bytes = slice
755                .get_mapped_range()
756                .map_err(|e| {
757                    PluginError::OperationFailed(format!(
758                        "GDeflate GPU output range map failed: {e}"
759                    ))
760                })?
761                .to_vec();
762            let size = tiles[i].uncompressed_size as usize;
763            results.push(out_bytes[..size.min(out_bytes.len())].to_vec());
764        }
765
766        Ok(results)
767    }
768
769    /// Inner dispatch loop for `GDeflate` tiles — batched for throughput.
770    ///
771    /// Processes tiles in chunks of `super::MAX_TILES_PER_BATCH`, checking for
772    /// cancellation between batches. Each batch is dispatched as a single
773    /// GPU submission to minimize host-GPU synchronization overhead.
774    fn decompress_tiles_gdeflate_inner(
775        &self,
776        tiles: &[CompressedTile],
777        cancel: &AtomicBool,
778    ) -> Result<Vec<Vec<u8>>> {
779        let mut results = Vec::with_capacity(tiles.len());
780        for batch in tiles.chunks(super::MAX_TILES_PER_BATCH) {
781            if cancel.load(Ordering::Relaxed) {
782                return Err(CrushError::Cancelled);
783            }
784            let batch_results = self.dispatch_batch_gdeflate(batch)?;
785            results.extend(batch_results);
786        }
787        Ok(results)
788    }
789
790    /// Inner dispatch loop, separated so `decompress_tiles` can wrap it in
791    /// `catch_unwind` to prevent wgpu panics from crashing the process.
792    fn decompress_tiles_inner(
793        &self,
794        tiles: &[CompressedTile],
795        cancel: &AtomicBool,
796    ) -> Result<Vec<Vec<u8>>> {
797        let mut results = Vec::with_capacity(tiles.len());
798
799        for (i, tile) in tiles.iter().enumerate() {
800            if cancel.load(Ordering::Relaxed) {
801                return Err(CrushError::Cancelled);
802            }
803
804            let tile_index =
805                u32::try_from(i).map_err(|e| PluginError::OperationFailed(e.to_string()))?;
806
807            let (raw_output, ss_lengths) = self.dispatch_tile(tile, tile_index)?;
808
809            let decompressed = super::deinterleave(
810                &raw_output,
811                &ss_lengths,
812                u32::from(tile.sub_stream_count),
813                tile.uncompressed_size,
814            );
815
816            results.push(decompressed);
817        }
818
819        Ok(results)
820    }
821}