mirror of
https://github.com/GraphiteEditor/Graphite.git
synced 2026-09-15 22:28:10 +08:00
Emit the async source tier from the node macro
This commit is contained in:
@@ -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)])
|
||||
|
||||
@@ -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, ®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),
|
||||
|
||||
@@ -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)]
|
||||
|
||||
@@ -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`
|
||||
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user