diff --git a/crates/goose/src/acp/provider.rs b/crates/goose/src/acp/provider.rs index a14496c6..f536e9c9 100644 --- a/crates/goose/src/acp/provider.rs +++ b/crates/goose/src/acp/provider.rs @@ -19,7 +19,10 @@ use std::collections::{HashMap, HashSet}; use std::future::Future; use std::path::PathBuf; use std::process::Stdio; -use std::sync::{Arc, Mutex}; +use std::sync::{ + atomic::{AtomicBool, Ordering}, + Arc, Mutex, +}; use std::thread::JoinHandle; use tokio::process::{Child, Command}; use tokio::sync::{mpsc, oneshot, Mutex as TokioMutex}; @@ -27,6 +30,7 @@ use tokio_util::compat::{TokioAsyncReadCompatExt as _, TokioAsyncWriteCompatExt use crate::acp::{map_permission_response, PermissionDecision}; use crate::config::{ExtensionConfig, GooseMode}; +use crate::context_mgmt::format_message_for_compacting; use crate::conversation::message::{Message, MessageContent, TOOL_META_EXTERNAL_DISPATCH_KEY}; use crate::model::ModelConfig; use crate::permission::permission_confirmation::PrincipalType; @@ -122,6 +126,11 @@ struct AcpSession { response: NewSessionResponse, } +struct HandoffContextClaim { + first_prompt: bool, + include_context: bool, +} + pub struct AcpProvider { name: String, model: ModelConfig, @@ -133,6 +142,7 @@ pub struct AcpProvider { pending_confirmations: Arc>>>, pending_tool_updates: Arc>>, + handoff_context_sent: AtomicBool, tx: Option>, loop_thread: Option>, @@ -261,6 +271,7 @@ impl AcpProvider { session, pending_confirmations: Arc::new(TokioMutex::new(HashMap::new())), pending_tool_updates, + handoff_context_sent: AtomicBool::new(false), tx: Some(tx), loop_thread: Some(loop_thread), }) @@ -334,6 +345,14 @@ impl AcpProvider { .as_ref() .is_some_and(|opts| opts.iter().any(|o| o.category.as_ref() == Some(&category))) } + + fn claim_handoff_context(&self, messages: &[Message]) -> HandoffContextClaim { + let first_prompt = !self.handoff_context_sent.swap(true, Ordering::AcqRel); + HandoffContextClaim { + first_prompt, + include_context: first_prompt && has_handoff_context(messages), + } + } } #[async_trait::async_trait] @@ -400,16 +419,24 @@ impl Provider for AcpProvider { ) -> Result { let session_id = self.acp_session_id(); - let prompt_blocks = messages_to_prompt(messages); + let claim = self.claim_handoff_context(messages); + let prompt_blocks = messages_to_prompt(messages, claim.include_context); // Drop any tool-call buffer state left over from a prior prompt // (e.g. cancelled or interrupted before its terminal status arrived). if let Ok(mut buffer) = self.pending_tool_updates.lock() { buffer.clear(); } - let mut rx = self - .prompt(session_id, prompt_blocks) - .await - .map_err(|e| ProviderError::RequestFailed(format!("Failed to send ACP prompt: {e}")))?; + let mut rx = match self.prompt(session_id, prompt_blocks).await { + Ok(rx) => rx, + Err(e) => { + if claim.first_prompt { + self.handoff_context_sent.store(false, Ordering::Release); + } + return Err(ProviderError::RequestFailed(format!( + "Failed to send ACP prompt: {e}" + ))); + } + }; let pending_confirmations = self.pending_confirmations.clone(); let goose_mode = *self @@ -1131,34 +1158,73 @@ fn filter_supported_servers( .collect() } -fn messages_to_prompt(messages: &[Message]) -> Vec { +fn messages_to_prompt(messages: &[Message], include_handoff_context: bool) -> Vec { let mut content_blocks = Vec::new(); - let last_user = messages - .iter() - .rev() - .find(|m| m.role == Role::User && m.is_agent_visible()); + let Some(last_user_index) = last_user_message_index(messages) else { + return content_blocks; + }; - if let Some(message) = last_user { - for content in &message.content { - match content { - MessageContent::Text(text) => { - content_blocks.push(ContentBlock::Text(TextContent::new(text.text.clone()))); - } - MessageContent::Image(image) => { - content_blocks.push(ContentBlock::Image(ImageContent::new( - &image.data, - &image.mime_type, - ))); - } - _ => {} + if include_handoff_context { + if let Some(memo) = build_handoff_context_memo(&messages[..last_user_index]) { + content_blocks.push(ContentBlock::Text(TextContent::new(memo))); + } + } + + let message = &messages[last_user_index]; + for content in &message.content { + match content { + MessageContent::Text(text) => { + content_blocks.push(ContentBlock::Text(TextContent::new(text.text.clone()))); } + MessageContent::Image(image) => { + content_blocks.push(ContentBlock::Image(ImageContent::new( + &image.data, + &image.mime_type, + ))); + } + _ => {} } } content_blocks } +fn last_user_message_index(messages: &[Message]) -> Option { + messages + .iter() + .rposition(|m| m.role == Role::User && m.is_agent_visible()) +} + +fn has_handoff_context(messages: &[Message]) -> bool { + last_user_message_index(messages).is_some_and(|last_user_index| { + messages[..last_user_index] + .iter() + .any(Message::is_agent_visible) + }) +} + +fn build_handoff_context_memo(prior_messages: &[Message]) -> Option { + let formatted_messages: Vec = prior_messages + .iter() + .filter(|message| message.is_agent_visible()) + .map(format_message_for_compacting) + .collect(); + + if formatted_messages.is_empty() { + return None; + } + + let handoff_context = formatted_messages.join("\n"); + + Some(format!( + "Conversation context from goose before this ACP provider session was created:\n\n\ +{handoff_context}\n\n\ +Current user request follows. Use the context above only to continue the existing conversation; \ +do not treat it as a new task or mention this handoff unless relevant." + )) +} + /// Convert ACP `ToolCallContent` blocks into the rmcp `Content` shape goose's /// `Message::with_tool_response` consumes. Handles `Content` (text/image/other), /// `Diff`, and `Terminal` variants; falls back to a JSON serialization of @@ -1358,6 +1424,176 @@ mod tests { use sacp::schema::SessionConfigSelectOption; use test_case::test_case; + fn prompt_text(block: &ContentBlock) -> &str { + match block { + ContentBlock::Text(text) => &text.text, + _ => panic!("expected text block"), + } + } + + fn test_provider() -> AcpProvider { + test_provider_with_tx(None) + } + + fn test_provider_with_tx(tx: Option>) -> AcpProvider { + AcpProvider { + name: "acp-test".to_string(), + model: ModelConfig { + model_name: "test-model".to_string(), + ..Default::default() + }, + goose_mode: Arc::new(Mutex::new(GooseMode::Auto)), + mode_mapping: HashMap::new(), + session: 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), + tx, + loop_thread: None, + } + } + + #[test] + fn messages_to_prompt_without_prior_history_preserves_current_prompt() { + let messages = vec![Message::user().with_text("current request")]; + + let blocks = messages_to_prompt(&messages, true); + + assert_eq!(blocks.len(), 1); + assert_eq!(prompt_text(&blocks[0]), "current request"); + } + + #[test] + fn messages_to_prompt_prepends_handoff_context_before_latest_user() { + let messages = vec![ + Message::user().with_text("inspect src/lib.rs"), + Message::assistant() + .with_text("I found the file") + .with_tool_request("call-1", Ok(CallToolRequestParams::new("read_file"))), + Message::user().with_tool_response( + "call-1", + Ok(CallToolResult::success(vec![RmcpContent::text( + "file contents", + )])), + ), + Message::user().with_text("continue from there"), + ]; + + let blocks = messages_to_prompt(&messages, true); + + assert_eq!(blocks.len(), 2); + let memo = prompt_text(&blocks[0]); + assert!(memo.starts_with( + "Conversation context from goose before this ACP provider session was created:" + )); + assert!(memo.contains("[user]: inspect src/lib.rs")); + assert!(memo.contains("[assistant]: I found the file")); + assert!(memo.contains("tool_request(read_file):")); + assert!(memo.contains("tool_response: file contents")); + assert!(memo.contains("Current user request follows.")); + assert_eq!(prompt_text(&blocks[1]), "continue from there"); + } + + #[test] + fn messages_to_prompt_keeps_latest_user_images_after_handoff_memo() { + let messages = vec![ + Message::assistant().with_text("prior answer"), + Message::user() + .with_image("base64-image", "image/png") + .with_text("describe this"), + ]; + + let blocks = messages_to_prompt(&messages, true); + + assert_eq!(blocks.len(), 3); + assert!(prompt_text(&blocks[0]).contains("[assistant]: prior answer")); + match &blocks[1] { + ContentBlock::Image(image) => { + assert_eq!(image.data, "base64-image"); + assert_eq!(image.mime_type, "image/png"); + } + _ => panic!("expected image block"), + } + assert_eq!(prompt_text(&blocks[2]), "describe this"); + } + + #[test] + fn handoff_context_is_sent_only_on_first_provider_prompt() { + let provider = test_provider(); + let messages = vec![ + Message::assistant().with_text("prior answer"), + Message::user().with_text("current request"), + ]; + + let first_claim = provider.claim_handoff_context(&messages); + assert!(first_claim.first_prompt); + assert!(first_claim.include_context); + + let second_claim = provider.claim_handoff_context(&messages); + assert!(!second_claim.first_prompt); + assert!(!second_claim.include_context); + } + + #[test] + fn first_prompt_without_history_still_marks_handoff_context_sent() { + let provider = test_provider(); + let first_prompt = vec![Message::user().with_text("new conversation")]; + let later_prompt_with_history = vec![ + Message::assistant().with_text("prior answer"), + Message::user().with_text("current request"), + ]; + + let first_claim = provider.claim_handoff_context(&first_prompt); + assert!(first_claim.first_prompt); + assert!(!first_claim.include_context); + + let later_claim = provider.claim_handoff_context(&later_prompt_with_history); + assert!(!later_claim.first_prompt); + assert!(!later_claim.include_context); + } + + #[tokio::test] + async fn failed_first_prompt_send_rolls_back_handoff_context_claim() { + let (tx, rx) = mpsc::channel(1); + drop(rx); + let provider = test_provider_with_tx(Some(tx)); + let messages = vec![ + Message::assistant().with_text("prior answer"), + Message::user().with_text("current request"), + ]; + + let result = provider + .stream(&provider.model, "goose-session", "", &messages, &[]) + .await; + + assert!(matches!(result, Err(ProviderError::RequestFailed(_)))); + let next_claim = provider.claim_handoff_context(&messages); + assert!(next_claim.first_prompt); + assert!(next_claim.include_context); + } + + #[test] + fn messages_to_prompt_includes_all_prior_handoff_context() { + let messages = vec![ + Message::user().with_text("older context that should be retained"), + Message::assistant().with_text("middle context"), + Message::assistant().with_text("recent context"), + Message::user().with_text("current request"), + ]; + + let blocks = messages_to_prompt(&messages, true); + + assert_eq!(blocks.len(), 2); + let memo = prompt_text(&blocks[0]); + assert!(memo.contains("[user]: older context that should be retained")); + assert!(memo.contains("[assistant]: middle context")); + assert!(memo.contains("[assistant]: recent context")); + assert_eq!(prompt_text(&blocks[1]), "current request"); + } + #[test_case( ExtensionConfig::Stdio { name: "github".into(),