Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
47 changes: 44 additions & 3 deletions rust/src/hooks.rs
Original file line number Diff line number Diff line change
Expand Up @@ -271,8 +271,9 @@ pub struct SessionEndInput {
/// The last assistant message.
#[serde(default)]
pub final_message: Option<String>,
/// Error message, if the session ended due to an error.
#[serde(default)]
/// Error message, if the session ended due to an error. Structured errors
/// use their string `message`, or JSON text if no string message is present.
#[serde(default, deserialize_with = "deserialize_optional_hook_error")]
pub error: Option<String>,
}

Expand Down Expand Up @@ -302,14 +303,54 @@ pub struct ErrorOccurredInput {
/// Working directory.
#[serde(rename = "cwd")]
pub working_directory: PathBuf,
/// The error message.
/// The error message. Structured errors use their string `message`, or
/// JSON text if no string message is present.
#[serde(deserialize_with = "deserialize_hook_error")]
pub error: String,
/// Context where the error occurred: `"model_call"`, `"tool_execution"`, `"system"`, `"user_input"`.
pub error_context: String,
/// Whether the error is recoverable.
pub recoverable: bool,
}

#[derive(Deserialize)]
#[serde(untagged)]
enum HookError {
Text(String),
Object(serde_json::Map<String, Value>),
}

impl HookError {
fn into_text(self) -> String {
match self {
Self::Text(text) => text,
Self::Object(object) => object
.get("message")
.and_then(Value::as_str)
.map(str::to_owned)
.unwrap_or_else(|| Value::Object(object).to_string()),
}
}
}

fn deserialize_hook_error<'de, Deserializer>(
deserializer: Deserializer,
) -> Result<String, Deserializer::Error>
where
Deserializer: serde::Deserializer<'de>,
{
HookError::deserialize(deserializer).map(HookError::into_text)
}

fn deserialize_optional_hook_error<'de, Deserializer>(
deserializer: Deserializer,
) -> Result<Option<String>, Deserializer::Error>
where
Deserializer: serde::Deserializer<'de>,
{
Option::<HookError>::deserialize(deserializer).map(|error| error.map(HookError::into_text))
}

/// Output for the `errorOccurred` hook.
#[derive(Debug, Clone, Default, Serialize)]
#[serde(rename_all = "camelCase")]
Expand Down
129 changes: 129 additions & 0 deletions rust/src/hooks/tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -147,6 +147,135 @@ async fn dispatch_unknown_hook_type() {
assert_eq!(result["output"], serde_json::json!({}));
}

struct ErrorHooks {
expected: Option<String>,
}

#[async_trait]
impl SessionHooks for ErrorHooks {
async fn on_error_occurred(
&self,
input: ErrorOccurredInput,
ctx: HookContext,
) -> Option<ErrorOccurredOutput> {
assert_eq!(ctx.session_id, SessionId::new("sess-1"));
assert_eq!(Some(input.error), self.expected);
Some(ErrorOccurredOutput {
suppress_output: Some(true),
error_handling: Some("abort".to_string()),
..Default::default()
})
}

async fn on_session_end(
&self,
input: SessionEndInput,
ctx: HookContext,
) -> Option<SessionEndOutput> {
assert_eq!(ctx.session_id, SessionId::new("sess-1"));
assert_eq!(input.error, self.expected);
Some(SessionEndOutput {
suppress_output: Some(true),
..Default::default()
})
}
}

fn error_hook_input() -> Value {
serde_json::json!({
"sessionId": "sess-1",
"timestamp": 1234567890,
"cwd": "/tmp",
"reason": "error",
"errorContext": "model_call",
"recoverable": true
})
}

#[tokio::test]
async fn dispatch_error_hooks_accept_strings_and_objects() {
for (error, expected) in [
(serde_json::json!("legacy error"), "legacy error"),
(
serde_json::json!({"name": "Error", "message": "model timeout", "stack": "trace"}),
"model timeout",
),
(serde_json::json!({"message": ""}), ""),
(serde_json::json!({"name": "Error"}), r#"{"name":"Error"}"#),
(serde_json::json!({"message": 42}), r#"{"message":42}"#),
(serde_json::json!({}), "{}"),
] {
for hook_type in ["errorOccurred", "sessionEnd"] {
let mut input = error_hook_input();
input["error"] = error.clone();
let hooks = ErrorHooks {
expected: Some(expected.to_string()),
};
let result = dispatch_hook(&hooks, &SessionId::new("sess-1"), hook_type, input)
.await
.unwrap_or_else(|error| panic!("{hook_type}: {error}"));
assert_eq!(result["output"]["suppressOutput"], true);
if hook_type == "errorOccurred" {
assert_eq!(result["output"]["errorHandling"], "abort");
}
}
}
}

#[tokio::test]
async fn dispatch_session_end_accepts_missing_and_null_errors() {
for error in [None, Some(Value::Null)] {
let mut input = error_hook_input();
if let Some(error) = error {
input["error"] = error;
}
let result = dispatch_hook(
&ErrorHooks { expected: None },
&SessionId::new("sess-1"),
"sessionEnd",
input,
)
.await
.unwrap();
assert_eq!(result["output"]["suppressOutput"], true);
}
}

#[tokio::test]
async fn dispatch_error_hooks_reject_unsupported_values() {
for error in [
serde_json::json!(42),
serde_json::json!(true),
serde_json::json!([]),
] {
for hook_type in ["errorOccurred", "sessionEnd"] {
let mut input = error_hook_input();
input["error"] = error.clone();
assert!(
dispatch_hook(&TestHooks, &SessionId::new("sess-1"), hook_type, input)
.await
.is_err()
);
}
}
for error in [None, Some(Value::Null)] {
let mut input = error_hook_input();
if let Some(error) = error {
input["error"] = error;
}
assert!(
dispatch_hook(
&TestHooks,
&SessionId::new("sess-1"),
"errorOccurred",
input
)
.await
.is_err()
);
}
}

#[tokio::test]
async fn dispatch_subagent_hooks_with_typed_input_and_output() {
struct SubagentHooks;
Expand Down