diff --git a/codex-rs/core/src/session/tests.rs b/codex-rs/core/src/session/tests.rs index 13fd3768b..6995ef8dc 100644 --- a/codex-rs/core/src/session/tests.rs +++ b/codex-rs/core/src/session/tests.rs @@ -66,9 +66,10 @@ use crate::tasks::execute_user_shell_command; use crate::tools::ToolRouter; use crate::tools::context::ToolInvocation; use crate::tools::context::ToolPayload; -use crate::tools::handlers::GoalHandler; +use crate::tools::handlers::CreateGoalHandler; +use crate::tools::handlers::ExecCommandHandler; use crate::tools::handlers::ShellHandler; -use crate::tools::handlers::UnifiedExecHandler; +use crate::tools::handlers::UpdateGoalHandler; use crate::tools::registry::ToolHandler; use crate::tools::router::ToolCallSource; use crate::turn_diff_tracker::TurnDiffTracker; @@ -8247,7 +8248,7 @@ async fn sample_rollout( async fn create_goal_tool_rejects_existing_goal() { let (session, turn_context, _rx, _codex_home) = make_goal_session_and_context_with_rx().await; let tracker = Arc::new(tokio::sync::Mutex::new(TurnDiffTracker::new())); - let handler = GoalHandler; + let handler = CreateGoalHandler; handler .handle(ToolInvocation { @@ -8309,9 +8310,10 @@ async fn create_goal_tool_rejects_existing_goal() { async fn update_goal_tool_rejects_pausing_goal() { let (session, turn_context, _rx, _codex_home) = make_goal_session_and_context_with_rx().await; let tracker = Arc::new(tokio::sync::Mutex::new(TurnDiffTracker::new())); - let handler = GoalHandler; + let create_handler = CreateGoalHandler; + let update_handler = UpdateGoalHandler; - handler + create_handler .handle(ToolInvocation { session: Arc::clone(&session), turn: Arc::clone(&turn_context), @@ -8331,7 +8333,7 @@ async fn update_goal_tool_rejects_pausing_goal() { .await .expect("initial create_goal should succeed"); - let response = handler + let response = update_handler .handle(ToolInvocation { session: Arc::clone(&session), turn: Arc::clone(&turn_context), @@ -8369,9 +8371,10 @@ async fn update_goal_tool_rejects_pausing_goal() { async fn update_goal_tool_marks_goal_complete() { let (session, turn_context, _rx, _codex_home) = make_goal_session_and_context_with_rx().await; let tracker = Arc::new(tokio::sync::Mutex::new(TurnDiffTracker::new())); - let handler = GoalHandler; + let create_handler = CreateGoalHandler; + let update_handler = UpdateGoalHandler; - handler + create_handler .handle(ToolInvocation { session: Arc::clone(&session), turn: Arc::clone(&turn_context), @@ -8391,7 +8394,7 @@ async fn update_goal_tool_marks_goal_complete() { .await .expect("initial create_goal should succeed"); - handler + update_handler .handle(ToolInvocation { session: Arc::clone(&session), turn: Arc::clone(&turn_context), @@ -8548,7 +8551,7 @@ async fn unified_exec_rejects_escalated_permissions_when_policy_not_on_request() let turn_context = Arc::new(turn_context_raw); let tracker = Arc::new(tokio::sync::Mutex::new(TurnDiffTracker::new())); - let handler = UnifiedExecHandler; + let handler = ExecCommandHandler; let resp = handler .handle(ToolInvocation { session: Arc::clone(&session), diff --git a/codex-rs/core/src/session/tests/guardian_tests.rs b/codex-rs/core/src/session/tests/guardian_tests.rs index ad7dbb105..4c61dad18 100644 --- a/codex-rs/core/src/session/tests/guardian_tests.rs +++ b/codex-rs/core/src/session/tests/guardian_tests.rs @@ -498,7 +498,7 @@ async fn guardian_allows_unified_exec_additional_permissions_requests_past_polic let turn_context = Arc::new(turn_context_raw); let tracker = Arc::new(tokio::sync::Mutex::new(TurnDiffTracker::new())); - let handler = UnifiedExecHandler; + let handler = ExecCommandHandler; let resp = handler .handle(ToolInvocation { session: Arc::clone(&session), diff --git a/codex-rs/core/src/tools/code_mode/execute_handler.rs b/codex-rs/core/src/tools/code_mode/execute_handler.rs index 6b99e09b5..42841d218 100644 --- a/codex-rs/core/src/tools/code_mode/execute_handler.rs +++ b/codex-rs/core/src/tools/code_mode/execute_handler.rs @@ -4,6 +4,7 @@ use crate::tools::context::ToolInvocation; use crate::tools::context::ToolPayload; use crate::tools::registry::ToolHandler; use crate::tools::registry::ToolKind; +use codex_tools::ToolName; use super::ExecContext; use super::PUBLIC_TOOL_NAME; @@ -78,6 +79,10 @@ impl CodeModeExecuteHandler { impl ToolHandler for CodeModeExecuteHandler { type Output = FunctionToolOutput; + fn tool_name(&self) -> ToolName { + ToolName::plain(PUBLIC_TOOL_NAME) + } + fn kind(&self) -> ToolKind { ToolKind::Function } diff --git a/codex-rs/core/src/tools/code_mode/wait_handler.rs b/codex-rs/core/src/tools/code_mode/wait_handler.rs index 70fa51251..8024c9586 100644 --- a/codex-rs/core/src/tools/code_mode/wait_handler.rs +++ b/codex-rs/core/src/tools/code_mode/wait_handler.rs @@ -6,6 +6,7 @@ use crate::tools::context::ToolInvocation; use crate::tools::context::ToolPayload; use crate::tools::registry::ToolHandler; use crate::tools::registry::ToolKind; +use codex_tools::ToolName; use super::DEFAULT_WAIT_YIELD_TIME_MS; use super::ExecContext; @@ -41,6 +42,10 @@ where impl ToolHandler for CodeModeWaitHandler { type Output = FunctionToolOutput; + fn tool_name(&self) -> ToolName { + ToolName::plain(WAIT_TOOL_NAME) + } + fn kind(&self) -> ToolKind { ToolKind::Function } diff --git a/codex-rs/core/src/tools/handlers/agent_jobs.rs b/codex-rs/core/src/tools/handlers/agent_jobs.rs index b3d9b481d..93380d5e4 100644 --- a/codex-rs/core/src/tools/handlers/agent_jobs.rs +++ b/codex-rs/core/src/tools/handlers/agent_jobs.rs @@ -19,6 +19,7 @@ use codex_protocol::protocol::AgentStatus; use codex_protocol::protocol::SessionSource; use codex_protocol::protocol::SubAgentSource; use codex_protocol::user_input::UserInput; +use codex_tools::ToolName; use codex_utils_absolute_path::AbsolutePathBuf; use futures::StreamExt; use futures::stream::FuturesUnordered; @@ -35,7 +36,8 @@ use tokio::time::Instant; use tokio::time::timeout; use uuid::Uuid; -pub struct BatchJobHandler; +pub struct SpawnAgentsOnCsvHandler; +pub struct ReportAgentJobResultHandler; const DEFAULT_AGENT_JOB_CONCURRENCY: usize = 16; const MAX_AGENT_JOB_CONCURRENCY: usize = 64; @@ -99,9 +101,13 @@ struct ActiveJobItem { status_rx: Option>, } -impl ToolHandler for BatchJobHandler { +impl ToolHandler for SpawnAgentsOnCsvHandler { type Output = FunctionToolOutput; + fn tool_name(&self) -> ToolName { + ToolName::plain("spawn_agents_on_csv") + } + fn kind(&self) -> ToolKind { ToolKind::Function } @@ -114,7 +120,6 @@ impl ToolHandler for BatchJobHandler { let ToolInvocation { session, turn, - tool_name, payload, .. } = invocation; @@ -128,13 +133,40 @@ impl ToolHandler for BatchJobHandler { } }; - match tool_name.name.as_str() { - "spawn_agents_on_csv" => spawn_agents_on_csv::handle(session, turn, arguments).await, - "report_agent_job_result" => report_agent_job_result::handle(session, arguments).await, - other => Err(FunctionCallError::RespondToModel(format!( - "unsupported agent job tool {other}" - ))), - } + spawn_agents_on_csv::handle(session, turn, arguments).await + } +} + +impl ToolHandler for ReportAgentJobResultHandler { + type Output = FunctionToolOutput; + + fn tool_name(&self) -> ToolName { + ToolName::plain("report_agent_job_result") + } + + fn kind(&self) -> ToolKind { + ToolKind::Function + } + + fn matches_kind(&self, payload: &ToolPayload) -> bool { + matches!(payload, ToolPayload::Function { .. }) + } + + async fn handle(&self, invocation: ToolInvocation) -> Result { + let ToolInvocation { + session, payload, .. + } = invocation; + + let arguments = match payload { + ToolPayload::Function { arguments } => arguments, + _ => { + return Err(FunctionCallError::RespondToModel( + "report_agent_job_result handler received unsupported payload".to_string(), + )); + } + }; + + report_agent_job_result::handle(session, arguments).await } } diff --git a/codex-rs/core/src/tools/handlers/apply_patch.rs b/codex-rs/core/src/tools/handlers/apply_patch.rs index e2e020a96..9766bfb57 100644 --- a/codex-rs/core/src/tools/handlers/apply_patch.rs +++ b/codex-rs/core/src/tools/handlers/apply_patch.rs @@ -47,6 +47,7 @@ use codex_sandboxing::policy_transforms::effective_file_system_sandbox_policy; use codex_sandboxing::policy_transforms::merge_permission_profiles; use codex_sandboxing::policy_transforms::normalize_additional_permissions; use codex_tools::ApplyPatchToolArgs; +use codex_tools::ToolName; use codex_utils_absolute_path::AbsolutePathBuf; const APPLY_PATCH_ARGUMENT_DIFF_BUFFER_INTERVAL: Duration = Duration::from_millis(500); @@ -292,6 +293,10 @@ async fn effective_patch_permissions( impl ToolHandler for ApplyPatchHandler { type Output = ApplyPatchToolOutput; + fn tool_name(&self) -> ToolName { + ToolName::plain("apply_patch") + } + fn kind(&self) -> ToolKind { ToolKind::Function } diff --git a/codex-rs/core/src/tools/handlers/dynamic.rs b/codex-rs/core/src/tools/handlers/dynamic.rs index eab1f0f80..549edd514 100644 --- a/codex-rs/core/src/tools/handlers/dynamic.rs +++ b/codex-rs/core/src/tools/handlers/dynamic.rs @@ -19,11 +19,23 @@ use std::time::Instant; use tokio::sync::oneshot; use tracing::warn; -pub struct DynamicToolHandler; +pub struct DynamicToolHandler { + tool_name: ToolName, +} + +impl DynamicToolHandler { + pub fn new(tool_name: ToolName) -> Self { + Self { tool_name } + } +} impl ToolHandler for DynamicToolHandler { type Output = FunctionToolOutput; + fn tool_name(&self) -> ToolName { + self.tool_name.clone() + } + fn kind(&self) -> ToolKind { ToolKind::Function } @@ -37,7 +49,6 @@ impl ToolHandler for DynamicToolHandler { session, turn, call_id, - tool_name, payload, .. } = invocation; @@ -52,13 +63,19 @@ impl ToolHandler for DynamicToolHandler { }; let args: Value = parse_arguments(&arguments)?; - let response = request_dynamic_tool(&session, turn.as_ref(), call_id, tool_name, args) - .await - .ok_or_else(|| { - FunctionCallError::RespondToModel( - "dynamic tool call was cancelled before receiving a response".to_string(), - ) - })?; + let response = request_dynamic_tool( + &session, + turn.as_ref(), + call_id, + self.tool_name.clone(), + args, + ) + .await + .ok_or_else(|| { + FunctionCallError::RespondToModel( + "dynamic tool call was cancelled before receiving a response".to_string(), + ) + })?; let DynamicToolResponse { content_items, diff --git a/codex-rs/core/src/tools/handlers/goal.rs b/codex-rs/core/src/tools/handlers/goal.rs index 74391d57b..6a7b304ce 100644 --- a/codex-rs/core/src/tools/handlers/goal.rs +++ b/codex-rs/core/src/tools/handlers/goal.rs @@ -8,8 +8,6 @@ use crate::function_tool::FunctionCallError; use crate::goals::CreateGoalRequest; use crate::goals::GoalRuntimeEvent; use crate::goals::SetGoalRequest; -use crate::session::session::Session; -use crate::session::turn_context::TurnContext; use crate::tools::context::FunctionToolOutput; use crate::tools::context::ToolInvocation; use crate::tools::context::ToolPayload; @@ -20,13 +18,15 @@ use codex_protocol::protocol::ThreadGoal; use codex_protocol::protocol::ThreadGoalStatus; use codex_tools::CREATE_GOAL_TOOL_NAME; use codex_tools::GET_GOAL_TOOL_NAME; +use codex_tools::ToolName; use codex_tools::UPDATE_GOAL_TOOL_NAME; use serde::Deserialize; use serde::Serialize; use std::fmt::Write as _; -use std::sync::Arc; -pub struct GoalHandler; +pub struct GetGoalHandler; +pub struct CreateGoalHandler; +pub struct UpdateGoalHandler; #[derive(Debug, Deserialize)] #[serde(rename_all = "snake_case")] @@ -76,9 +76,44 @@ impl GoalToolResponse { } } -impl ToolHandler for GoalHandler { +impl ToolHandler for GetGoalHandler { type Output = FunctionToolOutput; + fn tool_name(&self) -> ToolName { + ToolName::plain(GET_GOAL_TOOL_NAME) + } + + fn kind(&self) -> ToolKind { + ToolKind::Function + } + + async fn handle(&self, invocation: ToolInvocation) -> Result { + let ToolInvocation { + session, payload, .. + } = invocation; + + match payload { + ToolPayload::Function { .. } => { + let goal = session + .get_thread_goal() + .await + .map_err(|err| FunctionCallError::RespondToModel(format_goal_error(err)))?; + goal_response(goal, CompletionBudgetReport::Omit) + } + _ => Err(FunctionCallError::RespondToModel( + "get_goal handler received unsupported payload".to_string(), + )), + } + } +} + +impl ToolHandler for CreateGoalHandler { + type Output = FunctionToolOutput; + + fn tool_name(&self) -> ToolName { + ToolName::plain(CREATE_GOAL_TOOL_NAME) + } + fn kind(&self) -> ToolKind { ToolKind::Function } @@ -88,7 +123,6 @@ impl ToolHandler for GoalHandler { session, turn, payload, - tool_name, .. } = invocation; @@ -101,88 +135,89 @@ impl ToolHandler for GoalHandler { } }; - match tool_name.name.as_str() { - GET_GOAL_TOOL_NAME => handle_get_goal(session.as_ref()).await, - CREATE_GOAL_TOOL_NAME => { - handle_create_goal(session.as_ref(), turn.as_ref(), &arguments).await + let args: CreateGoalArgs = parse_arguments(&arguments)?; + let goal = session + .create_thread_goal( + turn.as_ref(), + CreateGoalRequest { + objective: args.objective, + token_budget: args.token_budget, + }, + ) + .await + .map_err(|err| { + if err + .chain() + .any(|cause| cause.to_string().contains("already has a goal")) + { + FunctionCallError::RespondToModel( + "cannot create a new goal because this thread already has a goal; use update_goal only when the existing goal is complete" + .to_string(), + ) + } else { + FunctionCallError::RespondToModel(format_goal_error(err)) + } + })?; + goal_response(Some(goal), CompletionBudgetReport::Omit) + } +} + +impl ToolHandler for UpdateGoalHandler { + type Output = FunctionToolOutput; + + fn tool_name(&self) -> ToolName { + ToolName::plain(UPDATE_GOAL_TOOL_NAME) + } + + fn kind(&self) -> ToolKind { + ToolKind::Function + } + + async fn handle(&self, invocation: ToolInvocation) -> Result { + let ToolInvocation { + session, + turn, + payload, + .. + } = invocation; + + let arguments = match payload { + ToolPayload::Function { arguments } => arguments, + _ => { + return Err(FunctionCallError::RespondToModel( + "update_goal handler received unsupported payload".to_string(), + )); } - UPDATE_GOAL_TOOL_NAME => handle_update_goal(&session, turn.as_ref(), &arguments).await, - other => Err(FunctionCallError::Fatal(format!( - "goal handler received unsupported tool: {other}" - ))), + }; + + let args: UpdateGoalArgs = parse_arguments(&arguments)?; + if args.status != ThreadGoalStatus::Complete { + return Err(FunctionCallError::RespondToModel( + "update_goal can only mark the existing goal complete; pause, resume, and budget-limited status changes are controlled by the user or system" + .to_string(), + )); } + session + .goal_runtime_apply(GoalRuntimeEvent::ToolCompletedGoal { + turn_context: turn.as_ref(), + }) + .await + .map_err(|err| FunctionCallError::RespondToModel(format_goal_error(err)))?; + let goal = session + .set_thread_goal( + turn.as_ref(), + SetGoalRequest { + objective: None, + status: Some(ThreadGoalStatus::Complete), + token_budget: None, + }, + ) + .await + .map_err(|err| FunctionCallError::RespondToModel(format_goal_error(err)))?; + goal_response(Some(goal), CompletionBudgetReport::Include) } } -async fn handle_get_goal(session: &Session) -> Result { - let goal = session - .get_thread_goal() - .await - .map_err(|err| FunctionCallError::RespondToModel(format_goal_error(err)))?; - goal_response(goal, CompletionBudgetReport::Omit) -} - -async fn handle_create_goal( - session: &Session, - turn_context: &TurnContext, - arguments: &str, -) -> Result { - let args: CreateGoalArgs = parse_arguments(arguments)?; - let goal = session - .create_thread_goal( - turn_context, - CreateGoalRequest { - objective: args.objective, - token_budget: args.token_budget, - }, - ) - .await - .map_err(|err| { - if err - .chain() - .any(|cause| cause.to_string().contains("already has a goal")) - { - FunctionCallError::RespondToModel( - "cannot create a new goal because this thread already has a goal; use update_goal only when the existing goal is complete" - .to_string(), - ) - } else { - FunctionCallError::RespondToModel(format_goal_error(err)) - } - })?; - goal_response(Some(goal), CompletionBudgetReport::Omit) -} - -async fn handle_update_goal( - session: &Arc, - turn_context: &TurnContext, - arguments: &str, -) -> Result { - let args: UpdateGoalArgs = parse_arguments(arguments)?; - if args.status != ThreadGoalStatus::Complete { - return Err(FunctionCallError::RespondToModel( - "update_goal can only mark the existing goal complete; pause, resume, and budget-limited status changes are controlled by the user or system" - .to_string(), - )); - } - session - .goal_runtime_apply(GoalRuntimeEvent::ToolCompletedGoal { turn_context }) - .await - .map_err(|err| FunctionCallError::RespondToModel(format_goal_error(err)))?; - let goal = session - .set_thread_goal( - turn_context, - SetGoalRequest { - objective: None, - status: Some(ThreadGoalStatus::Complete), - token_budget: None, - }, - ) - .await - .map_err(|err| FunctionCallError::RespondToModel(format_goal_error(err)))?; - goal_response(Some(goal), CompletionBudgetReport::Include) -} - fn format_goal_error(err: anyhow::Error) -> String { let mut message = err.to_string(); for cause in err.chain().skip(1) { diff --git a/codex-rs/core/src/tools/handlers/mcp.rs b/codex-rs/core/src/tools/handlers/mcp.rs index 568e45615..4dfcb44b1 100644 --- a/codex-rs/core/src/tools/handlers/mcp.rs +++ b/codex-rs/core/src/tools/handlers/mcp.rs @@ -13,12 +13,26 @@ use crate::tools::registry::PostToolUsePayload; use crate::tools::registry::PreToolUsePayload; use crate::tools::registry::ToolHandler; use crate::tools::registry::ToolKind; +use codex_tools::ToolName; use serde_json::Value; -pub struct McpHandler; +pub struct McpHandler { + tool_name: ToolName, +} + +impl McpHandler { + pub fn new(tool_name: ToolName) -> Self { + Self { tool_name } + } +} + impl ToolHandler for McpHandler { type Output = McpToolOutput; + fn tool_name(&self) -> ToolName { + self.tool_name.clone() + } + fn kind(&self) -> ToolKind { ToolKind::Mcp } @@ -29,7 +43,7 @@ impl ToolHandler for McpHandler { }; Some(PreToolUsePayload { - tool_name: HookToolName::new(invocation.tool_name.display()), + tool_name: HookToolName::new(self.tool_name.display()), tool_input: mcp_hook_tool_input(raw_arguments), }) } @@ -46,7 +60,7 @@ impl ToolHandler for McpHandler { let tool_response = result.post_tool_use_response(&invocation.call_id, &invocation.payload)?; Some(PostToolUsePayload { - tool_name: HookToolName::new(invocation.tool_name.display()), + tool_name: HookToolName::new(self.tool_name.display()), tool_use_id: invocation.call_id.clone(), tool_input: result.tool_input.clone(), tool_response, @@ -58,7 +72,6 @@ impl ToolHandler for McpHandler { session, turn, call_id, - tool_name: model_tool_name, payload, .. } = invocation; @@ -86,7 +99,7 @@ impl ToolHandler for McpHandler { call_id.clone(), server, tool, - model_tool_name.display(), + self.tool_name.display(), arguments_str, ) .await; @@ -134,9 +147,13 @@ mod tests { .to_string(), }; let (session, turn) = make_session_and_context().await; + let handler = McpHandler::new(codex_tools::ToolName::namespaced( + "mcp__memory__", + "create_entities", + )); assert_eq!( - McpHandler.pre_tool_use_payload(&ToolInvocation { + handler.pre_tool_use_payload(&ToolInvocation { session: session.into(), turn: turn.into(), cancellation_token: tokio_util::sync::CancellationToken::new(), @@ -185,6 +202,10 @@ mod tests { truncation_policy: codex_utils_output_truncation::TruncationPolicy::Bytes(1024), }; let (session, turn) = make_session_and_context().await; + let handler = McpHandler::new(codex_tools::ToolName::namespaced( + "mcp__filesystem__", + "read_file", + )); let invocation = ToolInvocation { session: session.into(), turn: turn.into(), @@ -196,7 +217,7 @@ mod tests { payload, }; assert_eq!( - McpHandler.post_tool_use_payload(&invocation, &output), + handler.post_tool_use_payload(&invocation, &output), Some(PostToolUsePayload { tool_name: HookToolName::new("mcp__filesystem__read_file"), tool_use_id: "call-mcp-post".to_string(), diff --git a/codex-rs/core/src/tools/handlers/mcp_resource.rs b/codex-rs/core/src/tools/handlers/mcp_resource.rs index 14f8db3a4..b03de53ae 100644 --- a/codex-rs/core/src/tools/handlers/mcp_resource.rs +++ b/codex-rs/core/src/tools/handlers/mcp_resource.rs @@ -30,8 +30,11 @@ use crate::tools::context::ToolPayload; use crate::tools::registry::ToolHandler; use crate::tools::registry::ToolKind; use codex_protocol::protocol::McpInvocation; +use codex_tools::ToolName; -pub struct McpResourceHandler; +pub struct ListMcpResourcesHandler; +pub struct ListMcpResourceTemplatesHandler; +pub struct ReadMcpResourceHandler; #[derive(Debug, Deserialize, Default)] struct ListResourcesArgs { @@ -178,9 +181,281 @@ struct ReadResourcePayload { result: ReadResourceResult, } -impl ToolHandler for McpResourceHandler { +impl ToolHandler for ListMcpResourcesHandler { type Output = FunctionToolOutput; + fn tool_name(&self) -> ToolName { + ToolName::plain("list_mcp_resources") + } + + fn kind(&self) -> ToolKind { + ToolKind::Function + } + + #[expect( + clippy::await_holding_invalid_type, + reason = "MCP resource listing reads through the session-owned manager guard" + )] + async fn handle(&self, invocation: ToolInvocation) -> Result { + let ToolInvocation { + session, + turn, + call_id, + payload, + .. + } = invocation; + + let arguments = match payload { + ToolPayload::Function { arguments } => arguments, + _ => { + return Err(FunctionCallError::RespondToModel( + "list_mcp_resources handler received unsupported payload".to_string(), + )); + } + }; + + let arguments = parse_arguments(arguments.as_str())?; + let args: ListResourcesArgs = parse_args_with_default(arguments.clone())?; + let ListResourcesArgs { server, cursor } = args; + let server = normalize_optional_string(server); + let cursor = normalize_optional_string(cursor); + + let invocation = McpInvocation { + server: server.clone().unwrap_or_else(|| "codex".to_string()), + tool: "list_mcp_resources".to_string(), + arguments: arguments.clone(), + }; + + emit_tool_call_begin(&session, turn.as_ref(), &call_id, invocation.clone()).await; + let start = Instant::now(); + + let payload_result: Result = async { + if let Some(server_name) = server.clone() { + let params = cursor.clone().map(|value| PaginatedRequestParams { + meta: None, + cursor: Some(value), + }); + let result = session + .list_resources(&server_name, params) + .await + .map_err(|err| { + FunctionCallError::RespondToModel(format!("resources/list failed: {err:#}")) + })?; + Ok(ListResourcesPayload::from_single_server( + server_name, + result, + )) + } else { + if cursor.is_some() { + return Err(FunctionCallError::RespondToModel( + "cursor can only be used when a server is specified".to_string(), + )); + } + + let resources = session + .services + .mcp_connection_manager + .read() + .await + .list_all_resources() + .await; + Ok(ListResourcesPayload::from_all_servers(resources)) + } + } + .await; + + match payload_result { + Ok(payload) => match serialize_function_output(payload) { + Ok(output) => { + let content = function_call_output_content_items_to_text(&output.body) + .unwrap_or_default(); + let duration = start.elapsed(); + emit_tool_call_end( + &session, + turn.as_ref(), + &call_id, + invocation, + duration, + Ok(call_tool_result_from_content(&content, output.success)), + ) + .await; + Ok(output) + } + Err(err) => { + let duration = start.elapsed(); + let message = err.to_string(); + emit_tool_call_end( + &session, + turn.as_ref(), + &call_id, + invocation, + duration, + Err(message.clone()), + ) + .await; + Err(err) + } + }, + Err(err) => { + let duration = start.elapsed(); + let message = err.to_string(); + emit_tool_call_end( + &session, + turn.as_ref(), + &call_id, + invocation, + duration, + Err(message.clone()), + ) + .await; + Err(err) + } + } + } +} + +impl ToolHandler for ListMcpResourceTemplatesHandler { + type Output = FunctionToolOutput; + + fn tool_name(&self) -> ToolName { + ToolName::plain("list_mcp_resource_templates") + } + + fn kind(&self) -> ToolKind { + ToolKind::Function + } + + #[expect( + clippy::await_holding_invalid_type, + reason = "MCP resource template listing reads through the session-owned manager guard" + )] + async fn handle(&self, invocation: ToolInvocation) -> Result { + let ToolInvocation { + session, + turn, + call_id, + payload, + .. + } = invocation; + + let arguments = match payload { + ToolPayload::Function { arguments } => arguments, + _ => { + return Err(FunctionCallError::RespondToModel( + "list_mcp_resource_templates handler received unsupported payload".to_string(), + )); + } + }; + + let arguments = parse_arguments(arguments.as_str())?; + let args: ListResourceTemplatesArgs = parse_args_with_default(arguments.clone())?; + let ListResourceTemplatesArgs { server, cursor } = args; + let server = normalize_optional_string(server); + let cursor = normalize_optional_string(cursor); + + let invocation = McpInvocation { + server: server.clone().unwrap_or_else(|| "codex".to_string()), + tool: "list_mcp_resource_templates".to_string(), + arguments: arguments.clone(), + }; + + emit_tool_call_begin(&session, turn.as_ref(), &call_id, invocation.clone()).await; + let start = Instant::now(); + + let payload_result: Result = async { + if let Some(server_name) = server.clone() { + let params = cursor.clone().map(|value| PaginatedRequestParams { + meta: None, + cursor: Some(value), + }); + let result = session + .list_resource_templates(&server_name, params) + .await + .map_err(|err| { + FunctionCallError::RespondToModel(format!( + "resources/templates/list failed: {err:#}" + )) + })?; + Ok(ListResourceTemplatesPayload::from_single_server( + server_name, + result, + )) + } else { + if cursor.is_some() { + return Err(FunctionCallError::RespondToModel( + "cursor can only be used when a server is specified".to_string(), + )); + } + + let templates = session + .services + .mcp_connection_manager + .read() + .await + .list_all_resource_templates() + .await; + Ok(ListResourceTemplatesPayload::from_all_servers(templates)) + } + } + .await; + + match payload_result { + Ok(payload) => match serialize_function_output(payload) { + Ok(output) => { + let content = function_call_output_content_items_to_text(&output.body) + .unwrap_or_default(); + let duration = start.elapsed(); + emit_tool_call_end( + &session, + turn.as_ref(), + &call_id, + invocation, + duration, + Ok(call_tool_result_from_content(&content, output.success)), + ) + .await; + Ok(output) + } + Err(err) => { + let duration = start.elapsed(); + let message = err.to_string(); + emit_tool_call_end( + &session, + turn.as_ref(), + &call_id, + invocation, + duration, + Err(message.clone()), + ) + .await; + Err(err) + } + }, + Err(err) => { + let duration = start.elapsed(); + let message = err.to_string(); + emit_tool_call_end( + &session, + turn.as_ref(), + &call_id, + invocation, + duration, + Err(message.clone()), + ) + .await; + Err(err) + } + } + } +} + +impl ToolHandler for ReadMcpResourceHandler { + type Output = FunctionToolOutput; + + fn tool_name(&self) -> ToolName { + ToolName::plain("read_mcp_resource") + } + fn kind(&self) -> ToolKind { ToolKind::Function } @@ -190,7 +465,6 @@ impl ToolHandler for McpResourceHandler { session, turn, call_id, - tool_name, payload, .. } = invocation; @@ -199,124 +473,80 @@ impl ToolHandler for McpResourceHandler { ToolPayload::Function { arguments } => arguments, _ => { return Err(FunctionCallError::RespondToModel( - "mcp_resource handler received unsupported payload".to_string(), + "read_mcp_resource handler received unsupported payload".to_string(), )); } }; - let arguments_value = parse_arguments(arguments.as_str())?; + let arguments = parse_arguments(arguments.as_str())?; + let args: ReadResourceArgs = parse_args(arguments.clone())?; + let ReadResourceArgs { server, uri } = args; + let server = normalize_required_string("server", server)?; + let uri = normalize_required_string("uri", uri)?; - match tool_name.name.as_str() { - "list_mcp_resources" => { - handle_list_resources( - Arc::clone(&session), - Arc::clone(&turn), - call_id.clone(), - arguments_value.clone(), - ) - .await - } - "list_mcp_resource_templates" => { - handle_list_resource_templates( - Arc::clone(&session), - Arc::clone(&turn), - call_id.clone(), - arguments_value.clone(), - ) - .await - } - "read_mcp_resource" => { - handle_read_resource( - Arc::clone(&session), - Arc::clone(&turn), - call_id, - arguments_value, - ) - .await - } - other => Err(FunctionCallError::RespondToModel(format!( - "unsupported MCP resource tool: {other}" - ))), - } - } -} + let invocation = McpInvocation { + server: server.clone(), + tool: "read_mcp_resource".to_string(), + arguments: arguments.clone(), + }; -#[expect( - clippy::await_holding_invalid_type, - reason = "MCP resource listing reads through the session-owned manager guard" -)] -async fn handle_list_resources( - session: Arc, - turn: Arc, - call_id: String, - arguments: Option, -) -> Result { - let args: ListResourcesArgs = parse_args_with_default(arguments.clone())?; - let ListResourcesArgs { server, cursor } = args; - let server = normalize_optional_string(server); - let cursor = normalize_optional_string(cursor); + emit_tool_call_begin(&session, turn.as_ref(), &call_id, invocation.clone()).await; + let start = Instant::now(); - let invocation = McpInvocation { - server: server.clone().unwrap_or_else(|| "codex".to_string()), - tool: "list_mcp_resources".to_string(), - arguments: arguments.clone(), - }; - - emit_tool_call_begin(&session, turn.as_ref(), &call_id, invocation.clone()).await; - let start = Instant::now(); - - let payload_result: Result = async { - if let Some(server_name) = server.clone() { - let params = cursor.clone().map(|value| PaginatedRequestParams { - meta: None, - cursor: Some(value), - }); + let payload_result: Result = async { let result = session - .list_resources(&server_name, params) + .read_resource( + &server, + ReadResourceRequestParams { + meta: None, + uri: uri.clone(), + }, + ) .await .map_err(|err| { - FunctionCallError::RespondToModel(format!("resources/list failed: {err:#}")) + FunctionCallError::RespondToModel(format!("resources/read failed: {err:#}")) })?; - Ok(ListResourcesPayload::from_single_server( - server_name, + + Ok(ReadResourcePayload { + server, + uri, result, - )) - } else { - if cursor.is_some() { - return Err(FunctionCallError::RespondToModel( - "cursor can only be used when a server is specified".to_string(), - )); - } - - let resources = session - .services - .mcp_connection_manager - .read() - .await - .list_all_resources() - .await; - Ok(ListResourcesPayload::from_all_servers(resources)) + }) } - } - .await; + .await; - match payload_result { - Ok(payload) => match serialize_function_output(payload) { - Ok(output) => { - let content = - function_call_output_content_items_to_text(&output.body).unwrap_or_default(); - let duration = start.elapsed(); - emit_tool_call_end( - &session, - turn.as_ref(), - &call_id, - invocation, - duration, - Ok(call_tool_result_from_content(&content, output.success)), - ) - .await; - Ok(output) - } + match payload_result { + Ok(payload) => match serialize_function_output(payload) { + Ok(output) => { + let content = function_call_output_content_items_to_text(&output.body) + .unwrap_or_default(); + let duration = start.elapsed(); + emit_tool_call_end( + &session, + turn.as_ref(), + &call_id, + invocation, + duration, + Ok(call_tool_result_from_content(&content, output.success)), + ) + .await; + Ok(output) + } + Err(err) => { + let duration = start.elapsed(); + let message = err.to_string(); + emit_tool_call_end( + &session, + turn.as_ref(), + &call_id, + invocation, + duration, + Err(message.clone()), + ) + .await; + Err(err) + } + }, Err(err) => { let duration = start.elapsed(); let message = err.to_string(); @@ -331,221 +561,6 @@ async fn handle_list_resources( .await; Err(err) } - }, - Err(err) => { - let duration = start.elapsed(); - let message = err.to_string(); - emit_tool_call_end( - &session, - turn.as_ref(), - &call_id, - invocation, - duration, - Err(message.clone()), - ) - .await; - Err(err) - } - } -} - -#[expect( - clippy::await_holding_invalid_type, - reason = "MCP resource template listing reads through the session-owned manager guard" -)] -async fn handle_list_resource_templates( - session: Arc, - turn: Arc, - call_id: String, - arguments: Option, -) -> Result { - let args: ListResourceTemplatesArgs = parse_args_with_default(arguments.clone())?; - let ListResourceTemplatesArgs { server, cursor } = args; - let server = normalize_optional_string(server); - let cursor = normalize_optional_string(cursor); - - let invocation = McpInvocation { - server: server.clone().unwrap_or_else(|| "codex".to_string()), - tool: "list_mcp_resource_templates".to_string(), - arguments: arguments.clone(), - }; - - emit_tool_call_begin(&session, turn.as_ref(), &call_id, invocation.clone()).await; - let start = Instant::now(); - - let payload_result: Result = async { - if let Some(server_name) = server.clone() { - let params = cursor.clone().map(|value| PaginatedRequestParams { - meta: None, - cursor: Some(value), - }); - let result = session - .list_resource_templates(&server_name, params) - .await - .map_err(|err| { - FunctionCallError::RespondToModel(format!( - "resources/templates/list failed: {err:#}" - )) - })?; - Ok(ListResourceTemplatesPayload::from_single_server( - server_name, - result, - )) - } else { - if cursor.is_some() { - return Err(FunctionCallError::RespondToModel( - "cursor can only be used when a server is specified".to_string(), - )); - } - - let templates = session - .services - .mcp_connection_manager - .read() - .await - .list_all_resource_templates() - .await; - Ok(ListResourceTemplatesPayload::from_all_servers(templates)) - } - } - .await; - - match payload_result { - Ok(payload) => match serialize_function_output(payload) { - Ok(output) => { - let content = - function_call_output_content_items_to_text(&output.body).unwrap_or_default(); - let duration = start.elapsed(); - emit_tool_call_end( - &session, - turn.as_ref(), - &call_id, - invocation, - duration, - Ok(call_tool_result_from_content(&content, output.success)), - ) - .await; - Ok(output) - } - Err(err) => { - let duration = start.elapsed(); - let message = err.to_string(); - emit_tool_call_end( - &session, - turn.as_ref(), - &call_id, - invocation, - duration, - Err(message.clone()), - ) - .await; - Err(err) - } - }, - Err(err) => { - let duration = start.elapsed(); - let message = err.to_string(); - emit_tool_call_end( - &session, - turn.as_ref(), - &call_id, - invocation, - duration, - Err(message.clone()), - ) - .await; - Err(err) - } - } -} - -async fn handle_read_resource( - session: Arc, - turn: Arc, - call_id: String, - arguments: Option, -) -> Result { - let args: ReadResourceArgs = parse_args(arguments.clone())?; - let ReadResourceArgs { server, uri } = args; - let server = normalize_required_string("server", server)?; - let uri = normalize_required_string("uri", uri)?; - - let invocation = McpInvocation { - server: server.clone(), - tool: "read_mcp_resource".to_string(), - arguments: arguments.clone(), - }; - - emit_tool_call_begin(&session, turn.as_ref(), &call_id, invocation.clone()).await; - let start = Instant::now(); - - let payload_result: Result = async { - let result = session - .read_resource( - &server, - ReadResourceRequestParams { - meta: None, - uri: uri.clone(), - }, - ) - .await - .map_err(|err| { - FunctionCallError::RespondToModel(format!("resources/read failed: {err:#}")) - })?; - - Ok(ReadResourcePayload { - server, - uri, - result, - }) - } - .await; - - match payload_result { - Ok(payload) => match serialize_function_output(payload) { - Ok(output) => { - let content = - function_call_output_content_items_to_text(&output.body).unwrap_or_default(); - let duration = start.elapsed(); - emit_tool_call_end( - &session, - turn.as_ref(), - &call_id, - invocation, - duration, - Ok(call_tool_result_from_content(&content, output.success)), - ) - .await; - Ok(output) - } - Err(err) => { - let duration = start.elapsed(); - let message = err.to_string(); - emit_tool_call_end( - &session, - turn.as_ref(), - &call_id, - invocation, - duration, - Err(message.clone()), - ) - .await; - Err(err) - } - }, - Err(err) => { - let duration = start.elapsed(); - let message = err.to_string(); - emit_tool_call_end( - &session, - turn.as_ref(), - &call_id, - invocation, - duration, - Err(message.clone()), - ) - .await; - Err(err) } } } diff --git a/codex-rs/core/src/tools/handlers/mod.rs b/codex-rs/core/src/tools/handlers/mod.rs index dea4b7dd8..7f9119583 100644 --- a/codex-rs/core/src/tools/handlers/mod.rs +++ b/codex-rs/core/src/tools/handlers/mod.rs @@ -38,20 +38,27 @@ pub use apply_patch::ApplyPatchHandler; use codex_protocol::models::AdditionalPermissionProfile; use codex_protocol::protocol::AskForApproval; pub use dynamic::DynamicToolHandler; -pub use goal::GoalHandler; +pub use goal::CreateGoalHandler; +pub use goal::GetGoalHandler; +pub use goal::UpdateGoalHandler; pub use mcp::McpHandler; -pub use mcp_resource::McpResourceHandler; +pub use mcp_resource::ListMcpResourceTemplatesHandler; +pub use mcp_resource::ListMcpResourcesHandler; +pub use mcp_resource::ReadMcpResourceHandler; pub use plan::PlanHandler; pub use request_permissions::RequestPermissionsHandler; pub use request_plugin_install::RequestPluginInstallHandler; pub use request_user_input::RequestUserInputHandler; +pub use shell::ContainerExecHandler; +pub use shell::LocalShellHandler; pub use shell::ShellCommandHandler; pub use shell::ShellHandler; pub use test_sync::TestSyncHandler; pub use tool_search::ToolSearchHandler; pub use unavailable_tool::UnavailableToolHandler; pub(crate) use unavailable_tool::unavailable_tool_message; -pub use unified_exec::UnifiedExecHandler; +pub use unified_exec::ExecCommandHandler; +pub use unified_exec::WriteStdinHandler; pub use view_image::ViewImageHandler; fn parse_arguments(arguments: &str) -> Result diff --git a/codex-rs/core/src/tools/handlers/multi_agents.rs b/codex-rs/core/src/tools/handlers/multi_agents.rs index 2d70d3e92..71ef84fd1 100644 --- a/codex-rs/core/src/tools/handlers/multi_agents.rs +++ b/codex-rs/core/src/tools/handlers/multi_agents.rs @@ -32,6 +32,7 @@ use codex_protocol::protocol::CollabResumeEndEvent; use codex_protocol::protocol::CollabWaitingBeginEvent; use codex_protocol::protocol::CollabWaitingEndEvent; use codex_protocol::user_input::UserInput; +use codex_tools::ToolName; use serde::Deserialize; use serde::Serialize; use serde_json::Value as JsonValue; diff --git a/codex-rs/core/src/tools/handlers/multi_agents/close_agent.rs b/codex-rs/core/src/tools/handlers/multi_agents/close_agent.rs index 0b308bb09..70d24c428 100644 --- a/codex-rs/core/src/tools/handlers/multi_agents/close_agent.rs +++ b/codex-rs/core/src/tools/handlers/multi_agents/close_agent.rs @@ -6,6 +6,10 @@ pub(crate) struct Handler; impl ToolHandler for Handler { type Output = CloseAgentResult; + fn tool_name(&self) -> ToolName { + ToolName::plain("close_agent") + } + fn kind(&self) -> ToolKind { ToolKind::Function } diff --git a/codex-rs/core/src/tools/handlers/multi_agents/resume_agent.rs b/codex-rs/core/src/tools/handlers/multi_agents/resume_agent.rs index 59a503893..8fa462261 100644 --- a/codex-rs/core/src/tools/handlers/multi_agents/resume_agent.rs +++ b/codex-rs/core/src/tools/handlers/multi_agents/resume_agent.rs @@ -8,6 +8,10 @@ pub(crate) struct Handler; impl ToolHandler for Handler { type Output = ResumeAgentResult; + fn tool_name(&self) -> ToolName { + ToolName::plain("resume_agent") + } + fn kind(&self) -> ToolKind { ToolKind::Function } diff --git a/codex-rs/core/src/tools/handlers/multi_agents/send_input.rs b/codex-rs/core/src/tools/handlers/multi_agents/send_input.rs index 1feb21b83..0994ba5e2 100644 --- a/codex-rs/core/src/tools/handlers/multi_agents/send_input.rs +++ b/codex-rs/core/src/tools/handlers/multi_agents/send_input.rs @@ -7,6 +7,10 @@ pub(crate) struct Handler; impl ToolHandler for Handler { type Output = SendInputResult; + fn tool_name(&self) -> ToolName { + ToolName::plain("send_input") + } + fn kind(&self) -> ToolKind { ToolKind::Function } diff --git a/codex-rs/core/src/tools/handlers/multi_agents/spawn.rs b/codex-rs/core/src/tools/handlers/multi_agents/spawn.rs index bc5dcd692..adfb926fe 100644 --- a/codex-rs/core/src/tools/handlers/multi_agents/spawn.rs +++ b/codex-rs/core/src/tools/handlers/multi_agents/spawn.rs @@ -13,6 +13,10 @@ pub(crate) struct Handler; impl ToolHandler for Handler { type Output = SpawnAgentResult; + fn tool_name(&self) -> ToolName { + ToolName::plain("spawn_agent") + } + fn kind(&self) -> ToolKind { ToolKind::Function } diff --git a/codex-rs/core/src/tools/handlers/multi_agents/wait.rs b/codex-rs/core/src/tools/handlers/multi_agents/wait.rs index 49b85dbfb..8d6c09193 100644 --- a/codex-rs/core/src/tools/handlers/multi_agents/wait.rs +++ b/codex-rs/core/src/tools/handlers/multi_agents/wait.rs @@ -18,6 +18,10 @@ pub(crate) struct Handler; impl ToolHandler for Handler { type Output = WaitAgentResult; + fn tool_name(&self) -> ToolName { + ToolName::plain("wait_agent") + } + fn kind(&self) -> ToolKind { ToolKind::Function } diff --git a/codex-rs/core/src/tools/handlers/multi_agents_v2.rs b/codex-rs/core/src/tools/handlers/multi_agents_v2.rs index b561c5acb..a477c25ca 100644 --- a/codex-rs/core/src/tools/handlers/multi_agents_v2.rs +++ b/codex-rs/core/src/tools/handlers/multi_agents_v2.rs @@ -22,6 +22,7 @@ use codex_protocol::protocol::CollabCloseEndEvent; use codex_protocol::protocol::CollabWaitingBeginEvent; use codex_protocol::protocol::CollabWaitingEndEvent; use codex_protocol::user_input::UserInput; +use codex_tools::ToolName; use serde::Deserialize; use serde::Serialize; use serde_json::Value as JsonValue; diff --git a/codex-rs/core/src/tools/handlers/multi_agents_v2/close_agent.rs b/codex-rs/core/src/tools/handlers/multi_agents_v2/close_agent.rs index c0a1bcbc5..f09bc7f34 100644 --- a/codex-rs/core/src/tools/handlers/multi_agents_v2/close_agent.rs +++ b/codex-rs/core/src/tools/handlers/multi_agents_v2/close_agent.rs @@ -6,6 +6,10 @@ pub(crate) struct Handler; impl ToolHandler for Handler { type Output = CloseAgentResult; + fn tool_name(&self) -> ToolName { + ToolName::plain("close_agent") + } + fn kind(&self) -> ToolKind { ToolKind::Function } diff --git a/codex-rs/core/src/tools/handlers/multi_agents_v2/followup_task.rs b/codex-rs/core/src/tools/handlers/multi_agents_v2/followup_task.rs index bcb3f49de..a5dfcb09d 100644 --- a/codex-rs/core/src/tools/handlers/multi_agents_v2/followup_task.rs +++ b/codex-rs/core/src/tools/handlers/multi_agents_v2/followup_task.rs @@ -9,6 +9,10 @@ pub(crate) struct Handler; impl ToolHandler for Handler { type Output = FunctionToolOutput; + fn tool_name(&self) -> ToolName { + ToolName::plain("followup_task") + } + fn kind(&self) -> ToolKind { ToolKind::Function } diff --git a/codex-rs/core/src/tools/handlers/multi_agents_v2/list_agents.rs b/codex-rs/core/src/tools/handlers/multi_agents_v2/list_agents.rs index 579c44199..dabfe72a7 100644 --- a/codex-rs/core/src/tools/handlers/multi_agents_v2/list_agents.rs +++ b/codex-rs/core/src/tools/handlers/multi_agents_v2/list_agents.rs @@ -6,6 +6,10 @@ pub(crate) struct Handler; impl ToolHandler for Handler { type Output = ListAgentsResult; + fn tool_name(&self) -> ToolName { + ToolName::plain("list_agents") + } + fn kind(&self) -> ToolKind { ToolKind::Function } diff --git a/codex-rs/core/src/tools/handlers/multi_agents_v2/message_tool.rs b/codex-rs/core/src/tools/handlers/multi_agents_v2/message_tool.rs index 12e443b81..dcf1a1e58 100644 --- a/codex-rs/core/src/tools/handlers/multi_agents_v2/message_tool.rs +++ b/codex-rs/core/src/tools/handlers/multi_agents_v2/message_tool.rs @@ -62,15 +62,7 @@ pub(crate) async fn handle_message_string_tool( target: String, message: String, ) -> Result { - handle_message_submission(invocation, mode, target, message_content(message)?).await -} - -async fn handle_message_submission( - invocation: ToolInvocation, - mode: MessageDeliveryMode, - target: String, - prompt: String, -) -> Result { + let prompt = message_content(message)?; let ToolInvocation { session, turn, diff --git a/codex-rs/core/src/tools/handlers/multi_agents_v2/send_message.rs b/codex-rs/core/src/tools/handlers/multi_agents_v2/send_message.rs index b327ccf52..e814c69f5 100644 --- a/codex-rs/core/src/tools/handlers/multi_agents_v2/send_message.rs +++ b/codex-rs/core/src/tools/handlers/multi_agents_v2/send_message.rs @@ -9,6 +9,10 @@ pub(crate) struct Handler; impl ToolHandler for Handler { type Output = FunctionToolOutput; + fn tool_name(&self) -> ToolName { + ToolName::plain("send_message") + } + fn kind(&self) -> ToolKind { ToolKind::Function } diff --git a/codex-rs/core/src/tools/handlers/multi_agents_v2/spawn.rs b/codex-rs/core/src/tools/handlers/multi_agents_v2/spawn.rs index 184ea36c5..8f09dbcaf 100644 --- a/codex-rs/core/src/tools/handlers/multi_agents_v2/spawn.rs +++ b/codex-rs/core/src/tools/handlers/multi_agents_v2/spawn.rs @@ -15,6 +15,10 @@ pub(crate) struct Handler; impl ToolHandler for Handler { type Output = SpawnAgentResult; + fn tool_name(&self) -> ToolName { + ToolName::plain("spawn_agent") + } + fn kind(&self) -> ToolKind { ToolKind::Function } diff --git a/codex-rs/core/src/tools/handlers/multi_agents_v2/wait.rs b/codex-rs/core/src/tools/handlers/multi_agents_v2/wait.rs index b86e237f5..706c0ad6a 100644 --- a/codex-rs/core/src/tools/handlers/multi_agents_v2/wait.rs +++ b/codex-rs/core/src/tools/handlers/multi_agents_v2/wait.rs @@ -10,6 +10,10 @@ pub(crate) struct Handler; impl ToolHandler for Handler { type Output = WaitAgentResult; + fn tool_name(&self) -> ToolName { + ToolName::plain("wait_agent") + } + fn kind(&self) -> ToolKind { ToolKind::Function } diff --git a/codex-rs/core/src/tools/handlers/plan.rs b/codex-rs/core/src/tools/handlers/plan.rs index 71636229e..ce217f457 100644 --- a/codex-rs/core/src/tools/handlers/plan.rs +++ b/codex-rs/core/src/tools/handlers/plan.rs @@ -1,6 +1,4 @@ use crate::function_tool::FunctionCallError; -use crate::session::session::Session; -use crate::session::turn_context::TurnContext; use crate::tools::context::ToolInvocation; use crate::tools::context::ToolOutput; use crate::tools::context::ToolPayload; @@ -11,6 +9,7 @@ use codex_protocol::models::FunctionCallOutputPayload; use codex_protocol::models::ResponseInputItem; use codex_protocol::plan_tool::UpdatePlanArgs; use codex_protocol::protocol::EventMsg; +use codex_tools::ToolName; use serde_json::Value as JsonValue; pub struct PlanHandler; @@ -46,6 +45,10 @@ impl ToolOutput for PlanToolOutput { impl ToolHandler for PlanHandler { type Output = PlanToolOutput; + fn tool_name(&self) -> ToolName { + ToolName::plain("update_plan") + } + fn kind(&self) -> ToolKind { ToolKind::Function } @@ -54,7 +57,7 @@ impl ToolHandler for PlanHandler { let ToolInvocation { session, turn, - call_id, + call_id: _, payload, .. } = invocation; @@ -68,33 +71,21 @@ impl ToolHandler for PlanHandler { } }; - handle_update_plan(session.as_ref(), turn.as_ref(), arguments, call_id).await?; + if turn.collaboration_mode.mode == ModeKind::Plan { + return Err(FunctionCallError::RespondToModel( + "update_plan is a TODO/checklist tool and is not allowed in Plan mode".to_string(), + )); + } + + let args = parse_update_plan_arguments(&arguments)?; + session + .send_event(turn.as_ref(), EventMsg::PlanUpdate(args)) + .await; Ok(PlanToolOutput) } } -/// This function doesn't do anything useful. However, it gives the model a structured way to record its plan that clients can read and render. -/// So it's the _inputs_ to this function that are useful to clients, not the outputs and neither are actually useful for the model other -/// than forcing it to come up and document a plan (TBD how that affects performance). -pub(crate) async fn handle_update_plan( - session: &Session, - turn_context: &TurnContext, - arguments: String, - _call_id: String, -) -> Result { - if turn_context.collaboration_mode.mode == ModeKind::Plan { - return Err(FunctionCallError::RespondToModel( - "update_plan is a TODO/checklist tool and is not allowed in Plan mode".to_string(), - )); - } - let args = parse_update_plan_arguments(&arguments)?; - session - .send_event(turn_context, EventMsg::PlanUpdate(args)) - .await; - Ok("Plan updated".to_string()) -} - fn parse_update_plan_arguments(arguments: &str) -> Result { serde_json::from_str::(arguments).map_err(|e| { FunctionCallError::RespondToModel(format!("failed to parse function arguments: {e}")) diff --git a/codex-rs/core/src/tools/handlers/request_permissions.rs b/codex-rs/core/src/tools/handlers/request_permissions.rs index 56facee65..7b49ec580 100644 --- a/codex-rs/core/src/tools/handlers/request_permissions.rs +++ b/codex-rs/core/src/tools/handlers/request_permissions.rs @@ -8,12 +8,17 @@ use crate::tools::context::ToolPayload; use crate::tools::handlers::parse_arguments_with_base_path; use crate::tools::registry::ToolHandler; use crate::tools::registry::ToolKind; +use codex_tools::ToolName; pub struct RequestPermissionsHandler; impl ToolHandler for RequestPermissionsHandler { type Output = FunctionToolOutput; + fn tool_name(&self) -> ToolName { + ToolName::plain("request_permissions") + } + fn kind(&self) -> ToolKind { ToolKind::Function } diff --git a/codex-rs/core/src/tools/handlers/request_plugin_install.rs b/codex-rs/core/src/tools/handlers/request_plugin_install.rs index 673a73bfb..879a8955b 100644 --- a/codex-rs/core/src/tools/handlers/request_plugin_install.rs +++ b/codex-rs/core/src/tools/handlers/request_plugin_install.rs @@ -13,6 +13,7 @@ use codex_tools::REQUEST_PLUGIN_INSTALL_PERSIST_KEY; use codex_tools::REQUEST_PLUGIN_INSTALL_TOOL_NAME; use codex_tools::RequestPluginInstallArgs; use codex_tools::RequestPluginInstallResult; +use codex_tools::ToolName; use codex_tools::all_requested_connectors_picked_up; use codex_tools::build_request_plugin_install_elicitation_request; use codex_tools::filter_request_plugin_install_discoverable_tools_for_client; @@ -37,6 +38,10 @@ pub struct RequestPluginInstallHandler; impl ToolHandler for RequestPluginInstallHandler { type Output = FunctionToolOutput; + fn tool_name(&self) -> ToolName { + ToolName::plain(REQUEST_PLUGIN_INSTALL_TOOL_NAME) + } + fn kind(&self) -> ToolKind { ToolKind::Function } diff --git a/codex-rs/core/src/tools/handlers/request_user_input.rs b/codex-rs/core/src/tools/handlers/request_user_input.rs index eea661276..cd00dc272 100644 --- a/codex-rs/core/src/tools/handlers/request_user_input.rs +++ b/codex-rs/core/src/tools/handlers/request_user_input.rs @@ -8,6 +8,7 @@ use crate::tools::registry::ToolKind; use codex_protocol::config_types::ModeKind; use codex_protocol::request_user_input::RequestUserInputArgs; use codex_tools::REQUEST_USER_INPUT_TOOL_NAME; +use codex_tools::ToolName; use codex_tools::normalize_request_user_input_args; use codex_tools::request_user_input_unavailable_message; @@ -18,6 +19,10 @@ pub struct RequestUserInputHandler { impl ToolHandler for RequestUserInputHandler { type Output = FunctionToolOutput; + fn tool_name(&self) -> ToolName { + ToolName::plain(REQUEST_USER_INPUT_TOOL_NAME) + } + fn kind(&self) -> ToolKind { ToolKind::Function } diff --git a/codex-rs/core/src/tools/handlers/shell.rs b/codex-rs/core/src/tools/handlers/shell.rs index d498f3653..dfc3c3e5b 100644 --- a/codex-rs/core/src/tools/handlers/shell.rs +++ b/codex-rs/core/src/tools/handlers/shell.rs @@ -40,8 +40,11 @@ use codex_protocol::models::AdditionalPermissionProfile; use codex_protocol::protocol::ExecCommandSource; use codex_shell_command::is_safe_command::is_known_safe_command; use codex_tools::ShellCommandBackendConfig; +use codex_tools::ToolName; pub struct ShellHandler; +pub struct ContainerExecHandler; +pub struct LocalShellHandler; #[derive(Clone, Copy, Debug, Eq, PartialEq)] enum ShellCommandBackend { @@ -53,16 +56,24 @@ pub struct ShellCommandHandler { backend: ShellCommandBackend, } -fn shell_payload_command(payload: &ToolPayload) -> Option { - match payload { - ToolPayload::Function { arguments } => parse_arguments::(arguments) - .ok() - .map(|params| codex_shell_command::parse_command::shlex_join(¶ms.command)), - ToolPayload::LocalShell { params } => Some(codex_shell_command::parse_command::shlex_join( - ¶ms.command, - )), - _ => None, - } +fn shell_function_payload_command(payload: &ToolPayload) -> Option { + let ToolPayload::Function { arguments } = payload else { + return None; + }; + + parse_arguments::(arguments) + .ok() + .map(|params| codex_shell_command::parse_command::shlex_join(¶ms.command)) +} + +fn local_shell_payload_command(payload: &ToolPayload) -> Option { + let ToolPayload::LocalShell { params } = payload else { + return None; + }; + + Some(codex_shell_command::parse_command::shlex_join( + ¶ms.command, + )) } fn shell_command_payload_command(payload: &ToolPayload) -> Option { @@ -182,31 +193,184 @@ impl From for ShellCommandHandler { impl ToolHandler for ShellHandler { type Output = FunctionToolOutput; + fn tool_name(&self) -> ToolName { + ToolName::plain("shell") + } + fn kind(&self) -> ToolKind { ToolKind::Function } fn matches_kind(&self, payload: &ToolPayload) -> bool { - matches!( - payload, - ToolPayload::Function { .. } | ToolPayload::LocalShell { .. } - ) + matches!(payload, ToolPayload::Function { .. }) } async fn is_mutating(&self, invocation: &ToolInvocation) -> bool { - match &invocation.payload { - ToolPayload::Function { arguments } => { - serde_json::from_str::(arguments) - .map(|params| !is_known_safe_command(¶ms.command)) - .unwrap_or(true) - } - ToolPayload::LocalShell { params } => !is_known_safe_command(¶ms.command), - _ => true, // unknown payloads => assume mutating - } + let ToolPayload::Function { arguments } = &invocation.payload else { + return true; + }; + + serde_json::from_str::(arguments) + .map(|params| !is_known_safe_command(¶ms.command)) + .unwrap_or(true) } fn pre_tool_use_payload(&self, invocation: &ToolInvocation) -> Option { - shell_payload_command(&invocation.payload).map(|command| PreToolUsePayload { + shell_function_pre_tool_use_payload(invocation) + } + + fn post_tool_use_payload( + &self, + invocation: &ToolInvocation, + result: &Self::Output, + ) -> Option { + shell_function_post_tool_use_payload(invocation, result) + } + + async fn handle(&self, invocation: ToolInvocation) -> Result { + let ToolInvocation { + session, + turn, + tracker, + call_id, + payload, + .. + } = invocation; + + let arguments = match payload { + ToolPayload::Function { arguments } => arguments, + _ => { + return Err(FunctionCallError::RespondToModel( + "unsupported payload for shell handler".to_string(), + )); + } + }; + + let cwd = resolve_workdir_base_path(&arguments, &turn.cwd)?; + let params: ShellToolCallParams = parse_arguments_with_base_path(&arguments, &cwd)?; + let prefix_rule = params.prefix_rule.clone(); + let exec_params = + ShellHandler::to_exec_params(¶ms, turn.as_ref(), session.conversation_id); + ShellHandler::run_exec_like(RunExecLikeArgs { + tool_name: "shell".to_string(), + exec_params, + hook_command: codex_shell_command::parse_command::shlex_join(¶ms.command), + additional_permissions: params.additional_permissions.clone(), + prefix_rule, + session, + turn, + tracker, + call_id, + freeform: false, + shell_runtime_backend: ShellRuntimeBackend::Generic, + }) + .await + } +} + +impl ToolHandler for ContainerExecHandler { + type Output = FunctionToolOutput; + + fn tool_name(&self) -> ToolName { + ToolName::plain("container.exec") + } + + fn kind(&self) -> ToolKind { + ToolKind::Function + } + + fn matches_kind(&self, payload: &ToolPayload) -> bool { + matches!(payload, ToolPayload::Function { .. }) + } + + async fn is_mutating(&self, invocation: &ToolInvocation) -> bool { + let ToolPayload::Function { arguments } = &invocation.payload else { + return true; + }; + + serde_json::from_str::(arguments) + .map(|params| !is_known_safe_command(¶ms.command)) + .unwrap_or(true) + } + + fn pre_tool_use_payload(&self, invocation: &ToolInvocation) -> Option { + shell_function_pre_tool_use_payload(invocation) + } + + fn post_tool_use_payload( + &self, + invocation: &ToolInvocation, + result: &Self::Output, + ) -> Option { + shell_function_post_tool_use_payload(invocation, result) + } + + async fn handle(&self, invocation: ToolInvocation) -> Result { + let ToolInvocation { + session, + turn, + tracker, + call_id, + payload, + .. + } = invocation; + + let arguments = match payload { + ToolPayload::Function { arguments } => arguments, + _ => { + return Err(FunctionCallError::RespondToModel( + "unsupported payload for container.exec handler".to_string(), + )); + } + }; + + let cwd = resolve_workdir_base_path(&arguments, &turn.cwd)?; + let params: ShellToolCallParams = parse_arguments_with_base_path(&arguments, &cwd)?; + let prefix_rule = params.prefix_rule.clone(); + let exec_params = + ShellHandler::to_exec_params(¶ms, turn.as_ref(), session.conversation_id); + ShellHandler::run_exec_like(RunExecLikeArgs { + tool_name: "container.exec".to_string(), + exec_params, + hook_command: codex_shell_command::parse_command::shlex_join(¶ms.command), + additional_permissions: params.additional_permissions.clone(), + prefix_rule, + session, + turn, + tracker, + call_id, + freeform: false, + shell_runtime_backend: ShellRuntimeBackend::Generic, + }) + .await + } +} + +impl ToolHandler for LocalShellHandler { + type Output = FunctionToolOutput; + + fn tool_name(&self) -> ToolName { + ToolName::plain("local_shell") + } + + fn kind(&self) -> ToolKind { + ToolKind::Function + } + + fn matches_kind(&self, payload: &ToolPayload) -> bool { + matches!(payload, ToolPayload::LocalShell { .. }) + } + + async fn is_mutating(&self, invocation: &ToolInvocation) -> bool { + let ToolPayload::LocalShell { params } = &invocation.payload else { + return true; + }; + + !is_known_safe_command(¶ms.command) + } + + fn pre_tool_use_payload(&self, invocation: &ToolInvocation) -> Option { + local_shell_payload_command(&invocation.payload).map(|command| PreToolUsePayload { tool_name: HookToolName::bash(), tool_input: serde_json::json!({ "command": command }), }) @@ -219,7 +383,7 @@ impl ToolHandler for ShellHandler { ) -> Option { let tool_response = result.post_tool_use_response(&invocation.call_id, &invocation.payload)?; - let command = shell_payload_command(&invocation.payload)?; + let command = local_shell_payload_command(&invocation.payload)?; Some(PostToolUsePayload { tool_name: HookToolName::bash(), tool_use_id: invocation.call_id.clone(), @@ -234,62 +398,63 @@ impl ToolHandler for ShellHandler { turn, tracker, call_id, - tool_name, payload, .. } = invocation; - match payload { - ToolPayload::Function { arguments } => { - let cwd = resolve_workdir_base_path(&arguments, &turn.cwd)?; - let params: ShellToolCallParams = parse_arguments_with_base_path(&arguments, &cwd)?; - let prefix_rule = params.prefix_rule.clone(); - let exec_params = - Self::to_exec_params(¶ms, turn.as_ref(), session.conversation_id); - Self::run_exec_like(RunExecLikeArgs { - tool_name: tool_name.display(), - exec_params, - hook_command: codex_shell_command::parse_command::shlex_join(¶ms.command), - additional_permissions: params.additional_permissions.clone(), - prefix_rule, - session, - turn, - tracker, - call_id, - freeform: false, - shell_runtime_backend: ShellRuntimeBackend::Generic, - }) - .await - } - ToolPayload::LocalShell { params } => { - let exec_params = - Self::to_exec_params(¶ms, turn.as_ref(), session.conversation_id); - Self::run_exec_like(RunExecLikeArgs { - tool_name: tool_name.display(), - exec_params, - hook_command: codex_shell_command::parse_command::shlex_join(¶ms.command), - additional_permissions: None, - prefix_rule: None, - session, - turn, - tracker, - call_id, - freeform: false, - shell_runtime_backend: ShellRuntimeBackend::Generic, - }) - .await - } - _ => Err(FunctionCallError::RespondToModel(format!( - "unsupported payload for shell handler: {}", - tool_name.display() - ))), - } + let ToolPayload::LocalShell { params } = payload else { + return Err(FunctionCallError::RespondToModel( + "unsupported payload for local_shell handler".to_string(), + )); + }; + + let exec_params = + ShellHandler::to_exec_params(¶ms, turn.as_ref(), session.conversation_id); + ShellHandler::run_exec_like(RunExecLikeArgs { + tool_name: "local_shell".to_string(), + exec_params, + hook_command: codex_shell_command::parse_command::shlex_join(¶ms.command), + additional_permissions: None, + prefix_rule: None, + session, + turn, + tracker, + call_id, + freeform: false, + shell_runtime_backend: ShellRuntimeBackend::Generic, + }) + .await } } +fn shell_function_pre_tool_use_payload(invocation: &ToolInvocation) -> Option { + shell_function_payload_command(&invocation.payload).map(|command| PreToolUsePayload { + tool_name: HookToolName::bash(), + tool_input: serde_json::json!({ "command": command }), + }) +} + +fn shell_function_post_tool_use_payload( + invocation: &ToolInvocation, + result: &FunctionToolOutput, +) -> Option { + let tool_response = result.post_tool_use_response(&invocation.call_id, &invocation.payload)?; + let command = shell_function_payload_command(&invocation.payload)?; + Some(PostToolUsePayload { + tool_name: HookToolName::bash(), + tool_use_id: invocation.call_id.clone(), + tool_input: serde_json::json!({ "command": command }), + tool_response, + }) +} + impl ToolHandler for ShellCommandHandler { type Output = FunctionToolOutput; + fn tool_name(&self) -> ToolName { + ToolName::plain("shell_command") + } + fn kind(&self) -> ToolKind { ToolKind::Function } @@ -348,7 +513,6 @@ impl ToolHandler for ShellCommandHandler { turn, tracker, call_id, - tool_name, payload, .. } = invocation; @@ -356,7 +520,7 @@ impl ToolHandler for ShellCommandHandler { let ToolPayload::Function { arguments } = payload else { return Err(FunctionCallError::RespondToModel(format!( "unsupported payload for shell_command handler: {}", - tool_name.display() + self.tool_name().display() ))); }; @@ -379,7 +543,7 @@ impl ToolHandler for ShellCommandHandler { turn.tools_config.allow_login_shell, )?; ShellHandler::run_exec_like(RunExecLikeArgs { - tool_name: tool_name.display(), + tool_name: self.tool_name().display(), exec_params, hook_command: params.command, additional_permissions: params.additional_permissions.clone(), diff --git a/codex-rs/core/src/tools/handlers/shell_tests.rs b/codex-rs/core/src/tools/handlers/shell_tests.rs index 49e2cf8f7..8a32e5404 100644 --- a/codex-rs/core/src/tools/handlers/shell_tests.rs +++ b/codex-rs/core/src/tools/handlers/shell_tests.rs @@ -16,8 +16,8 @@ use crate::tools::context::FunctionToolOutput; use crate::tools::context::ToolCallSource; use crate::tools::context::ToolInvocation; use crate::tools::context::ToolPayload; +use crate::tools::handlers::LocalShellHandler; use crate::tools::handlers::ShellCommandHandler; -use crate::tools::handlers::ShellHandler; use crate::tools::hook_names::HookToolName; use crate::tools::registry::ToolHandler; use crate::turn_diff_tracker::TurnDiffTracker; @@ -204,7 +204,7 @@ fn shell_command_handler_rejects_login_when_disallowed() { } #[tokio::test] -async fn shell_pre_tool_use_payload_uses_joined_command() { +async fn local_shell_pre_tool_use_payload_uses_joined_command() { let payload = ToolPayload::LocalShell { params: codex_protocol::models::ShellToolCallParams { command: vec![ @@ -215,13 +215,13 @@ async fn shell_pre_tool_use_payload_uses_joined_command() { workdir: None, timeout_ms: None, sandbox_permissions: None, - prefix_rule: None, additional_permissions: None, + prefix_rule: None, justification: None, }, }; let (session, turn) = make_session_and_context().await; - let handler = ShellHandler; + let handler = LocalShellHandler; assert_eq!( handler.pre_tool_use_payload(&ToolInvocation { @@ -230,7 +230,7 @@ async fn shell_pre_tool_use_payload_uses_joined_command() { cancellation_token: tokio_util::sync::CancellationToken::new(), tracker: Arc::new(Mutex::new(TurnDiffTracker::new())), call_id: "call-41".to_string(), - tool_name: codex_tools::ToolName::plain("shell"), + tool_name: codex_tools::ToolName::plain("local_shell"), source: crate::tools::context::ToolCallSource::Direct, payload, }), diff --git a/codex-rs/core/src/tools/handlers/test_sync.rs b/codex-rs/core/src/tools/handlers/test_sync.rs index ad2647243..e04400d17 100644 --- a/codex-rs/core/src/tools/handlers/test_sync.rs +++ b/codex-rs/core/src/tools/handlers/test_sync.rs @@ -15,6 +15,7 @@ use crate::tools::context::ToolPayload; use crate::tools::handlers::parse_arguments; use crate::tools::registry::ToolHandler; use crate::tools::registry::ToolKind; +use codex_tools::ToolName; pub struct TestSyncHandler; @@ -56,6 +57,10 @@ fn barrier_map() -> &'static tokio::sync::Mutex> { impl ToolHandler for TestSyncHandler { type Output = FunctionToolOutput; + fn tool_name(&self) -> ToolName { + ToolName::plain("test_sync_tool") + } + fn kind(&self) -> ToolKind { ToolKind::Function } diff --git a/codex-rs/core/src/tools/handlers/tool_search.rs b/codex-rs/core/src/tools/handlers/tool_search.rs index 67bc7b7f2..59deb5411 100644 --- a/codex-rs/core/src/tools/handlers/tool_search.rs +++ b/codex-rs/core/src/tools/handlers/tool_search.rs @@ -12,6 +12,7 @@ use bm25::SearchEngineBuilder; use codex_tools::LoadableToolSpec; use codex_tools::TOOL_SEARCH_DEFAULT_LIMIT; use codex_tools::TOOL_SEARCH_TOOL_NAME; +use codex_tools::ToolName; use codex_tools::coalesce_loadable_tool_specs; use std::collections::HashMap; @@ -44,6 +45,10 @@ impl ToolSearchHandler { impl ToolHandler for ToolSearchHandler { type Output = ToolSearchOutput; + fn tool_name(&self) -> ToolName { + ToolName::plain(TOOL_SEARCH_TOOL_NAME) + } + fn kind(&self) -> ToolKind { ToolKind::Function } diff --git a/codex-rs/core/src/tools/handlers/unavailable_tool.rs b/codex-rs/core/src/tools/handlers/unavailable_tool.rs index eb00cf8ff..64bb20058 100644 --- a/codex-rs/core/src/tools/handlers/unavailable_tool.rs +++ b/codex-rs/core/src/tools/handlers/unavailable_tool.rs @@ -4,8 +4,17 @@ use crate::tools::context::ToolInvocation; use crate::tools::context::ToolPayload; use crate::tools::registry::ToolHandler; use crate::tools::registry::ToolKind; +use codex_tools::ToolName; -pub struct UnavailableToolHandler; +pub struct UnavailableToolHandler { + tool_name: ToolName, +} + +impl UnavailableToolHandler { + pub fn new(tool_name: ToolName) -> Self { + Self { tool_name } + } +} pub(crate) fn unavailable_tool_message( tool_name: impl std::fmt::Display, @@ -19,19 +28,21 @@ pub(crate) fn unavailable_tool_message( impl ToolHandler for UnavailableToolHandler { type Output = FunctionToolOutput; + fn tool_name(&self) -> ToolName { + self.tool_name.clone() + } + fn kind(&self) -> ToolKind { ToolKind::Function } async fn handle(&self, invocation: ToolInvocation) -> Result { - let ToolInvocation { - tool_name, payload, .. - } = invocation; + let ToolInvocation { payload, .. } = invocation; match payload { ToolPayload::Function { .. } => Ok(FunctionToolOutput::from_text( unavailable_tool_message( - tool_name.display(), + self.tool_name.display(), "Retry after the tool becomes available or ask the user to re-enable it.", ), Some(false), diff --git a/codex-rs/core/src/tools/handlers/unified_exec.rs b/codex-rs/core/src/tools/handlers/unified_exec.rs index f109f22ac..c257240a4 100644 --- a/codex-rs/core/src/tools/handlers/unified_exec.rs +++ b/codex-rs/core/src/tools/handlers/unified_exec.rs @@ -33,6 +33,7 @@ use codex_protocol::models::AdditionalPermissionProfile; use codex_protocol::protocol::EventMsg; use codex_protocol::protocol::TerminalInteractionEvent; use codex_shell_command::is_safe_command::is_known_safe_command; +use codex_tools::ToolName; use codex_tools::UnifiedExecShellMode; use codex_utils_output_truncation::TruncationPolicy; use codex_utils_output_truncation::approx_token_count; @@ -40,7 +41,8 @@ use serde::Deserialize; use std::path::PathBuf; use std::sync::Arc; -pub struct UnifiedExecHandler; +pub struct ExecCommandHandler; +pub struct WriteStdinHandler; #[derive(Debug, Deserialize)] pub(crate) struct ExecCommandArgs { @@ -108,9 +110,13 @@ fn effective_max_output_tokens( resolve_max_tokens(max_output_tokens).min(truncation_policy.token_budget()) } -impl ToolHandler for UnifiedExecHandler { +impl ToolHandler for ExecCommandHandler { type Output = ExecCommandToolOutput; + fn tool_name(&self) -> ToolName { + ToolName::plain("exec_command") + } + fn kind(&self) -> ToolKind { ToolKind::Function } @@ -144,12 +150,6 @@ impl ToolHandler for UnifiedExecHandler { } fn pre_tool_use_payload(&self, invocation: &ToolInvocation) -> Option { - if invocation.tool_name.namespace.is_some() - || invocation.tool_name.name.as_str() != "exec_command" - { - return None; - } - let ToolPayload::Function { arguments } = &invocation.payload else { return None; }; @@ -167,23 +167,7 @@ impl ToolHandler for UnifiedExecHandler { invocation: &ToolInvocation, result: &Self::Output, ) -> Option { - let ToolPayload::Function { .. } = &invocation.payload else { - return None; - }; - - let command = result.hook_command.clone()?; - let tool_use_id = if result.event_call_id.is_empty() { - invocation.call_id.clone() - } else { - result.event_call_id.clone() - }; - let tool_response = result.post_tool_use_response(&tool_use_id, &invocation.payload)?; - Some(PostToolUsePayload { - tool_name: HookToolName::bash(), - tool_use_id, - tool_input: serde_json::json!({ "command": command }), - tool_response, - }) + post_unified_exec_tool_use_payload(invocation, result) } async fn handle(&self, invocation: ToolInvocation) -> Result { @@ -192,7 +176,6 @@ impl ToolHandler for UnifiedExecHandler { turn, tracker, call_id, - tool_name, payload, .. } = invocation; @@ -201,239 +184,292 @@ impl ToolHandler for UnifiedExecHandler { ToolPayload::Function { arguments } => arguments, _ => { return Err(FunctionCallError::RespondToModel( - "unified_exec handler received unsupported payload".to_string(), + "exec_command handler received unsupported payload".to_string(), )); } }; let manager: &UnifiedExecProcessManager = &session.services.unified_exec_manager; let context = UnifiedExecContext::new(session.clone(), turn.clone(), call_id.clone()); + let environment_args: ExecCommandEnvironmentArgs = parse_arguments(&arguments)?; + let Some(turn_environment) = + resolve_tool_environment(turn.as_ref(), environment_args.environment_id.as_deref())? + else { + return Err(FunctionCallError::RespondToModel( + "unified exec is unavailable in this session".to_string(), + )); + }; + let cwd = environment_args + .workdir + .as_deref() + .filter(|workdir| !workdir.is_empty()) + .map_or_else( + || turn_environment.cwd.clone(), + |workdir| turn_environment.cwd.join(workdir), + ); + let environment = Arc::clone(&turn_environment.environment); + let fs = environment.get_filesystem(); + let args: ExecCommandArgs = parse_arguments_with_base_path(&arguments, &cwd)?; + let hook_command = args.cmd.clone(); + maybe_emit_implicit_skill_invocation( + session.as_ref(), + context.turn.as_ref(), + &hook_command, + &cwd, + ) + .await; + let process_id = manager.allocate_process_id().await; + let command = get_command( + &args, + session.user_shell(), + &turn.tools_config.unified_exec_shell_mode, + turn.tools_config.allow_login_shell, + ) + .map_err(FunctionCallError::RespondToModel)?; + let command_for_display = codex_shell_command::parse_command::shlex_join(&command); - let response = match tool_name.name.as_str() { - "exec_command" => { - let environment_args: ExecCommandEnvironmentArgs = parse_arguments(&arguments)?; - let Some(turn_environment) = resolve_tool_environment( - turn.as_ref(), - environment_args.environment_id.as_deref(), - )? - else { - return Err(FunctionCallError::RespondToModel( - "unified exec is unavailable in this session".to_string(), - )); - }; - let cwd = environment_args - .workdir - .as_deref() - .filter(|workdir| !workdir.is_empty()) - .map_or_else( - || turn_environment.cwd.clone(), - |workdir| turn_environment.cwd.join(workdir), - ); - let environment = Arc::clone(&turn_environment.environment); - let fs = environment.get_filesystem(); - let args: ExecCommandArgs = parse_arguments_with_base_path(&arguments, &cwd)?; - let hook_command = args.cmd.clone(); - maybe_emit_implicit_skill_invocation( - session.as_ref(), - context.turn.as_ref(), - &hook_command, + let ExecCommandArgs { + tty, + yield_time_ms, + max_output_tokens, + sandbox_permissions, + additional_permissions, + justification, + prefix_rule, + .. + } = args; + let max_output_tokens = + effective_max_output_tokens(max_output_tokens, turn.truncation_policy); + + let exec_permission_approvals_enabled = + session.features().enabled(Feature::ExecPermissionApprovals); + let requested_additional_permissions = additional_permissions.clone(); + let effective_additional_permissions = apply_granted_turn_permissions( + context.session.as_ref(), + cwd.as_path(), + sandbox_permissions, + additional_permissions, + ) + .await; + let additional_permissions_allowed = exec_permission_approvals_enabled + || (session.features().enabled(Feature::RequestPermissionsTool) + && effective_additional_permissions.permissions_preapproved); + + // Sticky turn permissions have already been approved, so they should + // continue through the normal exec approval flow for the command. + if effective_additional_permissions + .sandbox_permissions + .requests_sandbox_override() + && !effective_additional_permissions.permissions_preapproved + && !matches!( + context.turn.approval_policy.value(), + codex_protocol::protocol::AskForApproval::OnRequest + ) + { + let approval_policy = context.turn.approval_policy.value(); + manager.release_process_id(process_id).await; + return Err(FunctionCallError::RespondToModel(format!( + "approval policy is {approval_policy:?}; reject command — you cannot ask for escalated permissions if the approval policy is {approval_policy:?}" + ))); + } + + let normalized_additional_permissions = match implicit_granted_permissions( + sandbox_permissions, + requested_additional_permissions.as_ref(), + &effective_additional_permissions, + ) + .map_or_else( + || { + normalize_and_validate_additional_permissions( + additional_permissions_allowed, + context.turn.approval_policy.value(), + effective_additional_permissions.sandbox_permissions, + effective_additional_permissions.additional_permissions, + effective_additional_permissions.permissions_preapproved, &cwd, ) - .await; - let process_id = manager.allocate_process_id().await; - let command = get_command( - &args, - session.user_shell(), - &turn.tools_config.unified_exec_shell_mode, - turn.tools_config.allow_login_shell, - ) - .map_err(FunctionCallError::RespondToModel)?; - let command_for_display = codex_shell_command::parse_command::shlex_join(&command); - - let ExecCommandArgs { - tty, - yield_time_ms, - max_output_tokens, - sandbox_permissions, - additional_permissions, - justification, - prefix_rule, - .. - } = args; - let max_output_tokens = - effective_max_output_tokens(max_output_tokens, turn.truncation_policy); - - let exec_permission_approvals_enabled = - session.features().enabled(Feature::ExecPermissionApprovals); - let requested_additional_permissions = additional_permissions.clone(); - let effective_additional_permissions = apply_granted_turn_permissions( - context.session.as_ref(), - cwd.as_path(), - sandbox_permissions, - additional_permissions, - ) - .await; - let additional_permissions_allowed = exec_permission_approvals_enabled - || (session.features().enabled(Feature::RequestPermissionsTool) - && effective_additional_permissions.permissions_preapproved); - - // Sticky turn permissions have already been approved, so they should - // continue through the normal exec approval flow for the command. - if effective_additional_permissions - .sandbox_permissions - .requests_sandbox_override() - && !effective_additional_permissions.permissions_preapproved - && !matches!( - context.turn.approval_policy.value(), - codex_protocol::protocol::AskForApproval::OnRequest - ) - { - let approval_policy = context.turn.approval_policy.value(); - manager.release_process_id(process_id).await; - return Err(FunctionCallError::RespondToModel(format!( - "approval policy is {approval_policy:?}; reject command — you cannot ask for escalated permissions if the approval policy is {approval_policy:?}" - ))); - } - - let normalized_additional_permissions = match implicit_granted_permissions( - sandbox_permissions, - requested_additional_permissions.as_ref(), - &effective_additional_permissions, - ) - .map_or_else( - || { - normalize_and_validate_additional_permissions( - additional_permissions_allowed, - context.turn.approval_policy.value(), - effective_additional_permissions.sandbox_permissions, - effective_additional_permissions.additional_permissions, - effective_additional_permissions.permissions_preapproved, - &cwd, - ) - }, - |permissions| Ok(Some(permissions)), - ) { - Ok(normalized) => normalized, - Err(err) => { - manager.release_process_id(process_id).await; - return Err(FunctionCallError::RespondToModel(err)); - } - }; - - if let Some(output) = intercept_apply_patch( - &command, - &cwd, - fs.as_ref(), - context.session.clone(), - context.turn.clone(), - Some(&tracker), - &context.call_id, - &tool_name.name, - ) - .await? - { - manager.release_process_id(process_id).await; - return Ok(ExecCommandToolOutput { - event_call_id: String::new(), - chunk_id: String::new(), - wall_time: std::time::Duration::ZERO, - raw_output: output.into_text().into_bytes(), - max_output_tokens: Some(max_output_tokens), - process_id: None, - exit_code: None, - original_token_count: None, - hook_command: None, - }); - } - - emit_unified_exec_tty_metric(&turn.session_telemetry, tty); - match manager - .exec_command( - ExecCommandRequest { - command, - hook_command: hook_command.clone(), - process_id, - yield_time_ms, - max_output_tokens: Some(max_output_tokens), - cwd, - environment, - network: context.turn.network.clone(), - tty, - sandbox_permissions: effective_additional_permissions - .sandbox_permissions, - additional_permissions: normalized_additional_permissions, - additional_permissions_preapproved: effective_additional_permissions - .permissions_preapproved, - justification, - prefix_rule, - }, - &context, - ) - .await - { - Ok(response) => response, - Err(UnifiedExecError::SandboxDenied { output, .. }) => { - let output_text = output.aggregated_output.text; - let original_token_count = approx_token_count(&output_text); - ExecCommandToolOutput { - event_call_id: context.call_id.clone(), - chunk_id: generate_chunk_id(), - wall_time: output.duration, - raw_output: output_text.into_bytes(), - max_output_tokens: Some(max_output_tokens), - // Sandbox denial is terminal, so there is no live - // process for write_stdin to resume. - process_id: None, - exit_code: Some(output.exit_code), - original_token_count: Some(original_token_count), - hook_command: Some(hook_command), - } - } - Err(err) => { - return Err(FunctionCallError::RespondToModel(format!( - "exec_command failed for `{command_for_display}`: {err:?}" - ))); - } - } - } - "write_stdin" => { - let args: WriteStdinArgs = parse_arguments(&arguments)?; - let max_output_tokens = - effective_max_output_tokens(args.max_output_tokens, turn.truncation_policy); - let response = manager - .write_stdin(WriteStdinRequest { - process_id: args.session_id, - input: &args.chars, - yield_time_ms: args.yield_time_ms, - max_output_tokens: Some(max_output_tokens), - }) - .await - .map_err(|err| { - FunctionCallError::RespondToModel(format!("write_stdin failed: {err}")) - })?; - - let interaction = TerminalInteractionEvent { - call_id: response.event_call_id.clone(), - process_id: args.session_id.to_string(), - stdin: args.chars.clone(), - }; - session - .send_event(turn.as_ref(), EventMsg::TerminalInteraction(interaction)) - .await; - - response - } - other => { - return Err(FunctionCallError::RespondToModel(format!( - "unsupported unified exec function {other}" - ))); + }, + |permissions| Ok(Some(permissions)), + ) { + Ok(normalized) => normalized, + Err(err) => { + manager.release_process_id(process_id).await; + return Err(FunctionCallError::RespondToModel(err)); } }; + if let Some(output) = intercept_apply_patch( + &command, + &cwd, + fs.as_ref(), + context.session.clone(), + context.turn.clone(), + Some(&tracker), + &context.call_id, + "exec_command", + ) + .await? + { + manager.release_process_id(process_id).await; + return Ok(ExecCommandToolOutput { + event_call_id: String::new(), + chunk_id: String::new(), + wall_time: std::time::Duration::ZERO, + raw_output: output.into_text().into_bytes(), + max_output_tokens: Some(max_output_tokens), + process_id: None, + exit_code: None, + original_token_count: None, + hook_command: None, + }); + } + + emit_unified_exec_tty_metric(&turn.session_telemetry, tty); + match manager + .exec_command( + ExecCommandRequest { + command, + hook_command: hook_command.clone(), + process_id, + yield_time_ms, + max_output_tokens: Some(max_output_tokens), + cwd, + environment, + network: context.turn.network.clone(), + tty, + sandbox_permissions: effective_additional_permissions.sandbox_permissions, + additional_permissions: normalized_additional_permissions, + additional_permissions_preapproved: effective_additional_permissions + .permissions_preapproved, + justification, + prefix_rule, + }, + &context, + ) + .await + { + Ok(response) => Ok(response), + Err(UnifiedExecError::SandboxDenied { output, .. }) => { + let output_text = output.aggregated_output.text; + let original_token_count = approx_token_count(&output_text); + Ok(ExecCommandToolOutput { + event_call_id: context.call_id.clone(), + chunk_id: generate_chunk_id(), + wall_time: output.duration, + raw_output: output_text.into_bytes(), + max_output_tokens: Some(max_output_tokens), + // Sandbox denial is terminal, so there is no live + // process for write_stdin to resume. + process_id: None, + exit_code: Some(output.exit_code), + original_token_count: Some(original_token_count), + hook_command: Some(hook_command), + }) + } + Err(err) => Err(FunctionCallError::RespondToModel(format!( + "exec_command failed for `{command_for_display}`: {err:?}" + ))), + } + } +} + +impl ToolHandler for WriteStdinHandler { + type Output = ExecCommandToolOutput; + + fn tool_name(&self) -> ToolName { + ToolName::plain("write_stdin") + } + + fn kind(&self) -> ToolKind { + ToolKind::Function + } + + fn matches_kind(&self, payload: &ToolPayload) -> bool { + matches!(payload, ToolPayload::Function { .. }) + } + + async fn is_mutating(&self, _invocation: &ToolInvocation) -> bool { + true + } + + fn post_tool_use_payload( + &self, + invocation: &ToolInvocation, + result: &Self::Output, + ) -> Option { + post_unified_exec_tool_use_payload(invocation, result) + } + + async fn handle(&self, invocation: ToolInvocation) -> Result { + let ToolInvocation { + session, + turn, + payload, + .. + } = invocation; + + let arguments = match payload { + ToolPayload::Function { arguments } => arguments, + _ => { + return Err(FunctionCallError::RespondToModel( + "write_stdin handler received unsupported payload".to_string(), + )); + } + }; + + let args: WriteStdinArgs = parse_arguments(&arguments)?; + let max_output_tokens = + effective_max_output_tokens(args.max_output_tokens, turn.truncation_policy); + let response = session + .services + .unified_exec_manager + .write_stdin(WriteStdinRequest { + process_id: args.session_id, + input: &args.chars, + yield_time_ms: args.yield_time_ms, + max_output_tokens: Some(max_output_tokens), + }) + .await + .map_err(|err| { + FunctionCallError::RespondToModel(format!("write_stdin failed: {err}")) + })?; + + let interaction = TerminalInteractionEvent { + call_id: response.event_call_id.clone(), + process_id: args.session_id.to_string(), + stdin: args.chars.clone(), + }; + session + .send_event(turn.as_ref(), EventMsg::TerminalInteraction(interaction)) + .await; + Ok(response) } } +fn post_unified_exec_tool_use_payload( + invocation: &ToolInvocation, + result: &ExecCommandToolOutput, +) -> Option { + let ToolPayload::Function { .. } = &invocation.payload else { + return None; + }; + + let command = result.hook_command.clone()?; + let tool_use_id = if result.event_call_id.is_empty() { + invocation.call_id.clone() + } else { + result.event_call_id.clone() + }; + let tool_response = result.post_tool_use_response(&tool_use_id, &invocation.payload)?; + Some(PostToolUsePayload { + tool_name: HookToolName::bash(), + tool_use_id, + tool_input: serde_json::json!({ "command": command }), + tool_response, + }) +} + fn emit_unified_exec_tty_metric(session_telemetry: &SessionTelemetry, tty: bool) { session_telemetry.counter( TOOL_CALL_UNIFIED_EXEC_METRIC, diff --git a/codex-rs/core/src/tools/handlers/unified_exec_tests.rs b/codex-rs/core/src/tools/handlers/unified_exec_tests.rs index 70b933bad..8818b2e34 100644 --- a/codex-rs/core/src/tools/handlers/unified_exec_tests.rs +++ b/codex-rs/core/src/tools/handlers/unified_exec_tests.rs @@ -184,7 +184,7 @@ async fn exec_command_pre_tool_use_payload_uses_raw_command() { arguments: serde_json::json!({ "cmd": "printf exec command" }).to_string(), }; let (session, turn) = make_session_and_context().await; - let handler = UnifiedExecHandler; + let handler = ExecCommandHandler; assert_eq!( handler.pre_tool_use_payload(&ToolInvocation { @@ -210,7 +210,7 @@ async fn exec_command_pre_tool_use_payload_skips_write_stdin() { arguments: serde_json::json!({ "chars": "echo hi" }).to_string(), }; let (session, turn) = make_session_and_context().await; - let handler = UnifiedExecHandler; + let handler = WriteStdinHandler; assert_eq!( handler.pre_tool_use_payload(&ToolInvocation { @@ -244,8 +244,9 @@ async fn exec_command_post_tool_use_payload_uses_output_for_noninteractive_one_s hook_command: Some("echo three".to_string()), }; let invocation = invocation_for_payload("exec_command", "call-43", payload).await; + let handler = ExecCommandHandler; assert_eq!( - UnifiedExecHandler.post_tool_use_payload(&invocation, &output), + handler.post_tool_use_payload(&invocation, &output), Some(crate::tools::registry::PostToolUsePayload { tool_name: HookToolName::bash(), tool_use_id: "call-43".to_string(), @@ -272,9 +273,10 @@ async fn exec_command_post_tool_use_payload_uses_output_for_interactive_completi hook_command: Some("echo three".to_string()), }; let invocation = invocation_for_payload("exec_command", "call-44", payload).await; + let handler = ExecCommandHandler; assert_eq!( - UnifiedExecHandler.post_tool_use_payload(&invocation, &output), + handler.post_tool_use_payload(&invocation, &output), Some(crate::tools::registry::PostToolUsePayload { tool_name: HookToolName::bash(), tool_use_id: "call-44".to_string(), @@ -301,10 +303,8 @@ async fn exec_command_post_tool_use_payload_skips_running_sessions() { hook_command: Some("echo three".to_string()), }; let invocation = invocation_for_payload("exec_command", "call-45", payload).await; - assert_eq!( - UnifiedExecHandler.post_tool_use_payload(&invocation, &output), - None - ); + let handler = ExecCommandHandler; + assert_eq!(handler.post_tool_use_payload(&invocation, &output), None); } #[tokio::test] @@ -328,9 +328,10 @@ async fn write_stdin_post_tool_use_payload_uses_original_exec_call_id_and_comman hook_command: Some("sleep 1; echo finished".to_string()), }; let invocation = invocation_for_payload("write_stdin", "write-stdin-call", payload).await; + let handler = WriteStdinHandler; assert_eq!( - UnifiedExecHandler.post_tool_use_payload(&invocation, &output), + handler.post_tool_use_payload(&invocation, &output), Some(crate::tools::registry::PostToolUsePayload { tool_name: HookToolName::bash(), tool_use_id: "exec-call-45".to_string(), @@ -369,10 +370,11 @@ async fn write_stdin_post_tool_use_payload_keeps_parallel_session_metadata_separ }; let invocation_b = invocation_for_payload("write_stdin", "write-call-b", payload.clone()).await; let invocation_a = invocation_for_payload("write_stdin", "write-call-a", payload).await; + let handler = WriteStdinHandler; let payloads = [ - UnifiedExecHandler.post_tool_use_payload(&invocation_b, &output_b), - UnifiedExecHandler.post_tool_use_payload(&invocation_a, &output_a), + handler.post_tool_use_payload(&invocation_b, &output_b), + handler.post_tool_use_payload(&invocation_a, &output_a), ]; assert_eq!( diff --git a/codex-rs/core/src/tools/handlers/view_image.rs b/codex-rs/core/src/tools/handlers/view_image.rs index d05807ef7..a7cbe7d97 100644 --- a/codex-rs/core/src/tools/handlers/view_image.rs +++ b/codex-rs/core/src/tools/handlers/view_image.rs @@ -19,6 +19,7 @@ use crate::tools::context::ToolPayload; use crate::tools::handlers::parse_arguments; use crate::tools::registry::ToolHandler; use crate::tools::registry::ToolKind; +use codex_tools::ToolName; pub struct ViewImageHandler; @@ -39,6 +40,10 @@ enum ViewImageDetail { impl ToolHandler for ViewImageHandler { type Output = ViewImageOutput; + fn tool_name(&self) -> ToolName { + ToolName::plain("view_image") + } + fn kind(&self) -> ToolKind { ToolKind::Function } diff --git a/codex-rs/core/src/tools/registry.rs b/codex-rs/core/src/tools/registry.rs index 87b36ff17..bdf18cf2f 100644 --- a/codex-rs/core/src/tools/registry.rs +++ b/codex-rs/core/src/tools/registry.rs @@ -44,6 +44,9 @@ pub enum ToolKind { pub trait ToolHandler: Send + Sync { type Output: ToolOutput + 'static; + /// The concrete tool name handled by this handler instance. + fn tool_name(&self) -> ToolName; + fn kind(&self) -> ToolKind; fn matches_kind(&self, payload: &ToolPayload) -> bool { @@ -227,10 +230,11 @@ impl ToolRegistry { } #[cfg(test)] - pub(crate) fn with_handler_for_test(name: ToolName, handler: Arc) -> Self + pub(crate) fn with_handler_for_test(handler: Arc) -> Self where T: ToolHandler + 'static, { + let name = handler.tool_name(); Self::new(HashMap::from([(name, handler as Arc)])) } @@ -250,14 +254,6 @@ impl ToolRegistry { self.handler(name)?.create_diff_consumer() } - // TODO(jif) for dynamic tools. - // pub fn register(&mut self, name: impl Into, handler: Arc) { - // let name = name.into(); - // if self.handlers.insert(name.clone(), handler).is_some() { - // warn!("overwriting handler for tool {name}"); - // } - // } - #[expect( clippy::await_holding_invalid_type, reason = "tool dispatch must keep active-turn accounting atomic" @@ -539,11 +535,11 @@ impl ToolRegistryBuilder { .push(ConfiguredToolSpec::new(spec, supports_parallel_tool_calls)); } - pub fn register_handler(&mut self, name: impl Into, handler: Arc) + pub fn register_handler(&mut self, handler: Arc) where H: ToolHandler + 'static, { - let name = name.into(); + let name = handler.tool_name(); let display_name = name.display(); let handler: Arc = handler; if self.handlers.insert(name, handler).is_some() { @@ -551,24 +547,6 @@ impl ToolRegistryBuilder { } } - // TODO(jif) for dynamic tools. - // pub fn register_many(&mut self, names: I, handler: Arc) - // where - // I: IntoIterator, - // I::Item: Into, - // { - // for name in names { - // let name = name.into(); - // if self - // .handlers - // .insert(name.clone(), handler.clone()) - // .is_some() - // { - // warn!("overwriting handler for tool {name}"); - // } - // } - // } - pub fn build(self) -> (Vec, ToolRegistry) { let registry = ToolRegistry::new(self.handlers); (self.specs, registry) diff --git a/codex-rs/core/src/tools/registry_tests.rs b/codex-rs/core/src/tools/registry_tests.rs index d44c3d0f9..ef7273999 100644 --- a/codex-rs/core/src/tools/registry_tests.rs +++ b/codex-rs/core/src/tools/registry_tests.rs @@ -1,12 +1,17 @@ use super::*; use pretty_assertions::assert_eq; -#[derive(Default)] -struct TestHandler; +struct TestHandler { + tool_name: codex_tools::ToolName, +} impl ToolHandler for TestHandler { type Output = crate::tools::context::FunctionToolOutput; + fn tool_name(&self) -> codex_tools::ToolName { + self.tool_name.clone() + } + fn kind(&self) -> ToolKind { ToolKind::Function } @@ -21,12 +26,16 @@ impl ToolHandler for TestHandler { #[test] fn handler_looks_up_namespaced_aliases_explicitly() { - let plain_handler = Arc::new(TestHandler) as Arc; - let namespaced_handler = Arc::new(TestHandler) as Arc; let namespace = "mcp__codex_apps__gmail"; let tool_name = "gmail_get_recent_emails"; let plain_name = codex_tools::ToolName::plain(tool_name); let namespaced_name = codex_tools::ToolName::namespaced(namespace, tool_name); + let plain_handler = Arc::new(TestHandler { + tool_name: plain_name.clone(), + }) as Arc; + let namespaced_handler = Arc::new(TestHandler { + tool_name: namespaced_name.clone(), + }) as Arc; let registry = ToolRegistry::new(HashMap::from([ (plain_name.clone(), Arc::clone(&plain_handler)), (namespaced_name.clone(), Arc::clone(&namespaced_handler)), diff --git a/codex-rs/core/src/tools/spec.rs b/codex-rs/core/src/tools/spec.rs index 83e8600e0..994e20ccd 100644 --- a/codex-rs/core/src/tools/spec.rs +++ b/codex-rs/core/src/tools/spec.rs @@ -1,6 +1,7 @@ use crate::shell::Shell; use crate::shell::ShellType; -use crate::tools::handlers::agent_jobs::BatchJobHandler; +use crate::tools::handlers::agent_jobs::ReportAgentJobResultHandler; +use crate::tools::handlers::agent_jobs::SpawnAgentsOnCsvHandler; use crate::tools::handlers::multi_agents_common::DEFAULT_WAIT_TIMEOUT_MS; use crate::tools::handlers::multi_agents_common::MAX_WAIT_TIMEOUT_MS; use crate::tools::handlers::multi_agents_common::MIN_WAIT_TIMEOUT_MS; @@ -76,11 +77,17 @@ pub(crate) fn build_specs_with_discoverable_tools( use crate::tools::handlers::ApplyPatchHandler; use crate::tools::handlers::CodeModeExecuteHandler; use crate::tools::handlers::CodeModeWaitHandler; + use crate::tools::handlers::ContainerExecHandler; + use crate::tools::handlers::CreateGoalHandler; use crate::tools::handlers::DynamicToolHandler; - use crate::tools::handlers::GoalHandler; + use crate::tools::handlers::ExecCommandHandler; + use crate::tools::handlers::GetGoalHandler; + use crate::tools::handlers::ListMcpResourceTemplatesHandler; + use crate::tools::handlers::ListMcpResourcesHandler; + use crate::tools::handlers::LocalShellHandler; use crate::tools::handlers::McpHandler; - use crate::tools::handlers::McpResourceHandler; use crate::tools::handlers::PlanHandler; + use crate::tools::handlers::ReadMcpResourceHandler; use crate::tools::handlers::RequestPermissionsHandler; use crate::tools::handlers::RequestPluginInstallHandler; use crate::tools::handlers::RequestUserInputHandler; @@ -89,8 +96,9 @@ pub(crate) fn build_specs_with_discoverable_tools( use crate::tools::handlers::TestSyncHandler; use crate::tools::handlers::ToolSearchHandler; use crate::tools::handlers::UnavailableToolHandler; - use crate::tools::handlers::UnifiedExecHandler; + use crate::tools::handlers::UpdateGoalHandler; use crate::tools::handlers::ViewImageHandler; + use crate::tools::handlers::WriteStdinHandler; use crate::tools::handlers::multi_agents::CloseAgentHandler; use crate::tools::handlers::multi_agents::ResumeAgentHandler; use crate::tools::handlers::multi_agents::SendInputHandler; @@ -150,30 +158,11 @@ pub(crate) fn build_specs_with_discoverable_tools( }, }, ); - let shell_handler = Arc::new(ShellHandler); - let unified_exec_handler = Arc::new(UnifiedExecHandler); - let plan_handler = Arc::new(PlanHandler); - let apply_patch_handler = Arc::new(ApplyPatchHandler); - let dynamic_tool_handler = Arc::new(DynamicToolHandler); - let goal_handler = Arc::new(GoalHandler); - let view_image_handler = Arc::new(ViewImageHandler); - let mcp_handler = Arc::new(McpHandler); - let mcp_resource_handler = Arc::new(McpResourceHandler); - let shell_command_handler = Arc::new(ShellCommandHandler::from(config.shell_command_backend)); - let request_permissions_handler = Arc::new(RequestPermissionsHandler); - let request_user_input_handler = Arc::new(RequestUserInputHandler { - available_modes: config.request_user_input_available_modes.clone(), - }); let deferred_dynamic_tools = dynamic_tools .iter() .filter(|tool| tool.defer_loading && (config.namespace_tools || tool.namespace.is_none())) .cloned() .collect::>(); - let mut tool_search_handler = None; - let request_plugin_install_handler = Arc::new(RequestPluginInstallHandler); - let code_mode_handler = Arc::new(CodeModeExecuteHandler); - let code_mode_wait_handler = Arc::new(CodeModeWaitHandler); - let unavailable_tool_handler = Arc::new(UnavailableToolHandler); let mut existing_spec_names = plan .specs .iter() @@ -191,113 +180,137 @@ pub(crate) fn build_specs_with_discoverable_tools( } for handler in plan.handlers { + let name = handler.name; match handler.kind { - ToolHandlerKind::AgentJobs => { - builder.register_handler(handler.name, Arc::new(BatchJobHandler)); - } ToolHandlerKind::ApplyPatch => { - builder.register_handler(handler.name, apply_patch_handler.clone()); + builder.register_handler(Arc::new(ApplyPatchHandler)); } ToolHandlerKind::CloseAgentV1 => { - builder.register_handler(handler.name, Arc::new(CloseAgentHandler)); + builder.register_handler(Arc::new(CloseAgentHandler)); } ToolHandlerKind::CloseAgentV2 => { - builder.register_handler(handler.name, Arc::new(CloseAgentHandlerV2)); + builder.register_handler(Arc::new(CloseAgentHandlerV2)); } ToolHandlerKind::CodeModeExecute => { - builder.register_handler(handler.name, code_mode_handler.clone()); + builder.register_handler(Arc::new(CodeModeExecuteHandler)); } ToolHandlerKind::CodeModeWait => { - builder.register_handler(handler.name, code_mode_wait_handler.clone()); + builder.register_handler(Arc::new(CodeModeWaitHandler)); + } + ToolHandlerKind::ContainerExec => { + builder.register_handler(Arc::new(ContainerExecHandler)); + } + ToolHandlerKind::CreateGoal => { + builder.register_handler(Arc::new(CreateGoalHandler)); } ToolHandlerKind::DynamicTool => { - builder.register_handler(handler.name, dynamic_tool_handler.clone()); + builder.register_handler(Arc::new(DynamicToolHandler::new(name))); + } + ToolHandlerKind::ExecCommand => { + builder.register_handler(Arc::new(ExecCommandHandler)); } ToolHandlerKind::FollowupTaskV2 => { - builder.register_handler(handler.name, Arc::new(FollowupTaskHandlerV2)); + builder.register_handler(Arc::new(FollowupTaskHandlerV2)); } - ToolHandlerKind::Goal => { - builder.register_handler(handler.name, goal_handler.clone()); + ToolHandlerKind::GetGoal => { + builder.register_handler(Arc::new(GetGoalHandler)); } ToolHandlerKind::ListAgentsV2 => { - builder.register_handler(handler.name, Arc::new(ListAgentsHandlerV2)); + builder.register_handler(Arc::new(ListAgentsHandlerV2)); + } + ToolHandlerKind::ListMcpResources => { + builder.register_handler(Arc::new(ListMcpResourcesHandler)); + } + ToolHandlerKind::ListMcpResourceTemplates => { + builder.register_handler(Arc::new(ListMcpResourceTemplatesHandler)); + } + ToolHandlerKind::LocalShell => { + builder.register_handler(Arc::new(LocalShellHandler)); } ToolHandlerKind::Mcp => { - builder.register_handler(handler.name, mcp_handler.clone()); - } - ToolHandlerKind::McpResource => { - builder.register_handler(handler.name, mcp_resource_handler.clone()); + builder.register_handler(Arc::new(McpHandler::new(name))); } ToolHandlerKind::Plan => { - builder.register_handler(handler.name, plan_handler.clone()); + builder.register_handler(Arc::new(PlanHandler)); + } + ToolHandlerKind::ReadMcpResource => { + builder.register_handler(Arc::new(ReadMcpResourceHandler)); + } + ToolHandlerKind::ReportAgentJobResult => { + builder.register_handler(Arc::new(ReportAgentJobResultHandler)); } ToolHandlerKind::RequestPermissions => { - builder.register_handler(handler.name, request_permissions_handler.clone()); + builder.register_handler(Arc::new(RequestPermissionsHandler)); } ToolHandlerKind::RequestUserInput => { - builder.register_handler(handler.name, request_user_input_handler.clone()); + builder.register_handler(Arc::new(RequestUserInputHandler { + available_modes: config.request_user_input_available_modes.clone(), + })); } ToolHandlerKind::ResumeAgentV1 => { - builder.register_handler(handler.name, Arc::new(ResumeAgentHandler)); + builder.register_handler(Arc::new(ResumeAgentHandler)); } ToolHandlerKind::SendInputV1 => { - builder.register_handler(handler.name, Arc::new(SendInputHandler)); + builder.register_handler(Arc::new(SendInputHandler)); } ToolHandlerKind::SendMessageV2 => { - builder.register_handler(handler.name, Arc::new(SendMessageHandlerV2)); + builder.register_handler(Arc::new(SendMessageHandlerV2)); } ToolHandlerKind::Shell => { - builder.register_handler(handler.name, shell_handler.clone()); + builder.register_handler(Arc::new(ShellHandler)); } ToolHandlerKind::ShellCommand => { - builder.register_handler(handler.name, shell_command_handler.clone()); + builder.register_handler(Arc::new(ShellCommandHandler::from( + config.shell_command_backend, + ))); + } + ToolHandlerKind::SpawnAgentsOnCsv => { + builder.register_handler(Arc::new(SpawnAgentsOnCsvHandler)); } ToolHandlerKind::SpawnAgentV1 => { - builder.register_handler(handler.name, Arc::new(SpawnAgentHandler)); + builder.register_handler(Arc::new(SpawnAgentHandler)); } ToolHandlerKind::SpawnAgentV2 => { - builder.register_handler(handler.name, Arc::new(SpawnAgentHandlerV2)); + builder.register_handler(Arc::new(SpawnAgentHandlerV2)); } ToolHandlerKind::TestSync => { - builder.register_handler(handler.name, Arc::new(TestSyncHandler)); + builder.register_handler(Arc::new(TestSyncHandler)); } ToolHandlerKind::ToolSearch => { - if tool_search_handler.is_none() { - let entries = build_tool_search_entries_for_config( - config, - deferred_mcp_tools.as_ref(), - &deferred_dynamic_tools, - ); - tool_search_handler = Some(Arc::new(ToolSearchHandler::new(entries))); - } - if let Some(tool_search_handler) = tool_search_handler.as_ref() { - builder.register_handler(handler.name, tool_search_handler.clone()); - } + let entries = build_tool_search_entries_for_config( + config, + deferred_mcp_tools.as_ref(), + &deferred_dynamic_tools, + ); + builder.register_handler(Arc::new(ToolSearchHandler::new(entries))); } ToolHandlerKind::RequestPluginInstall => { - builder.register_handler(handler.name, request_plugin_install_handler.clone()); + builder.register_handler(Arc::new(RequestPluginInstallHandler)); } - ToolHandlerKind::UnifiedExec => { - builder.register_handler(handler.name, unified_exec_handler.clone()); + ToolHandlerKind::UpdateGoal => { + builder.register_handler(Arc::new(UpdateGoalHandler)); } ToolHandlerKind::ViewImage => { - builder.register_handler(handler.name, view_image_handler.clone()); + builder.register_handler(Arc::new(ViewImageHandler)); } ToolHandlerKind::WaitAgentV1 => { - builder.register_handler(handler.name, Arc::new(WaitAgentHandler)); + builder.register_handler(Arc::new(WaitAgentHandler)); } ToolHandlerKind::WaitAgentV2 => { - builder.register_handler(handler.name, Arc::new(WaitAgentHandlerV2)); + builder.register_handler(Arc::new(WaitAgentHandlerV2)); + } + ToolHandlerKind::WriteStdin => { + builder.register_handler(Arc::new(WriteStdinHandler)); } } } if let Some(deferred_mcp_tools) = deferred_mcp_tools.as_ref() { - for (name, _) in deferred_mcp_tools.iter().filter(|(name, _)| { + for (_, tool) in deferred_mcp_tools.iter().filter(|(name, _)| { !mcp_tools .as_ref() .is_some_and(|tools| tools.contains_key(*name)) }) { - builder.register_handler(name.clone(), mcp_handler.clone()); + builder.register_handler(Arc::new(McpHandler::new(tool.canonical_tool_name()))); } } @@ -326,7 +339,7 @@ pub(crate) fn build_specs_with_discoverable_tools( }; builder.push_spec(spec); } - builder.register_handler(unavailable_tool, unavailable_tool_handler.clone()); + builder.register_handler(Arc::new(UnavailableToolHandler::new(unavailable_tool))); } builder } diff --git a/codex-rs/core/src/tools/tool_dispatch_trace_tests.rs b/codex-rs/core/src/tools/tool_dispatch_trace_tests.rs index 5f1181655..99d90c345 100644 --- a/codex-rs/core/src/tools/tool_dispatch_trace_tests.rs +++ b/codex-rs/core/src/tools/tool_dispatch_trace_tests.rs @@ -26,12 +26,17 @@ use crate::tools::registry::ToolKind; use crate::tools::registry::ToolRegistry; use crate::turn_diff_tracker::TurnDiffTracker; -#[derive(Default)] -struct TestHandler; +struct TestHandler { + tool_name: codex_tools::ToolName, +} impl ToolHandler for TestHandler { type Output = FunctionToolOutput; + fn tool_name(&self) -> codex_tools::ToolName { + self.tool_name.clone() + } + fn kind(&self) -> ToolKind { ToolKind::Function } @@ -53,10 +58,9 @@ async fn dispatch_lifecycle_trace_records_direct_and_code_mode_requesters() -> a "await tools.test_tool({})", ); - let registry = ToolRegistry::with_handler_for_test( - codex_tools::ToolName::plain("test_tool"), - Arc::new(TestHandler), - ); + let registry = ToolRegistry::with_handler_for_test(Arc::new(TestHandler { + tool_name: codex_tools::ToolName::plain("test_tool"), + })); let session = Arc::new(session); let turn = Arc::new(turn); @@ -165,10 +169,9 @@ async fn dispatch_lifecycle_trace_records_incompatible_payload_failures() -> any let (mut session, turn) = make_session_and_context().await; attach_test_trace(&mut session, &turn, temp.path())?; - let registry = ToolRegistry::with_handler_for_test( - codex_tools::ToolName::plain("test_tool"), - Arc::new(TestHandler), - ); + let registry = ToolRegistry::with_handler_for_test(Arc::new(TestHandler { + tool_name: codex_tools::ToolName::plain("test_tool"), + })); let session = Arc::new(session); let turn = Arc::new(turn); @@ -200,10 +203,7 @@ async fn missing_code_mode_wait_traces_only_the_wait_tool_call() -> anyhow::Resu let (mut session, turn) = make_session_and_context().await; attach_test_trace(&mut session, &turn, temp.path())?; - let registry = ToolRegistry::with_handler_for_test( - codex_tools::ToolName::plain(WAIT_TOOL_NAME), - Arc::new(CodeModeWaitHandler), - ); + let registry = ToolRegistry::with_handler_for_test(Arc::new(CodeModeWaitHandler)); let session = Arc::new(session); let turn = Arc::new(turn); diff --git a/codex-rs/tools/src/tool_registry_plan.rs b/codex-rs/tools/src/tool_registry_plan.rs index 687f99fcd..c777bd412 100644 --- a/codex-rs/tools/src/tool_registry_plan.rs +++ b/codex-rs/tools/src/tool_registry_plan.rs @@ -172,8 +172,8 @@ pub fn build_tool_registry_plan( /*supports_parallel_tool_calls*/ false, config.code_mode_enabled, ); - plan.register_handler("exec_command", ToolHandlerKind::UnifiedExec); - plan.register_handler("write_stdin", ToolHandlerKind::UnifiedExec); + plan.register_handler("exec_command", ToolHandlerKind::ExecCommand); + plan.register_handler("write_stdin", ToolHandlerKind::WriteStdin); } ConfigShellToolType::Disabled => {} ConfigShellToolType::ShellCommand => { @@ -193,8 +193,8 @@ pub fn build_tool_registry_plan( && config.shell_type != ConfigShellToolType::Disabled { plan.register_handler("shell", ToolHandlerKind::Shell); - plan.register_handler("container.exec", ToolHandlerKind::Shell); - plan.register_handler("local_shell", ToolHandlerKind::Shell); + plan.register_handler("container.exec", ToolHandlerKind::ContainerExec); + plan.register_handler("local_shell", ToolHandlerKind::LocalShell); plan.register_handler("shell_command", ToolHandlerKind::ShellCommand); } @@ -214,9 +214,12 @@ pub fn build_tool_registry_plan( /*supports_parallel_tool_calls*/ true, config.code_mode_enabled, ); - plan.register_handler("list_mcp_resources", ToolHandlerKind::McpResource); - plan.register_handler("list_mcp_resource_templates", ToolHandlerKind::McpResource); - plan.register_handler("read_mcp_resource", ToolHandlerKind::McpResource); + plan.register_handler("list_mcp_resources", ToolHandlerKind::ListMcpResources); + plan.register_handler( + "list_mcp_resource_templates", + ToolHandlerKind::ListMcpResourceTemplates, + ); + plan.register_handler("read_mcp_resource", ToolHandlerKind::ReadMcpResource); } plan.push_spec( @@ -231,19 +234,19 @@ pub fn build_tool_registry_plan( /*supports_parallel_tool_calls*/ false, config.code_mode_enabled, ); - plan.register_handler("get_goal", ToolHandlerKind::Goal); + plan.register_handler("get_goal", ToolHandlerKind::GetGoal); plan.push_spec( create_create_goal_tool(), /*supports_parallel_tool_calls*/ false, config.code_mode_enabled, ); - plan.register_handler("create_goal", ToolHandlerKind::Goal); + plan.register_handler("create_goal", ToolHandlerKind::CreateGoal); plan.push_spec( create_update_goal_tool(), /*supports_parallel_tool_calls*/ false, config.code_mode_enabled, ); - plan.register_handler("update_goal", ToolHandlerKind::Goal); + plan.register_handler("update_goal", ToolHandlerKind::UpdateGoal); } plan.push_spec( @@ -493,14 +496,17 @@ pub fn build_tool_registry_plan( /*supports_parallel_tool_calls*/ false, config.code_mode_enabled, ); - plan.register_handler("spawn_agents_on_csv", ToolHandlerKind::AgentJobs); + plan.register_handler("spawn_agents_on_csv", ToolHandlerKind::SpawnAgentsOnCsv); if config.agent_jobs_worker_tools { plan.push_spec( create_report_agent_job_result_tool(), /*supports_parallel_tool_calls*/ false, config.code_mode_enabled, ); - plan.register_handler("report_agent_job_result", ToolHandlerKind::AgentJobs); + plan.register_handler( + "report_agent_job_result", + ToolHandlerKind::ReportAgentJobResult, + ); } } diff --git a/codex-rs/tools/src/tool_registry_plan_types.rs b/codex-rs/tools/src/tool_registry_plan_types.rs index 260194253..0212cb53d 100644 --- a/codex-rs/tools/src/tool_registry_plan_types.rs +++ b/codex-rs/tools/src/tool_registry_plan_types.rs @@ -10,19 +10,26 @@ use std::collections::HashMap; #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum ToolHandlerKind { - AgentJobs, ApplyPatch, CloseAgentV1, CloseAgentV2, CodeModeExecute, CodeModeWait, + ContainerExec, + CreateGoal, DynamicTool, + ExecCommand, FollowupTaskV2, - Goal, + GetGoal, ListAgentsV2, + ListMcpResourceTemplates, + ListMcpResources, + LocalShell, Mcp, - McpResource, Plan, + ReadMcpResource, + ReportAgentJobResult, + RequestPluginInstall, RequestPermissions, RequestUserInput, ResumeAgentV1, @@ -30,15 +37,16 @@ pub enum ToolHandlerKind { SendMessageV2, Shell, ShellCommand, + SpawnAgentsOnCsv, SpawnAgentV1, SpawnAgentV2, TestSync, ToolSearch, - RequestPluginInstall, - UnifiedExec, + UpdateGoal, ViewImage, WaitAgentV1, WaitAgentV2, + WriteStdin, } #[derive(Debug, Clone, PartialEq, Eq)]