diff --git a/editor/src/messages/portfolio/document/utility_types/network_interface.rs b/editor/src/messages/portfolio/document/utility_types/network_interface.rs index 283925c587..35ca465a11 100644 --- a/editor/src/messages/portfolio/document/utility_types/network_interface.rs +++ b/editor/src/messages/portfolio/document/utility_types/network_interface.rs @@ -671,7 +671,7 @@ impl NodeNetworkInterface { resolved_type, name, description, - valid_types: self.valid_input_types(input_connector, network_path).iter().map(|ty| ty.to_string()).collect(), + valid_types: self.potential_valid_input_types(input_connector, network_path).iter().map(|ty| ty.to_string()).collect(), connected_to, }) } diff --git a/editor/src/messages/portfolio/document/utility_types/network_interface/resolved_types.rs b/editor/src/messages/portfolio/document/utility_types/network_interface/resolved_types.rs index 695a435c6d..ca58c5e654 100644 --- a/editor/src/messages/portfolio/document/utility_types/network_interface/resolved_types.rs +++ b/editor/src/messages/portfolio/document/utility_types/network_interface/resolved_types.rs @@ -2,7 +2,7 @@ use std::collections::{HashMap, HashSet}; use graph_craft::document::value::TaggedValue; use graph_craft::document::{DocumentNodeImplementation, InlineRust, NodeInput}; -use graph_craft::{Type, concrete}; +use graph_craft::{ProtoNodeIdentifier, Type, concrete}; use graphene_std::uuid::NodeId; use interpreted_executor::dynamic_executor::{NodeTypes, ResolvedDocumentNodeTypesDelta}; use interpreted_executor::node_registry::NODE_REGISTRY; @@ -43,6 +43,32 @@ pub enum TypeSource { } impl TypeSource { + /// The reduced set of frontend types for displaying color. + pub fn displayed_type(&self) -> FrontendGraphDataType { + match self.compiled_nested_type() { + Some(nested_type) => match TaggedValue::from_type_or_none(nested_type) { + TaggedValue::U32(_) + | TaggedValue::U64(_) + | TaggedValue::F32(_) + | TaggedValue::F64(_) + | TaggedValue::DVec2(_) + | TaggedValue::F64Array4(_) + | TaggedValue::VecF64(_) + | TaggedValue::VecDVec2(_) + | TaggedValue::DAffine2(_) => FrontendGraphDataType::Number, + TaggedValue::Artboard(_) => FrontendGraphDataType::Artboard, + TaggedValue::Graphic(_) => FrontendGraphDataType::Graphic, + TaggedValue::Raster(_) => FrontendGraphDataType::Raster, + TaggedValue::Vector(_) => FrontendGraphDataType::Vector, + TaggedValue::Color(_) => FrontendGraphDataType::Color, + TaggedValue::Gradient(_) | TaggedValue::GradientStops(_) | TaggedValue::GradientTable(_) => FrontendGraphDataType::Gradient, + TaggedValue::String(_) => FrontendGraphDataType::Typography, + _ => FrontendGraphDataType::General, + }, + None => FrontendGraphDataType::General, + } + } + pub fn compiled_nested_type(&self) -> Option<&Type> { match self { TypeSource::Compiled(compiled_type) => Some(compiled_type.nested_type()), @@ -71,14 +97,6 @@ impl TypeSource { TypeSource::Error(_) => "Error".to_string(), } } - - /// The reduced set of frontend types for displaying color. - pub fn displayed_type(&self) -> FrontendGraphDataType { - match self.compiled_nested_type() { - Some(nested_type) => FrontendGraphDataType::from_type(nested_type), - None => FrontendGraphDataType::General, - } - } } impl NodeNetworkInterface { @@ -126,7 +144,7 @@ impl NodeNetworkInterface { TypeSource::TaggedValue(value) => value, TypeSource::DocumentNodeDefinition(definition) => definition, TypeSource::Unknown => { - let mut valid_types = self.valid_input_types(input_connector, network_path); + let mut valid_types = self.potential_valid_input_types(input_connector, network_path); match valid_types.pop() { Some(valid_type) => valid_type, @@ -147,7 +165,8 @@ impl NodeNetworkInterface { TaggedValue::from_type_or_none(&guaranteed_type) } - pub fn valid_input_types(&mut self, input_connector: &InputConnector, network_path: &[NodeId]) -> Vec { + /// A list of all valid input types for this specific node. + pub fn potential_valid_input_types(&mut self, input_connector: &InputConnector, network_path: &[NodeId]) -> Vec { let InputConnector::Node { node_id, input_index } = input_connector else { // An export can have any type connected to it return vec![graph_craft::generic!(T)]; @@ -171,7 +190,7 @@ impl NodeNetworkInterface { let intersection: HashSet = inputs_from_import .clone() .iter() - .map(|input_connector| self.valid_input_types(input_connector, &nested_path).into_iter().collect::>()) + .map(|input_connector| self.potential_valid_input_types(input_connector, &nested_path).into_iter().collect::>()) .fold(None, |acc: Option>, set| match acc { Some(acc_set) => Some(acc_set.intersection(&set).cloned().collect()), None => Some(set), @@ -191,10 +210,9 @@ impl NodeNetworkInterface { .filter_map(|(node_io, _)| { let valid_implementation = (0..number_of_inputs).filter(|iterator_index| iterator_index != input_index).all(|iterator_index| { let input_type = self.input_type(&InputConnector::node(*node_id, iterator_index), network_path); - // Value inputs are stored as concrete, so they are compared to the nested type. Node inputs are stored as fn, so they are compared to the entire type. - // For example a node input of (Footprint) -> Vector would not be compatible with () -> Vector - node_io.inputs.get(iterator_index).map(|ty| ty.nested_type().clone()).as_ref() == input_type.compiled_nested_type() - || node_io.inputs.get(iterator_index) == input_type.compiled_nested_type() + // TODO: Fix type checking for different call arguments + // For example a node input of (Footprint) -> Vector would not be compatible with a node that is called with () and returns Vector + node_io.inputs.get(iterator_index).map(|ty| ty.nested_type()) == input_type.compiled_nested_type() }); if valid_implementation { node_io.inputs.get(*input_index).cloned() } else { None } }) @@ -207,6 +225,66 @@ impl NodeNetworkInterface { } } + /// Performs a downstream traversal to ensure input type will work in the full context of the graph. + pub fn complete_valid_input_types(&mut self, input_connector: &InputConnector, network_path: &[NodeId]) -> Result, String> { + match input_connector { + InputConnector::Node { node_id, input_index } => { + let Some(implementation) = self.implementation(node_id, network_path) else { + return Err(format!("Could not get node implementation for {:?} {} in valid_input_types", network_path, *node_id)); + }; + match implementation { + DocumentNodeImplementation::Network(_) => self.valid_output_types(&OutputConnector::Import(input_connector.input_index()), &[network_path, &[*node_id]].concat()), + DocumentNodeImplementation::ProtoNode(proto_node_identifier) => { + let Some(implementations) = NODE_REGISTRY.get(proto_node_identifier) else { + return Err(format!("Protonode {proto_node_identifier:?} not found in registry")); + }; + let valid_output_types = match self.valid_output_types(&OutputConnector::node(*node_id, 0), network_path) { + Ok(valid_types) => valid_types, + Err(e) => return Err(e), + }; + + let valid_types = implementations + .iter() + .filter_map(|(node_io, _)| { + if !valid_output_types.iter().any(|output_type| output_type.nested_type() == node_io.return_value.nested_type()) { + return None; + } + + let valid_inputs = (0..node_io.inputs.len()).filter(|iterator_index| iterator_index != input_index).all(|iterator_index| { + let input_type = self.input_type(&InputConnector::node(*node_id, iterator_index), network_path); + match input_type.compiled_nested_type() { + Some(input_type) => node_io.inputs.get(iterator_index).is_some_and(|node_io_input_type| node_io_input_type.nested_type() == input_type), + None => true, + } + }); + if valid_inputs { node_io.inputs.get(*input_index).cloned() } else { None } + }) + .collect::>(); + Ok(valid_types) + } + DocumentNodeImplementation::Extract => { + log::error!("Input types for extract node not supported"); + Ok(Vec::new()) + } + } + } + InputConnector::Export(export_index) => { + match network_path.split_last() { + Some((encapsulating_node, encapsulating_path)) => self.valid_output_types(&OutputConnector::node(*encapsulating_node, *export_index), encapsulating_path), + None => { + // Valid types for the export are all types that can be fed into the render node + // TODO: Use ::IDENTIFIER + let render_node = "graphene_std::wasm_application_io::RenderNode"; + let Some(implementations) = NODE_REGISTRY.get(&ProtoNodeIdentifier::new(render_node)) else { + return Err(format!("Protonode {render_node:?} not found in registry")); + }; + Ok(implementations.iter().map(|(types, _)| types.inputs[1].clone()).collect()) + } + } + } + } + } + pub fn output_type(&mut self, output_connector: &OutputConnector, network_path: &[NodeId]) -> TypeSource { match output_connector { OutputConnector::Node { node_id, output_index } => { @@ -243,7 +321,7 @@ impl NodeNetworkInterface { let intersection = inputs_from_import .clone() .iter() - .map(|input_connector| self.valid_input_types(input_connector, &network_path).into_iter().collect::>()) + .map(|input_connector| self.potential_valid_input_types(input_connector, &network_path).into_iter().collect::>()) .fold(None, |acc: Option>, set| match acc { Some(acc_set) => Some(acc_set.intersection(&set).cloned().collect()), None => Some(set), @@ -281,7 +359,7 @@ impl NodeNetworkInterface { log::error!("Protonode {proto_node_identifier:?} not found in registry"); return None; }; - implementations.keys().next().and_then(|node_io| node_io.inputs.get(input_connector.input_index())).cloned() + implementations.keys().min().and_then(|node_io| node_io.inputs.get(input_connector.input_index())).cloned() } DocumentNodeImplementation::Extract => None, } diff --git a/node-graph/gcore/src/types.rs b/node-graph/gcore/src/types.rs index b95b869496..b819f15839 100644 --- a/node-graph/gcore/src/types.rs +++ b/node-graph/gcore/src/types.rs @@ -77,7 +77,7 @@ macro_rules! fn_type_fut { }; } -#[derive(Clone, PartialEq, Eq, Hash, Default, serde::Serialize, serde::Deserialize)] +#[derive(Clone, PartialEq, Eq, Hash, Default, PartialOrd, Ord, serde::Serialize, serde::Deserialize)] pub struct NodeIOTypes { pub call_argument: Type, pub return_value: Type, @@ -229,8 +229,20 @@ impl PartialEq for TypeDescriptor { } } +impl Ord for TypeDescriptor { + fn cmp(&self, other: &Self) -> std::cmp::Ordering { + self.name.cmp(&other.name) + } +} + +impl PartialOrd for TypeDescriptor { + fn partial_cmp(&self, other: &Self) -> Option { + Some(self.cmp(other)) + } +} + /// Graph runtime type information used for type inference. -#[derive(Clone, PartialEq, Eq, Hash, specta::Type, serde::Serialize, serde::Deserialize)] +#[derive(Clone, PartialEq, Eq, Hash, PartialOrd, Ord, specta::Type, serde::Serialize, serde::Deserialize)] pub enum Type { /// A wrapper for some type variable used within the inference system. Resolved at inference time and replaced with a concrete type. Generic(Cow<'static, str>),