// src/agent/hooks/matcher.rs // // 工具级匹配器。 // Hook 可通过 `match_filter()` 声明只关心特定工具或参数模式, // 从而避免在热路径上被无关调用触发。 use serde_json::Value; /// 工具名匹配模式。 /// /// 三种策略: /// - `Exact("search_papers")` — 精确名称匹配 /// - `Prefix("file_")` — 前缀匹配(末尾 `*` 隐式) /// - `Wildcard` — 匹配所有工具(默认行为) /// /// # 示例 /// /// ``` /// use crate::agent::hooks::matcher::ToolNamePattern; /// let pat = ToolNamePattern::parse("file_*"); /// assert!(pat.matches("file_write")); /// assert!(!pat.matches("search_papers")); /// ``` #[derive(Debug, Clone, PartialEq, Eq)] pub enum ToolNamePattern { /// 精确匹配工具名 Exact(String), /// 前缀匹配 — `"file_*"` 匹配所有以 `"file_"` 开头的工具名 Prefix(String), /// 匹配所有工具 Wildcard, } impl ToolNamePattern { /// 从模式字符串解析。 /// /// - `"*"` 或空字符串 → `Wildcard` /// - `"xxx_*"` → `Prefix("xxx_")` /// - 其他 → `Exact(pattern)` pub fn parse(pattern: &str) -> Self { let trimmed = pattern.trim(); if trimmed.is_empty() || trimmed == "*" { return ToolNamePattern::Wildcard; } if let Some(prefix) = trimmed.strip_suffix('*') { if !prefix.is_empty() { return ToolNamePattern::Prefix(prefix.to_string()); } // 只有 "*" — 上面已经处理 return ToolNamePattern::Wildcard; } ToolNamePattern::Exact(trimmed.to_string()) } /// 若 `tool_name` 与此模式匹配则返回 true。 pub fn matches(&self, tool_name: &str) -> bool { match self { ToolNamePattern::Wildcard => true, ToolNamePattern::Exact(name) => name == tool_name, ToolNamePattern::Prefix(prefix) => tool_name.starts_with(prefix), } } } /// Hook 通过 `match_filter()` 声明的结构化过滤器,用于限制关注哪些工具调用。 /// /// 若 hook 返回非空 `ToolMatchFilter`,则仅当工具名和参数匹配时 /// 才调用其 `pre_tool_use` / `post_tool_use`。 /// /// 默认情况下(空过滤器,等价于 `[Wildcard]`),匹配所有工具。 /// /// # 示例 /// /// ``` /// use crate::agent::hooks::matcher::{ToolMatchFilter, ToolNamePattern}; /// /// // 仅匹配文件相关工具 /// let filter = ToolMatchFilter { /// name_patterns: vec![ToolNamePattern::parse("file_*")], /// content_patterns: vec![], /// }; /// assert!(filter.matches("file_read", &serde_json::json!({}))); /// assert!(!filter.matches("search_papers", &serde_json::json!({}))); /// ``` #[derive(Debug, Clone, Default)] pub struct ToolMatchFilter { /// 工具名模式 — 任意一个匹配即可通过。 /// 空 Vec = 匹配全部(等价于 `[Wildcard]`)。 pub name_patterns: Vec, /// 可选的内容级模式。格式:`"字段:子串"` 或直接 `"子串"`。 /// 例如 `"command:rm *"` 匹配 `run_bash` 调用中 command 字段以 "rm " 开头的情况。 /// 空 Vec = 不过滤内容。 pub content_patterns: Vec, } impl ToolMatchFilter { /// 若此过滤器匹配给定的工具调用则返回 true。 /// /// 检查流程: /// 1. 若 `name_patterns` 为空,所有工具名匹配 /// 2. 任一非通配模式必须匹配 tool_name /// 3. 任一 content_patterns 必须匹配 tool_args 中的某个字段 pub fn matches(&self, tool_name: &str, tool_args: &Value) -> bool { // 名称匹配 if !self.name_patterns.is_empty() { let name_match = self.name_patterns.iter().any(|pat| pat.matches(tool_name)); if !name_match { return false; } } // 内容匹配:每个模式必须至少匹配一个字段/值 if !self.content_patterns.is_empty() { for pattern in &self.content_patterns { if !content_matches(pattern, tool_args) { return false; } } } true } /// 快捷构造:创建仅匹配单个工具名的过滤器。 pub fn exact(tool_name: &str) -> Self { ToolMatchFilter { name_patterns: vec![ToolNamePattern::Exact(tool_name.to_string())], content_patterns: vec![], } } /// 快捷构造:创建前缀过滤器。 pub fn prefix(prefix: &str) -> Self { ToolMatchFilter { name_patterns: vec![ToolNamePattern::Prefix(prefix.to_string())], content_patterns: vec![], } } /// 若此过滤器为默认"匹配全部"(空 name_patterns + 空 content_patterns)则返回 true。 pub fn is_match_all(&self) -> bool { self.name_patterns.is_empty() && self.content_patterns.is_empty() } } /// 检查内容模式是否匹配 JSON 值中的任意字段。 /// /// 模式格式: /// - `"字段:值"` — 值必须是该字段字符串表示的子串 /// - `"值"` — 值必须出现在 JSON 字符串化后的任意位置 fn content_matches(pattern: &str, args: &Value) -> bool { if let Some((field, wanted)) = pattern.split_once(':') { // 匹配指定字段 if let Some(field_val) = args.get(field.trim()) { let field_str = match field_val { Value::String(s) => s.clone(), other => other.to_string(), }; return field_str.contains(wanted.trim()); } return false; } // 在 JSON 中任意位置匹配 let full_str = serde_json::to_string(args).unwrap_or_default(); full_str.contains(pattern.trim()) } // ── 测试 ── #[cfg(test)] mod tests { use super::*; // ── ToolNamePattern ── #[test] fn test_pattern_exact() { let pat = ToolNamePattern::parse("search_papers"); assert!(pat.matches("search_papers")); assert!(!pat.matches("download_paper")); } #[test] fn test_pattern_prefix() { let pat = ToolNamePattern::parse("file_*"); assert!(pat.matches("file_write")); assert!(pat.matches("file_read")); assert!(pat.matches("file_edit")); assert!(!pat.matches("search_papers")); assert!(!pat.matches("fi")); // 前缀比 "file_" 短 } #[test] fn test_pattern_wildcard() { let pat = ToolNamePattern::parse("*"); assert!(matches!(pat, ToolNamePattern::Wildcard)); assert!(pat.matches("anything")); assert!(pat.matches("search_papers")); let pat_empty = ToolNamePattern::parse(""); assert!(matches!(pat_empty, ToolNamePattern::Wildcard)); } #[test] fn test_pattern_prefix_no_trailing_wild() { // "file" 不带 * 应被解析为 Exact let pat = ToolNamePattern::parse("file"); assert!(matches!(pat, ToolNamePattern::Exact(_))); assert!(pat.matches("file")); assert!(!pat.matches("file_write")); } // ── ToolMatchFilter ── #[test] fn test_filter_default_matches_all() { let filter = ToolMatchFilter::default(); assert!(filter.matches("search_papers", &serde_json::json!({}))); assert!(filter.matches("run_bash", &serde_json::json!({"command": "rm -rf /"}))); } #[test] fn test_filter_exact_name_match() { let filter = ToolMatchFilter::exact("search_papers"); assert!(filter.matches("search_papers", &serde_json::json!({}))); assert!(!filter.matches("download_paper", &serde_json::json!({}))); } #[test] fn test_filter_prefix_name_match() { let filter = ToolMatchFilter::prefix("file_"); assert!(filter.matches("file_write", &serde_json::json!({}))); assert!(filter.matches("file_read", &serde_json::json!({}))); assert!(!filter.matches("run_bash", &serde_json::json!({}))); } #[test] fn test_filter_content_match_field() { let filter = ToolMatchFilter { name_patterns: vec![], content_patterns: vec!["command:rm".to_string()], }; assert!(filter.matches( "run_bash", &serde_json::json!({"command": "rm -rf /tmp/test"}) )); assert!(!filter.matches("run_bash", &serde_json::json!({"command": "ls -la"}))); assert!(!filter.matches("read_file", &serde_json::json!({"path": "/tmp/test"}))); } #[test] fn test_filter_content_match_anywhere() { let filter = ToolMatchFilter { name_patterns: vec![], content_patterns: vec!["dangerous".to_string()], }; assert!(filter.matches( "run_bash", &serde_json::json!({"command": "echo dangerous stuff"}) )); assert!(!filter.matches( "run_bash", &serde_json::json!({"command": "echo safe stuff"}) )); } #[test] fn test_filter_name_and_content_combined() { let filter = ToolMatchFilter { name_patterns: vec![ToolNamePattern::parse("file_*")], content_patterns: vec!["path:.env".to_string()], }; // 正确工具名 + 敏感路径 → 匹配 assert!(filter.matches("file_read", &serde_json::json!({"path": "/app/.env"}))); // 正确工具名 + 安全路径 → 不匹配 assert!(!filter.matches("file_read", &serde_json::json!({"path": "/app/README.md"}))); // 错误工具名 + 敏感路径 → 不匹配 assert!(!filter.matches("run_bash", &serde_json::json!({"path": "/app/.env"}))); } #[test] fn test_filter_empty_content_false() { // 内容模式指定了不存在的字段 → 不匹配 let filter = ToolMatchFilter { name_patterns: vec![], content_patterns: vec!["nonexistent:value".to_string()], }; assert!(!filter.matches("run_bash", &serde_json::json!({"command": "ls"}))); } #[test] fn test_is_match_all() { assert!(ToolMatchFilter::default().is_match_all()); assert!(!ToolMatchFilter::exact("foo").is_match_all()); assert!(!ToolMatchFilter { name_patterns: vec![], content_patterns: vec!["cmd:ls".to_string()] } .is_match_all()); } }