Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
286 changes: 269 additions & 17 deletions crates/libsy/src/algorithms/llm_class.rs
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,10 @@ use serde::{Deserialize, Deserializer};
use serde_json::Value;
use switchyard_protocol::{Category, ContentBlock, Message, Role};

mod decision;
use decision::DecisionClassifier;
pub use decision::DecisionJudgeConfig;

use super::escalation;
use super::fall_through::FallThrough;
use super::util::DEFAULT_JUDGE_MAX_OUTPUT_TOKENS;
Expand Down Expand Up @@ -326,6 +330,8 @@ pub struct TaskClassifierConfig {
pub enum CapabilityJudgeConfig {
/// A structured LLM verdict with a solve probability and capability boundary.
Llm(LlmCapabilityConfig),
/// A Choice score for the capable candidate succeeding where the efficient one fails.
Decision(DecisionJudgeConfig),
}

impl Default for CapabilityJudgeConfig {
Expand Down Expand Up @@ -467,6 +473,7 @@ impl TaskClassifierConfig {
fn validate(&self) -> Result<()> {
match &self.judge {
CapabilityJudgeConfig::Llm(config) => config.validate()?,
CapabilityJudgeConfig::Decision(config) => config.validate()?,
}
// Only `every_request` is rejected: it retains no target, so a fallback identity has nothing to key on. Both retaining triggers can key the retained target on a message hash when the caller sends no session id.
if self.message_hash_fallback && self.classify_trigger == ClassifyTrigger::EveryRequest {
Expand Down Expand Up @@ -679,25 +686,29 @@ impl LlmTaskClassifier {

fn build_capability(config: TaskClassifierConfig) -> Result<Self> {
config.validate()?;
let CapabilityJudgeConfig::Llm(judge) = &config.judge;
let contract = Self::load_capability_contract(&judge.contract)?;
let classify_trigger = config.classify_trigger;
let message_hash_fallback = config.message_hash_fallback;
let classifier: Arc<dyn Classifier<State>> = Arc::new(
JudgeClassifier::new(
StructuredJudge::new(
TaskInput {
recent_turn_window: config.recent_turn_window,
},
contract,
SerdeDecoder::new(),
JudgeRuntimeConfig::new(judge.max_output_tokens)?,
),
TaskClassifierPolicy::new(judge),
)
.with_error_recovery(config.fail_open)
.with_evidence(capability_evidence),
);
let input = TaskInput {
recent_turn_window: config.recent_turn_window,
};
let classifier: Arc<dyn Classifier<State>> = match config.judge {
CapabilityJudgeConfig::Llm(judge) => Arc::new(
JudgeClassifier::new(
StructuredJudge::new(
input,
Self::load_capability_contract(&judge.contract)?,
SerdeDecoder::new(),
JudgeRuntimeConfig::new(judge.max_output_tokens)?,
),
TaskClassifierPolicy::new(&judge),
)
.with_error_recovery(config.fail_open)
.with_evidence(capability_evidence),
),
CapabilityJudgeConfig::Decision(judge) => {
Arc::new(DecisionClassifier::new(judge, input, config.fail_open)?)
}
};
Self::from_classifier(
classifier,
ClassifierRouteConfig {
Expand Down Expand Up @@ -1270,6 +1281,247 @@ mod tests {
Ok(())
}

#[tokio::test]
async fn decision_judge_routes_on_relative_advantage() -> Result<()> {
use crate::{Call, RuntimeModels, drive};
use serde_json::json;
use std::collections::BTreeMap;
use std::sync::atomic::{AtomicUsize, Ordering};
use switchyard_protocol::{
BooleanEstimate, DecisionAnswer, DecisionResponse, DecisionValue, Probability,
ProviderConfidence,
};

let settings = DecisionJudgeConfig {
cutoff: 0.4,
instructions: None,
candidates: BTreeMap::from([
("a".into(), "capable".into()),
("b".into(), "efficient".into()),
("c".into(), "third".into()),
]),
evidence: json!({
"candidate_descriptions": {"a": "Profile A", "b": "Profile B", "c": "Profile C"},
"reference_cases": [{"task": "Independent example", "observed_solved": {"a": true, "b": null, "c": false}}],
"summaries": [],
}),
};
let mut request = classify_session_request();
request.llm_request.messages.push(Message {
role: Role::Assistant,
content: vec![
ContentBlock::Reasoning {
text: "private reasoning".into(),
signature: None,
details: vec![],
},
ContentBlock::Text {
text: "visible progress".into(),
},
],
});
request
.llm_request
.messages
.push(Message::text(Role::User, "latest follow-up"));
let original_messages = request.llm_request.messages.clone();
let expected_task = json!([
original_messages[0].clone(),
Message::text(Role::Assistant, "visible progress"),
Message::text(Role::User, "latest follow-up"),
]);

for (name, score, fail_open, expected) in [
("above", 0.45, true, Some("capable")),
("equal", 0.4, true, Some("efficient")),
("below", 0.39, false, Some("efficient")),
("missing answer", 0.0, true, Some("capable")),
("no distribution", 0.0, true, Some("capable")),
("wrong type", 0.0, true, Some("capable")),
("out of range", 1.1, true, Some("capable")),
("nonfinite", f64::NAN, false, Some("capable")),
("provider error", 0.0, true, Some("capable")),
("provider error", 0.0, false, None),
("dropped reply", 0.0, true, Some("capable")),
("missing candidate", 0.0, true, None),
] {
let mut config = settings.clone();
if name == "missing candidate" {
config.candidates.remove("a");
}
if name == "below" {
config.instructions = Some(
json!({"policy": "Judge capable-only success using the supplied evidence."}),
);
}
let override_instructions = config.instructions.clone();
let router: Arc<dyn Algorithm> =
Arc::new(LlmTaskClassifier::new(LlmClassifierConfig::Capability {
config: TaskClassifierConfig {
judge: CapabilityJudgeConfig::Decision(config),
fail_open,
recent_turn_window: Some(3),
classify_trigger: ClassifyTrigger::NewSession,
..TaskClassifierConfig::default()
},
})?);
let calls = AtomicUsize::new(0);
let serve = |call| {
let calls = &calls;
let expected_task = &expected_task;
let settings = &settings;
let override_instructions = &override_instructions;
async move {
let Call::Decision(call) = call else {
panic!("judge must use the decision step")
};
calls.fetch_add(1, Ordering::SeqCst);
assert_eq!(call.model, "judge");
assert_eq!(call.request.model, Some("judge".into()));
assert_eq!(
call.request.context,
json!({
"task": expected_task,
"candidates": ["a", "b", "c"],
"comparison": {"capable": "a", "efficient": "b"},
"evidence": settings.evidence,
})
);
assert_eq!(call.request.questions.len(), 1);
let question = &call.request.questions["route"];
let switchyard_protocol::DecisionKind::Choice { options } = &question.kind
else {
panic!("expected a choice")
};
assert_eq!(
options
.iter()
.map(|option| option.id.as_str())
.collect::<Vec<_>>(),
["advantage", "no_advantage"]
);
if let Some(instructions) = override_instructions {
assert_eq!(&question.instructions, instructions);
} else {
for key in ["setting", "evidence", "comparison", "boundary", "policy"] {
assert!(question.instructions[key].as_str().is_some());
}
}
if name == "provider error" {
return call.respond(Err(LibsyError::client_call(
"judge",
LlmClientError::General("private provider body".into()),
)));
}
if name == "dropped reply" {
drop(call);
return Ok(());
}
let value = if name == "wrong type" {
DecisionValue::Boolean(BooleanEstimate::Value(true))
} else {
DecisionValue::Choice {
selected: "no_advantage".into(),
probabilities: (name != "no distribution").then(|| {
BTreeMap::from([
("advantage".into(), Probability(score)),
("no_advantage".into(), Probability(1.0 - score)),
])
}),
}
};
let answers = if name == "missing answer" {
BTreeMap::new()
} else {
BTreeMap::from([(
"route".into(),
DecisionAnswer {
value,
provider_confidence: Some(ProviderConfidence(0.99)),
},
)])
};
call.respond(Ok(DecisionResponse {
id: None,
model: Some("provider-judge".into()),
answers,
usage: Default::default(),
}))
}
};
let models = Arc::new(RuntimeModels::new(runtime_models()));
let result = drive(router.clone(), request.clone(), models.clone(), &serve).await;
assert_eq!(
calls.load(Ordering::SeqCst),
usize::from(name != "missing candidate"),
"{name}"
);
let Some(expected) = expected else {
if name == "missing candidate" {
assert!(
matches!(result, Err(LibsyError::AlgorithmError { message }) if message.contains("candidate is missing"))
);
} else {
assert!(
matches!(result, Err(LibsyError::ClientCall { .. })),
"{name}"
);
}
continue;
};
let outcome = result?;
assert_eq!(outcome.selected_model_id()?, expected, "{name}");
assert!(outcome.response.is_none());
assert_eq!(outcome.request.llm_request.messages, original_messages);
let evidence = outcome
.metadata
.and_then(|metadata| metadata.evidence)
.expect("routing evidence");
if matches!(name, "above" | "equal" | "below") {
assert_eq!(evidence["source"], "decision_classifier");
assert_eq!(evidence["verdict"], "relative_advantage");
assert_eq!(evidence["threshold"], settings.cutoff);
assert_eq!(evidence["score"], score);
} else {
let reason = match name {
"provider error" => "client_error",
"dropped reply" => "call_error",
_ => "invalid_verdict",
};
assert_eq!(
evidence,
json!({"source": "fail_open", "reason_code": reason})
);
}
let retained = drive(router, request.clone(), models, &serve).await?;
assert_eq!(retained.selected_model_id()?, expected);
assert_eq!(
calls.load(Ordering::SeqCst),
1,
"affinity should skip the judge: {name}"
);
}
for cutoff in [-0.1, 1.1, f64::NAN] {
let mut config = settings.clone();
config.cutoff = cutoff;
assert!(
LlmTaskClassifier::new(LlmClassifierConfig::Capability {
config: TaskClassifierConfig {
judge: CapabilityJudgeConfig::Decision(config),
..TaskClassifierConfig::default()
},
})
.is_err()
);
}
let mut duplicate = settings.clone();
duplicate
.candidates
.insert("duplicate-a".into(), "capable".into());
assert!(duplicate.validate().is_err());
Ok(())
}

#[test]
fn the_threshold_boundary_is_inclusive() -> Result<()> {
let policy = policy();
Expand Down
Loading
Loading