1use crate::SubmissionEncoder;
2use std::collections::VecDeque;
3use std::sync::mpsc::{self, Receiver, TryRecvError};
4use std::sync::{Arc, OnceLock};
5use std::time::{Duration, Instant};
6use wgpu::{
7 Buffer, BufferAddress, BufferAsyncError, BufferDescriptor, BufferUsages, Device, MapMode,
8 PollType, Queue, SubmissionIndex,
9};
10
11const READBACK_TIMEOUT: Duration = Duration::from_secs(30);
12
13struct Pending {
14 sequence: u64,
15 bytes: BufferAddress,
16 submission: Arc<OnceLock<SubmissionIndex>>,
17 completion: Receiver<Result<(), BufferAsyncError>>,
18}
19
20struct Slot {
21 label: String,
22 staging: Buffer,
23 pending: Option<Pending>,
24}
25
26impl Slot {
27 fn new(device: &Device, label: &str, size: BufferAddress) -> Self {
28 Self {
29 label: label.to_owned(),
30 staging: device.create_buffer(&BufferDescriptor {
31 label: Some(label),
32 size,
33 usage: BufferUsages::COPY_DST | BufferUsages::MAP_READ,
34 mapped_at_creation: false,
35 }),
36 pending: None,
37 }
38 }
39
40 fn is_sealed(&self) -> bool {
41 self.pending
42 .as_ref()
43 .is_some_and(|pending| pending.submission.get().is_some())
44 }
45
46 fn record(
47 &mut self,
48 encoder: &mut SubmissionEncoder,
49 regions: &[(&Buffer, BufferAddress, BufferAddress)],
50 sequence: u64,
51 ) {
52 assert!(
53 self.pending.is_none(),
54 "readback {} is still in use",
55 self.label
56 );
57 let mut at = 0;
58 for (source, source_offset, bytes) in regions {
59 assert!(
60 *bytes > 0 && bytes.is_multiple_of(4),
61 "readback length must be positive and word aligned"
62 );
63 assert!(
64 source_offset.is_multiple_of(4),
65 "readback source offset must be word aligned"
66 );
67 assert!(
68 *source_offset <= source.size() && *bytes <= source.size() - source_offset,
69 "readback exceeds source buffer"
70 );
71 encoder.copy_buffer_to_buffer(source, *source_offset, &self.staging, at, *bytes);
72 at += bytes;
73 }
74 assert!(
75 at > 0 && at <= self.staging.size(),
76 "readback exceeds staging capacity"
77 );
78 let (sender, completion) = mpsc::channel();
79 encoder.map_buffer_on_submit(&self.staging, MapMode::Read, ..at, move |result| {
80 let _ = sender.send(result);
81 });
82 self.pending = Some(Pending {
83 sequence,
84 bytes: at,
85 submission: encoder.submission(),
86 completion,
87 });
88 }
89
90 fn collect(&mut self) -> Option<(u64, Vec<u8>)> {
91 let pending = self.pending.as_ref()?;
92 pending.submission.get()?;
93 match pending.completion.try_recv() {
94 Ok(result) => Some(self.consume(result)),
95 Err(TryRecvError::Empty) => None,
96 Err(TryRecvError::Disconnected) => {
97 panic!("readback {} completion was dropped", self.label)
98 }
99 }
100 }
101
102 fn wait(&mut self, device: &Device) -> (u64, Vec<u8>) {
103 let pending = self
104 .pending
105 .as_ref()
106 .expect("cannot wait for an idle readback");
107 let submission = pending.submission.get().cloned().unwrap_or_else(|| {
108 panic!(
109 "readback {} must be submitted before waiting or reusing its slot",
110 self.label
111 )
112 });
113 let deadline = Instant::now() + READBACK_TIMEOUT;
114 device
115 .poll(PollType::Wait {
116 submission_index: Some(submission),
117 timeout: Some(READBACK_TIMEOUT),
118 })
119 .unwrap_or_else(|error| panic!("readback {} submission failed: {error}", self.label));
120 let result = pending
121 .completion
122 .recv_timeout(deadline.saturating_duration_since(Instant::now()))
123 .unwrap_or_else(|error| panic!("readback {} completion failed: {error}", self.label));
124 self.consume(result)
125 }
126
127 fn consume(&mut self, result: Result<(), BufferAsyncError>) -> (u64, Vec<u8>) {
128 result.unwrap_or_else(|error| panic!("readback {} mapping failed: {error}", self.label));
129 let pending = self
130 .pending
131 .take()
132 .expect("readback completion requires a pending read");
133 let bytes = self
134 .staging
135 .slice(..pending.bytes)
136 .get_mapped_range()
137 .expect("completed readback must be mapped")
138 .to_vec();
139 self.staging.unmap();
140 (pending.sequence, bytes)
141 }
142}
143
144pub const EVENT_SLOTS: u32 = Readback::DEPTH as u32 + 2;
145
146pub struct Readback {
147 device: Device,
148 label: String,
149 size: BufferAddress,
150 slots: VecDeque<Slot>,
151 inflight: usize,
152 last_sequence: Option<u64>,
153}
154
155impl Readback {
156 pub const DEPTH: usize = 4;
157
158 pub fn new(device: &Device, label: &str, size: BufferAddress, depth: usize) -> Self {
159 assert!(
160 size > 0 && size.is_multiple_of(4),
161 "readback size must be positive and word aligned"
162 );
163 assert!(depth > 0, "a readback needs at least one staging slot");
164 Self {
165 device: device.clone(),
166 label: label.to_owned(),
167 size,
168 slots: (0..depth)
169 .map(|index| Slot::new(device, &format!("{label} slot {index}"), size))
170 .collect(),
171 inflight: 0,
172 last_sequence: None,
173 }
174 }
175
176 pub fn size(&self) -> BufferAddress {
177 self.size
178 }
179
180 pub fn is_idle(&self) -> bool {
181 self.inflight == 0
182 }
183
184 pub fn enqueue(
185 &mut self,
186 encoder: &mut SubmissionEncoder,
187 source: &Buffer,
188 source_offset: BufferAddress,
189 bytes: BufferAddress,
190 sequence: u64,
191 ) -> Option<(u64, Vec<u8>)> {
192 self.enqueue_regions(encoder, &[(source, source_offset, bytes)], sequence)
193 }
194
195 pub fn enqueue_regions(
196 &mut self,
197 encoder: &mut SubmissionEncoder,
198 regions: &[(&Buffer, BufferAddress, BufferAddress)],
199 sequence: u64,
200 ) -> Option<(u64, Vec<u8>)> {
201 assert!(
202 self.last_sequence
203 .is_none_or(|previous| sequence > previous),
204 "readback {} sequences must increase",
205 self.label
206 );
207 let displaced = self.reclaim();
208 self.slots[self.inflight].record(encoder, regions, sequence);
209 self.inflight += 1;
210 self.last_sequence = Some(sequence);
211 displaced
212 }
213
214 pub fn collect(&mut self) -> Vec<(u64, Vec<u8>)> {
215 let mut completed = Vec::new();
216 while self.inflight > 0 {
217 let Some(entry) = self.slots[0].collect() else {
218 break;
219 };
220 self.retire();
221 completed.push(entry);
222 }
223 completed
224 }
225
226 pub fn drain(&mut self) -> Vec<(u64, Vec<u8>)> {
227 let mut completed = Vec::with_capacity(self.inflight);
228 while self.inflight > 0 {
229 let entry = self.slots[0].wait(&self.device);
230 self.retire();
231 completed.push(entry);
232 }
233 completed
234 }
235
236 fn reclaim(&mut self) -> Option<(u64, Vec<u8>)> {
237 if self.inflight < self.slots.len() {
238 return None;
239 }
240 if !self.slots[0].is_sealed() {
241 let index = self.slots.len();
242 self.slots.push_back(Slot::new(
243 &self.device,
244 &format!("{} slot {index}", self.label),
245 self.size,
246 ));
247 return None;
248 }
249 let mut oldest = self.slots.pop_front().expect("readback holds a slot");
250 let entry = oldest.wait(&self.device);
251 self.slots.push_back(oldest);
252 self.inflight -= 1;
253 Some(entry)
254 }
255
256 fn retire(&mut self) {
257 let slot = self.slots.pop_front().expect("readback holds a slot");
258 self.slots.push_back(slot);
259 self.inflight -= 1;
260 if self.inflight == 0 {
261 self.last_sequence = None;
262 }
263 }
264}
265
266pub fn read_regions(
267 device: &Device,
268 queue: &Queue,
269 label: &str,
270 regions: &[(&Buffer, BufferAddress, BufferAddress)],
271) -> Vec<u8> {
272 let bytes: BufferAddress = regions.iter().map(|region| region.2).sum();
273 let mut readback = Readback::new(device, label, bytes, 1);
274 let mut encoder = SubmissionEncoder::new(device, label);
275 assert!(
276 readback.enqueue_regions(&mut encoder, regions, 0).is_none(),
277 "a one shot read requires an idle readback"
278 );
279 encoder.submit(queue);
280 readback
281 .drain()
282 .pop()
283 .expect("a one shot read retires exactly once")
284 .1
285}