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