1use std::sync::mpsc::{self, TryRecvError};
2use wgpu::{
3 Buffer, BufferAddress, BufferAsyncError, BufferDescriptor, BufferUsages, Device, MapMode,
4 PollType, Queue,
5};
6
7pub struct GpuBuffer {
8 buffer: Buffer,
9 size: BufferAddress,
10 usage: BufferUsages,
11}
12
13impl GpuBuffer {
14 pub fn new(device: &Device, label: &str, size: BufferAddress, usage: BufferUsages) -> Self {
15 let buffer = device.create_buffer(&BufferDescriptor {
16 label: Some(label),
17 size,
18 usage,
19 mapped_at_creation: false,
20 });
21 Self {
22 buffer,
23 size,
24 usage,
25 }
26 }
27
28 pub fn write(&self, queue: &Queue, bytes: &[u8]) {
29 assert!(
30 self.usage.contains(BufferUsages::COPY_DST),
31 "buffer write requires COPY_DST usage"
32 );
33 assert!(bytes.len() as u64 <= self.size, "write exceeds buffer size");
34 queue.write_buffer(&self.buffer, 0, bytes);
35 }
36
37 pub fn as_binding(&self) -> wgpu::BindingResource<'_> {
38 wgpu::BindingResource::Buffer(wgpu::BufferBinding {
39 buffer: &self.buffer,
40 offset: 0,
41 size: None,
42 })
43 }
44
45 pub fn as_indirect_args(&self) -> &Buffer {
46 &self.buffer
47 }
48
49 pub fn buffer(&self) -> &Buffer {
50 &self.buffer
51 }
52
53 pub fn size(&self) -> BufferAddress {
54 self.size
55 }
56}
57
58struct PendingRead {
59 sequence: u64,
60 channel: Option<mpsc::Receiver<Result<(), BufferAsyncError>>>,
61}
62
63pub struct GpuReadback {
64 staging: [Buffer; 2],
65 size: BufferAddress,
66 next: usize,
67 pending: [Option<PendingRead>; 2],
68}
69
70impl GpuReadback {
71 pub fn new(device: &Device, label: &str, size: BufferAddress) -> Self {
72 assert!(size > 0, "readback size must be positive");
73 let staging = std::array::from_fn(|index| {
74 let staging_label = format!("{label} staging {index}");
75 device.create_buffer(&BufferDescriptor {
76 label: Some(&staging_label),
77 size,
78 usage: BufferUsages::COPY_DST | BufferUsages::MAP_READ,
79 mapped_at_creation: false,
80 })
81 });
82 Self {
83 staging,
84 size,
85 next: 0,
86 pending: [None, None],
87 }
88 }
89
90 pub fn enqueue(
91 &mut self,
92 device: &Device,
93 encoder: &mut wgpu::CommandEncoder,
94 source: &Buffer,
95 sequence: u64,
96 ) -> Option<(u64, Vec<u8>)> {
97 let slot = self.next;
98 self.next = (self.next + 1) % 2;
99 let displaced = if self.pending[slot].is_some() {
100 Some(self.drain(device, slot))
101 } else {
102 None
103 };
104 encoder.copy_buffer_to_buffer(source, 0, &self.staging[slot], 0, self.size);
105 self.pending[slot] = Some(PendingRead {
106 sequence,
107 channel: None,
108 });
109 displaced
110 }
111
112 pub fn arm(&mut self) {
113 for slot in 0..2 {
114 if self.pending[slot].is_some() {
115 self.ensure_mapping(slot);
116 }
117 }
118 }
119
120 pub fn poll(&mut self, device: &Device) -> Vec<(u64, Vec<u8>)> {
121 for slot in 0..2 {
122 if self.pending[slot].is_some() {
123 self.ensure_mapping(slot);
124 }
125 }
126 device
127 .poll(PollType::Poll)
128 .expect("device lost while polling readback");
129 let mut completed = Vec::new();
130 for slot in 0..2 {
131 if let Some(entry) = self.try_consume(slot) {
132 completed.push(entry);
133 }
134 }
135 completed.sort_unstable_by_key(|(sequence, _)| *sequence);
136 completed
137 }
138
139 fn ensure_mapping(&mut self, slot: usize) {
140 let pending = self.pending[slot]
141 .as_mut()
142 .expect("pending readback just checked");
143 if pending.channel.is_none() {
144 let (sender, receiver) = mpsc::channel();
145 self.staging[slot]
146 .slice(..)
147 .map_async(MapMode::Read, move |result| {
148 let _ = sender.send(result);
149 });
150 pending.channel = Some(receiver);
151 }
152 }
153
154 fn try_consume(&mut self, slot: usize) -> Option<(u64, Vec<u8>)> {
155 let pending = self.pending[slot].as_ref()?;
156 let channel = pending.channel.as_ref()?;
157 match channel.try_recv() {
158 Ok(Ok(())) => {
159 let pending = self.pending[slot]
160 .take()
161 .expect("pending readback just checked");
162 let bytes = self.staging[slot]
163 .slice(..)
164 .get_mapped_range()
165 .expect("mapped range unavailable")
166 .to_vec();
167 self.staging[slot].unmap();
168 Some((pending.sequence, bytes))
169 }
170 Ok(Err(error)) => panic!("buffer mapping failed: {error}"),
171 Err(TryRecvError::Empty) => None,
172 Err(TryRecvError::Disconnected) => {
173 panic!("mapping callback was dropped")
174 }
175 }
176 }
177
178 fn drain(&mut self, device: &Device, slot: usize) -> (u64, Vec<u8>) {
179 self.ensure_mapping(slot);
180 loop {
181 device
182 .poll(PollType::wait_indefinitely())
183 .expect("device lost while awaiting readback");
184 if let Some(entry) = self.try_consume(slot) {
185 return entry;
186 }
187 std::thread::yield_now();
188 }
189 }
190}