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
62 changes: 57 additions & 5 deletions crates/libsy/src/algorithms/escalation.rs
Original file line number Diff line number Diff line change
Expand Up @@ -436,9 +436,10 @@ mod tests {
use std::sync::Arc;

use parking_lot::Mutex;
use serde_json::json;
use switchyard_protocol::{
ContentBlock, LlmClientError, LlmResponse, LlmResponseChunk, Metadata, ModelId, Request,
Response, completion_text, text_request, text_response,
ContentBlock, LlmClientError, LlmRequest, LlmResponse, LlmResponseChunk, Metadata, ModelId,
Request, Response, ToolCall, ToolResult, completion_text, text_request, text_response,
};

use super::*;
Expand Down Expand Up @@ -654,6 +655,8 @@ mod tests {
async fn config_overrides_the_packaged_prompt() -> Result<()> {
let prompts = Arc::new(Mutex::new(Vec::new()));
let recorded = Arc::clone(&prompts);
let judged = Arc::new(Mutex::new(Vec::new()));
let recorded_messages = Arc::clone(&judged);
let serve = move |target: ModelId, request: Request| {
if target == "judge" {
let prompt = request
Expand All @@ -667,6 +670,9 @@ mod tests {
})
});
recorded.lock().extend(prompt);
recorded_messages
.lock()
.push(format!("{:?}", request.llm_request.messages));
std::future::ready(Ok(reply(
r#"{"escalate":false,"category":"none","new_evidence":false,"reason":"progressing"}"#,
)))
Expand All @@ -675,17 +681,63 @@ mod tests {
}
};
let router = Arc::new(LlmTaskClassifier::new(LlmClassifierConfig::Escalation {
contract: ClassifierContractConfig::default().with_prompt("Custom trajectory rubric."),
contract: ClassifierContractConfig::default()
.with_prompt("Custom trajectory rubric.")
.with_prompt_suffix("Tool-step rubric."),
config: EscalationJudgeConfig {
confirmations: 1,
..EscalationJudgeConfig::default()
},
max_output_tokens: DEFAULT_JUDGE_MAX_OUTPUT_TOKENS,
})?);

test_drive_with_models(router, classify_request(), runtime_models(), serve).await?;
let tool_step = Request {
llm_request: LlmRequest {
messages: vec![
Message {
role: Role::User,
content: vec![ContentBlock::Text {
text: "fix the build".to_string(),
}],
},
Message {
role: Role::Assistant,
content: vec![ContentBlock::ToolCall(ToolCall {
id: "call-1".to_string(),
name: "bash".to_string(),
arguments: json!({"cmd": "cargo build"}),
})],
},
Message {
role: Role::Tool,
content: vec![ContentBlock::ToolResult(ToolResult {
tool_call_id: "call-1".to_string(),
content: vec![ContentBlock::Text {
text: "error[E0432]: unresolved import".to_string(),
}],
is_error: Some(true),
})],
},
],
..text_request(Some("auto".to_string()), "unused")
},
raw_request: None,
metadata: None,
};

test_drive_with_models(router, tool_step, runtime_models(), serve).await?;

assert_eq!(&*prompts.lock(), &["Custom trajectory rubric."]);
assert_eq!(
&*prompts.lock(),
&["Custom trajectory rubric.\n\nTool-step rubric."]
);
let judged = judged.lock();
assert!(
judged
.iter()
.any(|text| text.contains("error[E0432]: unresolved import")),
"judge never saw the tool result: {judged:?}"
);
Ok(())
}

Expand Down
113 changes: 112 additions & 1 deletion crates/libsy/src/algorithms/util/classifier_contract.rs
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,8 @@ pub struct ClassifierContractConfig {
#[serde(default)]
prompt: Option<String>,
#[serde(default)]
prompt_suffix: Option<String>,
#[serde(default)]
response_format_type: ClassifierResponseFormat,
}

Expand All @@ -43,6 +45,17 @@ impl ClassifierContractConfig {
self.prompt.as_deref()
}

/// Appends guidance to the packaged or overridden classifier prompt.
pub fn with_prompt_suffix(mut self, suffix: impl Into<String>) -> Self {
self.prompt_suffix = Some(suffix.into());
self
}

/// Returns the configured prompt suffix.
pub fn prompt_suffix(&self) -> Option<&str> {
self.prompt_suffix.as_deref()
}

/// Selects the provider-side structured-output mode.
pub fn with_response_format_type(
mut self,
Expand Down Expand Up @@ -77,7 +90,23 @@ impl ClassifierContract {
default_prompt: &str,
response_format_json: &str,
) -> Result<Self> {
let prompt_template = config.prompt().unwrap_or(default_prompt);
let base_prompt = config.prompt().unwrap_or(default_prompt);
// Validate the base first so a blank prompt cannot be masked by a suffix. A blank suffix
// leaves the prompt byte-identical; the suffix joins the template, so it precedes any
// schema block appended below.
validate_prompt(base_prompt)?;
let composed;
let prompt_template = match config
.prompt_suffix()
.map(str::trim)
.filter(|suffix| !suffix.is_empty())
{
Some(suffix) => {
composed = format!("{base_prompt}\n\n{suffix}");
composed.as_str()
}
None => base_prompt,
};
let response_format: Value =
serde_json::from_str(response_format_json).map_err(|error| {
LibsyError::AlgorithmError {
Expand Down Expand Up @@ -285,6 +314,88 @@ mod tests {
);
}

#[test]
fn prompt_suffix_composes_with_packaged_prompt() {
const SCHEMA: &str = r#"{"json_schema":{"schema":{"type":"object"}}}"#;
let default = ClassifierContractConfig::default;
// (name, config, response_format, expected system prompt or error substring)
let cases: Vec<(
&str,
ClassifierContractConfig,
&str,
std::result::Result<&str, &str>,
)> = vec![
("unset", default(), SCHEMA, Ok("packaged prompt")),
(
"suffix",
default().with_prompt_suffix("Step."),
SCHEMA,
Ok("packaged prompt\n\nStep."),
),
(
"override plus suffix",
default()
.with_prompt("Custom.")
.with_prompt_suffix(" Step.\n"),
SCHEMA,
Ok("Custom.\n\nStep."),
),
(
"blank suffix equals unset",
default().with_prompt_suffix(" "),
SCHEMA,
Ok("packaged prompt"),
),
(
"schema placeholder in suffix",
default().with_prompt_suffix("see {{RESPONSE_SCHEMA}}"),
SCHEMA,
Err("Switchyard supplies the schema automatically"),
),
(
"blank prompt plus suffix",
default().with_prompt(" ").with_prompt_suffix("x"),
SCHEMA,
Err("prompt must not be empty"),
),
];
for (name, config, format, expected) in cases {
let got = ClassifierContract::from_config(&config, "packaged prompt", format);
match expected {
Ok(prompt) => assert_eq!(
got.unwrap_or_else(|e| panic!("{name}: {e}"))
.system_prompt(),
prompt,
"{name}"
),
Err(needle) => {
let error = got
.err()
.unwrap_or_else(|| panic!("{name}: expected error"));
assert!(error.to_string().contains(needle), "{name}: {error}");
}
}
}

let contract = ClassifierContract::from_config(
&default()
.with_prompt_suffix("Step.")
.with_response_format_type(ClassifierResponseFormat::JsonObject),
"packaged prompt",
SCHEMA,
)
.expect("json_object contract");
let prompt = contract.system_prompt();
let suffix = prompt.find("Step.").expect("suffix in prompt");
let schema = prompt
.find("Return exactly one JSON object")
.expect("schema block in prompt");
assert!(
suffix < schema,
"suffix must precede schema block: {prompt}"
);
}

#[test]
fn a_custom_contract_wraps_and_validates_its_inner_schema() -> Result<()> {
let contract = ClassifierContract::from_inner_schema(
Expand Down
38 changes: 30 additions & 8 deletions crates/switchyard-runner/src/algorithm.rs
Original file line number Diff line number Diff line change
Expand Up @@ -108,6 +108,7 @@ struct CapabilityClassifierRouteConfig {
message_hash_fallback: bool,
recent_turn_window: Option<usize>,
prompt: Option<String>,
prompt_suffix: Option<String>,
response_format_type: ClassifierResponseFormat,
max_output_tokens: u64,
}
Expand All @@ -118,6 +119,7 @@ struct EscalationClassifierRouteConfig {
strong_target: String,
weak_target: String,
prompt: Option<String>,
prompt_suffix: Option<String>,
response_format_type: ClassifierResponseFormat,
max_output_tokens: u64,
judge: EscalationJudgeConfig,
Expand Down Expand Up @@ -251,6 +253,8 @@ pub struct LlmClassifierRouteConfig {
pub recent_turn_window: Option<usize>,
/// Replaces the packaged judge prompt. Required in custom mode.
pub prompt: Option<String>,
/// Appended to the packaged or overridden judge prompt. Not allowed in custom mode.
pub prompt_suffix: Option<String>,
/// How the judge is asked for structured output. Use `json_object` when the
/// provider cannot do JSON Schema.
pub response_format_type: ClassifierResponseFormat,
Expand Down Expand Up @@ -529,7 +533,7 @@ impl StageClassifierConfig {
judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig {
base_threshold: self.base_threshold,
threshold_step: self.threshold_step,
contract: classifier_contract(self.prompt.as_deref())
contract: classifier_contract(self.prompt.as_deref(), None)
.with_response_format_type(self.response_format_type),
max_output_tokens: self.max_output_tokens,
}),
Expand Down Expand Up @@ -890,6 +894,7 @@ impl LlmClassifierRouteConfig {
message_hash_fallback,
recent_turn_window,
prompt,
prompt_suffix,
response_format_type,
max_output_tokens,
escalation,
Expand Down Expand Up @@ -952,6 +957,7 @@ impl LlmClassifierRouteConfig {
message_hash_fallback: *message_hash_fallback,
recent_turn_window: *recent_turn_window,
prompt: prompt.clone(),
prompt_suffix: prompt_suffix.clone(),
response_format_type: *response_format_type,
max_output_tokens: *max_output_tokens,
},
Expand Down Expand Up @@ -995,6 +1001,7 @@ impl LlmClassifierRouteConfig {
weak_target,
)?,
prompt: prompt.clone(),
prompt_suffix: prompt_suffix.clone(),
response_format_type: *response_format_type,
max_output_tokens: *max_output_tokens,
judge: required_classifier_field(route_name, "escalation", escalation)?,
Expand All @@ -1008,6 +1015,7 @@ impl LlmClassifierRouteConfig {
|| base_threshold.is_some()
|| threshold_step.is_some()
|| escalation.is_some()
|| prompt_suffix.is_some()
|| *response_format_type != ClassifierResponseFormat::JsonSchema
{
return Err(AlgorithmConfigError::new(format!(
Expand Down Expand Up @@ -1247,8 +1255,11 @@ fn build_algorithm(
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),
contract: classifier_contract(
config.prompt.as_deref(),
config.prompt_suffix.as_deref(),
)
.with_response_format_type(config.response_format_type),
max_output_tokens: config.max_output_tokens,
}),
fail_open: config.fail_open,
Expand All @@ -1262,7 +1273,10 @@ fn build_algorithm(
}
LlmClassifierModeConfig::Escalation(config) => {
LlmTaskClassifier::new(LlmClassifierConfig::Escalation {
contract: classifier_contract(config.prompt.as_deref())
contract: classifier_contract(
config.prompt.as_deref(),
config.prompt_suffix.as_deref(),
)
.with_response_format_type(config.response_format_type),
config: config.judge,
max_output_tokens: config.max_output_tokens,
Expand Down Expand Up @@ -1491,10 +1505,18 @@ const fn default_fail_open() -> bool {
true
}

fn classifier_contract(prompt: Option<&str>) -> ClassifierContractConfig {
prompt.map_or_else(ClassifierContractConfig::default, |prompt| {
ClassifierContractConfig::default().with_prompt(prompt)
})
fn classifier_contract(
prompt: Option<&str>,
prompt_suffix: Option<&str>,
) -> ClassifierContractConfig {
let mut contract = ClassifierContractConfig::default();
if let Some(prompt) = prompt {
contract = contract.with_prompt(prompt);
}
if let Some(suffix) = prompt_suffix {
contract = contract.with_prompt_suffix(suffix);
}
contract
}

fn default_classifier_max_output_tokens() -> u64 {
Expand Down
Loading