1use crate::{Format, ImageDescriptor, Metadata, Op, PixelsError, Producer, Region, Result};
13use std::sync::Arc;
14use std::sync::atomic::{AtomicU64, Ordering};
15
16#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
22pub struct NodeId(u64);
23
24impl NodeId {
25 #[must_use]
27 pub const fn get(self) -> u64 {
28 self.0
29 }
30
31 fn next() -> Self {
33 static COUNTER: AtomicU64 = AtomicU64::new(0);
34 Self(COUNTER.fetch_add(1, Ordering::Relaxed))
35 }
36}
37
38#[derive(Debug, Clone)]
40enum NodeKind {
41 Source(Arc<dyn Producer>),
43 Op(Arc<dyn Op>),
45}
46
47#[derive(Debug)]
52pub struct Node {
53 id: NodeId,
54 kind: NodeKind,
55 inputs: Vec<Arc<Node>>,
56 descriptor: ImageDescriptor,
57}
58
59impl Node {
60 #[must_use]
62 pub const fn id(&self) -> NodeId {
63 self.id
64 }
65
66 #[must_use]
68 pub const fn descriptor(&self) -> ImageDescriptor {
69 self.descriptor
70 }
71
72 #[must_use]
74 pub fn inputs(&self) -> &[Arc<Node>] {
75 &self.inputs
76 }
77
78 #[must_use]
80 pub fn op(&self) -> Option<&Arc<dyn Op>> {
81 match &self.kind {
82 NodeKind::Op(op) => Some(op),
83 NodeKind::Source(_) => None,
84 }
85 }
86
87 #[must_use]
89 pub fn producer(&self) -> Option<&Arc<dyn Producer>> {
90 match &self.kind {
91 NodeKind::Source(producer) => Some(producer),
92 NodeKind::Op(_) => None,
93 }
94 }
95
96 #[must_use]
98 pub fn name(&self) -> &'static str {
99 match &self.kind {
100 NodeKind::Source(producer) => producer.name(),
101 NodeKind::Op(op) => op.name(),
102 }
103 }
104
105 #[must_use]
110 pub fn node_count(self: &Arc<Self>) -> usize {
111 let mut seen = std::collections::HashSet::new();
112 let mut stack = vec![Arc::clone(self)];
113 while let Some(node) = stack.pop() {
114 if !seen.insert(node.id) {
115 continue;
116 }
117 stack.extend(node.inputs.iter().map(Arc::clone));
118 }
119 seen.len()
120 }
121}
122
123impl Drop for Node {
124 fn drop(&mut self) {
135 let mut stack: Vec<Arc<Self>> = std::mem::take(&mut self.inputs);
136 while let Some(node) = stack.pop() {
137 if let Some(mut node) = Arc::into_inner(node) {
138 stack.append(&mut node.inputs);
141 }
142 }
143 }
144}
145
146#[derive(Debug, Clone)]
155pub struct Image {
156 node: Arc<Node>,
157 format: Format,
158}
159
160impl Image {
161 #[must_use]
167 pub fn from_producer(producer: Arc<dyn Producer>, format: Format) -> Self {
168 let descriptor = producer.descriptor();
169 Self {
170 node: Arc::new(Node {
171 id: NodeId::next(),
172 kind: NodeKind::Source(producer),
173 inputs: Vec::new(),
174 descriptor,
175 }),
176 format,
177 }
178 }
179
180 #[must_use]
182 pub fn node(&self) -> &Arc<Node> {
183 &self.node
184 }
185
186 #[must_use]
188 pub fn descriptor(&self) -> ImageDescriptor {
189 self.node.descriptor
190 }
191
192 pub fn metadata(&self) -> Result<Metadata> {
203 Ok(Metadata::new(&self.node.descriptor, self.format))
204 }
205
206 pub fn apply(&self, op: Arc<dyn Op>) -> Result<Self> {
216 Self::combine(std::slice::from_ref(self), op)
217 }
218
219 pub fn combine(inputs: &[Self], op: Arc<dyn Op>) -> Result<Self> {
230 let Some(first) = inputs.first() else {
231 return Err(PixelsError::graph(format!(
232 "op `{}` needs at least one input",
233 op.name()
234 )));
235 };
236 if inputs.len() != op.arity() {
237 return Err(PixelsError::graph(format!(
238 "op `{}` takes {} input(s), got {}",
239 op.name(),
240 op.arity(),
241 inputs.len()
242 )));
243 }
244 let descriptors: Vec<ImageDescriptor> = inputs.iter().map(Self::descriptor).collect();
245 let descriptor = op.output_descriptor(&descriptors)?;
246 let format = first.format;
247 Ok(Self {
248 node: Arc::new(Node {
249 id: NodeId::next(),
250 kind: NodeKind::Op(op),
251 inputs: inputs.iter().map(|image| Arc::clone(&image.node)).collect(),
252 descriptor,
253 }),
254 format,
255 })
256 }
257
258 #[must_use]
260 pub fn region(&self) -> Region {
261 self.node.descriptor.region()
262 }
263}
264
265#[cfg(test)]
266#[allow(
267 clippy::unwrap_used,
268 clippy::indexing_slicing,
269 reason = "tests operate on known-good values and assert shapes directly"
270)]
271mod tests {
272 use super::*;
273 use crate::testing::{ConstantOp, CountingProducer};
274 use crate::{AccessPattern, Op, PixelFormat, Tile, TileMut};
275
276 fn source(width: u32, height: u32) -> Image {
277 let producer =
278 CountingProducer::new(ImageDescriptor::new(width, height, PixelFormat::Gray8).unwrap());
279 Image::from_producer(Arc::new(producer), Format::Raw)
280 }
281
282 #[test]
283 fn chaining_builds_a_dag_without_touching_pixels() {
284 let producer = Arc::new(CountingProducer::new(
285 ImageDescriptor::new(4, 4, PixelFormat::Gray8).unwrap(),
286 ));
287 let image = Image::from_producer(Arc::clone(&producer) as Arc<dyn Producer>, Format::Raw);
288 let chained = image.apply(Arc::new(ConstantOp::new(7))).unwrap();
289 let _ = chained.apply(Arc::new(ConstantOp::new(9))).unwrap();
290 assert_eq!(
291 producer.produce_calls(),
292 0,
293 "graph construction must not pull pixels"
294 );
295 }
296
297 #[test]
298 fn metadata_is_available_without_evaluation() {
299 let producer = Arc::new(CountingProducer::new(
300 ImageDescriptor::new(6, 3, PixelFormat::Gray8).unwrap(),
301 ));
302 let image = Image::from_producer(Arc::clone(&producer) as Arc<dyn Producer>, Format::Raw);
303 let meta = image.metadata().unwrap();
304 assert_eq!((meta.width, meta.height), (6, 3));
305 assert_eq!(meta.format, Format::Raw);
306 assert_eq!(meta.pixel, PixelFormat::Gray8);
307 assert_eq!(producer.produce_calls(), 0);
308 }
309
310 #[test]
311 fn descriptors_flow_forward_through_the_chain() {
312 #[derive(Debug)]
314 struct Halve;
315 impl Op for Halve {
316 fn name(&self) -> &'static str {
317 "halve"
318 }
319 fn output_descriptor(&self, inputs: &[ImageDescriptor]) -> Result<ImageDescriptor> {
320 let input = inputs
321 .first()
322 .ok_or_else(|| PixelsError::graph("no input"))?;
323 input.resized(input.width / 2, input.height)
324 }
325 fn input_regions(&self, out: Region, _: &[ImageDescriptor]) -> Result<Vec<Region>> {
326 Ok(vec![out])
327 }
328 fn access_pattern(&self) -> AccessPattern {
329 AccessPattern::Sequential
330 }
331 fn compute(&self, _: &[Tile<'_>], _: &mut TileMut<'_>) -> Result<()> {
332 Ok(())
333 }
334 }
335 let image = source(16, 4)
336 .apply(Arc::new(Halve))
337 .unwrap()
338 .apply(Arc::new(Halve))
339 .unwrap();
340 assert_eq!(image.descriptor().width, 4);
341 assert_eq!(image.descriptor().height, 4);
342 }
343
344 #[test]
345 fn an_op_rejecting_its_input_fails_at_build_time() {
346 #[derive(Debug)]
348 struct Refuses;
349 impl Op for Refuses {
350 fn name(&self) -> &'static str {
351 "refuses"
352 }
353 fn output_descriptor(&self, _: &[ImageDescriptor]) -> Result<ImageDescriptor> {
354 Err(PixelsError::unsupported("never applicable"))
355 }
356 fn input_regions(&self, out: Region, _: &[ImageDescriptor]) -> Result<Vec<Region>> {
357 Ok(vec![out])
358 }
359 fn compute(&self, _: &[Tile<'_>], _: &mut TileMut<'_>) -> Result<()> {
360 Ok(())
361 }
362 }
363 let err = source(4, 4).apply(Arc::new(Refuses)).unwrap_err();
364 assert_eq!(err.code(), crate::ErrorCode::Unsupported);
365 }
366
367 #[test]
368 fn branches_share_their_common_prefix() {
369 let base = source(4, 4).apply(Arc::new(ConstantOp::new(1))).unwrap();
370 let left = base.apply(Arc::new(ConstantOp::new(2))).unwrap();
371 let right = base.apply(Arc::new(ConstantOp::new(3))).unwrap();
372 assert_eq!(left.node().inputs()[0].id(), right.node().inputs()[0].id());
374 assert_eq!(left.node().node_count(), 3);
376 }
377
378 #[test]
379 fn cloning_an_image_shares_the_node() {
380 let image = source(4, 4);
381 let clone = image.clone();
382 assert_eq!(image.node().id(), clone.node().id());
383 assert!(Arc::ptr_eq(image.node(), clone.node()));
384 }
385
386 #[test]
387 fn node_ids_are_unique() {
388 let a = source(2, 2);
389 let b = source(2, 2);
390 assert_ne!(a.node().id(), b.node().id());
391 assert_ne!(a.node().id().get(), b.node().id().get());
392 }
393
394 #[test]
395 fn arity_mismatch_is_a_graph_error() {
396 let image = source(4, 4);
397 let err =
398 Image::combine(&[image.clone(), image], Arc::new(ConstantOp::new(1))).unwrap_err();
399 assert_eq!(err.code(), crate::ErrorCode::Graph);
400 let err = Image::combine(&[], Arc::new(ConstantOp::new(1))).unwrap_err();
401 assert_eq!(err.code(), crate::ErrorCode::Graph);
402 }
403
404 #[test]
405 fn images_are_send_and_sync() {
406 const fn assert_send_sync<T: Send + Sync>() {}
407 assert_send_sync::<Image>();
408 assert_send_sync::<Arc<Node>>();
409 }
410}