From 68b1aa64e9e47f396373680d29ed899fb0a1c88a Mon Sep 17 00:00:00 2001 From: Hare Date: Thu, 3 Sep 2026 14:46:07 +0900 Subject: [PATCH] feat: type async interceptor failures --- crates/agen/examples/engine_cli.rs | 6 +- crates/agen/src/engine.rs | 77 ++++++- crates/agen/src/interceptor.rs | 149 ++++++++++++-- crates/agen/src/lib.rs | 4 +- crates/agen/tests/engine_state_test.rs | 202 ++++++++++++++++++- crates/agen/tests/parallel_execution_test.rs | 43 ++-- crates/session-store/tests/session_test.rs | 6 +- crates/worker/src/compact/worker.rs | 41 ++-- crates/worker/src/ipc/interceptor.rs | 96 ++++----- crates/worker/src/worker.rs | 4 + 10 files changed, 505 insertions(+), 123 deletions(-) diff --git a/crates/agen/examples/engine_cli.rs b/crates/agen/examples/engine_cli.rs index 0595cf95..2c50d354 100644 --- a/crates/agen/examples/engine_cli.rs +++ b/crates/agen/examples/engine_cli.rs @@ -40,7 +40,7 @@ use tracing_subscriber::EnvFilter; use agen::{ Engine, EngineRunExit, RunInterruptionReason, - interceptor::{Interceptor, PostToolAction, ToolResultInfo}, + interceptor::{Interceptor, InterceptorResult, PostToolAction, ToolResultInfo}, llm_client::{ LlmClient, capability::{CacheStrategy, ModelCapability, StructuredOutput, ToolCallingSupport}, @@ -280,7 +280,7 @@ impl ToolResultPrinterPolicy { #[async_trait] impl Interceptor for ToolResultPrinterPolicy { - async fn post_tool_call(&self, info: &mut ToolResultInfo) -> PostToolAction { + async fn post_tool_call(&self, info: &mut ToolResultInfo) -> InterceptorResult { let name = self .call_names .lock() @@ -294,7 +294,7 @@ impl Interceptor for ToolResultPrinterPolicy { println!(" Result ({}): ✅ {}", name, info.result.summary); } - PostToolAction::Continue + Ok(PostToolAction::Continue) } } diff --git a/crates/agen/src/engine.rs b/crates/agen/src/engine.rs index a3e79483..ae3279b2 100644 --- a/crates/agen/src/engine.rs +++ b/crates/agen/src/engine.rs @@ -15,8 +15,8 @@ use crate::{ }, handler::{ErrorKind, StatusKind, ToolUseBlockStart, UsageKind}, interceptor::{ - DefaultInterceptor, Interceptor, PostToolAction, PreRequestAction, PreToolAction, - PromptAction, ToolCallInfo, ToolResultInfo, TurnEndAction, + DefaultInterceptor, Interceptor, InterceptorFailure, InterceptorPoint, PostToolAction, + PreRequestAction, PreToolAction, PromptAction, ToolCallInfo, ToolResultInfo, TurnEndAction, }, llm_client::{ ClientError, ConfigWarning, LlmClient, Request, RequestConfig, ResponseStream, @@ -58,6 +58,9 @@ pub enum EngineError { /// A durable-history observer rejected an item before it entered history. #[error("History append failed: {0}")] HistoryAppend(String), + /// A trusted host interceptor callback failed. + #[error(transparent)] + Interceptor(#[from] InterceptorFailure), /// Tool terminalization lost its execution-attempt compare-and-set fence. #[error("Tool execution attempt fence failed: {0}")] ToolAttemptFence(String), @@ -153,6 +156,8 @@ pub enum EngineRunExit { /// A typed reason why an engine run could not finish normally. #[derive(Debug)] pub enum RunInterruptionReason { + /// A trusted host interceptor callback failed at a typed lifecycle point. + Interceptor(InterceptorFailure), LimitReached, ContextWindowExceeded, Cancelled, @@ -173,6 +178,9 @@ impl From> for EngineRunExit { } Err(EngineError::Cancelled) => Self::Interrupted(RunInterruptionReason::Cancelled), Err(EngineError::PauseRequested) => Self::Paused, + Err(EngineError::Interceptor(failure)) => { + Self::Interrupted(RunInterruptionReason::Interceptor(failure)) + } Err(error) => Self::Interrupted(RunInterruptionReason::Unexpected(error)), } } @@ -1092,7 +1100,9 @@ impl Engine { EngineError::Cancelled => "Cancelled".to_string(), _ => err.to_string(), }; - self.interceptor.on_abort(&reason).await; + if let Err(error) = self.interceptor.on_abort(&reason).await { + return Err(InterceptorFailure::new(InterceptorPoint::Abort, error).into()); + } Err(err) } } @@ -1175,7 +1185,17 @@ impl Engine { context, }; - match self.interceptor.pre_tool_call(&mut info).await { + let pre_tool_action = + self.interceptor + .pre_tool_call(&mut info) + .await + .map_err(|error| { + EngineError::from(InterceptorFailure::new( + InterceptorPoint::PreToolCall, + error, + )) + })?; + match pre_tool_action { PreToolAction::Continue => {} PreToolAction::Skip => { continue; @@ -1476,7 +1496,17 @@ impl Engine { context: context.clone(), }; - match self.interceptor.post_tool_call(&mut info).await { + let post_tool_action = + self.interceptor + .post_tool_call(&mut info) + .await + .map_err(|error| { + EngineError::from(InterceptorFailure::new( + InterceptorPoint::PostToolCall, + error, + )) + })?; + match post_tool_action { PostToolAction::Continue => {} PostToolAction::Abort(reason) => { abort_reason = Some(reason); @@ -1612,7 +1642,12 @@ impl Engine { .interceptor .pending_history_appends() .await - .map_err(EngineError::HistoryAppend)?; + .map_err(|error| { + EngineError::from(InterceptorFailure::new( + InterceptorPoint::PendingHistoryAppends, + error, + )) + })?; if !pending.is_empty() { self.append_history_items(history, pending, annotate)?; } @@ -1679,7 +1714,17 @@ impl Engine { } // Interceptor: pre_llm_request - match self.interceptor.pre_llm_request(&mut request_context).await { + let pre_request_action = self + .interceptor + .pre_llm_request(&mut request_context) + .await + .map_err(|error| { + EngineError::from(InterceptorFailure::new( + InterceptorPoint::PreLlmRequest, + error, + )) + })?; + match pre_request_action { PreRequestAction::Cancel(reason) => { info!(reason = %reason, "Aborted by interceptor"); for cb in &self.turn_end_cbs { @@ -1795,7 +1840,14 @@ impl Engine { if tool_calls.is_empty() { let turn_end_context = history.items_cloned(); - match self.interceptor.on_turn_end(&turn_end_context).await { + let turn_end_action = self + .interceptor + .on_turn_end(&turn_end_context) + .await + .map_err(|error| { + EngineError::from(InterceptorFailure::new(InterceptorPoint::TurnEnd, error)) + })?; + match turn_end_action { TurnEndAction::Finish => { return Ok(EngineResult::Finished); } @@ -2502,7 +2554,14 @@ impl Engine { // Supplying new user input abandons any paused/yielded logical run. self.active_run_turn_count = None; let mut user_item = Item::user_message(user_input); - let extras = match self.interceptor.on_prompt_submit(&mut user_item).await { + let prompt_action = match self.interceptor.on_prompt_submit(&mut user_item).await { + Ok(action) => action, + Err(error) => { + let error = InterceptorFailure::new(InterceptorPoint::PromptSubmit, error).into(); + return self.finalize_interruption(Err(error)).await; + } + }; + let extras = match prompt_action { PromptAction::Cancel(reason) => { return self .finalize_interruption(Err(EngineError::Aborted(reason))) diff --git a/crates/agen/src/interceptor.rs b/crates/agen/src/interceptor.rs index 5c5a307f..ecb1fd8f 100644 --- a/crates/agen/src/interceptor.rs +++ b/crates/agen/src/interceptor.rs @@ -11,6 +11,101 @@ use async_trait::async_trait; use crate::Item; use crate::tool::{Tool, ToolCall, ToolExecutionContext, ToolMeta, ToolResult}; +// ============================================================================= +// Failure Types +// ============================================================================= + +/// A typed failure returned by an [`Interceptor`] implementation. +/// +/// The Engine attaches the exact [`InterceptorPoint`] at which the failure was +/// observed before exposing it through the run termination boundary. +#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)] +#[error("{message}")] +pub struct InterceptorError { + message: String, +} + +impl InterceptorError { + /// Create an interceptor failure with a caller-defined message. + pub fn new(message: impl Into) -> Self { + Self { + message: message.into(), + } + } + + /// Return the failure message supplied by the interceptor. + pub fn message(&self) -> &str { + &self.message + } +} + +impl From for InterceptorError { + fn from(message: String) -> Self { + Self::new(message) + } +} + +impl From<&str> for InterceptorError { + fn from(message: &str) -> Self { + Self::new(message) + } +} + +/// The Engine lifecycle point at which an interceptor failed. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum InterceptorPoint { + PromptSubmit, + PendingHistoryAppends, + PreLlmRequest, + PreToolCall, + PostToolCall, + TurnEnd, + Abort, +} + +impl std::fmt::Display for InterceptorPoint { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + let name = match self { + Self::PromptSubmit => "prompt_submit", + Self::PendingHistoryAppends => "pending_history_appends", + Self::PreLlmRequest => "pre_llm_request", + Self::PreToolCall => "pre_tool_call", + Self::PostToolCall => "post_tool_call", + Self::TurnEnd => "turn_end", + Self::Abort => "abort", + }; + formatter.write_str(name) + } +} + +/// An interceptor failure bound to the exact Engine lifecycle point that ran it. +#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)] +#[error("{point} interceptor failed: {error}")] +pub struct InterceptorFailure { + point: InterceptorPoint, + #[source] + error: InterceptorError, +} + +impl InterceptorFailure { + pub(crate) fn new(point: InterceptorPoint, error: InterceptorError) -> Self { + Self { point, error } + } + + /// The lifecycle point that returned the failure. + pub fn point(&self) -> InterceptorPoint { + self.point + } + + /// The typed error returned by the interceptor. + pub fn error(&self) -> &InterceptorError { + &self.error + } +} + +/// Result returned by asynchronous interceptor lifecycle methods. +pub type InterceptorResult = Result; + // ============================================================================= // Action Enums // ============================================================================= @@ -130,14 +225,19 @@ pub struct ToolResultInfo { /// Intercepts the Engine execution loop at key decision points. /// -/// All methods have default implementations that let the Engine -/// proceed without intervention. Callers provide richer implementations for -/// approval flows, permission checks, etc. +/// Every lifecycle method is asynchronous and returns [`InterceptorResult`], +/// keeping implementation failure separate from the method's control-flow +/// action. The Engine reports a failure as a typed run interruption annotated +/// with the exact [`InterceptorPoint`] that failed. +/// +/// All methods have default implementations that let the Engine proceed +/// without intervention. Callers provide richer implementations for approval +/// flows, permission checks, and other trusted host adaptation. #[async_trait] pub trait Interceptor: Send + Sync { - /// Called after receiving user input, before adding to history. - async fn on_prompt_submit(&self, _item: &mut Item) -> PromptAction { - PromptAction::Continue + /// Called after receiving user input, before adding it to Engine history. + async fn on_prompt_submit(&self, _item: &mut Item) -> InterceptorResult { + Ok(PromptAction::Continue) } /// Items that should be **committed to `engine.history`** just @@ -158,7 +258,7 @@ pub trait Interceptor: Send + Sync { /// reproducible per-request transformations (pruning, content /// trimming, cache anchors) that depend only on the existing /// history. - async fn pending_history_appends(&self) -> Result, String> { + async fn pending_history_appends(&self) -> InterceptorResult> { Ok(Vec::new()) } @@ -170,27 +270,38 @@ pub trait Interceptor: Send + Sync { /// If an interceptor derives a human/model-visible nudge from the current /// request context, return [`PreRequestAction::ContinueWith`] so the Engine /// commits it to history before the request is sent. - async fn pre_llm_request(&self, _context: &mut Vec) -> PreRequestAction { - PreRequestAction::Continue + async fn pre_llm_request( + &self, + _context: &mut Vec, + ) -> InterceptorResult { + Ok(PreRequestAction::Continue) } /// Called before each tool is executed. - async fn pre_tool_call(&self, _info: &mut ToolCallInfo) -> PreToolAction { - PreToolAction::Continue + async fn pre_tool_call(&self, _info: &mut ToolCallInfo) -> InterceptorResult { + Ok(PreToolAction::Continue) } - /// Called after each tool completes. - async fn post_tool_call(&self, _info: &mut ToolResultInfo) -> PostToolAction { - PostToolAction::Continue + /// Called after each tool reaches one terminal result. + async fn post_tool_call( + &self, + _info: &mut ToolResultInfo, + ) -> InterceptorResult { + Ok(PostToolAction::Continue) } - /// Called when a turn ends with no tool calls. - async fn on_turn_end(&self, _history: &[Item]) -> TurnEndAction { - TurnEndAction::Finish + /// Called at the assistant boundary when a completed response has no tool calls. + /// + /// This is not the logical run termination observer. A host that needs that + /// boundary must inspect the returned [`crate::EngineRunExit`]. + async fn on_turn_end(&self, _history: &[Item]) -> InterceptorResult { + Ok(TurnEndAction::Finish) } - /// Called when execution is interrupted (abort or cancel). - async fn on_abort(&self, _reason: &str) {} + /// Called once when execution is interrupted (abort, cancellation, or failure). + async fn on_abort(&self, _reason: &str) -> InterceptorResult<()> { + Ok(()) + } } /// Default interceptor: no intervention. Engine proceeds through the loop diff --git a/crates/agen/src/lib.rs b/crates/agen/src/lib.rs index 96d4255c..52a7a246 100644 --- a/crates/agen/src/lib.rs +++ b/crates/agen/src/lib.rs @@ -26,7 +26,9 @@ pub use engine::{ }; pub use handler::ToolUseBlockStart; pub use history::{History, HistoryEntry}; -pub use interceptor::Interceptor; +pub use interceptor::{ + Interceptor, InterceptorError, InterceptorFailure, InterceptorPoint, InterceptorResult, +}; pub use message::{ContentPart, Item, Message, Role}; pub use tool::{ ToolCall, ToolExecutionContext, ToolExecutionHandle, ToolExecutionPolicy, diff --git a/crates/agen/tests/engine_state_test.rs b/crates/agen/tests/engine_state_test.rs index 527b8e94..c4d25428 100644 --- a/crates/agen/tests/engine_state_test.rs +++ b/crates/agen/tests/engine_state_test.rs @@ -10,7 +10,8 @@ use std::sync::{Arc, Mutex}; use agen::Item; use agen::interceptor::{ - Interceptor, PreRequestAction, PreToolAction, ToolCallInfo, TurnEndAction, + Interceptor, InterceptorError, InterceptorPoint, InterceptorResult, PostToolAction, + PreRequestAction, PreToolAction, PromptAction, ToolCallInfo, ToolResultInfo, TurnEndAction, }; use agen::llm_client::event::{Event, ResponseStatus, StatusEvent}; use agen::tool::{Tool, ToolDefinition, ToolError, ToolMeta, ToolOutput}; @@ -613,12 +614,15 @@ struct YieldOnce { #[async_trait] impl Interceptor for YieldOnce { - async fn pre_llm_request(&self, _context: &mut Vec) -> PreRequestAction { - if self.calls.fetch_add(1, Ordering::SeqCst) == 0 { + async fn pre_llm_request( + &self, + _context: &mut Vec, + ) -> InterceptorResult { + Ok(if self.calls.fetch_add(1, Ordering::SeqCst) == 0 { PreRequestAction::Yield } else { PreRequestAction::Continue - } + }) } } @@ -628,12 +632,12 @@ struct PauseToolOnce { #[async_trait] impl Interceptor for PauseToolOnce { - async fn pre_tool_call(&self, _info: &mut ToolCallInfo) -> PreToolAction { - if self.calls.fetch_add(1, Ordering::SeqCst) == 0 { + async fn pre_tool_call(&self, _info: &mut ToolCallInfo) -> InterceptorResult { + Ok(if self.calls.fetch_add(1, Ordering::SeqCst) == 0 { PreToolAction::Pause } else { PreToolAction::Continue - } + }) } } @@ -643,13 +647,193 @@ struct ContinueTurnOnce { #[async_trait] impl Interceptor for ContinueTurnOnce { - async fn on_turn_end(&self, _history: &[Item]) -> TurnEndAction { - if self.calls.fetch_add(1, Ordering::SeqCst) == 0 { + async fn on_turn_end(&self, _history: &[Item]) -> InterceptorResult { + Ok(if self.calls.fetch_add(1, Ordering::SeqCst) == 0 { TurnEndAction::ContinueWithMessages(vec![Item::system_message("continue")]) } else { TurnEndAction::Finish + }) + } +} + +#[derive(Debug, Clone)] +struct FailingLifecycleInterceptor { + failure: InterceptorPoint, + calls: Arc>>, +} + +impl FailingLifecycleInterceptor { + fn new(failure: InterceptorPoint) -> Self { + Self { + failure, + calls: Arc::new(Mutex::new(Vec::new())), } } + + fn record(&self, point: InterceptorPoint, action: T) -> InterceptorResult { + self.calls.lock().unwrap().push(point); + if self.failure == point { + Err(InterceptorError::new(format!("{point} rejected"))) + } else { + Ok(action) + } + } + + fn calls(&self) -> Vec { + self.calls.lock().unwrap().clone() + } +} + +#[async_trait] +impl Interceptor for FailingLifecycleInterceptor { + async fn on_prompt_submit(&self, _item: &mut Item) -> InterceptorResult { + tokio::task::yield_now().await; + self.record(InterceptorPoint::PromptSubmit, PromptAction::Continue) + } + + async fn pending_history_appends(&self) -> InterceptorResult> { + tokio::task::yield_now().await; + self.record(InterceptorPoint::PendingHistoryAppends, Vec::new()) + } + + async fn pre_llm_request( + &self, + _context: &mut Vec, + ) -> InterceptorResult { + tokio::task::yield_now().await; + if self.failure == InterceptorPoint::Abort { + self.calls + .lock() + .unwrap() + .push(InterceptorPoint::PreLlmRequest); + Ok(PreRequestAction::Cancel( + "trigger abort callback".to_string(), + )) + } else { + self.record(InterceptorPoint::PreLlmRequest, PreRequestAction::Continue) + } + } + + async fn pre_tool_call(&self, _info: &mut ToolCallInfo) -> InterceptorResult { + tokio::task::yield_now().await; + self.record(InterceptorPoint::PreToolCall, PreToolAction::Continue) + } + + async fn post_tool_call( + &self, + _info: &mut ToolResultInfo, + ) -> InterceptorResult { + tokio::task::yield_now().await; + self.record(InterceptorPoint::PostToolCall, PostToolAction::Continue) + } + + async fn on_turn_end(&self, _history: &[Item]) -> InterceptorResult { + tokio::task::yield_now().await; + self.record(InterceptorPoint::TurnEnd, TurnEndAction::Finish) + } + + async fn on_abort(&self, _reason: &str) -> InterceptorResult<()> { + tokio::task::yield_now().await; + self.record(InterceptorPoint::Abort, ()) + } +} + +fn expected_interceptor_calls(failure: InterceptorPoint) -> Vec { + use InterceptorPoint as Point; + + match failure { + Point::PromptSubmit => vec![Point::PromptSubmit, Point::Abort], + Point::PendingHistoryAppends => { + vec![ + Point::PromptSubmit, + Point::PendingHistoryAppends, + Point::Abort, + ] + } + Point::PreLlmRequest => vec![ + Point::PromptSubmit, + Point::PendingHistoryAppends, + Point::PreLlmRequest, + Point::Abort, + ], + Point::PreToolCall => vec![ + Point::PromptSubmit, + Point::PendingHistoryAppends, + Point::PreLlmRequest, + Point::PreToolCall, + Point::Abort, + ], + Point::PostToolCall => vec![ + Point::PromptSubmit, + Point::PendingHistoryAppends, + Point::PreLlmRequest, + Point::PreToolCall, + Point::PostToolCall, + Point::Abort, + ], + Point::TurnEnd => vec![ + Point::PromptSubmit, + Point::PendingHistoryAppends, + Point::PreLlmRequest, + Point::TurnEnd, + Point::Abort, + ], + Point::Abort => vec![ + Point::PromptSubmit, + Point::PendingHistoryAppends, + Point::PreLlmRequest, + Point::Abort, + ], + } +} + +#[tokio::test] +async fn interceptor_failures_are_typed_and_each_lifecycle_point_runs_once() { + use InterceptorPoint as Point; + + for failure_point in [ + Point::PromptSubmit, + Point::PendingHistoryAppends, + Point::PreLlmRequest, + Point::PreToolCall, + Point::PostToolCall, + Point::TurnEnd, + Point::Abort, + ] { + let interceptor = FailingLifecycleInterceptor::new(failure_point); + let needs_tool = matches!(failure_point, Point::PreToolCall | Point::PostToolCall); + let events = if needs_tool { + vec![ + Event::tool_use_start(0, "call-1", "count_tool"), + Event::tool_input_delta(0, "{}"), + Event::tool_use_stop(0), + Event::Status(StatusEvent { + status: ResponseStatus::Completed, + }), + ] + } else { + completed_text_events() + }; + let mut engine = Engine::new(MockLlmClient::new(events)); + engine.register_tool(CountingTool::new("count_tool").definition()); + engine.set_interceptor(interceptor.clone()); + let mut history = History::new(); + let mut engine = engine.lock(&history); + + let exit = engine.run(&mut history, "test").await; + let EngineRunExit::Interrupted(RunInterruptionReason::Interceptor(failure)) = exit else { + panic!("expected typed interceptor interruption at {failure_point}, got {exit:?}"); + }; + assert_eq!(failure.point(), failure_point); + assert_eq!( + failure.error().message(), + format!("{failure_point} rejected") + ); + assert_eq!( + interceptor.calls(), + expected_interceptor_calls(failure_point) + ); + } } #[tokio::test] diff --git a/crates/agen/tests/parallel_execution_test.rs b/crates/agen/tests/parallel_execution_test.rs index 88d2eda4..ead143b4 100644 --- a/crates/agen/tests/parallel_execution_test.rs +++ b/crates/agen/tests/parallel_execution_test.rs @@ -6,7 +6,9 @@ use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::{Arc, Mutex}; use std::time::{Duration, Instant}; -use agen::interceptor::{Interceptor, PostToolAction, PreToolAction, ToolCallInfo, ToolResultInfo}; +use agen::interceptor::{ + Interceptor, InterceptorResult, PostToolAction, PreToolAction, ToolCallInfo, ToolResultInfo, +}; use agen::llm_client::event::{Event, ResponseStatus, StatusEvent}; use agen::tool::{ Tool, ToolDefinition, ToolError, ToolExecutionContext, ToolMeta, ToolOutput, ToolResult, @@ -905,24 +907,27 @@ async fn test_tool_execution_context_for_skipped_and_synthetic_paths() { #[async_trait] impl Interceptor for ContextPolicy { - async fn pre_tool_call(&self, info: &mut ToolCallInfo) -> PreToolAction { + async fn pre_tool_call(&self, info: &mut ToolCallInfo) -> InterceptorResult { self.pre_contexts.lock().unwrap().push(info.context.clone()); - match info.call.name.as_str() { + Ok(match info.call.name.as_str() { "skip_tool" => PreToolAction::Skip, "synthetic_tool" => PreToolAction::SyntheticResult(ToolResult::from_output( &info.call.id, ToolOutput::from("synthetic result".to_string()), )), _ => PreToolAction::Continue, - } + }) } - async fn post_tool_call(&self, info: &mut ToolResultInfo) -> PostToolAction { + async fn post_tool_call( + &self, + info: &mut ToolResultInfo, + ) -> InterceptorResult { self.post_contexts .lock() .unwrap() .push(info.context.clone()); - PostToolAction::Continue + Ok(PostToolAction::Continue) } } @@ -994,12 +999,12 @@ async fn test_before_tool_call_skip() { #[async_trait] impl Interceptor for BlockingPolicy { - async fn pre_tool_call(&self, info: &mut ToolCallInfo) -> PreToolAction { - if info.call.name == "blocked_tool" { + async fn pre_tool_call(&self, info: &mut ToolCallInfo) -> InterceptorResult { + Ok(if info.call.name == "blocked_tool" { PreToolAction::Skip } else { PreToolAction::Continue - } + }) } } @@ -1081,10 +1086,13 @@ async fn test_post_tool_call_modification() { #[async_trait] impl Interceptor for ModifyingPolicy { - async fn post_tool_call(&self, info: &mut ToolResultInfo) -> PostToolAction { + async fn post_tool_call( + &self, + info: &mut ToolResultInfo, + ) -> InterceptorResult { info.result.summary = format!("[Modified] {}", info.result.summary); *self.modified_content.lock().unwrap() = Some(info.result.summary.clone()); - PostToolAction::Continue + Ok(PostToolAction::Continue) } } @@ -1143,11 +1151,11 @@ async fn test_before_tool_call_synthetic_result_committed() { #[async_trait] impl Interceptor for SyntheticPolicy { - async fn pre_tool_call(&self, info: &mut ToolCallInfo) -> PreToolAction { - PreToolAction::SyntheticResult(ToolResult::error( + async fn pre_tool_call(&self, info: &mut ToolCallInfo) -> InterceptorResult { + Ok(PreToolAction::SyntheticResult(ToolResult::error( info.call.id.clone(), "permission denied", - )) + ))) } } @@ -1184,8 +1192,11 @@ async fn post_tool_abort_commits_confirmed_result_before_stopping_run() { struct AbortAfterResult; #[async_trait] impl Interceptor for AbortAfterResult { - async fn post_tool_call(&self, _info: &mut ToolResultInfo) -> PostToolAction { - PostToolAction::Abort("policy stopped the run".to_string()) + async fn post_tool_call( + &self, + _info: &mut ToolResultInfo, + ) -> InterceptorResult { + Ok(PostToolAction::Abort("policy stopped the run".to_string())) } } engine.set_interceptor(AbortAfterResult); diff --git a/crates/session-store/tests/session_test.rs b/crates/session-store/tests/session_test.rs index 6f7ec3ee..2f0ffbeb 100644 --- a/crates/session-store/tests/session_test.rs +++ b/crates/session-store/tests/session_test.rs @@ -3,7 +3,7 @@ mod common; use std::ops::{Deref, DerefMut}; use std::sync::Arc; -use agen::interceptor::{Interceptor, TurnEndAction}; +use agen::interceptor::{Interceptor, InterceptorResult, TurnEndAction}; use agen::llm_client::event::{Event, ResponseStatus, StatusEvent}; use agen::llm_client::types::{Item, RequestConfig}; use agen::tool::{Tool, ToolDefinition, ToolError, ToolMeta, ToolOutput}; @@ -100,8 +100,8 @@ struct PausePolicy; #[async_trait] impl Interceptor for PausePolicy { - async fn on_turn_end(&self, _history: &[Item]) -> TurnEndAction { - TurnEndAction::Pause + async fn on_turn_end(&self, _history: &[Item]) -> InterceptorResult { + Ok(TurnEndAction::Pause) } } diff --git a/crates/worker/src/compact/worker.rs b/crates/worker/src/compact/worker.rs index 2f30407e..cd218ac0 100644 --- a/crates/worker/src/compact/worker.rs +++ b/crates/worker/src/compact/worker.rs @@ -22,7 +22,9 @@ use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; use std::sync::{Arc, Mutex}; use agen::Item; -use agen::interceptor::{Interceptor, PreRequestAction, PreToolAction, ToolCallInfo}; +use agen::interceptor::{ + Interceptor, InterceptorResult, PreRequestAction, PreToolAction, ToolCallInfo, +}; use agen::tool::{Tool, ToolDefinition, ToolError, ToolMeta, ToolOutput, ToolResult}; use async_trait::async_trait; use serde::Deserialize; @@ -398,14 +400,17 @@ impl CompactWorkerInterceptor { #[async_trait] impl Interceptor for CompactWorkerInterceptor { - async fn pre_llm_request(&self, context: &mut Vec) -> PreRequestAction { + async fn pre_llm_request( + &self, + context: &mut Vec, + ) -> InterceptorResult { let records = self.usage_tracker.records(); let estimate = agen::token_counter::total_tokens(context, &records); if estimate.tokens > self.max_input_tokens { - return PreRequestAction::Cancel(format!( + return Ok(PreRequestAction::Cancel(format!( "compact worker input occupancy exceeded {} tokens", self.max_input_tokens - )); + ))); } let remaining = self.max_input_tokens.saturating_sub(estimate.tokens); @@ -413,25 +418,25 @@ impl Interceptor for CompactWorkerInterceptor { .store(remaining, Ordering::Release); if let Some(item) = self.maybe_emit_warning(remaining) { self.usage_tracker.note_request(context.len() + 1); - return PreRequestAction::ContinueWith(vec![item]); + return Ok(PreRequestAction::ContinueWith(vec![item])); } self.usage_tracker.note_request(context.len()); - PreRequestAction::Continue + Ok(PreRequestAction::Continue) } - async fn pre_tool_call(&self, info: &mut ToolCallInfo) -> PreToolAction { + async fn pre_tool_call(&self, info: &mut ToolCallInfo) -> InterceptorResult { if self.final_reserve_tokens == 0 || info.call.name == "write_summary" { - return PreToolAction::Continue; + return Ok(PreToolAction::Continue); } let remaining = self.last_remaining_tokens.load(Ordering::Acquire); if remaining > self.final_reserve_tokens { - return PreToolAction::Continue; + return Ok(PreToolAction::Continue); } - PreToolAction::SyntheticResult(ToolResult::error( + Ok(PreToolAction::SyntheticResult(ToolResult::error( info.call.id.clone(), "compact worker final reserve reached; do not perform more exploratory tool reads. Call `write_summary` now.", - )) + ))) } } @@ -467,13 +472,13 @@ mod tests { let mut context = vec![Item::user_message("hello")]; assert!(matches!( - interceptor.pre_llm_request(&mut context).await, + interceptor.pre_llm_request(&mut context).await.unwrap(), PreRequestAction::Continue )); tracker.record_usage(&make_usage(100)); assert!(matches!( - interceptor.pre_llm_request(&mut context).await, + interceptor.pre_llm_request(&mut context).await.unwrap(), PreRequestAction::Continue )); tracker.record_usage(&make_usage(100)); @@ -481,7 +486,7 @@ mod tests { // Two 100-token requests would exceed a cumulative 150-token cap, but // current occupancy is still the latest 100-token measurement. assert!(matches!( - interceptor.pre_llm_request(&mut context).await, + interceptor.pre_llm_request(&mut context).await.unwrap(), PreRequestAction::Continue )); } @@ -503,13 +508,13 @@ mod tests { let mut context = vec![Item::user_message("hello")]; assert!(matches!( - interceptor.pre_llm_request(&mut context).await, + interceptor.pre_llm_request(&mut context).await.unwrap(), PreRequestAction::Continue )); tracker.record_usage(&make_usage(100)); assert!(matches!( - interceptor.pre_llm_request(&mut context).await, + interceptor.pre_llm_request(&mut context).await.unwrap(), PreRequestAction::ContinueWith(items) if items.len() == 1 && items[0].as_text().unwrap_or_default().contains("write_summary") )); @@ -523,13 +528,13 @@ mod tests { let mut context = vec![Item::user_message("hello")]; assert!(matches!( - interceptor.pre_llm_request(&mut context).await, + interceptor.pre_llm_request(&mut context).await.unwrap(), PreRequestAction::Continue )); tracker.record_usage(&make_usage(100)); assert!(matches!( - interceptor.pre_llm_request(&mut context).await, + interceptor.pre_llm_request(&mut context).await.unwrap(), PreRequestAction::Cancel(message) if message.contains("occupancy") )); } diff --git a/crates/worker/src/ipc/interceptor.rs b/crates/worker/src/ipc/interceptor.rs index f82444aa..a1468443 100644 --- a/crates/worker/src/ipc/interceptor.rs +++ b/crates/worker/src/ipc/interceptor.rs @@ -15,8 +15,8 @@ use std::sync::{Arc, Mutex}; use agen::Item; use agen::UsageRecord; use agen::interceptor::{ - Interceptor, PostToolAction, PreRequestAction, PreToolAction, PromptAction, ToolCallInfo, - ToolResultInfo, TurnEndAction, + Interceptor, InterceptorResult, PostToolAction, PreRequestAction, PreToolAction, PromptAction, + ToolCallInfo, ToolResultInfo, TurnEndAction, }; use agen::tool::ToolOutput; use arc_swap::ArcSwap; @@ -232,7 +232,7 @@ impl WorkerInterceptor { #[async_trait] impl Interceptor for WorkerInterceptor { - async fn on_prompt_submit(&self, item: &mut Item) -> PromptAction { + async fn on_prompt_submit(&self, item: &mut Item) -> InterceptorResult { let turn_index = self.next_turn_index.fetch_add(1, Ordering::Relaxed); self.tool_calls_this_turn.store(0, Ordering::Relaxed); @@ -243,7 +243,7 @@ impl Interceptor for WorkerInterceptor { for hook in &self.registry.on_prompt_submit { let action = hook.call(&info).await; if !matches!(action, HookPromptAction::Continue) { - return action.into(); + return Ok(action.into()); } } let mut extras: Vec = std::mem::take( @@ -252,7 +252,7 @@ impl Interceptor for WorkerInterceptor { .lock() .expect("pending_attachments poisoned"), ); - if extras.is_empty() { + Ok(if extras.is_empty() { PromptAction::Continue } else { // Commit the typed system items first, then hand the @@ -266,10 +266,10 @@ impl Interceptor for WorkerInterceptor { Ok(()) => PromptAction::ContinueWith(items), Err(error) => PromptAction::Cancel(format!("session persistence failed: {error}")), } - } + }) } - async fn pending_history_appends(&self) -> Result, String> { + async fn pending_history_appends(&self) -> InterceptorResult> { let drained = self.pending_notifies.drain(); if drained.is_empty() { return Ok(Vec::new()); @@ -295,7 +295,7 @@ impl Interceptor for WorkerInterceptor { Ok(system_item) => system_item, Err(error) => { self.pending_notifies.requeue_front(drained); - return Err(format!("failed to render notify_wrapper: {error}")); + return Err(format!("failed to render notify_wrapper: {error}").into()); } }; items.push(system_item.to_history_item()); @@ -303,15 +303,18 @@ impl Interceptor for WorkerInterceptor { } if let Err(error) = self.commit_system_items(&system_items) { self.pending_notifies.requeue_front(drained); - return Err(format!("session persistence failed: {error}")); + return Err(format!("session persistence failed: {error}").into()); } Ok(items) } - async fn pre_llm_request(&self, context: &mut Vec) -> PreRequestAction { + async fn pre_llm_request( + &self, + context: &mut Vec, + ) -> InterceptorResult { let initial_tokens = self.estimated_tokens(context); if self.request_threshold_exceeded(initial_tokens, context) { - return PreRequestAction::Yield; + return Ok(PreRequestAction::Yield); } let info = PreRequestInfo { item_count: context.len(), @@ -328,7 +331,7 @@ impl Interceptor for WorkerInterceptor { for hook in &self.registry.pre_llm_request { let action = hook.call(&hook_context).await; if !matches!(action, HookPreRequestAction::Continue) { - return action.into(); + return Ok(action.into()); } } @@ -353,28 +356,30 @@ impl Interceptor for WorkerInterceptor { if self.request_threshold_exceeded(current_tokens, effective_context.as_ref()) { if let Err(error) = self.commit_system_items(&system_items) { - return PreRequestAction::Cancel(format!("session persistence failed: {error}")); + return Ok(PreRequestAction::Cancel(format!( + "session persistence failed: {error}" + ))); } - return if appended_items.is_empty() { + return Ok(if appended_items.is_empty() { PreRequestAction::Yield } else { PreRequestAction::YieldWith(appended_items) - }; + }); } if let Some(usage_tracker) = self.usage_tracker.as_ref() { usage_tracker.note_request(effective_context.len()); } if system_items.is_empty() { - return PreRequestAction::Continue; + return Ok(PreRequestAction::Continue); } - match self.commit_system_items(&system_items) { + Ok(match self.commit_system_items(&system_items) { Ok(()) => PreRequestAction::ContinueWith(appended_items), Err(error) => PreRequestAction::Cancel(format!("session persistence failed: {error}")), - } + }) } - async fn pre_tool_call(&self, info: &mut ToolCallInfo) -> PreToolAction { + async fn pre_tool_call(&self, info: &mut ToolCallInfo) -> InterceptorResult { let summary = ToolCallSummary { call_id: info.call.id.clone(), tool_name: info.call.name.clone(), @@ -383,14 +388,14 @@ impl Interceptor for WorkerInterceptor { for hook in &self.registry.pre_tool_call { let action = hook.call(&summary).await; if !matches!(action, HookPreToolAction::Continue) { - return action.into_worker_action(summary.call_id.clone()); + return Ok(action.into_worker_action(summary.call_id.clone())); } } self.tool_calls_this_turn.fetch_add(1, Ordering::Relaxed); - PreToolAction::Continue + Ok(PreToolAction::Continue) } - async fn post_tool_call(&self, info: &mut ToolResultInfo) -> PostToolAction { + async fn post_tool_call(&self, info: &mut ToolResultInfo) -> InterceptorResult { let summary = ToolResultSummary { call_id: info.result.tool_use_id.clone(), tool_name: info.call.name.clone(), @@ -405,13 +410,13 @@ impl Interceptor for WorkerInterceptor { for hook in &self.registry.post_tool_call { let action = hook.call(&summary).await; if !matches!(action, HookPostToolAction::Continue) { - return action.into(); + return Ok(action.into()); } } - PostToolAction::Continue + Ok(PostToolAction::Continue) } - async fn on_turn_end(&self, history: &[Item]) -> TurnEndAction { + async fn on_turn_end(&self, history: &[Item]) -> InterceptorResult { let final_text_preview = history .iter() .rev() @@ -427,19 +432,20 @@ impl Interceptor for WorkerInterceptor { for hook in &self.registry.on_turn_end { let action = hook.call(&info).await; if !matches!(action, HookTurnEndAction::Finish) { - return action.into(); + return Ok(action.into()); } } - TurnEndAction::Finish + Ok(TurnEndAction::Finish) } - async fn on_abort(&self, reason: &str) { + async fn on_abort(&self, reason: &str) -> InterceptorResult<()> { let info = AbortInfo { reason: reason.to_string(), }; for hook in &self.registry.on_abort { hook.call(&info).await; } + Ok(()) } } @@ -623,7 +629,7 @@ mod tests { None, ); let mut ctx = ctx_items; - let action = interceptor.pre_llm_request(&mut ctx).await; + let action = interceptor.pre_llm_request(&mut ctx).await.unwrap(); assert!(matches!(action, PreRequestAction::Yield)); // Hook must not run when an internal mechanism short-circuits first. @@ -655,7 +661,7 @@ mod tests { })), ); let mut ctx = ctx_items; - let action = interceptor.pre_llm_request(&mut ctx).await; + let action = interceptor.pre_llm_request(&mut ctx).await.unwrap(); match action { PreRequestAction::YieldWith(items) => assert_eq!(items.len(), 1), @@ -692,7 +698,7 @@ mod tests { ) .with_usage_tracker(usage_tracker); let mut ctx = ctx_items; - let action = interceptor.pre_llm_request(&mut ctx).await; + let action = interceptor.pre_llm_request(&mut ctx).await.unwrap(); assert!(matches!(action, PreRequestAction::Yield)); } @@ -716,7 +722,7 @@ mod tests { None, ); let mut ctx = ctx_items; - let action = interceptor.pre_llm_request(&mut ctx).await; + let action = interceptor.pre_llm_request(&mut ctx).await.unwrap(); assert!(matches!(action, PreRequestAction::Continue)); assert_eq!(count.load(Ordering::Relaxed), 1); @@ -757,7 +763,7 @@ mod tests { None, ); let mut ctx = ctx_items; - let action = interceptor.pre_llm_request(&mut ctx).await; + let action = interceptor.pre_llm_request(&mut ctx).await.unwrap(); assert!(matches!(action, PreRequestAction::Continue)); assert_eq!(count.load(Ordering::Relaxed), 1); @@ -784,7 +790,7 @@ mod tests { None, ); let mut ctx = ctx_items; - let action = interceptor.pre_llm_request(&mut ctx).await; + let action = interceptor.pre_llm_request(&mut ctx).await.unwrap(); assert!(matches!(action, PreRequestAction::Continue)); assert_eq!(count.load(Ordering::Relaxed), 1); @@ -805,7 +811,7 @@ mod tests { None, ); let mut ctx: Vec = Vec::new(); - let action = interceptor.pre_llm_request(&mut ctx).await; + let action = interceptor.pre_llm_request(&mut ctx).await.unwrap(); assert!(matches!(action, PreRequestAction::Continue)); assert_eq!(count.load(Ordering::Relaxed), 1); @@ -834,7 +840,7 @@ mod tests { ); let mut ctx: Vec = Vec::new(); - let action = interceptor.pre_llm_request(&mut ctx).await; + let action = interceptor.pre_llm_request(&mut ctx).await.unwrap(); assert!(saw_handle.load(Ordering::Relaxed)); let PreRequestAction::ContinueWith(items) = action else { @@ -881,7 +887,7 @@ mod tests { ); let mut ctx: Vec = Vec::new(); - let action = interceptor.pre_llm_request(&mut ctx).await; + let action = interceptor.pre_llm_request(&mut ctx).await.unwrap(); assert!(!saw_handle.load(Ordering::Relaxed)); assert!(matches!(action, PreRequestAction::Continue)); @@ -938,7 +944,7 @@ mod tests { ); let mut info = task_tool_call_info("TaskList", serde_json::json!({"scope": "all"})); - let action = interceptor.pre_tool_call(&mut info).await; + let action = interceptor.pre_tool_call(&mut info).await.unwrap(); match action { PreToolAction::SyntheticResult(result) => { @@ -1000,7 +1006,7 @@ mod tests { context: info.context, }; - let action = interceptor.post_tool_call(&mut result_info).await; + let action = interceptor.post_tool_call(&mut result_info).await.unwrap(); assert_eq!(action, PostToolAction::Abort("post tool abort".to_string())); assert_eq!(count.load(Ordering::Relaxed), 1); @@ -1036,7 +1042,7 @@ mod tests { ); let history = vec![Item::user_message("hi"), Item::assistant_message("done")]; - let action = interceptor.on_turn_end(&history).await; + let action = interceptor.on_turn_end(&history).await.unwrap(); assert!(matches!(action, TurnEndAction::Pause)); assert_eq!(count.load(Ordering::Relaxed), 1); @@ -1073,7 +1079,7 @@ mod tests { let ctx_items = vec![Item::user_message("hi")]; for _ in 0..23 { let mut ctx = ctx_items.clone(); - let action = interceptor.pre_llm_request(&mut ctx).await; + let action = interceptor.pre_llm_request(&mut ctx).await.unwrap(); assert!(matches!(action, PreRequestAction::Continue)); usage_tracker.record_usage(&agen::event::UsageEvent { input_tokens: Some(10), @@ -1085,7 +1091,7 @@ mod tests { } let mut ctx = ctx_items.clone(); - let action = interceptor.pre_llm_request(&mut ctx).await; + let action = interceptor.pre_llm_request(&mut ctx).await.unwrap(); let appended_len = match action { PreRequestAction::ContinueWith(items) => items.len(), other => panic!("expected reminder append, got {other:?}"), @@ -1210,7 +1216,7 @@ mod tests { let error = interceptor.pending_history_appends().await.unwrap_err(); - assert!(error.contains("failed to render notify_wrapper")); + assert!(error.message().contains("failed to render notify_wrapper")); let requeued = buffer.drain(); assert_eq!(requeued.len(), 1); } @@ -1269,7 +1275,7 @@ mod tests { None, ); let mut ctx: Vec = vec![Item::user_message("hi")]; - let action = interceptor.pre_llm_request(&mut ctx).await; + let action = interceptor.pre_llm_request(&mut ctx).await.unwrap(); assert!(matches!(action, PreRequestAction::Continue)); assert_eq!(ctx.len(), 1, "pre_llm_request must not append notifies"); @@ -1299,7 +1305,7 @@ mod tests { None, ); let mut ctx: Vec = Vec::new(); - let action = interceptor.pre_llm_request(&mut ctx).await; + let action = interceptor.pre_llm_request(&mut ctx).await.unwrap(); assert!(matches!(action, PreRequestAction::Cancel(_))); assert!(first_called.load(Ordering::Relaxed)); diff --git a/crates/worker/src/worker.rs b/crates/worker/src/worker.rs index d8b84640..f98a9b5a 100644 --- a/crates/worker/src/worker.rs +++ b/crates/worker/src/worker.rs @@ -3733,6 +3733,7 @@ impl Worker { | EngineRunExit::Yielded | EngineRunExit::Interrupted(RunInterruptionReason::Cancelled) | EngineRunExit::Interrupted(RunInterruptionReason::ContextWindowExceeded) + | EngineRunExit::Interrupted(RunInterruptionReason::Interceptor(_)) | EngineRunExit::Interrupted(RunInterruptionReason::Unexpected(_)) ); let active_run_turn_count = self.engine.as_ref().unwrap().active_run_turn_count(); @@ -6129,12 +6130,14 @@ fn run_interruption_reason_error_code(reason: &RunInterruptionReason) -> ErrorCo RunInterruptionReason::Unexpected(EngineError::Tool(_)) => ErrorCode::ToolError, RunInterruptionReason::LimitReached | RunInterruptionReason::Cancelled + | RunInterruptionReason::Interceptor(_) | RunInterruptionReason::Unexpected( EngineError::Aborted(_) | EngineError::Cancelled | EngineError::PauseRequested | EngineError::ConfigWarnings(_) | EngineError::HistoryAppend(_) + | EngineError::Interceptor(_) | EngineError::ToolAttemptFence(_), ) => ErrorCode::Internal, } @@ -6145,6 +6148,7 @@ fn run_interruption_reason_message(reason: &RunInterruptionReason) -> String { RunInterruptionReason::LimitReached => "engine turn limit reached".to_string(), RunInterruptionReason::ContextWindowExceeded => "model context window reached".to_string(), RunInterruptionReason::Cancelled => "engine run cancelled".to_string(), + RunInterruptionReason::Interceptor(failure) => failure.to_string(), RunInterruptionReason::Unexpected(error) => format!("unexpected engine failure: {error}"), } }