// src/agent/hooks/mod.rs // // Agent 生命周期 Hooks 系统。 // // 参考 Claude Code 的 PreToolUse / PostToolUse / Stop hooks 设计, // 提供可扩展的事件回调链,支持 10 种生命周期事件 + 工具匹配过滤 + // 权限决策优先级 + 异步 fire-and-forget hook。 // // 子模块结构: // - types.rs — 所有数据类型定义(HookEvent、Contexts、Actions、Results) // - traits.rs — AgentHook + AsyncAgentHook traits // - matcher.rs — ToolNamePattern / ToolMatchFilter 工具匹配器 // - registry.rs — HookRegistry 结构体 + 基础方法 // - dispatch.rs — HookRegistry 调度方法(所有 run_*) // - builtins.rs — 内置 Hooks(CancellationHook、MetricsHook、AuditLogHook) pub mod builtins; pub mod dispatch; pub mod matcher; pub mod registry; pub mod traits; pub mod types; // ── 重导出:外部代码通过 `crate::agent::hooks::*` 访问 ── // 核心注册表 pub use registry::HookRegistry; // 类型 pub use types::{ event_label, BlockingError, HookAction, HookEvent, MetricsData, PermissionDecision, PermissionDenialSource, PermissionDeniedContext, PermissionRequestAction, PermissionRequestContext, PostCompactContext, PostToolUseAction, PostToolUseContext, PostToolUseFailureContext, PostToolUseResult, PreCompactContext, PreToolUseAction, PreToolUseContext, PreToolUseResult, SessionStartContext, SessionStopContext, StepCompleteContext, SubagentStartContext, SubagentStopContext, TaggedContext, DEFAULT_HOOK_TIMEOUT, }; // Traits pub use traits::{AgentHook, AsyncAgentHook}; // Matcher pub use matcher::{ToolMatchFilter, ToolNamePattern}; // Builtins pub use builtins::{AuditLogHook, CancellationHook, ContextDeduplicator, MetricsHook}; // ── 测试 ── #[cfg(test)] mod tests { use super::*; use async_trait::async_trait; use serde_json::{json, Value}; use sqlx::SqlitePool; use std::collections::HashSet; use std::sync::{Arc, Mutex}; struct TestHook { name: String, pre_called: std::sync::Mutex, } impl TestHook { fn new(name: &str) -> Self { TestHook { name: name.to_string(), pre_called: std::sync::Mutex::new(false), } } } #[async_trait] impl AgentHook for TestHook { fn name(&self) -> &str { &self.name } async fn pre_tool_use(&self, _ctx: &PreToolUseContext) -> PreToolUseAction { *self.pre_called.lock().unwrap() = true; PreToolUseAction::Continue } } #[tokio::test] async fn test_hook_registry_runs_all_hooks() { let mut registry = HookRegistry::new(); let hook1 = TestHook::new("test1"); let hook2 = TestHook::new("test2"); registry.add(Box::new(hook1)); registry.add(Box::new(hook2)); let ctx = PreToolUseContext { session_id: "test".into(), tool_name: "test_tool".into(), tool_args: serde_json::json!({}), step: 1, }; let result = registry.run_pre_tool_use(&ctx).await; assert!(!result.action.is_blocked()); } #[tokio::test] async fn test_blocking_hook_stops_chain() { struct BlockingHook; #[async_trait] impl AgentHook for BlockingHook { fn name(&self) -> &str { "blocker" } async fn pre_tool_use(&self, _ctx: &PreToolUseContext) -> PreToolUseAction { PreToolUseAction::Block { reason: "test block".into(), } } } let mut registry = HookRegistry::new(); registry.add(Box::new(BlockingHook)); let ctx = PreToolUseContext { session_id: "test".into(), tool_name: "test_tool".into(), tool_args: serde_json::json!({}), step: 1, }; let result = registry.run_pre_tool_use(&ctx).await; assert!(result.action.is_blocked()); assert_eq!(result.action.block_reason(), Some("test block")); } #[tokio::test] async fn test_mutate_input_accumulates_context() { struct MutateHook; #[async_trait] impl AgentHook for MutateHook { fn name(&self) -> &str { "mutator" } async fn pre_tool_use(&self, _ctx: &PreToolUseContext) -> PreToolUseAction { PreToolUseAction::MutateInput { updated_args: serde_json::json!({"key": "modified"}), additional_context: Some("injected context".to_string()), } } } let mut registry = HookRegistry::new(); registry.add(Box::new(MutateHook)); let ctx = PreToolUseContext { session_id: "test".into(), tool_name: "test_tool".into(), tool_args: serde_json::json!({"key": "original"}), step: 1, }; let result = registry.run_pre_tool_use(&ctx).await; assert_eq!(result.final_args, serde_json::json!({"key": "modified"})); assert_eq!( result.additional_context, Some("injected context".to_string()) ); } #[tokio::test] async fn test_post_tool_use_mutate_output() { struct MutateOutputHook; #[async_trait] impl AgentHook for MutateOutputHook { fn name(&self) -> &str { "output_mutator" } async fn post_tool_use(&self, _ctx: &PostToolUseContext) -> PostToolUseAction { PostToolUseAction::MutateOutput { updated_content: "modified output".to_string(), additional_context: None, } } } let mut registry = HookRegistry::new(); registry.add(Box::new(MutateOutputHook)); let ctx = PostToolUseContext { session_id: "test".into(), agent_name: "lead".into(), tool_name: "test_tool".into(), tool_args: serde_json::json!({}), output_content: "original output".into(), is_error: false, step: 1, elapsed_ms: 100, }; let result = registry.run_post_tool_use(&ctx).await; assert_eq!(result.final_content, "modified output"); } #[tokio::test] async fn test_cancellation_hook_blocks_when_cancelled() { let mut cancelled = HashSet::new(); cancelled.insert("test_session".to_string()); let cancelled_runs = Arc::new(std::sync::Mutex::new(cancelled)); let hook = CancellationHook::new(cancelled_runs); let ctx = PreToolUseContext { session_id: "test_session".into(), tool_name: "search_papers".into(), tool_args: serde_json::json!({}), step: 1, }; let action = hook.pre_tool_use(&ctx).await; assert!(action.is_blocked()); } #[tokio::test] async fn test_cancellation_hook_allows_when_not_cancelled() { let cancelled_runs = Arc::new(std::sync::Mutex::new(std::collections::HashSet::new())); let hook = CancellationHook::new(cancelled_runs); let ctx = PreToolUseContext { session_id: "test_session".into(), tool_name: "search_papers".into(), tool_args: serde_json::json!({}), step: 1, }; let action = hook.pre_tool_use(&ctx).await; assert!(!action.is_blocked()); } #[tokio::test] async fn test_metrics_hook_accumulates_counts() { let hook = MetricsHook::new(); let ctx = PostToolUseContext { session_id: "test".into(), agent_name: "lead".into(), tool_name: "search_papers".into(), tool_args: serde_json::json!({}), output_content: "result".into(), is_error: false, step: 1, elapsed_ms: 100, }; hook.post_tool_use(&ctx).await; let ctx2 = PostToolUseContext { session_id: "test".into(), agent_name: "lead".into(), tool_name: "search_papers".into(), tool_args: serde_json::json!({}), output_content: "result2".into(), is_error: false, step: 2, elapsed_ms: 200, }; hook.post_tool_use(&ctx2).await; let snapshot = hook.snapshot().expect("snapshot should succeed in test"); assert_eq!(snapshot.tool_call_counts.get("search_papers"), Some(&2)); assert_eq!(snapshot.total_steps, 2); } #[tokio::test] async fn test_session_start_hook_called() { struct StartTrackingHook { started: std::sync::Mutex>, } #[async_trait] impl AgentHook for StartTrackingHook { fn name(&self) -> &str { "start_tracker" } async fn on_session_start(&self, ctx: &SessionStartContext) { self.started.lock().unwrap().push(ctx.session_id.clone()); } } let hook = StartTrackingHook { started: std::sync::Mutex::new(Vec::new()), }; let mut registry = HookRegistry::new(); registry.add(Box::new(hook)); let ctx = SessionStartContext { session_id: "test_sid".into(), turn_index: 1, is_resume: false, }; registry.run_on_session_start(&ctx).await; } #[tokio::test] async fn test_new_lifecycle_events_called() { struct LifecycleTracker { subagent_start: std::sync::Mutex, subagent_stop: std::sync::Mutex, pre_compact: std::sync::Mutex, post_compact: std::sync::Mutex, } #[async_trait] impl AgentHook for LifecycleTracker { fn name(&self) -> &str { "lifecycle_tracker" } async fn on_subagent_start(&self, _ctx: &SubagentStartContext) { *self.subagent_start.lock().unwrap() = true; } async fn on_subagent_stop(&self, _ctx: &SubagentStopContext) { *self.subagent_stop.lock().unwrap() = true; } async fn on_pre_compact(&self, _ctx: &PreCompactContext) { *self.pre_compact.lock().unwrap() = true; } async fn on_post_compact(&self, _ctx: &PostCompactContext) { *self.post_compact.lock().unwrap() = true; } } let tracker = LifecycleTracker { subagent_start: std::sync::Mutex::new(false), subagent_stop: std::sync::Mutex::new(false), pre_compact: std::sync::Mutex::new(false), post_compact: std::sync::Mutex::new(false), }; let mut registry = HookRegistry::new(); registry.add(Box::new(tracker)); registry .run_on_subagent_start(&SubagentStartContext { parent_session_id: "s1".into(), subagent_name: "sub".into(), prompt: "test".into(), }) .await; registry .run_on_subagent_stop(&SubagentStopContext { parent_session_id: "s1".into(), subagent_name: "sub".into(), result_summary: "done".into(), steps: 3, is_error: false, }) .await; registry .run_on_pre_compact(&PreCompactContext { session_id: "s1".into(), message_count: 50, estimated_tokens: 10000, }) .await; registry .run_on_post_compact(&PostCompactContext { session_id: "s1".into(), new_message_count: 10, compression_method: "micro".into(), }) .await; // If no panic, all hooks were called successfully } #[tokio::test] async fn test_match_filter_skips_irrelevant() { use std::sync::atomic::{AtomicUsize, Ordering}; struct CountedHook { name: String, call_count: Arc, filter: ToolMatchFilter, } #[async_trait] impl AgentHook for CountedHook { fn name(&self) -> &str { &self.name } fn match_filter(&self) -> ToolMatchFilter { self.filter.clone() } async fn pre_tool_use(&self, _ctx: &PreToolUseContext) -> PreToolUseAction { self.call_count.fetch_add(1, Ordering::SeqCst); PreToolUseAction::Continue } async fn post_tool_use(&self, _ctx: &PostToolUseContext) -> PostToolUseAction { self.call_count.fetch_add(1, Ordering::SeqCst); PostToolUseAction::Continue } } let search_count = Arc::new(AtomicUsize::new(0)); let all_count = Arc::new(AtomicUsize::new(0)); let search_only = CountedHook { name: "search_only".into(), call_count: search_count.clone(), filter: ToolMatchFilter::exact("search_papers"), }; let all_match = CountedHook { name: "all_match".into(), call_count: all_count.clone(), filter: ToolMatchFilter::default(), }; let mut registry = HookRegistry::new(); registry.add(Box::new(search_only)); registry.add(Box::new(all_match)); let ctx = PreToolUseContext { session_id: "test".into(), tool_name: "search_papers".into(), tool_args: serde_json::json!({}), step: 1, }; registry.run_pre_tool_use(&ctx).await; assert_eq!(search_count.load(Ordering::SeqCst), 1); assert_eq!(all_count.load(Ordering::SeqCst), 1); let ctx2 = PreToolUseContext { session_id: "test".into(), tool_name: "download_paper".into(), tool_args: serde_json::json!({}), step: 2, }; registry.run_pre_tool_use(&ctx2).await; assert_eq!(search_count.load(Ordering::SeqCst), 1); assert_eq!(all_count.load(Ordering::SeqCst), 2); } #[tokio::test] async fn test_full_hook_pipeline_integration() { use std::sync::atomic::Ordering; struct IntegrationHook { name: String, filter: ToolMatchFilter, post_calls: Arc, } #[async_trait] impl AgentHook for IntegrationHook { fn name(&self) -> &str { &self.name } fn match_filter(&self) -> ToolMatchFilter { self.filter.clone() } async fn post_tool_use(&self, ctx: &PostToolUseContext) -> PostToolUseAction { self.post_calls .fetch_add(1, std::sync::atomic::Ordering::SeqCst); if self.name == "warn_hook" { PostToolUseAction::Warning { message: "test_warning".into(), truncate_output: false, } } else if self.name == "meta_hook" { PostToolUseAction::Metadata { key: "origin".into(), value: serde_json::Value::String("integration_test".into()), } } else { PostToolUseAction::MutateOutput { updated_content: format!("[{}] {}", self.name, ctx.output_content), additional_context: Some(format!("context_from_{}", self.name)), } } } async fn pre_tool_use(&self, _ctx: &PreToolUseContext) -> PreToolUseAction { PreToolUseAction::Continue } } let mut registry = HookRegistry::new(); let file_count = Arc::new(std::sync::atomic::AtomicUsize::new(0)); let search_count = Arc::new(std::sync::atomic::AtomicUsize::new(0)); registry.add(Box::new(IntegrationHook { name: "file_hook".into(), filter: ToolMatchFilter::prefix("file_"), post_calls: file_count.clone(), })); registry.add(Box::new(IntegrationHook { name: "search_hook".into(), filter: ToolMatchFilter::exact("search_papers"), post_calls: search_count.clone(), })); registry.add(Box::new(IntegrationHook { name: "warn_hook".into(), filter: ToolMatchFilter::default(), post_calls: Arc::new(std::sync::atomic::AtomicUsize::new(0)), })); registry.add(Box::new(IntegrationHook { name: "meta_hook".into(), filter: ToolMatchFilter::default(), post_calls: Arc::new(std::sync::atomic::AtomicUsize::new(0)), })); let ctx = PostToolUseContext { session_id: "test".into(), agent_name: "lead".into(), tool_name: "file_write".into(), tool_args: serde_json::json!({"path": "/tmp/test.txt"}), output_content: "file content".into(), is_error: false, step: 1, elapsed_ms: 100, }; let result = registry.run_post_tool_use(&ctx).await; assert_eq!(file_count.load(Ordering::SeqCst), 1); assert_eq!(search_count.load(Ordering::SeqCst), 0); assert!(result.final_content.contains("[file_hook]")); assert!(!result.tagged_contexts.is_empty()); assert!(result .tagged_contexts .iter() .any(|tc| tc.hook_name == "file_hook")); assert!(result.warnings.contains(&"test_warning".to_string())); assert_eq!( result.metadata.get("origin").and_then(|v| v.as_str()), Some("integration_test") ); let ctx2 = PostToolUseContext { session_id: "test".into(), agent_name: "lead".into(), tool_name: "search_papers".into(), tool_args: serde_json::json!({"query": "quasars"}), output_content: "search results".into(), is_error: false, step: 2, elapsed_ms: 200, }; let result2 = registry.run_post_tool_use(&ctx2).await; assert_eq!(search_count.load(Ordering::SeqCst), 1); assert_eq!(file_count.load(Ordering::SeqCst), 1); assert!(result2.final_content.contains("[search_hook]")); } }