Skip to main content

dynamis_gpu/
buffer.rs

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