From 15f4aa8f987c44bbdb17bf7260442a72514d0464 Mon Sep 17 00:00:00 2001 From: Dennis Kobert Date: Sun, 16 Aug 2026 11:27:13 +0000 Subject: [PATCH] Strip IList nesting from kernel returns at parse time --- node-graph/node-macro/src/codegen.rs | 4 +-- node-graph/node-macro/src/codegen/ir.rs | 34 +++++++++++++++++-- node-graph/node-macro/src/parsing.rs | 11 +++++- .../src/shader_nodes/per_pixel_adjust.rs | 1 + 4 files changed, 45 insertions(+), 5 deletions(-) diff --git a/node-graph/node-macro/src/codegen.rs b/node-graph/node-macro/src/codegen.rs index 03412df26c..9988d830ed 100644 --- a/node-graph/node-macro/src/codegen.rs +++ b/node-graph/node-macro/src/codegen.rs @@ -734,7 +734,7 @@ pub(crate) fn generate_node_impl(crate_ident: &CrateIdent, parsed: &ParsedNodeFn let subject_depth = node.inputs.iter().find(|input| input.subject).map_or(0, |input| input.shape.depth); let level_delta = node.output.shape.depth as i8 - subject_depth as i8; let pushed_levels = level_delta.max(0) as u8; - let output_row = crate::codegen::ir::strip_ilist(&slot_value_type(&parsed.output_type)).0; + let output_row = slot_value_type(&parsed.output_type); 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 = match ctx_param { @@ -1270,7 +1270,7 @@ pub(crate) fn generate_node_impl(crate_ident: &CrateIdent, parsed: &ParsedNodeFn // A bare `Attr` in the return type cannot elide its lifetime, so the // kernel gets a fresh one; reference-valued writes name their real // lifetime explicitly and pass through untouched. - let kernel_output = record_io.then(|| inject_attr_lifetimes(&crate::codegen::ir::strip_ilist(&parsed.output_type).0)).flatten(); + let kernel_output = record_io.then(|| inject_attr_lifetimes(&parsed.output_type)).flatten(); let attr_lifetime = kernel_output.is_some().then(|| quote!('__attr,)); let kernel_output = match derive_routing { true => { diff --git a/node-graph/node-macro/src/codegen/ir.rs b/node-graph/node-macro/src/codegen/ir.rs index d93a71b58e..013eee094a 100644 --- a/node-graph/node-macro/src/codegen/ir.rs +++ b/node-graph/node-macro/src/codegen/ir.rs @@ -82,8 +82,8 @@ fn subject(index: usize, field: &ParsedField, carrier_subject: bool, routing: Op } fn output(parsed: &ParsedNodeFn, generics: &[Ident]) -> Output { - // Strip the `IList` rank markers first so the write set is read from the row. - let (row, depth) = strip_ilist(&slot_value_type(&parsed.output_type)); + let row = slot_value_type(&parsed.output_type); + let depth = parsed.output_depth; let (element, writes, removes) = match record_writes(&row) { Some(RecordWrites { element, markers, removes }) => (element, markers, removes), None => (row, Vec::new(), Vec::new()), @@ -167,6 +167,36 @@ pub(crate) fn strip_ilist(ty: &Type) -> (Type, u8) { (element, depth) } +/// Strips `IList` rank nesting from the output's value position, preserving the +/// dialect wrapper (`Result`/`GPoll`), and returns the removed depth. +pub(crate) fn strip_output_rank(output: &Type) -> (Type, u8) { + use crate::codegen::classify::{KernelKind, kernel_kind}; + match kernel_kind(output) { + KernelKind::Plain => strip_ilist(output), + KernelKind::Interrupt(inner) | KernelKind::Poll(inner) => { + let (row, depth) = strip_ilist(&inner); + (replace_first_type_arg(output, row), depth) + } + KernelKind::Future(_) | KernelKind::FutureInterrupt(_) => (output.clone(), 0), + } +} + +fn replace_first_type_arg(ty: &Type, replacement: Type) -> Type { + let mut ty = ty.clone(); + if let Type::Path(path) = &mut ty + && let Some(segment) = path.path.segments.last_mut() + && let PathArguments::AngleBracketed(args) = &mut segment.arguments + { + for arg in args.args.iter_mut() { + if let GenericArgument::Type(inner) = arg { + *inner = replacement; + break; + } + } + } + ty +} + fn ilist_inner(ty: &Type) -> Option { let Type::Path(path) = ty else { return None }; let segment = path.path.segments.last()?; diff --git a/node-graph/node-macro/src/parsing.rs b/node-graph/node-macro/src/parsing.rs index 06e33034c2..b751bbe265 100644 --- a/node-graph/node-macro/src/parsing.rs +++ b/node-graph/node-macro/src/parsing.rs @@ -33,6 +33,7 @@ pub(crate) struct ParsedNodeFn { pub(crate) where_clause: Option, pub(crate) input: Input, pub(crate) output_type: Type, + pub(crate) output_depth: u8, pub(crate) is_async: bool, pub(crate) fields: Vec, pub(crate) body: TokenStream2, @@ -662,7 +663,7 @@ pub(crate) fn parse_node_fn(attr: TokenStream2, item: TokenStream2) -> syn::Resu let is_async = input_fn.sig.asyncness.is_some(); let (input, fields) = parse_inputs(&input_fn.sig.inputs)?; - let output_type = parse_output(&input_fn.sig.output)?; + 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 @@ -691,6 +692,7 @@ pub(crate) fn parse_node_fn(attr: TokenStream2, item: TokenStream2) -> syn::Resu fn_generics, input, output_type, + output_depth, is_async, fields, where_clause, @@ -1407,6 +1409,7 @@ mod tests { context_features: vec![], }, output_type: parse_quote!(f64), + output_depth: 0, is_async: false, fields: vec![ParsedField { pat_ident: pat_ident("b"), @@ -1485,6 +1488,7 @@ mod tests { context_features: vec![], }, output_type: parse_quote!(T), + output_depth: 0, is_async: false, fields: vec![ ParsedField { @@ -1578,6 +1582,7 @@ mod tests { context_features: vec![format_ident!("ExtractFootprint")], }, output_type: parse_quote!(Vector), + output_depth: 0, is_async: false, fields: vec![ParsedField { pat_ident: pat_ident("radius"), @@ -1652,6 +1657,7 @@ mod tests { context_features: vec![], }, output_type: parse_quote!(List>), + output_depth: 0, is_async: false, fields: vec![ParsedField { pat_ident: pat_ident("shadows"), @@ -1738,6 +1744,7 @@ mod tests { context_features: vec![], }, output_type: parse_quote!(f64), + output_depth: 0, is_async: false, fields: vec![ParsedField { pat_ident: pat_ident("b"), @@ -1827,6 +1834,7 @@ mod tests { context_features: vec![], }, output_type: parse_quote!(List>), + output_depth: 0, is_async: true, fields: vec![ParsedField { pat_ident: pat_ident("path"), @@ -1901,6 +1909,7 @@ mod tests { context_features: vec![], }, output_type: parse_quote!(i32), + output_depth: 0, is_async: false, fields: vec![], body: TokenStream2::new(), diff --git a/node-graph/node-macro/src/shader_nodes/per_pixel_adjust.rs b/node-graph/node-macro/src/shader_nodes/per_pixel_adjust.rs index 1e0907e880..3ff04d957f 100644 --- a/node-graph/node-macro/src/shader_nodes/per_pixel_adjust.rs +++ b/node-graph/node-macro/src/shader_nodes/per_pixel_adjust.rs @@ -317,6 +317,7 @@ impl PerPixelAdjustCodegen<'_> { context_features: self.parsed.input.context_features.clone(), }, output_type: raster_gpu, + output_depth: 0, is_async: false, fields, body,