diff --git a/crates/arc-agent/src/cli.rs b/crates/arc-agent/src/cli.rs index bff067b84..f628bcdd5 100644 --- a/crates/arc-agent/src/cli.rs +++ b/crates/arc-agent/src/cli.rs @@ -1,7 +1,8 @@ +use crate::config::ToolApprovalFn; use crate::{ subagent::{SessionFactory, SubAgentManager}, AgentEvent, AnthropicProfile, GeminiProfile, LocalSandbox, OpenAiProfile, ProviderProfile, - Session, SessionConfig, ToolApprovalFn, Turn, + Session, SessionConfig, Turn, }; use arc_llm::client::Client; use arc_llm::provider::{ModelId, Provider}; @@ -416,9 +417,11 @@ pub async fn run_with_args_and_client( let permissions = args.permissions.unwrap_or(PermissionLevel::ReadWrite); let is_interactive = std::io::stdin().is_terminal() && !args.auto_approve; let tool_approval = build_tool_approval(permissions, is_interactive, styles); + let tool_hooks: Arc = + Arc::new(crate::config::ToolApprovalAdapter(tool_approval)); let config = SessionConfig { - tool_approval: Some(tool_approval), + tool_hooks: Some(tool_hooks.clone()), skill_dirs: args.skills_dir.map(|d| vec![d]), mcp_servers, ..SessionConfig::default() @@ -432,7 +435,7 @@ pub async fn run_with_args_and_client( let factory_client = client.clone(); let factory_model = model.to_string(); let factory_env = Arc::clone(&env); - let factory_approval = config.tool_approval.clone(); + let factory_hooks = config.tool_hooks.clone(); let factory: SessionFactory = Arc::new(move || { let child_summarizer = build_summarizer(provider, Some(factory_client.clone())); let child_profile: Arc = match provider { @@ -458,7 +461,7 @@ pub async fn run_with_args_and_client( child_profile, Arc::clone(&factory_env), SessionConfig { - tool_approval: factory_approval.clone(), + tool_hooks: factory_hooks.clone(), ..SessionConfig::default() }, ) diff --git a/crates/arc-agent/src/config.rs b/crates/arc-agent/src/config.rs index cedd89201..672fea590 100644 --- a/crates/arc-agent/src/config.rs +++ b/crates/arc-agent/src/config.rs @@ -8,6 +8,55 @@ use arc_mcp::config::McpServerConfig; /// `Err(message)` to deny with the given message. pub type ToolApprovalFn = Arc Result<(), String> + Send + Sync>; +/// Decision returned by a [`ToolHookCallback`] before a tool executes. +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub enum ToolHookDecision { + /// Allow the tool call to proceed. + #[default] + Proceed, + /// Block the tool call with the given reason. + Block { reason: String }, +} + +/// Async callback trait invoked around tool execution. +#[async_trait::async_trait] +pub trait ToolHookCallback: Send + Sync { + /// Called before a tool executes. Return [`ToolHookDecision::Proceed`] to + /// allow or [`ToolHookDecision::Block`] to deny. + async fn pre_tool_use( + &self, + tool_name: &str, + tool_input: &serde_json::Value, + ) -> ToolHookDecision; + + /// Called after a tool executes successfully. + async fn post_tool_use(&self, tool_name: &str, tool_call_id: &str, tool_output: &str); + + /// Called after a tool execution fails. + async fn post_tool_use_failure(&self, tool_name: &str, tool_call_id: &str, error: &str); +} + +/// Adapter that wraps a [`ToolApprovalFn`] and implements [`ToolHookCallback`]. +pub struct ToolApprovalAdapter(pub ToolApprovalFn); + +#[async_trait::async_trait] +impl ToolHookCallback for ToolApprovalAdapter { + async fn pre_tool_use( + &self, + tool_name: &str, + tool_input: &serde_json::Value, + ) -> ToolHookDecision { + match (self.0)(tool_name, tool_input) { + Ok(()) => ToolHookDecision::Proceed, + Err(reason) => ToolHookDecision::Block { reason }, + } + } + + async fn post_tool_use(&self, _tool_name: &str, _tool_call_id: &str, _tool_output: &str) {} + + async fn post_tool_use_failure(&self, _tool_name: &str, _tool_call_id: &str, _error: &str) {} +} + #[derive(Clone)] pub struct SessionConfig { pub max_turns: usize, @@ -25,7 +74,8 @@ pub struct SessionConfig { pub max_subagent_depth: usize, pub git_root: Option, pub user_instructions: Option, - pub tool_approval: Option, + /// Async hook callbacks invoked around tool execution. + pub tool_hooks: Option>, pub enable_context_compaction: bool, pub compaction_threshold_percent: usize, pub compaction_preserve_turns: usize, @@ -58,8 +108,8 @@ impl std::fmt::Debug for SessionConfig { .field("git_root", &self.git_root) .field("user_instructions", &self.user_instructions) .field( - "tool_approval", - &self.tool_approval.as_ref().map(|_| ""), + "tool_hooks", + &self.tool_hooks.as_ref().map(|_| ""), ) .field("enable_context_compaction", &self.enable_context_compaction) .field( @@ -90,7 +140,7 @@ impl Default for SessionConfig { max_subagent_depth: 1, git_root: None, user_instructions: None, - tool_approval: None, + tool_hooks: None, enable_context_compaction: true, compaction_threshold_percent: 80, compaction_preserve_turns: 6, @@ -142,4 +192,41 @@ mod tests { assert_eq!(config.reasoning_effort, Some("high".into())); assert_eq!(config.max_tool_rounds_per_input, 200); } + + #[test] + fn tool_hook_decision_default_is_proceed() { + assert_eq!(ToolHookDecision::default(), ToolHookDecision::Proceed); + } + + #[tokio::test] + async fn tool_approval_adapter_allows() { + let approval: ToolApprovalFn = Arc::new(|_name, _args| Ok(())); + let adapter = ToolApprovalAdapter(approval); + let decision = adapter.pre_tool_use("shell", &serde_json::json!({})).await; + assert_eq!(decision, ToolHookDecision::Proceed); + } + + #[tokio::test] + async fn tool_approval_adapter_blocks() { + let approval: ToolApprovalFn = Arc::new(|_name, _args| Err("denied".to_string())); + let adapter = ToolApprovalAdapter(approval); + let decision = adapter.pre_tool_use("shell", &serde_json::json!({})).await; + assert_eq!( + decision, + ToolHookDecision::Block { + reason: "denied".to_string() + } + ); + } + + #[tokio::test] + async fn tool_approval_adapter_post_is_noop() { + let approval: ToolApprovalFn = Arc::new(|_name, _args| Ok(())); + let adapter = ToolApprovalAdapter(approval); + // These should not panic + adapter.post_tool_use("shell", "call_1", "output").await; + adapter + .post_tool_use_failure("shell", "call_1", "error") + .await; + } } diff --git a/crates/arc-agent/src/lib.rs b/crates/arc-agent/src/lib.rs index e6d5f5cdf..eeb097729 100644 --- a/crates/arc-agent/src/lib.rs +++ b/crates/arc-agent/src/lib.rs @@ -27,7 +27,7 @@ pub mod types; pub mod v4a_patch; pub use arc_mcp::config::McpServerConfig; -pub use config::{SessionConfig, ToolApprovalFn}; +pub use config::{SessionConfig, ToolApprovalAdapter, ToolHookCallback, ToolHookDecision}; #[cfg(feature = "docker")] pub use docker_sandbox::{DockerSandbox, DockerSandboxConfig}; pub use error::AgentError; diff --git a/crates/arc-agent/src/session.rs b/crates/arc-agent/src/session.rs index 126db8521..05570ccd1 100644 --- a/crates/arc-agent/src/session.rs +++ b/crates/arc-agent/src/session.rs @@ -610,7 +610,7 @@ impl Session { self.provider_profile.supports_parallel_tool_calls(), self.provider_profile.tool_registry(), self.sandbox.clone(), - self.config.tool_approval.as_ref(), + self.config.tool_hooks.as_ref(), &self.cancel_token, &self.config, &self.event_emitter, @@ -1489,7 +1489,9 @@ mod tests { ]; let config = SessionConfig { - tool_approval: Some(Arc::new(|_name, _args| Err("denied by policy".to_string()))), + tool_hooks: Some(Arc::new(crate::config::ToolApprovalAdapter(Arc::new( + |_name, _args| Err("denied by policy".to_string()), + )))), ..Default::default() }; @@ -1528,7 +1530,9 @@ mod tests { ]; let config = SessionConfig { - tool_approval: Some(Arc::new(|_name, _args| Ok(()))), + tool_hooks: Some(Arc::new(crate::config::ToolApprovalAdapter(Arc::new( + |_name, _args| Ok(()), + )))), ..Default::default() }; @@ -1562,10 +1566,12 @@ mod tests { ]; let config = SessionConfig { - tool_approval: Some(Arc::new(move |name, args| { - *captured_clone.lock().unwrap() = Some((name.to_string(), args.clone())); - Ok(()) - })), + tool_hooks: Some(Arc::new(crate::config::ToolApprovalAdapter(Arc::new( + move |name, args| { + *captured_clone.lock().unwrap() = Some((name.to_string(), args.clone())); + Ok(()) + }, + )))), ..Default::default() }; @@ -1591,7 +1597,7 @@ mod tests { ]; let config = SessionConfig { - tool_approval: None, + tool_hooks: None, ..Default::default() }; @@ -1622,7 +1628,9 @@ mod tests { ]; let config = SessionConfig { - tool_approval: Some(Arc::new(|_name, _args| Err("not allowed".to_string()))), + tool_hooks: Some(Arc::new(crate::config::ToolApprovalAdapter(Arc::new( + |_name, _args| Err("not allowed".to_string()), + )))), ..Default::default() }; diff --git a/crates/arc-agent/src/tool_execution.rs b/crates/arc-agent/src/tool_execution.rs index 97f96e8da..a4514aeda 100644 --- a/crates/arc-agent/src/tool_execution.rs +++ b/crates/arc-agent/src/tool_execution.rs @@ -1,4 +1,4 @@ -use crate::config::{SessionConfig, ToolApprovalFn}; +use crate::config::{SessionConfig, ToolHookCallback, ToolHookDecision}; use crate::event::EventEmitter; use crate::sandbox::Sandbox; use crate::tool_registry::ToolRegistry; @@ -8,6 +8,7 @@ use arc_llm::types::ToolResult; use std::collections::HashMap; use std::sync::Arc; use tokio_util::sync::CancellationToken; +use tracing::debug; /// Execute tool calls, choosing parallel or sequential based on `parallel` flag. #[allow(clippy::too_many_arguments)] @@ -16,7 +17,7 @@ pub async fn execute_tool_calls( parallel: bool, registry: &ToolRegistry, env: Arc, - tool_approval: Option<&ToolApprovalFn>, + tool_hooks: Option<&Arc>, cancel_token: &CancellationToken, config: &SessionConfig, emitter: &EventEmitter, @@ -28,7 +29,7 @@ pub async fn execute_tool_calls( tool_calls, registry, env, - tool_approval, + tool_hooks, cancel_token, config, emitter, @@ -41,7 +42,7 @@ pub async fn execute_tool_calls( tool_calls, registry, env, - tool_approval, + tool_hooks, cancel_token, config, emitter, @@ -57,7 +58,7 @@ async fn execute_tool_calls_sequential( tool_calls: &[arc_llm::types::ToolCall], registry: &ToolRegistry, env: Arc, - tool_approval: Option<&ToolApprovalFn>, + tool_hooks: Option<&Arc>, cancel_token: &CancellationToken, config: &SessionConfig, emitter: &EventEmitter, @@ -75,7 +76,7 @@ async fn execute_tool_calls_sequential( tc, registry, env.clone(), - tool_approval, + tool_hooks, cancel_token.child_token(), config, emitter, @@ -93,7 +94,7 @@ async fn execute_tool_calls_parallel( tool_calls: &[arc_llm::types::ToolCall], registry: &ToolRegistry, env: Arc, - tool_approval: Option<&ToolApprovalFn>, + tool_hooks: Option<&Arc>, cancel_token: &CancellationToken, config: &SessionConfig, emitter: &EventEmitter, @@ -110,7 +111,7 @@ async fn execute_tool_calls_parallel( let cancel_token = cancel_token.clone(); let tc = tc.clone(); let session_id = session_id.to_owned(); - let tool_approval = tool_approval.cloned(); + let tool_hooks = tool_hooks.cloned(); let tool_env = tool_env.clone(); // Look up the tool before spawning since ToolRegistry is not Send. let registered_tool = registry.get(&tc.name).cloned(); @@ -119,7 +120,7 @@ async fn execute_tool_calls_parallel( &tc, registered_tool.as_ref(), env, - tool_approval.as_ref(), + tool_hooks.as_ref(), cancel_token.child_token(), &config, &emitter, @@ -140,7 +141,7 @@ pub async fn execute_and_emit_one_tool( tc: &arc_llm::types::ToolCall, registry: &ToolRegistry, env: Arc, - tool_approval: Option<&ToolApprovalFn>, + tool_hooks: Option<&Arc>, cancel_token: CancellationToken, config: &SessionConfig, emitter: &EventEmitter, @@ -151,7 +152,7 @@ pub async fn execute_and_emit_one_tool( tc, registry.get(&tc.name), env, - tool_approval, + tool_hooks, cancel_token, config, emitter, @@ -167,7 +168,7 @@ async fn execute_and_emit_one_tool_with_lookup( tc: &arc_llm::types::ToolCall, registered_tool: Option<&crate::tool_registry::RegisteredTool>, env: Arc, - tool_approval: Option<&ToolApprovalFn>, + tool_hooks: Option<&Arc>, cancel_token: CancellationToken, config: &SessionConfig, emitter: &EventEmitter, @@ -183,15 +184,38 @@ async fn execute_and_emit_one_tool_with_lookup( }, ); - let result = execute_one_tool( - tc, - registered_tool, - env, - tool_approval, - cancel_token, - tool_env, - ) - .await; + // Pre-tool-use hook + if let Some(hooks) = tool_hooks { + debug!(tool = %tc.name, hook_event = "pre_tool_use", "Calling tool hook"); + let start = std::time::Instant::now(); + let decision = hooks.pre_tool_use(&tc.name, &tc.arguments).await; + let elapsed = start.elapsed().as_millis() as u64; + debug!(tool = %tc.name, hook_event = "pre_tool_use", ?decision, duration_ms = elapsed, "Tool hook complete"); + + if let ToolHookDecision::Block { reason } = decision { + let result = ToolResult::error(&tc.id, &reason); + + emitter.emit( + session_id.to_owned(), + AgentEvent::ToolCallOutputDelta { + delta: result.content.to_string(), + }, + ); + emitter.emit( + session_id.to_owned(), + AgentEvent::ToolCallCompleted { + tool_name: tc.name.clone(), + tool_call_id: tc.id.clone(), + output: result.content.clone(), + is_error: true, + }, + ); + + return truncate_tool_result(&result, &tc.name, config); + } + } + + let result = execute_one_tool(tc, registered_tool, env, cancel_token, tool_env).await; emitter.emit( session_id.to_owned(), @@ -210,6 +234,29 @@ async fn execute_and_emit_one_tool_with_lookup( }, ); + // Post-tool-use hooks + if let Some(hooks) = tool_hooks { + let fallback; + let content_str = match result.content.as_str() { + Some(s) => s, + None => { + fallback = result.content.to_string(); + &fallback + } + }; + if result.is_error { + debug!(tool = %tc.name, hook_event = "post_tool_use_failure", "Calling tool hook"); + hooks + .post_tool_use_failure(&tc.name, &tc.id, content_str) + .await; + debug!(tool = %tc.name, hook_event = "post_tool_use_failure", "Tool hook complete"); + } else { + debug!(tool = %tc.name, hook_event = "post_tool_use", "Calling tool hook"); + hooks.post_tool_use(&tc.name, &tc.id, content_str).await; + debug!(tool = %tc.name, hook_event = "post_tool_use", "Tool hook complete"); + } + } + truncate_tool_result(&result, &tc.name, config) } @@ -218,16 +265,9 @@ async fn execute_one_tool( tc: &arc_llm::types::ToolCall, registered_tool: Option<&crate::tool_registry::RegisteredTool>, env: Arc, - tool_approval: Option<&ToolApprovalFn>, cancel_token: CancellationToken, tool_env: Option<&HashMap>, ) -> ToolResult { - if let Some(approval_fn) = tool_approval { - if let Err(denial_message) = approval_fn(&tc.name, &tc.arguments) { - return ToolResult::error(&tc.id, denial_message); - } - } - match registered_tool { Some(tool) => { if let Err(validation_error) = @@ -300,3 +340,266 @@ pub fn validate_tool_args( )) } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::config::{ToolHookCallback, ToolHookDecision}; + use crate::event::EventEmitter; + use crate::tool_registry::{RegisteredTool, ToolContext, ToolRegistry}; + use arc_llm::types::{ToolCall, ToolDefinition}; + use std::sync::Mutex; + + fn make_echo_tool() -> RegisteredTool { + RegisteredTool { + definition: ToolDefinition { + name: "echo".to_string(), + description: "Echo input".to_string(), + parameters: serde_json::json!({ + "type": "object", + "properties": { + "text": {"type": "string"} + }, + "required": ["text"] + }), + }, + executor: Arc::new(|args: serde_json::Value, _ctx: ToolContext| { + Box::pin(async move { + let text = args["text"].as_str().unwrap_or("").to_string(); + Ok(format!("echo: {text}")) + }) + }), + } + } + + fn make_fail_tool() -> RegisteredTool { + RegisteredTool { + definition: ToolDefinition { + name: "fail_tool".to_string(), + description: "Always fails".to_string(), + parameters: serde_json::json!({}), + }, + executor: Arc::new(|_args: serde_json::Value, _ctx: ToolContext| { + Box::pin(async move { Err("tool failed".to_string()) }) + }), + } + } + + fn make_tool_call(name: &str, id: &str, args: serde_json::Value) -> ToolCall { + ToolCall { + id: id.to_string(), + name: name.to_string(), + arguments: args, + raw_arguments: None, + provider_metadata: None, + } + } + + struct MockHookCallback { + pre_decision: ToolHookDecision, + post_calls: Arc>>, + post_failure_calls: Arc>>, + } + + impl MockHookCallback { + fn new(decision: ToolHookDecision) -> Self { + Self { + pre_decision: decision, + post_calls: Arc::new(Mutex::new(Vec::new())), + post_failure_calls: Arc::new(Mutex::new(Vec::new())), + } + } + } + + #[async_trait::async_trait] + impl ToolHookCallback for MockHookCallback { + async fn pre_tool_use( + &self, + _tool_name: &str, + _tool_input: &serde_json::Value, + ) -> ToolHookDecision { + self.pre_decision.clone() + } + + async fn post_tool_use(&self, tool_name: &str, tool_call_id: &str, tool_output: &str) { + self.post_calls.lock().unwrap().push(( + tool_name.to_string(), + tool_call_id.to_string(), + tool_output.to_string(), + )); + } + + async fn post_tool_use_failure(&self, tool_name: &str, tool_call_id: &str, error: &str) { + self.post_failure_calls.lock().unwrap().push(( + tool_name.to_string(), + tool_call_id.to_string(), + error.to_string(), + )); + } + } + + fn make_sandbox() -> Arc { + Arc::new(crate::local_sandbox::LocalSandbox::new( + std::env::current_dir().unwrap(), + )) + } + + #[tokio::test] + async fn pre_tool_use_hook_blocks_execution() { + let mut registry = ToolRegistry::new(); + registry.register(make_echo_tool()); + + let hooks: Arc = + Arc::new(MockHookCallback::new(ToolHookDecision::Block { + reason: "blocked by hook".to_string(), + })); + + let tc = make_tool_call("echo", "call_1", serde_json::json!({"text": "hello"})); + let emitter = EventEmitter::new(); + let config = SessionConfig::default(); + + let result = execute_and_emit_one_tool( + &tc, + ®istry, + make_sandbox(), + Some(&hooks), + CancellationToken::new(), + &config, + &emitter, + "test-session", + None, + ) + .await; + + assert!(result.is_error); + let content = result.content.as_str().unwrap(); + assert!(content.contains("blocked by hook")); + } + + #[tokio::test] + async fn pre_tool_use_hook_proceeds() { + let mut registry = ToolRegistry::new(); + registry.register(make_echo_tool()); + + let hooks: Arc = + Arc::new(MockHookCallback::new(ToolHookDecision::Proceed)); + + let tc = make_tool_call("echo", "call_1", serde_json::json!({"text": "hello"})); + let emitter = EventEmitter::new(); + let config = SessionConfig::default(); + + let result = execute_and_emit_one_tool( + &tc, + ®istry, + make_sandbox(), + Some(&hooks), + CancellationToken::new(), + &config, + &emitter, + "test-session", + None, + ) + .await; + + assert!(!result.is_error); + let content = result.content.to_string(); + assert!(content.contains("echo: hello")); + } + + #[tokio::test] + async fn post_tool_use_hook_fires_on_success() { + let mut registry = ToolRegistry::new(); + registry.register(make_echo_tool()); + + let mock = Arc::new(MockHookCallback::new(ToolHookDecision::Proceed)); + let hooks: Arc = mock.clone(); + + let tc = make_tool_call("echo", "call_1", serde_json::json!({"text": "hello"})); + let emitter = EventEmitter::new(); + let config = SessionConfig::default(); + + execute_and_emit_one_tool( + &tc, + ®istry, + make_sandbox(), + Some(&hooks), + CancellationToken::new(), + &config, + &emitter, + "test-session", + None, + ) + .await; + + let calls = mock.post_calls.lock().unwrap(); + assert_eq!(calls.len(), 1); + assert_eq!(calls[0].0, "echo"); + assert_eq!(calls[0].1, "call_1"); + assert!(calls[0].2.contains("echo: hello")); + + let failure_calls = mock.post_failure_calls.lock().unwrap(); + assert!(failure_calls.is_empty()); + } + + #[tokio::test] + async fn post_tool_use_failure_hook_fires_on_error() { + let mut registry = ToolRegistry::new(); + registry.register(make_fail_tool()); + + let mock = Arc::new(MockHookCallback::new(ToolHookDecision::Proceed)); + let hooks: Arc = mock.clone(); + + let tc = make_tool_call("fail_tool", "call_1", serde_json::json!({})); + let emitter = EventEmitter::new(); + let config = SessionConfig::default(); + + execute_and_emit_one_tool( + &tc, + ®istry, + make_sandbox(), + Some(&hooks), + CancellationToken::new(), + &config, + &emitter, + "test-session", + None, + ) + .await; + + let failure_calls = mock.post_failure_calls.lock().unwrap(); + assert_eq!(failure_calls.len(), 1); + assert_eq!(failure_calls[0].0, "fail_tool"); + assert_eq!(failure_calls[0].1, "call_1"); + assert!(failure_calls[0].2.contains("tool failed")); + + let calls = mock.post_calls.lock().unwrap(); + assert!(calls.is_empty()); + } + + #[tokio::test] + async fn no_hooks_skips_all_callbacks() { + let mut registry = ToolRegistry::new(); + registry.register(make_echo_tool()); + + let tc = make_tool_call("echo", "call_1", serde_json::json!({"text": "hello"})); + let emitter = EventEmitter::new(); + let config = SessionConfig::default(); + + let result = execute_and_emit_one_tool( + &tc, + ®istry, + make_sandbox(), + None, + CancellationToken::new(), + &config, + &emitter, + "test-session", + None, + ) + .await; + + assert!(!result.is_error); + let content = result.content.to_string(); + assert!(content.contains("echo: hello")); + } +} diff --git a/crates/arc-workflows/src/cli/backend.rs b/crates/arc-workflows/src/cli/backend.rs index 2f7fe9542..8ae83871f 100644 --- a/crates/arc-workflows/src/cli/backend.rs +++ b/crates/arc-workflows/src/cli/backend.rs @@ -133,8 +133,17 @@ impl AgentApiBackend { &self, node: &Node, sandbox: &Arc, + tool_hooks: Option>, ) -> Result { - Self::create_session_for(&self.model, self.provider, node, sandbox, &self.env).await + Self::create_session_for( + &self.model, + self.provider, + node, + sandbox, + &self.env, + tool_hooks, + ) + .await } async fn create_session_for( @@ -143,6 +152,7 @@ impl AgentApiBackend { node: &Node, sandbox: &Arc, env: &HashMap, + tool_hooks: Option>, ) -> Result { let client = Client::from_env() .await @@ -153,6 +163,7 @@ impl AgentApiBackend { let config = SessionConfig { max_tokens: node.max_tokens(), reasoning_effort: Some(node.reasoning_effort().to_string()), + tool_hooks, ..SessionConfig::default() }; @@ -370,6 +381,7 @@ impl CodergenBackend for AgentApiBackend { emitter: &Arc, stage_dir: &std::path::Path, sandbox: &Arc, + tool_hooks: Option>, ) -> Result { let fidelity = context.fidelity(); let reuse_key = if fidelity == crate::context::keys::Fidelity::Full { @@ -384,10 +396,18 @@ impl CodergenBackend for AgentApiBackend { if let Some(s) = existing { (s, true) } else { - (self.create_session(node, sandbox).await?, false) + ( + self.create_session(node, sandbox, tool_hooks.clone()) + .await?, + false, + ) } } else { - (self.create_session(node, sandbox).await?, false) + ( + self.create_session(node, sandbox, tool_hooks.clone()) + .await?, + false, + ) }; tracing::debug!( @@ -460,6 +480,7 @@ impl CodergenBackend for AgentApiBackend { node, sandbox, &self.env, + tool_hooks.clone(), ) .await { diff --git a/crates/arc-workflows/src/cli/cli_backend.rs b/crates/arc-workflows/src/cli/cli_backend.rs index fc0ba12f6..62220503a 100644 --- a/crates/arc-workflows/src/cli/cli_backend.rs +++ b/crates/arc-workflows/src/cli/cli_backend.rs @@ -440,6 +440,7 @@ impl CodergenBackend for AgentCliBackend { emitter: &Arc, stage_dir: &Path, sandbox: &Arc, + _tool_hooks: Option>, ) -> Result { // 1. Snapshot git state before the CLI run let files_before = self.detect_changed_files(sandbox).await; @@ -704,17 +705,18 @@ impl CodergenBackend for BackendRouter { emitter: &Arc, stage_dir: &Path, sandbox: &Arc, + tool_hooks: Option>, ) -> Result { if self.should_use_cli(node) { self.cli_backend .run( - node, prompt, context, thread_id, emitter, stage_dir, sandbox, + node, prompt, context, thread_id, emitter, stage_dir, sandbox, tool_hooks, ) .await } else { self.api_backend .run( - node, prompt, context, thread_id, emitter, stage_dir, sandbox, + node, prompt, context, thread_id, emitter, stage_dir, sandbox, tool_hooks, ) .await } @@ -1133,6 +1135,7 @@ mod tests { _emitter: &Arc, _stage_dir: &Path, _sandbox: &Arc, + _tool_hooks: Option>, ) -> Result { Ok(CodergenResult::Text { text: "stub".to_string(), diff --git a/crates/arc-workflows/src/handler/agent.rs b/crates/arc-workflows/src/handler/agent.rs index 209228ce5..a219edd6d 100644 --- a/crates/arc-workflows/src/handler/agent.rs +++ b/crates/arc-workflows/src/handler/agent.rs @@ -38,6 +38,7 @@ pub trait CodergenBackend: Send + Sync { emitter: &Arc, stage_dir: &Path, sandbox: &Arc, + tool_hooks: Option>, ) -> Result; /// Run a single LLM call with no tools (one_shot mode). @@ -225,6 +226,17 @@ impl Handler for AgentHandler { // 3. Call LLM backend (agent loop) let thread_id = context.thread_id(); + let tool_hooks: Option> = + services.hook_runner.as_ref().map(|hr| { + Arc::new(crate::hook::bridge::WorkflowToolHookCallback { + hook_runner: Arc::clone(hr), + sandbox: Arc::clone(&services.sandbox), + run_id: context.run_id(), + workflow_name: graph.name.clone(), + work_dir: None, + node_id: node.id.clone(), + }) as Arc + }); let (response_text, stage_usage, backend_files_touched) = if let Some(backend) = &self.backend { let result = backend @@ -236,6 +248,7 @@ impl Handler for AgentHandler { &services.emitter, &stage_dir, &services.sandbox, + tool_hooks, ) .await; match result { @@ -494,6 +507,7 @@ mod tests { _emitter: &Arc, _stage_dir: &Path, _sandbox: &Arc, + _tool_hooks: Option>, ) -> Result { Ok(CodergenResult::Text { text: r#"Done. {"outcome": "success", "preferred_next_label": "approve"}"# @@ -603,6 +617,7 @@ mod tests { _emitter: &Arc, _stage_dir: &Path, _sandbox: &Arc, + _tool_hooks: Option>, ) -> Result { *self.captured_thread_id.lock().unwrap() = Some(thread_id.map(String::from)); Ok(CodergenResult::Text { @@ -654,6 +669,7 @@ mod tests { _emitter: &Arc, _stage_dir: &Path, _sandbox: &Arc, + _tool_hooks: Option>, ) -> Result { *self.captured_thread_id.lock().unwrap() = Some(thread_id.map(String::from)); Ok(CodergenResult::Text { @@ -700,6 +716,7 @@ mod tests { _emitter: &Arc, _stage_dir: &Path, _sandbox: &Arc, + _tool_hooks: Option>, ) -> Result { Err(ArcError::handler("Request timed out".to_string())) } @@ -843,6 +860,7 @@ Some text in between. _emitter: &Arc, _stage_dir: &Path, _sandbox: &Arc, + _tool_hooks: Option>, ) -> Result { Err(ArcError::Validation("bad config".to_string())) } @@ -881,6 +899,7 @@ Some text in between. _emitter: &Arc, _stage_dir: &std::path::Path, _sandbox: &Arc, + _tool_hooks: Option>, ) -> Result { *self.captured_prompt.lock().unwrap() = Some(prompt.to_string()); Ok(CodergenResult::Text { @@ -949,6 +968,7 @@ Some text in between. _emitter: &Arc, _stage_dir: &std::path::Path, _sandbox: &Arc, + _tool_hooks: Option>, ) -> Result { *self.captured_prompt.lock().unwrap() = Some(prompt.to_string()); Ok(CodergenResult::Text { diff --git a/crates/arc-workflows/src/handler/fan_in.rs b/crates/arc-workflows/src/handler/fan_in.rs index 438800774..1708e4e1f 100644 --- a/crates/arc-workflows/src/handler/fan_in.rs +++ b/crates/arc-workflows/src/handler/fan_in.rs @@ -228,6 +228,7 @@ async fn llm_evaluate( emitter, &stage_dir, sandbox, + None, ) .await { @@ -435,6 +436,7 @@ mod tests { _emitter: &Arc, _stage_dir: &std::path::Path, _sandbox: &Arc, + _tool_hooks: Option>, ) -> Result { // Return text that contains the ID "branch_b" Ok(CodergenResult::Text { diff --git a/crates/arc-workflows/src/handler/prompt.rs b/crates/arc-workflows/src/handler/prompt.rs index feb98dbb5..5dcda5c8f 100644 --- a/crates/arc-workflows/src/handler/prompt.rs +++ b/crates/arc-workflows/src/handler/prompt.rs @@ -211,6 +211,7 @@ mod tests { _emitter: &Arc, _stage_dir: &Path, _sandbox: &Arc, + _tool_hooks: Option>, ) -> Result { panic!("run() should not be called for prompt handler"); } @@ -272,6 +273,7 @@ mod tests { _emitter: &Arc, _stage_dir: &Path, _sandbox: &Arc, + _tool_hooks: Option>, ) -> Result { panic!("run() should not be called for prompt handler"); } diff --git a/crates/arc-workflows/src/hook/bridge.rs b/crates/arc-workflows/src/hook/bridge.rs new file mode 100644 index 000000000..cfc394271 --- /dev/null +++ b/crates/arc-workflows/src/hook/bridge.rs @@ -0,0 +1,265 @@ +use std::path::PathBuf; +use std::sync::Arc; + +use arc_agent::{Sandbox, ToolHookCallback, ToolHookDecision}; + +use super::runner::HookRunner; +use super::types::{HookContext, HookDecision, HookEvent}; + +/// Bridge between the workflow hook system and the agent tool-hook callback. +/// +/// Created per-node in the workflow engine, capturing the `HookRunner` and +/// context needed to build `HookContext` for tool-level events. +pub struct WorkflowToolHookCallback { + pub hook_runner: Arc, + pub sandbox: Arc, + pub run_id: String, + pub workflow_name: String, + pub work_dir: Option, + pub node_id: String, +} + +impl WorkflowToolHookCallback { + fn base_context(&self, event: HookEvent, tool_name: &str) -> HookContext { + let mut ctx = HookContext::new(event, self.run_id.clone(), self.workflow_name.clone()); + ctx.node_id = Some(self.node_id.clone()); + ctx.tool_name = Some(tool_name.to_string()); + ctx + } + + async fn run_hook(&self, ctx: &HookContext) -> HookDecision { + self.hook_runner + .run(ctx, self.sandbox.clone(), self.work_dir.as_deref()) + .await + } +} + +#[async_trait::async_trait] +impl ToolHookCallback for WorkflowToolHookCallback { + async fn pre_tool_use( + &self, + tool_name: &str, + tool_input: &serde_json::Value, + ) -> ToolHookDecision { + let mut ctx = self.base_context(HookEvent::PreToolUse, tool_name); + ctx.tool_input = Some(tool_input.clone()); + + match self.run_hook(&ctx).await { + HookDecision::Block { reason } => ToolHookDecision::Block { + reason: reason.unwrap_or_else(|| "Blocked by hook".to_string()), + }, + _ => ToolHookDecision::Proceed, + } + } + + async fn post_tool_use(&self, tool_name: &str, tool_call_id: &str, tool_output: &str) { + let mut ctx = self.base_context(HookEvent::PostToolUse, tool_name); + ctx.tool_call_id = Some(tool_call_id.to_string()); + ctx.tool_output = Some(tool_output.to_string()); + + self.run_hook(&ctx).await; + } + + async fn post_tool_use_failure(&self, tool_name: &str, tool_call_id: &str, error: &str) { + let mut ctx = self.base_context(HookEvent::PostToolUseFailure, tool_name); + ctx.tool_call_id = Some(tool_call_id.to_string()); + ctx.error_message = Some(error.to_string()); + + self.run_hook(&ctx).await; + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::hook::config::{HookConfig, HookDefinition}; + use crate::hook::executor::HookExecutor; + use crate::hook::types::{HookContext, HookResult}; + use std::path::Path; + use std::sync::Mutex; + + struct CapturingExecutor { + captured_contexts: Arc>>, + decision: HookDecision, + } + + #[async_trait::async_trait] + impl HookExecutor for CapturingExecutor { + async fn execute( + &self, + _definition: &HookDefinition, + context: &HookContext, + _sandbox: Arc, + _work_dir: Option<&Path>, + ) -> HookResult { + self.captured_contexts.lock().unwrap().push(context.clone()); + HookResult { + hook_name: None, + decision: self.decision.clone(), + duration_ms: 1, + } + } + } + + fn make_hook(event: HookEvent) -> HookDefinition { + HookDefinition { + name: Some("test-hook".into()), + event, + command: Some("echo test".into()), + hook_type: None, + matcher: None, + blocking: None, + timeout_ms: None, + sandbox: Some(false), + } + } + + fn make_sandbox() -> Arc { + Arc::new(arc_agent::LocalSandbox::new( + std::env::current_dir().unwrap(), + )) + } + + fn make_bridge( + hook_runner: Arc, + sandbox: Arc, + ) -> WorkflowToolHookCallback { + WorkflowToolHookCallback { + hook_runner, + sandbox, + run_id: "run-1".into(), + workflow_name: "test-wf".into(), + work_dir: None, + node_id: "plan".into(), + } + } + + #[tokio::test] + async fn pre_tool_use_builds_correct_context() { + let captured = Arc::new(Mutex::new(Vec::new())); + let executor = Arc::new(CapturingExecutor { + captured_contexts: captured.clone(), + decision: HookDecision::Proceed, + }); + let config = HookConfig { + hooks: vec![make_hook(HookEvent::PreToolUse)], + }; + let runner = Arc::new(HookRunner::with_executor(config, executor)); + let sandbox = make_sandbox(); + let bridge = make_bridge(runner, sandbox); + + bridge + .pre_tool_use("shell", &serde_json::json!({"command": "ls"})) + .await; + + let contexts = captured.lock().unwrap(); + assert_eq!(contexts.len(), 1); + assert_eq!(contexts[0].event, HookEvent::PreToolUse); + assert_eq!(contexts[0].tool_name.as_deref(), Some("shell")); + assert_eq!( + contexts[0].tool_input, + Some(serde_json::json!({"command": "ls"})) + ); + assert_eq!(contexts[0].run_id, "run-1"); + assert_eq!(contexts[0].node_id.as_deref(), Some("plan")); + } + + #[tokio::test] + async fn pre_tool_use_maps_block_decision() { + let executor = Arc::new(CapturingExecutor { + captured_contexts: Arc::new(Mutex::new(Vec::new())), + decision: HookDecision::Block { + reason: Some("forbidden".into()), + }, + }); + let config = HookConfig { + hooks: vec![make_hook(HookEvent::PreToolUse)], + }; + let runner = Arc::new(HookRunner::with_executor(config, executor)); + let sandbox = make_sandbox(); + let bridge = make_bridge(runner, sandbox); + + let decision = bridge.pre_tool_use("shell", &serde_json::json!({})).await; + assert_eq!( + decision, + ToolHookDecision::Block { + reason: "forbidden".to_string() + } + ); + } + + #[tokio::test] + async fn pre_tool_use_maps_proceed() { + let executor = Arc::new(CapturingExecutor { + captured_contexts: Arc::new(Mutex::new(Vec::new())), + decision: HookDecision::Proceed, + }); + let config = HookConfig { + hooks: vec![make_hook(HookEvent::PreToolUse)], + }; + let runner = Arc::new(HookRunner::with_executor(config, executor)); + let sandbox = make_sandbox(); + let bridge = make_bridge(runner, sandbox); + + let decision = bridge.pre_tool_use("shell", &serde_json::json!({})).await; + assert_eq!(decision, ToolHookDecision::Proceed); + } + + #[tokio::test] + async fn post_tool_use_builds_context_with_output() { + let captured = Arc::new(Mutex::new(Vec::new())); + let executor = Arc::new(CapturingExecutor { + captured_contexts: captured.clone(), + decision: HookDecision::Proceed, + }); + let config = HookConfig { + hooks: vec![make_hook(HookEvent::PostToolUse)], + }; + let runner = Arc::new(HookRunner::with_executor(config, executor)); + let sandbox = make_sandbox(); + let bridge = make_bridge(runner, sandbox); + + bridge + .post_tool_use("shell", "call_1", "file1.txt\nfile2.txt") + .await; + + let contexts = captured.lock().unwrap(); + assert_eq!(contexts.len(), 1); + assert_eq!(contexts[0].event, HookEvent::PostToolUse); + assert_eq!(contexts[0].tool_name.as_deref(), Some("shell")); + assert_eq!(contexts[0].tool_call_id.as_deref(), Some("call_1")); + assert_eq!( + contexts[0].tool_output.as_deref(), + Some("file1.txt\nfile2.txt") + ); + } + + #[tokio::test] + async fn post_tool_use_failure_builds_context_with_error() { + let captured = Arc::new(Mutex::new(Vec::new())); + let executor = Arc::new(CapturingExecutor { + captured_contexts: captured.clone(), + decision: HookDecision::Proceed, + }); + let config = HookConfig { + hooks: vec![make_hook(HookEvent::PostToolUseFailure)], + }; + let runner = Arc::new(HookRunner::with_executor(config, executor)); + let sandbox = make_sandbox(); + let bridge = make_bridge(runner, sandbox); + + bridge + .post_tool_use_failure("shell", "call_1", "command not found") + .await; + + let contexts = captured.lock().unwrap(); + assert_eq!(contexts.len(), 1); + assert_eq!(contexts[0].event, HookEvent::PostToolUseFailure); + assert_eq!(contexts[0].tool_name.as_deref(), Some("shell")); + assert_eq!(contexts[0].tool_call_id.as_deref(), Some("call_1")); + assert_eq!( + contexts[0].error_message.as_deref(), + Some("command not found") + ); + } +} diff --git a/crates/arc-workflows/src/hook/mod.rs b/crates/arc-workflows/src/hook/mod.rs index 4fad81372..76c85a393 100644 --- a/crates/arc-workflows/src/hook/mod.rs +++ b/crates/arc-workflows/src/hook/mod.rs @@ -1,8 +1,10 @@ +pub mod bridge; pub mod config; pub mod executor; pub mod runner; pub mod types; +pub use bridge::WorkflowToolHookCallback; pub use config::{HookConfig, HookDefinition, HookType, TlsMode}; pub use runner::HookRunner; pub use types::{HookContext, HookDecision, HookEvent}; diff --git a/crates/arc-workflows/src/hook/runner.rs b/crates/arc-workflows/src/hook/runner.rs index 0aa24f5e7..daa4fadd4 100644 --- a/crates/arc-workflows/src/hook/runner.rs +++ b/crates/arc-workflows/src/hook/runner.rs @@ -116,6 +116,7 @@ impl HookRunner { context.handler_type.as_deref(), context.edge_to.as_deref(), context.edge_from.as_deref(), + context.tool_name.as_deref(), ] .iter() .any(|field| field.is_some_and(|v| re.is_match(v))) @@ -333,6 +334,29 @@ mod tests { assert!(runner.filter_hooks(&ctx).is_empty()); } + #[tokio::test] + async fn matcher_filters_by_tool_name() { + let mut hook = make_hook(HookEvent::PreToolUse, "tool-filter"); + hook.matcher = Some("shell".into()); + let config = HookConfig { hooks: vec![hook] }; + let runner = HookRunner::with_executor( + config, + Arc::new(MockExecutor { + decision: HookDecision::Proceed, + }), + ); + + // Matches tool_name "shell" + let mut ctx = make_context(HookEvent::PreToolUse); + ctx.tool_name = Some("shell".into()); + assert_eq!(runner.filter_hooks(&ctx).len(), 1); + + // Does not match tool_name "read_file" + let mut ctx = make_context(HookEvent::PreToolUse); + ctx.tool_name = Some("read_file".into()); + assert!(runner.filter_hooks(&ctx).is_empty()); + } + #[tokio::test] async fn blocking_hook_block_decision() { let config = HookConfig { diff --git a/crates/arc-workflows/src/hook/types.rs b/crates/arc-workflows/src/hook/types.rs index a4b38e7e1..2a299b860 100644 --- a/crates/arc-workflows/src/hook/types.rs +++ b/crates/arc-workflows/src/hook/types.rs @@ -17,13 +17,19 @@ pub enum HookEvent { SandboxReady, SandboxCleanup, CheckpointSaved, + PreToolUse, + PostToolUse, + PostToolUseFailure, } impl HookEvent { /// Whether hooks for this event block execution by default. #[must_use] pub fn is_blocking_by_default(self) -> bool { - matches!(self, Self::RunStart | Self::StageStart | Self::EdgeSelected) + matches!( + self, + Self::RunStart | Self::StageStart | Self::EdgeSelected | Self::PreToolUse + ) } } @@ -43,6 +49,9 @@ impl std::fmt::Display for HookEvent { Self::SandboxReady => "sandbox_ready", Self::SandboxCleanup => "sandbox_cleanup", Self::CheckpointSaved => "checkpoint_saved", + Self::PreToolUse => "pre_tool_use", + Self::PostToolUse => "post_tool_use", + Self::PostToolUseFailure => "post_tool_use_failure", }) } } @@ -75,6 +84,16 @@ pub struct HookContext { pub attempt: Option, #[serde(default, skip_serializing_if = "Option::is_none")] pub max_attempts: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub tool_name: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub tool_input: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub tool_call_id: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub tool_output: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub error_message: Option, } impl HookContext { @@ -95,6 +114,11 @@ impl HookContext { failure_reason: None, attempt: None, max_attempts: None, + tool_name: None, + tool_input: None, + tool_call_id: None, + tool_output: None, + error_message: None, } } } @@ -173,6 +197,9 @@ mod tests { HookEvent::SandboxReady, HookEvent::SandboxCleanup, HookEvent::CheckpointSaved, + HookEvent::PreToolUse, + HookEvent::PostToolUse, + HookEvent::PostToolUseFailure, ]; for event in events { let json = serde_json::to_string(&event).unwrap(); @@ -226,6 +253,11 @@ mod tests { failure_reason: None, attempt: Some(1), max_attempts: Some(3), + tool_name: None, + tool_input: None, + tool_call_id: None, + tool_output: None, + error_message: None, }; let json = serde_json::to_string(&ctx).unwrap(); let back: HookContext = serde_json::from_str(&json).unwrap(); @@ -332,4 +364,50 @@ mod tests { assert!(!resp.ok); assert_eq!(resp.reason.as_deref(), Some("not ready")); } + + #[test] + fn pre_tool_use_serde_round_trip() { + let json = serde_json::to_string(&HookEvent::PreToolUse).unwrap(); + assert_eq!(json, "\"pre_tool_use\""); + let back: HookEvent = serde_json::from_str(&json).unwrap(); + assert_eq!(back, HookEvent::PreToolUse); + } + + #[test] + fn pre_tool_use_is_blocking_by_default() { + assert!(HookEvent::PreToolUse.is_blocking_by_default()); + } + + #[test] + fn post_tool_use_is_not_blocking_by_default() { + assert!(!HookEvent::PostToolUse.is_blocking_by_default()); + } + + #[test] + fn post_tool_use_failure_is_not_blocking_by_default() { + assert!(!HookEvent::PostToolUseFailure.is_blocking_by_default()); + } + + #[test] + fn hook_context_with_tool_fields() { + let mut ctx = HookContext::new(HookEvent::PreToolUse, "run-1".into(), "wf".into()); + ctx.tool_name = Some("shell".into()); + ctx.tool_input = Some(serde_json::json!({"command": "ls"})); + ctx.tool_call_id = Some("call_123".into()); + let json = serde_json::to_string(&ctx).unwrap(); + assert!(json.contains("\"tool_name\":\"shell\"")); + assert!(json.contains("\"tool_call_id\":\"call_123\"")); + assert!(json.contains("\"tool_input\"")); + } + + #[test] + fn hook_context_tool_output_serializes() { + let mut ctx = HookContext::new(HookEvent::PostToolUse, "run-1".into(), "wf".into()); + ctx.tool_name = Some("shell".into()); + ctx.tool_output = Some("file1.txt\nfile2.txt".into()); + let json = serde_json::to_string(&ctx).unwrap(); + assert!(json.contains("\"tool_output\"")); + // error_message should be omitted + assert!(!json.contains("\"error_message\"")); + } } diff --git a/crates/arc-workflows/tests/daytona_integration.rs b/crates/arc-workflows/tests/daytona_integration.rs index 0488ec7b0..540b7bce6 100644 --- a/crates/arc-workflows/tests/daytona_integration.rs +++ b/crates/arc-workflows/tests/daytona_integration.rs @@ -932,6 +932,7 @@ async fn run_daytona_cli_test(provider: Provider, model: &str, install_command: &emitter, dir.path(), &env, + None, ) .await; diff --git a/crates/arc-workflows/tests/integration.rs b/crates/arc-workflows/tests/integration.rs index 1444dd62e..e1269026e 100644 --- a/crates/arc-workflows/tests/integration.rs +++ b/crates/arc-workflows/tests/integration.rs @@ -1179,6 +1179,7 @@ impl CodergenBackend for MockCodergenBackend { _emitter: &Arc, _stage_dir: &std::path::Path, _sandbox: &Arc, + _tool_hooks: Option>, ) -> Result { Ok(CodergenResult::Text { text: format!( @@ -5822,6 +5823,7 @@ mod real_llm { _emitter: &Arc, _stage_dir: &std::path::Path, _sandbox: &Arc, + _tool_hooks: Option>, ) -> Result { let request = Request { model: self.model.clone(), @@ -9290,6 +9292,7 @@ async fn cli_backend_run_writes_prompt_and_calls_exec() { &emitter, dir.path(), &env, + None, ) .await .expect("CLI backend should succeed"); @@ -9361,6 +9364,7 @@ async fn cli_backend_run_detects_changed_files() { &emitter, dir.path(), &env, + None, ) .await .expect("CLI backend should succeed"); @@ -9394,6 +9398,7 @@ async fn cli_backend_run_with_codex_provider() { &emitter, dir.path(), &env, + None, ) .await .expect("CLI backend should succeed"); @@ -9557,6 +9562,7 @@ async fn cli_backend_run_fails_on_nonzero_exit() { &emitter, dir.path(), &failing_env, + None, ) .await; @@ -9594,6 +9600,7 @@ async fn cli_backend_run_fails_on_unparseable_output() { &emitter, dir.path(), &env, + None, ) .await; @@ -9627,7 +9634,16 @@ async fn cli_backend_run_uses_node_model_override() { let dir = tempfile::tempdir().unwrap(); backend - .run(&node, "test", &context, None, &emitter, dir.path(), &env) + .run( + &node, + "test", + &context, + None, + &emitter, + dir.path(), + &env, + None, + ) .await .expect("should succeed"); @@ -9668,7 +9684,16 @@ async fn cli_backend_run_uses_node_provider_override() { let dir = tempfile::tempdir().unwrap(); backend - .run(&node, "test", &context, None, &emitter, dir.path(), &env) + .run( + &node, + "test", + &context, + None, + &emitter, + dir.path(), + &env, + None, + ) .await .expect("should succeed"); @@ -9693,7 +9718,16 @@ async fn cli_backend_run_writes_provider_used_json() { let dir = tempfile::tempdir().unwrap(); backend - .run(&node, "test", &context, None, &emitter, dir.path(), &env) + .run( + &node, + "test", + &context, + None, + &emitter, + dir.path(), + &env, + None, + ) .await .expect("should succeed"); @@ -9742,6 +9776,7 @@ async fn backend_router_delegates_to_cli_for_cli_node() { &emitter, dir.path(), &env, + None, ) .await .expect("router should succeed"); @@ -9784,6 +9819,7 @@ async fn backend_router_delegates_to_api_for_normal_node() { &emitter, dir.path(), &env, + None, ) .await .expect("router should succeed"); @@ -9829,6 +9865,7 @@ async fn backend_router_delegates_to_cli_for_backend_attr() { &emitter, dir.path(), &env, + None, ) .await .expect("router should succeed"); @@ -10110,6 +10147,7 @@ async fn run_real_cli_test(provider: Provider, model: &str) { &emitter, dir.path(), &env, + None, ) .await .unwrap_or_else(|_| panic!("CLI backend ({provider}/{model}) should succeed"));