fix: per session working dir isolation (#6920)
This commit is contained in:
@@ -581,6 +581,7 @@ impl Agent {
|
||||
.dispatch_tool_call(
|
||||
&session.id,
|
||||
tool_call.clone(),
|
||||
Some(session.working_dir.as_path()),
|
||||
cancellation_token.unwrap_or_default(),
|
||||
)
|
||||
.await;
|
||||
|
||||
@@ -531,6 +531,7 @@ impl McpClientTrait for AppsManagerClient {
|
||||
session_id: &str,
|
||||
name: &str,
|
||||
arguments: Option<JsonObject>,
|
||||
_working_dir: Option<&str>,
|
||||
_cancel_token: CancellationToken,
|
||||
) -> Result<CallToolResult, Error> {
|
||||
let result = match name {
|
||||
|
||||
@@ -297,6 +297,7 @@ impl McpClientTrait for ChatRecallClient {
|
||||
session_id: &str,
|
||||
name: &str,
|
||||
arguments: Option<JsonObject>,
|
||||
_working_dir: Option<&str>,
|
||||
_cancellation_token: CancellationToken,
|
||||
) -> Result<CallToolResult, Error> {
|
||||
let content = match name {
|
||||
|
||||
@@ -262,7 +262,7 @@ fn create_tool_callback(
|
||||
arguments: args.and_then(|v| v.as_object().cloned()),
|
||||
};
|
||||
match manager
|
||||
.dispatch_tool_call(&session_id, tool_call, CancellationToken::new())
|
||||
.dispatch_tool_call(&session_id, tool_call, None, CancellationToken::new())
|
||||
.await
|
||||
{
|
||||
Ok(dispatch_result) => match dispatch_result.result.await {
|
||||
@@ -422,6 +422,7 @@ impl McpClientTrait for CodeExecutionClient {
|
||||
session_id: &str,
|
||||
name: &str,
|
||||
arguments: Option<JsonObject>,
|
||||
_working_dir: Option<&str>,
|
||||
_cancellation_token: CancellationToken,
|
||||
) -> Result<CallToolResult, Error> {
|
||||
let result = match name {
|
||||
|
||||
@@ -599,16 +599,6 @@ impl ExtensionManager {
|
||||
.await?;
|
||||
Box::new(client)
|
||||
} else {
|
||||
// Set GOOSE_WORKING_DIR in the current process for builtin extensions
|
||||
// since they run in-process and read from std::env::var
|
||||
if effective_working_dir.exists() && effective_working_dir.is_dir() {
|
||||
std::env::set_var("GOOSE_WORKING_DIR", &effective_working_dir);
|
||||
tracing::info!(
|
||||
"Set GOOSE_WORKING_DIR for builtin extension: {:?}",
|
||||
effective_working_dir
|
||||
);
|
||||
}
|
||||
|
||||
let (server_read, client_write) = tokio::io::duplex(65536);
|
||||
let (client_read, server_write) = tokio::io::duplex(65536);
|
||||
extension_fn(server_read, server_write);
|
||||
@@ -1189,6 +1179,7 @@ impl ExtensionManager {
|
||||
&self,
|
||||
session_id: &str,
|
||||
tool_call: CallToolRequestParams,
|
||||
working_dir: Option<&std::path::Path>,
|
||||
cancellation_token: CancellationToken,
|
||||
) -> Result<ToolCallResult> {
|
||||
// Some models strip the tool prefix, so auto-add it for known code_execution tools
|
||||
@@ -1248,16 +1239,24 @@ impl ExtensionManager {
|
||||
let client = client.clone();
|
||||
let notifications_receiver = client.lock().await.subscribe().await;
|
||||
let session_id = session_id.to_string();
|
||||
let working_dir_str = working_dir.map(|p| p.to_string_lossy().to_string());
|
||||
|
||||
let fut = async move {
|
||||
tracing::debug!(
|
||||
"dispatch_tool_call fut: calling client.call_tool tool={} session_id={}",
|
||||
"dispatch_tool_call fut: calling client.call_tool tool={} session_id={} working_dir={:?}",
|
||||
tool_name,
|
||||
session_id
|
||||
session_id,
|
||||
working_dir_str
|
||||
);
|
||||
let client_guard = client.lock().await;
|
||||
client_guard
|
||||
.call_tool(&session_id, &tool_name, arguments, cancellation_token)
|
||||
.call_tool(
|
||||
&session_id,
|
||||
&tool_name,
|
||||
arguments,
|
||||
working_dir_str.as_deref(),
|
||||
cancellation_token,
|
||||
)
|
||||
.await
|
||||
.map_err(|e| match e {
|
||||
ServiceError::McpError(error_data) => error_data,
|
||||
@@ -1591,6 +1590,7 @@ mod tests {
|
||||
_session_id: &str,
|
||||
name: &str,
|
||||
_arguments: Option<JsonObject>,
|
||||
_working_dir: Option<&str>,
|
||||
_cancellation_token: CancellationToken,
|
||||
) -> Result<CallToolResult, Error> {
|
||||
match name {
|
||||
@@ -1727,7 +1727,12 @@ mod tests {
|
||||
};
|
||||
|
||||
let result = extension_manager
|
||||
.dispatch_tool_call("test-session-id", tool_call, CancellationToken::default())
|
||||
.dispatch_tool_call(
|
||||
"test-session-id",
|
||||
tool_call,
|
||||
None,
|
||||
CancellationToken::default(),
|
||||
)
|
||||
.await;
|
||||
assert!(result.is_ok());
|
||||
|
||||
@@ -1739,7 +1744,12 @@ mod tests {
|
||||
};
|
||||
|
||||
let result = extension_manager
|
||||
.dispatch_tool_call("test-session-id", tool_call, CancellationToken::default())
|
||||
.dispatch_tool_call(
|
||||
"test-session-id",
|
||||
tool_call,
|
||||
None,
|
||||
CancellationToken::default(),
|
||||
)
|
||||
.await;
|
||||
assert!(result.is_ok());
|
||||
|
||||
@@ -1752,7 +1762,12 @@ mod tests {
|
||||
};
|
||||
|
||||
let result = extension_manager
|
||||
.dispatch_tool_call("test-session-id", tool_call, CancellationToken::default())
|
||||
.dispatch_tool_call(
|
||||
"test-session-id",
|
||||
tool_call,
|
||||
None,
|
||||
CancellationToken::default(),
|
||||
)
|
||||
.await;
|
||||
assert!(result.is_ok());
|
||||
|
||||
@@ -1765,7 +1780,12 @@ mod tests {
|
||||
};
|
||||
|
||||
let result = extension_manager
|
||||
.dispatch_tool_call("test-session-id", tool_call, CancellationToken::default())
|
||||
.dispatch_tool_call(
|
||||
"test-session-id",
|
||||
tool_call,
|
||||
None,
|
||||
CancellationToken::default(),
|
||||
)
|
||||
.await;
|
||||
assert!(result.is_ok());
|
||||
|
||||
@@ -1777,7 +1797,12 @@ mod tests {
|
||||
};
|
||||
|
||||
let result = extension_manager
|
||||
.dispatch_tool_call("test-session-id", tool_call, CancellationToken::default())
|
||||
.dispatch_tool_call(
|
||||
"test-session-id",
|
||||
tool_call,
|
||||
None,
|
||||
CancellationToken::default(),
|
||||
)
|
||||
.await;
|
||||
assert!(result.is_ok());
|
||||
|
||||
@@ -1793,6 +1818,7 @@ mod tests {
|
||||
.dispatch_tool_call(
|
||||
"test-session-id",
|
||||
invalid_tool_call,
|
||||
None,
|
||||
CancellationToken::default(),
|
||||
)
|
||||
.await
|
||||
@@ -1820,6 +1846,7 @@ mod tests {
|
||||
.dispatch_tool_call(
|
||||
"test-session-id",
|
||||
invalid_tool_call,
|
||||
None,
|
||||
CancellationToken::default(),
|
||||
)
|
||||
.await;
|
||||
@@ -1922,6 +1949,7 @@ mod tests {
|
||||
.dispatch_tool_call(
|
||||
"test-session-id",
|
||||
unavailable_tool_call,
|
||||
None,
|
||||
CancellationToken::default(),
|
||||
)
|
||||
.await;
|
||||
@@ -1947,6 +1975,7 @@ mod tests {
|
||||
.dispatch_tool_call(
|
||||
"test-session-id",
|
||||
available_tool_call,
|
||||
None,
|
||||
CancellationToken::default(),
|
||||
)
|
||||
.await;
|
||||
|
||||
@@ -435,6 +435,7 @@ impl McpClientTrait for ExtensionManagerClient {
|
||||
session_id: &str,
|
||||
name: &str,
|
||||
arguments: Option<JsonObject>,
|
||||
_working_dir: Option<&str>,
|
||||
_cancellation_token: CancellationToken,
|
||||
) -> Result<CallToolResult, Error> {
|
||||
let result = match name {
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
use crate::action_required_manager::ActionRequiredManager;
|
||||
use crate::agents::types::SharedProvider;
|
||||
use crate::session_context::SESSION_ID_HEADER;
|
||||
use crate::session_context::{SESSION_ID_HEADER, WORKING_DIR_HEADER};
|
||||
use rmcp::model::{
|
||||
Content, CreateElicitationRequestParams, CreateElicitationResult, ElicitationAction, ErrorCode,
|
||||
Extensions, JsonObject, Meta,
|
||||
@@ -51,6 +51,7 @@ pub trait McpClientTrait: Send + Sync {
|
||||
session_id: &str,
|
||||
name: &str,
|
||||
arguments: Option<JsonObject>,
|
||||
working_dir: Option<&str>,
|
||||
cancel_token: CancellationToken,
|
||||
) -> Result<CallToolResult, Error>;
|
||||
|
||||
@@ -105,7 +106,7 @@ pub trait McpClientTrait: Send + Sync {
|
||||
pub struct GooseClient {
|
||||
notification_handlers: Arc<Mutex<Vec<Sender<ServerNotification>>>>,
|
||||
provider: SharedProvider,
|
||||
// Single-slot because calls are serialized per MCP client; see send_request_with_session.
|
||||
// Single-slot because calls are serialized per MCP client.
|
||||
current_session_id: Arc<Mutex<Option<String>>>,
|
||||
}
|
||||
|
||||
@@ -386,13 +387,14 @@ impl McpClient {
|
||||
self.docker_container.as_deref()
|
||||
}
|
||||
|
||||
async fn send_request_with_session(
|
||||
async fn send_request_with_context(
|
||||
&self,
|
||||
session_id: &str,
|
||||
working_dir: Option<&str>,
|
||||
request: ClientRequest,
|
||||
cancel_token: CancellationToken,
|
||||
) -> Result<ServerResult, Error> {
|
||||
let request = inject_session_id_into_request(request, Some(session_id));
|
||||
let request = inject_session_context_into_request(request, Some(session_id), working_dir);
|
||||
// ExtensionManager serializes calls per MCP connection, so one current_session_id slot
|
||||
// is sufficient for mapping callbacks to the active request session.
|
||||
let handle = {
|
||||
@@ -473,8 +475,9 @@ impl McpClientTrait for McpClient {
|
||||
cancel_token: CancellationToken,
|
||||
) -> Result<ListResourcesResult, Error> {
|
||||
let res = self
|
||||
.send_request_with_session(
|
||||
.send_request_with_context(
|
||||
session_id,
|
||||
None,
|
||||
ClientRequest::ListResourcesRequest(ListResourcesRequest {
|
||||
params: Some(PaginatedRequestParams { meta: None, cursor }),
|
||||
method: Default::default(),
|
||||
@@ -497,8 +500,9 @@ impl McpClientTrait for McpClient {
|
||||
cancel_token: CancellationToken,
|
||||
) -> Result<ReadResourceResult, Error> {
|
||||
let res = self
|
||||
.send_request_with_session(
|
||||
.send_request_with_context(
|
||||
session_id,
|
||||
None,
|
||||
ClientRequest::ReadResourceRequest(ReadResourceRequest {
|
||||
params: ReadResourceRequestParams {
|
||||
meta: None,
|
||||
@@ -524,8 +528,9 @@ impl McpClientTrait for McpClient {
|
||||
cancel_token: CancellationToken,
|
||||
) -> Result<ListToolsResult, Error> {
|
||||
let res = self
|
||||
.send_request_with_session(
|
||||
.send_request_with_context(
|
||||
session_id,
|
||||
None,
|
||||
ClientRequest::ListToolsRequest(ListToolsRequest {
|
||||
params: Some(PaginatedRequestParams { meta: None, cursor }),
|
||||
method: Default::default(),
|
||||
@@ -546,6 +551,7 @@ impl McpClientTrait for McpClient {
|
||||
session_id: &str,
|
||||
name: &str,
|
||||
arguments: Option<JsonObject>,
|
||||
working_dir: Option<&str>,
|
||||
cancel_token: CancellationToken,
|
||||
) -> Result<CallToolResult, Error> {
|
||||
let request = ClientRequest::CallToolRequest(CallToolRequest {
|
||||
@@ -560,7 +566,7 @@ impl McpClientTrait for McpClient {
|
||||
});
|
||||
|
||||
let result = self
|
||||
.send_request_with_session(session_id, request, cancel_token)
|
||||
.send_request_with_context(session_id, working_dir, request, cancel_token)
|
||||
.await;
|
||||
|
||||
match result? {
|
||||
@@ -576,8 +582,9 @@ impl McpClientTrait for McpClient {
|
||||
cancel_token: CancellationToken,
|
||||
) -> Result<ListPromptsResult, Error> {
|
||||
let res = self
|
||||
.send_request_with_session(
|
||||
.send_request_with_context(
|
||||
session_id,
|
||||
None,
|
||||
ClientRequest::ListPromptsRequest(ListPromptsRequest {
|
||||
params: Some(PaginatedRequestParams { meta: None, cursor }),
|
||||
method: Default::default(),
|
||||
@@ -605,8 +612,9 @@ impl McpClientTrait for McpClient {
|
||||
_ => None,
|
||||
};
|
||||
let res = self
|
||||
.send_request_with_session(
|
||||
.send_request_with_context(
|
||||
session_id,
|
||||
None,
|
||||
ClientRequest::GetPromptRequest(GetPromptRequest {
|
||||
params: GetPromptRequestParams {
|
||||
meta: None,
|
||||
@@ -633,20 +641,24 @@ impl McpClientTrait for McpClient {
|
||||
}
|
||||
}
|
||||
|
||||
/// Injects the given session_id into Extensions._meta.
|
||||
/// None (or empty) removes any existing session id.
|
||||
fn inject_session_id_into_extensions(
|
||||
/// Injects the given session_id and working_dir into Extensions._meta.
|
||||
/// None (or empty) removes any existing values.
|
||||
fn inject_session_context_into_extensions(
|
||||
mut extensions: Extensions,
|
||||
session_id: Option<&str>,
|
||||
working_dir: Option<&str>,
|
||||
) -> Extensions {
|
||||
let session_id = session_id.filter(|id| !id.is_empty());
|
||||
let working_dir = working_dir.filter(|dir| !dir.is_empty());
|
||||
let mut meta_map = extensions
|
||||
.get::<Meta>()
|
||||
.map(|meta| meta.0.clone())
|
||||
.unwrap_or_default();
|
||||
|
||||
// JsonObject is case-sensitive, so we use retain for case-insensitive removal
|
||||
meta_map.retain(|k, _| !k.eq_ignore_ascii_case(SESSION_ID_HEADER));
|
||||
meta_map.retain(|k, _| {
|
||||
!k.eq_ignore_ascii_case(SESSION_ID_HEADER) && !k.eq_ignore_ascii_case(WORKING_DIR_HEADER)
|
||||
});
|
||||
|
||||
if let Some(session_id) = session_id {
|
||||
meta_map.insert(
|
||||
@@ -655,37 +667,51 @@ fn inject_session_id_into_extensions(
|
||||
);
|
||||
}
|
||||
|
||||
if let Some(working_dir) = working_dir {
|
||||
meta_map.insert(
|
||||
WORKING_DIR_HEADER.to_string(),
|
||||
Value::String(working_dir.to_string()),
|
||||
);
|
||||
}
|
||||
|
||||
extensions.insert(Meta(meta_map));
|
||||
extensions
|
||||
}
|
||||
|
||||
fn inject_session_id_into_request(
|
||||
fn inject_session_context_into_request(
|
||||
request: ClientRequest,
|
||||
session_id: Option<&str>,
|
||||
working_dir: Option<&str>,
|
||||
) -> ClientRequest {
|
||||
match request {
|
||||
ClientRequest::ListResourcesRequest(mut req) => {
|
||||
req.extensions = inject_session_id_into_extensions(req.extensions, session_id);
|
||||
req.extensions =
|
||||
inject_session_context_into_extensions(req.extensions, session_id, working_dir);
|
||||
ClientRequest::ListResourcesRequest(req)
|
||||
}
|
||||
ClientRequest::ReadResourceRequest(mut req) => {
|
||||
req.extensions = inject_session_id_into_extensions(req.extensions, session_id);
|
||||
req.extensions =
|
||||
inject_session_context_into_extensions(req.extensions, session_id, working_dir);
|
||||
ClientRequest::ReadResourceRequest(req)
|
||||
}
|
||||
ClientRequest::ListToolsRequest(mut req) => {
|
||||
req.extensions = inject_session_id_into_extensions(req.extensions, session_id);
|
||||
req.extensions =
|
||||
inject_session_context_into_extensions(req.extensions, session_id, working_dir);
|
||||
ClientRequest::ListToolsRequest(req)
|
||||
}
|
||||
ClientRequest::CallToolRequest(mut req) => {
|
||||
req.extensions = inject_session_id_into_extensions(req.extensions, session_id);
|
||||
req.extensions =
|
||||
inject_session_context_into_extensions(req.extensions, session_id, working_dir);
|
||||
ClientRequest::CallToolRequest(req)
|
||||
}
|
||||
ClientRequest::ListPromptsRequest(mut req) => {
|
||||
req.extensions = inject_session_id_into_extensions(req.extensions, session_id);
|
||||
req.extensions =
|
||||
inject_session_context_into_extensions(req.extensions, session_id, working_dir);
|
||||
ClientRequest::ListPromptsRequest(req)
|
||||
}
|
||||
ClientRequest::GetPromptRequest(mut req) => {
|
||||
req.extensions = inject_session_id_into_extensions(req.extensions, session_id);
|
||||
req.extensions =
|
||||
inject_session_context_into_extensions(req.extensions, session_id, working_dir);
|
||||
ClientRequest::GetPromptRequest(req)
|
||||
}
|
||||
other => other,
|
||||
@@ -814,7 +840,8 @@ mod tests {
|
||||
*slot = Some(session_id.to_string());
|
||||
}
|
||||
|
||||
let extensions = inject_session_id_into_extensions(Extensions::new(), ext_session);
|
||||
let extensions =
|
||||
inject_session_context_into_extensions(Extensions::new(), ext_session, None);
|
||||
|
||||
let resolved = client.resolve_session_id(&extensions).await;
|
||||
|
||||
@@ -830,8 +857,6 @@ mod tests {
|
||||
#[test_case(list_prompts_request; "list_prompts")]
|
||||
#[test_case(get_prompt_request; "get_prompt")]
|
||||
fn test_request_injects_session(request_builder: fn(Extensions) -> ClientRequest) {
|
||||
use serde_json::json;
|
||||
|
||||
let session_id = "test-session-id";
|
||||
let mut extensions = Extensions::new();
|
||||
extensions.insert(
|
||||
@@ -843,7 +868,7 @@ mod tests {
|
||||
);
|
||||
|
||||
let request = request_builder(extensions);
|
||||
let request = inject_session_id_into_request(request, Some(session_id));
|
||||
let request = inject_session_context_into_request(request, Some(session_id), None);
|
||||
let extensions = request_extensions(&request).expect("request should have extensions");
|
||||
let meta = extensions
|
||||
.get::<Meta>()
|
||||
@@ -861,10 +886,9 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn test_session_id_in_mcp_meta() {
|
||||
use serde_json::json;
|
||||
|
||||
let session_id = "test-session-789";
|
||||
let extensions = inject_session_id_into_extensions(Default::default(), Some(session_id));
|
||||
let extensions =
|
||||
inject_session_context_into_extensions(Default::default(), Some(session_id), None);
|
||||
let mcp_meta = extensions.get::<Meta>().unwrap();
|
||||
|
||||
assert_eq!(
|
||||
@@ -904,7 +928,7 @@ mod tests {
|
||||
expected_meta: serde_json::Value,
|
||||
) {
|
||||
use rmcp::model::Extensions;
|
||||
use serde_json::{from_value, json};
|
||||
use serde_json::from_value;
|
||||
|
||||
let mut extensions = Extensions::new();
|
||||
extensions.insert(
|
||||
@@ -916,7 +940,7 @@ mod tests {
|
||||
.unwrap(),
|
||||
);
|
||||
|
||||
let extensions = inject_session_id_into_extensions(extensions, session_id);
|
||||
let extensions = inject_session_context_into_extensions(extensions, session_id, None);
|
||||
let mcp_meta = extensions.get::<Meta>().unwrap();
|
||||
|
||||
assert_eq!(&mcp_meta.0, expected_meta.as_object().unwrap());
|
||||
|
||||
@@ -306,6 +306,7 @@ impl McpClientTrait for SkillsClient {
|
||||
_session_id: &str,
|
||||
name: &str,
|
||||
arguments: Option<JsonObject>,
|
||||
_working_dir: Option<&str>,
|
||||
_cancellation_token: CancellationToken,
|
||||
) -> Result<CallToolResult, Error> {
|
||||
let content = match name {
|
||||
|
||||
@@ -174,6 +174,7 @@ impl McpClientTrait for TodoClient {
|
||||
session_id: &str,
|
||||
name: &str,
|
||||
arguments: Option<JsonObject>,
|
||||
_working_dir: Option<&str>,
|
||||
_cancellation_token: CancellationToken,
|
||||
) -> Result<CallToolResult, Error> {
|
||||
let content = match name {
|
||||
|
||||
Reference in New Issue
Block a user