Fix cancellation during subagent waits

This commit is contained in:
Bryan Helmkamp 2026-07-23 18:04:58 -04:00
parent b558af070f
commit e78f5ff957
No known key found for this signature in database
2 changed files with 308 additions and 54 deletions

View file

@ -2073,7 +2073,7 @@ mod tests {
use super::*;
use crate::config::{ToolAccess, ToolAccessPolicy, ToolApprovalAdapter, ToolExposureMode};
use crate::skills::{Skill, make_use_skill_tool};
use crate::subagent::SubAgentStatus;
use crate::subagent::{SubAgentStatus, make_wait_tool};
use crate::test_support::*;
use crate::tool_registry::{RegisteredTool, ToolContext, ToolRegistry, ToolSource};
@ -4703,6 +4703,165 @@ mod tests {
);
}
async fn make_parent_waiting_on_blocked_subagent() -> (
Session,
Arc<AsyncMutex<SubAgentManager>>,
String,
CancellationToken,
) {
let block_until_cancelled = RegisteredTool {
definition: ToolDefinition {
name: "block_until_cancelled".into(),
description: "Waits until cancelled".into(),
parameters: serde_json::json!({"type": "object"}),
},
executor: Arc::new(|_args, ctx| {
Box::pin(async move {
ctx.cancel.cancelled().await;
Ok("cancelled".to_string())
})
}),
source: ToolSource::Native,
};
let mut child_registry = ToolRegistry::new();
child_registry.register(block_until_cancelled);
let child = make_session_with_tools(
vec![tool_call_response(
"block_until_cancelled",
"child_call",
serde_json::json!({}),
)],
child_registry,
)
.await;
let child_cancel = child.cancel_token();
let manager = Arc::new(AsyncMutex::new(SubAgentManager::new(3)));
let agent_id = manager
.lock()
.await
.spawn(child, "block until cancelled".into(), 0)
.unwrap();
let mut parent_registry = ToolRegistry::new();
parent_registry.register(make_wait_tool(manager.clone()));
let parent_provider = Arc::new(ScriptedStreamProvider::new(vec![
ScriptedStreamCall::Response(Box::new(tool_call_response(
"wait",
"parent_wait_call",
serde_json::json!({ "agent_id": agent_id }),
))),
ScriptedStreamCall::Response(Box::new(text_response("resumed"))),
]));
let client = make_client(parent_provider).await;
let profile = Arc::new(TestProfile::with_tools(parent_registry));
let env = Arc::new(MockSandbox::default());
let session = Session::new(
client,
profile,
env,
SessionOptions::default(),
Some(manager.clone()),
);
manager
.lock()
.await
.set_event_callback(session.sub_agent_event_callback());
(session, manager, agent_id, child_cancel)
}
async fn wait_for_agent_event(
rx: &mut broadcast::Receiver<SessionEvent>,
predicate: impl Fn(&AgentEvent) -> bool,
) {
loop {
let event = rx
.recv()
.await
.expect("session event stream should remain open");
if predicate(&event.event) {
return;
}
}
}
#[tokio::test]
async fn control_interrupt_during_subagent_wait_closes_child_and_resumes_after_steer() {
let (mut session, manager, agent_id, child_cancel) =
make_parent_waiting_on_blocked_subagent().await;
let control = session.control_handle();
let mut events = session.subscribe();
let controller = tokio::spawn(async move {
wait_for_agent_event(&mut events, |event| {
matches!(
event,
AgentEvent::ToolCallStarted { tool_name, .. } if tool_name == "wait"
)
})
.await;
control.interrupt(None);
wait_for_agent_event(&mut events, |event| {
matches!(event, AgentEvent::SubAgentClosed { .. })
})
.await;
control.steer("resume after interrupt".into(), None);
});
timeout(
Duration::from_secs(1),
session.process_input("wait for the child"),
)
.await
.expect("interrupt should unblock the subagent wait")
.unwrap();
controller.await.unwrap();
assert_eq!(session.state(), SessionState::Idle);
assert!(child_cancel.is_cancelled());
assert!(matches!(
manager.lock().await.status(&agent_id),
Some(SubAgentStatus::Closed)
));
}
#[tokio::test]
async fn terminal_cancel_during_subagent_wait_closes_child_and_session() {
let (mut session, manager, agent_id, child_cancel) =
make_parent_waiting_on_blocked_subagent().await;
let cancel = session.cancel_token();
let mut events = session.subscribe();
let controller = tokio::spawn(async move {
wait_for_agent_event(&mut events, |event| {
matches!(
event,
AgentEvent::ToolCallStarted { tool_name, .. } if tool_name == "wait"
)
})
.await;
cancel.cancel();
});
let result = timeout(
Duration::from_secs(1),
session.process_input("wait for the child"),
)
.await
.expect("terminal cancellation should unblock the subagent wait");
controller.await.unwrap();
assert!(matches!(
result,
Err(Error::Interrupted(InterruptReason::Cancelled))
));
assert_eq!(session.state(), SessionState::Closed);
assert!(child_cancel.is_cancelled());
assert!(matches!(
manager.lock().await.status(&agent_id),
Some(SubAgentStatus::Closed)
));
}
#[tokio::test]
async fn close_cleans_up_subagents_before_emitting_session_ended() {
use crate::subagent::SubAgentManager;

View file

@ -2,8 +2,9 @@ use std::collections::{HashMap, VecDeque};
use std::sync::{Arc, Mutex};
use fabro_llm::types::ToolDefinition;
use futures::future::{BoxFuture, FutureExt, Shared};
use tokio::sync::Mutex as AsyncMutex;
use tokio::task::JoinHandle;
use tokio::task::{AbortHandle, JoinHandle};
use tokio_util::sync::CancellationToken;
use crate::error::Error;
@ -29,6 +30,23 @@ pub struct SubAgentResult {
pub turns_used: usize,
}
type SubAgentTask = Shared<BoxFuture<'static, Result<SubAgentResult, Error>>>;
fn shared_subagent_task(
task: JoinHandle<Result<SubAgentResult, Error>>,
) -> (SubAgentTask, AbortHandle) {
let abort_handle = task.abort_handle();
let task = async move {
match task.await {
Ok(result) => result,
Err(e) => Err(Error::InvalidState(format!("Agent task panicked: {e}"))),
}
}
.boxed()
.shared();
(task, abort_handle)
}
#[derive(Debug, Clone)]
pub enum SubAgentStatus {
Running,
@ -37,7 +55,8 @@ pub enum SubAgentStatus {
}
pub struct SubAgent {
task: Option<JoinHandle<Result<SubAgentResult, Error>>>,
task: SubAgentTask,
abort_handle: AbortHandle,
followup_queue: Arc<Mutex<VecDeque<String>>>,
cancel_token: CancellationToken,
depth: usize,
@ -124,9 +143,11 @@ impl SubAgentManager {
turns_used: turns.len(),
})
});
let (task, abort_handle) = shared_subagent_task(task);
self.agents.insert(agent_id.clone(), SubAgent {
task: Some(task),
task,
abort_handle,
followup_queue,
cancel_token,
depth: depth + 1,
@ -168,45 +189,52 @@ impl SubAgentManager {
}
pub async fn wait(&mut self, agent_id: &str) -> Result<SubAgentResult, Error> {
// Phase 1: Check existence and current status
let agent = self.agents.get(agent_id);
let depth = match agent {
None => {
return Err(Error::InvalidState(format!(
"No agent found with id: {agent_id} (it was never spawned)"
)));
}
Some(a) => a.depth,
};
let task = self.result_future(agent_id)?;
let task_result = task.await;
self.record_result(agent_id, task_result)
}
match &self.agents[agent_id].status {
SubAgentStatus::Closed => {
return Err(Error::InvalidState(format!(
"Agent {agent_id} has been closed"
)));
}
SubAgentStatus::Finished(result) => {
return result.clone();
}
SubAgentStatus::Running => {}
fn result_future(&self, agent_id: &str) -> Result<SubAgentTask, Error> {
let agent = self.agents.get(agent_id).ok_or_else(|| {
Error::InvalidState(format!(
"No agent found with id: {agent_id} (it was never spawned)"
))
})?;
match &agent.status {
SubAgentStatus::Closed => Err(Error::InvalidState(format!(
"Agent {agent_id} has been closed"
))),
SubAgentStatus::Running | SubAgentStatus::Finished(_) => Ok(agent.task.clone()),
}
}
// Phase 2: Take the JoinHandle (brief mutable borrow, no await)
let join_handle = self
.agents
.get_mut(agent_id)
.expect("agent should still exist after status check")
.task
.take()
.ok_or_else(|| Error::InvalidState(format!("Agent {agent_id} has no running task")))?;
fn record_result(
&mut self,
agent_id: &str,
task_result: Result<SubAgentResult, Error>,
) -> Result<SubAgentResult, Error> {
let depth = {
let agent = self.agents.get_mut(agent_id).ok_or_else(|| {
Error::InvalidState(format!(
"No agent found with id: {agent_id} (it was never spawned)"
))
})?;
// Phase 3: Await the task (no borrow held)
let task_result = match join_handle.await {
Ok(result) => result,
Err(e) => Err(Error::InvalidState(format!("Agent task panicked: {e}"))),
match &agent.status {
SubAgentStatus::Closed => {
return Err(Error::InvalidState(format!(
"Agent {agent_id} has been closed"
)));
}
SubAgentStatus::Finished(result) => return result.clone(),
SubAgentStatus::Running => {}
}
agent.status = SubAgentStatus::Finished(task_result.clone());
agent.depth
};
// Phase 4: Emit event
match &task_result {
Ok(result) => {
self.emit_event(AgentEvent::SubAgentCompleted {
@ -225,17 +253,7 @@ impl SubAgentManager {
}
}
// Phase 5: Store result in status and return clone
let agent = self
.agents
.get_mut(agent_id)
.expect("agent should still exist when storing task result");
agent.status = SubAgentStatus::Finished(task_result);
match &agent.status {
SubAgentStatus::Finished(result) => result.clone(),
_ => unreachable!("agent status was just assigned to Finished on the line above"),
}
task_result
}
pub fn close(&mut self, agent_id: &str) -> Result<(), Error> {
@ -253,9 +271,7 @@ impl SubAgentManager {
}
SubAgentStatus::Running => {
agent.cancel_token.cancel();
if let Some(join_handle) = agent.task.take() {
join_handle.abort();
}
agent.abort_handle.abort();
}
SubAgentStatus::Finished(_) => {
// No task to cancel, just transition status
@ -415,13 +431,28 @@ pub fn make_wait_tool(manager: Arc<AsyncMutex<SubAgentManager>>) -> RegisteredTo
"required": ["agent_id"]
}),
},
executor: Arc::new(move |args, _ctx| {
executor: Arc::new(move |args, ctx| {
let manager = manager.clone();
Box::pin(async move {
let agent_id = required_str(&args, "agent_id")?;
let mut mgr = manager.lock().await;
let result = mgr.wait(agent_id).await.map_err(|e| e.to_string())?;
let task = {
let mgr = manager.lock().await;
mgr.result_future(agent_id).map_err(|e| e.to_string())?
};
let task_result = tokio::select! {
biased;
() = ctx.cancel.cancelled() => {
let _ = manager.lock().await.close(agent_id);
return Err("Cancelled".to_string());
}
result = task => result,
};
let result = manager
.lock()
.await
.record_result(agent_id, task_result)
.map_err(|e| e.to_string())?;
Ok(format!(
"Agent completed (success: {}, turns: {})\n\n{}",
result.success, result.turns_used, result.output
@ -471,6 +502,7 @@ mod tests {
use super::*;
use crate::config::SessionOptions;
use crate::test_support::*;
use crate::tool_registry::ToolContext;
// --- Tests ---
@ -610,6 +642,69 @@ mod tests {
));
}
#[tokio::test]
async fn wait_tool_returns_when_context_is_cancelled() {
let child_cancel = CancellationToken::new();
let child_cancel_probe = child_cancel.clone();
let task_cancel = child_cancel.clone();
let task = tokio::spawn(async move {
task_cancel.cancelled().await;
Ok(SubAgentResult {
output: String::new(),
success: false,
turns_used: 0,
})
});
let (task, abort_handle) = shared_subagent_task(task);
let agent_id = "blocked-agent".to_string();
let manager = Arc::new(AsyncMutex::new(SubAgentManager::new(3)));
manager
.lock()
.await
.agents
.insert(agent_id.clone(), SubAgent {
task,
abort_handle,
followup_queue: Arc::new(Mutex::new(VecDeque::new())),
cancel_token: child_cancel,
depth: 1,
status: SubAgentStatus::Running,
});
let tool = make_wait_tool(manager.clone());
let tool_cancel = CancellationToken::new();
let ctx = ToolContext {
env: Arc::new(MockSandbox::default()),
cancel: tool_cancel.clone(),
tool_env_provider: None,
session_id: None,
root_session_id: None,
tool_call_id: None,
agent_event_emitter: None,
};
let mut wait = (tool.executor)(serde_json::json!({ "agent_id": agent_id }), ctx);
assert!(
futures::poll!(wait.as_mut()).is_pending(),
"blocked subagent should leave the wait tool pending"
);
tool_cancel.cancel();
let result = time::timeout(std::time::Duration::from_millis(100), wait.as_mut()).await;
drop(wait);
manager.lock().await.close_all();
let result =
result.expect("wait tool should return promptly when its context is cancelled");
assert_eq!(result, Err("Cancelled".to_string()));
assert!(child_cancel_probe.is_cancelled());
assert!(matches!(
manager.lock().await.status(&agent_id),
Some(SubAgentStatus::Closed)
));
}
#[test]
fn tool_definitions_correct() {
let manager = Arc::new(AsyncMutex::new(SubAgentManager::new(3)));