1use super::NodeType;
2use crate::{Arity, Op};
3#[cfg(feature = "serde")]
4use serde::{Deserialize, Serialize, de::Deserializer, ser::Serializer};
5use std::collections::{BTreeMap, HashMap};
6use std::fmt::Debug;
7use std::sync::Arc;
8
9#[derive(Debug, Clone, PartialEq)]
10#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
11pub enum NodeValue<T> {
12 Bounded(T, Arity),
13 Unbound(T),
14}
15
16impl<T> NodeValue<T> {
17 pub fn value(&self) -> &T {
18 match self {
19 NodeValue::Bounded(value, _) => value,
20 NodeValue::Unbound(value) => value,
21 }
22 }
23
24 pub fn arity(&self) -> Option<Arity> {
25 match self {
26 NodeValue::Bounded(_, arity) => Some(*arity),
27 NodeValue::Unbound(_) => None,
28 }
29 }
30
31 pub fn allowed_node_types(&self) -> Vec<NodeType> {
32 match self.arity().unwrap_or(Arity::Any) {
33 Arity::Zero => vec![NodeType::Input, NodeType::Leaf],
34 Arity::Any => vec![NodeType::Output, NodeType::Root, NodeType::Vertex],
35 Arity::Exact(1) => vec![NodeType::Edge, NodeType::Vertex],
36 _ => vec![NodeType::Vertex],
37 }
38 }
39}
40
41macro_rules! impl_node_value {
42 ($($t:ty),*) => {
43 $(
44 impl From<$t> for NodeValue<$t> {
45 fn from(value: $t) -> Self {
46 NodeValue::Unbound(value)
47 }
48 }
49 )*
50 };
51}
52
53impl_node_value!(
54 u8,
55 u16,
56 u32,
57 u64,
58 u128,
59 i8,
60 i16,
61 i32,
62 i64,
63 i128,
64 f32,
65 f64,
66 String,
67 bool,
68 char,
69 usize,
70 isize,
71 &'static str
72);
73
74#[derive(Default)]
75pub struct NodeStore<T> {
76 values: Arc<BTreeMap<NodeType, Vec<NodeValue<T>>>>,
77}
78
79impl<T> NodeStore<T> {
80 pub fn new() -> Self {
81 NodeStore {
82 values: Arc::new(BTreeMap::new()),
83 }
84 }
85
86 pub fn ref_count(&self) -> usize {
87 Arc::strong_count(&self.values)
88 }
89
90 pub fn count_type(&self, node_type: NodeType) -> usize {
91 let values = &self.values;
92 if let Some(values) = values.get(&node_type) {
93 return values.len();
94 }
95
96 0
97 }
98
99 pub fn contains_type(&self, node_type: NodeType) -> bool {
100 let values = &self.values;
101 values.contains_key(&node_type)
102 && values
103 .get(&node_type)
104 .is_some_and(|values| !values.is_empty())
105 }
106
107 pub fn add(&mut self, values: Vec<T>)
108 where
109 T: Into<NodeValue<T>> + Clone,
110 {
111 let store = Arc::make_mut(&mut self.values);
112
113 for value in values {
114 let node_value = value.into();
115 for node_type in node_value.allowed_node_types() {
116 store.entry(node_type).or_default().push(node_value.clone());
117 }
118 }
119 }
120
121 pub fn insert<K>(&mut self, node_type: NodeType, values: Vec<K>)
122 where
123 T: Clone,
124 K: Into<NodeValue<T>>,
125 {
126 let store = Arc::make_mut(&mut self.values);
127 store.insert(node_type, values.into_iter().map(|x| x.into()).collect());
128 }
129
130 pub fn map<F, K>(&self, mapper: F) -> Option<K>
131 where
132 F: Fn(Vec<&NodeValue<T>>) -> K,
133 {
134 let values = &self.values;
135 let all_values = values.values().flatten().collect::<Vec<&NodeValue<T>>>();
136
137 if all_values.is_empty() {
138 return None;
139 }
140
141 Some(mapper(all_values))
142 }
143
144 pub fn map_by_type<F, K>(&self, node_type: NodeType, mapper: F) -> Option<K>
145 where
146 F: Fn(&[NodeValue<T>]) -> K,
147 {
148 let values = &self.values;
149 if let Some(values) = values.get(&node_type) {
150 return Some(mapper(values));
151 }
152
153 None
154 }
155}
156
157impl<T> From<HashMap<NodeType, Vec<T>>> for NodeStore<T>
158where
159 T: Into<NodeValue<T>> + Clone,
160{
161 fn from(values: HashMap<NodeType, Vec<T>>) -> Self {
162 let mut store = NodeStore::new();
163 for (node_type, ops) in values {
164 store.insert(node_type, ops);
165 }
166
167 store
168 }
169}
170
171impl<T> From<Vec<(NodeType, Vec<T>)>> for NodeStore<T>
172where
173 T: Into<NodeValue<T>> + Clone,
174{
175 fn from(values: Vec<(NodeType, Vec<T>)>) -> Self {
176 let mut store = NodeStore::new();
177 for (node_type, ops) in values {
178 store.insert(node_type, ops);
179 }
180
181 if !store.contains_type(NodeType::Leaf) && store.contains_type(NodeType::Input) {
182 let input_values = store
183 .map(|vals| {
184 vals.iter()
185 .filter_map(|v| match v.arity() {
186 Some(Arity::Zero) => Some((*v.value()).clone()),
187 _ => None,
188 })
189 .collect::<Vec<_>>()
190 })
191 .unwrap_or_default();
192
193 store.insert(NodeType::Leaf, input_values);
194 }
195
196 if !store.contains_type(NodeType::Root) && store.contains_type(NodeType::Output) {
197 let output_values = store
198 .map(|vals| {
199 vals.iter()
200 .filter_map(|v| match v.arity() {
201 Some(Arity::Any) | Some(Arity::Exact(_)) => Some((*v.value()).clone()),
202 _ => None,
203 })
204 .collect::<Vec<_>>()
205 })
206 .unwrap_or_default();
207
208 store.insert(NodeType::Root, output_values);
209 }
210
211 store
212 }
213}
214
215impl<T> From<Vec<T>> for NodeStore<T>
216where
217 T: Into<NodeValue<T>> + Clone,
218{
219 fn from(values: Vec<T>) -> Self {
220 let mut store = NodeStore::new();
221 store.add(values);
222 store
223 }
224}
225
226impl<T: Clone> From<Op<T>> for NodeStore<Op<T>> {
227 fn from(value: Op<T>) -> Self {
228 let mut store = NodeStore::new();
229
230 let input_values = vec![Op::var(0)];
231 let output_values = vec![value.clone()];
232 let edge_values = vec![Op::identity()];
233 let node_values = vec![value.clone()];
234
235 store.insert(NodeType::Input, input_values);
236 store.insert(NodeType::Output, output_values);
237 store.insert(NodeType::Edge, edge_values);
238 store.insert(NodeType::Vertex, node_values);
239
240 store
241 }
242}
243
244impl<T: Clone> From<&NodeStore<T>> for NodeStore<T> {
245 fn from(store: &NodeStore<T>) -> Self {
246 NodeStore {
247 values: Arc::clone(&store.values),
248 }
249 }
250}
251
252impl<T> Clone for NodeStore<T> {
253 fn clone(&self) -> Self {
254 NodeStore {
255 values: Arc::clone(&self.values),
256 }
257 }
258}
259
260impl<T: PartialEq> PartialEq for NodeStore<T> {
261 fn eq(&self, other: &Self) -> bool {
262 self.values == other.values
263 }
264}
265
266impl<T: Debug> Debug for NodeStore<T> {
267 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
268 for (node_type, values) in self.values.iter() {
269 writeln!(f, "{node_type:?}:")?;
270 for value in values {
271 writeln!(f, " {value:?}")?;
272 }
273 }
274
275 Ok(())
276 }
277}
278
279#[cfg(feature = "serde")]
280impl<T> Serialize for NodeStore<T>
281where
282 T: Serialize,
283{
284 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
285 where
286 S: Serializer,
287 {
288 let serializable = self.values.iter().collect::<Vec<_>>();
289
290 serializable.serialize(serializer)
291 }
292}
293
294#[cfg(feature = "serde")]
295impl<'de, T> Deserialize<'de> for NodeStore<T>
296where
297 T: Deserialize<'de>,
298{
299 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
300 where
301 D: Deserializer<'de>,
302 {
303 let values: Vec<(NodeType, Vec<NodeValue<T>>)> = Vec::deserialize(deserializer)?;
304
305 let mut map = BTreeMap::new();
306 for (node_type, node_values) in values {
307 map.insert(node_type, node_values);
308 }
309
310 Ok(NodeStore {
311 values: Arc::new(map),
312 })
313 }
314}
315
316#[macro_export]
317macro_rules! node_store {
318 ($($node_type:ident => $values:expr),+) => {
319 {
320 let mut store = NodeStore::new();
321 $(
322 store.insert(NodeType::$node_type, $values);
323 )*
324 store
325 }
326 };
327}
328
329#[cfg(test)]
330mod tests {
331 use super::*;
332 use crate::{Factory, Node, TreeNode, ops};
333
334 #[allow(dead_code)]
335 fn create_test_store() -> NodeStore<i32> {
336 let mut store = NodeStore::new();
337
338 store.insert(NodeType::Input, vec![1, 2, 3]);
339 store.insert(NodeType::Output, vec![4, 5]);
340 store.insert(NodeType::Vertex, vec![6, 7, 8, 9]);
341
342 let bounded_values = vec![
343 NodeValue::Bounded(10, Arity::Exact(2)),
344 NodeValue::Bounded(11, Arity::Zero),
345 ];
346
347 store.insert(
348 NodeType::Edge,
349 bounded_values
350 .into_iter()
351 .map(|v| v.value().clone())
352 .collect(),
353 );
354
355 store
356 }
357
358 #[test]
359 fn test_node_store() {
360 let mut store = NodeStore::<Op<f32>>::from(ops::all_ops());
361
362 store.add(Op::vars(0..3));
363
364 assert!(store.contains_type(NodeType::Input));
365 assert!(store.contains_type(NodeType::Output));
366 assert!(store.contains_type(NodeType::Edge));
367 assert!(store.contains_type(NodeType::Vertex));
368 assert!(store.contains_type(NodeType::Leaf));
369 assert!(store.contains_type(NodeType::Root));
370 }
371
372 #[test]
373 fn test_node_store_insert() {
374 let mut store = NodeStore::new();
375 let values = vec![1, 2, 3];
376 store.insert(NodeType::Input, values.clone());
377
378 assert!(store.contains_type(NodeType::Input));
379
380 for value in values {
381 assert!(
382 store
383 .map_by_type(NodeType::Input, |values| {
384 values.iter().any(|v| v.value() == &value)
385 })
386 .unwrap_or(false)
387 );
388 }
389 }
390
391 #[test]
392 fn test_node_store_macro() {
393 let store = node_store! {
394 Input => vec![1, 2, 3],
395 Output => vec![4, 5, 6],
396 Edge => vec![7, 8, 9],
397 Vertex => vec![10, 11, 12]
398 };
399
400 assert!(store.contains_type(NodeType::Input));
401 assert!(store.contains_type(NodeType::Output));
402 assert!(store.contains_type(NodeType::Edge));
403 assert!(store.contains_type(NodeType::Vertex));
404
405 let graph_node = store.new_instance((2, NodeType::Vertex)).unwrap();
406
407 assert_eq!(graph_node.index(), 2);
408 assert_eq!(graph_node.node_type(), NodeType::Vertex);
409
410 let tree_node: Option<TreeNode<i32>> = store.new_instance(NodeType::Vertex);
412 let tree_node = tree_node.unwrap();
413 assert_eq!(tree_node.node_type(), NodeType::Leaf);
414 assert!(tree_node.is_leaf());
415 }
416
417 #[test]
418 fn test_insert_and_contains() {
419 let mut store = NodeStore::new();
420
421 store.insert(NodeType::Input, vec![1, 2, 3]);
422 assert!(store.contains_type(NodeType::Input));
423
424 store.insert(NodeType::Output, vec![4, 5]);
425 assert!(store.contains_type(NodeType::Output));
426
427 assert!(!store.contains_type(NodeType::Vertex));
428 }
429
430 #[test]
431 fn test_new_store_is_empty() {
432 let store: NodeStore<i32> = NodeStore::new();
433 assert!(!store.contains_type(NodeType::Input));
434 assert!(!store.contains_type(NodeType::Output));
435 assert!(!store.contains_type(NodeType::Vertex));
436 }
437
438 #[test]
439 fn test_map_operation() {
440 let mut store = NodeStore::new();
441 store.insert(NodeType::Input, vec![1, 2, 3]);
442 store.insert(NodeType::Output, vec![4, 5]);
443
444 let total = store.map(|values| values.len()).unwrap();
446 assert_eq!(total, 5);
447
448 let sum: i32 = store
450 .map(|values| values.iter().map(|v| v.value()).sum())
451 .unwrap();
452 assert_eq!(sum, 15);
453 }
454
455 #[test]
456 fn test_map_by_type() {
457 let mut store = NodeStore::new();
458 store.insert(NodeType::Input, vec![1, 2, 3]);
459 store.insert(NodeType::Output, vec![4, 5]);
460
461 let input_sum: i32 = store
463 .map_by_type(NodeType::Input, |values| {
464 values.iter().map(|v| v.value()).sum()
465 })
466 .unwrap();
467 assert_eq!(input_sum, 6);
468
469 let output_sum: i32 = store
471 .map_by_type(NodeType::Output, |values| {
472 values.iter().map(|v| v.value()).sum()
473 })
474 .unwrap();
475 assert_eq!(output_sum, 9);
476
477 let result = store.map_by_type(NodeType::Vertex, |values| values.len());
479 assert!(result.is_none());
480 }
481
482 #[test]
483 fn test_from_hashmap() {
484 let mut map = HashMap::new();
485 map.insert(NodeType::Input, vec![1, 2, 3]);
486 map.insert(NodeType::Output, vec![4, 5]);
487
488 let store: NodeStore<i32> = map.into();
489
490 assert!(store.contains_type(NodeType::Input));
491 assert!(store.contains_type(NodeType::Output));
492 assert!(!store.contains_type(NodeType::Vertex));
493 }
494
495 #[test]
496 fn test_from_vec_of_tuples() {
497 let values = vec![
498 (NodeType::Input, vec![1, 2, 3]),
499 (NodeType::Output, vec![4, 5]),
500 ];
501
502 let store: NodeStore<i32> = values.into();
503
504 assert!(store.contains_type(NodeType::Input));
505 assert!(store.contains_type(NodeType::Output));
506 assert!(!store.contains_type(NodeType::Vertex));
507 }
508
509 #[test]
510 fn test_empty_map_returns_none() {
511 let store: NodeStore<i32> = NodeStore::new();
512
513 assert!(store.map(|_| 42).is_none());
515
516 assert!(store.map_by_type(NodeType::Input, |_| 42).is_none());
518 }
519
520 #[test]
521 fn test_insert_overwrites_existing() {
522 let mut store = NodeStore::new();
523
524 store.insert(NodeType::Input, vec![1, 2, 3]);
526
527 store.insert(NodeType::Input, vec![4, 5]);
529
530 let values: Vec<i32> = store
532 .map_by_type(NodeType::Input, |values: &[NodeValue<i32>]| {
533 values.iter().map(|v| v.value().clone()).collect()
534 })
535 .unwrap();
536
537 assert_eq!(values, vec![4, 5]);
538 }
539
540 #[test]
541 #[cfg(feature = "serde")]
542 fn test_serialize_deserialize_basic() {
543 let store = create_test_store();
544
545 let serialized = serde_json::to_string(&store).unwrap();
547
548 let deserialized: NodeStore<i32> = serde_json::from_str(&serialized).unwrap();
550
551 assert_eq!(store, deserialized);
553 }
554
555 #[test]
556 #[cfg(feature = "serde")]
557 fn test_serialize_deserialize_empty() {
558 let store: NodeStore<i32> = NodeStore::new();
559
560 let serialized = serde_json::to_string(&store).unwrap();
561 let deserialized: NodeStore<i32> = serde_json::from_str(&serialized).unwrap();
562
563 assert_eq!(store, deserialized);
564 }
565
566 #[test]
567 #[cfg(feature = "serde")]
568 fn test_serialize_deserialize_with_bounded_values() {
569 let mut store = NodeStore::new();
570
571 let bounded_values = vec![
573 NodeValue::Bounded(1, Arity::Exact(2)),
574 NodeValue::Bounded(2, Arity::Zero),
575 NodeValue::Bounded(3, Arity::Any),
576 ];
577
578 store.insert(
579 NodeType::Vertex,
580 bounded_values
581 .into_iter()
582 .map(|v| v.value().clone())
583 .collect(),
584 );
585
586 let serialized = serde_json::to_string(&store).unwrap();
587 let deserialized: NodeStore<i32> = serde_json::from_str(&serialized).unwrap();
588
589 assert_eq!(store, deserialized);
590 }
591
592 #[test]
593 #[cfg(feature = "serde")]
594 fn test_serialize_deserialize_with_unbound_values() {
595 let mut store = NodeStore::new();
596
597 let unbound_values = vec![1, 2, 3, 4, 5];
599 store.insert(NodeType::Vertex, unbound_values);
600
601 let serialized = serde_json::to_string(&store).unwrap();
602 let deserialized: NodeStore<i32> = serde_json::from_str(&serialized).unwrap();
603
604 assert_eq!(store, deserialized);
605 }
606
607 #[test]
608 #[cfg(feature = "serde")]
609 fn test_serialize_deserialize_mixed_values() {
610 let mut store = NodeStore::new();
611
612 let mixed_values = vec![
614 NodeValue::Bounded(1, Arity::Exact(2)),
615 NodeValue::Unbound(2),
616 NodeValue::Bounded(3, Arity::Zero),
617 NodeValue::Unbound(4),
618 ];
619
620 store.insert(
621 NodeType::Vertex,
622 mixed_values
623 .into_iter()
624 .map(|v| v.value().clone())
625 .collect(),
626 );
627
628 let serialized = serde_json::to_string(&store).unwrap();
629 let deserialized: NodeStore<i32> = serde_json::from_str(&serialized).unwrap();
630
631 assert_eq!(store, deserialized);
632 }
633
634 #[test]
635 #[cfg(feature = "serde")]
636 fn test_serialize_deserialize_all_node_types() {
637 let mut store = NodeStore::new();
638
639 store.insert(NodeType::Input, vec![1, 2]);
641 store.insert(NodeType::Output, vec![3, 4]);
642 store.insert(NodeType::Vertex, vec![5, 6]);
643 store.insert(NodeType::Edge, vec![7, 8]);
644 store.insert(NodeType::Leaf, vec![9, 10]);
645 store.insert(NodeType::Root, vec![11, 12]);
646
647 let serialized = serde_json::to_string(&store).unwrap();
648 let deserialized: NodeStore<i32> = serde_json::from_str(&serialized).unwrap();
649
650 assert_eq!(store, deserialized);
651 }
652
653 #[test]
654 #[cfg(feature = "serde")]
655 fn test_serialize_deserialize_complex_type() {
656 let mut store = NodeStore::new();
658
659 let values = vec![
660 NodeValue::Bounded("hello".to_string(), Arity::Exact(2)),
661 NodeValue::Unbound("world".to_string()),
662 NodeValue::Bounded("test".to_string(), Arity::Zero),
663 ];
664
665 store.insert(
666 NodeType::Vertex,
667 values.into_iter().map(|v| v.value().clone()).collect(),
668 );
669
670 let serialized = serde_json::to_string(&store).unwrap();
671 let deserialized: NodeStore<String> = serde_json::from_str(&serialized).unwrap();
672
673 assert_eq!(store, deserialized);
674 }
675}