mirror of
https://github.com/GraphiteEditor/Graphite.git
synced 2026-09-15 22:28:10 +08:00
2213 lines
72 KiB
Rust
2213 lines
72 KiB
Rust
use convert_case::{Case, Casing};
|
|
use indoc::{formatdoc, indoc};
|
|
use proc_macro2::TokenStream as TokenStream2;
|
|
use quote::{ToTokens, format_ident, quote};
|
|
use syn::parse::{Parse, ParseStream, Parser};
|
|
use syn::punctuated::Punctuated;
|
|
use syn::spanned::Spanned;
|
|
use syn::token::{Comma, RArrow};
|
|
use syn::{
|
|
AttrStyle, Attribute, Error, Expr, FnArg, GenericArgument, GenericParam, Ident, ItemFn, Lit, LitFloat, LitInt, LitStr, Meta, Pat, PatIdent, PatType, Path, PathArguments, ReturnType, TraitBound,
|
|
Type, TypeImplTrait, TypeParam, TypeParamBound, Visibility, WhereClause, parse_quote,
|
|
};
|
|
|
|
use crate::codegen::generate_node_code;
|
|
use crate::crate_ident::CrateIdent;
|
|
use crate::shader_nodes::ShaderNodeType;
|
|
|
|
#[derive(Clone, Debug)]
|
|
pub(crate) struct Implementation {
|
|
pub(crate) input: Type,
|
|
pub(crate) _arrow: RArrow,
|
|
pub(crate) output: Type,
|
|
}
|
|
|
|
#[derive(Debug)]
|
|
pub(crate) struct ParsedNodeFn {
|
|
pub(crate) vis: Visibility,
|
|
pub(crate) attributes: NodeFnAttributes,
|
|
pub(crate) fn_name: Ident,
|
|
pub(crate) struct_name: Ident,
|
|
pub(crate) mod_name: Ident,
|
|
pub(crate) fn_generics: Vec<GenericParam>,
|
|
pub(crate) where_clause: Option<WhereClause>,
|
|
pub(crate) input: Input,
|
|
pub(crate) output_type: Type,
|
|
pub(crate) output_depth: u8,
|
|
pub(crate) is_async: bool,
|
|
pub(crate) fields: Vec<ParsedField>,
|
|
/// The caller's frame claim, declared by a record-opaque kernel that
|
|
/// serves through it; not a wired input.
|
|
pub(crate) claim: Option<PatType>,
|
|
pub(crate) body: TokenStream2,
|
|
pub(crate) description: String,
|
|
}
|
|
|
|
/// An `Attr<Marker>` slot in a parameter's read tuple: a declared attribute
|
|
/// read on that input, not a wired input of its own.
|
|
#[derive(Clone, Debug)]
|
|
pub(crate) struct AttributeRead {
|
|
pub(crate) pat_ident: PatIdent,
|
|
pub(crate) marker: Type,
|
|
}
|
|
|
|
/// One attribute write slot: the marker, and whether it crosses as an owned
|
|
/// copy (`OwnedAttr<M>`) instead of an evaluation-lifetime value (`Attr<M>`).
|
|
pub(crate) struct AttrWrite {
|
|
pub(crate) marker: Type,
|
|
pub(crate) owned: bool,
|
|
}
|
|
|
|
/// The write half of a record kernel's return: the element type in the first
|
|
/// tuple slot, then the attribute markers written and the ones removed. `None`
|
|
/// unless the value is a well-formed write tuple (a non-marker element first,
|
|
/// then only `Attr`, `OwnedAttr` and `RemoveAttr` slots, at least one).
|
|
pub(crate) struct RecordWrites {
|
|
pub(crate) element: Type,
|
|
pub(crate) markers: Vec<AttrWrite>,
|
|
pub(crate) removes: Vec<Type>,
|
|
}
|
|
|
|
pub(crate) fn record_writes(value: &Type) -> Option<RecordWrites> {
|
|
let Type::Tuple(tuple) = value else { return None };
|
|
let mut slots = tuple.elems.iter();
|
|
let element = slots.next()?;
|
|
if attr_marker(element).is_some() || owned_attr_marker(element).is_some() || remove_attr_marker(element).is_some() {
|
|
return None;
|
|
}
|
|
let mut markers = Vec::new();
|
|
let mut removes = Vec::new();
|
|
for slot in slots {
|
|
if let Some(marker) = attr_marker(slot) {
|
|
markers.push(AttrWrite { marker, owned: false });
|
|
} else if let Some(marker) = owned_attr_marker(slot) {
|
|
markers.push(AttrWrite { marker, owned: true });
|
|
} else if let Some(marker) = remove_attr_marker(slot) {
|
|
removes.push(marker);
|
|
} else {
|
|
return None;
|
|
}
|
|
}
|
|
(!markers.is_empty() || !removes.is_empty()).then(|| RecordWrites {
|
|
element: element.clone(),
|
|
markers,
|
|
removes,
|
|
})
|
|
}
|
|
|
|
/// Returns the marker type of an `Attr<Marker>` type, if `ty` is one.
|
|
pub(crate) fn attr_marker(ty: &Type) -> Option<Type> {
|
|
marker_of(ty, "Attr")
|
|
}
|
|
|
|
/// Returns the marker type of an `OwnedAttr<Marker>` type, if `ty` is one.
|
|
pub(crate) fn owned_attr_marker(ty: &Type) -> Option<Type> {
|
|
marker_of(ty, "OwnedAttr")
|
|
}
|
|
|
|
/// Returns the marker type of a `RemoveAttr<Marker>` type, if `ty` is one.
|
|
pub(crate) fn remove_attr_marker(ty: &Type) -> Option<Type> {
|
|
marker_of(ty, "RemoveAttr")
|
|
}
|
|
|
|
/// Splits a `Named<X, V>` marker into its placeholder and value type. A write
|
|
/// of one takes its name from the input the placeholder is declared at rather
|
|
/// than from the marker, so the name folds at graph compile time.
|
|
pub(crate) fn named_marker(ty: &Type) -> Option<(Type, Type)> {
|
|
let mut args = named_arguments(ty)?.into_iter();
|
|
let (placeholder, value) = (args.next()?, args.next()?);
|
|
args.next().is_none().then_some((placeholder, value))
|
|
}
|
|
|
|
/// The placeholder a `Named<X>` parameter declares. Such a parameter is the
|
|
/// name source for every `Attr<Named<X, _>>` the signature writes, and crosses
|
|
/// the wire as constant text.
|
|
pub(crate) fn named_source(ty: &Type) -> Option<Type> {
|
|
let mut args = named_arguments(ty)?.into_iter();
|
|
let placeholder = args.next()?;
|
|
args.next().is_none().then_some(placeholder)
|
|
}
|
|
|
|
fn named_arguments(ty: &Type) -> Option<Vec<Type>> {
|
|
let Type::Path(path) = ty else { return None };
|
|
let segment = path.path.segments.last()?;
|
|
if segment.ident != "Named" {
|
|
return None;
|
|
}
|
|
let PathArguments::AngleBracketed(args) = &segment.arguments else { return None };
|
|
Some(
|
|
args.args
|
|
.iter()
|
|
.filter_map(|argument| match argument {
|
|
GenericArgument::Type(ty) => Some(ty.clone()),
|
|
_ => None,
|
|
})
|
|
.collect(),
|
|
)
|
|
}
|
|
|
|
fn marker_of(ty: &Type, wrapper: &str) -> Option<Type> {
|
|
let Type::Path(path) = ty else { return None };
|
|
let segment = path.path.segments.last()?;
|
|
if segment.ident != wrapper {
|
|
return None;
|
|
}
|
|
let PathArguments::AngleBracketed(args) = &segment.arguments else { return None };
|
|
let mut types = args.args.iter().filter_map(|argument| match argument {
|
|
GenericArgument::Type(ty) => Some(ty),
|
|
_ => None,
|
|
});
|
|
let marker = types.next()?;
|
|
types.next().is_none().then(|| marker.clone())
|
|
}
|
|
|
|
#[derive(Debug, Default, Clone)]
|
|
pub(crate) struct NodeFnAttributes {
|
|
pub(crate) category: Option<LitStr>,
|
|
pub(crate) display_name: Option<LitStr>,
|
|
pub(crate) path: Option<Path>,
|
|
pub(crate) skip_impl: bool,
|
|
pub(crate) properties_string: Option<LitStr>,
|
|
/// whether to `#[cfg]` gate the node implementation, defaults to None
|
|
pub(crate) cfg: Option<TokenStream2>,
|
|
/// if this node should get a gpu implementation, defaults to None
|
|
pub(crate) shader_node: Option<ShaderNodeType>,
|
|
/// Custom serialization function path (e.g., "my_module::custom_serialize")
|
|
pub(crate) serialize: Option<Path>,
|
|
/// Whether the preprocessor should add a Memoize node after this node in the generated subnetwork
|
|
pub(crate) memoize: bool,
|
|
/// Whether this node provides a scope
|
|
pub(crate) inject_scope: bool,
|
|
/// Function producing a stand-in value while an async source node's real value is in flight
|
|
pub(crate) placeholder: Option<Path>,
|
|
/// Function overriding the generated `extent` method
|
|
pub(crate) extent: Option<Path>,
|
|
/// Function overriding the generated `extent` method with the raw node/ctx/level form
|
|
pub(crate) extent_raw: Option<Path>,
|
|
/// Function overriding the generated `eval_batch` method
|
|
pub(crate) batch: Option<Path>,
|
|
/// Whether partial upstream values are mapped to `Pending` instead of flowing into this node
|
|
pub(crate) no_partial: bool,
|
|
/// Whether this node keeps the plain-input lowering during the record transition
|
|
pub(crate) plain: bool,
|
|
}
|
|
|
|
#[derive(Clone, Debug, Default)]
|
|
pub enum ParsedValueSource {
|
|
#[default]
|
|
None,
|
|
Default(TokenStream2),
|
|
Scope(Box<Expr>),
|
|
SourceId,
|
|
}
|
|
|
|
// #[widget(ParsedWidgetOverride::Hidden)]
|
|
// #[widget(ParsedWidgetOverride::String = "Some string")]
|
|
// #[widget(ParsedWidgetOverride::Custom = "Custom string")]
|
|
#[derive(Clone, Debug, Default)]
|
|
pub enum ParsedWidgetOverride {
|
|
#[default]
|
|
None,
|
|
Hidden,
|
|
String(LitStr),
|
|
Custom(LitStr),
|
|
}
|
|
|
|
impl Parse for ParsedWidgetOverride {
|
|
fn parse(input: ParseStream) -> syn::Result<Self> {
|
|
// Parse the full path (e.g., ParsedWidgetOverride::Hidden)
|
|
let path: Path = input.parse()?;
|
|
|
|
// Ensure the path starts with `ParsedWidgetOverride`
|
|
if path.segments.len() == 2 && path.segments[0].ident == "ParsedWidgetOverride" {
|
|
let variant = &path.segments[1].ident;
|
|
|
|
match variant.to_string().as_str() {
|
|
"Hidden" => Ok(ParsedWidgetOverride::Hidden),
|
|
"String" => {
|
|
input.parse::<syn::Token![=]>()?;
|
|
let lit: LitStr = input.parse()?;
|
|
Ok(ParsedWidgetOverride::String(lit))
|
|
}
|
|
"Custom" => {
|
|
input.parse::<syn::Token![=]>()?;
|
|
let lit: LitStr = input.parse()?;
|
|
Ok(ParsedWidgetOverride::Custom(lit))
|
|
}
|
|
_ => Err(Error::new(variant.span(), "Unknown ParsedWidgetOverride variant")),
|
|
}
|
|
} else {
|
|
Err(Error::new(input.span(), "Expected ParsedWidgetOverride::<variant>"))
|
|
}
|
|
}
|
|
}
|
|
|
|
#[derive(Clone, Debug)]
|
|
pub struct ParsedField {
|
|
pub pat_ident: PatIdent,
|
|
pub name: Option<LitStr>,
|
|
pub description: String,
|
|
pub widget_override: ParsedWidgetOverride,
|
|
pub ty: ParsedFieldType,
|
|
pub number_display_decimal_places: Option<LitInt>,
|
|
pub number_step: Option<LitFloat>,
|
|
pub unit: Option<LitStr>,
|
|
pub is_data_field: bool,
|
|
/// The attribute reads destructured from this input's tuple, resolved
|
|
/// against this input.
|
|
pub(crate) attribute_reads: Vec<AttributeRead>,
|
|
}
|
|
|
|
// Both variants are large parsed-syntax payloads (888/672 bytes), so boxing one still leaves the other large while forcing a
|
|
// deref on every pattern match across codegen; this is built once per node at compile time, never on a hot path
|
|
#[allow(clippy::large_enum_variant)]
|
|
#[derive(Clone, Debug)]
|
|
pub enum ParsedFieldType {
|
|
Regular(RegularParsedField),
|
|
Node(NodeParsedField),
|
|
}
|
|
|
|
/// A single numeric endpoint within a `#[soft(..)]` or `#[hard(..)]` bounds range.
|
|
/// Accepts both integer literals (e.g. `1`, `-1`) and float literals (e.g. `1.`, `-500.`).
|
|
#[derive(Clone, Debug)]
|
|
pub struct NumberBound {
|
|
is_negative: bool,
|
|
literal: NumberBoundLiteral,
|
|
}
|
|
|
|
#[derive(Clone, Debug)]
|
|
enum NumberBoundLiteral {
|
|
Float(LitFloat),
|
|
Int(LitInt),
|
|
}
|
|
|
|
impl NumberBound {
|
|
pub fn to_f64(&self) -> f64 {
|
|
let magnitude = match &self.literal {
|
|
NumberBoundLiteral::Float(lit) => lit.base10_parse::<f64>().unwrap_or_default(),
|
|
NumberBoundLiteral::Int(lit) => lit.base10_parse::<u64>().unwrap_or_default() as f64,
|
|
};
|
|
if self.is_negative { -magnitude } else { magnitude }
|
|
}
|
|
}
|
|
|
|
impl Parse for NumberBound {
|
|
fn parse(input: ParseStream) -> syn::Result<Self> {
|
|
let is_negative = input.peek(syn::Token![-]);
|
|
if is_negative {
|
|
let _: syn::Token![-] = input.parse()?;
|
|
}
|
|
|
|
let literal = if input.peek(LitFloat) {
|
|
NumberBoundLiteral::Float(input.parse()?)
|
|
} else if input.peek(LitInt) {
|
|
NumberBoundLiteral::Int(input.parse()?)
|
|
} else {
|
|
return Err(input.error("expected a numeric literal (integer or float)"));
|
|
};
|
|
|
|
Ok(NumberBound { is_negative, literal })
|
|
}
|
|
}
|
|
|
|
impl ToTokens for NumberBound {
|
|
fn to_tokens(&self, stream: &mut TokenStream2) {
|
|
match (&self.literal, self.is_negative) {
|
|
(NumberBoundLiteral::Float(lit), false) => lit.to_tokens(stream),
|
|
(NumberBoundLiteral::Float(lit), true) => stream.extend(quote!(-#lit)),
|
|
(NumberBoundLiteral::Int(lit), false) => stream.extend(quote!(#lit as f64)),
|
|
(NumberBoundLiteral::Int(lit), true) => stream.extend(quote!(-(#lit as f64))),
|
|
}
|
|
}
|
|
}
|
|
|
|
/// A pair of numeric bounds parsed from the `#[soft(a..b)]` and `#[hard(a..b)]` attributes.
|
|
/// Either endpoint may be omitted for an open-ended bound (`a..` or `..b`), and each endpoint
|
|
/// independently accepts an integer or float literal (each cast to `f64`), so a mixed range like
|
|
/// `0..3.14159` is valid.
|
|
///
|
|
/// The operator is always the bare `..`; both endpoints are treated as inclusive (clamping reaches them).
|
|
/// Unlike a Rust range there is no `..=` form, `..` is purely this attribute DSL's bounds operator.
|
|
#[derive(Clone, Debug)]
|
|
pub struct NumberRange {
|
|
start: Option<NumberBound>,
|
|
end: Option<NumberBound>,
|
|
}
|
|
|
|
impl Parse for NumberRange {
|
|
fn parse(input: ParseStream) -> syn::Result<Self> {
|
|
if input.is_empty() {
|
|
return Err(input.error("expected a range like `0..100`, `..100`, or `0..`"));
|
|
}
|
|
|
|
// A leading endpoint is present unless the range opens directly into the `..` operator.
|
|
let start = if input.peek(syn::Token![..=]) || input.peek(syn::Token![..]) {
|
|
None
|
|
} else {
|
|
Some(input.parse::<NumberBound>()?)
|
|
};
|
|
|
|
// Only the bare `..` is accepted. `..=` is rejected even though both endpoints are inclusive here:
|
|
// this DSL treats `..` as its own bounds operator, deliberately diverging from Rust's range semantics.
|
|
if input.peek(syn::Token![..=]) {
|
|
return Err(input.error("use `..` rather than `..=` for number bounds; both endpoints are always inclusive (e.g. `0..100`)"));
|
|
}
|
|
if !input.peek(syn::Token![..]) {
|
|
return Err(input.error("expected a range like `0..100`, `..100`, or `0..`"));
|
|
}
|
|
input.parse::<syn::Token![..]>()?;
|
|
|
|
let end = if input.is_empty() { None } else { Some(input.parse::<NumberBound>()?) };
|
|
|
|
if start.is_none() && end.is_none() {
|
|
return Err(input.error("a bounds range must specify at least a lower or upper bound"));
|
|
}
|
|
|
|
Ok(NumberRange { start, end })
|
|
}
|
|
}
|
|
|
|
/// a param of any kind, either a concrete type or a generic type with a set of possible types specified via
|
|
/// `#[implementation(type)]`
|
|
#[derive(Clone, Debug)]
|
|
pub struct RegularParsedField {
|
|
pub ty: Type,
|
|
/// The placeholder this parameter names, written `Named<X>`. Its `ty` is
|
|
/// rewritten to `String`, since the wire carries the name as constant text
|
|
/// while the kernel takes only the placeholder.
|
|
pub name_source: Option<Type>,
|
|
/// `IList` nesting stripped from `ty` at parse; `ty` holds the element row.
|
|
pub list_levels: u8,
|
|
/// The original reference tokens when the parameter was written `&T`; `ty` holds the peeled inner type.
|
|
pub lend: Option<syn::TypeReference>,
|
|
pub exposed: bool,
|
|
pub value_source: ParsedValueSource,
|
|
pub number_soft_min: Option<NumberBound>,
|
|
pub number_soft_max: Option<NumberBound>,
|
|
pub number_hard_min: Option<NumberBound>,
|
|
pub number_hard_max: Option<NumberBound>,
|
|
/// Whether the number input renders as a draggable slider (the `#[range]` attribute) rather than the default increment field.
|
|
pub number_mode_range: bool,
|
|
pub implementations: Punctuated<Type, Comma>,
|
|
pub gpu_image: bool,
|
|
}
|
|
|
|
/// a param of `impl Node` with `#[implementation(in -> out)]`
|
|
#[derive(Clone, Debug)]
|
|
pub struct NodeParsedField {
|
|
pub input_type: Type,
|
|
pub output_type: Type,
|
|
pub implementations: Punctuated<Implementation, Comma>,
|
|
}
|
|
|
|
#[derive(Clone, Debug)]
|
|
pub(crate) struct Input {
|
|
pub(crate) pat_ident: PatIdent,
|
|
pub(crate) ty: Type,
|
|
pub(crate) implementations: Punctuated<Type, Comma>,
|
|
pub(crate) context_features: Vec<ContextFeatureDecl>,
|
|
}
|
|
|
|
impl Parse for Implementation {
|
|
fn parse(input: ParseStream) -> syn::Result<Self> {
|
|
let input_type: Type = input.parse().map_err(|e| {
|
|
Error::new(
|
|
input.span(),
|
|
formatdoc!(
|
|
"Failed to parse input type for #[implementation(...)]. Expected a valid Rust type.
|
|
Error: {}",
|
|
e,
|
|
),
|
|
)
|
|
})?;
|
|
let arrow: RArrow = input.parse().map_err(|_| {
|
|
Error::new(
|
|
input.span(),
|
|
indoc!(
|
|
"Expected `->` arrow after input type in #[implementations(...)] on a field of type `impl Node`.
|
|
The correct syntax is `InputType -> OutputType`."
|
|
),
|
|
)
|
|
})?;
|
|
let output_type: Type = input.parse().map_err(|e| {
|
|
Error::new(
|
|
input.span(),
|
|
formatdoc!(
|
|
"Failed to parse output type for #[implementation(...)]. Expected a valid Rust type after `->`.
|
|
Error: {}",
|
|
e
|
|
),
|
|
)
|
|
})?;
|
|
|
|
Ok(Implementation {
|
|
input: input_type,
|
|
_arrow: arrow,
|
|
output: output_type,
|
|
})
|
|
}
|
|
}
|
|
|
|
impl Parse for NodeFnAttributes {
|
|
fn parse(input: ParseStream) -> syn::Result<Self> {
|
|
let mut category = None;
|
|
let mut display_name = None;
|
|
let mut path = None;
|
|
let mut skip_impl = false;
|
|
let mut properties_string = None;
|
|
let mut cfg = None;
|
|
let mut shader_node = None;
|
|
let mut serialize = None;
|
|
let mut memoize = false;
|
|
let mut inject_scope = false;
|
|
let mut placeholder = None;
|
|
let mut extent = None;
|
|
let mut extent_raw = None;
|
|
let mut batch = None;
|
|
let mut no_partial = false;
|
|
let mut plain = false;
|
|
|
|
let content = input;
|
|
// let content;
|
|
// syn::parenthesized!(content in input);
|
|
|
|
let nested = content.call(Punctuated::<Meta, Comma>::parse_terminated)?;
|
|
for meta in nested.iter() {
|
|
let name = meta.path().get_ident().ok_or_else(|| Error::new_spanned(meta.path(), "Node macro expects a known Ident, not a path"))?;
|
|
match name.to_string().as_str() {
|
|
// User-facing category in the node catalog. The empty string `category("")` hides the node from the catalog.
|
|
//
|
|
// Example usage:
|
|
// #[node_macro::node(..., category("Math: Arithmetic"), ...)]
|
|
"category" => {
|
|
let meta = meta.require_list()?;
|
|
if category.is_some() {
|
|
return Err(Error::new_spanned(meta, "Multiple 'category' attributes are not allowed"));
|
|
}
|
|
let lit: LitStr = meta
|
|
.parse_args()
|
|
.map_err(|_| Error::new_spanned(meta, "Expected a string literal for 'category', e.g., category(\"Value\")"))?;
|
|
category = Some(lit);
|
|
}
|
|
// Override for the display name in the node catalog in place of the auto-generated name taken from the function name with inferred Title Case formatting.
|
|
// Use this if capitalization or formatting needs to be overridden.
|
|
//
|
|
// Example usage:
|
|
// #[node_macro::node(..., name("Request URL"), ...)]
|
|
"name" => {
|
|
let meta = meta.require_list()?;
|
|
if display_name.is_some() {
|
|
return Err(Error::new_spanned(meta, "Multiple 'name' attributes are not allowed"));
|
|
}
|
|
let parsed_name: LitStr = meta.parse_args().map_err(|_| Error::new_spanned(meta, "Expected a string for 'name', e.g., name(\"Memoize\")"))?;
|
|
display_name = Some(parsed_name);
|
|
}
|
|
// Override for the fully qualified path used by Graphene to identify the node implementation.
|
|
// If not provided, the path will be inferred from the module path and function name.
|
|
// Use this if the node implementation has moved to a different module or crate but a migration to that new path is not desired.
|
|
//
|
|
// Example usage:
|
|
// #[node_macro::node(..., path(core_types::vector), ...)]
|
|
"path" => {
|
|
let meta = meta.require_list()?;
|
|
if path.is_some() {
|
|
return Err(Error::new_spanned(meta, "Multiple 'path' attributes are not allowed"));
|
|
}
|
|
let parsed_path: Path = meta
|
|
.parse_args()
|
|
.map_err(|_| Error::new_spanned(meta, "Expected a valid path for 'path', e.g., path(crate::MemoizeNode)"))?;
|
|
path = Some(parsed_path);
|
|
}
|
|
// Indicator that the node should allow generic type arguments but skip the automatic generation of concrete type implementations.
|
|
// It allows the type arguments in this node to not include the normally required `#[implementations(...)]` attribute on each generic parameter.
|
|
// Instead, concrete implementations must be manually listed in the Node Registry, or where impossible, produced at runtime by the compile server.
|
|
// This is used by a few advanced nodes that need to support many types where listing them all would be cumbersome or impossible.
|
|
//
|
|
// Example usage:
|
|
// #[node_macro::node(..., skip_impl, ...)]
|
|
"skip_impl" => {
|
|
let path = meta.require_path_only()?;
|
|
if skip_impl {
|
|
return Err(Error::new_spanned(path, "Multiple 'skip_impl' attributes are not allowed"));
|
|
}
|
|
skip_impl = true;
|
|
}
|
|
// Override UI layout generator function name defined in `node_properties.rs` that returns a custom Properties panel layout for this node.
|
|
// This is used to create custom UI for the input parameters of the node in cases where the defaults generated from the type and attributes are insufficient.
|
|
//
|
|
// Example usage:
|
|
// #[node_macro::node(..., properties("channel_mixer_properties"), ...)]
|
|
"properties" => {
|
|
let meta = meta.require_list()?;
|
|
if properties_string.is_some() {
|
|
return Err(Error::new_spanned(path, "Multiple 'properties' attributes are not allowed"));
|
|
}
|
|
let parsed_properties_string: LitStr = meta
|
|
.parse_args()
|
|
.map_err(|_| Error::new_spanned(meta, "Expected a string for 'properties', e.g., properties(\"channel_mixer_properties\")"))?;
|
|
|
|
properties_string = Some(parsed_properties_string);
|
|
}
|
|
// Conditional compilation tokens to gate when this node is included in the build.
|
|
//
|
|
// Example usage:
|
|
// #[node_macro::node(..., cfg(feature = "std"), ...)]
|
|
"cfg" => {
|
|
if cfg.is_some() {
|
|
return Err(Error::new_spanned(path, "Multiple 'cfg' attributes are not allowed"));
|
|
}
|
|
let meta = meta.require_list()?;
|
|
cfg = Some(meta.tokens.clone());
|
|
}
|
|
// Reference to a specific shader definition struct that is used to run the logic of this node on the GPU.
|
|
//
|
|
// Example usage:
|
|
// #[node_macro::node(..., shader_node(PerPixelAdjust), ...)]
|
|
"shader_node" => {
|
|
if shader_node.is_some() {
|
|
return Err(Error::new_spanned(path, "Multiple 'shader_node' attributes are not allowed"));
|
|
}
|
|
let meta = meta.require_list()?;
|
|
shader_node = Some(syn::parse2(meta.tokens.to_token_stream())?);
|
|
}
|
|
// Function name for custom serialization of this node's data. This is only used by the Monitor node.
|
|
//
|
|
// Example usage:
|
|
// #[node_macro::node(..., serialize(my_module::custom_serialize), ...)]
|
|
"serialize" => {
|
|
let meta = meta.require_list()?;
|
|
if serialize.is_some() {
|
|
return Err(Error::new_spanned(meta, "Multiple 'serialize' attributes are not allowed"));
|
|
}
|
|
let parsed_path: Path = meta
|
|
.parse_args()
|
|
.map_err(|_| Error::new_spanned(meta, "Expected a valid path for 'serialize', e.g., serialize(my_module::custom_serialize)"))?;
|
|
serialize = Some(parsed_path);
|
|
}
|
|
// Instructs the preprocessor to insert a Memoize node after this node in the generated subnetwork,
|
|
// caching its output across evaluations with identical inputs.
|
|
//
|
|
// Example usage:
|
|
// #[node_macro::node(..., memoize, ...)]
|
|
"memoize" => {
|
|
let path = meta.require_path_only()?;
|
|
if memoize {
|
|
return Err(Error::new_spanned(path, "Multiple 'memoize' attributes are not allowed"));
|
|
}
|
|
memoize = true;
|
|
}
|
|
// Instructs the preprocessor to make this node available as a scope.
|
|
// Other nodes can then access it with `#[scope(node::IDENTIFIER)]`.
|
|
//
|
|
// Example usage:
|
|
// #[node_macro::node(..., inject_scope, ...)]
|
|
"inject_scope" => {
|
|
let path = meta.require_path_only()?;
|
|
if inject_scope {
|
|
return Err(Error::new_spanned(path, "Multiple 'inject_scope' attributes are not allowed"));
|
|
}
|
|
inject_scope = true;
|
|
}
|
|
// Function producing a stand-in value for an async source node while the spawned future is in flight.
|
|
// The node reports `Partial` with the stand-in until the real value lands; without a placeholder it reports `Pending`.
|
|
//
|
|
// Example usage:
|
|
// #[node_macro::node(..., placeholder(empty_image), ...)]
|
|
"placeholder" => {
|
|
let meta = meta.require_list()?;
|
|
if placeholder.is_some() {
|
|
return Err(Error::new_spanned(meta, "Multiple 'placeholder' attributes are not allowed"));
|
|
}
|
|
let parsed_path: Path = meta
|
|
.parse_args()
|
|
.map_err(|_| Error::new_spanned(meta, "Expected a valid path for 'placeholder', e.g., placeholder(empty_image)"))?;
|
|
placeholder = Some(parsed_path);
|
|
}
|
|
// Function overriding the generated `extent` method, replacing the default meet over the node's inputs.
|
|
//
|
|
// Example usage:
|
|
// #[node_macro::node(..., extent(my_extent), ...)]
|
|
"extent" => {
|
|
let meta = meta.require_list()?;
|
|
if extent.is_some() {
|
|
return Err(Error::new_spanned(meta, "Multiple 'extent' attributes are not allowed"));
|
|
}
|
|
let parsed_path: Path = meta.parse_args().map_err(|_| Error::new_spanned(meta, "Expected a valid path for 'extent', e.g., extent(my_extent)"))?;
|
|
extent = Some(parsed_path);
|
|
}
|
|
// Escape hatch for extent overrides needing arbitrary context access: the raw
|
|
// `(node, ctx, level)` form instead of the typed `extent(fn)` input surface.
|
|
//
|
|
// Example usage:
|
|
// #[node_macro::node(..., extent_raw(my_extent), ...)]
|
|
"extent_raw" => {
|
|
let meta = meta.require_list()?;
|
|
if extent_raw.is_some() {
|
|
return Err(Error::new_spanned(meta, "Multiple 'extent_raw' attributes are not allowed"));
|
|
}
|
|
let parsed_path: Path = meta
|
|
.parse_args()
|
|
.map_err(|_| Error::new_spanned(meta, "Expected a valid path for 'extent_raw', e.g., extent_raw(my_extent)"))?;
|
|
extent_raw = Some(parsed_path);
|
|
}
|
|
// Function overriding the generated `eval_batch` method, replacing the trait's per-lane spec loop.
|
|
//
|
|
// Example usage:
|
|
// #[node_macro::node(..., batch(my_batch), ...)]
|
|
"batch" => {
|
|
let meta = meta.require_list()?;
|
|
if batch.is_some() {
|
|
return Err(Error::new_spanned(meta, "Multiple 'batch' attributes are not allowed"));
|
|
}
|
|
let parsed_path: Path = meta.parse_args().map_err(|_| Error::new_spanned(meta, "Expected a valid path for 'batch', e.g., batch(my_batch)"))?;
|
|
batch = Some(parsed_path);
|
|
}
|
|
// Keeps the plain-input lowering for this node during the record transition.
|
|
//
|
|
// Example usage:
|
|
// #[node_macro::node(..., plain, ...)]
|
|
"plain" => {
|
|
let path = meta.require_path_only()?;
|
|
if plain {
|
|
return Err(Error::new_spanned(path, "Multiple 'plain' attributes are not allowed"));
|
|
}
|
|
plain = true;
|
|
}
|
|
// Instructs the generated eval to report `Pending` instead of passing partial upstream values into this node.
|
|
//
|
|
// Example usage:
|
|
// #[node_macro::node(..., no_partial, ...)]
|
|
"no_partial" => {
|
|
let path = meta.require_path_only()?;
|
|
if no_partial {
|
|
return Err(Error::new_spanned(path, "Multiple 'no_partial' attributes are not allowed"));
|
|
}
|
|
no_partial = true;
|
|
}
|
|
_ => {
|
|
return Err(Error::new_spanned(
|
|
meta,
|
|
indoc!(
|
|
r#"
|
|
Unsupported attribute in `node`.
|
|
Supported attributes are 'category', 'name', 'path', 'skip_impl', 'properties', 'cfg', 'shader_node', 'serialize', 'memoize', 'inject_scope', 'placeholder', 'extent', 'extent_raw', 'batch', and 'no_partial'.
|
|
Example usage:
|
|
#[node_macro::node(..., name("Test Node"), ...)]
|
|
"#
|
|
),
|
|
));
|
|
}
|
|
}
|
|
}
|
|
|
|
if category.is_none() {
|
|
return Err(Error::new_spanned(
|
|
nested,
|
|
indoc!(
|
|
r#"
|
|
The attribute 'category' is required.
|
|
Example usage:
|
|
#[node_macro::node(..., category("Value"), ...)]
|
|
"#,
|
|
),
|
|
));
|
|
}
|
|
|
|
if let (Some(_), Some(raw)) = (&extent, &extent_raw) {
|
|
return Err(Error::new_spanned(raw, "'extent' and 'extent_raw' are mutually exclusive"));
|
|
}
|
|
|
|
Ok(NodeFnAttributes {
|
|
category,
|
|
display_name,
|
|
path,
|
|
skip_impl,
|
|
properties_string,
|
|
cfg,
|
|
shader_node,
|
|
serialize,
|
|
memoize,
|
|
inject_scope,
|
|
placeholder,
|
|
extent,
|
|
extent_raw,
|
|
batch,
|
|
no_partial,
|
|
plain,
|
|
})
|
|
}
|
|
}
|
|
|
|
pub(crate) fn parse_node_fn(attr: TokenStream2, item: TokenStream2) -> syn::Result<ParsedNodeFn> {
|
|
let attributes = syn::parse2::<NodeFnAttributes>(attr.clone()).map_err(|e| Error::new(e.span(), format!("Failed to parse node_fn attributes:\n{e}")))?;
|
|
let input_fn = syn::parse2::<ItemFn>(item.clone()).map_err(|e| Error::new(e.span(), format!("Failed to parse function: {e}. Make sure it's a valid Rust function.")))?;
|
|
|
|
let vis = input_fn.vis;
|
|
let fn_name = input_fn.sig.ident.clone();
|
|
let struct_name = format_ident!("{}", fn_name.to_string().to_case(Case::Pascal));
|
|
let mod_name = fn_name.clone();
|
|
let fn_generics = input_fn.sig.generics.params.into_iter().collect();
|
|
let is_async = input_fn.sig.asyncness.is_some();
|
|
|
|
let (input, fields, claim) = parse_inputs(&input_fn.sig.inputs)?;
|
|
let (output_type, output_depth) = crate::codegen::ir::strip_output_rank(&parse_output(&input_fn.sig.output)?);
|
|
let where_clause = input_fn.sig.generics.where_clause;
|
|
let body = input_fn.block.to_token_stream();
|
|
let description = input_fn
|
|
.attrs
|
|
.iter()
|
|
.filter_map(|a| {
|
|
if a.style != AttrStyle::Outer {
|
|
return None;
|
|
}
|
|
let Meta::NameValue(name_val) = &a.meta else { return None };
|
|
if name_val.path.get_ident().map(|x| x.to_string()) != Some("doc".into()) {
|
|
return None;
|
|
}
|
|
let Expr::Lit(expr_lit) = &name_val.value else { return None };
|
|
let Lit::Str(ref text) = expr_lit.lit else { return None };
|
|
Some(text.value().trim().to_string())
|
|
})
|
|
.fold(String::new(), |acc, b| acc + &b + "\n");
|
|
|
|
Ok(ParsedNodeFn {
|
|
vis,
|
|
attributes,
|
|
fn_name,
|
|
struct_name,
|
|
mod_name,
|
|
fn_generics,
|
|
input,
|
|
output_type,
|
|
output_depth,
|
|
is_async,
|
|
fields,
|
|
claim,
|
|
where_clause,
|
|
body,
|
|
description,
|
|
})
|
|
}
|
|
|
|
fn parse_inputs(inputs: &Punctuated<FnArg, Comma>) -> syn::Result<(Input, Vec<ParsedField>, Option<PatType>)> {
|
|
let mut fields = Vec::new();
|
|
let mut input = None;
|
|
let mut claim = None;
|
|
|
|
for (index, arg) in inputs.iter().enumerate() {
|
|
if let FnArg::Typed(PatType { pat, ty, attrs, .. }) = arg {
|
|
// Call argument
|
|
if index == 0 {
|
|
if extract_attribute(attrs, "default").is_some() {
|
|
return Err(Error::new_spanned(&attrs[0], "Call argument cannot be given a default value".to_string()));
|
|
}
|
|
if extract_attribute(attrs, "expose").is_some() {
|
|
return Err(Error::new_spanned(&attrs[0], "Call argument cannot be exposed".to_string()));
|
|
}
|
|
let pat_ident = match (**pat).clone() {
|
|
Pat::Ident(pat_ident) => pat_ident,
|
|
Pat::Wild(wild) => PatIdent {
|
|
attrs: wild.attrs,
|
|
by_ref: None,
|
|
mutability: None,
|
|
ident: wild.underscore_token.into(),
|
|
subpat: None,
|
|
},
|
|
_ => continue,
|
|
};
|
|
|
|
let implementations = extract_attribute(attrs, "implementations")
|
|
.map(|attr| parse_implementations(attr, &pat_ident.ident))
|
|
.transpose()?
|
|
.unwrap_or_default();
|
|
let context_features = parse_context_feature_idents(ty);
|
|
input = Some(Input {
|
|
pat_ident,
|
|
ty: (**ty).clone(),
|
|
implementations,
|
|
context_features,
|
|
});
|
|
} else if let Pat::Ident(pat_ident) = &**pat {
|
|
if attr_marker(ty).is_some() {
|
|
return Err(Error::new_spanned(pat_ident, "an attribute read binds to an input: destructure it as `(value, Attr<..>)`"));
|
|
}
|
|
// The claim is the caller's, not an input: it reaches the kernel
|
|
// from the serve the node is lowered into.
|
|
if is_frame_claim(ty) {
|
|
claim = Some(PatType {
|
|
attrs: Vec::new(),
|
|
pat: pat.clone(),
|
|
colon_token: Default::default(),
|
|
ty: ty.clone(),
|
|
});
|
|
continue;
|
|
}
|
|
let field = parse_field(pat_ident.clone(), (**ty).clone(), attrs).map_err(|e| Error::new_spanned(pat_ident, format!("Failed to parse argument '{}': {}", pat_ident.ident, e)))?;
|
|
fields.push(field);
|
|
} else if let Pat::Tuple(pat_tuple) = &**pat {
|
|
let field = parse_read_tuple(pat_tuple, ty, attrs, index)?;
|
|
fields.push(field);
|
|
} else if let Pat::Wild(wild) = &**pat {
|
|
let pat_ident = PatIdent {
|
|
attrs: wild.attrs.clone(),
|
|
by_ref: None,
|
|
mutability: None,
|
|
ident: format_ident!("_unit{}", index, span = wild.underscore_token.span),
|
|
subpat: None,
|
|
};
|
|
let field = parse_field(pat_ident, (**ty).clone(), attrs).map_err(|e| Error::new_spanned(pat, format!("Failed to parse argument: {e}")))?;
|
|
fields.push(field);
|
|
} else {
|
|
return Err(Error::new_spanned(pat, "Expected a simple identifier for the field name"));
|
|
}
|
|
} else {
|
|
return Err(Error::new_spanned(arg, "Expected a typed argument (e.g., `x: i32`)"));
|
|
}
|
|
}
|
|
|
|
let input = input.ok_or_else(|| Error::new_spanned(inputs, "Expected at least one input argument. The first argument should be the node input type."))?;
|
|
Ok((input, fields, claim))
|
|
}
|
|
|
|
/// Whether the parameter is the caller-provided frame claim a record-opaque
|
|
/// kernel serves through.
|
|
fn is_frame_claim(ty: &Type) -> bool {
|
|
matches!(ty, Type::Path(path) if path.path.segments.last().is_some_and(|segment| segment.ident == "FrameClaim"))
|
|
}
|
|
|
|
/// Splits a lazy input's `Output = (T, Attr<..>..)` tuple into the element
|
|
/// type (the input type) and the declared reads on that input. A tuple without
|
|
/// `Attr` slots is an ordinary tuple output and passes through untouched.
|
|
fn split_lazy_reads(output_type: Type) -> syn::Result<(Type, Vec<AttributeRead>)> {
|
|
let Type::Tuple(tuple) = &output_type else {
|
|
return Ok((output_type, Vec::new()));
|
|
};
|
|
if !tuple.elems.iter().any(|slot| attr_marker(slot).is_some()) {
|
|
return Ok((output_type, Vec::new()));
|
|
}
|
|
let spelling = "a lazy input with attribute reads declares `Output = (T, Attr<..>)`";
|
|
let mut slots = tuple.elems.iter();
|
|
let element = slots.next().ok_or_else(|| Error::new_spanned(tuple, spelling))?;
|
|
if attr_marker(element).is_some() {
|
|
return Err(Error::new_spanned(element, spelling));
|
|
}
|
|
let attribute_reads: Vec<AttributeRead> = slots
|
|
.enumerate()
|
|
.map(|(index, slot)| {
|
|
let marker = attr_marker(slot).ok_or_else(|| Error::new_spanned(slot, spelling))?;
|
|
Ok(AttributeRead {
|
|
pat_ident: PatIdent {
|
|
attrs: Vec::new(),
|
|
by_ref: None,
|
|
mutability: None,
|
|
ident: format_ident!("__lazy_read_{}", index, span = slot.span()),
|
|
subpat: None,
|
|
},
|
|
marker,
|
|
})
|
|
})
|
|
.collect::<syn::Result<_>>()?;
|
|
Ok((element.clone(), attribute_reads))
|
|
}
|
|
|
|
/// Parses a `(value, reads..): (T, Attr<..>..)` parameter: the value component
|
|
/// is an ordinary field of the value type, each `Attr` component a read bound
|
|
/// to this input.
|
|
fn parse_read_tuple(pat_tuple: &syn::PatTuple, ty: &Type, attrs: &[Attribute], index: usize) -> syn::Result<ParsedField> {
|
|
let spelling = "an input with attribute reads destructures as `(value, Attr<..>)` over `(T, Attr<..>)`";
|
|
let Type::Tuple(ty_tuple) = ty else {
|
|
return Err(Error::new_spanned(ty, spelling));
|
|
};
|
|
if pat_tuple.elems.len() != ty_tuple.elems.len() || ty_tuple.elems.len() < 2 {
|
|
return Err(Error::new_spanned(pat_tuple, spelling));
|
|
}
|
|
let mut slots = pat_tuple.elems.iter().zip(ty_tuple.elems.iter());
|
|
let (value_pat, value_ty) = slots.next().expect("length checked above");
|
|
if attr_marker(value_ty).is_some() {
|
|
return Err(Error::new_spanned(value_ty, spelling));
|
|
}
|
|
let value_ident = match value_pat {
|
|
Pat::Ident(pat_ident) => pat_ident.clone(),
|
|
Pat::Wild(wild) => PatIdent {
|
|
attrs: wild.attrs.clone(),
|
|
by_ref: None,
|
|
mutability: None,
|
|
ident: format_ident!("_value{}", index, span = wild.underscore_token.span),
|
|
subpat: None,
|
|
},
|
|
_ => return Err(Error::new_spanned(value_pat, "Expected a simple identifier for the value component")),
|
|
};
|
|
let attribute_reads: Vec<AttributeRead> = slots
|
|
.map(|(pat, ty)| {
|
|
let marker = attr_marker(ty).ok_or_else(|| Error::new_spanned(ty, spelling))?;
|
|
let Pat::Ident(pat_ident) = pat else {
|
|
return Err(Error::new_spanned(pat, "Expected a simple identifier for the attribute read"));
|
|
};
|
|
Ok(AttributeRead { pat_ident: pat_ident.clone(), marker })
|
|
})
|
|
.collect::<syn::Result<_>>()?;
|
|
let mut field = parse_field(value_ident.clone(), value_ty.clone(), attrs).map_err(|e| Error::new_spanned(&value_ident, format!("Failed to parse argument '{}': {}", value_ident.ident, e)))?;
|
|
field.attribute_reads = attribute_reads;
|
|
Ok(field)
|
|
}
|
|
|
|
/// A declared context feature; `ExtractIndex` carries the index level it reads.
|
|
#[derive(Debug, Clone, PartialEq)]
|
|
pub(crate) struct ContextFeatureDecl {
|
|
pub(crate) ident: Ident,
|
|
pub(crate) level: Option<u8>,
|
|
}
|
|
|
|
impl ContextFeatureDecl {
|
|
pub(crate) fn new(ident: Ident) -> Self {
|
|
Self { ident, level: None }
|
|
}
|
|
}
|
|
|
|
impl quote::ToTokens for ContextFeatureDecl {
|
|
fn to_tokens(&self, tokens: &mut proc_macro2::TokenStream) {
|
|
let ident = &self.ident;
|
|
match self.level {
|
|
Some(level) => tokens.extend(quote::quote!(#ident(#level))),
|
|
None => ident.to_tokens(tokens),
|
|
}
|
|
}
|
|
}
|
|
|
|
/// The level of an `ExtractIndex<N>` bound, defaulting to the innermost.
|
|
fn parse_index_level(segment: &syn::PathSegment) -> u8 {
|
|
let syn::PathArguments::AngleBracketed(arguments) = &segment.arguments else {
|
|
return 0;
|
|
};
|
|
for argument in &arguments.args {
|
|
if let syn::GenericArgument::Const(syn::Expr::Lit(syn::ExprLit { lit: syn::Lit::Int(int), .. })) = argument
|
|
&& let Ok(level) = int.base10_parse::<u8>()
|
|
{
|
|
return level;
|
|
}
|
|
}
|
|
0
|
|
}
|
|
|
|
/// Parse context feature identifiers from the trait bounds of a context parameter.
|
|
fn parse_context_feature_idents(ty: &Type) -> Vec<ContextFeatureDecl> {
|
|
let mut features = Vec::new();
|
|
|
|
// Check if this is an impl trait (impl Ctx + ...)
|
|
if let Type::ImplTrait(TypeImplTrait { bounds, .. }) = ty {
|
|
for bound in bounds {
|
|
if let TypeParamBound::Trait(TraitBound { path, .. }) = bound {
|
|
// Extract the last segment of the trait path
|
|
if let Some(segment) = path.segments.last() {
|
|
match segment.ident.to_string().as_str() {
|
|
"ExtractIndex" => features.push(ContextFeatureDecl {
|
|
ident: segment.ident.clone(),
|
|
level: Some(parse_index_level(segment)),
|
|
}),
|
|
// Reading the chain without a statically known level keeps every level.
|
|
"ExtractIndices" => features.push(ContextFeatureDecl {
|
|
ident: format_ident!("ExtractIndex"),
|
|
level: Some(u8::MAX),
|
|
}),
|
|
"ExtractFootprint"
|
|
| "ExtractRealTime"
|
|
| "ExtractAnimationTime"
|
|
| "ExtractPointerPosition"
|
|
| "ExtractPosition"
|
|
| "ExtractVarArgs"
|
|
| "InjectFootprint"
|
|
| "InjectRealTime"
|
|
| "InjectAnimationTime"
|
|
| "InjectPointerPosition"
|
|
| "InjectPosition"
|
|
| "InjectVarArgs" => {
|
|
features.push(ContextFeatureDecl::new(segment.ident.clone()));
|
|
}
|
|
// Modify* is conditionally transparent: the node rewrites the
|
|
// field only on its content's behalf, so it names no
|
|
// requirement of its own and the field nullifies early when
|
|
// nothing upstream reads it.
|
|
"ModifyFootprint" | "ModifyRealTime" | "ModifyAnimationTime" | "ModifyPointerPosition" | "ModifyPosition" | "ModifyIndex" | "ModifyVarArgs" => {}
|
|
// InjectIndex stays undeclared: a record node's injection
|
|
// re-addresses lanes derived from the incoming index, so it
|
|
// must not cancel the cone's index requirement in the
|
|
// nullification pass.
|
|
// Also ignore other traits like Ctx, ExtractAll, etc.
|
|
_ => {}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
features
|
|
}
|
|
|
|
fn parse_implementations(attr: &Attribute, name: &Ident) -> syn::Result<Punctuated<Type, Comma>> {
|
|
let content: TokenStream2 = attr.parse_args()?;
|
|
let parser = Punctuated::<Type, Comma>::parse_terminated;
|
|
parser.parse2(content.clone()).map_err(|e| {
|
|
let span = e.span(); // Get the span of the error
|
|
Error::new(span, format!("Failed to parse implementations for argument '{name}': {e}"))
|
|
})
|
|
}
|
|
|
|
fn parse_node_implementations<T: Parse>(attr: &Attribute, name: &Ident) -> syn::Result<Punctuated<T, Comma>> {
|
|
let content: TokenStream2 = attr.parse_args()?;
|
|
let parser = Punctuated::<T, Comma>::parse_terminated;
|
|
parser.parse2(content.clone()).map_err(|e| {
|
|
Error::new(
|
|
e.span(),
|
|
formatdoc!(
|
|
"Invalid #[implementations(...)] for argument `{}`.
|
|
Expected a comma-separated list of `InputType -> OutputType` pairs.
|
|
Example: #[implementations(i32 -> f64, String -> Vec<u8>)]
|
|
Error: {}",
|
|
name,
|
|
e
|
|
),
|
|
)
|
|
})
|
|
}
|
|
|
|
fn parse_field(pat_ident: PatIdent, ty: Type, attrs: &[Attribute]) -> syn::Result<ParsedField> {
|
|
let (ty, list_levels) = crate::codegen::ir::strip_ilist(&ty);
|
|
let ident = &pat_ident.ident;
|
|
|
|
// Checks for the #[data] attribute, indicating that this is a data field rather than an input parameter to the node.
|
|
// Data fields act as internal state, using interior mutability to cache data between node evaluations.
|
|
//
|
|
// Normally, an input parameter is a construction argument to the node that is stored as a field on the node struct.
|
|
// Specifically, its struct field stores the connected upstream node (an evaluatable lambda that returns data of the connection's type).
|
|
// By comparison, a data field is also stored as a field on the node struct, allowing it to persist state between evaluations.
|
|
// But it acts as internal state only, not exposed as a parameter in the UI or able to be wired to another node.
|
|
//
|
|
// Nodes implemented using a data field must ensure the persistent state is used in a manner that respects the invariant of idempotence,
|
|
// meaning the node's output is always deterministic whether or not the internal state is present.
|
|
let is_data_field = extract_attribute(attrs, "data").is_some();
|
|
|
|
let default_value = extract_attribute(attrs, "default")
|
|
.map(|attr| attr.parse_args().map_err(|e| Error::new_spanned(attr, format!("Invalid `default` value for argument '{ident}': {e}"))))
|
|
.transpose()?;
|
|
|
|
let scope = extract_attribute(attrs, "scope")
|
|
.map(|attr| attr.parse_args().map_err(|e| Error::new_spanned(attr, format!("Invalid `scope` value for argument '{ident}': {e}"))))
|
|
.transpose()?;
|
|
|
|
let name = extract_attribute(attrs, "name")
|
|
.map(|attr| attr.parse_args().map_err(|e| Error::new_spanned(attr, format!("Invalid `name` value for argument '{ident}': {e}"))))
|
|
.transpose()?;
|
|
|
|
let widget_override = extract_attribute(attrs, "widget")
|
|
.map(|attr| {
|
|
attr.parse_args()
|
|
.map_err(|e| Error::new_spanned(attr, format!("Invalid `widget override` value for argument '{ident}': {e}")))
|
|
})
|
|
.transpose()?
|
|
.unwrap_or_default();
|
|
|
|
let exposed = extract_attribute(attrs, "expose").is_some();
|
|
|
|
// Validate data field attributes
|
|
if is_data_field {
|
|
if default_value.is_some() {
|
|
return Err(Error::new_spanned(
|
|
&pat_ident,
|
|
"Data fields (#[data]) cannot have #[default] attribute. They are automatically initialized with Default::default()",
|
|
));
|
|
}
|
|
if scope.is_some() {
|
|
return Err(Error::new_spanned(&pat_ident, "Data fields (#[data]) cannot have #[scope] attribute"));
|
|
}
|
|
if exposed {
|
|
return Err(Error::new_spanned(
|
|
&pat_ident,
|
|
"Data fields (#[data]) cannot be exposed (#[expose]). They are internal state, not node parameters",
|
|
));
|
|
}
|
|
}
|
|
|
|
let value_source = match (default_value, scope) {
|
|
(Some(_), Some(_)) => return Err(Error::new_spanned(&pat_ident, "Cannot have both `default` and `scope` attributes")),
|
|
(Some(default_value), _) => ParsedValueSource::Default(default_value),
|
|
(_, Some(scope)) => ParsedValueSource::Scope(Box::new(scope)),
|
|
_ => ParsedValueSource::None,
|
|
};
|
|
|
|
// The slider's interactive extent (`#[soft(a..b)]`) and the enforced clamp (`#[hard(a..b)]`), each an
|
|
// optionally open-ended range. They decompose into the four bound values used by codegen and the UI.
|
|
let number_soft_bounds = extract_attribute(attrs, "soft")
|
|
.map(|attr| {
|
|
attr.parse_args::<NumberRange>()
|
|
.map_err(|e| Error::new_spanned(attr, format!("Invalid `soft` bounds for argument '{ident}': {e}\nUSAGE EXAMPLE: #[soft(0..100)]")))
|
|
})
|
|
.transpose()?;
|
|
let number_hard_bounds = extract_attribute(attrs, "hard")
|
|
.map(|attr| {
|
|
attr.parse_args::<NumberRange>()
|
|
.map_err(|e| Error::new_spanned(attr, format!("Invalid `hard` bounds for argument '{ident}': {e}\nUSAGE EXAMPLE: #[hard(0..100)]")))
|
|
})
|
|
.transpose()?;
|
|
let number_soft_min = number_soft_bounds.as_ref().and_then(|range| range.start.clone());
|
|
let number_soft_max = number_soft_bounds.as_ref().and_then(|range| range.end.clone());
|
|
let number_hard_min = number_hard_bounds.as_ref().and_then(|range| range.start.clone());
|
|
let number_hard_max = number_hard_bounds.as_ref().and_then(|range| range.end.clone());
|
|
|
|
// The `#[range]` marker selects the slider widget; its extent is derived from the soft (then hard) bounds.
|
|
let number_mode_range = extract_attribute(attrs, "range").is_some();
|
|
|
|
let unit = extract_attribute(attrs, "unit")
|
|
.map(|attr| attr.parse_args::<LitStr>().map_err(|_e| Error::new_spanned(attr, "Expected a unit type as string".to_string())))
|
|
.transpose()?;
|
|
|
|
let number_display_decimal_places = extract_attribute(attrs, "display_decimal_places")
|
|
.map(|attr| {
|
|
attr.parse_args::<LitInt>().map_err(|e| {
|
|
Error::new_spanned(
|
|
attr,
|
|
format!("Invalid `integer` for number of decimals for argument '{ident}': {e}\nUSAGE EXAMPLE: #[display_decimal_places(2)]"),
|
|
)
|
|
})
|
|
})
|
|
.transpose()?
|
|
.map(|f| {
|
|
if let Err(e) = f.base10_parse::<u32>() {
|
|
Err(Error::new_spanned(f, format!("Expected a `u32` for `display_decimal_places` for '{ident}': {e}")))
|
|
} else {
|
|
Ok(f)
|
|
}
|
|
})
|
|
.transpose()?;
|
|
let number_step = extract_attribute(attrs, "step")
|
|
.map(|attr| {
|
|
attr.parse_args::<LitFloat>()
|
|
.map_err(|e| Error::new_spanned(attr, format!("Invalid `step` for argument '{ident}': {e}\nUSAGE EXAMPLE: #[step(2.)]")))
|
|
})
|
|
.transpose()?;
|
|
let gpu_image = extract_attribute(attrs, "gpu_image").is_some();
|
|
|
|
let (is_node, node_input_type, node_output_type) = parse_node_type(&ty);
|
|
let description = attrs
|
|
.iter()
|
|
.filter_map(|a| {
|
|
if a.style != AttrStyle::Outer {
|
|
return None;
|
|
}
|
|
let Meta::NameValue(name_val) = &a.meta else { return None };
|
|
if name_val.path.get_ident().map(|x| x.to_string()) != Some("doc".into()) {
|
|
return None;
|
|
}
|
|
let Expr::Lit(expr_lit) = &name_val.value else { return None };
|
|
let Lit::Str(ref text) = expr_lit.lit else { return None };
|
|
Some(text.value().trim().to_string())
|
|
})
|
|
.fold(String::new(), |acc, b| acc + &b + "\n");
|
|
|
|
if is_node {
|
|
// Data fields cannot be impl Node types
|
|
if is_data_field {
|
|
return Err(Error::new_spanned(
|
|
&ty,
|
|
"Data fields (#[data]) cannot be of type `impl Node`. Data fields must be concrete types that implement Default",
|
|
));
|
|
}
|
|
|
|
let input_type = node_input_type.ok_or_else(|| Error::new_spanned(&ty, "Invalid Node type. Expected `impl Node<Input>` or `impl Node<Input, Output = OutputType>`"))?;
|
|
// A subject named without an output is a whole-record input: the kernel
|
|
// serves it through its own claim rather than reading an element.
|
|
let output_type = node_output_type.unwrap_or_else(|| syn::parse_quote!(Served<'_>));
|
|
if !matches!(&value_source, ParsedValueSource::None) {
|
|
return Err(Error::new_spanned(&ty, "No default values for `impl Node` allowed"));
|
|
}
|
|
let implementations = extract_attribute(attrs, "implementations")
|
|
.map(|attr| parse_node_implementations(attr, ident))
|
|
.transpose()?
|
|
.unwrap_or_default();
|
|
|
|
let (output_type, attribute_reads) = split_lazy_reads(output_type)?;
|
|
Ok(ParsedField {
|
|
pat_ident,
|
|
ty: ParsedFieldType::Node(NodeParsedField {
|
|
input_type,
|
|
output_type,
|
|
implementations,
|
|
}),
|
|
name,
|
|
description,
|
|
widget_override,
|
|
number_display_decimal_places,
|
|
number_step,
|
|
unit,
|
|
is_data_field,
|
|
attribute_reads,
|
|
})
|
|
} else {
|
|
let implementations = extract_attribute(attrs, "implementations")
|
|
.map(|attr| parse_implementations(attr, ident))
|
|
.transpose()?
|
|
.unwrap_or_default();
|
|
|
|
let (ty, lend) = match ty {
|
|
Type::Reference(reference) => ((*reference.elem).clone(), Some(reference)),
|
|
ty => (ty, None),
|
|
};
|
|
|
|
// Error if a float literal is given for a bound on an integer-typed field
|
|
if is_integer_type(&ty) {
|
|
let bound_attrs = [
|
|
(&number_soft_min, "soft", "lower"),
|
|
(&number_soft_max, "soft", "upper"),
|
|
(&number_hard_min, "hard", "lower"),
|
|
(&number_hard_max, "hard", "upper"),
|
|
];
|
|
for (bound, attr_name, end) in bound_attrs {
|
|
if let Some(NumberBound {
|
|
literal: NumberBoundLiteral::Float(_),
|
|
..
|
|
}) = bound
|
|
{
|
|
return Err(Error::new_spanned(
|
|
&pat_ident,
|
|
format!("The {end} `#[{attr_name}]` bound on `{ident}` is a float literal, but `{ident}` is an integer type. Use an integer literal without a decimal point."),
|
|
));
|
|
}
|
|
}
|
|
}
|
|
|
|
// A `Named<X>` parameter declares where `X`'s name is wired: the input
|
|
// carries constant text, the kernel takes only the placeholder.
|
|
let name_source = named_source(&ty);
|
|
let ty = match name_source {
|
|
Some(_) => parse_quote!(String),
|
|
None => ty,
|
|
};
|
|
|
|
Ok(ParsedField {
|
|
pat_ident,
|
|
ty: ParsedFieldType::Regular(RegularParsedField {
|
|
name_source,
|
|
exposed,
|
|
number_soft_min,
|
|
number_soft_max,
|
|
number_hard_min,
|
|
number_hard_max,
|
|
number_mode_range,
|
|
ty,
|
|
list_levels,
|
|
lend,
|
|
value_source,
|
|
implementations,
|
|
gpu_image,
|
|
}),
|
|
name,
|
|
description,
|
|
widget_override,
|
|
number_display_decimal_places,
|
|
number_step,
|
|
unit,
|
|
is_data_field,
|
|
attribute_reads: Vec::new(),
|
|
})
|
|
}
|
|
}
|
|
|
|
fn parse_node_type(ty: &Type) -> (bool, Option<Type>, Option<Type>) {
|
|
if let Type::ImplTrait(impl_trait) = ty {
|
|
for bound in &impl_trait.bounds {
|
|
if let syn::TypeParamBound::Trait(trait_bound) = bound
|
|
&& trait_bound.path.segments.last().is_some_and(|seg| seg.ident == "Node")
|
|
&& let syn::PathArguments::AngleBracketed(args) = &trait_bound.path.segments.last().unwrap().arguments
|
|
{
|
|
let input_type = args.args.iter().find_map(|arg| if let syn::GenericArgument::Type(ty) = arg { Some(ty.clone()) } else { None });
|
|
let output_type = args.args.iter().find_map(|arg| {
|
|
if let syn::GenericArgument::AssocType(assoc_type) = arg {
|
|
if assoc_type.ident == "Output" { Some(assoc_type.ty.clone()) } else { None }
|
|
} else {
|
|
None
|
|
}
|
|
});
|
|
return (true, input_type, output_type);
|
|
}
|
|
}
|
|
}
|
|
(false, None, None)
|
|
}
|
|
|
|
fn is_integer_type(ty: &Type) -> bool {
|
|
let Type::Path(type_path) = ty else { return false };
|
|
let Some(segment) = type_path.path.segments.last() else { return false };
|
|
matches!(
|
|
segment.ident.to_string().as_str(),
|
|
"u8" | "u16" | "u32" | "u64" | "u128" | "usize" | "i8" | "i16" | "i32" | "i64" | "i128" | "isize"
|
|
)
|
|
}
|
|
|
|
fn parse_output(output: &ReturnType) -> syn::Result<Type> {
|
|
match output {
|
|
ReturnType::Default => Ok(syn::parse_quote!(())),
|
|
ReturnType::Type(_, ty) => Ok((**ty).clone()),
|
|
}
|
|
}
|
|
|
|
fn extract_attribute<'a>(attrs: &'a [Attribute], name: &str) -> Option<&'a Attribute> {
|
|
attrs.iter().find(|attr| attr.path().is_ident(name))
|
|
}
|
|
|
|
// Modify the new_node_fn function to use the code generation
|
|
pub fn new_node_fn(attr: TokenStream2, item: TokenStream2) -> syn::Result<TokenStream2> {
|
|
let crate_ident = CrateIdent::default();
|
|
let mut parsed_node = parse_node_fn(attr, item.clone()).map_err(|e| Error::new(e.span(), format!("Failed to parse node function:\n{e}")))?;
|
|
parsed_node.replace_impl_trait_in_input();
|
|
if parsed_node.injects_async_source_fields() {
|
|
let core_types = crate_ident.gcore()?.clone();
|
|
parsed_node.inject_async_source_fields(&core_types);
|
|
}
|
|
crate::validation::validate_node_fn(&parsed_node).map_err(|e| Error::new(e.span(), format!("Validation error:\n{e}")))?;
|
|
generate_node_code(&crate_ident, &parsed_node).map_err(|e| Error::new(e.span(), format!("Failed to generate node code:\n{e}")))
|
|
}
|
|
|
|
impl ParsedNodeFn {
|
|
pub fn replace_impl_trait_in_input(&mut self) {
|
|
if let Type::ImplTrait(impl_trait) = self.input.ty.clone() {
|
|
let ident = Ident::new("_Input", impl_trait.span());
|
|
let mut bounds = impl_trait.bounds;
|
|
bounds.push(parse_quote!('n));
|
|
self.fn_generics.push(GenericParam::Type(TypeParam {
|
|
attrs: Default::default(),
|
|
ident: ident.clone(),
|
|
colon_token: Some(Default::default()),
|
|
bounds,
|
|
eq_token: None,
|
|
default: None,
|
|
}));
|
|
self.input.ty = parse_quote!(#ident);
|
|
if self.input.implementations.is_empty() {
|
|
self.input.implementations.push(parse_quote!(gcore::Context));
|
|
}
|
|
}
|
|
if self.input.pat_ident.ident == "_" {
|
|
self.input.pat_ident.ident = Ident::new("__ctx", self.input.pat_ident.ident.span());
|
|
}
|
|
}
|
|
|
|
pub fn injects_async_source_fields(&self) -> bool {
|
|
self.is_async || crate::codegen::is_source_kernel(&self.output_type)
|
|
}
|
|
|
|
pub fn inject_async_source_fields(&mut self, core_types: &TokenStream2) {
|
|
let hidden_field = |name: &str, ty: Type, value_source: ParsedValueSource| ParsedField {
|
|
pat_ident: PatIdent {
|
|
attrs: Vec::new(),
|
|
by_ref: None,
|
|
mutability: None,
|
|
ident: Ident::new(name, proc_macro2::Span::call_site()),
|
|
subpat: None,
|
|
},
|
|
name: None,
|
|
description: String::new(),
|
|
widget_override: ParsedWidgetOverride::Hidden,
|
|
ty: ParsedFieldType::Regular(RegularParsedField {
|
|
ty,
|
|
name_source: None,
|
|
list_levels: 0,
|
|
lend: None,
|
|
exposed: false,
|
|
value_source,
|
|
number_soft_min: None,
|
|
number_soft_max: None,
|
|
number_hard_min: None,
|
|
number_hard_max: None,
|
|
number_mode_range: false,
|
|
implementations: Default::default(),
|
|
gpu_image: false,
|
|
}),
|
|
number_display_decimal_places: None,
|
|
number_step: None,
|
|
unit: None,
|
|
is_data_field: false,
|
|
attribute_reads: Vec::new(),
|
|
};
|
|
self.fields.push(hidden_field(
|
|
"_runtime",
|
|
parse_quote!(#core_types::runtime::RuntimeHandle),
|
|
ParsedValueSource::Scope(Box::new(parse_quote!("graphene_std::runtime::RuntimeNode"))),
|
|
));
|
|
self.fields.push(hidden_field("_source", parse_quote!(#core_types::SourceId), ParsedValueSource::SourceId));
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
use proc_macro2::Span;
|
|
use quote::{quote, quote_spanned};
|
|
use syn::parse_quote;
|
|
fn pat_ident(name: &str) -> PatIdent {
|
|
PatIdent {
|
|
attrs: Vec::new(),
|
|
by_ref: None,
|
|
mutability: None,
|
|
ident: Ident::new(name, Span::call_site()),
|
|
subpat: None,
|
|
}
|
|
}
|
|
|
|
fn assert_parsed_node_fn(parsed: &ParsedNodeFn, expected: &ParsedNodeFn) {
|
|
assert_eq!(parsed.fn_name, expected.fn_name);
|
|
assert_eq!(parsed.struct_name, expected.struct_name);
|
|
assert_eq!(parsed.mod_name, expected.mod_name);
|
|
assert_eq!(parsed.is_async, expected.is_async);
|
|
assert_eq!(format!("{:?}", parsed.input), format!("{:?}", expected.input));
|
|
assert_eq!(format!("{:?}", parsed.output_type), format!("{:?}", expected.output_type));
|
|
assert_eq!(parsed.attributes.category, expected.attributes.category);
|
|
assert_eq!(parsed.attributes.display_name, expected.attributes.display_name);
|
|
assert_eq!(parsed.attributes.path, expected.attributes.path);
|
|
assert_eq!(parsed.attributes.skip_impl, expected.attributes.skip_impl);
|
|
assert_eq!(parsed.fields.len(), expected.fields.len());
|
|
assert_eq!(parsed.description, expected.description);
|
|
|
|
for (parsed_field, expected_field) in parsed.fields.iter().zip(expected.fields.iter()) {
|
|
match (parsed_field, expected_field) {
|
|
(
|
|
ParsedField {
|
|
pat_ident: p_name,
|
|
ty: ParsedFieldType::Regular(RegularParsedField {
|
|
ty: p_ty,
|
|
exposed: p_exp,
|
|
value_source: p_default,
|
|
..
|
|
}),
|
|
..
|
|
},
|
|
ParsedField {
|
|
pat_ident: e_name,
|
|
ty: ParsedFieldType::Regular(RegularParsedField {
|
|
ty: e_ty,
|
|
exposed: e_exp,
|
|
value_source: e_default,
|
|
..
|
|
}),
|
|
..
|
|
},
|
|
) => {
|
|
assert_eq!(p_name, e_name);
|
|
assert_eq!(p_exp, e_exp);
|
|
match (p_default, e_default) {
|
|
(ParsedValueSource::None, ParsedValueSource::None) => {}
|
|
(ParsedValueSource::Default(p), ParsedValueSource::Default(e)) => {
|
|
assert_eq!(p.to_token_stream().to_string(), e.to_token_stream().to_string());
|
|
}
|
|
(ParsedValueSource::Scope(p), ParsedValueSource::Scope(e)) => {
|
|
assert_eq!(p.to_token_stream().to_string(), e.to_token_stream().to_string());
|
|
}
|
|
_ => panic!("Mismatched default values"),
|
|
}
|
|
assert_eq!(format!("{p_ty:?}"), format!("{:?}", e_ty));
|
|
}
|
|
(
|
|
ParsedField {
|
|
pat_ident: p_name,
|
|
ty: ParsedFieldType::Node(NodeParsedField {
|
|
input_type: p_input,
|
|
output_type: p_output,
|
|
..
|
|
}),
|
|
..
|
|
},
|
|
ParsedField {
|
|
pat_ident: e_name,
|
|
ty: ParsedFieldType::Node(NodeParsedField {
|
|
input_type: e_input,
|
|
output_type: e_output,
|
|
..
|
|
}),
|
|
..
|
|
},
|
|
) => {
|
|
assert_eq!(p_name, e_name);
|
|
assert_eq!(format!("{p_input:?}"), format!("{:?}", e_input));
|
|
assert_eq!(format!("{p_output:?}"), format!("{:?}", e_output));
|
|
}
|
|
_ => panic!("Mismatched field types"),
|
|
}
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_basic_node() {
|
|
let attr = quote!(category("Math: Arithmetic"), path(core_types::TestNode), skip_impl);
|
|
let input = quote!(
|
|
/// Multi
|
|
/// Line
|
|
fn add(a: f64, b: f64) -> f64 {
|
|
a + b
|
|
}
|
|
);
|
|
|
|
let parsed = parse_node_fn(attr, input).unwrap();
|
|
let expected = ParsedNodeFn {
|
|
vis: Visibility::Inherited,
|
|
attributes: NodeFnAttributes {
|
|
category: Some(parse_quote!("Math: Arithmetic")),
|
|
display_name: None,
|
|
path: Some(parse_quote!(core_types::TestNode)),
|
|
skip_impl: true,
|
|
properties_string: None,
|
|
cfg: None,
|
|
shader_node: None,
|
|
serialize: None,
|
|
memoize: false,
|
|
inject_scope: false,
|
|
placeholder: None,
|
|
extent: None,
|
|
extent_raw: None,
|
|
batch: None,
|
|
no_partial: false,
|
|
plain: false,
|
|
},
|
|
fn_name: Ident::new("add", Span::call_site()),
|
|
struct_name: Ident::new("Add", Span::call_site()),
|
|
mod_name: Ident::new("add", Span::call_site()),
|
|
fn_generics: vec![],
|
|
where_clause: None,
|
|
input: Input {
|
|
pat_ident: pat_ident("a"),
|
|
ty: parse_quote!(f64),
|
|
implementations: Punctuated::new(),
|
|
context_features: vec![],
|
|
},
|
|
output_type: parse_quote!(f64),
|
|
output_depth: 0,
|
|
is_async: false,
|
|
claim: None,
|
|
fields: vec![ParsedField {
|
|
pat_ident: pat_ident("b"),
|
|
name: None,
|
|
description: String::new(),
|
|
widget_override: ParsedWidgetOverride::None,
|
|
ty: ParsedFieldType::Regular(RegularParsedField {
|
|
name_source: None,
|
|
lend: None,
|
|
list_levels: 0,
|
|
ty: parse_quote!(f64),
|
|
exposed: false,
|
|
value_source: ParsedValueSource::None,
|
|
number_soft_min: None,
|
|
number_soft_max: None,
|
|
number_hard_min: None,
|
|
number_hard_max: None,
|
|
number_mode_range: false,
|
|
implementations: Punctuated::new(),
|
|
gpu_image: false,
|
|
}),
|
|
number_display_decimal_places: None,
|
|
number_step: None,
|
|
unit: None,
|
|
is_data_field: false,
|
|
attribute_reads: Vec::new(),
|
|
}],
|
|
body: TokenStream2::new(),
|
|
description: String::from("Multi\nLine\n"),
|
|
};
|
|
|
|
assert_parsed_node_fn(&parsed, &expected);
|
|
}
|
|
|
|
#[test]
|
|
fn test_node_with_impl_node() {
|
|
let attr = quote!(category("General"));
|
|
let input = quote!(
|
|
/**
|
|
Hello
|
|
World
|
|
*/
|
|
fn transform<T: 'static>(footprint: Footprint, transform_target: impl Node<Footprint, Output = T>, translate: DVec2) -> T {
|
|
// Implementation details...
|
|
}
|
|
);
|
|
|
|
let parsed = parse_node_fn(attr, input).unwrap();
|
|
let expected = ParsedNodeFn {
|
|
vis: Visibility::Inherited,
|
|
attributes: NodeFnAttributes {
|
|
category: Some(parse_quote!("General")),
|
|
display_name: None,
|
|
path: None,
|
|
skip_impl: false,
|
|
properties_string: None,
|
|
cfg: None,
|
|
shader_node: None,
|
|
serialize: None,
|
|
memoize: false,
|
|
inject_scope: false,
|
|
placeholder: None,
|
|
extent: None,
|
|
extent_raw: None,
|
|
batch: None,
|
|
no_partial: false,
|
|
plain: false,
|
|
},
|
|
fn_name: Ident::new("transform", Span::call_site()),
|
|
struct_name: Ident::new("Transform", Span::call_site()),
|
|
mod_name: Ident::new("transform", Span::call_site()),
|
|
fn_generics: vec![parse_quote!(T: 'static)],
|
|
where_clause: None,
|
|
input: Input {
|
|
pat_ident: pat_ident("footprint"),
|
|
ty: parse_quote!(Footprint),
|
|
implementations: Punctuated::new(),
|
|
context_features: vec![],
|
|
},
|
|
output_type: parse_quote!(T),
|
|
output_depth: 0,
|
|
is_async: false,
|
|
claim: None,
|
|
fields: vec![
|
|
ParsedField {
|
|
pat_ident: pat_ident("transform_target"),
|
|
name: None,
|
|
description: String::new(),
|
|
widget_override: ParsedWidgetOverride::None,
|
|
ty: ParsedFieldType::Node(NodeParsedField {
|
|
input_type: parse_quote!(Footprint),
|
|
output_type: parse_quote!(T),
|
|
implementations: Punctuated::new(),
|
|
}),
|
|
number_display_decimal_places: None,
|
|
number_step: None,
|
|
unit: None,
|
|
is_data_field: false,
|
|
attribute_reads: Vec::new(),
|
|
},
|
|
ParsedField {
|
|
pat_ident: pat_ident("translate"),
|
|
name: None,
|
|
description: String::new(),
|
|
widget_override: ParsedWidgetOverride::None,
|
|
ty: ParsedFieldType::Regular(RegularParsedField {
|
|
name_source: None,
|
|
lend: None,
|
|
list_levels: 0,
|
|
ty: parse_quote!(DVec2),
|
|
exposed: false,
|
|
value_source: ParsedValueSource::None,
|
|
number_soft_min: None,
|
|
number_soft_max: None,
|
|
number_hard_min: None,
|
|
number_hard_max: None,
|
|
number_mode_range: false,
|
|
implementations: Punctuated::new(),
|
|
gpu_image: false,
|
|
}),
|
|
number_display_decimal_places: None,
|
|
number_step: None,
|
|
unit: None,
|
|
is_data_field: false,
|
|
attribute_reads: Vec::new(),
|
|
},
|
|
],
|
|
body: TokenStream2::new(),
|
|
description: String::from("Hello\n\t\t\t\tWorld\n"),
|
|
};
|
|
|
|
assert_parsed_node_fn(&parsed, &expected);
|
|
}
|
|
|
|
#[test]
|
|
fn test_node_with_default_values() {
|
|
let attr = quote!(category("Vector: Shape"));
|
|
let input = quote!(
|
|
/// Test
|
|
fn circle(_: impl Ctx + ExtractFootprint, #[default(50.)] radius: f64) -> Vector {
|
|
// Implementation details...
|
|
}
|
|
);
|
|
|
|
let parsed = parse_node_fn(attr, input).unwrap();
|
|
let expected = ParsedNodeFn {
|
|
vis: Visibility::Inherited,
|
|
attributes: NodeFnAttributes {
|
|
category: Some(parse_quote!("Vector: Shape")),
|
|
display_name: None,
|
|
path: None,
|
|
skip_impl: false,
|
|
properties_string: None,
|
|
cfg: None,
|
|
shader_node: None,
|
|
serialize: None,
|
|
memoize: false,
|
|
inject_scope: false,
|
|
placeholder: None,
|
|
extent: None,
|
|
extent_raw: None,
|
|
batch: None,
|
|
no_partial: false,
|
|
plain: false,
|
|
},
|
|
fn_name: Ident::new("circle", Span::call_site()),
|
|
struct_name: Ident::new("Circle", Span::call_site()),
|
|
mod_name: Ident::new("circle", Span::call_site()),
|
|
fn_generics: vec![],
|
|
where_clause: None,
|
|
input: Input {
|
|
pat_ident: pat_ident("_"),
|
|
ty: parse_quote!(impl Ctx + ExtractFootprint),
|
|
implementations: Punctuated::new(),
|
|
context_features: vec![ContextFeatureDecl::new(format_ident!("ExtractFootprint"))],
|
|
},
|
|
output_type: parse_quote!(Vector),
|
|
output_depth: 0,
|
|
is_async: false,
|
|
claim: None,
|
|
fields: vec![ParsedField {
|
|
pat_ident: pat_ident("radius"),
|
|
name: None,
|
|
description: String::new(),
|
|
widget_override: ParsedWidgetOverride::None,
|
|
ty: ParsedFieldType::Regular(RegularParsedField {
|
|
name_source: None,
|
|
lend: None,
|
|
list_levels: 0,
|
|
ty: parse_quote!(f64),
|
|
exposed: false,
|
|
value_source: ParsedValueSource::Default(quote!(50.)),
|
|
number_soft_min: None,
|
|
number_soft_max: None,
|
|
number_hard_min: None,
|
|
number_hard_max: None,
|
|
number_mode_range: false,
|
|
implementations: Punctuated::new(),
|
|
gpu_image: false,
|
|
}),
|
|
number_display_decimal_places: None,
|
|
number_step: None,
|
|
unit: None,
|
|
is_data_field: false,
|
|
attribute_reads: Vec::new(),
|
|
}],
|
|
body: TokenStream2::new(),
|
|
description: "Test\n".into(),
|
|
};
|
|
|
|
assert_parsed_node_fn(&parsed, &expected);
|
|
}
|
|
|
|
#[test]
|
|
fn test_node_with_implementations() {
|
|
let attr = quote!(category("Raster: Adjustment"));
|
|
let input = quote!(
|
|
fn levels<P: Pixel>(image: List<Raster<P>>, #[implementations(f32, f64)] shadows: f64) -> List<Raster<P>> {
|
|
// Implementation details...
|
|
}
|
|
);
|
|
|
|
let parsed = parse_node_fn(attr, input).unwrap();
|
|
let expected = ParsedNodeFn {
|
|
vis: Visibility::Inherited,
|
|
attributes: NodeFnAttributes {
|
|
category: Some(parse_quote!("Raster: Adjustment")),
|
|
display_name: None,
|
|
path: None,
|
|
skip_impl: false,
|
|
properties_string: None,
|
|
cfg: None,
|
|
shader_node: None,
|
|
serialize: None,
|
|
memoize: false,
|
|
inject_scope: false,
|
|
placeholder: None,
|
|
extent: None,
|
|
extent_raw: None,
|
|
batch: None,
|
|
no_partial: false,
|
|
plain: false,
|
|
},
|
|
fn_name: Ident::new("levels", Span::call_site()),
|
|
struct_name: Ident::new("Levels", Span::call_site()),
|
|
mod_name: Ident::new("levels", Span::call_site()),
|
|
fn_generics: vec![parse_quote!(P: Pixel)],
|
|
where_clause: None,
|
|
input: Input {
|
|
pat_ident: pat_ident("image"),
|
|
ty: parse_quote!(List<Raster<P>>),
|
|
implementations: Punctuated::new(),
|
|
context_features: vec![],
|
|
},
|
|
output_type: parse_quote!(List<Raster<P>>),
|
|
output_depth: 0,
|
|
is_async: false,
|
|
claim: None,
|
|
fields: vec![ParsedField {
|
|
pat_ident: pat_ident("shadows"),
|
|
name: None,
|
|
description: String::new(),
|
|
widget_override: ParsedWidgetOverride::None,
|
|
ty: ParsedFieldType::Regular(RegularParsedField {
|
|
name_source: None,
|
|
lend: None,
|
|
list_levels: 0,
|
|
ty: parse_quote!(f64),
|
|
exposed: false,
|
|
value_source: ParsedValueSource::None,
|
|
number_soft_min: None,
|
|
number_soft_max: None,
|
|
number_hard_min: None,
|
|
number_hard_max: None,
|
|
number_mode_range: false,
|
|
implementations: {
|
|
let mut p = Punctuated::new();
|
|
p.push(parse_quote!(f32));
|
|
p.push(parse_quote!(f64));
|
|
p
|
|
},
|
|
gpu_image: false,
|
|
}),
|
|
number_display_decimal_places: None,
|
|
number_step: None,
|
|
unit: None,
|
|
is_data_field: false,
|
|
attribute_reads: Vec::new(),
|
|
}],
|
|
body: TokenStream2::new(),
|
|
description: String::new(),
|
|
};
|
|
|
|
assert_parsed_node_fn(&parsed, &expected);
|
|
}
|
|
|
|
#[test]
|
|
fn test_number_min_max_range_mode() {
|
|
let attr = quote!(category("Math: Arithmetic"), path(core_types::TestNode));
|
|
let input = quote!(
|
|
fn add(
|
|
a: f64,
|
|
/// b
|
|
#[range]
|
|
#[soft(0..100)]
|
|
#[hard(-500..500)]
|
|
b: f64,
|
|
) -> f64 {
|
|
a + b
|
|
}
|
|
);
|
|
|
|
let parsed = parse_node_fn(attr, input).unwrap();
|
|
let expected = ParsedNodeFn {
|
|
vis: Visibility::Inherited,
|
|
attributes: NodeFnAttributes {
|
|
category: Some(parse_quote!("Math: Arithmetic")),
|
|
display_name: None,
|
|
path: Some(parse_quote!(core_types::TestNode)),
|
|
skip_impl: false,
|
|
properties_string: None,
|
|
cfg: None,
|
|
shader_node: None,
|
|
serialize: None,
|
|
memoize: false,
|
|
inject_scope: false,
|
|
placeholder: None,
|
|
extent: None,
|
|
extent_raw: None,
|
|
batch: None,
|
|
no_partial: false,
|
|
plain: false,
|
|
},
|
|
fn_name: Ident::new("add", Span::call_site()),
|
|
struct_name: Ident::new("Add", Span::call_site()),
|
|
mod_name: Ident::new("add", Span::call_site()),
|
|
fn_generics: vec![],
|
|
where_clause: None,
|
|
input: Input {
|
|
pat_ident: pat_ident("a"),
|
|
ty: parse_quote!(f64),
|
|
implementations: Punctuated::new(),
|
|
context_features: vec![],
|
|
},
|
|
output_type: parse_quote!(f64),
|
|
output_depth: 0,
|
|
is_async: false,
|
|
claim: None,
|
|
fields: vec![ParsedField {
|
|
pat_ident: pat_ident("b"),
|
|
name: None,
|
|
description: String::from("b"),
|
|
widget_override: ParsedWidgetOverride::None,
|
|
ty: ParsedFieldType::Regular(RegularParsedField {
|
|
name_source: None,
|
|
lend: None,
|
|
list_levels: 0,
|
|
ty: parse_quote!(f64),
|
|
exposed: false,
|
|
value_source: ParsedValueSource::None,
|
|
number_soft_min: Some(parse_quote!(0)),
|
|
number_soft_max: Some(parse_quote!(100)),
|
|
number_hard_min: Some(parse_quote!(-500)),
|
|
number_hard_max: Some(parse_quote!(500)),
|
|
number_mode_range: true,
|
|
implementations: Punctuated::new(),
|
|
gpu_image: false,
|
|
}),
|
|
number_display_decimal_places: None,
|
|
number_step: None,
|
|
unit: None,
|
|
is_data_field: false,
|
|
attribute_reads: Vec::new(),
|
|
}],
|
|
body: TokenStream2::new(),
|
|
description: String::new(),
|
|
};
|
|
|
|
assert_parsed_node_fn(&parsed, &expected);
|
|
}
|
|
|
|
#[test]
|
|
fn test_empty_bounds_range() {
|
|
let attr = quote!(category("Math: Arithmetic"));
|
|
let input = quote!(
|
|
fn add(a: f64, #[soft()] b: f64) -> f64 {
|
|
a + b
|
|
}
|
|
);
|
|
|
|
let result = parse_node_fn(attr, input);
|
|
assert!(result.is_err());
|
|
let error_message = result.unwrap_err().to_string();
|
|
assert!(error_message.contains("expected a range like `0..100`, `..100`, or `0..`"));
|
|
}
|
|
|
|
#[test]
|
|
fn test_async_node() {
|
|
let attr = quote!(category("IO"));
|
|
let input = quote!(
|
|
async fn load_image(api: &PlatformEditorApi, #[expose] path: String) -> List<Raster<CPU>> {
|
|
// Implementation details...
|
|
}
|
|
);
|
|
|
|
let parsed = parse_node_fn(attr, input).unwrap();
|
|
let expected = ParsedNodeFn {
|
|
vis: Visibility::Inherited,
|
|
attributes: NodeFnAttributes {
|
|
category: Some(parse_quote!("IO")),
|
|
display_name: None,
|
|
path: None,
|
|
skip_impl: false,
|
|
properties_string: None,
|
|
cfg: None,
|
|
shader_node: None,
|
|
serialize: None,
|
|
memoize: false,
|
|
inject_scope: false,
|
|
placeholder: None,
|
|
extent: None,
|
|
extent_raw: None,
|
|
batch: None,
|
|
no_partial: false,
|
|
plain: false,
|
|
},
|
|
fn_name: Ident::new("load_image", Span::call_site()),
|
|
struct_name: Ident::new("LoadImage", Span::call_site()),
|
|
mod_name: Ident::new("load_image", Span::call_site()),
|
|
fn_generics: vec![],
|
|
where_clause: None,
|
|
input: Input {
|
|
pat_ident: pat_ident("api"),
|
|
ty: parse_quote!(&PlatformEditorApi),
|
|
implementations: Punctuated::new(),
|
|
context_features: vec![],
|
|
},
|
|
output_type: parse_quote!(List<Raster<CPU>>),
|
|
output_depth: 0,
|
|
is_async: true,
|
|
claim: None,
|
|
fields: vec![ParsedField {
|
|
pat_ident: pat_ident("path"),
|
|
name: None,
|
|
description: String::new(),
|
|
widget_override: ParsedWidgetOverride::None,
|
|
ty: ParsedFieldType::Regular(RegularParsedField {
|
|
name_source: None,
|
|
lend: None,
|
|
list_levels: 0,
|
|
ty: parse_quote!(String),
|
|
exposed: true,
|
|
value_source: ParsedValueSource::None,
|
|
number_soft_min: None,
|
|
number_soft_max: None,
|
|
number_hard_min: None,
|
|
number_hard_max: None,
|
|
number_mode_range: false,
|
|
implementations: Punctuated::new(),
|
|
gpu_image: false,
|
|
}),
|
|
number_display_decimal_places: None,
|
|
number_step: None,
|
|
unit: None,
|
|
is_data_field: false,
|
|
attribute_reads: Vec::new(),
|
|
}],
|
|
body: TokenStream2::new(),
|
|
description: String::new(),
|
|
};
|
|
|
|
assert_parsed_node_fn(&parsed, &expected);
|
|
}
|
|
|
|
#[test]
|
|
fn test_node_with_custom_name() {
|
|
let attr = quote!(category("Custom"), name("CustomNode2"));
|
|
let input = quote!(
|
|
fn custom_node(input: i32) -> i32 {
|
|
input * 2
|
|
}
|
|
);
|
|
|
|
let parsed = parse_node_fn(attr, input).unwrap();
|
|
let expected = ParsedNodeFn {
|
|
vis: Visibility::Inherited,
|
|
attributes: NodeFnAttributes {
|
|
category: Some(parse_quote!("Custom")),
|
|
display_name: Some(parse_quote!("CustomNode2")),
|
|
path: None,
|
|
skip_impl: false,
|
|
properties_string: None,
|
|
cfg: None,
|
|
shader_node: None,
|
|
serialize: None,
|
|
memoize: false,
|
|
inject_scope: false,
|
|
placeholder: None,
|
|
extent: None,
|
|
extent_raw: None,
|
|
batch: None,
|
|
no_partial: false,
|
|
plain: false,
|
|
},
|
|
fn_name: Ident::new("custom_node", Span::call_site()),
|
|
struct_name: Ident::new("CustomNode", Span::call_site()),
|
|
mod_name: Ident::new("custom_node", Span::call_site()),
|
|
fn_generics: vec![],
|
|
where_clause: None,
|
|
input: Input {
|
|
pat_ident: pat_ident("input"),
|
|
ty: parse_quote!(i32),
|
|
implementations: Punctuated::new(),
|
|
context_features: vec![],
|
|
},
|
|
output_type: parse_quote!(i32),
|
|
output_depth: 0,
|
|
is_async: false,
|
|
claim: None,
|
|
fields: vec![],
|
|
body: TokenStream2::new(),
|
|
description: String::new(),
|
|
};
|
|
|
|
assert_parsed_node_fn(&parsed, &expected);
|
|
}
|
|
|
|
#[test]
|
|
#[should_panic(expected = "Multiple 'category' attributes are not allowed")]
|
|
fn test_multiple_categories() {
|
|
let attr = quote!(category("Math: Arithmetic"), category("General"));
|
|
let input = quote!(
|
|
fn add(a: i32, b: i32) -> i32 {
|
|
a + b
|
|
}
|
|
);
|
|
parse_node_fn(attr, input).unwrap();
|
|
}
|
|
|
|
#[test]
|
|
#[should_panic(expected = "Call argument cannot be given a default value")]
|
|
fn test_default_value_for_first_arg() {
|
|
let attr = quote!(category("Invalid"));
|
|
let input = quote!(
|
|
fn invalid_node(#[default(())] node: impl Node<(), Output = i32>) -> i32 {
|
|
node.eval(())
|
|
}
|
|
);
|
|
parse_node_fn(attr, input).unwrap();
|
|
}
|
|
|
|
#[test]
|
|
#[should_panic(expected = "No default values for `impl Node` allowed")]
|
|
fn test_default_value_for_impl_node() {
|
|
let attr = quote!(category("Invalid"));
|
|
let input = quote!(
|
|
fn invalid_node(_: (), #[default(())] node: impl Node<(), Output = i32>) -> i32 {
|
|
node.eval(())
|
|
}
|
|
);
|
|
parse_node_fn(attr, input).unwrap();
|
|
}
|
|
|
|
#[test]
|
|
#[should_panic(expected = "Unsupported attribute in `node`")]
|
|
fn test_unsupported_attribute() {
|
|
let attr = quote!(unsupported("Value"));
|
|
let input = quote!(
|
|
fn test_node(input: i32) -> i32 {
|
|
input
|
|
}
|
|
);
|
|
parse_node_fn(attr, input).unwrap();
|
|
}
|
|
|
|
#[test]
|
|
fn test_invalid_implementation_syntax() {
|
|
let attr = quote!(category("Test"));
|
|
let input = quote!(
|
|
fn test_node(_: (), #[implementations((Footprint, Color), (Footprint, List<Raster<CPU>>))] input: impl Node<Footprint, Output = T>) -> T {
|
|
// Implementation details...
|
|
}
|
|
);
|
|
|
|
let result = parse_node_fn(attr, input);
|
|
assert!(result.is_err());
|
|
let error = result.unwrap_err();
|
|
let error_message = error.to_string();
|
|
assert!(error_message.contains("Invalid #[implementations(...)] for argument `input`"));
|
|
assert!(error_message.contains("Expected a comma-separated list of `InputType -> OutputType` pairs"));
|
|
assert!(error_message.contains("Expected `->` arrow after input type in #[implementations(...)] on a field of type `impl Node`"));
|
|
}
|
|
|
|
#[test]
|
|
fn test_implementation_on_first_arg() {
|
|
let attr = quote!(category("Test"));
|
|
|
|
// Use quote_spanned! to attach a specific span to the problematic part
|
|
let problem_span = Span::call_site(); // You could create a custom span here if needed
|
|
let tuples = quote_spanned!(problem_span=> () ());
|
|
let input = quote! {
|
|
fn test_node(
|
|
#[implementations((), #tuples, Footprint)]
|
|
footprint: F,
|
|
#[implementations(
|
|
() -> List<Raster<CPU>>,
|
|
() -> List<Color>,
|
|
() -> List<GradientStops>,
|
|
Footprint -> List<Raster<CPU>>,
|
|
Footprint -> List<Color>,
|
|
Footprint -> List<GradientStops>,
|
|
)]
|
|
image: impl Node<F, Output = T>,
|
|
) -> T {
|
|
// Implementation details...
|
|
}
|
|
};
|
|
|
|
let result = parse_node_fn(attr, input);
|
|
assert!(result.is_err(), "Expected an error, but parsing succeeded");
|
|
|
|
let error = result.unwrap_err();
|
|
let error_string = error.to_string();
|
|
assert!(error_string.contains("Failed to parse implementations for argument 'footprint'"));
|
|
assert!(error_string.contains("expected `,`"));
|
|
|
|
// Instead of checking for exact line and column,
|
|
// verify that the error span is the one we specified
|
|
assert_eq!(error.span().start(), problem_span.start());
|
|
}
|
|
}
|