use std::sync::atomic::{AtomicUsize, Ordering};
use serde::{Deserialize, Serialize};
use super::credits::CreditManager;
use super::worker::{Worker, WorkerRegistry};
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize, Default)]
#[serde(rename_all = "snake_case")]
pub enum RoutingStrategy {
#[default]
BestPrice,
BestLatency,
BestAvailability,
RoundRobin,
Random,
WeightedCapacity,
}
impl RoutingStrategy {
pub fn as_str(&self) -> &'static str {
match self {
RoutingStrategy::BestPrice => "best_price",
RoutingStrategy::BestLatency => "best_latency",
RoutingStrategy::BestAvailability => "best_availability",
RoutingStrategy::RoundRobin => "round_robin",
RoutingStrategy::Random => "random",
RoutingStrategy::WeightedCapacity => "weighted_capacity",
}
}
pub fn description(&self) -> &'static str {
match self {
RoutingStrategy::BestPrice => "Lowest cost per compute",
RoutingStrategy::BestLatency => "Fastest response time",
RoutingStrategy::BestAvailability => "Most available resources",
RoutingStrategy::RoundRobin => "Even distribution",
RoutingStrategy::Random => "Random selection",
RoutingStrategy::WeightedCapacity => "Weighted by capacity",
}
}
}
impl std::str::FromStr for RoutingStrategy {
type Err = String;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s.to_lowercase().as_str() {
"best_price" | "price" | "cheap" | "cheapest" => Ok(RoutingStrategy::BestPrice),
"best_latency" | "latency" | "fast" | "fastest" => Ok(RoutingStrategy::BestLatency),
"best_availability" | "availability" | "available" => {
Ok(RoutingStrategy::BestAvailability)
}
"round_robin" | "robin" | "rr" => Ok(RoutingStrategy::RoundRobin),
"random" | "rand" => Ok(RoutingStrategy::Random),
"weighted_capacity" | "weighted" | "capacity" => Ok(RoutingStrategy::WeightedCapacity),
_ => Err(format!("Unknown routing strategy: {}", s)),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ResourceRequirements {
#[serde(default = "default_cpus")]
pub cpus: f64,
#[serde(default = "default_memory")]
pub memory_bytes: u64,
#[serde(default)]
pub gpus: u32,
#[serde(default)]
pub worker_type: Option<String>,
#[serde(default)]
pub tags: Vec<String>,
#[serde(default = "default_duration")]
pub estimated_duration_secs: f64,
#[serde(default)]
pub strategy: RoutingStrategy,
#[serde(default)]
pub timeout_secs: f64,
#[serde(default)]
pub remote_only: bool,
#[serde(default)]
pub budget_credits: Option<f64>,
#[serde(default)]
pub target_worker: Option<String>,
#[serde(default)]
pub target_node: Option<String>,
}
fn default_cpus() -> f64 {
1.0
}
fn default_memory() -> u64 {
1024 * 1024 * 1024
} fn default_duration() -> f64 {
1.0
}
impl Default for ResourceRequirements {
fn default() -> Self {
Self {
cpus: 1.0,
memory_bytes: 1024 * 1024 * 1024,
gpus: 0,
worker_type: None,
tags: Vec::new(),
estimated_duration_secs: 1.0,
strategy: RoutingStrategy::default(),
timeout_secs: 0.0, remote_only: false,
budget_credits: None,
target_worker: None,
target_node: None,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RoutingDecision {
pub worker: Worker,
pub estimated_cost: f64,
pub reason: String,
pub alternatives_count: usize,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RoutingError {
pub code: String,
pub message: String,
}
impl RoutingError {
pub fn no_workers() -> Self {
Self {
code: "NO_WORKERS".to_string(),
message: "No workers available".to_string(),
}
}
pub fn no_capacity(requirements: &ResourceRequirements) -> Self {
Self {
code: "NO_CAPACITY".to_string(),
message: format!(
"No worker has capacity for {} CPUs, {} bytes memory, {} GPUs",
requirements.cpus, requirements.memory_bytes, requirements.gpus
),
}
}
pub fn insufficient_credits(required: f64, available: f64) -> Self {
Self {
code: "INSUFFICIENT_CREDITS".to_string(),
message: format!(
"Insufficient credits: need {:.4}, have {:.4}",
required, available
),
}
}
pub fn rate_limited() -> Self {
Self {
code: "RATE_LIMITED".to_string(),
message: "Rate limit exceeded".to_string(),
}
}
pub fn worker_type_unavailable(worker_type: &str) -> Self {
Self {
code: "WORKER_TYPE_UNAVAILABLE".to_string(),
message: format!("No workers of type '{}' available", worker_type),
}
}
pub fn timeout_incompatible(timeout: f64) -> Self {
Self {
code: "TIMEOUT_INCOMPATIBLE".to_string(),
message: format!(
"No worker accepts timeout of {:.1}s. Reduce timeout or wait for a compatible worker.",
timeout
),
}
}
pub fn quota_exceeded() -> Self {
Self {
code: "QUOTA_EXCEEDED".to_string(),
message: "All workers have exceeded their request quota for this time window"
.to_string(),
}
}
}
pub struct Router {
round_robin_counter: AtomicUsize,
}
fn expected_completion_ms(w: &Worker) -> f64 {
w.avg_latency_ms * (1.0 + w.active_requests as f64)
}
impl Router {
pub fn new() -> Self {
Self {
round_robin_counter: AtomicUsize::new(0),
}
}
pub fn available_strategies() -> Vec<RoutingStrategy> {
vec![
RoutingStrategy::BestPrice,
RoutingStrategy::BestLatency,
RoutingStrategy::BestAvailability,
RoutingStrategy::RoundRobin,
RoutingStrategy::Random,
RoutingStrategy::WeightedCapacity,
]
}
fn pin_local(
&self,
_registry: &WorkerRegistry,
requirements: &ResourceRequirements,
) -> Option<Result<RoutingDecision, RoutingError>> {
let norm = |s: &str| s.strip_prefix("zc://").unwrap_or(s).to_string();
let bare_node = |s: &str| {
let n = norm(s);
n.strip_prefix("node-").unwrap_or(&n).to_string()
};
let tn = requirements.target_node.as_deref()?;
let ours = super::node_name_or_default();
if bare_node(tn) != bare_node(&ours) {
return Some(Err(RoutingError::no_workers()));
}
None
}
pub fn select_worker_no_checks(
&self,
registry: &WorkerRegistry,
requirements: &ResourceRequirements,
) -> Result<RoutingDecision, RoutingError> {
if let Some(pinned) = self.pin_local(registry, requirements) {
return pinned;
}
let candidates = registry.healthy();
if candidates.is_empty() {
return Err(RoutingError::no_workers());
}
let n = candidates.len();
let (chosen, reason) = self.select_by_strategy(&candidates, requirements);
let tied = candidates.iter().all(|w| {
self.strategy_key(w, requirements) == self.strategy_key(&chosen, requirements)
});
let (chosen, reason) = if tied && n > 1 {
let idx = self.round_robin_counter.fetch_add(1, Ordering::Relaxed) % n;
(
candidates[idx].clone(),
format!("{reason} (tied; spread to slot {idx})"),
)
} else {
(chosen, reason)
};
if registry.try_reserve_quota(&chosen.id) {
return Ok(RoutingDecision {
estimated_cost: 0.0,
worker: chosen,
reason,
alternatives_count: n - 1,
});
}
for w in candidates.iter() {
if w.id != chosen.id && registry.try_reserve_quota(&w.id) {
return Ok(RoutingDecision {
estimated_cost: 0.0,
worker: w.clone(),
reason: format!("{reason} (first choice was full)"),
alternatives_count: n - 1,
});
}
}
Err(RoutingError::no_workers())
}
pub fn select_adaptive(
&self,
registry: &WorkerRegistry,
requirements: &ResourceRequirements,
rng: &mut crate::broker::sampler::Rng,
weights: &crate::broker::policy::Weights,
) -> Result<RoutingDecision, RoutingError> {
use crate::broker::policy::{self, Candidate};
use crate::broker::sampler::AliasTable;
if let Some(pinned) = self.pin_local(registry, requirements) {
return pinned;
}
let workers = registry.healthy();
if workers.is_empty() {
return Err(RoutingError::no_workers());
}
let candidates: Vec<Candidate> = workers
.iter()
.map(|w| Candidate {
stats: w.stats,
in_flight: w.active_requests,
price_per_hour: w.pricing.price_per_hour,
})
.collect();
let table_weights: Vec<f64> = candidates
.iter()
.map(|c| policy::table_weight(c, weights))
.collect();
let Some(table) = AliasTable::build(&table_weights) else {
return Err(RoutingError::no_workers());
};
let Some(pick) = policy::choose(&candidates, &table, rng, 2, weights) else {
return Err(RoutingError::no_workers());
};
let n = workers.len();
let order = std::iter::once(pick).chain((0..n).filter(|i| *i != pick));
for i in order {
if registry.try_reserve_quota(&workers[i].id) {
let est = candidates[i]
.stats
.expected_service_ms()
.unwrap_or(weights.unmeasured_ms);
return Ok(RoutingDecision {
estimated_cost: 0.0,
reason: match i == pick {
true => format!("Adaptive: expected {est:.0}ms"),
false => format!("Adaptive: expected {est:.0}ms (first choice was full)"),
},
worker: workers[i].clone(),
alternatives_count: n - 1,
});
}
}
Err(RoutingError::no_workers())
}
fn strategy_key(&self, w: &Worker, requirements: &ResourceRequirements) -> u64 {
let v = match requirements.strategy {
RoutingStrategy::BestPrice => w
.pricing
.estimate_cost(requirements.estimated_duration_secs),
RoutingStrategy::BestLatency => match w.total_requests > 0 {
true => expected_completion_ms(w),
false => f64::INFINITY,
},
RoutingStrategy::BestAvailability => {
-(w.resources.cpus_available
+ w.resources.memory_available as f64 / (1024.0 * 1024.0 * 1024.0)
- w.active_requests as f64 * 0.5)
}
RoutingStrategy::RoundRobin
| RoutingStrategy::Random
| RoutingStrategy::WeightedCapacity => return 0,
};
(v * 1e6) as u64
}
pub fn select_local_worker(
&self,
registry: &WorkerRegistry,
own_ip: Option<&str>,
_requirements: &ResourceRequirements,
) -> Result<RoutingDecision, RoutingError> {
let all = registry.healthy();
let local: Vec<Worker> = all
.into_iter()
.filter(|w| {
w.uri.contains("127.0.0.1")
|| w.uri.contains("localhost")
|| own_ip.is_some_and(|ip| w.uri.contains(ip))
|| (w.explicit_local && w.source_node.is_none())
})
.collect();
if local.is_empty() {
return Err(RoutingError::no_workers());
}
let n = local.len();
let start = self.round_robin_counter.fetch_add(1, Ordering::Relaxed);
for i in 0..n {
let worker = local[(start + i) % n].clone();
if registry.try_reserve_quota(&worker.id) {
return Ok(RoutingDecision {
worker,
estimated_cost: 0.0,
reason: format!("Local worker (publish-lock): slot {}", (start + i) % n),
alternatives_count: n - 1,
});
}
}
Err(RoutingError::quota_exceeded())
}
pub fn select_worker(
&self,
registry: &WorkerRegistry,
credits: &CreditManager,
user_id: &str,
balance: f64,
requirements: &ResourceRequirements,
) -> Result<RoutingDecision, RoutingError> {
if !credits.check_rate_limit(user_id) {
return Err(RoutingError::rate_limited());
}
if let Some(pinned) = self.pin_local(registry, requirements) {
return pinned;
}
let workers = registry.healthy();
if workers.is_empty() {
return Err(RoutingError::no_workers());
}
let workers: Vec<Worker> = if let Some(ref wt) = requirements.worker_type {
let filtered: Vec<Worker> = workers
.into_iter()
.filter(|w| &w.worker_type == wt)
.collect();
if filtered.is_empty() {
return Err(RoutingError::worker_type_unavailable(wt));
}
filtered
} else {
workers
};
let workers: Vec<Worker> = if !requirements.tags.is_empty() {
workers
.into_iter()
.filter(|w| requirements.tags.iter().all(|t| w.tags.contains(t)))
.collect()
} else {
workers
};
let capable: Vec<Worker> = workers
.iter()
.filter(|w| {
w.can_handle(
requirements.cpus,
requirements.memory_bytes,
requirements.gpus,
)
})
.cloned()
.collect();
if capable.is_empty() {
return Err(RoutingError::no_capacity(requirements));
}
let capable: Vec<Worker> = if requirements.timeout_secs > 0.0 {
let filtered: Vec<Worker> = capable
.into_iter()
.filter(|w| {
w.max_timeout_secs <= 0.0 || w.max_timeout_secs >= requirements.timeout_secs
})
.collect();
if filtered.is_empty() {
return Err(RoutingError::timeout_incompatible(
requirements.timeout_secs,
));
}
filtered
} else {
capable
};
let mut candidates = capable;
loop {
if candidates.is_empty() {
return Err(RoutingError::quota_exceeded());
}
let alternatives_count = candidates.len();
let (best_worker, strategy_reason) = self.select_by_strategy(&candidates, requirements);
let estimated_cost = best_worker
.pricing
.estimate_cost(requirements.estimated_duration_secs);
if balance < estimated_cost {
return Err(RoutingError::insufficient_credits(estimated_cost, balance));
}
if registry.try_reserve_quota(&best_worker.id) {
let reason = format!(
"{} (checked {} workers)",
strategy_reason, alternatives_count
);
return Ok(RoutingDecision {
worker: best_worker,
estimated_cost,
reason,
alternatives_count,
});
}
let loser_id = best_worker.id.clone();
candidates.retain(|w| w.id != loser_id);
}
}
fn select_by_strategy(
&self,
workers: &[Worker],
requirements: &ResourceRequirements,
) -> (Worker, String) {
match requirements.strategy {
RoutingStrategy::BestPrice => {
let (idx, cost) = workers
.iter()
.enumerate()
.map(|(i, w)| {
(
i,
w.pricing
.estimate_cost(requirements.estimated_duration_secs),
)
})
.min_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal))
.expect("select_by_strategy is only called with a non-empty worker list");
(
workers[idx].clone(),
format!("Best price: {:.6} credits", cost),
)
}
RoutingStrategy::BestLatency => {
let (idx, latency) = workers
.iter()
.enumerate()
.map(|(i, w)| {
let measured = w.total_requests > 0;
(i, measured, expected_completion_ms(w))
})
.min_by(|a, b| {
b.1.cmp(&a.1).then_with(|| {
a.2.partial_cmp(&b.2).unwrap_or(std::cmp::Ordering::Equal)
})
})
.map(|(i, measured, latency)| (i, measured.then_some(latency)))
.expect("select_by_strategy is only called with a non-empty worker list");
(
workers[idx].clone(),
match latency {
Some(ms) => format!("Best latency: {:.1}ms", ms),
None => "Best latency: no measurement yet".to_string(),
},
)
}
RoutingStrategy::BestAvailability => {
let (idx, score) = workers
.iter()
.enumerate()
.map(|(i, w)| {
let cpu_score = w.resources.cpus_available;
let mem_score =
w.resources.memory_available as f64 / (1024.0 * 1024.0 * 1024.0);
let load_penalty = w.active_requests as f64 * 0.5;
(i, cpu_score + mem_score - load_penalty)
})
.min_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal))
.expect("select_by_strategy is only called with a non-empty worker list");
(
workers[idx].clone(),
format!("Best availability: score {:.1}", score),
)
}
RoutingStrategy::RoundRobin => {
let idx = self.round_robin_counter.fetch_add(1, Ordering::Relaxed);
let worker = workers[idx % workers.len()].clone();
(
worker,
format!("Round-robin: index {}", idx % workers.len()),
)
}
RoutingStrategy::Random => {
use std::time::{SystemTime, UNIX_EPOCH};
let seed = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.subsec_nanos() as usize;
let idx = seed % workers.len();
let worker = workers[idx].clone();
(worker, format!("Random: index {}", idx))
}
RoutingStrategy::WeightedCapacity => {
use std::time::{SystemTime, UNIX_EPOCH};
let weights: Vec<f64> = workers
.iter()
.map(|w| {
let cpu_weight = w.resources.cpus_available.max(0.1);
let mem_weight = (w.resources.memory_available as f64
/ (1024.0 * 1024.0 * 1024.0))
.max(0.1);
cpu_weight * mem_weight
})
.collect();
let total_weight: f64 = weights.iter().sum();
let seed = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.subsec_nanos() as f64;
let target = (seed / 1_000_000_000.0) * total_weight;
let mut cumulative = 0.0;
for (i, weight) in weights.iter().enumerate() {
cumulative += weight;
if cumulative >= target {
let worker = workers[i].clone();
return (worker, format!("Weighted capacity: weight {:.2}", weight));
}
}
let worker = workers[0].clone();
(worker, "Weighted capacity: fallback".to_string())
}
}
}
pub fn estimate_cost(
&self,
registry: &WorkerRegistry,
requirements: &ResourceRequirements,
) -> Option<(f64, f64)> {
let workers = registry.healthy();
if workers.is_empty() {
return None;
}
let capable: Vec<&Worker> = workers
.iter()
.filter(|w| {
w.can_handle(
requirements.cpus,
requirements.memory_bytes,
requirements.gpus,
)
})
.collect();
if capable.is_empty() {
return None;
}
let costs: Vec<f64> = capable
.iter()
.map(|w| {
w.pricing
.estimate_cost(requirements.estimated_duration_secs)
})
.collect();
let min_cost = costs.iter().cloned().fold(f64::INFINITY, f64::min);
let max_cost = costs.iter().cloned().fold(0.0, f64::max);
Some((min_cost, max_cost))
}
pub fn list_matching(
&self,
registry: &WorkerRegistry,
requirements: &ResourceRequirements,
) -> Vec<(Worker, f64)> {
let workers = registry.healthy();
let mut matching: Vec<(Worker, f64)> = workers
.iter()
.filter(|w| {
if let Some(ref wt) = requirements.worker_type {
if &w.worker_type != wt {
return false;
}
}
w.can_handle(
requirements.cpus,
requirements.memory_bytes,
requirements.gpus,
)
})
.map(|w| {
let cost = w
.pricing
.estimate_cost(requirements.estimated_duration_secs);
(w.clone(), cost)
})
.collect();
matching.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal));
matching
}
}
impl Default for Router {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::broker::credits::CreditManager;
use crate::broker::worker::{
HardwareInfo, WorkerPricing, WorkerRegistration, WorkerRegistry, WorkerResources,
};
pub(super) fn make_reg(
name: &str,
port: usize,
cpus: f64,
memory_gib: u64,
gpus: u32,
price_per_hour: f64,
) -> WorkerRegistration {
WorkerRegistration {
name: name.to_string(),
uri: format!("http://127.0.0.1:{}", 3960 + port),
worker_type: "zakuro".to_string(),
resources: WorkerResources {
cpus_total: cpus,
cpus_available: cpus,
memory_total: memory_gib * 1024 * 1024 * 1024,
memory_available: memory_gib * 1024 * 1024 * 1024,
gpus_total: gpus,
gpus_available: gpus,
},
pricing: WorkerPricing {
price_per_hour,
min_charge: 0.001,
},
tags: vec![],
max_timeout_secs: 0.0,
hardware: HardwareInfo::default(),
wireguard_ip: None,
is_docker: None,
source_node: None,
explicit_local: false,
provider_type: Default::default(),
served_models: vec![],
price_per_mtok: 0.0,
}
}
fn credits_for(user: &str, balance: f64) -> CreditManager {
let mgr = CreditManager::new();
mgr.get_or_create(user, balance);
mgr
}
pub(super) fn req_default() -> ResourceRequirements {
ResourceRequirements::default()
}
#[test]
fn local_selection_follows_ownership_not_uri_spelling() {
let router = Router::new();
let registry = WorkerRegistry::new();
let mut reg = make_reg("worker-sidecar", 0, 4.0, 4, 0, 0.001);
reg.uri = "zc://worker-7fed13b4a701206f-3960".to_string();
reg.source_node = None; reg.explicit_local = true; registry.register(reg);
let decision = router
.select_local_worker(®istry, None, &req_default())
.expect("an owned worker must be selectable as local");
assert_eq!(decision.worker.name, "worker-sidecar");
assert_eq!(decision.estimated_cost, 0.0);
}
#[test]
fn peer_synced_workers_are_never_selected_as_local() {
let router = Router::new();
let registry = WorkerRegistry::new();
let mut reg = make_reg("worker-remote", 1, 4.0, 4, 0, 0.001);
reg.uri = "zc://worker-remote-3960".to_string();
reg.explicit_local = true; reg.source_node = Some("node-i9".to_string());
registry.register(reg);
assert!(router
.select_local_worker(®istry, None, &req_default())
.is_err());
}
#[test]
fn test_no_healthy_workers_returns_no_workers() {
let router = Router::new();
let registry = WorkerRegistry::new();
let credits = credits_for("alice", 1000.0);
let err = router
.select_worker(®istry, &credits, "alice", 1000.0, &req_default())
.unwrap_err();
assert_eq!(err.code, "NO_WORKERS");
}
#[test]
fn test_client_target_worker_is_noop_for_routing() {
let router = Router::new();
let registry = WorkerRegistry::new();
registry.register(make_reg("worker-a", 0, 4.0, 4, 0, 0.001));
let req = ResourceRequirements {
target_worker: Some("zc://worker-ghost".to_string()),
..req_default()
};
let decision = router.select_worker_no_checks(®istry, &req).unwrap();
assert_eq!(decision.worker.name, "worker-a");
}
#[test]
fn test_pin_target_node_other_node_rejected() {
let router = Router::new();
let registry = WorkerRegistry::new();
registry.register(make_reg("worker-a", 0, 4.0, 4, 0, 0.001));
let req = ResourceRequirements {
target_node: Some("zc://node-somewhere-else".to_string()),
..req_default()
};
let err = router.select_worker_no_checks(®istry, &req).unwrap_err();
assert_eq!(err.code, "NO_WORKERS");
}
#[test]
fn test_no_capacity_returns_no_capacity() {
let router = Router::new();
let registry = WorkerRegistry::new();
registry.register(make_reg("small", 0, 1.0, 1, 0, 0.001));
let credits = credits_for("alice", 1000.0);
let req = ResourceRequirements {
cpus: 16.0,
..req_default()
};
let err = router
.select_worker(®istry, &credits, "alice", 1000.0, &req)
.unwrap_err();
assert_eq!(err.code, "NO_CAPACITY");
}
#[test]
fn test_insufficient_credits_returns_error() {
let router = Router::new();
let registry = WorkerRegistry::new();
registry.register(make_reg("expensive", 0, 4.0, 4, 0, 1.0));
let credits = credits_for("alice", 0.001);
let req = ResourceRequirements {
cpus: 1.0,
estimated_duration_secs: 3600.0, ..req_default()
};
let err = router
.select_worker(®istry, &credits, "alice", 0.001, &req)
.unwrap_err();
assert_eq!(err.code, "INSUFFICIENT_CREDITS");
}
#[test]
fn test_rate_limited_returns_error() {
let router = Router::new();
let registry = WorkerRegistry::new();
registry.register(make_reg("w", 0, 4.0, 4, 0, 0.001));
let mgr = CreditManager::new();
mgr.get_or_create("alice", 1000.0);
mgr.set_rate_limits("alice", Some(0), None, None);
let err = router
.select_worker(®istry, &mgr, "alice", 1000.0, &req_default())
.unwrap_err();
assert_eq!(err.code, "RATE_LIMITED");
}
#[test]
fn test_worker_type_mismatch_returns_error() {
let router = Router::new();
let registry = WorkerRegistry::new();
let mut reg = make_reg("ray-w", 0, 4.0, 4, 0, 0.001);
reg.worker_type = "ray".to_string();
registry.register(reg);
let credits = credits_for("alice", 1000.0);
let req = ResourceRequirements {
worker_type: Some("spark".to_string()),
..req_default()
};
let err = router
.select_worker(®istry, &credits, "alice", 1000.0, &req)
.unwrap_err();
assert_eq!(err.code, "WORKER_TYPE_UNAVAILABLE");
}
#[test]
fn test_timeout_incompatible_returns_error() {
let router = Router::new();
let registry = WorkerRegistry::new();
let mut reg = make_reg("limited", 0, 4.0, 4, 0, 0.001);
reg.max_timeout_secs = 30.0; registry.register(reg);
let credits = credits_for("alice", 1000.0);
let req = ResourceRequirements {
timeout_secs: 60.0,
..req_default()
};
let err = router
.select_worker(®istry, &credits, "alice", 1000.0, &req)
.unwrap_err();
assert_eq!(err.code, "TIMEOUT_INCOMPATIBLE");
}
#[test]
fn test_timeout_unlimited_worker_accepts_any_timeout() {
let router = Router::new();
let registry = WorkerRegistry::new();
let mut reg = make_reg("unlimited", 0, 4.0, 4, 0, 0.001);
reg.max_timeout_secs = 0.0; registry.register(reg);
let credits = credits_for("alice", 1000.0);
let req = ResourceRequirements {
timeout_secs: 3600.0,
..req_default()
};
assert!(router
.select_worker(®istry, &credits, "alice", 1000.0, &req)
.is_ok());
}
#[test]
fn test_tag_filter_requires_all_tags() {
let router = Router::new();
let registry = WorkerRegistry::new();
let mut reg = make_reg("gpu-w", 0, 8.0, 32, 2, 0.001);
reg.tags = vec!["gpu".to_string(), "a100".to_string()];
registry.register(reg);
let credits = credits_for("alice", 1000.0);
let req_match = ResourceRequirements {
tags: vec!["gpu".to_string(), "a100".to_string()],
..req_default()
};
assert!(router
.select_worker(®istry, &credits, "alice", 1000.0, &req_match)
.is_ok());
let req_no_match = ResourceRequirements {
tags: vec!["gpu".to_string(), "h100".to_string()],
..req_default()
};
assert!(router
.select_worker(®istry, &credits, "alice", 1000.0, &req_no_match)
.is_err());
}
#[test]
fn test_best_price_selects_cheapest_worker() {
let router = Router::new();
let registry = WorkerRegistry::new();
registry.register(make_reg("cheap", 0, 4.0, 4, 0, 0.001));
registry.register(make_reg("pricey", 1, 4.0, 4, 0, 1.0));
let credits = credits_for("alice", 10000.0);
let req = ResourceRequirements {
strategy: RoutingStrategy::BestPrice,
estimated_duration_secs: 3600.0,
..req_default()
};
let decision = router
.select_worker(®istry, &credits, "alice", 10000.0, &req)
.unwrap();
assert_eq!(decision.worker.name, "cheap");
}
#[test]
fn best_price_selects_cheapest_with_many_candidates() {
let router = Router::new();
let registry = WorkerRegistry::new();
registry.register(make_reg("w-a", 0, 4.0, 4, 0, 5.0));
registry.register(make_reg("w-b", 1, 4.0, 4, 0, 2.0)); registry.register(make_reg("w-c", 2, 4.0, 4, 0, 9.0));
let credits = credits_for("alice", 10000.0);
let req = ResourceRequirements {
strategy: RoutingStrategy::BestPrice,
estimated_duration_secs: 3600.0,
..req_default()
};
let decision = router
.select_worker(®istry, &credits, "alice", 10000.0, &req)
.unwrap();
assert_eq!(decision.worker.name, "w-b", "must pick the cheapest worker");
}
#[test]
fn test_best_latency_selects_fastest_worker() {
let router = Router::new();
let registry = WorkerRegistry::new();
let w_fast = registry.register(make_reg("fast", 0, 4.0, 4, 0, 0.001));
let w_slow = registry.register(make_reg("slow", 1, 4.0, 4, 0, 0.001));
registry.record_request(&w_fast.id, 5.0, true);
registry.record_request(&w_slow.id, 500.0, true);
let credits = credits_for("alice", 1000.0);
let req = ResourceRequirements {
strategy: RoutingStrategy::BestLatency,
..req_default()
};
let decision = router
.select_worker(®istry, &credits, "alice", 1000.0, &req)
.unwrap();
assert_eq!(decision.worker.name, "fast");
}
#[test]
fn best_latency_does_not_mistake_never_measured_for_instant() {
let router = Router::new();
let registry = WorkerRegistry::new();
let measured = registry.register(make_reg("measured", 0, 4.0, 4, 0, 0.001));
registry.register(make_reg("never-answered", 1, 4.0, 4, 0, 0.001));
registry.record_request(&measured.id, 50.0, true);
let credits = credits_for("alice", 1000.0);
let req = ResourceRequirements {
strategy: RoutingStrategy::BestLatency,
..req_default()
};
let decision = router
.select_worker(®istry, &credits, "alice", 1000.0, &req)
.unwrap();
assert_eq!(
decision.worker.name, "measured",
"a measured 5 ms beats an unknown"
);
}
#[test]
fn best_latency_still_routes_when_nothing_has_been_measured() {
let router = Router::new();
let registry = WorkerRegistry::new();
registry.register(make_reg("a", 0, 4.0, 4, 0, 0.001));
registry.register(make_reg("b", 1, 4.0, 4, 0, 0.001));
let credits = credits_for("alice", 1000.0);
let req = ResourceRequirements {
strategy: RoutingStrategy::BestLatency,
..req_default()
};
let decision = router
.select_worker(®istry, &credits, "alice", 1000.0, &req)
.unwrap();
assert!(
["a", "b"].contains(&decision.worker.name.as_str()),
"picked {}",
decision.worker.name
);
}
#[test]
fn the_unchecked_path_honours_the_strategy_it_was_given() {
let router = Router::new();
let registry = WorkerRegistry::new();
let fast = registry.register(make_reg("fast", 0, 4.0, 4, 0, 0.001));
let slow = registry.register(make_reg("slow", 1, 4.0, 4, 0, 0.001));
registry.record_request(&fast.id, 5.0, true);
registry.record_request(&slow.id, 900.0, true);
let req = ResourceRequirements {
strategy: RoutingStrategy::BestLatency,
..req_default()
};
for _ in 0..8 {
let d = router.select_worker_no_checks(®istry, &req).unwrap();
assert_eq!(d.worker.name, "fast", "reason was: {}", d.reason);
}
}
#[test]
fn workers_a_strategy_cannot_separate_still_share_the_load() {
let router = Router::new();
let registry = WorkerRegistry::new();
for i in 0..4 {
registry.register(make_reg(&format!("w{i}"), i, 4.0, 4, 0, 0.0));
}
let req = ResourceRequirements {
strategy: RoutingStrategy::BestPrice,
..req_default()
};
let mut seen = std::collections::HashSet::new();
for _ in 0..16 {
seen.insert(
router
.select_worker_no_checks(®istry, &req)
.unwrap()
.worker
.name,
);
}
assert_eq!(seen.len(), 4, "every worker took some of it, got {seen:?}");
}
#[test]
fn best_latency_accounts_for_what_is_already_queued() {
let router = Router::new();
let registry = WorkerRegistry::new();
let busy = registry.register(make_reg("busy-fast", 0, 4.0, 4, 0, 0.001));
let idle = registry.register(make_reg("idle-slower", 1, 4.0, 4, 0, 0.001));
registry.record_request(&busy.id, 10.0, true);
registry.record_request(&idle.id, 30.0, true);
for _ in 0..4 {
registry.increment_active(&busy.id);
}
let d = router
.select_worker_no_checks(
®istry,
&ResourceRequirements {
strategy: RoutingStrategy::BestLatency,
..req_default()
},
)
.unwrap();
assert_eq!(
d.worker.name, "idle-slower",
"an idle 30 ms worker beats a 10 ms one with four queued: {}",
d.reason
);
}
#[test]
fn the_decision_says_how_it_was_made() {
let router = Router::new();
let registry = WorkerRegistry::new();
let fast = registry.register(make_reg("fast", 0, 4.0, 4, 0, 0.001));
registry.register(make_reg("slow", 1, 4.0, 4, 0, 0.002));
registry.record_request(&fast.id, 5.0, true);
let latency = router
.select_worker_no_checks(
®istry,
&ResourceRequirements {
strategy: RoutingStrategy::BestLatency,
..req_default()
},
)
.unwrap();
assert!(
latency.reason.to_lowercase().contains("latency"),
"{}",
latency.reason
);
}
#[test]
fn test_round_robin_cycles_through_workers() {
let router = Router::new();
let registry = WorkerRegistry::new();
registry.register(make_reg("w1", 0, 4.0, 4, 0, 0.001));
registry.register(make_reg("w2", 1, 4.0, 4, 0, 0.001));
let credits = credits_for("alice", 10000.0);
let req = ResourceRequirements {
strategy: RoutingStrategy::RoundRobin,
..req_default()
};
let d1 = router
.select_worker(®istry, &credits, "alice", 10000.0, &req)
.unwrap();
let d2 = router
.select_worker(®istry, &credits, "alice", 10000.0, &req)
.unwrap();
assert_ne!(d1.worker.id, d2.worker.id);
}
#[test]
fn test_local_mode_returns_first_worker_free() {
let router = Router::new();
let registry = WorkerRegistry::new();
registry.register(make_reg("local-w", 0, 4.0, 4, 0, 0.001));
let req = req_default();
let decision = router.select_worker_no_checks(®istry, &req).unwrap();
assert_eq!(decision.estimated_cost, 0.0);
}
#[test]
fn test_local_mode_no_workers_returns_error() {
let router = Router::new();
let registry = WorkerRegistry::new();
let err = router
.select_worker_no_checks(®istry, &req_default())
.unwrap_err();
assert_eq!(err.code, "NO_WORKERS");
}
#[test]
fn test_estimate_cost_returns_min_and_max() {
let router = Router::new();
let registry = WorkerRegistry::new();
registry.register(make_reg("cheap", 0, 4.0, 4, 0, 0.001));
registry.register(make_reg("pricey", 1, 4.0, 4, 0, 0.1));
let req = ResourceRequirements {
cpus: 1.0,
estimated_duration_secs: 3600.0,
..req_default()
};
let (min, max) = router.estimate_cost(®istry, &req).unwrap();
assert!(min < max);
assert!((min - 0.001).abs() < 0.00001);
}
#[test]
fn test_estimate_cost_no_workers_returns_none() {
let router = Router::new();
let registry = WorkerRegistry::new();
assert!(router.estimate_cost(®istry, &req_default()).is_none());
}
#[test]
fn test_worker_type_match_routes_correctly() {
let router = Router::new();
let registry = WorkerRegistry::new();
let mut reg = make_reg("ray-w", 0, 4.0, 4, 0, 0.001);
reg.worker_type = "ray".to_string();
registry.register(reg);
let credits = credits_for("alice", 1000.0);
let req = ResourceRequirements {
worker_type: Some("ray".to_string()),
..req_default()
};
let decision = router
.select_worker(®istry, &credits, "alice", 1000.0, &req)
.unwrap();
assert_eq!(decision.worker.worker_type, "ray");
}
#[test]
fn test_price_change_affects_routing_decision() {
let router = Router::new();
let registry = WorkerRegistry::new();
let _w1 = registry.register(make_reg("w1", 0, 4.0, 4, 0, 0.001));
let _w2 = registry.register(make_reg("w2", 1, 4.0, 4, 0, 0.01));
let credits = credits_for("alice", 10000.0);
let req = ResourceRequirements {
strategy: RoutingStrategy::BestPrice,
estimated_duration_secs: 3600.0,
..req_default()
};
let d1 = router
.select_worker(®istry, &credits, "alice", 10000.0, &req)
.unwrap();
assert_eq!(d1.worker.name, "w1");
let (min, _max) = router.estimate_cost(®istry, &req).unwrap();
assert!((min - 0.001).abs() < 0.00001);
}
}
#[cfg(test)]
mod adaptive_tests {
use super::tests::{make_reg, req_default};
use super::*;
use crate::broker::policy::Weights;
use crate::broker::sampler::Rng;
#[test]
fn it_prefers_the_worker_that_has_been_fastest() {
let router = Router::new();
let registry = WorkerRegistry::new();
let fast = registry.register(make_reg("fast", 0, 4.0, 4, 0, 0.001));
let slow = registry.register(make_reg("slow", 1, 4.0, 4, 0, 0.001));
for _ in 0..40 {
registry.record_request(&fast.id, 20.0, true);
registry.record_request(&slow.id, 900.0, true);
}
let mut rng = Rng::seeded(4);
let w = Weights::default();
let picks = (0..200)
.filter(|_| {
router
.select_adaptive(®istry, &req_default(), &mut rng, &w)
.map(|d| d.worker.name == "fast")
.unwrap_or(false)
})
.count();
assert!(picks > 160, "chose the fast worker {picks}/200 times");
}
#[test]
fn a_brand_new_worker_is_still_tried() {
let router = Router::new();
let registry = WorkerRegistry::new();
let known = registry.register(make_reg("known", 0, 4.0, 4, 0, 0.001));
registry.register(make_reg("brand-new", 1, 4.0, 4, 0, 0.001));
for _ in 0..40 {
registry.record_request(&known.id, 200.0, true);
}
let mut rng = Rng::seeded(9);
let w = Weights::default();
let tried = (0..200)
.filter(|_| {
router
.select_adaptive(®istry, &req_default(), &mut rng, &w)
.map(|d| d.worker.name == "brand-new")
.unwrap_or(false)
})
.count();
assert!(
(10..190).contains(&tried),
"explored but not favoured blindly: {tried}/200"
);
}
#[test]
fn it_says_adaptive_in_the_reason_and_charges_nothing() {
let router = Router::new();
let registry = WorkerRegistry::new();
let w0 = registry.register(make_reg("w0", 0, 4.0, 4, 0, 0.001));
registry.record_request(&w0.id, 15.0, true);
let mut rng = Rng::seeded(2);
let d = router
.select_adaptive(®istry, &req_default(), &mut rng, &Weights::default())
.unwrap();
assert!(d.reason.starts_with("Adaptive:"), "{}", d.reason);
assert_eq!(d.estimated_cost, 0.0, "this path bills nothing");
}
#[test]
fn no_workers_is_an_error_not_a_panic() {
let router = Router::new();
let registry = WorkerRegistry::new();
let mut rng = Rng::seeded(1);
let err = router
.select_adaptive(®istry, &req_default(), &mut rng, &Weights::default())
.unwrap_err();
assert_eq!(err.code, "NO_WORKERS");
}
}