1use serde::Deserialize;
11use serde_json::Value;
12use switchyard_protocol::{ContentBlock, Message, ModelId, Role};
13
14use super::classifier_contract::{ClassifierContract, ClassifierContractConfig};
15use super::llm_judge::{
16 ClassifierInput, JudgeClassifier, JudgePolicy, JudgeRuntimeConfig, SerdeDecoder,
17 StructuredJudge,
18};
19use crate::core::classifier::{Classification, Score};
20use crate::core::state::State;
21use crate::{LibsyError, Result};
22use switchyard_protocol::Request;
23
24const PROMPT_TEMPLATE: &str = include_str!("../../prompts/escalation/prompt.md");
25const SCHEMA_TEMPLATE: &str = include_str!("../../prompts/escalation/schema.json");
26
27const TRIM_MARKER: &str = " ...[trimmed] ";
29
30const TRUNCATION_SUFFIX: &str = "...<truncated>";
32
33const SYSTEM_CHARS: usize = 1_000;
36
37const TASK_CHARS: usize = 4_000;
43
44const MAX_REQUEST_CHARS: usize = 18_000;
46
47#[derive(Clone, Debug, Deserialize)]
52#[serde(default, deny_unknown_fields)]
53pub struct EscalationJudgeConfig {
54 pub confirmations: u32,
59 pub recent_turn_window: usize,
61 pub window_message_chars: usize,
63 pub gate: Option<EscalationGateConfig>,
68}
69
70#[derive(Clone, Debug, Deserialize, PartialEq)]
77#[serde(deny_unknown_fields)]
78pub struct EscalationGateConfig {
79 pub base_threshold: f64,
81 #[serde(default)]
84 pub threshold_step: f64,
85 #[serde(default)]
87 pub prompt: Option<String>,
88 #[serde(default)]
94 pub classifier_target: Option<String>,
95}
96
97impl EscalationGateConfig {
98 fn validate(&self) -> Result<()> {
99 let reject = |message: String| Err(LibsyError::AlgorithmError { message });
100 if !(0.0..=1.0).contains(&self.base_threshold) {
101 return reject(format!(
102 "gate.base_threshold must be between 0 and 1, got {}",
103 self.base_threshold
104 ));
105 }
106 if !self.threshold_step.is_finite() || self.threshold_step < 0.0 {
107 return reject(format!(
108 "gate.threshold_step must be finite and non-negative, got {}",
109 self.threshold_step
110 ));
111 }
112 let unsupported_threshold = self.base_threshold + 2.0 * self.threshold_step;
113 if unsupported_threshold > 1.0 {
114 return reject(format!(
115 "gate.base_threshold + 2 * gate.threshold_step must be at most 1, got {unsupported_threshold}"
116 ));
117 }
118 Ok(())
119 }
120}
121
122impl EscalationJudgeConfig {
123 fn validate(&self) -> Result<()> {
125 let reject = |message: String| Err(LibsyError::AlgorithmError { message });
126 if self.confirmations == 0 {
127 return reject("confirmations must be at least 1".to_string());
128 }
129 if let Some(gate) = &self.gate {
130 gate.validate()?;
131 }
132 if self.recent_turn_window == 0 {
133 return reject("recent_turn_window must be at least 1".to_string());
134 }
135 if self.window_message_chars < 50 {
136 return reject(format!(
137 "window_message_chars must be at least 50, got {}",
138 self.window_message_chars
139 ));
140 }
141 Ok(())
142 }
143}
144
145impl Default for EscalationJudgeConfig {
146 fn default() -> Self {
147 Self {
148 confirmations: 2,
149 recent_turn_window: 28,
150 window_message_chars: 500,
151 gate: None,
152 }
153 }
154}
155
156#[derive(Deserialize)]
161pub(crate) struct EscalationVerdict {
162 escalate: bool,
163 #[serde(default)]
164 reason: String,
165}
166
167pub(crate) struct EscalationInput {
169 config: EscalationJudgeConfig,
170}
171
172impl ClassifierInput for EscalationInput {
173 fn build_messages(&self, _state: &State, request: &Request) -> Vec<Message> {
174 let messages = &request.llm_request.messages;
175 let summary = summarize_for_judge(messages, conversation_turn(request), &self.config);
176 vec![Message::text(Role::User, summary)]
177 }
178}
179
180pub(crate) type EscalationJudge = StructuredJudge<EscalationInput, SerdeDecoder<EscalationVerdict>>;
182
183pub(crate) struct EscalationPolicy {
189 capable: ModelId,
190 efficient: ModelId,
191}
192
193impl JudgePolicy for EscalationPolicy {
194 type Verdict = EscalationVerdict;
195
196 fn to_classification(&self, verdict: Option<&EscalationVerdict>) -> Classification {
197 if let Some(verdict) = verdict {
198 tracing::debug!(
199 escalate = verdict.escalate,
200 reason = %verdict.reason,
201 "escalation judge verdict"
202 );
203 }
204 match verdict {
205 Some(verdict) if verdict.escalate => Classification::Scores(vec![Score {
206 target: self.capable.clone(),
207 confidence: 1.0,
208 }]),
209 Some(_) => Classification::Scores(vec![Score {
210 target: self.efficient.clone(),
211 confidence: 1.0,
212 }]),
213 None => Classification::Ambiguous(Vec::new()),
214 }
215 }
216}
217
218fn escalation_evidence(
220 _policy: &EscalationPolicy,
221 verdict: Option<&EscalationVerdict>,
222) -> Option<Value> {
223 verdict.map(|verdict| {
224 serde_json::json!({
225 "source": "escalation",
226 "verdict": if verdict.escalate { "escalate" } else { "continue" },
227 })
228 })
229}
230
231pub(crate) fn build_judge(
236 judge_target: ModelId,
237 capable: ModelId,
238 efficient: ModelId,
239 contract_config: &ClassifierContractConfig,
240 config: EscalationJudgeConfig,
241 max_output_tokens: u64,
242) -> Result<JudgeClassifier<EscalationJudge, EscalationPolicy>> {
243 config.validate()?;
244 let contract =
245 ClassifierContract::from_config(contract_config, PROMPT_TEMPLATE, SCHEMA_TEMPLATE)?;
246 Ok(JudgeClassifier::new(
247 StructuredJudge::new(
248 EscalationInput { config },
249 contract,
250 SerdeDecoder::new(),
251 JudgeRuntimeConfig::new(max_output_tokens)?,
252 ),
253 judge_target,
254 EscalationPolicy { capable, efficient },
255 )
256 .with_evidence(escalation_evidence))
257}
258
259pub(crate) fn conversation_turn(request: &Request) -> usize {
268 request
269 .llm_request
270 .messages
271 .iter()
272 .filter(|message| message.role == Role::Assistant)
273 .count()
274}
275
276fn message_text(message: &Message) -> String {
282 let mut parts = Vec::new();
283 collect_text(&message.content, &mut parts);
284 parts.join(" ")
285}
286
287fn collect_text(content: &[ContentBlock], parts: &mut Vec<String>) {
289 for block in content {
290 match block {
291 ContentBlock::Text { text } | ContentBlock::Refusal { text } => {
292 parts.push(text.clone());
293 }
294 ContentBlock::ToolCall(call) => {
295 parts.push(format!("tool_call {}({})", call.name, call.arguments));
296 }
297 ContentBlock::ToolResult(result) => collect_text(&result.content, parts),
298 _ => {}
299 }
300 }
301}
302
303fn truncate_middle(text: &str, limit: usize) -> String {
308 let chars: Vec<char> = text.chars().collect();
309 if chars.len() <= limit {
310 return text.to_string();
311 }
312 let keep = limit
313 .saturating_sub(TRIM_MARKER.chars().count())
314 .max(20)
315 .min(chars.len());
316 let head = keep * 2 / 3;
317 let tail = keep - head;
318 let mut out: String = chars[..head].iter().collect();
319 out.push_str(TRIM_MARKER);
320 out.extend(chars[chars.len() - tail..].iter());
321 out
322}
323
324fn summarize_for_judge(
334 messages: &[Message],
335 turn: usize,
336 config: &EscalationJudgeConfig,
337) -> String {
338 let mut anchors: Vec<String> = Vec::new();
339 let mut window: Vec<String> = Vec::new();
340 let mut assistant_seen = false;
341
342 for message in messages {
343 let text = message_text(message);
344 match message.role {
345 Role::System | Role::Developer => anchors.push(format!(
346 "[{}] {}",
347 role_label(message.role),
348 truncate_middle(&text, SYSTEM_CHARS)
349 )),
350 Role::User if !assistant_seen => {
352 anchors.push(format!(
353 "[user (task)] {}",
354 truncate_middle(&text, TASK_CHARS)
355 ));
356 }
357 role => {
358 if role == Role::Assistant {
359 assistant_seen = true;
360 }
361 window.push(format!(
362 "[{}] {}",
363 role_label(role),
364 truncate_middle(&text, config.window_message_chars)
365 ));
366 }
367 }
368 }
369
370 if window.len() > config.recent_turn_window {
371 window.drain(..window.len() - config.recent_turn_window);
372 }
373
374 let assemble = |window: &[String]| {
375 let header = format!(
376 "Conversation turn {turn}; showing the last {} of {} messages after the task framing.",
377 window.len(),
378 messages.len(),
379 );
380 std::iter::once(header)
381 .chain(anchors.iter().cloned())
382 .chain(window.iter().cloned())
383 .collect::<Vec<_>>()
384 .join("\n")
385 };
386
387 let mut text = assemble(&window);
388 while text.chars().count() > MAX_REQUEST_CHARS && !window.is_empty() {
389 window.remove(0);
390 text = assemble(&window);
391 }
392 if text.chars().count() > MAX_REQUEST_CHARS {
393 let keep = MAX_REQUEST_CHARS.saturating_sub(TRUNCATION_SUFFIX.chars().count() + 1);
394 text = text.chars().take(keep).collect::<String>() + TRUNCATION_SUFFIX;
395 }
396 text
397}
398
399fn role_label(role: Role) -> &'static str {
401 match role {
402 Role::System => "system",
403 Role::Developer => "developer",
404 Role::User => "user",
405 Role::Assistant => "assistant",
406 Role::Tool => "tool",
407 }
408}
409
410#[cfg(test)]
415pub(crate) fn request_at_turn(session_id: Option<&str>, turn: usize) -> Request {
416 use switchyard_protocol::{LlmRequest, Metadata};
417
418 let mut messages = vec![Message::text(Role::User, "What is 2+2?")];
419 for attempt in 1..turn {
420 messages.push(Message::text(Role::Assistant, format!("attempt {attempt}")));
421 messages.push(Message::text(Role::User, format!("still wrong {attempt}")));
422 }
423 Request {
424 llm_request: LlmRequest {
425 model: Some("auto".to_string()),
426 messages,
427 ..LlmRequest::default()
428 },
429 raw_request: None,
430 metadata: session_id.map(|id| Metadata {
431 session_id: Some(id.to_string()),
432 ..Metadata::default()
433 }),
434 }
435}
436
437#[cfg(test)]
438mod tests {
439 use serde_json::json;
440 use switchyard_protocol::{ContentBlock, Message, Role, ToolCall, ToolResult};
441
442 use super::*;
443 use crate::algorithms::util::llm_judge::Judge;
444
445 fn escalation_judge(max_output_tokens: u64) -> Result<EscalationJudge> {
446 Ok(StructuredJudge::new(
447 EscalationInput {
448 config: EscalationJudgeConfig::default(),
449 },
450 ClassifierContract::from_config(
451 &ClassifierContractConfig::default(),
452 PROMPT_TEMPLATE,
453 SCHEMA_TEMPLATE,
454 )?,
455 SerdeDecoder::new(),
456 JudgeRuntimeConfig::new(max_output_tokens)?,
457 ))
458 }
459
460 #[test]
461 fn judge_request_is_rubric_plus_summary_under_a_completion_cap() -> Result<()> {
462 let judge = escalation_judge(super::super::DEFAULT_JUDGE_MAX_OUTPUT_TOKENS)?;
463
464 let mut judged = request_at_turn(None, 4);
466 judged
467 .llm_request
468 .messages
469 .push(Message::text(Role::Assistant, "this turn's reply"));
470 let built = judge.build_request(&State::default(), &judged);
471
472 assert_eq!(built.llm_request.instructions.len(), 1);
474 assert_eq!(built.llm_request.instructions[0].role, Role::System);
475 assert_eq!(built.llm_request.messages.len(), 1);
476 assert_eq!(built.llm_request.messages[0].role, Role::User);
477 assert!(
478 built.llm_request.messages[0]
479 .text_content("")
480 .is_some_and(|text| text.contains("Conversation turn 4"))
481 );
482 assert_eq!(
484 built.llm_request.output.max_output_tokens,
485 Some(super::super::DEFAULT_JUDGE_MAX_OUTPUT_TOKENS)
486 );
487 assert!(built.llm_request.output.response_format.is_some());
488 Ok(())
489 }
490
491 #[test]
492 fn judge_request_uses_the_configured_completion_cap() -> Result<()> {
493 let judge = escalation_judge(512)?;
494
495 let built = judge.build_request(&State::default(), &request_at_turn(None, 1));
496
497 assert_eq!(built.llm_request.output.max_output_tokens, Some(512));
498 Ok(())
499 }
500
501 #[test]
502 fn conversation_turn_counts_assistant_replies() {
503 for turn in [1, 5] {
506 let mut judged = request_at_turn(None, turn);
507 judged
508 .llm_request
509 .messages
510 .push(Message::text(Role::Assistant, "this turn's reply"));
511 assert_eq!(conversation_turn(&judged), turn);
512 }
513 }
514
515 #[test]
516 fn message_text_keeps_tool_calls_and_results() {
517 let call = Message {
518 role: Role::Assistant,
519 content: vec![
520 ContentBlock::Text {
521 text: "running it".to_string(),
522 },
523 ContentBlock::ToolCall(ToolCall {
524 id: "call-1".to_string(),
525 name: "bash".to_string(),
526 arguments: json!({"cmd": "ls"}),
527 }),
528 ],
529 };
530 let text = message_text(&call);
531 assert!(text.contains("running it"), "{text}");
532 assert!(text.contains(r#"tool_call bash({"cmd":"ls"})"#), "{text}");
533
534 let result = Message {
535 role: Role::Tool,
536 content: vec![ContentBlock::ToolResult(ToolResult {
537 tool_call_id: "call-1".to_string(),
538 content: vec![ContentBlock::Text {
539 text: "no such file".to_string(),
540 }],
541 is_error: Some(true),
542 })],
543 };
544 assert_eq!(message_text(&result), "no such file");
545 }
546
547 #[test]
548 fn truncate_middle_keeps_head_and_tail() {
549 let text = "a".repeat(40) + &"z".repeat(40);
550 let trimmed = truncate_middle(&text, 50);
551 assert!(trimmed.chars().count() <= 50, "{trimmed}");
552 assert!(trimmed.starts_with('a'));
553 assert!(trimmed.ends_with('z'));
554 assert!(trimmed.contains("[trimmed]"));
555
556 assert_eq!(truncate_middle("short", 50), "short");
558 }
559
560 #[test]
561 fn summary_keeps_anchors_and_the_recent_window() {
562 let mut messages = vec![
563 Message::text(Role::System, "you are a coding agent"),
564 Message::text(Role::User, "fix the failing test"),
565 ];
566 for i in 0..10 {
567 messages.push(Message::text(Role::Assistant, format!("step {i}")));
568 }
569 let config = EscalationJudgeConfig {
570 recent_turn_window: 3,
571 ..EscalationJudgeConfig::default()
572 };
573
574 let summary = summarize_for_judge(&messages, 11, &config);
575
576 assert!(
577 summary.contains("[system] you are a coding agent"),
578 "{summary}"
579 );
580 assert!(
581 summary.contains("[user (task)] fix the failing test"),
582 "{summary}"
583 );
584 assert!(summary.contains("Conversation turn 11; showing the last 3 of 12 messages"));
585 assert!(summary.contains("step 9"), "{summary}");
587 assert!(summary.contains("step 7"), "{summary}");
588 assert!(!summary.contains("step 6"), "{summary}");
589 }
590
591 #[test]
592 fn summary_anchors_every_user_message_before_the_first_reply() {
593 let mut messages = vec![
596 Message::text(
597 Role::Developer,
598 "<skills_instructions>...</skills_instructions>",
599 ),
600 Message::text(
601 Role::User,
602 "<environment_context><cwd>/app</cwd></environment_context>",
603 ),
604 Message::text(Role::User, "Implement RFC 5545 timezone interop in rrule."),
605 ];
606 for i in 0..40 {
607 messages.push(Message::text(Role::Assistant, format!("step {i}")));
608 messages.push(Message::text(Role::User, format!("later user note {i}")));
609 }
610 let config = EscalationJudgeConfig {
611 recent_turn_window: 3,
612 ..EscalationJudgeConfig::default()
613 };
614
615 let summary = summarize_for_judge(&messages, 40, &config);
616
617 assert!(
618 summary.contains("[user (task)] <environment_context>"),
619 "{summary}"
620 );
621 assert!(
622 summary.contains("[user (task)] Implement RFC 5545 timezone interop in rrule."),
623 "{summary}"
624 );
625 assert!(
627 !summary.contains("[user (task)] later user note"),
628 "{summary}"
629 );
630 assert!(summary.contains("[user] later user note 39"), "{summary}");
631 assert!(!summary.contains("later user note 0\n"), "{summary}");
632 }
633
634 #[test]
635 fn summary_drops_oldest_window_lines_under_the_char_cap() {
636 let mut messages = vec![
640 Message::text(Role::System, "framing"),
641 Message::text(Role::User, "task"),
642 ];
643 for i in 0..20 {
644 messages.push(Message::text(
645 Role::Assistant,
646 format!("{i} {}", "x".repeat(2_000)),
647 ));
648 }
649 let config = EscalationJudgeConfig {
650 window_message_chars: 2_000,
651 ..EscalationJudgeConfig::default()
652 };
653
654 let summary = summarize_for_judge(&messages, 21, &config);
655
656 assert!(
657 summary.chars().count() <= MAX_REQUEST_CHARS,
658 "{}",
659 summary.chars().count()
660 );
661 assert!(summary.contains("[system] framing"), "{summary}");
663 assert!(summary.contains("[user (task)] task"), "{summary}");
664 assert!(summary.contains("19 xxx"), "{summary}");
665 assert!(!summary.contains("0 xxx"), "{summary}");
666 }
667}