Skip to main content

oxicode_sdk/ports/inmem/
hook.rs

1//! In-memory [`HookRunner`] for tests and headless products.
2
3use std::pin::Pin;
4use std::sync::Arc;
5
6use parking_lot::Mutex;
7
8use crate::ports::{HookContext, HookEvent, HookOutcome, HookRunner};
9
10type Handler = Arc<dyn Fn(HookEvent, &HookContext) -> HookOutcome + Send + Sync>;
11
12/// Test hook runner — handlers are registered as plain closures.
13/// All handlers fire on every event; the first one that returns
14/// `block = true` short-circuits.
15#[derive(Default)]
16pub struct InMemoryHookRunner {
17    handlers: Mutex<Vec<Handler>>,
18}
19
20impl std::fmt::Debug for InMemoryHookRunner {
21    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
22        f.debug_struct("InMemoryHookRunner")
23            .field("handler_count", &self.handlers.lock().len())
24            .finish()
25    }
26}
27
28impl InMemoryHookRunner {
29    /// Create a runner with no handlers.
30    pub fn new() -> Self {
31        Self::default()
32    }
33
34    /// Register a handler. Handlers run in registration order.
35    pub fn on<F>(&self, f: F)
36    where
37        F: Fn(HookEvent, &HookContext) -> HookOutcome + Send + Sync + 'static,
38    {
39        self.handlers.lock().push(Arc::new(f));
40    }
41}
42
43impl HookRunner for InMemoryHookRunner {
44    fn run<'a>(
45        &'a self,
46        event: HookEvent,
47        ctx: &'a HookContext,
48    ) -> Pin<Box<dyn Future<Output = HookOutcome> + Send + 'a>> {
49        let handlers = self.handlers.lock().clone();
50        let ctx = ctx.clone();
51        Box::pin(async move {
52            let mut out = HookOutcome::default();
53            for h in &handlers {
54                let step = h(event, &ctx);
55                if step.block {
56                    return HookOutcome {
57                        block: true,
58                        reason: step.reason.or(out.reason),
59                        override_content: step.override_content.or(out.override_content),
60                    };
61                }
62                if step.override_content.is_some() {
63                    out.override_content = step.override_content;
64                }
65                if step.reason.is_some() {
66                    out.reason = step.reason;
67                }
68            }
69            out
70        })
71    }
72}
73
74#[cfg(test)]
75mod tests {
76    use super::*;
77
78    #[tokio::test]
79    async fn handler_fires_and_blocks() {
80        let runner = InMemoryHookRunner::new();
81        runner.on(|_, _| HookOutcome {
82            block: true,
83            reason: Some("blocked".into()),
84            ..Default::default()
85        });
86        let ctx = HookContext::default();
87        let out = runner.run(HookEvent::PreToolUse, &ctx).await;
88        assert!(out.block);
89        assert_eq!(out.reason.as_deref(), Some("blocked"));
90    }
91
92    #[tokio::test]
93    async fn empty_runner_returns_default() {
94        let runner = InMemoryHookRunner::new();
95        let ctx = HookContext::default();
96        let out = runner.run(HookEvent::PreToolUse, &ctx).await;
97        assert!(!out.block);
98    }
99
100    #[tokio::test]
101    async fn block_short_circuits() {
102        let runner = InMemoryHookRunner::new();
103        let counter = Arc::new(std::sync::atomic::AtomicUsize::new(0));
104        let c2 = Arc::clone(&counter);
105        let c3 = Arc::clone(&counter);
106        runner.on(move |_, _| {
107            c2.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
108            HookOutcome::default()
109        });
110        runner.on(|_, _| HookOutcome {
111            block: true,
112            ..Default::default()
113        });
114        runner.on(move |_, _| {
115            // Should not run.
116            c3.fetch_add(100, std::sync::atomic::Ordering::SeqCst);
117            HookOutcome::default()
118        });
119        let ctx = HookContext::default();
120        let out = runner.run(HookEvent::PreToolUse, &ctx).await;
121        assert!(out.block);
122        assert_eq!(counter.load(std::sync::atomic::Ordering::SeqCst), 1);
123    }
124}