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 buffer(&self) -> &Buffer {
88 &self.buffer
89 }
90
91 pub fn size(&self) -> BufferAddress {
92 self.size
93 }
94}
95
96struct PendingRead {
97 sequence: u64,
98 mapping_started: bool,
99 result: Arc<Mutex<Option<Result<(), BufferAsyncError>>>>,
100}
101
102pub struct GpuReadback {
103 staging: [Buffer; 2],
104 size: BufferAddress,
105 next: usize,
106 pending: [Option<PendingRead>; 2],
107}
108
109impl GpuReadback {
110 pub fn new(device: &Device, label: &str, size: BufferAddress) -> Self {
111 assert!(size > 0, "readback size must be positive");
112 let staging = std::array::from_fn(|index| {
113 let staging_label = format!("{label} staging {index}");
114 device.create_buffer(&BufferDescriptor {
115 label: Some(&staging_label),
116 size,
117 usage: BufferUsages::COPY_DST | BufferUsages::MAP_READ,
118 mapped_at_creation: false,
119 })
120 });
121 Self {
122 staging,
123 size,
124 next: 0,
125 pending: [None, None],
126 }
127 }
128
129 pub fn enqueue(
130 &mut self,
131 device: &Device,
132 encoder: &mut wgpu::CommandEncoder,
133 source: &Buffer,
134 sequence: u64,
135 ) -> Option<(u64, Vec<u8>)> {
136 let slot = self.next;
137 self.next = (self.next + 1) % 2;
138 let displaced = if self.pending[slot].is_some() {
139 Some(self.drain(device, slot))
140 } else {
141 None
142 };
143 encoder.copy_buffer_to_buffer(source, 0, &self.staging[slot], 0, self.size);
144 self.pending[slot] = Some(PendingRead {
145 sequence,
146 mapping_started: false,
147 result: Arc::new(Mutex::new(None)),
148 });
149 displaced
150 }
151
152 pub fn arm(&mut self) {
153 for slot in 0..2 {
154 if self.pending[slot].is_some() {
155 self.ensure_mapping(slot);
156 }
157 }
158 }
159
160 pub fn poll(&mut self, device: &Device) -> Vec<(u64, Vec<u8>)> {
161 for slot in 0..2 {
162 if self.pending[slot].is_some() {
163 self.ensure_mapping(slot);
164 }
165 }
166 device
167 .poll(PollType::Poll)
168 .expect("device lost while polling readback");
169 let mut completed = Vec::new();
170 for slot in 0..2 {
171 if let Some(entry) = self.try_consume(slot) {
172 completed.push(entry);
173 }
174 }
175 completed.sort_unstable_by_key(|(sequence, _)| *sequence);
176 completed
177 }
178
179 fn ensure_mapping(&mut self, slot: usize) {
180 let pending = self.pending[slot]
181 .as_mut()
182 .expect("pending readback just checked");
183 if !pending.mapping_started {
184 pending.mapping_started = true;
185 let result = pending.result.clone();
186 self.staging[slot]
187 .slice(..)
188 .map_async(MapMode::Read, move |value| {
189 *result.lock().unwrap() = Some(value);
190 });
191 }
192 }
193
194 fn try_consume(&mut self, slot: usize) -> Option<(u64, Vec<u8>)> {
195 let result = {
196 let pending = self.pending[slot].as_ref()?;
197 pending.result.lock().unwrap().take()?
198 };
199 result.unwrap_or_else(|error| panic!("buffer mapping failed: {error}"));
200 let pending = self.pending[slot]
201 .take()
202 .expect("pending readback just checked");
203 let bytes = self.staging[slot]
204 .slice(..)
205 .get_mapped_range()
206 .expect("mapped range unavailable")
207 .to_vec();
208 self.staging[slot].unmap();
209 Some((pending.sequence, bytes))
210 }
211
212 fn drain(&mut self, device: &Device, slot: usize) -> (u64, Vec<u8>) {
213 self.ensure_mapping(slot);
214 loop {
215 device
216 .poll(PollType::wait_indefinitely())
217 .expect("device lost while awaiting readback");
218 if let Some(entry) = self.try_consume(slot) {
219 return entry;
220 }
221 std::thread::yield_now();
222 }
223 }
224}