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