Skip to main content

xz_agent_hooks/
registry.rs

1//! Ordered hook registration and firing.
2
3use crate::contract::{
4    merge_outcomes, HookEvent, HookEventKind, HookOutcome, MergeMode, MergedEffect,
5};
6use async_trait::async_trait;
7use std::time::Duration;
8
9/// One registered hook handler.
10#[async_trait]
11pub trait HookHandler: Send + Sync {
12    /// Stable name for logging.
13    fn name(&self) -> &str;
14
15    /// Event kinds this handler listens to. Empty = all kinds.
16    fn event_kinds(&self) -> &[HookEventKind];
17
18    /// Optional tool-name matcher (`*`, `?`, `|` alternation). `None` = all tools.
19    fn tool_matcher(&self) -> Option<&str> {
20        None
21    }
22
23    /// Whether this handler should run for `event`.
24    fn matches(&self, event: &HookEvent) -> bool {
25        let kinds = self.event_kinds();
26        if !kinds.is_empty() && !kinds.contains(&event.kind) {
27            return false;
28        }
29        if let Some(pat) = self.tool_matcher() {
30            match event.tool.as_deref() {
31                Some(tool) => {
32                    if !glob_match(pat, tool) {
33                        return false;
34                    }
35                }
36                None => {
37                    if matches!(
38                        event.kind,
39                        HookEventKind::PreTool
40                            | HookEventKind::PostTool
41                            | HookEventKind::PermissionRequest
42                    ) {
43                        return false;
44                    }
45                }
46            }
47        }
48        true
49    }
50
51    /// Run the handler.
52    async fn on_event(&self, event: &HookEvent) -> Result<Vec<HookOutcome>, String>;
53}
54
55/// Aggregated result of firing all matching handlers.
56#[derive(Debug, Clone)]
57pub struct HookFireResult {
58    /// Outcomes in registration order (before merge).
59    pub outcomes: Vec<HookOutcome>,
60    /// Merged typed effect for the requested mode.
61    pub effect: MergedEffect,
62    /// Handler errors (fail-open); product may log.
63    pub errors: Vec<String>,
64}
65
66impl Default for HookFireResult {
67    fn default() -> Self {
68        Self {
69            outcomes: Vec::new(),
70            effect: MergedEffect::Inject(Default::default()),
71            errors: Vec::new(),
72        }
73    }
74}
75
76/// Ordered registry of hook handlers.
77///
78/// Registration order is preserved for deterministic mutate chains.
79pub struct HookRegistry {
80    handlers: Vec<Box<dyn HookHandler>>,
81    /// Default timeout hint for products spawning external work (not enforced here).
82    pub default_timeout: Duration,
83}
84
85impl Default for HookRegistry {
86    fn default() -> Self {
87        Self::new()
88    }
89}
90
91impl HookRegistry {
92    /// Empty registry.
93    pub fn new() -> Self {
94        Self {
95            handlers: Vec::new(),
96            default_timeout: Duration::from_secs(30),
97        }
98    }
99
100    /// Register a handler (appended).
101    pub fn register(&mut self, handler: Box<dyn HookHandler>) {
102        self.handlers.push(handler);
103    }
104
105    /// Remove all handlers with this name. Returns how many were removed.
106    pub fn remove(&mut self, name: &str) -> usize {
107        let before = self.handlers.len();
108        self.handlers.retain(|h| h.name() != name);
109        before - self.handlers.len()
110    }
111
112    /// Clear all handlers.
113    pub fn clear(&mut self) {
114        self.handlers.clear();
115    }
116
117    /// Number of handlers.
118    pub fn len(&self) -> usize {
119        self.handlers.len()
120    }
121
122    /// Whether no handlers are registered.
123    pub fn is_empty(&self) -> bool {
124        self.handlers.is_empty()
125    }
126
127    /// Handler names in order.
128    pub fn names(&self) -> Vec<String> {
129        self.handlers.iter().map(|h| h.name().to_string()).collect()
130    }
131
132    /// Fire all matching handlers; merge with `mode`.
133    ///
134    /// On PreTool, **stops invoking further handlers** once a Deny outcome is produced.
135    /// Errors from handlers are collected and treated as fail-open.
136    pub async fn fire(&self, event: &HookEvent, mode: MergeMode) -> HookFireResult {
137        let mut result = HookFireResult {
138            outcomes: Vec::new(),
139            effect: merge_outcomes(mode, &[]),
140            errors: Vec::new(),
141        };
142        for h in &self.handlers {
143            if !h.matches(event) {
144                continue;
145            }
146            match h.on_event(event).await {
147                Ok(outcomes) => {
148                    let denied = outcomes
149                        .iter()
150                        .any(|o| matches!(o, HookOutcome::Deny { .. }));
151                    result.outcomes.extend(outcomes);
152                    if denied && mode == MergeMode::PreTool {
153                        break;
154                    }
155                }
156                Err(e) => {
157                    tracing::warn!(hook = h.name(), error = %e, "hook handler error (fail-open)");
158                    result.errors.push(format!("{}: {e}", h.name()));
159                }
160            }
161        }
162        result.effect = merge_outcomes(mode, &result.outcomes);
163        result
164    }
165
166    /// Fire using the event kind's default merge mode.
167    pub async fn fire_default(&self, event: &HookEvent) -> HookFireResult {
168        self.fire(event, event.kind.default_merge_mode()).await
169    }
170
171    /// Convenience: PreTool fire + merge.
172    pub async fn fire_pre_tool(&self, event: &HookEvent) -> HookFireResult {
173        self.fire(event, MergeMode::PreTool).await
174    }
175
176    /// Convenience: PostTool fire + merge.
177    pub async fn fire_post_tool(&self, event: &HookEvent) -> HookFireResult {
178        self.fire(event, MergeMode::PostTool).await
179    }
180
181    /// Convenience: PermissionRequest fire + merge.
182    pub async fn fire_permission(&self, event: &HookEvent) -> HookFireResult {
183        self.fire(event, MergeMode::PermissionRequest).await
184    }
185}
186
187/// Simple glob: `*` any chars, `?` one char; `|` alternation.
188///
189/// Exact match (no wildcards) is case-insensitive. Patterns with `*`/`?` are case-sensitive.
190/// Claude-style `Bash(npm:*)` is reduced to the tool name before `(`.
191pub fn glob_match(pattern: &str, text: &str) -> bool {
192    if pattern == "*" || pattern.is_empty() {
193        return true;
194    }
195    if pattern.contains('|') {
196        return pattern.split('|').any(|p| glob_match(p.trim(), text));
197    }
198    let pat = pattern.split('(').next().unwrap_or(pattern).trim();
199    if !pat.contains('*') && !pat.contains('?') {
200        return pat.eq_ignore_ascii_case(text);
201    }
202    let pat: Vec<char> = pat.chars().collect();
203    let text: Vec<char> = text.chars().collect();
204    match_glob(&pat, &text)
205}
206
207fn match_glob(pat: &[char], text: &[char]) -> bool {
208    let mut pi = 0;
209    let mut ti = 0;
210    let mut star_p = None;
211    let mut star_t = 0;
212    while ti < text.len() {
213        if pi < pat.len() && (pat[pi] == '?' || pat[pi] == text[ti]) {
214            pi += 1;
215            ti += 1;
216        } else if pi < pat.len() && pat[pi] == '*' {
217            star_p = Some(pi);
218            star_t = ti;
219            pi += 1;
220        } else if let Some(sp) = star_p {
221            pi = sp + 1;
222            star_t += 1;
223            ti = star_t;
224        } else {
225            return false;
226        }
227    }
228    while pi < pat.len() && pat[pi] == '*' {
229        pi += 1;
230    }
231    pi == pat.len()
232}
233
234#[cfg(test)]
235mod tests {
236    use super::*;
237    use crate::contract::{PermissionDecision, PreToolEffect};
238    use serde_json::json;
239
240    struct StaticHook {
241        name: String,
242        kinds: Vec<HookEventKind>,
243        matcher: Option<String>,
244        outcomes: Vec<HookOutcome>,
245        fail: bool,
246    }
247
248    #[async_trait]
249    impl HookHandler for StaticHook {
250        fn name(&self) -> &str {
251            &self.name
252        }
253        fn event_kinds(&self) -> &[HookEventKind] {
254            &self.kinds
255        }
256        fn tool_matcher(&self) -> Option<&str> {
257            self.matcher.as_deref()
258        }
259        async fn on_event(&self, _event: &HookEvent) -> Result<Vec<HookOutcome>, String> {
260            if self.fail {
261                return Err("boom".into());
262            }
263            Ok(self.outcomes.clone())
264        }
265    }
266
267    fn hook(
268        name: &str,
269        kinds: &[HookEventKind],
270        matcher: Option<&str>,
271        outcomes: Vec<HookOutcome>,
272    ) -> StaticHook {
273        StaticHook {
274            name: name.into(),
275            kinds: kinds.to_vec(),
276            matcher: matcher.map(str::to_string),
277            outcomes,
278            fail: false,
279        }
280    }
281
282    #[tokio::test]
283    async fn ordered_mutate_and_deny_short_circuit() {
284        let mut reg = HookRegistry::new();
285        reg.register(Box::new(hook(
286            "a",
287            &[HookEventKind::PreTool],
288            None,
289            vec![HookOutcome::mutate_args(json!({"command": "one"}))],
290        )));
291        reg.register(Box::new(hook(
292            "b",
293            &[HookEventKind::PreTool],
294            None,
295            vec![HookOutcome::deny("stop")],
296        )));
297        reg.register(Box::new(hook(
298            "c",
299            &[HookEventKind::PreTool],
300            None,
301            vec![HookOutcome::mutate_args(json!({"command": "three"}))],
302        )));
303
304        let ev = HookEvent::pre_tool("shell", json!({}));
305        let r = reg.fire_pre_tool(&ev).await;
306        let PreToolEffect { deny, args, .. } = r.effect.as_pre_tool().cloned().unwrap_or_default();
307        assert_eq!(deny.as_deref(), Some("stop"));
308        assert_eq!(args, Some(json!({"command": "one"})));
309        assert_eq!(r.outcomes.len(), 2);
310    }
311
312    #[tokio::test]
313    async fn fail_open_on_handler_error() {
314        let mut reg = HookRegistry::new();
315        reg.register(Box::new(StaticHook {
316            name: "bad".into(),
317            kinds: vec![],
318            matcher: None,
319            outcomes: vec![],
320            fail: true,
321        }));
322        reg.register(Box::new(hook(
323            "good",
324            &[],
325            None,
326            vec![HookOutcome::context("ok")],
327        )));
328        let r = reg
329            .fire(
330                &HookEvent::unit(HookEventKind::SessionStart),
331                MergeMode::InjectOnly,
332            )
333            .await;
334        assert_eq!(r.errors.len(), 1);
335        assert_eq!(r.effect.contexts().len(), 1);
336    }
337
338    #[tokio::test]
339    async fn tool_matcher_filters() {
340        let mut reg = HookRegistry::new();
341        reg.register(Box::new(hook(
342            "only-shell",
343            &[HookEventKind::PreTool],
344            Some("shell|Bash"),
345            vec![HookOutcome::deny("x")],
346        )));
347        let r = reg
348            .fire_pre_tool(&HookEvent::pre_tool("FileRead", json!({})))
349            .await;
350        assert!(r.outcomes.is_empty());
351        let r = reg
352            .fire_pre_tool(&HookEvent::pre_tool("Bash", json!({})))
353            .await;
354        assert!(r.effect.as_pre_tool().is_some_and(|e| e.is_denied()));
355    }
356
357    #[tokio::test]
358    async fn kind_filter() {
359        let mut reg = HookRegistry::new();
360        reg.register(Box::new(hook(
361            "pre-only",
362            &[HookEventKind::PreTool],
363            None,
364            vec![HookOutcome::context("x")],
365        )));
366        let r = reg
367            .fire_default(&HookEvent::unit(HookEventKind::SessionStart))
368            .await;
369        assert!(r.outcomes.is_empty());
370    }
371
372    #[tokio::test]
373    async fn remove_and_clear() {
374        let mut reg = HookRegistry::new();
375        reg.register(Box::new(hook("a", &[], None, vec![])));
376        reg.register(Box::new(hook("b", &[], None, vec![])));
377        assert_eq!(reg.len(), 2);
378        assert_eq!(reg.remove("a"), 1);
379        assert_eq!(reg.names(), vec!["b".to_string()]);
380        reg.clear();
381        assert!(reg.is_empty());
382    }
383
384    #[tokio::test]
385    async fn post_tool_and_permission_helpers() {
386        let mut reg = HookRegistry::new();
387        reg.register(Box::new(hook(
388            "p",
389            &[HookEventKind::PostTool],
390            None,
391            vec![HookOutcome::replace_result("new")],
392        )));
393        reg.register(Box::new(hook(
394            "perm",
395            &[HookEventKind::PermissionRequest],
396            None,
397            vec![HookOutcome::Allow],
398        )));
399        let r = reg
400            .fire_post_tool(&HookEvent::post_tool("t", "old"))
401            .await;
402        assert_eq!(
403            r.effect
404                .as_post_tool()
405                .and_then(|e| e.replace_result.as_deref()),
406            Some("new")
407        );
408        let r = reg
409            .fire_permission(&HookEvent::permission_request("t", json!({})))
410            .await;
411        assert_eq!(
412            r.effect.as_permission().map(|e| &e.decision),
413            Some(&PermissionDecision::Allow)
414        );
415    }
416
417    #[test]
418    fn glob_cases() {
419        assert!(glob_match("Bash|Shell", "Shell"));
420        assert!(glob_match("File*", "FileRead"));
421        assert!(!glob_match("Git*", "FileRead"));
422        assert!(glob_match("bash", "Bash")); // exact case-insensitive
423        assert!(glob_match("*", "anything"));
424        assert!(glob_match("Bash(npm:*)", "Bash"));
425        assert!(glob_match("f?o", "foo"));
426        assert!(!glob_match("f?o", "fooo"));
427    }
428}