1use std::sync::mpsc::{self, TryRecvError};
2use wgpu::{
3 Buffer, BufferAddress, BufferDescriptor, BufferUsages, Device, MapMode, PollType, Queue,
4};
5
6pub struct GpuBuffer {
7 buffer: Buffer,
8 size: BufferAddress,
9}
10
11impl GpuBuffer {
12 pub fn new(device: &Device, label: &str, size: BufferAddress, usage: BufferUsages) -> Self {
13 let buffer = device.create_buffer(&BufferDescriptor {
14 label: Some(label),
15 size,
16 usage,
17 mapped_at_creation: false,
18 });
19 Self { buffer, size }
20 }
21
22 pub fn write(&self, queue: &Queue, bytes: &[u8]) {
23 assert!(bytes.len() as u64 <= self.size, "write exceeds buffer size");
24 queue.write_buffer(&self.buffer, 0, bytes);
25 }
26
27 pub fn as_entire_binding(&self) -> wgpu::BindingResource<'_> {
28 wgpu::BindingResource::Buffer(wgpu::BufferBinding {
29 buffer: &self.buffer,
30 offset: 0,
31 size: None,
32 })
33 }
34
35 pub fn as_indirect_target(&self) -> &Buffer {
36 &self.buffer
37 }
38
39 pub fn size(&self) -> BufferAddress {
40 self.size
41 }
42
43 pub fn buffer(&self) -> &Buffer {
44 &self.buffer
45 }
46
47 pub fn read_sync(&self, device: &Device, queue: &Queue) -> Vec<u8> {
48 let staging = device.create_buffer(&BufferDescriptor {
49 label: Some("readback staging"),
50 size: self.size,
51 usage: BufferUsages::COPY_DST | BufferUsages::MAP_READ,
52 mapped_at_creation: false,
53 });
54 let mut encoder = device.create_command_encoder(&wgpu::CommandEncoderDescriptor {
55 label: Some("readback encoder"),
56 });
57 encoder.copy_buffer_to_buffer(&self.buffer, 0, &staging, 0, self.size);
58 queue.submit([encoder.finish()]);
59
60 let slice = staging.slice(..);
61 let (sender, receiver) = mpsc::channel();
62 slice.map_async(MapMode::Read, move |result| {
63 let _ = sender.send(result);
64 });
65 loop {
66 device
67 .poll(PollType::wait_indefinitely())
68 .expect("device lost while awaiting readback");
69 match receiver.try_recv() {
70 Ok(Ok(())) => break,
71 Ok(Err(error)) => panic!("buffer mapping failed: {error}"),
72 Err(TryRecvError::Empty) => std::thread::yield_now(),
73 Err(TryRecvError::Disconnected) => panic!("mapping callback was dropped"),
74 }
75 }
76 let bytes = slice
77 .get_mapped_range()
78 .expect("mapped range unavailable")
79 .to_vec();
80 staging.unmap();
81 bytes
82 }
83}