1use ipfrs_core::error::{Error, Result};
31use ipfrs_core::{Block, Cid};
32use ipfrs_storage::traits::BlockStore;
33use serde::{Deserialize, Serialize};
34use std::collections::{HashMap, HashSet, VecDeque};
35use std::sync::Arc;
36use tokio::sync::RwLock;
37
38#[derive(Debug, Clone, Serialize, Deserialize)]
42#[serde(tag = "type", rename_all = "lowercase")]
43#[derive(Default)]
44pub enum Selector {
45 #[default]
47 All,
48 Fields { fields: Vec<String> },
50 RecursiveDepth { max_depth: usize },
52 RecursiveAll,
54 Index { index: usize },
56 Sequence { selectors: Vec<Selector> },
58 Matcher,
60}
61
62impl Selector {
63 pub fn from_json(json: &str) -> Result<Self> {
65 serde_json::from_str(json)
66 .map_err(|e| Error::InvalidInput(format!("Failed to parse selector: {}", e)))
67 }
68
69 pub fn validate(&self) -> Result<()> {
71 match self {
72 Selector::RecursiveDepth { max_depth } if *max_depth == 0 => {
73 return Err(Error::InvalidInput(
74 "max_depth must be greater than 0".to_string(),
75 ));
76 }
77 Selector::RecursiveDepth { .. } => {}
78 Selector::Sequence { selectors } => {
79 for sel in selectors {
80 sel.validate()?;
81 }
82 }
83 _ => {}
84 }
85 Ok(())
86 }
87
88 pub fn matches_all(&self) -> bool {
90 matches!(self, Selector::All | Selector::RecursiveAll)
91 }
92}
93
94#[derive(Debug, Clone, Copy, PartialEq, Eq)]
96pub enum TraversalMode {
97 BreadthFirst,
99 DepthFirst,
101}
102
103#[derive(Debug, Clone)]
105pub struct TraversalState {
106 pub root: Cid,
108 pub visited: HashSet<Cid>,
110 pub queue: VecDeque<(Cid, usize)>, pub current_depth: usize,
114 pub max_depth: Option<usize>,
116 pub blocks_fetched: usize,
118 pub bytes_fetched: u64,
120}
121
122impl TraversalState {
123 pub fn new(root: Cid, max_depth: Option<usize>) -> Self {
125 let mut queue = VecDeque::new();
126 queue.push_back((root, 0));
127
128 Self {
129 root,
130 visited: HashSet::new(),
131 queue,
132 current_depth: 0,
133 max_depth,
134 blocks_fetched: 0,
135 bytes_fetched: 0,
136 }
137 }
138
139 pub fn is_complete(&self) -> bool {
141 self.queue.is_empty()
142 }
143
144 pub fn next(&mut self, mode: TraversalMode) -> Option<(Cid, usize)> {
146 match mode {
147 TraversalMode::BreadthFirst => self.queue.pop_front(),
148 TraversalMode::DepthFirst => self.queue.pop_back(),
149 }
150 }
151
152 pub fn enqueue(&mut self, cid: Cid, depth: usize) {
154 if let Some(max) = self.max_depth {
155 if depth > max {
156 return;
157 }
158 }
159
160 if !self.visited.contains(&cid) {
161 self.queue.push_back((cid, depth));
162 }
163 }
164
165 pub fn mark_visited(&mut self, cid: Cid, size: u64) {
167 self.visited.insert(cid);
168 self.blocks_fetched += 1;
169 self.bytes_fetched += size;
170 }
171
172 pub fn checkpoint(&self) -> TraversalCheckpoint {
174 TraversalCheckpoint {
175 root: self.root,
176 visited: self.visited.clone(),
177 queue: self.queue.clone(),
178 max_depth: self.max_depth,
179 blocks_fetched: self.blocks_fetched,
180 bytes_fetched: self.bytes_fetched,
181 }
182 }
183
184 pub fn from_checkpoint(checkpoint: TraversalCheckpoint) -> Self {
186 Self {
187 root: checkpoint.root,
188 visited: checkpoint.visited,
189 queue: checkpoint.queue,
190 current_depth: 0,
191 max_depth: checkpoint.max_depth,
192 blocks_fetched: checkpoint.blocks_fetched,
193 bytes_fetched: checkpoint.bytes_fetched,
194 }
195 }
196}
197
198#[derive(Debug, Clone)]
200pub struct TraversalCheckpoint {
201 pub root: Cid,
203 pub visited: HashSet<Cid>,
205 pub queue: VecDeque<(Cid, usize)>,
207 pub max_depth: Option<usize>,
209 pub blocks_fetched: usize,
211 pub bytes_fetched: u64,
213}
214
215impl TraversalCheckpoint {
216 pub fn to_json(&self) -> Result<String> {
218 #[derive(Serialize)]
219 struct SerializableCheckpoint {
220 root: String,
221 visited: Vec<String>,
222 queue: Vec<(String, usize)>,
223 max_depth: Option<usize>,
224 blocks_fetched: usize,
225 bytes_fetched: u64,
226 }
227
228 let serializable = SerializableCheckpoint {
229 root: self.root.to_string(),
230 visited: self.visited.iter().map(|c| c.to_string()).collect(),
231 queue: self
232 .queue
233 .iter()
234 .map(|(c, d)| (c.to_string(), *d))
235 .collect(),
236 max_depth: self.max_depth,
237 blocks_fetched: self.blocks_fetched,
238 bytes_fetched: self.bytes_fetched,
239 };
240
241 serde_json::to_string(&serializable)
242 .map_err(|e| Error::Internal(format!("Failed to serialize checkpoint: {}", e)))
243 }
244
245 pub fn from_json(json: &str) -> Result<Self> {
247 #[derive(Deserialize)]
248 struct SerializableCheckpoint {
249 root: String,
250 visited: Vec<String>,
251 queue: Vec<(String, usize)>,
252 max_depth: Option<usize>,
253 blocks_fetched: usize,
254 bytes_fetched: u64,
255 }
256
257 let serializable: SerializableCheckpoint = serde_json::from_str(json)
258 .map_err(|e| Error::Internal(format!("Failed to deserialize checkpoint: {}", e)))?;
259
260 let root: Cid = serializable
261 .root
262 .parse()
263 .map_err(|e| Error::InvalidInput(format!("Invalid root CID: {}", e)))?;
264
265 let visited: Result<HashSet<Cid>> = serializable
266 .visited
267 .iter()
268 .map(|s| {
269 s.parse()
270 .map_err(|e| Error::InvalidInput(format!("Invalid CID: {}", e)))
271 })
272 .collect();
273
274 let queue: Result<VecDeque<(Cid, usize)>> = serializable
275 .queue
276 .iter()
277 .map(|(s, d)| {
278 s.parse()
279 .map(|c| (c, *d))
280 .map_err(|e| Error::InvalidInput(format!("Invalid CID: {}", e)))
281 })
282 .collect();
283
284 Ok(Self {
285 root,
286 visited: visited?,
287 queue: queue?,
288 max_depth: serializable.max_depth,
289 blocks_fetched: serializable.blocks_fetched,
290 bytes_fetched: serializable.bytes_fetched,
291 })
292 }
293}
294
295pub struct DagTraversal<S: BlockStore> {
297 store: Arc<S>,
299 mode: TraversalMode,
301 #[allow(dead_code)]
303 selector: Selector,
304 state: Arc<RwLock<TraversalState>>,
306}
307
308impl<S: BlockStore> DagTraversal<S> {
309 pub fn new(store: Arc<S>, root: Cid, selector: Selector, mode: TraversalMode) -> Result<Self> {
311 selector.validate()?;
312
313 let max_depth = match &selector {
314 Selector::RecursiveDepth { max_depth } => Some(*max_depth),
315 _ => None,
316 };
317
318 let state = TraversalState::new(root, max_depth);
319
320 Ok(Self {
321 store,
322 mode,
323 selector,
324 state: Arc::new(RwLock::new(state)),
325 })
326 }
327
328 pub fn from_checkpoint(
330 store: Arc<S>,
331 checkpoint: TraversalCheckpoint,
332 selector: Selector,
333 mode: TraversalMode,
334 ) -> Result<Self> {
335 selector.validate()?;
336 let state = TraversalState::from_checkpoint(checkpoint);
337
338 Ok(Self {
339 store,
340 mode,
341 selector,
342 state: Arc::new(RwLock::new(state)),
343 })
344 }
345
346 pub async fn next_block(&self) -> Result<Option<Block>> {
348 let mut state = self.state.write().await;
349
350 let (cid, depth) = match state.next(self.mode) {
352 Some(item) => item,
353 None => return Ok(None),
354 };
355
356 let block = match self.store.get(&cid).await? {
358 Some(b) => b,
359 None => return Err(Error::NotFound(format!("Block not found for CID: {}", cid))),
360 };
361
362 state.mark_visited(cid, block.data().len() as u64);
364 state.current_depth = depth;
365
366 if let Ok(links) = self.extract_links(&block) {
368 for link_cid in links {
369 state.enqueue(link_cid, depth + 1);
370 }
371 }
372
373 Ok(Some(block))
374 }
375
376 fn extract_links(&self, _block: &Block) -> Result<Vec<Cid>> {
378 Ok(Vec::new())
388 }
389
390 pub async fn is_complete(&self) -> bool {
392 self.state.read().await.is_complete()
393 }
394
395 pub async fn stats(&self) -> TraversalStats {
397 let state = self.state.read().await;
398 TraversalStats {
399 blocks_fetched: state.blocks_fetched,
400 bytes_fetched: state.bytes_fetched,
401 blocks_remaining: state.queue.len(),
402 current_depth: state.current_depth,
403 }
404 }
405
406 pub async fn checkpoint(&self) -> TraversalCheckpoint {
408 self.state.read().await.checkpoint()
409 }
410
411 pub async fn collect_all(&self) -> Result<Vec<Block>> {
413 let mut blocks = Vec::new();
414
415 while let Some(block) = self.next_block().await? {
416 blocks.push(block);
417 }
418
419 Ok(blocks)
420 }
421}
422
423#[derive(Debug, Clone)]
425pub struct TraversalStats {
426 pub blocks_fetched: usize,
428 pub bytes_fetched: u64,
430 pub blocks_remaining: usize,
432 pub current_depth: usize,
434}
435
436pub struct GraphSync<S: BlockStore> {
438 store: Arc<S>,
440 traversals: Arc<RwLock<HashMap<Cid, Arc<DagTraversal<S>>>>>,
442}
443
444impl<S: BlockStore> GraphSync<S> {
445 pub fn new(store: Arc<S>) -> Result<Self> {
447 Ok(Self {
448 store,
449 traversals: Arc::new(RwLock::new(HashMap::new())),
450 })
451 }
452
453 pub async fn start_traversal(
455 &self,
456 root: Cid,
457 selector: Selector,
458 mode: TraversalMode,
459 ) -> Result<Arc<DagTraversal<S>>> {
460 let traversal = Arc::new(DagTraversal::new(self.store.clone(), root, selector, mode)?);
461
462 let mut traversals = self.traversals.write().await;
463 traversals.insert(root, traversal.clone());
464
465 Ok(traversal)
466 }
467
468 pub async fn resume_traversal(
470 &self,
471 checkpoint: TraversalCheckpoint,
472 selector: Selector,
473 mode: TraversalMode,
474 ) -> Result<Arc<DagTraversal<S>>> {
475 let root = checkpoint.root;
476 let traversal = Arc::new(DagTraversal::from_checkpoint(
477 self.store.clone(),
478 checkpoint,
479 selector,
480 mode,
481 )?);
482
483 let mut traversals = self.traversals.write().await;
484 traversals.insert(root, traversal.clone());
485
486 Ok(traversal)
487 }
488
489 pub async fn get_traversal(&self, root: &Cid) -> Option<Arc<DagTraversal<S>>> {
491 self.traversals.read().await.get(root).cloned()
492 }
493
494 pub async fn remove_traversal(&self, root: &Cid) {
496 self.traversals.write().await.remove(root);
497 }
498
499 pub async fn active_count(&self) -> usize {
501 self.traversals.read().await.len()
502 }
503}
504
505#[derive(Debug, Clone, Serialize, Deserialize)]
507pub struct GradientMessage {
508 pub id: String,
510 pub data: Vec<u8>,
512 pub shape: Vec<usize>,
514 pub dtype: String,
516 pub checksum: u64,
518 pub metadata: HashMap<String, String>,
520}
521
522impl GradientMessage {
523 pub fn new(
525 id: impl Into<String>,
526 data: Vec<u8>,
527 shape: Vec<usize>,
528 dtype: impl Into<String>,
529 ) -> Self {
530 let checksum = Self::compute_checksum(&data);
531 Self {
532 id: id.into(),
533 data,
534 shape,
535 dtype: dtype.into(),
536 checksum,
537 metadata: HashMap::new(),
538 }
539 }
540
541 fn compute_checksum(data: &[u8]) -> u64 {
543 let mut hash: u64 = 0xcbf29ce484222325;
545 for &byte in data {
546 hash ^= byte as u64;
547 hash = hash.wrapping_mul(0x100000001b3);
548 }
549 hash
550 }
551
552 pub fn verify(&self) -> bool {
554 Self::compute_checksum(&self.data) == self.checksum
555 }
556
557 pub fn with_metadata(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
559 self.metadata.insert(key.into(), value.into());
560 self
561 }
562
563 pub fn num_elements(&self) -> usize {
565 self.shape.iter().product()
566 }
567
568 pub fn size_bytes(&self) -> usize {
570 self.data.len()
571 }
572}
573
574#[derive(Debug, Clone, Copy, PartialEq, Eq)]
576pub enum AggregationStrategy {
577 Average,
579 WeightedAverage,
581 Median,
583 FederatedAvg,
585}
586
587pub struct GradientAggregator {
589 strategy: AggregationStrategy,
591 gradients: Arc<RwLock<HashMap<String, Vec<GradientMessage>>>>,
593 expected_contributors: usize,
595 verify_checksums: bool,
597}
598
599impl GradientAggregator {
600 pub fn new(strategy: AggregationStrategy, expected_contributors: usize) -> Self {
602 Self {
603 strategy,
604 gradients: Arc::new(RwLock::new(HashMap::new())),
605 expected_contributors,
606 verify_checksums: true,
607 }
608 }
609
610 pub async fn add_gradient(&self, gradient: GradientMessage) -> Result<()> {
612 if self.verify_checksums && !gradient.verify() {
614 return Err(Error::InvalidInput(format!(
615 "Gradient checksum verification failed for {}",
616 gradient.id
617 )));
618 }
619
620 if gradient.num_elements() == 0 {
622 return Err(Error::InvalidInput(
623 "Gradient has zero elements".to_string(),
624 ));
625 }
626
627 let mut gradients = self.gradients.write().await;
628 gradients
629 .entry(gradient.id.clone())
630 .or_insert_with(Vec::new)
631 .push(gradient);
632
633 Ok(())
634 }
635
636 pub async fn is_ready(&self, layer_id: &str) -> bool {
638 let gradients = self.gradients.read().await;
639 gradients
640 .get(layer_id)
641 .map(|g| g.len() >= self.expected_contributors)
642 .unwrap_or(false)
643 }
644
645 pub async fn aggregate(&self, layer_id: &str) -> Result<GradientMessage> {
647 let gradients = self.gradients.read().await;
648 let layer_gradients = gradients
649 .get(layer_id)
650 .ok_or_else(|| Error::NotFound(format!("No gradients for layer: {}", layer_id)))?;
651
652 if layer_gradients.is_empty() {
653 return Err(Error::InvalidInput("No gradients to aggregate".to_string()));
654 }
655
656 let first_shape = &layer_gradients[0].shape;
658 for grad in layer_gradients.iter().skip(1) {
659 if &grad.shape != first_shape {
660 return Err(Error::InvalidInput("Gradient shape mismatch".to_string()));
661 }
662 }
663
664 match self.strategy {
665 AggregationStrategy::Average | AggregationStrategy::FederatedAvg => {
666 self.aggregate_average(layer_id, layer_gradients)
667 }
668 AggregationStrategy::WeightedAverage => {
669 self.aggregate_weighted(layer_id, layer_gradients)
670 }
671 AggregationStrategy::Median => self.aggregate_median(layer_id, layer_gradients),
672 }
673 }
674
675 fn aggregate_average(
677 &self,
678 layer_id: &str,
679 gradients: &[GradientMessage],
680 ) -> Result<GradientMessage> {
681 let n = gradients.len();
682 let size = gradients[0].data.len();
683
684 let mut sum = vec![0u8; size];
686 for grad in gradients {
687 for (i, &byte) in grad.data.iter().enumerate() {
688 sum[i] = sum[i].saturating_add(byte / n as u8);
689 }
690 }
691
692 Ok(GradientMessage::new(
693 layer_id,
694 sum,
695 gradients[0].shape.clone(),
696 gradients[0].dtype.clone(),
697 ))
698 }
699
700 fn aggregate_weighted(
702 &self,
703 layer_id: &str,
704 gradients: &[GradientMessage],
705 ) -> Result<GradientMessage> {
706 let weights: Vec<f32> = gradients
708 .iter()
709 .map(|g| {
710 g.metadata
711 .get("samples")
712 .and_then(|s| s.parse::<f32>().ok())
713 .unwrap_or(1.0)
714 })
715 .collect();
716
717 let total_weight: f32 = weights.iter().sum();
718 let size = gradients[0].data.len();
719
720 let mut weighted_sum = vec![0u8; size];
722 for (grad, &weight) in gradients.iter().zip(weights.iter()) {
723 let normalized_weight = weight / total_weight;
724 for (i, &byte) in grad.data.iter().enumerate() {
725 weighted_sum[i] =
726 weighted_sum[i].saturating_add((byte as f32 * normalized_weight) as u8);
727 }
728 }
729
730 Ok(GradientMessage::new(
731 layer_id,
732 weighted_sum,
733 gradients[0].shape.clone(),
734 gradients[0].dtype.clone(),
735 ))
736 }
737
738 fn aggregate_median(
745 &self,
746 layer_id: &str,
747 gradients: &[GradientMessage],
748 ) -> Result<GradientMessage> {
749 if gradients.is_empty() {
750 return Err(Error::InvalidInput(
751 "median aggregation: no gradients supplied".to_string(),
752 ));
753 }
754
755 let n = gradients.len();
756 let size = gradients[0].data.len();
757
758 for (i, g) in gradients.iter().enumerate() {
760 if g.data.len() != size {
761 return Err(Error::InvalidInput(format!(
762 "median aggregation: gradient {} has {} bytes, expected {}",
763 i,
764 g.data.len(),
765 size,
766 )));
767 }
768 }
769
770 let dtype = &gradients[0].dtype;
771
772 if dtype == "f32" || dtype == "float32" {
774 let element_size = 4usize;
775 if !size.is_multiple_of(element_size) {
776 return Err(Error::InvalidInput(format!(
777 "median aggregation: byte length {} is not a multiple of 4 for f32 dtype",
778 size
779 )));
780 }
781 let num_elements = size / element_size;
782 let mut out_data = vec![0u8; size];
783
784 for elem in 0..num_elements {
785 let byte_off = elem * element_size;
786 let mut values: Vec<f32> = gradients
787 .iter()
788 .map(|g| {
789 let bytes: [u8; 4] = g.data[byte_off..byte_off + element_size]
790 .try_into()
791 .unwrap_or([0u8; 4]);
792 f32::from_le_bytes(bytes)
793 })
794 .collect();
795
796 values.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
798 let median = if n % 2 == 1 {
799 values[n / 2]
800 } else {
801 (values[n / 2 - 1] + values[n / 2]) * 0.5
803 };
804
805 out_data[byte_off..byte_off + element_size].copy_from_slice(&median.to_le_bytes());
806 }
807
808 return Ok(GradientMessage::new(
809 layer_id,
810 out_data,
811 gradients[0].shape.clone(),
812 dtype.clone(),
813 ));
814 }
815
816 let mut out_data = vec![0u8; size];
819 for (byte_pos, out_byte) in out_data.iter_mut().enumerate() {
820 let mut col: Vec<u8> = gradients.iter().map(|g| g.data[byte_pos]).collect();
821 col.sort_unstable();
822 *out_byte = if n % 2 == 1 {
823 col[n / 2]
824 } else {
825 let lo = col[n / 2 - 1] as u16;
826 let hi = col[n / 2] as u16;
827 ((lo + hi) / 2) as u8
828 };
829 }
830
831 Ok(GradientMessage::new(
832 layer_id,
833 out_data,
834 gradients[0].shape.clone(),
835 dtype.clone(),
836 ))
837 }
838
839 pub async fn clear(&self, layer_id: &str) {
841 let mut gradients = self.gradients.write().await;
842 gradients.remove(layer_id);
843 }
844
845 pub async fn stats(&self) -> GradientAggregatorStats {
847 let gradients = self.gradients.read().await;
848 let total_gradients: usize = gradients.values().map(|v| v.len()).sum();
849 let layers_count = gradients.len();
850
851 GradientAggregatorStats {
852 total_gradients,
853 layers_count,
854 expected_contributors: self.expected_contributors,
855 }
856 }
857}
858
859#[derive(Debug, Clone)]
861pub struct GradientAggregatorStats {
862 pub total_gradients: usize,
864 pub layers_count: usize,
866 pub expected_contributors: usize,
868}
869
870pub struct GradientStream {
872 aggregator: Arc<GradientAggregator>,
874 outgoing: Arc<RwLock<VecDeque<GradientMessage>>>,
876 max_queue_size: usize,
878}
879
880impl GradientStream {
881 pub fn new(aggregator: Arc<GradientAggregator>, max_queue_size: usize) -> Self {
883 Self {
884 aggregator,
885 outgoing: Arc::new(RwLock::new(VecDeque::new())),
886 max_queue_size,
887 }
888 }
889
890 pub async fn push_gradient(&self, gradient: GradientMessage) -> Result<()> {
892 let mut outgoing = self.outgoing.write().await;
893 if outgoing.len() >= self.max_queue_size {
894 return Err(Error::Internal("Gradient queue is full".to_string()));
895 }
896 outgoing.push_back(gradient);
897 Ok(())
898 }
899
900 pub async fn pop_gradient(&self) -> Option<GradientMessage> {
902 self.outgoing.write().await.pop_front()
903 }
904
905 pub async fn receive_gradient(&self, gradient: GradientMessage) -> Result<()> {
907 self.aggregator.add_gradient(gradient).await
908 }
909
910 pub async fn queue_size(&self) -> usize {
912 self.outgoing.read().await.len()
913 }
914}
915
916#[cfg(test)]
917mod tests {
918 use super::*;
919
920 #[test]
921 fn test_selector_parse() {
922 let json = r#"{"type":"all"}"#;
923 let selector = Selector::from_json(json).expect("test: parse all-selector from JSON");
924 assert!(selector.matches_all());
925
926 let json2 = r#"{"type":"recursivedepth","max_depth":5}"#;
927 let selector2 =
928 Selector::from_json(json2).expect("test: parse recursive-depth selector from JSON");
929 match selector2 {
930 Selector::RecursiveDepth { max_depth } => assert_eq!(max_depth, 5),
931 _ => panic!("Wrong selector type"),
932 }
933 }
934
935 #[test]
936 fn test_selector_validate() {
937 let selector = Selector::RecursiveDepth { max_depth: 0 };
938 assert!(selector.validate().is_err());
939
940 let selector2 = Selector::RecursiveDepth { max_depth: 5 };
941 assert!(selector2.validate().is_ok());
942 }
943
944 #[test]
945 fn test_traversal_state() {
946 let cid: Cid = "bafybeigdyrzt5sfp7udm7hu76uh7y26nf3efuylqabf3oclgtqy55fbzdi"
947 .parse()
948 .expect("test: parse CID string");
949
950 let mut state = TraversalState::new(cid, Some(3));
951 assert!(!state.is_complete());
952
953 let (root_cid, depth) = state
955 .next(TraversalMode::BreadthFirst)
956 .expect("test: get next traversal item");
957 assert_eq!(root_cid, cid);
958 assert_eq!(depth, 0);
959
960 state.mark_visited(cid, 1024);
961 assert_eq!(state.blocks_fetched, 1);
962 assert_eq!(state.bytes_fetched, 1024);
963
964 assert!(state.is_complete());
965 }
966
967 #[test]
968 fn test_checkpoint() {
969 let cid: Cid = "bafybeigdyrzt5sfp7udm7hu76uh7y26nf3efuylqabf3oclgtqy55fbzdi"
970 .parse()
971 .expect("test: parse CID string");
972
973 let mut state = TraversalState::new(cid, Some(3));
974 state.mark_visited(cid, 1024);
975
976 let checkpoint = state.checkpoint();
977 assert_eq!(checkpoint.root, cid);
978 assert_eq!(checkpoint.blocks_fetched, 1);
979 assert_eq!(checkpoint.bytes_fetched, 1024);
980
981 let restored = TraversalState::from_checkpoint(checkpoint);
982 assert_eq!(restored.blocks_fetched, 1);
983 assert_eq!(restored.bytes_fetched, 1024);
984 }
985
986 #[test]
987 fn test_gradient_message() {
988 let data = vec![1, 2, 3, 4, 5];
989 let shape = vec![5];
990 let gradient = GradientMessage::new("layer1", data.clone(), shape.clone(), "f32");
991
992 assert_eq!(gradient.id, "layer1");
993 assert_eq!(gradient.data, data);
994 assert_eq!(gradient.shape, shape);
995 assert_eq!(gradient.dtype, "f32");
996 assert!(gradient.verify());
997 assert_eq!(gradient.num_elements(), 5);
998 assert_eq!(gradient.size_bytes(), 5);
999 }
1000
1001 #[test]
1002 fn test_gradient_checksum() {
1003 let data = vec![1, 2, 3, 4, 5];
1004 let mut gradient = GradientMessage::new("layer1", data, vec![5], "f32");
1005
1006 assert!(gradient.verify());
1008
1009 gradient.data[0] = 99;
1011
1012 assert!(!gradient.verify());
1014 }
1015
1016 #[tokio::test]
1017 async fn test_gradient_aggregator() {
1018 let aggregator = GradientAggregator::new(AggregationStrategy::Average, 2);
1019
1020 let grad1 = GradientMessage::new("layer1", vec![10, 20, 30], vec![3], "f32");
1021 let grad2 = GradientMessage::new("layer1", vec![20, 30, 40], vec![3], "f32");
1022
1023 aggregator
1024 .add_gradient(grad1)
1025 .await
1026 .expect("test: add first gradient to aggregator");
1027 aggregator
1028 .add_gradient(grad2)
1029 .await
1030 .expect("test: add second gradient to aggregator");
1031
1032 assert!(aggregator.is_ready("layer1").await);
1033
1034 let aggregated = aggregator
1035 .aggregate("layer1")
1036 .await
1037 .expect("test: aggregate layer1 gradients");
1038 assert_eq!(aggregated.shape, vec![3]);
1039 assert_eq!(aggregated.id, "layer1");
1040 }
1041
1042 #[tokio::test]
1043 async fn test_gradient_stream() {
1044 let aggregator = Arc::new(GradientAggregator::new(AggregationStrategy::Average, 1));
1045 let stream = GradientStream::new(aggregator, 10);
1046
1047 let grad = GradientMessage::new("layer1", vec![1, 2, 3], vec![3], "f32");
1048
1049 stream
1051 .push_gradient(grad.clone())
1052 .await
1053 .expect("test: push gradient to stream");
1054 assert_eq!(stream.queue_size().await, 1);
1055
1056 let popped = stream
1058 .pop_gradient()
1059 .await
1060 .expect("test: pop gradient from stream");
1061 assert_eq!(popped.id, "layer1");
1062 assert_eq!(stream.queue_size().await, 0);
1063
1064 stream
1066 .receive_gradient(grad)
1067 .await
1068 .expect("test: receive gradient into stream");
1069 }
1070}