use std::collections::{HashMap, HashSet};
use crate::kv_router::protocols::WorkerWithDpRank;
use crate::local_model::runtime_config::ModelRuntimeConfig;
use crate::lora::routing::RendezvousHasher;
use crate::lora::routing::table::LoraRoutingTable;
use crate::lora::state_tracker::LoraStateTracker;
type WorkerId = u64;
#[derive(Clone)]
pub struct LoraFilter {
routing_table: LoraRoutingTable,
state_tracker: LoraStateTracker,
}
impl LoraFilter {
pub fn new(routing_table: LoraRoutingTable, state_tracker: LoraStateTracker) -> Self {
Self {
routing_table,
state_tracker,
}
}
fn bounded_fallback(&self, lora_name: &str, available: &[u64]) -> Vec<u64> {
let loaded = self.state_tracker.get_loaded_workers(lora_name);
if !loaded.is_empty() {
let loaded_ids: HashSet<u64> = loaded.iter().map(|w| w.worker_id).collect();
let live_loaded: Vec<u64> = available
.iter()
.copied()
.filter(|id| loaded_ids.contains(id))
.collect();
if !live_loaded.is_empty() {
tracing::debug!(
lora = lora_name,
count = live_loaded.len(),
"Replica workers unavailable; narrowed to known-loaded live workers"
);
return live_loaded;
}
}
if let Some(pin) = available.iter().copied().max_by(|&a, &b| {
let sa = RendezvousHasher::compute_score(lora_name, WorkerWithDpRank::new(a, 0));
let sb = RendezvousHasher::compute_score(lora_name, WorkerWithDpRank::new(b, 0));
sa.cmp(&sb).then(a.cmp(&b))
}) {
tracing::debug!(
lora = lora_name,
worker_id = pin,
"Replica workers unavailable and adapter not loaded; bounded HRW pin (no scatter)"
);
return vec![pin];
}
Vec::new()
}
pub fn filter_worker_ids_for_lora(
&self,
lora_name: Option<&str>,
available: &[u64],
) -> Vec<u64> {
let Some(lora_name) = lora_name else {
return available.to_vec();
};
let Some(config) = self.routing_table.get_config(lora_name) else {
let loaded = self.state_tracker.get_loaded_workers(lora_name);
if !loaded.is_empty() {
let loaded_ids_set: HashSet<u64> = loaded.iter().map(|w| w.worker_id).collect();
let loaded_ids: Vec<u64> = available
.iter()
.copied()
.filter(|id| loaded_ids_set.contains(id))
.collect();
if !loaded_ids.is_empty() {
tracing::debug!(
lora = lora_name,
count = loaded_ids.len(),
"LoRA not in routing table; narrowed to known-loaded workers"
);
return loaded_ids;
}
}
tracing::debug!(
lora = lora_name,
"LoRA not in routing table and not known-loaded, returning all workers"
);
return available.to_vec();
};
let replica_id_set: HashSet<u64> = config.replica_set.iter().map(|w| w.worker_id).collect();
if config.is_active {
let loaded = self.state_tracker.get_loaded_workers(lora_name);
let loaded_ids: HashSet<u64> = loaded.iter().map(|w| w.worker_id).collect();
let loaded_in_set: Vec<u64> = available
.iter()
.copied()
.filter(|id| replica_id_set.contains(id) && loaded_ids.contains(id))
.collect();
if !loaded_in_set.is_empty() {
tracing::debug!(
lora = lora_name,
count = loaded_in_set.len(),
"Filtered to loaded workers in replica set"
);
return loaded_in_set;
}
let replica_set: Vec<u64> = available
.iter()
.copied()
.filter(|id| replica_id_set.contains(id))
.collect();
if !replica_set.is_empty() {
tracing::debug!(
lora = lora_name,
count = replica_set.len(),
"LoRA not loaded yet, returning full replica set for lazy load"
);
return replica_set;
}
tracing::warn!(
lora = lora_name,
"Replica set workers all unavailable; using bounded fallback (no scatter)"
);
self.bounded_fallback(lora_name, available)
} else {
if let Some(pin_id) = config.replica_set.first().map(|w| w.worker_id)
&& available.contains(&pin_id)
{
tracing::debug!(
lora = lora_name,
worker_id = pin_id,
"Cold-start: routing to HRW-pinned worker"
);
return vec![pin_id];
}
tracing::warn!(
lora = lora_name,
"Cold-start pin worker unavailable; using bounded fallback (no scatter)"
);
self.bounded_fallback(lora_name, available)
}
}
pub fn filter_workers_for_lora(
&self,
lora_name: Option<&str>,
workers: &HashMap<WorkerId, ModelRuntimeConfig>,
) -> HashMap<WorkerId, ModelRuntimeConfig> {
let available_ids: Vec<u64> = workers.keys().copied().collect();
let selected_ids = self.filter_worker_ids_for_lora(lora_name, &available_ids);
if selected_ids.len() == available_ids.len() {
return workers.clone();
}
let selected_set: HashSet<u64> = selected_ids.into_iter().collect();
workers
.iter()
.filter(|(wid, _)| selected_set.contains(wid))
.map(|(k, v)| (*k, v.clone()))
.collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::kv_router::protocols::WorkerWithDpRank;
use crate::lora::routing::table::{LoraReplicaConfig, LoraRoutingTable};
use crate::lora::state_tracker::LoraStateTracker;
use crate::model_card::LoraInfo;
use std::time::Instant;
fn make_workers_map(ids: &[u64]) -> HashMap<WorkerId, ModelRuntimeConfig> {
ids.iter()
.map(|&id| (id, ModelRuntimeConfig::default()))
.collect()
}
fn make_worker(id: u64) -> WorkerWithDpRank {
WorkerWithDpRank::new(id, 0)
}
fn make_lora_info(name: &str) -> LoraInfo {
LoraInfo {
name: name.to_string(),
max_gpu_lora_count: Some(4),
}
}
#[test]
fn test_no_lora_returns_all_workers() {
let rt = LoraRoutingTable::new();
let st = LoraStateTracker::new();
let filter = LoraFilter::new(rt, st);
let workers = make_workers_map(&[1, 2, 3]);
let result = filter.filter_workers_for_lora(None, &workers);
assert_eq!(result.len(), 3);
}
#[test]
fn test_not_in_routing_table_returns_all() {
let rt = LoraRoutingTable::new();
let st = LoraStateTracker::new();
let filter = LoraFilter::new(rt, st);
let workers = make_workers_map(&[1, 2, 3]);
let result = filter.filter_workers_for_lora(Some("unknown-lora"), &workers);
assert_eq!(result.len(), 3);
}
#[test]
fn test_not_in_routing_table_narrows_to_loaded_workers() {
let rt = LoraRoutingTable::new();
let st = LoraStateTracker::new();
st.handle_mdc_addition(make_worker(2), &make_lora_info("lora-a"));
let filter = LoraFilter::new(rt, st);
let workers = make_workers_map(&[1, 2, 3]);
let result = filter.filter_workers_for_lora(Some("lora-a"), &workers);
assert_eq!(result.len(), 1);
assert!(result.contains_key(&2));
}
#[test]
fn test_not_in_routing_table_loaded_worker_unavailable_falls_back_to_all() {
let rt = LoraRoutingTable::new();
let st = LoraStateTracker::new();
st.handle_mdc_addition(make_worker(9), &make_lora_info("lora-a"));
let filter = LoraFilter::new(rt, st);
let workers = make_workers_map(&[1, 2, 3]);
let result = filter.filter_workers_for_lora(Some("lora-a"), &workers);
assert_eq!(result.len(), 3);
}
#[test]
fn test_active_lora_filters_to_loaded_workers() {
let rt = LoraRoutingTable::new();
let st = LoraStateTracker::new();
rt.update_allocation(
"lora-a".to_string(),
LoraReplicaConfig {
lora_name: "lora-a".to_string(),
replica_factor: 2,
replica_set: vec![make_worker(1), make_worker(2)],
updated_at: Instant::now(),
is_active: true,
},
);
st.handle_mdc_addition(make_worker(1), &make_lora_info("lora-a"));
let filter = LoraFilter::new(rt, st);
let workers = make_workers_map(&[1, 2, 3]);
let result = filter.filter_workers_for_lora(Some("lora-a"), &workers);
assert_eq!(result.len(), 1);
assert!(result.contains_key(&1));
}
#[test]
fn test_active_lora_falls_back_to_replica_set() {
let rt = LoraRoutingTable::new();
let st = LoraStateTracker::new();
rt.update_allocation(
"lora-a".to_string(),
LoraReplicaConfig {
lora_name: "lora-a".to_string(),
replica_factor: 2,
replica_set: vec![make_worker(1), make_worker(2)],
updated_at: Instant::now(),
is_active: true,
},
);
let filter = LoraFilter::new(rt, st);
let workers = make_workers_map(&[1, 2, 3]);
let result = filter.filter_workers_for_lora(Some("lora-a"), &workers);
assert_eq!(result.len(), 2);
assert!(result.contains_key(&1));
assert!(result.contains_key(&2));
}
#[test]
fn test_inactive_lora_cold_start_pin() {
let rt = LoraRoutingTable::new();
let st = LoraStateTracker::new();
rt.update_allocation(
"lora-b".to_string(),
LoraReplicaConfig {
lora_name: "lora-b".to_string(),
replica_factor: 1,
replica_set: vec![make_worker(2)],
updated_at: Instant::now(),
is_active: false,
},
);
let filter = LoraFilter::new(rt, st);
let workers = make_workers_map(&[1, 2, 3]);
let result = filter.filter_workers_for_lora(Some("lora-b"), &workers);
assert_eq!(result.len(), 1);
assert!(result.contains_key(&2));
}
#[test]
fn test_inactive_pin_worker_unavailable_uses_bounded_pin() {
let rt = LoraRoutingTable::new();
let st = LoraStateTracker::new();
rt.update_allocation(
"lora-b".to_string(),
LoraReplicaConfig {
lora_name: "lora-b".to_string(),
replica_factor: 1,
replica_set: vec![make_worker(5)],
updated_at: Instant::now(),
is_active: false,
},
);
let filter = LoraFilter::new(rt, st);
let workers = make_workers_map(&[1, 2, 3]);
let result = filter.filter_workers_for_lora(Some("lora-b"), &workers);
assert_eq!(
result.len(),
1,
"must bound to one worker, not scatter to all three"
);
let again = filter.filter_workers_for_lora(Some("lora-b"), &workers);
assert_eq!(
result.keys().collect::<Vec<_>>(),
again.keys().collect::<Vec<_>>(),
"bounded HRW pin must be deterministic"
);
}
#[test]
fn test_active_all_replicas_unavailable_prefers_known_loaded() {
let rt = LoraRoutingTable::new();
let st = LoraStateTracker::new();
rt.update_allocation(
"lora-a".to_string(),
LoraReplicaConfig {
lora_name: "lora-a".to_string(),
replica_factor: 2,
replica_set: vec![make_worker(8), make_worker(9)], updated_at: Instant::now(),
is_active: true,
},
);
st.handle_mdc_addition(make_worker(2), &make_lora_info("lora-a"));
let filter = LoraFilter::new(rt, st);
let workers = make_workers_map(&[1, 2, 3]);
let result = filter.filter_workers_for_lora(Some("lora-a"), &workers);
assert_eq!(
result.len(),
1,
"must narrow to the known-loaded worker, not scatter"
);
assert!(result.contains_key(&2));
}
#[test]
fn test_active_all_replicas_unavailable_not_loaded_uses_bounded_pin() {
let rt = LoraRoutingTable::new();
let st = LoraStateTracker::new();
rt.update_allocation(
"lora-a".to_string(),
LoraReplicaConfig {
lora_name: "lora-a".to_string(),
replica_factor: 2,
replica_set: vec![make_worker(8), make_worker(9)],
updated_at: Instant::now(),
is_active: true,
},
);
let filter = LoraFilter::new(rt, st);
let workers = make_workers_map(&[1, 2, 3]);
let result = filter.filter_workers_for_lora(Some("lora-a"), &workers);
assert_eq!(
result.len(),
1,
"must bound to one worker, not scatter to all three"
);
}
}