use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum CascadeTier {
SelfReflect,
ToolVerify,
HumanExpert,
}
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
pub struct TierSpec {
pub tier: CascadeTier,
pub cost: f64,
pub expected_confidence: f64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CascadePolicy {
pub confidence_target: f64,
pub budget: f64,
pub tiers: Vec<TierSpec>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "decision", rename_all = "snake_case")]
pub enum CascadeOutcome {
AlreadyConfident { confidence: f64 },
Accept {
tier: CascadeTier,
confidence: f64,
cost_spent: f64,
},
Exhausted {
best_tier: Option<CascadeTier>,
confidence: f64,
cost_spent: f64,
},
}
pub fn decide_cascade(current_confidence: f64, policy: &CascadePolicy) -> CascadeOutcome {
let target = policy.confidence_target.clamp(0.0, 1.0);
let start = current_confidence.clamp(0.0, 1.0);
if start >= target {
return CascadeOutcome::AlreadyConfident { confidence: start };
}
let mut cumulative_cost = 0.0;
let mut best_tier: Option<CascadeTier> = None;
let mut best_confidence = start;
let mut best_cost = 0.0;
for spec in &policy.tiers {
let next_cost = cumulative_cost + spec.cost;
if next_cost > policy.budget {
break;
}
cumulative_cost = next_cost;
let conf = spec.expected_confidence.clamp(0.0, 1.0);
if conf >= target {
return CascadeOutcome::Accept {
tier: spec.tier,
confidence: conf,
cost_spent: cumulative_cost,
};
}
if conf > best_confidence {
best_confidence = conf;
best_tier = Some(spec.tier);
best_cost = cumulative_cost;
}
}
CascadeOutcome::Exhausted {
best_tier,
confidence: best_confidence,
cost_spent: best_cost,
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct TierResult {
pub knowledge: String,
pub confidence: f64,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct CascadeStep {
pub tier: CascadeTier,
pub confidence: f64,
pub cost: f64,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct CascadeRun {
pub accepted_tier: Option<CascadeTier>,
pub knowledge: Option<String>,
pub confidence: f64,
pub cost_spent: f64,
pub steps: Vec<CascadeStep>,
}
pub fn run_cascade(
current_confidence: f64,
policy: &CascadePolicy,
mut run: impl FnMut(CascadeTier) -> Result<TierResult, String>,
) -> Result<CascadeRun, String> {
let target = policy.confidence_target.clamp(0.0, 1.0);
let start = current_confidence.clamp(0.0, 1.0);
let mut result = CascadeRun {
accepted_tier: None,
knowledge: None,
confidence: start,
cost_spent: 0.0,
steps: Vec::new(),
};
if start >= target {
return Ok(result);
}
let mut cumulative_cost = 0.0;
let mut best_confidence = start;
for spec in &policy.tiers {
let next_cost = cumulative_cost + spec.cost;
if next_cost > policy.budget {
break; }
cumulative_cost = next_cost;
let outcome = run(spec.tier)?;
let conf = outcome.confidence.clamp(0.0, 1.0);
result.steps.push(CascadeStep {
tier: spec.tier,
confidence: conf,
cost: cumulative_cost,
});
result.cost_spent = cumulative_cost;
if conf >= best_confidence || result.knowledge.is_none() {
best_confidence = best_confidence.max(conf);
result.knowledge = Some(outcome.knowledge);
}
result.confidence = best_confidence;
if conf >= target {
result.accepted_tier = Some(spec.tier);
result.confidence = conf;
return Ok(result);
}
}
Ok(result)
}
pub async fn run_cascade_async<F, Fut>(
current_confidence: f64,
policy: &CascadePolicy,
mut run: F,
) -> Result<CascadeRun, String>
where
F: FnMut(CascadeTier) -> Fut,
Fut: std::future::Future<Output = Result<TierResult, String>>,
{
let target = policy.confidence_target.clamp(0.0, 1.0);
let start = current_confidence.clamp(0.0, 1.0);
let mut result = CascadeRun {
accepted_tier: None,
knowledge: None,
confidence: start,
cost_spent: 0.0,
steps: Vec::new(),
};
if start >= target {
return Ok(result);
}
let mut cumulative_cost = 0.0;
let mut best_confidence = start;
for spec in &policy.tiers {
let next_cost = cumulative_cost + spec.cost;
if next_cost > policy.budget {
break; }
cumulative_cost = next_cost;
let outcome = run(spec.tier).await?;
let conf = outcome.confidence.clamp(0.0, 1.0);
result.steps.push(CascadeStep {
tier: spec.tier,
confidence: conf,
cost: cumulative_cost,
});
result.cost_spent = cumulative_cost;
if conf >= best_confidence || result.knowledge.is_none() {
best_confidence = best_confidence.max(conf);
result.knowledge = Some(outcome.knowledge);
}
result.confidence = best_confidence;
if conf >= target {
result.accepted_tier = Some(spec.tier);
result.confidence = conf;
return Ok(result);
}
}
Ok(result)
}
#[cfg(test)]
mod tests {
use super::*;
fn policy(target: f64, budget: f64) -> CascadePolicy {
CascadePolicy {
confidence_target: target,
budget,
tiers: vec![
TierSpec {
tier: CascadeTier::SelfReflect,
cost: 1.0,
expected_confidence: 0.6,
},
TierSpec {
tier: CascadeTier::ToolVerify,
cost: 5.0,
expected_confidence: 0.85,
},
TierSpec {
tier: CascadeTier::HumanExpert,
cost: 50.0,
expected_confidence: 0.99,
},
],
}
}
#[test]
fn already_confident_spends_nothing() {
let out = decide_cascade(0.9, &policy(0.8, 100.0));
assert_eq!(out, CascadeOutcome::AlreadyConfident { confidence: 0.9 });
}
#[test]
fn picks_cheapest_tier_that_meets_target() {
let out = decide_cascade(0.2, &policy(0.6, 100.0));
assert_eq!(
out,
CascadeOutcome::Accept {
tier: CascadeTier::SelfReflect,
confidence: 0.6,
cost_spent: 1.0,
}
);
}
#[test]
fn escalates_when_cheap_tier_insufficient() {
let out = decide_cascade(0.2, &policy(0.8, 100.0));
assert_eq!(
out,
CascadeOutcome::Accept {
tier: CascadeTier::ToolVerify,
confidence: 0.85,
cost_spent: 6.0,
}
);
}
#[test]
fn escalates_to_human_for_high_target() {
let out = decide_cascade(0.1, &policy(0.95, 100.0));
assert_eq!(
out,
CascadeOutcome::Accept {
tier: CascadeTier::HumanExpert,
confidence: 0.99,
cost_spent: 56.0, }
);
}
#[test]
fn budget_caps_escalation_returns_best_affordable() {
let out = decide_cascade(0.1, &policy(0.95, 10.0));
assert_eq!(
out,
CascadeOutcome::Exhausted {
best_tier: Some(CascadeTier::ToolVerify),
confidence: 0.85,
cost_spent: 6.0,
}
);
}
#[test]
fn budget_too_small_for_any_tier() {
let out = decide_cascade(0.1, &policy(0.6, 0.5));
assert_eq!(
out,
CascadeOutcome::Exhausted {
best_tier: None,
confidence: 0.1,
cost_spent: 0.0,
}
);
}
#[test]
fn empty_tiers_is_exhausted() {
let p = CascadePolicy {
confidence_target: 0.8,
budget: 100.0,
tiers: vec![],
};
let out = decide_cascade(0.2, &p);
assert_eq!(
out,
CascadeOutcome::Exhausted {
best_tier: None,
confidence: 0.2,
cost_spent: 0.0,
}
);
}
fn runner(
observed: std::collections::HashMap<CascadeTier, f64>,
) -> impl FnMut(CascadeTier) -> Result<TierResult, String> {
move |tier| {
let confidence = *observed.get(&tier).unwrap_or(&0.0);
Ok(TierResult {
knowledge: format!("{tier:?} knowledge"),
confidence,
})
}
}
#[test]
fn run_already_confident_runs_nothing() {
let run = run_cascade(0.9, &policy(0.8, 100.0), runner(Default::default())).unwrap();
assert!(run.steps.is_empty());
assert_eq!(run.accepted_tier, None);
assert_eq!(run.knowledge, None);
assert_eq!(run.confidence, 0.9);
assert_eq!(run.cost_spent, 0.0);
}
#[test]
fn run_stops_at_cheapest_tier_that_observes_target() {
let observed = [(CascadeTier::SelfReflect, 0.9)].into_iter().collect();
let run = run_cascade(0.2, &policy(0.8, 100.0), runner(observed)).unwrap();
assert_eq!(run.accepted_tier, Some(CascadeTier::SelfReflect));
assert_eq!(run.steps.len(), 1);
assert_eq!(run.cost_spent, 1.0);
assert_eq!(run.confidence, 0.9);
assert_eq!(run.knowledge.as_deref(), Some("SelfReflect knowledge"));
}
#[test]
fn run_escalates_on_low_observed_confidence() {
let observed = [
(CascadeTier::SelfReflect, 0.4),
(CascadeTier::ToolVerify, 0.85),
]
.into_iter()
.collect();
let run = run_cascade(0.2, &policy(0.8, 100.0), runner(observed)).unwrap();
assert_eq!(run.accepted_tier, Some(CascadeTier::ToolVerify));
assert_eq!(run.steps.len(), 2);
assert_eq!(run.cost_spent, 6.0); assert_eq!(run.knowledge.as_deref(), Some("ToolVerify knowledge"));
}
#[test]
fn run_budget_caps_escalation_keeps_best_knowledge() {
let observed = [
(CascadeTier::SelfReflect, 0.5),
(CascadeTier::ToolVerify, 0.85),
(CascadeTier::HumanExpert, 0.99),
]
.into_iter()
.collect();
let run = run_cascade(0.1, &policy(0.95, 10.0), runner(observed)).unwrap();
assert_eq!(run.accepted_tier, None); assert_eq!(run.steps.len(), 2); assert_eq!(run.cost_spent, 6.0);
assert_eq!(run.confidence, 0.85); assert_eq!(run.knowledge.as_deref(), Some("ToolVerify knowledge"));
}
#[test]
fn run_propagates_runner_error() {
let err = run_cascade(0.2, &policy(0.8, 100.0), |_tier| {
Err("tool exploded".to_string())
});
assert_eq!(err, Err("tool exploded".to_string()));
}
fn async_runner(
observed: std::collections::HashMap<CascadeTier, f64>,
ran: std::rc::Rc<std::cell::RefCell<Vec<CascadeTier>>>,
) -> impl FnMut(
CascadeTier,
) -> std::pin::Pin<
Box<dyn std::future::Future<Output = Result<TierResult, String>>>,
> {
move |tier| {
let confidence = *observed.get(&tier).unwrap_or(&0.0);
ran.borrow_mut().push(tier);
Box::pin(async move {
Ok(TierResult {
knowledge: format!("{tier:?} knowledge"),
confidence,
})
})
}
}
#[tokio::test]
async fn async_stops_at_cheapest_tier_that_observes_target() {
let ran = std::rc::Rc::new(std::cell::RefCell::new(Vec::new()));
let observed = [(CascadeTier::SelfReflect, 0.9)].into_iter().collect();
let run = run_cascade_async(
0.2,
&policy(0.8, 100.0),
async_runner(observed, ran.clone()),
)
.await
.unwrap();
assert_eq!(run.accepted_tier, Some(CascadeTier::SelfReflect));
assert_eq!(run.cost_spent, 1.0);
assert_eq!(*ran.borrow(), vec![CascadeTier::SelfReflect]);
}
#[tokio::test]
async fn async_escalates_on_low_observed_confidence() {
let ran = std::rc::Rc::new(std::cell::RefCell::new(Vec::new()));
let observed = [
(CascadeTier::SelfReflect, 0.4),
(CascadeTier::ToolVerify, 0.85),
]
.into_iter()
.collect();
let run = run_cascade_async(
0.2,
&policy(0.8, 100.0),
async_runner(observed, ran.clone()),
)
.await
.unwrap();
assert_eq!(run.accepted_tier, Some(CascadeTier::ToolVerify));
assert_eq!(run.cost_spent, 6.0);
assert_eq!(
*ran.borrow(),
vec![CascadeTier::SelfReflect, CascadeTier::ToolVerify]
);
}
#[tokio::test]
async fn async_propagates_runner_error() {
let err = run_cascade_async(0.2, &policy(0.8, 100.0), |_tier| async {
Err::<TierResult, String>("tool exploded".to_string())
})
.await;
assert_eq!(err, Err("tool exploded".to_string()));
}
}