1use crate::{Image, Node, NodeId, PixelsError, Region, Result, Tile, TileBuf};
19use std::collections::HashMap;
20use std::sync::Arc;
21
22pub fn evaluate(image: &Image) -> Result<TileBuf> {
30 let mut cache: HashMap<NodeId, Arc<TileBuf>> = HashMap::new();
31 let buffer = evaluate_node(image.node(), &mut cache)?;
32 Ok(Arc::try_unwrap(buffer).unwrap_or_else(|shared| (*shared).clone()))
35}
36
37fn evaluate_node(
44 root: &Arc<Node>,
45 cache: &mut HashMap<NodeId, Arc<TileBuf>>,
46) -> Result<Arc<TileBuf>> {
47 enum Step {
49 Visit(Arc<Node>),
51 Compute(Arc<Node>),
53 }
54
55 let mut stack = vec![Step::Visit(Arc::clone(root))];
56 while let Some(step) = stack.pop() {
57 match step {
58 Step::Visit(node) => {
59 if cache.contains_key(&node.id()) {
60 continue;
61 }
62 stack.push(Step::Compute(Arc::clone(&node)));
64 for input in node.inputs() {
65 stack.push(Step::Visit(Arc::clone(input)));
66 }
67 }
68 Step::Compute(node) => {
69 if cache.contains_key(&node.id()) {
70 continue;
71 }
72 let buffer = compute_node(&node, cache)?;
73 cache.insert(node.id(), Arc::new(buffer));
74 }
75 }
76 }
77
78 cache
79 .get(&root.id())
80 .map(Arc::clone)
81 .ok_or_else(|| PixelsError::graph("graph evaluation produced no result for the root node"))
82}
83
84fn compute_node(node: &Arc<Node>, cache: &HashMap<NodeId, Arc<TileBuf>>) -> Result<TileBuf> {
86 let descriptor = node.descriptor();
87 let output_region = descriptor.region();
88 let mut output = TileBuf::for_image(&descriptor)?;
89
90 if let Some(producer) = node.producer() {
91 let mut tile = output.as_tile_mut()?;
92 producer.produce(output_region, &mut tile)?;
93 return Ok(output);
94 }
95
96 let op = node.op().ok_or_else(|| {
97 PixelsError::graph(format!("node `{}` is neither op nor source", node.name()))
98 })?;
99
100 let input_descriptors: Vec<_> = node.inputs().iter().map(|n| n.descriptor()).collect();
103 let requested = op.input_regions(output_region, &input_descriptors)?;
104 if requested.len() != node.inputs().len() {
105 return Err(PixelsError::graph(format!(
106 "op `{}` requested {} input region(s) for {} input(s)",
107 op.name(),
108 requested.len(),
109 node.inputs().len()
110 )));
111 }
112
113 let input_buffers: Vec<Arc<TileBuf>> = node
115 .inputs()
116 .iter()
117 .map(|input| {
118 cache.get(&input.id()).map(Arc::clone).ok_or_else(|| {
119 PixelsError::graph(format!("input `{}` was not evaluated first", input.name()))
120 })
121 })
122 .collect::<Result<_>>()?;
123
124 let mut input_tiles: Vec<Tile<'_>> = Vec::with_capacity(input_buffers.len());
125 for (buffer, region) in input_buffers.iter().zip(&requested) {
126 let tile = buffer.as_tile()?;
127 if !tile.region().contains(*region) {
128 return Err(PixelsError::graph(format!(
129 "op `{}` asked for {region}, outside its input {}",
130 op.name(),
131 tile.region()
132 )));
133 }
134 input_tiles.push(tile);
135 }
136
137 {
138 let mut tile = output.as_tile_mut()?;
139 op.compute(&input_tiles, &mut tile)?;
140 }
141 Ok(output)
142}
143
144pub fn evaluate_rows(
155 image: &Image,
156 mut consume: impl FnMut(u32, &[u8]) -> Result<()>,
157) -> Result<()> {
158 let buffer = evaluate(image)?;
159 let tile = buffer.as_tile()?;
160 let region = tile.region();
161 for y in region.y..region.y.saturating_add(region.height) {
162 let row = tile
163 .row(y)
164 .ok_or_else(|| PixelsError::graph(format!("evaluated buffer is missing row {y}")))?;
165 consume(y, row)?;
166 }
167 Ok(())
168}
169
170pub fn demand(image: &Image, output: Region) -> Result<Vec<(NodeId, Region)>> {
182 let mut out = Vec::new();
183 let mut stack = vec![(Arc::clone(image.node()), output)];
184 while let Some((node, region)) = stack.pop() {
185 out.push((node.id(), region));
186 let Some(op) = node.op() else { continue };
187 let descriptors: Vec<_> = node.inputs().iter().map(|n| n.descriptor()).collect();
188 let requested = op.input_regions(region, &descriptors)?;
189 for (input, region) in node.inputs().iter().zip(requested) {
190 stack.push((Arc::clone(input), region));
191 }
192 }
193 Ok(out)
194}
195
196#[cfg(test)]
197#[allow(
198 clippy::unwrap_used,
199 clippy::indexing_slicing,
200 reason = "tests operate on known-good values and assert shapes directly"
201)]
202mod tests {
203 use super::*;
204 use crate::testing::{ConstantOp, CountingProducer, RampProducer};
205 use crate::{Format, ImageDescriptor, PixelFormat, Producer};
206
207 fn ramp(width: u32, height: u32) -> Image {
208 let desc = ImageDescriptor::new(width, height, PixelFormat::Gray8).unwrap();
209 Image::from_producer(Arc::new(RampProducer::new(desc)), Format::Raw)
210 }
211
212 #[test]
213 fn a_bare_source_evaluates_to_its_pixels() {
214 let buffer = evaluate(&ramp(3, 2)).unwrap();
215 assert_eq!(buffer.bytes(), &[0, 1, 2, 3, 4, 5]);
216 }
217
218 #[test]
219 fn ops_run_in_dependency_order() {
220 let image = ramp(2, 2)
221 .apply(Arc::new(ConstantOp::new(5)))
222 .unwrap()
223 .apply(Arc::new(ConstantOp::new(9)))
224 .unwrap();
225 assert_eq!(evaluate(&image).unwrap().bytes(), &[9, 9, 9, 9]);
226 }
227
228 #[test]
229 fn a_shared_subgraph_is_evaluated_once() {
230 let desc = ImageDescriptor::new(2, 2, PixelFormat::Gray8).unwrap();
231 let producer = Arc::new(CountingProducer::new(desc));
232 let base = Image::from_producer(Arc::clone(&producer) as Arc<dyn Producer>, Format::Raw);
233 let left = base.apply(Arc::new(ConstantOp::new(1))).unwrap();
235 let right = base.apply(Arc::new(ConstantOp::new(2))).unwrap();
236 let joined = Image::combine(&[left, right], Arc::new(crate::testing::SumOp)).unwrap();
237 assert_eq!(evaluate(&joined).unwrap().bytes(), &[3, 3, 3, 3]);
238 assert_eq!(
239 producer.produce_calls(),
240 1,
241 "shared source pulled exactly once"
242 );
243 }
244
245 #[test]
246 fn nothing_is_pulled_until_a_terminal_runs() {
247 let desc = ImageDescriptor::new(4, 4, PixelFormat::Gray8).unwrap();
248 let producer = Arc::new(CountingProducer::new(desc));
249 let image = Image::from_producer(Arc::clone(&producer) as Arc<dyn Producer>, Format::Raw)
250 .apply(Arc::new(ConstantOp::new(1)))
251 .unwrap();
252 assert_eq!(producer.produce_calls(), 0);
253 let _ = image.metadata().unwrap();
254 assert_eq!(producer.produce_calls(), 0, "metadata must not pull pixels");
255 evaluate(&image).unwrap();
256 assert_eq!(producer.produce_calls(), 1);
257 }
258
259 #[test]
260 fn deep_chains_do_not_overflow_the_stack() {
261 let mut image = ramp(2, 2);
262 for _ in 0..10_000 {
263 image = image.apply(Arc::new(ConstantOp::new(3))).unwrap();
264 }
265 assert_eq!(evaluate(&image).unwrap().bytes(), &[3, 3, 3, 3]);
266 }
267
268 #[test]
269 fn a_failing_op_fails_the_whole_evaluation() {
270 let image = ramp(2, 2)
271 .apply(Arc::new(crate::testing::FailingOp))
272 .unwrap();
273 let err = evaluate(&image).unwrap_err();
274 assert_eq!(err.code(), crate::ErrorCode::Malformed);
275 }
276
277 #[test]
278 fn evaluate_rows_yields_rows_in_order() {
279 let mut seen = Vec::new();
280 evaluate_rows(&ramp(2, 3), |y, row| {
281 seen.push((y, row.to_vec()));
282 Ok(())
283 })
284 .unwrap();
285 assert_eq!(
286 seen,
287 vec![(0, vec![0, 1]), (1, vec![2, 3]), (2, vec![4, 5])]
288 );
289 }
290
291 #[test]
292 fn evaluate_rows_propagates_consumer_errors() {
293 let err = evaluate_rows(&ramp(2, 3), |_, _| {
294 Err(PixelsError::unsupported("sink refused"))
295 })
296 .unwrap_err();
297 assert_eq!(err.code(), crate::ErrorCode::Unsupported);
298 }
299
300 #[test]
301 fn demand_walks_the_inverse_mapping_without_computing() {
302 let desc = ImageDescriptor::new(4, 4, PixelFormat::Gray8).unwrap();
303 let producer = Arc::new(CountingProducer::new(desc));
304 let image = Image::from_producer(Arc::clone(&producer) as Arc<dyn Producer>, Format::Raw)
305 .apply(Arc::new(ConstantOp::new(1)))
306 .unwrap();
307 let pairs = demand(&image, Region::from_size(4, 4)).unwrap();
308 assert_eq!(pairs.len(), 2, "one op node and one source node");
309 assert_eq!(
310 producer.produce_calls(),
311 0,
312 "demand propagation touches no pixels"
313 );
314 }
315}