mirror of
https://github.com/fabro-sh/fabro.git
synced 2026-10-09 03:20:56 +00:00
Add PreToolUse, PostToolUse, PostToolUseFailure Hooks (#6)
* arc(01KK739V1KPSQJT6E492ZWR95G): implement (success) Arc-Run: 01KK739V1KPSQJT6E492ZWR95G Arc-Completed: 2 Arc-Checkpoint: 52e0f637b3c802833c0c6eb7e2e8b56aef071441 * arc(01KK739V1KPSQJT6E492ZWR95G): simplify (success) Arc-Run: 01KK739V1KPSQJT6E492ZWR95G Arc-Completed: 3 Arc-Checkpoint: c1d6abe508b30b269d019961988212076064fb9e * Merge main into PR branch: integrate MCP servers with server mode dispatch Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> --------- Co-authored-by: arc <arc@local> Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
parent
7f264e44c1
commit
7cfbe5f9d1
16 changed files with 912 additions and 55 deletions
|
|
@ -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<dyn crate::config::ToolHookCallback> =
|
||||
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<dyn ProviderProfile> = 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()
|
||||
},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -8,6 +8,55 @@ use arc_mcp::config::McpServerConfig;
|
|||
/// `Err(message)` to deny with the given message.
|
||||
pub type ToolApprovalFn = Arc<dyn Fn(&str, &serde_json::Value) -> 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<String>,
|
||||
pub user_instructions: Option<String>,
|
||||
pub tool_approval: Option<ToolApprovalFn>,
|
||||
/// Async hook callbacks invoked around tool execution.
|
||||
pub tool_hooks: Option<Arc<dyn ToolHookCallback>>,
|
||||
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(|_| "<fn>"),
|
||||
"tool_hooks",
|
||||
&self.tool_hooks.as_ref().map(|_| "<callback>"),
|
||||
)
|
||||
.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;
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -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<dyn Sandbox>,
|
||||
tool_approval: Option<&ToolApprovalFn>,
|
||||
tool_hooks: Option<&Arc<dyn ToolHookCallback>>,
|
||||
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<dyn Sandbox>,
|
||||
tool_approval: Option<&ToolApprovalFn>,
|
||||
tool_hooks: Option<&Arc<dyn ToolHookCallback>>,
|
||||
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<dyn Sandbox>,
|
||||
tool_approval: Option<&ToolApprovalFn>,
|
||||
tool_hooks: Option<&Arc<dyn ToolHookCallback>>,
|
||||
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<dyn Sandbox>,
|
||||
tool_approval: Option<&ToolApprovalFn>,
|
||||
tool_hooks: Option<&Arc<dyn ToolHookCallback>>,
|
||||
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<dyn Sandbox>,
|
||||
tool_approval: Option<&ToolApprovalFn>,
|
||||
tool_hooks: Option<&Arc<dyn ToolHookCallback>>,
|
||||
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<dyn Sandbox>,
|
||||
tool_approval: Option<&ToolApprovalFn>,
|
||||
cancel_token: CancellationToken,
|
||||
tool_env: Option<&HashMap<String, String>>,
|
||||
) -> 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<Mutex<Vec<(String, String, String)>>>,
|
||||
post_failure_calls: Arc<Mutex<Vec<(String, String, String)>>>,
|
||||
}
|
||||
|
||||
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<dyn Sandbox> {
|
||||
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<dyn ToolHookCallback> =
|
||||
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<dyn ToolHookCallback> =
|
||||
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<dyn ToolHookCallback> = 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<dyn ToolHookCallback> = 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"));
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -133,8 +133,17 @@ impl AgentApiBackend {
|
|||
&self,
|
||||
node: &Node,
|
||||
sandbox: &Arc<dyn Sandbox>,
|
||||
tool_hooks: Option<Arc<dyn arc_agent::ToolHookCallback>>,
|
||||
) -> Result<Session, ArcError> {
|
||||
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<dyn Sandbox>,
|
||||
env: &HashMap<String, String>,
|
||||
tool_hooks: Option<Arc<dyn arc_agent::ToolHookCallback>>,
|
||||
) -> Result<Session, ArcError> {
|
||||
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<crate::event::EventEmitter>,
|
||||
stage_dir: &std::path::Path,
|
||||
sandbox: &Arc<dyn Sandbox>,
|
||||
tool_hooks: Option<Arc<dyn arc_agent::ToolHookCallback>>,
|
||||
) -> Result<CodergenResult, ArcError> {
|
||||
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
|
||||
{
|
||||
|
|
|
|||
|
|
@ -440,6 +440,7 @@ impl CodergenBackend for AgentCliBackend {
|
|||
emitter: &Arc<EventEmitter>,
|
||||
stage_dir: &Path,
|
||||
sandbox: &Arc<dyn Sandbox>,
|
||||
_tool_hooks: Option<Arc<dyn arc_agent::ToolHookCallback>>,
|
||||
) -> Result<CodergenResult, ArcError> {
|
||||
// 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<EventEmitter>,
|
||||
stage_dir: &Path,
|
||||
sandbox: &Arc<dyn Sandbox>,
|
||||
tool_hooks: Option<Arc<dyn arc_agent::ToolHookCallback>>,
|
||||
) -> Result<CodergenResult, ArcError> {
|
||||
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<EventEmitter>,
|
||||
_stage_dir: &Path,
|
||||
_sandbox: &Arc<dyn Sandbox>,
|
||||
_tool_hooks: Option<Arc<dyn arc_agent::ToolHookCallback>>,
|
||||
) -> Result<CodergenResult, ArcError> {
|
||||
Ok(CodergenResult::Text {
|
||||
text: "stub".to_string(),
|
||||
|
|
|
|||
|
|
@ -38,6 +38,7 @@ pub trait CodergenBackend: Send + Sync {
|
|||
emitter: &Arc<EventEmitter>,
|
||||
stage_dir: &Path,
|
||||
sandbox: &Arc<dyn Sandbox>,
|
||||
tool_hooks: Option<Arc<dyn arc_agent::ToolHookCallback>>,
|
||||
) -> Result<CodergenResult, ArcError>;
|
||||
|
||||
/// 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<Arc<dyn arc_agent::ToolHookCallback>> =
|
||||
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<dyn arc_agent::ToolHookCallback>
|
||||
});
|
||||
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<EventEmitter>,
|
||||
_stage_dir: &Path,
|
||||
_sandbox: &Arc<dyn arc_agent::Sandbox>,
|
||||
_tool_hooks: Option<Arc<dyn arc_agent::ToolHookCallback>>,
|
||||
) -> Result<CodergenResult, ArcError> {
|
||||
Ok(CodergenResult::Text {
|
||||
text: r#"Done. {"outcome": "success", "preferred_next_label": "approve"}"#
|
||||
|
|
@ -603,6 +617,7 @@ mod tests {
|
|||
_emitter: &Arc<EventEmitter>,
|
||||
_stage_dir: &Path,
|
||||
_sandbox: &Arc<dyn Sandbox>,
|
||||
_tool_hooks: Option<Arc<dyn arc_agent::ToolHookCallback>>,
|
||||
) -> Result<CodergenResult, ArcError> {
|
||||
*self.captured_thread_id.lock().unwrap() = Some(thread_id.map(String::from));
|
||||
Ok(CodergenResult::Text {
|
||||
|
|
@ -654,6 +669,7 @@ mod tests {
|
|||
_emitter: &Arc<EventEmitter>,
|
||||
_stage_dir: &Path,
|
||||
_sandbox: &Arc<dyn Sandbox>,
|
||||
_tool_hooks: Option<Arc<dyn arc_agent::ToolHookCallback>>,
|
||||
) -> Result<CodergenResult, ArcError> {
|
||||
*self.captured_thread_id.lock().unwrap() = Some(thread_id.map(String::from));
|
||||
Ok(CodergenResult::Text {
|
||||
|
|
@ -700,6 +716,7 @@ mod tests {
|
|||
_emitter: &Arc<EventEmitter>,
|
||||
_stage_dir: &Path,
|
||||
_sandbox: &Arc<dyn Sandbox>,
|
||||
_tool_hooks: Option<Arc<dyn arc_agent::ToolHookCallback>>,
|
||||
) -> Result<CodergenResult, ArcError> {
|
||||
Err(ArcError::handler("Request timed out".to_string()))
|
||||
}
|
||||
|
|
@ -843,6 +860,7 @@ Some text in between.
|
|||
_emitter: &Arc<EventEmitter>,
|
||||
_stage_dir: &Path,
|
||||
_sandbox: &Arc<dyn Sandbox>,
|
||||
_tool_hooks: Option<Arc<dyn arc_agent::ToolHookCallback>>,
|
||||
) -> Result<CodergenResult, ArcError> {
|
||||
Err(ArcError::Validation("bad config".to_string()))
|
||||
}
|
||||
|
|
@ -881,6 +899,7 @@ Some text in between.
|
|||
_emitter: &Arc<EventEmitter>,
|
||||
_stage_dir: &std::path::Path,
|
||||
_sandbox: &Arc<dyn Sandbox>,
|
||||
_tool_hooks: Option<Arc<dyn arc_agent::ToolHookCallback>>,
|
||||
) -> Result<CodergenResult, ArcError> {
|
||||
*self.captured_prompt.lock().unwrap() = Some(prompt.to_string());
|
||||
Ok(CodergenResult::Text {
|
||||
|
|
@ -949,6 +968,7 @@ Some text in between.
|
|||
_emitter: &Arc<EventEmitter>,
|
||||
_stage_dir: &std::path::Path,
|
||||
_sandbox: &Arc<dyn Sandbox>,
|
||||
_tool_hooks: Option<Arc<dyn arc_agent::ToolHookCallback>>,
|
||||
) -> Result<CodergenResult, ArcError> {
|
||||
*self.captured_prompt.lock().unwrap() = Some(prompt.to_string());
|
||||
Ok(CodergenResult::Text {
|
||||
|
|
|
|||
|
|
@ -228,6 +228,7 @@ async fn llm_evaluate(
|
|||
emitter,
|
||||
&stage_dir,
|
||||
sandbox,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
{
|
||||
|
|
@ -435,6 +436,7 @@ mod tests {
|
|||
_emitter: &Arc<EventEmitter>,
|
||||
_stage_dir: &std::path::Path,
|
||||
_sandbox: &Arc<dyn Sandbox>,
|
||||
_tool_hooks: Option<Arc<dyn arc_agent::ToolHookCallback>>,
|
||||
) -> Result<CodergenResult, ArcError> {
|
||||
// Return text that contains the ID "branch_b"
|
||||
Ok(CodergenResult::Text {
|
||||
|
|
|
|||
|
|
@ -211,6 +211,7 @@ mod tests {
|
|||
_emitter: &Arc<crate::event::EventEmitter>,
|
||||
_stage_dir: &Path,
|
||||
_sandbox: &Arc<dyn Sandbox>,
|
||||
_tool_hooks: Option<Arc<dyn arc_agent::ToolHookCallback>>,
|
||||
) -> Result<CodergenResult, ArcError> {
|
||||
panic!("run() should not be called for prompt handler");
|
||||
}
|
||||
|
|
@ -272,6 +273,7 @@ mod tests {
|
|||
_emitter: &Arc<crate::event::EventEmitter>,
|
||||
_stage_dir: &Path,
|
||||
_sandbox: &Arc<dyn arc_agent::Sandbox>,
|
||||
_tool_hooks: Option<Arc<dyn arc_agent::ToolHookCallback>>,
|
||||
) -> Result<CodergenResult, ArcError> {
|
||||
panic!("run() should not be called for prompt handler");
|
||||
}
|
||||
|
|
|
|||
265
crates/arc-workflows/src/hook/bridge.rs
Normal file
265
crates/arc-workflows/src/hook/bridge.rs
Normal file
|
|
@ -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<HookRunner>,
|
||||
pub sandbox: Arc<dyn Sandbox>,
|
||||
pub run_id: String,
|
||||
pub workflow_name: String,
|
||||
pub work_dir: Option<PathBuf>,
|
||||
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<Mutex<Vec<HookContext>>>,
|
||||
decision: HookDecision,
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl HookExecutor for CapturingExecutor {
|
||||
async fn execute(
|
||||
&self,
|
||||
_definition: &HookDefinition,
|
||||
context: &HookContext,
|
||||
_sandbox: Arc<dyn Sandbox>,
|
||||
_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<dyn Sandbox> {
|
||||
Arc::new(arc_agent::LocalSandbox::new(
|
||||
std::env::current_dir().unwrap(),
|
||||
))
|
||||
}
|
||||
|
||||
fn make_bridge(
|
||||
hook_runner: Arc<HookRunner>,
|
||||
sandbox: Arc<dyn Sandbox>,
|
||||
) -> 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")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -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};
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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<usize>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub max_attempts: Option<usize>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub tool_name: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub tool_input: Option<serde_json::Value>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub tool_call_id: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub tool_output: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub error_message: Option<String>,
|
||||
}
|
||||
|
||||
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\""));
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -932,6 +932,7 @@ async fn run_daytona_cli_test(provider: Provider, model: &str, install_command:
|
|||
&emitter,
|
||||
dir.path(),
|
||||
&env,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
|
|
|
|||
|
|
@ -1179,6 +1179,7 @@ impl CodergenBackend for MockCodergenBackend {
|
|||
_emitter: &Arc<EventEmitter>,
|
||||
_stage_dir: &std::path::Path,
|
||||
_sandbox: &Arc<dyn arc_agent::Sandbox>,
|
||||
_tool_hooks: Option<Arc<dyn arc_agent::ToolHookCallback>>,
|
||||
) -> Result<CodergenResult, ArcError> {
|
||||
Ok(CodergenResult::Text {
|
||||
text: format!(
|
||||
|
|
@ -5822,6 +5823,7 @@ mod real_llm {
|
|||
_emitter: &Arc<EventEmitter>,
|
||||
_stage_dir: &std::path::Path,
|
||||
_sandbox: &Arc<dyn arc_agent::Sandbox>,
|
||||
_tool_hooks: Option<Arc<dyn arc_agent::ToolHookCallback>>,
|
||||
) -> Result<CodergenResult, ArcError> {
|
||||
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"));
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue