Skip to main content

kvbm_engine/leader/
state.rs

1// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4//! LeaderState - Coordination layer for managing workers.
5//!
6//! This module provides the leader's coordination state, including:
7//! - Worker registration and rank mapping
8//! - Remote leader tracking for cross-leader transfers
9//! - Routing strategies for asymmetric TP configurations
10
11use std::collections::HashMap;
12use std::sync::{Arc, RwLock};
13
14use ::velo::Messenger;
15use anyhow::Result;
16
17use crate::InstanceId;
18use crate::worker::{CoordinatedWorker, Worker};
19use kvbm_physical::manager::SerializedLayout;
20
21/// Info about a remote leader and its workers.
22#[derive(Debug)]
23pub struct RemoteLeaderInfo {
24    /// Instance ID of the remote leader process
25    pub instance_id: InstanceId,
26    /// Number of workers under the remote leader
27    pub worker_count: usize,
28    /// Cached metadata from remote workers (rank-ordered)
29    pub worker_metadata: Vec<SerializedLayout>,
30}
31
32/// Leader coordination state - owns workers and routing logic.
33///
34/// LeaderState manages:
35/// - Registration of workers during handshake phase
36/// - Coordination with remote leaders for cross-leader transfers
37/// - Routing strategies for asymmetric TP configurations
38pub struct LeaderState {
39    /// This leader's instance ID
40    instance_id: InstanceId,
41
42    /// Nova runtime for RPC
43    messenger: Arc<Messenger>,
44
45    /// Workers under this leader (rank-ordered)
46    workers: Vec<CoordinatedWorker>,
47
48    /// Known remote leaders (by their instance ID)
49    remote_leaders: RwLock<HashMap<InstanceId, RemoteLeaderInfo>>,
50}
51
52impl LeaderState {
53    /// Create a new LeaderState.
54    ///
55    /// # Arguments
56    /// * `instance_id` - This leader's unique identifier
57    /// * `nova` - Nova runtime for RPC communication
58    pub fn new(instance_id: InstanceId, messenger: Arc<Messenger>) -> Self {
59        Self {
60            instance_id,
61            messenger,
62            workers: Vec::new(),
63            remote_leaders: RwLock::new(HashMap::new()),
64        }
65    }
66
67    /// Get this leader's instance ID.
68    pub fn instance_id(&self) -> InstanceId {
69        self.instance_id
70    }
71
72    /// Get the Nova runtime.
73    pub fn nova(&self) -> &Arc<Messenger> {
74        &self.messenger
75    }
76
77    /// Register a worker during the handshake phase.
78    ///
79    /// Workers should be registered in rank order (0, 1, 2, ...).
80    ///
81    /// # Arguments
82    /// * `rank` - The worker's rank (0-indexed)
83    /// * `host_instance` - Instance ID of the process hosting this worker
84    /// * `worker` - The Worker implementation (DirectWorker or VeloWorkerClient)
85    pub fn register_worker(
86        &mut self,
87        rank: usize,
88        host_instance: InstanceId,
89        worker: Box<dyn Worker>,
90    ) {
91        let coordinated = CoordinatedWorker::new(worker, rank, host_instance);
92
93        // Ensure rank-ordered insertion
94        if rank == self.workers.len() {
95            // Sequential append (expected path)
96            self.workers.push(coordinated);
97        } else if rank < self.workers.len() {
98            // Re-registration or out-of-order within existing range
99            self.workers[rank] = coordinated;
100        } else {
101            panic!(
102                "Gap in worker ranks: rank {} but only {} workers registered",
103                rank,
104                self.workers.len()
105            );
106        }
107    }
108
109    /// Number of workers under this leader.
110    pub fn worker_count(&self) -> usize {
111        self.workers.len()
112    }
113
114    /// Get a worker by rank.
115    pub fn worker(&self, rank: usize) -> Option<&CoordinatedWorker> {
116        self.workers.get(rank)
117    }
118
119    /// Get a mutable worker by rank.
120    pub fn worker_mut(&mut self, rank: usize) -> Option<&mut CoordinatedWorker> {
121        self.workers.get_mut(rank)
122    }
123
124    /// Iterate over all workers.
125    pub fn workers(&self) -> impl Iterator<Item = &CoordinatedWorker> {
126        self.workers.iter()
127    }
128
129    /// Connect to a remote leader and distribute its worker metadata to our workers.
130    ///
131    /// This implements the routing strategy for cross-leader transfers:
132    /// - 1:1 mapping when TP sizes match
133    /// - Many-to-one when local TP > remote TP
134    /// - One-to-many when local TP < remote TP
135    ///
136    /// # Arguments
137    /// * `remote_leader_id` - Instance ID of the remote leader
138    /// * `remote_metadata` - Metadata from each remote worker (rank-ordered)
139    pub async fn import_remote_leader(
140        &self,
141        remote_leader_id: InstanceId,
142        remote_metadata: Vec<SerializedLayout>,
143    ) -> Result<()> {
144        let remote_count = remote_metadata.len();
145        let local_count = self.workers.len();
146
147        tracing::info!(
148            local_count,
149            remote_count,
150            %remote_leader_id,
151            "Importing remote leader metadata"
152        );
153
154        // Store remote leader info
155        {
156            let mut leaders = self.remote_leaders.write().unwrap();
157            leaders.insert(
158                remote_leader_id,
159                RemoteLeaderInfo {
160                    instance_id: remote_leader_id,
161                    worker_count: remote_count,
162                    worker_metadata: remote_metadata.clone(),
163                },
164            );
165        }
166
167        // Distribute metadata based on routing strategy
168        for (local_rank, worker) in self.workers.iter().enumerate() {
169            let target_remote_ranks = route_local_to_remote(local_rank, local_count, remote_count);
170
171            for remote_rank in target_remote_ranks {
172                tracing::debug!(
173                    local_rank,
174                    remote_rank,
175                    %remote_leader_id,
176                    "Importing remote metadata for local worker"
177                );
178
179                worker
180                    .import_remote_metadata(
181                        remote_leader_id,
182                        remote_rank,
183                        remote_metadata[remote_rank].clone(),
184                    )
185                    .await?;
186            }
187        }
188
189        Ok(())
190    }
191
192    /// Export this leader's workers' metadata for another leader to import.
193    ///
194    /// Returns metadata from each worker in rank order.
195    pub async fn export_worker_metadata(&self) -> Result<Vec<SerializedLayout>> {
196        let mut metadata = Vec::with_capacity(self.workers.len());
197
198        for worker in &self.workers {
199            let response = worker.inner().export_metadata()?;
200            metadata.push(response.await?);
201        }
202
203        Ok(metadata)
204    }
205
206    /// Check if we have imported metadata from a remote leader.
207    pub fn has_remote_leader(&self, remote_leader_id: InstanceId) -> bool {
208        self.remote_leaders
209            .read()
210            .unwrap()
211            .contains_key(&remote_leader_id)
212    }
213
214    /// Get info about a remote leader if known.
215    pub fn remote_leader_info(&self, remote_leader_id: InstanceId) -> Option<RemoteLeaderInfo> {
216        self.remote_leaders
217            .read()
218            .unwrap()
219            .get(&remote_leader_id)
220            .map(|info| RemoteLeaderInfo {
221                instance_id: info.instance_id,
222                worker_count: info.worker_count,
223                worker_metadata: info.worker_metadata.clone(),
224            })
225    }
226}
227
228/// Routing strategy: which local ranks receive from which remote ranks.
229///
230/// This function determines how metadata/transfers are routed when
231/// the local and remote TP sizes differ.
232///
233/// # Examples
234/// - TP=4 local, TP=4 remote: 1:1 mapping (rank 0→0, 1→1, 2→2, 3→3)
235/// - TP=4 local, TP=2 remote: 0→0, 1→0, 2→1, 3→1 (many-to-one)
236/// - TP=2 local, TP=4 remote: 0→\[0,1\], 1→\[2,3\] (one-to-many)
237pub fn route_local_to_remote(
238    local_rank: usize,
239    local_count: usize,
240    remote_count: usize,
241) -> Vec<usize> {
242    if local_count == remote_count {
243        // 1:1 mapping
244        vec![local_rank]
245    } else if local_count > remote_count {
246        // Many local → few remote: multiple locals share a remote
247        vec![local_rank % remote_count]
248    } else {
249        // Few local → many remote: each local gets multiple remotes
250        let remotes_per_local = remote_count / local_count;
251        let start = local_rank * remotes_per_local;
252        // Last local rank absorbs any remainder from non-divisible ratios
253        let end = if local_rank == local_count - 1 {
254            remote_count
255        } else {
256            start + remotes_per_local
257        };
258        (start..end).collect()
259    }
260}
261
262#[cfg(test)]
263mod tests {
264    use super::*;
265
266    #[test]
267    fn test_route_1_to_1() {
268        // Same TP size
269        assert_eq!(route_local_to_remote(0, 4, 4), vec![0]);
270        assert_eq!(route_local_to_remote(1, 4, 4), vec![1]);
271        assert_eq!(route_local_to_remote(2, 4, 4), vec![2]);
272        assert_eq!(route_local_to_remote(3, 4, 4), vec![3]);
273    }
274
275    #[test]
276    fn test_route_many_to_one() {
277        // Local TP=4, Remote TP=2
278        assert_eq!(route_local_to_remote(0, 4, 2), vec![0]);
279        assert_eq!(route_local_to_remote(1, 4, 2), vec![1]);
280        assert_eq!(route_local_to_remote(2, 4, 2), vec![0]);
281        assert_eq!(route_local_to_remote(3, 4, 2), vec![1]);
282    }
283
284    #[test]
285    fn test_route_one_to_many() {
286        // Local TP=2, Remote TP=4
287        assert_eq!(route_local_to_remote(0, 2, 4), vec![0, 1]);
288        assert_eq!(route_local_to_remote(1, 2, 4), vec![2, 3]);
289    }
290
291    #[test]
292    fn test_route_4_to_8() {
293        // Local TP=4, Remote TP=8
294        assert_eq!(route_local_to_remote(0, 4, 8), vec![0, 1]);
295        assert_eq!(route_local_to_remote(1, 4, 8), vec![2, 3]);
296        assert_eq!(route_local_to_remote(2, 4, 8), vec![4, 5]);
297        assert_eq!(route_local_to_remote(3, 4, 8), vec![6, 7]);
298    }
299
300    #[test]
301    fn test_route_non_divisible_remainder() {
302        // Local TP=2, Remote TP=5: last local rank absorbs remainder
303        assert_eq!(route_local_to_remote(0, 2, 5), vec![0, 1]);
304        assert_eq!(route_local_to_remote(1, 2, 5), vec![2, 3, 4]);
305
306        // Local TP=3, Remote TP=7: last rank gets extras
307        assert_eq!(route_local_to_remote(0, 3, 7), vec![0, 1]);
308        assert_eq!(route_local_to_remote(1, 3, 7), vec![2, 3]);
309        assert_eq!(route_local_to_remote(2, 3, 7), vec![4, 5, 6]);
310    }
311}