use std::collections::{HashMap, HashSet};
use super::LoraAllocator;
use dynamo_kv_router::protocols::WorkerWithDpRank;
pub struct RendezvousHasher;
impl RendezvousHasher {
pub fn compute_score(lora_name: &str, worker: WorkerWithDpRank) -> u64 {
let mut hasher = blake3::Hasher::new();
hasher.update(lora_name.as_bytes());
hasher.update(&worker.worker_id.to_le_bytes());
hasher.update(&worker.dp_rank.to_le_bytes());
let hash = hasher.finalize();
let hash_bytes = hash.as_bytes();
let mut bytes_array = [0u8; 8];
bytes_array.copy_from_slice(&hash_bytes[..8]);
u64::from_le_bytes(bytes_array)
}
pub fn rank_workers(
lora_name: &str,
workers: &[WorkerWithDpRank],
) -> Vec<(WorkerWithDpRank, u64)> {
let mut scores: Vec<_> = workers
.iter()
.map(|&w| (w, Self::compute_score(lora_name, w)))
.collect();
scores.sort_by(|(wa, sa), (wb, sb)| {
sb.cmp(sa)
.then_with(|| wa.worker_id.cmp(&wb.worker_id))
.then_with(|| wa.dp_rank.cmp(&wb.dp_rank))
});
scores
}
}
impl LoraAllocator for RendezvousHasher {
fn compute_replica_set(
&self,
lora_name: &str,
workers: &[WorkerWithDpRank],
replica_factor: usize,
) -> Vec<WorkerWithDpRank> {
if workers.is_empty() {
return Vec::new();
}
let ranked = Self::rank_workers(lora_name, workers);
ranked
.into_iter()
.take(replica_factor.min(workers.len()))
.map(|(w, _)| w)
.collect()
}
fn compute_replica_set_with_slots(
&self,
lora_name: &str,
workers: &[WorkerWithDpRank],
replica_factor: usize,
worker_slot_usage: &HashMap<WorkerWithDpRank, (usize, usize)>,
) -> Vec<WorkerWithDpRank> {
if workers.is_empty() || replica_factor == 0 {
return Vec::new();
}
let ranked = Self::rank_workers(lora_name, workers);
let mut result = Vec::with_capacity(replica_factor);
for (worker, _score) in &ranked {
if result.len() >= replica_factor {
break;
}
if let Some(&(used, cap)) = worker_slot_usage.get(worker)
&& used >= cap
{
continue;
}
result.push(*worker);
}
if result.is_empty() {
return self.compute_replica_set(lora_name, workers, replica_factor);
}
result
}
fn compute_replica_set_with_slots_sticky(
&self,
lora_name: &str,
workers: &[WorkerWithDpRank],
replica_factor: usize,
worker_slot_usage: &HashMap<WorkerWithDpRank, (usize, usize)>,
prior: &[WorkerWithDpRank],
) -> Vec<WorkerWithDpRank> {
if workers.is_empty() || replica_factor == 0 {
return Vec::new();
}
let ranked = Self::rank_workers(lora_name, workers);
let prior_set: HashSet<WorkerWithDpRank> = prior.iter().copied().collect();
let mut result = Vec::with_capacity(replica_factor);
let mut chosen: HashSet<WorkerWithDpRank> = HashSet::new();
let is_full = |w: &WorkerWithDpRank| {
worker_slot_usage
.get(w)
.map(|&(used, cap)| used >= cap)
.unwrap_or(false)
};
for (worker, _score) in &ranked {
if result.len() >= replica_factor {
break;
}
if prior_set.contains(worker) && !is_full(worker) {
result.push(*worker);
chosen.insert(*worker);
}
}
for (worker, _score) in &ranked {
if result.len() >= replica_factor {
break;
}
if chosen.contains(worker) || is_full(worker) {
continue;
}
result.push(*worker);
chosen.insert(*worker);
}
if result.is_empty() {
return self.compute_replica_set(lora_name, workers, replica_factor);
}
result
}
fn name(&self) -> &str {
"hrw"
}
}
#[cfg(test)]
mod tests {
use super::*;
fn make_workers(count: usize) -> Vec<WorkerWithDpRank> {
(0..count)
.map(|i| WorkerWithDpRank::new(i as u64, 0))
.collect()
}
#[test]
fn test_deterministic() {
let worker = WorkerWithDpRank::new(1, 0);
let lora_name = "test-lora";
let score1 = RendezvousHasher::compute_score(lora_name, worker);
let score2 = RendezvousHasher::compute_score(lora_name, worker);
assert_eq!(score1, score2, "Same inputs should produce same score");
}
#[test]
fn test_stability_adding_workers() {
let workers_before = make_workers(3);
let hasher = RendezvousHasher;
let replica_set_before = hasher.compute_replica_set("test-lora", &workers_before, 2);
assert_eq!(replica_set_before.len(), 2);
let workers_after = make_workers(5);
let replica_set_after = hasher.compute_replica_set("test-lora", &workers_after, 2);
assert_eq!(replica_set_after.len(), 2);
let top2_after: Vec<_> = replica_set_after.iter().map(|w| w.worker_id).collect();
let replica_set_after2 = hasher.compute_replica_set("test-lora", &workers_after, 2);
let top2_after2: Vec<_> = replica_set_after2.iter().map(|w| w.worker_id).collect();
assert_eq!(
top2_after, top2_after2,
"Same inputs should produce same outputs"
);
}
#[test]
fn test_stability_removing_workers() {
let hasher = RendezvousHasher;
let workers_5 = make_workers(5);
let set_5 = hasher.compute_replica_set("test-lora", &workers_5, 3);
assert_eq!(set_5.len(), 3);
let workers_4: Vec<_> = workers_5
.iter()
.filter(|w| w.worker_id != 2)
.copied()
.collect();
let set_4 = hasher.compute_replica_set("test-lora", &workers_4, 3);
assert_eq!(set_4.len(), 3);
if !set_5.iter().any(|w| w.worker_id == 2) {
for worker in &set_5 {
if workers_4.contains(worker) {
assert!(
set_4.contains(worker),
"Worker {} was in top 3 and is still available, should remain in top 3",
worker.worker_id
);
}
}
}
}
#[test]
fn test_compute_replica_set_more_replicas_than_workers() {
let hasher = RendezvousHasher;
let workers = make_workers(3);
let result = hasher.compute_replica_set("test-lora", &workers, 10);
assert_eq!(result.len(), 3);
}
#[test]
fn test_slot_aware_skips_full_workers() {
let hasher = RendezvousHasher;
let workers = make_workers(5);
let ranked = RendezvousHasher::rank_workers("test-lora", &workers);
let mut usage = HashMap::new();
for (w, _) in ranked.iter().take(2) {
usage.insert(*w, (4usize, 4usize));
}
for (w, _) in ranked.iter().skip(2) {
usage.insert(*w, (1usize, 4usize));
}
let result = hasher.compute_replica_set_with_slots("test-lora", &workers, 2, &usage);
assert_eq!(result.len(), 2);
for w in &result {
let (used, cap) = usage[w];
assert!(used < cap, "Selected worker {:?} should not be full", w);
}
}
#[test]
fn test_slot_aware_falls_back_when_all_full() {
let hasher = RendezvousHasher;
let workers = make_workers(3);
let mut usage = HashMap::new();
for w in &workers {
usage.insert(*w, (4usize, 4usize));
}
let result = hasher.compute_replica_set_with_slots("test-lora", &workers, 2, &usage);
assert_eq!(result.len(), 2);
}
#[test]
fn test_ranking_equal_scores_tie_break_is_deterministic() {
let workers_fwd = vec![
WorkerWithDpRank::new(10, 0),
WorkerWithDpRank::new(10, 1),
WorkerWithDpRank::new(7, 0),
WorkerWithDpRank::new(7, 1),
];
let mut workers_rev = workers_fwd.clone();
workers_rev.reverse();
let r1: Vec<_> = RendezvousHasher::rank_workers("lora-x", &workers_fwd)
.into_iter()
.map(|(w, _)| (w.worker_id, w.dp_rank))
.collect();
let r2: Vec<_> = RendezvousHasher::rank_workers("lora-x", &workers_rev)
.into_iter()
.map(|(w, _)| (w.worker_id, w.dp_rank))
.collect();
assert_eq!(r1, r2, "ranking must be invariant to worker input order");
}
#[test]
fn test_sticky_retains_prior_placement() {
let hasher = RendezvousHasher;
let workers = make_workers(5);
let usage: HashMap<_, _> = workers.iter().map(|w| (*w, (0usize, 4usize))).collect();
let natural = hasher.compute_replica_set_with_slots("lora-x", &workers, 2, &usage);
assert_eq!(natural.len(), 2);
let prior: Vec<_> = workers
.iter()
.filter(|w| !natural.contains(w))
.copied()
.take(2)
.collect();
assert_eq!(prior.len(), 2);
let sticky =
hasher.compute_replica_set_with_slots_sticky("lora-x", &workers, 2, &usage, &prior);
let sset: HashSet<_> = sticky.iter().copied().collect();
let pset: HashSet<_> = prior.iter().copied().collect();
assert_eq!(
sset, pset,
"sticky must retain the prior (non-full) placement"
);
}
#[test]
fn test_sticky_replaces_full_prior_worker() {
let hasher = RendezvousHasher;
let workers = make_workers(5);
let prior = vec![workers[0], workers[1]];
let mut usage: HashMap<_, _> = workers.iter().map(|w| (*w, (0usize, 4usize))).collect();
usage.insert(workers[0], (4, 4));
let sticky =
hasher.compute_replica_set_with_slots_sticky("lora-x", &workers, 2, &usage, &prior);
assert_eq!(sticky.len(), 2);
assert!(
!sticky.contains(&workers[0]),
"a full prior worker must be dropped"
);
assert!(
sticky.contains(&workers[1]),
"a non-full prior worker must be retained"
);
}
}