fix: Improved subagent escalation handling

This commit is contained in:
2026-08-14 16:53:37 -06:00
parent e1b6e3f8c6
commit f57bd21ee4
2 changed files with 239 additions and 28 deletions
+130 -27
View File
@@ -186,42 +186,29 @@ pub async fn eval_tool_calls(
}) })
.collect(); .collect();
for (idx, call, result) in future::join_all(futs).await { for (idx, call, result) in future::join_all(futs).await {
indexed_results.push((idx, ToolResult::new(call, normalize_tool_result(result?)))); let value = match result {
Ok(v) => normalize_tool_result(v),
Err(e) => json!({"tool_call_error": format!("{e}")}),
};
indexed_results.push((idx, ToolResult::new(call, value)));
} }
} }
for (idx, call) in sequential_calls { for (idx, call) in sequential_calls {
let result = call.eval(ctx).await?; let value = match call.eval(ctx).await {
indexed_results.push((idx, ToolResult::new(call, normalize_tool_result(result)))); Ok(v) => normalize_tool_result(v),
Err(e) => json!({
"tool_call_error": format!(
"{e}. This tool is not available or the call failed; use only tools listed in your catalog."
)
}),
};
indexed_results.push((idx, ToolResult::new(call, value)));
} }
indexed_results.sort_unstable_by_key(|(idx, _)| *idx); indexed_results.sort_unstable_by_key(|(idx, _)| *idx);
output = indexed_results.into_iter().map(|(_, r)| r).collect(); output = indexed_results.into_iter().map(|(_, r)| r).collect();
if !output.is_empty() {
let (has_escalations, summary) = if ctx.current_depth == 0
&& let Some(queue) = ctx.root_escalation_queue()
&& queue.has_pending()
{
(true, queue.pending_summary())
} else {
(false, vec![])
};
if has_escalations {
let notification = json!({
"pending_escalations": summary,
"instruction": "Child agents are BLOCKED waiting for your reply. Call agent__reply_escalation for each pending escalation to unblock them."
});
let synthetic_call = ToolCall::new(
"__escalation_notification".to_string(),
json!({}),
Some("escalation_check".to_string()),
);
output.push(ToolResult::new(synthetic_call, notification));
}
}
{ {
let max_chars = ctx let max_chars = ctx
.agent .agent
@@ -236,6 +223,14 @@ pub async fn eval_tool_calls(
} }
} }
if ctx.current_depth == 0
&& let Some(queue) = ctx.root_escalation_queue()
&& queue.has_pending()
&& let Some(last) = output.last_mut()
{
inject_escalation_notification(last, queue.pending_summary());
}
Ok(output) Ok(output)
} }
@@ -252,6 +247,24 @@ fn normalize_tool_result(result: Value) -> Value {
} }
} }
fn inject_escalation_notification(last: &mut ToolResult, summary: Vec<Value>) {
let instruction = "Child agents are BLOCKED waiting for your reply. \
Call agent__reply_escalation for each pending escalation to unblock them.";
match &mut last.output {
Value::Object(map) => {
map.insert("pending_escalations".into(), json!(summary));
map.insert("escalation_instruction".into(), json!(instruction));
}
other => {
*other = json!({
"output": other.take(),
"pending_escalations": summary,
"escalation_instruction": instruction,
});
}
}
}
#[derive(Debug, Clone, Deserialize, Serialize)] #[derive(Debug, Clone, Deserialize, Serialize)]
pub struct ToolResult { pub struct ToolResult {
pub call: ToolCall, pub call: ToolCall,
@@ -1741,7 +1754,10 @@ fn format_json_colored_keys(value: &serde_json::Value) -> String {
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
use crate::config::{AppState, WorkingMode};
use crate::supervisor::escalation::{EscalationQueue, EscalationRequest};
use serde_json::json; use serde_json::json;
use std::sync::Arc;
fn call(name: &str, id: Option<&str>) -> ToolCall { fn call(name: &str, id: Option<&str>) -> ToolCall {
ToolCall::new(name.to_string(), json!({}), id.map(|s| s.to_string())) ToolCall::new(name.to_string(), json!({}), id.map(|s| s.to_string()))
@@ -1751,11 +1767,98 @@ mod tests {
ToolCall::new(name.to_string(), args, Some("id1".to_string())) ToolCall::new(name.to_string(), args, Some("id1".to_string()))
} }
fn run_async<F: Future>(f: F) -> F::Output {
tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap()
.block_on(f)
}
fn submit_escalation(queue: &EscalationQueue, id: &str) {
let (tx, _rx) = tokio::sync::oneshot::channel();
queue.submit(EscalationRequest {
id: id.to_string(),
from_agent_id: "a1".into(),
from_agent_name: "explore".into(),
question: "What do?".into(),
options: None,
reply_tx: tx,
});
}
#[test] #[test]
fn normalize_tool_result_substitutes_done_for_null() { fn normalize_tool_result_substitutes_done_for_null() {
assert_eq!(normalize_tool_result(Value::Null), json!("DONE")); assert_eq!(normalize_tool_result(Value::Null), json!("DONE"));
} }
#[test]
fn inject_escalation_notification_extends_object_output() {
let mut result = ToolResult::new(call("t", Some("id-1")), json!({"status": "ok"}));
inject_escalation_notification(&mut result, vec![json!({"escalation_id": "esc_1"})]);
assert_eq!(result.output["status"], "ok");
assert_eq!(
result.output["pending_escalations"],
json!([{"escalation_id": "esc_1"}])
);
assert!(
result.output["escalation_instruction"]
.as_str()
.unwrap()
.contains("agent__reply_escalation")
);
assert!(result.text.is_none());
}
#[test]
fn inject_escalation_notification_wraps_non_object_output() {
let mut result = ToolResult::new(call("t", Some("id-1")), json!("DONE"));
inject_escalation_notification(&mut result, vec![json!({"escalation_id": "esc_2"})]);
assert_eq!(result.output["output"], json!("DONE"));
assert_eq!(
result.output["pending_escalations"],
json!([{"escalation_id": "esc_2"}])
);
assert!(result.output["escalation_instruction"].is_string());
}
#[test]
fn eval_tool_calls_soft_fails_unknown_tool() {
let mut ctx = RequestContext::new(Arc::new(AppState::test_default()), WorkingMode::Cmd);
let calls = vec![call("__escalation_notification", Some("id-1"))];
let results = run_async(eval_tool_calls(&mut ctx, calls)).unwrap();
assert_eq!(results.len(), 1);
let err = results[0].output["tool_call_error"].as_str().unwrap();
assert!(err.contains("Unexpected call"));
assert!(err.contains("use only tools listed in your catalog"));
}
#[test]
fn eval_tool_calls_injects_escalations_into_last_result() {
let mut ctx = RequestContext::new(Arc::new(AppState::test_default()), WorkingMode::Cmd);
let queue = ctx.ensure_root_escalation_queue();
submit_escalation(&queue, "esc_1");
let calls = vec![call("unknown_tool", Some("id-1"))];
let results = run_async(eval_tool_calls(&mut ctx, calls)).unwrap();
assert_eq!(results.len(), 1);
assert!(
results
.iter()
.all(|r| r.call.name != "__escalation_notification")
);
let out = &results[0].output;
assert!(out["tool_call_error"].is_string());
assert_eq!(out["pending_escalations"][0]["escalation_id"], "esc_1");
assert!(
out["escalation_instruction"]
.as_str()
.unwrap()
.contains("agent__reply_escalation")
);
}
#[test] #[test]
fn normalize_tool_result_preserves_non_null_values() { fn normalize_tool_result_preserves_non_null_values() {
assert_eq!( assert_eq!(
+109 -1
View File
@@ -85,7 +85,17 @@ pub fn check_pending_agents_guardrail(ctx: &mut RequestContext) -> GuardrailActi
} }
ctx.pending_agents_guardrail_count += 1; ctx.pending_agents_guardrail_count += 1;
GuardrailAction::Inject(build_pending_agents_guardrail_prompt(&pending)) let mut prompt = build_pending_agents_guardrail_prompt(&pending);
if let Some(queue) = ctx.root_escalation_queue()
&& queue.has_pending()
{
let summary = serde_json::to_string(&queue.pending_summary()).unwrap_or_default();
prompt.push_str(&format!(
"\n\nAdditionally, child agents have pending escalations blocking them. Reply to each \
via `agent__reply_escalation` first:\n{summary}"
));
}
GuardrailAction::Inject(prompt)
} }
pub fn escalation_function_declarations() -> Vec<FunctionDeclaration> { pub fn escalation_function_declarations() -> Vec<FunctionDeclaration> {
@@ -1996,4 +2006,102 @@ mod tests {
let q2 = ctx.ensure_root_escalation_queue(); let q2 = ctx.ensure_root_escalation_queue();
assert!(Arc::ptr_eq(&q1, &q2)); assert!(Arc::ptr_eq(&q1, &q2));
} }
#[test]
fn guardrail_prompt_mentions_pending_escalations() {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap();
rt.block_on(async {
let mut ctx = ctx_with_supervisor(4, 3);
let join_handle = tokio::spawn(async {
time::sleep(Duration::from_secs(60)).await;
Ok(AgentResult {
id: "slow".into(),
agent_name: "test".into(),
output: String::new(),
exit_status: AgentExitStatus::Completed,
})
});
let handle = AgentHandle {
id: "slow".into(),
agent_name: "test".into(),
depth: 1,
inbox: Arc::new(Inbox::new()),
abort_signal: create_abort_signal(),
join_handle,
child_supervisor: None,
};
ctx.supervisor
.as_ref()
.unwrap()
.write()
.register(handle)
.unwrap();
let queue = ctx.ensure_root_escalation_queue();
let (tx, _rx) = tokio::sync::oneshot::channel();
queue.submit(EscalationRequest {
id: "esc_9".into(),
from_agent_id: "a1".into(),
from_agent_name: "explore".into(),
question: "Which option?".into(),
options: None,
reply_tx: tx,
});
match check_pending_agents_guardrail(&mut ctx) {
GuardrailAction::Inject(prompt) => {
assert!(prompt.contains("agent__reply_escalation"));
assert!(prompt.contains("esc_9"));
}
_ => panic!("expected Inject action"),
}
});
}
#[test]
fn guardrail_prompt_omits_escalations_when_none_pending() {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap();
rt.block_on(async {
let mut ctx = ctx_with_supervisor(4, 3);
let join_handle = tokio::spawn(async {
time::sleep(Duration::from_secs(60)).await;
Ok(AgentResult {
id: "slow".into(),
agent_name: "test".into(),
output: String::new(),
exit_status: AgentExitStatus::Completed,
})
});
let handle = AgentHandle {
id: "slow".into(),
agent_name: "test".into(),
depth: 1,
inbox: Arc::new(Inbox::new()),
abort_signal: create_abort_signal(),
join_handle,
child_supervisor: None,
};
ctx.supervisor
.as_ref()
.unwrap()
.write()
.register(handle)
.unwrap();
match check_pending_agents_guardrail(&mut ctx) {
GuardrailAction::Inject(prompt) => {
assert!(!prompt.contains("agent__reply_escalation"));
}
_ => panic!("expected Inject action"),
}
});
}
} }