1use crate::contract::{
4 merge_outcomes, HookEvent, HookEventKind, HookOutcome, MergeMode, MergedEffect,
5};
6use async_trait::async_trait;
7use std::time::Duration;
8
9#[async_trait]
11pub trait HookHandler: Send + Sync {
12 fn name(&self) -> &str;
14
15 fn event_kinds(&self) -> &[HookEventKind];
17
18 fn tool_matcher(&self) -> Option<&str> {
20 None
21 }
22
23 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 async fn on_event(&self, event: &HookEvent) -> Result<Vec<HookOutcome>, String>;
53}
54
55#[derive(Debug, Clone)]
57pub struct HookFireResult {
58 pub outcomes: Vec<HookOutcome>,
60 pub effect: MergedEffect,
62 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
76pub struct HookRegistry {
80 handlers: Vec<Box<dyn HookHandler>>,
81 pub default_timeout: Duration,
83}
84
85impl Default for HookRegistry {
86 fn default() -> Self {
87 Self::new()
88 }
89}
90
91impl HookRegistry {
92 pub fn new() -> Self {
94 Self {
95 handlers: Vec::new(),
96 default_timeout: Duration::from_secs(30),
97 }
98 }
99
100 pub fn register(&mut self, handler: Box<dyn HookHandler>) {
102 self.handlers.push(handler);
103 }
104
105 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 pub fn clear(&mut self) {
114 self.handlers.clear();
115 }
116
117 pub fn len(&self) -> usize {
119 self.handlers.len()
120 }
121
122 pub fn is_empty(&self) -> bool {
124 self.handlers.is_empty()
125 }
126
127 pub fn names(&self) -> Vec<String> {
129 self.handlers.iter().map(|h| h.name().to_string()).collect()
130 }
131
132 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 pub async fn fire_default(&self, event: &HookEvent) -> HookFireResult {
168 self.fire(event, event.kind.default_merge_mode()).await
169 }
170
171 pub async fn fire_pre_tool(&self, event: &HookEvent) -> HookFireResult {
173 self.fire(event, MergeMode::PreTool).await
174 }
175
176 pub async fn fire_post_tool(&self, event: &HookEvent) -> HookFireResult {
178 self.fire(event, MergeMode::PostTool).await
179 }
180
181 pub async fn fire_permission(&self, event: &HookEvent) -> HookFireResult {
183 self.fire(event, MergeMode::PermissionRequest).await
184 }
185}
186
187pub 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")); 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}