From d68dc200fd229480efe95a217b8fd417efa93a2f Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?charlotte=20=F0=9F=8C=B8?= Date: Mon, 28 Sep 2026 21:56:39 -0700 Subject: [PATCH] Neighbor lists. --- crates/processing_core/src/constants.rs | 2 + crates/processing_ffi/src/lib.rs | 19 +++ crates/processing_pyo3/src/constants.rs | 12 +- crates/processing_pyo3/src/particles.rs | 77 +++++++--- .../shaders/processing/particles.wesl | 1 + crates/processing_render/src/lib.rs | 21 +-- .../processing_render/src/particles/emit.rs | 6 +- .../src/particles/kernels/mod.rs | 3 + .../particles/kernels/neighbors_build.wgsl | 76 ++++++++++ .../particles/kernels/neighbors_reduce.wgsl | 73 ++++++++++ crates/processing_render/src/particles/mod.rs | 34 ++++- .../src/particles/neighbors.rs | 136 ++++++++++++++++++ 12 files changed, 424 insertions(+), 36 deletions(-) create mode 100644 crates/processing_render/src/particles/kernels/neighbors_build.wgsl create mode 100644 crates/processing_render/src/particles/kernels/neighbors_reduce.wgsl create mode 100644 crates/processing_render/src/particles/neighbors.rs diff --git a/crates/processing_core/src/constants.rs b/crates/processing_core/src/constants.rs index bb616946..3fb7ba0c 100644 --- a/crates/processing_core/src/constants.rs +++ b/crates/processing_core/src/constants.rs @@ -73,6 +73,7 @@ pub const EXTRACT: &str = "extract"; pub const PACK: &str = "pack"; pub const GENERATE: &str = "generate"; pub const NEIGHBOR: &str = "neighbor"; +pub const FIND_NEIGHBORS: &str = "find_neighbors"; pub const COUNT: &str = "count"; pub const DENSITY: &str = "density"; @@ -82,6 +83,7 @@ pub const SMOOTHSTEP: &str = "smoothstep"; pub const QUADRATIC: &str = "quadratic"; pub const CUBIC: &str = "cubic"; pub const INVERSE: &str = "inverse"; +pub const INVERSE_SQUARE: &str = "inverse_square"; pub const AFFINE: &str = "affine"; pub const ABS: &str = "abs"; diff --git a/crates/processing_ffi/src/lib.rs b/crates/processing_ffi/src/lib.rs index 2e4b49fa..6eec696a 100644 --- a/crates/processing_ffi/src/lib.rs +++ b/crates/processing_ffi/src/lib.rs @@ -3892,6 +3892,25 @@ pub extern "C" fn processing_particles_grid_destroy(grid_id: u64) { error::check(|| grid_destroy(Entity::from_bits(grid_id))); } +/// Writes each particle's `neighbors` and `neighbor_count`, rebuilding the grid first. +#[unsafe(no_mangle)] +pub extern "C" fn processing_particles_find_neighbors( + particles_id: u64, + grid_id: u64, + radius: f32, + max: u32, +) { + error::clear_error(); + error::check(|| { + particles_find_neighbors( + Entity::from_bits(particles_id), + Entity::from_bits(grid_id), + radius, + max, + ) + }); +} + #[unsafe(no_mangle)] pub extern "C" fn processing_particles_draw(graphics_id: u64, particles_id: u64, geometry_id: u64) { error::clear_error(); diff --git a/crates/processing_pyo3/src/constants.rs b/crates/processing_pyo3/src/constants.rs index bd1a7436..99486f82 100644 --- a/crates/processing_pyo3/src/constants.rs +++ b/crates/processing_pyo3/src/constants.rs @@ -67,9 +67,17 @@ pub fn register(m: &Bound<'_, PyModule>) -> PyResult<()> { add!(m, SUB, MUL, DIV, POW); add!(m, LENGTH, SUM, SUMSQ, MEAN, MIN, MAX); add!(m, UNIFORM, SIGNED, GAUSSIAN); - add!(m, NEIGHBOR); + add!(m, NEIGHBOR, FIND_NEIGHBORS); add!(m, COUNT, DENSITY); - add!(m, CONSTANT, SMOOTHSTEP, QUADRATIC, CUBIC, INVERSE); + add!( + m, + CONSTANT, + SMOOTHSTEP, + QUADRATIC, + CUBIC, + INVERSE, + INVERSE_SQUARE + ); add!( m, NOISE, diff --git a/crates/processing_pyo3/src/particles.rs b/crates/processing_pyo3/src/particles.rs index eeb54e57..ea270ede 100644 --- a/crates/processing_pyo3/src/particles.rs +++ b/crates/processing_pyo3/src/particles.rs @@ -111,6 +111,7 @@ fn parse_falloff(s: &str) -> PyResult { _ if s.eq_ignore_ascii_case(c::QUADRATIC) => Ok(FALLOFF_QUADRATIC), _ if s.eq_ignore_ascii_case(c::CUBIC) => Ok(FALLOFF_CUBIC), _ if s.eq_ignore_ascii_case(c::INVERSE) => Ok(FALLOFF_INVERSE), + _ if s.eq_ignore_ascii_case(c::INVERSE_SQUARE) => Ok(FALLOFF_INVERSE_SQUARE), _ => Err(PyValueError::new_err(format!( "neighbor: unknown falloff {s:?}" ))), @@ -787,19 +788,56 @@ impl Particles { let scale = kw_f32(kwargs, "scale", 1.0)?; let offset = kw_f32(kwargs, "offset", 0.0)?; algebra_generate(out, comp, mode, seed, scale, offset).map_err(rt) - } else if name.eq_ignore_ascii_case(c::NEIGHBOR) { - reject_unknown_kwargs(kwargs, &["a", "out", "grid", "op", "radius", "falloff"])?; + } else if name.eq_ignore_ascii_case(c::FIND_NEIGHBORS) { + reject_unknown_kwargs(kwargs, &["grid", "radius", "max"])?; let grid = kw(kwargs, "grid") - .ok_or_else(|| PyRuntimeError::new_err("apply(neighbor): missing 'grid'"))? + .ok_or_else(|| PyTypeError::new_err("apply(find_neighbors): missing `grid=`"))? .extract::>() - .map_err(|_| PyRuntimeError::new_err("apply(neighbor): 'grid' must be a Grid"))?; + .map_err(|_| PyTypeError::new_err("apply(find_neighbors): `grid` must be a Grid"))? + .entity; + let cell = grid_get(grid).map_err(rt)?.params.cell_size; + let radius = kw_f32(kwargs, "radius", cell)?; + let max = kw_u32(kwargs, "max", 64)?; + particles_find_neighbors(self.entity, grid, radius, max).map_err(rt) + } else if name.eq_ignore_ascii_case(c::NEIGHBOR) { + reject_unknown_kwargs( + kwargs, + &["a", "out", "grid", "op", "radius", "falloff", "relative"], + )?; let op = kw_op(kwargs, NEIGHBOR_MEAN, parse_neighbor_op)?; + let relative = match kw(kwargs, "relative") { + Some(v) => v.extract::()?, + None => false, + }; + // `grid=` searches exactly; otherwise read the lists from find_neighbors + let grid = match kw(kwargs, "grid") { + Some(g) => { + if relative { + return Err(PyValueError::new_err( + "apply(neighbor): relative=True reads the neighbor lists, so drop grid=", + )); + } + let g = g.extract::>().map_err(|_| { + PyTypeError::new_err("apply(neighbor): `grid` must be a Grid") + })?; + Some(g.entity) + } + None => None, + }; + let (falloff_default, radius) = match grid { + Some(g) => { + let cell = grid_get(g).map_err(rt)?.params.cell_size; + ( + FALLOFF_SMOOTHSTEP, + kw_f32(kwargs, "radius", cell)?.min(cell), + ) + } + None => (FALLOFF_CONST, kw_f32(kwargs, "radius", f32::MAX)?), + }; let falloff = match kw(kwargs, "falloff") { Some(v) => parse_falloff(&v.extract::()?)?, - None => FALLOFF_SMOOTHSTEP, + None => falloff_default, }; - let cell = grid.cell_size()?; - let radius = kw_f32(kwargs, "radius", cell)?.min(cell); let (a, components) = if op == NEIGHBOR_COUNT { let a = match kw(kwargs, "a") { @@ -816,16 +854,21 @@ impl Particles { self.operand(kwargs, "a")? }; let out = self.output(kwargs, components)?; - particles_gather( - self.entity, - grid.entity, - a, - out, - op, - radius, - falloff, - components, - ) + match grid { + Some(grid) => { + particles_gather(self.entity, grid, a, out, op, radius, falloff, components) + } + None => particles_neighbor_reduce( + self.entity, + a, + out, + op, + radius, + falloff, + components, + relative, + ), + } .map_err(rt) } else { Err(PyValueError::new_err(format!( diff --git a/crates/processing_render/shaders/processing/particles.wesl b/crates/processing_render/shaders/processing/particles.wesl index 66984c73..331ffaed 100644 --- a/crates/processing_render/shaders/processing/particles.wesl +++ b/crates/processing_render/shaders/processing/particles.wesl @@ -6,6 +6,7 @@ fn falloff(d: f32, radius: f32, mode: u32) -> f32 { case 3u: { return n * n; } case 4u: { return n * n * n; } case 5u: { return radius / (d + radius); } + case 6u: { return 1.0 / max(d * d, 1e-6); } default: { return 1.0; } } } diff --git a/crates/processing_render/src/lib.rs b/crates/processing_render/src/lib.rs index 499f3a38..d215ea84 100644 --- a/crates/processing_render/src/lib.rs +++ b/crates/processing_render/src/lib.rs @@ -29,20 +29,23 @@ pub use particles::algebra::{ }; pub use particles::compact::compact; pub use particles::grid::{Grid, GridParams, grid_build, grid_create, grid_destroy, grid_get}; +pub use particles::neighbors::{ + particles_find_neighbors, particles_neighbor_lists, particles_neighbor_reduce, +}; pub use particles::reduce::{REDUCE_OP_MAX, REDUCE_OP_MIN, REDUCE_OP_SUM, reduce}; pub use particles::sort::bitonic_sort_by_key; pub use particles::{ BOUNDS_CLAMP, BOUNDS_REFLECT, BOUNDS_SOFT, BOUNDS_WRAP, COMBINE_ADD, COMBINE_DIV, COMBINE_MAX, COMBINE_MIN, COMBINE_MUL, COMBINE_POW, COMBINE_SUB, FALLOFF_CONST, FALLOFF_CUBIC, - FALLOFF_INVERSE, FALLOFF_LINEAR, FALLOFF_QUADRATIC, FALLOFF_SMOOTHSTEP, particles_apply, - particles_attribute_add, particles_attributes, particles_buffer, particles_capacity, - particles_connectivity_indirect, particles_create, particles_create_from_geometry, - particles_destroy, particles_emit, particles_emit_gpu, particles_ensure_attribute, - particles_flock, particles_gather, particles_kernel_age, particles_kernel_attr_combine, - particles_kernel_attr_linear, particles_kernel_attr_lookup1d, particles_kernel_attr_lookup2d, - particles_kernel_attr_mix, particles_kernel_attract, particles_kernel_bounds_box, - particles_kernel_bounds_geometry, particles_kernel_bounds_sphere, particles_kernel_drag, - particles_kernel_field, particles_kernel_flock, particles_kernel_force, + FALLOFF_INVERSE, FALLOFF_INVERSE_SQUARE, FALLOFF_LINEAR, FALLOFF_QUADRATIC, FALLOFF_SMOOTHSTEP, + particles_apply, particles_attribute_add, particles_attributes, particles_buffer, + particles_capacity, particles_connectivity_indirect, particles_create, + particles_create_from_geometry, particles_destroy, particles_emit, particles_emit_gpu, + particles_ensure_attribute, particles_flock, particles_gather, particles_kernel_age, + particles_kernel_attr_combine, particles_kernel_attr_linear, particles_kernel_attr_lookup1d, + particles_kernel_attr_lookup2d, particles_kernel_attr_mix, particles_kernel_attract, + particles_kernel_bounds_box, particles_kernel_bounds_geometry, particles_kernel_bounds_sphere, + particles_kernel_drag, particles_kernel_field, particles_kernel_flock, particles_kernel_force, particles_kernel_impulse, particles_kernel_integrate, particles_kernel_noise, particles_kernel_orient, particles_kernel_transform, particles_kernel_vortex, particles_reset_indices, particles_scatter_create, particles_scatter_volume_create, diff --git a/crates/processing_render/src/particles/emit.rs b/crates/processing_render/src/particles/emit.rs index 16e3ca27..fb5f30e1 100644 --- a/crates/processing_render/src/particles/emit.rs +++ b/crates/processing_render/src/particles/emit.rs @@ -252,13 +252,17 @@ pub fn particles_apply(particles_entity: Entity, compute_entity: Entity) -> erro let field = world .get::(particles_entity) .ok_or(error::ProcessingError::ParticlesNotFound)?; - let mut buffers: Vec<(String, Entity)> = Vec::with_capacity(field.buffers.len()); + let mut buffers: Vec<(String, Entity)> = Vec::with_capacity(field.buffers.len() + 2); for (&attr_entity, &buf_entity) in &field.buffers { let attr = world .get::(attr_entity) .ok_or(error::ProcessingError::InvalidEntity)?; buffers.push((attr.name.to_string(), buf_entity)); } + if let Some(lists) = field.neighbor_lists { + buffers.push(("neighbors".to_string(), lists.neighbors)); + buffers.push(("neighbor_count".to_string(), lists.count)); + } Ok((field.capacity, buffers)) })?; diff --git a/crates/processing_render/src/particles/kernels/mod.rs b/crates/processing_render/src/particles/kernels/mod.rs index a552af51..df9199fd 100644 --- a/crates/processing_render/src/particles/kernels/mod.rs +++ b/crates/processing_render/src/particles/kernels/mod.rs @@ -63,6 +63,8 @@ impl Plugin for ParticlesKernelsPlugin { embedded_asset!(app, "compact_scatter.wgsl"); embedded_asset!(app, "reduce.wgsl"); embedded_asset!(app, "neighbor.wgsl"); + embedded_asset!(app, "neighbors_build.wgsl"); + embedded_asset!(app, "neighbors_reduce.wgsl"); } } @@ -72,6 +74,7 @@ pub const FALLOFF_SMOOTHSTEP: u32 = 2; pub const FALLOFF_QUADRATIC: u32 = 3; pub const FALLOFF_CUBIC: u32 = 4; pub const FALLOFF_INVERSE: u32 = 5; +pub const FALLOFF_INVERSE_SQUARE: u32 = 6; pub const BOUNDS_CLAMP: u32 = 0; pub const BOUNDS_REFLECT: u32 = 1; diff --git a/crates/processing_render/src/particles/kernels/neighbors_build.wgsl b/crates/processing_render/src/particles/kernels/neighbors_build.wgsl new file mode 100644 index 00000000..ab3dae3c --- /dev/null +++ b/crates/processing_render/src/particles/kernels/neighbors_build.wgsl @@ -0,0 +1,76 @@ +import processing::particles::{Grid, cell_coords, cell_index}; + +// with radius == cell size, ~1 in 6.5 candidates is a neighbor; rounded up +const CANDIDATES_PER_NEIGHBOR: f32 = 7.0; + +struct Params { + radius: f32, + max_neighbors: u32, +} + +@group(0) @binding(0) var position: array; +@group(0) @binding(1) var grid_offsets: array; +@group(0) @binding(2) var grid_sorted: array; +@group(0) @binding(3) var grid: Grid; +@group(0) @binding(4) var neighbors: array; +@group(0) @binding(5) var neighbor_count: array; +@group(0) @binding(6) var params: Params; + +fn load_pos(i: u32) -> vec3 { + return vec3(position[i * 3u], position[i * 3u + 1u], position[i * 3u + 2u]); +} + +@compute @workgroup_size(64) +fn main(@builtin(global_invocation_id) gid: vec3) { + let i = gid.x; + if i >= arrayLength(&neighbor_count) { return; } + + let pos = load_pos(i); + let r2 = params.radius * params.radius; + let dims = grid.dims; + let base = cell_coords(pos, grid.origin, grid.cell_size, dims); + let reach = max(1, i32(ceil(params.radius / grid.cell_size))); + let lo = max(base - vec3(reach), vec3(0)); + let hi = min(base + vec3(reach), vec3(dims) - vec3(1)); + + // past the limit, sample every cell evenly so lists aren't biased toward the first cells + var total = 0u; + for (var cz = lo.z; cz <= hi.z; cz++) { + for (var cy = lo.y; cy <= hi.y; cy++) { + for (var cx = lo.x; cx <= hi.x; cx++) { + let cell = cell_index(vec3(u32(cx), u32(cy), u32(cz)), dims); + total += grid_offsets[cell + 1u] - grid_offsets[cell]; + } + } + } + let rate = min(1.0, f32(params.max_neighbors) * CANDIDATES_PER_NEIGHBOR / f32(max(total, 1u))); + // per-particle offset so particles in a cell don't all sample the same members + let phase = i * 2654435761u; + + let first = i * params.max_neighbors; + var found = 0u; + for (var cz = lo.z; cz <= hi.z && found < params.max_neighbors; cz++) { + for (var cy = lo.y; cy <= hi.y && found < params.max_neighbors; cy++) { + for (var cx = lo.x; cx <= hi.x && found < params.max_neighbors; cx++) { + let cell = cell_index(vec3(u32(cx), u32(cy), u32(cz)), dims); + let start = grid_offsets[cell]; + let in_cell = grid_offsets[cell + 1u] - start; + let take = min(in_cell, u32(ceil(f32(in_cell) * rate))); + for (var k = 0u; k < take && found < params.max_neighbors; k++) { + var s = start + k; + if take < in_cell { + s = start + (phase + u32(f32(k) * f32(in_cell) / f32(take))) % in_cell; + } + let j = grid_sorted[s]; + if j == i { continue; } + let diff = load_pos(j) - pos; + if dot(diff, diff) <= r2 { + neighbors[first + found] = j; + found += 1u; + } + } + } + } + } + neighbor_count[i] = found; +} diff --git a/crates/processing_render/src/particles/kernels/neighbors_reduce.wgsl b/crates/processing_render/src/particles/kernels/neighbors_reduce.wgsl new file mode 100644 index 00000000..09beee95 --- /dev/null +++ b/crates/processing_render/src/particles/kernels/neighbors_reduce.wgsl @@ -0,0 +1,73 @@ +import processing::particles::falloff; + +struct Params { + max_distance: f32, + op: u32, + falloff_mode: u32, + components: u32, + // 1 = reduce `source[j] - source[i]` + relative: u32, +} + +const OP_SUM: u32 = 0u; +const OP_MEAN: u32 = 1u; +const OP_COUNT: u32 = 2u; + +@group(0) @binding(0) var position: array; +@group(0) @binding(1) var source: array; +@group(0) @binding(2) var out: array; +@group(0) @binding(3) var neighbors: array; +@group(0) @binding(4) var neighbor_count: array; +@group(0) @binding(5) var params: Params; + +fn load_pos(i: u32) -> vec3 { + return vec3(position[i * 3u], position[i * 3u + 1u], position[i * 3u + 2u]); +} + +@compute @workgroup_size(64) +fn main(@builtin(global_invocation_id) gid: vec3) { + let i = gid.x; + let n = arrayLength(&neighbor_count); + if i >= n { return; } + + let pos = load_pos(i); + let r2 = params.max_distance * params.max_distance; + let comps = params.components; + + var own = array(0.0, 0.0, 0.0, 0.0); + if params.relative == 1u { + for (var c = 0u; c < comps; c++) { + own[c] = source[i * comps + c]; + } + } + + var value = array(0.0, 0.0, 0.0, 0.0); + var weight_sum = 0.0; + let first = i * (arrayLength(&neighbors) / n); + for (var k = 0u; k < neighbor_count[i]; k++) { + let j = neighbors[first + k]; + let diff = load_pos(j) - pos; + let d2 = dot(diff, diff); + if d2 > r2 { continue; } + let w = falloff(sqrt(d2), params.max_distance, params.falloff_mode); + weight_sum += w; + if params.op != OP_COUNT { + for (var c = 0u; c < comps; c++) { + value[c] += w * (source[j * comps + c] - own[c]); + } + } + } + + if params.op == OP_COUNT { + out[i] = weight_sum; + } else if params.op == OP_MEAN { + let inv = select(0.0, 1.0 / weight_sum, weight_sum > 0.0); + for (var c = 0u; c < comps; c++) { + out[i * comps + c] = value[c] * inv; + } + } else { + for (var c = 0u; c < comps; c++) { + out[i * comps + c] = value[c]; + } + } +} diff --git a/crates/processing_render/src/particles/mod.rs b/crates/processing_render/src/particles/mod.rs index ddc33a17..f0b20316 100644 --- a/crates/processing_render/src/particles/mod.rs +++ b/crates/processing_render/src/particles/mod.rs @@ -6,6 +6,7 @@ mod emit; pub mod grid; pub mod kernels; pub mod material; +pub mod neighbors; pub mod pack; pub mod point_render; pub mod reduce; @@ -27,13 +28,17 @@ pub use grid::{Grid, GridParams, grid_build, grid_create, grid_destroy, grid_get pub use kernels::{ BOUNDS_CLAMP, BOUNDS_REFLECT, BOUNDS_SOFT, BOUNDS_WRAP, COMBINE_ADD, COMBINE_DIV, COMBINE_MAX, COMBINE_MIN, COMBINE_MUL, COMBINE_POW, COMBINE_SUB, FALLOFF_CONST, FALLOFF_CUBIC, - FALLOFF_INVERSE, FALLOFF_LINEAR, FALLOFF_QUADRATIC, FALLOFF_SMOOTHSTEP, particles_kernel_age, - particles_kernel_attr_combine, particles_kernel_attr_linear, particles_kernel_attr_lookup1d, - particles_kernel_attr_lookup2d, particles_kernel_attr_mix, particles_kernel_attract, - particles_kernel_bounds_box, particles_kernel_bounds_geometry, particles_kernel_bounds_sphere, - particles_kernel_drag, particles_kernel_field, particles_kernel_flock, particles_kernel_force, - particles_kernel_impulse, particles_kernel_integrate, particles_kernel_noise, - particles_kernel_orient, particles_kernel_transform, particles_kernel_vortex, + FALLOFF_INVERSE, FALLOFF_INVERSE_SQUARE, FALLOFF_LINEAR, FALLOFF_QUADRATIC, FALLOFF_SMOOTHSTEP, + particles_kernel_age, particles_kernel_attr_combine, particles_kernel_attr_linear, + particles_kernel_attr_lookup1d, particles_kernel_attr_lookup2d, particles_kernel_attr_mix, + particles_kernel_attract, particles_kernel_bounds_box, particles_kernel_bounds_geometry, + particles_kernel_bounds_sphere, particles_kernel_drag, particles_kernel_field, + particles_kernel_flock, particles_kernel_force, particles_kernel_impulse, + particles_kernel_integrate, particles_kernel_noise, particles_kernel_orient, + particles_kernel_transform, particles_kernel_vortex, +}; +pub use neighbors::{ + particles_find_neighbors, particles_neighbor_lists, particles_neighbor_reduce, }; pub use reduce::{REDUCE_OP_MAX, REDUCE_OP_MIN, REDUCE_OP_SUM, reduce}; pub use scan::prefix_sum_u32; @@ -96,6 +101,15 @@ pub struct Particles { pub connectivity: Option, /// Ring-buffer write cursor; wraps at `capacity`. pub emit_head: u32, + pub neighbor_lists: Option, +} + +/// The `neighbors` (up to `max` indices per particle) and `neighbor_count` attributes. +#[derive(Clone, Copy)] +pub struct NeighborLists { + pub neighbors: Entity, + pub count: Entity, + pub max: u32, } #[derive(Clone, Copy)] @@ -150,6 +164,7 @@ pub fn create( raster_draw_entity: None, connectivity: None, emit_head: 0, + neighbor_lists: None, }) .id(); Ok(entity) @@ -231,6 +246,7 @@ pub fn create_from_geometry( raster_draw_entity: None, connectivity, emit_head: 0, + neighbor_lists: None, }) .id(); Ok(entity) @@ -387,6 +403,10 @@ pub fn destroy( commands.entity(connectivity.index_buffer).despawn(); commands.entity(connectivity.indirect_buffer).despawn(); } + if let Some(lists) = p.neighbor_lists { + commands.entity(lists.neighbors).despawn(); + commands.entity(lists.count).despawn(); + } commands.entity(entity).despawn(); Ok(()) } diff --git a/crates/processing_render/src/particles/neighbors.rs b/crates/processing_render/src/particles/neighbors.rs new file mode 100644 index 00000000..321870f4 --- /dev/null +++ b/crates/processing_render/src/particles/neighbors.rs @@ -0,0 +1,136 @@ +use std::sync::Mutex; + +use bevy::prelude::Entity; + +use processing_core::app_mut; +use processing_core::error::{ProcessingError, Result}; + +use crate::particles::grid::grid_build; +use crate::particles::{NeighborLists, Particles}; +use crate::shader_value::ShaderValue; +use crate::{ + buffer_create, buffer_destroy, compute_create, compute_dispatch_no_update, compute_set, + geometry_attribute_position, shader_load, +}; + +const BUILD_SHADER: &str = "embedded://processing_render/particles/kernels/neighbors_build.wgsl"; +const REDUCE_SHADER: &str = "embedded://processing_render/particles/kernels/neighbors_reduce.wgsl"; +const WG: u32 = 64; + +static COMPUTES: Mutex> = Mutex::new(None); + +fn computes() -> Result<(Entity, Entity)> { + let mut guard = COMPUTES.lock().unwrap(); + if let Some(v) = *guard { + return Ok(v); + } + let build = compute_create(shader_load(BUILD_SHADER)?)?; + let reduce = compute_create(shader_load(REDUCE_SHADER)?)?; + *guard = Some((build, reduce)); + Ok((build, reduce)) +} + +fn position_buffer(particles: Entity) -> Result { + crate::particles::particles_buffer(particles, geometry_attribute_position())?.ok_or_else(|| { + ProcessingError::InvalidArgument("neighbors need a `position` attribute".to_string()) + }) +} + +/// The `neighbors` and `neighbor_count` attributes, if `particles_find_neighbors` has run. +pub fn particles_neighbor_lists(particles: Entity) -> Result> { + app_mut(|app| { + Ok(app + .world() + .get::(particles) + .ok_or(ProcessingError::ParticlesNotFound)? + .neighbor_lists) + }) +} + +/// Rebuilds `grid` from the particles' positions, then writes each particle's `neighbors` +/// (up to `max` within `radius`) and `neighbor_count`. Past `max`, a spread sample is kept, +/// not the nearest. +pub fn particles_find_neighbors( + particles: Entity, + grid: Entity, + radius: f32, + max: u32, +) -> Result<()> { + if max == 0 { + return Err(ProcessingError::InvalidArgument( + "find_neighbors: max must be at least 1".to_string(), + )); + } + let capacity = crate::particles::particles_capacity(particles)?; + let lists = match particles_neighbor_lists(particles)? { + Some(lists) if lists.max == max => lists, + previous => { + if let Some(previous) = previous { + buffer_destroy(previous.neighbors)?; + buffer_destroy(previous.count)?; + } + let slots = capacity.max(1) as u64; + let lists = NeighborLists { + neighbors: buffer_create(slots * max as u64 * 4)?, + count: buffer_create(slots * 4)?, + max, + }; + app_mut(|app| { + app.world_mut() + .get_mut::(particles) + .ok_or(ProcessingError::ParticlesNotFound)? + .neighbor_lists = Some(lists); + Ok(()) + })?; + lists + } + }; + + let position = position_buffer(particles)?; + grid_build(grid, position)?; + let (build, _) = computes()?; + compute_set(build, "position", ShaderValue::Buffer(position))?; + compute_set(build, "grid", ShaderValue::Grid(grid))?; + compute_set(build, "neighbors", ShaderValue::Buffer(lists.neighbors))?; + compute_set(build, "neighbor_count", ShaderValue::Buffer(lists.count))?; + compute_set(build, "radius", ShaderValue::Float(radius))?; + compute_set(build, "max_neighbors", ShaderValue::UInt(max))?; + compute_dispatch_no_update(build, capacity.div_ceil(WG), 1, 1) +} + +/// Like `particles_gather`, but over each particle's `neighbors`, so a particle never counts +/// itself. With `relative`, reduces `source[j] - source[i]`. +#[allow(clippy::too_many_arguments)] +pub fn particles_neighbor_reduce( + particles: Entity, + source: Entity, + out: Entity, + op: u32, + radius: f32, + falloff_mode: u32, + components: u32, + relative: bool, +) -> Result<()> { + let lists = particles_neighbor_lists(particles)?.ok_or_else(|| { + ProcessingError::InvalidArgument("run find_neighbors before reading neighbors".to_string()) + })?; + let position = position_buffer(particles)?; + if out == source || out == position { + return Err(ProcessingError::InvalidArgument( + "neighbor: `out` must differ from the source and position".to_string(), + )); + } + let capacity = crate::particles::particles_capacity(particles)?; + let (_, reduce) = computes()?; + compute_set(reduce, "position", ShaderValue::Buffer(position))?; + compute_set(reduce, "source", ShaderValue::Buffer(source))?; + compute_set(reduce, "out", ShaderValue::Buffer(out))?; + compute_set(reduce, "neighbors", ShaderValue::Buffer(lists.neighbors))?; + compute_set(reduce, "neighbor_count", ShaderValue::Buffer(lists.count))?; + compute_set(reduce, "max_distance", ShaderValue::Float(radius))?; + compute_set(reduce, "op", ShaderValue::UInt(op))?; + compute_set(reduce, "falloff_mode", ShaderValue::UInt(falloff_mode))?; + compute_set(reduce, "components", ShaderValue::UInt(components))?; + compute_set(reduce, "relative", ShaderValue::UInt(relative as u32))?; + compute_dispatch_no_update(reduce, capacity.div_ceil(WG), 1, 1) +}