oxicode_sdk/middleware/
hook.rs1use 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
24pub 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 pub fn new(runner: Arc<dyn HookRunner>) -> Self {
44 Self {
45 runner,
46 session_id: None,
47 session_cwd: None,
48 }
49 }
50
51 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 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}