Skip to main content

spatialrust_render_wgpu/
runtime.rs

1use std::sync::{
2    atomic::{AtomicU64, Ordering},
3    Arc, Mutex,
4};
5
6use bytemuck::{Pod, Zeroable};
7use spatialrust_gpu::WgpuRuntime;
8use spatialrust_viz::{
9    DeviceIdentity, PointCloudView, TransferDirection, TransferEvent, TransferReceipt,
10    TriangleMeshView, VisualPrimitive, VisualResidency,
11};
12
13use crate::render::RenderPipelines;
14use crate::{GpuGeometry, GpuGeometryKind, RenderError, RenderResult};
15
16const MAX_CACHED_BUFFERS: usize = 32;
17static NEXT_RENDERER_ID: AtomicU64 = AtomicU64::new(1);
18
19#[repr(C)]
20#[derive(Clone, Copy, Debug, Pod, Zeroable)]
21struct PositionVertex {
22    position: [f32; 3],
23}
24
25#[repr(C)]
26#[derive(Clone, Copy, Debug, Pod, Zeroable)]
27struct Rgba8Vertex {
28    color: [u8; 4],
29}
30
31pub(crate) struct BufferSlot {
32    pub(crate) buffer: wgpu::Buffer,
33    pub(crate) capacity: u64,
34    pub(crate) logical_bytes: u64,
35    usage: wgpu::BufferUsages,
36}
37
38#[derive(Default)]
39struct RenderBufferPool {
40    buffers: Vec<BufferSlot>,
41}
42
43impl RenderBufferPool {
44    fn acquire(
45        &mut self,
46        device: &wgpu::Device,
47        logical_bytes: u64,
48        usage: wgpu::BufferUsages,
49        label: &'static str,
50    ) -> BufferSlot {
51        let required_capacity = logical_bytes.max(wgpu::COPY_BUFFER_ALIGNMENT);
52        if let Some(index) = self
53            .buffers
54            .iter()
55            .enumerate()
56            .filter(|(_, slot)| slot.usage == usage && slot.capacity >= required_capacity)
57            .min_by_key(|(_, slot)| slot.capacity)
58            .map(|(index, _)| index)
59        {
60            let mut slot = self.buffers.swap_remove(index);
61            slot.logical_bytes = logical_bytes;
62            return slot;
63        }
64        BufferSlot {
65            buffer: device.create_buffer(&wgpu::BufferDescriptor {
66                label: Some(label),
67                size: required_capacity,
68                usage,
69                mapped_at_creation: false,
70            }),
71            capacity: required_capacity,
72            logical_bytes,
73            usage,
74        }
75    }
76
77    fn recycle(&mut self, slot: BufferSlot) {
78        if self.buffers.len() < MAX_CACHED_BUFFERS {
79            self.buffers.push(slot);
80        } else {
81            slot.buffer.destroy();
82        }
83    }
84}
85
86/// Renderer runtime sharing one explicit `spatialrust-gpu` wgpu device.
87pub struct WgpuRenderer {
88    pub(crate) id: u64,
89    pub(crate) runtime: Arc<WgpuRuntime>,
90    device_identity: DeviceIdentity,
91    buffer_pool: Mutex<RenderBufferPool>,
92    pub(crate) render_pipelines: Mutex<RenderPipelines>,
93}
94
95impl WgpuRenderer {
96    /// Creates a renderer on a caller-selected wgpu runtime.
97    #[must_use]
98    pub fn new(runtime: Arc<WgpuRuntime>) -> Self {
99        let id = NEXT_RENDERER_ID.fetch_add(1, Ordering::Relaxed);
100        let adapter = runtime.adapter_info();
101        let backend = format!("wgpu-{}", adapter.backend);
102        let device = format!("{}#renderer-{id}", adapter.name);
103        let device_identity = DeviceIdentity::try_new(backend, device)
104            .expect("wgpu adapter identity and renderer id are non-empty");
105        Self {
106            id,
107            runtime,
108            device_identity,
109            buffer_pool: Mutex::new(RenderBufferPool::default()),
110            render_pipelines: Mutex::new(RenderPipelines::default()),
111        }
112    }
113
114    /// Identity of this exact renderer runtime.
115    #[must_use]
116    pub const fn device_identity(&self) -> &DeviceIdentity {
117        &self.device_identity
118    }
119
120    /// Number of buffers currently retained for reuse.
121    #[must_use]
122    pub fn cached_buffer_count(&self) -> usize {
123        self.buffer_pool.lock().expect("render buffer pool poisoned").buffers.len()
124    }
125
126    /// Explicitly uploads one borrowed visual primitive.
127    ///
128    /// Structure-of-arrays positions are packed into a vertex buffer as part of
129    /// this named operation. The returned receipt records every host/device
130    /// crossing and its exact logical byte count.
131    pub fn upload(
132        &self,
133        primitive: VisualPrimitive<'_>,
134    ) -> RenderResult<(GpuGeometry, TransferReceipt)> {
135        match primitive {
136            VisualPrimitive::Points(points) => self.upload_points(points),
137            VisualPrimitive::Lines(lines) => {
138                let vertices = pack_interleaved_positions(lines.positions_xyz)?;
139                self.upload_packed(GpuGeometryKind::Lines, vertices, None, None, None, "line")
140            }
141            VisualPrimitive::Triangles(mesh) => self.upload_triangles(mesh),
142        }
143    }
144
145    fn upload_points(
146        &self,
147        points: PointCloudView<'_>,
148    ) -> RenderResult<(GpuGeometry, TransferReceipt)> {
149        let vertex_count = checked_u32(points.positions.len(), "point count")?;
150        let mut positions = Vec::with_capacity(points.positions.len());
151        for index in 0..points.positions.len() {
152            positions.push(PositionVertex {
153                position: [
154                    points.positions.x[index],
155                    points.positions.y[index],
156                    points.positions.z[index],
157                ],
158            });
159        }
160        let rgb = points.rgb.map(|columns| {
161            (0..points.positions.len())
162                .map(|index| Rgba8Vertex {
163                    color: [columns.red[index], columns.green[index], columns.blue[index], u8::MAX],
164                })
165                .collect::<Vec<_>>()
166        });
167        let scalar = points.scalar.map(|column| column.values.to_vec());
168        self.upload_buffers(
169            GpuGeometryKind::Points,
170            vertex_count,
171            0,
172            &positions,
173            rgb.as_deref(),
174            scalar.as_deref(),
175            None,
176            "point",
177        )
178    }
179
180    fn upload_triangles(
181        &self,
182        mesh: TriangleMeshView<'_>,
183    ) -> RenderResult<(GpuGeometry, TransferReceipt)> {
184        let vertices = pack_interleaved_positions(mesh.positions_xyz)?;
185        let vertex_count = checked_u32(vertices.len(), "mesh vertex count")?;
186        let index_count = checked_u32(mesh.indices.len(), "mesh index count")?;
187        self.upload_buffers(
188            GpuGeometryKind::Triangles,
189            vertex_count,
190            index_count,
191            &vertices,
192            None,
193            None,
194            Some(mesh.indices),
195            "triangle",
196        )
197    }
198
199    fn upload_packed(
200        &self,
201        kind: GpuGeometryKind,
202        vertices: Vec<PositionVertex>,
203        rgb: Option<&[Rgba8Vertex]>,
204        scalar: Option<&[f32]>,
205        indices: Option<&[u32]>,
206        prefix: &'static str,
207    ) -> RenderResult<(GpuGeometry, TransferReceipt)> {
208        let vertex_count = checked_u32(vertices.len(), "vertex count")?;
209        let index_count = checked_u32(indices.map_or(0, <[u32]>::len), "index count")?;
210        self.upload_buffers(
211            kind,
212            vertex_count,
213            index_count,
214            &vertices,
215            rgb,
216            scalar,
217            indices,
218            prefix,
219        )
220    }
221
222    #[allow(clippy::too_many_arguments)]
223    fn upload_buffers(
224        &self,
225        kind: GpuGeometryKind,
226        vertex_count: u32,
227        index_count: u32,
228        positions: &[PositionVertex],
229        rgb: Option<&[Rgba8Vertex]>,
230        scalar: Option<&[f32]>,
231        indices: Option<&[u32]>,
232        prefix: &'static str,
233    ) -> RenderResult<(GpuGeometry, TransferReceipt)> {
234        let mut receipt = TransferReceipt::new();
235        let position_slot = self.upload_slice(
236            positions,
237            wgpu::BufferUsages::VERTEX | wgpu::BufferUsages::COPY_DST,
238            "spatialrust render positions",
239            &format!("{prefix}-positions-upload"),
240            &mut receipt,
241        )?;
242        let rgb_slot = rgb
243            .map(|values| {
244                self.upload_slice(
245                    values,
246                    wgpu::BufferUsages::VERTEX | wgpu::BufferUsages::COPY_DST,
247                    "spatialrust render rgb",
248                    &format!("{prefix}-rgb-upload"),
249                    &mut receipt,
250                )
251            })
252            .transpose()?;
253        let scalar_slot = scalar
254            .map(|values| {
255                self.upload_slice(
256                    values,
257                    wgpu::BufferUsages::VERTEX | wgpu::BufferUsages::COPY_DST,
258                    "spatialrust render scalar",
259                    &format!("{prefix}-scalar-upload"),
260                    &mut receipt,
261                )
262            })
263            .transpose()?;
264        let index_slot = indices
265            .map(|values| {
266                self.upload_slice(
267                    values,
268                    wgpu::BufferUsages::INDEX | wgpu::BufferUsages::COPY_DST,
269                    "spatialrust render indices",
270                    &format!("{prefix}-indices-upload"),
271                    &mut receipt,
272                )
273            })
274            .transpose()?;
275        Ok((
276            GpuGeometry::new(
277                self.id,
278                kind,
279                vertex_count,
280                index_count,
281                Some(position_slot),
282                rgb_slot,
283                scalar_slot,
284                index_slot,
285                self.device_identity.clone(),
286            ),
287            receipt,
288        ))
289    }
290
291    fn upload_slice<T: Pod>(
292        &self,
293        values: &[T],
294        usage: wgpu::BufferUsages,
295        label: &'static str,
296        stage: &str,
297        receipt: &mut TransferReceipt,
298    ) -> RenderResult<BufferSlot> {
299        let bytes = bytemuck::cast_slice(values);
300        let logical_bytes = u64::try_from(bytes.len())
301            .map_err(|_| RenderError::GeometrySize("buffer byte length exceeds u64".into()))?;
302        let mut slot = self.buffer_pool.lock().expect("render buffer pool poisoned").acquire(
303            self.runtime.device(),
304            logical_bytes,
305            usage,
306            label,
307        );
308        if !bytes.is_empty() {
309            self.runtime.queue().write_buffer(&slot.buffer, 0, bytes);
310        }
311        slot.logical_bytes = logical_bytes;
312        receipt.push(
313            TransferEvent::try_new(
314                stage,
315                TransferDirection::Upload,
316                VisualResidency::Host,
317                VisualResidency::Device(self.device_identity.clone()),
318                logical_bytes,
319            )
320            .map_err(|error| RenderError::Transfer(error.to_string()))?,
321        );
322        Ok(slot)
323    }
324
325    pub(crate) fn recycle_geometry(&self, geometry: &mut GpuGeometry) -> RenderResult<()> {
326        if geometry.renderer_id != self.id {
327            return Err(RenderError::RuntimeMismatch(
328                "geometry was uploaded by another renderer".into(),
329            ));
330        }
331        let mut pool = self.buffer_pool.lock().expect("render buffer pool poisoned");
332        for slot in [
333            geometry.positions.take(),
334            geometry.rgb.take(),
335            geometry.scalar.take(),
336            geometry.indices.take(),
337        ]
338        .into_iter()
339        .flatten()
340        {
341            pool.recycle(slot);
342        }
343        Ok(())
344    }
345}
346
347fn pack_interleaved_positions(values: &[f32]) -> RenderResult<Vec<PositionVertex>> {
348    if values.len() % 3 != 0 {
349        return Err(RenderError::GeometrySize(
350            "interleaved XYZ data must contain complete triples".into(),
351        ));
352    }
353    Ok(values
354        .chunks_exact(3)
355        .map(|xyz| PositionVertex { position: [xyz[0], xyz[1], xyz[2]] })
356        .collect())
357}
358
359fn checked_u32(value: usize, label: &str) -> RenderResult<u32> {
360    u32::try_from(value)
361        .map_err(|_| RenderError::GeometrySize(format!("{label} exceeds the u32 render limit")))
362}
363
364#[cfg(test)]
365mod tests {
366    use std::sync::Arc;
367
368    use spatialrust_gpu::WgpuRuntime;
369    use spatialrust_viz::{
370        PointCloudView, PositionColumns3, Rgb8Columns, ScalarColumn, TriangleMeshView,
371        VisualPrimitive, VisualResidency,
372    };
373
374    use super::{pack_interleaved_positions, WgpuRenderer};
375
376    #[test]
377    fn packs_interleaved_positions_without_reordering() {
378        let packed = pack_interleaved_positions(&[1.0, 2.0, 3.0, 4.0, 5.0, 6.0]).unwrap();
379        assert_eq!(packed[0].position, [1.0, 2.0, 3.0]);
380        assert_eq!(packed[1].position, [4.0, 5.0, 6.0]);
381        assert!(pack_interleaved_positions(&[1.0, 2.0]).is_err());
382    }
383
384    #[test]
385    fn explicit_upload_receipt_and_recycling() {
386        let Ok(runtime) = WgpuRuntime::new_headless() else {
387            eprintln!("skipping GPU upload test: no headless adapter");
388            return;
389        };
390        let renderer = WgpuRenderer::new(Arc::new(runtime));
391        let positions = PositionColumns3::try_new(&[1.0, 2.0], &[3.0, 4.0], &[5.0, 6.0]).unwrap();
392        let rgb = Rgb8Columns::try_new(&[1, 2], &[3, 4], &[5, 6], 2).unwrap();
393        let scalar = ScalarColumn::try_new("intensity", &[0.25, 0.75], 2).unwrap();
394        let points = PointCloudView::positions_only(positions)
395            .with_rgb(rgb)
396            .unwrap()
397            .with_scalar(scalar)
398            .unwrap();
399
400        let (geometry, receipt) = renderer.upload(VisualPrimitive::Points(points)).unwrap();
401        assert_eq!(geometry.vertex_count(), 2);
402        assert_eq!(geometry.position_bytes(), 24);
403        assert_eq!(geometry.rgb_bytes(), 8);
404        assert_eq!(geometry.scalar_bytes(), 8);
405        assert_eq!(receipt.total_bytes().unwrap(), 40);
406        assert_eq!(receipt.events().len(), 3);
407        assert_eq!(receipt.events()[0].stage, "point-positions-upload");
408        assert_eq!(
409            geometry.residency(),
410            &VisualResidency::Device(renderer.device_identity().clone())
411        );
412
413        geometry.recycle(&renderer).unwrap();
414        assert_eq!(renderer.cached_buffer_count(), 3);
415
416        let (reused, _) = renderer.upload(VisualPrimitive::Points(points)).unwrap();
417        assert_eq!(renderer.cached_buffer_count(), 0);
418        reused.recycle(&renderer).unwrap();
419
420        let mesh =
421            TriangleMeshView::try_new(&[0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0, 0.0], &[0, 1, 2])
422                .unwrap();
423        let (triangles, triangle_receipt) =
424            renderer.upload(VisualPrimitive::Triangles(mesh)).unwrap();
425        assert_eq!(triangles.vertex_count(), 3);
426        assert_eq!(triangles.index_count(), 3);
427        assert_eq!(triangles.position_bytes(), 36);
428        assert_eq!(triangles.index_bytes(), 12);
429        assert_eq!(triangle_receipt.total_bytes().unwrap(), 48);
430        assert_eq!(triangle_receipt.events().len(), 2);
431        triangles.recycle(&renderer).unwrap();
432    }
433
434    #[test]
435    fn recycle_rejects_another_renderer() {
436        let Ok(runtime) = WgpuRuntime::new_headless() else {
437            eprintln!("skipping GPU runtime identity test: no headless adapter");
438            return;
439        };
440        let runtime = Arc::new(runtime);
441        let first = WgpuRenderer::new(Arc::clone(&runtime));
442        let second = WgpuRenderer::new(runtime);
443        let positions = PositionColumns3::try_new(&[0.0], &[0.0], &[0.0]).unwrap();
444        let points = PointCloudView::positions_only(positions);
445        let (geometry, _) = first.upload(VisualPrimitive::Points(points)).unwrap();
446        assert!(geometry.recycle(&second).is_err());
447    }
448}