use std::{borrow::Cow, num::NonZeroU64};
use bevy::{
prelude::*,
render::{
Extract, ExtractSchedule, Render, RenderApp, RenderStartup, RenderSystems,
graph::CameraDriverLabel,
mesh::{RenderMesh, allocator::MeshAllocator},
render_asset::RenderAssets,
render_graph::{Node, NodeRunError, RenderGraph, RenderGraphContext, RenderLabel},
render_resource::{
BindGroup, BindGroupEntry, BindGroupLayout, BindGroupLayoutEntry, BindingType, Buffer,
BufferBinding, BufferBindingType, BufferDescriptor, BufferInitDescriptor, BufferUsages,
ComputePassDescriptor, ComputePipeline, PipelineLayoutDescriptor,
RawComputePipelineDescriptor, ShaderModuleDescriptor, ShaderSource, ShaderStages,
VertexAttribute,
},
renderer::{RenderContext, RenderDevice, RenderQueue},
sync_world::RenderEntity,
},
};
use finite_light_gpu_common::FrameUniforms;
use finite_light_math::PoincareTransform;
use crate::{PendingMeshData, Relativistic, RelativisticChild, RelativisticMetric};
pub(crate) struct RelativisticRenderPlugin;
impl Plugin for RelativisticRenderPlugin {
fn build(&self, app: &mut App) {
let Some(render_app) = app.get_sub_app_mut(RenderApp) else {
panic!("RelativisticRenderPlugin: no RenderApp found!");
};
render_app
.init_resource::<ExtractedCamera>()
.init_resource::<PreparedDispatches>()
.add_systems(
RenderStartup,
(enable_storage_on_vertex_buffers, init_pipeline),
)
.add_systems(ExtractSchedule, extract_relativistic)
.add_systems(
Render,
(
prepare_relativistic.in_set(RenderSystems::PrepareResources),
prepare_dispatches
.in_set(RenderSystems::Prepare)
.after(RenderSystems::PrepareBindGroups),
),
);
let mut render_graph = render_app.world_mut().resource_mut::<RenderGraph>();
render_graph.add_node(RelativisticComputeLabel, RelativisticComputeNode);
render_graph.add_node_edge(RelativisticComputeLabel, CameraDriverLabel);
}
}
fn enable_storage_on_vertex_buffers(mut allocator: ResMut<MeshAllocator>) {
allocator.extra_buffer_usages |= BufferUsages::STORAGE;
}
#[derive(Resource, Default)]
struct ExtractedCamera(Option<PoincareTransform>);
#[derive(Resource)]
struct RenderMetric(finite_light_math::Metric);
#[derive(Component)]
struct ExtractedWorldLine {
keyframes: Vec<PoincareTransform>,
child: Option<(finite_light_math::Vec3, Quat)>,
}
#[derive(Component)]
struct ExtractedMeshId(AssetId<Mesh>);
#[derive(Component)]
struct ExtractedMeshData {
vertices: Vec<finite_light_math::Vec4>,
normals: Vec<finite_light_math::Vec4>,
}
#[derive(Component)]
struct RenderBuffers {
vertices_local: Buffer,
normals_local: Buffer,
keyframes: Buffer,
keyframe_capacity: u32,
vertex_count: u32,
}
#[derive(Resource)]
struct RelativisticPipeline {
pipeline: ComputePipeline,
uniform_buffer: Buffer,
uniform_bind_group: BindGroup,
uniform_layout: BindGroupLayout,
uniform_capacity: usize,
aligned_stride: u64,
storage_layout: BindGroupLayout,
}
#[derive(Debug, Hash, PartialEq, Eq, Clone, RenderLabel)]
struct RelativisticComputeLabel;
struct PreparedDispatch {
storage_bind_group: BindGroup,
uniform_offset: u32,
workgroups: u32,
}
#[derive(Resource, Default)]
struct PreparedDispatches {
dispatches: Vec<PreparedDispatch>,
}
struct RelativisticComputeNode;
impl Node for RelativisticComputeNode {
fn run<'w>(
&self,
_graph: &mut RenderGraphContext,
render_context: &mut RenderContext<'w>,
world: &'w World,
) -> Result<(), NodeRunError> {
let Some(pipeline) = world.get_resource::<RelativisticPipeline>() else {
return Ok(());
};
let prepared = world.resource::<PreparedDispatches>();
if prepared.dispatches.is_empty() {
return Ok(());
}
let encoder = render_context.command_encoder();
let mut pass = encoder.begin_compute_pass(&ComputePassDescriptor {
label: Some("retarded_vertices"),
..default()
});
pass.set_pipeline(&pipeline.pipeline);
for dispatch in &prepared.dispatches {
pass.set_bind_group(0, &pipeline.uniform_bind_group, &[dispatch.uniform_offset]);
pass.set_bind_group(1, &dispatch.storage_bind_group, &[]);
pass.dispatch_workgroups(dispatch.workgroups, 1, 1);
}
Ok(())
}
}
fn extract_relativistic(
mut commands: Commands,
mut camera: ResMut<ExtractedCamera>,
render_buffers: Query<(), With<RenderBuffers>>,
mut extracted_world_lines: Query<&mut ExtractedWorldLine>,
camera_query: Extract<Query<&Relativistic, With<Camera3d>>>,
entity_query: Extract<
Query<(
RenderEntity,
&Relativistic,
&Mesh3d,
Option<&PendingMeshData>,
Option<&RelativisticChild>,
)>,
>,
source_query: Extract<Query<&Relativistic>>,
metric: Extract<Res<RelativisticMetric>>,
) {
let camera_relativistic = camera_query
.single()
.expect("exactly one Camera3d with Relativistic required");
camera.0 = Some(*camera_relativistic.transform());
commands.insert_resource(RenderMetric(metric.0));
for (render_entity, relativistic, mesh_handle, pending, child) in &entity_query {
let source_keyframes = if let Some(child) = child {
let Ok(source) = source_query.get(child.source) else {
continue;
};
source.world_line().keyframes()
} else {
relativistic.world_line().keyframes()
};
let child_data = child.map(|c| (c.offset, c.rotation));
if let Ok(mut extracted) = extracted_world_lines.get_mut(render_entity) {
extracted.keyframes.clear();
extracted.keyframes.extend(source_keyframes.iter().copied());
extracted.child = child_data;
} else {
commands.entity(render_entity).insert((
ExtractedWorldLine {
keyframes: source_keyframes.iter().copied().collect(),
child: child_data,
},
ExtractedMeshId(mesh_handle.id()),
));
}
if !render_buffers.contains(render_entity)
&& let Some(data) = pending
{
commands.entity(render_entity).insert(ExtractedMeshData {
vertices: data.vertices.clone(),
normals: data.normals.clone(),
});
}
}
}
fn prepare_relativistic(
mut commands: Commands,
render_device: Res<RenderDevice>,
render_queue: Res<RenderQueue>,
new_entities: Query<(Entity, &ExtractedMeshData, &ExtractedWorldLine), Without<RenderBuffers>>,
mut existing_entities: Query<(&ExtractedWorldLine, &mut RenderBuffers)>,
) {
for (entity, mesh_data, world_line) in &new_entities {
let vertex_count = mesh_data.vertices.len() as u32;
let vertices_local = render_device.create_buffer_with_data(&BufferInitDescriptor {
label: Some("vertices_local"),
contents: bytemuck::cast_slice(&mesh_data.vertices),
usage: BufferUsages::STORAGE,
});
let normals_local = render_device.create_buffer_with_data(&BufferInitDescriptor {
label: Some("normals_local"),
contents: bytemuck::cast_slice(&mesh_data.normals),
usage: BufferUsages::STORAGE,
});
let keyframe_capacity = (world_line.keyframes.len() as u32).next_power_of_two();
let keyframes_buffer = create_keyframe_buffer(&render_device, keyframe_capacity);
render_queue.write_buffer(
&keyframes_buffer,
0,
bytemuck::cast_slice(&world_line.keyframes),
);
commands.entity(entity).insert(RenderBuffers {
vertices_local,
normals_local,
keyframes: keyframes_buffer,
keyframe_capacity,
vertex_count,
});
commands.entity(entity).remove::<ExtractedMeshData>();
}
for (world_line, mut buffers) in &mut existing_entities {
if world_line.keyframes.is_empty() {
continue;
}
let needed = world_line.keyframes.len() as u32;
if needed > buffers.keyframe_capacity {
let new_capacity = needed.next_power_of_two();
buffers.keyframes = create_keyframe_buffer(&render_device, new_capacity);
buffers.keyframe_capacity = new_capacity;
}
render_queue.write_buffer(
&buffers.keyframes,
0,
bytemuck::cast_slice(&world_line.keyframes),
);
}
}
fn create_keyframe_buffer(render_device: &RenderDevice, capacity: u32) -> Buffer {
render_device.create_buffer(&BufferDescriptor {
label: Some("keyframes"),
size: capacity as u64 * std::mem::size_of::<PoincareTransform>() as u64,
usage: BufferUsages::STORAGE | BufferUsages::COPY_DST,
mapped_at_creation: false,
})
}
const UNIFORM_SIZE: u64 = std::mem::size_of::<FrameUniforms>() as u64;
const INITIAL_UNIFORM_CAPACITY: usize = 16;
fn init_pipeline(mut commands: Commands, render_device: Res<RenderDevice>) {
let module = naga::front::spv::parse_u8_slice(
finite_light_gpu::SPIRV,
&naga::front::spv::Options::default(),
)
.expect("failed to parse SPIR-V");
let info = naga::valid::Validator::new(
naga::valid::ValidationFlags::all(),
naga::valid::Capabilities::all(),
)
.validate(&module)
.expect("shader validation failed");
let wgsl =
naga::back::wgsl::write_string(&module, &info, naga::back::wgsl::WriterFlags::empty())
.expect("failed to convert to WGSL");
let shader = render_device.create_and_validate_shader_module(ShaderModuleDescriptor {
label: Some("retarded_vertices"),
source: ShaderSource::Wgsl(Cow::Owned(wgsl)),
});
let uniform_layout = render_device.create_bind_group_layout(
"retarded_vertices_uniforms",
&[BindGroupLayoutEntry {
binding: 0,
visibility: ShaderStages::COMPUTE,
ty: BindingType::Buffer {
ty: BufferBindingType::Uniform,
has_dynamic_offset: true,
min_binding_size: NonZeroU64::new(UNIFORM_SIZE),
},
count: None,
}],
);
let storage_layout = render_device.create_bind_group_layout(
"retarded_vertices_storage",
&[
BindGroupLayoutEntry {
binding: 0,
visibility: ShaderStages::COMPUTE,
ty: BindingType::Buffer {
ty: BufferBindingType::Storage { read_only: true },
has_dynamic_offset: false,
min_binding_size: None,
},
count: None,
},
BindGroupLayoutEntry {
binding: 1,
visibility: ShaderStages::COMPUTE,
ty: BindingType::Buffer {
ty: BufferBindingType::Storage { read_only: true },
has_dynamic_offset: false,
min_binding_size: None,
},
count: None,
},
BindGroupLayoutEntry {
binding: 2,
visibility: ShaderStages::COMPUTE,
ty: BindingType::Buffer {
ty: BufferBindingType::Storage { read_only: true },
has_dynamic_offset: false,
min_binding_size: None,
},
count: None,
},
BindGroupLayoutEntry {
binding: 3,
visibility: ShaderStages::COMPUTE,
ty: BindingType::Buffer {
ty: BufferBindingType::Storage { read_only: false },
has_dynamic_offset: false,
min_binding_size: None,
},
count: None,
},
],
);
let pipeline_layout = render_device.create_pipeline_layout(&PipelineLayoutDescriptor {
label: Some("retarded_vertices"),
bind_group_layouts: &[&uniform_layout, &storage_layout],
push_constant_ranges: &[],
});
let pipeline = render_device.create_compute_pipeline(&RawComputePipelineDescriptor {
label: Some("retarded_vertices"),
layout: Some(&pipeline_layout),
module: &shader,
entry_point: Some("retarded_vertices"),
compilation_options: Default::default(),
cache: None,
});
let alignment = render_device.limits().min_uniform_buffer_offset_alignment as u64;
let aligned_stride = UNIFORM_SIZE.next_multiple_of(alignment);
let uniform_buffer = render_device.create_buffer(&BufferDescriptor {
label: Some("retarded_vertices_uniforms"),
size: INITIAL_UNIFORM_CAPACITY as u64 * aligned_stride,
usage: BufferUsages::UNIFORM | BufferUsages::COPY_DST,
mapped_at_creation: false,
});
let uniform_bind_group =
create_uniform_bind_group(&render_device, &uniform_layout, &uniform_buffer);
commands.insert_resource(RelativisticPipeline {
pipeline,
uniform_buffer,
uniform_bind_group,
uniform_layout,
uniform_capacity: INITIAL_UNIFORM_CAPACITY,
aligned_stride,
storage_layout,
});
}
fn create_uniform_bind_group(
render_device: &RenderDevice,
layout: &BindGroupLayout,
buffer: &Buffer,
) -> BindGroup {
render_device.create_bind_group(
"retarded_vertices_uniforms",
layout,
&[BindGroupEntry {
binding: 0,
resource: bevy::render::render_resource::BindingResource::Buffer(BufferBinding {
buffer,
offset: 0,
size: NonZeroU64::new(UNIFORM_SIZE),
}),
}],
)
}
const BYTES_PER_F32: u64 = std::mem::size_of::<f32>() as u64;
fn prepare_dispatches(
render_device: Res<RenderDevice>,
render_queue: Res<RenderQueue>,
mut pipeline: ResMut<RelativisticPipeline>,
camera: Res<ExtractedCamera>,
metric: Res<RenderMetric>,
mesh_allocator: Res<MeshAllocator>,
render_meshes: Res<RenderAssets<RenderMesh>>,
query: Query<(&ExtractedMeshId, &ExtractedWorldLine, &RenderBuffers)>,
mut prepared: ResMut<PreparedDispatches>,
) {
prepared.dispatches.clear();
let Some(camera_poincare) = camera.0 else {
return;
};
let entity_count = query
.iter()
.filter(|(_, wl, _)| wl.keyframes.len() >= 2)
.count();
if entity_count == 0 {
return;
}
let aligned_stride = pipeline.aligned_stride;
if entity_count > pipeline.uniform_capacity {
let new_capacity = entity_count.next_power_of_two();
pipeline.uniform_buffer = render_device.create_buffer(&BufferDescriptor {
label: Some("retarded_vertices_uniforms"),
size: new_capacity as u64 * aligned_stride,
usage: BufferUsages::UNIFORM | BufferUsages::COPY_DST,
mapped_at_creation: false,
});
pipeline.uniform_bind_group = create_uniform_bind_group(
&render_device,
&pipeline.uniform_layout,
&pipeline.uniform_buffer,
);
pipeline.uniform_capacity = new_capacity;
}
let mut dispatch_index = 0u32;
for (mesh_id, world_line, buffers) in &query {
if world_line.keyframes.len() < 2 {
continue;
}
let Some(render_mesh) = render_meshes.get(mesh_id.0) else {
continue;
};
let Some(vertex_slice) = mesh_allocator.mesh_vertex_slice(&mesh_id.0) else {
continue;
};
let vertex_layout = render_mesh.layout.0.layout();
let stride_f32 = (vertex_layout.array_stride / BYTES_PER_F32) as u32;
let uniforms = FrameUniforms {
camera_transform: camera_poincare,
camera: camera_poincare.translation,
metric: metric.0,
num_vertices: buffers.vertex_count,
num_keyframes: world_line.keyframes.len() as u32,
vertex_base: vertex_slice.range.start * stride_f32,
vertex_stride: stride_f32,
position_offset: attribute_f32_offset(&vertex_layout.attributes, 0),
normal_offset: attribute_f32_offset(&vertex_layout.attributes, 1),
child_offset: world_line
.child
.map_or(finite_light_math::Vec3::ZERO, |c| c.0),
child_rotation: world_line.child.map_or(Quat::IDENTITY, |c| c.1),
..Default::default()
};
render_queue.write_buffer(
&pipeline.uniform_buffer,
dispatch_index as u64 * aligned_stride,
bytemuck::bytes_of(&uniforms),
);
let storage_bind_group = render_device.create_bind_group(
"retarded_vertices_storage",
&pipeline.storage_layout,
&[
BindGroupEntry {
binding: 0,
resource: buffers.keyframes.as_entire_binding(),
},
BindGroupEntry {
binding: 1,
resource: buffers.vertices_local.as_entire_binding(),
},
BindGroupEntry {
binding: 2,
resource: buffers.normals_local.as_entire_binding(),
},
BindGroupEntry {
binding: 3,
resource: vertex_slice.buffer.as_entire_binding(),
},
],
);
prepared.dispatches.push(PreparedDispatch {
storage_bind_group,
uniform_offset: (dispatch_index as u64 * aligned_stride) as u32,
workgroups: uniforms.num_vertices.div_ceil(64),
});
dispatch_index += 1;
}
}
fn attribute_f32_offset(attributes: &[VertexAttribute], location: u32) -> u32 {
let byte_offset = attributes
.iter()
.find(|a| a.shader_location == location)
.expect("required vertex attribute not found")
.offset;
(byte_offset / BYTES_PER_F32) as u32
}