use std::any::TypeId;
use std::collections::HashMap;
use std::sync::Arc;
use bevy::prelude::*;
use bevy::shader::Shader;
use serde::de::DeserializeOwned;
use serde_json::Value;
use ts_rs::TS;
use super::params::{ParamSlot, check_param_cap};
use crate::animations::ValueKind;
use crate::registry::{NamedEntry, register_entry};
use crate::ts_codegen::TsCollector;
pub trait ReactFilter: Send + Sync + Sized + 'static {
const NAME: &'static str;
const USES_TIME: bool = false;
fn identity_params() -> Option<Value> {
None
}
fn shader(assets: &AssetServer) -> Handle<Shader>;
fn outset(&self) -> Result<f32, String> {
Ok(0.0)
}
fn pack(&self) -> (Vec<Vec4>, Arc<[ParamSlot]>);
fn resolve(&self, assets: &AssetServer) -> Result<Vec<ResolvedFilterPass>, String> {
resolve_single_pass(self, assets)
}
}
pub fn resolve_single_pass<T: ReactFilter>(
filter: &T,
assets: &AssetServer,
) -> Result<Vec<ResolvedFilterPass>, String> {
let (params, layout) = filter.pack();
check_param_cap(T::NAME, params.len())?;
Ok(vec![ResolvedFilterPass {
shader: T::shader(assets),
params,
layout,
wire_index: 0,
}])
}
#[derive(Debug, Clone, PartialEq)]
pub struct ResolvedFilterPass {
pub shader: Handle<Shader>,
pub params: Vec<Vec4>,
pub layout: Arc<[ParamSlot]>,
pub wire_index: u8,
}
pub(super) fn rewrite_length_slots(pass: &mut ResolvedFilterPass, scale: f32) {
let layout = pass.layout.clone();
for slot in layout.iter().filter(|s| s.kind == ValueKind::Length) {
for comp in slot.comp..(slot.comp + slot.len).min(4) {
if let Some(vec) = pass.params.get_mut(slot.vec) {
vec[comp] *= scale;
}
}
}
}
pub(super) fn stamp_and_push(
passes: Vec<ResolvedFilterPass>,
wire_index: usize,
scale: f32,
out: &mut Vec<ResolvedFilterPass>,
) {
let wire_index = wire_index.min(u8::MAX as usize) as u8;
for mut pass in passes {
pass.wire_index = wire_index;
rewrite_length_slots(&mut pass, scale);
out.push(pass);
}
}
pub struct FilterRegistration {
type_id: TypeId,
pub(crate) resolve: fn(&Value, &AssetServer) -> Result<Vec<ResolvedFilterPass>, String>,
pub(crate) outset: fn(&Value) -> Result<f32, String>,
pub(crate) uses_time: bool,
pub(crate) identity: fn() -> Option<Value>,
pub(crate) ts_name: fn() -> String,
pub(crate) ts_collect: fn(&mut TsCollector),
}
impl NamedEntry for FilterRegistration {
fn type_id(&self) -> TypeId {
self.type_id
}
}
#[derive(Resource, Default)]
pub struct FilterRegistry {
pub(crate) entries: HashMap<&'static str, FilterRegistration>,
}
impl FilterRegistry {
pub fn register<T: ReactFilter + DeserializeOwned + TS>(&mut self) {
register_entry(
&mut self.entries,
T::NAME,
"filter",
FilterRegistration {
type_id: TypeId::of::<T>(),
resolve: |value, assets| {
let passes = decode_params::<T>(value)?.resolve(assets)?;
for pass in &passes {
check_param_cap(T::NAME, pass.params.len())?;
}
Ok(passes)
},
outset: |value| decode_params::<T>(value)?.outset(),
uses_time: T::USES_TIME,
identity: T::identity_params,
ts_name: <T as TS>::name,
ts_collect: |c| c.add::<T>(),
},
);
}
}
fn decode_params<T: ReactFilter + DeserializeOwned>(value: &Value) -> Result<T, String> {
T::deserialize(value).map_err(|e| format!("filter {:?} params: {e}", T::NAME))
}
#[cfg(test)]
mod tests {
use std::f32::consts::PI;
use serde::Deserialize;
use serde_json::json;
use super::super::test_util::asset_app;
use super::*;
use crate::filters::{
BlurParams, HueRotateParams, MAX_FILTER_PARAM_VECS, register_builtin_filters,
};
#[test]
fn over_cap_param_vecs_are_rejected() {
struct NineVecs;
impl ReactFilter for NineVecs {
const NAME: &'static str = "nineVecs";
fn shader(_assets: &AssetServer) -> Handle<Shader> {
Handle::default()
}
fn pack(&self) -> (Vec<Vec4>, Arc<[ParamSlot]>) {
(
vec![Vec4::ZERO; MAX_FILTER_PARAM_VECS + 1],
Arc::from(Vec::new()),
)
}
}
let app = asset_app();
let assets = app.world().resource::<AssetServer>();
let err = NineVecs.resolve(assets).expect_err("over cap must reject");
assert!(err.contains("nineVecs"), "error names the filter: {err}");
}
#[test]
fn registry_recheck_rejects_over_cap_custom_resolve() {
#[derive(Deserialize, ts_rs::TS)]
#[serde(deny_unknown_fields)]
struct SneakyResolve {}
impl ReactFilter for SneakyResolve {
const NAME: &'static str = "sneakyResolve";
fn shader(_assets: &AssetServer) -> Handle<Shader> {
Handle::default()
}
fn pack(&self) -> (Vec<Vec4>, Arc<[ParamSlot]>) {
(Vec::new(), Arc::from(Vec::new()))
}
fn resolve(&self, assets: &AssetServer) -> Result<Vec<ResolvedFilterPass>, String> {
Ok(vec![ResolvedFilterPass {
shader: Self::shader(assets),
params: vec![Vec4::ZERO; MAX_FILTER_PARAM_VECS + 1],
layout: Arc::from(Vec::new()),
wire_index: 0,
}])
}
}
let app = asset_app();
let assets = app.world().resource::<AssetServer>();
let mut registry = FilterRegistry::default();
registry.register::<SneakyResolve>();
let err = (registry.entries["sneakyResolve"].resolve)(&json!({}), assets)
.expect_err("over-cap custom resolve must reject");
assert!(
err.contains("sneakyResolve"),
"error names the filter: {err}"
);
}
#[test]
fn builtin_filters_register_all_ten() {
let mut app = App::new();
register_builtin_filters(&mut app);
let registry = app.world().resource::<FilterRegistry>();
let mut names: Vec<_> = registry.entries.keys().copied().collect();
names.sort_unstable();
assert_eq!(
names,
[
"bloom",
"blur",
"brightness",
"chromaticAberration",
"contrast",
"grayscale",
"hueRotate",
"invert",
"saturate",
"sepia",
]
);
assert!(registry.entries.values().all(|r| !r.uses_time));
register_builtin_filters(&mut app);
assert_eq!(app.world().resource::<FilterRegistry>().entries.len(), 10);
}
#[test]
fn builtin_filters_have_working_ts_slots() {
let mut app = App::new();
register_builtin_filters(&mut app);
let registry = app.world().resource::<FilterRegistry>();
for (name, reg) in ®istry.entries {
let ts = (reg.ts_name)();
let mut c = TsCollector::default();
(reg.ts_collect)(&mut c);
assert!(c.decls.contains_key(&ts), "{name}: no decl for {ts}");
}
let blur = ®istry.entries["blur"];
assert_eq!((blur.ts_name)(), "BlurParams");
let mut c = TsCollector::default();
(blur.ts_collect)(&mut c);
let decl = &c.decls["BlurParams"];
assert!(decl.contains("radius"), "{decl}");
assert!(decl.contains("number | string"), "{decl}");
assert_eq!((registry.entries["grayscale"].ts_name)(), "GrayscaleParams");
assert_eq!((registry.entries["hueRotate"].ts_name)(), "HueRotateParams");
let mut c = TsCollector::default();
(registry.entries["hueRotate"].ts_collect)(&mut c);
assert!(
c.decls["HueRotateParams"].contains("angle: number | string"),
"{}",
c.decls["HueRotateParams"]
);
}
#[test]
fn registry_resolve_end_to_end() {
let app = asset_app();
let assets = app.world().resource::<AssetServer>();
let mut registry = FilterRegistry::default();
registry.register::<BlurParams>();
registry.register::<HueRotateParams>();
let blur = ®istry.entries["blur"];
let passes = (blur.resolve)(&json!({ "radius": 8 }), assets).expect("blur resolves");
assert_eq!(passes.len(), 2);
assert_eq!(passes[0].params[0].x, 8.0);
assert_eq!(
(blur.outset)(&json!({ "radius": 8 })).expect("outset"),
24.0
);
let hue = ®istry.entries["hueRotate"];
let passes = (hue.resolve)(&json!({ "angle": "0.5turn" }), assets).expect("hue resolves");
assert_eq!(passes.len(), 1);
assert!((passes[0].params[1].z - PI).abs() < 1e-4);
assert_eq!((hue.outset)(&json!({})).expect("outset"), 0.0);
}
#[test]
fn resolve_returns_embedded_shader_handles() {
let mut app = asset_app();
register_builtin_filters(&mut app);
let world = app.world();
let assets = world.resource::<AssetServer>();
let registry = world.resource::<FilterRegistry>();
let shader_of = |name: &str| {
let passes =
(registry.entries[name].resolve)(&json!({}), assets).expect("filter resolves");
let first = passes[0].shader.clone();
assert!(
passes.iter().all(|p| p.shader == first),
"all of {name}'s passes share one shader"
);
assert_ne!(first, Handle::default(), "{name} has a real shader");
first
};
let path_of = |handle: &Handle<Shader>| handle.path().expect("embedded path").to_string();
let color = shader_of("brightness");
for name in [
"contrast",
"saturate",
"grayscale",
"sepia",
"invert",
"hueRotate",
] {
assert_eq!(shader_of(name), color, "{name} shares the color shader");
}
assert_eq!(
&path_of(&color),
"embedded://bevy_react/filters/builtin/color_matrix.wgsl"
);
let blur = shader_of("blur");
assert_ne!(blur, color);
assert_eq!(
&path_of(&blur),
"embedded://bevy_react/filters/builtin/blur.wgsl"
);
assert_eq!(
&path_of(&shader_of("chromaticAberration")),
"embedded://bevy_react/filters/builtin/chromatic_aberration.wgsl"
);
let passes =
(registry.entries["bloom"].resolve)(&json!({}), assets).expect("bloom resolves");
assert_eq!(passes.len(), 4);
assert_eq!(
&path_of(&passes[0].shader),
"embedded://bevy_react/filters/builtin/bloom.wgsl"
);
assert_eq!(passes[3].shader, passes[0].shader);
assert_eq!(passes[1].shader, blur, "middle passes reuse blur's shader");
assert_eq!(passes[2].shader, blur);
}
}