use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use uuid::Uuid;
use crate::hook::HookEvent;
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
#[serde(rename_all = "snake_case", tag = "action")]
pub enum HookResult {
#[default]
Continue,
ContinueWith {
modifications: HashMap<String, serde_json::Value>,
},
Block {
reason: String,
},
Ask {
question: String,
#[serde(default)]
default: Option<String>,
},
}
impl HookResult {
#[must_use]
pub(crate) fn continue_() -> Self {
Self::Continue
}
#[must_use]
pub(crate) fn continue_with(modifications: HashMap<String, serde_json::Value>) -> Self {
Self::ContinueWith { modifications }
}
#[must_use]
pub(crate) fn block(reason: impl Into<String>) -> Self {
Self::Block {
reason: reason.into(),
}
}
#[must_use]
pub(crate) fn ask(question: impl Into<String>) -> Self {
Self::Ask {
question: question.into(),
default: None,
}
}
#[must_use]
pub(crate) fn is_blocking(&self) -> bool {
matches!(self, Self::Block { .. })
}
#[must_use]
pub(crate) fn requires_interaction(&self) -> bool {
matches!(self, Self::Ask { .. })
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub(crate) struct HookContext {
pub invocation_id: Uuid,
pub event: HookEvent,
#[serde(default)]
pub session_id: Option<Uuid>,
#[serde(default)]
pub user_id: Option<Uuid>,
pub timestamp: DateTime<Utc>,
#[serde(default)]
pub data: HashMap<String, serde_json::Value>,
#[serde(default)]
pub previous_results: Vec<HookResult>,
}
impl HookContext {
#[must_use]
pub(crate) fn new(event: HookEvent) -> Self {
Self {
invocation_id: Uuid::new_v4(),
event,
session_id: None,
user_id: None,
timestamp: Utc::now(),
data: HashMap::new(),
previous_results: Vec::new(),
}
}
#[must_use]
pub(crate) fn with_session(mut self, session_id: Uuid) -> Self {
self.session_id = Some(session_id);
self
}
#[must_use]
pub(crate) fn with_user(mut self, user_id: Uuid) -> Self {
self.user_id = Some(user_id);
self
}
#[must_use]
pub(crate) fn with_data(mut self, key: impl Into<String>, value: serde_json::Value) -> Self {
self.data.insert(key.into(), value);
self
}
pub(crate) fn add_previous_result(&mut self, result: HookResult) {
self.previous_results.push(result);
}
#[must_use]
pub(crate) fn get_data(&self, key: &str) -> Option<&serde_json::Value> {
self.data.get(key)
}
#[must_use]
pub(crate) fn get_data_as<T: for<'de> Deserialize<'de>>(&self, key: &str) -> Option<T> {
self.data
.get(key)
.and_then(|v| serde_json::from_value(v.clone()).ok())
}
#[must_use]
pub(crate) fn was_blocked(&self) -> bool {
self.previous_results.iter().any(HookResult::is_blocking)
}
#[must_use]
pub(crate) fn to_json(&self) -> serde_json::Value {
serde_json::to_value(self).unwrap_or(serde_json::Value::Null)
}
#[must_use]
pub(crate) fn to_env_vars(&self) -> HashMap<String, String> {
let mut env = HashMap::new();
env.insert("ASTRID_HOOK_ID".to_string(), self.invocation_id.to_string());
env.insert("ASTRID_HOOK_EVENT".to_string(), self.event.to_string());
env.insert(
"ASTRID_HOOK_TIMESTAMP".to_string(),
self.timestamp.to_rfc3339(),
);
if let Some(session_id) = &self.session_id {
env.insert("ASTRID_SESSION_ID".to_string(), session_id.to_string());
}
if let Some(user_id) = &self.user_id {
env.insert("ASTRID_USER_ID".to_string(), user_id.to_string());
}
if !self.data.is_empty()
&& let Ok(json) = serde_json::to_string(&self.data)
{
env.insert("ASTRID_HOOK_DATA".to_string(), json);
}
env
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub(crate) struct HookExecution {
pub hook_id: Uuid,
pub invocation_id: Uuid,
pub started_at: DateTime<Utc>,
pub completed_at: DateTime<Utc>,
pub duration_ms: u64,
pub result: HookExecutionResult,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "snake_case", tag = "status")]
pub(crate) enum HookExecutionResult {
Success {
result: HookResult,
#[serde(default)]
stdout: Option<String>,
},
Failure {
error: String,
#[serde(default)]
stderr: Option<String>,
},
Timeout {
timeout_secs: u64,
},
Skipped {
reason: String,
},
}
impl HookExecutionResult {
#[must_use]
pub(crate) fn is_success(&self) -> bool {
matches!(self, Self::Success { .. })
}
#[must_use]
pub(crate) fn hook_result(&self) -> Option<&HookResult> {
match self {
Self::Success { result, .. } => Some(result),
_ => None,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_hook_result_continue() {
let result = HookResult::continue_();
assert!(!result.is_blocking());
assert!(!result.requires_interaction());
}
#[test]
fn test_hook_result_block() {
let result = HookResult::block("Policy violation");
assert!(result.is_blocking());
}
#[test]
fn test_hook_result_ask() {
let result = HookResult::ask("Are you sure?");
assert!(result.requires_interaction());
}
#[test]
fn test_hook_context_creation() {
let session_id = Uuid::new_v4();
let user_id = Uuid::new_v4();
let ctx = HookContext::new(HookEvent::PreToolCall)
.with_session(session_id)
.with_user(user_id)
.with_data("tool_name", serde_json::json!("read_file"));
assert_eq!(ctx.event, HookEvent::PreToolCall);
assert_eq!(ctx.session_id, Some(session_id));
assert_eq!(ctx.user_id, Some(user_id));
assert!(ctx.get_data("tool_name").is_some());
}
#[test]
fn test_hook_context_env_vars() {
let ctx = HookContext::new(HookEvent::SessionStart);
let env = ctx.to_env_vars();
assert!(env.contains_key("ASTRID_HOOK_ID"));
assert_eq!(
env.get("ASTRID_HOOK_EVENT"),
Some(&"session_start".to_string())
);
}
#[test]
fn test_hook_execution_result() {
let success = HookExecutionResult::Success {
result: HookResult::Continue,
stdout: Some("ok".to_string()),
};
assert!(success.is_success());
assert!(success.hook_result().is_some());
let failure = HookExecutionResult::Failure {
error: "command failed".to_string(),
stderr: None,
};
assert!(!failure.is_success());
assert!(failure.hook_result().is_none());
}
}