use std::collections::{BTreeMap, HashMap};
use bevy::{
math::{UVec2, Vec2, Vec3, Vec4},
reflect::{PartialReflect, Reflect, ReflectRef, structs::Struct},
};
use bevy_hanabi::{
Attribute, CpuValue, EffectAsset, ExprHandle, Gradient, Module, Value,
graph::expr::{Expr, LiteralExpr, PropertyExpr, PropertyHandle},
};
use crate::{
ModifierGroup,
bake::value_as_u32,
model::{
EditValue, EffectGraph, EffectHeader, ExprNode, GradientVec3, GradientVec4, GraphLink,
GraphNode, GraphStack, ImageBinding, InputSlot, ModifierNodeData, NodeId, NodePayload,
PortRef, PropertyDef, PropertyId, SharedStr, SlotId, TextureSlotDef,
},
schema::{ConfigKind, FieldRole, OUTPUT_PORT, modifier_schema},
};
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ImportWarning {
pub message: String,
}
impl ImportWarning {
fn new(message: impl Into<String>) -> Self {
Self {
message: message.into(),
}
}
}
impl std::fmt::Display for ImportWarning {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(&self.message)
}
}
pub fn import(asset: &EffectAsset) -> (EffectGraph, Vec<ImportWarning>) {
let module = asset.module();
let mut graph = EffectGraph {
header: EffectHeader {
name: asset.name.clone().into(),
capacity: asset.capacity(),
spawner: asset.spawner,
simulation_space: asset.simulation_space,
simulation_condition: asset.simulation_condition,
z_layer_2d: asset.z_layer_2d,
},
..EffectGraph::empty()
};
let mut props_by_name: HashMap<String, PropertyId> = HashMap::new();
for prop in module.properties() {
let id = graph.alloc_property_id();
graph.properties.push(PropertyDef {
id,
name: prop.name().into(),
default: *prop.default_value(),
exposed: true,
});
props_by_name.insert(prop.name().to_string(), id);
}
let mut slot_ids: Vec<SlotId> = Vec::new();
for slot in module.texture_layout().layout {
let id = graph.alloc_slot_id();
graph.texture_slots.push(TextureSlotDef {
id,
name: slot.name.into(),
});
slot_ids.push(id);
}
let mut importer = Importer {
graph: &mut graph,
module,
props_by_name: &props_by_name,
prop_ref_nodes: HashMap::new(),
slot_ids,
warnings: Vec::new(),
};
let init: Vec<NodeId> = asset
.init_modifiers()
.filter_map(|m| importer.import_modifier(m.as_reflect()))
.collect();
let update: Vec<NodeId> = asset
.update_modifiers()
.filter_map(|m| importer.import_modifier(m.as_reflect()))
.collect();
let render: Vec<NodeId> = asset
.render_modifiers()
.filter_map(|m| importer.import_modifier(m.as_modifier().as_reflect()))
.collect();
let warnings = std::mem::take(&mut importer.warnings);
for (group, members) in [
(ModifierGroup::Init, init),
(ModifierGroup::Update, update),
(ModifierGroup::Render, render),
] {
if members.is_empty() {
continue;
}
let id = graph.alloc_stack_id();
graph.stacks.push(GraphStack { id, group, members });
}
(graph, warnings)
}
struct Importer<'a> {
graph: &'a mut EffectGraph,
module: &'a Module,
props_by_name: &'a HashMap<String, PropertyId>,
prop_ref_nodes: HashMap<PropertyId, NodeId>,
slot_ids: Vec<SlotId>,
warnings: Vec<ImportWarning>,
}
impl Importer<'_> {
fn import_modifier(&mut self, reflect: &dyn Reflect) -> Option<NodeId> {
let type_path = reflect.reflect_type_path();
let Some(info) = reflect.get_represented_type_info() else {
self.warnings.push(ImportWarning::new(format!(
"modifier '{type_path}' has no type info; skipped"
)));
return None;
};
let Some(schema) = modifier_schema(info) else {
self.warnings.push(ImportWarning::new(format!(
"modifier '{type_path}' does not reflect as a struct; skipped"
)));
return None;
};
let node_id = self.graph.alloc_node_id();
let mut config: BTreeMap<SharedStr, EditValue> = BTreeMap::new();
for field in schema.config() {
match read_config_field(reflect, &field.name, &field.role, field.type_path) {
Ok(value) => {
config.insert(field.name.clone(), value);
}
Err(message) => self.warnings.push(ImportWarning::new(format!(
"modifier '{type_path}' field '{}': {message}",
field.name
))),
}
}
let mut inputs: Vec<InputSlot> = Vec::new();
let mut links: Vec<GraphLink> = Vec::new();
for field in schema.ports() {
if matches!(field.role, FieldRole::Texture) {
let binding = self.recover_image_binding(reflect, &field.name);
inputs.push(InputSlot {
name: field.name.clone(),
default: binding.into(),
});
continue;
}
let optional = matches!(field.role, FieldRole::ExprPort { optional: true });
let Some(handle) = read_expr_handle(reflect, &field.name, optional) else {
continue;
};
match self.recover_port(node_id, &field.name, handle) {
PortInput::Inline(value) => inputs.push(InputSlot {
name: field.name.clone(),
default: value.into(),
}),
PortInput::Link(link) => {
links.push(link);
if let Some(default) = handle_value_type_default(self.module, handle) {
inputs.push(InputSlot {
name: field.name.clone(),
default: default.into(),
});
}
}
}
}
self.graph.nodes.push(GraphNode {
id: node_id,
payload: NodePayload::Modifier(ModifierNodeData::Known {
type_path: type_path.into(),
config,
}),
inputs,
});
self.graph.links.extend(links);
Some(node_id)
}
fn recover_port(&mut self, node: NodeId, port: &SharedStr, handle: ExprHandle) -> PortInput {
match self.module.get(handle) {
Some(Expr::Literal(lit)) => match literal_value(lit) {
Some(value) => PortInput::Inline(value),
None => {
self.warnings.push(ImportWarning::new(format!(
"port '{port}': could not read literal value; reset to 0"
)));
PortInput::Inline(Value::from(0.0_f32))
}
},
Some(Expr::Property(pe)) => match self.property_ref(pe) {
Some(ref_node) => PortInput::Link(GraphLink {
from: PortRef {
node: ref_node,
port: OUTPUT_PORT.into(),
},
to: PortRef {
node,
port: port.clone(),
},
}),
None => {
self.warnings.push(ImportWarning::new(format!(
"port '{port}': references an unknown property; reset to 0"
)));
PortInput::Inline(Value::from(0.0_f32))
}
},
other => {
let kind = other.map(expr_kind).unwrap_or("missing");
self.warnings.push(ImportWarning::new(format!(
"port '{port}': {kind} expression input cannot be reversed; reset to default"
)));
let value = handle_value_type_default(self.module, handle)
.unwrap_or_else(|| Value::from(0.0_f32));
PortInput::Inline(value)
}
}
}
fn recover_image_binding(&mut self, reflect: &dyn Reflect, field: &SharedStr) -> ImageBinding {
let index = read_expr_handle(reflect, field, false).and_then(|handle| {
match self.module.get(handle) {
Some(Expr::Literal(lit)) => literal_value(lit).as_ref().and_then(value_as_u32),
_ => None,
}
});
match index.and_then(|i| self.slot_ids.get(i as usize).copied()) {
Some(id) => ImageBinding::Slot(id),
None => {
self.warnings.push(ImportWarning::new(format!(
"texture port '{field}': slot index could not be recovered; left unbound"
)));
ImageBinding::Unbound
}
}
}
fn property_ref(&mut self, pe: &PropertyExpr) -> Option<NodeId> {
let handle = property_handle(pe)?;
let name = self.module.get_property(handle)?.name();
let prop_id = *self.props_by_name.get(name)?;
if let Some(&existing) = self.prop_ref_nodes.get(&prop_id) {
return Some(existing);
}
let id = self.graph.alloc_node_id();
self.graph.nodes.push(GraphNode {
id,
payload: NodePayload::Expr(ExprNode::Property(prop_id)),
inputs: Vec::new(),
});
self.prop_ref_nodes.insert(prop_id, id);
Some(id)
}
}
enum PortInput {
Inline(Value),
Link(GraphLink),
}
fn read_config_field(
reflect: &dyn Reflect,
name: &str,
role: &FieldRole,
type_path: &str,
) -> Result<EditValue, String> {
let ReflectRef::Struct(s) = reflect.reflect_ref() else {
return Err("modifier does not reflect as a struct".to_string());
};
let field = s
.field(name)
.ok_or_else(|| format!("no such field '{name}'"))?;
let kind = match role {
FieldRole::Config(kind) => *kind,
FieldRole::Texture => {
return Err("texture bindings cannot be read back from a baked asset".to_string());
}
FieldRole::ExprPort { .. } => return Err("expression port is not config".to_string()),
};
match kind {
ConfigKind::Bool => downcast::<bool>(field).map(EditValue::Bool),
ConfigKind::U32 => downcast::<u32>(field).map(EditValue::U32),
ConfigKind::UVec2 => downcast::<UVec2>(field).map(EditValue::UVec2),
ConfigKind::Attribute => downcast::<Attribute>(field).map(EditValue::Attribute),
ConfigKind::CpuVec3 => downcast::<CpuValue<Vec3>>(field).map(EditValue::CpuVec3),
ConfigKind::CpuVec4 => downcast::<CpuValue<Vec4>>(field).map(EditValue::CpuVec4),
ConfigKind::Gradient3 => downcast::<Gradient<Vec3>>(field)
.map(|g| EditValue::Gradient3(GradientVec3::Analytical(g))),
ConfigKind::Gradient4 => downcast::<Gradient<Vec4>>(field)
.map(|g| EditValue::Gradient4(GradientVec4::Analytical(g))),
ConfigKind::Scalar => read_scalar(field),
ConfigKind::Enum => read_enum(field, type_path),
ConfigKind::Flags => read_flags(field, type_path),
ConfigKind::Raw => Err("unmodeled field type cannot be read back".to_string()),
}
}
fn downcast<T: Reflect + Clone>(field: &dyn PartialReflect) -> Result<T, String> {
field.try_downcast_ref::<T>().cloned().ok_or_else(|| {
format!(
"expected {}, found {}",
std::any::type_name::<T>(),
field.reflect_type_path()
)
})
}
fn read_scalar(field: &dyn PartialReflect) -> Result<EditValue, String> {
if let Some(v) = field.try_downcast_ref::<f32>() {
Ok(EditValue::Scalar(Value::from(*v)))
} else if let Some(v) = field.try_downcast_ref::<i32>() {
Ok(EditValue::Scalar(Value::from(*v)))
} else if let Some(v) = field.try_downcast_ref::<u32>() {
Ok(EditValue::Scalar(Value::from(*v)))
} else if let Some(v) = field.try_downcast_ref::<Vec2>() {
Ok(EditValue::Scalar(Value::from(*v)))
} else if let Some(v) = field.try_downcast_ref::<Vec3>() {
Ok(EditValue::Scalar(Value::from(*v)))
} else if let Some(v) = field.try_downcast_ref::<Vec4>() {
Ok(EditValue::Scalar(Value::from(*v)))
} else {
Err(format!(
"unsupported scalar field type {}",
field.reflect_type_path()
))
}
}
fn read_enum(field: &dyn PartialReflect, type_path: &str) -> Result<EditValue, String> {
let ReflectRef::Enum(e) = field.reflect_ref() else {
return Err("expected an enum field".to_string());
};
Ok(EditValue::Enum {
type_path: type_path.into(),
variant: e.variant_name().into(),
})
}
fn read_flags(field: &dyn PartialReflect, type_path: &str) -> Result<EditValue, String> {
let ReflectRef::TupleStruct(ts) = field.reflect_ref() else {
return Err("flags field is not a tuple struct".to_string());
};
let inner = ts.field(0).ok_or("flags newtype has no inner value")?;
let bits = if let Some(b) = inner.try_downcast_ref::<u8>() {
*b as u64
} else if let Some(b) = inner.try_downcast_ref::<u16>() {
*b as u64
} else if let Some(b) = inner.try_downcast_ref::<u32>() {
*b as u64
} else if let Some(b) = inner.try_downcast_ref::<u64>() {
*b
} else {
return Err(format!(
"unsupported flags integer type {}",
inner.reflect_type_path()
));
};
Ok(EditValue::Flags {
type_path: type_path.into(),
bits,
})
}
fn read_expr_handle(reflect: &dyn Reflect, name: &str, optional: bool) -> Option<ExprHandle> {
let ReflectRef::Struct(s) = reflect.reflect_ref() else {
return None;
};
let field = s.field(name)?;
if optional {
return field
.try_downcast_ref::<Option<ExprHandle>>()
.copied()
.flatten();
}
field.try_downcast_ref::<ExprHandle>().copied()
}
fn literal_value(lit: &LiteralExpr) -> Option<Value> {
lit.field("value")?.try_downcast_ref::<Value>().copied()
}
fn property_handle(pe: &PropertyExpr) -> Option<PropertyHandle> {
pe.field("property")?
.try_downcast_ref::<PropertyHandle>()
.copied()
}
fn handle_value_type_default(module: &Module, handle: ExprHandle) -> Option<Value> {
use bevy_hanabi::{ScalarType, ValueType};
Some(match module.get(handle)?.value_type()? {
ValueType::Scalar(ScalarType::Float) => Value::from(0.0_f32),
ValueType::Scalar(ScalarType::Int) => Value::from(0_i32),
ValueType::Scalar(ScalarType::Uint) => Value::from(0_u32),
ValueType::Scalar(ScalarType::Bool) => Value::from(false),
ValueType::Vector(v) => match (v.elem_type(), v.count()) {
(ScalarType::Float, 2) => Value::from(Vec2::ZERO),
(ScalarType::Float, 3) => Value::from(Vec3::ZERO),
(ScalarType::Float, 4) => Value::from(Vec4::ZERO),
(ScalarType::Uint, 2) => Value::from(UVec2::ZERO),
_ => return None,
},
_ => return None,
})
}
fn expr_kind(expr: &Expr) -> &'static str {
match expr {
Expr::BuiltIn(_) => "built-in",
Expr::Literal(_) => "literal",
Expr::Property(_) => "property",
Expr::Attribute(_) => "attribute",
Expr::ParentAttribute(_) => "parent-attribute",
Expr::Unary { .. } => "unary-operator",
Expr::Binary { .. } => "binary-operator",
Expr::Ternary { .. } => "ternary-operator",
Expr::Cast(_) => "cast",
Expr::TextureSample(_) => "texture-sample",
}
}
#[cfg(test)]
mod tests {
use bevy::prelude::*;
use super::*;
use crate::{bake::bake, demo::demo_graph, modifier_registry::ModifierRegistryPlugin};
#[test]
fn import_round_trips_demo_bake() {
let mut app = App::new();
app.add_plugins((MinimalPlugins, AssetPlugin::default()));
app.add_plugins(ModifierRegistryPlugin);
let registry = app.world().resource::<AppTypeRegistry>().read();
let asset = bake(&demo_graph(), ®istry).expect("demo bakes");
drop(registry);
let (graph, _warnings) = import(&asset);
assert_eq!(&*graph.header.name, "demo");
assert_eq!(graph.header.capacity, 8192);
let names: Vec<&str> = graph.properties.iter().map(|p| &*p.name).collect();
assert!(names.contains(&"gravity"), "gravity property imported");
assert!(
names.contains(&"spawn_speed"),
"spawn_speed property imported"
);
assert!(graph.properties.iter().all(|p| p.exposed));
let count = |g: ModifierGroup| {
graph
.stacks
.iter()
.find(|s| s.group == g)
.map(|s| s.members.len())
.unwrap_or(0)
};
assert_eq!(count(ModifierGroup::Init), 3);
assert_eq!(count(ModifierGroup::Update), 1);
assert_eq!(count(ModifierGroup::Render), 6);
let prop_ref_nodes = graph
.nodes
.iter()
.filter(|n| matches!(&n.payload, NodePayload::Expr(ExprNode::Property(_))))
.count();
assert_eq!(prop_ref_nodes, 2, "two property reference nodes");
assert_eq!(graph.links.len(), 2, "two property links");
}
#[test]
fn imported_graph_rebakes() {
let mut app = App::new();
app.add_plugins((MinimalPlugins, AssetPlugin::default()));
app.add_plugins(ModifierRegistryPlugin);
let registry = app.world().resource::<AppTypeRegistry>().read();
let asset = bake(&demo_graph(), ®istry).expect("demo bakes");
let (graph, _) = import(&asset);
let rebaked = bake(&graph, ®istry).expect("imported graph rebakes");
assert_eq!(rebaked.init_modifiers().count(), 3);
assert_eq!(rebaked.update_modifiers().count(), 1);
assert_eq!(rebaked.render_modifiers().count(), 6);
}
}