use super::estimator::Stats;
use super::sampler::{best_of_d, AliasTable, Rng};
#[derive(Debug, Clone)]
pub struct Candidate {
pub stats: Stats,
pub in_flight: u32,
pub price_per_hour: f64,
}
#[derive(Debug, Clone, Copy)]
pub struct Weights {
pub latency: f64,
pub price: f64,
pub retry: f64,
pub unmeasured_ms: f64,
}
impl Default for Weights {
fn default() -> Self {
Weights {
latency: 1.0,
price: 0.0,
retry: 2.0,
unmeasured_ms: 250.0,
}
}
}
fn standard_normal(rng: &mut Rng) -> f64 {
let u1 = rng.next_f64().max(f64::MIN_POSITIVE);
let u2 = rng.next_f64();
(-2.0 * u1.ln()).sqrt() * (std::f64::consts::TAU * u2).cos()
}
pub fn sample_service_ms(stats: &Stats, w: &Weights, rng: &mut Rng) -> f64 {
let n = stats.weight();
match (stats.median_service_ms(), n >= 1.0) {
(Some(median), true) => {
let sd = stats.log_sd().unwrap_or(0.5).max(0.05);
let spread = sd / n.sqrt();
(median.ln() + standard_normal(rng) * spread).exp()
}
_ => (w.unmeasured_ms.ln() + standard_normal(rng) * 0.8).exp(),
}
}
pub fn cost(candidate: &Candidate, sampled_service_ms: f64, w: &Weights) -> f64 {
let completion = sampled_service_ms * (1.0 + candidate.in_flight as f64);
let failure = 1.0 - candidate.stats.success_rate() * candidate.stats.trust();
let price = candidate.price_per_hour * sampled_service_ms / 3_600_000.0;
w.latency * completion + w.price * price + w.retry * failure * completion
}
pub fn choose(
candidates: &[Candidate],
table: &AliasTable,
rng: &mut Rng,
d: usize,
w: &Weights,
) -> Option<usize> {
if candidates.is_empty() || table.is_empty() {
return None;
}
let sampled = std::cell::RefCell::new(rng.clone());
let chosen = best_of_d(table, rng, d, |i| {
let c = &candidates[i];
let s = sample_service_ms(&c.stats, w, &mut sampled.borrow_mut());
cost(c, s, w)
});
chosen.filter(|i| *i < candidates.len())
}
pub fn table_weight(candidate: &Candidate, w: &Weights) -> f64 {
let service = candidate
.stats
.expected_service_ms()
.unwrap_or(w.unmeasured_ms);
let completion = service * (1.0 + candidate.in_flight as f64);
if completion <= 0.0 {
return 0.0;
}
let rate = 1.0 / completion;
rate * rate
}
#[cfg(test)]
mod tests {
use super::*;
fn measured(median_ms: f64, n: usize, ok: bool) -> Stats {
let mut s = Stats::new();
for _ in 0..n {
s.observe(median_ms, ok);
}
s
}
fn candidate(median_ms: f64, n: usize, in_flight: u32) -> Candidate {
Candidate {
stats: measured(median_ms, n, true),
in_flight,
price_per_hour: 3.6,
}
}
#[test]
fn a_well_measured_worker_is_sampled_close_to_what_it_measured() {
let w = Weights::default();
let mut rng = Rng::seeded(1);
let stats = measured(50.0, 200, true);
let draws: Vec<f64> = (0..500)
.map(|_| sample_service_ms(&stats, &w, &mut rng))
.collect();
let mean = draws.iter().sum::<f64>() / draws.len() as f64;
assert!(
(mean - 50.0).abs() < 10.0,
"200 observations should pin it near 50 ms, got {mean:.1}"
);
}
#[test]
fn a_barely_measured_worker_is_sampled_widely() {
let w = Weights::default();
let spread = |n: usize| {
let stats = measured(50.0, n, true);
let mut rng = Rng::seeded(7);
let draws: Vec<f64> = (0..2000)
.map(|_| sample_service_ms(&stats, &w, &mut rng))
.collect();
let mean = draws.iter().sum::<f64>() / draws.len() as f64;
let var = draws.iter().map(|d| (d - mean).powi(2)).sum::<f64>() / draws.len() as f64;
var.sqrt()
};
assert!(
spread(2) > spread(200) * 3.0,
"little evidence must sample widely: {:.1} vs {:.1}",
spread(2),
spread(200)
);
}
#[test]
fn queueing_counts_against_a_worker() {
let w = Weights::default();
let idle = candidate(30.0, 50, 0);
let busy = candidate(10.0, 50, 4);
assert!(
cost(&busy, 10.0, &w) > cost(&idle, 30.0, &w),
"four queued at 10 ms is worse than idle at 30 ms"
);
}
#[test]
fn failure_and_price_move_the_cost_the_way_they_should() {
let w = Weights::default();
let mut flaky = candidate(50.0, 0, 0);
for _ in 0..20 {
flaky.stats.observe(50.0, false);
}
let solid = candidate(50.0, 20, 0);
assert!(
cost(&flaky, 50.0, &w) > cost(&solid, 50.0, &w),
"a worker that keeps failing costs more than one that does not"
);
let priced = Weights {
price: 1e6,
..Weights::default()
};
let cheap = Candidate {
price_per_hour: 1.0,
..candidate(50.0, 20, 0)
};
let dear = Candidate {
price_per_hour: 100.0,
..candidate(50.0, 20, 0)
};
assert!(cost(&dear, 50.0, &priced) > cost(&cheap, 50.0, &priced));
}
#[test]
fn the_table_weights_a_fast_idle_worker_above_a_slow_busy_one() {
let w = Weights::default();
assert!(
table_weight(&candidate(10.0, 20, 0), &w) > table_weight(&candidate(10.0, 20, 4), &w)
);
assert!(
table_weight(&candidate(10.0, 20, 0), &w) > table_weight(&candidate(100.0, 20, 0), &w)
);
assert!(table_weight(&candidate(0.0, 0, 0), &w) > 0.0);
}
}
#[cfg(test)]
mod regret_tests {
use super::*;
struct Fleet {
truth_ms: Vec<f64>,
stats: Vec<Stats>,
in_flight: Vec<u32>,
}
impl Fleet {
fn new(truth_ms: Vec<f64>) -> Self {
let n = truth_ms.len();
Fleet {
truth_ms,
stats: vec![Stats::new(); n],
in_flight: vec![0; n],
}
}
fn candidates(&self) -> Vec<Candidate> {
(0..self.truth_ms.len())
.map(|i| Candidate {
stats: self.stats[i],
in_flight: self.in_flight[i],
price_per_hour: 3.6,
})
.collect()
}
fn run(&mut self, i: usize, rng: &mut Rng) -> f64 {
let noise = 0.75 + rng.next_f64() * 0.5;
let took = self.truth_ms[i] * noise;
self.stats[i].decay(0.98);
self.stats[i].observe(took, true);
took
}
}
fn run_policy(truth: Vec<f64>, tasks: usize, seed: u64, d: usize) -> f64 {
let mut fleet = Fleet::new(truth);
let mut rng = Rng::seeded(seed);
let w = Weights::default();
let mut total = 0.0;
for t in 0..tasks {
let cands = fleet.candidates();
let weights: Vec<f64> = cands.iter().map(|c| table_weight(c, &w)).collect();
let table = AliasTable::build(&weights).expect("a fleet");
let i = choose(&cands, &table, &mut rng, d, &w).expect("a choice");
total += fleet.run(i, &mut rng);
let _ = t;
}
total
}
fn run_round_robin(truth: Vec<f64>, tasks: usize, seed: u64) -> f64 {
let mut fleet = Fleet::new(truth);
let mut rng = Rng::seeded(seed);
let n = fleet.truth_ms.len();
(0..tasks).map(|t| fleet.run(t % n, &mut rng)).sum()
}
fn run_oracle(truth: Vec<f64>, tasks: usize, seed: u64) -> f64 {
let best = truth
.iter()
.enumerate()
.min_by(|a, b| a.1.partial_cmp(b.1).unwrap())
.map(|(i, _)| i)
.unwrap();
let mut fleet = Fleet::new(truth);
let mut rng = Rng::seeded(seed);
(0..tasks).map(|_| fleet.run(best, &mut rng)).sum()
}
fn mesh() -> Vec<f64> {
let mut t = vec![254.0; 4];
t.extend(std::iter::repeat_n(1000.0, 9));
t
}
#[test]
fn the_policy_beats_round_robin_by_a_wide_margin() {
let tasks = 3_000;
let policy = run_policy(mesh(), tasks, 11, 2);
let rr = run_round_robin(mesh(), tasks, 11);
let oracle = run_oracle(mesh(), tasks, 11);
let policy_regret = policy - oracle;
let rr_regret = rr - oracle;
println!(
" per task: oracle {:.0} ms, policy {:.0} ms, round-robin {:.0} ms",
oracle / tasks as f64,
policy / tasks as f64,
rr / tasks as f64
);
assert!(
policy_regret * 3.0 < rr_regret,
"policy regret {policy_regret:.0} ms vs round-robin {rr_regret:.0} ms"
);
}
#[test]
fn regret_per_task_shrinks_as_it_learns() {
let oracle_per_task = |tasks: usize| run_oracle(mesh(), tasks, 3) / tasks as f64;
let policy_per_task = |tasks: usize| run_policy(mesh(), tasks, 3, 2) / tasks as f64;
let early = policy_per_task(300) - oracle_per_task(300);
let late = policy_per_task(6_000) - oracle_per_task(6_000);
println!(" regret per task: early {early:.1} ms, late {late:.1} ms");
assert!(
late < early * 0.6,
"should keep improving: early {early:.1} ms, late {late:.1} ms"
);
}
#[test]
fn it_notices_a_slow_worker_that_becomes_fast() {
let mut fleet = Fleet::new(mesh());
let mut rng = Rng::seeded(33);
let w = Weights::default();
let mut dispatch = |fleet: &mut Fleet, rng: &mut Rng| -> usize {
let cands = fleet.candidates();
let weights: Vec<f64> = cands.iter().map(|c| table_weight(c, &w)).collect();
let table = AliasTable::build(&weights).expect("a fleet");
let i = choose(&cands, &table, rng, 2, &w).expect("a choice");
fleet.run(i, rng);
i
};
for _ in 0..1_500 {
dispatch(&mut fleet, &mut rng);
}
let before: usize = (0..300)
.map(|_| usize::from(dispatch(&mut fleet, &mut rng) == 12))
.sum();
fleet.truth_ms[12] = 20.0;
for _ in 0..4_000 {
dispatch(&mut fleet, &mut rng);
}
let after: usize = (0..300)
.map(|_| usize::from(dispatch(&mut fleet, &mut rng) == 12))
.sum();
println!(" share on the worker that got fast: {before}/300 → {after}/300");
assert!(before < 30, "it had learned to avoid it: {before}/300");
assert!(
after > 150,
"it must find a worker that improves, not write it off: {after}/300"
);
}
#[test]
fn it_follows_the_mesh_when_conditions_change() {
let mut fleet = Fleet::new(mesh());
let mut rng = Rng::seeded(21);
let w = Weights::default();
let mut dispatch = |fleet: &mut Fleet, rng: &mut Rng| -> usize {
let cands = fleet.candidates();
let weights: Vec<f64> = cands.iter().map(|c| table_weight(c, &w)).collect();
let table = AliasTable::build(&weights).expect("a fleet");
let i = choose(&cands, &table, rng, 2, &w).expect("a choice");
fleet.run(i, rng);
i
};
for _ in 0..1_500 {
dispatch(&mut fleet, &mut rng);
}
let before: usize = (0..300)
.map(|_| usize::from(dispatch(&mut fleet, &mut rng) < 4))
.sum();
for i in 0..4 {
fleet.truth_ms[i] = 2_000.0;
}
for _ in 0..1_500 {
dispatch(&mut fleet, &mut rng);
}
let after: usize = (0..300)
.map(|_| usize::from(dispatch(&mut fleet, &mut rng) < 4))
.sum();
println!(" share on the formerly-fast workers: {before}/300 → {after}/300");
assert!(before > 200, "it found the fast ones: {before}/300");
assert!(
after < 90,
"it moved off them once they slowed: {after}/300"
);
}
}