#![cfg_attr(docsrs, feature(doc_cfg))]
use proc_macro::TokenStream;
use quote::quote;
use syn::{DeriveInput, LitStr, Type, parse_macro_input};
#[proc_macro_attribute]
pub fn react_message(attr: TokenStream, item: TokenStream) -> TokenStream {
let name_override = match parse_name_only_attr(attr, "react_message") {
Ok(name) => name,
Err(e) => return e.to_compile_error().into(),
};
let input = parse_macro_input!(item as DeriveInput);
let PayloadParts {
ident,
impl_generics,
ty_generics,
where_clause,
name,
} = payload_parts(&input, name_override);
quote! {
#[derive(::serde::Deserialize, ::ts_rs::TS)]
#input
impl #impl_generics ::bevy::ecs::event::Event for #ident #ty_generics #where_clause {
type Trigger<'a> = ::bevy::ecs::event::GlobalTrigger;
}
impl #impl_generics ::bevy_react::ReactPayload for #ident #ty_generics #where_clause {
const NAME: &'static str = #name;
}
}
.into()
}
#[proc_macro_attribute]
pub fn react_request(attr: TokenStream, item: TokenStream) -> TokenStream {
let mut name_override: Option<String> = None;
let mut response: Option<Type> = None;
let arg_parser = syn::meta::parser(|meta| {
if try_parse_name_arg(&meta, &mut name_override)? {
Ok(())
} else if meta.path.is_ident("response") {
response = Some(meta.value()?.parse::<Type>()?);
Ok(())
} else {
Err(meta.error(
"unsupported `react_request` argument; expected `name = \"...\"` or `response = Type`",
))
}
});
parse_macro_input!(attr with arg_parser);
let response = match response {
Some(ty) => ty,
None => {
return syn::Error::new(
proc_macro2::Span::call_site(),
"`react_request` requires a `response = Type` argument",
)
.to_compile_error()
.into();
}
};
let input = parse_macro_input!(item as DeriveInput);
let PayloadParts {
ident,
impl_generics,
ty_generics,
where_clause,
name,
} = payload_parts(&input, name_override);
quote! {
#[derive(::serde::Deserialize, ::ts_rs::TS)]
#input
impl #impl_generics ::bevy_react::ReactRequest for #ident #ty_generics #where_clause {
const NAME: &'static str = #name;
type Response = #response;
}
}
.into()
}
#[proc_macro_attribute]
pub fn react_event(attr: TokenStream, item: TokenStream) -> TokenStream {
let name_override = match parse_name_only_attr(attr, "react_event") {
Ok(name) => name,
Err(e) => return e.to_compile_error().into(),
};
let input = parse_macro_input!(item as DeriveInput);
let PayloadParts {
ident,
impl_generics,
ty_generics,
where_clause,
name,
} = payload_parts(&input, name_override);
quote! {
#[derive(::serde::Serialize, ::ts_rs::TS)]
#input
impl #impl_generics ::bevy_react::ReactEvent for #ident #ty_generics #where_clause {
const NAME: &'static str = #name;
}
}
.into()
}
#[proc_macro_attribute]
pub fn react_filter(attr: TokenStream, item: TokenStream) -> TokenStream {
let mut name_override: Option<String> = None;
let mut shader: Option<LitStr> = None;
let mut outset: f32 = 0.0;
let mut time = false;
let arg_parser = syn::meta::parser(|meta| {
if try_parse_name_arg(&meta, &mut name_override)? {
Ok(())
} else if meta.path.is_ident("shader") {
shader = Some(meta.value()?.parse::<LitStr>()?);
Ok(())
} else if meta.path.is_ident("outset") {
outset = match meta.value()?.parse::<syn::Lit>()? {
syn::Lit::Float(f) => f.base10_parse()?,
syn::Lit::Int(i) => i.base10_parse()?,
other => {
return Err(syn::Error::new_spanned(
other,
"`outset` must be a number literal (logical px)",
));
}
};
Ok(())
} else if meta.path.is_ident("time") {
time = meta.value()?.parse::<syn::LitBool>()?.value;
Ok(())
} else {
Err(meta.error(
"unsupported `react_filter` argument; expected `name = \"...\"`, \
`shader = \"...\"`, `outset = <number>`, or `time = <bool>`",
))
}
});
parse_macro_input!(attr with arg_parser);
let Some(shader) = shader else {
return syn::Error::new(
proc_macro2::Span::call_site(),
"`react_filter` requires a `shader = \"path/to.wgsl\"` argument",
)
.to_compile_error()
.into();
};
let mut input = parse_macro_input!(item as DeriveInput);
let name = name_override.unwrap_or_else(|| lower_first(&input.ident.to_string()));
let syn::Data::Struct(data) = &mut input.data else {
return syn::Error::new_spanned(&input.ident, "`react_filter` requires a struct")
.to_compile_error()
.into();
};
let syn::Fields::Named(fields) = &mut data.fields else {
return syn::Error::new_spanned(
&input.ident,
"`react_filter` requires named fields (each field becomes a shader param)",
)
.to_compile_error()
.into();
};
let mut ts_overrides: Vec<(usize, &'static str)> = Vec::new();
let mut errors: Vec<proc_macro2::TokenStream> = Vec::new();
let mut slots: Vec<proc_macro2::TokenStream> = Vec::new();
let mut writes: Vec<proc_macro2::TokenStream> = Vec::new();
let mut length_checks: Vec<proc_macro2::TokenStream> = Vec::new();
let mut vec_i = 0usize;
let mut comp = 0usize;
for (field_i, field) in fields.named.iter().enumerate() {
let ident = field.ident.clone().expect("named field");
let field_name = ident.to_string();
let Some(param) = classify_filter_field(&field.ty) else {
errors.push(
syn::Error::new_spanned(
&field.ty,
format!(
"`react_filter` cannot pack field `{field_name}`: supported param types \
are f32, Vec2, Vec3, Vec4, [f32; 2..=4], Angle, Length, and FilterColor"
),
)
.to_compile_error(),
);
continue;
};
let len = param.len();
if comp + len > 4 {
vec_i += 1;
comp = 0;
}
let (v, c) = (vec_i, comp);
let kind = param.value_kind();
slots.push(quote! {
::bevy_react::filters::ParamSlot {
name: #field_name,
kind: #kind,
vec: #v,
comp: #c,
len: #len,
}
});
match ¶m {
FilterField::Scalar => writes.push(quote! { params[#v][#c] = self.#ident; }),
FilterField::Vector(n) => {
for (i, axis) in ["x", "y", "z", "w"].iter().take(*n).enumerate() {
let axis = syn::Ident::new(axis, proc_macro2::Span::call_site());
let ci = c + i;
writes.push(quote! { params[#v][#ci] = self.#ident.#axis; });
}
}
FilterField::Array(n) => {
for i in 0..*n {
let ci = c + i;
writes.push(quote! { params[#v][#ci] = self.#ident[#i]; });
}
}
FilterField::Angle => {
writes.push(quote! { params[#v][#c] = self.#ident.radians(); });
}
FilterField::Length => {
writes.push(quote! {
params[#v][#c] = ::bevy_react::filters::length_logical_px(
#name, #field_name, self.#ident,
)
.unwrap_or(0.0);
});
length_checks.push(quote! {
::bevy_react::filters::length_logical_px(#name, #field_name, self.#ident)?;
});
}
FilterField::Color => {
for i in 0..4usize {
let ci = c + i;
writes.push(quote! { params[#v][#ci] = self.#ident.0[#i]; });
}
}
}
comp += len;
if comp == 4 {
vec_i += 1;
comp = 0;
}
if let Some(ts) = param.ts_override() {
ts_overrides.push((field_i, ts));
}
}
if errors.is_empty() {
for (field_i, ts) in ts_overrides {
fields.named[field_i]
.attrs
.push(syn::parse_quote!(#[ts(type = #ts)]));
}
} else {
return quote! { #input #(#errors)* }.into();
}
let total_vecs = if comp == 0 { vec_i } else { vec_i + 1 };
let ident = &input.ident;
let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
let resolve_override = (!length_checks.is_empty()).then(|| {
quote! {
fn resolve(
&self,
assets: &::bevy::asset::AssetServer,
) -> ::std::result::Result<
::std::vec::Vec<::bevy_react::filters::ResolvedFilterPass>,
::std::string::String,
> {
#(#length_checks)*
::bevy_react::filters::resolve_single_pass(self, assets)
}
}
});
quote! {
#[derive(::serde::Deserialize, ::ts_rs::TS)]
#[serde(deny_unknown_fields)]
#input
impl #impl_generics ::bevy_react::filters::ReactFilter for #ident #ty_generics #where_clause {
const NAME: &'static str = #name;
const USES_TIME: bool = #time;
fn shader(
assets: &::bevy::asset::AssetServer,
) -> ::bevy::asset::Handle<::bevy::shader::Shader> {
assets.load(#shader)
}
fn outset(&self) -> ::std::result::Result<f32, ::std::string::String> {
#(#length_checks)*
::std::result::Result::Ok(#outset)
}
fn pack(
&self,
) -> (
::std::vec::Vec<::bevy::math::Vec4>,
::std::sync::Arc<[::bevy_react::filters::ParamSlot]>,
) {
static LAYOUT: ::std::sync::LazyLock<
::std::sync::Arc<[::bevy_react::filters::ParamSlot]>,
> = ::std::sync::LazyLock::new(|| {
::std::sync::Arc::from(::std::vec![#(#slots),*])
});
#[allow(unused_mut)]
let mut params = ::std::vec![::bevy::math::Vec4::ZERO; #total_vecs];
#(#writes)*
(params, LAYOUT.clone())
}
#resolve_override
}
}
.into()
}
enum FilterField {
Scalar,
Vector(usize),
Array(usize),
Angle,
Length,
Color,
}
impl FilterField {
fn len(&self) -> usize {
match self {
Self::Scalar | Self::Angle | Self::Length => 1,
Self::Vector(n) | Self::Array(n) => *n,
Self::Color => 4,
}
}
fn value_kind(&self) -> proc_macro2::TokenStream {
match self {
Self::Scalar | Self::Vector(_) | Self::Array(_) => {
quote!(::bevy_react::animations::ValueKind::Scalar)
}
Self::Angle => quote!(::bevy_react::animations::ValueKind::Angle),
Self::Length => quote!(::bevy_react::animations::ValueKind::Length),
Self::Color => quote!(::bevy_react::animations::ValueKind::Color),
}
}
fn ts_override(&self) -> Option<&'static str> {
match self {
Self::Vector(2) => Some("[number, number]"),
Self::Vector(3) => Some("[number, number, number]"),
Self::Vector(4) => Some("[number, number, number, number]"),
Self::Angle | Self::Length => Some("number | string"),
_ => None,
}
}
}
fn classify_filter_field(ty: &Type) -> Option<FilterField> {
match ty {
Type::Path(p) => {
let seg = p.path.segments.last()?;
if !seg.arguments.is_empty() {
return None;
}
match seg.ident.to_string().as_str() {
"f32" => Some(FilterField::Scalar),
"Vec2" => Some(FilterField::Vector(2)),
"Vec3" => Some(FilterField::Vector(3)),
"Vec4" => Some(FilterField::Vector(4)),
"Angle" => Some(FilterField::Angle),
"Length" => Some(FilterField::Length),
"FilterColor" => Some(FilterField::Color),
_ => None,
}
}
Type::Array(a) => {
let is_f32 = matches!(&*a.elem, Type::Path(p) if p.path.is_ident("f32"));
let syn::Expr::Lit(lit) = &a.len else {
return None;
};
let syn::Lit::Int(n) = &lit.lit else {
return None;
};
let n = n.base10_parse::<usize>().ok()?;
(is_f32 && (2..=4).contains(&n)).then_some(FilterField::Array(n))
}
_ => None,
}
}
fn try_parse_name_arg(
meta: &syn::meta::ParseNestedMeta,
out: &mut Option<String>,
) -> syn::Result<bool> {
if meta.path.is_ident("name") {
*out = Some(meta.value()?.parse::<LitStr>()?.value());
Ok(true)
} else {
Ok(false)
}
}
fn parse_name_only_attr(attr: TokenStream, macro_name: &str) -> syn::Result<Option<String>> {
let mut name_override: Option<String> = None;
let parser = syn::meta::parser(|meta| {
if try_parse_name_arg(&meta, &mut name_override)? {
Ok(())
} else {
Err(meta.error(format!(
"unsupported `{macro_name}` argument; expected `name = \"...\"`"
)))
}
});
syn::parse::Parser::parse(parser, attr)?;
Ok(name_override)
}
struct PayloadParts<'a> {
ident: &'a syn::Ident,
impl_generics: syn::ImplGenerics<'a>,
ty_generics: syn::TypeGenerics<'a>,
where_clause: Option<&'a syn::WhereClause>,
name: String,
}
fn payload_parts<'a>(input: &'a DeriveInput, name_override: Option<String>) -> PayloadParts<'a> {
let ident = &input.ident;
let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
PayloadParts {
ident,
impl_generics,
ty_generics,
where_clause,
name: name_override.unwrap_or_else(|| lower_first(&ident.to_string())),
}
}
fn lower_first(s: &str) -> String {
let mut chars = s.chars();
match chars.next() {
Some(first) => first.to_lowercase().collect::<String>() + chars.as_str(),
None => String::new(),
}
}