Skip to main content

radiate_gp/collections/
factory.rs

1use super::{GraphNode, NodeStore, NodeType, NodeValue, TreeNode};
2use crate::Arity;
3use radiate_core::random_provider;
4
5pub trait Factory<I, O> {
6    fn new_instance(&self, input: I) -> O;
7}
8
9impl<T> Factory<(), T> for NodeValue<T>
10where
11    T: Factory<(), T>,
12{
13    fn new_instance(&self, _: ()) -> T {
14        match self {
15            NodeValue::Bounded(value, _) => value.new_instance(()),
16            NodeValue::Unbound(value) => value.new_instance(()),
17        }
18    }
19}
20
21impl<T> Factory<NodeType, T> for NodeStore<T>
22where
23    T: Factory<(), T> + Default,
24{
25    fn new_instance(&self, input: NodeType) -> T {
26        self.map_by_type(input, |values| {
27            random_provider::choose(values).new_instance(())
28        })
29        .unwrap_or_default()
30    }
31}
32
33impl<T: Default + Clone> Factory<(usize, NodeType), Option<GraphNode<T>>> for NodeStore<T> {
34    fn new_instance(&self, (index, node_type): (usize, NodeType)) -> Option<GraphNode<T>> {
35        self.map_by_type(node_type, |values| {
36            let node_value = match node_type {
37                NodeType::Input => &values[index % values.len()],
38                _ => random_provider::choose(values),
39            };
40
41            match node_value {
42                NodeValue::Bounded(value, arity) => {
43                    (index, node_type, value.clone(), *arity).into()
44                }
45                NodeValue::Unbound(value) => (index, node_type, value.clone()).into(),
46            }
47        })
48    }
49}
50
51impl<T, F> Factory<(usize, NodeType, F), Option<GraphNode<T>>> for NodeStore<T>
52where
53    T: Default + Clone,
54    F: Fn(Arity) -> bool,
55{
56    fn new_instance(
57        &self,
58        (index, node_type, filter): (usize, NodeType, F),
59    ) -> Option<GraphNode<T>> {
60        self.map(|values| {
61            let mapped_values = values
62                .into_iter()
63                .filter(|value| match value {
64                    NodeValue::Bounded(_, arity) => filter(*arity),
65                    _ => false,
66                })
67                .collect::<Vec<&NodeValue<T>>>();
68
69            if mapped_values.is_empty() {
70                self.new_instance((index, node_type))
71            } else {
72                let node_value = random_provider::choose(&mapped_values);
73
74                match node_value {
75                    NodeValue::Bounded(value, arity) => Some(GraphNode::with_arity(
76                        index,
77                        node_type,
78                        value.clone(),
79                        *arity,
80                    )),
81                    NodeValue::Unbound(value) => {
82                        Some(GraphNode::new(index, node_type, value.clone()))
83                    }
84                }
85            }
86        })
87        .flatten()
88    }
89}
90
91impl<T, F> Factory<F, Option<TreeNode<T>>> for NodeStore<T>
92where
93    T: Default + Clone,
94    F: Fn(Arity) -> bool,
95{
96    fn new_instance(&self, input: F) -> Option<TreeNode<T>> {
97        self.map(|values| {
98            let mapped_values = values
99                .into_iter()
100                .filter(|value| match value {
101                    NodeValue::Bounded(_, arity) => input(*arity),
102                    _ => false,
103                })
104                .collect::<Vec<&NodeValue<T>>>();
105
106            if mapped_values.is_empty() {
107                TreeNode::new(T::default())
108            } else {
109                let node_value = random_provider::choose(&mapped_values);
110
111                match node_value {
112                    NodeValue::Bounded(value, arity) => TreeNode::with_arity(value.clone(), *arity),
113                    NodeValue::Unbound(value) => TreeNode::new(value.clone()),
114                }
115            }
116        })
117    }
118}
119
120impl<T: Clone + Default> Factory<NodeType, Option<TreeNode<T>>> for NodeStore<T> {
121    fn new_instance(&self, input: NodeType) -> Option<TreeNode<T>> {
122        self.map_by_type(input, |values| {
123            let node_value = random_provider::choose(values);
124
125            match node_value {
126                NodeValue::Bounded(value, arity) => TreeNode::with_arity(value.clone(), *arity),
127                NodeValue::Unbound(value) => TreeNode::new(value.clone()),
128            }
129        })
130    }
131}