use std::collections::{BTreeMap, HashMap, VecDeque};
use std::future::Future;
use std::time::{Duration, Instant};
use tokio::task::JoinSet;
use crate::protocol::{RpcError, RunResult, TranscriptSummary};
use crate::{Params, Trial};
const GROW_THRESHOLD: usize = 3;
const MAX_BACKOFF_STEPS: u32 = 6;
#[derive(Clone, Debug)]
pub struct CaseSpec {
pub eval: String,
pub sample: String,
pub target: String,
pub provider: String,
pub params: Params,
pub trial: Trial,
pub timeout: Option<Duration>,
}
impl CaseSpec {
pub fn key(&self) -> String {
format!("{}{}", self.logical_key(), self.trial.key_suffix())
}
pub fn logical_key(&self) -> String {
crate::case_key(&self.eval, &self.sample, &self.target, &self.params)
}
}
#[derive(Clone, Debug)]
pub struct Concurrency {
pub global: usize,
pub per_provider: BTreeMap<String, usize>,
pub default_per_provider: usize,
pub adaptive: bool,
pub max_retries: u32,
pub base_backoff: Duration,
}
impl Concurrency {
pub fn new(global: usize) -> Self {
let global = global.max(1);
Self {
global,
per_provider: BTreeMap::new(),
default_per_provider: global,
adaptive: true,
max_retries: 4,
base_backoff: Duration::from_millis(500),
}
}
pub fn provider(mut self, provider: impl Into<String>, limit: usize) -> Self {
self.per_provider.insert(provider.into(), limit.max(1));
self
}
}
impl Default for Concurrency {
fn default() -> Self {
Self::new(8)
}
}
#[derive(Debug)]
struct ProviderState {
in_flight: usize,
limit: usize,
ceiling: usize,
ok_streak: usize,
backoff_steps: u32,
backoff_until: Option<Instant>,
}
struct Limiter {
providers: HashMap<String, ProviderState>,
per_provider: BTreeMap<String, usize>,
default_per_provider: usize,
global_in_flight: usize,
global_max: usize,
adaptive: bool,
base_backoff: Duration,
}
impl Limiter {
fn new(cfg: &Concurrency) -> Self {
Self {
providers: HashMap::new(),
per_provider: cfg.per_provider.clone(),
default_per_provider: cfg.default_per_provider.max(1),
global_in_flight: 0,
global_max: cfg.global.max(1),
adaptive: cfg.adaptive,
base_backoff: cfg.base_backoff,
}
}
fn ceiling_for(&self, provider: &str) -> usize {
self.per_provider
.get(provider)
.copied()
.unwrap_or(self.default_per_provider)
.clamp(1, self.global_max)
}
fn state(&mut self, provider: &str) -> &mut ProviderState {
let ceiling = self.ceiling_for(provider);
self.providers
.entry(provider.to_string())
.or_insert_with(|| ProviderState {
in_flight: 0,
limit: ceiling,
ceiling,
ok_streak: 0,
backoff_steps: 0,
backoff_until: None,
})
}
fn can_start(&mut self, provider: &str, now: Instant) -> bool {
if self.global_in_flight >= self.global_max {
return false;
}
let st = self.state(provider);
st.in_flight < st.limit && st.backoff_until.is_none_or(|t| now >= t)
}
fn start(&mut self, provider: &str) {
self.global_in_flight += 1;
self.state(provider).in_flight += 1;
}
fn finish(&mut self, provider: &str, rate_limited: bool, now: Instant) {
let adaptive = self.adaptive;
let base = self.base_backoff;
self.global_in_flight = self.global_in_flight.saturating_sub(1);
let st = self.state(provider);
st.in_flight = st.in_flight.saturating_sub(1);
if !adaptive {
return;
}
if rate_limited {
st.limit = (st.limit / 2).max(1);
st.backoff_steps = (st.backoff_steps + 1).min(MAX_BACKOFF_STEPS);
let mult = 1u32 << (st.backoff_steps - 1);
st.backoff_until = Some(now + base * mult);
st.ok_streak = 0;
} else {
st.ok_streak += 1;
if st.ok_streak >= GROW_THRESHOLD {
st.ok_streak = 0;
st.backoff_steps = st.backoff_steps.saturating_sub(1);
if st.limit < st.ceiling {
st.limit += 1;
}
}
}
}
fn earliest_ready(&self, pending: &VecDeque<(CaseSpec, u32)>) -> Option<Instant> {
pending
.iter()
.filter_map(|(c, _)| self.providers.get(&c.provider))
.filter_map(|s| s.backoff_until)
.min()
}
}
fn outcome_rate_limited(res: &Result<RunResult, RpcError>) -> bool {
match res {
Err(e) => crate::is_rate_limited(&e.message),
Ok(r) => r
.transcript
.error
.as_deref()
.is_some_and(crate::is_rate_limited),
}
}
fn outcome_retryable(res: &Result<RunResult, RpcError>) -> bool {
match res {
Err(e) => e.retryable || crate::is_rate_limited(&e.message),
Ok(r) => {
r.transcript.error_kind == crate::ErrorKind::Infra
|| r.transcript
.error
.as_deref()
.is_some_and(crate::is_rate_limited)
}
}
}
fn failed_result(case: &CaseSpec, error: RpcError) -> RunResult {
let infra = error.retryable || crate::is_rate_limited(&error.message);
RunResult {
eval: case.eval.clone(),
sample: case.sample.clone(),
target: case.target.clone(),
params: case.params.clone(),
trial: case.trial.index,
trials: case.trial.count,
seed: case.trial.seed,
input: Vec::new(),
expected: None,
passed: false,
aggregate: 0.0,
scores: Vec::new(),
transcript: TranscriptSummary {
error: Some(error.message),
error_kind: if infra {
crate::ErrorKind::Infra
} else {
crate::ErrorKind::Subject
},
..Default::default()
},
skipped: false,
}
}
pub async fn run_cases<F, Fut>(
cases: Vec<CaseSpec>,
cfg: &Concurrency,
run: F,
mut on_done: impl FnMut(&CaseSpec, RunResult),
) where
F: Fn(CaseSpec) -> Fut,
Fut: Future<Output = Result<RunResult, RpcError>> + Send + 'static,
{
let mut limiter = Limiter::new(cfg);
let mut pending: VecDeque<(CaseSpec, u32)> = cases.into_iter().map(|c| (c, 0)).collect();
let mut tasks: JoinSet<Result<RunResult, RpcError>> = JoinSet::new();
let mut inflight: HashMap<tokio::task::Id, (CaseSpec, u32)> = HashMap::new();
loop {
loop {
let now = Instant::now();
let idx = pending
.iter()
.position(|(c, _)| limiter.can_start(&c.provider, now));
let Some(idx) = idx else { break };
let (case, attempts) = pending.remove(idx).expect("index in bounds");
limiter.start(&case.provider);
let task_case = case.clone();
let timeout = case.timeout;
let fut = run(case);
let task = async move {
match timeout {
Some(dur) => match tokio::time::timeout(dur, fut).await {
Ok(res) => res,
Err(_) => Err(RpcError::new(format!(
"timed out after {}s (target timeout)",
dur.as_secs()
))),
},
None => fut.await,
}
};
let id = tasks.spawn(task).id();
inflight.insert(id, (task_case, attempts));
}
if tasks.is_empty() {
if pending.is_empty() {
break;
}
match limiter.earliest_ready(&pending) {
Some(t) => {
tokio::time::sleep_until(t.into()).await;
continue;
}
None => break,
}
}
let Some(joined) = tasks.join_next_with_id().await else {
continue;
};
let (case, attempts, res) = match joined {
Ok((id, res)) => {
let (case, attempts) = inflight.remove(&id).expect("task id tracked");
(case, attempts, res)
}
Err(join_err) => {
let (case, attempts) = inflight.remove(&join_err.id()).expect("task id tracked");
(
case,
attempts,
Err(RpcError::new(format!("task panicked: {join_err}"))),
)
}
};
let rate_limited = outcome_rate_limited(&res);
limiter.finish(&case.provider, rate_limited, Instant::now());
if attempts < cfg.max_retries && outcome_retryable(&res) {
pending.push_back((case, attempts + 1));
continue;
}
let result = match res {
Ok(result) => result,
Err(error) => failed_result(&case, error),
};
on_done(&case, result);
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use tokio::sync::Mutex;
fn case(provider: &str, id: &str) -> CaseSpec {
CaseSpec {
eval: "e".into(),
sample: id.into(),
target: format!("{provider}/m"),
provider: provider.into(),
params: Params::new(),
trial: Trial::single(),
timeout: None,
}
}
fn ok_result(case: &CaseSpec) -> RunResult {
RunResult {
eval: case.eval.clone(),
sample: case.sample.clone(),
target: case.target.clone(),
params: case.params.clone(),
trial: case.trial.index,
trials: case.trial.count,
seed: case.trial.seed,
input: Vec::new(),
expected: None,
passed: true,
aggregate: 1.0,
scores: Vec::new(),
transcript: TranscriptSummary::default(),
skipped: false,
}
}
#[test]
fn limiter_buckets_per_provider() {
let cfg = Concurrency::new(10).provider("anthropic", 2);
let mut lim = Limiter::new(&cfg);
let now = Instant::now();
assert!(lim.can_start("anthropic", now));
lim.start("anthropic");
lim.start("anthropic");
assert!(!lim.can_start("anthropic", now));
assert!(lim.can_start("openai", now));
}
#[test]
fn rate_limit_halves_and_backs_off() {
let cfg = Concurrency::new(8); let mut lim = Limiter::new(&cfg);
let now = Instant::now();
lim.start("anthropic");
lim.finish("anthropic", true, now);
let st = &lim.providers["anthropic"];
assert_eq!(st.limit, 4); assert!(st.backoff_until.is_some());
assert!(!lim.can_start("anthropic", now));
}
#[test]
fn sustained_success_grows_back() {
let cfg = Concurrency::new(8);
let mut lim = Limiter::new(&cfg);
let now = Instant::now();
lim.start("anthropic");
lim.finish("anthropic", true, now); assert_eq!(lim.providers["anthropic"].limit, 4);
for _ in 0..GROW_THRESHOLD {
lim.start("anthropic");
lim.finish("anthropic", false, now);
}
assert_eq!(lim.providers["anthropic"].limit, 5);
}
#[tokio::test]
async fn runs_every_case_once() {
let cases: Vec<CaseSpec> = (0..20).map(|i| case("sim", &i.to_string())).collect();
let cfg = Concurrency::new(4);
let seen = Arc::new(AtomicUsize::new(0));
let seen2 = seen.clone();
let mut done = Vec::new();
run_cases(
cases,
&cfg,
move |c| {
let seen = seen2.clone();
async move {
seen.fetch_add(1, Ordering::SeqCst);
Ok(ok_result(&c))
}
},
|_, r| done.push(r),
)
.await;
assert_eq!(seen.load(Ordering::SeqCst), 20);
assert_eq!(done.len(), 20);
assert!(done.iter().all(|r| r.passed));
}
#[tokio::test]
async fn respects_global_concurrency_cap() {
let cases: Vec<CaseSpec> = (0..30).map(|i| case("sim", &i.to_string())).collect();
let cfg = Concurrency::new(3);
let active = Arc::new(AtomicUsize::new(0));
let peak = Arc::new(AtomicUsize::new(0));
let (a, p) = (active.clone(), peak.clone());
let mut done = 0usize;
run_cases(
cases,
&cfg,
move |c| {
let (a, p) = (a.clone(), p.clone());
async move {
let n = a.fetch_add(1, Ordering::SeqCst) + 1;
p.fetch_max(n, Ordering::SeqCst);
tokio::time::sleep(Duration::from_millis(5)).await;
a.fetch_sub(1, Ordering::SeqCst);
Ok(ok_result(&c))
}
},
|_, _| done += 1,
)
.await;
assert_eq!(done, 30);
assert!(peak.load(Ordering::SeqCst) <= 3, "peak exceeded global cap");
}
#[tokio::test]
async fn retries_rate_limited_case_then_succeeds() {
let cfg = Concurrency {
base_backoff: Duration::from_millis(1),
..Concurrency::new(2)
};
let attempts = Arc::new(Mutex::new(0usize));
let a = attempts.clone();
let mut results = Vec::new();
run_cases(
vec![case("anthropic", "x")],
&cfg,
move |c| {
let a = a.clone();
async move {
let mut n = a.lock().await;
*n += 1;
if *n == 1 {
Err(RpcError::new("HTTP 429 rate limit"))
} else {
Ok(ok_result(&c))
}
}
},
|_, r| results.push(r),
)
.await;
assert_eq!(*attempts.lock().await, 2);
assert_eq!(results.len(), 1);
assert!(results[0].passed);
}
#[tokio::test]
async fn retries_infra_errored_case_then_succeeds() {
let cfg = Concurrency::new(2);
let attempts = Arc::new(Mutex::new(0usize));
let a = attempts.clone();
let mut results = Vec::new();
run_cases(
vec![case("sim", "x")],
&cfg,
move |c| {
let a = a.clone();
async move {
let mut n = a.lock().await;
*n += 1;
if *n == 1 {
let mut r = ok_result(&c);
r.passed = false;
r.transcript.error = Some("provider 503 unavailable".into());
r.transcript.error_kind = crate::ErrorKind::Infra;
Ok(r)
} else {
Ok(ok_result(&c))
}
}
},
|_, r| results.push(r),
)
.await;
assert_eq!(*attempts.lock().await, 2); assert_eq!(results.len(), 1);
assert!(results[0].passed);
}
#[tokio::test]
async fn retries_retryable_rpc_error_then_succeeds() {
let cfg = Concurrency::new(2);
let attempts = Arc::new(Mutex::new(0usize));
let a = attempts.clone();
let mut results = Vec::new();
run_cases(
vec![case("sim", "x")],
&cfg,
move |c| {
let a = a.clone();
async move {
let mut n = a.lock().await;
*n += 1;
if *n == 1 {
Err(RpcError::new("provider outage").retryable())
} else {
Ok(ok_result(&c))
}
}
},
|_, r| results.push(r),
)
.await;
assert_eq!(*attempts.lock().await, 2); assert_eq!(results.len(), 1);
assert!(results[0].passed);
}
#[tokio::test]
async fn non_retryable_rpc_error_is_not_requeued() {
let cfg = Concurrency::new(2);
let count = Arc::new(AtomicUsize::new(0));
let c2 = count.clone();
let mut results = Vec::new();
run_cases(
vec![case("sim", "x")],
&cfg,
move |_| {
let c2 = c2.clone();
async move {
c2.fetch_add(1, Ordering::SeqCst);
Err(RpcError::new("bad run params"))
}
},
|_, r| results.push(r),
)
.await;
assert_eq!(count.load(Ordering::SeqCst), 1); assert_eq!(results.len(), 1);
assert!(!results[0].passed);
assert_eq!(results[0].transcript.error_kind, crate::ErrorKind::Subject);
}
#[tokio::test]
async fn panicking_case_is_recorded_and_frees_its_slot() {
let cfg = Concurrency::new(1);
let cases = vec![case("sim", "boom"), case("sim", "ok")];
let mut results = Vec::new();
run_cases(
cases,
&cfg,
move |c| async move {
if c.sample == "boom" {
panic!("subject blew up");
}
Ok(ok_result(&c))
},
|c, r| results.push((c.sample.clone(), r)),
)
.await;
assert_eq!(results.len(), 2);
let boom = results.iter().find(|(s, _)| s == "boom").unwrap();
assert!(!boom.1.passed);
assert!(
boom.1
.transcript
.error
.as_deref()
.unwrap()
.contains("panic")
);
let ok = results.iter().find(|(s, _)| s == "ok").unwrap();
assert!(ok.1.passed);
}
#[tokio::test]
async fn gives_up_after_max_retries() {
let cfg = Concurrency {
base_backoff: Duration::from_millis(1),
max_retries: 2,
..Concurrency::new(1)
};
let count = Arc::new(AtomicUsize::new(0));
let c2 = count.clone();
let mut results = Vec::new();
run_cases(
vec![case("anthropic", "x")],
&cfg,
move |_| {
let c2 = c2.clone();
async move {
c2.fetch_add(1, Ordering::SeqCst);
Err(RpcError::new("429 too many requests"))
}
},
|_, r| results.push(r),
)
.await;
assert_eq!(count.load(Ordering::SeqCst), 3);
assert_eq!(results.len(), 1);
assert!(!results[0].passed);
assert!(results[0].transcript.error.is_some());
}
#[tokio::test]
async fn times_out_slow_case_without_retrying() {
let cfg = Concurrency {
base_backoff: Duration::from_millis(1),
max_retries: 4,
..Concurrency::new(2)
};
let count = Arc::new(AtomicUsize::new(0));
let c2 = count.clone();
let slow = CaseSpec {
timeout: Some(Duration::from_millis(10)),
..case("sim", "slow")
};
let mut results = Vec::new();
run_cases(
vec![slow],
&cfg,
move |c| {
let c2 = c2.clone();
async move {
c2.fetch_add(1, Ordering::SeqCst);
tokio::time::sleep(Duration::from_secs(60)).await;
Ok(ok_result(&c))
}
},
|_, r| results.push(r),
)
.await;
assert_eq!(count.load(Ordering::SeqCst), 1, "no retry after timeout");
assert_eq!(results.len(), 1);
assert!(!results[0].passed);
let err = results[0].transcript.error.as_deref().unwrap();
assert!(err.contains("timed out"), "{err}");
assert_eq!(results[0].transcript.error_kind, crate::ErrorKind::Subject);
}
#[tokio::test]
async fn fast_case_under_timeout_passes() {
let cfg = Concurrency::new(2);
let fast = CaseSpec {
timeout: Some(Duration::from_secs(30)),
..case("sim", "fast")
};
let mut results = Vec::new();
run_cases(
vec![fast],
&cfg,
move |c| async move { Ok(ok_result(&c)) },
|_, r| results.push(r),
)
.await;
assert_eq!(results.len(), 1);
assert!(results[0].passed);
}
}