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}