Skip to main content

ipfrs_transport/
graphsync.rs

1//! GraphSync protocol for DAG traversal
2//!
3//! Implements efficient DAG traversal with:
4//! - IPLD selector parsing and execution
5//! - Incremental response streaming
6//! - Resume capability from partial transfers
7//! - Breadth-first and depth-first traversal
8//!
9//! # Example
10//!
11//! ```
12//! use ipfrs_transport::Selector;
13//!
14//! // Create a selector for recursive depth-limited traversal
15//! let selector = Selector::RecursiveDepth { max_depth: 5 };
16//!
17//! // Validate the selector
18//! assert!(selector.validate().is_ok());
19//!
20//! // Create a selector for specific fields
21//! let field_selector = Selector::Fields {
22//!     fields: vec!["data".to_string(), "links".to_string()]
23//! };
24//!
25//! // Parse from JSON
26//! let json = r#"{"type": "recursivedepth", "max_depth": 3}"#;
27//! let parsed = Selector::from_json(json).unwrap();
28//! ```
29
30use 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/// IPLD Selector
39///
40/// Selectors specify which parts of a DAG to traverse
41#[derive(Debug, Clone, Serialize, Deserialize)]
42#[serde(tag = "type", rename_all = "lowercase")]
43#[derive(Default)]
44pub enum Selector {
45    /// Match everything
46    #[default]
47    All,
48    /// Match specific fields by name
49    Fields { fields: Vec<String> },
50    /// Recursively traverse to a depth limit
51    RecursiveDepth { max_depth: usize },
52    /// Recursively traverse all links
53    RecursiveAll,
54    /// Match based on index
55    Index { index: usize },
56    /// Sequence of selectors
57    Sequence { selectors: Vec<Selector> },
58    /// Match the current node
59    Matcher,
60}
61
62impl Selector {
63    /// Parse a selector from JSON
64    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    /// Validate the selector
70    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    /// Check if this selector matches all
89    pub fn matches_all(&self) -> bool {
90        matches!(self, Selector::All | Selector::RecursiveAll)
91    }
92}
93
94/// Traversal mode
95#[derive(Debug, Clone, Copy, PartialEq, Eq)]
96pub enum TraversalMode {
97    /// Breadth-first search
98    BreadthFirst,
99    /// Depth-first search
100    DepthFirst,
101}
102
103/// DAG traversal state
104#[derive(Debug, Clone)]
105pub struct TraversalState {
106    /// Root CID
107    pub root: Cid,
108    /// Visited CIDs
109    pub visited: HashSet<Cid>,
110    /// Queue of CIDs to visit (for BFS) or stack (for DFS)
111    pub queue: VecDeque<(Cid, usize)>, // (CID, depth)
112    /// Current depth
113    pub current_depth: usize,
114    /// Maximum depth (if limited)
115    pub max_depth: Option<usize>,
116    /// Blocks fetched so far
117    pub blocks_fetched: usize,
118    /// Bytes fetched so far
119    pub bytes_fetched: u64,
120}
121
122impl TraversalState {
123    /// Create a new traversal state
124    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    /// Check if traversal is complete
140    pub fn is_complete(&self) -> bool {
141        self.queue.is_empty()
142    }
143
144    /// Get next CID to visit
145    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    /// Add a CID to the queue
153    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    /// Mark a CID as visited
166    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    /// Save checkpoint for resume
173    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    /// Restore from checkpoint
185    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/// Checkpoint for resuming traversal
199#[derive(Debug, Clone)]
200pub struct TraversalCheckpoint {
201    /// Root CID
202    pub root: Cid,
203    /// Visited CIDs
204    pub visited: HashSet<Cid>,
205    /// Queue state
206    pub queue: VecDeque<(Cid, usize)>,
207    /// Maximum depth
208    pub max_depth: Option<usize>,
209    /// Blocks fetched
210    pub blocks_fetched: usize,
211    /// Bytes fetched
212    pub bytes_fetched: u64,
213}
214
215impl TraversalCheckpoint {
216    /// Serialize to JSON using CID strings
217    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    /// Deserialize from JSON
246    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
295/// DAG traversal engine
296pub struct DagTraversal<S: BlockStore> {
297    /// Block store
298    store: Arc<S>,
299    /// Traversal mode
300    mode: TraversalMode,
301    /// Selector
302    #[allow(dead_code)]
303    selector: Selector,
304    /// Traversal state
305    state: Arc<RwLock<TraversalState>>,
306}
307
308impl<S: BlockStore> DagTraversal<S> {
309    /// Create a new DAG traversal
310    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    /// Resume from a checkpoint
329    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    /// Get the next block in the traversal
347    pub async fn next_block(&self) -> Result<Option<Block>> {
348        let mut state = self.state.write().await;
349
350        // Get next CID to visit
351        let (cid, depth) = match state.next(self.mode) {
352            Some(item) => item,
353            None => return Ok(None),
354        };
355
356        // Fetch the block
357        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        // Mark as visited
363        state.mark_visited(cid, block.data().len() as u64);
364        state.current_depth = depth;
365
366        // Extract links from the block and add to queue
367        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    /// Extract CID links from a block
377    fn extract_links(&self, _block: &Block) -> Result<Vec<Cid>> {
378        // Simple link extraction - in a real implementation, this would parse IPLD
379        // and extract CIDs based on the selector
380
381        // For now, we'll just return an empty vector
382        // In a real implementation, you would:
383        // 1. Parse the block data as IPLD
384        // 2. Apply the selector to determine which fields to follow
385        // 3. Extract CID links from those fields
386
387        Ok(Vec::new())
388    }
389
390    /// Check if traversal is complete
391    pub async fn is_complete(&self) -> bool {
392        self.state.read().await.is_complete()
393    }
394
395    /// Get traversal statistics
396    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    /// Create a checkpoint for resume
407    pub async fn checkpoint(&self) -> TraversalCheckpoint {
408        self.state.read().await.checkpoint()
409    }
410
411    /// Traverse all and collect blocks
412    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/// Traversal statistics
424#[derive(Debug, Clone)]
425pub struct TraversalStats {
426    /// Number of blocks fetched
427    pub blocks_fetched: usize,
428    /// Bytes fetched
429    pub bytes_fetched: u64,
430    /// Blocks remaining in queue
431    pub blocks_remaining: usize,
432    /// Current traversal depth
433    pub current_depth: usize,
434}
435
436/// GraphSync protocol handler
437pub struct GraphSync<S: BlockStore> {
438    /// Block store
439    store: Arc<S>,
440    /// Active traversals
441    traversals: Arc<RwLock<HashMap<Cid, Arc<DagTraversal<S>>>>>,
442}
443
444impl<S: BlockStore> GraphSync<S> {
445    /// Create a new GraphSync instance
446    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    /// Start a new traversal
454    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    /// Resume a traversal from checkpoint
469    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    /// Get an active traversal
490    pub async fn get_traversal(&self, root: &Cid) -> Option<Arc<DagTraversal<S>>> {
491        self.traversals.read().await.get(root).cloned()
492    }
493
494    /// Remove a completed traversal
495    pub async fn remove_traversal(&self, root: &Cid) {
496        self.traversals.write().await.remove(root);
497    }
498
499    /// Get number of active traversals
500    pub async fn active_count(&self) -> usize {
501        self.traversals.read().await.len()
502    }
503}
504
505/// Gradient message for federated learning
506#[derive(Debug, Clone, Serialize, Deserialize)]
507pub struct GradientMessage {
508    /// Gradient identifier (e.g., layer name or tensor CID)
509    pub id: String,
510    /// Gradient data (compressed)
511    pub data: Vec<u8>,
512    /// Shape of the gradient tensor
513    pub shape: Vec<usize>,
514    /// Data type (f32, f16, etc.)
515    pub dtype: String,
516    /// Checksum for verification
517    pub checksum: u64,
518    /// Metadata (e.g., learning rate, batch size)
519    pub metadata: HashMap<String, String>,
520}
521
522impl GradientMessage {
523    /// Create a new gradient message
524    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    /// Compute checksum for data verification
542    fn compute_checksum(data: &[u8]) -> u64 {
543        // Simple checksum using FNV-1a hash
544        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    /// Verify checksum
553    pub fn verify(&self) -> bool {
554        Self::compute_checksum(&self.data) == self.checksum
555    }
556
557    /// Add metadata
558    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    /// Get total elements in gradient
564    pub fn num_elements(&self) -> usize {
565        self.shape.iter().product()
566    }
567
568    /// Estimate size in bytes
569    pub fn size_bytes(&self) -> usize {
570        self.data.len()
571    }
572}
573
574/// Gradient aggregation strategy
575#[derive(Debug, Clone, Copy, PartialEq, Eq)]
576pub enum AggregationStrategy {
577    /// Simple averaging
578    Average,
579    /// Weighted average based on sample counts
580    WeightedAverage,
581    /// Median aggregation (robust to outliers)
582    Median,
583    /// Federated averaging (FedAvg)
584    FederatedAvg,
585}
586
587/// Gradient aggregator for federated learning
588pub struct GradientAggregator {
589    /// Aggregation strategy
590    strategy: AggregationStrategy,
591    /// Accumulated gradients per layer
592    gradients: Arc<RwLock<HashMap<String, Vec<GradientMessage>>>>,
593    /// Expected number of contributors
594    expected_contributors: usize,
595    /// Verification enabled
596    verify_checksums: bool,
597}
598
599impl GradientAggregator {
600    /// Create a new gradient aggregator
601    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    /// Add a gradient to the aggregator
611    pub async fn add_gradient(&self, gradient: GradientMessage) -> Result<()> {
612        // Verify checksum if enabled
613        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        // Verify dimensions
621        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    /// Check if ready to aggregate (all contributors submitted)
637    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    /// Aggregate gradients for a specific layer
646    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        // Verify all gradients have same shape
657        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    /// Simple averaging aggregation
676    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        // Sum all gradients (treating as bytes for now)
685        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    /// Weighted average aggregation
701    fn aggregate_weighted(
702        &self,
703        layer_id: &str,
704        gradients: &[GradientMessage],
705    ) -> Result<GradientMessage> {
706        // Extract weights from metadata (sample counts)
707        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        // Weighted sum
721        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    /// Median aggregation (robust to outliers).
739    ///
740    /// For each element position, all contributor values are collected and the
741    /// median is selected.  We handle `f32` gradients (4 bytes/element) directly;
742    /// all other dtypes fall back to a byte-level median which is dtype-agnostic
743    /// but preserves the robust property.
744    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        // Validate that every gradient has the same byte length.
759        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        // Fast path: f32 gradients – operate at float granularity.
773        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                // Sort to find the median.
797                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                    // Even count: average of the two middle values.
802                    (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        // Generic byte-level median fallback for other dtypes (f16, i8, etc.).
817        // The median is taken independently per byte position.
818        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    /// Clear gradients for a layer after aggregation
840    pub async fn clear(&self, layer_id: &str) {
841        let mut gradients = self.gradients.write().await;
842        gradients.remove(layer_id);
843    }
844
845    /// Get statistics
846    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/// Gradient aggregator statistics
860#[derive(Debug, Clone)]
861pub struct GradientAggregatorStats {
862    /// Total gradients received
863    pub total_gradients: usize,
864    /// Number of layers
865    pub layers_count: usize,
866    /// Expected contributors
867    pub expected_contributors: usize,
868}
869
870/// Bidirectional gradient stream
871pub struct GradientStream {
872    /// Gradient aggregator
873    aggregator: Arc<GradientAggregator>,
874    /// Outgoing gradient queue
875    outgoing: Arc<RwLock<VecDeque<GradientMessage>>>,
876    /// Maximum queue size
877    max_queue_size: usize,
878}
879
880impl GradientStream {
881    /// Create a new gradient stream
882    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    /// Push a gradient to send
891    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    /// Pop a gradient to send
901    pub async fn pop_gradient(&self) -> Option<GradientMessage> {
902        self.outgoing.write().await.pop_front()
903    }
904
905    /// Receive a gradient
906    pub async fn receive_gradient(&self, gradient: GradientMessage) -> Result<()> {
907        self.aggregator.add_gradient(gradient).await
908    }
909
910    /// Get queue size
911    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        // Get root
954        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        // Verify original
1007        assert!(gradient.verify());
1008
1009        // Corrupt data
1010        gradient.data[0] = 99;
1011
1012        // Should fail verification
1013        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        // Push gradient
1050        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        // Pop gradient
1057        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        // Receive gradient
1065        stream
1066            .receive_gradient(grad)
1067            .await
1068            .expect("test: receive gradient into stream");
1069    }
1070}