Skip to content
Open
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
2 changes: 2 additions & 0 deletions crates/libsy-llm-client/tests/observability.rs
Original file line number Diff line number Diff line change
Expand Up @@ -774,6 +774,7 @@ async fn stateful_escalation_warns_once_without_a_session_id() -> switchyard_lib
contract: ClassifierContractConfig::default(),
config: EscalationJudgeConfig::default(),
max_output_tokens: 64,
judge_deadline_ms: None,
})?) as Arc<dyn Algorithm>;
let client = Arc::new(JudgeClient {
judge_model: "warning-judge".into(),
Expand Down Expand Up @@ -824,6 +825,7 @@ async fn deescalation_evidence_stays_pending_until_confirmed() -> switchyard_lib
..EscalationJudgeConfig::default()
},
max_output_tokens: 64,
judge_deadline_ms: None,
})?) as Arc<dyn Algorithm>;
let client = |verdict| {
Arc::new(JudgeClient {
Expand Down
8 changes: 7 additions & 1 deletion crates/libsy/src/algorithms/escalation.rs
Original file line number Diff line number Diff line change
Expand Up @@ -89,6 +89,7 @@ pub(super) fn build_classifier(
contract_config: ClassifierContractConfig,
config: EscalationJudgeConfig,
max_output_tokens: u64,
judge_deadline_ms: Option<u64>,
) -> Result<Arc<dyn Classifier<State>>> {
let confirmations = config.confirmations;
let deescalation = match config.deescalation {
Expand All @@ -99,6 +100,7 @@ pub(super) fn build_classifier(
config.clone(),
Some(EvaluationPhase::Strong),
max_output_tokens,
judge_deadline_ms,
)?,
}),
None => None,
Expand All @@ -110,6 +112,7 @@ pub(super) fn build_classifier(
config,
is_phase_aware.then_some(EvaluationPhase::Efficient),
max_output_tokens,
judge_deadline_ms,
)?,
confirmations,
deescalation,
Expand Down Expand Up @@ -361,7 +364,7 @@ impl Classifier<State> for EscalationClassifier {
let verdict = self
.escalation_judge
.verdict(state, &judge_request, driver, judge_models)
.await;
.await?;

let held = count(state, STREAK_KEY);
let held_category = category(state).map(str::to_string);
Expand Down Expand Up @@ -539,6 +542,7 @@ mod tests {
..EscalationJudgeConfig::default()
},
max_output_tokens: DEFAULT_JUDGE_MAX_OUTPUT_TOKENS,
judge_deadline_ms: None,
},
)?))
}
Expand All @@ -553,6 +557,7 @@ mod tests {
..EscalationJudgeConfig::default()
},
max_output_tokens: DEFAULT_JUDGE_MAX_OUTPUT_TOKENS,
judge_deadline_ms: None,
},
)?))
}
Expand Down Expand Up @@ -681,6 +686,7 @@ mod tests {
..EscalationJudgeConfig::default()
},
max_output_tokens: DEFAULT_JUDGE_MAX_OUTPUT_TOKENS,
judge_deadline_ms: None,
})?);

test_drive_with_models(router, classify_request(), runtime_models(), serve).await?;
Expand Down
209 changes: 205 additions & 4 deletions crates/libsy/src/algorithms/llm_class.rs
Original file line number Diff line number Diff line change
Expand Up @@ -352,6 +352,10 @@ pub struct LlmCapabilityConfig {
pub contract: ClassifierContractConfig,
/// Maximum completion tokens available to the classifier verdict.
pub max_output_tokens: u64,
/// Whole-consultation bound on the judge call, in milliseconds. Covers the model
/// call and the response drain; on expiry the judge is treated as unavailable and
/// follows the route's `fail_open` setting. `None` is unbounded.
pub judge_deadline_ms: Option<u64>,
}

impl Default for LlmCapabilityConfig {
Expand All @@ -361,6 +365,7 @@ impl Default for LlmCapabilityConfig {
threshold_step: 0.0,
contract: ClassifierContractConfig::default(),
max_output_tokens: DEFAULT_JUDGE_MAX_OUTPUT_TOKENS,
judge_deadline_ms: None,
}
}
}
Expand All @@ -386,6 +391,8 @@ struct TaskClassifierConfigWire {
response_format_type: ClassifierResponseFormat,
#[serde(default = "default_judge_max_output_tokens")]
max_output_tokens: u64,
#[serde(default)]
judge_deadline_ms: Option<u64>,
}

impl<'de> Deserialize<'de> for TaskClassifierConfig {
Expand All @@ -405,6 +412,7 @@ impl<'de> Deserialize<'de> for TaskClassifierConfig {
threshold_step: wire.threshold_step,
contract,
max_output_tokens: wire.max_output_tokens,
judge_deadline_ms: wire.judge_deadline_ms,
}),
fail_open: wire.fail_open,
classify_trigger: wire.classify_trigger,
Expand Down Expand Up @@ -523,6 +531,8 @@ pub struct CustomClassifierConfig {
pub recent_turn_window: Option<usize>,
/// Maximum completion tokens available to the classifier verdict.
pub max_output_tokens: u64,
/// Whole-consultation bound on the judge call, in milliseconds. `None` is unbounded.
pub judge_deadline_ms: Option<u64>,
}

impl CustomClassifierConfig {
Expand All @@ -540,6 +550,7 @@ impl CustomClassifierConfig {
message_hash_fallback: false,
recent_turn_window: None,
max_output_tokens: DEFAULT_JUDGE_MAX_OUTPUT_TOKENS,
judge_deadline_ms: None,
}
}

Expand Down Expand Up @@ -652,6 +663,8 @@ pub enum LlmClassifierConfig {
config: EscalationJudgeConfig,
/// Maximum completion tokens available to the escalation verdict.
max_output_tokens: u64,
/// Whole-consultation bound on the escalation judge call, in milliseconds.
judge_deadline_ms: Option<u64>,
},
/// Routes among model categories using a user-supplied schema and policy.
Custom {
Expand All @@ -676,7 +689,8 @@ impl LlmTaskClassifier {
contract,
config,
max_output_tokens,
} => Self::build_escalation(contract, config, max_output_tokens),
judge_deadline_ms,
} => Self::build_escalation(contract, config, max_output_tokens, judge_deadline_ms),
LlmClassifierConfig::Custom {
default_target,
config,
Expand All @@ -698,7 +712,8 @@ impl LlmTaskClassifier {
input,
Self::load_capability_contract(&judge.contract)?,
SerdeDecoder::new(),
JudgeRuntimeConfig::new(judge.max_output_tokens)?,
JudgeRuntimeConfig::new(judge.max_output_tokens)?
.with_deadline_ms(judge.judge_deadline_ms)?,
),
TaskClassifierPolicy::new(&judge),
)
Expand Down Expand Up @@ -729,6 +744,7 @@ impl LlmTaskClassifier {
message_hash_fallback,
recent_turn_window,
max_output_tokens,
judge_deadline_ms,
} = config;
let contract = ClassifierContract::from_inner_schema(&prompt, response_schema)?;
let policy = match policy {
Expand All @@ -741,7 +757,7 @@ impl LlmTaskClassifier {
TaskInput { recent_turn_window },
contract,
JsonSchemaDecoder::new(),
JudgeRuntimeConfig::new(max_output_tokens)?,
JudgeRuntimeConfig::new(max_output_tokens)?.with_deadline_ms(judge_deadline_ms)?,
),
policy,
));
Expand All @@ -760,8 +776,14 @@ impl LlmTaskClassifier {
contract_config: ClassifierContractConfig,
config: EscalationJudgeConfig,
max_output_tokens: u64,
judge_deadline_ms: Option<u64>,
) -> Result<Self> {
let inner = escalation::build_classifier(contract_config, config, max_output_tokens)?;
let inner = escalation::build_classifier(
contract_config,
config,
max_output_tokens,
judge_deadline_ms,
)?;
Ok(Self {
route: FallThrough::<State>::new_with_state()
.with_name(ALGORITHM_NAME)
Expand Down Expand Up @@ -1008,6 +1030,60 @@ mod tests {
}
}

/// The judge stalls before answering; every other target answers normally.
fn slow_judge(delay: std::time::Duration) -> impl Serve {
move |model: ModelId, request: Request| {
let model = model.to_string();
let slow = model == "judge";
async move {
if slow {
tokio::time::sleep(delay).await;
}
Ok(Response {
llm_response: LlmResponse::Agg(text_response(
None,
format!("answer from {model}"),
)),
metadata: request.metadata,
upstream_headers: http::HeaderMap::new(),
})
}
}
}

/// The judge accepts the call but its stream never carries a chunk.
fn stalled_judge_stream() -> impl Serve {
use futures::StreamExt;
|model: ModelId, request: Request| async move {
let model = model.to_string();
let llm_response = if model == "judge" {
LlmResponse::Stream(futures::stream::pending().boxed())
} else {
LlmResponse::Agg(text_response(None, format!("answer from {model}")))
};
Ok(Response {
llm_response,
metadata: request.metadata,
upstream_headers: http::HeaderMap::new(),
})
}
}

fn deadline_router(fail_open: bool) -> Result<Arc<LlmTaskClassifier>> {
Ok(Arc::new(LlmTaskClassifier::new(
LlmClassifierConfig::Capability {
config: TaskClassifierConfig {
judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig {
judge_deadline_ms: Some(10),
..llm_config(TEST_THRESHOLD)
}),
fail_open,
..TaskClassifierConfig::default()
},
},
)?))
}

fn router() -> Result<Arc<LlmTaskClassifier>> {
Ok(Arc::new(LlmTaskClassifier::new(
LlmClassifierConfig::Capability {
Expand Down Expand Up @@ -1067,6 +1143,131 @@ mod tests {
Ok(())
}

#[tokio::test]
async fn a_judge_past_its_deadline_routes_capable() -> Result<()> {
let (selected_model, response) = test_drive_with_models(
deadline_router(true)?,
classify_request(),
runtime_models(),
slow_judge(std::time::Duration::from_millis(500)),
)
.await?;

assert_eq!(selected_model, "capable");
assert_eq!(
response.llm_response.as_agg().map(completion_text),
Some("answer from capable".to_string())
);
Ok(())
}

#[tokio::test]
async fn a_stalled_judge_stream_is_cut_by_the_deadline() -> Result<()> {
// The judge returned headers promptly, so only a bound on the whole
// consultation — the drain included — can end this turn.
let (selected_model, _) = test_drive_with_models(
deadline_router(true)?,
classify_request(),
runtime_models(),
stalled_judge_stream(),
)
.await?;

assert_eq!(selected_model, "capable");
Ok(())
}

#[tokio::test]
async fn a_judge_past_its_deadline_stops_when_not_failing_open() -> Result<()> {
let outcome = test_drive_with_models(
deadline_router(false)?,
classify_request(),
runtime_models(),
slow_judge(std::time::Duration::from_millis(500)),
)
.await;
let error = outcome
.err()
.expect("a deadline expiry with fail_open = false must stop the request");

assert!(
matches!(
error,
LibsyError::ClientCall {
source: LlmClientError::Timeout { .. },
..
}
),
"expected a client timeout, got {error:?}"
);
Ok(())
}

#[test]
fn a_zero_judge_deadline_is_rejected_in_every_mode() {
let capability = LlmClassifierConfig::Capability {
config: TaskClassifierConfig {
judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig {
judge_deadline_ms: Some(0),
..llm_config(TEST_THRESHOLD)
}),
..TaskClassifierConfig::default()
},
};
assert!(
LlmTaskClassifier::new(capability).is_err(),
"capability mode must reject judge_deadline_ms = 0"
);

let mut custom = CustomClassifierConfig::new(
"Route by topic.",
serde_json::json!({"type": "object"}),
CustomClassifierPolicy::target_selector("/target"),
);
custom.judge_deadline_ms = Some(0);
assert!(
LlmTaskClassifier::new(LlmClassifierConfig::Custom {
default_target: Category::Capable,
config: custom,
})
.is_err(),
"custom mode must reject judge_deadline_ms = 0"
);

let escalation = LlmClassifierConfig::Escalation {
contract: ClassifierContractConfig::default(),
config: EscalationJudgeConfig::default(),
max_output_tokens: DEFAULT_JUDGE_MAX_OUTPUT_TOKENS,
judge_deadline_ms: Some(0),
};
assert!(
LlmTaskClassifier::new(escalation).is_err(),
"escalation mode must reject judge_deadline_ms = 0"
);
}

#[test]
fn classifier_config_parses_a_judge_deadline() {
let config: TaskClassifierConfig = serde_json::from_value(serde_json::json!({
"base_threshold": 0.5,
"judge_deadline_ms": 250,
}))
.expect("a configured judge deadline should parse");
let CapabilityJudgeConfig::Llm(judge) = config.judge else {
panic!("expected the LLM judge variant");
};
assert_eq!(judge.judge_deadline_ms, Some(250));

let config: TaskClassifierConfig = serde_json::from_value(serde_json::json!({
"base_threshold": 0.5,
}))
.expect("an unset judge deadline should parse");
let CapabilityJudgeConfig::Llm(judge) = config.judge else {
panic!("expected the LLM judge variant");
};
assert_eq!(judge.judge_deadline_ms, None);
}

#[tokio::test]
async fn classifier_judges_each_request_without_affinity() -> Result<()> {
let recorder = Arc::new(Recorder::default());
Expand Down
3 changes: 2 additions & 1 deletion crates/libsy/src/algorithms/util/escalation.rs
Original file line number Diff line number Diff line change
Expand Up @@ -283,6 +283,7 @@ pub(crate) fn build_judge(
config: EscalationJudgeConfig,
phase: Option<EvaluationPhase>,
max_output_tokens: u64,
judge_deadline_ms: Option<u64>,
) -> Result<JudgeClassifier<EscalationJudge, EscalationPolicy>> {
config.validate()?;
let contract = build_contract(contract_config, phase.is_some())?;
Expand All @@ -291,7 +292,7 @@ pub(crate) fn build_judge(
EscalationInput { config, phase },
contract,
SerdeDecoder::new(),
JudgeRuntimeConfig::new(max_output_tokens)?,
JudgeRuntimeConfig::new(max_output_tokens)?.with_deadline_ms(judge_deadline_ms)?,
),
EscalationPolicy { phase },
)
Expand Down
Loading
Loading