diff --git a/Cargo.lock b/Cargo.lock index 77e6269..364c36e 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -61,7 +61,7 @@ dependencies = [ [[package]] name = "agentkit" -version = "0.10.8" +version = "0.10.11" dependencies = [ "agentkit-acp", "agentkit-adapter-completions", @@ -92,12 +92,13 @@ dependencies = [ [[package]] name = "agentkit-acp" -version = "0.10.8" +version = "0.10.11" dependencies = [ "agent-client-protocol", "agentkit-core", "agentkit-integration-tests", "agentkit-loop", + "agentkit-task-manager", "agentkit-tools-core", "async-trait", "base64 0.22.1", @@ -109,7 +110,7 @@ dependencies = [ [[package]] name = "agentkit-adapter-completions" -version = "0.10.6" +version = "0.10.7" dependencies = [ "agentkit-core", "agentkit-http", @@ -210,7 +211,7 @@ dependencies = [ [[package]] name = "agentkit-loop" -version = "0.10.7" +version = "0.10.10" dependencies = [ "agentkit-core", "agentkit-task-manager", @@ -261,7 +262,7 @@ dependencies = [ [[package]] name = "agentkit-provider-anthropic" -version = "0.10.7" +version = "0.10.8" dependencies = [ "agentkit-core", "agentkit-http", @@ -400,7 +401,7 @@ dependencies = [ [[package]] name = "agentkit-task-manager" -version = "0.10.6" +version = "0.10.7" dependencies = [ "agentkit-core", "agentkit-tools-core", diff --git a/crates/agentkit-acp/Cargo.toml b/crates/agentkit-acp/Cargo.toml index 5ee9288..993b3f6 100644 --- a/crates/agentkit-acp/Cargo.toml +++ b/crates/agentkit-acp/Cargo.toml @@ -4,7 +4,7 @@ homepage.workspace = true name = "agentkit-acp" readme = "README.md" repository.workspace = true -version = "0.10.8" +version = "0.10.11" edition.workspace = true license.workspace = true rust-version.workspace = true @@ -18,7 +18,7 @@ protocol-v2 = ["agent-client-protocol/unstable_protocol_v2"] [dependencies] agent-client-protocol = "=2.0.0" agentkit-core = { version = "0.10.5", path = "../agentkit-core" } -agentkit-loop = { version = "0.10.5", path = "../agentkit-loop" } +agentkit-loop = { version = "0.10.10", path = "../agentkit-loop" } agentkit-tools-core = { version = "0.10.5", path = "../agentkit-tools-core" } async-trait.workspace = true base64.workspace = true @@ -29,4 +29,5 @@ tracing.workspace = true [dev-dependencies] agentkit-integration-tests = { path = "../agentkit-integration-tests" } +agentkit-task-manager = { path = "../agentkit-task-manager" } tokio = { workspace = true, features = ["macros", "rt-multi-thread", "time"] } diff --git a/crates/agentkit-acp/src/v2.rs b/crates/agentkit-acp/src/v2.rs index ace92cb..ba5b8b3 100644 --- a/crates/agentkit-acp/src/v2.rs +++ b/crates/agentkit-acp/src/v2.rs @@ -14,10 +14,11 @@ use agentkit_core::{ CancellationController, CancellationHandle, DataRef, Delta, FilePart, FinishReason, Item, ItemKind, MediaPart, MetadataMap, Modality, Part, PartId, PartKind, SessionId as AgentkitSessionId, StructuredPart, TextPart, ToolCallPart, ToolOutput, - ToolResultPart, + ToolResultPart, TurnCancellation, TurnId, }; use agentkit_loop::{ - AgentEvent, LoopInterrupt, LoopObserver, LoopStep, ModelAdapter, ModelSession, ObservedEvent, + AgentEvent, LoopError, LoopInterrupt, LoopObserver, LoopStep, ModelAdapter, ModelSession, + ObservedEvent, }; use async_trait::async_trait; use serde_json::json; @@ -47,12 +48,30 @@ enum ClientMessage { #[derive(Clone)] struct ClientHandle { tx: mpsc::UnboundedSender, + #[cfg(test)] + before_update: Option>, } impl ClientHandle { fn channel() -> (Self, mpsc::UnboundedReceiver) { let (tx, rx) = mpsc::unbounded_channel(); - (Self { tx }, rx) + ( + Self { + tx, + #[cfg(test)] + before_update: None, + }, + rx, + ) + } + + #[cfg(test)] + fn channel_with_update_hook( + hook: impl Fn(&wire::SessionUpdate) + Send + Sync + 'static, + ) -> (Self, mpsc::UnboundedReceiver) { + let (mut client, rx) = Self::channel(); + client.before_update = Some(Arc::new(hook)); + (client, rx) } fn update( @@ -60,6 +79,10 @@ impl ClientHandle { session_id: wire::SessionId, update: wire::SessionUpdate, ) -> Result<(), AcpRuntimeError> { + #[cfg(test)] + if let Some(hook) = &self.before_update { + hook(&update); + } self.tx .send(ClientMessage::Update(Box::new( wire::UpdateSessionNotification::new(session_id, update), @@ -107,6 +130,21 @@ struct IntegrationSession { next_message: AtomicU64, current_messages: Mutex>, part_kinds: Mutex>, + unsupported_approval: Mutex>, + prompt_state: Mutex>, +} + +struct PromptState { + active_prompt: Arc, + lifecycle: Arc>, + pending_owner: Option, + turn_owners: HashMap, +} + +#[derive(Clone)] +struct PromptOwner { + id: u64, + cancellation: TurnCancellation, } #[derive(Default)] @@ -156,6 +194,8 @@ impl AcpIntegration { next_message: AtomicU64::new(1), current_messages: Mutex::new(None), part_kinds: Mutex::new(HashMap::new()), + unsupported_approval: Mutex::new(None), + prompt_state: Mutex::new(None), }), ); Ok(()) @@ -189,11 +229,45 @@ impl AcpIntegration { .ok_or_else(|| AcpRuntimeError::SessionNotFound(session_id.to_string())) } + fn install_prompt_state( + &self, + session_id: &wire::SessionId, + active_prompt: Arc, + lifecycle: Arc>, + ) -> Result<(), AcpRuntimeError> { + let session = self.session(session_id)?; + *session + .prompt_state + .lock() + .unwrap_or_else(|error| error.into_inner()) = Some(PromptState { + active_prompt, + lifecycle, + pending_owner: None, + turn_owners: HashMap::new(), + }); + Ok(()) + } + fn begin_prompt( &self, session_id: &wire::SessionId, + owner: u64, + cancellation: TurnCancellation, ) -> Result { let session = self.session(session_id)?; + { + let mut prompt_state = session + .prompt_state + .lock() + .unwrap_or_else(|error| error.into_inner()); + let prompt_state = prompt_state.as_mut().ok_or_else(|| { + AcpRuntimeError::Sdk("ACP v2 prompt state is not initialized".into()) + })?; + prompt_state.pending_owner = Some(PromptOwner { + id: owner, + cancellation, + }); + } let sequence = session.next_message.fetch_add(1, Ordering::Relaxed); finish_model_message(&session); Ok(wire::MessageId::new(format!( @@ -201,12 +275,96 @@ impl AcpIntegration { ))) } - fn finish_prompt(&self, session_id: &wire::SessionId) { + fn finish_prompt(&self, session_id: &wire::SessionId, owner: u64) { if let Ok(session) = self.session(session_id) { + if let Some(prompt_state) = session + .prompt_state + .lock() + .unwrap_or_else(|error| error.into_inner()) + .as_mut() + { + if prompt_state + .pending_owner + .as_ref() + .is_some_and(|pending| pending.id == owner) + { + prompt_state.pending_owner = None; + } + prompt_state + .turn_owners + .retain(|_, turn_owner| turn_owner.id != owner); + } finish_model_message(&session); } } + fn mark_unsupported_approval( + &self, + session_id: &wire::SessionId, + cancellation: CancellationHandle, + generation: u64, + ) { + if let Ok(session) = self.session(session_id) { + *session + .unsupported_approval + .lock() + .unwrap_or_else(|error| error.into_inner()) = Some((cancellation, generation)); + } + } + + fn clear_unsupported_approval(&self, session_id: &wire::SessionId) { + if let Ok(session) = self.session(session_id) { + session + .unsupported_approval + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take(); + } + } + + fn route_turn_finished( + session: &IntegrationSession, + result: &agentkit_loop::TurnResult, + unsupported_approval: Option<(CancellationHandle, u64)>, + ) { + let prompt_owner = session + .prompt_state + .lock() + .unwrap_or_else(|error| error.into_inner()) + .as_mut() + .and_then(|prompt_state| { + prompt_state + .turn_owners + .remove(&result.turn_id) + .map(|owner| (Arc::clone(&prompt_state.active_prompt), owner)) + }); + let Some((active_prompt, owner)) = prompt_owner else { + return; + }; + let prompt_cancelled = owner.cancellation.is_cancelled(); + let stop_reason = if prompt_cancelled { + wire::StopReason::Cancelled + } else { + match unsupported_approval { + Some((cancellation, generation)) + if !cancellation.is_cancelled_since(generation) => + { + error_stop_reason() + } + _ => finish_reason_to_stop_reason(&result.finish_reason), + } + }; + release_prompt(&active_prompt, owner.id); + if let Err(error) = session.client.update( + session.acp_session_id.clone(), + wire::SessionUpdate::StateUpdate(wire::StateUpdate::Idle( + wire::IdleStateUpdate::new().stop_reason(stop_reason), + )), + ) { + tracing::debug!(%error, "failed to queue ACP v2 idle update"); + } + } + fn route_event(&self, session_id: &AgentkitSessionId, event: AgentEvent) { let session = { let inner = self.inner.read().unwrap_or_else(|error| error.into_inner()); @@ -220,12 +378,53 @@ impl AcpIntegration { }; match &event { - AgentEvent::TurnStarted { .. } => { - start_model_message(&session); + AgentEvent::TurnStarted { turn_id, .. } => { + let prompt_owned = session + .prompt_state + .lock() + .unwrap_or_else(|error| error.into_inner()) + .as_mut() + .and_then(|prompt_state| { + let owner = prompt_state.pending_owner.take()?; + prompt_state.turn_owners.insert(turn_id.clone(), owner); + Some(()) + }) + .is_some(); + if prompt_owned { + start_model_message(&session); + if let Err(error) = session.client.update( + session.acp_session_id.clone(), + wire::SessionUpdate::StateUpdate(wire::StateUpdate::Running( + wire::RunningStateUpdate::new(), + )), + ) { + tracing::debug!(%error, "failed to queue ACP v2 running update"); + } + } return; } - AgentEvent::TurnFinished(_) => { + AgentEvent::TurnFinished(result) => { finish_model_message(&session); + let unsupported_approval = session + .unsupported_approval + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take(); + let lifecycle = session + .prompt_state + .lock() + .unwrap_or_else(|error| error.into_inner()) + .as_ref() + .and_then(|prompt_state| { + prompt_state + .turn_owners + .contains_key(&result.turn_id) + .then(|| Arc::clone(&prompt_state.lifecycle)) + }); + if let Some(lifecycle) = lifecycle { + let _lifecycle = lifecycle.lock().unwrap_or_else(|error| error.into_inner()); + Self::route_turn_finished(&session, result, unsupported_approval); + } return; } AgentEvent::ToolExecutionStarted(_) | AgentEvent::ToolResultReceived(_) => { @@ -558,9 +757,12 @@ struct SessionEntry { commands: mpsc::UnboundedSender, cancellation: CancellationController, info: wire::SessionInfo, - busy: Arc, + active_prompt: Arc, + driving_prompt: Arc, + cancelled_prompt: Arc, + next_prompt_owner: AtomicU64, closed: AtomicBool, - lifecycle: Mutex<()>, + lifecycle: Arc>, task: Mutex>>, drain_task: Mutex>>, } @@ -569,8 +771,10 @@ enum SessionCommand { Prompt { request: wire::PromptRequest, items: Vec, + prompt_cancellation: TurnCancellation, cancellation_generation: u64, - response: oneshot::Sender, AcpRuntimeError>>, + owner: u64, + start: oneshot::Receiver<()>, }, Shutdown, } @@ -640,11 +844,23 @@ where json!(request.additional_directories), ); + let active_prompt = Arc::new(AtomicU64::new(0)); + let driving_prompt = Arc::new(AtomicU64::new(0)); + let cancelled_prompt = Arc::new(AtomicU64::new(0)); + let lifecycle = Arc::new(Mutex::new(())); self.integration.bind( acp_session_id.clone(), agentkit_session_id.clone(), client.clone(), )?; + if let Err(error) = self.integration.install_prompt_state( + &acp_session_id, + Arc::clone(&active_prompt), + Arc::clone(&lifecycle), + ) { + let _ = self.integration.unbind(&acp_session_id); + return Err(error); + } let drain_task = tokio::spawn(drain_client_messages(client_messages, cx)); let ctx = AcpAgentFactoryContext { acp_session_id: acp_session_id.clone(), @@ -670,8 +886,10 @@ where }; let (commands, rx) = mpsc::unbounded_channel(); - let busy = Arc::new(AtomicBool::new(false)); - let worker_busy = Arc::clone(&busy); + let worker_active_prompt = Arc::clone(&active_prompt); + let worker_driving_prompt = Arc::clone(&driving_prompt); + let worker_cancelled_prompt = Arc::clone(&cancelled_prompt); + let worker_lifecycle = Arc::clone(&lifecycle); let integration = Arc::clone(&self.integration); let worker_session_id = acp_session_id.clone(); let worker_cancellation = cancellation.handle(); @@ -682,7 +900,10 @@ where client, integration, worker_cancellation, - worker_busy, + worker_active_prompt, + worker_driving_prompt, + worker_cancelled_prompt, + worker_lifecycle, rx, ) .await; @@ -691,9 +912,12 @@ where commands, cancellation, info, - busy, + active_prompt, + driving_prompt, + cancelled_prompt, + next_prompt_owner: AtomicU64::new(1), closed: AtomicBool::new(false), - lifecycle: Mutex::new(()), + lifecycle, task: Mutex::new(Some(task)), drain_task: Mutex::new(Some(drain_task)), }); @@ -774,7 +998,7 @@ where .get(&request.session_id) .cloned() .ok_or_else(|| AcpRuntimeError::SessionNotFound(request.session_id.to_string()))?; - let (tx, rx) = oneshot::channel(); + let (start_tx, start_rx) = oneshot::channel(); { let _lifecycle = entry .lifecycle @@ -785,31 +1009,38 @@ where request.session_id.to_string(), )); } + let mut owner = entry.next_prompt_owner.fetch_add(1, Ordering::Relaxed); + if owner == 0 { + owner = entry.next_prompt_owner.fetch_add(1, Ordering::Relaxed); + } if entry - .busy - .compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire) + .active_prompt + .compare_exchange(0, owner, Ordering::AcqRel, Ordering::Acquire) .is_err() { return Err(AcpRuntimeError::Unsupported( "session is already running a prompt".into(), )); } - let cancellation_generation = entry.cancellation.handle().generation(); + let prompt_cancellation = entry.cancellation.handle().checkpoint(); + let cancellation_generation = prompt_cancellation.generation(); if entry .commands .send(SessionCommand::Prompt { request, items, + prompt_cancellation, cancellation_generation, - response: tx, + owner, + start: start_rx, }) .is_err() { - entry.busy.store(false, Ordering::Release); + release_prompt(&entry.active_prompt, owner); return Err(AcpRuntimeError::ClientClosed); } } - rx.await.map_err(|_| AcpRuntimeError::ClientClosed)? + Ok(start_tx) } async fn cancel( @@ -827,8 +1058,18 @@ where .lifecycle .lock() .unwrap_or_else(|error| error.into_inner()); - if !entry.closed.load(Ordering::Acquire) && entry.busy.load(Ordering::Acquire) { - entry.cancellation.interrupt(); + if !entry.closed.load(Ordering::Acquire) { + let driving = entry.driving_prompt.load(Ordering::Acquire); + if driving != 0 { + entry.cancellation.interrupt(); + } else { + let queued = entry.active_prompt.load(Ordering::Acquire); + if queued != 0 { + entry.cancelled_prompt.store(queued, Ordering::Release); + } else { + entry.cancellation.interrupt(); + } + } } } Ok(()) @@ -928,93 +1169,222 @@ async fn stop_client(entry: Arc) { } } +fn release_prompt(active_prompt: &AtomicU64, owner: u64) { + let _ = active_prompt.compare_exchange(owner, 0, Ordering::AcqRel, Ordering::Acquire); +} + +fn clear_prompt_tracking( + active_prompt: &AtomicU64, + driving_prompt: &AtomicU64, + cancelled_prompt: &AtomicU64, + owner: u64, +) { + let _ = driving_prompt.compare_exchange(owner, 0, Ordering::AcqRel, Ordering::Acquire); + let _ = cancelled_prompt.compare_exchange(owner, 0, Ordering::AcqRel, Ordering::Acquire); + release_prompt(active_prompt, owner); +} + +#[allow(clippy::too_many_arguments)] +fn fail_accepted_prompt( + client: &ClientHandle, + integration: &AcpIntegration, + session_id: &wire::SessionId, + active_prompt: &AtomicU64, + driving_prompt: &AtomicU64, + cancelled_prompt: &AtomicU64, + lifecycle: &Mutex<()>, + owner: u64, + prompt_began: bool, +) { + let _lifecycle = lifecycle.lock().unwrap_or_else(|error| error.into_inner()); + if prompt_began { + integration.finish_prompt(session_id, owner); + } + clear_prompt_tracking(active_prompt, driving_prompt, cancelled_prompt, owner); + if let Err(error) = client.update( + session_id.clone(), + wire::SessionUpdate::StateUpdate(wire::StateUpdate::Idle( + wire::IdleStateUpdate::new().stop_reason(error_stop_reason()), + )), + ) { + tracing::debug!(%error, owner, "failed to queue ACP v2 error idle update"); + } +} + async fn session_worker( session_id: wire::SessionId, mut driver: agentkit_loop::LoopDriver, client: ClientHandle, integration: Arc, cancellation: CancellationHandle, - busy: Arc, + active_prompt: Arc, + driving_prompt: Arc, + cancelled_prompt: Arc, + lifecycle: Arc>, mut commands: mpsc::UnboundedReceiver, ) where S: ModelSession + Send + 'static, { - while let Some(command) = commands.recv().await { + enum SessionAction { + Command(Option), + LoopUpdate(Result<(), LoopError>), + } + + // Let one fresh prompt win a race with unrelated background work, then + // prefer the deferred update so repeated prompts cannot starve its delivery. + let mut prefer_loop_update = false; + loop { + let action = if prefer_loop_update { + tokio::select! { + biased; + wake = driver.wait_for_loop_update() => SessionAction::LoopUpdate(wake), + command = commands.recv() => SessionAction::Command(command), + } + } else { + tokio::select! { + biased; + command = commands.recv() => SessionAction::Command(command), + wake = driver.wait_for_loop_update() => SessionAction::LoopUpdate(wake), + } + }; + let command = match action { + SessionAction::Command(Some(command)) => command, + SessionAction::Command(None) => break, + SessionAction::LoopUpdate(wake) => { + if let Err(error) = wake { + tracing::debug!(%error, "failed waiting for an ACP v2 loop update"); + break; + } + prefer_loop_update = false; + drive_prompt( + &mut driver, + &integration, + &session_id, + &cancellation, + cancellation.generation(), + ) + .await; + if let Err(error) = client.flush().await { + tracing::debug!(%error, "failed to flush idle ACP v2 loop update"); + } + continue; + } + }; let SessionCommand::Prompt { request, items, + prompt_cancellation, cancellation_generation, - response, + owner, + start, } = command else { break; }; - if let Err(error) = driver - .submit_input(items) - .map_err(|error| AcpRuntimeError::Loop(error.to_string())) - { - busy.store(false, Ordering::Release); - let _ = response.send(Err(error)); + if start.await.is_err() { + let _lifecycle = lifecycle.lock().unwrap_or_else(|error| error.into_inner()); + clear_prompt_tracking(&active_prompt, &driving_prompt, &cancelled_prompt, owner); continue; } - let user_message_id = match integration.begin_prompt(&session_id) { - Ok(message_id) => message_id, - Err(error) => { - busy.store(false, Ordering::Release); - let _ = response.send(Err(error)); - continue; + let should_start = { + let _lifecycle = lifecycle.lock().unwrap_or_else(|error| error.into_inner()); + let still_active = active_prompt.load(Ordering::Acquire) == owner; + let was_cancelled = cancelled_prompt.load(Ordering::Acquire) == owner; + if !still_active || was_cancelled { + clear_prompt_tracking(&active_prompt, &driving_prompt, &cancelled_prompt, owner); + false + } else { + driving_prompt.store(owner, Ordering::Release); + true } }; - let (start_tx, start_rx) = oneshot::channel(); - if response.send(Ok(start_tx)).is_err() || start_rx.await.is_err() { - integration.finish_prompt(&session_id); - busy.store(false, Ordering::Release); + if !should_start { continue; } - if client - .update( - session_id.clone(), - wire::SessionUpdate::UserMessage( - wire::UserMessage::new(user_message_id).content(request.prompt), - ), - ) - .and_then(|()| { - client.update( - session_id.clone(), - wire::SessionUpdate::StateUpdate(wire::StateUpdate::Running( - wire::RunningStateUpdate::new(), - )), - ) - }) - .is_err() - { - integration.finish_prompt(&session_id); - busy.store(false, Ordering::Release); + if let Err(error) = driver.submit_input(items) { + tracing::debug!(%error, owner, "failed to submit accepted ACP v2 prompt"); + fail_accepted_prompt( + &client, + &integration, + &session_id, + &active_prompt, + &driving_prompt, + &cancelled_prompt, + &lifecycle, + owner, + false, + ); + continue; + } + let user_message_id = + match integration.begin_prompt(&session_id, owner, prompt_cancellation) { + Ok(message_id) => message_id, + Err(error) => { + tracing::debug!(%error, owner, "failed to begin accepted ACP v2 prompt"); + fail_accepted_prompt( + &client, + &integration, + &session_id, + &active_prompt, + &driving_prompt, + &cancelled_prompt, + &lifecycle, + owner, + false, + ); + continue; + } + }; + + if let Err(error) = client.update( + session_id.clone(), + wire::SessionUpdate::UserMessage( + wire::UserMessage::new(user_message_id).content(request.prompt), + ), + ) { + tracing::debug!(%error, owner, "failed to publish accepted ACP v2 prompt"); + fail_accepted_prompt( + &client, + &integration, + &session_id, + &active_prompt, + &driving_prompt, + &cancelled_prompt, + &lifecycle, + owner, + true, + ); continue; } - let stop_reason = drive_prompt(&mut driver, &cancellation, cancellation_generation).await; + drive_prompt( + &mut driver, + &integration, + &session_id, + &cancellation, + cancellation_generation, + ) + .await; + { + let _lifecycle = lifecycle.lock().unwrap_or_else(|error| error.into_inner()); + integration.finish_prompt(&session_id, owner); + clear_prompt_tracking(&active_prompt, &driving_prompt, &cancelled_prompt, owner); + } if let Err(error) = client.flush().await { tracing::debug!(%error, "failed to flush ACP v2 output"); } - integration.finish_prompt(&session_id); - busy.store(false, Ordering::Release); - let _ = client.update( - session_id.clone(), - wire::SessionUpdate::StateUpdate(wire::StateUpdate::Idle( - wire::IdleStateUpdate::new().stop_reason(stop_reason), - )), - ); + prefer_loop_update = true; } } async fn drive_prompt( driver: &mut agentkit_loop::LoopDriver, + integration: &AcpIntegration, + session_id: &wire::SessionId, cancellation: &CancellationHandle, generation: u64, -) -> wire::StopReason -where +) where S: ModelSession + Send + 'static, { loop { @@ -1026,37 +1396,21 @@ where if result.finish_reason == FinishReason::ToolCall { continue; } - return if cancellation.is_cancelled_since(generation) { - wire::StopReason::Cancelled - } else { - finish_reason_to_stop_reason(&result.finish_reason) - }; - } - Ok(LoopStep::Interrupt(LoopInterrupt::AwaitingInput(_))) => { - return if cancellation.is_cancelled_since(generation) { - wire::StopReason::Cancelled - } else { - wire::StopReason::EndTurn - }; + return; } + Ok(LoopStep::Interrupt(LoopInterrupt::AwaitingInput(_))) => return, Ok(LoopStep::Interrupt(LoopInterrupt::AfterToolResult(_))) => continue, Ok(LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(_))) => { + integration.mark_unsupported_approval(session_id, cancellation.clone(), generation); if let Err(error) = driver.cancel_pending_approvals().await { + integration.clear_unsupported_approval(session_id); tracing::debug!(%error, "failed to cancel unsupported ACP v2 approval"); } - return if cancellation.is_cancelled_since(generation) { - wire::StopReason::Cancelled - } else { - error_stop_reason() - }; + return; } Err(error) => { tracing::debug!(%error, "ACP v2 agent loop failed"); - return if cancellation.is_cancelled_since(generation) { - wire::StopReason::Cancelled - } else { - error_stop_reason() - }; + return; } } } @@ -1417,19 +1771,25 @@ fn headless_capabilities() -> wire::AgentCapabilities { #[cfg(test)] mod tests { use super::*; + use std::collections::VecDeque; use std::sync::atomic::AtomicUsize; use std::time::Duration; use agent_client_protocol::Channel; - use agentkit_core::{ItemKind, ToolCallId, ToolOutput, ToolResultPart, TurnCancellation}; + use agentkit_core::{ + ItemKind, ToolCallId, ToolOutput, ToolResultPart, TurnCancellation, TurnId, + }; use agentkit_integration_tests::mock_model::{MockAdapter, TurnScript}; + use agentkit_integration_tests::mock_tool::BlockingTool; use agentkit_loop::{ - Agent, LoopError, ModelSession, ModelTurn, ModelTurnEvent, ModelTurnResult, SessionConfig, - TurnRequest, + Agent, ModelSession, ModelTurn, ModelTurnEvent, ModelTurnResult, SessionConfig, + TurnRequest, TurnResult, }; + use agentkit_task_manager::{AsyncTaskManager, RoutingDecision}; use agentkit_tools_core::{ Tool, ToolContext, ToolError, ToolRegistry, ToolRequest, ToolResult, ToolSpec, }; + use tokio::sync::Notify; #[derive(Clone)] struct TestFactory { @@ -1520,6 +1880,12 @@ mod tests { tool: CancellationAwareTool, } + #[derive(Clone)] + struct BackgroundToolTestFactory { + adapter: A, + tool: BlockingTool, + } + #[async_trait] impl AcpAgentFactory for ToolTestFactory { async fn start( @@ -1542,6 +1908,150 @@ mod tests { } } + #[async_trait] + impl AcpAgentFactory for BackgroundToolTestFactory + where + A: ModelAdapter + Clone + Send + Sync + 'static, + A::Session: Send + 'static, + { + async fn start( + &self, + ctx: AcpAgentFactoryContext, + ) -> Result, AcpRuntimeError> { + Agent::builder() + .model(self.adapter.clone()) + .add_tool_source(ToolRegistry::new().with(self.tool.clone())) + .task_manager( + AsyncTaskManager::new().routing(|_: &ToolRequest| RoutingDecision::Background), + ) + .observer(ctx.integration.as_ref().clone()) + .cancellation(ctx.cancellation) + .build() + .map_err(|error| AcpRuntimeError::Loop(error.to_string()))? + .start(SessionConfig::new(ctx.agentkit_session_id).with_metadata(ctx.metadata)) + .await + .map_err(|error| AcpRuntimeError::Loop(error.to_string())) + } + } + + struct BlockingInferenceState { + scripts: Mutex>, + next_turn: AtomicUsize, + blocked_turn: usize, + entered: AtomicBool, + release: Notify, + cancelled: AtomicBool, + } + + #[derive(Clone)] + struct BlockingInferenceAdapter { + state: Arc, + } + + impl BlockingInferenceAdapter { + fn new(scripts: impl IntoIterator, blocked_turn: usize) -> Self { + Self { + state: Arc::new(BlockingInferenceState { + scripts: Mutex::new(scripts.into_iter().collect()), + next_turn: AtomicUsize::new(0), + blocked_turn, + entered: AtomicBool::new(false), + release: Notify::new(), + cancelled: AtomicBool::new(false), + }), + } + } + + async fn wait_until_blocked(&self) { + tokio::time::timeout(Duration::from_secs(2), async { + while !self.state.entered.load(Ordering::Acquire) { + tokio::task::yield_now().await; + } + }) + .await + .expect("model inference did not block"); + } + + fn release(&self) { + self.state.release.notify_one(); + } + + fn was_cancelled(&self) -> bool { + self.state.cancelled.load(Ordering::Acquire) + } + } + + struct BlockingInferenceSession { + state: Arc, + } + + struct BlockingInferenceTurn { + state: Arc, + events: VecDeque, + blocked: bool, + } + + #[async_trait] + impl ModelAdapter for BlockingInferenceAdapter { + type Session = BlockingInferenceSession; + + async fn start_session(&self, _config: SessionConfig) -> Result { + Ok(BlockingInferenceSession { + state: Arc::clone(&self.state), + }) + } + } + + #[async_trait] + impl ModelSession for BlockingInferenceSession { + type Turn = BlockingInferenceTurn; + + async fn begin_turn( + &mut self, + _request: TurnRequest, + _cancellation: Option, + ) -> Result { + let turn = self.state.next_turn.fetch_add(1, Ordering::AcqRel); + let script = self + .state + .scripts + .lock() + .unwrap() + .pop_front() + .ok_or_else(|| LoopError::InvalidState("missing blocking turn script".into()))?; + Ok(BlockingInferenceTurn { + state: Arc::clone(&self.state), + events: script.events.into(), + blocked: turn == self.state.blocked_turn, + }) + } + } + + #[async_trait] + impl ModelTurn for BlockingInferenceTurn { + async fn next_event( + &mut self, + cancellation: Option, + ) -> Result, LoopError> { + if self.blocked { + self.blocked = false; + self.state.entered.store(true, Ordering::Release); + if let Some(cancellation) = cancellation { + tokio::select! { + _ = self.state.release.notified() => {} + _ = cancellation.cancelled() => { + self.state.cancelled.store(true, Ordering::Release); + return Err(LoopError::Cancelled); + } + } + } else { + self.state.release.notified().await; + } + } + Ok(self.events.pop_front()) + } + } + fn tool_turn(call_id: &str) -> TurnScript { let call = ToolCallPart::new(ToolCallId::new(call_id), "blocking_tool", json!({})); TurnScript::new([ @@ -1779,6 +2289,28 @@ mod tests { session_updates[1], wire::SessionUpdate::StateUpdate(wire::StateUpdate::Running(_)) )); + assert_eq!( + session_updates + .iter() + .filter(|update| matches!( + update, + wire::SessionUpdate::StateUpdate(wire::StateUpdate::Running(_)) + )) + .count(), + 1, + "Running must come only from TurnStarted" + ); + assert_eq!( + session_updates + .iter() + .filter(|update| matches!( + update, + wire::SessionUpdate::StateUpdate(wire::StateUpdate::Idle(_)) + )) + .count(), + 1, + "Idle must come only from TurnFinished" + ); let user_id = match session_updates[0] { wire::SessionUpdate::UserMessage(message) => message.message_id.to_string(), _ => unreachable!(), @@ -1819,29 +2351,449 @@ mod tests { assert_ne!(agent_ids.last(), thought_ids.last()); assert!(matches!( session_updates.last(), - Some(wire::SessionUpdate::StateUpdate(wire::StateUpdate::Idle(_))) + Some(wire::SessionUpdate::StateUpdate(wire::StateUpdate::Idle(idle))) + if idle.stop_reason == Some(wire::StopReason::EndTurn) )); } } - #[test] - fn v2_tool_updates_include_visible_text_structured_parts_and_files() { - let outputs = [ - ToolOutput::text("plain text"), - ToolOutput::structured(json!({ "ok": true })), - ToolOutput::parts(vec![ - Part::text("part text"), - Part::structured(json!({ "part": true })), - ]), - ToolOutput::files(vec![ - FilePart::named("artifact.txt", DataRef::inline_text("artifact body")) - .with_mime_type("text/plain"), - FilePart::named("remote.txt", DataRef::uri("file:///tmp/remote.txt")), - ]), - ]; - let contents = outputs - .into_iter() - .map(|output| { + #[tokio::test] + async fn idle_session_drives_background_completion_without_another_prompt() { + let adapter = MockAdapter::new(); + adapter.enqueue(tool_turn("background-call")); + adapter.enqueue(streamed_text("background observed")); + let tool = BlockingTool::text("blocking_tool", "background done"); + let updates = Arc::new(Mutex::new(Vec::new())); + let (client_transport, agent_transport) = Channel::duplex(); + + let server = tokio::spawn({ + let factory = BackgroundToolTestFactory { + adapter, + tool: tool.clone(), + }; + async move { + AcpHeadlessRuntime::::builder() + .agent_factory(factory) + .serve(agent_transport) + .await + } + }); + + let client = agent_client_protocol::Client + .v2() + .on_receive_notification( + { + let updates = Arc::clone(&updates); + async move |notification: wire::UpdateSessionNotification, _cx| { + updates + .lock() + .unwrap() + .push((notification.session_id, notification.update)); + Ok(()) + } + }, + agent_client_protocol::on_receive_notification!(), + ) + .connect_with(client_transport, { + let updates = Arc::clone(&updates); + async move |cx| { + cx.send_request(wire::InitializeRequest::new( + wire::ProtocolVersion::V2, + wire::Implementation::new("test-client", "1"), + )) + .block_task() + .await?; + let cwd = std::env::current_dir() + .map_err(agent_client_protocol::Error::into_internal_error)?; + let session = cx + .send_request(wire::NewSessionRequest::new(cwd)) + .block_task() + .await?; + cx.send_request(wire::PromptRequest::new( + session.session_id.clone(), + vec![wire::ContentBlock::Text(wire::TextContent::new( + "start background work", + ))], + )) + .block_task() + .await?; + wait_for_idle(&updates, &session.session_id).await; + tool.wait_until_entered().await; + tool.release(); + + tokio::time::timeout(Duration::from_secs(2), async { + loop { + let updates = updates.lock().unwrap(); + let completed = updates.iter().any(|(id, update)| { + id == &session.session_id + && matches!( + update, + wire::SessionUpdate::ToolCallUpdate(update) + if update.tool_call_id.to_string() == "background-call" + && matches!( + update.status, + agent_client_protocol::schema::MaybeUndefined::Value( + wire::ToolCallStatus::Completed + ) + ) + ) + }); + let message_delivered = updates.iter().any(|(id, update)| { + id == &session.session_id + && matches!( + update, + wire::SessionUpdate::AgentMessageChunk(chunk) + if matches!( + &chunk.content, + wire::ContentBlock::Text(text) + if text.text == "background observed" + ) + ) + }); + if completed && message_delivered { + break; + } + drop(updates); + tokio::time::sleep(Duration::from_millis(5)).await; + } + }) + .await + .expect("background completion was not delivered while idle"); + cx.send_request(wire::CloseSessionRequest::new(session.session_id)) + .block_task() + .await?; + Ok(()) + } + }); + + tokio::time::timeout(Duration::from_secs(5), client) + .await + .expect("client timed out") + .expect("client run"); + server.abort(); + let _ = server.await; + + let updates = updates.lock().unwrap(); + assert_eq!( + updates + .iter() + .filter(|(_, update)| matches!(update, wire::SessionUpdate::UserMessage(_))) + .count(), + 1, + "the background continuation must not require or replay a prompt" + ); + assert!(updates.iter().any(|(_, update)| matches!( + update, + wire::SessionUpdate::AgentMessageChunk(chunk) + if matches!( + &chunk.content, + wire::ContentBlock::Text(text) if text.text == "background observed" + ) + ))); + assert_eq!( + updates + .iter() + .filter(|(_, update)| matches!( + update, + wire::SessionUpdate::StateUpdate(wire::StateUpdate::Running(_)) + )) + .count(), + 1, + "idle background completion must not re-enter Running" + ); + assert_eq!( + updates + .iter() + .filter(|(_, update)| matches!( + update, + wire::SessionUpdate::StateUpdate(wire::StateUpdate::Idle(_)) + )) + .count(), + 1, + "idle background completion must not emit another Idle" + ); + } + + #[tokio::test] + async fn cancel_without_prompt_owner_interrupts_blocked_autonomous_inference() { + let adapter = BlockingInferenceAdapter::new( + [ + tool_turn("background-call"), + streamed_text("autonomous complete"), + ], + 1, + ); + let tool = BlockingTool::text("blocking_tool", "background done"); + let updates = Arc::new(Mutex::new(Vec::new())); + let (client_transport, agent_transport) = Channel::duplex(); + + let server = tokio::spawn({ + let factory = BackgroundToolTestFactory { + adapter: adapter.clone(), + tool: tool.clone(), + }; + async move { + AcpHeadlessRuntime::::builder() + .agent_factory(factory) + .serve(agent_transport) + .await + } + }); + + let client = agent_client_protocol::Client + .v2() + .on_receive_notification( + { + let updates = Arc::clone(&updates); + async move |notification: wire::UpdateSessionNotification, _cx| { + updates + .lock() + .unwrap() + .push((notification.session_id, notification.update)); + Ok(()) + } + }, + agent_client_protocol::on_receive_notification!(), + ) + .connect_with(client_transport, { + let adapter = adapter.clone(); + let tool = tool.clone(); + let updates = Arc::clone(&updates); + async move |cx| { + cx.send_request(wire::InitializeRequest::new( + wire::ProtocolVersion::V2, + wire::Implementation::new("test-client", "1"), + )) + .block_task() + .await?; + let cwd = std::env::current_dir() + .map_err(agent_client_protocol::Error::into_internal_error)?; + let session = cx + .send_request(wire::NewSessionRequest::new(cwd)) + .block_task() + .await?; + cx.send_request(wire::PromptRequest::new( + session.session_id.clone(), + vec![wire::ContentBlock::Text(wire::TextContent::new("start"))], + )) + .block_task() + .await?; + wait_for_idle(&updates, &session.session_id).await; + tool.wait_until_entered().await; + tool.release(); + adapter.wait_until_blocked().await; + + cx.send_notification(wire::CancelSessionNotification::new( + session.session_id.clone(), + ))?; + cx.send_request(wire::ListSessionsRequest::new()) + .block_task() + .await?; + tokio::time::timeout(Duration::from_secs(2), async { + while !adapter.was_cancelled() { + tokio::task::yield_now().await; + } + }) + .await + .expect("session cancel did not interrupt autonomous inference"); + + cx.send_request(wire::CloseSessionRequest::new(session.session_id)) + .block_task() + .await?; + Ok(()) + } + }); + + tokio::time::timeout(Duration::from_secs(5), client) + .await + .expect("client timed out") + .expect("client run"); + server.abort(); + let _ = server.await; + } + + #[tokio::test] + async fn queued_prompt_is_acknowledged_without_cancelling_blocked_autonomous_inference() { + let adapter = BlockingInferenceAdapter::new( + [ + tool_turn("background-call"), + streamed_text("autonomous complete"), + streamed_text("queued prompt complete"), + ], + 1, + ); + let tool = BlockingTool::text("blocking_tool", "background done"); + let updates = Arc::new(Mutex::new(Vec::new())); + let (client_transport, agent_transport) = Channel::duplex(); + + let server = tokio::spawn({ + let factory = BackgroundToolTestFactory { + adapter: adapter.clone(), + tool: tool.clone(), + }; + async move { + AcpHeadlessRuntime::::builder() + .agent_factory(factory) + .serve(agent_transport) + .await + } + }); + + let client = agent_client_protocol::Client + .v2() + .on_receive_notification( + { + let updates = Arc::clone(&updates); + async move |notification: wire::UpdateSessionNotification, _cx| { + updates + .lock() + .unwrap() + .push((notification.session_id, notification.update)); + Ok(()) + } + }, + agent_client_protocol::on_receive_notification!(), + ) + .connect_with(client_transport, { + let adapter = adapter.clone(); + let tool = tool.clone(); + let updates = Arc::clone(&updates); + async move |cx| { + cx.send_request(wire::InitializeRequest::new( + wire::ProtocolVersion::V2, + wire::Implementation::new("test-client", "1"), + )) + .block_task() + .await?; + let cwd = std::env::current_dir() + .map_err(agent_client_protocol::Error::into_internal_error)?; + let session = cx + .send_request(wire::NewSessionRequest::new(cwd)) + .block_task() + .await?; + cx.send_request(wire::PromptRequest::new( + session.session_id.clone(), + vec![wire::ContentBlock::Text(wire::TextContent::new("start"))], + )) + .block_task() + .await?; + wait_for_idle(&updates, &session.session_id).await; + tool.wait_until_entered().await; + tool.release(); + adapter.wait_until_blocked().await; + + tokio::time::timeout( + Duration::from_millis(500), + cx.send_request(wire::PromptRequest::new( + session.session_id.clone(), + vec![wire::ContentBlock::Text(wire::TextContent::new( + "queued while idle", + ))], + )) + .block_task(), + ) + .await + .expect("prompt response waited for autonomous inference")?; + assert_eq!( + updates + .lock() + .unwrap() + .iter() + .filter(|(_, update)| matches!( + update, + wire::SessionUpdate::UserMessage(_) + )) + .count(), + 1, + "queued prompt update crossed the response gate" + ); + + cx.send_notification(wire::CancelSessionNotification::new( + session.session_id.clone(), + ))?; + cx.send_request(wire::ListSessionsRequest::new()) + .block_task() + .await?; + assert!( + tokio::time::timeout(Duration::from_millis(50), async { + while !adapter.was_cancelled() { + tokio::task::yield_now().await; + } + }) + .await + .is_err(), + "queued prompt cancellation interrupted autonomous inference" + ); + + adapter.release(); + tokio::time::timeout(Duration::from_secs(2), async { + loop { + if updates.lock().unwrap().iter().any(|(_, update)| { + matches!( + update, + wire::SessionUpdate::AgentMessageChunk(chunk) + if matches!( + &chunk.content, + wire::ContentBlock::Text(text) + if text.text == "autonomous complete" + ) + ) + }) { + break; + } + tokio::time::sleep(Duration::from_millis(5)).await; + } + }) + .await + .expect("autonomous inference did not complete"); + assert!(!adapter.was_cancelled()); + assert_eq!( + updates + .lock() + .unwrap() + .iter() + .filter(|(_, update)| matches!( + update, + wire::SessionUpdate::UserMessage(_) + )) + .count(), + 1, + "cancelled queued prompt was driven" + ); + + cx.send_request(wire::CloseSessionRequest::new(session.session_id)) + .block_task() + .await?; + Ok(()) + } + }); + + tokio::time::timeout(Duration::from_secs(5), client) + .await + .expect("client timed out") + .expect("client run"); + server.abort(); + let _ = server.await; + } + + #[test] + fn v2_tool_updates_include_visible_text_structured_parts_and_files() { + let outputs = [ + ToolOutput::text("plain text"), + ToolOutput::structured(json!({ "ok": true })), + ToolOutput::parts(vec![ + Part::text("part text"), + Part::structured(json!({ "part": true })), + ]), + ToolOutput::files(vec![ + FilePart::named("artifact.txt", DataRef::inline_text("artifact body")) + .with_mime_type("text/plain"), + FilePart::named("remote.txt", DataRef::uri("file:///tmp/remote.txt")), + ]), + ]; + let contents = outputs + .into_iter() + .map(|output| { let result = ToolResultPart::success(ToolCallId::new("call"), output); let update = tool_result_update(&result, wire::ToolCallStatus::Completed); serde_json::to_value(update).expect("serialize tool update")["content"].clone() @@ -1963,7 +2915,17 @@ mod tests { integration .bind(acp_id.clone(), agentkit_id.clone(), client) .expect("bind session"); - integration.begin_prompt(&acp_id).expect("begin prompt"); + integration + .install_prompt_state( + &acp_id, + Arc::new(AtomicU64::new(1)), + Arc::new(Mutex::new(())), + ) + .expect("install prompt state"); + let cancellation = CancellationController::new(); + integration + .begin_prompt(&acp_id, 1, cancellation.handle().checkpoint()) + .expect("begin prompt"); for event in [ AgentEvent::ContentDelta(Delta::BeginPart { @@ -2041,6 +3003,365 @@ mod tests { } } + #[test] + fn unsupported_approval_keeps_error_stop_reason_on_cancelled_turn() { + let integration = AcpIntegration::default(); + let (client, mut messages) = ClientHandle::channel(); + let acp_id = wire::SessionId::new("acp-session"); + let agentkit_id = AgentkitSessionId::new("agentkit-session"); + let cancellation = CancellationController::new(); + let generation = cancellation.handle().generation(); + integration + .bind(acp_id.clone(), agentkit_id.clone(), client) + .expect("bind session"); + integration + .install_prompt_state( + &acp_id, + Arc::new(AtomicU64::new(1)), + Arc::new(Mutex::new(())), + ) + .expect("install prompt state"); + integration + .begin_prompt(&acp_id, 1, cancellation.handle().checkpoint()) + .expect("begin prompt"); + + integration.route_event( + &agentkit_id, + AgentEvent::TurnStarted { + session_id: agentkit_id.clone(), + turn_id: TurnId::new("turn-1"), + }, + ); + integration.mark_unsupported_approval(&acp_id, cancellation.handle(), generation); + integration.route_event( + &agentkit_id, + AgentEvent::TurnFinished(TurnResult { + turn_id: TurnId::new("turn-1"), + finish_reason: FinishReason::Cancelled, + items: Vec::new(), + usage: None, + metadata: MetadataMap::new(), + }), + ); + + assert!(matches!( + messages.try_recv(), + Ok(ClientMessage::Update(notification)) + if matches!( + ¬ification.update, + wire::SessionUpdate::StateUpdate(wire::StateUpdate::Running(_)) + ) + )); + assert!(matches!( + messages.try_recv(), + Ok(ClientMessage::Update(notification)) + if matches!( + ¬ification.update, + wire::SessionUpdate::StateUpdate(wire::StateUpdate::Idle(idle)) + if idle.stop_reason.as_ref() + == Some(&wire::StopReason::Other("_error".into())) + ) + )); + assert!(messages.try_recv().is_err()); + } + + #[test] + fn prompt_cancellation_overrides_terminal_and_unsupported_stop_reasons() { + for (index, finish_reason, unsupported_approval) in [ + (1, FinishReason::Completed, false), + (2, FinishReason::Error, false), + (3, FinishReason::Completed, true), + ] { + let integration = AcpIntegration::default(); + let (client, mut messages) = ClientHandle::channel(); + let acp_id = wire::SessionId::new(format!("acp-session-{index}")); + let agentkit_id = AgentkitSessionId::new(format!("agentkit-session-{index}")); + let active_prompt = Arc::new(AtomicU64::new(index)); + let cancellation = CancellationController::new(); + let generation = cancellation.handle().generation(); + integration + .bind(acp_id.clone(), agentkit_id.clone(), client) + .expect("bind session"); + integration + .install_prompt_state(&acp_id, active_prompt, Arc::new(Mutex::new(()))) + .expect("install prompt state"); + integration + .begin_prompt(&acp_id, index, cancellation.handle().checkpoint()) + .expect("begin prompt"); + let turn_id = TurnId::new(format!("turn-{index}")); + integration.route_event( + &agentkit_id, + AgentEvent::TurnStarted { + session_id: agentkit_id.clone(), + turn_id: turn_id.clone(), + }, + ); + if unsupported_approval { + integration.mark_unsupported_approval(&acp_id, cancellation.handle(), generation); + } + cancellation.interrupt(); + integration.route_event( + &agentkit_id, + AgentEvent::TurnFinished(agentkit_loop::TurnResult { + turn_id, + finish_reason, + items: Vec::new(), + usage: None, + metadata: MetadataMap::new(), + }), + ); + + assert!(matches!( + messages.try_recv(), + Ok(ClientMessage::Update(notification)) + if matches!( + notification.update, + wire::SessionUpdate::StateUpdate(wire::StateUpdate::Running(_)) + ) + )); + assert!(matches!( + messages.try_recv(), + Ok(ClientMessage::Update(notification)) + if matches!( + ¬ification.update, + wire::SessionUpdate::StateUpdate(wire::StateUpdate::Idle(idle)) + if idle.stop_reason == Some(wire::StopReason::Cancelled) + ) + )); + } + } + + #[test] + fn accepted_prompt_failure_queues_error_idle_after_cleanup() { + let integration = AcpIntegration::default(); + let owner = 7; + let active_prompt = Arc::new(AtomicU64::new(owner)); + let driving_prompt = Arc::new(AtomicU64::new(owner)); + let cancelled_prompt = Arc::new(AtomicU64::new(owner)); + let lifecycle = Arc::new(Mutex::new(())); + let prompt_at_enqueue = Arc::clone(&active_prompt); + let lifecycle_at_enqueue = Arc::clone(&lifecycle); + let (client, mut messages) = ClientHandle::channel_with_update_hook(move |update| { + if matches!( + update, + wire::SessionUpdate::StateUpdate(wire::StateUpdate::Idle(_)) + ) { + assert_eq!(prompt_at_enqueue.load(Ordering::Acquire), 0); + assert!(lifecycle_at_enqueue.try_lock().is_err()); + } + }); + let acp_id = wire::SessionId::new("acp-failed-prompt"); + integration + .bind( + acp_id.clone(), + AgentkitSessionId::new("agentkit-failed-prompt"), + client.clone(), + ) + .expect("bind session"); + integration + .install_prompt_state(&acp_id, Arc::clone(&active_prompt), Arc::clone(&lifecycle)) + .expect("install prompt state"); + let cancellation = CancellationController::new(); + integration + .begin_prompt(&acp_id, owner, cancellation.handle().checkpoint()) + .expect("begin prompt"); + + fail_accepted_prompt( + &client, + &integration, + &acp_id, + &active_prompt, + &driving_prompt, + &cancelled_prompt, + &lifecycle, + owner, + true, + ); + + assert_eq!(active_prompt.load(Ordering::Acquire), 0); + assert_eq!(driving_prompt.load(Ordering::Acquire), 0); + assert_eq!(cancelled_prompt.load(Ordering::Acquire), 0); + let session = integration.session(&acp_id).expect("bound session"); + assert!( + session + .prompt_state + .lock() + .unwrap() + .as_ref() + .is_some_and(|state| state.pending_owner.is_none()) + ); + assert!(matches!( + messages.try_recv(), + Ok(ClientMessage::Update(notification)) + if matches!( + ¬ification.update, + wire::SessionUpdate::StateUpdate(wire::StateUpdate::Idle(idle)) + if idle.stop_reason == Some(error_stop_reason()) + ) + )); + } + + #[test] + fn idle_is_queued_after_releasing_its_prompt_owner() { + let integration = AcpIntegration::default(); + let active_prompt = Arc::new(AtomicU64::new(7)); + let lifecycle = Arc::new(Mutex::new(())); + let prompt_at_enqueue = Arc::clone(&active_prompt); + let lifecycle_at_enqueue = Arc::clone(&lifecycle); + let (client, mut messages) = ClientHandle::channel_with_update_hook(move |update| { + if matches!( + update, + wire::SessionUpdate::StateUpdate(wire::StateUpdate::Idle(_)) + ) { + assert_eq!( + prompt_at_enqueue.load(Ordering::Acquire), + 0, + "Idle was enqueued before prompt ownership was released" + ); + assert!( + lifecycle_at_enqueue.try_lock().is_err(), + "Idle was enqueued outside the prompt lifecycle lock" + ); + } + }); + let acp_id = wire::SessionId::new("acp-session"); + let agentkit_id = AgentkitSessionId::new("agentkit-session"); + integration + .bind(acp_id.clone(), agentkit_id.clone(), client) + .expect("bind session"); + integration + .install_prompt_state(&acp_id, Arc::clone(&active_prompt), Arc::clone(&lifecycle)) + .expect("install prompt state"); + let cancellation = CancellationController::new(); + integration + .begin_prompt(&acp_id, 7, cancellation.handle().checkpoint()) + .expect("begin prompt"); + let turn_id = TurnId::new("turn-1"); + integration.route_event( + &agentkit_id, + AgentEvent::TurnStarted { + session_id: agentkit_id.clone(), + turn_id: turn_id.clone(), + }, + ); + integration.route_event( + &agentkit_id, + AgentEvent::TurnFinished(agentkit_loop::TurnResult { + turn_id, + finish_reason: FinishReason::Completed, + items: Vec::new(), + usage: None, + metadata: MetadataMap::new(), + }), + ); + + assert_eq!(active_prompt.load(Ordering::Acquire), 0); + assert!( + std::iter::from_fn(|| messages.try_recv().ok()).any(|message| { + matches!( + message, + ClientMessage::Update(notification) + if matches!( + notification.update, + wire::SessionUpdate::StateUpdate(wire::StateUpdate::Idle(_)) + ) + ) + }) + ); + } + + #[test] + fn cancellation_while_turn_finish_waits_for_lifecycle_lock_wins() { + let integration = AcpIntegration::default(); + let active_prompt = Arc::new(AtomicU64::new(9)); + let lifecycle = Arc::new(Mutex::new(())); + let prompt_at_enqueue = Arc::clone(&active_prompt); + let lifecycle_at_enqueue = Arc::clone(&lifecycle); + let (client, mut messages) = ClientHandle::channel_with_update_hook(move |update| { + if matches!( + update, + wire::SessionUpdate::StateUpdate(wire::StateUpdate::Idle(_)) + ) { + assert_eq!(prompt_at_enqueue.load(Ordering::Acquire), 0); + assert!(lifecycle_at_enqueue.try_lock().is_err()); + } + }); + let acp_id = wire::SessionId::new("acp-race-session"); + let agentkit_id = AgentkitSessionId::new("agentkit-race-session"); + integration + .bind(acp_id.clone(), agentkit_id.clone(), client) + .expect("bind session"); + integration + .install_prompt_state(&acp_id, Arc::clone(&active_prompt), Arc::clone(&lifecycle)) + .expect("install prompt state"); + let cancellation = CancellationController::new(); + integration + .begin_prompt(&acp_id, 9, cancellation.handle().checkpoint()) + .expect("begin prompt"); + let turn_id = TurnId::new("turn-race"); + integration.route_event( + &agentkit_id, + AgentEvent::TurnStarted { + session_id: agentkit_id.clone(), + turn_id: turn_id.clone(), + }, + ); + + let lifecycle_guard = lifecycle.lock().unwrap(); + let ready = Arc::new(std::sync::Barrier::new(2)); + let finish_ready = Arc::clone(&ready); + let finish_integration = integration.clone(); + let finish_session = agentkit_id.clone(); + let finish = std::thread::spawn(move || { + finish_ready.wait(); + finish_integration.route_event( + &finish_session, + AgentEvent::TurnFinished(agentkit_loop::TurnResult { + turn_id, + finish_reason: FinishReason::Completed, + items: Vec::new(), + usage: None, + metadata: MetadataMap::new(), + }), + ); + }); + ready.wait(); + cancellation.interrupt(); + drop(lifecycle_guard); + finish.join().expect("turn finish thread"); + + assert!(matches!( + messages.try_recv(), + Ok(ClientMessage::Update(notification)) + if matches!( + notification.update, + wire::SessionUpdate::StateUpdate(wire::StateUpdate::Running(_)) + ) + )); + assert!(matches!( + messages.try_recv(), + Ok(ClientMessage::Update(notification)) + if matches!( + ¬ification.update, + wire::SessionUpdate::StateUpdate(wire::StateUpdate::Idle(idle)) + if idle.stop_reason == Some(wire::StopReason::Cancelled) + ) + )); + } + + #[test] + fn stale_prompt_owner_cannot_release_new_prompt() { + let active_prompt = AtomicU64::new(1); + release_prompt(&active_prompt, 1); + assert_eq!(active_prompt.load(Ordering::Acquire), 0); + + active_prompt.store(2, Ordering::Release); + release_prompt(&active_prompt, 1); + assert_eq!(active_prompt.load(Ordering::Acquire), 2); + release_prompt(&active_prompt, 2); + assert_eq!(active_prompt.load(Ordering::Acquire), 0); + } + #[tokio::test] async fn cancel_then_prompt_and_close_cleanup_active_tools() { let adapter = MockAdapter::new(); @@ -2260,7 +3581,7 @@ mod tests { } #[tokio::test] - async fn independent_prompts_are_accepted_immediately_and_cancel_separately() { + async fn same_session_prompt_race_admits_one_and_other_sessions_stay_independent() { let updates = Arc::new(Mutex::new(Vec::new())); let (client_transport, agent_transport) = Channel::duplex(); let server = tokio::spawn(async move { @@ -2308,18 +3629,42 @@ mod tests { .block_task() .await?; - for session_id in [&first.session_id, &second.session_id] { - tokio::time::timeout( - Duration::from_millis(250), - cx.send_request(wire::PromptRequest::new( - session_id.clone(), - vec![wire::ContentBlock::Text(wire::TextContent::new("wait"))], - )) - .block_task(), - ) + let first_prompt = cx + .send_request(wire::PromptRequest::new( + first.session_id.clone(), + vec![wire::ContentBlock::Text(wire::TextContent::new("first"))], + )) + .block_task(); + let competing_prompt = cx + .send_request(wire::PromptRequest::new( + first.session_id.clone(), + vec![wire::ContentBlock::Text(wire::TextContent::new( + "competing", + ))], + )) + .block_task(); + let (first_result, competing_result) = + tokio::time::timeout(Duration::from_millis(250), async { + tokio::join!(first_prompt, competing_prompt) + }) .await - .expect("prompt acceptance was blocked by another session")?; - } + .expect("same-session prompt admission timed out"); + assert_ne!( + first_result.is_ok(), + competing_result.is_ok(), + "exactly one same-session prompt must be admitted" + ); + + tokio::time::timeout( + Duration::from_millis(250), + cx.send_request(wire::PromptRequest::new( + second.session_id.clone(), + vec![wire::ContentBlock::Text(wire::TextContent::new("wait"))], + )) + .block_task(), + ) + .await + .expect("prompt acceptance was blocked by another session")?; cx.send_notification(wire::CancelSessionNotification::new( first.session_id.clone(), diff --git a/crates/agentkit-adapter-completions/Cargo.toml b/crates/agentkit-adapter-completions/Cargo.toml index 1adfdc5..cfd2225 100644 --- a/crates/agentkit-adapter-completions/Cargo.toml +++ b/crates/agentkit-adapter-completions/Cargo.toml @@ -7,7 +7,7 @@ edition.workspace = true license.workspace = true repository.workspace = true rust-version.workspace = true -version = "0.10.6" +version = "0.10.7" [dependencies] agentkit-core = { version = "0.10.5", path = "../agentkit-core" } diff --git a/crates/agentkit-adapter-completions/src/request.rs b/crates/agentkit-adapter-completions/src/request.rs index 73e36e5..ee2b487 100644 --- a/crates/agentkit-adapter-completions/src/request.rs +++ b/crates/agentkit-adapter-completions/src/request.rs @@ -591,7 +591,13 @@ mod tests { let body = build_request_body( &TestProvider::lenient(), &turn_request( - vec![Item::notification("background task done: ok")], + vec![Item::new( + ItemKind::Notification, + vec![ + Part::text("background task done: ok"), + Part::structured(serde_json::json!({ "typed": true })), + ], + )], Vec::new(), ), ) @@ -604,6 +610,7 @@ mod tests { assert!(text.starts_with("")); assert!(text.ends_with("")); assert!(text.contains("background task done: ok")); + assert!(text.contains("\"typed\": true")); } #[test] diff --git a/crates/agentkit-integration-tests/tests/background_results.rs b/crates/agentkit-integration-tests/tests/background_results.rs index 3d617e4..4cc6e60 100644 --- a/crates/agentkit-integration-tests/tests/background_results.rs +++ b/crates/agentkit-integration-tests/tests/background_results.rs @@ -182,7 +182,7 @@ async fn foreground_emits_single_tool_result_no_notification() { // ─── case 2: pure Background (control) ──────────────────────────────────── #[tokio::test] -async fn pure_background_completion_emits_single_tool_result_no_notification() { +async fn pure_background_emits_detached_placeholder_and_completion_notification() { let path = snapshot_path("bg_pure_background.ron"); let recording = SessionRecording::load_or_seed(&path, || SessionRecording { session_id: "pure-bg".into(), diff --git a/crates/agentkit-integration-tests/tests/snapshots/bg_detach_post_turn.ron b/crates/agentkit-integration-tests/tests/snapshots/bg_detach_post_turn.ron index dd8f1c0..0d2e1ea 100644 --- a/crates/agentkit-integration-tests/tests/snapshots/bg_detach_post_turn.ron +++ b/crates/agentkit-integration-tests/tests/snapshots/bg_detach_post_turn.ron @@ -302,7 +302,19 @@ SessionRecording( kind: Notification, parts: [ Text(TextPart( - text: "Background tool call call-d completed: detach-output", + text: "Background tool results: 1 total, 0 failed, 0 with metadata. call-d completed: text preview: detach-output", + metadata: {}, + )), + Structured(StructuredPart( + value: { + "call_id": "call-d", + "is_error": false, + "metadata": {}, + "output": { + "Text": "detach-output", + }, + }, + schema: None, metadata: {}, )), ], @@ -437,7 +449,19 @@ SessionRecording( kind: Notification, parts: [ Text(TextPart( - text: "Background tool call call-d completed: detach-output", + text: "Background tool results: 1 total, 0 failed, 0 with metadata. call-d completed: text preview: detach-output", + metadata: {}, + )), + Structured(StructuredPart( + value: { + "call_id": "call-d", + "is_error": false, + "metadata": {}, + "output": { + "Text": "detach-output", + }, + }, + schema: None, metadata: {}, )), ], diff --git a/crates/agentkit-integration-tests/tests/snapshots/bg_mixed.ron b/crates/agentkit-integration-tests/tests/snapshots/bg_mixed.ron index 7033370..1e671e5 100644 --- a/crates/agentkit-integration-tests/tests/snapshots/bg_mixed.ron +++ b/crates/agentkit-integration-tests/tests/snapshots/bg_mixed.ron @@ -390,7 +390,19 @@ SessionRecording( kind: Notification, parts: [ Text(TextPart( - text: "Background tool call call-bg completed: slow-out", + text: "Background tool results: 1 total, 0 failed, 0 with metadata. call-bg completed: text preview: slow-out", + metadata: {}, + )), + Structured(StructuredPart( + value: { + "call_id": "call-bg", + "is_error": false, + "metadata": {}, + "output": { + "Text": "slow-out", + }, + }, + schema: None, metadata: {}, )), ], @@ -563,7 +575,19 @@ SessionRecording( kind: Notification, parts: [ Text(TextPart( - text: "Background tool call call-bg completed: slow-out", + text: "Background tool results: 1 total, 0 failed, 0 with metadata. call-bg completed: text preview: slow-out", + metadata: {}, + )), + Structured(StructuredPart( + value: { + "call_id": "call-bg", + "is_error": false, + "metadata": {}, + "output": { + "Text": "slow-out", + }, + }, + schema: None, metadata: {}, )), ], diff --git a/crates/agentkit-integration-tests/tests/snapshots/bg_pure_background.ron b/crates/agentkit-integration-tests/tests/snapshots/bg_pure_background.ron index debdbf5..db8c272 100644 --- a/crates/agentkit-integration-tests/tests/snapshots/bg_pure_background.ron +++ b/crates/agentkit-integration-tests/tests/snapshots/bg_pure_background.ron @@ -166,7 +166,7 @@ SessionRecording( parts: [ ToolResult(ToolResultPart( call_id: ToolCallId("call-bg"), - output: Text("bg-output"), + output: Text("Tool bg-tool is now running in the background. The result will be delivered when it completes."), is_error: false, metadata: {}, )), @@ -176,6 +176,32 @@ SessionRecording( finish_reason: None, created_at: None, ), + Item( + id: None, + kind: Notification, + parts: [ + Text(TextPart( + text: "Background tool results: 1 total, 0 failed, 0 with metadata. call-bg completed: text preview: bg-output", + metadata: {}, + )), + Structured(StructuredPart( + value: { + "call_id": "call-bg", + "is_error": false, + "metadata": {}, + "output": { + "Text": "bg-output", + }, + }, + schema: None, + metadata: {}, + )), + ], + metadata: {}, + usage: None, + finish_reason: None, + created_at: None, + ), ], tools: [ ToolSpec( @@ -273,7 +299,7 @@ SessionRecording( parts: [ ToolResult(ToolResultPart( call_id: ToolCallId("call-bg"), - output: Text("bg-output"), + output: Text("Tool bg-tool is now running in the background. The result will be delivered when it completes."), is_error: false, metadata: {}, )), @@ -283,6 +309,32 @@ SessionRecording( finish_reason: None, created_at: None, ), + Item( + id: None, + kind: Notification, + parts: [ + Text(TextPart( + text: "Background tool results: 1 total, 0 failed, 0 with metadata. call-bg completed: text preview: bg-output", + metadata: {}, + )), + Structured(StructuredPart( + value: { + "call_id": "call-bg", + "is_error": false, + "metadata": {}, + "output": { + "Text": "bg-output", + }, + }, + schema: None, + metadata: {}, + )), + ], + metadata: {}, + usage: None, + finish_reason: None, + created_at: None, + ), Item( id: None, kind: Assistant, diff --git a/crates/agentkit-integration-tests/tests/snapshots/bg_two_in_order.ron b/crates/agentkit-integration-tests/tests/snapshots/bg_two_in_order.ron index d22c9a2..702bd62 100644 --- a/crates/agentkit-integration-tests/tests/snapshots/bg_two_in_order.ron +++ b/crates/agentkit-integration-tests/tests/snapshots/bg_two_in_order.ron @@ -390,7 +390,19 @@ SessionRecording( kind: Notification, parts: [ Text(TextPart( - text: "Background tool call call-a completed: out-a", + text: "Background tool results: 1 total, 0 failed, 0 with metadata. call-a completed: text preview: out-a", + metadata: {}, + )), + Structured(StructuredPart( + value: { + "call_id": "call-a", + "is_error": false, + "metadata": {}, + "output": { + "Text": "out-a", + }, + }, + schema: None, metadata: {}, )), ], @@ -563,7 +575,19 @@ SessionRecording( kind: Notification, parts: [ Text(TextPart( - text: "Background tool call call-a completed: out-a", + text: "Background tool results: 1 total, 0 failed, 0 with metadata. call-a completed: text preview: out-a", + metadata: {}, + )), + Structured(StructuredPart( + value: { + "call_id": "call-a", + "is_error": false, + "metadata": {}, + "output": { + "Text": "out-a", + }, + }, + schema: None, metadata: {}, )), ], @@ -591,7 +615,19 @@ SessionRecording( kind: Notification, parts: [ Text(TextPart( - text: "Background tool call call-b completed: out-b", + text: "Background tool results: 1 total, 0 failed, 0 with metadata. call-b completed: text preview: out-b", + metadata: {}, + )), + Structured(StructuredPart( + value: { + "call_id": "call-b", + "is_error": false, + "metadata": {}, + "output": { + "Text": "out-b", + }, + }, + schema: None, metadata: {}, )), ], @@ -764,7 +800,19 @@ SessionRecording( kind: Notification, parts: [ Text(TextPart( - text: "Background tool call call-a completed: out-a", + text: "Background tool results: 1 total, 0 failed, 0 with metadata. call-a completed: text preview: out-a", + metadata: {}, + )), + Structured(StructuredPart( + value: { + "call_id": "call-a", + "is_error": false, + "metadata": {}, + "output": { + "Text": "out-a", + }, + }, + schema: None, metadata: {}, )), ], @@ -792,7 +840,19 @@ SessionRecording( kind: Notification, parts: [ Text(TextPart( - text: "Background tool call call-b completed: out-b", + text: "Background tool results: 1 total, 0 failed, 0 with metadata. call-b completed: text preview: out-b", + metadata: {}, + )), + Structured(StructuredPart( + value: { + "call_id": "call-b", + "is_error": false, + "metadata": {}, + "output": { + "Text": "out-b", + }, + }, + schema: None, metadata: {}, )), ], diff --git a/crates/agentkit-integration-tests/tests/snapshots/bg_two_out_of_order.ron b/crates/agentkit-integration-tests/tests/snapshots/bg_two_out_of_order.ron index 5c9ad56..5b412ca 100644 --- a/crates/agentkit-integration-tests/tests/snapshots/bg_two_out_of_order.ron +++ b/crates/agentkit-integration-tests/tests/snapshots/bg_two_out_of_order.ron @@ -390,7 +390,19 @@ SessionRecording( kind: Notification, parts: [ Text(TextPart( - text: "Background tool call call-b completed: out-b", + text: "Background tool results: 1 total, 0 failed, 0 with metadata. call-b completed: text preview: out-b", + metadata: {}, + )), + Structured(StructuredPart( + value: { + "call_id": "call-b", + "is_error": false, + "metadata": {}, + "output": { + "Text": "out-b", + }, + }, + schema: None, metadata: {}, )), ], @@ -563,7 +575,19 @@ SessionRecording( kind: Notification, parts: [ Text(TextPart( - text: "Background tool call call-b completed: out-b", + text: "Background tool results: 1 total, 0 failed, 0 with metadata. call-b completed: text preview: out-b", + metadata: {}, + )), + Structured(StructuredPart( + value: { + "call_id": "call-b", + "is_error": false, + "metadata": {}, + "output": { + "Text": "out-b", + }, + }, + schema: None, metadata: {}, )), ], @@ -591,7 +615,19 @@ SessionRecording( kind: Notification, parts: [ Text(TextPart( - text: "Background tool call call-a completed: out-a", + text: "Background tool results: 1 total, 0 failed, 0 with metadata. call-a completed: text preview: out-a", + metadata: {}, + )), + Structured(StructuredPart( + value: { + "call_id": "call-a", + "is_error": false, + "metadata": {}, + "output": { + "Text": "out-a", + }, + }, + schema: None, metadata: {}, )), ], @@ -764,7 +800,19 @@ SessionRecording( kind: Notification, parts: [ Text(TextPart( - text: "Background tool call call-b completed: out-b", + text: "Background tool results: 1 total, 0 failed, 0 with metadata. call-b completed: text preview: out-b", + metadata: {}, + )), + Structured(StructuredPart( + value: { + "call_id": "call-b", + "is_error": false, + "metadata": {}, + "output": { + "Text": "out-b", + }, + }, + schema: None, metadata: {}, )), ], @@ -792,7 +840,19 @@ SessionRecording( kind: Notification, parts: [ Text(TextPart( - text: "Background tool call call-a completed: out-a", + text: "Background tool results: 1 total, 0 failed, 0 with metadata. call-a completed: text preview: out-a", + metadata: {}, + )), + Structured(StructuredPart( + value: { + "call_id": "call-a", + "is_error": false, + "metadata": {}, + "output": { + "Text": "out-a", + }, + }, + schema: None, metadata: {}, )), ], diff --git a/crates/agentkit-loop/Cargo.toml b/crates/agentkit-loop/Cargo.toml index ac245c4..e1ced83 100644 --- a/crates/agentkit-loop/Cargo.toml +++ b/crates/agentkit-loop/Cargo.toml @@ -4,7 +4,7 @@ homepage.workspace = true name = "agentkit-loop" readme = "README.md" repository.workspace = true -version = "0.10.7" +version = "0.10.10" edition.workspace = true license.workspace = true rust-version.workspace = true @@ -15,7 +15,7 @@ otel = ["dep:opentelemetry", "dep:tracing-opentelemetry"] [dependencies] agentkit-core = { version = "0.10.5", path = "../agentkit-core" } -agentkit-task-manager = { version = "0.10.5", path = "../agentkit-task-manager" } +agentkit-task-manager = { version = "0.10.7", path = "../agentkit-task-manager" } agentkit-tools-core = { version = "0.10.5", path = "../agentkit-tools-core" } async-trait.workspace = true serde.workspace = true diff --git a/crates/agentkit-loop/src/lib.rs b/crates/agentkit-loop/src/lib.rs index 531d56a..ec9288d 100644 --- a/crates/agentkit-loop/src/lib.rs +++ b/crates/agentkit-loop/src/lib.rs @@ -94,6 +94,9 @@ const INTERRUPTED_METADATA_KEY: &str = "agentkit.interrupted"; const INTERRUPT_REASON_METADATA_KEY: &str = "agentkit.interrupt_reason"; const INTERRUPT_STAGE_METADATA_KEY: &str = "agentkit.interrupt_stage"; const USER_CANCELLED_REASON: &str = "user_cancelled"; +const DETACHED_NOTIFICATION_TEXT_MAX_CHARS: usize = 512; +const DETACHED_TEXT_PREVIEW_MAX_CHARS: usize = 160; +const DETACHED_CALL_ID_MAX_CHARS: usize = 80; /// Metadata key used by adapters to retain provider-native finish reasons. pub const PROVIDER_FINISH_REASONS_METADATA_KEY: &str = "agentkit.provider_finish_reasons"; @@ -874,7 +877,7 @@ pub trait LoopMutator: Send + Sync { pub enum AgentEvent { /// The agent run has been initialised. RunStarted { session_id: SessionId }, - /// A new model turn is starting. + /// A new logical turn is starting. TurnStarted { session_id: SessionId, turn_id: agentkit_core::TurnId, @@ -935,7 +938,7 @@ pub enum AgentEvent { Warning { message: String }, /// The agent run has failed with an unrecoverable error. RunFailed { message: String }, - /// A turn has finished (successfully, via cancellation, etc.). + /// A logical turn has finished (successfully, via cancellation, etc.). TurnFinished(TurnResult), } @@ -1236,7 +1239,7 @@ struct PendingApprovalToolCall { request: ApprovalRequest, decision: Option, surfaced: bool, - turn_id: agentkit_core::TurnId, + presentation_turn_id: agentkit_core::TurnId, task_id: TaskId, call: ToolCallPart, tool_request: ToolRequest, @@ -1245,13 +1248,19 @@ struct PendingApprovalToolCall { #[derive(Clone, Default)] struct ActiveToolRound { - turn_id: agentkit_core::TurnId, + presentation_turn_id: agentkit_core::TurnId, + task_turn_id: agentkit_core::TurnId, pending_calls: VecDeque<(ToolCallPart, ToolRequest)>, cancellation: Option, background_pending: bool, foreground_progressed: bool, } +#[derive(Default)] +struct DriverLifecycle { + active_turn: Option, +} + /// A configured agent ready to start a session. /// /// Build one with [`Agent::builder`], supplying at minimum a [`ModelAdapter`]. @@ -1367,7 +1376,9 @@ where pending_approval_order: VecDeque::new(), active_tool_round: None, pending_round_resume: None, + pending_loop_updates: VecDeque::new(), next_turn_index: 1, + lifecycle: DriverLifecycle::default(), background_call_ids: HashSet::new(), detached_call_ids: HashSet::new(), interrupted_background_call_ids: HashSet::new(), @@ -1640,7 +1651,9 @@ where pending_approval_order: VecDeque, active_tool_round: Option, pending_round_resume: Option, + pending_loop_updates: VecDeque, next_turn_index: u64, + lifecycle: DriverLifecycle, /// Calls currently running in the background without a transcript result. background_call_ids: HashSet, /// Call ids whose original tool_use was already paired with a @@ -1762,6 +1775,37 @@ where !self.pending_approvals.is_empty() } + fn start_logical_turn(&mut self) -> agentkit_core::TurnId { + if let Some(turn_id) = &self.lifecycle.active_turn { + return turn_id.clone(); + } + let turn_id = agentkit_core::TurnId::new(format!("turn-{}", self.next_turn_index)); + self.next_turn_index += 1; + self.start_logical_turn_with(turn_id) + } + + fn start_logical_turn_with(&mut self, turn_id: agentkit_core::TurnId) -> agentkit_core::TurnId { + if let Some(active_turn) = &self.lifecycle.active_turn { + return active_turn.clone(); + } + self.lifecycle.active_turn = Some(turn_id.clone()); + self.emit(AgentEvent::TurnStarted { + session_id: self.session_id.clone(), + turn_id: turn_id.clone(), + }); + turn_id + } + + fn finish_logical_turn(&mut self, result: &TurnResult) { + if self.pending_round_resume.as_ref() == Some(&result.turn_id) { + self.pending_round_resume = None; + } + if self.lifecycle.active_turn.as_ref() == Some(&result.turn_id) { + self.lifecycle.active_turn = None; + self.emit(AgentEvent::TurnFinished(result.clone())); + } + } + fn emit_tool_catalog_events(&mut self, events: Vec) { for event in events { self.emit(AgentEvent::ToolCatalogChanged(event)); @@ -1770,7 +1814,7 @@ where fn enqueue_pending_approval( &mut self, - turn_id: &agentkit_core::TurnId, + presentation_turn_id: &agentkit_core::TurnId, task: TaskApproval, cancellation: Option, ) { @@ -1789,7 +1833,7 @@ where request: request.clone(), decision: None, surfaced: false, - turn_id: turn_id.clone(), + presentation_turn_id: presentation_turn_id.clone(), task_id: task.task_id, call, tool_request: task.tool_request, @@ -1843,7 +1887,7 @@ where fn queue_resolution_interrupt( &mut self, - turn_id: &agentkit_core::TurnId, + presentation_turn_id: &agentkit_core::TurnId, resolution: TaskResolution, cancellation: Option, ) -> Option { @@ -1853,18 +1897,28 @@ where None } TaskResolution::Approval(task) => { - self.enqueue_pending_approval(turn_id, task, cancellation); + self.enqueue_pending_approval(presentation_turn_id, task, cancellation); self.take_next_unsurfaced_approval_interrupt() } } } - async fn drain_pending_loop_updates(&mut self) -> Result<(bool, Option), LoopError> { - let PendingLoopUpdates { mut resolutions } = self + async fn collect_pending_loop_updates(&mut self) -> Result<(), LoopError> { + let PendingLoopUpdates { resolutions } = self .task_manager .take_pending_loop_updates() .await .map_err(|error| LoopError::Tool(ToolError::Internal(error.to_string())))?; + self.pending_loop_updates.extend(resolutions); + Ok(()) + } + + async fn drain_pending_loop_updates(&mut self) -> Result<(bool, Option), LoopError> { + self.collect_pending_loop_updates().await?; + let mut resolutions = std::mem::take(&mut self.pending_loop_updates); + if !resolutions.is_empty() { + self.start_logical_turn(); + } let mut saw_items = false; while let Some(resolution) = resolutions.pop_front() { match resolution { @@ -1873,7 +1927,8 @@ where saw_items = true; } TaskResolution::Approval(task) => { - self.enqueue_pending_approval(&task.tool_request.turn_id.clone(), task, None); + let turn_id = self.start_logical_turn(); + self.enqueue_pending_approval(&turn_id, task, None); } } } @@ -1948,30 +2003,30 @@ where } async fn continue_active_tool_round(&mut self) -> Result, LoopError> { - let Some(_) = self.active_tool_round.as_ref() else { + let Some((presentation_turn_id, task_turn_id, cancellation)) = + self.active_tool_round.as_ref().map(|active| { + ( + active.presentation_turn_id.clone(), + active.task_turn_id.clone(), + active.cancellation.clone(), + ) + }) + else { return Ok(None); }; loop { - let turn_id = self - .active_tool_round - .as_ref() - .map(|active| active.turn_id.clone()) - .ok_or_else(|| LoopError::InvalidState("missing active tool round".into()))?; - let cancellation = self - .active_tool_round - .as_ref() - .and_then(|active| active.cancellation.clone()); - if cancellation .as_ref() .is_some_and(TurnCancellation::is_cancelled) { self.task_manager - .on_turn_interrupted(&turn_id) + .on_turn_interrupted(&task_turn_id) .await .map_err(|error| LoopError::Tool(ToolError::Internal(error.to_string())))?; self.active_tool_round = None; - return self.finish_cancelled(turn_id, Vec::new()).map(Some); + return self + .finish_cancelled(presentation_turn_id, Vec::new()) + .map(Some); } let next_call = self @@ -1981,7 +2036,8 @@ where if let Some((call, tool_request)) = next_call { use tracing::Instrument; self.register_tool_cancellation(&call.id, cancellation.clone()); - let dispatch_span = self.execute_tool_span(&tool_request, &turn_id, "plain"); + let dispatch_span = + self.execute_tool_span(&tool_request, &presentation_turn_id, "plain"); match self .start_task_via_manager( None, @@ -2008,7 +2064,11 @@ where self.append_tool_result_item(item); } TaskResolution::Approval(task) => { - self.enqueue_pending_approval(&turn_id, task, cancellation.clone()); + self.enqueue_pending_approval( + &presentation_turn_id, + task, + cancellation.clone(), + ); } } continue; @@ -2016,7 +2076,7 @@ where TaskStartOutcome::Pending { kind, .. } => { self.emit(AgentEvent::ToolExecutionStarted(call.clone())); if kind == agentkit_task_manager::TaskKind::Background { - self.background_call_ids.insert(call.id.clone()); + self.append_detach_placeholder(call.id.clone(), &call.name); if let Some(active) = self.active_tool_round.as_mut() { active.background_pending = true; } @@ -2028,7 +2088,7 @@ where match self .task_manager - .wait_for_turn(&turn_id, cancellation.clone()) + .wait_for_turn(&task_turn_id, cancellation.clone()) .await .map_err(|error| LoopError::Tool(ToolError::Internal(error.to_string())))? { @@ -2042,51 +2102,16 @@ where self.append_tool_result_item(item); } TaskResolution::Approval(task) => { - self.enqueue_pending_approval(&turn_id, task, cancellation.clone()); + self.enqueue_pending_approval( + &presentation_turn_id, + task, + cancellation.clone(), + ); } } } Some(TurnTaskUpdate::Detached(snapshot)) => { - // The task was promoted to background. Push a synthetic - // tool result so the model knows the call is still - // running and can continue its turn. Track the - // call_id so when the real result arrives later via - // the task manager, we route it to a Notification - // item instead of emitting a second tool_result for - // the same call_id (which would violate the - // provider schema — exactly one tool_result per - // tool_use). - // Order matters: append the synthetic placeholder FIRST as - // a real Tool/ToolResult so the tool_use slot is filled - // (provider schemas require exactly one tool_result per - // tool_use). Only AFTER appending do we record the - // call_id in `detached_call_ids` — so the *next* item - // for this call_id (the real completion arriving later - // via the task manager) is the one converted to a - // Notification by `maybe_convert_detached`. - let detached_call_id = snapshot.call_id.clone(); - self.background_call_ids.insert(detached_call_id.clone()); - let detached_result = ToolResultPart { - call_id: detached_call_id.clone(), - output: ToolOutput::Text(format!( - "Tool {} is now running in the background. \ - The result will be delivered when it completes.", - snapshot.tool_name, - )), - is_error: false, - metadata: MetadataMap::new(), - }; - self.emit(AgentEvent::ToolExecutionProgress(detached_result.clone())); - self.append_item(Item { - id: None, - kind: ItemKind::Tool, - parts: vec![Part::ToolResult(detached_result)], - metadata: MetadataMap::new(), - usage: None, - finish_reason: None, - created_at: None, - }); - self.detached_call_ids.insert(detached_call_id); + self.append_detach_placeholder(snapshot.call_id, &snapshot.tool_name); if let Some(active) = self.active_tool_round.as_mut() { active.background_pending = true; active.foreground_progressed = true; @@ -2098,13 +2123,15 @@ where .is_some_and(TurnCancellation::is_cancelled) { self.task_manager - .on_turn_interrupted(&turn_id) + .on_turn_interrupted(&task_turn_id) .await .map_err(|error| { LoopError::Tool(ToolError::Internal(error.to_string())) })?; self.active_tool_round = None; - return self.finish_cancelled(turn_id, Vec::new()).map(Some); + return self + .finish_cancelled(presentation_turn_id, Vec::new()) + .map(Some); } let active = self.active_tool_round.take().ok_or_else(|| { LoopError::InvalidState("missing active tool round".into()) @@ -2126,10 +2153,10 @@ where // pending_round_resume. let info = ToolRoundInfo { session_id: self.session_id.clone(), - turn_id: turn_id.clone(), + turn_id: presentation_turn_id.clone(), transcript_len: self.transcript.len(), }; - self.pending_round_resume = Some(turn_id); + self.pending_round_resume = Some(presentation_turn_id); return Ok(Some(LoopStep::Interrupt(LoopInterrupt::AfterToolResult( info, )))); @@ -2156,7 +2183,6 @@ where async fn drive_turn( &mut self, turn_id: agentkit_core::TurnId, - emit_started: bool, mutation_point: MutationPoint, ) -> Result { let cancellation = self @@ -2187,16 +2213,10 @@ where usage: None, metadata: MetadataMap::new(), }; - self.emit(AgentEvent::TurnFinished(turn_result.clone())); + self.finish_logical_turn(&turn_result); return Ok(LoopStep::Finished(turn_result)); } - if emit_started { - self.emit(AgentEvent::TurnStarted { - session_id: self.session_id.clone(), - turn_id: turn_id.clone(), - }); - } if cancellation .as_ref() .is_some_and(TurnCancellation::is_cancelled) @@ -2437,7 +2457,8 @@ where }) .collect(); self.active_tool_round = Some(ActiveToolRound { - turn_id: turn_id.clone(), + presentation_turn_id: turn_id.clone(), + task_turn_id: turn_id.clone(), pending_calls, cancellation: cancellation.clone(), background_pending: false, @@ -2446,6 +2467,13 @@ where if let Some(step) = self.continue_active_tool_round().await? { return Ok(step); } + self.finish_logical_turn(&TurnResult { + turn_id, + finish_reason: result.finish_reason, + items: output_items, + usage: result.usage, + metadata: result.metadata, + }); return Ok(LoopStep::Interrupt(LoopInterrupt::AwaitingInput( InputRequest { session_id: self.session_id.clone(), @@ -2461,7 +2489,7 @@ where usage: result.usage, metadata: result.metadata, }; - self.emit(AgentEvent::TurnFinished(turn_result.clone())); + self.finish_logical_turn(&turn_result); Ok(LoopStep::Finished(turn_result)) } @@ -2478,14 +2506,17 @@ where ApprovalDecision::Approve => { use tracing::Instrument; self.emit(AgentEvent::ToolExecutionStarted(pending.call.clone())); - let dispatch_span = - self.execute_tool_span(&pending.tool_request, &pending.turn_id, "approved"); + let dispatch_span = self.execute_tool_span( + &pending.tool_request, + &pending.presentation_turn_id, + "approved", + ); let cancellation = self .cancellation .as_ref() .map(CancellationHandle::checkpoint); self.register_tool_cancellation(&pending.call.id, cancellation.clone()); - match self + let start = self .start_task_via_manager( Some(pending.task_id.clone()), pending.tool_request.clone(), @@ -2493,8 +2524,40 @@ where cancellation.clone(), ) .instrument(dispatch_span.clone()) - .await? - { + .await; + let outcome = match start { + Ok(outcome) => outcome, + Err(error) => { + self.append_tool_result_item(Item { + id: None, + kind: ItemKind::Tool, + parts: vec![Part::ToolResult(ToolResultPart { + call_id: pending.call.id.clone(), + output: ToolOutput::Text(format!( + "approved task failed to start: {error}" + )), + is_error: true, + metadata: pending.call.metadata.clone(), + })], + metadata: MetadataMap::new(), + usage: None, + finish_reason: None, + created_at: None, + }); + let turn_id = pending.tool_request.turn_id.clone(); + if let Err(cleanup_error) = + self.task_manager.on_turn_interrupted(&turn_id).await + { + tracing::debug!( + %cleanup_error, + %turn_id, + "failed to clean up turn after approved task start error" + ); + } + return Err(error); + } + }; + match outcome { TaskStartOutcome::Ready(resolution) => { let resolution = *resolution; if let TaskResolution::Item(item) = &resolution @@ -2503,7 +2566,7 @@ where dispatch_span.record("error.type", "tool_error"); } if let Some(step) = self.queue_resolution_interrupt( - &pending.turn_id, + &pending.presentation_turn_id, resolution, cancellation, ) { @@ -2512,7 +2575,19 @@ where } TaskStartOutcome::Pending { kind, .. } => { if kind == agentkit_task_manager::TaskKind::Background { - self.background_call_ids.insert(pending.call.id.clone()); + self.append_detach_placeholder( + pending.call.id.clone(), + &pending.call.name, + ); + } else { + self.active_tool_round = Some(ActiveToolRound { + presentation_turn_id: pending.presentation_turn_id.clone(), + task_turn_id: pending.tool_request.turn_id.clone(), + pending_calls: VecDeque::new(), + cancellation: cancellation.clone(), + background_pending: false, + foreground_progressed: false, + }); } } } @@ -2544,7 +2619,7 @@ where } else if let Some(step) = self.next_unresolved_approval_interrupt() { Ok(step) } else { - self.drive_turn(pending.turn_id, false, MutationPoint::AfterToolResult) + self.drive_turn(pending.presentation_turn_id, MutationPoint::AfterToolResult) .await } } @@ -2565,7 +2640,7 @@ where usage: None, metadata: interrupted_metadata("turn"), }; - self.emit(AgentEvent::TurnFinished(turn_result.clone())); + self.finish_logical_turn(&turn_result); Ok(LoopStep::Finished(turn_result)) } @@ -2706,7 +2781,11 @@ where call_id.0 ))); }; + let turn_id = pending.presentation_turn_id.clone(); self.reject_drained_approvals(vec![pending]); + if self.pending_approvals.is_empty() && self.active_tool_round.is_none() { + let _ = self.finish_cancelled(turn_id, Vec::new())?; + } Ok(()) } @@ -2723,17 +2802,44 @@ where .pending_approval_order .iter() .find_map(|call_id| self.pending_approvals.get(call_id)) - .map(|pending| pending.turn_id.clone()) + .map(|pending| pending.presentation_turn_id.clone()) else { return Ok(None); }; + let mut seen_turns = HashSet::new(); + let mut originating_turns = Vec::new(); + for pending in self.pending_approvals.values() { + let originating_turn = pending.tool_request.turn_id.clone(); + if seen_turns.insert(originating_turn.clone()) { + originating_turns.push(originating_turn); + } + } + let pending = self.drain_pending_approval_items(); self.active_tool_round = None; - self.task_manager - .on_turn_interrupted(&turn_id) - .await - .map_err(|error| LoopError::Tool(ToolError::Internal(error.to_string())))?; + let mut cleanup_error = None; + for originating_turn in originating_turns { + if let Err(error) = self + .task_manager + .on_turn_interrupted(&originating_turn) + .await + && cleanup_error.is_none() + { + cleanup_error = Some(LoopError::Tool(ToolError::Internal(error.to_string()))); + } + } self.reject_drained_approvals(pending); + if let Some(error) = cleanup_error { + self.close_interrupted_tool_calls(); + self.finish_logical_turn(&TurnResult { + turn_id, + finish_reason: FinishReason::Error, + items: Vec::new(), + usage: None, + metadata: MetadataMap::new(), + }); + return Err(error); + } self.finish_cancelled(turn_id, Vec::new()).map(Some) } @@ -2746,6 +2852,27 @@ where } } + /// Wait until an out-of-band update is available for the loop. + /// + /// This resolves immediately for updates already collected from the task + /// manager but deferred behind fresh input. It does not consume the update; + /// call [`next`](Self::next) after it resolves to append and drive the result. + pub fn wait_for_loop_update( + &self, + ) -> impl std::future::Future> + Send + 'static { + let has_collected_update = !self.pending_loop_updates.is_empty(); + let task_manager = self.task_manager.clone(); + async move { + if has_collected_update { + return Ok(()); + } + task_manager + .wait_for_loop_update() + .await + .map_err(|error| LoopError::Tool(ToolError::Internal(error.to_string()))) + } + } + /// Advance the loop by one step. /// /// This is the main method for driving the agent. It processes pending @@ -2767,6 +2894,76 @@ where /// Returns [`LoopError::InvalidState`] if called while an unresolved /// interrupt is pending, or propagates provider / tool / compaction errors. pub async fn next(&mut self) -> Result { + if self.lifecycle.active_turn.is_none() { + let continuation_turn = self + .pending_approval_order + .iter() + .find_map(|call_id| self.pending_approvals.get(call_id)) + .map(|pending| pending.presentation_turn_id.clone()) + .or_else(|| { + self.active_tool_round + .as_ref() + .map(|active| active.presentation_turn_id.clone()) + }) + .or_else(|| self.pending_round_resume.clone()); + if let Some(turn_id) = continuation_turn { + self.start_logical_turn_with(turn_id); + } else if !self.pending_input.is_empty() { + self.start_logical_turn(); + } + } + + let result = self.next_inner().await; + match &result { + Ok(LoopStep::Finished(turn)) => self.finish_logical_turn(turn), + Err(_) => { + if let Some(turn_id) = self.lifecycle.active_turn.clone() { + self.recover_from_next_error().await; + self.finish_logical_turn(&TurnResult { + turn_id, + finish_reason: FinishReason::Error, + items: Vec::new(), + usage: None, + metadata: MetadataMap::new(), + }); + } + } + _ => {} + } + result + } + + async fn recover_from_next_error(&mut self) { + let mut seen_turns = HashSet::new(); + let mut interrupted_turns = Vec::new(); + if let Some(active) = self.active_tool_round.take() + && seen_turns.insert(active.task_turn_id.clone()) + { + interrupted_turns.push(active.task_turn_id); + } + if let Some(turn_id) = self.pending_round_resume.take() + && seen_turns.insert(turn_id.clone()) + { + interrupted_turns.push(turn_id); + } + for pending in self.pending_approvals.values() { + let turn_id = pending.tool_request.turn_id.clone(); + if seen_turns.insert(turn_id.clone()) { + interrupted_turns.push(turn_id); + } + } + + let pending = self.drain_pending_approval_items(); + for turn_id in interrupted_turns { + if let Err(error) = self.task_manager.on_turn_interrupted(&turn_id).await { + tracing::debug!(%error, %turn_id, "failed to clean up turn after loop error"); + } + } + self.reject_drained_approvals(pending); + self.close_interrupted_tool_calls(); + } + + async fn next_inner(&mut self) -> Result { if let Some(pending) = self.take_next_resolved_approval() { return self.resume_after_approval(pending).await; } @@ -2787,6 +2984,22 @@ where return Ok(step); } + // A newly submitted user turn owns the next logical turn. Drive it + // before unrelated background completions so a delayed approval cannot + // bind itself to that turn's TurnStarted event. AfterToolResult resumes + // remain ordered ahead of fresh input below. + if self.pending_round_resume.is_none() && !self.pending_input.is_empty() { + // Take updates now to preserve the driver's once-per-step manager + // handoff, but defer presenting them until this input turn ends. + self.collect_pending_loop_updates().await?; + let turn_id = self.start_logical_turn(); + let drained: Vec = std::mem::take(&mut self.pending_input); + self.extend_transcript(drained); + return self + .drive_turn(turn_id, MutationPoint::AfterTurnEnded) + .await; + } + let (had_loop_updates, loop_step) = self.drain_pending_loop_updates().await?; if let Some(step) = loop_step { return Ok(step); @@ -2800,7 +3013,7 @@ where let drained: Vec = std::mem::take(&mut self.pending_input); self.extend_transcript(drained); return self - .drive_turn(turn_id, false, MutationPoint::AfterToolResult) + .drive_turn(turn_id, MutationPoint::AfterToolResult) .await; } @@ -2813,11 +3026,10 @@ where ))); } - let turn_id = agentkit_core::TurnId::new(format!("turn-{}", self.next_turn_index)); - self.next_turn_index += 1; + let turn_id = self.start_logical_turn(); let drained: Vec = std::mem::take(&mut self.pending_input); self.extend_transcript(drained); - self.drive_turn(turn_id, true, MutationPoint::AfterTurnEnded) + self.drive_turn(turn_id, MutationPoint::AfterTurnEnded) .await } @@ -2842,6 +3054,31 @@ where self.transcript.push(item); } + fn append_detach_placeholder(&mut self, call_id: ToolCallId, tool_name: &str) { + self.background_call_ids.insert(call_id.clone()); + if !self.detached_call_ids.insert(call_id.clone()) { + return; + } + let detached_result = ToolResultPart { + call_id: call_id.clone(), + output: ToolOutput::Text(format!( + "Tool {tool_name} is now running in the background. The result will be delivered when it completes." + )), + is_error: false, + metadata: MetadataMap::new(), + }; + self.emit(AgentEvent::ToolExecutionProgress(detached_result.clone())); + self.append_item(Item { + id: None, + kind: ItemKind::Tool, + parts: vec![Part::ToolResult(detached_result)], + metadata: MetadataMap::new(), + usage: None, + finish_reason: None, + created_at: None, + }); + } + /// Append a tool-result Item: emit one [`AgentEvent::ToolResultReceived`] /// per [`Part::ToolResult`] inside the Item, then funnel through /// [`Self::append_item`]. @@ -2931,7 +3168,7 @@ where } } - fn maybe_convert_detached(&mut self, item: Item) -> Item { + fn maybe_convert_detached(&mut self, mut item: Item) -> Item { if !matches!(item.kind, ItemKind::Tool) { return item; } @@ -2950,25 +3187,48 @@ where { return item; } - let mut text = String::new(); - for result in &results { + let structured_results = results + .iter() + .map(|result| { + Part::structured(serde_json::to_value(result).unwrap_or_else( + |error| serde_json::json!({ "serialization_error": error.to_string() }), + )) + }) + .collect::>(); + let failed = results.iter().filter(|result| result.is_error).count(); + let with_metadata = results + .iter() + .filter(|result| !result.metadata.is_empty()) + .count(); + let mut text = format!( + "Background tool results: {} total, {failed} failed, {with_metadata} with metadata. ", + results.len() + ); + for (index, result) in results.iter().enumerate() { self.detached_call_ids.remove(&result.call_id); self.interrupted_background_call_ids.remove(&result.call_id); - if !text.is_empty() { - text.push_str("\n\n"); + if text.chars().count() >= DETACHED_NOTIFICATION_TEXT_MAX_CHARS { + continue; + } + if index > 0 { + text.push_str("; "); } let label = if result.is_error { "failed" } else { "completed" }; + let call_id = truncate_chars(&result.call_id.0, DETACHED_CALL_ID_MAX_CHARS); let body = render_tool_output_brief(&result.output); - text.push_str(&format!( - "Background tool call {} {}: {body}", - result.call_id.0, label - )); + text.push_str(&format!("{call_id} {label}: {body}")); } - Item::notification(text) + let text = truncate_chars(&text, DETACHED_NOTIFICATION_TEXT_MAX_CHARS); + let mut notification_parts = Vec::with_capacity(1 + structured_results.len()); + notification_parts.push(Part::text(text)); + notification_parts.extend(structured_results); + item.kind = ItemKind::Notification; + item.parts = notification_parts; + item } /// Append several Items in order through [`Self::append_item`]. @@ -2987,11 +3247,24 @@ where fn render_tool_output_brief(output: &ToolOutput) -> String { match output { - ToolOutput::Text(t) => t.clone(), - ToolOutput::Structured(value) => value.to_string(), - ToolOutput::Parts(parts) => format!("[{} parts]", parts.len()), - ToolOutput::Files(files) => format!("[{} files]", files.len()), + ToolOutput::Text(text) => format!( + "text preview: {}", + truncate_chars(text, DETACHED_TEXT_PREVIEW_MAX_CHARS) + ), + ToolOutput::Structured(_) => "structured payload".into(), + ToolOutput::Parts(parts) => format!("parts payload ({} parts)", parts.len()), + ToolOutput::Files(files) => format!("files payload ({} files)", files.len()), + } +} + +fn truncate_chars(text: &str, max_chars: usize) -> String { + let mut chars = text.chars(); + let mut truncated = chars.by_ref().take(max_chars).collect::(); + if chars.next().is_some() && max_chars > 0 { + truncated.pop(); + truncated.push('…'); } + truncated } fn interrupted_metadata(stage: &str) -> MetadataMap { @@ -4023,8 +4296,8 @@ mod tests { ToolResultPart, }; use agentkit_task_manager::{ - AsyncTaskManager, RoutingDecision, TaskEvent, TaskManager, TaskManagerHandle, - TaskRoutingPolicy, + AsyncTaskManager, RoutingDecision, TaskEvent, TaskManager, TaskManagerError, + TaskManagerHandle, TaskRoutingPolicy, }; use agentkit_tools_core::{ FileSystemPermissionRequest, PermissionCode, PermissionDecision, PermissionDenial, Tool, @@ -4073,9 +4346,115 @@ mod tests { events: VecDeque, } + struct TestTaskManager { + inner: T, + start_error: Option<&'static str>, + approved_start_error: Option<&'static str>, + pending_update_error: Option<(usize, &'static str)>, + pending_update_calls: AtomicUsize, + interrupted: Option>>>, + interrupt_error: Option<&'static str>, + } + + impl TestTaskManager { + fn new(inner: T) -> Self { + Self { + inner, + start_error: None, + approved_start_error: None, + pending_update_error: None, + pending_update_calls: AtomicUsize::new(0), + interrupted: None, + interrupt_error: None, + } + } + + fn fail_start(mut self, message: &'static str) -> Self { + self.start_error = Some(message); + self + } + + fn fail_approved_start(mut self, message: &'static str) -> Self { + self.approved_start_error = Some(message); + self + } + + fn fail_pending_update_on(mut self, call: usize, message: &'static str) -> Self { + self.pending_update_error = Some((call, message)); + self + } + + fn record_interrupts( + mut self, + interrupted: StdArc>>, + ) -> Self { + self.interrupted = Some(interrupted); + self + } + + fn fail_interrupt(mut self, message: &'static str) -> Self { + self.interrupt_error = Some(message); + self + } + } + + #[async_trait] + impl TaskManager for TestTaskManager { + async fn start_task( + &self, + request: TaskLaunchRequest, + ctx: TaskStartContext, + ) -> Result { + if let Some(message) = self.start_error.or_else(|| { + matches!(&request.kind, TaskLaunchKind::Approved(_)) + .then_some(self.approved_start_error) + .flatten() + }) { + return Err(TaskManagerError::Internal(message.into())); + } + self.inner.start_task(request, ctx).await + } + + async fn wait_for_turn( + &self, + turn_id: &agentkit_core::TurnId, + cancellation: Option, + ) -> Result, TaskManagerError> { + self.inner.wait_for_turn(turn_id, cancellation).await + } + + async fn take_pending_loop_updates(&self) -> Result { + if let Some((call, message)) = self.pending_update_error + && self.pending_update_calls.fetch_add(1, Ordering::SeqCst) == call + { + return Err(TaskManagerError::Internal(message.into())); + } + self.inner.take_pending_loop_updates().await + } + + async fn on_turn_interrupted( + &self, + turn_id: &agentkit_core::TurnId, + ) -> Result<(), TaskManagerError> { + if let Some(interrupted) = &self.interrupted { + interrupted.lock().unwrap().push(turn_id.clone()); + } + if let Some(message) = self.interrupt_error { + return Err(TaskManagerError::Internal(message.into())); + } + self.inner.on_turn_interrupted(turn_id).await + } + + fn handle(&self) -> TaskManagerHandle { + self.inner.handle() + } + } + struct DelayedApprovalExecutor { entered: StdArc, release: StdArc, + approved_entered: Option>, + approved_release: Option>, cancellation: Option, spec: ToolSpec, } @@ -4085,6 +4464,8 @@ mod tests { Self { entered, release, + approved_entered: None, + approved_release: None, cancellation: None, spec: ToolSpec { name: ToolName::new("echo"), @@ -4108,6 +4489,16 @@ mod tests { self.cancellation = Some(controller); self } + + fn blocking_after_approval( + mut self, + entered: StdArc, + release: StdArc, + ) -> Self { + self.approved_entered = Some(entered); + self.approved_release = Some(release); + self + } } #[async_trait] @@ -4138,6 +4529,31 @@ mod tests { }), ) } + + async fn execute_approved( + &self, + request: ToolRequest, + approved_request: &ApprovalRequest, + ctx: &mut ToolContext<'_>, + ) -> ToolExecutionOutcome { + let (Some(entered), Some(release)) = (&self.approved_entered, &self.approved_release) + else { + return self.execute(request, ctx).await; + }; + let _ = approved_request; + entered.store(true, Ordering::SeqCst); + release.notified().await; + ToolExecutionOutcome::Completed(ToolResult { + result: ToolResultPart { + call_id: request.call_id, + output: ToolOutput::Text("approved-ok".into()), + is_error: false, + metadata: MetadataMap::new(), + }, + duration: None, + metadata: MetadataMap::new(), + }) + } } #[async_trait] @@ -4216,11 +4632,15 @@ mod tests { .iter() .rev() .find_map(|item| { - item.parts.iter().find_map(|part| match part { - Part::ToolResult(ToolResultPart { - output: ToolOutput::Text(text), - .. - }) => Some(text.clone()), + item.parts.iter().find_map(|part| match (item.kind, part) { + (ItemKind::Notification, Part::Text(text)) => Some(text.text.clone()), + ( + _, + Part::ToolResult(ToolResultPart { + output: ToolOutput::Text(text), + .. + }), + ) => Some(text.clone()), _ => None, }) }) @@ -4907,6 +5327,21 @@ mod tests { } } + fn turn_lifecycle_events( + events: &[AgentEvent], + ) -> Vec<(agentkit_core::TurnId, Option)> { + events + .iter() + .filter_map(|event| match event { + AgentEvent::TurnStarted { turn_id, .. } => Some((turn_id.clone(), None)), + AgentEvent::TurnFinished(turn) => { + Some((turn.turn_id.clone(), Some(turn.finish_reason.clone()))) + } + _ => None, + }) + .collect() + } + struct CatalogExecutor { version: AtomicUsize, events: StdMutex>, @@ -5080,11 +5515,15 @@ mod tests { #[tokio::test] async fn loop_continues_after_completed_tool_call() { + let events = StdArc::new(StdMutex::new(Vec::new())); let tools = ToolRegistry::new().with(EchoTool::default()); let agent = Agent::builder() .model(FakeAdapter) .add_tool_source(tools) .permissions(AllowAllPermissions) + .observer(RecordingObserver { + events: events.clone(), + }) .build() .unwrap(); @@ -5125,6 +5564,12 @@ mod tests { } other => panic!("unexpected loop step: {other:?}"), } + + let lifecycle = turn_lifecycle_events(&events.lock().unwrap()); + assert_eq!(lifecycle.len(), 2); + assert_eq!(lifecycle[0].0, lifecycle[1].0); + assert_eq!(lifecycle[0].1, None); + assert_eq!(lifecycle[1].1, Some(FinishReason::Completed)); } /// Test helper: drives the loop, transparently resuming non-blocking @@ -5187,27 +5632,235 @@ mod tests { ); } - #[test] - fn pending_input_requires_input_bearing_tail_role() { - assert!(!transcript_has_pending_input(&[])); - assert!(!transcript_has_pending_input(&[Item::text( - ItemKind::System, - "system" - )])); - assert!(!transcript_has_pending_input(&[Item::text( - ItemKind::Developer, - "developer" - )])); - assert!(!transcript_has_pending_input(&[Item::text( - ItemKind::Context, - "context" - )])); - assert!(!transcript_has_pending_input(&[Item::text( - ItemKind::Assistant, - "assistant" - )])); - - assert!(transcript_has_pending_input(&[Item::text( + #[tokio::test] + async fn no_work_awaiting_input_emits_no_turn_lifecycle() { + let events = StdArc::new(StdMutex::new(Vec::new())); + let agent = Agent::builder() + .model(SlowAdapter) + .observer(RecordingObserver { + events: events.clone(), + }) + .build() + .unwrap(); + let mut driver = agent + .start(SessionConfig::new("session-no-work")) + .await + .unwrap(); + + for _ in 0..2 { + assert!(matches!( + driver.next().await.unwrap(), + LoopStep::Interrupt(LoopInterrupt::AwaitingInput(_)) + )); + } + + assert!(turn_lifecycle_events(&events.lock().unwrap()).is_empty()); + } + + #[tokio::test] + async fn normal_turn_emits_one_matched_lifecycle_pair() { + let events = StdArc::new(StdMutex::new(Vec::new())); + let agent = Agent::builder() + .model(SlowAdapter) + .observer(RecordingObserver { + events: events.clone(), + }) + .build() + .unwrap(); + let mut driver = agent + .start(SessionConfig::new("session-normal-lifecycle")) + .await + .unwrap(); + driver + .submit_input(vec![Item::text(ItemKind::User, "ping")]) + .unwrap(); + + assert!(matches!( + driver.next().await.unwrap(), + LoopStep::Finished(TurnResult { + finish_reason: FinishReason::Completed, + .. + }) + )); + + let lifecycle = turn_lifecycle_events(&events.lock().unwrap()); + assert_eq!(lifecycle.len(), 2); + assert_eq!(lifecycle[0].0, lifecycle[1].0); + assert_eq!(lifecycle[0].1, None); + assert_eq!(lifecycle[1].1, Some(FinishReason::Completed)); + } + + #[tokio::test] + async fn post_start_error_emits_terminal_error_without_run_failed() { + let events = StdArc::new(StdMutex::new(Vec::new())); + let agent = Agent::builder() + .model(SlowAdapter) + .mutator(ErrorMutator) + .observer(RecordingObserver { + events: events.clone(), + }) + .build() + .unwrap(); + let mut driver = agent + .start(SessionConfig::new("session-error-lifecycle")) + .await + .unwrap(); + driver + .submit_input(vec![Item::text(ItemKind::User, "ping")]) + .unwrap(); + + assert!(matches!( + driver.next().await, + Err(LoopError::Mutator(message)) if message == "boom" + )); + + let events = events.lock().unwrap(); + let lifecycle = turn_lifecycle_events(&events); + assert_eq!(lifecycle.len(), 2); + assert_eq!(lifecycle[0].0, lifecycle[1].0); + assert_eq!(lifecycle[0].1, None); + assert_eq!(lifecycle[1].1, Some(FinishReason::Error)); + assert!( + !events + .iter() + .any(|event| matches!(event, AgentEvent::RunFailed { .. })) + ); + } + + #[tokio::test] + async fn active_tool_error_repairs_state_and_retry_uses_fresh_lifecycle() { + let events = StdArc::new(StdMutex::new(Vec::new())); + let interrupted = StdArc::new(StdMutex::new(Vec::new())); + let agent = Agent::builder() + .model(FakeAdapter) + .add_tool_source(ToolRegistry::new().with(EchoTool::default())) + .task_manager( + TestTaskManager::new(SimpleTaskManager::new()) + .fail_start("original start failure") + .record_interrupts(interrupted.clone()) + .fail_interrupt("cleanup failure"), + ) + .observer(RecordingObserver { + events: events.clone(), + }) + .build() + .unwrap(); + let mut driver = agent + .start(SessionConfig::new("session-active-tool-error")) + .await + .unwrap(); + driver + .submit_input(vec![Item::text(ItemKind::User, "first")]) + .unwrap(); + + let error = driver.next().await.unwrap_err(); + assert!(error.to_string().contains("original start failure")); + assert!(!error.to_string().contains("cleanup failure")); + assert!(driver.active_tool_round.is_none()); + assert!(driver.pending_round_resume.is_none()); + assert!(unanswered_tool_calls(&driver.snapshot().transcript).is_empty()); + validate_transcript_invariants(&driver.snapshot().transcript).unwrap(); + assert_eq!(interrupted.lock().unwrap().len(), 1); + + driver + .submit_input(vec![Item::text(ItemKind::User, "retry")]) + .unwrap(); + assert!(matches!( + driver.next().await.unwrap(), + LoopStep::Finished(TurnResult { + finish_reason: FinishReason::Completed, + .. + }) + )); + + let lifecycle = turn_lifecycle_events(&events.lock().unwrap()); + assert_eq!(lifecycle.len(), 4, "{lifecycle:?}"); + assert_eq!(lifecycle[0].0, lifecycle[1].0); + assert_eq!(lifecycle[1].1, Some(FinishReason::Error)); + assert_eq!(lifecycle[2].0, lifecycle[3].0); + assert_eq!(lifecycle[3].1, Some(FinishReason::Completed)); + assert_ne!(lifecycle[0].0, lifecycle[2].0); + } + + #[tokio::test] + async fn continuation_error_clears_resume_and_retry_uses_fresh_lifecycle() { + let events = StdArc::new(StdMutex::new(Vec::new())); + let interrupted = StdArc::new(StdMutex::new(Vec::new())); + let agent = Agent::builder() + .model(FakeAdapter) + .add_tool_source(ToolRegistry::new().with(EchoTool::default())) + .task_manager( + TestTaskManager::new(SimpleTaskManager::new()) + .fail_pending_update_on(1, "original continuation failure") + .record_interrupts(interrupted.clone()) + .fail_interrupt("cleanup failure"), + ) + .observer(RecordingObserver { + events: events.clone(), + }) + .build() + .unwrap(); + let mut driver = agent + .start(SessionConfig::new("session-continuation-error")) + .await + .unwrap(); + driver + .submit_input(vec![Item::text(ItemKind::User, "first")]) + .unwrap(); + + assert!(matches!( + driver.next().await.unwrap(), + LoopStep::Interrupt(LoopInterrupt::AfterToolResult(_)) + )); + let error = driver.next().await.unwrap_err(); + assert!(error.to_string().contains("original continuation failure")); + assert!(!error.to_string().contains("cleanup failure")); + assert!(driver.pending_round_resume.is_none()); + assert!(driver.active_tool_round.is_none()); + validate_transcript_invariants(&driver.snapshot().transcript).unwrap(); + assert_eq!(interrupted.lock().unwrap().len(), 1); + + driver + .submit_input(vec![Item::text(ItemKind::User, "retry")]) + .unwrap(); + assert!(matches!( + driver.next().await.unwrap(), + LoopStep::Finished(TurnResult { + finish_reason: FinishReason::Completed, + .. + }) + )); + + let lifecycle = turn_lifecycle_events(&events.lock().unwrap()); + assert_eq!(lifecycle.len(), 4, "{lifecycle:?}"); + assert_eq!(lifecycle[0].0, lifecycle[1].0); + assert_eq!(lifecycle[1].1, Some(FinishReason::Error)); + assert_eq!(lifecycle[2].0, lifecycle[3].0); + assert_eq!(lifecycle[3].1, Some(FinishReason::Completed)); + assert_ne!(lifecycle[0].0, lifecycle[2].0); + } + + #[test] + fn pending_input_requires_input_bearing_tail_role() { + assert!(!transcript_has_pending_input(&[])); + assert!(!transcript_has_pending_input(&[Item::text( + ItemKind::System, + "system" + )])); + assert!(!transcript_has_pending_input(&[Item::text( + ItemKind::Developer, + "developer" + )])); + assert!(!transcript_has_pending_input(&[Item::text( + ItemKind::Context, + "context" + )])); + assert!(!transcript_has_pending_input(&[Item::text( + ItemKind::Assistant, + "assistant" + )])); + + assert!(transcript_has_pending_input(&[Item::text( ItemKind::User, "user" )])); @@ -5236,6 +5889,18 @@ mod tests { /// strips an empty user prompt — leaving the transcript ending in an /// assistant message. struct DropTrailingUserMutator; + struct ErrorMutator; + + #[async_trait] + impl LoopMutator for ErrorMutator { + async fn mutate( + &self, + _cursor: &mut TranscriptCursor<'_>, + _ctx: LoopCtx<'_>, + ) -> Result<(), LoopError> { + Err(LoopError::Mutator("boom".into())) + } + } #[async_trait] impl LoopMutator for DropTrailingUserMutator { @@ -5312,11 +5977,15 @@ mod tests { #[tokio::test] async fn drive_does_not_dispatch_without_valid_trailing_input() { let saw_assistant_tail = StdArc::new(AtomicBool::new(false)); + let events = StdArc::new(StdMutex::new(Vec::new())); let agent = Agent::builder() .model(RejectAssistantPrefillAdapter { saw_assistant_tail: saw_assistant_tail.clone(), }) .mutator(DropTrailingUserMutator) + .observer(RecordingObserver { + events: events.clone(), + }) // Prior conversation ending in an assistant message — e.g. a cold // bootstrap that loaded a completed turn's history. .transcript(vec![ @@ -5348,6 +6017,11 @@ mod tests { message (outcome: {outcome:?}); with no valid trailing input the turn \ must finish instead of driving" ); + let lifecycle = turn_lifecycle_events(&events.lock().unwrap()); + assert_eq!(lifecycle.len(), 2); + assert_eq!(lifecycle[0].0, lifecycle[1].0); + assert_eq!(lifecycle[0].1, None); + assert_eq!(lifecycle[1].1, Some(FinishReason::Completed)); } #[tokio::test] @@ -5557,218 +6231,892 @@ mod tests { }]) .unwrap(); - let first = driver.next().await.unwrap(); - match first { - LoopStep::Interrupt(LoopInterrupt::AwaitingInput(_)) => {} - other => panic!("unexpected first loop step: {other:?}"), - } - - match wait_for_task_event(&handle).await { - TaskEvent::Started(snapshot) => assert_eq!(snapshot.tool_name, "background-wait"), - other => panic!("unexpected task event: {other:?}"), - } + let first = driver.next().await.unwrap(); + match first { + LoopStep::Interrupt(LoopInterrupt::AwaitingInput(_)) => {} + other => panic!("unexpected first loop step: {other:?}"), + } + + validate_transcript_invariants(&driver.snapshot().transcript).unwrap(); + let lifecycle = turn_lifecycle_events(&events.lock().unwrap()); + assert_eq!(lifecycle.len(), 2); + assert_eq!(lifecycle[0].0, lifecycle[1].0); + assert_eq!(lifecycle[1].1, Some(FinishReason::ToolCall)); + + match wait_for_task_event(&handle).await { + TaskEvent::Started(snapshot) => assert_eq!(snapshot.tool_name, "background-wait"), + other => panic!("unexpected task event: {other:?}"), + } + wait_until_entered(entered.as_ref()).await; + release.notify_waiters(); + + match wait_for_task_event(&handle).await { + TaskEvent::Completed(_, result) => { + assert_eq!(result.output, ToolOutput::Text("background-done".into())) + } + other => panic!("unexpected completion event: {other:?}"), + } + + let resumed = driver.next().await.unwrap(); + match resumed { + LoopStep::Finished(turn) => { + assert_eq!(turn.finish_reason, FinishReason::Completed); + match &turn.items[0].parts[0] { + Part::Text(text) => assert_eq!( + text.text, + "tool said: Background tool results: 1 total, 0 failed, 0 with metadata. \ + call-1 completed: text preview: background-done" + ), + other => panic!("unexpected part after resume: {other:?}"), + } + } + other => panic!("unexpected resumed step: {other:?}"), + } + + let events = events.lock().unwrap(); + let lifecycle = turn_lifecycle_events(&events); + assert_eq!(lifecycle.len(), 4); + assert_eq!(lifecycle[0].0, lifecycle[1].0); + assert_eq!(lifecycle[2].0, lifecycle[3].0); + assert_ne!(lifecycle[0].0, lifecycle[2].0); + assert_eq!(lifecycle[3].1, Some(FinishReason::Completed)); + + let terminal_results: Vec<_> = events + .iter() + .filter_map(|event| match event { + AgentEvent::ToolResultReceived(result) + if result.call_id == ToolCallId::new("call-1") => + { + Some(result) + } + _ => None, + }) + .collect(); + assert_eq!( + terminal_results.len(), + 1, + "background completion must emit one terminal result event per call: {events:?}" + ); + } + + #[tokio::test] + async fn detached_parts_notification_preserves_full_output_and_metadata() { + let agent = Agent::builder().model(FakeAdapter).build().unwrap(); + let mut driver = agent + .start(SessionConfig::new("session-detached-parts")) + .await + .unwrap(); + let call_id = ToolCallId::new("parts-call"); + driver.detached_call_ids.insert(call_id.clone()); + let parts = vec![ + Part::text("part text"), + Part::structured(json!({ + "nested": [1, 2, 3] + })), + ]; + let mut metadata = MetadataMap::new(); + metadata.insert("source".into(), json!("background")); + let result = ToolResultPart { + call_id, + output: ToolOutput::Parts(parts.clone()), + is_error: true, + metadata: metadata.clone(), + }; + let mut item_metadata = MetadataMap::new(); + item_metadata.insert("delivery".into(), json!("deferred")); + let item = Item::new(ItemKind::Tool, vec![Part::ToolResult(result.clone())]) + .with_metadata(item_metadata.clone()); + + let converted = driver.maybe_convert_detached(item); + let (text, structured) = match converted.parts.as_slice() { + [Part::Text(text), Part::Structured(structured)] => (text, structured), + other => panic!("unexpected converted parts: {other:?}"), + }; + assert_eq!(converted.kind, ItemKind::Notification); + assert_eq!(converted.metadata, item_metadata); + assert_eq!(structured.value, serde_json::to_value(&result).unwrap()); + assert_eq!( + text.text, + "Background tool results: 1 total, 1 failed, 1 with metadata. \ + parts-call failed: parts payload (2 parts)" + ); + assert!(!text.text.contains("part text")); + assert!(!text.text.contains("background")); + } + + #[tokio::test] + async fn detached_files_notification_preserves_full_output() { + let agent = Agent::builder().model(FakeAdapter).build().unwrap(); + let mut driver = agent + .start(SessionConfig::new("session-detached-files")) + .await + .unwrap(); + let call_id = ToolCallId::new("files-call"); + driver.detached_call_ids.insert(call_id.clone()); + let files = vec![ + agentkit_core::FilePart::named("report.txt", DataRef::inline_text("full file body")) + .with_mime_type("text/plain"), + agentkit_core::FilePart::named( + "remote.json", + DataRef::uri("https://example.test/remote.json"), + ), + ]; + let mut result_metadata = MetadataMap::new(); + result_metadata.insert("archive".into(), json!(true)); + let result = ToolResultPart::success(call_id, ToolOutput::Files(files.clone())) + .with_metadata(result_metadata); + let mut item_metadata = MetadataMap::new(); + item_metadata.insert("delivery".into(), json!("deferred")); + let item = Item::new(ItemKind::Tool, vec![Part::ToolResult(result.clone())]) + .with_metadata(item_metadata.clone()); + + let converted = driver.maybe_convert_detached(item); + let (text, structured) = match converted.parts.as_slice() { + [Part::Text(text), Part::Structured(structured)] => (text, structured), + other => panic!("unexpected converted files: {other:?}"), + }; + assert_eq!(converted.kind, ItemKind::Notification); + assert_eq!(converted.metadata, item_metadata); + assert_eq!(structured.value, serde_json::to_value(&result).unwrap()); + assert_eq!( + text.text, + "Background tool results: 1 total, 0 failed, 1 with metadata. \ + files-call completed: files payload (2 files)" + ); + assert!(!text.text.contains("full file body")); + assert!(!text.text.contains("remote.json")); + } + + #[test] + fn detached_result_summaries_are_bounded_and_do_not_serialize_structured_payloads() { + let long_text = "é".repeat(DETACHED_TEXT_PREVIEW_MAX_CHARS + 20); + let text_summary = render_tool_output_brief(&ToolOutput::Text(long_text.clone())); + assert_eq!( + text_summary.chars().count(), + "text preview: ".chars().count() + DETACHED_TEXT_PREVIEW_MAX_CHARS + ); + assert!(text_summary.ends_with('…')); + assert!(!text_summary.contains(&long_text)); + + let secret = "structured payload must remain out of notification text"; + let structured = ToolOutput::Structured(json!({ "secret": secret })); + assert_eq!(render_tool_output_brief(&structured), "structured payload"); + + let oversized = "x".repeat(DETACHED_NOTIFICATION_TEXT_MAX_CHARS + 20); + let bounded = truncate_chars(&oversized, DETACHED_NOTIFICATION_TEXT_MAX_CHARS); + assert_eq!( + bounded.chars().count(), + DETACHED_NOTIFICATION_TEXT_MAX_CHARS + ); + assert!(bounded.ends_with('…')); + } + + #[tokio::test] + async fn detached_tool_placeholder_is_progress_not_terminal_result() { + let events = StdArc::new(StdMutex::new(Vec::new())); + let entered = StdArc::new(AtomicBool::new(false)); + let release = StdArc::new(Notify::new()); + let task_manager = AsyncTaskManager::new().routing(NameRoutingPolicy::new([( + "detaching-wait", + RoutingDecision::ForegroundThenDetachAfter(Duration::from_millis(10)), + )])); + let handle = task_manager.handle(); + let tools = ToolRegistry::new().with(BlockingTool::new( + "detaching-wait", + entered.clone(), + release.clone(), + "detached-done", + )); + let agent = Agent::builder() + .model(FakeAdapter) + .add_tool_source(tools) + .permissions(AllowAllPermissions) + .task_manager(task_manager) + .observer(RecordingObserver { + events: events.clone(), + }) + .build() + .unwrap(); + + let mut driver = agent + .start(SessionConfig { + session_id: SessionId::new("session-detached-progress"), + metadata: MetadataMap::new(), + cache: None, + }) + .await + .unwrap(); + + driver + .submit_input(vec![Item::text(ItemKind::User, "ping")]) + .unwrap(); + + match driver.next().await.unwrap() { + LoopStep::Interrupt(LoopInterrupt::AfterToolResult(_)) => {} + other => panic!("unexpected detach step: {other:?}"), + } + + match wait_for_task_event(&handle).await { + TaskEvent::Started(snapshot) => assert_eq!(snapshot.tool_name, "detaching-wait"), + other => panic!("unexpected task event: {other:?}"), + } + match wait_for_task_event(&handle).await { + TaskEvent::Detached(snapshot) => assert_eq!(snapshot.tool_name, "detaching-wait"), + other => panic!("unexpected detach event: {other:?}"), + } + wait_until_entered(entered.as_ref()).await; + release.notify_waiters(); + + match wait_for_task_event(&handle).await { + TaskEvent::Completed(_, result) => { + assert_eq!(result.output, ToolOutput::Text("detached-done".into())) + } + other => panic!("unexpected completion event: {other:?}"), + } + + match driver.next().await.unwrap() { + LoopStep::Finished(turn) => assert_eq!(turn.finish_reason, FinishReason::Completed), + other => panic!("unexpected resumed step: {other:?}"), + } + + let events = events.lock().unwrap(); + assert!(events.iter().any(|event| matches!( + event, + AgentEvent::ToolExecutionProgress(result) + if result.call_id == ToolCallId::new("call-1") && !result.is_error + ))); + let terminal_results: Vec<_> = events + .iter() + .filter_map(|event| match event { + AgentEvent::ToolResultReceived(result) + if result.call_id == ToolCallId::new("call-1") => + { + Some(result) + } + _ => None, + }) + .collect(); + assert_eq!( + terminal_results.len(), + 1, + "detached call must emit one terminal result event: {events:?}" + ); + } + + #[tokio::test] + async fn cancelled_background_approval_auto_resolves_when_drained() { + let controller = CancellationController::new(); + let events = StdArc::new(StdMutex::new(Vec::new())); + let entered = StdArc::new(AtomicBool::new(false)); + let release = StdArc::new(Notify::new()); + let task_manager = AsyncTaskManager::new().routing(NameRoutingPolicy::new([( + "echo", + RoutingDecision::Background, + )])); + let handle = task_manager.handle(); + let agent = Agent::builder() + .model(FakeAdapter) + .tool_executor(DelayedApprovalExecutor::new( + entered.clone(), + release.clone(), + )) + .task_manager(task_manager) + .cancellation(controller.handle()) + .observer(RecordingObserver { + events: events.clone(), + }) + .build() + .unwrap(); + + let mut driver = agent + .start(SessionConfig { + session_id: SessionId::new("session-cancel-delayed-background-approval"), + metadata: MetadataMap::new(), + cache: None, + }) + .await + .unwrap(); + + driver + .submit_input(vec![Item::text(ItemKind::User, "ping")]) + .unwrap(); + + match driver.next().await.unwrap() { + LoopStep::Interrupt(LoopInterrupt::AwaitingInput(_)) => {} + other => panic!("unexpected first step: {other:?}"), + } + + match wait_for_task_event(&handle).await { + TaskEvent::Started(snapshot) => assert_eq!(snapshot.tool_name, "echo"), + other => panic!("unexpected task event: {other:?}"), + } + + wait_until_entered(entered.as_ref()).await; + controller.interrupt(); + release.notify_waiters(); + wait_until_completed(&handle).await; + + match driver.next().await.unwrap() { + LoopStep::Finished(turn) => assert_eq!(turn.finish_reason, FinishReason::Cancelled), + other => panic!("cancelled background approval should finish cancelled, got {other:?}"), + } + + let events = events.lock().unwrap(); + assert!( + events + .iter() + .any(|event| matches!(event, AgentEvent::ApprovalResolved { approved: false })) + ); + assert!(events.iter().any(|event| matches!( + event, + AgentEvent::ToolResultReceived(result) + if result.call_id == ToolCallId::new("call-1") && result.is_error + ))); + } + + #[tokio::test] + async fn approved_foreground_task_waits_for_result_before_model_continuation() { + let entered = StdArc::new(AtomicBool::new(false)); + let release = StdArc::new(Notify::new()); + let approved_entered = StdArc::new(AtomicBool::new(false)); + let approved_release = StdArc::new(Notify::new()); + let route_count = StdArc::new(AtomicUsize::new(0)); + let routing_count = route_count.clone(); + let task_manager = AsyncTaskManager::new().routing(move |_request: &ToolRequest| { + if routing_count.fetch_add(1, Ordering::SeqCst) == 0 { + RoutingDecision::Background + } else { + RoutingDecision::Foreground + } + }); + let handle = task_manager.handle(); + let agent = Agent::builder() + .model(FakeAdapter) + .tool_executor( + DelayedApprovalExecutor::new(entered.clone(), release.clone()) + .blocking_after_approval(approved_entered.clone(), approved_release.clone()), + ) + .task_manager(task_manager) + .build() + .unwrap(); + let mut driver = agent + .start(SessionConfig::new("session-approved-foreground")) + .await + .unwrap(); + driver + .submit_input(vec![Item::text(ItemKind::User, "ping")]) + .unwrap(); + + assert!(matches!( + driver.next().await.unwrap(), + LoopStep::Interrupt(LoopInterrupt::AwaitingInput(_)) + )); + let task_turn = match wait_for_task_event(&handle).await { + TaskEvent::Started(snapshot) => snapshot.turn_id, + other => panic!("unexpected task event: {other:?}"), + }; + wait_until_entered(entered.as_ref()).await; + release.notify_one(); + wait_until_completed(&handle).await; + + let pending = match driver.next().await.unwrap() { + LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(pending)) => pending, + other => panic!("unexpected delayed approval step: {other:?}"), + }; + let presentation_turn = driver.lifecycle.active_turn.clone().unwrap(); + assert_ne!(presentation_turn, task_turn); + pending.approve(&mut driver).unwrap(); + + let info = { + let next = driver.next(); + tokio::pin!(next); + tokio::select! { + () = wait_until_entered(approved_entered.as_ref()) => {} + result = &mut next => { + panic!("model continued before approved foreground result: {result:?}") + } + } + assert!( + timeout(Duration::from_millis(10), &mut next).await.is_err(), + "model continued while approved foreground work was blocked" + ); + approved_release.notify_one(); + let step = timeout(Duration::from_secs(1), &mut next) + .await + .expect("approved foreground result was not delivered") + .unwrap(); + match step { + LoopStep::Interrupt(LoopInterrupt::AfterToolResult(info)) => info, + other => panic!("unexpected approved foreground step: {other:?}"), + } + }; + assert_eq!(info.turn_id, presentation_turn); + + let turn = match driver.next().await.unwrap() { + LoopStep::Finished(turn) => turn, + other => panic!("model did not continue after approved result: {other:?}"), + }; + assert_eq!(turn.finish_reason, FinishReason::Completed); + assert_eq!(turn.turn_id, presentation_turn); + assert!(driver.snapshot().transcript.iter().any(|item| { + item.kind == ItemKind::Notification + && item.parts.iter().any( + |part| matches!(part, Part::Text(text) if text.text.contains("approved-ok")), + ) + })); + } + + #[tokio::test] + async fn approved_foreground_then_detach_waits_and_keeps_one_placeholder() { + let entered = StdArc::new(AtomicBool::new(false)); + let release = StdArc::new(Notify::new()); + let approved_entered = StdArc::new(AtomicBool::new(false)); + let approved_release = StdArc::new(Notify::new()); + let route_count = StdArc::new(AtomicUsize::new(0)); + let routing_count = route_count.clone(); + let task_manager = AsyncTaskManager::new().routing(move |_request: &ToolRequest| { + if routing_count.fetch_add(1, Ordering::SeqCst) == 0 { + RoutingDecision::Background + } else { + RoutingDecision::ForegroundThenDetachAfter(Duration::from_millis(10)) + } + }); + let handle = task_manager.handle(); + let agent = Agent::builder() + .model(FakeAdapter) + .tool_executor( + DelayedApprovalExecutor::new(entered.clone(), release.clone()) + .blocking_after_approval(approved_entered.clone(), approved_release.clone()), + ) + .task_manager(task_manager) + .build() + .unwrap(); + let mut driver = agent + .start(SessionConfig::new("session-approved-foreground-detach")) + .await + .unwrap(); + driver + .submit_input(vec![Item::text(ItemKind::User, "ping")]) + .unwrap(); + + assert!(matches!( + driver.next().await.unwrap(), + LoopStep::Interrupt(LoopInterrupt::AwaitingInput(_)) + )); + let _ = wait_for_task_event(&handle).await; + wait_until_entered(entered.as_ref()).await; + release.notify_one(); + wait_until_completed(&handle).await; + + let pending = match driver.next().await.unwrap() { + LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(pending)) => pending, + other => panic!("unexpected delayed approval step: {other:?}"), + }; + let presentation_turn = driver.lifecycle.active_turn.clone().unwrap(); + pending.approve(&mut driver).unwrap(); + + let step = timeout(Duration::from_secs(1), driver.next()) + .await + .expect("approved task did not detach") + .unwrap(); + assert!(approved_entered.load(Ordering::SeqCst)); + match step { + LoopStep::Interrupt(LoopInterrupt::AfterToolResult(info)) => { + assert_eq!(info.turn_id, presentation_turn); + } + other => panic!("unexpected approved detach step: {other:?}"), + } + let placeholders = driver + .snapshot() + .transcript + .iter() + .filter(|item| item.kind == ItemKind::Tool) + .flat_map(|item| &item.parts) + .filter(|part| { + matches!( + part, + Part::ToolResult(result) if result.call_id == ToolCallId::new("call-1") + ) + }) + .count(); + assert_eq!(placeholders, 1, "detach appended a second tool result"); + + approved_release.notify_one(); + wait_until_completed(&handle).await; + let turn = match driver.next().await.unwrap() { + LoopStep::Finished(turn) => turn, + other => panic!("model did not continue after detached result: {other:?}"), + }; + assert_eq!(turn.finish_reason, FinishReason::Completed); + assert_eq!(turn.turn_id, presentation_turn); + let transcript = driver.snapshot().transcript; + assert_eq!( + transcript + .iter() + .filter(|item| item.kind == ItemKind::Tool) + .flat_map(|item| &item.parts) + .filter(|part| { + matches!( + part, + Part::ToolResult(result) + if result.call_id == ToolCallId::new("call-1") + ) + }) + .count(), + 1 + ); + assert!(transcript.iter().any(|item| { + item.kind == ItemKind::Notification + && item.parts.iter().any( + |part| matches!(part, Part::Text(text) if text.text.contains("approved-ok")), + ) + })); + } + + #[tokio::test] + async fn approving_detached_background_call_keeps_one_placeholder() { + let entered = StdArc::new(AtomicBool::new(false)); + let release = StdArc::new(Notify::new()); + let task_manager = AsyncTaskManager::new().routing(NameRoutingPolicy::new([( + "echo", + RoutingDecision::Background, + )])); + let handle = task_manager.handle(); + let agent = Agent::builder() + .model(FakeAdapter) + .tool_executor(DelayedApprovalExecutor::new( + entered.clone(), + release.clone(), + )) + .task_manager(task_manager) + .build() + .unwrap(); + let mut driver = agent + .start(SessionConfig::new("session-approve-detached-background")) + .await + .unwrap(); + driver + .submit_input(vec![Item::text(ItemKind::User, "ping")]) + .unwrap(); + + assert!(matches!( + driver.next().await.unwrap(), + LoopStep::Interrupt(LoopInterrupt::AwaitingInput(_)) + )); + let _ = wait_for_task_event(&handle).await; + wait_until_entered(entered.as_ref()).await; + release.notify_waiters(); + wait_until_completed(&handle).await; + + let pending = match driver.next().await.unwrap() { + LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(pending)) => pending, + other => panic!("unexpected delayed approval step: {other:?}"), + }; + pending.approve(&mut driver).unwrap(); + release.notify_one(); + let _ = driver.next().await.unwrap(); + + let placeholders = driver + .snapshot() + .transcript + .iter() + .flat_map(|item| &item.parts) + .filter(|part| { + matches!( + part, + Part::ToolResult(result) if result.call_id == ToolCallId::new("call-1") + ) + }) + .count(); + assert_eq!(placeholders, 1, "approval appended a second detach result"); + } + + #[tokio::test] + async fn failed_background_approval_cleanup_clears_queued_resume() { + let events = StdArc::new(StdMutex::new(Vec::new())); + let entered = StdArc::new(AtomicBool::new(false)); + let release = StdArc::new(Notify::new()); + let inner = AsyncTaskManager::new().routing(NameRoutingPolicy::new([( + "echo", + RoutingDecision::ForegroundThenDetachAfter(Duration::from_millis(10)), + )])); + let handle = inner.handle(); + let agent = Agent::builder() + .model(FakeAdapter) + .tool_executor(DelayedApprovalExecutor::new( + entered.clone(), + release.clone(), + )) + .task_manager(TestTaskManager::new(inner).fail_interrupt("cleanup failure")) + .observer(RecordingObserver { + events: events.clone(), + }) + .build() + .unwrap(); + let mut driver = agent + .start(SessionConfig::new( + "session-failed-detached-background-approval-cleanup", + )) + .await + .unwrap(); + driver + .submit_input(vec![Item::text(ItemKind::User, "ping")]) + .unwrap(); + + let old_turn_id = match driver.next().await.unwrap() { + LoopStep::Interrupt(LoopInterrupt::AfterToolResult(info)) => info.turn_id, + other => panic!("unexpected detach step: {other:?}"), + }; + assert_eq!(driver.pending_round_resume.as_ref(), Some(&old_turn_id)); + driver + .submit_input(vec![Item::text(ItemKind::User, "fresh input")]) + .unwrap(); + let _ = wait_for_task_event(&handle).await; + wait_until_entered(entered.as_ref()).await; + release.notify_waiters(); + wait_until_completed(&handle).await; + + assert!(matches!( + driver.next().await.unwrap(), + LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(_)) + )); + let error = driver.cancel_pending_approvals().await.unwrap_err(); + assert!(error.to_string().contains("cleanup failure")); + + assert!(driver.lifecycle.active_turn.is_none()); + assert!(driver.pending_round_resume.is_none()); + assert_eq!(driver.snapshot().pending_input.len(), 1); + let lifecycle = turn_lifecycle_events(&events.lock().unwrap()); + let [started, finished] = &lifecycle[lifecycle.len() - 2..] else { + panic!("missing terminal lifecycle events: {lifecycle:?}"); + }; + assert_eq!(started.0, finished.0); + assert_eq!(finished.1, Some(FinishReason::Error)); + validate_transcript_invariants(&driver.snapshot().transcript).unwrap(); + + let fresh_turn = match driver.next().await.unwrap() { + LoopStep::Finished(turn) => turn, + other => panic!("fresh input did not start a new turn: {other:?}"), + }; + assert_ne!(fresh_turn.turn_id, old_turn_id); + } + + #[tokio::test] + async fn fresh_input_runs_before_delayed_background_approval() { + let entered = StdArc::new(AtomicBool::new(false)); + let release = StdArc::new(Notify::new()); + let task_manager = AsyncTaskManager::new().routing(NameRoutingPolicy::new([( + "echo", + RoutingDecision::Background, + )])); + let handle = task_manager.handle(); + let agent = Agent::builder() + .model(FakeAdapter) + .tool_executor(DelayedApprovalExecutor::new( + entered.clone(), + release.clone(), + )) + .task_manager(task_manager) + .build() + .unwrap(); + let mut driver = agent + .start(SessionConfig::new( + "session-input-before-background-approval", + )) + .await + .unwrap(); + driver + .submit_input(vec![Item::text(ItemKind::User, "ping")]) + .unwrap(); + + assert!(matches!( + driver.next().await.unwrap(), + LoopStep::Interrupt(LoopInterrupt::AwaitingInput(_)) + )); + let _ = wait_for_task_event(&handle).await; wait_until_entered(entered.as_ref()).await; release.notify_waiters(); + wait_until_completed(&handle).await; - match wait_for_task_event(&handle).await { - TaskEvent::Completed(_, result) => { - assert_eq!(result.output, ToolOutput::Text("background-done".into())) - } - other => panic!("unexpected completion event: {other:?}"), - } - - let resumed = driver.next().await.unwrap(); - match resumed { - LoopStep::Finished(turn) => { - assert_eq!(turn.finish_reason, FinishReason::Completed); - match &turn.items[0].parts[0] { - Part::Text(text) => assert_eq!(text.text, "tool said: background-done"), - other => panic!("unexpected part after resume: {other:?}"), - } - } - other => panic!("unexpected resumed step: {other:?}"), - } + driver + .submit_input(vec![Item::text(ItemKind::User, "fresh input")]) + .unwrap(); + let fresh_turn = match driver.next().await.unwrap() { + LoopStep::Finished(turn) => turn.turn_id, + other => panic!("fresh input was not driven first: {other:?}"), + }; + assert!(driver.pending_approvals.is_empty()); + assert!(driver.snapshot().pending_input.is_empty()); + timeout(Duration::from_millis(100), driver.wait_for_loop_update()) + .await + .expect("collected background update did not wake the loop") + .unwrap(); - let events = events.lock().unwrap(); - let terminal_results: Vec<_> = events - .iter() - .filter_map(|event| match event { - AgentEvent::ToolResultReceived(result) - if result.call_id == ToolCallId::new("call-1") => - { - Some(result) - } - _ => None, - }) - .collect(); + let approval = match driver.next().await.unwrap() { + LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(approval)) => approval, + other => panic!("delayed approval was not presented separately: {other:?}"), + }; + let approval_turn = driver.lifecycle.active_turn.clone().unwrap(); + assert_ne!(fresh_turn, approval_turn); assert_eq!( - terminal_results.len(), + driver + .snapshot() + .transcript + .iter() + .filter(|item| { + item.kind == ItemKind::User + && item.parts.iter().any( + |part| matches!(part, Part::Text(text) if text.text == "fresh input"), + ) + }) + .count(), 1, - "background completion must emit one terminal result event per call: {events:?}" + "fresh input must not be replayed while presenting the approval" ); + approval.deny(&mut driver).unwrap(); } #[tokio::test] - async fn detached_tool_placeholder_is_progress_not_terminal_result() { - let events = StdArc::new(StdMutex::new(Vec::new())); + async fn delayed_background_approval_interrupts_originating_task_turn() { let entered = StdArc::new(AtomicBool::new(false)); let release = StdArc::new(Notify::new()); - let task_manager = AsyncTaskManager::new().routing(NameRoutingPolicy::new([( - "detaching-wait", - RoutingDecision::ForegroundThenDetachAfter(Duration::from_millis(10)), + let interrupted = StdArc::new(StdMutex::new(Vec::new())); + let inner = AsyncTaskManager::new().routing(NameRoutingPolicy::new([( + "echo", + RoutingDecision::Background, )])); - let handle = task_manager.handle(); - let tools = ToolRegistry::new().with(BlockingTool::new( - "detaching-wait", - entered.clone(), - release.clone(), - "detached-done", - )); + let handle = inner.handle(); let agent = Agent::builder() .model(FakeAdapter) - .add_tool_source(tools) - .permissions(AllowAllPermissions) - .task_manager(task_manager) - .observer(RecordingObserver { - events: events.clone(), - }) + .tool_executor(DelayedApprovalExecutor::new( + entered.clone(), + release.clone(), + )) + .task_manager(TestTaskManager::new(inner).record_interrupts(interrupted.clone())) .build() .unwrap(); - let mut driver = agent - .start(SessionConfig { - session_id: SessionId::new("session-detached-progress"), - metadata: MetadataMap::new(), - cache: None, - }) + .start(SessionConfig::new( + "session-background-approval-origin-turn", + )) .await .unwrap(); - driver .submit_input(vec![Item::text(ItemKind::User, "ping")]) .unwrap(); - match driver.next().await.unwrap() { - LoopStep::Interrupt(LoopInterrupt::AfterToolResult(_)) => {} - other => panic!("unexpected detach step: {other:?}"), - } - - match wait_for_task_event(&handle).await { - TaskEvent::Started(snapshot) => assert_eq!(snapshot.tool_name, "detaching-wait"), - other => panic!("unexpected task event: {other:?}"), - } - match wait_for_task_event(&handle).await { - TaskEvent::Detached(snapshot) => assert_eq!(snapshot.tool_name, "detaching-wait"), - other => panic!("unexpected detach event: {other:?}"), - } + assert!(matches!( + driver.next().await.unwrap(), + LoopStep::Interrupt(LoopInterrupt::AwaitingInput(_)) + )); + let _ = wait_for_task_event(&handle).await; wait_until_entered(entered.as_ref()).await; release.notify_waiters(); + wait_until_completed(&handle).await; + assert!(matches!( + driver.next().await.unwrap(), + LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(_)) + )); - match wait_for_task_event(&handle).await { - TaskEvent::Completed(_, result) => { - assert_eq!(result.output, ToolOutput::Text("detached-done".into())) - } - other => panic!("unexpected completion event: {other:?}"), - } - - match driver.next().await.unwrap() { - LoopStep::Finished(turn) => assert_eq!(turn.finish_reason, FinishReason::Completed), - other => panic!("unexpected resumed step: {other:?}"), - } - - let events = events.lock().unwrap(); - assert!(events.iter().any(|event| matches!( - event, - AgentEvent::ToolExecutionProgress(result) - if result.call_id == ToolCallId::new("call-1") && !result.is_error - ))); - let terminal_results: Vec<_> = events - .iter() - .filter_map(|event| match event { - AgentEvent::ToolResultReceived(result) - if result.call_id == ToolCallId::new("call-1") => - { - Some(result) - } - _ => None, - }) - .collect(); - assert_eq!( - terminal_results.len(), - 1, - "detached call must emit one terminal result event: {events:?}" - ); + let presentation_turn = driver.lifecycle.active_turn.clone().unwrap(); + let task_turn = driver + .pending_approvals + .values() + .next() + .unwrap() + .tool_request + .turn_id + .clone(); + assert_ne!(presentation_turn, task_turn); + assert!(matches!( + driver.cancel_pending_approvals().await.unwrap(), + Some(LoopStep::Finished(TurnResult { + finish_reason: FinishReason::Cancelled, + .. + })) + )); + assert_eq!(interrupted.lock().unwrap().as_slice(), &[task_turn]); + assert!(driver.lifecycle.active_turn.is_none()); } #[tokio::test] - async fn cancelled_background_approval_auto_resolves_when_drained() { - let controller = CancellationController::new(); + async fn approved_background_start_error_interrupts_originating_turn() { let events = StdArc::new(StdMutex::new(Vec::new())); let entered = StdArc::new(AtomicBool::new(false)); let release = StdArc::new(Notify::new()); - let task_manager = AsyncTaskManager::new().routing(NameRoutingPolicy::new([( + let interrupted = StdArc::new(StdMutex::new(Vec::new())); + let inner = AsyncTaskManager::new().routing(NameRoutingPolicy::new([( "echo", RoutingDecision::Background, )])); - let handle = task_manager.handle(); + let handle = inner.handle(); let agent = Agent::builder() .model(FakeAdapter) .tool_executor(DelayedApprovalExecutor::new( entered.clone(), release.clone(), )) - .task_manager(task_manager) - .cancellation(controller.handle()) + .task_manager( + TestTaskManager::new(inner) + .fail_approved_start("original approved start failure") + .record_interrupts(interrupted.clone()) + .fail_interrupt("cleanup failure"), + ) .observer(RecordingObserver { events: events.clone(), }) .build() .unwrap(); - let mut driver = agent - .start(SessionConfig { - session_id: SessionId::new("session-cancel-delayed-background-approval"), - metadata: MetadataMap::new(), - cache: None, - }) + .start(SessionConfig::new("session-approved-start-error")) .await .unwrap(); - driver .submit_input(vec![Item::text(ItemKind::User, "ping")]) .unwrap(); - match driver.next().await.unwrap() { - LoopStep::Interrupt(LoopInterrupt::AwaitingInput(_)) => {} - other => panic!("unexpected first step: {other:?}"), - } - - match wait_for_task_event(&handle).await { - TaskEvent::Started(snapshot) => assert_eq!(snapshot.tool_name, "echo"), - other => panic!("unexpected task event: {other:?}"), - } - + assert!(matches!( + driver.next().await.unwrap(), + LoopStep::Interrupt(LoopInterrupt::AwaitingInput(_)) + )); + let _ = wait_for_task_event(&handle).await; wait_until_entered(entered.as_ref()).await; - controller.interrupt(); release.notify_waiters(); wait_until_completed(&handle).await; + let pending = match driver.next().await.unwrap() { + LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(pending)) => pending, + other => panic!("unexpected delayed approval step: {other:?}"), + }; + let task_turn = driver + .pending_approvals + .values() + .next() + .unwrap() + .tool_request + .turn_id + .clone(); + let call_id = pending.request.call_id.clone().expect("approval call id"); + assert!(driver.detached_call_ids.contains(&call_id)); + pending.approve(&mut driver).unwrap(); - match driver.next().await.unwrap() { - LoopStep::Finished(turn) => assert_eq!(turn.finish_reason, FinishReason::Cancelled), - other => panic!("cancelled background approval should finish cancelled, got {other:?}"), - } - - let events = events.lock().unwrap(); + let error = driver.next().await.unwrap_err(); assert!( - events - .iter() - .any(|event| matches!(event, AgentEvent::ApprovalResolved { approved: false })) + error + .to_string() + .contains("original approved start failure") ); - assert!(events.iter().any(|event| matches!( + assert!(!error.to_string().contains("cleanup failure")); + assert_eq!(interrupted.lock().unwrap().as_slice(), &[task_turn]); + assert!(driver.lifecycle.active_turn.is_none()); + assert!(!driver.detached_call_ids.contains(&call_id)); + assert!(!driver.background_call_ids.contains(&call_id)); + assert!(!driver.tool_cancellations.contains_key(&call_id)); + assert!(events.lock().unwrap().iter().any(|event| matches!( event, AgentEvent::ToolResultReceived(result) - if result.call_id == ToolCallId::new("call-1") && result.is_error + if result.call_id == call_id && result.is_error ))); + validate_transcript_invariants(&driver.snapshot().transcript).unwrap(); } #[tokio::test] @@ -6287,6 +7635,10 @@ mod tests { .filter(|event| matches!(event, AgentEvent::ToolExecutionStarted(_))) .count(); assert_eq!(started, 1); + let lifecycle = turn_lifecycle_events(&events.lock().unwrap()); + assert_eq!(lifecycle.len(), 2); + assert_eq!(lifecycle[0].0, lifecycle[1].0); + assert_eq!(lifecycle[1].1, Some(FinishReason::Completed)); } #[tokio::test] @@ -6352,6 +7704,43 @@ mod tests { validate_transcript_invariants(&driver.snapshot().transcript).unwrap(); } + #[tokio::test] + async fn cancelling_sole_foreground_approval_for_call_finishes_turn() { + let events = StdArc::new(StdMutex::new(Vec::new())); + let agent = Agent::builder() + .model(FakeAdapter) + .add_tool_source(ToolRegistry::new().with(EchoTool::default())) + .permissions(ApproveFsReads) + .observer(RecordingObserver { + events: events.clone(), + }) + .build() + .unwrap(); + let mut driver = agent + .start(SessionConfig::new("session-cancel-foreground-approval-for")) + .await + .unwrap(); + driver + .submit_input(vec![Item::text(ItemKind::User, "ping")]) + .unwrap(); + + let call_id = match driver.next().await.unwrap() { + LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(pending)) => { + pending.request.call_id.expect("approval call id") + } + other => panic!("unexpected loop step: {other:?}"), + }; + driver.cancel_pending_approval_for(call_id).unwrap(); + + assert!(driver.lifecycle.active_turn.is_none()); + assert!(driver.pending_approvals.is_empty()); + validate_transcript_invariants(&driver.snapshot().transcript).unwrap(); + let lifecycle = turn_lifecycle_events(&events.lock().unwrap()); + assert_eq!(lifecycle.len(), 2, "{lifecycle:?}"); + assert_eq!(lifecycle[0].0, lifecycle[1].0); + assert_eq!(lifecycle[1].1, Some(FinishReason::Cancelled)); + } + #[tokio::test] async fn resolved_approval_runs_even_if_cancellation_also_fired() { let controller = CancellationController::new(); @@ -6529,13 +7918,63 @@ mod tests { } #[tokio::test] - async fn cancelling_all_pending_approvals_pairs_every_tool_use() { + async fn failed_pending_approval_cleanup_repairs_and_finishes_error() { + let events = StdArc::new(StdMutex::new(Vec::new())); + let agent = Agent::builder() + .model(FakeAdapter) + .add_tool_source(ToolRegistry::new().with(EchoTool::default())) + .permissions(ApproveFsReads) + .task_manager( + TestTaskManager::new(SimpleTaskManager::new()) + .fail_interrupt("interrupt cleanup failed"), + ) + .observer(RecordingObserver { + events: events.clone(), + }) + .build() + .unwrap(); + let mut driver = agent + .start(SessionConfig::new("session-failed-approval-cleanup")) + .await + .unwrap(); + driver + .submit_input(vec![Item::text(ItemKind::User, "ping")]) + .unwrap(); + + assert!(matches!( + driver.next().await.unwrap(), + LoopStep::Interrupt(LoopInterrupt::ApprovalRequest(_)) + )); + let error = driver.cancel_pending_approvals().await.unwrap_err(); + assert!(error.to_string().contains("interrupt cleanup failed")); + assert!(driver.lifecycle.active_turn.is_none()); + validate_transcript_invariants(&driver.snapshot().transcript).unwrap(); + + let events = events.lock().unwrap(); + assert!(events.iter().any(|event| matches!( + event, + AgentEvent::ToolResultReceived(result) + if result.call_id == ToolCallId::new("call-1") && result.is_error + ))); + assert!(events.iter().any(|event| matches!( + event, + AgentEvent::TurnFinished(turn) if turn.finish_reason == FinishReason::Error + ))); + } + + #[tokio::test] + async fn cancelling_all_pending_approvals_interrupts_every_originating_turn() { let events = StdArc::new(StdMutex::new(Vec::new())); + let interrupted = StdArc::new(StdMutex::new(Vec::new())); let tools = ToolRegistry::new().with(EchoTool::default()); let agent = Agent::builder() .model(DualApprovalAdapter) .add_tool_source(tools) .permissions(ApproveFsReads) + .task_manager( + TestTaskManager::new(SimpleTaskManager::new()) + .record_interrupts(interrupted.clone()), + ) .observer(RecordingObserver { events: events.clone(), }) @@ -6578,6 +8017,21 @@ mod tests { } } + let first_origin = driver + .pending_approvals + .get(&ToolCallId::new("call-1")) + .unwrap() + .tool_request + .turn_id + .clone(); + let second_origin = agentkit_core::TurnId::new("second-originating-turn"); + driver + .pending_approvals + .get_mut(&ToolCallId::new("call-2")) + .unwrap() + .tool_request + .turn_id = second_origin.clone(); + match driver.cancel_pending_approvals().await.unwrap() { Some(LoopStep::Finished(turn)) => { assert_eq!(turn.finish_reason, FinishReason::Cancelled); @@ -6585,6 +8039,13 @@ mod tests { other => panic!("unexpected cancellation result: {other:?}"), } validate_transcript_invariants(&driver.snapshot().transcript).unwrap(); + let interrupted = interrupted + .lock() + .unwrap() + .iter() + .cloned() + .collect::>(); + assert_eq!(interrupted, HashSet::from([first_origin, second_origin])); let events = events.lock().unwrap(); let cancelled = events diff --git a/crates/agentkit-provider-anthropic/Cargo.toml b/crates/agentkit-provider-anthropic/Cargo.toml index 7707fca..b18fc70 100644 --- a/crates/agentkit-provider-anthropic/Cargo.toml +++ b/crates/agentkit-provider-anthropic/Cargo.toml @@ -7,7 +7,7 @@ edition.workspace = true license.workspace = true repository.workspace = true rust-version.workspace = true -version = "0.10.7" +version = "0.10.8" [dependencies] agentkit-core = { version = "0.10.5", path = "../agentkit-core" } diff --git a/crates/agentkit-provider-anthropic/src/request.rs b/crates/agentkit-provider-anthropic/src/request.rs index de552b3..5cf6c4d 100644 --- a/crates/agentkit-provider-anthropic/src/request.rs +++ b/crates/agentkit-provider-anthropic/src/request.rs @@ -727,7 +727,13 @@ mod tests { ))], ), Item::text(ItemKind::Assistant, "ok, kicked it off"), - Item::notification("Slack install completed: ok"), + Item::new( + ItemKind::Notification, + vec![ + Part::text("Slack install completed: ok"), + Part::structured(json!({ "typed": true })), + ], + ), ]; let body = build_request_body(&cfg(), &base_request(transcript)).unwrap(); let messages = body["messages"].as_array().unwrap(); @@ -737,6 +743,7 @@ mod tests { let text = messages[3]["content"][0]["text"].as_str().unwrap(); assert!(text.starts_with("")); assert!(text.contains("Slack install completed: ok")); + assert!(text.contains("\"typed\": true")); assert!(text.ends_with("")); // Critical: must NOT leak into top-level system blocks. let no_leak = body diff --git a/crates/agentkit-task-manager/Cargo.toml b/crates/agentkit-task-manager/Cargo.toml index 11e7bc3..198826a 100644 --- a/crates/agentkit-task-manager/Cargo.toml +++ b/crates/agentkit-task-manager/Cargo.toml @@ -4,7 +4,7 @@ homepage.workspace = true name = "agentkit-task-manager" readme = "README.md" repository.workspace = true -version = "0.10.6" +version = "0.10.7" edition.workspace = true license.workspace = true rust-version.workspace = true diff --git a/crates/agentkit-task-manager/src/lib.rs b/crates/agentkit-task-manager/src/lib.rs index 0e5572a..5705032 100644 --- a/crates/agentkit-task-manager/src/lib.rs +++ b/crates/agentkit-task-manager/src/lib.rs @@ -194,6 +194,14 @@ pub trait TaskManager: Send + Sync { async fn take_pending_loop_updates(&self) -> Result; + /// Wait until an update is available for delivery back into the agent loop. + /// + /// The default never resolves because custom managers that do not produce + /// out-of-band loop updates have nothing to wake a host for. + async fn wait_for_loop_update(&self) -> Result<(), TaskManagerError> { + std::future::pending().await + } + async fn on_turn_interrupted(&self, turn_id: &TurnId) -> Result<(), TaskManagerError>; fn handle(&self) -> TaskManagerHandle; @@ -753,6 +761,23 @@ impl TaskManager for AsyncTaskManager { }) } + async fn wait_for_loop_update(&self) -> Result<(), TaskManagerError> { + loop { + // Register before checking state so an update cannot land between + // the empty check and awaiting the notification. + let notified = self.inner.notify.notified(); + tokio::pin!(notified); + notified.as_mut().enable(); + { + let state = self.inner.state.lock().await; + if !state.pending_loop_updates.is_empty() { + return Ok(()); + } + } + notified.await; + } + } + async fn on_turn_interrupted(&self, turn_id: &TurnId) -> Result<(), TaskManagerError> { self.inner.interrupt_turn(turn_id).await; Ok(()) @@ -1595,6 +1620,67 @@ mod tests { assert!(handle.drain_ready_items().await.is_empty()); } + #[tokio::test] + async fn wait_for_loop_update_wakes_without_consuming_host_events() { + let release = StdArc::new(Notify::new()); + let entered = StdArc::new(AtomicBool::new(false)); + let executor: Arc = Arc::new(TestExecutor::new([( + "background", + TestBehavior::Block { + entered: entered.clone(), + release: release.clone(), + output: "wake-done", + }, + )])); + let manager = AsyncTaskManager::new().routing(NameRoutingPolicy::new([( + "background", + RoutingDecision::Background, + )])); + let handle = manager.handle(); + let request = make_request("background", "turn-1", "call-1"); + + manager + .start_task( + TaskLaunchRequest { + task_id: None, + request: request.clone(), + kind: TaskLaunchKind::Plain, + }, + make_context(executor, &request.turn_id, None), + ) + .await + .unwrap(); + wait_until_entered(entered.as_ref()).await; + + let waiting = manager.wait_for_loop_update(); + tokio::pin!(waiting); + assert!( + timeout(Duration::from_millis(20), &mut waiting) + .await + .is_err() + ); + release.notify_waiters(); + timeout(Duration::from_secs(1), &mut waiting) + .await + .expect("loop update wake timed out") + .unwrap(); + + assert!(matches!(next_event(&handle).await, TaskEvent::Started(_))); + assert!(matches!( + next_event(&handle).await, + TaskEvent::Completed(_, _) + )); + assert_eq!( + manager + .take_pending_loop_updates() + .await + .unwrap() + .resolutions + .len(), + 1 + ); + } + #[tokio::test] async fn wait_for_idle_returns_after_loop_updates_are_queued() { let release = StdArc::new(Notify::new()); diff --git a/crates/agentkit/Cargo.toml b/crates/agentkit/Cargo.toml index 234fda3..f6f296b 100644 --- a/crates/agentkit/Cargo.toml +++ b/crates/agentkit/Cargo.toml @@ -4,22 +4,22 @@ homepage.workspace = true name = "agentkit" readme = "README.md" repository.workspace = true -version = "0.10.8" +version = "0.10.11" edition.workspace = true license.workspace = true rust-version.workspace = true [dependencies] agentkit-capabilities = { version = "0.10.5", path = "../agentkit-capabilities", optional = true } -agentkit-acp = { version = "0.10.5", path = "../agentkit-acp", optional = true } +agentkit-acp = { version = "0.10.11", path = "../agentkit-acp", optional = true } agentkit-compaction = { version = "0.10.5", path = "../agentkit-compaction", optional = true } agentkit-context = { version = "0.10.5", path = "../agentkit-context", optional = true } agentkit-core = { version = "0.10.5", path = "../agentkit-core", optional = true } -agentkit-loop = { version = "0.10.7", path = "../agentkit-loop", optional = true } +agentkit-loop = { version = "0.10.10", path = "../agentkit-loop", optional = true } agentkit-mcp = { version = "0.10.5", path = "../agentkit-mcp", optional = true } agentkit-plugins = { version = "0.10.6", path = "../agentkit-plugins", optional = true } -agentkit-adapter-completions = { version = "0.10.6", path = "../agentkit-adapter-completions", optional = true } -agentkit-provider-anthropic = { version = "0.10.7", path = "../agentkit-provider-anthropic", optional = true } +agentkit-adapter-completions = { version = "0.10.7", path = "../agentkit-adapter-completions", optional = true } +agentkit-provider-anthropic = { version = "0.10.8", path = "../agentkit-provider-anthropic", optional = true } agentkit-provider-baseten = { version = "0.10.5", path = "../agentkit-provider-baseten", optional = true } agentkit-provider-cerebras = { version = "0.10.6", path = "../agentkit-provider-cerebras", optional = true } agentkit-provider-groq = { version = "0.10.5", path = "../agentkit-provider-groq", optional = true } @@ -29,7 +29,7 @@ agentkit-provider-openai = { version = "0.10.5", path = "../agentkit-provider-op agentkit-provider-openrouter = { version = "0.10.7", path = "../agentkit-provider-openrouter", optional = true } agentkit-provider-vllm = { version = "0.10.5", path = "../agentkit-provider-vllm", optional = true } agentkit-reporting = { version = "0.10.5", path = "../agentkit-reporting", optional = true } -agentkit-task-manager = { version = "0.10.5", path = "../agentkit-task-manager", optional = true } +agentkit-task-manager = { version = "0.10.7", path = "../agentkit-task-manager", optional = true } agentkit-tool-compose = { version = "0.10.5", path = "../agentkit-tool-compose", optional = true } agentkit-tool-fs = { version = "0.10.5", path = "../agentkit-tool-fs", optional = true } agentkit-tool-shell = { version = "0.10.5", path = "../agentkit-tool-shell", optional = true }