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