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::future::Future;
|
||||||
use std::path::PathBuf;
|
use std::path::PathBuf;
|
||||||
use std::process::Stdio;
|
use std::process::Stdio;
|
||||||
use std::sync::{Arc, Mutex};
|
use std::sync::{
|
||||||
|
atomic::{AtomicBool, Ordering},
|
||||||
|
Arc, Mutex,
|
||||||
|
};
|
||||||
use std::thread::JoinHandle;
|
use std::thread::JoinHandle;
|
||||||
use tokio::process::{Child, Command};
|
use tokio::process::{Child, Command};
|
||||||
use tokio::sync::{mpsc, oneshot, Mutex as TokioMutex};
|
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::acp::{map_permission_response, PermissionDecision};
|
||||||
use crate::config::{ExtensionConfig, GooseMode};
|
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::conversation::message::{Message, MessageContent, TOOL_META_EXTERNAL_DISPATCH_KEY};
|
||||||
use crate::model::ModelConfig;
|
use crate::model::ModelConfig;
|
||||||
use crate::permission::permission_confirmation::PrincipalType;
|
use crate::permission::permission_confirmation::PrincipalType;
|
||||||
@@ -122,6 +126,11 @@ struct AcpSession {
|
|||||||
response: NewSessionResponse,
|
response: NewSessionResponse,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
struct HandoffContextClaim {
|
||||||
|
first_prompt: bool,
|
||||||
|
include_context: bool,
|
||||||
|
}
|
||||||
|
|
||||||
pub struct AcpProvider {
|
pub struct AcpProvider {
|
||||||
name: String,
|
name: String,
|
||||||
model: ModelConfig,
|
model: ModelConfig,
|
||||||
@@ -133,6 +142,7 @@ pub struct AcpProvider {
|
|||||||
pending_confirmations:
|
pending_confirmations:
|
||||||
Arc<TokioMutex<HashMap<String, oneshot::Sender<PermissionConfirmation>>>>,
|
Arc<TokioMutex<HashMap<String, oneshot::Sender<PermissionConfirmation>>>>,
|
||||||
pending_tool_updates: Arc<Mutex<HashMap<String, AccumulatedToolCall>>>,
|
pending_tool_updates: Arc<Mutex<HashMap<String, AccumulatedToolCall>>>,
|
||||||
|
handoff_context_sent: AtomicBool,
|
||||||
|
|
||||||
tx: Option<mpsc::Sender<ClientRequest>>,
|
tx: Option<mpsc::Sender<ClientRequest>>,
|
||||||
loop_thread: Option<JoinHandle<()>>,
|
loop_thread: Option<JoinHandle<()>>,
|
||||||
@@ -261,6 +271,7 @@ impl AcpProvider {
|
|||||||
session,
|
session,
|
||||||
pending_confirmations: Arc::new(TokioMutex::new(HashMap::new())),
|
pending_confirmations: Arc::new(TokioMutex::new(HashMap::new())),
|
||||||
pending_tool_updates,
|
pending_tool_updates,
|
||||||
|
handoff_context_sent: AtomicBool::new(false),
|
||||||
tx: Some(tx),
|
tx: Some(tx),
|
||||||
loop_thread: Some(loop_thread),
|
loop_thread: Some(loop_thread),
|
||||||
})
|
})
|
||||||
@@ -334,6 +345,14 @@ impl AcpProvider {
|
|||||||
.as_ref()
|
.as_ref()
|
||||||
.is_some_and(|opts| opts.iter().any(|o| o.category.as_ref() == Some(&category)))
|
.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]
|
#[async_trait::async_trait]
|
||||||
@@ -400,16 +419,24 @@ impl Provider for AcpProvider {
|
|||||||
) -> Result<MessageStream, ProviderError> {
|
) -> Result<MessageStream, ProviderError> {
|
||||||
let session_id = self.acp_session_id();
|
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
|
// Drop any tool-call buffer state left over from a prior prompt
|
||||||
// (e.g. cancelled or interrupted before its terminal status arrived).
|
// (e.g. cancelled or interrupted before its terminal status arrived).
|
||||||
if let Ok(mut buffer) = self.pending_tool_updates.lock() {
|
if let Ok(mut buffer) = self.pending_tool_updates.lock() {
|
||||||
buffer.clear();
|
buffer.clear();
|
||||||
}
|
}
|
||||||
let mut rx = self
|
let mut rx = match self.prompt(session_id, prompt_blocks).await {
|
||||||
.prompt(session_id, prompt_blocks)
|
Ok(rx) => rx,
|
||||||
.await
|
Err(e) => {
|
||||||
.map_err(|e| ProviderError::RequestFailed(format!("Failed to send ACP prompt: {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 pending_confirmations = self.pending_confirmations.clone();
|
||||||
let goose_mode = *self
|
let goose_mode = *self
|
||||||
@@ -1131,34 +1158,73 @@ fn filter_supported_servers(
|
|||||||
.collect()
|
.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 mut content_blocks = Vec::new();
|
||||||
|
|
||||||
let last_user = messages
|
let Some(last_user_index) = last_user_message_index(messages) else {
|
||||||
.iter()
|
return content_blocks;
|
||||||
.rev()
|
};
|
||||||
.find(|m| m.role == Role::User && m.is_agent_visible());
|
|
||||||
|
|
||||||
if let Some(message) = last_user {
|
if include_handoff_context {
|
||||||
for content in &message.content {
|
if let Some(memo) = build_handoff_context_memo(&messages[..last_user_index]) {
|
||||||
match content {
|
content_blocks.push(ContentBlock::Text(TextContent::new(memo)));
|
||||||
MessageContent::Text(text) => {
|
}
|
||||||
content_blocks.push(ContentBlock::Text(TextContent::new(text.text.clone())));
|
}
|
||||||
}
|
|
||||||
MessageContent::Image(image) => {
|
let message = &messages[last_user_index];
|
||||||
content_blocks.push(ContentBlock::Image(ImageContent::new(
|
for content in &message.content {
|
||||||
&image.data,
|
match content {
|
||||||
&image.mime_type,
|
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
|
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
|
/// Convert ACP `ToolCallContent` blocks into the rmcp `Content` shape goose's
|
||||||
/// `Message::with_tool_response` consumes. Handles `Content` (text/image/other),
|
/// `Message::with_tool_response` consumes. Handles `Content` (text/image/other),
|
||||||
/// `Diff`, and `Terminal` variants; falls back to a JSON serialization of
|
/// `Diff`, and `Terminal` variants; falls back to a JSON serialization of
|
||||||
@@ -1358,6 +1424,176 @@ mod tests {
|
|||||||
use sacp::schema::SessionConfigSelectOption;
|
use sacp::schema::SessionConfigSelectOption;
|
||||||
use test_case::test_case;
|
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(
|
#[test_case(
|
||||||
ExtensionConfig::Stdio {
|
ExtensionConfig::Stdio {
|
||||||
name: "github".into(),
|
name: "github".into(),
|
||||||
|
|||||||
Reference in New Issue
Block a user