Skip to main content

nodedb_cluster/distributed_array/coordinator/
mod.rs

1// SPDX-License-Identifier: BUSL-1.1
2
3pub mod read;
4pub mod write;
5
6pub use read::{ArrayCoordParams, ArrayCoordinator, CoordAggResult, CoordSliceResult};
7pub use write::{ArrayWriteCoordParams, coord_delete, coord_put, coord_put_partitioned};
8
9#[cfg(test)]
10mod tests {
11    use std::sync::Arc;
12
13    use async_trait::async_trait;
14
15    use crate::circuit_breaker::{CircuitBreaker, CircuitBreakerConfig};
16    use crate::error::Result;
17    use crate::wire::{VShardEnvelope, VShardMessageType};
18
19    use super::super::merge::ArrayAggPartial;
20    use super::super::rpc::ShardRpcDispatch;
21    use super::super::wire::{
22        ArrayShardAggReq, ArrayShardAggResp, ArrayShardSliceReq, ArrayShardSliceResp,
23    };
24    use super::read::{ArrayCoordParams, ArrayCoordinator};
25
26    /// Mock dispatch that returns a pre-serialised `ArrayShardSliceResp`.
27    struct SliceEchoDispatch {
28        /// Rows to return from each shard.
29        rows: Vec<Vec<u8>>,
30    }
31
32    #[async_trait]
33    impl ShardRpcDispatch for SliceEchoDispatch {
34        async fn call(&self, req: VShardEnvelope, _timeout_ms: u64) -> Result<VShardEnvelope> {
35            let resp = ArrayShardSliceResp {
36                shard_id: req.vshard_id,
37                rows_msgpack: self.rows.clone(),
38                truncated: false,
39                truncated_before_horizon: false,
40            };
41            let payload = zerompk::to_msgpack_vec(&resp).unwrap();
42            Ok(VShardEnvelope::new(
43                VShardMessageType::ArrayShardSliceResp,
44                req.target_node,
45                req.source_node,
46                req.vshard_id,
47                payload,
48            ))
49        }
50    }
51
52    /// Mock dispatch that returns a pre-canned `ArrayShardAggResp`.
53    struct AggEchoDispatch {
54        partials: Vec<ArrayAggPartial>,
55    }
56
57    #[async_trait]
58    impl ShardRpcDispatch for AggEchoDispatch {
59        async fn call(&self, req: VShardEnvelope, _timeout_ms: u64) -> Result<VShardEnvelope> {
60            let resp = ArrayShardAggResp {
61                shard_id: req.vshard_id,
62                partials: self.partials.clone(),
63                truncated_before_horizon: false,
64            };
65            let payload = zerompk::to_msgpack_vec(&resp).unwrap();
66            Ok(VShardEnvelope::new(
67                VShardMessageType::ArrayShardSliceResp,
68                req.target_node,
69                req.source_node,
70                req.vshard_id,
71                payload,
72            ))
73        }
74    }
75
76    fn make_coordinator(
77        shard_ids: Vec<u32>,
78        dispatch: Arc<dyn ShardRpcDispatch>,
79    ) -> ArrayCoordinator {
80        ArrayCoordinator::new(
81            ArrayCoordParams {
82                source_node: 1,
83                shard_ids,
84                timeout_ms: 1000,
85                // Tests use prefix_bits=0 so shard-side routing validation
86                // is skipped — mock executors don't need to match Hilbert
87                // ownership.
88                prefix_bits: 0,
89                slice_hilbert_ranges: vec![],
90            },
91            dispatch,
92            Arc::new(CircuitBreaker::new(CircuitBreakerConfig::default())),
93        )
94    }
95
96    #[tokio::test]
97    async fn coord_slice_merges_rows_from_all_shards() {
98        let row_a = zerompk::to_msgpack_vec(&"row-a").unwrap();
99        let row_b = zerompk::to_msgpack_vec(&"row-b").unwrap();
100        let dispatch: Arc<dyn ShardRpcDispatch> = Arc::new(SliceEchoDispatch {
101            rows: vec![row_a.clone(), row_b.clone()],
102        });
103        let coord = make_coordinator(vec![0, 1, 2], dispatch);
104        let req = ArrayShardSliceReq {
105            array_id_msgpack: vec![],
106            slice_msgpack: vec![],
107            attr_projection: vec![],
108            limit: 100,
109            cell_filter_msgpack: vec![],
110            prefix_bits: 0,
111            slice_hilbert_ranges: vec![],
112            shard_hilbert_range: None,
113            system_time: nodedb_types::SystemTimeScope::Current,
114            valid_at_ms: None,
115        };
116
117        // 3 shards × 2 rows each = 6 merged rows.
118        let result = coord
119            .coord_slice(req, 0, nodedb_types::SystemTimeScope::Current)
120            .await
121            .expect("coord_slice should succeed");
122        assert_eq!(result.rows.len(), 6);
123        assert!(!result.truncated_before_horizon);
124    }
125
126    #[tokio::test]
127    async fn coord_slice_applies_coordinator_limit() {
128        let row = zerompk::to_msgpack_vec(&"row").unwrap();
129        let dispatch: Arc<dyn ShardRpcDispatch> = Arc::new(SliceEchoDispatch {
130            rows: vec![row.clone(), row.clone(), row.clone()],
131        });
132        // 2 shards × 3 rows = 6 total, but limit = 4.
133        let coord = make_coordinator(vec![0, 1], dispatch);
134        let req = ArrayShardSliceReq {
135            array_id_msgpack: vec![],
136            slice_msgpack: vec![],
137            attr_projection: vec![],
138            limit: 3,
139            cell_filter_msgpack: vec![],
140            prefix_bits: 0,
141            slice_hilbert_ranges: vec![],
142            shard_hilbert_range: None,
143            system_time: nodedb_types::SystemTimeScope::Current,
144            valid_at_ms: None,
145        };
146
147        let result = coord
148            .coord_slice(req, 4, nodedb_types::SystemTimeScope::Current)
149            .await
150            .expect("coord_slice with limit should succeed");
151        assert_eq!(result.rows.len(), 4);
152    }
153
154    fn make_agg_req() -> ArrayShardAggReq {
155        // Sum reducer c_enum = 0.
156        ArrayShardAggReq {
157            array_id_msgpack: vec![],
158            attr_idx: 0,
159            reducer_msgpack: vec![0x00],
160            group_by_dim: -1,
161            cell_filter_msgpack: vec![],
162            shard_hilbert_range: None,
163            system_as_of: None,
164            valid_at_ms: None,
165        }
166    }
167
168    #[tokio::test]
169    async fn coord_agg_merges_scalar_partials_from_shards() {
170        let dispatch: Arc<dyn ShardRpcDispatch> = Arc::new(AggEchoDispatch {
171            partials: vec![ArrayAggPartial::from_single(0, 10.0)],
172        });
173        // 3 shards each returning a partial with sum=10 → merged sum=30.
174        let coord = make_coordinator(vec![0, 1, 2], dispatch);
175        let merged = coord
176            .coord_agg(make_agg_req())
177            .await
178            .expect("coord_agg should succeed");
179
180        assert_eq!(merged.partials.len(), 1);
181        assert_eq!(merged.partials[0].count, 3);
182        assert!((merged.partials[0].sum - 30.0).abs() < f64::EPSILON);
183        assert!(!merged.truncated_before_horizon);
184    }
185
186    #[tokio::test]
187    async fn coord_agg_with_empty_shards_returns_empty() {
188        let dispatch: Arc<dyn ShardRpcDispatch> = Arc::new(AggEchoDispatch { partials: vec![] });
189        let coord = make_coordinator(vec![0, 1], dispatch);
190        let merged = coord
191            .coord_agg(make_agg_req())
192            .await
193            .expect("coord_agg with empty shards should succeed");
194        assert!(merged.partials.is_empty());
195    }
196
197    #[tokio::test]
198    async fn coord_agg_merges_grouped_partials_across_shards() {
199        // Shard 0 returns group_key=0 partial, shard 1 also group_key=0 + group_key=1.
200        struct GroupedDispatch {
201            shard0_partials: Vec<ArrayAggPartial>,
202            shard1_partials: Vec<ArrayAggPartial>,
203        }
204
205        #[async_trait]
206        impl ShardRpcDispatch for GroupedDispatch {
207            async fn call(&self, req: VShardEnvelope, _timeout_ms: u64) -> Result<VShardEnvelope> {
208                let partials = if req.vshard_id == 0 {
209                    self.shard0_partials.clone()
210                } else {
211                    self.shard1_partials.clone()
212                };
213                let resp = ArrayShardAggResp {
214                    shard_id: req.vshard_id,
215                    partials,
216                    truncated_before_horizon: false,
217                };
218                let payload = zerompk::to_msgpack_vec(&resp).unwrap();
219                Ok(VShardEnvelope::new(
220                    VShardMessageType::ArrayShardSliceResp,
221                    req.target_node,
222                    req.source_node,
223                    req.vshard_id,
224                    payload,
225                ))
226            }
227        }
228
229        let dispatch: Arc<dyn ShardRpcDispatch> = Arc::new(GroupedDispatch {
230            shard0_partials: vec![ArrayAggPartial::from_single(0, 5.0)],
231            shard1_partials: vec![
232                ArrayAggPartial::from_single(0, 15.0),
233                ArrayAggPartial::from_single(1, 20.0),
234            ],
235        });
236        let coord = make_coordinator(vec![0, 1], dispatch);
237        let merged = coord
238            .coord_agg(make_agg_req())
239            .await
240            .expect("grouped coord_agg should succeed");
241
242        // group_key=0: sum=5+15=20, count=2; group_key=1: sum=20, count=1.
243        assert_eq!(merged.partials.len(), 2);
244        let g0 = merged
245            .partials
246            .iter()
247            .find(|p| p.group_key == 0)
248            .expect("group 0");
249        let g1 = merged
250            .partials
251            .iter()
252            .find(|p| p.group_key == 1)
253            .expect("group 1");
254        assert!((g0.sum - 20.0).abs() < f64::EPSILON);
255        assert_eq!(g0.count, 2);
256        assert!((g1.sum - 20.0).abs() < f64::EPSILON);
257        assert_eq!(g1.count, 1);
258    }
259
260    #[tokio::test]
261    async fn coord_agg_or_reduces_truncated_before_horizon() {
262        // One shard reports below-horizon; the coordinator must OR-reduce the
263        // flag so the upstream caller can surface an incomplete-result signal.
264        // Dropping it here was a silent-partial-success bug.
265        struct HorizonDispatch;
266
267        #[async_trait]
268        impl ShardRpcDispatch for HorizonDispatch {
269            async fn call(&self, req: VShardEnvelope, _timeout_ms: u64) -> Result<VShardEnvelope> {
270                // Shard 1 is below horizon (zero partials); shard 0 has data.
271                let (partials, below) = if req.vshard_id == 0 {
272                    (vec![ArrayAggPartial::from_single(0, 10.0)], false)
273                } else {
274                    (vec![], true)
275                };
276                let resp = ArrayShardAggResp {
277                    shard_id: req.vshard_id,
278                    partials,
279                    truncated_before_horizon: below,
280                };
281                let payload = zerompk::to_msgpack_vec(&resp).unwrap();
282                Ok(VShardEnvelope::new(
283                    VShardMessageType::ArrayShardSliceResp,
284                    req.target_node,
285                    req.source_node,
286                    req.vshard_id,
287                    payload,
288                ))
289            }
290        }
291
292        let dispatch: Arc<dyn ShardRpcDispatch> = Arc::new(HorizonDispatch);
293        let coord = make_coordinator(vec![0, 1], dispatch);
294        let merged = coord
295            .coord_agg(make_agg_req())
296            .await
297            .expect("coord_agg should succeed");
298        assert!(
299            merged.truncated_before_horizon,
300            "coordinator must OR-reduce the below-horizon flag across shards"
301        );
302    }
303
304    #[tokio::test]
305    async fn coord_slice_zero_limit_returns_all() {
306        let row = zerompk::to_msgpack_vec(&"r").unwrap();
307        let dispatch: Arc<dyn ShardRpcDispatch> = Arc::new(SliceEchoDispatch {
308            rows: vec![row.clone(); 10],
309        });
310        let coord = make_coordinator(vec![0, 1], dispatch);
311        let req = ArrayShardSliceReq {
312            array_id_msgpack: vec![],
313            slice_msgpack: vec![],
314            attr_projection: vec![],
315            limit: 0,
316            cell_filter_msgpack: vec![],
317            prefix_bits: 0,
318            slice_hilbert_ranges: vec![],
319            shard_hilbert_range: None,
320            system_time: nodedb_types::SystemTimeScope::Current,
321            valid_at_ms: None,
322        };
323
324        // coordinator_limit = 0 → no cutoff → 20 rows.
325        let result = coord
326            .coord_slice(req, 0, nodedb_types::SystemTimeScope::Current)
327            .await
328            .expect("coord_slice unlimited should succeed");
329        assert_eq!(result.rows.len(), 20);
330    }
331
332    // ── coord_put / coord_delete tests ────────────────────────────────────
333
334    use super::super::wire::{ArrayShardDeleteResp, ArrayShardPutReq, ArrayShardPutResp};
335    use super::write::{ArrayWriteCoordParams, coord_delete, coord_put};
336    use crate::error::ClusterError;
337
338    /// Records which vShard IDs were called and echoes back an `ArrayShardPutResp`.
339    struct PutEchoDispatch;
340
341    #[async_trait]
342    impl ShardRpcDispatch for PutEchoDispatch {
343        async fn call(&self, req: VShardEnvelope, _timeout_ms: u64) -> Result<VShardEnvelope> {
344            let shard_req: ArrayShardPutReq = zerompk::from_msgpack(&req.payload).unwrap();
345            let resp = ArrayShardPutResp {
346                shard_id: req.vshard_id,
347                applied_lsn: shard_req.wal_lsn,
348            };
349            let payload = zerompk::to_msgpack_vec(&resp).unwrap();
350            Ok(VShardEnvelope::new(
351                VShardMessageType::ArrayShardSliceResp,
352                req.target_node,
353                req.source_node,
354                req.vshard_id,
355                payload,
356            ))
357        }
358    }
359
360    /// Dispatch that always returns a Codec error — used for failure-propagation tests.
361    struct FailDispatch;
362
363    #[async_trait]
364    impl ShardRpcDispatch for FailDispatch {
365        async fn call(&self, _req: VShardEnvelope, _timeout_ms: u64) -> Result<VShardEnvelope> {
366            Err(ClusterError::Codec {
367                detail: "injected failure".into(),
368            })
369        }
370    }
371
372    /// Echo dispatch for delete that returns an `ArrayShardDeleteResp`.
373    struct DeleteEchoDispatch;
374
375    #[async_trait]
376    impl ShardRpcDispatch for DeleteEchoDispatch {
377        async fn call(&self, req: VShardEnvelope, _timeout_ms: u64) -> Result<VShardEnvelope> {
378            use super::super::wire::ArrayShardDeleteReq;
379            let shard_req: ArrayShardDeleteReq = zerompk::from_msgpack(&req.payload).unwrap();
380            let resp = ArrayShardDeleteResp {
381                shard_id: req.vshard_id,
382                applied_lsn: shard_req.wal_lsn,
383            };
384            let payload = zerompk::to_msgpack_vec(&resp).unwrap();
385            Ok(VShardEnvelope::new(
386                VShardMessageType::ArrayShardSliceResp,
387                req.target_node,
388                req.source_node,
389                req.vshard_id,
390                payload,
391            ))
392        }
393    }
394
395    fn write_params() -> ArrayWriteCoordParams {
396        ArrayWriteCoordParams {
397            source_node: 1,
398            timeout_ms: 1000,
399        }
400    }
401
402    fn cb() -> Arc<CircuitBreaker> {
403        Arc::new(CircuitBreaker::new(CircuitBreakerConfig::default()))
404    }
405
406    #[tokio::test]
407    async fn coord_put_partitions_cells_by_tile() {
408        // prefix_bits=10, stride=1 → vshard == top-10-bit bucket.
409        // p0 → bucket 0 → vshard 0
410        // p1 → bucket 1 → vshard 1
411        // p2 → bucket 2 → vshard 2
412        let p0 = 0x0000_0000_0000_0000u64;
413        let p1 = 0x0040_0000_0000_0000u64;
414        let p2 = 0x0080_0000_0000_0000u64;
415
416        let cells = vec![
417            (p0, vec![0x01u8]),
418            (p1, vec![0x02u8]),
419            (p0, vec![0x03u8]),
420            (p2, vec![0x04u8]),
421            (p1, vec![0x05u8]),
422        ];
423
424        let dispatch: Arc<dyn ShardRpcDispatch> = Arc::new(PutEchoDispatch);
425        let mut resps = coord_put(&write_params(), vec![], 10, 42, &cells, &dispatch, &cb())
426            .await
427            .expect("coord_put should succeed");
428
429        resps.sort_by_key(|r| r.shard_id);
430        assert_eq!(resps.len(), 3, "should fan-out to 3 shards");
431        assert_eq!(resps[0].shard_id, 0);
432        assert_eq!(resps[1].shard_id, 1);
433        assert_eq!(resps[2].shard_id, 2);
434        // Each shard echoes back wal_lsn=42.
435        for r in &resps {
436            assert_eq!(r.applied_lsn, 42);
437        }
438    }
439
440    #[tokio::test]
441    async fn coord_put_aggregates_partial_failures() {
442        // A failing dispatch must surface as an error, not silent partial success.
443        let cells = vec![(0u64, vec![0xAAu8])];
444        let dispatch: Arc<dyn ShardRpcDispatch> = Arc::new(FailDispatch);
445        let err = coord_put(&write_params(), vec![], 10, 1, &cells, &dispatch, &cb())
446            .await
447            .expect_err("coord_put with failing shard should return error");
448        assert!(
449            matches!(err, ClusterError::Codec { .. }),
450            "expected Codec error, got {err:?}"
451        );
452    }
453
454    #[tokio::test]
455    async fn coord_delete_partitions_by_tile() {
456        let p0 = 0x0000_0000_0000_0000u64;
457        let p1 = 0x0040_0000_0000_0000u64;
458
459        let coords = vec![(p0, vec![0xAAu8]), (p1, vec![0xBBu8]), (p0, vec![0xCCu8])];
460
461        let dispatch: Arc<dyn ShardRpcDispatch> = Arc::new(DeleteEchoDispatch);
462        let mut resps = coord_delete(&write_params(), vec![], 10, 55, &coords, &dispatch, &cb())
463            .await
464            .expect("coord_delete should succeed");
465
466        resps.sort_by_key(|r| r.shard_id);
467        assert_eq!(resps.len(), 2, "should fan-out to 2 shards");
468        assert_eq!(resps[0].shard_id, 0);
469        assert_eq!(resps[1].shard_id, 1);
470        for r in &resps {
471            assert_eq!(r.applied_lsn, 55);
472        }
473    }
474}