From e4b2df8cd0d919895f39fff3f497d9a0cd8c1441 Mon Sep 17 00:00:00 2001 From: luke Date: Wed, 19 Aug 2026 19:39:22 -0400 Subject: [PATCH 01/12] feat(acp): prototype nested form elicitation, replay, and a Cursor transport Exploratory work behind the design discussion on aaif-goose/goose#11346. Not proposed for upstream: that issue has not reached Ready, and this branch deliberately spans three separable concerns, only the first of which is in its scope. 1. Nested ACP form relay. A form request from a managed ACP provider reaches the outer client and its typed response returns to the originating provider, under both the legacy loop and the state machine. The request is persisted before the response is accepted so an answer cannot sort ahead of its own question, and liveness check, persistence, and waiter consumption are one operation so two clients cannot both persist an answer while only one wins the channel. 2. Session-load replay. A non-goal of #11346 and explicitly deferred by #9797. Replays an unanswered question on load and distinguishes a live waiter from a dead one so a client can continue honestly. 3. Cursor ACP transport. A non-goal of #11346, and against the direction settled in #8391. Kept only because it is what surfaced the cursor/ask_question schema details reported there. --- crates/goose-provider-types/src/base.rs | 26 +- crates/goose/src/acp/cursor.rs | 337 +++++++ crates/goose/src/acp/mod.rs | 5 +- crates/goose/src/acp/provider.rs | 844 +++++++++++++++++- crates/goose/src/acp/server.rs | 15 +- crates/goose/src/acp/server/elicitation.rs | 270 ++++-- crates/goose/src/acp/server/load_session.rs | 120 ++- crates/goose/src/action_required_manager.rs | 13 + crates/goose/src/agents/agent.rs | 203 ++++- crates/goose/src/agents/mod.rs | 8 +- .../goose/src/agents/state_machine/ops_llm.rs | 44 +- .../agents/state_machine/tests/agent_reply.rs | 267 +++++- crates/goose/src/elicitation.rs | 23 +- crates/goose/src/providers/cursor_agent.rs | 372 +++++++- 14 files changed, 2358 insertions(+), 189 deletions(-) create mode 100644 crates/goose/src/acp/cursor.rs diff --git a/crates/goose-provider-types/src/base.rs b/crates/goose-provider-types/src/base.rs index b44254349441..12b2361d3d1d 100644 --- a/crates/goose-provider-types/src/base.rs +++ b/crates/goose-provider-types/src/base.rs @@ -1,6 +1,6 @@ use async_trait::async_trait; use futures::Stream; -use rmcp::model::Tool; +use rmcp::model::{ElicitationAction, Tool}; use serde::{Deserialize, Serialize}; use serde_json::Value; use std::collections::HashMap; @@ -462,6 +462,17 @@ pub trait Provider: Send + Sync { Ok(()) } + async fn prepare_session( + &self, + provider_session_id: Option<&str>, + _has_provider_history: bool, + ) -> Result<(), ProviderError> { + match provider_session_id { + Some(session_id) => self.resume(session_id).await, + None => Ok(()), + } + } + /// Primary streaming method that all providers must implement. async fn stream( &self, @@ -657,6 +668,19 @@ pub trait Provider: Send + Sync { ) -> bool { false } + + async fn handle_elicitation_response( + &self, + _request_id: &str, + _user_data: &Value, + _action: &ElicitationAction, + ) -> bool { + false + } + + async fn has_pending_elicitation(&self, _request_id: &str) -> bool { + false + } } #[cfg(test)] diff --git a/crates/goose/src/acp/cursor.rs b/crates/goose/src/acp/cursor.rs new file mode 100644 index 000000000000..935225b68ba9 --- /dev/null +++ b/crates/goose/src/acp/cursor.rs @@ -0,0 +1,337 @@ +use std::collections::{BTreeMap, HashSet}; + +use agent_client_protocol::schema::v1::{ + CreateElicitationRequest, CreateElicitationResponse, ElicitationAction, + ElicitationContentValue, ElicitationFormMode, ElicitationSchema, ElicitationSessionScope, + EnumOption, MultiSelectPropertySchema, StringPropertySchema, +}; +use agent_client_protocol::{JsonRpcRequest, JsonRpcResponse}; +use serde::{Deserialize, Serialize}; + +#[derive(Debug, Clone, Serialize, Deserialize, JsonRpcRequest)] +#[request(method = "cursor/ask_question", response = CursorAskQuestionResponse)] +#[serde(rename_all = "camelCase")] +pub(crate) struct CursorAskQuestionRequest { + pub(crate) tool_call_id: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub(crate) title: Option, + pub(crate) questions: Vec, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub(crate) struct CursorQuestion { + pub(crate) id: String, + pub(crate) prompt: String, + pub(crate) options: Vec, + #[serde(default)] + pub(crate) allow_multiple: bool, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub(crate) struct CursorQuestionOption { + pub(crate) id: String, + pub(crate) label: String, +} + +#[derive(Debug, Clone, Serialize, Deserialize, JsonRpcResponse)] +pub(crate) struct CursorAskQuestionResponse { + pub(crate) outcome: CursorAskQuestionOutcome, +} + +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] +#[serde(tag = "outcome", rename_all = "lowercase")] +pub(crate) enum CursorAskQuestionOutcome { + Answered { answers: Vec }, + Skipped, + Cancelled, +} + +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] +#[serde(rename_all = "camelCase")] +pub(crate) struct CursorQuestionAnswer { + pub(crate) question_id: String, + pub(crate) selected_option_ids: Vec, +} + +impl CursorAskQuestionRequest { + pub(crate) fn to_elicitation_request( + &self, + session_id: &str, + ) -> Option { + if self.questions.is_empty() { + return None; + } + + let mut question_ids = HashSet::new(); + let mut schema = ElicitationSchema::new().title(self.title.clone()); + + for question in &self.questions { + if question.id.is_empty() + || question.options.is_empty() + || !question_ids.insert(question.id.as_str()) + { + return None; + } + + let mut option_ids = HashSet::new(); + let options = question + .options + .iter() + .map(|option| { + option_ids + .insert(option.id.as_str()) + .then(|| EnumOption::new(option.id.clone(), option.label.clone())) + }) + .collect::>>()?; + + schema = if question.allow_multiple { + schema.property( + question.id.clone(), + MultiSelectPropertySchema::titled(options).title(question.prompt.clone()), + true, + ) + } else { + schema.property( + question.id.clone(), + StringPropertySchema::new() + .title(question.prompt.clone()) + .one_of(options), + true, + ) + }; + } + + let message = self.title.clone().unwrap_or_else(|| { + if self.questions.len() == 1 { + "Cursor has a question".to_string() + } else { + "Cursor has a few questions".to_string() + } + }); + + Some(CreateElicitationRequest::new( + ElicitationFormMode::new(ElicitationSessionScope::new(session_id.to_string()), schema), + message, + )) + } + + pub(crate) fn response_from_elicitation( + &self, + response: CreateElicitationResponse, + ) -> CursorAskQuestionResponse { + let outcome = match response.action { + ElicitationAction::Accept(accept) => { + let Some(content) = accept.content else { + return CursorAskQuestionResponse::cancelled(); + }; + let Some(answers) = self.answers_from_content(&content) else { + return CursorAskQuestionResponse::cancelled(); + }; + CursorAskQuestionOutcome::Answered { answers } + } + ElicitationAction::Decline => CursorAskQuestionOutcome::Skipped, + ElicitationAction::Cancel => CursorAskQuestionOutcome::Cancelled, + _ => CursorAskQuestionOutcome::Cancelled, + }; + + CursorAskQuestionResponse { outcome } + } + + fn answers_from_content( + &self, + content: &BTreeMap, + ) -> Option> { + self.questions + .iter() + .map(|question| { + let selected_option_ids = match content.get(&question.id)? { + ElicitationContentValue::String(value) if !question.allow_multiple => { + vec![value.clone()] + } + ElicitationContentValue::StringArray(values) if question.allow_multiple => { + values.clone() + } + _ => return None, + }; + + Some(CursorQuestionAnswer { + question_id: question.id.clone(), + selected_option_ids, + }) + }) + .collect() + } +} + +impl CursorAskQuestionResponse { + pub(crate) fn cancelled() -> Self { + Self { + outcome: CursorAskQuestionOutcome::Cancelled, + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use agent_client_protocol::schema::v1::{ + ElicitationAcceptAction, ElicitationPropertySchema, MultiSelectItems, + }; + + fn request() -> CursorAskQuestionRequest { + CursorAskQuestionRequest { + tool_call_id: "tool-1".to_string(), + title: Some("Shape the implementation".to_string()), + questions: vec![ + CursorQuestion { + id: "direction".to_string(), + prompt: "Which direction?".to_string(), + options: vec![ + CursorQuestionOption { + id: "native".to_string(), + label: "Native ACP".to_string(), + }, + CursorQuestionOption { + id: "prose".to_string(), + label: "Prose".to_string(), + }, + ], + allow_multiple: false, + }, + CursorQuestion { + id: "priorities".to_string(), + prompt: "What matters?".to_string(), + options: vec![ + CursorQuestionOption { + id: "quality".to_string(), + label: "Quality".to_string(), + }, + CursorQuestionOption { + id: "speed".to_string(), + label: "Speed".to_string(), + }, + ], + allow_multiple: true, + }, + ], + } + } + + #[test] + fn converts_multiple_questions_without_inventing_options() { + let request = request(); + let elicitation = request.to_elicitation_request("cursor-session").unwrap(); + let ElicitationFormMode { + requested_schema, .. + } = match elicitation.mode { + agent_client_protocol::schema::v1::ElicitationMode::Form(form) => form, + _ => panic!("expected form elicitation"), + }; + + let ElicitationPropertySchema::String(direction) = + &requested_schema.properties["direction"] + else { + panic!("expected single-select property"); + }; + assert_eq!( + direction + .one_of + .as_ref() + .unwrap() + .iter() + .map(|option| option.value.as_str()) + .collect::>(), + vec!["native", "prose"] + ); + + let ElicitationPropertySchema::Array(priorities) = + &requested_schema.properties["priorities"] + else { + panic!("expected multi-select property"); + }; + let MultiSelectItems::Titled(items) = &priorities.items else { + panic!("expected titled multi-select items"); + }; + assert_eq!( + items + .options + .iter() + .map(|option| option.value.as_str()) + .collect::>(), + vec!["quality", "speed"] + ); + } + + #[test] + fn maps_accepted_values_back_to_cursor_ids_in_question_order() { + let request = request(); + let response = CreateElicitationResponse::new(ElicitationAction::Accept( + ElicitationAcceptAction::new().content(BTreeMap::from([ + ( + "priorities".to_string(), + ElicitationContentValue::StringArray(vec![ + "quality".to_string(), + "speed".to_string(), + ]), + ), + ( + "direction".to_string(), + ElicitationContentValue::String("native".to_string()), + ), + ])), + )); + + assert_eq!( + request.response_from_elicitation(response).outcome, + CursorAskQuestionOutcome::Answered { + answers: vec![ + CursorQuestionAnswer { + question_id: "direction".to_string(), + selected_option_ids: vec!["native".to_string()], + }, + CursorQuestionAnswer { + question_id: "priorities".to_string(), + selected_option_ids: vec!["quality".to_string(), "speed".to_string()], + }, + ], + } + ); + } + + #[test] + fn rejects_ambiguous_question_and_option_ids() { + let mut duplicate_questions = request(); + duplicate_questions.questions[1].id = "direction".to_string(); + assert!(duplicate_questions + .to_elicitation_request("cursor-session") + .is_none()); + + let mut duplicate_options = request(); + duplicate_options.questions[0].options[1].id = "native".to_string(); + assert!(duplicate_options + .to_elicitation_request("cursor-session") + .is_none()); + } + + #[test] + fn maps_decline_and_cancel_to_cursor_outcomes() { + let request = request(); + assert_eq!( + request + .response_from_elicitation(CreateElicitationResponse::new( + ElicitationAction::Decline, + )) + .outcome, + CursorAskQuestionOutcome::Skipped + ); + assert_eq!( + request + .response_from_elicitation(CreateElicitationResponse::new( + ElicitationAction::Cancel, + )) + .outcome, + CursorAskQuestionOutcome::Cancelled + ); + } +} diff --git a/crates/goose/src/acp/mod.rs b/crates/goose/src/acp/mod.rs index 3e73d59c6b0b..295ca4eaf96c 100644 --- a/crates/goose/src/acp/mod.rs +++ b/crates/goose/src/acp/mod.rs @@ -1,4 +1,5 @@ mod common; +mod cursor; pub(crate) mod fs; mod handoff; mod mcp_app_proxy; @@ -13,8 +14,10 @@ pub mod transport; pub use common::{map_permission_response, PermissionDecision}; pub use goose_sdk_types::{custom_notifications, custom_requests}; pub use provider::{ - extension_configs_to_mcp_servers, AcpProvider, AcpProviderConfig, ACP_CURRENT_MODEL, + extension_configs_to_mcp_servers, AcpClientExtension, AcpProvider, AcpProviderConfig, + ACP_CURRENT_MODEL, }; +pub(crate) use provider::{is_provider_form_elicitation, ACP_PROVIDER_ELICITATION_ID_PREFIX}; /// `data.reason` on a prompt error raised because the agent's account is out of credits. /// Set by the ACP server, read by the provider to tell a spent account apart from a diff --git a/crates/goose/src/acp/provider.rs b/crates/goose/src/acp/provider.rs index f1505a4b708c..4eccfdd74910 100644 --- a/crates/goose/src/acp/provider.rs +++ b/crates/goose/src/acp/provider.rs @@ -1,13 +1,15 @@ use agent_client_protocol::schema::v1::{ Annotations as AcpAnnotations, ClientCapabilities, CloseSessionRequest, ContentBlock, - ContentChunk, EnvVariable, HttpHeader, ImageContent, InitializeRequest, InitializeResponse, - LoadSessionRequest, McpCapabilities, McpServer, McpServerHttp, McpServerStdio, - NewSessionRequest, NewSessionResponse, PromptRequest, PromptResponse, RequestPermissionOutcome, - RequestPermissionRequest, RequestPermissionResponse, Role as AcpRole, SessionConfigKind, - SessionConfigOption, SessionConfigOptionCategory, SessionConfigSelectOptions, SessionId, - SessionModeState, SessionNotification, SessionUpdate, SetSessionConfigOptionRequest, - SetSessionModeRequest, SetSessionModeResponse, StopReason, TextContent, ToolCallContent, - ToolCallStatus, ToolKind, + ContentChunk, CreateElicitationRequest, CreateElicitationResponse, ElicitationAcceptAction, + ElicitationAction as AcpElicitationAction, ElicitationCapabilities, ElicitationContentValue, + ElicitationFormCapabilities, ElicitationMode, EnvVariable, HttpHeader, ImageContent, + InitializeRequest, InitializeResponse, LoadSessionRequest, McpCapabilities, McpServer, + McpServerHttp, McpServerStdio, NewSessionRequest, NewSessionResponse, PromptRequest, + PromptResponse, RequestPermissionOutcome, RequestPermissionRequest, RequestPermissionResponse, + Role as AcpRole, SessionConfigKind, SessionConfigOption, SessionConfigOptionCategory, + SessionConfigSelectOptions, SessionId, SessionModeState, SessionNotification, SessionUpdate, + SetSessionConfigOptionRequest, SetSessionModeRequest, SetSessionModeResponse, StopReason, + TextContent, ToolCallContent, ToolCallStatus, ToolKind, }; use agent_client_protocol::schema::ProtocolVersion; use agent_client_protocol::{Agent, Client, ConnectionTo}; @@ -17,8 +19,11 @@ use anyhow::{Context, Result}; use async_stream::try_stream; use futures::future::BoxFuture; use goose_providers::conversation::token_usage::{ProviderUsage, Usage}; -use rmcp::model::{CallToolRequestParams, CallToolResult, ContentBlock as RmcpContent, Role, Tool}; -use std::collections::{HashMap, HashSet}; +use rmcp::model::{ + CallToolRequestParams, CallToolResult, ContentBlock as RmcpContent, + ElicitationAction as McpElicitationAction, Role, Tool, +}; +use std::collections::{BTreeMap, HashMap, HashSet}; use std::future::Future; use std::path::PathBuf; use std::process::Stdio; @@ -32,10 +37,13 @@ use tokio::process::{Child, Command}; use tokio::sync::{mpsc, oneshot, Mutex as TokioMutex}; use tokio_util::compat::{TokioAsyncReadCompatExt as _, TokioAsyncWriteCompatExt as _}; +use crate::acp::cursor::{CursorAskQuestionRequest, CursorAskQuestionResponse}; use crate::acp::handoff::{build_handoff_context_memo, memo_token_budget, prompt_token_cost}; use crate::acp::{map_permission_response, PermissionDecision}; use crate::config::{ExtensionConfig, GooseMode}; -use crate::conversation::message::{Message, MessageContent, TOOL_META_EXTERNAL_DISPATCH_KEY}; +use crate::conversation::message::{ + ActionRequiredData, Message, MessageContent, TOOL_META_EXTERNAL_DISPATCH_KEY, +}; use crate::permission::permission_confirmation::PrincipalType; use crate::permission::{Permission, PermissionConfirmation}; use crate::providers::base::{MessageStream, PermissionRouting, Provider}; @@ -47,6 +55,26 @@ use goose_providers::model::ModelConfig; /// Sentinel: resolved to the actual model name during connect(). pub const ACP_CURRENT_MODEL: &str = "current"; +pub(crate) const ACP_PROVIDER_ELICITATION_ID_PREFIX: &str = "acp-provider:"; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum AcpClientExtension { + CursorAskQuestion, +} + +pub(crate) fn is_provider_form_elicitation(message: &Message) -> bool { + message.content.iter().any(|content| { + matches!( + content, + MessageContent::ActionRequired(action) + if matches!( + &action.data, + ActionRequiredData::Elicitation { id, .. } + if id.starts_with(ACP_PROVIDER_ELICITATION_ID_PREFIX) + ) + ) + }) +} pub struct AcpProviderConfig { pub command: PathBuf, @@ -148,6 +176,10 @@ enum AcpUpdate { request: Box, response_tx: oneshot::Sender, }, + ElicitationRequest { + request: Box, + response_tx: oneshot::Sender, + }, Complete(StopReason, Option), Error(agent_client_protocol::Error), } @@ -239,6 +271,7 @@ pub struct AcpProvider { pending_confirmations: Arc>>>, + pending_elicitations: Arc>>>, pending_tool_updates: Arc>>, /// True after the first ACP prompt completes with the handoff context committed. /// Failed or abandoned first prompts reset this so the next prompt can retry it. @@ -268,6 +301,26 @@ impl std::fmt::Debug for AcpProvider { } } +fn cancel_pending_elicitations( + pending: &Mutex>>, +) { + if let Ok(mut pending) = pending.lock() { + for (_, response_tx) in pending.drain() { + let _ = response_tx.send(CreateElicitationResponse::new(AcpElicitationAction::Cancel)); + } + } +} + +struct PendingElicitationGuard { + pending: Arc>>>, +} + +impl Drop for PendingElicitationGuard { + fn drop(&mut self) { + cancel_pending_elicitations(&self.pending); + } +} + fn spawn_client_loop(fut: impl Future + Send + 'static) -> JoinHandle<()> { std::thread::spawn(move || { let rt = tokio::runtime::Builder::new_current_thread() @@ -283,11 +336,21 @@ impl AcpProvider { name: String, goose_mode: GooseMode, config: AcpProviderConfig, + ) -> Result { + Self::connect_with_client_extensions(name, goose_mode, config, Vec::new()).await + } + + pub async fn connect_with_client_extensions( + name: String, + goose_mode: GooseMode, + config: AcpProviderConfig, + client_extensions: Vec, ) -> Result { Self::start( name, goose_mode, config, + client_extensions, Box::new(|cl, rx, init_tx, mut cancel_rx| { Box::pin(async move { tokio::select! { @@ -307,11 +370,24 @@ impl AcpProvider { goose_mode: GooseMode, config: AcpProviderConfig, transport: impl agent_client_protocol::ConnectTo + 'static, + ) -> Result { + Self::connect_with_transport_and_extensions(name, goose_mode, config, transport, Vec::new()) + .await + } + + #[doc(hidden)] + pub async fn connect_with_transport_and_extensions( + name: String, + goose_mode: GooseMode, + config: AcpProviderConfig, + transport: impl agent_client_protocol::ConnectTo + 'static, + client_extensions: Vec, ) -> Result { Self::start( name, goose_mode, config, + client_extensions, Box::new(move |cl, mut rx, init_tx, mut cancel_rx| { Box::pin(async move { tokio::select! { @@ -333,6 +409,7 @@ impl AcpProvider { name: String, goose_mode: GooseMode, config: AcpProviderConfig, + client_extensions: Vec, run: ClientLoopFn, ) -> Result { let (tx, rx) = mpsc::channel(32); @@ -352,6 +429,7 @@ impl AcpProvider { let context_size = Arc::new(AtomicU64::new(0)); let client_loop = AcpClientLoop::new( config, + client_extensions, goose_mode_shared.clone(), pending_tool_updates.clone(), context_size.clone(), @@ -394,6 +472,7 @@ impl AcpProvider { mode_mapping, session: Mutex::new(session), pending_confirmations: Arc::new(TokioMutex::new(HashMap::new())), + pending_elicitations: Arc::new(Mutex::new(HashMap::new())), pending_tool_updates, handoff_context_sent: Arc::new(AtomicBool::new(false)), context_size, @@ -581,6 +660,8 @@ impl Provider for AcpProvider { return Ok(()); } + cancel_pending_elicitations(&self.pending_elicitations); + let previous_session_id = self.acp_session_id(); let loaded = self .load_session(SessionId::new(session_id)) @@ -659,6 +740,59 @@ impl Provider for AcpProvider { false } + async fn handle_elicitation_response( + &self, + request_id: &str, + user_data: &serde_json::Value, + action: &McpElicitationAction, + ) -> bool { + let action = match action { + McpElicitationAction::Accept => { + let Some(values) = user_data.as_object() else { + tracing::warn!(request_id, "ACP elicitation response is not an object"); + return false; + }; + let mut content = BTreeMap::new(); + for (key, value) in values { + let Ok(value) = + serde_json::from_value::(value.clone()) + else { + tracing::warn!( + request_id, + field = key, + "ACP elicitation response contains an unsupported value" + ); + return false; + }; + content.insert(key.clone(), value); + } + AcpElicitationAction::Accept(ElicitationAcceptAction::new().content(content)) + } + McpElicitationAction::Decline => AcpElicitationAction::Decline, + McpElicitationAction::Cancel => AcpElicitationAction::Cancel, + _ => AcpElicitationAction::Cancel, + }; + + let Some(response_tx) = self + .pending_elicitations + .lock() + .ok() + .and_then(|mut pending| pending.remove(request_id)) + else { + return false; + }; + + response_tx + .send(CreateElicitationResponse::new(action)) + .is_ok() + } + + async fn has_pending_elicitation(&self, request_id: &str) -> bool { + self.pending_elicitations + .lock() + .is_ok_and(|pending| pending.contains_key(request_id)) + } + async fn stream( &self, model_config: &ModelConfig, @@ -737,6 +871,7 @@ impl Provider for AcpProvider { bare_retry_blocks.map(|blocks| (self.tx.as_ref().unwrap().clone(), session_id, blocks)); let pending_confirmations = self.pending_confirmations.clone(); + let pending_elicitations = self.pending_elicitations.clone(); let goose_mode = *self .goose_mode .lock() @@ -746,6 +881,9 @@ impl Provider for AcpProvider { let model_name = model_config.model_name.clone(); Ok(Box::pin(try_stream! { + let _pending_elicitation_guard = PendingElicitationGuard { + pending: pending_elicitations.clone(), + }; let mut suppress_text = false; let mut bare_retry = bare_retry; let mut updates_seen = 0usize; @@ -877,6 +1015,32 @@ impl Provider for AcpProvider { } let _ = response_tx.send(map_permission_response(&request, decision)); } + AcpUpdate::ElicitationRequest { request, response_tx } => { + text_run = None; + thought_run = None; + + let Some((request_id, action_required)) = + build_action_required_elicitation_message(&request) + else { + let _ = response_tx.send(CreateElicitationResponse::new( + AcpElicitationAction::Cancel, + )); + continue; + }; + + match pending_elicitations.lock() { + Ok(mut pending) => { + pending.insert(request_id, response_tx); + } + Err(_) => { + let _ = response_tx.send(CreateElicitationResponse::new( + AcpElicitationAction::Cancel, + )); + continue; + } + } + yield (Some(action_required), None); + } AcpUpdate::Complete(reason, usage) => { // Prefer retrying context over silently losing it. A harness may have // ingested the memo before cancelling or refusing, so a retry can duplicate @@ -956,23 +1120,32 @@ impl Drop for AcpProvider { struct AcpClientLoop { config: AcpProviderConfig, + client_extensions: Vec, goose_mode: Arc>, - prompt_response_tx: Arc>>>, + active_prompt: Arc>>, pending_tool_updates: Arc>>, context_size: Arc, } +#[derive(Clone)] +struct ActivePrompt { + session_id: SessionId, + response_tx: mpsc::Sender, +} + impl AcpClientLoop { fn new( config: AcpProviderConfig, + client_extensions: Vec, goose_mode: Arc>, pending_tool_updates: Arc>>, context_size: Arc, ) -> Self { Self { config, + client_extensions, goose_mode, - prompt_response_tx: Arc::new(Mutex::new(None)), + active_prompt: Arc::new(Mutex::new(None)), pending_tool_updates, context_size, } @@ -1025,19 +1198,22 @@ impl AcpClientLoop { ) -> Result<()> { let AcpClientLoop { config, + client_extensions, goose_mode, - prompt_response_tx, + active_prompt, pending_tool_updates, context_size, } = self; let notification_callback = config.notification_callback.clone(); + let cursor_ask_question = + client_extensions.contains(&AcpClientExtension::CursorAskQuestion); let reverse_modes = reverse_mode_mapping(&config.mode_mapping); Client .builder() .on_receive_notification( { - let prompt_response_tx = prompt_response_tx.clone(); + let active_prompt = active_prompt.clone(); let reverse_modes = reverse_modes.clone(); let goose_mode = goose_mode.clone(); let pending_tool_updates = pending_tool_updates.clone(); @@ -1080,11 +1256,12 @@ impl AcpClientLoop { } _ => {} } - if let Some(tx) = prompt_response_tx + if let Some(tx) = active_prompt .lock() .ok() .as_ref() - .and_then(|g| g.as_ref()) + .and_then(|guard| guard.as_ref()) + .map(|prompt| &prompt.response_tx) { match notification.update { SessionUpdate::AgentMessageChunk(ContentChunk { @@ -1205,17 +1382,18 @@ impl AcpClientLoop { ) .on_receive_request( { - let prompt_response_tx = prompt_response_tx.clone(); + let active_prompt = active_prompt.clone(); async move |request: RequestPermissionRequest, responder, _connection_cx| { let (response_tx, response_rx) = oneshot::channel(); - let handler = prompt_response_tx + let handler = active_prompt .lock() .ok() .as_ref() .and_then(|g| g.as_ref().cloned()); - let tx = - handler.ok_or_else(agent_client_protocol::Error::internal_error)?; + let tx = handler + .map(|prompt| prompt.response_tx) + .ok_or_else(agent_client_protocol::Error::internal_error)?; if tx.is_closed() { return Err(agent_client_protocol::Error::internal_error()); @@ -1235,8 +1413,88 @@ impl AcpClientLoop { }, agent_client_protocol::on_receive_request!(), ) + .on_receive_request( + { + let active_prompt = active_prompt.clone(); + async move |request: CreateElicitationRequest, responder, _connection_cx| { + let (response_tx, response_rx) = oneshot::channel(); + + let handler = active_prompt + .lock() + .ok() + .as_ref() + .and_then(|g| g.as_ref().cloned()); + let tx = handler + .map(|prompt| prompt.response_tx) + .ok_or_else(agent_client_protocol::Error::internal_error)?; + + if tx.is_closed() { + return Err(agent_client_protocol::Error::internal_error()); + } + + tx.try_send(AcpUpdate::ElicitationRequest { + request: Box::new(request), + response_tx, + }) + .map_err(|_| agent_client_protocol::Error::internal_error())?; + + let response = response_rx.await.unwrap_or_else(|_| { + CreateElicitationResponse::new(AcpElicitationAction::Cancel) + }); + responder.respond(response) + } + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + { + let active_prompt = active_prompt.clone(); + async move |request: CursorAskQuestionRequest, responder, _connection_cx| { + if !cursor_ask_question { + return Err(agent_client_protocol::Error::method_not_found()); + } + + let handler = active_prompt + .lock() + .ok() + .as_ref() + .and_then(|guard| guard.as_ref().cloned()); + let Some(handler) = handler else { + return responder.respond(CursorAskQuestionResponse::cancelled()); + }; + let Some(elicitation_request) = + request.to_elicitation_request(handler.session_id.0.as_ref()) + else { + tracing::warn!( + tool_call_id = request.tool_call_id, + "Cursor question request is malformed" + ); + return responder.respond(CursorAskQuestionResponse::cancelled()); + }; + let (response_tx, response_rx) = oneshot::channel(); + let tx = handler.response_tx; + + if tx.is_closed() + || tx + .try_send(AcpUpdate::ElicitationRequest { + request: Box::new(elicitation_request), + response_tx, + }) + .is_err() + { + return responder.respond(CursorAskQuestionResponse::cancelled()); + } + + let response = response_rx.await.unwrap_or_else(|_| { + CreateElicitationResponse::new(AcpElicitationAction::Cancel) + }); + responder.respond(request.response_from_elicitation(response)) + } + }, + agent_client_protocol::on_receive_request!(), + ) .connect_with(transport, async move |cx: ConnectionTo| { - handle_requests(config, goose_mode, cx, rx, prompt_response_tx, init_tx).await + handle_requests(config, goose_mode, cx, rx, active_prompt, init_tx).await }) .await?; @@ -1338,12 +1596,13 @@ async fn handle_requests( goose_mode: Arc>, cx: ConnectionTo, rx: &mut mpsc::Receiver, - prompt_response_tx: Arc>>>, + active_prompt: Arc>>, init_tx: oneshot::Sender>, ) -> Result<(), agent_client_protocol::Error> { let mut init_tx = Some(init_tx); - let client_capabilities = ClientCapabilities::new(); + let client_capabilities = ClientCapabilities::new() + .elicitation(ElicitationCapabilities::new().form(ElicitationFormCapabilities::new())); let init_response: InitializeResponse = cx .send_request( InitializeRequest::new(ProtocolVersion::V1).client_capabilities(client_capabilities), @@ -1478,7 +1737,10 @@ async fn handle_requests( content, response_tx, } => { - *prompt_response_tx.lock().unwrap() = Some(response_tx.clone()); + *active_prompt.lock().unwrap() = Some(ActivePrompt { + session_id: session_id.clone(), + response_tx: response_tx.clone(), + }); let response: Result = cx .send_request(PromptRequest::new(session_id, content)) @@ -1500,7 +1762,7 @@ async fn handle_requests( } } - *prompt_response_tx.lock().unwrap() = None; + *active_prompt.lock().unwrap() = None; } } } @@ -1882,6 +2144,29 @@ fn build_action_required_message(request: &RequestPermissionRequest) -> Option Option<(String, Message)> { + let ElicitationMode::Form(form) = &request.mode else { + return None; + }; + let requested_schema = serde_json::to_value(&form.requested_schema).ok()?; + let request_id = format!( + "{}{}", + ACP_PROVIDER_ELICITATION_ID_PREFIX, + uuid::Uuid::new_v4() + ); + let message = Message::assistant() + .with_content(MessageContent::action_required_elicitation( + request_id.clone(), + request.message.clone(), + requested_schema, + )) + .user_only(); + + Some((request_id, message)) +} + fn extract_model_info_from_config_options( config_options: &[SessionConfigOption], ) -> Option<(String, Vec)> { @@ -1972,7 +2257,9 @@ mod tests { use super::*; use crate::agents::extension::Envs; use agent_client_protocol::schema::v1::{ - ErrorCode, SessionConfigSelectOption, SessionMode, SessionModeId, + AgentCapabilities, ElicitationFormMode, ElicitationSchema, ElicitationSessionScope, + ErrorCode, MultiSelectPropertySchema, SessionConfigSelectOption, SessionMode, + SessionModeId, }; use test_case::test_case; @@ -2088,6 +2375,507 @@ mod tests { assert_eq!(error.data, Some(serde_json::json!("sign in"))); } + #[tokio::test] + async fn stream_forwards_nested_form_elicitation_and_returns_response() { + use futures::StreamExt; + + let (tx, mut rx) = mpsc::channel(1); + let (provider, model) = test_provider_with_tx(Some(tx)); + let messages = vec![Message::user().with_text("ask me")]; + let mut stream = provider.stream(&model, "", &messages, &[]).await.unwrap(); + + let prompt_tx = match rx.recv().await.expect("expected ACP prompt request") { + ClientRequest::Prompt { response_tx, .. } => response_tx, + _ => panic!("expected ACP prompt request"), + }; + let request = CreateElicitationRequest::new( + ElicitationFormMode::new( + ElicitationSessionScope::new("nested-session"), + ElicitationSchema::new().string("answer", true), + ), + "Choose an answer", + ); + let (response_tx, response_rx) = oneshot::channel(); + prompt_tx + .send(AcpUpdate::ElicitationRequest { + request: Box::new(request), + response_tx, + }) + .await + .unwrap(); + + let (message, usage) = stream.next().await.unwrap().unwrap(); + assert!(usage.is_none()); + let message = message.expect("expected action-required message"); + assert!(is_provider_form_elicitation(&message)); + let MessageContent::ActionRequired(action_required) = &message.content[0] else { + panic!("expected action-required content"); + }; + let ActionRequiredData::Elicitation { id, message, .. } = &action_required.data else { + panic!("expected elicitation action-required content"); + }; + assert!(id.starts_with(ACP_PROVIDER_ELICITATION_ID_PREFIX)); + assert_eq!(message, "Choose an answer"); + + assert!( + provider + .handle_elicitation_response( + id, + &serde_json::json!({ "answer": "alpha" }), + &McpElicitationAction::Accept, + ) + .await + ); + + let response = response_rx.await.unwrap(); + let AcpElicitationAction::Accept(accept) = response.action else { + panic!("expected accepted elicitation response"); + }; + assert_eq!( + accept.content.unwrap().get("answer"), + Some(&ElicitationContentValue::String("alpha".to_string())) + ); + + prompt_tx + .send(AcpUpdate::Complete(StopReason::EndTurn, None)) + .await + .unwrap(); + assert!(stream.next().await.is_none()); + } + + #[tokio::test] + async fn invalid_elicitation_content_keeps_the_live_request_pending() { + let (provider, _) = test_provider(); + let (response_tx, response_rx) = oneshot::channel(); + provider + .pending_elicitations + .lock() + .unwrap() + .insert("request-1".to_string(), response_tx); + + assert!( + !provider + .handle_elicitation_response( + "request-1", + &serde_json::json!({ "answer": { "nested": true } }), + &McpElicitationAction::Accept, + ) + .await + ); + assert!(provider.has_pending_elicitation("request-1").await); + + assert!( + provider + .handle_elicitation_response( + "request-1", + &serde_json::json!({ "answer": "valid" }), + &McpElicitationAction::Accept, + ) + .await + ); + let response = response_rx.await.unwrap(); + let AcpElicitationAction::Accept(accept) = response.action else { + panic!("expected accepted response"); + }; + assert_eq!( + accept.content.unwrap().get("answer"), + Some(&ElicitationContentValue::String("valid".to_string())) + ); + } + + #[tokio::test] + async fn dropping_a_provider_stream_cancels_live_elicitation_waiters() { + let pending = Arc::new(Mutex::new(HashMap::new())); + let (response_tx, response_rx) = oneshot::channel(); + pending + .lock() + .unwrap() + .insert("request-1".to_string(), response_tx); + + drop(PendingElicitationGuard { + pending: pending.clone(), + }); + + assert!(pending.lock().unwrap().is_empty()); + assert!(matches!( + response_rx.await.unwrap().action, + AcpElicitationAction::Cancel + )); + } + + #[tokio::test] + async fn nested_acp_form_elicitation_round_trips_over_the_protocol() { + use futures::StreamExt; + + let (provider_read, agent_write) = tokio::io::duplex(64 * 1024); + let (agent_read, provider_write) = tokio::io::duplex(64 * 1024); + let provider_transport = agent_client_protocol::ByteStreams::new( + provider_write.compat_write(), + provider_read.compat(), + ); + let agent_transport = agent_client_protocol::ByteStreams::new( + agent_write.compat_write(), + agent_read.compat(), + ); + let (observed_response_tx, mut observed_response_rx) = mpsc::unbounded_channel(); + let (shutdown_tx, shutdown_rx) = oneshot::channel::<()>(); + + let nested_agent = tokio::spawn(async move { + Agent + .builder() + .on_receive_request( + async move |request: InitializeRequest, responder, _cx| { + assert!( + request + .client_capabilities + .elicitation + .as_ref() + .and_then(|elicitation| elicitation.form.as_ref()) + .is_some(), + "the nested provider must see form elicitation support" + ); + responder.respond( + InitializeResponse::new(request.protocol_version) + .agent_capabilities(AgentCapabilities::new()), + ) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async move |_request: NewSessionRequest, responder, _cx| { + responder.respond(NewSessionResponse::new("nested-session")) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async move |_request: PromptRequest, responder, cx| { + let cx_for_prompt = cx.clone(); + let observed_response_tx = observed_response_tx.clone(); + cx.spawn(async move { + let schema = ElicitationSchema::new() + .property( + "priorities", + MultiSelectPropertySchema::new(vec![ + "quality".to_string(), + "speed".to_string(), + "playfulness".to_string(), + ]), + true, + ) + .string("context", true); + let response = cx_for_prompt + .send_request(CreateElicitationRequest::new( + ElicitationFormMode::new( + ElicitationSessionScope::new("nested-session"), + schema, + ), + "Choose priorities and add context", + )) + .block_task() + .await?; + observed_response_tx + .send(response) + .map_err(|_| agent_client_protocol::Error::internal_error())?; + responder.respond(PromptResponse::new(StopReason::EndTurn)) + })?; + Ok(()) + }, + agent_client_protocol::on_receive_request!(), + ) + .connect_with(agent_transport, async move |_cx: ConnectionTo| { + let _ = shutdown_rx.await; + Ok(()) + }) + .await + }); + + let provider = tokio::time::timeout( + std::time::Duration::from_secs(5), + AcpProvider::connect_with_transport( + "nested-acp-test".to_string(), + GooseMode::Auto, + test_acp_config(HashMap::new(), None), + provider_transport, + ), + ) + .await + .expect("timed out connecting the nested ACP provider") + .unwrap(); + let model = ModelConfig::new("test-model"); + let messages = vec![Message::user().with_text("Interview me")]; + let mut stream = provider.stream(&model, "", &messages, &[]).await.unwrap(); + + let (message, usage) = + tokio::time::timeout(std::time::Duration::from_secs(5), stream.next()) + .await + .expect("timed out waiting for the nested form request") + .unwrap() + .unwrap(); + assert!(usage.is_none()); + let message = message.expect("expected forwarded form request"); + let MessageContent::ActionRequired(action_required) = &message.content[0] else { + panic!("expected action-required content"); + }; + let ActionRequiredData::Elicitation { + id, + message, + requested_schema, + } = &action_required.data + else { + panic!("expected elicitation content"); + }; + assert_eq!(message, "Choose priorities and add context"); + assert_eq!( + requested_schema["properties"]["priorities"]["type"], + "array" + ); + + assert!( + provider + .handle_elicitation_response( + id, + &serde_json::json!({ + "priorities": ["quality", "playfulness"], + "context": "Keep the Berd twist" + }), + &McpElicitationAction::Accept, + ) + .await + ); + + let response = tokio::time::timeout( + std::time::Duration::from_secs(5), + observed_response_rx.recv(), + ) + .await + .expect("timed out waiting for the nested form response") + .expect("nested agent should receive the form response"); + let AcpElicitationAction::Accept(accept) = response.action else { + panic!("expected accepted elicitation response"); + }; + let content = accept.content.expect("accepted response content"); + assert_eq!( + content.get("priorities"), + Some(&ElicitationContentValue::StringArray(vec![ + "quality".to_string(), + "playfulness".to_string(), + ])) + ); + assert_eq!( + content.get("context"), + Some(&ElicitationContentValue::String( + "Keep the Berd twist".to_string() + )) + ); + assert!( + tokio::time::timeout(std::time::Duration::from_secs(5), stream.next()) + .await + .expect("timed out waiting for the nested prompt to finish") + .is_none() + ); + + drop(stream); + drop(shutdown_tx); + let _ = tokio::time::timeout(std::time::Duration::from_secs(5), nested_agent).await; + drop(provider); + } + + #[tokio::test] + async fn cursor_questions_round_trip_through_standard_form_elicitation() { + use crate::acp::cursor::{ + CursorAskQuestionOutcome, CursorQuestion, CursorQuestionAnswer, CursorQuestionOption, + }; + use futures::StreamExt; + + let (provider_read, agent_write) = tokio::io::duplex(64 * 1024); + let (agent_read, provider_write) = tokio::io::duplex(64 * 1024); + let provider_transport = agent_client_protocol::ByteStreams::new( + provider_write.compat_write(), + provider_read.compat(), + ); + let agent_transport = agent_client_protocol::ByteStreams::new( + agent_write.compat_write(), + agent_read.compat(), + ); + let (observed_response_tx, mut observed_response_rx) = mpsc::unbounded_channel(); + let (shutdown_tx, shutdown_rx) = oneshot::channel::<()>(); + + let nested_agent = tokio::spawn(async move { + Agent + .builder() + .on_receive_request( + async move |request: InitializeRequest, responder, _cx| { + responder.respond( + InitializeResponse::new(request.protocol_version) + .agent_capabilities(AgentCapabilities::new()), + ) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async move |_request: NewSessionRequest, responder, _cx| { + responder.respond(NewSessionResponse::new("cursor-session")) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async move |_request: PromptRequest, responder, cx| { + let cx_for_prompt = cx.clone(); + let observed_response_tx = observed_response_tx.clone(); + cx.spawn(async move { + let response = cx_for_prompt + .send_request(CursorAskQuestionRequest { + tool_call_id: "tool-1".to_string(), + title: Some("Shape the implementation".to_string()), + questions: vec![ + CursorQuestion { + id: "direction".to_string(), + prompt: "Which direction?".to_string(), + options: vec![ + CursorQuestionOption { + id: "native".to_string(), + label: "Native ACP".to_string(), + }, + CursorQuestionOption { + id: "prose".to_string(), + label: "Prose".to_string(), + }, + ], + allow_multiple: false, + }, + CursorQuestion { + id: "priorities".to_string(), + prompt: "What matters?".to_string(), + options: vec![ + CursorQuestionOption { + id: "quality".to_string(), + label: "Quality".to_string(), + }, + CursorQuestionOption { + id: "speed".to_string(), + label: "Speed".to_string(), + }, + ], + allow_multiple: true, + }, + ], + }) + .block_task() + .await?; + observed_response_tx + .send(response) + .map_err(|_| agent_client_protocol::Error::internal_error())?; + responder.respond(PromptResponse::new(StopReason::EndTurn)) + })?; + Ok(()) + }, + agent_client_protocol::on_receive_request!(), + ) + .connect_with(agent_transport, async move |_cx: ConnectionTo| { + let _ = shutdown_rx.await; + Ok(()) + }) + .await + }); + + let provider = tokio::time::timeout( + std::time::Duration::from_secs(5), + AcpProvider::connect_with_transport_and_extensions( + "cursor-agent".to_string(), + GooseMode::Auto, + test_acp_config(HashMap::new(), None), + provider_transport, + vec![AcpClientExtension::CursorAskQuestion], + ), + ) + .await + .expect("timed out connecting the Cursor ACP provider") + .unwrap(); + let model = ModelConfig::new("auto"); + let messages = vec![Message::user().with_text("Interview me")]; + let mut stream = provider.stream(&model, "", &messages, &[]).await.unwrap(); + + let (message, usage) = + tokio::time::timeout(std::time::Duration::from_secs(5), stream.next()) + .await + .expect("timed out waiting for the Cursor question") + .unwrap() + .unwrap(); + assert!(usage.is_none()); + let message = message.expect("expected forwarded Cursor question"); + let MessageContent::ActionRequired(action_required) = &message.content[0] else { + panic!("expected action-required content"); + }; + let ActionRequiredData::Elicitation { + id, + message, + requested_schema, + } = &action_required.data + else { + panic!("expected elicitation content"); + }; + assert_eq!(message, "Shape the implementation"); + assert_eq!( + requested_schema["properties"]["direction"]["oneOf"][0]["const"], + "native" + ); + assert_eq!( + requested_schema["properties"]["priorities"]["items"]["anyOf"][1]["const"], + "speed" + ); + assert!( + !requested_schema.to_string().contains("Other"), + "the adapter must not invent a free-form option" + ); + + assert!( + provider + .handle_elicitation_response( + id, + &serde_json::json!({ + "direction": "native", + "priorities": ["quality", "speed"] + }), + &McpElicitationAction::Accept, + ) + .await + ); + + let response = tokio::time::timeout( + std::time::Duration::from_secs(5), + observed_response_rx.recv(), + ) + .await + .expect("timed out waiting for the Cursor response") + .expect("Cursor should receive the response"); + assert_eq!( + response.outcome, + CursorAskQuestionOutcome::Answered { + answers: vec![ + CursorQuestionAnswer { + question_id: "direction".to_string(), + selected_option_ids: vec!["native".to_string()], + }, + CursorQuestionAnswer { + question_id: "priorities".to_string(), + selected_option_ids: vec!["quality".to_string(), "speed".to_string()], + }, + ], + } + ); + + assert!( + tokio::time::timeout(std::time::Duration::from_secs(5), stream.next()) + .await + .expect("timed out waiting for the Cursor prompt to finish") + .is_none() + ); + drop(stream); + drop(shutdown_tx); + let _ = tokio::time::timeout(std::time::Duration::from_secs(5), nested_agent).await; + drop(provider); + } + #[test] fn prompt_auth_error_maps_to_provider_authentication() { let error = provider_error_from_acp(agent_client_protocol::Error::auth_required()); @@ -2119,6 +2907,7 @@ mod tests { response: NewSessionResponse::new("test-session"), }), pending_confirmations: Arc::new(TokioMutex::new(HashMap::new())), + pending_elicitations: Arc::new(Mutex::new(HashMap::new())), pending_tool_updates: Arc::new(Mutex::new(HashMap::new())), handoff_context_sent: Arc::new(AtomicBool::new(false)), context_size: Arc::new(AtomicU64::new(0)), @@ -3002,6 +3791,7 @@ mod tests { "acp-test".to_string(), GooseMode::Auto, test_acp_config(HashMap::new(), None), + Vec::new(), Box::new(move |_, _, init_tx, cancel_rx| { Box::pin(async move { let _init_tx = init_tx; diff --git a/crates/goose/src/acp/server.rs b/crates/goose/src/acp/server.rs index fdbe5ea6831a..c9843d708f0c 100644 --- a/crates/goose/src/acp/server.rs +++ b/crates/goose/src/acp/server.rs @@ -96,6 +96,7 @@ mod diagnostics; mod dictation; mod dispatch; mod elicitation; +use self::elicitation::FormElicitation; mod extensions; mod fork_session; mod list_sessions; @@ -1130,11 +1131,15 @@ impl GooseAcpAgent { } => { self.handle_form_elicitation( cx, - session_id, - id, - elicitation_message, - requested_schema, - message_meta_without_steer(message), + agent, + FormElicitation::new( + session_id.clone(), + id.clone(), + elicitation_message.clone(), + requested_schema.clone(), + message_meta_without_steer(message), + false, + ), ) .await?; } diff --git a/crates/goose/src/acp/server/elicitation.rs b/crates/goose/src/acp/server/elicitation.rs index 3593d88f973e..002cc1ae5c6c 100644 --- a/crates/goose/src/acp/server/elicitation.rs +++ b/crates/goose/src/acp/server/elicitation.rs @@ -11,36 +11,59 @@ use agent_client_protocol::{ use tracing::warn; use crate::action_required_manager::ElicitationOutcome; -use crate::session::SessionManager; +use crate::agents::Agent; + +pub(super) struct FormElicitation { + session_id: SessionId, + elicitation_id: String, + message: String, + requested_schema: serde_json::Value, + meta: Meta, + recovered: bool, +} + +impl FormElicitation { + pub(super) fn new( + session_id: SessionId, + elicitation_id: String, + message: String, + requested_schema: serde_json::Value, + meta: Meta, + recovered: bool, + ) -> Self { + Self { + session_id, + elicitation_id, + message, + requested_schema, + meta, + recovered, + } + } +} impl super::GooseAcpAgent { pub(super) async fn handle_form_elicitation( &self, cx: &ConnectionTo, - session_id: &SessionId, - elicitation_id: &str, - message: &str, - requested_schema: &serde_json::Value, - meta: Meta, + agent: &Arc, + elicitation: FormElicitation, ) -> Result<(), agent_client_protocol::Error> { if self.supports_acp_elicitation() { - self.send_form_elicitation( - cx, - session_id, - elicitation_id, - message, - requested_schema, - meta, - ) - .await?; + self.send_form_elicitation(cx, agent, elicitation).await?; } else { warn!( - session_id = %session_id.0.as_ref(), - elicitation_id = %elicitation_id, + session_id = %elicitation.session_id.0.as_ref(), + elicitation_id = %elicitation.elicitation_id, "ACP client does not support form elicitation" ); - self.cancel_form_elicitation(session_id.0.as_ref(), elicitation_id) - .await; + self.cancel_form_elicitation( + agent, + elicitation.session_id.0.as_ref(), + &elicitation.elicitation_id, + elicitation.recovered, + ) + .await; } Ok(()) @@ -49,15 +72,13 @@ impl super::GooseAcpAgent { async fn send_form_elicitation( &self, cx: &ConnectionTo, - session_id: &SessionId, - elicitation_id: &str, - message: &str, - requested_schema: &serde_json::Value, - meta: Meta, + agent: &Arc, + elicitation: FormElicitation, ) -> Result<(), agent_client_protocol::Error> { - let session_id = session_id.0.as_ref().to_string(); - let elicitation_id = elicitation_id.to_string(); - if requested_schema + let session_id = elicitation.session_id.0.as_ref().to_string(); + let elicitation_id = elicitation.elicitation_id; + if elicitation + .requested_schema .get("url") .and_then(|url| url.as_str()) .is_some() @@ -67,94 +88,135 @@ impl super::GooseAcpAgent { elicitation_id = %elicitation_id, "ACP URL elicitation is not supported" ); - record_acp_elicitation_response( - &self.session_manager, + finish_form_elicitation( + agent, &session_id, &elicitation_id, ElicitationOutcome::Cancel, + elicitation.recovered, ) .await; return Ok(()); } let requested_schema: ElicitationSchema = - match serde_json::from_value(requested_schema.clone()) { + match serde_json::from_value(elicitation.requested_schema) { Ok(schema) => schema, Err(error) => { - record_acp_elicitation_response( - &self.session_manager, + finish_form_elicitation( + agent, &session_id, &elicitation_id, ElicitationOutcome::Cancel, + elicitation.recovered, ) .await; return Err(agent_client_protocol::Error::internal_error() .data(format!("Failed to parse ACP elicitation schema: {error}"))); } }; + let has_live_waiter = agent + .has_pending_elicitation(&session_id, &elicitation_id) + .await; + let mut meta = elicitation.meta; + add_elicitation_meta( + &mut meta, + &elicitation_id, + elicitation.recovered, + if has_live_waiter { + "response" + } else { + "prompt" + }, + ); let request = CreateElicitationRequest::new( ElicitationFormMode::new( ElicitationSessionScope::new(session_id.clone()), requested_schema, ), - message.to_string(), + elicitation.message, ) .meta(meta); - let callback_session_manager = Arc::clone(&self.session_manager); + let callback_agent = Arc::clone(agent); let callback_session_id = session_id.clone(); let callback_elicitation_id = elicitation_id.clone(); - if let Err(error) = cx - .send_request(CreateElicitationRequestMessage(request)) + cx.send_request(CreateElicitationRequestMessage(request)) .on_receiving_result(move |result| async move { - let response = match result { - Ok(response) => elicitation_response_from_acp(response.0), + match result { + Ok(response) => { + finish_form_elicitation( + &callback_agent, + &callback_session_id, + &callback_elicitation_id, + elicitation_response_from_acp(response.0), + elicitation.recovered, + ) + .await; + } Err(error) => { warn!( error = %error, session_id = %callback_session_id, elicitation_id = %callback_elicitation_id, - "ACP elicitation request failed" + "ACP elicitation request disconnected; preserving pending response" ); - ElicitationOutcome::Cancel } - }; - - record_acp_elicitation_response( - &callback_session_manager, - &callback_session_id, - &callback_elicitation_id, - response, - ) - .await; + } Ok(()) - }) - { - record_acp_elicitation_response( - &self.session_manager, - &session_id, - &elicitation_id, - ElicitationOutcome::Cancel, - ) - .await; - return Err(error); - } + })?; Ok(()) } - async fn cancel_form_elicitation(&self, session_id: &str, elicitation_id: &str) { - record_acp_elicitation_response( - &self.session_manager, + async fn cancel_form_elicitation( + &self, + agent: &Arc, + session_id: &str, + elicitation_id: &str, + recovered: bool, + ) { + finish_form_elicitation( + agent, session_id, elicitation_id, ElicitationOutcome::Cancel, + recovered, ) .await; } } +fn add_elicitation_meta( + meta: &mut Meta, + elicitation_id: &str, + recovered: bool, + continuation: &str, +) { + // `response` means the original provider call is still blocked and this ACP + // response resumes it directly. `prompt` means only the persisted question + // survived, so the client must continue the session with a new user prompt. + let goose = meta + .entry("goose".to_string()) + .or_insert_with(|| serde_json::Value::Object(serde_json::Map::new())); + if !goose.is_object() { + *goose = serde_json::Value::Object(serde_json::Map::new()); + } + let goose = goose + .as_object_mut() + .expect("goose elicitation metadata was initialized as an object"); + goose.insert( + "elicitationId".to_string(), + serde_json::Value::String(elicitation_id.to_string()), + ); + goose.insert("recovered".to_string(), serde_json::Value::Bool(recovered)); + goose.insert( + "continuation".to_string(), + serde_json::Value::String(continuation.to_string()), + ); +} + #[derive(Debug, Clone)] struct CreateElicitationRequestMessage(CreateElicitationRequest); @@ -229,25 +291,79 @@ fn elicitation_response_from_acp(response: CreateElicitationResponse) -> Elicita } } -async fn record_acp_elicitation_response( - session_manager: &SessionManager, +async fn finish_form_elicitation( + agent: &Arc, session_id: &str, elicitation_id: &str, response: ElicitationOutcome, + recovered: bool, ) { - if let Err(error) = crate::elicitation::complete_elicitation_with_generated_message( - session_manager, - session_id, - elicitation_id, - response, - ) - .await + match agent + .submit_elicitation_response(session_id, elicitation_id, response.clone(), None) + .await { - warn!( - error = %error, - session_id = %session_id, - elicitation_id = %elicitation_id, - "Failed to record ACP elicitation response" + Ok(true) => {} + Ok(false) if recovered => { + if let Err(error) = agent + .record_recovered_elicitation_response(session_id, elicitation_id, &response) + .await + { + warn!( + error = %error, + session_id = %session_id, + elicitation_id = %elicitation_id, + "Failed to record recovered ACP elicitation response" + ); + } + } + Ok(false) => { + warn!( + session_id = %session_id, + elicitation_id = %elicitation_id, + "ACP elicitation response no longer has a live waiter" + ); + } + Err(error) => { + warn!( + error = %error, + session_id = %session_id, + elicitation_id = %elicitation_id, + "Failed to submit ACP elicitation response" + ); + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn recovered_elicitation_metadata_declares_prompt_continuation() { + let mut meta = Meta::new(); + add_elicitation_meta(&mut meta, "acp-provider:question-1", true, "prompt"); + + assert_eq!( + meta.get("goose"), + Some(&serde_json::json!({ + "elicitationId": "acp-provider:question-1", + "recovered": true, + "continuation": "prompt" + })) ); } + + #[test] + fn elicitation_metadata_preserves_existing_goose_fields() { + let mut meta = Meta::from_iter([( + "goose".to_string(), + serde_json::json!({ "messageId": "message-1" }), + )]); + add_elicitation_meta(&mut meta, "question-1", false, "response"); + + assert_eq!(meta["goose"]["messageId"], "message-1"); + assert_eq!(meta["goose"]["elicitationId"], "question-1"); + assert_eq!(meta["goose"]["recovered"], false); + assert_eq!(meta["goose"]["continuation"], "response"); + } } diff --git a/crates/goose/src/acp/server/load_session.rs b/crates/goose/src/acp/server/load_session.rs index f11bcfbfc71c..db2a5cdaf815 100644 --- a/crates/goose/src/acp/server/load_session.rs +++ b/crates/goose/src/acp/server/load_session.rs @@ -1,5 +1,6 @@ use super::message_meta::{ - content_chunk_for_message, merge_message_meta, populate_output_token_limit_content, + content_chunk_for_message, merge_message_meta, message_meta_without_steer, + populate_output_token_limit_content, }; use super::tool_calls::conversion::{ build_initial_tool_call_with_message_meta, tool_call_update_fields_from_response, @@ -47,6 +48,54 @@ fn active_turn_messages(conversation: &Conversation) -> &[Message] { .unwrap_or(messages) } +#[derive(Debug, Clone, PartialEq)] +struct PendingFormElicitation { + id: String, + message: String, + requested_schema: serde_json::Value, + meta: Meta, +} + +fn pending_form_elicitations(messages: &[Message]) -> Vec { + let answered = messages + .iter() + .flat_map(|message| &message.content) + .filter_map(|content| match content { + MessageContent::ActionRequired(action) => match &action.data { + ActionRequiredData::ElicitationResponse { id, .. } => Some(id.as_str()), + _ => None, + }, + _ => None, + }) + .collect::>(); + + messages + .iter() + .flat_map(|message| { + let answered = &answered; + message.content.iter().filter_map(move |content| { + let MessageContent::ActionRequired(action) = content else { + return None; + }; + let ActionRequiredData::Elicitation { + id, + message: elicitation_message, + requested_schema, + } = &action.data + else { + return None; + }; + (!answered.contains(id.as_str())).then(|| PendingFormElicitation { + id: id.clone(), + message: elicitation_message.clone(), + requested_schema: requested_schema.clone(), + meta: message_meta_without_steer(message), + }) + }) + }) + .collect() +} + fn send_replay_content_chunk( cx: &ConnectionTo, session_id: &SessionId, @@ -266,6 +315,38 @@ impl GooseAcpAgent { Ok(()) } + async fn resend_pending_form_elicitations( + &self, + cx: &ConnectionTo, + agent: &Arc, + session: &Session, + ) -> Result<(), agent_client_protocol::Error> { + let session_id = SessionId::new(session.id.clone()); + let messages = session + .conversation + .as_ref() + .map(active_turn_messages) + .unwrap_or(&[]); + + for pending in pending_form_elicitations(messages) { + self.handle_form_elicitation( + cx, + agent, + FormElicitation::new( + session_id.clone(), + pending.id, + pending.message, + pending.requested_schema, + pending.meta, + true, + ), + ) + .await?; + } + + Ok(()) + } + pub(super) async fn handle_load_session( &self, cx: &ConnectionTo, @@ -300,6 +381,8 @@ impl GooseAcpAgent { self.register_acp_session(session_id_str.clone(), agent.clone()) .await; self.resend_pending_tool_permissions(cx, &agent, &session)?; + self.resend_pending_form_elicitations(cx, &agent, &session) + .await?; session = self .session_manager @@ -483,4 +566,39 @@ mod tests { let no_kickoff = Conversation::new_unvalidated([approval("orphan")]); assert_eq!(active_turn_messages(&no_kickoff).len(), 1); } + + #[test] + fn pending_form_elicitations_recover_only_unanswered_current_turn_requests() { + let request = |id: &str, prompt: &str| { + Message::assistant().with_content(MessageContent::action_required_elicitation( + id.to_string(), + prompt.to_string(), + serde_json::json!({ + "type": "object", + "properties": { "answer": { "type": "string" } } + }), + )) + }; + let response = |id: &str| { + Message::user().with_content(MessageContent::action_required_elicitation_response( + id.to_string(), + serde_json::json!({ "answer": "done" }), + rmcp::model::ElicitationAction::Accept, + )) + }; + let conversation = Conversation::new_unvalidated([ + Message::user().with_text("old turn"), + request("old", "Old question"), + Message::user().with_text("current turn"), + request("answered", "Answered question"), + response("answered"), + request("pending", "Pending question"), + ]); + + let pending = pending_form_elicitations(active_turn_messages(&conversation)); + + assert_eq!(pending.len(), 1); + assert_eq!(pending[0].id, "pending"); + assert_eq!(pending[0].message, "Pending question"); + } } diff --git a/crates/goose/src/action_required_manager.rs b/crates/goose/src/action_required_manager.rs index 4cf109834f76..89e5bf96678d 100644 --- a/crates/goose/src/action_required_manager.rs +++ b/crates/goose/src/action_required_manager.rs @@ -157,6 +157,19 @@ impl ActionRequiredManager { .ok_or_else(|| anyhow::anyhow!("Request not found: {}", request_id)) } + pub(crate) async fn has_pending_response(&self, session_id: &str, request_id: &str) -> bool { + let pending = self.pending.read().await.get(request_id).cloned(); + let Some(pending) = pending else { + return false; + }; + let pending = pending.lock().await; + pending.session_id == session_id + && pending + .response_tx + .as_ref() + .is_some_and(|response_tx| !response_tx.is_closed()) + } + async fn wait_for_response( &self, request_id: &str, diff --git a/crates/goose/src/agents/agent.rs b/crates/goose/src/agents/agent.rs index 27f1f6190b89..dcedf0e6bc15 100644 --- a/crates/goose/src/agents/agent.rs +++ b/crates/goose/src/agents/agent.rs @@ -276,6 +276,7 @@ pub struct Agent { pub(super) tool_inspection_manager: ToolInspectionManager, pub(super) hook_manager: crate::hooks::HookManager, session_start_emitted: AtomicBool, + elicitation_response_lock: Mutex<()>, #[cfg(test)] pub(super) stop_hook_block_cap_override: Option, container: Mutex>, @@ -433,6 +434,7 @@ impl Agent { use_login_shell_path, ) }, + elicitation_response_lock: Mutex::new(()), session_start_emitted: AtomicBool::new(false), #[cfg(test)] stop_hook_block_cap_override: None, @@ -1801,7 +1803,23 @@ impl Agent { loop { tokio::select! { biased; - Some(event) = rx.recv() => yield event, + Some(event) = rx.recv() => { + let mut persist_result = Ok(()); + if let AgentEvent::Message(message) = &event { + if crate::acp::is_provider_form_elicitation(message) { + // Nested ACP providers stay blocked until this + // action is answered. Persist before yielding so + // the response can never sort ahead of its request. + persist_result = session_manager + .add_message(&session_id, message) + .await; + } + } + if let Err(error) = persist_result { + break Err(error); + } + yield event + }, result = &mut run => break result, } } @@ -1810,6 +1828,11 @@ impl Agent { // Without this the drain below never ends: `run` only borrows the emitter. drop(emit); while let Some(event) = rx.recv().await { + if let AgentEvent::Message(message) = &event { + if crate::acp::is_provider_form_elicitation(message) { + session_manager.add_message(&session_id, message).await?; + } + } yield event; } } @@ -1846,6 +1869,110 @@ impl Agent { Ok(Box::pin(events.map_ok(ensure_message_event_id))) } + pub(crate) async fn has_pending_elicitation( + &self, + session_id: &str, + elicitation_id: &str, + ) -> bool { + if elicitation_id.starts_with(crate::acp::ACP_PROVIDER_ELICITATION_ID_PREFIX) { + let provider = self.provider.lock().await.clone(); + return match provider { + Some(provider) => provider.has_pending_elicitation(elicitation_id).await, + None => false, + }; + } + + self.config + .session_manager + .action_required() + .has_pending_response(session_id, elicitation_id) + .await + } + + /// Persist and deliver a form response while its original waiter is still alive. + /// Returns `false` when the request was recovered from history and must continue as + /// a normal provider prompt instead of pretending the dead waiter can be reattached. + pub(crate) async fn submit_elicitation_response( + &self, + session_id: &str, + elicitation_id: &str, + response: ElicitationOutcome, + response_message: Option<&Message>, + ) -> Result { + // Checking liveness, persisting the response, and consuming the waiter must be one + // operation. Otherwise two clients can both observe a live request and persist the + // same answer before only one of them wins the response channel. + let _response_guard = self.elicitation_response_lock.lock().await; + if !self + .has_pending_elicitation(session_id, elicitation_id) + .await + { + return Ok(false); + } + + let generated_message; + let response_message = match response_message { + Some(message) => message, + None => { + generated_message = crate::elicitation::generated_elicitation_response_message( + elicitation_id, + &response, + ); + &generated_message + } + }; + + if elicitation_id.starts_with(crate::acp::ACP_PROVIDER_ELICITATION_ID_PREFIX) { + let provider = self + .provider + .lock() + .await + .clone() + .ok_or_else(|| anyhow!("Provider is not configured"))?; + self.config + .session_manager + .add_message(session_id, response_message) + .await?; + if !provider + .handle_elicitation_response( + elicitation_id, + &crate::elicitation::elicitation_response_user_data(&response), + &crate::elicitation::elicitation_response_action(&response), + ) + .await + { + return Err(anyhow!( + "ACP provider elicitation is no longer pending: {elicitation_id}" + )); + } + return Ok(true); + } + + crate::elicitation::complete_elicitation_with_message( + &self.config.session_manager, + session_id, + elicitation_id, + response, + response_message, + ) + .await?; + Ok(true) + } + + pub(crate) async fn record_recovered_elicitation_response( + &self, + session_id: &str, + elicitation_id: &str, + response: &ElicitationOutcome, + ) -> Result<()> { + let message = + crate::elicitation::generated_elicitation_response_message(elicitation_id, response); + self.config + .session_manager + .add_message(session_id, &message) + .await + } + async fn reply_impl( &self, user_message: Message, @@ -1885,18 +2012,21 @@ impl Agent { ElicitationAction::Cancel => ElicitationOutcome::Cancel, _ => ElicitationOutcome::Cancel, }; - crate::elicitation::complete_elicitation_with_message( - &session_manager, - &session_config.id, - id, - response, - &user_message, - ) - .await - .map_err(|e| { - error!("Failed to submit elicitation response: {}", e); - anyhow!("Failed to submit elicitation response: {}", e) - })?; + if !self + .submit_elicitation_response( + &session_config.id, + id, + response, + Some(&user_message), + ) + .await + .map_err(|e| { + error!("Failed to submit elicitation response: {}", e); + anyhow!("Failed to submit elicitation response: {}", e) + })? + { + return Err(anyhow!("Elicitation is no longer pending: {id}")); + } return Ok(Box::pin(futures::stream::empty())); } } @@ -2194,16 +2324,21 @@ impl Agent { let provider = self.provider().await?; let provider_name = provider.get_name().to_string(); - let saved_provider_session_id = - super::latest_provider_session_id(conversation.messages(), &provider_name); - if let Some(saved_provider_session_id) = saved_provider_session_id { - if let Err(error) = provider.resume(saved_provider_session_id).await { - warn!( - provider = provider_name, - %error, - "Could not resume provider session; continuing with a handoff" - ); - } + let saved_provider_inference = + super::latest_provider_inference(conversation.messages(), &provider_name); + if let Err(error) = provider + .prepare_session( + saved_provider_inference + .and_then(|inference| inference.provider_session_id.as_deref()), + saved_provider_inference.is_some(), + ) + .await + { + warn!( + provider = provider_name, + %error, + "Could not prepare provider session; continuing with a handoff" + ); } let requested_model = model_config.model_name.clone(); @@ -2551,6 +2686,18 @@ impl Agent { } else { response }; + let provider_form_elicitation = + crate::acp::is_provider_form_elicitation(&response); + if provider_form_elicitation { + // The provider stream remains blocked until ACP returns + // the user's answer. Persist this request before exposing + // it so the independently submitted response cannot sort + // ahead of it; skip the normal end-of-stream batch below. + session_manager + .add_message(&session_config.id, &response) + .await?; + conversation.push(response.clone()); + } surfaced_thinking_in_turn |= filtered_response.content.iter().any( |content| { @@ -2575,7 +2722,9 @@ impl Agent { if !text.is_empty() { last_assistant_text.push_str(&text); } - messages_to_add.push(response); + if !provider_form_elicitation { + messages_to_add.push(response); + } continue; } @@ -4027,11 +4176,13 @@ mod tests { ]; assert_eq!( - super::super::latest_provider_session_id(&messages, "claude-acp"), + super::super::latest_provider_inference(&messages, "claude-acp") + .and_then(|inference| inference.provider_session_id.as_deref()), Some("claude-session") ); assert_eq!( - super::super::latest_provider_session_id(&messages, "codex-acp"), + super::super::latest_provider_inference(&messages, "codex-acp") + .and_then(|inference| inference.provider_session_id.as_deref()), None ); } diff --git a/crates/goose/src/agents/mod.rs b/crates/goose/src/agents/mod.rs index 8c310f9e3d59..cf979ef6450e 100644 --- a/crates/goose/src/agents/mod.rs +++ b/crates/goose/src/agents/mod.rs @@ -38,15 +38,13 @@ pub use subagent_task_config::TaskConfig; pub use tool_execution::ToolCallContext; pub use types::{FrontendTool, RetryConfig, SessionConfig, SuccessCheck}; -fn latest_provider_session_id<'a>( +fn latest_provider_inference<'a>( messages: &'a [crate::conversation::message::Message], provider: &str, -) -> Option<&'a str> { +) -> Option<&'a crate::conversation::message::InferenceMetadata> { let inference = messages .iter() .rev() .find_map(|message| message.metadata.inference.as_ref())?; - (inference.provider == provider) - .then_some(inference.provider_session_id.as_deref()) - .flatten() + (inference.provider == provider).then_some(inference) } diff --git a/crates/goose/src/agents/state_machine/ops_llm.rs b/crates/goose/src/agents/state_machine/ops_llm.rs index 45ef245c96ae..0194b078a55a 100644 --- a/crates/goose/src/agents/state_machine/ops_llm.rs +++ b/crates/goose/src/agents/state_machine/ops_llm.rs @@ -439,17 +439,24 @@ impl Inference for InferenceRunner<'_> { .await .unwrap_or_else(|_| self.model_config.context_limit()); let provider_name = self.provider.get_name(); - if let Some(session_id) = super::super::latest_provider_session_id( + let saved_provider_inference = super::super::latest_provider_inference( conversation.messages(), provider_name, - ) { - if let Err(error) = self.provider.resume(session_id).await { - tracing::warn!( - provider = provider_name, - %error, - "Could not resume provider session; continuing with a handoff" - ); - } + ); + if let Err(error) = self + .provider + .prepare_session( + saved_provider_inference + .and_then(|inference| inference.provider_session_id.as_deref()), + saved_provider_inference.is_some(), + ) + .await + { + tracing::warn!( + provider = provider_name, + %error, + "Could not prepare provider session; continuing with a handoff" + ); } let turn = messages_since_kickoff(conversation)?; let turn_start = turn @@ -514,6 +521,7 @@ impl Inference for InferenceRunner<'_> { let mut accumulator = Conversation::empty(); let mut tool_request_ids = std::collections::HashSet::new(); + let mut saw_live_provider_elicitation = false; loop { tokio::select! { biased; @@ -552,15 +560,25 @@ impl Inference for InferenceRunner<'_> { continue; } let chunk = emit.message(chunk).await; - accumulator.push(chunk); + if crate::acp::is_provider_form_elicitation(&chunk) { + // `Agent::reply_with_state_machine` persists this before + // exposing it to ACP. The nested provider stream remains + // open while the user answers, so deferring persistence + // as a normal inference effect would put the response + // before the request (or duplicate the request). + saw_live_provider_elicitation = true; + } else { + accumulator.push(chunk); + } } } } } - let empty_response = !accumulator - .iter() - .any(|message| message.metadata.output_token_limit_reached) + let empty_response = !saw_live_provider_elicitation + && !accumulator + .iter() + .any(|message| message.metadata.output_token_limit_reached) && accumulator.iter().all(|message| { message.content.iter().all(|content| match content { MessageContent::Text(text) => text.text.trim().is_empty(), diff --git a/crates/goose/src/agents/state_machine/tests/agent_reply.rs b/crates/goose/src/agents/state_machine/tests/agent_reply.rs index a06be7797a0b..7ca1413e1c16 100644 --- a/crates/goose/src/agents/state_machine/tests/agent_reply.rs +++ b/crates/goose/src/agents/state_machine/tests/agent_reply.rs @@ -5,18 +5,98 @@ use std::sync::Arc; use std::time::Duration; use anyhow::Result; +use async_trait::async_trait; use futures::StreamExt; +use rmcp::model::{ElicitationAction, Tool}; +use serde_json::Value; +use tokio::sync::{oneshot, Mutex}; use tokio_util::sync::CancellationToken; use super::dummy_api::{DummyApi, ProviderFeatures}; +use crate::action_required_manager::ElicitationOutcome; use crate::agents::{Agent, AgentConfig, AgentEvent, GoosePlatform, SessionConfig}; use crate::config::permission::PermissionManager; use crate::config::GooseMode; -use crate::conversation::message::Message; -use crate::providers::base::Provider; +use crate::conversation::message::{ActionRequiredData, Message, MessageContent}; +use crate::providers::base::{MessageStream, Provider}; use crate::session::{SessionManager, SessionType}; +use goose_providers::errors::ProviderError; use goose_providers::model::ModelConfig; +const NESTED_ELICITATION_ID: &str = "acp-provider:state-machine-test"; + +struct NestedElicitationProvider { + pending: Arc>>>, +} + +impl NestedElicitationProvider { + fn new() -> Self { + Self { + pending: Arc::new(Mutex::new(None)), + } + } +} + +#[async_trait] +impl Provider for NestedElicitationProvider { + fn get_name(&self) -> &str { + "nested-acp-test" + } + + async fn stream( + &self, + _model_config: &ModelConfig, + _system: &str, + _messages: &[Message], + _tools: &[Tool], + ) -> Result { + let pending = self.pending.clone(); + Ok(Box::pin(async_stream::try_stream! { + let (response_tx, response_rx) = oneshot::channel(); + *pending.lock().await = Some(response_tx); + yield ( + Some( + Message::assistant() + .with_content(MessageContent::action_required_elicitation( + NESTED_ELICITATION_ID.to_string(), + "Choose a direction".to_string(), + serde_json::json!({ + "type": "object", + "properties": { "direction": { "type": "string" } }, + "required": ["direction"] + }), + )) + .user_only(), + ), + None, + ); + response_rx.await.map_err(|error| { + ProviderError::ExecutionError(format!("elicitation response dropped: {error}")) + })?; + yield (Some(Message::assistant().with_text("continued after answer")), None); + })) + } + + async fn has_pending_elicitation(&self, request_id: &str) -> bool { + request_id == NESTED_ELICITATION_ID && self.pending.lock().await.is_some() + } + + async fn handle_elicitation_response( + &self, + request_id: &str, + _user_data: &Value, + _action: &ElicitationAction, + ) -> bool { + if request_id != NESTED_ELICITATION_ID { + return false; + } + let Some(response_tx) = self.pending.lock().await.take() else { + return false; + }; + response_tx.send(()).is_ok() + } +} + async fn agent_with_dummy_api() -> Result<(Agent, Arc, String, tempfile::TempDir)> { let api = Arc::new(DummyApi::start(ProviderFeatures::default()).await); let api_client = goose_providers::api_client::ApiClient::new_with_tls( @@ -60,6 +140,33 @@ async fn agent_with_dummy_api() -> Result<(Agent, Arc, String, tempfil Ok((agent, api, session.id, temp_dir)) } +async fn agent_with_nested_elicitation_provider() -> Result<(Agent, String, tempfile::TempDir)> { + let provider: Arc = Arc::new(NestedElicitationProvider::new()); + let temp_dir = tempfile::tempdir()?; + let session_manager = Arc::new(SessionManager::new(temp_dir.path().to_path_buf())); + let session = session_manager + .create_session( + temp_dir.path().to_path_buf(), + "state-machine-elicitation".to_string(), + SessionType::Hidden, + GooseMode::Auto, + ) + .await?; + let agent = Agent::with_config(AgentConfig::new( + session_manager, + PermissionManager::instance(), + None, + GooseMode::Auto, + true, + GoosePlatform::GooseCli, + )); + agent + .update_provider(provider, ModelConfig::new("nested-acp-test"), &session.id) + .await?; + + Ok((agent, session.id, temp_dir)) +} + #[tokio::test] async fn reply_streams_the_turn_and_ends() -> Result<()> { let (agent, api, session_id, _temp_dir) = agent_with_dummy_api().await?; @@ -100,6 +207,162 @@ async fn reply_streams_the_turn_and_ends() -> Result<()> { Ok(()) } +async fn assert_nested_provider_elicitation_order(use_state_machine: bool) -> Result<()> { + let (agent, session_id, _temp_dir) = agent_with_nested_elicitation_provider().await?; + let session_config = SessionConfig { + id: session_id.clone(), + schedule_id: None, + max_turns: Some(2), + retry_config: None, + }; + let user_message = Message::user().with_text("interview me"); + let cancel = Some(CancellationToken::new()); + let stream = if use_state_machine { + agent + .reply_with_state_machine(user_message, session_config, cancel) + .await? + } else { + agent.reply(user_message, session_config, cancel).await? + }; + tokio::pin!(stream); + + let mut saw_request = false; + let mut saw_continuation = false; + while let Some(event) = stream.next().await { + let AgentEvent::Message(message) = event? else { + continue; + }; + if message.content.iter().any(|content| { + matches!( + content, + MessageContent::ActionRequired(action) + if matches!( + &action.data, + ActionRequiredData::Elicitation { id, .. } + if id == NESTED_ELICITATION_ID + ) + ) + }) { + saw_request = true; + let (first, second) = tokio::join!( + agent.submit_elicitation_response( + &session_id, + NESTED_ELICITATION_ID, + ElicitationOutcome::Accept(serde_json::json!({ + "direction": "local" + })), + None, + ), + agent.submit_elicitation_response( + &session_id, + NESTED_ELICITATION_ID, + ElicitationOutcome::Accept(serde_json::json!({ + "direction": "upstream" + })), + None, + ), + ); + assert_eq!( + [first?, second?] + .into_iter() + .filter(|submitted| *submitted) + .count(), + 1, + "only one concurrent response may consume the live elicitation" + ); + } + saw_continuation |= message.as_concat_text() == "continued after answer"; + } + + assert!(saw_request); + assert!(saw_continuation); + let session = agent + .config + .session_manager + .get_session(&session_id, true) + .await?; + let conversation = session.conversation.expect("conversation"); + let messages = conversation.messages(); + let request_index = messages + .iter() + .position(|message| { + message.content.iter().any(|content| { + matches!( + content, + MessageContent::ActionRequired(action) + if matches!( + &action.data, + ActionRequiredData::Elicitation { id, .. } + if id == NESTED_ELICITATION_ID + ) + ) + }) + }) + .expect("persisted elicitation request"); + let response_index = messages + .iter() + .position(|message| { + message.content.iter().any(|content| { + matches!( + content, + MessageContent::ActionRequired(action) + if matches!( + &action.data, + ActionRequiredData::ElicitationResponse { id, .. } + if id == NESTED_ELICITATION_ID + ) + ) + }) + }) + .expect("persisted elicitation response"); + let continuation_index = messages + .iter() + .position(|message| message.as_concat_text() == "continued after answer") + .expect("persisted continuation"); + + assert!(request_index < response_index); + assert!(response_index < continuation_index); + assert_eq!( + messages + .iter() + .filter(|message| crate::acp::is_provider_form_elicitation(message)) + .count(), + 1 + ); + assert_eq!( + messages + .iter() + .filter(|message| { + message.content.iter().any(|content| { + matches!( + content, + MessageContent::ActionRequired(action) + if matches!( + &action.data, + ActionRequiredData::ElicitationResponse { id, .. } + if id == NESTED_ELICITATION_ID + ) + ) + }) + }) + .count(), + 1 + ); + + Ok(()) +} + +#[tokio::test] +async fn nested_provider_elicitation_persists_before_response_in_state_machine() -> Result<()> { + assert_nested_provider_elicitation_order(true).await +} + +#[tokio::test] +async fn nested_provider_elicitation_persists_before_response_in_legacy_loop() -> Result<()> { + let _guard = env_lock::lock_env([("GOOSE_STATE_MACHINE", None::<&str>)]); + assert_nested_provider_elicitation_order(false).await +} + #[tokio::test] async fn bang_shell_uses_the_state_machine_when_the_flag_is_disabled() -> Result<()> { let _guard = env_lock::lock_env([("GOOSE_STATE_MACHINE", None::<&str>)]); diff --git a/crates/goose/src/elicitation.rs b/crates/goose/src/elicitation.rs index 435e26c5adbc..508120deb004 100644 --- a/crates/goose/src/elicitation.rs +++ b/crates/goose/src/elicitation.rs @@ -6,14 +6,14 @@ use crate::action_required_manager::ElicitationOutcome; use crate::conversation::message::{Message, MessageContent}; use crate::session::SessionManager; -fn elicitation_response_user_data(response: &ElicitationOutcome) -> Value { +pub(crate) fn elicitation_response_user_data(response: &ElicitationOutcome) -> Value { match response { ElicitationOutcome::Accept(user_data) => user_data.clone(), ElicitationOutcome::Decline | ElicitationOutcome::Cancel => serde_json::json!({}), } } -fn elicitation_response_action(response: &ElicitationOutcome) -> ElicitationAction { +pub(crate) fn elicitation_response_action(response: &ElicitationOutcome) -> ElicitationAction { match response { ElicitationOutcome::Accept(_) => ElicitationAction::Accept, ElicitationOutcome::Decline => ElicitationAction::Decline, @@ -21,7 +21,7 @@ fn elicitation_response_action(response: &ElicitationOutcome) -> ElicitationActi } } -fn generated_elicitation_response_message( +pub(crate) fn generated_elicitation_response_message( elicitation_id: &str, response: &ElicitationOutcome, ) -> Message { @@ -53,20 +53,3 @@ pub(crate) async fn complete_elicitation_with_message( claim.submit(response) } - -pub(crate) async fn complete_elicitation_with_generated_message( - session_manager: &SessionManager, - session_id: &str, - elicitation_id: &str, - response: ElicitationOutcome, -) -> Result<()> { - let response_message = generated_elicitation_response_message(elicitation_id, &response); - complete_elicitation_with_message( - session_manager, - session_id, - elicitation_id, - response, - &response_message, - ) - .await -} diff --git a/crates/goose/src/providers/cursor_agent.rs b/crates/goose/src/providers/cursor_agent.rs index 107880918a57..5d204e638e0d 100644 --- a/crates/goose/src/providers/cursor_agent.rs +++ b/crates/goose/src/providers/cursor_agent.rs @@ -1,20 +1,28 @@ use anyhow::Result; use async_trait::async_trait; -use rmcp::model::Role; +use rmcp::model::{ElicitationAction, Role}; use serde_json::{json, Value}; +use std::collections::HashMap; use std::path::PathBuf; use std::process::Stdio; +use std::sync::{Arc, Mutex}; use std::time::Duration; use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader}; use tokio::process::Command; use super::base::{ - stream_from_single_message, ConfigKey, MessageStream, Provider, ProviderDef, ProviderMetadata, + current_working_dir, stream_from_single_message, ConfigKey, MessageStream, PermissionRouting, + Provider, ProviderDef, ProviderMetadata, }; use super::catalog::ProviderSetupMetadata; use super::utils::filter_extensions_from_system_prompt; +use crate::acp::{ + extension_configs_to_mcp_servers, AcpClientExtension, AcpProvider, AcpProviderConfig, +}; use crate::config::search_path::SearchPaths; +use crate::config::{Config, ExtensionConfig, GooseMode}; use crate::conversation::message::{Message, MessageContent}; +use crate::permission::PermissionConfirmation; use crate::subprocess::configure_subprocess; use futures::future::BoxFuture; use goose_providers::conversation::token_usage::{ProviderUsage, Usage}; @@ -34,29 +42,132 @@ pub const CURSOR_AGENT_KNOWN_MODELS: &[&str] = &[ "composer-2.5-fast", ]; -pub const CURSOR_AGENT_DOC_URL: &str = "https://docs.cursor.com/en/cli/overview"; +pub const CURSOR_AGENT_DOC_URL: &str = "https://cursor.com/docs/cli/acp"; const CURSOR_AGENT_LIST_TIMEOUT: Duration = Duration::from_secs(10); +const STRUCTURED_ELICITATION_FALLBACK: &str = "This client is running Cursor without ACP interactive forms. If instructions ask you to use ask_user, request_user_input, or an equivalent interactive-question tool, ask the same question directly in plain prose with every choice, then stop and wait for the user's reply. Ask once and do not report a missing client tool."; -#[derive(Debug, serde::Serialize)] +#[derive(Debug)] pub struct CursorAgentProvider { command: PathBuf, - #[serde(skip)] name: String, + transport: Mutex, +} + +#[derive(Debug)] +enum CursorTransport { + Unprepared(Option>), + Acp(Arc), + Direct, + Unavailable(String), } impl CursorAgentProvider { pub async fn from_env( + extensions: Vec, + tls_config: Option, + ) -> Result { + Self::build(extensions, current_working_dir(), tls_config).await + } + + async fn build( + extensions: Vec, + working_dir: PathBuf, _tls_config: Option, ) -> Result { - let config = crate::config::Config::global(); + let config = Config::global(); let command: String = config.get_cursor_agent_command().unwrap_or_default().into(); let resolved_command = SearchPaths::builder().with_npm().resolve(&command)?; + let goose_mode = config.get_goose_mode().unwrap_or(GooseMode::Auto); + Ok(Self::build_with_command(resolved_command, extensions, working_dir, goose_mode).await) + } + + async fn build_with_command( + resolved_command: PathBuf, + extensions: Vec, + working_dir: PathBuf, + goose_mode: GooseMode, + ) -> Self { + let mode_mapping = HashMap::from([ + (GooseMode::Auto, vec!["agent".to_string()]), + (GooseMode::SmartApprove, vec!["agent".to_string()]), + (GooseMode::Approve, vec!["agent".to_string()]), + (GooseMode::Chat, vec!["ask".to_string(), "plan".to_string()]), + ]); + let provider_config = AcpProviderConfig { + command: resolved_command.clone(), + args: vec!["acp".to_string()], + env: vec![], + env_remove: vec![], + work_dir: working_dir, + mcp_servers: extension_configs_to_mcp_servers(&extensions), + session_mode_id: mode_mapping[&goose_mode].first().cloned(), + session_config_options: vec![], + model_config_option_id: Some("model".to_string()), + mode_mapping, + notification_callback: None, + }; + + let acp = match AcpProvider::connect_with_client_extensions( + CURSOR_AGENT_PROVIDER_NAME.to_string(), + goose_mode, + provider_config, + vec![AcpClientExtension::CursorAskQuestion], + ) + .await + { + Ok(provider) => Some(Arc::new(provider)), + Err(error) => { + tracing::warn!( + error = %error, + "Cursor ACP unavailable; using the direct CLI transport" + ); + None + } + }; - Ok(Self { + Self { command: resolved_command, name: CURSOR_AGENT_PROVIDER_NAME.to_string(), - }) + transport: Mutex::new(CursorTransport::Unprepared(acp)), + } + } + + fn acp_candidate(&self) -> Option> { + let transport = self.transport.lock().ok()?; + match &*transport { + CursorTransport::Unprepared(provider) => provider.clone(), + CursorTransport::Acp(provider) => Some(provider.clone()), + CursorTransport::Direct | CursorTransport::Unavailable(_) => None, + } + } + + fn selected_acp(&self) -> Option> { + let transport = self.transport.lock().ok()?; + match &*transport { + CursorTransport::Acp(provider) => Some(provider.clone()), + _ => None, + } + } + + fn selected_transport(&self) -> Result>, ProviderError> { + let mut transport = self.transport.lock().map_err(|_| { + ProviderError::RequestFailed("Cursor transport lock poisoned".to_string()) + })?; + + if let CursorTransport::Unprepared(provider) = &mut *transport { + *transport = match provider.take() { + Some(provider) => CursorTransport::Acp(provider), + None => CursorTransport::Direct, + }; + } + + match &*transport { + CursorTransport::Acp(provider) => Ok(Some(provider.clone())), + CursorTransport::Direct => Ok(None), + CursorTransport::Unavailable(error) => Err(ProviderError::RequestFailed(error.clone())), + CursorTransport::Unprepared(_) => unreachable!("Cursor transport was selected"), + } } /// Get authentication status from cursor-agent @@ -140,6 +251,8 @@ impl CursorAgentProvider { let filtered_system = filter_extensions_from_system_prompt(system); full_prompt.push_str(&filtered_system); full_prompt.push_str("\n\n"); + full_prompt.push_str(STRUCTURED_ELICITATION_FALLBACK); + full_prompt.push_str("\n\n"); // Add conversation history for message in messages { @@ -481,7 +594,8 @@ impl goose_providers::base::ProviderDescriptor for CursorAgentProvider { "cursor-agent", &["cursor-agent", "cursor_agent", "cursor"], ) - .with_docs_url("https://docs.cursor.com/en/cli/overview") + .with_acp() + .with_docs_url(CURSOR_AGENT_DOC_URL) .with_capabilities(true, true, true), ) } @@ -491,10 +605,18 @@ impl ProviderDef for CursorAgentProvider { type Provider = Self; fn from_env( - _extensions: Vec, + extensions: Vec, + tls_config: Option, + ) -> BoxFuture<'static, Result> { + Box::pin(Self::from_env(extensions, tls_config)) + } + + fn from_env_with_working_dir( + extensions: Vec, + working_dir: PathBuf, tls_config: Option, ) -> BoxFuture<'static, Result> { - Box::pin(Self::from_env(tls_config)) + Box::pin(Self::build(extensions, working_dir, tls_config)) } } @@ -504,6 +626,129 @@ impl Provider for CursorAgentProvider { &self.name } + fn provider_session_id(&self) -> Option { + self.selected_acp() + .and_then(|provider| provider.provider_session_id()) + } + + async fn resume(&self, session_id: &str) -> Result<(), ProviderError> { + match self.acp_candidate() { + Some(provider) => provider.resume(session_id).await, + None => Err(ProviderError::RequestFailed( + "Cursor ACP is unavailable for this saved session".to_string(), + )), + } + } + + async fn prepare_session( + &self, + provider_session_id: Option<&str>, + has_provider_history: bool, + ) -> Result<(), ProviderError> { + if has_provider_history && provider_session_id.is_none() { + let mut transport = self.transport.lock().map_err(|_| { + ProviderError::RequestFailed("Cursor transport lock poisoned".to_string()) + })?; + *transport = CursorTransport::Direct; + return Ok(()); + } + + let Some(session_id) = provider_session_id else { + self.selected_transport()?; + return Ok(()); + }; + + let Some(provider) = self.acp_candidate() else { + let error = + "Saved Cursor session requires ACP, but Cursor ACP is unavailable".to_string(); + if let Ok(mut transport) = self.transport.lock() { + *transport = CursorTransport::Unavailable(error.clone()); + } + return Err(ProviderError::RequestFailed(error)); + }; + + match provider.resume(session_id).await { + Ok(()) => { + let mut transport = self.transport.lock().map_err(|_| { + ProviderError::RequestFailed("Cursor transport lock poisoned".to_string()) + })?; + *transport = CursorTransport::Acp(provider); + Ok(()) + } + Err(error) => { + let message = format!("Could not resume the saved Cursor ACP session: {error}"); + if let Ok(mut transport) = self.transport.lock() { + *transport = CursorTransport::Unavailable(message.clone()); + } + Err(ProviderError::RequestFailed(message)) + } + } + } + + async fn get_context_limit(&self, model_config: &ModelConfig) -> Result { + match self.selected_acp() { + Some(provider) => provider.get_context_limit(model_config).await, + None => Ok(model_config.context_limit()), + } + } + + async fn update_mode(&self, session_id: &str, mode: GooseMode) -> Result<(), ProviderError> { + match self.selected_acp() { + Some(provider) => provider.update_mode(session_id, mode).await, + None => Ok(()), + } + } + + fn permission_routing(&self) -> PermissionRouting { + self.selected_acp() + .map_or(PermissionRouting::Noop, |provider| { + provider.permission_routing() + }) + } + + fn manages_own_context(&self) -> bool { + self.selected_acp() + .is_some_and(|provider| provider.manages_own_context()) + } + + async fn handle_permission_confirmation( + &self, + request_id: &str, + confirmation: &PermissionConfirmation, + ) -> bool { + match self.selected_acp() { + Some(provider) => { + provider + .handle_permission_confirmation(request_id, confirmation) + .await + } + None => false, + } + } + + async fn handle_elicitation_response( + &self, + request_id: &str, + user_data: &Value, + action: &ElicitationAction, + ) -> bool { + match self.selected_acp() { + Some(provider) => { + provider + .handle_elicitation_response(request_id, user_data, action) + .await + } + None => false, + } + } + + async fn has_pending_elicitation(&self, request_id: &str) -> bool { + match self.selected_acp() { + Some(provider) => provider.has_pending_elicitation(request_id).await, + None => false, + } + } + fn skip_canonical_filtering(&self) -> bool { // Cursor model IDs are CLI/account-specific and often absent from the // canonical registry. Keep the live list intact for inventory/config. @@ -511,6 +756,14 @@ impl Provider for CursorAgentProvider { } async fn fetch_supported_models(&self) -> Result, ProviderError> { + if let Some(provider) = self.acp_candidate() { + if let Ok(models) = provider.fetch_supported_models().await { + if !models.is_empty() { + return Ok(models); + } + } + } + match self.list_models_from_cli().await { Ok(models) if !models.is_empty() => Ok(models), Ok(_) => { @@ -536,6 +789,10 @@ impl Provider for CursorAgentProvider { messages: &[Message], tools: &[Tool], ) -> Result { + if let Some(provider) = self.selected_transport()? { + return provider.stream(model_config, system, messages, tools).await; + } + if super::cli_common::is_session_description_request(system) { let (message, provider_usage) = super::cli_common::generate_simple_session_description( &model_config.model_name, @@ -586,6 +843,9 @@ mod tests { &command, r#"#!/bin/sh record_dir=${0%/*} +if [ "$1" = "acp" ]; then + exit 64 +fi printf '%s\n' "$@" > "$record_dir/args" cat > "$record_dir/stdin" printf '%s\n' '{"type":"result","result":"ok"}' @@ -601,6 +861,7 @@ printf '%s\n' '{"type":"result","result":"ok"}' let provider = CursorAgentProvider { command: recording_cli(directory.path()), name: CURSOR_AGENT_PROVIDER_NAME.to_string(), + transport: Mutex::new(CursorTransport::Direct), }; let lines = provider @@ -618,6 +879,9 @@ printf '%s\n' '{"type":"result","result":"ok"}' let stdin = fs::read_to_string(directory.path().join("stdin")).unwrap(); assert!(!args.contains(SENTINEL)); assert!(stdin.contains(SENTINEL)); + assert!(stdin.contains("ask the same question directly in plain prose")); + assert!(stdin.contains("Ask once")); + assert!(stdin.contains("do not report a missing client tool")); assert!(!args.lines().any(|arg| arg == "-p")); assert!(args.contains("--model\nauto")); assert!(args.lines().any(|arg| arg == "--print")); @@ -640,6 +904,92 @@ printf '%s\n' '{"type":"result","result":"ok"}' .await; } + #[tokio::test] + async fn acp_startup_failure_selects_direct_transport_for_the_provider_instance() { + let directory = tempfile::tempdir().unwrap(); + let provider = CursorAgentProvider::build_with_command( + recording_cli(directory.path()), + vec![], + directory.path().to_path_buf(), + GooseMode::Auto, + ) + .await; + + assert!(matches!( + &*provider.transport.lock().unwrap(), + CursorTransport::Unprepared(None) + )); + provider.prepare_session(None, false).await.unwrap(); + assert!(matches!( + &*provider.transport.lock().unwrap(), + CursorTransport::Direct + )); + assert!(provider.provider_session_id().is_none()); + + let (message, _) = provider + .complete( + &ModelConfig::new(CURSOR_AGENT_DEFAULT_MODEL), + "system instructions", + &[Message::user().with_text(SENTINEL)], + &[], + ) + .await + .unwrap(); + assert_eq!(message.as_concat_text(), "ok"); + } + + #[tokio::test] + async fn historical_direct_session_remains_on_the_direct_transport() { + let directory = tempfile::tempdir().unwrap(); + let provider = CursorAgentProvider { + command: recording_cli(directory.path()), + name: CURSOR_AGENT_PROVIDER_NAME.to_string(), + transport: Mutex::new(CursorTransport::Unprepared(None)), + }; + + provider.prepare_session(None, true).await.unwrap(); + + assert!(matches!( + &*provider.transport.lock().unwrap(), + CursorTransport::Direct + )); + assert!(provider.provider_session_id().is_none()); + } + + #[tokio::test] + async fn saved_acp_session_never_falls_back_to_direct_transport() { + let directory = tempfile::tempdir().unwrap(); + let provider = CursorAgentProvider { + command: recording_cli(directory.path()), + name: CURSOR_AGENT_PROVIDER_NAME.to_string(), + transport: Mutex::new(CursorTransport::Unprepared(None)), + }; + + let error = provider + .prepare_session(Some("saved-cursor-session"), true) + .await + .unwrap_err(); + + assert!(error.to_string().contains("requires ACP")); + assert!(matches!( + &*provider.transport.lock().unwrap(), + CursorTransport::Unavailable(_) + )); + let stream_error = match provider + .stream( + &ModelConfig::new(CURSOR_AGENT_DEFAULT_MODEL), + "system instructions", + &[Message::user().with_text(SENTINEL)], + &[], + ) + .await + { + Ok(_) => panic!("saved ACP session must not fall back to direct transport"), + Err(error) => error, + }; + assert!(stream_error.to_string().contains("requires ACP")); + } + #[test] fn parse_models_output_extracts_ids_and_preserves_auto() { let stdout = r#" From e98a65e3cbba516e30c7cec330a80d86cb0e6bd4 Mon Sep 17 00:00:00 2001 From: luke Date: Wed, 19 Aug 2026 22:10:39 -0400 Subject: [PATCH 02/12] docs: state what the elicitation response lock actually guarantees The comment claimed liveness, persistence, and waiter consumption were one operation. The lock only serializes concurrent submitters; for a provider-owned request the waiter is still not reserved across the append, so a stream dropping in between can leave an answered transcript row the nested agent never received. Say so, and name the claim shape that would close it. --- crates/goose/src/agents/agent.rs | 10 +++++++--- 1 file changed, 7 insertions(+), 3 deletions(-) diff --git a/crates/goose/src/agents/agent.rs b/crates/goose/src/agents/agent.rs index dcedf0e6bc15..e4bc3370d91d 100644 --- a/crates/goose/src/agents/agent.rs +++ b/crates/goose/src/agents/agent.rs @@ -1899,9 +1899,13 @@ impl Agent { response: ElicitationOutcome, response_message: Option<&Message>, ) -> Result { - // Checking liveness, persisting the response, and consuming the waiter must be one - // operation. Otherwise two clients can both observe a live request and persist the - // same answer before only one of them wins the response channel. + // Serializes concurrent submitters: without this, two clients can both observe a + // live request and both persist the same answer while only one wins the channel. + // + // It is not a full claim. For a provider-owned request the waiter is not reserved + // across the append, so a stream that drops between the liveness check and the send + // can leave an accepted answer in history that the nested agent never received. + // `ActionRequiredManager::claim_response` is the shape this path still needs. let _response_guard = self.elicitation_response_lock.lock().await; if !self .has_pending_elicitation(session_id, elicitation_id) From 442e7e1aea8426292787185dfe9f6155f7c838ee Mon Sep 17 00:00:00 2001 From: luke Date: Thu, 20 Aug 2026 09:50:53 -0400 Subject: [PATCH 03/12] fix(acp): reserve a nested waiter before persisting its response MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Answering a nested form checked that the waiter was live, appended the response to the conversation, and only then delivered it. A stream dropping between the check and the send left an accepted answer recorded in history that the originating agent never received — and replay reads that response id as proof the question was answered, so it is not asked again. Claiming takes the waiter out of the provider's pending set, so nothing can cancel or consume it while the response is written. A claim is refused when the receiver has already closed, which is what stops a response being recorded against a request nothing is left to receive. If the append fails the waiter is released unanswered; a stream that drops while a claim is held releases it too, rather than leaving it hanging. This replaces the liveness probe on the response path — claiming subsumes it — and follows the shape ActionRequiredManager already uses for MCP-originated elicitations. --- crates/goose-provider-types/src/base.rs | 16 +++ crates/goose/src/acp/provider.rs | 119 ++++++++++++++++++++- crates/goose/src/agents/agent.rs | 35 +++--- crates/goose/src/providers/cursor_agent.rs | 13 +++ 4 files changed, 168 insertions(+), 15 deletions(-) diff --git a/crates/goose-provider-types/src/base.rs b/crates/goose-provider-types/src/base.rs index 12b2361d3d1d..32f7a6037c34 100644 --- a/crates/goose-provider-types/src/base.rs +++ b/crates/goose-provider-types/src/base.rs @@ -669,6 +669,22 @@ pub trait Provider: Send + Sync { false } + /// Reserve a live elicitation so nothing else can cancel or consume it. + /// + /// The response must be persisted before it is delivered, and the waiter + /// must not be able to disappear in between — a stream dropping between the + /// two would leave an answer recorded that the originating agent never + /// received. Claiming takes the waiter out of the provider's pending set; + /// `release_elicitation` puts it back if the response cannot be persisted. + async fn claim_elicitation(&self, _request_id: &str) -> bool { + false + } + + /// Return a claimed elicitation to the pending set unanswered. + async fn release_elicitation(&self, _request_id: &str) {} + + /// Deliver a response to a claimed elicitation. Returns false when the + /// originating agent is no longer there to receive it. async fn handle_elicitation_response( &self, _request_id: &str, diff --git a/crates/goose/src/acp/provider.rs b/crates/goose/src/acp/provider.rs index 4eccfdd74910..cada10970f75 100644 --- a/crates/goose/src/acp/provider.rs +++ b/crates/goose/src/acp/provider.rs @@ -272,6 +272,9 @@ pub struct AcpProvider { pending_confirmations: Arc>>>, pending_elicitations: Arc>>>, + /// Waiters reserved for delivery. Held out of `pending_elicitations` so no + /// other path can cancel or consume one while its response is persisted. + claimed_elicitations: Arc>>>, pending_tool_updates: Arc>>, /// True after the first ACP prompt completes with the handoff context committed. /// Failed or abandoned first prompts reset this so the next prompt can retry it. @@ -313,11 +316,15 @@ fn cancel_pending_elicitations( struct PendingElicitationGuard { pending: Arc>>>, + claimed: Arc>>>, } impl Drop for PendingElicitationGuard { fn drop(&mut self) { + // A claimed waiter is mid-delivery, but the stream carrying its agent is + // gone, so it has to be released too rather than left hanging. cancel_pending_elicitations(&self.pending); + cancel_pending_elicitations(&self.claimed); } } @@ -473,6 +480,7 @@ impl AcpProvider { session: Mutex::new(session), pending_confirmations: Arc::new(TokioMutex::new(HashMap::new())), pending_elicitations: Arc::new(Mutex::new(HashMap::new())), + claimed_elicitations: Arc::new(Mutex::new(HashMap::new())), pending_tool_updates, handoff_context_sent: Arc::new(AtomicBool::new(false)), context_size, @@ -661,6 +669,7 @@ impl Provider for AcpProvider { } cancel_pending_elicitations(&self.pending_elicitations); + cancel_pending_elicitations(&self.claimed_elicitations); let previous_session_id = self.acp_session_id(); let loaded = self @@ -773,11 +782,19 @@ impl Provider for AcpProvider { _ => AcpElicitationAction::Cancel, }; + // The waiter is normally reserved by `claim_elicitation`; fall back to + // the pending set so a caller that did not claim still behaves. let Some(response_tx) = self - .pending_elicitations + .claimed_elicitations .lock() .ok() - .and_then(|mut pending| pending.remove(request_id)) + .and_then(|mut claimed| claimed.remove(request_id)) + .or_else(|| { + self.pending_elicitations + .lock() + .ok() + .and_then(|mut pending| pending.remove(request_id)) + }) else { return false; }; @@ -791,6 +808,44 @@ impl Provider for AcpProvider { self.pending_elicitations .lock() .is_ok_and(|pending| pending.contains_key(request_id)) + || self + .claimed_elicitations + .lock() + .is_ok_and(|claimed| claimed.contains_key(request_id)) + } + + async fn claim_elicitation(&self, request_id: &str) -> bool { + let Ok(mut pending) = self.pending_elicitations.lock() else { + return false; + }; + let Some(response_tx) = pending.remove(request_id) else { + return false; + }; + if response_tx.is_closed() { + // The originating agent has already gone. Refusing the claim keeps a + // response from being recorded against a request nothing can receive. + return false; + } + match self.claimed_elicitations.lock() { + Ok(mut claimed) => { + claimed.insert(request_id.to_string(), response_tx); + true + } + Err(_) => false, + } + } + + async fn release_elicitation(&self, request_id: &str) { + let claimed = self + .claimed_elicitations + .lock() + .ok() + .and_then(|mut claimed| claimed.remove(request_id)); + if let Some(response_tx) = claimed { + if let Ok(mut pending) = self.pending_elicitations.lock() { + pending.insert(request_id.to_string(), response_tx); + } + } } async fn stream( @@ -872,6 +927,7 @@ impl Provider for AcpProvider { let pending_confirmations = self.pending_confirmations.clone(); let pending_elicitations = self.pending_elicitations.clone(); + let claimed_elicitations = self.claimed_elicitations.clone(); let goose_mode = *self .goose_mode .lock() @@ -883,6 +939,7 @@ impl Provider for AcpProvider { Ok(Box::pin(try_stream! { let _pending_elicitation_guard = PendingElicitationGuard { pending: pending_elicitations.clone(), + claimed: claimed_elicitations.clone(), }; let mut suppress_text = false; let mut bare_retry = bare_retry; @@ -2443,6 +2500,62 @@ mod tests { assert!(stream.next().await.is_none()); } + #[tokio::test] + async fn claiming_reserves_a_waiter_until_it_is_delivered_or_released() { + let (provider, _) = test_provider(); + let (response_tx, response_rx) = oneshot::channel(); + provider + .pending_elicitations + .lock() + .unwrap() + .insert("request-1".to_string(), response_tx); + + // A claim takes the waiter out of reach: nothing else can consume it, and + // a second submitter cannot claim it either. + assert!(provider.claim_elicitation("request-1").await); + assert!(!provider.claim_elicitation("request-1").await); + assert!(provider.has_pending_elicitation("request-1").await); + assert!(provider.pending_elicitations.lock().unwrap().is_empty()); + + // Releasing an unpersisted response puts it back, still unanswered. + provider.release_elicitation("request-1").await; + assert!(provider.claimed_elicitations.lock().unwrap().is_empty()); + assert!(provider.claim_elicitation("request-1").await); + + assert!( + provider + .handle_elicitation_response( + "request-1", + &serde_json::json!({ "answer": "delivered" }), + &McpElicitationAction::Accept, + ) + .await + ); + let AcpElicitationAction::Accept(accept) = response_rx.await.unwrap().action else { + panic!("expected accepted response"); + }; + assert_eq!( + accept.content.unwrap().get("answer"), + Some(&ElicitationContentValue::String("delivered".to_string())) + ); + } + + #[tokio::test] + async fn a_waiter_whose_agent_is_gone_cannot_be_claimed() { + let (provider, _) = test_provider(); + let (response_tx, response_rx) = oneshot::channel(); + provider + .pending_elicitations + .lock() + .unwrap() + .insert("request-1".to_string(), response_tx); + drop(response_rx); + + // Refusing the claim is what stops a response being recorded against a + // request that nothing is left to receive it. + assert!(!provider.claim_elicitation("request-1").await); + } + #[tokio::test] async fn invalid_elicitation_content_keeps_the_live_request_pending() { let (provider, _) = test_provider(); @@ -2494,6 +2607,7 @@ mod tests { drop(PendingElicitationGuard { pending: pending.clone(), + claimed: Arc::new(Mutex::new(HashMap::new())), }); assert!(pending.lock().unwrap().is_empty()); @@ -2908,6 +3022,7 @@ mod tests { }), pending_confirmations: Arc::new(TokioMutex::new(HashMap::new())), pending_elicitations: Arc::new(Mutex::new(HashMap::new())), + claimed_elicitations: Arc::new(Mutex::new(HashMap::new())), pending_tool_updates: Arc::new(Mutex::new(HashMap::new())), handoff_context_sent: Arc::new(AtomicBool::new(false)), context_size: Arc::new(AtomicU64::new(0)), diff --git a/crates/goose/src/agents/agent.rs b/crates/goose/src/agents/agent.rs index e4bc3370d91d..cd70c6b3ef75 100644 --- a/crates/goose/src/agents/agent.rs +++ b/crates/goose/src/agents/agent.rs @@ -1899,17 +1899,15 @@ impl Agent { response: ElicitationOutcome, response_message: Option<&Message>, ) -> Result { - // Serializes concurrent submitters: without this, two clients can both observe a - // live request and both persist the same answer while only one wins the channel. - // - // It is not a full claim. For a provider-owned request the waiter is not reserved - // across the append, so a stream that drops between the liveness check and the send - // can leave an accepted answer in history that the nested agent never received. - // `ActionRequiredManager::claim_response` is the shape this path still needs. + // Serializes concurrent submitters so two clients cannot both persist the same + // answer while only one wins the channel. let _response_guard = self.elicitation_response_lock.lock().await; - if !self - .has_pending_elicitation(session_id, elicitation_id) - .await + let provider_owned = + elicitation_id.starts_with(crate::acp::ACP_PROVIDER_ELICITATION_ID_PREFIX); + if !provider_owned + && !self + .has_pending_elicitation(session_id, elicitation_id) + .await { return Ok(false); } @@ -1926,17 +1924,28 @@ impl Agent { } }; - if elicitation_id.starts_with(crate::acp::ACP_PROVIDER_ELICITATION_ID_PREFIX) { + if provider_owned { let provider = self .provider .lock() .await .clone() .ok_or_else(|| anyhow!("Provider is not configured"))?; - self.config + // Reserve the waiter before writing anything. Claiming subsumes the liveness + // check and takes the waiter out of reach of cancellation, so the response + // cannot be recorded against a request that disappears mid-append. + if !provider.claim_elicitation(elicitation_id).await { + return Ok(false); + } + if let Err(error) = self + .config .session_manager .add_message(session_id, response_message) - .await?; + .await + { + provider.release_elicitation(elicitation_id).await; + return Err(error); + } if !provider .handle_elicitation_response( elicitation_id, diff --git a/crates/goose/src/providers/cursor_agent.rs b/crates/goose/src/providers/cursor_agent.rs index 5d204e638e0d..7329150eee96 100644 --- a/crates/goose/src/providers/cursor_agent.rs +++ b/crates/goose/src/providers/cursor_agent.rs @@ -742,6 +742,19 @@ impl Provider for CursorAgentProvider { } } + async fn claim_elicitation(&self, request_id: &str) -> bool { + match self.selected_acp() { + Some(provider) => provider.claim_elicitation(request_id).await, + None => false, + } + } + + async fn release_elicitation(&self, request_id: &str) { + if let Some(provider) = self.selected_acp() { + provider.release_elicitation(request_id).await; + } + } + async fn has_pending_elicitation(&self, request_id: &str) -> bool { match self.selected_acp() { Some(provider) => provider.has_pending_elicitation(request_id).await, From 3be2660a785a52a4dbac61a563b00e2970d08ad9 Mon Sep 17 00:00:00 2001 From: luke Date: Thu, 20 Aug 2026 09:53:34 -0400 Subject: [PATCH 04/12] test(acp): give the nested provider stub the same claim semantics MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The parity tests drive a stub provider whose elicitation methods stood in for a real one. Claiming now gates the response path, so the stub has to reserve its waiter the same way — otherwise both concurrent submitters are refused and the test proves the opposite of what it asserts. --- .../agents/state_machine/tests/agent_reply.rs | 34 +++++++++++++++++-- 1 file changed, 32 insertions(+), 2 deletions(-) diff --git a/crates/goose/src/agents/state_machine/tests/agent_reply.rs b/crates/goose/src/agents/state_machine/tests/agent_reply.rs index 7ca1413e1c16..bc6aa43b5c05 100644 --- a/crates/goose/src/agents/state_machine/tests/agent_reply.rs +++ b/crates/goose/src/agents/state_machine/tests/agent_reply.rs @@ -27,12 +27,16 @@ const NESTED_ELICITATION_ID: &str = "acp-provider:state-machine-test"; struct NestedElicitationProvider { pending: Arc>>>, + /// Mirrors the real provider: a claim moves the waiter out of `pending` so a + /// second submitter cannot take it while the first is persisting. + claimed: Arc>>>, } impl NestedElicitationProvider { fn new() -> Self { Self { pending: Arc::new(Mutex::new(None)), + claimed: Arc::new(Mutex::new(None)), } } } @@ -78,7 +82,28 @@ impl Provider for NestedElicitationProvider { } async fn has_pending_elicitation(&self, request_id: &str) -> bool { - request_id == NESTED_ELICITATION_ID && self.pending.lock().await.is_some() + request_id == NESTED_ELICITATION_ID + && (self.pending.lock().await.is_some() || self.claimed.lock().await.is_some()) + } + + async fn claim_elicitation(&self, request_id: &str) -> bool { + if request_id != NESTED_ELICITATION_ID { + return false; + } + let Some(response_tx) = self.pending.lock().await.take() else { + return false; + }; + *self.claimed.lock().await = Some(response_tx); + true + } + + async fn release_elicitation(&self, request_id: &str) { + if request_id != NESTED_ELICITATION_ID { + return; + } + if let Some(response_tx) = self.claimed.lock().await.take() { + *self.pending.lock().await = Some(response_tx); + } } async fn handle_elicitation_response( @@ -90,7 +115,12 @@ impl Provider for NestedElicitationProvider { if request_id != NESTED_ELICITATION_ID { return false; } - let Some(response_tx) = self.pending.lock().await.take() else { + let claimed = self.claimed.lock().await.take(); + let response_tx = match claimed { + Some(response_tx) => Some(response_tx), + None => self.pending.lock().await.take(), + }; + let Some(response_tx) = response_tx else { return false; }; response_tx.send(()).is_ok() From edad828b229aa592220fef9f54fa8fcd99f81fc2 Mon Sep 17 00:00:00 2001 From: luke Date: Thu, 20 Aug 2026 11:10:28 -0400 Subject: [PATCH 05/12] fix(acp): honor outer form elicitation support Only advertise ACP form elicitation to nested agents when the outer host explicitly negotiated it. This lets nested agents use their prose fallback instead of receiving a misleading cancellation from form-incapable clients. --- crates/goose/src/acp/provider.rs | 80 +++++++++++- crates/goose/src/acp/server.rs | 54 ++++++++ crates/goose/src/acp/server/providers.rs | 10 +- crates/goose/src/acp/server_factory.rs | 22 +++- crates/goose/src/providers/amp_acp.rs | 39 +++++- crates/goose/src/providers/base.rs | 41 ++++++ crates/goose/src/providers/claude_acp.rs | 39 +++++- crates/goose/src/providers/codex_acp.rs | 39 +++++- crates/goose/src/providers/copilot_acp.rs | 44 ++++++- crates/goose/src/providers/cursor_agent.rs | 65 +++++++++- crates/goose/src/providers/init.rs | 34 +++++ crates/goose/src/providers/mod.rs | 8 +- crates/goose/src/providers/pi_acp.rs | 44 ++++++- .../goose/src/providers/provider_registry.rs | 117 ++++++++++++++---- 14 files changed, 590 insertions(+), 46 deletions(-) diff --git a/crates/goose/src/acp/provider.rs b/crates/goose/src/acp/provider.rs index cada10970f75..f94ee5b2c55f 100644 --- a/crates/goose/src/acp/provider.rs +++ b/crates/goose/src/acp/provider.rs @@ -91,6 +91,9 @@ pub struct AcpProviderConfig { pub model_config_option_id: Option, pub mode_mapping: HashMap>, pub notification_callback: Option>, + /// Advertise forms only when the outer host explicitly negotiated support; + /// entry points without a known host capability leave this false. + pub supports_form_elicitation: bool, } enum ClientRequest { @@ -1658,8 +1661,12 @@ async fn handle_requests( ) -> Result<(), agent_client_protocol::Error> { let mut init_tx = Some(init_tx); - let client_capabilities = ClientCapabilities::new() - .elicitation(ElicitationCapabilities::new().form(ElicitationFormCapabilities::new())); + let client_capabilities = if config.supports_form_elicitation { + ClientCapabilities::new() + .elicitation(ElicitationCapabilities::new().form(ElicitationFormCapabilities::new())) + } else { + ClientCapabilities::new() + }; let init_response: InitializeResponse = cx .send_request( InitializeRequest::new(ProtocolVersion::V1).client_capabilities(client_capabilities), @@ -2794,6 +2801,74 @@ mod tests { drop(provider); } + #[tokio::test] + async fn nested_acp_omits_forms_when_the_outer_host_did_not_advertise_them() { + let (provider_read, agent_write) = tokio::io::duplex(64 * 1024); + let (agent_read, provider_write) = tokio::io::duplex(64 * 1024); + let provider_transport = agent_client_protocol::ByteStreams::new( + provider_write.compat_write(), + provider_read.compat(), + ); + let agent_transport = agent_client_protocol::ByteStreams::new( + agent_write.compat_write(), + agent_read.compat(), + ); + let (observed_tx, mut observed_rx) = mpsc::unbounded_channel(); + let (shutdown_tx, shutdown_rx) = oneshot::channel::<()>(); + + let nested_agent = tokio::spawn(async move { + Agent + .builder() + .on_receive_request( + async move |request: InitializeRequest, responder, _cx| { + let supports_forms = request + .client_capabilities + .elicitation + .as_ref() + .and_then(|elicitation| elicitation.form.as_ref()) + .is_some(); + let _ = observed_tx.send(supports_forms); + responder.respond( + InitializeResponse::new(request.protocol_version) + .agent_capabilities(AgentCapabilities::new()), + ) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async move |_request: NewSessionRequest, responder, _cx| { + responder.respond(NewSessionResponse::new("nested-session")) + }, + agent_client_protocol::on_receive_request!(), + ) + .connect_with(agent_transport, async move |_cx: ConnectionTo| { + let _ = shutdown_rx.await; + Ok(()) + }) + .await + }); + + let mut config = test_acp_config(HashMap::new(), None); + config.supports_form_elicitation = false; + let provider = tokio::time::timeout( + std::time::Duration::from_secs(5), + AcpProvider::connect_with_transport( + "nested-acp-test".to_string(), + GooseMode::Auto, + config, + provider_transport, + ), + ) + .await + .expect("timed out connecting the nested ACP provider") + .unwrap(); + + assert!(!observed_rx.recv().await.unwrap()); + drop(shutdown_tx); + let _ = tokio::time::timeout(std::time::Duration::from_secs(5), nested_agent).await; + drop(provider); + } + #[tokio::test] async fn cursor_questions_round_trip_through_standard_form_elicitation() { use crate::acp::cursor::{ @@ -3895,6 +3970,7 @@ mod tests { model_config_option_id: None, mode_mapping, notification_callback: None, + supports_form_elicitation: true, } } diff --git a/crates/goose/src/acp/server.rs b/crates/goose/src/acp/server.rs index c9843d708f0c..f066832d3344 100644 --- a/crates/goose/src/acp/server.rs +++ b/crates/goose/src/acp/server.rs @@ -123,6 +123,7 @@ pub type AcpProviderFactory = Arc< Vec, Option, bool, + crate::providers::base::ProviderHostCapabilities, ) -> BoxFuture<'static, Result>> + Send + Sync, @@ -662,6 +663,12 @@ impl GooseAcpAgent { .unwrap_or(false) } + fn provider_host_capabilities(&self) -> crate::providers::base::ProviderHostCapabilities { + crate::providers::base::ProviderHostCapabilities { + supports_form_elicitation: self.supports_acp_elicitation(), + } + } + // TODO: goose reads Paths::in_state_dir globally (e.g. RequestLog), ignoring this data_dir. pub async fn new(options: GooseAcpAgentOptions) -> Result { let session_manager = Arc::new(SessionManager::new(options.data_dir)); @@ -726,6 +733,7 @@ impl GooseAcpAgent { extensions, working_dir, use_default_model, + self.provider_host_capabilities(), ) .await } @@ -2824,6 +2832,52 @@ print(\"hello, world\") )); } + #[tokio::test] + async fn provider_factory_receives_the_outer_hosts_form_capability() { + let root = tempfile::tempdir().unwrap(); + let (observed_tx, mut observed_rx) = tokio::sync::mpsc::unbounded_channel(); + let provider_factory: AcpProviderFactory = Arc::new( + move |_provider_name, + _extensions, + _working_dir, + _use_default_model, + host_capabilities| { + let observed_tx = observed_tx.clone(); + Box::pin(async move { + observed_tx.send(host_capabilities).unwrap(); + Err(anyhow::anyhow!("capability observation only")) + }) + }, + ); + let agent = GooseAcpAgent::new(GooseAcpAgentOptions { + provider_factory, + builtin_selection: AcpBuiltinSelection::default(), + data_dir: root.path().to_path_buf(), + config_dir: root.path().to_path_buf(), + disable_session_naming: true, + goose_platform: GoosePlatform::GooseCli, + additional_source_roots: Vec::new(), + scheduler: None, + }) + .await + .unwrap(); + + agent + .on_initialize(InitializeRequest::new( + agent_client_protocol::schema::ProtocolVersion::V1, + )) + .await + .unwrap(); + let error = agent + .create_provider("test", Vec::new(), None, false) + .await + .err() + .expect("the observation provider should stop after recording capabilities"); + + assert_eq!(error.to_string(), "capability observation only"); + assert!(!observed_rx.recv().await.unwrap().supports_form_elicitation); + } + #[test] fn test_agent_capabilities_advertise_recipe_parameter_scopes() { assert_eq!( diff --git a/crates/goose/src/acp/server/providers.rs b/crates/goose/src/acp/server/providers.rs index 235aaae9d9ba..35e0dfc1e1c3 100644 --- a/crates/goose/src/acp/server/providers.rs +++ b/crates/goose/src/acp/server/providers.rs @@ -814,12 +814,20 @@ impl GooseAcpAgent { for refresh_job in refresh_plan.started.iter().cloned() { let provider_inventory = self.provider_inventory.clone(); let provider_factory = Arc::clone(&self.provider_factory); + let host_capabilities = self.provider_host_capabilities(); let provider_id = refresh_job.provider_id.clone(); let identity = refresh_job.identity.clone(); tokio::spawn(async move { let mut refresh_guard = provider_inventory.refresh_guard(&identity); let provider_result = AssertUnwindSafe(async { - provider_factory(provider_id.clone(), Vec::new(), None, true).await + provider_factory( + provider_id.clone(), + Vec::new(), + None, + true, + host_capabilities, + ) + .await }) .catch_unwind() .await; diff --git a/crates/goose/src/acp/server_factory.rs b/crates/goose/src/acp/server_factory.rs index 08e288ab1fc7..7b8e5da3f4fc 100644 --- a/crates/goose/src/acp/server_factory.rs +++ b/crates/goose/src/acp/server_factory.rs @@ -70,22 +70,34 @@ impl AcpServer { } let provider_factory: AcpProviderFactory = Arc::new( - move |provider_name, extensions, working_dir, use_default_model| { + move |provider_name, extensions, working_dir, use_default_model, host_capabilities| { Box::pin(async move { if use_default_model { - crate::providers::create_with_default_model(&provider_name, extensions) - .await + crate::providers::create_with_default_model_and_host_capabilities( + &provider_name, + extensions, + host_capabilities, + ) + .await } else { match working_dir { Some(working_dir) => { - crate::providers::create_with_working_dir( + crate::providers::create_with_working_dir_and_host_capabilities( &provider_name, extensions, working_dir, + host_capabilities, + ) + .await + } + None => { + crate::providers::create_with_host_capabilities( + &provider_name, + extensions, + host_capabilities, ) .await } - None => crate::providers::create(&provider_name, extensions).await, } } }) diff --git a/crates/goose/src/providers/amp_acp.rs b/crates/goose/src/providers/amp_acp.rs index 709b1dffd5a0..4ec72c4d3bd9 100644 --- a/crates/goose/src/providers/amp_acp.rs +++ b/crates/goose/src/providers/amp_acp.rs @@ -9,7 +9,8 @@ use crate::acp::{ use crate::config::search_path::SearchPaths; use crate::config::{Config, GooseMode}; use crate::providers::base::{ - current_working_dir, ProviderDef, ProviderDescriptor, ProviderMetadata, + current_working_dir, ProviderDef, ProviderDescriptor, ProviderHostCapabilities, + ProviderMetadata, }; use crate::providers::catalog::ProviderSetupMetadata; @@ -56,9 +57,36 @@ impl ProviderDef for AmpAcpProvider { } fn from_env_with_working_dir( + extensions: Vec, + working_dir: PathBuf, + tls_config: Option, + ) -> BoxFuture<'static, Result> { + Self::from_env_with_working_dir_and_host_capabilities( + extensions, + working_dir, + tls_config, + ProviderHostCapabilities::default(), + ) + } + + fn from_env_with_host_capabilities( + extensions: Vec, + tls_config: Option, + host_capabilities: ProviderHostCapabilities, + ) -> BoxFuture<'static, Result> { + Self::from_env_with_working_dir_and_host_capabilities( + extensions, + current_working_dir(), + tls_config, + host_capabilities, + ) + } + + fn from_env_with_working_dir_and_host_capabilities( extensions: Vec, working_dir: PathBuf, _tls_config: Option, + host_capabilities: ProviderHostCapabilities, ) -> BoxFuture<'static, Result> { Box::pin(async move { let config = Config::global(); @@ -86,10 +114,19 @@ impl ProviderDef for AmpAcpProvider { model_config_option_id: None, mode_mapping, notification_callback: None, + supports_form_elicitation: host_capabilities.supports_form_elicitation, }; let metadata = Self::metadata(); AcpProvider::connect(metadata.name, goose_mode, provider_config).await }) } + + fn from_env_with_default_model_and_host_capabilities( + extensions: Vec, + tls_config: Option, + host_capabilities: ProviderHostCapabilities, + ) -> BoxFuture<'static, Result> { + Self::from_env_with_host_capabilities(extensions, tls_config, host_capabilities) + } } diff --git a/crates/goose/src/providers/base.rs b/crates/goose/src/providers/base.rs index 2a5cfcfb5275..51660361acee 100644 --- a/crates/goose/src/providers/base.rs +++ b/crates/goose/src/providers/base.rs @@ -28,6 +28,13 @@ pub(crate) fn current_working_dir() -> PathBuf { std::env::current_dir().unwrap_or_else(|_| PathBuf::from(".")) } +/// Capabilities explicitly negotiated by the host constructing this provider. +/// Non-hosted entry points use the false-by-default values rather than guessing. +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] +pub struct ProviderHostCapabilities { + pub supports_form_elicitation: bool, +} + pub trait ProviderDef: ProviderDescriptor + Send + Sync { type Provider: Provider + 'static; @@ -49,6 +56,29 @@ pub trait ProviderDef: ProviderDescriptor + Send + Sync { Self::from_env(extensions, tls_config) } + fn from_env_with_host_capabilities( + extensions: Vec, + tls_config: Option, + _host_capabilities: ProviderHostCapabilities, + ) -> BoxFuture<'static, Result> + where + Self: Sized, + { + Self::from_env(extensions, tls_config) + } + + fn from_env_with_working_dir_and_host_capabilities( + extensions: Vec, + working_dir: PathBuf, + tls_config: Option, + _host_capabilities: ProviderHostCapabilities, + ) -> BoxFuture<'static, Result> + where + Self: Sized, + { + Self::from_env_with_working_dir(extensions, working_dir, tls_config) + } + fn from_env_with_default_model( extensions: Vec, tls_config: Option, @@ -58,4 +88,15 @@ pub trait ProviderDef: ProviderDescriptor + Send + Sync { { Self::from_env(extensions, tls_config) } + + fn from_env_with_default_model_and_host_capabilities( + extensions: Vec, + tls_config: Option, + _host_capabilities: ProviderHostCapabilities, + ) -> BoxFuture<'static, Result> + where + Self: Sized, + { + Self::from_env_with_default_model(extensions, tls_config) + } } diff --git a/crates/goose/src/providers/claude_acp.rs b/crates/goose/src/providers/claude_acp.rs index fe4680ad6502..307767871074 100644 --- a/crates/goose/src/providers/claude_acp.rs +++ b/crates/goose/src/providers/claude_acp.rs @@ -9,7 +9,8 @@ use crate::acp::{ use crate::config::search_path::SearchPaths; use crate::config::{Config, GooseMode}; use crate::providers::base::{ - current_working_dir, ProviderDef, ProviderDescriptor, ProviderMetadata, + current_working_dir, ProviderDef, ProviderDescriptor, ProviderHostCapabilities, + ProviderMetadata, }; use crate::providers::catalog::ProviderSetupMetadata; @@ -57,9 +58,36 @@ impl ProviderDef for ClaudeAcpProvider { } fn from_env_with_working_dir( + extensions: Vec, + working_dir: PathBuf, + tls_config: Option, + ) -> BoxFuture<'static, Result> { + Self::from_env_with_working_dir_and_host_capabilities( + extensions, + working_dir, + tls_config, + ProviderHostCapabilities::default(), + ) + } + + fn from_env_with_host_capabilities( + extensions: Vec, + tls_config: Option, + host_capabilities: ProviderHostCapabilities, + ) -> BoxFuture<'static, Result> { + Self::from_env_with_working_dir_and_host_capabilities( + extensions, + current_working_dir(), + tls_config, + host_capabilities, + ) + } + + fn from_env_with_working_dir_and_host_capabilities( extensions: Vec, working_dir: PathBuf, _tls_config: Option, + host_capabilities: ProviderHostCapabilities, ) -> BoxFuture<'static, Result> { Box::pin(async move { let config = Config::global(); @@ -96,10 +124,19 @@ impl ProviderDef for ClaudeAcpProvider { model_config_option_id: Some("model".to_string()), mode_mapping, notification_callback: None, + supports_form_elicitation: host_capabilities.supports_form_elicitation, }; let metadata = Self::metadata(); AcpProvider::connect(metadata.name, goose_mode, provider_config).await }) } + + fn from_env_with_default_model_and_host_capabilities( + extensions: Vec, + tls_config: Option, + host_capabilities: ProviderHostCapabilities, + ) -> BoxFuture<'static, Result> { + Self::from_env_with_host_capabilities(extensions, tls_config, host_capabilities) + } } diff --git a/crates/goose/src/providers/codex_acp.rs b/crates/goose/src/providers/codex_acp.rs index 29d618af5a2d..3d4ca613d21a 100644 --- a/crates/goose/src/providers/codex_acp.rs +++ b/crates/goose/src/providers/codex_acp.rs @@ -9,7 +9,8 @@ use crate::acp::{ use crate::config::search_path::SearchPaths; use crate::config::{Config, GooseMode}; use crate::providers::base::{ - current_working_dir, ProviderDef, ProviderDescriptor, ProviderMetadata, + current_working_dir, ProviderDef, ProviderDescriptor, ProviderHostCapabilities, + ProviderMetadata, }; use crate::providers::catalog::ProviderSetupMetadata; @@ -55,9 +56,36 @@ impl ProviderDef for CodexAcpProvider { } fn from_env_with_working_dir( + extensions: Vec, + working_dir: PathBuf, + tls_config: Option, + ) -> BoxFuture<'static, Result> { + Self::from_env_with_working_dir_and_host_capabilities( + extensions, + working_dir, + tls_config, + ProviderHostCapabilities::default(), + ) + } + + fn from_env_with_host_capabilities( + extensions: Vec, + tls_config: Option, + host_capabilities: ProviderHostCapabilities, + ) -> BoxFuture<'static, Result> { + Self::from_env_with_working_dir_and_host_capabilities( + extensions, + current_working_dir(), + tls_config, + host_capabilities, + ) + } + + fn from_env_with_working_dir_and_host_capabilities( extensions: Vec, working_dir: PathBuf, _tls_config: Option, + host_capabilities: ProviderHostCapabilities, ) -> BoxFuture<'static, Result> { Box::pin(async move { let config = Config::global(); @@ -87,10 +115,19 @@ impl ProviderDef for CodexAcpProvider { model_config_option_id: Some("model".to_string()), mode_mapping, notification_callback: None, + supports_form_elicitation: host_capabilities.supports_form_elicitation, }; let metadata = Self::metadata(); AcpProvider::connect(metadata.name, goose_mode, provider_config).await }) } + + fn from_env_with_default_model_and_host_capabilities( + extensions: Vec, + tls_config: Option, + host_capabilities: ProviderHostCapabilities, + ) -> BoxFuture<'static, Result> { + Self::from_env_with_host_capabilities(extensions, tls_config, host_capabilities) + } } diff --git a/crates/goose/src/providers/copilot_acp.rs b/crates/goose/src/providers/copilot_acp.rs index 543f98f40403..8802477ff398 100644 --- a/crates/goose/src/providers/copilot_acp.rs +++ b/crates/goose/src/providers/copilot_acp.rs @@ -10,7 +10,8 @@ use crate::acp::{ use crate::config::search_path::SearchPaths; use crate::config::{Config, GooseMode}; use crate::providers::base::{ - current_working_dir, ProviderDef, ProviderDescriptor, ProviderMetadata, + current_working_dir, ProviderDef, ProviderDescriptor, ProviderHostCapabilities, + ProviderMetadata, }; use crate::providers::catalog::ProviderSetupMetadata; @@ -55,6 +56,7 @@ impl CopilotAcpProvider { extensions: Vec, working_dir: PathBuf, use_default_model: bool, + host_capabilities: ProviderHostCapabilities, ) -> BoxFuture<'static, Result> { Box::pin(async move { let config = Config::global(); @@ -95,6 +97,7 @@ impl CopilotAcpProvider { model_config_option_id: Some("model".to_string()), mode_mapping, notification_callback: None, + supports_form_elicitation: host_capabilities.supports_form_elicitation, }; let metadata = Self::metadata(); @@ -118,13 +121,48 @@ impl ProviderDef for CopilotAcpProvider { working_dir: PathBuf, _tls_config: Option, ) -> BoxFuture<'static, Result> { - Self::create(extensions, working_dir, false) + Self::create( + extensions, + working_dir, + false, + ProviderHostCapabilities::default(), + ) } fn from_env_with_default_model( extensions: Vec, _tls_config: Option, ) -> BoxFuture<'static, Result> { - Self::create(extensions, current_working_dir(), true) + Self::create( + extensions, + current_working_dir(), + true, + ProviderHostCapabilities::default(), + ) + } + + fn from_env_with_host_capabilities( + extensions: Vec, + _tls_config: Option, + host_capabilities: ProviderHostCapabilities, + ) -> BoxFuture<'static, Result> { + Self::create(extensions, current_working_dir(), false, host_capabilities) + } + + fn from_env_with_working_dir_and_host_capabilities( + extensions: Vec, + working_dir: PathBuf, + _tls_config: Option, + host_capabilities: ProviderHostCapabilities, + ) -> BoxFuture<'static, Result> { + Self::create(extensions, working_dir, false, host_capabilities) + } + + fn from_env_with_default_model_and_host_capabilities( + extensions: Vec, + _tls_config: Option, + host_capabilities: ProviderHostCapabilities, + ) -> BoxFuture<'static, Result> { + Self::create(extensions, current_working_dir(), true, host_capabilities) } } diff --git a/crates/goose/src/providers/cursor_agent.rs b/crates/goose/src/providers/cursor_agent.rs index 7329150eee96..ca68724f0807 100644 --- a/crates/goose/src/providers/cursor_agent.rs +++ b/crates/goose/src/providers/cursor_agent.rs @@ -12,7 +12,7 @@ use tokio::process::Command; use super::base::{ current_working_dir, stream_from_single_message, ConfigKey, MessageStream, PermissionRouting, - Provider, ProviderDef, ProviderMetadata, + Provider, ProviderDef, ProviderHostCapabilities, ProviderMetadata, }; use super::catalog::ProviderSetupMetadata; use super::utils::filter_extensions_from_system_prompt; @@ -67,19 +67,33 @@ impl CursorAgentProvider { extensions: Vec, tls_config: Option, ) -> Result { - Self::build(extensions, current_working_dir(), tls_config).await + Self::build( + extensions, + current_working_dir(), + tls_config, + ProviderHostCapabilities::default(), + ) + .await } async fn build( extensions: Vec, working_dir: PathBuf, _tls_config: Option, + host_capabilities: ProviderHostCapabilities, ) -> Result { let config = Config::global(); let command: String = config.get_cursor_agent_command().unwrap_or_default().into(); let resolved_command = SearchPaths::builder().with_npm().resolve(&command)?; let goose_mode = config.get_goose_mode().unwrap_or(GooseMode::Auto); - Ok(Self::build_with_command(resolved_command, extensions, working_dir, goose_mode).await) + Ok(Self::build_with_command( + resolved_command, + extensions, + working_dir, + goose_mode, + host_capabilities, + ) + .await) } async fn build_with_command( @@ -87,6 +101,7 @@ impl CursorAgentProvider { extensions: Vec, working_dir: PathBuf, goose_mode: GooseMode, + host_capabilities: ProviderHostCapabilities, ) -> Self { let mode_mapping = HashMap::from([ (GooseMode::Auto, vec!["agent".to_string()]), @@ -106,6 +121,7 @@ impl CursorAgentProvider { model_config_option_id: Some("model".to_string()), mode_mapping, notification_callback: None, + supports_form_elicitation: host_capabilities.supports_form_elicitation, }; let acp = match AcpProvider::connect_with_client_extensions( @@ -611,12 +627,52 @@ impl ProviderDef for CursorAgentProvider { Box::pin(Self::from_env(extensions, tls_config)) } + fn from_env_with_host_capabilities( + extensions: Vec, + tls_config: Option, + host_capabilities: ProviderHostCapabilities, + ) -> BoxFuture<'static, Result> { + Box::pin(Self::build( + extensions, + current_working_dir(), + tls_config, + host_capabilities, + )) + } + fn from_env_with_working_dir( extensions: Vec, working_dir: PathBuf, tls_config: Option, ) -> BoxFuture<'static, Result> { - Box::pin(Self::build(extensions, working_dir, tls_config)) + Box::pin(Self::build( + extensions, + working_dir, + tls_config, + ProviderHostCapabilities::default(), + )) + } + + fn from_env_with_working_dir_and_host_capabilities( + extensions: Vec, + working_dir: PathBuf, + tls_config: Option, + host_capabilities: ProviderHostCapabilities, + ) -> BoxFuture<'static, Result> { + Box::pin(Self::build( + extensions, + working_dir, + tls_config, + host_capabilities, + )) + } + + fn from_env_with_default_model_and_host_capabilities( + extensions: Vec, + tls_config: Option, + host_capabilities: ProviderHostCapabilities, + ) -> BoxFuture<'static, Result> { + Self::from_env_with_host_capabilities(extensions, tls_config, host_capabilities) } } @@ -925,6 +981,7 @@ printf '%s\n' '{"type":"result","result":"ok"}' vec![], directory.path().to_path_buf(), GooseMode::Auto, + ProviderHostCapabilities::default(), ) .await; diff --git a/crates/goose/src/providers/init.rs b/crates/goose/src/providers/init.rs index b195f94921ef..4bcb91cb8013 100644 --- a/crates/goose/src/providers/init.rs +++ b/crates/goose/src/providers/init.rs @@ -266,6 +266,17 @@ pub async fn create(name: &str, extensions: Vec) -> Result, + host_capabilities: super::base::ProviderHostCapabilities, +) -> Result> { + let entry = get_from_registry(name).await?; + entry + .create_with_host_capabilities(extensions, host_capabilities) + .await +} + pub async fn create_with_working_dir( name: &str, extensions: Vec, @@ -275,6 +286,18 @@ pub async fn create_with_working_dir( entry.create_with_working_dir(extensions, working_dir).await } +pub async fn create_with_working_dir_and_host_capabilities( + name: &str, + extensions: Vec, + working_dir: PathBuf, + host_capabilities: super::base::ProviderHostCapabilities, +) -> Result> { + let entry = get_from_registry(name).await?; + entry + .create_with_working_dir_and_host_capabilities(extensions, working_dir, host_capabilities) + .await +} + pub async fn create_with_default_model( name: impl AsRef, extensions: Vec, @@ -285,6 +308,17 @@ pub async fn create_with_default_model( .await } +pub async fn create_with_default_model_and_host_capabilities( + name: impl AsRef, + extensions: Vec, + host_capabilities: super::base::ProviderHostCapabilities, +) -> Result> { + get_from_registry(name.as_ref()) + .await? + .create_with_default_model_and_host_capabilities(extensions, host_capabilities) + .await +} + pub async fn cleanup_provider(name: &str) -> Result<()> { let cleanup_fn = { let registry = get_registry().await.read().unwrap(); diff --git a/crates/goose/src/providers/mod.rs b/crates/goose/src/providers/mod.rs index 11243e5e8fe8..200b6b1b536b 100644 --- a/crates/goose/src/providers/mod.rs +++ b/crates/goose/src/providers/mod.rs @@ -93,8 +93,10 @@ pub mod xai; pub mod xai_oauth; pub use init::{ - cleanup_provider, create, create_with_default_model, create_with_named_model, - create_with_working_dir, get_from_registry, inventory_identity, providers, - refresh_custom_providers, + cleanup_provider, create, create_with_default_model, + create_with_default_model_and_host_capabilities, create_with_host_capabilities, + create_with_named_model, create_with_working_dir, + create_with_working_dir_and_host_capabilities, get_from_registry, inventory_identity, + providers, refresh_custom_providers, }; pub use retry::{retry_operation, RetryConfig}; diff --git a/crates/goose/src/providers/pi_acp.rs b/crates/goose/src/providers/pi_acp.rs index 038c27351d38..80a88cffedb0 100644 --- a/crates/goose/src/providers/pi_acp.rs +++ b/crates/goose/src/providers/pi_acp.rs @@ -10,7 +10,8 @@ use crate::acp::{ use crate::config::search_path::SearchPaths; use crate::config::{Config, GooseMode}; use crate::providers::base::{ - current_working_dir, ProviderDef, ProviderDescriptor, ProviderMetadata, + current_working_dir, ProviderDef, ProviderDescriptor, ProviderHostCapabilities, + ProviderMetadata, }; use crate::providers::catalog::ProviderSetupMetadata; @@ -50,6 +51,7 @@ impl PiAcpProvider { extensions: Vec, working_dir: PathBuf, use_default_model: bool, + host_capabilities: ProviderHostCapabilities, ) -> BoxFuture<'static, Result> { Box::pin(async move { let config = Config::global(); @@ -79,6 +81,7 @@ impl PiAcpProvider { model_config_option_id: Some("model".to_string()), mode_mapping: HashMap::new(), notification_callback: None, + supports_form_elicitation: host_capabilities.supports_form_elicitation, }; let metadata = Self::metadata(); @@ -102,13 +105,48 @@ impl ProviderDef for PiAcpProvider { working_dir: PathBuf, _tls_config: Option, ) -> BoxFuture<'static, Result> { - Self::create(extensions, working_dir, false) + Self::create( + extensions, + working_dir, + false, + ProviderHostCapabilities::default(), + ) } fn from_env_with_default_model( extensions: Vec, _tls_config: Option, ) -> BoxFuture<'static, Result> { - Self::create(extensions, current_working_dir(), true) + Self::create( + extensions, + current_working_dir(), + true, + ProviderHostCapabilities::default(), + ) + } + + fn from_env_with_host_capabilities( + extensions: Vec, + _tls_config: Option, + host_capabilities: ProviderHostCapabilities, + ) -> BoxFuture<'static, Result> { + Self::create(extensions, current_working_dir(), false, host_capabilities) + } + + fn from_env_with_working_dir_and_host_capabilities( + extensions: Vec, + working_dir: PathBuf, + _tls_config: Option, + host_capabilities: ProviderHostCapabilities, + ) -> BoxFuture<'static, Result> { + Self::create(extensions, working_dir, false, host_capabilities) + } + + fn from_env_with_default_model_and_host_capabilities( + extensions: Vec, + _tls_config: Option, + host_capabilities: ProviderHostCapabilities, + ) -> BoxFuture<'static, Result> { + Self::create(extensions, current_working_dir(), true, host_capabilities) } } diff --git a/crates/goose/src/providers/provider_registry.rs b/crates/goose/src/providers/provider_registry.rs index c67a24227baf..a64caae8ebd3 100644 --- a/crates/goose/src/providers/provider_registry.rs +++ b/crates/goose/src/providers/provider_registry.rs @@ -1,5 +1,8 @@ use super::api_client::TlsConfig; -use super::base::{ConfigKey, ModelInfo, Provider, ProviderDef, ProviderMetadata, ProviderType}; +use super::base::{ + ConfigKey, ModelInfo, Provider, ProviderDef, ProviderHostCapabilities, ProviderMetadata, + ProviderType, +}; use super::inventory::{InventoryIdentityInput, InventoryRegistration, InventoryResolvers}; use crate::config::{DeclarativeProviderConfig, ExtensionConfig}; use anyhow::Result; @@ -15,6 +18,7 @@ pub type ProviderConstructor = Arc< Option, Option, bool, + ProviderHostCapabilities, ) -> BoxFuture<'static, Result>> + Send + Sync, @@ -81,23 +85,73 @@ impl ProviderEntry { &self, extensions: Vec, ) -> Result> { - (self.constructor)(extensions, None, self.tls_config.clone(), true).await + self.create_with_default_model_and_host_capabilities( + extensions, + ProviderHostCapabilities::default(), + ) + .await + } + + pub async fn create_with_default_model_and_host_capabilities( + &self, + extensions: Vec, + host_capabilities: ProviderHostCapabilities, + ) -> Result> { + (self.constructor)( + extensions, + None, + self.tls_config.clone(), + true, + host_capabilities, + ) + .await } pub async fn create(&self, extensions: Vec) -> Result> { - (self.constructor)(extensions, None, self.tls_config.clone(), false).await + self.create_with_host_capabilities(extensions, ProviderHostCapabilities::default()) + .await + } + + pub async fn create_with_host_capabilities( + &self, + extensions: Vec, + host_capabilities: ProviderHostCapabilities, + ) -> Result> { + (self.constructor)( + extensions, + None, + self.tls_config.clone(), + false, + host_capabilities, + ) + .await } pub async fn create_with_working_dir( &self, extensions: Vec, working_dir: PathBuf, + ) -> Result> { + self.create_with_working_dir_and_host_capabilities( + extensions, + working_dir, + ProviderHostCapabilities::default(), + ) + .await + } + + pub async fn create_with_working_dir_and_host_capabilities( + &self, + extensions: Vec, + working_dir: PathBuf, + host_capabilities: ProviderHostCapabilities, ) -> Result> { (self.constructor)( extensions, Some(working_dir), self.tls_config.clone(), false, + host_capabilities, ) .await } @@ -140,19 +194,36 @@ impl ProviderRegistry { name, ProviderEntry { metadata, - constructor: Arc::new(|extensions, working_dir, tls_config, use_default_model| { - Box::pin(async move { - let provider = if use_default_model { - F::from_env_with_default_model(extensions, tls_config).await? - } else if let Some(working_dir) = working_dir { - F::from_env_with_working_dir(extensions, working_dir, tls_config) + constructor: Arc::new( + |extensions, working_dir, tls_config, use_default_model, host_capabilities| { + Box::pin(async move { + let provider = if use_default_model { + F::from_env_with_default_model_and_host_capabilities( + extensions, + tls_config, + host_capabilities, + ) + .await? + } else if let Some(working_dir) = working_dir { + F::from_env_with_working_dir_and_host_capabilities( + extensions, + working_dir, + tls_config, + host_capabilities, + ) + .await? + } else { + F::from_env_with_host_capabilities( + extensions, + tls_config, + host_capabilities, + ) .await? - } else { - F::from_env(extensions, tls_config).await? - }; - Ok(Arc::new(provider) as Arc) - }) - }), + }; + Ok(Arc::new(provider) as Arc) + }) + }, + ), inventory_identity: inventory.identity, inventory_configured: inventory.configured, cleanup: None, @@ -316,13 +387,15 @@ impl ProviderRegistry { config.name.clone(), ProviderEntry { metadata: custom_metadata, - constructor: Arc::new(move |_extensions, _working_dir, tls_config, _| { - let result = constructor(tls_config); - Box::pin(async move { - let provider = result?; - Ok(Arc::new(provider) as Arc) - }) - }), + constructor: Arc::new( + move |_extensions, _working_dir, tls_config, _, _host_capabilities| { + let result = constructor(tls_config); + Box::pin(async move { + let provider = result?; + Ok(Arc::new(provider) as Arc) + }) + }, + ), inventory_identity: Arc::new(inventory_identity), inventory_configured: inventory_configured.unwrap_or(default_inventory_configured), cleanup: None, From 1ec3c1c4931a915369b19010f7825a8c2dc614f8 Mon Sep 17 00:00:00 2001 From: luke Date: Thu, 20 Aug 2026 11:19:58 -0400 Subject: [PATCH 06/12] fix(acp): preserve relayed elicitation context Carry the nested tool call ID and ACP metadata through the persisted action-required record and back out to the outer client. This preserves forward-compatible relay context while retaining the Goose elicitation ID for response correlation. --- .../src/conversation/message.rs | 16 +++ crates/goose/src/acp/provider.rs | 44 +++++++-- crates/goose/src/acp/server.rs | 7 +- crates/goose/src/acp/server/elicitation.rs | 99 ++++++++++++++----- crates/goose/src/acp/server/load_session.rs | 32 +++++- 5 files changed, 162 insertions(+), 36 deletions(-) diff --git a/crates/goose-provider-types/src/conversation/message.rs b/crates/goose-provider-types/src/conversation/message.rs index 996ae6e63e42..14ff973fee14 100644 --- a/crates/goose-provider-types/src/conversation/message.rs +++ b/crates/goose-provider-types/src/conversation/message.rs @@ -199,6 +199,10 @@ pub enum ActionRequiredData { id: String, message: String, requested_schema: serde_json::Value, + #[serde(default, skip_serializing_if = "Option::is_none")] + tool_call_id: Option, + #[serde(default, rename = "_meta", skip_serializing_if = "Option::is_none")] + meta: Option, }, ElicitationResponse { id: String, @@ -513,12 +517,24 @@ impl MessageContentBlock { id: S, message: String, requested_schema: serde_json::Value, + ) -> Self { + Self::action_required_elicitation_with_context(id, message, requested_schema, None, None) + } + + pub fn action_required_elicitation_with_context>( + id: S, + message: String, + requested_schema: serde_json::Value, + tool_call_id: Option, + meta: Option, ) -> Self { MessageContentBlock::ActionRequired(ActionRequired { data: ActionRequiredData::Elicitation { id: id.into(), message, requested_schema, + tool_call_id, + meta, }, }) } diff --git a/crates/goose/src/acp/provider.rs b/crates/goose/src/acp/provider.rs index f94ee5b2c55f..3f126e960aee 100644 --- a/crates/goose/src/acp/provider.rs +++ b/crates/goose/src/acp/provider.rs @@ -2,9 +2,9 @@ use agent_client_protocol::schema::v1::{ Annotations as AcpAnnotations, ClientCapabilities, CloseSessionRequest, ContentBlock, ContentChunk, CreateElicitationRequest, CreateElicitationResponse, ElicitationAcceptAction, ElicitationAction as AcpElicitationAction, ElicitationCapabilities, ElicitationContentValue, - ElicitationFormCapabilities, ElicitationMode, EnvVariable, HttpHeader, ImageContent, - InitializeRequest, InitializeResponse, LoadSessionRequest, McpCapabilities, McpServer, - McpServerHttp, McpServerStdio, NewSessionRequest, NewSessionResponse, PromptRequest, + ElicitationFormCapabilities, ElicitationMode, ElicitationScope, EnvVariable, HttpHeader, + ImageContent, InitializeRequest, InitializeResponse, LoadSessionRequest, McpCapabilities, + McpServer, McpServerHttp, McpServerStdio, NewSessionRequest, NewSessionResponse, PromptRequest, PromptResponse, RequestPermissionOutcome, RequestPermissionRequest, RequestPermissionResponse, Role as AcpRole, SessionConfigKind, SessionConfigOption, SessionConfigOptionCategory, SessionConfigSelectOptions, SessionId, SessionModeState, SessionNotification, SessionUpdate, @@ -2215,16 +2215,25 @@ fn build_action_required_elicitation_message( return None; }; let requested_schema = serde_json::to_value(&form.requested_schema).ok()?; + let tool_call_id = match &form.scope { + ElicitationScope::Session(scope) => scope + .tool_call_id + .as_ref() + .map(|tool_call_id| tool_call_id.0.to_string()), + _ => None, + }; let request_id = format!( "{}{}", ACP_PROVIDER_ELICITATION_ID_PREFIX, uuid::Uuid::new_v4() ); let message = Message::assistant() - .with_content(MessageContent::action_required_elicitation( + .with_content(MessageContent::action_required_elicitation_with_context( request_id.clone(), request.message.clone(), requested_schema, + tool_call_id, + request.meta.clone(), )) .user_only(); @@ -2323,7 +2332,7 @@ mod tests { use agent_client_protocol::schema::v1::{ AgentCapabilities, ElicitationFormMode, ElicitationSchema, ElicitationSessionScope, ErrorCode, MultiSelectPropertySchema, SessionConfigSelectOption, SessionMode, - SessionModeId, + SessionModeId, ToolCallId, }; use test_case::test_case; @@ -2454,11 +2463,16 @@ mod tests { }; let request = CreateElicitationRequest::new( ElicitationFormMode::new( - ElicitationSessionScope::new("nested-session"), + ElicitationSessionScope::new("nested-session") + .tool_call_id(ToolCallId::new("nested-tool-call")), ElicitationSchema::new().string("answer", true), ), "Choose an answer", - ); + ) + .meta(serde_json::Map::from_iter([( + "nested".to_string(), + serde_json::json!({ "trace": "preserve-me" }), + )])); let (response_tx, response_rx) = oneshot::channel(); prompt_tx .send(AcpUpdate::ElicitationRequest { @@ -2475,11 +2489,23 @@ mod tests { let MessageContent::ActionRequired(action_required) = &message.content[0] else { panic!("expected action-required content"); }; - let ActionRequiredData::Elicitation { id, message, .. } = &action_required.data else { + let ActionRequiredData::Elicitation { + id, + message, + tool_call_id, + meta, + .. + } = &action_required.data + else { panic!("expected elicitation action-required content"); }; assert!(id.starts_with(ACP_PROVIDER_ELICITATION_ID_PREFIX)); assert_eq!(message, "Choose an answer"); + assert_eq!(tool_call_id.as_deref(), Some("nested-tool-call")); + assert_eq!( + meta.as_ref().and_then(|meta| meta.get("nested")), + Some(&serde_json::json!({ "trace": "preserve-me" })) + ); assert!( provider @@ -2741,6 +2767,7 @@ mod tests { id, message, requested_schema, + .. } = &action_required.data else { panic!("expected elicitation content"); @@ -2999,6 +3026,7 @@ mod tests { id, message, requested_schema, + .. } = &action_required.data else { panic!("expected elicitation content"); diff --git a/crates/goose/src/acp/server.rs b/crates/goose/src/acp/server.rs index f066832d3344..10fea52db9dd 100644 --- a/crates/goose/src/acp/server.rs +++ b/crates/goose/src/acp/server.rs @@ -77,7 +77,7 @@ use url::Url; use uuid::Uuid; use self::message_meta::{ - content_chunk_for_message, message_meta_without_steer, populate_output_token_limit_content, + content_chunk_for_message, merge_message_meta, populate_output_token_limit_content, }; use self::tool_calls::chain::{breaks_consecutive_tool_calls, ReadyToolChain, ToolChainTracker}; use self::tool_calls::conversion::{ @@ -1136,6 +1136,8 @@ impl GooseAcpAgent { id, message: elicitation_message, requested_schema, + tool_call_id, + meta, } => { self.handle_form_elicitation( cx, @@ -1145,7 +1147,8 @@ impl GooseAcpAgent { id.clone(), elicitation_message.clone(), requested_schema.clone(), - message_meta_without_steer(message), + tool_call_id.clone(), + merge_message_meta(meta.clone().unwrap_or_default(), message), false, ), ) diff --git a/crates/goose/src/acp/server/elicitation.rs b/crates/goose/src/acp/server/elicitation.rs index 002cc1ae5c6c..ebeea7d30f7a 100644 --- a/crates/goose/src/acp/server/elicitation.rs +++ b/crates/goose/src/acp/server/elicitation.rs @@ -2,7 +2,7 @@ use std::sync::Arc; use agent_client_protocol::schema::v1::{ CreateElicitationRequest, CreateElicitationResponse, ElicitationAction as AcpElicitationAction, - ElicitationFormMode, ElicitationSchema, ElicitationSessionScope, Meta, SessionId, + ElicitationFormMode, ElicitationSchema, ElicitationSessionScope, Meta, SessionId, ToolCallId, CLIENT_METHOD_NAMES, }; use agent_client_protocol::{ @@ -18,6 +18,7 @@ pub(super) struct FormElicitation { elicitation_id: String, message: String, requested_schema: serde_json::Value, + tool_call_id: Option, meta: Meta, recovered: bool, } @@ -28,6 +29,7 @@ impl FormElicitation { elicitation_id: String, message: String, requested_schema: serde_json::Value, + tool_call_id: Option, meta: Meta, recovered: bool, ) -> Self { @@ -36,10 +38,37 @@ impl FormElicitation { elicitation_id, message, requested_schema, + tool_call_id, meta, recovered, } } + + fn request( + &self, + session_id: &str, + requested_schema: ElicitationSchema, + has_live_waiter: bool, + ) -> CreateElicitationRequest { + let mut meta = self.meta.clone(); + add_elicitation_meta( + &mut meta, + &self.elicitation_id, + self.recovered, + if has_live_waiter { + "response" + } else { + "prompt" + }, + ); + let scope = ElicitationSessionScope::new(session_id.to_string()) + .tool_call_id(self.tool_call_id.as_deref().map(ToolCallId::new)); + CreateElicitationRequest::new( + ElicitationFormMode::new(scope, requested_schema), + self.message.clone(), + ) + .meta(meta) + } } impl super::GooseAcpAgent { @@ -76,7 +105,7 @@ impl super::GooseAcpAgent { elicitation: FormElicitation, ) -> Result<(), agent_client_protocol::Error> { let session_id = elicitation.session_id.0.as_ref().to_string(); - let elicitation_id = elicitation.elicitation_id; + let elicitation_id = elicitation.elicitation_id.clone(); if elicitation .requested_schema .get("url") @@ -100,7 +129,7 @@ impl super::GooseAcpAgent { } let requested_schema: ElicitationSchema = - match serde_json::from_value(elicitation.requested_schema) { + match serde_json::from_value(elicitation.requested_schema.clone()) { Ok(schema) => schema, Err(error) => { finish_form_elicitation( @@ -118,25 +147,7 @@ impl super::GooseAcpAgent { let has_live_waiter = agent .has_pending_elicitation(&session_id, &elicitation_id) .await; - let mut meta = elicitation.meta; - add_elicitation_meta( - &mut meta, - &elicitation_id, - elicitation.recovered, - if has_live_waiter { - "response" - } else { - "prompt" - }, - ); - let request = CreateElicitationRequest::new( - ElicitationFormMode::new( - ElicitationSessionScope::new(session_id.clone()), - requested_schema, - ), - elicitation.message, - ) - .meta(meta); + let request = elicitation.request(&session_id, requested_schema, has_live_waiter); let callback_agent = Arc::clone(agent); let callback_session_id = session_id.clone(); @@ -337,6 +348,50 @@ async fn finish_form_elicitation( #[cfg(test)] mod tests { use super::*; + use agent_client_protocol::schema::v1::{ElicitationMode, ElicitationScope}; + + #[test] + fn relayed_request_preserves_nested_tool_call_and_metadata() { + let elicitation = FormElicitation::new( + SessionId::new("outer-session"), + "acp-provider:question-1".to_string(), + "Choose a direction".to_string(), + serde_json::json!({}), + Some("nested-tool-call".to_string()), + Meta::from_iter([( + "nested".to_string(), + serde_json::json!({ "trace": "preserve-me" }), + )]), + false, + ); + + let request = elicitation.request( + "outer-session", + ElicitationSchema::new().string("direction", true), + true, + ); + let ElicitationMode::Form(form) = &request.mode else { + panic!("expected form elicitation"); + }; + let ElicitationScope::Session(scope) = &form.scope else { + panic!("expected session-scoped elicitation"); + }; + + assert_eq!(scope.session_id.0.as_ref(), "outer-session"); + assert_eq!( + scope + .tool_call_id + .as_ref() + .map(|tool_call_id| tool_call_id.0.as_ref()), + Some("nested-tool-call") + ); + let meta = request.meta.unwrap(); + assert_eq!( + meta.get("nested"), + Some(&serde_json::json!({ "trace": "preserve-me" })) + ); + assert_eq!(meta["goose"]["elicitationId"], "acp-provider:question-1"); + } #[test] fn recovered_elicitation_metadata_declares_prompt_continuation() { diff --git a/crates/goose/src/acp/server/load_session.rs b/crates/goose/src/acp/server/load_session.rs index db2a5cdaf815..bad8f501a58f 100644 --- a/crates/goose/src/acp/server/load_session.rs +++ b/crates/goose/src/acp/server/load_session.rs @@ -1,6 +1,5 @@ use super::message_meta::{ - content_chunk_for_message, merge_message_meta, message_meta_without_steer, - populate_output_token_limit_content, + content_chunk_for_message, merge_message_meta, populate_output_token_limit_content, }; use super::tool_calls::conversion::{ build_initial_tool_call_with_message_meta, tool_call_update_fields_from_response, @@ -53,6 +52,7 @@ struct PendingFormElicitation { id: String, message: String, requested_schema: serde_json::Value, + tool_call_id: Option, meta: Meta, } @@ -81,6 +81,8 @@ fn pending_form_elicitations(messages: &[Message]) -> Vec Vec Date: Thu, 20 Aug 2026 11:21:32 -0400 Subject: [PATCH 07/12] fix(cursor): keep session titles out of ACP transport Handle deterministic session-description requests before consulting the selected Cursor transport. This keeps title generation local even when the conversation transport is pinned to ACP or unavailable. --- crates/goose/src/providers/cursor_agent.rs | 32 +++++++++++++++++++--- 1 file changed, 28 insertions(+), 4 deletions(-) diff --git a/crates/goose/src/providers/cursor_agent.rs b/crates/goose/src/providers/cursor_agent.rs index ca68724f0807..630e12de20aa 100644 --- a/crates/goose/src/providers/cursor_agent.rs +++ b/crates/goose/src/providers/cursor_agent.rs @@ -858,10 +858,6 @@ impl Provider for CursorAgentProvider { messages: &[Message], tools: &[Tool], ) -> Result { - if let Some(provider) = self.selected_transport()? { - return provider.stream(model_config, system, messages, tools).await; - } - if super::cli_common::is_session_description_request(system) { let (message, provider_usage) = super::cli_common::generate_simple_session_description( &model_config.model_name, @@ -870,6 +866,10 @@ impl Provider for CursorAgentProvider { return Ok(stream_from_single_message(message, provider_usage)); } + if let Some(provider) = self.selected_transport()? { + return provider.stream(model_config, system, messages, tools).await; + } + let lines = self .execute_command(model_config, system, messages, tools) .await?; @@ -1060,6 +1060,30 @@ printf '%s\n' '{"type":"result","result":"ok"}' assert!(stream_error.to_string().contains("requires ACP")); } + #[tokio::test] + async fn session_description_does_not_enter_selected_transport() { + let directory = tempfile::tempdir().unwrap(); + let provider = CursorAgentProvider { + command: recording_cli(directory.path()), + name: CURSOR_AGENT_PROVIDER_NAME.to_string(), + transport: Mutex::new(CursorTransport::Unavailable( + "ACP transport should not be touched".to_string(), + )), + }; + + let (message, _) = provider + .complete( + &ModelConfig::new(CURSOR_AGENT_DEFAULT_MODEL), + "Generate a title in four words or less", + &[Message::user().with_text("Investigate nested Cursor ACP transport")], + &[], + ) + .await + .unwrap(); + + assert_eq!(message.as_concat_text(), "Investigate nested Cursor ACP"); + } + #[test] fn parse_models_output_extracts_ids_and_preserves_auto() { let stdout = r#" From bd69bd84c216973aad8236ff81cc1ec5781a1602 Mon Sep 17 00:00:00 2001 From: luke Date: Thu, 20 Aug 2026 11:30:01 -0400 Subject: [PATCH 08/12] fix(cursor): select transport at session activation Prepare the provider before either agent loop resolves context or other transport-sensitive behavior. This prevents Cursor first turns from being configured with direct-provider semantics and then executed on ACP. --- crates/goose/src/agents/agent.rs | 67 ++++--- .../goose/src/agents/state_machine/ops_llm.rs | 20 -- .../agents/state_machine/tests/agent_reply.rs | 183 +++++++++++++++++- 3 files changed, 225 insertions(+), 45 deletions(-) diff --git a/crates/goose/src/agents/agent.rs b/crates/goose/src/agents/agent.rs index cd70c6b3ef75..be50bc500c16 100644 --- a/crates/goose/src/agents/agent.rs +++ b/crates/goose/src/agents/agent.rs @@ -1725,7 +1725,7 @@ impl Agent { let cancel = cancel_token.unwrap_or_default(); let session_id = session_config.id.clone(); - let entry_session = session_manager.get_session(&session_id, false).await?; + let entry_session = session_manager.get_session(&session_id, true).await?; if let Some(schedule_id) = session_config.schedule_id.clone() { session_manager .update(&session_id) @@ -1743,6 +1743,28 @@ impl Agent { .await .clone() .ok_or_else(|| anyhow!("Provider not set"))?; + let provider_name = provider.get_name().to_string(); + let saved_provider_inference = + entry_session + .conversation + .as_ref() + .and_then(|conversation| { + super::latest_provider_inference(conversation.messages(), &provider_name) + }); + if let Err(error) = provider + .prepare_session( + saved_provider_inference + .and_then(|inference| inference.provider_session_id.as_deref()), + saved_provider_inference.is_some(), + ) + .await + { + warn!( + provider = provider_name, + %error, + "Could not prepare provider session; continuing with a handoff" + ); + } if !self.config.disable_session_naming { let manager = session_manager.clone(); @@ -2219,14 +2241,27 @@ impl Agent { .conversation .clone() .ok_or_else(|| anyhow::anyhow!("Session {} has no conversation", session_config.id))?; + let provider = self.provider().await?; + let provider_name = provider.get_name().to_string(); + let saved_provider_inference = + super::latest_provider_inference(conversation.messages(), &provider_name); + if let Err(error) = provider + .prepare_session( + saved_provider_inference + .and_then(|inference| inference.provider_session_id.as_deref()), + saved_provider_inference.is_some(), + ) + .await + { + warn!( + provider = provider_name, + %error, + "Could not prepare provider session; continuing with a handoff" + ); + } - let needs_auto_compact = check_if_compaction_needed( - self.provider().await?.as_ref(), - &conversation, - None, - &session, - ) - .await?; + let needs_auto_compact = + check_if_compaction_needed(provider.as_ref(), &conversation, None, &session).await?; let conversation_to_compact = conversation.clone(); let reply_span = tracing::Span::current(); @@ -2337,22 +2372,6 @@ impl Agent { let provider = self.provider().await?; let provider_name = provider.get_name().to_string(); - let saved_provider_inference = - super::latest_provider_inference(conversation.messages(), &provider_name); - if let Err(error) = provider - .prepare_session( - saved_provider_inference - .and_then(|inference| inference.provider_session_id.as_deref()), - saved_provider_inference.is_some(), - ) - .await - { - warn!( - provider = provider_name, - %error, - "Could not prepare provider session; continuing with a handoff" - ); - } let requested_model = model_config.model_name.clone(); let resolved_model = provider diff --git a/crates/goose/src/agents/state_machine/ops_llm.rs b/crates/goose/src/agents/state_machine/ops_llm.rs index 0194b078a55a..0ac0e48784d8 100644 --- a/crates/goose/src/agents/state_machine/ops_llm.rs +++ b/crates/goose/src/agents/state_machine/ops_llm.rs @@ -438,26 +438,6 @@ impl Inference for InferenceRunner<'_> { .get_context_limit(&self.model_config) .await .unwrap_or_else(|_| self.model_config.context_limit()); - let provider_name = self.provider.get_name(); - let saved_provider_inference = super::super::latest_provider_inference( - conversation.messages(), - provider_name, - ); - if let Err(error) = self - .provider - .prepare_session( - saved_provider_inference - .and_then(|inference| inference.provider_session_id.as_deref()), - saved_provider_inference.is_some(), - ) - .await - { - tracing::warn!( - provider = provider_name, - %error, - "Could not prepare provider session; continuing with a handoff" - ); - } let turn = messages_since_kickoff(conversation)?; let turn_start = turn .first() diff --git a/crates/goose/src/agents/state_machine/tests/agent_reply.rs b/crates/goose/src/agents/state_machine/tests/agent_reply.rs index bc6aa43b5c05..718ccc4e7264 100644 --- a/crates/goose/src/agents/state_machine/tests/agent_reply.rs +++ b/crates/goose/src/agents/state_machine/tests/agent_reply.rs @@ -1,6 +1,7 @@ //! Covers `Agent::reply_with_state_machine`, the entry point the CLI and desktop //! reach when the state machine is enabled. +use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; use std::sync::Arc; use std::time::Duration; @@ -18,7 +19,7 @@ use crate::agents::{Agent, AgentConfig, AgentEvent, GoosePlatform, SessionConfig use crate::config::permission::PermissionManager; use crate::config::GooseMode; use crate::conversation::message::{ActionRequiredData, Message, MessageContent}; -use crate::providers::base::{MessageStream, Provider}; +use crate::providers::base::{MessageStream, PermissionRouting, Provider}; use crate::session::{SessionManager, SessionType}; use goose_providers::errors::ProviderError; use goose_providers::model::ModelConfig; @@ -127,6 +128,94 @@ impl Provider for NestedElicitationProvider { } } +struct ActivationSensitiveProvider { + armed: AtomicBool, + prepared: AtomicBool, + prepare_calls: AtomicUsize, + queried_before_prepare: AtomicBool, + context_queried_after_prepare: AtomicBool, +} + +impl ActivationSensitiveProvider { + fn new() -> Self { + Self { + armed: AtomicBool::new(false), + prepared: AtomicBool::new(false), + prepare_calls: AtomicUsize::new(0), + queried_before_prepare: AtomicBool::new(false), + context_queried_after_prepare: AtomicBool::new(false), + } + } + + fn arm(&self) { + self.armed.store(true, Ordering::SeqCst); + } + + fn record_capability_query(&self) -> bool { + let prepared = self.prepared.load(Ordering::SeqCst); + if self.armed.load(Ordering::SeqCst) && !prepared { + self.queried_before_prepare.store(true, Ordering::SeqCst); + } + prepared + } +} + +#[async_trait] +impl Provider for ActivationSensitiveProvider { + fn get_name(&self) -> &str { + "cursor-activation-test" + } + + async fn prepare_session( + &self, + _provider_session_id: Option<&str>, + _has_provider_history: bool, + ) -> Result<(), ProviderError> { + self.prepare_calls.fetch_add(1, Ordering::SeqCst); + self.prepared.store(true, Ordering::SeqCst); + Ok(()) + } + + async fn stream( + &self, + _model_config: &ModelConfig, + _system: &str, + _messages: &[Message], + _tools: &[Tool], + ) -> Result { + if !self.prepared.load(Ordering::SeqCst) { + return Err(ProviderError::RequestFailed( + "stream entered before session activation".to_string(), + )); + } + Ok(Box::pin(async_stream::try_stream! { + yield (Some(Message::assistant().with_text("activated")), None); + })) + } + + async fn get_context_limit(&self, _model_config: &ModelConfig) -> Result { + if self.record_capability_query() { + self.context_queried_after_prepare + .store(true, Ordering::SeqCst); + Ok(262_144) + } else { + Ok(8_192) + } + } + + fn manages_own_context(&self) -> bool { + self.record_capability_query() + } + + fn permission_routing(&self) -> PermissionRouting { + if self.record_capability_query() { + PermissionRouting::ActionRequired + } else { + PermissionRouting::Noop + } + } +} + async fn agent_with_dummy_api() -> Result<(Agent, Arc, String, tempfile::TempDir)> { let api = Arc::new(DummyApi::start(ProviderFeatures::default()).await); let api_client = goose_providers::api_client::ApiClient::new_with_tls( @@ -197,6 +286,43 @@ async fn agent_with_nested_elicitation_provider() -> Result<(Agent, String, temp Ok((agent, session.id, temp_dir)) } +async fn agent_with_activation_sensitive_provider() -> Result<( + Agent, + Arc, + String, + tempfile::TempDir, +)> { + let provider = Arc::new(ActivationSensitiveProvider::new()); + let temp_dir = tempfile::tempdir()?; + let session_manager = Arc::new(SessionManager::new(temp_dir.path().to_path_buf())); + let session = session_manager + .create_session( + temp_dir.path().to_path_buf(), + "cursor-session-activation".to_string(), + SessionType::Hidden, + GooseMode::Auto, + ) + .await?; + let agent = Agent::with_config(AgentConfig::new( + session_manager, + PermissionManager::instance(), + None, + GooseMode::Auto, + true, + GoosePlatform::GooseCli, + )); + agent + .update_provider( + provider.clone(), + ModelConfig::new("cursor-activation-test"), + &session.id, + ) + .await?; + provider.arm(); + + Ok((agent, provider, session.id, temp_dir)) +} + #[tokio::test] async fn reply_streams_the_turn_and_ends() -> Result<()> { let (agent, api, session_id, _temp_dir) = agent_with_dummy_api().await?; @@ -393,6 +519,61 @@ async fn nested_provider_elicitation_persists_before_response_in_legacy_loop() - assert_nested_provider_elicitation_order(false).await } +async fn assert_cursor_transport_precedes_capability_queries( + use_state_machine: bool, +) -> Result<()> { + let (agent, provider, session_id, _temp_dir) = + agent_with_activation_sensitive_provider().await?; + let session_config = SessionConfig { + id: session_id, + schedule_id: None, + max_turns: Some(2), + retry_config: None, + }; + let user_message = Message::user().with_text("use the selected transport"); + let cancel = Some(CancellationToken::new()); + let stream = if use_state_machine { + agent + .reply_with_state_machine(user_message, session_config, cancel) + .await? + } else { + agent.reply(user_message, session_config, cancel).await? + }; + let saw_activated = tokio::time::timeout(Duration::from_secs(30), async move { + tokio::pin!(stream); + while let Some(event) = stream.next().await { + if let AgentEvent::Message(message) = event? { + if message.as_concat_text() == "activated" { + return anyhow::Ok(true); + } + } + } + anyhow::Ok(false) + }) + .await??; + + assert!(saw_activated); + assert!(provider.prepared.load(Ordering::SeqCst)); + assert_eq!(provider.prepare_calls.load(Ordering::SeqCst), 1); + assert!(provider + .context_queried_after_prepare + .load(Ordering::SeqCst)); + assert!(!provider.queried_before_prepare.load(Ordering::SeqCst)); + + Ok(()) +} + +#[tokio::test] +async fn cursor_transport_precedes_capability_queries_in_state_machine() -> Result<()> { + assert_cursor_transport_precedes_capability_queries(true).await +} + +#[tokio::test] +async fn cursor_transport_precedes_capability_queries_in_legacy_loop() -> Result<()> { + let _guard = env_lock::lock_env([("GOOSE_STATE_MACHINE", None::<&str>)]); + assert_cursor_transport_precedes_capability_queries(false).await +} + #[tokio::test] async fn bang_shell_uses_the_state_machine_when_the_flag_is_disabled() -> Result<()> { let _guard = env_lock::lock_env([("GOOSE_STATE_MACHINE", None::<&str>)]); From 95de653c21a0a46caca53ade00b85b111be25617 Mon Sep 17 00:00:00 2001 From: luke Date: Thu, 20 Aug 2026 11:34:08 -0400 Subject: [PATCH 09/12] fix(cursor): fail closed when ACP resume fails Propagate provider session preparation failures from both agent loops instead of logging a handoff that never occurs. A failed saved-session resume now stops before any provider stream can run. --- crates/goose/src/agents/agent.rs | 22 +-- .../agents/state_machine/tests/agent_reply.rs | 145 +++++++++++++++++- 2 files changed, 148 insertions(+), 19 deletions(-) diff --git a/crates/goose/src/agents/agent.rs b/crates/goose/src/agents/agent.rs index be50bc500c16..eebd659bab5e 100644 --- a/crates/goose/src/agents/agent.rs +++ b/crates/goose/src/agents/agent.rs @@ -1751,20 +1751,13 @@ impl Agent { .and_then(|conversation| { super::latest_provider_inference(conversation.messages(), &provider_name) }); - if let Err(error) = provider + provider .prepare_session( saved_provider_inference .and_then(|inference| inference.provider_session_id.as_deref()), saved_provider_inference.is_some(), ) - .await - { - warn!( - provider = provider_name, - %error, - "Could not prepare provider session; continuing with a handoff" - ); - } + .await?; if !self.config.disable_session_naming { let manager = session_manager.clone(); @@ -2245,20 +2238,13 @@ impl Agent { let provider_name = provider.get_name().to_string(); let saved_provider_inference = super::latest_provider_inference(conversation.messages(), &provider_name); - if let Err(error) = provider + provider .prepare_session( saved_provider_inference .and_then(|inference| inference.provider_session_id.as_deref()), saved_provider_inference.is_some(), ) - .await - { - warn!( - provider = provider_name, - %error, - "Could not prepare provider session; continuing with a handoff" - ); - } + .await?; let needs_auto_compact = check_if_compaction_needed(provider.as_ref(), &conversation, None, &session).await?; diff --git a/crates/goose/src/agents/state_machine/tests/agent_reply.rs b/crates/goose/src/agents/state_machine/tests/agent_reply.rs index 718ccc4e7264..a345ce4d14ce 100644 --- a/crates/goose/src/agents/state_machine/tests/agent_reply.rs +++ b/crates/goose/src/agents/state_machine/tests/agent_reply.rs @@ -18,7 +18,9 @@ use crate::action_required_manager::ElicitationOutcome; use crate::agents::{Agent, AgentConfig, AgentEvent, GoosePlatform, SessionConfig}; use crate::config::permission::PermissionManager; use crate::config::GooseMode; -use crate::conversation::message::{ActionRequiredData, Message, MessageContent}; +use crate::conversation::message::{ + ActionRequiredData, InferenceMetadata, Message, MessageContent, +}; use crate::providers::base::{MessageStream, PermissionRouting, Provider}; use crate::session::{SessionManager, SessionType}; use goose_providers::errors::ProviderError; @@ -216,6 +218,57 @@ impl Provider for ActivationSensitiveProvider { } } +struct PrepareFailureProvider { + prepare_calls: AtomicUsize, + stream_calls: AtomicUsize, + saw_saved_session: AtomicBool, +} + +impl PrepareFailureProvider { + fn new() -> Self { + Self { + prepare_calls: AtomicUsize::new(0), + stream_calls: AtomicUsize::new(0), + saw_saved_session: AtomicBool::new(false), + } + } +} + +#[async_trait] +impl Provider for PrepareFailureProvider { + fn get_name(&self) -> &str { + "cursor-resume-failure-test" + } + + async fn prepare_session( + &self, + provider_session_id: Option<&str>, + has_provider_history: bool, + ) -> Result<(), ProviderError> { + self.prepare_calls.fetch_add(1, Ordering::SeqCst); + self.saw_saved_session.store( + provider_session_id == Some("saved-cursor-session") && has_provider_history, + Ordering::SeqCst, + ); + Err(ProviderError::RequestFailed( + "Cursor ACP resume failed".to_string(), + )) + } + + async fn stream( + &self, + _model_config: &ModelConfig, + _system: &str, + _messages: &[Message], + _tools: &[Tool], + ) -> Result { + self.stream_calls.fetch_add(1, Ordering::SeqCst); + Err(ProviderError::RequestFailed( + "stream must not follow a failed resume".to_string(), + )) + } +} + async fn agent_with_dummy_api() -> Result<(Agent, Arc, String, tempfile::TempDir)> { let api = Arc::new(DummyApi::start(ProviderFeatures::default()).await); let api_client = goose_providers::api_client::ApiClient::new_with_tls( @@ -323,6 +376,55 @@ async fn agent_with_activation_sensitive_provider() -> Result<( Ok((agent, provider, session.id, temp_dir)) } +async fn agent_with_prepare_failure_provider() -> Result<( + Agent, + Arc, + String, + tempfile::TempDir, +)> { + let provider = Arc::new(PrepareFailureProvider::new()); + let temp_dir = tempfile::tempdir()?; + let session_manager = Arc::new(SessionManager::new(temp_dir.path().to_path_buf())); + let session = session_manager + .create_session( + temp_dir.path().to_path_buf(), + "cursor-resume-failure".to_string(), + SessionType::Hidden, + GooseMode::Auto, + ) + .await?; + session_manager + .add_message( + &session.id, + &Message::assistant() + .with_text("prior Cursor response") + .with_inference(InferenceMetadata { + provider: "cursor-resume-failure-test".to_string(), + requested_model: "cursor-resume-failure-test".to_string(), + resolved_model: None, + provider_session_id: Some("saved-cursor-session".to_string()), + }), + ) + .await?; + let agent = Agent::with_config(AgentConfig::new( + session_manager, + PermissionManager::instance(), + None, + GooseMode::Auto, + true, + GoosePlatform::GooseCli, + )); + agent + .update_provider( + provider.clone(), + ModelConfig::new("cursor-resume-failure-test"), + &session.id, + ) + .await?; + + Ok((agent, provider, session.id, temp_dir)) +} + #[tokio::test] async fn reply_streams_the_turn_and_ends() -> Result<()> { let (agent, api, session_id, _temp_dir) = agent_with_dummy_api().await?; @@ -574,6 +676,47 @@ async fn cursor_transport_precedes_capability_queries_in_legacy_loop() -> Result assert_cursor_transport_precedes_capability_queries(false).await } +async fn assert_cursor_resume_failure_is_terminal(use_state_machine: bool) -> Result<()> { + let (agent, provider, session_id, _temp_dir) = agent_with_prepare_failure_provider().await?; + let session_config = SessionConfig { + id: session_id, + schedule_id: None, + max_turns: Some(2), + retry_config: None, + }; + let user_message = Message::user().with_text("resume the Cursor session"); + let cancel = Some(CancellationToken::new()); + let result = if use_state_machine { + agent + .reply_with_state_machine(user_message, session_config, cancel) + .await + } else { + agent.reply(user_message, session_config, cancel).await + }; + let error = match result { + Ok(_) => panic!("failed session preparation must stop the reply"), + Err(error) => error, + }; + + assert!(error.to_string().contains("Cursor ACP resume failed")); + assert_eq!(provider.prepare_calls.load(Ordering::SeqCst), 1); + assert_eq!(provider.stream_calls.load(Ordering::SeqCst), 0); + assert!(provider.saw_saved_session.load(Ordering::SeqCst)); + + Ok(()) +} + +#[tokio::test] +async fn cursor_resume_failure_is_terminal_in_state_machine() -> Result<()> { + assert_cursor_resume_failure_is_terminal(true).await +} + +#[tokio::test] +async fn cursor_resume_failure_is_terminal_in_legacy_loop() -> Result<()> { + let _guard = env_lock::lock_env([("GOOSE_STATE_MACHINE", None::<&str>)]); + assert_cursor_resume_failure_is_terminal(false).await +} + #[tokio::test] async fn bang_shell_uses_the_state_machine_when_the_flag_is_disabled() -> Result<()> { let _guard = env_lock::lock_env([("GOOSE_STATE_MACHINE", None::<&str>)]); From 438ebb5fa1dbcf026581a43ef650afcabf41fea3 Mon Sep 17 00:00:00 2001 From: luke Date: Thu, 20 Aug 2026 11:37:29 -0400 Subject: [PATCH 10/12] refactor(acp): remove obsolete message meta helper The elicitation relay now merges incoming metadata directly, leaving the no-steer wrapper without production callers. Remove the dead wrapper while retaining its serialization assertion against the underlying helper. --- crates/goose/src/acp/server/message_meta.rs | 6 +----- 1 file changed, 1 insertion(+), 5 deletions(-) diff --git a/crates/goose/src/acp/server/message_meta.rs b/crates/goose/src/acp/server/message_meta.rs index 85765a8c4126..3b10d20a9066 100644 --- a/crates/goose/src/acp/server/message_meta.rs +++ b/crates/goose/src/acp/server/message_meta.rs @@ -60,10 +60,6 @@ fn message_meta(message: &Message) -> Meta { message_meta_with_steer(message, message.metadata.steer) } -pub(super) fn message_meta_without_steer(message: &Message) -> Meta { - message_meta_with_steer(message, false) -} - pub(super) fn merge_message_meta(mut meta: Meta, message: &Message) -> Meta { extend_message_meta(&mut meta, message, message.metadata.steer); meta @@ -126,7 +122,7 @@ mod tests { })), ); assert_eq!( - message_meta_without_steer(&steer_message).get("goose"), + message_meta_with_steer(&steer_message, false).get("goose"), Some(&serde_json::json!({ "created": 1_700_000_000, "messageId": "msg_live", From a7a3babe8bc6445d16879b23c1f6ffbee69adcbc Mon Sep 17 00:00:00 2001 From: luke Date: Thu, 20 Aug 2026 13:09:29 -0400 Subject: [PATCH 11/12] fix(cli): match the widened elicitation action shape Preserving the nested request's `toolCallId` and `_meta` through the relay added two fields to `ActionRequiredData::Elicitation`, which breaks every exhaustive pattern over it. The CLI's extractor was one, so the workspace stopped building even though `-p goose --lib` stayed green. The CLI renders the prompt itself and has no use for relay correlation fields, so it ignores them explicitly rather than threading them through. --- crates/goose-cli/src/session/mod.rs | 3 +++ 1 file changed, 3 insertions(+) diff --git a/crates/goose-cli/src/session/mod.rs b/crates/goose-cli/src/session/mod.rs index b4fec8ca74ad..cc1146517e8b 100644 --- a/crates/goose-cli/src/session/mod.rs +++ b/crates/goose-cli/src/session/mod.rs @@ -2382,6 +2382,9 @@ fn find_elicitation_request(message: &Message) -> Option<(String, String, Value) id, message, requested_schema, + // The CLI renders the prompt itself; relay correlation fields + // are only meaningful to the ACP forwarding path. + .. } = &action.data { return Some((id.clone(), message.clone(), requested_schema.clone())); From 26da1da0dc17897b65ad75048c64ab82dec8ab9a Mon Sep 17 00:00:00 2001 From: luke Date: Thu, 20 Aug 2026 13:21:54 -0400 Subject: [PATCH 12/12] fix(tests): update ACP fixtures for the widened provider factory Honouring the outer client's form capability gave the provider factory a `ProviderHostCapabilities` parameter and `AcpProviderConfig` a `supports_form_elicitation` field. Both are the right shape, and both ripple into every fixture that constructs them. Seven call sites in `crates/goose/tests/` still used the old arity, so the workspace did not build even though `-p goose --lib` stayed green. The integration fixtures now match, and the host stub advertises form support because that is what it stands in for. --- crates/goose/tests/acp_common_tests/mod.rs | 2 +- crates/goose/tests/acp_custom_requests_test.rs | 6 +++--- crates/goose/tests/acp_fixtures/mod.rs | 6 +++++- crates/goose/tests/acp_fixtures/provider.rs | 2 ++ crates/goose/tests/acp_secret_cache_invalidation_test.rs | 2 +- 5 files changed, 12 insertions(+), 6 deletions(-) diff --git a/crates/goose/tests/acp_common_tests/mod.rs b/crates/goose/tests/acp_common_tests/mod.rs index aa97f1df3db8..b5689ec00323 100644 --- a/crates/goose/tests/acp_common_tests/mod.rs +++ b/crates/goose/tests/acp_common_tests/mod.rs @@ -447,7 +447,7 @@ pub async fn run_fs_write_text_file_true() { pub async fn run_initialize_doesnt_hit_provider() { let provider_factory: AcpProviderFactory = - Arc::new(|_, _, _, _| Box::pin(async { Err(anyhow::anyhow!("no provider configured")) })); + Arc::new(|_, _, _, _, _| Box::pin(async { Err(anyhow::anyhow!("no provider configured")) })); let openai = OpenAiFixture::new(vec![], C::expected_session_id()).await; let config = TestConnectionConfig { diff --git a/crates/goose/tests/acp_custom_requests_test.rs b/crates/goose/tests/acp_custom_requests_test.rs index 446f088b7041..1f43ec7fdf12 100644 --- a/crates/goose/tests/acp_custom_requests_test.rs +++ b/crates/goose/tests/acp_custom_requests_test.rs @@ -1175,7 +1175,7 @@ fn test_custom_provider_supported_models_lists_raw_provider_models() { run_test(async move { let openai = OpenAiFixture::new(vec![], Arc::new(EnforceSessionId::default())).await; let provider_factory: AcpProviderFactory = Arc::new( - |provider_name, _extensions, _working_dir, use_default_model| { + |provider_name, _extensions, _working_dir, use_default_model, _host_capabilities| { assert!(use_default_model); Box::pin(async move { Ok(Arc::new(MockProvider { @@ -1226,7 +1226,7 @@ fn test_custom_provider_supported_models_maps_not_configured_error() { write_acp_global_config(DEFAULT_ACP_TEST_CONFIG); run_test(async move { let openai = OpenAiFixture::new(vec![], Arc::new(EnforceSessionId::default())).await; - let provider_factory: AcpProviderFactory = Arc::new(|provider_name, _, _, _| { + let provider_factory: AcpProviderFactory = Arc::new(|provider_name, _, _, _, _| { Box::pin(async move { Ok(Arc::new(MockProvider { name: provider_name, @@ -1263,7 +1263,7 @@ fn test_custom_provider_supported_models_maps_authentication_error() { write_acp_global_config(DEFAULT_ACP_TEST_CONFIG); run_test(async move { let openai = OpenAiFixture::new(vec![], Arc::new(EnforceSessionId::default())).await; - let provider_factory: AcpProviderFactory = Arc::new(|provider_name, _, _, _| { + let provider_factory: AcpProviderFactory = Arc::new(|provider_name, _, _, _, _| { Box::pin(async move { Ok(Arc::new(MockProvider { name: provider_name, diff --git a/crates/goose/tests/acp_fixtures/mod.rs b/crates/goose/tests/acp_fixtures/mod.rs index 2f42b71479d2..ede8b4d89949 100644 --- a/crates/goose/tests/acp_fixtures/mod.rs +++ b/crates/goose/tests/acp_fixtures/mod.rs @@ -373,7 +373,11 @@ pub async fn spawn_acp_server_in_process( let provider_factory = provider_factory.unwrap_or_else(|| { let base_url = openai_base_url.to_string(); Arc::new( - move |_provider_name, _extensions, _working_dir, _use_default_model| { + move |_provider_name, + _extensions, + _working_dir, + _use_default_model, + _host_capabilities| { let base_url = base_url.clone(); Box::pin(async move { let api_client = ApiClient::new_with_tls( diff --git a/crates/goose/tests/acp_fixtures/provider.rs b/crates/goose/tests/acp_fixtures/provider.rs index ef28fcbd1667..2e09fd9610c2 100644 --- a/crates/goose/tests/acp_fixtures/provider.rs +++ b/crates/goose/tests/acp_fixtures/provider.rs @@ -176,6 +176,8 @@ impl Connection for AcpProviderConnection { let session_models: SessionModels = Arc::new(std::sync::Mutex::new(HashMap::new())); let sink_clone = notification_sink.clone(); let provider_config = AcpProviderConfig { + // The fixture stands in for a host that can render forms. + supports_form_elicitation: true, command: "unused".into(), args: vec![], env: vec![], diff --git a/crates/goose/tests/acp_secret_cache_invalidation_test.rs b/crates/goose/tests/acp_secret_cache_invalidation_test.rs index 10e3f13bbb45..bcec3052ccd2 100644 --- a/crates/goose/tests/acp_secret_cache_invalidation_test.rs +++ b/crates/goose/tests/acp_secret_cache_invalidation_test.rs @@ -46,7 +46,7 @@ impl Provider for MockProvider { fn mock_provider_factory() -> goose::acp::server::AcpProviderFactory { Arc::new( - |provider_name, _extensions, _working_dir, _use_default_model| { + |provider_name, _extensions, _working_dir, _use_default_model, _host_capabilities| { Box::pin(async move { Ok(Arc::new(MockProvider { name: provider_name,