From 4f51f15ab8824f02a650e663223f8e884c481ce0 Mon Sep 17 00:00:00 2001 From: ChethanUK Date: Thu, 1 Oct 2026 20:52:36 +0200 Subject: [PATCH] feat(classifier): add prompt_suffix to extend judge prompts Signed-off-by: ChethanUK --- crates/libsy/src/algorithms/escalation.rs | 62 +++++++++- .../algorithms/util/classifier_contract.rs | 113 +++++++++++++++++- crates/switchyard-runner/src/algorithm.rs | 38 ++++-- crates/switchyard-runner/src/config.rs | 36 +++++- docs/reference/toml_schema.md | 2 + .../escalation_router_routing.md | 5 +- .../llm_classifier_routing.md | 1 + 7 files changed, 239 insertions(+), 18 deletions(-) diff --git a/crates/libsy/src/algorithms/escalation.rs b/crates/libsy/src/algorithms/escalation.rs index 7e3d15b83..4a641485f 100644 --- a/crates/libsy/src/algorithms/escalation.rs +++ b/crates/libsy/src/algorithms/escalation.rs @@ -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::*; @@ -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 @@ -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"}"#, ))) @@ -675,7 +681,9 @@ 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() @@ -683,9 +691,53 @@ mod tests { 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(()) } diff --git a/crates/libsy/src/algorithms/util/classifier_contract.rs b/crates/libsy/src/algorithms/util/classifier_contract.rs index 6cdc42551..f5cca6e84 100644 --- a/crates/libsy/src/algorithms/util/classifier_contract.rs +++ b/crates/libsy/src/algorithms/util/classifier_contract.rs @@ -28,6 +28,8 @@ pub struct ClassifierContractConfig { #[serde(default)] prompt: Option, #[serde(default)] + prompt_suffix: Option, + #[serde(default)] response_format_type: ClassifierResponseFormat, } @@ -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) -> 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, @@ -77,7 +90,23 @@ impl ClassifierContract { default_prompt: &str, response_format_json: &str, ) -> Result { - 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 { @@ -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( diff --git a/crates/switchyard-runner/src/algorithm.rs b/crates/switchyard-runner/src/algorithm.rs index 5cc838209..6635d73e6 100644 --- a/crates/switchyard-runner/src/algorithm.rs +++ b/crates/switchyard-runner/src/algorithm.rs @@ -108,6 +108,7 @@ struct CapabilityClassifierRouteConfig { message_hash_fallback: bool, recent_turn_window: Option, prompt: Option, + prompt_suffix: Option, response_format_type: ClassifierResponseFormat, max_output_tokens: u64, } @@ -118,6 +119,7 @@ struct EscalationClassifierRouteConfig { strong_target: String, weak_target: String, prompt: Option, + prompt_suffix: Option, response_format_type: ClassifierResponseFormat, max_output_tokens: u64, judge: EscalationJudgeConfig, @@ -251,6 +253,8 @@ pub struct LlmClassifierRouteConfig { pub recent_turn_window: Option, /// Replaces the packaged judge prompt. Required in custom mode. pub prompt: Option, + /// Appended to the packaged or overridden judge prompt. Not allowed in custom mode. + pub prompt_suffix: Option, /// How the judge is asked for structured output. Use `json_object` when the /// provider cannot do JSON Schema. pub response_format_type: ClassifierResponseFormat, @@ -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, }), @@ -890,6 +894,7 @@ impl LlmClassifierRouteConfig { message_hash_fallback, recent_turn_window, prompt, + prompt_suffix, response_format_type, max_output_tokens, escalation, @@ -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, }, @@ -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)?, @@ -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!( @@ -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, @@ -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, @@ -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 { diff --git a/crates/switchyard-runner/src/config.rs b/crates/switchyard-runner/src/config.rs index 1c976a47b..c8ef2c744 100644 --- a/crates/switchyard-runner/src/config.rs +++ b/crates/switchyard-runner/src/config.rs @@ -1257,13 +1257,13 @@ new = ["send_message"] fn classifier_prompts_are_configurable_in_both_modes() -> RunnerResult<()> { let capability = VALID_CONFIG.replace( "base_threshold = 0.5", - "base_threshold = 0.5\nprompt = \"custom capability rubric\"", + "base_threshold = 0.5\nprompt = \"custom capability rubric\"\nprompt_suffix = \"tool-step rubric\"", ); runner_from_toml(&capability)?; let escalation = VALID_CONFIG.replace( "base_threshold = 0.5", - "base_threshold = 0.5\nprompt = \"custom trajectory rubric\"\nescalation = { confirmations = 2 }", + "base_threshold = 0.5\nprompt = \"custom trajectory rubric\"\nprompt_suffix = \"tool-step rubric\"\nescalation = { confirmations = 2 }", ); runner_from_toml(&escalation)?; @@ -1284,6 +1284,38 @@ new = ["send_message"] Ok(()) } + #[test] + fn prompt_suffix_reaches_both_classifier_modes() { + let cases = [ + ("capability", "prompt_suffix = \"{{RESPONSE_SCHEMA}}\""), + ( + "escalation", + "prompt_suffix = \"{{RESPONSE_SCHEMA}}\"\nescalation = { confirmations = 2 }", + ), + ]; + for (name, fields) in cases { + let config = VALID_CONFIG.replace( + "base_threshold = 0.5", + &format!("base_threshold = 0.5\n{fields}"), + ); + assert!( + error_message(&config).contains("Switchyard supplies the schema automatically"), + "{name}" + ); + } + } + + #[test] + fn mode_custom_rejects_prompt_suffix() { + let mixed = + with_subagent_llm_classifier(VALID_CONFIG, "passthrough", "\nprompt_suffix = \"x\""); + + assert!( + error_message(&mixed) + .contains("mode custom cannot use capability or escalation fields") + ); + } + #[test] fn mode_custom_rejects_capability_fields() { let mixed = VALID_CONFIG.replace( diff --git a/docs/reference/toml_schema.md b/docs/reference/toml_schema.md index c77c6c89e..ac9344426 100644 --- a/docs/reference/toml_schema.md +++ b/docs/reference/toml_schema.md @@ -262,6 +262,7 @@ Capability mode classifies before serving. See | `message_hash_fallback` | No | `false` | Retains the target against a hash of the first user message when a request carries no session ID. Requires `classify_trigger = "new_session"` or `"user_turn"`. | | `recent_turn_window` | No | unset | When unset, the judge sees the opening task and latest user follow-up, when present. When set, it also sees trailing turns. | | `prompt` | No | packaged prompt | Replaces the capability prompt. The packaged schema is sent separately as structured-output configuration. | +| `prompt_suffix` | No | unset | Appended to the packaged or overridden prompt after a blank line. Blank values are ignored. Not allowed in custom mode. | Escalation mode serves the weak target first and judges the completed turn. See [Escalation-Router Routing](../routing_algorithms/escalation_router_routing.md). @@ -271,6 +272,7 @@ Escalation mode serves the weak target first and judges the completed turn. See | `strong_target` | Yes | — | Target used after the session latches. | | `weak_target` | Yes | — | Target served before the latch. | | `prompt` | No | packaged prompt | Replaces the trajectory-judge prompt. | +| `prompt_suffix` | No | unset | Appended to the packaged or overridden prompt after a blank line. Blank values are ignored. Not allowed in custom mode. | | `escalation.confirmations` | No | `2` | Consecutive fresh-evidence verdicts for the same failure category required to latch. Above `1` needs a stable session ID. | | `escalation.recent_turn_window` | No | `28` | Trailing messages shown to the judge. | | `escalation.window_message_chars` | No | `500` | Per-message cap inside that window. | diff --git a/docs/routing_algorithms/escalation_router_routing.md b/docs/routing_algorithms/escalation_router_routing.md index 994359a08..7e5d53cfe 100644 --- a/docs/routing_algorithms/escalation_router_routing.md +++ b/docs/routing_algorithms/escalation_router_routing.md @@ -57,8 +57,9 @@ Choose the route id deliberately for the surface you want the efficient tier to work on, and keep it stable across runs you intend to compare, because Switchyard cannot change the client's choice from the server side. -The route-level `prompt` key replaces the packaged trajectory-judge prompt. It -uses the escalation verdict schema rather than the capability verdict schema. +The route-level `prompt` key replaces the packaged trajectory-judge prompt; +`prompt_suffix` appends guidance to it without replacing it. The judge prompt uses +the escalation verdict schema rather than the capability verdict schema. Switchyard supplies that schema according to the route's `response_format_type`: through the structured-output request in the default `json_schema` mode, or in the prompt in `json_object` mode. The verdict includes an `escalate` decision, a diff --git a/docs/routing_algorithms/llm_classifier_routing.md b/docs/routing_algorithms/llm_classifier_routing.md index de35015ff..d57087a5b 100644 --- a/docs/routing_algorithms/llm_classifier_routing.md +++ b/docs/routing_algorithms/llm_classifier_routing.md @@ -123,6 +123,7 @@ for the server merge behavior. | `classify_trigger` | `every_request` | When the judge runs. `every_request` judges every request, tool continuations included. `user_turn` judges each new user message and holds that target across the tool calls between. `new_session` judges once and reuses that target for the session. | | `message_hash_fallback` | `false` | When session metadata is absent, keys affinity from the first user-message text. Requires `classify_trigger = "new_session"` or `"user_turn"`. | | `prompt` | packaged capability prompt | Replaces the classifier's system prompt. The packaged verdict schema and routing policy remain active. | +| `prompt_suffix` | unset | Appends guidance to the packaged or overridden prompt after a blank line. Blank values are ignored. Not allowed in custom mode. | | `response_format_type` | `json_schema` | Structured-output mode for capability and escalation judges. Use `json_object` for providers without JSON Schema support. | | `max_output_tokens` | `4096` | Maximum completion tokens available to the classifier verdict. Must be at least `1`. |