1use std::ops::Not;
13
14use bevy::{
15 asset::RenderAssetUsages,
16 color::palettes::tailwind::{RED_400, SKY_400},
17 core_pipeline::schedule::camera_driver,
18 mesh::Indices,
19 platform::collections::HashSet,
20 prelude::*,
21 render::{
22 extract_component::{ExtractComponent, ExtractComponentPlugin},
23 mesh::allocator::{MeshAllocator, MeshAllocatorSettings},
24 render_resource::{
25 binding_types::{storage_buffer, uniform_buffer},
26 *,
27 },
28 renderer::{RenderContext, RenderGraph, RenderQueue},
29 Render, RenderApp, RenderStartup,
30 },
31};
32
33const SHADER_ASSET_PATH: &str = "shaders/compute_mesh.wesl";
35
36fn main() {
37 App::new()
38 .add_plugins((
39 DefaultPlugins,
40 ComputeShaderMeshGeneratorPlugin,
41 ExtractComponentPlugin::<GenerateMesh>::default(),
42 ))
43 .insert_resource(ClearColor(Color::BLACK))
44 .add_systems(Startup, setup)
45 .run();
46}
47
48struct ComputeShaderMeshGeneratorPlugin;
50impl Plugin for ComputeShaderMeshGeneratorPlugin {
51 fn build(&self, app: &mut App) {
52 let Some(render_app) = app.get_sub_app_mut(RenderApp) else {
53 return;
54 };
55
56 render_app
57 .init_resource::<ChunksToProcess>()
58 .add_systems(RenderStartup, init_compute_pipeline)
59 .add_systems(Render, prepare_chunks)
60 .add_systems(RenderGraph, compute_mesh.before(camera_driver));
61 }
62 fn finish(&self, app: &mut App) {
63 let Some(render_app) = app.get_sub_app_mut(RenderApp) else {
64 return;
65 };
66 render_app
67 .world_mut()
68 .resource_mut::<MeshAllocatorSettings>()
69 .extra_buffer_usages = BufferUsages::STORAGE;
74 }
75}
76
77#[derive(Component, ExtractComponent, Clone)]
80#[extract_app(RenderApp)]
81struct GenerateMesh(Handle<Mesh>);
82
83fn setup(
84 mut commands: Commands,
85 mut meshes: ResMut<Assets<Mesh>>,
86 mut materials: ResMut<Assets<StandardMaterial>>,
87) {
88 let empty_mesh = {
103 let mut mesh = Mesh::new(
104 PrimitiveTopology::TriangleList,
105 RenderAssetUsages::RENDER_WORLD,
106 )
107 .with_inserted_attribute(Mesh::ATTRIBUTE_POSITION, vec![[0.; 3]; 50])
108 .with_inserted_attribute(Mesh::ATTRIBUTE_NORMAL, vec![[0.; 3]; 50])
109 .with_inserted_attribute(Mesh::ATTRIBUTE_UV_0, vec![[0.; 2]; 50])
110 .with_inserted_indices(Indices::U32(vec![0; 50]));
111
112 mesh.asset_usage = RenderAssetUsages::RENDER_WORLD;
113 mesh
114 };
115
116 let handle = meshes.add(empty_mesh);
117
118 commands.spawn((
122 GenerateMesh(handle.clone()),
123 Mesh3d(handle.clone()),
124 MeshMaterial3d(materials.add(StandardMaterial {
125 base_color: RED_400.into(),
126 ..default()
127 })),
128 Transform::from_xyz(-2.5, 1.5, 0.),
129 ));
130
131 commands.spawn((
132 Mesh3d(handle),
133 MeshMaterial3d(materials.add(StandardMaterial {
134 base_color: SKY_400.into(),
135 ..default()
136 })),
137 Transform::from_xyz(2.5, 1.5, 0.),
138 ));
139
140 commands.spawn((
145 Mesh3d(meshes.add(Circle::new(4.0))),
146 MeshMaterial3d(materials.add(Color::WHITE)),
147 Transform::from_rotation(Quat::from_rotation_x(-std::f32::consts::FRAC_PI_2)),
148 ));
149 commands.spawn((
150 PointLight {
151 shadow_maps_enabled: true,
152 ..default()
153 },
154 Transform::from_xyz(4.0, 8.0, 4.0),
155 ));
156 commands.spawn((
158 Camera3d::default(),
159 Transform::from_xyz(-2.5, 4.5, 9.0).looking_at(Vec3::ZERO, Vec3::Y),
160 ));
161}
162
163#[derive(Resource, Default)]
167struct ChunksToProcess(Vec<AssetId<Mesh>>);
168
169fn prepare_chunks(
175 meshes_to_generate: Query<&GenerateMesh>,
176 mut chunks: ResMut<ChunksToProcess>,
177 pipeline_cache: Res<PipelineCache>,
178 pipeline: Res<ComputePipeline>,
179 mut processed: Local<HashSet<AssetId<Mesh>>>,
180) {
181 if pipeline_cache
185 .get_compute_pipeline(pipeline.pipeline)
186 .is_some()
187 {
188 let chunk_data: Vec<AssetId<Mesh>> = meshes_to_generate
192 .iter()
193 .filter_map(|gmesh| {
194 let id = gmesh.0.id();
195 processed.contains(&id).not().then_some(id)
196 })
197 .collect();
198
199 for id in &chunk_data {
201 processed.insert(*id);
202 }
203
204 chunks.0 = chunk_data;
205 }
206}
207
208#[derive(Resource)]
209struct ComputePipeline {
210 layout: BindGroupLayoutDescriptor,
211 pipeline: CachedComputePipelineId,
212}
213
214fn init_compute_pipeline(
216 mut commands: Commands,
217 asset_server: Res<AssetServer>,
218 pipeline_cache: Res<PipelineCache>,
219) {
220 let layout = BindGroupLayoutDescriptor::new(
221 "",
222 &BindGroupLayoutEntries::sequential(
223 ShaderStages::COMPUTE,
224 (
225 uniform_buffer::<DataRanges>(false),
227 storage_buffer::<Vec<f32>>(false),
229 storage_buffer::<Vec<u32>>(false),
231 ),
232 ),
233 );
234 let shader = asset_server.load(SHADER_ASSET_PATH);
235 let pipeline = pipeline_cache.queue_compute_pipeline(ComputePipelineDescriptor {
236 label: Some("Mesh generation compute shader".into()),
237 layout: vec![layout.clone()],
238 shader: shader.clone(),
239 ..default()
240 });
241 commands.insert_resource(ComputePipeline { layout, pipeline });
242}
243
244#[derive(ShaderType)]
247struct DataRanges {
248 vertex_start: u32,
249 vertex_end: u32,
250 index_start: u32,
251 index_end: u32,
252}
253
254fn compute_mesh(
255 mut render_context: RenderContext,
256 chunks: Res<ChunksToProcess>,
257 mesh_allocator: Res<MeshAllocator>,
258 pipeline_cache: Res<PipelineCache>,
259 pipeline: Res<ComputePipeline>,
260 render_queue: Res<RenderQueue>,
261) {
262 let Some(init_pipeline) = pipeline_cache.get_compute_pipeline(pipeline.pipeline) else {
263 return;
264 };
265
266 for mesh_id in &chunks.0 {
267 info!(?mesh_id, "processing mesh");
268
269 let vertex_buffer_slice = mesh_allocator.mesh_vertex_slice(mesh_id).unwrap();
274 let index_buffer_slice = mesh_allocator.mesh_index_slice(mesh_id).unwrap();
275
276 let first = DataRanges {
277 vertex_start: vertex_buffer_slice.range.start * 8,
282 vertex_end: vertex_buffer_slice.range.end * 8,
283 index_start: index_buffer_slice.range.start,
286 index_end: index_buffer_slice.range.end,
287 };
288
289 let mut uniforms = UniformBuffer::from(first);
290 uniforms.write_buffer(render_context.render_device(), &render_queue);
291
292 let bind_group = render_context.render_device().create_bind_group(
295 None,
296 &pipeline_cache.get_bind_group_layout(&pipeline.layout),
297 &BindGroupEntries::sequential((
298 &uniforms,
299 vertex_buffer_slice.buffer.as_entire_buffer_binding(),
300 index_buffer_slice.buffer.as_entire_buffer_binding(),
301 )),
302 );
303
304 let mut pass =
305 render_context
306 .command_encoder()
307 .begin_compute_pass(&ComputePassDescriptor {
308 label: Some("Mesh generation compute pass"),
309 ..default()
310 });
311 pass.push_debug_group("compute_mesh");
312
313 pass.set_bind_group(0, &bind_group, &[]);
314 pass.set_pipeline(init_pipeline);
315 pass.dispatch_workgroups(1, 1, 1);
318
319 pass.pop_debug_group();
320 }
321}