fix: harden interceptor lifecycle contracts

This commit is contained in:
2026-09-03 16:46:02 +09:00
parent e62c7cf4f5
commit eac4a0c071
11 changed files with 877 additions and 190 deletions
+134 -30
View File
@@ -10,9 +10,11 @@ use std::sync::{Arc, Mutex};
use agen::Item;
use agen::interceptor::{
AssistantTurnEndContext, Interceptor, InterceptorError, InterceptorPoint, InterceptorResult,
PostToolAction, PreLlmRequestContext, PreRequestAction, PreToolAction, PromptAction,
PromptSubmitContext, RunExitContext, ToolCallInfo, ToolResultInfo, TurnEndAction,
AssistantTurnEndContext, Interceptor, InterceptorError, InterceptorErrorCategory,
InterceptorPhase as InterceptorPoint, InterceptorResult, MAX_INTERCEPTOR_DIAGNOSTIC_BYTES,
PendingHistoryAppendsContext, PostToolAction, PreLlmRequestContext, PreRequestAction,
PreToolAction, PromptAction, PromptSubmitContext, RunExitContext, ToolCallInfo, ToolResultInfo,
TurnEndAction,
};
use agen::llm_client::{
ClientError, LlmClient, Request, ResponseStream,
@@ -620,7 +622,7 @@ struct YieldOnce {
impl Interceptor for YieldOnce {
async fn pre_llm_request(
&self,
_context: PreLlmRequestContext<'_>,
_context: PreLlmRequestContext<'_, ()>,
) -> InterceptorResult<PreRequestAction> {
Ok(if self.calls.fetch_add(1, Ordering::SeqCst) == 0 {
PreRequestAction::Yield
@@ -636,7 +638,10 @@ struct PauseToolOnce {
#[async_trait]
impl Interceptor for PauseToolOnce {
async fn pre_tool_call(&self, _info: &mut ToolCallInfo) -> InterceptorResult<PreToolAction> {
async fn pre_tool_call(
&self,
_info: &mut ToolCallInfo<'_, ()>,
) -> InterceptorResult<PreToolAction> {
Ok(if self.calls.fetch_add(1, Ordering::SeqCst) == 0 {
PreToolAction::Pause
} else {
@@ -653,7 +658,7 @@ struct ContinueTurnOnce {
impl Interceptor for ContinueTurnOnce {
async fn on_assistant_turn_end(
&self,
_context: AssistantTurnEndContext<'_>,
_context: AssistantTurnEndContext<'_, ()>,
) -> InterceptorResult<TurnEndAction> {
Ok(if self.calls.fetch_add(1, Ordering::SeqCst) == 0 {
TurnEndAction::ContinueWithMessages(vec![Item::system_message("continue")])
@@ -680,7 +685,10 @@ impl FailingLifecycleInterceptor {
fn record<T>(&self, point: InterceptorPoint, action: T) -> InterceptorResult<T> {
self.calls.lock().unwrap().push(point);
if self.failure == point {
Err(InterceptorError::new(format!("{point} rejected")))
Err(InterceptorError::new(
InterceptorErrorCategory::Policy,
format!("{point} rejected"),
))
} else {
Ok(action)
}
@@ -695,47 +703,56 @@ impl FailingLifecycleInterceptor {
impl Interceptor for FailingLifecycleInterceptor {
async fn on_prompt_submit(
&self,
_context: PromptSubmitContext<'_>,
_context: PromptSubmitContext<'_, ()>,
) -> InterceptorResult<PromptAction> {
tokio::task::yield_now().await;
self.record(InterceptorPoint::PromptSubmit, PromptAction::Continue)
}
async fn pending_history_appends(&self) -> InterceptorResult<Vec<Item>> {
async fn pending_history_appends(
&self,
_context: PendingHistoryAppendsContext<'_, ()>,
) -> InterceptorResult<Vec<Item>> {
tokio::task::yield_now().await;
self.record(InterceptorPoint::PendingHistoryAppends, Vec::new())
}
async fn pre_llm_request(
&self,
_context: PreLlmRequestContext<'_>,
_context: PreLlmRequestContext<'_, ()>,
) -> InterceptorResult<PreRequestAction> {
tokio::task::yield_now().await;
self.record(InterceptorPoint::PreLlmRequest, PreRequestAction::Continue)
}
async fn pre_tool_call(&self, _info: &mut ToolCallInfo) -> InterceptorResult<PreToolAction> {
async fn pre_tool_call(
&self,
_info: &mut ToolCallInfo<'_, ()>,
) -> InterceptorResult<PreToolAction> {
tokio::task::yield_now().await;
self.record(InterceptorPoint::PreToolCall, PreToolAction::Continue)
}
async fn post_tool_call(&self, _info: &ToolResultInfo) -> InterceptorResult<PostToolAction> {
async fn post_tool_call(
&self,
_info: &ToolResultInfo<'_, ()>,
) -> InterceptorResult<PostToolAction> {
tokio::task::yield_now().await;
self.record(InterceptorPoint::PostToolCall, PostToolAction::Continue)
}
async fn on_assistant_turn_end(
&self,
context: AssistantTurnEndContext<'_>,
context: AssistantTurnEndContext<'_, ()>,
) -> InterceptorResult<TurnEndAction> {
tokio::task::yield_now().await;
assert!(context.history.ends_with(context.assistant_items));
assert!(context.history.ends_with(context.assistant_entries));
if !context.tool_calls.is_empty() {
assert_eq!(
context
.assistant_items
.assistant_entries
.iter()
.filter(|item| matches!(item, Item::ToolCall { .. }))
.filter(|entry| matches!(&entry.item, Item::ToolCall { .. }))
.count(),
context.tool_calls.len()
);
@@ -743,7 +760,7 @@ impl Interceptor for FailingLifecycleInterceptor {
self.record(InterceptorPoint::AssistantTurnEnd, TurnEndAction::Finish)
}
async fn on_run_exit(&self, _context: RunExitContext<'_>) -> InterceptorResult<()> {
async fn on_run_exit(&self, _context: RunExitContext<'_, ()>) -> InterceptorResult<()> {
tokio::task::yield_now().await;
self.record(InterceptorPoint::RunExit, ())
}
@@ -795,7 +812,7 @@ fn expected_interceptor_calls(failure: InterceptorPoint) -> Vec<InterceptorPoint
}
#[tokio::test]
async fn interceptor_failures_are_typed_unexpected_and_each_lifecycle_point_runs_once() {
async fn interceptor_failures_are_typed_and_terminal_observer_preserves_original_exit() {
use InterceptorPoint as Point;
for failure_point in [
@@ -828,15 +845,23 @@ async fn interceptor_failures_are_typed_unexpected_and_each_lifecycle_point_runs
let mut engine = engine.lock(&history);
let exit = engine.run(&mut history, "test").await;
let EngineRunExit::Interrupted(RunInterruptionReason::Unexpected(
EngineError::Interceptor(failure),
)) = exit
else {
panic!("expected typed interceptor interruption at {failure_point}, got {exit:?}");
let failure = if failure_point == Point::RunExit {
assert!(matches!(exit, EngineRunExit::Finished));
engine
.last_run_exit_observer_failure()
.expect("terminal observer diagnostic should be retained")
} else {
let EngineRunExit::Interrupted(RunInterruptionReason::Unexpected(
EngineError::Interceptor(failure),
)) = &exit
else {
panic!("expected typed interceptor interruption at {failure_point}, got {exit:?}");
};
failure
};
assert_eq!(failure.point(), failure_point);
assert_eq!(failure.phase(), failure_point);
assert_eq!(
failure.error().message(),
failure.error().diagnostic(),
format!("{failure_point} rejected")
);
assert_eq!(
@@ -854,6 +879,85 @@ async fn interceptor_failures_are_typed_unexpected_and_each_lifecycle_point_runs
}
}
#[test]
fn interceptor_error_keeps_typed_category_and_bounded_utf8_diagnostic() {
let error = InterceptorError::new(
InterceptorErrorCategory::Dependency,
"".repeat(MAX_INTERCEPTOR_DIAGNOSTIC_BYTES),
);
assert_eq!(error.category(), InterceptorErrorCategory::Dependency);
assert!(error.diagnostic().len() <= MAX_INTERCEPTOR_DIAGNOSTIC_BYTES);
assert!(
error
.diagnostic()
.is_char_boundary(error.diagnostic().len())
);
}
struct FailingRunExitObserver {
pause: bool,
}
#[async_trait]
impl Interceptor for FailingRunExitObserver {
async fn on_assistant_turn_end(
&self,
_context: AssistantTurnEndContext<'_, ()>,
) -> InterceptorResult<TurnEndAction> {
Ok(if self.pause {
TurnEndAction::Pause
} else {
TurnEndAction::Finish
})
}
async fn on_run_exit(&self, _context: RunExitContext<'_, ()>) -> InterceptorResult<()> {
Err(InterceptorError::new(
InterceptorErrorCategory::Dependency,
"terminal audit unavailable",
))
}
}
#[tokio::test]
async fn terminal_observer_failure_preserves_paused_and_interrupted_exits() {
let mut paused_engine = Engine::new(MockLlmClient::new(completed_text_events()));
paused_engine.set_interceptor(FailingRunExitObserver { pause: true });
let mut paused_history = History::new();
let mut paused_engine = paused_engine.lock(&paused_history);
assert!(matches!(
paused_engine.run(&mut paused_history, "pause").await,
EngineRunExit::Paused
));
assert_eq!(
paused_engine
.last_run_exit_observer_failure()
.expect("paused observer diagnostic")
.error()
.category(),
InterceptorErrorCategory::Dependency
);
let mut interrupted_engine = Engine::new(MockLlmClient::new(completed_text_events()));
interrupted_engine.set_max_turns(Some(0));
interrupted_engine.set_interceptor(FailingRunExitObserver { pause: false });
let mut interrupted_history = History::new();
let mut interrupted_engine = interrupted_engine.lock(&interrupted_history);
assert!(matches!(
interrupted_engine
.run(&mut interrupted_history, "limit")
.await,
EngineRunExit::Interrupted(RunInterruptionReason::LimitReached)
));
assert_eq!(
interrupted_engine
.last_run_exit_observer_failure()
.expect("interrupted observer diagnostic")
.phase(),
InterceptorPoint::RunExit
);
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum TerminalMode {
Finish,
@@ -886,7 +990,7 @@ impl RecordingTerminalInterceptor {
impl Interceptor for RecordingTerminalInterceptor {
async fn pre_llm_request(
&self,
_context: PreLlmRequestContext<'_>,
_context: PreLlmRequestContext<'_, ()>,
) -> InterceptorResult<PreRequestAction> {
Ok(if self.mode == TerminalMode::Yield {
PreRequestAction::Yield
@@ -897,11 +1001,11 @@ impl Interceptor for RecordingTerminalInterceptor {
async fn on_assistant_turn_end(
&self,
context: AssistantTurnEndContext<'_>,
context: AssistantTurnEndContext<'_, ()>,
) -> InterceptorResult<TurnEndAction> {
assert!(!context.assistant_items.is_empty());
assert!(!context.assistant_entries.is_empty());
assert!(
context.history.ends_with(context.assistant_items),
context.history.ends_with(context.assistant_entries),
"assistant-turn callback must observe committed terminal items"
);
let turn = self.assistant_turns.fetch_add(1, Ordering::SeqCst);
@@ -912,7 +1016,7 @@ impl Interceptor for RecordingTerminalInterceptor {
})
}
async fn on_run_exit(&self, context: RunExitContext<'_>) -> InterceptorResult<()> {
async fn on_run_exit(&self, context: RunExitContext<'_, ()>) -> InterceptorResult<()> {
let kind = match context.exit {
EngineRunExit::Finished => "finished",
EngineRunExit::Paused => "paused",