fix(acp): seed provider handoff history (#8941)

Signed-off-by: Matt Toohey <contact@matttoohey.com>
This commit is contained in:
Matt Toohey
2026-05-02 08:09:38 +10:00
committed by GitHub
parent e76640c8c4
commit c365e7b950
+260 -24
View File
@@ -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<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<()>>,
@@ -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<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);
}
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<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());
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<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");
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<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(),