use ezu_style as spec;
use crate::graph::{BuildError, Graph, GraphBuilder};
use crate::port::PortKind;
use crate::registry::{FactoryCtx, FactoryError, NodeRegistry};
#[derive(Debug, thiserror::Error)]
pub enum BuildGraphError {
#[error("unknown op `{op}` on node `{node}`")]
UnknownOp { node: String, op: String },
#[error("factory error on node `{node}`: {source}")]
Factory {
node: String,
#[source]
source: FactoryError,
},
#[error(transparent)]
Expand(#[from] spec::ExpandError),
#[error(
"call `{call}` of `{func}`: input `{input}` expects {expected}, but `@{src}` produces {got}"
)]
FuncInputKind {
call: String,
func: String,
input: String,
expected: PortKind,
src: String,
got: PortKind,
},
#[error("call `{call}` of `{func}`: declared output-kind is {declared}, but the body produces {got}")]
FuncOutputKind {
call: String,
func: String,
declared: PortKind,
got: PortKind,
},
#[error(transparent)]
Graph(#[from] BuildError),
}
fn port_kind(k: spec::FuncKind) -> PortKind {
match k {
spec::FuncKind::Features => PortKind::Features,
spec::FuncKind::Raster => PortKind::Raster,
spec::FuncKind::Sprite => PortKind::Sprite,
spec::FuncKind::Brush => PortKind::Brush,
spec::FuncKind::Scalar => PortKind::Scalar,
spec::FuncKind::ScalarField => PortKind::ScalarField,
}
}
pub fn build_graph(
doc: &spec::Document,
registry: &NodeRegistry,
) -> Result<Graph, BuildGraphError> {
let expanded = spec::expand_functions(doc)?;
let (doc, kind_checks) = match &expanded {
Some(e) => (&e.doc, e.kind_checks.as_slice()),
None => (doc, &[][..]),
};
let ctx = FactoryCtx {
params: &doc.params,
sources: &doc.sources,
};
let mut gb = GraphBuilder::new();
let mut pending: Vec<(String, Vec<crate::registry::Connection>)> = Vec::new();
for (id, spec) in &doc.nodes {
let factory = registry
.get(&spec.op)
.ok_or_else(|| BuildGraphError::UnknownOp {
node: id.clone(),
op: spec.op.clone(),
})?;
let built = factory
.build(&spec.fields, &ctx)
.map_err(|e| BuildGraphError::Factory {
node: id.clone(),
source: e,
})?;
gb.add_node(id.clone(), built.node);
pending.push((id.clone(), built.connections));
}
for (dst, conns) in pending {
for c in conns {
gb.connect(c.src, dst.clone(), c.port);
}
}
gb.set_output(doc.output.as_str().to_string());
let graph = gb.build()?;
for check in kind_checks {
let Some(ix) = graph.index_of(&check.node) else {
continue;
};
let got = graph.output_kind(ix);
let expected = port_kind(check.declared);
if got != expected {
return Err(match &check.input {
Some(input) => BuildGraphError::FuncInputKind {
call: check.call.clone(),
func: check.func.clone(),
input: input.clone(),
expected,
src: check.node.clone(),
got,
},
None => BuildGraphError::FuncOutputKind {
call: check.call.clone(),
func: check.func.clone(),
declared: expected,
got,
},
});
}
}
Ok(graph)
}