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