From cf66f8396b27cc24e89b736b21658b650daedff5 Mon Sep 17 00:00:00 2001 From: Dennis Kobert Date: Sun, 26 Jul 2026 22:58:40 +0000 Subject: [PATCH] Emit the async source tier from the node macro --- node-graph/node-macro/src/codegen.rs | 13 +- node-graph/node-macro/src/gcodegen.rs | 134 +++++++++++++++++- node-graph/node-macro/src/parsing.rs | 41 ++++++ .../src/shader_nodes/per_pixel_adjust.rs | 7 +- node-graph/node-macro/src/validation.rs | 35 +++++ 5 files changed, 219 insertions(+), 11 deletions(-) diff --git a/node-graph/node-macro/src/codegen.rs b/node-graph/node-macro/src/codegen.rs index da3bc8cbe5..d48a4bf6e2 100644 --- a/node-graph/node-macro/src/codegen.rs +++ b/node-graph/node-macro/src/codegen.rs @@ -15,6 +15,8 @@ pub(crate) fn generate_node_code(crate_ident: &CrateIdent, parsed: &ParsedNodeFn mod_name, fn_generics, input, + output_type, + is_async, fields, description, .. @@ -109,7 +111,11 @@ pub(crate) fn generate_node_code(crate_ident: &CrateIdent, parsed: &ParsedNodeFn quote! { pub(super) #name: #r#gen } }); - let struct_fields = data_field_defs.chain(regular_field_defs); + let slot_value_type = crate::gcodegen::slot_value_type(output_type); + let slot_field = is_async + .then(|| quote! { pub(super) slot: std::sync::Arc>>>> }) + .into_iter(); + let struct_fields = data_field_defs.chain(regular_field_defs).chain(slot_field); // Only regular fields have UI metadata (data fields are internal state) let widget_override: Vec<_> = regular_fields @@ -217,10 +223,11 @@ pub(crate) fn generate_node_code(crate_ident: &CrateIdent, parsed: &ParsedNodeFn let regular_inits = regular_field_names.iter().map(|name| { quote! { #name } }); - let all_field_inits = data_inits.chain(regular_inits); + let slot_init = is_async.then(|| quote! { slot: Default::default() }).into_iter(); + let all_field_inits = data_inits.chain(regular_inits).chain(slot_init); // Data fields may not implement Copy, PartialEq, etc., so only derive Debug and Clone - let struct_derives = if data_fields.is_empty() { + let struct_derives = if data_fields.is_empty() && !is_async { quote!(#[derive(Debug, Copy, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]) } else { quote!(#[derive(Debug, Clone)]) diff --git a/node-graph/node-macro/src/gcodegen.rs b/node-graph/node-macro/src/gcodegen.rs index 1a953c4868..0c20febdd0 100644 --- a/node-graph/node-macro/src/gcodegen.rs +++ b/node-graph/node-macro/src/gcodegen.rs @@ -18,7 +18,16 @@ pub(crate) fn generate_gnode_code(crate_ident: &CrateIdent, parsed: &ParsedNodeF Some(ctx_param) => ctx_param.ident.clone(), None => format_ident!("__Ctx"), }; - let ctx_bounds: Vec = match ctx_param { + let async_source = parsed.is_async; + if async_source && parsed.fields.iter().any(|field| matches!(field.ty, ParsedFieldType::Node(_))) { + return Ok(GNodeTokens { + in_mod: quote!(), + top_level: quote!(), + }); + } + let snapshot_ctx = async_source && 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 { Some(ctx_param) => ctx_param .bounds .iter() @@ -29,6 +38,20 @@ pub(crate) fn generate_gnode_code(crate_ident: &CrateIdent, parsed: &ParsedNodeF .collect(), None => vec![quote!(#core_types::Ctx)], }; + if async_source { + ctx_bounds.push(quote!(#core_types::CacheHash)); + } + if snapshot_ctx { + ctx_bounds.extend([ + quote!(#core_types::context::DeriveCtx), + quote!(#core_types::context::ExtractFootprint), + quote!(#core_types::context::ExtractRealTime), + quote!(#core_types::context::ExtractAnimationTime), + quote!(#core_types::context::ExtractPointerPosition), + quote!(#core_types::context::ExtractIndex), + quote!(#core_types::context::ExtractPosition), + ]); + } let derives = ctx_param.is_some_and(|ctx_param| { ctx_param.bounds.iter().any(|bound| match bound { @@ -121,6 +144,22 @@ pub(crate) fn generate_gnode_code(crate_ident: &CrateIdent, parsed: &ParsedNodeF } }); + let async_bounds = match async_source { + false => Vec::new(), + true => { + let output_clone = std::iter::once(quote!(#trait_output: Clone)); + let value_clones = regular_fields.iter().filter_map(|field| match &field.ty { + ParsedFieldType::Regular(RegularParsedField { ty, .. }) => Some(quote!(#ty: Clone)), + _ => None, + }); + let data_clones = data_fields.iter().filter_map(|field| match &field.ty { + ParsedFieldType::Regular(RegularParsedField { ty, .. }) => Some(quote!(#ty: Clone)), + _ => None, + }); + output_clone.chain(value_clones).chain(data_clones).collect() + } + }; + let clampable_bounds = regular_fields.iter().filter_map(|field| { let ParsedFieldType::Regular(RegularParsedField { ty, number_hard_min, number_hard_max, .. }) = &field.ty else { return None; @@ -215,9 +254,39 @@ pub(crate) fn generate_gnode_code(crate_ident: &CrateIdent, parsed: &ParsedNodeF let fn_where = &parsed.where_clause; let body = &parsed.body; let vis = &parsed.vis; - let kernel = quote! { - #[allow(clippy::too_many_arguments)] - #vis fn #fn_name<#(#generics,)*>(#ctx_pat: &#ctx_ident #(, #data_params)* #(, #kernel_params)*) -> #output_type #fn_where #body + let injected = |field: &&&ParsedField| async_source && (field.pat_ident.ident == "_runtime" || field.pat_ident.ident == "_source"); + let kernel_fields: Vec<&&ParsedField> = regular_fields.iter().filter(|field| !injected(field)).collect(); + let kernel = match async_source { + false => quote! { + #[allow(clippy::too_many_arguments)] + #vis fn #fn_name<#(#generics,)*>(#ctx_pat: &#ctx_ident #(, #data_params)* #(, #kernel_params)*) -> #output_type #fn_where #body + }, + true => { + let kernel_generics = parsed.fn_generics.iter().filter(|param| match param { + GenericParam::Type(type_param) => Some(&type_param.ident) != ctx_param.map(|ctx_param| &ctx_param.ident), + _ => true, + }); + let snapshot_param = snapshot_ctx.then(|| quote!(#ctx_pat: #core_types::context::CtxSnapshot)).into_iter(); + let data_kernel_params = data_fields.iter().map(|field| { + let pat = &field.pat_ident; + let ParsedFieldType::Regular(RegularParsedField { ty, .. }) = &field.ty else { + unreachable!("data fields are regular types"); + }; + quote!(#pat: #ty) + }); + let value_kernel_params = kernel_fields.iter().map(|field| { + let pat = &field.pat_ident; + let ParsedFieldType::Regular(RegularParsedField { ty, .. }) = &field.ty else { + unreachable!("async source fields are eager values"); + }; + quote!(#pat: #ty) + }); + let params = snapshot_param.chain(data_kernel_params).chain(value_kernel_params); + quote! { + #[allow(clippy::too_many_arguments)] + #vis async fn #fn_name<#(#kernel_generics,)*>(#(#params),*) -> #output_type #fn_where #body + } + } }; let cell_constructor = match parsed.attributes.no_partial { true => quote!(#core_types::gnode::StatusCell::no_partial()), @@ -235,6 +304,53 @@ pub(crate) fn generate_gnode_code(crate_ident: &CrateIdent, parsed: &ParsedNodeF KernelKind::Plain => quote!(__cell.finish(#kernel_call)), }; + let eval_tail = match async_source { + false => lift, + true => { + let kernel_value_names: Vec<&Ident> = kernel_fields.iter().map(|field| &field.pat_ident.ident).collect(); + let inflight = match &parsed.attributes.placeholder { + Some(path) => quote!(__cell.merge(#core_types::gpoll::GPoll::Partial(#path(#(&#kernel_value_names),*)))), + None => quote!(#core_types::gpoll::GPoll::Pending), + }; + let snapshot_binding = snapshot_ctx.then(|| quote!(let __snapshot = #core_types::context::CtxSnapshot::capture(__input);)).into_iter(); + let snapshot_arg = snapshot_ctx.then(|| quote!(__snapshot)).into_iter(); + let future_args = snapshot_arg + .chain(data_names.iter().map(|name| quote!(self.#name.clone()))) + .chain(kernel_value_names.iter().map(|name| quote!(#name.clone()))); + let completion = match kernel_kind(&parsed.output_type) { + KernelKind::Plain => quote!(#core_types::gpoll::GPoll::Final(__future.await)), + KernelKind::Poll(_) => quote!(__future.await), + KernelKind::Interrupt(_) => quote! { + match __future.await { + Ok(value) => #core_types::gpoll::GPoll::Final(value), + Err(interrupt) => interrupt.into(), + } + }, + }; + quote! { + let __key = #core_types::wire::cache_key(__input); + { + let __entries = self.slot.lock().unwrap(); + if let Some(__state) = __entries.get(&__key) { + return match __state { + Some(value) => __cell.merge(value.clone()), + None => #inflight, + }; + } + } + self.slot.lock().unwrap().insert(__key, None); + let __slot = std::sync::Arc::clone(&self.slot); + #(#snapshot_binding)* + let __future = self::#fn_name(#(#future_args),*); + _runtime.0.spawn(_source, Box::pin(async move { + let __value = #completion; + __slot.lock().unwrap().insert(__key, Some(__value)); + })); + #inflight + } + } + }; + let wire = entries_tokens(parsed, &struct_name, &data_field_generic_idents, ®ular_fields); let cfg = crate::shader_nodes::modify_cfg(&parsed.attributes); let wire_reexport = match wire.is_empty() { @@ -258,6 +374,7 @@ pub(crate) fn generate_gnode_code(crate_ident: &CrateIdent, parsed: &ParsedNodeF where #(#node_bounds,)* #(#clampable_bounds,)* + #(#async_bounds,)* #(#where_predicates,)* { type Output = #trait_output; @@ -266,7 +383,7 @@ pub(crate) fn generate_gnode_code(crate_ident: &CrateIdent, parsed: &ParsedNodeF let __cell = #cell_constructor; #(#eval_values)* #(#clamps)* - #lift + #eval_tail } #extent_impl @@ -285,6 +402,13 @@ pub(crate) fn generate_gnode_code(crate_ident: &CrateIdent, parsed: &ParsedNodeF }) } +pub(crate) fn slot_value_type(output: &Type) -> Type { + match kernel_kind(output) { + KernelKind::Plain => output.clone(), + KernelKind::Poll(inner) | KernelKind::Interrupt(inner) => inner, + } +} + enum KernelKind { Plain, Interrupt(Type), diff --git a/node-graph/node-macro/src/parsing.rs b/node-graph/node-macro/src/parsing.rs index 61b8e95301..213b84e861 100644 --- a/node-graph/node-macro/src/parsing.rs +++ b/node-graph/node-macro/src/parsing.rs @@ -997,6 +997,10 @@ pub fn new_node_fn(attr: TokenStream2, item: TokenStream2) -> syn::Result { description: "".to_string(), widget_override: Default::default(), ty: ParsedFieldType::Regular(RegularParsedField { - ty: parse_quote!(&'a WgpuExecutor), + ty: parse_quote!(std::sync::Arc), exposed: true, - value_source: ParsedValueSource::Scope(parse_quote!("graphene_std::platform_application_io::WgpuExecutorNode")), + value_source: ParsedValueSource::Scope(parse_quote!("graphene_std::platform_application_io::WgpuExecutorArcNode")), number_soft_min: None, number_soft_max: None, number_hard_min: None, @@ -305,7 +305,7 @@ impl PerPixelAdjustCodegen<'_> { fn_name: self.shader_node_mod.clone(), struct_name: format_ident!("{}", self.shader_node_mod.to_string().to_case(Case::Pascal)), mod_name: self.shader_node_mod.clone(), - fn_generics: vec![parse_quote!('a: 'n)], + fn_generics: Vec::new(), where_clause: None, input: Input { pat_ident: self.parsed.input.pat_ident.clone(), @@ -320,6 +320,7 @@ impl PerPixelAdjustCodegen<'_> { description: self.parsed.description.clone(), }; parsed_node_fn.replace_impl_trait_in_input(); + parsed_node_fn.inject_async_source_fields(self.crate_ident.gcore()?); let gpu_node_impl = crate::codegen::generate_node_code(self.crate_ident, &parsed_node_fn)?; // wrap node in `mod #gpu_node_mod` diff --git a/node-graph/node-macro/src/validation.rs b/node-graph/node-macro/src/validation.rs index 7b51b31efa..3b2f27534e 100644 --- a/node-graph/node-macro/src/validation.rs +++ b/node-graph/node-macro/src/validation.rs @@ -11,6 +11,7 @@ pub fn validate_node_fn(parsed: &ParsedNodeFn) -> syn::Result<()> { validate_primary_input_expose, validate_min_max, validate_range_slider_bounds, + validate_async_source, ]; for validator in validators { @@ -20,6 +21,40 @@ pub fn validate_node_fn(parsed: &ParsedNodeFn) -> syn::Result<()> { Ok(()) } +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")); + if !parsed.is_async { + if snapshot_ctx { + emit_error!(parsed.input.pat_ident.span(), "`CtxSnapshot` is the async source context; synchronous nodes take `impl Ctx` and read through extract bounds"); + } + return; + } + for field in &parsed.fields { + if matches!(field.ty, ParsedFieldType::Node(_)) { + emit_error!( + field.pat_ident.span(), + "async source nodes cannot take `impl Node` inputs: the spawned future outlives any borrow of the graph, so it cannot evaluate other nodes; declare the input as an eager value instead" + ); + } + } + let ctx_ident = match &parsed.input.ty { + Type::Path(path) => path.path.get_ident(), + _ => None, + }; + for param in &parsed.fn_generics { + let GenericParam::Type(type_param) = param else { continue }; + if Some(&type_param.ident) == ctx_ident { + continue; + } + if crate::codegen::type_contains_ident(&parsed.output_type, &type_param.ident) { + emit_error!( + parsed.output_type.span(), + "async source nodes do not support generic output types yet; the slot map needs the output type stated per implementation row, which is not wired up until a node requires it" + ); + } + } +} + fn validate_min_max(parsed: &ParsedNodeFn) { for field in &parsed.fields { if let ParsedField {