From 166241b0e8fc055b257c9511f54676e3b1d70fae Mon Sep 17 00:00:00 2001 From: Dennis Kobert Date: Thu, 6 Aug 2026 13:36:46 +0000 Subject: [PATCH] Flip nested-generic nodes onto record wires by unifying row types against field types --- node-graph/node-macro/src/codegen.rs | 82 ++++++++++++++++++++++------ 1 file changed, 66 insertions(+), 16 deletions(-) diff --git a/node-graph/node-macro/src/codegen.rs b/node-graph/node-macro/src/codegen.rs index 93002e5958..ca40e157c1 100644 --- a/node-graph/node-macro/src/codegen.rs +++ b/node-graph/node-macro/src/codegen.rs @@ -804,9 +804,20 @@ pub(crate) fn generate_node_impl(crate_ident: &CrateIdent, parsed: &ParsedNodeFn _ => true, }) .map(|param| match param { - // Flipped kernels clone lazy elements out of their records, so - // every element generic carries the bound wire values satisfy. - GenericParam::Type(type_param) if flip && Some(&type_param.ident) != ctx_param.map(|ctx_param| &ctx_param.ident) => { + // Flipped kernels clone bare-typed elements out of their records, + // so those generics carry the bound wire values satisfy; a + // generic only nested in a field's type stays as declared. + GenericParam::Type(type_param) + if flip + && Some(&type_param.ident) != ctx_param.map(|ctx_param| &ctx_param.ident) + && parsed.fields.iter().filter(|field| !field.is_data_field).any(|field| { + let ty = match &field.ty { + ParsedFieldType::Regular(RegularParsedField { ty, .. }) => ty, + ParsedFieldType::Node(NodeParsedField { output_type, .. }) => output_type, + }; + matches!(ty, Type::Path(path) if path.qself.is_none() && path.path.get_ident() == Some(&type_param.ident)) + }) => + { let mut bounded = type_param.clone(); bounded.bounds.push(syn::parse_quote!(::core::clone::Clone)); quote!(#bounded) @@ -1898,17 +1909,18 @@ pub(crate) fn record_flip(parsed: &ParsedNodeFn) -> bool { for param in &parsed.fn_generics { match param { GenericParam::Type(type_param) if Some(&type_param.ident) == ctx_ident.as_ref() => {} - // Registry rows assign a generic from a field it names bare, so a - // generic without such a position keeps the plain lowering. + // Registry rows assign a generic by unifying a field's type with + // the row's, so a generic without an extractable position keeps + // the plain lowering. GenericParam::Type(type_param) => { - let bare = parsed.fields.iter().any(|field| { + let extractable = parsed.fields.iter().filter(|field| !field.is_data_field).any(|field| { let ty = match &field.ty { ParsedFieldType::Regular(RegularParsedField { ty, .. }) => ty, ParsedFieldType::Node(NodeParsedField { output_type, .. }) => output_type, }; - matches!(ty, Type::Path(path) if path.qself.is_none() && path.path.get_ident() == Some(&type_param.ident)) + generic_extractable(ty, &type_param.ident) }); - if !bare { + if !extractable { return false; } } @@ -1918,6 +1930,43 @@ pub(crate) fn record_flip(parsed: &ParsedNodeFn) -> bool { true } +/// Whether unifying a value of `field_ty`'s shape can bind `generic`: the +/// generic sits bare or under path type arguments, the shapes +/// [`generic_assignment`] walks. +fn generic_extractable(field_ty: &Type, generic: &Ident) -> bool { + match field_ty { + Type::Path(path) if path.qself.is_none() && path.path.get_ident() == Some(generic) => true, + Type::Path(path) => path.path.segments.iter().any(|segment| match &segment.arguments { + PathArguments::AngleBracketed(args) => args.args.iter().any(|argument| match argument { + GenericArgument::Type(inner) => generic_extractable(inner, generic), + _ => false, + }), + _ => false, + }), + _ => false, + } +} + +/// Binds `generic` by unifying `field_ty` against `row_ty`: where the field +/// names the generic, the row's corresponding subtree is the assignment. +fn generic_assignment(field_ty: &Type, row_ty: &Type, generic: &Ident) -> Option { + if matches!(field_ty, Type::Path(path) if path.qself.is_none() && path.path.get_ident() == Some(generic)) { + return Some(row_ty.clone()); + } + let (Type::Path(field_path), Type::Path(row_path)) = (field_ty, row_ty) else { + return None; + }; + let field_segment = field_path.path.segments.last()?; + let row_segment = row_path.path.segments.last()?; + let (PathArguments::AngleBracketed(field_args), PathArguments::AngleBracketed(row_args)) = (&field_segment.arguments, &row_segment.arguments) else { + return None; + }; + field_args.args.iter().zip(row_args.args.iter()).find_map(|(field_arg, row_arg)| match (field_arg, row_arg) { + (GenericArgument::Type(field_inner), GenericArgument::Type(row_inner)) => generic_assignment(field_inner, row_inner, generic), + _ => None, + }) +} + pub(crate) fn is_record_value(ty: &Type) -> bool { matches!(ty, Type::Path(path) if path.path.segments.last().is_some_and(|segment| segment.ident == "RecordValue")) } @@ -2206,6 +2255,10 @@ fn flip_entries_tokens(parsed: &ParsedNodeFn, struct_name: &Ident, regular_field } let output = slot_value_type(&parsed.output_type); + let field_type = |field: &ParsedField| match &field.ty { + ParsedFieldType::Regular(RegularParsedField { ty, .. }) => ty.clone(), + ParsedFieldType::Node(NodeParsedField { output_type, .. }) => output_type.clone(), + }; let ctx_ident = context_param(parsed).map(|ctx| ctx.ident.clone()); let generic_positions: Option> = parsed .fn_generics @@ -2217,13 +2270,7 @@ fn flip_entries_tokens(parsed: &ParsedNodeFn, struct_name: &Ident, regular_field .map(|generic| { regular_fields .iter() - .position(|field| { - let ty = match &field.ty { - ParsedFieldType::Regular(RegularParsedField { ty, .. }) => ty, - ParsedFieldType::Node(NodeParsedField { output_type, .. }) => output_type, - }; - matches!(ty, Type::Path(path) if path.qself.is_none() && path.path.get_ident() == Some(generic)) - }) + .position(|field| generic_extractable(&field_type(field), generic)) .map(|index| (generic.clone(), index)) }) .collect(); @@ -2287,7 +2334,10 @@ fn flip_entries_tokens(parsed: &ParsedNodeFn, struct_name: &Ident, regular_field }; let entries = rows.iter().filter_map(|row| { - let assignments: Vec<(Ident, Type)> = generic_positions.iter().map(|(generic, index)| (generic.clone(), row[*index].clone())).collect(); + let assignments: Vec<(Ident, Type)> = generic_positions + .iter() + .map(|(generic, index)| generic_assignment(&field_type(regular_fields[*index]), &row[*index], generic).map(|assigned| (generic.clone(), assigned))) + .collect::>()?; if type_disqualifies(&substitute_ident_types(&output, &assignments)) { return None; }