Skip to main content

oxicode_sdk/middleware/
hook.rs

1//! [`HookMiddleware`](crate::middleware::HookMiddleware) — bridge [`HookRunner`](crate::ports::HookRunner) into the existing
2//! [`MiddlewarePipeline`] so Pre/PostToolUse hooks fire through the
3//! same path as audit/authorizer middlewares.
4//!
5//! SubagentStop is fired here as a side effect: when an `AfterTool` call
6//! has `tool_name == "subagent"`, we additionally invoke
7//! `runner.run(SubagentStop, ctx)` so users only need a single matcher
8//! rule. SessionStart / SessionEnd / Stop / Notification are NOT
9//! fired here — those are product-lifecycle events owned by the
10//! composition root.
11
12use std::path::PathBuf;
13use std::pin::Pin;
14use std::sync::Arc;
15
16use crate::middleware::{
17    Middleware, MiddlewareAction, MiddlewareContext, MiddlewareData, MiddlewarePhase,
18    MiddlewareResult,
19};
20use crate::ports::{HookContext, HookEvent, HookRunner};
21
22const SUBAGENT_TOOL_NAME: &str = "subagent";
23
24/// Middleware that routes `BeforeTool` / `AfterTool` phases through the
25/// registered [`HookRunner`] as `PreToolUse` / `PostToolUse` events.
26pub struct HookMiddleware {
27    runner: Arc<dyn HookRunner>,
28    session_id: Option<String>,
29    session_cwd: Option<PathBuf>,
30}
31
32impl std::fmt::Debug for HookMiddleware {
33    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
34        f.debug_struct("HookMiddleware")
35            .field("session_id", &self.session_id)
36            .finish()
37    }
38}
39
40impl HookMiddleware {
41    /// Wrap a [`HookRunner`] as a middleware. The runner is usually the
42    /// engine's registered `CommandHookRunner` (port #16).
43    pub fn new(runner: Arc<dyn HookRunner>) -> Self {
44        Self {
45            runner,
46            session_id: None,
47            session_cwd: None,
48        }
49    }
50
51    /// Tag every emitted [`HookContext`] with a session id + cwd so hook
52    /// scripts can identify which session triggered them.
53    pub fn with_session(mut self, id: String, cwd: PathBuf) -> Self {
54        self.session_id = Some(id);
55        self.session_cwd = Some(cwd);
56        self
57    }
58}
59
60impl Middleware for HookMiddleware {
61    fn name(&self) -> &str {
62        "HookMiddleware"
63    }
64
65    fn phases(&self) -> Vec<MiddlewarePhase> {
66        vec![MiddlewarePhase::BeforeTool, MiddlewarePhase::AfterTool]
67    }
68
69    fn handle<'a>(
70        &'a self,
71        ctx: &'a MiddlewareContext,
72    ) -> Pin<Box<dyn Future<Output = MiddlewareResult> + Send + 'a>> {
73        let (event, tool_name, args, result) = match &ctx.data {
74            MiddlewareData::BeforeTool { tool_name, params } => (
75                HookEvent::PreToolUse,
76                tool_name.clone(),
77                params.clone(),
78                None,
79            ),
80            MiddlewareData::AfterTool {
81                tool_name,
82                params: _,
83                result,
84            } => (
85                HookEvent::PostToolUse,
86                tool_name.clone(),
87                serde_json::Value::Null,
88                Some(result.clone()),
89            ),
90            _ => return Box::pin(async { MiddlewareResult::pass() }),
91        };
92
93        let runner = Arc::clone(&self.runner);
94        let session_id = self.session_id.clone();
95        let session_cwd = self.session_cwd.clone();
96        let is_after = matches!(ctx.phase, MiddlewarePhase::AfterTool);
97
98        Box::pin(async move {
99            let hook_ctx = HookContext {
100                event,
101                tool_name: Some(tool_name.clone()),
102                tool_args: if args.is_null() { None } else { Some(args) },
103                tool_result: result,
104                is_error: None,
105                session_id,
106                session_cwd,
107                extra: None,
108            };
109            let outcome = runner.run(event, &hook_ctx).await;
110            if outcome.block {
111                return MiddlewareResult {
112                    action: MiddlewareAction::Block,
113                    modified_data: None,
114                    reason: outcome.reason.or(Some(format!(
115                        "hook {:?} denied tool `{}`",
116                        event, tool_name
117                    ))),
118                };
119            }
120
121            // SubagentStop is fired as a side effect of the `subagent`
122            // tool completing. We don't block on it (SubagentStop is
123            // notification-only by design); we just route through so
124            // users can react to subagent completion.
125            if is_after && tool_name == SUBAGENT_TOOL_NAME {
126                let sub_ctx = HookContext {
127                    event: HookEvent::SubagentStop,
128                    ..hook_ctx
129                };
130                let _ = runner.run(HookEvent::SubagentStop, &sub_ctx).await;
131            }
132
133            MiddlewareResult::pass()
134        })
135    }
136}
137
138#[cfg(test)]
139mod tests {
140    use super::*;
141    use crate::ports::HookOutcome;
142    use crate::ports::inmem::InMemoryHookRunner;
143    use serde_json::json;
144
145    fn before_tool_ctx(tool_name: &str) -> MiddlewareContext {
146        MiddlewareContext::new(
147            MiddlewarePhase::BeforeTool,
148            "agent-1",
149            MiddlewareData::BeforeTool {
150                tool_name: tool_name.into(),
151                params: json!({"command": "ls"}),
152            },
153        )
154    }
155
156    fn after_tool_ctx(tool_name: &str, result: &str) -> MiddlewareContext {
157        MiddlewareContext::new(
158            MiddlewarePhase::AfterTool,
159            "agent-1",
160            MiddlewareData::AfterTool {
161                tool_name: tool_name.into(),
162                params: json!({}),
163                result: result.into(),
164            },
165        )
166    }
167
168    #[tokio::test]
169    async fn before_tool_pass_through_when_no_handlers() {
170        let mw = HookMiddleware::new(Arc::new(InMemoryHookRunner::new()));
171        let result = mw.handle(&before_tool_ctx("bash")).await;
172        assert!(result.is_continue());
173    }
174
175    #[tokio::test]
176    async fn before_tool_block_short_circuits_tool() {
177        let runner = InMemoryHookRunner::new();
178        runner.on(|_, _| HookOutcome {
179            block: true,
180            reason: Some("deny".into()),
181            ..Default::default()
182        });
183        let mw = HookMiddleware::new(Arc::new(runner));
184        let result = mw.handle(&before_tool_ctx("bash")).await;
185        assert!(matches!(result.action, MiddlewareAction::Block));
186        assert_eq!(result.reason.as_deref(), Some("deny"));
187    }
188
189    #[tokio::test]
190    async fn after_subagent_fires_subagent_stop() {
191        let runner = InMemoryHookRunner::new();
192        let counter = Arc::new(std::sync::atomic::AtomicUsize::new(0));
193        let c = Arc::clone(&counter);
194        runner.on(move |event, _| {
195            if event == HookEvent::SubagentStop {
196                c.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
197            }
198            HookOutcome::default()
199        });
200        let mw = HookMiddleware::new(Arc::new(runner));
201        mw.handle(&after_tool_ctx("subagent", "{}")).await;
202        assert_eq!(counter.load(std::sync::atomic::Ordering::SeqCst), 1);
203    }
204
205    #[tokio::test]
206    async fn after_non_subagent_does_not_fire_subagent_stop() {
207        let runner = InMemoryHookRunner::new();
208        let counter = Arc::new(std::sync::atomic::AtomicUsize::new(0));
209        let c = Arc::clone(&counter);
210        runner.on(move |event, _| {
211            if event == HookEvent::SubagentStop {
212                c.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
213            }
214            HookOutcome::default()
215        });
216        let mw = HookMiddleware::new(Arc::new(runner));
217        mw.handle(&after_tool_ctx("read", "ok")).await;
218        assert_eq!(counter.load(std::sync::atomic::Ordering::SeqCst), 0);
219    }
220}