feat(goose-acp): enable parallel sessions with isolated agent state (#6392)
Signed-off-by: Adrian Cole <adrian@tetrate.io>
This commit is contained in:
@@ -33,7 +33,7 @@ use super::tool_execution::ToolCallResult;
|
||||
use super::types::SharedProvider;
|
||||
use crate::agents::extension::{Envs, ProcessExit};
|
||||
use crate::agents::extension_malware_check;
|
||||
use crate::agents::mcp_client::{McpClient, McpClientTrait};
|
||||
use crate::agents::mcp_client::{McpClient, McpClientTrait, McpMeta};
|
||||
use crate::config::search_path::SearchPaths;
|
||||
use crate::config::{get_all_extensions, Config};
|
||||
use crate::oauth::oauth_flow;
|
||||
@@ -93,7 +93,7 @@ impl Extension {
|
||||
/// Manages goose extensions / MCP clients and their interactions
|
||||
pub struct ExtensionManager {
|
||||
extensions: Mutex<HashMap<String, Extension>>,
|
||||
context: Mutex<PlatformExtensionContext>,
|
||||
context: PlatformExtensionContext,
|
||||
provider: SharedProvider,
|
||||
tools_cache: Mutex<Option<Arc<Vec<Tool>>>>,
|
||||
tools_cache_version: AtomicU64,
|
||||
@@ -210,12 +210,6 @@ pub fn get_parameter_names(tool: &Tool) -> Vec<String> {
|
||||
names
|
||||
}
|
||||
|
||||
impl Default for ExtensionManager {
|
||||
fn default() -> Self {
|
||||
Self::new(Arc::new(Mutex::new(None)))
|
||||
}
|
||||
}
|
||||
|
||||
async fn child_process_client(
|
||||
mut command: Command,
|
||||
timeout: &Option<u64>,
|
||||
@@ -446,45 +440,36 @@ async fn create_streamable_http_client(
|
||||
}
|
||||
|
||||
impl ExtensionManager {
|
||||
pub fn new(provider: SharedProvider) -> Self {
|
||||
pub fn new(
|
||||
provider: SharedProvider,
|
||||
session_manager: Arc<crate::session::SessionManager>,
|
||||
) -> Self {
|
||||
Self {
|
||||
extensions: Mutex::new(HashMap::new()),
|
||||
context: Mutex::new(PlatformExtensionContext {
|
||||
session_id: None,
|
||||
context: PlatformExtensionContext {
|
||||
extension_manager: None,
|
||||
}),
|
||||
session_manager,
|
||||
},
|
||||
provider,
|
||||
tools_cache: Mutex::new(None),
|
||||
tools_cache_version: AtomicU64::new(0),
|
||||
}
|
||||
}
|
||||
|
||||
/// Create a new ExtensionManager with no provider (useful for tests)
|
||||
pub fn new_without_provider() -> Self {
|
||||
Self::new(Arc::new(Mutex::new(None)))
|
||||
#[cfg(test)]
|
||||
pub fn new_without_provider(data_dir: std::path::PathBuf) -> Self {
|
||||
let session_manager = Arc::new(crate::session::SessionManager::new(data_dir));
|
||||
Self::new(Arc::new(Mutex::new(None)), session_manager)
|
||||
}
|
||||
|
||||
pub async fn set_context(&self, context: PlatformExtensionContext) {
|
||||
*self.context.lock().await = context;
|
||||
}
|
||||
|
||||
pub async fn get_context(&self) -> PlatformExtensionContext {
|
||||
self.context.lock().await.clone()
|
||||
pub fn get_context(&self) -> &PlatformExtensionContext {
|
||||
&self.context
|
||||
}
|
||||
|
||||
/// Resolve the working directory for an extension.
|
||||
/// Priority: session working_dir > current_dir
|
||||
/// Falls back to current_dir when working_dir is not available.
|
||||
async fn resolve_working_dir(&self) -> PathBuf {
|
||||
// Try to get working_dir from session via context
|
||||
if let Some(ref session_id) = self.context.lock().await.session_id {
|
||||
if let Ok(session) =
|
||||
crate::session::SessionManager::get_session(session_id, false).await
|
||||
{
|
||||
return session.working_dir;
|
||||
}
|
||||
}
|
||||
|
||||
// Fall back to current_dir
|
||||
// Fall back to current_dir - working_dir is passed through the call chain from session
|
||||
std::env::current_dir().unwrap_or_default()
|
||||
}
|
||||
|
||||
@@ -496,7 +481,7 @@ impl ExtensionManager {
|
||||
.any(|ext| ext.supports_resources())
|
||||
}
|
||||
|
||||
pub async fn add_extension(&self, config: ExtensionConfig) -> ExtensionResult<()> {
|
||||
pub async fn add_extension(self: &Arc<Self>, config: ExtensionConfig) -> ExtensionResult<()> {
|
||||
let config_name = config.key().to_string();
|
||||
let sanitized_name = normalize(&config_name);
|
||||
|
||||
@@ -564,32 +549,23 @@ impl ExtensionManager {
|
||||
Box::new(client)
|
||||
}
|
||||
ExtensionConfig::Builtin { name, timeout, .. } => {
|
||||
let cmd = std::env::current_exe()
|
||||
.and_then(|path| {
|
||||
path.to_str().map(|s| s.to_string()).ok_or_else(|| {
|
||||
std::io::Error::new(
|
||||
std::io::ErrorKind::InvalidData,
|
||||
"Invalid UTF-8 in executable path",
|
||||
)
|
||||
})
|
||||
})
|
||||
.map_err(|e| {
|
||||
ExtensionError::ConfigError(format!(
|
||||
"Failed to resolve executable path: {}",
|
||||
e
|
||||
))
|
||||
let timeout_duration = Duration::from_secs(timeout.unwrap_or(300));
|
||||
let def = goose_mcp::BUILTIN_EXTENSIONS
|
||||
.get(name.as_str())
|
||||
.ok_or_else(|| {
|
||||
ExtensionError::ConfigError(format!("Unknown builtin extension: {}", name))
|
||||
})?;
|
||||
let command = Command::new(cmd).configure(|command| {
|
||||
command.arg("mcp").arg(name);
|
||||
});
|
||||
let client = child_process_client(
|
||||
command,
|
||||
timeout,
|
||||
self.provider.clone(),
|
||||
Some(&effective_working_dir),
|
||||
let (server_read, client_write) = tokio::io::duplex(65536);
|
||||
let (client_read, server_write) = tokio::io::duplex(65536);
|
||||
(def.spawn_server)(server_read, server_write);
|
||||
Box::new(
|
||||
McpClient::connect(
|
||||
(client_read, client_write),
|
||||
timeout_duration,
|
||||
self.provider.clone(),
|
||||
)
|
||||
.await?,
|
||||
)
|
||||
.await?;
|
||||
Box::new(client)
|
||||
}
|
||||
ExtensionConfig::Platform { name, .. } => {
|
||||
let normalized_key = normalize(name);
|
||||
@@ -598,7 +574,8 @@ impl ExtensionManager {
|
||||
.ok_or_else(|| {
|
||||
ExtensionError::ConfigError(format!("Unknown platform extension: {}", name))
|
||||
})?;
|
||||
let context = self.get_context().await;
|
||||
let mut context = self.context.clone();
|
||||
context.extension_manager = Some(Arc::downgrade(self));
|
||||
(def.client_factory)(context)
|
||||
}
|
||||
ExtensionConfig::InlinePython {
|
||||
@@ -1132,6 +1109,7 @@ impl ExtensionManager {
|
||||
|
||||
pub async fn dispatch_tool_call(
|
||||
&self,
|
||||
session_id: &str,
|
||||
tool_call: CallToolRequestParam,
|
||||
cancellation_token: CancellationToken,
|
||||
) -> Result<ToolCallResult> {
|
||||
@@ -1191,11 +1169,18 @@ impl ExtensionManager {
|
||||
let arguments = tool_call.arguments.clone();
|
||||
let client = client.clone();
|
||||
let notifications_receiver = client.lock().await.subscribe().await;
|
||||
let session_id = session_id.to_string();
|
||||
|
||||
let fut = async move {
|
||||
tracing::debug!(
|
||||
"dispatch_tool_call fut: calling client.call_tool tool={} session_id={}",
|
||||
tool_name,
|
||||
session_id
|
||||
);
|
||||
let client_guard = client.lock().await;
|
||||
let meta = McpMeta::new(&session_id);
|
||||
client_guard
|
||||
.call_tool(&tool_name, arguments, cancellation_token)
|
||||
.call_tool(&tool_name, arguments, meta, cancellation_token)
|
||||
.await
|
||||
.map_err(|e| match e {
|
||||
ServiceError::McpError(error_data) => error_data,
|
||||
@@ -1376,7 +1361,11 @@ impl ExtensionManager {
|
||||
.map(|ext| ext.get_client())
|
||||
}
|
||||
|
||||
pub async fn collect_moim(&self, working_dir: &std::path::Path) -> Option<String> {
|
||||
pub async fn collect_moim(
|
||||
&self,
|
||||
session_id: &str,
|
||||
working_dir: &std::path::Path,
|
||||
) -> Option<String> {
|
||||
// Use minute-level granularity to prevent conversation changes every second
|
||||
let timestamp = chrono::Local::now().format("%Y-%m-%d %H:%M:00").to_string();
|
||||
let mut content = format!(
|
||||
@@ -1401,7 +1390,7 @@ impl ExtensionManager {
|
||||
|
||||
for (name, client) in platform_clients {
|
||||
let client_guard = client.lock().await;
|
||||
if let Some(moim_content) = client_guard.get_moim().await {
|
||||
if let Some(moim_content) = client_guard.get_moim(session_id).await {
|
||||
tracing::debug!("MOIM content from {}: {} chars", name, moim_content.len());
|
||||
content.push('\n');
|
||||
content.push_str(&moim_content);
|
||||
@@ -1517,6 +1506,7 @@ mod tests {
|
||||
&self,
|
||||
name: &str,
|
||||
_arguments: Option<JsonObject>,
|
||||
_meta: McpMeta,
|
||||
_cancellation_token: CancellationToken,
|
||||
) -> Result<CallToolResult, Error> {
|
||||
match name {
|
||||
@@ -1554,7 +1544,9 @@ mod tests {
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_get_client_for_tool() {
|
||||
let extension_manager = ExtensionManager::new_without_provider();
|
||||
let temp_dir = tempfile::tempdir().unwrap();
|
||||
let extension_manager =
|
||||
ExtensionManager::new_without_provider(temp_dir.path().to_path_buf());
|
||||
|
||||
// Add some mock clients using the helper method
|
||||
extension_manager
|
||||
@@ -1614,7 +1606,9 @@ mod tests {
|
||||
async fn test_dispatch_tool_call() {
|
||||
// test that dispatch_tool_call parses out the sanitized name correctly, and extracts
|
||||
// tool_names
|
||||
let extension_manager = ExtensionManager::new_without_provider();
|
||||
let temp_dir = tempfile::tempdir().unwrap();
|
||||
let extension_manager =
|
||||
ExtensionManager::new_without_provider(temp_dir.path().to_path_buf());
|
||||
|
||||
// Add some mock clients using the helper method
|
||||
extension_manager
|
||||
@@ -1645,7 +1639,7 @@ mod tests {
|
||||
};
|
||||
|
||||
let result = extension_manager
|
||||
.dispatch_tool_call(tool_call, CancellationToken::default())
|
||||
.dispatch_tool_call("test-session-id", tool_call, CancellationToken::default())
|
||||
.await;
|
||||
assert!(result.is_ok());
|
||||
|
||||
@@ -1655,7 +1649,7 @@ mod tests {
|
||||
};
|
||||
|
||||
let result = extension_manager
|
||||
.dispatch_tool_call(tool_call, CancellationToken::default())
|
||||
.dispatch_tool_call("test-session-id", tool_call, CancellationToken::default())
|
||||
.await;
|
||||
assert!(result.is_ok());
|
||||
|
||||
@@ -1666,7 +1660,7 @@ mod tests {
|
||||
};
|
||||
|
||||
let result = extension_manager
|
||||
.dispatch_tool_call(tool_call, CancellationToken::default())
|
||||
.dispatch_tool_call("test-session-id", tool_call, CancellationToken::default())
|
||||
.await;
|
||||
assert!(result.is_ok());
|
||||
|
||||
@@ -1677,7 +1671,7 @@ mod tests {
|
||||
};
|
||||
|
||||
let result = extension_manager
|
||||
.dispatch_tool_call(tool_call, CancellationToken::default())
|
||||
.dispatch_tool_call("test-session-id", tool_call, CancellationToken::default())
|
||||
.await;
|
||||
assert!(result.is_ok());
|
||||
|
||||
@@ -1687,7 +1681,7 @@ mod tests {
|
||||
};
|
||||
|
||||
let result = extension_manager
|
||||
.dispatch_tool_call(tool_call, CancellationToken::default())
|
||||
.dispatch_tool_call("test-session-id", tool_call, CancellationToken::default())
|
||||
.await;
|
||||
assert!(result.is_ok());
|
||||
|
||||
@@ -1698,7 +1692,11 @@ mod tests {
|
||||
};
|
||||
|
||||
let result = extension_manager
|
||||
.dispatch_tool_call(invalid_tool_call, CancellationToken::default())
|
||||
.dispatch_tool_call(
|
||||
"test-session-id",
|
||||
invalid_tool_call,
|
||||
CancellationToken::default(),
|
||||
)
|
||||
.await
|
||||
.unwrap()
|
||||
.result
|
||||
@@ -1719,7 +1717,11 @@ mod tests {
|
||||
};
|
||||
|
||||
let result = extension_manager
|
||||
.dispatch_tool_call(invalid_tool_call, CancellationToken::default())
|
||||
.dispatch_tool_call(
|
||||
"test-session-id",
|
||||
invalid_tool_call,
|
||||
CancellationToken::default(),
|
||||
)
|
||||
.await;
|
||||
if let Err(err) = result {
|
||||
let tool_err = err.downcast_ref::<ErrorData>().expect("Expected ErrorData");
|
||||
@@ -1731,7 +1733,9 @@ mod tests {
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_tool_availability_filtering() {
|
||||
let extension_manager = ExtensionManager::new_without_provider();
|
||||
let temp_dir = tempfile::tempdir().unwrap();
|
||||
let extension_manager =
|
||||
ExtensionManager::new_without_provider(temp_dir.path().to_path_buf());
|
||||
|
||||
// Only "available_tool" should be available to the LLM
|
||||
let available_tools = vec!["available_tool".to_string()];
|
||||
@@ -1759,7 +1763,9 @@ mod tests {
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_tool_availability_defaults_to_available() {
|
||||
let extension_manager = ExtensionManager::new_without_provider();
|
||||
let temp_dir = tempfile::tempdir().unwrap();
|
||||
let extension_manager =
|
||||
ExtensionManager::new_without_provider(temp_dir.path().to_path_buf());
|
||||
|
||||
extension_manager
|
||||
.add_mock_extension_with_tools(
|
||||
@@ -1784,7 +1790,9 @@ mod tests {
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_dispatch_unavailable_tool_returns_error() {
|
||||
let extension_manager = ExtensionManager::new_without_provider();
|
||||
let temp_dir = tempfile::tempdir().unwrap();
|
||||
let extension_manager =
|
||||
ExtensionManager::new_without_provider(temp_dir.path().to_path_buf());
|
||||
|
||||
let available_tools = vec!["available_tool".to_string()];
|
||||
|
||||
@@ -1803,7 +1811,11 @@ mod tests {
|
||||
};
|
||||
|
||||
let result = extension_manager
|
||||
.dispatch_tool_call(unavailable_tool_call, CancellationToken::default())
|
||||
.dispatch_tool_call(
|
||||
"test-session-id",
|
||||
unavailable_tool_call,
|
||||
CancellationToken::default(),
|
||||
)
|
||||
.await;
|
||||
|
||||
// Should return RESOURCE_NOT_FOUND error
|
||||
@@ -1822,7 +1834,11 @@ mod tests {
|
||||
};
|
||||
|
||||
let result = extension_manager
|
||||
.dispatch_tool_call(available_tool_call, CancellationToken::default())
|
||||
.dispatch_tool_call(
|
||||
"test-session-id",
|
||||
available_tool_call,
|
||||
CancellationToken::default(),
|
||||
)
|
||||
.await;
|
||||
|
||||
assert!(result.is_ok());
|
||||
@@ -1893,10 +1909,11 @@ mod tests {
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_collect_moim_uses_minute_granularity() {
|
||||
let em = ExtensionManager::new_without_provider();
|
||||
let temp_dir = tempfile::tempdir().unwrap();
|
||||
let em = ExtensionManager::new_without_provider(temp_dir.path().to_path_buf());
|
||||
let working_dir = std::path::Path::new("/tmp");
|
||||
|
||||
if let Some(moim) = em.collect_moim(working_dir).await {
|
||||
if let Some(moim) = em.collect_moim("test-session-id", working_dir).await {
|
||||
// Timestamp should end with :00 (seconds fixed to 00)
|
||||
assert!(
|
||||
moim.contains(":00\n"),
|
||||
@@ -1907,7 +1924,9 @@ mod tests {
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_tools_cache_invalidated_on_add_extension() {
|
||||
let extension_manager = ExtensionManager::new_without_provider();
|
||||
let temp_dir = tempfile::tempdir().unwrap();
|
||||
let extension_manager =
|
||||
ExtensionManager::new_without_provider(temp_dir.path().to_path_buf());
|
||||
|
||||
extension_manager
|
||||
.add_mock_extension(
|
||||
@@ -1942,7 +1961,9 @@ mod tests {
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_tools_cache_invalidated_on_remove_extension() {
|
||||
let extension_manager = ExtensionManager::new_without_provider();
|
||||
let temp_dir = tempfile::tempdir().unwrap();
|
||||
let extension_manager =
|
||||
ExtensionManager::new_without_provider(temp_dir.path().to_path_buf());
|
||||
|
||||
extension_manager
|
||||
.add_mock_extension(
|
||||
@@ -1972,7 +1993,9 @@ mod tests {
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_get_prefixed_tools_excluding() {
|
||||
let extension_manager = ExtensionManager::new_without_provider();
|
||||
let temp_dir = tempfile::tempdir().unwrap();
|
||||
let extension_manager =
|
||||
ExtensionManager::new_without_provider(temp_dir.path().to_path_buf());
|
||||
|
||||
extension_manager
|
||||
.add_mock_extension(
|
||||
@@ -1999,7 +2022,9 @@ mod tests {
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_get_prefixed_tools_by_extension_name() {
|
||||
let extension_manager = ExtensionManager::new_without_provider();
|
||||
let temp_dir = tempfile::tempdir().unwrap();
|
||||
let extension_manager =
|
||||
ExtensionManager::new_without_provider(temp_dir.path().to_path_buf());
|
||||
|
||||
extension_manager
|
||||
.add_mock_extension(
|
||||
|
||||
Reference in New Issue
Block a user