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
124 changes: 99 additions & 25 deletions crates/switchyard-runner/src/algorithm.rs
Original file line number Diff line number Diff line change
Expand Up @@ -12,12 +12,14 @@ use std::sync::Arc;
use libsy::{
AdvisorGate, AdvisorGateConfig, Algorithm, CapabilityJudgeConfig, ClassifierContractConfig,
ClassifierResponseFormat, ClassifyTrigger, CompositeRouter, CompositeRouterConfig,
CustomClassifierConfig, CustomClassifierPolicy, EscalationJudgeConfig, GateTrigger,
HandoffNoteConfig, LlmCapabilityConfig, LlmClassifierConfig, LlmFallback, LlmTaskClassifier,
Noop, Passthrough, PickerMode, PlanExecute, PlanExecuteConfig, Random, StageRouter,
StageRouterConfig, SubagentRouter, SubagentRouterConfig, TaskClassifierConfig, ToolSemantics,
CustomClassifierConfig, CustomClassifierPolicy, DecisionJudgeConfig, EscalationJudgeConfig,
GateTrigger, HandoffNoteConfig, LlmCapabilityConfig, LlmClassifierConfig, LlmFallback,
LlmTaskClassifier, Noop, Passthrough, PickerMode, PlanExecute, PlanExecuteConfig, Random,
StageRouter, StageRouterConfig, SubagentRouter, SubagentRouterConfig, TaskClassifierConfig,
ToolSemantics,
};
use serde::Deserialize;
use serde_json::Value;
use switchyard_protocol::{Category, ModelId};

/// Error returned when an algorithm description cannot be constructed.
Expand Down Expand Up @@ -102,14 +104,16 @@ struct CapabilityClassifierRouteConfig {
fail_open: bool,
strong_target: String,
weak_target: String,
base_threshold: f64,
threshold_step: f64,
judge: CapabilityJudgeRouteConfig,
classify_trigger: ClassifyTrigger,
message_hash_fallback: bool,
recent_turn_window: Option<usize>,
prompt: Option<String>,
response_format_type: ClassifierResponseFormat,
max_output_tokens: u64,
}

#[derive(Clone, Debug)]
enum CapabilityJudgeRouteConfig {
Llm(LlmCapabilityConfig),
Decision(DecisionJudgeRouteConfig),
}

#[derive(Clone, Debug)]
Expand Down Expand Up @@ -219,6 +223,20 @@ impl CategoryModelConfig {
}
}

/// Relative-advantage judgment using a decision model instead of an LLM prompt.
#[derive(Clone, Debug, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct DecisionJudgeRouteConfig {
/// Route capable only when its advantage score is strictly above this cutoff.
pub cutoff: f64,
/// Replaces the default relative-advantage instructions.
pub instructions: Option<Value>,
/// Anonymous evidence labels mapped to configured LLM target names.
pub candidates: BTreeMap<String, String>,
/// Candidate descriptions and reference outcomes, passed unchanged to the judge.
pub evidence: Value,
}

/// Settings for an `llm_classifier` route. Which fields are required depends on
/// the [`ClassifierMode`]; using a field from the wrong mode is an error.
#[derive(Clone, Debug, Default, Deserialize)]
Expand Down Expand Up @@ -260,6 +278,8 @@ pub struct LlmClassifierRouteConfig {
/// Escalation mode: how many escalate verdicts latch the session, and how
/// much of the transcript the judge sees.
pub escalation: Option<EscalationJudgeConfig>,
/// Capability mode: use a decision target as the judge, with relative-advantage scoring.
pub decision: Option<DecisionJudgeRouteConfig>,
/// Custom mode: runtime model groups.
pub models: Option<CategoryModelConfig>,
/// Custom mode: category used when the judge fails or its verdict cannot be routed.
Expand Down Expand Up @@ -542,6 +562,16 @@ impl StageClassifierConfig {
}

impl AlgorithmSpec {
pub(crate) fn decision_judge(&self) -> Option<(&str, &DecisionJudgeRouteConfig)> {
match self {
Self::LlmClassifier { config, .. } => config
.decision
.as_ref()
.map(|decision| (config.classifier_target.as_str(), decision)),
_ => None,
}
}

/// Completion targets in algorithm order; judge-only targets are excluded.
pub fn routing_target_names(&self) -> Vec<&str> {
match self {
Expand Down Expand Up @@ -893,6 +923,7 @@ impl LlmClassifierRouteConfig {
response_format_type,
max_output_tokens,
escalation,
decision,
models,
default_target,
response_schema,
Expand All @@ -905,6 +936,12 @@ impl LlmClassifierRouteConfig {
(None, false) => ClassifierMode::Capability,
};

if decision.is_some() && !matches!(selected_mode, ClassifierMode::Capability) {
return Err(AlgorithmConfigError::new(format!(
"llm_classifier route {route_name}: decision requires capability mode"
)));
}

if !matches!(selected_mode, ClassifierMode::Capability) && *fail_open == Some(true) {
return Err(AlgorithmConfigError::new(format!(
"llm_classifier route {route_name}: fail_open = true requires capability mode"
Expand All @@ -928,6 +965,31 @@ impl LlmClassifierRouteConfig {
response_schema,
policy,
)?;
let judge = if let Some(decision) = decision {
if base_threshold.is_some()
|| threshold_step.is_some()
|| prompt.is_some()
|| *response_format_type != ClassifierResponseFormat::JsonSchema
|| *max_output_tokens != default_classifier_max_output_tokens()
{
return Err(AlgorithmConfigError::new(format!(
"llm_classifier route {route_name}: decision cannot use LLM judge settings; use decision.cutoff and decision.instructions"
)));
}
CapabilityJudgeRouteConfig::Decision(decision.clone())
} else {
CapabilityJudgeRouteConfig::Llm(LlmCapabilityConfig {
base_threshold: required_classifier_field(
route_name,
"base_threshold",
base_threshold,
)?,
threshold_step: threshold_step.unwrap_or_default(),
contract: classifier_contract(prompt.as_deref())
.with_response_format_type(*response_format_type),
max_output_tokens: *max_output_tokens,
})
};
Ok(LlmClassifierModeConfig::Capability(
CapabilityClassifierRouteConfig {
classifier_target: classifier_target.clone(),
Expand All @@ -942,18 +1004,10 @@ impl LlmClassifierRouteConfig {
"weak_target",
weak_target,
)?,
base_threshold: required_classifier_field(
route_name,
"base_threshold",
base_threshold,
)?,
threshold_step: threshold_step.unwrap_or_default(),
judge,
classify_trigger: *classify_trigger,
message_hash_fallback: *message_hash_fallback,
recent_turn_window: *recent_turn_window,
prompt: prompt.clone(),
response_format_type: *response_format_type,
max_output_tokens: *max_output_tokens,
},
))
}
Expand Down Expand Up @@ -1243,14 +1297,34 @@ fn build_algorithm(
let mode = classifier_config.validated_classifier_mode(route_name)?;
let algorithm = match mode {
LlmClassifierModeConfig::Capability(config) => {
let judge = match config.judge {
CapabilityJudgeRouteConfig::Llm(judge) => CapabilityJudgeConfig::Llm(judge),
CapabilityJudgeRouteConfig::Decision(judge) => {
for target in [&config.strong_target, &config.weak_target] {
if !judge.candidates.values().any(|name| name == target) {
return Err(AlgorithmConfigError::new(format!(
"llm_classifier route {route_name}: decision.candidates must include target {target}"
)));
}
}
let candidates = judge
.candidates
.into_iter()
.map(|(label, name)| {
resolve_target_model_id(route_name, &name, targets)
.map(|model| (label, model))
})
.collect::<AlgorithmResult<_>>()?;
CapabilityJudgeConfig::Decision(DecisionJudgeConfig {
cutoff: judge.cutoff,
instructions: judge.instructions,
candidates,
evidence: judge.evidence,
})
}
};
let classifier_config = TaskClassifierConfig {
judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig {
base_threshold: config.base_threshold,
threshold_step: config.threshold_step,
contract: classifier_contract(config.prompt.as_deref())
.with_response_format_type(config.response_format_type),
max_output_tokens: config.max_output_tokens,
}),
judge,
fail_open: config.fail_open,
classify_trigger: config.classify_trigger,
message_hash_fallback: config.message_hash_fallback,
Expand Down
Loading
Loading