Emit the async source tier from the node macro

This commit is contained in:
Dennis Kobert
2026-07-26 22:58:40 +00:00
parent be5c78421a
commit cf66f8396b
5 changed files with 219 additions and 11 deletions

View File

@@ -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<std::sync::Mutex<std::collections::HashMap<u64, Option<gcore::gpoll::GPoll<#slot_value_type>>>>> })
.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)])

View File

@@ -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<TokenStream2> = 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<TokenStream2> = 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, &regular_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),

View File

@@ -997,6 +997,10 @@ pub fn new_node_fn(attr: TokenStream2, item: TokenStream2) -> syn::Result<TokenS
let crate_ident = CrateIdent::default();
let mut parsed_node = parse_node_fn(attr, item.clone()).map_err(|e| Error::new(e.span(), format!("Failed to parse node function:\n{e}")))?;
parsed_node.replace_impl_trait_in_input();
if parsed_node.is_async {
let core_types = crate_ident.gcore()?.clone();
parsed_node.inject_async_source_fields(&core_types);
}
crate::validation::validate_node_fn(&parsed_node).map_err(|e| Error::new(e.span(), format!("Validation error:\n{e}")))?;
generate_node_code(&crate_ident, &parsed_node).map_err(|e| Error::new(e.span(), format!("Failed to generate node code:\n{e}")))
}
@@ -1024,6 +1028,43 @@ impl ParsedNodeFn {
self.input.pat_ident.ident = Ident::new("__ctx", self.input.pat_ident.ident.span());
}
}
pub fn inject_async_source_fields(&mut self, core_types: &TokenStream2) {
let hidden_field = |name: &str, ty: Type, value_source: ParsedValueSource| ParsedField {
pat_ident: PatIdent {
attrs: Vec::new(),
by_ref: None,
mutability: None,
ident: Ident::new(name, proc_macro2::Span::call_site()),
subpat: None,
},
name: None,
description: String::new(),
widget_override: ParsedWidgetOverride::Hidden,
ty: ParsedFieldType::Regular(RegularParsedField {
ty,
exposed: false,
value_source,
number_soft_min: None,
number_soft_max: None,
number_hard_min: None,
number_hard_max: None,
number_mode_range: false,
implementations: Default::default(),
gpu_image: false,
}),
number_display_decimal_places: None,
number_step: None,
unit: None,
is_data_field: false,
};
self.fields.push(hidden_field(
"_runtime",
parse_quote!(#core_types::runtime::RuntimeHandle),
ParsedValueSource::Scope(parse_quote!("graphene_std::runtime::RuntimeNode")),
));
self.fields.push(hidden_field("_source", parse_quote!(#core_types::SourceId), ParsedValueSource::None));
}
}
#[cfg(test)]

View File

@@ -231,9 +231,9 @@ impl PerPixelAdjustCodegen<'_> {
description: "".to_string(),
widget_override: Default::default(),
ty: ParsedFieldType::Regular(RegularParsedField {
ty: parse_quote!(&'a WgpuExecutor),
ty: parse_quote!(std::sync::Arc<WgpuExecutor>),
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`

View File

@@ -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 {