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}
64
65impl EscalationJudgeConfig {
66 fn validate(&self) -> Result<()> {
68 let reject = |message: String| Err(LibsyError::AlgorithmError { message });
69 if self.confirmations == 0 {
70 return reject("confirmations must be at least 1".to_string());
71 }
72 if self.recent_turn_window == 0 {
73 return reject("recent_turn_window must be at least 1".to_string());
74 }
75 if self.window_message_chars < 50 {
76 return reject(format!(
77 "window_message_chars must be at least 50, got {}",
78 self.window_message_chars
79 ));
80 }
81 Ok(())
82 }
83}
84
85impl Default for EscalationJudgeConfig {
86 fn default() -> Self {
87 Self {
88 confirmations: 2,
89 recent_turn_window: 28,
90 window_message_chars: 500,
91 }
92 }
93}
94
95#[derive(Deserialize)]
100pub(crate) struct EscalationVerdict {
101 escalate: bool,
102 #[serde(default)]
103 reason: String,
104}
105
106pub(crate) struct EscalationInput {
108 config: EscalationJudgeConfig,
109}
110
111impl ClassifierInput for EscalationInput {
112 fn build_messages(&self, _state: &State, request: &Request) -> Vec<Message> {
113 let messages = &request.llm_request.messages;
114 let summary = summarize_for_judge(messages, conversation_turn(request), &self.config);
115 vec![Message::text(Role::User, summary)]
116 }
117}
118
119pub(crate) type EscalationJudge = StructuredJudge<EscalationInput, SerdeDecoder<EscalationVerdict>>;
121
122pub(crate) struct EscalationPolicy {
128 capable: ModelId,
129 efficient: ModelId,
130}
131
132impl JudgePolicy for EscalationPolicy {
133 type Verdict = EscalationVerdict;
134
135 fn to_classification(&self, verdict: Option<&EscalationVerdict>) -> Classification {
136 if let Some(verdict) = verdict {
137 tracing::debug!(
138 escalate = verdict.escalate,
139 reason = %verdict.reason,
140 "escalation judge verdict"
141 );
142 }
143 match verdict {
144 Some(verdict) if verdict.escalate => Classification::Scores(vec![Score {
145 target: self.capable.clone(),
146 confidence: 1.0,
147 }]),
148 Some(_) => Classification::Scores(vec![Score {
149 target: self.efficient.clone(),
150 confidence: 1.0,
151 }]),
152 None => Classification::Ambiguous(Vec::new()),
153 }
154 }
155}
156
157fn escalation_evidence(
159 _policy: &EscalationPolicy,
160 verdict: Option<&EscalationVerdict>,
161) -> Option<Value> {
162 verdict.map(|verdict| {
163 serde_json::json!({
164 "source": "escalation",
165 "verdict": if verdict.escalate { "escalate" } else { "continue" },
166 })
167 })
168}
169
170pub(crate) fn build_judge(
175 judge_target: ModelId,
176 capable: ModelId,
177 efficient: ModelId,
178 contract_config: &ClassifierContractConfig,
179 config: EscalationJudgeConfig,
180 max_output_tokens: u64,
181) -> Result<JudgeClassifier<EscalationJudge, EscalationPolicy>> {
182 config.validate()?;
183 let contract =
184 ClassifierContract::from_config(contract_config, PROMPT_TEMPLATE, SCHEMA_TEMPLATE)?;
185 Ok(JudgeClassifier::new(
186 StructuredJudge::new(
187 EscalationInput { config },
188 contract,
189 SerdeDecoder::new(),
190 JudgeRuntimeConfig::new(max_output_tokens)?,
191 ),
192 judge_target,
193 EscalationPolicy { capable, efficient },
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}