use std::hash::Hash;
use bevy::prelude::*;
use serde::{Deserialize, Serialize};
use crate::{
calc_func_id,
expr::PropertyHandle,
graph::{BuiltInExpr, EvalContext, ExprError},
Attribute, BoxedModifier, ExprHandle, Modifier, ModifierContext, Module, ShaderWriter,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Reflect, Serialize, Deserialize)]
pub struct AccelModifier {
accel: ExprHandle,
}
impl AccelModifier {
pub fn new(accel: ExprHandle) -> Self {
Self { accel }
}
pub fn via_property(module: &mut Module, property: PropertyHandle) -> Self {
Self {
accel: module.prop(property),
}
}
pub fn constant(module: &mut Module, acceleration: Vec3) -> Self {
Self {
accel: module.lit(acceleration),
}
}
}
impl Modifier for AccelModifier {
fn context(&self) -> ModifierContext {
ModifierContext::Update
}
fn attributes(&self) -> &[Attribute] {
&[Attribute::VELOCITY]
}
fn boxed_clone(&self) -> BoxedModifier {
Box::new(*self)
}
fn apply(&self, module: &mut Module, context: &mut ShaderWriter) -> Result<(), ExprError> {
let attr = module.attr(Attribute::VELOCITY);
let attr = context.eval(module, attr)?;
let expr = context.eval(module, self.accel)?;
let dt = BuiltInExpr::new(crate::graph::BuiltInOperator::DeltaTime).eval(context)?;
context.main_code += &format!("{} += ({}) * {};", attr, expr, dt);
Ok(())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Reflect, Serialize, Deserialize)]
pub struct RadialAccelModifier {
origin: ExprHandle,
accel: ExprHandle,
}
impl RadialAccelModifier {
pub fn new(origin: ExprHandle, accel: ExprHandle) -> Self {
Self { origin, accel }
}
pub fn via_property(module: &mut Module, origin: Vec3, property: PropertyHandle) -> Self {
Self {
origin: module.lit(origin),
accel: module.prop(property),
}
}
pub fn constant(module: &mut Module, origin: Vec3, acceleration: f32) -> Self {
Self {
origin: module.lit(origin),
accel: module.lit(acceleration),
}
}
}
impl Modifier for RadialAccelModifier {
fn context(&self) -> ModifierContext {
ModifierContext::Update
}
fn attributes(&self) -> &[Attribute] {
&[Attribute::POSITION, Attribute::VELOCITY]
}
fn boxed_clone(&self) -> BoxedModifier {
Box::new(*self)
}
fn apply(&self, module: &mut Module, context: &mut ShaderWriter) -> Result<(), ExprError> {
let func_id = calc_func_id(self);
let func_name = format!("radial_accel_{0:016X}", func_id);
context.make_fn(
&func_name,
"particle: ptr<function, Particle>",
module,
&mut |m: &mut Module, ctx: &mut dyn EvalContext| -> Result<String, ExprError> {
let origin = ctx.eval(m, self.origin)?;
let accel = ctx.eval(m, self.accel)?;
Ok(format!(
r##"let radial = normalize((*particle).{} - {});
(*particle).{} += radial * (({}) * sim_params.delta_time);
"##,
Attribute::POSITION.name(),
origin,
Attribute::VELOCITY.name(),
accel,
))
},
)?;
context.main_code += &format!("{}(&particle);\n", func_name);
Ok(())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Reflect, Serialize, Deserialize)]
pub struct TangentAccelModifier {
origin: ExprHandle,
axis: ExprHandle,
accel: ExprHandle,
}
impl TangentAccelModifier {
pub fn new(origin: ExprHandle, axis: ExprHandle, accel: ExprHandle) -> Self {
Self {
origin,
axis,
accel,
}
}
pub fn via_property(
module: &mut Module,
origin: Vec3,
axis: Vec3,
property: PropertyHandle,
) -> Self {
Self {
origin: module.lit(origin),
axis: module.lit(axis),
accel: module.prop(property),
}
}
pub fn constant(module: &mut Module, origin: Vec3, axis: Vec3, acceleration: f32) -> Self {
Self {
origin: module.lit(origin),
axis: module.lit(axis),
accel: module.lit(acceleration),
}
}
}
impl Modifier for TangentAccelModifier {
fn context(&self) -> ModifierContext {
ModifierContext::Update
}
fn attributes(&self) -> &[Attribute] {
&[Attribute::POSITION, Attribute::VELOCITY]
}
fn boxed_clone(&self) -> BoxedModifier {
Box::new(*self)
}
fn apply(&self, module: &mut Module, context: &mut ShaderWriter) -> Result<(), ExprError> {
let func_id = calc_func_id(self);
let func_name = format!("tangent_accel_{0:016X}", func_id);
let origin = context.eval(module, self.origin)?;
let axis = context.eval(module, self.axis)?;
let accel = context.eval(module, self.accel)?;
context.extra_code += &format!(
r##"fn {}(particle: ptr<function, Particle>) {{
let radial = normalize((*particle).{} - {});
let tangent = normalize(cross({}, radial));
(*particle).{} += tangent * (({}) * sim_params.delta_time);
}}
"##,
func_name,
Attribute::POSITION.name(),
origin,
axis,
Attribute::VELOCITY.name(),
accel,
);
context.main_code += &format!("{}(&particle);\n", func_name);
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{ParticleLayout, Property, PropertyLayout, ToWgslString};
#[test]
fn mod_accel() {
let mut module = Module::default();
let accel = Vec3::new(1., 2., 3.);
let modifier = AccelModifier::constant(&mut module, accel);
let property_layout = PropertyLayout::default();
let particle_layout = ParticleLayout::default();
let mut context =
ShaderWriter::new(ModifierContext::Update, &property_layout, &particle_layout);
assert!(modifier.apply(&mut module, &mut context).is_ok());
assert!(context.main_code.contains(&accel.to_wgsl_string()));
}
#[test]
fn mod_radial_accel() {
let mut module = Module::default();
let property_layout = PropertyLayout::new(&[Property::new("my_prop", 3.)]);
let particle_layout = ParticleLayout::default();
let origin = Vec3::new(-1.2, 5.3, -8.5);
let accel = 6.;
let modifier = RadialAccelModifier::constant(&mut module, origin, accel);
let mut context =
ShaderWriter::new(ModifierContext::Update, &property_layout, &particle_layout);
assert!(modifier.apply(&mut module, &mut context).is_ok());
assert!(context.extra_code.contains(&accel.to_wgsl_string()));
let origin = module.attr(Attribute::POSITION);
let my_prop = module.add_property("my_prop", 3.0.into());
let accel = module.prop(my_prop);
let modifier = RadialAccelModifier::new(origin, accel);
let mut context =
ShaderWriter::new(ModifierContext::Update, &property_layout, &particle_layout);
assert!(modifier.apply(&mut module, &mut context).is_ok());
assert!(context.extra_code.contains(Attribute::POSITION.name()));
assert!(context.extra_code.contains("my_prop"));
}
#[test]
fn mod_tangent_accel() {
let mut module = Module::default();
let origin = Vec3::new(-1.2, 5.3, -8.5);
let axis = Vec3::Y;
let accel = 6.;
let modifier = TangentAccelModifier::constant(&mut module, origin, axis, accel);
let property_layout = PropertyLayout::default();
let particle_layout = ParticleLayout::default();
let mut context =
ShaderWriter::new(ModifierContext::Update, &property_layout, &particle_layout);
assert!(modifier.apply(&mut module, &mut context).is_ok());
assert!(context.extra_code.contains(&accel.to_wgsl_string()));
}
}