Skip to main content

dynamis_gpu/
buffer.rs

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