diff --git a/crates/agen/examples/engine_cli.rs b/crates/agen/examples/engine_cli.rs index 10eb28ef..ee90ca7e 100644 --- a/crates/agen/examples/engine_cli.rs +++ b/crates/agen/examples/engine_cli.rs @@ -280,7 +280,10 @@ impl ToolResultPrinterPolicy { #[async_trait] impl Interceptor for ToolResultPrinterPolicy { - async fn post_tool_call(&self, info: &ToolResultInfo) -> InterceptorResult { + async fn post_tool_call( + &self, + info: &ToolResultInfo<'_, ()>, + ) -> InterceptorResult { let name = self .call_names .lock() diff --git a/crates/agen/src/engine.rs b/crates/agen/src/engine.rs index 276ee3a4..644cf161 100644 --- a/crates/agen/src/engine.rs +++ b/crates/agen/src/engine.rs @@ -15,10 +15,12 @@ use crate::{ }, handler::{ErrorKind, StatusKind, ToolUseBlockStart, UsageKind}, interceptor::{ - AssistantTurnEndContext, DefaultInterceptor, Interceptor, InterceptorFailure, - InterceptorPoint, PostToolAction, PreLlmRequestContext, PreRequestAction, PreToolAction, - PromptAction, PromptSubmitContext, RunExitContext, ToolCallInfo, ToolResultInfo, - TurnEndAction, + AssistantTurnEndContext, DefaultInterceptor, Interceptor, InterceptorCallId, + InterceptorCounter, InterceptorCounters, InterceptorError, InterceptorErrorCategory, + InterceptorFailure, InterceptorInvocation, InterceptorPhase, InterceptorRunId, + InterceptorTurnId, PendingHistoryAppendsContext, PostToolAction, PreLlmRequestContext, + PreRequestAction, PreToolAction, PromptAction, PromptSubmitContext, RunExitContext, + ToolCallInfo, ToolResultInfo, TurnEndAction, }, llm_client::{ ClientError, ConfigWarning, LlmClient, Request, RequestConfig, ResponseStream, @@ -186,7 +188,7 @@ impl From> for EngineRunExit { /// Result of [`Engine::run`] or [`Engine::resume`]. /// /// Contains the `Locked` Engine (ready for subsequent runs) and the outcome. -pub struct EngineRunOutput { +pub struct EngineRunOutput { /// The Engine, now in Locked state. pub engine: Engine, /// Outcome of the turn. @@ -310,7 +312,7 @@ enum StreamCompletion { Interrupted { reason: String }, } -pub struct Engine { +pub struct Engine { /// LLM client client: C, /// Retry policy for opening an LLM response stream. @@ -327,7 +329,7 @@ pub struct Engine { /// Tool server handle tool_server: ToolServerHandle, /// Interceptor for control-flow decisions - interceptor: Box, + interceptor: Box>, /// System prompt system_prompt: Option, /// History length at lock time (only meaningful in Locked state) @@ -346,6 +348,11 @@ pub struct Engine { /// `max_turns` is enforced against this run-scoped count rather than the /// cumulative `turn_count` above. active_run_turn_count: Option, + /// Identity retained across pause/yield and resume. + active_run_id: Option, + next_run_id: u64, + interceptor_invocation_count: usize, + last_run_exit_observer_failure: Option, /// LlmCall count (per-Engine running counter, monotonic). Unlike /// `turn_count` this never collapses retries. llm_call_count: usize, @@ -426,18 +433,57 @@ pub struct Engine { _state: PhantomData<(S, A)>, } -impl Engine { +impl Engine { fn start_logical_run(&mut self) { self.active_run_turn_count = Some(0); + self.active_run_id = Some(InterceptorRunId(self.next_run_id)); + self.next_run_id = self.next_run_id.wrapping_add(1).max(1); + self.interceptor_invocation_count = 0; + self.last_run_exit_observer_failure = None; } fn ensure_logical_run(&mut self) { self.active_run_turn_count.get_or_insert(0); + if self.active_run_id.is_none() { + self.active_run_id = Some(InterceptorRunId(self.next_run_id)); + self.next_run_id = self.next_run_id.wrapping_add(1).max(1); + self.interceptor_invocation_count = 0; + } + } + + fn interceptor_invocation( + &mut self, + phase: InterceptorPhase, + turn_id: Option, + call_id: Option, + tool_call: usize, + ) -> InterceptorInvocation { + let invocation = self.interceptor_invocation_count; + self.interceptor_invocation_count = self.interceptor_invocation_count.saturating_add(1); + InterceptorInvocation { + run_id: self + .active_run_id + .expect("logical run identity must exist before interception"), + turn_id: turn_id.map(|value| InterceptorTurnId(value as u64)), + call_id, + phase, + counters: InterceptorCounters { + invocation: InterceptorCounter::from_usize(invocation), + engine_turn: InterceptorCounter::from_usize(self.turn_count), + run_turn: InterceptorCounter::from_usize( + self.active_run_turn_count.unwrap_or_default(), + ), + llm_call: InterceptorCounter::from_usize(self.llm_call_count), + tool_batch: InterceptorCounter::from_usize(self.tool_execution_batch_count), + tool_call: InterceptorCounter::from_usize(tool_call), + }, + } } fn finish_logical_run(&mut self, exit: &EngineRunExit) { if !matches!(exit, EngineRunExit::Paused | EngineRunExit::Yielded) { self.active_run_turn_count = None; + self.active_run_id = None; } } @@ -743,7 +789,7 @@ impl Engine { /// The interceptor governs approval, skip, pause, and abort decisions /// at key points in the execution loop. If not set, the default /// interceptor is used (all Continue / Finish). - pub fn set_interceptor(&mut self, interceptor: impl Interceptor + 'static) { + pub fn set_interceptor(&mut self, interceptor: impl Interceptor + 'static) { self.interceptor = Box::new(interceptor); } @@ -844,6 +890,10 @@ impl Engine { /// /// `Some` is retained only while Pause or Yield permits a later /// [`resume`](Self::resume). Terminal outcomes return this to `None`. + pub fn last_run_exit_observer_failure(&self) -> Option<&InterceptorFailure> { + self.last_run_exit_observer_failure.as_ref() + } + pub fn active_run_turn_count(&self) -> Option { self.active_run_turn_count } @@ -855,6 +905,13 @@ impl Engine { /// [`resume`](Self::resume) starts a fresh budget. pub fn set_active_run_turn_count(&mut self, turn_count: Option) { self.active_run_turn_count = turn_count; + if turn_count.is_none() { + self.active_run_id = None; + } else if self.active_run_id.is_none() { + self.active_run_id = Some(InterceptorRunId(self.next_run_id)); + self.next_run_id = self.next_run_id.wrapping_add(1).max(1); + self.interceptor_invocation_count = 0; + } } /// Get the current LlmCall count (per-Engine running counter, never @@ -1082,19 +1139,24 @@ impl Engine { async fn finalize_run_exit( &mut self, + history: &History, result: Result, ) -> EngineRunExit { let exit = EngineRunExit::from(result); - let exit = match self + let invocation = self.interceptor_invocation(InterceptorPhase::RunExit, None, None, 0); + self.last_run_exit_observer_failure = None; + if let Err(error) = self .interceptor - .on_run_exit(RunExitContext { exit: &exit }) + .on_run_exit(RunExitContext { + invocation, + exit: &exit, + history: history.entries(), + }) .await { - Ok(()) => exit, - Err(error) => EngineRunExit::Interrupted(RunInterruptionReason::Unexpected( - InterceptorFailure::new(InterceptorPoint::RunExit, error).into(), - )), - }; + self.last_run_exit_observer_failure = + Some(InterceptorFailure::new(InterceptorPhase::RunExit, error)); + } self.finish_logical_run(&exit); exit } @@ -1167,9 +1229,18 @@ impl Engine { // Phase 1: Apply pre_tool_call interceptor (determine skip/abort/synthetic result) let mut approved_calls = Vec::new(); for (call_index, mut tool_call) in tool_calls.into_iter().enumerate() { + let expected_tool_use_id = tool_call.id.clone(); let context = ToolExecutionContext::new(&tool_call.id, &batch_id, call_index); if let Some((meta, tool)) = self.tool_server.get_tool(&tool_call.name) { + let invocation = self.interceptor_invocation( + InterceptorPhase::PreToolCall, + Some(self.turn_count.saturating_sub(1)), + Some(InterceptorCallId::Tool(expected_tool_use_id.clone())), + call_index, + ); let mut info = ToolCallInfo { + invocation, + history: history.entries(), call: tool_call.clone(), meta, tool, @@ -1182,16 +1253,36 @@ impl Engine { .await .map_err(|error| { EngineError::from(InterceptorFailure::new( - InterceptorPoint::PreToolCall, + InterceptorPhase::PreToolCall, error, )) })?; + if info.call.id != expected_tool_use_id { + return Err(InterceptorFailure::new( + InterceptorPhase::PreToolCall, + InterceptorError::new( + InterceptorErrorCategory::ContractViolation, + "pre-tool interceptor changed immutable tool call identity", + ), + ) + .into()); + } match pre_tool_action { PreToolAction::Continue => {} PreToolAction::Skip => { continue; } PreToolAction::SyntheticResult(result) => { + if result.tool_use_id != expected_tool_use_id { + return Err(InterceptorFailure::new( + InterceptorPhase::PreToolCall, + InterceptorError::new( + InterceptorErrorCategory::ContractViolation, + "synthetic tool result changed immutable tool call identity", + ), + ) + .into()); + } let tool_call = info.call; let mut context = info.context; context.call_id = tool_call.id.clone(); @@ -1538,7 +1629,15 @@ impl Engine { self.emit_tool_result(&tool_result); if let Some((tool_call, meta, tool, context)) = call_info { + let invocation = self.interceptor_invocation( + InterceptorPhase::PostToolCall, + Some(self.turn_count.saturating_sub(1)), + Some(InterceptorCallId::Tool(tool_call.id.clone())), + context.call_index, + ); let info = ToolResultInfo { + invocation, + history: history.entries(), call: tool_call.clone(), result: tool_result, meta: meta.clone(), @@ -1551,7 +1650,7 @@ impl Engine { .await .map_err(|error| { EngineError::from(InterceptorFailure::new( - InterceptorPoint::PostToolCall, + InterceptorPhase::PostToolCall, error, )) })?; @@ -1622,13 +1721,22 @@ impl Engine { // These are committed *before* the per-request clone so they // participate in the LLM request below and get persisted by // the caller that owns durable history. + let pending_invocation = self.interceptor_invocation( + InterceptorPhase::PendingHistoryAppends, + Some(current_turn), + None, + 0, + ); let pending = self .interceptor - .pending_history_appends() + .pending_history_appends(PendingHistoryAppendsContext { + invocation: pending_invocation, + history: history.entries(), + }) .await .map_err(|error| { EngineError::from(InterceptorFailure::new( - InterceptorPoint::PendingHistoryAppends, + InterceptorPhase::PendingHistoryAppends, error, )) })?; @@ -1698,15 +1806,23 @@ impl Engine { } // Interceptor: pre_llm_request + let request_invocation = self.interceptor_invocation( + InterceptorPhase::PreLlmRequest, + Some(current_turn), + Some(InterceptorCallId::Llm(self.llm_call_count as u64)), + 0, + ); let pre_request_action = self .interceptor .pre_llm_request(PreLlmRequestContext { + invocation: request_invocation, items: &mut request_context, + history: history.entries(), }) .await .map_err(|error| { EngineError::from(InterceptorFailure::new( - InterceptorPoint::PreLlmRequest, + InterceptorPhase::PreLlmRequest, error, )) })?; @@ -1822,21 +1938,29 @@ impl Engine { let tool_calls = self.tool_call_collector.take_collected(); let assistant_items = self.build_assistant_items(&reasoning_items, &text_blocks, &tool_calls); - let committed_assistant_items = assistant_items.clone(); + let assistant_start = history.len(); self.append_history_items(history, assistant_items, annotate)?; - let assistant_turn_history = history.items_cloned(); + let assistant_invocation = self.interceptor_invocation( + InterceptorPhase::AssistantTurnEnd, + Some(current_turn), + Some(InterceptorCallId::Llm( + self.llm_call_count.saturating_sub(1) as u64, + )), + 0, + ); let assistant_turn_action = self .interceptor .on_assistant_turn_end(AssistantTurnEndContext { - assistant_items: &committed_assistant_items, - history: &assistant_turn_history, + invocation: assistant_invocation, + assistant_entries: &history.entries()[assistant_start..], + history: history.entries(), tool_calls: &tool_calls, }) .await .map_err(|error| { EngineError::from(InterceptorFailure::new( - InterceptorPoint::AssistantTurnEnd, + InterceptorPhase::AssistantTurnEnd, error, )) })?; @@ -2145,7 +2269,7 @@ impl Engine { } } -impl Engine { +impl Engine { /// Create a new annotated Engine (in Mutable state). pub fn new_annotated(client: C) -> Self { let text_block_collector = TextBlockCollector::new(); @@ -2173,6 +2297,10 @@ impl Engine { locked_prefix_len: 0, turn_count: 0, active_run_turn_count: None, + active_run_id: None, + next_run_id: 1, + interceptor_invocation_count: 0, + last_run_exit_observer_failure: None, llm_call_count: 0, tool_execution_batch_count: 0, max_turns: None, @@ -2448,6 +2576,10 @@ impl Engine { locked_prefix_len, turn_count: self.turn_count, active_run_turn_count: self.active_run_turn_count, + active_run_id: self.active_run_id, + next_run_id: self.next_run_id, + interceptor_invocation_count: self.interceptor_invocation_count, + last_run_exit_observer_failure: self.last_run_exit_observer_failure, llm_call_count: self.llm_call_count, tool_execution_batch_count: self.tool_execution_batch_count, max_turns: self.max_turns, @@ -2524,7 +2656,7 @@ impl Engine { } } -impl Engine { +impl Engine { /// Execute a turn /// /// Adds a new user message to history and sends a request to the LLM. @@ -2538,7 +2670,7 @@ impl Engine { let result = self .run_result_with_annotation(history, user_input.into(), annotate) .await; - self.finalize_run_exit(result).await + self.finalize_run_exit(history, result).await } async fn run_result_with_annotation( @@ -2549,16 +2681,21 @@ impl Engine { ) -> Result { // Supplying new user input abandons any paused/yielded logical run. self.active_run_turn_count = None; + self.active_run_id = None; + self.start_logical_run(); let mut user_item = Item::user_message(user_input); + let invocation = self.interceptor_invocation(InterceptorPhase::PromptSubmit, None, None, 0); let prompt_action = self .interceptor .on_prompt_submit(PromptSubmitContext { + invocation, item: &mut user_item, + history: history.entries(), }) .await .map_err(|error| { EngineError::from(InterceptorFailure::new( - InterceptorPoint::PromptSubmit, + InterceptorPhase::PromptSubmit, error, )) })?; @@ -2571,7 +2708,6 @@ impl Engine { if !extras.is_empty() { self.append_history_items(history, extras, annotate)?; } - self.start_logical_run(); match self.run_turn_loop(history, annotate).await { Err(EngineError::PauseRequested) => Ok(EngineResult::Paused), other => other, @@ -2585,7 +2721,7 @@ impl Engine { annotate: &mut impl FnMut(&Item) -> Result, ) -> EngineRunExit { let result = self.resume_result_with_annotation(history, annotate).await; - self.finalize_run_exit(result).await + self.finalize_run_exit(history, result).await } async fn resume_result_with_annotation( @@ -2623,6 +2759,10 @@ impl Engine { locked_prefix_len: 0, turn_count: self.turn_count, active_run_turn_count: self.active_run_turn_count, + active_run_id: self.active_run_id, + next_run_id: self.next_run_id, + interceptor_invocation_count: self.interceptor_invocation_count, + last_run_exit_observer_failure: self.last_run_exit_observer_failure, llm_call_count: self.llm_call_count, tool_execution_batch_count: self.tool_execution_batch_count, max_turns: self.max_turns, diff --git a/crates/agen/src/interceptor.rs b/crates/agen/src/interceptor.rs index b85e0d9e..28b2afa6 100644 --- a/crates/agen/src/interceptor.rs +++ b/crates/agen/src/interceptor.rs @@ -10,51 +10,73 @@ use async_trait::async_trait; use crate::Item; use crate::engine::EngineRunExit; +use crate::history::HistoryEntry; use crate::tool::{Tool, ToolCall, ToolExecutionContext, ToolMeta, ToolResult}; // ============================================================================= -// Failure Types +// Typed lifecycle metadata and failures // ============================================================================= -/// 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. +/// Maximum UTF-8 byte length retained for interceptor diagnostics. +pub const MAX_INTERCEPTOR_DIAGNOSTIC_BYTES: usize = 1024; + +/// Stable category for the source of an interceptor failure. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum InterceptorErrorCategory { + Policy, + Dependency, + ContractViolation, + Internal, +} + +impl std::fmt::Display for InterceptorErrorCategory { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str(match self { + Self::Policy => "policy", + Self::Dependency => "dependency", + Self::ContractViolation => "contract_violation", + Self::Internal => "internal", + }) + } +} + +/// A typed, bounded failure returned by an [`Interceptor`] implementation. #[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)] -#[error("{message}")] +#[error("{category}: {diagnostic}")] pub struct InterceptorError { - message: String, + category: InterceptorErrorCategory, + diagnostic: String, } impl InterceptorError { - /// Create an interceptor failure with a caller-defined message. - pub fn new(message: impl Into) -> Self { + pub fn new(category: InterceptorErrorCategory, diagnostic: impl Into) -> Self { + let mut diagnostic = diagnostic.into(); + if diagnostic.len() > MAX_INTERCEPTOR_DIAGNOSTIC_BYTES { + let mut end = MAX_INTERCEPTOR_DIAGNOSTIC_BYTES; + while !diagnostic.is_char_boundary(end) { + end -= 1; + } + diagnostic.truncate(end); + } Self { - message: message.into(), + category, + diagnostic, } } - /// Return the failure message supplied by the interceptor. - pub fn message(&self) -> &str { - &self.message + pub fn category(&self) -> InterceptorErrorCategory { + self.category + } + + pub fn diagnostic(&self) -> &str { + &self.diagnostic } } -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 { +/// The lifecycle phase at which an interceptor callback executes. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)] +pub enum InterceptorPhase { + #[default] PromptSubmit, PendingHistoryAppends, PreLlmRequest, @@ -64,9 +86,9 @@ pub enum InterceptorPoint { RunExit, } -impl std::fmt::Display for InterceptorPoint { +impl std::fmt::Display for InterceptorPhase { fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - let name = match self { + formatter.write_str(match self { Self::PromptSubmit => "prompt_submit", Self::PendingHistoryAppends => "pending_history_appends", Self::PreLlmRequest => "pre_llm_request", @@ -74,66 +96,113 @@ impl std::fmt::Display for InterceptorPoint { Self::PostToolCall => "post_tool_call", Self::AssistantTurnEnd => "assistant_turn_end", Self::RunExit => "run_exit", - }; - formatter.write_str(name) + }) } } -/// An interceptor failure bound to the exact Engine lifecycle point that ran it. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)] +pub struct InterceptorRunId(pub u64); + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub struct InterceptorTurnId(pub u64); + +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +pub enum InterceptorCallId { + Llm(u64), + Tool(String), +} + +/// Saturating public counter used by interceptor contexts. +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Default)] +pub struct InterceptorCounter(u32); + +impl InterceptorCounter { + pub fn from_usize(value: usize) -> Self { + Self(u32::try_from(value).unwrap_or(u32::MAX)) + } + + pub fn get(self) -> u32 { + self.0 + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)] +pub struct InterceptorCounters { + pub invocation: InterceptorCounter, + pub engine_turn: InterceptorCounter, + pub run_turn: InterceptorCounter, + pub llm_call: InterceptorCounter, + pub tool_batch: InterceptorCounter, + pub tool_call: InterceptorCounter, +} + +/// Identity, phase, and bounded counters common to every lifecycle callback. +#[derive(Debug, Clone, PartialEq, Eq, Default)] +pub struct InterceptorInvocation { + pub run_id: InterceptorRunId, + pub turn_id: Option, + pub call_id: Option, + pub phase: InterceptorPhase, + pub counters: InterceptorCounters, +} + +/// An interceptor failure bound to the exact Engine lifecycle phase that ran it. #[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)] -#[error("{point} interceptor failed: {error}")] +#[error("{phase} interceptor failed: {error}")] pub struct InterceptorFailure { - point: InterceptorPoint, + phase: InterceptorPhase, #[source] error: InterceptorError, } impl InterceptorFailure { - pub(crate) fn new(point: InterceptorPoint, error: InterceptorError) -> Self { - Self { point, error } + pub(crate) fn new(phase: InterceptorPhase, error: InterceptorError) -> Self { + Self { phase, error } } - /// The lifecycle point that returned the failure. - pub fn point(&self) -> InterceptorPoint { - self.point + pub fn phase(&self) -> InterceptorPhase { + self.phase } - /// 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; // ============================================================================= // Lifecycle Contexts // ============================================================================= -/// Mutable prompt input presented before it is committed to Engine history. -pub struct PromptSubmitContext<'a> { +pub struct PromptSubmitContext<'a, A = ()> { + pub invocation: InterceptorInvocation, pub item: &'a mut Item, + pub history: &'a [HistoryEntry], } -/// Mutable provider-visible item projection presented before an LLM request. -pub struct PreLlmRequestContext<'a> { +pub struct PendingHistoryAppendsContext<'a, A = ()> { + pub invocation: InterceptorInvocation, + pub history: &'a [HistoryEntry], +} + +pub struct PreLlmRequestContext<'a, A = ()> { + pub invocation: InterceptorInvocation, pub items: &'a mut Vec, + pub history: &'a [HistoryEntry], } -/// A terminalized and committed assistant response at the next-phase boundary. -pub struct AssistantTurnEndContext<'a> { - /// The exact assistant items committed for this response. - pub assistant_items: &'a [Item], - /// The committed Engine history after the assistant items were appended. - pub history: &'a [Item], - /// Terminal tool calls collected from the response, if any. +pub struct AssistantTurnEndContext<'a, A = ()> { + pub invocation: InterceptorInvocation, + pub assistant_entries: &'a [HistoryEntry], + pub history: &'a [HistoryEntry], pub tool_calls: &'a [ToolCall], } -/// The one terminal outcome produced by a public Engine run or resume call. -pub struct RunExitContext<'a> { +pub struct RunExitContext<'a, A = ()> { + pub invocation: InterceptorInvocation, pub exit: &'a EngineRunExit, + pub history: &'a [HistoryEntry], } // ============================================================================= @@ -224,8 +293,9 @@ pub enum TurnEndAction { // ============================================================================= /// Context for pre-tool-call decisions. -pub struct ToolCallInfo { - /// Tool call information (modifiable). +pub struct ToolCallInfo<'a, A = ()> { + pub invocation: InterceptorInvocation, + pub history: &'a [HistoryEntry], pub call: ToolCall, /// Tool meta information. pub meta: ToolMeta, @@ -236,8 +306,9 @@ pub struct ToolCallInfo { } /// Context for post-tool-call decisions. -pub struct ToolResultInfo { - /// Original tool call. +pub struct ToolResultInfo<'a, A = ()> { + pub invocation: InterceptorInvocation, + pub history: &'a [HistoryEntry], pub call: ToolCall, /// Committed terminal tool execution result. pub result: ToolResult, @@ -258,17 +329,17 @@ pub struct ToolResultInfo { /// 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. +/// with the exact [`InterceptorPhase`] 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 { +pub trait Interceptor: Send + Sync { /// Called after receiving user input, before adding it to Engine history. async fn on_prompt_submit( &self, - _context: PromptSubmitContext<'_>, + _context: PromptSubmitContext<'_, A>, ) -> InterceptorResult { Ok(PromptAction::Continue) } @@ -291,7 +362,10 @@ 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) -> InterceptorResult> { + async fn pending_history_appends( + &self, + _context: PendingHistoryAppendsContext<'_, A>, + ) -> InterceptorResult> { Ok(Vec::new()) } @@ -305,18 +379,24 @@ pub trait Interceptor: Send + Sync { /// commits it to history before the request is sent. async fn pre_llm_request( &self, - _context: PreLlmRequestContext<'_>, + _context: PreLlmRequestContext<'_, A>, ) -> InterceptorResult { Ok(PreRequestAction::Continue) } /// Called before each tool is executed. - async fn pre_tool_call(&self, _info: &mut ToolCallInfo) -> InterceptorResult { + async fn pre_tool_call( + &self, + _info: &mut ToolCallInfo<'_, A>, + ) -> InterceptorResult { Ok(PreToolAction::Continue) } /// Called after each tool reaches one terminal result and that result is committed. - async fn post_tool_call(&self, _info: &ToolResultInfo) -> InterceptorResult { + async fn post_tool_call( + &self, + _info: &ToolResultInfo<'_, A>, + ) -> InterceptorResult { Ok(PostToolAction::Continue) } @@ -324,13 +404,13 @@ pub trait Interceptor: Send + Sync { /// the Engine decides whether to execute tools, continue, or finish. async fn on_assistant_turn_end( &self, - _context: AssistantTurnEndContext<'_>, + _context: AssistantTurnEndContext<'_, A>, ) -> InterceptorResult { Ok(TurnEndAction::Finish) } /// Called once for the terminal outcome of each public run or resume call. - async fn on_run_exit(&self, _context: RunExitContext<'_>) -> InterceptorResult<()> { + async fn on_run_exit(&self, _context: RunExitContext<'_, A>) -> InterceptorResult<()> { Ok(()) } } @@ -340,4 +420,4 @@ pub trait Interceptor: Send + Sync { pub(crate) struct DefaultInterceptor; #[async_trait] -impl Interceptor for DefaultInterceptor {} +impl Interceptor for DefaultInterceptor {} diff --git a/crates/agen/src/lib.rs b/crates/agen/src/lib.rs index 27bed76e..89bdaee9 100644 --- a/crates/agen/src/lib.rs +++ b/crates/agen/src/lib.rs @@ -27,8 +27,11 @@ pub use engine::{ pub use handler::ToolUseBlockStart; pub use history::{History, HistoryEntry}; pub use interceptor::{ - AssistantTurnEndContext, Interceptor, InterceptorError, InterceptorFailure, InterceptorPoint, - InterceptorResult, PreLlmRequestContext, PromptSubmitContext, RunExitContext, + AssistantTurnEndContext, Interceptor, InterceptorCallId, InterceptorCounter, + InterceptorCounters, InterceptorError, InterceptorErrorCategory, InterceptorFailure, + InterceptorInvocation, InterceptorPhase, InterceptorResult, InterceptorRunId, + InterceptorTurnId, MAX_INTERCEPTOR_DIAGNOSTIC_BYTES, PendingHistoryAppendsContext, + PreLlmRequestContext, PromptSubmitContext, RunExitContext, }; pub use message::{ContentPart, Item, Message, Role}; pub use tool::{ diff --git a/crates/agen/tests/annotated_history_test.rs b/crates/agen/tests/annotated_history_test.rs index 09a5019d..cf1be3c0 100644 --- a/crates/agen/tests/annotated_history_test.rs +++ b/crates/agen/tests/annotated_history_test.rs @@ -1,8 +1,15 @@ mod common; +use agen::interceptor::{ + AssistantTurnEndContext, Interceptor, InterceptorCallId, InterceptorInvocation, + InterceptorPhase, InterceptorResult, PendingHistoryAppendsContext, PreLlmRequestContext, + PreRequestAction, PromptAction, PromptSubmitContext, RunExitContext, TurnEndAction, +}; use agen::llm_client::event::{Event, ResponseStatus, StatusEvent}; use agen::{Engine, EngineError, History, HistoryEntry, Item, Role}; +use async_trait::async_trait; use common::MockLlmClient; +use std::sync::{Arc, Mutex}; fn completed_text_events(text: &str) -> Vec { vec![ @@ -47,6 +54,125 @@ async fn run_preserves_item_annotations_without_projecting_them() { assert_eq!(history.items_cloned().len(), 2); } +#[derive(Clone)] +struct AnnotationObservingInterceptor { + observed: Arc)>>>, +} + +impl AnnotationObservingInterceptor { + fn record(&self, invocation: &InterceptorInvocation, history: &[HistoryEntry]) { + self.observed.lock().unwrap().push(( + invocation.clone(), + history + .iter() + .map(|entry| entry.annotation.clone()) + .collect(), + )); + } +} + +#[async_trait] +impl Interceptor for AnnotationObservingInterceptor { + async fn on_prompt_submit( + &self, + context: PromptSubmitContext<'_, String>, + ) -> InterceptorResult { + self.record(&context.invocation, context.history); + Ok(PromptAction::Continue) + } + + async fn pending_history_appends( + &self, + context: PendingHistoryAppendsContext<'_, String>, + ) -> InterceptorResult> { + self.record(&context.invocation, context.history); + Ok(Vec::new()) + } + + async fn pre_llm_request( + &self, + context: PreLlmRequestContext<'_, String>, + ) -> InterceptorResult { + self.record(&context.invocation, context.history); + Ok(PreRequestAction::Continue) + } + + async fn on_assistant_turn_end( + &self, + context: AssistantTurnEndContext<'_, String>, + ) -> InterceptorResult { + assert_eq!(context.assistant_entries.len(), 1); + assert_eq!(context.assistant_entries[0].annotation, "2:assistant"); + self.record(&context.invocation, context.history); + Ok(TurnEndAction::Finish) + } + + async fn on_run_exit(&self, context: RunExitContext<'_, String>) -> InterceptorResult<()> { + self.record(&context.invocation, context.history); + Ok(()) + } +} + +#[tokio::test] +async fn interceptor_contexts_preserve_annotations_and_typed_lifecycle_identity() { + let client = MockLlmClient::new(completed_text_events("assistant reply")); + let mut engine = Engine::<_, agen::state::Mutable, String>::new_annotated(client); + let observed = Arc::new(Mutex::new(Vec::new())); + engine.set_interceptor(AnnotationObservingInterceptor { + observed: observed.clone(), + }); + let mut history = History::::new(); + let mut next = 0usize; + let mut annotate = |item: &Item| { + next += 1; + let kind = if item.is_assistant_message() { + "assistant" + } else { + "user" + }; + Ok(format!("{next}:{kind}")) + }; + + let output = engine + .run_with_annotation(&mut history, "hello", &mut annotate) + .await; + assert!(matches!(output.result, agen::EngineRunExit::Finished)); + + let observed = observed.lock().unwrap(); + let phases: Vec<_> = observed + .iter() + .map(|(invocation, _)| invocation.phase) + .collect(); + assert_eq!( + phases, + [ + InterceptorPhase::PromptSubmit, + InterceptorPhase::PendingHistoryAppends, + InterceptorPhase::PreLlmRequest, + InterceptorPhase::AssistantTurnEnd, + InterceptorPhase::RunExit, + ] + ); + assert!( + observed + .iter() + .all(|(invocation, _)| invocation.run_id == observed[0].0.run_id) + ); + assert_eq!( + observed + .iter() + .map(|(invocation, _)| invocation.counters.invocation.get()) + .collect::>(), + [0, 1, 2, 3, 4] + ); + assert_eq!(observed[2].0.call_id, Some(InterceptorCallId::Llm(0))); + assert_eq!(observed[3].0.call_id, Some(InterceptorCallId::Llm(0))); + assert_eq!(observed[1].1, ["1:user"]); + assert_eq!(observed[2].1, ["1:user"]); + assert_eq!(observed[3].1, ["1:user", "2:assistant"]); + assert_eq!(observed[4].1, ["1:user", "2:assistant"]); +} + #[test] fn append_failure_does_not_make_item_live() { let client = MockLlmClient::new(vec![]); diff --git a/crates/agen/tests/engine_state_test.rs b/crates/agen/tests/engine_state_test.rs index ab810e9e..5060015d 100644 --- a/crates/agen/tests/engine_state_test.rs +++ b/crates/agen/tests/engine_state_test.rs @@ -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 { 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 { + async fn pre_tool_call( + &self, + _info: &mut ToolCallInfo<'_, ()>, + ) -> InterceptorResult { 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 { 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(&self, point: InterceptorPoint, action: T) -> InterceptorResult { 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 { tokio::task::yield_now().await; self.record(InterceptorPoint::PromptSubmit, PromptAction::Continue) } - async fn pending_history_appends(&self) -> InterceptorResult> { + async fn pending_history_appends( + &self, + _context: PendingHistoryAppendsContext<'_, ()>, + ) -> InterceptorResult> { tokio::task::yield_now().await; self.record(InterceptorPoint::PendingHistoryAppends, Vec::new()) } async fn pre_llm_request( &self, - _context: PreLlmRequestContext<'_>, + _context: PreLlmRequestContext<'_, ()>, ) -> InterceptorResult { tokio::task::yield_now().await; self.record(InterceptorPoint::PreLlmRequest, PreRequestAction::Continue) } - async fn pre_tool_call(&self, _info: &mut ToolCallInfo) -> InterceptorResult { + 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: &ToolResultInfo) -> InterceptorResult { + async fn post_tool_call( + &self, + _info: &ToolResultInfo<'_, ()>, + ) -> InterceptorResult { tokio::task::yield_now().await; self.record(InterceptorPoint::PostToolCall, PostToolAction::Continue) } async fn on_assistant_turn_end( &self, - context: AssistantTurnEndContext<'_>, + context: AssistantTurnEndContext<'_, ()>, ) -> InterceptorResult { 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, + ) -> InterceptorResult { + 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 { 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 { - 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", diff --git a/crates/agen/tests/parallel_execution_test.rs b/crates/agen/tests/parallel_execution_test.rs index f63b20a1..c79c0a2e 100644 --- a/crates/agen/tests/parallel_execution_test.rs +++ b/crates/agen/tests/parallel_execution_test.rs @@ -7,14 +7,17 @@ use std::sync::{Arc, Mutex}; use std::time::{Duration, Instant}; use agen::interceptor::{ - Interceptor, InterceptorResult, PostToolAction, PreToolAction, ToolCallInfo, ToolResultInfo, + Interceptor, InterceptorErrorCategory, InterceptorPhase, InterceptorResult, PostToolAction, + PreToolAction, ToolCallInfo, ToolResultInfo, }; use agen::llm_client::event::{Event, ResponseStatus, StatusEvent}; use agen::tool::{ Tool, ToolDefinition, ToolError, ToolExecutionContext, ToolMeta, ToolOutput, ToolResult, ToolResultDisposition, }; -use agen::{Engine, History, Item, ToolExecutionPolicy}; +use agen::{ + Engine, EngineError, EngineRunExit, History, Item, RunInterruptionReason, ToolExecutionPolicy, +}; use async_trait::async_trait; mod common; @@ -907,7 +910,10 @@ 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) -> InterceptorResult { + async fn pre_tool_call( + &self, + info: &mut ToolCallInfo<'_, ()>, + ) -> InterceptorResult { self.pre_contexts.lock().unwrap().push(info.context.clone()); Ok(match info.call.name.as_str() { "skip_tool" => PreToolAction::Skip, @@ -919,7 +925,10 @@ async fn test_tool_execution_context_for_skipped_and_synthetic_paths() { }) } - async fn post_tool_call(&self, info: &ToolResultInfo) -> InterceptorResult { + async fn post_tool_call( + &self, + info: &ToolResultInfo<'_, ()>, + ) -> InterceptorResult { self.post_contexts .lock() .unwrap() @@ -996,7 +1005,10 @@ async fn test_before_tool_call_skip() { #[async_trait] impl Interceptor for BlockingPolicy { - async fn pre_tool_call(&self, info: &mut ToolCallInfo) -> InterceptorResult { + async fn pre_tool_call( + &self, + info: &mut ToolCallInfo<'_, ()>, + ) -> InterceptorResult { Ok(if info.call.name == "blocked_tool" { PreToolAction::Skip } else { @@ -1083,7 +1095,19 @@ async fn test_post_tool_call_observes_committed_result() { #[async_trait] impl Interceptor for ObservingPolicy { - async fn post_tool_call(&self, info: &ToolResultInfo) -> InterceptorResult { + async fn post_tool_call( + &self, + info: &ToolResultInfo<'_, ()>, + ) -> InterceptorResult { + assert_eq!(info.invocation.phase, InterceptorPhase::PostToolCall); + assert_eq!( + info.invocation.call_id, + Some(agen::InterceptorCallId::Tool(info.call.id.clone())) + ); + assert!(matches!( + info.history.last().map(|entry| &entry.item), + Some(Item::ToolResult { call_id, .. }) if call_id == &info.call.id + )); *self.observed_content.lock().unwrap() = Some(info.result.summary.clone()); Ok(PostToolAction::Continue) } @@ -1144,7 +1168,10 @@ async fn test_before_tool_call_synthetic_result_committed() { #[async_trait] impl Interceptor for SyntheticPolicy { - async fn pre_tool_call(&self, info: &mut ToolCallInfo) -> InterceptorResult { + async fn pre_tool_call( + &self, + info: &mut ToolCallInfo<'_, ()>, + ) -> InterceptorResult { Ok(PreToolAction::SyntheticResult(ToolResult::error( info.call.id.clone(), "permission denied", @@ -1168,6 +1195,80 @@ async fn test_before_tool_call_synthetic_result_committed() { ))); } +#[derive(Clone, Copy)] +enum InvalidIdentityMode { + ContinuedCall, + SyntheticResult, +} + +struct InvalidIdentityPolicy(InvalidIdentityMode); + +#[async_trait] +impl Interceptor for InvalidIdentityPolicy { + async fn pre_tool_call( + &self, + info: &mut ToolCallInfo<'_, ()>, + ) -> InterceptorResult { + assert_eq!(info.invocation.phase, InterceptorPhase::PreToolCall); + assert_eq!( + info.invocation.call_id, + Some(agen::InterceptorCallId::Tool("call_1".to_string())) + ); + assert!(matches!( + info.history.last().map(|entry| &entry.item), + Some(Item::ToolCall { call_id, .. }) if call_id == "call_1" + )); + Ok(match self.0 { + InvalidIdentityMode::ContinuedCall => { + info.call.id = "different-call".to_string(); + PreToolAction::Continue + } + InvalidIdentityMode::SyntheticResult => PreToolAction::SyntheticResult( + ToolResult::error("different-call", "invalid synthetic result"), + ), + }) + } +} + +#[tokio::test] +async fn interceptor_cannot_change_tool_call_identity() { + for mode in [ + InvalidIdentityMode::ContinuedCall, + InvalidIdentityMode::SyntheticResult, + ] { + let client = MockLlmClient::new(vec![ + Event::tool_use_start(0, "call_1", "echo"), + Event::tool_input_delta(0, r#"{}"#), + Event::tool_use_stop(0), + Event::Status(StatusEvent { + status: ResponseStatus::Completed, + }), + ]); + let mut engine = Engine::new(client); + engine.register_tool(SlowTool::new("echo", 1).definition()); + engine.set_interceptor(InvalidIdentityPolicy(mode)); + let mut history = History::new(); + + let result = engine.run(&mut history, "identity").await; + let EngineRunExit::Interrupted(RunInterruptionReason::Unexpected( + EngineError::Interceptor(failure), + )) = result.result + else { + panic!("invalid tool identity must interrupt with a typed failure"); + }; + assert_eq!(failure.phase(), InterceptorPhase::PreToolCall); + assert_eq!( + failure.error().category(), + InterceptorErrorCategory::ContractViolation + ); + assert!( + !history + .items() + .any(|item| matches!(item, Item::ToolResult { .. })) + ); + } +} + #[tokio::test] async fn post_tool_abort_commits_confirmed_result_before_stopping_run() { let client = MockLlmClient::new(vec![ @@ -1190,7 +1291,7 @@ async fn post_tool_abort_commits_confirmed_result_before_stopping_run() { impl Interceptor for AbortAfterResult { async fn post_tool_call( &self, - _info: &ToolResultInfo, + _info: &ToolResultInfo<'_, ()>, ) -> InterceptorResult { self.lifecycle.lock().unwrap().push("post_tool_call"); Ok(PostToolAction::Abort("policy stopped the run".to_string())) diff --git a/crates/worker/src/compact/worker.rs b/crates/worker/src/compact/worker.rs index 158c68ff..5aef61df 100644 --- a/crates/worker/src/compact/worker.rs +++ b/crates/worker/src/compact/worker.rs @@ -400,10 +400,10 @@ impl CompactWorkerInterceptor { } #[async_trait] -impl Interceptor for CompactWorkerInterceptor { +impl Interceptor for CompactWorkerInterceptor { async fn pre_llm_request( &self, - context: PreLlmRequestContext<'_>, + context: PreLlmRequestContext<'_, A>, ) -> InterceptorResult { let context = context.items; let records = self.usage_tracker.records(); @@ -427,7 +427,10 @@ impl Interceptor for CompactWorkerInterceptor { Ok(PreRequestAction::Continue) } - async fn pre_tool_call(&self, info: &mut ToolCallInfo) -> InterceptorResult { + async fn pre_tool_call( + &self, + info: &mut ToolCallInfo<'_, A>, + ) -> InterceptorResult { if self.final_reserve_tokens == 0 || info.call.name == "write_summary" { return Ok(PreToolAction::Continue); } @@ -475,8 +478,10 @@ mod tests { assert!(matches!( interceptor - .pre_llm_request(PreLlmRequestContext { + .pre_llm_request(PreLlmRequestContext::<()> { + invocation: Default::default(), items: &mut context, + history: &[], }) .await .unwrap(), @@ -486,8 +491,10 @@ mod tests { assert!(matches!( interceptor - .pre_llm_request(PreLlmRequestContext { + .pre_llm_request(PreLlmRequestContext::<()> { + invocation: Default::default(), items: &mut context, + history: &[], }) .await .unwrap(), @@ -499,8 +506,10 @@ mod tests { // current occupancy is still the latest 100-token measurement. assert!(matches!( interceptor - .pre_llm_request(PreLlmRequestContext { + .pre_llm_request(PreLlmRequestContext::<()> { + invocation: Default::default(), items: &mut context, + history: &[], }) .await .unwrap(), @@ -526,8 +535,10 @@ mod tests { assert!(matches!( interceptor - .pre_llm_request(PreLlmRequestContext { + .pre_llm_request(PreLlmRequestContext::<()> { + invocation: Default::default(), items: &mut context, + history: &[], }) .await .unwrap(), @@ -537,8 +548,10 @@ mod tests { assert!(matches!( interceptor - .pre_llm_request(PreLlmRequestContext { + .pre_llm_request(PreLlmRequestContext::<()> { + invocation: Default::default(), items: &mut context, + history: &[], }) .await .unwrap(), @@ -556,8 +569,10 @@ mod tests { assert!(matches!( interceptor - .pre_llm_request(PreLlmRequestContext { + .pre_llm_request(PreLlmRequestContext::<()> { + invocation: Default::default(), items: &mut context, + history: &[], }) .await .unwrap(), @@ -567,8 +582,10 @@ mod tests { assert!(matches!( interceptor - .pre_llm_request(PreLlmRequestContext { + .pre_llm_request(PreLlmRequestContext::<()> { + invocation: Default::default(), items: &mut context, + history: &[], }) .await .unwrap(), diff --git a/crates/worker/src/feature.rs b/crates/worker/src/feature.rs index f0231d15..5698e84b 100644 --- a/crates/worker/src/feature.rs +++ b/crates/worker/src/feature.rs @@ -1795,7 +1795,7 @@ impl FeatureRegistryBuilder { } /// Install modules into the existing Engine tool path and hook builder. - pub(crate) fn install_into_engine( + pub(crate) fn install_into_engine( self, worker: &mut Engine, hook_builder: &mut HookRegistryBuilder, diff --git a/crates/worker/src/ipc/interceptor.rs b/crates/worker/src/ipc/interceptor.rs index 8d380ea7..2706f715 100644 --- a/crates/worker/src/ipc/interceptor.rs +++ b/crates/worker/src/ipc/interceptor.rs @@ -15,7 +15,8 @@ use std::sync::{Arc, Mutex}; use agen::Item; use agen::UsageRecord; use agen::interceptor::{ - AssistantTurnEndContext, Interceptor, InterceptorResult, PostToolAction, PreLlmRequestContext, + AssistantTurnEndContext, Interceptor, InterceptorError, InterceptorErrorCategory, + InterceptorResult, PendingHistoryAppendsContext, PostToolAction, PreLlmRequestContext, PreRequestAction, PreToolAction, PromptAction, PromptSubmitContext, ToolCallInfo, ToolResultInfo, TurnEndAction, }; @@ -232,10 +233,10 @@ impl WorkerInterceptor { } #[async_trait] -impl Interceptor for WorkerInterceptor { +impl Interceptor for WorkerInterceptor { async fn on_prompt_submit( &self, - context: PromptSubmitContext<'_>, + context: PromptSubmitContext<'_, SessionHistoryMetadata>, ) -> InterceptorResult { let item = context.item; let turn_index = self.next_turn_index.fetch_add(1, Ordering::Relaxed); @@ -274,7 +275,10 @@ impl Interceptor for WorkerInterceptor { }) } - async fn pending_history_appends(&self) -> InterceptorResult> { + async fn pending_history_appends( + &self, + _context: PendingHistoryAppendsContext<'_, SessionHistoryMetadata>, + ) -> InterceptorResult> { let drained = self.pending_notifies.drain(); if drained.is_empty() { return Ok(Vec::new()); @@ -300,7 +304,10 @@ 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}").into()); + return Err(InterceptorError::new( + InterceptorErrorCategory::Dependency, + format!("failed to render notify_wrapper: {error}"), + )); } }; items.push(system_item.to_history_item()); @@ -308,14 +315,17 @@ 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}").into()); + return Err(InterceptorError::new( + InterceptorErrorCategory::Dependency, + format!("session persistence failed: {error}"), + )); } Ok(items) } async fn pre_llm_request( &self, - context: PreLlmRequestContext<'_>, + context: PreLlmRequestContext<'_, SessionHistoryMetadata>, ) -> InterceptorResult { let context = context.items; let initial_tokens = self.estimated_tokens(context); @@ -385,7 +395,10 @@ impl Interceptor for WorkerInterceptor { }) } - async fn pre_tool_call(&self, info: &mut ToolCallInfo) -> InterceptorResult { + async fn pre_tool_call( + &self, + info: &mut ToolCallInfo<'_, SessionHistoryMetadata>, + ) -> InterceptorResult { let summary = ToolCallSummary { call_id: info.call.id.clone(), tool_name: info.call.name.clone(), @@ -401,7 +414,10 @@ impl Interceptor for WorkerInterceptor { Ok(PreToolAction::Continue) } - async fn post_tool_call(&self, info: &ToolResultInfo) -> InterceptorResult { + async fn post_tool_call( + &self, + info: &ToolResultInfo<'_, SessionHistoryMetadata>, + ) -> InterceptorResult { let summary = ToolResultSummary { call_id: info.result.tool_use_id.clone(), tool_name: info.call.name.clone(), @@ -424,14 +440,14 @@ impl Interceptor for WorkerInterceptor { async fn on_assistant_turn_end( &self, - context: AssistantTurnEndContext<'_>, + context: AssistantTurnEndContext<'_, SessionHistoryMetadata>, ) -> InterceptorResult { let history = context.history; let final_text_preview = history .iter() .rev() - .find(|i| i.is_assistant_message()) - .and_then(extract_message_text) + .find(|entry| entry.item.is_assistant_message()) + .and_then(|entry| extract_message_text(&entry.item)) .map(|t| preview(&t, FINAL_TEXT_PREVIEW_LIMIT)) .unwrap_or_default(); let info = TurnEndInfo { @@ -515,6 +531,7 @@ mod tests { Hook, HookPostToolAction, HookPreRequestAction, HookPreToolAction, HookRegistryBuilder, HookTurnEndAction, OnTurnEnd, PostToolCall, PreLlmRequest, PreToolCall, }; + use crate::session_history::{WorkerHistoryProvenance, history_entry}; fn test_prompts() -> Arc> { Arc::new(ArcSwap::from(PromptCatalog::builtins_only().unwrap())) @@ -574,7 +591,10 @@ mod tests { } } - fn task_tool_call_info(name: &str, input: serde_json::Value) -> ToolCallInfo { + fn task_tool_call_info( + name: &str, + input: serde_json::Value, + ) -> ToolCallInfo<'static, SessionHistoryMetadata> { let def = crate::feature::builtin::task::task_tools( crate::feature::builtin::task::TaskStore::new(), ) @@ -586,6 +606,8 @@ mod tests { .expect("task tool definition"); let (meta, tool) = def(); ToolCallInfo { + invocation: Default::default(), + history: &[], call: agen::tool::ToolCall { id: "call-id".into(), name: name.into(), @@ -630,7 +652,11 @@ mod tests { ); let mut ctx = ctx_items; let action = interceptor - .pre_llm_request(PreLlmRequestContext { items: &mut ctx }) + .pre_llm_request(PreLlmRequestContext { + invocation: Default::default(), + items: &mut ctx, + history: &[], + }) .await .unwrap(); @@ -665,7 +691,11 @@ mod tests { ); let mut ctx = ctx_items; let action = interceptor - .pre_llm_request(PreLlmRequestContext { items: &mut ctx }) + .pre_llm_request(PreLlmRequestContext { + invocation: Default::default(), + items: &mut ctx, + history: &[], + }) .await .unwrap(); @@ -705,7 +735,11 @@ mod tests { .with_usage_tracker(usage_tracker); let mut ctx = ctx_items; let action = interceptor - .pre_llm_request(PreLlmRequestContext { items: &mut ctx }) + .pre_llm_request(PreLlmRequestContext { + invocation: Default::default(), + items: &mut ctx, + history: &[], + }) .await .unwrap(); @@ -732,7 +766,11 @@ mod tests { ); let mut ctx = ctx_items; let action = interceptor - .pre_llm_request(PreLlmRequestContext { items: &mut ctx }) + .pre_llm_request(PreLlmRequestContext { + invocation: Default::default(), + items: &mut ctx, + history: &[], + }) .await .unwrap(); @@ -776,7 +814,11 @@ mod tests { ); let mut ctx = ctx_items; let action = interceptor - .pre_llm_request(PreLlmRequestContext { items: &mut ctx }) + .pre_llm_request(PreLlmRequestContext { + invocation: Default::default(), + items: &mut ctx, + history: &[], + }) .await .unwrap(); @@ -806,7 +848,11 @@ mod tests { ); let mut ctx = ctx_items; let action = interceptor - .pre_llm_request(PreLlmRequestContext { items: &mut ctx }) + .pre_llm_request(PreLlmRequestContext { + invocation: Default::default(), + items: &mut ctx, + history: &[], + }) .await .unwrap(); @@ -830,7 +876,11 @@ mod tests { ); let mut ctx: Vec = Vec::new(); let action = interceptor - .pre_llm_request(PreLlmRequestContext { items: &mut ctx }) + .pre_llm_request(PreLlmRequestContext { + invocation: Default::default(), + items: &mut ctx, + history: &[], + }) .await .unwrap(); @@ -862,7 +912,11 @@ mod tests { let mut ctx: Vec = Vec::new(); let action = interceptor - .pre_llm_request(PreLlmRequestContext { items: &mut ctx }) + .pre_llm_request(PreLlmRequestContext { + invocation: Default::default(), + items: &mut ctx, + history: &[], + }) .await .unwrap(); @@ -912,7 +966,11 @@ mod tests { let mut ctx: Vec = Vec::new(); let action = interceptor - .pre_llm_request(PreLlmRequestContext { items: &mut ctx }) + .pre_llm_request(PreLlmRequestContext { + invocation: Default::default(), + items: &mut ctx, + history: &[], + }) .await .unwrap(); @@ -1017,7 +1075,9 @@ mod tests { None, ); let info = task_tool_call_info("TaskList", serde_json::json!({})); - let mut result_info = ToolResultInfo { + let result_info = ToolResultInfo { + invocation: Default::default(), + history: &[], call: info.call, result: agen::tool::ToolResult::from_output( "call-id", @@ -1033,7 +1093,7 @@ mod tests { context: info.context, }; - let action = interceptor.post_tool_call(&mut result_info).await.unwrap(); + let action = interceptor.post_tool_call(&result_info).await.unwrap(); assert_eq!(action, PostToolAction::Abort("post tool abort".to_string())); assert_eq!(count.load(Ordering::Relaxed), 1); @@ -1067,11 +1127,20 @@ mod tests { test_prompts(), None, ); - let history = vec![Item::user_message("hi"), Item::assistant_message("done")]; - + let history = vec![ + history_entry( + Item::user_message("hi"), + WorkerHistoryProvenance::LegacyUnknown, + ), + history_entry( + Item::assistant_message("done"), + WorkerHistoryProvenance::LegacyUnknown, + ), + ]; let action = interceptor .on_assistant_turn_end(AssistantTurnEndContext { - assistant_items: &[], + invocation: Default::default(), + assistant_entries: &history[1..], history: &history, tool_calls: &[], }) @@ -1114,7 +1183,11 @@ mod tests { for _ in 0..23 { let mut ctx = ctx_items.clone(); let action = interceptor - .pre_llm_request(PreLlmRequestContext { items: &mut ctx }) + .pre_llm_request(PreLlmRequestContext { + invocation: Default::default(), + items: &mut ctx, + history: &[], + }) .await .unwrap(); assert!(matches!(action, PreRequestAction::Continue)); @@ -1129,7 +1202,11 @@ mod tests { let mut ctx = ctx_items.clone(); let action = interceptor - .pre_llm_request(PreLlmRequestContext { items: &mut ctx }) + .pre_llm_request(PreLlmRequestContext { + invocation: Default::default(), + items: &mut ctx, + history: &[], + }) .await .unwrap(); let appended_len = match action { @@ -1204,7 +1281,13 @@ mod tests { )); buffer.push_notify("updated".to_string(), false); - let appends = interceptor.pending_history_appends().await.unwrap(); + let appends = interceptor + .pending_history_appends(PendingHistoryAppendsContext { + invocation: Default::default(), + history: &[], + }) + .await + .unwrap(); assert_eq!(appends.len(), 1); assert!(format!("{:?}", appends[0]).contains("CURRENT-PROJECTION updated")); let committed = committed.lock().unwrap(); @@ -1254,9 +1337,19 @@ mod tests { )); buffer.push_notify("must persist".to_string(), false); - let error = interceptor.pending_history_appends().await.unwrap_err(); + let error = interceptor + .pending_history_appends(PendingHistoryAppendsContext { + invocation: Default::default(), + history: &[], + }) + .await + .unwrap_err(); - assert!(error.message().contains("failed to render notify_wrapper")); + assert!( + error + .diagnostic() + .contains("failed to render notify_wrapper") + ); let requeued = buffer.drain(); assert_eq!(requeued.len(), 1); } @@ -1278,7 +1371,13 @@ mod tests { None, ); - let items = interceptor.pending_history_appends().await.unwrap(); + let items = interceptor + .pending_history_appends(PendingHistoryAppendsContext { + invocation: Default::default(), + history: &[], + }) + .await + .unwrap(); assert_eq!(items.len(), 2); let first = items[0].as_text().unwrap_or_default(); let second = items[1].as_text().unwrap_or_default(); @@ -1292,7 +1391,13 @@ mod tests { ); // Empty buffer → empty Vec (no synthesised items). - let again = interceptor.pending_history_appends().await.unwrap(); + let again = interceptor + .pending_history_appends(PendingHistoryAppendsContext { + invocation: Default::default(), + history: &[], + }) + .await + .unwrap(); assert!(again.is_empty()); } @@ -1316,7 +1421,11 @@ mod tests { ); let mut ctx: Vec = vec![Item::user_message("hi")]; let action = interceptor - .pre_llm_request(PreLlmRequestContext { items: &mut ctx }) + .pre_llm_request(PreLlmRequestContext { + invocation: Default::default(), + items: &mut ctx, + history: &[], + }) .await .unwrap(); @@ -1349,7 +1458,11 @@ mod tests { ); let mut ctx: Vec = Vec::new(); let action = interceptor - .pre_llm_request(PreLlmRequestContext { items: &mut ctx }) + .pre_llm_request(PreLlmRequestContext { + invocation: Default::default(), + items: &mut ctx, + history: &[], + }) .await .unwrap(); diff --git a/crates/worker/src/worker.rs b/crates/worker/src/worker.rs index 332ab2e5..db3536ca 100644 --- a/crates/worker/src/worker.rs +++ b/crates/worker/src/worker.rs @@ -5996,7 +5996,7 @@ where /// Note: `system_prompt` is intentionally not applied here. It is a /// minijinja template that is parsed by `Worker::from_manifest` and /// rendered once at first turn in `ensure_system_prompt_materialized`. -pub fn apply_worker_manifest( +pub fn apply_worker_manifest( worker: &mut Engine, wm: &manifest::EngineManifest, ) {