mirror of
https://github.com/GraphiteEditor/Graphite.git
synced 2026-09-15 22:28:10 +08:00
Add the attribute census, node macro attribute io over list wires, and the rank-0 record tier
This commit is contained in:
@@ -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, ®ular_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 {
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -317,6 +317,7 @@ impl PerPixelAdjustCodegen<'_> {
|
||||
output_type: raster_gpu,
|
||||
is_async: false,
|
||||
fields,
|
||||
attribute_reads: Vec::new(),
|
||||
body,
|
||||
description: self.parsed.description.clone(),
|
||||
};
|
||||
|
||||
@@ -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(),
|
||||
|
||||
Reference in New Issue
Block a user