use std::path::PathBuf;
use std::time::Instant;
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use crate::types::SessionId;
#[derive(Debug, Clone)]
pub struct HookContext {
pub session_id: SessionId,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct PreToolUseInput {
pub session_id: String,
pub timestamp: f64,
#[serde(rename = "cwd")]
pub working_directory: PathBuf,
pub tool_name: String,
pub tool_args: Value,
}
#[derive(Debug, Clone, Default, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct PreToolUseOutput {
#[serde(skip_serializing_if = "Option::is_none")]
pub permission_decision: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub permission_decision_reason: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub modified_args: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub additional_context: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub suppress_output: Option<bool>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct PreMcpToolCallInput {
pub session_id: String,
pub timestamp: f64,
#[serde(rename = "cwd")]
pub working_directory: PathBuf,
pub server_name: String,
pub tool_name: String,
pub arguments: Value,
#[serde(default)]
pub tool_call_id: Option<String>,
#[serde(default, rename = "_meta")]
pub meta: Option<Value>,
}
#[derive(Debug, Clone, Default, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct PreMcpToolCallOutput {
#[serde(skip_serializing_if = "Option::is_none")]
pub meta_to_use: Option<Value>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct PostToolUseInput {
pub session_id: String,
pub timestamp: f64,
#[serde(rename = "cwd")]
pub working_directory: PathBuf,
pub tool_name: String,
pub tool_args: Value,
pub tool_result: Value,
}
#[derive(Debug, Clone, Default, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct PostToolUseOutput {
#[serde(skip_serializing_if = "Option::is_none")]
pub modified_result: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub additional_context: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub suppress_output: Option<bool>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct PostToolUseFailureInput {
pub session_id: String,
pub timestamp: f64,
#[serde(rename = "cwd")]
pub working_directory: PathBuf,
pub tool_name: String,
pub tool_args: Value,
pub error: String,
}
#[derive(Debug, Clone, Default, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct PostToolUseFailureOutput {
#[serde(skip_serializing_if = "Option::is_none")]
pub additional_context: Option<String>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct UserPromptSubmittedInput {
pub session_id: String,
pub timestamp: f64,
#[serde(rename = "cwd")]
pub working_directory: PathBuf,
pub prompt: String,
}
#[derive(Debug, Clone, Default, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct UserPromptSubmittedOutput {
#[serde(skip_serializing_if = "Option::is_none")]
pub modified_prompt: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub additional_context: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub suppress_output: Option<bool>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct UserPromptTransformedInput {
pub session_id: String,
pub timestamp: f64,
#[serde(rename = "cwd")]
pub working_directory: PathBuf,
pub prompt: String,
pub transformed_prompt: String,
}
#[derive(Debug, Clone, Default, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct UserPromptTransformedOutput {
#[serde(skip_serializing_if = "Option::is_none")]
pub modified_transformed_prompt: Option<String>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct SessionStartInput {
pub session_id: String,
pub timestamp: f64,
#[serde(rename = "cwd")]
pub working_directory: PathBuf,
pub source: String,
#[serde(default)]
pub initial_prompt: Option<String>,
}
#[derive(Debug, Clone, Default, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct SessionStartOutput {
#[serde(skip_serializing_if = "Option::is_none")]
pub additional_context: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub modified_config: Option<Value>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct SessionEndInput {
pub session_id: String,
pub timestamp: f64,
#[serde(rename = "cwd")]
pub working_directory: PathBuf,
pub reason: String,
#[serde(default)]
pub final_message: Option<String>,
#[serde(default)]
pub error: Option<String>,
}
#[derive(Debug, Clone, Default, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct SessionEndOutput {
#[serde(skip_serializing_if = "Option::is_none")]
pub suppress_output: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub cleanup_actions: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub session_summary: Option<String>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct ErrorOccurredInput {
pub session_id: String,
pub timestamp: f64,
#[serde(rename = "cwd")]
pub working_directory: PathBuf,
pub error: String,
pub error_context: String,
pub recoverable: bool,
}
#[derive(Debug, Clone, Default, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct ErrorOccurredOutput {
#[serde(skip_serializing_if = "Option::is_none")]
pub suppress_output: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub error_handling: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub retry_count: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub user_notification: Option<String>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct AgentStopInput {
pub session_id: String,
pub timestamp: f64,
#[serde(rename = "cwd")]
pub working_directory: PathBuf,
#[serde(default)]
pub stop_reason: Option<String>,
#[serde(default)]
pub transcript_path: Option<PathBuf>,
#[serde(default, rename = "stop_hook_active")]
pub stop_hook_active: Option<bool>,
}
#[derive(Debug, Clone, Default, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct AgentStopOutput {
#[serde(skip_serializing_if = "Option::is_none")]
pub decision: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub reason: Option<String>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct SubagentStartInput {
pub session_id: String,
pub timestamp: f64,
#[serde(rename = "cwd")]
pub working_directory: PathBuf,
pub transcript_path: PathBuf,
pub agent_name: String,
#[serde(default)]
pub agent_display_name: Option<String>,
#[serde(default)]
pub agent_description: Option<String>,
}
#[derive(Debug, Clone, Default, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct SubagentStartOutput {
#[serde(skip_serializing_if = "Option::is_none")]
pub additional_context: Option<String>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct SubagentStopInput {
pub session_id: String,
pub timestamp: f64,
#[serde(rename = "cwd")]
pub working_directory: PathBuf,
pub transcript_path: PathBuf,
pub agent_name: String,
pub agent_type: String,
#[serde(default)]
pub agent_id: Option<String>,
#[serde(default)]
pub agent_display_name: Option<String>,
#[serde(default)]
pub agent_description: Option<String>,
pub stop_reason: String,
pub response: String,
}
#[derive(Debug, Clone, Default, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct SubagentStopOutput {
#[serde(skip_serializing_if = "Option::is_none")]
pub decision: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub reason: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub modified_response: Option<String>,
}
#[non_exhaustive]
#[derive(Debug)]
pub enum HookEvent {
PreToolUse {
input: PreToolUseInput,
ctx: HookContext,
},
PreMcpToolCall {
input: PreMcpToolCallInput,
ctx: HookContext,
},
PostToolUse {
input: PostToolUseInput,
ctx: HookContext,
},
PostToolUseFailure {
input: PostToolUseFailureInput,
ctx: HookContext,
},
UserPromptSubmitted {
input: UserPromptSubmittedInput,
ctx: HookContext,
},
UserPromptTransformed {
input: UserPromptTransformedInput,
ctx: HookContext,
},
SessionStart {
input: SessionStartInput,
ctx: HookContext,
},
SessionEnd {
input: SessionEndInput,
ctx: HookContext,
},
ErrorOccurred {
input: ErrorOccurredInput,
ctx: HookContext,
},
AgentStop {
input: AgentStopInput,
ctx: HookContext,
},
SubagentStart {
input: SubagentStartInput,
ctx: HookContext,
},
SubagentStop {
input: SubagentStopInput,
ctx: HookContext,
},
}
#[non_exhaustive]
#[derive(Debug)]
pub enum HookOutput {
None,
PreToolUse(PreToolUseOutput),
PreMcpToolCall(PreMcpToolCallOutput),
PostToolUse(PostToolUseOutput),
PostToolUseFailure(PostToolUseFailureOutput),
UserPromptSubmitted(UserPromptSubmittedOutput),
UserPromptTransformed(UserPromptTransformedOutput),
SessionStart(SessionStartOutput),
SessionEnd(SessionEndOutput),
ErrorOccurred(ErrorOccurredOutput),
AgentStop(AgentStopOutput),
SubagentStart(SubagentStartOutput),
SubagentStop(SubagentStopOutput),
}
impl HookOutput {
fn variant_name(&self) -> &'static str {
match self {
Self::None => "None",
Self::PreToolUse(_) => "PreToolUse",
Self::PreMcpToolCall(_) => "PreMcpToolCall",
Self::PostToolUse(_) => "PostToolUse",
Self::PostToolUseFailure(_) => "PostToolUseFailure",
Self::UserPromptSubmitted(_) => "UserPromptSubmitted",
Self::UserPromptTransformed(_) => "UserPromptTransformed",
Self::SessionStart(_) => "SessionStart",
Self::SessionEnd(_) => "SessionEnd",
Self::ErrorOccurred(_) => "ErrorOccurred",
Self::AgentStop(_) => "AgentStop",
Self::SubagentStart(_) => "SubagentStart",
Self::SubagentStop(_) => "SubagentStop",
}
}
}
#[async_trait]
pub trait SessionHooks: Send + Sync + 'static {
async fn on_hook(&self, event: HookEvent) -> HookOutput {
match event {
HookEvent::PreToolUse { input, ctx } => self
.on_pre_tool_use(input, ctx)
.await
.map(HookOutput::PreToolUse)
.unwrap_or(HookOutput::None),
HookEvent::PreMcpToolCall { input, ctx } => self
.on_pre_mcp_tool_call(input, ctx)
.await
.map(HookOutput::PreMcpToolCall)
.unwrap_or(HookOutput::None),
HookEvent::PostToolUse { input, ctx } => self
.on_post_tool_use(input, ctx)
.await
.map(HookOutput::PostToolUse)
.unwrap_or(HookOutput::None),
HookEvent::PostToolUseFailure { input, ctx } => self
.on_post_tool_use_failure(input, ctx)
.await
.map(HookOutput::PostToolUseFailure)
.unwrap_or(HookOutput::None),
HookEvent::UserPromptSubmitted { input, ctx } => self
.on_user_prompt_submitted(input, ctx)
.await
.map(HookOutput::UserPromptSubmitted)
.unwrap_or(HookOutput::None),
HookEvent::UserPromptTransformed { input, ctx } => self
.on_user_prompt_transformed(input, ctx)
.await
.map(HookOutput::UserPromptTransformed)
.unwrap_or(HookOutput::None),
HookEvent::SessionStart { input, ctx } => self
.on_session_start(input, ctx)
.await
.map(HookOutput::SessionStart)
.unwrap_or(HookOutput::None),
HookEvent::SessionEnd { input, ctx } => self
.on_session_end(input, ctx)
.await
.map(HookOutput::SessionEnd)
.unwrap_or(HookOutput::None),
HookEvent::ErrorOccurred { input, ctx } => self
.on_error_occurred(input, ctx)
.await
.map(HookOutput::ErrorOccurred)
.unwrap_or(HookOutput::None),
HookEvent::AgentStop { input, ctx } => self
.on_agent_stop(input, ctx)
.await
.map(HookOutput::AgentStop)
.unwrap_or(HookOutput::None),
HookEvent::SubagentStart { input, ctx } => self
.on_subagent_start(input, ctx)
.await
.map(HookOutput::SubagentStart)
.unwrap_or(HookOutput::None),
HookEvent::SubagentStop { input, ctx } => self
.on_subagent_stop(input, ctx)
.await
.map(HookOutput::SubagentStop)
.unwrap_or(HookOutput::None),
}
}
async fn on_pre_tool_use(
&self,
_input: PreToolUseInput,
_ctx: HookContext,
) -> Option<PreToolUseOutput> {
None
}
async fn on_pre_mcp_tool_call(
&self,
_input: PreMcpToolCallInput,
_ctx: HookContext,
) -> Option<PreMcpToolCallOutput> {
None
}
async fn on_post_tool_use(
&self,
_input: PostToolUseInput,
_ctx: HookContext,
) -> Option<PostToolUseOutput> {
None
}
async fn on_post_tool_use_failure(
&self,
_input: PostToolUseFailureInput,
_ctx: HookContext,
) -> Option<PostToolUseFailureOutput> {
None
}
async fn on_user_prompt_submitted(
&self,
_input: UserPromptSubmittedInput,
_ctx: HookContext,
) -> Option<UserPromptSubmittedOutput> {
None
}
async fn on_user_prompt_transformed(
&self,
_input: UserPromptTransformedInput,
_ctx: HookContext,
) -> Option<UserPromptTransformedOutput> {
None
}
async fn on_session_start(
&self,
_input: SessionStartInput,
_ctx: HookContext,
) -> Option<SessionStartOutput> {
None
}
async fn on_session_end(
&self,
_input: SessionEndInput,
_ctx: HookContext,
) -> Option<SessionEndOutput> {
None
}
async fn on_error_occurred(
&self,
_input: ErrorOccurredInput,
_ctx: HookContext,
) -> Option<ErrorOccurredOutput> {
None
}
async fn on_agent_stop(
&self,
_input: AgentStopInput,
_ctx: HookContext,
) -> Option<AgentStopOutput> {
None
}
async fn on_subagent_start(
&self,
_input: SubagentStartInput,
_ctx: HookContext,
) -> Option<SubagentStartOutput> {
None
}
async fn on_subagent_stop(
&self,
_input: SubagentStopInput,
_ctx: HookContext,
) -> Option<SubagentStopOutput> {
None
}
}
pub(crate) async fn dispatch_hook(
hooks: &dyn SessionHooks,
session_id: &SessionId,
hook_type: &str,
raw_input: Value,
) -> Result<Value, crate::Error> {
let ctx = HookContext {
session_id: session_id.clone(),
};
let event = match hook_type {
"preToolUse" => {
let input: PreToolUseInput = serde_json::from_value(raw_input)?;
HookEvent::PreToolUse { input, ctx }
}
"preMcpToolCall" => {
let input: PreMcpToolCallInput = serde_json::from_value(raw_input)?;
HookEvent::PreMcpToolCall { input, ctx }
}
"postToolUse" => {
let input: PostToolUseInput = serde_json::from_value(raw_input)?;
HookEvent::PostToolUse { input, ctx }
}
"postToolUseFailure" => {
let input: PostToolUseFailureInput = serde_json::from_value(raw_input)?;
HookEvent::PostToolUseFailure { input, ctx }
}
"userPromptSubmitted" => {
let input: UserPromptSubmittedInput = serde_json::from_value(raw_input)?;
HookEvent::UserPromptSubmitted { input, ctx }
}
"userPromptTransformed" => {
let input: UserPromptTransformedInput = serde_json::from_value(raw_input)?;
HookEvent::UserPromptTransformed { input, ctx }
}
"sessionStart" => {
let input: SessionStartInput = serde_json::from_value(raw_input)?;
HookEvent::SessionStart { input, ctx }
}
"sessionEnd" => {
let input: SessionEndInput = serde_json::from_value(raw_input)?;
HookEvent::SessionEnd { input, ctx }
}
"errorOccurred" => {
let input: ErrorOccurredInput = serde_json::from_value(raw_input)?;
HookEvent::ErrorOccurred { input, ctx }
}
"agentStop" => {
let input: AgentStopInput = serde_json::from_value(raw_input)?;
HookEvent::AgentStop { input, ctx }
}
"subagentStart" => {
let input: SubagentStartInput = serde_json::from_value(raw_input)?;
HookEvent::SubagentStart { input, ctx }
}
"subagentStop" => {
let input: SubagentStopInput = serde_json::from_value(raw_input)?;
HookEvent::SubagentStop { input, ctx }
}
_ => {
tracing::warn!(
hook_type = hook_type,
session_id = %session_id,
"unknown hook type"
);
return Ok(serde_json::json!({ "output": {} }));
}
};
let dispatch_start = Instant::now();
let output = hooks.on_hook(event).await;
tracing::debug!(
elapsed_ms = dispatch_start.elapsed().as_millis(),
session_id = %session_id,
hook_type = hook_type,
"SessionHooks::on_hook dispatch"
);
let output_value = match (hook_type, &output) {
(_, HookOutput::None) => None,
("preToolUse", HookOutput::PreToolUse(o)) => Some(serde_json::to_value(o)?),
("preMcpToolCall", HookOutput::PreMcpToolCall(o)) => Some(serde_json::to_value(o)?),
("postToolUse", HookOutput::PostToolUse(o)) => Some(serde_json::to_value(o)?),
("postToolUseFailure", HookOutput::PostToolUseFailure(o)) => Some(serde_json::to_value(o)?),
("userPromptSubmitted", HookOutput::UserPromptSubmitted(o)) => {
Some(serde_json::to_value(o)?)
}
("userPromptTransformed", HookOutput::UserPromptTransformed(o)) => {
Some(serde_json::to_value(o)?)
}
("sessionStart", HookOutput::SessionStart(o)) => Some(serde_json::to_value(o)?),
("sessionEnd", HookOutput::SessionEnd(o)) => Some(serde_json::to_value(o)?),
("errorOccurred", HookOutput::ErrorOccurred(o)) => Some(serde_json::to_value(o)?),
("agentStop", HookOutput::AgentStop(o)) => Some(serde_json::to_value(o)?),
("subagentStart", HookOutput::SubagentStart(o)) => Some(serde_json::to_value(o)?),
("subagentStop", HookOutput::SubagentStop(o)) => Some(serde_json::to_value(o)?),
_ => {
tracing::warn!(
hook_type = hook_type,
session_id = %session_id,
output_variant = output.variant_name(),
"hook returned mismatched output variant, treating as unregistered"
);
None
}
};
Ok(serde_json::json!({ "output": output_value.unwrap_or(Value::Object(Default::default())) }))
}
#[cfg(test)]
mod tests;