1use std::path::PathBuf;
9use std::time::Instant;
10
11use async_trait::async_trait;
12use serde::{Deserialize, Serialize};
13use serde_json::Value;
14
15use crate::types::SessionId;
16
17#[derive(Debug, Clone)]
19pub struct HookContext {
20 pub session_id: SessionId,
22}
23
24#[derive(Debug, Clone, Deserialize)]
26#[serde(rename_all = "camelCase")]
27pub struct PreToolUseInput {
28 pub session_id: String,
30 pub timestamp: f64,
32 #[serde(rename = "cwd")]
34 pub working_directory: PathBuf,
35 pub tool_name: String,
37 pub tool_args: Value,
39}
40
41#[derive(Debug, Clone, Default, Serialize)]
43#[serde(rename_all = "camelCase")]
44pub struct PreToolUseOutput {
45 #[serde(skip_serializing_if = "Option::is_none")]
47 pub permission_decision: Option<String>,
48 #[serde(skip_serializing_if = "Option::is_none")]
50 pub permission_decision_reason: Option<String>,
51 #[serde(skip_serializing_if = "Option::is_none")]
53 pub modified_args: Option<Value>,
54 #[serde(skip_serializing_if = "Option::is_none")]
56 pub additional_context: Option<String>,
57 #[serde(skip_serializing_if = "Option::is_none")]
59 pub suppress_output: Option<bool>,
60}
61
62#[derive(Debug, Clone, Deserialize)]
64#[serde(rename_all = "camelCase")]
65pub struct PreMcpToolCallInput {
66 pub session_id: String,
68 pub timestamp: f64,
70 #[serde(rename = "cwd")]
72 pub working_directory: PathBuf,
73 pub server_name: String,
75 pub tool_name: String,
77 pub arguments: Value,
79 #[serde(default)]
81 pub tool_call_id: Option<String>,
82 #[serde(default, rename = "_meta")]
84 pub meta: Option<Value>,
85}
86
87#[derive(Debug, Clone, Default, Serialize)]
94#[serde(rename_all = "camelCase")]
95pub struct PreMcpToolCallOutput {
96 #[serde(skip_serializing_if = "Option::is_none")]
98 pub meta_to_use: Option<Value>,
99}
100
101#[derive(Debug, Clone, Deserialize)]
103#[serde(rename_all = "camelCase")]
104pub struct PostToolUseInput {
105 pub session_id: String,
107 pub timestamp: f64,
109 #[serde(rename = "cwd")]
111 pub working_directory: PathBuf,
112 pub tool_name: String,
114 pub tool_args: Value,
116 pub tool_result: Value,
118}
119
120#[derive(Debug, Clone, Default, Serialize)]
122#[serde(rename_all = "camelCase")]
123pub struct PostToolUseOutput {
124 #[serde(skip_serializing_if = "Option::is_none")]
126 pub modified_result: Option<Value>,
127 #[serde(skip_serializing_if = "Option::is_none")]
129 pub additional_context: Option<String>,
130 #[serde(skip_serializing_if = "Option::is_none")]
132 pub suppress_output: Option<bool>,
133}
134
135#[derive(Debug, Clone, Deserialize)]
143#[serde(rename_all = "camelCase")]
144pub struct PostToolUseFailureInput {
145 pub session_id: String,
147 pub timestamp: f64,
149 #[serde(rename = "cwd")]
151 pub working_directory: PathBuf,
152 pub tool_name: String,
154 pub tool_args: Value,
156 pub error: String,
158}
159
160#[derive(Debug, Clone, Default, Serialize)]
165#[serde(rename_all = "camelCase")]
166pub struct PostToolUseFailureOutput {
167 #[serde(skip_serializing_if = "Option::is_none")]
169 pub additional_context: Option<String>,
170}
171
172#[derive(Debug, Clone, Deserialize)]
174#[serde(rename_all = "camelCase")]
175pub struct UserPromptSubmittedInput {
176 pub session_id: String,
178 pub timestamp: f64,
180 #[serde(rename = "cwd")]
182 pub working_directory: PathBuf,
183 pub prompt: String,
185}
186
187#[derive(Debug, Clone, Default, Serialize)]
189#[serde(rename_all = "camelCase")]
190pub struct UserPromptSubmittedOutput {
191 #[serde(skip_serializing_if = "Option::is_none")]
193 pub modified_prompt: Option<String>,
194 #[serde(skip_serializing_if = "Option::is_none")]
196 pub additional_context: Option<String>,
197 #[serde(skip_serializing_if = "Option::is_none")]
199 pub suppress_output: Option<bool>,
200}
201
202#[derive(Debug, Clone, Deserialize)]
204#[serde(rename_all = "camelCase")]
205pub struct UserPromptTransformedInput {
206 pub session_id: String,
208 pub timestamp: f64,
210 #[serde(rename = "cwd")]
212 pub working_directory: PathBuf,
213 pub prompt: String,
215 pub transformed_prompt: String,
217}
218
219#[derive(Debug, Clone, Default, Serialize)]
221#[serde(rename_all = "camelCase")]
222pub struct UserPromptTransformedOutput {
223 #[serde(skip_serializing_if = "Option::is_none")]
225 pub modified_transformed_prompt: Option<String>,
226}
227
228#[derive(Debug, Clone, Deserialize)]
230#[serde(rename_all = "camelCase")]
231pub struct SessionStartInput {
232 pub session_id: String,
234 pub timestamp: f64,
236 #[serde(rename = "cwd")]
238 pub working_directory: PathBuf,
239 pub source: String,
241 #[serde(default)]
243 pub initial_prompt: Option<String>,
244}
245
246#[derive(Debug, Clone, Default, Serialize)]
248#[serde(rename_all = "camelCase")]
249pub struct SessionStartOutput {
250 #[serde(skip_serializing_if = "Option::is_none")]
252 pub additional_context: Option<String>,
253 #[serde(skip_serializing_if = "Option::is_none")]
255 pub modified_config: Option<Value>,
256}
257
258#[derive(Debug, Clone, Deserialize)]
260#[serde(rename_all = "camelCase")]
261pub struct SessionEndInput {
262 pub session_id: String,
264 pub timestamp: f64,
266 #[serde(rename = "cwd")]
268 pub working_directory: PathBuf,
269 pub reason: String,
271 #[serde(default)]
273 pub final_message: Option<String>,
274 #[serde(default)]
276 pub error: Option<String>,
277}
278
279#[derive(Debug, Clone, Default, Serialize)]
281#[serde(rename_all = "camelCase")]
282pub struct SessionEndOutput {
283 #[serde(skip_serializing_if = "Option::is_none")]
285 pub suppress_output: Option<bool>,
286 #[serde(skip_serializing_if = "Option::is_none")]
288 pub cleanup_actions: Option<Vec<String>>,
289 #[serde(skip_serializing_if = "Option::is_none")]
291 pub session_summary: Option<String>,
292}
293
294#[derive(Debug, Clone, Deserialize)]
296#[serde(rename_all = "camelCase")]
297pub struct ErrorOccurredInput {
298 pub session_id: String,
300 pub timestamp: f64,
302 #[serde(rename = "cwd")]
304 pub working_directory: PathBuf,
305 pub error: String,
307 pub error_context: String,
309 pub recoverable: bool,
311}
312
313#[derive(Debug, Clone, Default, Serialize)]
315#[serde(rename_all = "camelCase")]
316pub struct ErrorOccurredOutput {
317 #[serde(skip_serializing_if = "Option::is_none")]
319 pub suppress_output: Option<bool>,
320 #[serde(skip_serializing_if = "Option::is_none")]
322 pub error_handling: Option<String>,
323 #[serde(skip_serializing_if = "Option::is_none")]
325 pub retry_count: Option<u32>,
326 #[serde(skip_serializing_if = "Option::is_none")]
328 pub user_notification: Option<String>,
329}
330
331#[derive(Debug, Clone, Deserialize)]
333#[serde(rename_all = "camelCase")]
334pub struct AgentStopInput {
335 pub session_id: String,
337 pub timestamp: f64,
339 #[serde(rename = "cwd")]
341 pub working_directory: PathBuf,
342 #[serde(default)]
344 pub stop_reason: Option<String>,
345 #[serde(default)]
347 pub transcript_path: Option<PathBuf>,
348 #[serde(default, rename = "stop_hook_active")]
350 pub stop_hook_active: Option<bool>,
351}
352
353#[derive(Debug, Clone, Default, Serialize)]
355#[serde(rename_all = "camelCase")]
356pub struct AgentStopOutput {
357 #[serde(skip_serializing_if = "Option::is_none")]
359 pub decision: Option<String>,
360 #[serde(skip_serializing_if = "Option::is_none")]
362 pub reason: Option<String>,
363}
364
365#[non_exhaustive]
371#[derive(Debug)]
372pub enum HookEvent {
373 PreToolUse {
375 input: PreToolUseInput,
377 ctx: HookContext,
379 },
380 PreMcpToolCall {
382 input: PreMcpToolCallInput,
384 ctx: HookContext,
386 },
387 PostToolUse {
389 input: PostToolUseInput,
391 ctx: HookContext,
393 },
394 PostToolUseFailure {
398 input: PostToolUseFailureInput,
400 ctx: HookContext,
402 },
403 UserPromptSubmitted {
405 input: UserPromptSubmittedInput,
407 ctx: HookContext,
409 },
410 UserPromptTransformed {
412 input: UserPromptTransformedInput,
414 ctx: HookContext,
416 },
417 SessionStart {
419 input: SessionStartInput,
421 ctx: HookContext,
423 },
424 SessionEnd {
426 input: SessionEndInput,
428 ctx: HookContext,
430 },
431 ErrorOccurred {
433 input: ErrorOccurredInput,
435 ctx: HookContext,
437 },
438 AgentStop {
440 input: AgentStopInput,
442 ctx: HookContext,
444 },
445}
446
447#[non_exhaustive]
452#[derive(Debug)]
453pub enum HookOutput {
454 None,
456 PreToolUse(PreToolUseOutput),
458 PreMcpToolCall(PreMcpToolCallOutput),
460 PostToolUse(PostToolUseOutput),
462 PostToolUseFailure(PostToolUseFailureOutput),
464 UserPromptSubmitted(UserPromptSubmittedOutput),
466 UserPromptTransformed(UserPromptTransformedOutput),
468 SessionStart(SessionStartOutput),
470 SessionEnd(SessionEndOutput),
472 ErrorOccurred(ErrorOccurredOutput),
474 AgentStop(AgentStopOutput),
476}
477
478impl HookOutput {
479 fn variant_name(&self) -> &'static str {
480 match self {
481 Self::None => "None",
482 Self::PreToolUse(_) => "PreToolUse",
483 Self::PreMcpToolCall(_) => "PreMcpToolCall",
484 Self::PostToolUse(_) => "PostToolUse",
485 Self::PostToolUseFailure(_) => "PostToolUseFailure",
486 Self::UserPromptSubmitted(_) => "UserPromptSubmitted",
487 Self::UserPromptTransformed(_) => "UserPromptTransformed",
488 Self::SessionStart(_) => "SessionStart",
489 Self::SessionEnd(_) => "SessionEnd",
490 Self::ErrorOccurred(_) => "ErrorOccurred",
491 Self::AgentStop(_) => "AgentStop",
492 }
493 }
494}
495
496#[async_trait]
514pub trait SessionHooks: Send + Sync + 'static {
515 async fn on_hook(&self, event: HookEvent) -> HookOutput {
519 match event {
520 HookEvent::PreToolUse { input, ctx } => self
521 .on_pre_tool_use(input, ctx)
522 .await
523 .map(HookOutput::PreToolUse)
524 .unwrap_or(HookOutput::None),
525 HookEvent::PreMcpToolCall { input, ctx } => self
526 .on_pre_mcp_tool_call(input, ctx)
527 .await
528 .map(HookOutput::PreMcpToolCall)
529 .unwrap_or(HookOutput::None),
530 HookEvent::PostToolUse { input, ctx } => self
531 .on_post_tool_use(input, ctx)
532 .await
533 .map(HookOutput::PostToolUse)
534 .unwrap_or(HookOutput::None),
535 HookEvent::PostToolUseFailure { input, ctx } => self
536 .on_post_tool_use_failure(input, ctx)
537 .await
538 .map(HookOutput::PostToolUseFailure)
539 .unwrap_or(HookOutput::None),
540 HookEvent::UserPromptSubmitted { input, ctx } => self
541 .on_user_prompt_submitted(input, ctx)
542 .await
543 .map(HookOutput::UserPromptSubmitted)
544 .unwrap_or(HookOutput::None),
545 HookEvent::UserPromptTransformed { input, ctx } => self
546 .on_user_prompt_transformed(input, ctx)
547 .await
548 .map(HookOutput::UserPromptTransformed)
549 .unwrap_or(HookOutput::None),
550 HookEvent::SessionStart { input, ctx } => self
551 .on_session_start(input, ctx)
552 .await
553 .map(HookOutput::SessionStart)
554 .unwrap_or(HookOutput::None),
555 HookEvent::SessionEnd { input, ctx } => self
556 .on_session_end(input, ctx)
557 .await
558 .map(HookOutput::SessionEnd)
559 .unwrap_or(HookOutput::None),
560 HookEvent::ErrorOccurred { input, ctx } => self
561 .on_error_occurred(input, ctx)
562 .await
563 .map(HookOutput::ErrorOccurred)
564 .unwrap_or(HookOutput::None),
565 HookEvent::AgentStop { input, ctx } => self
566 .on_agent_stop(input, ctx)
567 .await
568 .map(HookOutput::AgentStop)
569 .unwrap_or(HookOutput::None),
570 }
571 }
572
573 async fn on_pre_tool_use(
576 &self,
577 _input: PreToolUseInput,
578 _ctx: HookContext,
579 ) -> Option<PreToolUseOutput> {
580 None
581 }
582
583 async fn on_pre_mcp_tool_call(
586 &self,
587 _input: PreMcpToolCallInput,
588 _ctx: HookContext,
589 ) -> Option<PreMcpToolCallOutput> {
590 None
591 }
592
593 async fn on_post_tool_use(
597 &self,
598 _input: PostToolUseInput,
599 _ctx: HookContext,
600 ) -> Option<PostToolUseOutput> {
601 None
602 }
603
604 async fn on_post_tool_use_failure(
609 &self,
610 _input: PostToolUseFailureInput,
611 _ctx: HookContext,
612 ) -> Option<PostToolUseFailureOutput> {
613 None
614 }
615
616 async fn on_user_prompt_submitted(
620 &self,
621 _input: UserPromptSubmittedInput,
622 _ctx: HookContext,
623 ) -> Option<UserPromptSubmittedOutput> {
624 None
625 }
626
627 async fn on_user_prompt_transformed(
630 &self,
631 _input: UserPromptTransformedInput,
632 _ctx: HookContext,
633 ) -> Option<UserPromptTransformedOutput> {
634 None
635 }
636
637 async fn on_session_start(
640 &self,
641 _input: SessionStartInput,
642 _ctx: HookContext,
643 ) -> Option<SessionStartOutput> {
644 None
645 }
646
647 async fn on_session_end(
650 &self,
651 _input: SessionEndInput,
652 _ctx: HookContext,
653 ) -> Option<SessionEndOutput> {
654 None
655 }
656
657 async fn on_error_occurred(
660 &self,
661 _input: ErrorOccurredInput,
662 _ctx: HookContext,
663 ) -> Option<ErrorOccurredOutput> {
664 None
665 }
666
667 async fn on_agent_stop(
670 &self,
671 _input: AgentStopInput,
672 _ctx: HookContext,
673 ) -> Option<AgentStopOutput> {
674 None
675 }
676}
677
678pub(crate) async fn dispatch_hook(
684 hooks: &dyn SessionHooks,
685 session_id: &SessionId,
686 hook_type: &str,
687 raw_input: Value,
688) -> Result<Value, crate::Error> {
689 let ctx = HookContext {
690 session_id: session_id.clone(),
691 };
692
693 let event = match hook_type {
694 "preToolUse" => {
695 let input: PreToolUseInput = serde_json::from_value(raw_input)?;
696 HookEvent::PreToolUse { input, ctx }
697 }
698 "preMcpToolCall" => {
699 let input: PreMcpToolCallInput = serde_json::from_value(raw_input)?;
700 HookEvent::PreMcpToolCall { input, ctx }
701 }
702 "postToolUse" => {
703 let input: PostToolUseInput = serde_json::from_value(raw_input)?;
704 HookEvent::PostToolUse { input, ctx }
705 }
706 "postToolUseFailure" => {
707 let input: PostToolUseFailureInput = serde_json::from_value(raw_input)?;
708 HookEvent::PostToolUseFailure { input, ctx }
709 }
710 "userPromptSubmitted" => {
711 let input: UserPromptSubmittedInput = serde_json::from_value(raw_input)?;
712 HookEvent::UserPromptSubmitted { input, ctx }
713 }
714 "userPromptTransformed" => {
715 let input: UserPromptTransformedInput = serde_json::from_value(raw_input)?;
716 HookEvent::UserPromptTransformed { input, ctx }
717 }
718 "sessionStart" => {
719 let input: SessionStartInput = serde_json::from_value(raw_input)?;
720 HookEvent::SessionStart { input, ctx }
721 }
722 "sessionEnd" => {
723 let input: SessionEndInput = serde_json::from_value(raw_input)?;
724 HookEvent::SessionEnd { input, ctx }
725 }
726 "errorOccurred" => {
727 let input: ErrorOccurredInput = serde_json::from_value(raw_input)?;
728 HookEvent::ErrorOccurred { input, ctx }
729 }
730 "agentStop" => {
731 let input: AgentStopInput = serde_json::from_value(raw_input)?;
732 HookEvent::AgentStop { input, ctx }
733 }
734 _ => {
735 tracing::warn!(
736 hook_type = hook_type,
737 session_id = %session_id,
738 "unknown hook type"
739 );
740 return Ok(serde_json::json!({ "output": {} }));
741 }
742 };
743
744 let dispatch_start = Instant::now();
745 let output = hooks.on_hook(event).await;
746 tracing::debug!(
747 elapsed_ms = dispatch_start.elapsed().as_millis(),
748 session_id = %session_id,
749 hook_type = hook_type,
750 "SessionHooks::on_hook dispatch"
751 );
752
753 let output_value = match (hook_type, &output) {
758 (_, HookOutput::None) => None,
759 ("preToolUse", HookOutput::PreToolUse(o)) => Some(serde_json::to_value(o)?),
760 ("preMcpToolCall", HookOutput::PreMcpToolCall(o)) => Some(serde_json::to_value(o)?),
761 ("postToolUse", HookOutput::PostToolUse(o)) => Some(serde_json::to_value(o)?),
762 ("postToolUseFailure", HookOutput::PostToolUseFailure(o)) => Some(serde_json::to_value(o)?),
763 ("userPromptSubmitted", HookOutput::UserPromptSubmitted(o)) => {
764 Some(serde_json::to_value(o)?)
765 }
766 ("userPromptTransformed", HookOutput::UserPromptTransformed(o)) => {
767 Some(serde_json::to_value(o)?)
768 }
769 ("sessionStart", HookOutput::SessionStart(o)) => Some(serde_json::to_value(o)?),
770 ("sessionEnd", HookOutput::SessionEnd(o)) => Some(serde_json::to_value(o)?),
771 ("errorOccurred", HookOutput::ErrorOccurred(o)) => Some(serde_json::to_value(o)?),
772 ("agentStop", HookOutput::AgentStop(o)) => Some(serde_json::to_value(o)?),
773 _ => {
774 tracing::warn!(
775 hook_type = hook_type,
776 session_id = %session_id,
777 output_variant = output.variant_name(),
778 "hook returned mismatched output variant, treating as unregistered"
779 );
780 None
781 }
782 };
783
784 Ok(serde_json::json!({ "output": output_value.unwrap_or(Value::Object(Default::default())) }))
785}
786
787#[cfg(test)]
788mod tests {
789 use super::*;
790
791 struct TestHooks;
792
793 #[async_trait]
794 impl SessionHooks for TestHooks {
795 async fn on_hook(&self, event: HookEvent) -> HookOutput {
796 match event {
797 HookEvent::PreToolUse { input, .. } => {
798 if input.tool_name == "dangerous_tool" {
799 HookOutput::PreToolUse(PreToolUseOutput {
800 permission_decision: Some("deny".to_string()),
801 permission_decision_reason: Some("blocked by policy".to_string()),
802 ..Default::default()
803 })
804 } else {
805 HookOutput::None
806 }
807 }
808 HookEvent::UserPromptSubmitted { input, .. } => {
809 HookOutput::UserPromptSubmitted(UserPromptSubmittedOutput {
810 modified_prompt: Some(format!("[prefixed] {}", input.prompt)),
811 ..Default::default()
812 })
813 }
814 HookEvent::UserPromptTransformed { input, .. } => {
815 HookOutput::UserPromptTransformed(UserPromptTransformedOutput {
816 modified_transformed_prompt: Some(format!(
817 "[transformed] {}",
818 input.transformed_prompt
819 )),
820 })
821 }
822 _ => HookOutput::None,
823 }
824 }
825 }
826
827 #[tokio::test]
828 async fn dispatch_pre_tool_use_deny() {
829 let hooks = TestHooks;
830 let input = serde_json::json!({
831 "sessionId": "sess-1",
832 "timestamp": 1234567890,
833 "cwd": "/tmp",
834 "toolName": "dangerous_tool",
835 "toolArgs": {}
836 });
837 let result = dispatch_hook(&hooks, &SessionId::new("sess-1"), "preToolUse", input)
838 .await
839 .unwrap();
840 let output = &result["output"];
841 assert_eq!(output["permissionDecision"], "deny");
842 assert_eq!(output["permissionDecisionReason"], "blocked by policy");
843 }
844
845 #[tokio::test]
846 async fn dispatch_pre_tool_use_passthrough() {
847 let hooks = TestHooks;
848 let input = serde_json::json!({
849 "sessionId": "sess-1",
850 "timestamp": 1234567890,
851 "cwd": "/tmp",
852 "toolName": "safe_tool",
853 "toolArgs": {"key": "value"}
854 });
855 let result = dispatch_hook(&hooks, &SessionId::new("sess-1"), "preToolUse", input)
856 .await
857 .unwrap();
858 assert_eq!(result["output"], serde_json::json!({}));
860 }
861
862 #[tokio::test]
863 async fn dispatch_user_prompt_submitted() {
864 let hooks = TestHooks;
865 let input = serde_json::json!({
866 "sessionId": "sess-1",
867 "timestamp": 1234567890,
868 "cwd": "/tmp",
869 "prompt": "hello world"
870 });
871 let result = dispatch_hook(
872 &hooks,
873 &SessionId::new("sess-1"),
874 "userPromptSubmitted",
875 input,
876 )
877 .await
878 .unwrap();
879 assert_eq!(result["output"]["modifiedPrompt"], "[prefixed] hello world");
880 }
881
882 #[tokio::test]
883 async fn dispatch_user_prompt_transformed() {
884 let hooks = TestHooks;
885 let input = serde_json::json!({
886 "sessionId": "sess-1",
887 "timestamp": 1234567890,
888 "cwd": "/tmp",
889 "prompt": "hello world",
890 "transformedPrompt": "<current_datetime>now</current_datetime>\nhello world"
891 });
892 let result = dispatch_hook(
893 &hooks,
894 &SessionId::new("sess-1"),
895 "userPromptTransformed",
896 input,
897 )
898 .await
899 .unwrap();
900 assert_eq!(
901 result["output"]["modifiedTransformedPrompt"],
902 "[transformed] <current_datetime>now</current_datetime>\nhello world"
903 );
904 }
905
906 #[tokio::test]
907 async fn dispatch_unregistered_hook_returns_empty() {
908 let hooks = TestHooks;
909 let input = serde_json::json!({
910 "sessionId": "sess-1",
911 "timestamp": 1234567890,
912 "cwd": "/tmp",
913 "reason": "complete"
914 });
915 let result = dispatch_hook(&hooks, &SessionId::new("sess-1"), "sessionEnd", input)
917 .await
918 .unwrap();
919 assert_eq!(result["output"], serde_json::json!({}));
920 }
921
922 #[tokio::test]
923 async fn dispatch_unknown_hook_type() {
924 let hooks = TestHooks;
925 let input = serde_json::json!({});
926 let result = dispatch_hook(&hooks, &SessionId::new("sess-1"), "unknownHook", input)
927 .await
928 .unwrap();
929 assert_eq!(result["output"], serde_json::json!({}));
930 }
931
932 #[tokio::test]
933 async fn dispatch_mismatched_output_returns_empty() {
934 struct MismatchHooks;
935 #[async_trait]
936 impl SessionHooks for MismatchHooks {
937 async fn on_hook(&self, _event: HookEvent) -> HookOutput {
938 HookOutput::SessionEnd(SessionEndOutput {
940 session_summary: Some("oops".to_string()),
941 ..Default::default()
942 })
943 }
944 }
945
946 let hooks = MismatchHooks;
947 let input = serde_json::json!({
948 "sessionId": "sess-1",
949 "timestamp": 1234567890,
950 "cwd": "/tmp",
951 "toolName": "some_tool",
952 "toolArgs": {}
953 });
954 let result = dispatch_hook(&hooks, &SessionId::new("sess-1"), "preToolUse", input)
956 .await
957 .unwrap();
958 assert_eq!(result["output"], serde_json::json!({}));
959 }
960
961 #[tokio::test]
962 async fn dispatch_post_tool_use_default() {
963 let hooks = TestHooks;
964 let input = serde_json::json!({
965 "sessionId": "sess-1",
966 "timestamp": 1234567890,
967 "cwd": "/tmp",
968 "toolName": "some_tool",
969 "toolArgs": {},
970 "toolResult": "success"
971 });
972 let result = dispatch_hook(&hooks, &SessionId::new("sess-1"), "postToolUse", input)
973 .await
974 .unwrap();
975 assert_eq!(result["output"], serde_json::json!({}));
976 }
977
978 #[tokio::test]
979 async fn dispatch_post_tool_use_failure_default() {
980 let hooks = TestHooks;
982 let input = serde_json::json!({
983 "sessionId": "sess-1",
984 "timestamp": 1234567890,
985 "cwd": "/tmp",
986 "toolName": "some_tool",
987 "toolArgs": {"key": "value"},
988 "error": "boom"
989 });
990 let result = dispatch_hook(
991 &hooks,
992 &SessionId::new("sess-1"),
993 "postToolUseFailure",
994 input,
995 )
996 .await
997 .unwrap();
998 assert_eq!(result["output"], serde_json::json!({}));
999 }
1000
1001 #[tokio::test]
1002 async fn dispatch_post_tool_use_failure_returns_additional_context() {
1003 struct FailureHooks;
1004 #[async_trait]
1005 impl SessionHooks for FailureHooks {
1006 async fn on_post_tool_use_failure(
1007 &self,
1008 input: PostToolUseFailureInput,
1009 _ctx: HookContext,
1010 ) -> Option<PostToolUseFailureOutput> {
1011 assert_eq!(input.session_id, "sess-1");
1012 assert_eq!(input.tool_name, "some_tool");
1013 assert_eq!(input.error, "boom");
1014 assert_eq!(input.working_directory, PathBuf::from("/tmp"));
1015 Some(PostToolUseFailureOutput {
1016 additional_context: Some(format!(
1017 "tool {} failed: {}",
1018 input.tool_name, input.error
1019 )),
1020 })
1021 }
1022 }
1023
1024 let input = serde_json::json!({
1025 "sessionId": "sess-1",
1026 "timestamp": 1234567890,
1027 "cwd": "/tmp",
1028 "toolName": "some_tool",
1029 "toolArgs": {},
1030 "error": "boom"
1031 });
1032 let result = dispatch_hook(
1033 &FailureHooks,
1034 &SessionId::new("sess-1"),
1035 "postToolUseFailure",
1036 input,
1037 )
1038 .await
1039 .unwrap();
1040 assert_eq!(
1041 result["output"]["additionalContext"],
1042 "tool some_tool failed: boom"
1043 );
1044 }
1045
1046 #[tokio::test]
1047 async fn dispatch_post_tool_use_failure_invalid_input_errors() {
1048 let hooks = TestHooks;
1051 let input = serde_json::json!({
1052 "sessionId": "sess-1",
1053 "timestamp": 1234567890,
1054 "cwd": "/tmp",
1055 "toolName": "some_tool",
1056 "toolArgs": {}
1057 });
1058 let err = dispatch_hook(
1059 &hooks,
1060 &SessionId::new("sess-1"),
1061 "postToolUseFailure",
1062 input,
1063 )
1064 .await
1065 .unwrap_err();
1066 let msg = err.to_string().to_ascii_lowercase();
1067 assert!(
1068 msg.contains("error") || msg.contains("missing field"),
1069 "unexpected error: {msg}"
1070 );
1071 }
1072
1073 #[tokio::test]
1074 async fn dispatch_session_start() {
1075 struct StartHooks;
1076 #[async_trait]
1077 impl SessionHooks for StartHooks {
1078 async fn on_hook(&self, event: HookEvent) -> HookOutput {
1079 match event {
1080 HookEvent::SessionStart { .. } => {
1081 HookOutput::SessionStart(SessionStartOutput {
1082 additional_context: Some("extra context".to_string()),
1083 ..Default::default()
1084 })
1085 }
1086 _ => HookOutput::None,
1087 }
1088 }
1089 }
1090
1091 let hooks = StartHooks;
1092 let input = serde_json::json!({
1093 "sessionId": "sess-1",
1094 "timestamp": 1234567890,
1095 "cwd": "/tmp",
1096 "source": "new"
1097 });
1098 let result = dispatch_hook(&hooks, &SessionId::new("sess-1"), "sessionStart", input)
1099 .await
1100 .unwrap();
1101 assert_eq!(result["output"]["additionalContext"], "extra context");
1102 }
1103
1104 #[tokio::test]
1105 async fn dispatch_error_occurred() {
1106 struct ErrorHooks;
1107 #[async_trait]
1108 impl SessionHooks for ErrorHooks {
1109 async fn on_hook(&self, event: HookEvent) -> HookOutput {
1110 match event {
1111 HookEvent::ErrorOccurred { .. } => {
1112 HookOutput::ErrorOccurred(ErrorOccurredOutput {
1113 error_handling: Some("retry".to_string()),
1114 retry_count: Some(3),
1115 ..Default::default()
1116 })
1117 }
1118 _ => HookOutput::None,
1119 }
1120 }
1121 }
1122
1123 let hooks = ErrorHooks;
1124 let input = serde_json::json!({
1125 "sessionId": "sess-1",
1126 "timestamp": 1234567890,
1127 "cwd": "/tmp",
1128 "error": "model timeout",
1129 "errorContext": "model_call",
1130 "recoverable": true
1131 });
1132 let result = dispatch_hook(&hooks, &SessionId::new("sess-1"), "errorOccurred", input)
1133 .await
1134 .unwrap();
1135 assert_eq!(result["output"]["errorHandling"], "retry");
1136 assert_eq!(result["output"]["retryCount"], 3);
1137 }
1138
1139 #[tokio::test]
1140 async fn dispatch_agent_stop_block() {
1141 struct AgentStopHooks;
1142 #[async_trait]
1143 impl SessionHooks for AgentStopHooks {
1144 async fn on_agent_stop(
1145 &self,
1146 input: AgentStopInput,
1147 ctx: HookContext,
1148 ) -> Option<AgentStopOutput> {
1149 assert_eq!(ctx.session_id, SessionId::new("sess-1"));
1150 assert_eq!(input.session_id, "sess-1");
1151 assert_eq!(input.stop_reason.as_deref(), Some("end_turn"));
1152 assert_eq!(
1153 input.transcript_path,
1154 Some(PathBuf::from("/tmp/transcript.jsonl"))
1155 );
1156 assert_eq!(input.stop_hook_active, Some(true));
1157 Some(AgentStopOutput {
1158 decision: Some("block".to_string()),
1159 reason: Some("finish the remaining work".to_string()),
1160 })
1161 }
1162 }
1163
1164 let input = serde_json::json!({
1165 "sessionId": "sess-1",
1166 "timestamp": 1234567890,
1167 "cwd": "/tmp",
1168 "stopReason": "end_turn",
1169 "transcriptPath": "/tmp/transcript.jsonl",
1170 "stop_hook_active": true
1171 });
1172 let result = dispatch_hook(
1173 &AgentStopHooks,
1174 &SessionId::new("sess-1"),
1175 "agentStop",
1176 input,
1177 )
1178 .await
1179 .unwrap();
1180
1181 assert_eq!(result["output"]["decision"], "block");
1182 assert_eq!(result["output"]["reason"], "finish the remaining work");
1183 }
1184}