Add the attribute census, node macro attribute io over list wires, and the rank-0 record tier

This commit is contained in:
Dennis Kobert
2026-08-04 22:05:02 +00:00
parent 62b7ccf40c
commit 6154e4d219
11 changed files with 1620 additions and 41 deletions

View File

@@ -605,6 +605,14 @@ pub(crate) fn generate_node_impl(crate_ident: &CrateIdent, parsed: &ParsedNodeFn
top_level: quote!(),
});
}
let record = record_shape(parsed);
if record.is_none() && has_record_io(parsed) {
return Ok(NodeImplTokens {
in_mod: quote!(),
top_level: quote!(),
});
}
let routing = routing_io(parsed);
let snapshot_ctx = async_fn && matches!(&parsed.input.ty, Type::Path(path) if path.path.segments.last().is_some_and(|segment| segment.ident == "CtxSnapshot"));
let mut ctx_bounds: Vec<TokenStream2> = match ctx_param {
@@ -665,26 +673,44 @@ pub(crate) fn generate_node_impl(crate_ident: &CrateIdent, parsed: &ParsedNodeFn
true => quote!(#ctx_ident),
false => quote!(#ctx_ident: #(#ctx_bounds)+*),
};
let mut generics: Vec<TokenStream2> = parsed
let generic_tokens = |param: &GenericParam| match param {
GenericParam::Type(type_param) if Some(&type_param.ident) == ctx_param.map(|ctx_param| &ctx_param.ident) => ctx_generic.clone(),
param => quote!(#param),
};
let mut generics: Vec<TokenStream2> = parsed.fn_generics.iter().map(&generic_tokens).collect();
let mut impl_generics: Vec<TokenStream2> = parsed
.fn_generics
.iter()
.map(|param| match param {
GenericParam::Type(type_param) if Some(&type_param.ident) == ctx_param.map(|ctx_param| &ctx_param.ident) => ctx_generic.clone(),
param => quote!(#param),
.filter(|param| match (param, &routing) {
(GenericParam::Type(type_param), Some(routing)) => type_param.ident != routing.generic,
_ => true,
})
.map(&generic_tokens)
.collect();
if ctx_param.is_none() {
generics.push(ctx_generic);
generics.push(ctx_generic.clone());
impl_generics.push(ctx_generic);
}
if let Some(lifetime) = &introduced_lend_lifetime {
generics.insert(0, quote!(#lifetime));
impl_generics.insert(0, quote!(#lifetime));
}
if routing.is_some() {
impl_generics.insert(0, quote!('__record));
}
let fn_name = &parsed.fn_name;
let mod_name = format_ident!("_{}_mod", parsed.mod_name);
let struct_name = format_ident!("{}Node", parsed.struct_name);
let output_type = &parsed.output_type;
let trait_output = slot_value_type(&parsed.output_type);
let trait_output = match (&record, &routing) {
(Some(shape), _) => {
let element_out = &shape.element_out;
syn::parse_quote!(#core_types::list::List<#element_out>)
}
(None, Some(_)) => syn::parse_quote!(#core_types::record::RecordValue<'__record>),
(None, None) => slot_value_type(&parsed.output_type),
};
let raw_lazy = matches!(kernel_kind(&parsed.output_type), KernelKind::Poll(_));
let injected_name = |ident: &Ident| async_source && (ident == "_runtime" || ident == "_source");
let where_predicates: Vec<TokenStream2> = parsed.where_clause.iter().flat_map(|clause| clause.predicates.iter()).map(|predicate| quote!(#predicate)).collect();
@@ -723,28 +749,49 @@ pub(crate) fn generate_node_impl(crate_ident: &CrateIdent, parsed: &ParsedNodeFn
false => quote!(#core_types::node::Node<#ctx_ident, Output = #output_type>),
};
let kernel_params = regular_fields.iter().filter(|field| !injected_name(&field.pat_ident.ident)).map(|field| {
let pat = &field.pat_ident;
match &field.ty {
ParsedFieldType::Regular(RegularParsedField { ty, lend: Some(_), .. }) => quote!(#pat: &#ty),
ParsedFieldType::Regular(RegularParsedField { ty, .. }) => quote!(#pat: #ty),
ParsedFieldType::Node(NodeParsedField { output_type, .. }) if raw_lazy => {
let bound = lazy_bound(output_type);
quote!(#pat: &impl #bound)
}
ParsedFieldType::Node(NodeParsedField { output_type, .. }) => {
let bound = lazy_bound(output_type);
quote!(#pat: #core_types::node::LazyInput<'_, impl #bound>)
}
}
let attr_kernel_params = parsed.attribute_reads.iter().map(|read| {
let pat = &read.pat_ident;
let marker = &read.marker;
quote!(#pat: #core_types::attribute::Attr<#marker>)
});
let kernel_params = regular_fields
.iter()
.filter(|field| !injected_name(&field.pat_ident.ident))
.map(|field| {
let pat = &field.pat_ident;
match &field.ty {
ParsedFieldType::Regular(RegularParsedField { ty, lend: Some(_), .. }) => quote!(#pat: &#ty),
ParsedFieldType::Regular(RegularParsedField { ty, .. }) => quote!(#pat: #ty),
ParsedFieldType::Node(NodeParsedField { output_type, .. }) if raw_lazy => {
let bound = lazy_bound(output_type);
quote!(#pat: &impl #bound)
}
ParsedFieldType::Node(NodeParsedField { output_type, .. }) => {
let bound = lazy_bound(output_type);
quote!(#pat: #core_types::node::LazyInput<'_, impl #bound>)
}
}
})
.chain(attr_kernel_params);
let node_bounds = regular_fields.iter().zip(&node_generics).map(|(field, node_generic)| match &field.ty {
let routing_source = |ty: &Type| matches!((&routing, ty), (Some(routing), Type::Path(path)) if path.path.get_ident() == Some(&routing.generic));
let record_value_ty: Type = syn::parse_quote!(#core_types::record::RecordValue<'__record>);
let node_bounds = regular_fields.iter().enumerate().zip(&node_generics).map(|((index, field), node_generic)| match &field.ty {
ParsedFieldType::Regular(RegularParsedField { ty, lend: Some(_), .. }) => {
let lifetime = lend_lifetime.as_ref().expect("lend fields imply the lend lifetime");
quote!(#node_generic: #core_types::node::Node<#ctx_ident, Output = &#lifetime #ty>)
}
ParsedFieldType::Regular(RegularParsedField { ty, .. }) if record.is_some() && index == 0 => {
quote!(#node_generic: #core_types::node::Node<#ctx_ident, Output = #core_types::list::List<#ty>>)
}
ParsedFieldType::Regular(RegularParsedField { ty, .. }) if routing_source(ty) => {
quote!(#node_generic: #core_types::node::Node<#ctx_ident, Output = #record_value_ty>)
}
ParsedFieldType::Regular(RegularParsedField { ty, .. }) => quote!(#node_generic: #core_types::node::Node<#ctx_ident, Output = #ty>),
ParsedFieldType::Node(NodeParsedField { output_type, .. }) if routing_source(output_type) => {
let bound = lazy_bound(&record_value_ty);
quote!(#node_generic: #bound)
}
ParsedFieldType::Node(NodeParsedField { output_type, .. }) => {
let bound = lazy_bound(output_type);
quote!(#node_generic: #bound)
@@ -801,6 +848,12 @@ pub(crate) fn generate_node_impl(crate_ident: &CrateIdent, parsed: &ParsedNodeFn
let eval_values = regular_fields.iter().enumerate().map(|(index, field)| {
let name = &field.pat_ident.ident;
match &field.ty {
ParsedFieldType::Regular(_) if record.is_some() && index == 0 => quote! {
let mut __record_list = match __cell.eval_input(#index, &self.#name, __input) {
Ok(value) => value,
Err(interrupt) => return interrupt.into(),
};
},
ParsedFieldType::Regular(_) => quote! {
let #name = match __cell.eval_input(#index, &self.#name, __input) {
Ok(value) => value,
@@ -994,8 +1047,109 @@ pub(crate) fn generate_node_impl(crate_ident: &CrateIdent, parsed: &ParsedNodeFn
#fallback
}
};
let record_tail = record.as_ref().map(|shape| {
let record_call_args = regular_fields.iter().enumerate().map(|(index, field)| {
let name = &field.pat_ident.ident;
match (index, &field.ty) {
(0, _) => quote!(__element),
(_, ParsedFieldType::Regular(RegularParsedField { lend: Some(_), .. })) => quote!(#name),
_ => quote!(#name.clone()),
}
});
let attr_call_args = parsed.attribute_reads.iter().map(|read| {
let pat = &read.pat_ident.ident;
quote!(#pat)
});
let record_kernel_call = quote!(self::#fn_name(__input #(, &self.#data_names)* #(, #record_call_args)* #(, #attr_call_args)*));
let read_columns = parsed.attribute_reads.iter().enumerate().map(|(index, read)| {
let marker = &read.marker;
let column = format_ident!("__read_{index}");
quote! {
let #column = __record_list
.iter_attribute_values::<<#marker as #core_types::attribute::Attribute>::Value>(<#marker as #core_types::attribute::Attribute>::NAME)
.map(|__values| __values.cloned().collect::<::std::vec::Vec<_>>());
}
});
let read_bindings = parsed.attribute_reads.iter().enumerate().map(|(index, read)| {
let pat = &read.pat_ident;
let marker = &read.marker;
let column = format_ident!("__read_{index}");
quote! {
let #pat = #core_types::attribute::Attr::<#marker>(match &#column {
Some(__values) => __values[__index].clone(),
None => <#marker as #core_types::attribute::Attribute>::default(),
});
}
});
let write_markers = &shape.write_markers;
let write_columns: Vec<Ident> = (0..write_markers.len()).map(|index| format_ident!("__write_{index}")).collect();
let write_column_decls = write_markers.iter().zip(&write_columns).map(|(marker, column)| {
quote! {
let mut #column: ::std::vec::Vec<<#marker as #core_types::attribute::Attribute>::Value> = ::std::vec::Vec::with_capacity(__len);
}
});
let write_pats: Vec<Ident> = (0..write_markers.len()).map(|index| format_ident!("__written_{index}")).collect();
let destructure = match write_markers.is_empty() {
true => quote!(let __element_out = __kernel_value;),
false => quote!(let (__element_out #(, #core_types::attribute::Attr(#write_pats))*) = __kernel_value;),
};
let write_pushes = write_columns.iter().zip(&write_pats).map(|(column, pat)| quote!(#column.push(#pat);));
let written_names = write_markers.iter().map(|marker| quote!(<#marker as #core_types::attribute::Attribute>::NAME));
let write_inserts = write_markers.iter().zip(&write_columns).map(|(marker, column)| {
quote! {
__out.insert_attribute_dyn(
<#marker as #core_types::attribute::Attribute>::NAME,
#core_types::list::AttributeDyn(::std::boxed::Box::new(#core_types::list::Attribute(#column))),
);
}
});
let kernel_value = match shape.dialect {
RecordDialect::Plain => quote!(#record_kernel_call),
RecordDialect::Interrupt => quote! {
match #record_kernel_call {
Ok(__value) => __value,
Err(__interrupt) => return __interrupt.into(),
}
},
};
let element_out = &shape.element_out;
quote! {
let __len = __record_list.len();
#(#read_columns)*
let __written_names: &[&str] = &[#(#written_names),*];
let __carried_keys: ::std::vec::Vec<::std::string::String> = __record_list
.attribute_keys()
.filter(|__key| !__written_names.contains(__key))
.map(::std::string::String::from)
.collect();
let mut __carried: ::std::vec::Vec<(::std::string::String, #core_types::list::AttributeDyn)> = ::std::vec::Vec::with_capacity(__carried_keys.len());
for __key in __carried_keys {
if let Some(__column) = __record_list.take_attribute_dyn(&__key) {
__carried.push((__key, __column));
}
}
#(#write_column_decls)*
let mut __out_elements: ::std::vec::Vec<#element_out> = ::std::vec::Vec::with_capacity(__len);
for (__index, __element) in __record_list.into_element_values().into_iter().enumerate() {
#(#read_bindings)*
let __kernel_value = #kernel_value;
#destructure
__out_elements.push(__element_out);
#(#write_pushes)*
}
let mut __out = #core_types::list::List::from_element_values(__out_elements);
for (__key, __column) in __carried {
__out.insert_attribute_dyn(__key, __column);
}
#(#write_inserts)*
__cell.finish(__out)
}
});
let eval_tail = match (async_fn, future_kernel) {
(false, false) => lift,
(false, false) => match record_tail {
Some(tail) => tail,
None => lift,
},
(true, _) => {
let kernel_value_names: Vec<&Ident> = kernel_fields.iter().map(|field| &field.pat_ident.ident).collect();
let snapshot_binding = snapshot_ctx.then(|| quote!(let __snapshot = #core_types::context::CtxSnapshot::capture(__input);)).into_iter();
@@ -1044,18 +1198,32 @@ pub(crate) fn generate_node_impl(crate_ident: &CrateIdent, parsed: &ParsedNodeFn
}
};
let record_bounds: Vec<TokenStream2> = match &record {
Some(_) => regular_fields
.iter()
.skip(1)
.filter_map(|field| match &field.ty {
ParsedFieldType::Regular(RegularParsedField { lend: Some(_), .. }) => None,
ParsedFieldType::Regular(RegularParsedField { ty, .. }) => Some(quote!(#ty: Clone)),
_ => None,
})
.collect(),
None => Vec::new(),
};
let entries = entries_tokens(parsed, &struct_name, &data_field_generic_idents, &regular_fields);
let cfg = crate::shader_nodes::modify_cfg(&parsed.attributes);
let top_level = quote! {
#cfg
#[automatically_derived]
impl<#(#generics,)* #(#node_generics,)*> #core_types::node::Node<#ctx_ident> for #mod_name::#struct_name<#(#struct_type_params,)*>
impl<#(#impl_generics,)* #(#node_generics,)*> #core_types::node::Node<#ctx_ident> for #mod_name::#struct_name<#(#struct_type_params,)*>
where
#(#node_bounds,)*
#(#lend_outlives,)*
#(#clampable_bounds,)*
#(#async_bounds,)*
#(#record_bounds,)*
#(#where_predicates,)*
{
type Output = #trait_output;
@@ -1085,6 +1253,140 @@ pub(crate) fn generate_node_impl(crate_ident: &CrateIdent, parsed: &ParsedNodeFn
})
}
pub(crate) enum RecordDialect {
Plain,
Interrupt,
}
/// The record io of a node fn: the output element type and the written
/// markers. Present exactly when the signature declares attribute reads or
/// writes in a shape the driver supports (the carrier is the first field);
/// malformed record io is reported by validation and generates no node impl.
pub(crate) struct RecordShape {
pub(crate) element_out: Type,
pub(crate) write_markers: Vec<Type>,
pub(crate) dialect: RecordDialect,
}
pub(crate) fn has_record_io(parsed: &ParsedNodeFn) -> bool {
!parsed.attribute_reads.is_empty() || record_writes(&slot_value_type(&parsed.output_type)).is_some()
}
pub(crate) fn record_shape(parsed: &ParsedNodeFn) -> Option<RecordShape> {
let (value, dialect) = match kernel_kind(&parsed.output_type) {
KernelKind::Plain => (parsed.output_type.clone(), RecordDialect::Plain),
KernelKind::Interrupt(inner) => (inner, RecordDialect::Interrupt),
_ => return None,
};
let writes = record_writes(&value);
if parsed.attribute_reads.is_empty() && writes.is_none() {
return None;
}
if parsed.is_async {
return None;
}
let carrier = parsed.fields.first()?;
if carrier.is_data_field {
return None;
}
let ParsedFieldType::Regular(RegularParsedField { ty, lend: None, .. }) = &carrier.ty else {
return None;
};
if matches!(ty, Type::Tuple(tuple) if tuple.elems.is_empty()) {
return None;
}
if parsed.fields.iter().any(|field| matches!(field.ty, ParsedFieldType::Node(_))) {
return None;
}
let (element_out, write_markers) = match writes {
Some(RecordWrites { element, markers }) => (element, markers),
None => (value, Vec::new()),
};
Some(RecordShape {
element_out,
write_markers,
dialect,
})
}
pub(crate) fn is_poll_kernel(output: &Type) -> bool {
matches!(kernel_kind(output), KernelKind::Poll(_))
}
/// A routing family: an unbounded generic shared by lazy inputs (and
/// optionally the first parameter) and returned whole, instantiated at
/// `RecordValue` so opaque records flow through the kernel. Detected only
/// when the family's fields carry no implementations lists, so the existing
/// per-type row spelling keeps its meaning.
pub(crate) struct RoutingIo {
pub(crate) generic: Ident,
}
pub(crate) fn routing_io(parsed: &ParsedNodeFn) -> Option<RoutingIo> {
if has_record_io(parsed) || parsed.is_async {
return None;
}
if !matches!(kernel_kind(&parsed.output_type), KernelKind::Plain | KernelKind::Interrupt(_)) {
return None;
}
let value = slot_value_type(&parsed.output_type);
let Type::Path(path) = &value else { return None };
let ident = path.path.get_ident()?.clone();
let ctx_ident = context_param(parsed).map(|ctx| ctx.ident.clone());
parsed
.fn_generics
.iter()
.find(|param| matches!(param, GenericParam::Type(type_param) if type_param.ident == ident && type_param.bounds.is_empty() && Some(&type_param.ident) != ctx_ident.as_ref()))?;
if let Some(where_clause) = &parsed.where_clause
&& tokens_contain_ident(where_clause.to_token_stream(), &ident)
{
return None;
}
let mut sources = 0;
for (index, field) in parsed.fields.iter().enumerate() {
match &field.ty {
ParsedFieldType::Node(NodeParsedField {
output_type,
input_type,
implementations,
}) => {
if bare_ident(output_type) == Some(&ident) {
if !implementations.is_empty() || type_contains_ident(input_type, &ident) {
return None;
}
sources += 1;
} else if type_contains_ident(output_type, &ident) || type_contains_ident(input_type, &ident) {
return None;
}
}
ParsedFieldType::Regular(RegularParsedField { ty, implementations, lend, .. }) => {
if bare_ident(ty) == Some(&ident) {
if index != 0 || field.is_data_field || !implementations.is_empty() || lend.is_some() {
return None;
}
sources += 1;
} else if type_contains_ident(ty, &ident) {
return None;
}
}
}
}
(sources > 0).then(|| RoutingIo { generic: ident })
}
fn bare_ident(ty: &Type) -> Option<&Ident> {
let Type::Path(path) = ty else { return None };
path.path.get_ident()
}
fn tokens_contain_ident(tokens: TokenStream2, ident: &Ident) -> bool {
tokens.into_iter().any(|token| match token {
proc_macro2::TokenTree::Ident(candidate) => &candidate == ident,
proc_macro2::TokenTree::Group(group) => tokens_contain_ident(group.stream(), ident),
_ => false,
})
}
pub(crate) fn slot_value_type(output: &Type) -> Type {
match kernel_kind(output) {
KernelKind::Plain => output.clone(),
@@ -1252,14 +1554,17 @@ fn entries_tokens(parsed: &ParsedNodeFn, struct_name: &Ident, data_field_generic
.map(|field| matches!(&field.ty, ParsedFieldType::Regular(RegularParsedField { lend: Some(_), .. })))
.collect();
let record_carrier = record_shape(parsed).is_some();
let entries = rows.iter().map(|row| {
let input_types = row.iter().zip(&lend_flags).map(|(ty, lend)| match lend {
true => quote!(gcore::registry::lend_edge_type::<#ty>()),
false => quote!(gcore::registry::edge_type::<#ty>()),
let input_types = row.iter().enumerate().zip(&lend_flags).map(|((index, ty), lend)| match (lend, record_carrier && index == 0) {
(true, _) => quote!(gcore::registry::lend_edge_type::<#ty>()),
(false, true) => quote!(gcore::registry::edge_type::<gcore::list::List<#ty>>()),
(false, false) => quote!(gcore::registry::edge_type::<#ty>()),
});
let edge_types = row.iter().zip(&lend_flags).map(|(ty, lend)| match lend {
true => quote!(gcore::registry::SharedEdge<gcore::registry::ErasedLendNode<#ty>>),
false => quote!(gcore::registry::SharedEdge<gcore::registry::ErasedNode<#ty>>),
let edge_types = row.iter().enumerate().zip(&lend_flags).map(|((index, ty), lend)| match (lend, record_carrier && index == 0) {
(true, _) => quote!(gcore::registry::SharedEdge<gcore::registry::ErasedLendNode<#ty>>),
(false, true) => quote!(gcore::registry::SharedEdge<gcore::registry::ErasedNode<gcore::list::List<#ty>>>),
(false, false) => quote!(gcore::registry::SharedEdge<gcore::registry::ErasedNode<#ty>>),
});
let output = quote!(<#struct_name<#(#edge_types),*> as gcore::node::Node<gcore::context::ContextImpl<'static>>>::Output);
let (io_output, construct) = match &ref_output_inner {
@@ -1272,9 +1577,10 @@ fn entries_tokens(parsed: &ParsedNodeFn, struct_name: &Ident, data_field_generic
quote!(Ok(gcore::registry::EdgeHandle::new(::std::sync::Arc::new(#struct_name::new(#(#names),*)) as ::std::sync::Arc<gcore::registry::ErasedNode<#output>>))),
),
};
let downcasts = names.iter().zip(row.iter()).zip(&lend_flags).map(|((name, ty), lend)| match lend {
true => quote!(let #name = inputs.next().unwrap().downcast_lend::<#ty>()?;),
false => quote!(let #name = inputs.next().unwrap().downcast::<#ty>()?;),
let downcasts = names.iter().zip(row.iter().enumerate()).zip(&lend_flags).map(|((name, (index, ty)), lend)| match (lend, record_carrier && index == 0) {
(true, _) => quote!(let #name = inputs.next().unwrap().downcast_lend::<#ty>()?;),
(false, true) => quote!(let #name = inputs.next().unwrap().downcast::<gcore::list::List<#ty>>()?;),
(false, false) => quote!(let #name = inputs.next().unwrap().downcast::<#ty>()?;),
});
quote! {
gcore::registry::RegistryEntry {

View File

@@ -7,8 +7,8 @@ use syn::punctuated::Punctuated;
use syn::spanned::Spanned;
use syn::token::{Comma, RArrow};
use syn::{
AttrStyle, Attribute, Error, Expr, FnArg, GenericParam, Ident, ItemFn, Lit, LitFloat, LitInt, LitStr, Meta, Pat, PatIdent, PatType, Path, ReturnType, TraitBound, Type, TypeImplTrait, TypeParam,
TypeParamBound, Visibility, WhereClause, parse_quote,
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;
@@ -35,10 +35,59 @@ pub(crate) struct ParsedNodeFn {
pub(crate) output_type: Type,
pub(crate) is_async: bool,
pub(crate) fields: Vec<ParsedField>,
pub(crate) attribute_reads: Vec<AttributeRead>,
pub(crate) body: TokenStream2,
pub(crate) description: String,
}
/// An `Attr<Marker>` parameter: a declared attribute read on the carrier's
/// items, not a wired input.
#[derive(Clone, Debug)]
pub(crate) struct AttributeRead {
pub(crate) pat_ident: PatIdent,
pub(crate) marker: Type,
}
/// The write half of a record kernel's return: the element type in the first
/// tuple slot and the attribute markers written after it. `None` unless the
/// value is a well-formed write tuple (a non-`Attr` element first, then only
/// `Attr` slots, at least one).
pub(crate) struct RecordWrites {
pub(crate) element: Type,
pub(crate) markers: 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() {
return None;
}
let markers: Option<Vec<Type>> = slots.map(attr_marker).collect();
let markers = markers?;
(!markers.is_empty()).then(|| RecordWrites {
element: element.clone(),
markers,
})
}
/// Returns the marker type of an `Attr<Marker>` type, if `ty` is one.
pub(crate) fn attr_marker(ty: &Type) -> Option<Type> {
let Type::Path(path) = ty else { return None };
let segment = path.path.segments.last()?;
if segment.ident != "Attr" {
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>,
@@ -577,7 +626,7 @@ fn parse_node_fn(attr: TokenStream2, item: TokenStream2) -> syn::Result<ParsedNo
let fn_generics = input_fn.sig.generics.params.into_iter().collect();
let is_async = input_fn.sig.asyncness.is_some();
let (input, fields) = parse_inputs(&input_fn.sig.inputs)?;
let (input, fields, attribute_reads) = parse_inputs(&input_fn.sig.inputs)?;
let output_type = parse_output(&input_fn.sig.output)?;
let where_clause = input_fn.sig.generics.where_clause;
let body = input_fn.block.to_token_stream();
@@ -609,14 +658,16 @@ fn parse_node_fn(attr: TokenStream2, item: TokenStream2) -> syn::Result<ParsedNo
output_type,
is_async,
fields,
attribute_reads,
where_clause,
body,
description,
})
}
fn parse_inputs(inputs: &Punctuated<FnArg, Comma>) -> syn::Result<(Input, Vec<ParsedField>)> {
fn parse_inputs(inputs: &Punctuated<FnArg, Comma>) -> syn::Result<(Input, Vec<ParsedField>, Vec<AttributeRead>)> {
let mut fields = Vec::new();
let mut attribute_reads = Vec::new();
let mut input = None;
for (index, arg) in inputs.iter().enumerate() {
@@ -653,8 +704,18 @@ fn parse_inputs(inputs: &Punctuated<FnArg, Comma>) -> syn::Result<(Input, Vec<Pa
context_features,
});
} else if let Pat::Ident(pat_ident) = &**pat {
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);
if let Some(marker) = attr_marker(ty) {
if !attrs.iter().all(|attr| attr.path().is_ident("doc")) {
return Err(Error::new_spanned(pat_ident, "attribute parameters take no field attributes"));
}
attribute_reads.push(AttributeRead {
pat_ident: pat_ident.clone(),
marker,
});
} else {
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 {
return Err(Error::new_spanned(pat, "Expected a simple identifier for the field name"));
}
@@ -664,7 +725,7 @@ fn parse_inputs(inputs: &Punctuated<FnArg, Comma>) -> syn::Result<(Input, Vec<Pa
}
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))
Ok((input, fields, attribute_reads))
}
/// Parse context feature identifiers from the trait bounds of a context parameter.
@@ -1221,6 +1282,7 @@ mod tests {
},
output_type: parse_quote!(f64),
is_async: false,
attribute_reads: vec![],
fields: vec![ParsedField {
pat_ident: pat_ident("b"),
name: None,
@@ -1296,6 +1358,7 @@ mod tests {
},
output_type: parse_quote!(T),
is_async: false,
attribute_reads: vec![],
fields: vec![
ParsedField {
pat_ident: pat_ident("transform_target"),
@@ -1385,6 +1448,7 @@ mod tests {
},
output_type: parse_quote!(Vector),
is_async: false,
attribute_reads: vec![],
fields: vec![ParsedField {
pat_ident: pat_ident("radius"),
name: None,
@@ -1456,6 +1520,7 @@ mod tests {
},
output_type: parse_quote!(List<Raster<P>>),
is_async: false,
attribute_reads: vec![],
fields: vec![ParsedField {
pat_ident: pat_ident("shadows"),
name: None,
@@ -1539,6 +1604,7 @@ mod tests {
},
output_type: parse_quote!(f64),
is_async: false,
attribute_reads: vec![],
fields: vec![ParsedField {
pat_ident: pat_ident("b"),
name: None,
@@ -1625,6 +1691,7 @@ mod tests {
},
output_type: parse_quote!(List<Raster<CPU>>),
is_async: true,
attribute_reads: vec![],
fields: vec![ParsedField {
pat_ident: pat_ident("path"),
name: None,
@@ -1696,6 +1763,7 @@ mod tests {
},
output_type: parse_quote!(i32),
is_async: false,
attribute_reads: vec![],
fields: vec![],
body: TokenStream2::new(),
description: String::new(),

View File

@@ -317,6 +317,7 @@ impl PerPixelAdjustCodegen<'_> {
output_type: raster_gpu,
is_async: false,
fields,
attribute_reads: Vec::new(),
body,
description: self.parsed.description.clone(),
};

View File

@@ -1,6 +1,6 @@
use crate::parsing::{Implementation, NodeParsedField, ParsedField, ParsedFieldType, ParsedNodeFn, RegularParsedField};
use crate::parsing::{Implementation, NodeParsedField, ParsedField, ParsedFieldType, ParsedNodeFn, RegularParsedField, attr_marker, record_writes};
use proc_macro_error2::emit_error;
use quote::quote;
use quote::{ToTokens, quote};
use syn::spanned::Spanned;
use syn::{GenericParam, Type};
@@ -13,6 +13,7 @@ pub fn validate_node_fn(parsed: &ParsedNodeFn) -> syn::Result<()> {
validate_range_slider_bounds,
validate_async_source,
validate_lend_fields,
validate_record_io,
];
for validator in validators {
@@ -22,6 +23,74 @@ pub fn validate_node_fn(parsed: &ParsedNodeFn) -> syn::Result<()> {
Ok(())
}
fn validate_record_io(parsed: &ParsedNodeFn) {
let value = crate::codegen::slot_value_type(&parsed.output_type);
if let Type::Tuple(tuple) = &value {
let has_attr_slot = tuple.elems.iter().any(|slot| attr_marker(slot).is_some());
if has_attr_slot && record_writes(&value).is_none() {
emit_error!(
parsed.output_type.span(),
"a record return tuple is the element first, then only `Attr<..>` writes"
);
}
} else if attr_marker(&value).is_some() {
emit_error!(parsed.output_type.span(), "an `Attr<..>` write needs an element in the first tuple slot, e.g. `(T, Attr<..>)`");
}
let writes = record_writes(&value);
if parsed.attribute_reads.is_empty() && writes.is_none() {
return;
}
if parsed.is_async || crate::codegen::is_source_kernel(&parsed.output_type) {
emit_error!(parsed.output_type.span(), "attribute io is not supported on async source kernels");
}
if crate::codegen::is_poll_kernel(&parsed.output_type) {
emit_error!(parsed.output_type.span(), "attribute io needs a plain or `Result<_, Interrupt>` kernel, not a `GPoll` one");
}
match parsed.fields.first() {
None => emit_error!(
parsed.fn_name.span(),
"attribute io needs a value carrier as the first parameter after the context"
),
Some(carrier) => {
let valid = !carrier.is_data_field
&& matches!(&carrier.ty, ParsedFieldType::Regular(RegularParsedField { ty, lend: None, .. }) if !matches!(ty, Type::Tuple(tuple) if tuple.elems.is_empty()));
if !valid {
emit_error!(
carrier.pat_ident.span(),
"attribute io needs a value carrier as the first parameter after the context: an owned element type, not `()`, `#[data]`, `&T`, or `impl Node`"
);
}
}
}
for field in parsed.fields.iter().skip(1) {
if matches!(field.ty, ParsedFieldType::Node(_)) {
emit_error!(field.pat_ident.span(), "record nodes take no lazy inputs yet");
}
}
let mut seen_reads: Vec<String> = Vec::new();
for read in &parsed.attribute_reads {
let marker = read.marker.to_token_stream().to_string();
if seen_reads.contains(&marker) {
emit_error!(read.pat_ident.span(), "attribute `{}` is read twice", marker);
}
seen_reads.push(marker);
}
if let Some(writes) = &writes {
let mut seen_writes: Vec<String> = Vec::new();
for marker in &writes.markers {
let written = marker.to_token_stream().to_string();
if seen_writes.contains(&written) {
emit_error!(parsed.output_type.span(), "attribute `{}` is written twice", written);
}
seen_writes.push(written);
}
}
}
fn validate_async_source(parsed: &ParsedNodeFn) {
let snapshot_ctx = matches!(&parsed.input.ty, Type::Path(path) if path.path.segments.last().is_some_and(|segment| segment.ident == "CtxSnapshot"));
let future_kernel = crate::codegen::is_source_kernel(&parsed.output_type);
@@ -206,6 +275,8 @@ fn validate_primary_input_expose(parsed: &ParsedNodeFn) {
fn validate_implementations_for_generics(parsed: &ParsedNodeFn) {
let has_skip_impl = parsed.attributes.skip_impl;
let routing = crate::codegen::routing_io(parsed);
let routing_source = |ty: &Type| matches!((&routing, ty), (Some(routing), Type::Path(path)) if path.path.get_ident() == Some(&routing.generic));
if !has_skip_impl && !parsed.fn_generics.is_empty() {
for field in &parsed.fields {
@@ -217,6 +288,9 @@ fn validate_implementations_for_generics(parsed: &ParsedNodeFn) {
let pat_ident = &field.pat_ident;
match &field.ty {
ParsedFieldType::Regular(RegularParsedField { ty, implementations, .. }) => {
if routing_source(ty) {
continue;
}
if contains_generic_param(ty, &parsed.fn_generics) && implementations.is_empty() {
emit_error!(
ty.span(),
@@ -234,6 +308,9 @@ fn validate_implementations_for_generics(parsed: &ParsedNodeFn) {
implementations,
..
}) => {
if routing_source(output_type) {
continue;
}
if (contains_generic_param(input_type, &parsed.fn_generics) || contains_generic_param(output_type, &parsed.fn_generics)) && implementations.is_empty() {
emit_error!(
pat_ident.span(),