use crate::contract::{
merge_outcomes, HookEvent, HookEventKind, HookOutcome, MergeMode, MergedEffect,
};
use async_trait::async_trait;
use std::time::Duration;
#[async_trait]
pub trait HookHandler: Send + Sync {
fn name(&self) -> &str;
fn event_kinds(&self) -> &[HookEventKind];
fn tool_matcher(&self) -> Option<&str> {
None
}
fn matches(&self, event: &HookEvent) -> bool {
let kinds = self.event_kinds();
if !kinds.is_empty() && !kinds.contains(&event.kind) {
return false;
}
if let Some(pat) = self.tool_matcher() {
match event.tool.as_deref() {
Some(tool) => {
if !glob_match(pat, tool) {
return false;
}
}
None => {
if matches!(
event.kind,
HookEventKind::PreTool
| HookEventKind::PostTool
| HookEventKind::PermissionRequest
) {
return false;
}
}
}
}
true
}
async fn on_event(&self, event: &HookEvent) -> Result<Vec<HookOutcome>, String>;
}
#[derive(Debug, Clone)]
pub struct HookFireResult {
pub outcomes: Vec<HookOutcome>,
pub effect: MergedEffect,
pub errors: Vec<String>,
}
impl Default for HookFireResult {
fn default() -> Self {
Self {
outcomes: Vec::new(),
effect: MergedEffect::Inject(Default::default()),
errors: Vec::new(),
}
}
}
pub struct HookRegistry {
handlers: Vec<Box<dyn HookHandler>>,
pub default_timeout: Duration,
}
impl Default for HookRegistry {
fn default() -> Self {
Self::new()
}
}
impl HookRegistry {
pub fn new() -> Self {
Self {
handlers: Vec::new(),
default_timeout: Duration::from_secs(30),
}
}
pub fn register(&mut self, handler: Box<dyn HookHandler>) {
self.handlers.push(handler);
}
pub fn remove(&mut self, name: &str) -> usize {
let before = self.handlers.len();
self.handlers.retain(|h| h.name() != name);
before - self.handlers.len()
}
pub fn clear(&mut self) {
self.handlers.clear();
}
pub fn len(&self) -> usize {
self.handlers.len()
}
pub fn is_empty(&self) -> bool {
self.handlers.is_empty()
}
pub fn names(&self) -> Vec<String> {
self.handlers.iter().map(|h| h.name().to_string()).collect()
}
pub async fn fire(&self, event: &HookEvent, mode: MergeMode) -> HookFireResult {
let mut result = HookFireResult {
outcomes: Vec::new(),
effect: merge_outcomes(mode, &[]),
errors: Vec::new(),
};
for h in &self.handlers {
if !h.matches(event) {
continue;
}
match h.on_event(event).await {
Ok(outcomes) => {
let denied = outcomes
.iter()
.any(|o| matches!(o, HookOutcome::Deny { .. }));
result.outcomes.extend(outcomes);
if denied && mode == MergeMode::PreTool {
break;
}
}
Err(e) => {
tracing::warn!(hook = h.name(), error = %e, "hook handler error (fail-open)");
result.errors.push(format!("{}: {e}", h.name()));
}
}
}
result.effect = merge_outcomes(mode, &result.outcomes);
result
}
pub async fn fire_default(&self, event: &HookEvent) -> HookFireResult {
self.fire(event, event.kind.default_merge_mode()).await
}
pub async fn fire_pre_tool(&self, event: &HookEvent) -> HookFireResult {
self.fire(event, MergeMode::PreTool).await
}
pub async fn fire_post_tool(&self, event: &HookEvent) -> HookFireResult {
self.fire(event, MergeMode::PostTool).await
}
pub async fn fire_permission(&self, event: &HookEvent) -> HookFireResult {
self.fire(event, MergeMode::PermissionRequest).await
}
}
pub fn glob_match(pattern: &str, text: &str) -> bool {
if pattern == "*" || pattern.is_empty() {
return true;
}
if pattern.contains('|') {
return pattern.split('|').any(|p| glob_match(p.trim(), text));
}
let pat = pattern.split('(').next().unwrap_or(pattern).trim();
if !pat.contains('*') && !pat.contains('?') {
return pat.eq_ignore_ascii_case(text);
}
let pat: Vec<char> = pat.chars().collect();
let text: Vec<char> = text.chars().collect();
match_glob(&pat, &text)
}
fn match_glob(pat: &[char], text: &[char]) -> bool {
let mut pi = 0;
let mut ti = 0;
let mut star_p = None;
let mut star_t = 0;
while ti < text.len() {
if pi < pat.len() && (pat[pi] == '?' || pat[pi] == text[ti]) {
pi += 1;
ti += 1;
} else if pi < pat.len() && pat[pi] == '*' {
star_p = Some(pi);
star_t = ti;
pi += 1;
} else if let Some(sp) = star_p {
pi = sp + 1;
star_t += 1;
ti = star_t;
} else {
return false;
}
}
while pi < pat.len() && pat[pi] == '*' {
pi += 1;
}
pi == pat.len()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::contract::{PermissionDecision, PreToolEffect};
use serde_json::json;
struct StaticHook {
name: String,
kinds: Vec<HookEventKind>,
matcher: Option<String>,
outcomes: Vec<HookOutcome>,
fail: bool,
}
#[async_trait]
impl HookHandler for StaticHook {
fn name(&self) -> &str {
&self.name
}
fn event_kinds(&self) -> &[HookEventKind] {
&self.kinds
}
fn tool_matcher(&self) -> Option<&str> {
self.matcher.as_deref()
}
async fn on_event(&self, _event: &HookEvent) -> Result<Vec<HookOutcome>, String> {
if self.fail {
return Err("boom".into());
}
Ok(self.outcomes.clone())
}
}
fn hook(
name: &str,
kinds: &[HookEventKind],
matcher: Option<&str>,
outcomes: Vec<HookOutcome>,
) -> StaticHook {
StaticHook {
name: name.into(),
kinds: kinds.to_vec(),
matcher: matcher.map(str::to_string),
outcomes,
fail: false,
}
}
#[tokio::test]
async fn ordered_mutate_and_deny_short_circuit() {
let mut reg = HookRegistry::new();
reg.register(Box::new(hook(
"a",
&[HookEventKind::PreTool],
None,
vec![HookOutcome::mutate_args(json!({"command": "one"}))],
)));
reg.register(Box::new(hook(
"b",
&[HookEventKind::PreTool],
None,
vec![HookOutcome::deny("stop")],
)));
reg.register(Box::new(hook(
"c",
&[HookEventKind::PreTool],
None,
vec![HookOutcome::mutate_args(json!({"command": "three"}))],
)));
let ev = HookEvent::pre_tool("shell", json!({}));
let r = reg.fire_pre_tool(&ev).await;
let PreToolEffect { deny, args, .. } = r.effect.as_pre_tool().cloned().unwrap_or_default();
assert_eq!(deny.as_deref(), Some("stop"));
assert_eq!(args, Some(json!({"command": "one"})));
assert_eq!(r.outcomes.len(), 2);
}
#[tokio::test]
async fn fail_open_on_handler_error() {
let mut reg = HookRegistry::new();
reg.register(Box::new(StaticHook {
name: "bad".into(),
kinds: vec![],
matcher: None,
outcomes: vec![],
fail: true,
}));
reg.register(Box::new(hook(
"good",
&[],
None,
vec![HookOutcome::context("ok")],
)));
let r = reg
.fire(
&HookEvent::unit(HookEventKind::SessionStart),
MergeMode::InjectOnly,
)
.await;
assert_eq!(r.errors.len(), 1);
assert_eq!(r.effect.contexts().len(), 1);
}
#[tokio::test]
async fn tool_matcher_filters() {
let mut reg = HookRegistry::new();
reg.register(Box::new(hook(
"only-shell",
&[HookEventKind::PreTool],
Some("shell|Bash"),
vec![HookOutcome::deny("x")],
)));
let r = reg
.fire_pre_tool(&HookEvent::pre_tool("FileRead", json!({})))
.await;
assert!(r.outcomes.is_empty());
let r = reg
.fire_pre_tool(&HookEvent::pre_tool("Bash", json!({})))
.await;
assert!(r.effect.as_pre_tool().is_some_and(|e| e.is_denied()));
}
#[tokio::test]
async fn kind_filter() {
let mut reg = HookRegistry::new();
reg.register(Box::new(hook(
"pre-only",
&[HookEventKind::PreTool],
None,
vec![HookOutcome::context("x")],
)));
let r = reg
.fire_default(&HookEvent::unit(HookEventKind::SessionStart))
.await;
assert!(r.outcomes.is_empty());
}
#[tokio::test]
async fn remove_and_clear() {
let mut reg = HookRegistry::new();
reg.register(Box::new(hook("a", &[], None, vec![])));
reg.register(Box::new(hook("b", &[], None, vec![])));
assert_eq!(reg.len(), 2);
assert_eq!(reg.remove("a"), 1);
assert_eq!(reg.names(), vec!["b".to_string()]);
reg.clear();
assert!(reg.is_empty());
}
#[tokio::test]
async fn post_tool_and_permission_helpers() {
let mut reg = HookRegistry::new();
reg.register(Box::new(hook(
"p",
&[HookEventKind::PostTool],
None,
vec![HookOutcome::replace_result("new")],
)));
reg.register(Box::new(hook(
"perm",
&[HookEventKind::PermissionRequest],
None,
vec![HookOutcome::Allow],
)));
let r = reg
.fire_post_tool(&HookEvent::post_tool("t", "old"))
.await;
assert_eq!(
r.effect
.as_post_tool()
.and_then(|e| e.replace_result.as_deref()),
Some("new")
);
let r = reg
.fire_permission(&HookEvent::permission_request("t", json!({})))
.await;
assert_eq!(
r.effect.as_permission().map(|e| &e.decision),
Some(&PermissionDecision::Allow)
);
}
#[test]
fn glob_cases() {
assert!(glob_match("Bash|Shell", "Shell"));
assert!(glob_match("File*", "FileRead"));
assert!(!glob_match("Git*", "FileRead"));
assert!(glob_match("bash", "Bash")); assert!(glob_match("*", "anything"));
assert!(glob_match("Bash(npm:*)", "Bash"));
assert!(glob_match("f?o", "foo"));
assert!(!glob_match("f?o", "fooo"));
}
}