Skip to main content

oxirs_stream/distributed/
shard_manager.rs

1//! # Shard manager
2//!
3//! Tracks the assignment of stream shards to cluster nodes and produces
4//! deterministic rebalance plans on node join/leave events.
5//!
6//! Each shard is identified by an integer in `[0, n_shards)`. A
7//! [`ShardAssignment`] is a snapshot mapping `shard_id -> node_id`. The
8//! manager keeps the latest assignment in memory and exposes operations to:
9//!
10//! * [`ShardManager::add_node`] / [`ShardManager::remove_node`] —
11//!   register / deregister a node and produce a [`RebalancePlan`] describing
12//!   what shard moves are needed to keep the assignment balanced.
13//! * [`ShardManager::owner_of`] — look up the owning node for a shard.
14//! * [`ShardManager::shards_owned_by`] — enumerate the shards a given node
15//!   currently owns.
16//! * [`ShardManager::current_assignment`] — clone the latest snapshot for
17//!   downstream propagation (e.g. through Raft via
18//!   [`super::coordinator::DistributedStreamCoordinator`]).
19//!
20//! The placement policy is **balanced round-robin** — when there are `K`
21//! nodes and `S` shards, every node owns either `floor(S/K)` or `ceil(S/K)`
22//! shards. This is deterministic for a given node ordering, so the plan a
23//! coordinator commits through Raft can be replayed by every node.
24
25use std::collections::{BTreeMap, HashMap, HashSet};
26use std::sync::Arc;
27
28use parking_lot::RwLock;
29use serde::{Deserialize, Serialize};
30use thiserror::Error;
31use tracing::debug;
32
33// ─── Errors ─────────────────────────────────────────────────────────────────
34
35/// Errors raised by [`ShardManager`].
36#[derive(Debug, Error)]
37pub enum ShardManagerError {
38    /// Tried to remove a node that is not registered.
39    #[error("unknown node: {0}")]
40    UnknownNode(String),
41    /// Tried to add a node that is already registered.
42    #[error("node already registered: {0}")]
43    NodeAlreadyExists(String),
44    /// `n_shards` was zero, which is rejected to avoid degenerate plans.
45    #[error("n_shards must be >= 1")]
46    NoShards,
47    /// All nodes were removed, leaving the cluster with no owners.
48    #[error("no nodes available to assign shards")]
49    NoNodes,
50}
51
52/// Convenience alias.
53pub type ShardManagerResult<T> = std::result::Result<T, ShardManagerError>;
54
55// ─── Types ─────────────────────────────────────────────────────────────────
56
57/// Stable shard identifier. Always `< n_shards`.
58pub type ShardId = u32;
59
60/// Stable node identifier (logical). The mapping to physical addresses is
61/// outside this module's responsibility.
62pub type NodeId = String;
63
64/// Snapshot of the shard → node mapping.
65#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
66pub struct ShardAssignment {
67    /// Map from shard id to node id.
68    pub map: BTreeMap<ShardId, NodeId>,
69}
70
71impl ShardAssignment {
72    /// Build an assignment from a flat vector indexed by shard id.
73    pub fn from_vec(nodes_per_shard: Vec<NodeId>) -> Self {
74        let map = nodes_per_shard
75            .into_iter()
76            .enumerate()
77            .map(|(i, n)| (i as ShardId, n))
78            .collect();
79        Self { map }
80    }
81
82    /// Total number of shards.
83    pub fn n_shards(&self) -> usize {
84        self.map.len()
85    }
86
87    /// Owner of a given shard.
88    pub fn owner_of(&self, shard: ShardId) -> Option<&NodeId> {
89        self.map.get(&shard)
90    }
91
92    /// Counts shards per owner.
93    pub fn counts(&self) -> HashMap<NodeId, usize> {
94        let mut counts = HashMap::new();
95        for owner in self.map.values() {
96            *counts.entry(owner.clone()).or_insert(0) += 1;
97        }
98        counts
99    }
100}
101
102/// A single shard reassignment.
103#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
104pub struct ShardMove {
105    pub shard: ShardId,
106    pub from: Option<NodeId>,
107    pub to: NodeId,
108}
109
110/// Result of a rebalance: the new full assignment plus the per-shard moves
111/// that need to take effect to transition from the old assignment to the new
112/// one.
113#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
114pub struct RebalancePlan {
115    /// The new assignment (after applying `moves`).
116    pub new_assignment: ShardAssignment,
117    /// Ordered list of shard moves.
118    pub moves: Vec<ShardMove>,
119}
120
121// ─── ShardManager ───────────────────────────────────────────────────────────
122
123/// Configuration for [`ShardManager`].
124#[derive(Debug, Clone, Serialize, Deserialize)]
125pub struct ShardManagerConfig {
126    /// Number of shards in the topology.
127    pub n_shards: u32,
128}
129
130impl Default for ShardManagerConfig {
131    fn default() -> Self {
132        Self { n_shards: 8 }
133    }
134}
135
136/// Tracks shard ownership and produces rebalance plans.
137pub struct ShardManager {
138    config: ShardManagerConfig,
139    /// Sorted list of registered nodes. `BTreeMap` preserves order so the
140    /// rebalance is deterministic.
141    nodes: RwLock<BTreeMap<NodeId, NodeMeta>>,
142    /// Latest assignment snapshot.
143    assignment: RwLock<ShardAssignment>,
144    /// Number of plans produced so far.
145    plans_emitted: Arc<RwLock<u64>>,
146}
147
148impl std::fmt::Debug for ShardManager {
149    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
150        f.debug_struct("ShardManager")
151            .field("config", &self.config)
152            .field("nodes", &self.nodes.read().keys().collect::<Vec<_>>())
153            .field("plans_emitted", &*self.plans_emitted.read())
154            .finish()
155    }
156}
157
158#[derive(Debug, Clone)]
159struct NodeMeta {
160    /// Insertion order (used to break ties in the round-robin allocation).
161    seq: u64,
162}
163
164impl ShardManager {
165    /// Build an empty manager. The first [`ShardManager::add_node`] call will
166    /// produce an initial plan.
167    pub fn new(config: ShardManagerConfig) -> ShardManagerResult<Self> {
168        if config.n_shards == 0 {
169            return Err(ShardManagerError::NoShards);
170        }
171        Ok(Self {
172            config,
173            nodes: RwLock::new(BTreeMap::new()),
174            assignment: RwLock::new(ShardAssignment::default()),
175            plans_emitted: Arc::new(RwLock::new(0)),
176        })
177    }
178
179    /// Build a manager pre-populated with the provided node ids.
180    pub fn with_nodes(
181        config: ShardManagerConfig,
182        nodes: impl IntoIterator<Item = impl Into<NodeId>>,
183    ) -> ShardManagerResult<Self> {
184        let mgr = Self::new(config)?;
185        for n in nodes {
186            let _ = mgr.add_node(n.into())?;
187        }
188        Ok(mgr)
189    }
190
191    /// Number of plans emitted so far.
192    pub fn plans_emitted(&self) -> u64 {
193        *self.plans_emitted.read()
194    }
195
196    /// Owner of a shard in the latest assignment.
197    pub fn owner_of(&self, shard: ShardId) -> Option<NodeId> {
198        self.assignment.read().owner_of(shard).cloned()
199    }
200
201    /// Shards currently owned by a node.
202    pub fn shards_owned_by(&self, node_id: &str) -> Vec<ShardId> {
203        self.assignment
204            .read()
205            .map
206            .iter()
207            .filter(|(_, owner)| owner.as_str() == node_id)
208            .map(|(s, _)| *s)
209            .collect()
210    }
211
212    /// Snapshot of the current assignment.
213    pub fn current_assignment(&self) -> ShardAssignment {
214        self.assignment.read().clone()
215    }
216
217    /// Total number of currently registered nodes.
218    pub fn node_count(&self) -> usize {
219        self.nodes.read().len()
220    }
221
222    /// Adds a node and returns the rebalance plan that brings the assignment
223    /// back in balance.
224    pub fn add_node(&self, node_id: NodeId) -> ShardManagerResult<RebalancePlan> {
225        {
226            let mut nodes = self.nodes.write();
227            if nodes.contains_key(&node_id) {
228                return Err(ShardManagerError::NodeAlreadyExists(node_id));
229            }
230            let seq = nodes.len() as u64;
231            nodes.insert(node_id.clone(), NodeMeta { seq });
232        }
233        let plan = self.recompute_plan()?;
234        debug!(node = %node_id, moves = plan.moves.len(), "shard manager: add_node");
235        Ok(plan)
236    }
237
238    /// Removes a node and returns the rebalance plan.
239    pub fn remove_node(&self, node_id: &str) -> ShardManagerResult<RebalancePlan> {
240        {
241            let mut nodes = self.nodes.write();
242            if nodes.remove(node_id).is_none() {
243                return Err(ShardManagerError::UnknownNode(node_id.to_string()));
244            }
245        }
246        let plan = self.recompute_plan()?;
247        debug!(node = %node_id, moves = plan.moves.len(), "shard manager: remove_node");
248        Ok(plan)
249    }
250
251    /// Apply an externally-supplied assignment (e.g. one that was committed
252    /// through Raft). Returns the diff against the current snapshot.
253    pub fn install_assignment(&self, new_assignment: ShardAssignment) -> RebalancePlan {
254        let old = self.assignment.read().clone();
255        let moves = compute_moves(&old, &new_assignment);
256        *self.assignment.write() = new_assignment.clone();
257        *self.plans_emitted.write() += 1;
258        RebalancePlan {
259            new_assignment,
260            moves,
261        }
262    }
263
264    fn recompute_plan(&self) -> ShardManagerResult<RebalancePlan> {
265        let nodes_snap = self.nodes.read().clone();
266        if nodes_snap.is_empty() {
267            // Special case: empty assignment.
268            let empty = ShardAssignment::default();
269            let old = self.assignment.read().clone();
270            let moves: Vec<ShardMove> = old
271                .map
272                .iter()
273                .map(|(shard, owner)| ShardMove {
274                    shard: *shard,
275                    from: Some(owner.clone()),
276                    to: String::new(),
277                })
278                .collect();
279            *self.assignment.write() = empty.clone();
280            *self.plans_emitted.write() += 1;
281            return Ok(RebalancePlan {
282                new_assignment: empty,
283                moves,
284            });
285        }
286
287        let nodes: Vec<NodeId> = {
288            let mut by_seq: Vec<(u64, NodeId)> = nodes_snap
289                .iter()
290                .map(|(id, m)| (m.seq, id.clone()))
291                .collect();
292            by_seq.sort();
293            by_seq.into_iter().map(|(_, id)| id).collect()
294        };
295
296        let new_assignment = balanced_assignment(self.config.n_shards, &nodes);
297        let old = self.assignment.read().clone();
298        let moves = compute_moves(&old, &new_assignment);
299        *self.assignment.write() = new_assignment.clone();
300        *self.plans_emitted.write() += 1;
301        Ok(RebalancePlan {
302            new_assignment,
303            moves,
304        })
305    }
306}
307
308/// Build a deterministic balanced assignment for `n_shards` over the provided
309/// node ordering. Each node receives either `floor(n_shards / N)` or
310/// `ceil(n_shards / N)` shards.
311fn balanced_assignment(n_shards: u32, nodes: &[NodeId]) -> ShardAssignment {
312    if nodes.is_empty() {
313        return ShardAssignment::default();
314    }
315    let n = nodes.len() as u32;
316    let mut map = BTreeMap::new();
317    for shard in 0..n_shards {
318        let owner = &nodes[(shard % n) as usize];
319        map.insert(shard, owner.clone());
320    }
321    ShardAssignment { map }
322}
323
324fn compute_moves(old: &ShardAssignment, new_assignment: &ShardAssignment) -> Vec<ShardMove> {
325    let mut moves = Vec::new();
326    let all_shards: HashSet<ShardId> = old
327        .map
328        .keys()
329        .chain(new_assignment.map.keys())
330        .cloned()
331        .collect();
332    let mut shards: Vec<ShardId> = all_shards.into_iter().collect();
333    shards.sort();
334    for shard in shards {
335        let from = old.map.get(&shard).cloned();
336        let to = new_assignment.map.get(&shard).cloned();
337        match (from, to) {
338            (Some(f), Some(t)) if f == t => {}
339            (Some(f), Some(t)) => moves.push(ShardMove {
340                shard,
341                from: Some(f),
342                to: t,
343            }),
344            (None, Some(t)) => moves.push(ShardMove {
345                shard,
346                from: None,
347                to: t,
348            }),
349            (Some(f), None) => moves.push(ShardMove {
350                shard,
351                from: Some(f),
352                to: String::new(),
353            }),
354            (None, None) => {}
355        }
356    }
357    moves
358}
359
360// ─── Tests ──────────────────────────────────────────────────────────────────
361
362#[cfg(test)]
363mod tests {
364    use super::*;
365
366    #[test]
367    fn balanced_assignment_round_robins() {
368        let assignment =
369            balanced_assignment(6, &["n1".to_string(), "n2".to_string(), "n3".to_string()]);
370        let counts = assignment.counts();
371        for c in counts.values() {
372            assert_eq!(*c, 2);
373        }
374    }
375
376    #[test]
377    fn add_node_initial_plan() {
378        let mgr = ShardManager::new(ShardManagerConfig { n_shards: 4 }).expect("ok");
379        let plan = mgr.add_node("n1".into()).expect("add");
380        assert_eq!(plan.new_assignment.n_shards(), 4);
381        for owner in plan.new_assignment.map.values() {
382            assert_eq!(owner, "n1");
383        }
384    }
385
386    #[test]
387    fn add_node_balances_existing() {
388        let mgr = ShardManager::new(ShardManagerConfig { n_shards: 6 }).expect("ok");
389        mgr.add_node("n1".into()).expect("ok");
390        let plan = mgr.add_node("n2".into()).expect("ok");
391        let counts = plan.new_assignment.counts();
392        assert_eq!(counts.get("n1"), Some(&3));
393        assert_eq!(counts.get("n2"), Some(&3));
394        assert_eq!(plan.moves.len(), 3);
395    }
396
397    #[test]
398    fn remove_node_redistributes() {
399        let mgr = ShardManager::new(ShardManagerConfig { n_shards: 6 }).expect("ok");
400        mgr.add_node("n1".into()).expect("ok");
401        mgr.add_node("n2".into()).expect("ok");
402        mgr.add_node("n3".into()).expect("ok");
403        let plan = mgr.remove_node("n2").expect("ok");
404        let counts = plan.new_assignment.counts();
405        assert!(!counts.contains_key("n2"));
406        let total: usize = counts.values().sum();
407        assert_eq!(total, 6);
408    }
409
410    #[test]
411    fn empty_node_list_returns_empty_assignment() {
412        let mgr = ShardManager::new(ShardManagerConfig { n_shards: 3 }).expect("ok");
413        mgr.add_node("n1".into()).expect("ok");
414        let plan = mgr.remove_node("n1").expect("ok");
415        assert!(plan.new_assignment.map.is_empty());
416        assert_eq!(plan.moves.len(), 3);
417    }
418
419    #[test]
420    fn install_assignment_overrides_state() {
421        let mgr = ShardManager::new(ShardManagerConfig { n_shards: 2 }).expect("ok");
422        let new_assignment = ShardAssignment::from_vec(vec!["nA".into(), "nB".into()]);
423        let plan = mgr.install_assignment(new_assignment.clone());
424        assert_eq!(plan.new_assignment, new_assignment);
425        assert_eq!(mgr.owner_of(0), Some("nA".to_string()));
426        assert_eq!(mgr.owner_of(1), Some("nB".to_string()));
427        assert_eq!(plan.moves.len(), 2);
428    }
429
430    #[test]
431    fn duplicate_add_rejected() {
432        let mgr = ShardManager::new(ShardManagerConfig { n_shards: 2 }).expect("ok");
433        mgr.add_node("n1".into()).expect("ok");
434        let err = mgr.add_node("n1".into()).expect_err("should fail");
435        assert!(matches!(err, ShardManagerError::NodeAlreadyExists(_)));
436    }
437
438    #[test]
439    fn unknown_remove_rejected() {
440        let mgr = ShardManager::new(ShardManagerConfig { n_shards: 2 }).expect("ok");
441        let err = mgr.remove_node("ghost").expect_err("should fail");
442        assert!(matches!(err, ShardManagerError::UnknownNode(_)));
443    }
444
445    #[test]
446    fn n_shards_zero_rejected() {
447        let err = ShardManager::new(ShardManagerConfig { n_shards: 0 }).expect_err("should fail");
448        assert!(matches!(err, ShardManagerError::NoShards));
449    }
450
451    #[test]
452    fn shards_owned_by_returns_correct_subset() {
453        let mgr =
454            ShardManager::with_nodes(ShardManagerConfig { n_shards: 4 }, ["n1", "n2"]).expect("ok");
455        let s1 = mgr.shards_owned_by("n1");
456        let s2 = mgr.shards_owned_by("n2");
457        // 0,2 → n1; 1,3 → n2 in deterministic order.
458        assert_eq!(s1, vec![0, 2]);
459        assert_eq!(s2, vec![1, 3]);
460    }
461}