Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
286 changes: 261 additions & 25 deletions crates/goose/src/acp/provider.rs
Original file line number Diff line number Diff line change
Expand Up @@ -19,14 +19,18 @@ 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};
use tokio_util::compat::{TokioAsyncReadCompatExt as _, TokioAsyncWriteCompatExt as _};

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;
Expand Down Expand Up @@ -122,6 +126,11 @@ struct AcpSession {
response: NewSessionResponse,
}

struct HandoffContextClaim {
first_prompt: bool,
include_context: bool,
}

pub struct AcpProvider {
name: String,
model: ModelConfig,
Expand All @@ -133,6 +142,7 @@ pub struct AcpProvider {
pending_confirmations:
Arc<TokioMutex<HashMap<String, oneshot::Sender<PermissionConfirmation>>>>,
pending_tool_updates: Arc<Mutex<HashMap<String, AccumulatedToolCall>>>,
handoff_context_sent: AtomicBool,

tx: Option<mpsc::Sender<ClientRequest>>,
loop_thread: Option<JoinHandle<()>>,
Expand Down Expand Up @@ -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),
})
Expand Down Expand Up @@ -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]
Expand Down Expand Up @@ -400,16 +419,24 @@ impl Provider for AcpProvider {
) -> Result<MessageStream, ProviderError> {
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);
}
Comment on lines +432 to +434

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Badge Reset handoff claim when first prompt fails asynchronously

Roll back handoff_context_sent not only when prompt() enqueueing fails, but also when the first prompt later fails via AcpUpdate::Error. Right now the flag is set in claim_handoff_context, and only the immediate send error path resets it; if ACP accepts the request locally but returns an error (e.g., transport hiccup or prompt rejection) before producing a response, retries will skip the handoff memo even though no successful first prompt completed, dropping prior conversation context on the next attempt.

Useful? React with 👍 / 👎.

return Err(ProviderError::RequestFailed(format!(
"Failed to send ACP prompt: {e}"
)));
}
};

let pending_confirmations = self.pending_confirmations.clone();
let goose_mode = *self
Expand Down Expand Up @@ -1131,34 +1158,73 @@ fn filter_supported_servers(
.collect()
}

fn messages_to_prompt(messages: &[Message]) -> Vec<ContentBlock> {
fn messages_to_prompt(messages: &[Message], include_handoff_context: bool) -> Vec<ContentBlock> {
let mut content_blocks = Vec::new();

let last_user = messages
.iter()
.rev()
.find(|m| m.role == Role::User && m.is_agent_visible());

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,
)));
}
_ => {}
let Some(last_user_index) = last_user_message_index(messages) else {
return content_blocks;
};

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<usize> {
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<String> {
let formatted_messages: Vec<String> = 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");

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Badge Bound handoff memo size before building first ACP prompt

The new handoff path includes all prior agent-visible messages in a single memo, but there is no character/token cap before join("\n"). On long sessions (especially with large tool responses), the first ACP prompt can become too large and fail immediately with context/request-size errors, which blocks provider handoff for exactly the conversations this feature targets. Please cap or truncate the serialized history (e.g., by token/char budget and recency) before constructing the memo.

Useful? React with 👍 / 👎.


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
Expand Down Expand Up @@ -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<mpsc::Sender<ClientRequest>>) -> 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(),
Expand Down
Loading