use std::collections::HashMap;
use std::sync::{Arc, RwLock};
use ::velo::Messenger;
use anyhow::Result;
use crate::InstanceId;
use crate::worker::{CoordinatedWorker, Worker};
use kvbm_physical::manager::SerializedLayout;
#[derive(Debug)]
pub struct RemoteLeaderInfo {
pub instance_id: InstanceId,
pub worker_count: usize,
pub worker_metadata: Vec<SerializedLayout>,
}
pub struct LeaderState {
instance_id: InstanceId,
messenger: Arc<Messenger>,
workers: Vec<CoordinatedWorker>,
remote_leaders: RwLock<HashMap<InstanceId, RemoteLeaderInfo>>,
}
impl LeaderState {
pub fn new(instance_id: InstanceId, messenger: Arc<Messenger>) -> Self {
Self {
instance_id,
messenger,
workers: Vec::new(),
remote_leaders: RwLock::new(HashMap::new()),
}
}
pub fn instance_id(&self) -> InstanceId {
self.instance_id
}
pub fn nova(&self) -> &Arc<Messenger> {
&self.messenger
}
pub fn register_worker(
&mut self,
rank: usize,
host_instance: InstanceId,
worker: Box<dyn Worker>,
) {
let coordinated = CoordinatedWorker::new(worker, rank, host_instance);
if rank == self.workers.len() {
self.workers.push(coordinated);
} else if rank < self.workers.len() {
self.workers[rank] = coordinated;
} else {
panic!(
"Gap in worker ranks: rank {} but only {} workers registered",
rank,
self.workers.len()
);
}
}
pub fn worker_count(&self) -> usize {
self.workers.len()
}
pub fn worker(&self, rank: usize) -> Option<&CoordinatedWorker> {
self.workers.get(rank)
}
pub fn worker_mut(&mut self, rank: usize) -> Option<&mut CoordinatedWorker> {
self.workers.get_mut(rank)
}
pub fn workers(&self) -> impl Iterator<Item = &CoordinatedWorker> {
self.workers.iter()
}
pub async fn import_remote_leader(
&self,
remote_leader_id: InstanceId,
remote_metadata: Vec<SerializedLayout>,
) -> Result<()> {
let remote_count = remote_metadata.len();
let local_count = self.workers.len();
tracing::info!(
local_count,
remote_count,
%remote_leader_id,
"Importing remote leader metadata"
);
{
let mut leaders = self.remote_leaders.write().unwrap();
leaders.insert(
remote_leader_id,
RemoteLeaderInfo {
instance_id: remote_leader_id,
worker_count: remote_count,
worker_metadata: remote_metadata.clone(),
},
);
}
for (local_rank, worker) in self.workers.iter().enumerate() {
let target_remote_ranks = route_local_to_remote(local_rank, local_count, remote_count);
for remote_rank in target_remote_ranks {
tracing::debug!(
local_rank,
remote_rank,
%remote_leader_id,
"Importing remote metadata for local worker"
);
worker
.import_remote_metadata(
remote_leader_id,
remote_rank,
remote_metadata[remote_rank].clone(),
)
.await?;
}
}
Ok(())
}
pub async fn export_worker_metadata(&self) -> Result<Vec<SerializedLayout>> {
let mut metadata = Vec::with_capacity(self.workers.len());
for worker in &self.workers {
let response = worker.inner().export_metadata()?;
metadata.push(response.await?);
}
Ok(metadata)
}
pub fn has_remote_leader(&self, remote_leader_id: InstanceId) -> bool {
self.remote_leaders
.read()
.unwrap()
.contains_key(&remote_leader_id)
}
pub fn remote_leader_info(&self, remote_leader_id: InstanceId) -> Option<RemoteLeaderInfo> {
self.remote_leaders
.read()
.unwrap()
.get(&remote_leader_id)
.map(|info| RemoteLeaderInfo {
instance_id: info.instance_id,
worker_count: info.worker_count,
worker_metadata: info.worker_metadata.clone(),
})
}
}
pub fn route_local_to_remote(
local_rank: usize,
local_count: usize,
remote_count: usize,
) -> Vec<usize> {
if local_count == remote_count {
vec![local_rank]
} else if local_count > remote_count {
vec![local_rank % remote_count]
} else {
let remotes_per_local = remote_count / local_count;
let start = local_rank * remotes_per_local;
let end = if local_rank == local_count - 1 {
remote_count
} else {
start + remotes_per_local
};
(start..end).collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_route_1_to_1() {
assert_eq!(route_local_to_remote(0, 4, 4), vec![0]);
assert_eq!(route_local_to_remote(1, 4, 4), vec![1]);
assert_eq!(route_local_to_remote(2, 4, 4), vec![2]);
assert_eq!(route_local_to_remote(3, 4, 4), vec![3]);
}
#[test]
fn test_route_many_to_one() {
assert_eq!(route_local_to_remote(0, 4, 2), vec![0]);
assert_eq!(route_local_to_remote(1, 4, 2), vec![1]);
assert_eq!(route_local_to_remote(2, 4, 2), vec![0]);
assert_eq!(route_local_to_remote(3, 4, 2), vec![1]);
}
#[test]
fn test_route_one_to_many() {
assert_eq!(route_local_to_remote(0, 2, 4), vec![0, 1]);
assert_eq!(route_local_to_remote(1, 2, 4), vec![2, 3]);
}
#[test]
fn test_route_4_to_8() {
assert_eq!(route_local_to_remote(0, 4, 8), vec![0, 1]);
assert_eq!(route_local_to_remote(1, 4, 8), vec![2, 3]);
assert_eq!(route_local_to_remote(2, 4, 8), vec![4, 5]);
assert_eq!(route_local_to_remote(3, 4, 8), vec![6, 7]);
}
#[test]
fn test_route_non_divisible_remainder() {
assert_eq!(route_local_to_remote(0, 2, 5), vec![0, 1]);
assert_eq!(route_local_to_remote(1, 2, 5), vec![2, 3, 4]);
assert_eq!(route_local_to_remote(0, 3, 7), vec![0, 1]);
assert_eq!(route_local_to_remote(1, 3, 7), vec![2, 3]);
assert_eq!(route_local_to_remote(2, 3, 7), vec![4, 5, 6]);
}
}