From 5e7430fefaadd99792e6db2ae212e4227212487b Mon Sep 17 00:00:00 2001 From: John Matthew Tennant Date: Mon, 10 Aug 2026 06:08:34 -0400 Subject: [PATCH] fix(acp): resume provider-native sessions (#10379) Co-authored-by: John Tennant Co-authored-by: Douwe M Osinga --- crates/goose-provider-types/src/base.rs | 8 + .../goose-provider-types/src/conversation.rs | 1 + .../src/conversation/message.rs | 2 + crates/goose/src/acp/provider.rs | 178 ++++++++++++++++-- crates/goose/src/agents/agent.rs | 53 +++++- crates/goose/src/agents/mod.rs | 13 ++ .../goose/src/agents/state_machine/ops_llm.rs | 27 ++- ui/desktop/src/types/message.ts | 1 + 8 files changed, 257 insertions(+), 26 deletions(-) diff --git a/crates/goose-provider-types/src/base.rs b/crates/goose-provider-types/src/base.rs index beee620a4..f8cad1708 100644 --- a/crates/goose-provider-types/src/base.rs +++ b/crates/goose-provider-types/src/base.rs @@ -426,6 +426,14 @@ pub trait Provider: Send + Sync { /// Get the name of this provider instance fn get_name(&self) -> &str; + fn provider_session_id(&self) -> Option { + None + } + + async fn resume(&self, _session_id: &str) -> Result<(), ProviderError> { + Ok(()) + } + /// Primary streaming method that all providers must implement. async fn stream( &self, diff --git a/crates/goose-provider-types/src/conversation.rs b/crates/goose-provider-types/src/conversation.rs index 53572b6ee..1ba7cea67 100644 --- a/crates/goose-provider-types/src/conversation.rs +++ b/crates/goose-provider-types/src/conversation.rs @@ -1856,6 +1856,7 @@ mod tests { provider: "test-provider".to_string(), requested_model: "test-model".to_string(), resolved_model: None, + provider_session_id: None, }; let mut limited = Message::assistant() .with_id("turn-1") diff --git a/crates/goose-provider-types/src/conversation/message.rs b/crates/goose-provider-types/src/conversation/message.rs index 65b062810..8cc004890 100644 --- a/crates/goose-provider-types/src/conversation/message.rs +++ b/crates/goose-provider-types/src/conversation/message.rs @@ -673,6 +673,8 @@ pub struct InferenceMetadata { pub requested_model: String, #[serde(skip_serializing_if = "Option::is_none")] pub resolved_model: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub provider_session_id: Option, } #[derive(Clone, PartialEq, Serialize, Deserialize, Debug, Default)] diff --git a/crates/goose/src/acp/provider.rs b/crates/goose/src/acp/provider.rs index f2b0b4e0d..52a18a4c6 100644 --- a/crates/goose/src/acp/provider.rs +++ b/crates/goose/src/acp/provider.rs @@ -1,8 +1,8 @@ use agent_client_protocol::schema::v1::{ Annotations as AcpAnnotations, ClientCapabilities, CloseSessionRequest, ContentBlock, ContentChunk, EnvVariable, HttpHeader, ImageContent, InitializeRequest, InitializeResponse, - McpCapabilities, McpServer, McpServerHttp, McpServerStdio, NewSessionRequest, - NewSessionResponse, PromptRequest, PromptResponse, RequestPermissionOutcome, + LoadSessionRequest, McpCapabilities, McpServer, McpServerHttp, McpServerStdio, + NewSessionRequest, NewSessionResponse, PromptRequest, PromptResponse, RequestPermissionOutcome, RequestPermissionRequest, RequestPermissionResponse, Role as AcpRole, SessionConfigKind, SessionConfigOption, SessionConfigOptionCategory, SessionConfigSelectOptions, SessionId, SessionModeState, SessionNotification, SessionUpdate, SetSessionConfigOptionRequest, @@ -69,6 +69,13 @@ enum ClientRequest { NewSession { response_tx: oneshot::Sender>, }, + LoadSession { + session_id: SessionId, + response_tx: oneshot::Sender>, + }, + CloseSession { + session_id: SessionId, + }, SetMode { session_id: SessionId, mode_id: String, @@ -171,7 +178,7 @@ pub struct AcpProvider { goose_mode: Arc>, mode_mapping: HashMap>, - session: AcpSession, + session: Mutex, pending_confirmations: Arc>>>, @@ -304,7 +311,7 @@ impl AcpProvider { name, goose_mode: goose_mode_shared, mode_mapping, - session, + session: Mutex::new(session), pending_confirmations: Arc::new(TokioMutex::new(HashMap::new())), pending_tool_updates, handoff_context_sent: AtomicBool::new(false), @@ -318,7 +325,25 @@ impl AcpProvider { } fn acp_session_id(&self) -> SessionId { - self.session.id.clone() + self.session.lock().unwrap().id.clone() + } + + async fn load_session(&self, session_id: SessionId) -> Result { + let (response_tx, response_rx) = oneshot::channel(); + self.tx + .as_ref() + .unwrap() + .send(ClientRequest::LoadSession { + session_id, + response_tx, + }) + .await + .context("ACP client is unavailable")?; + let response = response_rx.await.context("ACP session load cancelled")??; + Ok(AcpSession { + id: response.session_id.clone(), + response, + }) } pub(crate) async fn send_set_mode(&self, _goose_id: &str, mode_id: String) -> Result<()> { @@ -413,6 +438,8 @@ impl AcpProvider { fn session_has_config_option(&self, category: SessionConfigOptionCategory) -> bool { self.session + .lock() + .unwrap() .response .config_options .as_ref() @@ -441,6 +468,33 @@ impl Provider for AcpProvider { &self.name } + fn provider_session_id(&self) -> Option { + Some(self.acp_session_id().to_string()) + } + + async fn resume(&self, session_id: &str) -> Result<(), ProviderError> { + if self.acp_session_id().0.as_ref() == session_id { + return Ok(()); + } + + let previous_session_id = self.acp_session_id(); + let loaded = self + .load_session(SessionId::new(session_id)) + .await + .map_err(|error| ProviderError::RequestFailed(error.to_string()))?; + *self.session.lock().unwrap() = loaded; + self.handoff_context_sent.store(true, Ordering::Release); + let _ = self + .tx + .as_ref() + .unwrap() + .send(ClientRequest::CloseSession { + session_id: previous_session_id, + }) + .await; + Ok(()) + } + async fn get_context_limit(&self, model_config: &ModelConfig) -> Result { let size = self.context_size.load(Ordering::Relaxed); if size > 0 { @@ -451,8 +505,9 @@ impl Provider for AcpProvider { async fn update_mode(&self, session_id: &str, mode: GooseMode) -> Result<(), ProviderError> { if let Some(candidates) = self.mode_mapping.get(&mode) { - let mode_str = select_mode_id(candidates, self.session.response.modes.as_ref()) - .ok_or_else(|| { + let session = self.session.lock().unwrap().clone(); + let mode_str = + select_mode_id(candidates, session.response.modes.as_ref()).ok_or_else(|| { ProviderError::RequestFailed(format!( "None of the mode ids [{}] are offered by the agent", candidates.join(", ") @@ -708,7 +763,8 @@ impl Provider for AcpProvider { } async fn fetch_supported_models(&self) -> Result, ProviderError> { - let (_, available) = resolve_model_info(&self.name, &self.session.response)?; + let session = self.session.lock().unwrap().clone(); + let (_, available) = resolve_model_info(&self.name, &session.response)?; Ok(available) } } @@ -1125,6 +1181,7 @@ async fn handle_requests( .session_capabilities .close .is_some(); + let supports_load = init_response.agent_capabilities.load_session; let mcp_capabilities = init_response.agent_capabilities.mcp_capabilities.clone(); if let Some(tx) = init_tx.take() { log_undelivered(tx.send(Ok(init_response)), AGENT_METHOD_NAMES.initialize); @@ -1156,6 +1213,52 @@ async fn handle_requests( }; log_undelivered(response_tx.send(result), AGENT_METHOD_NAMES.session_new); } + ClientRequest::LoadSession { + session_id, + response_tx, + } => { + let result = if supports_load { + let mcp_servers = + filter_supported_servers(&config.mcp_servers, &mcp_capabilities); + cx.send_request( + LoadSessionRequest::new(session_id.clone(), config.work_dir.clone()) + .mcp_servers(mcp_servers), + ) + .block_task() + .await + .map(|response| { + NewSessionResponse::new(session_id.clone()) + .modes(response.modes) + .config_options(response.config_options) + .meta(response.meta) + }) + .map_err(anyhow::Error::from) + } else { + Err(anyhow::anyhow!("ACP agent does not support session/load")) + }; + let result = match result { + Ok(session) => { + session_ids.push(session.session_id.clone()); + apply_session_config_options(&config, &cx, session.session_id.clone()) + .await?; + apply_session_mode(&config, &goose_mode, &cx, session).await + } + Err(error) => Err(error), + }; + log_undelivered(response_tx.send(result), AGENT_METHOD_NAMES.session_load); + } + ClientRequest::CloseSession { session_id } => { + if supports_close { + if let Err(error) = cx + .send_request(CloseSessionRequest::new(session_id.clone())) + .block_task() + .await + { + tracing::debug!(method = AGENT_METHOD_NAMES.session_close, session_id = %session_id, %error, "failed to close replaced ACP session"); + } + } + session_ids.retain(|id| id != &session_id); + } ClientRequest::SetMode { session_id, mode_id, @@ -1223,7 +1326,7 @@ async fn handle_requests( } } - if supports_close { + if supports_close && !supports_load { for session_id in session_ids { if let Err(e) = cx .send_request(CloseSessionRequest::new(session_id.clone())) @@ -1752,10 +1855,10 @@ mod tests { name: "acp-test".to_string(), goose_mode: Arc::new(Mutex::new(GooseMode::Auto)), mode_mapping: HashMap::new(), - session: AcpSession { + session: Mutex::new(AcpSession { id: SessionId::new("test-session"), response: NewSessionResponse::new("test-session"), - }, + }), pending_confirmations: Arc::new(TokioMutex::new(HashMap::new())), pending_tool_updates: Arc::new(Mutex::new(HashMap::new())), handoff_context_sent: AtomicBool::new(false), @@ -2081,6 +2184,50 @@ mod tests { assert!(!later_claim.include_context); } + #[tokio::test] + async fn resume_replaces_session_and_skips_handoff() { + let (tx, mut rx) = mpsc::channel(2); + let (provider, _) = test_provider_with_tx(Some(tx)); + + let handle = tokio::spawn(async move { + provider.resume("saved-session").await.unwrap(); + provider + }); + + let ClientRequest::LoadSession { + session_id, + response_tx, + } = rx.recv().await.expect("expected session/load") + else { + panic!("expected session/load"); + }; + assert_eq!(session_id.to_string(), "saved-session"); + response_tx + .send(Ok(NewSessionResponse::new("saved-session"))) + .unwrap(); + + let ClientRequest::CloseSession { session_id } = + rx.recv().await.expect("expected temporary session close") + else { + panic!("expected temporary session close"); + }; + assert_eq!(session_id.to_string(), "test-session"); + + let provider = handle.await.unwrap(); + assert_eq!( + provider.provider_session_id().as_deref(), + Some("saved-session") + ); + + let messages = vec![ + Message::assistant().with_text("prior answer"), + Message::user().with_text("current request"), + ]; + let claim = provider.claim_handoff_context(&messages); + assert!(!claim.first_prompt); + assert!(!claim.include_context); + } + #[tokio::test] async fn get_context_limit_surfaces_captured_context_size() { let (provider, model) = test_provider(); @@ -2324,10 +2471,9 @@ mod tests { GooseMode::Auto, vec!["full-access".to_string(), "agent-full-access".to_string()], )]); - provider.session.response = NewSessionResponse::new("test-session").modes(mode_state( - "read-only", - &["read-only", "agent", "agent-full-access"], - )); + provider.session.lock().unwrap().response = NewSessionResponse::new("test-session").modes( + mode_state("read-only", &["read-only", "agent", "agent-full-access"]), + ); let handle = tokio::spawn(async move { provider @@ -2358,7 +2504,7 @@ mod tests { let (tx, mut rx) = mpsc::channel(1); let (mut provider, _) = test_provider_with_tx(Some(tx)); provider.mode_mapping = HashMap::from([(GooseMode::Chat, vec!["read-only".to_string()])]); - provider.session.response = NewSessionResponse::new("test-session") + provider.session.lock().unwrap().response = NewSessionResponse::new("test-session") .modes(mode_state("agent", &["agent", "agent-full-access"])); let result = provider.update_mode("session", GooseMode::Chat).await; diff --git a/crates/goose/src/agents/agent.rs b/crates/goose/src/agents/agent.rs index 69591818d..32174d3d2 100644 --- a/crates/goose/src/agents/agent.rs +++ b/crates/goose/src/agents/agent.rs @@ -2163,17 +2163,33 @@ impl Agent { let provider = self.provider().await?; let provider_name = provider.get_name().to_string(); + let saved_provider_session_id = + super::latest_provider_session_id(conversation.messages(), &provider_name); + if let Some(saved_provider_session_id) = saved_provider_session_id { + if let Err(error) = provider.resume(saved_provider_session_id).await { + warn!( + provider = provider_name, + %error, + "Could not resume provider session; continuing with a handoff" + ); + } + } + let requested_model = model_config.model_name.clone(); - let inference = provider + let resolved_model = provider .fetch_model_info(&requested_model) .await .ok() - .and_then(|model_info| model_info.resolved_model) - .map(|resolved_model| InferenceMetadata { + .and_then(|model_info| model_info.resolved_model); + let provider_session_id = provider.provider_session_id(); + let inference = (resolved_model.is_some() || provider_session_id.is_some()).then(|| { + InferenceMetadata { provider: provider_name.clone(), requested_model, - resolved_model: Some(resolved_model), - }); + resolved_model, + provider_session_id, + } + }); let session_manager = self.config.session_manager.clone(); let session_id = session_config.id.clone(); if !self.config.disable_session_naming { @@ -3941,6 +3957,33 @@ mod tests { use std::sync::atomic::{AtomicUsize, Ordering}; use tempfile::TempDir; + #[test] + fn provider_session_id_comes_from_latest_inference() { + let messages = vec![ + Message::assistant().with_inference(InferenceMetadata { + provider: "codex-acp".to_string(), + requested_model: "current".to_string(), + resolved_model: None, + provider_session_id: Some("codex-session".to_string()), + }), + Message::assistant().with_inference(InferenceMetadata { + provider: "claude-acp".to_string(), + requested_model: "current".to_string(), + resolved_model: None, + provider_session_id: Some("claude-session".to_string()), + }), + ]; + + assert_eq!( + super::super::latest_provider_session_id(&messages, "claude-acp"), + Some("claude-session") + ); + assert_eq!( + super::super::latest_provider_session_id(&messages, "codex-acp"), + None + ); + } + #[test] fn recipe_history_excludes_turn_context_events() { use crate::conversation::message::MessageMetadata; diff --git a/crates/goose/src/agents/mod.rs b/crates/goose/src/agents/mod.rs index 0ac624fb4..7a4011408 100644 --- a/crates/goose/src/agents/mod.rs +++ b/crates/goose/src/agents/mod.rs @@ -36,3 +36,16 @@ pub use subagent_handler::SUBAGENT_TOOL_REQUEST_TYPE; pub use subagent_task_config::TaskConfig; pub use tool_execution::ToolCallContext; pub use types::{FrontendTool, RetryConfig, SessionConfig, SuccessCheck}; + +fn latest_provider_session_id<'a>( + messages: &'a [crate::conversation::message::Message], + provider: &str, +) -> Option<&'a str> { + let inference = messages + .iter() + .rev() + .find_map(|message| message.metadata.inference.as_ref())?; + (inference.provider == provider) + .then_some(inference.provider_session_id.as_deref()) + .flatten() +} diff --git a/crates/goose/src/agents/state_machine/ops_llm.rs b/crates/goose/src/agents/state_machine/ops_llm.rs index b1e2d2904..0ba660f89 100644 --- a/crates/goose/src/agents/state_machine/ops_llm.rs +++ b/crates/goose/src/agents/state_machine/ops_llm.rs @@ -434,6 +434,19 @@ impl Inference for InferenceRunner<'_> { .get_context_limit(&self.model_config) .await .unwrap_or_else(|_| self.model_config.context_limit()); + let provider_name = self.provider.get_name(); + if let Some(session_id) = super::super::latest_provider_session_id( + conversation.messages(), + provider_name, + ) { + if let Err(error) = self.provider.resume(session_id).await { + tracing::warn!( + provider = provider_name, + %error, + "Could not resume provider session; continuing with a handoff" + ); + } + } let turn = messages_since_kickoff(conversation)?; let turn_start = turn .first() @@ -481,17 +494,21 @@ impl Inference for InferenceRunner<'_> { }; let requested_model = self.model_config.model_name.clone(); - let inference = self + let resolved_model = self .provider .fetch_model_info(&requested_model) .await .ok() - .and_then(|model_info| model_info.resolved_model) - .map(|resolved_model| InferenceMetadata { + .and_then(|model_info| model_info.resolved_model); + let provider_session_id = self.provider.provider_session_id(); + let inference = (resolved_model.is_some() || provider_session_id.is_some()).then(|| { + InferenceMetadata { provider: self.provider.get_name().to_string(), requested_model, - resolved_model: Some(resolved_model), - }); + resolved_model, + provider_session_id, + } + }); let mut accumulator = Conversation::empty(); let mut tool_request_ids = std::collections::HashSet::new(); diff --git a/ui/desktop/src/types/message.ts b/ui/desktop/src/types/message.ts index 86f14b855..6fbc43344 100644 --- a/ui/desktop/src/types/message.ts +++ b/ui/desktop/src/types/message.ts @@ -168,6 +168,7 @@ export type InferenceMetadata = { provider: string; requestedModel: string; resolvedModel?: string | null; + providerSessionId?: string | null; }; /** Mirrors the backend `MessageUsage` schema (camelCase). */