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 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 let total = store.map(|values| values.len()).unwrap();
466 assert_eq!(total, 5);
467
468 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 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 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 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 assert!(store.map(|_| 42).is_none());
535
536 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 store.insert(NodeType::Input, vec![1, 2, 3]);
546
547 store.insert(NodeType::Input, vec![4, 5]);
549
550 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 let serialized = serde_json::to_string(&store).unwrap();
567
568 let deserialized: NodeStore<i32> = serde_json::from_str(&serialized).unwrap();
570
571 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 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 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 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 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 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}