From 55dd84775b419032b7fb06f81c85bcdfa0b1c2f0 Mon Sep 17 00:00:00 2001 From: Lifei Zhou Date: Wed, 17 Jun 2026 15:47:50 +1000 Subject: [PATCH] fixed acp cancel race condition (#9804) --- crates/goose/src/acp/server.rs | 185 +++++++++++------- crates/goose/src/acp/server/dispatch.rs | 2 +- crates/goose/src/acp/server/extensions.rs | 4 +- crates/goose/src/acp/server/load_session.rs | 1 + .../goose/src/acp/server/manage_sessions.rs | 2 +- crates/goose/src/acp/server/resources.rs | 2 +- crates/goose/src/acp/server/tools.rs | 4 +- 7 files changed, 120 insertions(+), 80 deletions(-) diff --git a/crates/goose/src/acp/server.rs b/crates/goose/src/acp/server.rs index 3ae3aa713..15da2943c 100644 --- a/crates/goose/src/acp/server.rs +++ b/crates/goose/src/acp/server.rs @@ -169,8 +169,11 @@ struct GooseAcpSession { /// Tool_call_ids of chains that have already had a summary task fired. /// Idempotence guard so we summarize each chain at most once. summarized_chains: HashSet, - cancel_token: Option, - active_run_id: Option, +} + +struct ActivePromptRun { + run_id: String, + cancel_token: CancellationToken, } /// A run of consecutive ToolRequest blocks within one assistant message, @@ -197,6 +200,8 @@ pub struct GooseAcpAgentOptions { pub struct GooseAcpAgent { sessions: Arc>>, + active_prompt_runs: Arc>>, + closed_session_ids: Arc>>, agent_manager: Arc, provider_factory: AcpProviderFactory, builtins: Vec, @@ -877,6 +882,8 @@ impl GooseAcpAgent { Ok(Self { sessions: Arc::new(Mutex::new(HashMap::new())), + active_prompt_runs: Arc::new(Mutex::new(HashMap::new())), + closed_session_ids: Arc::new(Mutex::new(HashSet::new())), agent_manager, provider_factory: options.provider_factory, builtins: options.builtins, @@ -1152,8 +1159,6 @@ impl GooseAcpAgent { chain_membership: HashMap::new(), responded_tool_ids: HashSet::new(), summarized_chains: HashSet::new(), - cancel_token: None, - active_run_id: None, }; self.sessions.lock().await.insert(session_id, acp_session); } @@ -2123,19 +2128,21 @@ impl GooseAcpAgent { self.handle_new_session(cx, args).await } - /// Look up the session's agent. Optionally sets a cancellation token on - /// the session (needed by `on_prompt`). + /// Look up the session's agent. async fn get_session_agent( &self, session_id: &str, - cancel_token: Option, ) -> Result, agent_client_protocol::Error> { + if self.closed_session_ids.lock().await.contains(session_id) { + return Err(agent_client_protocol::Error::resource_not_found(Some( + session_id.to_string(), + )) + .data(format!("Session not found: {}", session_id))); + } + { - let mut sessions = self.sessions.lock().await; - if let Some(session) = sessions.get_mut(session_id) { - if let Some(token) = cancel_token { - session.cancel_token = Some(token); - } + let sessions = self.sessions.lock().await; + if let Some(session) = sessions.get(session_id) { return Ok(session.agent.clone()); } } @@ -2155,13 +2162,6 @@ impl GooseAcpAgent { let (agent, _) = self .activate_acp_session(cx, &session, HashMap::new()) .await?; - - if let Some(token) = cancel_token { - let mut sessions = self.sessions.lock().await; - if let Some(session) = sessions.get_mut(session_id) { - session.cancel_token = Some(token); - } - } Ok(agent) } @@ -2171,37 +2171,59 @@ impl GooseAcpAgent { run_id: String, cancel_token: CancellationToken, ) -> Result<(), agent_client_protocol::Error> { - let mut sessions = self.sessions.lock().await; - let session = sessions.get_mut(session_id).ok_or_else(|| { - agent_client_protocol::Error::resource_not_found(Some(session_id.to_string())) - .data(format!("Session not found: {}", session_id)) - })?; + if self.closed_session_ids.lock().await.contains(session_id) { + return Err(agent_client_protocol::Error::resource_not_found(Some( + session_id.to_string(), + )) + .data(format!("Session not found: {}", session_id))); + } - if let Some(active_run_id) = &session.active_run_id { + let mut active_prompt_runs = self.active_prompt_runs.lock().await; + if let Some(active_run) = active_prompt_runs.get(session_id) { return Err(agent_client_protocol::Error::invalid_params().data(format!( - "session already has active run `{active_run_id}`; use _goose/unstable/session/steer" + "session already has active run `{}`; use _goose/unstable/session/steer", + active_run.run_id.as_str() ))); } - session.cancel_token = Some(cancel_token); - session.active_run_id = Some(run_id); + active_prompt_runs.insert( + session_id.to_string(), + ActivePromptRun { + run_id, + cancel_token, + }, + ); Ok(()) } async fn clear_active_run(&self, session_id: &str, run_id: &str) { - let agent = { - let mut sessions = self.sessions.lock().await; - let Some(session) = sessions.get_mut(session_id) else { + { + let mut active_prompt_runs = self.active_prompt_runs.lock().await; + let Some(active_run) = active_prompt_runs.get(session_id) else { return; }; - if session.active_run_id.as_deref() != Some(run_id) { + + if active_run.run_id != run_id { return; } - session.cancel_token = None; - session.active_run_id = None; - session.agent.clone() + + active_prompt_runs.remove(session_id); + } + + let agent = { + let sessions = self.sessions.lock().await; + sessions + .get(session_id) + .map(|session| session.agent.clone()) }; - agent.discard_pending_steers(session_id).await; + if let Some(agent) = agent { + agent.discard_pending_steers(session_id).await; + } + + if self.closed_session_ids.lock().await.contains(session_id) { + self.sessions.lock().await.remove(session_id); + let _ = self.agent_manager.remove_session(session_id).await; + } } async fn require_active_run( @@ -2214,26 +2236,23 @@ impl GooseAcpAgent { .data("expectedRunId must not be empty")); } - let sessions = self.sessions.lock().await; - let session = sessions.get(session_id).ok_or_else(|| { - agent_client_protocol::Error::resource_not_found(Some(session_id.to_string())) - .data(format!("Session not found: {}", session_id)) - })?; - let active_run_id = session.active_run_id.as_ref().ok_or_else(|| { + let active_prompt_runs = self.active_prompt_runs.lock().await; + let active_run = active_prompt_runs.get(session_id).ok_or_else(|| { agent_client_protocol::Error::invalid_params().data("no active run to steer") })?; - if active_run_id != expected_run_id { + if active_run.run_id != expected_run_id { return Err( agent_client_protocol::Error::invalid_params().data(serde_json::json!({ "message": format!( - "expected active run id `{expected_run_id}` but found `{active_run_id}`" + "expected active run id `{expected_run_id}` but found `{}`", + active_run.run_id.as_str() ), "expectedRunId": expected_run_id, - "actualRunId": active_run_id, + "actualRunId": active_run.run_id.as_str(), })), ); } - Ok(active_run_id.clone()) + Ok(active_run.run_id.clone()) } fn active_run_meta(active_run_id: Option<&str>) -> Meta { @@ -2345,20 +2364,26 @@ impl GooseAcpAgent { let cancel_token = CancellationToken::new(); self.start_active_run(&session_id, run_id.clone(), cancel_token.clone()) .await?; + + let agent = match self.get_session_agent(&session_id).await { + Ok(agent) => agent, + Err(error) => { + self.clear_active_run(&session_id, &run_id).await; + return Err(error); + } + }; + + if cancel_token.is_cancelled() { + self.clear_active_run(&session_id, &run_id).await; + Self::send_active_run_update(cx, &args.session_id, None)?; + return Ok(PromptResponse::new(StopReason::Cancelled)); + } + if let Err(error) = Self::send_active_run_update(cx, &args.session_id, Some(&run_id)) { self.clear_active_run(&session_id, &run_id).await; return Err(error); } - let agent = match self.get_session_agent(&session_id, None).await { - Ok(agent) => agent, - Err(error) => { - self.clear_active_run(&session_id, &run_id).await; - let _ = Self::send_active_run_update(cx, &args.session_id, None); - return Err(error); - } - }; - let user_message = Self::convert_acp_prompt_to_message(&args.prompt); let message_text = user_message.as_concat_text(); @@ -2590,7 +2615,7 @@ impl GooseAcpAgent { self.require_active_run(&req.session_id, &req.expected_run_id) .await?; - let agent = self.get_session_agent(&req.session_id, None).await?; + let agent = self.get_session_agent(&req.session_id).await?; let active_run_id = self .require_active_run(&req.session_id, &req.expected_run_id) .await?; @@ -2627,14 +2652,17 @@ impl GooseAcpAgent { debug!(?args, "cancel request"); let session_id = args.session_id.0.to_string(); - let mut sessions = self.sessions.lock().await; + let token = { + let active_prompt_runs = self.active_prompt_runs.lock().await; + active_prompt_runs + .get(&session_id) + .map(|active_run| active_run.cancel_token.clone()) + }; - if let Some(session) = sessions.get_mut(&session_id) { - if let Some(ref token) = session.cancel_token { - info!(session_id = %session_id, "prompt cancelled"); - token.cancel(); - } - } else { + if let Some(token) = token { + info!(session_id = %session_id, "prompt cancelled"); + token.cancel(); + } else if !self.sessions.lock().await.contains_key(&session_id) { warn!(session_id = %session_id, "cancel request for unknown session"); } @@ -2646,7 +2674,7 @@ impl GooseAcpAgent { session_id: &str, model_id: &str, ) -> Result { - let agent = self.get_session_agent(session_id, None).await?; + let agent = self.get_session_agent(session_id).await?; let current_provider = agent .provider() .await @@ -2675,7 +2703,7 @@ impl GooseAcpAgent { .get_session(&session_id.0, false) .await .internal_err()?; - let agent = self.get_session_agent(&session_id.0, None).await?; + let agent = self.get_session_agent(&session_id.0).await?; let provider = agent .provider() .await @@ -2720,7 +2748,7 @@ impl GooseAcpAgent { .data(format!("Invalid mode: {}", mode_id)) })?; - let agent = self.get_session_agent(session_id, None).await?; + let agent = self.get_session_agent(session_id).await?; agent .update_goose_mode(mode, session_id) .await @@ -2742,7 +2770,7 @@ impl GooseAcpAgent { agent_client_protocol::Error::invalid_params() .data(format!("Invalid thinking effort: {}", effort_id)) })?; - let agent = self.get_session_agent(session_id, None).await?; + let agent = self.get_session_agent(session_id).await?; agent .update_thinking_effort(session_id, effort) .await @@ -2760,7 +2788,7 @@ impl GooseAcpAgent { request_params: Option>, ) -> Result<(), agent_client_protocol::Error> { let config = self.config()?; - let agent = self.get_session_agent(session_id, None).await?; + let agent = self.get_session_agent(session_id).await?; let current_provider = agent .provider() .await @@ -2817,12 +2845,23 @@ impl GooseAcpAgent { &self, session_id: &str, ) -> Result { - let mut sessions = self.sessions.lock().await; - if let Some(session) = sessions.get(session_id) { - if let Some(ref token) = session.cancel_token { - token.cancel(); - } + self.closed_session_ids + .lock() + .await + .insert(session_id.to_string()); + + let active_run_token = { + let active_prompt_runs = self.active_prompt_runs.lock().await; + active_prompt_runs + .get(session_id) + .map(|active_run| active_run.cancel_token.clone()) + }; + + if let Some(token) = active_run_token { + token.cancel(); } + + let mut sessions = self.sessions.lock().await; sessions.remove(session_id); drop(sessions); diff --git a/crates/goose/src/acp/server/dispatch.rs b/crates/goose/src/acp/server/dispatch.rs index a38492188..f01285df8 100644 --- a/crates/goose/src/acp/server/dispatch.rs +++ b/crates/goose/src/acp/server/dispatch.rs @@ -175,7 +175,7 @@ impl HandleDispatchFrom for GooseAcpHandler { let provider_result: Result> = AssertUnwindSafe(async { let session_agent = - agent_bg.get_session_agent(&session_id_bg.0, None).await?; + agent_bg.get_session_agent(&session_id_bg.0).await?; let provider = session_agent .provider() .await diff --git a/crates/goose/src/acp/server/extensions.rs b/crates/goose/src/acp/server/extensions.rs index ff58cb21d..37afe56f2 100644 --- a/crates/goose/src/acp/server/extensions.rs +++ b/crates/goose/src/acp/server/extensions.rs @@ -12,7 +12,7 @@ impl GooseAcpAgent { let config: ExtensionConfig = serde_json::from_value(req.config).map_err(|e| { agent_client_protocol::Error::invalid_params().data(format!("bad config: {e}")) })?; - let agent = self.get_session_agent(&req.session_id, None).await?; + let agent = self.get_session_agent(&req.session_id).await?; agent .add_extension(config, session_id) .await @@ -25,7 +25,7 @@ impl GooseAcpAgent { req: RemoveExtensionRequest, ) -> Result { let session_id = &req.session_id; - let agent = self.get_session_agent(&req.session_id, None).await?; + let agent = self.get_session_agent(&req.session_id).await?; agent .remove_extension(&req.name, session_id) .await diff --git a/crates/goose/src/acp/server/load_session.rs b/crates/goose/src/acp/server/load_session.rs index b22ccabe6..9600ec431 100644 --- a/crates/goose/src/acp/server/load_session.rs +++ b/crates/goose/src/acp/server/load_session.rs @@ -246,6 +246,7 @@ impl GooseAcpAgent { ms = t_start.elapsed().as_millis() as u64, "perf: load_session_refactor done" ); + self.closed_session_ids.lock().await.remove(&session_id_str); Ok(response) } } diff --git a/crates/goose/src/acp/server/manage_sessions.rs b/crates/goose/src/acp/server/manage_sessions.rs index a3382c208..6ba9309a6 100644 --- a/crates/goose/src/acp/server/manage_sessions.rs +++ b/crates/goose/src/acp/server/manage_sessions.rs @@ -42,7 +42,7 @@ impl GooseAcpAgent { ); } - let agent = self.get_session_agent(session_id, None).await?; + let agent = self.get_session_agent(session_id).await?; match req.mode { SessionSystemPromptMode::Set => { if req.text.trim().is_empty() { diff --git a/crates/goose/src/acp/server/resources.rs b/crates/goose/src/acp/server/resources.rs index ba2f07bff..8022b7c2c 100644 --- a/crates/goose/src/acp/server/resources.rs +++ b/crates/goose/src/acp/server/resources.rs @@ -6,7 +6,7 @@ impl GooseAcpAgent { req: ReadResourceRequest, ) -> Result { let session_id = &req.session_id; - let agent = self.get_session_agent(&req.session_id, None).await?; + let agent = self.get_session_agent(&req.session_id).await?; let cancel_token = CancellationToken::new(); let result = agent .extension_manager diff --git a/crates/goose/src/acp/server/tools.rs b/crates/goose/src/acp/server/tools.rs index 6204f327b..24b58acac 100644 --- a/crates/goose/src/acp/server/tools.rs +++ b/crates/goose/src/acp/server/tools.rs @@ -8,7 +8,7 @@ impl GooseAcpAgent { req: GetToolsRequest, ) -> Result { let session_id = &req.session_id; - let agent = self.get_session_agent(&req.session_id, None).await?; + let agent = self.get_session_agent(&req.session_id).await?; let tools = agent.list_tools(session_id, None).await; let tools_json = tools .into_iter() @@ -23,7 +23,7 @@ impl GooseAcpAgent { req: GooseToolCallRequest, ) -> Result { let session_id = &req.session_id; - let agent = self.get_session_agent(&req.session_id, None).await?; + let agent = self.get_session_agent(&req.session_id).await?; let tools = agent.list_tools(session_id, None).await; let Some(tool) = tools.iter().find(|t| *t.name == req.name) else {