1use serde::Deserialize;
11use serde_json::Value;
12use switchyard_protocol::{Category, ContentBlock, Message, Role};
13
14use super::classifier_contract::{ClassifierContract, ClassifierContractConfig};
15use super::llm_judge::{
16 ClassifierInput, JudgeClassifier, JudgePolicy, JudgeRuntimeConfig, SerdeDecoder,
17 StructuredJudge,
18};
19use crate::core::algorithm::Driver;
20use crate::core::classifier::{Classification, Score};
21use crate::core::state::State;
22use crate::{LibsyError, Result};
23use switchyard_protocol::Request;
24
25const PROMPT_TEMPLATE: &str = include_str!("../../prompts/escalation/prompt.md");
26const SCHEMA_TEMPLATE: &str = include_str!("../../prompts/escalation/schema.json");
27
28const TRIM_MARKER: &str = " ...[trimmed] ";
30
31const TRUNCATION_SUFFIX: &str = "...<truncated>";
33
34const SYSTEM_CHARS: usize = 1_000;
37
38const TASK_CHARS: usize = 4_000;
44
45const MAX_REQUEST_CHARS: usize = 18_000;
47
48#[derive(Clone, Debug, Deserialize)]
53#[serde(default, deny_unknown_fields)]
54pub struct EscalationJudgeConfig {
55 pub confirmations: u32,
60 pub recent_turn_window: usize,
62 pub window_message_chars: usize,
64}
65
66impl EscalationJudgeConfig {
67 fn validate(&self) -> Result<()> {
69 let reject = |message: String| Err(LibsyError::AlgorithmError { message });
70 if self.confirmations == 0 {
71 return reject("confirmations must be at least 1".to_string());
72 }
73 if self.recent_turn_window == 0 {
74 return reject("recent_turn_window must be at least 1".to_string());
75 }
76 if self.window_message_chars < 50 {
77 return reject(format!(
78 "window_message_chars must be at least 50, got {}",
79 self.window_message_chars
80 ));
81 }
82 Ok(())
83 }
84}
85
86impl Default for EscalationJudgeConfig {
87 fn default() -> Self {
88 Self {
89 confirmations: 2,
90 recent_turn_window: 28,
91 window_message_chars: 500,
92 }
93 }
94}
95
96#[derive(Deserialize)]
101pub(crate) struct EscalationVerdict {
102 escalate: bool,
103 #[serde(default)]
104 reason: String,
105}
106
107pub(crate) struct EscalationInput {
109 config: EscalationJudgeConfig,
110}
111
112impl ClassifierInput for EscalationInput {
113 fn build_messages(&self, _state: &State, request: &Request) -> Vec<Message> {
114 let messages = &request.llm_request.messages;
115 let summary = summarize_for_judge(messages, conversation_turn(request), &self.config);
116 vec![Message::text(Role::User, summary)]
117 }
118}
119
120pub(crate) type EscalationJudge = StructuredJudge<EscalationInput, SerdeDecoder<EscalationVerdict>>;
122
123pub(crate) struct EscalationPolicy;
129
130impl JudgePolicy for EscalationPolicy {
131 type Verdict = EscalationVerdict;
132
133 fn to_classification(
134 &self,
135 verdict: Option<&EscalationVerdict>,
136 driver: &Driver,
137 ) -> Result<Classification> {
138 if let Some(verdict) = verdict {
139 tracing::debug!(
140 escalate = verdict.escalate,
141 reason = %verdict.reason,
142 "escalation judge verdict"
143 );
144 }
145 match verdict {
146 Some(verdict) if verdict.escalate => Ok(Classification::Scores(vec![Score {
147 target: driver.first_model_for(&Category::Capable)?.clone(),
148 confidence: 1.0,
149 category: Some(Category::Capable),
150 }])),
151 Some(_) => Ok(Classification::Scores(vec![Score {
152 target: driver.first_model_for(&Category::Efficient)?.clone(),
153 confidence: 1.0,
154 category: Some(Category::Efficient),
155 }])),
156 None => Ok(Classification::Ambiguous(Vec::new())),
157 }
158 }
159}
160
161fn escalation_evidence(
163 _policy: &EscalationPolicy,
164 verdict: Option<&EscalationVerdict>,
165) -> Option<Value> {
166 verdict.map(|verdict| {
167 serde_json::json!({
168 "source": "escalation",
169 "verdict": if verdict.escalate { "escalate" } else { "continue" },
170 })
171 })
172}
173
174pub(crate) fn build_judge(
179 contract_config: &ClassifierContractConfig,
180 config: EscalationJudgeConfig,
181 max_output_tokens: u64,
182) -> Result<JudgeClassifier<EscalationJudge, EscalationPolicy>> {
183 config.validate()?;
184 let contract =
185 ClassifierContract::from_config(contract_config, PROMPT_TEMPLATE, SCHEMA_TEMPLATE)?;
186 Ok(JudgeClassifier::new(
187 StructuredJudge::new(
188 EscalationInput { config },
189 contract,
190 SerdeDecoder::new(),
191 JudgeRuntimeConfig::new(max_output_tokens)?,
192 ),
193 EscalationPolicy,
194 )
195 .with_evidence(escalation_evidence))
196}
197
198pub(crate) fn conversation_turn(request: &Request) -> usize {
207 request
208 .llm_request
209 .messages
210 .iter()
211 .filter(|message| message.role == Role::Assistant)
212 .count()
213}
214
215fn message_text(message: &Message) -> String {
221 let mut parts = Vec::new();
222 collect_text(&message.content, &mut parts);
223 parts.join(" ")
224}
225
226fn collect_text(content: &[ContentBlock], parts: &mut Vec<String>) {
228 for block in content {
229 match block {
230 ContentBlock::Text { text } | ContentBlock::Refusal { text } => {
231 parts.push(text.clone());
232 }
233 ContentBlock::ToolCall(call) => {
234 parts.push(format!("tool_call {}({})", call.name, call.arguments));
235 }
236 ContentBlock::ToolResult(result) => collect_text(&result.content, parts),
237 _ => {}
238 }
239 }
240}
241
242fn truncate_middle(text: &str, limit: usize) -> String {
247 let chars: Vec<char> = text.chars().collect();
248 if chars.len() <= limit {
249 return text.to_string();
250 }
251 let keep = limit
252 .saturating_sub(TRIM_MARKER.chars().count())
253 .max(20)
254 .min(chars.len());
255 let head = keep * 2 / 3;
256 let tail = keep - head;
257 let mut out: String = chars[..head].iter().collect();
258 out.push_str(TRIM_MARKER);
259 out.extend(chars[chars.len() - tail..].iter());
260 out
261}
262
263fn summarize_for_judge(
273 messages: &[Message],
274 turn: usize,
275 config: &EscalationJudgeConfig,
276) -> String {
277 let mut anchors: Vec<String> = Vec::new();
278 let mut window: Vec<String> = Vec::new();
279 let mut assistant_seen = false;
280
281 for message in messages {
282 let text = message_text(message);
283 match message.role {
284 Role::System | Role::Developer => anchors.push(format!(
285 "[{}] {}",
286 role_label(message.role),
287 truncate_middle(&text, SYSTEM_CHARS)
288 )),
289 Role::User if !assistant_seen => {
291 anchors.push(format!(
292 "[user (task)] {}",
293 truncate_middle(&text, TASK_CHARS)
294 ));
295 }
296 role => {
297 if role == Role::Assistant {
298 assistant_seen = true;
299 }
300 window.push(format!(
301 "[{}] {}",
302 role_label(role),
303 truncate_middle(&text, config.window_message_chars)
304 ));
305 }
306 }
307 }
308
309 if window.len() > config.recent_turn_window {
310 window.drain(..window.len() - config.recent_turn_window);
311 }
312
313 let assemble = |window: &[String]| {
314 let header = format!(
315 "Conversation turn {turn}; showing the last {} of {} messages after the task framing.",
316 window.len(),
317 messages.len(),
318 );
319 std::iter::once(header)
320 .chain(anchors.iter().cloned())
321 .chain(window.iter().cloned())
322 .collect::<Vec<_>>()
323 .join("\n")
324 };
325
326 let mut text = assemble(&window);
327 while text.chars().count() > MAX_REQUEST_CHARS && !window.is_empty() {
328 window.remove(0);
329 text = assemble(&window);
330 }
331 if text.chars().count() > MAX_REQUEST_CHARS {
332 let keep = MAX_REQUEST_CHARS.saturating_sub(TRUNCATION_SUFFIX.chars().count() + 1);
333 text = text.chars().take(keep).collect::<String>() + TRUNCATION_SUFFIX;
334 }
335 text
336}
337
338fn role_label(role: Role) -> &'static str {
340 match role {
341 Role::System => "system",
342 Role::Developer => "developer",
343 Role::User => "user",
344 Role::Assistant => "assistant",
345 Role::Tool => "tool",
346 }
347}
348
349#[cfg(test)]
354pub(crate) fn request_at_turn(session_id: Option<&str>, turn: usize) -> Request {
355 use switchyard_protocol::{LlmRequest, Metadata};
356
357 let mut messages = vec![Message::text(Role::User, "What is 2+2?")];
358 for attempt in 1..turn {
359 messages.push(Message::text(Role::Assistant, format!("attempt {attempt}")));
360 messages.push(Message::text(Role::User, format!("still wrong {attempt}")));
361 }
362 Request {
363 llm_request: LlmRequest {
364 model: Some("auto".to_string()),
365 messages,
366 ..LlmRequest::default()
367 },
368 raw_request: None,
369 metadata: session_id.map(|id| Metadata {
370 session_id: Some(id.to_string()),
371 ..Metadata::default()
372 }),
373 }
374}
375
376#[cfg(test)]
377mod tests {
378 use serde_json::json;
379 use switchyard_protocol::{ContentBlock, Message, Role, ToolCall, ToolResult};
380
381 use super::*;
382 use crate::algorithms::util::llm_judge::Judge;
383
384 fn escalation_judge(max_output_tokens: u64) -> Result<EscalationJudge> {
385 Ok(StructuredJudge::new(
386 EscalationInput {
387 config: EscalationJudgeConfig::default(),
388 },
389 ClassifierContract::from_config(
390 &ClassifierContractConfig::default(),
391 PROMPT_TEMPLATE,
392 SCHEMA_TEMPLATE,
393 )?,
394 SerdeDecoder::new(),
395 JudgeRuntimeConfig::new(max_output_tokens)?,
396 ))
397 }
398
399 #[test]
400 fn judge_request_is_rubric_plus_summary_under_a_completion_cap() -> Result<()> {
401 let judge = escalation_judge(super::super::DEFAULT_JUDGE_MAX_OUTPUT_TOKENS)?;
402
403 let mut judged = request_at_turn(None, 4);
405 judged
406 .llm_request
407 .messages
408 .push(Message::text(Role::Assistant, "this turn's reply"));
409 let built = judge.build_request(&State::default(), &judged);
410
411 assert_eq!(built.llm_request.instructions.len(), 1);
413 assert_eq!(built.llm_request.instructions[0].role, Role::System);
414 assert_eq!(built.llm_request.messages.len(), 1);
415 assert_eq!(built.llm_request.messages[0].role, Role::User);
416 assert!(
417 built.llm_request.messages[0]
418 .text_content("")
419 .is_some_and(|text| text.contains("Conversation turn 4"))
420 );
421 assert_eq!(
423 built.llm_request.output.max_output_tokens,
424 Some(super::super::DEFAULT_JUDGE_MAX_OUTPUT_TOKENS)
425 );
426 assert!(built.llm_request.output.response_format.is_some());
427 Ok(())
428 }
429
430 #[test]
431 fn judge_request_uses_the_configured_completion_cap() -> Result<()> {
432 let judge = escalation_judge(512)?;
433
434 let built = judge.build_request(&State::default(), &request_at_turn(None, 1));
435
436 assert_eq!(built.llm_request.output.max_output_tokens, Some(512));
437 Ok(())
438 }
439
440 #[test]
441 fn conversation_turn_counts_assistant_replies() {
442 for turn in [1, 5] {
445 let mut judged = request_at_turn(None, turn);
446 judged
447 .llm_request
448 .messages
449 .push(Message::text(Role::Assistant, "this turn's reply"));
450 assert_eq!(conversation_turn(&judged), turn);
451 }
452 }
453
454 #[test]
455 fn message_text_keeps_tool_calls_and_results() {
456 let call = Message {
457 role: Role::Assistant,
458 content: vec![
459 ContentBlock::Text {
460 text: "running it".to_string(),
461 },
462 ContentBlock::ToolCall(ToolCall {
463 id: "call-1".to_string(),
464 name: "bash".to_string(),
465 arguments: json!({"cmd": "ls"}),
466 }),
467 ],
468 };
469 let text = message_text(&call);
470 assert!(text.contains("running it"), "{text}");
471 assert!(text.contains(r#"tool_call bash({"cmd":"ls"})"#), "{text}");
472
473 let result = Message {
474 role: Role::Tool,
475 content: vec![ContentBlock::ToolResult(ToolResult {
476 tool_call_id: "call-1".to_string(),
477 content: vec![ContentBlock::Text {
478 text: "no such file".to_string(),
479 }],
480 is_error: Some(true),
481 })],
482 };
483 assert_eq!(message_text(&result), "no such file");
484 }
485
486 #[test]
487 fn truncate_middle_keeps_head_and_tail() {
488 let text = "a".repeat(40) + &"z".repeat(40);
489 let trimmed = truncate_middle(&text, 50);
490 assert!(trimmed.chars().count() <= 50, "{trimmed}");
491 assert!(trimmed.starts_with('a'));
492 assert!(trimmed.ends_with('z'));
493 assert!(trimmed.contains("[trimmed]"));
494
495 assert_eq!(truncate_middle("short", 50), "short");
497 }
498
499 #[test]
500 fn summary_keeps_anchors_and_the_recent_window() {
501 let mut messages = vec![
502 Message::text(Role::System, "you are a coding agent"),
503 Message::text(Role::User, "fix the failing test"),
504 ];
505 for i in 0..10 {
506 messages.push(Message::text(Role::Assistant, format!("step {i}")));
507 }
508 let config = EscalationJudgeConfig {
509 recent_turn_window: 3,
510 ..EscalationJudgeConfig::default()
511 };
512
513 let summary = summarize_for_judge(&messages, 11, &config);
514
515 assert!(
516 summary.contains("[system] you are a coding agent"),
517 "{summary}"
518 );
519 assert!(
520 summary.contains("[user (task)] fix the failing test"),
521 "{summary}"
522 );
523 assert!(summary.contains("Conversation turn 11; showing the last 3 of 12 messages"));
524 assert!(summary.contains("step 9"), "{summary}");
526 assert!(summary.contains("step 7"), "{summary}");
527 assert!(!summary.contains("step 6"), "{summary}");
528 }
529
530 #[test]
531 fn summary_anchors_every_user_message_before_the_first_reply() {
532 let mut messages = vec![
535 Message::text(
536 Role::Developer,
537 "<skills_instructions>...</skills_instructions>",
538 ),
539 Message::text(
540 Role::User,
541 "<environment_context><cwd>/app</cwd></environment_context>",
542 ),
543 Message::text(Role::User, "Implement RFC 5545 timezone interop in rrule."),
544 ];
545 for i in 0..40 {
546 messages.push(Message::text(Role::Assistant, format!("step {i}")));
547 messages.push(Message::text(Role::User, format!("later user note {i}")));
548 }
549 let config = EscalationJudgeConfig {
550 recent_turn_window: 3,
551 ..EscalationJudgeConfig::default()
552 };
553
554 let summary = summarize_for_judge(&messages, 40, &config);
555
556 assert!(
557 summary.contains("[user (task)] <environment_context>"),
558 "{summary}"
559 );
560 assert!(
561 summary.contains("[user (task)] Implement RFC 5545 timezone interop in rrule."),
562 "{summary}"
563 );
564 assert!(
566 !summary.contains("[user (task)] later user note"),
567 "{summary}"
568 );
569 assert!(summary.contains("[user] later user note 39"), "{summary}");
570 assert!(!summary.contains("later user note 0\n"), "{summary}");
571 }
572
573 #[test]
574 fn summary_drops_oldest_window_lines_under_the_char_cap() {
575 let mut messages = vec![
579 Message::text(Role::System, "framing"),
580 Message::text(Role::User, "task"),
581 ];
582 for i in 0..20 {
583 messages.push(Message::text(
584 Role::Assistant,
585 format!("{i} {}", "x".repeat(2_000)),
586 ));
587 }
588 let config = EscalationJudgeConfig {
589 window_message_chars: 2_000,
590 ..EscalationJudgeConfig::default()
591 };
592
593 let summary = summarize_for_judge(&messages, 21, &config);
594
595 assert!(
596 summary.chars().count() <= MAX_REQUEST_CHARS,
597 "{}",
598 summary.chars().count()
599 );
600 assert!(summary.contains("[system] framing"), "{summary}");
602 assert!(summary.contains("[user (task)] task"), "{summary}");
603 assert!(summary.contains("19 xxx"), "{summary}");
604 assert!(!summary.contains("0 xxx"), "{summary}");
605 }
606}