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