Skip to main content

radiate_gp/collections/trees/
chromosome.rs

1use crate::{NodeStore, TreeNode};
2use radiate_core::{Chromosome, Valid};
3#[cfg(feature = "serde")]
4use serde::{Deserialize, Deserializer, Serialize, Serializer};
5use std::{fmt::Debug, hash::Hash, sync::Arc};
6
7type Constraint<N> = Arc<dyn Fn(&N) -> bool + Send + Sync>;
8
9#[derive(Clone, Default)]
10pub struct TreeChromosome<T> {
11    nodes: Vec<TreeNode<T>>,
12    store: Option<NodeStore<T>>,
13    constraint: Option<Constraint<TreeNode<T>>>,
14}
15
16impl<T> TreeChromosome<T> {
17    pub fn new(
18        nodes: Vec<TreeNode<T>>,
19        store: Option<NodeStore<T>>,
20        constraint: Option<Constraint<TreeNode<T>>>,
21    ) -> Self {
22        TreeChromosome {
23            nodes,
24            store,
25            constraint,
26        }
27    }
28
29    pub fn root(&self) -> &TreeNode<T> {
30        &self.nodes[0]
31    }
32
33    pub fn root_mut(&mut self) -> &mut TreeNode<T> {
34        &mut self.nodes[0]
35    }
36
37    pub fn get_store(&self) -> Option<NodeStore<T>> {
38        self.store.clone()
39    }
40}
41
42impl<T> Chromosome for TreeChromosome<T>
43where
44    T: Clone + PartialEq,
45{
46    type Gene = TreeNode<T>;
47
48    fn as_slice(&self) -> &[Self::Gene] {
49        &self.nodes
50    }
51
52    fn as_mut_slice(&mut self) -> &mut [Self::Gene] {
53        &mut self.nodes
54    }
55}
56
57impl<T> Valid for TreeChromosome<T> {
58    fn is_valid(&self) -> bool {
59        for gene in &self.nodes {
60            if let Some(constraint) = &self.constraint {
61                if !constraint(gene) {
62                    return false;
63                }
64            } else if !gene.is_valid() {
65                return false;
66            }
67        }
68
69        true
70    }
71}
72
73impl<T> From<Vec<TreeNode<T>>> for TreeChromosome<T> {
74    fn from(nodes: Vec<TreeNode<T>>) -> Self {
75        TreeChromosome {
76            nodes,
77            store: None,
78            constraint: None,
79        }
80    }
81}
82
83impl<T> FromIterator<TreeNode<T>> for TreeChromosome<T> {
84    fn from_iter<I: IntoIterator<Item = TreeNode<T>>>(iter: I) -> Self {
85        let nodes: Vec<TreeNode<T>> = iter.into_iter().collect();
86        TreeChromosome {
87            nodes,
88            store: None,
89            constraint: None,
90        }
91    }
92}
93
94impl<T> AsRef<[TreeNode<T>]> for TreeChromosome<T> {
95    fn as_ref(&self) -> &[TreeNode<T>] {
96        &self.nodes
97    }
98}
99
100impl<T> AsMut<[TreeNode<T>]> for TreeChromosome<T> {
101    fn as_mut(&mut self) -> &mut [TreeNode<T>] {
102        &mut self.nodes
103    }
104}
105
106impl<T: PartialEq> PartialEq for TreeChromosome<T> {
107    fn eq(&self, other: &Self) -> bool {
108        self.nodes == other.nodes
109    }
110}
111
112impl<T> IntoIterator for TreeChromosome<T> {
113    type Item = TreeNode<T>;
114    type IntoIter = std::vec::IntoIter<TreeNode<T>>;
115
116    fn into_iter(self) -> Self::IntoIter {
117        self.nodes.into_iter()
118    }
119}
120
121impl<T: Hash> Hash for TreeChromosome<T> {
122    fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
123        for node in self.as_ref() {
124            node.hash(state);
125        }
126    }
127}
128
129impl<T: Debug> Debug for TreeChromosome<T> {
130    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
131        f.debug_struct("TreeChromosome")
132            .field("nodes", &self.nodes)
133            .field("store", &self.store)
134            .field("constraint", &self.constraint.is_some())
135            .finish()
136    }
137}
138
139#[cfg(feature = "serde")]
140impl<T> Serialize for TreeChromosome<T>
141where
142    T: Serialize + Clone + PartialEq,
143{
144    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
145    where
146        S: Serializer,
147    {
148        (&self.nodes, &self.store).serialize(serializer)
149    }
150}
151
152#[cfg(feature = "serde")]
153impl<'de, T> Deserialize<'de> for TreeChromosome<T>
154where
155    T: Deserialize<'de> + Clone + PartialEq,
156{
157    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
158    where
159        D: Deserializer<'de>,
160    {
161        let (nodes, store): (Vec<TreeNode<T>>, Option<NodeStore<T>>) =
162            Deserialize::deserialize(deserializer)?;
163
164        Ok(TreeChromosome {
165            nodes,
166            store,
167            constraint: None, // There is no good way to serialize constraints directly
168        })
169    }
170}
171
172#[cfg(test)]
173mod tests {
174    use super::*;
175    use crate::{Node, NodeType};
176
177    fn create_test_chromosome() -> TreeChromosome<i32> {
178        let store = NodeStore::new();
179        store.insert(NodeType::Vertex, vec![1, 2, 3]);
180        store.insert(NodeType::Leaf, vec![4, 5]);
181
182        // Create a simple tree: 1 -> (2 -> (4, 5), 3)
183        let root = TreeNode::with_children(
184            1,
185            vec![
186                TreeNode::with_children(2, vec![TreeNode::new(4), TreeNode::new(5)]),
187                TreeNode::new(3),
188            ],
189        );
190
191        TreeChromosome::new(vec![root], Some(store), None)
192    }
193
194    #[test]
195    fn test_new_chromosome() {
196        let chromosome = TreeChromosome::new(vec![TreeNode::new(42)], None, None);
197
198        assert_eq!(chromosome.nodes.len(), 1);
199        assert_eq!(chromosome.store, None);
200        assert!(chromosome.constraint.is_none());
201        assert_eq!(chromosome.root().value(), &42);
202    }
203
204    #[test]
205    fn test_root_access() {
206        let chromosome = create_test_chromosome();
207
208        assert_eq!(chromosome.root().value(), &1);
209
210        let mut chromosome = chromosome;
211        let root_mut = chromosome.root_mut();
212
213        assert_eq!(root_mut.value(), &1);
214
215        *root_mut.value_mut() = 10;
216
217        assert_eq!(chromosome.root().value(), &10);
218    }
219
220    #[test]
221    fn test_store_access() {
222        let chromosome = create_test_chromosome();
223        let store = chromosome.get_store();
224
225        assert!(store.is_some());
226
227        let store = store.unwrap();
228        assert!(store.contains_type(NodeType::Vertex));
229        assert!(store.contains_type(NodeType::Leaf));
230    }
231
232    #[test]
233    fn test_constraint_validation() {
234        let constraint = Arc::new(|node: &TreeNode<i32>| node.value() % 2 == 0);
235        let chromosome = TreeChromosome::new(
236            vec![TreeNode::with_children(
237                2,
238                vec![TreeNode::new(4), TreeNode::new(6)],
239            )],
240            None,
241            Some(constraint.clone()),
242        );
243
244        assert!(chromosome.is_valid());
245
246        let invalid_chromosome = TreeChromosome::new(
247            vec![TreeNode::with_children(
248                1,
249                vec![TreeNode::new(4), TreeNode::new(6)],
250            )],
251            None,
252            Some(constraint),
253        );
254        assert!(!invalid_chromosome.is_valid());
255    }
256
257    #[test]
258    fn test_partial_eq() {
259        let chromosome1 = create_test_chromosome();
260        let chromosome2 = create_test_chromosome();
261        let chromosome3 = TreeChromosome::new(vec![TreeNode::new(42)], None, None);
262
263        assert_eq!(chromosome1, chromosome2);
264        assert_ne!(chromosome1, chromosome3);
265    }
266
267    #[test]
268    #[cfg(feature = "serde")]
269    fn test_serialize_deserialize_basic() {
270        let chromosome = create_test_chromosome();
271
272        let serialized = serde_json::to_string(&chromosome).unwrap();
273        let deserialized: TreeChromosome<i32> = serde_json::from_str(&serialized).unwrap();
274
275        assert_eq!(chromosome.nodes, deserialized.nodes);
276        assert!(deserialized.store.is_some());
277
278        let store = deserialized.store.unwrap();
279
280        assert!(store.contains_type(NodeType::Vertex));
281        assert!(store.contains_type(NodeType::Leaf));
282        assert!(deserialized.constraint.is_none());
283    }
284
285    #[test]
286    #[cfg(feature = "serde")]
287    fn test_serialize_deserialize_with_complex_type() {
288        use crate::Op;
289        let store = NodeStore::new();
290        store.insert(NodeType::Vertex, vec![Op::add(), Op::sub(), Op::mul()]);
291        store.insert(NodeType::Leaf, vec![Op::constant(1.0), Op::constant(2.0)]);
292
293        let root = TreeNode::with_children(
294            Op::add(),
295            vec![
296                TreeNode::with_children(
297                    Op::mul(),
298                    vec![
299                        TreeNode::new(Op::constant(1.0)),
300                        TreeNode::new(Op::constant(2.0)),
301                    ],
302                ),
303                TreeNode::new(Op::constant(3.0)),
304            ],
305        );
306
307        let chromosome = TreeChromosome::new(vec![root], Some(store), None);
308
309        let serialized = serde_json::to_string(&chromosome).unwrap();
310        let deserialized: TreeChromosome<Op<f32>> = serde_json::from_str(&serialized).unwrap();
311
312        assert_eq!(chromosome.nodes, deserialized.nodes);
313    }
314
315    #[test]
316    #[cfg(feature = "serde")]
317    fn test_serialize_deserialize_empty() {
318        let chromosome: TreeChromosome<i32> = TreeChromosome::new(vec![], None, None);
319
320        let serialized = serde_json::to_string(&chromosome).unwrap();
321        let deserialized: TreeChromosome<i32> = serde_json::from_str(&serialized).unwrap();
322
323        assert_eq!(chromosome, deserialized);
324    }
325
326    #[test]
327    #[cfg(feature = "serde")]
328    fn test_serialize_deserialize_with_constraint() {
329        let constraint = Arc::new(|node: &TreeNode<i32>| node.value() > &0);
330        let chromosome = TreeChromosome::new(vec![TreeNode::new(42)], None, Some(constraint));
331
332        let serialized = serde_json::to_string(&chromosome).unwrap();
333        let deserialized: TreeChromosome<i32> = serde_json::from_str(&serialized).unwrap();
334
335        assert_eq!(chromosome.nodes, deserialized.nodes);
336        assert!(deserialized.constraint.is_none());
337    }
338}