Skip to main content

ursus_core/render/gfx/descriptor/
allocator.rs

1use 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
12/// A single entry point for creating and populating regular (non-bindless) descriptor sets.
13pub 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}