use crate::{Image, Node, NodeId, PixelsError, Region, Result, Tile, TileBuf};
use std::collections::HashMap;
use std::sync::Arc;
pub fn evaluate(image: &Image) -> Result<TileBuf> {
let mut cache: HashMap<NodeId, Arc<TileBuf>> = HashMap::new();
let buffer = evaluate_node(image.node(), &mut cache)?;
Ok(Arc::try_unwrap(buffer).unwrap_or_else(|shared| (*shared).clone()))
}
fn evaluate_node(
root: &Arc<Node>,
cache: &mut HashMap<NodeId, Arc<TileBuf>>,
) -> Result<Arc<TileBuf>> {
enum Step {
Visit(Arc<Node>),
Compute(Arc<Node>),
}
let mut stack = vec![Step::Visit(Arc::clone(root))];
while let Some(step) = stack.pop() {
match step {
Step::Visit(node) => {
if cache.contains_key(&node.id()) {
continue;
}
stack.push(Step::Compute(Arc::clone(&node)));
for input in node.inputs() {
stack.push(Step::Visit(Arc::clone(input)));
}
}
Step::Compute(node) => {
if cache.contains_key(&node.id()) {
continue;
}
let buffer = compute_node(&node, cache)?;
cache.insert(node.id(), Arc::new(buffer));
}
}
}
cache
.get(&root.id())
.map(Arc::clone)
.ok_or_else(|| PixelsError::graph("graph evaluation produced no result for the root node"))
}
fn compute_node(node: &Arc<Node>, cache: &HashMap<NodeId, Arc<TileBuf>>) -> Result<TileBuf> {
let descriptor = node.descriptor();
let output_region = descriptor.region();
let mut output = TileBuf::for_image(&descriptor)?;
if let Some(producer) = node.producer() {
let mut tile = output.as_tile_mut()?;
producer.produce(output_region, &mut tile)?;
return Ok(output);
}
let op = node.op().ok_or_else(|| {
PixelsError::graph(format!("node `{}` is neither op nor source", node.name()))
})?;
let input_descriptors: Vec<_> = node.inputs().iter().map(|n| n.descriptor()).collect();
let requested = op.input_regions(output_region, &input_descriptors)?;
if requested.len() != node.inputs().len() {
return Err(PixelsError::graph(format!(
"op `{}` requested {} input region(s) for {} input(s)",
op.name(),
requested.len(),
node.inputs().len()
)));
}
let input_buffers: Vec<Arc<TileBuf>> = node
.inputs()
.iter()
.map(|input| {
cache.get(&input.id()).map(Arc::clone).ok_or_else(|| {
PixelsError::graph(format!("input `{}` was not evaluated first", input.name()))
})
})
.collect::<Result<_>>()?;
let mut input_tiles: Vec<Tile<'_>> = Vec::with_capacity(input_buffers.len());
for (buffer, region) in input_buffers.iter().zip(&requested) {
let tile = buffer.as_tile()?;
if !tile.region().contains(*region) {
return Err(PixelsError::graph(format!(
"op `{}` asked for {region}, outside its input {}",
op.name(),
tile.region()
)));
}
input_tiles.push(tile);
}
{
let mut tile = output.as_tile_mut()?;
op.compute(&input_tiles, &mut tile)?;
}
Ok(output)
}
pub fn evaluate_rows(
image: &Image,
mut consume: impl FnMut(u32, &[u8]) -> Result<()>,
) -> Result<()> {
let buffer = evaluate(image)?;
let tile = buffer.as_tile()?;
let region = tile.region();
for y in region.y..region.y.saturating_add(region.height) {
let row = tile
.row(y)
.ok_or_else(|| PixelsError::graph(format!("evaluated buffer is missing row {y}")))?;
consume(y, row)?;
}
Ok(())
}
pub fn demand(image: &Image, output: Region) -> Result<Vec<(NodeId, Region)>> {
let mut out = Vec::new();
let mut stack = vec![(Arc::clone(image.node()), output)];
while let Some((node, region)) = stack.pop() {
out.push((node.id(), region));
let Some(op) = node.op() else { continue };
let descriptors: Vec<_> = node.inputs().iter().map(|n| n.descriptor()).collect();
let requested = op.input_regions(region, &descriptors)?;
for (input, region) in node.inputs().iter().zip(requested) {
stack.push((Arc::clone(input), region));
}
}
Ok(out)
}
#[cfg(test)]
#[allow(
clippy::unwrap_used,
clippy::indexing_slicing,
reason = "tests operate on known-good values and assert shapes directly"
)]
mod tests {
use super::*;
use crate::testing::{ConstantOp, CountingProducer, RampProducer};
use crate::{Format, ImageDescriptor, PixelFormat, Producer};
fn ramp(width: u32, height: u32) -> Image {
let desc = ImageDescriptor::new(width, height, PixelFormat::Gray8).unwrap();
Image::from_producer(Arc::new(RampProducer::new(desc)), Format::Raw)
}
#[test]
fn a_bare_source_evaluates_to_its_pixels() {
let buffer = evaluate(&ramp(3, 2)).unwrap();
assert_eq!(buffer.bytes(), &[0, 1, 2, 3, 4, 5]);
}
#[test]
fn ops_run_in_dependency_order() {
let image = ramp(2, 2)
.apply(Arc::new(ConstantOp::new(5)))
.unwrap()
.apply(Arc::new(ConstantOp::new(9)))
.unwrap();
assert_eq!(evaluate(&image).unwrap().bytes(), &[9, 9, 9, 9]);
}
#[test]
fn a_shared_subgraph_is_evaluated_once() {
let desc = ImageDescriptor::new(2, 2, PixelFormat::Gray8).unwrap();
let producer = Arc::new(CountingProducer::new(desc));
let base = Image::from_producer(Arc::clone(&producer) as Arc<dyn Producer>, Format::Raw);
let left = base.apply(Arc::new(ConstantOp::new(1))).unwrap();
let right = base.apply(Arc::new(ConstantOp::new(2))).unwrap();
let joined = Image::combine(&[left, right], Arc::new(crate::testing::SumOp)).unwrap();
assert_eq!(evaluate(&joined).unwrap().bytes(), &[3, 3, 3, 3]);
assert_eq!(
producer.produce_calls(),
1,
"shared source pulled exactly once"
);
}
#[test]
fn nothing_is_pulled_until_a_terminal_runs() {
let desc = ImageDescriptor::new(4, 4, PixelFormat::Gray8).unwrap();
let producer = Arc::new(CountingProducer::new(desc));
let image = Image::from_producer(Arc::clone(&producer) as Arc<dyn Producer>, Format::Raw)
.apply(Arc::new(ConstantOp::new(1)))
.unwrap();
assert_eq!(producer.produce_calls(), 0);
let _ = image.metadata().unwrap();
assert_eq!(producer.produce_calls(), 0, "metadata must not pull pixels");
evaluate(&image).unwrap();
assert_eq!(producer.produce_calls(), 1);
}
#[test]
fn deep_chains_do_not_overflow_the_stack() {
let mut image = ramp(2, 2);
for _ in 0..10_000 {
image = image.apply(Arc::new(ConstantOp::new(3))).unwrap();
}
assert_eq!(evaluate(&image).unwrap().bytes(), &[3, 3, 3, 3]);
}
#[test]
fn a_failing_op_fails_the_whole_evaluation() {
let image = ramp(2, 2)
.apply(Arc::new(crate::testing::FailingOp))
.unwrap();
let err = evaluate(&image).unwrap_err();
assert_eq!(err.code(), crate::ErrorCode::Malformed);
}
#[test]
fn evaluate_rows_yields_rows_in_order() {
let mut seen = Vec::new();
evaluate_rows(&ramp(2, 3), |y, row| {
seen.push((y, row.to_vec()));
Ok(())
})
.unwrap();
assert_eq!(
seen,
vec![(0, vec![0, 1]), (1, vec![2, 3]), (2, vec![4, 5])]
);
}
#[test]
fn evaluate_rows_propagates_consumer_errors() {
let err = evaluate_rows(&ramp(2, 3), |_, _| {
Err(PixelsError::unsupported("sink refused"))
})
.unwrap_err();
assert_eq!(err.code(), crate::ErrorCode::Unsupported);
}
#[test]
fn demand_walks_the_inverse_mapping_without_computing() {
let desc = ImageDescriptor::new(4, 4, PixelFormat::Gray8).unwrap();
let producer = Arc::new(CountingProducer::new(desc));
let image = Image::from_producer(Arc::clone(&producer) as Arc<dyn Producer>, Format::Raw)
.apply(Arc::new(ConstantOp::new(1)))
.unwrap();
let pairs = demand(&image, Region::from_size(4, 4)).unwrap();
assert_eq!(pairs.len(), 2, "one op node and one source node");
assert_eq!(
producer.produce_calls(),
0,
"demand propagation touches no pixels"
);
}
}