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