diff --git a/crates/goose/src/acp/server.rs b/crates/goose/src/acp/server.rs index c3554b6019b0..30831ad2306a 100644 --- a/crates/goose/src/acp/server.rs +++ b/crates/goose/src/acp/server.rs @@ -613,6 +613,23 @@ impl GooseAcpAgent { ) } + pub(super) async fn prepare_session_setup_by_id( + &self, + session_id: &str, + ) -> Result<(Session, SessionUsageTotals), agent_client_protocol::Error> { + let session = self + .session_manager + .get_session(session_id, false) + .await + .internal_err_ctx("Failed to load session for setup notifications")?; + let totals = self + .session_manager + .get_session_usage_totals(session_id) + .await + .unwrap_or_default(); + Ok((session, totals)) + } + pub(super) fn supports_recipe_param_requests(&self) -> bool { self.client_supports_recipe_param_requests .get() diff --git a/crates/goose/src/acp/server/dispatch.rs b/crates/goose/src/acp/server/dispatch.rs index 64e7513b1a27..faa6c73e900b 100644 --- a/crates/goose/src/acp/server/dispatch.rs +++ b/crates/goose/src/acp/server/dispatch.rs @@ -44,7 +44,31 @@ impl HandleDispatchFrom for GooseAcpHandler { let agent = agent.clone(); let cx_clone = cx.clone(); cx.spawn(async move { - responder.respond_with_result(agent.on_new_session(&cx_clone, req).await)?; + match agent.on_new_session(&cx_clone, req).await { + Ok(response) => { + let session_id = response.session_id.0.to_string(); + let session_setup = + agent.prepare_session_setup_by_id(&session_id).await; + responder.respond(response)?; + if let Err(error) = session_setup.and_then(|(session, totals)| { + send_session_setup_notifications( + &cx_clone, + &session, + &totals, + agent.supports_goose_custom_notifications(), + ) + }) { + tracing::warn!( + session_id = %session_id, + error = ?error, + "Failed to send ACP session setup notifications" + ); + } + } + Err(error) => { + responder.respond_with_error(error)?; + } + } Ok(()) })?; Ok(()) @@ -377,7 +401,31 @@ impl HandleDispatchFrom for GooseAcpHandler { |req: ForkSessionRequest, responder: Responder| async move { let cx_spawn = cx.clone(); cx.spawn(async move { - responder.respond_with_result(agent.on_fork_session(&cx_spawn, req).await)?; + match agent.on_fork_session(&cx_spawn, req).await { + Ok(response) => { + let session_id = response.session_id.0.to_string(); + let session_setup = + agent.prepare_session_setup_by_id(&session_id).await; + responder.respond(response)?; + if let Err(error) = session_setup.and_then(|(session, totals)| { + send_session_setup_notifications( + &cx_spawn, + &session, + &totals, + agent.supports_goose_custom_notifications(), + ) + }) { + tracing::warn!( + session_id = %session_id, + error = ?error, + "Failed to send ACP forked session setup notifications" + ); + } + } + Err(error) => { + responder.respond_with_error(error)?; + } + } Ok(()) })?; Ok(()) diff --git a/crates/goose/src/acp/server/fork_session.rs b/crates/goose/src/acp/server/fork_session.rs index 892ef0789950..46202445fb00 100644 --- a/crates/goose/src/acp/server/fork_session.rs +++ b/crates/goose/src/acp/server/fork_session.rs @@ -72,7 +72,6 @@ impl GooseAcpAgent { if let Some(co) = config_options { response = response.config_options(co); } - self.notify_session_setup(cx, &goose_session).await?; Ok(response) } } diff --git a/crates/goose/src/acp/server/new_session.rs b/crates/goose/src/acp/server/new_session.rs index c968e42c2994..b3c192dc11c9 100644 --- a/crates/goose/src/acp/server/new_session.rs +++ b/crates/goose/src/acp/server/new_session.rs @@ -79,7 +79,6 @@ impl GooseAcpAgent { let response = self .build_new_session_response(&reloaded_session, &extension_results) .await?; - self.notify_session_setup(cx, &reloaded_session).await?; Ok(response) } diff --git a/crates/goose/tests/acp_fixtures/server.rs b/crates/goose/tests/acp_fixtures/server.rs index dc86a8ac5e6c..98a8fcc60d0e 100644 --- a/crates/goose/tests/acp_fixtures/server.rs +++ b/crates/goose/tests/acp_fixtures/server.rs @@ -15,6 +15,8 @@ use agent_client_protocol::schema::v1::{ use agent_client_protocol::schema::ProtocolVersion; use agent_client_protocol::{Agent, Client, ConnectionTo}; use async_trait::async_trait; +use futures::io::BufReader; +use futures::{AsyncBufReadExt, AsyncWriteExt, StreamExt}; use goose::config::PermissionManager; use goose_test_support::{ExpectedSessionId, IgnoreSessionId}; use std::sync::{Arc, Mutex}; @@ -103,6 +105,78 @@ impl AcpServerConnection { } } +pub async fn assert_session_response_precedes_available_commands( + transport: super::DuplexTransport, + method: &str, + params: serde_json::Value, +) { + let agent_client_protocol::ByteStreams { + mut outgoing, + incoming, + } = transport; + let mut incoming = BufReader::new(incoming).lines(); + + let initialize = serde_json::json!({ + "jsonrpc": "2.0", + "id": 1, + "method": "initialize", + "params": { + "protocolVersion": 1, + "clientCapabilities": {} + } + }); + outgoing + .write_all(format!("{initialize}\n").as_bytes()) + .await + .unwrap(); + outgoing.flush().await.unwrap(); + let initialize_response = incoming.next().await.unwrap().unwrap(); + let initialize_response: serde_json::Value = + serde_json::from_str(&initialize_response).unwrap(); + assert_eq!(initialize_response["id"], 1); + + let request = serde_json::json!({ + "jsonrpc": "2.0", + "id": 2, + "method": method, + "params": params, + }); + outgoing + .write_all(format!("{request}\n").as_bytes()) + .await + .unwrap(); + outgoing.flush().await.unwrap(); + + let response = incoming.next().await.unwrap().unwrap(); + let response: serde_json::Value = serde_json::from_str(&response).unwrap(); + assert_eq!(response["id"], 2); + let session_id = response["result"]["sessionId"] + .as_str() + .unwrap() + .to_string(); + + let deadline = tokio::time::Instant::now() + Duration::from_secs(5); + loop { + let message = tokio::time::timeout_at(deadline, incoming.next()) + .await + .expect("timed out waiting for available commands") + .expect("ACP connection closed") + .unwrap(); + let message: serde_json::Value = serde_json::from_str(&message).unwrap(); + if message["params"]["update"]["sessionUpdate"] != "available_commands_update" { + continue; + } + + assert_eq!(message["params"]["sessionId"], session_id); + assert!(message["params"]["update"]["availableCommands"] + .as_array() + .unwrap() + .iter() + .any(|command| command["name"] == "goal")); + break; + } +} + #[async_trait] impl Connection for AcpServerConnection { type Session = AcpServerSession; @@ -343,6 +417,20 @@ impl Connection for AcpServerConnection { _work_dir: work_dir, }; let models = extract_model_state_from_config_options(response.config_options.as_deref()); + let deadline = tokio::time::Instant::now() + Duration::from_secs(5); + while !self.updates.lock().unwrap().iter().any(|notification| { + notification.session_id == response.session_id + && matches!( + ¬ification.update, + SessionUpdate::AvailableCommandsUpdate(_) + ) + }) { + assert!( + tokio::time::Instant::now() < deadline, + "timed out waiting for initial available commands" + ); + tokio::task::yield_now().await; + } self.updates.lock().unwrap().clear(); Ok(SessionData { session, diff --git a/crates/goose/tests/acp_fork_session_test.rs b/crates/goose/tests/acp_fork_session_test.rs index 82c181991585..3d26f1ec03ff 100644 --- a/crates/goose/tests/acp_fork_session_test.rs +++ b/crates/goose/tests/acp_fork_session_test.rs @@ -4,8 +4,12 @@ mod common_tests; use agent_client_protocol::schema::v1::{ForkSessionRequest, ForkSessionResponse, SessionId}; -use common_tests::fixtures::server::AcpServerConnection; -use common_tests::fixtures::{run_test, Connection, OpenAiFixture, TestConnectionConfig}; +use common_tests::fixtures::server::{ + assert_session_response_precedes_available_commands, AcpServerConnection, +}; +use common_tests::fixtures::{ + run_test, spawn_acp_server_in_process, Connection, OpenAiFixture, TestConnectionConfig, +}; use goose::config::GooseMode; use goose::conversation::message::{Message, MessageContent}; use goose::session::{SessionManager, SessionType}; @@ -92,6 +96,42 @@ fn conversation_before_meta(timestamp: i64) -> serde_json::Map::expected_session_id(), + ) + .await; + let (transport, _handle, _permission_manager) = spawn_acp_server_in_process( + openai.uri(), + &[], + data_root.path(), + GooseMode::default(), + None, + goose_test_support::TEST_MODEL, + true, + ) + .await; + + assert_session_response_precedes_available_commands( + transport, + "session/fork", + serde_json::json!({ + "sessionId": source.id, + "cwd": cwd.path(), + "mcpServers": [] + }), + ) + .await; + }); +} + #[test] fn fork_session_conversation_before_matches_rest_cutoff() { run_test(async { diff --git a/crates/goose/tests/acp_server_test.rs b/crates/goose/tests/acp_server_test.rs index e3dd70edb855..a75b348f1969 100644 --- a/crates/goose/tests/acp_server_test.rs +++ b/crates/goose/tests/acp_server_test.rs @@ -8,8 +8,12 @@ use agent_client_protocol::schema::v1::{ SetSessionConfigOptionRequest, }; use agent_client_protocol::ErrorCode; -use common_tests::fixtures::server::AcpServerConnection; -use common_tests::fixtures::{run_test, Connection, OpenAiFixture, Session, TestConnectionConfig}; +use common_tests::fixtures::server::{ + assert_session_response_precedes_available_commands, AcpServerConnection, +}; +use common_tests::fixtures::{ + run_test, spawn_acp_server_in_process, Connection, OpenAiFixture, Session, TestConnectionConfig, +}; #[cfg(feature = "code-mode")] use common_tests::run_prompt_codemode; use common_tests::{ @@ -631,6 +635,38 @@ fn test_new_session_uses_current_config_mode() { run_test(async { run_new_session_uses_current_config_mode::().await }); } +#[test] +fn test_new_session_response_precedes_available_commands() { + run_test(async { + let data_root = tempfile::tempdir().unwrap(); + let work_dir = tempfile::tempdir().unwrap(); + let openai = OpenAiFixture::new( + vec![], + ::expected_session_id(), + ) + .await; + let (transport, _handle, _permission_manager) = spawn_acp_server_in_process( + openai.uri(), + &[], + data_root.path(), + GooseMode::default(), + None, + goose_test_support::TEST_MODEL, + true, + ) + .await; + assert_session_response_precedes_available_commands( + transport, + "session/new", + serde_json::json!({ + "cwd": work_dir.path(), + "mcpServers": [] + }), + ) + .await; + }); +} + #[test] fn test_new_session_honors_recipe_model_without_recipe_provider() { run_test(async {