From d5fe71783997ac053762eba7a079d620ce513daa Mon Sep 17 00:00:00 2001 From: Zhangyi Yuan Date: Sat, 10 Oct 2026 15:58:05 +0800 Subject: [PATCH] fix(rust): decode structured lifecycle hook errors --- rust/src/hooks.rs | 47 ++++++++++++++- rust/src/hooks/tests.rs | 129 ++++++++++++++++++++++++++++++++++++++++ 2 files changed, 173 insertions(+), 3 deletions(-) diff --git a/rust/src/hooks.rs b/rust/src/hooks.rs index 3b2af37cb1..239e852eb7 100644 --- a/rust/src/hooks.rs +++ b/rust/src/hooks.rs @@ -271,8 +271,9 @@ pub struct SessionEndInput { /// The last assistant message. #[serde(default)] pub final_message: Option, - /// 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, } @@ -302,7 +303,9 @@ 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, @@ -310,6 +313,44 @@ pub struct ErrorOccurredInput { pub recoverable: bool, } +#[derive(Deserialize)] +#[serde(untagged)] +enum HookError { + Text(String), + Object(serde_json::Map), +} + +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 +where + Deserializer: serde::Deserializer<'de>, +{ + HookError::deserialize(deserializer).map(HookError::into_text) +} + +fn deserialize_optional_hook_error<'de, Deserializer>( + deserializer: Deserializer, +) -> Result, Deserializer::Error> +where + Deserializer: serde::Deserializer<'de>, +{ + Option::::deserialize(deserializer).map(|error| error.map(HookError::into_text)) +} + /// Output for the `errorOccurred` hook. #[derive(Debug, Clone, Default, Serialize)] #[serde(rename_all = "camelCase")] diff --git a/rust/src/hooks/tests.rs b/rust/src/hooks/tests.rs index e8a8179c41..6a26cfbaf8 100644 --- a/rust/src/hooks/tests.rs +++ b/rust/src/hooks/tests.rs @@ -147,6 +147,135 @@ async fn dispatch_unknown_hook_type() { assert_eq!(result["output"], serde_json::json!({})); } +struct ErrorHooks { + expected: Option, +} + +#[async_trait] +impl SessionHooks for ErrorHooks { + async fn on_error_occurred( + &self, + input: ErrorOccurredInput, + ctx: HookContext, + ) -> Option { + 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 { + 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;