Skip to main content

dynamis_gpu/
readback.rs

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}