Skip to main content

radiate_gp/collections/
store.rs

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