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())); diff --git a/crates/goose-provider-types/src/base.rs b/crates/goose-provider-types/src/base.rs index b44254349441..32f7a6037c34 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,35 @@ pub trait Provider: Send + Sync { ) -> bool { 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, + _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-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/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..3f126e960aee 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, 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, + 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, @@ -63,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 { @@ -148,6 +179,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 +274,10 @@ 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. @@ -268,6 +307,30 @@ 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>>>, + 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); + } +} + fn spawn_client_loop(fut: impl Future + Send + 'static) -> JoinHandle<()> { std::thread::spawn(move || { let rt = tokio::runtime::Builder::new_current_thread() @@ -283,11 +346,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 +380,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 +419,7 @@ impl AcpProvider { name: String, goose_mode: GooseMode, config: AcpProviderConfig, + client_extensions: Vec, run: ClientLoopFn, ) -> Result { let (tx, rx) = mpsc::channel(32); @@ -352,6 +439,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 +482,8 @@ impl AcpProvider { mode_mapping, 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, @@ -581,6 +671,9 @@ impl Provider for AcpProvider { return Ok(()); } + cancel_pending_elicitations(&self.pending_elicitations); + cancel_pending_elicitations(&self.claimed_elicitations); + let previous_session_id = self.acp_session_id(); let loaded = self .load_session(SessionId::new(session_id)) @@ -659,6 +752,105 @@ 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, + }; + + // 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 + .claimed_elicitations + .lock() + .ok() + .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; + }; + + 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)) + || 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( &self, model_config: &ModelConfig, @@ -737,6 +929,8 @@ 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 claimed_elicitations = self.claimed_elicitations.clone(); let goose_mode = *self .goose_mode .lock() @@ -746,6 +940,10 @@ 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(), + claimed: claimed_elicitations.clone(), + }; let mut suppress_text = false; let mut bare_retry = bare_retry; let mut updates_seen = 0usize; @@ -877,6 +1075,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 +1180,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 +1258,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 +1316,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 +1442,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 +1473,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 +1656,17 @@ 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 = 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), @@ -1478,7 +1801,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 +1826,7 @@ async fn handle_requests( } } - *prompt_response_tx.lock().unwrap() = None; + *active_prompt.lock().unwrap() = None; } } } @@ -1882,6 +2208,38 @@ 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 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_context( + request_id.clone(), + request.message.clone(), + requested_schema, + tool_call_id, + request.meta.clone(), + )) + .user_only(); + + Some((request_id, message)) +} + fn extract_model_info_from_config_options( config_options: &[SessionConfigOption], ) -> Option<(String, Vec)> { @@ -1972,7 +2330,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, ToolCallId, }; use test_case::test_case; @@ -2088,6 +2448,651 @@ 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") + .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 { + 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, + 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 + .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 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(); + 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(), + claimed: Arc::new(Mutex::new(HashMap::new())), + }); + + 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 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::{ + 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 +3124,8 @@ mod tests { response: NewSessionResponse::new("test-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: Arc::new(Mutex::new(HashMap::new())), handoff_context_sent: Arc::new(AtomicBool::new(false)), context_size: Arc::new(AtomicU64::new(0)), @@ -2991,6 +3998,7 @@ mod tests { model_config_option_id: None, mode_mapping, notification_callback: None, + supports_form_elicitation: true, } } @@ -3002,6 +4010,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..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::{ @@ -96,6 +96,7 @@ mod diagnostics; mod dictation; mod dispatch; mod elicitation; +use self::elicitation::FormElicitation; mod extensions; mod fork_session; mod list_sessions; @@ -122,6 +123,7 @@ pub type AcpProviderFactory = Arc< Vec, Option, bool, + crate::providers::base::ProviderHostCapabilities, ) -> BoxFuture<'static, Result>> + Send + Sync, @@ -661,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)); @@ -725,6 +733,7 @@ impl GooseAcpAgent { extensions, working_dir, use_default_model, + self.provider_host_capabilities(), ) .await } @@ -1127,14 +1136,21 @@ impl GooseAcpAgent { id, message: elicitation_message, requested_schema, + tool_call_id, + meta, } => { 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(), + tool_call_id.clone(), + merge_message_meta(meta.clone().unwrap_or_default(), message), + false, + ), ) .await?; } @@ -2819,6 +2835,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/elicitation.rs b/crates/goose/src/acp/server/elicitation.rs index 3593d88f973e..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::{ @@ -11,36 +11,88 @@ 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, + tool_call_id: Option, + meta: Meta, + recovered: bool, +} + +impl FormElicitation { + pub(super) fn new( + session_id: SessionId, + elicitation_id: String, + message: String, + requested_schema: serde_json::Value, + tool_call_id: Option, + meta: Meta, + recovered: bool, + ) -> Self { + Self { + session_id, + 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 { 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 +101,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.clone(); + if elicitation + .requested_schema .get("url") .and_then(|url| url.as_str()) .is_some() @@ -67,94 +117,117 @@ 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.clone()) { 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 request = CreateElicitationRequest::new( - ElicitationFormMode::new( - ElicitationSessionScope::new(session_id.clone()), - requested_schema, - ), - message.to_string(), - ) - .meta(meta); + let has_live_waiter = agent + .has_pending_elicitation(&session_id, &elicitation_id) + .await; + let request = elicitation.request(&session_id, requested_schema, has_live_waiter); - 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 +302,123 @@ 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::*; + 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() { + 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..bad8f501a58f 100644 --- a/crates/goose/src/acp/server/load_session.rs +++ b/crates/goose/src/acp/server/load_session.rs @@ -47,6 +47,58 @@ 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, + tool_call_id: Option, + 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, + tool_call_id, + meta, + } = &action.data + else { + return None; + }; + (!answered.contains(id.as_str())).then(|| PendingFormElicitation { + id: id.clone(), + message: elicitation_message.clone(), + requested_schema: requested_schema.clone(), + tool_call_id: tool_call_id.clone(), + meta: merge_message_meta(meta.clone().unwrap_or_default(), message), + }) + }) + }) + .collect() +} + fn send_replay_content_chunk( cx: &ConnectionTo, session_id: &SessionId, @@ -266,6 +318,39 @@ 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.tool_call_id, + pending.meta, + true, + ), + ) + .await?; + } + + Ok(()) + } + pub(super) async fn handle_load_session( &self, cx: &ConnectionTo, @@ -300,6 +385,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 +570,59 @@ 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 pending_request = Message::assistant().with_content( + MessageContent::action_required_elicitation_with_context( + "pending".to_string(), + "Pending question".to_string(), + serde_json::json!({ + "type": "object", + "properties": { "answer": { "type": "string" } } + }), + Some("nested-tool-call".to_string()), + Some(Meta::from_iter([( + "nested".to_string(), + serde_json::json!({ "trace": "preserve-me" }), + )])), + ), + ); + 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"), + pending_request, + ]); + + 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"); + assert_eq!(pending[0].tool_call_id.as_deref(), Some("nested-tool-call")); + assert_eq!( + pending[0].meta.get("nested"), + Some(&serde_json::json!({ "trace": "preserve-me" })) + ); + } } 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", 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/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..eebd659bab5e 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, @@ -1723,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) @@ -1741,6 +1743,21 @@ 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) + }); + provider + .prepare_session( + saved_provider_inference + .and_then(|inference| inference.provider_session_id.as_deref()), + saved_provider_inference.is_some(), + ) + .await?; if !self.config.disable_session_naming { let manager = session_manager.clone(); @@ -1801,7 +1818,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 +1843,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 +1884,123 @@ 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 { + // 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; + 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); + } + + 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 provider_owned { + let provider = self + .provider + .lock() + .await + .clone() + .ok_or_else(|| anyhow!("Provider is not configured"))?; + // 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 + { + provider.release_elicitation(elicitation_id).await; + return Err(error); + } + 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 +2040,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())); } } @@ -2076,14 +2234,20 @@ 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); + provider + .prepare_session( + saved_provider_inference + .and_then(|inference| inference.provider_session_id.as_deref()), + saved_provider_inference.is_some(), + ) + .await?; - 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(); @@ -2194,17 +2358,6 @@ 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 requested_model = model_config.model_name.clone(); let resolved_model = provider @@ -2551,6 +2704,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 +2740,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 +4194,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..0ac0e48784d8 100644 --- a/crates/goose/src/agents/state_machine/ops_llm.rs +++ b/crates/goose/src/agents/state_machine/ops_llm.rs @@ -438,19 +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(); - if let Some(session_id) = super::super::latest_provider_session_id( - 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" - ); - } - } let turn = messages_since_kickoff(conversation)?; let turn_start = turn .first() @@ -514,6 +501,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 +540,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..a345ce4d14ce 100644 --- a/crates/goose/src/agents/state_machine/tests/agent_reply.rs +++ b/crates/goose/src/agents/state_machine/tests/agent_reply.rs @@ -1,22 +1,274 @@ //! 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; 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, InferenceMetadata, Message, MessageContent, +}; +use crate::providers::base::{MessageStream, PermissionRouting, 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>>>, + /// 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)), + } + } +} + +#[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() || 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( + &self, + request_id: &str, + _user_data: &Value, + _action: &ElicitationAction, + ) -> bool { + if request_id != NESTED_ELICITATION_ID { + return false; + } + 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() + } +} + +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 + } + } +} + +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( @@ -60,6 +312,119 @@ 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)) +} + +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)) +} + +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?; @@ -100,6 +465,258 @@ 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 +} + +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 +} + +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>)]); 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/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 107880918a57..630e12de20aa 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, ProviderHostCapabilities, 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,148 @@ 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, + ProviderHostCapabilities::default(), + ) + .await + } + + async fn build( + extensions: Vec, + working_dir: PathBuf, _tls_config: Option, + host_capabilities: ProviderHostCapabilities, ) -> 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, + host_capabilities, + ) + .await) + } - Ok(Self { + async fn build_with_command( + resolved_command: PathBuf, + extensions: Vec, + working_dir: PathBuf, + goose_mode: GooseMode, + host_capabilities: ProviderHostCapabilities, + ) -> 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, + supports_form_elicitation: host_capabilities.supports_form_elicitation, + }; + + 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 + } + }; + + 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 +267,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 +610,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 +621,58 @@ 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_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, + 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> { - Box::pin(Self::from_env(tls_config)) + Self::from_env_with_host_capabilities(extensions, tls_config, host_capabilities) } } @@ -504,6 +682,142 @@ 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 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, + 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 +825,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(_) => { @@ -544,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?; @@ -586,6 +912,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 +930,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 +948,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 +973,117 @@ 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, + ProviderHostCapabilities::default(), + ) + .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")); + } + + #[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#" 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, 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,