diff --git a/node-graph/node-macro/src/codegen.rs b/node-graph/node-macro/src/codegen.rs index ff789593e6..db5c8fca90 100644 --- a/node-graph/node-macro/src/codegen.rs +++ b/node-graph/node-macro/src/codegen.rs @@ -12,7 +12,7 @@ use crate::shader_nodes::{ShaderCodegen, ShaderTokens}; mod classify; mod entries; -mod ir; +pub(crate) mod ir; mod metadata; pub(crate) use classify::*; use entries::entries_tokens; @@ -681,7 +681,7 @@ pub(crate) struct NodeFields<'a> { pub(crate) struct_type_params: Vec, } -pub(crate) fn generate_node_impl(crate_ident: &CrateIdent, parsed: &ParsedNodeFn, model: &Option, fields: NodeFields) -> syn::Result { +pub(crate) fn generate_node_impl(crate_ident: &CrateIdent, parsed: &ParsedNodeFn, model: &Option, fields: NodeFields) -> syn::Result { let core_types = crate_ident.gcore()?; let ctx_param = context_param(parsed); @@ -692,8 +692,8 @@ pub(crate) fn generate_node_impl(crate_ident: &CrateIdent, parsed: &ParsedNodeFn let Some(model) = model.as_ref() else { return Ok(NodePlan::default()); }; - let async_fn = matches!(model.dialect, Dialect::AsyncFn); - let future_kernel = matches!(model.dialect, Dialect::Future | Dialect::FutureInterrupt); + let async_fn = matches!(*model, Dialect::AsyncFn); + let future_kernel = matches!(*model, Dialect::Future | Dialect::FutureInterrupt); let async_source = async_fn || future_kernel; let node = crate::codegen::ir::build(parsed); let kind = crate::codegen::ir::node_kind(&node); @@ -826,7 +826,7 @@ pub(crate) fn generate_node_impl(crate_ident: &CrateIdent, parsed: &ParsedNodeFn (false, None) if flip => syn::parse_quote!(#core_types::record::RecordValue<'__record>), (false, None) => slot_value_type(&parsed.output_type), }; - let raw_lazy = matches!(model.dialect, Dialect::Poll); + let raw_lazy = matches!(*model, Dialect::Poll); let injected_name = |ident: &Ident| async_source && (ident == "_runtime" || ident == "_source"); let where_predicates: Vec = parsed.where_clause.iter().flat_map(|clause| clause.predicates.iter()).map(|predicate| quote!(#predicate)).collect(); @@ -1294,7 +1294,7 @@ pub(crate) fn generate_node_impl(crate_ident: &CrateIdent, parsed: &ParsedNodeFn false => quote!(#core_types::node::StatusCell::new()), }; let kernel_call = quote!(self::#fn_name(__input #(, &self.#data_names)* #(, #call_args)*)); - let lift = match model.dialect { + let lift = match *model { Dialect::Interrupt => quote! { match #kernel_call { Ok(value) => __cell.finish(value), @@ -1443,7 +1443,7 @@ pub(crate) fn generate_node_impl(crate_ident: &CrateIdent, parsed: &ParsedNodeFn true => Vec::new(), false => reads_of(0).into_iter().map(|(slot, read)| read_binding(slot, read, quote!(__src_rec))).collect(), }; - let kernel_value = match model.dialect { + let kernel_value = match *model { Dialect::Interrupt => quote! { match #record_kernel_call { Ok(__value) => __value, @@ -1507,7 +1507,7 @@ pub(crate) fn generate_node_impl(crate_ident: &CrateIdent, parsed: &ParsedNodeFn } }); let flip_tail = flip.then(|| { - if matches!(model.dialect, Dialect::Poll) { + if matches!(*model, Dialect::Poll) { return match &carried_prelude { Some(prelude) => quote! { #prelude @@ -1523,7 +1523,7 @@ pub(crate) fn generate_node_impl(crate_ident: &CrateIdent, parsed: &ParsedNodeFn }, }; } - let kernel_value = match model.dialect { + let kernel_value = match *model { Dialect::Interrupt => quote! { match #kernel_call { Ok(value) => value, @@ -1594,7 +1594,7 @@ pub(crate) fn generate_node_impl(crate_ident: &CrateIdent, parsed: &ParsedNodeFn ), None => (quote!(), pending_return.clone()), }; - let acquire = match model.dialect { + let acquire = match *model { Dialect::FutureInterrupt => quote! { let __future = match #kernel_call { Ok(future) => future, diff --git a/node-graph/node-macro/src/codegen/classify.rs b/node-graph/node-macro/src/codegen/classify.rs index ef050174b9..eff9ac62b6 100644 --- a/node-graph/node-macro/src/codegen/classify.rs +++ b/node-graph/node-macro/src/codegen/classify.rs @@ -1,53 +1,32 @@ use super::*; -/// How a record node's primary input lowers. +/// How a record node's primary input lowers: `None` writes a fresh record, +/// `Token` carries the element bytes through as `ElToken`, `Read` reads a +/// concrete element at offset 0. The element and write set fold from the IR. #[derive(Clone)] pub(crate) enum RecordCarrier { - /// `_: ()`: no carrier edge, the kernel writes a fresh record. None, - /// An unbounded generic returned in the element position: the element - /// bytes carry through the copy plan and the kernel sees `ElToken`. - Token(Ident), - /// An element type read at offset 0, monomorphized per its - /// implementations list where generic. - Read(Type), + Token, + Read, } -/// The record io of a node fn: how the carrier lowers, the element write, -/// and the markers written and removed. Present exactly when the signature -/// declares attribute reads or writes in a shape the record tier supports; -/// malformed record io is reported by validation and generates no node impl. +/// A well-formed record-io node: only the carrier form is retained, so +/// [`skips_carrier`] can gate the fresh-record path. Malformed record io yields +/// `None` from [`record_shape`] and generates no node impl. #[derive(Clone)] pub(crate) struct RecordShape { pub(crate) carrier: RecordCarrier, - pub(crate) element_write: Option, - pub(crate) write_markers: Vec, - pub(crate) removes: Vec, } impl RecordShape { pub(crate) fn skips_carrier(&self) -> bool { matches!(self.carrier, RecordCarrier::None) } - - pub(crate) fn carries_element(&self) -> bool { - self.element_write.is_none() - } -} - -/// The record-tier lowering a node fn resolves to. Exactly one class per node, -/// computed once by [`analyze`]; every downstream fragment reads the class -/// instead of recomputing the classification predicates. -pub(crate) enum Class { - RecordIo(RecordShape), - Routing(RoutingIo), - Flip { carrier: bool }, - Opaque, } /// The effect/return axis of a node's kernel, resolved once from the signature. -/// Orthogonal to [`Class`]: it selects the eval tail (finish / merge / spawn) -/// and the kernel signature wrapping across every class. +/// It selects the eval tail (finish / merge / spawn) and the kernel signature +/// wrapping across every node kind. #[derive(Clone, Copy, PartialEq)] pub(crate) enum Dialect { Sync, @@ -71,33 +50,22 @@ pub(crate) fn dialect(parsed: &ParsedNodeFn) -> Dialect { } } -/// The result of classifying a node fn. A node with no supported lowering -/// (an async node with lazy inputs, malformed record io, or a signature no -/// class accepts) yields `None` and generates a struct and metadata but no -/// `Node` impl. -pub(crate) struct NodeModel { - pub(crate) class: Class, - pub(crate) dialect: Dialect, -} - -pub(crate) fn analyze(parsed: &ParsedNodeFn) -> Option { +/// The dialect of a node fn that lowers to a `Node` impl, or `None` when no +/// lowering supports the signature (an async node with lazy inputs, malformed +/// record io, or a shape no kind accepts). The kind itself is derived from the +/// intent IR ([`crate::codegen::ir::node_kind`]); this only gates support. +pub(crate) fn analyze(parsed: &ParsedNodeFn) -> Option { if parsed.is_async && parsed.fields.iter().any(|field| matches!(field.ty, ParsedFieldType::Node(_))) { return None; } - let class = if let Some(shape) = record_shape(parsed) { - Class::RecordIo(shape) + let supported = if record_shape(parsed).is_some() { + true } else if has_record_io(parsed) { return None; - } else if let Some(routing) = routing_io(parsed) { - Class::Routing(routing) - } else if record_flip(parsed) { - Class::Flip { carrier: flip_carrier(parsed) } - } else if record_opaque(parsed) { - Class::Opaque } else { - return None; + routing_io(parsed).is_some() || record_flip(parsed) || record_opaque(parsed) }; - Some(NodeModel { class, dialect: dialect(parsed) }) + supported.then(|| dialect(parsed)) } /// The tail form of a node's eval, selected from its class and dialect: forward @@ -317,43 +285,43 @@ pub(crate) fn record_shape(parsed: &ParsedNodeFn) -> Option { let ParsedFieldType::Regular(RegularParsedField { ty, lend: None, implementations, .. }) = &carrier_field.ty else { return None; }; - let carrier = match ty { - Type::Tuple(tuple) if tuple.elems.is_empty() => RecordCarrier::None, + let token = match ty { + Type::Tuple(tuple) if tuple.elems.is_empty() => None, ty => match implementations.is_empty().then(|| unbounded_generic(parsed, ty)).flatten() { - Some(token) => RecordCarrier::Token(token), + Some(token) => Some(token), None => { if contains_open_generic(parsed, ty) { return None; } - RecordCarrier::Read(ty.clone()) + None } }, }; - let (element, write_markers, removes) = match writes { + let carrier = match ty { + Type::Tuple(tuple) if tuple.elems.is_empty() => RecordCarrier::None, + _ if token.is_some() => RecordCarrier::Token, + _ => RecordCarrier::Read, + }; + let (element, _, removes) = match writes { Some(RecordWrites { element, markers, removes }) => (element, markers, removes), None => (value, Vec::new(), Vec::new()), }; - let element_write = match &carrier { - RecordCarrier::Token(token) => match bare_ident(&element) { - Some(ident) if ident == token => None, - _ => return None, - }, - _ => { + match &token { + Some(token) => { + if !matches!(bare_ident(&element), Some(ident) if ident == token) { + return None; + } + } + None => { if contains_open_generic(parsed, &element) { return None; } - Some(element) } - }; + } if matches!(carrier, RecordCarrier::None) && !removes.is_empty() { return None; } - Some(RecordShape { - carrier, - element_write, - write_markers, - removes, - }) + Some(RecordShape { carrier }) } pub(crate) fn is_poll_kernel(output: &Type) -> bool { diff --git a/node-graph/node-macro/src/codegen/ir.rs b/node-graph/node-macro/src/codegen/ir.rs index e1b67a14e8..6e54d9cb7b 100644 --- a/node-graph/node-macro/src/codegen/ir.rs +++ b/node-graph/node-macro/src/codegen/ir.rs @@ -391,7 +391,7 @@ pub(crate) enum Effect { #[cfg(test)] mod tests { use super::*; - use crate::codegen::classify::{Class, Dialect, analyze, context_param, dialect}; + use crate::codegen::classify::{Dialect, analyze, context_param, dialect, record_flip, record_opaque, unbounded_generic}; use crate::parsing::parse_node_fn; use proc_macro2::TokenStream as TokenStream2; use quote::{ToTokens, quote}; @@ -427,53 +427,87 @@ mod tests { } } - fn facts_from_class(class: &Class, fields: &[&ParsedField]) -> Facts { + /// The kinds a supported node resolves to, from the classify predicates in + /// `analyze`'s order; the frozen oracle the IR's `node_kind` must reproduce. + struct Kinds { + record_io: bool, + routing: bool, + flip: bool, + opaque: bool, + } + + fn kinds(parsed: &ParsedNodeFn) -> Kinds { + let record_io = record_shape(parsed).is_some(); + let routing = !record_io && routing_io(parsed).is_some(); + let flip = !record_io && !routing && record_flip(parsed); + let opaque = !record_io && !routing && !flip && record_opaque(parsed); + Kinds { record_io, routing, flip, opaque } + } + + fn skips_carrier(parsed: &ParsedNodeFn) -> bool { + record_shape(parsed).is_some_and(|shape| shape.skips_carrier()) + } + + fn routing_generic(parsed: &ParsedNodeFn) -> Option { + kinds(parsed).routing.then(|| routing_io(parsed).map(|routing| routing.generic)).flatten() + } + + fn token_carrier(parsed: &ParsedNodeFn) -> bool { + let element = record_writes(&slot_value_type(&parsed.output_type)).map_or_else(|| slot_value_type(&parsed.output_type), |writes| writes.element); + kinds(parsed).record_io && unbounded_generic(parsed, &element).is_some() + } + + fn facts_from_signature(parsed: &ParsedNodeFn) -> Facts { + let fields: Vec<&ParsedField> = parsed.fields.iter().filter(|field| !field.is_data_field).collect(); let source_ty = |field: &ParsedField| match &field.ty { ParsedFieldType::Node(NodeParsedField { output_type, .. }) => output_type.clone(), ParsedFieldType::Regular(RegularParsedField { ty, .. }) => ty.clone(), }; - match class { - Class::Flip { carrier } => Facts { - sources: if *carrier { vec![0] } else { vec![] }, + let kinds = kinds(parsed); + if kinds.flip { + Facts { + sources: if flip_carrier(parsed) { vec![0] } else { vec![] }, carried: false, writes: vec![], removes: vec![], delta: 0, - }, - Class::Opaque => { - let record = fields.iter().position(|field| matches!(&field.ty, ParsedFieldType::Node(NodeParsedField { output_type, .. }) if is_record_value(output_type))); - Facts { - sources: record.into_iter().collect(), - carried: true, - writes: vec![], - removes: vec![], - delta: 0, - } } - Class::Routing(routing) => Facts { - sources: fields.iter().enumerate().filter(|(_, field)| bare_ident(&source_ty(field)) == Some(&routing.generic)).map(|(index, _)| index).collect(), + } else if kinds.opaque { + let record = fields.iter().position(|field| matches!(&field.ty, ParsedFieldType::Node(NodeParsedField { output_type, .. }) if is_record_value(output_type))); + Facts { + sources: record.into_iter().collect(), carried: true, writes: vec![], removes: vec![], delta: 0, - }, - Class::RecordIo(shape) => Facts { - sources: if shape.skips_carrier() { vec![] } else { vec![0] }, - carried: shape.carries_element(), - writes: markers(&shape.write_markers), - removes: markers(&shape.removes), + } + } else if kinds.routing { + let generic = routing_generic(parsed).expect("routing has a generic"); + Facts { + sources: fields.iter().enumerate().filter(|(_, field)| bare_ident(&source_ty(field)) == Some(&generic)).map(|(index, _)| index).collect(), + carried: true, + writes: vec![], + removes: vec![], delta: 0, - }, + } + } else { + let (write_markers, removes) = record_writes(&slot_value_type(&parsed.output_type)).map_or((Vec::new(), Vec::new()), |writes| (writes.markers, writes.removes)); + Facts { + sources: if skips_carrier(parsed) { vec![] } else { vec![0] }, + carried: token_carrier(parsed), + writes: markers(write_markers.iter()), + removes: markers(removes.iter()), + delta: 0, + } } } fn assert_bridge(attr: TokenStream2, item: TokenStream2) -> Node { let mut parsed = parse_node_fn(attr, item).unwrap(); parsed.replace_impl_trait_in_input(); - let model = analyze(&parsed).expect("representative resolves to a class"); - let fields: Vec<&ParsedField> = parsed.fields.iter().filter(|field| !field.is_data_field).collect(); + analyze(&parsed).expect("representative resolves to a supported node"); let node = build(&parsed); - assert_eq!(facts_from_ir(&node), facts_from_class(&model.class, &fields)); + assert_eq!(facts_from_ir(&node), facts_from_signature(&parsed)); node } @@ -536,15 +570,18 @@ mod tests { } /// The frozen `field_role` classification the IR bindings must reproduce. - fn reference_label(parsed: &ParsedNodeFn, class: &Class, raw: bool, index: usize, field: &ParsedField) -> &'static str { - let record = matches!(class, Class::RecordIo(_)); - let skips_carrier = matches!(class, Class::RecordIo(shape) if shape.skips_carrier()); - let carrier_flip = matches!(class, Class::Flip { carrier: true }); - let flip = matches!(class, Class::Flip { .. }); - let opaque = matches!(class, Class::Opaque); - let routing = matches!(class, Class::Routing(_)); + fn reference_label(parsed: &ParsedNodeFn, raw: bool, index: usize, field: &ParsedField) -> &'static str { + let Kinds { + record_io: record, + routing, + flip, + opaque, + } = kinds(parsed); + let skips_carrier = skips_carrier(parsed); + let carrier_flip = flip && flip_carrier(parsed); let derives = ctx_derives(parsed); - let routing_source = |ty: &Type| matches!(class, Class::Routing(routing) if bare_ident(ty) == Some(&routing.generic)); + let generic = routing_generic(parsed); + let routing_source = |ty: &Type| generic.as_ref().is_some_and(|generic| bare_ident(ty) == Some(generic)); match &field.ty { ParsedFieldType::Regular(RegularParsedField { ty, lend, .. }) => { if record && !skips_carrier && index == 0 { @@ -602,14 +639,18 @@ mod tests { fn assert_bindings(attr: TokenStream2, item: TokenStream2) { let mut parsed = parse_node_fn(attr, item).unwrap(); parsed.replace_impl_trait_in_input(); - let model = analyze(&parsed).expect("representative resolves to a class"); + analyze(&parsed).expect("representative resolves to a supported node"); let raw = matches!(dialect(&parsed), Dialect::Poll); let node = build(&parsed); - let expected_kind = match &model.class { - Class::RecordIo(_) => "record-io", - Class::Flip { .. } => "flip", - Class::Routing(_) => "routing", - Class::Opaque => "opaque", + let kinds = kinds(&parsed); + let expected_kind = if kinds.record_io { + "record-io" + } else if kinds.routing { + "routing" + } else if kinds.flip { + "flip" + } else { + "opaque" }; let actual_kind = match node_kind(&node) { NodeKind::RecordIo => "record-io", @@ -622,7 +663,7 @@ mod tests { for (index, field) in fields.iter().enumerate() { assert_eq!( ir_label(&node, index, field, raw), - reference_label(&parsed, &model.class, raw, index, field), + reference_label(&parsed, raw, index, field), "field {index} of {}", parsed.fn_name ); diff --git a/node-graph/node-macro/src/validation.rs b/node-graph/node-macro/src/validation.rs index a712733316..4716e38ade 100644 --- a/node-graph/node-macro/src/validation.rs +++ b/node-graph/node-macro/src/validation.rs @@ -393,10 +393,11 @@ 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 record_token = crate::codegen::record_shape(parsed).and_then(|shape| match shape.carrier { - crate::codegen::RecordCarrier::Token(token) => Some(token), + let node = crate::codegen::ir::build(parsed); + let record_token = match (crate::codegen::ir::node_kind(&node), &node.output.shape.element) { + (crate::codegen::ir::NodeKind::RecordIo, crate::codegen::ir::Element::Generic(ident)) => Some(ident.clone()), _ => None, - }); + }; let opaque_record_generic = |ty: &Type| { let ident = match ty { Type::Path(path) => path.path.get_ident(),