diff --git a/node-graph/libraries/core-types/src/gpoll.rs b/node-graph/libraries/core-types/src/gpoll.rs index 547103b9d8..071010fb16 100644 --- a/node-graph/libraries/core-types/src/gpoll.rs +++ b/node-graph/libraries/core-types/src/gpoll.rs @@ -176,6 +176,26 @@ impl Extent { (Extent::Exactly(n), Extent::Exactly(m)) => GPoll::fallback(Extent::Exactly(n.min(m)), "extent mismatch"), }) } + + /// The product of two extents, used to compose nested-level counts; an + /// unbounded operand leaves the product unbounded. + pub fn mul(a: GPoll, b: GPoll) -> GPoll { + a.zip(b).map(|(a, b)| match (a, b) { + (Extent::Exactly(n), Extent::Exactly(m)) => Extent::Exactly(n * m), + _ => Extent::Free, + }) + } +} + +/// A query over a node's nesting levels: one level, the product below or above +/// it, or the whole domain. The composite [`Node::extent`](crate::node::Node::extent) +/// derives these from the per-level [`extent_at`](crate::node::Node::extent_at). +#[derive(Clone, Copy, Debug)] +pub enum Level { + At(u8), + Below(u8), + Above(u8), + Total, } #[derive(Clone, Copy, Debug, PartialEq, Eq)] diff --git a/node-graph/libraries/core-types/src/node.rs b/node-graph/libraries/core-types/src/node.rs index f49c801652..577ea89560 100644 --- a/node-graph/libraries/core-types/src/node.rs +++ b/node-graph/libraries/core-types/src/node.rs @@ -1,5 +1,5 @@ use crate::context::InjectIndex; -use crate::gpoll::{Extent, Finality, GPoll, GraphError, Interrupt}; +use crate::gpoll::{Extent, Finality, GPoll, GraphError, Interrupt, Level}; use std::cell::Cell; use std::mem::MaybeUninit; use std::ops::Range; @@ -177,18 +177,27 @@ pub trait Node { fn eval(&self, input: &Input) -> GPoll; - fn extent(&self, _input: &Input) -> GPoll { - GPoll::Final(Extent::Free) - } - /// The count of items at one absolute nesting level (innermost `0`). The /// leveled primitive a structure node overrides to report a pushed level's /// size; the scalar base is one item at every level. Uncertainty rides the - /// `GPoll` status axis, as with [`extent`](Node::extent). + /// `GPoll` status axis. fn extent_at(&self, _input: &Input, _level: u8) -> GPoll { GPoll::Final(Extent::Exactly(1)) } + /// The composite domain query derived from [`extent_at`](Node::extent_at): + /// one level, the product of the levels below or above it, or the whole + /// domain's flat count. Consumers query this; nodes only write `extent_at`. + fn extent(&self, input: &Input, at: Level) -> GPoll { + let product = |range: core::ops::Range| range.fold(GPoll::Final(Extent::Exactly(1)), |acc, level| Extent::mul(acc, self.extent_at(input, level))); + match at { + Level::At(level) => self.extent_at(input, level), + Level::Below(level) => product(0..level), + Level::Above(level) => product((level + 1)..self.depth()), + Level::Total => product(0..self.depth()), + } + } + /// The node's domain depth (number of nesting levels; `0` = scalar), baked /// into the record layout at wiring. fn depth(&self) -> u8 { @@ -262,10 +271,6 @@ where (**self).eval(input) } - fn extent(&self, input: &Input) -> GPoll { - (**self).extent(input) - } - fn extent_at(&self, input: &Input, level: u8) -> GPoll { (**self).extent_at(input, level) } @@ -296,10 +301,6 @@ where (**self).eval(input) } - fn extent(&self, input: &Input) -> GPoll { - (**self).extent(input) - } - fn extent_at(&self, input: &Input, level: u8) -> GPoll { (**self).extent_at(input, level) } @@ -330,10 +331,6 @@ where (**self).eval(input) } - fn extent(&self, input: &Input) -> GPoll { - (**self).extent(input) - } - fn extent_at(&self, input: &Input, level: u8) -> GPoll { (**self).extent_at(input, level) } diff --git a/node-graph/libraries/core-types/src/registry.rs b/node-graph/libraries/core-types/src/registry.rs index bbbaa1d71c..2e02bc5548 100644 --- a/node-graph/libraries/core-types/src/registry.rs +++ b/node-graph/libraries/core-types/src/registry.rs @@ -165,11 +165,6 @@ where result } - fn extent(&self, input: &Input) -> crate::gpoll::GPoll { - // SAFETY: as in eval. - unsafe { self.ptr.as_ref() }.extent(input) - } - fn extent_at(&self, input: &Input, level: u8) -> crate::gpoll::GPoll { // SAFETY: as in eval. unsafe { self.ptr.as_ref() }.extent_at(input, level) diff --git a/node-graph/node-macro/src/codegen.rs b/node-graph/node-macro/src/codegen.rs index dac91119b9..710f71adbb 100644 --- a/node-graph/node-macro/src/codegen.rs +++ b/node-graph/node-macro/src/codegen.rs @@ -1214,31 +1214,16 @@ pub(crate) fn generate_node_impl(crate_ident: &CrateIdent, parsed: &ParsedNodeFn } }); - let value_field_names: Vec<&Ident> = regular_fields - .iter() - .filter(|field| matches!(field.ty, ParsedFieldType::Regular(_))) - .map(|field| &field.pat_ident.ident) - .collect(); - + // The extent override is the leveled `extent_at`; consumers query the + // composite `extent(ctx, Level)`, which the trait derives from it. A node + // without `extent = fn` keeps the scalar default (one item at every level). let extent_impl = match &parsed.attributes.extent { Some(path) => quote! { - fn extent(&self, __input: &#ctx_ident) -> #core_types::gpoll::GPoll<#core_types::gpoll::Extent> { - #path(self, __input) + fn extent_at(&self, __input: &#ctx_ident, __level: u8) -> #core_types::gpoll::GPoll<#core_types::gpoll::Extent> { + #path(self, __input, __level) } }, - None if value_field_names.is_empty() => quote!(), - None => { - let first = value_field_names[0]; - let mut meet = quote!(self.#first.extent(__input)); - for name in &value_field_names[1..] { - meet = quote!(#core_types::gpoll::Extent::meet(#meet, self.#name.extent(__input))); - } - quote! { - fn extent(&self, __input: &#ctx_ident) -> #core_types::gpoll::GPoll<#core_types::gpoll::Extent> { - #meet - } - } - } + None => quote!(), }; let serialize_impl = match &parsed.attributes.serialize { diff --git a/node-graph/nodes/gcore/src/memo.rs b/node-graph/nodes/gcore/src/memo.rs index ce2ef37b08..a7ac02527a 100644 --- a/node-graph/nodes/gcore/src/memo.rs +++ b/node-graph/nodes/gcore/src/memo.rs @@ -45,11 +45,11 @@ fn memoize<'e>( result } -fn memoize_extent(node: &MemoizeNode, ctx: &C) -> GPoll +fn memoize_extent(node: &MemoizeNode, ctx: &C, level: u8) -> GPoll where NodeContent: Node, { - node.content.extent(ctx) + node.content.extent_at(ctx, level) } #[node_macro::node(category(""), path(graphene_core::memo), extent(frame_memo_extent))] @@ -95,11 +95,11 @@ fn frame_memo<'e>( } } -fn frame_memo_extent(node: &FrameMemoNode, ctx: &C) -> GPoll +fn frame_memo_extent(node: &FrameMemoNode, ctx: &C, level: u8) -> GPoll where NodeContent: Node, { - node.content.extent(ctx) + node.content.extent_at(ctx, level) } type MonitorValue = Arc>>>;