Skip to main content

ranim_render/primitives/
mesh_items.rs

1use crate::utils::{WgpuContext, WgpuVecBuffer};
2use bytemuck::{Pod, Zeroable};
3use glam::Vec3;
4use ranim_core::{components::rgba::Rgba, core_item::mesh_item::MeshItem};
5
6#[repr(C)]
7#[derive(Debug, Default, Clone, Copy, Pod, Zeroable)]
8pub struct MeshTransform {
9    pub transform: [[f32; 4]; 4],
10}
11
12#[derive(bevy_ecs::prelude::Resource)]
13pub struct MeshItemsBuffer {
14    /// Per-vertex positions (vertex buffer)
15    pub(crate) vertices_buffer: WgpuVecBuffer<Vec3>,
16    /// Per-vertex mesh id (vertex buffer)
17    pub(crate) mesh_ids_buffer: WgpuVecBuffer<u32>,
18    /// Per-vertex colors (vertex buffer)
19    pub(crate) vertex_colors_buffer: WgpuVecBuffer<Rgba>,
20    /// Per-vertex normals (vertex buffer) — all-zero → flat shading fallback
21    pub(crate) vertex_normals_buffer: WgpuVecBuffer<Vec3>,
22    /// Merged triangle indices (index buffer)
23    pub(crate) indices_buffer: WgpuVecBuffer<u32>,
24
25    /// Per-mesh transform matrices (storage buffer, indexed by mesh_id)
26    pub(crate) transforms_buffer: WgpuVecBuffer<MeshTransform>,
27    /// Per-mesh global scene order (storage buffer, indexed by mesh_id),
28    /// consumed by the fragment depth-order bias.
29    pub(crate) orders_buffer: WgpuVecBuffer<f32>,
30
31    pub(crate) item_count: u32,
32    pub(crate) total_vertices: u32,
33    pub(crate) total_indices: u32,
34
35    pub(crate) render_bind_group: Option<wgpu::BindGroup>,
36}
37
38impl MeshItemsBuffer {
39    pub fn new(ctx: &WgpuContext) -> Self {
40        let vertex_usage = wgpu::BufferUsages::VERTEX | wgpu::BufferUsages::COPY_DST;
41        let index_usage = wgpu::BufferUsages::INDEX | wgpu::BufferUsages::COPY_DST;
42        let storage_ro = wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_DST;
43
44        Self {
45            vertices_buffer: WgpuVecBuffer::new(ctx, Some("MeshVertices"), vertex_usage, 1),
46            mesh_ids_buffer: WgpuVecBuffer::new(ctx, Some("MeshIds"), vertex_usage, 1),
47            vertex_colors_buffer: WgpuVecBuffer::new(
48                ctx,
49                Some("MeshVertexColors"),
50                vertex_usage,
51                1,
52            ),
53            vertex_normals_buffer: WgpuVecBuffer::new(
54                ctx,
55                Some("MeshVertexNormals"),
56                vertex_usage,
57                1,
58            ),
59            indices_buffer: WgpuVecBuffer::new(ctx, Some("MeshIndices"), index_usage, 1),
60            transforms_buffer: WgpuVecBuffer::new(ctx, Some("MeshTransforms"), storage_ro, 1),
61            orders_buffer: WgpuVecBuffer::new(ctx, Some("MeshOrders"), storage_ro, 1),
62            item_count: 0,
63            total_vertices: 0,
64            total_indices: 0,
65            render_bind_group: None,
66        }
67    }
68
69    pub fn update<'a, I>(&mut self, ctx: &WgpuContext, mesh_items: I)
70    where
71        I: IntoIterator<Item = (f32, &'a MeshItem)>,
72        I::IntoIter: ExactSizeIterator + Clone,
73    {
74        let mesh_items = mesh_items.into_iter();
75        if mesh_items.len() == 0 {
76            self.item_count = 0;
77            self.total_vertices = 0;
78            self.total_indices = 0;
79            return;
80        }
81
82        let item_count = mesh_items.len();
83        let total_vertices: usize = mesh_items.clone().map(|(_, m)| m.points.len()).sum();
84        let total_indices: usize = mesh_items
85            .clone()
86            .map(|(_, m)| m.triangle_indices.len())
87            .sum();
88
89        let mut transforms = Vec::with_capacity(item_count);
90        let mut all_orders = Vec::with_capacity(item_count);
91        let mut all_vertices = Vec::with_capacity(total_vertices);
92        let mut all_mesh_ids = Vec::with_capacity(total_vertices);
93        let mut all_vertex_colors = Vec::with_capacity(total_vertices);
94        let mut all_vertex_normals = Vec::with_capacity(total_vertices);
95        let mut all_indices = Vec::with_capacity(total_indices);
96
97        let mut vertex_offset: u32 = 0;
98
99        for (mesh_idx, (order, mesh)) in mesh_items.enumerate() {
100            let vc = mesh.points.len() as u32;
101
102            transforms.push(MeshTransform {
103                transform: mesh.transform.to_cols_array_2d(),
104            });
105            all_orders.push(order);
106
107            all_vertices.extend_from_slice(&mesh.points);
108            all_mesh_ids.extend(std::iter::repeat_n(mesh_idx as u32, vc as usize));
109            all_vertex_colors.extend_from_slice(&mesh.vertex_colors);
110
111            // Pad normals with zero if shorter than points (flat shading fallback)
112            let normals = &mesh.vertex_normals;
113            let normals_len = normals.len();
114            if normals_len >= vc as usize {
115                all_vertex_normals.extend_from_slice(&normals[..vc as usize]);
116            } else {
117                all_vertex_normals.extend_from_slice(normals);
118                all_vertex_normals
119                    .extend(std::iter::repeat_n(Vec3::ZERO, vc as usize - normals_len));
120            }
121
122            all_indices.extend(mesh.triangle_indices.iter().map(|&i| i + vertex_offset));
123
124            vertex_offset += vc;
125        }
126
127        self.item_count = item_count as u32;
128        self.total_vertices = total_vertices as u32;
129        self.total_indices = total_indices as u32;
130
131        // Vertex/index buffers (no bind group dependency)
132        self.vertices_buffer.set(ctx, &all_vertices);
133        self.mesh_ids_buffer.set(ctx, &all_mesh_ids);
134        self.vertex_colors_buffer.set(ctx, &all_vertex_colors);
135        self.vertex_normals_buffer.set(ctx, &all_vertex_normals);
136        self.indices_buffer.set(ctx, &all_indices);
137
138        // Storage buffers (bind group recreated on realloc)
139        let mut any_realloc = self.transforms_buffer.set(ctx, &transforms);
140        any_realloc |= self.orders_buffer.set(ctx, &all_orders);
141
142        if any_realloc || self.render_bind_group.is_none() {
143            self.render_bind_group = Some(Self::create_render_bind_group(ctx, self));
144        }
145    }
146
147    pub fn item_count(&self) -> u32 {
148        self.item_count
149    }
150
151    pub fn total_indices(&self) -> u32 {
152        self.total_indices
153    }
154
155    pub fn vertex_buffer_layouts() -> [wgpu::VertexBufferLayout<'static>; 4] {
156        [
157            // Slot 0: positions (vec3<f32>)
158            wgpu::VertexBufferLayout {
159                array_stride: std::mem::size_of::<Vec3>() as u64,
160                step_mode: wgpu::VertexStepMode::Vertex,
161                attributes: &[wgpu::VertexAttribute {
162                    format: wgpu::VertexFormat::Float32x3,
163                    offset: 0,
164                    shader_location: 0,
165                }],
166            },
167            // Slot 1: mesh_id (u32)
168            wgpu::VertexBufferLayout {
169                array_stride: std::mem::size_of::<u32>() as u64,
170                step_mode: wgpu::VertexStepMode::Vertex,
171                attributes: &[wgpu::VertexAttribute {
172                    format: wgpu::VertexFormat::Uint32,
173                    offset: 0,
174                    shader_location: 1,
175                }],
176            },
177            // Slot 2: vertex_color (vec4<f32>)
178            wgpu::VertexBufferLayout {
179                array_stride: std::mem::size_of::<Rgba>() as u64,
180                step_mode: wgpu::VertexStepMode::Vertex,
181                attributes: &[wgpu::VertexAttribute {
182                    format: wgpu::VertexFormat::Float32x4,
183                    offset: 0,
184                    shader_location: 2,
185                }],
186            },
187            // Slot 3: vertex_normal (vec3<f32>)
188            wgpu::VertexBufferLayout {
189                array_stride: std::mem::size_of::<Vec3>() as u64,
190                step_mode: wgpu::VertexStepMode::Vertex,
191                attributes: &[wgpu::VertexAttribute {
192                    format: wgpu::VertexFormat::Float32x3,
193                    offset: 0,
194                    shader_location: 3,
195                }],
196            },
197        ]
198    }
199
200    pub fn render_bind_group_layout(ctx: &WgpuContext) -> wgpu::BindGroupLayout {
201        ctx.device
202            .create_bind_group_layout(&wgpu::BindGroupLayoutDescriptor {
203                label: Some("MeshItems Render BGL"),
204                entries: &[
205                    // binding 0: transforms (per-mesh, vertex stage)
206                    bgl_storage_entry(0, wgpu::ShaderStages::VERTEX),
207                    // binding 1: orders (per-mesh, fragment stage, depth bias)
208                    bgl_storage_entry(1, wgpu::ShaderStages::FRAGMENT),
209                ],
210            })
211    }
212
213    fn create_render_bind_group(ctx: &WgpuContext, this: &Self) -> wgpu::BindGroup {
214        ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
215            label: Some("MeshItems Render BG"),
216            layout: &Self::render_bind_group_layout(ctx),
217            entries: &[
218                bg_entry(0, &this.transforms_buffer.buffer),
219                bg_entry(1, &this.orders_buffer.buffer),
220            ],
221        })
222    }
223}
224
225fn bgl_storage_entry(binding: u32, visibility: wgpu::ShaderStages) -> wgpu::BindGroupLayoutEntry {
226    wgpu::BindGroupLayoutEntry {
227        binding,
228        visibility,
229        ty: wgpu::BindingType::Buffer {
230            ty: wgpu::BufferBindingType::Storage { read_only: true },
231            has_dynamic_offset: false,
232            min_binding_size: None,
233        },
234        count: None,
235    }
236}
237
238fn bg_entry(binding: u32, buffer: &wgpu::Buffer) -> wgpu::BindGroupEntry<'_> {
239    wgpu::BindGroupEntry {
240        binding,
241        resource: wgpu::BindingResource::Buffer(buffer.as_entire_buffer_binding()),
242    }
243}
244
245#[cfg(test)]
246mod tests {
247    use super::*;
248    use crate::primitives::test_output_path;
249    use crate::{Renderer, world::RenderFrame};
250    use glam::{Mat4, Vec3};
251    use pollster::block_on;
252    use ranim_core::{components::rgba::Rgba, core_item::CoreItem};
253
254    fn create_triangle_mesh(color: Rgba, offset: Vec3) -> MeshItem {
255        MeshItem {
256            points: vec![
257                Vec3::new(0.0, 1.0, 0.0) + offset,
258                Vec3::new(-1.0, -1.0, 0.0) + offset,
259                Vec3::new(1.0, -1.0, 0.0) + offset,
260            ],
261            triangle_indices: vec![0, 1, 2],
262            transform: Mat4::IDENTITY,
263            vertex_colors: vec![color; 3],
264            vertex_normals: vec![Vec3::ZERO; 3],
265        }
266    }
267
268    fn create_quad_mesh(color: Rgba, offset: Vec3) -> MeshItem {
269        MeshItem {
270            points: vec![
271                Vec3::new(-1.0, 1.0, 0.0) + offset,
272                Vec3::new(1.0, 1.0, 0.0) + offset,
273                Vec3::new(1.0, -1.0, 0.0) + offset,
274                Vec3::new(-1.0, -1.0, 0.0) + offset,
275            ],
276            triangle_indices: vec![0, 1, 2, 0, 2, 3],
277            transform: Mat4::IDENTITY,
278            vertex_colors: vec![color; 4],
279            vertex_normals: vec![Vec3::ZERO; 4],
280        }
281    }
282
283    fn create_sphere_mesh(color: Rgba, radius: f32, position: Vec3) -> MeshItem {
284        let mut points = Vec::new();
285        let mut indices = Vec::new();
286
287        // Simple UV sphere
288        let lat_segments = 20;
289        let lon_segments = 20;
290
291        for lat in 0..=lat_segments {
292            let theta = lat as f32 * std::f32::consts::PI / lat_segments as f32;
293            let sin_theta = theta.sin();
294            let cos_theta = theta.cos();
295
296            for lon in 0..=lon_segments {
297                let phi = lon as f32 * 2.0 * std::f32::consts::PI / lon_segments as f32;
298                let sin_phi = phi.sin();
299                let cos_phi = phi.cos();
300
301                let x = sin_theta * cos_phi;
302                let y = sin_theta * sin_phi;
303                let z = cos_theta;
304
305                points.push(Vec3::new(x * radius, y * radius, z * radius) + position);
306            }
307        }
308
309        for lat in 0..lat_segments {
310            for lon in 0..lon_segments {
311                let first = lat * (lon_segments + 1) + lon;
312                let second = first + lon_segments + 1;
313
314                indices.push(first);
315                indices.push(second);
316                indices.push(first + 1);
317
318                indices.push(second);
319                indices.push(second + 1);
320                indices.push(first + 1);
321            }
322        }
323
324        let vertex_colors = vec![color; points.len()];
325        let vertex_normals = points.iter().map(|p| (*p - position).normalize()).collect();
326
327        MeshItem {
328            points,
329            triangle_indices: indices,
330            transform: Mat4::IDENTITY,
331            vertex_colors,
332            vertex_normals,
333        }
334    }
335
336    #[test]
337    #[ignore = "requires a GPU"]
338    fn render_mesh_items() {
339        use ranim_core::core_item::camera_frame::CameraFrame;
340
341        let ctx = block_on(WgpuContext::new());
342
343        let width = 800u32;
344        let height = 600u32;
345
346        let mut renderer = Renderer::new(&ctx, width, height, 8);
347        let mut render_textures = renderer.new_render_textures(&ctx);
348        let mut store = RenderFrame::new();
349
350        let red = Rgba(glam::Vec4::new(1.0, 0.0, 0.0, 1.0));
351        let green = Rgba(glam::Vec4::new(0.0, 1.0, 0.0, 1.0));
352        let blue = Rgba(glam::Vec4::new(0.0, 0.0, 1.0, 0.8));
353        let yellow = Rgba(glam::Vec4::new(1.0, 1.0, 0.0, 0.9));
354
355        let camera_frame = CameraFrame::default();
356        let triangle1 = create_triangle_mesh(red, Vec3::new(-2.0, 0.0, 0.0));
357        let triangle2 = create_triangle_mesh(green, Vec3::new(2.0, 0.0, 0.0));
358        let quad1 = create_quad_mesh(blue, Vec3::new(0.0, 2.0, 0.0));
359        let quad2 = create_quad_mesh(yellow, Vec3::new(0.0, -2.0, 0.0));
360
361        store.update(
362            [
363                ((0, 0), CoreItem::CameraFrame(camera_frame)),
364                ((1, 0), CoreItem::MeshItem(triangle1)),
365                ((1, 1), CoreItem::MeshItem(triangle2)),
366                ((2, 0), CoreItem::MeshItem(quad1)),
367                ((3, 1), CoreItem::MeshItem(quad2)),
368            ]
369            .into_iter(),
370        );
371
372        let clear_color = wgpu::Color {
373            r: 0.1,
374            g: 0.1,
375            b: 0.1,
376            a: 1.0,
377        };
378
379        renderer.render_frame(&mut render_textures, clear_color, &store);
380
381        ctx.device
382            .poll(wgpu::PollType::wait_indefinitely())
383            .unwrap();
384
385        let buffer = render_textures.get_rendered_texture_img_buffer(&ctx);
386
387        let output_path = test_output_path("mesh_items_render.png");
388        buffer.save(&output_path).expect("Failed to save image");
389
390        println!("Rendered image saved to: {:?}", output_path);
391        println!("Open it to see the mesh rendering result!");
392
393        assert!(output_path.exists(), "Image file should be created");
394    }
395
396    #[test]
397    #[ignore = "requires a GPU"]
398    fn test_nested_transparent_spheres() {
399        use ranim_core::core_item::camera_frame::CameraFrame;
400
401        let ctx = block_on(WgpuContext::new());
402        let width = 800u32;
403        let height = 600u32;
404
405        let mut renderer = Renderer::new(&ctx, width, height, 8);
406        let mut render_textures = renderer.new_render_textures(&ctx);
407        let mut store = RenderFrame::new();
408
409        // Create nested spheres:
410        // 1. Outer transparent sphere (blue, alpha=0.3, radius=2.0)
411        // 2. Middle opaque sphere (red, alpha=1.0, radius=1.5)
412        // 3. Inner transparent sphere (green, alpha=0.5, radius=1.0)
413
414        let outer_transparent = Rgba(glam::Vec4::new(0.0, 0.0, 1.0, 0.3));
415        let middle_opaque = Rgba(glam::Vec4::new(1.0, 0.0, 0.0, 1.0));
416        let inner_transparent = Rgba(glam::Vec4::new(0.0, 1.0, 0.0, 0.5));
417
418        let outer_sphere = create_sphere_mesh(outer_transparent, 2.0, Vec3::ZERO);
419        let middle_sphere = create_sphere_mesh(middle_opaque, 1.5, Vec3::ZERO);
420        let inner_sphere = create_sphere_mesh(inner_transparent, 1.0, Vec3::ZERO);
421
422        let camera_frame = CameraFrame::default();
423
424        store.update(
425            [
426                ((0, 0), CoreItem::CameraFrame(camera_frame)),
427                ((1, 0), CoreItem::MeshItem(outer_sphere)),
428                ((2, 0), CoreItem::MeshItem(middle_sphere)),
429                ((3, 0), CoreItem::MeshItem(inner_sphere)),
430            ]
431            .into_iter(),
432        );
433
434        let clear_color = wgpu::Color {
435            r: 0.1,
436            g: 0.1,
437            b: 0.1,
438            a: 1.0,
439        };
440
441        renderer.render_frame(&mut render_textures, clear_color, &store);
442
443        ctx.device
444            .poll(wgpu::PollType::wait_indefinitely())
445            .unwrap();
446
447        // Analyze depth buffer
448        let depth_data = render_textures.get_depth_texture_data(&ctx);
449        let mut min_depth = f32::MAX;
450        let mut max_depth = f32::MIN;
451        let mut depth_histogram: std::collections::HashMap<u32, usize> =
452            std::collections::HashMap::new();
453
454        for &d in depth_data {
455            if (d - 1.0).abs() > 0.001 {
456                min_depth = min_depth.min(d);
457                max_depth = max_depth.max(d);
458                let bucket = (d * 10000.0) as u32;
459                *depth_histogram.entry(bucket).or_insert(0) += 1;
460            }
461        }
462
463        println!("\n=== Nested Spheres Depth Test ===");
464        println!("Depth buffer analysis:");
465        println!("  Min depth: {}", min_depth);
466        println!("  Max depth: {}", max_depth);
467        println!("\nDepth histogram (top 10 buckets):");
468        let mut buckets: Vec<_> = depth_histogram.iter().collect();
469        buckets.sort_by_key(|(k, _)| *k);
470        for (bucket, count) in buckets.iter().take(10) {
471            println!(
472                "    depth ~{:.4}: {} pixels",
473                **bucket as f32 / 10000.0,
474                count
475            );
476        }
477
478        let buffer = render_textures.get_rendered_texture_img_buffer(&ctx);
479
480        // Sample some pixels to see actual colors
481        println!("\nColor samples (center region):");
482        let center_x = width / 2;
483        let center_y = height / 2;
484        for dy in [-50, 0, 50].iter() {
485            for dx in [-50, 0, 50].iter() {
486                let x = (center_x as i32 + dx) as u32;
487                let y = (center_y as i32 + dy) as u32;
488                if x < width && y < height {
489                    let pixel = buffer.get_pixel(x, y);
490                    println!(
491                        "  ({:3}, {:3}): R={:3} G={:3} B={:3} A={:3}",
492                        dx, dy, pixel[0], pixel[1], pixel[2], pixel[3]
493                    );
494                }
495            }
496        }
497
498        let buffer = render_textures.get_rendered_texture_img_buffer(&ctx);
499        let output_path = test_output_path("nested_spheres_render.png");
500        buffer.save(&output_path).expect("Failed to save image");
501
502        let depth_buffer = render_textures.get_depth_texture_img_buffer(&ctx);
503        let depth_path = test_output_path("nested_spheres_depth.png");
504        depth_buffer
505            .save(&depth_path)
506            .expect("Failed to save depth image");
507
508        println!("\nImages saved to output/");
509        println!("\nExpected behavior:");
510        println!("  - Outer transparent blue sphere should be visible");
511        println!("  - Middle opaque red sphere should occlude inner green sphere");
512        println!("  - Inner green sphere should NOT be visible from outside");
513        println!("  - Depth buffer should show opaque red sphere's depth");
514
515        assert!(output_path.exists(), "Image file should be created");
516        assert!(depth_path.exists(), "Depth image file should be created");
517    }
518}