ursus_core/render/gfx/descriptor/
allocator.rs1use crate::render::gfx::descriptor::{BindingKind, DescriptorBindingDesc, DescriptorSetDesc};
2use crate::render::gfx::types::DescriptorSetId;
3use ash::vk;
4
5pub(crate) struct StoredDescriptorSet {
6 layout: vk::DescriptorSetLayout,
7 set: vk::DescriptorSet,
8 pool: vk::DescriptorPool,
9 bindings: Vec<DescriptorBindingDesc>,
10}
11
12pub struct DescriptorAllocator {
14 sets: Vec<StoredDescriptorSet>,
15 device: ash::Device,
16}
17
18impl DescriptorAllocator {
19 pub fn new(device: ash::Device) -> Self {
20 Self { sets: Vec::new(), device }
21 }
22
23 pub fn create_set(&mut self, desc: DescriptorSetDesc) -> anyhow::Result<DescriptorSetId> {
24 let has_bindless = desc.bindings.iter().any(|b| b.bindless);
25
26 let vk_bindings: Vec<vk::DescriptorSetLayoutBinding> = desc
27 .bindings
28 .iter()
29 .map(|b| {
30 let mut vb = vk::DescriptorSetLayoutBinding::default()
31 .binding(b.binding)
32 .descriptor_type(to_vk_type(b.kind))
33 .descriptor_count(b.count)
34 .stage_flags(b.stage.to_vk());
35 if let Some(sampler) = &b.immutable_sampler {
36 vb = vb.immutable_samplers(std::slice::from_ref(sampler));
37 }
38 vb
39 })
40 .collect();
41
42 let binding_flags: Vec<vk::DescriptorBindingFlags> = desc
43 .bindings
44 .iter()
45 .map(|b| {
46 if b.bindless {
47 vk::DescriptorBindingFlags::PARTIALLY_BOUND
48 | vk::DescriptorBindingFlags::VARIABLE_DESCRIPTOR_COUNT
49 | vk::DescriptorBindingFlags::UPDATE_AFTER_BIND
50 } else {
51 vk::DescriptorBindingFlags::empty()
52 }
53 })
54 .collect();
55
56 let mut flags_info = vk::DescriptorSetLayoutBindingFlagsCreateInfo::default().binding_flags(&binding_flags);
57
58 let mut layout_info = vk::DescriptorSetLayoutCreateInfo::default().bindings(&vk_bindings);
59 if has_bindless {
60 layout_info = layout_info
61 .flags(vk::DescriptorSetLayoutCreateFlags::UPDATE_AFTER_BIND_POOL)
62 .push_next(&mut flags_info);
63 }
64
65 let layout = unsafe { self.device.create_descriptor_set_layout(&layout_info, None)? };
66
67 let pool_sizes: Vec<vk::DescriptorPoolSize> = desc
68 .bindings
69 .iter()
70 .map(|b| vk::DescriptorPoolSize { ty: to_vk_type(b.kind), descriptor_count: b.count })
71 .collect();
72
73 let mut pool_info = vk::DescriptorPoolCreateInfo::default().pool_sizes(&pool_sizes).max_sets(1);
74 if has_bindless {
75 pool_info = pool_info.flags(vk::DescriptorPoolCreateFlags::UPDATE_AFTER_BIND);
76 }
77 let pool = unsafe { self.device.create_descriptor_pool(&pool_info, None)? };
78
79 let variable_count: Option<u32> = desc.bindings.iter().find(|b| b.bindless).map(|b| b.count);
80 let variable_count_value: u32 = variable_count.unwrap_or(0);
81 let mut var_count_info = variable_count.map(|_| {
82 vk::DescriptorSetVariableDescriptorCountAllocateInfo::default()
83 .descriptor_counts(std::slice::from_ref(&variable_count_value))
84 });
85
86 let mut alloc_info =
87 vk::DescriptorSetAllocateInfo::default().descriptor_pool(pool).set_layouts(std::slice::from_ref(&layout));
88 if let Some(vci) = var_count_info.as_mut() {
89 alloc_info = alloc_info.push_next(vci);
90 }
91
92 let set = unsafe { self.device.allocate_descriptor_sets(&alloc_info)?[0] };
93
94 let id = DescriptorSetId(self.sets.len() as u32);
95 self.sets.push(StoredDescriptorSet { layout, set, pool, bindings: desc.bindings });
96 Ok(id)
97 }
98
99 pub fn write_sampled_image_array(
100 &self,
101 set: DescriptorSetId,
102 binding: u32,
103 array_element: u32,
104 view: vk::ImageView,
105 ) {
106 let stored = &self.sets[set.0 as usize];
107 let image_info =
108 vk::DescriptorImageInfo::default().image_view(view).image_layout(vk::ImageLayout::SHADER_READ_ONLY_OPTIMAL);
109 let write = vk::WriteDescriptorSet::default()
110 .dst_set(stored.set)
111 .dst_binding(binding)
112 .dst_array_element(array_element)
113 .descriptor_type(vk::DescriptorType::SAMPLED_IMAGE)
114 .image_info(std::slice::from_ref(&image_info));
115 unsafe { self.device.update_descriptor_sets(std::slice::from_ref(&write), &[]) };
116 }
117
118 pub(crate) fn layout(&self, id: DescriptorSetId) -> vk::DescriptorSetLayout {
119 self.sets[id.0 as usize].layout
120 }
121
122 pub fn handle(&self, id: DescriptorSetId) -> vk::DescriptorSet {
123 self.sets[id.0 as usize].set
124 }
125
126 pub fn bind_uniform_buffer(
127 &self,
128 set: DescriptorSetId,
129 binding: u32,
130 buffer: vk::Buffer,
131 size: vk::DeviceSize,
132 ) -> anyhow::Result<()> {
133 self.check_kind(set, binding, "UniformBuffer", |k| matches!(k, BindingKind::UniformBuffer { .. }))?;
134
135 let stored = &self.sets[set.0 as usize];
136 let buf_info = vk::DescriptorBufferInfo::default().buffer(buffer).offset(0).range(size);
137 let write = vk::WriteDescriptorSet::default()
138 .dst_set(stored.set)
139 .dst_binding(binding)
140 .descriptor_type(vk::DescriptorType::UNIFORM_BUFFER)
141 .buffer_info(std::slice::from_ref(&buf_info));
142
143 unsafe { self.device.update_descriptor_sets(std::slice::from_ref(&write), &[]) };
144 Ok(())
145 }
146
147 pub fn bind_mapped_uniform_buffer<T: Copy>(
148 &self,
149 set: DescriptorSetId,
150 binding: u32,
151 mapped: &crate::vulkan::MappedGpuBuffer<T>,
152 ) -> anyhow::Result<()> {
153 self.bind_uniform_buffer(set, binding, mapped.buffer, mapped.size())
154 }
155
156 pub fn bind_storage_buffer(
157 &self,
158 set: DescriptorSetId,
159 binding: u32,
160 buffer: vk::Buffer,
161 size: vk::DeviceSize,
162 ) -> anyhow::Result<()> {
163 self.check_kind(set, binding, "StorageBuffer", |k| matches!(k, BindingKind::StorageBuffer { .. }))?;
164
165 let stored = &self.sets[set.0 as usize];
166 let buf_info = vk::DescriptorBufferInfo::default().buffer(buffer).offset(0).range(size);
167 let write = vk::WriteDescriptorSet::default()
168 .dst_set(stored.set)
169 .dst_binding(binding)
170 .descriptor_type(vk::DescriptorType::STORAGE_BUFFER)
171 .buffer_info(std::slice::from_ref(&buf_info));
172
173 unsafe { self.device.update_descriptor_sets(std::slice::from_ref(&write), &[]) };
174 Ok(())
175 }
176
177 pub fn bind_mapped_storage_buffer<T: Copy>(
178 &self,
179 set: DescriptorSetId,
180 binding: u32,
181 mapped: &crate::vulkan::MappedGpuBuffer<T>,
182 ) -> anyhow::Result<()> {
183 self.bind_storage_buffer(set, binding, mapped.buffer, mapped.size())
184 }
185
186 pub fn bind_sampled_image(
187 &self,
188 set: DescriptorSetId,
189 binding: u32,
190 view: vk::ImageView,
191 layout: vk::ImageLayout,
192 sampler: vk::Sampler,
193 ) -> anyhow::Result<()> {
194 self.check_kind(set, binding, "CombinedImageSampler", |k| matches!(k, BindingKind::CombinedImageSampler))?;
195
196 let stored = &self.sets[set.0 as usize];
197 let image_info = vk::DescriptorImageInfo::default().image_view(view).image_layout(layout).sampler(sampler);
198 let write = vk::WriteDescriptorSet::default()
199 .dst_set(stored.set)
200 .dst_binding(binding)
201 .descriptor_type(vk::DescriptorType::COMBINED_IMAGE_SAMPLER)
202 .image_info(std::slice::from_ref(&image_info));
203
204 unsafe { self.device.update_descriptor_sets(std::slice::from_ref(&write), &[]) };
205 Ok(())
206 }
207
208 fn check_kind(
209 &self,
210 set: DescriptorSetId,
211 binding: u32,
212 expected_name: &str,
213 matches: impl Fn(BindingKind) -> bool,
214 ) -> anyhow::Result<()> {
215 let stored = self
216 .sets
217 .get(set.0 as usize)
218 .ok_or_else(|| anyhow::anyhow!("DescriptorAllocator: DescriptorSetId {:?} not found", set))?;
219
220 match stored.bindings.iter().find(|b| b.binding == binding) {
221 Some(b) if matches(b.kind) => Ok(()),
222 Some(b) => anyhow::bail!(
223 "DescriptorAllocator: binding {} in set {:?} is declared as {:?}, expected {}",
224 binding,
225 set,
226 b.kind,
227 expected_name
228 ),
229 None => anyhow::bail!("DescriptorAllocator: binding {} is not declared in set {:?}", binding, set),
230 }
231 }
232}
233
234fn to_vk_type(kind: BindingKind) -> vk::DescriptorType {
235 match kind {
236 BindingKind::CombinedImageSampler => vk::DescriptorType::COMBINED_IMAGE_SAMPLER,
237 BindingKind::UniformBuffer { .. } => vk::DescriptorType::UNIFORM_BUFFER,
238 BindingKind::StorageBuffer { .. } => vk::DescriptorType::STORAGE_BUFFER,
239 BindingKind::Sampler => vk::DescriptorType::SAMPLER,
240 BindingKind::SampledImageArray => vk::DescriptorType::SAMPLED_IMAGE,
241 }
242}
243
244impl Drop for DescriptorAllocator {
245 fn drop(&mut self) {
246 unsafe {
247 for ds in &self.sets {
248 self.device.destroy_descriptor_pool(ds.pool, None);
249 self.device.destroy_descriptor_set_layout(ds.layout, None);
250 }
251 }
252 }
253}