diff --git a/crates/goose-acp/src/server.rs b/crates/goose-acp/src/server.rs index ee927308..260d919b 100644 --- a/crates/goose-acp/src/server.rs +++ b/crates/goose-acp/src/server.rs @@ -789,19 +789,7 @@ impl GooseAcpAgent { sacp::Error::internal_error().data(format!("Failed to set provider: {}", e)) })?; - for mcp_server in args.mcp_servers { - let config = match mcp_server_to_extension_config(mcp_server) { - Ok(c) => c, - Err(msg) => { - return Err(sacp::Error::invalid_params().data(msg)); - } - }; - let name = config.name().to_string(); - if let Err(e) = agent.add_extension(config, &goose_session.id).await { - return Err(sacp::Error::internal_error() - .data(format!("Failed to add MCP server '{}': {}", name, e))); - } - } + Self::add_mcp_extensions(&agent, args.mcp_servers, &goose_session.id).await?; let session = GooseAcpSession { agent, @@ -841,6 +829,27 @@ impl GooseAcpAgent { Ok(provider) } + async fn add_mcp_extensions( + agent: &Agent, + mcp_servers: Vec, + session_id: &str, + ) -> Result<(), sacp::Error> { + for mcp_server in mcp_servers { + let config = match mcp_server_to_extension_config(mcp_server) { + Ok(c) => c, + Err(msg) => { + return Err(sacp::Error::invalid_params().data(msg)); + } + }; + let name = config.name().to_string(); + if let Err(e) = agent.add_extension(config, session_id).await { + return Err(sacp::Error::internal_error() + .data(format!("Failed to add MCP server '{}': {}", name, e))); + } + } + Ok(()) + } + async fn on_load_session( &self, cx: &JrConnectionCx, @@ -875,6 +884,8 @@ impl GooseAcpAgent { sacp::Error::internal_error().data(format!("Failed to set provider: {}", e)) })?; + Self::add_mcp_extensions(&agent, args.mcp_servers, &session_id).await?; + let conversation = goose_session.conversation.ok_or_else(|| { sacp::Error::internal_error() .data(format!("Session {} has no conversation data", session_id)) diff --git a/crates/goose-acp/tests/common_tests/mod.rs b/crates/goose-acp/tests/common_tests/mod.rs index 2290dae4..5de5837b 100644 --- a/crates/goose-acp/tests/common_tests/mod.rs +++ b/crates/goose-acp/tests/common_tests/mod.rs @@ -243,7 +243,7 @@ pub async fn run_load_model() { assert_eq!(output.text, "2"); let session_id = session.session_id().0.to_string(); - let (_, models) = conn.load_session(&session_id).await; + let (_, models) = conn.load_session(&session_id, vec![]).await; assert_eq!(&*models.unwrap().current_model_id.0, "o4-mini"); } @@ -497,6 +497,61 @@ pub async fn run_prompt_image_attachment() { expected_session_id.assert_matches(&session.session_id().0); } +pub async fn run_load_session_mcp() { + let expected_session_id = ExpectedSessionId::default(); + let prompt = "Use the get_code tool and output only its result."; + let mcp = McpFixture::new(Some(expected_session_id.clone())).await; + let mcp_url = mcp.url.clone(); + + // Two rounds of tool call + tool result: one for new session, one for loaded session. + let openai = OpenAiFixture::new( + vec![ + ( + prompt.to_string(), + include_str!("../test_data/openai_tool_call.txt"), + ), + ( + format!(r#""content":"{FAKE_CODE}""#), + include_str!("../test_data/openai_tool_result.txt"), + ), + ( + prompt.to_string(), + include_str!("../test_data/openai_tool_call.txt"), + ), + ( + format!(r#""content":"{FAKE_CODE}""#), + include_str!("../test_data/openai_tool_result.txt"), + ), + ], + expected_session_id.clone(), + ) + .await; + + let mcp_servers = vec![McpServer::Http(McpServerHttp::new("mcp-fixture", &mcp_url))]; + + let config = TestConnectionConfig { + mcp_servers: mcp_servers.clone(), + ..Default::default() + }; + let mut conn = C::new(config, openai).await; + let (mut session, _) = conn.new_session().await; + expected_session_id.set(session.session_id().0.to_string()); + + // First prompt: tool should work in the new session. + let output = session.prompt(prompt, PermissionDecision::Cancel).await; + assert_eq!(output.text, FAKE_CODE, "tool call failed in new session"); + + // Load the same session with MCP servers re-specified. + let session_id = session.session_id().0.to_string(); + let (mut loaded_session, _) = conn.load_session(&session_id, mcp_servers).await; + + // Second prompt: tool should work in the loaded session. + let output = loaded_session + .prompt(prompt, PermissionDecision::Cancel) + .await; + assert_eq!(output.text, FAKE_CODE, "tool call failed in loaded session"); +} + pub async fn run_prompt_mcp() { let expected_session_id = ExpectedSessionId::default(); let mcp = McpFixture::new(Some(expected_session_id.clone())).await; diff --git a/crates/goose-acp/tests/fixtures/mod.rs b/crates/goose-acp/tests/fixtures/mod.rs index 5c25eb61..bef807f3 100644 --- a/crates/goose-acp/tests/fixtures/mod.rs +++ b/crates/goose-acp/tests/fixtures/mod.rs @@ -299,6 +299,7 @@ pub trait Connection: Sized { async fn load_session( &mut self, session_id: &str, + mcp_servers: Vec, ) -> (Self::Session, Option); fn auth_methods(&self) -> &[AuthMethod]; fn reset_openai(&self); diff --git a/crates/goose-acp/tests/fixtures/provider.rs b/crates/goose-acp/tests/fixtures/provider.rs index 232068a6..7256b8cd 100644 --- a/crates/goose-acp/tests/fixtures/provider.rs +++ b/crates/goose-acp/tests/fixtures/provider.rs @@ -12,7 +12,7 @@ use goose::permission::permission_confirmation::PrincipalType; use goose::permission::{Permission, PermissionConfirmation}; use goose::providers::base::Provider; use goose_test_support::TEST_MODEL; -use sacp::schema::{AuthMethod, SessionModelState, ToolCallStatus}; +use sacp::schema::{AuthMethod, McpServer, SessionModelState, ToolCallStatus}; use std::sync::Arc; use tokio::sync::Mutex; @@ -180,6 +180,7 @@ impl Connection for ClientToProviderConnection { async fn load_session( &mut self, _session_id: &str, + _mcp_servers: Vec, ) -> (ClientToProviderSession, Option) { unimplemented!("TODO: implement load_session in ACP provider") } diff --git a/crates/goose-acp/tests/fixtures/server.rs b/crates/goose-acp/tests/fixtures/server.rs index 9b59cb03..95c6319d 100644 --- a/crates/goose-acp/tests/fixtures/server.rs +++ b/crates/goose-acp/tests/fixtures/server.rs @@ -250,13 +250,17 @@ impl Connection for ClientToAgentConnection { async fn load_session( &mut self, session_id: &str, + mcp_servers: Vec, ) -> (ClientToAgentSession, Option) { self.updates.lock().unwrap().clear(); let work_dir = tempfile::tempdir().unwrap(); let session_id = sacp::schema::SessionId::new(session_id.to_string()); let response = self .cx - .send_request(LoadSessionRequest::new(session_id.clone(), work_dir.path())) + .send_request( + LoadSessionRequest::new(session_id.clone(), work_dir.path()) + .mcp_servers(mcp_servers), + ) .block_task() .await .unwrap(); diff --git a/crates/goose-acp/tests/provider_test.rs b/crates/goose-acp/tests/provider_test.rs index 15d72c76..d0ab8587 100644 --- a/crates/goose-acp/tests/provider_test.rs +++ b/crates/goose-acp/tests/provider_test.rs @@ -6,8 +6,9 @@ use common_tests::fixtures::run_test; use common_tests::{ run_config_mcp, run_fs_read_text_file_true, run_fs_write_text_file_false, run_fs_write_text_file_true, run_initialize_doesnt_hit_provider, run_load_model, - run_model_list, run_model_set, run_permission_persistence, run_prompt_basic, - run_prompt_codemode, run_prompt_image, run_prompt_image_attachment, run_prompt_mcp, + run_load_session_mcp, run_model_list, run_model_set, run_permission_persistence, + run_prompt_basic, run_prompt_codemode, run_prompt_image, run_prompt_image_attachment, + run_prompt_mcp, }; #[test] @@ -43,6 +44,12 @@ fn test_provider_load_model() { run_test(async { run_load_model::().await }); } +#[test] +#[ignore = "TODO: implement load_session in ACP provider"] +fn test_provider_load_session_mcp() { + run_test(async { run_load_session_mcp::().await }); +} + #[test] fn test_provider_model_list() { run_test(async { run_model_list::().await }); diff --git a/crates/goose-acp/tests/server_test.rs b/crates/goose-acp/tests/server_test.rs index 026bbf62..bd0d5754 100644 --- a/crates/goose-acp/tests/server_test.rs +++ b/crates/goose-acp/tests/server_test.rs @@ -4,9 +4,9 @@ use common_tests::fixtures::server::ClientToAgentConnection; use common_tests::{ run_config_mcp, run_fs_read_text_file_true, run_fs_write_text_file_false, run_fs_write_text_file_true, run_initialize_doesnt_hit_provider, - run_initialize_without_provider, run_load_model, run_model_list, run_model_set, - run_permission_persistence, run_prompt_basic, run_prompt_codemode, run_prompt_image, - run_prompt_image_attachment, run_prompt_mcp, + run_initialize_without_provider, run_load_model, run_load_session_mcp, run_model_list, + run_model_set, run_permission_persistence, run_prompt_basic, run_prompt_codemode, + run_prompt_image, run_prompt_image_attachment, run_prompt_mcp, }; #[test] @@ -79,6 +79,11 @@ fn test_prompt_image_attachment() { run_test(async { run_prompt_image_attachment::().await }); } +#[test] +fn test_load_session_mcp() { + run_test(async { run_load_session_mcp::().await }); +} + #[test] fn test_prompt_mcp() { run_test(async { run_prompt_mcp::().await });