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
48pub struct Readback {
49 staging: [Buffer; 2],
50 size: BufferAddress,
51 submitted: u64,
52}
53
54impl Readback {
55 pub fn new(device: &Device, label: &str, size: BufferAddress) -> Self {
56 assert!(size > 0, "readback size must be positive");
57 let staging = std::array::from_fn(|index| {
58 let staging_label = format!("{label} staging {index}");
59 device.create_buffer(&BufferDescriptor {
60 label: Some(&staging_label),
61 size,
62 usage: BufferUsages::COPY_DST | BufferUsages::MAP_READ,
63 mapped_at_creation: false,
64 })
65 });
66 Self {
67 staging,
68 size,
69 submitted: 0,
70 }
71 }
72
73 pub fn enqueue(&mut self, encoder: &mut wgpu::CommandEncoder, source: &Buffer) {
74 let index = (self.submitted % 2) as usize;
75 encoder.copy_buffer_to_buffer(source, 0, &self.staging[index], 0, self.size);
76 self.submitted += 1;
77 }
78
79 pub fn read(&mut self, device: &Device) -> Vec<u8> {
80 assert!(
81 self.submitted > 0,
82 "no readback has been submitted to the GPU"
83 );
84 let index = ((self.submitted - 1) % 2) as usize;
85 let staging = &self.staging[index];
86 let slice = staging.slice(..);
87 let (sender, receiver) = mpsc::channel();
88 slice.map_async(MapMode::Read, move |result| {
89 let _ = sender.send(result);
90 });
91 loop {
92 device
93 .poll(PollType::wait_indefinitely())
94 .expect("device lost while awaiting readback");
95 match receiver.try_recv() {
96 Ok(Ok(())) => break,
97 Ok(Err(error)) => panic!("buffer mapping failed: {error}"),
98 Err(TryRecvError::Empty) => std::thread::yield_now(),
99 Err(TryRecvError::Disconnected) => panic!("mapping callback was dropped"),
100 }
101 }
102 let bytes = slice
103 .get_mapped_range()
104 .expect("mapped range unavailable")
105 .to_vec();
106 staging.unmap();
107 bytes
108 }
109}