Skip to main content

radiate_gp/collections/
store.rs

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        // hmmmm
411        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        // Test map operation that counts total values
445        let total = store.map(|values| values.len()).unwrap();
446        assert_eq!(total, 5);
447
448        // Test map operation that sums all values
449        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        // Test map_by_type for Input
462        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        // Test map_by_type for Output
470        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        // Test map_by_type for non-existent type
478        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        // map should return None for empty store
514        assert!(store.map(|_| 42).is_none());
515
516        // map_by_type should return None for empty store
517        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        // First insert
525        store.insert(NodeType::Input, vec![1, 2, 3]);
526
527        // Second insert should overwrite
528        store.insert(NodeType::Input, vec![4, 5]);
529
530        // Verify only new values exist
531        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        // Serialize to JSON
546        let serialized = serde_json::to_string(&store).unwrap();
547
548        // Deserialize back
549        let deserialized: NodeStore<i32> = serde_json::from_str(&serialized).unwrap();
550
551        // Verify the contents
552        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        // Create a store with only bounded values
572        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        // Create a store with only unbound values
598        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        // Mix of bounded and unbound values
613        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        // Add values for all node types
640        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        // Test with a more complex type (String)
657        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}