1use 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#[derive(Debug)]
23pub struct RemoteLeaderInfo {
24 pub instance_id: InstanceId,
26 pub worker_count: usize,
28 pub worker_metadata: Vec<SerializedLayout>,
30}
31
32pub struct LeaderState {
39 instance_id: InstanceId,
41
42 messenger: Arc<Messenger>,
44
45 workers: Vec<CoordinatedWorker>,
47
48 remote_leaders: RwLock<HashMap<InstanceId, RemoteLeaderInfo>>,
50}
51
52impl LeaderState {
53 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 pub fn instance_id(&self) -> InstanceId {
69 self.instance_id
70 }
71
72 pub fn nova(&self) -> &Arc<Messenger> {
74 &self.messenger
75 }
76
77 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 if rank == self.workers.len() {
95 self.workers.push(coordinated);
97 } else if rank < self.workers.len() {
98 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 pub fn worker_count(&self) -> usize {
111 self.workers.len()
112 }
113
114 pub fn worker(&self, rank: usize) -> Option<&CoordinatedWorker> {
116 self.workers.get(rank)
117 }
118
119 pub fn worker_mut(&mut self, rank: usize) -> Option<&mut CoordinatedWorker> {
121 self.workers.get_mut(rank)
122 }
123
124 pub fn workers(&self) -> impl Iterator<Item = &CoordinatedWorker> {
126 self.workers.iter()
127 }
128
129 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 {
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 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 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 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 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
228pub 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 vec![local_rank]
245 } else if local_count > remote_count {
246 vec![local_rank % remote_count]
248 } else {
249 let remotes_per_local = remote_count / local_count;
251 let start = local_rank * remotes_per_local;
252 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 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 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 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 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 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 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}