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