fix(acp): seed provider handoff history (#8941)
Signed-off-by: Matt Toohey <contact@matttoohey.com>
This commit is contained in:
@@ -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(),
|
||||
|
||||
Reference in New Issue
Block a user