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